From 0fd44259cdd5d7929c07c0fa82c92fd649433d72 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Sun, 10 Dec 2023 16:31:52 -0800 Subject: [PATCH] bf16 fix + cleanups from mixtral (#2698) * bf16 fix + cleanups from mixtral * generic bf16 cast --- examples/coder.py | 9 ++++++--- examples/llama.py | 3 +++ test/unit/test_disk_tensor.py | 3 ++- tinygrad/device.py | 4 ++-- tinygrad/helpers.py | 1 + tinygrad/nn/state.py | 10 ++-------- tinygrad/runtime/ops_hip.py | 3 ++- tinygrad/tensor.py | 8 ++++++-- 8 files changed, 24 insertions(+), 17 deletions(-) diff --git a/examples/coder.py b/examples/coder.py index 9cd5b429e5..597dd13011 100644 --- a/examples/coder.py +++ b/examples/coder.py @@ -4,7 +4,7 @@ sys.path.append(os.getcwd()) from io import StringIO from contextlib import redirect_stdout -from tinygrad import Tensor, nn +from tinygrad import Tensor, nn, Device, dtypes from tinygrad.helpers import Timing, colored, getenv, fetch from extra.models.llama import Transformer, convert_from_huggingface from sentencepiece import SentencePieceProcessor @@ -30,9 +30,12 @@ if __name__ == "__main__": part1 = nn.state.torch_load(fetch("https://huggingface.co/teknium/OpenHermes-2.5-Mistral-7B/resolve/main/pytorch_model-00001-of-00002.bin?download=true")) part2 = nn.state.torch_load(fetch("https://huggingface.co/teknium/OpenHermes-2.5-Mistral-7B/resolve/main/pytorch_model-00002-of-00002.bin?download=true")) + # fix bf16, TODO: check if device supports bf16 + def fix_bf16(weights): return {k:v.to(Device.DEFAULT).cast(dtypes.float16) if v.dtype == dtypes.bfloat16 else v for k,v in weights.items()} + with Timing("weights -> model: "): - nn.state.load_state_dict(model, convert_from_huggingface(part1, model, 32, 8), strict=False) - nn.state.load_state_dict(model, convert_from_huggingface(part2, model, 32, 8), strict=False) + nn.state.load_state_dict(model, fix_bf16(convert_from_huggingface(part1, model, 32, 8)), strict=False) + nn.state.load_state_dict(model, fix_bf16(convert_from_huggingface(part2, model, 32, 8)), strict=False) if not os.path.isfile("/tmp/tokenizer.model"): create_fixed_tokenizer("/tmp/tokenizer.model") spp = SentencePieceProcessor(model_file="/tmp/tokenizer.model") diff --git a/examples/llama.py b/examples/llama.py index 6cf183ffdd..3a8d17c75c 100755 --- a/examples/llama.py +++ b/examples/llama.py @@ -165,6 +165,9 @@ class LLaMa: if "model.embed_tokens.weight" in weights: weights = convert_from_huggingface(weights, model, params["args"]["n_heads"], params["args"].get("n_kv_heads", params["args"]["n_heads"])) + # fix bf16, TODO: check if device supports bf16 + weights = {k:v.to(Device.DEFAULT).cast(dtypes.float16) if v.dtype == dtypes.bfloat16 else v for k,v in weights.items()} + if quantize: weights = AbsmaxQuantizedLinear.quantize(weights) for _,v in weights.items(): v.realize() diff --git a/test/unit/test_disk_tensor.py b/test/unit/test_disk_tensor.py index 6b11831a4a..f13c5f23d5 100644 --- a/test/unit/test_disk_tensor.py +++ b/test/unit/test_disk_tensor.py @@ -1,7 +1,7 @@ import pathlib import unittest import numpy as np -from tinygrad.tensor import Tensor, Device +from tinygrad.tensor import Tensor, Device, dtypes from tinygrad.nn.state import safe_load, safe_save, get_state_dict, torch_load from tinygrad.helpers import CI, fetch, temp from tinygrad.helpers import Timing @@ -13,6 +13,7 @@ def compare_weights_both(url): torch_weights = get_state_dict(torch.load(fn, map_location=torch.device('cpu')), tensor_type=torch.Tensor) assert list(tg_weights.keys()) == list(torch_weights.keys()) for k in tg_weights: + if tg_weights[k].dtype == dtypes.bfloat16: tg_weights[k] = torch_weights[k].float() # numpy doesn't support bfloat16 if torch_weights[k].dtype == torch.bfloat16: torch_weights[k] = torch_weights[k].float() # numpy doesn't support bfloat16 np.testing.assert_equal(tg_weights[k].numpy(), torch_weights[k].numpy(), err_msg=f"mismatch at {k}, {tg_weights[k].shape}") print(f"compared {len(tg_weights)} weights") diff --git a/tinygrad/device.py b/tinygrad/device.py index 1405af0975..56a7de6848 100644 --- a/tinygrad/device.py +++ b/tinygrad/device.py @@ -56,7 +56,7 @@ def update_stats(name:str, op_estimate:sint, mem_estimate:sint, var_vals: Option GlobalCounters.global_mem += mem_estimate if et is not None: GlobalCounters.time_sum_s += et if DEBUG >= 2: - print(f"{colored(f'*** {GlobalCounters.kernel_count:4d}', ('magenta' if num_kernels == 1 else 'CYAN') if jit else None)} {name+' '*(37-ansilen(name))} arg {buf_count:3d} sz {str(lra.get('global_size', '') if lra else ''):18s} dev {device:7s} OPs {int(op_estimate/1e6):6d}M/{GlobalCounters.global_ops/1e9:7.2f}G mem {GlobalCounters.mem_used/1e9:5.2f} GB " + + print(f"{colored(f'*** {GlobalCounters.kernel_count:4d}', ('magenta' if num_kernels == 1 else 'CYAN') if jit else None)} {name+' '*(37-ansilen(name))} arg {buf_count:3d} sz {str(lra.get('global_size', '') if lra else ''):18s} dev {device[:10]:10s} OPs {int(op_estimate/1e6):6d}M/{GlobalCounters.global_ops/1e9:7.2f}G mem {GlobalCounters.mem_used/1e9:5.2f} GB " + (str() if et is None else f"tm {et*1e6:9.2f}us/{GlobalCounters.time_sum_s*1e3:9.2f}ms ({op_estimate/((et or 1e-20)*1e9):8.2f} GFLOPS, {mem_estimate/((et or 1e-20)*1e9):7.2f} GB/s)")) # **************** Buffer / Allocator **************** @@ -123,7 +123,7 @@ class _BufferCopy(JITRunner): if wait or DEBUG >= 2: Device[dest.device].synchronize() et = time.perf_counter() - st - update_stats(colored(f"copy {dest.device:7s} <- {src.device:7s}", "yellow"), 0, dest.size*dest.dtype.itemsize, {}, et, 2, jit, lra={"global_size": dest.size}, device=dest.device) + update_stats(colored(f"copy {dest.device[:10]:10s} <- {src.device[:10]:10s}", "yellow"), 0, dest.size*dest.dtype.itemsize, {}, et, 2, jit, lra={"global_size": dest.size}, device=dest.device) BufferCopy = _BufferCopy() # TODO: size, dest, src are the same type. can we enforce this? diff --git a/tinygrad/helpers.py b/tinygrad/helpers.py index c37d3dae12..bcc70cfbca 100644 --- a/tinygrad/helpers.py +++ b/tinygrad/helpers.py @@ -268,6 +268,7 @@ def cpu_time_execution(cb, enable): # *** ctypes helpers +# TODO: make this work with read only memoryviews (if possible) def from_mv(mv, to_type=ctypes.c_char): return ctypes.cast(ctypes.addressof(to_type.from_buffer(mv)), ctypes.POINTER(to_type)) def to_char_p_p(options: List[bytes], to_type=ctypes.c_char): return (ctypes.POINTER(to_type) * len(options))(*[ctypes.cast(ctypes.create_string_buffer(o), ctypes.POINTER(to_type)) for o in options]) @functools.lru_cache(maxsize=None) diff --git a/tinygrad/nn/state.py b/tinygrad/nn/state.py index 0882e24676..4d57648571 100644 --- a/tinygrad/nn/state.py +++ b/tinygrad/nn/state.py @@ -4,7 +4,6 @@ from typing import Dict, Union, List, Optional, Any, Tuple from tinygrad.tensor import Tensor from tinygrad.helpers import dtypes, prod, argsort, DEBUG, Timing, GlobalCounters, CI, unwrap from tinygrad.shape.view import strides_for_shape -from tinygrad import Device safe_dtypes = {"F16": dtypes.float16, "F32": dtypes.float32, "U8": dtypes.uint8, "I8": dtypes.int8, "I32": dtypes.int32, "I64": dtypes.int64} inverse_safe_dtypes = {v:k for k,v in safe_dtypes.items()} @@ -72,13 +71,7 @@ def torch_load(fn:str): lens[storage[2]] = storage[4] * storage[1].itemsize if storage[2] not in offsets: return None byte_offset = offsets[storage[2]]+storage_offset*storage[1].itemsize - ret = t[byte_offset:byte_offset+prod(size)] - # convert bfloat16 -> float16 using LLVM for Llama 2 - # upstream LLaMA also does this conversion: - # https://github.com/facebookresearch/llama/blob/6c7fe276574e78057f917549435a2554000a876d/llama/generation.py#L95 - # TODO: should this be done in the example instead? or maybe we don't need this anymore with better bfloat16 support - if storage[1] == dtypes.bfloat16: ret = ret.cast(dtypes.uint16).to(Device.DEFAULT).cast(dtypes.uint32).mul(1<<16).contiguous().bitcast(dtypes.float32).half() - else: ret = ret.cast(storage[1]) + ret = t[byte_offset:byte_offset+prod(size)].cast(storage[1]) # 7 lines to deal with permuted tensors. NOTE: this currently requires reading off the disk shape_strides = [(s, st) for s,st in zip(size, stride) if s != 1] @@ -87,6 +80,7 @@ def torch_load(fn:str): intermediate_shape = tuple([shape_strides[x][0] for x in argsort(permute_indexes)]) assert tuple([shape_strides[i][1] for i in argsort(permute_indexes)]) == strides_for_shape(intermediate_shape), "nonpermutable strides" if DEBUG >= 3: print(f"WARNING: this torch load is slow. CPU to permute {intermediate_shape} with {permute_indexes}") + assert storage[1] != dtypes.bfloat16, "can't CPU permute BF16" # TODO: find a nice way to support all shapetracker on disktensors ret = ret.cpu().reshape(intermediate_shape).permute(permute_indexes) diff --git a/tinygrad/runtime/ops_hip.py b/tinygrad/runtime/ops_hip.py index 685ded4fe0..54e401c3ca 100644 --- a/tinygrad/runtime/ops_hip.py +++ b/tinygrad/runtime/ops_hip.py @@ -49,7 +49,8 @@ class HIPAllocator(LRUAllocator): def _free(self, opaque:T): check(hip.hipFree(opaque)) def copyin(self, dest:T, src: memoryview): check(hip.hipSetDevice(self.device)) - check(hip.hipMemcpyAsync(dest, from_mv(src), len(src), hip.hipMemcpyHostToDevice, None)) + # TODO: have to make sure src isn't freed to make this async + check(hip.hipMemcpy(dest, from_mv(src), len(src), hip.hipMemcpyHostToDevice)) def copyout(self, dest:memoryview, src:T): check(hip.hipSetDevice(self.device)) check(hip.hipMemcpy(from_mv(dest), src, len(dest), hip.hipMemcpyDeviceToHost)) diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 9baf972ca5..af16824b6a 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -112,7 +112,8 @@ class Tensor: self.contiguous().realize().lazydata.realized.copyin(x.numpy().data) return self if x.__class__ is not Tensor: x = Tensor(x, device=self.device, dtype=self.dtype) - assert self.shape == x.shape and self.device == x.device, f"assign shape mismatch {self.shape} != {x.shape} or device mismatch {self.device} != {x.device}" + # NOTE: we allow cross device assign + assert self.shape == x.shape, f"assign shape mismatch {self.shape} != {x.shape}" assert not x.requires_grad # self requires_grad is okay? if DEBUG >= 4: print(f"assign {self.lazydata} <- {x.lazydata}") if self.dtype == x.dtype and self.lazydata.realized is not None and not getenv("DISALLOW_ASSIGN"): x.lazydata.output_buffer = self.lazydata.realized @@ -855,7 +856,10 @@ class Tensor: # ***** cast ops ***** - def cast(self, dtype:DType) -> Tensor: return mlops.Cast.apply(self, dtype=dtype) if self.dtype != dtype else self + def cast(self, dtype:DType) -> Tensor: + # hack for devices that don't support bfloat16 + if self.dtype == dtypes.bfloat16: return self.bitcast(dtypes.uint16).cast(dtypes.uint32).mul(1<<16).contiguous().bitcast(dtypes.float32).cast(dtype) + return mlops.Cast.apply(self, dtype=dtype) if self.dtype != dtype else self def bitcast(self, dtype:DType) -> Tensor: assert self.dtype.itemsize == dtype.itemsize, "can't bitcast mismatched dtype itemsizes" return mlops.Cast.apply(self, dtype=dtype, bitcast=True) if self.dtype != dtype else self