finetune_moss-sortformer / docs /reduce_permutation_complexity.md
czyhust's picture
Full upload
74e4281 verified
|
Raw
History Blame Contribute Delete
16.6 kB

降低 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_targetsget_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 原理

分两阶段工作:

  1. 粗匹配(coarse matching):用 N×N score matrix(同方案一)做一次贪心匹配,得到初始排列
  2. 精调(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

八、总结

  1. 根本原因:排列爆炸源于代码在 T(时间帧)维度上扩展排列,而实际上 PIT 匹配分数只需 N×N 的 pairwise 统计量
  2. 最佳方案:匈牙利算法在 O(N³) 时间、O(N²) 内存内精确求解,彻底消除 N! 依赖
  3. 最简方案:随机排列采样仅需替换 speaker_permutations 的生成方式,1 行改动即可启用
  4. 长期方向:Sinkhorn 可微排列提供端到端梯度流,适合与自监督 / 联合训练结合