Remap fm_decoder.encoders.{stack}.layers.{layer}.* to fm_decoder.layers.{flat}.* Also handle fm_decoder.encoders.{stack}.encoder.layers.{layer}.* for downsampled stacks.
(key: str)
| 22 | |
| 23 | |
| 24 | def flatten_fm_key(key: str) -> str: |
| 25 | """Remap fm_decoder.encoders.{stack}.layers.{layer}.* to fm_decoder.layers.{flat}.* |
| 26 | Also handle fm_decoder.encoders.{stack}.encoder.layers.{layer}.* for downsampled stacks. |
| 27 | """ |
| 28 | import re |
| 29 | |
| 30 | # Pattern 1: fm_decoder.encoders.{stack}.layers.{layer}.{rest} |
| 31 | m = re.match(r'fm_decoder\.encoders\.(\d+)\.layers\.(\d+)\.(.*)', key) |
| 32 | if m: |
| 33 | stack = int(m.group(1)) |
| 34 | layer = int(m.group(2)) |
| 35 | rest = m.group(3) |
| 36 | flat = sum(STACK_SIZES[:stack]) + layer |
| 37 | return f'fm_decoder.layers.{flat}.{rest}' |
| 38 | |
| 39 | # Pattern 2: fm_decoder.encoders.{stack}.encoder.layers.{layer}.{rest} (downsampled stacks) |
| 40 | m = re.match(r'fm_decoder\.encoders\.(\d+)\.encoder\.layers\.(\d+)\.(.*)', key) |
| 41 | if m: |
| 42 | stack = int(m.group(1)) |
| 43 | layer = int(m.group(2)) |
| 44 | rest = m.group(3) |
| 45 | flat = sum(STACK_SIZES[:stack]) + layer |
| 46 | return f'fm_decoder.layers.{flat}.{rest}' |
| 47 | |
| 48 | # Pattern 3: fm_decoder.encoders.{stack}.encoder.time_emb.{rest} → fm_decoder.stack_time_emb.{stack}.{rest} |
| 49 | m = re.match(r'fm_decoder\.encoders\.(\d+)\.encoder\.time_emb\.(.*)', key) |
| 50 | if m: |
| 51 | stack = int(m.group(1)) |
| 52 | rest = m.group(2) |
| 53 | return f'fm_decoder.stack_time_emb.{stack}.{rest}' |
| 54 | |
| 55 | # Pattern 4: fm_decoder.encoders.{stack}.time_emb.{rest} → fm_decoder.stack_time_emb.{stack}.{rest} |
| 56 | m = re.match(r'fm_decoder\.encoders\.(\d+)\.time_emb\.(.*)', key) |
| 57 | if m: |
| 58 | stack = int(m.group(1)) |
| 59 | rest = m.group(2) |
| 60 | return f'fm_decoder.stack_time_emb.{stack}.{rest}' |
| 61 | |
| 62 | # Pattern 5: fm_decoder.encoders.{stack}.downsample.{rest} → fm_decoder.downsample.{stack}.{rest} |
| 63 | m = re.match(r'fm_decoder\.encoders\.(\d+)\.downsample\.(.*)', key) |
| 64 | if m: |
| 65 | stack = int(m.group(1)) |
| 66 | rest = m.group(2) |
| 67 | return f'fm_decoder.downsample.{stack}.{rest}' |
| 68 | |
| 69 | # Pattern 6: fm_decoder.encoders.{stack}.out_combiner.{rest} → fm_decoder.out_combiner.{stack}.{rest} |
| 70 | m = re.match(r'fm_decoder\.encoders\.(\d+)\.out_combiner\.(.*)', key) |
| 71 | if m: |
| 72 | stack = int(m.group(1)) |
| 73 | rest = m.group(2) |
| 74 | return f'fm_decoder.out_combiner.{stack}.{rest}' |
| 75 | |
| 76 | # Similarly for text_encoder |
| 77 | m = re.match(r'text_encoder\.encoders\.0\.layers\.(\d+)\.(.*)', key) |
| 78 | if m: |
| 79 | layer = int(m.group(1)) |
| 80 | rest = m.group(2) |
| 81 | return f'text_encoder.layers.{layer}.{rest}' |