Spatial-BEATs / docs /0429_v11a_with_dynamic.md
dieKarotte's picture
Add files using upload-large-folder tool
bf04039 verified
|
Raw
History Blame Contribute Delete
22.5 kB

0429 · v11a_with_dynamic_10hz —— 动态 DOA 监督 + 真实/QA 数据扩展

记录当前 ov1_local_spatial_v11a_with_dynamic_10hz preset 的完整链路: 数据 → loader → 模型 → loss;并标明相对于 v11a_real_balanced_10hz 的每一处差异。 对应代码入口:

  • Preset: train_spatial_beats.py::make_ov1_local_spatial_v11a_with_dynamic_10hz_config
  • Run script: run_ov1_v11a_with_dynamic_10hz.sh
  • DCASE 转换器: tools/dcase_starss_to_jsonl.py

1. 一句话定位

v11a_with_dynamic_10hz = v11a_real_balanced_10hz 的训练数据扩展 + loader/loss 升级到逐帧 target。

  • 模型结构不变:仍然是 local_spatial_track + SourceQueryDecoder (K=4) + spatial_head_demixer (v11a 的新组件),10 Hz token rate,ov123 top4 目录。
  • 预测侧不变FrameTrackPredictionOutput 仍然是 [B, K, T_s, ...] 的逐帧四元组 (activity, class, direction, distance)
  • 监督侧升级:target 张量从 [B, N_gt] 扩展为 [B, N_gt, T_s]——静态源沿 T_s 轴广播 (和旧行为一致),动态源按每帧轨迹线性插值到 10 Hz 栅格。
  • 数据扩展:新增 5 个训练 manifest(qa_moving / qa_counting / qa_lr_pair / qa_same_doa / dcase_starss_foa.train)和 1 个验证 manifest(dcase_starss_foa.valid)。
  • Hot-start:默认从 v11a_real_balanced_10hz/03_ov123_top4/best.pt 继续训练,strict=False 且不继承 optimizer/epoch/best。上游 ov123 静态 clip 的 epoch 0 loss 应与 v11a 吻合 (逐帧 target 对静态源退化为标量广播)。

2. 数据:新增的 manifest 与样本格式

2.1 训练集组成

manifest 记录数 类型 DOA 来源 distance 有效 复制次数
ov1_foa.jsonl (sim) ov1 静态 scalar 1
ov2_foa.jsonl (sim) ov2 静态 scalar 3
ov3_foa.jsonl (sim) ov3 静态 scalar 3
ov1_real_static_foa_mapped.jsonl ov1 real 静态 scalar ✗ (null) 4
ov2_real_static_foa_mapped.jsonl ov2 real 静态 scalar 8
ov3_real_static_foa_mapped.jsonl ov3 real 静态 scalar 8
qa_moving.jsonl 19 597 QA sim 动态(单源平滑轨迹) frames[] per-frame 2
qa_counting.jsonl 2 428 QA sim 静态(多源 2-5 个) scalar 1
qa_lr_pair.jsonl 6 631 QA sim 静态(左右成对) scalar 1
qa_same_doa.jsonl 7 896 QA sim 静态(同一方向多源) scalar 1
dcase_starss_foa.train.jsonl 12 805 DCASE 真录 20s 多源动态 frames[] per-frame ✗ (-1) 2

对应的 train_manifest_replication = (1, 3, 3, 4, 8, 8, 2, 1, 1, 1, 2)

  • 总训练 clip 数(未复制)= ov123sim + ov123real + 5 × 新 manifest ≈ ov123 基础 + 49 357。
  • 验证集相比 v11a 增加 dcase_starss_foa.valid.jsonl (4 560 clips),作为真实录音的统一评估入口。

2.2 manifest schema(以 qa_moving / DCASE 为例)

qa_moving.jsonl(来自 build_qa_foa_moving.py 合成管道)

