281 lines
9.1 KiB
Python
281 lines
9.1 KiB
Python
import contextlib
|
|
import functools
|
|
|
|
from ..latent import ImageBatch
|
|
from .types import Empty
|
|
from .util import torch
|
|
|
|
|
|
class Arg:
|
|
__slots__ = ("default", "name", "validator")
|
|
|
|
def __init__(self, name, default=Empty, *, validator=None):
|
|
self.name = name
|
|
self.default = default
|
|
self.validator = validator
|
|
|
|
def __call__(self, _key, value, *args, **kwargs):
|
|
return self.validate(value, *args, **kwargs)
|
|
|
|
def validate(self, value):
|
|
if value is Empty:
|
|
if self.default is Empty:
|
|
raise ValueError(f"Missing value for argument {self.name}")
|
|
return self.default
|
|
try:
|
|
return self.validator(self.name, value) if self.validator else value
|
|
except ValidateError as exc:
|
|
raise ValidateError(f"Failed to validate argument {self.name}: {exc}")
|
|
|
|
@classmethod
|
|
def tensor(cls, name):
|
|
return cls(name, validator=ValidateArg.validate_tensor)
|
|
|
|
@classmethod
|
|
def image(cls, name):
|
|
return cls(name, validator=ValidateArg.validate_image)
|
|
|
|
@classmethod
|
|
def numeric(cls, name, default=Empty):
|
|
return cls(name, default=default, validator=ValidateArg.validate_numeric)
|
|
|
|
@classmethod
|
|
def numeric_scalar(cls, name, default=Empty):
|
|
return cls(name, default=default, validator=ValidateArg.validate_numeric_scalar)
|
|
|
|
@classmethod
|
|
def integer(cls, name, default=Empty):
|
|
return cls(name, default=default, validator=ValidateArg.validate_integer)
|
|
|
|
@classmethod
|
|
def numscalar_sequence(cls, name, default=Empty):
|
|
return cls(
|
|
name, default=default, validator=ValidateArg.validate_numscalar_sequence
|
|
)
|
|
|
|
@classmethod
|
|
def numeric_sequence(cls, name, default=Empty):
|
|
return cls(
|
|
name, default=default, validator=ValidateArg.validate_numeric_sequence
|
|
)
|
|
|
|
@classmethod
|
|
def numscalar_sequence_or_single(cls, name, default=Empty):
|
|
return cls.one_of(
|
|
name,
|
|
(
|
|
ValidateArg.validate_numscalar_sequence,
|
|
ValidateArg.validate_numeric_scalar,
|
|
),
|
|
default=default,
|
|
)
|
|
|
|
@classmethod
|
|
def tensor_slice(cls, name, default=Empty):
|
|
return cls(name, default=default, validator=ValidateArg.validate_tensor_slice)
|
|
|
|
@classmethod
|
|
def sequence(cls, name, default=Empty, *, item_validator=None):
|
|
return cls(
|
|
name,
|
|
default=default,
|
|
validator=functools.partial(
|
|
ValidateArg.validate_sequence, item_validator=item_validator
|
|
),
|
|
)
|
|
|
|
@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)
|
|
|
|
@classmethod
|
|
def boolean(cls, name, default=Empty):
|
|
return cls(name, default=default, validator=ValidateArg.validate_boolean)
|
|
|
|
@classmethod
|
|
def present(cls, name, default=Empty):
|
|
return cls(name, default=default, validator=ValidateArg.validate_passthrough)
|
|
|
|
@classmethod
|
|
def one_of(cls, name, validators, *, default=Empty):
|
|
def validate(idx, val):
|
|
for validator in validators:
|
|
try:
|
|
return validator(idx, val)
|
|
except ValidateError:
|
|
continue
|
|
raise ValidateError(
|
|
f"Failed to validate argument at {idx} of type {type(val)}"
|
|
)
|
|
|
|
return cls(name, default=default, validator=validate)
|
|
|
|
|
|
class ValidateError(Exception):
|
|
pass
|
|
|
|
|
|
class ValidateArg:
|
|
__slots__ = ("groupfun", "kwargs", "kwargslist", "valfuns")
|
|
|
|
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
|
|
if not isinstance(name, (list, tuple)):
|
|
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")
|
|
self.groupfun = group
|
|
self.kwargs = kwargs
|
|
self.kwargslist = kwargslist if kwargslist is not None else {}
|
|
|
|
def __call__(self, *args, **kwargs):
|
|
kalen = len(self.kwargslist)
|
|
return self.groupfun(
|
|
vf(
|
|
*args,
|
|
**(self.kwargslist if idx < kalen else {}),
|
|
**self.kwargs,
|
|
)
|
|
for idx, vf in enumerate(self.valfuns)
|
|
)
|
|
|
|
@staticmethod
|
|
def validate_numeric(idx, val):
|
|
if not isinstance(val, (int, float, torch.Tensor)):
|
|
raise ValidateError(
|
|
f"Expected numeric or tensor argument at {idx}, got {type(val)}"
|
|
)
|
|
return val
|
|
|
|
@classmethod
|
|
def validate_numeric_scalar(cls, idx, val):
|
|
if not isinstance(val, (int, float)):
|
|
raise ValidateError(f"Expected numeric argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@classmethod
|
|
def validate_tensor_slice_item(cls, idx, val):
|
|
with contextlib.suppress(ValidateError):
|
|
ok = (
|
|
val in {Ellipsis, None}
|
|
or isinstance(val, (int, slice))
|
|
or cls.validate_sequence(
|
|
idx, val, item_validator=ValidateArg.validate_integer
|
|
)
|
|
)
|
|
if ok:
|
|
return val
|
|
raise ValidateError(
|
|
f"Expected none, int, slice, tuple of int or ellipsis argument at {idx}, got {type(val)}"
|
|
)
|
|
|
|
@classmethod
|
|
def validate_integer(cls, idx, val):
|
|
if not isinstance(val, int):
|
|
raise ValidateError(f"Expected integer argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@staticmethod
|
|
def validate_tensor(idx, val):
|
|
if not isinstance(val, torch.Tensor):
|
|
raise ValidateError(f"Expected tensor argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@staticmethod
|
|
def validate_image(idx, val):
|
|
if not isinstance(val, ImageBatch):
|
|
raise ValidateError(
|
|
f"Expected PIL Image argument at {idx}, got {type(val)}"
|
|
)
|
|
return val
|
|
|
|
@staticmethod
|
|
def validate_sequence(idx, val, *, item_validator=None):
|
|
if not isinstance(val, (list, tuple)):
|
|
raise ValidateError(f"Expected sequence argument at {idx}, got {type(val)}")
|
|
if item_validator is None:
|
|
return val
|
|
try:
|
|
return tuple(item_validator(iidx, v) for iidx, v in enumerate(val))
|
|
except ValidateError as 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):
|
|
return cls.validate_sequence(
|
|
idx, val, item_validator=cls.validate_numeric_scalar
|
|
)
|
|
|
|
@classmethod
|
|
def validate_numeric_sequence(cls, idx, val):
|
|
return cls.validate_sequence(idx, val, item_validator=cls.validate_numeric)
|
|
|
|
@classmethod
|
|
def validate_tensor_slice(cls, idx, val):
|
|
return cls.validate_sequence(
|
|
idx, val, item_validator=cls.validate_tensor_slice_item
|
|
)
|
|
|
|
@classmethod
|
|
def validate_string(cls, idx, val):
|
|
if not isinstance(val, str):
|
|
raise ValidateError(f"Expected string argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@classmethod
|
|
def validate_dict(cls, idx, val):
|
|
if not isinstance(val, dict):
|
|
raise ValidateError(f"Expected dict argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@classmethod
|
|
def validate_boolean(cls, idx, val):
|
|
if val is not True and val is not False:
|
|
raise ValidateError(f"Expected boolean argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@classmethod
|
|
def validate_none(cls, idx, val):
|
|
if val is not None:
|
|
raise ValidateError(f"Expected none argument at {idx}, got {type(val)}")
|
|
return val
|
|
|
|
@classmethod
|
|
def validate_passthrough(cls, idx, val):
|
|
return val
|