File size: 5,260 Bytes
74e4281
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Sortformer Output Layer Extension: Script Architecture

## Background

Extending the Sortformer from N to M speakers involves three sequential operations:

| Stage | Operation | Input | Output |
|---|---|---|---|
| 1 | Extend weights (SVD orthogonal init) + create split head | N-spk unified model | M-spk split-head model |
| 2 | Fix config (recursively update all `num_spks`) | M-spk model with stale config fields | M-spk model with consistent config |
| 3 | Merge split head back to unified head | Post-finetune split-head model | Standard unified inference model |

## Architecture Decision

```
extend_output_layer.py     →  Stages 1 + 2  (extension + config fix)
merge_split_head.py        →  Stage 3        (post-finetune merge)
```

### Why merge Stages 1+2

- Stage 1 and Stage 2 are **always executed together** after extension. There is no valid use case for extending weights without fixing the config — doing so produces a checkpoint with stale `num_spks` fields in `train_ds`/`validation_ds`/`test_ds`, which causes `hungarian_get_perm` IndexError during training or inference.
- Running both in a single script avoids intermediate `.nemo` file I/O (load + save ~10-15s for a 500MB checkpoint).
- Eliminates the "forgot to run step 2" class of bugs entirely.

### Why Stage 3 is separate

- Stage 1+2 runs **before** finetuning. Stage 3 runs **after** finetuning.
- The two workflows are temporally separated (days/weeks apart in practice) and have completely different inputs: Stage 1+2 takes a pretrained N-spk checkpoint, Stage 3 takes a finetuned split-head checkpoint.
- Keeping Stage 3 as its own script makes it discoverable and self-documenting — a user looking at `src/finetune_pipeline/scripts/` sees `merge_split_head.py` and immediately understands its purpose.

## Usage

```bash
# --- Pre-finetune: extend 4spk → 10spk (Stages 1+2) ---
python src/finetune_pipeline/scripts/extend_output_layer.py \
    --src checkpoint_4spk.nemo \
    --dst-spk 10 \
    --out checkpoint_10spk_extended.nemo

# --- Finetune with split head (differential LR) ---
bash train.sh \
    --init_nemo_path checkpoint_10spk_extended.nemo ...

# --- Post-finetune: merge split head → unified inference model ---
python src/finetune_pipeline/scripts/merge_split_head.py \
    --src exp/.../checkpoints/model--val_loss=xxx.nemo \
    --out checkpoint_10spk_inference.nemo
```

## Implementation Details

### `extend_output_layer.py` (Stages 1+2)

1. **Stage 1 (weight extension):**
   - Load source model, extract unified weight matrix from either unified or already-split head via `get_unified_output_weights()`
   - Generate new rows via `orthogonal_extend_weight()` (SVD-based: uses right singular vectors as directions for new rows, falls back to random-magnitude noise when Vh is exhausted)
   - Extend bias similarly with random values centered at source bias statistics
   - Build target model with split head: `single_hidden_to_spks_base` (n_base_spks=N) + `single_hidden_to_spks_new` (n_new=M-N)
   - Copy matching-shape weights from source; inject extended base/new weights

2. **Stage 2 (config fix):**
   - `_update_all_num_spks(cfg, n_dst)` recursively traverses the OmegaConf tree, replacing every `max_num_of_spks` and `num_spks` with the target value
   - This covers `train_ds`, `validation_ds`, `test_ds`, and any nested config sections
   - Sets `sortformer_modules.n_base_spks = n_src` to enable differential LR during subsequent finetuning

### `merge_split_head.py` (Stage 3)

1. Load finetuned split-head model; verify `n_base_spks > 0` and split keys exist
2. Concatenate `base.weight` with `new.weight` (and biases) along dim=0
3. Build target config with `n_base_spks=0` → model uses unified `single_hidden_to_spks` path
4. Copy all state dict entries from finetuned model, excluding the 4 split head keys, then assign the concatenated weights to `single_hidden_to_spks.{weight,bias}`
5. Save unified inference checkpoint

### Shared Constants

Both scripts share the same state dict key constants:
```python
SK_BASE_W = "sortformer_modules.single_hidden_to_spks_base.weight"
SK_BASE_B = "sortformer_modules.single_hidden_to_spks_base.bias"
SK_NEW_W  = "sortformer_modules.single_hidden_to_spks_new.weight"
SK_NEW_B  = "sortformer_modules.single_hidden_to_spks_new.bias"
SK_UNI_W  = "sortformer_modules.single_hidden_to_spks.weight"
SK_UNI_B  = "sortformer_modules.single_hidden_to_spks.bias"
```

And `_update_all_num_spks()`:

```python
def _update_all_num_spks(cfg, target: int):
    """Replace every num_spks / max_num_of_spks occurrence in the config tree."""
    if OmegaConf.is_list(cfg):
        return
    if "max_num_of_spks" in cfg:
        cfg.max_num_of_spks = target
    if "num_spks" in cfg:
        cfg.num_spks = target
    for key in list(cfg.keys()):
        val = cfg[key]
        if OmegaConf.is_dict(val) or OmegaConf.is_config(val):
            _update_all_num_spks(val, target)
```

## Existing Script Disposition

| Script | Disposition |
|---|---|
| `extend_output_layer.py` | Rewritten — now includes Stage 2 config fix |
| `convert_extended_sortformer.py` | **Deleted** — logic absorbed into `extend_output_layer.py` |
| `merge_split_head.py` | **New** — standalone Stage 3 script |