XoneLM-1.0-Paper / train_example.py
cloverxion's picture
feat: publish XoneLM architecture and LuminaV optimizer
bbf30b7 verified
Raw
History Blame Contribute Delete
2.82 kB
import time
import torch
import torch.nn.functional as F
from modeling_xonelm import XoneLM, HardwareContext, create_universal_document_boundary_mask
from luminav import LuminaV
from tokenizer import build_xonelm_tokenizer
def run_train_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)
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,
chunk_size=1024,
).to(device)
total_params = sum(p.numel() for p in model.parameters())
print(f"Total Parameters: {total_params / 1e6:.2f}M")
optimizer = LuminaV(
model.parameters(),
lr=8e-4,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=8e-2,
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
batch_size = 2
seq_len = 512
num_steps = 5
model.train()
optimizer.zero_grad()
start_time = time.time()
for step in range(num_steps):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
y = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
doc_mask = create_universal_document_boundary_mask(
x_tokens=x,
hub_size=model.hub_size,
past_k_len=model.hub_size,
eod_token_id=4,
is_dense_with_hub=True,
)
with HardwareContext.get_autocast_context(device):
output = model(x, labels=y, attn_mask=doc_mask)
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"Step [{step+1}/{num_steps}] | Loss: {loss.item():.4f} | Z-Loss: {output.z_loss.item():.4f}")
elapsed = time.time() - start_time
print(f"Demo training completed in {elapsed:.2f}s!")
if __name__ == "__main__":
run_train_demo()