Instructions to use czyhust/finetune_moss-sortformer with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- NeMo
How to use czyhust/finetune_moss-sortformer with NeMo:
# tag did not correspond to a valid NeMo domain.
- Notebooks
- Google Colab
- Kaggle
降低 PIT/ATS 排列复杂度的方案
当前截断排列(truncate permutation)方案直接将后
N - max_perm_spks个说话人排除在排列优化之外,导致其标签顺序完全随机,损害模型的排列不变性。本文从算法与工程两个维度,提出多种替代方案。
一、问题重述
核心矛盾:
# 全排列方案(不可行)
permed_labels = labels[:, :, speaker_permutations] # (B, T, N!, N) → 10人时 ~73GB
# 截断排列方案(当前,有缺陷)
permed_labels = labels[:, :, :6][:, :, speaker_permutations[:720]] # 仅前6人参与排列,后4人原样保留
关键观察:排列搜索本质上只需要一个 N×N 的 pairwise 匹配分数矩阵,而非完整的 (B, T, N!, N) 张量。排列数量 N! 的爆炸可以被完全绕过。
二、方案一(推荐):匈牙利算法 + N×N Score Matrix
2.1 数学原理
PIT 的匹配分数可分解:
match_score[k] = Σ_t Σ_i labels[t, σ_k(i)] · preds[t, i]
= Σ_i ( Σ_t labels[t, σ_k(i)] · preds[t, i] )
= Σ_i S[σ_k(i), i]
其中 S[j, i] = Σ_t labels[t, j] · preds[t, i] 是一个 N×N 的 pairwise score 矩阵,表示 ground truth 说话人 j 与预测轨道 i 的匹配度。
问题变为:在 N×N 的 score 矩阵上,找一个排列 σ 最大化 Σ_i S[σ(i), i]。这正是经典线性分配问题(Linear Assignment Problem),可用匈牙利算法在 O(N³) 时间内精确求解。
2.2 算法流程
输入: labels (B, T, N), preds (B, T, N)
输出: targets_pil (B, T, N) — 按最优排列重排后的 labels
Step 1: 计算 pairwise score matrix
S[b, j, i] = Σ_t labels[b, t, j] · preds[b, t, i] # (B, N, N)
内存: B × N² × 4 bytes → 10人时仅 400 bytes/batch
Step 2: 对每个 batch 元素,求解线性分配问题
for b in range(B):
σ_b = hungarian_algorithm(S[b]) # 最大化 Σ_i S[σ(i), i]
时间复杂度: O(B × N³) → 10人时约 1000 次操作/样本
Step 3: 按最优排列重排 labels
targets_pil[b, :, :] = labels[b, :, σ_b]
2.3 对比当前方案
| 指标 | 全排列 (N=10) | 截断排列 (当前) | 匈牙利算法 (推荐) |
|---|---|---|---|
| 中间张量形状 | (B, T, 3.6M, 10) | (B, T, 720, 6) | (B, N, N) |
| N=10 显存 | ~73 GB | ~9 MB | ~400 bytes |
| 排列完备性 | 全部 10 人 | 仅前 6 人 | 全部 10 人 |
| 最优性 | 全局最优 | 不保证最优 | 全局最优 |
| 时间复杂度 | O(N! · T) | O(720 · T) | O(N³ + N² · T) |
2.4 代码实现
import torch
from scipy.optimize import linear_sum_assignment
def get_pit_targets_hungarian(labels: torch.Tensor, preds: torch.Tensor) -> torch.Tensor:
"""
使用匈牙利算法的 PIT targets 计算,完全不产生 O(N!) 中间张量。
Args:
labels: (B, T, N) ground truth
preds: (B, T, N) model predictions (sigmoid outputs)
Returns:
targets_pil: (B, T, N) labels reordered by optimal permutation
"""
B, T, N = labels.shape
# Step 1: 计算 N×N pairwise score matrix
# S[b, j, i] = Σ_t labels[b, t, j] · preds[b, t, i]
# 通过 einsum 高效计算: 'btj,bti->bji'
S = torch.einsum('btj,bti->bji', labels, preds) # (B, N, N)
best_perms = []
for b in range(B):
# Step 2: 匈牙利算法求最大权重匹配
# linear_sum_assignment 最小化 cost,所以取负号
cost = -S[b].detach().cpu().numpy()
row_ind, col_ind = linear_sum_assignment(cost)
# row_ind[j] = i 意味着 gt speaker j → pred slot i
# 我们需要 pred slot i → gt speaker j 的映射
perm = torch.zeros(N, dtype=torch.long)
perm[col_ind] = torch.tensor(row_ind)
best_perms.append(perm)
best_perms = torch.stack(best_perms).to(labels.device) # (B, N)
# Step 3: 按最优排列重排 labels
batch_perm_inds_exp = best_perms.unsqueeze(1).expand(-1, T, -1)
targets_pil = torch.gather(labels, 2, batch_perm_inds_exp)
return targets_pil
def get_ats_targets_hungarian(
labels: torch.Tensor,
preds: torch.Tensor,
thres: float = 0.5,
tolerance: float = 0
) -> torch.Tensor:
"""
使用匈牙利算法 + 到达时间约束的 ATS targets 计算。
与 PIT 的区别:只有到达时间顺序一致的 (j,i) 配对才被允许参与分配。
"""
B, T, N = labels.shape
# 计算 pairwise score matrix
S = torch.einsum('btj,bti->bji', labels, preds) # (B, N, N)
# 计算每个说话人的首次发言帧(到达时间)
nonzero_ind = find_first_nonzero(labels, max_cap_val=T, thres=thres) # (B, N)
sorted_values = torch.sort(nonzero_ind, dim=1)[0] # (B, N)
best_perms = []
for b in range(B):
cost = -S[b].detach().cpu().numpy() # (N, N)
# ATS 约束:仅允许 arrival_time(gt j) ≈ sorted_arrival_time(pred slot i) 的配对
for j in range(N): # ground truth speaker
for i in range(N): # prediction slot (by arrival order)
if abs(nonzero_ind[b, j].item() - sorted_values[b, i].item()) > tolerance:
cost[j, i] = 1e9 # 禁止此配对
row_ind, col_ind = linear_sum_assignment(cost)
perm = torch.zeros(N, dtype=torch.long)
perm[col_ind] = torch.tensor(row_ind)
best_perms.append(perm)
best_perms = torch.stack(best_perms).to(labels.device)
batch_perm_inds_exp = best_perms.unsqueeze(1).expand(-1, T, -1)
targets_ats = torch.gather(labels, 2, batch_perm_inds_exp)
return targets_ats
2.5 集成到现有代码
修改 asr_multispeaker_utils.py 中的 get_pil_targets 和 get_ats_targets,新增基于匈牙利算法的路径:
def get_pil_targets(labels, preds, speaker_permutations, use_hungarian: bool = False):
if use_hungarian:
return _get_pil_targets_hungarian(labels, preds)
# ... 原有全排列逻辑(向后兼容)
配置入口(YAML):
model:
pit_algorithm: hungarian # "full_permutation" | "truncated_permutation" | "hungarian"
max_perm_spks: 6 # 仅在 truncated_permutation 模式下生效
2.6 关于匈牙利算法实现的注意事项
scipy.optimize.linear_sum_assignment是 CPU 实现的,每次调用需要将 N×N 矩阵从 GPU 拷贝到 CPU。对于 N=10 这是可忽略的开销(复制 400 bytes)- 如需纯 GPU 实现,可参考 PyTorch 版 Jonker-Volgenant 算法 或自行实现基于
torch的匈牙利算法 - 训练时建议使用
torch.no_grad()包裹排列搜索部分(排列本身不需要梯度),梯度通过重排后的 labels 流回 BCE loss
三、方案二:Stochastic Permutation Sampling(随机排列采样)
3.1 原理
不枚举所有 N! 种排列,而是随机采样 K 种排列(K << N!),从中选最佳匹配。这在统计上是无偏的,且随着 K 增大趋近于全排列结果。
3.2 实现
def sample_random_permutations(N: int, K: int, device: torch.device) -> torch.Tensor:
"""随机采样 K 种不同的 N 元排列,避免重复"""
perms = set()
while len(perms) < min(K, math.factorial(N)):
perm = tuple(torch.randperm(N).tolist())
perms.add(perm)
return torch.tensor(list(perms), device=device) # (K, N)
def get_pit_targets_sampled(labels, preds, K: int = 720):
B, T, N = labels.shape
sampled_perms = sample_random_permutations(N, K, labels.device) # (K, N)
permed_labels = labels[:, :, sampled_perms] # (B, T, K, N)
preds_rep = preds.unsqueeze(2).repeat(1, 1, K, 1) # (B, T, K, N)
match_score = (permed_labels * preds_rep).sum(1).sum(2) # (B, K)
best_perm_inds = find_best_permutation(match_score, sampled_perms)
return reconstruct_labels(labels, best_perm_inds)
3.3 适用场景
| 优点 | 缺点 |
|---|---|
| 实现极简,完全复用现有流程 | 不保证全局最优(采样运气成分) |
| K 可动态调整(初期大 K,后期小 K) | N 很大时采样效率下降 |
可无缝替换 speaker_permutations |
排列重复检测有开销 |
经验值:K=720 (6!) 在 10 人场景中通常能找到接近最优的排列,因为大量排列的匹配分数本就极低。
四、方案三:Progressive Permutation(渐进排列)
4.1 原理
分两阶段工作:
- 粗匹配(coarse matching):用 N×N score matrix(同方案一)做一次贪心匹配,得到初始排列
- 精调(local refinement):在初始排列的邻域内做有限的全排列搜索(如交换 2-3 个说话人)
4.2 实现
def get_pit_targets_progressive(labels, preds, max_swap: int = 3):
B, T, N = labels.shape
targets_all = []
for b in range(B):
# Stage 1: 贪心匹配得到初始排列
S = torch.einsum('tj,ti->ji', labels[b], preds[b]) # (N, N)
init_perm = greedy_matching(S) # 或匈牙利算法
# Stage 2: 局部邻域搜索(交换至多 max_swap 个位置)
local_perms = generate_swap_permutations(init_perm, max_swap)
# local_perms: (K, N) where K = C(N, max_swap) * max_swap!
permed_labels = labels[b:b+1, :, local_perms] # (1, T, K, N)
preds_rep = preds[b:b+1].unsqueeze(2).repeat(1, 1, K, 1)
match_score = (permed_labels * preds_rep).sum(1).sum(2) # (1, K)
best_local = local_perms[match_score.argmax()]
targets_all.append(labels[b, :, best_local])
return torch.stack(targets_all, dim=0)
4.3 适用场景
当匈牙利算法的 O(N³) 并不构成瓶颈时,此方案相比方案一无明显优势。但在 N 非常大(>50)且需要极致内存效率时,贪心 + 局部优化的组合可进一步降低复杂度。
五、方案四:Pairwise Score Cache + Chunked Enumeration
5.1 原理
不一次性生成所有排列的中间张量,而是在循环中逐批处理排列,每次只计算 (B, T, chunk_size, N) 的 permed_labels,累积 match_score 后取最大值。核心 trick 是预先计算 N×N 的 pairwise score matrix 以加速循环内的计算。
5.2 实现
def get_pit_targets_chunked(
labels, preds, speaker_permutations, chunk_size: int = 1000
):
B, T, N = labels.shape
n_perm = speaker_permutations.shape[0]
# Pre-compute pairwise scores: S[b, j, i] = Σ_t labels[t,j] · preds[t,i]
S = torch.einsum('btj,bti->bji', labels, preds) # (B, N, N)
best_score = torch.full((B,), -float('inf'), device=labels.device)
best_perms = torch.zeros((B, N), dtype=torch.long, device=labels.device)
for start in range(0, n_perm, chunk_size):
end = min(start + chunk_size, n_perm)
perms_chunk = speaker_permutations[start:end] # (chunk, N)
chunk_len = perms_chunk.shape[0]
# 利用预计算的 S 矩阵计算匹配分数(O(chunk × N) 而非 O(chunk × T × N))
# match_score[b, k] = Σ_i S[b, perms_chunk[k, i], i]
match_chunk = torch.zeros((B, chunk_len), device=labels.device)
for b in range(B):
for k in range(chunk_len):
for i in range(N):
match_chunk[b, k] += S[b, perms_chunk[k, i], i]
# 更新最佳排列
chunk_best_idx = match_chunk.argmax(dim=1)
for b in range(B):
if match_chunk[b, chunk_best_idx[b]] > best_score[b]:
best_score[b] = match_chunk[b, chunk_best_idx[b]]
best_perms[b] = perms_chunk[chunk_best_idx[b]]
return reconstruct_labels(labels, best_perms)
关键优化:通过预计算 N×N pairwise score matrix,chunk 循环内的计算从 O(chunk × T × N) 降为 O(chunk × N),因为去除了 T 维度上的求和。
然而,此方案的循环仍然需要遍历所有 N! 排列(对于 N=10 是 360 万次),即便每次只需要 O(N) 操作,总计 3.6M × 10 = 36M 次浮点运算——理论上可行但相比匈牙利算法仍显低效。因此此方案不推荐作为首选,仅作为无匈牙利算法实现时的备选。
六、方案五:Sinkhorn-based Differentiable Assignment
6.1 原理
将离散排列搜索松弛为连续的双随机矩阵(doubly stochastic matrix),通过 Sinkhorn 迭代逼近最优匹配。这是一种完全可微的方法,梯度可以流过排列过程。
6.2 实现草图
def sinkhorn_pit_targets(labels, preds, n_iter=20, tau=0.1):
"""使用 Sinkhorn 算法计算软排列下的 PIT targets"""
B, T, N = labels.shape
S = torch.einsum('btj,bti->bji', labels, preds) # (B, N, N)
# Sinkhorn: 将 score matrix 转化为双随机矩阵
P = torch.softmax(S / tau, dim=-1) # (B, N, N) 行归一化初始化
for _ in range(n_iter):
P = P / P.sum(dim=1, keepdim=True) # 行归一化
P = P / P.sum(dim=2, keepdim=True) # 列归一化
# 软排列重排 labels
# targets_soft = Σ_j labels[:,:,j] · P[b,j,i]
targets_soft = torch.einsum('btj,bji->bti', labels, P)
return targets_soft, P
6.3 优缺点
| 优点 | 缺点 |
|---|---|
| 完全可微,梯度自然流动 | 软排列不够"硬",训练初期可能模糊 |
| 无需求解离散优化问题 | Sinkhorn 迭代有计算开销 |
| 理论上可处理极大规模 | 温度参数 τ 需要调优 |
| 内存仅 O(N²) | 训练后期需退火至硬排列 |
七、方案对比与推荐
| 维度 | 匈牙利算法 | 随机采样 | 渐进排列 | Chunked Enum | Sinkhorn |
|---|---|---|---|---|---|
| 内存复杂度 | O(N²) | O(K·N) | O(K·N) | O(N²) | O(N²) |
| 时间复杂度的主导项 | O(N³) | O(K·T·N) | O(N³ + K·T·N) | O(N!·N) | O(N²·iter) |
| 排列最优性 | ✅ 精确 | ⚠️ 近似 | ⚠️ 近最优 | ✅ 精确 | ⚠️ 软分配 |
| 实现复杂度 | 中 | 低 | 中-高 | 低 | 中 |
| 与现有代码兼容 | 需重写 | 即插即用 | 需重写 | 需重写 | 需重写 |
| GPU 友好 | ⚠️ 需 CPU | ✅ | ✅ | ✅ | ✅ |
| 可微性 | ❌ | ❌ | ❌ | ❌ | ✅ |
推荐实施路径
短期(立即可用):
└── 方案二(随机采样): 改动量最小,将 speaker_permutations 替换为随机采样的排列即可
中期(最优方案):
└── 方案一(匈牙利算法): 内存 O(N²),全局最优,彻底解决排列爆炸问题
配套修改:
1. asr_multispeaker_utils.py: 新增 get_pit_targets_hungarian / get_ats_targets_hungarian
2. sortformer_diar_models.py: 移除 speaker_permutations 生成,传递 N 即可
3. YAML: 新增 pit_algorithm: hungarian 配置项
可选增强:
└── 方案五(Sinkhorn): 如果将来需要完全端到端、可微的排列学习
工程实现建议
匈牙利算法的 fallback 兼容:同时保留原有的全排列路径(N≤6 时),仅当 N>6 时启用匈牙利算法:
class SortformerEncLabelModel:
def __init__(self, cfg, trainer=None):
...
self.n_spk = cfg.max_num_of_spks
self.pit_algorithm = cfg.get("pit_algorithm", "full_permutation")
if self.pit_algorithm == "hungarian" or self.n_spk > 6:
# 使用匈牙利算法,无需预先生成排列矩阵
self.speaker_permutations = None
self._get_pil_fn = get_pit_targets_hungarian
self._get_ats_fn = get_ats_targets_hungarian
else:
# 原有全排列路径
self.speaker_permutations = torch.tensor(list(itertools.permutations(range(self.n_spk))))
self._get_pil_fn = get_pil_targets
self._get_ats_fn = get_ats_targets
八、总结
- 根本原因:排列爆炸源于代码在 T(时间帧)维度上扩展排列,而实际上 PIT 匹配分数只需 N×N 的 pairwise 统计量
- 最佳方案:匈牙利算法在 O(N³) 时间、O(N²) 内存内精确求解,彻底消除 N! 依赖
- 最简方案:随机排列采样仅需替换
speaker_permutations的生成方式,1 行改动即可启用 - 长期方向:Sinkhorn 可微排列提供端到端梯度流,适合与自监督 / 联合训练结合