SNN / train_multi.py
sensen0025's picture
Upload train_multi.py with huggingface_hub
ccfdd3c verified
Raw
History Blame Contribute Delete
11.2 kB
from torch.utils.checkpoint import checkpoint
import os, time, torch, torch.nn as nn, torch.optim as optim, torch.distributed as dist
from torch.utils.data import DataLoader, DistributedSampler
from torch.nn.parallel import DistributedDataParallel as DDP
from dataset import get_or_create_tokenizer, PackedChineseDataset
from srcn_model import SRCNv3_1B
def train():
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
device = torch.device(f"cuda:{local_rank}")
C_default, M_default, K_default = 160, 384, 8
C = int(os.environ.get("SRCN_C", str(C_default)))
M = int(os.environ.get("SRCN_M", str(M_default)))
K = int(os.environ.get("SRCN_K", str(K_default)))
B = int(os.environ.get("SRCN_B", "64"))
seq_len = 512
bptt_steps = 32
pool_steps = 4
lr = 3e-4
lr_enc = 3e-4
lr_w = 5e-5
lr_w_wd = 5e-4
grad_clip = 0.3
save_interval = 1800 # 30 min
script_dir = os.path.dirname(os.path.abspath(__file__))
corpus = os.path.join(script_dir, "annotated_corpus.jsonl")
tokenizer = get_or_create_tokenizer(corpus, os.path.join(script_dir, "vocab_tokenizer_v3.pkl"))
V_size = tokenizer.vocab_size
dataset = PackedChineseDataset(corpus, tokenizer, chunk_len=seq_len, cache_path=os.path.join(script_dir, "packed_dataset_340m.pkl"))
sampler = DistributedSampler(dataset, shuffle=True)
loader = DataLoader(dataset, batch_size=B, sampler=sampler, drop_last=True, num_workers=2, pin_memory=True)
if rank == 0:
print(f"World: {world_size} GPUs | B/GPU: {B} | Eff B: {B*world_size} | bptt: {bptt_steps}")
print(f"Dataset: {len(dataset)} chunks | Batches/epoch: {len(loader)}")
model = SRCNv3_1B(vocab_size=V_size, num_columns=C, neurons_per_column=M, num_partners=K, num_motor_pool_steps=pool_steps).to(device)
model = DDP(model, device_ids=[local_rank], find_unused_parameters=False)
total = sum(p.numel() for p in model.module.parameters())
if rank == 0:
print(f"Params: {total:,} ({total/1e9:.3f}B) | C={C} M={M} K={K}")
motor_start = model.module.motor_start_col
num_motor = model.module.num_motor_neurons
# 3 param groups: encoder (high LR), W_raw (slow), MLP (no wd)
enc_params = []
w_raw_params = []
mlp_params = []
for name, param in model.named_parameters():
if 'W_raw' in name:
w_raw_params.append(param)
elif 'vocab_head' in name:
mlp_params.append(param)
else:
enc_params.append(param)
opt = optim.AdamW([
{'params': enc_params, 'lr': lr_enc, 'weight_decay': 1e-4},
{'params': w_raw_params, 'lr': lr_w, 'weight_decay': lr_w_wd},
{'params': mlp_params, 'lr': lr, 'weight_decay': 0.0},
])
if rank == 0:
print(f"LR: enc={lr_enc}, mlp={lr}, W_raw={lr_w}(wd={lr_w_wd})")
criterion = nn.CrossEntropyLoss(ignore_index=tokenizer.pad_id)
V_th_persist = None
ckpt_path = os.path.join(script_dir, "checkpoint.pt")
start_epoch = 0
start_batch = 0
if os.path.exists(ckpt_path):
ckpt = torch.load(ckpt_path, map_location=device, weights_only=True)
model.module.load_state_dict(ckpt["model"])
opt.load_state_dict(ckpt["optimizer"])
start_epoch = ckpt["epoch"]
start_batch = ckpt.get("batch_idx", 0)
if "V_th_persist" in ckpt:
V_th_persist = ckpt["V_th_persist"].to(device)
if rank == 0:
print(f"Resumed from epoch {start_epoch+1} batch {start_batch}")
if rank == 0:
print("Starting training...\n")
torch.cuda.synchronize()
t0 = time.time()
last_save_time = time.time()
total_tokens = 0
num_epochs = 20
for epoch in range(start_epoch, num_epochs):
sampler.set_epoch(epoch)
epoch_loss = 0.0
n_batches = 0
for batch_idx, batch_x in enumerate(loader):
# Skip batches already processed before checkpoint
if batch_idx < start_batch:
continue
start_batch = 0 # reset after first epoch
batch_x = batch_x.to(device)
# Precompute W once per batch (saves 377MB MP16 recompute per timestep)
W_fp16 = model.module.precompute_W()
S = torch.zeros(B, C, M, device=device)
V = torch.zeros(B, C, M, device=device)
if V_th_persist is None or V_th_persist.shape[0] != B:
V_th_persist = torch.full((B, C, M), 2.0, device=device)
V_th = V_th_persist.detach()
I_ampa = torch.zeros(B, C, M, device=device)
I_nmda = torch.zeros(B, C, M, device=device)
I_psc = torch.zeros(B, num_motor, device=device)
window_loss = 0.0
n_windows = 0
total_spikes = 0.0
total_steps = 0
t_ranges = list(range(0, seq_len - 1, bptt_steps))
for wi, t_start in enumerate(t_ranges):
t_end = min(t_start + bptt_steps, seq_len - 1)
win_tokens = batch_x[:, t_start:t_end]
target = batch_x[:, t_start + 1:t_end + 1]
opt.zero_grad(set_to_none=True)
W_fp16 = model.module.precompute_W()
# Flat iteration: all tokens in window processed sequentially
pooled_list = []
num_win_tokens = win_tokens.shape[1]
for t in range(num_win_tokens):
token = win_tokens[:, t]
ts_start = t_start + t * pool_steps
ts_start_tensor = torch.tensor(ts_start, device=device)
S, V, V_th, I_ampa, I_nmda, I_psc, pooled_token, spikes_sum = checkpoint(
model.module.forward_token_with_psc,
S, V, V_th, I_ampa, I_nmda, I_psc, token, ts_start_tensor, W_fp16,
use_reentrant=True
)
total_spikes += spikes_sum.item()
total_steps += pool_steps
pooled_list.append(pooled_token)
window_pooled = torch.stack(pooled_list, dim=1)
logits = model.module.vocab_head(window_pooled)
loss = criterion(logits.view(-1, V_size), target.reshape(-1))
if not torch.isfinite(loss):
# Reset all state to zero — detach cuts old graph
S = torch.zeros(B, C, M, device=device)
V = torch.zeros(B, C, M, device=device)
I_ampa = torch.zeros(B, C, M, device=device)
I_nmda = torch.zeros(B, C, M, device=device)
I_psc = torch.zeros(B, num_motor, device=device)
# If V_th_persist is NaN, reinit from scratch
if not torch.isfinite(V_th_persist).all():
V_th_persist = torch.full((B, C, M), 2.0, device=device)
V_th = V_th_persist.detach()
# DO NOT backward — NaN grads poison DDP bucket
# Old graph freed when loss is overwritten next iteration
opt.zero_grad(set_to_none=True)
window_loss += float('nan')
n_windows += 1
if rank == 0:
t = time.time() - t0
print(f" [!] NaN/Inf at E{epoch+1}B{batch_idx+1:04d} W{wi} | VRAM: {torch.cuda.memory_allocated()/1e9:.2f}GB | time={t:.0f}s")
continue
if wi < len(t_ranges) - 1:
with model.no_sync():
loss.backward()
else:
loss.backward()
# Check for NaN/Inf in gradients to prevent parameter poisoning
grad_ok = True
for name, param in model.named_parameters():
if param.grad is not None:
if not torch.isfinite(param.grad).all():
grad_ok = False
if rank == 0:
print(f" [!] Infinite/NaN gradient in {name} at E{epoch+1}B{batch_idx+1:04d} W{wi}")
break
if not grad_ok:
opt.zero_grad(set_to_none=True)
S = torch.zeros(B, C, M, device=device)
V = torch.zeros(B, C, M, device=device)
I_ampa = torch.zeros(B, C, M, device=device)
I_nmda = torch.zeros(B, C, M, device=device)
I_psc = torch.zeros(B, num_motor, device=device)
V_th = V_th_persist.detach()
continue
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=grad_clip)
opt.step()
S, V, V_th, I_ampa, I_nmda, I_psc = [x.detach() for x in [S, V, V_th, I_ampa, I_nmda, I_psc]]
window_loss += loss.item()
n_windows += 1
avg = window_loss / max(n_windows, 1)
epoch_loss += avg
n_batches += 1
total_tokens += B * seq_len * world_size
elapsed = time.time() - t0
V_th_persist = V_th.detach()
cur_mem = torch.cuda.memory_allocated() / 1e9
max_mem = torch.cuda.max_memory_allocated() / 1e9
if rank == 0:
spike_rate = total_spikes / max(total_steps * B * C * M, 1) if total_steps > 0 else 0.0
print(f"E{epoch+1}B{batch_idx+1:04d} | Loss: {avg:.6f} | Tok/s: {total_tokens/elapsed:.0f} | VRAM: {cur_mem:.2f}/{max_mem:.2f}GB | SR: {spike_rate:.4f}")
# Periodic checkpoint
if rank == 0 and time.time() - last_save_time > save_interval:
torch.save({"model": model.module.state_dict(), "optimizer": opt.state_dict(),
"epoch": epoch, "batch_idx": batch_idx + 1,
"V_th_persist": V_th_persist.detach().cpu()}, ckpt_path)
last_save_time = time.time()
print(f" [Checkpoint saved at {time.time()-t0:.0f}s] E{epoch+1}B{batch_idx+1}")
# Defragment CUDA allocator every 100 batches
if batch_idx > 0 and batch_idx % 100 == 0:
torch.cuda.empty_cache()
avg_epoch = epoch_loss / max(n_batches, 1)
if rank == 0:
print(f"\n=== Epoch {epoch+1} done | Avg loss: {avg_epoch:.4f} | Elapsed: {time.time()-t0:.0f}s ===\n")
torch.save({"model": model.module.state_dict(), "optimizer": opt.state_dict(),
"epoch": epoch, "batch_idx": 0,
"V_th_persist": V_th_persist.detach().cpu()}, ckpt_path)
last_save_time = time.time()
start_batch = 0 # Ensure start_batch is reset if the previous epoch was completed early or skipped
dist.destroy_process_group()
if rank == 0:
print(f"Done in {time.time()-t0:.0f}s | Total tokens: {total_tokens:,}")
if __name__ == "__main__":
train()