diff --git a/dependency/Megatron-LM-core_v0.16.1/megatron/core/transformer/attention.py b/dependency/Megatron-LM-core_v0.16.1/megatron/core/transformer/attention.py index 2200b558..a10ff8e3 100644 --- a/dependency/Megatron-LM-core_v0.16.1/megatron/core/transformer/attention.py +++ b/dependency/Megatron-LM-core_v0.16.1/megatron/core/transformer/attention.py @@ -1069,6 +1069,23 @@ def forward( value = value.squeeze(1) nvtx_range_pop(suffix="adjust_key_value") + # ##### [PS-diag] OFF pre-RoPE Q/K/V dump(侵入式,post-squeeze、pre-RoPE)##### + import os as _ps_os + if _ps_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: + try: + from prefix_sharing.tools.diagnostic_dump_verl080 import ( + dump_rope_preqk_verl080, dump_build_kv_input_v_on, dump_hidden_states_on, + ) + dump_rope_preqk_verl080(self.layer_number, query, key, + self.config.num_layers) + dump_build_kv_input_v_on(self.layer_number, value, + self.config.num_layers) + dump_hidden_states_on(self.layer_number, hidden_states, + self.config.num_layers) + except Exception as exc: + print(f"[PS-diag] OFF preqkv dump failed: {exc}", flush=True) + # ##### [PS-diag] OFF pre-RoPE Q/K/V dump end ##### + # ================================================ # relative positional embedding (rotary embedding) # ================================================ @@ -1126,6 +1143,20 @@ def forward( # value_layer = apply_rotary_pos_emb(value_layer, k_pos_emb) nvtx_range_pop(suffix="rotary_pos_emb") + # ##### [PS-diag] OFF post-RoPE Q/K + full_kv dump(侵入式,rotary block 之后)##### + if _ps_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: + try: + from prefix_sharing.tools.diagnostic_dump_verl080 import ( + dump_rope_postqk_verl080, dump_full_kv_off, + ) + dump_rope_postqk_verl080(self.layer_number, query, key, + self.config.num_layers) + dump_full_kv_off(self.layer_number, key, value, + self.config.num_layers) + except Exception as exc: + print(f"[PS-diag] OFF postqk/full_kv dump failed: {exc}", flush=True) + # ##### [PS-diag] OFF post-RoPE Q/K + full_kv dump end ##### + # ================================== # core attention computation # ================================== diff --git a/dependency/verl_cdd9014f/verl/trainer/ppo/ray_trainer.py b/dependency/verl_cdd9014f/verl/trainer/ppo/ray_trainer.py index a645779d..81244666 100644 --- a/dependency/verl_cdd9014f/verl/trainer/ppo/ray_trainer.py +++ b/dependency/verl_cdd9014f/verl/trainer/ppo/ray_trainer.py @@ -1354,8 +1354,26 @@ def fit(self): max_prompt_length=self.config.data.max_prompt_length, max_response_length=self.config.data.max_response_length, ) + # Inject baseline synthetic data when PREFIX_SHARING_BASELINE_SYNTHETIC env is set + baseline_json = os.environ.get("PREFIX_SHARING_BASELINE_SYNTHETIC", None) + if baseline_json: + from prefix_sharing.tools.inject_baseline_synthetic import patch_baseline_synthetic + + _num_seq = int(os.environ.get("PREFIX_SHARING_BASELINE_NUM_SEQ", "1")) + _stack = int(os.environ.get("PREFIX_SHARING_BASELINE_STACK", "1")) + _seed = int(os.environ.get("PREFIX_SHARING_BASELINE_SEED", "42")) + patch_baseline_synthetic( + self, + json_path=baseline_json, + max_prompt_length=self.config.data.max_prompt_length, + max_response_length=self.config.data.max_response_length, + num_seq=_num_seq, + stack=_stack, + seed=_seed, + ) #####prefix-sharing:inject data######## + current_epoch = self.global_steps // len(self.train_dataloader) # perform validation before training diff --git a/prefix-sharing/prefix_sharing/backends/factory.py b/prefix-sharing/prefix_sharing/backends/factory.py index a98ed65b..c761d3c2 100644 --- a/prefix-sharing/prefix_sharing/backends/factory.py +++ b/prefix-sharing/prefix_sharing/backends/factory.py @@ -20,19 +20,19 @@ def get_backend_instance( """ if backend is not None: return backend - + if config.backend == "torch_ref": from prefix_sharing.backends.torch_ref import TorchReferenceBackend return TorchReferenceBackend() - + if config.backend == "flash_atten_gpu": from prefix_sharing.backends.flash_atten_gpu import GpuFlashAttentionBackend return GpuFlashAttentionBackend() - + if config.backend == "flash_atten_npu": from prefix_sharing.backends.flash_atten_npu import NpuFlashAttentionBackend return NpuFlashAttentionBackend() - + raise ValueError( f"Unknown backend '{config.backend}'. " f"Supported: torch_ref, flash_atten_gpu, flash_atten_npu" diff --git a/prefix-sharing/prefix_sharing/core/prefix_detector.py b/prefix-sharing/prefix_sharing/core/prefix_detector.py index 3d40a37d..fdf9edc8 100644 --- a/prefix-sharing/prefix_sharing/core/prefix_detector.py +++ b/prefix-sharing/prefix_sharing/core/prefix_detector.py @@ -180,6 +180,20 @@ def detect(self, input_ids: Sequence[TokenSequence]) -> PrefixDetectionResult: group_key_to_id: dict[tuple[int, int], int] = {} group_members: dict[int, list[int]] = {} + # 诊断开关:跳过 trie 检测,直接返回 0-prefix(所有序列当 provider,无复用)。 + # 用法见 PREFIX_SHARING_FORCE_ZERO_PREFIX 说明(配合 build 层的 has_sharing 旁路)。 + import os as _os + if _os.environ.get("PREFIX_SHARING_FORCE_ZERO_PREFIX"): + return PrefixDetectionResult( + batch_size=batch_size, + reuse_specs=(), + groups=(), + group_ids=tuple(group_ids), + provider_index=tuple(provider_index), + prefix_lens=tuple(prefix_lens), + is_provider=tuple(is_provider), + ) + for index, seq in enumerate(input_ids): node = root matched = 0 diff --git a/prefix-sharing/prefix_sharing/integrations/megatron_runtime.py b/prefix-sharing/prefix_sharing/integrations/megatron_runtime.py index 01ddb36d..bbbf06ca 100644 --- a/prefix-sharing/prefix_sharing/integrations/megatron_runtime.py +++ b/prefix-sharing/prefix_sharing/integrations/megatron_runtime.py @@ -106,6 +106,15 @@ def prefix_attention( f"built expanded kv: expanded_key_shape={tuple(expanded_key.shape)}, expanded_value_shape={tuple(expanded_value.shape)}" ) + ##### [PS-diag] ON expanded K/V dump(build_kv 输出)##### + try: + from prefix_sharing.tools.diagnostic_dump_verl080 import dump_expanded_kv_on + dump_expanded_kv_on(layer_id, expanded_key, expanded_value, + attention_module.config.num_layers) + except Exception as exc: + print(f"expanded_kv dump failed: {exc}", flush=True) + ##### [PS-diag] ON expanded K/V dump end ##### + # 注意力计算 core_attn_out = attention_backend.attention( query, @@ -119,15 +128,15 @@ def prefix_attention( core_attn_out = core_attn_out.reshape(core_attn_out.size(0), 1, -1) output = attention_module.linear_proj(core_attn_out) # (tensor, bias) tuple - ######### prefix-sharing diag: ON attention_output (per-layer) ######### + ##### [PS-diag] ON attention_output (per-layer) ##### try: from prefix_sharing.tools.diagnostic_dump import dump_attn_on dump_attn_on(output[0], packed_seq_params, prefix_sharing_context.prefix_sharing_plan, attention_module.layer_number, attention_module.config.num_layers) - except Exception as e: - print(f"last-attn dump (ON) failed: {e}") - ######### prefix-sharing diag: ON attention_output (per-layer) ######### + except Exception as exc: + print(f"last-attn dump (ON) failed: {exc}") + ##### [PS-diag] ON attention_output (per-layer) end ##### # --- return output @@ -173,6 +182,32 @@ def _apply_positioned_rope( if q_pos_emb is not None and max_needed > q_pos_emb.shape[0]: dim_half = q_pos_emb.shape[-1] // 2 step = q_pos_emb[1:2, :, :, :dim_half] - q_pos_emb[0:1, :, :, :dim_half] + + # ── [PS-diag] RoPE extrapolation: 验证相邻位置 step 是否恒定 ── + _layer = getattr(attention_module, 'layer_number', -1) + _num_extra = int(max_needed - q_pos_emb.shape[0]) + _step_vals = step.detach().flatten() + # 关键诊断: step01 == step12 ? (pos_emb[p] 对 p 是否线性) + _step01_vs_step12_diff = None + if q_pos_emb.shape[0] >= 3: + _step12 = q_pos_emb[2:3, :, :, :dim_half] - q_pos_emb[1:2, :, :, :dim_half] + _step12_vals = _step12.detach().flatten() + _diff = (_step_vals - _step12_vals).abs() + _step01_vs_step12_diff = _diff.max().item() + _is_linear = (_step01_vs_step12_diff is not None and _step01_vs_step12_diff < 1e-8) + print( + f"[PS][RoPE-extrapolate-Q] layer={_layer} " + f"max_needed={max_needed} precomputed={q_pos_emb.shape[0]} extra={_num_extra} " + f"step_min={_step_vals.min().item():.8f} step_max={_step_vals.max().item():.8f} " + f"step_mean={_step_vals.mean().item():.8f} step_std={_step_vals.std().item():.8f} " + f"step01_vs_step12_maxdiff={_step01_vs_step12_diff} is_linear={_is_linear}", + flush=True, + ) + # ── [PS-diag] end ── + + # RoPE 具有线性性质:freqs[p] = p * inv_freq。 + # 因此可以通过 pos_emb[1] - pos_emb[0] 恢复出 step(即 inv_freq), + # 从而生成缺失的高位置频率向量。 extra_positions = torch.arange( q_pos_emb.shape[0], max_needed, device=q_pos_emb.device, dtype=q_pos_emb.dtype, @@ -183,6 +218,28 @@ def _apply_positioned_rope( if k_pos_emb is not None and max_needed > k_pos_emb.shape[0]: dim_half = k_pos_emb.shape[-1] // 2 step = k_pos_emb[1:2, :, :, :dim_half] - k_pos_emb[0:1, :, :, :dim_half] + + # ── [PS-diag] RoPE extrapolation: 验证相邻位置 step 是否恒定 ── + _layer = getattr(attention_module, 'layer_number', -1) + _num_extra = int(max_needed - k_pos_emb.shape[0]) + _step_vals = step.detach().flatten() + _step01_vs_step12_diff = None + if k_pos_emb.shape[0] >= 3: + _step12 = k_pos_emb[2:3, :, :, :dim_half] - k_pos_emb[1:2, :, :, :dim_half] + _step12_vals = _step12.detach().flatten() + _diff = (_step_vals - _step12_vals).abs() + _step01_vs_step12_diff = _diff.max().item() + _is_linear = (_step01_vs_step12_diff is not None and _step01_vs_step12_diff < 1e-8) + print( + f"[PS][RoPE-extrapolate-K] layer={_layer} " + f"max_needed={max_needed} precomputed={k_pos_emb.shape[0]} extra={_num_extra} " + f"step_min={_step_vals.min().item():.8f} step_max={_step_vals.max().item():.8f} " + f"step_mean={_step_vals.mean().item():.8f} step_std={_step_vals.std().item():.8f} " + f"step01_vs_step12_maxdiff={_step01_vs_step12_diff} is_linear={_is_linear}", + flush=True, + ) + # ── [PS-diag] end ── + extra_positions = torch.arange( k_pos_emb.shape[0], max_needed, device=k_pos_emb.device, dtype=k_pos_emb.dtype, @@ -191,6 +248,54 @@ def _apply_positioned_rope( extra_emb = torch.cat([extra_angles, extra_angles], dim=-1) k_pos_emb = torch.cat([k_pos_emb, extra_emb], dim=0) + # ── [PS-diag] RoPE ground-truth probe ── + # 当 PREFIX_SHARING_DIAG_ROPE_GROUND_TRUTH=1 时,从 inv_freq 直接计算完整 + # 频率表(跳过线性外推),用于验证外推是否引入数值偏差。 + # 如果启用此开关后 ON vs OFF 结果一致,RoPE 外推就是根因。 + import os as _os_gt + if _os_gt.environ.get("PREFIX_SHARING_DIAG_ROPE_GROUND_TRUTH"): + _layer_gt = getattr(attention_module, 'layer_number', -1) + # 尝试从 attention_module 获取 inv_freq + _inv_freq = None + _rotary_emb = getattr(attention_module, 'rotary_pos_emb', None) + if _rotary_emb is not None: + _inv_freq = getattr(_rotary_emb, 'inv_freq', None) + if _inv_freq is None: + # fallback: 从 config 读取 RoPE 参数 + _cfg = attention_module.config + _dim = getattr(_cfg, 'hidden_size', 4096) // getattr(_cfg, 'num_attention_heads', 32) + _base = getattr(_cfg, 'rope_theta', 10000.0) + _inv_freq = 1.0 / (_base ** (torch.arange( + 0, _dim, 2, device=q_pos_emb.device if q_pos_emb is not None + else k_pos_emb.device).float() / _dim)) + + _device = q_pos_emb.device if q_pos_emb is not None else k_pos_emb.device + _dtype = q_pos_emb.dtype if q_pos_emb is not None else k_pos_emb.dtype + _all_positions = torch.arange(0, max_needed, device=_device, dtype=torch.float) + _freqs = torch.outer(_all_positions, _inv_freq.to(_device).float()) # [max_needed, dim/2] + _emb_gt = torch.cat([_freqs, _freqs], dim=-1) # [max_needed, dim] + # reshape 匹配 pos_emb 维度 [max_needed, 1, 1, dim] + _emb_gt = _emb_gt.unsqueeze(1).unsqueeze(1).to(_dtype) + + _extra_old = max(0, int(max_needed - ( + q_pos_emb.shape[0] if q_pos_emb is not None else max_needed))) + _is_q_truncated = q_pos_emb is not None and q_pos_emb.shape[0] < max_needed + _is_k_truncated = k_pos_emb is not None and k_pos_emb.shape[0] < max_needed + + if q_pos_emb is not None: + q_pos_emb = _emb_gt + if k_pos_emb is not None: + k_pos_emb = _emb_gt + + print( + f"[PS][RoPE-ground-truth] layer={_layer_gt} " + f"max_needed={max_needed} emb_shape={_emb_gt.shape} " + f"was_extrapolated_q={_is_q_truncated} was_extrapolated_k={_is_k_truncated} " + f"old_extra_count={_extra_old}", + flush=True, + ) + # ── [PS-diag] end ── + # Build kwargs for apply_rotary_pos_emb. # Only include version-specific params when they're provided, # to maintain backward compat with v070 (mcore <= 0.15.x). @@ -210,14 +315,14 @@ def _rope_kwargs(_unused_cu_seqlens: Any | None) -> dict[str, Any]: if q_pos_emb is not None: q_freqs = q_pos_emb.index_select(0, positions) - ######### prefix-sharing diag: ON rope_freqs (per-layer) ######### + ##### [PS-diag] ON rope_freqs (per-layer) ##### try: - from prefix_sharing.tools.diagnostic_dump import dump_rope_freqs_on - dump_rope_freqs_on(q_freqs, attention_module.layer_number, + from prefix_sharing.tools.diagnostic_dump import dump_rope_freqs + dump_rope_freqs(q_freqs, attention_module.layer_number, attention_module.config.num_layers) - except Exception as e: - print(f"rope_freqs_on dump failed: {e}") - ######### prefix-sharing diag: ON rope_freqs (per-layer) ######### + except Exception as exc: + print(f"rope_freqs dump failed: {exc}") + ##### [PS-diag] ON rope_freqs (per-layer) end ##### query = apply_rotary_pos_emb( query.unsqueeze(1), q_freqs, @@ -230,6 +335,20 @@ def _rope_kwargs(_unused_cu_seqlens: Any | None) -> dict[str, Any]: k_freqs, **_rope_kwargs(cu_seqlens_kv), ).squeeze(1) + + ##### [PS-diag] ON post-RoPE Q/K dump (per-layer) ##### + try: + from prefix_sharing.tools.diagnostic_dump_verl080 import dump_rope_postqk_verl080 + dump_rope_postqk_verl080( + attention_module.layer_number, + query, key, + attention_module.config.num_layers, + positions=packed_position_ids, + ) + except Exception as exc: + print(f"rope_postqk_layer dump failed: {exc}") + ##### [PS-diag] ON post-RoPE Q/K dump end ##### + return query, key diff --git a/prefix-sharing/prefix_sharing/integrations/verl_mcore.py b/prefix-sharing/prefix_sharing/integrations/verl_mcore.py index f968d29e..4b4cb228 100644 --- a/prefix-sharing/prefix_sharing/integrations/verl_mcore.py +++ b/prefix-sharing/prefix_sharing/integrations/verl_mcore.py @@ -35,7 +35,6 @@ from prefix_sharing.integrations.parallel_info import get_megatron_parallel_info from prefix_sharing.integrations.patch_manager import PatchHandle - @dataclass(frozen=True) class PrefixSharingRuntimeState: prefix_sharing_plan: PrefixSharingPlan @@ -390,6 +389,16 @@ def restore_via_2d_unfold_verl080( if ctx is None: return output plan = ctx.prefix_sharing_plan + # ── [PS-diag] skip-restore probe ── + # PREFIX_SHARING_DIAG_SKIP_RESTORE=1 时跳过 restore,直接返回 forward 原始输出。 + # 用于隔离:如果 skip restore 后 suffix logprobs 与 OFF 一致,偏差在 restore; + # 如果 suffix logprobs 仍不一致,偏差在 forward(attention/build_kv)。 + import os as _os_skip + if _os_skip.environ.get("PREFIX_SHARING_DIAG_SKIP_RESTORE"): + print("[PS][diag] SKIP_RESTORE: 跳过 restore,直接返回 forward 原始输出") + return output + # ── [PS-diag] end ── + # Guard on reuser presence, not on prefix_last_restore_indices: a batch # whose reusers all have suffix_len == 0 has no prefix-last spec but still # needs interior prefix columns restored. @@ -517,7 +526,7 @@ def _fold_2d_to_nested(tensor_2d: Any, original_lengths: list[int]) -> Any: import torch rows = [tensor_2d[seq_idx, :original_lengths[seq_idx]] for seq_idx in range(len(original_lengths))] - return torch.nested.nested_tensor(rows, layout=torch.jagged) + return torch.nested.as_nested_tensor(rows, layout=torch.jagged) def _clone_batch(batch: Any) -> Any: @@ -636,7 +645,18 @@ def build_prefix_sharing_micro_batch_verl080( # ── 阶段 4: 前缀共享规划 ── plan = PrefixSharingPlanner(ps_config).plan(sequences) - if not plan.has_sharing: + + # 诊断开关 PREFIX_SHARING_FORCE_ZERO_PREFIX:detect() 被旁路返回 0-prefix,此处 + # plan.has_sharing=False。仍要让 ON pipeline 跑(prefix_attention 用 mode-3 跑全 + # 序列、无裁剪/注入),故旁路 has_sharing 早返回。用于隔离 "kernel-mode 差异": + # ON(0-prefix, mode 3, 全序列) vs OFF(mode 2, 全序列) + # ≈ → kernel mode 不是根因,偏差来自裁剪/注入;偏差大 → kernel mode 是根因。 + import os as _os_force_zero + _force_zero_prefix = bool(_os_force_zero.environ.get("PREFIX_SHARING_FORCE_ZERO_PREFIX")) + if _force_zero_prefix: + print("[PS][prepare] FORCE_ZERO_PREFIX: 跳过 has_sharing 早返回," + "ON pipeline 用 mode-3 跑全序列(无裁剪/注入)") + elif not plan.has_sharing: print("[PS][prepare] no prefix sharing detected") return batch, None @@ -702,13 +722,13 @@ def _trim_nested_batch(batch: Any, plan: PrefixSharingPlan) -> Any: # 裁剪 input_ids NestedTensor trimmed_ids_seqs = _slice_nested_sequences(input_ids, plan) - new_input_ids = torch.nested.nested_tensor(trimmed_ids_seqs, layout=torch.jagged) + new_input_ids = torch.nested.as_nested_tensor(trimmed_ids_seqs, layout=torch.jagged) trimmed_batch["input_ids"] = new_input_ids # 裁剪 position_ids NestedTensor if _is_nested_tensor(position_ids): trimmed_pos_seqs = _slice_nested_sequences(position_ids, plan) - new_position_ids = torch.nested.nested_tensor(trimmed_pos_seqs, layout=torch.jagged) + new_position_ids = torch.nested.as_nested_tensor(trimmed_pos_seqs, layout=torch.jagged) else: # position_ids 是 2D tensor → 需要用 attention_mask 的 # valid_indices 切片(keep_range 是序列偏移,不是列索引) @@ -728,14 +748,14 @@ def _trim_nested_batch(batch: Any, plan: PrefixSharingPlan) -> Any: trimmed_pos_seqs = _slice_2d_position_rows( position_ids, plan, attention_mask_bool, ) - new_position_ids = torch.nested.nested_tensor(trimmed_pos_seqs, layout=torch.jagged) + new_position_ids = torch.nested.as_nested_tensor(trimmed_pos_seqs, layout=torch.jagged) trimmed_batch["position_ids"] = new_position_ids # loss_mask 也需要裁剪(如果存在) loss_mask = batch.get("loss_mask") if loss_mask is not None and _is_nested_tensor(loss_mask): trimmed_loss_seqs = _slice_nested_sequences(loss_mask, plan) - trimmed_batch["loss_mask"] = torch.nested.nested_tensor( + trimmed_batch["loss_mask"] = torch.nested.as_nested_tensor( trimmed_loss_seqs, layout=torch.jagged ) @@ -789,8 +809,8 @@ def _trim_plain_batch_thd(batch: Any, plan: PrefixSharingPlan) -> Any: # 用裁剪后的序列构建 NestedTensor(jagged layout) # 这样 preprocess_thd_engine 会从 offsets 正确计算 cu_seqlens trimmed_batch = _clone_batch(batch) - trimmed_batch["input_ids"] = torch.nested.nested_tensor(kept_id_rows, layout=torch.jagged) - trimmed_batch["position_ids"] = torch.nested.nested_tensor(kept_pos_rows, layout=torch.jagged) + trimmed_batch["input_ids"] = torch.nested.as_nested_tensor(kept_id_rows, layout=torch.jagged) + trimmed_batch["position_ids"] = torch.nested.as_nested_tensor(kept_pos_rows, layout=torch.jagged) # loss_mask loss_mask = batch.get("loss_mask") @@ -801,7 +821,7 @@ def _trim_plain_batch_thd(batch: Any, plan: PrefixSharingPlan) -> Any: keep_start, keep_end = plan.input_keep_ranges[row] kept_indices = indices[keep_start:keep_end] kept_loss_rows.append(loss_mask[row, kept_indices]) - trimmed_batch["loss_mask"] = torch.nested.nested_tensor( + trimmed_batch["loss_mask"] = torch.nested.as_nested_tensor( kept_loss_rows, layout=torch.jagged ) diff --git a/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/attention.py b/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/attention.py index 97ef0c1e..2fb7fbb0 100644 --- a/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/attention.py +++ b/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/attention.py @@ -36,7 +36,11 @@ def patched_forward( ctx = current_prefix_sharing_context() if ctx is None: # ── normal path: 调用原始 forward ── - _result = original_forward( + # post-RoPE Q/K / full_kv / preqk 由 Megatron attention.py 侵入式 dump 写入; + # patch 层只负责 attn_outputs + rope_freqs(侵入式未覆盖的)。 + import os as _diag_os + diag_enabled = _diag_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None + forward_result = original_forward( self, hidden_states, attention_mask, @@ -51,34 +55,48 @@ def patched_forward( sequence_len_offset=sequence_len_offset, inference_params=inference_params, ) - # ##### [PS-diag] OFF attn_outputs + rope_freqs_off dump ##### - # OFF 走原始 forward,不经 prefix_attention/_apply_positioned_rope, - # 所以 ON 路径里的 dump_attn_on/dump_rope_freqs_on 不会触发。 - # 这里在 OFF 分支补 dump,让 cmp_diag 的 attn/RoPE 对比有 OFF ground truth。 - # v070 是直接改 megatron attention 源码在 forward 内部 dump;v080 用 patch - # wrapper 在 forward 返回后 dump output + 入参 rotary_pos_emb 解包出 angle table, - # 语义等价(唯一拿不到的是 rope_emb rotated q/k,在 forward 内部,但 rope_freqs - # angle table 已够验证 RoPE)。 - import os as _os - if _os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: + # ##### [PS-diag] OFF attn_outputs + rope_freqs dump ##### + if diag_enabled: + import torch # 仅用于构造 positions(freqs 切 per-token + Q/K debug) from prefix_sharing.tools.diagnostic_dump import ( - dump_attn_off, dump_rope_freqs_off, + dump_attn_off, dump_rope_freqs, ) + # rope_postqk / preqk / full_kv 已由 Megatron 侵入式 dump 覆盖 from prefix_sharing.integrations.megatron_runtime import _unpack_rotary_pos_emb - _attn_out = _result[0] if isinstance(_result, tuple) else _result - _bs = ( + attn_output = forward_result[0] if isinstance(forward_result, tuple) else forward_result + batch_size = ( len(packed_seq_params.cu_seqlens_q_padded) - 1 if (packed_seq_params is not None and hasattr(packed_seq_params, "cu_seqlens_q_padded")) else 0 ) - dump_attn_off(_attn_out, packed_seq_params, - self.layer_number, _bs, self.config.num_layers) + dump_attn_off(attn_output, packed_seq_params, + self.layer_number, batch_size, self.config.num_layers) if rotary_pos_emb is not None: - _q_pos_emb, _ = _unpack_rotary_pos_emb(rotary_pos_emb) - dump_rope_freqs_off(_q_pos_emb, self.layer_number, self.config.num_layers) - # ##### [PS-diag] OFF attn_outputs + rope_freqs_off dump end ##### - return _result + q_pos_emb, k_pos_emb = _unpack_rotary_pos_emb(rotary_pos_emb) + # OFF 标准 positions(每 segment 内 0..seg-1):切 per-token freqs + Q/K debug + per_token_positions = None + if (packed_seq_params is not None + and hasattr(packed_seq_params, "cu_seqlens_q_padded")): + cu_seqlens_tensor = packed_seq_params.cu_seqlens_q_padded + # device 必须显式到 cu_seqlens_tensor.device(GPU):arange 默认 CPU,否则后面 + # q_pos_emb.index_select(0, per_token_positions) 会 device 不匹配崩 forward。 + per_token_positions = torch.cat([ + torch.arange(int(cu_seqlens_tensor[i + 1] - cu_seqlens_tensor[i]), + device=cu_seqlens_tensor.device) + for i in range(len(cu_seqlens_tensor) - 1) + ]).long() + # rope_freqs:存 per-token 角度(与 ON 同款),统一 rope_freqs.pt + if per_token_positions is not None: + dump_rope_freqs( + q_pos_emb.index_select(0, per_token_positions), + self.layer_number, self.config.num_layers, + ) + # rope_postqk + full_kv 已由 megatron attention.py 侵入式 dump + # (rotary block 之后),不在 patch 层重复——避免 hook 二次 flush + # 覆盖侵入式已写好的完整 24 层文件。 + # ##### [PS-diag] OFF attn_outputs + rope_freqs dump end ##### + return forward_result # ── prefix-sharing path ── # phase 1: training, THD, no fusion, no output gate @@ -95,6 +113,33 @@ def patched_forward( key = key.squeeze(1) value = value.squeeze(1) + # ##### [PS-diag] ON pre-RoPE Q/K/V 统一 dump(get_qkv 之后、RoPE 之前)##### + # 全部在此点 dump(squeeze 后、_apply_positioned_rope / build_kv 之前), + # 与 OFF baseline(hook 在 get_qkv 输出处截)同口径,集中对比,避免分散。 + import os as _diag_os + if _diag_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: + from prefix_sharing.tools.diagnostic_dump_verl080 import ( + dump_rope_preqk_verl080, dump_build_kv_input_v_on, + ) + try: + dump_rope_preqk_verl080(self.layer_number, query, key, + self.config.num_layers) + except Exception as exc: + print(f"rope_preqk (pre-RoPE Q/K) dump failed: {exc}", flush=True) + try: + dump_build_kv_input_v_on(self.layer_number, value, + self.config.num_layers) + except Exception as exc: + print(f"build_kv_input_v (pre-RoPE V) dump failed: {exc}", flush=True) + # [PS-diag] dump hidden_states for input-level comparison + try: + from prefix_sharing.tools.diagnostic_dump_verl080 import dump_hidden_states_on + dump_hidden_states_on(self.layer_number, hidden_states, + self.config.num_layers) + except Exception as exc: + print(f"hidden_states dump failed: {exc}", flush=True) + # ##### [PS-diag] ON pre-RoPE Q/K/V dump end ##### + # delegate to verified integrations code from prefix_sharing.integrations.megatron_runtime import ( prefix_attention, diff --git a/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/forward_step.py b/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/forward_step.py index f6243a7c..863d847f 100644 --- a/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/forward_step.py +++ b/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/forward_step.py @@ -122,8 +122,8 @@ def patched_forward_step( # ON: prefix_lens / original_lengths 取自 plan; # OFF: prefix_lens 全0、original_lengths 从 input_ids NestedTensor offsets diff 推。 # cu_seqlens 取送进 forward 的 input_ids NestedTensor offsets(ON=裁剪后 packed 边界, OFF=完整)。 - import os as _os - if _os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: + import os as _diag_os + if _diag_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: from prefix_sharing.tools.diagnostic_dump_verl080 import ( dump_meta_verl080, dump_attention_mask_verl080, dump_label_mask_verl080, @@ -131,48 +131,49 @@ def patched_forward_step( nested_offsets_to_cu, ) from prefix_sharing.integrations.verl_mcore import _is_nested_tensor - _ids_nested = batch_for_forward["input_ids"] - if _is_nested_tensor(_ids_nested): + input_ids_nested = batch_for_forward["input_ids"] + if _is_nested_tensor(input_ids_nested): if ps_state is not None: - _plan = ps_state.prefix_sharing_plan - _prefix_lens = list(_plan.prefix_lens) - _orig_lens = list(_plan.original_lengths) + sharing_plan = ps_state.prefix_sharing_plan + prefix_lens_list = list(sharing_plan.prefix_lens) + original_lengths = list(sharing_plan.original_lengths) else: - _diffs = _ids_nested.offsets().diff().tolist() - _orig_lens = [int(d) for d in _diffs] - _prefix_lens = [0] * len(_orig_lens) - dump_meta_verl080(_prefix_lens, nested_offsets_to_cu(_ids_nested)) + seq_lengths = input_ids_nested.offsets().diff().tolist() + original_lengths = [int(length) for length in seq_lengths] + prefix_lens_list = [0] * len(original_lengths) + dump_meta_verl080(prefix_lens_list, + nested_offsets_to_cu(input_ids_nested)) # attention_mask + label_mask:两种 log_probs 对比范围,都不含越界预测位 # (POS L_i-1,其 logp 预测不存在的 token[L_i])。对齐 restore 后 log_probs # 的 [B, L_max] 紧凑坐标系。 # attention_mask:[0:L_i-1) prompt 区+prompt-last+response 区(整体 restore 验证) # label_mask:[prompt-last:L_i-1) prompt-last+response 区(PPO loss 范围) - _Lmax_lm = max(_orig_lens) if _orig_lens else 0 + max_seq_len = max(original_lengths) if original_lengths else 0 # tag 与 logprobs 一致(接入点2 用 model.training 区分 old/train), # 保证 mask 和 logprobs_{tag} 来自同一 forward(同 batch、同 L_max)。 - _tag_lm = "train" if model.training else "old" - # attention_mask 仅依赖 _orig_lens,不需要 loss_mask。 + diag_tag = "train" if model.training else "old" + # attention_mask 仅依赖 original_lengths,不需要 loss_mask。 dump_attention_mask_verl080( - build_attention_mask_2d(_orig_lens, _Lmax_lm), _tag_lm) - # label_mask 用 response_lens(每行 response token 数)。verl080 padding 后 + build_attention_mask_2d(original_lengths, max_seq_len), diag_tag) + # label_mask 用 response_lengths(每行 response token 数)。verl080 padding 后 # loss_mask = response_mask 是 2D left-right padded(非 NestedTensor,见 # verl padding.py:71),不能走 nested_to_2d_full;但 response token 数 = # loss_mask 行 sum,与坐标系无关,据此推 prompt_len 最稳(2D/NestedTensor 均适用)。 - _lm = original_batch.get("loss_mask") - if _lm is not None: - if _is_nested_tensor(_lm): - _lm_off = _lm.offsets() - _lm_val = _lm.values() - _response_lens = [ - int(_lm_val[_lm_off[i]:_lm_off[i + 1]].sum()) - for i in range(len(_orig_lens))] + loss_mask_tensor = original_batch.get("loss_mask") + if loss_mask_tensor is not None: + if _is_nested_tensor(loss_mask_tensor): + loss_offsets = loss_mask_tensor.offsets() + loss_values = loss_mask_tensor.values() + response_lengths = [ + int(loss_values[loss_offsets[i]:loss_offsets[i + 1]].sum()) + for i in range(len(original_lengths))] else: # .long() 免 import torch(本文件顶部未导入 torch); # .cpu() 防御 on-device tensor 的 tolist() - _response_lens = _lm.sum(dim=-1).long().cpu().tolist() + response_lengths = loss_mask_tensor.sum(dim=-1).long().cpu().tolist() dump_label_mask_verl080( - build_label_mask_2d(_response_lens, _orig_lens, _Lmax_lm), - _tag_lm) + build_label_mask_2d(response_lengths, original_lengths, max_seq_len), + diag_tag) # ##### [PS-diag] dump 元数据 + masks end ##### # ── 构造修改后的 iterator 喂回原始 forward_step ── @@ -201,6 +202,9 @@ def patched_forward_step( # forward_step 返回 (output_dict, partial(postprocess_func)), # 解包处理 output_dict 再重包。restore_via_2d_unfold_verl080 内部 # 会检查 context / restore_indices,无 restore 需求时 early return。 + # + # PP guard: 非末 stage 返回的是 tensor(hidden_states),不是 dict; + # restore 只在末 stage(output 是含 log_probs 的 dict)才有意义。 if ps_state is not None: from prefix_sharing.integrations.verl_mcore import restore_via_2d_unfold_verl080 from prefix_sharing.integrations.context import current_prefix_sharing_context @@ -209,22 +213,25 @@ def patched_forward_step( vocab_parallel_log_probs_from_logits, ) output_dict, postprocess_fn = output - output_dict = restore_via_2d_unfold_verl080( - output_dict, - vocab_parallel_log_probs_from_logits, - vocab_parallel_entropy, - ) - # 释放 vocab 维 logits(占用大,只在 context 生命周期内持有, - # restore 已消费完毕)。clear 职责在此,不在包装函数内。 - ctx = current_prefix_sharing_context() - if ctx is not None: - ctx.prefix_last_logits_saved.clear() + if isinstance(output_dict, dict): + output_dict = restore_via_2d_unfold_verl080( + output_dict, + vocab_parallel_log_probs_from_logits, + vocab_parallel_entropy, + ) + # 释放 vocab 维 logits(占用大,只在 context 生命周期内持有, + # restore 已消费完毕)。clear 职责在此,不在包装函数内。 + ctx = current_prefix_sharing_context() + if ctx is not None: + ctx.prefix_last_logits_saved.clear() output = (output_dict, postprocess_fn) # ##### [PS-diag] dump 2D logprobs/entropy(ON=restore后, OFF=原始) ##### # restore 后(ON)或原始 forward(OFF)的 log_probs/entropy 都是 NestedTensor, # 每行长度 = original_lengths[i],展开到统一 [B, L_max] 供 cmp_diag.cmp_2d 逐元素对比。 - import os as _os2 - if _os2.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: + # + # PP guard: 非末 stage 返回 tensor 而非 dict;跳过 get() 避免 AttributeError。 + import os as _diag_os + if _diag_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None: from prefix_sharing.tools.diagnostic_dump_verl080 import ( nested_to_2d_full, dump_logprobs_2d_verl080, dump_entropy_2d_verl080, ) @@ -234,19 +241,27 @@ def patched_forward_step( # (eval_mode→training=False→"old" 对应 old_logp 阶段; # train_mode→training=True→"train" 对应 update_actor 阶段)。 # 这样一次 run 自动产出 logprobs_old + logprobs_train 两份,不互相覆盖。 - _tag = "train" if model.training else "old" - _out_dict, _ = output - _lp = _out_dict.get("log_probs") - if _is_nested_tensor(_lp): - if ps_state is not None: - _ol = list(ps_state.prefix_sharing_plan.original_lengths) - else: - _ol = [int(d) for d in _lp.offsets().diff().tolist()] - _Lmax = max(_ol) if _ol else 0 - dump_logprobs_2d_verl080(nested_to_2d_full(_lp, _ol, _Lmax), _tag) - _ent = _out_dict.get("entropy") - if _is_nested_tensor(_ent): - dump_entropy_2d_verl080(nested_to_2d_full(_ent, _ol, _Lmax), _tag) + diag_tag = "train" if model.training else "old" + output_dict_raw, _ = output + if isinstance(output_dict_raw, dict): + log_probs_nested = output_dict_raw.get("log_probs") + if _is_nested_tensor(log_probs_nested): + if ps_state is not None: + original_lengths = list( + ps_state.prefix_sharing_plan.original_lengths) + else: + original_lengths = [ + int(length) for length + in log_probs_nested.offsets().diff().tolist()] + max_seq_len = max(original_lengths) if original_lengths else 0 + dump_logprobs_2d_verl080( + nested_to_2d_full(log_probs_nested, original_lengths, max_seq_len), + diag_tag, scope="pp_last") + entropy_nested = output_dict_raw.get("entropy") + if _is_nested_tensor(entropy_nested): + dump_entropy_2d_verl080( + nested_to_2d_full(entropy_nested, original_lengths, max_seq_len), + diag_tag, scope="pp_last") # ##### [PS-diag] dump 2D logprobs/entropy end ##### _ps_forward_step_probe("after_original_forward_step") return output diff --git a/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/vocab_logprobs.py b/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/vocab_logprobs.py index 65d7b746..5ebcfef8 100644 --- a/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/vocab_logprobs.py +++ b/prefix-sharing/prefix_sharing/setup/patches/verl080_mcore0161_ms0160/vocab_logprobs.py @@ -26,9 +26,9 @@ def patched_fn(logits, labels): # ##### [PS-diag] dump logits(ON/OFF 都 dump,必须在 original_fn 之前) ##### # logits 形态 [N, V//tp](或 [N,1,V//tp]),cmp_diag.cmp_logits_packed 会 # reshape 成 token-major [N,V] 再用 cu_seqlens+prefix_lens 对齐。 - import os as _os - _diag_on = _os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None - if _diag_on: + import os as _diag_os + diag_enabled = _diag_os.environ.get("PREFIX_SHARING_DIAG_DUMP") is not None + if diag_enabled: from prefix_sharing.tools.diagnostic_dump_verl080 import dump_logits_verl080 dump_logits_verl080(logits) # ##### [PS-diag] dump logits end ##### @@ -50,21 +50,21 @@ def patched_fn(logits, labels): logits_2d = logits.view(-1, logits.size(-1)) # ##### [PS-diag] 验证 packed 坐标对齐(logits N 是 valid 还是 padded) ##### - if _diag_on: - _layout = ctx.packed_batch_layout + if diag_enabled: + packed_layout = ctx.packed_batch_layout print( f"[PS-diag][packed-align] logits_N={logits_2d.shape[0]} " - f"total_padded={_layout.total_padded_length} " - f"total_valid={_layout.total_valid_length} " - f"has_padding={_layout.has_padding}", + f"total_padded={packed_layout.total_padded_length} " + f"total_valid={packed_layout.total_valid_length} " + f"has_padding={packed_layout.has_padding}", flush=True, ) - for _idx in ctx.prefix_last_restore_indices: + for restore_index in ctx.prefix_last_restore_indices: print( - f"[PS-diag][packed-align] reuser={_idx.reuse_idx_in_batch} " - f"provider={_idx.provider_idx_in_batch} " - f"provider_1d_pos={_idx.provider_1d_pos} " - f"target_2d_pos={_idx.target_2d_pos}", + f"[PS-diag][packed-align] reuser={restore_index.reuse_idx_in_batch} " + f"provider={restore_index.provider_idx_in_batch} " + f"provider_1d_pos={restore_index.provider_1d_pos} " + f"target_2d_pos={restore_index.target_2d_pos}", flush=True, ) # ##### [PS-diag] 验证 packed 坐标对齐 end ##### diff --git a/prefix-sharing/prefix_sharing/tools/assemble_dump.py b/prefix-sharing/prefix_sharing/tools/assemble_dump.py new file mode 100644 index 00000000..d0a23c66 --- /dev/null +++ b/prefix-sharing/prefix_sharing/tools/assemble_dump.py @@ -0,0 +1,196 @@ +"""Multi-rank diagnostic dump assembler — merges PP/TP-sharded files into a +single-card-compatible flat directory. + +Reads a raw dump directory produced by the diagnostic dump infrastructure under +TP/PP parallelism (with ``_pp{p}`` / ``_tp{r}`` file suffixes) and assembles a +clean flat directory that looks exactly like a single-card dump. The output can +then be fed directly to the **unmodified** ``cmp_diag_verl080.py``. + +Usage:: + + python assemble_dump.py --input-dir /path/to/raw_dump --output-dir /path/to/assembled + +Single-card dumps (tp==1, pp==1) are a fast path: all files are copied verbatim. +""" + +from __future__ import annotations + +import argparse +import json +import os +import shutil +from typing import Any + +import torch + +# ── Per-layer dict files that may be PP-sharded ────────────────── +# These glob patterns match the stem (without _pp suffix or .pt extension). +_PP_STAGE_STEMS: list[str] = [ + "attn_outputs", + "rope_preqk", + "rope_postqk", + "rope_freqs", + "expanded_kv", + "full_kv", + "build_kv_input_v", + "hidden_states", +] + +# ── Global (non-sharded) files — copied verbatim ───────────────── +_GLOBAL_FILES: list[str] = [ + "parallel_info.json", + "cu_seqlens_q.pt", + "cu_seqlens_q_logits.pt", + "prefix_lens.pt", +] + +# ── Optional tag-suffixed files (glob-matched) ─────────────────── +_GLOB_TAG_FILES: list[str] = [ + "attention_mask_", + "label_mask_", + "logprobs_", + "entropy_", +] + + +def _strip_ext(fname: str) -> tuple[str, str]: + """Split ``stem.ext`` → ``(stem, ext)``. ``ext`` includes the dot.""" + idx = fname.rfind(".") + if idx == -1: + return fname, "" + return fname[:idx], fname[idx:] + + +def assemble(input_dir: str, output_dir: str) -> None: + """Assemble a multi-rank dump into a single-card-compatible directory. + + Reads ``parallel_info.json`` to determine topology, then: + - copies global metadata files verbatim + - merges per-PP-stage layer dicts into single dicts + - concatenates TP-sharded logits along the vocab dimension + - copies 2D files verbatim + """ + os.makedirs(output_dir, exist_ok=True) + + # ── Load topology ───────────────────────────────────────────── + manifest_path = os.path.join(input_dir, "parallel_info.json") + if not os.path.exists(manifest_path): + print(f"[assemble] ERROR: {manifest_path} not found — is this a diagnostic dump?") + return + + with open(manifest_path, encoding="utf-8") as f: + manifest = json.load(f) + + tp_size: int = manifest.get("tp_size", 1) + pp_size: int = manifest.get("pp_size", 1) + scopes: dict[str, str] = manifest.get("scopes", {}) + + is_multi_rank = tp_size > 1 or pp_size > 1 + if not is_multi_rank: + print("[assemble] tp=1 pp=1 — fast path: copying all files") + _copy_tree(input_dir, output_dir) + return + + print(f"[assemble] tp_size={tp_size} pp_size={pp_size}") + + # ── Copy global metadata files ───────────────────────────────── + for fname in _GLOBAL_FILES: + src = os.path.join(input_dir, fname) + if os.path.exists(src): + shutil.copy2(src, os.path.join(output_dir, fname)) + print(f" [copy] {fname}") + + # ── Copy glob-tagged files (attention_mask_old.pt, etc.) ──────── + for prefix in _GLOB_TAG_FILES: + if not os.path.isdir(input_dir): + continue + for fname in os.listdir(input_dir): + if fname.startswith(prefix) and fname.endswith(".pt"): + src = os.path.join(input_dir, fname) + shutil.copy2(src, os.path.join(output_dir, fname)) + print(f" [copy] {fname}") + + # ── Merge per-PP-stage layer dicts ───────────────────────────── + for stem in _PP_STAGE_STEMS: + fname = f"{stem}.pt" + # Determine if this file is PP-sharded from manifest scopes + scope = scopes.get(stem, "") + if scope != "pp_stage" or pp_size <= 1: + # Single file — copy verbatim + src = os.path.join(input_dir, fname) + if os.path.exists(src): + shutil.copy2(src, os.path.join(output_dir, fname)) + print(f" [copy] {fname}") + continue + + # PP-sharded: load and merge + merged: dict[int, Any] = {} + found_any = False + for p in range(pp_size): + stem_ext = f"{stem}_pp{p}.pt" + src = os.path.join(input_dir, stem_ext) + if os.path.exists(src): + d = torch.load(src, weights_only=True) + if isinstance(d, dict): + merged.update(d) + found_any = True + if found_any: + torch.save(merged, os.path.join(output_dir, fname)) + print(f" [merge] {fname} ← {pp_size} stage(s), {len(merged)} layers") + + # ── Concat TP-sharded logits ────────────────────────────────── + if scopes.get("logits", "") == "tp_vocab" and tp_size > 1: + shards = [] + for t in range(tp_size): + src = os.path.join(input_dir, f"logits_tp{t}.pt") + if os.path.exists(src): + shards.append(torch.load(src, weights_only=True)) + if shards: + full = torch.cat(shards, dim=-1) + torch.save(full, os.path.join(output_dir, "logits.pt")) + print(f" [concat] logits.pt ← {len(shards)} tp shards, shape {list(full.shape)}") + else: + # tp_size==1: copy logits.pt verbatim + src = os.path.join(input_dir, "logits.pt") + if os.path.exists(src): + shutil.copy2(src, os.path.join(output_dir, "logits.pt")) + print(" [copy] logits.pt") + + # ── Summary ─────────────────────────────────────────────────── + n_files = len(os.listdir(output_dir)) + print(f"\n[assemble] done — {n_files} files written to {output_dir}") + + +def _copy_tree(src_dir: str, dst_dir: str) -> None: + """Copy all .pt and .json files from src_dir to dst_dir (fast path for single-card).""" + if not os.path.isdir(src_dir): + return + for fname in os.listdir(src_dir): + if fname.endswith(".pt") or fname.endswith(".json"): + src = os.path.join(src_dir, fname) + if os.path.isfile(src): + shutil.copy2(src, os.path.join(dst_dir, fname)) + + +# ── CLI ─────────────────────────────────────────────────────────── + +def main() -> None: + ap = argparse.ArgumentParser( + description="Assemble multi-rank (TP/PP) diagnostic dump into single-card format", + ) + ap.add_argument( + "--input-dir", "-i", + required=True, + help="Raw multi-rank dump directory (with _pp{p}/_tp{r} suffixes)", + ) + ap.add_argument( + "--output-dir", "-o", + required=True, + help="Assembled single-card-compatible directory", + ) + args = ap.parse_args() + assemble(args.input_dir, args.output_dir) + + +if __name__ == "__main__": + main() diff --git a/prefix-sharing/prefix_sharing/tools/cmp_baseline_cross_batch.py b/prefix-sharing/prefix_sharing/tools/cmp_baseline_cross_batch.py new file mode 100644 index 00000000..3c06ddd1 --- /dev/null +++ b/prefix-sharing/prefix_sharing/tools/cmp_baseline_cross_batch.py @@ -0,0 +1,679 @@ +"""GEMM Precision Baseline — Cross-Batch-Size Comparison. + +Compares the SAME data processed at DIFFERENT batch sizes to quantify +GEMM floating-point noise. Reuses comparison metrics and printing +functions from ``cmp_diag_verl080`` for consistent output. + +Usage:: + + # Run 1: single copy + export PREFIX_SHARING_DIAG_DUMP=/dump_single + # forward with batch=[A] + + # Run 2: stacked copies + export PREFIX_SHARING_DIAG_DUMP=/dump_stacked + # forward with batch=[A x N] + + python cmp_baseline_cross_batch.py \ + --dir-single /dump_single --dir-stacked /dump_stacked --num-copies 4 +""" + +from __future__ import annotations + +import argparse +import os + +import torch + +from prefix_sharing.tools.cmp_diag_verl080 import ( + CheckResult, + _SEP_DOUBLE, + _SEP_SINGLE, + _CHECK, + _CROSS, + _cosine_sim, + _pearson_r, + _dump_json, + _load_tensor, + _print_logits_packed, + _print_2d_result, + _print_topk_vec, + _print_topk_2d, + _print_summary, + _print_shapes, +) + + +# ══════════════════════════════════════════════════════════════════ +# I/O helpers +# ══════════════════════════════════════════════════════════════════ + +def _load_per_layer_dict(directory: str, filename: str) -> dict | None: + """Load a ``{layer_index: tensor_or_dict}`` file.""" + filepath = os.path.join(directory, filename) + if not os.path.exists(filepath): + return None + data = torch.load(filepath, weights_only=True) + return data if isinstance(data, dict) else None + + +def _load_cu_seqlens(directory: str) -> torch.Tensor | None: + """Load ``cu_seqlens_q.pt`` (cumulative token boundaries).""" + filepath = os.path.join(directory, "cu_seqlens_q.pt") + if not os.path.exists(filepath): + return None + return torch.load(filepath, weights_only=True) + + +def _slice_sequence(packed_tensor: torch.Tensor, + cu_seqlens: torch.Tensor, + sequence_index: int) -> torch.Tensor: + """Slice one sequence from a packed tensor using cu_seqlens.""" + start = int(cu_seqlens[sequence_index]) + end = int(cu_seqlens[sequence_index + 1]) + return packed_tensor[start:end] + + +def _sorted_layer_keys(data: dict) -> list[int]: + return sorted(int(k) for k in data.keys()) + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — plain per-layer dicts +# ══════════════════════════════════════════════════════════════════ + +def _compare_plain_per_layer( + dir_single: str, dir_stacked: str, + filename: str, + cu_seqlens_single: torch.Tensor, + cu_seqlens_multi: torch.Tensor, + stack_count: int, + filter_layer: int | None, + label: str, +) -> CheckResult | None: + """Compare ``{layer: [total_tokens, ...]}`` across batch sizes. + + Matches each sequence in the single dump against its *stack_count* + copies in the stacked dump (located at index + ``seq_index + copy_index * num_sequences``). + """ + single_data = _load_per_layer_dict(dir_single, filename) + multi_data = _load_per_layer_dict(dir_stacked, filename) + if single_data is None or multi_data is None: + return None + + layers = _sorted_layer_keys(single_data) + if filter_layer is not None: + layers = [l for l in layers if l == filter_layer] + if not layers: + return None + + num_sequences = cu_seqlens_single.numel() - 1 + per_layer: dict = {} + worst_max_diff = 0.0 + worst_cos_min = 1.0 + + for layer_index in layers: + single_tensor = single_data[layer_index].float() + multi_tensor = multi_data[layer_index].float() + layer_max_diff = 0.0 + layer_cos_min = 1.0 + all_cos_values: list[float] = [] + + for seq_index in range(num_sequences): + single_seq = _slice_sequence(single_tensor, cu_seqlens_single, seq_index) + tokens_in_seq = single_seq.shape[0] + if tokens_in_seq == 0: + continue + single_flat = single_seq.reshape(tokens_in_seq, -1) + + for copy_index in range(stack_count): + multi_offset = seq_index + copy_index * num_sequences + multi_seq = _slice_sequence(multi_tensor, cu_seqlens_multi, multi_offset) + if multi_seq.shape[0] != tokens_in_seq: + continue + multi_flat = multi_seq.reshape(tokens_in_seq, -1) + + per_token_cos = _cosine_sim(single_flat, multi_flat, dim=-1) + all_cos_values.extend(per_token_cos.tolist()) + diff_i = float((single_flat - multi_flat).abs().max()) + cos_min_i = float(per_token_cos.min()) + layer_max_diff = max(layer_max_diff, diff_i) + layer_cos_min = min(layer_cos_min, cos_min_i) + + per_layer[layer_index] = { + "max_diff": layer_max_diff, + "cos_avg": sum(all_cos_values) / len(all_cos_values) if all_cos_values else 0.0, + "cos_min": layer_cos_min, + "n_tokens": single_tensor.shape[0], + "on_T": single_tensor.shape[0], + "off_T": multi_tensor.shape[0], + } + worst_max_diff = max(worst_max_diff, layer_max_diff) + worst_cos_min = min(worst_cos_min, layer_cos_min) + + passed = worst_max_diff < 1e-5 + result_name = f"{label}_L{filter_layer}" if filter_layer is not None else label + return CheckResult( + name=result_name, passed=passed, + metrics={"layers": per_layer, "max_diff": worst_max_diff, + "cos_min": worst_cos_min, "num_layers": len(layers)}, + ) + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — per-layer KV dicts (rope_preqk / rope_postqk / +# full_kv) +# ══════════════════════════════════════════════════════════════════ + +def _compare_kv_per_layer( + dir_single: str, dir_stacked: str, + filename: str, + cu_seqlens_single: torch.Tensor, + cu_seqlens_multi: torch.Tensor, + stack_count: int, + filter_layer: int | None, + label: str, + field_first: str, + field_second: str, +) -> CheckResult | None: + """Compare ``{layer: {field_first, field_second}}`` across batch sizes.""" + single_data = _load_per_layer_dict(dir_single, filename) + multi_data = _load_per_layer_dict(dir_stacked, filename) + if single_data is None or multi_data is None: + return None + + layers = _sorted_layer_keys(single_data) + if filter_layer is not None: + layers = [l for l in layers if l == filter_layer] + if not layers: + return None + + num_sequences = cu_seqlens_single.numel() - 1 + per_layer: dict = {} + + for layer_index in layers: + single_first = single_data[layer_index][field_first].float() + single_second = single_data[layer_index][field_second].float() + multi_first = multi_data[layer_index][field_first].float() + multi_second = multi_data[layer_index][field_second].float() + + first_max_diff, second_max_diff = 0.0, 0.0 + first_cos_min, second_cos_min = 1.0, 1.0 + first_cos_list, second_cos_list = [], [] + + for seq_index in range(num_sequences): + single_seq_f = _slice_sequence(single_first, cu_seqlens_single, seq_index) + single_seq_s = _slice_sequence(single_second, cu_seqlens_single, seq_index) + tokens_in_seq = single_seq_f.shape[0] + if tokens_in_seq == 0: + continue + single_flat_f = single_seq_f.reshape(tokens_in_seq, -1) + single_flat_s = single_seq_s.reshape(tokens_in_seq, -1) + + for copy_index in range(stack_count): + multi_offset = seq_index + copy_index * num_sequences + multi_seq_f = _slice_sequence(multi_first, cu_seqlens_multi, multi_offset) + multi_seq_s = _slice_sequence(multi_second, cu_seqlens_multi, multi_offset) + if multi_seq_f.shape[0] != tokens_in_seq: + continue + multi_flat_f = multi_seq_f.reshape(tokens_in_seq, -1) + multi_flat_s = multi_seq_s.reshape(tokens_in_seq, -1) + + cos_f = _cosine_sim(single_flat_f, multi_flat_f, dim=-1) + cos_s = _cosine_sim(single_flat_s, multi_flat_s, dim=-1) + first_cos_list.extend(cos_f.tolist()) + second_cos_list.extend(cos_s.tolist()) + first_max_diff = max(first_max_diff, float((single_flat_f - multi_flat_f).abs().max())) + second_max_diff = max(second_max_diff, float((single_flat_s - multi_flat_s).abs().max())) + first_cos_min = min(first_cos_min, float(cos_f.min())) + second_cos_min = min(second_cos_min, float(cos_s.min())) + + per_layer[layer_index] = { + "Q_max_diff": first_max_diff, "K_max_diff": second_max_diff, + "Q_cos_avg": sum(first_cos_list) / len(first_cos_list) if first_cos_list else 0.0, + "Q_cos_min": first_cos_min, + "K_cos_avg": sum(second_cos_list) / len(second_cos_list) if second_cos_list else 0.0, + "K_cos_min": second_cos_min, + "n_tokens": len(first_cos_list), + } + + result_name = f"{label}_L{filter_layer}" if filter_layer is not None else label + return CheckResult(name=result_name, passed=True, metrics={"layers": per_layer}) + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — logits (single packed tensor) +# ══════════════════════════════════════════════════════════════════ + +def _compare_logits_cross_batch( + dir_single: str, dir_stacked: str, + cu_seqlens_single: torch.Tensor, + cu_seqlens_multi: torch.Tensor, + stack_count: int, +) -> CheckResult | None: + """Compare packed logits across batch sizes.""" + single_path = os.path.join(dir_single, "logits.pt") + multi_path = os.path.join(dir_stacked, "logits.pt") + if not os.path.exists(single_path) or not os.path.exists(multi_path): + return None + + single_logits = torch.load(single_path, weights_only=True).float() + multi_logits = torch.load(multi_path, weights_only=True).float() + single_logits = single_logits.reshape(-1, single_logits.size(-1)) + multi_logits = multi_logits.reshape(-1, multi_logits.size(-1)) + + num_sequences = cu_seqlens_single.numel() - 1 + worst_max_diff = 0.0 + worst_cos_min = 1.0 + all_cos_values: list[float] = [] + total_tokens = 0 + + for seq_index in range(num_sequences): + single_seq = _slice_sequence(single_logits, cu_seqlens_single, seq_index) + if single_seq.shape[0] == 0: + continue + total_tokens += single_seq.shape[0] + + for copy_index in range(stack_count): + multi_offset = seq_index + copy_index * num_sequences + multi_seq = _slice_sequence(multi_logits, cu_seqlens_multi, multi_offset) + if multi_seq.shape[0] != single_seq.shape[0]: + continue + per_token_cos = _cosine_sim(single_seq, multi_seq, dim=-1) + all_cos_values.extend(per_token_cos.tolist()) + worst_max_diff = max(worst_max_diff, float((single_seq - multi_seq).abs().max())) + worst_cos_min = min(worst_cos_min, float(per_token_cos.min())) + + return CheckResult( + name="logits", passed=worst_max_diff < 1e-5, + metrics={ + "n_tokens": total_tokens, + "cos_avg": sum(all_cos_values) / len(all_cos_values) if all_cos_values else 0.0, + "cos_min": worst_cos_min, + }, + ) + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — 2D (logprobs / entropy) +# ══════════════════════════════════════════════════════════════════ + +def _compare_2d_cross_batch( + dir_single: str, dir_stacked: str, + filename: str, label: str, + stack_count: int, num_sequences: int, + atol: float = 1e-5, +) -> tuple[CheckResult | None, torch.Tensor | None, torch.Tensor | None]: + """Compare 2D [batch, L_max] — row ``i`` from single vs rows + ``i + k * num_sequences`` from stacked.""" + single_2d = _load_tensor(dir_single, filename) + multi_2d = _load_tensor(dir_stacked, filename) + if single_2d is None or multi_2d is None: + return None, single_2d, multi_2d + + single_2d = single_2d.float() + multi_2d = multi_2d.float() + if single_2d.dim() < 2 or multi_2d.dim() < 2: + return None, single_2d, multi_2d + + worst_max_diff = 0.0 + worst_cos_min = 1.0 + worst_rel_max = 0.0 + all_abs_diffs: list[float] = [] + all_rel_diffs: list[float] = [] + pearson_pairs: list[tuple[torch.Tensor, torch.Tensor]] = [] + + for seq_index in range(num_sequences): + if seq_index >= single_2d.shape[0]: + break + single_row = single_2d[seq_index].reshape(-1) + + for copy_index in range(stack_count): + multi_offset = seq_index + copy_index * num_sequences + if multi_offset >= multi_2d.shape[0]: + continue + multi_row = multi_2d[multi_offset].reshape(-1) + + abs_diff = (single_row - multi_row).abs() + rel_diff = abs_diff / single_row.abs().clamp(min=1e-8) + all_abs_diffs.extend(abs_diff.tolist()) + all_rel_diffs.extend(rel_diff.tolist()) + + worst_max_diff = max(worst_max_diff, float(abs_diff.max())) + worst_rel_max = max(worst_rel_max, float(rel_diff.max())) + worst_cos_min = min(worst_cos_min, float(_cosine_sim(single_row, multi_row, dim=-1))) + pearson_pairs.append((single_row, multi_row)) + + # pearson: average over up to 10 pairs + pearson_values = [_pearson_r(a, b) for a, b in pearson_pairs[:10]] + pearson_avg = (sum(pearson_values) / len(pearson_values) + if pearson_values else 0.0) + + result = CheckResult( + name=label, passed=worst_max_diff <= atol, + metrics={ + "shape": tuple(single_2d.shape), + "active": single_2d[:num_sequences].numel(), + "abs_max": worst_max_diff, + "abs_mean": sum(all_abs_diffs) / len(all_abs_diffs) if all_abs_diffs else 0.0, + "rel_max": worst_rel_max, + "rel_mean": sum(all_rel_diffs) / len(all_rel_diffs) if all_rel_diffs else 0.0, + "pearson_r": pearson_avg, + "atol": atol, + }, + ) + return result, single_2d, multi_2d + + +# ══════════════════════════════════════════════════════════════════ +# Top-K helpers +# ══════════════════════════════════════════════════════════════════ + +def _print_topk_plain( + dir_single: str, dir_stacked: str, + cu_seqlens_single: torch.Tensor, + cu_seqlens_multi: torch.Tensor, + stack_count: int, + filename: str, label: str, + topk: int, sort_err: str, +): + """Print top-K worst dimensions for a plain per-layer file.""" + single_data = _load_per_layer_dict(dir_single, filename) + multi_data = _load_per_layer_dict(dir_stacked, filename) + if single_data is None or multi_data is None: + return + + last_layer = max(int(k) for k in single_data.keys()) + single_tensor = single_data[last_layer].float() + multi_tensor = multi_data[last_layer].float() + num_sequences = cu_seqlens_single.numel() - 1 + + worst_max_diff = 0.0 + worst_single_flat = worst_multi_flat = None + + for seq_index in range(num_sequences): + single_seq = _slice_sequence(single_tensor, cu_seqlens_single, seq_index) + if single_seq.shape[0] == 0: + continue + single_flat = single_seq.reshape(single_seq.shape[0], -1) + + for copy_index in range(stack_count): + multi_offset = seq_index + copy_index * num_sequences + multi_seq = _slice_sequence(multi_tensor, cu_seqlens_multi, multi_offset) + if multi_seq.shape[0] != single_seq.shape[0]: + continue + multi_flat = multi_seq.reshape(multi_seq.shape[0], -1) + max_diff = float((single_flat - multi_flat).abs().max()) + if max_diff > worst_max_diff: + worst_max_diff = max_diff + worst_single_flat = single_flat + worst_multi_flat = multi_flat + + if worst_single_flat is not None: + _print_topk_vec(worst_single_flat[0].cpu(), worst_multi_flat[0].cpu(), + topk, sort_err, f"{label}_L{last_layer}_token0") + + +def _print_topk_kv( + dir_single: str, dir_stacked: str, + cu_seqlens_single: torch.Tensor, + cu_seqlens_multi: torch.Tensor, + stack_count: int, + filename: str, field_first: str, field_second: str, + label: str, topk: int, sort_err: str, +): + """Print top-K worst dimensions for a KV-style per-layer file.""" + single_data = _load_per_layer_dict(dir_single, filename) + multi_data = _load_per_layer_dict(dir_stacked, filename) + if single_data is None or multi_data is None: + return + + last_layer = max(int(k) for k in single_data.keys()) + num_sequences = cu_seqlens_single.numel() - 1 + + for field, tag in [(field_first, f"{label}_{field_first}"), + (field_second, f"{label}_{field_second}")]: + try: + single_field = single_data[last_layer][field].float() + multi_field = multi_data[last_layer][field].float() + except (KeyError, TypeError, AttributeError) as exc: + print(f" [top-K skip] {label} {field}: {exc}") + continue + + worst_max_diff = 0.0 + worst_single_flat = worst_multi_flat = None + + for seq_index in range(num_sequences): + single_seq = _slice_sequence(single_field, cu_seqlens_single, seq_index) + if single_seq.shape[0] == 0: + continue + single_flat = single_seq.reshape(single_seq.shape[0], -1) + + for copy_index in range(stack_count): + multi_offset = seq_index + copy_index * num_sequences + multi_seq = _slice_sequence(multi_field, cu_seqlens_multi, multi_offset) + if multi_seq.shape[0] != single_seq.shape[0]: + continue + multi_flat = multi_seq.reshape(multi_seq.shape[0], -1) + max_diff = float((single_flat - multi_flat).abs().max()) + if max_diff > worst_max_diff: + worst_max_diff = max_diff + worst_single_flat = single_flat + worst_multi_flat = multi_flat + + if worst_single_flat is not None: + _print_topk_vec(worst_single_flat[0].cpu(), worst_multi_flat[0].cpu(), + topk, sort_err, f"{tag}_L{last_layer}_token0") + + +# ══════════════════════════════════════════════════════════════════ +# Print wrappers (Single/Stacked labels instead of ON/OFF) +# ══════════════════════════════════════════════════════════════════ + +def _print_table_baseline(result: CheckResult): + """Print a per-layer comparison table with Single/Stacked labels.""" + print(_SEP_SINGLE + f"\n [{result.name}] Single vs Stacked(per-seq aligned)") + print(_SEP_SINGLE) + metrics = result.metrics + if "error" in metrics: + print(f" {_CROSS} {metrics['error']}\n") + return + + layers = metrics.get("layers", {}) + header = (f" {'LAYER':>6s} {'MAXDIFF':>12s} {'COS_AVG':>10s} {'COS_MIN':>10s} " + f"{'SNG_T':>8s} {'STK_T':>8s} {'STATUS':>8s}") + print(header) + print(f" {'─' * 6} {'─' * 12} {'─' * 10} {'─' * 10} {'─' * 8} {'─' * 8} {'─' * 8}") + + for layer_index in sorted(layers): + entry = layers[layer_index] + if "max_diff" not in entry: + print(f" {layer_index:>6d} {entry.get('error', '')}") + continue + max_diff = entry["max_diff"] + cos_avg = entry.get("cos_avg", 0.0) + cos_min = entry.get("cos_min", 1.0) + single_tokens = entry.get("on_T", "—") + stacked_tokens = entry.get("off_T", "—") + ok = max_diff < 1e-5 + print(f" {layer_index:>6d} {max_diff:>12.3e} {cos_avg:>10.6f} {cos_min:>10.6f} " + f"{single_tokens:>8} {stacked_tokens:>8} " + f"{'OK' if ok else 'DIFF':>8s}") + + print(f"\n max_diff={metrics.get('max_diff')} cos_min={metrics.get('cos_min')} " + f"{_CHECK if result.passed else _CROSS}") + print() + + +def _print_kv_table_baseline(result: CheckResult, label: str): + """Print a per-layer Q/K or K/V comparison table.""" + print(_SEP_SINGLE + f"\n [{result.name}] {label} Single vs Stacked") + print(_SEP_SINGLE) + layers = result.metrics.get("layers") + if not isinstance(layers, dict): + return + + header = (f" {'LAYER':>6s} {'Q_MAXDIFF':>12s} {'Q_COS_AVG':>12s} {'Q_COS_MIN':>12s} " + f"{'K_MAXDIFF':>12s} {'K_COS_AVG':>12s} {'K_COS_MIN':>12s} " + f"{'TOKENS':>8s}") + print(header) + print(f" {'─' * 6} {'─' * 12} {'─' * 12} {'─' * 12} " + f"{'─' * 12} {'─' * 12} {'─' * 12} {'─' * 8}") + + for layer_index in sorted(layers.keys()): + entry = layers[layer_index] + if "error" in entry: + print(f" {layer_index:>6d} {entry['error']}") + continue + print(f" {layer_index:>6d} " + f"{entry.get('Q_max_diff', 0.0):>12.3e} {entry.get('Q_cos_avg', 0.0):>12.6f} " + f"{entry.get('Q_cos_min', 0.0):>12.6f} " + f"{entry.get('K_max_diff', 0.0):>12.3e} {entry.get('K_cos_avg', 0.0):>12.6f} " + f"{entry.get('K_cos_min', 0.0):>12.6f} {entry.get('n_tokens', '—'):>8}") + print() + + +# ══════════════════════════════════════════════════════════════════ +# Main +# ══════════════════════════════════════════════════════════════════ + +def main(): + parser = argparse.ArgumentParser( + description="GEMM precision baseline — cross-batch-size (single vs N copies)", + epilog=__doc__, + ) + parser.add_argument("--dir-single", required=True, + help="Single-copy dump directory") + parser.add_argument("--dir-stacked", required=True, + help="Stacked-copies dump directory") + parser.add_argument("--num-copies", type=int, required=True, + help="Number of stacked copies (stack count)") + parser.add_argument("--layer", type=int, default=None, + help="Compare specific layer (1-indexed, default: all)") + parser.add_argument("--tag", default="old", + help="2D file tag for logprobs/entropy (default: old)") + parser.add_argument("--atol", type=float, default=1e-5, + help="Absolute tolerance for 2D (default: 1e-5)") + parser.add_argument("--topk", type=int, default=0, + help="top-K worst dims for packed token (0=disabled)") + parser.add_argument("--sort-err", choices=["abs", "rel", "val"], default="abs", + help="top-K sort order: abs / rel / val") + parser.add_argument("--output", "-o", default=None, + help="Write JSON report to this path") + args = parser.parse_args() + + cu_seqlens_single = _load_cu_seqlens(args.dir_single) + cu_seqlens_multi = _load_cu_seqlens(args.dir_stacked) + if cu_seqlens_single is None or cu_seqlens_multi is None: + print(f"{_CROSS} cu_seqlens_q.pt missing") + return 1 + + stack_count = args.num_copies + num_sequences = cu_seqlens_single.numel() - 1 + total_tokens_single = int(cu_seqlens_single[-1]) + + print(_SEP_DOUBLE) + print(" GEMM Precision Baseline — Cross-Batch-Size Comparison") + print(f" Single : {args.dir_single} ({num_sequences} seqs, {total_tokens_single} tokens)") + print(f" Stacked: {args.dir_stacked} ({num_sequences * stack_count} seqs, " + f"{total_tokens_single * stack_count} tokens, {stack_count}x stack)") + print(_SEP_DOUBLE) + + # Shape diagnostics + _print_shapes(args.dir_single, args.dir_stacked, args.tag) + + all_results: list[CheckResult] = [] + + # ── Per-layer plain dicts ── + for filename, label in [ + ("hidden_states.pt", "hidden_states"), + ("build_kv_input_v.pt", "build_kv_input_v"), + ("rope_freqs.pt", "rope_freqs"), + ("attn_outputs.pt", "attn_outputs"), + ]: + result = _compare_plain_per_layer( + args.dir_single, args.dir_stacked, filename, + cu_seqlens_single, cu_seqlens_multi, stack_count, + args.layer, label, + ) + if result: + all_results.append(result) + _print_table_baseline(result) + + # ── Per-layer KV dicts ── + for filename, label, field_a, field_b in [ + ("rope_preqk.pt", "rope_preqk", "query", "key"), + ("rope_postqk.pt", "rope_postqk", "query", "key"), + ("full_kv.pt", "full_kv", "key", "value"), + ]: + result = _compare_kv_per_layer( + args.dir_single, args.dir_stacked, filename, + cu_seqlens_single, cu_seqlens_multi, stack_count, + args.layer, label, field_a, field_b, + ) + if result: + all_results.append(result) + _print_kv_table_baseline(result, label) + + # ── Logits ── + result = _compare_logits_cross_batch( + args.dir_single, args.dir_stacked, + cu_seqlens_single, cu_seqlens_multi, stack_count, + ) + if result: + all_results.append(result) + _print_logits_packed(result) + + # ── 2D ── + _2d_tensors: list[tuple[str, torch.Tensor, torch.Tensor]] = [] + for file_tag, compare_name in [("logprobs", "logp"), ("entropy", "entropy")]: + filename = f"{file_tag}_{args.tag}.pt" + result, single_2d, multi_2d = _compare_2d_cross_batch( + args.dir_single, args.dir_stacked, filename, + f"{compare_name}_{args.tag}", stack_count, num_sequences, args.atol, + ) + if result: + all_results.append(result) + _print_2d_result(result) + if single_2d is not None and multi_2d is not None: + _2d_tensors.append((f"{compare_name}_{args.tag}", single_2d, multi_2d)) + + # ── Top-K ── + if args.topk > 0: + _print_topk_plain( + args.dir_single, args.dir_stacked, + cu_seqlens_single, cu_seqlens_multi, stack_count, + "build_kv_input_v.pt", "build_kv_input_v", + args.topk, args.sort_err, + ) + _print_topk_kv( + args.dir_single, args.dir_stacked, + cu_seqlens_single, cu_seqlens_multi, stack_count, + "rope_preqk.pt", "query", "key", "rope_preqk", + args.topk, args.sort_err, + ) + _print_topk_kv( + args.dir_single, args.dir_stacked, + cu_seqlens_single, cu_seqlens_multi, stack_count, + "rope_postqk.pt", "query", "key", "rope_postqk", + args.topk, args.sort_err, + ) + for label, single_2d, multi_2d in _2d_tensors: + if (single_2d.dim() >= 2 and multi_2d.dim() >= 2 + and single_2d.shape[1] == multi_2d.shape[1]): + t1 = single_2d[:num_sequences] + t2 = multi_2d[:num_sequences] + if t1.shape == t2.shape: + _print_topk_2d(t1.cpu(), t2.cpu(), None, + args.topk, args.sort_err, label) + + _print_summary(all_results) + + if args.output: + _dump_json(all_results, args.output, args.dir_single, args.dir_stacked, + tag=f"cross_batch_N{stack_count}", dir_off2=None) + + +if __name__ == "__main__": + main() diff --git a/prefix-sharing/prefix_sharing/tools/cmp_baseline_within_batch.py b/prefix-sharing/prefix_sharing/tools/cmp_baseline_within_batch.py new file mode 100644 index 00000000..c3ce5c7a --- /dev/null +++ b/prefix-sharing/prefix_sharing/tools/cmp_baseline_within_batch.py @@ -0,0 +1,538 @@ +"""GEMM Precision Baseline — Within-Batch Pairwise Comparison. + +Compares N identical copies WITHIN a single forward pass to verify +same-GEMM-kernel bit-identical reproduction. Reuses ``cmp_diag_verl080`` +printing for consistent output. + +Usage:: + + export PREFIX_SHARING_DIAG_DUMP=/dump_multi + # forward with batch=[A x N] + + python cmp_baseline_within_batch.py --dir-multi /dump_multi --num-seq 4 +""" + +from __future__ import annotations + +import argparse +import os + +import torch + +from prefix_sharing.tools.cmp_diag_verl080 import ( + CheckResult, + _SEP_DOUBLE, + _SEP_SINGLE, + _CHECK, + _CROSS, + _cosine_sim, + _pearson_r, + _dump_json, + _load_tensor, + _print_logits_packed, + _print_2d_result, + _print_topk_vec, + _print_topk_2d, + _print_summary, + _print_shapes, +) + + +# ══════════════════════════════════════════════════════════════════ +# I/O helpers +# ══════════════════════════════════════════════════════════════════ + +def _load_per_layer_dict(directory: str, filename: str) -> dict | None: + filepath = os.path.join(directory, filename) + if not os.path.exists(filepath): + return None + data = torch.load(filepath, weights_only=True) + return data if isinstance(data, dict) else None + + +def _load_cu_seqlens(directory: str) -> torch.Tensor | None: + filepath = os.path.join(directory, "cu_seqlens_q.pt") + if not os.path.exists(filepath): + return None + return torch.load(filepath, weights_only=True) + + +def _slice_sequence(packed_tensor: torch.Tensor, + cu_seqlens: torch.Tensor, + sequence_index: int) -> torch.Tensor: + return packed_tensor[int(cu_seqlens[sequence_index]): + int(cu_seqlens[sequence_index + 1])] + + +def _sorted_layer_keys(data: dict) -> list[int]: + return sorted(int(k) for k in data.keys()) + + +def _group_by_token_count(tensors: list[torch.Tensor] + ) -> dict[int, list[torch.Tensor]]: + """Group tensors by their leading dimension (token count).""" + groups: dict[int, list[torch.Tensor]] = {} + for tensor in tensors: + groups.setdefault(tensor.shape[0], []).append(tensor) + return groups + + +# ══════════════════════════════════════════════════════════════════ +# Pairwise comparison within a group of same-length tensors +# ══════════════════════════════════════════════════════════════════ + +def _pairwise_metrics(copies: list[torch.Tensor]) -> dict: + """Compare all pairs within a group; return worst max_diff, cos_min, + and average cos_avg.""" + worst_max_diff = 0.0 + worst_cos_min = 1.0 + all_cos_values: list[float] = [] + + for i in range(len(copies)): + for j in range(i + 1, len(copies)): + flat_i = copies[i].reshape(copies[i].shape[0], -1).float() + flat_j = copies[j].reshape(copies[j].shape[0], -1).float() + + per_token_cos = _cosine_sim(flat_i, flat_j, dim=-1) + all_cos_values.extend(per_token_cos.tolist()) + worst_max_diff = max(worst_max_diff, float((flat_i - flat_j).abs().max())) + worst_cos_min = min(worst_cos_min, float(per_token_cos.min())) + + return { + "max_diff": worst_max_diff, + "cos_avg": sum(all_cos_values) / len(all_cos_values) if all_cos_values else 0.0, + "cos_min": worst_cos_min, + } + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — plain per-layer dicts +# ══════════════════════════════════════════════════════════════════ + +def _compare_plain_within( + directory: str, filename: str, + cu_seqlens: torch.Tensor, total_sequences: int, + filter_layer: int | None, label: str, +) -> CheckResult | None: + """Pairwise-compare copies within a single dump for ``filename``.""" + data = _load_per_layer_dict(directory, filename) + if data is None: + return None + + layers = _sorted_layer_keys(data) + if filter_layer is not None: + layers = [l for l in layers if l == filter_layer] + if not layers: + return None + + per_layer: dict = {} + worst_max_diff = 0.0 + worst_cos_min = 1.0 + + for layer_index in layers: + multi_tensor = data[layer_index].float() + copies = [_slice_sequence(multi_tensor, cu_seqlens, seq_index) + for seq_index in range(total_sequences)] + groups = _group_by_token_count(copies) + + layer_max_diff = 0.0 + layer_cos_min = 1.0 + group_cos_avgs: list[float] = [] + + for group_copies in groups.values(): + if len(group_copies) < 2: + continue + metrics = _pairwise_metrics(group_copies) + layer_max_diff = max(layer_max_diff, metrics["max_diff"]) + layer_cos_min = min(layer_cos_min, metrics["cos_min"]) + group_cos_avgs.append(metrics["cos_avg"]) + + per_layer[layer_index] = { + "max_diff": layer_max_diff, + "cos_avg": (sum(group_cos_avgs) / len(group_cos_avgs) + if group_cos_avgs else 0.0), + "cos_min": layer_cos_min, + "n_tokens": multi_tensor.shape[0], + "on_T": multi_tensor.shape[0], + "off_T": multi_tensor.shape[0], + } + worst_max_diff = max(worst_max_diff, layer_max_diff) + worst_cos_min = min(worst_cos_min, layer_cos_min) + + passed = worst_max_diff == 0.0 + result_name = f"{label}_L{filter_layer}" if filter_layer is not None else label + return CheckResult( + name=result_name, passed=passed, + metrics={"layers": per_layer, "max_diff": worst_max_diff, + "cos_min": worst_cos_min, "num_layers": len(layers)}, + ) + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — per-layer KV dicts +# ══════════════════════════════════════════════════════════════════ + +def _compare_kv_within( + directory: str, filename: str, + cu_seqlens: torch.Tensor, total_sequences: int, + filter_layer: int | None, label: str, + field_first: str, field_second: str, +) -> CheckResult | None: + """Pairwise-compare ``{layer: {field_first, field_second}}`` within a dump.""" + data = _load_per_layer_dict(directory, filename) + if data is None: + return None + + layers = _sorted_layer_keys(data) + if filter_layer is not None: + layers = [l for l in layers if l == filter_layer] + if not layers: + return None + + per_layer: dict = {} + + for layer_index in layers: + multi_first = data[layer_index][field_first].float() + multi_second = data[layer_index][field_second].float() + first_copies = [_slice_sequence(multi_first, cu_seqlens, seq_index) + for seq_index in range(total_sequences)] + second_copies = [_slice_sequence(multi_second, cu_seqlens, seq_index) + for seq_index in range(total_sequences)] + + first_groups = _group_by_token_count(first_copies) + second_groups = _group_by_token_count(second_copies) + + first_worst = {"max_diff": 0.0, "cos_min": 1.0, "cos_avg": 0.0} + second_worst = {"max_diff": 0.0, "cos_min": 1.0, "cos_avg": 0.0} + first_avgs, second_avgs = [], [] + + for group in first_groups.values(): + if len(group) >= 2: + metrics = _pairwise_metrics(group) + first_worst["max_diff"] = max(first_worst["max_diff"], metrics["max_diff"]) + first_worst["cos_min"] = min(first_worst["cos_min"], metrics["cos_min"]) + first_avgs.append(metrics["cos_avg"]) + + for group in second_groups.values(): + if len(group) >= 2: + metrics = _pairwise_metrics(group) + second_worst["max_diff"] = max(second_worst["max_diff"], metrics["max_diff"]) + second_worst["cos_min"] = min(second_worst["cos_min"], metrics["cos_min"]) + second_avgs.append(metrics["cos_avg"]) + + per_layer[layer_index] = { + "Q_max_diff": first_worst["max_diff"], + "K_max_diff": second_worst["max_diff"], + "Q_cos_avg": sum(first_avgs) / len(first_avgs) if first_avgs else 0.0, + "Q_cos_min": first_worst["cos_min"], + "K_cos_avg": sum(second_avgs) / len(second_avgs) if second_avgs else 0.0, + "K_cos_min": second_worst["cos_min"], + "n_tokens": sum(g[0].shape[0] * len(g) for g in first_groups.values()), + } + + result_name = f"{label}_L{filter_layer}" if filter_layer is not None else label + return CheckResult(name=result_name, passed=True, metrics={"layers": per_layer}) + + +# ══════════════════════════════════════════════════════════════════ +# Comparison logic — logits +# ══════════════════════════════════════════════════════════════════ + +def _compare_logits_within( + directory: str, + cu_seqlens: torch.Tensor, + total_sequences: int, +) -> CheckResult | None: + """Pairwise-compare packed logits within a dump.""" + filepath = os.path.join(directory, "logits.pt") + if not os.path.exists(filepath): + return None + + multi_logits = torch.load(filepath, weights_only=True).float() + multi_logits = multi_logits.reshape(-1, multi_logits.size(-1)) + copies = [_slice_sequence(multi_logits, cu_seqlens, seq_index) + for seq_index in range(total_sequences)] + groups = _group_by_token_count(copies) + + worst_max_diff = 0.0 + worst_cos_min = 1.0 + all_cos_avgs: list[float] = [] + + for group in groups.values(): + if len(group) >= 2: + metrics = _pairwise_metrics(group) + worst_max_diff = max(worst_max_diff, metrics["max_diff"]) + worst_cos_min = min(worst_cos_min, metrics["cos_min"]) + all_cos_avgs.append(metrics["cos_avg"]) + + return CheckResult( + name="logits", passed=worst_max_diff == 0.0, + metrics={ + "n_tokens": copies[0].shape[0] if copies else 0, + "cos_avg": (sum(all_cos_avgs) / len(all_cos_avgs) + if all_cos_avgs else 0.0), + "cos_min": worst_cos_min, + }, + ) + + +# ══════════════════════════════════════════════════════════════════ +# Print wrappers (within-batch labels) +# ══════════════════════════════════════════════════════════════════ + +def _print_table_baseline(result: CheckResult): + """Print a per-layer comparison table for within-batch.""" + print(_SEP_SINGLE + f"\n [{result.name}] within-batch pairwise") + print(_SEP_SINGLE) + metrics = result.metrics + if "error" in metrics: + print(f" {_CROSS} {metrics['error']}\n") + return + + layers = metrics.get("layers", {}) + header = (f" {'LAYER':>6s} {'MAXDIFF':>12s} {'COS_AVG':>10s} {'COS_MIN':>10s} " + f"{'TOKENS':>8s} {'STATUS':>8s}") + print(header) + print(f" {'─' * 6} {'─' * 12} {'─' * 10} {'─' * 10} {'─' * 8} {'─' * 8}") + + for layer_index in sorted(layers): + entry = layers[layer_index] + if "max_diff" not in entry: + print(f" {layer_index:>6d} {entry.get('error', '')}") + continue + max_diff = entry["max_diff"] + cos_avg = entry.get("cos_avg", 0.0) + cos_min = entry.get("cos_min", 1.0) + tokens = entry.get("n_tokens", "—") + ok = max_diff == 0.0 + print(f" {layer_index:>6d} {max_diff:>12.3e} {cos_avg:>10.6f} {cos_min:>10.6f} " + f"{tokens:>8} {'PASS' if ok else 'DIFF':>8s}") + + print(f"\n max_diff={metrics.get('max_diff')} cos_min={metrics.get('cos_min')} " + f"{_CHECK if result.passed else _CROSS}") + print() + + +def _print_kv_table_baseline(result: CheckResult, label: str): + """Print a per-layer Q/K or K/V comparison table for within-batch.""" + print(_SEP_SINGLE + f"\n [{result.name}] {label} within-batch pairwise") + print(_SEP_SINGLE) + layers = result.metrics.get("layers") + if not isinstance(layers, dict): + return + + header = (f" {'LAYER':>6s} {'Q_MAXDIFF':>12s} {'Q_COS_AVG':>12s} {'Q_COS_MIN':>12s} " + f"{'K_MAXDIFF':>12s} {'K_COS_AVG':>12s} {'K_COS_MIN':>12s} " + f"{'TOKENS':>8s}") + print(header) + print(f" {'─' * 6} {'─' * 12} {'─' * 12} {'─' * 12} " + f"{'─' * 12} {'─' * 12} {'─' * 12} {'─' * 8}") + + for layer_index in sorted(layers.keys()): + entry = layers[layer_index] + if "error" in entry: + print(f" {layer_index:>6d} {entry['error']}") + continue + print(f" {layer_index:>6d} " + f"{entry.get('Q_max_diff', 0.0):>12.3e} {entry.get('Q_cos_avg', 0.0):>12.6f} " + f"{entry.get('Q_cos_min', 0.0):>12.6f} " + f"{entry.get('K_max_diff', 0.0):>12.3e} {entry.get('K_cos_avg', 0.0):>12.6f} " + f"{entry.get('K_cos_min', 0.0):>12.6f} {entry.get('n_tokens', '—'):>8}") + print() + + +# ══════════════════════════════════════════════════════════════════ +# Main +# ══════════════════════════════════════════════════════════════════ + +def main(): + parser = argparse.ArgumentParser( + description="GEMM precision baseline — within-batch pairwise comparison", + epilog=__doc__, + ) + parser.add_argument("--dir-multi", required=True, + help="Multi-copy dump directory") + parser.add_argument("--num-seq", type=int, default=None, + help="Number of distinct sequences (required for 2D within-batch)") + parser.add_argument("--layer", type=int, default=None, + help="Compare specific layer (1-indexed, default: all)") + parser.add_argument("--tag", default="old", + help="2D file tag for logprobs/entropy (default: old)") + parser.add_argument("--atol", type=float, default=1e-5, + help="Absolute tolerance for 2D (default: 1e-5)") + parser.add_argument("--topk", type=int, default=0, + help="top-K worst dims (0=disabled)") + parser.add_argument("--sort-err", choices=["abs", "rel", "val"], default="abs", + help="top-K sort order: abs / rel / val") + parser.add_argument("--output", "-o", default=None, + help="Write JSON report to this path") + args = parser.parse_args() + + cu_seqlens = _load_cu_seqlens(args.dir_multi) + if cu_seqlens is None: + print(f"{_CROSS} cu_seqlens_q.pt missing") + return 1 + + total_sequences = cu_seqlens.numel() - 1 + lengths = [int(cu_seqlens[i + 1]) - int(cu_seqlens[i]) + for i in range(total_sequences)] + + print(_SEP_DOUBLE) + print(" GEMM Precision Baseline — Within-Batch Pairwise Comparison") + print(f" Directory: {args.dir_multi} ({total_sequences} sequences)") + print(f" Sequence lengths: {lengths}") + print(_SEP_DOUBLE) + + _print_shapes(args.dir_multi, args.dir_multi, args.tag) + + all_results: list[CheckResult] = [] + + # ── Per-layer plain dicts ── + for filename, label in [ + ("hidden_states.pt", "hidden_states"), + ("build_kv_input_v.pt", "build_kv_input_v"), + ("rope_freqs.pt", "rope_freqs"), + ("attn_outputs.pt", "attn_outputs"), + ]: + result = _compare_plain_within( + args.dir_multi, filename, cu_seqlens, total_sequences, + args.layer, label, + ) + if result: + all_results.append(result) + _print_table_baseline(result) + + # ── Per-layer KV dicts ── + for filename, label, field_a, field_b in [ + ("rope_preqk.pt", "rope_preqk", "query", "key"), + ("rope_postqk.pt", "rope_postqk", "query", "key"), + ("full_kv.pt", "full_kv", "key", "value"), + ]: + result = _compare_kv_within( + args.dir_multi, filename, cu_seqlens, total_sequences, + args.layer, label, field_a, field_b, + ) + if result: + all_results.append(result) + _print_kv_table_baseline(result, label) + + # ── Logits ── + result = _compare_logits_within(args.dir_multi, cu_seqlens, total_sequences) + if result: + all_results.append(result) + _print_logits_packed(result) + + # ── 2D ── + _2d_tensors: list[tuple[str, torch.Tensor]] = [] + num_sequences = args.num_seq or total_sequences + + for file_tag, compare_name in [("logprobs", "logp"), ("entropy", "entropy")]: + filename = f"{file_tag}_{args.tag}.pt" + tensor_2d = _load_tensor(args.dir_multi, filename) + if tensor_2d is None or tensor_2d.dim() < 2: + continue + tensor_2d = tensor_2d.float() + batch_size = tensor_2d.shape[0] + stack = (batch_size // num_sequences + if num_sequences > 0 and batch_size % num_sequences == 0 else 1) + + all_abs_diffs: list[float] = [] + all_rel_diffs: list[float] = [] + worst_max_diff = 0.0 + worst_cos_min = 1.0 + worst_rel_max = 0.0 + + # Only compare stack copies of the SAME logical sequence + for seq_index in range(num_sequences): + if seq_index >= batch_size: + break + row_a = tensor_2d[seq_index].reshape(-1) + for copy_index in range(1, stack): + row_b_offset = seq_index + copy_index * num_sequences + if row_b_offset >= batch_size: + continue + row_b = tensor_2d[row_b_offset].reshape(-1) + + abs_diff = (row_a - row_b).abs() + rel_diff = abs_diff / row_a.abs().clamp(min=1e-8) + all_abs_diffs.extend(abs_diff.tolist()) + all_rel_diffs.extend(rel_diff.tolist()) + worst_max_diff = max(worst_max_diff, float(abs_diff.max())) + worst_rel_max = max(worst_rel_max, float(rel_diff.max())) + worst_cos_min = min(worst_cos_min, + float(_cosine_sim(row_a, row_b, dim=-1))) + + # Pearson on first two stack copies + pearson_val = 1.0 + if batch_size >= 2 * num_sequences: + pearson_val = _pearson_r( + tensor_2d[:num_sequences].float().reshape(-1), + tensor_2d[num_sequences:2 * num_sequences].float().reshape(-1), + ) + + result = CheckResult( + name=f"{compare_name}_{args.tag}", + passed=worst_max_diff == 0.0, + metrics={ + "shape": tuple(tensor_2d.shape), + "active": num_sequences * tensor_2d.shape[1], + "abs_max": worst_max_diff, + "abs_mean": (sum(all_abs_diffs) / len(all_abs_diffs) + if all_abs_diffs else 0.0), + "rel_max": worst_rel_max, + "rel_mean": (sum(all_rel_diffs) / len(all_rel_diffs) + if all_rel_diffs else 0.0), + "pearson_r": pearson_val, + "atol": args.atol, + }, + ) + all_results.append(result) + _print_2d_result(result) + _2d_tensors.append((f"{compare_name}_{args.tag}", tensor_2d)) + + # ── Top-K ── + if args.topk > 0: + # 2D top-K: compare row 0 vs row num_seq (stack copies of same seq) + for label, tensor_2d in _2d_tensors: + ns = (num_sequences if num_sequences > 0 + and tensor_2d.shape[0] % num_sequences == 0 + else tensor_2d.shape[0]) + if tensor_2d.shape[0] >= ns + 1: + _print_topk_2d(tensor_2d[:1].cpu(), tensor_2d[ns:ns + 1].cpu(), + None, args.topk, args.sort_err, label) + + # build_kv_input_v top-K + data = _load_per_layer_dict(args.dir_multi, "build_kv_input_v.pt") + if data: + last_layer = max(int(k) for k in data.keys()) + multi_tensor = data[last_layer].float() + copies = [_slice_sequence(multi_tensor, cu_seqlens, seq_index) + for seq_index in range(total_sequences)] + groups = _group_by_token_count(copies) + worst_max_diff = 0.0 + worst_a = worst_b = None + for group in groups.values(): + if len(group) < 2: + continue + for i in range(len(group)): + for j in range(i + 1, len(group)): + max_diff = float((group[i].reshape(-1) - + group[j].reshape(-1)).abs().max()) + if max_diff > worst_max_diff: + worst_max_diff = max_diff + worst_a = group[i].reshape(-1) + worst_b = group[j].reshape(-1) + if worst_a is not None: + _print_topk_vec(worst_a.cpu(), worst_b.cpu(), + args.topk, args.sort_err, + f"build_kv_input_v_L{last_layer}_token0") + + _print_summary(all_results) + + if args.output: + _dump_json(all_results, args.output, "—", args.dir_multi, + tag=f"within_batch_N{total_sequences}", dir_off2=None) + + +if __name__ == "__main__": + main() diff --git a/prefix-sharing/prefix_sharing/tools/cmp_diag.py b/prefix-sharing/prefix_sharing/tools/cmp_diag.py index 7f410db5..80c0a411 100644 --- a/prefix-sharing/prefix_sharing/tools/cmp_diag.py +++ b/prefix-sharing/prefix_sharing/tools/cmp_diag.py @@ -200,7 +200,7 @@ def _first_token_metrics(a_vec: torch.Tensor, b_vec: torch.Tensor) -> dict: "cos": cos, "pearson": pr} -def _logits_first_token(lo: torch.Tensor, lf: torch.Tensor +def _logits_first_token(logits_on: torch.Tensor, logits_off: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor] | None: """Extract first token's full-vocab vector from ON/OFF logits. @@ -209,21 +209,21 @@ def _logits_first_token(lo: torch.Tensor, lf: torch.Tensor output_layer → [S, B, V//tp] → model returns [B, S, V//tp]). """ # Flatten all leading batch/token dims into N, keep V as last dim - lo_2d = lo.reshape(-1, lo.size(-1)) - lf_2d = lf.reshape(-1, lf.size(-1)) + on_flat = logits_on.reshape(-1, logits_on.size(-1)) + off_flat = logits_off.reshape(-1, logits_off.size(-1)) # First token = first row → full-vocab vector [V] - return lo_2d[0, :].contiguous(), lf_2d[0, :].contiguous() + return on_flat[0, :].contiguous(), off_flat[0, :].contiguous() -def _logits_ensure_token_major(lo: torch.Tensor, lf: torch.Tensor +def _logits_ensure_token_major(logits_on: torch.Tensor, logits_off: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor]: """Ensure logits are 2D [N, V] for packed alignment. Vocab is always the last dim (verified from Megatron model forward). Batch dim on leading axes is flattened into N. """ - return (lo.reshape(-1, lo.size(-1)).contiguous(), - lf.reshape(-1, lf.size(-1)).contiguous()) + return (logits_on.reshape(-1, logits_on.size(-1)).contiguous(), + logits_off.reshape(-1, logits_off.size(-1)).contiguous()) # ══════════════════════════════════════════════════════════════════ @@ -231,30 +231,30 @@ def _logits_ensure_token_major(lo: torch.Tensor, lf: torch.Tensor # ══════════════════════════════════════════════════════════════════ def _load_tensor(dir_path: str, filename: str) -> torch.Tensor | None: - fp = os.path.join(dir_path, filename) - return torch.load(fp, weights_only=True).float() if os.path.exists(fp) else None + filepath = os.path.join(dir_path, filename) + return torch.load(filepath, weights_only=True).float() if os.path.exists(filepath) else None def _load_packed_meta(dir_path: str, cu_fname: str = "cu_seqlens_q.pt") -> dict | None: - fp = os.path.join(dir_path, cu_fname) - if not os.path.exists(fp): - fp = os.path.join(dir_path, "cu_seqlens_q.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, cu_fname) + if not os.path.exists(filepath): + filepath = os.path.join(dir_path, "cu_seqlens_q.pt") + if not os.path.exists(filepath): return None - pl_fp = os.path.join(dir_path, "prefix_lens.pt") - if not os.path.exists(pl_fp): + prefix_lens_filepath = os.path.join(dir_path, "prefix_lens.pt") + if not os.path.exists(prefix_lens_filepath): return None - return {"cu_seqlens": torch.load(fp, weights_only=True), - "prefix_lens": torch.load(pl_fp, weights_only=True)} + return {"cu_seqlens": torch.load(filepath, weights_only=True), + "prefix_lens": torch.load(prefix_lens_filepath, weights_only=True)} def _load_attn_output(dir_path: str, layer: int) -> torch.Tensor | None: """Load a single layer's attn_output from attn_outputs.pt dict.""" - fp = os.path.join(dir_path, "attn_outputs.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, "attn_outputs.pt") + if not os.path.exists(filepath): return None - d = torch.load(fp, weights_only=True) + d = torch.load(filepath, weights_only=True) return d.get(layer) if isinstance(d, dict) else None @@ -270,17 +270,17 @@ def _load_attention_mask_2d(dir_path: str) -> torch.Tensor | None: runs). ON mask is a strict subset of OFF mask per row, which ``_build_alignment_mask_from_2d`` relies on. """ - fp = os.path.join(dir_path, "attention_mask.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, "attention_mask.pt") + if not os.path.exists(filepath): return None - return torch.load(fp, weights_only=True).to(torch.bool) + return torch.load(filepath, weights_only=True).to(torch.bool) def _get_num_layers(dir_path: str) -> int: - fp = os.path.join(dir_path, "attn_outputs.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, "attn_outputs.pt") + if not os.path.exists(filepath): return 0 - d = torch.load(fp, weights_only=True) + d = torch.load(filepath, weights_only=True) return max(d.keys()) if isinstance(d, dict) and d else 0 @@ -386,25 +386,111 @@ def cmp_position_ids(dir_a: str, dir_b: str) -> CheckResult | None: # 2. RoPE encoding — absolute equality # ══════════════════════════════════════════════════════════════════ -def cmp_rope_emb(dir_a: str, dir_b: str) -> CheckResult | None: - fa = os.path.join(dir_a, "rope_emb.pt") - fb = os.path.join(dir_b, "rope_emb.pt") - if not os.path.exists(fa) or not os.path.exists(fb): +def cmp_rope_postqk(dir_on: str, dir_off: str) -> CheckResult | None: + """Compare post-RoPE Q/K between ON and OFF, aligned by position IDs. + + ON positions are absolute (preserved from original input). + OFF positions are per-segment relative (0..L-1 for each row). + + Comparison strategy: + 1. First packed token (row 0, pos 0): always a provider, must match exactly. + 2. Within the first row: match tokens by position ID (0..L-1 for both paths). + 3. For reuser rows: ON uses absolute positions (prefix_len..), OFF uses + relative (0..). Positions differ by design — skip direct comparison. + """ + filepath_on = os.path.join(dir_on, "rope_postqk.pt") + filepath_off = os.path.join(dir_off, "rope_postqk.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): return None - a = torch.load(fa, weights_only=True) - b = torch.load(fb, weights_only=True) + a = torch.load(filepath_on, weights_only=True) + b = torch.load(filepath_off, weights_only=True) if not isinstance(a, dict) or not isinstance(b, dict): return None - la, lb = set(a.keys()), set(b.keys()) - if la != lb: - return CheckResult(name="rope_emb", passed=False, + layers_on, layers_off =set(a.keys()), set(b.keys()) + if layers_on != layers_off: + return CheckResult(name="rope_postqk", passed=False, metrics={"error": "layer set mismatch"}) - md = 0.0 - for lyr in sorted(la): - for k in ("query", "key"): - md = max(md, float((a[lyr][k] - b[lyr][k]).abs().max())) - return CheckResult(name="rope_emb", passed=md == 0.0, - metrics={"max_diff": md, "num_layers": len(la)}) + + max_diff_q = 0.0 + max_diff_k = 0.0 + first_token_q_diff = 0.0 + first_token_k_diff = 0.0 + first_row_q_diff = 0.0 + first_row_k_diff = 0.0 + first_row_len = 0 + num_layers = len(layers_on) + + for layer_idx in sorted(layers_on): + on_entry = a[layer_idx] + off_entry = b[layer_idx] + on_q, on_k = on_entry["query"], on_entry["key"] + off_q, off_k = off_entry["query"], off_entry["key"] + on_pos = on_entry.get("positions") + off_pos = off_entry.get("positions") + + # ── Check 1: first packed token (row 0, position 0) ── + first_token_q_diff = max(first_token_q_diff, + float((on_q[0] - off_q[0]).abs().max())) + first_token_k_diff = max(first_token_k_diff, + float((on_k[0] - off_k[0]).abs().max())) + + # ── Check 2: first row — match by ON position IDs ── + if on_pos is not None and off_pos is not None: + on_pos_t = on_pos.long() + off_pos_t = off_pos.long() + # Find first-row extent in ON: tokens before the first position reset + # (position decreases or jumps to prefix_start) + on_row1_end = 1 + for j in range(1, len(on_pos_t)): + if on_pos_t[j] <= on_pos_t[j - 1]: + break + on_row1_end = j + 1 + on_row1_len = on_row1_end # positions 0..L-1 + + # First row in OFF: positions go 0..L-1 (find matching extent) + off_row1_end = 1 + for j in range(1, len(off_pos_t)): + if off_pos_t[j] <= off_pos_t[j - 1]: + break + off_row1_end = j + 1 + off_row1_len = off_row1_end + + compare_length = min(on_row1_len, off_row1_len) + if compare_length > 0: + first_row_len = max(first_row_len, compare_length) + # Match by position ID within the first row + for position_id in range(compare_length): + on_index = position_id # ON row 1 starts at packed index 0 + off_index = position_id # OFF row 1 starts at packed index 0 + first_row_q_diff = max(first_row_q_diff, + float((on_q[on_index] - off_q[off_index]).abs().max())) + first_row_k_diff = max(first_row_k_diff, + float((on_k[on_index] - off_k[off_index]).abs().max())) + max_diff_q = max(max_diff_q, first_row_q_diff) + max_diff_k = max(max_diff_k, first_row_k_diff) + else: + # Fallback: no positions available, compare first row by direct offset + compare_length = min(on_q.shape[0], off_q.shape[0], 128) + first_row_len = compare_length + first_row_q_diff = float((on_q[:compare_length] - off_q[:compare_length]).abs().max()) + first_row_k_diff = float((on_k[:compare_length] - off_k[:compare_length]).abs().max()) + max_diff_q = first_row_q_diff + max_diff_k = first_row_k_diff + + # Pass if all diffs are near zero + threshold = 1e-7 + passed = (first_token_q_diff < threshold and first_token_k_diff < threshold + and first_row_q_diff < threshold and first_row_k_diff < threshold) + return CheckResult(name="rope_postqk", passed=passed, metrics={ + "num_layers": num_layers, + "first_row_len": first_row_len, + "first_token_q_maxdiff": first_token_q_diff, + "first_token_k_maxdiff": first_token_k_diff, + "first_row_q_maxdiff": first_row_q_diff, + "first_row_k_maxdiff": first_row_k_diff, + "max_diff_q": max_diff_q, + "max_diff_k": max_diff_k, + }) def cmp_rope_freqs(dir_on: str, dir_off: str) -> CheckResult | None: @@ -419,40 +505,40 @@ def cmp_rope_freqs(dir_on: str, dir_off: str) -> CheckResult | None: The two are aligned to suffix-only via the same attention_mask dual-pointer logic used for attn_outputs / logits. """ - fa = os.path.join(dir_on, "rope_freqs_on.pt") - fb = os.path.join(dir_off, "rope_freqs_off.pt") - if not os.path.exists(fa) or not os.path.exists(fb): + filepath_on = os.path.join(dir_on, "rope_freqs_on.pt") + filepath_off = os.path.join(dir_off, "rope_freqs_off.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): return None - on_dict = torch.load(fa, weights_only=True) - off_dict = torch.load(fb, weights_only=True) + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): return None - la, lb = set(on_dict.keys()), set(off_dict.keys()) - if la != lb: + layers_on, layers_off =set(on_dict.keys()), set(off_dict.keys()) + if layers_on != layers_off: return CheckResult(name="rope_freqs", passed=False, metrics={"error": "layer set mismatch", - "on_layers": sorted(la), - "off_layers": sorted(lb)}) + "on_layers": sorted(layers_on), + "off_layers": sorted(layers_off)}) # Load OFF metadata for per-token reconstruction + alignment - mb = _load_packed_meta(dir_off) - if mb is None: + meta_off = _load_packed_meta(dir_off) + if meta_off is None: return CheckResult(name="rope_freqs", passed=False, metrics={"error": "OFF cu_seqlens missing"}) - cu_off = mb["cu_seqlens"] + cu_off = meta_off["cu_seqlens"] T_off = int(cu_off[-1]) if cu_off.numel() > 0 else 0 # Build alignment mask (prefer 2D attention_mask) mask_on_2d = _load_attention_mask_2d(dir_on) mask_off_2d = _load_attention_mask_2d(dir_off) - ma = _load_packed_meta(dir_on) + meta_on = _load_packed_meta(dir_on) if mask_on_2d is not None and mask_off_2d is not None: align_mask = _build_alignment_mask_from_2d( mask_on_2d, mask_off_2d, cu_off, T_off) - elif ma is not None: + elif meta_on is not None: align_mask = _build_alignment_mask( - cu_off, ma["prefix_lens"], T_off) + cu_off, meta_on["prefix_lens"], T_off) else: return CheckResult(name="rope_freqs", passed=False, metrics={"error": "cannot build alignment mask"}) @@ -462,18 +548,18 @@ def cmp_rope_freqs(dir_on: str, dir_off: str) -> CheckResult | None: max_diff = 0.0 mismatches: list[dict] = [] # [{layer, token_idx, dim, on_val, off_val, diff}] - for lyr in sorted(la): - on_freqs = on_dict[lyr] # [T_on, 1, 1, D] + for layer_idx in sorted(layers_on): + on_freqs = on_dict[layer_idx] # [T_on, 1, 1, D] # Reconstruct OFF per-token for this layer off_freqs = torch.cat( - [off_dict[lyr][:s, :, :, :] for s in seqlens], dim=0) # [T_off, 1, 1, D] + [off_dict[layer_idx][:s, :, :, :] for s in seqlens], dim=0) # [T_off, 1, 1, D] try: on_aligned, off_aligned = _align_packed( on_freqs, off_freqs, align_mask) except ValueError as e: return CheckResult(name="rope_freqs", passed=False, - metrics={"error": f"align failed L{lyr}: {e}"}) + metrics={"error": f"align failed L{layer_idx}: {e}"}) diff = (on_aligned - off_aligned).abs() # [N, 1, 1, D] md = float(diff.max()) @@ -487,7 +573,7 @@ def cmp_rope_freqs(dir_on: str, dir_off: str) -> CheckResult | None: t = int(t) d = int(token_diff.indices[t]) mismatches.append({ - "layer": lyr, + "layer": layer_idx, "token_idx": t, "dim": d, "on_val": float(on_aligned[t, 0, 0, d]), @@ -495,7 +581,7 @@ def cmp_rope_freqs(dir_on: str, dir_off: str) -> CheckResult | None: "diff": float(token_diff.values[t]), }) - metrics: dict = {"max_diff": max_diff, "num_layers": len(la)} + metrics: dict = {"max_diff": max_diff, "num_layers": len(layers_on)} if mismatches: metrics["mismatches"] = mismatches[:20] # cap to top 20 metrics["total_mismatches"] = len(mismatches) @@ -515,7 +601,7 @@ def _per_layer_cos(dir_on: str, dir_off: str, layer: int | None) -> dict | None: extract the matching suffix region from OFF before comparison. """ - def _cos_for_layer(a, b, lyr, need_align, align_mask): + def _cos_for_layer(a, b, layer_idx, need_align, align_mask): if a.dim() == 3: a, b = a.squeeze(1), b.squeeze(1) if need_align and a.shape[0] != b.shape[0]: @@ -534,46 +620,46 @@ def _cos_for_layer(a, b, lyr, need_align, align_mask): b = _load_attn_output(dir_off, layer) if a is None or b is None: return None - ma = _load_packed_meta(dir_on) + meta_on = _load_packed_meta(dir_on) align_mask = None - need_align = (ma is not None and a.shape[0] != b.shape[0]) + need_align = (meta_on is not None and a.shape[0] != b.shape[0]) if need_align: - mb = _load_packed_meta(dir_off) - T = int(mb["cu_seqlens"][-1]) if mb and mb["cu_seqlens"].numel() > 0 else b.shape[0] + meta_off = _load_packed_meta(dir_off) + T = int(meta_off["cu_seqlens"][-1]) if meta_off and meta_off["cu_seqlens"].numel() > 0 else b.shape[0] # Prefer 2D attention_mask.pt for exact suffix alignment (dp aligned) mask_on_2d = _load_attention_mask_2d(dir_on) mask_off_2d = _load_attention_mask_2d(dir_off) - if mask_on_2d is not None and mask_off_2d is not None and mb is not None: + if mask_on_2d is not None and mask_off_2d is not None and meta_off is not None: align_mask = _build_alignment_mask_from_2d( - mask_on_2d, mask_off_2d, mb["cu_seqlens"], T) + mask_on_2d, mask_off_2d, meta_off["cu_seqlens"], T) else: align_mask = _build_alignment_mask( - mb["cu_seqlens"], ma["prefix_lens"], T) + meta_off["cu_seqlens"], meta_on["prefix_lens"], T) d = _cos_for_layer(a, b, layer, need_align, align_mask) d["layer"] = layer return d # All-layers mode - fa = os.path.join(dir_on, "attn_outputs.pt") - fb = os.path.join(dir_off, "attn_outputs.pt") - if not os.path.exists(fa) or not os.path.exists(fb): + filepath_on = os.path.join(dir_on, "attn_outputs.pt") + filepath_off = os.path.join(dir_off, "attn_outputs.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): return None - da = torch.load(fa, weights_only=True) - db = torch.load(fb, weights_only=True) - if not isinstance(da, dict) or not isinstance(db, dict): + attn_dict_on = torch.load(filepath_on, weights_only=True) + attn_dict_off = torch.load(filepath_off, weights_only=True) + if not isinstance(attn_dict_on, dict) or not isinstance(attn_dict_off, dict): return None # Build alignment mask once from ON metadata - ma = _load_packed_meta(dir_on) + meta_on = _load_packed_meta(dir_on) align_mask = None need_align = False - if ma is not None: - mb = _load_packed_meta(dir_off) - T = int(mb["cu_seqlens"][-1]) if mb and mb["cu_seqlens"].numel() > 0 else 0 + if meta_on is not None: + meta_off = _load_packed_meta(dir_off) + T = int(meta_off["cu_seqlens"][-1]) if meta_off and meta_off["cu_seqlens"].numel() > 0 else 0 if T > 0: # Check if any layer has shape mismatch - for lyr in da: - if lyr in db and da[lyr].shape != db[lyr].shape: + for layer_idx in attn_dict_on: + if layer_idx in attn_dict_off and attn_dict_on[layer_idx].shape != attn_dict_off[layer_idx].shape: need_align = True break if need_align: @@ -582,14 +668,14 @@ def _cos_for_layer(a, b, lyr, need_align, align_mask): mask_off_2d = _load_attention_mask_2d(dir_off) if mask_on_2d is not None and mask_off_2d is not None: align_mask = _build_alignment_mask_from_2d( - mask_on_2d, mask_off_2d, mb["cu_seqlens"], T) + mask_on_2d, mask_off_2d, meta_off["cu_seqlens"], T) else: align_mask = _build_alignment_mask( - mb["cu_seqlens"], ma["prefix_lens"], T) + meta_off["cu_seqlens"], meta_on["prefix_lens"], T) results = {} - for lyr in sorted(set(da.keys()) & set(db.keys())): - results[lyr] = _cos_for_layer(da[lyr], db[lyr], lyr, need_align, align_mask) + for layer_idx in sorted(set(attn_dict_on.keys()) & set(attn_dict_off.keys())): + results[layer_idx] = _cos_for_layer(attn_dict_on[layer_idx], attn_dict_off[layer_idx], layer_idx, need_align, align_mask) return results @@ -599,8 +685,8 @@ def cmp_attn_layer(dir_on: str, dir_off: str, if r is None: return None if "cos_avg" in r: - lyr = r["layer"] - return CheckResult(name=f"attn_L{lyr}", + layer_idx = r["layer"] + return CheckResult(name=f"attn_L{layer_idx}", passed=r["cos_avg"] > 0.9999 and r["cos_min"] > 0.999, metrics=r) return CheckResult(name="attn_per_layer", passed=True, @@ -625,22 +711,22 @@ def cmp_first_token(dir_on: str, dir_off: str) -> list[CheckResult]: a = _load_attn_output(dir_on, last) b = _load_attn_output(dir_off, last) if a is not None and b is not None: - a0 = a.squeeze(1) if a.dim() == 3 else a - b0 = b.squeeze(1) if b.dim() == 3 else b - ft = _first_token_metrics(a0[0], b0[0]) - results.append(CheckResult(name="first_token_attn", metrics=ft)) + on_token0 = a.squeeze(1) if a.dim() == 3 else a + off_token0 = b.squeeze(1) if b.dim() == 3 else b + metrics = _first_token_metrics(on_token0[0], off_token0[0]) + results.append(CheckResult(name="first_token_attn", metrics=metrics)) # logits — packed[0], auto-detect [N,V] vs [V,N] format - lo = _load_tensor(dir_on, "logits.pt") - lf = _load_tensor(dir_off, "logits.pt") - if lo is not None and lf is not None: - lo_first, lf_first = _logits_first_token(lo, lf) - if lo_first is not None: - ft = _first_token_metrics(lo_first, lf_first) - results.append(CheckResult(name="first_token_logits", metrics=ft)) + logits_on = _load_tensor(dir_on, "logits.pt") + logits_off = _load_tensor(dir_off, "logits.pt") + if logits_on is not None and logits_off is not None: + logits_on_first, logits_off_first = _logits_first_token(logits_on, logits_off) + if logits_on_first is not None: + metrics = _first_token_metrics(logits_on_first, logits_off_first) + results.append(CheckResult(name="first_token_logits", metrics=metrics)) else: _log.warning("first_token_logits skipped: cannot determine token dim " - "(ON %s, OFF %s)", _fmt_shape(lo.shape), _fmt_shape(lf.shape)) + "(ON %s, OFF %s)", _fmt_shape(logits_on.shape), _fmt_shape(logits_off.shape)) return results @@ -656,22 +742,22 @@ def cmp_logits_packed(dir_on: str, dir_off: str) -> CheckResult | None: with OFF (full-sequence) logits, same alignment logic as attn_output per-layer comparison. """ - lo = _load_tensor(dir_on, "logits.pt") - lf = _load_tensor(dir_off, "logits.pt") - if lo is None or lf is None: + logits_on = _load_tensor(dir_on, "logits.pt") + logits_off = _load_tensor(dir_off, "logits.pt") + if logits_on is None or logits_off is None: return None # Logits may be [N,V] or [V,N] — ensure token-major [N,V] for alignment - lo, lf = _logits_ensure_token_major(lo, lf) + logits_on, logits_off = _logits_ensure_token_major(logits_on, logits_off) # Metadata for logits uses cu_seqlens_q_logits.pt - ma = _load_packed_meta(dir_on, "cu_seqlens_q_logits.pt") - mb = _load_packed_meta(dir_off, "cu_seqlens_q_logits.pt") - if ma is None or mb is None: + meta_on = _load_packed_meta(dir_on, "cu_seqlens_q_logits.pt") + meta_off = _load_packed_meta(dir_off, "cu_seqlens_q_logits.pt") + if meta_on is None or meta_off is None: return None - T_off = int(mb["cu_seqlens"][-1]) if mb["cu_seqlens"].numel() > 0 else 0 - if T_off == 0 or lo.shape[0] == 0 or lf.shape[0] == 0: + total_off_tokens = int(meta_off["cu_seqlens"][-1]) if meta_off["cu_seqlens"].numel() > 0 else 0 + if total_off_tokens == 0 or logits_on.shape[0] == 0 or logits_off.shape[0] == 0: return None # Build alignment mask (same logic as attn_output all-layers) @@ -679,19 +765,19 @@ def cmp_logits_packed(dir_on: str, dir_off: str) -> CheckResult | None: mask_off_2d = _load_attention_mask_2d(dir_off) if mask_on_2d is not None and mask_off_2d is not None: align_mask = _build_alignment_mask_from_2d( - mask_on_2d, mask_off_2d, mb["cu_seqlens"], T_off) + mask_on_2d, mask_off_2d, meta_off["cu_seqlens"], total_off_tokens) else: align_mask = _build_alignment_mask( - mb["cu_seqlens"], ma["prefix_lens"], T_off) + meta_off["cu_seqlens"], meta_on["prefix_lens"], total_off_tokens) # Align: extract suffix-only region from OFF try: - on_aligned, off_aligned = _align_packed(lo, lf, align_mask) + on_aligned, off_aligned = _align_packed(logits_on, logits_off, align_mask) except ValueError as e: return CheckResult(name="logits", passed=False, metrics={"error": str(e), - "n_on": lo.shape[0], - "n_off": lf.shape[0]}) + "n_on": logits_on.shape[0], + "n_off": logits_off.shape[0]}) n_tokens = on_aligned.shape[0] @@ -775,7 +861,7 @@ def _print_pos_ids(r: CheckResult): def _print_rope(r: CheckResult): - print(_SEP_SINGLE + f"\n [rope_emb] {_CHECK if r.passed else _CROSS} {'PASS' if r.passed else 'FAIL'}") + print(_SEP_SINGLE + f"\n [rope_postqk] {_CHECK if r.passed else _CROSS} {'PASS' if r.passed else 'FAIL'}") print(_SEP_SINGLE) m = r.metrics if "error" in m: @@ -818,13 +904,13 @@ def _print_per_layer(r: CheckResult): print(f" {'LAYER':>6s} {'COS_AVG':>14s} {'COS_MIN':>14s} {'TOKENS':>8s} {'STATUS':>8s}") print(f" {'─'*6} {'─'*14} {'─'*14} {'─'*8} {'─'*8}") bad = [] - for lyr in sorted(layers.keys()): - d = layers[lyr] + for layer_idx in sorted(layers.keys()): + d = layers[layer_idx] ok = d["cos_avg"] > 0.9999 and d["cos_min"] > 0.999 - print(f" {lyr:>6d} {d['cos_avg']:>14.6e} {d['cos_min']:>14.6e} " + print(f" {layer_idx:>6d} {d['cos_avg']:>14.6e} {d['cos_min']:>14.6e} " f"{d['n_tokens']:>8d} {'PASS' if ok else 'WARN':>8s}") if not ok: - bad.append(lyr) + bad.append(layer_idx) if bad: print(f"\n ⚠ First deviating layer: {bad[0]}") elif "cos_avg" in r.metrics: @@ -1024,7 +1110,7 @@ def main(): # ── ②b RoPE encoding (post-apply rotated Q/K) ── if not stop: - r = cmp_rope_emb(args.dir_on, args.dir_off) + r = cmp_rope_postqk(args.dir_on, args.dir_off) if r: all_results.append(r) _print_rope(r) @@ -1052,20 +1138,20 @@ def main(): a = _load_attn_output(args.dir_on, last) b = _load_attn_output(args.dir_off, last) if a is not None and b is not None: - a0 = (a.squeeze(1) if a.dim() == 3 else a)[0].cpu() - b0 = (b.squeeze(1) if b.dim() == 3 else b)[0].cpu() - _print_topk_vec(a0, b0, args.topk, "val", + on_token0 = (a.squeeze(1) if a.dim() == 3 else a)[0].cpu() + off_token0 = (b.squeeze(1) if b.dim() == 3 else b)[0].cpu() + _print_topk_vec(on_token0, off_token0, args.topk, "val", "first_token_attn") # first_token_logits — per-dim top-K (real = ON signed value, # so positive logits — the tokens actually selectable by # sampling — surface first, not the large-magnitude negatives) - lo = _load_tensor(args.dir_on, "logits.pt") - lf = _load_tensor(args.dir_off, "logits.pt") - if lo is not None and lf is not None: - ft = _logits_first_token(lo, lf) - if ft is not None: - _print_topk_vec(ft[0].cpu(), ft[1].cpu(), args.topk, - "real", "first_token_logits") + logits_on = _load_tensor(args.dir_on, "logits.pt") + logits_off = _load_tensor(args.dir_off, "logits.pt") + if logits_on is not None and logits_off is not None: + first_token_pair = _logits_first_token(logits_on, logits_off) + if first_token_pair is not None: + _print_topk_vec(first_token_pair[0].cpu(), first_token_pair[1].cpu(), + args.topk, "real", "first_token_logits") # ── ④b Logits (full packed alignment via attention_mask + dual pointer) ── if not stop: diff --git a/prefix-sharing/prefix_sharing/tools/cmp_diag_verl080.py b/prefix-sharing/prefix_sharing/tools/cmp_diag_verl080.py index 259ffb86..c30f0bfd 100644 --- a/prefix-sharing/prefix_sharing/tools/cmp_diag_verl080.py +++ b/prefix-sharing/prefix_sharing/tools/cmp_diag_verl080.py @@ -7,7 +7,7 @@ **packed(suffix 对齐)** - attention_output per-layer cos 每层 attention 输出余弦相似度 - - first_token packed[0](attn[0] + logits[0]) + - packed_token packed[pos](attn[pos] + logits[pos],--token 指定 pos,默认 0) - logits packed 全 packed logits suffix 对齐对比 **2D(v080 特有,restore 后 ``[B, L_max]``)** @@ -25,21 +25,22 @@ label_mask_{tag}.pt [B, L_max] bool [prompt-last,L_i-1) PPO loss 范围 logits.pt [N, V//tp] packed logits(ON 裁剪后 / OFF 完整) attn_outputs.pt dict {layer: [N, hidden]} per-layer packed attn output - rope_freqs_on.pt dict {layer: [T_on,1,1,D]} ON per-token RoPE 角度 - rope_freqs_off.pt dict {layer: [L0,1,1,D]} OFF raw 角度表(freqs[p]=p*inv_freq) + rope_freqs.pt dict {layer: [T,1,1,D]} per-token RoPE 角度(ON/OFF 同款) + rope_preqk.pt dict {layer: [T,H,D]} 旋转前 Q/K(pre-RoPE) + rope_postqk.pt dict {layer: [T,H,D]} 旋转后 Q/K(post-RoPE) prefix_lens.pt [B] ON=plan.prefix_lens / OFF=全0 cu_seqlens_q.pt [B+1] NestedTensor offsets(ON 裁剪后 / OFF 完整) cu_seqlens_q_logits.pt [B+1] logits packed 边界(同上) Usage: - # 完整对比(attn per-layer + first_token + logits + logprobs + entropy) + # 完整对比(attn per-layer + packed_token + logits + logprobs + entropy) python cmp_diag_verl080.py --dir-on ./dump_on --dir-off ./dump_off --tag old # 只看某一层 attention(1-indexed) python cmp_diag_verl080.py --dir-on ./dump_on --dir-off ./dump_off \\ --tag old --layer 12 - # top-K 误差最大位置(2D + first_token) + # top-K 误差最大位置(2D + packed_token) python cmp_diag_verl080.py --dir-on ./dump_on --dir-off ./dump_off \\ --tag old --topk 20 @@ -132,7 +133,7 @@ def _pearson_r(t1: torch.Tensor, t2: torch.Tensor, return float("nan") if sx == 0 or sy == 0 else float(cov / (sx * sy)) -def _first_token_metrics(a_vec: torch.Tensor, b_vec: torch.Tensor) -> dict: +def _vec_metrics(a_vec: torch.Tensor, b_vec: torch.Tensor) -> dict: err = _error_abs_rel(a_vec, b_vec) cos = float(_cosine_sim(a_vec, b_vec, dim=-1)) pr = _pearson_r(a_vec, b_vec) @@ -146,14 +147,23 @@ def _first_token_metrics(a_vec: torch.Tensor, b_vec: torch.Tensor) -> dict: # ════════════════════════════════════════════════════════════════ def _load_tensor(dir_path: str, filename: str) -> torch.Tensor | None: - fp = os.path.join(dir_path, filename) - return torch.load(fp, weights_only=True).float() if os.path.exists(fp) else None + filepath = os.path.join(dir_path, filename) + return torch.load(filepath, weights_only=True).float() if os.path.exists(filepath) else None + + +def _load_logits(dir_path: str) -> torch.Tensor | None: + """Load packed logits from ``logits.pt``. + + Multi-rank assembly (tp vocab concat) is done by ``assemble_dump.py`` + before cmp is called; cmp works on flat single-card data only. + """ + return _load_tensor(dir_path, "logits.pt") def _load_manifest(dir_path: str) -> dict | None: """Load ``parallel_info.json`` written by the dump layer (topology + scopes). - Returns None when absent (single-card or pre-manifest dumps) → callers fall + Returns None when absent (single-card or pre-manifest dumps) -> callers fall back to tp_size==1 behavior (plain filenames, single-card compatible). """ fp = os.path.join(dir_path, "parallel_info.json") @@ -166,60 +176,36 @@ def _load_manifest(dir_path: str) -> dict | None: return None -def _load_logits(dir_path: str, manifest: dict | None = None) -> torch.Tensor | None: - """Load packed logits, gathering tp vocab shards to full vocab when tp>1. - - tp_size==1 (or no manifest) → single ``logits.pt`` (single-card compatible). - tp_size>1 → concat ``logits_tp{0..tp-1}.pt`` on the vocab (last) dim, - reconstructing ``[N, V]`` so ON-vs-OFF compares on the same - full-vocab coordinate system as single-card. A missing shard - aborts the reconstruction (returns None) rather than silently - comparing partial vocab. - """ - if manifest is None: - manifest = _load_manifest(dir_path) - tp_size = (manifest or {}).get("tp_size", 1) - if tp_size <= 1: - return _load_tensor(dir_path, "logits.pt") - shards = [] - for t in range(tp_size): - s = _load_tensor(dir_path, f"logits_tp{t}.pt") - if s is None: - return None - shards.append(s) - return torch.cat(shards, dim=-1) - - def _load_packed_meta(dir_path: str, cu_fname: str = "cu_seqlens_q.pt") -> dict | None: """加载 cu_seqlens + prefix_lens(suffix 对齐所需)。""" - fp = os.path.join(dir_path, cu_fname) - if not os.path.exists(fp): - fp = os.path.join(dir_path, "cu_seqlens_q.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, cu_fname) + if not os.path.exists(filepath): + filepath = os.path.join(dir_path, "cu_seqlens_q.pt") + if not os.path.exists(filepath): return None - pl_fp = os.path.join(dir_path, "prefix_lens.pt") - if not os.path.exists(pl_fp): + prefix_lens_filepath = os.path.join(dir_path, "prefix_lens.pt") + if not os.path.exists(prefix_lens_filepath): return None - return {"cu_seqlens": torch.load(fp, weights_only=True), - "prefix_lens": torch.load(pl_fp, weights_only=True)} + return {"cu_seqlens": torch.load(filepath, weights_only=True), + "prefix_lens": torch.load(prefix_lens_filepath, weights_only=True)} def _load_attn_output(dir_path: str, layer: int) -> torch.Tensor | None: """加载单层 attn_output(attn_outputs.pt = dict {layer: tensor})。""" - fp = os.path.join(dir_path, "attn_outputs.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, "attn_outputs.pt") + if not os.path.exists(filepath): return None - d = torch.load(fp, weights_only=True) - return d.get(layer) if isinstance(d, dict) else None + attn_dict = torch.load(filepath, weights_only=True) + return attn_dict.get(layer) if isinstance(attn_dict, dict) else None def _get_num_layers(dir_path: str) -> int: - fp = os.path.join(dir_path, "attn_outputs.pt") - if not os.path.exists(fp): + filepath = os.path.join(dir_path, "attn_outputs.pt") + if not os.path.exists(filepath): return 0 - d = torch.load(fp, weights_only=True) - return max(d.keys()) if isinstance(d, dict) and d else 0 + attn_dict = torch.load(filepath, weights_only=True) + return max(attn_dict.keys()) if isinstance(attn_dict, dict) and attn_dict else 0 # ════════════════════════════════════════════════════════════════ @@ -268,23 +254,52 @@ def _align_packed(on_tensor: torch.Tensor, off_tensor: torch.Tensor, # Logits helpers # ════════════════════════════════════════════════════════════════ -def _logits_first_token(lo: torch.Tensor, lf: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - """取 packed[0] 的 full-vocab 向量(logits 词表恒在最后一维)。""" - lo_2d = lo.reshape(-1, lo.size(-1)) - lf_2d = lf.reshape(-1, lf.size(-1)) - return lo_2d[0, :].contiguous(), lf_2d[0, :].contiguous() +def _aligned_vec_at_pos( + on_tensor: torch.Tensor | None, + off_tensor: torch.Tensor | None, + is_attn: bool, + pos: int, + align_mask: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor] | None: + """ON(suffix-only)/OFF(full-packed) 的 packed 张量 **suffix 对齐后** 取 [pos]。 + + ON 物理裁剪后只含 suffix,OFF 含完整序列,两者 token 不直接对应——必须先用 + align_mask 把 OFF 的 suffix 段抽出来与 ON 对齐,再取 [pos]。pos 索引的是 + 对齐后的 suffix-packed 空间(ON/OFF 一致,指向同一个 token)。 + + - is_attn=True:attn_output ``[T,1,hidden]`` → ``[T,hidden]``。 + - is_attn=False:logits → ``[N, V]``(vocab 恒在最后一维)。 + 返回 (on_vec, off_vec)(同 token、同向量长度),或 None(数据缺失 / pos 越界 / + 对齐失败)。 + """ + if on_tensor is None or off_tensor is None: + return None + if is_attn: + on = on_tensor.squeeze(1) if on_tensor.dim() == 3 else on_tensor + off = off_tensor.squeeze(1) if off_tensor.dim() == 3 else off_tensor + else: + on = on_tensor.reshape(-1, on_tensor.size(-1)) + off = off_tensor.reshape(-1, off_tensor.size(-1)) + if align_mask is not None and on.shape[0] != off.shape[0]: + try: + on, off = _align_packed(on, off, align_mask) + except ValueError: + return None + n = min(on.shape[0], off.shape[0]) + if pos < 0 or pos >= n: + return None + return on[pos].contiguous(), off[pos].contiguous() -def _logits_ensure_token_major(lo: torch.Tensor, lf: torch.Tensor +def _logits_ensure_token_major(logits_on: torch.Tensor, logits_off: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor]: """确保 logits 为 2D [N, V](token-major),vocab 在最后一维。""" - return (lo.reshape(-1, lo.size(-1)).contiguous(), - lf.reshape(-1, lf.size(-1)).contiguous()) + return (logits_on.reshape(-1, logits_on.size(-1)).contiguous(), + logits_off.reshape(-1, logits_off.size(-1)).contiguous()) # ════════════════════════════════════════════════════════════════ -# Packed compare: attention_output / first_token / logits +# Packed compare: attention_output / packed_token / logits # ════════════════════════════════════════════════════════════════ def _cos_for_layer(a: torch.Tensor, b: torch.Tensor, @@ -301,15 +316,15 @@ def _cos_for_layer(a: torch.Tensor, b: torch.Tensor, def _build_attn_align_mask(dir_on: str, dir_off: str) -> torch.Tensor | None: """从 OFF cu_seqlens + ON prefix_lens 构建 suffix 对齐 mask(None=无法构建)。""" - ma = _load_packed_meta(dir_on) - mb = _load_packed_meta(dir_off) - if ma is None or mb is None: + meta_on = _load_packed_meta(dir_on) + meta_off = _load_packed_meta(dir_off) + if meta_on is None or meta_off is None: return None - cu_off = mb["cu_seqlens"] + cu_off = meta_off["cu_seqlens"] T = int(cu_off[-1]) if cu_off.numel() > 0 else 0 if T == 0: return None - return _build_alignment_mask(cu_off, ma["prefix_lens"], T) + return _build_alignment_mask(cu_off, meta_on["prefix_lens"], T) def cmp_attn_layer(dir_on: str, dir_off: str, @@ -336,79 +351,309 @@ def cmp_attn_layer(dir_on: str, dir_off: str, passed=d["cos_avg"] > _COS_AVG_PASS and d["cos_min"] > _COS_MIN_PASS, metrics=d) - fa = os.path.join(dir_on, "attn_outputs.pt") - fb = os.path.join(dir_off, "attn_outputs.pt") - if not os.path.exists(fa) or not os.path.exists(fb): + filepath_on = os.path.join(dir_on, "attn_outputs.pt") + filepath_off = os.path.join(dir_off, "attn_outputs.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): return None - da = torch.load(fa, weights_only=True) - db = torch.load(fb, weights_only=True) - if not isinstance(da, dict) or not isinstance(db, dict): + attn_dict_on = torch.load(filepath_on, weights_only=True) + attn_dict_off = torch.load(filepath_off, weights_only=True) + if not isinstance(attn_dict_on, dict) or not isinstance(attn_dict_off, dict): return None results = {} - for lyr in sorted(set(da.keys()) & set(db.keys())): - a, b = da[lyr], db[lyr] + for layer_idx in sorted(set(attn_dict_on.keys()) & set(attn_dict_off.keys())): + a, b = attn_dict_on[layer_idx], attn_dict_off[layer_idx] need = align_mask is not None and a.shape[0] != b.shape[0] try: - results[lyr] = _cos_for_layer(a, b, align_mask if need else None) + results[layer_idx] = _cos_for_layer(a, b, align_mask if need else None) except ValueError as e: - results[lyr] = {"error": str(e)} + results[layer_idx] = {"error": str(e)} return CheckResult(name="attn_per_layer", passed=True, metrics={"layers": results}) -def cmp_first_token(dir_on: str, dir_off: str) -> list[CheckResult]: - """packed[0] 对比:最后一层 attn[0] + logits[0]。 +def cmp_packed_token(dir_on: str, dir_off: str, + pos: int = 0, layer: int | None = None, + align_mask: torch.Tensor | None = None) -> list[CheckResult]: + """packed[pos] 对比(**suffix 对齐后**):attn[pos](可指定层)+ logits[pos](仅最后一层)。 + + ON 是裁剪后的 suffix-only packed,OFF 是完整 packed,两者 token **不直接对应**—— + 必须先用 align_mask(OFF cu_seqlens + ON prefix_lens)把 OFF 的 suffix 段抽出来 + 与 ON 对齐,再取 [pos]。pos 索引的是对齐后的 suffix-packed 空间(ON/OFF 一致)。 - 第一个序列(row 0)永远是 provider(完整序列),packed[0] 是完整 suffix token, - 可直接对比无需对齐。 + - pos:对齐后 suffix-packed 里的位置(单个 int,默认 0)。 + - attn:用 *layer*(默认最后一层)。对比第 1 层可区分 + "结构错(第 1 层就偏)" vs "数值累积(第 1 层完美、深层才偏)"。 + - logits:永远最后一层。 + - align_mask:可选,复用调用方已构建的;None 则内部构建。 """ + if align_mask is None: + align_mask = _build_attn_align_mask(dir_on, dir_off) results: list[CheckResult] = [] - last = _get_num_layers(dir_on) or _get_num_layers(dir_off) - - if last: - a = _load_attn_output(dir_on, last) - b = _load_attn_output(dir_off, last) - if a is not None and b is not None: - a0 = a.squeeze(1) if a.dim() == 3 else a - b0 = b.squeeze(1) if b.dim() == 3 else b + attn_layer = layer if layer is not None else ( + _get_num_layers(dir_on) or _get_num_layers(dir_off)) + + if attn_layer: + a = _load_attn_output(dir_on, attn_layer) + b = _load_attn_output(dir_off, attn_layer) + vecs = _aligned_vec_at_pos(a, b, True, pos, align_mask) + if vecs is None: results.append(CheckResult( - name="first_token_attn", - metrics=_first_token_metrics(a0[0], b0[0]))) + name=f"attn_L{attn_layer}_pos{pos}", + metrics={"error": f"无法对齐或 pos {pos} 越界"})) + else: + results.append(CheckResult( + name=f"attn_L{attn_layer}_pos{pos}", + metrics=_vec_metrics(vecs[0], vecs[1]))) - lo = _load_logits(dir_on) - lf = _load_logits(dir_off) - if lo is not None and lf is not None: - lo_first, lf_first = _logits_first_token(lo, lf) + logits_on = _load_logits(dir_on) + logits_off = _load_logits(dir_off) + vecs = _aligned_vec_at_pos(logits_on, logits_off, False, pos, align_mask) + if vecs is None: + results.append(CheckResult( + name=f"logits_pos{pos}", + metrics={"error": f"无法对齐或 pos {pos} 越界"})) + else: results.append(CheckResult( - name="first_token_logits", - metrics=_first_token_metrics(lo_first, lf_first))) + name=f"logits_pos{pos}", + metrics=_vec_metrics(vecs[0], vecs[1]))) return results +# ══════════════════════════════════════════════════════════════════ +# Post-RoPE Q/K compare: per-layer + packed_token +# ══════════════════════════════════════════════════════════════════ + +# RoPE 对比阶段:**先 pre(旋转前,rope_preqk.pt)后 post(旋转后,rope_postqk.pt)**。 +# (stage, fname, label) — label 用作结果名前缀与打印 section 头。 +_ROPE_STAGES: list[tuple[str, str, str]] = [ + ("pre", "rope_preqk.pt", "rope_preqk"), + ("post", "rope_postqk.pt", "rope_postqk"), +] + + +def _load_rope_postqk(dir_path: str, layer: int, fname: str = "rope_postqk.pt" + ) -> tuple[torch.Tensor | None, torch.Tensor | None]: + """Load Q/K for a single layer from ``fname`` (rope_postqk.pt=post, rope_preqk.pt=pre). + + Returns ``(query, key)`` or ``(None, None)``. + """ + filepath = os.path.join(dir_path, fname) + if not os.path.exists(filepath): + return None, None + d = torch.load(filepath, weights_only=True) + if not isinstance(d, dict): + return None, None + entry = d.get(layer) + if entry is None: + return None, None + return entry.get("query"), entry.get("key") + + +def _rope_postqk_cos_for_layer(qa: torch.Tensor, ka: torch.Tensor, + qb: torch.Tensor, kb: torch.Tensor, + align_mask: torch.Tensor | None = None) -> dict: + """单层 Q/K suffix 对齐 + per-token cosine(Q 和 K 分别算)。""" + # Q/K shape: [T, H, D] → 压平 head*dim 维度算 cosine + qa_flat = qa.reshape(qa.shape[0], -1) + qb_flat = qb.reshape(qb.shape[0], -1) + ka_flat = ka.reshape(ka.shape[0], -1) + kb_flat = kb.reshape(kb.shape[0], -1) + + if align_mask is not None and qa.shape[0] != qb.shape[0]: + qa_flat, qb_flat = _align_packed(qa_flat, qb_flat, align_mask) + ka_flat, kb_flat = _align_packed(ka_flat, kb_flat, align_mask) + + q_cos = _cosine_sim(qa_flat, qb_flat, dim=-1) + k_cos = _cosine_sim(ka_flat, kb_flat, dim=-1) + return { + "n_tokens": qa_flat.shape[0], + "Q_cos_avg": float(q_cos.mean()), "Q_cos_min": float(q_cos.min()), + "K_cos_avg": float(k_cos.mean()), "K_cos_min": float(k_cos.min()), + "Q_max_diff": float((qa_flat - qb_flat).abs().max()), + "K_max_diff": float((ka_flat - kb_flat).abs().max()), + } + + +def _cmp_rope_stage_layer(dir_on: str, dir_off: str, layer: int | None, + fname: str, label: str) -> CheckResult | None: + """单 stage(fname/label)的 Q/K per-layer cosine(suffix 对齐)。""" + align_mask = _build_attn_align_mask(dir_on, dir_off) + + if layer is not None: + q_on, k_on = _load_rope_postqk(dir_on, layer, fname) + q_off, k_off = _load_rope_postqk(dir_off, layer, fname) + if q_on is None or q_off is None: + return None + need = align_mask is not None and q_on.shape[0] != q_off.shape[0] + try: + d = _rope_postqk_cos_for_layer(q_on, k_on, q_off, k_off, + align_mask if need else None) + except ValueError as e: + return CheckResult(name=f"{label}_L{layer}", passed=False, + metrics={"error": str(e)}) + d["layer"] = layer + ok = (d["Q_cos_avg"] > _COS_AVG_PASS and d["Q_cos_min"] > _COS_MIN_PASS + and d["K_cos_avg"] > _COS_AVG_PASS and d["K_cos_min"] > _COS_MIN_PASS) + return CheckResult(name=f"{label}_L{layer}", passed=ok, metrics=d) + + # All layers + filepath_on = os.path.join(dir_on, fname) + filepath_off = os.path.join(dir_off, fname) + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): + return None + attn_dict_on = torch.load(filepath_on, weights_only=True) + attn_dict_off = torch.load(filepath_off, weights_only=True) + if not isinstance(attn_dict_on, dict) or not isinstance(attn_dict_off, dict): + return None + + results = {} + for layer_idx in sorted(set(attn_dict_on.keys()) & set(attn_dict_off.keys())): + ea, eb = attn_dict_on[layer_idx], attn_dict_off[layer_idx] + qa, ka = ea.get("query"), ea.get("key") + qb, kb = eb.get("query"), eb.get("key") + if qa is None or qb is None: + continue + need = align_mask is not None and qa.shape[0] != qb.shape[0] + try: + results[layer_idx] = _rope_postqk_cos_for_layer(qa, ka, qb, kb, + align_mask if need else None) + except ValueError as e: + results[layer_idx] = {"error": str(e)} + return CheckResult(name=f"{label}_per_layer", passed=True, + metrics={"layers": results}) + + +def cmp_rope_postqk_layer(dir_on: str, dir_off: str, layer: int | None, + stage: str = "post") -> CheckResult | None: + """Q/K per-layer cosine(suffix 对齐),单 stage。 + + stage="pre" → rope_preqk.pt(旋转前),stage="post" → rope_postqk.pt(旋转后)。 + 调用方按 pre → rope_freqs → post 顺序分别调用,便于定位分歧出现在 RoPE 哪一步。 + """ + if stage == "pre": + return _cmp_rope_stage_layer(dir_on, dir_off, layer, "rope_preqk.pt", "rope_preqk") + return _cmp_rope_stage_layer(dir_on, dir_off, layer, "rope_postqk.pt", "rope_postqk") + + +def _rope_postqk_vec_at_pos(q_on: torch.Tensor | None, k_on: torch.Tensor | None, + q_off: torch.Tensor | None, k_off: torch.Tensor | None, + pos: int, + align_mask: torch.Tensor | None + ) -> tuple[torch.Tensor, torch.Tensor, + torch.Tensor, torch.Tensor] | None: + """Q/K suffix 对齐后取 [pos],返回 (q_on, q_off, k_on, k_off) 四向量。 + + 每个向量压平 [H*D],可直接做 vec_metrics 对比。 + """ + if q_on is None or q_off is None: + return None + q_on_f = q_on.reshape(q_on.shape[0], -1) + q_off_f = q_off.reshape(q_off.shape[0], -1) + k_on_f = k_on.reshape(k_on.shape[0], -1) if k_on is not None else None + k_off_f = k_off.reshape(k_off.shape[0], -1) if k_off is not None else None + + if align_mask is not None and q_on.shape[0] != q_off.shape[0]: + try: + q_on_f, q_off_f = _align_packed(q_on_f, q_off_f, align_mask) + if k_on_f is not None: + k_on_f, k_off_f = _align_packed(k_on_f, k_off_f, align_mask) + except ValueError: + return None + n = min(q_on_f.shape[0], q_off_f.shape[0]) + if pos < 0 or pos >= n: + return None + qo = q_on_f[pos].contiguous() + qf = q_off_f[pos].contiguous() + ko = k_on_f[pos].contiguous() if k_on_f is not None else None + kf = k_off_f[pos].contiguous() if k_off_f is not None else None + return qo, qf, ko, kf + + +def _diag_rope_pos_fail(q_on: torch.Tensor | None, q_off: torch.Tensor | None, + pos: int, align_mask: torch.Tensor | None) -> str: + """rope_postqk packed_token 取 [pos] 失败时的诊断串:区分 缺失 / 对齐失败 / pos 越界。""" + if q_on is None or q_off is None: + return f"rope_postqk 该层在 {'ON' if q_on is None else 'OFF'} 侧缺失" + n_on, n_off = q_on.shape[0], q_off.shape[0] + if align_mask is not None and n_on != n_off: + msum = int(align_mask.sum()) + return (f"对齐失败: n_on={n_on} n_off={n_off} " + f"align_mask(len={align_mask.shape[0]}, sum={msum}); " + f"需 ON tokens==sum({msum}) 且 mask_len==n_off({n_off})") + post = min(n_on, n_off) + return f"pos {pos} 越界: 对齐后 token 数={post} (n_on={n_on}, n_off={n_off})" + + +def _cmp_rope_stage_token(dir_on: str, dir_off: str, pos: int, layer: int | None, + align_mask: torch.Tensor | None, fname: str, + label: str) -> list[CheckResult]: + """单 stage(fname/label)的 Q/K packed[pos](suffix 对齐后)。""" + if align_mask is None: + align_mask = _build_attn_align_mask(dir_on, dir_off) + results: list[CheckResult] = [] + rope_layer = layer if layer is not None else ( + _get_num_layers(dir_on) or _get_num_layers(dir_off)) + if rope_layer: + q_on, k_on = _load_rope_postqk(dir_on, rope_layer, fname) + q_off, k_off = _load_rope_postqk(dir_off, rope_layer, fname) + vecs = _rope_postqk_vec_at_pos(q_on, k_on, q_off, k_off, pos, align_mask) + if vecs is None: + results.append(CheckResult( + name=f"{label}_L{rope_layer}_pos{pos}", + metrics={"error": _diag_rope_pos_fail(q_on, q_off, pos, align_mask)})) + else: + qo, qf, ko, kf = vecs + results.append(CheckResult( + name=f"{label}_L{rope_layer}_Q_pos{pos}", + metrics=_vec_metrics(qo, qf))) + if ko is not None and kf is not None: + results.append(CheckResult( + name=f"{label}_L{rope_layer}_K_pos{pos}", + metrics=_vec_metrics(ko, kf))) + return results + + +def cmp_rope_postqk_token(dir_on: str, dir_off: str, + pos: int = 0, layer: int | None = None, + align_mask: torch.Tensor | None = None, + stage: str = "post") -> list[CheckResult]: + """Q/K packed[pos] 对比(**suffix 对齐后**),单 stage。 + + stage="pre" → rope_preqk.pt(旋转前),stage="post" → rope_postqk.pt(旋转后)。 + 对 Q、K 分别输出 {label}_L{layer_idx}_Q_pos{pos} / {label}_L{layer_idx}_K_pos{pos}。 + 调用方按 pre → rope_freqs → post 顺序分别调用。 + """ + if stage == "pre": + return _cmp_rope_stage_token(dir_on, dir_off, pos, layer, align_mask, + "rope_preqk.pt", "rope_preqk") + return _cmp_rope_stage_token(dir_on, dir_off, pos, layer, align_mask, + "rope_postqk.pt", "rope_postqk") + + def cmp_logits_packed(dir_on: str, dir_off: str) -> CheckResult | None: """全 packed logits suffix 对齐 + per-token cosine。""" - lo = _load_logits(dir_on) - lf = _load_logits(dir_off) - if lo is None or lf is None: + logits_on = _load_logits(dir_on) + logits_off = _load_logits(dir_off) + if logits_on is None or logits_off is None: return None - lo, lf = _logits_ensure_token_major(lo, lf) + logits_on, logits_off = _logits_ensure_token_major(logits_on, logits_off) - ma = _load_packed_meta(dir_on, "cu_seqlens_q_logits.pt") - mb = _load_packed_meta(dir_off, "cu_seqlens_q_logits.pt") - if ma is None or mb is None: + meta_on = _load_packed_meta(dir_on, "cu_seqlens_q_logits.pt") + meta_off = _load_packed_meta(dir_off, "cu_seqlens_q_logits.pt") + if meta_on is None or meta_off is None: return None - T_off = int(mb["cu_seqlens"][-1]) if mb["cu_seqlens"].numel() > 0 else 0 - if T_off == 0 or lo.shape[0] == 0 or lf.shape[0] == 0: + total_off_tokens = int(meta_off["cu_seqlens"][-1]) if meta_off["cu_seqlens"].numel() > 0 else 0 + if total_off_tokens == 0 or logits_on.shape[0] == 0 or logits_off.shape[0] == 0: return None - align_mask = _build_alignment_mask(mb["cu_seqlens"], ma["prefix_lens"], T_off) + align_mask = _build_alignment_mask(meta_off["cu_seqlens"], meta_on["prefix_lens"], total_off_tokens) try: - on_aligned, off_aligned = _align_packed(lo, lf, align_mask) + on_aligned, off_aligned = _align_packed(logits_on, logits_off, align_mask) except ValueError as e: return CheckResult(name="logits", passed=False, metrics={"error": str(e), - "n_on": lo.shape[0], "n_off": lf.shape[0]}) + "n_on": logits_on.shape[0], "n_off": logits_off.shape[0]}) cos = _cosine_sim(on_aligned, off_aligned, dim=-1) cos_avg, cos_min = float(cos.mean()), float(cos.min()) @@ -418,86 +663,420 @@ def cmp_logits_packed(dir_on: str, dir_off: str) -> CheckResult | None: "cos_avg": cos_avg, "cos_min": cos_min}) -def cmp_rope_freqs(dir_on: str, dir_off: str) -> CheckResult | None: - """对比 pre-RoPE 角度表(angle table,非 cos/sin)— suffix 对齐。 +def _align_rope_freqs_layer(on_freqs: torch.Tensor, off_freqs: torch.Tensor, + align_mask: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor] | None: + """单层 rope_freqs(per-token [T,1,1,D])suffix 对齐。 - ON ``rope_freqs_on.pt``: per-token 角度 dict {layer: [T_on, 1, 1, D]} - (已 index_select 到 packed_position_ids,每 token 实际旋转角度) - OFF ``rope_freqs_off.pt``: raw 角度表 dict {layer: [L0, 1, 1, D]} - (freqs[p] = p * inv_freq,未切片) + 返回 (on_aligned, off_aligned) [N,1,1,D];对齐失败返回 None。 + 供 cmp_rope_freqs(per-layer max_diff)与 cmp_rope_freqs_token([pos] 角度向量)复用。 + ON/OFF 现在都是 per-token,直接对齐即可(不再从 raw 表重建)。 + """ + try: + return _align_packed(on_freqs, off_freqs, align_mask) + except ValueError: + return None - OFF per-token 角度从 raw 表按 cu_seqlens_off 重建(每段取 ``[:seg_len]``), - 再与 ON 用同一 suffix 对齐(cu_seqlens + prefix_lens)后逐元素比 max_diff。 - 角度是 RoPE 的输入,应精确相等(``max_diff == 0``)。 + +def _load_rope_freqs_vec_at_pos(dir_on: str, dir_off: str, layer: int, pos: int, + align_mask: torch.Tensor | None = None + ) -> tuple[torch.Tensor | None, torch.Tensor | None]: + """加载 rope_freqs 对齐后 [pos] 的角度向量 [D],返回 (on_vec, off_vec) 或 (None, None)。 + + 供 top-K 跨 stage 对齐用(freqs dim = Q/K dim % D,角度按 head_dim 共享)。 + """ + filepath_on = os.path.join(dir_on, "rope_freqs.pt") + filepath_off = os.path.join(dir_off, "rope_freqs.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): + return None, None + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) + if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): + return None, None + if layer not in on_dict or layer not in off_dict: + return None, None + if align_mask is None: + align_mask = _build_attn_align_mask(dir_on, dir_off) + if align_mask is None: + return None, None + aligned_result = _align_rope_freqs_layer(on_dict[layer], off_dict[layer], align_mask) + if aligned_result is None: + return None, None + on_a, off_a = aligned_result + if pos < 0 or pos >= on_a.shape[0]: + return None, None + return on_a[pos].reshape(-1), off_a[pos].reshape(-1) + + +def cmp_rope_freqs(dir_on: str, dir_off: str, + layer: int | None = None) -> CheckResult | None: + """对比 per-token RoPE 角度 — suffix 对齐,应精确相等 max_diff==0。 + + ON/OFF 都存 per-token 角度 ``rope_freqs.pt`` {layer: [T,1,1,D]}(cos/sin 之前), + suffix 对齐后逐元素比。角度是 RoPE 输入,应精确相等(max_diff==0)。 + ``layer`` 给定则只比该层。 """ - fa = os.path.join(dir_on, "rope_freqs_on.pt") - fb = os.path.join(dir_off, "rope_freqs_off.pt") - if not os.path.exists(fa) or not os.path.exists(fb): + filepath_on = os.path.join(dir_on, "rope_freqs.pt") + filepath_off = os.path.join(dir_off, "rope_freqs.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): return None - on_dict = torch.load(fa, weights_only=True) - off_dict = torch.load(fb, weights_only=True) + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): return None - la, lb = set(on_dict.keys()), set(off_dict.keys()) - if la != lb: - return CheckResult(name="rope_freqs", passed=False, - metrics={"error": "layer set mismatch", - "on_layers": sorted(la), - "off_layers": sorted(lb)}) - - mb = _load_packed_meta(dir_off) - if mb is None: - return CheckResult(name="rope_freqs", passed=False, - metrics={"error": "OFF cu_seqlens missing"}) - cu_off = mb["cu_seqlens"] - T_off = int(cu_off[-1]) if cu_off.numel() > 0 else 0 - - ma = _load_packed_meta(dir_on) - if ma is None: - return CheckResult(name="rope_freqs", passed=False, - metrics={"error": "ON prefix_lens missing"}) - align_mask = _build_alignment_mask(cu_off, ma["prefix_lens"], T_off) + layers = sorted(set(on_dict.keys()) & set(off_dict.keys())) + if layer is not None: + layers = [l for l in layers if l == layer] + result_name = f"rope_freqs_L{layer}" if layer is not None else "rope_freqs" + if not layers: + return CheckResult(name=result_name, passed=False, + metrics={"error": f"layer {layer} 不在双方 rope_freqs 中"}) - seqlens = (cu_off[1:] - cu_off[:-1]).tolist() + align_mask = _build_attn_align_mask(dir_on, dir_off) + if align_mask is None: + return CheckResult(name=result_name, passed=False, + metrics={"error": "cu_seqlens/prefix_lens 缺失"}) max_diff = 0.0 mismatches: list[dict] = [] - for lyr in sorted(la): - on_freqs = on_dict[lyr] # [T_on, 1, 1, D] - # 从 raw 表重建 OFF per-token:每段 [:seg_len] - off_freqs = torch.cat( - [off_dict[lyr][:s, :, :, :] for s in seqlens], dim=0) # [T_off, 1, 1, D] - - try: - on_aligned, off_aligned = _align_packed( - on_freqs, off_freqs, align_mask) - except ValueError as e: - return CheckResult(name="rope_freqs", passed=False, - metrics={"error": f"align failed L{lyr}: {e}"}) - - diff = (on_aligned - off_aligned).abs() # [N, 1, 1, D] + for layer_idx in layers: + aligned_result = _align_rope_freqs_layer(on_dict[layer_idx], off_dict[layer_idx], align_mask) + if aligned_result is None: + continue + on_a, off_a = aligned_result + diff = (on_a - off_a).abs() # [N,1,1,D] md = float(diff.max()) max_diff = max(max_diff, md) - if md > 0: - token_diff = diff.squeeze(1).squeeze(1).max(dim=-1) # values [N], indices [N] + token_diff = diff.squeeze(1).squeeze(1).max(dim=-1) # values [N], indices [N] bad_mask = token_diff.values > 0 for t in bad_mask.nonzero(as_tuple=True)[0].tolist(): t = int(t) d = int(token_diff.indices[t]) mismatches.append({ - "layer": lyr, "token_idx": t, "dim": d, - "on_val": float(on_aligned[t, 0, 0, d]), - "off_val": float(off_aligned[t, 0, 0, d]), + "layer": layer_idx, "token_idx": t, "dim": d, + "on_val": float(on_a[t, 0, 0, d]), + "off_val": float(off_a[t, 0, 0, d]), "diff": float(token_diff.values[t]), }) - metrics: dict = {"max_diff": max_diff, "num_layers": len(la)} + metrics: dict = {"max_diff": max_diff, "num_layers": len(layers)} if mismatches: metrics["mismatches"] = mismatches[:20] metrics["total_mismatches"] = len(mismatches) - return CheckResult(name="rope_freqs", passed=max_diff == 0.0, metrics=metrics) + return CheckResult(name=result_name, passed=max_diff == 0.0, metrics=metrics) + + +def cmp_rope_freqs_token(dir_on: str, dir_off: str, pos: int, + layer: int | None = None, + align_mask: torch.Tensor | None = None) -> CheckResult | None: + """rope_freqs 在对齐后 suffix-packed 位置 [pos] 的角度向量对比(应精确相等)。 + + 取 ``layer``(默认最后一层)对齐后第 ``pos`` 个 token 的角度向量 [D],比 ON/OFF。 + 角度是 RoPE 输入,应逐元素相等 → max_abs 应为 0。 + """ + filepath_on = os.path.join(dir_on, "rope_freqs.pt") + filepath_off = os.path.join(dir_off, "rope_freqs.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): + return None + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) + if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): + return None + common = set(on_dict.keys()) & set(off_dict.keys()) + rf_layer = layer if layer is not None else (max(common) if common else 0) + result_name = f"rope_freqs_L{rf_layer}_pos{pos}" + if rf_layer not in on_dict or rf_layer not in off_dict: + return CheckResult(name=result_name, metrics={"error": f"layer {rf_layer} 缺失"}) + + if align_mask is None: + align_mask = _build_attn_align_mask(dir_on, dir_off) + if align_mask is None: + return CheckResult(name=result_name, metrics={"error": "cu_seqlens/prefix_lens 缺失"}) + + aligned_result = _align_rope_freqs_layer(on_dict[rf_layer], off_dict[rf_layer], align_mask) + if aligned_result is None: + return CheckResult(name=result_name, metrics={"error": "对齐失败"}) + on_a, off_a = aligned_result + n = on_a.shape[0] + if pos < 0 or pos >= n: + return CheckResult(name=result_name, + metrics={"error": f"pos {pos} 越界: 对齐后 token 数={n}"}) + on_vec = on_a[pos].reshape(-1) + off_vec = off_a[pos].reshape(-1) + m = _vec_metrics(on_vec, off_vec) + return CheckResult(name=result_name, passed=m["max_abs"] == 0.0, metrics=m) + + +# ════════════════════════════════════════════════════════════════ +# Attention KV: ON expanded_kv vs OFF full_kv(prefix 复用校验) +# ════════════════════════════════════════════════════════════════ + +def _load_attn_kv(dir_path: str, layer: int, + fname: str) -> tuple[torch.Tensor | None, torch.Tensor | None]: + """load {key, value} for a layer from fname. Returns (key, value) or (None, None).""" + filepath = os.path.join(dir_path, fname) + if not os.path.exists(filepath): + return None, None + kv_dict = torch.load(filepath, weights_only=True) + if not isinstance(kv_dict, dict): + return None, None + entry = kv_dict.get(layer) + if entry is None: + return None, None + return entry.get("key"), entry.get("value") + + +def cmp_attn_kv(dir_on: str, dir_off: str, + layer: int | None = None) -> CheckResult | None: + """对比 ON expanded_kv vs OFF full_kv(K/V 分别),逐元素 max_diff + cos。 + + 两者都应是 full(prefix+suffix)且**逐元素相同**(prefix-sharing 的 KV 展开应精确还原 + 完整 KV)。相同 → attention 输入一致,attention_output 差异必来自 attention 计算/mask; + 不同 → bug 在 build_kv 的 prefix 复用(store/expand)。 + """ + filepath_on = os.path.join(dir_on, "expanded_kv.pt") + filepath_off = os.path.join(dir_off, "full_kv.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): + return None + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) + if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): + return None + layers = sorted(set(on_dict.keys()) & set(off_dict.keys())) + if layer is not None: + layers = [l for l in layers if l == layer] + result_name = f"attn_kv_L{layer}" if layer is not None else "attn_kv" + if not layers: + return CheckResult(name=result_name, passed=False, + metrics={"error": f"layer {layer} 不在双方 attn_kv 中"}) + + per_layer: dict = {} + worst = {"max_diff": 0.0, "cos_min": 1.0} + for layer_idx in layers: + on_key = on_dict[layer_idx].get("key") + on_value = on_dict[layer_idx].get("value") + off_key = off_dict[layer_idx].get("key") + off_value = off_dict[layer_idx].get("value") + entry_result: dict = {} + for kv_type, (on_kv, off_kv) in [("K", (on_key, off_key)), ("V", (on_value, off_value))]: + if on_kv is None or off_kv is None: + entry_result[kv_type] = {"error": "缺失"} + continue + if on_kv.shape != off_kv.shape: + entry_result[kv_type] = { + "error": f"shape mismatch ON{tuple(on_kv.shape)} vs OFF{tuple(off_kv.shape)}"} + continue + on_flat = on_kv.reshape(on_kv.shape[0], -1).float() + off_flat = off_kv.reshape(off_kv.shape[0], -1).float() + element_diff = (on_flat - off_flat).abs() + token_cos = _cosine_sim(on_flat, off_flat, dim=-1) + max_elem_diff = float(element_diff.max()) + entry_result[kv_type] = { + "max_diff": max_elem_diff, + "cos_avg": float(token_cos.mean()), "cos_min": float(token_cos.min()), + "n_tokens": on_flat.shape[0]} + worst["max_diff"] = max(worst["max_diff"], max_elem_diff) + worst["cos_min"] = min(worst["cos_min"], float(token_cos.min())) + per_layer[layer_idx] = entry_result + # expanded 应精确等于 full → 阈值极严 + passed = worst["max_diff"] < 1e-5 and worst["cos_min"] > 0.9999 + return CheckResult(name=result_name, passed=passed, + metrics={"layers": per_layer, "max_diff": worst["max_diff"], + "cos_min": worst["cos_min"], "num_layers": len(layers)}) + + +def _print_attn_kv(r: CheckResult): + print(_SEP_SINGLE + f"\n [{r.name}] ON expanded_kv vs OFF full_kv(K/V 逐元素)") + print(_SEP_SINGLE) + m = r.metrics + if "error" in m: + print(f" {_CROSS} {m['error']}\n"); return + layers = m.get("layers", {}) + print(f" {'LAYER':>6s} {'K_MAXDIFF':>12s} {'K_COS':>10s} " + f"{'V_MAXDIFF':>12s} {'V_COS':>10s} {'STATUS':>8s}") + print(f" {'─' * 6} {'─' * 12} {'─' * 10} {'─' * 12} {'─' * 10} {'─' * 8}") + bad = [] + for layer_idx in sorted(layers): + d = layers[layer_idx] + kd, vd = d.get("K", {}), d.get("V", {}) + if "error" in kd or "error" in vd: + print(f" {layer_idx:>6d} K:{kd.get('error','')} V:{vd.get('error','')}") + bad.append(layer_idx); continue + kmd, kcos = kd["max_diff"], kd["cos_avg"] + vmd, vcos = vd["max_diff"], vd["cos_avg"] + ok = kmd < 1e-5 and vmd < 1e-5 + if not ok: + bad.append(layer_idx) + print(f" {layer_idx:>6d} {kmd:>12.3e} {kcos:>10.6f} " + f"{vmd:>12.3e} {vcos:>10.6f} {'OK' if ok else 'DIFF':>8s}") + print(f"\n max_diff={m.get('max_diff')} cos_min={m.get('cos_min')} " + f"{_CHECK if r.passed else _CROSS} " + f"{'PASS(KV 一致)' if r.passed else 'FAIL(KV 不一致 → build_kv prefix 复用)'}") + if bad: + print(f" ⚠ 首个 KV 不一致层: {bad[0]}") + print() + + +def cmp_build_kv_input_v(dir_on: str, dir_off: str, + layer: int | None = None) -> CheckResult | None: + """对比 ON build_kv_input_v vs OFF build_kv_input_v(suffix 对齐)。 + + 两边都存 get_qkv 后、build_kv/RoPE 前的 raw V(``{layer: tensor}``)。同源对比, + 应逐元素相同——若不同则问题在 QKV 投影阶段(hidden_states / QKV 权重)。 + ON_T vs OFF_T 还能看出 ON 有没有把 hidden_states 裁成 suffix-only。 + """ + filepath_on = os.path.join(dir_on, "build_kv_input_v.pt") + filepath_off = os.path.join(dir_off, "build_kv_input_v.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): + return None + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) + if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): + return None + layers = sorted(set(on_dict.keys()) & set(off_dict.keys())) + if layer is not None: + layers = [l for l in layers if l == layer] + result_name = f"build_kv_input_v_L{layer}" if layer is not None else "build_kv_input_v" + if not layers: + return CheckResult(name=result_name, passed=False, metrics={"error": "no layers"}) + + align_mask = _build_attn_align_mask(dir_on, dir_off) + per_layer: dict = {} + worst_md = 0.0 + worst_cos = 1.0 + for layer_idx in layers: + on_v = on_dict[layer_idx] + off_v = off_dict[layer_idx] + if on_v is None or off_v is None: + per_layer[layer_idx] = {"error": "缺失"}; continue + on_f = on_v.reshape(on_v.shape[0], -1).float() + off_f = off_v.reshape(off_v.shape[0], -1).float() + on_T, off_T = int(on_v.shape[0]), int(off_v.shape[0]) + if align_mask is not None and on_f.shape[0] != off_f.shape[0]: + try: + on_f, off_f = _align_packed(on_f, off_f, align_mask) + except ValueError as e: + per_layer[layer_idx] = {"error": str(e), "on_T": on_T, "off_T": off_T} + continue + diff = (on_f - off_f).abs() + cos = _cosine_sim(on_f, off_f, dim=-1) + md = float(diff.max()) + per_layer[layer_idx] = {"max_diff": md, "cos_avg": float(cos.mean()), + "cos_min": float(cos.min()), "n_tokens": on_f.shape[0], + "on_T": on_T, "off_T": off_T} + worst_md = max(worst_md, md) + worst_cos = min(worst_cos, float(cos.min())) + passed = worst_md < 1e-5 + return CheckResult(name=result_name, passed=passed, + metrics={"layers": per_layer, "max_diff": worst_md, "cos_min": worst_cos}) + + +def _print_build_kv_input_v(r: CheckResult): + print(_SEP_SINGLE + f"\n [{r.name}] ON build_kv_input_v vs OFF build_kv_input_v(suffix 对齐)") + print(_SEP_SINGLE) + m = r.metrics + if "error" in m: + print(f" {_CROSS} {m['error']}\n"); return + layers = m.get("layers", {}) + print(f" {'LAYER':>6s} {'MAXDIFF':>12s} {'COS':>10s} " + f"{'ON_T':>8s} {'OFF_T':>8s} {'STATUS':>8s}") + print(f" {'─' * 6} {'─' * 12} {'─' * 10} {'─' * 8} {'─' * 8} {'─' * 8}") + for layer_idx in sorted(layers): + d = layers[layer_idx] + if "max_diff" not in d: + print(f" {layer_idx:>6d} {d.get('error', '')} ON_T={d.get('on_T')} OFF_T={d.get('off_T')}") + continue + md, cos = d["max_diff"], d["cos_avg"] + ok = md < 1e-5 + _crop = " (cropped)" if d.get("on_T") != d.get("off_T") else "" + print(f" {layer_idx:>6d} {md:>12.3e} {cos:>10.6f} " + f"{d.get('on_T', '—'):>8} {d.get('off_T', '—'):>8} " + f"{'OK' if ok else 'DIFF':>8s}{_crop}") + print(f"\n max_diff={m.get('max_diff')} cos_min={m.get('cos_min')} " + f"{_CHECK if r.passed else _CROSS} " + f"{'PASS(build_kv 前 V 一致 → 偏由 build_kv 引入)' if r.passed else 'FAIL(build_kv 前 V 已偏 → 根因在 get_qkv/hidden_states)'}") + print() + + +def cmp_hidden_states(dir_on: str, dir_off: str, + layer: int | None = None) -> CheckResult | None: + """对比 ON vs OFF hidden_states(suffix 对齐,注意力层入口)。 + + 这是 QKV 投影的 INPUT。如果 hidden_states 一致但 V 不一致 → GEMM 精度差异; + 如果 hidden_states 就不一致 → 根因在上游(embedding / input_layernorm)。 + """ + filepath_on = os.path.join(dir_on, "hidden_states.pt") + filepath_off = os.path.join(dir_off, "hidden_states.pt") + if not os.path.exists(filepath_on) or not os.path.exists(filepath_off): + return None + on_dict = torch.load(filepath_on, weights_only=True) + off_dict = torch.load(filepath_off, weights_only=True) + if not isinstance(on_dict, dict) or not isinstance(off_dict, dict): + return None + layers = sorted(set(on_dict.keys()) & set(off_dict.keys())) + if layer is not None: + layers = [l for l in layers if l == layer] + result_name = f"hidden_states_L{layer}" if layer is not None else "hidden_states" + if not layers: + return CheckResult(name=result_name, passed=False, metrics={"error": "no layers"}) + + align_mask = _build_attn_align_mask(dir_on, dir_off) + per_layer: dict = {} + worst_md = 0.0 + worst_cos = 1.0 + for layer_idx in layers: + on_hs = on_dict[layer_idx] + off_hs = off_dict[layer_idx] + if on_hs is None or off_hs is None: + per_layer[layer_idx] = {"error": "缺失"}; continue + on_f = on_hs.reshape(on_hs.shape[0], -1).float() + off_f = off_hs.reshape(off_hs.shape[0], -1).float() + on_T, off_T = int(on_hs.shape[0]), int(off_hs.shape[0]) + if align_mask is not None and on_f.shape[0] != off_f.shape[0]: + try: + on_f, off_f = _align_packed(on_f, off_f, align_mask) + except ValueError as e: + per_layer[layer_idx] = {"error": str(e), "on_T": on_T, "off_T": off_T} + continue + diff = (on_f - off_f).abs() + cos = _cosine_sim(on_f, off_f, dim=-1) + md = float(diff.max()) + per_layer[layer_idx] = {"max_diff": md, "cos_avg": float(cos.mean()), + "cos_min": float(cos.min()), "n_tokens": on_f.shape[0], + "on_T": on_T, "off_T": off_T} + worst_md = max(worst_md, md) + worst_cos = min(worst_cos, float(cos.min())) + passed = worst_md < 1e-5 + return CheckResult(name=result_name, passed=passed, + metrics={"layers": per_layer, "max_diff": worst_md, "cos_min": worst_cos}) + + +def _print_hidden_states(r: CheckResult): + print(_SEP_SINGLE + f"\n [{r.name}] ON vs OFF hidden_states(suffix 对齐,注意力入口)") + print(_SEP_SINGLE) + m = r.metrics + if "error" in m: + print(f" {_CROSS} {m['error']}\n"); return + layers = m.get("layers", {}) + print(f" {'LAYER':>6s} {'MAXDIFF':>12s} {'COS':>10s} " + f"{'ON_T':>8s} {'OFF_T':>8s} {'STATUS':>8s}") + print(f" {'─' * 6} {'─' * 12} {'─' * 10} {'─' * 8} {'─' * 8} {'─' * 8}") + for layer_idx in sorted(layers): + d = layers[layer_idx] + if "max_diff" not in d: + print(f" {layer_idx:>6d} {d.get('error', '')} ON_T={d.get('on_T')} OFF_T={d.get('off_T')}") + continue + md, cos = d["max_diff"], d["cos_avg"] + ok = md < 1e-5 + print(f" {layer_idx:>6d} {md:>12.3e} {cos:>10.6f} " + f"{d.get('on_T', '—'):>8} {d.get('off_T', '—'):>8} " + f"{'OK' if ok else 'DIFF':>8s}") + print(f"\n max_diff={m.get('max_diff')} cos_min={m.get('cos_min')} " + f"{_CHECK if r.passed else _CROSS} " + f"{'PASS(hidden_states 一致 → V 差异在 GEMM)' if r.passed else 'FAIL(hidden_states 不一致 → 根因上游)'}") + print() # ════════════════════════════════════════════════════════════════ @@ -509,10 +1088,10 @@ def _load_mask_2d(dir_path: str, mask_kind: str, tag: str) -> torch.Tensor | Non if mask_kind == "none": return None fname = f"{mask_kind}_mask_{tag}.pt" # label_mask_{tag} / attention_mask_{tag} - fp = os.path.join(dir_path, fname) - if not os.path.exists(fp): + filepath = os.path.join(dir_path, fname) + if not os.path.exists(filepath): return None - return torch.load(fp, weights_only=True).to(torch.bool) + return torch.load(filepath, weights_only=True).to(torch.bool) def _resolve_mask(dir_off: str, mask_kind: str, tag: str, @@ -578,21 +1157,30 @@ def cmp_2d(dir_on: str, dir_off: str, filename: str, name: str, # ════════════════════════════════════════════════════════════════ def _shape_of(dir_path: str, filename: str) -> str: - fp = os.path.join(dir_path, filename) - if not os.path.exists(fp): + filepath = os.path.join(dir_path, filename) + if not os.path.exists(filepath): return "(missing)" try: - obj = torch.load(fp, weights_only=True) + obj = torch.load(filepath, weights_only=True) if isinstance(obj, dict): # per-layer dict(attn_outputs / rope_freqs_*):显示层数 + 首层 shape sample = next(iter(obj.values())) if obj else None - sample_shape = f",{tuple(sample.shape)}" if sample is not None else "" + # rope_postqk.pt:每层值是 {"query","key"[,"positions"]} dict,取 query 的 shape 代表 + if isinstance(sample, dict): + query_tensor = sample.get("query") + sample_shape = f",Q{tuple(query_tensor.shape)}" if query_tensor is not None else "" + elif sample is not None: + sample_shape = f",{tuple(sample.shape)}" + else: + sample_shape = "" return f"(dict,{len(obj)}L{sample_shape})" return str(tuple(obj.shape)) except Exception: return "(error)" + + def _logits_shape(dir_path: str, manifest: dict | None) -> str: """Shape string for logits, manifest-aware: tp>1 → show shard shape tagged. @@ -640,6 +1228,13 @@ def _print_shapes(dir_on: str, dir_off: str, tag: str, f"attention_mask_{tag}.pt", "logits.pt", "attn_outputs.pt", + "rope_postqk.pt", + "rope_preqk.pt", + "rope_freqs.pt", + "expanded_kv.pt", + "full_kv.pt", + "build_kv_input_v.pt", + "hidden_states.pt", "prefix_lens.pt", "cu_seqlens_q.pt", ] @@ -647,7 +1242,6 @@ def _print_shapes(dir_on: str, dir_off: str, tag: str, print(f" {'─' * 28} {'─' * 16} {'─' * 16} {'─' * 10}") for fname in files: if fname == "logits.pt": - # TP-sharded: per-rank logits_tp{r}.pt, not a plain logits.pt s_on = _logits_shape(dir_on, manifest_on) s_off = _logits_shape(dir_off, manifest_off) else: @@ -713,17 +1307,17 @@ def _print_per_layer(r: CheckResult): f"{'TOKENS':>8s} {'STATUS':>8s}") print(f" {'─' * 6} {'─' * 14} {'─' * 14} {'─' * 8} {'─' * 8}") bad = [] - for lyr in sorted(layers.keys()): - d = layers[lyr] + for layer_idx in sorted(layers.keys()): + d = layers[layer_idx] if "error" in d: - print(f" {lyr:>6d} {d['error']}") - bad.append(lyr) + print(f" {layer_idx:>6d} {d['error']}") + bad.append(layer_idx) continue ok = d["cos_avg"] > _COS_AVG_PASS and d["cos_min"] > _COS_MIN_PASS - print(f" {lyr:>6d} {d['cos_avg']:>14.6e} {d['cos_min']:>14.6e} " + print(f" {layer_idx:>6d} {d['cos_avg']:>14.6e} {d['cos_min']:>14.6e} " f"{d['n_tokens']:>8d} {'PASS' if ok else 'WARN':>8s}") if not ok: - bad.append(lyr) + bad.append(layer_idx) if bad: print(f"\n ⚠ First deviating layer: {bad[0]}") elif "cos_avg" in r.metrics: @@ -736,10 +1330,56 @@ def _print_per_layer(r: CheckResult): print() -def _print_first_token(r: CheckResult): - print(_SEP_SINGLE + f"\n [first_token] {r.name} (packed position [0])") +def _print_rope_postqk_per_layer(r: CheckResult): + _sec = "rope_preqk" if "preqk" in r.name else "rope_postqk" + _stage = "Pre-RoPE" if "preqk" in r.name else "Post-RoPE" + print(_SEP_SINGLE + f"\n [{_sec}] {_stage} Q/K Per-Layer Cosine Similarity") + print(_SEP_SINGLE) + layers = r.metrics.get("layers") + if isinstance(layers, dict): + print(f" {'LAYER':>6s} {'Q_MAXDIFF':>12s} {'Q_COS_AVG':>12s} " + f"{'K_MAXDIFF':>12s} {'K_COS_AVG':>12s} " + f"{'TOKENS':>8s} {'STATUS':>8s}") + print(f" {'─' * 6} {'─' * 12} {'─' * 12} {'─' * 12} {'─' * 12} " + f"{'─' * 8} {'─' * 8}") + bad = [] + for layer_idx in sorted(layers.keys()): + d = layers[layer_idx] + if "error" in d: + print(f" {layer_idx:>6d} {d['error']}") + bad.append(layer_idx) + continue + ok = (d["Q_cos_avg"] > _COS_AVG_PASS and d["Q_cos_min"] > _COS_MIN_PASS + and d["K_cos_avg"] > _COS_AVG_PASS and d["K_cos_min"] > _COS_MIN_PASS) + print(f" {layer_idx:>6d} {d.get('Q_max_diff', 0.0):>12.3e} " + f"{d['Q_cos_avg']:>12.6e} " + f"{d.get('K_max_diff', 0.0):>12.3e} {d['K_cos_avg']:>12.6e} " + f"{d['n_tokens']:>8d} {'PASS' if ok else 'WARN':>8s}") + if not ok: + bad.append(layer_idx) + if bad: + print(f"\n ⚠ First deviating layer: {bad[0]}") + print(" (Q/K max_diff 与 build_kv_input_v 的 V max_diff 同口径,可直接对比)") + elif "Q_cos_avg" in r.metrics: + d = r.metrics + ok = (d["Q_cos_avg"] > _COS_AVG_PASS and d["Q_cos_min"] > _COS_MIN_PASS + and d["K_cos_avg"] > _COS_AVG_PASS and d["K_cos_min"] > _COS_MIN_PASS) + print(f" L{d['layer']} Q_maxdiff={d.get('Q_max_diff', 0.0):.3e} " + f"Q_cos_avg={d['Q_cos_avg']:.6e} " + f"K_maxdiff={d.get('K_max_diff', 0.0):.3e} " + f"K_cos_avg={d['K_cos_avg']:.6e} {'PASS' if ok else 'WARN'}") + elif "error" in r.metrics: + print(f" {_CROSS} {r.metrics['error']}") + print() + + +def _print_packed_token(r: CheckResult): + print(_SEP_SINGLE + f"\n [packed_token] {r.name}") print(_SEP_SINGLE) m = r.metrics + if "error" in m: + print(f" {_CROSS} {m['error']}\n") + return for k in ["mean_abs", "max_abs", "rel_max", "rel_mean", "cos", "pearson"]: v = m.get(k) if v is not None: @@ -778,23 +1418,58 @@ def _print_2d_result(r: CheckResult): def _print_topk_vec(on_vec: torch.Tensor, off_vec: torch.Tensor, - topk: int, sort_by: str, label: str): - """1D 向量 top-K(first_token per-dim)。""" + topk: int, sort_by: str, label: str, show_rel: bool = True): + """1D 向量 top-K(packed_token per-dim)。 + + show_rel=False 时省略 REL_ERR 列(用于 logits 表,只看 val/abs)。 + """ abs_err = (on_vec - off_vec).abs() rel_err = abs_err / torch.maximum(on_vec.abs(), off_vec.abs()).clamp(min=1e-8) if sort_by == "abs": sort_key = abs_err elif sort_by == "rel": sort_key = rel_err - else: # "val" - sort_key = torch.maximum(on_vec.abs(), off_vec.abs()) + else: # "val" —— 带符号的实际值,不是绝对值 + # 对 logits:绝对值大但符号为负的 logit,softmax 后概率极低、不会被选中。 + # 按 abs 排会把这种"必不选"的 token 顶到表头,掩盖真正的高 logit 候选。 + # 改用 max(on, off) 带符号值,让真正的高 logit(候选 token)排前面。 + sort_key = torch.maximum(on_vec, off_vec) _, idx = sort_key.topk(min(topk, sort_key.numel())) idx = idx.to(torch.long) print(f"\n [{label}] top-{topk} dims (sort by {sort_by})") - print(f" {'DIM':>6s} {'ON':>14s} {'OFF':>14s} {'ABS_ERR':>12s} {'REL_ERR':>12s}") - for i in idx.tolist(): - print(f" {i:>6d} {float(on_vec[i]):>14.6e} {float(off_vec[i]):>14.6e}" - f" {float(abs_err[i]):>12.6e} {float(rel_err[i]):>12.6e}") + if show_rel: + print(f" {'DIM':>6s} {'ON':>14s} {'OFF':>14s} {'ABS_ERR':>12s} {'REL_ERR':>12s}") + for i in idx.tolist(): + print(f" {i:>6d} {float(on_vec[i]):>14.6e} {float(off_vec[i]):>14.6e}" + f" {float(abs_err[i]):>12.6e} {float(rel_err[i]):>12.6e}") + else: + print(f" {'DIM':>6s} {'ON':>14s} {'OFF':>14s} {'ABS_ERR':>12s}") + for i in idx.tolist(): + print(f" {i:>6d} {float(on_vec[i]):>14.6e} {float(off_vec[i]):>14.6e}" + f" {float(abs_err[i]):>12.6e}") + return idx.tolist() + + +def _print_vec_at_dims(on_vec: torch.Tensor, off_vec: torch.Tensor, + dims, label: str, show_rel: bool = True): + """在指定 dims 上打印 ON/OFF/ABS_ERR(不排序),跨 stage 对齐同一批 dim。 + + 供 rope 流水线 top-K 对齐:dims 取自 rope_postqk 的 sort-err top-K, + 在 rope_preqk / rope_freqs 上显示同样的 dim,逐 dim 追溯误差来源。 + """ + abs_err = (on_vec - off_vec).abs() + rel_err = abs_err / torch.maximum(on_vec.abs(), off_vec.abs()).clamp(min=1e-8) + print(f"\n [{label}] at {len(dims)} dims") + if show_rel: + print(f" {'DIM':>6s} {'ON':>14s} {'OFF':>14s} {'ABS_ERR':>12s} {'REL_ERR':>12s}") + for i in dims: + print(f" {i:>6d} {float(on_vec[i]):>14.6e} {float(off_vec[i]):>14.6e}" + f" {float(abs_err[i]):>12.6e} {float(rel_err[i]):>12.6e}") + else: + print(f" {'DIM':>6s} {'ON':>14s} {'OFF':>14s} {'ABS_ERR':>12s}") + for i in dims: + print(f" {i:>6d} {float(on_vec[i]):>14.6e} {float(off_vec[i]):>14.6e}" + f" {float(abs_err[i]):>12.6e}") def _print_topk_2d(on_t: torch.Tensor, off_t: torch.Tensor, @@ -875,7 +1550,11 @@ def main(): ap.add_argument("--mask", choices=["label", "attention", "none"], default="label", help="2D mask type (default: label)") ap.add_argument("--layer", type=int, default=None, - help="Compare specific attn layer 1-indexed (default: all)") + help="Compare specific attn layer 1-indexed (default: all). " + "Also used by packed_token attn (default: last layer).") + ap.add_argument("--token", type=int, default=0, + help="Packed token position for packed_token compare " + "(single int index, default: 0)") ap.add_argument("--atol", type=float, default=1e-5, help="Absolute tolerance for 2D (default: 1e-5)") ap.add_argument("--topk", type=int, default=0, @@ -888,13 +1567,14 @@ def main(): _print_header(args.dir_on, args.dir_off, args.dir_off2, args.tag, args.mask, args.layer) - # ── parallel topology (manifest-driven: TP shards, future SP/PP) ── + # ── topology diagnostics ── manifest_on = _load_manifest(args.dir_on) manifest_off = _load_manifest(args.dir_off) _print_topology(manifest_on, manifest_off) # ── shape diagnostics ── - _print_shapes(args.dir_on, args.dir_off, args.tag, manifest_on, manifest_off) + _print_shapes(args.dir_on, args.dir_off, args.tag, + manifest_on=manifest_on, manifest_off=manifest_off) # ── resolve 2D mask ── ref = _load_tensor(args.dir_off, f"logprobs_{args.tag}.pt") @@ -907,39 +1587,75 @@ def main(): all_results: list[CheckResult] = [] - # ── packed: rope 角度(suffix 对齐,应精确相等 max_diff==0) ── - r = cmp_rope_freqs(args.dir_on, args.dir_off) + # ── RoPE pipeline per-layer:pre Q/K → rope 角度 → post Q/K ── + # 按计算顺序串联:旋转前 Q/K → 每token旋转角度(freqs) → 旋转后 Q/K, + # 定位分歧出现在 RoPE 哪一步(pre 就偏=上游;freqs 偏=角度表;post 才偏=旋转应用)。 + r = cmp_rope_postqk_layer(args.dir_on, args.dir_off, args.layer, stage="pre") + if r: + all_results.append(r) + _print_rope_postqk_per_layer(r) + + r = cmp_rope_freqs(args.dir_on, args.dir_off, layer=args.layer) if r: all_results.append(r) _print_rope_freqs(r) - # ── packed: attention_output per-layer cos ── + r = cmp_rope_postqk_layer(args.dir_on, args.dir_off, args.layer, stage="post") + if r: + all_results.append(r) + _print_rope_postqk_per_layer(r) + + # ── packed: attention KV(ON expanded vs OFF full,prefix 复用校验)── + r = cmp_attn_kv(args.dir_on, args.dir_off, args.layer) + if r: + all_results.append(r) + _print_attn_kv(r) + + # ── packed: build_kv 输入 V — ON vs OFF 同源对比 ── + r = cmp_build_kv_input_v(args.dir_on, args.dir_off, args.layer) + if r: + all_results.append(r) + _print_build_kv_input_v(r) + + # ── packed: hidden_states(注意力入口)— 隔离 QKV 投影 vs 上游 ── + r = cmp_hidden_states(args.dir_on, args.dir_off, args.layer) + if r: + all_results.append(r) + _print_hidden_states(r) + + # ── packed: attention_output per-layer cos(RoPE 下游)── r = cmp_attn_layer(args.dir_on, args.dir_off, args.layer) if r: all_results.append(r) _print_per_layer(r) - # ── packed: first_token(attn[0] + logits[0]) ── - ft_results = cmp_first_token(args.dir_on, args.dir_off) - for r in ft_results: + # ── packed: packed_token(attn[pos] + logits[pos],suffix 对齐后) ── + # pos 由 --token 指定(默认 0,索引对齐后的 suffix-packed 空间); + # attn 用 --layer 指定的层(默认最后一层);logits 永远最后一层。 + pos = args.token + align_mask = _build_attn_align_mask(args.dir_on, args.dir_off) + pt_results = cmp_packed_token(args.dir_on, args.dir_off, pos, args.layer, + align_mask=align_mask) + for r in pt_results: all_results.append(r) - _print_first_token(r) - # first_token top-K(per-dim) - if args.topk > 0 and ft_results: - last = _get_num_layers(args.dir_on) or _get_num_layers(args.dir_off) - if last: - a = _load_attn_output(args.dir_on, last) - b = _load_attn_output(args.dir_off, last) - if a is not None and b is not None: - a0 = (a.squeeze(1) if a.dim() == 3 else a)[0].cpu() - b0 = (b.squeeze(1) if b.dim() == 3 else b)[0].cpu() - _print_topk_vec(a0, b0, args.topk, "val", "first_token_attn") - lo = _load_logits(args.dir_on) - lf = _load_logits(args.dir_off) - if lo is not None and lf is not None: - lo_f, lf_f = _logits_first_token(lo, lf) - _print_topk_vec(lo_f.cpu(), lf_f.cpu(), args.topk, - args.sort_err, "first_token_logits") + _print_packed_token(r) + # packed_token top-K(per-dim,同样先对齐再取 [pos]) + if args.topk > 0 and pt_results: + attn_layer = args.layer if args.layer is not None else ( + _get_num_layers(args.dir_on) or _get_num_layers(args.dir_off)) + if attn_layer: + a = _load_attn_output(args.dir_on, attn_layer) + b = _load_attn_output(args.dir_off, attn_layer) + vecs = _aligned_vec_at_pos(a, b, True, pos, align_mask) + if vecs is not None: + _print_topk_vec(vecs[0].cpu(), vecs[1].cpu(), args.topk, "val", + f"attn_L{attn_layer}_pos{pos}") + logits_on = _load_logits(args.dir_on) + logits_off = _load_logits(args.dir_off) + vecs = _aligned_vec_at_pos(logits_on, logits_off, False, pos, align_mask) + if vecs is not None: + _print_topk_vec(vecs[0].cpu(), vecs[1].cpu(), args.topk, + "val", f"logits_pos{pos}", show_rel=False) # ── packed: logits(suffix 对齐) ── r = cmp_logits_packed(args.dir_on, args.dir_off) @@ -947,6 +1663,57 @@ def main(): all_results.append(r) _print_logits_packed(r) + # ── RoPE pipeline packed_token:pre Q/K → rope_freqs → post Q/K(指定 pos)── + rope_pt_results: list[CheckResult] = [] + for r in cmp_rope_postqk_token(args.dir_on, args.dir_off, pos, args.layer, + align_mask=align_mask, stage="pre"): + all_results.append(r); rope_pt_results.append(r); _print_packed_token(r) + _rf = cmp_rope_freqs_token(args.dir_on, args.dir_off, pos, args.layer, + align_mask=align_mask) + if _rf is not None: + all_results.append(_rf); rope_pt_results.append(_rf); _print_packed_token(_rf) + for r in cmp_rope_postqk_token(args.dir_on, args.dir_off, pos, args.layer, + align_mask=align_mask, stage="post"): + all_results.append(r); rope_pt_results.append(r); _print_packed_token(r) + # rope packed_token top-K —— dim 跨 stage 对齐:以 rope_postqk 的 sort-err top-K dim 为基准, + # rope_preqk 显示同样 dim,rope_freqs 显示 dim%D(角度按 head_dim 共享),逐 dim 追溯误差。 + if args.topk > 0 and rope_pt_results: + rope_layer = args.layer if args.layer is not None else ( + _get_num_layers(args.dir_on) or _get_num_layers(args.dir_off)) + if rope_layer: + pre_q_on, pre_k_on = _load_rope_postqk(args.dir_on, rope_layer, "rope_preqk.pt") + pre_q_off, pre_k_off = _load_rope_postqk(args.dir_off, rope_layer, "rope_preqk.pt") + post_q_on, post_k_on = _load_rope_postqk(args.dir_on, rope_layer, "rope_postqk.pt") + post_q_off, post_k_off = _load_rope_postqk(args.dir_off, rope_layer, "rope_postqk.pt") + pre_vecs = _rope_postqk_vec_at_pos(pre_q_on, pre_k_on, pre_q_off, pre_k_off, pos, align_mask) + post_vecs = _rope_postqk_vec_at_pos(post_q_on, post_k_on, post_q_off, post_k_off, pos, align_mask) + freq_on, freq_off = _load_rope_freqs_vec_at_pos( + args.dir_on, args.dir_off, rope_layer, pos, align_mask) + if post_vecs is not None: + pqo, pqf, pko, pkf = post_vecs + # Q: postqk sort-err top-K → preqk / freqs 同 dim + q_dims = _print_topk_vec(pqo.cpu(), pqf.cpu(), args.topk, args.sort_err, + f"rope_postqk_L{rope_layer}_Q_pos{pos}") + if pre_vecs is not None: + _print_vec_at_dims(pre_vecs[0].cpu(), pre_vecs[1].cpu(), q_dims, + f"rope_preqk_L{rope_layer}_Q_pos{pos} (same dims)") + if freq_on is not None and freq_off is not None: + _D = freq_on.numel() + _print_vec_at_dims(freq_on.cpu(), freq_off.cpu(), + [d % _D for d in q_dims], + f"rope_freqs_L{rope_layer}_Q_pos{pos} (dim%D)") + # K: 同样 + if pko is not None and pkf is not None: + k_dims = _print_topk_vec(pko.cpu(), pkf.cpu(), args.topk, args.sort_err, + f"rope_postqk_L{rope_layer}_K_pos{pos}") + if pre_vecs is not None and pre_vecs[2] is not None and pre_vecs[3] is not None: + _print_vec_at_dims(pre_vecs[2].cpu(), pre_vecs[3].cpu(), k_dims, + f"rope_preqk_L{rope_layer}_K_pos{pos} (same dims)") + if freq_on is not None and freq_off is not None: + _print_vec_at_dims(freq_on.cpu(), freq_off.cpu(), + [d % _D for d in k_dims], + f"rope_freqs_L{rope_layer}_K_pos{pos} (dim%D)") + # ── 2D: logprobs + entropy ── for fname, cname in [("logprobs", "logp"), ("entropy", "entropy")]: fn = f"{fname}_{args.tag}.pt" diff --git a/prefix-sharing/prefix_sharing/tools/diagnostic_dump.py b/prefix-sharing/prefix_sharing/tools/diagnostic_dump.py index aa0541cc..5f57e04a 100644 --- a/prefix-sharing/prefix_sharing/tools/diagnostic_dump.py +++ b/prefix-sharing/prefix_sharing/tools/diagnostic_dump.py @@ -10,7 +10,7 @@ 2. Positional encoding: - ``position_ids.pt`` — packed position ids [N] - - ``rope_emb.pt`` — per-layer RoPE encoding dict + - ``rope_postqk.pt`` — per-layer RoPE encoding dict {layer_idx: {"query": rotated_q, "key": rotated_k}} 3. 2D format: @@ -39,8 +39,7 @@ _META_SAVED: set[str] = set() # saved metadata keys (dedup per key) _ATTN_BUFFER: dict[int, torch.Tensor] | None = None # {layer_idx: tensor} _ROPE_BUFFER: dict[int, dict] | None = None # {layer_idx: {"query": q, "key": k}} -_ROPE_FREQS_ON_BUFFER: dict[int, torch.Tensor] | None = None # {layer_idx: q_freqs per-token} -_ROPE_FREQS_OFF_BUFFER: dict[int, torch.Tensor] | None = None # {layer_idx: q_pos_emb table} +_ROPE_FREQS_BUFFER: dict[int, torch.Tensor] | None = None # {layer_idx: per-token RoPE 角度} def _get_dump_dir() -> str | None: @@ -87,7 +86,15 @@ def _rank0_only() -> bool: # cmp_diag knows how to gather each tensor without guessing. _TENSOR_SCOPES: dict[str, str] = { - "logits": "tp_vocab", # packed logits [N, V//tp] — vocab-sharded under TP + "logits": "tp_vocab", # packed logits [N, V//tp] — vocab-sharded under TP + "attn_outputs": "pp_stage", # per-layer attn output dict — PP-sharded + "rope_postqk": "pp_stage", # per-layer post-RoPE Q/K dict — PP-sharded + "rope_preqk": "pp_stage", # per-layer pre-RoPE Q/K dict — PP-sharded + "rope_freqs": "pp_stage", # per-layer RoPE freqs dict — PP-sharded + "expanded_kv": "pp_stage", # per-layer expanded KV dict — PP-sharded (ON only) + "full_kv": "pp_stage", # per-layer full KV dict — PP-sharded (OFF only) + "build_kv_input_v": "pp_stage", # per-layer V dict — PP-sharded + "hidden_states": "pp_stage", # per-layer hidden_states dict — PP-sharded } _PARALLEL_INFO_CACHE: Any = None @@ -113,6 +120,63 @@ def _cached_parallel_info() -> Any: return _PARALLEL_INFO_CACHE +def _stage_last_layer(num_layers_global: int) -> int: + """Return the last global layer number owned by the current PP stage. + + Under PP, ``num_layers_global`` is split across stages. + Non-last stages never reach ``layer_number == num_layers_global``, + so the standard auto-flush condition silently drops their data. + This function computes the stage-local last layer so every stage + flushes independently. + + Uniform split (standard Megatron): + layers_per_stage = num_layers // pp_size + remainder stages get 1 extra layer. + """ + pi = _cached_parallel_info() + if pi is None or pi.pp_size <= 1: + return num_layers_global + base = num_layers_global // pi.pp_size + rem = num_layers_global % pi.pp_size + if pi.pp_rank < rem: + return (pi.pp_rank + 1) * (base + 1) + else: + return rem * (base + 1) + (pi.pp_rank - rem + 1) * base + + +def _pp_suffix() -> str: + """Return ``'_pp{r}'`` when ``pp_size > 1``, else ``''``.""" + pi = _cached_parallel_info() + if pi is not None and pi.pp_size > 1: + return f"_pp{pi.pp_rank}" + return "" + + +def _should_write_for_scope(scope: str) -> bool: + """Return True if this rank should dump data for the given scope. + + Gate logic: + - ``"global"`` → rank 0 only (all ranks have identical data) + - ``"tp_vocab"`` → every tp rank dumps (each has different shard) + - ``"pp_last"`` → tp_rank==0 on pp_last stage only + - ``"pp_stage"`` → tp_rank==0 within each PP stage + """ + pi = _cached_parallel_info() + if scope == "global": + return _rank0_only() + if scope == "tp_vocab": + return True + if scope == "pp_last": + if pi is not None and not pi.is_pipeline_last_stage: + return False + return pi is None or pi.tp_rank == 0 + if scope == "pp_stage": + if pi is None or pi.pp_size <= 1: + return _rank0_only() + return pi.tp_rank == 0 + return _rank0_only() + + def _with_suffix(name: str, suffix: str) -> str: """Insert a rank suffix before the extension: logits.pt → logits_tp0.pt.""" if not suffix: @@ -131,7 +195,13 @@ def _ensure_manifest(dump_dir: str, pi: Any) -> None: if dump_dir in _MANIFEST_WRITTEN: return _MANIFEST_WRITTEN.add(dump_dir) - if not _rank0_only(): + # Under PP, each stage is a separate process; allow tp_rank==0 on any + # PP stage to write the manifest (content is identical across stages). + # This prevents data loss if pp0's dump path is never reached. + if pi is not None and pi.pp_size > 1: + if pi.tp_rank != 0: + return + elif not _rank0_only(): return import json manifest = { @@ -159,21 +229,28 @@ def _save_tensor(name: str, tensor: torch.Tensor, dump_dir: str, See the ``_TENSOR_SCOPES`` block above for the scope semantics. ``scope`` defaults to ``"global"`` (rank-0-only, plain filename) so existing callers - are unchanged; per-rank tensors opt in via ``scope="tp_vocab"`` (etc.). + are unchanged; per-rank tensors opt in via ``scope="tp_vocab"`` / ``"pp_last"`` + (etc.). + + Gate logic (delegated to ``_should_write_for_scope``): + - ``"global"`` → rank 0 only + - ``"tp_vocab"`` → every tp rank dumps + - ``"pp_last"`` → tp_rank==0 on pp_last only """ pi = _cached_parallel_info() _ensure_manifest(dump_dir, pi) + if not _should_write_for_scope(scope): + return False + + # Determine filename suffix from scope if scope == "tp_vocab" and pi is not None and pi.tp_size > 1: # Every tp rank dumps its own vocab shard; no rank-0 gate, no comm. fname = _with_suffix(name, f"_tp{pi.tp_rank}") else: - # global scope, OR tp_vocab with tp_size==1 (single shard == full): - # rank-0 dumps once under the plain name (single-card compatible). + # global / pp_last / tp_vocab with tp_size==1: plain filename. if scope == "tp_vocab": scope = "global" # tp==1 → behaves as global for logging - if not _rank0_only(): - return False fname = name try: path = os.path.join(dump_dir, fname) @@ -232,25 +309,31 @@ def _add_to_attn_buffer(layer_number: int, tensor: torch.Tensor) -> None: def _flush_attn_buffer(dump_dir: str) -> None: - """Write accumulated attn_outputs dict to disk and clear buffer.""" + """Write accumulated attn_outputs dict to disk and clear buffer. + + Under PP, each stage's tp_rank==0 writes with ``_pp{r}`` suffix; + assembly later merges all stage files. + """ global _ATTN_BUFFER if _ATTN_BUFFER is None: return - if not _rank0_only(): + if not _should_write_for_scope("pp_stage"): _ATTN_BUFFER = None return try: # Move every tensor to CPU (dict has no .detach()/.cpu()/.clone()) moved = {k: v.detach().cpu().clone() for k, v in _ATTN_BUFFER.items()} - torch.save(moved, os.path.join(dump_dir, "attn_outputs.pt")) - _log.warning("attn_outputs.pt saved (%d layers)", len(moved)) + fname = f"attn_outputs{_pp_suffix()}.pt" + torch.save(moved, os.path.join(dump_dir, fname)) + _log.warning("%s saved (%d layers)", fname, len(moved)) _ATTN_BUFFER = None except Exception as e: _log.warning("attn_outputs.pt save failed: %s", e) def _add_to_rope_buffer(layer_number: int, rotated_query: torch.Tensor, - rotated_key: torch.Tensor) -> None: + rotated_key: torch.Tensor, + positions: torch.Tensor | None = None) -> None: """Accumulate one layer's RoPE encoding into the global buffer.""" global _ROPE_BUFFER if _ROPE_BUFFER is None: @@ -258,79 +341,60 @@ def _add_to_rope_buffer(layer_number: int, rotated_query: torch.Tensor, _ROPE_BUFFER[layer_number] = { "query": rotated_query.detach().cpu().clone(), "key": rotated_key.detach().cpu().clone(), + "positions": positions.detach().cpu().clone() if positions is not None else None, } def _flush_rope_buffer(dump_dir: str) -> None: - """Write accumulated rope_emb dict to disk and clear buffer.""" + """Write accumulated rope_postqk dict to disk and clear buffer. + + Under PP, each stage's tp_rank==0 writes with ``_pp{r}`` suffix; + assembly later merges all stage files. + """ global _ROPE_BUFFER if _ROPE_BUFFER is None: return - if not _rank0_only(): + if not _should_write_for_scope("pp_stage"): _ROPE_BUFFER = None return try: - torch.save(_ROPE_BUFFER, os.path.join(dump_dir, "rope_emb.pt")) - _log.warning("rope_emb.pt saved (%d layers)", len(_ROPE_BUFFER)) + fname = f"rope_postqk{_pp_suffix()}.pt" + torch.save(_ROPE_BUFFER, os.path.join(dump_dir, fname)) + _log.warning("%s saved (%d layers)", fname, len(_ROPE_BUFFER)) _ROPE_BUFFER = None except Exception as e: - _log.warning("rope_emb.pt save failed: %s", e) + _log.warning("rope_postqk.pt save failed: %s", e) # ── RoPE angle dump (pre-apply, per-layer) ────────────────────── -def dump_rope_freqs_on(q_freqs: torch.Tensor, layer_number: int, - num_layers: int) -> None: - """Accumulate ON-mode per-token RoPE angles. Auto-flush on last layer. - - ``q_freqs`` is the result of ``q_pos_emb.index_select(0, packed_position_ids)`` - — shape [T_on, 1, 1, D], each token's actual rotation angles **before - cos/sin**. Stored as ``rope_freqs_on.pt`` (dict {layer_idx: tensor}). - """ - global _ROPE_FREQS_ON_BUFFER - dump_dir = _get_dump_dir() - if dump_dir is None: - return - if _ROPE_FREQS_ON_BUFFER is None: - _ROPE_FREQS_ON_BUFFER = {} - _ROPE_FREQS_ON_BUFFER[layer_number] = q_freqs.detach().cpu().clone() - if layer_number == num_layers: - if _rank0_only(): - try: - torch.save(_ROPE_FREQS_ON_BUFFER, - os.path.join(dump_dir, "rope_freqs_on.pt")) - _log.warning("rope_freqs_on.pt saved (%d layers)", - len(_ROPE_FREQS_ON_BUFFER)) - except Exception as e: - _log.warning("rope_freqs_on.pt save failed: %s", e) - _ROPE_FREQS_ON_BUFFER = None - - -def dump_rope_freqs_off(q_pos_emb: torch.Tensor, layer_number: int, - num_layers: int) -> None: - """Accumulate OFF-mode raw RoPE angle table. Auto-flush on last layer. +def dump_rope_freqs(q_freqs: torch.Tensor, layer_number: int, + num_layers: int) -> None: + """Accumulate per-token RoPE angles (ON/OFF 共用). Auto-flush on last layer. - ``q_pos_emb`` is the raw angle table (before per-token slicing) — - shape [L0, 1, 1, D], where ``freqs[p] = p * inv_freq``. - Stored as ``rope_freqs_off.pt`` (dict {layer_idx: tensor}). + ``q_freqs`` = 每 token 实际旋转角度(cos/sin 之前),shape [T, 1, 1, D]。 + ON 由 ``q_pos_emb.index_select(0, packed_position_ids)`` 得到;OFF 由 raw 表 + 按 cu_seqlens 切 per-token(每段 0..seg-1)得到。两边语义统一,写到各自 dump + 目录的 ``rope_freqs.pt``,cmp 侧 suffix 对齐后比 max_diff(角度应精确相等)。 """ - global _ROPE_FREQS_OFF_BUFFER + global _ROPE_FREQS_BUFFER dump_dir = _get_dump_dir() if dump_dir is None: return - if _ROPE_FREQS_OFF_BUFFER is None: - _ROPE_FREQS_OFF_BUFFER = {} - _ROPE_FREQS_OFF_BUFFER[layer_number] = q_pos_emb.detach().cpu().clone() - if layer_number == num_layers: - if _rank0_only(): + if _ROPE_FREQS_BUFFER is None: + _ROPE_FREQS_BUFFER = {} + _ROPE_FREQS_BUFFER[layer_number] = q_freqs.detach().cpu().clone() + if layer_number == _stage_last_layer(num_layers): + if _should_write_for_scope("pp_stage"): try: - torch.save(_ROPE_FREQS_OFF_BUFFER, - os.path.join(dump_dir, "rope_freqs_off.pt")) - _log.warning("rope_freqs_off.pt saved (%d layers)", - len(_ROPE_FREQS_OFF_BUFFER)) + fname = f"rope_freqs{_pp_suffix()}.pt" + torch.save(_ROPE_FREQS_BUFFER, + os.path.join(dump_dir, fname)) + _log.warning("%s saved (%d layers)", + fname, len(_ROPE_FREQS_BUFFER)) except Exception as e: - _log.warning("rope_freqs_off.pt save failed: %s", e) - _ROPE_FREQS_OFF_BUFFER = None + _log.warning("rope_freqs.pt save failed: %s", e) + _ROPE_FREQS_BUFFER = None # ── Position IDs dump ─────────────────────────────────────────── @@ -349,8 +413,9 @@ def dump_position_ids(position_ids: torch.Tensor) -> None: # ── RoPE encoding dump (per layer) ────────────────────────────── -def dump_rope_emb_layer(layer_number: int, rotated_query: torch.Tensor, - rotated_key: torch.Tensor, num_layers: int) -> None: +def dump_rope_postqk_layer(layer_number: int, rotated_query: torch.Tensor, + rotated_key: torch.Tensor, num_layers: int, + positions: torch.Tensor | None = None) -> None: """Accumulate one layer's post-RoPE query/key. Auto-flush on last layer. Call after RoPE is applied in each attention layer, for both ON and OFF modes. @@ -358,8 +423,8 @@ def dump_rope_emb_layer(layer_number: int, rotated_query: torch.Tensor, dump_dir = _get_dump_dir() if dump_dir is None: return - _add_to_rope_buffer(layer_number, rotated_query, rotated_key) - if layer_number == num_layers: + _add_to_rope_buffer(layer_number, rotated_query, rotated_key, positions) + if layer_number == _stage_last_layer(num_layers): _flush_rope_buffer(dump_dir) @@ -389,8 +454,8 @@ def dump_attn_on( _save_meta(packed_seq_params, list(prefix_sharing_plan.prefix_lens), dump_dir, meta_key="attn") - # Flush on last layer - if layer_number == num_layers: + # Flush on last layer of this PP stage + if layer_number == _stage_last_layer(num_layers): _flush_attn_buffer(dump_dir) @@ -419,8 +484,8 @@ def dump_attn_off( if layer_number == 1: _save_meta(packed_seq_params, [0] * batch_size, dump_dir, meta_key="attn") - # Flush on last layer - if layer_number == num_layers: + # Flush on last layer of this PP stage + if layer_number == _stage_last_layer(num_layers): _flush_attn_buffer(dump_dir) diff --git a/prefix-sharing/prefix_sharing/tools/diagnostic_dump_verl080.py b/prefix-sharing/prefix_sharing/tools/diagnostic_dump_verl080.py index 19838145..2fdf0243 100644 --- a/prefix-sharing/prefix_sharing/tools/diagnostic_dump_verl080.py +++ b/prefix-sharing/prefix_sharing/tools/diagnostic_dump_verl080.py @@ -22,6 +22,7 @@ from __future__ import annotations +import contextlib from typing import Any import torch @@ -30,6 +31,8 @@ from prefix_sharing.tools.diagnostic_dump import ( _get_dump_dir, _save_tensor, + _stage_last_layer, + _pp_suffix, ) @@ -158,12 +161,16 @@ def dump_logits_verl080(logits: torch.Tensor) -> None: _save_tensor("logits.pt", logits, dump_dir, scope="tp_vocab") -def dump_logprobs_2d_verl080(logp_2d: torch.Tensor, tag: str) -> None: - """存 2D log_probs ``[B, L_max]``,文件名 ``logprobs_{tag}.pt``。""" +def dump_logprobs_2d_verl080(logp_2d: torch.Tensor, tag: str, + scope: str = "global") -> None: + """存 2D log_probs ``[B, L_max]``,文件名 ``logprobs_{tag}.pt``。 + + ``scope="pp_last"`` 时仅最后一个 PP stage 落盘(多卡下 logprobs 只在末 stage 产生)。 + """ dump_dir = _get_dump_dir() if dump_dir is None: return - _save_tensor(f"logprobs_{tag}.pt", logp_2d, dump_dir) + _save_tensor(f"logprobs_{tag}.pt", logp_2d, dump_dir, scope=scope) def dump_attention_mask_verl080(mask_2d: torch.Tensor, tag: str) -> None: @@ -197,9 +204,248 @@ def dump_label_mask_verl080(mask_2d: torch.Tensor, tag: str) -> None: _save_tensor(f"label_mask_{tag}.pt", mask_2d.to(torch.bool), dump_dir) -def dump_entropy_2d_verl080(ent_2d: torch.Tensor | None, tag: str) -> None: - """存 2D entropy ``[B, L_max]``,文件名 ``entropy_{tag}.pt``。``None`` 时跳过。""" +def dump_entropy_2d_verl080(ent_2d: torch.Tensor | None, tag: str, + scope: str = "global") -> None: + """存 2D entropy ``[B, L_max]``,文件名 ``entropy_{tag}.pt``。``None`` 时跳过。 + + ``scope="pp_last"`` 时仅最后一个 PP stage 落盘(多卡下 entropy 只在末 stage 产生)。 + """ dump_dir = _get_dump_dir() if dump_dir is None or ent_2d is None: return - _save_tensor(f"entropy_{tag}.pt", ent_2d, dump_dir) + _save_tensor(f"entropy_{tag}.pt", ent_2d, dump_dir, scope=scope) + + +# ════════════════════════════════════════════════════════════════ +# Post-RoPE Q/K dump (per layer) +# ════════════════════════════════════════════════════════════════ + +_ROPE_POSTQK_BUFFER: dict[int, dict] | None = None + + +def dump_rope_postqk_verl080(layer_number: int, + rotated_query: torch.Tensor, + rotated_key: torch.Tensor, + num_layers: int, + positions: torch.Tensor | None = None) -> None: + """Accumulate one layer's post-RoPE Q/K. Auto-flush to ``rope_postqk.pt`` on last layer. + + Format: ``{layer_idx: {"query": [T, H, D], "key": [T, H, D], "positions": [T] or None}}`` + ON packed 只含 suffix(裁剪后),OFF packed 含完整序列。 + cmp 侧用 prefix_lens + cu_seqlens 做 suffix 对齐后对比(同 attn_output 模式)。 + positions 可选,用于手动排查时的位置回溯。 + """ + global _ROPE_POSTQK_BUFFER + dump_dir = _get_dump_dir() + if dump_dir is None: + return + if _ROPE_POSTQK_BUFFER is None: + _ROPE_POSTQK_BUFFER = {} + entry = { + "query": rotated_query.detach().cpu().clone(), + "key": rotated_key.detach().cpu().clone(), + } + if positions is not None: + entry["positions"] = positions.detach().cpu().clone() + _ROPE_POSTQK_BUFFER[layer_number] = entry + if layer_number == _stage_last_layer(num_layers): + _flush_dict_buffer("rope_postqk.pt", _ROPE_POSTQK_BUFFER, dump_dir) + _ROPE_POSTQK_BUFFER = None + + +def _flush_dict_buffer(fname: str, buffer: dict, dump_dir: str) -> None: + """rank0 直接 torch.save 一个 dict buffer。PP-aware gating + suffix。 + + 不能用 _save_tensor:它对入参做 .detach().cpu().clone(),dict 没 .detach() → + AttributeError 被其 except 吞掉,文件永不写盘(rope_postqk.pt 曾因此丢失)。 + entries 应在插入时已 detach().cpu().clone()。仿 _flush_attn_buffer。 + + Under PP, each stage's tp_rank==0 writes with ``_pp{r}`` suffix; + assembly later merges all stage files. + """ + import os as _os + from prefix_sharing.tools.diagnostic_dump import _should_write_for_scope + if not _should_write_for_scope("pp_stage"): + return + try: + stem, sep, ext = fname.rpartition(".") + pp_sfx = _pp_suffix() + fname_pp = f"{stem}{pp_sfx}{sep}{ext}" if sep else f"{fname}{pp_sfx}" + torch.save(buffer, _os.path.join(dump_dir, fname_pp)) + except Exception as _e: + print(f"[PS-diag] {fname} save failed: {_e}", flush=True) + + +_ROPE_PREQK_BUFFER: dict[int, dict] | None = None + + +def dump_rope_preqk_verl080(layer_number: int, + query: torch.Tensor, + key: torch.Tensor, + num_layers: int) -> None: + """Accumulate one layer's pre-RoPE Q/K. Auto-flush to ``rope_preqk.pt`` on last layer. + + Format: ``{layer_idx: {"query": [T, H, D], "key": [T, H, D]}}``。 + 旋转前的 Q/K(纯 QKV 投影输出,未加位置编码)。ON/OFF 应逐元素相同—— + 同 hidden_states、同 QKV 权重。用来隔离:pre-RoPE 相同但 post-RoPE 不同 → + 问题在 RoPE;pre 就不同 → 问题在上游(hidden_states / 投影)。 + """ + global _ROPE_PREQK_BUFFER + dump_dir = _get_dump_dir() + if dump_dir is None: + return + if _ROPE_PREQK_BUFFER is None: + _ROPE_PREQK_BUFFER = {} + _ROPE_PREQK_BUFFER[layer_number] = { + "query": query.detach().cpu().clone(), + "key": key.detach().cpu().clone(), + } + if layer_number == _stage_last_layer(num_layers): + _flush_dict_buffer("rope_preqk.pt", _ROPE_PREQK_BUFFER, dump_dir) + _ROPE_PREQK_BUFFER = None + + +_EXPANDED_KV_BUFFER: dict[int, dict] | None = None + + +def dump_expanded_kv_on(layer_number: int, expanded_key: torch.Tensor, + expanded_value: torch.Tensor, num_layers: int) -> None: + """ON: 累加 build_kv 输出(expanded K/V = prefix 复用 + suffix,全量),满层 flush ``expanded_kv.pt``。 + + Format: ``{layer_idx: {"key": [T,H,D], "value": [T,H,D]}}``。这是 ON attention 实际用的 + 完整 K/V,应与 OFF ``full_kv.pt`` 逐元素相同——验证 prefix-sharing 的 KV 展开/复用 + 是否正确还原了完整 KV。 + """ + global _EXPANDED_KV_BUFFER + dump_dir = _get_dump_dir() + if dump_dir is None: + return + if _EXPANDED_KV_BUFFER is None: + _EXPANDED_KV_BUFFER = {} + _EXPANDED_KV_BUFFER[layer_number] = { + "key": expanded_key.detach().cpu().clone(), + "value": expanded_value.detach().cpu().clone(), + } + if layer_number == _stage_last_layer(num_layers): + _flush_dict_buffer("expanded_kv.pt", _EXPANDED_KV_BUFFER, dump_dir) + _EXPANDED_KV_BUFFER = None + + +_FULL_KV_BUFFER: dict[int, dict] | None = None + + +def dump_full_kv_off(layer_number: int, key: torch.Tensor, value: torch.Tensor, + num_layers: int) -> None: + """OFF: 累加完整 K(post-RoPE)/ V,满层 flush ``full_kv.pt``。 + + Format: ``{layer_idx: {"key": [T,H,D], "value": [T,H,D]}}``。key 应为 post-RoPE 完整 K + (与 ON expanded_key 同语义),value 为完整 V。供与 ON expanded_kv 逐元素对比。 + """ + global _FULL_KV_BUFFER + dump_dir = _get_dump_dir() + if dump_dir is None: + return + if _FULL_KV_BUFFER is None: + _FULL_KV_BUFFER = {} + _FULL_KV_BUFFER[layer_number] = { + "key": key.detach().cpu().clone(), + "value": value.detach().cpu().clone(), + } + if layer_number == _stage_last_layer(num_layers): + _flush_dict_buffer("full_kv.pt", _FULL_KV_BUFFER, dump_dir) + _FULL_KV_BUFFER = None + + +_BUILD_KV_INPUT_V_BUFFER: dict[int, torch.Tensor] | None = None + + +def dump_build_kv_input_v_on(layer_number: int, value: torch.Tensor, + num_layers: int) -> None: + """ON: build_kv 输入的 V(get_qkv 出来、build_kv 之前的 raw V)。满层 flush ``build_kv_input_v.pt``。 + + Format: ``{layer_idx: [T_on, ...]}``。供与 OFF ``full_kv.pt`` 的 V 做 suffix 对比—— + 定位 V 是在 build_kv 之前(get_qkv/hidden_states)就偏,还是 build_kv 引入。 + T_on vs T_off 还能看出 ON 有没有把 hidden_states 裁剪成 suffix-only。 + """ + global _BUILD_KV_INPUT_V_BUFFER + dump_dir = _get_dump_dir() + if dump_dir is None: + return + if _BUILD_KV_INPUT_V_BUFFER is None: + _BUILD_KV_INPUT_V_BUFFER = {} + _BUILD_KV_INPUT_V_BUFFER[layer_number] = value.detach().cpu().clone() + if layer_number == _stage_last_layer(num_layers): + _flush_dict_buffer("build_kv_input_v.pt", _BUILD_KV_INPUT_V_BUFFER, dump_dir) + _BUILD_KV_INPUT_V_BUFFER = None + + +_HIDDEN_STATES_BUFFER: dict[int, torch.Tensor] | None = None + + +def dump_hidden_states_on(layer_number: int, hidden_states: torch.Tensor, + num_layers: int) -> None: + """Accumulate hidden_states at attention entrance. Auto-flush ``hidden_states.pt`` on last layer. + + Used to verify whether ON/OFF hidden_states are bit-identical for suffix tokens. + Format: ``{layer_idx: [T, H]}``. + """ + global _HIDDEN_STATES_BUFFER + dump_dir = _get_dump_dir() + if dump_dir is None: + return + if _HIDDEN_STATES_BUFFER is None: + _HIDDEN_STATES_BUFFER = {} + _HIDDEN_STATES_BUFFER[layer_number] = hidden_states.detach().cpu().clone() + if layer_number == _stage_last_layer(num_layers): + _flush_dict_buffer("hidden_states.pt", _HIDDEN_STATES_BUFFER, dump_dir) + _HIDDEN_STATES_BUFFER = None + + +@contextlib.contextmanager +def capture_rope_qk(attention_module): + """Hook apply_rotary_pos_emb(post-RoPE Q/K)+ get_query_key_value_tensors(pre-RoPE Q/K/V)。 + + 两个 hook,全 in-context: + - get_qkv 返回值统一截 Q/K/V(pre-squeeze,与 ON 侧 get_qkv+squeeze 之后 dump 同源)。 + - apply_rotary_pos_emb 截 post-RoPE Q/K(返回值)。 + + Q/K/V 都从 get_qkv 一次返回取(pre-RoPE),避免 ON 侧 Q/K/V 也走两个不对称来源。 + 用法:: + + with capture_rope_qk(self) as (qk_caps, qkv_caps): + result = original_forward(...) + # qk_caps[0]={"post":Q_post}, qk_caps[1]={"post":K_post}(apply_rotary 返回) + # qkv_caps[-1] = (Q_pre, K_pre, V_pre)(get_qkv 返回,pre-squeeze) + """ + import megatron.core.transformer.attention as _attn_mod + import types as _types + + _orig_arpe = _attn_mod.apply_rotary_pos_emb + _orig_get_qkv = type(attention_module).get_query_key_value_tensors + qk_captures: list = [] # apply_rotary 返回(post-RoPE) + qkv_captures: list = [] # get_qkv 返回 (Q, K, V) pre-squeeze + + def _capturing_arpe(t, *args, **kwargs): + out = _orig_arpe(t, *args, **kwargs) + qk_captures.append({"post": out}) + return out + + def _capturing_get_qkv(self_, *args, **kwargs): + out = _orig_get_qkv(self_, *args, **kwargs) + try: + qkv_captures.append(tuple(out[:3])) # (Q, K, V) + except Exception: + pass + return out + + _attn_mod.apply_rotary_pos_emb = _capturing_arpe + attention_module.get_query_key_value_tensors = _types.MethodType( + _capturing_get_qkv, attention_module) + try: + yield qk_captures, qkv_captures + finally: + _attn_mod.apply_rotary_pos_emb = _orig_arpe + try: + del attention_module.get_query_key_value_tensors + except AttributeError: + pass diff --git a/prefix-sharing/prefix_sharing/tools/inject_baseline_synthetic.py b/prefix-sharing/prefix_sharing/tools/inject_baseline_synthetic.py new file mode 100644 index 00000000..9c9f335a --- /dev/null +++ b/prefix-sharing/prefix_sharing/tools/inject_baseline_synthetic.py @@ -0,0 +1,115 @@ +"""Inject synthetic data for GEMM precision baseline — reuse nested-prefix +batch builder, then shuffle sequence order and optionally stack. + +Usage:: + + PREFIX_SHARING_BASELINE_SYNTHETIC=/path/to/data.json + PREFIX_SHARING_BASELINE_NUM_SEQ=4 # gen_batch_size (sequences to build) + PREFIX_SHARING_BASELINE_STACK=3 # stack the shuffled batch N times + PREFIX_SHARING_BASELINE_SEED=42 # shuffle seed + + # Single-copy: stack=1 → 4 sequences (no stacking) + # Multi-copy: stack=3 → 12 sequences (4 shuffled × 3 stacked) +""" + +import random + +import torch + + +def patch_baseline_synthetic( + trainer, + json_path: str, + max_prompt_length: int, + max_response_length: int, + num_seq: int = 1, + stack: int = 1, + seed: int = 42, + num_workers: int = 8, +): + """Monkey-patch generate_sequences to return baseline synthetic data. + + Builds ``num_seq`` sequences via the existing :func:`_build_synthetic_batch`, + deterministically shuffles their order, then stacks the shuffled batch + ``stack`` times. + + Args: + trainer: RayPPOTrainer instance. + json_path: Path to JSON with input_ids. + max_prompt_length: Prompt length. + max_response_length: Response length. + num_seq: Number of distinct sequences (env: BASELINE_NUM_SEQ). + stack: How many times to tile the shuffled batch (env: BASELINE_STACK). + seed: Shuffle seed (env: BASELINE_SEED). + num_workers: Agent loop workers for chunk() divisibility. + """ + from prefix_sharing.tools.inject_synthetic_prefix import _build_synthetic_batch, _load_base_tokens + from verl.protocol import DataProto + + base_tokens = _load_base_tokens(json_path) + + # ── 1. Build the batch using the existing nested-prefix builder ── + batch = _build_synthetic_batch( + base_tokens=base_tokens, + batch_size=num_seq, + max_prompt_length=max_prompt_length, + max_response_length=max_response_length, + ) + + # ── 2. Shuffle sequence order (deterministic) ── + rng = random.Random(seed) + n = batch["input_ids"].shape[0] + idx = list(range(n)) + rng.shuffle(idx) + for k in batch: + if isinstance(batch[k], torch.Tensor) and batch[k].shape[0] == n: + batch[k] = batch[k][idx] + + # ── 3. Stack (tile) the shuffled batch ── + if stack > 1: + for k in batch: + if isinstance(batch[k], torch.Tensor) and batch[k].shape[0] == n: + batch[k] = batch[k].repeat(stack, *([1] * (batch[k].dim() - 1))) + + total_bs = n * stack + print( + f"[BaselineSynthetic] num_seq={num_seq} stack={stack} total_bs={total_bs}" + f" P={max_prompt_length} R={max_response_length} seed={seed}" + f" shuffle_idx={idx}" + ) + + # ── 4. Multi-modal placeholder (verl >= 0.8.0) ── + non_tensors = None + try: + import verl + from packaging.version import parse as parse_version + + if parse_version(verl.__version__) > parse_version("0.7.99"): + import numpy as np + non_tensors = {"multi_modal_inputs": np.array([{}] * total_bs, dtype=object)} + except Exception: + pass + + fixed_data = DataProto.from_dict(batch, non_tensors=non_tensors) + + # Pad to be divisible by num_workers + rem = len(fixed_data) % num_workers + if rem: + fixed_data.padding(num_workers - rem, "last") + print(f"[BaselineSynthetic] Padded {len(fixed_data) - (num_workers - rem)} -> {len(fixed_data)}") + + def _patched(batch, **kwargs): + print( + f"[BaselineSynthetic] Returning synthetic baseline data " + f"(num_seq={num_seq}, stack={stack}, total_bs={total_bs}, " + f"P={max_prompt_length}, R={max_response_length}, seed={seed})." + ) + fixed_data.meta_info["timing"] = {} + return fixed_data + + trainer.actor_rollout_wg.generate_sequences = _patched + print("[BaselineSynthetic] Patched actor_rollout_wg.generate_sequences.") + + if hasattr(trainer, "async_rollout_manager") and trainer.async_rollout_manager is not None: + trainer.async_rollout_manager.generate_sequences = _patched + print("[BaselineSynthetic] Patched async_rollout_manager.generate_sequences.")