Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled

Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
freedakgmail
2026-07-09 22:38:58 +08:00
commit 6116b1f3c6
3683 changed files with 990984 additions and 0 deletions
View File
+298
View File
@@ -0,0 +1,298 @@
"""FunASR Server — unified vLLM-based inference service.
Provides OpenAI-compatible API (/v1/audio/transcriptions) and REST API (/asr).
Uses vLLM for Fun-ASR-Nano (GPU) or falls back to AutoModel for non-LLM models (SenseVoice/Paraformer).
"""
import io
import os
import re
import time
import logging
import tempfile
from typing import Optional
import numpy as np
import soundfile as sf
try:
from fastapi import FastAPI, UploadFile, File, Form, HTTPException
from fastapi.responses import JSONResponse
except ImportError:
raise ImportError(
"funasr-server requires additional packages. Install with: pip install vllm fastapi uvicorn python-multipart"
)
logger = logging.getLogger("funasr.server")
def prepare_audio_for_inference(audio_data, sr, target_sr=16000):
"""Return mono float32 audio at target_sr for ASR inference."""
audio_data = np.asarray(audio_data)
if audio_data.ndim > 1:
channel_axis = -1 if audio_data.shape[-1] <= audio_data.shape[0] else 0
audio_data = audio_data.mean(axis=channel_axis)
if sr != target_sr:
import librosa
audio_data = librosa.resample(audio_data, orig_sr=sr, target_sr=target_sr)
sr = target_sr
return audio_data.astype(np.float32), sr
def create_app(device: str = "cuda", preload_model: str = "auto") -> FastAPI:
if preload_model == "auto":
preload_model = "fun-asr-nano" if device.startswith("cuda") else "sensevoice"
app = FastAPI(title="FunASR Server", version="1.3.6")
app.state.device = device
app.state.engine = None
app.state.vad_model = None
app.state.fallback_models = {}
# Non-LLM model configs (use AutoModel, no vLLM)
FALLBACK_CONFIGS = {
"sensevoice": {
"model": "iic/SenseVoiceSmall",
"vad_model": "fsmn-vad",
"vad_kwargs": {"max_single_segment_time": 30000},
},
"paraformer": {
"model": "paraformer-zh",
"vad_model": "fsmn-vad",
"punc_model": "ct-punc",
},
}
def _load_vllm_engine():
"""Load Fun-ASR-Nano vLLM engine. Falls back to AutoModel if vLLM unavailable."""
if app.state.engine is not None:
return
try:
from funasr.models.fun_asr_nano.inference_vllm import FunASRNanoVLLM
from funasr import AutoModel as _AutoModel
logger.info("Loading Fun-ASR-Nano vLLM engine...")
t0 = time.time()
app.state.engine = FunASRNanoVLLM.from_pretrained(
model="FunAudioLLM/Fun-ASR-Nano-2512",
hub="hf",
device=device,
dtype="bf16",
max_model_len=4096,
gpu_memory_utilization=0.5,
)
logger.info(f"vLLM engine ready in {time.time()-t0:.1f}s")
app.state.use_vllm = True
logger.info("Loading VAD model...")
app.state.vad_model = _AutoModel(model="fsmn-vad", device=device, disable_update=True)
logger.info("VAD ready.")
except Exception as e:
logger.warning(f"vLLM failed ({e}), falling back to AutoModel for fun-asr-nano")
app.state.use_vllm = False
from funasr import AutoModel
cfg = {
"model": "FunAudioLLM/Fun-ASR-Nano-2512",
"hub": "hf",
"trust_remote_code": True,
"vad_model": "fsmn-vad",
"vad_kwargs": {"max_single_segment_time": 30000},
"device": device,
"disable_update": True,
}
app.state.fallback_models["fun-asr-nano"] = AutoModel(**cfg)
logger.info("Fallback AutoModel loaded for fun-asr-nano.")
def _load_fallback(name: str):
"""Load non-LLM model via AutoModel."""
if name in app.state.fallback_models:
return app.state.fallback_models[name]
if name not in FALLBACK_CONFIGS:
return None
from funasr import AutoModel
cfg = FALLBACK_CONFIGS[name].copy()
cfg["device"] = device
cfg["disable_update"] = True
logger.info(f"Loading fallback model '{name}'...")
model = AutoModel(**cfg)
app.state.fallback_models[name] = model
return model
def _process_vllm(audio_data, sr, language=None, hotwords=None, use_spk=False):
"""Process audio with vLLM engine (Fun-ASR-Nano)."""
audio_data, sr = prepare_audio_for_inference(audio_data, sr)
# VAD
vad_res = app.state.vad_model.generate(input=audio_data, fs=sr)
segments = vad_res[0]["value"] if vad_res and vad_res[0].get("value") else [[0, int(len(audio_data)*1000/sr)]]
seg_audios = []
seg_times = []
for seg in segments:
s0 = int(seg[0] * sr / 1000)
s1 = int(seg[1] * sr / 1000)
seg_audio = audio_data[s0:s1]
if len(seg_audio) > sr * 0.3:
seg_audios.append(seg_audio)
seg_times.append((seg[0], seg[1]))
if not seg_audios:
return {"text": "", "segments": [], "duration": len(audio_data)/sr}
# repetition_penalty is left at the neutral 1.0: the Fun-ASR-Nano vLLM
# engine runs in prompt-embeds mode, where any other value crashes the
# CUDA kernel (see issue #2948 and fun_asr_nano.vllm_utils).
gen_kwargs = {"max_new_tokens": 500, "repetition_penalty": 1.0}
if language:
gen_kwargs["language"] = language
if hotwords:
gen_kwargs["hotwords"] = hotwords
results = app.state.engine.generate(inputs=seg_audios, **gen_kwargs)
output_segments = []
full_text_parts = []
for r, (start_ms, end_ms) in zip(results, seg_times):
text = r["text"]
seg_info = {"text": text, "start": start_ms/1000, "end": end_ms/1000}
if "timestamps" in r:
offset = start_ms / 1000
seg_info["words"] = [
{"word": ts["token"], "start": ts["start_time"]+offset, "end": ts["end_time"]+offset}
for ts in r["timestamps"]
]
output_segments.append(seg_info)
full_text_parts.append(text)
return {
"text": "".join(full_text_parts),
"segments": output_segments,
"duration": len(audio_data) / sr,
}
def _process_fallback(model_name, audio_path, language=None):
"""Process with non-LLM model (SenseVoice/Paraformer)."""
model = _load_fallback(model_name)
kwargs = {"input": audio_path, "batch_size": 1}
if language:
kwargs["language"] = language
result = model.generate(**kwargs)
text = re.sub(r'<\|[^|]*\|>', '', result[0]["text"]).strip()
segments = []
if "sentence_info" in result[0]:
for s in result[0]["sentence_info"]:
segments.append({
"start": s.get("start", 0)/1000,
"end": s.get("end", 0)/1000,
"text": re.sub(r'<\|[^|]*\|>', '', s.get("text", "")).strip(),
"speaker": s.get("spk"),
})
return {"text": text, "segments": segments}
# Pre-load
if preload_model == "fun-asr-nano":
_load_vllm_engine()
else:
_load_fallback(preload_model)
@app.post("/v1/audio/transcriptions")
async def transcribe(
file: UploadFile = File(...),
model: str = Form(default="fun-asr-nano"),
language: Optional[str] = Form(default=None),
response_format: Optional[str] = Form(default="json"),
spk: bool = Form(default=False),
):
content = await file.read()
t0 = time.perf_counter()
if model == "fun-asr-nano":
_load_vllm_engine()
if app.state.use_vllm:
audio_data, sr = sf.read(io.BytesIO(content))
result = _process_vllm(audio_data, sr, language=language, use_spk=spk)
else:
suffix = os.path.splitext(file.filename)[1] if file.filename else ".wav"
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
tmp.write(content)
tmp_path = tmp.name
try:
result = _process_fallback("fun-asr-nano", tmp_path, language=language)
finally:
os.unlink(tmp_path)
elif model in FALLBACK_CONFIGS:
suffix = os.path.splitext(file.filename)[1] if file.filename else ".wav"
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
tmp.write(content)
tmp_path = tmp.name
try:
result = _process_fallback(model, tmp_path, language=language)
finally:
os.unlink(tmp_path)
else:
raise HTTPException(400, f"Unknown model '{model}'. Available: fun-asr-nano, {', '.join(FALLBACK_CONFIGS.keys())}")
t1 = time.perf_counter()
if response_format == "verbose_json":
return JSONResponse({
"task": "transcribe",
"language": language or "zh",
"duration": result.get("duration", 0),
"text": result["text"],
"segments": [
{"id": i, "start": s["start"], "end": s["end"], "text": s["text"], "words": s.get("words", [])}
for i, s in enumerate(result["segments"])
],
})
elif response_format == "text":
return JSONResponse(result["text"])
else:
return JSONResponse({"text": result["text"]})
@app.post("/asr")
async def asr_endpoint(
file: UploadFile = File(...),
language: Optional[str] = Form(default=None),
hotwords: str = Form(default=""),
spk: bool = Form(default=False),
):
"""Full-featured ASR endpoint with timestamps and speaker diarization."""
content = await file.read()
_load_vllm_engine()
hw_list = [w.strip() for w in hotwords.split(",") if w.strip()] if hotwords else None
t0 = time.perf_counter()
if app.state.use_vllm:
audio_data, sr = sf.read(io.BytesIO(content))
result = _process_vllm(audio_data, sr, language=language, hotwords=hw_list, use_spk=spk)
else:
suffix = ".wav"
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
tmp.write(content)
tmp_path = tmp.name
try:
result = _process_fallback("fun-asr-nano", tmp_path, language=language)
finally:
os.unlink(tmp_path)
t1 = time.perf_counter()
result["processing_time"] = round(t1 - t0, 3)
result["rtf"] = round((t1 - t0) / result["duration"], 4) if result.get("duration", 0) > 0 else 0
return JSONResponse(result)
@app.get("/v1/models")
async def list_models():
all_models = ["fun-asr-nano"] + list(FALLBACK_CONFIGS.keys())
return JSONResponse({"object": "list", "data": [{"id": n, "object": "model"} for n in all_models]})
@app.get("/health")
async def health():
loaded = []
if app.state.engine is not None:
loaded.append("fun-asr-nano (vLLM)")
loaded.extend(app.state.fallback_models.keys())
return {"status": "ok", "device": device, "models_loaded": loaded}
return app
+146
View File
@@ -0,0 +1,146 @@
import os
import json
import numpy as np
import torch
import hydra
import logging
from omegaconf import DictConfig, OmegaConf
from funasr.register import tables
from funasr.download.download_model_from_hub import download_model
from funasr.train_utils.set_all_random_seed import set_all_random_seed
@hydra.main(config_name=None, version_base=None)
def main_hydra(kwargs: DictConfig):
"""Main hydra.
Args:
kwargs: Additional keyword arguments.
"""
if kwargs.get("debug", False):
import pdb
pdb.set_trace()
assert "model" in kwargs
if "model_conf" not in kwargs:
logging.info("download models from model hub: {}".format(kwargs.get("hub", "ms")))
kwargs = download_model(is_training=kwargs.get("is_training", True), **kwargs)
main(**kwargs)
def main(**kwargs):
"""Main.
Args:
**kwargs: Additional keyword arguments.
"""
print(kwargs)
# set random seed
# tables.print()
set_all_random_seed(kwargs.get("seed", 0))
torch.backends.cudnn.enabled = kwargs.get("cudnn_enabled", torch.backends.cudnn.enabled)
torch.backends.cudnn.benchmark = kwargs.get("cudnn_benchmark", torch.backends.cudnn.benchmark)
torch.backends.cudnn.deterministic = kwargs.get("cudnn_deterministic", True)
tokenizer = kwargs.get("tokenizer", None)
# build frontend if frontend is none None
frontend = kwargs.get("frontend", None)
if frontend is not None:
frontend_class = tables.frontend_classes.get(frontend)
frontend = frontend_class(**kwargs["frontend_conf"])
kwargs["frontend"] = frontend
kwargs["input_size"] = frontend.output_size()
# dataset
dataset_class = tables.dataset_classes.get(kwargs.get("dataset", "AudioDataset"))
dataset_train = dataset_class(
kwargs.get("train_data_set_list"),
frontend=frontend,
tokenizer=None,
is_training=False,
**kwargs.get("dataset_conf"),
)
# dataloader
batch_sampler = kwargs["dataset_conf"].get("batch_sampler", "BatchSampler")
batch_sampler_class = tables.batch_sampler_classes.get(batch_sampler)
dataset_conf = kwargs.get("dataset_conf")
dataset_conf["batch_type"] = "example"
dataset_conf["batch_size"] = 1
dataset_conf["num_workers"] = os.cpu_count() or 32
batch_sampler_train = batch_sampler_class(dataset_train, is_training=False, **dataset_conf)
dataloader_train = torch.utils.data.DataLoader(
dataset_train, collate_fn=dataset_train.collator, **batch_sampler_train
)
total_frames = 0
for batch_idx, batch in enumerate(dataloader_train):
iter_stop = int(kwargs.get("scale", -1.0) * len(dataloader_train))
log_step = iter_stop // 100
if batch_idx % log_step == 0:
logging.info(f"prcessed: {batch_idx}/{iter_stop}")
if batch_idx >= iter_stop and iter_stop > 0.0:
logging.info(f"prcessed: {iter_stop}/{iter_stop}")
break
fbank = batch["speech"].numpy()[0, :, :]
if total_frames == 0:
mean_stats = np.sum(fbank, axis=0)
var_stats = np.sum(np.square(fbank), axis=0)
else:
mean_stats += np.sum(fbank, axis=0)
var_stats += np.sum(np.square(fbank), axis=0)
total_frames += fbank.shape[0]
cmvn_info = {
"mean_stats": mean_stats.tolist(),
"var_stats": var_stats.tolist(),
"total_frames": total_frames,
}
cmvn_file = kwargs.get("cmvn_file", "cmvn.json")
# import pdb;pdb.set_trace()
with open(cmvn_file, "w") as fout:
fout.write(json.dumps(cmvn_info))
mean = -1.0 * mean_stats / total_frames
var = 1.0 / np.sqrt(var_stats / total_frames - mean * mean)
dims = mean.shape[0]
am_mvn = os.path.dirname(cmvn_file) + "/am.mvn"
with open(am_mvn, "w") as fout:
fout.write(
"<Nnet>"
+ "\n"
+ "<Splice> "
+ str(dims)
+ " "
+ str(dims)
+ "\n"
+ "[ 0 ]"
+ "\n"
+ "<AddShift> "
+ str(dims)
+ " "
+ str(dims)
+ "\n"
)
fout.write("<LearnRateCoef> 0 [ " + " ".join([str(item) for item in mean]) + " ]\n")
fout.write("<Rescale> " + str(dims) + " " + str(dims) + "\n")
fout.write("<LearnRateCoef> 0 [ " + " ".join([str(item) for item in var]) + " ]\n")
fout.write("</Nnet>" + "\n")
"""
python funasr/bin/compute_audio_cmvn.py \
--config-path "/Users/zhifu/funasr1.0/examples/aishell/paraformer/conf" \
--config-name "train_asr_paraformer_conformer_12e_6d_2048_256.yaml" \
++train_data_set_list="/Users/zhifu/funasr1.0/data/list/audio_datasets.jsonl" \
++cmvn_file="/Users/zhifu/funasr1.0/data/list/cmvn.json" \
++dataset_conf.num_workers=0
"""
if __name__ == "__main__":
main_hydra()
+52
View File
@@ -0,0 +1,52 @@
import os
import hydra
import logging
from omegaconf import DictConfig, OmegaConf, ListConfig
from funasr.auto.auto_model import AutoModel
@hydra.main(config_name=None, version_base=None)
def main_hydra(cfg: DictConfig):
"""Main hydra.
Args:
cfg: Configuration overrides.
"""
def to_plain_list(cfg_item):
"""To plain list.
Args:
cfg_item: TODO.
"""
if isinstance(cfg_item, ListConfig):
return OmegaConf.to_container(cfg_item, resolve=True)
elif isinstance(cfg_item, DictConfig):
return {k: to_plain_list(v) for k, v in cfg_item.items()}
else:
return cfg_item
kwargs = to_plain_list(cfg)
if kwargs.get("debug", False):
import pdb
pdb.set_trace()
if "device" not in kwargs:
kwargs["device"] = "cpu"
model = AutoModel(**kwargs)
res = model.export(
input=kwargs.get("input", None),
type=kwargs.get("type", "onnx"),
quantize=kwargs.get("quantize", False),
fallback_num=kwargs.get("fallback-num", 5),
calib_num=kwargs.get("calib_num", 100),
opset_version=kwargs.get("opset_version", 14),
)
print(res)
if __name__ == "__main__":
main_hydra()
+40
View File
@@ -0,0 +1,40 @@
import hydra
import logging
from omegaconf import DictConfig, OmegaConf, ListConfig
from funasr.auto.auto_model import AutoModel
@hydra.main(config_name=None, version_base=None)
def main_hydra(cfg: DictConfig):
"""Main hydra.
Args:
cfg: Configuration overrides.
"""
def to_plain_list(cfg_item):
"""To plain list.
Args:
cfg_item: TODO.
"""
if isinstance(cfg_item, ListConfig):
return OmegaConf.to_container(cfg_item, resolve=True)
elif isinstance(cfg_item, DictConfig):
return {k: to_plain_list(v) for k, v in cfg_item.items()}
else:
return cfg_item
kwargs = to_plain_list(cfg)
if kwargs.get("debug", False):
import pdb
pdb.set_trace()
model = AutoModel(**kwargs)
res = model.generate(input=kwargs["input"])
print(res)
if __name__ == "__main__":
main_hydra()
+66
View File
@@ -0,0 +1,66 @@
"""
FunASR Server — OpenAI-compatible speech recognition API.
Usage:
funasr-server # default: sensevoice on cuda:0, port 8000
funasr-server --device cpu --port 9000
funasr-server --model paraformer
"""
import argparse
import sys
def main():
parser = argparse.ArgumentParser(
description="FunASR OpenAI-Compatible API Server",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
funasr-server # Start with SenseVoice on GPU
funasr-server --device cpu # Start on CPU
funasr-server --model paraformer # Use Paraformer model
funasr-server --port 9000 # Custom port
Then use with OpenAI SDK:
from openai import OpenAI
client = OpenAI(base_url="http://localhost:8000/v1", api_key="x")
result = client.audio.transcriptions.create(model="sensevoice", file=open("a.wav","rb"))
""",
)
parser.add_argument("--host", default="0.0.0.0", help="Bind address (default: 0.0.0.0)")
parser.add_argument("--port", type=int, default=8000, help="Port (default: 8000)")
parser.add_argument("--device", default="cuda", help="Device: cuda, cpu, mps (default: cuda)")
parser.add_argument("--model", default="auto", help="Pre-load model: auto (GPU=fun-asr-nano, CPU=sensevoice), sensevoice, paraformer, fun-asr-nano")
args = parser.parse_args()
try:
import uvicorn
import fastapi
except ImportError:
print("Error: funasr-server requires additional packages.")
print("Install with: pip install vllm fastapi uvicorn python-multipart")
sys.exit(1)
# Import and configure the app
import os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', '..', 'examples', 'openai_api'))
# Use inline app to avoid path issues
from funasr.bin._server_app import create_app
app = create_app(device=args.device, preload_model=args.model)
print(f"╔══════════════════════════════════════════════╗")
print(f"║ FunASR Server v1.3.6 ║")
print(f"║ Device: {args.device:<8}")
print(f"║ Model: {args.model:<12}")
print(f"║ URL: http://{args.host}:{args.port}/v1 ║")
print(f"║ Docs: http://{args.host}:{args.port}/docs ║")
print(f"╚══════════════════════════════════════════════╝")
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()
+307
View File
@@ -0,0 +1,307 @@
#!/usr/bin/env python3
import argparse
from collections import Counter
import logging
from pathlib import Path
import sys
from typing import List
from typing import Optional
try:
from funasr.utils.cli_utils import get_commandline_args
except ImportError:
def get_commandline_args():
return {}
from funasr.tokenizer.build_tokenizer import build_tokenizer
from funasr.tokenizer.cleaner import TextCleaner
from funasr.tokenizer.phoneme_tokenizer import g2p_classes
from funasr.utils.types import str2bool
from funasr.utils.types import str_or_none
def field2slice(field: Optional[str]) -> slice:
"""Convert field string to slice
Note that field string accepts 1-based integer.
Examples:
>>> field2slice("1-")
slice(0, None, None)
>>> field2slice("1-3")
slice(0, 3, None)
>>> field2slice("-3")
slice(None, 3, None)
"""
field = field.strip()
try:
if "-" in field:
# e.g. "2-" or "2-5" or "-7"
s1, s2 = field.split("-", maxsplit=1)
if s1.strip() == "":
s1 = None
else:
s1 = int(s1)
if s1 == 0:
raise ValueError("1-based string")
if s2.strip() == "":
s2 = None
else:
s2 = int(s2)
else:
# e.g. "2"
s1 = int(field)
s2 = s1 + 1
if s1 == 0:
raise ValueError("must be 1 or more value")
except ValueError:
raise RuntimeError(f"Format error: e.g. '2-', '2-5', or '-5': {field}")
if s1 is None:
slic = slice(None, s2)
else:
# -1 because of 1-based integer following "cut" command
# e.g "1-3" -> slice(0, 3)
slic = slice(s1 - 1, s2)
return slic
def tokenize(
input: str,
output: str,
field: Optional[str],
delimiter: Optional[str],
token_type: str,
space_symbol: str,
non_linguistic_symbols: Optional[str],
bpemodel: Optional[str],
log_level: str,
write_vocabulary: bool,
vocabulary_size: int,
remove_non_linguistic_symbols: bool,
cutoff: int,
add_symbol: List[str],
cleaner: Optional[str],
g2p: Optional[str],
):
"""Tokenize.
Args:
input: Input audio/text data.
output: TODO.
field: TODO.
delimiter: TODO.
token_type: TODO.
space_symbol: TODO.
non_linguistic_symbols: TODO.
bpemodel: TODO.
log_level: TODO.
write_vocabulary: TODO.
vocabulary_size: Size/dimension parameter.
remove_non_linguistic_symbols: TODO.
cutoff: TODO.
add_symbol: TODO.
cleaner: TODO.
g2p: TODO.
"""
logging.basicConfig(
level=log_level,
format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
)
if input == "-":
fin = sys.stdin
else:
fin = Path(input).open("r", encoding="utf-8")
if output == "-":
fout = sys.stdout
else:
p = Path(output)
p.parent.mkdir(parents=True, exist_ok=True)
fout = p.open("w", encoding="utf-8")
cleaner = TextCleaner(cleaner)
tokenizer = build_tokenizer(
token_type=token_type,
bpemodel=bpemodel,
delimiter=delimiter,
space_symbol=space_symbol,
non_linguistic_symbols=non_linguistic_symbols,
remove_non_linguistic_symbols=remove_non_linguistic_symbols,
g2p_type=g2p,
)
counter = Counter()
if field is not None:
field = field2slice(field)
for line in fin:
line = line.rstrip()
if field is not None:
# e.g. field="2-"
# uttidA hello world!! -> hello world!!
tokens = line.split(delimiter)
tokens = tokens[field]
if delimiter is None:
line = " ".join(tokens)
else:
line = delimiter.join(tokens)
line = cleaner(line)
tokens = tokenizer.text2tokens(line)
if not write_vocabulary:
fout.write(" ".join(tokens) + "\n")
else:
for t in tokens:
counter[t] += 1
if not write_vocabulary:
return
## FIXME
## del duplicate add_symbols in counter
for symbol_and_id in add_symbol:
# e.g symbol="<blank>:0"
try:
symbol, idx = symbol_and_id.split(":")
except ValueError:
raise RuntimeError(f"Format error: e.g. '<blank>:0': {symbol_and_id}")
symbol = symbol.strip()
if symbol in counter:
del counter[symbol]
# ======= write_vocabulary mode from here =======
# Sort by the number of occurrences in descending order
# and filter lower frequency words than cutoff value
words_and_counts = list(
filter(lambda x: x[1] > cutoff, sorted(counter.items(), key=lambda x: -x[1]))
)
# Restrict the vocabulary size
if vocabulary_size > 0:
if vocabulary_size < len(add_symbol):
raise RuntimeError(f"vocabulary_size is too small: {vocabulary_size}")
words_and_counts = words_and_counts[: vocabulary_size - len(add_symbol)]
# Parse the values of --add_symbol
for symbol_and_id in add_symbol:
# e.g symbol="<blank>:0"
try:
symbol, idx = symbol_and_id.split(":")
idx = int(idx)
except ValueError:
raise RuntimeError(f"Format error: e.g. '<blank>:0': {symbol_and_id}")
symbol = symbol.strip()
# e.g. idx=0 -> append as the first symbol
# e.g. idx=-1 -> append as the last symbol
if idx < 0:
idx = len(words_and_counts) + 1 + idx
words_and_counts.insert(idx, (symbol, None))
# Write words
for w, c in words_and_counts:
fout.write(w + "\n")
# Logging
total_count = sum(counter.values())
invocab_count = sum(c for w, c in words_and_counts if c is not None)
logging.info(f"OOV rate = {(total_count - invocab_count) / total_count * 100} %")
def get_parser() -> argparse.ArgumentParser:
"""Get parser."""
parser = argparse.ArgumentParser(
description="Tokenize texts",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument(
"--log_level",
type=lambda x: x.upper(),
default="INFO",
choices=("CRITICAL", "ERROR", "WARNING", "INFO", "DEBUG", "NOTSET"),
help="The verbose level of logging",
)
parser.add_argument("--input", "-i", required=True, help="Input text. - indicates sys.stdin")
parser.add_argument("--output", "-o", required=True, help="Output text. - indicates sys.stdout")
parser.add_argument(
"--field",
"-f",
help="The target columns of the input text as 1-based integer. e.g 2-",
)
parser.add_argument(
"--token_type",
"-t",
default="char",
choices=["char", "bpe", "word", "phn"],
help="Token type",
)
parser.add_argument("--delimiter", "-d", default=None, help="The delimiter")
parser.add_argument("--space_symbol", default="<space>", help="The space symbol")
parser.add_argument("--bpemodel", default=None, help="The bpemodel file path")
parser.add_argument(
"--non_linguistic_symbols",
type=str_or_none,
help="non_linguistic_symbols file path",
)
parser.add_argument(
"--remove_non_linguistic_symbols",
type=str2bool,
default=False,
help="Remove non-language-symbols from tokens",
)
parser.add_argument(
"--cleaner",
type=str_or_none,
choices=[None, "tacotron", "jaconv", "vietnamese", "korean_cleaner"],
default=None,
help="Apply text cleaning",
)
parser.add_argument(
"--g2p",
type=str_or_none,
choices=g2p_classes,
default=None,
help="Specify g2p method if --token_type=phn",
)
group = parser.add_argument_group("write_vocabulary mode related")
group.add_argument(
"--write_vocabulary",
type=str2bool,
default=False,
help="Write tokens list instead of tokenized text per line",
)
group.add_argument("--vocabulary_size", type=int, default=0, help="Vocabulary size")
group.add_argument(
"--cutoff",
default=0,
type=int,
help="cut-off frequency used for write-vocabulary mode",
)
group.add_argument(
"--add_symbol",
type=str,
default=[],
action="append",
help="Append symbol e.g. --add_symbol '<blank>:0' --add_symbol '<unk>:1'",
)
return parser
def main(cmd=None):
"""Main.
Args:
cmd: TODO.
"""
print(get_commandline_args(), file=sys.stderr)
parser = get_parser()
args = parser.parse_args(cmd)
kwargs = vars(args)
tokenize(**kwargs)
if __name__ == "__main__":
main()
+289
View File
@@ -0,0 +1,289 @@
#!/usr/bin/env python3
# -*- encoding: utf-8 -*-
import os
import functools
import sys
import torch
import torch.nn as nn
import hydra
import logging
import time
import argparse
from io import BytesIO
from contextlib import nullcontext
import torch.distributed as dist
from omegaconf import DictConfig, OmegaConf
from torch.cuda.amp import autocast, GradScaler
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.algorithms.join import Join
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
from tensorboardX import SummaryWriter
from funasr.train_utils.average_nbest_models import average_checkpoints
from funasr.register import tables
from funasr.optimizers import optim_classes
from funasr.train_utils.trainer import Trainer
from funasr.schedulers import scheduler_classes
from funasr.train_utils.initialize import initialize
from funasr.download.download_model_from_hub import download_model
from funasr.models.lora.utils import mark_only_lora_as_trainable
from funasr.train_utils.set_all_random_seed import set_all_random_seed
from funasr.train_utils.load_pretrained_model import load_pretrained_model
from funasr.utils.misc import prepare_model_dir
from funasr.train_utils.model_summary import model_summary
from funasr import AutoModel
@hydra.main(config_name=None, version_base=None)
def main_hydra(kwargs: DictConfig):
"""Main hydra.
Args:
kwargs: Additional keyword arguments.
"""
if kwargs.get("debug", False):
import pdb
pdb.set_trace()
assert "model" in kwargs
if "model_conf" not in kwargs:
logging.info("download models from model hub: {}".format(kwargs.get("hub", "ms")))
kwargs = download_model(is_training=kwargs.get("is_training", True), **kwargs)
main(**kwargs)
def main(**kwargs):
# set random seed
"""Main.
Args:
**kwargs: Additional keyword arguments.
"""
set_all_random_seed(kwargs.get("seed", 0))
torch.backends.cudnn.enabled = kwargs.get("cudnn_enabled", torch.backends.cudnn.enabled)
torch.backends.cudnn.benchmark = kwargs.get("cudnn_benchmark", torch.backends.cudnn.benchmark)
torch.backends.cudnn.deterministic = kwargs.get("cudnn_deterministic", True)
# open tf32
torch.backends.cuda.matmul.allow_tf32 = kwargs.get("enable_tf32", True)
local_rank = int(os.environ.get("LOCAL_RANK", 0))
if local_rank == 0:
tables.print()
# Check if we are using DDP or FSDP
use_ddp = "WORLD_SIZE" in os.environ and int(os.environ["WORLD_SIZE"]) > 1
use_fsdp = kwargs.get("use_fsdp", False)
# use_ddp = False if use_fsdp else use_fsdp
if use_ddp or use_fsdp:
dist.init_process_group(backend=kwargs.get("backend", "nccl"), init_method="env://")
torch.cuda.set_device(local_rank)
logging.info("Build model, frontend, tokenizer")
device = kwargs.get("device", "cuda")
kwargs["device"] = "cpu"
model = AutoModel(**kwargs)
# save config.yaml
if (
(use_ddp or use_fsdp)
and dist.get_rank() == 0
or not (use_ddp or use_fsdp)
and local_rank == 0
):
prepare_model_dir(**kwargs)
# parse kwargs
kwargs = model.kwargs
kwargs["device"] = device
tokenizer = kwargs["tokenizer"]
frontend = kwargs["frontend"]
model = model.model
del kwargs["model"]
# freeze_param
freeze_param = kwargs.get("freeze_param", None)
if freeze_param is not None:
if "," in freeze_param:
freeze_param = freeze_param.split(",")
if not isinstance(freeze_param, (list, tuple)):
freeze_param = (freeze_param,)
logging.info("freeze_param is not None: %s", freeze_param)
for t in freeze_param:
for k, p in model.named_parameters():
if k.startswith(t + ".") or k == t:
logging.info(f"Setting {k}.requires_grad = False")
p.requires_grad = False
lora_only = kwargs.get("lora_only", False)
if lora_only:
lora_bias = kwargs.get("lora_bias", "none")
logging.info("Enable LoRA-only training with bias=%s", lora_bias)
mark_only_lora_as_trainable(model, bias=lora_bias)
if local_rank == 0:
logging.info(f"{model_summary(model)}")
if use_ddp:
model = model.cuda(local_rank)
model = DDP(
model,
device_ids=[local_rank],
find_unused_parameters=kwargs.get("train_conf", {}).get(
"find_unused_parameters", False
),
)
elif use_fsdp:
# model = FSDP(model).cuda(local_rank)
def custom_auto_wrap_policy(
module: nn.Module,
recurse: bool,
nonwrapped_numel: int,
# Additional custom arguments
min_num_params: int = int(1e8),
) -> bool:
# 根据自定义逻辑决定是否包装模块
"""Custom auto wrap policy.
Args:
module: TODO.
recurse: TODO.
nonwrapped_numel: TODO.
min_num_params: TODO.
"""
is_large = nonwrapped_numel >= min_num_params
requires_grad_uniform = len({p.requires_grad for p in module.parameters()}) == 1
return is_large and requires_grad_uniform
# Configure a custom `min_num_params`
my_auto_wrap_policy = functools.partial(custom_auto_wrap_policy, min_num_params=int(1e5))
torch.cuda.set_device(local_rank)
model = FSDP(
model,
auto_wrap_policy=custom_auto_wrap_policy,
mixed_precision=None,
device_id=torch.cuda.current_device(),
)
else:
model = model.to(device=kwargs.get("device", "cuda"))
kwargs["device"] = next(model.parameters()).device
# optim
logging.info("Build optim")
optim = kwargs.get("optim", "adam")
assert optim in optim_classes
optim_class = optim_classes.get(optim)
optim = optim_class(model.parameters(), **kwargs.get("optim_conf"))
# scheduler
logging.info("Build scheduler")
scheduler = kwargs.get("scheduler", "warmuplr")
assert scheduler in scheduler_classes
scheduler_class = scheduler_classes.get(scheduler)
scheduler = scheduler_class(optim, **kwargs.get("scheduler_conf"))
# dataset
logging.info("Build dataloader")
dataloader_class = tables.dataloader_classes.get(
kwargs["dataset_conf"].get("dataloader", "DataloaderMapStyle")
)
dataloader = dataloader_class(**kwargs)
# dataloader_tr, dataloader_val = dataloader_class(**kwargs)
trainer = Trainer(
local_rank=local_rank,
use_ddp=use_ddp,
use_fsdp=use_fsdp,
device=kwargs["device"],
output_dir=kwargs.get("output_dir", "./exp"),
**kwargs.get("train_conf"),
)
scaler = GradScaler(enabled=trainer.use_fp16) if trainer.use_fp16 else None
scaler = ShardedGradScaler(enabled=trainer.use_fp16) if trainer.use_fsdp else scaler
trainer.resume_checkpoint(
model=model,
optim=optim,
scheduler=scheduler,
scaler=scaler,
)
tensorboard_dir = os.path.join(kwargs.get("output_dir"), "tensorboard")
os.makedirs(tensorboard_dir, exist_ok=True)
try:
writer = SummaryWriter(tensorboard_dir) # if trainer.rank == 0 else None
except:
writer = None
dataloader_tr, dataloader_val = None, None
for epoch in range(trainer.start_epoch, trainer.max_epoch):
time1 = time.perf_counter()
for data_split_i in range(trainer.start_data_split_i, dataloader.data_split_num):
time_slice_i = time.perf_counter()
dataloader_tr, dataloader_val = dataloader.build_iter(
epoch, data_split_i=data_split_i, start_step=trainer.start_step
)
trainer.train_epoch(
model=model,
optim=optim,
scheduler=scheduler,
scaler=scaler,
dataloader_train=dataloader_tr,
dataloader_val=dataloader_val,
epoch=epoch,
writer=writer,
data_split_i=data_split_i,
data_split_num=dataloader.data_split_num,
start_step=trainer.start_step,
)
trainer.start_step = 0
device = next(model.parameters()).device
if device.type == "cuda":
with torch.cuda.device(device):
torch.cuda.empty_cache()
time_escaped = (time.perf_counter() - time_slice_i) / 3600.0
logging.info(
f"rank: {local_rank}, "
f"time_escaped_epoch: {time_escaped:.3f} hours, "
f"estimated to finish {dataloader.data_split_num} data_slices, remaining: {dataloader.data_split_num-data_split_i} slices, {(dataloader.data_split_num-data_split_i)*time_escaped:.3f} hours, "
f"epoch: {trainer.max_epoch - epoch} epochs, {((trainer.max_epoch - epoch - 1)*dataloader.data_split_num + dataloader.data_split_num-data_split_i)*time_escaped:.3f} hours\n"
)
trainer.start_data_split_i = 0
trainer.validate_epoch(
model=model, dataloader_val=dataloader_val, epoch=epoch + 1, writer=writer
)
scheduler.step()
trainer.step_in_epoch = 0
trainer.save_checkpoint(
epoch + 1, model=model, optim=optim, scheduler=scheduler, scaler=scaler
)
time2 = time.perf_counter()
time_escaped = (time2 - time1) / 3600.0
logging.info(
f"rank: {local_rank}, "
f"time_escaped_epoch: {time_escaped:.3f} hours, "
f"estimated to finish {trainer.max_epoch} "
f"epoch: {(trainer.max_epoch - epoch) * time_escaped:.3f} hours\n"
)
trainer.train_acc_avg = 0.0
trainer.train_loss_avg = 0.0
if trainer.rank == 0:
average_checkpoints(trainer.output_dir, trainer.avg_nbest_model)
trainer.close()
if __name__ == "__main__":
main_hydra()
+259
View File
@@ -0,0 +1,259 @@
#!/usr/bin/env python3
# -*- encoding: utf-8 -*-
import os
import sys
import torch
import torch.nn as nn
import hydra
import logging
import time
import argparse
from io import BytesIO
from contextlib import nullcontext
import torch.distributed as dist
from omegaconf import DictConfig, OmegaConf
from torch.cuda.amp import autocast, GradScaler
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.algorithms.join import Join
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
from funasr.train_utils.average_nbest_models import average_checkpoints
from funasr.register import tables
from funasr.optimizers import optim_classes
from funasr.train_utils.trainer_ds import Trainer
from funasr.schedulers import scheduler_classes
from funasr.train_utils.initialize import initialize
from funasr.download.download_model_from_hub import download_model
from funasr.models.lora.utils import mark_only_lora_as_trainable
from funasr.train_utils.set_all_random_seed import set_all_random_seed
from funasr.train_utils.load_pretrained_model import load_pretrained_model
from funasr.utils.misc import prepare_model_dir
from funasr.train_utils.model_summary import model_summary
from funasr import AutoModel
try:
import deepspeed
except:
deepspeed = None
@hydra.main(config_name=None, version_base=None)
def main_hydra(kwargs: DictConfig):
"""Main hydra.
Args:
kwargs: Additional keyword arguments.
"""
if kwargs.get("debug", False):
import pdb
pdb.set_trace()
assert "model" in kwargs
if "model_conf" not in kwargs:
logging.info("download models from model hub: {}".format(kwargs.get("hub", "ms")))
kwargs = download_model(is_training=kwargs.get("is_training", True), **kwargs)
main(**kwargs)
def main(**kwargs):
# set random seed
"""Main.
Args:
**kwargs: Additional keyword arguments.
"""
set_all_random_seed(kwargs.get("seed", 0))
torch.backends.cudnn.enabled = kwargs.get("cudnn_enabled", torch.backends.cudnn.enabled)
torch.backends.cudnn.benchmark = kwargs.get("cudnn_benchmark", torch.backends.cudnn.benchmark)
torch.backends.cudnn.deterministic = kwargs.get("cudnn_deterministic", True)
# open tf32
torch.backends.cuda.matmul.allow_tf32 = kwargs.get("enable_tf32", True)
rank = int(os.environ.get("RANK", 0))
local_rank = int(os.environ.get("LOCAL_RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
if local_rank == 0:
tables.print()
use_ddp = world_size > 1
use_fsdp = kwargs.get("use_fsdp", False)
use_deepspeed = kwargs.get("use_deepspeed", False)
if use_deepspeed:
logging.info(f"use_deepspeed: {use_deepspeed}")
deepspeed.init_distributed(dist_backend=kwargs.get("backend", "nccl"))
elif use_ddp or use_fsdp:
logging.info(f"use_ddp: {use_ddp}, use_fsdp: {use_fsdp}")
dist.init_process_group(
backend=kwargs.get("backend", "nccl"),
init_method="env://",
)
torch.cuda.set_device(local_rank)
# rank = dist.get_rank()
logging.info("Build model, frontend, tokenizer")
device = kwargs.get("device", "cuda")
kwargs["device"] = "cpu"
model = AutoModel(**kwargs)
# save config.yaml
if rank == 0:
prepare_model_dir(**kwargs)
# parse kwargs
kwargs = model.kwargs
kwargs["device"] = device
tokenizer = kwargs["tokenizer"]
frontend = kwargs["frontend"]
model = model.model
del kwargs["model"]
# freeze_param
freeze_param = kwargs.get("freeze_param", None)
if freeze_param is not None:
if "," in freeze_param:
freeze_param = freeze_param.split(",")
if not isinstance(freeze_param, (list, tuple)):
freeze_param = (freeze_param,)
logging.info("freeze_param is not None: %s", freeze_param)
for t in freeze_param:
for k, p in model.named_parameters():
if k.startswith(t + ".") or k == t:
logging.info(f"Setting {k}.requires_grad = False")
p.requires_grad = False
lora_only = kwargs.get("lora_only", False)
if lora_only:
lora_bias = kwargs.get("lora_bias", "none")
logging.info("Enable LoRA-only training with bias=%s", lora_bias)
mark_only_lora_as_trainable(model, bias=lora_bias)
if local_rank == 0:
logging.info(f"{model_summary(model)}")
trainer = Trainer(
rank=rank,
local_rank=local_rank,
world_size=world_size,
use_ddp=use_ddp,
use_fsdp=use_fsdp,
device=kwargs["device"],
excludes=kwargs.get("excludes", None),
output_dir=kwargs.get("output_dir", "./exp"),
**kwargs.get("train_conf"),
)
model = trainer.warp_model(model, **kwargs)
kwargs["device"] = int(os.environ.get("LOCAL_RANK", 0))
trainer.device = int(os.environ.get("LOCAL_RANK", 0))
model, optim, scheduler = trainer.warp_optim_scheduler(model, **kwargs)
# dataset
logging.info("Build dataloader")
dataloader_class = tables.dataloader_classes.get(
kwargs["dataset_conf"].get("dataloader", "DataloaderMapStyle")
)
dataloader = dataloader_class(**kwargs)
# dataloader_tr, dataloader_val = dataloader_class(**kwargs)
scaler = GradScaler(enabled=True) if trainer.use_fp16 else None
scaler = ShardedGradScaler(enabled=trainer.use_fp16) if trainer.use_fsdp else scaler
trainer.resume_checkpoint(
model=model,
optim=optim,
scheduler=scheduler,
scaler=scaler,
)
early_stopping_patience = kwargs.get("train_conf", {}).get("early_stopping_patience", 0)
best_val_loss = float("inf")
epochs_no_improve = 0
dataloader_tr, dataloader_val = None, None
for epoch in range(trainer.start_epoch, trainer.max_epoch):
time1 = time.perf_counter()
for data_split_i in range(trainer.start_data_split_i, dataloader.data_split_num):
time_slice_i = time.perf_counter()
dataloader_tr, dataloader_val = dataloader.build_iter(
epoch, data_split_i=data_split_i, start_step=trainer.start_step
)
trainer.train_epoch(
model=model,
optim=optim,
scheduler=scheduler,
scaler=scaler,
dataloader_train=dataloader_tr,
dataloader_val=dataloader_val,
epoch=epoch,
data_split_i=data_split_i,
data_split_num=dataloader.data_split_num,
start_step=trainer.start_step,
)
trainer.start_step = 0
device = next(model.parameters()).device
if device.type == "cuda":
with torch.cuda.device(device):
torch.cuda.empty_cache()
time_escaped = (time.perf_counter() - time_slice_i) / 3600.0
logging.info(
f"\n\nrank: {local_rank}, "
f"time_escaped_epoch: {time_escaped:.3f} hours, "
f"estimated to finish {dataloader.data_split_num} data_slices, remaining: {dataloader.data_split_num-data_split_i} slices, {(dataloader.data_split_num-data_split_i)*time_escaped:.3f} hours, "
f"epoch: {trainer.max_epoch - epoch} epochs, {((trainer.max_epoch - epoch - 1)*dataloader.data_split_num + dataloader.data_split_num-data_split_i)*time_escaped:.3f} hours\n"
)
trainer.start_data_split_i = 0
trainer.validate_epoch(model=model, dataloader_val=dataloader_val, epoch=epoch + 1)
current_val = trainer.val_loss_avg
if current_val < best_val_loss:
logging.info(f"current_val: {current_val}, best_val_loss: {best_val_loss}")
best_val_loss = current_val
epochs_no_improve = 0
else:
epochs_no_improve += 1
logging.info(f"No val_loss improvement for {epochs_no_improve}/{early_stopping_patience} epochs")
if early_stopping_patience > 0 and epochs_no_improve >= early_stopping_patience:
logging.info(f"Early stopping triggered at epoch {epoch+1}")
break
trainer.step_in_epoch = 0
trainer.save_checkpoint(
epoch + 1, model=model, optim=optim, scheduler=scheduler, scaler=scaler
)
time2 = time.perf_counter()
time_escaped = (time2 - time1) / 3600.0
logging.info(
f"\n\nrank: {local_rank}, "
f"time_escaped_epoch: {time_escaped:.3f} hours, "
f"estimated to finish {trainer.max_epoch} "
f"epoch: {(trainer.max_epoch - epoch) * time_escaped:.3f} hours\n"
)
trainer.train_acc_avg = 0.0
trainer.train_loss_avg = 0.0
if trainer.rank == 0:
average_checkpoints(
trainer.output_dir, trainer.avg_nbest_model, use_deepspeed=trainer.use_deepspeed
)
trainer.close()
if __name__ == "__main__":
main_hydra()