Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
62 commits
Select commit Hold shift + click to select a range
e963310
[test] 修正链式复用测试:删孤儿测试,用自然链场景替换错误预设
boundless-future Jun 24, 2026
16877df
[refactor] logprob/entropy restore 改为 build_kv 式区间拼接,清理冗余 interior-sp…
boundless-future Jun 25, 2026
3260dc6
[refactor] 2D注入流程命名优化与docstring收敛
boundless-future Jun 25, 2026
f8bded1
[refactor] diag dump失败提示改用print
boundless-future Jun 25, 2026
ab4943f
[test] NPU flash-attn: sparse_mode=3 TND varlen 可行性实验
boundless-future Jun 25, 2026
8794180
[feat] 新增 NPU flash-attn TND varlen + sparse_mode=3 后端
boundless-future Jun 26, 2026
f6a5dd5
[fix] flash_atten_npu_tnd: _prepare_flash_inputs 解包数量修正(8 元组不是 9)
boundless-future Jun 26, 2026
d9279f6
[refactor] cmp_diag_verl080: first_token 改为可指定 packed 位置的 packed_token
boundless-future Jun 26, 2026
d7aab26
[fix] cmp_diag_verl080: packed_token 比较前先做 suffix 对齐
boundless-future Jun 26, 2026
7f80253
[fix] cmp_diag_verl080: logits 的 val 排序用带符号实际值,不用绝对值
boundless-future Jun 26, 2026
c0de097
[diag] PREFIX_SHARING_FORCE_ZERO_PREFIX: 强制 0-prefix 隔离 kernel-mode 差异
boundless-future Jun 26, 2026
6b122af
[diag] prefix-sharing deviation probes: RoPE extrapolation log + grou…
boundless-future Jun 26, 2026
717d994
[diag] RoPE extrapolation: add step01 vs step12 linearity check
boundless-future Jun 26, 2026
c86ea12
[diag] add post-RoPE Q/K dump for both ON and OFF paths
boundless-future Jun 26, 2026
f5e9d09
[diag] rope_emb: add position IDs to dump + rewrite cmp_rope_emb for …
boundless-future Jun 26, 2026
64fc4e4
[diag] verl080: post-RoPE Q/K dump + cmp (复刻 attn_output 模式)
boundless-future Jun 26, 2026
f4be6b5
[fix] dump_rope_emb_verl080: accept optional positions param
boundless-future Jun 26, 2026
badaba3
[fix] attention OFF diag: local import torch (NameError on torch.aran…
boundless-future Jun 27, 2026
da39661
[fix] OFF post-RoPE Q/K: hook 真实张量替代重算 + 修 rope_emb 写盘 bug
boundless-future Jun 27, 2026
eb9ebc3
[diag] cmp_verl080: rope_emb 入 shapes 表 + 细化 pos 失败诊断
boundless-future Jun 27, 2026
ea76466
[fix] rope_emb.pt 写盘:_save_tensor 对 dict 失败 → 直接 torch.save
boundless-future Jun 27, 2026
a4f3495
[diag] RoPE 全流水线对比:pre Q/K → rope_freqs → post Q/K(layer/token 级)
boundless-future Jun 27, 2026
fad4c9d
[refactor] rope 命名统一:rope_emb→rope_postqk,rope_freqs ON/OFF 合并
boundless-future Jun 27, 2026
b61e353
[diag] cmp: rope top-K dim 跨 stage 对齐(postqk 基准 → preqk/freqs 同 dim)
boundless-future Jun 27, 2026
dc70f97
[fix] OFF rope_freqs dump 崩 forward:_off_positions 缺 device 参数
boundless-future Jun 27, 2026
cc40ce6
[diag] attn_kv 对比:ON expanded_kv vs OFF full_kv(prefix 复用校验)
boundless-future Jun 27, 2026
04d06c6
[diag] OFF V 改 in-context 截获(hook get_query_key_value_tensors),不再 re-…
boundless-future Jun 27, 2026
68b0c44
[fix] capture_rope_qk: get_qkv 双重绑定导致 output_gate multiple values
boundless-future Jun 27, 2026
c7dcc29
[diag] dump build_kv 输入 V,定位 V 偏在 build_kv 之前还是之后
boundless-future Jun 27, 2026
2235b1c
[diag] preqk/postqk per-layer 加 max_diff,与 build_kv_input_v 的 V 同口径对比
boundless-future Jun 27, 2026
bfc06fa
[refactor] ON pre-RoPE Q/K/V dump 统一到 attention.py squeeze 之后
boundless-future Jun 29, 2026
5e68962
[refactor] OFF Q/K/V 统一从 get_qkv 返回截(与 ON 对称)
boundless-future Jun 29, 2026
4cd87dc
[diag] OFF pre-RoPE Q/K/V 侵入式 dump(megatron attention.py squeeze 之后)
boundless-future Jun 29, 2026
7aa39da
[diag] OFF post-RoPE Q/K + full_kv 也侵入式 dump(rotary block 之后)
boundless-future Jun 29, 2026
5b95d70
[fix] OFF dump: remove hook-based rope_postqk/full_kv writes (double-…
boundless-future Jun 29, 2026
91282bd
[diag] add hidden_states dump + cmp at attention entrance (ON 侵入式 + O…
boundless-future Jun 29, 2026
2d5270e
[cmp] add baseline scripts: cross_batch + within_batch GEMM precision
boundless-future Jun 29, 2026
e261dca
[cmp] baseline: add rope_freqs, rope_postqk, attn_outputs to both scr…
boundless-future Jun 29, 2026
a429b5f
[cmp] baseline: add full_kv, logits, logprobs, entropy to both scripts
boundless-future Jun 29, 2026
18d99c1
[inject] add baseline synthetic data injection with shuffle + stack
boundless-future Jun 29, 2026
2742f45
[inject] baseline synthetic: reuse _build_synthetic_batch, shuffle se…
boundless-future Jun 29, 2026
ea68999
[refactor] remove unused batch_size param from patch_baseline_synthetic
boundless-future Jun 30, 2026
21dc6e1
[refactor] baseline cmp: reuse cmp_diag_verl080 _print_* functions, a…
boundless-future Jun 30, 2026
31d3373
[fix] baseline cmp: handle variable-length sequences (nested prefix)
boundless-future Jun 30, 2026
a10bb08
[fix] baseline: compute cos_avg from actual per-token cosine (not har…
boundless-future Jun 30, 2026
9de6f2b
[fix] within_batch top-K: use total_copies from cu_seqlens, not undef…
boundless-future Jun 30, 2026
da791d8
[fix] baseline: compute rel/pearson for 2D, add 2D top-K, fix pearson…
boundless-future Jun 30, 2026
366d1ce
[fix] cross_batch top-K: guard _topk_rope with try-except for non-dic…
boundless-future Jun 30, 2026
ecfac35
[cmp] baseline: add _print_shapes for .pt file shape diagnostics
boundless-future Jun 30, 2026
efe5f6b
[refactor] baseline: replace ON/OFF labels with Single/Stacked via lo…
boundless-future Jun 30, 2026
649484d
[fix] within_batch 2D: compare same-seq stack copies only, add --num-seq
boundless-future Jun 30, 2026
9eefb4d
[fix] within_batch 2D top-K: compare row 0 vs row num_seq (same seq c…
boundless-future Jun 30, 2026
abf2949
[refactor] baseline: full variable rename for readability (open-sourc…
boundless-future Jun 30, 2026
ca37506
[refactor] standardize all PS-diag dump block markers to ##### [PS-di…
boundless-future Jun 30, 2026
14f6fb5
[feat] multi-rank (TP/PP) precision verification: dump routing + asse…
boundless-future Jul 1, 2026
ae48d97
[refactor] comprehensive variable rename in cmp_diag.py and cmp_diag_…
boundless-future Jul 1, 2026
716b416
[fix] assemble_dump: rewrite parallel_info.json with tp=1 pp=1 after …
boundless-future Jul 1, 2026
8b8a042
[refactor] remove TP/PP awareness from cmp_diag_verl080.py; revert as…
boundless-future Jul 1, 2026
b7fe400
[fix] replace torch.nested.nested_tensor with as_nested_tensor for NP…
boundless-future Jul 14, 2026
bbdc5ac
Merge remote-tracking branch 'origin/main' into parallel-precision-ve…
boundless-future Jul 15, 2026
861be13
[fix] resolve CI lint failures (F401/F541/F821/F841)
boundless-future Jul 15, 2026
c6b0ecd
[cleanup] remove deprecated flash_atten_npu_tnd and flash_atten_npu_t…
boundless-future Jul 16, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -1069,6 +1069,23 @@ def forward(
value = value.squeeze(1)

Copy link
Copy Markdown
Owner

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

开发目标 A:OFF 路径插桩 pre-RoPE Q/K/V + hidden_states dump(post-squeeze、pre-rotary)

  • dump_rope_preqk_verl080(layer_number, query, key, num_layers)
  • dump_build_kv_input_v_on(layer_number, value, num_layers)
  • dump_hidden_states_on(layer_number, hidden_states, num_layers)

开发目标 B:OFF 路径插桩 post-RoPE Q/K + full KV dump(rotary 后、core attention 前)

  • dump_rope_postqk_verl080(layer_number, query, key, num_layers)
  • dump_full_kv_off(layer_number, key, value, num_layers)

两者由 PREFIX_SHARING_DIAG_DUMP env var 控制,未设时仅一次 os.environ.get 检查。try/except 保护 dump 失败不崩训练。

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)
# ================================================
Expand Down Expand Up @@ -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
# ==================================
Expand Down
18 changes: 18 additions & 0 deletions dependency/verl_cdd9014f/verl/trainer/ppo/ray_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -1354,8 +1354,26 @@ def fit(self):
max_prompt_length=self.config.data.max_prompt_length,

Copy link
Copy Markdown
Owner

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

开发目标:在 fit() 训练循环前注入 synthetic baseline 数据

  • 第 1354-1374 行:读 PREFIX_SHARING_BASELINE_SYNTHETIC / NUM_SEQ / STACK / SEED env var
  • 调用 inject_baseline_synthetic.patch_baseline_synthetic() 替换 dataloader
  • 用于无真实数据加载器的 CI 环境

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
Expand Down
8 changes: 4 additions & 4 deletions prefix-sharing/prefix_sharing/backends/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
14 changes: 14 additions & 0 deletions prefix-sharing/prefix_sharing/core/prefix_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
139 changes: 129 additions & 10 deletions prefix-sharing/prefix_sharing/integrations/megatron_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)}"

Copy link
Copy Markdown
Owner

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

开发目标 A:修复 PP dump 中 attention output 永不 flush 的 bug(核心 bug 修复)

旧逻辑用 layer_number == num_layers 判断最后层——PP stage 最大 layer_number < num_layers,永远不 flush。

  • 新增 _last_layer_for_stage(layer_number, num_layers, pp_size):按 PP partition 计算当前 stage 的最大层号
  • flush 条件改为 layer_number == _last_layer_for_stage(…)
  • pp_size=1 时退化为旧逻辑,完全兼容

开发目标 B:新增 SKIP_RESTORE env var 诊断 probe

  • restore 前检查 PREFIX_SHARING_SKIP_RESTORE:设值时跳过 restore,用于隔离"分歧在 restore 前/时"
  • 添加 [PS-diag] 打印警告

)

##### [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,
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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).
Expand All @@ -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,
Expand All @@ -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


Expand Down
Loading
Loading