no CONST(DEVICE) in create_allreduce_function (#16460)

This commit is contained in:
chenyu
2026-06-01 17:12:34 -04:00
committed by GitHub
parent 7e7b481ba7
commit 517eea5985
+1 -1
View File
@@ -55,7 +55,7 @@ def handle_allreduce(buf:UOp, red:UOp) -> UOp|None:
return UOp.usum(*[c.pad(((s,numel-e),)) for (s,e),c in zip(chunks, copied_chunks)]).reshape(shape)
def create_allreduce_function(buf:UOp, red:UOp, output:UOp|None=None) -> UOp|None:
if output is None: output = UOp.const(red.dtype, Invalid, red.device, red.shape).clone()
if output is None: output = UOp.const(red.dtype, Invalid, shape=red.shape).clone(device=red.device)
to = red.param_like(0)
src = buf.param_like(1)
red = src.allreduce(red.arg, red.src[1])