Download loopq_quantization/scripts/loopq/ouro_calibration.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 25.7 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/ouro_calibration.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/ouro_calibration.py
-
curl -L -o ouro_calibration.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/ouro_calibration.py
25.7 kB
| """Real-data Ouro calibration primitives for the LQ7 driver.""" | |
| from __future__ import annotations | |
| from dataclasses import asdict, dataclass | |
| from pathlib import Path | |
| from collections import defaultdict | |
| from collections.abc import Mapping | |
| from typing import Any | |
| import torch | |
| import torch.nn.functional as F | |
| from torch import nn | |
| from adapters.ouro import PAPER_GROUPS | |
| from .cta import CrossLoopTransitionAdapter | |
| from .las import LoopAwareActivationScales | |
| from .objective import AdaptiveMuCache, TrajectoryLoss, trajectory_aware_loss | |
| from .quantization import quantize_weight | |
| from .sharing_gap import SharingGapStatistics | |
| from .transforms import FlatQuantSVDKroneckerTransform, SharedKroneckerTransform | |
| from loopq.paths import pinned_snapshot | |
| PINNED_OURO_SNAPSHOT = pinned_snapshot("ouro") | |
| PILE_DATASET = "mit-han-lab/pile-val-backup" | |
| PILE_DATASET_REVISION = "2f5e46ae6a69cf0dce4b12f78241c408936ca0e4" | |
| class OuroCalibrationDataConfig: | |
| samples: int = 1024 | |
| max_length: int = 256 | |
| split: str = "validation" | |
| text_field: str = "text" | |
| smoke_small: bool = False | |
| def paper(cls) -> "OuroCalibrationDataConfig": | |
| return cls() | |
| def smoke(cls, samples: int = 2) -> "OuroCalibrationDataConfig": | |
| if not 1 <= samples <= 2: | |
| raise ValueError("--smoke-small permits only 1 or 2 samples") | |
| # A smoke validates the real-model forward/backward/export path, not | |
| # sequence-length scaling. Keeping 256 tokens here produces roughly | |
| # 200 GiB of exact saved tensors and turns a pipeline check into an | |
| # hour-scale calibration. Paper runs retain the 256-token contract. | |
| return cls(samples=samples, max_length=16, smoke_small=True) | |
| def metadata(self) -> dict[str, Any]: | |
| value = asdict(self) | |
| value.update({ | |
| "dataset": PILE_DATASET, | |
| "dataset_revision": PILE_DATASET_REVISION, | |
| "streaming": True, | |
| "paper_calibration": not self.smoke_small and self.samples == 1024 and self.max_length == 256, | |
| "warning": ( | |
| None if not self.smoke_small else | |
| "Pipeline validation only; not a paper calibration result." | |
| ), | |
| }) | |
| return value | |
| class OuroLASStatisticsCollector: | |
| """Online per-module/loop/group absmax collector with bounded memory.""" | |
| def __init__(self, *, bits: int, loop_count: int = 4, group_size: int = 32) -> None: | |
| if bits not in (4, 8): | |
| raise ValueError("activation bits must be 4 or 8") | |
| self.bits = bits | |
| self.loop_count = loop_count | |
| self.group_size = group_size | |
| self._absmax: dict[str, torch.Tensor] = {} | |
| self._calls: dict[str, int] = {} | |
| self._sample_calls: dict[str, int] | None = None | |
| def begin_sample(self) -> None: | |
| if self._sample_calls is not None: | |
| raise RuntimeError("LAS sample already active") | |
| self._sample_calls = {} | |
| def end_sample(self, *, expected_modules: set[str]) -> None: | |
| calls, self._sample_calls = self._sample_calls, None | |
| if calls is None or set(calls) != expected_modules or any(n != self.loop_count for n in calls.values()): | |
| raise ValueError("LAS sample must visit every module exactly four times") | |
| def observe(self, module_key: str, value: torch.Tensor) -> int: | |
| if self._sample_calls is not None: | |
| n = self._sample_calls.get(module_key, 0) | |
| if n >= self.loop_count: | |
| raise ValueError("LAS sample loop routing cannot wrap") | |
| self._sample_calls[module_key] = n + 1 | |
| loop = self._calls.get(module_key, 0) % self.loop_count | |
| self._calls[module_key] = self._calls.get(module_key, 0) + 1 | |
| width = value.shape[-1] | |
| groups = (width + self.group_size - 1) // self.group_size | |
| if module_key not in self._absmax: | |
| self._absmax[module_key] = torch.zeros(self.loop_count, groups) | |
| elif self._absmax[module_key].shape[1] != groups: | |
| raise ValueError("module feature width changed during LAS collection") | |
| work = value.detach().abs() | |
| work = F.pad(work, (0, groups * self.group_size - width)) | |
| maximum = work.reshape(-1, groups, self.group_size).amax(dim=(0, 2)).cpu() | |
| self._absmax[module_key][loop] = torch.maximum(self._absmax[module_key][loop], maximum) | |
| return loop | |
| def build_las(self) -> LoopAwareActivationScales: | |
| if self._sample_calls is not None: | |
| raise ValueError("LAS sample still active") | |
| if not self._absmax: | |
| raise ValueError("no activation statistics were collected") | |
| expected_calls = {name: count for name, count in self._calls.items() if count % self.loop_count} | |
| if expected_calls: | |
| raise ValueError(f"incomplete recurrent trajectories for modules: {expected_calls}") | |
| # Section 4.1 and Appendix B.3 specify one LAS scalar per module and | |
| # loop (O(TL)), while activation quantization remains group-wise. The | |
| # observed tensors establish routing/shape coverage; runtime group | |
| # absmax values are dynamic and LAS learns their clipping multiplier. | |
| return LoopAwareActivationScales.dynamic_for_modules( | |
| sorted(self._absmax), loop_count=self.loop_count, | |
| group_size=self.group_size, initial_clip=1.0 | |
| ) | |
| def observe_activation_pre_hook( | |
| collector: OuroLASStatisticsCollector, | |
| module_key: str, | |
| inputs: tuple[torch.Tensor, ...], | |
| ) -> None: | |
| """Collect statistics while preserving the module's original inputs. | |
| PyTorch treats a non-``None`` forward-pre-hook return value as replacement | |
| inputs. ``observe`` intentionally returns the recurrence index for direct | |
| callers, so registering it as the callback would replace a tensor with an | |
| integer. | |
| """ | |
| collector.observe(module_key, inputs[0]) | |
| return None | |
| def load_pinned_teacher_student(device: str): | |
| """Load two independent BF16 Ouro models from the pinned local snapshot.""" | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| common = dict( | |
| pretrained_model_name_or_path=str(PINNED_OURO_SNAPSHOT), | |
| trust_remote_code=True, | |
| local_files_only=True, | |
| torch_dtype=torch.bfloat16, | |
| ) | |
| teacher = AutoModelForCausalLM.from_pretrained(**common).to(device).eval() | |
| student = AutoModelForCausalLM.from_pretrained(**common).to(device).eval() | |
| for parameter in teacher.parameters(): | |
| parameter.requires_grad_(False) | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| PINNED_OURO_SNAPSHOT, trust_remote_code=True, local_files_only=True | |
| ) | |
| return teacher, student, tokenizer | |
| def load_pile_texts(config: OuroCalibrationDataConfig) -> list[str]: | |
| from datasets import load_dataset | |
| dataset = load_dataset( | |
| PILE_DATASET, | |
| split=config.split, | |
| revision=PILE_DATASET_REVISION, | |
| streaming=True, | |
| ) | |
| texts = [] | |
| for row in dataset: | |
| text = row[config.text_field] | |
| if isinstance(text, str) and text.strip(): | |
| texts.append(text) | |
| if len(texts) == config.samples: | |
| break | |
| if len(texts) != config.samples: | |
| raise ValueError(f"requested {config.samples} texts but found {len(texts)}") | |
| return texts | |
| def _ste(original: torch.Tensor, quantized: torch.Tensor) -> torch.Tensor: | |
| """Quantized forward plus an identity activation-gradient path. | |
| Unlike ``original + (quantized-original).detach()``, this form deliberately | |
| retains gradients to learned LAS scales and transforms through the QDQ | |
| expression while adding the standard identity STE for rounding. | |
| """ | |
| return quantized + (original - original.detach()) | |
| def _weight_ste(original: torch.Tensor, quantized: torch.Tensor) -> torch.Tensor: | |
| """RTN weight forward with the standard identity transform gradient. | |
| Weight scales are inferred rather than learned LoopQ parameters. Detaching | |
| their QDQ graph avoids retaining full-matrix rounding intermediates while | |
| the identity path still optimizes every transform affecting ``original``. | |
| """ | |
| return original + (quantized - original).detach() | |
| def qdq_linear(value, module, transform, las, key, loop, bits, *, checkpoint=False, activation_ste="identity"): | |
| """Optionally recompute a deterministic linear during backward. | |
| Capture the exact loop transform now: routing counters and module hooks | |
| must never be replayed, especially after statistics-mode cleanup. | |
| """ | |
| def compute(x): | |
| transformed = transform(x) | |
| activation_q = las.quantize(key, loop, transformed, bits=bits, | |
| rounding_ste=activation_ste == "rounding").dequantized | |
| if activation_ste == "identity": | |
| activation_q = _ste(transformed, activation_q) | |
| folded = transform.fold_weight(module.weight) | |
| weight_q = quantize_weight(folded.detach()).dequantized | |
| return F.linear(activation_q, _weight_ste(folded, weight_q), module.bias) | |
| if checkpoint: | |
| from torch.utils.checkpoint import checkpoint as recompute | |
| return recompute(compute, value, use_reentrant=False, preserve_rng_state=False) | |
| return compute(value) | |
| def ouro_projection_sites(model: nn.Module) -> dict[str, tuple[nn.Module, ...]]: | |
| """Resolve all 24x4 paper groups without depending on concrete HF classes.""" | |
| sites: dict[str, tuple[nn.Module, ...]] = {} | |
| for layer_index, layer in enumerate(model.model.layers): | |
| for group, spec in PAPER_GROUPS.items(): | |
| modules = [] | |
| for path in spec["hf_weights"]: | |
| current = layer | |
| for part in path.split("."): | |
| current = getattr(current, part) | |
| modules.append(current) | |
| sites[f"model.layers.{layer_index}.{group}"] = tuple(modules) | |
| return sites | |
| class OuroTrajectoryCapture: | |
| """Capture true recurrent pre-CTA states and feed CTA into the next loop.""" | |
| def __init__(self, norm: nn.Module, cta: CrossLoopTransitionAdapter | None) -> None: | |
| self.norm = norm | |
| self.cta = cta | |
| self.pre_cta: list[torch.Tensor] = [] | |
| self.adapted: list[torch.Tensor] = [] | |
| self._handle = None | |
| def __enter__(self) -> "OuroTrajectoryCapture": | |
| self.pre_cta.clear() | |
| self.adapted.clear() | |
| def hook(_module, _inputs, output): | |
| loop = len(self.pre_cta) | |
| if loop >= 4: | |
| raise RuntimeError("Ouro norm was invoked more than four recurrent loops") | |
| self.pre_cta.append(output) | |
| if loop < 3 and self.cta is not None: | |
| output = self.cta(output, loop) | |
| self.adapted.append(output) | |
| return output | |
| self._handle = self.norm.register_forward_hook(hook) | |
| return self | |
| def __exit__(self, exception_type, _exception, _traceback) -> None: | |
| self._handle.remove() | |
| self._handle = None | |
| # Do not hide the primary forward/backward failure with a secondary | |
| # trajectory-length assertion during cleanup. | |
| if exception_type is None and len(self.pre_cta) != 4: | |
| raise RuntimeError(f"expected four recurrent states, captured {len(self.pre_cta)}") | |
| class OuroDifferentiableQDQ(nn.Module): | |
| """Differentiable HF calibration hooks for every Ouro paper group. | |
| The backbone weights remain owned by the model and are never mutated. A | |
| selected group gets four independent transform copies; all other groups | |
| route through one shared transform. | |
| """ | |
| def __init__( | |
| self, | |
| *, | |
| model: nn.Module, | |
| las: LoopAwareActivationScales, | |
| activation_bits: int, | |
| factor_by_width: Mapping[int, tuple[int, int]], | |
| checkpoint_linears: bool = False, | |
| activation_ste: str = "identity", | |
| statistics_coordinate: str = "svd_parameters", | |
| ) -> None: | |
| super().__init__() | |
| self.model = model | |
| self.las = las | |
| self.activation_bits = activation_bits | |
| if activation_ste not in {"identity", "rounding"}: | |
| raise ValueError("unknown activation STE") | |
| self.activation_ste = activation_ste | |
| if statistics_coordinate not in {"svd_parameters", "effective_factors"}: | |
| raise ValueError("statistics coordinate must be svd_parameters or effective_factors") | |
| self.statistics_coordinate = statistics_coordinate | |
| self.checkpoint_linears = checkpoint_linears | |
| self.sites = ouro_projection_sites(model) | |
| self._encoded = {key: key.replace(".", "__") for key in self.sites} | |
| transforms = {} | |
| for key, modules in self.sites.items(): | |
| width = modules[0].in_features | |
| factors = factor_by_width.get(width) | |
| if factors is None or factors[0] * factors[1] != width: | |
| raise ValueError(f"missing valid factors for feature width {width}") | |
| transforms[self._encoded[key]] = FlatQuantSVDKroneckerTransform(*factors) | |
| self.shared_transforms = nn.ModuleDict(transforms) | |
| self.selected_loop_transforms = nn.ModuleDict() | |
| self.statistics_loop_transforms = nn.ModuleDict() | |
| self._statistics_mode = False | |
| self._handles: list[Any] = [] | |
| self._calls: dict[int, int] = defaultdict(int) | |
| self._records: dict[tuple[str, int], list[torch.Tensor]] = defaultdict(list) | |
| self._execution_views: dict[int, object] = {} | |
| for parameter in model.parameters(): | |
| parameter.requires_grad_(False) | |
| def select_group(self, key: str) -> None: | |
| if key not in self.sites: | |
| raise KeyError(key) | |
| encoded = self._encoded[key] | |
| if encoded in self.selected_loop_transforms: | |
| return | |
| base = self.shared_transforms[encoded] | |
| loops = nn.ModuleList([base.fresh_copy() for _ in range(4)]) | |
| self.selected_loop_transforms[encoded] = loops | |
| def _base_transform_for(self, key: str, loop: int) -> SharedKroneckerTransform: | |
| encoded = self._encoded[key] | |
| if encoded in self.selected_loop_transforms: | |
| return self.selected_loop_transforms[encoded][loop] | |
| return self.shared_transforms[encoded] | |
| def transform_for(self, key: str, loop: int) -> SharedKroneckerTransform: | |
| encoded = self._encoded[key] | |
| if self._statistics_mode: | |
| return self.statistics_loop_transforms[encoded][loop] | |
| return self._base_transform_for(key, loop) | |
| def _build_statistics_transforms(self) -> None: | |
| transforms = {} | |
| for key in sorted(self.sites): | |
| copies = [] | |
| for loop in range(4): | |
| base = self._base_transform_for(key, loop) | |
| copy = (base.fresh_copy() if self.statistics_coordinate == "svd_parameters" | |
| else SharedKroneckerTransform.from_export_state(base.export_state()).to( | |
| device=base.left.device, dtype=base.left.dtype)) | |
| copies.append(copy) | |
| transforms[self._encoded[key]] = nn.ModuleList(copies) | |
| self.statistics_loop_transforms = nn.ModuleDict(transforms) | |
| def begin(self, *, statistics_mode: bool = False) -> None: | |
| if self._handles: | |
| raise RuntimeError("QDQ hooks are already active") | |
| self._statistics_mode = statistics_mode | |
| if statistics_mode: | |
| self._build_statistics_transforms() | |
| self._calls.clear() | |
| self._records.clear() | |
| self._execution_views.clear() | |
| for key, modules in self.sites.items(): | |
| for module in modules: | |
| def pre_hook(current, inputs, key=key): | |
| call = self._calls[id(current)] | |
| self._calls[id(current)] += 1 | |
| if call >= 4: | |
| raise RuntimeError("projection invoked more than four times; loop routing cannot wrap") | |
| loop = call | |
| activation = inputs[0] | |
| transform_module = self.transform_for(key, loop) | |
| transform = self._execution_views.get(id(transform_module)) | |
| if transform is None: | |
| transform = transform_module.materialize() | |
| self._execution_views[id(transform_module)] = transform | |
| current._loopq_override = qdq_linear( | |
| activation, current, transform, self.las, key, loop, | |
| self.activation_bits, checkpoint=self.checkpoint_linears, | |
| activation_ste=self.activation_ste, | |
| ) | |
| # The original linear output is replaced below; its graph | |
| # is unused. Keep its execution free of saved tensors. | |
| return tuple(x.detach() if isinstance(x, torch.Tensor) else x for x in inputs) | |
| def post_hook(current, _inputs, output): | |
| replacement = current._loopq_override | |
| del current._loopq_override | |
| return replacement | |
| self._handles.append(module.register_forward_pre_hook(pre_hook)) | |
| self._handles.append(module.register_forward_hook(post_hook)) | |
| def end(self, *, validate: bool = True) -> None: | |
| for handle in self._handles: | |
| handle.remove() | |
| self._handles.clear() | |
| self._execution_views.clear() | |
| incomplete = [key for key, modules in self.sites.items() | |
| for module in modules if self._calls[id(module)] != 4] | |
| for modules in self.sites.values(): | |
| for module in modules: | |
| if hasattr(module, "_loopq_override"): | |
| del module._loopq_override | |
| self._statistics_mode = False | |
| self.statistics_loop_transforms = nn.ModuleDict() | |
| if validate and incomplete: | |
| raise RuntimeError(f"incomplete projection trajectories: {incomplete[:4]}") | |
| def sharing_gap_statistics(self, loss: torch.Tensor, *, fisher_loss: torch.Tensor | None = None) -> dict[str, SharingGapStatistics]: | |
| """Collect Eq.8 transform VJPs and their diagonal-Fisher estimate.""" | |
| if not self._statistics_mode: | |
| raise RuntimeError("sharing-gap statistics require statistics_mode") | |
| result = {} | |
| transforms = [ | |
| self.transform_for(key, loop) | |
| for key in sorted(self.sites) for loop in range(4) | |
| ] | |
| all_parameters = tuple( | |
| parameter for transform in transforms for parameter in transform.parameters() | |
| ) | |
| all_gradients = torch.autograd.grad( | |
| loss, all_parameters, retain_graph=True, allow_unused=True | |
| ) | |
| gradient_by_id = { | |
| id(parameter): gradient | |
| for parameter, gradient in zip(all_parameters, all_gradients) | |
| } | |
| fisher_gradients = (all_gradients if fisher_loss is None else torch.autograd.grad( | |
| fisher_loss, all_parameters, retain_graph=True, allow_unused=True)) | |
| fisher_by_id = dict(zip(map(id, all_parameters), fisher_gradients)) | |
| for key in sorted(self.sites): | |
| per_loop = [] | |
| fisher_per_loop = [] | |
| for loop in range(4): | |
| transform = self.transform_for(key, loop) | |
| params = tuple(transform.parameters()) | |
| accumulated = [ | |
| (gradient_by_id[id(parameter)].detach() | |
| if gradient_by_id[id(parameter)] is not None | |
| else torch.zeros_like(parameter)) | |
| for parameter in params | |
| ] | |
| per_loop.append(torch.cat([item.flatten() for item in accumulated]).cpu()) | |
| fisher_per_loop.append(torch.cat([ | |
| (fisher_by_id[id(parameter)].detach() if fisher_by_id[id(parameter)] is not None | |
| else torch.zeros_like(parameter)).flatten() for parameter in params | |
| ]).cpu().to(torch.float64)) | |
| gradients = torch.stack(per_loop) | |
| result[key] = SharingGapStatistics( | |
| gradients=gradients, | |
| # Square in FP64: finite FP32 VJPs can exceed sqrt(FP32_MAX). | |
| fisher_diagonal=torch.stack(fisher_per_loop).square().mean(dim=0), | |
| parameter_count=gradients.shape[1], | |
| ) | |
| return result | |
| def export_shared(self) -> dict[str, dict[str, object]]: | |
| return {key: self.shared_transforms[self._encoded[key]].export_state() | |
| for key in sorted(self.sites)} | |
| def export_selected(self) -> dict[str, dict[str, dict[str, object]]]: | |
| return { | |
| key: {str(loop): transform.export_state() for loop, transform in enumerate( | |
| self.selected_loop_transforms[self._encoded[key]])} | |
| for key in sorted(self.sites) | |
| if self._encoded[key] in self.selected_loop_transforms | |
| } | |
| def _extract_logits(output: Any) -> torch.Tensor: | |
| if hasattr(output, "logits"): | |
| return output.logits | |
| if isinstance(output, (tuple, list)): | |
| return output[0] | |
| raise TypeError("model output must expose .logits or place logits first") | |
| def ouro_trajectory_loss( | |
| *, | |
| teacher: nn.Module, | |
| student: nn.Module, | |
| student_qdq: OuroDifferentiableQDQ, | |
| cta: CrossLoopTransitionAdapter, | |
| inputs: Mapping[str, torch.Tensor] | torch.Tensor, | |
| step: int, | |
| mu_cache: AdaptiveMuCache, | |
| collect_statistics: bool = True, | |
| offload_saved_tensors: bool = False, | |
| saved_tensor_keep_bytes: int | None = None, | |
| asynchronous_offload: bool = False, | |
| kl_mode: str = "conditional_topk", | |
| fisher_estimator: str = "trajectory_gradient_square", | |
| ) -> tuple[TrajectoryLoss, dict[str, SharingGapStatistics]]: | |
| """Run one bounded true four-loop teacher/student calibration example.""" | |
| if fisher_estimator not in {"trajectory_gradient_square", "model_score_mc"}: | |
| raise ValueError("unknown Fisher estimator") | |
| call = (lambda model: model(inputs)) if isinstance(inputs, torch.Tensor) else ( | |
| lambda model: model(**inputs, use_cache=False, exit_at_step=3) | |
| ) | |
| with torch.no_grad(), OuroTrajectoryCapture(teacher.model.norm, None) as teacher_trace: | |
| teacher_output = call(teacher) | |
| teacher_hidden = torch.stack(teacher_trace.pre_cta) | |
| student_qdq.begin(statistics_mode=collect_statistics) | |
| failed = True | |
| try: | |
| saved_tensors = ( | |
| # Weight-QDQ internals are detached by _weight_ste, bounding this | |
| # exact saved-tensor set so pinned asynchronous restoration is | |
| # practical. The former differentiable full-weight quantizer graph | |
| # exceeded the pinned-host allocator and is not retained here. | |
| torch.autograd.graph.save_on_cpu(pin_memory=True, device_type="cuda") | |
| if offload_saved_tensors else torch.autograd.graph.saved_tensors_hooks( | |
| lambda tensor: tensor, lambda tensor: tensor | |
| ) | |
| ) | |
| if saved_tensor_keep_bytes is not None: | |
| if not offload_saved_tensors: | |
| raise ValueError("budgeted placement requires offload_saved_tensors") | |
| from .saved_tensors import BudgetedSavedTensors | |
| saved_tensors = BudgetedSavedTensors( | |
| (p for p in student.parameters() if not p.requires_grad), | |
| keep_bytes=saved_tensor_keep_bytes) | |
| if asynchronous_offload: | |
| if not offload_saved_tensors or saved_tensor_keep_bytes is not None: | |
| raise ValueError("asynchronous offload requires CPU placement without retention budget") | |
| from .saved_tensors import AsyncSavedTensors | |
| saved_tensors = AsyncSavedTensors() | |
| with saved_tensors: | |
| with OuroTrajectoryCapture(student.model.norm, cta) as student_trace: | |
| student_output = call(student) | |
| student_hidden = torch.stack(student_trace.pre_cta) | |
| adapted = torch.stack(student_trace.adapted) | |
| mu = mu_cache.get(step, teacher_hidden, student_hidden) | |
| loss = trajectory_aware_loss( | |
| teacher_logits=_extract_logits(teacher_output), | |
| student_logits=_extract_logits(student_output), | |
| teacher_hidden=teacher_hidden, | |
| student_hidden=student_hidden, | |
| adapted_transitions=adapted, | |
| teacher_next_inputs=teacher_hidden[:-1], | |
| mu=mu, | |
| include_transition=cta.enabled, | |
| kl_mode=kl_mode, | |
| ) | |
| fisher_loss = None | |
| if collect_statistics and fisher_estimator == "model_score_mc": | |
| # One Monte Carlo draw of the joint categorical outputs | |
| # at the fixed teacher-forced contexts. This is a Fisher | |
| # estimator, not the square of the trajectory-loss gradient. | |
| logits = _extract_logits(student_output).float() | |
| logp = logits.log_softmax(-1) | |
| labels = torch.multinomial(logp.detach().exp().reshape(-1, logits.shape[-1]), 1) | |
| fisher_loss = -logp.reshape(-1, logits.shape[-1]).gather(-1, labels).sum() | |
| statistics = ( | |
| student_qdq.sharing_gap_statistics(loss.total, fisher_loss=fisher_loss) | |
| if collect_statistics else {} | |
| ) | |
| failed = False | |
| finally: | |
| student_qdq.end(validate=not failed) | |
| return loss, statistics | |