{
  "scene_id": "...",
  "output_foa_path": "/abs/.../foa.wav",
  "output_duration_seconds": 10.0,
  "sample_rate": 16000,
  "frame_rate": 10.0,
  "num_frames": 100,
  "sources": [
    {
      "source_index": 0,
      "is_moving": true,
      "mono_target_label": "speech",       // FSD50K 63-class 名字
      "active_time": [0.0, 10.0],
      "doa": null,                          // ← 动态源 scalar doa 为空
      "distance_cm": 150.0,                 // clip 级 fallback 距离
      "trajectory": "...",                  // sweep_arc / lshape / ...
      "frames": [
        {"frame_idx": 0, "doa": {"azimuth_deg": 175.6, "elevation_deg": -3.2},
         "distance_cm": 150.2},
        {"frame_idx": 1, "doa": {"azimuth_deg": 177.9, "elevation_deg": -3.1},
         "distance_cm": 150.1},
        ...
      ]
    }
  ]
}

dcase_starss_foa.{train,valid,test}.jsonl(由 tools/dcase_starss_to_jsonl.py 生成)

{
  "scene_id": "fold1_starss22__fold4_room10_mix001_0",
  "dataset_source": "dcase_starss",
  "split": "train",
  "output_foa_path": "/abs/.../foa.wav",
  "output_duration_seconds": 20.0,
  "sample_rate": 16000,
  "frame_rate": 10.0,
  "sources": [
    {
      "source_index": 0,
      "is_moving": true,
      "dcase_class_idx": 5,
      "dcase_source_idx": 3,
      "mono_target_label": "speech",         // 经 DCASE_TO_FSD50K 重映射
      "mono_primary_label": "male_speech",
      "active_time": [1.3, 4.8],
      "full_time": [0.0, 20.0],
      "doa": null,
      "distance_cm": -1,                     // DCASE 没有距离
      "distance_valid": false,
      "frames": [
        {"frame_idx": 13, "time_s": 1.3,
         "doa": {"azimuth_deg": -45.0, "elevation_deg": 10.0},
         "distance_cm": -1},
        ...
      ]
    },
    ...   // 可能 4+ 个 track,但逐帧同时 active 的 ≤ 4(K=4 由 matcher 保证)
  ]
}

2.3 生成 DCASE manifest 的一次性步骤

python tools/dcase_starss_to_jsonl.py \
    --dcase-root /apdcephfs_cq10/.../DCASE2024_seld_baseline/prepared_datasets/starss23_foa_plus_29cls_20s \
    --output /apdcephfs_cq10/.../data/metadata/dcase_starss_foa.jsonl \
    --per-split-output
  • 扫描 metadata_dev/<dataset>/<stem>.csv,每行 frame_idx, class_idx, source_idx, az_deg, el_deg, dist_cm
  • (class_idx, source_idx) 分 track,若相邻 labelled frame 间隔 > gap_split_frames (默认 50 帧 = 5s), 就把 track 切成多段 SourceEvent,避免在静默区间乱插值。
  • 类别空间压缩:DCASE 29 类 → FSD50K 63 类的语义最近邻映射 (DCASE_TO_FSD50K dict,见文件头部);碰不上 FSD50K 词表的 DCASE 类(如 unknown_*)整 track 丢弃。
  • 输出统计:18 061 CSV → train 12 805 / valid 4 560 / test 505。

2.4 FSD50K 63 类别名

qa_*/DCASE manifest 里可能出现细粒度标签(male_singing / female_singing), 但 v11a 的 vocab 只有压缩后的 63 类(含 singing 不含性别变体)。在 spatial_dataset.py_resolve_class_index / _resolve_class_label 之前先跑一次 _LABEL_ALIASES.get(raw, raw) 归一化:

_LABEL_ALIASES = {
    "male_singing": "singing",
    "female_singing": "singing",
}

这样不用重新生成 jsonl,就能把 508 条 male_singing + 527 条 female_singing 折到 singing 上。


3. Loader:spatial_dataset.py 的逐帧化改造

3.1 SourceEvent 新增 5 个可选字段

@dataclass
class SourceEvent:
    class_index: int
    class_label: str
    azimuth_deg: float        # 静态 scalar;动态时是 frames[0] 的 fallback
    elevation_deg: float
    distance: float           # 动态时是第一个 valid frame 的 fallback
    distance_valid: bool
    start_time_seconds: float
    end_time_seconds: float
    # ---- 动态轨迹(仅动态源设置)----
    frame_times_s: Optional[Tensor] = None      # [N_f] 秒,相对 clip 起点
    frame_azi_deg: Optional[Tensor] = None      # [N_f] 度,未 unwrap
    frame_ele_deg: Optional[Tensor] = None      # [N_f] 度
    frame_distance_m: Optional[Tensor] = None   # [N_f] 米
    frame_distance_valid: Optional[Tensor] = None  # [N_f] bool

