"""Check exact logical ownership, independent of local SRAM packing."""
import collections
import csv
import json
from pathlib import Path

ROOT = Path(__file__).parent
CASES = {
    "up": (1, 729, 1152, 4304),
    "qkv": (1, 729, 1152, 3456),
    "down": (1, 729, 4304, 1152),
    "qk": (16, 729, 72, 729),
    "pv": (16, 729, 729, 72),
}
results = {}
for name, (groups, m, k, n) in CASES.items():
    mapping = collections.defaultdict(lambda: collections.defaultdict(list))
    for row in csv.DictReader((ROOT / f"{name}-mapping.csv").open()):
        mapping[row["operand"]][int(row["tile"])].append(
            (int(row["begin"]), int(row["end"])))
    results[name] = {}
    for operand, tiles in mapping.items():
        rows, columns = (m, k) if operand == "lhs" else (k, n)
        rect = linear = 0
        examples = []
        all_intervals = []
        for tile, intervals in sorted(tiles.items()):
            intervals.sort()
            all_intervals.extend(intervals)
            count = sum(b - a for a, b in intervals)
            lo = [groups, rows, columns]
            hi = [-1, -1, -1]
            row_spans = []
            for begin, end in intervals:
                while begin < end:
                    g, r = divmod(begin // columns, rows)
                    c = begin % columns
                    stop = min(end, (begin // columns + 1) * columns)
                    last_c = c + stop - begin - 1
                    for axis, low, high in [(0, g, g), (1, r, r), (2, c, last_c)]:
                        lo[axis] = min(lo[axis], low)
                        hi[axis] = max(hi[axis], high)
                    if len(row_spans) < 12:
                        row_spans.append([g, r, c, last_c + 1])
                    begin = stop
            volume = 1
            for a, b in zip(lo, hi):
                volume *= b - a + 1
            is_rect = volume == count
            is_linear = intervals[-1][1] - intervals[0][0] == count
            rect += is_rect
            linear += is_linear
            if not is_rect and len(examples) < 3:
                examples.append(dict(tile=tile, elements=count, bounds=[lo, hi],
                                     first_row_spans=row_spans))
        # Ensure one complete, unreplicated assignment; no holes or overlaps.
        cursor = 0
        for a, b in sorted(all_intervals):
            assert a == cursor, (name, operand, a, cursor)
            cursor = b
        assert cursor == groups * rows * columns
        results[name][operand] = dict(tiles=len(tiles), rectangular=rect,
                                     contiguous_logical_linear=linear,
                                     examples=examples)
print(json.dumps(results, indent=2))
