diff --git a/tinygrad/engine/schedule.py b/tinygrad/engine/schedule.py index af79a7b879..e1abe507d3 100644 --- a/tinygrad/engine/schedule.py +++ b/tinygrad/engine/schedule.py @@ -206,8 +206,10 @@ def create_kernel(ctx:KernelContext, x:UOp): if x not in ctx.realizes: return None assert isinstance(x.device, str), f"buf device in kernel must be string {x.device}" b = x.buf_uop if x.op is Ops.ASSIGN else UOp.new_buffer(x.device, x.size, x.dtype) + output_st = ShapeTracker.from_shape(x.shape) # KERNEL nodes become: ASSIGN(VIEW(BUFFER), KERNEL) - return b.view(ShapeTracker.from_shape(x.shape)).assign(UOp(Ops.KERNEL, src=x.src, arg=Kernel(x, (m,) if (m:=ctx.ops_metadata.get(x)) else ()))) + # TODO: this should be ASSIGN(BUFFER, KERNEL) followed by the output ShapeTracker + return b.view(output_st).assign(UOp(Ops.KERNEL, src=(b,)+x.src, arg=Kernel(x, (m,) if (m:=ctx.ops_metadata.get(x)) else ()))) def append_to_kernel(ctx:KernelContext, x:UOp): new_srcs: list[UOp] = []