File size: 8,645 Bytes
ecc81b3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
"""How full is the lattice, really — answered before a model is built.

A lattice is dense when every combination of its axes exists and sparse when
most of them do not, and which one you have decides whether the sparse
machinery is doing anything for you. That is a property of the *data*, not a
setting, so it is measured rather than declared:

    report = td.data.sparsity(table)
    print(report)          # a summary, including the percentage
    report.percent_sparse  # 32.1

The per-axis breakdown is the part worth reading twice. A lattice can be 30%
sparse because observations are scattered, or because one station in twelve
reports nothing at all — the first is ordinary, the second is usually a join
that went wrong upstream, and `empty_slices` tells them apart.
"""

from __future__ import annotations

from collections.abc import Sequence
from dataclasses import dataclass, field

import torch

from torch_dimensions.lattice import Lattice

__all__ = ["SparsityReport", "sparsity"]


@dataclass
class SparsityReport:
    """What a pre-run over the data found about lattice occupancy."""

    shape: tuple[int, ...]
    names: tuple[str, ...]
    present: int
    total: int
    per_axis: dict[str, list[int]] = field(default_factory=dict)
    """Per axis, how many cells are present at each index along it."""
    empty_slices: dict[str, list[int]] = field(default_factory=dict)
    """Per axis, the indices with no present cell at all."""
    observed: float | None = None
    """Fraction of (time, cell, feature) entries actually observed, when the
    data carried a time axis. ``None`` when only structure was inspected."""

    @property
    def absent(self) -> int:
        return self.total - self.present

    @property
    def fraction_present(self) -> float:
        return self.present / self.total if self.total else 1.0

    @property
    def percent_sparse(self) -> float:
        """The headline number: what percentage of the lattice is absent."""
        return 100.0 * (1.0 - self.fraction_present)

    @property
    def dense(self) -> bool:
        return self.present == self.total

    def summary(self) -> str:
        head = (
            f"lattice {' × '.join(str(s) for s in self.shape) or '—'}  "
            f"{self.present}/{self.total} cells present  "
            f"({self.percent_sparse:.1f}% sparse)"
        )
        if self.dense:
            head += "  — dense; a validity mask would do nothing"
        lines = [head]
        if self.observed is not None:
            lines.append(f"  observed entries: {100.0 * self.observed:.1f}% of the series")
        for name in self.names:
            counts = self.per_axis.get(name, [])
            if not counts:
                continue
            per = self.total // len(counts) if counts else 0
            worst = min(counts) if counts else 0
            empty = self.empty_slices.get(name, [])
            note = f"  {name:<12} {min(counts)}–{max(counts)} of {per} per index"
            if empty:
                note += f"   ⚠ {len(empty)} empty: {empty[:6]}{'…' if len(empty) > 6 else ''}"
            elif worst == per:
                note += "   (full)"
            lines.append(note)
        return "\n".join(lines)

    def __repr__(self) -> str:
        return (
            f"SparsityReport(shape={self.shape}, present={self.present}/{self.total}, "
            f"percent_sparse={self.percent_sparse:.1f})"
        )


def _lattice_mask(lattice: Lattice) -> torch.Tensor:
    """The presence mask in the lattice's own shape.

    Not ``lattice.mask()``, which is broadcast-shaped ``(1, 1, *shape, 1)`` for
    multiplying against data — right for arithmetic, wrong here, where the
    singleton batch/time/feature axes would be reported as lattice axes and
    push the real names off the end.
    """
    if lattice.valid is None:
        return torch.ones(tuple(lattice.shape), dtype=torch.bool)
    return lattice.valid.bool()


