Scott/Codex commited on
Commit
78c46ca
Β·
1 Parent(s): 979693c

Add structured sublinear masks and trainer heartbeat

Browse files
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], h, causal, 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,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], h2, smask, 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,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], hn, None, 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())
 
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 _sublinear_attention(self, q, k, v, attn_mask=None):
 
 
 
 
 
 
 
 
 
 
 
 
 
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 attn_mask is not None and attn_mask.size(-1) == k_len and attn_mask.size(-2) >= q_end:
 
 
 
 
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.use_relpos and rel_bias_tokens is not None:
 
 
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(blk, x, mask, use_reentrant=False)
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 --delta_every_steps 25000 --delta_max_keep 1 --max_ckpts 1
 
 
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