| 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 |