fix map()

This commit is contained in:
mcDandy
2025-12-29 21:34:12 +01:00
parent 744931c6e3
commit 47f4d1309d
2 changed files with 17 additions and 24 deletions
+1 -3
View File
@@ -85,9 +85,7 @@ You can also get the node from comfy manager under the name of More math.
### Advanced Tensor Operations (Tensor Only)
- `map(tensor, c1, ...)`: Remaps `tensor` using source coordinates.
- Up to 3 coordinate mapping functions can be provided which map to the last (up to 3) dimensions of the tensor.
- If less than 3 functions are provided and shape of tensor >= 3, the remaining dimensions are assumed to be identity functions.
- That means map(tensor4,func) where tensor4 has shape [a,b,c,d] will map 2nd dimension and map(tensor4,func,func) will map 2nd and 3rd dimension. map(tensor4,func,func,func) will map last 3 dimensions.
- Up to 3 coordinate mapping functions can be provided which map to the last (up to 3) dimensions of the tensor. Rest uses identity mapping.
- `conv(tensor, kw, [kh], [kd], k_expr)`: Applies a convolution to `tensor`.
- `k_expr` can be a math expression (using `kX`, `kY`, `kZ`) or a list literal.
+16 -21
View File
@@ -74,7 +74,7 @@ class TensorEvalVisitor(MathExprVisitor):
return tsr.reshape(*new_shape)
def visitNumberExp(self, ctx):
return torch.tensor(float(ctx.getText()), device=self.device)
return torch.full(self.shape,float(ctx.getText()), device=self.device)
def visitConstantExp(self, ctx):
name = ctx.getText().lower()
@@ -388,48 +388,43 @@ class TensorEvalVisitor(MathExprVisitor):
batch_size = 1
for s in leading_shape: batch_size *= s
input_view = tensor.reshape(batch_size, 1, *spatial_in_shape)
norm_coords_list = []
for i in range(num_coords):
dim_size = spatial_in_shape[-(i+1)]
dim_size = spatial_in_shape[i]
norm = self._normalize_coord(coords[i], dim_size)
norm_coords_list.append(norm)
broadcasted = torch.broadcast_tensors(*norm_coords_list)
grid = torch.stack(broadcasted, dim=-1)
grid = torch.stack(norm_coords_list[::-1], dim=-1)
grid_spatial_shape = grid.shape[:-1]
if grid.numel() // max(2, num_coords) >= batch_size and \
grid.shape[:len(leading_shape)] == leading_shape:
grid_view = grid.reshape(batch_size, -1, num_coords)
else:
try:
grid_view = grid.reshape(batch_size, *grid_spatial_shape[-(num_coords):], num_coords)
except RuntimeError:
print("Reshape failed in map(); attempting expand workaround.")
grid_view = grid.expand(batch_size, *([-1] * len(grid_spatial_shape)), -1)
grid_view = grid_view.reshape(batch_size, -1, num_coords)
grid_view = grid_view.reshape(batch_size, *grid_view.shape[-(num_coords+1):-1], num_coords)
if num_coords == 1:
y_zeros = torch.zeros_like(grid_view[..., :1])
grid_final = torch.cat([grid_view, y_zeros], dim=-1).reshape(batch_size, 1, -1, 2)
input_final = input_view.reshape(batch_size, 1, 1, -1)
y_zeros = torch.zeros_like(grid_view[..., :1])
grid_final = torch.cat([grid_view, y_zeros], dim=-1).unsqueeze(1)
output = torch.nn.functional.grid_sample(input_final, grid_final, align_corners=True)
elif num_coords == 2:
grid_final = grid_view.reshape(batch_size, 1, -1, 2)
grid_final = grid_view.reshape(batch_size, *grid_view.shape[-3:-1], 2)
output = torch.nn.functional.grid_sample(input_view, grid_final, align_corners=True)
else: # 3D
grid_final = grid_view.reshape(batch_size, 1, 1, -1, 3)
else:
grid_final = grid_view.reshape(batch_size, *grid_view.shape[-4:-1], 3)
output = torch.nn.functional.grid_sample(input_view, grid_final, align_corners=True)
actual_spatial = grid_spatial_shape
if len(grid_spatial_shape) >= len(leading_shape) and \
grid_spatial_shape[:len(leading_shape)] == leading_shape:
actual_spatial = grid_spatial_shape[len(leading_shape):]
actual_spatial = grid_view.shape[1:-1]
final_shape = list(leading_shape) + list(actual_spatial)
return output.reshape(final_shape)
def _kernel_coords(self, size, device):
half = size // 2
return torch.arange(size, device=device).float() - half