"""A shape- and meaning-checked pipeline of small linear programs.""" from __future__ import annotations from dataclasses import dataclass import math from typing import Iterable, Sequence Vector = tuple[float, ...] Matrix = tuple[Vector, ...] def _require_name(value: object, label: str) -> str: if not isinstance(value, str) or not value.strip(): raise ValueError(f"{label} must be a non-empty string") return value def _finite_number(value: object, label: str) -> float: if isinstance(value, bool) or not isinstance(value, (int, float)): raise ValueError(f"{label} must be a real number") converted = float(value) if not math.isfinite(converted): raise ValueError(f"{label} must be finite") return converted def _freeze_axes(values: Sequence[str], label: str) -> tuple[str, ...]: if isinstance(values, (str, bytes)) or not isinstance(values, Sequence): raise ValueError(f"{label} must be a sequence") axes = tuple(_require_name(value, f"{label}[{index}]") for index, value in enumerate(values)) if not axes: raise ValueError(f"{label} must not be empty") if len(axes) != len(set(axes)): raise ValueError(f"{label} must be unique") return axes def _freeze_vector(values: Sequence[float], label: str) -> Vector: if isinstance(values, (str, bytes)) or not isinstance(values, Sequence): raise ValueError(f"{label} must be a numeric sequence") vector = tuple( _finite_number(value, f"{label}[{index}]") for index, value in enumerate(values) ) if not vector: raise ValueError(f"{label} must not be empty") return vector def _freeze_matrix(values: Sequence[Sequence[float]], label: str) -> Matrix: if isinstance(values, (str, bytes)) or not isinstance(values, Sequence): raise ValueError(f"{label} must be a sequence of rows") rows = tuple( _freeze_vector(row, f"{label}[{index}]") for index, row in enumerate(values) ) if not rows: raise ValueError(f"{label} must not be empty") width = len(rows[0]) if any(len(row) != width for row in rows): raise ValueError(f"{label} must be rectangular") return rows def _dot(left: Vector, right: Vector) -> float: if len(left) != len(right): raise ValueError("dot product dimension mismatch") products = tuple(a * b for a, b in zip(left, right)) if any(not math.isfinite(product) for product in products): raise ValueError("dot product produced a non-finite intermediate") return _finite_number(math.fsum(products), "dot product result") def _matmul(left: Matrix, right: Matrix) -> Matrix: if len(left[0]) != len(right): raise ValueError("matrix multiplication dimension mismatch") right_columns = tuple(zip(*right)) return tuple( tuple(_dot(row, tuple(column)) for column in right_columns) for row in left ) @dataclass(frozen=True) class LinearMap: name: str input_axes: tuple[str, ...] output_axes: tuple[str, ...] weights: Matrix def __post_init__(self) -> None: _require_name(self.name, "map name") if not isinstance(self.input_axes, tuple) or not isinstance( self.output_axes, tuple ): raise ValueError("map axes must be immutable tuples") if not isinstance(self.weights, tuple) or any( not isinstance(row, tuple) for row in self.weights ): raise ValueError("weights must be immutable tuples") checked_inputs = _freeze_axes(self.input_axes, "input_axes") checked_outputs = _freeze_axes(self.output_axes, "output_axes") checked_weights = _freeze_matrix(self.weights, "weights") if checked_inputs != self.input_axes or checked_outputs != self.output_axes: raise ValueError("map axes failed normalization") if checked_weights != self.weights: raise ValueError("weights must contain normalized floats") if len(self.weights) != len(self.output_axes): raise ValueError("weight row count must match output axes") if len(self.weights[0]) != len(self.input_axes): raise ValueError("weight column count must match input axes") @classmethod def capture( cls, *, name: str, input_axes: Sequence[str], output_axes: Sequence[str], weights: Sequence[Sequence[float]], ) -> "LinearMap": frozen_inputs = _freeze_axes(input_axes, "input_axes") frozen_outputs = _freeze_axes(output_axes, "output_axes") frozen_weights = _freeze_matrix(weights, "weights") if len(frozen_weights) != len(frozen_outputs): raise ValueError("weight row count must match output axes") if len(frozen_weights[0]) != len(frozen_inputs): raise ValueError("weight column count must match input axes") return cls( name=_require_name(name, "map name"), input_axes=frozen_inputs, output_axes=frozen_outputs, weights=frozen_weights, ) @property def shape(self) -> tuple[int, int]: return len(self.output_axes), len(self.input_axes) def apply(self, values: Sequence[float]) -> Vector: vector = _freeze_vector(values, "input vector") if len(vector) != len(self.input_axes): raise ValueError( f"{self.name} expected {len(self.input_axes)} inputs, got {len(vector)}" ) return tuple(_dot(row, vector) for row in self.weights) @dataclass(frozen=True) class StageTrace: stage: str axes: tuple[str, ...] values: Vector @dataclass(frozen=True) class PipelineRun: pipeline: str input_values: Vector stages: tuple[StageTrace, ...] @property def output(self) -> Vector: return self.stages[-1].values @dataclass(frozen=True) class LinearMapPipeline: name: str stages: tuple[LinearMap, ...] def __post_init__(self) -> None: _require_name(self.name, "pipeline name") if not isinstance(self.stages, tuple) or not self.stages: raise ValueError("pipeline stages must be a non-empty immutable tuple") if not all(isinstance(stage, LinearMap) for stage in self.stages): raise ValueError("pipeline stages must be LinearMap values") names = tuple(stage.name for stage in self.stages) if len(names) != len(set(names)): raise ValueError("pipeline stage names must be unique") for previous, current in zip(self.stages, self.stages[1:]): if previous.output_axes != current.input_axes: raise ValueError( f"semantic axis mismatch: {previous.name} outputs " f"{previous.output_axes}, but {current.name} expects " f"{current.input_axes}" ) @classmethod def capture( cls, name: str, stages: Iterable[LinearMap] ) -> "LinearMapPipeline": frozen_stages = tuple(stages) if not frozen_stages: raise ValueError("pipeline requires at least one stage") if not all(isinstance(stage, LinearMap) for stage in frozen_stages): raise ValueError("pipeline stages must be LinearMap values") stage_names: set[str] = set() for stage in frozen_stages: if stage.name in stage_names: raise ValueError(f"duplicate stage name: {stage.name}") stage_names.add(stage.name) for previous, current in zip(frozen_stages, frozen_stages[1:]): if previous.output_axes != current.input_axes: raise ValueError( f"semantic axis mismatch: {previous.name} outputs " f"{previous.output_axes}, but {current.name} expects " f"{current.input_axes}" ) return cls(name=_require_name(name, "pipeline name"), stages=frozen_stages) def run(self, values: Sequence[float]) -> PipelineRun: frozen_input = _freeze_vector(values, "pipeline input") if len(frozen_input) != len(self.stages[0].input_axes): raise ValueError( f"pipeline expected {len(self.stages[0].input_axes)} inputs, " f"got {len(frozen_input)}" ) current = frozen_input traces: list[StageTrace] = [] for stage in self.stages: current = stage.apply(current) traces.append(StageTrace(stage.name, stage.output_axes, current)) return PipelineRun(self.name, frozen_input, tuple(traces)) def compile(self) -> LinearMap: """Compose stages without changing their left-to-right execution order.""" combined = self.stages[0].weights for stage in self.stages[1:]: combined = _matmul(stage.weights, combined) return LinearMap.capture( name=f"{self.name}:compiled", input_axes=self.stages[0].input_axes, output_axes=self.stages[-1].output_axes, weights=combined, ) ROTATE = LinearMap.capture( name="rotate", input_axes=("raw_x", "raw_y"), output_axes=("rotated_x", "rotated_y"), weights=((0.0, -1.0), (1.0, 0.0)), ) SCALE = LinearMap.capture( name="scale", input_axes=("rotated_x", "rotated_y"), output_axes=("scaled_x", "scaled_y"), weights=((2.0, 0.0), (0.0, 0.5)), ) PROJECT = LinearMap.capture( name="project", input_axes=("scaled_x", "scaled_y"), output_axes=("decision_signal",), weights=((1.0, 2.0),), ) EXAMPLE_PIPELINE = LinearMapPipeline.capture( "rotate-scale-project", (ROTATE, SCALE, PROJECT) ) if __name__ == "__main__": run = EXAMPLE_PIPELINE.run((3.0, 1.0)) compiled = EXAMPLE_PIPELINE.compile() print(f"pipeline={EXAMPLE_PIPELINE.name}") print(f"compiled_shape={compiled.shape[0]}x{compiled.shape[1]}") print("stage_outputs=" + ";".join( f"{trace.stage}:{trace.values}" for trace in run.stages )) print(f"compiled_matches={compiled.apply(run.input_values) == run.output}")