Scott/Codex commited on
Commit Β·
78c46ca
1
Parent(s): 979693c
Add structured sublinear masks and trainer heartbeat
Browse files- README.md +9 -0
- dblocks_train.py +5 -5
- nB300_agillm4_vram_dblock.py +126 -13
- relaunch_agillm4_dblock.sh +2 -1
README.md
CHANGED
|
@@ -46,6 +46,8 @@ whose released code is ViT/classification only.
|
|
| 46 |
- SAT now uses fused vocab-streaming CE in the dblock path, and the dblock step releases AR/SAT activations before moving to the next objective.
|
| 47 |
- DBlock now uses loss-balanced block scheduling after warmup, per-block EMA diagnostics, sigma-range curriculum, objective weights, and peak VRAM logging.
|
| 48 |
- The folded-in DBlock path now builds the dense causal/SAT masks once per objective instead of once per layer, and NAT obeys `--nat_max_tokens` so long-context AR does not force full-context NAT memory.
|
|
|
|
|
|
|
| 49 |
|
| 50 |
## Honest findings
|
| 51 |
- DiffusionBlocks and gradient-checkpointing are **substitutes** for activation
|
|
@@ -67,4 +69,11 @@ per-block loss/VRAM logging, single-build masks per objective, and NAT token cap
|
|
| 67 |
These are meant to preserve the VRAM breakthrough while making block-wise training
|
| 68 |
less brittle over long runs.
|
| 69 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
License: Apache-2.0 (matching the upstream method).
|
|
|
|
| 46 |
- SAT now uses fused vocab-streaming CE in the dblock path, and the dblock step releases AR/SAT activations before moving to the next objective.
|
| 47 |
- DBlock now uses loss-balanced block scheduling after warmup, per-block EMA diagnostics, sigma-range curriculum, objective weights, and peak VRAM logging.
|
| 48 |
- The folded-in DBlock path now builds the dense causal/SAT masks once per objective instead of once per layer, and NAT obeys `--nat_max_tokens` so long-context AR does not force full-context NAT memory.
|
| 49 |
+
- Sublinear attention now supports structured causal, SAT block-causal, and unrestricted/NAT rules directly, and computes ALiBi only for gathered local/anchor candidates instead of allocating dense `[H x T x T]` bias.
|
| 50 |
+
- The trainer prints lightweight heartbeat lines and clears the CUDA cache after checkpoint load so reserved VRAM does not stay inflated by transient load tensors.
|
| 51 |
|
| 52 |
## Honest findings
|
| 53 |
- DiffusionBlocks and gradient-checkpointing are **substitutes** for activation
|
|
|
|
| 69 |
These are meant to preserve the VRAM breakthrough while making block-wise training
|
| 70 |
less brittle over long runs.
|
| 71 |
|
| 72 |
+
Structured-mask update 2026-05-29: the sublinear backend now accepts symbolic causal,
|
| 73 |
+
SAT block-causal, and unrestricted/NAT mask rules. This removes dense O(T^2) mask
|
| 74 |
+
allocation for long context, and also gathers ALiBi bias directly for selected
|
| 75 |
+
local/anchor keys instead of materializing dense `[heads x T x T]` bias tensors.
|
| 76 |
+
A trainer heartbeat and post-checkpoint CUDA cache clear were added for easier
|
| 77 |
+
long-running Vast monitoring.
|
| 78 |
+
|
| 79 |
License: Apache-2.0 (matching the upstream method).
|
dblocks_train.py
CHANGED
|
@@ -149,13 +149,13 @@ def _dblock_step(core, ar_h, sat_h, nat_h, opt, scaler, args, ids, state):
|
|
| 149 |
nat_val = 0.0
|
| 150 |
|
| 151 |
if ar_weight > 0.0:
|
| 152 |
-
causal = M.causal_mask(T)
|
| 153 |
with M.amp(args.amp):
|
| 154 |
emb = core.emb(ids)
|
| 155 |
zt = emb + sig[:, None, None] * torch.randn_like(emb)
|
| 156 |
h = ci * zt
|
| 157 |
for li in layers:
|
| 158 |
-
h = _ck.checkpoint(core.blocks[li]
|
| 159 |
Dn = core.ln(cs * zt + co * h)
|
| 160 |
ar = ar_weight * w * fused_ce(Dn[:, :-1].contiguous(), ar_h.proj.weight, ids[:, 1:].contiguous())
|
| 161 |
ar_val = float(ar.detach())
|
|
@@ -166,13 +166,13 @@ def _dblock_step(core, ar_h, sat_h, nat_h, opt, scaler, args, ids, state):
|
|
| 166 |
int(getattr(args, "sat_every", 1)) <= 1 or ((int(state.get("step", 0)) + 1) % int(getattr(args, "sat_every", 1)) == 0)
|
| 167 |
)
|
| 168 |
if sat_weight > 0.0 and do_sat:
|
| 169 |
-
smask = M.sat_mask(T)
|
| 170 |
with M.amp(args.amp):
|
| 171 |
emb2 = core.emb(ids)
|
| 172 |
zt2 = emb2 + sig[:, None, None] * torch.randn_like(emb2)
|
| 173 |
h2 = ci * zt2
|
| 174 |
for li in layers:
|
| 175 |
-
h2 = _ck.checkpoint(core.blocks[li]
|
| 176 |
Ds = core.ln(cs * zt2 + co * h2)
|
| 177 |
last = Ds[:, -SATB:]
|
| 178 |
satf = fused_ce(last.contiguous(), sat_h.proj.weight, ids[:, 1 : SATB + 1].contiguous())
|
|
@@ -211,7 +211,7 @@ def _dblock_step(core, ar_h, sat_h, nat_h, opt, scaler, args, ids, state):
|
|
| 211 |
nat_in[m] = M.BLANK
|
| 212 |
hn = core.emb(nat_in)
|
| 213 |
for li in layers:
|
| 214 |
-
hn = _ck.checkpoint(core.blocks[li]
|
| 215 |
Dnat = core.ln(hn)
|
| 216 |
nat = nat_weight * fused_ce(Dnat[m], nat_h.proj.weight, nat_ids[m])
|
| 217 |
nat_val = float(nat.detach())
|
|
|
|
| 149 |
nat_val = 0.0
|
| 150 |
|
| 151 |
if ar_weight > 0.0:
|
| 152 |
+
causal = M.causal_mask(T, structured=M.use_structured_masks(args))
|
| 153 |
with M.amp(args.amp):
|
| 154 |
emb = core.emb(ids)
|
| 155 |
zt = emb + sig[:, None, None] * torch.randn_like(emb)
|
| 156 |
h = ci * zt
|
| 157 |
for li in layers:
|
| 158 |
+
h = _ck.checkpoint(lambda y, block=core.blocks[li]: block(y, causal), h, use_reentrant=False)
|
| 159 |
Dn = core.ln(cs * zt + co * h)
|
| 160 |
ar = ar_weight * w * fused_ce(Dn[:, :-1].contiguous(), ar_h.proj.weight, ids[:, 1:].contiguous())
|
| 161 |
ar_val = float(ar.detach())
|
|
|
|
| 166 |
int(getattr(args, "sat_every", 1)) <= 1 or ((int(state.get("step", 0)) + 1) % int(getattr(args, "sat_every", 1)) == 0)
|
| 167 |
)
|
| 168 |
if sat_weight > 0.0 and do_sat:
|
| 169 |
+
smask = M.sat_mask(T, structured=M.use_structured_masks(args))
|
| 170 |
with M.amp(args.amp):
|
| 171 |
emb2 = core.emb(ids)
|
| 172 |
zt2 = emb2 + sig[:, None, None] * torch.randn_like(emb2)
|
| 173 |
h2 = ci * zt2
|
| 174 |
for li in layers:
|
| 175 |
+
h2 = _ck.checkpoint(lambda y, block=core.blocks[li]: block(y, smask), h2, use_reentrant=False)
|
| 176 |
Ds = core.ln(cs * zt2 + co * h2)
|
| 177 |
last = Ds[:, -SATB:]
|
| 178 |
satf = fused_ce(last.contiguous(), sat_h.proj.weight, ids[:, 1 : SATB + 1].contiguous())
|
|
|
|
| 211 |
nat_in[m] = M.BLANK
|
| 212 |
hn = core.emb(nat_in)
|
| 213 |
for li in layers:
|
| 214 |
+
hn = _ck.checkpoint(lambda y, block=core.blocks[li]: block(y, None), hn, use_reentrant=False)
|
| 215 |
Dnat = core.ln(hn)
|
| 216 |
nat = nat_weight * fused_ce(Dnat[m], nat_h.proj.weight, nat_ids[m])
|
| 217 |
nat_val = float(nat.detach())
|
nB300_agillm4_vram_dblock.py
CHANGED
|
@@ -1194,6 +1194,48 @@ def alibi_bias(n_heads: int, n_tokens: int):
|
|
| 1194 |
dist = (j - i).clamp_min(0)
|
| 1195 |
return -_alibi_slopes(n_heads) * dist
|
| 1196 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1197 |
# βββββββββββββββββββββββββ Model components βββββββββββββββββββββββββ
|
| 1198 |
class KVBuffer:
|
| 1199 |
"""Preallocated K/V cache for decode. Replaces torch.cat-based growth.
|
|
@@ -1340,7 +1382,20 @@ class TuneableAttentionMHA(nn.Module):
|
|
| 1340 |
self._metric_cache_shape = (-1, -1)
|
| 1341 |
return super().train(mode)
|
| 1342 |
|
| 1343 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1344 |
"""Local-window + landmark attention: O(N * (window + N/stride))."""
|
| 1345 |
bsz, heads, q_len, _ = q.shape
|
| 1346 |
k_len = k.size(2)
|
|
@@ -1348,6 +1403,9 @@ class TuneableAttentionMHA(nn.Module):
|
|
| 1348 |
query_base = max(0, k_len - q_len)
|
| 1349 |
outputs = []
|
| 1350 |
scale = 1.0 / math.sqrt(self.dk)
|
|
|
|
|
|
|
|
|
|
| 1351 |
|
| 1352 |
anchor_start = self.sublinear_stride - 1
|
| 1353 |
if self.sublinear_stride > 0 and self.sublinear_max_anchors > 0 and anchor_start < k_len:
|
|
@@ -1388,10 +1446,18 @@ class TuneableAttentionMHA(nn.Module):
|
|
| 1388 |
idx = local_idx
|
| 1389 |
valid = local_valid
|
| 1390 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1391 |
k_sel = k[:, :, idx, :]
|
| 1392 |
scores = (q[:, :, q_start:q_end, :].unsqueeze(-2) * k_sel).sum(dim=-1) * scale
|
| 1393 |
|
| 1394 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1395 |
mask_q = attn_mask[..., q_start:q_end, :]
|
| 1396 |
gather_idx = idx.view(1, 1, cur, -1).expand(mask_q.size(0), mask_q.size(1), cur, idx.size(1))
|
| 1397 |
scores = scores + torch.gather(mask_q, -1, gather_idx)
|
|
@@ -1428,7 +1494,9 @@ class TuneableAttentionMHA(nn.Module):
|
|
| 1428 |
else:
|
| 1429 |
k, v = k_new, v_new
|
| 1430 |
attn_mask = mask
|
| 1431 |
-
if self.
|
|
|
|
|
|
|
| 1432 |
rel = alibi_bias(self.h, rel_bias_tokens)[:, :, -q.size(2):, :]
|
| 1433 |
attn_mask = rel if attn_mask is None else attn_mask + rel
|
| 1434 |
if self.attn_backend == "sdpa":
|
|
@@ -1445,7 +1513,7 @@ class TuneableAttentionMHA(nn.Module):
|
|
| 1445 |
q_scaled = q * math.sqrt(q.size(-1) / self.dk)
|
| 1446 |
z = F.scaled_dot_product_attention(q_scaled, k, v, attn_mask=attn_mask, dropout_p=0.0)
|
| 1447 |
elif self.attn_backend == "sublinear":
|
| 1448 |
-
z = self._sublinear_attention(q, k, v, attn_mask=attn_mask)
|
| 1449 |
else:
|
| 1450 |
att = (q @ k.transpose(-1, -2)) / math.sqrt(self.dk)
|
| 1451 |
if attn_mask is not None:
|
|
@@ -1560,7 +1628,7 @@ class Encoder(nn.Module):
|
|
| 1560 |
if not use_cache:
|
| 1561 |
for i, blk in enumerate(self.blocks):
|
| 1562 |
if self.grad_checkpoint and self.training:
|
| 1563 |
-
x = torch_checkpoint.checkpoint(
|
| 1564 |
else:
|
| 1565 |
x = blk(x, mask)
|
| 1566 |
if self.anchor is not None and i == self.anchor_position:
|
|
@@ -1622,17 +1690,23 @@ class SATHead(nn.Module):
|
|
| 1622 |
|
| 1623 |
|
| 1624 |
# βββββββββββββββββββββββββ Masks βββββββββββββββββββββββββ
|
| 1625 |
-
def causal_mask(n):
|
|
|
|
|
|
|
| 1626 |
return torch.triu(torch.full((1, 1, n, n), float("-inf"), device=DEV), 1)
|
| 1627 |
|
| 1628 |
-
def sat_mask(n, block=SAT_BLOCK):
|
|
|
|
|
|
|
| 1629 |
idx = torch.arange(n, device=DEV)
|
| 1630 |
grp = idx.unsqueeze(0) // block
|
| 1631 |
allow = (grp.T == grp) | (grp.T > grp)
|
| 1632 |
return torch.where(allow, 0.0, float("-inf")).unsqueeze(0).unsqueeze(0)
|
| 1633 |
|
| 1634 |
-
def sat_mask_cached(new_len: int, cached_len: int, block=SAT_BLOCK):
|
| 1635 |
total_len = cached_len + new_len
|
|
|
|
|
|
|
| 1636 |
q_idx = torch.arange(cached_len, total_len, device=DEV).unsqueeze(1)
|
| 1637 |
k_idx = torch.arange(total_len, device=DEV).unsqueeze(0)
|
| 1638 |
q_grp = q_idx // block
|
|
@@ -2156,6 +2230,7 @@ def _train_phase(
|
|
| 2156 |
now_wall = time.time()
|
| 2157 |
last_save_mono = time.monotonic() - (now_wall - (resume_wall_time or now_wall))
|
| 2158 |
last_delta_step = start_step
|
|
|
|
| 2159 |
print(f"[{phase_name}] Starting. Goal: {total_tokens_needed:,} tokens. Batch={BATCH}, Block={BLOCK}")
|
| 2160 |
print(
|
| 2161 |
f"[{phase_name}] AR_ONLY={args.ar_only}, SAT_EVERY={args.sat_every}, "
|
|
@@ -2171,6 +2246,19 @@ def _train_phase(
|
|
| 2171 |
except (ValueError, OSError):
|
| 2172 |
pass
|
| 2173 |
_DBS = _dblock_init(core, args) if getattr(args,'dblock',False) else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2174 |
while seen_tok < total_tokens_needed:
|
| 2175 |
try:
|
| 2176 |
while len(buf) < BLOCK:
|
|
@@ -2190,7 +2278,7 @@ def _train_phase(
|
|
| 2190 |
loss_value = _dblock_step(core, ar_h, sat_h, nat_h, opt, scaler, args, ids, _DBS)
|
| 2191 |
else:
|
| 2192 |
with amp(args.amp):
|
| 2193 |
-
h_ar = core(ids, causal_mask(ids.size(1)))
|
| 2194 |
logits_ar = ar_h(h_ar)[:, :-1]
|
| 2195 |
loss_ar = ce_tok(logits_ar.reshape(-1, VOCAB), tgt_ar[:, 1:].reshape(-1))
|
| 2196 |
loss_value = float(loss_ar.detach().item())
|
|
@@ -2201,7 +2289,7 @@ def _train_phase(
|
|
| 2201 |
# Same AR+SAT objective as a summed loss, but sequential backward keeps
|
| 2202 |
# only one core-forward activation graph live at a time on 24GB cards.
|
| 2203 |
with amp(args.amp):
|
| 2204 |
-
h_sat = core(ids, sat_mask(ids.size(1)))
|
| 2205 |
logits_sat, gate = sat_h(h_sat[:, -SAT_BLOCK:])
|
| 2206 |
tgt_sat = ids[:, 1:SAT_BLOCK+1]
|
| 2207 |
loss_sat = ce_tok(logits_sat.reshape(-1, VOCAB), tgt_sat.reshape(-1))
|
|
@@ -2288,6 +2376,26 @@ def _train_phase(
|
|
| 2288 |
seen_tok += toks_processed
|
| 2289 |
pbar.set_postfix(loss=f"{loss_value:.3f}", B=BATCH, L=BLOCK)
|
| 2290 |
pbar.update(toks_processed)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2291 |
_flush_sentinel = pathlib.Path(args.save_dir) / "FLUSH_NOW"
|
| 2292 |
if _flush_flag[0] or _flush_sentinel.exists():
|
| 2293 |
_flush_flag[0] = False
|
|
@@ -2642,7 +2750,7 @@ def infer(args):
|
|
| 2642 |
print(f"{Colors.INFO}Generating ({mode_str})...{Colors.RESET}")
|
| 2643 |
start = time.time()
|
| 2644 |
if args.mode == "ar":
|
| 2645 |
-
h, kvs = core(ids, causal_mask(ids.size(1)), use_cache=True, total_seq_len=ids.size(1))
|
| 2646 |
for _ in range(args.max_new):
|
| 2647 |
logits = ar_h(h)[:, -1]
|
| 2648 |
logits = _apply_penalties(logits, ids, args.penalty_last_n, args.repetition_penalty, args.presence_penalty, args.frequency_penalty)
|
|
@@ -2676,7 +2784,7 @@ def infer(args):
|
|
| 2676 |
ids[0, pos] = int(pred[0, pos])
|
| 2677 |
else:
|
| 2678 |
cached_len = ids.size(1)
|
| 2679 |
-
h, kvs = core(ids, sat_mask(ids.size(1)), use_cache=True, total_seq_len=cached_len)
|
| 2680 |
h_buffer = h[:, -SAT_BLOCK:]
|
| 2681 |
added = 0
|
| 2682 |
stop = False
|
|
@@ -2707,7 +2815,7 @@ def infer(args):
|
|
| 2707 |
if added >= args.max_new: break
|
| 2708 |
if added >= args.max_new: break
|
| 2709 |
new_ids = torch.cat(new_tokens, dim=1)
|
| 2710 |
-
mask = sat_mask_cached(new_ids.size(1), cached_len)
|
| 2711 |
h, kvs = core(new_ids, mask, kv_caches=kvs, use_cache=True, total_seq_len=ids.size(1))
|
| 2712 |
cached_len = ids.size(1)
|
| 2713 |
h_buffer = torch.cat([h_buffer, h], dim=1)[:, -SAT_BLOCK:]
|
|
@@ -2768,6 +2876,8 @@ def main():
|
|
| 2768 |
help="For --attn_backend sublinear, cap landmark candidates per query chunk.")
|
| 2769 |
tr.add_argument("--sublinear_chunk", type=int, default=DEFAULT_SUBLINEAR_CHUNK,
|
| 2770 |
help="For --attn_backend sublinear, query chunk size controlling peak gather memory.")
|
|
|
|
|
|
|
| 2771 |
tr.add_argument("--anchor_memory", action="store_true",
|
| 2772 |
help="Enable anchor-memory long-context augmentation (one AnchorMemoryLayer at mid-stack).")
|
| 2773 |
tr.add_argument("--anchor_stride", type=int, default=DEFAULT_ANCHOR_STRIDE,
|
|
@@ -2781,6 +2891,8 @@ def main():
|
|
| 2781 |
tr.add_argument("--optimizer", choices=["adamw", "adamw8bit", "paged_adamw8bit"], default="adamw",
|
| 2782 |
help="Optimizer backend. 8-bit options reduce VRAM on 24GB production runs.")
|
| 2783 |
tr.add_argument("--save_every_sec", type=int, default=DEFAULT_SAVE_SEC)
|
|
|
|
|
|
|
| 2784 |
tr.add_argument("--delta_every_steps", type=int, default=DEFAULT_DELTA_STEPS, help="Weight-only delta save every N steps (0=off)")
|
| 2785 |
tr.add_argument("--delta_max_keep", type=int, default=DEFAULT_MAX_DELTAS, help="Max delta checkpoints to keep")
|
| 2786 |
tr.add_argument("--resume_delta", type=str, help="Resume from a delta (weight-only, no optimizer state)")
|
|
@@ -2872,6 +2984,7 @@ def main():
|
|
| 2872 |
inf.add_argument("--sublinear_stride", type=int, default=DEFAULT_SUBLINEAR_STRIDE)
|
| 2873 |
inf.add_argument("--sublinear_max_anchors", type=int, default=DEFAULT_SUBLINEAR_MAX_ANCHORS)
|
| 2874 |
inf.add_argument("--sublinear_chunk", type=int, default=DEFAULT_SUBLINEAR_CHUNK)
|
|
|
|
| 2875 |
inf.add_argument("--nat_expand", type=int, default=2)
|
| 2876 |
inf.add_argument("--nat_passes", type=int, default=1)
|
| 2877 |
st = sub.add_parser("status", help="Read-only training status")
|
|
|
|
| 1194 |
dist = (j - i).clamp_min(0)
|
| 1195 |
return -_alibi_slopes(n_heads) * dist
|
| 1196 |
|
| 1197 |
+
|
| 1198 |
+
class StructuredAttentionMask:
|
| 1199 |
+
"""Symbolic attention rules for sublinear attention.
|
| 1200 |
+
|
| 1201 |
+
Dense masks are O(T^2). This object carries the rule so sublinear attention can
|
| 1202 |
+
apply it only to the gathered local/anchor candidate keys: O(T * candidates).
|
| 1203 |
+
"""
|
| 1204 |
+
|
| 1205 |
+
__slots__ = ("kind", "q_len", "k_len", "query_base", "block")
|
| 1206 |
+
|
| 1207 |
+
def __init__(self, kind: str, q_len: int, k_len: int = None, query_base: int = 0, block: int = 1):
|
| 1208 |
+
self.kind = (kind or "none").lower()
|
| 1209 |
+
self.q_len = int(q_len)
|
| 1210 |
+
self.k_len = int(k_len if k_len is not None else q_len)
|
| 1211 |
+
self.query_base = int(query_base)
|
| 1212 |
+
self.block = max(1, int(block))
|
| 1213 |
+
|
| 1214 |
+
def to_dense(self, device=None, dtype=torch.float32):
|
| 1215 |
+
device = device or DEV
|
| 1216 |
+
if self.kind in {"none", "nat", "bidirectional", "unrestricted"}:
|
| 1217 |
+
return None
|
| 1218 |
+
q_pos = torch.arange(self.query_base, self.query_base + self.q_len, device=device, dtype=torch.long).view(self.q_len, 1)
|
| 1219 |
+
k_pos = torch.arange(self.k_len, device=device, dtype=torch.long).view(1, self.k_len)
|
| 1220 |
+
if self.kind == "causal":
|
| 1221 |
+
allow = k_pos <= q_pos
|
| 1222 |
+
elif self.kind in {"sat", "block_causal", "block-causal"}:
|
| 1223 |
+
allow = (k_pos // self.block) <= (q_pos // self.block)
|
| 1224 |
+
else:
|
| 1225 |
+
raise ValueError(f"unknown structured attention mask kind: {self.kind}")
|
| 1226 |
+
zeros = torch.zeros((self.q_len, self.k_len), device=device, dtype=dtype)
|
| 1227 |
+
neg = torch.full_like(zeros, float("-inf"))
|
| 1228 |
+
return torch.where(allow, zeros, neg).unsqueeze(0).unsqueeze(0)
|
| 1229 |
+
|
| 1230 |
+
|
| 1231 |
+
def _is_structured_attention_mask(mask) -> bool:
|
| 1232 |
+
return isinstance(mask, StructuredAttentionMask)
|
| 1233 |
+
|
| 1234 |
+
|
| 1235 |
+
def use_structured_masks(args=None, backend: str = None) -> bool:
|
| 1236 |
+
backend = (backend or getattr(args, "attn_backend", "") or "").lower()
|
| 1237 |
+
return backend == "sublinear" and not bool(getattr(args, "no_structured_masks", False))
|
| 1238 |
+
|
| 1239 |
# βββββββββββββββββββββββββ Model components βββββββββββββββββββββββββ
|
| 1240 |
class KVBuffer:
|
| 1241 |
"""Preallocated K/V cache for decode. Replaces torch.cat-based growth.
|
|
|
|
| 1382 |
self._metric_cache_shape = (-1, -1)
|
| 1383 |
return super().train(mode)
|
| 1384 |
|
| 1385 |
+
def _structured_valid(self, attn_mask, q_pos, idx):
|
| 1386 |
+
if not _is_structured_attention_mask(attn_mask):
|
| 1387 |
+
return None
|
| 1388 |
+
kind = attn_mask.kind
|
| 1389 |
+
if kind in {"none", "nat", "bidirectional", "unrestricted"}:
|
| 1390 |
+
return torch.ones_like(idx, dtype=torch.bool)
|
| 1391 |
+
if kind == "causal":
|
| 1392 |
+
return idx <= q_pos[:, None]
|
| 1393 |
+
if kind in {"sat", "block_causal", "block-causal"}:
|
| 1394 |
+
block = max(1, int(attn_mask.block))
|
| 1395 |
+
return (idx // block) <= (q_pos[:, None] // block)
|
| 1396 |
+
raise ValueError(f"unknown structured attention mask kind: {kind}")
|
| 1397 |
+
|
| 1398 |
+
def _sublinear_attention(self, q, k, v, attn_mask=None, rel_bias_tokens=None):
|
| 1399 |
"""Local-window + landmark attention: O(N * (window + N/stride))."""
|
| 1400 |
bsz, heads, q_len, _ = q.shape
|
| 1401 |
k_len = k.size(2)
|
|
|
|
| 1403 |
query_base = max(0, k_len - q_len)
|
| 1404 |
outputs = []
|
| 1405 |
scale = 1.0 / math.sqrt(self.dk)
|
| 1406 |
+
slopes = None
|
| 1407 |
+
if self.use_relpos and rel_bias_tokens is not None:
|
| 1408 |
+
slopes = _alibi_slopes(self.h).to(device=device, dtype=torch.float32)
|
| 1409 |
|
| 1410 |
anchor_start = self.sublinear_stride - 1
|
| 1411 |
if self.sublinear_stride > 0 and self.sublinear_max_anchors > 0 and anchor_start < k_len:
|
|
|
|
| 1446 |
idx = local_idx
|
| 1447 |
valid = local_valid
|
| 1448 |
|
| 1449 |
+
structured_valid = self._structured_valid(attn_mask, q_pos, idx)
|
| 1450 |
+
if structured_valid is not None:
|
| 1451 |
+
valid = valid & structured_valid
|
| 1452 |
+
|
| 1453 |
k_sel = k[:, :, idx, :]
|
| 1454 |
scores = (q[:, :, q_start:q_end, :].unsqueeze(-2) * k_sel).sum(dim=-1) * scale
|
| 1455 |
|
| 1456 |
+
if slopes is not None:
|
| 1457 |
+
dist = (idx.view(1, 1, cur, -1) - q_pos.view(1, 1, cur, 1)).clamp_min(0).to(torch.float32)
|
| 1458 |
+
scores = scores + (-slopes * dist).to(scores.dtype)
|
| 1459 |
+
|
| 1460 |
+
if torch.is_tensor(attn_mask) and attn_mask.size(-1) == k_len and attn_mask.size(-2) >= q_end:
|
| 1461 |
mask_q = attn_mask[..., q_start:q_end, :]
|
| 1462 |
gather_idx = idx.view(1, 1, cur, -1).expand(mask_q.size(0), mask_q.size(1), cur, idx.size(1))
|
| 1463 |
scores = scores + torch.gather(mask_q, -1, gather_idx)
|
|
|
|
| 1494 |
else:
|
| 1495 |
k, v = k_new, v_new
|
| 1496 |
attn_mask = mask
|
| 1497 |
+
if self.attn_backend != "sublinear" and _is_structured_attention_mask(attn_mask):
|
| 1498 |
+
attn_mask = attn_mask.to_dense(device=q.device, dtype=q.dtype)
|
| 1499 |
+
if self.attn_backend != "sublinear" and self.use_relpos and rel_bias_tokens is not None:
|
| 1500 |
rel = alibi_bias(self.h, rel_bias_tokens)[:, :, -q.size(2):, :]
|
| 1501 |
attn_mask = rel if attn_mask is None else attn_mask + rel
|
| 1502 |
if self.attn_backend == "sdpa":
|
|
|
|
| 1513 |
q_scaled = q * math.sqrt(q.size(-1) / self.dk)
|
| 1514 |
z = F.scaled_dot_product_attention(q_scaled, k, v, attn_mask=attn_mask, dropout_p=0.0)
|
| 1515 |
elif self.attn_backend == "sublinear":
|
| 1516 |
+
z = self._sublinear_attention(q, k, v, attn_mask=attn_mask, rel_bias_tokens=rel_bias_tokens)
|
| 1517 |
else:
|
| 1518 |
att = (q @ k.transpose(-1, -2)) / math.sqrt(self.dk)
|
| 1519 |
if attn_mask is not None:
|
|
|
|
| 1628 |
if not use_cache:
|
| 1629 |
for i, blk in enumerate(self.blocks):
|
| 1630 |
if self.grad_checkpoint and self.training:
|
| 1631 |
+
x = torch_checkpoint.checkpoint(lambda y, block=blk: block(y, mask), x, use_reentrant=False)
|
| 1632 |
else:
|
| 1633 |
x = blk(x, mask)
|
| 1634 |
if self.anchor is not None and i == self.anchor_position:
|
|
|
|
| 1690 |
|
| 1691 |
|
| 1692 |
# βββββββββββββββββββββββββ Masks βββββββββββββββββββββββββ
|
| 1693 |
+
def causal_mask(n, structured: bool = False):
|
| 1694 |
+
if structured:
|
| 1695 |
+
return StructuredAttentionMask("causal", q_len=n, k_len=n, query_base=0)
|
| 1696 |
return torch.triu(torch.full((1, 1, n, n), float("-inf"), device=DEV), 1)
|
| 1697 |
|
| 1698 |
+
def sat_mask(n, block=SAT_BLOCK, structured: bool = False):
|
| 1699 |
+
if structured:
|
| 1700 |
+
return StructuredAttentionMask("sat", q_len=n, k_len=n, query_base=0, block=block)
|
| 1701 |
idx = torch.arange(n, device=DEV)
|
| 1702 |
grp = idx.unsqueeze(0) // block
|
| 1703 |
allow = (grp.T == grp) | (grp.T > grp)
|
| 1704 |
return torch.where(allow, 0.0, float("-inf")).unsqueeze(0).unsqueeze(0)
|
| 1705 |
|
| 1706 |
+
def sat_mask_cached(new_len: int, cached_len: int, block=SAT_BLOCK, structured: bool = False):
|
| 1707 |
total_len = cached_len + new_len
|
| 1708 |
+
if structured:
|
| 1709 |
+
return StructuredAttentionMask("sat", q_len=new_len, k_len=total_len, query_base=cached_len, block=block)
|
| 1710 |
q_idx = torch.arange(cached_len, total_len, device=DEV).unsqueeze(1)
|
| 1711 |
k_idx = torch.arange(total_len, device=DEV).unsqueeze(0)
|
| 1712 |
q_grp = q_idx // block
|
|
|
|
| 2230 |
now_wall = time.time()
|
| 2231 |
last_save_mono = time.monotonic() - (now_wall - (resume_wall_time or now_wall))
|
| 2232 |
last_delta_step = start_step
|
| 2233 |
+
last_heartbeat_mono = time.monotonic()
|
| 2234 |
print(f"[{phase_name}] Starting. Goal: {total_tokens_needed:,} tokens. Batch={BATCH}, Block={BLOCK}")
|
| 2235 |
print(
|
| 2236 |
f"[{phase_name}] AR_ONLY={args.ar_only}, SAT_EVERY={args.sat_every}, "
|
|
|
|
| 2246 |
except (ValueError, OSError):
|
| 2247 |
pass
|
| 2248 |
_DBS = _dblock_init(core, args) if getattr(args,'dblock',False) else None
|
| 2249 |
+
if DEV.type == "cuda":
|
| 2250 |
+
try:
|
| 2251 |
+
torch.cuda.empty_cache()
|
| 2252 |
+
torch.cuda.reset_peak_memory_stats()
|
| 2253 |
+
print(
|
| 2254 |
+
f"[vram] training-start cache cleared: "
|
| 2255 |
+
f"alloc={torch.cuda.memory_allocated() / (1024**3):.2f}GB "
|
| 2256 |
+
f"reserved={torch.cuda.memory_reserved() / (1024**3):.2f}GB "
|
| 2257 |
+
f"structured_masks={use_structured_masks(args)}",
|
| 2258 |
+
flush=True,
|
| 2259 |
+
)
|
| 2260 |
+
except Exception:
|
| 2261 |
+
pass
|
| 2262 |
while seen_tok < total_tokens_needed:
|
| 2263 |
try:
|
| 2264 |
while len(buf) < BLOCK:
|
|
|
|
| 2278 |
loss_value = _dblock_step(core, ar_h, sat_h, nat_h, opt, scaler, args, ids, _DBS)
|
| 2279 |
else:
|
| 2280 |
with amp(args.amp):
|
| 2281 |
+
h_ar = core(ids, causal_mask(ids.size(1), structured=use_structured_masks(args)))
|
| 2282 |
logits_ar = ar_h(h_ar)[:, :-1]
|
| 2283 |
loss_ar = ce_tok(logits_ar.reshape(-1, VOCAB), tgt_ar[:, 1:].reshape(-1))
|
| 2284 |
loss_value = float(loss_ar.detach().item())
|
|
|
|
| 2289 |
# Same AR+SAT objective as a summed loss, but sequential backward keeps
|
| 2290 |
# only one core-forward activation graph live at a time on 24GB cards.
|
| 2291 |
with amp(args.amp):
|
| 2292 |
+
h_sat = core(ids, sat_mask(ids.size(1), structured=use_structured_masks(args)))
|
| 2293 |
logits_sat, gate = sat_h(h_sat[:, -SAT_BLOCK:])
|
| 2294 |
tgt_sat = ids[:, 1:SAT_BLOCK+1]
|
| 2295 |
loss_sat = ce_tok(logits_sat.reshape(-1, VOCAB), tgt_sat.reshape(-1))
|
|
|
|
| 2376 |
seen_tok += toks_processed
|
| 2377 |
pbar.set_postfix(loss=f"{loss_value:.3f}", B=BATCH, L=BLOCK)
|
| 2378 |
pbar.update(toks_processed)
|
| 2379 |
+
heartbeat_every = int(getattr(args, "heartbeat_every_sec", 300) or 0)
|
| 2380 |
+
now_mono = time.monotonic()
|
| 2381 |
+
if heartbeat_every > 0 and now_mono - last_heartbeat_mono >= heartbeat_every:
|
| 2382 |
+
mem = ""
|
| 2383 |
+
if DEV.type == "cuda":
|
| 2384 |
+
try:
|
| 2385 |
+
mem = (
|
| 2386 |
+
f" gpu_alloc={torch.cuda.memory_allocated() / (1024**3):.2f}GB"
|
| 2387 |
+
f" gpu_reserved={torch.cuda.memory_reserved() / (1024**3):.2f}GB"
|
| 2388 |
+
f" gpu_peak={torch.cuda.max_memory_allocated() / (1024**3):.2f}GB"
|
| 2389 |
+
)
|
| 2390 |
+
except Exception:
|
| 2391 |
+
mem = ""
|
| 2392 |
+
print(
|
| 2393 |
+
f"[heartbeat] phase={phase_name} pid={os.getpid()} step={step} "
|
| 2394 |
+
f"seen_tok={seen_tok} loss={loss_value:.3f} B={BATCH} L={BLOCK} "
|
| 2395 |
+
f"dblock={bool(getattr(args, 'dblock', False))} structured_masks={use_structured_masks(args)}{mem}",
|
| 2396 |
+
flush=True,
|
| 2397 |
+
)
|
| 2398 |
+
last_heartbeat_mono = now_mono
|
| 2399 |
_flush_sentinel = pathlib.Path(args.save_dir) / "FLUSH_NOW"
|
| 2400 |
if _flush_flag[0] or _flush_sentinel.exists():
|
| 2401 |
_flush_flag[0] = False
|
|
|
|
| 2750 |
print(f"{Colors.INFO}Generating ({mode_str})...{Colors.RESET}")
|
| 2751 |
start = time.time()
|
| 2752 |
if args.mode == "ar":
|
| 2753 |
+
h, kvs = core(ids, causal_mask(ids.size(1), structured=use_structured_masks(args)), use_cache=True, total_seq_len=ids.size(1))
|
| 2754 |
for _ in range(args.max_new):
|
| 2755 |
logits = ar_h(h)[:, -1]
|
| 2756 |
logits = _apply_penalties(logits, ids, args.penalty_last_n, args.repetition_penalty, args.presence_penalty, args.frequency_penalty)
|
|
|
|
| 2784 |
ids[0, pos] = int(pred[0, pos])
|
| 2785 |
else:
|
| 2786 |
cached_len = ids.size(1)
|
| 2787 |
+
h, kvs = core(ids, sat_mask(ids.size(1), structured=use_structured_masks(args)), use_cache=True, total_seq_len=cached_len)
|
| 2788 |
h_buffer = h[:, -SAT_BLOCK:]
|
| 2789 |
added = 0
|
| 2790 |
stop = False
|
|
|
|
| 2815 |
if added >= args.max_new: break
|
| 2816 |
if added >= args.max_new: break
|
| 2817 |
new_ids = torch.cat(new_tokens, dim=1)
|
| 2818 |
+
mask = sat_mask_cached(new_ids.size(1), cached_len, structured=use_structured_masks(args))
|
| 2819 |
h, kvs = core(new_ids, mask, kv_caches=kvs, use_cache=True, total_seq_len=ids.size(1))
|
| 2820 |
cached_len = ids.size(1)
|
| 2821 |
h_buffer = torch.cat([h_buffer, h], dim=1)[:, -SAT_BLOCK:]
|
|
|
|
| 2876 |
help="For --attn_backend sublinear, cap landmark candidates per query chunk.")
|
| 2877 |
tr.add_argument("--sublinear_chunk", type=int, default=DEFAULT_SUBLINEAR_CHUNK,
|
| 2878 |
help="For --attn_backend sublinear, query chunk size controlling peak gather memory.")
|
| 2879 |
+
tr.add_argument("--no_structured_masks", action="store_true",
|
| 2880 |
+
help="Disable structured causal/SAT masks for sublinear attention and fall back to dense masks.")
|
| 2881 |
tr.add_argument("--anchor_memory", action="store_true",
|
| 2882 |
help="Enable anchor-memory long-context augmentation (one AnchorMemoryLayer at mid-stack).")
|
| 2883 |
tr.add_argument("--anchor_stride", type=int, default=DEFAULT_ANCHOR_STRIDE,
|
|
|
|
| 2891 |
tr.add_argument("--optimizer", choices=["adamw", "adamw8bit", "paged_adamw8bit"], default="adamw",
|
| 2892 |
help="Optimizer backend. 8-bit options reduce VRAM on 24GB production runs.")
|
| 2893 |
tr.add_argument("--save_every_sec", type=int, default=DEFAULT_SAVE_SEC)
|
| 2894 |
+
tr.add_argument("--heartbeat_every_sec", type=int, default=300,
|
| 2895 |
+
help="Print lightweight trainer heartbeat/status lines every N seconds; 0 disables.")
|
| 2896 |
tr.add_argument("--delta_every_steps", type=int, default=DEFAULT_DELTA_STEPS, help="Weight-only delta save every N steps (0=off)")
|
| 2897 |
tr.add_argument("--delta_max_keep", type=int, default=DEFAULT_MAX_DELTAS, help="Max delta checkpoints to keep")
|
| 2898 |
tr.add_argument("--resume_delta", type=str, help="Resume from a delta (weight-only, no optimizer state)")
|
|
|
|
| 2984 |
inf.add_argument("--sublinear_stride", type=int, default=DEFAULT_SUBLINEAR_STRIDE)
|
| 2985 |
inf.add_argument("--sublinear_max_anchors", type=int, default=DEFAULT_SUBLINEAR_MAX_ANCHORS)
|
| 2986 |
inf.add_argument("--sublinear_chunk", type=int, default=DEFAULT_SUBLINEAR_CHUNK)
|
| 2987 |
+
inf.add_argument("--no_structured_masks", action="store_true")
|
| 2988 |
inf.add_argument("--nat_expand", type=int, default=2)
|
| 2989 |
inf.add_argument("--nat_passes", type=int, default=1)
|
| 2990 |
st = sub.add_parser("status", help="Read-only training status")
|
relaunch_agillm4_dblock.sh
CHANGED
|
@@ -19,4 +19,5 @@ exec python -u nB300_agillm4.py train --preset agillm4_floor --resume "$CKPT" \
|
|
| 19 |
--batch_size 1 --block "${AGILLM4_BLOCK:-1280}" --amp --attn_backend "${AGILLM_ATTN_BACKEND}" --grad_checkpoint \
|
| 20 |
--optimizer paged_adamw8bit --sat_every 1 --nat_every 1 --nat_max_tokens 768 --nat_mask_ratio 0.5 \
|
| 21 |
--token_param_ratio 100 --save_dir "$SAVE_DIR" \
|
| 22 |
-
--save_every_sec 86400 --
|
|
|
|
|
|
| 19 |
--batch_size 1 --block "${AGILLM4_BLOCK:-1280}" --amp --attn_backend "${AGILLM_ATTN_BACKEND}" --grad_checkpoint \
|
| 20 |
--optimizer paged_adamw8bit --sat_every 1 --nat_every 1 --nat_max_tokens 768 --nat_mask_ratio 0.5 \
|
| 21 |
--token_param_ratio 100 --save_dir "$SAVE_DIR" \
|
| 22 |
+
--save_every_sec 86400 --heartbeat_every_sec "${AGILLM4_HEARTBEAT_EVERY_SEC:-300}" \
|
| 23 |
+
--delta_every_steps 25000 --delta_max_keep 1 --max_ckpts 1
|