- 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
87 lines
2.9 KiB
Python
87 lines
2.9 KiB
Python
import json
|
|
import os
|
|
|
|
import pytest
|
|
import safetensors.torch as st
|
|
import torch
|
|
|
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
|
from astrai.model.transformer import AutoRegressiveLM
|
|
from tests.helpers import TINY_CONFIG
|
|
|
|
|
|
def test_tie_weight_init(base_test_env):
|
|
config_path = base_test_env["config_path"]
|
|
|
|
config_data = TINY_CONFIG.copy()
|
|
config_data["tie_word_embeddings"] = True
|
|
|
|
with open(config_path, "w") as f:
|
|
json.dump(config_data, f)
|
|
|
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
|
model = AutoRegressiveLM(config)
|
|
|
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
|
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
|
|
|
|
original_weight = model.embed_tokens.weight.clone()
|
|
model.embed_tokens.weight.data[0, 0] = 100.0
|
|
|
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
|
assert not torch.equal(model.lm_head.weight, original_weight)
|
|
|
|
config_data["tie_word_embeddings"] = False
|
|
|
|
with open(config_path, "w") as f:
|
|
json.dump(config_data, f)
|
|
|
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
|
model = AutoRegressiveLM(config)
|
|
|
|
assert not torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
|
assert model.lm_head.weight.data_ptr() != model.embed_tokens.weight.data_ptr()
|
|
|
|
original_weight = model.embed_tokens.weight.clone()
|
|
model.embed_tokens.weight.data[0, 0] = 100.0
|
|
|
|
assert not torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
|
assert not torch.equal(model.lm_head.weight, original_weight)
|
|
|
|
|
|
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 = TINY_CONFIG.copy()
|
|
config_data["tie_word_embeddings"] = True
|
|
config_path = os.path.join(test_dir, "config.json")
|
|
|
|
with open(config_path, "w") as f:
|
|
json.dump(config_data, f)
|
|
|
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
|
original_model = AutoRegressiveLM(config)
|
|
|
|
st.save_file(original_model.state_dict(), model_path)
|
|
|
|
loaded_config = AutoRegressiveLMConfig.from_file(config_path)
|
|
model = AutoRegressiveLM(loaded_config)
|
|
model.load_state_dict(st.load_file(model_path))
|
|
|
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
|
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
|
|
assert "lm_head.weight" not in model.state_dict()
|
|
|
|
config_data["tie_word_embeddings"] = False
|
|
with open(config_path, "w") as f:
|
|
json.dump(config_data, f)
|
|
|
|
loaded_config = AutoRegressiveLMConfig.from_file(config_path)
|
|
model = AutoRegressiveLM(loaded_config)
|
|
model.load_state_dict(st.load_file(model_path))
|
|
|
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
|
assert model.lm_head.weight.data_ptr() != model.embed_tokens.weight.data_ptr()
|
|
assert "lm_head.weight" in model.state_dict()
|