From 29e4a0e3d029d2ac54587d4adaec71dbab09c0e9 Mon Sep 17 00:00:00 2001 From: Daoyuan Li <94409450+DaoyuanLi2816@users.noreply.github.com> Date: Wed, 8 Jul 2026 22:34:33 -0700 Subject: [PATCH] Add Sequential Temperature Scaling for confidence-head calibration The DSpark paper (Section 3.2.1) calibrates the confidence head with Sequential Temperature Scaling (STS): for each block position, left to right, a 1D grid search picks the temperature that minimizes the ECE of the cumulative prefix-survival probability, keeping earlier positions fixed. The evaluator already measures exactly that quantity (per-position cumulative-product ECE / AUROC / reliability diagrams), but the calibration itself was missing, and raw confidence scores are noticeably overconfident. This adds the full offline STS workflow: - eval.py --confidence-dump-dir writes per-proposal confidence records ((logits, prefix-accept labels) per dataset) for offline fitting. - scripts/fit_confidence_calibration.py fits per-position temperatures from dumped records via the paper's left-to-right grid search and reports cumulative ECE before/after at each position. - eval.py --confidence-calibration-path applies the fitted temperatures to the confidence logits during evaluation, so both the reported calibration metrics and the --confidence-threshold early stop consume calibrated probabilities. With deepseek-ai/dspark_qwen3_4b_block7 on Qwen/Qwen3-4B, fitting on gsm8k and evaluating held-out (temperature 1.0, threshold 0, identical seeds so generation is unchanged): alpaca ece_mean 0.0893 -> 0.0315 and mt-bench 0.0697 -> 0.0224, with AUROC unchanged (0.8530 -> 0.8516 and 0.8693 -> 0.8686), confirming the calibration is order-preserving. Fitted temperatures are all > 1, matching the paper's overconfidence observation. --- deepspec/eval/dspark/confidence_head.py | 45 +++++ deepspec/eval/dspark/draft_ops.py | 7 + deepspec/eval/dspark/evaluator.py | 31 +++ eval.py | 21 ++ scripts/fit_confidence_calibration.py | 258 ++++++++++++++++++++++++ 5 files changed, 362 insertions(+) create mode 100644 scripts/fit_confidence_calibration.py diff --git a/deepspec/eval/dspark/confidence_head.py b/deepspec/eval/dspark/confidence_head.py index 36b1c32c..0cc25310 100644 --- a/deepspec/eval/dspark/confidence_head.py +++ b/deepspec/eval/dspark/confidence_head.py @@ -322,6 +322,7 @@ def __init__( tensorboard_dir: str | None, step: int | None, artifact_root: Path | None, + records_dir: Path | None = None, ): self.device = device self.max_proposal_tokens = int(max_proposal_tokens) @@ -331,7 +332,9 @@ def __init__( self.tensorboard_dir = tensorboard_dir self.step = step self.artifact_root = artifact_root + self.records_dir = records_dir self.dataset_metrics: PerPositionConfidenceMetrics | None = None + self.records: list[dict] = [] self.rows: list[dict] = [] def start(self) -> None: @@ -341,6 +344,7 @@ def start(self) -> None: num_fine_bins=self.num_fine_bins, device=self.device, ) + self.records = [] def observe( self, @@ -373,6 +377,39 @@ def observe( probs=cumprod_pred, targets=prefix_label, ) + if self.records_dir is not None: + self.records.append( + { + "logits": confidence_logits[0, :effective_length] + .detach() + .to(torch.float32) + .cpu() + .tolist(), + "labels": prefix_label.long().cpu().tolist(), + } + ) + + def _gather_records(self) -> list[dict]: + records = self.records + self.records = [] + if self.records_dir is None: + return [] + gathered: list[list[dict] | None] = [None] * dist.get_world_size() + dist.all_gather_object(gathered, records) + merged: list[dict] = [] + for rank_records in gathered: + merged.extend(rank_records or []) + return merged + + def _write_records(self, *, dataset_name: str, records: list[dict]) -> Path: + assert self.records_dir is not None + dataset_dir = self.records_dir / dataset_name + dataset_dir.mkdir(parents=True, exist_ok=True) + records_path = dataset_dir / "confidence_records.jsonl" + with records_path.open("w", encoding="utf-8") as handle: + for record in records: + handle.write(json.dumps(record, ensure_ascii=False) + "\n") + return records_path def finish( self, @@ -384,10 +421,18 @@ def finish( dataset_metrics = self.dataset_metrics dataset_metrics.all_reduce() self.dataset_metrics = None + records = self._gather_records() if dist.get_rank() != 0 or int(metric_summary["sample_count"]) == 0: return None + if self.records_dir is not None and records: + records_path = self._write_records( + dataset_name=dataset_name, + records=records, + ) + print(f"Wrote confidence records to {records_path}", flush=True) + row = self.build_dataset_row( dataset_name=dataset_name, metric_summary=metric_summary, diff --git a/deepspec/eval/dspark/draft_ops.py b/deepspec/eval/dspark/draft_ops.py index 753a56d9..fcd98982 100644 --- a/deepspec/eval/dspark/draft_ops.py +++ b/deepspec/eval/dspark/draft_ops.py @@ -101,6 +101,7 @@ def build_dspark_proposal( block_size: int, temperature: float, confidence_threshold: float, + confidence_temperatures: torch.Tensor | None = None, ) -> DSparkDraftProposal: assert draft_input_ids.size(0) == 1, "build_dspark_proposal requires batch_size=1" proposal_hidden_states = block_hidden[:, :block_size, :] @@ -124,6 +125,12 @@ def build_dspark_proposal( ) if confidence_logits is None: return _empty_dspark_proposal(draft_input_ids) + if confidence_temperatures is not None: + # Sequential Temperature Scaling: divide each position's logit by + # its fitted temperature so downstream consumers (the threshold + # early-stop and the calibration metrics) see calibrated + # probabilities. + confidence_logits = confidence_logits / confidence_temperatures proposal_draft_tokens = _confident_prefix_length( confidence_logits, block_size=block_size, diff --git a/deepspec/eval/dspark/evaluator.py b/deepspec/eval/dspark/evaluator.py index eba2b34d..12ec162f 100644 --- a/deepspec/eval/dspark/evaluator.py +++ b/deepspec/eval/dspark/evaluator.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from pathlib import Path from types import SimpleNamespace @@ -35,12 +36,37 @@ class Qwen3DSparkEvaluator(BaseEvaluator): def __init__(self, local_rank: int, args): super().__init__(local_rank, args) + self.confidence_temperatures = self._load_confidence_temperatures() self.confidence_head_recorder = self._build_confidence_head_recorder() @property def max_proposal_tokens(self) -> int: return int(self.draft_model.block_size) + def _load_confidence_temperatures(self) -> torch.Tensor | None: + path = getattr(self.args, "confidence_calibration_path", None) + if path is None: + return None + assert self.draft_model.confidence_head is not None, ( + "--confidence-calibration-path requires a draft model with a " + "confidence head." + ) + with Path(path).open("r", encoding="utf-8") as handle: + payload = json.load(handle) + temperatures = [float(value) for value in payload["temperatures"]] + assert len(temperatures) == self.max_proposal_tokens, ( + f"Calibration file {path} has {len(temperatures)} temperatures, " + f"expected block_size={self.max_proposal_tokens}." + ) + assert all(value > 0.0 for value in temperatures), ( + f"Calibration temperatures must be positive, got {temperatures}." + ) + return torch.tensor( + temperatures, + dtype=torch.float32, + device=self.device, + ) + def _build_confidence_head_recorder(self) -> ConfidenceHeadRecorder | None: if self.draft_model.confidence_head is None: return None @@ -54,6 +80,9 @@ def _build_confidence_head_recorder(self) -> ConfidenceHeadRecorder | None: / "artifacts" / f"step_{self.args.step}" ) + records_dir = None + if getattr(self.args, "confidence_dump_dir", None) is not None: + records_dir = Path(self.args.confidence_dump_dir) return ConfidenceHeadRecorder( device=self.device, max_proposal_tokens=self.max_proposal_tokens, @@ -63,6 +92,7 @@ def _build_confidence_head_recorder(self) -> ConfidenceHeadRecorder | None: tensorboard_dir=self.args.tensorboard_dir, step=self.args.step, artifact_root=artifact_root, + records_dir=records_dir, ) def build_models(self) -> tuple[object, Qwen3DSparkModel, AutoTokenizer]: @@ -129,6 +159,7 @@ def _propose( block_size=self.max_proposal_tokens, temperature=float(self.args.temperature), confidence_threshold=float(self.args.confidence_threshold), + confidence_temperatures=self.confidence_temperatures, ) def _update( diff --git a/eval.py b/eval.py index a35e7ea3..85bebba8 100644 --- a/eval.py +++ b/eval.py @@ -39,6 +39,27 @@ def parse_args(): default=0.0, help=("Confidence-head early-stop threshold. Confidence calibration metrics are collected only when this is 0.0."), ) + parser.add_argument( + "--confidence-dump-dir", + type=str, + default=None, + help=( + "Dump per-proposal confidence records to " + "//confidence_records.jsonl for offline calibration " + "fitting (see scripts/fit_confidence_calibration.py). Requires " + "--confidence-threshold 0." + ), + ) + parser.add_argument( + "--confidence-calibration-path", + type=str, + default=None, + help=( + "JSON file with per-position confidence temperatures fitted by " + "scripts/fit_confidence_calibration.py; the temperatures are " + "applied to the confidence logits during evaluation." + ), + ) parser.add_argument("--tensorboard-dir", type=str, default=None) parser.add_argument("--step", type=int, default=None,help=("step for tensorboard logging"),) parser.add_argument("--seed", type=int, default=980406) diff --git a/scripts/fit_confidence_calibration.py b/scripts/fit_confidence_calibration.py new file mode 100644 index 00000000..7c754123 --- /dev/null +++ b/scripts/fit_confidence_calibration.py @@ -0,0 +1,258 @@ +"""Fit Sequential Temperature Scaling (STS) for the DSpark confidence head. + +Implements the post-hoc calibration described in the DSpark paper +(Section 3.2.1): for each block position k (left to right), a 1D grid +search selects a temperature scalar that minimizes the Expected +Calibration Error (ECE) of the cumulative prefix-survival probability + + a_k = prod_{i <= k} sigmoid(logit_i / T_i), + +keeping the already-fitted temperatures of all preceding positions +fixed. Temperature scaling is order-preserving, so it rectifies the +absolute probability magnitudes without changing the confidence head's +token ranking. + +Input records are produced by evaluation runs with +``eval.py --confidence-dump-dir `` (run them with +``--confidence-threshold 0`` and without ``--confidence-calibration-path`` +so the dumped logits are uncalibrated and blocks are untruncated). Fit on +a held-out split, then pass the resulting JSON to +``eval.py --confidence-calibration-path``. + +Example: + python scripts/fit_confidence_calibration.py \ + --records confidence_records/gsm8k/confidence_records.jsonl \ + --output confidence_calibration.json +""" + +from __future__ import annotations + +import argparse +import json +import math +from pathlib import Path + +import torch + + +EPS_PROB = 1e-8 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=( + "Fit per-position confidence temperatures (Sequential " + "Temperature Scaling) from dumped confidence records." + ), + ) + parser.add_argument( + "--records", + nargs="+", + required=True, + help="One or more confidence_records.jsonl files to fit on.", + ) + parser.add_argument( + "--output", + required=True, + help="Output JSON path for the fitted calibration.", + ) + parser.add_argument( + "--num-bins", + type=int, + default=20, + help="Equal-width ECE bins; matches the evaluator's coarse bins.", + ) + parser.add_argument( + "--grid-min", + type=float, + default=0.05, + help="Smallest temperature in the search grid.", + ) + parser.add_argument( + "--grid-max", + type=float, + default=20.0, + help="Largest temperature in the search grid.", + ) + parser.add_argument( + "--grid-size", + type=int, + default=256, + help="Number of log-spaced grid points between grid-min and grid-max.", + ) + return parser.parse_args() + + +def load_records(paths: list[str]) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Load ragged (logits, labels) rows into padded tensors plus a valid mask.""" + logits_rows: list[list[float]] = [] + label_rows: list[list[int]] = [] + for path in paths: + with Path(path).open("r", encoding="utf-8") as handle: + for line_number, line in enumerate(handle, start=1): + line = line.strip() + if not line: + continue + record = json.loads(line) + logits = [float(value) for value in record["logits"]] + labels = [int(value) for value in record["labels"]] + assert len(logits) == len(labels) and len(logits) > 0, ( + f"{path}:{line_number} has mismatched or empty " + f"logits/labels ({len(logits)} vs {len(labels)})." + ) + logits_rows.append(logits) + label_rows.append(labels) + assert logits_rows, f"No records found in: {', '.join(paths)}" + + block_size = max(len(row) for row in logits_rows) + num_records = len(logits_rows) + logits = torch.zeros((num_records, block_size), dtype=torch.float64) + labels = torch.zeros((num_records, block_size), dtype=torch.float64) + valid = torch.zeros((num_records, block_size), dtype=torch.bool) + for row_idx, (logit_row, label_row) in enumerate(zip(logits_rows, label_rows)): + length = len(logit_row) + logits[row_idx, :length] = torch.tensor(logit_row, dtype=torch.float64) + labels[row_idx, :length] = torch.tensor(label_row, dtype=torch.float64) + valid[row_idx, :length] = True + return logits, labels, valid + + +def expected_calibration_error( + probs: torch.Tensor, + targets: torch.Tensor, + *, + num_bins: int, +) -> float: + """Equal-width-bin ECE; mirrors PerPositionConfidenceMetrics.""" + probs = probs.clamp(EPS_PROB, 1.0 - EPS_PROB) + bin_idx = (probs * num_bins).long().clamp_(0, num_bins - 1) + ones = torch.ones_like(probs) + count = torch.zeros(num_bins, dtype=torch.float64).scatter_add_(0, bin_idx, ones) + pred = torch.zeros(num_bins, dtype=torch.float64).scatter_add_(0, bin_idx, probs) + target = torch.zeros(num_bins, dtype=torch.float64).scatter_add_( + 0, bin_idx, targets + ) + total = float(count.sum().item()) + if total <= 0.0: + return float("nan") + denom = count.clamp_min(1e-12) + bin_err = (pred / denom - target / denom).abs() + return float((bin_err * count).sum().item() / total) + + +def fit_sequential_temperatures( + *, + logits: torch.Tensor, + labels: torch.Tensor, + valid: torch.Tensor, + grid: torch.Tensor, + num_bins: int, +) -> tuple[list[float], list[float], list[float]]: + block_size = logits.shape[1] + temperatures: list[float] = [] + ece_before: list[float] = [] + ece_after: list[float] = [] + calibrated_cumprod = torch.ones(logits.shape[0], dtype=torch.float64) + raw_cumprod = torch.ones(logits.shape[0], dtype=torch.float64) + + for position in range(block_size): + position_valid = valid[:, position] + num_valid = int(position_valid.sum().item()) + if num_valid == 0: + temperatures.append(1.0) + ece_before.append(float("nan")) + ece_after.append(float("nan")) + continue + + position_logits = logits[position_valid, position] + position_labels = labels[position_valid, position] + prefix = calibrated_cumprod[position_valid] + + raw_probs = raw_cumprod[position_valid] * torch.sigmoid(position_logits) + ece_before.append( + expected_calibration_error( + raw_probs, + position_labels, + num_bins=num_bins, + ) + ) + + best_temperature = 1.0 + best_ece = float("inf") + for temperature in grid.tolist(): + probs = prefix * torch.sigmoid(position_logits / temperature) + ece = expected_calibration_error( + probs, + position_labels, + num_bins=num_bins, + ) + if ece < best_ece: + best_ece = ece + best_temperature = float(temperature) + temperatures.append(best_temperature) + ece_after.append(best_ece) + + calibrated_step = torch.sigmoid(logits[:, position] / best_temperature) + raw_step = torch.sigmoid(logits[:, position]) + calibrated_cumprod = torch.where( + position_valid, + calibrated_cumprod * calibrated_step, + calibrated_cumprod, + ) + raw_cumprod = torch.where( + position_valid, + raw_cumprod * raw_step, + raw_cumprod, + ) + print( + f"position {position}: n={num_valid} T={best_temperature:.4f} " + f"cumulative ECE {ece_before[-1]:.4f} -> {best_ece:.4f}", + flush=True, + ) + return temperatures, ece_before, ece_after + + +def main() -> None: + args = parse_args() + logits, labels, valid = load_records(args.records) + + grid = torch.logspace( + math.log10(args.grid_min), + math.log10(args.grid_max), + steps=int(args.grid_size), + dtype=torch.float64, + ) + # Always include the identity temperature so calibration can no-op. + grid = torch.unique(torch.cat([grid, torch.ones(1, dtype=torch.float64)])) + + temperatures, ece_before, ece_after = fit_sequential_temperatures( + logits=logits, + labels=labels, + valid=valid, + grid=grid, + num_bins=int(args.num_bins), + ) + + payload = { + "temperatures": temperatures, + "block_size": logits.shape[1], + "num_bins": int(args.num_bins), + "grid": { + "min": float(args.grid_min), + "max": float(args.grid_max), + "size": int(args.grid_size), + }, + "num_records": logits.shape[0], + "sources": [str(path) for path in args.records], + "cumulative_ece_before": ece_before, + "cumulative_ece_after": ece_after, + } + output_path = Path(args.output) + output_path.parent.mkdir(parents=True, exist_ok=True) + with output_path.open("w", encoding="utf-8") as handle: + json.dump(payload, handle, ensure_ascii=False, indent=2) + print(f"Wrote calibration to {output_path}", flush=True) + + +if __name__ == "__main__": + main()