* feat(phase3): PyTorch→MLX weight converter + safetensors output 365 keys mapped (367 source - 2 num_batches_tracked skipped). Conv weights transposed (O,I,H,W→O,H,W,I), refiner stem remapped, 4ch patch embed preserved. 12 conversion tests + diagnostic report. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix(phase3): conv allowlist, unused var, test fixture - Replace _is_conv_weight heuristic with explicit CONV_WEIGHT_KEYS frozenset (15 keys) - Remove unused `skipped` list from convert_state_dict - Module-scoped pytest fixture eliminates 11 redundant checkpoint loads Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2.3 KiB
| title | type | date |
|---|---|---|
| Fix converter review feedback (items 4, 1, 2) | fix | 2026-03-01 |
Fix converter review feedback
Three fixes from PR #1 review, ordered by simplicity.
Fix 4: Remove unused skipped list
File: src/corridorkey_mlx/convert/converter.py:105,110
Delete the skipped list variable and its .append() call. The skip logic already works via continue — the list serves no purpose.
Fix 1: Explicit conv key allowlist
File: src/corridorkey_mlx/convert/converter.py:60-66
Replace heuristic _is_conv_weight() with an explicit CONV_WEIGHT_KEYS: frozenset[str] allowlist. This matches the "no regex guessing, no silent fallbacks" philosophy.
The 15 conv keys (from checkpoint analysis + tests):
encoder.model.patch_embed.proj.weight— (112,4,7,7)alpha_decoder.linear_fuse.weight— (256,1024,1,1)fg_decoder.linear_fuse.weight— (256,1024,1,1)refiner.stem.0.weight(pre-remap) — (64,7,3,3)refiner.res{1-4}.conv{1-2}.weight— 8 keys, all (64,64,3,3)refiner.final.weight— (4,64,1,1)
Note: Use pre-remap key names since _is_conv_weight is called before remapping in the loop. Total: 15 keys.
Replace _is_conv_weight(key, arr) → simple key in CONV_WEIGHT_KEYS check. Remove the arr parameter.
Fix 2: Module-scoped test fixture
File: tests/test_conversion.py
Add a @pytest.fixture(scope="module") that loads + converts once, shared by all tests. This avoids 11 redundant load_pytorch_checkpoint + convert_state_dict calls.
@pytest.fixture(scope="module")
def converted_checkpoint():
if not CHECKPOINT_PATH.exists():
pytest.skip("Checkpoint not found")
state_dict = load_pytorch_checkpoint(CHECKPOINT_PATH)
converted, diagnostics = convert_state_dict(state_dict)
return state_dict, converted, diagnostics
Each test method takes converted_checkpoint as param, destructures what it needs. Remove _skip_if_no_checkpoint() helper.
Acceptance Criteria
skippedlist removed fromconvert_state_dictCONV_WEIGHT_KEYSfrozenset with 15 keys replaces_is_conv_weightheuristic- Module-scoped fixture eliminates repeated checkpoint loading
uv run pytest tests/test_conversion.py -v— 12/12 passuv run ruff check .— clean