Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled
Update API Documentation / build-api-docs (push) Has been cancelled
Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
@@ -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
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Executable
+307
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user