3.2 _parse_frame_trajectory

从 manifest 的 frames[] 里抽出 5 个 1D tensor,同时支持两种 layout:

  1. qa_moving:每帧带 frame_idx(不带 time_s),clip 级 frame_rate 用于换算 time_s = frame_idx / frame_rate
  2. DCASE 转换器输出:每帧直接给 time_s,跳过 frame_idx / frame_rate 换算。

距离单位处理:优先读 distance_cm(除以 100 得米),-1 或缺失标记为 distance_valid=False; 其次读 distance_m

3.3 _build_source_event_from_nested_entry 的 fallback

  • 动态源 top-level doa 通常是 null,所以把 _get_float 换成 _maybe_get_float, 再用 frames[0] 的 DOA 补 azimuth_deg / elevation_deg(scalar fallback,只有在 loss 层碰到静态路径时才会用到)。
  • 距离同理:若 source-level distance_valid=False,但 frames[] 里有至少一个 distance_cm >= 0, 就用第一个 valid frame 的距离作为 scalar fallback;否则保留 distance_valid=False

3.4 _maybe_crop_sample 的轨迹裁剪

随机/中心裁剪时,除了裁 waveform 和更新 start/end_time_seconds,还要:

  • new_start/new_end 窗口过滤 frame_times_s,把留下来的帧时间重置到新 clip 起点(- crop_start_seconds)。
  • frame_azi_deg / frame_ele_deg / frame_distance_m / frame_distance_valid 一起按索引截断。

保证裁剪后的 SourceEvent 时间轴仍和 waveform 同源。

3.5 Collate:[B, N_gt, T_s] 逐帧 target

collate_spatial_batch 相对旧实现的关键变化:

t_s_max = int(target_num_steps.max())          # batch 内最大 token 数
source_azimuth_deg   = zeros(B, N_gt_max, t_s_max)   # 原来是 (B, N_gt_max)
source_elevation_deg = zeros(B, N_gt_max, t_s_max)
source_distance      = zeros(B, N_gt_max, t_s_max)
source_distance_valid = ones (B, N_gt_max, t_s_max, dtype=bool)  # 默认 True

for b, sample in enumerate(samples):
    t_axis = arange(t_s_i) / target_token_rate        # 该 sample 的有效时间轴
    for s, source in enumerate(sample.sources):
        azi_row, ele_row, dist_row, dist_valid_row = _build_per_frame_targets(
            source=source, t_axis=t_axis, t_s_max=t_s_max,
        )
        source_azimuth_deg[b, s]   = azi_row
        source_elevation_deg[b, s] = ele_row
        source_distance[b, s]      = dist_row
        source_distance_valid[b, s] = dist_valid_row

_build_per_frame_targets 的两条路径:

  • 静态源frame_times_s is None):在 [0:t_s_i) 填入 scalar;[t_s_i:t_s_max) 填零(padding)。 行为等价于旧版广播。
  • 动态源:对 t_axis 做线性插值。方位角先用 _unwrap_deg 去掉 ±180° 的跳变 (qa_moving / DCASE 都可能有 170° → -170° 这样跨接的情况),插值后再 wrap 回 [-180, 180]; elevation / distance 直接线性插值;distance_valid两端都 valid 才 valid的逻辑 (_linear_interp_valid_mask),避免在未知距离段里猜出假的 valid。

3.6 SpatialBatch 的契约变化

@dataclass
class SpatialBatch:
    ...
    source_azimuth_deg:     Tensor   # [B, N_gt_max, T_s_max]     ← 原 [B, N_gt_max]
    source_elevation_deg:   Tensor   # [B, N_gt_max, T_s_max]
    source_distance:        Tensor   # [B, N_gt_max, T_s_max]
    source_distance_valid:  Tensor   # [B, N_gt_max, T_s_max]  新字段
    source_class_indices:   Tensor   # [B, N_gt_max]  (class 仍是 clip 级)
    source_start_time_seconds: Tensor  # [B, N_gt_max]
    source_end_time_seconds:   Tensor  # [B, N_gt_max]
    source_valid_mask:         Tensor  # [B, N_gt_max]

