fix logspace not using base
This commit is contained in:
@@ -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"""
|
||||
|
||||
Reference in New Issue
Block a user