# 降低 PIT/ATS 排列复杂度的方案 > 当前截断排列(truncate permutation)方案直接将后 `N - max_perm_spks` 个说话人排除在排列优化之外,导致其标签顺序完全随机,损害模型的排列不变性。本文从算法与工程两个维度,提出多种替代方案。 --- ## 一、问题重述 **核心矛盾**: ```python # 全排列方案(不可行) 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 代码实现 ```python 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`,新增基于匈牙利算法的路径: ```python def get_pil_targets(labels, preds, speaker_permutations, use_hungarian: bool = False): if use_hungarian: return _get_pil_targets_hungarian(labels, preds) # ... 原有全排列逻辑(向后兼容) ``` 配置入口(YAML): ```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 算法](https://github.com/ivan-chai/hungarian-algorithm-pytorch) 或自行实现基于 `torch` 的匈牙利算法 - 训练时建议使用 `torch.no_grad()` 包裹排列搜索部分(排列本身不需要梯度),梯度通过重排后的 labels 流回 BCE loss --- ## 三、方案二:Stochastic Permutation Sampling(随机排列采样) ### 3.1 原理 不枚举所有 N! 种排列,而是**随机采样 K 种排列**(K << N!),从中选最佳匹配。这在统计上是无偏的,且随着 K 增大趋近于全排列结果。 ### 3.2 实现 ```python 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 原理 分两阶段工作: 1. **粗匹配(coarse matching)**:用 N×N score matrix(同方案一)做一次贪心匹配,得到初始排列 2. **精调(local refinement)**:在初始排列的邻域内做有限的全排列搜索(如交换 2-3 个说话人) ### 4.2 实现 ```python 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 实现 ```python 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 实现草图 ```python 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 时启用匈牙利算法: ```python 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 ``` --- ## 八、总结 1. **根本原因**:排列爆炸源于代码在 T(时间帧)维度上扩展排列,而实际上 PIT 匹配分数只需 N×N 的 pairwise 统计量 2. **最佳方案**:匈牙利算法在 O(N³) 时间、O(N²) 内存内精确求解,彻底消除 N! 依赖 3. **最简方案**:随机排列采样仅需替换 `speaker_permutations` 的生成方式,1 行改动即可启用 4. **长期方向**:Sinkhorn 可微排列提供端到端梯度流,适合与自监督 / 联合训练结合