File size: 6,737 Bytes
9e874de dd6cefc 9e874de dd6cefc 9e874de dd6cefc 9e874de dd6cefc 9e874de | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | """Tests for Modal LoRA fine-tuning scaffolding helpers."""
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from scripts import finetune_lora
class FakeTokenizer:
pad_token = "<pad>"
eos_token = "</s>"
def apply_chat_template(
self,
messages: list[dict[str, str]],
*,
tokenize: bool,
add_generation_prompt: bool,
) -> str:
text = "".join(
f"<{message['role']}>{message['content']}</{message['role']}>"
for message in messages
)
if add_generation_prompt:
text += "<assistant>"
return text
def __call__(
self,
text: str,
*,
truncation: bool,
max_length: int,
padding: bool,
add_special_tokens: bool = False,
) -> dict[str, list[int]]:
del padding, add_special_tokens
ids = [ord(character) % 251 + 1 for character in text]
if truncation:
ids = ids[:max_length]
return {"input_ids": ids, "attention_mask": [1] * len(ids)}
def _valid_record() -> dict[str, object]:
return {
"id": "sft-preview-0001",
"messages": [
{"role": "system", "content": "You are Objectverse Diary."},
{"role": "user", "content": "Create a persona."},
{"role": "assistant", "content": "{\"persona\": {}, \"diary\": {}}"},
],
}
class FinetuneLoraToolingTest(unittest.TestCase):
def test_load_sft_records_rejects_missing_messages(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / "bad.jsonl"
path.write_text(json.dumps({"id": "bad"}) + "\n", encoding="utf-8")
with self.assertRaises(ValueError):
finetune_lora.load_sft_records(path)
def test_load_sft_records_rejects_malformed_messages(self) -> None:
bad_record = {"id": "bad", "messages": [{"role": "user"}]}
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / "bad.jsonl"
path.write_text(json.dumps(bad_record) + "\n", encoding="utf-8")
with self.assertRaises(ValueError):
finetune_lora.load_sft_records(path)
def test_record_to_training_text_is_non_empty(self) -> None:
text = finetune_lora.record_to_training_text(_valid_record())
self.assertIn("system:", text)
self.assertIn("user:", text)
self.assertIn("assistant:", text)
self.assertIn("Objectverse Diary", text)
def test_default_training_config_uses_safe_qwen_lora_defaults(self) -> None:
config = finetune_lora.TrainingConfig()
self.assertEqual(config.base_model, "Qwen/Qwen2.5-1.5B-Instruct")
self.assertEqual(config.lora_r, 16)
self.assertEqual(config.lora_alpha, 32)
self.assertEqual(config.lora_dropout, 0.05)
self.assertEqual(config.max_steps, 80)
self.assertEqual(config.num_train_epochs, 3.0)
self.assertEqual(config.per_device_train_batch_size, 1)
self.assertEqual(config.gradient_accumulation_steps, 4)
self.assertEqual(config.eval_ratio, 0.1)
self.assertTrue(config.assistant_only_loss)
self.assertIn("q_proj", config.target_modules)
self.assertIn("down_proj", config.target_modules)
def test_training_config_serializes_v2_experiment_settings(self) -> None:
config = finetune_lora.TrainingConfig(
max_steps=0,
num_train_epochs=4.0,
per_device_train_batch_size=2,
gradient_accumulation_steps=8,
eval_ratio=0.2,
eval_steps=25,
lora_r=32,
lora_alpha=64,
assistant_only_loss=False,
)
payload = config.as_remote_dict()
self.assertEqual(payload["num_train_epochs"], 4.0)
self.assertEqual(payload["per_device_train_batch_size"], 2)
self.assertEqual(payload["gradient_accumulation_steps"], 8)
self.assertEqual(payload["eval_ratio"], 0.2)
self.assertEqual(payload["eval_steps"], 25)
self.assertEqual(payload["lora_r"], 32)
self.assertEqual(payload["lora_alpha"], 64)
self.assertFalse(payload["assistant_only_loss"])
def test_dry_run_does_not_call_remote_runner(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / "records.jsonl"
path.write_text(json.dumps(_valid_record()) + "\n", encoding="utf-8")
def fail_remote_call(
records: list[dict[str, object]],
config: finetune_lora.TrainingConfig,
) -> dict[str, object]:
raise AssertionError("dry-run should not call remote training")
summary = finetune_lora.run_training_entrypoint(
dataset=path,
config=finetune_lora.TrainingConfig(),
dry_run=True,
allow_remote=False,
remote_runner=fail_remote_call,
)
self.assertEqual(summary["mode"], "dry-run")
self.assertEqual(summary["record_count"], 1)
self.assertEqual(summary["base_model"], "Qwen/Qwen2.5-1.5B-Instruct")
self.assertEqual(summary["train_record_count"], 1)
self.assertEqual(summary["eval_record_count"], 0)
def test_dry_run_reports_eval_split_for_larger_datasets(self) -> None:
records = [_valid_record() for _ in range(20)]
summary = finetune_lora._dry_run_summary(
Path("records.jsonl"),
records,
finetune_lora.TrainingConfig(eval_ratio=0.2),
)
self.assertEqual(summary["train_record_count"], 16)
self.assertEqual(summary["eval_record_count"], 4)
def test_assistant_only_tokenization_masks_prompt_labels(self) -> None:
tokenized = finetune_lora._tokenize_training_example(
_valid_record(),
FakeTokenizer(),
max_length=512,
assistant_only_loss=True,
)
labels = tokenized["labels"]
self.assertIn(-100, labels)
self.assertTrue(any(label != -100 for label in labels))
first_unmasked = next(index for index, label in enumerate(labels) if label != -100)
self.assertGreater(first_unmasked, 0)
def test_full_loss_tokenization_keeps_all_labels(self) -> None:
tokenized = finetune_lora._tokenize_training_example(
_valid_record(),
FakeTokenizer(),
max_length=512,
assistant_only_loss=False,
)
self.assertNotIn(-100, tokenized["labels"])
if __name__ == "__main__":
unittest.main()
|