Add t_new_like expression function

This commit is contained in:
blepping
2025-05-19 01:48:54 -06:00
parent f514274656
commit 6fc85f55f9
2 changed files with 54 additions and 2 deletions
+37 -2
View File
@@ -66,6 +66,16 @@ class Arg:
),
)
@classmethod
def nested_sequence(cls, name, default=Empty, *, item_validator=None):
return cls(
name,
default=default,
validator=functools.partial(
ValidateArg.validate_nested_sequence, item_validator=item_validator
),
)
@classmethod
def string(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_string)
@@ -102,7 +112,10 @@ class ValidateArg:
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
if not isinstance(name, (list, tuple)):
return self.__init__((name,), (args,), group=group, kwargslist=kwargs)
name = (name,)
args = ((args,),)
kwargslist = kwargs
kwargs = {}
self.valfuns = (getattr(self, f"validate_{n}", None) for n in name)
if not all(self.valfuns):
raise ValueError("Unknown validator")
@@ -172,7 +185,29 @@ class ValidateArg:
try:
return tuple(item_validator(iidx, v) for iidx, v in enumerate(val))
except ValidateError as exc:
raise ValidateError(f"Item validation failed for in sequence: {exc}")
raise ValidateError(
f"Item validation failed for sequence argument at {idx}: {exc}"
)
@classmethod
def validate_nested_sequence(cls, idx, val, *, item_validator=None, depth=0):
if not isinstance(val, (list, tuple)):
raise ValidateError(
f"Expected nested sequence argument at {idx}, depth {depth} but got {type(val)}"
)
try:
return tuple(
cls.validate_nested_sequence(
idx, v, item_validator=item_validator, depth=depth + 1
)
if isinstance(v, (list, tuple))
else (item_validator(iidx, v) if item_validator is not None else v)
for iidx, v in enumerate(val)
)
except ValidateError as exc:
raise ValidateError(
f"Item validation failed for nested sequence argument at {idx}, depth {depth}: {exc}"
)
@classmethod
def validate_numscalar_sequence(cls, idx, val):
+17
View File
@@ -147,6 +147,22 @@ class IndexedCopyHandler(NormHandler):
return result
class NewLikeHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.nested_sequence(
"values",
default=(),
item_validator=expr.ValidateArg.validate_numeric_scalar,
),
)
def handle(self, obj, getter):
tensor, values = self.safe_get_all(obj, getter)
print(f"\nNEW TENSOR: {values}")
return torch.tensor(values, dtype=tensor.dtype, device=tensor.device)
class MeanHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
@@ -617,6 +633,7 @@ TENSOR_OP_HANDLERS = {
"t_cat": CatHandler(),
"t_stack": StackHandler(),
"t_indexed_copy": IndexedCopyHandler(),
"t_new_like": NewLikeHandler(),
"t_mean": MeanHandler(),
"t_std": StdHandler(),
"t_blend": BlendHandler(),