MCPcopy Create free account
hub / github.com/evilsocket/cake / flatten_fm_key

Function flatten_fm_key

scripts/convert_luxtts.py:24–83  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

22
23
24def 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}'

Callers 1

convert_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected