refactor: eliminate test duplication via shared helpers
- Add tests/helpers.py with shared config, dataset, tokenizer, executor, and assertion helpers - Replace 15 copies of device one-liner with session-scoped fixture - Collapse 5 near-identical Dataset subclasses into RandomTokenDataset - Remove duplicate _make_config/_make_model/_make_frozen and FakeTokenizer/FakeExecutor definitions - Make test_callbacks and test_early_stopping use existing train_config_factory - Replace 6 duplicate meta.json read blocks with load_shard_meta - Fix mkdtemp leaks in test_lora.py with TemporaryDirectory
This commit is contained in:
+15
-106
@@ -6,11 +6,10 @@ import tempfile
|
||||
import pytest
|
||||
import torch
|
||||
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
from tests.helpers import TINY_CONFIG, RandomTokenDataset, make_tiny_config
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
@@ -19,6 +18,12 @@ def pytest_configure(config):
|
||||
config.addinivalue_line("markers", "unit: fast unit tests")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def device():
|
||||
"""Session-scoped device string (``"cuda"`` if available, else ``"cpu"``)."""
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
||||
"""Create a simple tokenizer for testing purposes."""
|
||||
tokenizer = Tokenizer(models.BPE())
|
||||
@@ -33,69 +38,6 @@ def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
||||
return auto_tokenizer
|
||||
|
||||
|
||||
class RandomDataset(Dataset):
|
||||
"""Random dataset for testing purposes."""
|
||||
|
||||
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
||||
self.length = length or int(torch.randint(100, 200, (1,)).item())
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {
|
||||
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
}
|
||||
|
||||
|
||||
class MultiTurnDataset(Dataset):
|
||||
"""Multi-turn dataset with loss mask for SFT training tests."""
|
||||
|
||||
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
||||
self.length = length or int(torch.randint(100, 200, (1,)).item())
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
input_ids = torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
target_ids = torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
loss_mask = torch.randint(0, 1, (self.max_length,))
|
||||
|
||||
return {
|
||||
"input_ids": input_ids,
|
||||
"target_ids": target_ids,
|
||||
"loss_mask": loss_mask,
|
||||
}
|
||||
|
||||
|
||||
class EarlyStoppingDataset(Dataset):
|
||||
"""Dataset that triggers early stopping after consuming a specified number of samples."""
|
||||
|
||||
def __init__(self, length=10, stop_after=5):
|
||||
self.length = length
|
||||
self.stop_after = stop_after
|
||||
self.count = 0
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
self.count += 1
|
||||
if self.count == self.stop_after:
|
||||
raise RuntimeError("Simulated early stopping")
|
||||
|
||||
return {
|
||||
"input_ids": torch.randint(0, 1000, (64,)),
|
||||
"target_ids": torch.randint(0, 1000, (64,)),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def test_tokenizer():
|
||||
"""Session-scoped tokenizer, created once for the entire test run."""
|
||||
@@ -103,50 +45,20 @@ def test_tokenizer():
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def test_model():
|
||||
def test_model(device):
|
||||
"""Session-scoped small AutoRegressiveLM model, created once."""
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=1000,
|
||||
hidden_size=8,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=16,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
config = make_tiny_config()
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"device": device,
|
||||
"config": config,
|
||||
}
|
||||
return {"model": model, "device": device, "config": config}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def base_test_env(test_model, test_tokenizer):
|
||||
"""Function-scoped test environment with isolated temp directory.
|
||||
|
||||
Composes session-scoped model and tokenizer with a per-test temp dir.
|
||||
"""
|
||||
"""Function-scoped test environment with isolated temp directory."""
|
||||
test_dir = tempfile.mkdtemp()
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(
|
||||
{
|
||||
"vocab_size": 1000,
|
||||
"hidden_size": 8,
|
||||
"num_attention_heads": 2,
|
||||
"num_key_value_heads": 1,
|
||||
"intermediate_size": 16,
|
||||
"max_position_embeddings": 64,
|
||||
"num_hidden_layers": 2,
|
||||
"rms_norm_eps": 1e-5,
|
||||
},
|
||||
f,
|
||||
)
|
||||
json.dump(TINY_CONFIG, f)
|
||||
|
||||
yield {
|
||||
"device": test_model["device"],
|
||||
@@ -162,17 +74,14 @@ def base_test_env(test_model, test_tokenizer):
|
||||
|
||||
@pytest.fixture
|
||||
def random_dataset():
|
||||
dataset = RandomDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset(length=None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def multi_turn_dataset():
|
||||
dataset = MultiTurnDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset(length=None, with_loss_mask=True)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def early_stopping_dataset():
|
||||
dataset = EarlyStoppingDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset(length=10, stop_after=5)
|
||||
|
||||
@@ -16,6 +16,7 @@ from tests.data.conftest import (
|
||||
make_dpo_chat_config,
|
||||
make_grpo_no_template_config,
|
||||
)
|
||||
from tests.helpers import load_shard_meta
|
||||
|
||||
|
||||
def test_filter_by_length():
|
||||
@@ -68,10 +69,7 @@ def test_full_chat_pipeline(temp_dir, chat_tokenizer_dir):
|
||||
tokenizer_path=chat_tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "sequence" in meta
|
||||
assert "loss_mask" in meta
|
||||
assert meta["sequence"]["dtype"] == "int32"
|
||||
@@ -112,10 +110,7 @@ def test_full_text_pipeline(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "sequence" in meta
|
||||
assert "loss_mask" not in meta
|
||||
|
||||
@@ -158,10 +153,7 @@ def test_full_instruction_pipeline(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "sequence" in meta
|
||||
assert "loss_mask" in meta
|
||||
|
||||
@@ -187,9 +179,7 @@ def test_dtype_override(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert meta["sequence"]["dtype"] == "int32"
|
||||
assert meta["loss_mask"]["dtype"] == "bool"
|
||||
|
||||
@@ -221,10 +211,7 @@ def test_dpo_pipeline(temp_dir, chat_tokenizer_dir):
|
||||
tokenizer_path=chat_tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "chosen" in meta
|
||||
assert "rejected" in meta
|
||||
assert "chosen_mask" in meta
|
||||
@@ -254,10 +241,7 @@ def test_grpo_pipeline(temp_dir, tokenizer_dir):
|
||||
tokenizer_path=tokenizer_dir,
|
||||
).run()
|
||||
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path)
|
||||
with open(meta_path, "r") as f:
|
||||
meta = json.load(f)
|
||||
meta = load_shard_meta(out_dir)
|
||||
assert "prompts" in meta
|
||||
assert "responses" in meta
|
||||
assert "masks" in meta
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
"""Shared test helpers for the AstrAI test suite."""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=1000,
|
||||
hidden_size=8,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=16,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
CHAT_TEMPLATE = (
|
||||
"{% for message in messages %}"
|
||||
"{% if message['role'] == 'system' %}SYSTEM: {{ message['content'] }}\n{% endif %}"
|
||||
"{% if message['role'] == 'user' %}USER: {{ message['content'] }}\n{% endif %}"
|
||||
"{% if message['role'] == 'assistant' %}ASSISTANT: {{ message['content'] }}\n{% endif %}"
|
||||
"{% endfor %}"
|
||||
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||
)
|
||||
|
||||
|
||||
def make_tiny_config(**overrides):
|
||||
"""Create a tiny ``AutoRegressiveLMConfig`` for tests.
|
||||
|
||||
All keyword arguments override ``TINY_CONFIG`` defaults.
|
||||
"""
|
||||
return AutoRegressiveLMConfig(**{**TINY_CONFIG, **overrides})
|
||||
|
||||
|
||||
def make_rollout_config(vocab_size=200, max_position_embeddings=64, **kwargs):
|
||||
"""Create a tiny config sized for rollout / strategy tests."""
|
||||
return make_tiny_config(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def make_model(device, **cfg_overrides):
|
||||
"""Create a tiny ``AutoRegressiveLM`` on *device* and return ``(model, config)``."""
|
||||
cfg = make_rollout_config(**cfg_overrides)
|
||||
model = AutoRegressiveLM(cfg).to(device=device)
|
||||
model.eval()
|
||||
return model, cfg
|
||||
|
||||
|
||||
def make_frozen(model, device):
|
||||
"""Create a frozen, eval-mode copy of *model* with identical weights."""
|
||||
cfg = make_rollout_config()
|
||||
copy = AutoRegressiveLM(cfg).to(device=device)
|
||||
copy.load_state_dict(model.state_dict())
|
||||
copy.requires_grad_(False)
|
||||
copy.eval()
|
||||
return copy
|
||||
|
||||
|
||||
class RandomTokenDataset(Dataset):
|
||||
"""Random token dataset combining all test dataset variants.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
length : int or None
|
||||
Fixed length, or ``None`` for a random length in [100, 200).
|
||||
max_length : int
|
||||
Sequence length per sample.
|
||||
vocab_size : int
|
||||
Upper bound for random token ids.
|
||||
with_loss_mask : bool
|
||||
Include a ``loss_mask`` key in each sample.
|
||||
stop_after : int or None
|
||||
Raise ``RuntimeError`` after this many samples (for early-stopping tests).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
length=100,
|
||||
max_length=64,
|
||||
vocab_size=1000,
|
||||
*,
|
||||
with_loss_mask=False,
|
||||
stop_after=None,
|
||||
):
|
||||
self.length = (
|
||||
length if length is not None else int(torch.randint(100, 200, (1,)).item())
|
||||
)
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
self.with_loss_mask = with_loss_mask
|
||||
self.stop_after = stop_after
|
||||
self._count = 0
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.stop_after is not None:
|
||||
self._count += 1
|
||||
if self._count == self.stop_after:
|
||||
raise RuntimeError("Simulated early stopping")
|
||||
|
||||
item = {
|
||||
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
}
|
||||
if self.with_loss_mask:
|
||||
item["loss_mask"] = torch.randint(0, 1, (self.max_length,))
|
||||
return item
|
||||
|
||||
|
||||
class FakeTokenizer:
|
||||
"""Minimal stub tokenizer with optional chat-template support."""
|
||||
|
||||
stop_ids = [2]
|
||||
|
||||
def __init__(self, *, with_chat_template=False):
|
||||
if with_chat_template:
|
||||
from astrai.tokenize.chat_template import ChatTemplate
|
||||
|
||||
self._chat_template = ChatTemplate.from_string(CHAT_TEMPLATE)
|
||||
else:
|
||||
self._chat_template = None
|
||||
|
||||
def encode(self, texts, **_):
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
return [[b for b in t.encode("utf-8")] for t in texts]
|
||||
|
||||
def decode(self, ids, skip_special_tokens=True):
|
||||
if isinstance(ids, list):
|
||||
return bytes(b for b in ids if b > 2 or not skip_special_tokens).decode(
|
||||
"utf-8", errors="ignore"
|
||||
)
|
||||
return str(ids)
|
||||
|
||||
def apply_chat_template(
|
||||
self, messages, tokenize=True, add_generation_prompt=True, **_
|
||||
):
|
||||
if self._chat_template is None:
|
||||
raise RuntimeError("Chat template not configured")
|
||||
rendered = self._chat_template.render(
|
||||
messages=messages, add_generation_prompt=add_generation_prompt
|
||||
)
|
||||
if tokenize:
|
||||
return (
|
||||
self.encode(rendered)[0]
|
||||
if isinstance(rendered, str)
|
||||
else [self.encode(t)[0] for t in rendered]
|
||||
)
|
||||
return rendered
|
||||
|
||||
|
||||
class FakeExecutor:
|
||||
"""Executor stub tracking ``sync_gradients`` and providing ``unwrap_model``."""
|
||||
|
||||
def __init__(self, sync_gradients=True):
|
||||
self._sync_gradients = sync_gradients
|
||||
|
||||
@property
|
||||
def sync_gradients(self):
|
||||
return self._sync_gradients
|
||||
|
||||
def unwrap_model(self, model):
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
def find_checkpoint_meta(ckpt_dir):
|
||||
"""Walk *ckpt_dir* and return the path to the first ``meta.json`` found."""
|
||||
for root, _dirs, files in os.walk(ckpt_dir):
|
||||
if "meta.json" in files:
|
||||
return os.path.join(root, "meta.json")
|
||||
return None
|
||||
|
||||
|
||||
def load_checkpoint_meta(ckpt_dir):
|
||||
"""Find and load the first checkpoint ``meta.json`` under *ckpt_dir*."""
|
||||
meta_path = find_checkpoint_meta(ckpt_dir)
|
||||
assert meta_path is not None, f"No checkpoint meta.json found in {ckpt_dir}"
|
||||
with open(meta_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def load_shard_meta(out_dir):
|
||||
"""Load ``meta.json`` from the default shard output directory."""
|
||||
meta_path = os.path.join(out_dir, "__default__", "shard_0000", "meta.json")
|
||||
assert os.path.exists(meta_path), f"Shard meta.json not found at {meta_path}"
|
||||
with open(meta_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def assert_state_dicts_equal(a, b):
|
||||
"""Assert two state dicts have identical keys and equal tensor values."""
|
||||
assert set(a.keys()) == set(b.keys()), f"Key mismatch: {set(a) ^ set(b)}"
|
||||
for key in a:
|
||||
assert torch.equal(a[key], b[key]), f"Tensor mismatch at key: {key}"
|
||||
@@ -7,6 +7,8 @@ import pytest
|
||||
import torch
|
||||
|
||||
from astrai.inference import InferenceScheduler
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from tests.helpers import FakeTokenizer, make_rollout_config
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -195,34 +197,9 @@ def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
|
||||
|
||||
def _make_real_scheduler(device):
|
||||
"""Build a scheduler backed by a tiny real model for run_batch tests."""
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
class _Tok:
|
||||
stop_ids = [2]
|
||||
|
||||
def encode(self, texts, **_):
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
return [[b for b in t.encode("utf-8")] for t in texts]
|
||||
|
||||
def decode(self, ids, skip_special_tokens=True):
|
||||
return bytes(b for b in ids if b > 2 or not skip_special_tokens).decode(
|
||||
"utf-8", errors="ignore"
|
||||
)
|
||||
|
||||
cfg = AutoRegressiveLMConfig(
|
||||
vocab_size=200,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
cfg = make_rollout_config(max_position_embeddings=64)
|
||||
model = AutoRegressiveLM(cfg).to(device=device).eval()
|
||||
tokenizer = _Tok()
|
||||
tokenizer = FakeTokenizer()
|
||||
scheduler = InferenceScheduler(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
@@ -232,8 +209,7 @@ def _make_real_scheduler(device):
|
||||
return scheduler, tokenizer, model
|
||||
|
||||
|
||||
def test_run_batch_returns_token_sequences():
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
def test_run_batch_returns_token_sequences(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30], [5, 6, 7, 8]]
|
||||
@@ -247,9 +223,8 @@ def test_run_batch_returns_token_sequences():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_return_logprobs_aligned():
|
||||
def test_run_batch_return_logprobs_aligned(device):
|
||||
"""return_logprobs=True gives (token_ids, logprobs) tuples with equal len."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30, 40]]
|
||||
@@ -264,8 +239,7 @@ def test_run_batch_return_logprobs_aligned():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_respects_max_tokens():
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
def test_run_batch_respects_max_tokens(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30]]
|
||||
@@ -275,9 +249,8 @@ def test_run_batch_respects_max_tokens():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_stop_id_terminates():
|
||||
def test_run_batch_stop_id_terminates(device):
|
||||
"""A token matching stop_ids terminates generation for that prompt."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30]]
|
||||
@@ -290,9 +263,8 @@ def test_run_batch_stop_id_terminates():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_empty_prompts():
|
||||
def test_run_batch_empty_prompts(device):
|
||||
"""Empty prompt list yields empty result list."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
assert scheduler.run_batch([], max_tokens=4) == []
|
||||
@@ -300,9 +272,8 @@ def test_run_batch_empty_prompts():
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_too_long_prompt_skipped():
|
||||
def test_run_batch_too_long_prompt_skipped(device):
|
||||
"""A prompt longer than max_seq_len yields an empty result slot."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
long = list(range(100)) # > max_seq_len=64
|
||||
|
||||
@@ -9,34 +9,22 @@ import torch
|
||||
from astrai.config.model_config import EncoderConfig
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.encoder import EmbeddingEncoder
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=128,
|
||||
hidden_size=8,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=16,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
_device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
from tests.helpers import TINY_CONFIG, assert_state_dicts_equal
|
||||
|
||||
|
||||
def _make_model(**kwargs):
|
||||
def _make_model(device, **kwargs):
|
||||
config = EncoderConfig(**{**TINY_CONFIG, **kwargs})
|
||||
return EmbeddingEncoder(config).to(device=_device)
|
||||
return EmbeddingEncoder(config).to(device=device)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("pooling_type", ["mean", "cls", "last"])
|
||||
def test_encoder_forward_pooling(pooling_type):
|
||||
model = _make_model(pooling_type=pooling_type)
|
||||
def test_encoder_forward_pooling(pooling_type, device):
|
||||
model = _make_model(device, pooling_type=pooling_type)
|
||||
model.eval()
|
||||
|
||||
batch_size, seq_len = 2, 8
|
||||
input_ids = torch.randint(
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=_device
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=device
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -46,15 +34,15 @@ def test_encoder_forward_pooling(pooling_type):
|
||||
assert not torch.isnan(output).any()
|
||||
|
||||
|
||||
def test_encoder_forward_with_padding():
|
||||
model = _make_model()
|
||||
def test_encoder_forward_with_padding(device):
|
||||
model = _make_model(device)
|
||||
model.eval()
|
||||
|
||||
batch_size, seq_len = 2, 8
|
||||
input_ids = torch.randint(
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=_device
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=device
|
||||
)
|
||||
input_mask = torch.ones(batch_size, seq_len, dtype=torch.bool, device=_device)
|
||||
input_mask = torch.ones(batch_size, seq_len, dtype=torch.bool, device=device)
|
||||
input_mask[:, 4:] = False
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -64,13 +52,13 @@ def test_encoder_forward_with_padding():
|
||||
assert not torch.isnan(output).any()
|
||||
|
||||
|
||||
def test_encoder_normalize():
|
||||
model = _make_model(pooling_type="mean", normalize_embeddings=True)
|
||||
def test_encoder_normalize(device):
|
||||
model = _make_model(device, pooling_type="mean", normalize_embeddings=True)
|
||||
model.eval()
|
||||
|
||||
batch_size, seq_len = 2, 8
|
||||
input_ids = torch.randint(
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=_device
|
||||
0, TINY_CONFIG["vocab_size"], (batch_size, seq_len), device=device
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -86,26 +74,24 @@ def test_encoder_register():
|
||||
assert cls is EmbeddingEncoder
|
||||
|
||||
|
||||
def test_encoder_from_transformer_checkpoint():
|
||||
model = _make_model()
|
||||
def test_encoder_from_transformer_checkpoint(device):
|
||||
model = _make_model(device)
|
||||
state_dict = model.state_dict()
|
||||
state_dict["lm_head.weight"] = torch.randn(
|
||||
TINY_CONFIG["vocab_size"], TINY_CONFIG["hidden_size"], device=_device
|
||||
TINY_CONFIG["vocab_size"], TINY_CONFIG["hidden_size"], device=device
|
||||
)
|
||||
|
||||
new_model = _make_model()
|
||||
new_model = _make_model(device)
|
||||
new_model.load_state_dict(state_dict, strict=True)
|
||||
|
||||
for key in model.state_dict():
|
||||
assert torch.equal(new_model.state_dict()[key], model.state_dict()[key])
|
||||
assert_state_dicts_equal(new_model.state_dict(), model.state_dict())
|
||||
|
||||
|
||||
def test_encoder_save_load():
|
||||
test_dir = tempfile.mkdtemp(prefix="encoder_test_")
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
weights_path = os.path.join(test_dir, "model.safetensors")
|
||||
def test_encoder_save_load(device):
|
||||
with tempfile.TemporaryDirectory(prefix="encoder_test_") as test_dir:
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
weights_path = os.path.join(test_dir, "model.safetensors")
|
||||
|
||||
try:
|
||||
config_data = {**TINY_CONFIG, "pooling_type": "mean"}
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
@@ -117,10 +103,4 @@ def test_encoder_save_load():
|
||||
loaded = EmbeddingEncoder(config)
|
||||
loaded.load_state_dict(st.load_file(weights_path))
|
||||
|
||||
for key in original.state_dict():
|
||||
assert torch.equal(original.state_dict()[key], loaded.state_dict()[key])
|
||||
finally:
|
||||
if os.path.exists(test_dir):
|
||||
for f in os.listdir(test_dir):
|
||||
os.remove(os.path.join(test_dir, f))
|
||||
os.rmdir(test_dir)
|
||||
assert_state_dicts_equal(original.state_dict(), loaded.state_dict())
|
||||
|
||||
@@ -1,20 +1,8 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
TINY_CONFIG = dict(
|
||||
vocab_size=128,
|
||||
hidden_size=8,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=16,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
from tests.helpers import TINY_CONFIG
|
||||
|
||||
CONFIGS = [
|
||||
pytest.param(
|
||||
@@ -70,9 +58,10 @@ CONFIGS = [
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_kwargs", CONFIGS)
|
||||
def test_model_forward(config_kwargs):
|
||||
def test_model_forward(config_kwargs, device):
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
|
||||
config = AutoRegressiveLMConfig(**config_kwargs)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
model.eval()
|
||||
|
||||
@@ -97,9 +86,10 @@ def test_model_forward(config_kwargs):
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_kwargs", CONFIGS)
|
||||
def test_model_forward_with_padding(config_kwargs):
|
||||
def test_model_forward_with_padding(config_kwargs, device):
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
|
||||
config = AutoRegressiveLMConfig(**config_kwargs)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
model.eval()
|
||||
|
||||
|
||||
+28
-29
@@ -249,17 +249,17 @@ def test_save_load_roundtrip():
|
||||
with torch.no_grad():
|
||||
out_src = model(x)["logits"].clone()
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, cfg)
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
load_lora(model2, tmpdir)
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
load_lora(model2, tmpdir)
|
||||
|
||||
with torch.no_grad():
|
||||
out_dst = model2(x)["logits"]
|
||||
with torch.no_grad():
|
||||
out_dst = model2(x)["logits"]
|
||||
|
||||
torch.testing.assert_close(out_src, out_dst)
|
||||
torch.testing.assert_close(out_src, out_dst)
|
||||
|
||||
|
||||
def test_save_after_merge_raises():
|
||||
@@ -271,13 +271,13 @@ def test_save_after_merge_raises():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
merge_lora(model)
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, cfg)
|
||||
merge_lora(model)
|
||||
|
||||
tmpdir2 = tempfile.mkdtemp()
|
||||
with pytest.raises(RuntimeError, match="No LoRA parameters"):
|
||||
save_lora(model, tmpdir2, cfg)
|
||||
with tempfile.TemporaryDirectory() as tmpdir2:
|
||||
with pytest.raises(RuntimeError, match="No LoRA parameters"):
|
||||
save_lora(model, tmpdir2, cfg)
|
||||
|
||||
|
||||
def test_load_lora_on_already_injected():
|
||||
@@ -289,16 +289,15 @@ def test_load_lora_on_already_injected():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, LoRAConfig(r=4, alpha=8, target_modules=("q_proj",)))
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, LoRAConfig(r=4, alpha=8, target_modules=("q_proj",)))
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
# load onto already-injected model
|
||||
load_lora(model2, tmpdir)
|
||||
assert _get_lora_count(model2) > 0
|
||||
load_lora(model2, tmpdir)
|
||||
assert _get_lora_count(model2) > 0
|
||||
|
||||
|
||||
def test_load_lora_mismatched_r_raises():
|
||||
@@ -310,15 +309,15 @@ def test_load_lora_mismatched_r_raises():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
save_lora(model, tmpdir, cfg)
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with pytest.raises(RuntimeError, match="size mismatch"):
|
||||
load_lora(model2, tmpdir) # strict=False, only lora keys
|
||||
with pytest.raises(RuntimeError, match="size mismatch"):
|
||||
load_lora(model2, tmpdir)
|
||||
|
||||
|
||||
def test_merge_preserves_output():
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import safetensors.torch as st
|
||||
@@ -8,43 +7,13 @@ import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from tests.helpers import TINY_CONFIG
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def transformer_test_env():
|
||||
test_dir = tempfile.mkdtemp(prefix="transformer_test_")
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
def test_tie_weight_init(base_test_env):
|
||||
config_path = base_test_env["config_path"]
|
||||
|
||||
config = {
|
||||
"vocab_size": 1000,
|
||||
"hidden_size": 8,
|
||||
"num_attention_heads": 2,
|
||||
"num_key_value_heads": 1,
|
||||
"intermediate_size": 16,
|
||||
"max_position_embeddings": 64,
|
||||
"num_hidden_layers": 2,
|
||||
"rms_norm_eps": 1e-5,
|
||||
}
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config, f)
|
||||
|
||||
yield {"test_dir": test_dir, "config_path": config_path, "config": config}
|
||||
|
||||
if os.path.exists(test_dir):
|
||||
try:
|
||||
for file in os.listdir(test_dir):
|
||||
os.remove(os.path.join(test_dir, file))
|
||||
os.rmdir(test_dir)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def test_tie_weight_init(transformer_test_env):
|
||||
config_path = transformer_test_env["config_path"]
|
||||
config_data = transformer_test_env["config"].copy()
|
||||
|
||||
# case 1: tie weight
|
||||
config_data = TINY_CONFIG.copy()
|
||||
config_data["tie_word_embeddings"] = True
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
@@ -62,7 +31,6 @@ def test_tie_weight_init(transformer_test_env):
|
||||
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
||||
assert not torch.equal(model.lm_head.weight, original_weight)
|
||||
|
||||
# case 2: not tie weight
|
||||
config_data["tie_word_embeddings"] = False
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
@@ -81,13 +49,11 @@ def test_tie_weight_init(transformer_test_env):
|
||||
assert not torch.equal(model.lm_head.weight, original_weight)
|
||||
|
||||
|
||||
def test_model_save_load_with_tie_weight(transformer_test_env):
|
||||
test_dir = transformer_test_env["test_dir"]
|
||||
def test_model_save_load_with_tie_weight(base_test_env):
|
||||
test_dir = base_test_env["test_dir"]
|
||||
model_path = os.path.join(test_dir, "model.safetensors")
|
||||
|
||||
config_data = transformer_test_env["config"].copy()
|
||||
|
||||
# case 1: tie weight
|
||||
config_data = TINY_CONFIG.copy()
|
||||
config_data["tie_word_embeddings"] = True
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
|
||||
@@ -107,7 +73,6 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
|
||||
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
|
||||
assert "lm_head.weight" not in model.state_dict()
|
||||
|
||||
# case 2: not tie weight (form tie-weight state dict load)
|
||||
config_data["tie_word_embeddings"] = False
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_data, f)
|
||||
|
||||
@@ -2,33 +2,15 @@ import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
|
||||
|
||||
class TrainerDataset(Dataset):
|
||||
"""Base dataset for trainer tests with consistent interface."""
|
||||
|
||||
def __init__(self, length=100, max_length=64, vocab_size=1000):
|
||||
self.length = length
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {
|
||||
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
}
|
||||
from tests.helpers import RandomTokenDataset
|
||||
|
||||
|
||||
def create_train_config(
|
||||
model_fn,
|
||||
dataset: Dataset,
|
||||
dataset,
|
||||
test_dir: str,
|
||||
device: str,
|
||||
strategy: str = "seq",
|
||||
@@ -40,25 +22,7 @@ def create_train_config(
|
||||
random_seed: int = 42,
|
||||
**kwargs,
|
||||
):
|
||||
"""Factory function to create common TrainConfig for tests.
|
||||
|
||||
Args:
|
||||
model_fn: Model factory (callable returning nn.Module)
|
||||
dataset: Training dataset
|
||||
test_dir: Checkpoint directory
|
||||
device: Device type ("cuda" or "cpu")
|
||||
strategy: Training strategy type (default: "seq")
|
||||
n_epoch: Number of epochs (default: 1)
|
||||
batch_per_device: Batch size per device (default: 2)
|
||||
grad_accum_steps: Gradient accumulation steps (default: 1)
|
||||
max_grad_norm: Maximum gradient norm for clipping (default: 1.0)
|
||||
ckpt_interval: Checkpoint save interval in optimizer steps (default: 5)
|
||||
random_seed: Random seed for reproducibility (default: 42)
|
||||
**kwargs: Additional arguments passed to TrainConfig
|
||||
|
||||
Returns:
|
||||
TrainConfig instance configured for testing
|
||||
"""
|
||||
"""Factory function to create common TrainConfig for tests."""
|
||||
|
||||
def optimizer_fn(m):
|
||||
return torch.optim.AdamW(m.parameters(), lr=0.001)
|
||||
@@ -89,16 +53,11 @@ def create_train_config(
|
||||
|
||||
@pytest.fixture
|
||||
def train_config_factory():
|
||||
"""Fixture that provides the create_train_config factory function.
|
||||
|
||||
This fixture can be used by tests to create consistent TrainConfig
|
||||
instances with sensible defaults for testing.
|
||||
"""
|
||||
"""Fixture providing the ``create_train_config`` factory function."""
|
||||
return create_train_config
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def trainer_dataset():
|
||||
"""Fixture providing a dataset for trainer tests."""
|
||||
dataset = TrainerDataset()
|
||||
yield dataset
|
||||
return RandomTokenDataset()
|
||||
|
||||
@@ -1,10 +1,6 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
from astrai.trainer.train_callback import GradientCheckpointingCallback, TrainCallback
|
||||
from astrai.trainer.trainer import Trainer
|
||||
|
||||
@@ -94,69 +90,35 @@ def test_gradient_checkpointing_backward(test_model):
|
||||
assert p.grad is None or p.grad.sum().item() == 0, f"{name} grad not zeroed"
|
||||
|
||||
|
||||
def test_gradient_checkpointing_trainer_integration(base_test_env, random_dataset):
|
||||
def test_gradient_checkpointing_trainer_integration(
|
||||
base_test_env, random_dataset, train_config_factory, device
|
||||
):
|
||||
"""Gradient checkpointing runs end-to-end via Trainer."""
|
||||
|
||||
def optimizer_fn(model):
|
||||
return torch.optim.AdamW(model.parameters())
|
||||
|
||||
def scheduler_fn(optim):
|
||||
return SchedulerFactory.create(
|
||||
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
|
||||
)
|
||||
|
||||
train_config = TrainConfig(
|
||||
train_config = train_config_factory(
|
||||
model_fn=lambda: base_test_env["model"],
|
||||
strategy="seq",
|
||||
dataset=random_dataset,
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=base_test_env["test_dir"],
|
||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||
n_epoch=1,
|
||||
batch_per_device=2,
|
||||
test_dir=base_test_env["test_dir"],
|
||||
device=device,
|
||||
ckpt_interval=3,
|
||||
grad_accum_steps=1,
|
||||
max_grad_norm=1.0,
|
||||
random_seed=42,
|
||||
device_type=base_test_env["device"],
|
||||
gradient_checkpointing_modules=[DecoderBlock],
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
trainer.train()
|
||||
# no crash = callback correctly enabled/disabled
|
||||
|
||||
|
||||
def test_callback_integration(base_test_env, random_dataset):
|
||||
def test_callback_integration(
|
||||
base_test_env, random_dataset, train_config_factory, device
|
||||
):
|
||||
"""Test that all callbacks are properly integrated"""
|
||||
|
||||
def optimizer_fn(model):
|
||||
return torch.optim.AdamW(model.parameters())
|
||||
|
||||
def scheduler_fn(optim):
|
||||
return SchedulerFactory.create(
|
||||
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
|
||||
)
|
||||
|
||||
train_config = TrainConfig(
|
||||
train_config = train_config_factory(
|
||||
model_fn=lambda: base_test_env["model"],
|
||||
strategy="seq",
|
||||
dataset=random_dataset,
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=base_test_env["test_dir"],
|
||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||
n_epoch=1,
|
||||
batch_per_device=2,
|
||||
test_dir=base_test_env["test_dir"],
|
||||
device=device,
|
||||
ckpt_interval=3,
|
||||
grad_accum_steps=1,
|
||||
max_grad_norm=1.0,
|
||||
random_seed=42,
|
||||
device_type=base_test_env["device"],
|
||||
)
|
||||
|
||||
# Create custom callbacks to track calls
|
||||
callback_calls = []
|
||||
|
||||
class TrackingCallback(TrainCallback):
|
||||
@@ -170,10 +132,8 @@ def test_callback_integration(base_test_env, random_dataset):
|
||||
callback_calls.append("on_epoch_end")
|
||||
|
||||
trainer = Trainer(train_config, callbacks=[TrackingCallback()])
|
||||
|
||||
trainer.train()
|
||||
|
||||
# Verify callbacks were called
|
||||
assert "on_train_begin" in callback_calls
|
||||
assert "on_batch_end" in callback_calls
|
||||
assert "on_epoch_end" in callback_calls
|
||||
|
||||
@@ -1,43 +1,25 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
from astrai.trainer.trainer import Trainer
|
||||
from tests.helpers import load_checkpoint_meta
|
||||
|
||||
|
||||
def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
|
||||
def test_early_stopping_simulation(
|
||||
base_test_env, early_stopping_dataset, train_config_factory, device
|
||||
):
|
||||
"""Simulate early stopping behavior"""
|
||||
|
||||
def optimizer_fn(model):
|
||||
return torch.optim.AdamW(model.parameters())
|
||||
|
||||
def scheduler_fn(optim):
|
||||
return SchedulerFactory.create(
|
||||
"cosine", optim, warmup_steps=10, lr_decay_steps=10, min_rate=0.05
|
||||
)
|
||||
|
||||
train_config = TrainConfig(
|
||||
strategy="seq",
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
train_config = train_config_factory(
|
||||
model_fn=lambda: base_test_env["model"],
|
||||
dataset=early_stopping_dataset,
|
||||
ckpt_dir=base_test_env["test_dir"],
|
||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||
test_dir=base_test_env["test_dir"],
|
||||
device=device,
|
||||
n_epoch=2,
|
||||
batch_per_device=2,
|
||||
ckpt_interval=1,
|
||||
grad_accum_steps=2,
|
||||
random_seed=np.random.randint(1e4),
|
||||
device_type=base_test_env["device"],
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
|
||||
# Should handle early stopping gracefully
|
||||
try:
|
||||
trainer.train()
|
||||
except Exception:
|
||||
@@ -50,8 +32,5 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
|
||||
|
||||
# Verify checkpoint was saved at expected step
|
||||
load_dir = os.path.join(base_test_env["test_dir"], "epoch_1_step_5")
|
||||
import json
|
||||
|
||||
with open(os.path.join(load_dir, "meta.json")) as f:
|
||||
meta = json.load(f)
|
||||
meta = load_checkpoint_meta(load_dir)
|
||||
assert meta["consumed_samples"] == 20
|
||||
|
||||
@@ -1,35 +1,9 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.strategy import GRPOStrategy
|
||||
|
||||
|
||||
class _FakeExecutor:
|
||||
"""Minimal executor stub providing ``unwrap_model`` for ref model creation."""
|
||||
|
||||
def unwrap_model(self, model):
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_position_embeddings=64):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(device):
|
||||
config = _make_config()
|
||||
model = AutoRegressiveLM(config).to(device=device)
|
||||
return model, config
|
||||
from tests.helpers import FakeExecutor, make_frozen, make_model, make_rollout_config
|
||||
|
||||
|
||||
def _make_batch(
|
||||
@@ -44,9 +18,7 @@ def _make_batch(
|
||||
responses = torch.randint(
|
||||
0, 200, (batch_size, group_size, response_len), device=device
|
||||
)
|
||||
# All response tokens valid.
|
||||
masks = torch.ones(batch_size, group_size, response_len, device=device)
|
||||
# Distinct rewards per group member so std > 0.
|
||||
rewards = torch.randn(batch_size, group_size, device=device)
|
||||
return {
|
||||
"prompts": prompts,
|
||||
@@ -56,23 +28,12 @@ def _make_batch(
|
||||
}
|
||||
|
||||
|
||||
def _make_frozen_copy(model, device):
|
||||
"""Create a frozen copy of ``model`` with independent weights loaded."""
|
||||
config = _make_config()
|
||||
copy = AutoRegressiveLM(config).to(device=device)
|
||||
copy.load_state_dict(model.state_dict())
|
||||
copy.requires_grad_(False)
|
||||
copy.eval()
|
||||
return copy
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def grpo_strategy():
|
||||
def grpo_strategy(device):
|
||||
"""Build a GRPOStrategy with a small real model and fake executor."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model, config = _make_model(device)
|
||||
old_model = _make_frozen_copy(model, device)
|
||||
ref_model = _make_frozen_copy(model, device)
|
||||
model, config = make_model(device)
|
||||
old_model = make_frozen(model, device)
|
||||
ref_model = make_frozen(model, device)
|
||||
|
||||
strategy = GRPOStrategy(
|
||||
model=model,
|
||||
@@ -83,7 +44,7 @@ def grpo_strategy():
|
||||
kl_coef=0.01,
|
||||
group_size=4,
|
||||
model_fn=lambda c=config: AutoRegressiveLM(c).to(device=device),
|
||||
executor=_FakeExecutor(),
|
||||
executor=FakeExecutor(),
|
||||
)
|
||||
return strategy, device
|
||||
|
||||
@@ -103,7 +64,6 @@ def test_grpo_loss_backward(grpo_strategy):
|
||||
batch = _make_batch(device=device)
|
||||
loss = strategy.compute_loss(batch)
|
||||
loss.backward()
|
||||
# At least some parameter should receive a gradient.
|
||||
has_grad = any(
|
||||
p.grad is not None and p.grad.abs().sum().item() > 0
|
||||
for p in strategy.model.parameters()
|
||||
@@ -136,32 +96,27 @@ def test_grpo_prompt_tokens_masked(grpo_strategy):
|
||||
the policy loss should be zero (no valid tokens contribute)."""
|
||||
strategy, device = grpo_strategy
|
||||
batch = _make_batch(device=device)
|
||||
# Zero out all response masks → no response token contributes.
|
||||
batch["masks"] = torch.zeros_like(batch["masks"])
|
||||
loss = strategy.compute_loss(batch)
|
||||
# With no valid tokens, policy_loss term is 0 and KL term is 0.
|
||||
assert loss.item() == pytest.approx(0.0, abs=1e-6)
|
||||
|
||||
|
||||
def test_grpo_identical_rewards_zero_advantage(grpo_strategy):
|
||||
"""When all group rewards are identical, advantage is 0 → policy_loss is 0.
|
||||
"""When all group rewards are identical, advantage is 0 -> policy_loss is 0.
|
||||
Only the KL term remains (which is 0 when policy == ref at init)."""
|
||||
strategy, device = grpo_strategy
|
||||
batch = _make_batch(device=device)
|
||||
batch["rewards"] = torch.ones(batch["rewards"].shape, device=device)
|
||||
loss = strategy.compute_loss(batch)
|
||||
# At init policy == old == ref, so ratio == 1, KL == 0; advantage == 0.
|
||||
assert loss.item() == pytest.approx(0.0, abs=1e-5)
|
||||
|
||||
|
||||
def test_grpo_sync_old_model(grpo_strategy):
|
||||
"""sync_old_model copies current policy weights into old_model."""
|
||||
strategy, device = grpo_strategy
|
||||
# Perturb policy model so it differs from old.
|
||||
with torch.no_grad():
|
||||
for p in strategy.model.parameters():
|
||||
p.add_(0.05)
|
||||
# old_model should still hold original weights (differ from policy).
|
||||
policy_sd = strategy.model.state_dict()
|
||||
old_sd = strategy.old_model.state_dict()
|
||||
differs_before = any(
|
||||
@@ -195,14 +150,12 @@ def test_grpo_clipping_effect(grpo_strategy):
|
||||
"""After diverging policy from ref, ratio should be clipped to [1-eps, 1+eps]
|
||||
on the surrogate. Verify loss is finite and non-zero for distinct rewards."""
|
||||
strategy, device = grpo_strategy
|
||||
# Diverge policy from ref.
|
||||
with torch.no_grad():
|
||||
for p in strategy.model.parameters():
|
||||
p.add_(0.3)
|
||||
batch = _make_batch(device=device)
|
||||
loss = strategy.compute_loss(batch)
|
||||
assert torch.isfinite(loss).item()
|
||||
# With distinct rewards and diverged policy, loss should be non-trivial.
|
||||
assert loss.abs().item() > 1e-4
|
||||
|
||||
|
||||
|
||||
@@ -12,19 +12,7 @@ from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.rollout import BaseRewardModel
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
from astrai.trainer.trainer import Trainer
|
||||
|
||||
_CHAT_TEMPLATE = (
|
||||
"{% for message in messages %}"
|
||||
"{% if message['role'] == 'system' %}"
|
||||
"SYSTEM: {{ message['content'] }}\n"
|
||||
"{% elif message['role'] == 'user' %}"
|
||||
"USER: {{ message['content'] }}\n"
|
||||
"{% elif message['role'] == 'assistant' %}"
|
||||
"ASSISTANT: {{ message['content'] }}\n"
|
||||
"{% endif %}"
|
||||
"{% endfor %}"
|
||||
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||
)
|
||||
from tests.helpers import CHAT_TEMPLATE
|
||||
|
||||
|
||||
class InstructionDataset(Dataset):
|
||||
@@ -97,7 +85,7 @@ def test_online_dpo_end_to_end(base_test_env):
|
||||
|
||||
# Equip tokenizer with a chat template so RolloutGenerator can
|
||||
# render instruction/input via apply_chat_template.
|
||||
tokenizer.set_chat_template(_CHAT_TEMPLATE)
|
||||
tokenizer.set_chat_template(CHAT_TEMPLATE)
|
||||
tokenizer.save_pretrained(test_dir)
|
||||
|
||||
model_fn = partial(_model_fn, model_config)
|
||||
|
||||
@@ -9,7 +9,6 @@ the per-strategy ``prepare_from_rollout`` mappings for both
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.rollout import RolloutResult
|
||||
from astrai.trainer.strategy import (
|
||||
@@ -17,47 +16,7 @@ from astrai.trainer.strategy import (
|
||||
GRPOStrategy,
|
||||
StrategyFactory,
|
||||
)
|
||||
|
||||
|
||||
class _FakeExecutor:
|
||||
"""Executor stub tracking ``sync_gradients`` and providing unwrap_model."""
|
||||
|
||||
def __init__(self, sync_gradients=True):
|
||||
self._sync_gradients = sync_gradients
|
||||
|
||||
@property
|
||||
def sync_gradients(self):
|
||||
return self._sync_gradients
|
||||
|
||||
def unwrap_model(self, model):
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_position_embeddings=64):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(device):
|
||||
cfg = _make_config()
|
||||
return AutoRegressiveLM(cfg).to(device=device), cfg
|
||||
|
||||
|
||||
def _make_frozen(model, device):
|
||||
cfg = _make_config()
|
||||
copy = AutoRegressiveLM(cfg).to(device=device)
|
||||
copy.load_state_dict(model.state_dict())
|
||||
copy.requires_grad_(False)
|
||||
copy.eval()
|
||||
return copy
|
||||
from tests.helpers import FakeExecutor, make_frozen, make_model, make_rollout_config
|
||||
|
||||
|
||||
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
|
||||
@@ -75,7 +34,7 @@ class _RecordingRunner:
|
||||
"""Fake RolloutRunner returning a fixed result with freshness tracking.
|
||||
|
||||
Freshness is ``True`` on the first call after construction or after
|
||||
:meth:`swap_result`; ``False`` on subsequent cached calls — mirroring
|
||||
:meth:`swap_result`; ``False`` on subsequent cached calls -- mirroring
|
||||
the real ``RolloutRunner`` contract without invoking generation.
|
||||
"""
|
||||
|
||||
@@ -99,15 +58,10 @@ class _RecordingRunner:
|
||||
self._fresh = True
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def device():
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def _make_grpo(device, executor=None):
|
||||
model, _ = _make_model(device)
|
||||
old_model = _make_frozen(model, device)
|
||||
ref_model = _make_frozen(model, device)
|
||||
model, _ = make_model(device)
|
||||
old_model = make_frozen(model, device)
|
||||
ref_model = make_frozen(model, device)
|
||||
return GRPOStrategy(
|
||||
model=model,
|
||||
device=device,
|
||||
@@ -116,22 +70,22 @@ def _make_grpo(device, executor=None):
|
||||
clip_eps=0.2,
|
||||
kl_coef=0.01,
|
||||
group_size=4,
|
||||
model_fn=lambda c=_make_config(): AutoRegressiveLM(c).to(device=device),
|
||||
executor=executor or _FakeExecutor(),
|
||||
model_fn=lambda c=make_rollout_config(): AutoRegressiveLM(c).to(device=device),
|
||||
executor=executor or FakeExecutor(),
|
||||
)
|
||||
|
||||
|
||||
def _make_dpo(device, executor=None):
|
||||
model, _ = _make_model(device)
|
||||
ref_model = _make_frozen(model, device)
|
||||
model, _ = make_model(device)
|
||||
ref_model = make_frozen(model, device)
|
||||
return DPOStrategy(
|
||||
model=model,
|
||||
device=device,
|
||||
ref_model=ref_model,
|
||||
beta=0.1,
|
||||
reduction="sum",
|
||||
model_fn=lambda c=_make_config(): AutoRegressiveLM(c).to(device=device),
|
||||
executor=executor or _FakeExecutor(),
|
||||
model_fn=lambda c=make_rollout_config(): AutoRegressiveLM(c).to(device=device),
|
||||
executor=executor or FakeExecutor(),
|
||||
)
|
||||
|
||||
|
||||
@@ -295,7 +249,7 @@ def test_dpo_no_sync_hook_when_new_rollout_result(device):
|
||||
|
||||
|
||||
def test_step_not_called_when_sync_gradients_false(device):
|
||||
executor = _FakeExecutor(sync_gradients=False)
|
||||
executor = FakeExecutor(sync_gradients=False)
|
||||
strat = _make_grpo(device, executor=executor)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
@@ -304,7 +258,7 @@ def test_step_not_called_when_sync_gradients_false(device):
|
||||
|
||||
|
||||
def test_step_called_when_sync_gradients_true(device):
|
||||
executor = _FakeExecutor(sync_gradients=True)
|
||||
executor = FakeExecutor(sync_gradients=True)
|
||||
strat = _make_grpo(device, executor=executor)
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
|
||||
@@ -3,9 +3,7 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.trainer.rollout import (
|
||||
BaseRewardModel,
|
||||
RawRollout,
|
||||
@@ -13,50 +11,7 @@ from astrai.trainer.rollout import (
|
||||
RolloutResult,
|
||||
RolloutRunner,
|
||||
)
|
||||
|
||||
_CHAT_TEMPLATE = (
|
||||
"{% for message in messages %}"
|
||||
"{% if message['role'] == 'system' %}SYSTEM: {{ message['content'] }}\n{% endif %}"
|
||||
"{% if message['role'] == 'user' %}USER: {{ message['content'] }}\n{% endif %}"
|
||||
"{% if message['role'] == 'assistant' %}ASSISTANT: {{ message['content'] }}\n{% endif %}"
|
||||
"{% endfor %}"
|
||||
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||
)
|
||||
|
||||
|
||||
class FakeTokenizer:
|
||||
"""Minimal stub tokenizer with a chat template for rollout tests."""
|
||||
|
||||
stop_ids = [2]
|
||||
|
||||
def __init__(self):
|
||||
from astrai.tokenize.chat_template import ChatTemplate
|
||||
|
||||
self._chat_template = ChatTemplate.from_string(_CHAT_TEMPLATE)
|
||||
|
||||
def encode(self, texts, **_):
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
return [[b for b in t.encode("utf-8")] for t in texts]
|
||||
|
||||
def decode(self, ids, skip_special_tokens=True):
|
||||
if isinstance(ids, list):
|
||||
return bytes(b for b in ids if b > 2).decode("utf-8", errors="ignore")
|
||||
return str(ids)
|
||||
|
||||
def apply_chat_template(
|
||||
self, messages, tokenize=True, add_generation_prompt=True, **_
|
||||
):
|
||||
rendered = self._chat_template.render(
|
||||
messages=messages, add_generation_prompt=add_generation_prompt
|
||||
)
|
||||
if tokenize:
|
||||
return (
|
||||
self.encode(rendered)[0]
|
||||
if isinstance(rendered, str)
|
||||
else [self.encode(t)[0] for t in rendered]
|
||||
)
|
||||
return rendered
|
||||
from tests.helpers import FakeTokenizer, make_model
|
||||
|
||||
|
||||
class ConstantRewardModel(BaseRewardModel):
|
||||
@@ -83,26 +38,6 @@ class NonFiniteRewardModel(BaseRewardModel):
|
||||
return torch.full((B, G), float("nan"))
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_position_embeddings=128):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
hidden_size=16,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=32,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(device):
|
||||
cfg = _make_config()
|
||||
m = AutoRegressiveLM(cfg).to(device=device)
|
||||
m.eval()
|
||||
return m, cfg
|
||||
|
||||
|
||||
def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128):
|
||||
return InferenceScheduler(
|
||||
model=model,
|
||||
@@ -160,14 +95,9 @@ def test_constant_reward_model_shape():
|
||||
assert torch.all(out == 0.5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def device():
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def _make_generator(device, **kw):
|
||||
model, _ = _make_model(device)
|
||||
tokenizer = FakeTokenizer()
|
||||
model, _ = make_model(device, max_position_embeddings=128)
|
||||
tokenizer = FakeTokenizer(with_chat_template=True)
|
||||
scheduler = _make_scheduler(
|
||||
model,
|
||||
tokenizer,
|
||||
@@ -229,7 +159,7 @@ def test_rollout_generator_mask_matches_responses(device):
|
||||
|
||||
|
||||
def test_rollout_generator_logprobs_are_nonpositive(device):
|
||||
"""Behaviour-policy logprobs of sampled tokens should be ≤ 0."""
|
||||
"""Behaviour-policy logprobs of sampled tokens should be <= 0."""
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
||||
batch = _make_instruction_batch(n=1)
|
||||
r = gen.generate(batch)
|
||||
@@ -241,7 +171,7 @@ def test_rollout_generator_logprobs_are_nonpositive(device):
|
||||
|
||||
|
||||
def test_rollout_generator_instruction_role_mapping(device):
|
||||
"""instruction → system, input → user, output → assistant."""
|
||||
"""instruction -> system, input -> user, output -> assistant."""
|
||||
gen, _ = _make_generator(device, group_size=1, max_tokens=4)
|
||||
batch = {
|
||||
"instruction": ["Be helpful"],
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import signal
|
||||
@@ -10,15 +9,15 @@ import torch.optim as optim
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
from astrai.parallel.signal_handler import register_signal_handlers
|
||||
from astrai.trainer import Trainer
|
||||
from astrai.trainer.schedule import SchedulerFactory
|
||||
from astrai.trainer.train_context import TrainContext
|
||||
from tests.helpers import load_checkpoint_meta, make_tiny_config
|
||||
|
||||
|
||||
class _PicklableDataset(Dataset):
|
||||
class PicklableDataset(Dataset):
|
||||
def __init__(self, length=200, max_length=64, vocab_size=1000):
|
||||
self.length = length
|
||||
self.max_length = max_length
|
||||
@@ -35,16 +34,7 @@ class _PicklableDataset(Dataset):
|
||||
|
||||
|
||||
def _build_model():
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=1000,
|
||||
hidden_size=8,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
intermediate_size=16,
|
||||
max_position_embeddings=64,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-5,
|
||||
)
|
||||
config = make_tiny_config()
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
return AutoRegressiveLM(config).to(device=device)
|
||||
|
||||
@@ -61,7 +51,7 @@ class _ReadyCallback:
|
||||
|
||||
|
||||
def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file):
|
||||
dataset = _PicklableDataset()
|
||||
dataset = PicklableDataset()
|
||||
|
||||
def model_fn():
|
||||
return _build_model()
|
||||
@@ -147,17 +137,7 @@ def test_sigterm_triggers_checkpoint_save(base_test_env):
|
||||
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGTERM)
|
||||
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
||||
|
||||
ckpt_dir = base_test_env["test_dir"]
|
||||
meta_files = []
|
||||
for root, dirs, files in os.walk(ckpt_dir):
|
||||
for f in files:
|
||||
if f == "meta.json":
|
||||
meta_files.append(os.path.join(root, f))
|
||||
|
||||
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
|
||||
|
||||
with open(meta_files[-1]) as f:
|
||||
meta = json.load(f)
|
||||
meta = load_checkpoint_meta(base_test_env["test_dir"])
|
||||
assert "consumed_samples" in meta
|
||||
assert meta["consumed_samples"] >= 0
|
||||
|
||||
@@ -167,11 +147,6 @@ def test_sigint_triggers_checkpoint_save(base_test_env):
|
||||
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGINT)
|
||||
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
||||
|
||||
ckpt_dir = base_test_env["test_dir"]
|
||||
meta_files = []
|
||||
for root, dirs, files in os.walk(ckpt_dir):
|
||||
for f in files:
|
||||
if f == "meta.json":
|
||||
meta_files.append(os.path.join(root, f))
|
||||
|
||||
assert len(meta_files) > 0, f"No checkpoint meta.json found in {ckpt_dir}"
|
||||
meta = load_checkpoint_meta(base_test_env["test_dir"])
|
||||
assert "consumed_samples" in meta
|
||||
assert meta["consumed_samples"] >= 0
|
||||
|
||||
Reference in New Issue
Block a user