cwLeeDev's picture
Correct shared-state lineage and add composite torch exports
c07793c verified
Raw History Blame
4.47 kB
"""Math Ink 0.6์˜ online/raster ๊ฒฝ๋กœ๋ฅผ torch.export์™€ LiteRT ์นœํ™” ์ถœ๋ ฅ์œผ๋กœ ๊ณ ์ •ํ•œ๋‹ค."""
from __future__ import annotations
from typing import Iterable
import torch
from torch import Tensor, nn
from .math_ink_06 import MathInk06Model, fuse_raster_logits06, virtual_features06
class OnlineExportWrapper06(nn.Module):
"""ํ•„์š” ๋ณ€์ˆ˜: 0.6 ๋ชจ๋ธยทonline adapter. ์ž‘๋™ ์›๋ฆฌ: ์‹ค์ œ composite ๊ฒฝ๋กœ์˜ exact/family logits๋ฅผ ๋ฐ˜ํ™˜ํ•œ๋‹ค."""
def __init__(self, model: MathInk06Model, adapter: nn.Module | None = None) -> None:
super().__init__()
self.model = model
self.adapter = adapter if adapter is not None else nn.Identity()
def forward(self, sequence: Tensor) -> tuple[Tensor, Tensor]:
"""ํ•„์š” ๋ณ€์ˆ˜: Bร—128ร—19 canonical trajectory. ์ž‘๋™ ์›๋ฆฌ: shared encoder์˜ ๋‘ ๋ถ„๋ฅ˜ head๋ฅผ ์ง์ ‘ ์‹คํ–‰ํ•œ๋‹ค."""
return self.model.forward_online(self.adapter(sequence))
class RasterExportWrapper06(nn.Module):
"""ํ•„์š” ๋ณ€์ˆ˜: 0.6 ๋ชจ๋ธยทraster adapterยทfusion ์ƒ์ˆ˜. ์ž‘๋™ ์›๋ฆฌ: top-4๋ฅผ composite trajectory ๊ฒฝ๋กœ๋กœ ๋ถ„๋ฅ˜ํ•œ๋‹ค."""
def __init__(
self, model: MathInk06Model, *, adapter: nn.Module | None = None,
fusion_mode: str, score_weight: float,
) -> None:
super().__init__()
self.model = model
self.adapter = adapter if adapter is not None else nn.Identity()
self.fusion_mode = fusion_mode
self.score_weight = float(score_weight)
def forward(self, raster: Tensor) -> Tensor:
"""ํ•„์š” ๋ณ€์ˆ˜: Bร—1ร—128ร—128 raster. ์ž‘๋™ ์›๋ฆฌ: direct raster-label shortcut ์—†์ด shared trajectory ๋ถ„๋ฅ˜๋ฅผ ๊ฒฐํ•ฉํ•œ๋‹ค."""
coordinates, states, progress, hypothesis_scores = self.model.decode_raster_trajectories(raster)
features = virtual_features06(
coordinates, states,
None if self.model.raster_architecture == "spatial_flat_v1" else progress,
contract=self.model.virtual_contract,
)
batch, hypotheses, steps, channels = features.shape
if self.model.use_virtual_adapter:
raw_features = features
internal = self.model.virtual_adapter(
features.reshape(batch * hypotheses, steps, channels),
).reshape(batch, hypotheses, steps, channels)
features = raw_features + self.model.virtual_adapter_weight * (internal - raw_features)
flat_features = self.adapter(features.reshape(batch * hypotheses, steps, channels))
exact, family = self.model.classify_trajectory(flat_features)
output = {
"hypothesis_scores": hypothesis_scores,
"exact_logits": exact.reshape(batch, hypotheses, -1),
"family_logits": family.reshape(batch, hypotheses, -1),
}
fused, _selected = fuse_raster_logits06(
output, mode=self.fusion_mode, score_weight=self.score_weight,
)
return fused
def exported_equivalence06(
eager: nn.Module, exported: torch.export.ExportedProgram, inputs: Iterable[tuple[Tensor, ...]],
) -> dict[str, float | int | bool]:
"""ํ•„์š” ๋ณ€์ˆ˜: eager/export ๋ชจ๋ธยท๋Œ€ํ‘œ ์ž…๋ ฅ. ์ž‘๋™ ์›๋ฆฌ: ๋ชจ๋“  ์ถœ๋ ฅ tensor์˜ top-1 ์ผ์น˜์™€ ์ตœ๋Œ€ logit ์˜ค์ฐจ๋ฅผ ๊ณ„์‚ฐํ•œ๋‹ค."""
exported_module = exported.module()
samples = top1_matches = 0
max_error = 0.0
eager.eval()
with torch.inference_mode():
for arguments in inputs:
eager_output = eager(*arguments)
export_output = exported_module(*arguments)
eager_values = eager_output if isinstance(eager_output, tuple) else (eager_output,)
export_values = export_output if isinstance(export_output, tuple) else (export_output,)
if len(eager_values) != len(export_values):
raise ValueError("eager/export ์ถœ๋ ฅ ๊ฐœ์ˆ˜๊ฐ€ ๋‹ค๋ฆ…๋‹ˆ๋‹ค.")
for eager_value, export_value in zip(eager_values, export_values):
max_error = max(max_error, float((eager_value - export_value).abs().max()))
samples += int(eager_values[0].shape[0])
top1_matches += int((eager_values[0].argmax(dim=-1) == export_values[0].argmax(dim=-1)).sum())
return {
"samples": samples, "top1_matches": top1_matches,
"top1_agreement": top1_matches / max(samples, 1), "max_absolute_logit_error": max_error,
"gate_passed": top1_matches == samples and max_error <= 0.02,
}