Add t_shape expression function

Fix width/height order for t_scale in absolute mode
This commit is contained in:
blepping
2024-08-08 14:10:25 -06:00
parent 764f04a6a7
commit 442bde294a
2 changed files with 13 additions and 1 deletions
+2
View File
@@ -135,6 +135,8 @@ Available in model filters, with the exception of the `input` filter.
| <td colspan=3 align=left>Rolls a tensor along the specified dimensions. If amount is >= -1.0 and < 1.0 this will be interpreted as a percentage. <br/> **Example:** `t_roll(some_tensor, 10, (-2,))`</td> |
|⬤| `t_scale` | tensor:`T`, scale:`SN \| NS`, mode:`SY(bicubic)`, absolute_scale:`B(false)` | `T` |
| <td colspan=3 align=left>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. <br/> Example: `t_scale(some_tensor, (0.75, 0.5), 'bilinear)`</td> |
|⬤| `t_shape` | tensor:`T` | `SN` |
| <td colspan=3 align=left>Returns a tensor's shape as a tuple. <br/> Example: `shp := t_shape(some_tensor); width := shp[-1]; height := shp[-2]`</td> |
|⬤| `t_std` | tensor:`T`, dim:`SN(-3, -2, -1)` | `T` |
| <td colspan=3 align=left>Tensor std, second argument is dimensions. <br/> **Example:** `t_std(some_tensor, (-2, -1))`</td> |
|⬤| `unsafe_tensor_method` | `T`, `SY`, `*`\* | `*` |
+11 -1
View File
@@ -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(),
}