tiny-hinglish-turn-detector / tests /test_train_warmstart.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
1.79 kB
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
try:
import torch
except ImportError: # pragma: no cover
torch = None
@unittest.skipUnless(torch is not None, "PyTorch is not installed")
class WarmStartTest(unittest.TestCase):
def test_weights_only_warm_start_is_strict_and_described(self) -> None:
from scripts.train import _warm_start_model
from turn_detection.models.tiny_tcn import TinyTCNConfig, TinyTurnDetector
config = TinyTCNConfig(
channels=16,
num_blocks=1,
kernel_size=3,
dilation_cycle=(1,),
attention_channels=8,
head_hidden=8,
dropout=0.0,
)
torch.manual_seed(1)
source = TinyTurnDetector(config)
torch.manual_seed(2)
target = TinyTurnDetector(config)
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "best.pt"
torch.save(
{
"model_config": source.model_config(),
"model_state": source.state_dict(),
"epoch": 3,
"metadata": {"run_name": "source"},
},
path,
)
evidence = _warm_start_model(target, path, torch)
for expected, actual in zip(source.parameters(), target.parameters(), strict=True):
torch.testing.assert_close(actual, expected)
self.assertEqual(evidence["mode"], "weights_only_fresh_optimizer")
self.assertEqual(evidence["selected_epoch"], 3)
self.assertEqual(evidence["source_run"], "source")
self.assertEqual(len(evidence["sha256"]), 64)
if __name__ == "__main__":
unittest.main()