diff --git a/docs/expression.md b/docs/expression.md index 5f8482c..3fd1bbf 100644 --- a/docs/expression.md +++ b/docs/expression.md @@ -135,6 +135,8 @@ Available in model filters, with the exception of the `input` filter. | Rolls a tensor along the specified dimensions. If amount is >= -1.0 and < 1.0 this will be interpreted as a percentage.
**Example:** `t_roll(some_tensor, 10, (-2,))` | |⬤| `t_scale` | tensor:`T`, scale:`SN \| NS`, mode:`SY(bicubic)`, absolute_scale:`B(false)` | `T` | | Scales a tensor. If scale is a tuple, it will be interpreted as `(height, width)`. When `absolute_scale` is not set, the scales will be interpreted as percentages otherwise absolute values will be used.
Example: `t_scale(some_tensor, (0.75, 0.5), 'bilinear)` | + |⬤| `t_shape` | tensor:`T` | `SN` | + | Returns a tensor's shape as a tuple.
Example: `shp := t_shape(some_tensor); width := shp[-1]; height := shp[-2]` | |⬤| `t_std` | tensor:`T`, dim:`SN(-3, -2, -1)` | `T` | | Tensor std, second argument is dimensions.
**Example:** `t_std(some_tensor, (-2, -1))` | |⬤| `unsafe_tensor_method` | `T`, `SY`, `*`\* | `*` | diff --git a/py/expression_handlers.py b/py/expression_handlers.py index 3cd3d61..d0a2dd8 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -159,7 +159,8 @@ class ScaleHandler(NormHandler): if abs_scale: scale = tuple(int(v) for v in scale) else: - scale = (int(t.shape[-1] * scale[0]), int(t.shape[-2] * scale[1])) + scale = (int(t.shape[-2] * scale[0]), int(t.shape[-1] * scale[1])) + print("SCALE", t.shape[-2:], "->", scale) if not all(v > 0 for v in scale): raise ValueError(f"Invalid scale: scale values must be > 0, got: {scale!r}") return latent.scale_samples(t, scale[1], scale[0], mode=mode) @@ -181,6 +182,14 @@ class NoiseHandler(NormHandler): return ns(s, sn) +class ShapeHandler(expr.BaseHandler): + input_validators = (expr.Arg.tensor("tensor"),) + + def handle(self, obj, getter): + t = self.safe_get("tensor", obj, getter) + return expr.types.ExpTuple((*t.shape,)) + + class UnsafeTorchTensorMethodHandler(NormHandler): input_validators = ( expr.Arg.tensor("__tensor"), @@ -613,6 +622,7 @@ TENSOR_OP_HANDLERS = { "t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(), "t_scale": ScaleHandler(), "t_noise": NoiseHandler(), + "t_shape": ShapeHandler(), "unsafe_tensor_method": UnsafeTorchTensorMethodHandler(), "unsafe_torch": UnsafeTorchHandler(), }