Cortazar-Dice / scripts /speak.py
macayaven's picture
Upload folder using huggingface_hub
335196b verified
Raw
History Blame Contribute Delete
17.4 kB
import argparse
import hashlib
import os
import random
import warnings
from pathlib import Path
import soundfile as sf
import torch
from f5_tts.infer.utils_infer import infer_process, load_checkpoint, load_vocoder, preprocess_ref_audio_text
from f5_tts.model import CFM, DiT
from f5_tts.model.utils import get_tokenizer
_REF_TEXT_BY_AUDIO_PATH = None
# Suppress noisy torch custom_ops kernel override warning (DGX Spark base image).
warnings.filterwarnings(
"ignore",
message="Warning only once for all operators, other operators may also be overridden.*",
category=UserWarning,
)
def _merge_short_chunks(chunks: list[str], min_chars: int) -> list[str]:
if min_chars <= 0:
return [c for c in chunks if c]
merged: list[str] = []
for chunk in chunks:
chunk = chunk.strip()
if not chunk:
continue
if not merged:
merged.append(chunk)
continue
if len(chunk.encode("utf-8")) < min_chars:
merged[-1] = (merged[-1].rstrip() + " " + chunk.lstrip()).strip()
else:
merged.append(chunk)
if len(merged) > 1 and len(merged[-1].encode("utf-8")) < min_chars:
merged[-2] = (merged[-2].rstrip() + " " + merged[-1].lstrip()).strip()
merged.pop()
return [c for c in merged if c]
def _lookup_ref_text(ref_audio_path: str) -> str | None:
"""
Best-effort lookup: if ref_audio matches a prepared dataset entry, use its text and avoid ASR.
This keeps inference deterministic and prevents wasting GPU/CPU on Whisper during prompt setup.
"""
dataset_dir = os.environ.get("REF_TEXT_DATASET_DIR", "/workspace/dataset_prepared/raw")
try:
from datasets import load_from_disk
except Exception:
return None
global _REF_TEXT_BY_AUDIO_PATH
if _REF_TEXT_BY_AUDIO_PATH is None:
try:
ds = load_from_disk(dataset_dir)
_REF_TEXT_BY_AUDIO_PATH = {row["audio_path"]: row["text"] for row in ds}
except Exception:
_REF_TEXT_BY_AUDIO_PATH = {}
# Try direct match first (fast path)
direct = _REF_TEXT_BY_AUDIO_PATH.get(ref_audio_path)
if direct:
return direct
# Normalise common path variants used by shell wrappers (relative paths run from /workspace)
try:
p = Path(ref_audio_path)
if not p.is_absolute():
maybe = str((Path("/workspace") / p).resolve())
direct = _REF_TEXT_BY_AUDIO_PATH.get(maybe)
if direct:
return direct
else:
maybe = str(p.resolve())
direct = _REF_TEXT_BY_AUDIO_PATH.get(maybe)
if direct:
return direct
except Exception:
# Path normalization is best-effort; if it fails, fall back to ASR or provided ref_text.
pass
return None
def main():
parser = argparse.ArgumentParser(description="Generate speech using fine-tuned F5-TTS")
parser.add_argument("--project_name", required=True, help="Name of the project")
parser.add_argument(
"--exp_name",
default=os.environ.get("EXP_NAME", "F5TTS_v1_Base"),
help="Experiment name (F5TTS_v1_Base, F5TTS_Base, E2TTS_Base)",
)
parser.add_argument("--text", required=True, help="Text to speak or path to a text file")
parser.add_argument("--ref_audio", help="Reference audio file (optional, random from dataset if not provided)")
parser.add_argument("--ref_text", default="", help="Reference text (optional, auto-transcribed if not provided)")
parser.add_argument("--ref_text_file", default="", help="Path to reference text file (optional)")
parser.add_argument("--output", default="generated.wav", help="Output filename")
parser.add_argument("--checkpoint", help="Specific checkpoint path (optional, uses latest if not provided)")
parser.add_argument("--speed", type=float, default=1.0, help="Speech speed (default: 1.0, <1.0 for slower)")
parser.add_argument("--nfe_step", type=int, default=None, help="Number of diffusion/flow steps (quality vs speed)")
parser.add_argument("--cfg_strength", type=float, default=None, help="Classifier-free guidance strength")
parser.add_argument("--sway_sampling_coef", type=float, default=None, help="Sway sampling coefficient")
parser.add_argument("--device", default="auto", choices=["auto", "cuda", "cpu"], help="Device for inference")
parser.add_argument("--tokenizer", default="pinyin", choices=["pinyin", "char", "custom"], help="Tokenizer type")
parser.add_argument("--tokenizer_path", help="Path to custom tokenizer vocab file")
parser.add_argument("--seed", type=int, default=None, help="Random seed for deterministic inference")
parser.add_argument(
"--max_chars",
type=int,
default=None,
help="Override max chars per chunk (higher => fewer splits). Ignored if not set.",
)
parser.add_argument(
"--min_chars",
type=int,
default=None,
help="Merge short chunks below this byte length (default: env F5_TTS_MIN_CHARS or 0).",
)
parser.add_argument(
"--no_chunk",
action="store_true",
help="Do not split the text; attempt a single long chunk (may be slower or fail on very long text).",
)
parser.add_argument(
"--save_used_text",
action="store_true",
help=(
"Save the exact target text used for generation next to the output wav "
"(<output>.used_text.txt). Disabled by default; enable for debugging."
),
)
args = parser.parse_args()
# ... (rest of imports/args)
project_name = args.project_name
if args.device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
else:
if args.device == "cuda" and not torch.cuda.is_available():
print("[WARN] CUDA requested but not available. Falling back to CPU.")
device = "cpu"
else:
device = args.device
if args.seed is not None:
random.seed(args.seed)
torch.manual_seed(args.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(args.seed)
# Ensure utils_infer uses the same device for ASR and helpers
import f5_tts.infer.utils_infer as infer_utils
infer_utils.device = device
# Paths
f5_tts_dir = Path(__file__).parent.parent / "F5-TTS"
ckpt_dir = f5_tts_dir / "ckpts" / project_name
if args.checkpoint:
ckpt_path = args.checkpoint
else:
# Find latest checkpoint
# Priority: model_last.pt > model_XXXX.pt (highest)
last_ckpt = ckpt_dir / "model_last.pt"
if last_ckpt.exists():
ckpt_path = str(last_ckpt)
else:
ckpts = list(ckpt_dir.glob("model_*.pt"))
if not ckpts:
print(f"No checkpoints found in {ckpt_dir}")
return
# sort by step count assuming model_1234.pt format
ckpts.sort(key=lambda x: int(x.stem.split("_")[1]) if x.stem.split("_")[1].isdigit() else 0)
ckpt_path = str(ckpts[-1])
print(f"Loading checkpoint: {ckpt_path}")
# Model Config
if args.exp_name == "F5TTS_Base":
model_cfg = dict(
dim=1024,
depth=22,
heads=16,
ff_mult=2,
text_dim=512,
text_mask_padding=False,
conv_layers=4,
pe_attn_head=1,
)
elif args.exp_name == "E2TTS_Base":
model_cfg = dict(
dim=1024,
depth=24,
heads=16,
ff_mult=4,
text_dim=512,
text_mask_padding=False,
pe_attn_head=1,
)
else:
model_cfg = dict(
dim=1024,
depth=22,
heads=16,
ff_mult=2,
text_dim=512,
conv_layers=4,
)
mel_spec_kwargs = dict(
n_fft=1024,
hop_length=256,
win_length=1024,
n_mel_channels=100,
target_sample_rate=24000,
mel_spec_type="vocos",
)
# Tokenizer
if args.tokenizer == "custom":
if not args.tokenizer_path:
raise ValueError("Custom tokenizer selected but no path provided")
tokenizer_path = args.tokenizer_path
else:
tokenizer_path = project_name
vocab_char_map, vocab_size = get_tokenizer(tokenizer_path, args.tokenizer)
# Load Model
model = CFM(
transformer=DiT(**model_cfg, text_num_embeds=vocab_size, mel_dim=100),
mel_spec_kwargs=mel_spec_kwargs,
vocab_char_map=vocab_char_map,
).to(device)
load_checkpoint(model, ckpt_path, device=device)
vocoder = load_vocoder().to(device)
# Reference Audio
if args.ref_audio:
ref_audio = args.ref_audio
else:
# Prioritize using processed chunks from temp dir if available
# These are better candidates (short, clean speech) than the raw input file
temp_wavs_dir = f5_tts_dir.parent / "temp" / project_name / "raw" / "wavs"
candidates = []
if temp_wavs_dir.exists():
candidates = list(temp_wavs_dir.glob("*.wav"))
if candidates:
print(f"Index: Found {len(candidates)} processed clips in {temp_wavs_dir}")
if not candidates:
# Fallback to raw_audio
data_dir = Path("raw_audio")
if not data_dir.exists():
data_dir = Path("/home/carlos/workspace/voices/cortazar/raw_audio")
if data_dir.exists():
candidates = list(data_dir.glob("*.wav")) + list(data_dir.glob("*.mp3"))
if candidates:
ref_audio = str(random.choice(candidates))
print(f"Using random reference audio: {ref_audio}")
else:
print("No reference audio found. Please provide --ref_audio")
return
# Preprocess Reference Audio/Text
# This handles silence removal and ASR if text is missing
print("Preprocessing reference audio/text...")
ref_text_input = args.ref_text
if not ref_text_input.strip() and args.ref_text_file:
try:
ref_text_input = Path(args.ref_text_file).read_text(encoding="utf-8").rstrip("\r\n")
except Exception as exc:
print(f"[WARN] Failed to read ref_text_file={args.ref_text_file}: {exc}")
if not ref_text_input.strip():
looked_up = _lookup_ref_text(ref_audio)
if looked_up:
print("[INFO] Found reference text in prepared dataset (skipping ASR).")
ref_text_input = looked_up
ref_audio, ref_text = preprocess_ref_audio_text(ref_audio, ref_text_input)
print(f"Reference Text: {ref_text}")
# Determine text to speak
text_to_speak = args.text
try:
# Check if argument is a file path (heuristic: length check to avoid OSError on long text)
path_obj = Path(args.text)
if len(args.text) < 255 and path_obj.exists() and path_obj.is_file():
print(f"Reading text from file: {args.text}")
text_to_speak = path_obj.read_text(encoding="utf-8").rstrip("\r\n")
except Exception:
# If any path error occurs (e.g. name too long), assume it's raw text
pass
# Optional debugging: persist the exact text used for generation next to the wav.
save_used_text_env = os.environ.get("F5_TTS_SAVE_USED_TEXT", "").strip().lower() in {"1", "true", "yes", "y"}
if args.save_used_text or save_used_text_env:
try:
out_path = Path(args.output)
text_hash = hashlib.sha256(text_to_speak.encode("utf-8")).hexdigest()[:16]
used_text_path = out_path.with_suffix(out_path.suffix + ".used_text.txt")
used_text_path.write_text(text_to_speak, encoding="utf-8")
print(f"[INFO] Saved used text to: {used_text_path} (sha256[:16]={text_hash})")
except Exception as exc:
print(f"[WARN] Failed to save used text: {exc}")
# Inference
print("[INFO] Generating ONLY the provided text; reference text is used only to condition voice/style.")
print(f"Generating: '{text_to_speak}'")
infer_kwargs = {"mel_spec_type": "vocos", "speed": args.speed, "device": device}
if args.nfe_step is not None:
infer_kwargs["nfe_step"] = int(args.nfe_step)
if args.cfg_strength is not None:
infer_kwargs["cfg_strength"] = float(args.cfg_strength)
if args.sway_sampling_coef is not None:
infer_kwargs["sway_sampling_coef"] = float(args.sway_sampling_coef)
auto_custom = len(text_to_speak.encode("utf-8")) > 320
use_custom_chunking = args.no_chunk or args.max_chars is not None or args.min_chars is not None or auto_custom
if use_custom_chunking:
if args.max_chars is not None and args.max_chars <= 0:
raise ValueError("--max_chars must be a positive integer")
if args.min_chars is not None and args.min_chars < 0:
raise ValueError("--min_chars must be a non-negative integer")
# Custom chunking path so we can control split size (or disable splitting entirely).
audio_np, sr = sf.read(ref_audio)
if audio_np.ndim == 1:
audio = torch.from_numpy(audio_np).float().unsqueeze(0)
else:
audio = torch.from_numpy(audio_np).float().t()
min_chars_env = os.environ.get("F5_TTS_MIN_CHARS", "")
min_chars = args.min_chars if args.min_chars is not None else int(min_chars_env or 0)
if args.no_chunk:
gen_text_batches = [text_to_speak]
max_chars = None
else:
from f5_tts.infer.utils_infer import chunk_text
if args.max_chars is not None:
max_chars = int(args.max_chars)
else:
# Heuristic: keep chunks large enough to avoid tiny fragments that cause audible seams.
duration = float(audio.shape[-1] / sr) if sr else 0.0
ref_len = len(ref_text.encode("utf-8"))
if duration <= 0.0 or ref_len <= 0:
max_chars = 240
else:
raw = (ref_len / duration) * max(6.0, 22.0 - duration) * float(args.speed)
max_chars = int(max(140, min(raw, 420)))
gen_text_batches = chunk_text(text_to_speak, max_chars=max_chars)
if min_chars > 0:
gen_text_batches = _merge_short_chunks(gen_text_batches, min_chars)
if not gen_text_batches:
gen_text_batches = [text_to_speak]
if max_chars is None:
print(f"[INFO] Custom chunking enabled: {len(gen_text_batches)} batch(es) (no_chunk)")
else:
print(
f"[INFO] Custom chunking enabled: {len(gen_text_batches)} batch(es) (max_chars={max_chars}, min_chars={min_chars})"
)
if os.environ.get("F5_TTS_LOG_CHUNKS", "").strip().lower() in {"1", "true", "yes", "y"}:
for idx, chunk in enumerate(gen_text_batches):
print(f"[CHUNK {idx:02d}] {chunk}")
if os.environ.get("F5_TTS_SAVE_CHUNKS_TEXT", "").strip().lower() in {"1", "true", "yes", "y"}:
try:
out_path = Path(args.output)
chunks_path = out_path.with_suffix(out_path.suffix + ".chunks.txt")
chunks_path.write_text(
"\n\n".join([f"[{i:02d}] {c}" for i, c in enumerate(gen_text_batches)]),
encoding="utf-8",
)
print(f"[INFO] Saved chunk text to: {chunks_path}")
except Exception as exc:
print(f"[WARN] Failed to save chunk text: {exc}")
nfe_step = int(args.nfe_step) if args.nfe_step is not None else infer_utils.nfe_step
cfg_strength = float(args.cfg_strength) if args.cfg_strength is not None else infer_utils.cfg_strength
sway_sampling_coef = (
float(args.sway_sampling_coef) if args.sway_sampling_coef is not None else infer_utils.sway_sampling_coef
)
try:
audio, sr, spectrogram = next(
infer_utils.infer_batch_process(
(audio, sr),
ref_text,
gen_text_batches,
model,
vocoder,
mel_spec_type="vocos",
progress=None,
target_rms=infer_utils.target_rms,
cross_fade_duration=infer_utils.cross_fade_duration,
nfe_step=nfe_step,
cfg_strength=cfg_strength,
sway_sampling_coef=sway_sampling_coef,
speed=args.speed,
fix_duration=infer_utils.fix_duration,
device=device,
)
)
except StopIteration as exc:
raise RuntimeError(
"infer_batch_process produced no output. This can occur if the input text is empty "
"or if chunking produced no valid batches."
) from exc
else:
audio, sr, spectrogram = infer_process(ref_audio, ref_text, text_to_speak, model, vocoder, **infer_kwargs)
sf.write(args.output, audio, sr)
print(f"Saved to {args.output}")
def create_project_data_name(project_name):
return f"{project_name}_pinyin"
if __name__ == "__main__":
main()