【Bug已解决】[Feature Proposal] Add FunASR ASR Distributed Inference Examples 解决方案

一、现象长什么样

FunASR(阿里巴巴开源的语音识别工具箱)做批量 / 长音频推理时,单张 GPU 成了明显瓶颈:

  • 几千个音频文件排队,单卡串行跑要数小时,GPU 利用率还不到 50%(因为前后处理是 CPU 密集)。
  • 想「多卡并行」却没有官方分布式推理示例,于是用户各显神通:有的用 multiprocessing 起 N 个进程各加载一份模型(VRAM 爆了),有的手动把音频列表按 rank 切分但聚合结果错位(第 3 卡的转写混到第 1 卡的输出里)。

最容易踩的「非崩溃但错」现象:

# 多进程各加载模型,结果文件互相覆盖 / 顺序错乱
rank0 写出 result.jsonl  (含它负责的 1000 条)
rank1 也写出 result.jsonl  (同名覆盖,或 append 乱序)

本质:FunASR 官方示例几乎都是单卡 model.generate(),缺少「如何把一批音频正确、省显存地分到多卡上推理并聚合」的分布式范式。用户缺范例就只能自己 hack,于是要么显存爆、要么结果错位。

二、背景

FunASR 的推理入口是 AutoModel + model.generate(audio_in=...)。单卡模式下,你把所有音频塞进一个列表,模型逐个(或分批)generate。问题在规模:

  • 计算 vs 前后处理不匹配:ASR 的 VAD(语音活动检测)、特征提取、后处理很多在 CPU/模型外,单卡 GPU 在等这些时就空闲。多卡能把「不同音频」分到不同卡,整体吞吐上去。
  • 没有分布式范例,用户不知道该「按音频分片、每卡一份模型副本、各自写各自的结果分片、最后合并」,还是「一份模型靠 tensor parallel」。对 ASR 这种数据并行友好的任务,前者(按音频分片)最简单高效,但没人示范。

于是出现两类错误实践:

  1. 每进程全量加载模型:N 个进程各自 AutoModel(...),每个吃一份完整 VRAM,N 张卡显存都爆。正确做法是「每卡一个进程、一份模型」,本来就该这样,但用户常误以为要「主进程加载、子进程共享」,在 FunASR 里共享 CUDA 上下文很麻烦,反而写错。
  2. 结果聚合错位:分片推理后,各 rank 把结果写同一文件,顺序/覆盖混乱。正确做法是「各 rank 写独立分片文件,最后按全局索引合并」。

这个 Feature Proposal 就是要补上「用 Accelerate / torch.distributed 做 FunASR 数据并行推理」的范例。

三、根因(能力缺口分析)

把这个缺口当 bug 分析,根因是 FunASR 缺分布式推理范式,导致用户错用多进程(显存爆)或错聚合结果(错位),三层:

第一层(主因):缺「按音频数据并行的标准范式」。 用户不知道应该 rank = process_index; 负责 audio_list[rank::world_size],于是要么全量加载(爆显存),要么手动切分得不对。范例缺失是直接原因。

第二层:结果聚合无规范。 多卡各自产出结果,怎么汇成一份有序输出没有示范,用户用「同文件名写」导致覆盖/乱序。这是数据并行最常见的坑,但缺范例时必踩。

第三层:模型加载与 rank 生命周期不匹配。 用户常在「fork 出的子进程里加载模型」,但 FunASR/CUDA 要求模型在目标 rank 的进程里、init_process_group 之后才加载。顺序错了会报 CUDA error 或上下文错乱。范例应明确「先 init、后加载」的顺序。

一句话:缺分布式推理范例 → 用户错用多进程(显存爆)/错聚合(结果错位)/错加载顺序(CUDA 上下文乱)。

四、最小可运行复现

下面用纯 Python 模拟「错误聚合(同文件覆盖)vs 正确聚合(分片文件后合并)」的差异,不需要 GPU:

from dataclasses import dataclass
from typing import List, Dict


@dataclass
class Audio:
    gid: int       # 全局索引
    path: str


def shard(audios: List[Audio], rank: int, world: int) -> List[Audio]:
    return audios[rank::world]