def _mask_from_values(
    values: torch.Tensor, shape: tuple[int, ...], missing: float | None
) -> tuple[torch.Tensor, float]:
    """Reduce a data tensor to a per-cell presence mask.

    The lattice axes are located as a contiguous run inside ``values.shape``;
    everything before them is time and everything after is features, both of
    which are reduced away — a cell counts as present when *any* observation of
    it exists. Requiring the caller to state the shape rather than guessing it
    is deliberate: a (6, 8) lattice inside a (10, 6, 8, 1) tensor has exactly
    one sensible reading, but a (6, 6) one does not, and silently picking is
    how a transposed axis survives to training.
    """
    rank = len(shape)
    dims = values.shape
    start = None
    for s in range(len(dims) - rank + 1):
        if tuple(dims[s : s + rank]) == shape:
            if start is not None:
                raise ValueError(
                    f"lattice shape {shape} appears more than once in tensor shape "
                    f"{tuple(dims)}; slice the tensor so the placement is unambiguous"
                )
            start = s
    if start is None:
        raise ValueError(f"lattice shape {shape} does not appear in tensor shape {tuple(dims)}")

    observed = (
        torch.isfinite(values)
        if values.is_floating_point()
        else torch.ones_like(values, dtype=torch.bool)
    )
    if missing is not None:
        observed = observed & (values != missing)

    reduce_dims = [d for d in range(values.ndim) if not (start <= d < start + rank)]
    fraction = float(observed.float().mean()) if observed.numel() else 1.0
    mask = observed.any(dim=reduce_dims) if reduce_dims else observed
    return mask.bool(), fraction


def sparsity(
    data,
    *,
    shape: Sequence[int] | None = None,
    names: Sequence[str] | None = None,
    missing: float | None = None,
) -> SparsityReport:
    """Measure how much of a lattice its data actually occupies.

    Args:
        data: a :class:`~torch_dimensions.Lattice`, a
            :class:`~torch_dimensions.data.LatticeTable`, a boolean presence
            mask shaped like the lattice, or a data tensor (with ``shape``
            given) whose non-finite entries mark absence.
        shape: the lattice shape, required only for the data-tensor form.
        names: axis names, when the source does not carry them.
        missing: an additional sentinel counted as absent (e.g. ``0.0`` for a
            table that filled gaps with zeros rather than NaNs).

    Returns:
        A :class:`SparsityReport`; ``report.percent_sparse`` is the headline.
    """
    from torch_dimensions.data.table import LatticeTable

    observed: float | None = None

    if isinstance(data, LatticeTable):
        lattice = data.lattice
        mask = _lattice_mask(lattice)
        _, observed = _mask_from_values(data.series, tuple(lattice.shape), missing)
        axis_names = tuple(lattice.names or ())
    elif isinstance(data, Lattice):
        lattice = data
        mask = _lattice_mask(lattice)
        axis_names = tuple(lattice.names or ())
    else:
        values = torch.as_tensor(data)
        if values.dtype == torch.bool and shape is None:
            mask = values
        else:
            if shape is None:
                raise ValueError(
                    "measuring a data tensor needs the lattice shape: "
                    "sparsity(values, shape=(6, 8), names=('h', 'w'))"
                )
            mask, observed = _mask_from_values(values, tuple(shape), missing)
        axis_names = tuple(names or ())

    mask = mask.bool()
    shape_t = tuple(mask.shape)
    if names is not None:
        axis_names = tuple(names)
    # Lattice names include the time axis in front of the spatial ones; only
    # the spatial names describe the mask's dimensions.
    if len(axis_names) > len(shape_t):
        axis_names = axis_names[len(axis_names) - len(shape_t) :]
    if len(axis_names) != len(shape_t):
        axis_names = tuple(f"dim{i}" for i in range(len(shape_t)))

    per_axis: dict[str, list[int]] = {}
    empty: dict[str, list[int]] = {}
    for i, name in enumerate(axis_names):
        others = [d for d in range(mask.ndim) if d != i]
        counts = mask.sum(dim=others).tolist() if others else mask.long().tolist()
        per_axis[name] = [int(c) for c in counts]
        empty[name] = [j for j, c in enumerate(counts) if int(c) == 0]

    return SparsityReport(
        shape=shape_t,
        names=axis_names,
        present=int(mask.sum()),
        total=int(mask.numel()),
        per_axis=per_axis,
        empty_slices=empty,
        observed=observed,
    )