diff --git a/.github/workflows/ci-tests.yml b/.github/workflows/ci-tests.yml index 8dc68e6..f422298 100644 --- a/.github/workflows/ci-tests.yml +++ b/.github/workflows/ci-tests.yml @@ -2,7 +2,9 @@ name: CI testing on: push: + branches: [main] pull_request: + branches: [main] permissions: contents: read @@ -15,7 +17,7 @@ jobs: uses: ./.github/workflows/build-package.yml - test: + tests-cpu: name: Pytest runs-on: ${{ matrix.os }} timeout-minutes: 30 @@ -41,5 +43,32 @@ jobs: PYTHONHASHSEED: '0' PYTHONPATH: . run: | - uv run --no-sync python -m coverage run --branch --source=cutie -m pytest tests -q + uv run --no-sync python -m coverage run --branch --source=cutie -m pytest tests uv run --no-sync python -m coverage report --show-missing + + + tests-gpu: + name: Pytest (GPU) + runs-on: Roboflow-GPU-VM-Runner + timeout-minutes: 30 + env: + UV_TORCH_BACKEND: auto + steps: + - name: Print GPU information + run: nvidia-smi + - name: Check out source + uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + - name: Set up UV and Python + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 + with: + version: '0.9.26' + python-version: '3.12' + activate-environment: true + - name: Install all extras and test dependency group + run: uv pip install --group tests --strict '.[inference,evaluation,train,gui,video,data]' + - name: Run GPU-accelerated device-parity tests + env: + PYTEST_DISABLE_PLUGIN_AUTOLOAD: '1' + PYTHONHASHSEED: '0' + PYTHONPATH: . + run: uv run --no-sync pytest tests diff --git a/.gitignore b/.gitignore index c563304..4089bab 100644 --- a/.gitignore +++ b/.gitignore @@ -148,6 +148,7 @@ dmypy.json .idea/ # AI assets +.developments/ .reports/ # env files diff --git a/cutie/inference/memory_manager.py b/cutie/inference/memory_manager.py index 81de79a..126f386 100644 --- a/cutie/inference/memory_manager.py +++ b/cutie/inference/memory_manager.py @@ -244,11 +244,8 @@ def add_memory( for obj_id, obj in enumerate(objects): if obj in self.obj_v: # Keep embedding sums and counts for the object transformer's streaming average. - last_acc = self.obj_v[obj][:, :, -1] - new_acc = last_acc + obj_value[:, obj_id, :, -1] - - self.obj_v[obj][:, :, :-1] = self.obj_v[obj][:, :, :-1] + obj_value[:, obj_id, :, :-1] - self.obj_v[obj][:, :, -1] = new_acc + self.obj_v[obj][:, :, :-1].add_(obj_value[:, obj_id, :, :-1]) + self.obj_v[obj][:, :, -1].add_(obj_value[:, obj_id, :, -1]) else: self.obj_v[obj] = obj_value[:, obj_id] diff --git a/cutie/model/big_modules.py b/cutie/model/big_modules.py index 843143a..91730e5 100644 --- a/cutie/model/big_modules.py +++ b/cutie/model/big_modules.py @@ -80,7 +80,7 @@ def __init__(self, model_cfg: DictConfig): def forward(self, x: torch.Tensor, *, need_s: bool, need_e: bool) -> (torch.Tensor, torch.Tensor, torch.Tensor): x = self.pix_feat_proj(x) - shrinkage = self.d_proj(x) ** 2 + 1 if (need_s) else None + shrinkage = self.d_proj(x).pow(2).add_(1) if (need_s) else None selection = torch.sigmoid(self.e_proj(x)) if (need_e) else None return self.key_proj(x), shrinkage, selection diff --git a/cutie/model/channel_attn.py b/cutie/model/channel_attn.py index 7763bac..903db4e 100644 --- a/cutie/model/channel_attn.py +++ b/cutie/model/channel_attn.py @@ -31,4 +31,4 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: w = self.pool(x).view(b, 1, c) w = self.conv(w).transpose(-1, -2).unsqueeze(-1).sigmoid() # B*C*1*1 - return x * w + self.downsample(r) if self.residual else x * w + return torch.addcmul(self.downsample(r), x, w) if self.residual else x * w diff --git a/cutie/model/cutie.py b/cutie/model/cutie.py index 6916fd6..fc701f5 100644 --- a/cutie/model/cutie.py +++ b/cutie/model/cutie.py @@ -54,7 +54,8 @@ def _get_others(self, masks: torch.Tensor) -> torch.Tensor: return (masks.sum(dim=1, keepdim=True) - masks).clamp(0, 1) if num_objects >= 1 else torch.zeros_like(masks) def encode_image(self, image: torch.Tensor) -> (Iterable[torch.Tensor], torch.Tensor): - image = (image - self.pixel_mean) / self.pixel_std + # sub() copies (never mutates the caller's frame); only the fresh copy is div_'d in place + image = image.sub(self.pixel_mean).div_(self.pixel_std) ms_image_feat = self.pixel_encoder(image) return ms_image_feat, self.pix_feat_proj(ms_image_feat[0]) @@ -69,7 +70,7 @@ def encode_mask( chunk_size: int = -1, need_weights: bool = False, ) -> (torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor): - image = (image - self.pixel_mean) / self.pixel_std + image = image.sub(self.pixel_mean).div_(self.pixel_std) others = self._get_others(masks) mask_value, new_sensory = self.mask_encoder( image, diff --git a/cutie/model/group_modules.py b/cutie/model/group_modules.py index 6895da7..546f14f 100644 --- a/cutie/model/group_modules.py +++ b/cutie/model/group_modules.py @@ -89,7 +89,7 @@ def forward(self, x: torch.Tensor, g: torch.Tensor, skip_expand: bool = False) - elif self.method == 'mulcat': g = torch.cat([x * g, g], dim=2) elif self.method == 'muladd': - g = x * g + g + g = torch.addcmul(g, x, g) else: raise NotImplementedError diff --git a/cutie/model/utils/memory_utils.py b/cutie/model/utils/memory_utils.py index b1c1257..e94e60d 100644 --- a/cutie/model/utils/memory_utils.py +++ b/cutie/model/utils/memory_utils.py @@ -30,16 +30,18 @@ def get_similarity( # See XMem's appendix for derivation mk = mk.transpose(1, 2) a_sq = mk.pow(2) @ qe - two_ab = 2 * (mk @ (qk * qe)) + two_ab = (mk @ (qk * qe)).mul_(2) b_sq = (qe * qk.pow(2)).sum(1, keepdim=True) - similarity = -a_sq + two_ab - b_sq + similarity = two_ab.sub_(a_sq).sub_(b_sq) else: # similar to STCN if we don't have the selection term a_sq = mk.pow(2).sum(1).unsqueeze(2) - two_ab = 2 * (mk.transpose(1, 2) @ qk) - similarity = -a_sq + two_ab + two_ab = (mk.transpose(1, 2) @ qk).mul_(2) + similarity = two_ab.sub_(a_sq) - return similarity * ms / math.sqrt(CK) if ms is not None else similarity / math.sqrt(CK) # B*N*HW + if ms is not None: + return similarity.mul_(ms).div_(math.sqrt(CK)) + return similarity.div_(math.sqrt(CK)) # B*N*HW def do_softmax( @@ -63,7 +65,10 @@ def do_softmax( affinity = torch.zeros_like(similarity).scatter_(1, indices, x_exp) # B*N*HW else: maxes = torch.max(similarity, dim=1, keepdim=True)[0] - x_exp = torch.exp(similarity - maxes) + # exp_() is safe (sub() output is a fresh throwaway tensor); div_() is NOT — + # Exp's backward needs its own output value preserved, so the final + # normalization must stay out-of-place (confirmed by backward-parity test). + x_exp = similarity.sub(maxes).exp_() x_exp_sum = torch.sum(x_exp, dim=1, keepdim=True) affinity = x_exp / x_exp_sum indices = None diff --git a/pyproject.toml b/pyproject.toml index 1718235..c502d03 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,7 +14,11 @@ quote-style = "single" ignore-words-list = ["MOSE", "concurent", "indx"] [tool.pytest.ini_options] -testpaths = ["tests"] +testpaths = ["cutie", "tests"] +addopts = [ + "--color=yes", + "--doctest-modules", +] [tool.ruff.lint] # Ruff-only baseline. `PL` enables Ruff's Pylint-origin rules; explicit codes diff --git a/scripts/bench_vram_mps.py b/scripts/bench_vram_mps.py new file mode 100644 index 0000000..f7a2e10 --- /dev/null +++ b/scripts/bench_vram_mps.py @@ -0,0 +1,85 @@ +"""Standalone MPS VRAM bench for the memory-read hot path. + +Not a pytest assertion — MPS allocator memory numbers are noisy run-to-run, so +this is evidence to eyeball (before/after the in-place-op refactor), not a CI +gate. Correctness is covered by tests/test_mps_parity.py; this script only +measures peak Apple-GPU memory and wall-clock for cutie.model.utils.memory_utils +get_similarity + do_softmax at a realistic memory-bank size. + +Usage: + python scripts/bench_vram_mps.py [--frames N] [--hw N] [--ck N] [--iters N] +""" + +import argparse +import time + +import torch +from torch import mps + +from cutie.model.utils.memory_utils import do_softmax, get_similarity + + +def _build_inputs( + *, batch: int, ck: int, num_memory_frames: int, hw: int, device: str +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + n = num_memory_frames * hw + mk = torch.randn(batch, ck, n, device=device) + ms = torch.rand(batch, 1, n, device=device) + qk = torch.randn(batch, ck, hw, device=device) + qe = torch.rand(batch, ck, hw, device=device) + return mk, ms, qk, qe + + +def _run_once(mk: torch.Tensor, ms: torch.Tensor, qk: torch.Tensor, qe: torch.Tensor) -> torch.Tensor: + similarity = get_similarity(mk, ms, qk, qe) + return do_softmax(similarity) + + +def bench(*, batch: int, ck: int, num_memory_frames: int, hw: int, iters: int) -> None: + if iters < 1: + raise ValueError('iters must be >= 1') + if not torch.backends.mps.is_available(): + print('MPS not available on this machine — nothing to bench.') + return + + device = 'mps' + mk, ms, qk, qe = _build_inputs(batch=batch, ck=ck, num_memory_frames=num_memory_frames, hw=hw, device=device) + + # warm up (first MPS dispatch pays kernel-compile cost, not representative) + _run_once(mk, ms, qk, qe) + torch.mps.synchronize() + + mps.empty_cache() + baseline_allocated = mps.current_allocated_memory() + + start = time.perf_counter() + for _ in range(iters): + affinity = _run_once(mk, ms, qk, qe) + torch.mps.synchronize() + elapsed = time.perf_counter() - start + + peak_allocated = mps.driver_allocated_memory() + + print('MPS VRAM bench — cutie.model.utils.memory_utils (get_similarity + do_softmax)') + print(f' shapes: batch={batch} ck={ck} memory_frames={num_memory_frames} hw={hw} -> N={num_memory_frames * hw}') + print(f' iters: {iters}') + print(f' baseline allocated (post-warmup, pre-loop): {baseline_allocated / 2**20:.2f} MiB') + print(f' driver allocated (peak, post-loop): {peak_allocated / 2**20:.2f} MiB') + print(f' wall-clock: {elapsed:.4f}s total, {elapsed / iters * 1000:.3f}ms/iter') + print(f' output shape: {tuple(affinity.shape)}') + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--batch', type=int, default=1) + parser.add_argument('--ck', type=int, default=64, help='key channel dim') + parser.add_argument('--frames', type=int, default=20, dest='num_memory_frames', help='accumulated memory frames') + parser.add_argument('--hw', type=int, default=30 * 54, help='flattened spatial size (H*W/patch)') + parser.add_argument('--iters', type=int, default=50) + args = parser.parse_args() + + bench(batch=args.batch, ck=args.ck, num_memory_frames=args.num_memory_frames, hw=args.hw, iters=args.iters) + + +if __name__ == '__main__': + main() diff --git a/tests/test_device_parity.py b/tests/test_device_parity.py new file mode 100644 index 0000000..6b7a026 --- /dev/null +++ b/tests/test_device_parity.py @@ -0,0 +1,129 @@ +"""CPU-vs-accelerator parity guardrails for pure-math ops slated for VRAM optimization. + +These tests prove `cutie/model/utils/memory_utils.py` and +`cutie/model/channel_attn.py` produce numerically consistent results on every +GPU backend PyTorch supports on this codebase's target platforms - CUDA and +Apple's MPS - before any device-related refactor touches them. Every existing +test in this suite otherwise runs CPU-only; this file is the first to actually +execute tensor ops on an accelerator, so each backend's tests are guarded by a +real `is_available()` check rather than being unconditionally skipped. A +machine with only one backend (e.g. this MacBook has MPS, no CUDA) still runs +that backend's cases for real and simply skips the other - both are exercised +wherever hardware allows, neither is assumed absent. +""" + +import copy + +import pytest +import torch + +from cutie.model.channel_attn import CAResBlock +from cutie.model.utils.memory_utils import do_softmax, get_similarity + +# Backends to check parity against CPU, each independently skipped when its +# hardware isn't present - never assume one backend stands in for the other. +_ACCELERATOR_DEVICES = [ + pytest.param( + 'cuda', + marks=pytest.mark.skipif(not torch.cuda.is_available(), reason='CUDA not available on this machine'), + ), + pytest.param( + 'mps', + marks=pytest.mark.skipif(not torch.backends.mps.is_available(), reason='MPS not available on this machine'), + ), +] + +# Both CUDA and MPS accumulate in float32 and may reduce in a different order +# than CPU BLAS, so results are close but not bit-identical. These tolerances +# were determined empirically against this repo's ops (see module docstring); +# widen only with a concrete numerical justification. +_ATOL = 1e-5 +_RTOL = 1e-4 + + +def _build_similarity_inputs() -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Build small mk/ms/qk/qe tensors on CPU via the seeded RNG.""" + mk = torch.randn(2, 4, 3) + ms = torch.rand(2, 1, 3) + 0.5 + qk = torch.randn(2, 4, 5) + qe = torch.rand(2, 4, 5) + 0.5 + return mk, ms, qk, qe + + +class TestGetSimilarityDeviceParity: + """Guardrail: `get_similarity` must agree between CPU and each accelerator backend.""" + + @pytest.mark.parametrize('device', _ACCELERATOR_DEVICES) + @pytest.mark.parametrize( + 'use_qe', + [ + pytest.param(True, id='qe_present'), + pytest.param(False, id='qe_none'), + ], + ) + def test_get_similarity_matches_between_cpu_and_device(self, device: str, use_qe: bool) -> None: + """Same CPU-built tensors moved to the accelerator give results close to the CPU run.""" + mk_cpu, ms_cpu, qk_cpu, qe_cpu = _build_similarity_inputs() + qe_cpu = qe_cpu if use_qe else None + + mk_dev, ms_dev, qk_dev = mk_cpu.to(device), ms_cpu.to(device), qk_cpu.to(device) + qe_dev = qe_cpu.to(device) if qe_cpu is not None else None + + cpu_result = get_similarity(mk_cpu, ms_cpu, qk_cpu, qe_cpu) + device_result = get_similarity(mk_dev, ms_dev, qk_dev, qe_dev) + + torch.testing.assert_close(cpu_result, device_result.cpu(), atol=_ATOL, rtol=_RTOL) + + +class TestDoSoftmaxDeviceParity: + """Guardrail: `do_softmax` must agree between CPU and each accelerator backend.""" + + @pytest.mark.parametrize('device', _ACCELERATOR_DEVICES) + @pytest.mark.parametrize( + 'top_k', + [ + pytest.param(None, id='dense_softmax'), + pytest.param(2, id='top_k_subset'), + ], + ) + def test_do_softmax_matches_between_cpu_and_device(self, device: str, top_k: int | None) -> None: + """Same CPU-built similarity tensor moved to the accelerator gives a close affinity map.""" + similarity_cpu = torch.randn(2, 5, 3) + similarity_dev = similarity_cpu.to(device) + + cpu_result = do_softmax(similarity_cpu, top_k=top_k, inplace=False) + device_result = do_softmax(similarity_dev, top_k=top_k, inplace=False) + + torch.testing.assert_close(cpu_result, device_result.cpu(), atol=_ATOL, rtol=_RTOL) + + +class TestCAResBlockDeviceParity: + """Guardrail: `CAResBlock` forward pass must agree between CPU and each accelerator backend.""" + + @pytest.mark.parametrize('device', _ACCELERATOR_DEVICES) + def test_ca_res_block_forward_matches_between_cpu_and_device(self, device: str) -> None: + """Identical weights on CPU vs the accelerator produce a forward output within tolerance.""" + module_cpu = CAResBlock(4, 4).eval() + # deepcopy before moving preserves exact weights - two independently + # constructed modules would diverge even under the seed fixture, since + # Conv2d init happens at construction time. + module_dev = copy.deepcopy(module_cpu).to(device).eval() + x_cpu = torch.randn(1, 4, 8, 8) + x_dev = x_cpu.to(device) + + with torch.no_grad(): + cpu_result = module_cpu(x_cpu) + device_result = module_dev(x_dev) + + torch.testing.assert_close(cpu_result, device_result.cpu(), atol=_ATOL, rtol=_RTOL) + + +@pytest.mark.parametrize('device', _ACCELERATOR_DEVICES) +def test_get_similarity_actually_executes_on_device(device: str) -> None: + """Result tensor stays on the accelerator - proves no silent CPU fallback occurred.""" + mk_cpu, ms_cpu, qk_cpu, qe_cpu = _build_similarity_inputs() + mk, ms, qk, qe = (t.to(device) for t in (mk_cpu, ms_cpu, qk_cpu, qe_cpu)) + + result = get_similarity(mk, ms, qk, qe) + + assert result.device.type == device diff --git a/tests/test_memory_utils.py b/tests/test_memory_utils.py index 431d222..46dd54e 100644 --- a/tests/test_memory_utils.py +++ b/tests/test_memory_utils.py @@ -1,5 +1,14 @@ -"""Tests for deterministic memory similarity, affinity, and readout math.""" +"""Tests for deterministic memory similarity, affinity, and readout math. +The tests below characterize the CURRENT behavior of +``cutie/model/utils/memory_utils.py`` as a safety net ahead of an in-place-op +VRAM refactor. They pin exact numeric parity (not just finiteness) so any +future in-place rewrite can be checked against these baselines. +""" + +import math + +import pytest import torch from cutie.model.utils.memory_utils import do_softmax, get_affinity, get_similarity, readout @@ -41,3 +50,252 @@ def test_top_k_softmax_and_readout_keep_only_selected_memory_entries() -> None: assert torch.equal(usage, torch.tensor([[0.0, 1.0, 1.0]])) memory_values = torch.tensor([[[[[10.0, 20.0]]]]]) assert torch.equal(readout(affinity[:, :2], memory_values), torch.tensor([[[[20.0, 0.0]]]])) + + +def _reference_similarity( + mk: torch.Tensor, + ms: torch.Tensor | None, + qk: torch.Tensor, + qe: torch.Tensor | None, +) -> torch.Tensor: + """Elementwise reimplementation of get_similarity's formula, independent of its matmul path.""" + b_dim, ck_dim, n_dim = mk.shape + hw_dim = qk.shape[2] + out = torch.zeros(b_dim, n_dim, hw_dim, dtype=mk.dtype) + for b in range(b_dim): + for n in range(n_dim): + for hw in range(hw_dim): + a_sq = 0.0 + two_ab = 0.0 + b_sq = 0.0 + for ck in range(ck_dim): + mk_val = mk[b, ck, n].item() + qk_val = qk[b, ck, hw].item() + if qe is not None: + qe_val = qe[b, ck, hw].item() + a_sq += (mk_val**2) * qe_val + two_ab += 2 * mk_val * qk_val * qe_val + b_sq += qe_val * (qk_val**2) + else: + a_sq += mk_val**2 + two_ab += 2 * mk_val * qk_val + similarity = -a_sq + two_ab - b_sq + similarity /= math.sqrt(ck_dim) + if ms is not None: + similarity *= ms[b, 0, n].item() + out[b, n, hw] = similarity + return out + + +@pytest.mark.parametrize( + 'has_qe, has_ms, ck_dim, n_dim', + [ + pytest.param(True, True, 2, 2, id='qe-present_ms-present_ck2_n2'), + pytest.param(True, True, 1, 2, id='qe-present_ms-present_ck1-boundary'), + pytest.param(True, True, 2, 1, id='qe-present_ms-present_n1-single-key'), + pytest.param(False, True, 2, 2, id='qe-none_ms-present'), + pytest.param(True, False, 2, 2, id='qe-present_ms-none'), + pytest.param(False, False, 2, 2, id='qe-none_ms-none'), + ], +) +def test_get_similarity_matches_independent_reference_implementation( + has_qe: bool, + has_ms: bool, + ck_dim: int, + n_dim: int, +) -> None: + """get_similarity matches an elementwise reference across qe/ms presence and CK/N boundaries.""" + hw_dim = 2 + mk = torch.randn(1, ck_dim, n_dim) + ms = torch.rand(1, 1, n_dim) + 0.5 if has_ms else None + qk = torch.randn(1, ck_dim, hw_dim) + qe = torch.rand(1, ck_dim, hw_dim) + 0.1 if has_qe else None + + similarity = get_similarity(mk, ms, qk, qe) + + expected = _reference_similarity(mk, ms, qk, qe) + torch.testing.assert_close(similarity, expected) + + +def test_get_similarity_add_batch_dim_matches_manual_unsqueeze() -> None: + """add_batch_dim=True on un-batched tensors matches manually unsqueezing then calling normally.""" + mk = torch.randn(2, 2) + ms = torch.rand(1, 2) + 0.5 + qk = torch.randn(2, 2) + qe = torch.rand(2, 2) + 0.1 + + via_flag = get_similarity(mk, ms, qk, qe, add_batch_dim=True) + via_manual_unsqueeze = get_similarity(mk.unsqueeze(0), ms.unsqueeze(0), qk.unsqueeze(0), qe.unsqueeze(0)) + + torch.testing.assert_close(via_flag, via_manual_unsqueeze) + + +def test_do_softmax_dense_branch_matches_torch_softmax() -> None: + """do_softmax with top_k=None reduces to a plain softmax over the memory dimension.""" + similarity = torch.randn(2, 3, 4) + + affinity = do_softmax(similarity, top_k=None) + + torch.testing.assert_close(affinity, torch.softmax(similarity, dim=1)) + + +def test_do_softmax_top_k_equals_full_size_matches_dense_softmax() -> None: + """Requesting top_k equal to the memory size reduces to full dense softmax regardless of ordering.""" + similarity = torch.randn(1, 3, 2) + + affinity = do_softmax(similarity, top_k=3) + + torch.testing.assert_close(affinity, torch.softmax(similarity, dim=1)) + + +def test_do_softmax_inplace_flag_controls_tensor_identity() -> None: + """inplace=True mutates and returns the input tensor; inplace=False returns an equal-valued new tensor.""" + base_similarity = torch.tensor([[[0.0, 0.0], [2.0, 1.0], [1.0, 3.0]]]) + similarity_for_inplace = base_similarity.clone() + similarity_for_out_of_place = base_similarity.clone() + + affinity_inplace = do_softmax(similarity_for_inplace, top_k=1, inplace=True) + affinity_out_of_place = do_softmax(similarity_for_out_of_place, top_k=1, inplace=False) + + assert affinity_inplace is similarity_for_inplace + assert affinity_out_of_place is not similarity_for_out_of_place + torch.testing.assert_close(affinity_inplace, affinity_out_of_place) + + +def test_do_softmax_top_k_tie_break_selects_first_matching_index() -> None: + """Tied similarity values resolve via torch.topk's current lowest-index-first tie-break.""" + similarity = torch.tensor([[[1.0, 1.0], [1.0, 1.0], [1.0, 1.0]]]) + + affinity = do_softmax(similarity, top_k=1) + + assert torch.equal(affinity, torch.tensor([[[1.0, 1.0], [0.0, 0.0], [0.0, 0.0]]])) + + +def test_do_softmax_dense_branch_large_magnitude_input_is_finite_and_normalized() -> None: + """Large-magnitude similarity values stay finite and normalize to 1 via the max-subtraction trick.""" + similarity = torch.tensor([[[0.0, 1e4], [1e4, 0.0], [-1e4, -1e4]]]) + + affinity = do_softmax(similarity, top_k=None) + + assert torch.isfinite(affinity).all() + torch.testing.assert_close(affinity.sum(dim=1), torch.ones(1, 2)) + + +def _reference_readout(affinity: torch.Tensor, mv: torch.Tensor) -> torch.Tensor: + """Independent, elementwise reimplementation of readout's flatten+bmm+reshape composition.""" + b_dim, cv_dim, t_dim, h_dim, w_dim = mv.shape + hw_dim = affinity.shape[2] + out = torch.zeros(b_dim, cv_dim, hw_dim, dtype=mv.dtype) + for b in range(b_dim): + for cv in range(cv_dim): + for hw in range(hw_dim): + acc = 0.0 + n = 0 + for t in range(t_dim): + for h in range(h_dim): + for w in range(w_dim): + acc += mv[b, cv, t, h, w].item() * affinity[b, n, hw].item() + n += 1 + out[b, cv, hw] = acc + return out.view(b_dim, cv_dim, h_dim, w_dim) + + +@pytest.mark.parametrize( + 'num_frames', + [ + pytest.param(1, id='single-memory-frame'), + pytest.param(3, id='multiple-memory-frames'), + ], +) +def test_readout_matches_independent_reference_across_frame_counts(num_frames: int) -> None: + """readout's batched matmul matches an elementwise reference for T=1 and T>1 memory frames.""" + height, width, channels = 1, 2, 2 + memory_values = torch.randn(1, channels, num_frames, height, width) + affinity = torch.softmax(torch.randn(1, num_frames * height * width, height * width), dim=1) + + mem = readout(affinity, memory_values) + + expected = _reference_readout(affinity, memory_values) + torch.testing.assert_close(mem, expected) + + +def test_get_affinity_equals_do_softmax_of_get_similarity() -> None: + """get_affinity's composition matches calling do_softmax directly on get_similarity's output.""" + memory_key = torch.randn(1, 2, 2) + memory_shrinkage = torch.rand(1, 1, 2) + 0.5 + query_key = torch.randn(1, 2, 2) + query_selection = torch.rand(1, 2, 2) + 0.1 + + affinity_via_shorthand = get_affinity(memory_key, memory_shrinkage, query_key, query_selection) + affinity_via_composed_calls = do_softmax(get_similarity(memory_key, memory_shrinkage, query_key, query_selection)) + + torch.testing.assert_close(affinity_via_shorthand, affinity_via_composed_calls) + + +@pytest.mark.parametrize( + 'has_qe', + [ + pytest.param(True, id='qe-present'), + pytest.param(False, id='qe-none'), + ], +) +def test_get_similarity_backward_produces_finite_gradients(has_qe: bool) -> None: + """Backprop through get_similarity yields finite gradients for mk, ms, qk (and qe when present).""" + mk = torch.randn(1, 2, 2, requires_grad=True) + ms = torch.randn(1, 1, 2, requires_grad=True) + qk = torch.randn(1, 2, 2, requires_grad=True) + qe = torch.randn(1, 2, 2, requires_grad=True) if has_qe else None + + similarity = get_similarity(mk, ms, qk, qe) + similarity.sum().backward() + + assert mk.grad is not None + assert torch.isfinite(mk.grad).all() + assert ms.grad is not None + assert torch.isfinite(ms.grad).all() + assert qk.grad is not None + assert torch.isfinite(qk.grad).all() + if has_qe: + assert qe.grad is not None + assert torch.isfinite(qe.grad).all() + + +def test_get_affinity_backward_produces_finite_gradients() -> None: + """Backprop through get_affinity's dense-softmax path yields finite gradients for all inputs.""" + mk = torch.randn(1, 2, 2, requires_grad=True) + ms = torch.randn(1, 1, 2, requires_grad=True) + qk = torch.randn(1, 2, 2, requires_grad=True) + qe = torch.randn(1, 2, 2, requires_grad=True) + + affinity = get_affinity(mk, ms, qk, qe) + affinity.sum().backward() + + assert mk.grad is not None + assert torch.isfinite(mk.grad).all() + assert ms.grad is not None + assert torch.isfinite(ms.grad).all() + assert qk.grad is not None + assert torch.isfinite(qk.grad).all() + assert qe.grad is not None + assert torch.isfinite(qe.grad).all() + + +@pytest.mark.parametrize( + 'inplace', + [ + pytest.param(True, id='inplace-true'), + pytest.param(False, id='inplace-false'), + ], +) +def test_do_softmax_top_k_branch_backward_raises_inplace_version_error(inplace: bool) -> None: + """Pins current behavior: values.exp_() in the top-k path breaks backward regardless of inplace.""" + mk = torch.randn(1, 2, 2, requires_grad=True) + ms = torch.randn(1, 1, 2, requires_grad=True) + qk = torch.randn(1, 2, 2, requires_grad=True) + qe = torch.randn(1, 2, 2, requires_grad=True) + similarity = get_similarity(mk, ms, qk, qe) + + affinity = do_softmax(similarity, top_k=1, inplace=inplace) + + with pytest.raises(RuntimeError, match='modified by an inplace operation'): + affinity.sum().backward() diff --git a/tests/test_model_ops_vram.py b/tests/test_model_ops_vram.py new file mode 100644 index 0000000..71a2c33 --- /dev/null +++ b/tests/test_model_ops_vram.py @@ -0,0 +1,524 @@ +"""Characterization tests pinning today's (pre-refactor) numeric behavior of three +small model ops that are about to be rewritten into fused/in-place forms as part +of a VRAM-reduction pass: ``CAResBlock.forward``, ``MainToGroupDistributor.forward`` +(``method='muladd'`` branch), and ``KeyProjection.forward`` (``need_s=True`` branch). + +Expected tensors below were captured by running the current, unmodified +implementations with the repo's autouse ``torch.manual_seed(7)`` fixture in +effect, in the exact same arrange-order used by each test. Any future refactor +that changes these numeric outputs should fail these tests, by design. +""" + +import torch +from omegaconf import DictConfig, OmegaConf + +from cutie.model.big_modules import KeyProjection +from cutie.model.channel_attn import CAResBlock +from cutie.model.group_modules import MainToGroupDistributor + +import pytest + +# --------------------------------------------------------------------------- +# Golden snapshots (captured from current, unmodified source; see module docstring) +# --------------------------------------------------------------------------- + +_CARESBLOCK_IDENTITY_EXPECTED = [ + [ + [ + [-0.7206540703773499, -0.48447108268737793, 1.4976285696029663, -2.2937800884246826], # noqa: E501 + [1.0130069255828857, 0.3014580309391022, 2.0042061805725098, 0.5107771754264832], + [2.6103992462158203, 0.6673257946968079, 0.9283835291862488, 0.6571156978607178], + [0.6664162874221802, -0.8899217247962952, 0.16833892464637756, 0.5963094234466553], + ], + [ + [-0.8231000304222107, -1.1227807998657227, 2.6791787147521973, 0.14100490510463715], + [0.8509199023246765, -1.5944148302078247, -0.5495172739028931, -1.5795214176177979], + [1.0521636009216309, -1.4985541105270386, 0.9299943447113037, 0.42296963930130005], + [1.959437370300293, 1.0618263483047485, 0.0570513978600502, 0.10532163083553314], + ], + [ + [-0.6470261812210083, -1.645531177520752, 1.2063766717910767, 0.33335402607917786], + [0.13356362283229828, 0.7508739829063416, -0.36858347058296204, 0.31745144724845886], + [0.11065501719713211, -0.33934926986694336, 1.0882991552352905, -0.03698182478547096], + [-0.20804888010025024, -0.6829233169555664, 1.4042248725891113, 0.40697717666625977], + ], + [ + [0.02453736960887909, 1.7491830587387085, 0.7938931584358215, -0.5703772306442261], + [1.2537891864776611, -0.530758798122406, 0.4854661822319031, 0.9110860228538513], + [-1.1689016819000244, 0.43948251008987427, 1.4955304861068726, 0.2895548939704895], + [-1.7287875413894653, 0.5662024021148682, -0.9128029346466064, -0.1867765337228775], + ], + ] +] + +_CARESBLOCK_CONV_DOWNSAMPLE_EXPECTED = [ + [ + [ + [-0.7988987565040588, -0.8760892152786255, -0.16340935230255127, -0.37766095995903015], # noqa: E501 + [-0.3989090919494629, -0.7659040689468384, -0.7555376887321472, -0.20786331593990326], + [-1.006700873374939, 0.9164913892745972, -0.8373125791549683, -0.8962278962135315], + [0.011420249938964844, -0.4846363663673401, -0.8128257393836975, -0.695214569568634], + ], + [ + [-0.6941392421722412, -0.3061423897743225, -0.877033531665802, -1.126333475112915], + [-0.46585813164711, -0.7083226442337036, -0.38309425115585327, -0.9920669198036194], + [-1.0767713785171509, -0.6609583497047424, -0.49366840720176697, -0.07011915743350983], + [-1.1873834133148193, -0.31060564517974854, -1.3123273849487305, -0.879603385925293], + ], + [ + [1.0059151649475098, 0.3442264199256897, 0.9926256537437439, -0.2178320288658142], + [-0.29888027906417847, -0.1858750730752945, -0.08467313647270203, 0.0785488709807396], + [0.23613089323043823, -1.1202203035354614, 0.6520887613296509, 0.011555083096027374], + [0.24948400259017944, 0.6392955780029297, 1.3669610023498535, 0.6639381051063538], + ], + [ + [0.2068508267402649, -0.11143504083156586, -0.2776601016521454, -0.10822432488203049], + [0.16328930854797363, 0.7057285308837891, 0.6807976961135864, 0.058768875896930695], + [-1.3220728635787964, 0.2698323726654053, -0.04276278614997864, 0.7890651226043701], + [-0.700471043586731, 0.7764428853988647, -0.9466810822486877, -0.78248131275177], + ], + [ + [-0.46799108386039734, 0.262423574924469, -1.2755674123764038, -0.03028191439807415], + [0.2538577914237976, 0.9549093842506409, 0.7396650910377502, -0.3241986334323883], + [-0.4326525032520294, -0.45118892192840576, -0.2296057641506195, 0.7507773041725159], + [-1.371410608291626, -0.10546045750379562, -1.240556240081787, -0.882111668586731], + ], + [ + [-0.8674231767654419, -1.159828782081604, -0.312688410282135, -0.04612858593463898], + [-0.4642990827560425, -0.5516917705535889, -0.671933114528656, 0.08076949417591095], + [-1.5400004386901855, 1.4260551929473877, -1.3906091451644897, -1.0571593046188354], + [0.2224826216697693, -0.4511127173900604, -1.1908737421035767, -1.1668423414230347], + ], + [ + [-0.04571156948804855, 0.44293633103370667, -0.12959255278110504, 1.333909273147583], + [0.6900953054428101, 1.1973319053649902, 0.7281767129898071, 0.8715765476226807], + [0.8518197536468506, 0.9327398538589478, -0.18876105546951294, 0.2804362177848816], + [0.5176968574523926, -0.13629235327243805, 0.13538305461406708, 0.2745157480239868], + ], + [ + [0.26096439361572266, 0.40914738178253174, 0.8838558197021484, 0.5552221536636353], + [0.5665378570556641, -0.3622337877750397, -0.29275283217430115, 0.5856112837791443], + [1.501960277557373, 1.0697489976882935, 0.4497874081134796, -0.20012272894382477], + [1.4402728080749512, -0.11722021549940109, 1.0891470909118652, 1.2733904123306274], + ], + ] +] + +_CARESBLOCK_NO_RESIDUAL_EXPECTED = [ + [ + [ + [-0.10764598846435547, -0.1224059984087944, -0.05224497243762016, -0.05897151306271553], # noqa: E501 + [-0.042520247399806976, -0.11571522802114487, -0.0017320869956165552, -0.13556385040283203], + [-0.05730715021491051, -0.03567095845937729, -0.08553246408700943, -0.06843295693397522], + [-0.033901263028383255, -0.10138627141714096, -0.06513384729623795, -0.07943180203437805], + ], + [ + [0.03450392931699753, 0.0027291467413306236, 0.09042184799909592, 0.037953753024339676], + [-0.047239698469638824, 0.0001550402375869453, 0.06085469573736191, -0.013033552095293999], + [0.006835754960775375, -0.06763991713523865, -0.04564031958580017, -0.011223092675209045], + [-0.046970807015895844, -0.058978691697120667, -0.013011818751692772, -0.0007400549366138875], + ], + [ + [0.0015307003632187843, -0.013993693515658379, -0.029154637828469276, -0.015532853081822395], + [-0.04315483197569847, 0.014886012300848961, 0.039306093007326126, 0.002822999842464924], + [-0.0823882594704628, -0.045322950929403305, -0.04231299087405205, 0.033640120178461075], + [-0.00015583087224513292, -0.008149852976202965, 0.036061353981494904, 0.013531940057873726], + ], + [ + [0.12083804607391357, 0.014927023090422153, 0.1056428775191307, 0.006749439984560013], + [0.07568039745092392, 0.037382371723651886, 0.08347615599632263, -0.0004902100190520287], + [0.11134612560272217, 0.06914503127336502, -0.006513969972729683, 0.05894852802157402], + [0.024193676188588142, 0.03006836771965027, 0.03717657923698425, 0.06020074710249901], + ], + ] +] + +_DISTRIBUTOR_NUM_OBJECTS_1_EXPAND_EXPECTED = [ + [ + [ + [ + [-0.03205069899559021, -0.3621646463871002, -0.027511660009622574, 0.1491047739982605], # noqa: E501 + [-2.4708313941955566, -1.3629722595214844, -0.11090141534805298, 0.0993567705154419], + [3.1293156147003174, 0.0007280065910890698, 0.0698843002319336, -1.1382774114608765], + [1.9562900066375732, -0.7692430019378662, -0.43711814284324646, 1.5763871669769287], + ], + [ + [-0.07763387262821198, 0.12329724431037903, -1.5060032606124878, 0.536338210105896], + [-1.4898933172225952, -3.0611588954925537, -1.1875540018081665, -0.16952049732208252], + [-0.7003640532493591, 0.9392507076263428, -0.38132262229919434, -0.09424584358930588], + [0.8619248867034912, -0.6471694707870483, 0.335304856300354, -0.21299368143081665], + ], + [ + [-2.0691936016082764, 0.3948063254356384, 0.6719764471054077, 1.0152740478515625], + [-0.24639272689819336, -0.32385724782943726, 0.30052220821380615, 0.3862622380256653], + [0.15988200902938843, -0.39782416820526123, 1.998871922492981, -0.05590343475341797], + [-0.6094374656677246, 0.12244312465190887, 0.7524683475494385, 1.7899457216262817], + ], + ] + ] +] + +_DISTRIBUTOR_NUM_OBJECTS_2_EXPAND_EXPECTED = [ + [ + [ + [ + [-0.03205069899559021, -0.3621646463871002, -0.027511660009622574, 0.1491047739982605], # noqa: E501 + [-2.4708313941955566, -1.3629722595214844, -0.11090141534805298, 0.0993567705154419], + [3.1293156147003174, 0.0007280065910890698, 0.0698843002319336, -1.1382774114608765], + [1.9562900066375732, -0.7692430019378662, -0.43711814284324646, 1.5763871669769287], + ], + [ + [-0.07763387262821198, 0.12329724431037903, -1.5060032606124878, 0.536338210105896], + [-1.4898933172225952, -3.0611588954925537, -1.1875540018081665, -0.16952049732208252], + [-0.7003640532493591, 0.9392507076263428, -0.38132262229919434, -0.09424584358930588], + [0.8619248867034912, -0.6471694707870483, 0.335304856300354, -0.21299368143081665], + ], + [ + [-2.0691936016082764, 0.3948063254356384, 0.6719764471054077, 1.0152740478515625], + [-0.24639272689819336, -0.32385724782943726, 0.30052220821380615, 0.3862622380256653], + [0.15988200902938843, -0.39782416820526123, 1.998871922492981, -0.05590343475341797], + [-0.6094374656677246, 0.12244312465190887, 0.7524683475494385, 1.7899457216262817], + ], + ], + [ + [ + [-0.09655582904815674, 1.1568140983581543, 2.007277011871338, 0.4290674924850464], + [0.1061609759926796, -0.02434452436864376, -0.14617294073104858, 0.05654314160346985], + [0.739094078540802, 0.8883645534515381, 0.19398140907287598, 0.9561367630958557], + [0.8663599491119385, -0.31718286871910095, 0.5663720369338989, -2.410396099090576], + ], + [ + [0.6553231477737427, -0.21238300204277039, 1.8317945003509521, 0.06083924323320389], + [-0.7022362351417542, -0.8724789619445801, 1.3046162128448486, 0.16788628697395325], + [-0.6451551914215088, 0.9074205756187439, 0.9875812530517578, 0.332057923078537], + [0.4575645625591278, 1.026653528213501, -0.09151521325111389, -0.4264882802963257], + ], + [ + [0.521377444267273, 0.13017569482326508, 0.6576604843139648, -2.4247493743896484], + [-0.4383692741394043, -0.2689738869667053, -0.36607158184051514, -0.29545196890830994], + [-0.5832324028015137, -3.541163444519043, -0.8753631114959717, -0.03120255470275879], + [0.09360139071941376, -1.1362227201461792, -1.370066523551941, 1.6022694110870361], + ], + ], + ] +] + +_DISTRIBUTOR_NUM_OBJECTS_1_SKIP_EXPAND_EXPECTED = [ + [ + [ + [ + [-0.03205069899559021, -0.3621646463871002, -0.027511660009622574, 0.1491047739982605], # noqa: E501 + [-2.4708313941955566, -1.3629722595214844, -0.11090141534805298, 0.0993567705154419], + [3.1293156147003174, 0.0007280065910890698, 0.0698843002319336, -1.1382774114608765], + [1.9562900066375732, -0.7692430019378662, -0.43711814284324646, 1.5763871669769287], + ], + [ + [-0.07763387262821198, 0.12329724431037903, -1.5060032606124878, 0.536338210105896], + [-1.4898933172225952, -3.0611588954925537, -1.1875540018081665, -0.16952049732208252], + [-0.7003640532493591, 0.9392507076263428, -0.38132262229919434, -0.09424584358930588], + [0.8619248867034912, -0.6471694707870483, 0.335304856300354, -0.21299368143081665], + ], + [ + [-2.0691936016082764, 0.3948063254356384, 0.6719764471054077, 1.0152740478515625], + [-0.24639272689819336, -0.32385724782943726, 0.30052220821380615, 0.3862622380256653], + [0.15988200902938843, -0.39782416820526123, 1.998871922492981, -0.05590343475341797], + [-0.6094374656677246, 0.12244312465190887, 0.7524683475494385, 1.7899457216262817], + ], + ] + ] +] + +_DISTRIBUTOR_NUM_OBJECTS_2_SKIP_EXPAND_EXPECTED = [ + [ + [ + [ + [-0.09655582904815674, 1.1568140983581543, 2.007277011871338, 0.4290674924850464], # noqa: E501 + [0.1061609759926796, -0.02434452436864376, -0.14617294073104858, 0.05654314160346985], + [0.739094078540802, 0.8883645534515381, 0.19398140907287598, 0.9561367630958557], + [0.8663599491119385, -0.31718286871910095, 0.5663720369338989, -2.410396099090576], + ], + [ + [0.6553231477737427, -0.21238300204277039, 1.8317945003509521, 0.06083924323320389], + [-0.7022362351417542, -0.8724789619445801, 1.3046162128448486, 0.16788628697395325], + [-0.6451551914215088, 0.9074205756187439, 0.9875812530517578, 0.332057923078537], + [0.4575645625591278, 1.026653528213501, -0.09151521325111389, -0.4264882802963257], + ], + [ + [0.521377444267273, 0.13017569482326508, 0.6576604843139648, -2.4247493743896484], + [-0.4383692741394043, -0.2689738869667053, -0.36607158184051514, -0.29545196890830994], + [-0.5832324028015137, -3.541163444519043, -0.8753631114959717, -0.03120255470275879], + [0.09360139071941376, -1.1362227201461792, -1.370066523551941, 1.6022694110870361], + ], + ], + [ + [ + [0.2148890197277069, -0.11064700782299042, 1.1356507539749146, -0.6967417001724243], + [-4.104326248168945, 0.03213697671890259, 0.5557848811149597, 0.14293473958969116], + [0.17750436067581177, -0.5497691631317139, -0.2317337840795517, -0.0598011240363121], + [1.55190110206604, -0.66949462890625, -0.3418583869934082, 0.23257191479206085], + ], + [ + [0.7730319499969482, -1.2142809629440308, -0.038366325199604034, 0.2589553892612457], + [-0.34108293056488037, 0.2990848422050476, 0.06909536570310593, -0.2752949893474579], + [-0.07691073417663574, -1.5652966499328613, -1.014060378074646, 0.1978754848241806], + [-0.14204496145248413, -0.07347895205020905, -0.03530174493789673, 0.23751643300056458], + ], + [ + [0.5915085077285767, -1.191149353981018, 1.2868471145629883, 0.2695143222808838], + [0.6128371953964233, -0.029170632362365723, 0.8040453791618347, 2.892679214477539], + [0.18484124541282654, 0.6278330087661743, 2.6938729286193848, -1.5402262210845947], + [-0.37571096420288086, -1.4885963201522827, -1.484975814819336, -0.010609875433146954], + ], + ], + ] +] + +_KEYPROJECTION_SHRINKAGE_EXPECTED = [ + [ + [ + [ + 1.1253504753112793, + 1.0510914325714111, + 1.006724238395691, + 1.003408432006836, + 1.0146713256835938, # noqa: E501 + 1.0462661981582642, + ], + [ + 1.028773307800293, + 1.0125497579574585, + 1.0988126993179321, + 1.0008161067962646, + 1.002332091331482, + 1.0287643671035767, + ], + [ + 1.037503957748413, + 1.0087380409240723, + 1.2560498714447021, + 1.3191046714782715, + 1.0037364959716797, + 1.045958399772644, + ], + [ + 1.0111522674560547, + 1.0996208190917969, + 1.2209157943725586, + 1.1320141553878784, + 1.4356380701065063, + 1.0058521032333374, + ], + [ + 1.2221016883850098, + 1.0041530132293701, + 1.3392685651779175, + 1.4854915142059326, + 1.5084608793258667, + 1.0615326166152954, + ], + [ + 1.0018631219863892, + 1.0065165758132935, + 1.0162720680236816, + 1.0001444816589355, + 1.0031205415725708, + 1.0041496753692627, + ], + ] + ] +] + + +def _make_key_projection_cfg() -> DictConfig: + """Build the minimal OmegaConf config accepted by ``KeyProjection.__init__``.""" + return OmegaConf.create({'pixel_encoder': {'ms_dims': [8]}, 'pixel_dim': 4, 'key_dim': 4}) + + +# --------------------------------------------------------------------------- +# CAResBlock +# --------------------------------------------------------------------------- + +_CARESBLOCK_FORWARD_CASES = [ + pytest.param(4, 4, True, _CARESBLOCK_IDENTITY_EXPECTED, id='identity-downsample'), + pytest.param(4, 8, True, _CARESBLOCK_CONV_DOWNSAMPLE_EXPECTED, id='conv-downsample'), + pytest.param(4, 4, False, _CARESBLOCK_NO_RESIDUAL_EXPECTED, id='no-residual'), +] + +_CARESBLOCK_BACKWARD_CASES = [ + pytest.param(4, 4, True, 'identity', id='identity-downsample'), + pytest.param(4, 8, True, 'conv', id='conv-downsample'), + pytest.param(4, 4, False, 'none', id='no-residual'), +] + + +class TestCAResBlock: + """Pin CAResBlock.forward's current numeric output and gradient flow.""" + + @pytest.mark.parametrize('in_dim,out_dim,residual,expected', _CARESBLOCK_FORWARD_CASES) + def test_forward_matches_pinned_snapshot(self, in_dim: int, out_dim: int, residual: bool, expected: list) -> None: + """Forward output for each residual/downsample variant matches today's pinned snapshot.""" + model = CAResBlock(in_dim=in_dim, out_dim=out_dim, residual=residual) + x = torch.randn(1, in_dim, 4, 4) + + output = model(x) + + torch.testing.assert_close(output, torch.tensor(expected, dtype=torch.float32)) + + @pytest.mark.parametrize('in_dim,out_dim,residual,downsample_kind', _CARESBLOCK_BACKWARD_CASES) + def test_backward_populates_finite_gradients( + self, in_dim: int, out_dim: int, residual: bool, downsample_kind: str + ) -> None: + """Backward pass yields finite gradients on the input and on every conv weight in use.""" + model = CAResBlock(in_dim=in_dim, out_dim=out_dim, residual=residual) + x = torch.randn(1, in_dim, 4, 4, requires_grad=True) + + output = model(x) + output.sum().backward() + + assert x.grad is not None + assert torch.isfinite(x.grad).all() + assert model.conv1.weight.grad is not None + assert torch.isfinite(model.conv1.weight.grad).all() + assert model.conv2.weight.grad is not None + assert torch.isfinite(model.conv2.weight.grad).all() + if downsample_kind == 'conv': + assert model.downsample.weight.grad is not None + assert torch.isfinite(model.downsample.weight.grad).all() + + +# --------------------------------------------------------------------------- +# MainToGroupDistributor (method='muladd') +# --------------------------------------------------------------------------- + +_DISTRIBUTOR_FORWARD_CASES = [ + pytest.param( + False, + (1, 3, 4, 4), + (1, 1, 3, 4, 4), + _DISTRIBUTOR_NUM_OBJECTS_1_EXPAND_EXPECTED, + id='num-objects-1-expand', + ), + pytest.param( + False, + (1, 3, 4, 4), + (1, 2, 3, 4, 4), + _DISTRIBUTOR_NUM_OBJECTS_2_EXPAND_EXPECTED, + id='num-objects-2-expand', + ), + pytest.param( + True, + (1, 1, 3, 4, 4), + (1, 1, 3, 4, 4), + _DISTRIBUTOR_NUM_OBJECTS_1_SKIP_EXPAND_EXPECTED, + id='num-objects-1-skip-expand', + ), + pytest.param( + True, + (1, 2, 3, 4, 4), + (1, 2, 3, 4, 4), + _DISTRIBUTOR_NUM_OBJECTS_2_SKIP_EXPAND_EXPECTED, + id='num-objects-2-skip-expand', + ), +] + +_DISTRIBUTOR_BACKWARD_CASES = [ + pytest.param(False, (1, 3, 4, 4), (1, 1, 3, 4, 4), id='num-objects-1-expand'), + pytest.param(False, (1, 3, 4, 4), (1, 2, 3, 4, 4), id='num-objects-2-expand'), + pytest.param(True, (1, 1, 3, 4, 4), (1, 1, 3, 4, 4), id='num-objects-1-skip-expand'), + pytest.param(True, (1, 2, 3, 4, 4), (1, 2, 3, 4, 4), id='num-objects-2-skip-expand'), +] + + +class TestMainToGroupDistributorMulAdd: + """Pin MainToGroupDistributor.forward's 'muladd' branch (g = x * g + g).""" + + @pytest.mark.parametrize('skip_expand,x_shape,g_shape,expected', _DISTRIBUTOR_FORWARD_CASES) + def test_forward_matches_pinned_snapshot( + self, skip_expand: bool, x_shape: tuple, g_shape: tuple, expected: list + ) -> None: + """Forward output for each num_objects/skip_expand combination matches today's pinned snapshot.""" + model = MainToGroupDistributor(method='muladd') + x = torch.randn(*x_shape) + g = torch.randn(*g_shape) + + output = model(x, g, skip_expand=skip_expand) + + torch.testing.assert_close(output, torch.tensor(expected, dtype=torch.float32)) + + @pytest.mark.parametrize('skip_expand,x_shape,g_shape', _DISTRIBUTOR_BACKWARD_CASES) + def test_backward_populates_finite_gradients(self, skip_expand: bool, x_shape: tuple, g_shape: tuple) -> None: + """Backward pass yields finite gradients on both the main feature and the group feature inputs.""" + model = MainToGroupDistributor(method='muladd') + x = torch.randn(*x_shape, requires_grad=True) + g = torch.randn(*g_shape, requires_grad=True) + + output = model(x, g, skip_expand=skip_expand) + output.sum().backward() + + assert x.grad is not None + assert torch.isfinite(x.grad).all() + assert g.grad is not None + assert torch.isfinite(g.grad).all() + + +# --------------------------------------------------------------------------- +# KeyProjection (need_s=True shrinkage branch) +# --------------------------------------------------------------------------- + + +class TestKeyProjectionShrinkage: + """Pin KeyProjection.forward's shrinkage branch (shrinkage = d_proj(x) ** 2 + 1).""" + + def test_forward_need_s_matches_pinned_snapshot(self) -> None: + """need_s=True returns a shrinkage map matching today's pinned snapshot; selection stays None.""" + cfg = _make_key_projection_cfg() + model = KeyProjection(cfg) + x = torch.randn(1, 8, 6, 6) + + _, shrinkage, selection = model(x, need_s=True, need_e=False) + + assert selection is None + assert shrinkage.shape == (1, 1, 6, 6) + torch.testing.assert_close(shrinkage, torch.tensor(_KEYPROJECTION_SHRINKAGE_EXPECTED, dtype=torch.float32)) + + def test_forward_need_s_shrinkage_is_bounded_below_by_one(self) -> None: + """shrinkage = d_proj(x) ** 2 + 1 is mathematically >= 1 for any input, by construction.""" + cfg = _make_key_projection_cfg() + model = KeyProjection(cfg) + x = torch.randn(1, 8, 6, 6) + + _, shrinkage, _ = model(x, need_s=True, need_e=False) + + assert torch.all(shrinkage >= 1.0) + + def test_backward_populates_finite_gradient_on_d_proj_weight(self) -> None: + """Backpropagating through the shrinkage branch yields a finite gradient on d_proj.weight.""" + cfg = _make_key_projection_cfg() + model = KeyProjection(cfg) + x = torch.randn(1, 8, 6, 6) + + _, shrinkage, _ = model(x, need_s=True, need_e=False) + shrinkage.sum().backward() + + assert model.d_proj.weight.grad is not None + assert torch.isfinite(model.d_proj.weight.grad).all() + + +class TestImageNormalizeCallerSafety: + """Pins the ``image.sub(mean).div_(std)`` idiom used by ``CUTIE.encode_image``/ + ``encode_mask`` (cutie/model/cutie.py): the caller's frame tensor must never be + mutated, even though the ``std`` division happens in place on a fresh copy. + """ + + def test_normalize_matches_out_of_place_reference_and_leaves_caller_tensor_untouched(self) -> None: + """sub().div_() matches (image - mean) / std and does not mutate the input tensor.""" + image = torch.randn(1, 3, 8, 8) + image_snapshot = image.clone() + mean = torch.tensor([0.485, 0.456, 0.406]).view(-1, 1, 1) + std = torch.tensor([0.229, 0.224, 0.225]).view(-1, 1, 1) + + expected = (image - mean) / std + normalized = image.sub(mean).div_(std) + + torch.testing.assert_close(normalized, expected) + torch.testing.assert_close(image, image_snapshot)