Skip to content

math_spec.dimensions

Static dim-set checking — a type system whose type is a set of dim names.

Every node's dim set is computable before any data is bound, so this pass runs at load on the resolved tree. The per-node rules are the "Dim algebra" table in docs/reference/language/expressions.md; a constraint's two sides together must equal its dims, and a where or a bound may not exceed the frame.

check_schema(schema, program) #

Check every declaration's dim rules, on the trees program holds for schema.

RAISES DESCRIPTION
DimensionError

On the first declaration that breaks one.

Source code in src/math_spec/dimensions.py
def check_schema(schema: Spec, program: Program) -> None:
    """Check every declaration's dim rules, on the trees *program* holds for *schema*.

    Raises:
        DimensionError: On the first declaration that breaks one.
    """
    for vname, vdef in schema.variables.items():
        frame = frozenset(vdef.dims)
        context = f"Variable '{vname}'"
        _check_where_dims(program.variables[vname].where, frame, context)
        for side in ('lower', 'upper'):
            bound = getattr(vdef.bounds, side)
            if isinstance(bound, str):
                bdims = frozenset(schema.parameters[bound].dims)
                if not bdims <= frame:
                    raise DimensionError(
                        f"{context}: bounds.{side} parameter '{bound}' has dims "
                        f"{sorted(bdims - frame)} outside the variable's dims "
                        f'{sorted(frame)}.'
                    )

    for ename, entry in program.expressions.items():
        if not isinstance(entry.expression, Cases):
            continue
        block = schema.expressions[ename]
        frame = frozenset(block.dims or [])
        for region, label in zip(entry.expression.regions, [*block.cases, None], strict=True):
            context = case_context(ename, label)
            if label is not None:
                _check_where_dims(region.when, frame, context)
            _check_value_dims(region.value, schema, frame, context)

    for cname, constraint in program.constraints.items():
        frame = frozenset(constraint.dims)
        context = f"Constraint '{cname}'"
        _check_where_dims(constraint.where, frame, context)
        got = dims_of(constraint.lhs, schema, context) | dims_of(constraint.rhs, schema, context)
        if got != frame:
            stray, missing = sorted(got - frame), sorted(frame - got)
            detail = (
                f'carries dims {stray} that are not in its dims: {sorted(frame)} — every '
                f'stray dim multiplies the rows this constraint builds; add it to '
                f'dims: if that is intended, or sum it out'
                if stray
                else f'does not carry {missing}, which its dims: declares — the same row '
                f'would be repeated across {missing}; drop it from dims:, or use it '
                f'in the expression'
            )
            raise DimensionError(f'{context}: the expression {detail}.')

    if program.objective is not None:
        context = 'The objective'
        got = dims_of(program.objective.expression, schema, context)
        if got:
            raise DimensionError(
                f'{context}: the expression carries dims {sorted(got)}, and an objective is one '
                f'number. Wrap each additive term in its own sum(): '
                f'`sum(p * cost) + sum(p_nom * capex)`.'
            )

dims_of(node, schema, context) #

The dim set of a resolved expression, checking every rule on the way.

RAISES DESCRIPTION
DimensionError

On the first rule broken.

Source code in src/math_spec/dimensions.py
def dims_of(node: Expression, schema: Spec, context: str) -> frozenset[str]:
    """The dim set of a resolved expression, checking every rule on the way.

    Raises:
        DimensionError: On the first rule broken.
    """
    if isinstance(node, Constant):
        return frozenset()

    if isinstance(node, Parameter):
        return frozenset(schema.parameters[node.name].dims)

    if isinstance(node, Variable):
        return frozenset(schema.variables[node.name].dims)

    if isinstance(node, Dual):
        return frozenset(schema.constraints[node.constraint].dims)

    if isinstance(node, Named):
        return _named_dims(node, schema, context)

    if isinstance(node, Cases):
        return frozenset().union(*(dims_of(region.value, schema, context) for region in node.regions))

    if isinstance(node, Negate | Add | Multiply | Power | Divide):
        return frozenset().union(*(dims_of(child, schema, context) for child in children(node)))

    inner = dims_of(node.operand, schema, context)
    if isinstance(node, Sum):
        return _sum_dims(node, inner, context)
    if isinstance(node, GroupSum):
        return _group_sum_dims(node, inner, context)
    if isinstance(node, Pullback):
        return _at_dims(node, inner, context)
    if isinstance(node, Translate | WindowSum):
        return _translation_dims(node, inner, schema, context)

    assert_never(node)

pulled_back_dims(direction, inner, context, operand) #

The dims inner has once at reads it through direction, an expression's or a predicate's alike.

RAISES DESCRIPTION
DimensionError

operand does not carry a dim the read consumes or joins on, or already carries one it lands on.

Source code in src/math_spec/dimensions.py
def pulled_back_dims(direction: Direction, inner: frozenset[str], context: str, operand: str) -> frozenset[str]:
    """The dims *inner* has once ``at`` reads it through *direction*, an expression's or a predicate's alike.

    Raises:
        DimensionError: *operand* does not carry a dim the read consumes or
            joins on, or already carries one it lands on.
    """
    if absent := sorted(set(direction.consumed_dims) - inner):
        raise DimensionError(
            f'{context}: at(by={direction.name}) reads through '
            f'{absent}, which {operand} does not carry (dims '
            f'{sorted(inner)}). A pullback needs the coarse dims to read *from* — '
            f'sum is the direction that produces them.'
        )
    return _read_dims(f'at(by={direction.name})', direction, inner, context)