From c17592ca88f7e41b0f44cf65697e1eda0944f3b6 Mon Sep 17 00:00:00 2001 From: "vladimir.talpa" Date: Wed, 26 Aug 2026 13:37:11 +0300 Subject: [PATCH 1/5] Fix ValueError in FireRedTTS3Instruct.generate_tts MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The text-frontend wrapper unpacked three values from the backend generate_tts, which returns only (gen_audio, gen_audio_sr) — matching the two-value form documented in the README. Every call to fireredtts3.core.FireRedTTS3Instruct.generate_tts therefore raised ValueError: not enough values to unpack (expected 3, got 2) regardless of device. Co-Authored-By: Claude Opus 5 --- fireredtts3/core.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fireredtts3/core.py b/fireredtts3/core.py index 5b4578c..5b19d50 100644 --- a/fireredtts3/core.py +++ b/fireredtts3/core.py @@ -381,7 +381,7 @@ def generate_tts( gen_audio_sr = None segments: List[torch.Tensor] = [] for sent in sentences: - seg, seg_sr, _ = super().generate_tts( + seg, seg_sr = super().generate_tts( prompt_text=prompt_text, prompt_audio=prompt_audio, prompt_audio_sr=prompt_audio_sr, From 7cc41f74414f8b8ee63e372adc3835a5df302c01 Mon Sep 17 00:00:00 2001 From: "vladimir.talpa" Date: Wed, 26 Aug 2026 13:37:27 +0300 Subject: [PATCH 2/5] Support MPS (Apple Silicon) and CPU in addition to CUDA Device, compute dtype, attention implementation and FFT capability are now resolved once in fireredtts3/utils/device.py instead of being hard-coded to CUDA, so the model code stays free of per-backend branches. Replaced CUDA-only constructs: - torch.device('cuda') -> get_device() - @torch.autocast(device_type='cuda', ...) -> @autocast - @autocast('cuda', enabled=False) -> @disable_autocast - attn_implementation='flash_attention_2' -> get_attn_implementation() (flash-attn has no macOS wheel; SDPA covers MPS/CPU) - unguarded torch.cuda.manual_seed* -> guarded, + torch.mps Backends without FFT / complex kernels (MPS on older torch) now compute the RedAE ISTFT and the CAM++ kaldi fbank on CPU, detected by probe rather than by version check. Added FIRERED_WEIGHT_DTYPE so checkpoints can be loaded in half precision, which halves resident memory (11.44 -> 5.73 GiB) at equal speaker similarity. Autocast alone cannot do this, as it only casts activations. Making that work required the dtype to follow the weights instead of being pinned to fp32 at the pipeline entry points (prompt latents, speaker embedding, flow-matching noise and time span); the ISTFT stays in fp32 since complex half is unusable. requirements-mps.txt drops flash_attn and adds soundfile, because torchcodec 0.7 does not load against FFmpeg 8 and torchaudio then has no I/O backend at all. Verified on an M4 Max (macOS 26, torch 2.8) across all five paths: Base and Instruct zero-shot cloning, voice design, acoustic edit and semantic edit. Co-Authored-By: Claude Opus 5 --- .gitignore | 5 +- README.md | 36 +++++ fireredtts3/campp/campp.py | 3 + fireredtts3/llm/fireredtts3_base.py | 20 +-- fireredtts3/llm/fireredtts3_instruct.py | 18 +-- fireredtts3/llm/rotary_embedding.py | 6 +- fireredtts3/redae/redae.py | 26 ++-- fireredtts3/utils/device.py | 124 +++++++++++++++++ fireredtts3/utils/utils.py | 7 +- requirements-mps.txt | 14 ++ scripts/mps_smoke_test.py | 178 ++++++++++++++++++++++++ scripts/spk_similarity.py | 27 ++++ 12 files changed, 433 insertions(+), 31 deletions(-) create mode 100644 fireredtts3/utils/device.py create mode 100644 requirements-mps.txt create mode 100644 scripts/mps_smoke_test.py create mode 100644 scripts/spk_similarity.py diff --git a/.gitignore b/.gitignore index 70e8e64..c2a55a3 100644 --- a/.gitignore +++ b/.gitignore @@ -14,4 +14,7 @@ pretrained_models/ # Eval data -data/ \ No newline at end of file +data/ +# Local venv / test output +.venv/ +outputs/ diff --git a/README.md b/README.md index 7fbd3c0..b989d0d 100644 --- a/README.md +++ b/README.md @@ -79,6 +79,42 @@ pip install -r requirements.txt ``` +### Installation on Apple Silicon (MPS) + +`flash_attn` has no macOS wheel, and `torchcodec` needs FFmpeg <= 7, so use the +macOS requirement set instead (attention falls back to PyTorch SDPA, audio I/O +to `soundfile`): + +```sh +uv venv --python 3.12 .venv +uv pip install --python .venv/bin/python -r requirements-mps.txt +``` + +The device and compute dtype are resolved once in `fireredtts3/utils/device.py` +(CUDA -> MPS -> CPU) and can be overridden: + +| Variable | Values | Default | +| --- | --- | --- | +| `FIRERED_DEVICE` | `cuda` / `mps` / `cpu` | first available | +| `FIRERED_DTYPE` | autocast dtype (activations) | `bfloat16` on CUDA, `float32` elsewhere | +| `FIRERED_WEIGHT_DTYPE` | dtype the checkpoints are loaded in | `float32` | + +`FIRERED_DTYPE` only casts activations, so it does **not** shrink resident +memory; `FIRERED_WEIGHT_DTYPE` is the knob that does. Measured on an M4 Max +(torch 2.8, macOS 26), Base zero-shot cloning of one short sentence: + +| Weights | Resident weights | Peak MPS | RTF | Speaker sim | +| --- | --- | --- | --- | --- | +| `float32` (default) | 11.44 GiB | 13.87 GiB | 1.15 | 0.895 | +| `bfloat16` | 5.73 GiB | 6.92 GiB | 1.47 | 0.913 | +| `float16` | 5.73 GiB | 6.92 GiB | 1.85 | 0.896 | + +Half weights halve memory at equal quality but are *slower* on MPS, and +autocast bf16 over fp32 weights buys nothing (RTF 0.88-0.93 vs 0.85-0.91) — +the DiT flow head, the real bottleneck, sits outside the autocast region. Hence +fp32 weights + no autocast is the default off CUDA. On longer multi-sentence +input the same fp32 path reaches RTF 0.85-0.91, i.e. faster than real time. + ### Model Download Download the pretrained model from Hugging Face with the `hf` CLI: diff --git a/fireredtts3/campp/campp.py b/fireredtts3/campp/campp.py index 586e853..47e2433 100644 --- a/fireredtts3/campp/campp.py +++ b/fireredtts3/campp/campp.py @@ -2,10 +2,13 @@ import torchaudio import torchaudio.compliance.kaldi as kaldi from fireredtts3.campp.DTDNN import CAMPPlus +from fireredtts3.utils.device import fft_device def extract_kaldi_mel(audio: torch.Tensor, audio_sr: int): audio = audio[:1] + # kaldi.fbank runs an FFT internally; fall back to CPU on devices without it + audio = fft_device(audio) if audio_sr != 16000: audio = torchaudio.functional.resample(audio, audio_sr, 16000) audio_sr = 16000 diff --git a/fireredtts3/llm/fireredtts3_base.py b/fireredtts3/llm/fireredtts3_base.py index a37bf9e..1050544 100644 --- a/fireredtts3/llm/fireredtts3_base.py +++ b/fireredtts3/llm/fireredtts3_base.py @@ -11,6 +11,7 @@ from fireredtts3.redae.redae import RedAE from fireredtts3.campp.campp import CamppEmbedding from fireredtts3.utils.utils import fix_seed +from fireredtts3.utils.device import get_device, get_attn_implementation, get_weight_dtype, autocast from fireredtts3.utils.text_tokenizer import ( load_text_tokenizer, MULTI_LANG_TAGS, MULTI_DIALECT_TAGS, @@ -46,7 +47,7 @@ "use_cache": True, "use_sliding_window": False, "vocab_size": 151936, - "attn_implementation": "flash_attention_2", + "attn_implementation": get_attn_implementation(), } @@ -140,7 +141,7 @@ def _init_weights(self, module): module.rope_init() # Backbone Transformer AR wrapper - @torch.autocast(device_type='cuda', dtype=torch.bfloat16) + @autocast def _backbone_one_step(self, input_embeds: torch.Tensor, cache = None): outs = self.backbone_llm.forward( inputs_embeds=input_embeds, @@ -162,7 +163,8 @@ def _flow_one_step( inference_cfg: float, ): # Compose input - x0 = torch.randn(1, self.patch_size, self.redae_dim, device=hist_latents.device) + x0 = torch.randn(1, self.patch_size, self.redae_dim, + device=hist_latents.device, dtype=hist_latents.dtype) xt = torch.cat([hist_latents, x0], dim=1) # History clean + current noise cond = torch.cat([ @@ -213,7 +215,7 @@ def generate( input_embeds = torch.cat([spk_embs_llm.unsqueeze(1), input_embeds, patch_prompt_latents], dim=1) # Prepare DiT decode - t_span = torch.linspace(0, 1, n_timesteps + 1).to(device) + t_span = torch.linspace(0, 1, n_timesteps + 1).to(device=device, dtype=input_embeds.dtype) t_span = 1 - torch.cos(t_span * 0.5 * torch.pi) # (n_timesteps+1,) latents_gen = F.pad(prompt_latents, (0, 0, self.history_length, 0)) @@ -264,16 +266,16 @@ def generate( # RedAE + TextTokenizer + TTS3Core class FireRedTTS3Base(object): def __init__(self, pretrained_model_dir: str): - self.device = torch.device('cuda') + self.device = get_device() # RedAE redae_model_dir = os.path.join(pretrained_model_dir, 'redae') assert os.path.exists(redae_model_dir), f'{redae_model_dir} not found' - self.redae = RedAE.from_pretrained(redae_model_dir) + self.redae = RedAE.from_pretrained(redae_model_dir, dtype=get_weight_dtype()) self.redae.to(self.device) # LLM-DiT tts_model_dir = os.path.join(pretrained_model_dir, 'fireredtts3_base') assert os.path.exists(tts_model_dir), f'{tts_model_dir} not found' - self.tts_core = FireRedTTS3BaseCore.from_pretrained(tts_model_dir) + self.tts_core = FireRedTTS3BaseCore.from_pretrained(tts_model_dir, dtype=get_weight_dtype()) self.tts_core.to(self.device) # Text Tokenizer text_tok_dir = os.path.join(pretrained_model_dir, 'text_tokenizer') @@ -320,10 +322,10 @@ def generate( prompt_audio = self.redae.pad_to_multiple_of(prompt_audio, self.redae.downsample_rate*self.tts_core.patch_size) prompt_audio = prompt_audio.to(self.device) prompt_latents = self.redae.encode(prompt_audio, prompt_audio_sr) - prompt_latents = prompt_latents.to(torch.float32) + prompt_latents = prompt_latents.to(self.tts_core.dtype) # Spk emb spk_emb = self.spk_extractor.forward(prompt_audio, prompt_audio_sr) - spk_emb = spk_emb.to(self.device) + spk_emb = spk_emb.to(device=self.device, dtype=self.tts_core.dtype) # TTS if seed is not None: fix_seed(seed) diff --git a/fireredtts3/llm/fireredtts3_instruct.py b/fireredtts3/llm/fireredtts3_instruct.py index b93c445..822acdc 100644 --- a/fireredtts3/llm/fireredtts3_instruct.py +++ b/fireredtts3/llm/fireredtts3_instruct.py @@ -17,6 +17,7 @@ from fireredtts3.llm.dit import DiT from fireredtts3.redae.redae import RedAE from fireredtts3.utils.utils import fix_seed +from fireredtts3.utils.device import get_device, get_weight_dtype, autocast from fireredtts3.utils.text_tokenizer import load_text_tokenizer from fireredtts3.llm.fireredtts3_base import Qwen3_1_7B_ConfigDict from fireredtts3.utils.chatml import ( @@ -139,7 +140,7 @@ def _init_weights(self, module): module.rope_init() # Backbone Transformer AR wrapper - @torch.autocast(device_type='cuda', dtype=torch.bfloat16) + @autocast def _backbone_one_step(self, input_embeds: torch.Tensor, cache = None): outs = self.backbone_llm.model.forward( inputs_embeds=input_embeds, @@ -160,7 +161,8 @@ def _flow_one_step( inference_cfg: float, ): # Compose input - x0 = torch.randn(1, self.patch_size, self.redae_dim, device=hist_latents.device) + x0 = torch.randn(1, self.patch_size, self.redae_dim, + device=hist_latents.device, dtype=hist_latents.dtype) xt = torch.cat([hist_latents, x0], dim=1) # History clean + current noise cond = backbone_cond.repeat_interleave(self.patch_size, dim=1) # Correspond backbone cond @@ -228,11 +230,11 @@ def generate( latents_patch_out.reshape(-1).to(input_embeds), ) else: - latents_out = torch.zeros(1, 0, self.config.redae_dim, device=device) + latents_out = torch.zeros(1, 0, self.config.redae_dim, device=device, dtype=input_embeds.dtype) latents_patch_out = None # Prepare DiT decode - t_span = torch.linspace(0, 1, n_timesteps + 1).to(device) + t_span = torch.linspace(0, 1, n_timesteps + 1).to(device=device, dtype=input_embeds.dtype) t_span = 1 - torch.cos(t_span * 0.5 * torch.pi) # (n_timesteps+1,) # Init Backbone states @@ -316,16 +318,16 @@ def generate( # RedAE + TextTokenizer + TTS3Core class FireRedTTS3Instruct(object): def __init__(self, pretrained_model_dir: str): - self.device = torch.device('cuda') + self.device = get_device() # RedAE redae_model_dir = os.path.join(pretrained_model_dir, 'redae') assert os.path.exists(redae_model_dir), f'{redae_model_dir} not found' - self.redae = RedAE.from_pretrained(redae_model_dir) + self.redae = RedAE.from_pretrained(redae_model_dir, dtype=get_weight_dtype()) self.redae.to(self.device) # LLM-DiT tts_model_dir = os.path.join(pretrained_model_dir, 'fireredtts3_instruct') assert os.path.exists(tts_model_dir), f'{tts_model_dir} not found' - self.tts_core = FireRedTTS3InstructCore.from_pretrained(tts_model_dir) + self.tts_core = FireRedTTS3InstructCore.from_pretrained(tts_model_dir, dtype=get_weight_dtype()) self.tts_core.to(self.device) # Text Tokenizer text_tok_dir = os.path.join(pretrained_model_dir, 'text_tokenizer') @@ -347,7 +349,7 @@ def _tokenize_audio(self, audio: torch.Tensor, audio_sr: int): audio = self.redae.pad_to_multiple_of(audio, self.redae.downsample_rate*self.tts_core.patch_size) audio = audio.to(self.device) latents = self.redae.encode(audio, audio_sr) * REDAE_SCALE - latents = latents.to(torch.float32) + latents = latents.to(self.tts_core.dtype) return latents # --- Inference Interface diff --git a/fireredtts3/llm/rotary_embedding.py b/fireredtts3/llm/rotary_embedding.py index 3d70110..a7b7858 100644 --- a/fireredtts3/llm/rotary_embedding.py +++ b/fireredtts3/llm/rotary_embedding.py @@ -1,8 +1,8 @@ import torch from torch.nn import Module -from torch.amp import autocast from torch import cat, stack, arange from einops import rearrange +from fireredtts3.utils.device import disable_autocast class RotaryEmbedding(Module): @@ -38,7 +38,7 @@ def forward_from_seq_len(self, seq_len): t = arange(seq_len, device = device) return self.forward(t) - @autocast('cuda', enabled = False) + @disable_autocast def forward(self, t, offset = 0): if t.ndim == 1: t = rearrange(t, 'n -> 1 n') @@ -57,7 +57,7 @@ def rotate_half(x): return rearrange(x, '... d r -> ... (d r)') -@autocast('cuda', enabled = False) +@disable_autocast def apply_rotary_pos_emb(t, freqs, scale = 1): rot_dim, seq_len, orig_dtype = freqs.shape[-1], t.shape[-2], t.dtype diff --git a/fireredtts3/redae/redae.py b/fireredtts3/redae/redae.py index 30de619..b91992d 100644 --- a/fireredtts3/redae/redae.py +++ b/fireredtts3/redae/redae.py @@ -3,6 +3,7 @@ import torchaudio import torch.nn as nn import torch.nn.functional as F +from fireredtts3.utils.device import get_attn_implementation, autocast, fft_device from transformers import ( Qwen3Config, Qwen3Model, PretrainedConfig, PreTrainedModel @@ -34,7 +35,7 @@ def __init__( max_position_embeddings=max_position_embeddings, num_attention_heads=num_attention_heads, num_key_value_heads=num_key_value_heads, - attn_implementation='flash_attention_2', + attn_implementation=get_attn_implementation(), ) self.qwen3 = Qwen3Model(self.qwen3_config) self.cls_tok = torch.nn.Parameter(torch.ones(1, 1, hidden_size)) @@ -119,7 +120,7 @@ def __init__( num_key_value_heads=num_key_value_heads, sliding_window=sliding_window, use_sliding_window=use_sliding_window, - attn_implementation='flash_attention_2', + attn_implementation=get_attn_implementation(), ) self.qwen3 = Qwen3Model(self.qwen3_config) if self.extra_downsample_rate > 1: @@ -197,9 +198,16 @@ def __init__( self.register_buffer("window", window) def forward(self, spec: torch.Tensor) -> torch.Tensor: + # Inverse FFT / complex math is unavailable on some backends (e.g. MPS on + # older torch); compute on an FFT-capable device and move the result back. + out_device = spec.device + spec = fft_device(spec) + window = self.window.to(spec.device) + if self.padding == "center": # Fallback to pytorch native implementation - return torch.istft(spec, self.n_fft, self.hop_length, self.win_length, self.window, center=True) + audio = torch.istft(spec, self.n_fft, self.hop_length, self.win_length, window, center=True) + return audio.to(out_device) elif self.padding == "same": pad = (self.win_length - self.hop_length) // 2 else: @@ -210,7 +218,7 @@ def forward(self, spec: torch.Tensor) -> torch.Tensor: # Inverse FFT ifft = torch.fft.irfft(spec, self.n_fft, dim=1, norm="backward") - ifft = ifft * self.window[None, :, None] + ifft = ifft * window[None, :, None] # Overlap and Add output_size = (T - 1) * self.hop_length + self.win_length @@ -219,7 +227,7 @@ def forward(self, spec: torch.Tensor) -> torch.Tensor: )[:, 0, 0, pad:-pad] # Window envelope - window_sq = self.window.square().expand(1, T, -1).transpose(1, 2) + window_sq = window.square().expand(1, T, -1).transpose(1, 2) window_envelope = torch.nn.functional.fold( window_sq, output_size=(1, output_size), kernel_size=(1, self.win_length), stride=(1, self.hop_length), ).squeeze()[pad:-pad] @@ -228,7 +236,7 @@ def forward(self, spec: torch.Tensor) -> torch.Tensor: assert (window_envelope > 1e-11).all() y = y / window_envelope - return y + return y.to(out_device) class ISTFTHead(nn.Module): @@ -246,7 +254,8 @@ def __init__( self.istft = ISTFT(n_fft=n_fft, hop_length=hop_length, win_length=n_fft, padding=padding) def forward(self, x: torch.Tensor) -> torch.Tensor: - x_pred = self.out(x) + # complex half is not usable -> run the spectrum/ISTFT in fp32 + x_pred = self.out(x).float() x_pred = x_pred.transpose(1, 2) mag, p = x_pred.chunk(2, dim=1) mag = torch.exp(mag) @@ -467,7 +476,7 @@ def pad_to_multiple_of(audio: torch.Tensor, multiple_of: int): audio = F.pad(audio, (pad_len, 0)) # NOTE left pad return audio - @torch.autocast(device_type='cuda', dtype=torch.bfloat16) + @autocast @torch.no_grad() def encode(self, audio: torch.Tensor, audio_sr:int): """ @@ -480,6 +489,7 @@ def encode(self, audio: torch.Tensor, audio_sr:int): audio = audio[:1] audio = torchaudio.functional.resample(audio, audio_sr, self.sample_rate) audio = self.pad_to_multiple_of(audio, self.downsample_rate) + audio = audio.to(self.dtype) latents = self.encoder.forward(audio) return latents diff --git a/fireredtts3/utils/device.py b/fireredtts3/utils/device.py new file mode 100644 index 0000000..430f7cd --- /dev/null +++ b/fireredtts3/utils/device.py @@ -0,0 +1,124 @@ +"""Device / dtype abstraction so the pipeline runs on CUDA, MPS or CPU. + +The upstream code hard-codes CUDA (``torch.device('cuda')``, ``autocast('cuda')``, +``flash_attention_2``). This module resolves everything once at import time so the +model code stays free of ``if cuda ... else ...`` branches. + +Environment overrides: + +- ``FIRERED_DEVICE`` — force a device, e.g. ``cuda`` / ``mps`` / ``cpu``. +- ``FIRERED_DTYPE`` — autocast dtype: ``bfloat16`` / ``float16`` / ``float32`` + (``float32`` disables autocast). Default: bfloat16 on CUDA, float32 elsewhere. +- ``FIRERED_WEIGHT_DTYPE`` — dtype the checkpoints are *loaded* in. Autocast only + casts activations, so this is the knob that actually halves resident memory. + Default: ``float32`` (the checkpoint dtype). +""" + +import os +import torch + + +_DTYPES = { + "bfloat16": torch.bfloat16, + "bf16": torch.bfloat16, + "float16": torch.float16, + "fp16": torch.float16, + "half": torch.float16, + "float32": torch.float32, + "fp32": torch.float32, + "none": torch.float32, +} + + +def _resolve_device() -> torch.device: + override = os.environ.get("FIRERED_DEVICE", "").strip() + if override: + return torch.device(override) + if torch.cuda.is_available(): + return torch.device("cuda") + if torch.backends.mps.is_available(): + return torch.device("mps") + return torch.device("cpu") + + +def _resolve_autocast_dtype(device_type: str) -> torch.dtype: + override = os.environ.get("FIRERED_DTYPE", "").strip().lower() + if override: + if override not in _DTYPES: + raise ValueError(f"invalid FIRERED_DTYPE={override!r}, expected one of {sorted(_DTYPES)}") + return _DTYPES[override] + # MPS/CPU: bf16 autocast is either unsupported or emulated (slow, and the + # flow-matching head is numerically touchy) -> stay in fp32 by default. + return torch.bfloat16 if device_type == "cuda" else torch.float32 + + +def _resolve_weight_dtype() -> torch.dtype: + override = os.environ.get("FIRERED_WEIGHT_DTYPE", "").strip().lower() + if not override: + return torch.float32 + if override not in _DTYPES: + raise ValueError(f"invalid FIRERED_WEIGHT_DTYPE={override!r}, expected one of {sorted(_DTYPES)}") + return _DTYPES[override] + + +def _probe_fft(device_type: str) -> bool: + """Whether complex tensors + inverse FFT work on this device (MPS: no).""" + if device_type == "cpu": + return True + try: + spec = torch.zeros(1, 3, 2, dtype=torch.complex64, device=device_type) + torch.fft.irfft(spec, 4, dim=1) + return True + except Exception: + return False + + +DEVICE: torch.device = _resolve_device() +DEVICE_TYPE: str = DEVICE.type +AUTOCAST_DTYPE: torch.dtype = _resolve_autocast_dtype(DEVICE_TYPE) +AUTOCAST_ENABLED: bool = AUTOCAST_DTYPE != torch.float32 +WEIGHT_DTYPE: torch.dtype = _resolve_weight_dtype() +# MPS has no FFT / complex kernels -> ISTFT and kaldi fbank must run on CPU. +FFT_ON_DEVICE: bool = _probe_fft(DEVICE_TYPE) + + +def get_device() -> torch.device: + return DEVICE + + +def get_weight_dtype() -> torch.dtype: + return WEIGHT_DTYPE + + +def get_attn_implementation() -> str: + """flash-attn wheels are CUDA-only; SDPA covers MPS/CPU.""" + return "flash_attention_2" if DEVICE_TYPE == "cuda" else "sdpa" + + +def autocast(func): + """Decorator replacing ``@torch.autocast(device_type='cuda', dtype=bfloat16)``.""" + if not AUTOCAST_ENABLED: + return func + return torch.autocast(device_type=DEVICE_TYPE, dtype=AUTOCAST_DTYPE)(func) + + +def disable_autocast(func): + """Decorator replacing ``@autocast('cuda', enabled=False)``.""" + if not AUTOCAST_ENABLED: + return func + return torch.autocast(device_type=DEVICE_TYPE, enabled=False)(func) + + +def fft_device(tensor: torch.Tensor) -> torch.Tensor: + """Move a tensor to a device that can actually run FFT / complex math.""" + if FFT_ON_DEVICE: + return tensor + return tensor.cpu() + + +def describe() -> str: + return ( + f"device={DEVICE}, weights={WEIGHT_DTYPE}, " + f"autocast={'off' if not AUTOCAST_ENABLED else AUTOCAST_DTYPE}, " + f"attn={get_attn_implementation()}, fft_on_device={FFT_ON_DEVICE}" + ) diff --git a/fireredtts3/utils/utils.py b/fireredtts3/utils/utils.py index 20a5133..4493da1 100644 --- a/fireredtts3/utils/utils.py +++ b/fireredtts3/utils/utils.py @@ -5,5 +5,8 @@ def fix_seed(seed: int): random.seed(seed) torch.manual_seed(seed) - torch.cuda.manual_seed(seed) - torch.cuda.manual_seed_all(seed) \ No newline at end of file + if torch.cuda.is_available(): + torch.cuda.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + if torch.backends.mps.is_available(): + torch.mps.manual_seed(seed) diff --git a/requirements-mps.txt b/requirements-mps.txt new file mode 100644 index 0000000..3ca1080 --- /dev/null +++ b/requirements-mps.txt @@ -0,0 +1,14 @@ +# macOS / Apple Silicon (MPS) dependency set. +# Same as requirements.txt minus flash_attn (CUDA-only wheel); attention falls +# back to PyTorch SDPA, see fireredtts3/utils/device.py. +torch==2.8.0 +torchaudio==2.8.0 +torchcodec==0.7.0 +transformers==5.6.2 +einops==0.8.2 +dotenv +regex +wetext +fasttext +faster-whisper +soundfile # torchaudio I/O backend on macOS (torchcodec needs ffmpeg<=7) diff --git a/scripts/mps_smoke_test.py b/scripts/mps_smoke_test.py new file mode 100644 index 0000000..d3b23be --- /dev/null +++ b/scripts/mps_smoke_test.py @@ -0,0 +1,178 @@ +"""Smoke test for FireRedTTS3 on Apple Silicon (MPS). + +Runs the Base zero-shot cloning path and, optionally, the Instruct paths +(cloning / voice design / semantic edit / acoustic edit), timing each call and +writing wavs to an output directory. + +Usage: + python scripts/mps_smoke_test.py \ + --refs /path/to/persons_refs \ + --out outputs/mps_test \ + --language Russian \ + --tasks base,instruct_tts,voice_design,semantic_edit,acoustic_edit +""" + +import argparse +import os +import time + +import torch +import torchaudio + +from fireredtts3.utils.device import describe, get_device + + +TEXTS = { + "Russian": "Он поднял голову и посмотрел на далёкие пики, скрытые в облаках. " + "До входа в долину оставалось три дня пути, и никто из них не знал, " + "чем закончится это путешествие.", + "English": "He raised his head and looked at the distant peaks hidden in the clouds. " + "Three days of travel remained before the valley.", + "Chinese": "他抬起头,望着远处隐没在云雾中的山峰,谁也不知道这次旅程会如何结束。", +} + + +def load_refs(refs_dir: str, max_seconds: float): + """Load (name, wav, sr, transcript) for every audio file with a .txt sibling.""" + items = [] + for fn in sorted(os.listdir(refs_dir)): + stem, ext = os.path.splitext(fn) + if ext.lower() not in (".mp3", ".wav", ".flac", ".m4a", ".ogg"): + continue + txt_path = os.path.join(refs_dir, stem + ".txt") + if not os.path.exists(txt_path): + print(f"[SKIP] {fn}: no transcript") + continue + with open(txt_path, encoding="utf-8") as f: + transcript = f.read().strip() + audio, sr = torchaudio.load(os.path.join(refs_dir, fn)) + audio = audio[:1] # mono + if max_seconds > 0: + audio = audio[:, : int(max_seconds * sr)] + items.append((stem, audio, sr, transcript)) + return items + + +def timed(label: str, fn): + t0 = time.perf_counter() + out = fn() + dt = time.perf_counter() - t0 + return out, dt, label + + +def report(label, dt, audio, sr): + dur = audio.shape[-1] / sr + rtf = dt / dur if dur > 0 else float("nan") + print(f"[OK] {label}: {dur:.2f}s audio in {dt:.1f}s wall (RTF {rtf:.2f})", flush=True) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--refs", required=True, help="directory with .mp3 + .txt") + ap.add_argument("--out", default="outputs/mps_test") + ap.add_argument("--models", default="pretrained_models") + ap.add_argument("--language", default="Russian") + ap.add_argument("--text", default=None) + ap.add_argument("--max-ref-seconds", type=float, default=10.0, + help="trim reference audio (0 = keep full length)") + ap.add_argument("--max-voices", type=int, default=0, help="0 = all voices") + ap.add_argument("--acoustic-instruction", default="adjust the speed to 0.8x", + help="templates: 'adjust the speed to X' / 'shift the pitch by N steps' / 'adjust the volume to X'") + ap.add_argument("--semantic-instruction", default="Replace 'Возможно' with 'Вероятно'.") + ap.add_argument("--tasks", default="base", + help="comma list: base,instruct_tts,voice_design,semantic_edit,acoustic_edit") + args = ap.parse_args() + + tasks = [t.strip() for t in args.tasks.split(",") if t.strip()] + os.makedirs(args.out, exist_ok=True) + print(f"[INFO] {describe()}") + print(f"[INFO] torch {torch.__version__}, tasks={tasks}") + + refs = load_refs(args.refs, args.max_ref_seconds) + if args.max_voices: + refs = refs[: args.max_voices] + print(f"[INFO] {len(refs)} reference voice(s): {[r[0] for r in refs]}") + text = args.text or TEXTS.get(args.language, TEXTS["English"]) + + # ---------------- Base: zero-shot cloning ---------------- + if "base" in tasks: + from fireredtts3.core import FireRedTTS3 + + (tts, load_dt, _) = timed("load base", lambda: FireRedTTS3( + args.models, use_fasttext=False, use_wetext=True, use_llm_tn=False)) + print(f"[OK] base model loaded in {load_dt:.1f}s", flush=True) + + for name, audio, sr, transcript in refs: + (out, dt, _) = timed(name, lambda: tts.generate( + language=args.language, + prompt_text=transcript, + prompt_audio=audio, + prompt_audio_sr=sr, + text=text, + )) + gen_audio, gen_sr = out + path = os.path.join(args.out, f"base_{name}.wav") + torchaudio.save(path, gen_audio.cpu(), gen_sr) + report(f"base clone {name} -> {path}", dt, gen_audio, gen_sr) + + del tts + if get_device().type == "mps": + torch.mps.empty_cache() + + # ---------------- Instruct ---------------- + instruct_tasks = [t for t in tasks if t != "base"] + if instruct_tasks: + from fireredtts3.core import FireRedTTS3Instruct + + (ins, load_dt, _) = timed("load instruct", lambda: FireRedTTS3Instruct( + args.models, use_fasttext=False, use_wetext=True, use_llm_tn=False)) + print(f"[OK] instruct model loaded in {load_dt:.1f}s", flush=True) + + if "instruct_tts" in instruct_tasks: + name, audio, sr, transcript = refs[0] + (out, dt, _) = timed("instruct_tts", lambda: ins.generate_tts( + prompt_text=transcript, prompt_audio=audio, prompt_audio_sr=sr, + text=text, language=args.language)) + gen_audio, gen_sr = out + path = os.path.join(args.out, f"instruct_clone_{name}.wav") + torchaudio.save(path, gen_audio.cpu(), gen_sr) + report(f"instruct clone {name} -> {path}", dt, gen_audio, gen_sr) + + if "voice_design" in instruct_tasks: + instruction = ("A calm middle-aged male narrator with a deep, warm voice, " + "speaking slowly and clearly, like an audiobook reader.") + (out, dt, _) = timed("voice_design", lambda: ins.generate_voice_design( + instruction=instruction, text=text, language=args.language)) + gen_audio, gen_sr, gen_text = out + path = os.path.join(args.out, "instruct_voice_design.wav") + torchaudio.save(path, gen_audio.cpu(), gen_sr) + report(f"voice design -> {path}", dt, gen_audio, gen_sr) + print(f" voice plan: {gen_text}") + + if "semantic_edit" in instruct_tasks or "acoustic_edit" in instruct_tasks: + # edits re-render the whole utterance -> use the shortest reference + name, audio, sr, transcript = min(refs, key=lambda r: r[1].shape[-1]) + print(f"[INFO] editing {name} ({audio.shape[-1] / sr:.1f}s): {transcript[:60]}...") + + if "acoustic_edit" in instruct_tasks: + (out, dt, _) = timed("acoustic_edit", lambda: ins.generate_acoustic_edit( + instruction=args.acoustic_instruction, audio_in=audio, audio_in_sr=sr)) + gen_audio, gen_sr = out[0], out[1] + path = os.path.join(args.out, f"acoustic_edit_{name}.wav") + torchaudio.save(path, gen_audio.cpu(), gen_sr) + report(f"acoustic edit -> {path}", dt, gen_audio, gen_sr) + + if "semantic_edit" in instruct_tasks: + (out, dt, _) = timed("semantic_edit", lambda: ins.generate_semantic_edit( + instruction=args.semantic_instruction, audio_in=audio, audio_in_sr=sr)) + gen_audio, gen_sr, gen_text = out + path = os.path.join(args.out, f"semantic_edit_{name}.wav") + torchaudio.save(path, gen_audio.cpu(), gen_sr) + report(f"semantic edit -> {path}", dt, gen_audio, gen_sr) + print(f" edited text: {gen_text}") + + print("[DONE]") + + +if __name__ == "__main__": + main() diff --git a/scripts/spk_similarity.py b/scripts/spk_similarity.py new file mode 100644 index 0000000..effab2e --- /dev/null +++ b/scripts/spk_similarity.py @@ -0,0 +1,27 @@ +"""Speaker-similarity check: cosine sim between each generated clone and its +reference, using the CAM++ extractor that ships with the model.""" +import os +import sys +import torch +import torchaudio +import torch.nn.functional as F + +from fireredtts3.campp.campp import CamppEmbedding + +refs_dir, out_dir = sys.argv[1], sys.argv[2] +spk = CamppEmbedding(os.path.join('pretrained_models', 'campp', 'campplus_voxceleb.bin')) + +def emb(path): + a, sr = torchaudio.load(path) + return F.normalize(spk.forward(a[:1], sr), dim=-1) + +names = sorted(n[:-4] for n in os.listdir(refs_dir) if n.endswith('.mp3')) +ref = {n: emb(os.path.join(refs_dir, n + '.mp3')) for n in names} +gen = {n: emb(os.path.join(out_dir, f'base_{n}.wav')) for n in names + if os.path.exists(os.path.join(out_dir, f'base_{n}.wav'))} + +print(f'{"generated":<14}' + ''.join(f'{n:>12}' for n in names) + ' <- reference') +for g in gen: + row = ''.join(f'{float(gen[g] @ ref[n].T):>12.3f}' for n in names) + print(f'{g:<14}{row}') +print('\nmatched pairs (diagonal) should clearly dominate their row/column') From 482e45961c6aaf06ff26613bf3ba37d1ba14dedc Mon Sep 17 00:00:00 2001 From: "vladimir.talpa" Date: Wed, 26 Aug 2026 15:43:53 +0300 Subject: [PATCH 3/5] Correct the MPS dtype benchmark in the README The earlier table timed a single short sentence, where fixed per-call overhead dominates, and wrongly concluded that half weights are slower on MPS. Measured on a 3-sentence paragraph across two voices, 3 interleaved runs each, bfloat16 weights are consistently ~20% *faster* than float32 (RTF 0.73-0.76 vs 0.87-0.96) as well as half the size, at unchanged speaker similarity. Co-Authored-By: Claude Opus 5 --- README.md | 29 ++++++++++++++++------------- 1 file changed, 16 insertions(+), 13 deletions(-) diff --git a/README.md b/README.md index b989d0d..5af2976 100644 --- a/README.md +++ b/README.md @@ -101,19 +101,22 @@ The device and compute dtype are resolved once in `fireredtts3/utils/device.py` `FIRERED_DTYPE` only casts activations, so it does **not** shrink resident memory; `FIRERED_WEIGHT_DTYPE` is the knob that does. Measured on an M4 Max -(torch 2.8, macOS 26), Base zero-shot cloning of one short sentence: - -| Weights | Resident weights | Peak MPS | RTF | Speaker sim | -| --- | --- | --- | --- | --- | -| `float32` (default) | 11.44 GiB | 13.87 GiB | 1.15 | 0.895 | -| `bfloat16` | 5.73 GiB | 6.92 GiB | 1.47 | 0.913 | -| `float16` | 5.73 GiB | 6.92 GiB | 1.85 | 0.896 | - -Half weights halve memory at equal quality but are *slower* on MPS, and -autocast bf16 over fp32 weights buys nothing (RTF 0.88-0.93 vs 0.85-0.91) — -the DiT flow head, the real bottleneck, sits outside the autocast region. Hence -fp32 weights + no autocast is the default off CUDA. On longer multi-sentence -input the same fp32 path reaches RTF 0.85-0.91, i.e. faster than real time. +(torch 2.8, macOS 26), Base zero-shot cloning of a 3-sentence paragraph over +two reference voices, RTF as the range across 3 interleaved runs: + +| Weights | Resident weights | RTF | Speaker sim | +| --- | --- | --- | --- | +| `float32` (default) | 11.42 GiB | 0.87-0.96 | 0.947 / 0.935 | +| `bfloat16` | 5.71 GiB | **0.73-0.76** | 0.941 / 0.933 | + +**On MPS, `FIRERED_WEIGHT_DTYPE=bfloat16` is the setting to want:** half the +memory *and* ~20% faster, at effectively unchanged speaker similarity. The +default stays `float32` only because it matches the checkpoints and leaves CUDA +untouched. + +Autocast, by contrast, buys nothing on MPS (RTF 0.88-0.93 vs 0.85-0.91 over +fp32 weights) — the DiT flow head, the real bottleneck, sits outside the +autocast region — so autocast is off by default off CUDA. ### Model Download From 48f9049109e6d8089c29c6276e84338f885906c2 Mon Sep 17 00:00:00 2001 From: "vladimir.talpa" Date: Wed, 26 Aug 2026 15:57:27 +0300 Subject: [PATCH 4/5] Let FIRERED_WEIGHT_DTYPE default to the checkpoint's own dtype MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Defaulting to float32 overrode what transformers does by itself: with no dtype argument it loads weights in the dtype the checkpoint stores. For the official fp32 weights the result is identical, but the hard default would silently upcast any half-precision checkpoint back to fp32 — exactly defeating the purpose of converting one. Unset now means "as stored". Co-Authored-By: Claude Opus 5 --- README.md | 10 +++++----- fireredtts3/utils/device.py | 15 +++++++++------ 2 files changed, 14 insertions(+), 11 deletions(-) diff --git a/README.md b/README.md index 5af2976..8c8976b 100644 --- a/README.md +++ b/README.md @@ -97,7 +97,7 @@ The device and compute dtype are resolved once in `fireredtts3/utils/device.py` | --- | --- | --- | | `FIRERED_DEVICE` | `cuda` / `mps` / `cpu` | first available | | `FIRERED_DTYPE` | autocast dtype (activations) | `bfloat16` on CUDA, `float32` elsewhere | -| `FIRERED_WEIGHT_DTYPE` | dtype the checkpoints are loaded in | `float32` | +| `FIRERED_WEIGHT_DTYPE` | dtype the checkpoints are loaded in | as stored (fp32 officially) | `FIRERED_DTYPE` only casts activations, so it does **not** shrink resident memory; `FIRERED_WEIGHT_DTYPE` is the knob that does. Measured on an M4 Max @@ -106,13 +106,13 @@ two reference voices, RTF as the range across 3 interleaved runs: | Weights | Resident weights | RTF | Speaker sim | | --- | --- | --- | --- | -| `float32` (default) | 11.42 GiB | 0.87-0.96 | 0.947 / 0.935 | +| `float32` (official checkpoints) | 11.42 GiB | 0.87-0.96 | 0.947 / 0.935 | | `bfloat16` | 5.71 GiB | **0.73-0.76** | 0.941 / 0.933 | **On MPS, `FIRERED_WEIGHT_DTYPE=bfloat16` is the setting to want:** half the -memory *and* ~20% faster, at effectively unchanged speaker similarity. The -default stays `float32` only because it matches the checkpoints and leaves CUDA -untouched. +memory *and* ~20% faster, at effectively unchanged speaker similarity. Left +unset, the loader keeps whatever the checkpoint stores — fp32 for the official +weights — so nothing changes unless you ask for it. Autocast, by contrast, buys nothing on MPS (RTF 0.88-0.93 vs 0.85-0.91 over fp32 weights) — the DiT flow head, the real bottleneck, sits outside the diff --git a/fireredtts3/utils/device.py b/fireredtts3/utils/device.py index 430f7cd..b615c5c 100644 --- a/fireredtts3/utils/device.py +++ b/fireredtts3/utils/device.py @@ -11,7 +11,8 @@ (``float32`` disables autocast). Default: bfloat16 on CUDA, float32 elsewhere. - ``FIRERED_WEIGHT_DTYPE`` — dtype the checkpoints are *loaded* in. Autocast only casts activations, so this is the knob that actually halves resident memory. - Default: ``float32`` (the checkpoint dtype). + Unset means "whatever the checkpoint stores", which is what transformers does + on its own — fp32 for the official weights. """ import os @@ -52,10 +53,11 @@ def _resolve_autocast_dtype(device_type: str) -> torch.dtype: return torch.bfloat16 if device_type == "cuda" else torch.float32 -def _resolve_weight_dtype() -> torch.dtype: +def _resolve_weight_dtype(): + """Returns None when unset, so from_pretrained keeps the checkpoint's dtype.""" override = os.environ.get("FIRERED_WEIGHT_DTYPE", "").strip().lower() if not override: - return torch.float32 + return None if override not in _DTYPES: raise ValueError(f"invalid FIRERED_WEIGHT_DTYPE={override!r}, expected one of {sorted(_DTYPES)}") return _DTYPES[override] @@ -77,7 +79,7 @@ def _probe_fft(device_type: str) -> bool: DEVICE_TYPE: str = DEVICE.type AUTOCAST_DTYPE: torch.dtype = _resolve_autocast_dtype(DEVICE_TYPE) AUTOCAST_ENABLED: bool = AUTOCAST_DTYPE != torch.float32 -WEIGHT_DTYPE: torch.dtype = _resolve_weight_dtype() +WEIGHT_DTYPE = _resolve_weight_dtype() # None => as stored in the checkpoint # MPS has no FFT / complex kernels -> ISTFT and kaldi fbank must run on CPU. FFT_ON_DEVICE: bool = _probe_fft(DEVICE_TYPE) @@ -86,7 +88,8 @@ def get_device() -> torch.device: return DEVICE -def get_weight_dtype() -> torch.dtype: +def get_weight_dtype(): + """dtype to load checkpoints in; None means keep whatever they store.""" return WEIGHT_DTYPE @@ -118,7 +121,7 @@ def fft_device(tensor: torch.Tensor) -> torch.Tensor: def describe() -> str: return ( - f"device={DEVICE}, weights={WEIGHT_DTYPE}, " + f"device={DEVICE}, weights={WEIGHT_DTYPE or 'as-stored'}, " f"autocast={'off' if not AUTOCAST_ENABLED else AUTOCAST_DTYPE}, " f"attn={get_attn_implementation()}, fft_on_device={FFT_ON_DEVICE}" ) From 6a5049025792de94bfd96ab008ef48b260b6104c Mon Sep 17 00:00:00 2001 From: "vladimir.talpa" Date: Wed, 26 Aug 2026 16:35:19 +0300 Subject: [PATCH 5/5] Add scripts/convert_to_bf16.py and document both bf16 routes Half-precision weights can be had two ways, and the README now spells out both: FIRERED_WEIGHT_DTYPE=bfloat16 casts the official fp32 checkpoints while loading and leaves the disk alone, or convert_to_bf16.py halves them on disk once (19 GB -> 9.7 GB) after which they load as bfloat16 with no variable set. The script verifies every tensor of the new file against the fp32 original for exact bfloat16 rounding before replacing it, skips components that are already converted, and takes --keep-fp32 for machines with room for both. Co-Authored-By: Claude Opus 5 --- README.md | 37 ++++++++++-- scripts/convert_to_bf16.py | 120 +++++++++++++++++++++++++++++++++++++ 2 files changed, 153 insertions(+), 4 deletions(-) create mode 100644 scripts/convert_to_bf16.py diff --git a/README.md b/README.md index 8c8976b..a4d91b5 100644 --- a/README.md +++ b/README.md @@ -109,10 +109,39 @@ two reference voices, RTF as the range across 3 interleaved runs: | `float32` (official checkpoints) | 11.42 GiB | 0.87-0.96 | 0.947 / 0.935 | | `bfloat16` | 5.71 GiB | **0.73-0.76** | 0.941 / 0.933 | -**On MPS, `FIRERED_WEIGHT_DTYPE=bfloat16` is the setting to want:** half the -memory *and* ~20% faster, at effectively unchanged speaker similarity. Left -unset, the loader keeps whatever the checkpoint stores — fp32 for the official -weights — so nothing changes unless you ask for it. +**On MPS, bfloat16 weights are what you want:** half the memory *and* ~20% +faster, at effectively unchanged speaker similarity. There are two ways to get +them, and they are equivalent at inference time: + +**a) Cast on the fly** — keep the official fp32 checkpoints on disk and convert +them while loading: + +```sh +FIRERED_WEIGHT_DTYPE=bfloat16 python your_script.py +``` + +Nothing on disk changes (19 GB stays 19 GB), and you can go back to fp32 by +dropping the variable. The cast is redone on every start. + +**b) Convert the checkpoints once** — halve them on disk too, then feed those +weights with no variable set at all: + +```sh +python scripts/convert_to_bf16.py # in place: 19 GB -> 9.7 GB +python your_script.py # loads bfloat16 by itself +``` + +Each tensor of the new file is verified against the fp32 original for exact +bfloat16 rounding before the original is replaced, and the script skips +components that are already converted. Loading gets faster too (3.3 s vs +5.7-9.2 s here), since half as much is read from disk. Pass `--keep-fp32` to +write `.bf16` alongside instead of replacing, and note the cast is one-way +— recovering fp32 means re-downloading. + +Option (b) is also the way to move a converted model to another machine: copy +`pretrained_models/` over, or run the script there after downloading. Either +way, leaving `FIRERED_WEIGHT_DTYPE` unset means the loader keeps whatever the +checkpoint stores, so nothing changes unless you ask for it. Autocast, by contrast, buys nothing on MPS (RTF 0.88-0.93 vs 0.85-0.91 over fp32 weights) — the DiT flow head, the real bottleneck, sits outside the diff --git a/scripts/convert_to_bf16.py b/scripts/convert_to_bf16.py new file mode 100644 index 0000000..8f917c9 --- /dev/null +++ b/scripts/convert_to_bf16.py @@ -0,0 +1,120 @@ +"""Convert the FireRedTTS3 checkpoints in place from fp32 to bfloat16. + +Halves them on disk (19 GB -> 9.7 GB) and speeds up loading, since the loader +then picks up bfloat16 from each config.json on its own. Measured speaker +similarity is unchanged; see the MPS section of the README. + +Every tensor of the new file is checked against the fp32 original for exact +bfloat16 rounding BEFORE the original is replaced. The cast is lossy and +one-way: to get fp32 back, re-download the checkpoints. + +Usage: + python scripts/convert_to_bf16.py # all components + python scripts/convert_to_bf16.py redae # just one + python scripts/convert_to_bf16.py --keep-fp32 # write .bf16, keep originals +""" + +import argparse +import json +import os +import shutil +import sys + +import torch +from safetensors import safe_open + +from fireredtts3.llm.fireredtts3_base import FireRedTTS3BaseCore +from fireredtts3.llm.fireredtts3_instruct import FireRedTTS3InstructCore +from fireredtts3.redae.redae import RedAE + +# smallest first, so disk is freed progressively on a nearly-full volume +COMPONENTS = { + "redae": RedAE, + "fireredtts3_base": FireRedTTS3BaseCore, + "fireredtts3_instruct": FireRedTTS3InstructCore, +} + + +def dir_size(path: str) -> int: + return sum( + os.path.getsize(os.path.join(dp, f)) + for dp, _, fs in os.walk(path) + for f in fs + ) + + +def verify(old_dir: str, new_dir: str) -> int: + """Assert every tensor equals the fp32 original rounded to bf16.""" + old_f = os.path.join(old_dir, "model.safetensors") + new_f = os.path.join(new_dir, "model.safetensors") + with safe_open(old_f, framework="pt") as a, safe_open(new_f, framework="pt") as b: + ka, kb = set(a.keys()), set(b.keys()) + if ka != kb: + raise AssertionError( + f"key mismatch: missing {sorted(ka - kb)[:5]}, extra {sorted(kb - ka)[:5]}" + ) + for k in ka: + t_old, t_new = a.get_tensor(k), b.get_tensor(k) + if t_old.shape != t_new.shape: + raise AssertionError(f"{k}: shape {t_old.shape} != {t_new.shape}") + if t_new.dtype != torch.bfloat16: + raise AssertionError(f"{k}: dtype is {t_new.dtype}, expected bfloat16") + if not torch.equal(t_old.to(torch.bfloat16), t_new): + raise AssertionError(f"{k}: does not match exact bf16 rounding") + return len(ka) + + +def convert(root: str, comp: str, keep_fp32: bool) -> None: + src = os.path.join(root, comp) + if not os.path.isdir(src): + print(f"[{comp}] SKIP: {src} not found") + return + + cfg_path = os.path.join(src, "config.json") + if json.load(open(cfg_path)).get("dtype") == "bfloat16": + print(f"[{comp}] SKIP: already bfloat16") + return + + tmp = src + ".bf16" + if os.path.exists(tmp): + shutil.rmtree(tmp) + + before = dir_size(src) + print(f"[{comp}] fp32 on disk: {before / 1024**3:.2f} GiB", flush=True) + + model = COMPONENTS[comp].from_pretrained(src, dtype=torch.bfloat16) + model.save_pretrained(tmp) + del model + after = dir_size(tmp) + print(f"[{comp}] bf16 written: {after / 1024**3:.2f} GiB " + f"({before / after:.2f}x smaller)", flush=True) + + n = verify(src, tmp) + print(f"[{comp}] verified {n} tensors: exact bf16 rounding") + + if keep_fp32: + print(f"[{comp}] kept fp32; bf16 left in {tmp}") + return + shutil.rmtree(src) + os.rename(tmp, src) + print(f"[{comp}] replaced in place") + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("components", nargs="*", choices=list(COMPONENTS), default=None, + help="components to convert (default: all)") + ap.add_argument("--models", default="pretrained_models") + ap.add_argument("--keep-fp32", action="store_true", + help="write .bf16 and leave the originals alone " + "(needs room for both)") + args = ap.parse_args() + + for comp in (args.components or list(COMPONENTS)): + convert(args.models, comp, args.keep_fp32) + print("[DONE]") + return 0 + + +if __name__ == "__main__": + sys.exit(main())