Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 47 additions & 0 deletions src/math_spec/dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,6 +222,7 @@ def _dims_call(
raise DimensionError(
f'{context}: {node.name}(over={over.name}) but the expression has dims {sorted(inner)}.'
)
_check_named_amount(node, over.name, schema, context)
partition = node.kwargs.get('by')
if partition is not None:
assert isinstance(partition, LookupNode)
Expand All @@ -245,6 +246,52 @@ def _dims_call(
raise DimensionError(msg)


#: The kwarg each axis-walking operator takes its amount in, and the word its
#: errors call that amount. ``shift`` reaches by an offset and a window reaches
#: over a width, but both count positions along the axis, so both obey the two
#: rules below.
_AMOUNTS = {'shift': ('offset', 'offset'), 'sum_back': ('within', 'width')}

#: Why a named amount that varies along the axis it walks is not the thing its
#: operator claims to be.
_VARIES = {
'shift': 'a permutation rather than a lag',
'sum_back': 'a different window at every position, which is no longer "the last n"',
}


def _check_named_amount(node: FunctionCallNode, over: str, schema: Model, context: str) -> None:
"""The two rules that hold of an ``offset=`` or ``within=`` naming a parameter.

Both are about the *amount*, not about a dim set, but they live here
because here is where the schema is in hand — a parameter's ``dtype`` and
its ``dims`` are read off the same declaration, and splitting the pair
across two passes would give one rule of a documented pair two voices.

A literal is checked by the grammar already: it parses as a number and a
number carries no dims, so only a named amount can break either rule.
"""
kwarg, noun = _AMOUNTS[node.name]
amount = node.kwargs[kwarg]
if not isinstance(amount, ParameterNode):
return
declared = schema.parameters[amount.name]
if declared.dtype != 'int':
raise DimensionError(
f'{context}: {node.name}({kwarg}={amount.name}) counts positions along '
f"'{over}', but '{amount.name}' is declared dtype: {declared.dtype}. A count of "
f'positions is integral — declare it dtype: int, which binds only an integer '
f'column, so a fractional {noun} has nowhere to arrive from.'
)
if over in declared.dims:
raise DimensionError(
f'{context}: {node.name}({kwarg}={amount.name}) walks '
f"'{over}', but '{amount.name}' is declared over {sorted(declared.dims)}, which "
f'carries it. A named {noun} that varies along the axis it walks is {_VARIES[node.name]} '
f"— declare '{amount.name}' over dims '{over}' is not one of."
)


# ---------------------------------------------------------------------------
# declaration-level rules
# ---------------------------------------------------------------------------
Expand Down
30 changes: 30 additions & 0 deletions tests/test_dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,10 @@
'p_max': {'dims': ['generator']},
'cost': {'dims': ['generator']},
'load': {'dims': ['snapshot', 'bus']},
# a named offset or width counts positions, so the two rules about it
# need a parameter that obeys them and one that spans the axis walked
'spinup': {'dims': ['generator'], 'dtype': 'int'},
'horizon': {'dims': ['snapshot'], 'dtype': 'int'},
},
'variables': {'p': {'foreach': ['snapshot', 'generator'], 'bounds': {'lower': 0, 'upper': 'p_max'}}},
'constraints': {
Expand Down Expand Up @@ -80,6 +84,8 @@ def test_the_base_model_typechecks():
('sum(p * cost, over=generator)', {'snapshot'}),
('sum(p, by=gen_bus)', {'snapshot', 'bus'}),
("shift(p, over=snapshot, offset=1, edge='wrap')", {'snapshot', 'generator'}),
("shift(p, over=snapshot, offset=spinup, edge='wrap')", {'snapshot', 'generator'}),
('sum_back(p, over=snapshot, within=spinup)', {'snapshot', 'generator'}),
],
)
def test_dim_inference(expr, expected):
Expand Down Expand Up @@ -127,6 +133,30 @@ def test_dim_inference(expr, expected):
r'shift\(over=snapshot\) but the expression has dims',
id='shift-requires-the-dim',
),
# A named offset or width counts positions along the axis its operator
# walks. Both rules below were documented as load errors and enforced
# nowhere (#58), so a fractional lag or a width that changed along the
# very axis it measured rendered as though it were neither.
pytest.param(
"shift(p, over=snapshot, offset=cost, edge='wrap')",
r'declared dtype: float',
id='a-named-offset-is-integral',
),
pytest.param(
"shift(p, over=snapshot, offset=horizon, edge='wrap')",
r'varies along the axis it walks is a permutation rather than a lag',
id='a-named-offset-does-not-span-the-axis-it-walks',
),
pytest.param(
'sum_back(p, over=snapshot, within=cost)',
r'declared dtype: float',
id='a-named-width-is-integral',
),
pytest.param(
'sum_back(p, over=snapshot, within=horizon)',
r'no longer "the last n"',
id='a-named-width-does-not-span-the-summed-axis',
),
],
)
def test_an_ill_dimensioned_expression_is_rejected(expr, match):
Expand Down
Loading