Skip to content

Commit 1d75ab7

Browse files
committed
update
1 parent b086b6d commit 1d75ab7

1 file changed

Lines changed: 53 additions & 7 deletions

File tree

utils/make_tiny_model.py

Lines changed: 53 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,33 @@
4646
import re
4747

4848

49-
LAYER_PARAM_PATTERN = re.compile(r"^(num_.*layers?|n_layers)$")
49+
LAYER_PARAM_PATTERN = re.compile(r"^(num_.*layers?|n_layers|n_refiner_layers)$")
50+
51+
DIM_PARAM_PATTERNS = {
52+
re.compile(r"^num_attention_heads$"): 2,
53+
re.compile(r"^num_.*attention_heads$"): 2,
54+
re.compile(r"^num_key_value_heads$"): 2,
55+
re.compile(r"^num_kv_heads$"): 1,
56+
re.compile(r"^n_heads$"): 2,
57+
re.compile(r"^n_kv_heads$"): 2,
58+
re.compile(r"^attention_head_dim$"): 8,
59+
re.compile(r"^.*attention_head_dim$"): 4,
60+
re.compile(r"^cross_attention_dim.*$"): 8,
61+
re.compile(r"^joint_attention_dim$"): 32,
62+
re.compile(r"^pooled_projection_dim$"): 32,
63+
re.compile(r"^caption_projection_dim$"): 32,
64+
re.compile(r"^caption_channels$"): 8,
65+
re.compile(r"^cap_feat_dim$"): 16,
66+
re.compile(r"^hidden_size$"): 16,
67+
re.compile(r"^dim$"): 16,
68+
re.compile(r"^.*embed_dim$"): 16,
69+
re.compile(r"^.*embed_.*dim$"): 16,
70+
re.compile(r"^text_dim$"): 16,
71+
re.compile(r"^time_embed_dim$"): 4,
72+
re.compile(r"^ffn_dim$"): 32,
73+
re.compile(r"^intermediate_size$"): 32,
74+
re.compile(r"^sample_size$"): 32,
75+
}
5076

5177

5278
def parse_args():
@@ -60,6 +86,11 @@ def parse_args():
6086
)
6187
parser.add_argument("--subfolder", type=str, default=None, help="Subfolder within the model repo.")
6288
parser.add_argument("--num_layers", type=int, default=2, help="Number of layers to use for the tiny model.")
89+
parser.add_argument(
90+
"--shrink_dims",
91+
action="store_true",
92+
help="Also reduce dimension parameters (attention heads, hidden size, embedding dims, etc.).",
93+
)
6394
parser.add_argument("--push_to_hub", action="store_true", help="Push the tiny model to the HuggingFace Hub.")
6495
parser.add_argument(
6596
"--token", type=str, default=None, help="HuggingFace token. Defaults to $HF_TOKEN env var if not provided."
@@ -89,6 +120,8 @@ def launch_job(args):
89120
]
90121
if args.subfolder:
91122
script_args.extend(["--subfolder", args.subfolder])
123+
if args.shrink_dims:
124+
script_args.append("--shrink_dims")
92125
if args.push_to_hub:
93126
script_args.append("--push_to_hub")
94127

@@ -104,7 +137,9 @@ def launch_job(args):
104137
return job
105138

106139

107-
def make_tiny_model(model_repo_id, output_repo_id, subfolder=None, num_layers=2, push_to_hub=False, token=None):
140+
def make_tiny_model(
141+
model_repo_id, output_repo_id, subfolder=None, num_layers=2, shrink_dims=False, push_to_hub=False, token=None
142+
):
108143
from diffusers import AutoModel
109144

110145
config_kwargs = {}
@@ -119,14 +154,24 @@ def make_tiny_model(model_repo_id, output_repo_id, subfolder=None, num_layers=2,
119154
modified_keys[key] = (value, num_layers)
120155
config[key] = num_layers
121156

157+
if shrink_dims:
158+
for key, value in config.items():
159+
if not isinstance(value, int) or key.startswith("_"):
160+
continue
161+
for pattern, tiny_value in DIM_PARAM_PATTERNS.items():
162+
if pattern.match(key) and value > tiny_value:
163+
modified_keys[key] = (value, tiny_value)
164+
config[key] = tiny_value
165+
break
166+
122167
if not modified_keys:
123-
print(f"WARNING: No layer parameters found matching pattern '{LAYER_PARAM_PATTERN.pattern}' in config.")
168+
print("WARNING: No config parameters were modified.")
124169
print(f"Config keys: {[k for k in config if not k.startswith('_')]}")
125170
return
126-
else:
127-
print("Modified layer parameters:")
128-
for key, (old, new) in modified_keys.items():
129-
print(f" {key}: {old} -> {new}")
171+
172+
print("Modified config parameters:")
173+
for key, (old, new) in modified_keys.items():
174+
print(f" {key}: {old} -> {new}")
130175

131176
model = AutoModel.from_config(config)
132177
total_params = sum(p.numel() for p in model.parameters())
@@ -155,6 +200,7 @@ def main():
155200
output_repo_id=args.output_repo_id,
156201
subfolder=args.subfolder,
157202
num_layers=args.num_layers,
203+
shrink_dims=args.shrink_dims,
158204
push_to_hub=args.push_to_hub,
159205
token=args.token,
160206
)

0 commit comments

Comments
 (0)