J-space / math_jlens /explore.py
ayh015's picture
Upload folder using huggingface_hub
85b17bd verified
Raw
History Blame Contribute Delete
20.5 kB
"""Interactive, on-demand J-lens explorer for held-out DAPO Math samples."""
from __future__ import annotations
import argparse
import json
import random
import re
import shlex
from dataclasses import dataclass
from pathlib import Path
import torch
from .model import load_qwen
def extract_boxed(text: str) -> str | None:
"""Return the content of the last balanced ``\\boxed{...}``."""
marker = r"\boxed{"
start = text.rfind(marker)
if start < 0:
return None
content_start = start + len(marker)
depth = 1
for index in range(content_start, len(text)):
if text[index] == "{":
depth += 1
elif text[index] == "}":
depth -= 1
if depth == 0:
return text[content_start:index]
return None
def normalize_answer(answer: str | None) -> str | None:
if answer is None:
return None
value = answer.strip().strip("$")
value = re.sub(r"\\(?:,|!|;|:|quad|qquad)", "", value)
value = value.replace(" ", "").replace(",", "")
return value
def response_answer(text: str) -> str | None:
boxed = extract_boxed(text)
if boxed is not None:
return boxed
matches = re.findall(r"(?im)^\s*Answer\s*:\s*(.+?)\s*$", text)
return matches[-1] if matches else None
@dataclass(frozen=True)
class DapoSample:
original_index: int
messages: list[dict]
label: str
def load_held_out_sample(path: str, *, seed: int, fit_count: int, sample_index: int) -> DapoSample:
records: list[tuple[int, dict]] = []
with Path(path).open(encoding="utf-8") as handle:
for index, line in enumerate(handle):
records.append((index, json.loads(line)))
random.Random(seed).shuffle(records)
held_out_index = fit_count + sample_index
if not 0 <= held_out_index < len(records):
raise IndexError(f"held-out sample {sample_index} is outside the dataset")
original_index, record = records[held_out_index]
prompt = record["prompt"]
messages = prompt if isinstance(prompt, list) else [{"role": "user", "content": prompt}]
return DapoSample(original_index, messages, str(record["label"]))
def find_last_subsequence(sequence: list[int], pattern: list[int]) -> int | None:
if not pattern or len(pattern) > len(sequence):
return None
for start in range(len(sequence) - len(pattern), -1, -1):
if sequence[start:start + len(pattern)] == pattern:
return start
return None
class Explorer:
def __init__(self, args: argparse.Namespace) -> None:
self.args = args
self.adapter = load_qwen(args.model, device=args.device, dtype=torch.bfloat16)
self.model = self.adapter.model
self.tokenizer = self.adapter.tokenizer
lens_state = torch.load(args.lens, map_location="cpu", weights_only=True)
if lens_state["d_model"] != self.adapter.d_model:
raise ValueError(
f"lens width {lens_state['d_model']} does not match model width {self.adapter.d_model}"
)
self.layer_numbers = sorted(lens_state["J"])
self.jacobians = {
int(layer): matrix.to(args.device)
for layer, matrix in lens_state["J"].items()
}
self.cache_dir = Path(args.cache_dir)
self.cache_dir.mkdir(parents=True, exist_ok=True)
self.activations: dict[int, torch.Tensor] = {}
self.sample: DapoSample | None = None
self.full_ids: torch.Tensor | None = None
self.prompt_length = 0
self.response_text = ""
self.current_position = 0
self.load_sample(args.sample_index)
def _cache_path(self, sample_index: int) -> Path:
return self.cache_dir / f"seed-{self.args.seed}-fit-{self.args.fit_count}-sample-{sample_index}.json"
def _chat_prompt_ids(self, messages: list[dict]) -> torch.Tensor:
ids = self.tokenizer.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_tensors="pt",
)
return ids.to(self.args.device)
@torch.inference_mode()
def _generate(self, sample: DapoSample) -> tuple[list[int], str]:
prompt_ids = self._chat_prompt_ids(sample.messages)
attention_mask = torch.ones_like(prompt_ids)
generated = self.model.generate(
input_ids=prompt_ids,
attention_mask=attention_mask,
do_sample=False,
max_new_tokens=self.args.max_new_tokens,
use_cache=True,
pad_token_id=self.tokenizer.eos_token_id,
)
response_ids = generated[0, prompt_ids.shape[1]:].tolist()
response = self.tokenizer.decode(response_ids, skip_special_tokens=True)
return response_ids, response
def _load_or_generate(self, sample_index: int, sample: DapoSample) -> tuple[list[int], str]:
path = self._cache_path(sample_index)
if path.exists() and not self.args.regenerate:
cached = json.loads(path.read_text(encoding="utf-8"))
if cached["original_index"] != sample.original_index:
raise ValueError(f"cache identity mismatch at {path}")
return cached["response_ids"], cached["response_text"]
response_ids, response = self._generate(sample)
path.write_text(
json.dumps(
{
"sample_index": sample_index,
"original_index": sample.original_index,
"label": sample.label,
"response_ids": response_ids,
"response_text": response,
"max_new_tokens": self.args.max_new_tokens,
"do_sample": False,
},
ensure_ascii=False,
indent=2,
),
encoding="utf-8",
)
return response_ids, response
@torch.inference_mode()
def _capture(self, full_ids: torch.Tensor) -> None:
self.activations.clear()
handles = []
for layer, block in enumerate(self.adapter.layers):
def capture(_module, _inputs, output, layer=layer):
tensor = output if torch.is_tensor(output) else output[0]
self.activations[layer] = tensor[0].detach()
handles.append(block.register_forward_hook(capture))
try:
self.adapter.forward(full_ids)
finally:
for handle in handles:
handle.remove()
def load_sample(self, sample_index: int) -> None:
sample = load_held_out_sample(
self.args.data,
seed=self.args.seed,
fit_count=self.args.fit_count,
sample_index=sample_index,
)
prompt_ids = self._chat_prompt_ids(sample.messages)
response_ids, response = self._load_or_generate(sample_index, sample)
full_ids = torch.cat(
[prompt_ids, torch.tensor([response_ids], device=self.args.device)], dim=1
)
self._capture(full_ids)
self.args.sample_index = sample_index
self.sample = sample
self.full_ids = full_ids[0]
self.prompt_length = prompt_ids.shape[1]
self.response_text = response
self.current_position = max(0, self.prompt_length - 1)
print(
f"Loaded held-out sample {sample_index} (dataset row {sample.original_index})\n"
f"prompt_tokens={self.prompt_length} response_tokens={len(response_ids)} "
f"total_tokens={len(self.full_ids)}"
)
self.print_result()
def problem_text(self) -> str:
assert self.sample is not None
parts = []
for message in self.sample.messages:
role = str(message.get("role", "unknown")).upper()
content = str(message.get("content", ""))
parts.append(f"[{role}]\n{content}")
return "\n\n".join(parts)
def print_problem(self) -> None:
print("=== Problem ===")
print(self.problem_text())
def print_output(self) -> None:
print("=== Model output ===")
print(self.response_text)
def print_result(self) -> None:
assert self.sample is not None
predicted = response_answer(self.response_text)
normalized_prediction = normalize_answer(predicted)
normalized_gold = normalize_answer(self.sample.label)
correct = normalized_prediction == normalized_gold
print("=== Answers ===")
print(f"Extracted: {predicted!r}")
print(f"Ground truth: {self.sample.label!r}")
print(f"Correct: {correct}")
def print_info(self) -> None:
assert self.sample is not None
print(
f"=== Sample ===\nheld-out index: {self.args.sample_index}\n"
f"dataset row: {self.sample.original_index}"
)
self.print_problem()
self.print_output()
self.print_result()
def _check_position(self, position: int) -> int:
assert self.full_ids is not None
if position < 0:
position += len(self.full_ids)
if not 0 <= position < len(self.full_ids) - 1:
raise ValueError(f"position must be in [0, {len(self.full_ids) - 2}]")
return position
@torch.inference_mode()
def logits(self, layer: int, position: int, *, use_jacobian: bool = True) -> torch.Tensor:
residual = self.activations[layer][position]
if use_jacobian:
residual = residual @ self.jacobians[layer].T
residual = self.adapter.decoder.norm(residual)
return self.model.lm_head(residual).float()
def token_label(self, token_id: int) -> str:
return repr(self.tokenizer.decode([int(token_id)]))
def print_tokens(self, start: int | None = None, end: int | None = None) -> None:
assert self.full_ids is not None
if start is None:
start = self.prompt_length
if end is None:
end = min(len(self.full_ids), start + 80)
start = max(0, start)
end = min(len(self.full_ids), end)
for position in range(start, end):
marker = ">" if position == self.current_position else " "
region = "R" if position >= self.prompt_length else "P"
print(f"{marker} {position:5d} {region} {self.token_label(self.full_ids[position])}")
def inspect(self, position: int, top_k: int = 5) -> None:
position = self._check_position(position)
self.current_position = position
assert self.full_ids is not None
print(
f"position={position} current={self.token_label(self.full_ids[position])} "
f"predicts={self.token_label(self.full_ids[position + 1])}"
)
for layer in self.layer_numbers:
logits = self.logits(layer, position)
values, ids = logits.topk(top_k)
decoded = " ".join(
f"{self.token_label(token_id)}({value:.2f})"
for token_id, value in zip(ids.tolist(), values.tolist(), strict=True)
)
print(f"L{layer:02d} J: {decoded}")
def compare(self, position: int, layer: int, top_k: int = 10) -> None:
position = self._check_position(position)
if layer not in self.layer_numbers:
raise ValueError(f"layer must be one of {self.layer_numbers}")
for name, enabled in (("J-lens", True), ("logit", False)):
values, ids = self.logits(layer, position, use_jacobian=enabled).topk(top_k)
decoded = " ".join(
f"{self.token_label(token_id)}({value:.2f})"
for token_id, value in zip(ids.tolist(), values.tolist(), strict=True)
)
print(f"{name:7s}: {decoded}")
@torch.inference_mode()
def save_interval(
self,
start: int,
end: int,
layer: int | None = None,
top_k: int = 5,
output_path: str | None = None,
) -> Path:
"""Export original tokens and their J-lens mappings for ``[start, end)``."""
assert self.full_ids is not None and self.sample is not None
if start < 0 or end <= start or end > len(self.full_ids) - 1:
raise ValueError(
f"require 0 <= START < END <= {len(self.full_ids) - 1}; "
"END is exclusive"
)
if top_k < 1:
raise ValueError("TOP_K must be positive")
layers = self.layer_numbers if layer is None else [layer]
if any(item not in self.layer_numbers for item in layers):
raise ValueError(f"layer must be one of {self.layer_numbers}, or 'all'")
layer_label = "all" if layer is None else str(layer)
if output_path is None:
output = Path(self.args.export_dir) / (
f"sample-{self.args.sample_index}-tokens-{start}-{end}-layer-{layer_label}.txt"
)
else:
output = Path(output_path)
output.parent.mkdir(parents=True, exist_ok=True)
lines = [
"J-lens token interval export",
f"held_out_sample: {self.args.sample_index}",
f"dataset_row: {self.sample.original_index}",
f"ground_truth: {self.sample.label!r}",
f"token_interval: [{start}, {end})",
f"layers: {layer_label}",
f"top_k: {top_k}",
"convention: activation at position t predicts token at t+1",
"",
]
for position in range(start, end):
current_id = int(self.full_ids[position])
next_id = int(self.full_ids[position + 1])
region = "response" if position >= self.prompt_length else "prompt"
lines.extend(
[
f"POSITION {position} ({region})",
f" original: id={current_id} token={self.token_label(current_id)}",
f" predicts: id={next_id} token={self.token_label(next_id)}",
]
)
for layer_number in layers:
values, ids = self.logits(layer_number, position).topk(top_k)
mapped = " | ".join(
f"rank={rank} id={token_id} token={self.token_label(token_id)} logit={value:.4f}"
for rank, (token_id, value) in enumerate(
zip(ids.tolist(), values.tolist(), strict=True), start=1
)
)
lines.append(f" L{layer_number:02d}: {mapped}")
lines.append("")
output.write_text("\n".join(lines), encoding="utf-8")
print(f"saved {end - start} positions × {len(layers)} layers to {output}")
return output
def _token_ids(self, text: str) -> list[int]:
return self.tokenizer.encode(text, add_special_tokens=False)
def trace(self, text: str, position: int | None = None) -> None:
position = self.current_position if position is None else self._check_position(position)
ids = self._token_ids(text)
if not ids:
raise ValueError("text tokenized to no tokens")
print(f"trace text={text!r} ids={ids} at position={position}")
for layer in self.layer_numbers:
logits = self.logits(layer, position)
details = []
for token_id in ids:
rank = 1 + int((logits > logits[token_id]).sum().item())
details.append(f"{self.token_label(token_id)} rank={rank} logit={logits[token_id]:.2f}")
print(f"L{layer:02d}: {'; '.join(details)}")
def answer_trace(self) -> None:
assert self.sample is not None and self.full_ids is not None
full_list = self.full_ids.tolist()
candidates = [self.sample.label, " " + self.sample.label]
match = None
matched_ids = None
for candidate in candidates:
ids = self._token_ids(candidate)
found = find_last_subsequence(full_list[self.prompt_length:], ids)
if found is not None:
match = self.prompt_length + found
matched_ids = ids
break
if match is None or matched_ids is None:
print(f"Gold answer {self.sample.label!r} was not found as a token sequence in the response")
return
activation_position = match - 1
print(
f"gold answer starts at token position {match}; inspecting position "
f"{activation_position}, which predicts its first token"
)
self.trace(self.tokenizer.decode([matched_ids[0]]), activation_position)
def repl(self) -> None:
print("Type 'help' for commands.")
while True:
try:
raw = input("jlens> ").strip()
if not raw:
continue
parts = shlex.split(raw)
command, values = parts[0].lower(), parts[1:]
if command in {"quit", "exit", "q"}:
return
if command == "help":
print(
"tokens [start] [end]\ninspect POSITION [TOP_K]\n"
"compare POSITION LAYER [TOP_K]\ntrace TEXT [POSITION]\n"
"save START END [LAYER|all] [TOP_K] [FILE]\n"
"problem\noutput\nresult\ninfo\nanswer\n"
"sample INDEX\nnext\nquit"
)
elif command == "tokens":
self.print_tokens(*(int(value) for value in values))
elif command == "inspect":
self.inspect(int(values[0]), int(values[1]) if len(values) > 1 else 5)
elif command == "compare":
self.compare(int(values[0]), int(values[1]), int(values[2]) if len(values) > 2 else 10)
elif command == "trace":
self.trace(values[0], int(values[1]) if len(values) > 1 else None)
elif command == "save":
if len(values) < 2:
raise ValueError(
"usage: save START END [LAYER|all] [TOP_K] [FILE]"
)
chosen_layer = None
if len(values) > 2 and values[2].lower() != "all":
chosen_layer = int(values[2])
chosen_top_k = int(values[3]) if len(values) > 3 else 5
chosen_file = values[4] if len(values) > 4 else None
self.save_interval(
int(values[0]), int(values[1]), chosen_layer,
chosen_top_k, chosen_file,
)
elif command == "answer":
self.answer_trace()
elif command == "problem":
self.print_problem()
elif command == "output":
self.print_output()
elif command == "result":
self.print_result()
elif command == "sample":
self.load_sample(int(values[0]))
elif command == "next":
self.load_sample(self.args.sample_index + 1)
elif command == "info":
self.print_info()
else:
print(f"unknown command: {command}")
except (IndexError, ValueError) as exc:
print(f"error: {exc}")
except KeyboardInterrupt:
print("\nUse 'quit' to exit.")
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", required=True)
parser.add_argument("--data", required=True)
parser.add_argument("--lens", required=True)
parser.add_argument("--fit-count", type=int, required=True)
parser.add_argument("--seed", type=int, default=17)
parser.add_argument("--sample-index", type=int, default=0)
parser.add_argument("--device", default="cuda:0")
parser.add_argument("--max-new-tokens", type=int, default=2048)
parser.add_argument("--cache-dir", default="outputs/explorer-cache")
parser.add_argument("--export-dir", default="outputs/jlens-exports")
parser.add_argument("--regenerate", action="store_true")
return parser
def main() -> None:
Explorer(build_parser().parse_args()).repl()
if __name__ == "__main__":
main()