Download serving/reference_forward.py from perplexity-ai/PII-Tracer-vLLM: direct link, hf CLI and curl.
- Browser
- Download file 6.13 kB
-
https://huggingface.co/perplexity-ai/PII-Tracer-vLLM/resolve/main/serving/reference_forward.py
- Command line
-
hf download hf://perplexity-ai/PII-Tracer-vLLM/serving/reference_forward.py
-
curl -L -o reference_forward.py https://huggingface.co/perplexity-ai/PII-Tracer-vLLM/resolve/main/serving/reference_forward.py
6.13 kB
| #!/usr/bin/env python3 | |
| """Independent reference forward pass for the pplx-pii-masking checkpoint, used | |
| to verify what vLLM serves. | |
| Runs the ORIGINAL checkpoint (`backbone.*` + `token_cls_head` + | |
| `sensitivity_head`, straight off HuggingFace) in plain fp32 torch on CPU, with | |
| no vLLM and no transformers modeling code: embed -> 28 x (RMSNorm, GQA attention | |
| with q/k per-head RMSNorm and RoPE, SwiGLU MLP) -> final norm -> heads. Attention | |
| is fully bidirectional, which the vLLM deployment reproduces with | |
| `is_causal: false` in its config.json. | |
| Compare against the served path with --scoring-url; it fetches the same text from | |
| the /v1/scoring adapter and reports max abs delta, cosine, and label agreement. | |
| docker run --rm --user "$(id -u):$(id -g)" --entrypoint python3 --network host \ | |
| -v "$PWD":/work:ro -v ~/models/pplx-pii-masking:/src:ro \ | |
| vllm/vllm-openai:v0.26.0 /work/reference_forward.py --src /src --text "..." | |
| """ | |
| import argparse | |
| import json | |
| import math | |
| from pathlib import Path | |
| import torch | |
| from safetensors.torch import load_file | |
| def rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: | |
| var = x.pow(2).mean(-1, keepdim=True) | |
| return x * torch.rsqrt(var + eps) * weight | |
| def rope_tables(seq_len: int, head_dim: int, theta: float): | |
| inv = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) | |
| pos = torch.arange(seq_len, dtype=torch.float32) | |
| freqs = torch.outer(pos, inv) # [T, hd/2] | |
| emb = torch.cat([freqs, freqs], dim=-1) # [T, hd] (HF rotate_half layout) | |
| return emb.cos(), emb.sin() | |
| def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| # x: [T, H, D]; rotate_half convention, matching HF Qwen3. | |
| half = x.shape[-1] // 2 | |
| rotated = torch.cat([-x[..., half:], x[..., :half]], dim=-1) | |
| return x * cos[:, None, :] + rotated * sin[:, None, :] | |
| def forward(sd: dict, cfg: dict, ids: list[int]) -> tuple[torch.Tensor, float]: | |
| bb = cfg["backbone"] | |
| n_layers = bb["num_hidden_layers"] | |
| n_heads, n_kv = bb["num_attention_heads"], bb["num_key_value_heads"] | |
| head_dim, eps = bb["head_dim"], bb["rms_norm_eps"] | |
| theta = bb["rope_parameters"]["rope_theta"] | |
| rep = n_heads // n_kv | |
| seq_len = len(ids) | |
| g = lambda k: sd[k].to(torch.float32) # noqa: E731 - terse weight getter | |
| h = g("backbone.embed_tokens.weight")[torch.tensor(ids)] # [T, C] | |
| cos, sin = rope_tables(seq_len, head_dim, theta) | |
| for i in range(n_layers): | |
| p = f"backbone.layers.{i}." | |
| x = rms_norm(h, g(p + "input_layernorm.weight"), eps) | |
| q = (x @ g(p + "self_attn.q_proj.weight").T).view(seq_len, n_heads, head_dim) | |
| k = (x @ g(p + "self_attn.k_proj.weight").T).view(seq_len, n_kv, head_dim) | |
| v = (x @ g(p + "self_attn.v_proj.weight").T).view(seq_len, n_kv, head_dim) | |
| # Qwen3 normalises per head, before RoPE. | |
| q = rms_norm(q, g(p + "self_attn.q_norm.weight"), eps) | |
| k = rms_norm(k, g(p + "self_attn.k_norm.weight"), eps) | |
| q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin) | |
| k = k.repeat_interleave(rep, dim=1) # GQA | |
| v = v.repeat_interleave(rep, dim=1) | |
| # [H, T, D]; no mask at all -- bidirectional/encoder-only attention. | |
| scores = torch.einsum("qhd,khd->hqk", q, k) / math.sqrt(head_dim) | |
| ctx = torch.einsum("hqk,khd->qhd", scores.softmax(-1), v) | |
| attn = ctx.reshape(seq_len, n_heads * head_dim) @ g(p + "self_attn.o_proj.weight").T | |
| h = h + attn | |
| x = rms_norm(h, g(p + "post_attention_layernorm.weight"), eps) | |
| gate = x @ g(p + "mlp.gate_proj.weight").T | |
| up = x @ g(p + "mlp.up_proj.weight").T | |
| h = h + (torch.nn.functional.silu(gate) * up) @ g(p + "mlp.down_proj.weight").T | |
| h = rms_norm(h, g("backbone.norm.weight"), eps) | |
| token_logits = h @ g("token_cls_head.weight").T + g("token_cls_head.bias") | |
| # Sensitivity head is applied to the MEAN-POOLED hidden state in the original. | |
| sens = (h.mean(0) @ g("sensitivity_head.weight").T + g("sensitivity_head.bias")).item() | |
| return token_logits, sens | |
| def main() -> None: | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--src", default="/src", help="original HF checkpoint folder") | |
| ap.add_argument("--text", default="My name is John Smith and my SSN is 123-45-6789.") | |
| ap.add_argument("--scoring-url", default="http://127.0.0.1:8002") | |
| ap.add_argument("--out", help="write comparison JSON here") | |
| a = ap.parse_args() | |
| src = Path(a.src) | |
| cfg = json.loads((src / "config.json").read_text()) | |
| sd = load_file(str(src / "model.safetensors")) | |
| import urllib.request | |
| req = urllib.request.Request( | |
| f"{a.scoring_url}/v1/scoring", | |
| data=json.dumps({"model": "pii-masking-latest", "sequences": [a.text]}).encode(), | |
| headers={"Content-Type": "application/json"}) | |
| served = json.loads(urllib.request.urlopen(req, timeout=120).read()) | |
| ids = served["token_ids"][0] | |
| served_logits = torch.tensor(served["pruning_scores"][0], dtype=torch.float32) | |
| served_sens = served["ranking_scores"][0][0] | |
| ref_logits, ref_sens = forward(sd, cfg, ids) | |
| assert ref_logits.shape == served_logits.shape, (ref_logits.shape, served_logits.shape) | |
| delta = (ref_logits - served_logits).abs() | |
| cos = torch.nn.functional.cosine_similarity(ref_logits, served_logits, dim=-1) | |
| agree = int((ref_logits.argmax(-1) == served_logits.argmax(-1)).sum()) | |
| result = { | |
| "tokens": len(ids), | |
| "max_abs_delta": round(delta.max().item(), 5), | |
| "mean_abs_delta": round(delta.mean().item(), 6), | |
| "min_cosine": round(cos.min().item(), 8), | |
| "argmax_agreement": f"{agree}/{len(ids)}", | |
| "sensitivity_ref": round(ref_sens, 5), | |
| "sensitivity_served": round(served_sens, 5), | |
| "sensitivity_delta": round(abs(ref_sens - served_sens), 6), | |
| } | |
| print(json.dumps(result, indent=2)) | |
| if a.out: | |
| Path(a.out).write_text(json.dumps({"text": a.text, **result}, indent=2) + "\n") | |
| if __name__ == "__main__": | |
| main() | |