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 " + "