source_class_indices 保持 clip 级:v11a 没有「同一 source 换类」的需求,且对应 track 内 class 恒定。


4. 模型:和 v11a 完全一致

没改,为了让 hot-start 生效。这里简要记录一下 v11a 已有的配置,便于对照:

                        FOA 4-ch waveform @ 16 kHz
                                  │
                                  ▼
                     SpatialBEATsPreprocessor  (mel, iv feat)
                                  │
              ┌───────────────────┴──────────────────┐
              ▼                                      ▼
     SpatialPatchEmbedding                SpatialDeltaPatchAdapter
     (mel → 768-d patch tokens)           (+IV residual contribution)
              └───────────────┬──────────────────────┘
                              ▼
                BEATs TransformerEncoder (12 层,冻结)
                              │
                              ▼
               LocalSpatialEncoder  (IV-aware conv over (T_p, F_p))
                              │
                              ▼
            TemporalResampler → fused_spatial_embeddings [B, T_s, 768]
                   (T_s @ 10 Hz ,cfg.target_token_rate=10)
                              │
                              ▼
      SourceQueryDecoder  (K=4 queries × T_s 次 decode,两段式)
        • track_latents:  [B, K, D]
        • track_time_feat:[B, K, T_s, D]
                              │
                              ▼
   FrameTrackHeads (+ SpatialHeadDemixer 1 层 attn refine, heads=8)
        • pred_activity:          [B, K, T_s]
        • pred_class_logits:      [B, K, T_s, 63]
        • pred_direction:         [B, K, T_s, 3]  L2-normed
        • pred_distance:          [B, K, T_s]     softplus 米
        • pred_num_active_logits: [B, T_s, K+1]   (v10 num_active head)

v11a 相对 v9 的关键新组件(均保留):

  • use_spatial_head_demixer=True(1 层 self-attn,8 heads,dropout 0.1)—— 在 FrameTrack head 输出后做一次 track 维解相关。
  • local_spatial_lr_scale=1.0 —— LocalSpatialEncoder 和 head 用相同 LR(v9 默认 0.3 偏低)。

5. Loss:逐帧 target + distance valid mask

入口 compute_frame_track_losses(prediction_output, batch, temporal_padding_mask, config)

5.1 target 抽取

targets = _frame_source_target_tensors(batch, t_s_max, device)
# 返回:
#   window_mask:           [B, N_gt, T_s]  (active_time 内为 True)
#   source_valid:          [B, N_gt]
#   source_class:          [B, N_gt]
#   source_direction:      [B, N_gt, T_s, 3]  ← 逐帧 unit vector
#   source_distance:       [B, N_gt, T_s]     ← 逐帧米
#   source_distance_valid: [B, N_gt, T_s]     ← 逐帧 bool

对于来自 loader 的 source_azimuth_deg / source_elevation_deg_align_t 处理长度不匹配: batch 内 t_s_max 可能与 loader 构造时的大小不同(不同 DataLoader 的 collate 边界), 短则 pad 末帧的值,长则截断。

5.2 Hungarian 匹配(代价按每帧)

_match_frame_tracksper_framesegment 两种策略,preset 里走 segment) 在 [B, N, K, T] 的代价张量上做匹配:

cost[b, n, k, t] =   class_cost_w * NLL(pred_class, target_class[b, n])
                   + dir_cost_w   * (1 - pred_direction[b,k,t] · target_direction[b,n,t])
                   + dist_cost_w  * |pred_distance[b,k,t] - target_distance[b,n,t]|
                   + (1 - σ(pred_activity[b,k,t]))         # include_activity_cost

关键点target_direction / target_distance[B, N, T, *] 广播到 [B, N, 1, T, *](之前是 clip 级标量), 代价按每帧独立累加,所以动态源在不同帧的 best-match track 可以不同。segment matching 额外加了一个 −2.0 的 continuity bonus,让同一 GT 在连续的相同 active-set segment 里尽量停留在同一 track。