def wrong_aggregate(rank_results: Dict[int, List[str]], out_file: str):
    # 错误:所有 rank 写同一个文件,后写覆盖先写
    with open(out_file, "w") as f:
        f.write("\n".join(rank_results[max(rank_results)]))


def correct_aggregate(rank_results: Dict[int, Dict[int, str]], out_file: str):
    # 正确:按全局 gid 合并各 rank 分片
    merged = {}
    for res in rank_results.values():
        merged.update(res)
    with open(out_file, "w") as f:
        for gid in sorted(merged):
            f.write(f"{gid}: {merged[gid]}\n")


def main():
    audios = [Audio(i, f"a{i}.wav") for i in range(8)]
    world = 2
    r0 = {a.gid: f"text{a.gid}" for a in shard(audios, 0, world)}
    r1 = {a.gid: f"text{a.gid}" for a in shard(audios, 1, world)}

    # 正确聚合:合并后应有 8 条、顺序正确
    correct_aggregate({0: r0, 1: r1}, "result.txt")
    print("正确聚合后条目数:", len(r0) + len(r1))   # 8


if __name__ == "__main__":
    main()

跑出来正确聚合得到 8 条有序结果——演示了「分片推理 + 按 gid 合并」为何能对。错误聚合则会丢一半(覆盖)。

五、解决方案(第一层:最小直接修复)

最省事的落地:用 Accelerate 做「按音频数据并行」的标准范式,每 rank 加载一份模型、负责 1/N 音频、写独立分片、最后合并:

from accelerate import Accelerator
from funasr import AutoModel

accelerator = Accelerator()

# 先 init(Accelerator 已帮做),再在每 rank 加载模型 —— 顺序关键
model = AutoModel(
    model="paraformer-zh",
    device=accelerator.device,
    disable_update=True,
)

# 读取完整音频列表,按 rank 分片
all_audios = load_audio_manifest("manifest.txt")   # 全局有序列表
my_audios = all_audios[accelerator.process_index::accelerator.num_processes]

results = {}
for a in my_audios:
    text = model.generate(input=a.path, batch_size_s=300)[0]["text"]
    results[a.gid] = text                            # 用全局 gid 作 key

# 每 rank 写独立分片,绝不共用文件名
with open(f"result_rank{accelerator.process_index}.jsonl", "w") as f:
    for gid, text in results.items():
        f.write(f"{gid}\t{text}\n")

最后在所有 rank 结束后,用一个小脚本按 gid 合并各 result_rank*.jsonl。这是避免覆盖/乱序的关键。

六、解决方案(第二层:结构性改进)

第一层是「手动范式」,第二层是「封装一个 FunASR 分布式推理工具,把分片/加载/聚合都固化」,从设计上消灭错用:

from dataclasses import dataclass
from typing import List, Dict, Callable


@dataclass
class DistASRConfig:
    model_name: str
    manifest: str
    out_dir: str


class FunASRDistributor:
    """FunASR 数据并行推理的规范范式(单一事实来源)。"""

    def __init__(self, accelerator, model_builder: Callable):
        self.acc = accelerator
        # 关键:init 之后、在当前 rank 进程里加载模型
        self.model = model_builder(device=accelerator.device)

    def run(self, audios, infer_fn: Callable) -> Dict[int, str]:
        # 1) 按 rank 分片
        my = audios[self.acc.process_index::self.acc.num_processes]
        # 2) 各自推理
        out: Dict[int, str] = {}
        for a in my:
            out[a.gid] = infer_fn(self.model, a.path)
        # 3) 写独立分片
        self._dump_shard(out)
        return out

    def _dump_shard(self, out: Dict[int, str]):
        rank = self.acc.process_index
        path = f"{self.out_dir}/result_rank{rank}.jsonl"
        with open(path, "w") as f:
            for gid, text in out.items():
                f.write(f"{gid}\t{text}\n")

    @staticmethod
    def merge(out_dir: str, total: int) -> List[str]:
        # 4) 合并:按全局 gid 排序
        merged: Dict[int, str] = {}
        import glob
        for fp in glob.glob(f"{out_dir}/result_rank*.jsonl"):
            for line in open(fp):
                gid, text = line.rstrip("\n").split("\t", 1)
                merged[int(gid)] = text
        return [merged[i] for i in range(total)]


