From 1376a23bda5e6fe94148cb57ca1447907c6bb21e Mon Sep 17 00:00:00 2001 From: tslil Date: Tue, 14 Oct 2025 20:48:03 +0100 Subject: Sketching the implementation of the monad and selections --- python/src/types/selection.py | 161 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 161 insertions(+) create mode 100644 python/src/types/selection.py (limited to 'python/src/types/selection.py') diff --git a/python/src/types/selection.py b/python/src/types/selection.py new file mode 100644 index 0000000..478b5b9 --- /dev/null +++ b/python/src/types/selection.py @@ -0,0 +1,161 @@ +from dataclasses import dataclass, field +from typing import Callable + +from .buffer import Buffer + + +@dataclass +class Position: + pos: int + + +@dataclass +class Range: + start: Position + end: Position + + +@dataclass +class Ranges: + ranges: list[Range] + + +@dataclass +class MultiRanges: + multi_ranges: dict[Buffer, Ranges] + + +@dataclass +class Selection: + selection: Position | Range | Ranges | MultiRanges + buffers: set[Buffer] = field(default_factory=set) + + def __post_init__(self): + if type(self.selection) is MultiRanges: + self.buffers = set(self.selection.multi_ranges) + else: + assert len(self.buffers) == 1 + + @classmethod + def empty(cls) -> "Selection": + return cls(selection=Position(0), buffers=set()) + + def promote(self, to_rank: int) -> "Selection": + if to_rank < 0 or to_rank > 3: + raise ValueError(f"Invalid rank: {to_rank}") + if self.rank > to_rank: + raise ValueError( + f"Cannot coerce selection of rank {self.rank} to rank {to_rank}" + ) + if self.rank == to_rank: + return self + match self.selection: + case Position(_): + as_range = Range( + start=Position(self.selection.pos), + end=Position(self.selection.pos + 1), + ) + as_ranges = Ranges(ranges=[as_range]) + if to_rank == 1: + return Selection(selection=as_range, buffers=self.buffers) + if to_rank == 2: + return Selection(selection=as_ranges, buffers=self.buffers) + if to_rank == 3: + return Selection( + MultiRanges({buffer: as_ranges for buffer in self.buffers}) + ) + case Range(_, _): + as_ranges = Ranges(ranges=[self.selection]) + if to_rank == 2: + return Selection(selection=as_ranges, buffers=self.buffers) + if to_rank == 3: + return Selection( + MultiRanges({buffer: as_ranges for buffer in self.buffers}) + ) + case Ranges(_): + return Selection( + MultiRanges({buffer: self.selection for buffer in self.buffers}) + ) + + raise RuntimeError("Unexpected promotion issue") + + @classmethod + def union(cls, *selections: "Selection") -> "Selection": + if not selections: + return cls.empty() + + all_buffers = set.union(*[s.buffers for s in selections]) + + if len(all_buffers) == 0: + return cls.empty() + + if len(all_buffers) > 1: + max_rank = 3 + else: + max_rank = max(s.rank for s in selections) + if max_rank < 2: + max_rank = 2 + + promoted = [s.promote(max_rank) for s in selections] + + if max_rank == 2: + all_ranges = [] + for sel in promoted: + if not type(sel.selection) is Ranges: + raise ValueError( + f"Selection of rank {sel.rank} is not a range. Should have been promoted to rank 2." + ) + all_ranges.extend(sel.selection.ranges) + return cls(selection=Ranges(all_ranges), buffers=all_buffers) + + if max_rank == 3: + buffer_ranges: dict[Buffer, list[Range]] = {buf: [] for buf in all_buffers} + for sel in promoted: + if not type(sel.selection) is MultiRanges: + raise ValueError( + f"Selection of rank {sel.rank} is not a multi-range. Should have been promoted to rank 3." + ) + for buffer, ranges in sel.selection.multi_ranges.items(): + buffer_ranges[buffer].extend(ranges.ranges) + + multi = MultiRanges( + {buffer: Ranges(ranges) for buffer, ranges in buffer_ranges.items()} + ) + return cls(selection=multi) + + raise RuntimeError("Unexpected union issue") + + @property + def rank(self) -> int: + match self.selection: + case Position(_): + return 0 + case Range(_, _): + return 1 + case Ranges(_): + return 2 + case MultiRanges(_): + return 3 + + def broadcast( + self, + rank_zero: Callable[[Position, Buffer], "Selection"], + rank_one: Callable[[Range, Buffer], "Selection"], + ) -> "Selection": + match self.selection: + case Position(_): + return rank_zero(self.selection, list(self.buffers)[0]) + + case Range(_, _) as r: + return rank_one(r, list(self.buffers)[0]) + + case Ranges(ranges): + buffer = list(self.buffers)[0] + return Selection.union(*[rank_one(r, buffer) for r in ranges]) + + case MultiRanges(multi): + all_results = [] + for buffer, ranges in multi.items(): + for r in ranges.ranges: + all_results.append(rank_one(r, buffer)) + return Selection.union(*all_results) -- cgit v1.2.3