Piko-9b / scripts /diagnose_generation.py
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
Raw
History Blame Contribute Delete
5.8 kB
#!/usr/bin/env python3
"""Diagnose degenerate generation: is it the checkpoint or the placement?
The composed checkpoint emits a constant token under `device_map="auto"` with
CPU offload. This script isolates the cause by checking, for one short prompt:
* whether hidden states or logits contain NaN/Inf
* where in the layer stack the first non-finite value appears
* whether the result changes without CPU offload (4-bit, fully resident)
Run it against the composed checkpoint and against the language-only source
checkpoint to tell a composition bug apart from an environment bug.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Any
import torch
def load(model_path: str, mode: str) -> tuple[Any, Any]:
from transformers import AutoModelForMultimodalLM, AutoTokenizer
kwargs: dict[str, Any] = {"dtype": torch.bfloat16}
if mode == "offload":
kwargs["device_map"] = "auto"
elif mode == "gpu4bit":
from transformers import BitsAndBytesConfig
kwargs["quantization_config"] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
kwargs["device_map"] = {"": 0}
elif mode == "cpu":
kwargs["device_map"] = {"": "cpu"}
kwargs["dtype"] = torch.float32
else:
raise ValueError(mode)
model = AutoModelForMultimodalLM.from_pretrained(model_path, **kwargs)
model.eval()
tokenizer = AutoTokenizer.from_pretrained(model_path)
return model, tokenizer
def encode(tokenizer: Any, prompt: str, device: Any) -> torch.Tensor:
"""Return a plain input_ids tensor across transformers 4.x / 5.x behaviour."""
encoded = tokenizer.apply_chat_template(
[{"role": "user", "content": prompt}],
add_generation_prompt=True,
return_tensors="pt",
)
if not isinstance(encoded, torch.Tensor):
encoded = encoded["input_ids"]
return encoded.to(device)
def probe(model: Any, tokenizer: Any, prompt: str) -> dict[str, Any]:
ids = encode(tokenizer, prompt, model.device)
with torch.inference_mode():
out = model(input_ids=ids, output_hidden_states=True, use_cache=False)
logits = out.logits[0, -1].float()
hidden = out.hidden_states
first_bad = None
layer_stats = []
for index, state in enumerate(hidden):
tensor = state.float()
finite = bool(torch.isfinite(tensor).all())
layer_stats.append(
{
"layer": index,
"finite": finite,
"absmax": float(tensor.abs().max()) if finite else None,
"std": float(tensor.std()) if finite else None,
}
)
if not finite and first_bad is None:
first_bad = index
top = torch.topk(logits, 5)
return {
"prompt": prompt,
"prompt_tokens": int(ids.shape[1]),
"logits_finite": bool(torch.isfinite(logits).all()),
"logits_absmax": float(logits.abs().max()) if bool(torch.isfinite(logits).all()) else None,
"logits_nan_count": int(torch.isnan(logits).sum()),
"logits_inf_count": int(torch.isinf(logits).sum()),
"first_nonfinite_hidden_layer": first_bad,
"hidden_layer_count": len(hidden),
"top5_token_ids": [int(i) for i in top.indices],
"top5_tokens": [tokenizer.decode([int(i)]) for i in top.indices],
"top5_logits": [float(v) for v in top.values],
"layer_stats": layer_stats,
}
def short_generate(model: Any, tokenizer: Any, prompt: str, n: int = 24) -> str:
ids = encode(tokenizer, prompt, model.device)
with torch.inference_mode():
out = model.generate(ids, max_new_tokens=n, do_sample=False)
return tokenizer.decode(out[0][ids.shape[1] :], skip_special_tokens=True).strip()
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", required=True)
parser.add_argument("--mode", default="gpu4bit", choices=["offload", "gpu4bit", "cpu"])
parser.add_argument("--label", default=None)
parser.add_argument("--prompt", default="What is the capital of France?")
parser.add_argument("--output", type=Path, default=Path("reports/generation_diagnosis.json"))
args = parser.parse_args()
print(f"Loading {args.model} [{args.mode}]", flush=True)
model, tokenizer = load(args.model, args.mode)
record = {
"model": args.model,
"label": args.label or Path(args.model).name,
"mode": args.mode,
"device_map": str(getattr(model, "hf_device_map", None)),
"probe": probe(model, tokenizer, args.prompt),
"generation": short_generate(model, tokenizer, args.prompt),
}
print(json.dumps({k: v for k, v in record.items() if k != "probe"}, indent=2))
p = record["probe"]
print(
f"logits finite={p['logits_finite']} nan={p['logits_nan_count']} "
f"inf={p['logits_inf_count']} first_bad_hidden={p['first_nonfinite_hidden_layer']}"
)
pairs = list(zip(p["top5_tokens"], [round(v, 2) for v in p["top5_logits"]], strict=True))
print(f"top5: {pairs}")
args.output.parent.mkdir(parents=True, exist_ok=True)
existing = []
if args.output.is_file():
existing = json.loads(args.output.read_text(encoding="utf-8"))
existing.append(record)
args.output.write_text(json.dumps(existing, indent=2) + "\n", encoding="utf-8")
print(f"\nAppended to {args.output}")
if __name__ == "__main__":
main()