5.3 监督张量的构建

matched_track: [B, N_gt, T_s] (匹配结果 k∈[0,K) 或 -1)
valid_assign = matched_track >= 0
idx_b, idx_gt, idx_t = valid_assign.nonzero(as_tuple=True)
idx_k = matched_track[idx_b, idx_gt, idx_t]

activity_target[idx_b, idx_k, idx_t] = 1.0
class_target    [idx_b, idx_k, idx_t] = targets["source_class"][idx_b, idx_gt]
direction_target[idx_b, idx_k, idx_t] = targets["source_direction"][idx_b, idx_gt, idx_t]   # ← 3D 索引
distance_target [idx_b, idx_k, idx_t] = targets["source_distance" ][idx_b, idx_gt, idx_t]   # ← 3D 索引
dist_supervise_mask[idx_b, idx_k, idx_t] = targets["source_distance_valid"][idx_b, idx_gt, idx_t]

相对旧版(targets["source_direction"][idx_b, idx_gt][M, 3])的差别是把第三个 axis 替换为具体 的 idx_t,真正拿到逐帧 GT。supervise_mask 是 activity-winning 的 mask,dist_supervise_mask 在其基础上再 AND 一个逐帧 distance validity:STARSS/DCASE 整源为 False 时, distance loss 就不会回传任何梯度。

5.4 各项损失

公式 mask
activity BCE_with_logits(pred_activity, activity_target, pos_weight=dyn) valid_time 扩到 [B, K, T_s]
num_active (v10) CE(pred_num_active_logits, active_count) valid_time
class CE(pred_class_logits, class_target) + 可选 ontology smoothing supervise_mask
direction mean(1 - pred · target) supervise_mask
distance smooth_l1(pred, target) dist_supervise_mask逐帧 validity

ADPIT duplicate & nonwinner soft activity(v9/v10 的两个辅助)同样全部走逐帧 source_direction[..., t]source_distance_valid[..., t] 索引;旧版的 batch.source_azimuth_deg[:, 0] 之类的 2D 访问被替换成 [:, 0, 0](共 44 处)以避免 shape 冲突。

最终汇总:

loss_total = λ_act  · loss_activity
           + λ_cls  · loss_class
           + λ_dir  · loss_direction
           + λ_dist · loss_distance
           + λ_na   · loss_num_active

λ 与 v11a 完全相同(在 SpatialLossConfig 里走 v10 phase-2 的基线数值)。

5.5 静态源的退化等价性

因为 loader 把静态源沿 T_s 轴广播,target_direction[b, gt, 0:T_s_i] 每一帧都一致, Hungarian 代价 (1 - pred·target) 和 per-clip 版本逐项相等;distance 同理。因此 ov123 sim/real clip 的 epoch 0 loss 与 v11a 数值吻合,这也是为什么可以直接从 v11a best.pt 热启动。


6. 训练配置(preset diff)

def make_ov1_local_spatial_v11a_with_dynamic_10hz_config(...):
    cfg = make_ov1_local_spatial_v11a_real_balanced_10hz_config(...)

    # —— 只改了数据和轮次 ——
    cfg.train_manifest_paths = (
        ov1_sim, ov2_sim, ov3_sim,
        ov1_real, ov2_real, ov3_real,
        qa_moving, qa_counting, qa_lr_pair, qa_same_doa,
        dcase_starss_train,
    )
    cfg.train_manifest_replication = (1, 3, 3, 4, 8, 8, 2, 1, 1, 1, 2)

    cfg.val_manifest_paths = (
        ov1_sim, ov2_sim, ov3_sim,
        ov1_real, ov2_real, ov3_real,
        dcase_starss_valid,
    )
    cfg.test_manifest_paths = cfg.val_manifest_paths

    cfg.num_epochs = 15
    cfg.output_dir = "checkpoints/spatial_beats_ov1_local_spatial_v11a_with_dynamic_10hz_exp/03_ov123_top4"
    return cfg

运行入口(run_ov1_v11a_with_dynamic_10hz.sh)默认:

