| """Focused production seams for learned local continuation.""" |
|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import torch |
| from torch import nn |
|
|
| from music3lab.editing.learned_audio_continuation import ( |
| Rank4QKVContinuationAdapter, |
| apply_generated_residual, |
| continuation_residual_target, |
| flow_interpolate, |
| analyze_tail, |
| composed_seam_metrics, |
| deterministic_pair_offset, |
| shared_condition_variant_inputs, |
| ) |
| from music3lab.editing.learned_audio_continuation_data import ( |
| load_learned_continuation_config, |
| load_pair_records, |
| ) |
| from music3lab.editing.learned_audio_continuation_runner import ( |
| cosine_learning_rate, |
| ) |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| CORPUS = Path( |
| "/home/ubuntu/minimax-laion-corpus/versions/" |
| "interim_tranche_678_due_systemic_bot_auth" |
| ) |
|
|
|
|
| class _Attention(nn.Module): |
| def __init__(self, hidden: int) -> None: |
| super().__init__() |
| self.to_q = nn.Linear(hidden, hidden, bias=False) |
| self.to_k = nn.Linear(hidden, hidden, bias=False) |
| self.to_v = nn.Linear(hidden, hidden, bias=False) |
|
|
|
|
| class _Block(nn.Module): |
| def __init__(self, hidden: int) -> None: |
| super().__init__() |
| self.attn = _Attention(hidden) |
|
|
|
|
| class _Flow(nn.Module): |
| def __init__(self, layers: int, hidden: int) -> None: |
| super().__init__() |
| self.transformer_blocks = nn.ModuleList( |
| [_Block(hidden) for _ in range(layers)] |
| ) |
|
|
| def forward( |
| self, |
| *, |
| hidden_states: torch.Tensor, |
| timestep: torch.Tensor, |
| encoder_hidden_states: torch.Tensor, |
| return_dict: bool, |
| ) -> tuple[torch.Tensor]: |
| del timestep, return_dict |
| probe = encoder_hidden_states[:, :, :8] |
| for block in self.transformer_blocks: |
| probe = ( |
| block.attn.to_q(probe) |
| + block.attn.to_k(probe) |
| + block.attn.to_v(probe) |
| ) / 3 |
| update = probe.mean(dim=(1, 2)).view(-1, 1, 1) |
| return (hidden_states + update,) |
|
|
|
|
| def test_exact_config_and_frozen_542_68_68_source_inventory() -> None: |
| loaded = load_learned_continuation_config( |
| ROOT / "configs" / "learned-audio-continuation-v1.yaml" |
| ) |
| records = load_pair_records(CORPUS, loaded) |
| assert {key: len(value) for key, value in records.items()} == { |
| "train": 542, |
| "validation": 68, |
| "heldout": 68, |
| } |
| assert loaded.config.projector.target_visible_to_conditioner is False |
| assert loaded.config.projector.champion_eligible is False |
| assert loaded.config.claims.handcrafted_baseline_reclassified is False |
| assert loaded.config.training.batch_size == 16 |
| assert ( |
| loaded.config.target_parameterization |
| == "target_minus_repeat_tail" |
| ) |
| assert loaded.config.generated_latent_parameterization == ( |
| "repeat_tail_plus_generated_residual" |
| ) |
|
|
|
|
| def test_residual_coordinate_reconstructs_next_and_shares_anchor_noise() -> None: |
| anchor = torch.tensor( |
| [[[1.0, 2.0], [3.0, 4.0]]], |
| dtype=torch.float32, |
| ) |
| target = torch.tensor( |
| [[[5.0, 7.0], [11.0, 13.0]]], |
| dtype=torch.float32, |
| ) |
| noise = torch.tensor( |
| [[[0.5, -0.5], [1.5, -1.5]]], |
| dtype=torch.float32, |
| ) |
| residual = continuation_residual_target(target, anchor) |
| terminal = flow_interpolate(noise, residual, torch.ones(1)) |
| assert torch.equal(terminal, residual) |
| assert torch.equal(apply_generated_residual(anchor, terminal), target) |
| assert torch.equal( |
| apply_generated_residual(anchor, torch.zeros_like(anchor)), |
| anchor, |
| ) |
| conditions = { |
| "conditional": torch.ones(1, 2, 3), |
| "zero_context": torch.zeros(1, 2, 3), |
| "unrelated_context": -torch.ones(1, 2, 3), |
| } |
| variants = shared_condition_variant_inputs( |
| context_anchor=anchor, |
| noise_latent=noise, |
| conditions=conditions, |
| ) |
| assert set(variants) == set(conditions) |
| assert all( |
| shared_anchor is anchor and shared_noise is noise |
| for shared_anchor, shared_noise, _ in variants.values() |
| ) |
| assert all( |
| torch.equal( |
| apply_generated_residual(shared_anchor, torch.zeros_like(shared_anchor)), |
| anchor, |
| ) |
| for shared_anchor, _, _ in variants.values() |
| ) |
|
|
|
|
| def test_rank_qkv_hooks_are_only_trainables_and_context_changes_same_noise() -> None: |
| flow = _Flow(2, 8) |
| adapter = Rank4QKVContinuationAdapter( |
| flow, layers=2, hidden_size=8, rank=2 |
| ) |
| try: |
| assert adapter.trainable_parameter_count() == 2 * 3 * 2 * 8 * 2 |
| assert all(not value.requires_grad for value in flow.parameters()) |
| noise = torch.zeros(1, 128, 86) |
| time = torch.zeros(1) |
| zero = torch.zeros(1, 86, 2048) |
| context = torch.ones_like(zero) |
| left = adapter.predict_cfg_velocity( |
| noise, time, zero, guidance_scale=1.7 |
| ) |
| with torch.no_grad(): |
| adapter.q_up[0].fill_(0.1) |
| right = adapter.predict_cfg_velocity( |
| noise, time, context, guidance_scale=1.7 |
| ) |
| assert not torch.equal(left, right) |
| finally: |
| adapter.close() |
|
|
|
|
| def test_exact_zero_trim_retains_near_zero_and_composed_metrics_are_finite() -> None: |
| source = torch.ones(1, 2, 5000) * 1e-12 |
| source[:, :, -4:] = 0 |
| assert analyze_tail(source).shape[-1] == 4996 |
| append = torch.ones(1, 2, 44032) * 0.01 |
| metric = composed_seam_metrics( |
| source[:, :, :4996], |
| append, |
| overlap_samples=1024, |
| derivative_absolute_floor=1e-5, |
| rms_absolute_floor=1e-4, |
| ) |
| assert metric.boundary_derivative_ratio >= 0 |
| assert metric.overlap_rms_log_error >= 0 |
|
|
|
|
| def test_crop_and_schedule_endpoints_are_deterministic() -> None: |
| digest = "1" * 64 |
| first = deterministic_pair_offset( |
| digest, frame_count=1_000_000, guard_samples=220_500 |
| ) |
| second = deterministic_pair_offset( |
| digest, frame_count=1_000_000, guard_samples=220_500 |
| ) |
| assert first == second |
| assert cosine_learning_rate(0, 1200, 5e-5, 5e-6) == 5e-5 |
| assert abs( |
| cosine_learning_rate(1199, 1200, 5e-5, 5e-6) - 5e-6 |
| ) < 1e-12 |
|
|