fugu-lite / tests /test_tiny_end_to_end.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame Contribute Delete
3.77 kB
from pathlib import Path
from tokenizers import Tokenizer
from tokenizers.models import WordLevel
from tokenizers.pre_tokenizers import Whitespace
from transformers import BertConfig, BertModel, PreTrainedTokenizerFast
from fugu_lite.config import (
ESConfig,
ESTrainingConfig,
ModelConfig,
RLConfig,
RLTrainingConfig,
SFTConfig,
SFTTrainingConfig,
)
from fugu_lite.evaluate import evaluate_checkpoint
from fugu_lite.schemas import RewardRecord
from fugu_lite.train_es import train_es
from fugu_lite.train_rl import train_rl
from fugu_lite.train_sft import train_sft
def _tiny_backbone(path: Path) -> None:
words = [
"[PAD]",
"[UNK]",
"[CLS]",
"[SEP]",
"[MASK]",
"Select",
"best",
"worker",
"Domain",
"math",
"code",
"general",
"Task",
"number",
"python",
"capital",
]
vocabulary = {word: index for index, word in enumerate(words)}
tokenizer_object = Tokenizer(WordLevel(vocabulary, unk_token="[UNK]"))
tokenizer_object.pre_tokenizer = Whitespace()
tokenizer = PreTrainedTokenizerFast(
tokenizer_object=tokenizer_object,
unk_token="[UNK]",
pad_token="[PAD]",
cls_token="[CLS]",
sep_token="[SEP]",
mask_token="[MASK]",
)
tokenizer.save_pretrained(path)
model = BertModel(
BertConfig(
vocab_size=len(vocabulary),
hidden_size=24,
num_hidden_layers=1,
num_attention_heads=4,
intermediate_size=48,
max_position_embeddings=128,
pad_token_id=vocabulary["[PAD]"],
)
)
model.save_pretrained(path, safe_serialization=True)
def _records() -> list[RewardRecord]:
rows = []
domains = ["math", "code", "general"] * 4
splits = ["train"] * 6 + ["validation"] * 3 + ["test"] * 3
for index, (domain, split) in enumerate(zip(domains, splits)):
rewards = [1.0, 0.0] if domain == "math" else [0.0, 1.0]
rows.append(
RewardRecord(
task_id=f"tiny-{index}",
prompt=f"A {domain} task number {index}",
domain=domain,
split=split,
worker_ids=["worker_a", "worker_b"],
rewards=rewards,
)
)
return rows
def test_tiny_sft_rl_es_checkpoint_cycle(tmp_path: Path):
backbone = tmp_path / "tiny-backbone"
backbone.mkdir()
_tiny_backbone(backbone)
model_config = ModelConfig(
base_model=str(backbone),
max_length=64,
dropout=0.0,
dtype="float32",
)
records = _records()
sft_dir = tmp_path / "sft"
train_sft(
records,
SFTConfig(
model=model_config,
training=SFTTrainingConfig(epochs=1, batch_size=2, log_every=100),
),
sft_dir,
)
assert (sft_dir / "router_head.safetensors").exists()
rl_dir = tmp_path / "rl"
train_rl(
records,
RLConfig(
model=model_config,
training=RLTrainingConfig(
epochs=1,
batch_size=2,
estimator="expected_reward",
log_every=100,
),
),
rl_dir,
checkpoint=str(sft_dir),
)
report = evaluate_checkpoint(records, rl_dir, split="test")
assert report["examples"] == 3
es_dir = tmp_path / "es"
train_es(
records,
ESConfig(
model=model_config,
training=ESTrainingConfig(generations=2, population_size=4, sigma=0.01),
),
es_dir,
checkpoint=str(rl_dir),
)
assert (es_dir / "training_report.json").exists()