PII-Tracer-vLLM / serving /reference_forward.py
kaiyuanzh's picture
initial commit
e4a672d verified
Raw History Blame Contribute Delete
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()