From 97bc72353827f7415577e0df3e49758975d1f205 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Sat, 22 Feb 2025 22:16:23 +0800 Subject: [PATCH] torch backend works for ResNet-18 (#9200) * torch backend progress, a few more functions * resnet works * pillow * tv --- .github/workflows/test.yml | 5 +- extra/torch_backend/.gitignore | 1 + extra/torch_backend/backend.py | 96 +++++++++++++++++++++++++++------- extra/torch_backend/example.py | 19 +++++++ 4 files changed, 101 insertions(+), 20 deletions(-) create mode 100644 extra/torch_backend/.gitignore create mode 100644 extra/torch_backend/example.py diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 58acf64880..c8186a0eac 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -153,14 +153,17 @@ jobs: - name: Setup Environment uses: ./.github/actions/setup-tinygrad with: - key: torch-backend + key: torch-backend-pillow-torchvision deps: testing_minimal + pydeps: "pillow torchvision" - name: Install ninja run: | sudo apt update || true sudo apt install -y --no-install-recommends ninja-build - name: Test one op run: PYTHONPATH=. FORWARD_ONLY=1 TINY_BACKEND=1 python3 test/test_ops.py TestOps.test_add + - name: Test ResNet-18 + run: PYTHONPATH=. python3 extra/torch_backend/example.py - name: Test Ops with TINY_BACKEND (expect failure) run: PYTHONPATH=. TINY_BACKEND=1 pytest test/test_ops.py || true - name: Test beautiful_mnist in torch with TINY_BACKEND (expect failure) diff --git a/extra/torch_backend/.gitignore b/extra/torch_backend/.gitignore new file mode 100644 index 0000000000..1269488f7f --- /dev/null +++ b/extra/torch_backend/.gitignore @@ -0,0 +1 @@ +data diff --git a/extra/torch_backend/backend.py b/extra/torch_backend/backend.py index 3c9d111e58..19ba7bd780 100644 --- a/extra/torch_backend/backend.py +++ b/extra/torch_backend/backend.py @@ -1,5 +1,6 @@ from tinygrad import Tensor, dtypes -from tinygrad.helpers import DEBUG, getenv +from tinygrad.helpers import DEBUG, getenv, prod +TORCH_DEBUG = getenv("TORCH_DEBUG") import torch, pathlib torch.autograd.grad_mode.set_multithreading_enabled(False) @@ -46,31 +47,44 @@ def masked_select(self, mask): return wrap(Tensor(self.cpu().numpy()[mask.cpu().numpy()])) @torch.library.impl("aten::as_strided", "privateuseone") -def as_strided(tensor, size, stride, storage_offset=None): - if size == [] and storage_offset is not None: - # TODO: is this right? - return wrap(unwrap(tensor).flatten()[storage_offset:storage_offset+1].reshape(())) - # broadcast - if len(tensor.shape) == 0: return wrap(unwrap(tensor).reshape((1,)*len(size)).expand(size)) - print("******* NOTE: this as_strided is wrong ***********\n", tensor.shape, size, stride, storage_offset) - return wrap(Tensor.zeros(*size)) +def as_strided(tensor:torch.Tensor, size, stride, storage_offset=None): + #return tensor.cpu().as_strided(size, stride).tiny() + if TORCH_DEBUG >= 1: print("** NOTE: this as_strided is wrong", tensor.shape, size, stride, storage_offset) + + if tuple(x for x in tensor.shape if x != 1) == tuple(x for x in size if x != 1): + # this is squeeze/unsqueeze + return tensor.reshape(size) + + # TODO: how do i know this is permute? + if tensor.shape == (1000, 512) and size == [512, 1000] and stride == [0, 1]: + return wrap(unwrap(tensor).permute(1,0)) + + #print(tensor.cpu().numpy()) raise NotImplementedError("fix as_strided") @torch.library.impl("aten::empty_strided", "privateuseone") -def empty_strided(size, stride, dtype, layout, device, pin_memory=False): - if DEBUG >= 2: print(f"empty_strided {size=} {stride=} {dtype=} {layout=} {device=} {pin_memory=}") +def empty_strided(size, stride, dtype, layout=None, device=None, pin_memory=False): + if TORCH_DEBUG: print(f"empty_strided {size=} {stride=} {dtype=} {layout=} {device=} {pin_memory=}") ret = Tensor.empty(*size, dtype=torch_to_tiny_dtype[dtype]) return wrap(ret) @torch.library.impl("aten::empty.memory_format", "privateuseone") def empty_memory_format(size, dtype=None, layout=None, device=None, pin_memory=False, memory_format=None): - if DEBUG >= 2: print(f"empty.memory_format {size=} {dtype=} {layout=} {device=} {pin_memory=} {memory_format=}") - ret = Tensor.empty(*size, dtype=torch_to_tiny_dtype[dtype]) + if TORCH_DEBUG: print(f"empty.memory_format {size=} {dtype=} {layout=} {device=} {pin_memory=} {memory_format=}") + ret = Tensor.empty(*size, dtype=torch_to_tiny_dtype[dtype or torch.get_default_dtype()]) return wrap(ret) +@torch.library.impl("aten::max_pool2d_with_indices", "privateuseone") +def max_pool2d_with_indices(self:Tensor, kernel_size, stride=None, padding=0, dilation=1, ceil_mode=False): + # TODO: support return_indices in tinygrad + ret = unwrap(self).max_pool2d(kernel_size, stride, dilation, padding, ceil_mode) + # TODO: this is wrong + return (wrap(ret), wrap(Tensor.zeros_like(ret, dtype=dtypes.int64))) + @torch.library.impl("aten::convolution_overrideable", "privateuseone") def convolution_overrideable(input, weight, bias, stride, padding, dilation, transposed, output_padding, groups): - #print(f"{input.shape=} {weight.shape=} {bias.shape=} {stride=} {padding=} {dilation=} {transposed=} {output_padding=} {groups=}") + if TORCH_DEBUG >= 1: + print(f"convolution {input.shape=} {weight.shape=} {stride=} {padding=} {dilation=} {transposed=} {output_padding=} {groups=}") return wrap(unwrap(input).conv2d(unwrap(weight), unwrap(bias) if bias is not None else None, groups=groups, stride=stride, dilation=dilation, padding=padding)) #raise NotImplementedError("need convolution") @@ -94,29 +108,69 @@ def cat_out(tensors, out, dim=0): unwrap(out).replace(Tensor.cat(*[unwrap(x) for @torch.library.impl("aten::index.Tensor", "privateuseone") def index_tensor(x, y): return wrap(unwrap(x)[y[0].tolist()]) +# register some decompositions +from torch._decomp import get_decompositions +aten = torch.ops.aten +decomps = [ + aten.native_batch_norm, + aten.addmm, + # NOTE: many of these don't work or cause infinite loops + #aten.var_mean, + #aten.var, + #aten.rsqrt, + #aten.max_pool2d_with_indices, +] +for k,v in get_decompositions(decomps).items(): + key = str(k._schema).split("(")[0] + if TORCH_DEBUG >= 2: print("register decomp for", k) + torch.library.impl(key, "privateuseone")(v) + tiny_backend = { "aten.view": Tensor.reshape, "aten.add.Tensor": Tensor.add, "aten.sub.Tensor": Tensor.sub, "aten.mul.Tensor": Tensor.mul, "aten.div.Tensor": Tensor.div, - "aten.add_.Tensor": lambda x,y: x.assign(x.add(y)), + "aten.add_.Tensor": lambda x,y,alpha=1: x.assign(x.add(y)*alpha), "aten.pow.Tensor_Scalar": Tensor.pow, "aten.bitwise_and.Tensor": Tensor.bitwise_and, "aten.eq.Tensor": Tensor.eq, "aten.eq.Scalar": Tensor.eq, "aten.ne.Tensor": Tensor.ne, "aten.ne.Scalar": Tensor.ne, "aten.gt.Tensor": Tensor.__gt__, "aten.gt.Scalar": Tensor.__gt__, "aten.lt.Tensor": Tensor.__lt__, "aten.lt.Scalar": Tensor.__lt__, + "aten.le.Tensor": Tensor.__le__, "aten.le.Scalar": Tensor.__le__, + "aten.abs": Tensor.abs, + "aten.exp": Tensor.exp, "aten.exp2": Tensor.exp2, "aten.min": Tensor.min, "aten.max": Tensor.max, "aten.relu": Tensor.relu, + "aten.relu_": lambda x: x.assign(x.relu()), "aten.mean": Tensor.mean, + "aten.mean.dim": Tensor.mean, "aten.neg": Tensor.neg, + "aten.reciprocal": Tensor.reciprocal, + "aten.sqrt": Tensor.sqrt, + "aten.rsqrt": Tensor.rsqrt, "aten.mm": Tensor.matmul, + "aten.var.correction": Tensor.var, + # TODO: support var_mean in tinygrad + "aten.var_mean.correction": lambda self, dims, keepdim=False, correction=1: (self.var(dims, keepdim, correction), self.mean(dims, keepdim)), + # NOTE: axis=[] in torch means all, change tinygrad? + "aten.sum.IntList_out": lambda self,axis,keepdim=False,out=None: + out.replace(Tensor.sum(self, axis if len(axis) else None, keepdim), allow_shape_mismatch=True), + "aten.argmax": Tensor.argmax, + "aten.scatter.value": Tensor.scatter, + "aten.gather": Tensor.gather, + "aten.where.self": Tensor.where, + "aten._log_softmax": lambda self,dim,half_to_float: self.softmax(dim), + "aten.random_": lambda self: + self.assign(Tensor.randint(*self.shape, low=dtypes.min(self.dtype), high=dtypes.max(self.dtype), device=self.device, dtype=self.dtype)), + "aten.uniform_": lambda self, low=0, high=1: self.assign(Tensor.uniform(*self.shape, low=low, high=high)), + "aten.normal_": lambda self, low=0, high=1: self.assign(Tensor.normal(*self.shape, low=low, high=high)), } -# there's earlier things to hook here +# NOTE: there's earlier things to hook these, so the .out form isn't needed #"aten.add.out": lambda x,y,out: out.replace(x+y, allow_shape_mismatch=True), #"aten.abs.out": lambda x,out: out.replace(x.abs(), allow_shape_mismatch=True), #"aten.ceil.out": lambda x,out: out.replace(x.ceil(), allow_shape_mismatch=True), @@ -124,15 +178,19 @@ tiny_backend = { def wrap_fxn(k,f): def nf(*args, **kwargs): - #print(k, len(args), kwargs.keys()) + if TORCH_DEBUG: print(k, len(args), [x.shape if isinstance(x, torch.Tensor) else x for x in args], + {k:v.shape if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()}) args = [unwrap(x) if isinstance(x, torch.Tensor) else x for x in args] kwargs = {k:unwrap(v) if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()} - return wrap(f(*args, **kwargs)) + out = f(*args, **kwargs) + if isinstance(out, Tensor): return wrap(out) + elif isinstance(out, tuple): return tuple(wrap(x) for x in out) + else: raise RuntimeError(f"unknown output type {type(out)}") return nf for k,v in tiny_backend.items(): torch.library.impl(k.replace("aten.", "aten::"), "privateuseone")(wrap_fxn(k,v)) -if getenv("TORCH_DEBUG"): +if TORCH_DEBUG: from torch.utils._python_dispatch import TorchDispatchMode class DispatchLog(TorchDispatchMode): def __torch_dispatch__(self, func, types, args, kwargs=None): diff --git a/extra/torch_backend/example.py b/extra/torch_backend/example.py new file mode 100644 index 0000000000..570f5b2e9a --- /dev/null +++ b/extra/torch_backend/example.py @@ -0,0 +1,19 @@ +from PIL import Image +import torch, torchvision, pathlib +import torchvision.transforms as transforms +import extra.torch_backend.backend +device = "tiny" +torch.set_default_device(device) + +if __name__ == "__main__": + img = Image.open(pathlib.Path(__file__).parent.parent.parent / "test/models/efficientnet/Chicken.jpg").convert('RGB') + transform = transforms.Compose([ + transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), + transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) + ]) + img = transform(img).unsqueeze(0).to(device) + + model = torchvision.models.resnet18(weights=torchvision.models.ResNet18_Weights.DEFAULT).eval() + out = model(img).detach().cpu().numpy() + print("output:", out.shape, out.argmax()) + assert out.argmax() == 7 # cock