# 用法
dist = FunASRDistributor(accelerator, lambda device: AutoModel(model="paraformer-zh", device=device))
audios = [Audio(i, f"a{i}.wav") for i in range(8000)]
dist.run(audios, lambda m, p: m.generate(input=p)[0]["text"])
# 合并(任意一 rank 或单独脚本)
if accelerator.is_main_process:
    final = FunASRDistributor.merge("out", total=8000)

这样分片、加载顺序、独立分片写、合并都固化,用户不会再踩显存爆/结果错位。

七、解决方案(第三层:断言 / CI 守护)

把「分片不重叠不遗漏」「独立分片写不覆盖」「合并有序」固化成测试:

import pytest


def test_shard_no_overlap_no_missing():
    audios = [Audio(i, f"a{i}") for i in range(8)]
    world = 2
    s0 = shard(audios, 0, world)
    s1 = shard(audios, 1, world)
    gids = {a.gid for a in s0} | {a.gid for a in s1}
    assert gids == set(range(8))          # 不遗漏
    assert not ({a.gid for a in s0} & {a.gid for a in s1})  # 不重叠


def test_merge_ordered():
    r0 = {0: "t0", 2: "t2", 4: "t4", 6: "t6"}
    r1 = {1: "t1", 3: "t3", 5: "t5", 7: "t7"}
    out = FunASRDistributor.merge_from_dicts([r0, r1], total=8)
    assert out == [f"t{i}" for i in range(8)]


def test_one_model_per_rank():
    # 每 rank 只加载一次模型(避免 N 进程 N 份全量加载)
    loads = []
    def builder(device):
        loads.append(device)
        return "model"
    acc = FakeAccelerator(process_index=0, num_processes=4)
    d = FunASRDistributor(acc, builder)
    assert len(loads) == 1     # 仅一次加载


def test_no_shared_output_file():
    # 每 rank 写独立文件,不会同文件名覆盖
    paths = [f"result_rank{r}.jsonl" for r in range(4)]
    assert len(set(paths)) == 4


def test_full_pipeline_distributed():
    audios = [Audio(i, f"a{i}") for i in range(16)]
    world = 4
    parts = [shard(audios, r, world) for r in range(world)]
    all_gids = set()
    for p in parts:
        all_gids |= {a.gid for a in p}
    assert all_gids == set(range(16))

八、排查清单

  1. 看多卡 ASR 是否显存爆(每进程全量加载)或结果文件覆盖/乱序 → 是缺分布式范式。
  2. 确认是否「先 init_process_group / Accelerator 再加载模型」,顺序错会 CUDA 上下文乱。
  3. 临时救火:用 Accelerator 按 process_index::num_processes 分片音频,每 rank 一份模型、独立分片文件。
  4. 检查合并脚本是否按全局 gid 排序,避免覆盖(各 rank 写不同文件名)。
  5. 长期方向:采用封装好的 FunASRDistributor 范式(分片/加载/写/合并一体化)。
  6. 升级 FunASR / 参考合了该范例的版本,并跑上面的「分片不重叠」「合并有序」用例。
  7. 若音频极长,单条就超 VRAM,应在分片基础上再做「按 chunk 切分单条音频」,而非单纯加卡。

九、小结

FunASR 分布式推理范例缺口,不是工具不能用,而是缺「按音频数据并行」的标准范式,导致用户错用多进程(显存爆)、错聚合(结果覆盖/乱序)、错加载顺序(CUDA 上下文乱)。最小修复是用 Accelerate 按 rank 分片音频、每 rank 一份模型、写独立分片文件、最后按 gid 合并;结构性方向是封装 FunASRDistributor 把分片/加载/写/合并固化;最后用 pytest 把「分片不重叠不遗漏」「独立分片不覆盖」「合并有序」锁死。抓住「ASR 是数据并行友好的任务、按音频分片 + 独立分片写 + 全局索引合并」这条,所有批量 ASR 的分布式化都能照此落地。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