| import time | |
| from typing import Dict, List | |
| import torch | |
| from torch.utils.data import DataLoader, Dataset | |
| from modeling_xonelm import XoneLM, HardwareContext | |
| from luminav import LuminaV | |
| from tokenizer import ( | |
| build_xonelm_tokenizer, | |
| MultiTurnConversationFormatter, | |
| SpecialTokenConfig, | |
| ) | |
| class SafeSFTCollator: | |
| def __init__(self, max_seq_len: int = 512, pad_token_id: int = 0): | |
| self.max_seq_len = max_seq_len | |
| self.pad_token_id = pad_token_id | |
| def __call__(self, samples: List[Dict[str, List[int]]]) -> Dict[str, torch.Tensor]: | |
| batch_inputs = [] | |
| batch_labels = [] | |
| for item in samples: | |
| inp = item["input_ids"][: self.max_seq_len] | |
| lbl = item["labels"][: self.max_seq_len] | |
| pad_len = self.max_seq_len - len(inp) | |
| batch_inputs.append( | |
| torch.tensor(inp + [self.pad_token_id] * pad_len, dtype=torch.long) | |
| ) | |
| batch_labels.append( | |
| torch.tensor(lbl + [-100] * pad_len, dtype=torch.long) | |
| ) | |
| return { | |
| "input_ids": torch.stack(batch_inputs), | |
| "labels": torch.stack(batch_labels), | |
| } | |
| class ConversationDataset(Dataset): | |
| def __init__(self, data: List[Dict[str, List[int]]]): | |
| self.data = data | |
| def __len__(self) -> int: | |
| return len(self.data) | |
| def __getitem__(self, idx: int) -> Dict[str, List[int]]: | |
| return self.data[idx] | |
| def run_sft_demo(): | |
| device = HardwareContext.get_optimal_device() | |
| autocast_dtype = HardwareContext.get_optimal_autocast_dtype(device) | |
| print("Compute Device :", device) | |
| print("Autocast Dtype :", autocast_dtype) | |
| tokenizer = build_xonelm_tokenizer() | |
| vocab_size = len(tokenizer) | |
| token_cfg = SpecialTokenConfig( | |
| pad_token_id=0, | |
| bos_token_id=1, | |
| eos_token_id=2, | |
| unk_token_id=3, | |
| eod_token_id=4, | |
| ) | |
| formatter = MultiTurnConversationFormatter(tokenizer, token_cfg) | |
| sample_dialogues = [ | |
| [ | |
| {"role": "system", "content": "You are a precise reasoning assistant."}, | |
| {"role": "user", "content": "Lily found a wooden box. What did she open?"}, | |
| {"role": "assistant", "content": "She opened the wooden box to see what was inside."}, | |
| ], | |
| [ | |
| {"role": "system", "content": "You are a polite companion."}, | |
| {"role": "user", "content": "Hello! How can we optimize memory bandwidth?"}, | |
| {"role": "assistant", "content": "We can compress Key-Value caches using low-rank latent projections."}, | |
| ], | |
| [ | |
| {"role": "system", "content": "You are a creative writer."}, | |
| {"role": "user", "content": "Tell me a story about a kitten in the garden."}, | |
| {"role": "assistant", "content": "Once upon a time, a tiny kitten chased a butterfly across the grass."}, | |
| ], | |
| ] | |
| formatted_samples = [formatter.format_conversation(dialogue) for dialogue in sample_dialogues] | |
| dataset = ConversationDataset(formatted_samples) | |
| collator = SafeSFTCollator(max_seq_len=256, pad_token_id=token_cfg.pad_token_id) | |
| loader = DataLoader(dataset, batch_size=2, shuffle=True, collate_fn=collator) | |
| model = XoneLM( | |
| vocab_size=vocab_size, | |
| dim=512, | |
| num_layers=12, | |
| num_heads=8, | |
| kv_latent_dim=64, | |
| hub_size=512, | |
| num_specialized_hubs=12, | |
| num_terminals=32, | |
| slots_per_terminal=16, | |
| ).to(device) | |
| optimizer = LuminaV( | |
| model.parameters(), | |
| lr=2e-4, | |
| betas=(0.9, 0.999), | |
| eps=1e-8, | |
| weight_decay=1e-3, | |
| tau=0.8, | |
| buffer=2, | |
| cautious=True, | |
| execution="auto", | |
| ) | |
| use_scaler = (device.type == "cuda" and autocast_dtype == torch.float16) | |
| scaler = torch.amp.GradScaler("cuda", enabled=True) if use_scaler else None | |
| model.train() | |
| optimizer.zero_grad() | |
| start_time = time.time() | |
| for epoch in range(2): | |
| for step, batch in enumerate(loader): | |
| x = batch["input_ids"].to(device, non_blocking=True) | |
| y = batch["labels"].to(device, non_blocking=True) | |
| with HardwareContext.get_autocast_context(device): | |
| output = model(x, labels=y, is_sft=True) | |
| loss = output.loss | |
| if scaler is not None: | |
| scaler.scale(loss).backward() | |
| scaler.unscale_(optimizer) | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| else: | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) | |
| optimizer.step() | |
| optimizer.zero_grad() | |
| print(f"Epoch [{epoch+1}/2] | Step [{step+1}/{len(loader)}] | SFT Loss: {loss.item():.4f}") | |
| elapsed = time.time() - start_time | |
| print(f"[+] SFT Training Demo completed successfully in {elapsed:.2f}s!") | |
| if __name__ == "__main__": | |
| run_sft_demo() |