fix logspace not using base

This commit is contained in:
Daniel Martinek
2026-04-05 13:00:50 +02:00
parent f0721e09c9
commit 58287392fb
+2 -2
View File
@@ -3498,14 +3498,14 @@ class UnifiedMathVisitor(MathExprVisitor):
start_val = yield ctx.expr(0)
end_val = yield ctx.expr(1)
steps_val = yield ctx.expr(2)
base_val = yield ctx.expr(2)
base_val = yield ctx.expr(3)
start = float(start_val.item()) if self._is_tensor(start_val) else float(start_val)
end = float(end_val.item()) if self._is_tensor(end_val) else float(end_val)
base = float(base_val.item()) if self._is_tensor(base_val) else float(base_val)
steps = int(steps_val.item()) if self._is_tensor(steps_val) else int(steps_val)
return torch.logspace(start, end, steps, device=self.device)
return torch.logspace(start, end, steps, base=base, device=self.device)
def visitRollFunc(self, ctx):
"""roll(x, shift, [dim]) - circular shift of elements"""