GPUS=8  BATCH_SIZE=4  NUM_WORKERS=8
SPATIAL_EPOCHS=15  SPATIAL_LR=1.5e-5  AMP=fp32
RESUME_CKPT=checkpoints/spatial_beats_ov1_local_spatial_v11a_real_balanced_10hz_exp/03_ov123_top4/best.pt
  --no-resume-optimizer --reset-epoch-on-resume --reset-best-on-resume

7. 与 v11a_real_balanced_10hz 的差异一览

维度 v11a_real_balanced_10hz v11a_with_dynamic_10hz
训练 manifest 数 6 (ov123 sim + ov123 real) 11(+5 个 QA/DCASE)
训练 clip 量(未复制) ~O(20k) +49 357
真实录音监督 ov123 real 静态 scalar +DCASE STARSS 20s 逐帧
动态源监督 qa_moving + DCASE 动态 track
SourceEvent 字段 frame_* +5 个 frame_times_s/azi/ele/dist/distance_valid
Loader target shape source_*: [B, N_gt] **source_*: [B, N_gt, T_s]**(静态广播,动态插值)
source_distance_valid [B, N_gt] **[B, N_gt, T_s]**(逐帧)
Hungarian 代价 dir/dist 代价用 clip 级标量 target 用逐帧 [B, N, T, *] target 广播
distance 监督 mask 按源 distance_valid **按 (源, 帧) dist_supervise_mask**,STARSS/DCASE 不回传梯度
Class 词表别名 _LABEL_ALIASES 处理 male_singing/female_singing → singing
验证集 ov123 sim+real + DCASE valid (4 560 clips)
num_epochs 15 (继承 v9) 15(不变)
模型结构 local_spatial_track + K=4 + demixer 完全一致
热启动 v9_real_balanced_10hz best.pt v11a_real_balanced_10hz best.pt, strict=False
输出目录 .../v11a_real_balanced_10hz_exp/03_ov123_top4 .../v11a_with_dynamic_10hz_exp/03_ov123_top4

8. 已知坑 & 兼容性备忘

  1. pred_* 的 t_s 与 loader 的 T_s_max 有可能不一致。loss 侧 _align_t 会 pad/truncate GT 的最后一维; loader 侧已保证 target_num_steps = round(duration × target_token_rate) 与模型 temporal resampler 对齐, 正常只有在 batch 内最长样本决定 t_s_max 而 frame-track head 输出更短时需要截断。
  2. 方位角跨 ±180°:qa_moving 里观测到 175.6° → -177.3° 的自然轨迹, _unwrap_deg 会把它解包成 175.6° → 182.7° 插值,最后 wrap 回 [-180, 180],不会出现 -340° 的错误跨越。
  3. DCASE 一个 clip 可以出现 6+ 个 track,但任意帧同时 active 的 ≤ 4(DCASE 规范约束); _match_frame_tracks_per_frameactive_count.clamp(max=K) 是一层防御。
  4. distance=-1 的 clip 在 loss 内走 dist_supervise_mask 全 False 的分支: loss_distance = pred_distance.sum() * 0.0,梯度为 0,这一 clip 只贡献 activity/class/direction loss。
  5. male_singing/female_singing 是 loader 侧别名,若以后又有新细粒度标签冒出来,直接在 spatial_dataset.py_LABEL_ALIASES 里加映射即可,不需要重跑数据。
  6. ov123 静态 clip 的 epoch 0 loss 必须与 v11a 一致——这是验证 loader/loss 升级没有引入 回归的烟雾测试;曾经跑过 smoke script 确认 variance=0 for 静态 target、variance>0 for qa_moving target。

9. 相关脚本索引

  • tools/dcase_starss_to_jsonl.py —— DCASE CSV → jsonl + FSD50K 类别映射
  • run_ov1_v11a_with_dynamic_10hz.sh —— 训练入口(含 6 个动态 manifest 存在性 warning)
  • spatial_dataset.py::_parse_frame_trajectory / _build_per_frame_targets —— 动态 target 构建
  • spatial_loss.py::_frame_source_target_tensors / compute_frame_track_losses —— 逐帧 loss
  • Preset: train_spatial_beats.py::make_ov1_local_spatial_v11a_with_dynamic_10hz_config (CLI 名: --preset ov1_local_spatial_v11a_with_dynamic_10hz)