File size: 2,818 Bytes
bbf30b7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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()