返回 F5-TTS
convert_checkpoint.py
根目录 / src / f5_tts / runtime / triton_trtllm / scripts / convert_checkpoint.py
1 import argparse
2 import json
3 import os
4 import re
5 import time
6 import traceback
7 from concurrent.futures import ThreadPoolExecutor, as_completed
8
9 import safetensors.torch
10 import torch
11 from tensorrt_llm import str_dtype_to_torch
12 from tensorrt_llm.mapping import Mapping
13 from tensorrt_llm.models.convert_utils import split, split_matrix_tp
14
15
16 def split_q_tp(v, n_head, n_hidden, tensor_parallel, rank):
17 split_v = split(v, tensor_parallel, rank, dim=1)
18 return split_v.contiguous()
19
20
21 def split_q_bias_tp(v, n_head, n_hidden, tensor_parallel, rank):
22 split_v = split(v, tensor_parallel, rank, dim=0)
23 return split_v.contiguous()
24
25
26 def parse_arguments():
27 parser = argparse.ArgumentParser()
28 parser.add_argument("--pytorch_ckpt", type=str, default="./ckpts/model_last.pt")
29 parser.add_argument(
30 "--output_dir", type=str, default="./tllm_checkpoint", help="The path to save the TensorRT-LLM checkpoint"
31 )
32 parser.add_argument("--tp_size", type=int, default=1, help="N-way tensor parallelism size")
33 parser.add_argument("--cp_size", type=int, default=1, help="Context parallelism size")
34 parser.add_argument("--pp_size", type=int, default=1, help="N-way pipeline parallelism size")
35 parser.add_argument("--dtype", type=str, default="float16", choices=["float32", "bfloat16", "float16"])
36 parser.add_argument("--fp8_linear", action="store_true", help="Whether use FP8 for linear layers")
37 parser.add_argument(
38 "--workers", type=int, default=1, help="The number of workers for converting checkpoint in parallel"
39 )
40 parser.add_argument(
41 "--model_name",
42 type=str,
43 default="F5TTS_Custom",
44 choices=[
45 "F5TTS_v1_Base",
46 "F5TTS_Base",
47 "F5TTS_v1_Small",
48 "F5TTS_Small",
49 ], # if set, overwrite the below hyperparams
50 )
51 parser.add_argument("--hidden_size", type=int, default=1024, help="The hidden size of DiT")
52 parser.add_argument("--depth", type=int, default=22, help="The number of DiTBlock layers")
53 parser.add_argument("--num_heads", type=int, default=16, help="The number of heads of attention module")
54 parser.add_argument("--dim_head", type=int, default=64, help="The dimension of attention head")
55 parser.add_argument("--ff_mult", type=int, default=2, help="The FFN intermediate dimension multiplier")
56 parser.add_argument("--text_dim", type=int, default=512, help="The output dimension of text encoder")
57 parser.add_argument(
58 "--text_mask_padding",
59 type=lambda x: x.lower() == "true",
60 choices=[True, False],
61 default=True,
62 help="Whether apply padding mask for conv layers in text encoder",
63 )
64 parser.add_argument("--conv_layers", type=int, default=4, help="The number of conv layers of text encoder")
65 parser.add_argument("--pe_attn_head", type=int, default=None, help="The number of attn head that apply pos emb")
66 args = parser.parse_args()
67
68 # overwrite if --model_name ordered
69 if args.model_name == "F5TTS_v1_Base":
70 args.hidden_size = 1024
71 args.depth = 22
72 args.num_heads = 16
73 args.dim_head = 64
74 args.ff_mult = 2
75 args.text_dim = 512
76 args.text_mask_padding = True
77 args.conv_layers = 4
78 args.pe_attn_head = None
79 elif args.model_name == "F5TTS_Base":
80 args.hidden_size = 1024
81 args.depth = 22
82 args.num_heads = 16
83 args.dim_head = 64
84 args.ff_mult = 2
85 args.text_dim = 512
86 args.text_mask_padding = False
87 args.conv_layers = 4
88 args.pe_attn_head = 1
89 elif args.model_name == "F5TTS_v1_Small":
90 args.hidden_size = 768
91 args.depth = 18
92 args.num_heads = 12
93 args.dim_head = 64
94 args.ff_mult = 2
95 args.text_dim = 512
96 args.text_mask_padding = True
97 args.conv_layers = 4
98 args.pe_attn_head = None
99 elif args.model_name == "F5TTS_Small":
100 args.hidden_size = 768
101 args.depth = 18
102 args.num_heads = 12
103 args.dim_head = 64
104 args.ff_mult = 2
105 args.text_dim = 512
106 args.text_mask_padding = False
107 args.conv_layers = 4
108 args.pe_attn_head = 1
109
110 return args
111
112
113 def convert_pytorch_dit_to_trtllm_weight(args, mapping, dtype="float32", use_ema=True):
114 weights = {}
115 tik = time.time()
116 torch_dtype = str_dtype_to_torch(dtype)
117 tensor_parallel = mapping.tp_size
118
119 ckpt_path = args.pytorch_ckpt
120 ckpt_type = ckpt_path.split(".")[-1]
121 if ckpt_type == "safetensors":
122 from safetensors.torch import load_file
123
124 model_params = load_file(ckpt_path)
125 else:
126 ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=True)
127 model_params = ckpt["ema_model_state_dict"] if use_ema else ckpt["model_state_dict"]
128
129 prefix = "ema_model.transformer." if use_ema else "transformer."
130 if any(k.startswith(prefix) for k in model_params.keys()):
131 model_params = {
132 key[len(prefix) :] if key.startswith(prefix) else key: value
133 for key, value in model_params.items()
134 if key.startswith(prefix)
135 }
136
137 pytorch_to_trtllm_name = {
138 r"^time_embed\.time_mlp\.0\.(weight|bias)$": r"time_embed.mlp1.\1",
139 r"^time_embed\.time_mlp\.2\.(weight|bias)$": r"time_embed.mlp2.\1",
140 r"^input_embed\.conv_pos_embed\.conv1d\.0\.(weight|bias)$": r"input_embed.conv_pos_embed.conv1d1.\1",
141 r"^input_embed\.conv_pos_embed\.conv1d\.2\.(weight|bias)$": r"input_embed.conv_pos_embed.conv1d2.\1",
142 r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(weight|bias)$": r"transformer_blocks.\1.attn.to_out.\2",
143 r"^transformer_blocks\.(\d+)\.ff\.ff\.0\.0\.(weight|bias)$": r"transformer_blocks.\1.ff.project_in.\2",
144 r"^transformer_blocks\.(\d+)\.ff\.ff\.2\.(weight|bias)$": r"transformer_blocks.\1.ff.ff.\2",
145 }
146
147 def get_trtllm_name(pytorch_name):
148 for pytorch_name_pattern, trtllm_name_replacement in pytorch_to_trtllm_name.items():
149 trtllm_name_if_matched = re.sub(pytorch_name_pattern, trtllm_name_replacement, pytorch_name)
150 if trtllm_name_if_matched != pytorch_name:
151 return trtllm_name_if_matched
152 return pytorch_name
153
154 weights = dict()
155 for name, param in model_params.items():
156 if name == "input_embed.conv_pos_embed.conv1d.0.weight" or name == "input_embed.conv_pos_embed.conv1d.2.weight":
157 weights[get_trtllm_name(name)] = param.contiguous().to(torch_dtype).unsqueeze(-1)
158 else:
159 weights[get_trtllm_name(name)] = param.contiguous().to(torch_dtype)
160
161 assert len(weights) == len(model_params)
162
163 # new_prefix = "f5_transformer."
164 new_prefix = ""
165 weights = {new_prefix + key: value for key, value in weights.items()}
166 import math
167
168 scale_factor = math.pow(64, -0.25)
169 for k, v in weights.items():
170 if re.match("^transformer_blocks.*.attn.to_k.weight$", k):
171 weights[k] *= scale_factor
172 weights[k] = split_q_tp(v, args.num_heads, args.hidden_size, tensor_parallel, mapping.tp_rank)
173
174 elif re.match("^transformer_blocks.*.attn.to_k.bias$", k):
175 weights[k] *= scale_factor
176 weights[k] = split_q_bias_tp(v, args.num_heads, args.hidden_size, tensor_parallel, mapping.tp_rank)
177
178 elif re.match("^transformer_blocks.*.attn.to_q.weight$", k):
179 weights[k] = split_q_tp(v, args.num_heads, args.hidden_size, tensor_parallel, mapping.tp_rank)
180 weights[k] *= scale_factor
181
182 elif re.match("^transformer_blocks.*.attn.to_q.bias$", k):
183 weights[k] = split_q_bias_tp(v, args.num_heads, args.hidden_size, tensor_parallel, mapping.tp_rank)
184 weights[k] *= scale_factor
185
186 elif re.match("^transformer_blocks.*.attn.to_v.weight$", k):
187 weights[k] = split_q_tp(v, args.num_heads, args.hidden_size, tensor_parallel, mapping.tp_rank)
188
189 elif re.match("^transformer_blocks.*.attn.to_v.bias$", k):
190 weights[k] = split_q_bias_tp(v, args.num_heads, args.hidden_size, tensor_parallel, mapping.tp_rank)
191
192 elif re.match("^transformer_blocks.*.attn.to_out.weight$", k):
193 weights[k] = split_matrix_tp(v, tensor_parallel, mapping.tp_rank, dim=1)
194
195 tok = time.time()
196 t = time.strftime("%H:%M:%S", time.gmtime(tok - tik))
197 print(f"Weights loaded. Total time: {t}")
198 return weights
199
200
201 def save_config(args):
202 if not os.path.exists(args.output_dir):
203 os.makedirs(args.output_dir)
204 config = {
205 "architecture": "F5TTS", # set the same as in ../patch/__init__.py
206 "dtype": args.dtype,
207 "hidden_size": args.hidden_size,
208 "num_hidden_layers": args.depth,
209 "num_attention_heads": args.num_heads,
210 "dim_head": args.dim_head,
211 "dropout": 0.0, # inference-only
212 "ff_mult": args.ff_mult,
213 "mel_dim": 100,
214 "text_dim": args.text_dim,
215 "text_mask_padding": args.text_mask_padding,
216 "conv_layers": args.conv_layers,
217 "pe_attn_head": args.pe_attn_head,
218 "mapping": {
219 "world_size": args.cp_size * args.tp_size * args.pp_size,
220 "cp_size": args.cp_size,
221 "tp_size": args.tp_size,
222 "pp_size": args.pp_size,
223 },
224 }
225 if args.fp8_linear:
226 config["quantization"] = {
227 "quant_algo": "FP8",
228 # TODO: add support for exclude modules.
229 # "exclude_modules": "*final_layer*",
230 }
231
232 with open(os.path.join(args.output_dir, "config.json"), "w") as f:
233 json.dump(config, f, indent=4)
234
235
236 def covert_and_save(args, rank):
237 if rank == 0:
238 save_config(args)
239
240 mapping = Mapping(
241 world_size=args.cp_size * args.tp_size * args.pp_size,
242 rank=rank,
243 cp_size=args.cp_size,
244 tp_size=args.tp_size,
245 pp_size=args.pp_size,
246 )
247
248 weights = convert_pytorch_dit_to_trtllm_weight(args, mapping, dtype=args.dtype)
249
250 safetensors.torch.save_file(weights, os.path.join(args.output_dir, f"rank{rank}.safetensors"))
251
252
253 def execute(workers, func, args):
254 if workers == 1:
255 for rank, f in enumerate(func):
256 f(args, rank)
257 else:
258 with ThreadPoolExecutor(max_workers=workers) as p:
259 futures = [p.submit(f, args, rank) for rank, f in enumerate(func)]
260 exceptions = []
261 for future in as_completed(futures):
262 try:
263 future.result()
264 except Exception as e:
265 traceback.print_exc()
266 exceptions.append(e)
267 assert len(exceptions) == 0, "Checkpoint conversion failed, please check error log."
268
269
270 def main():
271 args = parse_arguments()
272 world_size = args.cp_size * args.tp_size * args.pp_size
273
274 assert args.pp_size == 1, "PP is not supported yet."
275
276 tik = time.time()
277 if args.pytorch_ckpt is None:
278 return
279 print("Start execute")
280 execute(args.workers, [covert_and_save] * world_size, args)
281
282 tok = time.time()
283 t = time.strftime("%H:%M:%S", time.gmtime(tok - tik))
284 print(f"Total time of converting checkpoints: {t}")
285
286
287 if __name__ == "__main__":
288 main()
289
289 lines PYTHON