| 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() |
|
|
|
|