More expression tensor operations Make the return expression handler actually work
441 lines
13 KiB
Python
441 lines
13 KiB
Python
import operator
|
|
import traceback
|
|
|
|
from .types import Empty, ExpDict, ExpOp, ExpReturn
|
|
from .util import torch
|
|
from .validation import Arg, ValidateArg, ValidateError
|
|
|
|
|
|
class HandlerError(Exception):
|
|
pass
|
|
|
|
|
|
class HandlerContext:
|
|
def __init__(self, handlers=None, constants=None, variables=None):
|
|
self.handlers = handlers if handlers is not None else {}
|
|
self.constants = constants if constants is not None else {}
|
|
self.variables = variables if variables is not None else {}
|
|
|
|
def get_handler(self, k, default=Empty):
|
|
return self.handlers.get(k, default)
|
|
|
|
def get_var(self, k, default=Empty):
|
|
result = self.constants.get(k, Empty)
|
|
if result is Empty:
|
|
result = self.variables.get(k, Empty)
|
|
return default if result is Empty else result
|
|
|
|
def set_var(self, k, v):
|
|
if k in self.constants:
|
|
raise KeyError(
|
|
f"Cannot set variable with key {k}: already exists as a constant"
|
|
)
|
|
self.variables[k] = v
|
|
|
|
def unset_var(self, k):
|
|
if k in self.variables:
|
|
del self.variables[k]
|
|
return True
|
|
return False
|
|
|
|
def __contains__(self, k):
|
|
return any(
|
|
k in coll for coll in (self.handlers, self.constants, self.variables)
|
|
)
|
|
|
|
def clone(self, *, handlers=Empty, constants=Empty, variables=Empty):
|
|
return self.__class__(
|
|
self.handlers if handlers is Empty else handlers,
|
|
self.constants if constants is Empty else constants,
|
|
self.variables if variables is Empty else variables,
|
|
)
|
|
|
|
|
|
class BaseHandler:
|
|
input_validators = ()
|
|
|
|
def __init__(self):
|
|
self.input_validators_by_key = {
|
|
v.name: (idx, v) for idx, v in enumerate(self.input_validators)
|
|
}
|
|
|
|
def __call__(self, obj, *, getter):
|
|
try:
|
|
val = self.handle(obj, getter)
|
|
except ExpReturn:
|
|
raise
|
|
except Exception as exc:
|
|
tb = traceback.format_exc()
|
|
raise HandlerError(f'Error evaluating "{obj.name}": {exc!s}\n{tb}') from exc
|
|
return self.validate_output(obj, val)
|
|
|
|
def safe_get(self, key, obj, getter=None, *, default=Empty):
|
|
str_key = isinstance(key, str)
|
|
if str_key:
|
|
argidx, validator = self.input_validators_by_key.get(key, (-1, None))
|
|
else:
|
|
argidx, validator = (
|
|
key,
|
|
(
|
|
self.input_validators[key]
|
|
if key < len(self.input_validators)
|
|
else None
|
|
),
|
|
)
|
|
default = (
|
|
default
|
|
if default is not Empty or validator is None
|
|
else getattr(validator, "default", Empty)
|
|
)
|
|
if argidx >= 0 and argidx < len(obj.args):
|
|
eff_key = argidx
|
|
str_eff_key = False
|
|
elif str_key:
|
|
eff_key = key
|
|
str_eff_key = True
|
|
else:
|
|
raise ValidateError(
|
|
f"Error validating input argument {key} for {obj.name}, out of range for actual function arguments"
|
|
)
|
|
if getter is None:
|
|
if str_eff_key:
|
|
val = obj.kwargs.get(eff_key)
|
|
else:
|
|
val = default if eff_key > len(obj.args) else obj.args[eff_key]
|
|
else:
|
|
val = getter(eff_key, default=default)
|
|
if validator is None:
|
|
return val
|
|
try:
|
|
return validator(key, val)
|
|
except ValidateError as exc:
|
|
raise ValidateError(
|
|
f"Error validating input argument {key} for {obj.name}, type {type(val)}: {exc!r}"
|
|
)
|
|
|
|
def safe_get_multi(self, keys, obj, getter=None, *, default=Empty):
|
|
return (self.safe_get(k, obj, getter, default=default) for k in keys)
|
|
|
|
def safe_get_all(self, obj, getter=None, *, default=Empty):
|
|
return self.safe_get_multi(
|
|
(v.name for v in self.input_validators), obj, getter, default=default
|
|
)
|
|
|
|
def handle(self, obj, getter):
|
|
raise NotImplementedError
|
|
|
|
def validate_output(self, obj, value):
|
|
return value
|
|
|
|
|
|
class BinopLogicHandler(BaseHandler):
|
|
input_validators = (
|
|
Arg.present("lhs"),
|
|
Arg.present("rhs"),
|
|
)
|
|
|
|
def validate_output(self, obj, value):
|
|
return operator.truth(value)
|
|
|
|
|
|
class OrHandler(BinopLogicHandler):
|
|
def handle(self, obj, getter):
|
|
return operator.truth(
|
|
self.safe_get("lhs", obj, getter=getter)
|
|
) or operator.truth(self.safe_get("rhs", obj, getter=getter))
|
|
|
|
|
|
class AndHandler(BinopLogicHandler):
|
|
def handle(self, obj, getter):
|
|
return operator.truth(
|
|
self.safe_get("lhs", obj, getter=getter)
|
|
) and operator.truth(self.safe_get("rhs", obj, getter=getter))
|
|
|
|
|
|
class AllHandler(BinopLogicHandler):
|
|
input_validators = ()
|
|
|
|
def handle(self, obj, getter):
|
|
return all(
|
|
operator.truth(self.safe_get(idx, obj, getter=getter))
|
|
for idx in range(len(obj.args))
|
|
) and all(
|
|
operator.truth(self.safe_get(key, obj, getter=getter)) for key in obj.kwargs
|
|
)
|
|
|
|
|
|
class AnyHandler(BinopLogicHandler):
|
|
def handle(self, obj, getter):
|
|
return any(
|
|
operator.truth(self.safe_get(idx, obj, getter=getter))
|
|
for idx in range(len(obj.args))
|
|
) or any(
|
|
operator.truth(self.safe_get(key, obj, getter=getter)) for key in obj.kwargs
|
|
)
|
|
|
|
|
|
class EqHandler(BinopLogicHandler):
|
|
def handle(self, obj, getter):
|
|
a1, a2 = self.safe_get_all(obj, getter)
|
|
if isinstance(a1, torch.Tensor) and isinstance(a2, torch.Tensor):
|
|
return torch.equal(a1, a2)
|
|
return a1 == a2
|
|
|
|
|
|
class NeqHandler(EqHandler):
|
|
def handle(self, *args, **kwargs):
|
|
return not super().handle(*args, **kwargs)
|
|
|
|
|
|
class NotHandler(BinopLogicHandler):
|
|
input_validators = (Arg.present("value"),)
|
|
|
|
def handle(self, obj, getter):
|
|
return not operator.truth(self.safe_get("value", obj, getter=getter))
|
|
|
|
|
|
class IfHandler(BaseHandler):
|
|
input_validators = (
|
|
Arg.present("condition"),
|
|
Arg.present("then"),
|
|
Arg.present("else"),
|
|
)
|
|
|
|
def handle(self, obj, getter):
|
|
if operator.truth(self.safe_get("condition", obj, getter=getter)):
|
|
return self.safe_get("then", obj, getter=getter)
|
|
return self.safe_get("else", obj, getter=getter)
|
|
|
|
|
|
class BetweenHandler(BaseHandler): # Inclusive
|
|
input_validators = (
|
|
Arg.numeric("value"),
|
|
Arg.numeric("from", 0.0),
|
|
Arg.numeric("to"),
|
|
)
|
|
|
|
def handle(self, obj, getter):
|
|
value, low, high = self.safe_get_all(obj, getter)
|
|
if low > high:
|
|
low, high = high, low
|
|
return low <= value <= high
|
|
|
|
|
|
class SimpleMathHandler(BaseHandler):
|
|
input_validators = (Arg.numeric("lhs"), Arg.numeric("rhs"))
|
|
|
|
def __init__(self, handler):
|
|
super().__init__()
|
|
self.handler = handler
|
|
|
|
def validate_output(self, obj, value):
|
|
return ValidateArg.validate_numeric(-1, value)
|
|
|
|
def handle(self, obj, getter):
|
|
args = (
|
|
self.safe_get(idx, obj, getter=getter)
|
|
for idx in range(len(self.input_validators))
|
|
)
|
|
return self.handler(*args)
|
|
|
|
|
|
class MinusHandler(SimpleMathHandler):
|
|
input_validators = (Arg.numeric("lhs"), Arg.numeric("rhs", default=Empty))
|
|
|
|
__init__ = BaseHandler.__init__
|
|
|
|
def handle(self, obj, getter):
|
|
lhs, rhs = self.safe_get_all(obj, getter)
|
|
if rhs is Empty:
|
|
return operator.neg(lhs)
|
|
return operator.sub(lhs, rhs)
|
|
|
|
|
|
class RelComparisonHandler(SimpleMathHandler):
|
|
def validate_output(self, obj, value):
|
|
return operator.truth(value)
|
|
|
|
|
|
class UnarySimpleMathHandler(SimpleMathHandler):
|
|
input_validators = (Arg.numeric("lhs"),)
|
|
|
|
|
|
class IsSetHandler(BaseHandler):
|
|
input_validators = (Arg.string("name"),)
|
|
|
|
def handle(self, obj, getter):
|
|
key = self.safe_get(0, obj, getter=getter)
|
|
return key in getter.ctx
|
|
|
|
def validate_output(self, obj, value):
|
|
return operator.truth(value)
|
|
|
|
|
|
class GetHandler(BaseHandler):
|
|
input_validators = (
|
|
Arg.string("name"),
|
|
Arg.present("fallback"),
|
|
)
|
|
|
|
def handle(self, obj, getter):
|
|
key = self.safe_get("name", obj, getter=getter)
|
|
result = getter.ctx.get_var(key)
|
|
if result is Empty:
|
|
return self.safe_get("fallback", obj, getter=getter)
|
|
return ExpOp(key).eval(getter.ctx, *getter.args, **getter.kwargs)
|
|
|
|
|
|
class S_Handler(BaseHandler):
|
|
input_validators = (
|
|
Arg.one_of(
|
|
"start",
|
|
(ValidateArg.validate_none, ValidateArg.validate_integer),
|
|
default=None,
|
|
),
|
|
Arg.one_of(
|
|
"end",
|
|
(ValidateArg.validate_none, ValidateArg.validate_integer),
|
|
default=None,
|
|
),
|
|
Arg.integer("step", 1),
|
|
)
|
|
|
|
def handle(self, obj, getter):
|
|
return slice(*self.safe_get_all(obj, getter=getter))
|
|
|
|
|
|
class IndexHandler(BaseHandler):
|
|
input_validators = (
|
|
Arg.present("index"),
|
|
Arg.one_of(
|
|
"value", (ValidateArg.validate_sequence, ValidateArg.validate_tensor)
|
|
),
|
|
)
|
|
|
|
def handle(self, obj, getter):
|
|
idx, value = self.safe_get_all(obj, getter=getter)
|
|
return value[idx]
|
|
|
|
|
|
class MinHandler(BaseHandler):
|
|
input_validators = (Arg.numscalar_sequence("values"),)
|
|
|
|
def handle(self, obj, getter):
|
|
return min(*self.safe_get("values", obj, getter))
|
|
|
|
def validate_output(self, obj, value):
|
|
return ValidateArg.validate_numeric(-1, value)
|
|
|
|
|
|
class MaxHandler(MinHandler):
|
|
def handle(self, obj, getter):
|
|
return max(*self.safe_get("values", obj, getter))
|
|
|
|
|
|
class UnsafeCallHandler(BaseHandler):
|
|
input_validators = (Arg.present("__callable"),)
|
|
|
|
def handle(self, obj, getter):
|
|
if "__callable" in obj.kwargs:
|
|
raise ValueError(
|
|
"unsafe_call does not support passing the callable via keyword arg"
|
|
)
|
|
fun = self.safe_get("__callable", obj, getter)
|
|
if not callable(fun):
|
|
raise ValueError("Cannot call supplied value: not a callable")
|
|
args = (self.safe_get(idx, obj, getter) for idx in range(1, len(obj.args)))
|
|
kwargs = {k: self.safe_get(k, obj, getter) for k in obj.kwargs}
|
|
return fun(*args, **kwargs)
|
|
|
|
|
|
class DictHandler(BaseHandler):
|
|
def handle(self, obj, getter):
|
|
if len(obj.args):
|
|
raise ValueError("Non-KV items passed to dict constructor")
|
|
return ExpDict({k: self.safe_get(k, obj, getter) for k in obj.kwargs.keys()})
|
|
|
|
|
|
class CommentHandler(BaseHandler):
|
|
def handle(self, obj, getter):
|
|
return None
|
|
|
|
|
|
class SetVarHandler(BaseHandler):
|
|
input_validators = (Arg.string("lhs"), Arg.present("rhs"))
|
|
|
|
def handle(self, obj, getter):
|
|
key, val = self.safe_get_all(obj, getter)
|
|
getter.ctx.set_var(key, val)
|
|
return val
|
|
|
|
|
|
class ReturnHandler(BaseHandler):
|
|
input_validators = (Arg.present("expression"),)
|
|
|
|
def handle(self, obj, getter):
|
|
raise ExpReturn(self.safe_get("expression", obj, getter))
|
|
|
|
|
|
LOGIC_HANDLERS = {
|
|
"||": OrHandler(),
|
|
"&&": AndHandler(),
|
|
"==": EqHandler(),
|
|
"!=": NeqHandler(),
|
|
"not": NotHandler(),
|
|
"if": IfHandler(),
|
|
"all": AllHandler(),
|
|
"any": AnyHandler(),
|
|
}
|
|
for k, alias in (
|
|
("||", "or"),
|
|
("&&", "and"),
|
|
("==", "eq"),
|
|
("!=", "neq"),
|
|
):
|
|
LOGIC_HANDLERS[alias] = LOGIC_HANDLERS[k]
|
|
|
|
|
|
MATH_HANDLERS = {
|
|
"+": SimpleMathHandler(operator.add),
|
|
"-": MinusHandler(),
|
|
"*": SimpleMathHandler(operator.mul),
|
|
"/": SimpleMathHandler(operator.truediv),
|
|
"//": SimpleMathHandler(operator.floordiv),
|
|
"**": SimpleMathHandler(operator.pow),
|
|
"mod": SimpleMathHandler(operator.mod),
|
|
"neg": UnarySimpleMathHandler(operator.neg),
|
|
"between": BetweenHandler(),
|
|
"<": RelComparisonHandler(operator.lt),
|
|
"<=": RelComparisonHandler(operator.le),
|
|
">": RelComparisonHandler(operator.gt),
|
|
">=": RelComparisonHandler(operator.ge),
|
|
"min": MinHandler(),
|
|
"max": MaxHandler(),
|
|
"float": UnarySimpleMathHandler(handler=float),
|
|
"int": UnarySimpleMathHandler(handler=int),
|
|
"bool": UnarySimpleMathHandler(handler=bool),
|
|
}
|
|
for k, alias in (
|
|
("+", "add"),
|
|
("-", "sub"),
|
|
("*", "mul"),
|
|
("/", "div"),
|
|
("//", "idiv"),
|
|
("**", "pow"),
|
|
):
|
|
MATH_HANDLERS[alias] = MATH_HANDLERS[k]
|
|
|
|
MISC_HANDLERS = {
|
|
"is_set": IsSetHandler(),
|
|
"get": GetHandler(),
|
|
"index": IndexHandler(),
|
|
"s_": S_Handler(),
|
|
"unsafe_call": UnsafeCallHandler(),
|
|
"dict": DictHandler(),
|
|
"comment": CommentHandler(),
|
|
"set_var": SetVarHandler(),
|
|
"return": ReturnHandler(),
|
|
}
|
|
|
|
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
|