test: deduplicate suites and prune low-value cases

- extract shared helpers for dataset writers, scheduler construction, thread interleaving, hf roundtrips, and moe configs
- remove about 20 cases whose only assertions were format checks, restated declarations, fake-taxonomy duplicates, or test-local scaffolding
- strengthen weak cases into exact reference comparisons, positional mask checks, and deterministic outcomes
- replace two schedule factory smoke tests with cosine/sgdr formula assertions
- delete root-level CLI tests whose merge-priority facts are covered by tests/config/test_cli.py
- suite shrinks from 857 to 826 items; ruff format, import order, and pytest all green
This commit is contained in:
2026-09-03 21:54:14 +08:00
parent 28d11f1610
commit 9d3ae76683
22 changed files with 407 additions and 1027 deletions
+87 -189
View File
@@ -26,11 +26,22 @@ from astrai.serialization import (
)
from tests.data.factories import make_grpo_config
SIMPLE_CHAT_TEMPLATE = (
"{% for message in messages %}"
"{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
def _rand_seq(length, vocab=1000):
return torch.randint(0, vocab, (length,), dtype=torch.int64)
def _dump_jsonl(path, records):
with open(path, "w", encoding="utf-8") as f:
for record in records:
f.write(json.dumps(record, ensure_ascii=False) + "\n")
def _save_test_tokenizer(test_dir, tokenizer):
tokenizer_path = os.path.join(test_dir, "tokenizer")
os.makedirs(tokenizer_path, exist_ok=True)
@@ -38,19 +49,19 @@ def _save_test_tokenizer(test_dir, tokenizer):
return tokenizer_path
def _write_jsonl_dataset(test_dir, tokenizer_path, records, config_overrides=None):
data_dir = os.path.join(test_dir, "jsonl_data")
def _write_text_dataset(
test_dir, dirname, tokenizer_path, records, config_overrides=None
):
"""Write a JSONL dataset directory with a text-section default config."""
data_dir = os.path.join(test_dir, dirname)
os.makedirs(data_dir, exist_ok=True)
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for record in records:
f.write(json.dumps(record, ensure_ascii=False) + "\n")
_dump_jsonl(os.path.join(data_dir, "data.jsonl"), records)
config = {
"tokenizer_path": tokenizer_path,
"version": 1,
"input": {"sections": [{"field": "text", "action": "train"}]},
"preprocessing": {"max_seq_len": 128},
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
"output": {"position_ids_mode": "continuous"},
}
if config_overrides:
@@ -71,6 +82,27 @@ def _fake_fetch_record(self, idx, keys):
return {k: self._data[k][idx] for k in keys}
def _grpo_fake_store(prompts, responses, masks, rewards):
"""Fake GRPO record store matching real Store semantics."""
return type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": len(prompts),
"token_count": 0,
"_data": {
"prompts": prompts,
"responses": responses,
"masks": masks,
"rewards": rewards,
},
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
},
)()
def _make_seq_dataset(
test_dir, name="data", seq_length=200, train_type="seq", data=None, **load_kwargs
):
@@ -85,32 +117,6 @@ def _make_seq_dataset(
)
def test_dataset_loader_random_paths(base_test_env):
"""Test dataset loader with multiple random paths"""
test_dir = base_test_env["test_dir"]
loaded_dataset = None
num_files = np.random.randint(2, 5)
for i in range(num_files):
seq_length = np.random.randint(200, 400)
dummy_data = {"sequence": [_rand_seq(seq_length) for _ in range(10)]}
sub_dir = os.path.join(test_dir, f"sub_{i}")
os.makedirs(sub_dir, exist_ok=True)
loaded_dataset = _make_seq_dataset(
sub_dir, f"data_{i}", seq_length, data=dummy_data
)
assert loaded_dataset is not None
assert len(loaded_dataset) > 0
# Test that we can get items without errors
for i in range(len(loaded_dataset)):
item = loaded_dataset[i]
assert "input_ids" in item
assert "target_ids" in item
assert item["input_ids"].shape == item["target_ids"].shape
assert item["input_ids"].shape[0] == 64
def test_dpo_strategy_with_random_data(base_test_env):
"""Test DPO strategy with randomized preference data"""
test_dir = base_test_env["test_dir"]
@@ -345,25 +351,12 @@ def test_normalize_mixed_empty_key():
def test_grpo_dataset_dtype(base_test_env):
"""GRPO dataset returns correct dtypes for per-record structured data."""
G = 4
store = type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": 1,
"token_count": 0,
"_data": {
"prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)],
"responses": [
[torch.randint(0, 100, (5,), dtype=torch.int32) for _ in range(G)]
],
"masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
"rewards": [torch.rand(G, dtype=torch.float32)],
},
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
},
)()
store = _grpo_fake_store(
prompts=[torch.randint(0, 100, (10,), dtype=torch.int32)],
responses=[[torch.randint(0, 100, (5,), dtype=torch.int32) for _ in range(G)]],
masks=[[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
rewards=[torch.rand(G, dtype=torch.float32)],
)
dataset = GRPODataset(store=store)
item = dataset[0]
@@ -378,23 +371,12 @@ def test_grpo_dataset_load(base_test_env):
G = 3
prompt_len = 8
resp_lens = [5, 7, 4]
store = type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": 1,
"token_count": 0,
"_data": {
"prompts": [torch.randint(0, 100, (prompt_len,))],
"responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
"masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
"rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
},
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
},
)()
store = _grpo_fake_store(
prompts=[torch.randint(0, 100, (prompt_len,))],
responses=[[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
masks=[[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
rewards=[torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
)
dataset = GRPODataset(store=store)
assert len(dataset) == 1
@@ -465,62 +447,28 @@ def test_dataset_load_explicit_storage_type(base_test_env):
assert dataset.token_count == 200
def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None):
"""Write JSONL dataset — one JSON object per line."""
data_dir = os.path.join(test_dir, "json_data")
os.makedirs(data_dir, exist_ok=True)
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
config = {
"tokenizer_path": tokenizer_path,
"version": 1,
"input": {"sections": [{"field": "text", "action": "train"}]},
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
"output": {"position_ids_mode": "continuous"},
}
if config_overrides:
config.update(config_overrides)
with open(
os.path.join(data_dir, "dataset_config.json"), "w", encoding="utf-8"
) as f:
json.dump(config, f, ensure_ascii=False, indent=2)
return data_dir
@pytest.mark.parametrize(
"use_jsonl",
[True, False],
)
@pytest.mark.parametrize("use_jsonl", [True, False])
def test_detect_format_data_dir(base_test_env, use_jsonl):
"""detect_format returns 'jsonl' for dirs of .jsonl or .json files."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
if use_jsonl:
data_dir = _write_jsonl_dataset(
test_dir,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz"}],
)
else:
data_dir = _write_json_dataset(
test_dir,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz qux"}],
)
data_dir = _write_text_dataset(
test_dir,
"jsonl_data" if use_jsonl else "json_data",
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz"}],
)
assert detect_format(data_dir) == "jsonl"
def test_json_store_seq(base_test_env):
"""JsonlStore loads .json array correctly."""
@pytest.mark.parametrize("dirname", ["json_data", "jsonl_data"])
def test_json_store_seq(base_test_env, dirname):
"""JsonlStore loads a text JSONL dataset and feeds seq training."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
data_dir = _write_json_dataset(
data_dir = _write_text_dataset(
test_dir,
dirname,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz qux"}],
)
@@ -535,15 +483,14 @@ def test_json_store_seq(base_test_env):
item = dataset[0]
assert "input_ids" in item
assert "target_ids" in item
assert item["input_ids"].dtype == torch.long
def test_json_store_no_tokenizer_path(base_test_env):
"""JsonlStore uses dataset dir as tokenizer_path when omitted."""
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
tokenizer.set_chat_template(SIMPLE_CHAT_TEMPLATE)
data_dir = os.path.join(test_dir, "self_contained")
os.makedirs(data_dir, exist_ok=True)
@@ -560,9 +507,7 @@ def test_json_store_no_tokenizer_path(base_test_env):
]
}
]
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
_dump_jsonl(os.path.join(data_dir, "data.jsonl"), records)
# dataset_config.json WITHOUT tokenizer_path
config = {
@@ -587,38 +532,14 @@ def test_json_store_no_tokenizer_path(base_test_env):
assert "loss_mask" in store.keys
def test_jsonl_store_seq(base_test_env):
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
data_dir = _write_jsonl_dataset(
test_dir,
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar baz qux"}],
config_overrides={"preprocessing": {"max_seq_len": 128, "min_chars": 0}},
)
store = StoreFactory.create("jsonl")
store.load(data_dir, transform=_build_jsonl_transform(data_dir))
assert len(store) > 0
assert "sequence" in store.keys
dataset = DatasetFactory.load("seq", data_dir, window_size=8)
assert len(dataset) > 0
item = dataset[0]
assert "input_ids" in item
assert "target_ids" in item
assert item["input_ids"].dtype == torch.long
def test_jsonl_store_sft(base_test_env):
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
tokenizer.set_chat_template(SIMPLE_CHAT_TEMPLATE)
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
data_dir = _write_jsonl_dataset(
data_dir = _write_text_dataset(
test_dir,
"sft_jsonl",
tokenizer_path,
[
{
@@ -661,9 +582,7 @@ def test_sft_jsonl_default_messages_config(base_test_env):
"""
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
tokenizer.set_chat_template(SIMPLE_CHAT_TEMPLATE)
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
data_dir = os.path.join(test_dir, "jsonl_data")
@@ -683,9 +602,7 @@ def test_sft_jsonl_default_messages_config(base_test_env):
]
},
]
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for record in records:
f.write(json.dumps(record, ensure_ascii=False) + "\n")
_dump_jsonl(os.path.join(data_dir, "data.jsonl"), records)
dataset = DatasetFactory.load(
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
@@ -706,13 +623,12 @@ def test_sft_jsonl_explicit_config_takes_priority(base_test_env):
"""When dataset_config.json exists, it overrides the default messages config."""
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer.set_chat_template(
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
)
tokenizer.set_chat_template(SIMPLE_CHAT_TEMPLATE)
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
data_dir = _write_jsonl_dataset(
data_dir = _write_text_dataset(
test_dir,
"sft_explicit",
tokenizer_path,
[
{
@@ -748,10 +664,7 @@ def _write_grpo_jsonl(test_dir, tokenizer_path, records):
"""Write a GRPO JSONL dataset directory with config."""
data_dir = os.path.join(test_dir, "grpo_jsonl")
os.makedirs(data_dir, exist_ok=True)
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
_dump_jsonl(os.path.join(data_dir, "data.jsonl"), records)
config = {
"tokenizer_path": tokenizer_path,
@@ -927,28 +840,15 @@ def test_grpo_multiple_records(base_test_env):
[torch.randint(0, 100, (np.random.randint(3, 8),)) for _ in range(G)]
for _ in range(n_records)
]
store = type(
"FakeStore",
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": n_records,
"token_count": 0,
"_data": {
"prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)],
"responses": dummy_responses,
"masks": [
[torch.ones(r.shape[0], dtype=torch.int64) for r in resps]
for resps in dummy_responses
],
"rewards": [
torch.rand(G, dtype=torch.float32) for _ in range(n_records)
],
},
"fetch_record": _fake_fetch_record,
"__len__": lambda self: self.num_records,
},
)()
store = _grpo_fake_store(
prompts=[torch.randint(0, 100, (10,)) for _ in range(n_records)],
responses=dummy_responses,
masks=[
[torch.ones(r.shape[0], dtype=torch.int64) for r in resps]
for resps in dummy_responses
],
rewards=[torch.rand(G, dtype=torch.float32) for _ in range(n_records)],
)
dataset = GRPODataset(store=store)
assert len(dataset) == n_records
@@ -965,9 +865,7 @@ def test_grpo_multiple_records(base_test_env):
def _write_dpo_jsonl(test_dir, records):
"""Write a raw DPO JSONL file (no dataset_config.json)."""
path = os.path.join(test_dir, "dpo.jsonl")
with open(path, "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
_dump_jsonl(path, records)
return path
@@ -1075,12 +973,12 @@ def test_jsonl_store_eager_len_returns_token_count(base_test_env):
"""JsonlStore in eager mode: num_records reflects per-record count."""
test_dir = base_test_env["test_dir"]
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
data_dir = _write_jsonl_dataset(
data_dir = _write_text_dataset(
test_dir,
"jsonl_data",
tokenizer_path,
[{"text": "hello world"}, {"text": "foo bar"}],
config_overrides={
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
"output": {"position_ids_mode": "none"},
},
)
+29 -92
View File
@@ -14,7 +14,6 @@ from astrai.preprocessing.builder import (
)
from tests.data.factories import (
CHAT_SECTIONS,
INSTRUCTION_SECTIONS,
TEXT_SECTIONS,
make_chat_config,
make_dpo_chat_config,
@@ -35,9 +34,6 @@ def test_chat_simple(chat_tokenizer, builder):
}
result = builder.build(item, config, chat_tokenizer)
assert result is not None
assert "sequence" in result
assert "loss_mask" in result
assert len(result["sequence"]) == len(result["loss_mask"])
ids = chat_tokenizer.decode(result["sequence"], skip_special_tokens=False)
assert "system" in ids.lower() or "<|im_start|>system" in ids
@@ -51,21 +47,26 @@ def test_chat_simple(chat_tokenizer, builder):
def test_chat_mask_only_assistant(chat_tokenizer, builder):
config = make_chat_config()
item = {
"messages": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
]
}
result = builder.build(item, config, chat_tokenizer)
messages = [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
]
result = builder.build({"messages": messages}, config, chat_tokenizer)
mask = result["loss_mask"]
ids = result["sequence"]
assert len(ids) == len(mask)
trained = [i for i, m in enumerate(mask) if m == 1]
masked = [i for i, m in enumerate(mask) if m == 0]
assert len(trained) > 0
assert len(masked) > 0
def template_ids(message):
rendered = chat_tokenizer.apply_chat_template(
[message], tokenize=False, add_generation_prompt=False
)
return chat_tokenizer.encode(rendered, add_special_tokens=False)
user_len = len(template_ids(messages[0]))
assistant_len = len(template_ids(messages[1]))
bos = 1 if chat_tokenizer.bos_token_id is not None else 0
assert len(mask) == bos + user_len + assistant_len
assert all(m == 0 for m in mask[: bos + user_len])
assert all(m == 1 for m in mask[bos + user_len :])
def test_chat_batch_matches_single(chat_tokenizer, builder):
@@ -166,14 +167,6 @@ def test_chat_truncation(chat_tokenizer, builder):
assert len(result["loss_mask"]) == len(result["sequence"])
def test_instruction_basic(test_tokenizer, builder):
config = make_instruction_config()
item = {"prompt": "Translate to French: Hello", "response": "Bonjour"}
result = builder.build(item, config, test_tokenizer)
assert result is not None
assert len(result["sequence"]) == len(result["loss_mask"])
def test_instruction_batch_matches_single(test_tokenizer, builder):
config = make_instruction_config()
items = [
@@ -224,7 +217,6 @@ def test_text_basic(test_tokenizer, builder):
item = {"text": "Hello world. This is a test document."}
result = builder.build(item, config, test_tokenizer)
assert result is not None
assert "sequence" in result
assert len(result["sequence"]) > 0
assert "loss_mask" not in result
@@ -253,72 +245,17 @@ def test_text_truncation(test_tokenizer, builder):
assert len(result["sequence"]) <= 3
def test_sectioned_chat(chat_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=CHAT_SECTIONS),
mask={"system": "mask", "user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
)
item = {
"messages": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
]
}
result = builder.build(item, config, chat_tokenizer)
assert result is not None
assert len(result["sequence"]) == len(result["loss_mask"])
assert sum(result["loss_mask"]) > 0
assert 0 in result["loss_mask"]
def test_sectioned_instruction(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=INSTRUCTION_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=0),
)
item = {"prompt": "Q: Why?", "response": "A: Because."}
result = builder.build(item, config, test_tokenizer)
assert result is not None
mask = result["loss_mask"]
assert mask[0] == 0
assert mask[-1] == 1
def test_sectioned_text(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=1),
)
item = {"text": "Hello world, this is a test."}
result = builder.build(item, config, test_tokenizer)
assert result is not None
assert "loss_mask" not in result
def test_sectioned_text_too_short(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=100),
)
assert builder.build({"text": "short"}, config, test_tokenizer) is None
def test_factory_registered():
names = MaskBuilderFactory.list_registered()
assert "single" in names
assert "multi" in names
assert "sectioned" in names
def test_factory_create():
single = MaskBuilderFactory.create("single")
assert isinstance(single, SingleOutputMaskBuilder)
multi = MaskBuilderFactory.create("multi")
assert isinstance(multi, MultiOutputMaskBuilder)
sectioned = MaskBuilderFactory.create("sectioned")
assert isinstance(sectioned, SectionedMaskBuilder)
@pytest.mark.parametrize(
("name", "builder_cls"),
[
("single", SingleOutputMaskBuilder),
("multi", MultiOutputMaskBuilder),
("sectioned", SectionedMaskBuilder),
],
)
def test_factory_create(name, builder_cls):
assert name in MaskBuilderFactory.list_registered()
assert isinstance(MaskBuilderFactory.create(name), builder_cls)
def test_dpo_chat_basic(chat_tokenizer, builder):
-43
View File
@@ -115,49 +115,6 @@ def test_full_text_pipeline(temp_dir, tokenizer_dir):
assert "loss_mask" not in meta
def test_full_instruction_pipeline(temp_dir, tokenizer_dir):
jsonl_path = os.path.join(temp_dir, "instruct.jsonl")
with open(jsonl_path, "w", encoding="utf-8") as f:
f.write(
json.dumps(
{
"prompt": "Tell me a joke",
"response": "Why did the chicken cross the road?",
}
)
+ "\n"
)
f.write(
json.dumps(
{
"prompt": "What is AI?",
"response": "Artificial Intelligence is a field of computer science.",
}
)
+ "\n"
)
config = PipelineConfig(
input=InputConfig(sections=INSTRUCTION_SECTIONS),
mask={"prompt": "mask", "response": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
output=OutputConfig(storage_format="bin"),
)
out_dir = os.path.join(temp_dir, "output")
Pipeline(
config=config,
input_paths=[jsonl_path],
output_dir=out_dir,
tokenizer_path=tokenizer_dir,
).run()
meta = load_shard_meta(out_dir)
assert "sequence" in meta
assert "loss_mask" in meta
def test_dtype_override(temp_dir, tokenizer_dir):
jsonl_path = os.path.join(temp_dir, "data.jsonl")
with open(jsonl_path, "w", encoding="utf-8") as f: