Clover-Image-Tiny / coreml-tools /test_multi_lora_math.py
neonforestmist's picture
Add object-aware inpainting v3 training and evaluation
e952093
Raw
History Blame Contribute Delete
1.51 kB
"""Numerical contract tests for block-concatenated Core ML LoRA slots."""
import torch
def test_block_concatenation_matches_weighted_adapter_sum() -> None:
generator = torch.Generator().manual_seed(20260813)
batch, pixels, input_width, output_width, rank = 2, 11, 7, 9, 4
hidden = torch.randn(batch, pixels, input_width, generator=generator)
downs = [
torch.randn(rank, input_width, generator=generator)
for _ in range(3)
]
ups = [
torch.randn(output_width, rank, generator=generator)
for _ in range(3)
]
strengths = [0.25, 1.0, 1.35]
expected = sum(
strength * ((hidden @ down.T) @ up.T)
for strength, down, up in zip(strengths, downs, ups)
)
state_down = torch.cat(downs, dim=0)
state_up = torch.cat(
[strength * up for strength, up in zip(strengths, ups)],
dim=1,
)
actual = (hidden @ state_down.T) @ state_up.T
torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-5)
def test_unused_slots_are_exactly_zero() -> None:
generator = torch.Generator().manual_seed(20260813)
hidden = torch.randn(1, 5, 6, generator=generator)
down = torch.randn(3, 6, generator=generator)
up = torch.randn(8, 3, generator=generator)
state_down = torch.zeros(9, 6)
state_up = torch.zeros(8, 9)
state_down[:3] = down
state_up[:, :3] = up
torch.testing.assert_close(
(hidden @ state_down.T) @ state_up.T,
(hidden @ down.T) @ up.T,
)