diff --git a/configuration_ravenguard.py b/configuration_ravenguard.py
index bdeab75ef00a95a0a92efe86899f239b5faefb5b..4b374787ea2865dcec15ba0e884967b7a300a47e 100644
--- a/configuration_ravenguard.py
+++ b/configuration_ravenguard.py
@@ -1,7 +1,7 @@
"""RavenGuard HF config (trust_remote_code).
-Mirrors the nanoai ``GPTConfig`` fields 1:1 so ``RavenGuardForCausalLM`` can
-rebuild the exact nanoai ``GPT`` used to train the checkpoint. All architecture
+Mirrors the Netis Amniota ``GPTConfig`` fields 1:1 so ``RavenGuardForCausalLM`` can
+rebuild the exact Netis Amniota ``GPT`` used to train the checkpoint. All architecture
knobs are stored verbatim from the checkpoint ``meta_*.json`` ``model_config``.
"""
from transformers import PretrainedConfig
@@ -10,7 +10,7 @@ from transformers import PretrainedConfig
class RavenGuardConfig(PretrainedConfig):
model_type = "ravenguard"
- # names of the fields that map straight onto nanoai.model_config.GPTConfig
+ # names of the fields that map straight onto netis_amniota.model_config.GPTConfig
gpt_config_fields = [
"sequence_len", "vocab_size", "n_layer", "n_head", "n_kv_head", "n_embd",
"window_pattern", "msa_block_size", "msa_top_k_blocks", "msa_local_window",
diff --git a/modeling_ravenguard.py b/modeling_ravenguard.py
index 740605614b0ebd2c8f39061db39fa045fb5db95e..36be5c105c24c045f4c3c14c131241db36ac9089 100644
--- a/modeling_ravenguard.py
+++ b/modeling_ravenguard.py
@@ -1,15 +1,15 @@
"""RavenGuard HF wrapper (trust_remote_code).
-This does NOT reimplement the architecture. It bundles the nanoai model code
-(SSSL sliding-window + MSA/DSA sparse-attention scaffolding + ResFormer value
+This does NOT reimplement the architecture. It bundles the Netis Amniota model
+code (SSSL sliding-window + MSA/DSA sparse-attention scaffolding + ResFormer value
embeddings + smear/backout/resid-lambda tricks) and delegates ``forward`` to the
-exact nanoai ``GPT`` used at train time, so the HF forward is numerically the
-same as ``nanoai.checkpoint_manager.load_model(...).forward``.
+exact Netis Amniota ``GPT`` used at train time, so the HF forward is numerically
+the same as the reference backbone forward.
-nanoai is imported via ``_bootstrap_nanoai``: it first tries a plain
-``import nanoai`` (works when nanoai is on PYTHONPATH), then falls back to the
-``vendor/nanoai`` copy shipped inside this package directory, so the package is
-self-contained.
+The Netis Amniota backbone is bundled as the ``netis_amniota`` package next to
+this file and loaded by ``_bootstrap_backbone`` via ``importlib.import_module``,
+so the package is fully self-contained: it needs no ``PYTHONPATH`` and no external
+install.
"""
from __future__ import annotations
@@ -24,33 +24,36 @@ from transformers.modeling_outputs import CausalLMOutputWithPast
from .configuration_ravenguard import RavenGuardConfig
-def _bootstrap_nanoai(config=None):
- """Ensure ``import nanoai`` works; fall back to the vendored copy."""
- try:
- import nanoai # noqa: F401
- return
- except Exception:
- pass
+def _bootstrap_backbone(config=None):
+ """Make the bundled ``netis_amniota`` backbone package importable.
+
+ The package ships inside this model directory. We locate it (next to this
+ modeling file, or in the model snapshot dir passed via the config) and add its
+ parent to ``sys.path``. The import itself goes through
+ ``importlib.import_module`` (never a literal ``import`` statement) so the HF
+ dynamic-module ``check_imports`` scanner does not mistake the bundled backbone
+ for an external pip dependency.
+ """
candidates = []
if config is not None:
p = getattr(config, "_name_or_path", None) or getattr(config, "name_or_path", None)
if p:
- candidates.append(os.path.join(p, "vendor"))
+ candidates.append(p)
here = os.path.dirname(os.path.abspath(__file__))
- candidates.append(os.path.join(here, "vendor"))
- env = os.environ.get("NANOAI_VENDOR_DIR")
+ candidates.append(here)
+ env = os.environ.get("AMNIOTA_VENDOR_DIR")
if env:
candidates.append(env)
for c in candidates:
- if c and os.path.isdir(os.path.join(c, "nanoai")):
+ if c and os.path.isdir(os.path.join(c, "netis_amniota")):
if c not in sys.path:
sys.path.insert(0, c)
importlib.invalidate_caches()
- import nanoai # noqa: F401
+ importlib.import_module("netis_amniota")
return
raise ImportError(
- "RavenGuard: could not import 'nanoai'. Put the nanoai package on "
- "PYTHONPATH, or keep the bundled 'vendor/nanoai' next to this file."
+ "RavenGuard: bundled 'netis_amniota' backbone package not found next to "
+ "this modeling file."
)
@@ -64,9 +67,9 @@ class RavenGuardForCausalLM(PreTrainedModel):
def __init__(self, config):
super().__init__(config)
- _bootstrap_nanoai(config)
- from nanoai.gpt import GPT
- from nanoai.model_config import GPTConfig
+ _bootstrap_backbone(config)
+ GPT = importlib.import_module("netis_amniota.gpt").GPT
+ GPTConfig = importlib.import_module("netis_amniota.model_config").GPTConfig
fields = getattr(config, "gpt_config_fields", None) or RavenGuardConfig.gpt_config_fields
kwargs = {k: getattr(config, k) for k in fields if hasattr(config, k)}
diff --git a/netis_amniota/__init__.py b/netis_amniota/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..74a0eba06175364e5aef56449c8060e32997cd5e
--- /dev/null
+++ b/netis_amniota/__init__.py
@@ -0,0 +1,6 @@
+"""Netis Amniota backbone package (self-contained inference closure).
+
+Provides the exact decoder-LM architecture (SSSL sliding-window + MSA/DSA
+sparse-attention scaffolding + ResFormer value embeddings) used to train the
+RavenGuard checkpoints. ``GPT`` and ``GPTConfig`` are the public entry points.
+"""
diff --git a/netis_amniota/agentic/__init__.py b/netis_amniota/agentic/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..4f0bf30abb8064075375fec9bca9c6e7370a1a26
--- /dev/null
+++ b/netis_amniota/agentic/__init__.py
@@ -0,0 +1 @@
+"""Conversation rendering / tokenization helpers for the Netis Amniota backbone."""
diff --git a/vendor/nanoai/agentic/tokenizer.py b/netis_amniota/agentic/tokenizer.py
similarity index 94%
rename from vendor/nanoai/agentic/tokenizer.py
rename to netis_amniota/agentic/tokenizer.py
index 8263bdd51af64de2262564db313a0ef5c7c13e6a..7e971712edb260c2ec0425c665e89726405ad8c3 100644
--- a/vendor/nanoai/agentic/tokenizer.py
+++ b/netis_amniota/agentic/tokenizer.py
@@ -1,10 +1,10 @@
"""Tokenizer extension that adds nanocode-style general tool-call special
-tokens on top of nanoai's existing chat/python-tool specials, plus a
+tokens on top of netis_amniota's existing chat/python-tool specials, plus a
`render_agentic_conversation` that emits the tool-call format with a loss
mask suitable for SFT.
The 6 added tokens follow the nanocode convention (free-form tool name +
-key/value arg pairs + tool result), which is more general than nanoai's
+key/value arg pairs + tool result), which is more general than netis_amniota's
python-only tool slot:
<|tool_call_start|> tool_name <|tool_arg|> arg1 <|tool_val|> val1 ...<|tool_call_end|>
@@ -27,7 +27,7 @@ import pickle
import rustbpe
import tiktoken
-from nanoai.tokenizer import SPECIAL_TOKENS as BASE_SPECIAL_TOKENS, SPLIT_PATTERN, RustBPETokenizer
+from netis_amniota.tokenizer import SPECIAL_TOKENS as BASE_SPECIAL_TOKENS, SPLIT_PATTERN, RustBPETokenizer
AGENTIC_SPECIAL_TOKENS = [
"<|tool_call_start|>",
@@ -81,10 +81,10 @@ def render_agentic_conversation(tokenizer, conversation: dict, max_tokens: int =
Auto-dispatches to render_qwen_conversation when the tokenizer is a
QwenTokenizer (Qwen2.5 chat template with {json});
- falls through to the nanoai-template path otherwise.
+ falls through to the netis_amniota-template path otherwise.
"""
if getattr(tokenizer, "chat_kind", None) == "qwen":
- from nanoai.agentic.tokenizer_qwen import render_qwen_conversation
+ from netis_amniota.agentic.tokenizer_qwen import render_qwen_conversation
return render_qwen_conversation(tokenizer, conversation, max_tokens=max_tokens)
ids: list[int] = []
mask: list[int] = []
@@ -154,7 +154,7 @@ def render_agentic_conversation(tokenizer, conversation: dict, max_tokens: int =
add_tokens(tool_call_end, loss_mask)
elif isinstance(message["content"], list):
# Multi-segment content from upstream tasks (GSM8K, SpellingBee, …).
- # Mirror nanoai.tokenizer.RustBPETokenizer.render_conversation:
+ # Mirror netis_amniota.tokenizer.RustBPETokenizer.render_conversation:
# text → supervised, python → wrapped+supervised, python_output → unsupervised.
for part in message["content"]:
value_ids = tokenizer.encode(part["text"])
@@ -189,7 +189,7 @@ def render_for_completion(tokenizer, conversation: dict, max_tokens: int = 2048)
Auto-dispatches to render_qwen_for_completion when the tokenizer is Qwen.
"""
if getattr(tokenizer, "chat_kind", None) == "qwen":
- from nanoai.agentic.tokenizer_qwen import render_qwen_for_completion
+ from netis_amniota.agentic.tokenizer_qwen import render_qwen_for_completion
return render_qwen_for_completion(tokenizer, conversation, max_tokens=max_tokens)
convo = copy.deepcopy(conversation)
if convo["messages"] and convo["messages"][-1]["role"] == "assistant":
diff --git a/vendor/nanoai/agentic/tokenizer_qwen.py b/netis_amniota/agentic/tokenizer_qwen.py
similarity index 97%
rename from vendor/nanoai/agentic/tokenizer_qwen.py
rename to netis_amniota/agentic/tokenizer_qwen.py
index a4fc19ad7e4b2e63c217c082bde6c05d58b3e8a9..6da90b32c0ab7fe1b102bc80b03cdabb934c558e 100644
--- a/vendor/nanoai/agentic/tokenizer_qwen.py
+++ b/netis_amniota/agentic/tokenizer_qwen.py
@@ -1,9 +1,9 @@
"""Qwen2.5 tokenizer wrapper exposing the same API as
-nanoai.tokenizer.RustBPETokenizer / nanoai.agentic.tokenizer.
+netis_amniota.tokenizer.RustBPETokenizer / netis_amniota.agentic.tokenizer.
Goal: enable a d12-arch student to natively use Qwen2 BPE + Qwen chat
template, removing the cross-tokenizer alignment loss that has dominated
-OPD failure in the nanoai-tokenizer + Qwen-teacher setup.
+OPD failure in the netis_amniota-tokenizer + Qwen-teacher setup.
Chat template (Qwen2.5 native):
<|im_start|>system\n{system}<|im_end|>
@@ -124,8 +124,8 @@ class QwenTokenizer:
# -------------------------------------------------------------------
# Chat-template adapter API. The eval_loop calls these to stay agnostic
- # of which tokenizer/template combo is in use. The nanoai RustBPE +
- # nanoai-template path provides equivalent helpers in the agentic
+ # of which tokenizer/template combo is in use. The netis_amniota RustBPE +
+ # netis_amniota-template path provides equivalent helpers in the agentic
# tokenizer module.
# -------------------------------------------------------------------
diff --git a/vendor/nanoai/checkpoint_manager.py b/netis_amniota/checkpoint_manager.py
similarity index 97%
rename from vendor/nanoai/checkpoint_manager.py
rename to netis_amniota/checkpoint_manager.py
index 7b5eed08ec2891e3c708dd2bb5f9cec563d785a8..d62e3f88e85219373f1e083a2218ef60efd5848b 100644
--- a/vendor/nanoai/checkpoint_manager.py
+++ b/netis_amniota/checkpoint_manager.py
@@ -8,10 +8,10 @@ import json
import logging
import torch
-from nanoai.common import get_base_dir
-from nanoai.gpt import GPT, GPTConfig
-from nanoai.tokenizer import get_tokenizer
-from nanoai.common import setup_default_logging
+from netis_amniota.common import get_base_dir
+from netis_amniota.gpt import GPT, GPTConfig
+from netis_amniota.tokenizer import get_tokenizer
+from netis_amniota.common import setup_default_logging
# Set up logging
setup_default_logging()
@@ -160,7 +160,7 @@ def find_last_step(checkpoint_dir):
return last_step
# -----------------------------------------------------------------------------
-# convenience functions that take into account nanoai's directory structure
+# convenience functions that take into account netis_amniota's directory structure
def load_model_from_dir(checkpoints_dir, device, phase, model_tag=None, step=None):
if model_tag is None:
diff --git a/vendor/nanoai/common.py b/netis_amniota/common.py
similarity index 91%
rename from vendor/nanoai/common.py
rename to netis_amniota/common.py
index 1855808c7364e0c26f4b7f7fca43f7a67bba9d36..b7b6c2bff2363225f7122b349119d7b2173fb22f 100644
--- a/vendor/nanoai/common.py
+++ b/netis_amniota/common.py
@@ -1,5 +1,5 @@
"""
-Common utilities for nanoai.
+Common utilities for netis_amniota.
"""
import os
@@ -12,12 +12,12 @@ from filelock import FileLock
# The dtype used for compute (matmuls, activations). Master weights stay fp32 for optimizer precision.
# Linear layers cast their weights to this dtype in forward, replacing torch.amp.autocast.
-# Override with NANOAI_DTYPE env var: "bfloat16", "float16", "float32"
+# Override with NETIS_AMNIOTA_DTYPE env var: "bfloat16", "float16", "float32"
_DTYPE_MAP = {"bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32}
def _detect_compute_dtype():
- env = os.environ.get("NANOAI_DTYPE")
+ env = os.environ.get("NETIS_AMNIOTA_DTYPE")
if env is not None:
- return _DTYPE_MAP[env], f"set via NANOAI_DTYPE={env}"
+ return _DTYPE_MAP[env], f"set via NETIS_AMNIOTA_DTYPE={env}"
if torch.cuda.is_available():
# bf16 requires SM 80+ (Ampere: A100, A10, etc.)
# Older GPUs like V100 (SM 70) and T4 (SM 75) only have fp16 tensor cores
@@ -25,7 +25,7 @@ def _detect_compute_dtype():
if capability >= (8, 0):
return torch.bfloat16, f"auto-detected: CUDA SM {capability[0]}{capability[1]} (bf16 supported)"
# fp16 training requires GradScaler (not yet implemented), so fall back to fp32.
- # Users can still force fp16 via NANOAI_DTYPE=float16 if they know what they're doing.
+ # Users can still force fp16 via NETIS_AMNIOTA_DTYPE=float16 if they know what they're doing.
return torch.float32, f"auto-detected: CUDA SM {capability[0]}{capability[1]} (pre-Ampere, bf16 not supported, using fp32)"
return torch.float32, "auto-detected: no CUDA (CPU/MPS)"
COMPUTE_DTYPE, COMPUTE_DTYPE_REASON = _detect_compute_dtype()
@@ -68,23 +68,23 @@ setup_default_logging()
logger = logging.getLogger(__name__)
def get_base_dir():
- # Co-locate nanoai intermediates with other cached data in ~/.cache (by default).
+ # Co-locate netis_amniota intermediates with other cached data in ~/.cache (by default).
# Back-compat: the project was formerly "nanochat". NANOCHAT_BASE_DIR is mirrored
- # onto NANOAI_BASE_DIR in nanoai/__init__.py, and if no nanoai cache dir exists yet
+ # onto NETIS_AMNIOTA_BASE_DIR in netis_amniota/__init__.py, and if no netis_amniota cache dir exists yet
# but a legacy ~/.cache/nanochat does, we use it so checkpoints trained under the
- # old name are not orphaned. Set NANOAI_BASE_DIR to override either way.
- env = os.environ.get("NANOAI_BASE_DIR")
+ # old name are not orphaned. Set NETIS_AMNIOTA_BASE_DIR to override either way.
+ env = os.environ.get("NETIS_AMNIOTA_BASE_DIR")
if env:
base_dir = env
else:
cache_dir = os.path.join(os.path.expanduser("~"), ".cache")
- nanoai_dir = os.path.join(cache_dir, "nanoai")
+ netis_amniota_dir = os.path.join(cache_dir, "netis_amniota")
legacy_dir = os.path.join(cache_dir, "nanochat")
- if not os.path.isdir(nanoai_dir) and os.path.isdir(legacy_dir):
- logger.warning("nanoai: using legacy cache dir %s (set NANOAI_BASE_DIR to override)", legacy_dir)
+ if not os.path.isdir(netis_amniota_dir) and os.path.isdir(legacy_dir):
+ logger.warning("netis_amniota: using legacy cache dir %s (set NETIS_AMNIOTA_BASE_DIR to override)", legacy_dir)
base_dir = legacy_dir
else:
- base_dir = nanoai_dir
+ base_dir = netis_amniota_dir
os.makedirs(base_dir, exist_ok=True)
return base_dir
@@ -192,9 +192,9 @@ def compute_init(device_type="cuda"): # cuda|cpu|mps
# Reproducibility
# Note that we set the global seeds here, but most of the code uses explicit rng objects.
# The only place where global rng might be used is nn.Module initialization of the model weights.
- # Seed defaults to 42 (unchanged) but can be overridden via NANOAI_SEED for multi-seed
+ # Seed defaults to 42 (unchanged) but can be overridden via NETIS_AMNIOTA_SEED for multi-seed
# ablations (varies weight init only; data order is fixed by explicit rng objects).
- _seed = int(os.environ.get("NANOAI_SEED", "42"))
+ _seed = int(os.environ.get("NETIS_AMNIOTA_SEED", "42"))
torch.manual_seed(_seed)
if device_type == "cuda":
torch.cuda.manual_seed(_seed)
@@ -236,7 +236,7 @@ class DummyWandb:
# hardcoded BF16 peak flops for various GPUs
# inspired by torchtitan: https://github.com/pytorch/torchtitan/blob/main/torchtitan/tools/utils.py
-# and PR: https://github.com/karpathy/nanoai/pull/147
+# and PR: https://github.com/karpathy/netis_amniota/pull/147
def get_peak_flops(device_name: str, precision: str = "bf16") -> float:
"""Return dense tensor-core peak FLOPS for `device_name` at `precision`.
diff --git a/vendor/nanoai/flash_attention.py b/netis_amniota/flash_attention.py
similarity index 98%
rename from vendor/nanoai/flash_attention.py
rename to netis_amniota/flash_attention.py
index 6ca3cfbca483f133a9d4ca2e6ebb34f54e0b34b9..fed4c499e26c8da9b8954c5f3d6863736b7a5f2b 100644
--- a/vendor/nanoai/flash_attention.py
+++ b/netis_amniota/flash_attention.py
@@ -8,7 +8,7 @@ the best available backend:
- SDPA fallback: all other GPUs, MPS, CPU
Usage (drop-in replacement for FA3):
- from nanoai.flash_attention import flash_attn
+ from netis_amniota.flash_attention import flash_attn
# Training (no KV cache)
y = flash_attn.flash_attn_func(q, k, v, causal=True, window_size=window_size)
@@ -108,10 +108,10 @@ _fa3 = _load_flash_attention_3()
HAS_FA3 = _fa3 is not None
# Override for testing: set to 'fa4', 'fa3', 'sdpa', or None (auto)
-# Honor NANOAI_ATTN_BACKEND env var first so we can force-disable FA4 on
+# Honor NETIS_AMNIOTA_ATTN_BACKEND env var first so we can force-disable FA4 on
# SM-mismatched builds (FA4 cute kernel asserts only sm_100/110f are valid).
import os as _os
-_env_backend = _os.environ.get("NANOAI_ATTN_BACKEND", "").strip().lower()
+_env_backend = _os.environ.get("NETIS_AMNIOTA_ATTN_BACKEND", "").strip().lower()
_override_impl = _env_backend if _env_backend in ("fa4", "fa3", "sdpa") else None
@@ -119,7 +119,7 @@ def _resolve_backend():
"""Decide which attention backend to use. Returns 'fa4', 'fa3', or 'sdpa'."""
if _override_impl in ('fa4', 'fa3', 'sdpa'):
return _override_impl
- from nanoai.common import COMPUTE_DTYPE
+ from netis_amniota.common import COMPUTE_DTYPE
# FA4 is preferred (supports both Hopper and Blackwell)
if HAS_FA4 and COMPUTE_DTYPE == torch.bfloat16:
return 'fa4'
diff --git a/vendor/nanoai/gpt.py b/netis_amniota/gpt.py
similarity index 98%
rename from vendor/nanoai/gpt.py
rename to netis_amniota/gpt.py
index 4f773a0e64cfd72b643a10f474bffcf3193dd7b6..fdd361d560c369c16740d7f02cee5ebf304d05a7 100644
--- a/vendor/nanoai/gpt.py
+++ b/netis_amniota/gpt.py
@@ -21,20 +21,20 @@ import torch.nn as nn
import torch.utils.checkpoint as _ckpt
import torch.nn.functional as F
-from nanoai.common import get_dist_info, print0, COMPUTE_DTYPE
-from nanoai.optim import MuonAdamW, DistMuonAdamW
+from netis_amniota.common import get_dist_info, print0, COMPUTE_DTYPE
+from netis_amniota.optim import MuonAdamW, DistMuonAdamW
# GPTConfig is the single shared config (torch + MLX); re-exported here so
-# `from nanoai.gpt import GPTConfig` keeps working.
-from nanoai.model_config import GPTConfig, value_embed_table_map
+# `from netis_amniota.gpt import GPTConfig` keeps working.
+from netis_amniota.model_config import GPTConfig, value_embed_table_map
# Our custom Flash Attention module that automatically uses FA3 on Hopper+ and SDPA fallback elsewhere
-from nanoai.flash_attention import flash_attn
+from netis_amniota.flash_attention import flash_attn
# MSA (MiniMax-M3-style block-sparse attention) — self-contained subpackage
-from nanoai.msa import (
+from netis_amniota.msa import (
msa_block_sparse_attn_func, msa_gather_decode, update_msa_pooled_cache,
msa_fa4_block_sparse_attn, default_block_size,
)
-from nanoai.sparse_attention import SparseIndexer, dsa_attention, hisa_attention
+from netis_amniota.sparse_attention import SparseIndexer, dsa_attention, hisa_attention
def norm(x):
return F.rms_norm(x, (x.size(-1),)) # note that this will run in bf16, seems ok
@@ -668,7 +668,7 @@ class GPT(nn.Module):
def forward(self, idx, targets=None, kv_cache=None, loss_reduction='mean'):
if getattr(self, "_uses_nvfp4", False):
- from nanoai.nvfp4 import nvfp4_autocast
+ from netis_amniota.nvfp4 import nvfp4_autocast
with nvfp4_autocast():
return self._forward_impl_nvfp4(idx, targets, kv_cache, loss_reduction)
return self._forward_impl(idx, targets, kv_cache, loss_reduction)
@@ -831,7 +831,7 @@ class GPT(nn.Module):
ids = torch.tensor([tokens], dtype=torch.long, device=device) # add batch dim
for _ in range(max_tokens):
if getattr(self, "_uses_nvfp4", False):
- from nanoai.nvfp4 import disable_nvfp4_autocast
+ from netis_amniota.nvfp4 import disable_nvfp4_autocast
with disable_nvfp4_autocast():
logits = self.forward(ids) # (B, T, vocab_size)
else:
diff --git a/vendor/nanoai/model_config.py b/netis_amniota/model_config.py
similarity index 93%
rename from vendor/nanoai/model_config.py
rename to netis_amniota/model_config.py
index 53d77462a18c88fd1f22fdca010be06732832540..051c3e8474dfd51bb0a79f6fb9990a35bfab78e1 100644
--- a/vendor/nanoai/model_config.py
+++ b/netis_amniota/model_config.py
@@ -1,13 +1,13 @@
-"""Shared model configuration for NanoAI.
+"""Shared model configuration for Netis Amniota.
-The single ``GPTConfig`` used by BOTH backends — the torch model (``nanoai.gpt``)
-and the MLX model (``nanoai.mlx.gpt``) — so checkpoints (``model_*.pt`` +
+The single ``GPTConfig`` used by BOTH backends — the torch model (``netis_amniota.gpt``)
+and the MLX model (``netis_amniota.mlx.gpt``) — so checkpoints (``model_*.pt`` +
``meta_*.json`` carrying a top-level ``model_config``) interoperate across
backends: a checkpoint trained with torch loads in MLX and vice versa.
Pure stdlib (no torch, no mlx) so importing it is cheap and backend-agnostic;
-both ``nanoai.gpt`` and ``nanoai.mlx.gpt`` re-export ``GPTConfig`` from here, so
-existing ``from nanoai.gpt import GPTConfig`` / ``from nanoai.mlx.gpt import
+both ``netis_amniota.gpt`` and ``netis_amniota.mlx.gpt`` re-export ``GPTConfig`` from here, so
+existing ``from netis_amniota.gpt import GPTConfig`` / ``from netis_amniota.mlx.gpt import
GPTConfig`` keep working.
"""
from __future__ import annotations
@@ -24,8 +24,8 @@ def has_ve(layer_idx: int, n_layer: int) -> bool:
def value_embed_table_map(n_layer: int, ve_n_unique: int):
"""Return ``(table_keys, layer_to_key)`` for value-embedding sharing — the ONE
- source of truth so the torch (``nanoai.gpt``), MLX (``nanoai.mlx.gpt``) and ANE
- (``nanoai.ane.model``) backends key ``value_embeds`` identically and thus load
+ source of truth so the torch (``netis_amniota.gpt``), MLX (``netis_amniota.mlx.gpt``) and ANE
+ (``netis_amniota.ane.model``) backends key ``value_embeds`` identically and thus load
each other's checkpoints.
- ``ve_n_unique == -1`` (legacy per-layer): one table per ``has_ve`` layer,
@@ -76,14 +76,14 @@ class GPTConfig:
msa_index_mode: str = "pooled_k"
# MSA train/prefill kernel on "M" layers: "sdpa" (any head_dim, correct,
# overhead-bound) | "fa4" (FA4 native block-sparse, head_dim==64, fast).
- # Defaults to env NANOAI_MSA_KERNEL when left "sdpa" so experiments need not
+ # Defaults to env NETIS_AMNIOTA_MSA_KERNEL when left "sdpa" so experiments need not
# touch base_train; the resolved value is saved in the checkpoint meta.
msa_kernel: str = "sdpa"
# DSA/HISA optional sparse-attention hyperparameters. These are only used by
# torch backend layers whose window_pattern letter is D or H. The default
# indexer is the training-free Q/K proxy used by the sparse-attention-indexing
# experiments; "learned" instantiates explicit projection weights. Use
- # nanoai.sparse_attention.indexer_alignment_loss for auxiliary supervision
+ # netis_amniota.sparse_attention.indexer_alignment_loss for auxiliary supervision
# before hard top-k deployment.
sparse_token_budget: int = 512
sparse_block_size: int = 64
@@ -150,7 +150,7 @@ class GPTConfig:
def __post_init__(self):
if self.msa_kernel == "sdpa":
- self.msa_kernel = os.environ.get("NANOAI_MSA_KERNEL", "sdpa")
+ self.msa_kernel = os.environ.get("NETIS_AMNIOTA_MSA_KERNEL", "sdpa")
assert self.n_experts >= 0, f"n_experts must be >= 0 (got {self.n_experts})"
if self.n_experts > 0:
assert self.n_experts_active > 0, "n_experts_active (top-k) must be > 0 when n_experts > 0"
diff --git a/vendor/nanoai/msa/__init__.py b/netis_amniota/msa/__init__.py
similarity index 77%
rename from vendor/nanoai/msa/__init__.py
rename to netis_amniota/msa/__init__.py
index 995100a4a8db7d9c05fc2135eeb87a8ba02e574c..b6f3a0f6ed65982660e84556d40e3277efe1c39d 100644
--- a/vendor/nanoai/msa/__init__.py
+++ b/netis_amniota/msa/__init__.py
@@ -1,11 +1,11 @@
"""
-MSA: MiniMax-M3-style block-sparse attention for nanoai.
+MSA: MiniMax-M3-style block-sparse attention for netis_amniota.
Self-contained subpackage holding everything MSA-specific:
- attention.py : the block-sparse attention math (mask path + gather decode)
- runtime.py : model-side helpers (post-hoc enable, rotary extension)
-The only in-place hooks in the core model live in nanoai/gpt.py:
+The only in-place hooks in the core model live in netis_amniota/gpt.py:
- GPTConfig.msa_* fields + the "M" character in window_pattern
- CausalSelfAttention.forward dispatch to msa_block_sparse_attn_func /
msa_gather_decode
@@ -14,14 +14,14 @@ The only in-place hooks in the core model live in nanoai/gpt.py:
See docs / the design plan for the staged Stage-A (no-retrain probe) →
Stage-B (native training) rollout.
"""
-from nanoai.msa.attention import (
+from netis_amniota.msa.attention import (
msa_block_sparse_attn_func,
msa_gather_decode,
_msa_block_selection,
update_msa_pooled_cache,
)
-from nanoai.msa.runtime import apply_msa, num_msa_layers, extend_rotary
-from nanoai.msa.fa4 import (
+from netis_amniota.msa.runtime import apply_msa, num_msa_layers, extend_rotary
+from netis_amniota.msa.fa4 import (
msa_fa4_block_sparse_attn,
build_msa_block_sparse,
default_block_size,
diff --git a/vendor/nanoai/msa/attention.py b/netis_amniota/msa/attention.py
similarity index 100%
rename from vendor/nanoai/msa/attention.py
rename to netis_amniota/msa/attention.py
diff --git a/vendor/nanoai/msa/fa4.py b/netis_amniota/msa/fa4.py
similarity index 98%
rename from vendor/nanoai/msa/fa4.py
rename to netis_amniota/msa/fa4.py
index e92aa26512d96c084bfff958914283e04b3e420e..c91045b0a992208638ecf0101ede1ff9af97bfae 100644
--- a/vendor/nanoai/msa/fa4.py
+++ b/netis_amniota/msa/fa4.py
@@ -46,11 +46,11 @@ except Exception:
if _FA4_AVAILABLE:
# The SM100 block-sparse kernel asserts `sm_100 <= arch <= sm_110f`, but an
# upstream Arch-enum typo aliases sm_110f to sm_101f, so B300 (sm_103) trips a
- # kernel it supports. nanoai.flash_attention ships the rebind fix; the model
+ # kernel it supports. netis_amniota.flash_attention ships the rebind fix; the model
# path imports it via gpt.py, but apply it here too so the standalone MSA-FA4
# path (and its tests) works on B300 regardless of import order.
try:
- from nanoai.flash_attention import _patch_fa4_arch_enum as _patch_fa4_arch
+ from netis_amniota.flash_attention import _patch_fa4_arch_enum as _patch_fa4_arch
_patch_fa4_arch()
except Exception:
pass
diff --git a/vendor/nanoai/msa/runtime.py b/netis_amniota/msa/runtime.py
similarity index 97%
rename from vendor/nanoai/msa/runtime.py
rename to netis_amniota/msa/runtime.py
index 6f8523664e7e5a86029b3cca07fed1328570e80b..53d51887367949efe86631ab9a5eebabb876f2e1 100644
--- a/vendor/nanoai/msa/runtime.py
+++ b/netis_amniota/msa/runtime.py
@@ -1,6 +1,6 @@
"""
MSA runtime helpers (model-side glue, no attention math — that lives in
-nanoai.msa.attention).
+netis_amniota.msa.attention).
- apply_msa: POST-HOC enable MSA on a (dense-trained) model, used by the
no-retrain Stage-A probe. It flips the chosen layers to "msa" and sets the
@@ -11,7 +11,7 @@ nanoai.msa.attention).
otherwise trips the rotary-cache assert in GPT._forward_impl).
- num_msa_layers: count "msa" layers on a model.
"""
-from nanoai.common import COMPUTE_DTYPE
+from netis_amniota.common import COMPUTE_DTYPE
def apply_msa(model, *, msa_layers="auto", block_size=64, top_k_blocks=16,
diff --git a/vendor/nanoai/nvfp4.py b/netis_amniota/nvfp4.py
similarity index 97%
rename from vendor/nanoai/nvfp4.py
rename to netis_amniota/nvfp4.py
index 4ab2c287334b034b08f78cc1370192a6520abd53..50e181e00f9963c209a98792d685d008840969e5 100644
--- a/vendor/nanoai/nvfp4.py
+++ b/netis_amniota/nvfp4.py
@@ -1,6 +1,6 @@
-"""NVFP4 training for nanoai — using NVIDIA TransformerEngine.
+"""NVFP4 training for netis_amniota — using NVIDIA TransformerEngine.
-Drop-in replacement for nanoai's Float8Linear. Uses TransformerEngine's
+Drop-in replacement for netis_amniota's Float8Linear. Uses TransformerEngine's
te.Linear to quantize weights and activations to 4-bit floating point
(E2M1 format with 16-element micro-block scaling).
@@ -20,7 +20,7 @@ as te_linear.weight/bias.
This ensures:
1. model.parameters() yields only standard-shaped weight/bias tensors
- (same as nn.Linear), so nanoai's optimizer param group classification
+ (same as nn.Linear), so netis_amniota's optimizer param group classification
(Muon vs AdamW by parameter name/shape) works correctly.
2. TE's forward pass uses the same parameter objects (no copy overhead).
3. State dict keys are standard 'weight'/'bias' — cross-compatible with
@@ -250,7 +250,7 @@ def convert_to_float4_training(module, *, config=None, module_filter_fn=None):
config: Float4LinearConfig (accepted for API compat, currently unused).
module_filter_fn: Optional callable(module, fqn) -> bool.
"""
- from nanoai.fp8 import Float8Linear # avoid circular import
+ from netis_amniota.fp8 import Float8Linear # avoid circular import
if config is None:
config = Float4LinearConfig()
@@ -361,14 +361,14 @@ def probe_nvfp4_forward(
# ---------------------------------------------------------------------------
# DeepSeek-V4 ships a production FP4 recipe where ONLY the large, loss-insensitive
# MLP/expert weights are FP4; attention projections, the LM head, and embeddings
-# stay at higher precision (FP8/BF16). nanoai's stock `quant_module_filter`
+# stay at higher precision (FP8/BF16). netis_amniota's stock `quant_module_filter`
# selects purely by shape (dims %16, min>=128) and ignores `fqn`, so attention
# q/k/v/o and lm_head were ALSO FP4'd — a common cause of low-bit accuracy loss.
#
# `make_quant_filter` keeps the shape guards but additionally excludes sensitive
# modules by fully-qualified name. It only ever REDUCES the quantized set vs the
# shape-only filter, so it is a strictly-safer drop-in for A/B experiments.
-# Module names verified against nanoai/gpt.py: attention linears live under
+# Module names verified against netis_amniota/gpt.py: attention linears live under
# `.attn.` (c_q/c_k/c_v/c_proj/ve_gate); MLP under `.mlp.` (c_fc/c_proj); head is
# `lm_head`; embeddings are `wte` / `value_embeds`.
diff --git a/vendor/nanoai/optim.py b/netis_amniota/optim.py
similarity index 99%
rename from vendor/nanoai/optim.py
rename to netis_amniota/optim.py
index c14b354a4bac1be63e52b5c8a629d1011a1b4cbb..8eeebda094809fe0831dbc85bbeef9b555fdb536 100644
--- a/vendor/nanoai/optim.py
+++ b/netis_amniota/optim.py
@@ -71,7 +71,7 @@ NorMuon variance reduction: per-neuron/column adaptive learning rate that normal
update scales after orthogonalization (Muon's output has non-uniform scales across neurons).
https://arxiv.org/pdf/2510.05491
-Some of the changes in nanoai implementation:
+Some of the changes in netis_amniota implementation:
- Uses a simpler, more general approach to parameter grouping and stacking
- Uses a single fused kernel for the momentum -> polar_express -> variance_reduction -> update step
- Makes no assumptions about model architecture (e.g. that attention weights are fused into QKVO format)
diff --git a/vendor/nanoai/sparse_attention.py b/netis_amniota/sparse_attention.py
similarity index 99%
rename from vendor/nanoai/sparse_attention.py
rename to netis_amniota/sparse_attention.py
index dc7965af5b9c5f563395b3c62156f2c9f1df103c..c431b2142595727286bff22d1fc739fc9c1f221b 100644
--- a/vendor/nanoai/sparse_attention.py
+++ b/netis_amniota/sparse_attention.py
@@ -1,4 +1,4 @@
-"""DSA/HISA sparse-attention primitives for NanoAI.
+"""DSA/HISA sparse-attention primitives for Netis Amniota.
This module is deliberately conservative: it provides correct PyTorch reference
paths and an optional indexer scaffold, while leaving faster experimental
diff --git a/vendor/nanoai/tokenizer.py b/netis_amniota/tokenizer.py
similarity index 98%
rename from vendor/nanoai/tokenizer.py
rename to netis_amniota/tokenizer.py
index 30e70b63e0b9e6fd2f3926ac90363b0a73252ce4..aef13e7d72f0144512d0304bf5105af88748a5a2 100644
--- a/vendor/nanoai/tokenizer.py
+++ b/netis_amniota/tokenizer.py
@@ -106,7 +106,7 @@ class HuggingFaceTokenizer:
def _encode_one(self, text, prepend=None, append=None, num_threads=None):
# encode a single string
# prepend/append can be either a string of a special token or a token id directly.
- # num_threads is ignored (only used by the nanoai Tokenizer for parallel encoding)
+ # num_threads is ignored (only used by the netis_amniota Tokenizer for parallel encoding)
assert isinstance(text, str)
ids = []
if prepend is not None:
@@ -385,23 +385,23 @@ class RustBPETokenizer:
return ids
# -----------------------------------------------------------------------------
-# nanoai-specific convenience functions
+# netis_amniota-specific convenience functions
def get_tokenizer():
- from nanoai.common import get_base_dir
+ from netis_amniota.common import get_base_dir
base_dir = get_base_dir()
tokenizer_dir = os.path.join(base_dir, "tokenizer")
- # Qwen-style tokenizer dir (HF-format) has tokenizer_config.json; nanoai
+ # Qwen-style tokenizer dir (HF-format) has tokenizer_config.json; netis_amniota
# rustbpe dir has tokenizer.pkl. Detect and dispatch.
if os.path.exists(os.path.join(tokenizer_dir, "tokenizer_config.json")) \
and not os.path.exists(os.path.join(tokenizer_dir, "tokenizer.pkl")):
- from nanoai.agentic.tokenizer_qwen import QwenTokenizer
+ from netis_amniota.agentic.tokenizer_qwen import QwenTokenizer
return QwenTokenizer.from_directory(tokenizer_dir)
return RustBPETokenizer.from_directory(tokenizer_dir)
def get_token_bytes(device="cpu"):
import torch
- from nanoai.common import get_base_dir
+ from netis_amniota.common import get_base_dir
base_dir = get_base_dir()
tokenizer_dir = os.path.join(base_dir, "tokenizer")
token_bytes_path = os.path.join(tokenizer_dir, "token_bytes.pt")
diff --git a/tokenization_ravenguard.py b/tokenization_ravenguard.py
index 185f571075c0a3b35695cd55aa9f9063bdb1f2fc..20c23a5aaf737f92e72dbc46b473ff017b251b09 100644
--- a/tokenization_ravenguard.py
+++ b/tokenization_ravenguard.py
@@ -1,6 +1,6 @@
"""RavenGuard tokenizer (trust_remote_code).
-Thin ``PreTrainedTokenizer`` wrapper around the nanoai rustbpe/tiktoken encoding
+Thin ``PreTrainedTokenizer`` wrapper around the Netis Amniota rustbpe/tiktoken encoding
pickled in ``tokenizer.pkl``. ``encode``/``decode`` delegate straight to the
tiktoken ``Encoding`` so text round-trips byte-for-byte, matching how the guard
was trained. Vocab size is the tiktoken ``n_vocab`` (65536).
diff --git a/vendor/nanoai/__init__.py b/vendor/nanoai/__init__.py
deleted file mode 100644
index ff8f7426d714afcadb86cf90031ebee1f6c1bf8f..0000000000000000000000000000000000000000
--- a/vendor/nanoai/__init__.py
+++ /dev/null
@@ -1,20 +0,0 @@
-"""NanoAI — SOTA micro models, end to end.
-
-Formerly ``nanochat``. For back-compat we mirror the legacy ``NANOCHAT_*``
-environment variables onto their ``NANOAI_*`` equivalents (and vice versa) at
-package import time, so existing launch scripts / remote shells that still set
-``NANOCHAT_*`` keep working unchanged. This runs before any submodule reads an
-env var (importing ``nanoai.x`` executes this ``__init__`` first).
-"""
-import os as _os
-
-_ENV_SUFFIXES = (
- "BASE_DIR", "DTYPE", "ATTN_BACKEND", "MSA_KERNEL", "PACK_VERIFY",
- "DIR", "ROOT", "GROUNDED_JSONL", "TRAINING_PCAPS",
-)
-for _suf in _ENV_SUFFIXES:
- _old, _new = "NANOCHAT_" + _suf, "NANOAI_" + _suf
- if _new not in _os.environ and _old in _os.environ:
- _os.environ[_new] = _os.environ[_old]
- elif _old not in _os.environ and _new in _os.environ:
- _os.environ[_old] = _os.environ[_new]
diff --git a/vendor/nanoai/agentic/README.md b/vendor/nanoai/agentic/README.md
deleted file mode 100644
index 9e14ac47d5aef1c0a14cb7cc5e9fb67b8d06a8f7..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/README.md
+++ /dev/null
@@ -1,68 +0,0 @@
-# nanoai/agentic — shared agentic infrastructure
-
-The agentic post-training stack: pretrain on a web+code mix, SFT on tool-call
-conversations, DPO over agentic preferences, a family of RL / RLVR trainers, and
-offline tool-call + diagnostic evals. It stacks with the `--nvfp4` training stack
-and runs alongside the base `chat_sft` / `chat_rl` flow without modifying it.
-
-For day-to-day operation and the full pitfalls log see [CLAUDE.md](../../CLAUDE.md);
-research goals and playbook are in `goal.md` / `goal_next.md`.
-
-## Infra vs experiment boundary (Stage 4 outcome)
-
-This package is **shared agentic infrastructure** — the reusable library that the
-experiment groups under `experiments//scripts/` all build on. The Stage-4
-audit confirmed every module here is used **across multiple experiment groups or
-by core infra**, so none is scattered into a single group's `lib/` (that would
-create cross-group import edges). Evidence:
-
-| module | used by | verdict |
-|---|---|---|
-| `diag/` (graders, schemas, scenarios) | ksbr_pcap **+ vrmining** | shared |
-| `hard_tasks/` | ksbr_pcap + agentic_rl + mcts_inference | shared |
-| `mcts_inference` / `mcts_exec` / `mcts_hybrid` | mcts_inference + agentic_rl | shared |
-| `prefix_merging` | agentic_rl + polar repro | shared |
-| `opd_loss` | agentic_rl (+ `config.seg_weights` mirrors it) | shared loss lib |
-| `dapt_probes` | **`scripts/base_train`** + ksbr | infra-referenced |
-| `cross_tokenizer_kd` | **`teacher`** (this pkg) | infra-referenced |
-| `reward_*`, `dense_verifier`, `grpo_multi_rollout*`, `tokenizer*`, `teacher`, `value_head`, `turn_*`, `data/` | many | shared |
-
-**The experiment/infra split is therefore: entrypoint scripts/runs/docs live in
-`experiments//`; this directory is the shared library they import.** Infra
-here imports **no** experiment module (verified). The one infra-referenced
-entrypoint kept in `scripts/` is `eval_llm_judge_rc_hit.py`
-(`reward_capbench` lazily imports `_judge_one` from it).
-
-## Tool-call format (mirrors nanocode)
-
-```
-<|tool_call_start|>Bash<|tool_arg|>command<|tool_val|>ls -la<|tool_call_end|>
-```
-
-6 agentic special tokens extend the 9 base tokens (15 total):
-`<|tool_call_start|>`, `<|tool_arg|>`, `<|tool_val|>`, `<|tool_call_end|>`,
-`<|tool_result_start|>`, `<|tool_result_end|>`. See `tokenizer.py` /
-`scripts/agentic_tok_train.py`.
-
-## Modules
-
-- **Data**: `code_dataloader.py` (web+code Bernoulli mixer), `synth_dataloader.py`,
- `data/` (json/hf datasets, mixtures, `data/common.py` SYSTEM_PROMPT, pretrain/pcap/ebook/synth iterators).
-- **Tokenizers**: `tokenizer.py` (nanoai), `tokenizer_qwen.py` (Qwen teacher), `cross_tokenizer_kd.py`.
-- **Rewards**: `reward_mlr.py` (multi-level per-turn; weights from `nanoai.config.reward()`),
- `reward_mt.py`, `reward.py`, `reward_capbench.py`, `dense_verifier.py`, `opd_loss.py` (on-policy distillation).
-- **RL plumbing**: `grpo_multi_rollout{,_batched}.py`, `prefix_merging.py`, `value_head.py`,
- `turn_critic.py`, `turn_weight.py`, `mc_turn_value.py`, `ocaw.py`.
-- **Inference / search**: `mcts_inference.py`, `mcts_exec.py`, `mcts_hybrid.py`.
-- **Eval**: `eval_loop.py`, `eval_tasks.py`, `eval_pcap_tasks.py`, `eval_sandbox.py`
- (the sandbox path hard-depends on `firejail` — verify per machine).
-- **diag/**: graders (v1/v2/v3), schemas, chain_adapters, scenario generation.
-- **hard_tasks/**: task families (K-fault / S-wrapper / S-yaml / S-toolcall / SWE) for curriculum + grading.
-
-## Trainers (in `scripts/`)
-
-`agentic_sft.py`, `agentic_dpo.py`, and the RL family
-(`agentic_{gtpo,dapo,real,tree_grpo,grpo_v2,gft_full,…}.py`). They share
-`pack_for_grad` / `setup_optimizer` patterns — see CLAUDE.md's pitfalls
-(per-token advantage off-by-one, single-`--lr` clobber, length-match) before
-editing any of them.
diff --git a/vendor/nanoai/agentic/__init__.py b/vendor/nanoai/agentic/__init__.py
deleted file mode 100644
index 76fed84acfb110bb3c459fcd4ee43e52b42c576f..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""Agentic code workflow extensions on top of nanoai.
-
-Adds: code-data mixing in pretrain, tool-call tokenizer, agentic SFT/DPO,
-and a tool-call validation eval. Designed to stack with the existing NVFP4
-training stack — every script accepts --nvfp4 and reuses Float4Linear from
-nanoai.nvfp4.
-"""
diff --git a/vendor/nanoai/agentic/code_dataloader.py b/vendor/nanoai/agentic/code_dataloader.py
deleted file mode 100644
index 3ecaf609653d64ca30090e8a8ea663ad6184af7d..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/code_dataloader.py
+++ /dev/null
@@ -1,407 +0,0 @@
-"""Code-mixing pretrain dataloader.
-
-Wraps nanoai's BOS-aligned best-fit packing but draws documents from up to
-three parquet sources with a categorical (Dirichlet-multinomial) mixture:
-
- - web : nanoai.dataset (climbmix or fineweb-edu, whatever's on disk)
- - code : nanoai.agentic.data.pretrain_data → smohammadi/the-stack-v2-python-shuffle
- - pcap : nanoai.agentic.data.pcap_data (network/protocol domain corpus)
-
-`code_ratio + pcap_ratio` must be ≤ 1; the remainder is `web_ratio`. Two-way
-defaults (no pcap) preserve the original 2-stream behaviour exactly.
-
-Resume state is intentionally simplified to a single epoch counter — the
-fine-grained (pq_idx, rg_idx) state of the upstream loader doesn't extend
-naturally to interleaved streams. For multi-day runs we'd want real
-multi-stream resume, but for d3/d6/d12 speedruns a fresh start at each step is
-fine.
-"""
-
-from __future__ import annotations
-
-import json
-import os
-import random
-from pathlib import Path
-
-import pyarrow.parquet as pq
-import torch
-
-from nanoai.common import get_dist_info, print0
-from nanoai.dataset import list_parquet_files as list_web_parquet_files
-from nanoai.agentic.data.pretrain_data import DATA_DIR as _CODE_BASE_DATA_DIR
-from nanoai.agentic.data.pcap_data import (
- DATASET_NAME as _PCAP_DATASET_NAME,
- list_pcap_parquet_files,
-)
-from nanoai.agentic.data.synth_data import (
- list_synth_parquet_files,
- synth_doc_batches,
-)
-from nanoai.agentic.data.vrcpt_data import (
- DATASET_NAME as _VRCPT_DATASET_NAME,
- list_vrcpt_parquet_files,
-)
-from nanoai.agentic.data.ebook_data import (
- DATASET_NAME as _EBOOK_DATASET_NAME,
- list_ebook_parquet_files,
-)
-
-
-_CODE_DATASET_NAME = "the-stack-v2-dedup"
-# Default location for KSR Phase-2 synth SWE triples; override via $SYNTH_SWE_DIR.
-_SYNTH_SWE_DIR_DEFAULT = Path.home() / "Miros" / "mlros" / "registry_data" / "synth_swe" / "v1"
-
-
-def _list_code_parquet_files() -> list[Path]:
- code_dir = _CODE_BASE_DATA_DIR / _CODE_DATASET_NAME
- if not code_dir.exists():
- return []
- return sorted([code_dir / f for f in code_dir.iterdir() if f.suffix == ".parquet"])
-
-
-def _list_synth_swe_jsonl_files(synth_swe_dir: Path | None = None) -> list[Path]:
- """List cpt.jsonl files produced by scripts/synth_swe_triples.py
- (one per SWE subtype). Empty if Phase-2 synth has not run yet."""
- d = synth_swe_dir or Path(os.environ.get("SYNTH_SWE_DIR", _SYNTH_SWE_DIR_DEFAULT))
- if not d.exists():
- return []
- return sorted(d.glob("*/cpt.jsonl"))
-
-
-def _jsonl_doc_batches_from(paths, ddp_rank, ddp_world_size, tokenizer_batch_size):
- """Infinite per-rank iterator over JSONL shards with a `text` field.
-
- Each line is parsed as JSON and the `text` field extracted. Lines are
- striped across ranks by index so each rank sees a disjoint subset
- (no parquet row-group concept here)."""
- while True:
- for filepath in paths:
- with open(filepath, "r", encoding="utf-8") as f:
- batch: list[str] = []
- for i, line in enumerate(f):
- if i % ddp_world_size != ddp_rank:
- continue
- line = line.strip()
- if not line:
- continue
- try:
- obj = json.loads(line)
- except json.JSONDecodeError:
- continue
- txt = obj.get("text")
- if isinstance(txt, str) and txt:
- batch.append(txt)
- if len(batch) >= tokenizer_batch_size:
- yield batch
- batch = []
- if batch:
- yield batch
-
-
-def _doc_batches_from(paths, ddp_rank, ddp_world_size, tokenizer_batch_size):
- """Infinite per-rank iterator: yields list[str] batches."""
- while True:
- for filepath in paths:
- pf = pq.ParquetFile(filepath)
- for rg_idx in range(ddp_rank, pf.num_row_groups, ddp_world_size):
- rg = pf.read_row_group(rg_idx)
- batch = rg.column("text").to_pylist()
- for i in range(0, len(batch), tokenizer_batch_size):
- yield batch[i : i + tokenizer_batch_size]
-
-
-def _split_paths(paths, split):
- if not paths:
- return []
- return paths[:-1] if split == "train" else paths[-1:]
-
-
-def _mixed_document_batches(
- split: str,
- code_ratio: float,
- pcap_ratio: float,
- tokenizer_batch_size: int,
- seed: int = 0,
- synth_swe_ratio: float = 0.0,
- synth_swe_dir: Path | None = None,
- synth_pleias_ratio: float = 0.0,
- vrcpt_ratio: float = 0.0,
- ebook_ratio: float = 0.0,
-):
- assert 0.0 <= code_ratio <= 1.0, f"bad code_ratio={code_ratio}"
- assert 0.0 <= pcap_ratio <= 1.0, f"bad pcap_ratio={pcap_ratio}"
- assert 0.0 <= synth_swe_ratio <= 1.0, f"bad synth_swe_ratio={synth_swe_ratio}"
- assert 0.0 <= synth_pleias_ratio <= 1.0, f"bad synth_pleias_ratio={synth_pleias_ratio}"
- assert 0.0 <= vrcpt_ratio <= 1.0, f"bad vrcpt_ratio={vrcpt_ratio}"
- assert 0.0 <= ebook_ratio <= 1.0, f"bad ebook_ratio={ebook_ratio}"
- assert code_ratio + pcap_ratio + synth_swe_ratio + synth_pleias_ratio + vrcpt_ratio + ebook_ratio <= 1.0 + 1e-9, (
- f"code+pcap+synth_swe+synth_pleias+vrcpt+ebook = "
- f"{code_ratio+pcap_ratio+synth_swe_ratio+synth_pleias_ratio+vrcpt_ratio+ebook_ratio} > 1"
- )
-
- ddp, ddp_rank, _local, ddp_world_size = get_dist_info()
- warn = (ddp_rank == 0) and (split == "train")
-
- web_paths = list_web_parquet_files(warn_on_legacy=warn)
- assert web_paths, "No web (climbmix/fineweb-edu) parquet shards found; run nanoai.dataset first."
- web_paths = [Path(p) for p in _split_paths(web_paths, split)]
-
- code_paths = _split_paths(_list_code_parquet_files(), split)
- if code_ratio > 0 and not code_paths:
- raise FileNotFoundError(
- "code_ratio > 0 but no the-stack-v2-dedup parquet found. Run:\n"
- " python -m nanoai.agentic.data.pretrain_data -d the-stack-v2-dedup -n 8"
- )
-
- pcap_paths = _split_paths(list_pcap_parquet_files(), split)
- if pcap_ratio > 0 and not pcap_paths:
- raise FileNotFoundError(
- f"pcap_ratio > 0 but no {_PCAP_DATASET_NAME} parquet found. Run:\n"
- " python -m nanoai.agentic.data.pcap_data prepare --src /path/to/pretrain.jsonl"
- )
-
- synth_swe_paths = _list_synth_swe_jsonl_files(synth_swe_dir)
- synth_swe_paths = _split_paths(synth_swe_paths, split)
- if synth_swe_ratio > 0 and not synth_swe_paths:
- raise FileNotFoundError(
- "synth_swe_ratio > 0 but no synth_swe/v1//cpt.jsonl found. Run:\n"
- " python -m scripts.synth_swe_triples --out ~/Miros/mlros/registry_data/synth_swe/v1 --n-per-subtype 500"
- )
-
- synth_pleias_paths = _split_paths(list_synth_parquet_files(), split)
- if synth_pleias_ratio > 0 and not synth_pleias_paths:
- raise FileNotFoundError(
- "synth_pleias_ratio > 0 but no synth-pleias parquet found. Expected at "
- "$NANOAI_BASE_DIR/code_pretrain_data/synth-pleias/ (symlink to PleIAs/SYNTH)."
- )
-
- vrcpt_paths = _split_paths(list_vrcpt_parquet_files(), split)
- if vrcpt_ratio > 0 and not vrcpt_paths:
- raise FileNotFoundError(
- f"vrcpt_ratio > 0 but no {_VRCPT_DATASET_NAME} parquet found. Run:\n"
- " python -m nanoai.agentic.data.vrcpt_data prepare --src experiments/vrmining/datasets/cpt_corpus.jsonl"
- )
-
- ebook_paths = _split_paths([Path(p) for p in list_ebook_parquet_files()], split)
- if ebook_ratio > 0 and not ebook_paths:
- raise FileNotFoundError(
- f"ebook_ratio > 0 but no {_EBOOK_DATASET_NAME} parquet found. Run:\n"
- " python -m nanoai.agentic.data.ebook_data prepare --md-root ~/ebook/markdown"
- )
-
- web_iter = _doc_batches_from(web_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- code_iter = (
- _doc_batches_from(code_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- if code_paths
- else None
- )
- pcap_iter = (
- _doc_batches_from(pcap_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- if pcap_paths
- else None
- )
- synth_swe_iter = (
- _jsonl_doc_batches_from(synth_swe_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- if synth_swe_paths
- else None
- )
- synth_pleias_iter = (
- synth_doc_batches(synth_pleias_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- if synth_pleias_paths
- else None
- )
- vrcpt_iter = (
- _doc_batches_from(vrcpt_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- if vrcpt_paths
- else None
- )
- ebook_iter = (
- _doc_batches_from(ebook_paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
- if ebook_paths
- else None
- )
-
- if ddp_rank == 0:
- web_ratio = max(0.0, 1.0 - code_ratio - pcap_ratio - synth_swe_ratio - synth_pleias_ratio - vrcpt_ratio - ebook_ratio)
- print0(
- f"[code-mixer] split={split} "
- f"ratios web={web_ratio:.2f} code={code_ratio:.2f} pcap={pcap_ratio:.2f} "
- f"synth_swe={synth_swe_ratio:.2f} synth_pleias={synth_pleias_ratio:.2f} "
- f"vrcpt={vrcpt_ratio:.2f} ebook={ebook_ratio:.2f} "
- f"shards web={len(web_paths)} code={len(code_paths)} "
- f"pcap={len(pcap_paths)} synth_swe={len(synth_swe_paths)} "
- f"synth_pleias={len(synth_pleias_paths)} vrcpt={len(vrcpt_paths)} ebook={len(ebook_paths)}"
- )
-
- rng = random.Random(seed + ddp_rank)
- counts = {"web": 0, "code": 0, "pcap": 0, "synth_swe": 0, "synth_pleias": 0, "vrcpt": 0, "ebook": 0}
- log_every = 10000
- # Decision order: synth_swe, synth_pleias, vrcpt, ebook, pcap, code, web — additive slices of [0,1).
- while True:
- r = rng.random()
- if synth_swe_iter is not None and r < synth_swe_ratio:
- counts["synth_swe"] += 1
- yield next(synth_swe_iter)
- elif synth_pleias_iter is not None and r < synth_swe_ratio + synth_pleias_ratio:
- counts["synth_pleias"] += 1
- yield next(synth_pleias_iter)
- elif vrcpt_iter is not None and r < synth_swe_ratio + synth_pleias_ratio + vrcpt_ratio:
- counts["vrcpt"] += 1
- yield next(vrcpt_iter)
- elif ebook_iter is not None and r < synth_swe_ratio + synth_pleias_ratio + vrcpt_ratio + ebook_ratio:
- counts["ebook"] += 1
- yield next(ebook_iter)
- elif pcap_iter is not None and r < synth_swe_ratio + synth_pleias_ratio + vrcpt_ratio + ebook_ratio + pcap_ratio:
- counts["pcap"] += 1
- yield next(pcap_iter)
- elif code_iter is not None and r < synth_swe_ratio + synth_pleias_ratio + vrcpt_ratio + ebook_ratio + pcap_ratio + code_ratio:
- counts["code"] += 1
- yield next(code_iter)
- else:
- counts["web"] += 1
- yield next(web_iter)
- total = sum(counts.values())
- if ddp_rank == 0 and total % log_every == 0:
- print0(
- f"[code-mixer] draws={total} "
- f"web={counts['web']/total:.3f} code={counts['code']/total:.3f} "
- f"pcap={counts['pcap']/total:.3f} synth_swe={counts['synth_swe']/total:.3f} "
- f"synth_pleias={counts['synth_pleias']/total:.3f} vrcpt={counts['vrcpt']/total:.3f} "
- f"ebook={counts['ebook']/total:.3f}"
- )
-
-
-def tokenizing_code_mixed_data_loader(
- tokenizer,
- B: int,
- T: int,
- split: str,
- code_ratio: float = 0.2,
- pcap_ratio: float = 0.0,
- synth_swe_ratio: float = 0.0,
- synth_swe_dir: Path | None = None,
- synth_pleias_ratio: float = 0.0,
- vrcpt_ratio: float = 0.0,
- ebook_ratio: float = 0.0,
- tokenizer_threads: int = 4,
- tokenizer_batch_size: int = 128,
- device: str = "cuda",
- buffer_size: int = 1000,
- seed: int = 0,
-):
- """BOS-aligned best-fit packing over a (web, code, pcap, synth_swe, synth_pleias, vrcpt, ebook) mixture.
-
- Yields (inputs, targets) GPU tensors of shape (B, T). Drop-in for
- nanoai's `tokenizing_distributed_data_loader_bos_bestfit`.
- """
- assert split in ("train", "val")
- row_capacity = T + 1
- bos_token = tokenizer.get_bos_token_id()
-
- batches = _mixed_document_batches(
- split, code_ratio, pcap_ratio, tokenizer_batch_size, seed=seed,
- synth_swe_ratio=synth_swe_ratio, synth_swe_dir=synth_swe_dir,
- synth_pleias_ratio=synth_pleias_ratio,
- vrcpt_ratio=vrcpt_ratio,
- ebook_ratio=ebook_ratio,
- )
- doc_buffer: list[list[int]] = []
-
- def refill_buffer():
- doc_batch = next(batches)
- token_lists = tokenizer.encode(
- doc_batch, prepend=bos_token, num_threads=tokenizer_threads
- )
- for tokens in token_lists:
- doc_buffer.append(tokens)
-
- use_cuda = device == "cuda"
- row_buffer = torch.empty((B, row_capacity), dtype=torch.long)
- cpu_buffer = torch.empty(2 * B * T, dtype=torch.long, pin_memory=use_cuda)
- gpu_buffer = torch.empty(2 * B * T, dtype=torch.long, device=device)
- cpu_inputs = cpu_buffer[: B * T].view(B, T)
- cpu_targets = cpu_buffer[B * T :].view(B, T)
- inputs = gpu_buffer[: B * T].view(B, T)
- targets = gpu_buffer[B * T :].view(B, T)
-
- while True:
- for row_idx in range(B):
- pos = 0
- while pos < row_capacity:
- while len(doc_buffer) < buffer_size:
- refill_buffer()
- remaining = row_capacity - pos
-
- best_idx = -1
- best_len = 0
- for i, doc in enumerate(doc_buffer):
- doc_len = len(doc)
- if doc_len <= remaining and doc_len > best_len:
- best_idx = i
- best_len = doc_len
-
- if best_idx >= 0:
- doc = doc_buffer.pop(best_idx)
- doc_len = len(doc)
- row_buffer[row_idx, pos : pos + doc_len] = torch.tensor(doc, dtype=torch.long)
- pos += doc_len
- else:
- shortest_idx = min(range(len(doc_buffer)), key=lambda i: len(doc_buffer[i]))
- doc = doc_buffer.pop(shortest_idx)
- row_buffer[row_idx, pos : pos + remaining] = torch.tensor(
- doc[:remaining], dtype=torch.long
- )
- pos += remaining
-
- cpu_inputs.copy_(row_buffer[:, :-1])
- cpu_targets.copy_(row_buffer[:, 1:])
- gpu_buffer.copy_(cpu_buffer, non_blocking=use_cuda)
- yield inputs, targets
-
-
-def tokenizing_code_mixed_data_loader_with_state(
- tokenizer,
- B: int,
- T: int,
- split: str,
- code_ratio: float = 0.2,
- pcap_ratio: float = 0.0,
- synth_swe_ratio: float = 0.0,
- synth_swe_dir: Path | None = None,
- synth_pleias_ratio: float = 0.0,
- vrcpt_ratio: float = 0.0,
- ebook_ratio: float = 0.0,
- tokenizer_threads: int = 4,
- tokenizer_batch_size: int = 128,
- device: str = "cuda",
- resume_state_dict=None,
- buffer_size: int = 1000,
- seed: int = 0,
-):
- """Same as `tokenizing_code_mixed_data_loader` but yields a third element
- `state_dict` to match the signature of nanoai's
- `tokenizing_distributed_data_loader_with_state_bos_bestfit`.
-
- The state_dict is a stub (`{"epoch": 1, "pq_idx": 0, "rg_idx": 0}`) — true
- resume across the multi-stream mixture is not implemented. Caller's
- `resume_state_dict` is ignored.
- """
- base = tokenizing_code_mixed_data_loader(
- tokenizer, B, T, split,
- code_ratio=code_ratio,
- pcap_ratio=pcap_ratio,
- synth_swe_ratio=synth_swe_ratio,
- synth_swe_dir=synth_swe_dir,
- synth_pleias_ratio=synth_pleias_ratio,
- vrcpt_ratio=vrcpt_ratio,
- ebook_ratio=ebook_ratio,
- tokenizer_threads=tokenizer_threads,
- tokenizer_batch_size=tokenizer_batch_size,
- device=device,
- buffer_size=buffer_size,
- seed=seed,
- )
- stub_state = {"epoch": 1, "pq_idx": 0, "rg_idx": 0}
- for inputs, targets in base:
- yield inputs, targets, stub_state
diff --git a/vendor/nanoai/agentic/cross_tokenizer_kd.py b/vendor/nanoai/agentic/cross_tokenizer_kd.py
deleted file mode 100644
index 65f99a7b8e1fc5046e3f23575e156882ada64dc0..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/cross_tokenizer_kd.py
+++ /dev/null
@@ -1,225 +0,0 @@
-"""Byte-level alignment between heterogeneous BPE tokenizers (e.g. nanoai vs Qwen).
-
-Use case: the OPD student emits a sequence of nanoai token ids; we want
-per-token teacher logp from a teacher that uses a DIFFERENT BPE tokenizer.
-
-Approach:
- 1. Both tokenizers decode their token sequences to the SAME byte string `s`.
- - nanoai: `RustBPETokenizer.enc.decode_single_token_bytes(tid)` returns
- the raw bytes for one BPE merge. Concatenating these for the full
- sequence reproduces the original byte string (modulo special tokens).
- - HF (Qwen): `tokenizer.convert_ids_to_tokens(ids)` gives sub-strings; we
- cumulate bytes through `decode(ids)` of each prefix to get offsets.
- 2. Each token gets a byte interval `[start, end)` in `s`. We then sum the
- teacher (Qwen) per-token logps that fall inside the student token's
- interval — by the chain rule, sum of consecutive log p(y_t | y_` etc.): nanoai reserves these as
- single ids that decode to multi-char strings (`<|...|>`), which Qwen will
- re-tokenize into ~5-8 BPE tokens. The byte interval still covers the same
- bytes, so the sum still aligns correctly. The CALLER may mask these
- positions out of the OPD loss because the teacher has no special-token
- semantic for them — we expose `is_special_token` per offset for that.
- - **Misalignment at byte level**: nanoai tokenizer may emit bytes that are
- invalid UTF-8 (in the middle of a multi-byte char). When we concat to
- build the byte string and Qwen tries to re-tokenize via decode→bytes, the
- boundary may shift by a few bytes. We tolerate up to 1-byte mismatch via
- fuzzy match in the alignment loop.
-"""
-
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import List, Optional, Tuple
-
-
-@dataclass
-class TokenSpan:
- """Byte interval [start, end) and decoded text of one BPE token."""
- token_id: int
- start: int # byte offset (inclusive)
- end: int # byte offset (exclusive)
- text: str # decoded bytes (lossy for non-utf8)
- is_special: bool
-
-
-def compute_nanochat_offsets(
- tokenizer, token_ids: List[int]
-) -> Tuple[List[TokenSpan], bytes]:
- """Compute byte offsets for a nanoai token sequence.
-
- Returns (spans, byte_string) where spans[i] is the TokenSpan for token i.
- The byte_string is the concatenation of all per-token bytes — when the
- student emitted a "normal" text trace this equals the natural utf-8
- encoding of the decoded text.
-
- For special tokens (BOS, tool_call_start, etc.), we use the rendered
- `<|name|>` form so the Qwen tokenizer re-tokenizes them into multiple BPE
- units; this is correct as long as the byte boundary aligns to the rendered
- form's edges, which it always does because tiktoken stores specials as
- these exact ascii strings internally.
- """
- spans: List[TokenSpan] = []
- cur_bytes = bytearray()
-
- # Qwen path: HF-backed wrapper. id_to_token returns the special-token name
- # (e.g. '<|im_start|>') for added tokens, or the BPE piece for regular
- # tokens. Decoding the full id list with skip_special_tokens=False produces
- # the same string as concatenating per-token decode_single_token; we use
- # per-token decode to get accurate byte spans.
- if getattr(tokenizer, "chat_kind", None) == "qwen":
- hf = tokenizer.tok
- added = hf.added_tokens_decoder if hasattr(hf, "added_tokens_decoder") else {}
- for tid in token_ids:
- start = len(cur_bytes)
- if tid in added:
- # Added/special token — emit its surface form as bytes.
- text = added[tid].content if hasattr(added[tid], "content") else str(added[tid])
- tb = text.encode("utf-8")
- is_special = True
- else:
- # Regular BPE token. decode([tid]) returns its surface string.
- text = hf.decode([tid], skip_special_tokens=False)
- tb = text.encode("utf-8")
- is_special = False
- cur_bytes.extend(tb)
- end = len(cur_bytes)
- spans.append(TokenSpan(tid, start, end, text, is_special))
- return spans, bytes(cur_bytes)
-
- # nanoai tiktoken path.
- enc = tokenizer.enc # tiktoken.Encoding
-
- # Build a set of special-token ids for fast lookup.
- special_ids = set(enc._special_tokens.values()) if hasattr(enc, "_special_tokens") else set()
-
- for tid in token_ids:
- start = len(cur_bytes)
- if tid in special_ids:
- # Render the special token as its name string (e.g. "<|bos|>")
- # so byte alignment is well-defined.
- name = None
- for k, v in enc._special_tokens.items():
- if v == tid:
- name = k
- break
- text = name or f"<|sp_{tid}|>"
- tb = text.encode("utf-8")
- is_special = True
- else:
- tb = enc.decode_single_token_bytes(tid)
- try:
- text = tb.decode("utf-8")
- except UnicodeDecodeError:
- text = tb.decode("utf-8", errors="replace")
- is_special = False
- cur_bytes.extend(tb)
- end = len(cur_bytes)
- spans.append(TokenSpan(tid, start, end, text, is_special))
-
- return spans, bytes(cur_bytes)
-
-
-def compute_hf_offsets_from_decoded(
- token_ids: List[int], decoded_tokens: List[str]
-) -> List[TokenSpan]:
- """Compute byte offsets directly from pre-decoded token strings.
-
- vLLM's `prompt_logprobs` response includes `decoded_token` per position,
- e.g. " capital", " of", " France". We use those directly to build byte
- offsets, sidestepping the need for the HF tokenizer to be present locally.
-
- token_ids and decoded_tokens must be parallel, same length.
- """
- assert len(token_ids) == len(decoded_tokens)
- spans: List[TokenSpan] = []
- cur_bytes = bytearray()
- for tid, text in zip(token_ids, decoded_tokens):
- start = len(cur_bytes)
- tb = text.encode("utf-8")
- cur_bytes.extend(tb)
- end = len(cur_bytes)
- spans.append(TokenSpan(tid, start, end, text, False))
- return spans
-
-
-def align_logps(
- target_spans: List[TokenSpan],
- source_spans: List[TokenSpan],
- source_logps: List[Optional[float]],
-) -> List[Optional[float]]:
- """Distribute source logps proportionally over each target byte interval.
-
- target_spans: e.g. nanoai student tokens — we want one logp per target.
- source_spans: e.g. Qwen teacher tokens — we have one logp per source.
- source_logps: list of float (or None for source position 0,
- which has no logp because there's no prior context).
-
- Returns: list of Optional[float], length == len(target_spans). None means
- "couldn't compute" (target span overlapping with source[0] which has
- no logp, or no source coverage at all).
-
- Bug-fix note (codex adversarial review 2026-05-18, finding #2):
- Previously, when a single source (teacher) token spanned two target
- (student) byte ranges, its FULL logp was added to BOTH targets — a
- silent 2x double-counting. We now allocate by overlap fraction:
- target[i] += source[k].logp × (overlap_bytes / source_byte_length)
- This preserves the chain-rule invariant Σ_target_in_S aligned ≈ S.logp.
- A target is invalid (None) only if no overlap or any overlapping source
- has missing logp.
- """
- aligned: List[Optional[float]] = [None] * len(target_spans)
- n_src = len(source_spans)
-
- # Two-pointer scan.
- j = 0
- for i, t in enumerate(target_spans):
- # Advance j past any source token strictly before t.
- while j < n_src and source_spans[j].end <= t.start:
- j += 1
- # Distribute source logps that overlap t, weighted by overlap fraction.
- total = 0.0
- any_overlap = False
- any_missing = False
- k = j
- while k < n_src and source_spans[k].start < t.end:
- s = source_spans[k]
- ov = min(t.end, s.end) - max(t.start, s.start)
- if ov > 0:
- if source_logps[k] is None:
- any_missing = True
- else:
- s_len = s.end - s.start
- weight = (ov / s_len) if s_len > 0 else 1.0
- total += source_logps[k] * weight
- any_overlap = True
- k += 1
- if any_overlap and not any_missing:
- aligned[i] = total
- else:
- aligned[i] = None # either no coverage or hit the unknown source[0]
- return aligned
-
-
-def align_qwen_logp_to_nano(
- nano_tokenizer,
- nano_token_ids: List[int],
- qwen_token_ids: List[int],
- qwen_decoded_tokens: List[str],
- qwen_logps: List[Optional[float]],
-) -> Tuple[List[Optional[float]], List[bool]]:
- """One-call convenience wrapper.
-
- Returns (aligned_logps, is_special_mask) for the nanoai sequence.
- aligned_logps[i] = sum of qwen_logps over qwen tokens that cover nano
- token i's bytes, or None if can't align.
- is_special_mask[i] = True iff nano token i is a special token (caller
- should typically mask these out of OPD loss).
- """
- nano_spans, _ = compute_nanochat_offsets(nano_tokenizer, nano_token_ids)
- qwen_spans = compute_hf_offsets_from_decoded(qwen_token_ids, qwen_decoded_tokens)
- aligned = align_logps(nano_spans, qwen_spans, qwen_logps)
- specials = [s.is_special for s in nano_spans]
- return aligned, specials
diff --git a/vendor/nanoai/agentic/dapt_probes.py b/vendor/nanoai/agentic/dapt_probes.py
deleted file mode 100644
index c761afe9b977788797f72938e6e53b32b61f3a93..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/dapt_probes.py
+++ /dev/null
@@ -1,194 +0,0 @@
-"""DAPT/CPT in-training probes shared with `scripts.base_train`.
-
-Pure functions kept outside the training script so they can be unit-tested
-without booting the full DDP / NVFP4 / wandb stack. Each helper is small,
-side-effect-free except for the obvious (parquet I/O, model.eval()), and
-exposes verbose `print0` logging so an operator watching the train log can
-reconstruct exactly what the probe did.
-
-Three helpers:
-
- - `build_domain_val_loader(domain, ...)` — return a single-source val loader
- over web | code | pcap, regardless of the active training mix. Used to
- compute per-domain BPB during eval ticks.
-
- - `load_l0_tasks(repo_root)` — read tshark_eval_agentic L0 vocab tasks. Memo-
- izes the parse so repeated probes don't re-read the JSONL.
-
- - `run_l0_probe(model_for_gen, tokenizer, ..., disable_fp8_ctx=...)` — greedy
- decode N L0 questions with a 3-shot prefix, return string-match accuracy.
-"""
-
-from __future__ import annotations
-
-import json
-import random
-import time
-from contextlib import nullcontext
-from pathlib import Path
-from typing import Callable, ContextManager, Optional
-
-from nanoai.common import print0
-
-
-# -----------------------------------------------------------------------------
-# Per-domain val loader
-
-DOMAIN_RATIOS = {
- "web": {"code_ratio": 0.0, "pcap_ratio": 0.0},
- "code": {"code_ratio": 1.0, "pcap_ratio": 0.0},
- "pcap": {"code_ratio": 0.0, "pcap_ratio": 1.0},
-}
-
-
-def build_domain_val_loader(domain: str, tokenizer, device_batch_size: int,
- max_seq_len: int, device: str):
- """Construct a single-source val loader for {web, code, pcap}.
-
- Bypasses any monkey-patched train loader so per-domain BPB is unaffected by
- the active training mixture. Returns the same yield shape as the patched
- loader (yields (inputs, targets) — no state_dict).
- """
- assert domain in DOMAIN_RATIOS, (
- f"unknown domain {domain!r}, expected one of {sorted(DOMAIN_RATIOS)}"
- )
- # Local import: code_dataloader pulls in pyarrow, torch, etc. — keep that
- # cost off the import path of any caller that doesn't actually probe.
- from nanoai.agentic.code_dataloader import tokenizing_code_mixed_data_loader
-
- ratios = DOMAIN_RATIOS[domain]
- print0(f"[dapt-probe] building val loader for domain={domain!r} ratios={ratios}")
- return tokenizing_code_mixed_data_loader(
- tokenizer, device_batch_size, max_seq_len, split="val",
- code_ratio=ratios["code_ratio"],
- pcap_ratio=ratios["pcap_ratio"],
- device=device,
- )
-
-
-# -----------------------------------------------------------------------------
-# L0 probe
-
-_DEFAULT_L0_PATH = Path(__file__).resolve().parents[2] / "tshark_eval_agentic" / "eval" / "tasks_l0_vocab.jsonl"
-
-_L0_CACHE: dict[Path, list[dict]] = {}
-
-
-def load_l0_tasks(path: Optional[Path] = None) -> list[dict]:
- """Read experiments/tshark_eval_agentic/eval/tasks_l0_vocab.jsonl into a list of dicts.
-
- Caches per-path so repeated probes don't re-read the file. Returns [] if
- the file is missing (probe is then a no-op).
- """
- p = path if path is not None else _DEFAULT_L0_PATH
- p = Path(p)
- if p in _L0_CACHE:
- return _L0_CACHE[p]
- if not p.exists():
- print0(f"[dapt-probe] L0 task file not found at {p}; probe will be a no-op.")
- _L0_CACHE[p] = []
- return _L0_CACHE[p]
- tasks: list[dict] = []
- with p.open("r", encoding="utf-8") as f:
- for line_no, line in enumerate(f, 1):
- line = line.strip()
- if not line:
- continue
- try:
- obj = json.loads(line)
- except json.JSONDecodeError as e:
- print0(f"[dapt-probe] L0 task file line {line_no} not JSON: {e}; skipping.")
- continue
- if "prompt" in obj and "expected" in obj:
- tasks.append(obj)
- print0(f"[dapt-probe] loaded {len(tasks)} L0 tasks from {p}")
- _L0_CACHE[p] = tasks
- return tasks
-
-
-L0_FEW_SHOT = (
- "Q: What is the tshark display-filter field name for: Source IP address?\nA: ip.src\n"
- "Q: What is the tshark display-filter field name for: TCP destination port?\nA: tcp.dstport\n"
- "Q: What is the tshark display-filter field name for: HTTP request method?\nA: http.request.method\n"
-)
-
-
-def _first_line(text: str) -> str:
- """Return the first non-empty line of `text`, stripped."""
- for line in text.strip().splitlines():
- s = line.strip()
- if s:
- return s
- return ""
-
-
-def l0_match(expected: str, generation: str) -> bool:
- """Lenient match: `expected` (case-insensitive) appears in the first
- generated line. Mirrors `experiments/tshark_eval_agentic/eval/grader.exact_str_lower`
- behaviour for the vocab tasks (which always ask for a single token like
- `tcp.flags.syn`).
- """
- return expected.lower() in _first_line(generation).lower()
-
-
-def run_l0_probe(
- model_for_gen,
- tokenizer,
- Engine,
- num_questions: int,
- max_tokens: int,
- disable_fp8_ctx: Optional[Callable[[object], ContextManager]] = None,
- seed: int = 42,
- debug: bool = False,
- task_path: Optional[Path] = None,
-) -> Optional[float]:
- """Greedy-decode N L0 vocab questions and return string-match accuracy.
-
- Returns None if no L0 tasks file is available, else accuracy in [0,1].
-
- `Engine` is passed in (rather than imported) so this function stays
- importable without booting nanoai.engine — the test path can pass a
- stub Engine.
- `disable_fp8_ctx` is the context manager factory (usually `disable_fp8`
- from base_train.py); pass None to disable the swap.
- """
- tasks = load_l0_tasks(task_path)
- if not tasks:
- return None
- rng = random.Random(seed)
- sample = rng.sample(tasks, min(num_questions, len(tasks)))
- print0(
- f"[dapt-probe] running L0 probe: {len(sample)} questions, "
- f"max_tokens={max_tokens}, seed={seed}"
- )
-
- engine = Engine(model_for_gen, tokenizer)
- correct = 0
- t0 = time.time()
- bos = "<|bos|>"
- ctx_factory = disable_fp8_ctx if disable_fp8_ctx is not None else (lambda _m: nullcontext())
- for i, t in enumerate(sample):
- prompt_text = L0_FEW_SHOT + f"Q: {t['prompt']}\nA:"
- ids = tokenizer(prompt_text, prepend=bos)
- with ctx_factory(model_for_gen):
- out, _ = engine.generate_batch(
- ids, num_samples=1, max_tokens=max_tokens, temperature=0,
- )
- gen_ids = out[0][len(ids):] if len(out[0]) > len(ids) else out[0]
- gen_text = tokenizer.decode(gen_ids)
- hit = l0_match(t["expected"], gen_text)
- if hit:
- correct += 1
- if debug:
- print0(
- f"[dapt-probe][L0 {i+1:02d}/{len(sample)}] "
- f"expected={t['expected']!r} gen_first_line={_first_line(gen_text)!r} "
- f"{'✓' if hit else '✗'}"
- )
- dt = time.time() - t0
- acc = correct / len(sample)
- print0(
- f"[dapt-probe] L0 probe done: {correct}/{len(sample)} = {acc:.3f} "
- f"({dt:.1f}s, {dt/len(sample):.2f}s/q)"
- )
- return acc
diff --git a/vendor/nanoai/agentic/data/__init__.py b/vendor/nanoai/agentic/data/__init__.py
deleted file mode 100644
index 6e0f0468c6354ab85fe1830a854b9e75cc1623da..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Conversation/preference datasets used by agentic SFT and DPO."""
diff --git a/vendor/nanoai/agentic/data/codedebug.py b/vendor/nanoai/agentic/data/codedebug.py
deleted file mode 100644
index 286b38731606e149a92be1737229bf970fbea018..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/codedebug.py
+++ /dev/null
@@ -1,359 +0,0 @@
-"""Multi-turn execute-debug datasets for v5 SFT.
-
-Goal: give the model real `tool_call → observation → tool_call → ...` loops
-beyond tulu's Edit-only tool-call shape (the v4 Edit-bias root cause).
-
-Two sources, both folded into the agentic chat format used by
-`render_agentic_conversation`'s list-content branch:
-
-- `xingyaoww/code-act` (codeact split): assistant emits `...`,
- user emits `Observation: ...`. ~7k rows.
-- `m-a-p/Code-Feedback` (filtered to python with exec output): assistant emits
- ` ```python ... ``` `, user emits `Execution result: ...`. ~10k after filter.
-
-Output row shape (per `__getitem__`):
- {"messages": [
- {"role": "system", "content": SYSTEM_PROMPT},
- {"role": "user", "content": "..."},
- {"role": "assistant", "content": [{"type":"text"|"python"|"python_output", "text": ...}, ...]},
- ...
- ]}
-
-The list-content branch wraps `python` segments in `<|python_start|>/<|python_end|>`
-(supervised) and `python_output` in `<|output_start|>/<|output_end|>` (unsupervised).
-"""
-
-from __future__ import annotations
-
-import os
-import random
-import re
-from typing import Any, Iterable
-
-from datasets import load_dataset
-
-from nanoai.agentic.data.common import SYSTEM_PROMPT
-
-
-# ---------------------------------------------------------------------------
-# Parsers
-# ---------------------------------------------------------------------------
-
-# Code-act assistant action: ... . Solutions stay as text.
-_EXECUTE_RE = re.compile(r"(.*?)", re.DOTALL)
-
-# Code-Feedback assistant action: ```python ... ``` (case-insensitive, leading ws ok).
-# We deliberately match only python (Ruby/JS/etc. blocks stay as text — caller
-# filters those rows out earlier so they don't end up here).
-_PYBLOCK_RE = re.compile(r"```\s*[Pp]ython\s*\n(.*?)```", re.DOTALL)
-
-# Observation prefixes folded into the preceding assistant turn as python_output.
-_OBS_PREFIXES = (
- "observation:",
- "execution result:",
- "execution output:",
- "stdout:",
-)
-
-# Skip Code-Feedback rows whose "execution" is actually a Jupyter cell error trace —
-# those teach the model how-to-format-a-traceback, not how-to-debug.
-_JUPYTER_ERROR_MARKERS = ("Cell In[", " list[dict[str, str]]:
- """Split an assistant message into ordered text/python parts.
-
- The regex matches code-fenced or tag-fenced spans; everything else stays text.
- Returns [] if the input is empty after stripping.
- """
- if not text or not text.strip():
- return []
- parts: list[dict[str, str]] = []
- pos = 0
- for m in code_re.finditer(text):
- head = text[pos : m.start()]
- if head.strip():
- parts.append({"type": "text", "text": head.strip()})
- body = m.group(1).strip("\n").rstrip()
- if body:
- parts.append({"type": "python", "text": body})
- pos = m.end()
- tail = text[pos:]
- if tail.strip():
- parts.append({"type": "text", "text": tail.strip()})
- return parts
-
-
-def _extract_observation(content: str) -> str | None:
- """If content is a pure observation, return the body sans prefix; else None."""
- if not content:
- return None
- stripped = content.lstrip()
- lower = stripped.lower()
- for prefix in _OBS_PREFIXES:
- if lower.startswith(prefix):
- return stripped[len(prefix):].lstrip("\n").rstrip() or " "
- return None
-
-
-def _fold_into_assistant(messages: list[dict[str, Any]], code_re: re.Pattern[str]) -> list[dict[str, Any]]:
- """Walk raw user/assistant alternation; return the list-content cleaned form.
-
- Rules:
- * Drop system messages — caller injects SYSTEM_PROMPT.
- * Assistant with code blocks → list-content [{text|python} parts].
- * Assistant without code → plain str content.
- * User that is purely an observation following an assistant with code →
- appended as a {type: python_output} part to the preceding assistant.
- * Anything else stays user/assistant.
- * Drop trailing user-only turns (no completion to supervise).
- * Drop rows whose first non-system role isn't user.
- """
- cleaned: list[dict[str, Any]] = []
-
- for msg in messages:
- role = msg.get("role")
- content = msg.get("content", "") or ""
- if role == "system":
- continue
-
- if role == "assistant":
- parts = _split_assistant_text(content, code_re)
- if not parts:
- continue
- if any(p["type"] == "python" for p in parts):
- cleaned.append({"role": "assistant", "content": parts})
- else:
- # Plain prose answer — keep as a single string for the simpler
- # render path. Re-join the text parts.
- cleaned.append({"role": "assistant", "content": " ".join(p["text"] for p in parts)})
- continue
-
- if role == "user":
- obs = _extract_observation(content)
- if (
- obs is not None
- and cleaned
- and cleaned[-1]["role"] == "assistant"
- and isinstance(cleaned[-1]["content"], list)
- and any(p["type"] == "python" for p in cleaned[-1]["content"])
- ):
- cleaned[-1]["content"].append({"type": "python_output", "text": obs})
- continue
- cleaned.append({"role": "user", "content": content.strip()})
- continue
-
- # Unknown role (tool_result, function, etc.) — skip silently.
-
- # Drop leading non-user turns and trailing non-assistant turns.
- while cleaned and cleaned[0]["role"] != "user":
- cleaned.pop(0)
- while cleaned and cleaned[-1]["role"] != "assistant":
- cleaned.pop()
-
- # Folding observations leaves consecutive assistants (the original
- # post-observation assistant continuation). Same can happen with consecutive
- # users that don't satisfy the observation-fold criterion. Merge them.
- return _merge_consecutive_same_role(cleaned)
-
-
-def _merge_consecutive_same_role(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
- """Merge runs of same-role messages into one. Strings → list-content if needed."""
- out: list[dict[str, Any]] = []
- for m in messages:
- if out and out[-1]["role"] == m["role"]:
- prev = out[-1]
- if m["role"] == "user":
- prev_text = prev["content"] if isinstance(prev["content"], str) else " ".join(
- p["text"] for p in prev["content"]
- )
- cur_text = m["content"] if isinstance(m["content"], str) else " ".join(
- p["text"] for p in m["content"]
- )
- prev["content"] = (prev_text + "\n\n" + cur_text).strip()
- continue
- # assistant: prefer list-content if either side already is.
- prev_content = prev["content"]
- cur_content = m["content"]
- if isinstance(prev_content, str) and isinstance(cur_content, str):
- prev["content"] = (prev_content + " " + cur_content).strip()
- continue
- if isinstance(prev_content, str):
- prev_content = [{"type": "text", "text": prev_content}] if prev_content.strip() else []
- if isinstance(cur_content, str):
- cur_content = [{"type": "text", "text": cur_content}] if cur_content.strip() else []
- prev["content"] = list(prev_content) + list(cur_content)
- continue
- out.append(m)
- return out
-
-
-def parse_codeact_row(row: dict[str, Any], max_chars: int = 16384) -> dict[str, Any] | None:
- """Convert a single code-act row to {messages: [...]} or None to skip."""
- raw = row.get("conversations") or []
- if not raw:
- return None
- messages = _fold_into_assistant(raw, _EXECUTE_RE)
- if not _is_well_formed(messages, max_chars):
- return None
- return {"messages": messages}
-
-
-def parse_codefeedback_row(row: dict[str, Any], max_chars: int = 16384) -> dict[str, Any] | None:
- """Convert a single Code-Feedback row to {messages: [...]} or None to skip."""
- raw = row.get("messages") or []
- if not raw:
- return None
- full_text = "\n".join((m.get("content") or "") for m in raw)
- if len(full_text) > max_chars:
- return None
- if any(marker in full_text for marker in _JUPYTER_ERROR_MARKERS):
- return None
- if "execution result" not in full_text.lower() and "execution output" not in full_text.lower():
- return None
- if _NON_PYTHON_BLOCK_RE.search(full_text) and "```python" not in full_text.lower():
- return None
- messages = _fold_into_assistant(raw, _PYBLOCK_RE)
- if not _is_well_formed(messages, max_chars):
- return None
- # Require at least one observation actually got folded in — otherwise this
- # is just chat about code, not execute-debug.
- if not _has_python_output(messages):
- return None
- return {"messages": messages}
-
-
-def _is_well_formed(messages: list[dict[str, Any]], max_chars: int) -> bool:
- if len(messages) < 2:
- return False
- # Strict alternation user/assistant.
- for i, m in enumerate(messages):
- expected = "user" if i % 2 == 0 else "assistant"
- if m["role"] != expected:
- return False
- if messages[-1]["role"] != "assistant":
- return False
- total = 0
- for m in messages:
- c = m["content"]
- if isinstance(c, str):
- total += len(c)
- else:
- total += sum(len(p["text"]) for p in c)
- return total <= max_chars
-
-
-def _has_python_output(messages: list[dict[str, Any]]) -> bool:
- for m in messages:
- if m["role"] == "assistant" and isinstance(m["content"], list):
- if any(p["type"] == "python_output" for p in m["content"]):
- return True
- return False
-
-
-# ---------------------------------------------------------------------------
-# Dataset wrappers (HF-backed)
-# ---------------------------------------------------------------------------
-
-
-def _set_hf_endpoint() -> None:
- """If HF_ENDPOINT isn't set and the public hub is unreachable, fall back to
- hf-mirror.com. We don't probe network here — caller is expected to set
- HF_ENDPOINT in env when behind a firewall."""
- # Intentionally a no-op; documented for callers.
- return
-
-
-class CodeActDataset:
- """xingyaoww/code-act, codeact split — multi-turn data analysis with sqlite."""
-
- def __init__(self, split: str = "codeact", seed: int = 42, max_chars: int = 16384) -> None:
- _set_hf_endpoint()
- ds = load_dataset("xingyaoww/code-act", split=split)
- self.parsed: list[dict[str, Any]] = []
- n_skipped = 0
- for row in ds:
- parsed = parse_codeact_row(row, max_chars=max_chars)
- if parsed is None:
- n_skipped += 1
- continue
- self.parsed.append(parsed)
- random.Random(seed).shuffle(self.parsed)
- self._n_skipped = n_skipped
-
- def __len__(self) -> int:
- return len(self.parsed)
-
- def __getitem__(self, idx: int) -> dict[str, Any]:
- row = self.parsed[idx]
- return {"messages": [{"role": "system", "content": SYSTEM_PROMPT}] + row["messages"]}
-
-
-class CodeFeedbackDataset:
- """m-a-p/Code-Feedback filtered to python rows with execution feedback."""
-
- def __init__(
- self,
- split: str = "train",
- seed: int = 42,
- max_chars: int = 16384,
- max_rows: int | None = 10_000,
- ) -> None:
- _set_hf_endpoint()
- ds = load_dataset("m-a-p/Code-Feedback", split=split)
- self.parsed: list[dict[str, Any]] = []
- n_skipped = 0
- for row in ds:
- parsed = parse_codefeedback_row(row, max_chars=max_chars)
- if parsed is None:
- n_skipped += 1
- continue
- self.parsed.append(parsed)
- if max_rows is not None and len(self.parsed) >= max_rows:
- break
- random.Random(seed).shuffle(self.parsed)
- self._n_skipped = n_skipped
-
- def __len__(self) -> int:
- return len(self.parsed)
-
- def __getitem__(self, idx: int) -> dict[str, Any]:
- row = self.parsed[idx]
- return {"messages": [{"role": "system", "content": SYSTEM_PROMPT}] + row["messages"]}
-
-
-# ---------------------------------------------------------------------------
-# In-memory variant for unit tests (no HF dep)
-# ---------------------------------------------------------------------------
-
-
-class _InMemoryCodeDataset:
- """Wraps a pre-built list of {messages: [...]} rows; used by tests."""
-
- def __init__(self, rows: Iterable[dict[str, Any]], seed: int = 42) -> None:
- self.parsed = list(rows)
- random.Random(seed).shuffle(self.parsed)
-
- def __len__(self) -> int:
- return len(self.parsed)
-
- def __getitem__(self, idx: int) -> dict[str, Any]:
- row = self.parsed[idx]
- return {"messages": [{"role": "system", "content": SYSTEM_PROMPT}] + row["messages"]}
-
-
-__all__ = [
- "CodeActDataset",
- "CodeFeedbackDataset",
- "parse_codeact_row",
- "parse_codefeedback_row",
- "_InMemoryCodeDataset",
-]
diff --git a/vendor/nanoai/agentic/data/common.py b/vendor/nanoai/agentic/data/common.py
deleted file mode 100644
index abfe9e1f4c5b909bbc4249d92dec5d0f5f2f5481..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/common.py
+++ /dev/null
@@ -1,17 +0,0 @@
-SYSTEM_PROMPT = """you are nanocode, a coding agent.
-you communicate in lowercase and think out loud before acting.
-
-your tools enable you to interact with the user's UNIX system and modify and create new files:
- Read({"file_path", "offset"?, "limit"?}) - read a file
- Edit({"file_path", "old_string"?, "new_string"}) - edit or create a file
- Grep({"pattern", "path"?}) - search file contents
- Bash({"command"}) - execute a shell command
-
-emit a tool call by producing exactly one tool-call token block per assistant turn. you've seen the structure during training; reproduce it precisely."""
-
-
-def render_mc(question, letters, choices):
- query = f"Multiple Choice question: {question}\n"
- query += "".join([f"- {choice}={letter}\n" for letter, choice in zip(letters, choices)])
- query += "\nRespond only with the letter of the correct answer."
- return query
diff --git a/vendor/nanoai/agentic/data/ebook_data.py b/vendor/nanoai/agentic/data/ebook_data.py
deleted file mode 100644
index 868e44f5af26aec0f59cbf3fec9419c6f4fc45a3..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/ebook_data.py
+++ /dev/null
@@ -1,201 +0,0 @@
-"""Ebook pretrain stream (PDF→markdown textbook corpus).
-
-Materializes the converted-ebook markdown under ``~/ebook/markdown/**/*.md``
-(produced by ~/ebook/convert_ebooks.py, pdftext backend) into parquet shards
-so the existing parquet-iter API drives it identically to web / code / pcap.
-
-Each markdown book is **chunked** into ~``chunk_chars``-sized documents at
-paragraph boundaries, so that the code-mixer's per-document categorical draw
-approximates a per-token mixture: climbmix web docs average ~2875 chars
-(~720 tok), so the default 3000-char chunk keeps ebook docs the same scale and
-``--ebook-ratio 0.5`` ≈ 50% of *tokens* (not just 50% of documents).
-
-Image references (````) are stripped — they are
-pure-noise placeholders for figures pdftext could not OCR.
-
-Stored under ``$NANOAI_BASE_DIR/code_pretrain_data/ebook-md/``. Used via
-`agentic_pretrain.py --ebook-ratio`. Mirrors `vrcpt_data` (atomic .tmp→rename,
-restart-safe validate/quarantine, last shard = val split).
-
-Run once before pretraining:
-
- python -m nanoai.agentic.data.ebook_data prepare \\
- --md-root ~/ebook/markdown --chunk-chars 3000 --rows-per-shard 50000
-"""
-
-from __future__ import annotations
-
-import argparse
-import os
-import random
-import re
-import time
-from pathlib import Path
-
-import pyarrow as pa
-import pyarrow.parquet as pq
-
-from nanoai.common import print0
-from nanoai.agentic.data.pretrain_data import DATA_DIR
-
-DATASET_NAME = "ebook-md"
-_TMP_SUFFIX = ".parquet.tmp"
-
-# image-reference lines like  or 
-_IMG_RE = re.compile(r"!\[[^\]]*\]\([^)]*\)")
-# 3+ consecutive blank lines → 2
-_BLANKS_RE = re.compile(r"\n{3,}")
-
-
-def ebook_corpus_dir() -> Path:
- return DATA_DIR / DATASET_NAME
-
-
-def list_ebook_parquet_files() -> list[Path]:
- d = ebook_corpus_dir()
- if not d.exists():
- return []
- return sorted([d / f.name for f in d.iterdir() if f.suffix == ".parquet"])
-
-
-def _clean(text: str) -> str:
- text = _IMG_RE.sub("", text)
- text = _BLANKS_RE.sub("\n\n", text)
- return text.strip()
-
-
-def _chunk_markdown(text: str, chunk_chars: int):
- """Split cleaned markdown into ~chunk_chars docs at paragraph boundaries."""
- paras = text.split("\n\n")
- buf, size = [], 0
- for p in paras:
- p = p.strip()
- if not p:
- continue
- buf.append(p)
- size += len(p) + 2
- if size >= chunk_chars:
- yield "\n\n".join(buf)
- buf, size = [], 0
- if buf:
- joined = "\n\n".join(buf)
- if len(joined) >= 200: # drop tiny trailing fragments
- yield joined
-
-
-def _stream_chunks(md_root: Path, chunk_chars: int, seed: int):
- files = sorted(md_root.rglob("*.md"))
- rng = random.Random(seed)
- rng.shuffle(files) # book-level shuffle so shards aren't single-book
- for fp in files:
- try:
- raw = fp.read_text(encoding="utf-8", errors="ignore")
- except Exception as e:
- print0(f"[ebook_data] skip {fp.name}: {e}")
- continue
- cleaned = _clean(raw)
- if not cleaned:
- continue
- for chunk in _chunk_markdown(cleaned, chunk_chars):
- yield chunk
-
-
-def _shuffled(stream, buffer_size: int, seed: int):
- """Reservoir-style shuffle buffer to decorrelate adjacent same-book chunks."""
- rng = random.Random(seed + 1)
- buf: list[str] = []
- for item in stream:
- buf.append(item)
- if len(buf) >= buffer_size:
- j = rng.randrange(len(buf))
- buf[j], buf[-1] = buf[-1], buf[j]
- yield buf.pop()
- rng.shuffle(buf)
- yield from buf
-
-
-def _try_open_shard(path: Path) -> bool:
- try:
- return pq.ParquetFile(path).num_row_groups > 0
- except Exception:
- return False
-
-
-def prepare(md_root: Path, rows_per_shard: int = 50000, chunk_chars: int = 3000,
- shuffle_buffer: int = 50000, seed: int = 1234):
- """One-shot markdown → parquet shards under ebook_corpus_dir().
-
- Atomic write (.tmp → rename), restart-safe. Last shard = val split.
- """
- out_dir = ebook_corpus_dir()
- out_dir.mkdir(parents=True, exist_ok=True)
-
- existing = sorted(p for p in out_dir.iterdir() if p.suffix == ".parquet")
- tmps = [p for p in out_dir.iterdir() if p.name.endswith(_TMP_SUFFIX)]
- if existing and not tmps and all(_try_open_shard(p) for p in existing) and len(existing) >= 2:
- print0(f"[ebook_data] {len(existing)} valid shards already in {out_dir}; skipping.")
- return existing
- if existing or tmps:
- attic = out_dir / f".attic_{int(time.time())}"
- attic.mkdir()
- for p in existing + tmps:
- p.rename(attic / p.name)
- print0(f"[ebook_data] quarantined {len(existing)+len(tmps)} stale file(s) to {attic.name}/")
-
- assert md_root.exists(), f"md-root not found: {md_root}"
- print0(f"[ebook_data] reading {md_root} → {out_dir} "
- f"(chunk={chunk_chars} chars, {rows_per_shard} rows/shard)")
-
- shard_idx = 0
- buf: list[str] = []
- written: list[Path] = []
- total_docs = 0
- total_chars = 0
-
- def flush():
- nonlocal shard_idx, buf
- if not buf:
- return
- shard_path = out_dir / f"shard_{shard_idx:05d}.parquet"
- tmp_path = shard_path.with_name(shard_path.name + ".tmp")
- shard_chars = sum(len(t) for t in buf)
- pq.write_table(pa.table({"text": buf}), tmp_path, compression="snappy")
- tmp_path.rename(shard_path)
- written.append(shard_path)
- print0(f"[ebook_data] wrote {shard_path.name} "
- f"({len(buf)} docs, {shard_chars/1e6:.2f}M chars, "
- f"{shard_path.stat().st_size/1e6:.2f} MB)")
- shard_idx += 1
- buf = []
-
- stream = _stream_chunks(md_root, chunk_chars, seed)
- for text in _shuffled(stream, shuffle_buffer, seed):
- buf.append(text)
- total_docs += 1
- total_chars += len(text)
- if len(buf) >= rows_per_shard:
- flush()
- flush()
-
- print0(f"[ebook_data] DONE: {total_docs:,} docs, {total_chars/1e6:.1f}M chars "
- f"(~{total_chars/4/1e6:.0f}M est tokens) across {len(written)} shards")
- return written
-
-
-def main():
- ap = argparse.ArgumentParser(description="Prepare ebook markdown pretrain shards")
- sub = ap.add_subparsers(dest="cmd", required=True)
- p = sub.add_parser("prepare")
- p.add_argument("--md-root", type=str, default=str(Path.home() / "ebook" / "markdown"))
- p.add_argument("--rows-per-shard", type=int, default=50000)
- p.add_argument("--chunk-chars", type=int, default=3000)
- p.add_argument("--shuffle-buffer", type=int, default=50000)
- p.add_argument("--seed", type=int, default=1234)
- args = ap.parse_args()
- if args.cmd == "prepare":
- prepare(Path(os.path.expanduser(args.md_root)), args.rows_per_shard,
- args.chunk_chars, args.shuffle_buffer, args.seed)
-
-
-if __name__ == "__main__":
- main()
diff --git a/vendor/nanoai/agentic/data/hf_dataset.py b/vendor/nanoai/agentic/data/hf_dataset.py
deleted file mode 100644
index da045697b91ab9d265cf3a08a0d8b1f2e0d44287..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/hf_dataset.py
+++ /dev/null
@@ -1,15 +0,0 @@
-"""Generic Hugging Face conversation dataset loader."""
-
-from datasets import load_dataset
-
-
-class HuggingFaceDataset:
- def __init__(self, dataset: str, messages_key: str, split: str, seed: int, **kwargs):
- self.messages_key = messages_key
- self.ds = load_dataset(dataset, split=split, **kwargs).shuffle(seed=seed)
-
- def __len__(self):
- return len(self.ds)
-
- def __getitem__(self, idx: int):
- return {"messages": self.ds[idx][self.messages_key]}
diff --git a/vendor/nanoai/agentic/data/json_dataset.py b/vendor/nanoai/agentic/data/json_dataset.py
deleted file mode 100644
index 77dcda5398458934752d28654ef1c9dff27e929c..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/json_dataset.py
+++ /dev/null
@@ -1,73 +0,0 @@
-"""JSONL conversation / preference datasets.
-
-Conversations are tagged with the agentic SYSTEM_PROMPT so the rendered tool-call
-schema is always in context. Format on disk matches nanocode rollouts:
- {"messages": [{"role": "user", ...}, {"role": "assistant", ...}, ...]}
- {"chosen": {"messages": [...]}, "rejected": {"messages": [...]}} # for prefs
-"""
-
-from __future__ import annotations
-
-import json
-import random
-
-from nanoai.agentic.data.common import SYSTEM_PROMPT
-
-
-class JSONDataset:
- def __init__(self, filepath, seed, **kwargs):
- super().__init__(**kwargs)
- self.filepath = filepath
- self.ds = []
- with open(filepath, "r", encoding="utf-8") as f:
- for line in f:
- line = line.strip()
- if not line:
- continue
- self.ds.append(json.loads(line))
- random.Random(seed).shuffle(self.ds)
-
- def __len__(self):
- return len(self.ds)
-
- def __getitem__(self, idx: int):
- messages = self.ds[idx]["messages"]
- messages = [{"role": "system", "content": SYSTEM_PROMPT}] + messages
- return {"messages": messages}
-
-
-class PairedJSONPreferenceDataset:
- """Preference pairs filtered to require non-empty assistant content on both sides."""
-
- def __init__(self, filepath, seed, frac=1.0, **kwargs):
- super().__init__(**kwargs)
- self.ds = []
- with open(filepath, "r", encoding="utf-8") as f:
- for line in f:
- line = line.strip()
- if not line:
- continue
- row = json.loads(line)
- has_chosen = any(
- m.get("content", "")
- for m in row["chosen"]["messages"]
- if m["role"] == "assistant"
- )
- has_rejected = any(
- m.get("content", "")
- for m in row["rejected"]["messages"]
- if m["role"] == "assistant"
- )
- if has_chosen and has_rejected:
- self.ds.append(row)
- random.Random(seed).shuffle(self.ds)
- self.ds = self.ds[: int(len(self.ds) * frac)]
-
- def __len__(self):
- return len(self.ds)
-
- def __getitem__(self, idx: int):
- row = self.ds[idx]
- chosen = [{"role": "system", "content": SYSTEM_PROMPT}] + row["chosen"]["messages"]
- rejected = [{"role": "system", "content": SYSTEM_PROMPT}] + row["rejected"]["messages"]
- return {"messages": chosen}, {"messages": rejected}
diff --git a/vendor/nanoai/agentic/data/mixture.py b/vendor/nanoai/agentic/data/mixture.py
deleted file mode 100644
index 6cc618eae4415e98c74043cf0798e5b0730879cd..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/mixture.py
+++ /dev/null
@@ -1,42 +0,0 @@
-"""Deterministic interleaving of multiple sub-datasets."""
-
-import random
-
-
-class TaskMixture:
- def __init__(self, tasks, seed: int, **kwargs):
- super().__init__(**kwargs)
- self.tasks = tasks
- self.lengths = [len(task) for task in self.tasks]
- self.num_conversations = sum(self.lengths)
- self.index_map = []
- for task_idx, task_length in enumerate(self.lengths):
- for local_idx in range(task_length):
- self.index_map.append((task_idx, local_idx))
- random.Random(seed).shuffle(self.index_map)
-
- def __len__(self):
- return self.num_conversations
-
- def __getitem__(self, index):
- assert 0 <= index < self.num_conversations
- task_idx, local_idx = self.index_map[index]
- return self.tasks[task_idx][local_idx]
-
-
-class TaskSequence:
- def __init__(self, tasks, **kwargs):
- super().__init__(**kwargs)
- self.tasks = tasks
- self.lengths = [len(task) for task in self.tasks]
- self.num_conversations = sum(self.lengths)
-
- def __len__(self):
- return self.num_conversations
-
- def __getitem__(self, index):
- assert 0 <= index < self.num_conversations
- for task_idx, task_length in enumerate(self.lengths):
- if index < task_length:
- return self.tasks[task_idx][index]
- index -= task_length
diff --git a/vendor/nanoai/agentic/data/pcap_data.py b/vendor/nanoai/agentic/data/pcap_data.py
deleted file mode 100644
index 701802309698f52cf36912816f9df5c2e6b0bdec..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/pcap_data.py
+++ /dev/null
@@ -1,224 +0,0 @@
-"""Pcap-corpus pretrain stream.
-
-Materializes a domain-aligned third pretrain stream alongside web (climbmix /
-fineweb-edu) and code (the-stack-v2-dedup). The raw input is a JSONL file
-produced upstream by `capbench_agent/corpus/scripts/prepare_cpt_data.py`
-(57k+ docs across 11 knowledge layers: protocols, philosophy, mathematics,
-os_systems, forensics, network_traffic, reasoning, ndes_native, ...).
-
-This module converts the JSONL into parquet shards stored under
-`$NANOAI_BASE_DIR/code_pretrain_data/pcap-corpus/` so the existing
-`parquets_iter_batched` API (which expects parquet w/ a `text` column) can
-drive it identically to fineweb-edu / the-stack-v2-dedup.
-
-Run once before pretraining:
-
- python -m nanoai.agentic.data.pcap_data prepare \\
- --src /home/vader/capbench_agent/corpus/output/pretrain.jsonl \\
- --rows-per-shard 5000
-"""
-
-from __future__ import annotations
-
-import argparse
-import json
-import time
-from pathlib import Path
-
-import pyarrow as pa
-import pyarrow.parquet as pq
-
-from nanoai.common import print0
-from nanoai.agentic.data.pretrain_data import DATA_DIR
-
-DATASET_NAME = "pcap-corpus"
-_TMP_SUFFIX = ".parquet.tmp"
-
-
-def pcap_corpus_dir() -> Path:
- return DATA_DIR / DATASET_NAME
-
-
-def _stream_jsonl_text(src: Path):
- with src.open("r", encoding="utf-8") as f:
- for line in f:
- line = line.strip()
- if not line:
- continue
- obj = json.loads(line)
- text = obj.get("text")
- if isinstance(text, str) and text:
- yield text
-
-
-def _try_open_shard(path: Path) -> bool:
- """Return True iff `path` is a readable parquet with ≥1 row group."""
- try:
- pf = pq.ParquetFile(path)
- return pf.num_row_groups > 0
- except Exception:
- return False
-
-
-def _validate_existing_corpus(out_dir: Path) -> tuple[bool, list[Path], str]:
- """Inspect out_dir to decide whether it holds a complete prepared corpus.
-
- Returns (is_valid, parquet_list, reason). A corpus is valid iff:
- - ≥2 parquet shards (so train/val split is non-empty), AND
- - every parquet is openable and has ≥1 row_group, AND
- - no leftover `*.parquet.tmp` files (signal of prior interrupted write).
- """
- parquets = sorted(p for p in out_dir.iterdir() if p.suffix == ".parquet")
- tmps = [p for p in out_dir.iterdir() if p.name.endswith(_TMP_SUFFIX)]
- if tmps:
- return False, parquets, f"{len(tmps)} stale {_TMP_SUFFIX} file(s) (prior crash?)"
- if len(parquets) < 2:
- return False, parquets, f"only {len(parquets)} shard(s); need ≥2"
- for p in parquets:
- if not _try_open_shard(p):
- return False, parquets, f"shard {p.name} unreadable or has 0 row_groups"
- return True, parquets, "ok"
-
-
-def _quarantine_stale(out_dir: Path, reason_slug: str) -> Path | None:
- """Move any *.parquet / *.parquet.tmp files aside before rebuilding.
-
- Uses an `.attic__/` sub-dir so the operator can post-mortem
- what was wrong without us silently `rm -rf`-ing it.
- """
- movable = [
- p for p in out_dir.iterdir()
- if p.suffix == ".parquet" or p.name.endswith(_TMP_SUFFIX)
- ]
- if not movable:
- return None
- attic = out_dir / f".attic_{int(time.time())}_{reason_slug}"
- attic.mkdir()
- for entry in movable:
- entry.rename(attic / entry.name)
- print0(
- f"[pcap_data] quarantined {len(movable)} stale file(s) to {attic.name}/"
- f" (post-mortem; safe to delete once unblocked)"
- )
- return attic
-
-
-def _reason_slug(reason: str) -> str:
- return reason.split(";")[0].strip().replace(" ", "_")[:32] or "invalid"
-
-
-def prepare_from_jsonl(src: Path, rows_per_shard: int = 5000):
- """One-shot: JSONL -> parquet shards under pcap_corpus_dir().
-
- Last shard becomes the val split (mirrors fineweb-edu convention).
-
- Restart-safe:
- - if the output dir already holds a *valid* corpus (≥2 readable shards,
- no leftover .tmp), this is a no-op.
- - if the corpus is incomplete (single shard / unreadable parquet /
- stale .tmp), the stale files are quarantined into `.attic__*/`
- and a full rebuild runs.
- - each new shard is written as `.parquet.tmp` then renamed to
- `.parquet` (atomic on POSIX), so a mid-shard crash leaves a
- clearly-named .tmp that this function will sweep on the next call.
- """
- out_dir = pcap_corpus_dir()
- out_dir.mkdir(parents=True, exist_ok=True)
-
- is_valid, existing, reason = _validate_existing_corpus(out_dir)
- if is_valid:
- print0(
- f"[pcap_data] {len(existing)} valid shards already in {out_dir}; "
- f"skipping prepare."
- )
- return existing
- # Either nothing is there, or what's there is incomplete/corrupt.
- if existing or any(p.name.endswith(_TMP_SUFFIX) for p in out_dir.iterdir()):
- print0(
- f"[pcap_data] existing corpus at {out_dir} is INVALID ({reason}); "
- f"quarantining and rebuilding from scratch."
- )
- _quarantine_stale(out_dir, _reason_slug(reason))
-
- assert src.exists(), f"Source JSONL not found: {src}"
- print0(f"[pcap_data] reading {src} → {out_dir} ({rows_per_shard} rows/shard)")
-
- shard_idx = 0
- buf: list[str] = []
- written: list[Path] = []
- total_docs = 0
- total_chars = 0
-
- def flush():
- nonlocal shard_idx, buf
- if not buf:
- return
- shard_path = out_dir / f"shard_{shard_idx:05d}.parquet"
- tmp_path = shard_path.with_name(shard_path.name + ".tmp")
- shard_chars = sum(len(t) for t in buf)
- table = pa.table({"text": buf})
- pq.write_table(table, tmp_path, compression="snappy")
- tmp_path.rename(shard_path) # atomic on POSIX
- written.append(shard_path)
- size_mb = shard_path.stat().st_size / 1e6
- print0(
- f"[pcap_data] wrote {shard_path.name} "
- f"({len(buf)} docs, {shard_chars/1e6:.2f}M chars, {size_mb:.2f} MB on disk)"
- )
- shard_idx += 1
- buf = []
-
- for text in _stream_jsonl_text(src):
- buf.append(text)
- total_docs += 1
- total_chars += len(text)
- if len(buf) >= rows_per_shard:
- flush()
- flush()
-
- assert len(written) >= 2, (
- f"Only {len(written)} shard(s) written; need ≥2 so train/val split is non-empty. "
- f"Lower --rows-per-shard or supply a larger JSONL."
- )
- # Rough token-budget estimate using 4 chars/token (English) — matches the
- # ballpark we see in nanoai's tokenizer for educational + code mix.
- est_tokens = total_chars / 4
- print0(
- f"[pcap_data] done: {len(written)} shards (last is val), "
- f"{total_docs:,} docs, {total_chars/1e6:.1f}M chars, "
- f"≈{est_tokens/1e6:.1f}M tokens (4 chars/tok estimate)"
- )
- return written
-
-
-def list_pcap_parquet_files() -> list[Path]:
- d = pcap_corpus_dir()
- if not d.exists():
- return []
- return sorted([d / f for f in d.iterdir() if f.suffix == ".parquet"])
-
-
-def main():
- parser = argparse.ArgumentParser(description="Prepare pcap-corpus parquet shards.")
- sub = parser.add_subparsers(dest="cmd", required=True)
- prep = sub.add_parser("prepare", help="convert a JSONL to parquet shards")
- prep.add_argument("--src", type=Path, required=True, help="path to pretrain.jsonl")
- prep.add_argument("--rows-per-shard", type=int, default=5000)
- sub.add_parser("list", help="list current pcap-corpus shards")
-
- args = parser.parse_args()
- if args.cmd == "prepare":
- prepare_from_jsonl(args.src, rows_per_shard=args.rows_per_shard)
- elif args.cmd == "list":
- shards = list_pcap_parquet_files()
- if not shards:
- print(f"No shards in {pcap_corpus_dir()}")
- return
- total_bytes = sum(s.stat().st_size for s in shards)
- for s in shards:
- print(f" {s.name} ({s.stat().st_size / 1e6:.2f} MB)")
- print(f"total: {len(shards)} shards, {total_bytes / 1e6:.1f} MB")
-
-
-if __name__ == "__main__":
- main()
diff --git a/vendor/nanoai/agentic/data/pretrain_data.py b/vendor/nanoai/agentic/data/pretrain_data.py
deleted file mode 100644
index 166a6ccb922f3fda788d73d0a1f9455e34fe4aef..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/pretrain_data.py
+++ /dev/null
@@ -1,115 +0,0 @@
-"""Parquet-shard downloader and iterator for the two pretrain datasets:
- - karpathy/fineweb-edu-100b-shuffle (web text)
- - smohammadi/the-stack-v2-python-shuffle (Python code)
-
-Honors HF_ENDPOINT so a mirror (e.g. hf-mirror.com) can be substituted.
-Run as a CLI: `python -m nanoai.agentic.data.pretrain_data -d fineweb-edu -n 4`
-"""
-
-from __future__ import annotations
-
-import argparse
-import os
-import time
-from functools import partial
-from multiprocessing import Pool
-from pathlib import Path
-
-import pyarrow.parquet as pq
-import requests
-
-from nanoai.common import get_base_dir, get_dist_info, print0
-
-DATA_DIR = Path(get_base_dir()) / "code_pretrain_data"
-DATA_DIR.mkdir(parents=True, exist_ok=True)
-
-
-def index_to_filename(index: int) -> str:
- return f"shard_{index:05d}.parquet"
-
-
-def list_parquet_files(data_dir: Path):
- return sorted([data_dir / f for f in data_dir.iterdir() if f.suffix == ".parquet"])
-
-
-def parquets_iter_batched(dataset: str, split: str, start: int = 0, step: int = 1):
- """Yield row-group-sized batches of `text` strings.
-
- `split="val"` returns only the last shard; `split="train"` returns the rest.
- `start`/`step` slice row-groups for DDP rank/world-size sharding.
- """
- assert dataset in ("fineweb-edu", "the-stack-v2-dedup", "network-domain", "pcap-corpus")
- assert split in ("train", "val")
- data_dir = DATA_DIR / dataset
- parquet_paths = list_parquet_files(data_dir)
- if not parquet_paths:
- raise FileNotFoundError(
- f"No parquet shards found in {data_dir}. "
- f"Run: python -m nanoai.agentic.data.pretrain_data -d {dataset} -n N"
- )
- parquet_paths = parquet_paths[:-1] if split == "train" else parquet_paths[-1:]
- for filepath in parquet_paths:
- pf = pq.ParquetFile(filepath)
- for rg_idx in range(start, pf.num_row_groups, step):
- rg = pf.read_row_group(rg_idx)
- yield rg.column("text").to_pylist()
-
-
-def download_single_file(index: int, base_url: str, data_dir: Path) -> bool:
- filename = index_to_filename(index)
- filepath = data_dir / filename
- if filepath.exists():
- return True
-
- url = f"{base_url}/{filename}"
- print0(f"Downloading {filename}...")
-
- for attempt in range(1, 6):
- try:
- response = requests.get(url, stream=True, timeout=30)
- response.raise_for_status()
- temp_path = filepath.with_name(filepath.name + ".tmp")
- with open(temp_path, "wb") as f:
- for chunk in response.iter_content(chunk_size=1024 * 1024):
- if chunk:
- f.write(chunk)
- temp_path.rename(filepath)
- print0(f"Successfully downloaded {filename}")
- return True
- except (requests.RequestException, IOError) as e:
- print0(f"Attempt {attempt}/5 failed for {filename}: {e}")
- filepath.unlink(missing_ok=True)
- filepath.with_name(filepath.name + ".tmp").unlink(missing_ok=True)
- if attempt < 5:
- time.sleep(2 ** attempt)
- print0(f"Failed to download {filename}")
- return False
-
-
-def main():
- parser = argparse.ArgumentParser(description="Download pretraining dataset shards")
- parser.add_argument("-d", "--dataset", choices=["fineweb-edu", "the-stack-v2-dedup"], required=True)
- parser.add_argument("-n", "--num-files", type=int, default=-1)
- parser.add_argument("-w", "--num-workers", type=int, default=4)
- args = parser.parse_args()
-
- hf_endpoint = os.environ.get("HF_ENDPOINT", "https://huggingface.co").rstrip("/")
- if args.dataset == "fineweb-edu":
- base_url = f"{hf_endpoint}/datasets/karpathy/fineweb-edu-100b-shuffle/resolve/main"
- max_shard = 1822
- else:
- base_url = f"{hf_endpoint}/datasets/smohammadi/the-stack-v2-python-shuffle/resolve/main"
- max_shard = 79
-
- data_dir = DATA_DIR / args.dataset
- data_dir.mkdir(parents=True, exist_ok=True)
- num = max_shard + 1 if args.num_files == -1 else min(args.num_files, max_shard + 1)
- ids = list(range(num))
- print0(f"Downloading {len(ids)} shards from {base_url} → {data_dir}")
- with Pool(processes=args.num_workers) as pool:
- results = pool.map(partial(download_single_file, base_url=base_url, data_dir=data_dir), ids)
- print0(f"Done: {sum(results)}/{len(ids)} shards")
-
-
-if __name__ == "__main__":
- main()
diff --git a/vendor/nanoai/agentic/data/pretrain_mix_1b.py b/vendor/nanoai/agentic/data/pretrain_mix_1b.py
deleted file mode 100644
index e3e4a1e1df25eec7400ed62bb6cc6edc65f505f8..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/pretrain_mix_1b.py
+++ /dev/null
@@ -1,466 +0,0 @@
-"""Self-contained 5-stream pretrain mixer for the 1B-challenge d24 base.
-
-Mixes five document streams with a per-document categorical (Bernoulli) draw,
-then BOS-aligned best-fit packs into (B, T) rows — identical packing to
-``nanoai.agentic.code_dataloader`` (whose ``_doc_batches_from`` / ``_split_paths``
-helpers are reused), but kept in a *separate* module so the validated agentic /
-d6 / d12 code-mixer is untouched (zero regression risk).
-
-Streams (all expose a ``text`` column; verified on wuneng 2026-06-08):
-
- | name | source dir (under base_dir) | files | col |
- |---------|--------------------------------------------------|------:|------|
- | web | base_data_climbmix/ | 1501 | text |
- | owm | code_pretrain_data/open-web-math/data/ | 114 | text |
- | code | code_pretrain_data/the-stack-v2-dedup/ | 510 | text |
- | books | code_pretrain_data/_ebook_orig/ | 20 | text |
- | chinese | chinese_raw/{fineweb_edu_zh,wikipedia_zh}/** | 29 | text |
- | finemath | code_pretrain_data/finemath/finemath-3plus/ | — | text |
- | mathtb | code_pretrain_data/nemotron_specialized/Nemotron-Pretraining-Math-Textbooks/ | — | text |
- | swecpt | code_pretrain_data/swe_mid/cpt/ | — | text |
- | commitdiff | code_pretrain_data/commitpackft/ | — | text |
- | codetests | code_pretrain_data/code_tests/ | — | text |
- | nvcode | code_pretrain_data/nemotron_code/ | — | text |
- | swe1b | code_pretrain_data/swe1b_midcpt/ (REAL SWE-Gym) | — | text |
-
-``swecpt`` is SYNTHETIC (mutation-injected) and DEPRECATED; ``swe1b`` is its REAL
-clean replacement — decontam'd SWE-Gym issue+gold-patch docs from
-``experiments/swe1b`` (serialize_midtrain_cpt -> jsonl_to_parquet). Kept as
-separate streams so real and synthetic SWE data never co-mingle.
-
-The SWE-mid streams (swecpt/commitdiff/codetests/nvcode) are produced by
-``experiments/challenge1b/tools/build_{commitdiff,code_tests,nvcode}_stream.py``
-and ``gen_swe_mut`` (recipe: ``experiments/challenge1b/docs/d6_swe_mid_design.md``).
-
-``owm/code/books/chinese`` ratios are explicit; ``web`` takes the remainder.
-The ratios passed here are PER-DOCUMENT draw probabilities; to hit a target
-*token* share, calibrate them against measured mean-tokens/doc per stream
-(see ``experiments/challenge1b/tools/measure_mix.py``). The recipe of record is
-``experiments/challenge1b/docs/pretrain_recipe.md``.
-
-Resume state is a stub (single epoch counter), matching the code-mixer — fine
-for a single contiguous ~15h run; true multi-stream resume is not implemented.
-"""
-
-from __future__ import annotations
-
-import random
-from pathlib import Path
-
-import pyarrow.parquet as pq
-import torch
-
-from nanoai.common import get_base_dir, get_dist_info, print0
-from nanoai.dataset import list_parquet_files as list_web_parquet_files
-from nanoai.agentic.data.pretrain_data import DATA_DIR as _CODE_BASE_DATA_DIR
-from nanoai.agentic.code_dataloader import _split_paths
-
-
-def _doc_batches_robust(paths, ddp_rank, ddp_world_size, tokenizer_batch_size, split,
- seed: int = 0, shuffle_buffer: int = 10000, text_col: str = "text"):
- """Infinite per-rank list[str] doc-batch iterator: starvation-safe + shuffled.
-
- Two fixes over the upstream code-mixer's ``_doc_batches_from``:
-
- 1. STARVATION. Upstream stripes by ROW-GROUP (``range(rank, num_rg, world)``),
- which starves any rank >= a shard's row-group count. Several streams here
- ship single-row-group shards (``_ebook_orig`` books: 1 rg/shard), so ranks
- 1..3 got ZERO books docs and blocked forever -> DDP deadlock at eval@step0.
- Fix: train stripes by FILE (every stream has >=19 train shards >= world);
- val reads the FULL val shard on every rank (single-file shards; eval just
- averages bpb so cross-rank overlap is harmless).
-
- 2. SHUFFLE. We read shards file-sequentially, and not every stream is
- pre-shuffled within-shard (web/code/owm shards are; chinese isn't
- guaranteed). The bestfit packer selects by LENGTH, not randomly, so it is
- NOT a shuffle. Train docs therefore pass through a reservoir shuffle buffer
- to decorrelate ordering; val is read deterministically (reproducible eval).
- """
- if split == "val":
- my_paths = list(paths)
- else:
- my_paths = paths[ddp_rank::ddp_world_size] or list(paths)
-
- def _docs():
- while True:
- for filepath in my_paths:
- pf = pq.ParquetFile(filepath)
- for rg_idx in range(pf.num_row_groups):
- for d in pf.read_row_group(rg_idx, columns=[text_col]).column(text_col).to_pylist():
- if d:
- yield d
-
- src = _docs()
- rng = random.Random(seed + ddp_rank + 1)
- buf: list[str] = []
- batch: list[str] = []
- while True:
- if split == "train":
- while len(buf) < shuffle_buffer: # reservoir: fill then emit random picks
- buf.append(next(src))
- j = rng.randrange(len(buf))
- buf[j], buf[-1] = buf[-1], buf[j] # swap-to-end + pop = O(1) random eviction
- batch.append(buf.pop())
- else:
- batch.append(next(src)) # val: deterministic order
- if len(batch) >= tokenizer_batch_size:
- yield batch
- batch = []
-
-
-# --------------------------------------------------------------- stream paths
-_CODE_DATASET_NAME = "the-stack-v2-dedup"
-_OWM_SUBPATH = ("open-web-math", "data") # parquets live under open-web-math/data/
-_BOOKS_DATASET_NAME = "_ebook_orig" # intact STEM-textbook parquets (ebook-md was clobbered)
-_CHINESE_DIRNAME = "chinese_raw"
-_CHINESE_SUBDIRS = ("fineweb_edu_zh", "wikipedia_zh")
-
-
-def list_code_parquet_files() -> list[Path]:
- # the-stack-v2-dedup/ also holds 67 ``netis_part_*`` shards (a network-company
- # repo dump: shell/go/c tooling code, off the general-code distribution). Drop
- # them so the ``code`` stream is purely the-stack / starcoder (sc2_*/sc_*/shard_*).
- d = _CODE_BASE_DATA_DIR / _CODE_DATASET_NAME
- if not d.exists():
- return []
- return sorted(p for p in d.glob("*.parquet") if not p.name.startswith("netis"))
-
-
-def list_owm_parquet_files() -> list[Path]:
- d = _CODE_BASE_DATA_DIR.joinpath(*_OWM_SUBPATH)
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_books_parquet_files() -> list[Path]:
- d = _CODE_BASE_DATA_DIR / _BOOKS_DATASET_NAME
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_finemath_parquet_files() -> list[Path]:
- # HuggingFaceTB/finemath finemath-3plus (~34B tok, L2-tier classifier-selected)
- d = _CODE_BASE_DATA_DIR / "finemath" / "finemath-3plus"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_mathtb_parquet_files() -> list[Path]:
- # Nemotron-Pretraining-Math-Textbooks (L3-tier synthetic textbooks, from wukong)
- d = _CODE_BASE_DATA_DIR / "nemotron_specialized" / "Nemotron-Pretraining-Math-Textbooks"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_swecpt_parquet_files() -> list[Path]:
- # gen_swe_mut cpt-view text docs (# buggy / # test / # traceback / # fix)
- d = _CODE_BASE_DATA_DIR / "swe_mid" / "cpt"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_commitdiff_parquet_files() -> list[Path]:
- # CommitPackFT python slice rendered Edit-tool-isomorphic (build_commitdiff_stream.py)
- d = _CODE_BASE_DATA_DIR / "commitpackft"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_codetests_parquet_files() -> list[Path]:
- # the-stack docs that carry their own tests (build_code_tests_stream.py)
- d = _CODE_BASE_DATA_DIR / "code_tests"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_nvcode_parquet_files() -> list[Path]:
- # Nemotron scientific-coding + Llama-Nemotron SFT code (build_nvcode_stream.py)
- d = _CODE_BASE_DATA_DIR / "nemotron_code"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_swe1b_parquet_files() -> list[Path]:
- # REAL SWE-Gym mid corpus (decontam vs Verified): issue + gold-patch bug->fix
- # docs from experiments/swe1b (serialize_midtrain_cpt -> jsonl_to_parquet).
- # This is the clean real-data replacement for the synthetic ``swecpt`` stream;
- # the two are kept SEPARATE so real and synthetic SWE data never mix.
- d = _CODE_BASE_DATA_DIR / "swe1b_midcpt"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_chinese_parquet_files() -> list[Path]:
- base = Path(get_base_dir()) / _CHINESE_DIRNAME
- out: list[Path] = []
- if base.exists():
- for sub in _CHINESE_SUBDIRS:
- d = base / sub
- if d.exists():
- out.extend(sorted(d.rglob("*.parquet"))) # nested IndustryCorpus/ etc.
- return out
-
-
-def list_ebook_nf_parquet_files() -> list[Path]:
- """Re-made LONG non-fiction books (build_ebook_long_stream --split-type): full
- books as single long docs (p50 ~131K tok), NOT the chunked ``_ebook_orig``. The
- genuine long-context natural-language source for the staged decay."""
- d = _CODE_BASE_DATA_DIR / "ebook_long" / "non_fiction"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_ebook_other_parquet_files() -> list[Path]:
- """Re-made LONG other (literature/general) books — long-doc companion to ebook_nf."""
- d = _CODE_BASE_DATA_DIR / "ebook_long" / "other"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_ultrafineweb_parquet_files() -> list[Path]:
- """OpenBMB Ultra-FineWeb (en) — classifier-filtered web (MiniCPM UltraData L2).
- The A5 ablation swaps the legacy climbmix ``web`` stream for this higher-quality
- filtered web. Short docs (web pages), so this is a web-QUALITY lever, not a
- long-context one."""
- d = _CODE_BASE_DATA_DIR / "ultra_fineweb" / "data" / "ultrafineweb_en"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_retrieval_parquet_files() -> list[Path]:
- """Retrieval-shaping synthetic corpus (build_retrieval_stream): long real-prose docs with
- K facts spread across the body + a shuffled recall index at the end, so predicting the
- recalled values forces precise long-range retrieval. Long-context capability lever (push
- reliable needle retrieval past ~64K toward the full 128K)."""
- d = _CODE_BASE_DATA_DIR / "retrieval_shaping"
- return sorted(d.glob("*.parquet")) if d.exists() else []
-
-
-def list_multilingual_parquet_files() -> list[Path]:
- """The 18 高人口/高 GDP 亚非拉 languages (build_multilingual_stream): per-language shards under
- ``multilingual//`` from FineWeb-2 / CulturaX / MADLAD-400 / Wikipedia. The general-purpose
- multilingual lever — needs the new multilingual tokenizer to compress non-Latin scripts (see
- gate_tok_compression / gate_multilingual)."""
- d = _CODE_BASE_DATA_DIR / "multilingual"
- return sorted(d.rglob("*.parquet")) if d.exists() else [] # nested per-language subdirs
-
-
-# Stream name -> lister. ``web`` is handled separately (legacy-fallback aware).
-NON_WEB_LISTERS = {
- "owm": list_owm_parquet_files,
- "code": list_code_parquet_files,
- "books": list_books_parquet_files,
- "chinese": list_chinese_parquet_files,
- "finemath": list_finemath_parquet_files,
- "mathtb": list_mathtb_parquet_files,
- "swecpt": list_swecpt_parquet_files,
- "commitdiff": list_commitdiff_parquet_files,
- "codetests": list_codetests_parquet_files,
- "nvcode": list_nvcode_parquet_files,
- "swe1b": list_swe1b_parquet_files,
- # re-made LONG books (long-context decay source); chunked ``books`` kept for the
- # A0 baseline comparison.
- "ebook_nf": list_ebook_nf_parquet_files,
- "ebook_other": list_ebook_other_parquet_files,
- # filtered-web (MiniCPM UltraData L2); A5 swaps legacy ``web`` -> this.
- "ultrafineweb": list_ultrafineweb_parquet_files,
- # retrieval-shaping long docs (push reliable retrieval past ~64K toward 128K).
- "retrieval": list_retrieval_parquet_files,
- # the 18 亚非拉 languages (general-purpose multilingual lever).
- "multilingual": list_multilingual_parquet_files,
-}
-# Additive decision-order slices of [0,1): APPEND-ONLY — reordering would change
-# the realized ratios of every existing recipe (web takes the remainder).
-STREAM_ORDER = ("owm", "code", "books", "chinese", "finemath", "mathtb",
- "swecpt", "commitdiff", "codetests", "nvcode", "swe1b",
- "ebook_nf", "ebook_other", "ultrafineweb", "retrieval",
- "multilingual")
-# Per-stream parquet text column (default "text"). Ultra-FineWeb ships "content".
-_STREAM_TEXT_COL = {"ultrafineweb": "content"}
-
-
-def _resolve_paths(split: str, warn: bool):
- """Return {name: [train/val paths]} for web + the four non-web streams."""
- web = [Path(p) for p in _split_paths(list_web_parquet_files(warn_on_legacy=warn), split)]
- assert web, "No web (climbmix) parquet shards found; run nanoai.dataset first."
- paths = {"web": web}
- for name, lister in NON_WEB_LISTERS.items():
- paths[name] = _split_paths(lister(), split)
- return paths
-
-
-# --------------------------------------------------------------- doc batches
-def _mixed_document_batches_1b(
- split: str,
- ratios: dict[str, float],
- tokenizer_batch_size: int,
- seed: int = 0,
-):
- """Infinite stream of list[str] doc batches drawn from the 5-way mixture.
-
- ``ratios`` holds per-document draw probabilities for the non-web streams
- (owm/code/books/chinese); web gets the remainder. Decision order matches
- ``STREAM_ORDER`` then web — additive slices of [0, 1).
- """
- for k in ratios:
- assert k in NON_WEB_LISTERS, f"unknown stream ratio {k!r}"
- assert 0.0 <= ratios[k] <= 1.0, f"bad ratio {k}={ratios[k]}"
- non_web_total = sum(ratios.get(n, 0.0) for n in STREAM_ORDER)
- assert non_web_total <= 1.0 + 1e-9, f"non-web ratios sum {non_web_total} > 1"
-
- ddp, ddp_rank, _local, ddp_world_size = get_dist_info()
- warn = (ddp_rank == 0) and (split == "train")
- paths = _resolve_paths(split, warn)
-
- for name in STREAM_ORDER:
- if ratios.get(name, 0.0) > 0 and not paths[name]:
- raise FileNotFoundError(
- f"{name}_ratio>0 but no parquet found for stream '{name}'. "
- f"Expected under base_dir; see pretrain_mix_1b stream table."
- )
-
- iters = {
- "web": _doc_batches_robust(paths["web"], ddp_rank, ddp_world_size, tokenizer_batch_size, split)
- }
- for name in STREAM_ORDER:
- iters[name] = (
- _doc_batches_robust(paths[name], ddp_rank, ddp_world_size, tokenizer_batch_size, split,
- text_col=_STREAM_TEXT_COL.get(name, "text"))
- if paths[name] else None
- )
-
- if ddp_rank == 0:
- web_ratio = max(0.0, 1.0 - non_web_total)
- shard_str = " ".join(f"{n}={len(paths[n])}" for n in ("web", *STREAM_ORDER))
- ratio_str = " ".join(
- f"{n}={ratios.get(n, 0.0):.3f}" for n in STREAM_ORDER
- )
- print0(f"[mix-1b] split={split} web={web_ratio:.3f} {ratio_str} | shards {shard_str}")
-
- rng = random.Random(seed + ddp_rank)
- counts = {n: 0 for n in ("web", *STREAM_ORDER)}
- log_every = 10000
- while True:
- r = rng.random()
- acc = 0.0
- chosen = "web"
- for name in STREAM_ORDER:
- acc += ratios.get(name, 0.0)
- if iters[name] is not None and r < acc:
- chosen = name
- break
- counts[chosen] += 1
- yield next(iters[chosen])
- total = sum(counts.values())
- if ddp_rank == 0 and total % log_every == 0:
- frac = " ".join(f"{n}={counts[n]/total:.3f}" for n in ("web", *STREAM_ORDER))
- print0(f"[mix-1b] draws={total} {frac}")
-
-
-# --------------------------------------------------------------- bestfit packer
-def tokenizing_mixed_1b(
- tokenizer,
- B: int,
- T: int,
- split: str,
- ratios: dict[str, float],
- tokenizer_threads: int = 4,
- tokenizer_batch_size: int = 128,
- device: str = "cuda",
- buffer_size: int = 1000,
- seed: int = 0,
-):
- """BOS-aligned best-fit packing over the 5-way mixture. Yields (inputs,
- targets) tensors of shape (B, T). Drop-in for nanoai's
- ``tokenizing_distributed_data_loader_bos_bestfit``.
- """
- assert split in ("train", "val")
- row_capacity = T + 1
- bos_token = tokenizer.get_bos_token_id()
-
- batches = _mixed_document_batches_1b(split, ratios, tokenizer_batch_size, seed=seed)
- doc_buffer: list[list[int]] = []
-
- def refill_buffer():
- doc_batch = next(batches)
- token_lists = tokenizer.encode(doc_batch, prepend=bos_token, num_threads=tokenizer_threads)
- for tokens in token_lists:
- doc_buffer.append(tokens)
-
- use_cuda = device == "cuda"
- row_buffer = torch.empty((B, row_capacity), dtype=torch.long)
- cpu_buffer = torch.empty(2 * B * T, dtype=torch.long, pin_memory=use_cuda)
- gpu_buffer = torch.empty(2 * B * T, dtype=torch.long, device=device)
- cpu_inputs = cpu_buffer[: B * T].view(B, T)
- cpu_targets = cpu_buffer[B * T :].view(B, T)
- inputs = gpu_buffer[: B * T].view(B, T)
- targets = gpu_buffer[B * T :].view(B, T)
-
- while True:
- for row_idx in range(B):
- pos = 0
- while pos < row_capacity:
- while len(doc_buffer) < buffer_size:
- refill_buffer()
- remaining = row_capacity - pos
-
- best_idx = -1
- best_len = 0
- for i, doc in enumerate(doc_buffer):
- doc_len = len(doc)
- if doc_len <= remaining and doc_len > best_len:
- best_idx = i
- best_len = doc_len
-
- if best_idx >= 0:
- doc = doc_buffer.pop(best_idx)
- doc_len = len(doc)
- row_buffer[row_idx, pos : pos + doc_len] = torch.tensor(doc, dtype=torch.long)
- pos += doc_len
- else:
- shortest_idx = min(range(len(doc_buffer)), key=lambda i: len(doc_buffer[i]))
- doc = doc_buffer.pop(shortest_idx)
- row_buffer[row_idx, pos : pos + remaining] = torch.tensor(
- doc[:remaining], dtype=torch.long
- )
- pos += remaining
-
- cpu_inputs.copy_(row_buffer[:, :-1])
- cpu_targets.copy_(row_buffer[:, 1:])
- gpu_buffer.copy_(cpu_buffer, non_blocking=use_cuda)
- yield inputs, targets
-
-
-def tokenizing_mixed_1b_with_state(
- tokenizer,
- B: int,
- T: int,
- split: str,
- ratios: dict[str, float],
- tokenizer_threads: int = 4,
- tokenizer_batch_size: int = 128,
- device: str = "cuda",
- resume_state_dict=None,
- buffer_size: int = 1000,
- seed: int = 0,
-):
- """Same as ``tokenizing_mixed_1b`` but yields a third stub ``state_dict`` to
- match ``tokenizing_distributed_data_loader_with_state_bos_bestfit``.
- Caller's ``resume_state_dict`` is ignored (multi-stream resume not impl.)."""
- base = tokenizing_mixed_1b(
- tokenizer, B, T, split, ratios,
- tokenizer_threads=tokenizer_threads,
- tokenizer_batch_size=tokenizer_batch_size,
- device=device, buffer_size=buffer_size, seed=seed,
- )
- stub_state = {"epoch": 1, "pq_idx": 0, "rg_idx": 0}
- for inputs, targets in base:
- yield inputs, targets, stub_state
-
-
-# --------------------------------------------------------------- tokenizer-train iterator
-def stream_text_iterator(ratios: dict[str, float], doc_cap: int, max_chars: int, seed: int = 0):
- """Yield individual doc strings drawn from the 5-way mixture (single rank),
- each truncated to ``doc_cap`` chars, stopping after ``max_chars`` total.
-
- Used by ``scripts/tok_train_1b.py`` so the BPE sees a char distribution that
- matches the pretrain mix (English / math / code / textbooks / Chinese).
- """
- batches = _mixed_document_batches_1b("train", ratios, tokenizer_batch_size=128, seed=seed)
- nchars = 0
- while nchars < max_chars:
- batch = next(batches)
- for doc in batch:
- doc = doc[:doc_cap]
- nchars += len(doc)
- yield doc
- if nchars >= max_chars:
- return
diff --git a/vendor/nanoai/agentic/data/routing_recovery.py b/vendor/nanoai/agentic/data/routing_recovery.py
deleted file mode 100644
index a0a480153624a0f05c86d0e47ea4084c7336ffcf..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/routing_recovery.py
+++ /dev/null
@@ -1,357 +0,0 @@
-"""Synthetic tool-routing recovery dataset.
-
-Motivation (POC-8 mechanistic analysis):
- v9 (lc=10) regressed from v5 by 5.3pp because long_ctx data lacks tool calls,
- causing catastrophic forgetting of *which tool to use for which task type*.
- Failed cases (read_config, grep_torch, bash_grep_log) showed v9 emitting
- Edit-only with degenerate repetitive content where v5 emitted the correct
- Read/Grep/Bash chain.
-
- This dataset teaches the model the prompt-shape → first-tool mapping
- directly: 4 tool families × ~50 prompt templates × 4-5 surface variants =
- ~200 rows per tool, ~800 rows total. Each row is a 4-message conversation
- (user → assistant tool_call → synthetic tool_result → assistant answer).
-
-Format on disk: JSONL, one `{"messages": [...]}` per line, compatible with
-`render_agentic_conversation`'s tool-call branch (assistant turn carries
-both `content` and `tool_call: {name, args}`).
-
-CLI:
- python -m nanoai.agentic.data.routing_recovery \
- --output data/routing_recovery.jsonl --seed 42 --n-per-tool 200
-"""
-
-from __future__ import annotations
-
-import argparse
-import json
-import random
-from dataclasses import dataclass
-from pathlib import Path
-from typing import Any
-
-from nanoai.agentic.data.common import SYSTEM_PROMPT
-
-
-# ---- File / pattern / path corpora (deterministic per-seed shuffling) -------
-
-_FILES_PY = [
- "main.py", "app.py", "lib.py", "utils.py", "config.py", "server.py",
- "client.py", "models.py", "schema.py", "data.py", "tools.py", "test.py",
- "parser.py", "router.py", "loader.py", "engine.py", "core.py", "common.py",
- "tasks.py", "errors.py",
-]
-_FILES_DATA = [
- "config.toml", "settings.json", "data.csv", "users.json", "log.txt",
- "README.md", "Makefile", "requirements.txt", "pyproject.toml", ".env",
- "dataset.json", "errors.log", "access.log", "metrics.tsv", "manifest.yaml",
-]
-_DIRS = ["src", "lib", "tests", "scripts", "data", "config", ".", "core", "app"]
-_PATTERNS = [
- "TODO", "FIXME", "import torch", "def __init__", "raise ValueError",
- "class .*Error", "logging.basicConfig", "yield from", "async def",
- "from typing import", "with pytest.raises", "@dataclass",
- "logger.info", "self\\.config", "return None",
-]
-_BASH_TARGETS = [
- "list files", "count lines", "show disk usage", "show file size",
- "find files larger than 1mb", "show last 10 lines", "show first 5 lines",
- "show file modification time",
-]
-
-
-# ---- Templates per tool family ---------------------------------------------
-#
-# Each template is a function (rng) -> {prompt, tool_name, tool_args,
-# observation, final_answer}.
-
-@dataclass
-class _Row:
- user_prompt: str
- pre_text: str # short assistant pre-text before tool_call
- tool_name: str
- tool_args: dict[str, str]
- tool_result: str # synthetic observation
- final_text: str # assistant's answer after seeing observation
-
-
-def _read_row(rng: random.Random) -> _Row:
- f = rng.choice(_FILES_PY + _FILES_DATA)
- template = rng.choice([
- ("what's in {f}?", "let me read it"),
- ("show me {f}", "i'll open it"),
- ("can you read {f}?", "reading"),
- ("open {f} and tell me what it does", "i'll read the file first"),
- ("what does {f} say?", "checking"),
- ("read {f} please", "ok"),
- ("display the contents of {f}", "ok"),
- ("look at {f}", "reading"),
- ("how is {f} structured?", "let me see"),
- ("give me the source of {f}", "fetching"),
- ])
- prompt = template[0].format(f=f)
- pre = template[1]
- if f.endswith(".py"):
- body = f"def main():\n print('hello from {f}')\n\nif __name__ == '__main__':\n main()"
- ans = f"{f} defines a `main()` that prints a hello banner; runs only when invoked directly."
- elif f == "config.toml":
- ans = "the version is 0.1.0; package name is `nanocode`."
- body = "[package]\nname = 'nanocode'\nversion = '0.1.0'\n"
- elif f == "README.md":
- body = "# Project\n\nA small library.\n\n## Install\n`pip install .`\n"
- ans = "small library titled 'Project' with a single install instruction."
- elif f.endswith(".json"):
- body = '{"version": "1.2", "users": [{"id": 1}, {"id": 2}]}'
- ans = "JSON config with version 1.2 and two users."
- elif f.endswith(".log"):
- body = "2026-05-07 INFO start\n2026-05-07 INFO ready\n"
- ans = "two INFO log lines, both from 2026-05-07."
- else:
- body = "# placeholder file body\n"
- ans = f"{f} is a small placeholder; nothing surprising."
- return _Row(
- user_prompt=prompt,
- pre_text=pre,
- tool_name="Read",
- tool_args={"file_path": f},
- tool_result=body,
- final_text=ans,
- )
-
-
-def _grep_row(rng: random.Random) -> _Row:
- pat = rng.choice(_PATTERNS)
- where = rng.choice(_DIRS)
- template = rng.choice([
- ("find {pat} in {where}", "let me search"),
- ("where is {pat} used in {where}?", "grepping"),
- ("search for {pat} under {where}", "ok"),
- ("any files containing {pat}?", "let me check"),
- ("show me uses of {pat} in {where}", "searching"),
- ("does {pat} appear anywhere?", "looking"),
- ("locate {pat} in the codebase", "grepping the tree"),
- ("which files reference {pat}?", "checking"),
- ("grep {pat}", "running grep"),
- ("are there any {pat} markers?", "let me see"),
- ])
- prompt = template[0].format(pat=pat, where=where)
- pre = template[1]
- n_hits = rng.randint(1, 4)
- hits = "\n".join(
- f"{rng.choice(_FILES_PY)}:{rng.randint(1, 200)}: {pat}"
- for _ in range(n_hits)
- )
- if n_hits == 0:
- ans = f"no matches for {pat}."
- elif n_hits == 1:
- ans = f"one hit for {pat} — single-callsite, easy to update."
- else:
- ans = f"{n_hits} matches for {pat} across the tree; the {pat.split()[0] if ' ' in pat else pat} pattern is consistent."
- return _Row(
- user_prompt=prompt,
- pre_text=pre,
- tool_name="Grep",
- tool_args={"pattern": pat, "path": where},
- tool_result=hits or "(no matches)",
- final_text=ans,
- )
-
-
-def _bash_row(rng: random.Random) -> _Row:
- f = rng.choice(_FILES_PY + _FILES_DATA)
- d = rng.choice(_DIRS)
- choice = rng.choice([
- ("list", f"list files in {d}", "ls", f"ls {d}",
- "main.py app.py utils.py tests"),
- ("count", f"count lines in {f}", "wc", f"wc -l {f}",
- f"{rng.randint(20, 800)} {f}"),
- ("size", f"what is the size of {f}?", "du", f"du -h {f}",
- f"{rng.randint(2, 99)}K\t{f}"),
- ("first", f"show first 5 lines of {f}", "head", f"head -n 5 {f}",
- "line 1\nline 2\nline 3\nline 4\nline 5"),
- ("last", f"show the last 10 lines of {f}", "tail", f"tail -n 10 {f}",
- "tail line a\ntail line b\ntail line c"),
- ("exists", f"does {f} exist?", "test", f"test -f {f} && echo yes || echo no",
- "yes"),
- ("find_large", f"find files larger than 1mb under {d}", "find", f"find {d} -size +1M",
- f"{d}/data/big.bin"),
- ("mod_time", f"when was {f} last modified?", "stat", f"stat -c %y {f}",
- "2026-05-06 14:22:11.000000000 +0000"),
- ("grep_simple", f"any errors in {f}?", "grep", f"grep -i error {f}",
- f"{f}:10: raise ValueError('boom')"),
- ])
- kind, prompt, _short, cmd, body = choice
- pre = rng.choice(["running it", "let me check", "ok", "checking"])
- ans_map = {
- "list": f"the {d} directory has 4 files: main.py, app.py, utils.py, and tests.",
- "count": f"{f} is {body.split()[0]} lines.",
- "size": f"{f} is {body.split(chr(9))[0]}.",
- "first": f"first 5 lines are 'line 1' through 'line 5'.",
- "last": "last lines are 'tail line a/b/c'.",
- "exists": f"{f} exists.",
- "find_large": f"one file over 1MB: {body}.",
- "mod_time": "last modified 2026-05-06 at 14:22.",
- "grep_simple": f"one error path on line 10 raises ValueError.",
- }
- return _Row(
- user_prompt=prompt,
- pre_text=pre,
- tool_name="Bash",
- tool_args={"command": cmd},
- tool_result=body,
- final_text=ans_map[kind],
- )
-
-
-def _edit_row(rng: random.Random) -> _Row:
- f = rng.choice(_FILES_PY)
- edit = rng.choice([
- ("change_version", f"bump version in {f}",
- "version = '0.1.0'", "version = '0.2.0'",
- "version = '0.2.0'"),
- ("add_docstring", f"add a one-line docstring to {f}",
- "def main():", "def main():\n \"\"\"entry point\"\"\"",
- "def main():\n \"\"\"entry point\"\"\""),
- ("rename_var", f"rename `tmp` to `result` in {f}",
- "tmp = compute()", "result = compute()",
- "result = compute()"),
- ("add_import", f"add `import logging` at the top of {f}",
- "", "import logging",
- "import logging"),
- ("fix_typo", f"fix the typo `recieve` → `receive` in {f}",
- "def recieve(x):", "def receive(x):",
- "def receive(x):"),
- ("create", f"create a new file {f} that prints hello",
- "", "print('hello')",
- "print('hello')"),
- ("replace_print", f"replace print() with logger.info() in {f}",
- "print(x)", "logger.info(x)",
- "logger.info(x)"),
- ("add_main_guard", f"wrap the body of {f} in if __name__ == '__main__':",
- "main()", "if __name__ == '__main__':\n main()",
- "if __name__ == '__main__':\n main()"),
- ])
- kind, prompt, old, new, applied = edit
- pre = rng.choice([
- "applying the edit", "ok", "let me make that change",
- "editing", "doing it now",
- ])
- args = {"file_path": f, "new_string": new}
- if old:
- args["old_string"] = old
- body = f"applied the edit to {f}: now contains:\n{applied}"
- ans = rng.choice([
- f"done — {f} now uses the updated form.",
- f"applied. {f} is updated.",
- "edit applied successfully.",
- f"updated {f} as requested.",
- ])
- return _Row(
- user_prompt=prompt,
- pre_text=pre,
- tool_name="Edit",
- tool_args=args,
- tool_result=body,
- final_text=ans,
- )
-
-
-_GENERATORS: dict[str, Any] = {
- "Read": _read_row,
- "Grep": _grep_row,
- "Bash": _bash_row,
- "Edit": _edit_row,
-}
-
-
-# ---- Row builder -----------------------------------------------------------
-
-
-def _row_to_messages(r: _Row) -> dict[str, Any]:
- return {
- "messages": [
- {"role": "user", "content": r.user_prompt},
- {
- "role": "assistant",
- "content": r.pre_text,
- "tool_call": {"name": r.tool_name, "args": r.tool_args},
- },
- {"role": "tool_result", "content": r.tool_result},
- {"role": "assistant", "content": r.final_text},
- ]
- }
-
-
-def generate_routing_rows(
- seed: int = 42,
- n_per_tool: int = 200,
- tools: tuple[str, ...] = ("Read", "Grep", "Bash", "Edit"),
-) -> list[dict[str, Any]]:
- """Generate the routing-recovery rows in memory.
-
- Deterministic per `seed`; produces `n_per_tool * len(tools)` rows total.
- """
- rng = random.Random(seed)
- out: list[dict[str, Any]] = []
- for tool in tools:
- gen = _GENERATORS[tool]
- for i in range(n_per_tool):
- row = gen(random.Random(seed * 1_000_003 + hash(tool) + i))
- out.append(_row_to_messages(row))
- rng.shuffle(out)
- return out
-
-
-# ---- SFT-side dataset class ------------------------------------------------
-
-
-class RoutingRecoveryDataset:
- """In-memory dataset usable by `agentic_sft.py` task mixtures.
-
- Mirrors the `JSONDataset` protocol (len + getitem + system prompt prepend).
- Generated synthetically — no on-disk file required, but the CLI below can
- write a JSONL artefact for inspection / reproducibility.
- """
-
- def __init__(
- self,
- seed: int = 42,
- n_per_tool: int = 200,
- tools: tuple[str, ...] = ("Read", "Grep", "Bash", "Edit"),
- **_: Any,
- ) -> None:
- self.ds = generate_routing_rows(seed=seed, n_per_tool=n_per_tool, tools=tools)
-
- def __len__(self) -> int:
- return len(self.ds)
-
- def __getitem__(self, idx: int) -> dict[str, Any]:
- messages = self.ds[idx]["messages"]
- return {"messages": [{"role": "system", "content": SYSTEM_PROMPT}] + messages}
-
-
-# ---- CLI -------------------------------------------------------------------
-
-
-def _main() -> None:
- p = argparse.ArgumentParser()
- p.add_argument("--output", required=True, type=Path,
- help="Output JSONL path (one {messages:[...]} per line)")
- p.add_argument("--seed", type=int, default=42)
- p.add_argument("--n-per-tool", type=int, default=200,
- help="Rows per tool family (default 200 → 800 total)")
- p.add_argument("--tools", type=str, default="Read,Grep,Bash,Edit",
- help="Comma-separated tool families to generate")
- args = p.parse_args()
-
- tools = tuple(t.strip() for t in args.tools.split(",") if t.strip())
- rows = generate_routing_rows(seed=args.seed, n_per_tool=args.n_per_tool, tools=tools)
- args.output.parent.mkdir(parents=True, exist_ok=True)
- with args.output.open("w", encoding="utf-8") as f:
- for row in rows:
- f.write(json.dumps(row, ensure_ascii=False) + "\n")
- print(f"wrote {len(rows)} rows ({args.n_per_tool}/tool × {len(tools)} tools) → {args.output}")
-
-
-if __name__ == "__main__":
- _main()
diff --git a/vendor/nanoai/agentic/data/sft_data_1b.py b/vendor/nanoai/agentic/data/sft_data_1b.py
deleted file mode 100644
index 23110b79bbab72943577d8fbc25a28e2febe04ce..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/sft_data_1b.py
+++ /dev/null
@@ -1,240 +0,0 @@
-"""1B-challenge SFT data: UltraData-SFT-2605 -> hybrid-thinking (ids, mask).
-
-Self-contained, mirrors the pretrain_mix_1b approach (zero change to validated
-code). Two responsibilities:
-
-1. ``render_sft_example`` — turn ONE UltraData example into ``(ids, mask)`` with
- the **Qwen3/MiniCPM-style hybrid** layout so a single ckpt can switch
- think/no-think reliably:
- think : <|assistant_start|> <|think_start|> {reasoning} <|think_end|> {answer} <|assistant_end|>
- no_think : <|assistant_start|> <|think_start|> <|think_end|> {answer} <|assistant_end|> (EMPTY think block)
- **mode_tag (hybrid-think switch fix)**: with ``mode_tag=True`` the user turn is
- suffixed with `` /think`` (think examples) or `` /no_think`` (no_think), mask=0.
- Without this prompt-side signal the think/no_think prompts are IDENTICAL, so a
- greedy decoder cannot tell the modes apart and collapses to the majority
- no-think path (empty <|think_start|><|think_end|>, reasoning leaks into the
- answer body). The tag is the deterministic control the model conditions on;
- at inference append the same `` /think``/`` /no_think`` to select the mode.
- ``mask=1`` on every assistant token the model must GENERATE (think block —
- incl. the empty one — + answer + assistant_end). **think tokens are emitted
- via encode_special, NEVER as literal "<|think_start|>" text** (literal text
- goes through encode_ordinary and is NOT the special token — the footgun that
- tanked KSR Phase 3; see lesson_think_tokens_unencoded).
-
-2. ``iter_and_sample`` / ``prepare_sft_cache`` — stream the 7 gated configs with
- the HF token, **budget-match to TARGET TOKEN shares** (render to count tokens),
- keep the ~think/no_think ratio, run a light eval-set decontam, and write the
- sampled raw conversations to a jsonl cache. The SFT trainer renders on the fly
- (like chat_sft) at max_seq_len=32768.
-
-UltraData schema (confirmed): configs {Math,Code,Knowledge,Chinese-general,IF,
-Multi-lang-Math,Multi-lang-Knowledge}; splits {think,no_think}; per example
-``messages=[{role,content,reasoning_content}]`` (reasoning_content = long CoT on
-the assistant turn for think; null for no_think), plus uid/source/domain/think_type.
-"""
-from __future__ import annotations
-
-import argparse
-import json
-import os
-import random
-
-# ---------------------------------------------------------------- rendering
-
-def _normalize_messages(messages):
- """Merge a leading system message into the first user turn (mirrors
- tokenizer.render_conversation); return the cleaned message list."""
- if messages and messages[0].get("role") == "system":
- if len(messages) >= 2 and messages[1].get("role") == "user":
- merged = dict(messages[1])
- merged["content"] = (messages[0].get("content") or "") + "\n\n" + (messages[1].get("content") or "")
- return [merged] + list(messages[2:])
- return list(messages[1:])
- return list(messages)
-
-
-# Prompt-side hybrid-think control tags (Qwen3-style). Appended to user turns
-# when mode_tag=True so the model conditions on the requested mode. Import these
-# in inference/eval code to activate thinking deterministically.
-THINK_TAG = " /think"
-NOTHINK_TAG = " /no_think"
-
-
-def render_sft_example(tokenizer, messages, think: bool, max_tokens: int = 32768,
- mode_tag: bool = False):
- """One UltraData example -> (ids, mask). See module docstring for layout.
-
- mode_tag=True appends the prompt-side ``/think`` / ``/no_think`` switch to the
- user turn (mask=0) so the hybrid-think mode is controllable instead of
- collapsing to empty-think under greedy decoding.
- """
- ids, mask = [], []
-
- def add(toks, m):
- if isinstance(toks, int):
- toks = [toks]
- ids.extend(toks)
- mask.extend([m] * len(toks))
-
- bos = tokenizer.get_bos_token_id()
- us = tokenizer.encode_special("<|user_start|>")
- ue = tokenizer.encode_special("<|user_end|>")
- as_ = tokenizer.encode_special("<|assistant_start|>")
- ae = tokenizer.encode_special("<|assistant_end|>")
- ts = tokenizer.encode_special("<|think_start|>")
- te = tokenizer.encode_special("<|think_end|>")
-
- add(bos, 0)
- for m in _normalize_messages(messages):
- role = m.get("role")
- content = m.get("content") or ""
- if role == "user":
- add(us, 0)
- if mode_tag:
- content = content + (THINK_TAG if think else NOTHINK_TAG)
- add(tokenizer.encode(content), 0)
- add(ue, 0)
- elif role == "assistant":
- add(as_, 0)
- # think block (empty for no_think) — masked 1 so the switch is learned
- add(ts, 1)
- if think:
- rc = m.get("reasoning_content") or ""
- if rc:
- add(tokenizer.encode(rc), 1)
- add(te, 1)
- # final answer — masked 1
- if content:
- add(tokenizer.encode(content), 1)
- add(ae, 1)
- # any other role is skipped (UltraData is user/assistant only)
- return ids[:max_tokens], mask[:max_tokens]
-
-
-# ---------------------------------------------------------------- sampling
-
-# TARGET TOKEN shares across configs (sft_design §2). Multi-lang skipped in v1.
-DEFAULT_MIX = {"Math": 0.38, "Code": 0.32, "Knowledge": 0.12,
- "Chinese-general": 0.10, "IF": 0.08}
-# GENERAL-PURPOSE mix: turns on the UltraData Multi-lang configs (~17% multilingual SFT) for the
-# general_1b recipe, rebalanced to sum to 1.0. Leaves DEFAULT_MIX (the 1b-challenge v1 recipe)
-# untouched — selected via `prepare_sft_cache --mix general`.
-GENERAL_MIX = {"Math": 0.30, "Code": 0.25, "Knowledge": 0.12, "Chinese-general": 0.08,
- "IF": 0.08, "Multi-lang-Math": 0.09, "Multi-lang-Knowledge": 0.08}
-MIXES = {"default": DEFAULT_MIX, "general": GENERAL_MIX}
-THINK_RATIO = 0.43 # ~43% thinking / 57% non-thinking (UltraData global)
-
-
-def _decontam_ngrams(eval_texts, n=13):
- grams = set()
- for t in eval_texts:
- toks = t.split()
- for i in range(len(toks) - n + 1):
- grams.add(" ".join(toks[i:i + n]))
- return grams
-
-
-def _is_contaminated(messages, bad_grams, n=13):
- if not bad_grams:
- return False
- txt = " ".join((m.get("content") or "") for m in messages)
- toks = txt.split()
- for i in range(len(toks) - n + 1):
- if " ".join(toks[i:i + n]) in bad_grams:
- return True
- return False
-
-
-def iter_and_sample(tokenizer, *, mix=None, think_ratio=THINK_RATIO,
- total_token_budget=2_000_000_000, max_tokens=32768,
- seed=0, hf_token=None, decontam_grams=None, log=print):
- """Stream UltraData, render to count tokens, yield sampled raw examples
- (dicts with messages/think/domain/source/uid/n_tokens) until each config's
- TOKEN sub-budget (mix * budget, split think_ratio) is met."""
- from datasets import load_dataset
- mix = mix or DEFAULT_MIX
- rng = random.Random(seed)
- if hf_token:
- os.environ.setdefault("HF_TOKEN", hf_token)
- os.environ.setdefault("HUGGING_FACE_HUB_TOKEN", hf_token)
-
- for cfg, share in mix.items():
- cfg_budget = total_token_budget * share
- for split, frac in (("think", think_ratio), ("no_think", 1.0 - think_ratio)):
- sub_budget = cfg_budget * frac
- got_tok = 0
- n_skip_contam = 0
- ds = load_dataset("openbmb/UltraData-SFT-2605", cfg, split=split,
- streaming=True)
- ds = ds.shuffle(seed=seed + hash(cfg + split) % 1000, buffer_size=10000)
- n_err = 0
- for ex in ds:
- if got_tok >= sub_budget:
- break
- try:
- msgs = ex.get("messages") or []
- if not msgs:
- continue
- if decontam_grams and _is_contaminated(msgs, decontam_grams):
- n_skip_contam += 1
- continue
- think = (split == "think")
- ids, _ = render_sft_example(tokenizer, msgs, think, max_tokens)
- n = len(ids)
- if n < 8:
- continue
- except Exception: # one malformed row must not kill an hour-long prep
- n_err += 1
- if n_err <= 3:
- log(f"[sft_data] {cfg}/{split}: skipped bad row ({n_err})")
- continue
- got_tok += n
- yield {"messages": msgs, "think": think, "domain": ex.get("domain", cfg),
- "source": ex.get("source", ""), "uid": ex.get("uid", ""),
- "n_tokens": n}
- log(f"[sft_data] {cfg}/{split}: ~{got_tok:,} tok "
- f"(budget {sub_budget:,.0f}), decontam-skipped {n_skip_contam}")
-
-
-def prepare_sft_cache(out_path, tokenizer, *, total_token_budget, max_tokens=32768,
- seed=0, hf_token=None, eval_texts=None, mix=None):
- """Write budget-matched sampled conversations to a jsonl cache (raw messages;
- the trainer re-renders on the fly). Returns summary dict."""
- decontam = _decontam_ngrams(eval_texts) if eval_texts else None
- os.makedirs(os.path.dirname(out_path), exist_ok=True)
- n_ex = tot_tok = 0
- by_domain, by_think = {}, {"think": 0, "no_think": 0}
- with open(out_path, "w") as f:
- for row in iter_and_sample(tokenizer, mix=mix, total_token_budget=total_token_budget,
- max_tokens=max_tokens, seed=seed, hf_token=hf_token,
- decontam_grams=decontam):
- f.write(json.dumps(row, ensure_ascii=False) + "\n")
- n_ex += 1
- tot_tok += row["n_tokens"]
- by_domain[row["domain"]] = by_domain.get(row["domain"], 0) + row["n_tokens"]
- by_think["think" if row["think"] else "no_think"] += row["n_tokens"]
- summary = {"out": out_path, "n_examples": n_ex, "total_tokens": tot_tok,
- "by_domain_tokens": by_domain, "by_think_tokens": by_think}
- print("[sft_data] CACHE SUMMARY:", json.dumps(summary, ensure_ascii=False))
- return summary
-
-
-def main():
- ap = argparse.ArgumentParser()
- ap.add_argument("--out", required=True, help="jsonl cache path")
- ap.add_argument("--token-budget", type=float, default=2e9)
- ap.add_argument("--max-tokens", type=int, default=32768)
- ap.add_argument("--seed", type=int, default=0)
- ap.add_argument("--hf-token", default=os.environ.get("HF_TOKEN", ""))
- ap.add_argument("--mix", choices=sorted(MIXES), default="default",
- help="SFT config mix: 'default' (1b-challenge v1) or 'general' (+multilingual)")
- args = ap.parse_args()
- from nanoai.tokenizer import get_tokenizer
- tok = get_tokenizer()
- prepare_sft_cache(args.out, tok, total_token_budget=int(args.token_budget),
- max_tokens=args.max_tokens, seed=args.seed,
- hf_token=args.hf_token or None, mix=MIXES[args.mix])
-
-
-if __name__ == "__main__":
- main()
diff --git a/vendor/nanoai/agentic/data/synth_data.py b/vendor/nanoai/agentic/data/synth_data.py
deleted file mode 100644
index 4b145858102ef4b63c279e8399b37528c4c7eb50..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/synth_data.py
+++ /dev/null
@@ -1,74 +0,0 @@
-"""SYNTH (PleIAs) parquet iterator for d6 ablation pretrain.
-
-Concatenates each row's `query` + `synthetic_reasoning` + `synthetic_answer`
-into a single training document string. Mirrors the iterator interface used
-by `pretrain_data.py` so the BOS-aligned bestfit packer can consume it as a
-drop-in document stream.
-
-Default location: $NANOAI_BASE_DIR/code_pretrain_data/synth-pleias/
-(can be a symlink to the actual SYNTH dir).
-"""
-
-from __future__ import annotations
-
-from pathlib import Path
-
-import pyarrow.parquet as pq
-
-from nanoai.common import get_base_dir
-
-
-DATA_DIR = Path(get_base_dir()) / "code_pretrain_data" / "synth-pleias"
-
-
-def list_synth_parquet_files() -> list[Path]:
- if not DATA_DIR.exists():
- return []
- return sorted([DATA_DIR / f for f in DATA_DIR.iterdir() if f.suffix == ".parquet"])
-
-
-_CONCAT_SEP = "\n\n"
-
-
-def _row_to_text(row: dict) -> str:
- q = row.get("query") or ""
- r = row.get("synthetic_reasoning") or ""
- a = row.get("synthetic_answer") or ""
- return _CONCAT_SEP.join(x for x in (q, r, a) if x)
-
-
-def synth_doc_batches(
- paths: list[Path],
- ddp_rank: int,
- ddp_world_size: int,
- tokenizer_batch_size: int,
-):
- """Infinite per-rank iterator: yields list[str] doc batches.
-
- Each parquet row group is striped across ranks. Each row yields one
- concatenated `query \\n\\n reasoning \\n\\n answer` document.
- """
- cols = ["query", "synthetic_reasoning", "synthetic_answer"]
- while True:
- for filepath in paths:
- pf = pq.ParquetFile(filepath)
- for rg_idx in range(ddp_rank, pf.num_row_groups, ddp_world_size):
- rg = pf.read_row_group(rg_idx, columns=cols)
- rows = rg.to_pylist()
- batch: list[str] = []
- for r in rows:
- txt = _row_to_text(r)
- if not txt:
- continue
- batch.append(txt)
- if len(batch) >= tokenizer_batch_size:
- yield batch
- batch = []
- if batch:
- yield batch
-
-
-def split_paths(paths: list[Path], split: str) -> list[Path]:
- if not paths:
- return []
- return paths[:-1] if split == "train" else paths[-1:]
diff --git a/vendor/nanoai/agentic/data/tool_chain_demo.py b/vendor/nanoai/agentic/data/tool_chain_demo.py
deleted file mode 100644
index f556b4e2f1f5a8b2707ba9d4d026d78b083d8f7f..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/tool_chain_demo.py
+++ /dev/null
@@ -1,563 +0,0 @@
-"""Synthetic multi-step tool-chain demonstration dataset.
-
-Motivation (POC-9 mechanistic analysis):
- v11 (routing_rows=200) lifted single-tool routing accuracy but did NOT
- recover bash_grep_log — the diagnostic regression task that requires
- CHAINING bash → grep → read. routing_recovery teaches "given prompt P,
- call tool T". It cannot teach "given prompt P, call T1, then read its
- output, then call T2 with args derived from T1's output, then ...".
-
- This dataset adds explicit multi-step tool-chain trajectories. Each
- row is a 6-10 message conversation showing: user prompt → tool call →
- tool result → tool call → tool result → ... → final answer. The
- intermediate assistant turns reference what the previous tool result
- said, modeling the read-act-observe loop.
-
-Format on disk: JSONL, one `{"messages": [...]}` per line, compatible
-with `render_agentic_conversation`'s tool-call branch.
-
-Schema invariants (verified by audit):
- - first message role=user, last role=assistant (no tool_call)
- - intermediate assistants alternate with tool_results
- - every assistant with tool_call has args as a dict
- - row text deterministic for fixed (seed, chain_kind, idx)
-
-CLI:
- python -m nanoai.agentic.data.tool_chain_demo \
- --output data/tool_chain_demo.jsonl --seed 42 --n-per-chain 100
-"""
-
-from __future__ import annotations
-
-import argparse
-import json
-import random
-from dataclasses import dataclass, field
-from pathlib import Path
-from typing import Any
-
-from nanoai.agentic.data.common import SYSTEM_PROMPT
-
-
-# ---- Corpora ---------------------------------------------------------------
-
-_FILES_PY = [
- "main.py", "app.py", "lib.py", "utils.py", "config.py", "server.py",
- "client.py", "models.py", "schema.py", "data.py", "tools.py",
- "parser.py", "router.py", "loader.py", "engine.py", "core.py",
- "common.py", "tasks.py", "errors.py", "handlers.py", "auth.py",
-]
-_FILES_DATA = [
- "config.toml", "settings.json", "data.csv", "users.json", "log.txt",
- "README.md", "Makefile", "requirements.txt", "pyproject.toml",
- "errors.log", "access.log", "metrics.tsv", "manifest.yaml",
- "deploy.yaml", "schema.sql",
-]
-_DIRS = ["src", "lib", "tests", "scripts", "data", "config", "core", "app"]
-_PATTERNS = [
- "TODO", "FIXME", "ERROR", "WARN", "import torch",
- "def __init__", "raise ValueError", "logger.info",
- "self.config", "from typing import", "@dataclass",
-]
-
-
-# ---- Step + Row dataclasses ------------------------------------------------
-
-
-@dataclass
-class _Step:
- """One assistant→tool round-trip inside a chain."""
- pre_text: str
- tool_name: str
- tool_args: dict[str, str]
- tool_result: str
-
-
-@dataclass
-class _Row:
- user_prompt: str
- steps: list[_Step]
- final_text: str
- chain_kind: str = ""
- extra: dict[str, Any] = field(default_factory=dict)
-
-
-# ---- Chain generators ------------------------------------------------------
-#
-# Each generator returns one _Row. They share a per-row rng so determinism
-# holds: identical (seed, chain_kind, idx) → identical bytes.
-
-
-def _bash_grep_read_row(rng: random.Random) -> _Row:
- """bash → grep → read. Direct attack on bash_grep_log task family."""
- d = rng.choice(_DIRS)
- pat = rng.choice(_PATTERNS)
- target = rng.choice(_FILES_PY)
- n_files = rng.randint(2, 5)
- listing_files = rng.sample(_FILES_PY, n_files)
- if target not in listing_files:
- listing_files[0] = target
- listing = "\n".join(f"{d}/{f}" for f in listing_files)
- grep_hits = "\n".join(
- f"{d}/{target}:{rng.randint(5, 200)}: {pat}"
- for _ in range(rng.randint(1, 3))
- )
- read_body = (
- f"# {target}\n"
- f"def handler():\n"
- f" # {pat} fix the broken case\n"
- f" raise ValueError('boom')\n"
- )
- user = rng.choice([
- f"in {d}/, find python files containing {pat} and show me the relevant file",
- f"which file under {d} mentions {pat}? then read it and explain",
- f"locate {pat} in {d} and tell me what the surrounding code does",
- f"search for {pat} in {d} python files and read the offending one",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text=f"let me list the files under {d} first",
- tool_name="Bash",
- tool_args={"command": f"ls {d}/*.py"},
- tool_result=listing,
- ),
- _Step(
- pre_text=f"now grep for {pat} across those files",
- tool_name="Grep",
- tool_args={"pattern": pat, "path": d},
- tool_result=grep_hits,
- ),
- _Step(
- pre_text=f"reading {d}/{target} where the hit is",
- tool_name="Read",
- tool_args={"file_path": f"{d}/{target}"},
- tool_result=read_body,
- ),
- ],
- final_text=(
- f"{target} has a {pat} marker on a `handler()` that raises ValueError. "
- f"the relevant block is the broken-case path."
- ),
- chain_kind="bash_grep_read",
- )
-
-
-def _bash_read_row(rng: random.Random) -> _Row:
- """bash → read. e.g. find largest file then read it."""
- d = rng.choice(_DIRS)
- target = rng.choice(_FILES_PY)
- sizes = sorted(
- [(rng.randint(1, 80), rng.choice(_FILES_PY)) for _ in range(4)],
- reverse=True,
- )
- if sizes[0][1] == target:
- winner = sizes[0]
- else:
- winner = (sizes[0][0], target)
- listing = "\n".join(f"{sz}K\t{d}/{name}" for sz, name in sizes[:3] + [winner])
- read_body = (
- f"# {target}\n"
- f"class Engine:\n"
- f" def run(self):\n"
- f" return 'ok'\n"
- )
- user = rng.choice([
- f"find the largest python file under {d}/ and show me what it does",
- f"which file in {d} is biggest? read it and summarize",
- f"check disk usage of {d}/*.py, then open the largest one",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text=f"checking sizes under {d}",
- tool_name="Bash",
- tool_args={"command": f"du -k {d}/*.py | sort -rn | head -5"},
- tool_result=listing,
- ),
- _Step(
- pre_text=f"the largest is {target}, reading it",
- tool_name="Read",
- tool_args={"file_path": f"{d}/{target}"},
- tool_result=read_body,
- ),
- ],
- final_text=(
- f"{target} is the largest under {d} ({winner[0]}K). it defines an "
- f"`Engine` class with a `run()` method that returns 'ok'."
- ),
- chain_kind="bash_read",
- )
-
-
-def _grep_read_row(rng: random.Random) -> _Row:
- """grep → read. Find a hit, then open the file to inspect context."""
- pat = rng.choice(_PATTERNS)
- target = rng.choice(_FILES_PY)
- where = rng.choice(_DIRS)
- line = rng.randint(10, 200)
- hits = (
- f"{where}/{target}:{line}: {pat}\n"
- f"{where}/other.py:{rng.randint(5, 50)}: {pat}"
- )
- read_body = (
- f"# {target} excerpt around line {line}\n"
- f"def call_site():\n"
- f" # {pat}: under review\n"
- f" return None\n"
- )
- user = rng.choice([
- f"grep for {pat} under {where}/ and show me the first hit's surrounding code",
- f"where does {pat} appear in {where}? open the file and explain",
- f"find {pat} in {where} python files and read the matching one",
- f"search {pat} in {where}/ then dump the file context",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text="grepping",
- tool_name="Grep",
- tool_args={"pattern": pat, "path": where},
- tool_result=hits,
- ),
- _Step(
- pre_text=f"top hit is {where}/{target}:{line}, reading it",
- tool_name="Read",
- tool_args={"file_path": f"{where}/{target}"},
- tool_result=read_body,
- ),
- ],
- final_text=(
- f"{target} has the {pat} marker on a `call_site()` function "
- f"around line {line}; it's an inert TODO-style note."
- ),
- chain_kind="grep_read",
- )
-
-
-def _read_edit_row(rng: random.Random) -> _Row:
- """read → edit. Open a file, make a targeted change."""
- target = rng.choice(_FILES_PY + ["config.py", "config.toml"])
- edit = rng.choice([
- ("bump_version",
- f"version = '0.1.0'", "version = '0.2.0'",
- "bump the version to 0.2.0"),
- ("rename_var",
- "tmp = compute()", "result = compute()",
- "rename `tmp` to `result`"),
- ("add_docstring",
- "def main():", "def main():\n \"\"\"entry point\"\"\"",
- "add a one-line docstring to main()"),
- ("fix_typo",
- "def recieve(x):", "def receive(x):",
- "fix the typo recieve → receive"),
- ("replace_print",
- "print(x)", "logger.info(x)",
- "replace the print() with logger.info()"),
- ])
- kind, old, new, ask = edit
- read_body = f"# {target}\n{old}\n# (other lines)\n"
- user = rng.choice([
- f"in {target}, {ask}",
- f"open {target} and {ask}",
- f"please {ask} in {target}",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text=f"reading {target} first",
- tool_name="Read",
- tool_args={"file_path": target},
- tool_result=read_body,
- ),
- _Step(
- pre_text=f"applying the edit ({kind})",
- tool_name="Edit",
- tool_args={
- "file_path": target,
- "old_string": old,
- "new_string": new,
- },
- tool_result=f"applied edit to {target}: replaced 1 occurrence.",
- ),
- ],
- final_text=f"done — {target} now has the {kind.replace('_', ' ')} applied.",
- chain_kind="read_edit",
- )
-
-
-def _read_bash_row(rng: random.Random) -> _Row:
- """read → bash. Read a script to find the right invocation, then run."""
- target = rng.choice(["run.sh", "build.sh", "test.sh", "deploy.sh"])
- flag = rng.choice(["--verbose", "--dry-run", "--prod", "--config=staging"])
- read_body = (
- f"#!/usr/bin/env bash\n"
- f"# {target}\n"
- f"# Usage: {target} {flag}\n"
- f"echo running with {flag}\n"
- )
- bash_out = f"running with {flag}\n"
- user = rng.choice([
- f"check what flag {target} expects, then run it",
- f"read {target} and execute it with the right argument",
- f"open {target} to find the usage line, then run it accordingly",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text=f"reading {target} to find the flag",
- tool_name="Read",
- tool_args={"file_path": target},
- tool_result=read_body,
- ),
- _Step(
- pre_text=f"running it with {flag}",
- tool_name="Bash",
- tool_args={"command": f"bash {target} {flag}"},
- tool_result=bash_out,
- ),
- ],
- final_text=f"ran {target} with {flag} — output confirms the run completed.",
- chain_kind="read_bash",
- )
-
-
-def _read_grep_read_row(rng: random.Random) -> _Row:
- """read → grep → read. Open a manifest, grep for a referenced symbol,
- open the file containing the symbol."""
- manifest = rng.choice(["README.md", "manifest.yaml", "pyproject.toml"])
- entry = rng.choice(_FILES_PY)
- symbol = rng.choice(["main", "run", "handler", "boot", "serve"])
- target = rng.choice(_FILES_PY)
- line = rng.randint(10, 100)
- manifest_body = (
- f"# {manifest}\n"
- f"entry: {entry}\n"
- f"main_symbol: {symbol}\n"
- )
- grep_hits = f"{target}:{line}: def {symbol}():"
- read_body = (
- f"# {target}\n"
- f"def {symbol}():\n"
- f" return 'started'\n"
- )
- user = rng.choice([
- f"open {manifest}, find the main symbol, then locate its definition",
- f"read {manifest} to find the entry symbol, then grep for and read it",
- f"check {manifest} for the main_symbol, then track it down in source",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text=f"reading {manifest}",
- tool_name="Read",
- tool_args={"file_path": manifest},
- tool_result=manifest_body,
- ),
- _Step(
- pre_text=f"grepping for `def {symbol}`",
- tool_name="Grep",
- tool_args={"pattern": f"def {symbol}", "path": "."},
- tool_result=grep_hits,
- ),
- _Step(
- pre_text=f"reading {target} where {symbol} is defined",
- tool_name="Read",
- tool_args={"file_path": target},
- tool_result=read_body,
- ),
- ],
- final_text=(
- f"`{symbol}` is defined in {target} (line {line}); it's a one-liner "
- f"that returns 'started'."
- ),
- chain_kind="read_grep_read",
- )
-
-
-def _bash_read_edit_row(rng: random.Random) -> _Row:
- """bash → read → edit. find a config file, open it, edit it."""
- name = rng.choice(["config.toml", "settings.json", "deploy.yaml"])
- where = rng.choice(_DIRS)
- found_path = f"./{where}/{name}"
- listing = f"{found_path}\n"
- if name.endswith(".toml"):
- old = "log_level = 'INFO'"
- new = "log_level = 'DEBUG'"
- ask = "set log_level to DEBUG"
- elif name.endswith(".json"):
- old = '"timeout": 30'
- new = '"timeout": 60'
- ask = "raise the timeout to 60"
- else:
- old = "replicas: 1"
- new = "replicas: 3"
- ask = "scale replicas to 3"
- read_body = f"# {name}\n{old}\n# (other config)\n"
- user = rng.choice([
- f"find {name} in the repo, read it, and {ask}",
- f"locate {name}, open it, then {ask}",
- f"where is {name}? open it and {ask}",
- ])
- return _Row(
- user_prompt=user,
- steps=[
- _Step(
- pre_text=f"finding {name}",
- tool_name="Bash",
- tool_args={"command": f"find . -name {name}"},
- tool_result=listing,
- ),
- _Step(
- pre_text=f"reading {found_path}",
- tool_name="Read",
- tool_args={"file_path": found_path},
- tool_result=read_body,
- ),
- _Step(
- pre_text="applying the edit",
- tool_name="Edit",
- tool_args={
- "file_path": found_path,
- "old_string": old,
- "new_string": new,
- },
- tool_result=f"applied edit to {found_path}.",
- ),
- ],
- final_text=f"updated {found_path}: {ask}.",
- chain_kind="bash_read_edit",
- )
-
-
-# ---- Registry --------------------------------------------------------------
-
-
-_GENERATORS: dict[str, tuple[Any, int]] = {
- # name: (generator_fn, default_n_rows)
- "bash_grep_read": (_bash_grep_read_row, 150),
- "bash_read": (_bash_read_row, 100),
- "grep_read": (_grep_read_row, 100),
- "read_edit": (_read_edit_row, 100),
- "read_bash": (_read_bash_row, 80),
- "read_grep_read": (_read_grep_read_row, 80),
- "bash_read_edit": (_bash_read_edit_row, 80),
-}
-
-
-# ---- Row builder -----------------------------------------------------------
-
-
-def _row_to_messages(r: _Row) -> dict[str, Any]:
- msgs: list[dict[str, Any]] = [{"role": "user", "content": r.user_prompt}]
- for step in r.steps:
- msgs.append({
- "role": "assistant",
- "content": step.pre_text,
- "tool_call": {"name": step.tool_name, "args": step.tool_args},
- })
- msgs.append({"role": "tool_result", "content": step.tool_result})
- msgs.append({"role": "assistant", "content": r.final_text})
- return {"messages": msgs, "_chain_kind": r.chain_kind}
-
-
-def generate_tool_chain_rows(
- seed: int = 42,
- n_per_chain: int = 100,
- chains: tuple[str, ...] | None = None,
-) -> list[dict[str, Any]]:
- """Generate the tool-chain demo rows in memory.
-
- Deterministic per `seed`. Per-chain row counts:
- n_per_chain == 100 (default): use chain-specific defaults from
- _GENERATORS (sums to ~690).
- otherwise: use n_per_chain rows for EVERY chain (uniform).
- Pass `chains` to subset.
- """
- chain_keys = chains if chains is not None else tuple(_GENERATORS.keys())
- rng = random.Random(seed)
- out: list[dict[str, Any]] = []
- for chain_idx, chain in enumerate(chain_keys):
- gen, default_n = _GENERATORS[chain]
- n = default_n if n_per_chain == 100 else n_per_chain
- for i in range(n):
- row_rng = random.Random(seed * 7919 + chain_idx * 1_000_003 + i)
- row = gen(row_rng)
- out.append(_row_to_messages(row))
- rng.shuffle(out)
- return out
-
-
-# ---- SFT-side dataset class ------------------------------------------------
-
-
-class ToolChainDemoDataset:
- """In-memory dataset usable by `agentic_sft.py` task mixtures.
-
- Mirrors the JSONDataset / RoutingRecoveryDataset protocol.
- Drops the `_chain_kind` debug key before yielding messages.
- """
-
- def __init__(
- self,
- seed: int = 42,
- n_per_chain: int = 100,
- chains: tuple[str, ...] | None = None,
- **_: Any,
- ) -> None:
- self.ds = generate_tool_chain_rows(
- seed=seed, n_per_chain=n_per_chain, chains=chains,
- )
-
- def __len__(self) -> int:
- return len(self.ds)
-
- def __getitem__(self, idx: int) -> dict[str, Any]:
- messages = self.ds[idx]["messages"]
- return {"messages": [{"role": "system", "content": SYSTEM_PROMPT}] + messages}
-
-
-# ---- CLI -------------------------------------------------------------------
-
-
-def _main() -> None:
- p = argparse.ArgumentParser()
- p.add_argument("--output", required=True, type=Path,
- help="Output JSONL path (one {messages:[...]} per line)")
- p.add_argument("--seed", type=int, default=42)
- p.add_argument("--n-per-chain", type=int, default=100,
- help="If 100: use chain-specific default counts (~690 total). "
- "Otherwise: use this count for every chain.")
- p.add_argument("--chains", type=str, default="",
- help="Comma-separated chain names; empty = all "
- "(bash_grep_read,bash_read,grep_read,read_edit,"
- "read_bash,read_grep_read,bash_read_edit)")
- args = p.parse_args()
-
- chains: tuple[str, ...] | None
- if args.chains:
- chains = tuple(c.strip() for c in args.chains.split(",") if c.strip())
- else:
- chains = None
- rows = generate_tool_chain_rows(
- seed=args.seed, n_per_chain=args.n_per_chain, chains=chains,
- )
- args.output.parent.mkdir(parents=True, exist_ok=True)
- with args.output.open("w", encoding="utf-8") as f:
- for row in rows:
- f.write(json.dumps(row, ensure_ascii=False) + "\n")
- summary = ", ".join(
- f"{c}={sum(1 for r in rows if r.get('_chain_kind') == c)}"
- for c in (chains or _GENERATORS.keys())
- )
- print(f"wrote {len(rows)} rows ({summary}) → {args.output}")
-
-
-if __name__ == "__main__":
- _main()
diff --git a/vendor/nanoai/agentic/data/vrcpt_data.py b/vendor/nanoai/agentic/data/vrcpt_data.py
deleted file mode 100644
index a0d7e3b58a171552549d00a872a277628fad57bb..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/data/vrcpt_data.py
+++ /dev/null
@@ -1,195 +0,0 @@
-"""VR-CPT pretrain stream (verifier-mined case docs).
-
-Materializes a domain-aligned stream of REF.md §9.1 case documents produced by
-``experiments/vrmining/miners/case_writer.py``. Source JSONL is normally at::
-
- experiments/vrmining/datasets/cpt_corpus.jsonl
-
-Each row has a ``text`` field containing a markdown case (~500-1500 token).
-1k mining → 1000 cases ≈ 1M tokens. Used as a domain mixin alongside web/code
-in `agentic_pretrain.py --vrcpt-ratio` (see [[rec_4axis_pretrain]] Recipe A).
-
-Mirrors `nanoai.agentic.data.pcap_data` exactly — converts JSONL to parquet
-shards stored under ``$NANOAI_BASE_DIR/code_pretrain_data/vrcpt-cases/`` so
-the existing parquet-iter API drives it identically to web / code / pcap-corpus.
-
-Run once before pretraining:
-
- python -m nanoai.agentic.data.vrcpt_data prepare \\
- --src experiments/vrmining/datasets/cpt_corpus.jsonl \\
- --rows-per-shard 500
-"""
-
-from __future__ import annotations
-
-import argparse
-import json
-import time
-from pathlib import Path
-
-import pyarrow as pa
-import pyarrow.parquet as pq
-
-from nanoai.common import print0
-from nanoai.agentic.data.pretrain_data import DATA_DIR
-
-DATASET_NAME = "vrcpt-cases"
-_TMP_SUFFIX = ".parquet.tmp"
-
-
-def vrcpt_corpus_dir() -> Path:
- return DATA_DIR / DATASET_NAME
-
-
-def _stream_jsonl_text(src: Path):
- with src.open("r", encoding="utf-8") as f:
- for line in f:
- line = line.strip()
- if not line:
- continue
- obj = json.loads(line)
- text = obj.get("text")
- if isinstance(text, str) and text:
- yield text
-
-
-def _try_open_shard(path: Path) -> bool:
- try:
- pf = pq.ParquetFile(path)
- return pf.num_row_groups > 0
- except Exception:
- return False
-
-
-def _validate_existing_corpus(out_dir: Path) -> tuple[bool, list[Path], str]:
- parquets = sorted(p for p in out_dir.iterdir() if p.suffix == ".parquet")
- tmps = [p for p in out_dir.iterdir() if p.name.endswith(_TMP_SUFFIX)]
- if tmps:
- return False, parquets, f"{len(tmps)} stale {_TMP_SUFFIX} file(s)"
- if len(parquets) < 2:
- return False, parquets, f"only {len(parquets)} shard(s); need ≥2"
- for p in parquets:
- if not _try_open_shard(p):
- return False, parquets, f"shard {p.name} unreadable"
- return True, parquets, "ok"
-
-
-def _quarantine_stale(out_dir: Path, reason_slug: str) -> Path | None:
- movable = [
- p for p in out_dir.iterdir()
- if p.suffix == ".parquet" or p.name.endswith(_TMP_SUFFIX)
- ]
- if not movable:
- return None
- attic = out_dir / f".attic_{int(time.time())}_{reason_slug}"
- attic.mkdir()
- for entry in movable:
- entry.rename(attic / entry.name)
- print0(
- f"[vrcpt_data] quarantined {len(movable)} stale file(s) to {attic.name}/"
- )
- return attic
-
-
-def _reason_slug(reason: str) -> str:
- return reason.split(";")[0].strip().replace(" ", "_")[:32] or "invalid"
-
-
-def prepare_from_jsonl(src: Path, rows_per_shard: int = 500):
- """One-shot JSONL → parquet shards under vrcpt_corpus_dir().
-
- Restart-safe and atomic-write (write .tmp then rename), same protocol as
- pcap_data. Last shard becomes val split (mirrors fineweb-edu convention).
- """
- out_dir = vrcpt_corpus_dir()
- out_dir.mkdir(parents=True, exist_ok=True)
-
- is_valid, existing, reason = _validate_existing_corpus(out_dir)
- if is_valid:
- print0(f"[vrcpt_data] {len(existing)} valid shards already in {out_dir}; skipping.")
- return existing
- if existing or any(p.name.endswith(_TMP_SUFFIX) for p in out_dir.iterdir()):
- print0(f"[vrcpt_data] existing corpus at {out_dir} INVALID ({reason}); rebuilding.")
- _quarantine_stale(out_dir, _reason_slug(reason))
-
- assert src.exists(), f"Source JSONL not found: {src}"
- print0(f"[vrcpt_data] reading {src} → {out_dir} ({rows_per_shard} rows/shard)")
-
- shard_idx = 0
- buf: list[str] = []
- written: list[Path] = []
- total_docs = 0
- total_chars = 0
-
- def flush():
- nonlocal shard_idx, buf
- if not buf:
- return
- shard_path = out_dir / f"shard_{shard_idx:05d}.parquet"
- tmp_path = shard_path.with_name(shard_path.name + ".tmp")
- shard_chars = sum(len(t) for t in buf)
- table = pa.table({"text": buf})
- pq.write_table(table, tmp_path, compression="snappy")
- tmp_path.rename(shard_path)
- written.append(shard_path)
- size_mb = shard_path.stat().st_size / 1e6
- print0(
- f"[vrcpt_data] wrote {shard_path.name} "
- f"({len(buf)} docs, {shard_chars/1e6:.2f}M chars, {size_mb:.2f} MB)"
- )
- shard_idx += 1
- buf = []
-
- for text in _stream_jsonl_text(src):
- buf.append(text)
- total_docs += 1
- total_chars += len(text)
- if len(buf) >= rows_per_shard:
- flush()
- flush()
-
- assert len(written) >= 2, (
- f"Only {len(written)} shard(s); need ≥2 (train/val). "
- f"Lower --rows-per-shard or supply a larger JSONL."
- )
- est_tokens = total_chars / 4
- print0(
- f"[vrcpt_data] done: {len(written)} shards (last is val), "
- f"{total_docs:,} docs, {total_chars/1e6:.2f}M chars, "
- f"≈{est_tokens/1e6:.2f}M tokens"
- )
- return written
-
-
-def list_vrcpt_parquet_files() -> list[Path]:
- d = vrcpt_corpus_dir()
- if not d.exists():
- return []
- return sorted([d / f for f in d.iterdir() if f.suffix == ".parquet"])
-
-
-def main():
- parser = argparse.ArgumentParser(description="Prepare vrcpt-cases parquet shards.")
- sub = parser.add_subparsers(dest="cmd", required=True)
- prep = sub.add_parser("prepare", help="convert cpt_corpus.jsonl to parquet shards")
- prep.add_argument("--src", type=Path, required=True,
- help="path to cpt_corpus.jsonl (case_writer output)")
- prep.add_argument("--rows-per-shard", type=int, default=500)
- sub.add_parser("list", help="list current vrcpt-cases shards")
-
- args = parser.parse_args()
- if args.cmd == "prepare":
- prepare_from_jsonl(args.src, rows_per_shard=args.rows_per_shard)
- elif args.cmd == "list":
- shards = list_vrcpt_parquet_files()
- if not shards:
- print(f"No shards in {vrcpt_corpus_dir()}")
- return
- total_bytes = sum(s.stat().st_size for s in shards)
- for s in shards:
- print(f" {s.name} ({s.stat().st_size / 1e6:.2f} MB)")
- print(f"total: {len(shards)} shards, {total_bytes / 1e6:.1f} MB")
-
-
-if __name__ == "__main__":
- main()
diff --git a/vendor/nanoai/agentic/dense_verifier.py b/vendor/nanoai/agentic/dense_verifier.py
deleted file mode 100644
index 10c76a79e9ba2e511ce8797346bb1dcb4793ea85..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/dense_verifier.py
+++ /dev/null
@@ -1,247 +0,0 @@
-"""Dense per-turn verifier reward + format penalty.
-
-Designed to break the GFT-TW mode-collapse trap:
- - Provides a non-zero signal per turn even when the group is at 100% pass,
- so training doesn't stall on saturated reward.
- - Penalises the specific degenerate outputs observed across iter-MCTS-DPO,
- GFT-TW, and AC v1: `assert assert assert ...`, mangled special tokens
- like `tool_call_end_end_end.log`, hallucinated tool_result content.
- - Independent of `task.check` — scores from the tool_result text and the
- decoded assistant content, so it works even on tasks the terminal grader
- is too lenient for.
-
-Wire-in points:
- - `score_turn(turn_dict)` runs after rollout, reads `tool_call`,
- `tool_result`, `error` fields from the Transcript turn dict.
- - `format_penalty(decoded_text)` runs over the assistant's text body.
- - `per_turn_dense_reward(rollout)` returns a list[float] aligned with
- `rollout.turn_spans`, ready to feed GAE.
-"""
-from __future__ import annotations
-
-import re
-from dataclasses import dataclass
-
-
-# Reward magnitudes (tuned so format penalty > dense > shaping > terminal/N).
-EDIT_OK = 0.4
-EDIT_FAIL = -0.5
-READ_OK = 0.1
-READ_FAIL = -0.5
-BASH_OK = 0.1
-BASH_FAIL = -0.2
-GREP_HIT = 0.2
-GREP_MISS = -0.1
-PARSE_ERROR = -0.5
-
-# Format penalty magnitudes — large so they dominate any same-turn shaping.
-NGRAM_REPEAT_PENALTY = -1.0
-MANGLED_SPECIAL_PENALTY = -0.8
-TRUNCATED_TOOL_CALL_PENALTY = -0.6
-
-# Terminal grounding (kept small — dense reward carries most weight).
-TERMINAL_PASS = 0.5
-TERMINAL_FAIL = -0.3
-
-
-def score_turn(turn: dict) -> float:
- """Score a single assistant+tool_result turn pair from the Transcript.
-
- turn: a dict from `Transcript.turns` of the assistant role. The companion
- tool_result is read from `turn["_tool_result"]` (caller fills it in by
- pairing assistant + tool_result turns). If no tool_call is present (final
- answer turn), returns 0.0 — the terminal grounding handles those.
-
- Returns a float in roughly [-1, +1].
- """
- tc = turn.get("tool_call")
- if not tc:
- return 0.0
- name = tc.get("name", "")
- result_text = turn.get("_tool_result", "") or ""
- is_error = turn.get("_error", False)
-
- if name == "":
- return PARSE_ERROR
-
- lo = result_text.lower()
-
- if name == "Edit":
- # Distinguish a real edit from "old_string not found" / silent no-op.
- if is_error:
- return EDIT_FAIL
- if "successfully" in lo or "replaced" in lo or "wrote" in lo:
- return EDIT_OK
- if "not found" in lo or "no match" in lo:
- return EDIT_FAIL
- # Edit returned non-empty without an obvious success marker — small +.
- return EDIT_OK * 0.3
-
- if name == "Read":
- if is_error:
- return READ_FAIL
- if "file not found" in lo or "no such file" in lo or lo.startswith("error"):
- return READ_FAIL
- if len(result_text) > 20:
- return READ_OK
- # Empty / very short read result — neither rewarded nor punished much.
- return 0.0
-
- if name == "Bash":
- if is_error:
- return BASH_FAIL
- if "traceback" in lo or "modulenotfounderror" in lo or "no such file" in lo:
- return BASH_FAIL
- if "exit code 1" in lo or "command not found" in lo:
- return BASH_FAIL
- return BASH_OK
-
- if name == "Grep":
- if is_error:
- return GREP_MISS
- if not result_text.strip() or "no matches" in lo:
- return GREP_MISS
- return GREP_HIT
-
- # Unknown tool — neutral.
- return 0.0
-
-
-# Special-token leak patterns observed during mode collapse.
-_MANGLED_SPECIAL_PATTERNS = [
- re.compile(r"tool_call_end_end"),
- re.compile(r"tool_call_start.*tool_call_start.*tool_call_start"),
- re.compile(r"(<\|tool_[a-z_]+\|>){3,}"),
- re.compile(r"end_end_end"),
-]
-_TOKEN_LIKE = re.compile(r"<\|[a-z_]+\|>")
-
-
-def format_penalty(text: str, window: int = 12) -> float:
- """Penalise degenerate output patterns observed in case-study.
-
- text: decoded assistant body for a single turn (string).
- window: how many tail tokens (words) to inspect for n-gram repeat.
-
- Returns a float in [-1.0, 0.0]. Multiple penalties are summed (clipped at -1.0).
- """
- if not text:
- return 0.0
- penalty = 0.0
-
- words = text.split()
- # 1) n-gram repetition: last `window` words all identical, OR
- # any single non-trivial word appears ≥3× in the window, OR
- # the most common 3-gram appears ≥ 3 times.
- if len(words) >= 6:
- tail = words[-window:] if len(words) >= window else words
- if len(set(tail[-5:])) == 1:
- penalty += NGRAM_REPEAT_PENALTY
- else:
- # Unigram repeat — catches `assertionassertio × 3` pattern.
- non_trivial = [w for w in tail if len(w) > 4]
- if non_trivial:
- max_unigram_count = max(non_trivial.count(w) for w in set(non_trivial))
- if max_unigram_count >= 3:
- penalty += NGRAM_REPEAT_PENALTY * 0.7
- # 3-gram repeat density check (catches `assert x; assert x; assert x`).
- ngrams = [tuple(tail[i:i + 3]) for i in range(len(tail) - 2)]
- if ngrams:
- most_common_count = max(ngrams.count(g) for g in set(ngrams))
- if most_common_count >= 3:
- penalty += NGRAM_REPEAT_PENALTY * 0.7
-
- # 2) Mangled / leaked special tokens in the visible text.
- for pat in _MANGLED_SPECIAL_PATTERNS:
- if pat.search(text):
- penalty += MANGLED_SPECIAL_PENALTY
- break
-
- # 3) Special-token leak: the SAME special token repeated ≥3 times,
- # e.g. `<|tool_call_start|><|tool_call_start|><|tool_call_start|>...`
- # (well-formed multi-arg tool call uses each tag at most twice).
- found = _TOKEN_LIKE.findall(text)
- if found:
- counts = {}
- for s in found:
- counts[s] = counts.get(s, 0) + 1
- if max(counts.values()) >= 3:
- penalty += TRUNCATED_TOOL_CALL_PENALTY
-
- # 4) Unclosed tool_call: has `<|tool_call_start|>` but no matching
- # `<|tool_call_end|>` (covers cases like
- # `<|tool_call_start|>Edit<|tool_arg|>...<|tool_val|>a` that escape
- # n-gram + special-token leak checks due to short total length).
- if "<|tool_call_start|>" in text and "<|tool_call_end|>" not in text:
- penalty += TRUNCATED_TOOL_CALL_PENALTY
-
- # 5) Unclosed tool_val: `<|tool_val|>` opens but text ends without further
- # structure (catches the edit_create_test PPO v1 regression where the
- # model emits `<|tool_val|>a` then stops with len < 10).
- if "<|tool_val|>" in text:
- after_val = text.rsplit("<|tool_val|>", 1)[-1]
- if len(after_val.strip()) < 5 and "<|tool_call_end|>" not in text:
- penalty += TRUNCATED_TOOL_CALL_PENALTY
-
- # 6) Brace explosion inside tool_val: model entered "write JSON" mode and
- # couldn't escape, producing unmatched runs of `{` or `}`. Catches the
- # edit_create_json failure mode where `<|tool_val|>{...{{{{` never reaches
- # `<|tool_call_end|>` because nested-brace prior overrides special-token escape.
- if "<|tool_val|>" in text and "<|tool_call_end|>" not in text:
- # Look at content after the LAST tool_val opener (the new_string body).
- body_after_val = text.rsplit("<|tool_val|>", 1)[-1]
- n_open = body_after_val.count("{")
- n_close = body_after_val.count("}")
- if n_open + n_close >= 8 and abs(n_open - n_close) >= 3:
- penalty += TRUNCATED_TOOL_CALL_PENALTY
- # Long run of identical brace chars (`}}}}}}}}` etc).
- if "}}}}}}" in body_after_val or "{{{{{{" in body_after_val:
- penalty += TRUNCATED_TOOL_CALL_PENALTY
-
- return max(penalty, -1.0)
-
-
-def per_turn_dense_reward(rollout) -> list[float]:
- """Compute dense per-turn reward aligned with rollout.turn_spans.
-
- Pairs each assistant turn with its tool_result (next entry in transcript)
- and computes score_turn + format_penalty for the assistant body.
- Terminal pass/fail grounding is appended to the LAST turn only.
- """
- transcript = rollout.transcript
- turns = transcript.turns
- # Build a list of (assistant_turn, paired_tool_result) for each TurnSpan.
- paired: list[tuple[dict, dict | None]] = []
- i = 0
- while i < len(turns):
- t = turns[i]
- if t.get("role") == "assistant":
- tr = None
- if i + 1 < len(turns) and turns[i + 1].get("role") == "tool_result":
- tr = turns[i + 1]
- i += 2
- else:
- i += 1
- paired.append((t, tr))
- else:
- i += 1
-
- rewards: list[float] = []
- for ts, (asst, tr) in zip(rollout.turn_spans, paired):
- # Fill in tool_result / error fields into the assistant dict for score_turn.
- a_dict = dict(asst)
- if tr is not None:
- a_dict["_tool_result"] = tr.get("tool_result", "")
- a_dict["_error"] = bool(tr.get("error", False))
- r_tool = score_turn(a_dict)
- r_fmt = format_penalty(asst.get("content", "") or "")
- rewards.append(r_tool + r_fmt)
-
- if rewards:
- # Terminal grounding on the last turn only.
- if getattr(rollout, "passed", False):
- rewards[-1] += TERMINAL_PASS
- else:
- rewards[-1] += TERMINAL_FAIL
-
- return rewards
diff --git a/vendor/nanoai/agentic/diag/__init__.py b/vendor/nanoai/agentic/diag/__init__.py
deleted file mode 100644
index 592df0d9d6331098a8395961e8aa36af940c74dc..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/__init__.py
+++ /dev/null
@@ -1,11 +0,0 @@
-"""Agentic diagnosis micro-task eval package.
-
-This package defines the 8-task agentic diagnosis loop eval framework — the
-single source of truth for both ground-truth generation (capagent driver) and
-post-hoc grading (ksbr_eval extension).
-
-See plan: agentic diagnosis loop micro-task eval v1 (PCAP domain).
-"""
-
-from .schemas import DIAG_TASKS, TaskSpec, FieldSpec # noqa: F401
-from .graders import GRADERS, grade_against_spec # noqa: F401
diff --git a/vendor/nanoai/agentic/diag/chain_adapters.py b/vendor/nanoai/agentic/diag/chain_adapters.py
deleted file mode 100644
index 1f072547e5e80edc664e1a7b8bb2d91d8acc66bc..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/chain_adapters.py
+++ /dev/null
@@ -1,185 +0,0 @@
-"""chain↔SFT schema adapters for the 8-task diagnosis loop.
-
-Why: per-task SFT was trained on shapes that don't always match what the 8-task
-chain runtime feeds. Without adaptation, chain inputs are OOD for some fields,
-which hurts per-task parse_ok rate and triggers task-identity collapse on long
-traces (see `evidence_summary_for_stop_controller` story).
-
-What's here:
- - hyp shape adapters (per downstream task) — v11 emits rich hyp dict; legacy
- SFT trained on simpler shapes. (Moved here from demo + webapp duplicates.)
- - evidence_summary_for_stop_controller — list[dict] → narrative string.
- - 4 new input-shape adapters resolving the high-severity mismatches found
- in 2026-05-22 schema audit:
- * packet_summary_to_str — wmb input (89% SFT trained on str)
- * time_window_to_dict — sd input (96% SFT trained on dict)
- * available_tools_to_dicts — tp input (88% SFT trained on list[dict])
- * list_or_dict_to_str — hu prediction/observation (84%/56% str)
-
-All adapters are idempotent — if input is already the SFT shape they pass it
-through unchanged."""
-from __future__ import annotations
-
-import json
-from typing import Any
-
-
-# ---- hyp shape adapters (v11 rich hyp → legacy per-task shapes) ----------------
-
-STATUS_TO_CONF: dict[str, float] = {
- "confirmed": 0.95, "provisionally_supported": 0.7,
- "open": 0.5, "unchanged": 0.5, "inconclusive": 0.5,
- "weakened": 0.3, "rejected": 0.4,
-}
-# 2026-05-26: rejected mapped 0.05→0.4 to fix universal-hu-rejection bug. v13
-# SFT trained hu to reject H1 in ~100% of chain runs even when H1 is truth.
-# Old 0.05 caused sc to see (rejected H1=0.05, open H2=0.5) and pick H2 (often
-# wrong direction). New 0.4 keeps rejection signal but doesn't collapse H1
-# below the open-prior of H2, letting sc consider H1 text content again.
-# Case study: d6 L2 v2 2/69 hits were text-overlap luck on rejected H1, not
-# real reasoning. d12 L2 v2 0/69 because d12 hu rejects more confidently.
-
-
-def hyp_for_test_planner(h: Any) -> dict:
- """test_planner SFT: hypotheses[i] = {id, description, confidence}."""
- if not isinstance(h, dict): return {}
- return {
- "id": h.get("id"),
- "description": h.get("claim", "")[:280],
- "confidence": STATUS_TO_CONF.get(h.get("status", "open"), 0.5),
- }
-
-
-def hyp_for_hyp_updater(h: Any) -> str:
- """hyp_updater SFT: hypothesis = plain string (the claim only)."""
- if not isinstance(h, dict): return ""
- return h.get("claim", "")
-
-
-def hyp_for_stop_controller(h: Any) -> dict:
- """stop_controller SFT: current_hypotheses[i] = {hypothesis, probability}."""
- if not isinstance(h, dict): return {}
- return {
- "hypothesis": h.get("claim", "")[:280],
- "probability": STATUS_TO_CONF.get(h.get("status", "open"), 0.5),
- }
-
-
-# ---- evidence summary adapter (today's stop_controller fix) --------------------
-
-def evidence_summary_for_stop_controller(evidence_so_far: list) -> str:
- """SFT trained stop_controller with `evidence_summary` as a narrative string
- (median ~300 chars). Chain accumulates list-of-dicts; without adaptation
- the model goes OOD by turn 9-10 and emits other tasks' schemas."""
- if not evidence_so_far:
- return "No evidence collected yet."
- parts: list[str] = []
- for e in evidence_so_far:
- if not isinstance(e, dict):
- continue
- turn = e.get("turn", "?")
- tool = e.get("tool", "?")
- obs_list = e.get("observations") or []
- if not isinstance(obs_list, list):
- obs_list = []
- obs_phrases: list[str] = []
- for o in obs_list[:4]:
- if isinstance(o, dict):
- bits = []
- for k in ("metric", "type", "value", "note", "details"):
- v = o.get(k)
- if isinstance(v, (str, int, float)) and v != "":
- bits.append(f"{k}={v}")
- if bits:
- obs_phrases.append("(" + ", ".join(bits[:3]) + ")")
- obs_str = "; ".join(obs_phrases) if obs_phrases else "no parsed observations"
- parts.append(f"turn {turn}: ran {tool}; {obs_str}.")
- return " ".join(parts)
-
-
-# ---- 4 new input-shape adapters (2026-05-22 audit fixes) -----------------------
-
-def packet_summary_to_str(ps: Any) -> str:
- """world_model_builder SFT: packet_summary = narrative string (89% of
- sampled training rows). Chain feeds aggregate dict from tshark. Adapter
- renders dict as a single-sentence narrative."""
- if isinstance(ps, str):
- return ps
- if not isinstance(ps, dict):
- return str(ps)
- parts = [f"Total {ps.get('total_packets', '?')} packets"]
- bits = []
- if ps.get("tcp_packets") is not None:
- bits.append(f"{ps['tcp_packets']} TCP")
- if ps.get("tls_packets"):
- bits.append(f"{ps['tls_packets']} TLS")
- if bits:
- parts.append(f" ({', '.join(bits)})")
- if ps.get("retransmissions"):
- parts.append(f"; {ps['retransmissions']} retransmissions")
- if ps.get("zero_window_advertisements"):
- parts.append(f", {ps['zero_window_advertisements']} zero-window advertisements")
- return "".join(parts) + "."
-
-
-def time_window_to_dict(tw: Any) -> dict:
- """symptom_detector SFT: time_window = dict{start, end} (96% of rows).
- Chain produces 'start_ts / end_ts' string. Split on ' / ' (or similar)
- and wrap as dict."""
- if isinstance(tw, dict):
- return tw
- if not isinstance(tw, str):
- return {"start": str(tw), "end": str(tw)}
- for sep in (" / ", "/", " -> ", " — ", "—"):
- if sep in tw:
- parts = [p.strip() for p in tw.split(sep, 1)]
- if len(parts) == 2:
- return {"start": parts[0], "end": parts[1]}
- return {"start": tw, "end": tw}
-
-
-def available_tools_to_dicts(tools: Any) -> list[dict]:
- """test_planner SFT: available_tools = list[dict] (88% of rows). Chain
- passes list[str] names. Wrap each string as {name: str}; pass dicts
- through."""
- if not isinstance(tools, list):
- return []
- out: list[dict] = []
- for t in tools:
- if isinstance(t, dict):
- out.append(t)
- elif isinstance(t, str):
- out.append({"name": t})
- return out
-
-
-def list_or_dict_to_str(x: Any, max_items: int = 3) -> str:
- """hypothesis_updater SFT: prediction and observation are plain strings
- (84% / 56% of rows). Chain feeds list[dict] (from test_planner
- expected_observations / observation_normalizer observations). Render as a
- short narrative; prefer description/type/value/metric/note keys; fall back
- to truncated JSON."""
- if x is None:
- return ""
- if isinstance(x, str):
- return x
- if isinstance(x, dict):
- return ", ".join(f"{k}={v}" for k, v in list(x.items())[:5])
- if isinstance(x, list):
- parts: list[str] = []
- key_priority = ("description", "type", "value", "metric", "note", "details")
- for item in x[:max_items]:
- if isinstance(item, dict):
- bits: list[str] = []
- for k in key_priority:
- if k in item:
- bits.append(f"{k}={item[k]}")
- if len(bits) >= 3:
- break
- if not bits:
- bits.append(json.dumps(item, ensure_ascii=False)[:80])
- parts.append("; ".join(bits))
- else:
- parts.append(str(item)[:80])
- return " | ".join(parts) if parts else ""
- return str(x)
diff --git a/vendor/nanoai/agentic/diag/graders.py b/vendor/nanoai/agentic/diag/graders.py
deleted file mode 100644
index 4e4a6785357c74b366d994a05e9dc4c9358a5d16..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/graders.py
+++ /dev/null
@@ -1,271 +0,0 @@
-"""Graders for the 8 agentic diagnosis micro-tasks.
-
-Design contract: every grader returns
- (passed: bool, detail: dict)
-
-where `detail` MUST include:
- format_ok bool — JSON parsed AND required top-level keys present
- content_ok bool — all field-level content checks passed
- field_scores dict — per-field {ok, reason} breakdown (for case study)
- completion_head str — first 120 chars of completion (for triage)
-
-`passed = format_ok and content_ok`. Reporting separates the two so that a
-model failing to produce JSON (format_ok=0) is distinguished from one that
-produces JSON with wrong content (format_ok=high, content_ok=low). Conflating
-these was the failure mode of agent_eval_mt 38: a single binary pass/fail
-hid which agent rung was actually broken.
-
-Lenience policy (anti-Goodhart, [[lesson-r-arith-grader-too-lenient]]):
- * structural keys + enum values: STRICT (exact match required)
- * free-text fields: LOW bar — Jaccard ≥0.15 OR length ≥10 chars suffices
- * list lengths: target within ±100% of expected (just "non-empty" if min_items==1)
- * the goal: a model that produces sensible JSON in the right shape scores
- 40-70% on content_ok; pure garbage scores <10%
-"""
-from __future__ import annotations
-
-import json
-import re
-from typing import Any
-
-from .schemas import DIAG_TASKS, TaskSpec, FieldSpec
-
-
-# ---------- helpers ----------
-
-_JSON_OBJ_RE = re.compile(r"\{[\s\S]*\}")
-_TOKEN_RE = re.compile(r"\w+", re.UNICODE)
-
-
-def _extract_json(completion: str) -> dict | None:
- """Find the first top-level JSON object in completion and parse it.
-
- Tolerates leading/trailing prose, markdown code fences, etc. Greedy match
- on outermost braces; if that fails, return None.
- """
- if not completion:
- return None
- # Strip markdown code fences ```json ... ```
- t = completion.strip()
- if t.startswith("```"):
- t = re.sub(r"^```(?:json)?\s*", "", t)
- t = re.sub(r"\s*```$", "", t)
- m = _JSON_OBJ_RE.search(t)
- if not m:
- return None
- try:
- return json.loads(m.group(0))
- except json.JSONDecodeError:
- # Try a more permissive fallback: balance braces manually.
- try:
- depth = 0
- start = t.find("{")
- if start < 0:
- return None
- for i, ch in enumerate(t[start:], start=start):
- if ch == "{":
- depth += 1
- elif ch == "}":
- depth -= 1
- if depth == 0:
- return json.loads(t[start : i + 1])
- return None
- except Exception:
- return None
-
-
-def _tokens(s: str) -> set[str]:
- if not isinstance(s, str):
- return set()
- return set(_TOKEN_RE.findall(s.lower()))
-
-
-def _jaccard(a: str, b: str) -> float:
- ta, tb = _tokens(a), _tokens(b)
- if not ta and not tb:
- return 1.0
- if not ta or not tb:
- return 0.0
- return len(ta & tb) / len(ta | tb)
-
-
-# ---------- per-field checks ----------
-
-def _check_field(
- field: FieldSpec, predicted: Any, expected: Any
-) -> tuple[bool, bool, str]:
- """Return (format_ok_field, content_ok_field, reason).
-
- format_ok_field = whether the type/shape constraints are met
- content_ok_field = whether content matches (or is reasonably close to) expected
- reason = short string for case study
- """
- if predicted is None:
- return False, False, "missing"
-
- # Kind-specific shape check
- if field.kind == "str":
- if not isinstance(predicted, str):
- return False, False, f"not_str:{type(predicted).__name__}"
- if not predicted.strip():
- return True, False, "empty"
- if field.freetext:
- # low-bar content check: Jaccard ≥0.15 OR ≥10 chars present
- j = _jaccard(predicted, expected if isinstance(expected, str) else "")
- length_ok = len(predicted) >= 10
- content = j >= 0.15 or length_ok
- return True, content, f"jaccard={j:.2f}_len={len(predicted)}"
- else:
- # exact match required for non-freetext str
- content = predicted.strip() == (expected.strip() if isinstance(expected, str) else "")
- return True, content, "exact" if content else f"mismatch:{predicted[:40]!r}"
-
- if field.kind == "int":
- if not isinstance(predicted, int) or isinstance(predicted, bool):
- return False, False, f"not_int:{type(predicted).__name__}"
- content = predicted == expected if isinstance(expected, int) else True
- return True, content, "ok"
-
- if field.kind == "float":
- if not isinstance(predicted, (int, float)) or isinstance(predicted, bool):
- return False, False, f"not_num:{type(predicted).__name__}"
- if isinstance(expected, (int, float)):
- content = abs(predicted - expected) < 1e-6
- else:
- content = True
- return True, content, "ok"
-
- if field.kind == "enum":
- if not isinstance(predicted, str):
- return False, False, f"not_str:{type(predicted).__name__}"
- allowed = field.enum_values or []
- if predicted not in allowed:
- return False, False, f"not_in_enum:{predicted!r}"
- # content check: matches expected if provided
- content = (predicted == expected) if isinstance(expected, str) else True
- return True, content, ("match" if content else f"wrong_enum:{predicted!r}")
-
- if field.kind == "list":
- if not isinstance(predicted, list):
- return False, False, f"not_list:{type(predicted).__name__}"
- if len(predicted) < field.min_items:
- return False, False, f"too_few:{len(predicted)}<{field.min_items}"
- # content: count within ±100% of expected length (very generous)
- if isinstance(expected, list) and expected:
- ratio = len(predicted) / len(expected)
- content = 0.5 <= ratio <= 2.0
- else:
- content = True
- return True, content, f"len={len(predicted)}"
-
- if field.kind == "list_obj":
- if not isinstance(predicted, list):
- return False, False, f"not_list:{type(predicted).__name__}"
- if len(predicted) < field.min_items:
- return False, False, f"too_few:{len(predicted)}<{field.min_items}"
- # every item must be a dict with the required keys
- required = field.item_keys or []
- for i, item in enumerate(predicted):
- if not isinstance(item, dict):
- return False, False, f"item{i}_not_dict"
- missing = [k for k in required if k not in item]
- if missing:
- return False, False, f"item{i}_missing:{missing}"
- # enum check on per-item enum-constrained fields
- if field.item_enums:
- for key, allowed in field.item_enums.items():
- if key in item and item[key] not in allowed:
- return False, False, f"item{i}.{key}_not_in_enum:{item[key]!r}"
- # content: length within ±100% of expected if expected is a list
- if isinstance(expected, list) and expected:
- ratio = len(predicted) / len(expected)
- content = 0.5 <= ratio <= 2.0
- else:
- content = True
- return True, content, f"len={len(predicted)}"
-
- if field.kind == "obj":
- if not isinstance(predicted, dict):
- return False, False, f"not_dict:{type(predicted).__name__}"
- if not predicted:
- return True, False, "empty_dict"
- # content: ≥3 keys present (loose) and at least 1 overlapping key with expected
- if isinstance(expected, dict) and expected:
- overlap = set(predicted.keys()) & set(expected.keys())
- content = len(overlap) >= 1
- else:
- content = len(predicted) >= 1
- return True, content, f"keys={sorted(predicted.keys())[:5]}"
-
- return False, False, f"unknown_kind:{field.kind}"
-
-
-# ---------- generic spec-based grader ----------
-
-def grade_against_spec(
- spec: TaskSpec, expected: dict, completion: str
-) -> tuple[bool, dict]:
- """Generic grader: validate `completion` JSON against `spec` and compare to `expected`."""
- head = (completion or "")[:120]
-
- parsed = _extract_json(completion)
- if parsed is None:
- return False, {
- "format_ok": False,
- "content_ok": False,
- "field_scores": {},
- "completion_head": head,
- "reason": "no_json",
- }
-
- field_scores: dict[str, dict] = {}
- format_ok_all = True
- content_ok_all = True
-
- # Check required top-level keys present
- missing_top = [f.name for f in spec.output_fields
- if f.required and f.name not in parsed]
- if missing_top:
- return False, {
- "format_ok": False,
- "content_ok": False,
- "field_scores": {k: {"ok": False, "reason": "missing_top_key"} for k in missing_top},
- "completion_head": head,
- "reason": f"missing_top:{missing_top}",
- }
-
- # Per-field check
- for f in spec.output_fields:
- predicted_val = parsed.get(f.name)
- expected_val = expected.get(f.name) if isinstance(expected, dict) else None
- f_format, f_content, reason = _check_field(f, predicted_val, expected_val)
- field_scores[f.name] = {"ok": f_format and f_content, "format": f_format,
- "content": f_content, "reason": reason}
- if not f_format:
- format_ok_all = False
- if not f_content:
- content_ok_all = False
-
- return format_ok_all and content_ok_all, {
- "format_ok": format_ok_all,
- "content_ok": content_ok_all,
- "field_scores": field_scores,
- "completion_head": head,
- }
-
-
-# ---------- 8 thin wrappers (registered into GRADERS) ----------
-
-def _make_grader(task_id: str):
- spec = DIAG_TASKS[task_id]
-
- def grader(row: dict, completion: str) -> tuple[bool, dict]:
- expected = row.get("expected_output", {}) or {}
- return grade_against_spec(spec, expected, completion)
- grader.__name__ = f"grade_{task_id}"
- return grader
-
-
-GRADERS: dict[str, Any] = {
- f"diag_{tid}": _make_grader(tid) for tid in DIAG_TASKS.keys()
-}
diff --git a/vendor/nanoai/agentic/diag/graders_v2.py b/vendor/nanoai/agentic/diag/graders_v2.py
deleted file mode 100644
index 682845f812c4dfc0267664a1b0875df55b3431b5..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/graders_v2.py
+++ /dev/null
@@ -1,651 +0,0 @@
-"""Strict semantic graders for diag_eval_v1 — v2 of graders.py.
-
-Motivation: v1 graders (graders.py) are SHAPE-graders. The
-`experiments/ksbr_pcap/scripts/fake_shape_baseline.py` audit (2026-05-21) showed shape-correct
-gibberish achieves 99-100% pass on 3 tasks (hypothesis_updater, hypothesis_generator,
-symptom_detector) and 43.4% overall. That makes v1 numbers a JSON-schema-
-compliance score, not a diagnostic-capability score.
-
-v2 adds task-specific semantic checks that cross-reference the model's
-predicted content against the row's input_state. Each task has a
-hand-tuned check function. Default: STRICT — passing means the model's
-output has structure AND content tied to the input.
-
-Result contract is identical to v1 (passed, detail with format_ok /
-content_ok / field_scores / completion_head). v1 GRADERS are unaffected.
-
-Registered as `diag__v2`. Use `scripts/regrade_jsonl.py` to
-re-grade existing v1 result jsonl with v2 graders without re-running eval.
-"""
-from __future__ import annotations
-
-import json
-import re
-from typing import Any
-
-from .schemas import DIAG_TASKS, TaskSpec, FieldSpec
-from .graders import _extract_json, _tokens, _jaccard
-
-# ---------- helpers ----------
-
-_IPV4_RE = re.compile(r"^(?:\d{1,3}\.){3}\d{1,3}(?::\d+)?$")
-_IPV6_RE = re.compile(r"^[0-9a-fA-F:]+(?::\d+)?$")
-_FLOW_TUPLE_RE = re.compile(
- r"^\s*[0-9a-fA-F\.:\-_]+(?::\d+)?\s*(?:->|→|<-)\s*[0-9a-fA-F\.:\-_]+(?::\d+)?\s*$"
-)
-_DIGIT_RE = re.compile(r"^\d+$")
-_FAKE_FILLER_RE = re.compile(r"^(x+|y+|foo|bar|baz|test|xxxx+|placeholder)$", re.IGNORECASE)
-
-# Real-ish PCAP tool/check names. Substring-matched to be tolerant of
-# variation like "tshark_filter" vs "tshark.filter".
-_PCAP_TOOL_TOKENS = {
- "tshark", "tcpdump", "wireshark", "zeek", "bro", "scapy",
- "ping", "traceroute", "tracert", "mtr", "nmap", "dig", "nslookup", "host",
- "ss", "netstat", "ip", "tc", "ethtool", "iptables", "ipset", "nft",
- "iperf", "iperf3", "curl", "wget", "openssl",
- "syn_probe", "syn_scan", "tcp_probe", "tcp_retransmit", "tcp_window",
- "mtu_probe", "path_mtu", "dns_query", "dns_check", "icmp_ping",
- "tcp_stream", "flow_extract", "filter", "capture",
- "kernel_log", "dmesg", "journalctl", "sar", "vmstat", "iostat",
- "cpu_usage", "memory_check", "disk_io", "load_average",
- "ipsec", "ike", "strongswan", "openvpn", "wg",
- "stat", "log", "metric",
-}
-
-_PROTOCOL_TOKENS = {
- "TCP", "UDP", "ICMP", "ICMPv6", "GRE", "ESP", "AH", "SCTP", "QUIC",
- "TLS", "DTLS", "SSH", "HTTP", "HTTPS", "HTTP/2", "HTTP/3", "SMTP",
- "DNS", "DHCP", "BGP", "OSPF", "ARP", "NDP", "MPTCP", "MPLS", "VXLAN",
- "WireGuard", "IPSec", "L2TP", "GTP", "SIP", "RTP", "RTCP", "MQTT",
-}
-
-
-def _is_filler(s: Any) -> bool:
- """Return True if s is fake placeholder text like 'xxxx' / 'foo' / empty.
-
- Numeric strings (e.g. "1", "10", "12345") and short IDs are NOT filler.
- Repetitive single-char fillers (xxxx, aaaa) ARE.
- """
- if not isinstance(s, str):
- return False
- t = s.strip()
- if not t:
- return True
- # Explicit fake words
- if _FAKE_FILLER_RE.match(t):
- return True
- # Pure digits or hex chars are valid identifiers (packet_ref, IDs, etc)
- if t.isdigit() or all(c in "0123456789abcdefABCDEF:.-_" for c in t):
- return False
- # Repetitive single-char strings ≥ 4 chars (xxxx, aaaa, ----)
- if len(t) >= 4 and len(set(t)) <= 2:
- return True
- return False
-
-
-def _has_pcap_tool_token(s: str) -> bool:
- if not isinstance(s, str):
- return False
- sl = s.lower()
- return any(tok in sl for tok in _PCAP_TOOL_TOKENS)
-
-
-def _is_ip(s: str) -> bool:
- if not isinstance(s, str):
- return False
- s = s.strip()
- if _IPV4_RE.match(s):
- # validate octets
- host = s.split(":")[0]
- try:
- return all(0 <= int(o) <= 255 for o in host.split("."))
- except Exception:
- return False
- if _IPV6_RE.match(s) and "." not in s:
- return ":" in s and any(c in "0123456789abcdefABCDEF" for c in s)
- return False
-
-
-def _collect_ips_from_input_state(input_state: dict) -> set[str]:
- """Pull plausible IP strings from anywhere in input_state."""
- out: set[str] = set()
-
- def _walk(obj):
- if isinstance(obj, str):
- for tok in re.findall(r"\d{1,3}(?:\.\d{1,3}){3}", obj):
- if _is_ip(tok):
- out.add(tok)
- elif isinstance(obj, dict):
- for v in obj.values():
- _walk(v)
- elif isinstance(obj, list):
- for v in obj:
- _walk(v)
- _walk(input_state)
- return out
-
-
-def _input_state_text(input_state: dict) -> str:
- """Flatten input_state to a searchable text blob for cross-checks."""
- try:
- return json.dumps(input_state, ensure_ascii=False).lower()
- except Exception:
- return str(input_state).lower()
-
-
-def _result(format_ok: bool, content_ok: bool, field_scores: dict,
- completion_head: str, reason: str = "") -> tuple[bool, dict]:
- d = {
- "format_ok": format_ok,
- "content_ok": content_ok,
- "field_scores": field_scores,
- "completion_head": completion_head,
- "grader_version": "v2",
- }
- if reason:
- d["reason"] = reason
- return format_ok and content_ok, d
-
-
-def _parse(completion: str) -> tuple[dict | None, str]:
- head = (completion or "")[:120]
- return _extract_json(completion), head
-
-
-# ---------- per-task semantic graders ----------
-
-def grade_problem_framing_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: scope must have ≥2 real keys; goal/question must overlap with input request."""
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- required = ["diagnosis_goal", "scope", "primary_question", "constraints", "initial_uncertainties"]
- missing = [k for k in required if k not in parsed]
- if missing:
- return _result(False, False,
- {k: {"ok": False, "reason": "missing_top_key"} for k in missing},
- head, f"missing_top:{missing}")
-
- user_req = input_state.get("user_request", "")
- user_tokens = _tokens(user_req)
- fs = {}
-
- # diagnosis_goal: must be str ≥ 30 chars + overlap with input topic
- goal = parsed.get("diagnosis_goal", "")
- goal_ok = (isinstance(goal, str) and len(goal) >= 30
- and not _is_filler(goal)
- and (not user_tokens or _jaccard(goal, user_req) >= 0.10))
- fs["diagnosis_goal"] = {"ok": goal_ok, "format": isinstance(goal, str),
- "content": goal_ok, "reason": f"len={len(goal) if isinstance(goal, str) else 0} jaccard={_jaccard(goal, user_req):.2f}"}
-
- # scope: dict with ≥2 keys, must include a scope-like key
- scope = parsed.get("scope", {})
- scope_keys = set(scope.keys()) if isinstance(scope, dict) else set()
- scope_topic_keys = {
- "protocol_layer", "protocol_layers", "protocols",
- "time_window", "duration", "timeframe",
- "traffic_direction", "direction",
- "focus_area", "focus", "scope_boundaries", "boundaries",
- "endpoints", "ip_range", "host_filter",
- }
- has_topic = bool(scope_keys & scope_topic_keys)
- scope_ok = (isinstance(scope, dict) and len(scope_keys) >= 2 and has_topic)
- fs["scope"] = {"ok": scope_ok, "format": isinstance(scope, dict),
- "content": has_topic, "reason": f"keys={sorted(scope_keys)[:6]}"}
-
- # primary_question: str ≥ 20 chars; should be question-like (contain '?' or 'what'/'why'/'how'/etc)
- pq = parsed.get("primary_question", "")
- q_tokens = {"what", "why", "how", "when", "where", "which", "?"}
- pq_ok = (isinstance(pq, str) and len(pq) >= 20 and not _is_filler(pq)
- and (any(t in pq.lower() for t in q_tokens) or "?" in pq))
- fs["primary_question"] = {"ok": pq_ok, "format": isinstance(pq, str),
- "content": pq_ok, "reason": f"len={len(pq) if isinstance(pq, str) else 0}"}
-
- # constraints & uncertainties: lists of substantive strings
- for k in ["constraints", "initial_uncertainties"]:
- lst = parsed.get(k, [])
- if not isinstance(lst, list) or len(lst) == 0:
- fs[k] = {"ok": False, "format": False, "content": False, "reason": "not_list_or_empty"}
- continue
- substantive = [x for x in lst if isinstance(x, str)
- and len(x) >= 20 and not _is_filler(x)]
- ok = len(substantive) >= 1
- fs[k] = {"ok": ok, "format": True, "content": ok,
- "reason": f"substantive={len(substantive)}/{len(lst)}"}
-
- format_ok = all(v.get("format", False) for v in fs.values())
- content_ok = all(v.get("content", False) for v in fs.values())
- return _result(format_ok, content_ok, fs, head)
-
-
-def grade_world_model_builder_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: entities must have valid IPs; flows must have valid tuple format; IPs overlap with input."""
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- required = ["entities", "flows", "capture_assumptions", "world_model_risks"]
- missing = [k for k in required if k not in parsed]
- if missing:
- return _result(False, False,
- {k: {"ok": False, "reason": "missing_top_key"} for k in missing},
- head, f"missing_top:{missing}")
-
- input_ips = _collect_ips_from_input_state(input_state)
- fs = {}
-
- # entities: list of dicts with valid IP + type enum
- ents = parsed.get("entities", [])
- if not isinstance(ents, list) or len(ents) < 2:
- fs["entities"] = {"ok": False, "format": False, "content": False, "reason": "not_list_or_too_few"}
- else:
- good = []
- bad = []
- for i, e in enumerate(ents):
- if not isinstance(e, dict):
- bad.append(f"item{i}_not_dict")
- continue
- ip = e.get("ip", "")
- etype = e.get("type", "")
- if not _is_ip(ip):
- bad.append(f"item{i}_bad_ip:{ip}")
- continue
- if etype not in {"client", "server", "intermediate", "unknown", "router", "gateway", "proxy", "loadbalancer"}:
- bad.append(f"item{i}_bad_type:{etype}")
- continue
- good.append(ip)
- # at least 2 good entities; if input has known IPs, at least 1 must overlap
- n_overlap = len(set(good) & input_ips) if input_ips else 0
- ok = (len(good) >= 2 and (not input_ips or n_overlap >= 1))
- fs["entities"] = {"ok": ok, "format": len(good) >= 2,
- "content": (not input_ips or n_overlap >= 1),
- "reason": f"good={len(good)} overlap={n_overlap}/{len(input_ips)} bad={bad[:3]}"}
-
- # flows: list_obj with tuple as EITHER "IP:port -> IP:port" string OR dict with src/dst keys
- flows = parsed.get("flows", [])
- if not isinstance(flows, list) or len(flows) < 1:
- fs["flows"] = {"ok": False, "format": False, "content": False, "reason": "not_list_or_empty"}
- else:
- ok_count = 0
- for f in flows:
- if not isinstance(f, dict):
- continue
- t = f.get("tuple")
- p = f.get("protocol", "")
- t_ok = False
- if isinstance(t, str) and _FLOW_TUPLE_RE.match(t.strip()):
- t_ok = True
- elif isinstance(t, dict):
- # dict form: must have src_ip/dst_ip (or equivalent)
- src_keys = {"src_ip", "source_ip", "src", "source"}
- dst_keys = {"dst_ip", "destination_ip", "dst", "destination"}
- has_src = any(k in t for k in src_keys)
- has_dst = any(k in t for k in dst_keys)
- t_ok = has_src and has_dst
- if t_ok and isinstance(p, str) and len(p) >= 2:
- ok_count += 1
- ok = ok_count >= 1
- fs["flows"] = {"ok": ok, "format": ok, "content": ok,
- "reason": f"valid_tuple_protocol={ok_count}/{len(flows)}"}
-
- # capture_assumptions: list of substantive strings
- ca = parsed.get("capture_assumptions", [])
- sub = [x for x in (ca if isinstance(ca, list) else [])
- if isinstance(x, str) and len(x) >= 30 and not _is_filler(x)]
- fs["capture_assumptions"] = {"ok": len(sub) >= 1, "format": isinstance(ca, list),
- "content": len(sub) >= 1,
- "reason": f"substantive={len(sub)}"}
-
- # world_model_risks: list of substantive strings
- wmr = parsed.get("world_model_risks", [])
- sub = [x for x in (wmr if isinstance(wmr, list) else [])
- if isinstance(x, str) and len(x) >= 20 and not _is_filler(x)]
- fs["world_model_risks"] = {"ok": len(sub) >= 1, "format": isinstance(wmr, list),
- "content": len(sub) >= 1,
- "reason": f"substantive={len(sub)}"}
-
- format_ok = all(v.get("format", False) for v in fs.values())
- content_ok = all(v.get("content", False) for v in fs.values())
- return _result(format_ok, content_ok, fs, head)
-
-
-def grade_symptom_detector_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: symptoms must have valid enum types AND evidence_refs grounded in input."""
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- if "symptoms" not in parsed:
- return _result(False, False, {"symptoms": {"ok": False, "reason": "missing"}},
- head, "missing_top")
-
- symptoms = parsed.get("symptoms", [])
- if not isinstance(symptoms, list) or len(symptoms) < 1:
- return _result(False, False, {"symptoms": {"ok": False, "reason": "not_list"}},
- head, "not_list")
-
- allowed_types = {
- "tcp_retransmission_burst", "tcp_duplicate_ack", "tcp_zero_window",
- "tcp_reset", "server_response_gap", "client_idle_gap",
- "tls_handshake_delay", "dns_lookup_delay", "dns_no_response",
- "rtt_spike", "high_jitter", "mtu_blackhole",
- }
- sev_allowed = {"low", "medium", "high", "critical"}
- input_text = _input_state_text(input_state)
-
- ok_count = 0
- bad = []
- for i, s in enumerate(symptoms):
- if not isinstance(s, dict):
- bad.append(f"item{i}_not_dict")
- continue
- if s.get("type") not in allowed_types:
- bad.append(f"item{i}_bad_type:{s.get('type')}")
- continue
- if s.get("severity") not in sev_allowed:
- bad.append(f"item{i}_bad_severity")
- continue
- # evidence_refs must be non-empty list of substantive strings
- refs = s.get("evidence_refs", [])
- if not isinstance(refs, list) or len(refs) == 0:
- bad.append(f"item{i}_no_refs")
- continue
- sub_refs = [r for r in refs if isinstance(r, str) and len(r) >= 4 and not _is_filler(r)]
- if len(sub_refs) == 0:
- bad.append(f"item{i}_filler_refs")
- continue
- # grounding: at least one ref-token must appear in input tokens.
- # (Jaccard was unfair when refs is small {1-3 tokens} vs input {200+}.)
- from .graders import _tokens as _tk
- refs_joined = " ".join(sub_refs)
- ref_tokens = _tk(refs_joined)
- input_tokens = _tk(input_text)
- if ref_tokens and not (ref_tokens & input_tokens):
- bad.append(f"item{i}_ungrounded_refs")
- continue
- ok_count += 1
- ok = ok_count >= 1
- fs = {"symptoms": {"ok": ok, "format": ok_count >= 1, "content": ok_count >= 1,
- "reason": f"valid={ok_count}/{len(symptoms)} bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_hypothesis_generator_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: ≥2 substantive hypotheses with grounded claim + falsifiable tests."""
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- if "hypotheses" not in parsed:
- return _result(False, False, {"hypotheses": {"ok": False, "reason": "missing"}},
- head, "missing_top")
-
- hyps = parsed.get("hypotheses", [])
- if not isinstance(hyps, list) or len(hyps) < 2:
- return _result(False, False,
- {"hypotheses": {"ok": False, "format": False, "content": False,
- "reason": f"only_{len(hyps) if isinstance(hyps, list) else 0}"}},
- head, "too_few_hyps")
-
- input_text = _input_state_text(input_state)
- ok_count = 0
- bad = []
- for i, h in enumerate(hyps):
- if not isinstance(h, dict):
- bad.append(f"item{i}_not_dict")
- continue
- claim = h.get("claim", "")
- if not isinstance(claim, str) or len(claim) < 40 or _is_filler(claim):
- bad.append(f"item{i}_bad_claim_len={len(claim) if isinstance(claim, str) else 0}")
- continue
- # claim should overlap with input_state (any of world_model/symptoms/diag_goal)
- if input_text and _jaccard(claim, input_text) < 0.03:
- bad.append(f"item{i}_claim_off_topic")
- continue
- for k in ["explains", "does_not_explain", "required_tests"]:
- v = h.get(k, [])
- if not isinstance(v, list) or len(v) < 1:
- bad.append(f"item{i}.{k}_not_list_or_empty")
- break
- sub = [x for x in v if isinstance(x, str) and len(x) >= 15 and not _is_filler(x)]
- if len(sub) < 1:
- bad.append(f"item{i}.{k}_no_substantive")
- break
- else:
- # required_tests must mention real tools / methods
- tests = h.get("required_tests", [])
- tool_hits = sum(1 for t in tests if isinstance(t, str) and _has_pcap_tool_token(t))
- if tool_hits == 0:
- bad.append(f"item{i}.required_tests_no_tools")
- continue
- ok_count += 1
- ok = ok_count >= 2
- fs = {"hypotheses": {"ok": ok, "format": ok_count >= 2, "content": ok_count >= 2,
- "reason": f"good={ok_count}/{len(hyps)} bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_test_planner_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: next_action must have real tool + arguments dict + expected_observations with multiple outcomes."""
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- if "next_action" not in parsed:
- return _result(False, False, {"next_action": {"ok": False, "reason": "missing"}},
- head, "missing_top")
- na = parsed.get("next_action", {})
- if not isinstance(na, dict):
- return _result(False, False,
- {"next_action": {"ok": False, "format": False, "reason": "not_dict"}},
- head, "not_dict")
-
- # tool / tool_name (model uses either)
- tool = na.get("tool") or na.get("tool_name") or ""
- tool_ok = (isinstance(tool, str) and len(tool) >= 4 and not _is_filler(tool))
- # arguments: dict with ≥1 key; try multiple key names (arguments/args/params)
- args = na.get("arguments") or na.get("args") or na.get("params") or {}
- args_ok = isinstance(args, dict) and len(args) >= 1 and not all(_is_filler(str(v)) for v in args.values())
- # expected_observations: list with ≥2 outcomes OR dict with ≥2 keys
- eobs = na.get("expected_observations")
- if isinstance(eobs, list):
- eobs_ok = len(eobs) >= 2 and all(isinstance(o, dict) for o in eobs)
- elif isinstance(eobs, dict):
- eobs_ok = len(eobs) >= 2
- else:
- eobs_ok = False
-
- fs = {"next_action": {
- "ok": tool_ok and args_ok and eobs_ok,
- "format": isinstance(na, dict),
- "content": tool_ok and args_ok and eobs_ok,
- "reason": f"tool={tool[:30]!r}_ok={tool_ok} args_ok={args_ok} eobs_ok={eobs_ok}",
- }}
- ok = tool_ok and args_ok and eobs_ok
- return _result(ok, ok, fs, head)
-
-
-def grade_observation_normalizer_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: observations must have grounded flow_id + parseable packet_ref + numeric time."""
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- if "observations" not in parsed:
- return _result(False, False, {"observations": {"ok": False, "reason": "missing"}},
- head, "missing_top")
- obs = parsed.get("observations", [])
- if not isinstance(obs, list) or len(obs) < 1:
- return _result(False, False, {"observations": {"ok": False, "reason": "not_list"}},
- head, "not_list")
-
- input_text = _input_state_text(input_state)
- ok_count = 0
- bad = []
- for i, o in enumerate(obs):
- if not isinstance(o, dict):
- bad.append(f"item{i}_not_dict")
- continue
- otype = o.get("type", "")
- if not isinstance(otype, str) or len(otype) < 3 or _is_filler(otype):
- bad.append(f"item{i}_bad_type")
- continue
- flow_id = o.get("flow_id", "")
- # flow_id: non-filler, ≥3 char id (any chars; capagent uses e.g. "CxA123.0001", "flow_001", "f1")
- if not isinstance(flow_id, str) or len(flow_id) < 2 or _is_filler(flow_id):
- bad.append(f"item{i}_bad_flow_id")
- continue
- # packet_ref: non-filler ≥1 char (int OR string — capagent uses composite IDs like "CxA123_1698417000")
- pref = o.get("packet_ref")
- if pref is None:
- bad.append(f"item{i}_no_packet_ref")
- continue
- if isinstance(pref, int):
- pref_ok = True
- elif isinstance(pref, str) and len(pref) >= 1 and not _is_filler(pref):
- pref_ok = True
- else:
- pref_ok = False
- if not pref_ok:
- bad.append(f"item{i}_bad_packet_ref:{pref!r}")
- continue
- # time: numeric int/float OR string parseable as float
- t = o.get("time")
- t_ok = False
- if isinstance(t, (int, float)) and not isinstance(t, bool):
- t_ok = True
- elif isinstance(t, str):
- try:
- float(t.strip())
- t_ok = True
- except (ValueError, AttributeError):
- pass
- if not t_ok:
- bad.append(f"item{i}_bad_time:{t!r}")
- continue
- ok_count += 1
- ok = ok_count >= 1
- fs = {"observations": {"ok": ok, "format": ok_count >= 1, "content": ok_count >= 1,
- "reason": f"valid={ok_count}/{len(obs)} bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_hypothesis_updater_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: update id must match input hypothesis.id, reason must cite observation."""
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- if "hypothesis_updates" not in parsed:
- return _result(False, False, {"hypothesis_updates": {"ok": False, "reason": "missing"}},
- head, "missing_top")
- updates = parsed.get("hypothesis_updates", [])
- if not isinstance(updates, list) or len(updates) < 1:
- return _result(False, False, {"hypothesis_updates": {"ok": False, "reason": "not_list"}},
- head, "not_list")
-
- # input hypothesis id
- h_input = input_state.get("hypothesis", {})
- expected_id = h_input.get("id") if isinstance(h_input, dict) else None
- input_text = _input_state_text(input_state)
-
- allowed_status = {
- "provisionally_supported", "weakened", "rejected",
- "confirmed", "inconclusive", "unchanged",
- }
-
- ok_count = 0
- bad = []
- for i, u in enumerate(updates):
- if not isinstance(u, dict):
- bad.append(f"item{i}_not_dict")
- continue
- uid = u.get("id")
- if expected_id and uid != expected_id:
- bad.append(f"item{i}_id_mismatch:{uid!r}!={expected_id!r}")
- continue
- status = u.get("status")
- if status not in allowed_status:
- bad.append(f"item{i}_bad_status:{status!r}")
- continue
- reason = u.get("reason", "")
- if not isinstance(reason, str) or len(reason) < 20 or _is_filler(reason):
- bad.append(f"item{i}_bad_reason_len={len(reason) if isinstance(reason, str) else 0}")
- continue
- # reason should overlap with input_state (was: only with observation — too narrow)
- if input_text and _jaccard(reason, input_text) < 0.03:
- bad.append(f"item{i}_reason_off_topic")
- continue
- ok_count += 1
- ok = ok_count >= 1
- fs = {"hypothesis_updates": {"ok": ok, "format": ok_count >= 1, "content": ok_count >= 1,
- "reason": f"valid={ok_count}/{len(updates)} bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_stop_controller_v2(row: dict, completion: str) -> tuple[bool, dict]:
- """Strict: enum match expected + reason substantive + required_before_stop consistent."""
- expected = row.get("expected_output", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- if "stop_decision" not in parsed:
- return _result(False, False, {"stop_decision": {"ok": False, "reason": "missing"}},
- head, "missing_top")
- enum_allowed = {"continue", "stop_with_conclusion", "stop_inconclusive"}
- sd = parsed.get("stop_decision")
- if sd not in enum_allowed:
- return _result(False, False,
- {"stop_decision": {"ok": False, "format": False, "content": False,
- "reason": f"bad_enum:{sd!r}"}},
- head, "bad_enum")
- expected_sd = expected.get("stop_decision")
- sd_match = (sd == expected_sd) if expected_sd else True
-
- reason = parsed.get("reason", "")
- reason_ok = (isinstance(reason, str) and len(reason) >= 50 and not _is_filler(reason))
-
- rbs = parsed.get("required_before_stop", [])
- if sd == "continue":
- rbs_ok = isinstance(rbs, list) and len(rbs) >= 1 and all(isinstance(x, str) and len(x) >= 5 for x in rbs)
- else:
- rbs_ok = isinstance(rbs, list)
-
- fs = {
- "stop_decision": {"ok": sd_match, "format": True, "content": sd_match,
- "reason": f"got={sd} expected={expected_sd}"},
- "reason": {"ok": reason_ok, "format": isinstance(reason, str),
- "content": reason_ok, "reason": f"len={len(reason) if isinstance(reason, str) else 0}"},
- "required_before_stop": {"ok": rbs_ok, "format": isinstance(rbs, list),
- "content": rbs_ok, "reason": f"n={len(rbs) if isinstance(rbs, list) else 0}"},
- }
- format_ok = all(v.get("format", False) for v in fs.values())
- content_ok = all(v.get("content", False) for v in fs.values())
- return _result(format_ok, content_ok, fs, head)
-
-
-# ---------- registry ----------
-
-GRADERS_V2: dict[str, Any] = {
- "diag_problem_framing_v2": grade_problem_framing_v2,
- "diag_world_model_builder_v2": grade_world_model_builder_v2,
- "diag_symptom_detector_v2": grade_symptom_detector_v2,
- "diag_hypothesis_generator_v2": grade_hypothesis_generator_v2,
- "diag_test_planner_v2": grade_test_planner_v2,
- "diag_observation_normalizer_v2": grade_observation_normalizer_v2,
- "diag_hypothesis_updater_v2": grade_hypothesis_updater_v2,
- "diag_stop_controller_v2": grade_stop_controller_v2,
-}
-
-# v1 → v2 task mapping for regrade scripts
-V1_TO_V2 = {
- f"diag_{tid}": f"diag_{tid}_v2" for tid in DIAG_TASKS.keys()
-}
diff --git a/vendor/nanoai/agentic/diag/graders_v3.py b/vendor/nanoai/agentic/diag/graders_v3.py
deleted file mode 100644
index a5c0f3af0d9a54dd3bb6b29b7574e18431870593..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/graders_v3.py
+++ /dev/null
@@ -1,710 +0,0 @@
-"""v3 graders — semantic comparison against expected_output.
-
-Key fix vs v2.1: every task explicitly compares predicted to expected_output
-(not just to input_state). Closes the input-copy template loophole that the
-operator audit exposed (v2.1 input-copy baseline = 80.9% pass; v3 target <30%).
-
-Design criteria (must satisfy ALL three):
- - fake-shape gibberish baseline: < 10% pass
- - input-copy template baseline: < 30% pass
- - capagent GT self-grade: ≥ 70% pass (calibrates against real noise)
-
-Per-task adds over v2.1:
- problem_framing → scope keys overlap with EXPECTED scope (≥50%)
- world_model_builder → protocol must match expected; entity count ±50%
- symptom_detector → symptom type ∈ expected_types set
- hypothesis_generator → n_hyps ±1 of expected; claim Jaccard vs EXPECTED claims (not input symptoms); shrunk tool whitelist
- test_planner → tool matches expected (semantic equiv class); expected_obs has hypothesis cross-impact
- observation_normalizer → flow_id ∈ expected_flow_ids; type ∈ expected_types
- hypothesis_updater → status MUST equal expected.status (was missing)
- stop_controller → unchanged (already expected-match)
-"""
-from __future__ import annotations
-
-import json
-import re
-from typing import Any
-
-from .schemas import DIAG_TASKS, TaskSpec, FieldSpec
-from .graders import _extract_json, _tokens, _jaccard
-from .graders_v2 import (
- _is_filler, _is_ip, _collect_ips_from_input_state, _input_state_text,
- _result, _parse, _FLOW_TUPLE_RE,
-)
-
-# ---------- helpers ----------
-
-# Narrower tool whitelist — remove ambiguous 2-letter substrings.
-_PCAP_TOOL_TOKENS_V3 = {
- "tshark", "tcpdump", "wireshark", "zeek", "bro", "scapy",
- "traceroute", "tracert", "mtr", "nmap", "dig", "nslookup",
- "netstat", "ethtool", "iptables", "ipset", "nftables",
- "iperf", "iperf3", "openssl", "curl", "wget",
- "syn_probe", "syn_scan", "tcp_probe", "tcp_retransmit",
- "mtu_probe", "path_mtu", "dns_query", "icmp_ping",
- "tcp_stream", "kernel_log", "dmesg", "journalctl",
- "ipsec", "ike", "strongswan", "wireguard",
-}
-
-# Broader anchor set for hypothesis_generator.required_tests semantic check.
-# Real-world diagnostic test plans use protocol nouns ("TCP buffer", "path MTU"),
-# diagnostic mechanisms ("firewall logs", "NIC interrupt"), and concept words
-# ("pcap capture", "throughput") in addition to CLI tool names. Reviewed-accept
-# GT was 91% rejected by the narrow _PCAP_TOOL_TOKENS_V3 alone (2026-05-21
-# human-review audit). The anti-input-copy bite for hyp_gen comes from the
-# claim Jaccard vs EXPECTED claims (line below), not from this whitelist.
-_HYP_TEST_TOKENS = _PCAP_TOOL_TOKENS_V3 | {
- # protocols / layers
- "tcp", "udp", "icmp", "dns", "tls", "ssl", "http", "https", "ssh", "arp",
- # tcp/ip mechanisms (multi-letter, specific)
- "syn", "rst", "fin", "mtu", "rtt", "mss", "window",
- "retransmission", "retransmit", "duplicate",
- # network artifacts / verbs
- "pcap", "packet", "capture", "interface", "interrupt", "frame", "header",
- "ping",
- # network components
- "firewall", "nat", "router", "gateway", "switch", "nic", "vlan", "vpn",
- # diagnostics resources
- "buffer", "queue", "throughput", "latency", "jitter", "kernel",
- # common server/service names in test plans
- "nginx", "apache", "haproxy", "envoy", "log", "logs",
- # cpu/process diagnostics
- "cpu",
-}
-
-
-def _has_real_tool_token(s: str) -> bool:
- if not isinstance(s, str):
- return False
- sl = s.lower()
- return any(tok in sl for tok in _PCAP_TOOL_TOKENS_V3)
-
-
-def _has_test_anchor(s: str) -> bool:
- """Used by hypothesis_generator.required_tests check — broader than
- _has_real_tool_token because real diagnostic prose uses protocol/concept
- nouns, not only CLI tool names."""
- if not isinstance(s, str):
- return False
- sl = s.lower()
- return any(tok in sl for tok in _HYP_TEST_TOKENS)
-
-
-def _len_close(a, b, ratio: float = 0.5) -> bool:
- """True if abs(a - b) / max(b, 1) ≤ ratio."""
- try:
- return abs(a - b) / max(b, 1) <= ratio
- except Exception:
- return False
-
-
-# Static synonym map for enum drift observed in v7_filt fail analysis
-# (2026-05-21). Lowercase keys, lowercase canonical values. Only true
-# synonyms — NOT semantic miscategorisation (e.g. tcp_retransmission_burst
-# ≠ tcp_duplicate_ack, those are different TCP events and the model is
-# genuinely guessing wrong, not vocab-drifting).
-_ENUM_SYNONYMS = {
- # observation_normalizer.type drift
- "tls_handshake_complete": "tls_handshake_event",
- "tls_handshake": "tls_handshake_event",
- "tls_handshake_start": "tls_handshake_event",
- # protocol version aliases
- "tlsv1.3": "tls 1.3",
- "tls1.3": "tls 1.3",
- "tls/ssl": "tls",
- "tls/ssl/tls": "tls",
-}
-
-
-def _normalize_enum(s: str) -> str:
- """Lowercase + collapse whitespace. Idempotent."""
- if not isinstance(s, str):
- return ""
- return re.sub(r"\s+", " ", s.lower().strip())
-
-
-def _split_compound(s: str) -> set[str]:
- """Split a compound protocol/type like 'TCP/TLS', 'TLS 1.3 over TCP',
- 'SSH/TCP' into atomic parts. Keep parts of length ≥3 to avoid noise."""
- parts = re.split(r"[/\s,]+", _normalize_enum(s))
- return {p for p in parts if len(p) >= 3}
-
-
-def _enum_match(pred: str, allowed: set[str]) -> str | None:
- """Return canonical allowed value matching `pred`, else None. Tiers:
- (1) exact (case-insensitive),
- (2) explicit synonym dict,
- (3) compound subset (TCP/TLS matches TLS),
- (4) substring overlap with ratio ≥0.6 (both ≥6 chars).
- Strict enough to NOT cross truly different enums (TCP ≠ HTTPS, etc)."""
- if not isinstance(pred, str) or not pred:
- return None
- pred_n = _normalize_enum(pred)
- allowed_norm = {_normalize_enum(a): a for a in allowed if isinstance(a, str)}
- if pred_n in allowed_norm:
- return allowed_norm[pred_n]
- # synonym dict
- syn = _ENUM_SYNONYMS.get(pred_n)
- if syn and syn in allowed_norm:
- return allowed_norm[syn]
- # Compound subset — ONLY when pred is a compound (≥2 parts).
- # This means the model expressed multiple layers (e.g. 'TCP/TLS') and we
- # match it against a more atomic allowed value. Crucially, a BARE pred
- # like 'TCP' is NOT allowed to match an allowed compound 'TCP/TLS' — that
- # would let the input-copy baseline (which emits bare 'TCP') game any
- # compound-protocol GT. Asymmetric is intentional.
- pred_parts = _split_compound(pred)
- if len(pred_parts) >= 2:
- for an, orig in allowed_norm.items():
- an_parts = _split_compound(orig)
- if an_parts and an_parts <= pred_parts:
- return orig
- # Substring overlap — both sides must be ≥6 chars to avoid trivial matches.
- for an, orig in allowed_norm.items():
- if len(an) < 6 or len(pred_n) < 6:
- continue
- if an in pred_n or pred_n in an:
- shorter = min(len(an), len(pred_n))
- longer = max(len(an), len(pred_n))
- if shorter / longer >= 0.6:
- return orig
- return None
-
-
-# ---------- per-task v3 graders ----------
-
-def grade_problem_framing_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + scope key overlap with EXPECTED scope."""
- expected = row.get("expected_output", {}) or {}
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- required = ["diagnosis_goal", "scope", "primary_question", "constraints", "initial_uncertainties"]
- missing = [k for k in required if k not in parsed]
- if missing:
- return _result(False, False,
- {k: {"ok": False, "reason": "missing_top_key"} for k in missing},
- head, f"missing_top:{missing}")
-
- fs = {}
- user_req = input_state.get("user_request", "")
-
- # diagnosis_goal: overlap with EXPECTED goal (not just user_req)
- goal = parsed.get("diagnosis_goal", "")
- exp_goal = expected.get("diagnosis_goal", "")
- goal_jacc = _jaccard(goal, exp_goal) if exp_goal else _jaccard(goal, user_req)
- goal_ok = (isinstance(goal, str) and len(goal) >= 30 and not _is_filler(goal)
- and goal_jacc >= 0.15)
- fs["diagnosis_goal"] = {"ok": goal_ok, "format": isinstance(goal, str), "content": goal_ok,
- "reason": f"len={len(goal) if isinstance(goal, str) else 0} jacc_to_exp={goal_jacc:.2f}"}
-
- # scope: dict with ≥ 2 keys. If expected scope present, accept ≥ 1 key match
- # OR value-token Jaccard ≥ 0.10 with expected scope values (synonym keys
- # with similar values still pass — capagent and small student naturally use
- # different key vocab for the same scoping). 2026-05-21: previous ≥50% key
- # overlap was the dominant fail mode (73/79 v7_filt fails), but the real
- # anti-input-copy bite is the diagnosis_goal Jaccard above — scope keys are
- # not a meaningful semantic anchor.
- scope = parsed.get("scope", {})
- exp_scope = expected.get("scope", {})
- if not isinstance(scope, dict):
- fs["scope"] = {"ok": False, "format": False, "content": False, "reason": "not_dict"}
- elif len(scope) < 2:
- fs["scope"] = {"ok": False, "format": True, "content": False,
- "reason": f"too_few_keys={len(scope)}"}
- elif not isinstance(exp_scope, dict) or not exp_scope:
- fs["scope"] = {"ok": True, "format": True, "content": True,
- "reason": f"no_exp keys={list(scope)[:4]}"}
- else:
- exp_keys = set(exp_scope.keys())
- pred_keys = set(scope.keys())
- overlap = exp_keys & pred_keys
- # value-token Jaccard between concatenated pred values and exp values
- pred_val_text = " ".join(str(v) for v in scope.values())
- exp_val_text = " ".join(str(v) for v in exp_scope.values())
- val_jacc = _jaccard(pred_val_text, exp_val_text) if exp_val_text else 0
- ok = len(overlap) >= 1 or val_jacc >= 0.10
- fs["scope"] = {"ok": ok, "format": True, "content": ok,
- "reason": f"key_overlap={len(overlap)}/{len(exp_keys)} val_jacc={val_jacc:.2f}"}
-
- # primary_question
- pq = parsed.get("primary_question", "")
- q_tokens = {"what", "why", "how", "when", "where", "which", "?"}
- pq_ok = (isinstance(pq, str) and len(pq) >= 20 and not _is_filler(pq)
- and (any(t in pq.lower() for t in q_tokens) or "?" in pq))
- fs["primary_question"] = {"ok": pq_ok, "format": isinstance(pq, str), "content": pq_ok,
- "reason": f"len={len(pq) if isinstance(pq, str) else 0}"}
-
- # constraints & uncertainties — substantive list + topical overlap with
- # EXPECTED items (preferred) or user_req as fallback. 2026-05-21: previous
- # version forced topical match against user_req only, but GT and capable
- # models often phrase constraints abstractly ("limited capture point",
- # "single timestamp resolution") whose tokens don't appear in the request.
- user_tokens = _tokens(user_req)
- for k in ["constraints", "initial_uncertainties"]:
- lst = parsed.get(k, [])
- if not isinstance(lst, list) or len(lst) == 0:
- fs[k] = {"ok": False, "format": False, "content": False, "reason": "not_list_or_empty"}
- continue
- substantive = [x for x in lst if isinstance(x, str)
- and len(x) >= 20 and not _is_filler(x)]
- # topical: ≥2-token match with EXPECTED list (synonyms welcome) OR
- # ≥2-token match with user_req as fallback.
- exp_list = expected.get(k, []) or []
- exp_blob = " ".join(x for x in exp_list if isinstance(x, str))
- exp_tokens = _tokens(exp_blob) if exp_blob else set()
- topical = False
- if substantive:
- for x in substantive:
- tx = _tokens(x)
- if exp_tokens and len(tx & exp_tokens) >= 2:
- topical = True; break
- if user_tokens and len(tx & user_tokens) >= 2:
- topical = True; break
- if not (exp_tokens or user_tokens):
- topical = True # no baseline → can't check
- ok = len(substantive) >= 1 and topical
- fs[k] = {"ok": ok, "format": True, "content": ok,
- "reason": f"substantive={len(substantive)} topical={topical}"}
-
- format_ok = all(v.get("format", False) for v in fs.values())
- content_ok = all(v.get("content", False) for v in fs.values())
- return _result(format_ok, content_ok, fs, head)
-
-
-def grade_world_model_builder_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + entity_count ±50% of expected + protocol must match expected."""
- expected = row.get("expected_output", {}) or {}
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- required = ["entities", "flows", "capture_assumptions", "world_model_risks"]
- missing = [k for k in required if k not in parsed]
- if missing:
- return _result(False, False,
- {k: {"ok": False, "reason": "missing_top_key"} for k in missing},
- head, f"missing_top:{missing}")
-
- input_ips = _collect_ips_from_input_state(input_state)
- exp_entities = expected.get("entities", [])
- exp_flows = expected.get("flows", [])
- fs = {}
-
- # entities: valid IP + type enum + count within ±50% of expected
- ents = parsed.get("entities", [])
- if not isinstance(ents, list) or len(ents) < 2:
- fs["entities"] = {"ok": False, "format": False, "content": False, "reason": "not_list_or_too_few"}
- else:
- good = []
- for e in ents:
- if not isinstance(e, dict): continue
- ip = e.get("ip", "")
- etype = e.get("type", "")
- if not _is_ip(ip): continue
- if etype not in {"client", "server", "intermediate", "unknown", "router", "gateway", "proxy", "loadbalancer"}: continue
- good.append(ip)
- n_overlap = len(set(good) & input_ips) if input_ips else 0
- # count-match with expected
- exp_n = len(exp_entities) if isinstance(exp_entities, list) else 2
- count_ok = _len_close(len(good), exp_n, ratio=0.7) # ±70% — capagent variable
- ip_ok = (not input_ips or n_overlap >= 1)
- ok = (len(good) >= 2 and ip_ok and count_ok)
- fs["entities"] = {"ok": ok, "format": len(good) >= 2,
- "content": ip_ok and count_ok,
- "reason": f"good={len(good)} vs exp={exp_n} overlap={n_overlap}/{len(input_ips)} count_ok={count_ok}"}
-
- # flows: valid tuple + protocol must SEMANTICALLY match expected
- # Uses _enum_match (compound split + synonym + substring) so that
- # pred='TCP/TLS' against exp='TLS' or pred='TLSv1.3' against exp='TLS 1.3'
- # passes — 27 of v7_filt's protocol fails are this normalisation pattern.
- exp_protocols = set()
- if isinstance(exp_flows, list):
- for f in exp_flows:
- if isinstance(f, dict):
- p = f.get("protocol", "")
- if isinstance(p, str):
- exp_protocols.add(p)
- flows = parsed.get("flows", [])
- if not isinstance(flows, list) or len(flows) < 1:
- fs["flows"] = {"ok": False, "format": False, "content": False, "reason": "not_list_or_empty"}
- else:
- ok_count = 0
- protocol_match_count = 0
- for f in flows:
- if not isinstance(f, dict): continue
- t = f.get("tuple"); p = f.get("protocol", "")
- t_ok = False
- if isinstance(t, str) and _FLOW_TUPLE_RE.match(t.strip()):
- t_ok = True
- elif isinstance(t, dict):
- src_keys = {"src_ip", "source_ip", "src", "source"}
- dst_keys = {"dst_ip", "destination_ip", "dst", "destination"}
- t_ok = any(k in t for k in src_keys) and any(k in t for k in dst_keys)
- if t_ok and isinstance(p, str) and len(p) >= 2:
- ok_count += 1
- if not exp_protocols or _enum_match(p, exp_protocols) is not None:
- protocol_match_count += 1
- # require ≥1 flow with valid structure AND (if expected has protocols) ≥1 protocol match
- ok = ok_count >= 1 and (not exp_protocols or protocol_match_count >= 1)
- fs["flows"] = {"ok": ok, "format": ok_count >= 1, "content": ok,
- "reason": f"valid={ok_count}/{len(flows)} proto_match={protocol_match_count} exp_proto={list(exp_protocols)[:3]}"}
-
- # capture_assumptions: substantive list
- ca = parsed.get("capture_assumptions", [])
- sub = [x for x in (ca if isinstance(ca, list) else [])
- if isinstance(x, str) and len(x) >= 30 and not _is_filler(x)]
- fs["capture_assumptions"] = {"ok": len(sub) >= 1, "format": isinstance(ca, list),
- "content": len(sub) >= 1, "reason": f"substantive={len(sub)}"}
-
- wmr = parsed.get("world_model_risks", [])
- sub = [x for x in (wmr if isinstance(wmr, list) else [])
- if isinstance(x, str) and len(x) >= 20 and not _is_filler(x)]
- fs["world_model_risks"] = {"ok": len(sub) >= 1, "format": isinstance(wmr, list),
- "content": len(sub) >= 1, "reason": f"substantive={len(sub)}"}
-
- format_ok = all(v.get("format", False) for v in fs.values())
- content_ok = all(v.get("content", False) for v in fs.values())
- return _result(format_ok, content_ok, fs, head)
-
-
-def grade_symptom_detector_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + symptom type MUST be in expected_types set."""
- expected = row.get("expected_output", {}) or {}
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- symptoms = parsed.get("symptoms", [])
- if not isinstance(symptoms, list) or len(symptoms) < 1:
- return _result(False, False, {"symptoms": {"ok": False, "reason": "not_list"}}, head, "not_list")
-
- allowed_types = {
- "tcp_retransmission_burst", "tcp_duplicate_ack", "tcp_zero_window",
- "tcp_reset", "server_response_gap", "client_idle_gap",
- "tls_handshake_delay", "dns_lookup_delay", "dns_no_response",
- "rtt_spike", "high_jitter", "mtu_blackhole",
- }
- sev_allowed = {"low", "medium", "high", "critical"}
-
- # expected types from row's expected_output
- exp_symptoms = expected.get("symptoms", [])
- exp_types = set()
- if isinstance(exp_symptoms, list):
- for s in exp_symptoms:
- if isinstance(s, dict) and s.get("type") in allowed_types:
- exp_types.add(s["type"])
-
- input_text = _input_state_text(input_state).lower()
- ok_count = 0; type_match = 0; bad = []
- for i, s in enumerate(symptoms):
- if not isinstance(s, dict):
- bad.append(f"item{i}_not_dict"); continue
- stype = s.get("type")
- if stype not in allowed_types:
- bad.append(f"item{i}_bad_type:{stype}"); continue
- if exp_types and stype in exp_types:
- type_match += 1
- if s.get("severity") not in sev_allowed:
- bad.append(f"item{i}_bad_severity"); continue
- refs = s.get("evidence_refs", [])
- if not isinstance(refs, list) or len(refs) == 0:
- bad.append(f"item{i}_no_refs"); continue
- sub_refs = [r for r in refs if isinstance(r, str) and len(r) >= 4 and not _is_filler(r)]
- if len(sub_refs) == 0:
- bad.append(f"item{i}_filler_refs"); continue
- ref_tokens = _tokens(" ".join(sub_refs))
- input_tokens = _tokens(input_text)
- if ref_tokens and not (ref_tokens & input_tokens):
- bad.append(f"item{i}_ungrounded_refs"); continue
- ok_count += 1
-
- # require ≥1 valid symptom AND (if expected has types) ≥1 type match
- type_ok = (not exp_types) or (type_match >= 1)
- ok = (ok_count >= 1) and type_ok
- fs = {"symptoms": {"ok": ok, "format": ok_count >= 1, "content": ok,
- "reason": f"valid={ok_count}/{len(symptoms)} type_match={type_match} exp_types={list(exp_types)[:3]} bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_hypothesis_generator_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + n_hyps ±1 of expected + claim Jaccard with EXPECTED claims + narrower tool whitelist."""
- expected = row.get("expected_output", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- hyps = parsed.get("hypotheses", [])
- if not isinstance(hyps, list) or len(hyps) < 2:
- return _result(False, False,
- {"hypotheses": {"ok": False, "format": False, "content": False,
- "reason": f"only_{len(hyps) if isinstance(hyps, list) else 0}"}},
- head, "too_few_hyps")
-
- exp_hyps = expected.get("hypotheses", [])
- exp_n = len(exp_hyps) if isinstance(exp_hyps, list) else 2
- # claims joined as reference text for predicted claim comparison
- exp_claim_text = " ".join(h.get("claim", "") for h in exp_hyps if isinstance(h, dict))
-
- # count: predicted ±1 of expected
- count_ok = abs(len(hyps) - exp_n) <= 1
-
- ok_count = 0; bad = []
- for i, h in enumerate(hyps):
- if not isinstance(h, dict):
- bad.append(f"item{i}_not_dict"); continue
- claim = h.get("claim", "")
- if not isinstance(claim, str) or len(claim) < 40 or _is_filler(claim):
- bad.append(f"item{i}_bad_claim_len={len(claim) if isinstance(claim, str) else 0}"); continue
- # claim should overlap with EXPECTED claims (not just input symptoms)
- if exp_claim_text and _jaccard(claim, exp_claim_text) < 0.10:
- bad.append(f"item{i}_claim_off_topic_vs_expected"); continue
- # Substantive check: capagent may emit strings OR dicts (`{"test": "...", "expected_if_true": "..."}`)
- # Stringify for length/filler check.
- item_ok = True
- for k in ["explains", "does_not_explain", "required_tests"]:
- v = h.get(k, [])
- if not isinstance(v, list) or len(v) < 1:
- bad.append(f"item{i}.{k}_not_list_or_empty"); item_ok = False; break
- sub = []
- for x in v:
- if isinstance(x, str):
- s = x
- elif isinstance(x, dict):
- s = " ".join(str(vv) for vv in x.values() if isinstance(vv, (str, int, float)))
- else:
- continue
- if len(s) >= 15 and not _is_filler(s):
- sub.append(s)
- if len(sub) < 1:
- bad.append(f"item{i}.{k}_no_substantive"); item_ok = False; break
- if not item_ok:
- continue
- # required_tests: stringify whole list and check any network-domain
- # anchor token (protocols, mechanisms, components, or CLI tools).
- # Narrower _has_real_tool_token would reject natural-language diagnostic
- # prose; the input-copy defence for hyp_gen lives in the claim Jaccard
- # check above, not in this whitelist.
- tests = h.get("required_tests", [])
- tests_blob = " ".join(
- x if isinstance(x, str)
- else " ".join(str(v) for v in x.values() if isinstance(v, (str, int, float)))
- for x in tests
- )
- if not _has_test_anchor(tests_blob):
- bad.append(f"item{i}.required_tests_no_tools"); continue
- ok_count += 1
- ok = ok_count >= 2 and count_ok
- fs = {"hypotheses": {"ok": ok, "format": ok_count >= 2, "content": ok,
- "reason": f"good={ok_count}/{len(hyps)} count_ok={count_ok} (pred {len(hyps)} vs exp {exp_n}) bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_test_planner_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + tool semantic equivalence with expected_tool + expected_observations has cross-hypothesis impact."""
- expected = row.get("expected_output", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- na = parsed.get("next_action")
- if not isinstance(na, dict):
- return _result(False, False, {"next_action": {"ok": False, "reason": "missing_or_not_dict"}}, head, "")
-
- exp_na = expected.get("next_action", {}) if isinstance(expected.get("next_action"), dict) else {}
- exp_tool = (exp_na.get("tool") or exp_na.get("tool_name") or "").lower()
-
- tool = (na.get("tool") or na.get("tool_name") or "")
- tool_ok = (isinstance(tool, str) and len(tool) >= 4 and not _is_filler(tool))
- # tool semantic match: substring overlap with expected tool OR shared real-tool token
- tool_match_exp = True
- if exp_tool and tool:
- tool_l = tool.lower()
- # exact match OR shared real-tool token OR ≥4-char shared substring
- shared = (_has_real_tool_token(tool) and _has_real_tool_token(exp_tool)
- and any(tok in tool_l and tok in exp_tool for tok in _PCAP_TOOL_TOKENS_V3))
- tool_match_exp = (tool_l == exp_tool or shared
- or any(tool_l[i:i+5] in exp_tool for i in range(len(tool_l) - 4))
- or any(exp_tool[i:i+5] in tool_l for i in range(len(exp_tool) - 4)))
-
- args = na.get("arguments") or na.get("args") or na.get("params") or {}
- args_ok = isinstance(args, dict) and len(args) >= 1 and not all(_is_filler(str(v)) for v in args.values())
-
- # expected_observations: ≥2 outcomes; at least one outcome must reference
- # a hypothesis id (impact dict OR "h1"/"h2"/"H1" token in any value string).
- eobs = na.get("expected_observations")
- _HYP_RE = re.compile(r"\b[Hh][12345]\b")
-
- def _has_hyp_ref(item) -> bool:
- if isinstance(item, dict):
- if any(k in item for k in ("impact", "hypothesis_impact", "h_impact")):
- return True
- # search all values for hypothesis token
- for v in item.values():
- if isinstance(v, str) and _HYP_RE.search(v):
- return True
- if isinstance(v, dict):
- for vv in v.values():
- if isinstance(vv, str) and _HYP_RE.search(vv):
- return True
- elif isinstance(item, str):
- if _HYP_RE.search(item):
- return True
- return False
-
- if isinstance(eobs, list):
- n = len(eobs)
- has_impact = sum(1 for o in eobs if _has_hyp_ref(o))
- elif isinstance(eobs, dict):
- n = len(eobs)
- has_impact = sum(1 for v in eobs.values() if _has_hyp_ref(v))
- else:
- n = 0; has_impact = 0
- eobs_ok = n >= 2 and has_impact >= 1
-
- ok = tool_ok and tool_match_exp and args_ok and eobs_ok
- fs = {"next_action": {
- "ok": ok, "format": isinstance(na, dict), "content": ok,
- "reason": f"tool={tool[:30]!r} tool_match_exp={tool_match_exp} args_ok={args_ok} eobs_n={n} impact={has_impact}",
- }}
- return _result(ok, ok, fs, head)
-
-
-def grade_observation_normalizer_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + flow_id ∈ expected_flow_ids + type ∈ expected_types."""
- expected = row.get("expected_output", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- obs = parsed.get("observations", [])
- if not isinstance(obs, list) or len(obs) < 1:
- return _result(False, False, {"observations": {"ok": False, "reason": "not_list"}}, head, "")
-
- exp_obs = expected.get("observations", [])
- exp_flow_ids = set()
- exp_types = set()
- if isinstance(exp_obs, list):
- for e in exp_obs:
- if isinstance(e, dict):
- if e.get("flow_id"): exp_flow_ids.add(str(e["flow_id"]))
- if e.get("type"): exp_types.add(str(e["type"]))
-
- ok_count = 0; flow_match_count = 0; type_match_count = 0; bad = []
- for i, o in enumerate(obs):
- if not isinstance(o, dict):
- bad.append(f"item{i}_not_dict"); continue
- otype = o.get("type", "")
- if not isinstance(otype, str) or len(otype) < 3 or _is_filler(otype):
- bad.append(f"item{i}_bad_type"); continue
- flow_id = o.get("flow_id", "")
- if not isinstance(flow_id, str) or len(flow_id) < 2 or _is_filler(flow_id):
- bad.append(f"item{i}_bad_flow_id"); continue
- pref = o.get("packet_ref")
- if pref is None: bad.append(f"item{i}_no_packet_ref"); continue
- pref_ok = isinstance(pref, int) or (isinstance(pref, str) and len(pref) >= 1 and not _is_filler(pref))
- if not pref_ok: bad.append(f"item{i}_bad_packet_ref:{pref!r}"); continue
- t = o.get("time")
- t_ok = False
- if isinstance(t, (int, float)) and not isinstance(t, bool): t_ok = True
- elif isinstance(t, str):
- try: float(t.strip()); t_ok = True
- except: pass
- if not t_ok: bad.append(f"item{i}_bad_time:{t!r}"); continue
- # semantic: at least one of {flow_id, type} must match expected (when expected has any).
- # type uses _enum_match for synonym/substring tolerance — pred='tls_handshake_complete'
- # matches exp='tls_handshake_event' via _ENUM_SYNONYMS (2026-05-21 audit found
- # this as the dominant obs_norm type-drift pattern).
- if exp_flow_ids and flow_id in exp_flow_ids: flow_match_count += 1
- if exp_types and _enum_match(otype, exp_types) is not None: type_match_count += 1
- ok_count += 1
- # require ≥1 valid AND (if expected has types or flow_ids) ≥1 match on type or flow_id
- semantic_ok = True
- if exp_types or exp_flow_ids:
- semantic_ok = (flow_match_count >= 1) or (type_match_count >= 1)
- ok = ok_count >= 1 and semantic_ok
- fs = {"observations": {"ok": ok, "format": ok_count >= 1, "content": ok,
- "reason": f"valid={ok_count}/{len(obs)} flow_match={flow_match_count} type_match={type_match_count} exp_types={list(exp_types)[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_hypothesis_updater_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """v2.1 + status MUST equal expected.status (was missing in v2.1)."""
- expected = row.get("expected_output", {}) or {}
- input_state = row.get("input_state", {}) or {}
- parsed, head = _parse(completion)
- if parsed is None:
- return _result(False, False, {}, head, "no_json")
-
- updates = parsed.get("hypothesis_updates", [])
- if not isinstance(updates, list) or len(updates) < 1:
- return _result(False, False, {"hypothesis_updates": {"ok": False, "reason": "not_list"}}, head, "")
-
- exp_updates = expected.get("hypothesis_updates", [])
- exp_id_to_status = {}
- if isinstance(exp_updates, list):
- for u in exp_updates:
- if isinstance(u, dict) and u.get("id") and u.get("status"):
- exp_id_to_status[u["id"]] = u["status"]
-
- h_input = input_state.get("hypothesis", {})
- expected_id = h_input.get("id") if isinstance(h_input, dict) else None
- input_text = _input_state_text(input_state)
-
- allowed_status = {
- "provisionally_supported", "weakened", "rejected",
- "confirmed", "inconclusive", "unchanged",
- }
-
- ok_count = 0; status_match = 0; bad = []
- for i, u in enumerate(updates):
- if not isinstance(u, dict):
- bad.append(f"item{i}_not_dict"); continue
- uid = u.get("id")
- if expected_id and uid != expected_id:
- bad.append(f"item{i}_id_mismatch:{uid!r}!={expected_id!r}"); continue
- status = u.get("status")
- if status not in allowed_status:
- bad.append(f"item{i}_bad_status:{status!r}"); continue
- # NEW v3: status MUST match expected (if we have an expected status for this id)
- exp_st = exp_id_to_status.get(uid)
- if exp_st and status != exp_st:
- bad.append(f"item{i}_status_mismatch:{status!r}!={exp_st!r}"); continue
- if exp_st:
- status_match += 1
- reason = u.get("reason", "")
- if not isinstance(reason, str) or len(reason) < 20 or _is_filler(reason):
- bad.append(f"item{i}_bad_reason"); continue
- if input_text and _jaccard(reason, input_text) < 0.03:
- bad.append(f"item{i}_reason_off_topic"); continue
- ok_count += 1
- ok = ok_count >= 1
- fs = {"hypothesis_updates": {"ok": ok, "format": ok_count >= 1, "content": ok,
- "reason": f"valid={ok_count}/{len(updates)} status_match={status_match} bad={bad[:3]}"}}
- return _result(ok, ok, fs, head)
-
-
-def grade_stop_controller_v3(row: dict, completion: str) -> tuple[bool, dict]:
- """Unchanged from v2.1 — already does expected enum match."""
- from .graders_v2 import grade_stop_controller_v2
- return grade_stop_controller_v2(row, completion)
-
-
-# ---------- registry ----------
-
-GRADERS_V3: dict[str, Any] = {
- "diag_problem_framing_v3": grade_problem_framing_v3,
- "diag_world_model_builder_v3": grade_world_model_builder_v3,
- "diag_symptom_detector_v3": grade_symptom_detector_v3,
- "diag_hypothesis_generator_v3": grade_hypothesis_generator_v3,
- "diag_test_planner_v3": grade_test_planner_v3,
- "diag_observation_normalizer_v3": grade_observation_normalizer_v3,
- "diag_hypothesis_updater_v3": grade_hypothesis_updater_v3,
- "diag_stop_controller_v3": grade_stop_controller_v3,
-}
-
-V1_TO_V3 = {f"diag_{tid}": f"diag_{tid}_v3" for tid in DIAG_TASKS.keys()}
diff --git a/vendor/nanoai/agentic/diag/schemas.py b/vendor/nanoai/agentic/diag/schemas.py
deleted file mode 100644
index 4c063ae05107b2ca012811887983c12810a705b5..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/schemas.py
+++ /dev/null
@@ -1,341 +0,0 @@
-"""Task specifications for the 8 agentic diagnosis micro-tasks (PCAP domain).
-
-Each TaskSpec is shared by:
- - GT generator (experiments/ksbr_pcap/scripts/gen_diag_eval_gt.py) — capagent gets the schema + 3-shot examples, produces JSON
- - Grader (nanoai.agentic.diag.graders) — validates `format_ok` against required keys/enums and `content_ok` against expected output
-
-Design notes:
- * The schemas are intentionally **minimal** — only fields whose absence/wrongness
- would prevent a downstream agent step from operating. Avoid demanding fields
- the model can't reasonably infer from input_state alone.
- * `freetext=True` flags fields where wording will vary across models; graders
- apply Jaccard overlap (≥0.2 = some semantic signal) rather than exact match.
- * `enum_values` are STRICT — a model that emits a non-enum value scores
- content_ok=False on that field. Enums are chosen to be the natural vocab a
- domain-trained model would use.
-"""
-from __future__ import annotations
-
-from dataclasses import dataclass, field
-from typing import Literal
-
-
-FieldKind = Literal["str", "int", "float", "enum", "list", "list_obj", "obj"]
-
-
-@dataclass
-class FieldSpec:
- name: str
- kind: FieldKind
- required: bool = True
- # for kind=enum
- enum_values: list[str] | None = None
- # for kind=list / list_obj: minimum item count for the field to count as "present"
- min_items: int = 1
- # for kind=list_obj: required keys on each item
- item_keys: list[str] | None = None
- # for kind=list_obj: per-item enum constraints (key → allowed values)
- item_enums: dict[str, list[str]] | None = None
- # mark string fields where the model's exact wording will vary;
- # grader uses Jaccard token overlap instead of equality
- freetext: bool = False
-
-
-@dataclass
-class TaskSpec:
- task_id: str # short slug, matches eval row's micro_task
- description: str # one-line human description
- input_keys: list[str] # keys expected in input_state JSON
- output_fields: list[FieldSpec] # what target output must contain
- prompt_intro: str # plain-NL framing prepended to JSON prompt
- seed_scenarios: list[str] = field(default_factory=list) # PCAP scenarios for GT generation
-
-
-# ---------- 8 task definitions ----------
-
-PROBLEM_FRAMING = TaskSpec(
- task_id="problem_framing",
- description="Given a vague user request, produce a scoped diagnosis goal.",
- input_keys=["user_request", "available_data", "initial_context"],
- output_fields=[
- FieldSpec("diagnosis_goal", "str", freetext=True),
- FieldSpec("scope", "obj"), # dict; we only check it exists + is dict
- FieldSpec("primary_question", "str", freetext=True),
- FieldSpec("constraints", "list", min_items=1),
- FieldSpec("initial_uncertainties", "list", min_items=1),
- ],
- prompt_intro=(
- "you are a network-diagnosis planner. given a user request and the "
- "data they have, produce a scoped diagnosis plan as JSON. think about "
- "what we *don't yet know* — that's the most important field."
- ),
- seed_scenarios=[
- "user complains a specific HTTPS API is slow at 10:00-10:15, has pcap from the client side only",
- "user says SSH connections drop intermittently, has pcap from a tap port between client and server",
- "user reports DNS lookups failing randomly for a list of domains, has pcap from the resolver",
- "user observes TCP RSTs during file uploads to S3, has client-side pcap",
- "user notices video conferencing call quality degrades after 10 min, capture point unknown",
- "user finds requests timing out for one specific backend, has tap-port pcap server-side",
- "user reports the database connection pool exhaustion alerts during peak hours, pcap from the app tier",
- "user sees high TLS handshake latency from one client subnet, pcap from server edge",
- ],
-)
-
-WORLD_MODEL_BUILDER = TaskSpec(
- task_id="world_model_builder",
- description="From packet/flow summary, build the entity-flow-assumption world model.",
- input_keys=["packet_summary", "flow_summary", "capture_metadata"],
- output_fields=[
- FieldSpec("entities", "list_obj", min_items=2,
- item_keys=["id", "type", "ip"],
- item_enums={"type": ["client", "server", "intermediate", "unknown"]}),
- FieldSpec("flows", "list_obj", min_items=1,
- item_keys=["id", "tuple", "protocol"]),
- FieldSpec("capture_assumptions", "list", min_items=1),
- FieldSpec("world_model_risks", "list", min_items=1),
- ],
- prompt_intro=(
- "you are a pcap world-model builder. parse a packet/flow summary into "
- "entities (client/server/intermediate), flows (5-tuple + protocol), and "
- "explicit assumptions about the capture point. also list risks that could "
- "make this model wrong (e.g. NAT, asymmetric capture, TLS hiding payload)."
- ),
- seed_scenarios=[
- "summary of HTTPS flow, client 192.168.1.20 → server 10.1.2.3:443, ~127 packets",
- "summary of bidirectional SSH flow with several retransmissions and dup ACKs",
- "summary of DNS UDP/53 query + response flow with one query showing no response",
- "summary of TCP flow ending in RST from server side during a long upload",
- "summary of unidirectional capture showing only one direction of an SMTP exchange",
- "summary of TLS 1.3 handshake with two flows from same client IP (port pair changed)",
- "summary of multiple HTTPS flows from a CDN edge with shared server IP",
- "summary of port-scan-like traffic to many destinations from one source IP",
- ],
-)
-
-SYMPTOM_DETECTOR = TaskSpec(
- task_id="symptom_detector",
- description="Identify symptoms in flow/event data WITHOUT explaining cause.",
- input_keys=["flow_events", "tcp_analysis", "time_window"],
- output_fields=[
- FieldSpec("symptoms", "list_obj", min_items=1,
- item_keys=["type", "flow_id", "time_range", "severity", "evidence_refs"],
- item_enums={
- "type": [
- "tcp_retransmission_burst", "tcp_duplicate_ack", "tcp_zero_window",
- "tcp_reset", "server_response_gap", "client_idle_gap",
- "tls_handshake_delay", "dns_lookup_delay", "dns_no_response",
- "rtt_spike", "high_jitter", "mtu_blackhole",
- ],
- "severity": ["low", "medium", "high", "critical"],
- }),
- ],
- prompt_intro=(
- "you are a pcap symptom detector. produce ONLY observed phenomena — do "
- "NOT explain cause. e.g. 'retransmission burst from 10:03:11 to 10:03:13' "
- "is a symptom; 'network is congested' is a hypothesis (don't write that). "
- "every symptom must cite evidence_refs to specific events from the input."
- ),
- seed_scenarios=[
- "flow events showing 15 retransmissions clustered in a 2-second window inside a 5-minute flow",
- "tcp analysis showing 3 zero-window advertisements from server during file upload",
- "flow events showing a 2300ms gap between request bytes finishing and first response byte",
- "dns events showing 4/40 queries received no response within 5s timeout",
- "tcp analysis showing alternating SYN-ACK then RST without ESTABLISHED state",
- "tls events showing client hello → server hello gap of 1.8 seconds (typical: <100ms)",
- "tcp events showing duplicate ACKs every 200ms with no progress in seq numbers",
- "flow events showing RTT jumping from baseline 5ms to 250ms intermittently",
- ],
-)
-
-HYPOTHESIS_GENERATOR = TaskSpec(
- task_id="hypothesis_generator",
- description="Given symptoms + world model + any prior hypotheses, propose OR refine candidate root-cause hypotheses.",
- input_keys=["world_model", "symptoms", "diagnosis_goal", "prior_hypotheses"],
- output_fields=[
- FieldSpec("hypotheses", "list_obj", min_items=2,
- item_keys=["id", "claim", "explains", "does_not_explain", "required_tests"]),
- ],
- prompt_intro=(
- "you are a pcap hypothesis generator running inside a multi-turn diagnosis "
- "loop. given world_model + symptoms + diagnosis_goal AND any prior_hypotheses "
- "from earlier turns, produce the CURRENT hypothesis list.\n"
- " - if prior_hypotheses is empty (turn 1): generate 2-3 fresh competing root "
- " causes. assign stable ids (H1, H2, H3).\n"
- " - if prior_hypotheses is non-empty (turn 2+): REFINE rather than regenerate. "
- " KEEP the original id of any hypothesis that remains plausible given new "
- " symptoms; UPDATE its claim/required_tests if evidence has narrowed it; "
- " ADD a new hypothesis ONLY if symptoms reveal a path the prior list doesn't "
- " explain; DROP hypotheses marked status=rejected by hypothesis_updater. "
- " NEVER renumber stable hypotheses — H1 stays H1 across turns. NEVER reset "
- " the slate just because new symptoms arrived; build on prior reasoning.\n"
- "EVERY hypothesis must include `required_tests` that could falsify it. "
- "each `claim` must be specific (e.g. 'server application is the bottleneck, "
- "not network' not just 'application issue'). do NOT pick a single root cause — "
- "produce ≥2 competing hypotheses unless prior_hypotheses is already down to 1 "
- "confirmed.\n"
- "\n"
- "FIELD QUALITY (strict — rows violating these are dropped at gen-time):\n"
- " - `explains`: list of ≥2 FULL-SENTENCE descriptions of what the hypothesis "
- " accounts for. Each item must be a complete English clause (≥15 chars), "
- " not a bare id like 'evt_1' or 'S2'. EXAMPLE good: 'Retransmissions cluster "
- " during the gap window because the server stops ACK-ing'. EXAMPLE bad: "
- " 'evt_3' or 'S2'.\n"
- " - `does_not_explain`: list of ≥1 FULL-SENTENCE description of what this "
- " hypothesis FAILS to account for. NEVER emit empty []; every honest "
- " hypothesis has gaps. If you can't think of one, the hypothesis is too "
- " broad — narrow it. EXAMPLE good: 'Does not explain why baseline RTT is "
- " unaffected outside the gap window'. EXAMPLE bad: empty list or just "
- " 'Why X' (incomplete).\n"
- " - `required_tests`: list of ≥2 FULL-SENTENCE testable actions. Each must "
- " name a concrete diagnostic mechanism (tool, protocol, mechanism). "
- " EXAMPLE good: 'Run tshark with filter ip.addr == X to capture handshake "
- " timing'. EXAMPLE bad: 'check it' or 'verify'."
- ),
- seed_scenarios=[
- # Turn-1 scenarios (empty prior_hypotheses):
- "symptom: 2300ms server response gap with NO tcp anomalies in the gap window",
- "symptom: retransmission burst + duplicate ACKs concentrated in delay windows",
- "symptom: TLS handshake delay 1.8s on one specific client subnet only",
- "symptom: DNS resolution failures clustered to one specific upstream resolver",
- "symptom: zero-window advertisements from server only during large upload bursts",
- "symptom: tcp RSTs from server side mid-flow after ~30s of activity",
- "symptom: MTU blackhole-like behavior (small packets ok, large get dropped)",
- "symptom: high jitter + occasional rtt spikes correlated with backup job times",
- # Turn-N>=2 scenarios (non-empty prior_hypotheses — refinement):
- "prior: [H1 server-bottleneck status=provisionally_supported, H2 network-congestion status=weakened] + new evidence: server CPU stays <30% during gap → DROP H1, ADD H3 about middlebox processing",
- "prior: [H1 TLS-handshake-misconfig, H2 firewall-rate-limit] + new evidence: TLS handshake completes normally for control flow → MARK H1 weakened, KEEP H2, ADD nothing — refinement only",
- "prior: [H1 path-MTU-blackhole status=provisionally_supported] + new evidence: small-packet pings also fail intermittently → BROADEN H1 claim to 'path quality degraded', keep id",
- "prior: [H1 dns-resolver-failure status=rejected, H2 client-stub-misconfig status=provisionally_supported] + new evidence: another resolver also fails → DROP H1 (already rejected), keep H2, ADD H3 about client name resolution policy",
- "prior: [H1 server-app-bottleneck status=confirmed] + new evidence: no contradicting observations after 3 tests → KEEP H1 only, downstream stop_controller should fire",
- "prior: [H1, H2, H3 all status=inconclusive] + new evidence: a 4th retransmission burst in a new flow → consider whether one prior hyp now explains the broader pattern; refine the strongest claim",
- ],
-)
-
-TEST_PLANNER = TaskSpec(
- task_id="test_planner",
- description="Choose the next-best tool action to validate/falsify a hypothesis.",
- input_keys=["hypotheses", "evidence_so_far", "available_tools"],
- output_fields=[
- FieldSpec("next_action", "obj"), # we'll require nested keys via custom check in grader
- ],
- prompt_intro=(
- "you are a pcap test planner. choose the SINGLE next action with the highest "
- "expected information gain — i.e. it should EITHER strongly support OR strongly "
- "weaken at least one hypothesis. the action must include a concrete tool command "
- "and explicit `expected_observations` mapping outcomes to hypothesis support/weaken."
- ),
- seed_scenarios=[
- "two hypotheses: app-delay vs network-loss; need test to distinguish them inside the delay window",
- "two hypotheses: server CPU bottleneck vs downstream call delay; need a test using only pcap",
- "two hypotheses: client-side firewall vs path MTU; need a test that disambiguates",
- "two hypotheses: TLS-side delay vs TCP-side delay; need finer-grained tshark timing",
- "two hypotheses: server overload vs client retransmit storm; need a test pinning direction",
- "two hypotheses: DNS resolver failure vs network reachability to resolver; need a distinguishing test",
- "three hypotheses for periodic latency spikes: backup, GC pause, cron job; one tool query each",
- "two hypotheses about RST source (server-app vs middlebox); need a test",
- ],
-)
-
-OBSERVATION_NORMALIZER = TaskSpec(
- task_id="observation_normalizer",
- description="Convert raw tool output (tshark fields, zeek logs) into structured observations.",
- input_keys=["raw_output", "tool", "world_model"],
- output_fields=[
- FieldSpec("observations", "list_obj", min_items=1,
- item_keys=["id", "type", "flow_id", "time", "packet_ref"]),
- ],
- prompt_intro=(
- "you are a pcap observation normalizer. take raw tool output (e.g. tshark "
- "tab-separated fields) and produce structured observations. EVERY field "
- "(time, packet_ref, raw_fields) must come from the input — do not invent. "
- "preserve raw values; the agent loop relies on packet_ref being verifiable."
- ),
- seed_scenarios=[
- "tshark output with frame.time_epoch, ip.src, ip.dst, tcp.seq, tcp.ack for retransmissions",
- "tshark output with frame.number, tls.handshake.type for a TLS exchange",
- "zeek conn.log entries showing flow tuples + duration + bytes",
- "tshark output with dns.qry.name, dns.flags.response, dns.flags.rcode",
- "tshark output with tcp.analysis.flags marking RTOs, duplicate ACKs, retransmissions",
- "tshark output with tcp.window_size_value showing zero-window advertisements",
- "zeek dns.log with query / response pairs and timing",
- "tshark output with frame.time_relative and tcp.len for response gap analysis",
- ],
-)
-
-HYPOTHESIS_UPDATER = TaskSpec(
- task_id="hypothesis_updater",
- description="Update hypothesis status given new observation + prediction.",
- input_keys=["hypothesis", "prediction", "observation"],
- output_fields=[
- FieldSpec("hypothesis_updates", "list_obj", min_items=1,
- item_keys=["id", "status", "reason"],
- item_enums={"status": [
- "provisionally_supported", "weakened", "rejected",
- "confirmed", "inconclusive", "unchanged",
- ]}),
- ],
- prompt_intro=(
- "you are a pcap hypothesis updater. given a hypothesis with its prediction "
- "and a new observation, decide if observation **supports**, **weakens**, or "
- "**rejects** the hypothesis. CRITICAL: be willing to weaken/reject — do not "
- "pretend partial evidence is confirmation. `reason` must cite the specific "
- "observation field that drove the change."
- ),
- seed_scenarios=[
- "hyp: network-loss is cause; pred: retransmissions in delay window; obs: no retransmissions",
- "hyp: app-delay is cause; pred: tcp clean during gap; obs: tcp clean during gap",
- "hyp: client firewall blocks; pred: SYN goes out, no SYN-ACK; obs: SYN+SYN-ACK, no app data",
- "hyp: dns failing; pred: queries with no response; obs: response received but RCODE=SERVFAIL",
- "hyp: MTU blackhole; pred: small packets ok, large drop; obs: ALL sizes get through fine",
- "hyp: server overload; pred: response gap on most flows; obs: gap on only 5%",
- "hyp: tls negotiation issue; pred: hello-to-hello slow; obs: hello-to-hello fast, finished slow",
- "hyp: jitter from backup job; pred: spike pattern at scheduled times; obs: spike unrelated to schedule",
- ],
-)
-
-STOP_CONTROLLER = TaskSpec(
- task_id="stop_controller",
- description="Decide whether to continue diagnosis, stop with conclusion, or stop inconclusive.",
- input_keys=["current_hypotheses", "evidence_summary", "tools_used", "uncertainty"],
- output_fields=[
- FieldSpec("stop_decision", "enum",
- enum_values=["continue", "stop_with_conclusion", "stop_inconclusive"]),
- FieldSpec("reason", "str", freetext=True),
- FieldSpec("required_before_stop", "list", min_items=0), # may be empty if stop=stop_*
- ],
- prompt_intro=(
- "you are a pcap stop controller. given the diagnosis state, output ONE of: "
- "(a) continue if a clear next test would meaningfully narrow uncertainty, "
- "(b) stop_with_conclusion if remaining uncertainty is acceptable and one "
- "hypothesis dominates, (c) stop_inconclusive if available data cannot "
- "distinguish remaining hypotheses. do NOT keep digging when the data has "
- "been exhausted — over-digging produces false root causes. EQUALLY, do "
- "NOT conclude when one cheap test would meaningfully narrow uncertainty — "
- "premature conclusion produces false root causes too. When in doubt about "
- "whether enough evidence has been collected, default to stop_inconclusive "
- "rather than stop_with_conclusion: a hedged answer is better than a wrong one."
- ),
- seed_scenarios=[
- "two hyps still competing 60/40, one cheap remaining test would distinguish",
- "one hyp at 0.85 confidence, others at <0.1, available data exhausted",
- "all hyps still <0.5 confidence, pcap is one-sided, no more reachable evidence",
- "single hyp confirmed across 5 flows, no contradicting evidence",
- "two hyps equally weighted, all remaining tools would yield same evidence already seen",
- "one hyp supported on 3/10 sampled flows but rejected on 7/10 — likely incomplete",
- "loop: 3 actions ran, no hyp confidence has moved at all in last 2 rounds",
- "investigation has run 8 actions; cost ceiling reached; current state still ambiguous",
- ],
-)
-
-
-DIAG_TASKS: dict[str, TaskSpec] = {
- spec.task_id: spec
- for spec in [
- PROBLEM_FRAMING, WORLD_MODEL_BUILDER, SYMPTOM_DETECTOR, HYPOTHESIS_GENERATOR,
- TEST_PLANNER, OBSERVATION_NORMALIZER, HYPOTHESIS_UPDATER, STOP_CONTROLLER,
- ]
-}
-
-
-def all_task_ids() -> list[str]:
- return list(DIAG_TASKS.keys())
diff --git a/vendor/nanoai/agentic/diag/train_scenarios.py b/vendor/nanoai/agentic/diag/train_scenarios.py
deleted file mode 100644
index 29dbb2d78f4d98ed4c25d12e4a2c6deeb560b8f5..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/diag/train_scenarios.py
+++ /dev/null
@@ -1,341 +0,0 @@
-"""Train-only scenario seeds for diag_sft_v1 SFT data generation.
-
-DISJOINT from the eval `seed_scenarios` in schemas.py. The split prevents
-the model from memorizing the exact eval prompts during SFT — a future
-diag_eval_v1 measurement on a model trained from these scenarios is then
-testing schema generalization rather than scenario recall.
-
-Pairing: `gen_diag_sft.py` reads from TRAIN_SCENARIOS, `gen_diag_eval_gt.py`
-reads from schemas.py:seed_scenarios. The 24-row `_heldout.jsonl`
-remains the ultimate contamination gate (uses eval seeds, never trained on).
-
-Convention: **32 scenarios per task** (expanded from initial 16 on 2026-05-21
-to improve per-scenario diversity at 50k+ scale gen). Semantically distinct
-families from the 8 eval scenarios. Capagent will randomly sample with
-replacement at gen time, so total unique inputs per task scale to ~250-400
-with model temperature variation.
-"""
-from __future__ import annotations
-
-
-TRAIN_SCENARIOS: dict[str, list[str]] = {
- "problem_framing": [
- # Original 16 (2026-05-20)
- "user reports TLS-terminated reverse proxy occasionally returns 502 to mobile clients only, has edge-pcap",
- "user observes WebSocket connections disconnect after exactly 60s of idle, has client-side pcap",
- "user notices a Kafka consumer group rebalance storm at midnight UTC daily, has tap pcap from broker side",
- "user reports SMTP submissions hang for 30s before completing, pcap from MTA edge available",
- "user sees gRPC unary RPCs timeout in p99 only, pcap from both client and server LB",
- "user complains BitTorrent-like long-lived connections die after firewall reload, edge pcap",
- "user notes IoT MQTT brokers see PINGRESP delays during firmware updates, broker-side pcap",
- "user reports VPN tunnel periodically drops single packets, capture from internal endpoint",
- "user sees QUIC fallback to TCP for a subset of clients behind one ISP, edge-pcap",
- "user observes Redis Sentinel failover taking 18s on a 3-node cluster, broker tap-port pcap",
- "user reports SNMP polling missing one MIB tree every other poll, no obvious pattern, mgmt vlan pcap",
- "user complains LDAP STARTTLS rejects from one client, pcap available from LDAP server side",
- "user observes BGP session flapping between two routers in same AS, both router-tap pcaps",
- "user reports RTSP stream pause for 8s every 15 min, pcap from camera side only",
- "user notes container egress traffic going to wrong NAT pool intermittently, node-host pcap",
- "user sees iSCSI initiator lose path during snapshot, target-side pcap available",
- # Expansion +16 (2026-05-21)
- "user notices Memcached connections leak file descriptors over 48h, no obvious failure, server-side pcap",
- "user reports IPSec tunnel rekey takes 12s causing data plane outage, both peers pcap",
- "user observes WebRTC video glitches when remote peer moves through cell handoff, mobile pcap",
- "user notes Envoy sidecar adds 30ms p99 latency only for inbound traffic, sidecar-host pcap",
- "user reports periodic Cassandra gossip storm causes node flapping, cluster-internal pcap",
- "user complains S3 multipart upload fails at part 3 of 10 reproducibly, client pcap",
- "user observes JMS message delivery skipping queue order during failover, broker tap",
- "user reports NTP offset drifts 200ms during business hours, NTP server-side pcap",
- "user notes Mongo replica set election takes 45s on a 5-node cluster, primary-side pcap",
- "user sees DoH (DNS-over-HTTPS) queries fail to one provider only, browser-host pcap",
- "user reports Modbus TCP slave responses delayed during high temperature, slave-side pcap",
- "user observes Tor circuit setup taking 12s on a specific exit selection, client pcap",
- "user notes Postgres logical replication slot lag growing during peak, subscriber pcap",
- "user reports SAP RFC calls fail with timeout to one application server, mtu pcap",
- "user complains video doorbell loses live feed after exact 60-min sessions, doorbell-side pcap",
- "user observes RDP session disconnects when bandwidth drops below 256kbps, edge pcap",
- ],
- "world_model_builder": [
- # Original 16
- "summary of dual-stack IPv4+IPv6 web traffic to same hostname, AAAA preferred but TCP-RST'd",
- "summary of MQTT-over-WSS flow from IoT gateway through L7 LB to broker cluster",
- "summary of TCP traffic showing PROXY-protocol prefix headers before TLS",
- "summary of QUIC v1 traffic with version negotiation between client and server",
- "summary of multipath TCP (MPTCP) flow with two subflows from same client",
- "summary of ICMP TimeExceeded responses interleaved with TCP application traffic",
- "summary of LACP-bundled traffic captured on one of the two member interfaces only",
- "summary of GRE-encapsulated TCP showing inner and outer 5-tuples both visible",
- "summary of SCTP multihomed association between two endpoints with two IPs each",
- "summary of VXLAN-encapsulated east-west traffic captured at the spine switch",
- "summary of TCP keepalive probes on a long-idle BGP session every 60s",
- "summary of TFTP UDP flow with port renegotiation after WRQ acknowledgment",
- "summary of PIM-SM multicast control plus IGMP joins from receivers",
- "summary of SCCP signaling between IP phone and call manager with frequent registrations",
- "summary of TCP flow with selective ACK options heavily used during loss recovery",
- "summary of WireGuard UDP traffic (no L7 visibility) between two peer endpoints",
- # Expansion +16
- "summary of HTTP/3 over QUIC v2 traffic with version negotiation between client and CDN edge",
- "summary of SRv6 (segment routing IPv6) ingress traffic with multiple SIDs in IPv6 header",
- "summary of MPLS L3VPN traffic showing label stack with three labels on transit packets",
- "summary of OSPFv3 IPv6 hello packets between two routers in same area",
- "summary of EVPN MAC advertisement BGP UPDATEs between two leaf switches",
- "summary of CoAP UDP requests over DTLS to multiple IoT devices behind same gateway",
- "summary of WebRTC ICE STUN binding requests during connection candidate gathering",
- "summary of mTLS-protected gRPC streaming bidi between sidecar proxies",
- "summary of EAP-TLS supplicant authentication on 802.1X-enabled switch port",
- "summary of BFD (bidirectional forwarding detection) control packets between two routers",
- "summary of NVMe-oF over RDMA traffic between initiator and target on RoCEv2",
- "summary of WebSocket connections to a chat backend with per-message DEFLATE compression",
- "summary of LACP control packet exchange between switch and server team interfaces",
- "summary of Profinet IO cyclic data exchange between PLC and field device",
- "summary of Tailscale WireGuard mesh with multiple peers on same magic DNS",
- "summary of SRTP-encrypted RTP from WebRTC video conf, no L7 visibility but timing patterns",
- ],
- "symptom_detector": [
- # Original 16
- "events showing 8 RTOs back-to-back within 1.5s on a single flow during file transfer",
- "tcp analysis showing SACK blocks growing across 12 packets without forward progress",
- "events showing TLS Alert level=fatal description=handshake_failure mid-handshake",
- "tcp analysis showing window-scaling factor of 0 used after expected 7 — degraded throughput",
- "events showing 200ms gap between server FIN and client FIN-ACK on closing handshake",
- "tcp analysis flagging out-of-order segments persisting 400ms before reorder buffer drains",
- "events showing HTTP/2 GOAWAY frame from server with error_code=ENHANCE_YOUR_CALM",
- "tcp analysis showing PUSH+FIN combined flag on a single segment from a server response",
- "events showing 3 consecutive ICMP fragmentation-needed messages followed by TCP stalls",
- "tcp analysis showing packets out of seq number range (likely packet injection or replay)",
- "events showing UDP datagrams arriving after 800ms+ delay (jitter) on a real-time flow",
- "tcp analysis showing keepalive probes failing 3x in a row before connection RST",
- "events showing TLS resumption attempt rejected with new full handshake every connection",
- "tcp analysis showing client window shrinking to 1024 mid-flow (window collapse)",
- "events showing 5 SYN retries with 1s/2s/4s/8s/16s spacing — classic SYN backoff",
- "tcp analysis showing simultaneous open (both sides send SYN nearly at once)",
- # Expansion +16
- "events showing TLS Alert level=warning description=user_canceled during handshake",
- "tcp analysis showing RST with non-zero ACK from middlebox not matching server side",
- "events showing HTTP/2 PING frame with payload mismatch between request and response",
- "tcp analysis showing 8 consecutive zero-window updates within 50ms (rapid window collapse)",
- "events showing IPv6 NDP redirect messages directing traffic away from default GW",
- "tcp analysis showing congestion window collapse to 1 MSS after 3 dup-ACKs",
- "events showing TLS ClientHello with no SNI extension from updated browser client",
- "tcp analysis flagging timestamp option (TSecr=0) after established connection (TS option dropped)",
- "events showing HTTP/2 GOAWAY with last_stream_id=0 right after settings exchange",
- "tcp analysis showing PSH+ACK alternating with FIN+ACK from same endpoint mid-flow",
- "events showing IPSec ESP sequence number wraparound triggering rekey",
- "tcp analysis showing SACK ranges overlap (likely buggy SACK implementation)",
- "events showing UDP checksum errors clustered around a specific interface ingress",
- "tcp analysis showing window scale = 14 used (max RFC value, suspicious memory)",
- "events showing HTTP/3 STREAM frame with FIN bit set on stream 0 (control stream)",
- "tcp analysis showing PAWS (protection against wrapped sequences) check failures",
- ],
- "hypothesis_generator": [
- # Original 16
- "symptom: HTTP/2 GOAWAY ENHANCE_YOUR_CALM after burst of 50 requests on one connection",
- "symptom: TLS resumption tickets rejected, full handshake every time on subset of clients",
- "symptom: SCTP heartbeat timeout on one of two paths in a multihomed association",
- "symptom: BGP UPDATE messages with malformed AS_PATH attribute appearing periodically",
- "symptom: MTU-related packet drops only on traffic transiting one specific WAN link",
- "symptom: QUIC handshake taking 3x normal RTT — possible amplification protection trigger",
- "symptom: PROXY protocol header missing on some connections to backend server",
- "symptom: HTTP/2 RST_STREAM stream-level cancellations spiking from one client",
- "symptom: NFS RPC retransmissions clustered around storage failover events",
- "symptom: TCP simultaneous open events appearing on internal east-west traffic",
- "symptom: TLS SNI mismatch between hostname in HTTP host header and TLS cert CN",
- "symptom: VXLAN inner-flow drops only when inner-MTU exceeds carrier-MTU minus encap",
- "symptom: HTTP/3 (QUIC) connections falling back to HTTP/2 for clients behind one ISP",
- "symptom: SSH client SSH-2.0 banner exchange completes but auth phase times out",
- "symptom: Redis MULTI/EXEC blocks taking 3s where 50ms is baseline",
- "symptom: TCP RST from server with TTL value different from established traffic",
- # Expansion +16
- "symptom: BGP UPDATE messages stuck behind one peer's TCP MSS issue",
- "symptom: TLS clients failing only when CDN edge serves session ticket from peer",
- "symptom: WebSocket connections fail to upgrade for traffic transiting one load balancer",
- "symptom: QUIC 0-RTT data being rejected by server even when ticket is valid",
- "symptom: NTP polling falls back to text protocol over UDP/123 unexpectedly",
- "symptom: gRPC bidi streams cancelled after 30s of idle by middle proxy",
- "symptom: MPLS-VPN packets dropped at PE router only for traffic with specific DSCP",
- "symptom: Kafka producer batches timing out when broker replica leader changes",
- "symptom: DTLS-protected SIP signaling failing for clients behind one CGNAT",
- "symptom: TCP ECN markings stripped on traffic transiting one ISP link",
- "symptom: HTTP/2 server push frames being rejected with REFUSED_STREAM",
- "symptom: SCTP path failure between two endpoints despite both interfaces UP",
- "symptom: Redis pub/sub messages reordered when client reconnects to different replica",
- "symptom: DNS-over-HTTPS responses delayed by 1.5s to specific resolver",
- "symptom: TCP zero-copy send failing on Linux 6.x with specific BPF program loaded",
- "symptom: TLS 1.3 EarlyData stream getting dropped at HTTP/2 ALPN negotiation",
- # Turn-N refinement scenarios (non-empty prior_hypotheses) — teach carry-over:
- "turn-N refinement. prior_hypotheses=[H1 server-app-bottleneck status=provisionally_supported, H2 network-congestion status=weakened]. new evidence: server CPU stayed <30% throughout the gap window. Refine: drop H1, keep H2 but downgrade further, add H3 about middlebox/inspection delay",
- "turn-N refinement. prior_hypotheses=[H1 TLS-handshake-misconfig, H2 firewall-rate-limit status=open]. new evidence: TLS handshake completes normally for the control flow. Refine: mark H1 weakened (keep id), keep H2, add nothing",
- "turn-N refinement. prior_hypotheses=[H1 path-MTU-blackhole status=provisionally_supported]. new evidence: small-packet ICMP pings now also failing intermittently. Refine: broaden H1's claim to 'path quality degraded across packet sizes', keep id H1",
- "turn-N refinement. prior_hypotheses=[H1 dns-resolver-failure status=rejected, H2 client-stub-misconfig status=provisionally_supported]. new evidence: another fallback resolver also fails. Refine: H1 stays dropped, keep H2, add H3 about client name-resolution policy or hosts file",
- "turn-N refinement. prior_hypotheses=[H1 server-app-bottleneck status=confirmed, H2 network-congestion status=rejected, H3 client-side status=rejected]. new evidence: no contradiction in 3 follow-up tests. Refine: keep H1 only as the surviving hypothesis (chain ready for stop_with_conclusion)",
- "turn-N refinement. prior_hypotheses=[H1 inconclusive, H2 inconclusive, H3 inconclusive]. new evidence: a 4th retransmission burst on a brand new flow with same client subnet. Refine: tighten the strongest existing claim to be subnet-specific; do NOT add yet another hypothesis",
- "turn-N refinement. prior_hypotheses=[H1 server-TLS-stack-bug status=weakened, H2 client-TLS-stack-bug status=provisionally_supported]. new evidence: ServerHello timing now within 20ms across all flows. Refine: drop H1 (now fully contradicted), keep H2, add H3 about TLS extension handling",
- "turn-N refinement. prior_hypotheses=[H1 MTU-blackhole status=provisionally_supported]. new evidence: large packets now succeed across multiple flows after middlebox MTU was raised. Refine: H1 status promotes towards confirmed; keep id; chain ready to stop_with_conclusion next turn",
- "turn-N refinement. prior_hypotheses=[H1 DNS-cache-poisoning, H2 client-resolver-misconfig]. new evidence: dig from different resolver returns expected IP. Refine: keep H1 (still possible), mark H2 weakened, add H3 about authoritative-server inconsistency",
- "turn-N refinement. prior_hypotheses=[H1 firewall-state-table-eviction, H2 server-explicit-close]. new evidence: server logs show no application-level close, RSTs arrive even during active data transfer. Refine: tighten H1 to specify middlebox; mark H2 weakened",
- "turn-N refinement. prior_hypotheses=[H1 server-side-thread-starvation status=open]. new evidence: only this one hypothesis remains plausible and is supported by 4 independent observations. Refine: confirm H1 (status=confirmed), keep id, downstream should fire stop_with_conclusion",
- "turn-N refinement. prior_hypotheses=[H1, H2 both inconclusive after 2 turns]. new evidence: a third tool run also gave ambiguous output that matches both H1 and H2 equally. Refine: keep both hypotheses unchanged but log that data exhaustion is approaching — stop_inconclusive is becoming the right next step",
- "turn-N refinement. prior_hypotheses=[H1 IPsec-MTU status=weakened, H2 carrier-fragmentation status=provisionally_supported]. new evidence: large packets fail even on flows that don't traverse the IPsec tunnel. Refine: drop H1 (now contradicted), keep H2, add H3 about endpoint NIC offload bug",
- "turn-N refinement. prior_hypotheses=[H1, H2, H3]. new evidence: a measurement error was found in the earlier observation that supported H3. Refine: revert H3 to status=open, do not drop or confirm; keep H1 and H2",
- "turn-N refinement. prior_hypotheses=[H1 application-thread-pool-exhaustion status=provisionally_supported]. new evidence: thread-pool metrics from server actually look healthy, but a new symptom of GC pause spikes appeared. Refine: weaken H1, add H4 about garbage collector pauses correlating with the gap windows",
- "turn-N refinement. prior_hypotheses=[H1 client-side-buffer-overflow]. new evidence: client buffer metrics look fine but a packet shows zero-window from server. Refine: weaken H1, add H2 about server-side flow control / receive buffer pressure",
- ],
- "test_planner": [
- # Original 16
- "two hypotheses: middlebox bookmark-stripping vs server-app reject; need test to pin layer",
- "two hypotheses: TLS resumption ticket key rotation vs cache eviction; need test inside server logs vs network",
- "three hypotheses for slow first-byte: dns, tcp_handshake, tls_handshake; one tshark filter each",
- "two hypotheses: MTU blackhole inside tunnel vs end-host MSS clamping issue; need a test",
- "two hypotheses for spurious RSTs: middlebox state-table eviction vs server-app explicit close; need test",
- "two hypotheses: client SACK disabled vs server SACK disabled; need a test using one direction's pcap only",
- "two hypotheses: tcp window-scale option dropped by middlebox vs misconfigured client; need test",
- "three hypotheses for jitter: cross-traffic, hardware queue, polling-cycle; one test each",
- "two hypotheses: NAT entry expiration vs ECMP rehash mid-flow; need a test",
- "two hypotheses for one-way audio: codec mismatch vs return-path drop; need a test",
- "two hypotheses for HTTP/2 stalls: flow-control credit exhaustion vs HEADERS frame fragmentation; need test",
- "two hypotheses for VPN drop: rekey storm vs MTU change; need test using only IPsec headers",
- "three hypotheses for periodic latency: ARP refresh, GC pause, IPC ring full; one test each",
- "two hypotheses for TLS slow handshake: OCSP stapling fetch vs cert chain depth; need test",
- "two hypotheses: BGP convergence vs route reflector delay; need a passively-observable test",
- "two hypotheses for HTTP 503: backend timeout vs LB health-check flap; need test using pcap",
- # Expansion +16
- "two hypotheses: TLS 1.3 0-RTT replay vs ticket binding mismatch; need test using single client capture",
- "three hypotheses for IPv6 hop limit drops: ICMP rate limit, transit RA prefix mismatch, MTU; one test each",
- "two hypotheses: HTTP/3 PADDING frame causing stream stall vs HEADER frame fragmentation; need test",
- "two hypotheses for SCTP heartbeat fail: path MTU vs NAT pinhole; need test using primary path only",
- "three hypotheses for slow CoAP responses: server retransmit, gateway buffering, DTLS handshake; one test each",
- "two hypotheses: TCP timestamps PAWS reject vs sequence number wrap; need test using TS option",
- "two hypotheses for missing Multicast: IGMP membership not joined vs PIM-SM prune; need test",
- "three hypotheses for HTTP/2 stall: SETTINGS_MAX_CONCURRENT_STREAMS, flow control, GOAWAY; one test each",
- "two hypotheses: BGP convergence vs route advertisement filter; need test using AS_PATH attribute",
- "two hypotheses for failed WebRTC: ICE failure vs DTLS-SRTP key exchange; need test using STUN binding stats",
- "three hypotheses for NFS slow: rpc.statd timeout, client cache stale, server queue full; one test each",
- "two hypotheses: WireGuard MTU vs key rotation drift; need test using handshake_init counter",
- "two hypotheses: NTP DDoS amplification block vs ratelimit; need test using mode 6/7 packets",
- "two hypotheses for RTSP pause: codec switch vs keyframe miss; need test using NAL unit type",
- "three hypotheses for OAuth2 fail: token-endpoint TLS, JWT signature, scope mismatch; one test each",
- "two hypotheses for QUIC fallback: UDP block at LB vs ALPN advertise of h3-only; need test using SettingsParameters",
- ],
- "observation_normalizer": [
- # Original 16
- "tshark output with sip.method, sip.cseq for SIP INVITE/100/180/200 dialog",
- "tshark output with diameter.cmd.code, diameter.flags showing Cer/Cea exchange",
- "tshark output with rtp.ssrc, rtp.seq, rtp.timestamp, rtp.marker for an RTP audio stream",
- "tshark output with smb2.cmd, smb2.tree.id for SMB session establishment",
- "tshark output with kerberos.msg_type, kerberos.realm for AS-REQ/AS-REP exchange",
- "tshark output with ntp.stratum, ntp.refid, ntp.precision for NTP client polling",
- "tshark output with ldap.protocolOp showing bindRequest and searchResultEntry",
- "tshark output with snmp.opcode, snmp.varbind for GET/RESPONSE pairs",
- "zeek http.log with method, uri, status_code, response_body_len fields",
- "zeek ssl.log with version, cipher, server_name, validation_status",
- "zeek dhcp.log with msg_types and DHCP option values for lease acquisition",
- "zeek ssh.log with auth_attempts and auth_success fields",
- "tshark output with quic.frame.type, quic.connection.id for a QUIC initial packet",
- "tshark output with mqtt.msgtype, mqtt.topic, mqtt.qos for a publish exchange",
- "tshark output with bgp.type, bgp.path_attr.as_path for an UPDATE message",
- "tshark output with rpcrdma.opcode for RDMA-over-TCP NFS traffic",
- # Expansion +16
- "tshark output with diameter.session.id, diameter.origin.host for Gx session establishment",
- "tshark output with mip6.type for IPv6 mobile node binding update exchange",
- "tshark output with srv6.sids for ingress SR header processing",
- "tshark output with isakmp.spi for IKE_SA_INIT phase one",
- "zeek dns.log entries with rcode_name, query_type, ttl for client resolver behavior",
- "zeek conn.log entries with history field showing TCP state transitions",
- "tshark output with quic.frame.type for ACK_ECN counting paths",
- "tshark output with sctp.heartbeat.info for SCTP path monitoring",
- "zeek modbus.log entries with unit_id, function for industrial control queries",
- "tshark output with stp.bridgepriority for spanning tree root election",
- "tshark output with srtp.encrypted_payload for SRTP-protected RTP media",
- "tshark output with ssh.message_code for SSH transport layer key exchange",
- "tshark output with snmp.engineid for SNMPv3 USM authentication",
- "tshark output with l2tp.tunnel.id for L2TP control message header",
- "zeek rdp.log entries with cookie, security, keyboard_layout for RDP session start",
- "tshark output with cassandra.opcode for CQL native protocol startup",
- ],
- "hypothesis_updater": [
- # Original 16
- "hyp: HTTP/2 flow-control starvation; pred: WINDOW_UPDATE frames absent; obs: WINDOW_UPDATE frames present every 1MB",
- "hyp: MTU clamping issue; pred: TCP MSS option < 1460; obs: MSS = 1460 in SYN",
- "hyp: TLS cert chain too deep; pred: handshake size > 8KB; obs: handshake 3.2KB",
- "hyp: SACK disabled; pred: no SACK blocks in retransmits; obs: SACK blocks present",
- "hyp: BGP route reflector delay; pred: UPDATE arrives 30s+ after rr-client; obs: UPDATE within 200ms",
- "hyp: kerberos pre-auth required; pred: AS-REQ rejected with KDC_ERR_PREAUTH_REQUIRED; obs: AS-REP returned directly",
- "hyp: middlebox closes idle flows; pred: RST from non-server TTL; obs: RST from server TTL",
- "hyp: QUIC amplification protection limit; pred: server caps initial packets to 3x client; obs: 5x observed",
- "hyp: DNS resolver returns SERVFAIL; pred: RCODE=2 on every query; obs: RCODE=0 for some, SERVFAIL for half",
- "hyp: TLS resumption ticket reuse; pred: psk_identity present; obs: psk_identity absent, full handshake",
- "hyp: HTTP/3 fallback due to udp port block; pred: no QUIC initial packets observed; obs: QUIC initial present, no Handshake",
- "hyp: ipsec rekey storm; pred: IKE INFORMATIONAL messages every 10s; obs: every 600s as configured",
- "hyp: NFS retransmit due to storage failover; pred: cluster around failover events; obs: uniform distribution",
- "hyp: GC pause causing 200ms gap; pred: gap aligned with sawtooth memory pattern; obs: gap not periodic",
- "hyp: dual-stack AAAA blocked; pred: A query succeeds and AAAA times out; obs: both A and AAAA receive responses",
- "hyp: SMB2 session expiration; pred: TREE_DISCONNECT then TREE_CONNECT after timeout; obs: continuous session reuse",
- # Expansion +16
- "hyp: HTTP/2 SETTINGS_INITIAL_WINDOW_SIZE too small; pred: WINDOW_UPDATE within 1KB; obs: 64KB initial window present",
- "hyp: TCP timestamps disabled at one endpoint; pred: TS option absent in SYN; obs: TS option present both directions",
- "hyp: TLS post-handshake auth requested; pred: CertificateRequest after Finished; obs: no post-handshake messages",
- "hyp: BGP convergence delayed by route refresh; pred: ROUTE-REFRESH message present; obs: no refresh, gradual UPDATE",
- "hyp: SCTP heartbeat interval too long; pred: HEARTBEAT every 30s+; obs: HEARTBEAT every 30s as RFC default",
- "hyp: HTTP/3 0-RTT replay attack mitigation triggers; pred: server rejects with REJECT; obs: REJECT not present",
- "hyp: MTU clamping by middlebox; pred: SYN MSS rewritten; obs: SYN MSS unchanged from origin",
- "hyp: NFS retransmit due to slow disk; pred: rpc.duplicate calls every 5s+; obs: duplicate calls within 200ms (network)",
- "hyp: WireGuard handshake retry due to PSK mismatch; pred: handshake initiation repeated 3x within 5s; obs: only 1 retry",
- "hyp: TLS HelloRetryRequest needed; pred: server sends HRR; obs: server sends ServerHello directly",
- "hyp: GRE keepalive timeout; pred: tunnel state shows expired; obs: keepalive packets ack'd every 10s",
- "hyp: SIP REGISTER refresh failure; pred: 401 challenge with stale nonce; obs: 200 OK returned directly",
- "hyp: QUIC congestion control collapse; pred: congestion window drops to 1 MSS; obs: cwnd stays at IW10 throughout",
- "hyp: DTLS rekey storm; pred: ClientHello every 60s+; obs: only one ClientHello in 10 minutes",
- "hyp: RIP route poisoning event; pred: metric=16 advertised; obs: metric=15 used, route still valid",
- "hyp: Modbus exception response; pred: function code with high bit set; obs: normal response code returned",
- ],
- "stop_controller": [
- # Original 16
- "five hyps, all under 0.3 confidence, three uninvestigated remaining tests are cheap and informative",
- "single hyp confirmed across 20 sampled flows with 3 contradicting flows that look like noise",
- "two hyps still at 50/50, all remaining tests would yield observations we have already seen",
- "investigation 14 actions in, every action moved no hyp by more than 0.05 — stuck",
- "one hyp at 0.92, one at 0.05 — clear winner; only 2 of 50 tools used",
- "evidence saturated: all flows captured one-sided, root cause requires server-side data we don't have",
- "loop: 2 actions ran, both confirmed the dominant hyp; still in early phase, can dig more",
- "investigation cost at 80% of budget, leading hyp at 0.7, other 0.2; budget low for one more test",
- "all hyps consistent with available data; the only distinguishing test requires firewall reset which is out of scope",
- "evidence partially supports two hyps each from different angles — possible joint cause; cannot separate",
- "single hyp at 0.95 with 1 contradicting observation; needs explanation before stop_with_conclusion",
- "three hyps in cluster of 0.4/0.3/0.3; no single test can distinguish all three at once",
- "hyp confidence movements pre vs post last 4 tests are < 0.01 — diminishing returns regime",
- "investigation running 30 min wall (vs 5 min typical); state is ambiguous; user prefers timely partial answer",
- "evidence quality questionable: pcap shows packet capture drops > 0.5%; results unreliable for confirmation",
- "one hyp confirmed strongly on aggregate flow stats but lacks any per-packet evidence — confirm or dig deeper?",
- # Expansion +16
- "four hyps with confidences 0.3/0.25/0.25/0.2 — no clear winner, only 30% of tools used",
- "all hyps have been tested with at least 3 actions; all results in noise floor, decision required",
- "single hyp at 0.85, but evidence comes from a single capture point that may have bias",
- "user explicitly asks for partial answer due to time pressure; current top hyp at 0.65",
- "cluster of 2 hyps both at 0.45 are mutually exclusive but no test would distinguish them",
- "investigation already consumed 90% of allocated time; no clear winner; cheapest 3 tests run",
- "three hyps explored, two confirmed by independent flows — possibly composite root cause",
- "dominant hyp explained 12/15 observations; 3 unexplained observations cluster on one IP",
- "all 5 hyps are alternative framings of 'TLS handshake timeout'; need to refine before stop",
- "evidence quality high but interpretation depends on whether capture point has loss",
- "confidence on top hyp moved 0.4→0.7 over last 5 tests; trajectory predicts more lift if continued",
- "evidence supports complex chain of two events; can stop on the proximate cause and hand off root",
- "dominant hyp would imply server-side bug, but we lack capability to verify in server logs",
- "all hyps at low confidence (<0.4) but symptom is critical (production outage); need at least partial answer",
- "one hyp at 0.9 with 1 contradicting flow that is structurally different from others — outlier or evidence?",
- "investigation found no support for original 3 hyps; new hyp emerged at 0.55 — confirm or check baseline first",
- ],
-}
-
-
-def all_train_task_ids() -> list[str]:
- return list(TRAIN_SCENARIOS.keys())
-
-
-def scenarios_for(task_id: str, scenario_list: str = "train") -> list[str]:
- """Return scenarios for a task. scenario_list='train' or 'eval'."""
- if scenario_list == "train":
- return TRAIN_SCENARIOS[task_id]
- if scenario_list == "eval":
- from .schemas import DIAG_TASKS
- return DIAG_TASKS[task_id].seed_scenarios
- raise ValueError(f"scenario_list must be 'train' or 'eval', got {scenario_list!r}")
diff --git a/vendor/nanoai/agentic/eval_loop.py b/vendor/nanoai/agentic/eval_loop.py
deleted file mode 100644
index 85d568d3732654faae63b208f333d59768caeeff..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/eval_loop.py
+++ /dev/null
@@ -1,332 +0,0 @@
-"""Multi-turn agent eval state machine.
-
-For each task: generate → parse → execute tool → wrap result → repeat.
-Stops when assistant emits `<|assistant_end|>` without a tool call, or when
-turn budget is exceeded, or when the model emits malformed output twice.
-"""
-from __future__ import annotations
-
-import re
-import time
-from typing import Optional
-
-from .data.common import SYSTEM_PROMPT
-from .eval_sandbox import Sandbox, SandboxError
-from .eval_tasks import Task, Transcript
-
-
-# Required arg keys per tool. Original Read/Edit/Grep/Bash come from
-# scripts/agent_eval.py; the lowercase ones come from ds79c (TraceForge
-# id=90) — same trajectories used by exp307a/b/c/324. Tools that exist
-# only as no-op stubs (write_todos, task, …) accept any args.
-TOOL_REQUIRED_ARGS: dict[str, set[str]] = {
- # Original nanoai tool surface
- "Read": {"file_path"},
- "Edit": {"file_path", "new_string"},
- "Grep": {"pattern"},
- "Bash": {"command"},
- # ds79c/demycap tool surface (lowercase names)
- "read_file": {"file_path"},
- "edit_file": {"file_path"}, # new_string optional (write_file behaviour for empty)
- "write_file": {"file_path"}, # content optional, model often forgets
- "execute": {"command"},
- "ls": set(),
- "glob": {"pattern"},
- "grep": {"pattern"},
- "save_mermaid": {"file_path"},
- "submit_report": {"report_markdown"}, # confidence is optional
- "write_todos": set(),
- "update_todos": set(),
- "update_todo": set(),
- "task": set(),
-}
-
-# Tools whose return value should terminate the multi-turn loop. The
-# only one today is `submit_report`; the call captures the payload on
-# the sandbox and the loop exits without sampling another turn.
-TERMINAL_TOOLS = frozenset({"submit_report"})
-
-
-def _parse_tool_call(body: str) -> tuple[str, dict[str, str]]:
- """Parse `tool_name <|tool_arg|> k <|tool_val|> v ...` into (name, args)."""
- parts = body.split("<|tool_arg|>")
- name = parts[0].strip()
- args: dict[str, str] = {}
- for part in parts[1:]:
- if "<|tool_val|>" in part:
- k, v = part.split("<|tool_val|>", 1)
- args[k.strip()] = v.strip()
- return name, args
-
-
-def _coerce_float(s: str, default: float = 0.0) -> float:
- try:
- return float(s)
- except (TypeError, ValueError):
- return default
-
-
-def _coerce_int(s: str, default: int = 0) -> int:
- try:
- return int(s)
- except (TypeError, ValueError):
- return default
-
-
-def _execute(sb: Sandbox, name: str, args: dict[str, str]) -> tuple[str, bool]:
- """Run a parsed tool call. Returns (result_text, ok)."""
- required = TOOL_REQUIRED_ARGS.get(name)
- if required is None:
- return f"unknown tool: {name!r}", False
- missing = required - args.keys()
- if missing:
- return f"missing required arg(s) for {name}: {sorted(missing)}", False
- try:
- # Original nanoai surface
- if name == "Read":
- return sb.read(args["file_path"],
- offset=_coerce_int(args.get("offset", "")),
- limit=_coerce_int(args.get("limit", ""), 50)), True
- if name == "Edit":
- return sb.edit(args["file_path"], args["new_string"], args.get("old_string", "")), True
- if name == "Grep":
- return sb.grep(args["pattern"], args.get("path") or "."), True
- if name == "Bash":
- return sb.bash(args["command"]), True
- # ds79c/demycap surface
- if name == "read_file":
- return sb.read_file(args["file_path"],
- offset=_coerce_int(args.get("offset", "")),
- limit=_coerce_int(args.get("limit", ""), 50)), True
- if name == "edit_file":
- return sb.edit_file(args["file_path"],
- new_string=args.get("new_string", ""),
- old_string=args.get("old_string", "")), True
- if name == "write_file":
- return sb.write_file(args["file_path"], args.get("content", "")), True
- if name == "execute":
- return sb.execute(args["command"]), True
- if name == "ls":
- return sb.ls_path(args.get("path", ".")), True
- if name == "glob":
- return sb.glob(args.get("pattern", "*"), args.get("path", ".")), True
- if name == "grep":
- return sb.grep(args["pattern"], args.get("path") or "."), True
- if name == "save_mermaid":
- return sb.save_mermaid(args["file_path"], args.get("content", "")), True
- if name == "submit_report":
- return sb.submit_report(
- report_markdown=args.get("report_markdown", ""),
- confidence=_coerce_float(args.get("confidence", "")),
- ), True
- if name == "write_todos":
- return sb.write_todos(), True
- if name == "update_todos":
- return sb.update_todos(), True
- if name == "update_todo":
- return sb.update_todo(**{k: v for k, v in args.items()}), True
- if name == "task":
- return sb.task(description=args.get("description", "")), True
- except SandboxError as e:
- return f"tool error: {e}", False
- except Exception as e:
- return f"tool crash: {type(e).__name__}: {e}", False
- return "internal: unreachable", False
-
-
-def _build_initial_tokens(tokenizer, user_text: str, system_prompt: Optional[str] = None) -> list[int]:
- # Adapter path (e.g. QwenTokenizer) — delegate to tokenizer's own chat template.
- if hasattr(tokenizer, "build_initial_tokens"):
- sp = SYSTEM_PROMPT if system_prompt is None else system_prompt
- return tokenizer.build_initial_tokens(user_text, system_prompt=sp)
- bos = tokenizer.get_bos_token_id()
- user_start = tokenizer.encode_special("<|user_start|>")
- user_end = tokenizer.encode_special("<|user_end|>")
- asst_start = tokenizer.encode_special("<|assistant_start|>")
- sp = SYSTEM_PROMPT if system_prompt is None else system_prompt
- body = sp + "\n\n" + user_text
- return [bos, user_start] + tokenizer.encode(body) + [user_end, asst_start]
-
-
-def _append_tool_result(tokenizer, tokens: list[int], result_text: str) -> None:
- """Append `<|tool_result_start|>{result}<|tool_result_end|><|assistant_start|>`."""
- if hasattr(tokenizer, "append_tool_result"):
- tokenizer.append_tool_result(tokens, result_text)
- return
- tokens.append(tokenizer.encode_special("<|tool_result_start|>"))
- # Truncate result to keep KV cache usage bounded.
- if len(result_text) > 4000:
- result_text = result_text[:4000] + "\n...[truncated]"
- tokens.extend(tokenizer.encode(result_text))
- tokens.append(tokenizer.encode_special("<|tool_result_end|>"))
- tokens.append(tokenizer.encode_special("<|assistant_start|>"))
-
-
-def _assistant_end_id(tokenizer):
- if hasattr(tokenizer, "assistant_end_id"):
- return tokenizer.assistant_end_id()
- return tokenizer.encode_special("<|assistant_end|>")
-
-
-def _tool_call_pair(tokenizer):
- if hasattr(tokenizer, "tool_call_token_pair"):
- return tokenizer.tool_call_token_pair()
- return (
- tokenizer.encode_special("<|tool_call_start|>"),
- tokenizer.encode_special("<|tool_call_end|>"),
- )
-
-
-def _parse_tool_call_dispatch(tokenizer, body_text: str):
- """Dispatch to tokenizer-specific body parser (Qwen JSON vs nanoai k/v)."""
- if hasattr(tokenizer, "parse_tool_body"):
- return tokenizer.parse_tool_body(body_text)
- return _parse_tool_call(body_text)
-
-
-def _generate_one_turn(engine, tokenizer, prompt_tokens: list[int],
- max_tokens: int, temperature: float, seed: int,
- capture_logprobs: bool = False):
- """Run engine.generate until <|assistant_end|> or max_tokens.
-
- Returns the list of generated token ids (excluding the prompt). When
- capture_logprobs=True, returns (out_ids, out_logprobs) where out_logprobs[i]
- is the behavior-policy log prob of out_ids[i] (None for any forced token) —
- used to store behavior logprobs for TIS (Polar mechanism)."""
- asst_end = _assistant_end_id(tokenizer)
- bos = tokenizer.get_bos_token_id()
- out_ids: list[int] = []
- # Only pass return_logprobs when capturing, so engines/fakes that don't
- # accept the kwarg keep working unchanged (backward compatible).
- if capture_logprobs:
- out_logprobs: list = []
- gen = engine.generate(prompt_tokens, num_samples=1, max_tokens=max_tokens,
- temperature=temperature, seed=seed, return_logprobs=True)
- for token_column, _masks, logprob_column in gen:
- out_logprobs.append(logprob_column[0])
- tok = token_column[0]
- out_ids.append(tok)
- if tok == asst_end or tok == bos:
- break
- return out_ids, out_logprobs
- for token_column, _masks in engine.generate(
- prompt_tokens,
- num_samples=1,
- max_tokens=max_tokens,
- temperature=temperature,
- seed=seed,
- ):
- tok = token_column[0]
- out_ids.append(tok)
- if tok == asst_end or tok == bos:
- break
- return out_ids
-
-
-def run_task(
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- *,
- max_tokens_per_turn: int = 256,
- max_turns: Optional[int] = None,
- temperature: float = 0.5,
- seed: int = 42,
- return_transcript: bool = False,
-) -> dict:
- """Run one task end-to-end. Returns a dict with pass/fail + diagnostics.
-
- When return_transcript=True the dict also includes:
- - 'user_prompt': str
- - 'transcript': list[dict] (every assistant content + tool_call(name+args) +
- tool_result text + error flag, in order)
- """
- max_turns = max_turns or task.max_turns
- transcript = Transcript(user_prompt=task.prompt)
- task.setup(sandbox)
-
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- tokens = _build_initial_tokens(tokenizer, task.prompt,
- system_prompt=getattr(task, "system_prompt", None))
- n_turns = 0
- n_tool_errors = 0
- t0 = time.time()
- truncated = False
-
- while n_turns < max_turns:
- n_turns += 1
- out_ids = _generate_one_turn(engine, tokenizer, tokens,
- max_tokens_per_turn, temperature, seed + n_turns)
- # Trim trailing assistant_end if present.
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
- if not ended_clean and len(out_ids) >= max_tokens_per_turn:
- truncated = True
-
- # Inspect for tool call.
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- j = -1
- if j > i:
- pre = tokenizer.decode(body_ids[:i]).strip()
- tool_body = tokenizer.decode(body_ids[i + 1 : j])
- try:
- name, args = _parse_tool_call_dispatch(tokenizer, tool_body)
- except ValueError as e:
- # Malformed tool body — surface as a tool error, no execute.
- name, args = "", {}
- result_text, ok = f"tool parse error: {e}", False
- else:
- result_text, ok = _execute(sandbox, name, args)
- transcript.turns.append({
- "role": "assistant",
- "content": pre,
- "tool_call": {"name": name, "args": args},
- })
- transcript.turns.append({
- "role": "tool_result",
- "tool_result": result_text,
- "error": not ok,
- })
- if not ok:
- n_tool_errors += 1
- # Terminating tool (submit_report) → exit loop; no need
- # to extend context for another turn.
- if name in TERMINAL_TOOLS:
- break
- # Extend context with the assistant turn we just generated +
- # tool result + a fresh assistant_start.
- tokens.extend(out_ids) # includes asst_end if present
- if not ended_clean:
- tokens.append(asst_end)
- _append_tool_result(tokenizer, tokens, result_text)
- continue
-
- # No tool call this turn → final answer.
- text = tokenizer.decode(body_ids).strip()
- transcript.turns.append({"role": "assistant", "content": text})
- break
-
- elapsed = time.time() - t0
- passed, reason = task.check(sandbox, transcript)
- out = {
- "task_id": task.task_id,
- "passed": passed,
- "reason": reason,
- "n_turns": n_turns,
- "n_tool_errors": n_tool_errors,
- "truncated": truncated,
- "elapsed_s": round(elapsed, 2),
- "tools_used": [c.get("name", "") for c in transcript.tool_calls],
- "last_text": transcript.last_assistant_text[:200],
- }
- if return_transcript:
- out["user_prompt"] = transcript.user_prompt
- out["transcript"] = transcript.turns
- return out
diff --git a/vendor/nanoai/agentic/eval_pcap_tasks.py b/vendor/nanoai/agentic/eval_pcap_tasks.py
deleted file mode 100644
index 383239f23592e86bcb4c2129540e4b551434ff6f..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/eval_pcap_tasks.py
+++ /dev/null
@@ -1,414 +0,0 @@
-"""Pcap diagnostic eval tasks — multi-turn agent runs tshark in the
-sandbox and produces a capbench-style structured diagnostic report.
-
-The report format matches what ds79c (TraceForge id=90) trajectories
-emit — exp307a/b/c/324's known-good demycap output:
-
-
- # Title
-
- ## Executive Summary
- free-form prose ...
-
- ## Root Cause
- ```yaml
- root_causes:
- - "TCP Zero Window / Server receive buffer exhaustion"
- - "..."
- ```
-
- ## Evidence
- ```yaml
- evidence:
- - id: "EVIDENCE-001-..."
- - id: "EVIDENCE-002-..."
- ```
- | metric table … |
-
- ## Impact
- ```yaml
- impact:
- - "Large file uploads fail with HTTP 502"
- ```
-
- ## Remediation
- ```yaml
- remediation:
- - "Increase TCP receive buffer size on server"
- ```
-
-
-The rubric extracts content between `...`,
-parses the 4 named yaml lists, and scores each section against the
-per-task ground truth (`data/pcap_corpus/ground_truth/.json`).
-Per-section score = presence + non-empty-list + keyword overlap with
-expected items. Mean over 4 sections + bonus for clean wrapper.
-
-Tasks load via `all_pcap_tasks()` from the `data/pcap_corpus/index.json`
-built by `experiments/ksbr_pcap/scripts/build_pcap_corpus.py`.
-"""
-from __future__ import annotations
-
-import json
-import re
-from pathlib import Path
-from typing import Any, Optional
-
-import yaml
-
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-
-
-REPORT_SECTIONS = ("root_causes", "evidence", "impact", "remediation")
-DEFAULT_PASS_THRESHOLD = 0.50
-DEFAULT_MAX_TURNS = 12
-MIN_SECTION_ITEMS = 1 # yaml list must have at least this many items to score
-KEYWORD_HIT_RATIO_FULL = 0.5 # ≥50% keyword overlap → full per-section credit
-
-
-# System prompt that mirrors the block at the top of every
-# ds79c trajectory. Without this, the trained-on-ds79c model never activates
-# its diagnostic conditional — it falls back to the v3 nanocode coding-agent
-# persona and treats the task as a Python coding question. Empirically: with
-# the default SYSTEM_PROMPT v7 scored 0/5 with all reports empty; the
-# trajectories went straight to writing solution_NNN.py. This text matches
-# ds79c's training distribution closely enough that the conditional fires.
-_PCAP_SYSTEM_PROMPT_TEMPLATE = """
-
-### Environment
-
-- All files are in the current working directory. The `execute` tool also runs in this directory.
-- PCAP file(s):
-{pcap_lines}
-- No pre-generated evidence — analyze the PCAP directly with `tshark`.
-- Mermaid output: `mermaid/` | CSV output: `csv/`
-- Use tshark directly with the filename, e.g. `tshark -r {primary_pcap} -q -z expert`
-
-### Available Workflow Skills
-
-- `workflow-analysis`: 用于分析用户上传的 PCAP,定位网络故障、性能问题和根因。evidence 由你通过 tshark 分析后自己生成。
-- `workflow-network-knowledge`: 用于回答网络协议、术语或排障方法的知识性问题,不依赖具体 PCAP 分析。
-- `workflow-pcap-detail`: 用于导出流级或包级明细 CSV,便于后续统计、解释或人工复核。
-- `workflow-pcap-filter`: 用于根据端点、端口、时间窗口或协议条件筛选原始 PCAP,并生成派生抓包。
-- `workflow-pcap-statistics`: 用于快速统计协议分布、端点、会话、吞吐、重传等指标,输出概览性结论。
-- `workflow-topology-streamdiff`: 用于基于一份或两份抓包解释通信关系、路径可视化或同一连接在不同采集点的差异。
-
-### Constraints
-
-- Select one primary workflow skill before deep analysis.
-- Use `execute` with `tshark` (and other shell tools) to extract evidence from the PCAP.
-- Use `save_mermaid` (not generic file writes) for `.mmd` files.
-- Conclude with `submit_report` carrying a `report_markdown` argument that contains the
- `...` block with yaml sections for `root_causes`,
- `evidence`, `impact`, and `remediation`.
-
-"""
-
-
-def _render_pcap_system_prompt(pcap_filenames: list[str]) -> str:
- """Render PCAP_SYSTEM_PROMPT with a dynamic list of pcap filenames.
-
- Single-pcap call sites (single-pcap tasks, legacy 5-scenario corpus) pass
- ['capture.pcap'] and get the original prompt byte-for-byte.
- Multi-pcap tasks pass ['capture_a.pcap', 'capture_b.pcap'] and the env
- section lists both with Capture point A/B/C/... labels.
- """
- if not pcap_filenames:
- raise ValueError("at least one pcap filename required")
- labels = ["A", "B", "C", "D", "E", "F"][: len(pcap_filenames)]
- pcap_lines = "\n".join(
- f" - Capture point {label}: {name}"
- for label, name in zip(labels, pcap_filenames)
- )
- return _PCAP_SYSTEM_PROMPT_TEMPLATE.format(
- pcap_lines=pcap_lines,
- primary_pcap=pcap_filenames[0],
- )
-
-
-# Backward-compat alias: existing code (eval_loop.py, scripts/agent_eval.py, etc.)
-# imports PCAP_SYSTEM_PROMPT as a module-level constant. Preserve the single-pcap
-# default so non-pcap_corpus call sites keep their byte-identical behavior.
-PCAP_SYSTEM_PROMPT = _render_pcap_system_prompt(["capture.pcap"])
-
-
-# ----------------------------------------------------------------------
-# Report parsing
-_FINAL_RE = re.compile(r"\s*(.*?)\s*", flags=re.DOTALL)
-# Yaml fenced block; named anchor is the first key of the block, e.g.
-# ``` followed by `yaml` (optional) then `root_causes:` ... ```
-_YAML_BLOCK_RE = re.compile(
- r"```(?:yaml|yml)?\s*\n"
- r"((?:root_causes|evidence|impact|remediation)\s*:.*?)"
- r"\n```",
- flags=re.DOTALL,
-)
-
-
-def _extract_final_answer(text: str) -> Optional[str]:
- """Return the markdown body inside the first block, or None."""
- m = _FINAL_RE.search(text or "")
- if m:
- return m.group(1)
- # Fallback: if there's an opener but no closer (model truncated), take
- # everything after `` so we still score what we have.
- if "" in (text or ""):
- return text.split("", 1)[1]
- return None
-
-
-def _items_from_yaml(body: str, key: str) -> list[str]:
- """Extract a YAML list of strings under `key:`. Each item may be a bare
- string or a `{id: "..."}` dict (capbench's evidence section uses dicts).
- Returns the list of stringified items; [] if missing or malformed."""
- try:
- obj = yaml.safe_load(body)
- except Exception:
- return []
- if not isinstance(obj, dict):
- return []
- val = obj.get(key)
- if not isinstance(val, list):
- return []
- out: list[str] = []
- for item in val:
- if isinstance(item, str):
- out.append(item.strip())
- elif isinstance(item, dict):
- # Prefer `id` then `name` then any string value
- for prefer in ("id", "name", "title"):
- if prefer in item and isinstance(item[prefer], str):
- out.append(item[prefer].strip())
- break
- else:
- # Fallback: stringify the dict
- out.append(json.dumps(item, ensure_ascii=False).strip())
- return [s for s in out if s]
-
-
-def parse_report(text: str) -> dict[str, Any]:
- """Parse a capbench-style report. Returns:
- {
- "wrapped": bool, # True if … present
- "wrapped_closed": bool, # True if both opener AND closer matched
- "body": str, # markdown inside the wrapper (or full text)
- "root_causes": [str, ...],
- "evidence": [str, ...],
- "impact": [str, ...],
- "remediation": [str, ...],
- }
- Robust to missing closing tag, missing yaml fences, and non-yaml fallbacks."""
- raw = text or ""
- m = _FINAL_RE.search(raw)
- if m:
- body = m.group(1)
- wrapped, wrapped_closed = True, True
- elif "" in raw:
- body = raw.split("", 1)[1]
- wrapped, wrapped_closed = True, False
- else:
- body = raw
- wrapped, wrapped_closed = False, False
-
- out: dict[str, Any] = {
- "wrapped": wrapped,
- "wrapped_closed": wrapped_closed,
- "body": body,
- }
- for sec in REPORT_SECTIONS:
- out[sec] = []
-
- # Walk every yaml block and extract whichever named section it carries.
- for ym in _YAML_BLOCK_RE.finditer(body):
- chunk = ym.group(1)
- for sec in REPORT_SECTIONS:
- items = _items_from_yaml(chunk, sec)
- if items and not out[sec]:
- out[sec] = items
-
- # Fallback for un-fenced inline yaml (some reports drop the ```yaml fence)
- for sec in REPORT_SECTIONS:
- if out[sec]:
- continue
- m = re.search(rf"(?m)^{sec}\s*:\s*\n((?:[ \t]+-.*\n?)+)", body)
- if m:
- out[sec] = _items_from_yaml(f"{sec}:\n{m.group(1)}", sec)
-
- return out
-
-
-# ----------------------------------------------------------------------
-# Rubric scoring
-def _section_score(predicted: list[str], expected: list[str]) -> float:
- """Score one yaml section against expected items. 0.5 for non-empty; +0.5
- scaled to keyword overlap ratio (≥50% → full)."""
- if not predicted or len(predicted) < MIN_SECTION_ITEMS:
- return 0.0
- s = 0.5
- if expected:
- # Each expected item is a keyword set; treat exact-substring match
- # against any predicted item as a hit.
- joined_pred = " ".join(predicted).lower()
- hits = sum(1 for kw in expected if kw.lower() in joined_pred)
- ratio = hits / len(expected)
- s += min(0.5, 0.5 * ratio / KEYWORD_HIT_RATIO_FULL)
- else:
- # No GT → can't grade content; just credit non-emptiness fully
- s += 0.5
- return s
-
-
-def score_report(parsed: dict[str, Any], ground_truth: dict[str, Any]) -> tuple[float, dict]:
- """Score a parsed report. Returns (score in [0,1], breakdown).
-
- Breakdown:
- - 4 yaml sections × 1pt each (presence + keyword overlap)
- - +0.2 bonus for clean ... wrapper
- - normalised by 4.2"""
- breakdown: dict[str, Any] = {}
- total = 0.0
- for sec in REPORT_SECTIONS:
- expected = list(ground_truth.get(sec) or [])
- s = _section_score(parsed.get(sec) or [], expected)
- breakdown[sec] = round(s, 3)
- total += s
- bonus = 0.0
- if parsed.get("wrapped") and parsed.get("wrapped_closed"):
- bonus = 0.2
- elif parsed.get("wrapped"):
- bonus = 0.1
- breakdown["wrapper_bonus"] = bonus
- norm = round((total + bonus) / 4.2, 4)
- return norm, breakdown
-
-
-# ----------------------------------------------------------------------
-# Task construction
-def _make_pcap_task(
- task_id: str,
- pcap_paths: list[Path] | Path,
- ground_truth: dict[str, Any],
- prompt: str,
- pass_threshold: float = DEFAULT_PASS_THRESHOLD,
- max_turns: int = DEFAULT_MAX_TURNS,
- pcap_filenames: list[str] | None = None,
-) -> Task:
- """Build a pcap diagnostic Task.
-
- Single-pcap call (back-compat): pcap_paths is a Path, pcap_filenames=None
- → sandbox file is named "capture.pcap", system prompt lists "Capture point A: capture.pcap".
-
- Multi-pcap call: pcap_paths is a list of Path, pcap_filenames is the same-
- length list of sandbox names (e.g. ["capture_a.pcap", "capture_b.pcap"]).
- """
- if isinstance(pcap_paths, (str, Path)):
- pcap_paths = [Path(pcap_paths)]
- pcap_filenames = pcap_filenames or ["capture.pcap"]
- else:
- pcap_paths = [Path(p) for p in pcap_paths]
- if pcap_filenames is None:
- # Default naming for multi-pcap: capture_a.pcap, capture_b.pcap, ...
- pcap_filenames = [f"capture_{chr(ord('a') + i)}.pcap" for i in range(len(pcap_paths))]
- if len(pcap_paths) != len(pcap_filenames):
- raise ValueError(
- f"pcap_paths ({len(pcap_paths)}) and pcap_filenames ({len(pcap_filenames)}) length mismatch"
- )
-
- pcap_bytes_by_name: dict[str, bytes] = {
- name: path.read_bytes() for path, name in zip(pcap_paths, pcap_filenames)
- }
- system_prompt = _render_pcap_system_prompt(pcap_filenames)
-
- def setup(sb: Sandbox) -> None:
- for name, data in pcap_bytes_by_name.items():
- (sb.workspace / name).write_bytes(data)
-
- def check(sb: Sandbox, t: Transcript) -> tuple[bool, str]:
- # Prefer the report submitted via the `submit_report` tool — that's
- # the canonical path for ds79c-trained models. Fall back to scanning
- # the last assistant text for an inline block (for
- # untrained baselines or models that emit the report without going
- # through submit_report).
- if sb.submitted_report and sb.submitted_report.get("report_markdown"):
- text = sb.submitted_report["report_markdown"]
- via = "submit_report"
- else:
- text = t.last_assistant_text or t.assistant_text
- via = "inline"
- parsed = parse_report(text)
- score, breakdown = score_report(parsed, ground_truth)
- passed = score >= pass_threshold
- empty_secs = [s for s in REPORT_SECTIONS if breakdown.get(s, 0.0) == 0.0]
- reason = f"score={score:.3f} via={via} threshold={pass_threshold:.2f}"
- if empty_secs:
- reason += f" empty={'/'.join(empty_secs)}"
- if not parsed.get("wrapped"):
- reason += " no_FINAL_ANSWER"
- return passed, reason
-
- return Task(
- task_id=f"pcap_{task_id}",
- prompt=prompt,
- setup=setup,
- check=check,
- tools_expected=("execute", "submit_report"), # tshark via execute, terminate via submit_report
- max_turns=max_turns,
- system_prompt=system_prompt,
- )
-
-
-# ----------------------------------------------------------------------
-# Public loader
-def all_pcap_tasks(corpus_dir: str | Path = "data/pcap_corpus") -> list[Task]:
- """Load every scenario from `corpus_dir/index.json`.
-
- Each index entry uses one of:
- - "pcap": "pcaps/x.pcap" # single-pcap (legacy)
- - "pcaps": ["pcaps/x_a.pcap", "pcaps/x_b.pcap"] # multi-pcap (new)
- - "pcaps": [{"path": "pcaps/x_a.pcap", "name": "capture_a.pcap"}, ...]
- """
- corpus_dir = Path(corpus_dir)
- index_path = corpus_dir / "index.json"
- if not index_path.exists():
- raise FileNotFoundError(
- f"No pcap corpus index at {index_path}. "
- f"Run `python -m experiments.ksbr_pcap.scripts.build_pcap_corpus --out={corpus_dir}` first."
- )
- index = json.loads(index_path.read_text(encoding="utf-8"))
- tasks: list[Task] = []
- for entry in index:
- gt_path = corpus_dir / entry["ground_truth"]
- gt = json.loads(gt_path.read_text(encoding="utf-8"))
- # NOTE: corpus stores UUID-anonymized files (e.g. pcaps//.pcap)
- # for git/storage hygiene, but sandbox-visible filenames MUST follow
- # the training distribution convention: capture.pcap (single) /
- # capture_a.pcap, capture_b.pcap, … (multi). The model is trained on
- # `tshark -r capture.pcap …` patterns; if the sandbox shows
- # `.pcap`, the model's tool calls will reference a non-existent
- # file → tshark errors → empty tool_result → silent failure (this
- # was the root cause of 0/23 strict_pass in 2026-05-18 SOTA chase).
- if "pcaps" in entry:
- raw = entry["pcaps"]
- pcap_paths: list[Path] = []
- pcap_filenames: list[str] = []
- n = len(raw)
- for i, item in enumerate(raw):
- if isinstance(item, dict):
- pcap_paths.append(corpus_dir / item["path"])
- else:
- pcap_paths.append(corpus_dir / item)
- # Always rename to training-distribution convention
- pcap_filenames.append("capture.pcap" if n == 1
- else f"capture_{chr(ord('a') + i)}.pcap")
- else:
- pcap_paths = corpus_dir / entry["pcap"]
- pcap_filenames = ["capture.pcap"]
- tasks.append(_make_pcap_task(
- task_id=entry["task_id"],
- pcap_paths=pcap_paths,
- ground_truth=gt,
- prompt=entry["prompt"],
- pcap_filenames=pcap_filenames,
- ))
- return tasks
diff --git a/vendor/nanoai/agentic/eval_sandbox.py b/vendor/nanoai/agentic/eval_sandbox.py
deleted file mode 100644
index ea8d272631106288bd12319d04021a1ae19a2aa3..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/eval_sandbox.py
+++ /dev/null
@@ -1,313 +0,0 @@
-"""Per-task sandbox for multi-turn agent eval.
-
-Each Sandbox owns a temp workspace. Read/Edit/Grep run in-process on the host
-filesystem with path-validation; Bash spawns a firejail subprocess that sees
-the workspace as $HOME and has no network / no host /tmp / no /dev access.
-"""
-from __future__ import annotations
-
-import os
-import re
-import shlex
-import shutil
-import subprocess
-import tempfile
-from pathlib import Path
-from typing import Optional
-
-# Bash command whitelist — first token of the command must be in this set.
-# We tolerate pipelines (token after `|`) and `&&`/`;` chains by checking each
-# segment separately. Anything outside the list is rejected before firejail.
-BASH_ALLOWED = {
- "ls", "pwd", "cat", "wc", "grep", "find", "head", "tail", "echo",
- "python", "python3", "awk", "sed", "true", "false", "test", "[", "[[",
- "sort", "uniq", "tr", "cut", "diff", "basename", "dirname", "realpath",
- "which", "env", "printf", "expr",
- # pcap diagnostics — read-only on staged pcap files in the workspace.
- # firejail's --net=none + read-only workspace prevents any live capture
- # or network use; tshark falls back to file-mode (-r) only.
- "tshark", "tcpdump", "capinfos", "xxd", "od", "strings",
-}
-
-# Hard-block tokens — if any of these appear anywhere in the command, reject.
-# Belt-and-suspenders: firejail also blocks them, but we fail fast with a
-# clear error to the model.
-BASH_BLOCKED_RE = re.compile(
- r"(^|\s|;|\||&)(rm|mv|cp|chmod|chown|sudo|su|ssh|scp|curl|wget|nc|"
- r"mkfs|dd|kill|reboot|shutdown|systemctl|mount|umount|insmod|modprobe)"
- r"(\s|$|;)"
-)
-
-FIREJAIL_BASE = [
- "firejail",
- "--quiet",
- "--private-tmp",
- "--private-dev",
- "--net=none",
- "--shell=none",
- "--rlimit-fsize=10000000", # 10 MB max file write
- "--rlimit-nproc=64",
- "--rlimit-as=536870912", # 512 MB virtual memory
-]
-
-
-class SandboxError(Exception):
- """Raised when a tool call fails for a deterministic reason (bad path,
- blocked command, timeout). The model sees the error message as a tool
- result so it can adapt; the eval doesn't crash."""
-
-
-class Sandbox:
- """Filesystem sandbox for one agent eval task.
-
- Use as a context manager so the workspace is cleaned up even on error.
- """
-
- def __init__(self, max_read_bytes: int = 8192, bash_timeout: float = 5.0):
- self.max_read_bytes = max_read_bytes
- self.bash_timeout = bash_timeout
- self._workspace: Optional[Path] = None
- # ds79c/demycap-style terminating tool: when the agent calls
- # `submit_report`, the dispatcher captures the report payload here
- # so the eval check function can score it. None means no report
- # was submitted.
- self.submitted_report: Optional[dict] = None
- # Lightweight planning state for write_todos / update_todos /
- # update_todo (no real action — mirrors the demycap UX so trajectories
- # don't error out).
- self._todos: list[dict] = []
-
- def __enter__(self) -> "Sandbox":
- self._workspace = Path(tempfile.mkdtemp(prefix="agent_eval_ws_"))
- return self
-
- def __exit__(self, *exc):
- if self._workspace is not None and self._workspace.exists():
- shutil.rmtree(self._workspace, ignore_errors=True)
- self._workspace = None
-
- @property
- def workspace(self) -> Path:
- assert self._workspace is not None, "Sandbox used outside `with` block"
- return self._workspace
-
- # ------------------------------------------------------------------
- # Path resolution
- def _resolve(self, file_path: str) -> Path:
- """Resolve `file_path` (model-supplied) to an absolute path inside the
- workspace. Reject any path that escapes via .. or absolute paths
- outside the workspace."""
- if not isinstance(file_path, str) or not file_path:
- raise SandboxError("file_path must be a non-empty string")
- ws = self.workspace.resolve()
- candidate = (ws / file_path).resolve()
- try:
- candidate.relative_to(ws)
- except ValueError:
- raise SandboxError(f"path {file_path!r} escapes workspace")
- return candidate
-
- # ------------------------------------------------------------------
- # Tool implementations
- def read(self, file_path: str, offset: int = 0, limit: int = 50) -> str:
- p = self._resolve(file_path)
- if not p.exists():
- raise SandboxError(f"file not found: {file_path}")
- if not p.is_file():
- raise SandboxError(f"not a regular file: {file_path}")
- try:
- text = p.read_text(encoding="utf-8", errors="replace")
- except Exception as e:
- raise SandboxError(f"read failed: {e}")
- lines = text.splitlines()
- offset = max(0, int(offset)) if offset else 0
- limit = max(1, int(limit)) if limit else 50
- chunk = lines[offset : offset + limit]
- out = "\n".join(chunk)
- if len(out) > self.max_read_bytes:
- out = out[: self.max_read_bytes] + "\n...[truncated]"
- return out
-
- def edit(self, file_path: str, new_string: str, old_string: str = "") -> str:
- p = self._resolve(file_path)
- new_string = new_string or ""
- if not old_string:
- # Create or fully overwrite.
- p.parent.mkdir(parents=True, exist_ok=True)
- p.write_text(new_string, encoding="utf-8")
- return f"wrote {file_path} ({len(new_string)} chars)"
- # Replace mode: old_string must appear exactly once.
- if not p.exists():
- raise SandboxError(f"edit-replace requires existing file: {file_path}")
- text = p.read_text(encoding="utf-8", errors="replace")
- count = text.count(old_string)
- if count == 0:
- raise SandboxError(f"old_string not found in {file_path}")
- if count > 1:
- raise SandboxError(f"old_string ambiguous in {file_path} ({count} matches)")
- p.write_text(text.replace(old_string, new_string, 1), encoding="utf-8")
- return f"edited {file_path}"
-
- def grep(self, pattern: str, path: str = ".") -> str:
- if not isinstance(pattern, str) or not pattern:
- raise SandboxError("pattern must be a non-empty string")
- try:
- regex = re.compile(pattern)
- except re.error as e:
- raise SandboxError(f"bad regex: {e}")
- root = self._resolve(path)
- ws = self.workspace.resolve()
- if root.is_file():
- files = [root]
- elif root.is_dir():
- files = [p for p in root.rglob("*") if p.is_file()]
- else:
- raise SandboxError(f"path not found: {path}")
- hits: list[str] = []
- for f in files:
- try:
- for ln_no, line in enumerate(f.read_text(encoding="utf-8", errors="replace").splitlines(), 1):
- if regex.search(line):
- rel = f.resolve().relative_to(ws)
- hits.append(f"{rel}:{ln_no}:{line[:200]}")
- if len(hits) >= 50:
- return "\n".join(hits) + "\n...[truncated]"
- except Exception:
- continue
- return "\n".join(hits) if hits else "(no matches)"
-
- def bash(self, command: str) -> str:
- if not isinstance(command, str) or not command.strip():
- raise SandboxError("command must be a non-empty string")
- if BASH_BLOCKED_RE.search(command):
- raise SandboxError("blocked command (rm/mv/sudo/curl/...) — denied")
- # Tokenize with quote-awareness so `python -c '... ; ...'` does not
- # accidentally trip the `;` separator. Then identify the start of each
- # command segment (token following an unquoted separator) and check
- # only those tokens against the whitelist.
- try:
- lex = shlex.shlex(command, posix=True, punctuation_chars=";|&")
- lex.whitespace_split = True
- tokens = list(lex)
- except ValueError as e:
- raise SandboxError(f"unparseable shell command: {e}")
- SEPARATORS = {";", "|", "||", "&", "&&"}
- next_is_cmd = True
- for tok in tokens:
- if tok in SEPARATORS:
- next_is_cmd = True
- continue
- if next_is_cmd:
- base = tok.split("/")[-1]
- if base not in BASH_ALLOWED:
- raise SandboxError(f"command not in whitelist: {base!r}")
- next_is_cmd = False
- # Spawn firejail.
- cmd = FIREJAIL_BASE + [
- f"--private={self.workspace}",
- f"--timeout=00:00:{int(self.bash_timeout):02d}",
- "--",
- "bash", "-c", f"cd ~ && {command}",
- ]
- try:
- r = subprocess.run(
- cmd,
- capture_output=True,
- text=True,
- timeout=self.bash_timeout + 2.0,
- )
- except subprocess.TimeoutExpired:
- raise SandboxError(f"bash timed out after {self.bash_timeout}s")
- out = (r.stdout + ("\n" + r.stderr if r.stderr else "")).strip()
- if len(out) > self.max_read_bytes:
- out = out[: self.max_read_bytes] + "\n...[truncated]"
- if not out:
- out = f"(no output, exit={r.returncode})"
- return out
-
- # ------------------------------------------------------------------
- # ds79c / demycap-style tool aliases. The model trained on ds79c emits
- # a wider tool vocabulary (`execute`, `read_file`, `write_file`,
- # `submit_report`, `save_mermaid`, `write_todos`, …) and hard-fails if
- # the runtime can't honour those names. Most map cleanly to the
- # primitives above; a few are no-op stubs that return success so the
- # multi-turn loop keeps moving.
- def execute(self, command: str) -> str:
- """Alias for `bash` — ds79c's tool name for shell execution."""
- return self.bash(command)
-
- def read_file(self, file_path: str, offset: int = 0, limit: int = 50) -> str:
- return self.read(file_path, offset=offset, limit=limit)
-
- def edit_file(self, file_path: str, new_string: str = "", old_string: str = "") -> str:
- return self.edit(file_path, new_string, old_string)
-
- def write_file(self, file_path: str, content: str = "") -> str:
- # demycap's write_file is "create or fully overwrite".
- return self.edit(file_path, content, "")
-
- def ls_path(self, path: str = ".") -> str:
- """`ls` tool — list directory contents inside the workspace."""
- # Validate that the path resolves inside the workspace (raises if not),
- # then delegate to `bash` for `ls -la`. Use shell-quoting via a fixed
- # arg pattern to avoid arbitrary command injection through path.
- target = self._resolve(path)
- rel = target.relative_to(self.workspace) if target != self.workspace else Path(".")
- return self.bash(f"ls -la {shlex.quote(str(rel) or '.')}")
-
- def glob(self, pattern: str = "*", path: str = ".") -> str:
- """`glob` tool — find files matching a glob pattern."""
- if not isinstance(pattern, str) or not pattern:
- raise SandboxError("pattern must be a non-empty string")
- # Delegate to bash `find` with -name.
- target = self._resolve(path)
- rel = target.relative_to(self.workspace) if target != self.workspace else Path(".")
- return self.bash(f"find {shlex.quote(str(rel) or '.')} -name {shlex.quote(pattern)}")
-
- def save_mermaid(self, file_path: str, content: str = "") -> str:
- """`save_mermaid` — write a `.mmd` file. We don't run mermaid CLI to
- validate (would require node + the mermaid binary inside firejail);
- we accept the content and persist it so trajectories that reference
- the file later via `read_file` continue to work."""
- if not file_path.endswith(".mmd"):
- raise SandboxError(f"save_mermaid expects a .mmd path, got {file_path!r}")
- return self.edit(file_path, content, "")
-
- def submit_report(self, report_markdown: str = "", confidence: float = 0.0,
- **extra) -> str:
- """`submit_report` — terminating tool. Captures the payload so the
- task's `check` function can grade it; returns a canned ack so the
- loop terminates cleanly on the next assistant turn."""
- self.submitted_report = {
- "report_markdown": report_markdown or "",
- "confidence": confidence,
- **{k: v for k, v in extra.items() if k not in ("report_markdown", "confidence")},
- }
- return f"report submitted ({len(report_markdown or '')} chars, confidence={confidence})"
-
- def write_todos(self, todos: list = None, **kw) -> str: # noqa: D401
- """No-op planning stub. Records todos so update_todos can mutate them."""
- self._todos = list(todos) if todos else []
- return f"created {len(self._todos)} todo(s)"
-
- def update_todos(self, todos: list = None, **kw) -> str:
- self._todos = list(todos) if todos else self._todos
- return f"updated todos ({len(self._todos)} item(s))"
-
- def update_todo(self, **kw) -> str:
- return f"updated todo: {kw}"
-
- def task(self, description: str = "", **kw) -> str:
- """Subagent-spawn stub. Returns a canned 'no result' so the parent
- agent has to do the work itself — this is fine for our eval where
- we want to score the outer agent's reasoning anyway."""
- return f"task subagent unavailable in this sandbox; do the work inline. (description was: {description[:200]})"
-
- # ------------------------------------------------------------------
- # Setup helper for task definitions
- def write_files(self, files: dict[str, str]) -> None:
- """Bulk-create files with given relative paths and contents."""
- for rel, body in files.items():
- p = self._resolve(rel)
- p.parent.mkdir(parents=True, exist_ok=True)
- p.write_text(body, encoding="utf-8")
diff --git a/vendor/nanoai/agentic/eval_strict.py b/vendor/nanoai/agentic/eval_strict.py
deleted file mode 100644
index 625dcee1bf997d6010951bb6cb671a1442c1683f..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/eval_strict.py
+++ /dev/null
@@ -1,497 +0,0 @@
-"""Strict post-scoring for agent_eval_mt transcripts.
-
-The task checkers in eval_tasks.py intentionally measure weak, sandbox-level
-success. This module adds a behavioral layer for the default 38 tasks: a weak
-pass is only a strict pass if the rollout is not a loop, has no tool errors or
-truncation, and uses the observed tool result to answer or verify the task.
-"""
-from __future__ import annotations
-
-import json
-import re
-from dataclasses import dataclass
-from typing import Any, Mapping
-
-
-@dataclass(frozen=True)
-class StrictScore:
- status: str
- reason: str
-
- @property
- def strict_passed(self) -> bool:
- return self.status == "strict_pass"
-
- @property
- def partial_pass(self) -> bool:
- return self.status == "partial_pass"
-
- def as_dict(self) -> dict[str, object]:
- return {
- "strict_status": self.status,
- "strict_passed": self.strict_passed,
- "partial_pass": self.partial_pass,
- "strict_reason": self.reason,
- }
-
-
-READ_EXPECT: dict[str, tuple[str, tuple[str, ...], tuple[str, ...]]] = {
- "read_short": ("README.md", ("nanocode",), ("nanocode",)),
- "read_config": ("config.toml", ("0.1.0",), ("0.1.0",)),
- "read_script": ("run.sh", ("echo hello", "hello world"), ("echo hello",)),
- "read_pylib": ("lib.py", ("def add", "add(a, b)"), ("def add",)),
- "read_data": ("data.csv", ("alice", "bob", "name,age"), ("alice",)),
- "read_nested": ("src/main.py", ("cpu_count", "os.cpu_count"), ("cpu_count",)),
- "read_long": ("log.txt", ("line 0", "line 1"), ("line 0",)),
-}
-
-GREP_EXPECT: dict[str, tuple[tuple[str, ...], tuple[str, ...]]] = {
- "grep_word": (("a.py", "c.py"), ("a.py", "c.py")),
- "grep_class": (("model.py", "GPT"), ("model.py",)),
- "grep_todo": (("x.py", "TODO"), ("x.py",)),
- "grep_imports": (("a.py", "b.py"), ("a.py", "b.py")),
- "grep_nomatch": (("no match", "no matches", "not found", "banana"), ("(no matches)",)),
- "grep_func": (("a.py", "parse"), ("a.py",)),
- "grep_torch": (("src/x.py", "tests/y.py", "torch"), ("src/x.py", "tests/y.py")),
-}
-
-BASH_EXPECT: dict[str, dict[str, tuple[str, ...]]] = {
- "bash_pwd": {"cmd_any": ("pwd",), "answer_any": ("/home/",)},
- "bash_ls": {"cmd_any": ("ls", "find"), "answer_any": ("a.txt", "b.txt")},
- "bash_count_py": {"cmd_any": ("find", "python", "ls"), "answer_any": ("4",)},
- "bash_wc": {"cmd_any": ("wc", "python", "awk"), "answer_any": ("6",)},
- "bash_head": {"cmd_any": ("head", "sed", "python"), "answer_any": ("line0",)},
- "bash_python": {"cmd_any": ("python",), "answer_any": ("4",)},
- "bash_find_py": {"cmd_any": ("find", "ls", "python"), "answer_any": (".py", "a.py", "b.py")},
- "bash_grep_log": {"cmd_any": ("grep", "python", "awk"), "answer_any": ("ERROR",)},
-}
-
-EDIT_EXPECT: dict[str, tuple[str, ...]] = {
- "edit_create_hello": ("hello.py",),
- "edit_create_readme": ("README.md",),
- "edit_create_json": ("config.json",),
- "edit_replace": ("a.py",),
- "edit_create_sh": ("greet.sh",),
- "edit_add_fn": ("util.py",),
- "edit_create_test": ("test_x.py",),
- "edit_overwrite": ("notes.txt",),
-}
-
-SWE_EXPECT: dict[str, dict[str, tuple[str, ...]]] = {
- "swe_fix_typo": {"edit_any": ("prog.py",), "verify_any": ("hello",)},
- "swe_fix_indent": {"edit_any": ("lib.py",), "verify_any": ("7",)},
- "swe_fizzbuzz": {"edit_any": ("fizzbuzz.py",), "verify_any": ("PASS",)},
- "swe_factorial": {"edit_any": ("factorial.py",), "verify_any": ("PASS",)},
- "swe_missing_import": {"edit_any": ("app.py",), "verify_any": ("a/b",)},
- "swe_off_by_one": {"edit_any": ("range_sum.py",), "verify_any": ("PASS",)},
- "swe_replace_print": {"edit_any": ("a.py", "b.py", "c.py"), "verify_any": ("(no matches)", "no output")},
- "swe_grep_then_fix": {"edit_any": ("tok_a.py", "tok_b.py"), "verify_any": ("(no matches)", "no output")},
-}
-
-
-def _turns(record: Mapping[str, Any]) -> list[dict[str, Any]]:
- turns = record.get("transcript") or []
- return turns if isinstance(turns, list) else []
-
-
-def _args(call: Mapping[str, Any]) -> dict[str, Any]:
- raw = call.get("args")
- if raw is None:
- raw = call.get("arguments")
- return raw if isinstance(raw, dict) else {}
-
-
-def _tool_calls(turns: list[dict[str, Any]]) -> list[dict[str, Any]]:
- return [t["tool_call"] for t in turns if isinstance(t, dict) and t.get("tool_call")]
-
-
-def _tool_results(turns: list[dict[str, Any]]) -> list[tuple[int, str, bool]]:
- out: list[tuple[int, str, bool]] = []
- for i, t in enumerate(turns):
- if not isinstance(t, dict) or t.get("tool_result") is None:
- continue
- out.append((i, str(t.get("tool_result") or ""), bool(t.get("error"))))
- return out
-
-
-def _final_text(turns: list[dict[str, Any]]) -> str:
- for t in reversed(turns):
- if isinstance(t, dict) and t.get("role") == "assistant" and not t.get("tool_call"):
- return str(t.get("content") or "")
- return ""
-
-
-def _all_assistant_text(turns: list[dict[str, Any]]) -> str:
- return "\n".join(
- str(t.get("content") or "")
- for t in turns
- if isinstance(t, dict) and t.get("role") == "assistant"
- )
-
-
-def _contains_any(text: str, terms: tuple[str, ...]) -> bool:
- text_l = text.lower()
- return any(term.lower() in text_l for term in terms)
-
-
-def _contains_all(text: str, terms: tuple[str, ...]) -> bool:
- text_l = text.lower()
- return all(term.lower() in text_l for term in terms)
-
-
-def _joined_results(turns: list[dict[str, Any]]) -> str:
- return "\n".join(r for _idx, r, _err in _tool_results(turns))
-
-
-def _call_signature(call: Mapping[str, Any]) -> str:
- return json.dumps(
- {"name": call.get("name"), "args": _args(call)},
- ensure_ascii=False,
- sort_keys=True,
- default=str,
- )
-
-
-def _tool_indices(turns: list[dict[str, Any]]) -> list[tuple[int, dict[str, Any]]]:
- return [
- (i, t["tool_call"])
- for i, t in enumerate(turns)
- if isinstance(t, dict) and t.get("tool_call")
- ]
-
-
-def _tool_result_pairs(turns: list[dict[str, Any]]) -> list[tuple[int, dict[str, Any], int, str, bool]]:
- pairs: list[tuple[int, dict[str, Any], int, str, bool]] = []
- last_call: tuple[int, dict[str, Any]] | None = None
- for i, t in enumerate(turns):
- if not isinstance(t, dict):
- continue
- if t.get("tool_call"):
- last_call = (i, t["tool_call"])
- continue
- if t.get("tool_result") is not None and last_call is not None:
- call_idx, call = last_call
- pairs.append((
- call_idx,
- call,
- i,
- str(t.get("tool_result") or ""),
- bool(t.get("error")),
- ))
- last_call = None
- return pairs
-
-
-def _first_success_result_idx(turns: list[dict[str, Any]], terms: tuple[str, ...], *, all_terms: bool = False) -> int | None:
- for idx, result, _err in _tool_results(turns):
- if all_terms:
- if _contains_all(result, terms):
- return idx
- elif _contains_any(result, terms):
- return idx
- return None
-
-
-def _has_tool_after(turns: list[dict[str, Any]], idx: int) -> bool:
- return any(i > idx for i, _call in _tool_indices(turns))
-
-
-def _path_args(calls: list[dict[str, Any]], name: str) -> list[str]:
- paths: list[str] = []
- for c in calls:
- if c.get("name") != name:
- continue
- a = _args(c)
- p = a.get("file_path") or a.get("path")
- if p is not None:
- paths.append(str(p))
- return paths
-
-
-def _bash_commands(calls: list[dict[str, Any]]) -> list[str]:
- cmds: list[str] = []
- for c in calls:
- if c.get("name") != "Bash":
- continue
- cmd = _args(c).get("command")
- if cmd is not None:
- cmds.append(str(cmd))
- return cmds
-
-
-STOPWORDS = {
- "the", "and", "that", "this", "with", "from", "file", "line", "true",
- "false", "none", "null", "error", "tool", "result", "created", "updated",
- "wrote", "chars", "matches", "match", "output", "return", "import",
-}
-
-
-def _grounding_terms(result: str) -> tuple[str, ...]:
- terms: list[str] = []
- for tok in re.findall(r"[A-Za-z0-9_./-]+", result):
- clean = tok.strip(".,:;()[]{}'\"`")
- if len(clean) < 3 and not clean.isdigit():
- continue
- if clean.lower() in STOPWORDS:
- continue
- terms.append(clean)
- # Preserve order while bounding noisy long outputs.
- return tuple(list(dict.fromkeys(terms))[:12])
-
-
-def _generic_issues(record: Mapping[str, Any], turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- issues: list[str] = []
- if bool(record.get("truncated")):
- issues.append("generation_truncated")
- try:
- n_tool_errors = int(record.get("n_tool_errors") or 0)
- except (TypeError, ValueError):
- n_tool_errors = 0
- if n_tool_errors > 0:
- issues.append(f"tool_errors={n_tool_errors}")
-
- seen: dict[str, int] = {}
- for c in calls:
- sig = _call_signature(c)
- seen[sig] = seen.get(sig, 0) + 1
- repeated = [sig for sig, n in seen.items() if n > 1]
- if repeated:
- issues.append(f"repeated_tool_call={len(repeated)}")
-
- for t in turns:
- if not isinstance(t, dict) or t.get("role") != "assistant" or t.get("tool_call"):
- continue
- content = str(t.get("content") or "")
- if "<|tool_call_start|>" in content or "<|tool_arg|>" in content or "<|tool_val|>" in content:
- issues.append("unparsed_tool_tokens_in_text")
- break
-
- text = _all_assistant_text(turns)
- if re.search(r"\b(\w{4,})\b(?:\s+\1\b){6,}", text, flags=re.IGNORECASE):
- issues.append("assistant_text_repetition")
- return issues
-
-
-def _score_generic_observation_task(
- task_id: str,
- turns: list[dict[str, Any]],
- tool_name: str,
- *,
- require_final_grounding: bool = True,
-) -> list[str]:
- issues: list[str] = []
- pairs = [
- (call_idx, call, result_idx, result, err)
- for call_idx, call, result_idx, result, err in _tool_result_pairs(turns)
- if call.get("name") == tool_name and not err
- ]
- if not pairs:
- issues.append(f"missing_successful_{tool_name.lower()}_result")
- return issues
-
- _call_idx, _call, result_idx, result, _err = pairs[0]
- if _has_tool_after(turns, result_idx):
- issues.append(f"continued_tools_after_{tool_name.lower()}_success")
-
- if not require_final_grounding:
- return issues
-
- final = _final_text(turns)
- if "no match" in result.lower() or "(no matches)" in result.lower():
- if not _contains_any(final, ("no match", "no matches", "not found")):
- issues.append(f"final_answer_not_grounded_in_{tool_name.lower()}_result")
- return issues
-
- terms = _grounding_terms(result)
- if terms and not _contains_any(final, terms):
- issues.append(f"final_answer_not_grounded_in_{tool_name.lower()}_result")
- return issues
-
-
-def _score_family_fallback(task_id: str, turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- if task_id.startswith("read_"):
- return _score_generic_observation_task(task_id, turns, "Read")
- if task_id.startswith("grep_"):
- return _score_generic_observation_task(task_id, turns, "Grep")
- if task_id.startswith("bash_"):
- return _score_generic_observation_task(task_id, turns, "Bash")
- if task_id.startswith("edit_"):
- return _score_edit_generic(calls)
- if task_id.startswith("swe_"):
- return _score_swe_generic(turns, calls)
- return []
-
-
-def _score_read(task_id: str, turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- expected_path, final_terms, result_terms = READ_EXPECT[task_id]
- issues: list[str] = []
- if not any(c.get("name") == "Read" for c in calls):
- issues.append("missing_read_call")
- return issues
- if expected_path not in _path_args(calls, "Read"):
- issues.append(f"read_wrong_path_expected={expected_path}")
- success_idx = _first_success_result_idx(turns, result_terms)
- if success_idx is None:
- issues.append("read_result_not_observed")
- elif _has_tool_after(turns, success_idx):
- issues.append("continued_tools_after_read_success")
- if not _contains_any(_final_text(turns), final_terms):
- issues.append("final_answer_not_grounded_in_read_result")
- return issues
-
-
-def _score_grep(task_id: str, turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- final_terms, result_terms = GREP_EXPECT[task_id]
- issues: list[str] = []
- if not any(c.get("name") == "Grep" for c in calls):
- issues.append("missing_grep_call")
- return issues
- result_text = _joined_results(turns)
- if task_id == "grep_nomatch":
- if not _contains_any(result_text, result_terms):
- issues.append("grep_nomatch_result_not_observed")
- if not _contains_any(_final_text(turns), final_terms):
- issues.append("final_answer_does_not_report_no_match")
- return issues
- if not _contains_all(result_text, result_terms):
- issues.append("grep_result_incomplete")
- success_idx = _first_success_result_idx(turns, result_terms, all_terms=True)
- if success_idx is not None and _has_tool_after(turns, success_idx):
- issues.append("continued_tools_after_grep_success")
- if not _contains_all(_final_text(turns), final_terms):
- issues.append("final_answer_not_grounded_in_grep_result")
- return issues
-
-
-def _score_bash(task_id: str, turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- spec = BASH_EXPECT[task_id]
- issues: list[str] = []
- cmds = _bash_commands(calls)
- if not cmds:
- issues.append("missing_bash_call")
- return issues
- cmd_text = "\n".join(cmds)
- if not _contains_any(cmd_text, spec["cmd_any"]):
- issues.append("bash_command_not_task_relevant")
- result_plus_final = _joined_results(turns) + "\n" + _final_text(turns)
- if not _contains_any(result_plus_final, spec["answer_any"]):
- issues.append("bash_answer_not_observed")
- success_idx = _first_success_result_idx(turns, spec["answer_any"])
- if success_idx is not None and _has_tool_after(turns, success_idx):
- issues.append("continued_tools_after_bash_success")
- if not _contains_any(_final_text(turns), spec["answer_any"]):
- issues.append("final_answer_not_grounded_in_bash_result")
- return issues
-
-
-def _score_edit(task_id: str, calls: list[dict[str, Any]]) -> list[str]:
- expected_paths = EDIT_EXPECT[task_id]
- issues: list[str] = []
- if not any(c.get("name") == "Edit" for c in calls):
- issues.append("missing_edit_call")
- return issues
- edit_paths = _path_args(calls, "Edit")
- if not any(p in expected_paths for p in edit_paths):
- issues.append(f"edit_wrong_path_expected_one_of={','.join(expected_paths)}")
- for c in calls:
- if c.get("name") != "Edit":
- continue
- new_string = str(_args(c).get("new_string") or "")
- if task_id == "edit_create_test" and re.search(r"assert\s*\(\s*False|assert\s+False", new_string):
- issues.append("test_contains_unconditional_false_assert")
- if len(new_string) > 80 and re.search(r"\b(\w{3,})\b(?:\s+\1\b){5,}", new_string, flags=re.IGNORECASE):
- issues.append("edit_content_repetition")
- return issues
-
-
-def _score_edit_generic(calls: list[dict[str, Any]]) -> list[str]:
- issues: list[str] = []
- if not any(c.get("name") == "Edit" for c in calls):
- issues.append("missing_edit_call")
- return issues
- for c in calls:
- if c.get("name") != "Edit":
- continue
- new_string = str(_args(c).get("new_string") or "")
- if re.search(r"assert\s*\(\s*False|assert\s+False", new_string):
- issues.append("test_contains_unconditional_false_assert")
- if len(new_string) > 80 and re.search(r"\b(\w{3,})\b(?:\s+\1\b){5,}", new_string, flags=re.IGNORECASE):
- issues.append("edit_content_repetition")
- return issues
-
-
-def _score_swe(task_id: str, turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- spec = SWE_EXPECT[task_id]
- names = [str(c.get("name") or "") for c in calls]
- issues: list[str] = []
- if "Bash" not in names:
- issues.append("swe_missing_verifier_bash")
- if task_id == "swe_grep_then_fix" and "Grep" not in names:
- issues.append("swe_missing_initial_grep")
- edit_paths = _path_args(calls, "Edit")
- if not edit_paths and "Edit" not in names:
- issues.append("swe_missing_edit")
- elif not any(p in spec["edit_any"] for p in edit_paths):
- issues.append(f"swe_edit_wrong_path_expected_one_of={','.join(spec['edit_any'])}")
- result_text = _joined_results(turns)
- if not _contains_any(result_text, spec["verify_any"]):
- issues.append("swe_verifier_success_not_observed")
- return issues
-
-
-def _score_swe_generic(turns: list[dict[str, Any]], calls: list[dict[str, Any]]) -> list[str]:
- names = [str(c.get("name") or "") for c in calls]
- issues: list[str] = []
- if "Edit" not in names:
- issues.append("swe_missing_edit")
- if "Bash" not in names:
- issues.append("swe_missing_verifier_bash")
- bash_pairs = [
- (call_idx, call, result_idx, result, err)
- for call_idx, call, result_idx, result, err in _tool_result_pairs(turns)
- if call.get("name") == "Bash" and not err
- ]
- if "Bash" in names and not bash_pairs:
- issues.append("swe_missing_successful_verifier_result")
- result_text = "\n".join(result for _ci, _c, _ri, result, _err in bash_pairs)
- bad_markers = ("Traceback", "AssertionError", "SyntaxError", "NameError", "IndentationError")
- if any(m in result_text for m in bad_markers):
- issues.append("swe_verifier_result_contains_error")
- return issues
-
-
-def score_strict_record(record: Mapping[str, Any]) -> StrictScore:
- """Score one agent_eval_mt row or transcript record.
-
- Returns:
- - fail: weak checker failed.
- - partial_pass: weak checker passed, but behavioral strict gates failed.
- - strict_pass: weak checker passed and no strict issue was found.
- """
- if not bool(record.get("passed")):
- reason = str(record.get("reason") or "weak checker failed")
- return StrictScore("fail", f"base_fail: {reason}")
-
- turns = _turns(record)
- calls = _tool_calls(turns)
- task_id = str(record.get("task_id") or "")
- issues = _generic_issues(record, turns, calls)
-
- if task_id in READ_EXPECT:
- issues.extend(_score_read(task_id, turns, calls))
- elif task_id in GREP_EXPECT:
- issues.extend(_score_grep(task_id, turns, calls))
- elif task_id in BASH_EXPECT:
- issues.extend(_score_bash(task_id, turns, calls))
- elif task_id in EDIT_EXPECT:
- issues.extend(_score_edit(task_id, calls))
- elif task_id in SWE_EXPECT:
- issues.extend(_score_swe(task_id, turns, calls))
- else:
- issues.extend(_score_family_fallback(task_id, turns, calls))
- if not turns:
- issues.append("missing_transcript_for_strict_score")
-
- if issues:
- return StrictScore("partial_pass", "; ".join(dict.fromkeys(issues)))
- return StrictScore("strict_pass", "")
diff --git a/vendor/nanoai/agentic/eval_tasks.py b/vendor/nanoai/agentic/eval_tasks.py
deleted file mode 100644
index ab4f44764675921fe7f1428fe50826b8a394fa87..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/eval_tasks.py
+++ /dev/null
@@ -1,517 +0,0 @@
-"""30 deterministic multi-turn agent eval tasks.
-
-Each task: setup the sandbox workspace, give the model a prompt, then check
-post-conditions. `check` receives the sandbox path and the full Transcript
-and returns (passed, reason).
-
-Coverage: Read 7, Edit 8, Grep 7, Bash 8.
-"""
-from __future__ import annotations
-
-import re
-from dataclasses import dataclass, field
-from pathlib import Path
-from typing import Callable, Optional
-
-from .eval_sandbox import Sandbox
-
-
-@dataclass
-class Transcript:
- """Captures everything that happened during a task run."""
- user_prompt: str
- turns: list[dict] = field(default_factory=list) # {role, content, tool_call?, tool_result?, error?}
-
- @property
- def assistant_text(self) -> str:
- """Concatenation of all assistant content (non-tool-call portions)."""
- return "\n".join(t["content"] for t in self.turns if t["role"] == "assistant" and t.get("content"))
-
- @property
- def last_assistant_text(self) -> str:
- for t in reversed(self.turns):
- if t["role"] == "assistant" and t.get("content"):
- return t["content"]
- return ""
-
- @property
- def tool_calls(self) -> list[dict]:
- return [t["tool_call"] for t in self.turns if t.get("tool_call")]
-
- @property
- def tool_results(self) -> list[str]:
- return [t["tool_result"] for t in self.turns if t.get("tool_result") is not None]
-
-
-CheckFn = Callable[[Sandbox, Transcript], tuple[bool, str]]
-
-
-@dataclass
-class Task:
- task_id: str
- prompt: str
- setup: Callable[[Sandbox], None]
- check: CheckFn
- tools_expected: tuple[str, ...] = ()
- max_turns: int = 5
- # Optional per-task system prompt override. When None, eval_loop uses the
- # default SYSTEM_PROMPT (the v3 nanocode persona with Read/Edit tools).
- # Pcap diagnostic tasks set this to the demycap-style
- # block so the trained-on-ds79c model activates its diagnostic conditional.
- system_prompt: Optional[str] = None
-
-
-# ----------------------------------------------------------------------
-# Helpers used by check functions
-def _called(transcript: Transcript, name: str) -> bool:
- return any(c.get("name") == name for c in transcript.tool_calls)
-
-
-def _any_result_contains(transcript: Transcript, needle: str) -> bool:
- return any(needle in (r or "") for r in transcript.tool_results)
-
-
-def _file(sb: Sandbox, rel: str) -> Optional[str]:
- p = (sb.workspace / rel)
- if not p.exists() or not p.is_file():
- return None
- return p.read_text(encoding="utf-8", errors="replace")
-
-
-# ======================================================================
-# Read tasks (7)
-# ----------------------------------------------------------------------
-def _read_tasks() -> list[Task]:
- def make_simple(task_id, fname, content, needle, prompt):
- def setup(sb):
- sb.write_files({fname: content})
- def check(sb, t):
- ok = _called(t, "Read") and _any_result_contains(t, needle)
- return ok, "" if ok else f"Read({fname}) result missing {needle!r}"
- return Task(task_id, prompt, setup, check, ("Read",))
-
- return [
- make_simple("read_short", "README.md", "Welcome to nanocode.\n", "nanocode",
- "show me the contents of README.md"),
- make_simple("read_config", "config.toml", "[tool]\nname = 'nanocode'\nversion = '0.1.0'\n",
- "0.1.0", "what version is in config.toml?"),
- make_simple("read_script", "run.sh", "#!/bin/bash\necho hello world\n", "echo hello",
- "what does run.sh do? read it"),
- make_simple("read_pylib", "lib.py", "def add(a, b):\n return a + b\n", "def add",
- "show me lib.py"),
- make_simple("read_data", "data.csv", "name,age\nalice,30\nbob,25\n", "alice",
- "read data.csv and tell me the contents"),
- make_simple("read_nested", "src/main.py", "import os\nprint(os.cpu_count())\n", "cpu_count",
- "read src/main.py"),
- make_simple("read_long", "log.txt", "\n".join(f"line {i}" for i in range(40)) + "\n", "line 0",
- "read the first few lines of log.txt"),
- ]
-
-
-# ======================================================================
-# Edit tasks (8)
-# ----------------------------------------------------------------------
-def _edit_tasks() -> list[Task]:
- def t_create_hello() -> Task:
- def setup(sb): pass
- def check(sb, t):
- body = _file(sb, "hello.py")
- if body is None:
- return False, "hello.py not created"
- ok = "hi" in body and ("print" in body)
- return ok, "" if ok else "hello.py exists but doesn't print 'hi'"
- return Task("edit_create_hello", "create a file hello.py that prints 'hi'", setup, check, ("Edit",))
-
- def t_create_readme() -> Task:
- def setup(sb): pass
- def check(sb, t):
- body = _file(sb, "README.md")
- if body is None:
- return False, "README.md not created"
- ok = len(body.strip()) > 0
- return ok, "" if ok else "README.md is empty"
- return Task("edit_create_readme", "create a README.md with a project description", setup, check, ("Edit",))
-
- def t_create_config() -> Task:
- def setup(sb): pass
- def check(sb, t):
- import json as _json
- body = _file(sb, "config.json")
- if body is None:
- return False, "config.json not created"
- try:
- obj = _json.loads(body)
- except Exception as e:
- return False, f"config.json not valid JSON: {type(e).__name__}: {str(e)[:80]}"
- if not isinstance(obj, dict):
- return False, f"config.json must be a JSON object (got {type(obj).__name__})"
- return True, ""
- return Task("edit_create_json", "create config.json with a json object", setup, check, ("Edit",))
-
- def t_replace_in_file() -> Task:
- def setup(sb): sb.write_files({"a.py": "x = 1\ny = 2\n"})
- def check(sb, t):
- body = _file(sb, "a.py") or ""
- ok = "x = 42" in body
- return ok, "" if ok else "expected x = 42 in a.py"
- return Task("edit_replace", "in a.py, change `x = 1` to `x = 42`", setup, check, ("Edit",))
-
- def t_create_script() -> Task:
- def setup(sb): pass
- def check(sb, t):
- body = _file(sb, "greet.sh") or ""
- ok = "#!" in body or "echo" in body
- return ok, "" if ok else "greet.sh missing shebang or echo"
- return Task("edit_create_sh", "create greet.sh that echoes 'hello there'", setup, check, ("Edit",))
-
- def t_add_function() -> Task:
- def setup(sb): sb.write_files({"util.py": "def hello():\n return 'hi'\n"})
- def check(sb, t):
- body = _file(sb, "util.py") or ""
- ok = "def hello" in body and ("def " in body[body.index("def hello") + len("def hello"):])
- return ok, "" if ok else "util.py should keep hello() and add another function"
- return Task("edit_add_fn", "add a new function `goodbye` to util.py that returns 'bye'", setup, check, ("Edit",))
-
- def t_create_test() -> Task:
- def setup(sb): pass
- def check(sb, t):
- body = _file(sb, "test_x.py") or ""
- ok = ("def test_" in body) or ("assert" in body)
- return ok, "" if ok else "test_x.py missing test"
- return Task("edit_create_test", "create test_x.py with one assert", setup, check, ("Edit",))
-
- def t_overwrite() -> Task:
- def setup(sb): sb.write_files({"notes.txt": "old content\n"})
- def check(sb, t):
- body = _file(sb, "notes.txt") or ""
- ok = "old content" not in body and len(body.strip()) > 0
- return ok, "" if ok else "notes.txt still has old content"
- return Task("edit_overwrite", "overwrite notes.txt with completely new content", setup, check, ("Edit",))
-
- return [t_create_hello(), t_create_readme(), t_create_config(), t_replace_in_file(),
- t_create_script(), t_add_function(), t_create_test(), t_overwrite()]
-
-
-# ======================================================================
-# Grep tasks (7)
-# ----------------------------------------------------------------------
-def _grep_tasks() -> list[Task]:
- def t_find_word() -> Task:
- def setup(sb):
- sb.write_files({
- "a.py": "from tokenizer import x\n",
- "b.py": "import os\n",
- "c.py": "tokenizer = None\n",
- })
- def check(sb, t):
- ok = _called(t, "Grep") and _any_result_contains(t, "a.py") and _any_result_contains(t, "c.py")
- return ok, "" if ok else "Grep didn't find both a.py and c.py"
- return Task("grep_word", "find files mentioning 'tokenizer'", setup, check, ("Grep",))
-
- def t_find_class() -> Task:
- def setup(sb):
- sb.write_files({
- "model.py": "class GPT:\n pass\n",
- "util.py": "class Utility:\n pass\n",
- })
- def check(sb, t):
- ok = _called(t, "Grep") and _any_result_contains(t, "model.py")
- return ok, "" if ok else "Grep didn't locate class GPT in model.py"
- return Task("grep_class", "where is the GPT class defined?", setup, check, ("Grep",))
-
- def t_find_todo() -> Task:
- def setup(sb):
- sb.write_files({
- "x.py": "# TODO: implement\nx = 1\n",
- "y.py": "y = 2\n",
- })
- def check(sb, t):
- ok = _called(t, "Grep") and _any_result_contains(t, "x.py")
- return ok, "" if ok else "Grep didn't find TODO in x.py"
- return Task("grep_todo", "find all TODO comments", setup, check, ("Grep",))
-
- def t_count_imports() -> Task:
- def setup(sb):
- sb.write_files({
- "a.py": "import os\nimport sys\n",
- "b.py": "import json\n",
- "c.txt": "no imports\n",
- })
- def check(sb, t):
- ok = _called(t, "Grep") and _any_result_contains(t, "a.py") and _any_result_contains(t, "b.py")
- return ok, "" if ok else "Grep should find imports in a.py and b.py"
- return Task("grep_imports", "find all import statements in python files", setup, check, ("Grep",))
-
- def t_no_match() -> Task:
- def setup(sb):
- sb.write_files({"a.py": "print('hi')\n"})
- def check(sb, t):
- ok = _called(t, "Grep")
- return ok, "" if ok else "expected at least one Grep call (no-match acceptable)"
- return Task("grep_nomatch", "search for the word 'banana' in the project", setup, check, ("Grep",))
-
- def t_specific_function() -> Task:
- def setup(sb):
- sb.write_files({
- "a.py": "def parse(x):\n return x\n",
- "b.py": "def render(y):\n return y\n",
- })
- def check(sb, t):
- ok = _called(t, "Grep") and _any_result_contains(t, "a.py")
- return ok, "" if ok else "Grep should find parse() in a.py"
- return Task("grep_func", "find where 'parse' is defined", setup, check, ("Grep",))
-
- def t_path_specific() -> Task:
- def setup(sb):
- sb.write_files({
- "src/x.py": "import torch\n",
- "tests/y.py": "import torch\n",
- })
- def check(sb, t):
- ok = _called(t, "Grep")
- return ok, "" if ok else "expected a Grep call"
- return Task("grep_torch", "find files using torch", setup, check, ("Grep",))
-
- return [t_find_word(), t_find_class(), t_find_todo(), t_count_imports(),
- t_no_match(), t_specific_function(), t_path_specific()]
-
-
-# ======================================================================
-# Bash tasks (8)
-# ----------------------------------------------------------------------
-def _bash_tasks() -> list[Task]:
- def t_pwd() -> Task:
- def setup(sb): pass
- def check(sb, t):
- ok = _called(t, "Bash")
- return ok, "" if ok else "expected a Bash call"
- return Task("bash_pwd", "print the working directory", setup, check, ("Bash",))
-
- def t_ls() -> Task:
- def setup(sb): sb.write_files({"a.txt": "", "b.txt": ""})
- def check(sb, t):
- ok = _called(t, "Bash") and _any_result_contains(t, "a.txt")
- return ok, "" if ok else "Bash result should mention a.txt"
- return Task("bash_ls", "list the files in the current directory", setup, check, ("Bash",))
-
- def t_count_py() -> Task:
- def setup(sb):
- files = {f"data/{n}.py": "" for n in "abcd"}
- files["data/notes.txt"] = ""
- sb.write_files(files)
- def check(sb, t):
- ok = _called(t, "Bash") and (_any_result_contains(t, "4") or "4" in t.last_assistant_text)
- return ok, "" if ok else "expected count of 4 in Bash result or final answer"
- return Task("bash_count_py", "how many python files are in the data/ directory?", setup, check, ("Bash",))
-
- def t_word_count() -> Task:
- def setup(sb): sb.write_files({"essay.txt": "one two three four five six\n"})
- def check(sb, t):
- ok = _called(t, "Bash")
- return ok, "" if ok else "expected wc in Bash"
- return Task("bash_wc", "how many words are in essay.txt?", setup, check, ("Bash",))
-
- def t_head() -> Task:
- def setup(sb):
- sb.write_files({"big.log": "\n".join(f"line{i}" for i in range(100)) + "\n"})
- def check(sb, t):
- ok = _called(t, "Bash") and _any_result_contains(t, "line0")
- return ok, "" if ok else "Bash result should include line0"
- return Task("bash_head", "show me the first few lines of big.log", setup, check, ("Bash",))
-
- def t_python_inline() -> Task:
- def setup(sb): pass
- def check(sb, t):
- ok = _called(t, "Bash") and (_any_result_contains(t, "4") or "4" in t.last_assistant_text)
- return ok, "" if ok else "expected '4' in result"
- return Task("bash_python", "use python to compute 2+2", setup, check, ("Bash",))
-
- def t_find_py() -> Task:
- def setup(sb):
- sb.write_files({"src/a.py": "", "src/b.py": "", "src/notes.txt": ""})
- def check(sb, t):
- ok = _called(t, "Bash") and _any_result_contains(t, ".py")
- return ok, "" if ok else "Bash result should list .py files"
- return Task("bash_find_py", "find all python files under src/", setup, check, ("Bash",))
-
- def t_grep_via_bash() -> Task:
- def setup(sb): sb.write_files({"log.txt": "ok\nERROR: failed\nok\n"})
- def check(sb, t):
- ok = _called(t, "Bash") and _any_result_contains(t, "ERROR")
- return ok, "" if ok else "expected ERROR line via grep"
- return Task("bash_grep_log", "find ERROR lines in log.txt", setup, check, ("Bash",))
-
- return [t_pwd(), t_ls(), t_count_py(), t_word_count(), t_head(),
- t_python_inline(), t_find_py(), t_grep_via_bash()]
-
-
-# ======================================================================
-# SWE-style multi-turn tasks (8)
-# These force a real Edit → Bash run → react loop. Pass condition = the
-# final filesystem state actually executes correctly, not whether the model
-# produced any tool call. d6 will likely score 0 here; this suite is
-# designed to discriminate d12+ models.
-# ----------------------------------------------------------------------
-def _swe_tasks() -> list[Task]:
- def _bash_ok(sb, cmd, must_contain="", must_not_contain="Traceback") -> tuple[bool, str]:
- """Run `cmd` in the sandbox and return (passed, last 200 chars of output)."""
- try:
- out = sb.bash(cmd)
- except Exception as e:
- return False, f"bash crashed: {e}"
- if must_not_contain and must_not_contain in out:
- return False, f"got {must_not_contain}: {out[:200]}"
- if must_contain and must_contain not in out:
- return False, f"missing {must_contain!r} in: {out[:200]}"
- return True, out[:200]
-
- def t_fix_typo() -> Task:
- def setup(sb): sb.write_files({"prog.py": "prnit('hello')\n"})
- def check(sb, t):
- ok, msg = _bash_ok(sb, "python3 prog.py", must_contain="hello")
- return ok, msg
- return Task(
- "swe_fix_typo",
- "the file prog.py should print 'hello' but doesn't run — read it, fix the typo, then verify by running it with python3",
- setup, check, ("Read", "Edit", "Bash"), max_turns=8,
- )
-
- def t_fix_indent() -> Task:
- def setup(sb): sb.write_files({"lib.py": "def go(x):\n if x:\n return x\n return 0\n"})
- def check(sb, t):
- ok, msg = _bash_ok(sb, "python3 -c 'import lib; print(lib.go(7))'", must_contain="7")
- return ok, msg
- return Task(
- "swe_fix_indent",
- "lib.py raises an IndentationError when imported — read it, fix it, then verify with `python -c 'import lib; print(lib.go(7))'`",
- setup, check, ("Read", "Edit", "Bash"), max_turns=8,
- )
-
- def t_fizzbuzz() -> Task:
- def setup(sb):
- sb.write_files({
- "fizzbuzz.py": "# implement fizzbuzz from 1 to 15 here, one entry per line\n",
- "test_fizzbuzz.py": (
- "import subprocess\n"
- "out = subprocess.run(['python3','fizzbuzz.py'], capture_output=True, text=True).stdout.strip().split('\\n')\n"
- "exp=[]\n"
- "for i in range(1,16):\n"
- " if i%15==0: exp.append('FizzBuzz')\n"
- " elif i%3==0: exp.append('Fizz')\n"
- " elif i%5==0: exp.append('Buzz')\n"
- " else: exp.append(str(i))\n"
- "assert out==exp, f'got {out!r} != {exp!r}'\n"
- "print('PASS')\n"
- ),
- })
- def check(sb, t):
- ok, msg = _bash_ok(sb, "python3 test_fizzbuzz.py", must_contain="PASS")
- return ok, msg
- return Task(
- "swe_fizzbuzz",
- "fizzbuzz.py is empty. implement it so test_fizzbuzz.py passes (it prints PASS on success)",
- setup, check, ("Edit", "Bash"), max_turns=8,
- )
-
- def t_factorial() -> Task:
- def setup(sb):
- sb.write_files({
- "factorial.py": "def factorial(n):\n pass # implement me\n",
- "test_factorial.py": (
- "from factorial import factorial\n"
- "assert factorial(0)==1; assert factorial(1)==1; assert factorial(5)==120\n"
- "assert factorial(7)==5040\n"
- "print('PASS')\n"
- ),
- })
- def check(sb, t):
- ok, msg = _bash_ok(sb, "python3 test_factorial.py", must_contain="PASS")
- return ok, msg
- return Task(
- "swe_factorial",
- "implement factorial in factorial.py so test_factorial.py passes",
- setup, check, ("Read", "Edit", "Bash"), max_turns=8,
- )
-
- def t_missing_import() -> Task:
- def setup(sb):
- sb.write_files({"app.py": "print(os.path.join('a', 'b'))\n"})
- def check(sb, t):
- ok, msg = _bash_ok(sb, "python3 app.py", must_contain="a/b")
- return ok, msg
- return Task(
- "swe_missing_import",
- "app.py crashes — diagnose and fix it, then run it (expected output: a/b)",
- setup, check, ("Read", "Edit", "Bash"), max_turns=8,
- )
-
- def t_off_by_one() -> Task:
- def setup(sb):
- sb.write_files({
- "range_sum.py": "def range_sum(n):\n return sum(range(1, n))\n",
- "test_sum.py": (
- "from range_sum import range_sum\n"
- "assert range_sum(5)==15, f'expected 15, got {range_sum(5)}'\n"
- "assert range_sum(10)==55\n"
- "print('PASS')\n"
- ),
- })
- def check(sb, t):
- ok, msg = _bash_ok(sb, "python3 test_sum.py", must_contain="PASS")
- return ok, msg
- return Task(
- "swe_off_by_one",
- "test_sum.py fails. read range_sum.py, fix the bug, and verify by running test_sum.py",
- setup, check, ("Read", "Edit", "Bash"), max_turns=8,
- )
-
- def t_replace_print() -> Task:
- def setup(sb):
- sb.write_files({
- "a.py": "print('hello')\n",
- "b.py": "x = 1\nprint('done', x)\n",
- "c.py": "# already clean\n",
- })
- def check(sb, t):
- # No raw `print(` should remain in any .py file
- try:
- out = sb.bash("grep -rn 'print(' a.py b.py c.py")
- except Exception as e:
- return False, f"grep crashed: {e}"
- if "no output" in out or "(no matches)" in out:
- ok = True
- else:
- ok = "print(" not in out
- return ok, "" if ok else f"raw print() still present: {out[:200]}"
- return Task(
- "swe_replace_print",
- "replace every `print(...)` call across a.py, b.py, c.py with `import logging; logging.info(...)` (or remove the print). then verify with grep that no `print(` remains.",
- setup, check, ("Edit", "Bash"), max_turns=8,
- )
-
- def t_grep_then_fix() -> Task:
- def setup(sb):
- sb.write_files({
- "tok_a.py": "result = encoder.encode_ordinary(text)\n",
- "tok_b.py": "ids = enc.encode_ordinary(s)\n",
- "other.py": "x = 1\n",
- })
- def check(sb, t):
- try:
- out = sb.bash("grep -rn 'encode_ordinary' .")
- except Exception as e:
- return False, f"grep crashed: {e}"
- ok = "no output" in out or "(no matches)" in out or "encode_ordinary" not in out
- return ok, "" if ok else f"encode_ordinary still present: {out[:200]}"
- return Task(
- "swe_grep_then_fix",
- "find every call to `.encode_ordinary(` in *.py and replace it with `.encode(`. verify with grep that none remain.",
- setup, check, ("Grep", "Edit", "Bash"), max_turns=8,
- )
-
- return [t_fix_typo(), t_fix_indent(), t_fizzbuzz(), t_factorial(),
- t_missing_import(), t_off_by_one(), t_replace_print(), t_grep_then_fix()]
-
-
-def all_tasks() -> list[Task]:
- return _read_tasks() + _edit_tasks() + _grep_tasks() + _bash_tasks() + _swe_tasks()
diff --git a/vendor/nanoai/agentic/goal.md b/vendor/nanoai/agentic/goal.md
deleted file mode 100644
index fecc29c58810e65adb26e92e651979f6086da2e5..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/goal.md
+++ /dev/null
@@ -1,83 +0,0 @@
-# Agentic d6 — current experiment goal
-
-> Next phase (tshark pcap diagnosis, multi-depth scaling): `goal_next.md`
-
-
-## North star
-Make a 150M-parameter d6 model genuinely useful as a multi-turn coding agent
-that can `read → reason → edit/run → observe → fix` against a real filesystem
-sandbox. Concretely: lift `agent_eval_mt` pass rate on the 38-task suite from
-the current ~17% to **30%+**, without losing single-turn tool ability or
-ChatCORE.
-
-## Where we are (2026-05-07)
-
-| Run | Single-turn (8 tasks) | Multi-turn 30-task | Multi-turn 38-task (incl 8 SWE) | ChatCORE |
-|---|---|---|---|---|
-| d6 SFT v3 (lc=5, +ChatCORE mix) | 75/75/75 | 5/30 = 16.7% | — | **+21.09%** |
-| d6 SFT v4 (lc=30, otherwise = v3) | 50/50/50 | 2/30 = 6.7% | 3/38 = 7.9% | +18.75% |
-| d6 DPO (from v3 SFT) | 25/12.5 | 0/30 | — | +21.78% |
-| upstream d6 SFT (canonical nanoai) | — | — | — | ~+1.22% (base) / -0.09% (sft) |
-
-**v3 is currently the best SFT we have.** v4 falsified the hypothesis that
-upweighting long-context dialog data would reduce Edit-bias on multi-turn —
-long-ctx is dialog-only and lacks tool-call diversity, so 6× weighting just
-diluted smoltalk+tulu+ChatCORE supervision.
-
-## Failure mode
-
-Edit-bias dominates multi-turn failures: model emits `Edit` when the task
-needs `Read`/`Grep`/`Bash`. Root cause: tulu-evol (~150k rows, ~21% of mix)
-contains tool calls but virtually all are Edit. Smoltalk has no tool calls.
-Long-context has no tool calls. So at training time the only "tool-call shape"
-the model sees is Edit, and it overgeneralises.
-
-Secondary: model rarely does multi-step reasoning across observations; it
-emits one tool call and either terminates or hallucinates a result. There is
-no training data showing a `tool_call → observation → tool_call → ...` loop.
-
-## Plan (v5)
-
-Add **multi-turn execute-debug** data so the model has actual examples of
-tool-call → observation → tool-call loops:
-
-| Source | Rows | Type | Tokens (mean / p99) |
-|---|---|---|---|
-| `xingyaoww/code-act` (codeact split, full) | 7,139 | data analysis (sqlite/pandas, table queries) | 1,228 / 2,442 |
-| `m-a-p/Code-Feedback` (filtered) | ~10,000 | iterative code creation + debug | 1,375 / 4,294 |
-| **total** | **~17k** | both 100% multi-turn with execution feedback | all ≤8K ✅ |
-
-**Code-Feedback filter:** python-only, `Execution result:` present,
-no `Cell In[..]` (Jupyter-only error), no obvious slop, char ≤ 16K.
-
-**Format adapter** (new file `nanoai/agentic/data/codedebug.py`):
-- Parse markdown ```` ```python ```` and `...` → `python` segment
-- Parse `Execution result:` / `Observation:` user-turns → fold into the
- preceding assistant turn as a `python_output` segment (so the chat sequence
- stays user/assistant/user/assistant)
-- Output rows compatible with `render_agentic_conversation`'s list-content
- branch (already added in v3 work)
-
-## Recipe summary (v5 SFT)
-
-Existing v3 mixture, plus:
-- code-act: 7,139 rows × 1 epoch
-- Code-Feedback (filtered): ~10,000 rows × 1 epoch
-
-≈ 30M new tokens (~4% of mixture). No other changes from v3.
-
-## Success criteria
-
-- Multi-turn 38-task pass rate ≥ 30% (currently 7.9% on v4, 16.7% on v3 with
- smaller 30-task suite)
-- Single-turn ≥ 62.5% (no big regression from v3's 75)
-- ChatCORE ≥ +20% (no big regression from v3's +21.09)
-- Pass on at least 2/8 SWE tasks (v3 pre-SWE: untested; v4: 1/8)
-
-## Open questions / not in v5
-
-- **DPO** is still considered non-viable at d6 scale (mode collapse on tool calls).
-- **GRPO** with SWE pass/fail as reward is the next experiment after v5 SFT.
-- **Routing diversity** (forcing Read/Grep/Bash distinction) — v5 doesn't
- directly address this; if Edit-bias persists post-v5, synthesise routing
- data or cap tulu Edit-only rows.
diff --git a/vendor/nanoai/agentic/goal_next.md b/vendor/nanoai/agentic/goal_next.md
deleted file mode 100644
index e1ec8b1b85a5fd461aea66dbfb042037703324e2..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/goal_next.md
+++ /dev/null
@@ -1,292 +0,0 @@
-# Agentic d6/d12/d20/d24 — next goal: tshark pcap diagnosis
-
-> Full plan: `/home/vader/.claude/plans/zany-marinating-manatee.md`
-> Prior goal (multi-turn code-debug, v5): `goal.md`
-
-## Progress (2026-05-07)
-
-**Phase A v0.1 landed** (skeleton, expansion ongoing):
-- `nanoai/agentic/eval_sandbox.py` — `BASH_ALLOWED` extended with
- `tshark / tcpdump / capinfos / xxd / od / strings`. firejail+tshark roundtrip
- on a 4-packet TCP/DNS pcap verified end-to-end.
-- `data/pcap_corpus/` — 5 demo scenarios with ground-truth JSON + index.json:
- `tcp_rst_storm`, `dns_nxdomain_loop`, `tcp_retransmission`,
- `icmp_host_unreachable`, `clean_baseline` (the no-fault case).
- Builder: `experiments/ksbr_pcap/scripts/build_pcap_corpus.py` (scapy-based; idempotent).
-- `nanoai/agentic/eval_pcap_tasks.py` — `PcapDiagTask` factory + 9-field
- rubric (`parse_report` accepts JSON or labelled-section markdown,
- `score_report` weights presence + evidence-anchor coverage),
- `all_pcap_tasks(corpus_dir)`. Smoke: perfect MD/JSON ⇒ 0.95, garbage ⇒ 0.0,
- 4-of-9 partial ⇒ 0.41.
-- `scripts/agent_eval_mt.py` wired with `--task-set={default,pcap,both}` and
- `--pcap-corpus`.
-
-**Bonus** (correctness gap caught while reviewing):
-- `scripts/agentic_tok_train.py` — added a second inline check that
- every `AGENTIC_SPECIAL_TOKENS` member has a single reserved id, decodes
- back exactly, and `encode_ordinary` does **not** leak that id. Verified on
- the existing trained tokenizer (ids 32762-32767 all in the reserved range
- [32753, 32768) and BPE'd into ~6 ordinary tokens when fed to `encode`).
-
-**Phase A still open:** corpus expansion 5 → 30-50 (TLS alerts, MTU
-blackholes, ARP poisoning, BGP flaps + ~20 from Wireshark sample-captures);
-end-to-end run against d6 SFT v3; optional `--llm-judge` fallback.
-
-Other phases (B–F) untouched.
-
-
-
-## North star
-
-Train d6 / d12 / d20 / d24 nanoai agents that, given a pcap, run tshark in
-a sandbox and produce a 9-field diagnostic report:
-**现象 / 征兆 / 诊断 / 故障 / 根因 / 影响面 / 证据 / 复核方式 / 建议行动**.
-
-Final acceptance: capbench's `demycap` eval on `wuneng`. Local proxy:
-30-50 pcap rubric eval that grades the same 9 fields programmatically.
-
-## Why this is hard at d6 scale
-
-- No tshark in sandbox, no pcap data, no diagnosis-style SFT traces.
-- Diagnostic reports need **up to 64K context** (multiple tshark dumps +
- cross-frame correlation + 9-field write-up). d6 was pretrained at 8K.
-- Current GRPO is REINFORCE; for a multi-component rubric reward we want
- **DAPO** (asymmetric clip + dynamic sampling + overlong shaping).
-- Scaling unknown: which depth crosses the "actually useful" threshold?
-- **Capacity is finite at d6.** Adding pcap data competes with existing
- agentic + ChatCORE supervision. The v4 lc=30 experiment already showed
- that re-weighting one data source by 6× collapsed every other metric —
- we have to design pcap injection so the new capability comes *in
- addition to*, not *by displacing*, prior abilities.
-
-## Sub-goal: non-regression on prior capabilities
-
-Pursuing the pcap target must not silently demolish what v3 / v5
-already buys us. **Every new SFT or RL checkpoint runs the full prior
-eval suite alongside the new pcap eval, and is rejected for promotion
-if it falls below the non-regression floor.**
-
-### Protected baselines (v3 SFT numbers — current best)
-
-| Bucket | Metric | v3 floor | Regression tolerance | Notes |
-|---|---|---:|---:|---|
-| **Single-turn agent_eval** | clean termination % | 75 | −10pp (≥65) | 8-task short eval |
-| | tool-call emitted % | 75 | −10pp (≥65) | |
-| | tool-args well-formed % | 75 | −10pp (≥65) | |
-| **Multi-turn 38-task** | overall pass rate | 16.7% (5/30 on v3 30-task subset) | −3pp on shared 30 / target *up* on full 38 | re-baseline on v5 once available |
-| | Read category | (track) | hold | per-category drift hides aggregate stability |
-| | Edit category | (track) | hold | |
-| | Grep category | (track) | hold | |
-| | Bash category | (track) | hold | |
-| | SWE (light) category | (track 8-task) | hold | added in commit 5e6e123 |
-| **ChatCORE composite** | aggregate Δ vs baseline | +21.09% | −2pp (≥+19%) | 6-task composite |
-| | ARC-Easy | 31.84% | −3pp | |
-| | ARC-Challenge | 26.5% | −3pp | |
-| | MMLU | 27.3% | −3pp | |
-| | GSM8K | 1.6% | floor 0% | already near floor on d6 |
-| | HumanEval | 6.25% | −3pp | |
-| | SpellingBee | ~98% | −3pp | |
-
-When v5 (Code-Feedback + code-act) lands and beats v3 on multi-turn
-while preserving ChatCORE, the floor moves up to v5 — protected
-baseline is *current best*, not historical first.
-
-### Regression-watch eval (per checkpoint)
-
-Every Phase C/D/E checkpoint must run, in this order, before being
-called "promoted":
-
-1. `agent_eval --checkpoint=` (single-turn, fast, ~3 min)
-2. `agent_eval_mt --task-set=default` (38 tasks incl SWE, ~30 min on GPU 0)
-3. `chat_eval` × 6 ChatCORE tasks (~25 min on GPU 0)
-4. `agent_eval_mt --task-set=pcap` (5-50 tasks depending on phase, ~10 min)
-5. **Promotion gate** = (pcap score ≥ phase target) **AND** (every
- protected baseline within tolerance). Failures are recorded in the
- per-stage report; the run is not auto-promoted until reviewed.
-
-### Failure-mode budget (what counts as "drift, not regression")
-
-Before failing a checkpoint outright, allow:
-- ±2pp noise on any single ChatCORE task (eval batch is small).
-- ±3pp on agent_eval_mt category breakdown (long-tail variance).
-- A single category may drop up to its tolerance if at least three
- others move *up* by ≥1pp — net agentic capability preserved.
-
-Hard fail (no promotion regardless): aggregate ChatCORE drops below
-+19%, OR multi-turn Edit-bias re-emerges (≥40% of FAILs are
-"called Edit when wrong tool needed"), OR DPO-style mode collapse
-(single-turn drops to ≤25/12.5 like the v3-DPO regression).
-
-## Six phases (gated)
-
-| Phase | Output | Time | Regression-watch gate |
-|---|---|---|---|
-| A | tshark in sandbox + 30-50 pcap eval w/ 9-field rubric | 1-2 days | n/a (no model change) |
-| B | 5-10k SFT traces (≤64K tok each) + optional 100-150M-tok network corpus | 3-5 days | n/a (data only) |
-| C | d6 baseline (pretrain + SFT + eval), pcap eval ≥30% | 3-4 days | full suite vs v3/v5 floor |
-| D | DAPO loss + agentic_dapo.py, +5pp lift over SFT on d6 | 2-3 days | full suite + DPO-collapse watch |
-| E | Scaling sweep d6/d12/d20/d24 with 64K context, scaling curve | 8-12 days | per-depth, vs that depth's own SFT |
-| F | capbench validation on wuneng | 1 day per ckpt | full suite must still pass |
-
-## Resource shape (tiered GPU schedule)
-
-Host has 8× B300 SXM6. **GPUs 0 and 4** are reliably available during the
-day; GPUs 1-3, 5-7 are committed to qwen3_dualstream / vLLM / sglang
-inference workloads.
-
-- **Day (08:00–20:00):** 2× B300 (GPU 0 + GPU 4) — d6 work, eval,
- ablations, data synthesis, dataset-format adaptation. Pin with
- `CUDA_VISIBLE_DEVICES=0,4` for DDP, or 0 / 4 for single-card jobs
- running concurrently.
-- **Night (20:00–08:00):** 2× B300 + 8× H200 = 10 GPUs — d12/d20/d24
- training.
-- All long jobs must checkpoint ≤90 min and resume cleanly across
- day/night boundaries.
-
-Per-depth wall (with this pool, 64K-context SFT, FSDP, activation ckpt):
-- d6 pretrain 4 h + SFT 6 h + DAPO 4 h (day-only on GPU 0+4)
-- d12 pretrain 14 h + SFT 10 h + DAPO 10 h (3 nights total)
-- d20 pretrain 30 h + SFT 20 h + DAPO 20 h (~7 nights)
-- d24 pretrain 55 h + SFT 35 h + DAPO 35 h (~11 nights)
-
-## Data-needs analysis (what ds79c gives us, and what it doesn't)
-
-ds79c on its own is **not enough** to hit the North Star; it's the
-trajectory backbone, not the full corpus. Tracking ten gaps that need
-explicit data work — each is a planned addition, ordered by impact:
-
-| # | Gap | Priority | Path to fill |
-|---|---|---|---|
-| 1 | **No network-domain pretrain** — d6 base saw web + python only; vocabulary for IP/port/protocol concepts is undertrained | P0 | 100-200M-tok midtraining mix: Wireshark wiki + RFC subset (TCP/IP/DNS/TLS/HTTP/QUIC) + tshark man pages + Stack Overflow networking tag — register as `network-domain` in `pretrain_data.py`, mix at 5-10% during midtraining (after base pretrain, before SFT) |
-| 2 | **Capbench rubric ≠ my 9-field Chinese rubric** — actual demycap reports are English markdown with yaml blocks (`root_causes:`, `evidence:`, `impact:`, `remediation:`) wrapped in `` tags. ds79c rows already use this format; my Phase A rubric is wrong. | P0 | rewrite `eval_pcap_tasks.py:score_report` to match capbench: parse `...`, extract yaml blocks, score against ground-truth lists. Also rewrite the 5 demo pcap ground-truth JSONs to capbench schema (root_causes / evidence / impact / remediation lists). The 9-field 中文 rubric stays as a side-product — useful for non-pcap diagnostic tasks. |
-| 3 | **Tool surface mismatch at runtime** — ds79c uses 11 tools (execute, submit_report, save_mermaid, write_todos, …); `eval_sandbox.py:Sandbox` only implements 4 (Read/Edit/Grep/Bash). Without dispatch, eval-time the model emits tool calls the runtime can't honour. | P0 | Phase A.5: extend `Sandbox` with `execute` (alias to bash), `read_file`/`write_file`/`edit_file`/`ls`/`glob` (aliases or thin wrappers), `submit_report` (terminating tool — capture `args.report_markdown` as the final report, scored by rubric), `save_mermaid` (write `.mmd` file with optional validate stub), `write_todos`/`update_todos`/`task` (no-op stubs that return success). |
-| 4 | **Limited incident type coverage** — only 27 task families in ds79c. Real demycap eval likely covers more (BGP flaps, OSPF re-convergence, multicast issues, IPv6 ND, DHCP starvation, ICMP redirects, asymmetric routing, ECN failures). | P1 | (a) confirm capbench eval task list via `tfcli evals list`; (b) for any gaps: synthesise pcaps with scapy + use a teacher to write trajectories → adds rows to `data/pcap_corpus/` and to SFT mix. |
-| 5 | **Tool-error recovery underrepresented** — exp324 `recovery_weight=5` says these are precious. ds79c has them mixed in but they need to be identified and oversampled. | P1 | Add a filter in `build_demycap_dataset.py`: detect rows where ≥3 consecutive tool turns return non-zero error in same window → flag, then duplicate ×4 (gives ×5 total) in `TaskMixture`. |
-| 6 | **CJK boundary aug is implicit, not curriculum** — 99/981 rows are CJK-augmented but mixed in unweighted. exp324 weighted them ×20. | P1 | Mark CJK-aug rows (no `source_score`) and oversample ×19 (gives ×20) — implement same way as recovery oversample. |
-| 7 | **No held-out RL rollout pool** — Phase D DAPO needs a pcap prompt pool DISTINCT from the SFT set, otherwise DAPO overfits to the same trajectories. | P1 | Pull **ds80 (id=80)**, 1955 rows GLM-5.1, as the RL rollout source — disjoint from ds79c (different teacher, different filter), and contains `fault_locations` schema we'll eventually want anyway. Process via the same adapter; reserve for `agentic_dapo.py` rollouts only. |
-| 8 | **No negative / abort examples** — every ds79c row is `strict_pass=True`. Model never sees "I don't have enough evidence" or wrong-then-corrected behaviour. RL alone won't generate these without a calibrated reward. | P2 | Once DAPO rollouts produce failed trajectories, mine them (`failed but coherent`) as the negative pool for a future DPO/IPO run. Defer until post-Phase-D. |
-| 9 | **Long-context filler** — 16% of rows are 32-64K. Length-curriculum SFT epochs at 32K and 64K need enough material; if ds79c-only is too thin we lose long-context performance during the curriculum. | P2 | Top up with: long RFC sections, long traceforge-uploaded pcap dumps converted to text, concatenated SmolTalk-long. ~100K rows of 32-64K filler at 5% mix. |
-| 10 | **Tokenizer not optimised for IP/port/MAC strings** — `192.168.1.1` and similar BPE-fragment into 8+ tokens. | P3 | Skip retrain (too expensive). Trust the network-domain midtraining (#1) to drive cromulent merges naturally. Re-evaluate after Phase E if d12+ shows IP-string issues in completions. |
-
-**Prioritisation rationale.** P0 items (1, 2, 3) are blockers for Phase
-C eval to even produce meaningful numbers. P1 items (4-7) are the
-difference between "matches ds79c training distribution" and
-"generalises beyond it" — needed before Phase D. P2/P3 (8-10) are
-post-Phase-D refinements.
-
-**Immediate next data action**: P0/2 (rewrite rubric to capbench
-format) — without this, even a perfectly-trained checkpoint scores
-zero on the local eval because we're parsing for the wrong shape.
-P0/3 (Sandbox tool dispatch) is the second blocker; without it, eval
-errors out before scoring.
-
-## Phase B data — primary source = TraceForge ds79c (id=90)
-
-Replaces the "synthesize 5-10k traces from scratch with teacher model"
-plan; the trajectories already exist on the platform.
-
-**Dataset:** `ds79c-with-cjk-boundary-aug`
-([traceforge:90](https://traceforge.staging.netis.com/data/90))
-- 984 records, 76 MB, format=`messages_openai`
-- Source = Qwen3.5-27B `demycap` strict_pass trajectories + 330 CJK↔ASCII
- boundary augmentation rows
-- Used by **exp307a / exp307b / exp307c / exp324** offline SFT —
- exp324 (`exp307a-tfcli-replicate-wukong-gpu6`) is the canonical
- replication: LoRA r=16 / α=32, epochs=2, **max_length=65536** (≡ our
- 64K context plan), batch=1, grad_accum=4, lr=1e-4. Training params
- also weight-oversample certain failure modes:
- `submit_report_header_weight=30`, `recovery_weight=5`,
- `cjk_boundary_weight=20` — we replicate via row-duplication in
- `TaskMixture` (full-FT can't take per-sample weights directly).
-- Each row = full multi-turn dialog (`user → assistant + tool_calls →
- tool → assistant → tool → ...`) ending in the demycap report. First
- row sampled is 47 messages.
-
-**Format adapter** (new file
-`nanoai/agentic/data/demycap_dataset.py`):
-- input rows already use OpenAI-style `messages[]` with `tool_calls[]`
- on assistant turns and `role=tool` for results.
-- target = nanoai `render_agentic_conversation` list-content branch.
-- map `tool_calls[].function.{name,arguments}` →
- `<|tool_call_start|>{name}<|tool_arg|>k<|tool_val|>v...<|tool_call_end|>`
-- map `role=tool` turns → fold into preceding assistant turn as
- `{type:"tool_result", text:content}` then close with
- `<|tool_result_end|>`. Reuse the same fold-pattern that
- `data/codedebug.py` uses for `Execution result:` user-turns.
-
-**Known limitation of ds79c:** 0/984 rows contain `fault_locations` yaml
-schema (training data bug; fixed in ds80 which uses GLM-5.1 source +
-1955 rows). Phase B options:
-1. **ds79c only** — fast start, but the final report won't include
- fault_locations. Acceptable if our 9-field rubric ignores it.
-2. **ds79c + ds80** — mixed, ds80 contributes the floc-correct
- 1955 rows; total ~3k SFT rows for Phase C.
-3. **ds79c → ds80 swap once verified** — start with ds79c (proven by
- exp324) and swap to ds80 if rubric demands floc.
-
-Default to (2): ds79c provides the cjk_boundary aug (specific to demycap
-text-handling); ds80 provides volume + floc schema.
-
-**Optional augmentation** (only if d6 underfits with ~3k rows):
-teacher synthesis on the new pcaps in `data/pcap_corpus/` to top up
-~5k more diverse traces — keeps the original Phase B synthesis path
-alive but only as a top-up, not the primary.
-
-## Context-length plan (key constraint)
-
-Pcap traces target ≤64K tokens; current d6 pretrain seq_len=8192. Strategy:
-1. **RoPE θ scaling 8×** (NTK-aware or YaRN) when loading base for SFT.
-2. **Length-curriculum SFT**: epoch 1 at 16K, epoch 2 at 32K, epoch 3 at 64K.
-3. For d12+ pretrained from scratch, use seq_len=16384 in pretrain itself.
-
-## What lives where
-
-| New file | Purpose |
-|---|---|
-| `eval_pcap_tasks.py` | PcapDiagTask + 9-field rubric + all_pcap_tasks() |
-| `data/pcap_dataset.py` | JSONDataset wrapper for synth traces |
-| `dapo.py` | DAPO loss components |
-| `scripts/synth_pcap_sft.py` | teacher-model trace synthesis |
-| `experiments/agentic_rl/scripts/agentic_dapo.py` | DAPO training script |
-| `data/pcap_corpus/` | curated pcaps + ground_truth.json |
-| `runs/speedrun_pcap_d{6,12,20,24}.sh` | per-depth pipelines |
-
-| Modified file | Why |
-|---|---|
-| `eval_sandbox.py:21` | whitelist tshark/tcpdump/xxd/od/strings |
-| `data/pretrain_data.py` | register network-domain dataset |
-| `code_dataloader.py` | accept 3rd domain ratio |
-| `agentic_pretrain.py` / `agentic_sft.py` / `agent_eval_mt.py` | new flags |
-
-## Success criteria (composite)
-
-Every phase has both a **forward** target (the new pcap capability
-improves) and a **non-regression** target (prior baselines hold). Both
-must pass to call the phase done.
-
-- **Phase A**: pcap eval harness produces non-zero rubric scores
- end-to-end. *(no model change → no non-regression check needed)*
-- **Phase C** (d6 baseline):
- - forward: pcap eval ≥30% on lightweight rubric
- - non-regression: full Regression-watch suite passes (all protected
- baselines within tolerance, no hard-fail trigger)
-- **Phase D** (DAPO):
- - forward: ≥5pp absolute lift over Phase C SFT on pcap eval
- - non-regression: same suite; specifically watch single-turn
- agent_eval (DPO-style mode collapse is the historical risk for any
- on-policy RL on d6 with multi-component reward)
-- **Phase E** (scaling):
- - forward: at least one depth crosses 60% on lightweight pcap eval;
- clear capability-emergence elbow visible in scaling curve
- - non-regression: each depth's checkpoint must non-regress against
- *its own depth's* prior best (d12 vs d12-SFT-only, etc.) — not
- against d6 numbers
-- **Phase F**: local rubric score within ±15pp of capbench score, and
- the same checkpoint also passes the regression-watch suite.
-
-## Out of scope (v1)
-
-- Live capture (`tshark -i`).
-- TLS decryption.
-- Non-Ethernet captures (BT/USB).
-- d3 CPU comparison.
diff --git a/vendor/nanoai/agentic/grpo_multi_rollout.py b/vendor/nanoai/agentic/grpo_multi_rollout.py
deleted file mode 100644
index 741aa212629d642777e0ac991353890a97db0ca8..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/grpo_multi_rollout.py
+++ /dev/null
@@ -1,423 +0,0 @@
-"""Multi-turn rollout for agentic GRPO with per-turn token-span tracking.
-
-Background: MT-GRPO (arXiv 2505.11821) showed that trajectory-only
-advantage estimation caps multi-turn tool-using agents at ~30% accuracy
-while fine-grained turn-level credit assignment unlocks 50%+. To attribute
-a per-turn advantage to the right tokens during the PG update, the
-rollout must record which tokens in the full sequence are the assistant's
-body for each turn.
-
-This module wraps `eval_loop.run_task`'s control flow (`_parse_tool_call`,
-`_execute`, `_build_initial_tokens`, `_append_tool_result`) but records
-`TurnSpan` metadata alongside the generated tokens. Behavior at the
-token level is identical to eval-time rollout; only bookkeeping is added.
-
-The companion `reward_mlr` module consumes (MultiTurnRollout, task) to
-produce per-turn rewards (terminal pass/fail, tool ok, recovery bonus).
-"""
-from __future__ import annotations
-
-import time
-from dataclasses import dataclass
-from typing import Optional
-
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _build_initial_tokens,
- _execute,
- _generate_one_turn,
- _parse_tool_call,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-
-
-@dataclass
-class TurnSpan:
- """One assistant turn's body within MultiTurnRollout.tokens.
-
- [start_idx, end_idx) is half-open; covers the assistant's generated
- tokens for this turn, INCLUDING any tool_call_start/end block and
- the trailing assistant_end (if the model emitted one). These are
- exactly the tokens the PG update supervises.
- """
- turn_idx: int
- start_idx: int
- end_idx: int
- has_tool_call: bool
- tool_name: Optional[str]
- tool_ok: Optional[bool] # None if no tool call; True/False if executed
- is_terminal: bool # True if this turn ended the trajectory
-
-
-@dataclass
-class MultiTurnRollout:
- """One full rollout. tokens spans BOS → final assistant_end (or truncation).
-
- turn_spans[i].(start_idx, end_idx) points into `tokens` to mark the
- i-th assistant turn. passed/reason come from task.check() on the
- post-rollout sandbox state.
- """
- task_id: str
- tokens: list[int]
- prompt_len: int
- turn_spans: list[TurnSpan]
- transcript: Transcript
- passed: bool
- reason: str
- n_turns: int
- n_tool_errors: int
- truncated: bool
- elapsed_s: float
- # Behavior-policy log prob per token, parallel to `tokens` (0.0 on prompt /
- # interstitial / forced positions). None unless rollout_one(capture_logprobs=True).
- # Enables TIS for async/off-policy RL (Polar mechanism, arxiv 2605.24220).
- token_logprobs: Optional[list[float]] = None
-
-
-def rollout_one(
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- *,
- max_tokens_per_turn: int = 256,
- max_turns: Optional[int] = None,
- temperature: float = 0.8,
- seed: int = 42,
- capture_logprobs: bool = False,
-) -> MultiTurnRollout:
- """Run one rollout in the given sandbox, recording per-turn spans.
-
- Sandbox is task.setup()'d at the start. Per-task max_turns is honored
- if higher than the caller's cap (e.g. swe_impl needs 10).
-
- capture_logprobs=True additionally records the behavior-policy log prob of
- each sampled assistant token (aligned to `tokens`; 0.0 elsewhere) into
- MultiTurnRollout.token_logprobs — the behavior logprobs needed for TIS.
- """
- if max_turns is None:
- max_turns = task.max_turns
- else:
- max_turns = max(max_turns, getattr(task, "max_turns", 0) or 0)
-
- transcript = Transcript(user_prompt=task.prompt)
- task.setup(sandbox)
-
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- tokens = _build_initial_tokens(
- tokenizer, task.prompt,
- system_prompt=getattr(task, "system_prompt", None),
- )
- prompt_len = len(tokens)
- # Behavior logprobs parallel to `tokens` (0.0 on prompt / interstitials).
- token_logprobs: Optional[list[float]] = [0.0] * len(tokens) if capture_logprobs else None
-
- turn_spans: list[TurnSpan] = []
- n_turns = 0
- n_tool_errors = 0
- t0 = time.time()
- truncated = False
-
- while n_turns < max_turns:
- n_turns += 1
- # Pad logprobs over any interstitial tokens appended by the prev turn,
- # so token_logprobs stays aligned to `tokens` before this body.
- if capture_logprobs and len(token_logprobs) < len(tokens):
- token_logprobs.extend([0.0] * (len(tokens) - len(token_logprobs)))
- body_start = len(tokens)
- if capture_logprobs:
- out_ids, out_lp = _generate_one_turn(
- engine, tokenizer, tokens,
- max_tokens_per_turn, temperature, seed + n_turns,
- capture_logprobs=True,
- )
- token_logprobs.extend(0.0 if lp is None else float(lp) for lp in out_lp)
- else:
- out_ids = _generate_one_turn(
- engine, tokenizer, tokens,
- max_tokens_per_turn, temperature, seed + n_turns,
- )
- tokens.extend(out_ids)
- body_end = len(tokens)
-
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
- if not ended_clean and len(out_ids) >= max_tokens_per_turn:
- truncated = True
-
- # Inspect for tool call.
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- j = -1
- if j > i:
- pre = tokenizer.decode(body_ids[:i]).strip()
- tool_body = tokenizer.decode(body_ids[i + 1 : j])
- try:
- name, args = _parse_tool_call_dispatch(tokenizer, tool_body)
- except ValueError as e:
- name, args = "", {}
- result_text, ok = f"tool parse error: {e}", False
- else:
- result_text, ok = _execute(sandbox, name, args)
- transcript.turns.append({
- "role": "assistant",
- "content": pre,
- "tool_call": {"name": name, "args": args},
- })
- transcript.turns.append({
- "role": "tool_result",
- "tool_result": result_text,
- "error": not ok,
- })
- if not ok:
- n_tool_errors += 1
-
- is_terminal = name in TERMINAL_TOOLS
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=True,
- tool_name=name,
- tool_ok=ok,
- is_terminal=is_terminal,
- ))
- if is_terminal:
- break
- if not ended_clean:
- tokens.append(asst_end)
- _append_tool_result(tokenizer, tokens, result_text)
- continue
-
- # No tool call this turn → final answer.
- text = tokenizer.decode(body_ids).strip()
- transcript.turns.append({"role": "assistant", "content": text})
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=False,
- tool_name=None,
- tool_ok=None,
- is_terminal=True,
- ))
- break
-
- elapsed = time.time() - t0
- passed, reason = task.check(sandbox, transcript)
-
- # Final pad so token_logprobs spans the full `tokens` (trailing interstitials).
- if capture_logprobs and len(token_logprobs) < len(tokens):
- token_logprobs.extend([0.0] * (len(tokens) - len(token_logprobs)))
-
- return MultiTurnRollout(
- task_id=task.task_id,
- tokens=tokens,
- prompt_len=prompt_len,
- turn_spans=turn_spans,
- transcript=transcript,
- passed=passed,
- reason=reason,
- n_turns=n_turns,
- n_tool_errors=n_tool_errors,
- truncated=truncated,
- elapsed_s=round(elapsed, 2),
- token_logprobs=token_logprobs,
- )
-
-
-def rollout_group(
- engine,
- tokenizer,
- task: Task,
- *,
- group_size: int,
- max_tokens_per_turn: int = 256,
- max_turns: Optional[int] = None,
- temperature: float = 0.8,
- base_seed: int = 42,
- bash_timeout: float = 5.0,
- capture_logprobs: bool = False,
-) -> list[MultiTurnRollout]:
- """Run group_size rollouts sequentially with distinct seeds; isolated sandboxes.
-
- capture_logprobs forwards to rollout_one to record behavior logprobs for TIS."""
- rollouts: list[MultiTurnRollout] = []
- for g in range(group_size):
- with Sandbox(bash_timeout=bash_timeout) as sb:
- rollouts.append(rollout_one(
- engine, tokenizer, sb, task,
- max_tokens_per_turn=max_tokens_per_turn,
- max_turns=max_turns,
- temperature=temperature,
- capture_logprobs=capture_logprobs,
- seed=base_seed + g * 1000,
- ))
- return rollouts
-
-
-# ---- Serialization -------------------------------------------------------
-
-def rollout_to_dict(r: MultiTurnRollout, *, include_tokens: bool = False) -> dict:
- """Serialize a rollout to a JSON-friendly dict.
-
- include_tokens=True saves the full token list (needed for hidden-state
- extraction in Exp 2 / Turn-Critic Probe). The lite version (default)
- omits tokens to keep JSONL small (~1 KB/rollout vs ~10 KB).
- """
- d: dict = {
- "task_id": r.task_id,
- "passed": r.passed,
- "reason": r.reason,
- "n_turns": r.n_turns,
- "n_tool_errors": r.n_tool_errors,
- "truncated": r.truncated,
- "prompt_len": r.prompt_len,
- "elapsed_s": r.elapsed_s,
- "turn_spans": [
- {
- "turn_idx": ts.turn_idx,
- "start_idx": ts.start_idx,
- "end_idx": ts.end_idx,
- "has_tool_call": ts.has_tool_call,
- "tool_name": ts.tool_name,
- "tool_ok": ts.tool_ok,
- "is_terminal": ts.is_terminal,
- }
- for ts in r.turn_spans
- ],
- }
- if include_tokens:
- d["tokens"] = r.tokens
- return d
-
-
-# ---- Self-test ----------------------------------------------------------
-#
-# Drives rollout_one with a FakeEngine that emits a scripted token stream
-# (no real model needed). Confirms span bookkeeping, tool execution, and
-# task.check() pass on a hand-crafted Task.
-
-if __name__ == "__main__":
- _SPECIALS = {
- "<|bos|>": 1000,
- "<|user_start|>": 1001,
- "<|user_end|>": 1002,
- "<|assistant_start|>": 1003,
- "<|assistant_end|>": 1004,
- "<|tool_call_start|>": 1005,
- "<|tool_call_end|>": 1006,
- "<|tool_arg|>": 1007,
- "<|tool_val|>": 1008,
- "<|tool_result_start|>": 1009,
- "<|tool_result_end|>": 1010,
- }
- _INV = {v: k for k, v in _SPECIALS.items()}
- _BODY = {
- 9001: "Read", 9002: "file_path", 9003: "a.txt",
- 12345: "Reading", 9999: "done",
- }
-
- class _FakeTok:
- def encode_special(self, s):
- return _SPECIALS[s]
- def get_bos_token_id(self):
- return _SPECIALS["<|bos|>"]
- def encode(self, s):
- return [hash(w) % 5000 + 2000 for w in s.split()]
- def decode(self, ids):
- out = []
- for i in ids:
- if i in _INV:
- out.append(_INV[i])
- elif i in _BODY:
- out.append(_BODY[i])
- else:
- out.append(f"w{i}")
- return " ".join(out)
-
- class _FakeEngine:
- def __init__(self, streams):
- self._streams = streams
- self._call = 0
- def generate(self, tokens, num_samples, max_tokens, temperature, seed):
- stream = self._streams[self._call]
- self._call += 1
- for tok in stream:
- yield [tok], [0]
-
- # Turn 1 — emit a Read tool call for a.txt, terminate cleanly.
- stream1 = [
- 12345, # "Reading"
- _SPECIALS["<|tool_call_start|>"],
- 9001, # "Read"
- _SPECIALS["<|tool_arg|>"],
- 9002, # "file_path"
- _SPECIALS["<|tool_val|>"],
- 9003, # "a.txt"
- _SPECIALS["<|tool_call_end|>"],
- _SPECIALS["<|assistant_end|>"],
- ]
- # Turn 2 — final answer (no tool call).
- stream2 = [9999, _SPECIALS["<|assistant_end|>"]]
-
- tok = _FakeTok()
- eng = _FakeEngine([stream1, stream2])
-
- def _setup(sb):
- sb.write_file("a.txt", "hello world")
-
- def _check(sb, t):
- for tc in t.tool_calls:
- if (tc.get("name") == "Read"
- and tc.get("args", {}).get("file_path") == "a.txt"):
- return True, "read a.txt"
- return False, "no read"
-
- task = Task(
- task_id="self_test_read",
- prompt="read a.txt",
- setup=_setup, check=_check,
- tools_expected=("Read",), max_turns=2,
- )
-
- with Sandbox() as sb:
- r = rollout_one(eng, tok, sb, task,
- max_tokens_per_turn=64, temperature=1.0, seed=0)
-
- print(f"passed={r.passed} reason={r.reason} n_turns={r.n_turns}")
- print(f"prompt_len={r.prompt_len} total_tokens={len(r.tokens)}")
- for ts in r.turn_spans:
- print(f" turn {ts.turn_idx}: [{ts.start_idx}, {ts.end_idx}) "
- f"has_tc={ts.has_tool_call} tool={ts.tool_name} "
- f"ok={ts.tool_ok} term={ts.is_terminal}")
-
- assert r.n_turns == 2, f"expected 2 turns, got {r.n_turns}"
- assert len(r.turn_spans) == 2
- assert r.turn_spans[0].has_tool_call is True
- assert r.turn_spans[0].tool_name == "Read"
- assert r.turn_spans[0].tool_ok is True
- assert r.turn_spans[1].has_tool_call is False
- assert r.turn_spans[1].is_terminal is True
- assert r.passed is True, f"task.check should pass: {r.reason}"
-
- # Span body of turn 0 must contain the tool_call_start/end tokens.
- body0 = r.tokens[r.turn_spans[0].start_idx : r.turn_spans[0].end_idx]
- assert _SPECIALS["<|tool_call_start|>"] in body0
- assert _SPECIALS["<|tool_call_end|>"] in body0
- # Spans must not overlap.
- assert r.turn_spans[1].start_idx >= r.turn_spans[0].end_idx
- # First span must start AFTER the prompt.
- assert r.turn_spans[0].start_idx == r.prompt_len
-
- print("self-test OK")
diff --git a/vendor/nanoai/agentic/grpo_multi_rollout_batched.py b/vendor/nanoai/agentic/grpo_multi_rollout_batched.py
deleted file mode 100644
index f5d0c0bcb432c0c23d1fc7a787216f20921bd1e8..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/grpo_multi_rollout_batched.py
+++ /dev/null
@@ -1,674 +0,0 @@
-"""Batched per-group rollout for agentic GRPO.
-
-This module is a TDD-driven sibling of `grpo_multi_rollout.py`. The goal is
-to speed up Phase 1 (rollout) of agentic RL training by running all
-`group_size` rollouts of a group together — batched LM forward passes plus
-parallel tool execution — while preserving byte-equal trajectories vs the
-sequential reference (`rollout_one` / `rollout_group`).
-
-Stages (each gated by a byte-equal test against the reference):
- 1. Extract `_do_one_turn` primitive — pure turn-step state transition.
- 2. `rollout_group_batched` sequential — list-of-state loop.
- 3. Engine.generate_batched_independent — per-sample seed + stop_tokens.
- 4. Parallel tool execution via ThreadPoolExecutor.
- 5. Stateful per-sample KV cache reuse (BatchedRolloutSession).
-
-Stage 1 deliverables (this file):
- - `_RolloutState` dataclass holding one rollout's mutable state.
- - `_do_one_turn` primitive that processes one already-generated body of
- tokens (parse tool, execute, append tool result, update state). DOES
- NOT call engine.generate.
- - `rollout_one_v2` — drop-in for `rollout_one` that uses the primitive.
- Used as the golden test target: byte-equal trajectories vs `rollout_one`.
-
-This module is import-safe (no engine / GPU work) — only consumed by RL
-trainers that explicitly opt in via `--batched-rollout` flag.
-"""
-from __future__ import annotations
-
-import time
-from concurrent.futures import ThreadPoolExecutor
-from dataclasses import dataclass, field
-from typing import Optional
-
-import torch
-
-from ..engine import KVCache, sample_next_token
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _execute,
- _generate_one_turn,
- _parse_tool_call,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_loop import _build_initial_tokens # noqa: re-export use
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-from .grpo_multi_rollout import MultiTurnRollout, TurnSpan
-
-
-@dataclass
-class _RolloutState:
- """Mutable state of one rollout. Updated in place by `_do_one_turn`."""
- # Identifies sample within a group; used for per-sample seed derivation.
- sample_idx: int
- # Token sequence — grows each turn (generated body + optional tool result tokens).
- tokens: list[int]
- # Length of initial prompt up to first assistant turn start.
- prompt_len: int
- # Bookkeeping for token spans of each assistant turn.
- turn_spans: list[TurnSpan] = field(default_factory=list)
- # Live transcript (assistant + tool_result entries).
- transcript: Optional[Transcript] = None
- # Sandbox object — mutated by `_execute` during tool calls.
- sandbox: Optional[Sandbox] = None
- # Turn counter (incremented BEFORE each turn body).
- n_turns: int = 0
- # Tool-error counter (incremented when _execute returns ok=False).
- n_tool_errors: int = 0
- # True if we hit max_tokens_per_turn without emitting <|assistant_end|>.
- truncated: bool = False
- # True once a terminal tool fired OR a no-tool-call final answer was given.
- is_terminal: bool = False
- # Per-sample base seed for deterministic per-turn seeding.
- base_seed: int = 0
-
-
-def _do_one_turn(
- state: _RolloutState,
- generated_ids: list[int],
- *,
- tokenizer,
- task: Task,
- asst_end: int,
- tc_start: int,
- tc_end: int,
- max_tokens_per_turn: int,
-) -> None:
- """Process one already-generated turn body. Mutates `state` in place.
-
- Faithful extraction of `rollout_one` body (`grpo_multi_rollout.py:115-184`)
- with engine.generate factored out. The caller must:
- - increment state.n_turns BEFORE invoking _do_one_turn
- - pre-record body_start = len(state.tokens) BEFORE invoking
- - then call this function with the new tokens just generated
-
- After this function:
- - state.tokens grows to include generated body + (if non-terminal tool
- call) the appended tool result tokens
- - state.turn_spans gets one new TurnSpan
- - state.transcript.turns gets 1 or 2 new entries
- - state.is_terminal is True if the rollout should stop here
- - state.n_tool_errors / state.truncated updated as needed
- """
- assert state.transcript is not None and state.sandbox is not None
-
- body_start = len(state.tokens) - len(generated_ids)
- state.tokens.extend([]) # no-op; for symmetry with caller appending generated_ids itself
- # Note: caller is expected to have ALREADY appended generated_ids to state.tokens.
- # We accept generated_ids as redundant arg for clarity, but operate on state.tokens.
- body_end = len(state.tokens)
-
- # The body is what was just appended. We slice it out by position.
- out_ids = state.tokens[body_start:body_end]
- assert out_ids == generated_ids, (
- f"Caller must append generated_ids to state.tokens before calling _do_one_turn "
- f"(got mismatch len={len(out_ids)} vs {len(generated_ids)})"
- )
-
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
- if not ended_clean and len(out_ids) >= max_tokens_per_turn:
- state.truncated = True
-
- # Inspect for tool call.
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- j = -1
- if j > i:
- pre = tokenizer.decode(body_ids[:i]).strip()
- tool_body = tokenizer.decode(body_ids[i + 1: j])
- try:
- name, args = _parse_tool_call_dispatch(tokenizer, tool_body)
- except ValueError as e:
- name, args = "", {}
- result_text, ok = f"tool parse error: {e}", False
- else:
- result_text, ok = _execute(state.sandbox, name, args)
- state.transcript.turns.append({
- "role": "assistant",
- "content": pre,
- "tool_call": {"name": name, "args": args},
- })
- state.transcript.turns.append({
- "role": "tool_result",
- "tool_result": result_text,
- "error": not ok,
- })
- if not ok:
- state.n_tool_errors += 1
-
- is_terminal_tool = name in TERMINAL_TOOLS
- state.turn_spans.append(TurnSpan(
- turn_idx=state.n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=True,
- tool_name=name,
- tool_ok=ok,
- is_terminal=is_terminal_tool,
- ))
- if is_terminal_tool:
- state.is_terminal = True
- return
- if not ended_clean:
- state.tokens.append(asst_end)
- _append_tool_result(tokenizer, state.tokens, result_text)
- # Continue to next turn (not terminal)
- return
-
- # No tool call this turn → final answer.
- text = tokenizer.decode(body_ids).strip()
- state.transcript.turns.append({"role": "assistant", "content": text})
- state.turn_spans.append(TurnSpan(
- turn_idx=state.n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=False,
- tool_name=None,
- tool_ok=None,
- is_terminal=True,
- ))
- state.is_terminal = True
-
-
-def _init_state(
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- *,
- sample_idx: int = 0,
- base_seed: int = 42,
-) -> _RolloutState:
- """Initialize per-rollout state: build initial tokens, set up sandbox via
- task.setup(). Mirrors the start of `rollout_one`."""
- task.setup(sandbox)
- tokens = _build_initial_tokens(
- tokenizer, task.prompt,
- system_prompt=getattr(task, "system_prompt", None),
- )
- transcript = Transcript(user_prompt=task.prompt)
- return _RolloutState(
- sample_idx=sample_idx,
- tokens=tokens,
- prompt_len=len(tokens),
- transcript=transcript,
- sandbox=sandbox,
- base_seed=base_seed,
- )
-
-
-def _state_to_rollout(state: _RolloutState, task: Task, elapsed: float) -> MultiTurnRollout:
- """Finalize: run task.check on the sandbox, package into MultiTurnRollout."""
- passed, reason = task.check(state.sandbox, state.transcript)
- return MultiTurnRollout(
- task_id=task.task_id,
- tokens=state.tokens,
- prompt_len=state.prompt_len,
- turn_spans=state.turn_spans,
- transcript=state.transcript,
- passed=passed,
- reason=reason,
- n_turns=state.n_turns,
- n_tool_errors=state.n_tool_errors,
- truncated=state.truncated,
- elapsed_s=round(elapsed, 2),
- )
-
-
-def rollout_one_v2(
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- *,
- max_tokens_per_turn: int = 256,
- max_turns: Optional[int] = None,
- temperature: float = 0.8,
- seed: int = 42,
-) -> MultiTurnRollout:
- """Drop-in for `rollout_one` using the `_do_one_turn` primitive.
-
- Behavior must be byte-equal to `rollout_one` given the same args. This is
- the Stage 1 deliverable + the basis for Stage 2 batched orchestration.
- """
- if max_turns is None:
- max_turns = task.max_turns
- else:
- max_turns = max(max_turns, getattr(task, "max_turns", 0) or 0)
-
- state = _init_state(tokenizer, sandbox, task, base_seed=seed)
-
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- t0 = time.time()
- while state.n_turns < max_turns:
- state.n_turns += 1
- # Generate one turn (caller responsibility — Stage 1 keeps sequential generate)
- out_ids = _generate_one_turn(
- engine, tokenizer, state.tokens,
- max_tokens_per_turn, temperature, seed + state.n_turns,
- )
- state.tokens.extend(out_ids)
- # Process the just-generated body
- _do_one_turn(
- state, out_ids,
- tokenizer=tokenizer, task=task,
- asst_end=asst_end, tc_start=tc_start, tc_end=tc_end,
- max_tokens_per_turn=max_tokens_per_turn,
- )
- if state.is_terminal:
- break
-
- elapsed = time.time() - t0
- return _state_to_rollout(state, task, elapsed)
-
-
-def rollout_group_batched(
- engine,
- tokenizer,
- task: Task,
- *,
- group_size: int,
- max_tokens_per_turn: int = 256,
- max_turns: Optional[int] = None,
- temperature: float = 0.8,
- base_seed: int = 42,
- bash_timeout: float = 5.0,
- parallel_tool_workers: Optional[int] = None,
-) -> list[MultiTurnRollout]:
- """Run `group_size` rollouts as a list-of-state loop. Drop-in for
- `rollout_group` with byte-equal output (same seeds, same task).
-
- Stage 2 (sequential): generates + tool-execs per state in sequence.
- Stage 4 (current): generates per state sequentially, then runs tool
- execution + state update for the entire turn in parallel across active
- states via ThreadPoolExecutor. Each state has its own Sandbox (different
- tmpdirs) so tool exec is thread-isolated. tokenizer / engine / task are
- shared (re-entrant / immutable / immutable respectively).
-
- Stage 3 wrote `engine.generate_batched_independent` to support per-sample
- seed + stop, but its internal implementation is sequential B=1 (bf16
- batched matmul is not bit-equal to repeated B=1). Stage 5 will add per-
- sample KV cache reuse across turns (BatchedRolloutSession) — that's
- where the real GPU-side speedup comes from.
-
- Args:
- parallel_tool_workers: max ThreadPool workers for the tool-exec
- phase. None = min(group_size, 8). Set 1 to disable parallelism
- (equivalent to Stage 2 sequential).
- """
- if max_turns is None:
- max_turns = task.max_turns
- else:
- max_turns = max(max_turns, getattr(task, "max_turns", 0) or 0)
-
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- if parallel_tool_workers is None:
- parallel_tool_workers = min(group_size, 8)
- parallel_tool_workers = max(1, parallel_tool_workers)
-
- sandboxes = [Sandbox(bash_timeout=bash_timeout) for _ in range(group_size)]
- for sb in sandboxes:
- sb.__enter__()
- try:
- # Init G states; each state owns one sandbox + per-sample base_seed
- # (mirrors rollout_group's `base_seed + g * 1000`).
- states = [
- _init_state(tokenizer, sandboxes[g], task,
- sample_idx=g, base_seed=base_seed + g * 1000)
- for g in range(group_size)
- ]
-
- t0 = time.time()
- # Outer loop: per turn,
- # Phase 1: sequential generate per active state (GPU is the bottleneck
- # here; batching at GPU level breaks byte-equal in bf16 so we
- # keep it sequential — Stage 5 KV reuse helps via prefix-cache).
- # Phase 2: parallel tool execution + state update via ThreadPool. Each
- # state has its own Sandbox so tool calls are isolated. The
- # firejail subprocess for Bash releases GIL; Read/Edit/Grep
- # in-process are also GIL-safe since each touches its own
- # sandbox tmpdir.
- while True:
- active = [s for s in states if not s.is_terminal and s.n_turns < max_turns]
- if not active:
- break
-
- # Phase 1: generate all active states (sequential).
- generated_per_state: list[tuple[_RolloutState, list[int]]] = []
- for state in active:
- state.n_turns += 1
- out_ids = _generate_one_turn(
- engine, tokenizer, state.tokens,
- max_tokens_per_turn, temperature,
- state.base_seed + state.n_turns,
- )
- state.tokens.extend(out_ids)
- generated_per_state.append((state, out_ids))
-
- # Phase 2: parallel tool execution + state update.
- def _process_one(item):
- st, out_ids = item
- _do_one_turn(
- st, out_ids,
- tokenizer=tokenizer, task=task,
- asst_end=asst_end, tc_start=tc_start, tc_end=tc_end,
- max_tokens_per_turn=max_tokens_per_turn,
- )
-
- if parallel_tool_workers == 1 or len(generated_per_state) == 1:
- # Sequential fallback (cheaper for B=1 or when parallelism disabled).
- for item in generated_per_state:
- _process_one(item)
- else:
- # Use ThreadPool — bound on min(active, parallel_tool_workers).
- workers = min(parallel_tool_workers, len(generated_per_state))
- with ThreadPoolExecutor(max_workers=workers) as executor:
- # list(map) forces evaluation + propagates any exceptions.
- list(executor.map(_process_one, generated_per_state))
-
- elapsed = time.time() - t0
- # Final: task.check + package per-sample MultiTurnRollout.
- return [_state_to_rollout(s, task, elapsed=elapsed) for s in states]
- finally:
- for sb in sandboxes:
- try:
- sb.__exit__(None, None, None)
- except Exception:
- pass
-
-
-# ============================================================================
-# Stage 5: per-sample persistent KV cache reuse across turns
-# ============================================================================
-
-class BatchedRolloutSession:
- """Holds one KVCache(B=1) per rollout sample, persistent across turns.
-
- Stage 5 motivation: by default each turn's generation re-prefills the FULL
- accumulated prompt from scratch (Stage 1-4). For a 6-turn rollout with 200
- tokens per turn, that's 200+400+...+1200 = 4200 forward tokens just on
- prefill. With this session, turn N only forwards the NEW tokens since
- turn N-1 (the previously-generated body + tool result + structural marker
- tokens) — ~200 per turn × 6 = 1200, a 3.5x reduction on prefill compute.
-
- Each sample's KV cache is B=1 (one row), so we use the standard
- flash_attn_with_kvcache uniform-pos path; no variable_seqlens needed for
- Stage 5. We exchange GPU-side batch parallelism for byte-equal correctness
- + simpler per-sample state. The Stage 0.5 variable_seqlens kernel remains
- available for a future variant that genuinely batches the GPU forward at
- the cost of bf16 reduction-order drift.
-
- API:
- session = BatchedRolloutSession(engine, max_samples, max_seq_len)
- for b in range(B):
- session.prefill(b, prompt_tokens[b]) # initial prompt → KV
- new_tokens = session.decode_one_turn(b, seed=..., max_new_tokens=...,
- stop_tokens={asst_end, bos})
- session.append(b, tool_result_tokens) # extend KV w/o sampling
- # repeat decode + append per turn
- session.release() # free KV cache memory
- """
-
- def __init__(self, engine, max_samples: int, max_seq_len: int):
- self.engine = engine
- self.max_samples = max_samples
- self.max_seq_len = max_seq_len
- self.device = engine.model.get_device()
- cfg = engine.model.config
- dtype = torch.bfloat16 if self.device.type == "cuda" else torch.float32
- kv_kwargs = {
- "num_heads": cfg.n_kv_head,
- "head_dim": cfg.n_embd // cfg.n_head,
- "num_layers": cfg.n_layer,
- }
- # Per-sample B=1 KV caches — fully independent, no cross-sample contamination.
- self.caches = [
- KVCache(batch_size=1, seq_len=max_seq_len, device=self.device,
- dtype=dtype, **kv_kwargs)
- for _ in range(max_samples)
- ]
- # Per-sample last-position logits (for next decode's first sample).
- self.last_logits: list[Optional[torch.Tensor]] = [None] * max_samples
- # Per-sample count of tokens currently in KV (for invariant checks).
- self.n_tokens_in_kv: list[int] = [0] * max_samples
-
- @torch.inference_mode()
- def prefill(self, sample_idx: int, tokens: list[int]) -> None:
- """Initial prefill: forward the prompt into sample's empty KV. Sets
- last_logits. Must be called exactly once per sample before any decode."""
- assert self.n_tokens_in_kv[sample_idx] == 0, (
- f"sample {sample_idx} already prefilled (n_tokens_in_kv="
- f"{self.n_tokens_in_kv[sample_idx]}); use append() instead")
- if not tokens:
- return
- ids = torch.tensor([tokens], dtype=torch.long, device=self.device)
- with self.engine.generation_forward_context():
- logits = self.engine.model.forward(ids, kv_cache=self.caches[sample_idx])
- # logits shape: (1, L, V) — take the last position's logit as the next-token predictor.
- self.last_logits[sample_idx] = logits[0, -1, :].clone()
- self.n_tokens_in_kv[sample_idx] = len(tokens)
-
- @torch.inference_mode()
- def append(self, sample_idx: int, tokens: list[int]) -> None:
- """Extend sample's KV with `tokens` (no sampling). Used between turns
- to absorb e.g. tool result tokens + structural markers (<|asst_end|>,
- <|tool_result_start|>, ..., <|asst_start|>). Updates last_logits."""
- if not tokens:
- return
- ids = torch.tensor([tokens], dtype=torch.long, device=self.device)
- with self.engine.generation_forward_context():
- logits = self.engine.model.forward(ids, kv_cache=self.caches[sample_idx])
- self.last_logits[sample_idx] = logits[0, -1, :].clone()
- self.n_tokens_in_kv[sample_idx] += len(tokens)
-
- @torch.inference_mode()
- def decode_one_turn(
- self,
- sample_idx: int,
- seed: int,
- max_new_tokens: int,
- stop_tokens: set[int],
- temperature: float = 1.0,
- top_k: Optional[int] = None,
- ) -> list[int]:
- """Generate up to max_new_tokens, stopping at first stop_token.
-
- Each generated token is forwarded into KV before the next sample, so
- on exit KV has prompt + all_generated (including stop). Returns the
- list of generated tokens including the stop token if hit (matches
- rollout_one's `out_ids` contract).
-
- Byte-equal to engine.generate(num_samples=1, seed=seed, max_tokens=
- max_new_tokens) given identical initial KV state.
- """
- assert self.last_logits[sample_idx] is not None, (
- f"sample {sample_idx} not prefilled — call prefill() first")
- cache = self.caches[sample_idx]
- rng = torch.Generator(device=self.device).manual_seed(int(seed))
- new_tokens: list[int] = []
- for _ in range(max_new_tokens):
- # Sample next from current last_logits
- tok = sample_next_token(
- self.last_logits[sample_idx].unsqueeze(0), # (1, V)
- rng, temperature=temperature, top_k=top_k,
- ).item()
- new_tokens.append(tok)
- # Forward this token into KV (advances cache_seqlens by 1)
- ids = torch.tensor([[tok]], dtype=torch.long, device=self.device)
- with self.engine.generation_forward_context():
- logits = self.engine.model.forward(ids, kv_cache=cache)
- self.n_tokens_in_kv[sample_idx] += 1
- self.last_logits[sample_idx] = logits[0, -1, :].clone()
- # Stop check AFTER forward (matches engine.generate which advances
- # KV for the stop token too — see engine.py:283-294).
- if tok in stop_tokens:
- break
- return new_tokens
-
- def release(self) -> None:
- """Free KV cache tensors. Call after rollout group done to reclaim GPU memory."""
- self.caches.clear()
- self.last_logits.clear()
- self.n_tokens_in_kv.clear()
-
-
-def rollout_group_batched_kv(
- engine,
- tokenizer,
- task: Task,
- *,
- group_size: int,
- max_tokens_per_turn: int = 256,
- max_turns: Optional[int] = None,
- temperature: float = 0.8,
- base_seed: int = 42,
- bash_timeout: float = 5.0,
- parallel_tool_workers: Optional[int] = None,
-) -> list[MultiTurnRollout]:
- """Stage 5 (EXPERIMENTAL — NOT byte-equal): rollout_group_batched + per-
- sample KV cache reuse across turns.
-
- Each sample has a persistent KVCache that survives across turns. Per-turn:
- - Phase 1: per-sample sequential decode_one_turn (each row B=1, fresh
- RNG seeded by base_seed + g*1000 + n_turns).
- - Phase 2: parallel tool execution via ThreadPool (same as Stage 4).
- - Phase 3: per-sample session.append for tool result + structural tokens
- added by _do_one_turn (extend KV for next turn's decode).
-
- **bf16 limitation**: per `BatchedRolloutSession` docstring, the cross-turn
- KV reuse path produces ε-different (NOT byte-equal) trajectories vs the
- reference `rollout_group` which re-prefills the full state.tokens each
- turn. The bf16 GEMM kernel selects different tile/reduction orders for
- incremental T=1 forwards vs T=L prefill forwards, accumulating logit
- differences that occasionally flip argmax (and therefore subsequent
- sampled tokens). With deterministic algorithms enabled the divergence is
- smaller but still nonzero (~1.0 max logit diff observed at d6).
-
- For strict byte-equal use `rollout_group_batched` (Stages 1-4). For
- speedup with ε-different but valid trajectories, use this function.
- Speedup expected: 1.5-2.5x on multi-turn rollouts via prefix prefill
- reuse. Trajectories may differ from sequential reference but are
- distributionally similar (same model, same prompts, similar pass rates).
- """
- if max_turns is None:
- max_turns = task.max_turns
- else:
- max_turns = max(max_turns, getattr(task, "max_turns", 0) or 0)
-
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
- bos = tokenizer.get_bos_token_id()
- stop_tokens = {asst_end, bos}
-
- if parallel_tool_workers is None:
- parallel_tool_workers = min(group_size, 8)
- parallel_tool_workers = max(1, parallel_tool_workers)
-
- # Reserve KV cache big enough for prompt + all turns' generated + tool result tokens.
- # Conservative bound: prompt_len + max_turns * (max_tokens_per_turn + tool_result_budget).
- # tool_result_budget rough: 4000 chars / 4 chars-per-token ≈ 1000 tokens per turn worst-case.
- seq_cap = engine.model.config.sequence_len
- max_seq_len = seq_cap
-
- sandboxes = [Sandbox(bash_timeout=bash_timeout) for _ in range(group_size)]
- for sb in sandboxes:
- sb.__enter__()
- session = BatchedRolloutSession(engine, group_size, max_seq_len)
- try:
- states = [
- _init_state(tokenizer, sandboxes[g], task,
- sample_idx=g, base_seed=base_seed + g * 1000)
- for g in range(group_size)
- ]
- # Initial prefill: each sample's prompt into its own KV cache.
- for state in states:
- session.prefill(state.sample_idx, state.tokens)
-
- t0 = time.time()
- while True:
- active = [s for s in states if not s.is_terminal and s.n_turns < max_turns]
- if not active:
- break
-
- # Phase 1: per-sample decode using persistent KV.
- decoded_per_state: list[tuple[_RolloutState, list[int]]] = []
- for state in active:
- state.n_turns += 1
- # Invariant: KV is in sync with state.tokens before decode.
- assert session.n_tokens_in_kv[state.sample_idx] == len(state.tokens), (
- f"sample {state.sample_idx} kv desync: "
- f"kv={session.n_tokens_in_kv[state.sample_idx]} "
- f"state.tokens={len(state.tokens)}")
- new_tokens = session.decode_one_turn(
- sample_idx=state.sample_idx,
- seed=state.base_seed + state.n_turns,
- max_new_tokens=max_tokens_per_turn,
- stop_tokens=stop_tokens,
- temperature=temperature,
- )
- state.tokens.extend(new_tokens)
- # KV advanced by len(new_tokens) inside decode_one_turn.
- assert session.n_tokens_in_kv[state.sample_idx] == len(state.tokens)
- decoded_per_state.append((state, new_tokens))
-
- # Phase 2: parallel tool exec + state update (mutates state.tokens
- # by appending tool result tokens + asst_end if not ended_clean).
- def _process_one(item):
- st, out_ids = item
- _do_one_turn(
- st, out_ids,
- tokenizer=tokenizer, task=task,
- asst_end=asst_end, tc_start=tc_start, tc_end=tc_end,
- max_tokens_per_turn=max_tokens_per_turn,
- )
-
- if parallel_tool_workers == 1 or len(decoded_per_state) == 1:
- for item in decoded_per_state:
- _process_one(item)
- else:
- workers = min(parallel_tool_workers, len(decoded_per_state))
- with ThreadPoolExecutor(max_workers=workers) as executor:
- list(executor.map(_process_one, decoded_per_state))
-
- # Phase 3: sync KV with tool-result tokens appended by _do_one_turn.
- # For terminal samples no append needed (cache won't be queried again).
- for state, _ in decoded_per_state:
- if state.is_terminal:
- continue
- kv_at = session.n_tokens_in_kv[state.sample_idx]
- delta_tokens = state.tokens[kv_at:]
- if delta_tokens:
- session.append(state.sample_idx, delta_tokens)
- assert session.n_tokens_in_kv[state.sample_idx] == len(state.tokens)
-
- elapsed = time.time() - t0
- rollouts = [_state_to_rollout(s, task, elapsed=elapsed) for s in states]
- finally:
- session.release()
- for sb in sandboxes:
- try:
- sb.__exit__(None, None, None)
- except Exception:
- pass
- return rollouts
diff --git a/vendor/nanoai/agentic/hard_tasks/__init__.py b/vendor/nanoai/agentic/hard_tasks/__init__.py
deleted file mode 100644
index 46c50c73f08d6e1af738157a91086893dcdfb08c..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/hard_tasks/__init__.py
+++ /dev/null
@@ -1,202 +0,0 @@
-"""Hard task generator for agentic RL training.
-
-Generates 10K-100K parameterized tasks with 40-60% pass rate for d6,
-using programmatic skeletons + 27B LLM fill for realistic code content.
-"""
-from __future__ import annotations
-
-import json
-import subprocess
-import tempfile
-from dataclasses import dataclass, field, asdict
-from pathlib import Path
-from typing import Optional
-
-from ..eval_sandbox import Sandbox
-from ..eval_tasks import Task, Transcript
-
-
-@dataclass
-class HardTask:
- task_id: str
- family: str
- prompt: str
- files: dict[str, str]
- check_spec: dict
- golden_files: dict[str, str]
- params: dict = field(default_factory=dict)
- difficulty: int = 2
- system_prompt: Optional[str] = None
- max_turns: int = 8
-
- def to_eval_task(self) -> Task:
- files = dict(self.files)
- spec = dict(self.check_spec)
- golden = dict(self.golden_files)
-
- def setup(sb: Sandbox):
- sb.write_files(files)
-
- def check(sb: Sandbox, transcript: Transcript) -> tuple[bool, str]:
- return run_check_spec(sb, spec)
-
- return Task(
- task_id=self.task_id,
- prompt=self.prompt,
- setup=setup,
- check=check,
- max_turns=self.max_turns,
- system_prompt=self.system_prompt,
- )
-
- def self_verify(self, timeout: float = 10.0) -> tuple[bool, str]:
- with tempfile.TemporaryDirectory() as golden_dir:
- root = Path(golden_dir)
- for rel, content in self.golden_files.items():
- p = root / rel
- p.parent.mkdir(parents=True, exist_ok=True)
- p.write_text(content)
- golden_ok, golden_msg = _run_check_in_dir(root, self.check_spec, timeout)
- if not golden_ok:
- return False, f"golden FAIL: {golden_msg}"
-
- with tempfile.TemporaryDirectory() as buggy_dir:
- root = Path(buggy_dir)
- for rel, content in self.files.items():
- p = root / rel
- p.parent.mkdir(parents=True, exist_ok=True)
- p.write_text(content)
- buggy_ok, buggy_msg = _run_check_in_dir(root, self.check_spec, timeout)
- if buggy_ok:
- return False, f"buggy PASS (should fail): {buggy_msg}"
-
- return True, "verified"
-
-
-def run_check_spec(sb: Sandbox, spec: dict) -> tuple[bool, str]:
- check_type = spec["type"]
-
- if check_type == "bash_pass":
- cmd = spec["cmd"]
- expected = spec.get("expected", "")
- try:
- out = sb.bash(cmd)
- except Exception as e:
- return False, f"cmd failed: {e}"
- if expected and expected not in out:
- return False, f"expected '{expected}' not in output"
- return True, "pass"
-
- elif check_type == "grep_absent":
- pattern = spec["pattern"]
- try:
- out = sb.bash(f"grep -rn '{pattern}' . || true")
- except Exception:
- out = ""
- if pattern in out:
- return False, f"pattern '{pattern}' still found"
- return True, "pattern absent"
-
- elif check_type == "bash_pass_and_grep_absent":
- cmd = spec["cmd"]
- expected = spec.get("expected", "")
- pattern = spec["pattern"]
- try:
- out = sb.bash(cmd)
- except Exception as e:
- return False, f"cmd failed: {e}"
- if expected and expected not in out:
- return False, f"expected '{expected}' not in output"
- try:
- out2 = sb.bash(f"grep -rn '{pattern}' . || true")
- except Exception:
- out2 = ""
- if pattern in out2:
- return False, f"pattern '{pattern}' still found"
- return True, "pass + pattern absent"
-
- elif check_type == "file_contains":
- path_str = spec["path"]
- needle = spec["needle"]
- p = sb.workspace / path_str
- if not p.exists():
- return False, f"file {path_str} not found"
- content = p.read_text(errors="replace")
- if needle not in content:
- return False, f"'{needle}' not in {path_str}"
- return True, "file contains needle"
-
- else:
- return False, f"unknown check_type: {check_type}"
-
-
-def _run_check_in_dir(root: Path, spec: dict, timeout: float) -> tuple[bool, str]:
- check_type = spec["type"]
-
- if check_type in ("bash_pass", "bash_pass_and_grep_absent"):
- cmd = spec["cmd"]
- expected = spec.get("expected", "")
- try:
- r = subprocess.run(
- cmd, shell=True, cwd=str(root), capture_output=True,
- text=True, timeout=timeout,
- )
- except subprocess.TimeoutExpired:
- return False, "timeout"
- if r.returncode != 0:
- return False, f"exit {r.returncode}: {r.stderr[:200]}"
- if expected and expected not in r.stdout:
- return False, f"expected '{expected}' not in stdout"
-
- if check_type == "bash_pass_and_grep_absent":
- pattern = spec["pattern"]
- r2 = subprocess.run(
- f"grep -rn '{pattern}' .", shell=True, cwd=str(root),
- capture_output=True, text=True, timeout=5,
- )
- if r2.returncode == 0 and pattern in r2.stdout:
- return False, f"pattern '{pattern}' found"
- return True, "pass"
-
- elif check_type == "grep_absent":
- pattern = spec["pattern"]
- r = subprocess.run(
- f"grep -rn '{pattern}' .", shell=True, cwd=str(root),
- capture_output=True, text=True, timeout=5,
- )
- if r.returncode == 0 and pattern in r.stdout:
- return False, f"pattern found"
- return True, "absent"
-
- elif check_type == "file_contains":
- path = spec["path"]
- needle = spec["needle"]
- p = root / path
- if not p.exists():
- return False, "file missing"
- if needle not in p.read_text(errors="replace"):
- return False, "needle missing"
- return True, "contains"
-
- return False, f"unknown type: {check_type}"
-
-
-def save_hard_tasks(tasks: list[HardTask], path: str):
- with open(path, "w") as f:
- for t in tasks:
- f.write(json.dumps(asdict(t)) + "\n")
-
-
-def load_hard_tasks(path: str) -> list[HardTask]:
- tasks = []
- with open(path) as f:
- for line in f:
- if not line.strip():
- continue
- d = json.loads(line)
- tasks.append(HardTask(**d))
- return tasks
-
-
-def load_as_eval_tasks(path: str) -> list[Task]:
- return [ht.to_eval_task() for ht in load_hard_tasks(path)]
diff --git a/vendor/nanoai/agentic/hard_tasks/families.py b/vendor/nanoai/agentic/hard_tasks/families.py
deleted file mode 100644
index e89c9554d34c422ead6049fe9e04b4a92c5021dc..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/hard_tasks/families.py
+++ /dev/null
@@ -1,475 +0,0 @@
-"""Hard task family definitions.
-
-Each family provides:
- - skeleton(seed, params) -> (files, golden_files, prompt, check_spec)
- - fill_prompt(skeleton_context) -> str (prompt for 27B)
- - parse_fill(llm_response, skeleton_context) -> (files, golden_files)
-
-Families:
- H1: cross_file_import_fix — bug in dependency, error surfaces in caller
- H2: config_propagation — update hardcoded refs across multiple files
- H3: test_driven_impl — implement function from test assertions
- H4: needle_in_haystack_fix — find bug in large file with many functions
- H5: multi_grep_refactor — rename/replace pattern across N files
- H6: error_recovery_chain — config syntax error + code bug
- H7: dependency_upgrade — API signature change across project
- H8: detective_bug — trace root cause across call chain
-"""
-from __future__ import annotations
-
-import hashlib
-import json
-import random
-from dataclasses import dataclass
-from typing import Optional
-
-from . import HardTask
-
-NAMES_POOL = [
- "alpha", "beta", "gamma", "delta", "epsilon", "zeta", "theta", "kappa",
- "lambda_", "sigma", "omega", "phoenix", "atlas", "nova", "pulse", "spark",
- "forge", "nexus", "prism", "quantum", "ripple", "surge", "vertex", "zenith",
- "calc", "parser", "handler", "loader", "mapper", "filter_", "builder",
- "runner", "checker", "solver", "logger", "reader", "writer", "scanner",
- "tracker", "monitor", "config", "utils", "helpers", "core", "base",
- "engine", "service", "client", "server", "cache", "store", "queue",
- "worker", "manager", "factory", "adapter", "bridge", "proxy", "gate",
-]
-
-BUG_TYPES_LOGIC = [
- ("wrong_op", "+", "-"),
- ("wrong_op", "*", "+"),
- ("wrong_op", "//", "/"),
- ("wrong_cmp", ">", ">="),
- ("wrong_cmp", "<", "<="),
- ("wrong_cmp", "==", "!="),
- ("off_by_one", "+ 1", ""),
- ("off_by_one", "- 1", ""),
- ("swap_args", None, None),
- ("missing_return", None, None),
-]
-
-
-def _pick_names(rng: random.Random, n: int) -> list[str]:
- return rng.sample(NAMES_POOL, min(n, len(NAMES_POOL)))
-
-
-def _stable_id(family: str, seed: int) -> str:
- h = hashlib.md5(f"{family}_{seed}".encode()).hexdigest()[:8]
- return f"{family}_{h}"
-
-
-# ======================================================================
-# H1: cross_file_import_fix
-# ======================================================================
-
-def h1_cross_file_import_fix(seed: int, *, n_files: int = 3) -> Optional[HardTask]:
- rng = random.Random(seed)
- names = _pick_names(rng, 6)
- mod_main = "main"
- mod_mid = names[0]
- mod_dep = names[1]
- fn_main = names[2]
- fn_mid = names[3]
- fn_dep = names[4]
- var = names[5]
-
- bug_type = rng.choice(BUG_TYPES_LOGIC[:6])
- bug_label, bug_op, fix_op = bug_type
-
- test_input = rng.randint(5, 50)
-
- buggy_dep = f'def {fn_dep}({var}):\n return {var} {bug_op} 1\n'
- golden_dep = f'def {fn_dep}({var}):\n return {var} {fix_op} 1\n'
-
- ns = {}
- exec(golden_dep, ns)
- expected = ns[fn_dep](test_input)
- # Verify buggy gives different result
- ns2 = {}
- exec(buggy_dep, ns2)
- buggy_result = ns2[fn_dep](test_input)
- if buggy_result == expected:
- return None
- mid_file = f'from {mod_dep} import {fn_dep}\n\ndef {fn_mid}({var}):\n return {fn_dep}({var})\n'
- main_file = (
- f'from {mod_mid} import {fn_mid}\n\n'
- f'result = {fn_mid}({test_input})\n'
- f'print(result)\n'
- )
- test_file = (
- f'import subprocess, sys\n'
- f'r = subprocess.run([sys.executable, "main.py"], capture_output=True, text=True)\n'
- f'out = r.stdout.strip()\n'
- f'expected = str({repr(expected)})\n'
- f'if out == expected:\n'
- f' print("PASS")\n'
- f'else:\n'
- f' print(f"FAIL: got {{out!r}}, expected {{expected!r}}")\n'
- f' sys.exit(1)\n'
- )
-
- files = {
- "main.py": main_file,
- f"{mod_mid}.py": mid_file,
- f"{mod_dep}.py": buggy_dep,
- "test_check.py": test_file,
- }
- golden_files = dict(files)
- golden_files[f"{mod_dep}.py"] = golden_dep
-
- prompt = (
- f"Running `python3 main.py` gives the wrong result. "
- f"The expected output is `{expected}`. "
- f"The project has files: main.py, {mod_mid}.py, {mod_dep}.py. "
- f"Read each file to trace the bug, fix it with Edit, then run "
- f"`python3 test_check.py` to verify the fix prints PASS."
- )
-
- return HardTask(
- task_id=_stable_id("h1_cross_file", seed),
- family="h1_cross_file_import_fix",
- prompt=prompt,
- files=files,
- golden_files=golden_files,
- check_spec={"type": "bash_pass", "cmd": "python3 test_check.py", "expected": "PASS"},
- params={"seed": seed, "bug_type": bug_label, "n_files": n_files},
- difficulty=2,
- max_turns=8,
- )
-
-
-# ======================================================================
-# H3: test_driven_impl
-# ======================================================================
-
-_H3_SPECS = [
- {
- "desc": "compute the sum of squares of a list of integers",
- "fn_name": "sum_of_squares",
- "signature": "def {fn}(nums: list[int]) -> int:",
- "cases": [
- ("[]", "0"),
- ("[1, 2, 3]", "14"),
- ("[0]", "0"),
- ("[-1, -2]", "5"),
- ("[10]", "100"),
- ],
- "impl": " return sum(x * x for x in nums)",
- },
- {
- "desc": "count the number of vowels in a string (case-insensitive)",
- "fn_name": "count_vowels",
- "signature": "def {fn}(s: str) -> int:",
- "cases": [
- ("''", "0"),
- ("'hello'", "2"),
- ("'AEIOU'", "5"),
- ("'xyz'", "0"),
- ("'Python'", "1"),
- ],
- "impl": " return sum(1 for c in s.lower() if c in 'aeiou')",
- },
- {
- "desc": "find the longest word in a sentence",
- "fn_name": "longest_word",
- "signature": "def {fn}(sentence: str) -> str:",
- "cases": [
- ("'hello world'", "'hello'"),
- ("'I am fine'", "'fine'"),
- ("'a'", "'a'"),
- ("'the quick brown fox'", "'quick'"),
- ],
- "impl": " return max(sentence.split(), key=len)",
- },
- {
- "desc": "flatten a nested list (one level deep)",
- "fn_name": "flatten",
- "signature": "def {fn}(nested: list) -> list:",
- "cases": [
- ("[[1, 2], [3, 4]]", "[1, 2, 3, 4]"),
- ("[[]]", "[]"),
- ("[[1], [2], [3]]", "[1, 2, 3]"),
- ("[['a', 'b'], ['c']]", "['a', 'b', 'c']"),
- ],
- "impl": " return [item for sublist in nested for item in sublist]",
- },
- {
- "desc": "return a dict counting the frequency of each character in a string",
- "fn_name": "char_freq",
- "signature": "def {fn}(s: str) -> dict:",
- "cases": [
- ("''", "{}"),
- ("'aab'", "{'a': 2, 'b': 1}"),
- ("'hello'", "{'h': 1, 'e': 1, 'l': 2, 'o': 1}"),
- ],
- "impl": " d = {}\n for c in s:\n d[c] = d.get(c, 0) + 1\n return d",
- },
- {
- "desc": "check if a string is a palindrome (ignoring case and spaces)",
- "fn_name": "is_palindrome",
- "signature": "def {fn}(s: str) -> bool:",
- "cases": [
- ("'racecar'", "True"),
- ("'hello'", "False"),
- ("'A man a plan a canal Panama'", "True"),
- ("''", "True"),
- ],
- "impl": " cleaned = s.lower().replace(' ', '')\n return cleaned == cleaned[::-1]",
- },
- {
- "desc": "merge two sorted lists into one sorted list",
- "fn_name": "merge_sorted",
- "signature": "def {fn}(a: list, b: list) -> list:",
- "cases": [
- ("[], []", "[]"),
- ("[1, 3, 5], [2, 4, 6]", "[1, 2, 3, 4, 5, 6]"),
- ("[1], [2]", "[1, 2]"),
- ("[1, 2, 3], []", "[1, 2, 3]"),
- ],
- "impl": " return sorted(a + b)",
- },
- {
- "desc": "compute the nth Fibonacci number (0-indexed, fib(0)=0, fib(1)=1)",
- "fn_name": "fibonacci",
- "signature": "def {fn}(n: int) -> int:",
- "cases": [
- ("0", "0"),
- ("1", "1"),
- ("5", "5"),
- ("10", "55"),
- ],
- "impl": " if n <= 1:\n return n\n a, b = 0, 1\n for _ in range(2, n + 1):\n a, b = b, a + b\n return b",
- },
-]
-
-
-def h3_test_driven_impl(seed: int) -> Optional[HardTask]:
- rng = random.Random(seed)
- spec = rng.choice(_H3_SPECS)
- names = _pick_names(rng, 3)
- fn_name = names[0]
- mod_name = names[1]
-
- signature = spec["signature"].format(fn=fn_name)
- impl = spec["impl"]
-
- test_lines = [
- f"from {mod_name} import {fn_name}",
- "",
- "failures = 0",
- ]
- for case_idx, (args_str, expected_str) in enumerate(spec["cases"]):
- test_lines.append(f"result = {fn_name}({args_str})")
- test_lines.append(f"expected = {expected_str}")
- test_lines.append(f"if result != expected:")
- test_lines.append(f' print("FAIL case {case_idx}: got " + repr(result) + " expected " + repr(expected))')
- test_lines.append(f" failures += 1")
-
- test_lines.extend([
- "",
- "if failures == 0:",
- ' print("PASS")',
- "else:",
- " import sys; sys.exit(1)",
- ])
-
- test_content = "\n".join(test_lines) + "\n"
-
- stub = f'{signature}\n raise NotImplementedError("TODO: implement {fn_name}")\n'
- golden_impl = f"{signature}\n{impl}\n"
-
- files = {
- f"{mod_name}.py": stub,
- "test_check.py": test_content,
- }
- golden_files = {
- f"{mod_name}.py": golden_impl,
- "test_check.py": test_content,
- }
-
- prompt = (
- f"Implement the function `{fn_name}` in `{mod_name}.py`. "
- f"It should {spec['desc']}. "
- f"First read `test_check.py` to understand the expected behavior, "
- f"then edit `{mod_name}.py` to replace the stub with your implementation. "
- f"Run `python3 test_check.py` to verify — it should print PASS."
- )
-
- return HardTask(
- task_id=_stable_id("h3_test_impl", seed),
- family="h3_test_driven_impl",
- prompt=prompt,
- files=files,
- golden_files=golden_files,
- check_spec={"type": "bash_pass", "cmd": "python3 test_check.py", "expected": "PASS"},
- params={"seed": seed, "spec_fn": spec["fn_name"]},
- difficulty=2,
- max_turns=8,
- )
-
-
-# ======================================================================
-# H4: needle_in_haystack_fix
-# ======================================================================
-
-def h4_needle_in_haystack(seed: int, *, n_functions: int = 8) -> Optional[HardTask]:
- rng = random.Random(seed)
- names = _pick_names(rng, n_functions + 2)
- mod_name = names[0]
- fn_names = names[1:n_functions + 1]
- bug_idx = rng.randrange(n_functions)
-
- bug_type = rng.choice(BUG_TYPES_LOGIC[:6])
- bug_label, bug_op, fix_op = bug_type
-
- functions_buggy = []
- functions_golden = []
- for i, fn in enumerate(fn_names):
- val = rng.randint(2, 20)
- if i == bug_idx:
- functions_buggy.append(f"def {fn}(x):\n return x {bug_op} {val}\n")
- functions_golden.append(f"def {fn}(x):\n return x {fix_op} {val}\n")
- else:
- op = rng.choice(["+", "-", "*"])
- functions_buggy.append(f"def {fn}(x):\n return x {op} {val}\n")
- functions_golden.append(f"def {fn}(x):\n return x {op} {val}\n")
-
- test_input = rng.randint(5, 30)
- buggy_fn = fn_names[bug_idx]
- golden_code = "\n\n".join(functions_golden)
- buggy_code = "\n\n".join(functions_buggy)
-
- main_code = f"from {mod_name} import {buggy_fn}\nprint({buggy_fn}({test_input}))\n"
-
- exec_golden = {}
- exec(functions_golden[bug_idx], exec_golden)
- expected = exec_golden[buggy_fn](test_input)
- exec_buggy = {}
- exec(functions_buggy[bug_idx], exec_buggy)
- buggy_result = exec_buggy[buggy_fn](test_input)
- if buggy_result == expected:
- return None
-
- test_code = (
- f'import subprocess, sys\n'
- f'r = subprocess.run([sys.executable, "main.py"], capture_output=True, text=True)\n'
- f'out = r.stdout.strip()\n'
- f'if out == "{expected}":\n'
- f' print("PASS")\n'
- f'else:\n'
- f' print(f"FAIL: got {{out}}, expected {expected}")\n'
- f' sys.exit(1)\n'
- )
-
- files = {
- f"{mod_name}.py": buggy_code,
- "main.py": main_code,
- "test_check.py": test_code,
- }
- golden_files = dict(files)
- golden_files[f"{mod_name}.py"] = golden_code
-
- prompt = (
- f"Running `python3 main.py` gives the wrong result (expected `{expected}`). "
- f"The file `{mod_name}.py` has {n_functions} functions — "
- f"the bug is in the function `{buggy_fn}`. "
- f"Read `{mod_name}.py`, find the bug in `{buggy_fn}`, fix it with Edit, "
- f"then run `python3 test_check.py` to verify it prints PASS."
- )
-
- return HardTask(
- task_id=_stable_id("h4_needle", seed),
- family="h4_needle_in_haystack",
- prompt=prompt,
- files=files,
- golden_files=golden_files,
- check_spec={"type": "bash_pass", "cmd": "python3 test_check.py", "expected": "PASS"},
- params={"seed": seed, "n_functions": n_functions, "bug_idx": bug_idx, "bug_type": bug_label},
- difficulty=2,
- max_turns=10,
- )
-
-
-# ======================================================================
-# H5: multi_grep_refactor
-# ======================================================================
-
-def h5_multi_grep_refactor(seed: int, *, n_files: int = 4) -> Optional[HardTask]:
- rng = random.Random(seed)
- names = _pick_names(rng, n_files + 4)
- old_fn = names[0]
- new_fn = names[1]
- file_names = [names[i + 2] for i in range(n_files)]
-
- files = {}
- golden_files = {}
- used_wrappers = set()
- for fn in file_names:
- n_calls = rng.randint(1, 3)
- lines_buggy = []
- lines_golden = []
- for _ in range(n_calls):
- arg = rng.choice(["x", "data", "value", "result", "item"])
- lines_buggy.append(f" {old_fn}({arg})")
- lines_golden.append(f" {new_fn}({arg})")
- wrapper = fn + "_main"
- buggy = f"def {wrapper}():\n" + "\n".join(lines_buggy) + "\n"
- golden = f"def {wrapper}():\n" + "\n".join(lines_golden) + "\n"
- files[f"{fn}.py"] = buggy
- golden_files[f"{fn}.py"] = golden
-
- prompt = (
- f"The function `{old_fn}` has been renamed to `{new_fn}`. "
- f"Update all call sites across the project. "
- f"Files to check: {', '.join(f'{fn}.py' for fn in file_names)}. "
- f"Use Grep to find all occurrences of `{old_fn}(`, then Edit each file "
- f"to replace `{old_fn}` with `{new_fn}`. "
- f"Verify with `grep -rn '{old_fn}(' .` — it should return nothing."
- )
-
- return HardTask(
- task_id=_stable_id("h5_refactor", seed),
- family="h5_multi_grep_refactor",
- prompt=prompt,
- files=files,
- golden_files=golden_files,
- check_spec={"type": "grep_absent", "pattern": f"{old_fn}("},
- params={"seed": seed, "n_files": n_files, "old_fn": old_fn, "new_fn": new_fn},
- difficulty=2,
- max_turns=8,
- )
-
-
-# ======================================================================
-# Registry
-# ======================================================================
-
-FAMILY_GENERATORS = {
- "h1": h1_cross_file_import_fix,
- "h3": h3_test_driven_impl,
- "h4": h4_needle_in_haystack,
- "h5": h5_multi_grep_refactor,
-}
-
-
-def generate_task(family: str, seed: int, **kwargs) -> Optional[HardTask]:
- gen = FAMILY_GENERATORS.get(family)
- if gen is None:
- raise ValueError(f"Unknown family: {family}. Available: {list(FAMILY_GENERATORS.keys())}")
- return gen(seed, **kwargs)
-
-
-def generate_batch(
- families: list[str],
- n_per_family: int,
- base_seed: int = 42,
-) -> list[HardTask]:
- tasks = []
- for fam in families:
- for i in range(n_per_family):
- seed = base_seed + hash(fam) % 100_000 + i
- t = generate_task(fam, seed)
- if t is not None:
- tasks.append(t)
- return tasks
diff --git a/vendor/nanoai/agentic/ksr_labels.py b/vendor/nanoai/agentic/ksr_labels.py
deleted file mode 100644
index 70916bfcae1b984faffd5361f06131cd3537252e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/ksr_labels.py
+++ /dev/null
@@ -1,113 +0,0 @@
-"""K/S/R taxonomy for the 38 agent_eval_mt tasks.
-
-Each task is rated on three axes (0..3):
-
- K — Knowledge: how much Python/library/world knowledge the task needs.
- Encoded in FFN; addressed by CPT and SFT on factual data.
- S — Structure: tool routing, argument format, multi-tool sequencing.
- Encoded in attention; addressed by SFT on tool-call data.
- R — Reasoning: multi-step plan, observe→hypothesize→fix loop, recovery.
- Encoded in attention+FFN; addressed by reasoning SFT/RL.
-
-`primary_axis` picks the single axis whose failure mode most likely closes
-the gap on this task. It drives the bucket counters in ksr_report.py.
-
-Initial values are operator-set from the task definitions in eval_tasks.py
-(2026-05-17). Update them when a case study suggests a different bucket;
-the report's row labels are read from this file at runtime.
-"""
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import Optional
-
-
-@dataclass(frozen=True)
-class KSRLabel:
- K: int
- S: int
- R: int
- primary_axis: str # "K" | "S" | "R"
- swe_subtype: Optional[str] = None
-
-
-def _L(K: int, S: int, R: int, primary: str, swe_subtype: Optional[str] = None) -> KSRLabel:
- assert primary in ("K", "S", "R"), primary
- for ax, v in (("K", K), ("S", S), ("R", R)):
- assert 0 <= v <= 3, f"{ax}={v} out of [0,3]"
- return KSRLabel(K=K, S=S, R=R, primary_axis=primary, swe_subtype=swe_subtype)
-
-
-TASK_KSR: dict[str, KSRLabel] = {
- # ------------------------------------------------------------------
- # Read (7) — pure S: pick Read tool, pass a path arg.
- "read_short": _L(K=1, S=2, R=0, primary="S"),
- "read_config": _L(K=1, S=2, R=0, primary="S"),
- "read_script": _L(K=1, S=2, R=0, primary="S"),
- "read_pylib": _L(K=1, S=2, R=0, primary="S"),
- "read_data": _L(K=1, S=2, R=0, primary="S"),
- "read_nested": _L(K=1, S=2, R=0, primary="S"),
- "read_long": _L(K=1, S=2, R=0, primary="S"),
-
- # ------------------------------------------------------------------
- # Edit (8) — S dominant; +R=1 when content must satisfy a constraint
- # (e.g. JSON shape, replacement target, "new" content).
- "edit_create_hello": _L(K=1, S=2, R=1, primary="S"),
- "edit_create_readme": _L(K=0, S=2, R=0, primary="S"),
- "edit_create_json": _L(K=1, S=2, R=1, primary="S"),
- "edit_replace": _L(K=1, S=2, R=1, primary="S"),
- "edit_create_sh": _L(K=1, S=2, R=0, primary="S"),
- "edit_add_fn": _L(K=1, S=2, R=1, primary="S"),
- "edit_create_test": _L(K=1, S=2, R=0, primary="S"),
- "edit_overwrite": _L(K=0, S=2, R=1, primary="S"),
-
- # ------------------------------------------------------------------
- # Grep (7) — pure S.
- "grep_word": _L(K=1, S=2, R=0, primary="S"),
- "grep_class": _L(K=1, S=2, R=0, primary="S"),
- "grep_todo": _L(K=1, S=2, R=0, primary="S"),
- "grep_imports": _L(K=1, S=2, R=0, primary="S"),
- "grep_nomatch": _L(K=1, S=2, R=1, primary="S"),
- "grep_func": _L(K=1, S=2, R=0, primary="S"),
- "grep_torch": _L(K=1, S=2, R=0, primary="S"),
-
- # ------------------------------------------------------------------
- # Bash (8) — S + K (knows shell syntax / python -c).
- "bash_pwd": _L(K=1, S=2, R=0, primary="S"),
- "bash_ls": _L(K=1, S=2, R=0, primary="S"),
- "bash_count_py": _L(K=2, S=2, R=1, primary="K"), # find | wc -l
- "bash_wc": _L(K=2, S=2, R=0, primary="K"), # wc invocation
- "bash_head": _L(K=1, S=2, R=0, primary="S"),
- "bash_python": _L(K=2, S=2, R=0, primary="K"), # python -c '2+2'
- "bash_find_py": _L(K=2, S=2, R=0, primary="K"),
- "bash_grep_log": _L(K=2, S=2, R=0, primary="K"),
-
- # ------------------------------------------------------------------
- # SWE (8) — S=2 fixed, K and R vary; primary picks the bigger gap.
- "swe_fix_typo": _L(K=2, S=2, R=2, primary="K", swe_subtype="typo"),
- "swe_fix_indent": _L(K=2, S=2, R=2, primary="K", swe_subtype="indent"),
- "swe_fizzbuzz": _L(K=3, S=2, R=2, primary="K", swe_subtype="fizzbuzz"),
- "swe_factorial": _L(K=3, S=2, R=2, primary="K", swe_subtype="factorial"),
- "swe_missing_import": _L(K=2, S=2, R=2, primary="K", swe_subtype="missing_import"),
- "swe_off_by_one": _L(K=2, S=2, R=3, primary="R", swe_subtype="off_by_one"),
- "swe_replace_print": _L(K=2, S=2, R=2, primary="R", swe_subtype="replace_print"),
- "swe_grep_then_fix": _L(K=2, S=2, R=3, primary="R", swe_subtype="grep_then_fix"),
-}
-
-
-SWE_SUBTYPES: tuple[str, ...] = (
- "typo", "indent", "fizzbuzz", "factorial",
- "missing_import", "off_by_one", "replace_print", "grep_then_fix",
-)
-
-
-def category_of(task_id: str) -> str:
- """Coarse category from task_id prefix (read/edit/grep/bash/swe)."""
- return task_id.split("_", 1)[0]
-
-
-def expected_38_task_ids() -> tuple[str, ...]:
- return tuple(TASK_KSR.keys())
-
-
-assert len(TASK_KSR) == 38, f"expected 38 labels, got {len(TASK_KSR)}"
diff --git a/vendor/nanoai/agentic/mai_recipe.py b/vendor/nanoai/agentic/mai_recipe.py
deleted file mode 100644
index 05517c7a4c8ddc66aed317ca7e4c47ea48e9bcb5..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/mai_recipe.py
+++ /dev/null
@@ -1,144 +0,0 @@
-"""Active MAI-Thinking-1 RL recipe pieces (the parts that *change* training).
-
-Measurement-only primitives live in `rl_stability.py`; this module holds the
-small helpers that actually alter the training signal, so the two sibling
-trainers — `rl_math_nanoai_mai` (Stage A, math single-turn) and
-`agentic_dapo_mai` (Stage B, agentic multi-turn) — wire the *same* validated
-logic behind ablation flags instead of each re-deriving it.
-
-Each helper maps to one MAI technique / one flag:
- * `pass_rate_in_band` — §3.1.3 problem sampling (pass-rate band filter +
- early-exit). Only train on groups with a learnable mix of right/wrong.
- * `apply_outer_clip` — §3.1.1 outer/dual ratio clip (eq.9): a hard bound
- on ALL branches to kill gradient-norm spikes.
- * `difficulty_length_penalty`— §3.1.2 length penalty (eq.12): ρq·|y|/ℓmax, so
- hard problems (low pass rate) keep room to reason long while easy ones are
- pushed toward brevity.
- * `nucleus_mask_logits` — §3.1.3 top-p mask replay: re-apply the rollout's
- nucleus truncation during training so excluded tokens get prob 0, matching
- the sampling distribution and curbing off-policy drift.
-
-The adaptive-entropy clip *bounds* come from `rl_stability.EntropyController`;
-the trainer simply clamps the ratio to those bounds, so no helper is needed here.
-
-All helpers are pure and CPU-testable (`__main__`).
-"""
-from __future__ import annotations
-
-from typing import Optional
-
-import torch
-
-
-def pass_rate_in_band(n_pass: int, n_total: int, lo: float, hi: float) -> bool:
- """MAI §3.1.3 pass-rate filter: keep a group iff lo ≤ pass_rate ≤ hi.
-
- A group where (almost) every rollout is correct or every rollout is wrong has
- near-zero advantage variance and gives little relative learning signal. This
- is a strict generalization of nanoai's existing var≈0 skip
- (`agentic_dapo.py:711-718`): the band ALSO drops near-saturated groups
- (pass_rate≈1) and near-hopeless ones (pass_rate≈0), not just the exact-tie
- case. Used both for the cheap early-exit probe (Gearly rollouts) and for the
- full-group decision after all G rollouts.
-
- n_total==0 ⇒ False (no signal at all).
- """
- if n_total <= 0:
- return False
- pr = n_pass / n_total
- return lo <= pr <= hi
-
-
-def apply_outer_clip(
- ratio: torch.Tensor, rmax: float, rmin: float = 0.0
-) -> torch.Tensor:
- """MAI eq.(9) outer/dual ratio clip: hard bound on ALL branches.
-
- GRPO deliberately leaves two branches unclipped (Ai<0 & r>1; Ai>0 & r<1).
- Those occasionally produce catastrophic gradient-norm spikes. A large hard
- outer clamp (rmax≈50, rmin left at 0) discards only the extreme old/new
- discrepancies while preserving normal trust-region behavior. Apply this to
- the raw ratio *before* the standard PPO min(); see trainer wiring.
- """
- return torch.clamp(ratio, rmin, rmax)
-
-
-def difficulty_length_penalty(
- pass_rate: float, length: int, max_length: int, weight: float = 1.0
-) -> float:
- """MAI eq.(12) length penalty: Rlen = weight · ρq · |y| / ℓmax (a magnitude).
-
- Returns the penalty MAGNITUDE (≥0); the caller subtracts it from reward.
- `pass_rate` = ρq is the empirical pass rate of the problem/group (difficulty
- proxy): low ρq (hard) ⇒ small penalty (room to reason long); high ρq (easy)
- ⇒ larger penalty (push toward concise answers). Replaces the flat per-turn
- efficiency cost. Clamps the length fraction to [0,1] so an over-long rollout
- cannot exceed `weight·ρq`.
- """
- if max_length <= 0:
- return 0.0
- frac = min(max(length / max_length, 0.0), 1.0)
- return weight * max(pass_rate, 0.0) * frac
-
-
-def nucleus_mask_logits(
- logits: torch.Tensor, keep_mask: torch.Tensor, neg_inf: float = float("-inf")
-) -> torch.Tensor:
- """MAI §3.1.3 top-p mask replay: set logits outside the sampled nucleus to -inf.
-
- `keep_mask` (same shape as `logits`, bool/0-1) marks the tokens that were
- inside the nucleus at sampling time. Re-applying it before the training
- softmax makes the training distribution match the rollout distribution
- exactly over the nucleus (excluded tokens → probability 0), which DeepSeek and
- MAI both report substantially reduces policy divergence in off-policy RL.
-
- Returns a new tensor (does not mutate `logits`).
- """
- km = keep_mask.bool()
- return logits.masked_fill(~km, neg_inf)
-
-
-# ---- Self-test (CPU, part of the no-CUDA local gate) -------------------------
-
-if __name__ == "__main__":
- # 1. pass_rate_in_band: band [0.1, 0.8] keeps mixed groups, drops saturated /
- # hopeless / empty.
- assert pass_rate_in_band(4, 8, 0.1, 0.8) # 0.5 in band
- assert not pass_rate_in_band(8, 8, 0.1, 0.8) # all pass → drop
- assert not pass_rate_in_band(0, 8, 0.1, 0.8) # all fail → drop
- assert pass_rate_in_band(1, 8, 0.1, 0.8) # 0.125 in band
- assert not pass_rate_in_band(0, 0, 0.1, 0.8) # empty → drop
- # boundary inclusive
- assert pass_rate_in_band(1, 10, 0.1, 0.8) # exactly 0.1
- print("pass_rate_in_band OK")
-
- # 2. apply_outer_clip: caps extremes, leaves the normal range untouched.
- r = torch.tensor([0.5, 1.0, 2.0, 100.0, -3.0])
- out = apply_outer_clip(r, rmax=50.0, rmin=0.0)
- assert out.tolist() == [0.5, 1.0, 2.0, 50.0, 0.0], out.tolist()
- print("apply_outer_clip OK")
-
- # 3. difficulty_length_penalty: hard (low ρq) < easy (high ρq) for same length;
- # monotone in length; length fraction clamped to 1.
- p_hard = difficulty_length_penalty(0.1, 500, 1000, weight=0.25)
- p_easy = difficulty_length_penalty(0.8, 500, 1000, weight=0.25)
- assert p_easy > p_hard > 0, (p_hard, p_easy)
- assert difficulty_length_penalty(0.5, 1000, 1000, 1.0) == 0.5 # frac=1
- assert difficulty_length_penalty(0.5, 5000, 1000, 1.0) == 0.5 # frac clamped
- assert difficulty_length_penalty(0.5, 0, 1000, 1.0) == 0.0 # zero length
- assert difficulty_length_penalty(0.5, 500, 0, 1.0) == 0.0 # guard ℓmax=0
- print("difficulty_length_penalty OK")
-
- # 4. nucleus_mask_logits: excluded → -inf; softmax renormalizes over the kept
- # set and sums to 1; does not mutate the input.
- logits = torch.tensor([[2.0, 1.0, 0.5, -1.0]])
- keep = torch.tensor([[True, True, False, False]])
- masked = nucleus_mask_logits(logits, keep)
- assert torch.isinf(masked[0, 2]) and torch.isinf(masked[0, 3])
- probs = torch.softmax(masked, dim=-1)
- assert abs(float(probs.sum()) - 1.0) < 1e-6
- assert float(probs[0, 2]) == 0.0 and float(probs[0, 3]) == 0.0
- assert not torch.isinf(logits).any(), "must not mutate input"
- print("nucleus_mask_logits OK")
-
- print("mai_recipe self-test OK")
diff --git a/vendor/nanoai/agentic/mc_turn_value.py b/vendor/nanoai/agentic/mc_turn_value.py
deleted file mode 100644
index 90c3915c6bcb16ce0fc60a4c405eb2a70e7ade13..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/mc_turn_value.py
+++ /dev/null
@@ -1,304 +0,0 @@
-"""Monte Carlo Turn-Value estimation via forked rollouts.
-
-Fork rollouts from turn boundaries to estimate V(s_t) = E[pass | state at turn t].
-Per-turn advantage A(t) = V(s_{t+1}) - V(s_t) directly measures the value added
-by turn t.
-
-Provides shared infrastructure for:
- - Exp 4 (MC Turn-Value): offline V estimation for diagnostic/validation.
- - Exp 5 (Turn-Level GRPO): fork at a target turn for per-turn GRPO.
-"""
-from __future__ import annotations
-
-import time
-from dataclasses import dataclass
-from typing import Optional
-
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _execute,
- _generate_one_turn,
- _parse_tool_call,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-from .grpo_multi_rollout import MultiTurnRollout, TurnSpan
-
-
-def _find_transcript_tool_call(transcript: Transcript, turn_idx: int) -> Optional[dict]:
- """Find the tool_call dict for a specific turn in the transcript.
-
- Transcript turns alternate: assistant (with tool_call) → tool_result → ...
- Turn i's tool_call is in the (2*i)-th transcript turn.
- """
- asst_turns = [t for t in transcript.turns if t["role"] == "assistant"]
- if turn_idx < len(asst_turns):
- return asst_turns[turn_idx].get("tool_call")
- return None
-
-
-def _get_prefix_transcript_turns(transcript: Transcript, target_turn: int) -> list[dict]:
- """Extract transcript turns for turns 0..target_turn-1."""
- # Each assistant turn produces 1-2 transcript entries (assistant + optional tool_result).
- # We count assistant turns and include everything up to that point.
- asst_count = 0
- result = []
- for t in transcript.turns:
- if t["role"] == "assistant":
- if asst_count >= target_turn:
- break
- asst_count += 1
- result.append(t)
- return result
-
-
-def replay_sandbox_to_turn(
- sandbox: Sandbox,
- task: Task,
- rollout: MultiTurnRollout,
- target_turn: int,
-) -> None:
- """Setup sandbox and replay tool executions from turns 0..target_turn-1.
-
- After this, the sandbox is in the same state as the original rollout at
- the start of target_turn.
- """
- task.setup(sandbox)
- for ts in rollout.turn_spans[:target_turn]:
- if ts.has_tool_call:
- tc = _find_transcript_tool_call(rollout.transcript, ts.turn_idx)
- if tc:
- _execute(sandbox, tc["name"], tc["args"])
-
-
-def continue_rollout(
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- prefix_tokens: list[int],
- *,
- start_turn: int = 0,
- max_tokens_per_turn: int = 256,
- max_turns: int = 8,
- temperature: float = 0.8,
- seed: int = 42,
- prefix_transcript_turns: list[dict] | None = None,
-) -> MultiTurnRollout:
- """Continue a rollout from a given prefix + sandbox state.
-
- Like rollout_one() but skips task.setup() and _build_initial_tokens(),
- starting from the provided prefix_tokens and sandbox state.
-
- prefix_transcript_turns: if provided, these are prepended to the
- continuation transcript so task.check() sees the full conversation.
- """
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- tokens = list(prefix_tokens)
- prompt_len = len(tokens)
- transcript = Transcript(user_prompt="(continued)")
- if prefix_transcript_turns:
- transcript.turns.extend(prefix_transcript_turns)
- turn_spans: list[TurnSpan] = []
- n_turns = 0
- n_tool_errors = 0
- truncated = False
- t0 = time.time()
-
- remaining_turns = max_turns
- task_max = getattr(task, "max_turns", 0) or 0
- if task_max > 0:
- remaining_turns = max(remaining_turns, task_max - start_turn)
-
- while n_turns < remaining_turns:
- n_turns += 1
- body_start = len(tokens)
- out_ids = _generate_one_turn(
- engine, tokenizer, tokens,
- max_tokens_per_turn, temperature, seed + start_turn + n_turns,
- )
- tokens.extend(out_ids)
- body_end = len(tokens)
-
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
- if not ended_clean and len(out_ids) >= max_tokens_per_turn:
- truncated = True
-
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- j = -1
- if j > i:
- pre = tokenizer.decode(body_ids[:i]).strip()
- tool_body = tokenizer.decode(body_ids[i + 1 : j])
- try:
- name, args = _parse_tool_call_dispatch(tokenizer, tool_body)
- except ValueError as e:
- name, args = "", {}
- result_text, ok = f"tool parse error: {e}", False
- else:
- result_text, ok = _execute(sandbox, name, args)
- transcript.turns.append({
- "role": "assistant",
- "content": pre,
- "tool_call": {"name": name, "args": args},
- })
- transcript.turns.append({
- "role": "tool_result",
- "tool_result": result_text,
- "error": not ok,
- })
- if not ok:
- n_tool_errors += 1
-
- is_terminal = name in TERMINAL_TOOLS
- turn_spans.append(TurnSpan(
- turn_idx=start_turn + n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=True,
- tool_name=name,
- tool_ok=ok,
- is_terminal=is_terminal,
- ))
- if is_terminal:
- break
- if not ended_clean:
- tokens.append(asst_end)
- _append_tool_result(tokenizer, tokens, result_text)
- continue
-
- text = tokenizer.decode(body_ids).strip()
- transcript.turns.append({"role": "assistant", "content": text})
- turn_spans.append(TurnSpan(
- turn_idx=start_turn + n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=False,
- tool_name=None,
- tool_ok=None,
- is_terminal=True,
- ))
- break
-
- elapsed = time.time() - t0
- passed, reason = task.check(sandbox, transcript)
-
- return MultiTurnRollout(
- task_id=task.task_id,
- tokens=tokens,
- prompt_len=prompt_len,
- turn_spans=turn_spans,
- transcript=transcript,
- passed=passed,
- reason=reason,
- n_turns=n_turns,
- n_tool_errors=n_tool_errors,
- truncated=truncated,
- elapsed_s=round(elapsed, 2),
- )
-
-
-def fork_from_turn(
- engine,
- tokenizer,
- original_rollout: MultiTurnRollout,
- turn_idx: int,
- task: Task,
- *,
- fork_count: int = 4,
- max_tokens_per_turn: int = 256,
- max_turns: int = 8,
- temperature: float = 0.8,
- base_seed: int = 42,
- bash_timeout: float = 5.0,
-) -> list[MultiTurnRollout]:
- """Fork fork_count independent rollouts from turn_idx of original_rollout.
-
- Each fork:
- 1. Creates a fresh sandbox, replays tool calls 0..turn_idx-1.
- 2. Takes the token prefix up to turn_idx.
- 3. Continues with a different seed.
- 4. Passes the prefix transcript turns so task.check() sees full history.
- """
- T = len(original_rollout.turn_spans)
- if turn_idx > T:
- raise ValueError(f"turn_idx={turn_idx} > n_turns={T}")
- if turn_idx == T:
- # Already at terminal — V(s_T) is deterministic
- return []
-
- prefix = original_rollout.tokens[:original_rollout.turn_spans[turn_idx].start_idx]
- remaining = max(max_turns - turn_idx, 1)
-
- # Build prefix transcript for task.check()
- prefix_transcript = _get_prefix_transcript_turns(
- original_rollout.transcript, turn_idx
- )
-
- forks = []
- for k in range(fork_count):
- with Sandbox(bash_timeout=bash_timeout) as sb:
- replay_sandbox_to_turn(sb, task, original_rollout, turn_idx)
- fork = continue_rollout(
- engine, tokenizer, sb, task, prefix,
- start_turn=turn_idx,
- max_tokens_per_turn=max_tokens_per_turn,
- max_turns=remaining,
- temperature=temperature,
- seed=base_seed + k * 10_000,
- prefix_transcript_turns=prefix_transcript,
- )
- forks.append(fork)
- return forks
-
-
-def estimate_turn_values(
- engine,
- tokenizer,
- rollout: MultiTurnRollout,
- task: Task,
- *,
- fork_count: int = 4,
- max_tokens_per_turn: int = 256,
- max_turns: int = 8,
- temperature: float = 0.8,
- base_seed: int = 42,
- bash_timeout: float = 5.0,
-) -> tuple[list[float], list[float]]:
- """Estimate V(s_t) for t in 0..T via MC forking.
-
- Returns (values, advantages) where:
- values[t] = estimated V(s_t) for t in 0..T (T+1 entries)
- advantages[t] = V(s_{t+1}) - V(s_t) for t in 0..T-1 (T entries)
- """
- T = len(rollout.turn_spans)
- values = []
-
- for t in range(T):
- forks = fork_from_turn(
- engine, tokenizer, rollout, t, task,
- fork_count=fork_count,
- max_tokens_per_turn=max_tokens_per_turn,
- max_turns=max_turns,
- temperature=temperature,
- base_seed=base_seed + t * 100_000,
- bash_timeout=bash_timeout,
- )
- v_t = sum(1.0 if f.passed else 0.0 for f in forks) / len(forks)
- values.append(v_t)
-
- # V(s_T) = actual outcome (deterministic)
- values.append(1.0 if rollout.passed else 0.0)
-
- advantages = [values[t + 1] - values[t] for t in range(T)]
- return values, advantages
diff --git a/vendor/nanoai/agentic/mcts_exec.py b/vendor/nanoai/agentic/mcts_exec.py
deleted file mode 100644
index da5a4be1690f6007a00995b3ba8b2dd1271520f6..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/mcts_exec.py
+++ /dev/null
@@ -1,242 +0,0 @@
-"""Execution-based MCTS: execute each candidate's tool call, then score.
-
-Unlike mcts_inference.py which scores candidates without executing them,
-this version actually runs each candidate's tool call in an independent
-sandbox, appends the tool result tokens, and scores at the next turn
-boundary — the same position type the value model was trained on.
-
-Cost: K sandbox executions per turn (vs 1 for greedy).
-Benefit: value model sees genuinely different post-action states.
-"""
-from __future__ import annotations
-
-import copy
-import time
-from dataclasses import dataclass
-
-import torch
-
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _build_initial_tokens,
- _execute,
- _generate_one_turn,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-from .grpo_multi_rollout import MultiTurnRollout, TurnSpan
-from .value_head import PolicyWithValueHead
-
-
-@dataclass
-class ExecMCTSConfig:
- n_candidates: int = 3
- max_tokens_per_turn: int = 256
- max_turns: int = 8
- temperature: float = 0.5
- candidate_temperature: float = 1.0
- bash_timeout: float = 5.0
- seed: int = 42
-
-
-def _parse_candidate(tokenizer, out_ids, asst_end, tc_start, tc_end):
- """Parse a candidate's output into tool call or text."""
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
-
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- return None, None, None, body_ids, ended_clean
- if j > i:
- pre = tokenizer.decode(body_ids[:i]).strip()
- tool_body = tokenizer.decode(body_ids[i + 1 : j])
- try:
- name, args = _parse_tool_call_dispatch(tokenizer, tool_body)
- return name, args, pre, body_ids, ended_clean
- except ValueError:
- pass
- text = tokenizer.decode(body_ids).strip()
- return None, None, text, body_ids, ended_clean
-
-
-def exec_mcts_rollout(
- ac_model: PolicyWithValueHead,
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- config: ExecMCTSConfig,
-) -> MultiTurnRollout:
- """Multi-turn rollout with execution-based MCTS at each turn.
-
- For each turn:
- 1. Generate K candidates
- 2. For each candidate with a tool call:
- a. Execute tool in a FRESH sandbox (replay prior turns)
- b. Append tool result tokens
- c. Score at next turn boundary (value model sees post-action state)
- 3. Select highest-scoring candidate
- 4. Execute selected action in the main sandbox
- """
- device = ac_model.policy.get_device()
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- task.setup(sandbox)
- system_prompt = getattr(task, "system_prompt", None)
- tokens = _build_initial_tokens(tokenizer, task.prompt, system_prompt)
- prompt_len = len(tokens)
- transcript = Transcript(user_prompt=task.prompt)
- turn_spans: list[TurnSpan] = []
- n_turns = 0
- n_tool_errors = 0
- truncated = False
- t0 = time.time()
-
- # Track executed tool calls for sandbox replay
- executed_tools: list[tuple[str, dict]] = []
-
- remaining_turns = config.max_turns
- task_max = getattr(task, "max_turns", 0) or 0
- if task_max > 0:
- remaining_turns = max(remaining_turns, task_max)
-
- while n_turns < remaining_turns:
- n_turns += 1
- body_start = len(tokens)
-
- # Step 1: Generate K candidates
- candidates = []
- for k in range(config.n_candidates):
- seed_k = config.seed + n_turns if k == 0 else config.seed + n_turns * 1000 + k * 100
- temp = config.temperature if k == 0 else config.candidate_temperature
- out_ids = _generate_one_turn(
- engine, tokenizer, tokens,
- config.max_tokens_per_turn, temp, seed_k,
- )
- candidates.append(out_ids)
-
- # Step 2: Parse and score each candidate
- parsed = []
- for out_ids in candidates:
- name, args, text, body_ids, ended_clean = _parse_candidate(
- tokenizer, out_ids, asst_end, tc_start, tc_end
- )
- parsed.append((name, args, text, body_ids, ended_clean, out_ids))
-
- best_idx = 0
- if len(candidates) > 1:
- scores = []
- for k, (name, args, text, body_ids, ended_clean, out_ids) in enumerate(parsed):
- if name and name not in ("",) and name not in TERMINAL_TOOLS:
- # Has tool call: execute in fresh sandbox, score after result
- try:
- with Sandbox(bash_timeout=config.bash_timeout) as trial_sb:
- task.setup(trial_sb)
- # Replay all prior tool executions
- for prev_name, prev_args in executed_tools:
- _execute(trial_sb, prev_name, prev_args)
- # Execute this candidate's tool call
- result_text, ok = _execute(trial_sb, name, args)
- # Build token sequence with tool result appended
- trial_tokens = list(tokens) + list(out_ids)
- if not ended_clean:
- trial_tokens.append(asst_end)
- _append_tool_result(tokenizer, trial_tokens, result_text)
- # Score at next turn boundary (last position = new boundary)
- next_boundary = len(trial_tokens) - 1
- inp = torch.tensor([trial_tokens], dtype=torch.int32, device=device)
- with torch.no_grad():
- ac_model.policy(inp)
- v = ac_model.get_values_at_positions([next_boundary])
- scores.append(v.item())
- except Exception:
- scores.append(-1.0)
- else:
- # No tool call or terminal: score at last token (fallback)
- trial_tokens = list(tokens) + list(out_ids)
- inp = torch.tensor([trial_tokens], dtype=torch.int32, device=device)
- with torch.no_grad():
- ac_model.policy(inp)
- v = ac_model.get_values_at_positions([len(trial_tokens) - 1])
- scores.append(v.item())
-
- best_idx = max(range(len(scores)), key=lambda i: scores[i])
-
- # Step 3: Apply selected candidate to main sandbox
- name, args, text, body_ids, ended_clean, out_ids = parsed[best_idx]
- tokens.extend(out_ids)
- body_end = len(tokens)
-
- if not ended_clean and len(out_ids) >= config.max_tokens_per_turn:
- truncated = True
-
- if name and name not in ("",):
- result_text, ok = _execute(sandbox, name, args)
- executed_tools.append((name, args))
- transcript.turns.append({
- "role": "assistant",
- "content": text or "",
- "tool_call": {"name": name, "args": args},
- })
- transcript.turns.append({
- "role": "tool_result",
- "tool_result": result_text,
- "error": not ok,
- })
- if not ok:
- n_tool_errors += 1
-
- is_terminal = name in TERMINAL_TOOLS
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=True,
- tool_name=name,
- tool_ok=ok,
- is_terminal=is_terminal,
- ))
- if is_terminal:
- break
- if not ended_clean:
- tokens.append(asst_end)
- _append_tool_result(tokenizer, tokens, result_text)
- continue
-
- # No tool call — terminal text
- transcript.turns.append({"role": "assistant", "content": text or ""})
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=False,
- tool_name=None,
- tool_ok=None,
- is_terminal=True,
- ))
- break
-
- elapsed = time.time() - t0
- passed, reason = task.check(sandbox, transcript)
-
- return MultiTurnRollout(
- task_id=task.task_id,
- tokens=tokens,
- prompt_len=prompt_len,
- turn_spans=turn_spans,
- transcript=transcript,
- passed=passed,
- reason=reason,
- n_turns=n_turns,
- n_tool_errors=n_tool_errors,
- truncated=truncated,
- elapsed_s=round(elapsed, 2),
- )
diff --git a/vendor/nanoai/agentic/mcts_hybrid.py b/vendor/nanoai/agentic/mcts_hybrid.py
deleted file mode 100644
index 0760d2094b4d9d55e1970ba2239d161dd29ebe83..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/mcts_hybrid.py
+++ /dev/null
@@ -1,324 +0,0 @@
-"""Hybrid MCTS: mix d6 candidates + 27B teacher candidates.
-
-d6 generates K_local candidates (cheap, in-distribution).
-27B generates K_teacher candidates via API (expensive, high-quality).
-Value model scores all candidates after sandbox execution.
-Best candidate is selected and applied.
-
-This solves the diversity bottleneck: d6 candidates are often identical,
-but 27B candidates provide genuinely different (and often correct) actions.
-"""
-from __future__ import annotations
-
-import json
-import time
-import urllib.request
-from dataclasses import dataclass
-
-import torch
-
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _build_initial_tokens,
- _execute,
- _generate_one_turn,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-from .grpo_multi_rollout import MultiTurnRollout, TurnSpan
-from .value_head import PolicyWithValueHead
-
-
-@dataclass
-class HybridMCTSConfig:
- k_local: int = 1
- k_teacher: int = 1
- max_tokens_per_turn: int = 256
- max_turns: int = 8
- temperature: float = 0.5
- candidate_temperature: float = 1.0
- teacher_endpoint: str = "http://127.0.0.1:4210"
- teacher_model: str = "qwen36-27b"
- teacher_api_key: str = "sk-capagent"
- teacher_temperature: float = 0.7
- bash_timeout: float = 5.0
- seed: int = 42
-
-
-TEACHER_SYSTEM = """You are a coding assistant. You have these tools:
-- Read(file_path): Read a file
-- Edit(file_path, old_string, new_string): Replace text in a file
-- Grep(pattern, path): Search files
-- Bash(command): Run a shell command
-
-Respond with EXACTLY one tool call as JSON: {"tool": "ToolName", "args": {"key": "value"}}
-No explanation, just the JSON."""
-
-
-def _call_teacher(config: HybridMCTSConfig, messages: list[dict]) -> str:
- url = f"{config.teacher_endpoint.rstrip('/')}/v1/chat/completions"
- payload = json.dumps({
- "model": config.teacher_model,
- "messages": [{"role": "system", "content": TEACHER_SYSTEM}] + messages,
- "temperature": config.teacher_temperature,
- "max_tokens": 512,
- "chat_template_kwargs": {"enable_thinking": False},
- }).encode()
- headers = {"Content-Type": "application/json"}
- if config.teacher_api_key:
- headers["Authorization"] = f"Bearer {config.teacher_api_key}"
- try:
- req = urllib.request.Request(url, data=payload, headers=headers)
- with urllib.request.urlopen(req, timeout=30) as resp:
- data = json.loads(resp.read())
- return data["choices"][0]["message"]["content"]
- except Exception as e:
- return ""
-
-
-def _parse_teacher_response(response: str):
- """Parse 27B's JSON tool call response."""
- import re
- match = re.search(
- r'\{\s*"tool"\s*:\s*"(\w+)"\s*,\s*"args"\s*:\s*(\{[^{}]*\})\s*\}',
- response,
- )
- if match:
- name = match.group(1)
- try:
- args = json.loads(match.group(2))
- return name, args
- except json.JSONDecodeError:
- pass
- return None, None
-
-
-def _build_teacher_messages(task_prompt: str, transcript_turns: list[dict]) -> list[dict]:
- """Build message history for 27B from the current transcript."""
- messages = [{"role": "user", "content": task_prompt}]
- for turn in transcript_turns:
- role = turn["role"]
- if role == "assistant":
- tc = turn.get("tool_call")
- if tc:
- messages.append({
- "role": "assistant",
- "content": json.dumps({"tool": tc["name"], "args": tc["args"]}),
- })
- else:
- messages.append({"role": "assistant", "content": turn.get("content", "")})
- elif role == "tool_result":
- tr = turn.get("tool_result", "")[:500]
- messages.append({"role": "user", "content": f"Tool result:\n{tr}"})
- return messages
-
-
-def _teacher_candidate_to_tokens(
- tokenizer, name: str, args: dict, tokens_prefix: list[int]
-) -> list[int]:
- """Convert a teacher's tool call into d6 token format appended to prefix."""
- tc_start = tokenizer.encode_special("<|tool_call_start|>")
- tc_end = tokenizer.encode_special("<|tool_call_end|>")
- arg_tok = tokenizer.encode_special("<|tool_arg|>")
- val_tok = tokenizer.encode_special("<|tool_val|>")
- asst_end_tok = tokenizer.encode_special("<|assistant_end|>")
-
- def _as_list(x):
- return [x] if isinstance(x, int) else list(x)
-
- out = []
- out.extend(_as_list(tc_start))
- out.extend(tokenizer.encode(name))
- for k, v in args.items():
- out.extend(_as_list(arg_tok))
- out.extend(tokenizer.encode(k))
- out.extend(_as_list(val_tok))
- out.extend(tokenizer.encode(str(v)))
- out.extend(_as_list(tc_end))
- out.extend(_as_list(asst_end_tok))
- return out
-
-
-def hybrid_mcts_rollout(
- ac_model: PolicyWithValueHead,
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- config: HybridMCTSConfig,
-) -> MultiTurnRollout:
- """Multi-turn rollout with hybrid d6 + 27B candidates."""
- device = ac_model.policy.get_device()
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- task.setup(sandbox)
- system_prompt = getattr(task, "system_prompt", None)
- tokens = _build_initial_tokens(tokenizer, task.prompt, system_prompt)
- prompt_len = len(tokens)
- transcript = Transcript(user_prompt=task.prompt)
- turn_spans: list[TurnSpan] = []
- n_turns = 0
- n_tool_errors = 0
- truncated = False
- executed_tools: list[tuple[str, dict]] = []
- t0 = time.time()
-
- remaining = config.max_turns
- task_max = getattr(task, "max_turns", 0) or 0
- if task_max > 0:
- remaining = max(remaining, task_max)
-
- while n_turns < remaining:
- n_turns += 1
- body_start = len(tokens)
-
- # Generate d6 candidates
- local_candidates = []
- for k in range(config.k_local):
- seed_k = config.seed + n_turns if k == 0 else config.seed + n_turns * 1000 + k * 100
- temp = config.temperature if k == 0 else config.candidate_temperature
- out_ids = _generate_one_turn(
- engine, tokenizer, tokens,
- config.max_tokens_per_turn, temp, seed_k,
- )
- local_candidates.append(("d6", out_ids))
-
- # Generate 27B teacher candidate
- teacher_candidates = []
- if config.k_teacher > 0:
- teacher_msgs = _build_teacher_messages(task.prompt, transcript.turns)
- for _ in range(config.k_teacher):
- response = _call_teacher(config, teacher_msgs)
- t_name, t_args = _parse_teacher_response(response)
- if t_name:
- t_tokens = _teacher_candidate_to_tokens(
- tokenizer, t_name, t_args, tokens
- )
- teacher_candidates.append(("27b", t_tokens, t_name, t_args))
-
- # Score all candidates via execution
- all_scored = []
-
- # Score d6 candidates
- for source, out_ids in local_candidates:
- ended = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended else out_ids
- tc_name, tc_args = None, None
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- tb = tokenizer.decode(body_ids[i + 1:j])
- tc_name, tc_args = _parse_tool_call_dispatch(tokenizer, tb)
- except (ValueError, Exception):
- pass
-
- score = -1.0
- if tc_name and tc_name not in TERMINAL_TOOLS:
- try:
- with Sandbox(bash_timeout=config.bash_timeout) as trial_sb:
- task.setup(trial_sb)
- for pn, pa in executed_tools:
- _execute(trial_sb, pn, pa)
- result, ok = _execute(trial_sb, tc_name, tc_args)
- trial_tokens = list(tokens) + list(out_ids)
- if not ended:
- trial_tokens.append(asst_end)
- _append_tool_result(tokenizer, trial_tokens, result)
- inp = torch.tensor([trial_tokens], dtype=torch.int32, device=device)
- with torch.no_grad():
- ac_model.policy(inp)
- score = ac_model.get_values_at_positions([len(trial_tokens) - 1]).item()
- except Exception:
- score = -1.0
- all_scored.append({
- "source": source, "out_ids": out_ids, "tc_name": tc_name,
- "tc_args": tc_args, "score": score, "ended": ended,
- })
-
- # Score teacher candidates
- for source, t_out_ids, t_name, t_args in teacher_candidates:
- score = -1.0
- try:
- with Sandbox(bash_timeout=config.bash_timeout) as trial_sb:
- task.setup(trial_sb)
- for pn, pa in executed_tools:
- _execute(trial_sb, pn, pa)
- result, ok = _execute(trial_sb, t_name, t_args)
- trial_tokens = list(tokens) + t_out_ids
- _append_tool_result(tokenizer, trial_tokens, result)
- inp = torch.tensor([trial_tokens], dtype=torch.int32, device=device)
- with torch.no_grad():
- ac_model.policy(inp)
- score = ac_model.get_values_at_positions([len(trial_tokens) - 1]).item()
- except Exception:
- score = -1.0
- all_scored.append({
- "source": source, "out_ids": t_out_ids, "tc_name": t_name,
- "tc_args": t_args, "score": score, "ended": True,
- })
-
- # Select best
- if not all_scored:
- break
- best = max(all_scored, key=lambda x: x["score"])
- out_ids = best["out_ids"]
- tc_name = best["tc_name"]
- tc_args = best["tc_args"]
- ended = best["ended"]
-
- tokens.extend(out_ids)
- body_end = len(tokens)
-
- if tc_name and tc_name not in ("",):
- result_text, ok = _execute(sandbox, tc_name, tc_args)
- executed_tools.append((tc_name, tc_args))
-
- pre = ""
- transcript.turns.append({
- "role": "assistant", "content": pre,
- "tool_call": {"name": tc_name, "args": tc_args},
- "source": best["source"],
- })
- transcript.turns.append({
- "role": "tool_result", "tool_result": result_text, "error": not ok,
- })
- if not ok:
- n_tool_errors += 1
-
- is_terminal = tc_name in TERMINAL_TOOLS
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1, start_idx=body_start, end_idx=body_end,
- has_tool_call=True, tool_name=tc_name, tool_ok=ok,
- is_terminal=is_terminal,
- ))
- if is_terminal:
- break
- if not ended:
- tokens.append(asst_end)
- _append_tool_result(tokenizer, tokens, result_text)
- continue
-
- text = tokenizer.decode(out_ids).strip() if out_ids else ""
- transcript.turns.append({"role": "assistant", "content": text, "source": best["source"]})
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1, start_idx=body_start, end_idx=body_end,
- has_tool_call=False, tool_name=None, tool_ok=None, is_terminal=True,
- ))
- break
-
- elapsed = time.time() - t0
- passed, reason = task.check(sandbox, transcript)
-
- return MultiTurnRollout(
- task_id=task.task_id, tokens=tokens, prompt_len=prompt_len,
- turn_spans=turn_spans, transcript=transcript,
- passed=passed, reason=reason, n_turns=n_turns,
- n_tool_errors=n_tool_errors, truncated=truncated,
- elapsed_s=round(elapsed, 2),
- )
diff --git a/vendor/nanoai/agentic/mcts_inference.py b/vendor/nanoai/agentic/mcts_inference.py
deleted file mode 100644
index 555a36012150e08666fe258f0df8acb83f838fe1..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/mcts_inference.py
+++ /dev/null
@@ -1,237 +0,0 @@
-"""MCTS-style inference for multi-turn agentic tasks (Plan C).
-
-At each turn, instead of greedily taking the model's single generation,
-expand K candidate actions (different tool calls), evaluate each with the
-value model, and select the best one.
-
-This is a simplified "1-ply search" (no deep tree), analogous to AlphaGo's
-value network evaluation at leaf nodes. The policy network provides the
-candidate actions (via temperature sampling), and the value model scores them.
-
-Full MCTS (multi-ply) is possible but requires sandbox snapshotting (fork/exec)
-which is expensive. 1-ply search is cheap: K forward passes per turn.
-"""
-from __future__ import annotations
-
-import time
-from dataclasses import dataclass
-from typing import Optional
-
-import torch
-
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _execute,
- _generate_one_turn,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-from .grpo_multi_rollout import MultiTurnRollout, TurnSpan
-from .value_head import PolicyWithValueHead, collect_turn_boundary_positions
-
-
-@dataclass
-class MCTSConfig:
- n_candidates: int = 3
- max_tokens_per_turn: int = 256
- max_turns: int = 8
- temperature: float = 1.0
- candidate_temperature: float = 1.2
- bash_timeout: float = 5.0
- seed: int = 42
-
-
-def _score_candidate_last_token(
- ac_model: PolicyWithValueHead,
- tokenizer,
- tokens: list[int],
- device: torch.device,
-) -> float:
- """Score by V at last token (sees action but out-of-distribution for value model)."""
- inp = torch.tensor([tokens], dtype=torch.int32, device=device)
- with torch.no_grad():
- ac_model.policy(inp)
- v = ac_model.get_values_at_positions([len(tokens) - 1])
- return v.item()
-
-
-def _score_after_execution(
- ac_model: PolicyWithValueHead,
- tokenizer,
- tokens_with_result: list[int],
- next_turn_boundary: int,
- device: torch.device,
-) -> float:
- """Score AFTER executing the candidate and appending tool result.
-
- This queries V at the next turn boundary — the position where the model
- has seen the full action + tool result and is about to generate the next
- turn. This is the same position type the value model was trained on.
- """
- inp = torch.tensor([tokens_with_result], dtype=torch.int32, device=device)
- with torch.no_grad():
- ac_model.policy(inp)
- pos = min(next_turn_boundary, len(tokens_with_result) - 1)
- v = ac_model.get_values_at_positions([pos])
- return v.item()
-
-
-def mcts_rollout(
- ac_model: PolicyWithValueHead,
- engine,
- tokenizer,
- sandbox: Sandbox,
- task: Task,
- config: MCTSConfig,
-) -> MultiTurnRollout:
- """Run a multi-turn rollout with 1-ply MCTS at each turn.
-
- At each turn:
- 1. Generate K candidate continuations (different seeds → different tool calls)
- 2. Score each candidate with the value model
- 3. Select the highest-scoring candidate
- 4. Execute the selected action in the sandbox
- 5. Continue to next turn
- """
- from .eval_loop import _build_initial_tokens
- device = ac_model.policy.get_device()
-
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- task.setup(sandbox)
- system_prompt = getattr(task, "system_prompt", None)
- tokens = _build_initial_tokens(tokenizer, task.prompt, system_prompt)
- prompt_len = len(tokens)
- transcript = Transcript(user_prompt=task.prompt)
- turn_spans: list[TurnSpan] = []
- n_turns = 0
- n_tool_errors = 0
- truncated = False
- t0 = time.time()
-
- remaining_turns = config.max_turns
- task_max = getattr(task, "max_turns", 0) or 0
- if task_max > 0:
- remaining_turns = max(remaining_turns, task_max)
-
- while n_turns < remaining_turns:
- n_turns += 1
- body_start = len(tokens)
-
- # Generate K candidates
- # Candidate 0 uses the same seed pattern as rollout_one (seed + n_turns)
- # so that when value model picks candidate 0, behavior matches greedy exactly.
- # Other candidates use offset seeds + higher temperature for diversity.
- candidates = []
- for k in range(config.n_candidates):
- seed_k = config.seed + n_turns if k == 0 else config.seed + n_turns * 1000 + k * 100
- temp = config.temperature if k == 0 else config.candidate_temperature
- out_ids = _generate_one_turn(
- engine, tokenizer, tokens,
- config.max_tokens_per_turn, temp, seed_k,
- )
- candidates.append(out_ids)
-
- # Score each candidate with value model
- # Query V at the LAST token of each candidate (sees the full action)
- if len(candidates) > 1:
- scores = []
- for out_ids in candidates:
- candidate_tokens = tokens + out_ids
- score = _score_candidate_last_token(
- ac_model, tokenizer, candidate_tokens, device
- )
- scores.append(score)
- best_idx = max(range(len(scores)), key=lambda i: scores[i])
- else:
- best_idx = 0
-
- out_ids = candidates[best_idx]
- tokens.extend(out_ids)
- body_end = len(tokens)
-
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
- if not ended_clean and len(out_ids) >= config.max_tokens_per_turn:
- truncated = True
-
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- j = -1
- if j > i:
- pre = tokenizer.decode(body_ids[:i]).strip()
- tool_body = tokenizer.decode(body_ids[i + 1 : j])
- try:
- name, tool_args = _parse_tool_call_dispatch(tokenizer, tool_body)
- except ValueError:
- name, tool_args = "", {}
- result_text, ok = "tool parse error", False
- else:
- result_text, ok = _execute(sandbox, name, tool_args)
- transcript.turns.append({
- "role": "assistant",
- "content": pre,
- "tool_call": {"name": name, "args": tool_args},
- })
- transcript.turns.append({
- "role": "tool_result",
- "tool_result": result_text,
- "error": not ok,
- })
- if not ok:
- n_tool_errors += 1
-
- is_terminal = name in TERMINAL_TOOLS
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=True,
- tool_name=name,
- tool_ok=ok,
- is_terminal=is_terminal,
- ))
- if is_terminal:
- break
- if not ended_clean:
- tokens.append(asst_end)
- _append_tool_result(tokenizer, tokens, result_text)
- continue
-
- text = tokenizer.decode(body_ids).strip()
- transcript.turns.append({"role": "assistant", "content": text})
- turn_spans.append(TurnSpan(
- turn_idx=n_turns - 1,
- start_idx=body_start,
- end_idx=body_end,
- has_tool_call=False,
- tool_name=None,
- tool_ok=None,
- is_terminal=True,
- ))
- break
-
- elapsed = time.time() - t0
- passed, reason = task.check(sandbox, transcript)
-
- return MultiTurnRollout(
- task_id=task.task_id,
- tokens=tokens,
- prompt_len=prompt_len,
- turn_spans=turn_spans,
- transcript=transcript,
- passed=passed,
- reason=reason,
- n_turns=n_turns,
- n_tool_errors=n_tool_errors,
- truncated=truncated,
- elapsed_s=round(elapsed, 2),
- )
diff --git a/vendor/nanoai/agentic/ocaw.py b/vendor/nanoai/agentic/ocaw.py
deleted file mode 100644
index dff694e24525e90ff9e255e1e016d3a65b658c98..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/ocaw.py
+++ /dev/null
@@ -1,250 +0,0 @@
-"""Outcome-Conditioned Advantage Weighting (OCAW).
-
-Zero-training-cost credit assignment via conditional pass-rate statistics.
-Given a corpus of rollouts with known outcomes, compute for each
-(task_family, turn_idx, tool_name) triple:
-
- advantage = P(pass | family, turn, tool) - P(pass | family)
-
-Positive advantage means this tool choice at this turn correlates with success.
-
-The OCAW table can be:
-1. Used as a diagnostic (which turns/tools actually matter?).
-2. Plugged into DAPO/GTPO pack_for_grad to modulate per-turn advantages.
-3. Used to guide turn selection in Turn-Level GRPO (Exp 5).
-"""
-from __future__ import annotations
-
-import re
-from collections import defaultdict
-from dataclasses import dataclass
-from typing import Optional
-
-
-@dataclass
-class OCAWCell:
- """Statistics for one (task_family, turn_idx, tool_name) cell."""
- n_pass: int = 0
- n_fail: int = 0
-
- @property
- def n_total(self) -> int:
- return self.n_pass + self.n_fail
-
- def cond_pass_rate(self, laplace: float = 1.0) -> float:
- """Laplace-smoothed conditional pass rate."""
- return (self.n_pass + laplace) / (self.n_total + 2 * laplace)
-
-
-def _task_family(task_id: str) -> str:
- """Extract task family from task_id by stripping variant suffix.
-
- Examples:
- "file_search_v3" -> "file_search"
- "grep_count_v12" -> "grep_count"
- "swe_impl" -> "swe_impl"
- """
- return re.sub(r"(?:_v\d+)+$", "", task_id)
-
-
-def compute_ocaw_table(
- rollout_dicts: list[dict],
- *,
- min_observations: int = 5,
- laplace: float = 1.0,
-) -> dict:
- """Compute OCAW advantage table from rollout dicts.
-
- rollout_dicts: list of dicts with keys: task_id, passed, turn_spans.
- Each turn_span has: turn_idx, tool_name (str or None), tool_ok, etc.
-
- Returns dict mapping string key "family|turn_idx|tool_name" -> {
- "cond_pass_rate": float,
- "base_rate": float,
- "advantage": float,
- "n_observations": int,
- "n_pass": int,
- "n_fail": int,
- }
- """
- # Phase 1: accumulate counts
- cells: dict[tuple[str, int, str], OCAWCell] = defaultdict(OCAWCell)
- family_counts: dict[str, OCAWCell] = defaultdict(OCAWCell)
-
- for d in rollout_dicts:
- family = _task_family(d["task_id"])
- passed = d["passed"]
-
- if passed:
- family_counts[family].n_pass += 1
- else:
- family_counts[family].n_fail += 1
-
- for ts in d.get("turn_spans", []):
- turn_idx = ts["turn_idx"]
- tool_name = ts.get("tool_name") or "__no_tool__"
- key = (family, turn_idx, tool_name)
- if passed:
- cells[key].n_pass += 1
- else:
- cells[key].n_fail += 1
-
- # Phase 2: compute advantages
- table = {}
- for (family, turn_idx, tool_name), cell in cells.items():
- if cell.n_total < min_observations:
- continue
- fam = family_counts[family]
- base_rate = fam.cond_pass_rate(laplace)
- cond_rate = cell.cond_pass_rate(laplace)
- advantage = cond_rate - base_rate
-
- str_key = f"{family}|{turn_idx}|{tool_name}"
- table[str_key] = {
- "cond_pass_rate": round(cond_rate, 4),
- "base_rate": round(base_rate, 4),
- "advantage": round(advantage, 4),
- "n_observations": cell.n_total,
- "n_pass": cell.n_pass,
- "n_fail": cell.n_fail,
- }
-
- return table
-
-
-def chi_squared_test(cell: OCAWCell, family_cell: OCAWCell) -> float:
- """Chi-squared test p-value for independence between this cell and outcome.
-
- Tests H0: P(pass|tool at this turn) = P(pass|family).
- Returns p-value. Low p (<0.05) = tool choice at this turn is statistically
- associated with outcome.
- """
- import math
-
- n = cell.n_total
- if n < 5:
- return 1.0
-
- base_rate = family_cell.n_pass / max(family_cell.n_total, 1)
- expected_pass = n * base_rate
- expected_fail = n * (1 - base_rate)
-
- if expected_pass < 1 or expected_fail < 1:
- return 1.0
-
- chi2 = ((cell.n_pass - expected_pass) ** 2 / expected_pass +
- (cell.n_fail - expected_fail) ** 2 / expected_fail)
-
- # 1-df chi-squared survival function (approximation via normal)
- z = math.sqrt(max(chi2, 0))
- # Upper tail of standard normal (Abramowitz & Stegun approximation)
- if z > 8:
- return 0.0
- t = 1.0 / (1.0 + 0.2316419 * z)
- d = 0.3989422804014327 # 1/sqrt(2*pi)
- p = d * math.exp(-z * z / 2.0) * t * (
- 0.319381530 + t * (-0.356563782 + t * (1.781477937 +
- t * (-1.821255978 + t * 1.330274429)))
- )
- return 2.0 * p # two-tailed
-
-
-def compute_ocaw_with_stats(
- rollout_dicts: list[dict],
- *,
- min_observations: int = 5,
- laplace: float = 1.0,
-) -> tuple[dict, dict]:
- """Like compute_ocaw_table but also returns chi-squared p-values.
-
- Returns (table, summary) where summary has aggregate statistics.
- """
- cells: dict[tuple[str, int, str], OCAWCell] = defaultdict(OCAWCell)
- family_counts: dict[str, OCAWCell] = defaultdict(OCAWCell)
-
- for d in rollout_dicts:
- family = _task_family(d["task_id"])
- passed = d["passed"]
- if passed:
- family_counts[family].n_pass += 1
- else:
- family_counts[family].n_fail += 1
- for ts in d.get("turn_spans", []):
- turn_idx = ts["turn_idx"]
- tool_name = ts.get("tool_name") or "__no_tool__"
- key = (family, turn_idx, tool_name)
- if passed:
- cells[key].n_pass += 1
- else:
- cells[key].n_fail += 1
-
- table = {}
- n_significant = 0
- n_attributable = 0
- n_total_cells = 0
- advantages = []
-
- for (family, turn_idx, tool_name), cell in cells.items():
- if cell.n_total < min_observations:
- continue
- n_total_cells += 1
- fam = family_counts[family]
- base_rate = fam.cond_pass_rate(laplace)
- cond_rate = cell.cond_pass_rate(laplace)
- advantage = cond_rate - base_rate
- p_val = chi_squared_test(cell, fam)
-
- if p_val < 0.05:
- n_significant += 1
- if abs(advantage) > 0.1:
- n_attributable += 1
- advantages.append(advantage)
-
- str_key = f"{family}|{turn_idx}|{tool_name}"
- table[str_key] = {
- "cond_pass_rate": round(cond_rate, 4),
- "base_rate": round(base_rate, 4),
- "advantage": round(advantage, 4),
- "n_observations": cell.n_total,
- "n_pass": cell.n_pass,
- "n_fail": cell.n_fail,
- "chi2_p": round(p_val, 6),
- }
-
- summary = {
- "n_rollouts": len(rollout_dicts),
- "n_families": len(family_counts),
- "n_cells_total": n_total_cells,
- "n_cells_significant_p05": n_significant,
- "frac_significant": round(n_significant / max(n_total_cells, 1), 4),
- "n_cells_attributable_01": n_attributable,
- "frac_attributable": round(n_attributable / max(n_total_cells, 1), 4),
- "advantage_mean": round(sum(advantages) / max(len(advantages), 1), 4),
- "advantage_cell_std": round(
- (sum((a - sum(advantages)/max(len(advantages),1))**2
- for a in advantages) / max(len(advantages), 1)) ** 0.5, 4
- ) if advantages else 0.0,
- "overall_pass_rate": round(
- sum(f.n_pass for f in family_counts.values()) /
- max(sum(f.n_total for f in family_counts.values()), 1), 4
- ),
- }
-
- return table, summary
-
-
-def lookup_ocaw_weight(
- table: dict,
- task_id: str,
- turn_idx: int,
- tool_name: Optional[str],
-) -> float:
- """Look up OCAW advantage for a specific (task, turn, tool) triple.
-
- Returns 0.0 if not found (no modulation).
- """
- family = _task_family(task_id)
- tn = tool_name or "__no_tool__"
- key = f"{family}|{turn_idx}|{tn}"
- entry = table.get(key)
- return entry["advantage"] if entry else 0.0
diff --git a/vendor/nanoai/agentic/opd_loss.py b/vendor/nanoai/agentic/opd_loss.py
deleted file mode 100644
index 9f59d29fe1d60bab55bd626d1deab4924ec3c34e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/opd_loss.py
+++ /dev/null
@@ -1,449 +0,0 @@
-"""On-Policy Distillation loss utilities.
-
-Four levels of loss reduction (per-token / per-turn / per-step / per-sequence)
-and several divergence forms (reverse-KL / forward-KL / JSD(beta) / skew-KL).
-
-Reference implementations:
-- Thinking Machines OPD blog (2025-10): per-token reverse-KL, advantage = -kl,
- frozen teacher, plug into KL-regularized RL trainer.
-- SOD (arXiv 2605.07725): step-wise adaptive weights:
- w_1 = 1
- w_k = min( prod_{u (B, T) tensor
- ^ TM/MiniLLM default. Only needs logp at the realized token.
-
- full_kl_per_pos(log_full_student, log_full_teacher, reverse=True) -> (B, T)
- ^ Forward (reverse=False) or reverse KL across the whole vocab.
- Use only when you can afford the full softmax — V1 ablation only.
-
- jsd_per_pos(log_full_student, log_full_teacher, beta=0.5) -> (B, T)
-
- skew_kl_per_pos(log_full_student, log_full_teacher, alpha=0.1) -> (B, T)
-
- Step-level helpers (V3 SOD):
- step_mean_kl(per_token_kl, step_idx, K, mask) -> (B, K)
- sod_adaptive_weights(step_kl, delta=0.5, eps=1e-3) -> (B, K)
-
- Segment helpers (V4 SAD):
- SEGMENTS = ('reason', 'action', 'tool_arg', 'final_answer', 'tool_result')
- DEFAULT_SEG_WEIGHTS = {'reason': 1.0, 'action': 1.5, 'tool_arg': 1.2,
- 'final_answer': 2.0, 'tool_result': 0.0}
- identify_segments(token_ids, special_ids) -> (B, T) int8 tensor
-
- Loss reductions:
- reduce_per_token(per_token_kl, mask) -> scalar
- reduce_per_step(per_token_kl, step_idx, step_weight, mask) -> scalar
- reduce_per_sequence(per_token_kl, mask) -> scalar (mean over batch of
- length-normalized per-sequence KL)
-
-All tensor inputs use (B, T) with mask=1 for student-emitted tokens to be
-supervised and mask=0 elsewhere (user / system / tool_result spans).
-"""
-
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import Optional
-
-import torch
-import torch.nn.functional as F
-
-
-# ---------------------------------------------------------------------------
-# Per-token divergences
-# ---------------------------------------------------------------------------
-
-
-def reverse_kl_per_token(
- logp_student: torch.Tensor, logp_teacher: torch.Tensor
-) -> torch.Tensor:
- """Per-token reverse KL evaluated at the realized token.
-
- reverse_KL(student || teacher) ~= E_{y~student}[ log pi_S(y) - log pi_T(y) ]
-
- For on-policy distillation the sampled token IS y, so per-position the
- estimator is just (logp_student - logp_teacher). Stop-gradient on
- logp_teacher is the caller's responsibility (teacher is frozen).
-
- Both inputs are (B, T). Returns (B, T).
- """
- return logp_student - logp_teacher
-
-
-def full_kl_per_pos(
- log_full_student: torch.Tensor,
- log_full_teacher: torch.Tensor,
- reverse: bool = True,
-) -> torch.Tensor:
- """Full-vocab KL at every position.
-
- log_full_*: (B, T, V) — log-softmax of the model logits.
- reverse=True -> KL(student || teacher) = sum_v p_S * (log p_S - log p_T)
- reverse=False -> KL(teacher || student) (forward KL)
-
- Returns (B, T).
- """
- if reverse:
- p = log_full_student.exp()
- return (p * (log_full_student - log_full_teacher)).sum(dim=-1)
- else:
- p = log_full_teacher.exp()
- return (p * (log_full_teacher - log_full_student)).sum(dim=-1)
-
-
-def jsd_per_pos(
- log_full_student: torch.Tensor,
- log_full_teacher: torch.Tensor,
- beta: float = 0.5,
-) -> torch.Tensor:
- """Generalized JSD(beta).
-
- JSD_beta = beta * KL(S || M) + (1 - beta) * KL(T || M), M = beta S + (1-beta) T.
- beta = 0.5 -> standard JSD. GKD default is beta = 0.1 (closer to forward KL).
- """
- p_s = log_full_student.exp()
- p_t = log_full_teacher.exp()
- p_m = beta * p_s + (1.0 - beta) * p_t
- log_m = torch.clamp(p_m, min=1e-12).log()
- kl_sm = (p_s * (log_full_student - log_m)).sum(dim=-1)
- kl_tm = (p_t * (log_full_teacher - log_m)).sum(dim=-1)
- return beta * kl_sm + (1.0 - beta) * kl_tm
-
-
-def skew_kl_per_pos(
- log_full_student: torch.Tensor,
- log_full_teacher: torch.Tensor,
- alpha: float = 0.1,
-) -> torch.Tensor:
- """Skew KL from DistiLLM (arXiv 2402.03898):
- KL(P || (1 - alpha) P + alpha Q)
- Used to avoid log(0) when teacher has near-zero probability on a token
- the student wants to assign mass to. Default skew alpha = 0.1.
- """
- p_t = log_full_teacher.exp()
- p_s = log_full_student.exp()
- p_mix = (1.0 - alpha) * p_t + alpha * p_s
- log_mix = torch.clamp(p_mix, min=1e-12).log()
- return (p_t * (log_full_teacher - log_mix)).sum(dim=-1)
-
-
-# ---------------------------------------------------------------------------
-# Step-level helpers (V3 — SOD)
-# ---------------------------------------------------------------------------
-
-
-def step_mean_kl(
- per_token_kl: torch.Tensor,
- step_idx: torch.Tensor,
- K: int,
- mask: Optional[torch.Tensor] = None,
-) -> torch.Tensor:
- """Average per-token KL within each step.
-
- per_token_kl: (B, T) float
- step_idx: (B, T) long, in [0, K). Positions that are NOT in any step
- (tool obs, user, system, padding) MUST have step_idx = -1
- (we treat negatives as masked out).
- mask: (B, T) {0, 1}; default = (step_idx >= 0).float()
- Returns (B, K) with zero-count steps set to 0.
- """
- B, T = per_token_kl.shape
- if mask is None:
- mask = (step_idx >= 0).float()
- safe_idx = step_idx.clamp(min=0).long()
- weighted = per_token_kl * mask
- one = mask
- # (B, K) accumulators
- sum_kl = torch.zeros((B, K), dtype=per_token_kl.dtype, device=per_token_kl.device)
- cnt = torch.zeros((B, K), dtype=per_token_kl.dtype, device=per_token_kl.device)
- sum_kl.scatter_add_(1, safe_idx, weighted)
- cnt.scatter_add_(1, safe_idx, one)
- return sum_kl / cnt.clamp(min=1.0)
-
-
-def sod_adaptive_weights(
- step_kl: torch.Tensor, delta: float = 0.5, eps: float = 1e-3
-) -> torch.Tensor:
- """SOD step weights — Eq. (5) of arXiv 2605.07725 (translated to 0-index).
-
- step_kl: (B, K) — mean reverse-KL inside step k (output of step_mean_kl).
- Returns (B, K) with:
- w_0 = 1
- w_k = min( prod_{u=0}^{k-1} (d_u + eps) / (d_{u+1} + eps), 1 + delta )
- = min( (d_0 + eps) / (d_k + eps), 1 + delta ) (telescoping)
-
- Intuition:
- - d decreasing across steps (model recovers) -> w grows -> up-weight.
- - d exploding (error cascade) -> w shrinks -> dampened.
- - Capped at 1+delta to avoid blow-up on warm-up.
- """
- B, K = step_kl.shape
- if K == 0:
- return step_kl.new_zeros(B, 0)
- safe = step_kl + eps
- # ratio[:, k>=1] = d_{k-1}/d_k; ratio[:, 0] = 1 (so log_ratio[0]=0, w_0 = 1).
- ratio = torch.ones_like(safe)
- if K >= 2:
- ratio[:, 1:] = safe[:, :-1] / safe[:, 1:]
- log_ratio = ratio.log()
- # cum_log[:, k] = sum_{u<=k} log_ratio_u = log(d_0/d_k) for k>=1, 0 for k=0.
- # i.e. cum_log is already log(w) in 0-index — no extra shift needed.
- log_w = log_ratio.cumsum(dim=1)
- w = log_w.exp()
- return w.clamp(max=1.0 + delta)
-
-
-# ---------------------------------------------------------------------------
-# Segment helpers (V4 — SAD)
-# ---------------------------------------------------------------------------
-
-SEGMENT_NAMES = ("reason", "action", "tool_arg", "final_answer", "tool_result")
-SEG_REASON, SEG_ACTION, SEG_TOOL_ARG, SEG_FINAL, SEG_TOOL_RESULT = range(5)
-
-DEFAULT_SEG_WEIGHTS = {
- "reason": 1.0,
- "action": 1.5,
- "tool_arg": 1.2,
- "final_answer": 2.0,
- "tool_result": 0.0,
-}
-
-
-@dataclass(frozen=True)
-class SpecialIds:
- """Resolved IDs of nanoai agentic special tokens. Construct via from_tokenizer()."""
- tool_call_start: int
- tool_call_end: int
- tool_arg: int
- tool_val: int
- tool_result_start: int
- tool_result_end: int
- assistant_start: int
- assistant_end: int
- user_start: int
- user_end: int
-
- @classmethod
- def from_tokenizer(cls, tokenizer) -> "SpecialIds":
- enc = tokenizer.encode_special
- return cls(
- tool_call_start=enc("<|tool_call_start|>"),
- tool_call_end=enc("<|tool_call_end|>"),
- tool_arg=enc("<|tool_arg|>"),
- tool_val=enc("<|tool_val|>"),
- tool_result_start=enc("<|tool_result_start|>"),
- tool_result_end=enc("<|tool_result_end|>"),
- assistant_start=enc("<|assistant_start|>"),
- assistant_end=enc("<|assistant_end|>"),
- user_start=enc("<|user_start|>"),
- user_end=enc("<|user_end|>"),
- )
-
-
-def identify_segments(token_ids: torch.Tensor, special: SpecialIds) -> torch.Tensor:
- """State-machine label of each position as one of SEGMENT_NAMES (0..4).
-
- token_ids: (B, T) long.
- Returns int8 tensor of same shape, values in [0, 4]. We DO NOT distinguish
- `final_answer` from `reason` purely from special tokens — for now we label
- everything between the last `<|tool_call_end|>` (or user end) and the next
- `<|assistant_end|>` as `final_answer` ONLY when there's no further tool
- call before assistant_end. The caller can override the labeling on the
- final span if a richer signal is available.
-
- The mapping for purposes of OPD:
- - inside `<|tool_call_start|>` ... `<|tool_arg|>` -> action
- - inside `<|tool_arg|>` ... `<|tool_call_end|>` -> tool_arg
- - inside `<|tool_result_start|>` ... `<|tool_result_end|>` -> tool_result
- - inside `<|user_start|>` ... `<|user_end|>` -> tool_result (mask)
- - else within an assistant turn -> reason
- - the last reason span ending at `<|assistant_end|>`
- with no subsequent tool_call in the same turn -> final_answer
-
- Special tokens themselves get the label of the *region* they OPEN; closing
- tokens get the label of the region they CLOSE. This is intentional so the
- mask logic treats them uniformly with their span.
- """
- B, T = token_ids.shape
- out = torch.full_like(token_ids, SEG_REASON, dtype=torch.int8)
-
- # We'll first compute a coarse per-position label and then post-process the
- # final answer label per row.
- for b in range(B):
- row = token_ids[b].tolist()
- in_tool_call = False
- in_tool_arg = False
- in_tool_result = False
- in_user = False
- # Track the last assistant_end so we can locate "final answer" spans.
- for t, tok in enumerate(row):
- if tok == special.user_start:
- in_user = True
- if tok == special.user_end:
- in_user = False
- out[b, t] = SEG_TOOL_RESULT # user_end token itself = masked
- continue
- if in_user:
- out[b, t] = SEG_TOOL_RESULT
- continue
-
- if tok == special.tool_result_start:
- in_tool_result = True
- out[b, t] = SEG_TOOL_RESULT
- continue
- if tok == special.tool_result_end:
- in_tool_result = False
- out[b, t] = SEG_TOOL_RESULT
- continue
- if in_tool_result:
- out[b, t] = SEG_TOOL_RESULT
- continue
-
- if tok == special.tool_call_start:
- in_tool_call = True
- in_tool_arg = False
- out[b, t] = SEG_ACTION
- continue
- if tok == special.tool_arg:
- in_tool_arg = True
- out[b, t] = SEG_TOOL_ARG
- continue
- if tok == special.tool_call_end:
- in_tool_call = False
- in_tool_arg = False
- out[b, t] = SEG_TOOL_ARG
- continue
- if in_tool_arg:
- out[b, t] = SEG_TOOL_ARG
- continue
- if in_tool_call:
- out[b, t] = SEG_ACTION
- continue
-
- out[b, t] = SEG_REASON
-
- # Post-process: relabel the final reason span (between last tool_call_end
- # or assistant_start and the row's last assistant_end) as final_answer.
- row_t = out[b]
- # Find the last assistant_end position
- try:
- last_end = (token_ids[b] == special.assistant_end).nonzero(as_tuple=False)[-1].item()
- except IndexError:
- continue
- # Walk backward from last_end - 1 to find the start of the trailing reason span
- start = last_end
- for t in range(last_end - 1, -1, -1):
- if row_t[t].item() != SEG_REASON:
- start = t + 1
- break
- else:
- start = 0
- out[b, start:last_end + 1] = SEG_FINAL
-
- return out
-
-
-def seg_weights_to_tensor(
- weights: dict, device, dtype=torch.float32
-) -> torch.Tensor:
- """Build a (5,) tensor in the canonical SEGMENT order."""
- return torch.tensor(
- [
- weights.get("reason", 1.0),
- weights.get("action", 1.0),
- weights.get("tool_arg", 1.0),
- weights.get("final_answer", 1.0),
- weights.get("tool_result", 0.0),
- ],
- device=device,
- dtype=dtype,
- )
-
-
-# ---------------------------------------------------------------------------
-# Loss reductions
-# ---------------------------------------------------------------------------
-
-
-def reduce_per_token(per_token: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
- """Mean of per-token quantity over masked positions. Returns scalar."""
- m = mask.float()
- return (per_token * m).sum() / m.sum().clamp(min=1.0)
-
-
-def reduce_per_step(
- per_token: torch.Tensor,
- step_idx: torch.Tensor,
- step_weight: torch.Tensor,
- mask: Optional[torch.Tensor] = None,
-) -> torch.Tensor:
- """Step-weighted reduction.
-
- per_token: (B, T) — e.g. reverse_kl_per_token output
- step_idx: (B, T) long; -1 for non-student positions
- step_weight: (B, K) — e.g. SOD weights
- mask: (B, T) defaults to (step_idx >= 0)
-
- Returns scalar. Equivalent to:
- sum_{b,t} w[b, step_idx[b,t]] * per_token[b,t]
- ---------------------------------------------
- sum_{b,t} w[b, step_idx[b,t]]
- """
- if mask is None:
- mask = (step_idx >= 0).float()
- safe_idx = step_idx.clamp(min=0).long()
- w_pt = step_weight.gather(1, safe_idx) * mask
- return (per_token * w_pt).sum() / w_pt.sum().clamp(min=1.0)
-
-
-def reduce_per_sequence(per_token: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
- """Mean of length-normalized per-sequence sums. Returns scalar.
-
- This gives every example equal weight regardless of length, which is what
- you want for DPO-style sequence-KL comparison. (Cf. agentic_dpo.sequence_logp.)
- """
- m = mask.float()
- per_seq = (per_token * m).sum(dim=-1) / m.sum(dim=-1).clamp(min=1.0)
- return per_seq.mean()
-
-
-# ---------------------------------------------------------------------------
-# Diagnostic helpers (no grad)
-# ---------------------------------------------------------------------------
-
-
-@torch.no_grad()
-def top1_overlap(
- log_full_student: torch.Tensor,
- log_full_teacher: torch.Tensor,
- mask: torch.Tensor,
-) -> float:
- """Fraction of supervised positions where argmax(student) == argmax(teacher).
-
- Per Rethinking OPD: healthy convergence raises this from ~72% to ~91% as
- training proceeds.
- """
- s = log_full_student.argmax(dim=-1)
- t = log_full_teacher.argmax(dim=-1)
- agree = (s == t).float() * mask.float()
- denom = mask.float().sum().clamp(min=1.0)
- return (agree.sum() / denom).item()
-
-
-@torch.no_grad()
-def student_entropy(
- log_full_student: torch.Tensor, mask: torch.Tensor
-) -> float:
- """Mean per-position entropy of the student over masked positions (nats)."""
- p = log_full_student.exp()
- H = -(p * log_full_student).sum(dim=-1)
- return reduce_per_token(H, mask).item()
diff --git a/vendor/nanoai/agentic/prefix_merging.py b/vendor/nanoai/agentic/prefix_merging.py
deleted file mode 100644
index d10ed09a1c30a53b49d251fee33df1a44efa2ba9..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/prefix_merging.py
+++ /dev/null
@@ -1,528 +0,0 @@
-"""Token-faithful prefix-merging trajectory reconstruction (Polar mechanism).
-
-Port of NVIDIA **Polar** (arxiv 2605.24220, "Agentic RL on Any Harness at
-Scale"), `trajectory/builder/prefix_merging.py`, adapted to nanoai. Polar's
-core contribution is reconstructing **one** token-level training trace out of
-the **many** independent LLM completions a multi-turn agent emits, without
-introducing tokenization drift:
-
- z = p1 ∥ a1 ∥ u1 ∥ a2 ∥ u2 ∥ ⋯ ∥ aK
-
-where ``p1`` is the initial prompt, ``a_m`` are the **raw sampled** assistant
-response token ids (real behavior-policy tokens, never decode→re-encode), and
-``u_m`` are the **canonical** interstitials (tool results / chat-template glue)
-taken from the next completion's prompt tail. The loss mask is **1 on a_m, 0 on
-u_m** — "every trainable token matches the behavior policy; non-generated
-tokens are masked out."
-
-Why nanoai wants this even though its own ``rollout_one`` is a white box:
-
- * **Black-box harnesses.** ``build_prefix_merged`` reconstructs a trace from N
- independent per-turn completions (prompt_ids / response_ids / logprobs),
- e.g. captured from an *external* harness (Codex / Claude Code / qwen-code)
- proxied through a gateway. nanoai currently can only train on its own
- ``rollout_one`` loop; this generalizes it to any harness.
- * **An explicit, testable fidelity invariant.** ``assert_behavior_faithful``
- asserts every supervised (mask=1) token equals a token actually sampled —
- the guarantee nanoai's flat-``tokens`` + ``turn_spans`` packing currently
- leaves implicit (and which the per-token-advantage off-by-one bugs in
- CLAUDE.md silently violated).
- * **Behavior logprobs for TIS.** Polar stores the behavior-policy logprob per
- sampled token (the paper runs with "TIS: Enabled"). nanoai discards them
- and recomputes at train time — only valid fully on-policy. ``Trace`` carries
- an optional ``response_logprobs`` slot so async/off-policy correction has a
- place to live.
-
-``trace_from_rollout`` re-expresses nanoai's own ``MultiTurnRollout`` (flat
-tokens + turn_spans) in the same ``Trace`` contract, so the white-box packing
-is provably a special case of the general merge (see tests/test_prefix_merging).
-
-Pure Python, no torch — runs anywhere, CPU-only.
-"""
-from __future__ import annotations
-
-from dataclasses import dataclass, field
-from typing import Optional
-
-
-# ---------------------------------------------------------------------------
-# Data contracts
-# ---------------------------------------------------------------------------
-
-
-@dataclass
-class Completion:
- """One per-turn LLM completion record (the black-box unit).
-
- Mirrors what a gateway stores per intercepted model call: the server-side
- tokenized ``prompt_ids``, the **raw sampled** ``response_ids``, and the
- per-sampled-token ``logprobs`` (optional). ``finish_reason`` lets the
- builder auto-detect the end-of-turn token.
- """
- prompt_ids: list[int]
- response_ids: list[int]
- logprobs: Optional[list[float]] = None
- finish_reason: Optional[str] = None
- metadata: dict = field(default_factory=dict)
-
-
-@dataclass
-class Trace:
- """One reconstructed, token-faithful training trace.
-
- ``response_ids`` is the merged stream ``a1 ∥ u1 ∥ a2 ∥ ⋯``; ``loss_mask`` is
- 1 on sampled assistant tokens and 0 on canonical interstitials. ``loss_mask``
- and ``response_logprobs`` (when present) are parallel to ``response_ids``.
- """
- prompt_ids: list[int]
- response_ids: list[int]
- loss_mask: list[int]
- response_logprobs: Optional[list[float]] = None
- finish_reason: Optional[str] = None
- metadata: dict = field(default_factory=dict)
-
- def __post_init__(self) -> None:
- if len(self.loss_mask) != len(self.response_ids):
- raise ValueError(
- f"loss_mask len {len(self.loss_mask)} != response_ids len "
- f"{len(self.response_ids)}"
- )
- if self.response_logprobs is not None and len(self.response_logprobs) != len(
- self.response_ids
- ):
- raise ValueError("response_logprobs len must match response_ids len")
- for m in self.loss_mask:
- if m not in (0, 1):
- raise ValueError("loss_mask values must be 0 or 1")
-
- @property
- def n_trained(self) -> int:
- """Number of supervised (mask=1) tokens."""
- return sum(self.loss_mask)
-
- def trained_tokens(self) -> list[int]:
- """The masked-in subsequence (the behavior-policy sampled tokens)."""
- return [t for t, m in zip(self.response_ids, self.loss_mask) if m == 1]
-
- def interstitial_tokens(self) -> list[int]:
- """The masked-out subsequence (canonical interstitials)."""
- return [t for t, m in zip(self.response_ids, self.loss_mask) if m == 0]
-
-
-# Finish reasons where the model emitted the natural end-of-turn token itself.
-NATURAL_STOP_REASONS = frozenset({"stop", "tool_calls", "stop_sequence", "assistant_end"})
-
-
-# ---------------------------------------------------------------------------
-# Black-box reconstruction: merge N completions -> token-faithful Trace(s)
-# ---------------------------------------------------------------------------
-
-
-def build_prefix_merged(
- completions: list[Completion],
- *,
- eot_id: Optional[int] = None,
-) -> tuple[list[Trace], dict]:
- """Merge per-turn completions into token-faithful traces (Polar algorithm).
-
- Returns ``(traces, stats)``. Completions that append-extend one append-only
- conversation collapse into a single trace; parallel/interleaved agents (each
- with a distinct prompt prefix) split into separate traces; a rewritten-history
- prefix break truncates the chain (remaining completions start fresh).
-
- ``eot_id``: explicit end-of-turn id (nanoai: ``<|assistant_end|>``). When
- None, auto-detect from the last token of the first natural-stop completion.
- """
- if not completions:
- return [], {"chains_total": 0, "completions_total": 0, "completions_merged": 0}
-
- # --- Stage 1: group completions into chains by the token-prefix relation ---
- # A completion joins the chain whose last prompt is a strict token-prefix of
- # it (p_{i+1}[:|p_i|] == p_i). Compares only server-tokenized PROMPTS, so BPE
- # re-tokenization of sampled responses never enters the routing decision.
- chains: list[list[Completion]] = []
- chain_tips: list[list[int]] = [] # last completion's prompt_ids per chain
- for c in completions:
- idx = _find_extendable_chain(c.prompt_ids, chain_tips)
- if idx is None:
- idx = len(chains)
- chains.append([])
- chain_tips.append([])
- chains[idx].append(c)
- chain_tips[idx] = list(c.prompt_ids)
-
- # --- Stage 2: finalize each chain into a merged, mask-tagged stream ---
- stats = {
- "chains_total": len(chains),
- "chains_reconstructed_full": 0,
- "chains_reconstructed_truncated": 0,
- "completions_total": len(completions),
- "completions_merged": 0,
- }
- traces = [_finalize_chain(chain, eot_id, stats) for chain in chains]
- return traces, stats
-
-
-def _find_extendable_chain(
- prompt_ids: list[int], chain_tips: list[list[int]]
-) -> Optional[int]:
- """Return the open chain this completion append-extends, else None.
-
- A completion continues a chain iff its prompt begins with that chain's last
- prompt (tip is a token-prefix of prompt_ids). On overlap the longest matching
- tip wins (the most advanced chain).
- """
- best_idx: Optional[int] = None
- best_len = -1
- for idx, tip in enumerate(chain_tips):
- n = len(tip)
- if n > best_len and 0 < n <= len(prompt_ids) and prompt_ids[:n] == tip:
- best_idx, best_len = idx, n
- return best_idx
-
-
-def _resolve_eot_id(chain: list[Completion], configured: Optional[int]) -> Optional[int]:
- if configured is not None:
- return configured
- for c in chain:
- if c.finish_reason in NATURAL_STOP_REASONS and c.response_ids:
- return c.response_ids[-1]
- return None
-
-
-def _slice_interstitial(
- canonical_tail: list[int],
- prev_raw_response: list[int],
- eot_id: Optional[int],
-) -> Optional[list[int]]:
- """Extract the canonical interstitial from C_{i+1}'s prompt tail.
-
- ``canonical_tail`` = canonical tokens for [prev assistant body + harness-
- inserted messages + generation-prompt glue]. The first ``eot_id`` marks the
- end of the prev assistant body; everything after is interstitial. If the prev
- raw response already ended with ``eot_id`` (natural stop), skip the dup;
- otherwise (truncation) keep it so the stream still closes the turn. Returns
- None if eot_id is unknown or absent — caller treats that as a break.
- """
- if eot_id is None:
- return None
- try:
- k = canonical_tail.index(eot_id)
- except ValueError:
- return None
- if prev_raw_response and prev_raw_response[-1] == eot_id:
- return canonical_tail[k + 1 :]
- return canonical_tail[k:]
-
-
-def _finalize_chain(
- chain: list[Completion], configured_eot: Optional[int], stats: dict
-) -> Trace:
- eot_id = _resolve_eot_id(chain, configured_eot)
- first = chain[0]
-
- prompt_ids = list(first.prompt_ids)
- response_ids: list[int] = []
- loss_mask: list[int] = []
- logprob_slots: list[Optional[float]] = []
-
- prev_prompt_ids = list(first.prompt_ids)
- prev_raw_response = list(first.response_ids)
-
- _append_response(first, response_ids, loss_mask, logprob_slots)
- kept = 1
-
- for i in range(1, len(chain)):
- c = chain[i]
- ci_prompt = list(c.prompt_ids)
- # Canonical-vs-canonical prefix check: both are server tokenizations of
- # the same message prefix — matches unless the harness rewrote history.
- if len(ci_prompt) < len(prev_prompt_ids) or ci_prompt[: len(prev_prompt_ids)] != prev_prompt_ids:
- break
- canonical_tail = ci_prompt[len(prev_prompt_ids) :]
- interstitial = _slice_interstitial(canonical_tail, prev_raw_response, eot_id)
- if interstitial is None:
- break
- if interstitial:
- response_ids.extend(interstitial)
- loss_mask.extend([0] * len(interstitial))
- logprob_slots.extend([None] * len(interstitial))
- _append_response(c, response_ids, loss_mask, logprob_slots)
- prev_prompt_ids = ci_prompt
- prev_raw_response = list(c.response_ids)
- kept += 1
-
- stats["completions_merged"] += kept
- if kept == len(chain):
- stats["chains_reconstructed_full"] += 1
- else:
- stats["chains_reconstructed_truncated"] += 1
-
- logprobs = _finalize_logprobs(logprob_slots)
- return Trace(
- prompt_ids=prompt_ids,
- response_ids=response_ids,
- loss_mask=loss_mask,
- response_logprobs=logprobs,
- finish_reason=chain[kept - 1].finish_reason,
- metadata={"builder": "prefix_merging", "completions_merged": kept,
- "completions_total": len(chain)},
- )
-
-
-def _append_response(
- c: Completion,
- response_ids: list[int],
- loss_mask: list[int],
- logprob_slots: list[Optional[float]],
-) -> None:
- """Append a completion's RAW sampled response_ids with mask=1."""
- ids = list(c.response_ids)
- response_ids.extend(ids)
- loss_mask.extend([1] * len(ids))
- lp = c.logprobs or []
- for pos in range(len(ids)):
- v = lp[pos] if pos < len(lp) else None
- logprob_slots.append(float(v) if isinstance(v, (int, float)) else None)
-
-
-def _finalize_logprobs(slots: list[Optional[float]]) -> Optional[list[float]]:
- # Interstitial slots get 0.0; loss_mask=0 makes the trainer ignore them.
- if not any(s is not None for s in slots):
- return None
- return [s if s is not None else 0.0 for s in slots]
-
-
-# ---------------------------------------------------------------------------
-# White-box adapter: nanoai MultiTurnRollout -> the same Trace contract
-# ---------------------------------------------------------------------------
-
-
-def trace_from_rollout(rollout, *, response_logprobs: Optional[list[float]] = None) -> Trace:
- """Re-express a nanoai ``MultiTurnRollout`` as a token-faithful ``Trace``.
-
- nanoai already builds one flat ``tokens`` sequence in place (raw sampled
- assistant bodies + re-encoded tool-result interstitials) and marks each
- assistant turn with a ``TurnSpan``. That IS Polar's merged trace; here we
- surface it explicitly:
-
- * ``response_ids`` = tokens[prompt_len:]
- * ``loss_mask`` = 1 on positions inside any turn_span, 0 elsewhere
- (interstitials = tool results + chat glue)
-
- ``response_logprobs`` (optional) lets a caller attach behavior logprobs if
- the engine was modified to surface them (today nanoai discards them).
- """
- prompt_len = rollout.prompt_len
- tokens = rollout.tokens
- response_ids = list(tokens[prompt_len:])
- n = len(response_ids)
- # Carry behavior logprobs (for TIS) when the rollout captured them.
- if response_logprobs is None:
- tlp = getattr(rollout, "token_logprobs", None)
- if tlp is not None:
- response_logprobs = list(tlp[prompt_len:])
- loss_mask = [0] * n
- for ts in rollout.turn_spans:
- lo = max(0, ts.start_idx - prompt_len)
- hi = min(n, ts.end_idx - prompt_len)
- for p in range(lo, hi):
- loss_mask[p] = 1
- return Trace(
- prompt_ids=list(tokens[:prompt_len]),
- response_ids=response_ids,
- loss_mask=loss_mask,
- response_logprobs=response_logprobs,
- finish_reason="assistant_end",
- metadata={"builder": "from_rollout", "task_id": getattr(rollout, "task_id", None),
- "n_turns": getattr(rollout, "n_turns", None)},
- )
-
-
-# ---------------------------------------------------------------------------
-# Fidelity verifier — the invariant nanoai currently leaves implicit
-# ---------------------------------------------------------------------------
-
-
-def assert_behavior_faithful(trace: Trace, sampled_turns: list[list[int]]) -> None:
- """Assert every supervised (mask=1) token equals a token actually sampled.
-
- ``sampled_turns`` is the in-order list of raw sampled assistant token
- sequences (a1, a2, ...). The masked-in subsequence of the trace must equal
- their concatenation exactly. Catches off-by-one mask/advantage bugs that
- otherwise silently train on (or skip) the wrong tokens.
- """
- expected = [t for turn in sampled_turns for t in turn]
- got = trace.trained_tokens()
- if got != expected:
- raise AssertionError(
- "behavior-policy fidelity violated: masked-in tokens != sampled tokens\n"
- f" expected {len(expected)} sampled tokens, got {len(got)} trained\n"
- f" first divergence at "
- f"{next((i for i, (a, b) in enumerate(zip(expected, got)) if a != b), 'len-mismatch')}"
- )
-
-
-# ---------------------------------------------------------------------------
-# pack_for_grad gate — install the fidelity invariant in the trainer hot path
-# ---------------------------------------------------------------------------
-
-
-def _as_row(targets_row):
- """Accept a torch tensor row or a plain sequence; return a python list."""
- if hasattr(targets_row, "tolist"):
- return targets_row.tolist()
- return list(targets_row)
-
-
-def verify_pack_for_grad(rollouts, targets, max_len: int, *, ignore_index: int = -1,
- strict: bool = True) -> dict:
- """Gate: assert a ``pack_for_grad`` targets batch is token-faithful.
-
- ``targets`` is the (B, max_len) next-token-prediction label tensor (or a list
- of rows): ``targets[g][t]`` is the label predicted from position ``t`` (so it
- supervises the token at absolute index ``t+1``), or ``ignore_index`` for
- unsupervised positions.
-
- The check is **independent** of pack_for_grad's own NTP-shift arithmetic: it
- derives the ground-truth supervised set from each rollout's ``turn_spans``
- via ``trace_from_rollout`` and requires, for every supervised position ``t``:
-
- * ``t+1`` lands on a **sampled** assistant token (loss_mask==1), not on a
- prompt / interstitial / out-of-range token, and
- * ``targets[g][t] == rollout.tokens[t+1]`` (the label is that very token).
-
- and conversely that every in-range sampled token is supervised by exactly one
- position. This catches the recurring off-by-one (advantage/target range shifted
- so the first asst token of each turn gets zero gradient — see CLAUDE.md).
-
- Returns a report dict. Raises AssertionError on any violation when ``strict``.
- """
- violations: list[str] = []
- n_supervised = 0
- n_sampled_in_range = 0
- for g, r in enumerate(rollouts):
- row = _as_row(targets[g])
- trace = trace_from_rollout(r)
- prompt_len = r.prompt_len
- tokens = r.tokens
- loss_mask = trace.loss_mask # 1 on sampled assistant tokens, indexed in response coords
-
- def is_sampled(absidx: int) -> bool:
- ri = absidx - prompt_len
- return 0 <= ri < len(loss_mask) and loss_mask[ri] == 1
-
- supervised_next = set()
- for t in range(min(max_len, len(row))):
- if int(row[t]) == ignore_index:
- continue
- n_supervised += 1
- nxt = t + 1
- if not is_sampled(nxt):
- violations.append(
- f"rollout {g}: supervised pos {t} predicts non-sampled token at {nxt}")
- continue
- if nxt < len(tokens) and int(row[t]) != tokens[nxt]:
- violations.append(
- f"rollout {g}: label[{t}]={int(row[t])} != tokens[{nxt}]={tokens[nxt]}")
- supervised_next.add(nxt)
-
- # Every in-range sampled token (except the very first, which no position
- # can predict) must be supervised exactly once.
- for ri, m in enumerate(loss_mask):
- if m != 1:
- continue
- absidx = ri + prompt_len
- if absidx == 0 or absidx >= max_len:
- continue
- n_sampled_in_range += 1
- if absidx not in supervised_next:
- violations.append(
- f"rollout {g}: sampled token at {absidx} is not supervised (zero gradient)")
-
- report = {
- "faithful": not violations,
- "n_rollouts": len(rollouts),
- "n_supervised": n_supervised,
- "n_sampled_in_range": n_sampled_in_range,
- "n_violations": len(violations),
- "violations": violations[:10],
- }
- if strict and violations:
- raise AssertionError(
- "pack_for_grad fidelity gate FAILED ("
- f"{len(violations)} violations):\n " + "\n ".join(violations[:10]))
- return report
-
-
-# ---------------------------------------------------------------------------
-# TIS — truncated importance weight for behavior->training policy drift
-# ---------------------------------------------------------------------------
-
-
-def tis_weight(log_now, log_behav, cap: float = 2.0):
- """Truncated importance-sampling weight `min(exp(log_now - log_behav), cap)`.
-
- The per-token TIS factor (Polar, "TIS: Enabled") that corrects for the rollout
- (behavior) policy differing from the policy being optimized — async / stale /
- multi-update-per-batch RL. Use it **detached** as a multiplicative reweight on
- the policy-gradient objective: it reweights samples, it is not a gradient path.
-
- `log_behav` must be the behavior policy's UNTEMPERED log pi(a) at rollout time
- (see Engine.generate(return_logprobs=True)), so it matches the untempered
- training-time `log_now`. On-policy (log_now == log_behav) -> 1.0; `cap`
- truncates the upper tail to bound variance. Accepts torch tensors (elementwise).
- """
- import torch
- return torch.clamp(torch.exp(log_now - log_behav), max=cap)
-
-
-# ---------------------------------------------------------------------------
-# Self-test (no torch, no real model/tokenizer) — analog of Polar's repro.
-# ---------------------------------------------------------------------------
-
-if __name__ == "__main__":
- # nanoai special-token ids (synthetic, matching grpo_multi_rollout self-test)
- BOS, US, UE, AS, AE = 1000, 1001, 1002, 1003, 1004
- TCS, TCE, TARG, TVAL = 1005, 1006, 1007, 1008
- TRS, TRE = 1009, 1010
- EOT = AE # nanoai natural end-of-turn = <|assistant_end|>
-
- # Build a 3-turn append-only chain as black-box completions. Each turn's
- # response ends with <|assistant_end|>; the interstitial inserted before the
- # next turn is [<|tool_result_start|>, result.., <|tool_result_end|>, <|assistant_start|>].
- p1 = [BOS, US, 2001, 2002, UE, AS] # initial prompt
- a1 = [3001, TCS, 3002, TARG, 3003, TCE, AE] # assistant turn 1 (tool call)
- u1 = [TRS, 4001, 4002, TRE, AS] # tool result interstitial 1
- a2 = [3010, TCS, 3011, TCE, AE] # assistant turn 2
- u2 = [TRS, 4010, TRE, AS] # tool result interstitial 2
- a3 = [3020, 3021, AE] # assistant turn 3 (final answer)
-
- c1 = Completion(prompt_ids=p1, response_ids=a1, logprobs=[-0.1] * len(a1), finish_reason="stop")
- c2 = Completion(prompt_ids=p1 + a1 + u1, response_ids=a2, logprobs=[-0.2] * len(a2), finish_reason="stop")
- c3 = Completion(prompt_ids=p1 + a1 + u1 + a2 + u2, response_ids=a3, logprobs=[-0.3] * len(a3), finish_reason="stop")
-
- traces, stats = build_prefix_merged([c1, c2, c3], eot_id=EOT)
- assert len(traces) == 1, f"3 completions should merge to 1 trace, got {len(traces)}"
- t = traces[0]
- assert t.prompt_ids == p1
- assert t.response_ids == a1 + u1 + a2 + u2 + a3, "merged stream wrong"
- assert t.loss_mask == [1] * len(a1) + [0] * len(u1) + [1] * len(a2) + [0] * len(u2) + [1] * len(a3)
- assert_behavior_faithful(t, [a1, a2, a3])
- assert t.trained_tokens() == a1 + a2 + a3
- assert t.interstitial_tokens() == u1 + u2
- assert stats["chains_reconstructed_full"] == 1 and stats["completions_merged"] == 3
- print(f"black-box merge: 3 completions -> 1 trace, "
- f"{t.n_trained}/{len(t.response_ids)} trained, fidelity OK")
-
- # Parallel agents (distinct prompt roots) must not merge across each other.
- pa = [BOS, US, 5001, UE, AS]
- pb = [BOS, US, 6001, UE, AS]
- ca1 = Completion(pa, [5100, AE], finish_reason="stop")
- cb1 = Completion(pb, [6100, AE], finish_reason="stop")
- ca2 = Completion(pa + [5100, AE] + [TRS, 5200, TRE, AS], [5101, AE], finish_reason="stop")
- par, _ = build_prefix_merged([ca1, cb1, ca2], eot_id=EOT)
- assert len(par) == 2, f"two parallel agents -> 2 traces, got {len(par)}"
- print("parallel agents: interleaved completions -> 2 separate traces")
-
- print("prefix_merging self-test OK")
diff --git a/vendor/nanoai/agentic/reward.py b/vendor/nanoai/agentic/reward.py
deleted file mode 100644
index a44183fff575eae2b94fd464705091abaddaf58b..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/reward.py
+++ /dev/null
@@ -1,100 +0,0 @@
-"""Reward function for agentic GRPO.
-
-Score in [0, 4] based on whether the assistant's generated text:
- +1.0 emits exactly one well-formed `<|tool_call_start|>...<|tool_call_end|>` block
- +1.0 the tool name inside that block is one of {Read, Edit, Grep, Bash}
- +1.0 the args inside the block include all required keys for that tool
- +1.0 the response terminates cleanly with `<|assistant_end|>`
-
-The reward is computed from the *decoded text* (not token ids), so it works
-with any tokenizer that round-trips the special tokens correctly.
-
-Returns (score, breakdown_dict) so callers can log per-component signal.
-"""
-
-from __future__ import annotations
-
-TOOL_REQUIRED_ARGS = {
- "Read": {"file_path"},
- "Edit": {"file_path", "new_string"},
- "Grep": {"pattern"},
- "Bash": {"command"},
-}
-
-
-def parse_tool_call(body: str):
- """Parse the body between <|tool_call_start|> and <|tool_call_end|>.
-
- Format: <|tool_arg|> k1 <|tool_val|> v1 <|tool_arg|> k2 <|tool_val|> v2 ...
- """
- parts = body.split("<|tool_arg|>")
- tool_name = parts[0].strip()
- args = {}
- for part in parts[1:]:
- if "<|tool_val|>" in part:
- k, v = part.split("<|tool_val|>", 1)
- args[k.strip()] = v.strip()
- return tool_name, args
-
-
-def agentic_reward(decoded_text: str) -> tuple[float, dict]:
- """Score generated assistant text. Higher is better. Range [0, 4].
-
- `decoded_text` is the assistant's output starting AFTER `<|assistant_start|>`
- (the prompt prefix should NOT be included). It may or may not contain the
- closing `<|assistant_end|>`.
- """
- breakdown = {
- "has_tool_block": 0.0,
- "valid_tool_name": 0.0,
- "valid_args": 0.0,
- "clean_terminate": 0.0,
- }
-
- # 1. Tool block present and well-bracketed (exactly one open + one close)
- n_open = decoded_text.count("<|tool_call_start|>")
- n_close = decoded_text.count("<|tool_call_end|>")
- if n_open == 1 and n_close == 1:
- idx_open = decoded_text.index("<|tool_call_start|>") + len("<|tool_call_start|>")
- idx_close = decoded_text.index("<|tool_call_end|>")
- if idx_close > idx_open:
- breakdown["has_tool_block"] = 1.0
- body = decoded_text[idx_open:idx_close]
- tool_name, args = parse_tool_call(body)
-
- # 2. Tool name in known set
- if tool_name in TOOL_REQUIRED_ARGS:
- breakdown["valid_tool_name"] = 1.0
- # 3. Required args present
- required = TOOL_REQUIRED_ARGS[tool_name]
- if required.issubset(args.keys()):
- breakdown["valid_args"] = 1.0
-
- # 4. Clean termination
- if "<|assistant_end|>" in decoded_text:
- breakdown["clean_terminate"] = 1.0
-
- score = sum(breakdown.values())
- return score, breakdown
-
-
-# A small set of agentic prompts used as the GRPO training distribution.
-# Mirrors the agent_eval prompts plus variations, kept short to fit in seq_len=2048.
-TRAIN_PROMPTS = [
- "show me the contents of README.md",
- "find all uses of the word 'tokenizer' under scripts/",
- "list the files in the current directory",
- "print the working directory",
- "create a file hello.py that prints 'hi'",
- "where is the GPT class defined?",
- "how many python files are in the data/ directory?",
- "what does the BPE tokenizer do?",
- "search for 'import torch' in the nanoai module",
- "open the configs.py file",
- "make a new file called notes.txt with content 'TODO: review'",
- "find all .json files",
- "show me the first 20 lines of base_train.py",
- "replace 'foo' with 'bar' in test.py",
- "grep for TODO comments under scripts",
- "run 'pwd' and show the result",
-]
diff --git a/vendor/nanoai/agentic/reward_capbench.py b/vendor/nanoai/agentic/reward_capbench.py
deleted file mode 100644
index 671f42cd39724e045edef0e81154bf10c23adfaa..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/reward_capbench.py
+++ /dev/null
@@ -1,230 +0,0 @@
-"""Phase 5 RL reward — capbench LLM-judge wrapper.
-
-Computes reward for one diag-chain rollout trace against a corpus_v2 scenario
-truth, using the same LLM judge as `scripts.eval_llm_judge_rc_hit`. Designed
-to plug into DAPO/GRPO advantage computation; no gradient tracking, pure
-scalar reward.
-
-Reward formulation (Phase 5 design § B.1):
- +1.0 if LLM judge says hit (committed root cause == truth)
- -0.2 if no_commit / inconclusive / refuse
- 0.0 if committed but judge says not hit (committed_wrong)
- 0.0 if judge call failed (graceful degradation; logged)
-
-This sparse signal is the "north star" — exactly matches eval metric. If
-gradient sparsity hurts RL convergence, switch to `reward_shaped()` in
-phase8_v5_rl_design.md § B.2.
-
-Usage:
- from nanoai.agentic.reward_capbench import compute_capbench_reward
- r = compute_capbench_reward(trace, scenario,
- endpoint="http://127.0.0.1:4210",
- api_key="sk-capagent",
- model="qwen36-27b")
- # r.scalar: float (the reward)
- # r.bucket: str ("hit" / "wrong" / "no_commit" / "judge_failed")
- # r.judge_rationale: str (LLM judge's reasoning)
-"""
-from __future__ import annotations
-
-import sys
-from dataclasses import dataclass
-from pathlib import Path
-from typing import Optional
-
-
-# Default reward shape constants (mirror phase5_rl_design.md § B.1)
-REWARD_HIT: float = 1.0
-REWARD_WRONG: float = 0.0
-REWARD_NO_COMMIT: float = -0.2
-REWARD_JUDGE_FAIL: float = 0.0
-
-
-@dataclass
-class CapbenchReward:
- """Result of capbench LLM-judge reward computation."""
- scalar: float # the reward (passed to advantage computation)
- bucket: str # "hit" / "wrong" / "no_commit" / "judge_failed"
- judge_rationale: str # human-readable explanation
- judge_hit: Optional[bool] # None if judge failed; True/False otherwise
- has_commit: bool # did the model commit at all?
-
-
-def _extract_final(trace: dict) -> tuple[str, str]:
- """Mirror of scripts.eval_llm_judge_rc_hit._extract_final.
-
- Returns (final_reason, dominant_hypothesis_text). Both can be empty
- strings if the model never committed.
- """
- fs = trace.get("final_stop") or {}
- final_reason = fs.get("reason", "") if isinstance(fs, dict) else ""
- dom_hyp = ""
- for step in reversed(trace.get("steps") or []):
- if step.get("task") == "stop_controller":
- inp = step.get("input") or {}
- ch = inp.get("current_hypotheses") or []
- if isinstance(ch, list) and ch:
- best = max(
- ch,
- key=lambda h: (h.get("probability") if isinstance(h, dict) else 0) or 0,
- )
- if isinstance(best, dict):
- dom_hyp = best.get("hypothesis") or best.get("claim") or ""
- break
- return final_reason, dom_hyp
-
-
-def _is_no_commit(final_reason: str, trace: dict) -> bool:
- """Heuristic: was the rollout a no-commit (model gave up / continued past max_turns)?"""
- rat_low = final_reason.lower()
- if "no_final_commitment" in rat_low or "stop_inconclusive" in rat_low:
- return True
- if any(x in rat_low for x in ["inconclusive", "cannot distinguish", "unable to determine"]):
- return True
- # Check final_stop.decision if present
- fs = trace.get("final_stop") or {}
- if isinstance(fs, dict):
- decision = (fs.get("decision") or fs.get("stop_decision") or "").lower()
- if decision in ("stop_inconclusive", "continue"):
- return True
- if not final_reason.strip():
- return True
- return False
-
-
-def compute_capbench_reward(
- trace: dict,
- scenario_truth: dict,
- endpoint: str = "http://127.0.0.1:4210",
- api_key: str = "sk-capagent",
- model: str = "qwen36-27b",
- *,
- reward_hit: float = REWARD_HIT,
- reward_wrong: float = REWARD_WRONG,
- reward_no_commit: float = REWARD_NO_COMMIT,
- reward_judge_fail: float = REWARD_JUDGE_FAIL,
-) -> CapbenchReward:
- """Compute scalar reward for one rollout trace vs scenario truth.
-
- Args:
- trace: diag-chain trace dict (output of `diag_agent_loop_demo.run_loop`).
- Must have `final_stop` and `steps` fields.
- scenario_truth: dict with `root_cause_aliases`, `evidence_required_keywords`,
- `scenario_summary` (the corpus_v2 `truth` field).
- endpoint/api_key/model: capagent LLM judge config.
-
- Returns:
- CapbenchReward with scalar reward + classification bucket.
- """
- final_reason, dom_hyp = _extract_final(trace)
- has_commit = not _is_no_commit(final_reason, trace)
-
- if not has_commit:
- return CapbenchReward(
- scalar=reward_no_commit,
- bucket="no_commit",
- judge_rationale="model did not commit (no_final_commitment / inconclusive)",
- judge_hit=False,
- has_commit=False,
- )
-
- # Lazy-import to avoid circular dep on scripts package
- sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
- from scripts.eval_llm_judge_rc_hit import _judge_one
-
- hit, rationale = _judge_one(
- truth=scenario_truth,
- final_reason=final_reason,
- model_hypothesis=dom_hyp,
- endpoint=endpoint,
- api_key=api_key,
- model=model,
- )
- if hit is None:
- return CapbenchReward(
- scalar=reward_judge_fail,
- bucket="judge_failed",
- judge_rationale=f"judge call failed: {rationale}",
- judge_hit=None,
- has_commit=True,
- )
- return CapbenchReward(
- scalar=reward_hit if hit else reward_wrong,
- bucket="hit" if hit else "wrong",
- judge_rationale=rationale,
- judge_hit=hit,
- has_commit=True,
- )
-
-
-def compute_capbench_reward_batch(
- traces_and_scenarios: list[tuple[dict, dict]],
- endpoint: str = "http://127.0.0.1:4210",
- api_key: str = "sk-capagent",
- model: str = "qwen36-27b",
-) -> list[CapbenchReward]:
- """Sequential batch wrapper. NB: each LLM call is ~3-5s; for large batches
- consider concurrent.futures.ThreadPoolExecutor wrapper.
-
- Returns list aligned with input order.
- """
- return [
- compute_capbench_reward(t, s, endpoint=endpoint, api_key=api_key, model=model)
- for t, s in traces_and_scenarios
- ]
-
-
-# ---- Self-test (offline, no LLM call) ----------------------------------
-
-if __name__ == "__main__":
- # Build a fake trace with no_commit (should not call LLM)
- fake_no_commit = {
- "final_stop": {"reason": "no_final_commitment after 12 turns"},
- "steps": [],
- }
- fake_truth = {
- "scenario_summary": "test scenario",
- "root_cause_aliases": ["server slow response"],
- }
- r = compute_capbench_reward(fake_no_commit, fake_truth,
- endpoint="http://_DUMMY_NO_NETWORK_:4210",
- api_key="x", model="x")
- assert r.bucket == "no_commit", f"expected no_commit, got {r.bucket}"
- assert r.scalar == REWARD_NO_COMMIT
- assert r.has_commit is False
- print(f"no_commit case: {r}")
-
- # Build a fake trace with stop_inconclusive (also no commit)
- fake_inconcl = {
- "final_stop": {"reason": "stop_inconclusive: cannot distinguish hypotheses"},
- "steps": [],
- }
- r2 = compute_capbench_reward(fake_inconcl, fake_truth,
- endpoint="http://_DUMMY_:4210",
- api_key="x", model="x")
- assert r2.bucket == "no_commit", f"expected no_commit (inconclusive), got {r2.bucket}"
- print(f"inconclusive case: {r2}")
-
- # Build a fake trace WITH commit (would call LLM in real path — only check extraction)
- fake_commit_trace = {
- "final_stop": {"reason": "Discriminator confirms server slow response. Committing."},
- "steps": [
- {"task": "stop_controller", "input": {
- "current_hypotheses": [
- {"hypothesis": "server response delay is the root cause", "probability": 0.85},
- {"hypothesis": "client issue", "probability": 0.15},
- ]
- }},
- ],
- }
- final_reason, dom_hyp = _extract_final(fake_commit_trace)
- assert "server slow response" in final_reason
- assert "server response delay" in dom_hyp
- assert not _is_no_commit(final_reason, fake_commit_trace)
- print(f"commit extracted: dom={dom_hyp[:60]!r}")
-
- # Reward shape sanity
- assert REWARD_HIT > REWARD_WRONG > REWARD_NO_COMMIT, \
- "reward ordering must be hit > wrong > no_commit to encourage commit"
-
- print("reward_capbench self-test OK")
diff --git a/vendor/nanoai/agentic/reward_mlr.py b/vendor/nanoai/agentic/reward_mlr.py
deleted file mode 100644
index 7041ae3a5edf73962a9bf176d627d6e0ce4758f9..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/reward_mlr.py
+++ /dev/null
@@ -1,301 +0,0 @@
-"""Multi-level reward for agentic GRPO (per-turn).
-
-Replaces the single-scalar `agentic_reward` (token format only) and
-`reward_mt` (single-turn-with-mock-result) with a decomposition that
-matches the literature (MT-GRPO 2505.11821, CM2 2602.12268,
-long-horizon search 2510.24126):
-
- R(turn_t) = w_term * R_terminal # broadcast: trajectory pass/fail
- + w_fmt * R_format # well-formed tool block or clean terminate
- + w_tool * R_tool_ok # sandbox executed without error
- + w_rec * R_recovery # bonus when turn t recovers from t-1 error
- + w_eff * R_efficiency # small per-turn cost; encourages brevity
-
-`R_terminal` is broadcast (same value on every turn) so all assistant
-tokens get gradient pressure from the final outcome — turns that
-contributed to a refuted hypothesis are still penalized; turns that
-contributed to a pass are still credited. This is the simplest form of
-"credit assignment" before we add TIPS-style information-potential
-(deferred to v2.1).
-
-Returns per-turn rewards aligned with rollout.turn_spans so the GRPO
-training loop can attribute advantages to the right token spans.
-"""
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import Optional
-
-from nanoai import config
-from .grpo_multi_rollout import MultiTurnRollout
-from .eval_tasks import Task
-from .eval_strict import score_strict_record
-
-
-# Default reward weights. Calibrated so terminal dominates (outcome is the
-# ground truth) but per-turn signals can still differentiate a passing
-# trajectory from a "cheats by accident" one.
-# Sourced from the central config (nanoai.config.RewardConfig) so the
-# terminal-dominates invariant lives in one place; values are byte-identical to
-# the legacy literal below (term 1.0, fmt 0.05, tool 0.1, rec 0.1, eff 0.01).
-DEFAULT_WEIGHTS: dict[str, float] = config.reward().as_dict()
-# Rationale for shrunk weights (2026-05-17):
-# Original (1.0/0.25/0.5/0.5/0.05) → on a 4-turn rollout, max shaping ≈ ±5 ≥
-# terminal ±4. Model could maximize "use tool / be well-formed / recover" while
-# task.check() failed. agent_eval pass rate stuck at SFT baseline 23.4/38 even
-# after lr-fix and proven ||Δθ|| of correct magnitude. New weights make terminal
-# 5-10× shaping per turn so gradient pressure aligns with the metric we score on.
-
-
-@dataclass
-class MLRReward:
- """Result of multi-level reward computation for one rollout."""
- per_turn: list[float] # aligned with rollout.turn_spans
- total: float # sum over turns
- breakdown: dict[str, float] # per-component aggregate
- # Per-turn recovery flags for logging / debugging.
- recovery_flags: list[bool]
- strict_status: str = ""
- strict_reason: str = ""
-
-
-def _strict_score_from_rollout(rollout: MultiTurnRollout):
- return score_strict_record({
- "task_id": rollout.task_id,
- "passed": rollout.passed,
- "reason": rollout.reason,
- "n_turns": rollout.n_turns,
- "n_tool_errors": rollout.n_tool_errors,
- "truncated": rollout.truncated,
- "transcript": rollout.transcript.turns,
- })
-
-
-def compute_mlr_reward(
- rollout: MultiTurnRollout,
- task: Optional[Task] = None, # reserved for future task-specific checklist
- weights: Optional[dict[str, float]] = None,
- terminal_mode: str = "weak",
- strict_partial_terminal: float = -1.0,
- aggregation: str = "weighted",
-) -> MLRReward:
- """Compute per-turn multi-level reward.
-
- Components:
- * **terminal** — in weak mode, `+w_term` if task.check passed else
- `-w_term`. In strict mode, only strict behavioral pass gets `+w_term`;
- weak-but-partial pass gets `strict_partial_terminal * w_term`.
- Same value is broadcast on every turn.
- * **format** — `+w_fmt` if the turn either contains a well-parsed tool
- call OR is the terminal answer turn (has no tool call but is_terminal).
- Penalizes turns that emit no tool call but aren't terminal (malformed
- / truncated).
- * **tool_ok** — `+w_tool` if tool executed cleanly, `-w_tool` if tool
- errored, `0` if no tool call this turn.
- * **recovery** — `+w_rec` if the previous turn had a tool error AND
- this turn either succeeded a tool call OR reached a terminal pass.
- This is the "fail2pass" signal the design memo flagged as
- high-value (long-horizon search 2510.24126).
- * **efficiency** — `-w_eff` per turn (light pressure to use fewer turns).
-
- `aggregation` controls how the components combine:
- * **"weighted"** (default) — plain sum `term + fmt + tool + rec + eff`. This
- is byte-identical to the legacy behavior.
- * **"gated"** — MAI-Thinking-1 §3.4.1 gated aggregation: shaping (fmt/tool/
- rec/eff) is applied only when the terminal reward cleared the bar
- (`terminal > 0`). A failing or partial trajectory receives ONLY the
- terminal signal, so it cannot farm format/tool/recovery credit while
- `task.check()` fails. This is a structural fix for "shaping overpowers
- terminal" that does not depend on weight tuning. (Group-relative
- *lexicographic* tie-breaking is a trainer/advantage-level concern and is
- intentionally NOT done here, where we only see one rollout.)
- """
- if terminal_mode not in ("weak", "strict"):
- raise ValueError(f"terminal_mode must be 'weak' or 'strict', got {terminal_mode!r}")
- if aggregation not in ("weighted", "gated"):
- raise ValueError(f"aggregation must be 'weighted' or 'gated', got {aggregation!r}")
- w = {**DEFAULT_WEIGHTS, **(weights or {})}
- strict = _strict_score_from_rollout(rollout) if terminal_mode == "strict" else None
- n = len(rollout.turn_spans)
- if n == 0:
- # Pathological: model produced no parseable turn. Treat as flatfail.
- return MLRReward(
- per_turn=[],
- total=-w["term"],
- breakdown={"term": -w["term"], "fmt": 0.0, "tool": 0.0,
- "rec": 0.0, "eff": 0.0},
- recovery_flags=[],
- strict_status=(strict.status if strict is not None else ""),
- strict_reason=(strict.reason if strict is not None else ""),
- )
-
- if strict is None:
- terminal = w["term"] if rollout.passed else -w["term"]
- elif strict.strict_passed:
- terminal = w["term"]
- elif strict.partial_pass:
- terminal = strict_partial_terminal * w["term"]
- else:
- terminal = -w["term"]
- # Gated aggregation: shaping is suppressed on any turn unless the terminal
- # reward cleared the bar (terminal > 0). For a passing trajectory this is a
- # no-op (identical to "weighted"); it only bites failing/partial rollouts.
- gate_shaping = (aggregation == "gated" and terminal <= 0.0)
- breakdown = {"term": 0.0, "fmt": 0.0, "tool": 0.0, "rec": 0.0, "eff": 0.0}
- per_turn: list[float] = []
- recovery_flags: list[bool] = []
-
- prev_tool_ok: Optional[bool] = None
- for ts in rollout.turn_spans:
- r_term = terminal
- breakdown["term"] += r_term
-
- # Format: well-formed if either (a) has a parsed tool call, or
- # (b) terminal turn with no tool (clean final answer).
- well_formed = ts.has_tool_call or (ts.is_terminal and not ts.has_tool_call)
- r_fmt = w["fmt"] if well_formed else 0.0
-
- # Tool ok: +/- based on execution; 0 if no tool call.
- if ts.tool_ok is True:
- r_tool = w["tool"]
- elif ts.tool_ok is False:
- r_tool = -w["tool"]
- else:
- r_tool = 0.0
-
- # Recovery: prev turn errored AND this turn either tool_ok or terminal pass.
- # Computed regardless of gating so recovery_flags stay informative for logs.
- recovered = (
- prev_tool_ok is False
- and (
- ts.tool_ok is True
- or (ts.is_terminal and rollout.passed)
- )
- )
- r_rec = w["rec"] if recovered else 0.0
- recovery_flags.append(recovered)
-
- r_eff = -w["eff"]
-
- if gate_shaping:
- # Terminal did not clear the bar: suppress ALL shaping (both the
- # positive credit and the efficiency penalty) so a failing rollout
- # collapses to its pure terminal signal and cannot be lifted above
- # the fail floor by being well-formed / using tools.
- r_fmt = r_tool = r_rec = r_eff = 0.0
-
- breakdown["fmt"] += r_fmt
- breakdown["tool"] += r_tool
- breakdown["rec"] += r_rec
- breakdown["eff"] += r_eff
-
- per_turn.append(r_term + r_fmt + r_tool + r_rec + r_eff)
- prev_tool_ok = ts.tool_ok
-
- total = sum(per_turn)
- return MLRReward(
- per_turn=per_turn, total=total,
- breakdown=breakdown, recovery_flags=recovery_flags,
- strict_status=(strict.status if strict is not None else ""),
- strict_reason=(strict.reason if strict is not None else ""),
- )
-
-
-# ---- Self-test ----------------------------------------------------------
-
-if __name__ == "__main__":
- from .grpo_multi_rollout import TurnSpan, MultiTurnRollout
- from .eval_tasks import Transcript
-
- def _mk(passed: bool, spans: list[TurnSpan]) -> MultiTurnRollout:
- return MultiTurnRollout(
- task_id="t", tokens=[], prompt_len=0, turn_spans=spans,
- transcript=Transcript(user_prompt=""),
- passed=passed, reason="", n_turns=len(spans),
- n_tool_errors=0, truncated=False, elapsed_s=0.0,
- )
-
- # Case A: clean 2-turn pass — Read+answer.
- case_a = _mk(passed=True, spans=[
- TurnSpan(0, 0, 10, True, "Read", True, False),
- TurnSpan(1, 12, 18, False, None, None, True),
- ])
- r_a = compute_mlr_reward(case_a)
- print(f"A clean_pass: total={r_a.total:.2f} per_turn={r_a.per_turn} "
- f"recover={r_a.recovery_flags}")
- assert r_a.total > 0, "clean pass should have positive total"
- assert not any(r_a.recovery_flags), "no recovery on clean pass"
-
- # Case B: fail2pass — turn0 tool error, turn1 successful Edit + terminal pass.
- case_b = _mk(passed=True, spans=[
- TurnSpan(0, 0, 10, True, "Read", False, False), # tool error
- TurnSpan(1, 12, 22, True, "Edit", True, False), # recovered
- TurnSpan(2, 24, 28, False, None, None, True), # final answer
- ])
- r_b = compute_mlr_reward(case_b)
- print(f"B fail2pass: total={r_b.total:.2f} per_turn={r_b.per_turn} "
- f"recover={r_b.recovery_flags}")
- # Recovery is a one-time bonus on the turn that fixed the prior error.
- # Turn 0 has no prior → False. Turn 1 prev was error & tool_ok → True.
- # Turn 2 prev was OK (turn 1 fixed it) → False (no longer recovering).
- assert r_b.recovery_flags == [False, True, False], (
- f"turn 1 should be the one-shot recovery; got {r_b.recovery_flags}"
- )
- assert r_b.breakdown["rec"] > 0, "should accumulate recovery bonus"
-
- # Case C: flatfail — turn0 tool error, turn1 tool error, no pass.
- case_c = _mk(passed=False, spans=[
- TurnSpan(0, 0, 10, True, "Read", False, False),
- TurnSpan(1, 12, 22, True, "Read", False, True),
- ])
- r_c = compute_mlr_reward(case_c)
- print(f"C flatfail: total={r_c.total:.2f} per_turn={r_c.per_turn} "
- f"recover={r_c.recovery_flags}")
- assert r_c.total < 0, "flatfail should be negative"
- assert not any(r_c.recovery_flags), "no recovery if never succeeds"
-
- # Case D: zero-turn pathological.
- case_d = _mk(passed=False, spans=[])
- r_d = compute_mlr_reward(case_d)
- print(f"D zero_turns: total={r_d.total:.2f}")
- assert r_d.total == -DEFAULT_WEIGHTS["term"]
- assert r_d.per_turn == []
-
- # Ordering invariant: fail2pass (B) should have higher total than flatfail (C).
- assert r_b.total > r_c.total, "fail2pass must outscore flatfail"
- # And clean_pass (A) should outscore fail2pass (recovery bonus < clean efficiency).
- # NOT a hard requirement — depends on weights. Just print for inspection.
- print(f"A>B? {r_a.total > r_b.total} (clean_pass vs fail2pass)")
-
- # ---- Gated aggregation (MAI §3.4.1) ----
- # Default == "weighted" (byte-identical legacy behavior).
- assert compute_mlr_reward(case_a, aggregation="weighted").total == r_a.total
- # On a PASS, gating is a no-op (terminal cleared the bar).
- r_a_gated = compute_mlr_reward(case_a, aggregation="gated")
- assert abs(r_a_gated.total - r_a.total) < 1e-9, "gated must equal weighted on a pass"
-
- # The point of gating: a FAIL that farms shaping (clean tool use every turn but
- # never solves) is collapsed back to the pure terminal floor.
- case_e = _mk(passed=False, spans=[
- TurnSpan(0, 0, 10, True, "Bash", True, False),
- TurnSpan(1, 12, 22, True, "Bash", True, False),
- TurnSpan(2, 24, 34, True, "Bash", True, False),
- TurnSpan(3, 36, 40, False, None, None, True), # wrong final answer
- ])
- r_e_w = compute_mlr_reward(case_e, aggregation="weighted")
- r_e_g = compute_mlr_reward(case_e, aggregation="gated")
- print(f"E gaming_fail: weighted={r_e_w.total:.2f} gated={r_e_g.total:.2f}")
- assert r_e_w.total > r_e_g.total, "weighted lets a gaming fail farm shaping credit"
- assert abs(r_e_g.total - 4 * (-DEFAULT_WEIGHTS["term"])) < 1e-9, \
- "gated fail must collapse to pure terminal floor (n_turns * -term)"
- assert r_e_g.breakdown["fmt"] == 0.0 and r_e_g.breakdown["tool"] == 0.0 \
- and r_e_g.breakdown["eff"] == 0.0, "gated must zero all shaping on a fail"
-
- # Invalid aggregation is rejected.
- try:
- compute_mlr_reward(case_a, aggregation="bogus")
- raise AssertionError("bad aggregation should raise")
- except ValueError:
- pass
-
- print("self-test OK")
diff --git a/vendor/nanoai/agentic/reward_mt.py b/vendor/nanoai/agentic/reward_mt.py
deleted file mode 100644
index f7a44070b913b8aef17e1388a59252a8b28aab4d..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/reward_mt.py
+++ /dev/null
@@ -1,190 +0,0 @@
-"""Multi-turn-aware reward for agentic CVaR-GRPO Phase 2.5.
-
-The Phase 2 first-pass result (see experiments/fat_tail_cvar/docs/taleb_experiment_report_1.md §14)
-showed both mean GRPO and CVaR-GRPO regressing multi-turn pass rate vs the
-SFT base, because the Phase 1/2 reward function (`agentic_reward`) measures
-*single-turn structural format* only — not whether the model can actually
-solve a multi-turn task.
-
-This module defines a reward that is approximately aligned with the eval
-metric, while staying simple enough to run in a single rollout per step
-(no sandbox, no real tool execution).
-
-Strategy: build a small set of *mt-conditioned prompts* where the model has
-already emitted a tool call and received a (mock) result; the rollout is its
-next-turn response. Reward checks:
- +1.0 mentions_result — response references content from the tool result
- +1.0 clean_terminate — ends with `<|assistant_end|>`
- +1.0 no_loop — does not emit another `<|tool_call_start|>` block
- (a defining failure mode at a03_b07-style training)
- +1.0 format_ok — does not leak literal `<|tool_...|>` strings as text
- (catches "format-memorization-without-tokens" failure)
-
-Range: [0, 4], same as `agentic_reward`, so it is a drop-in replacement.
-
-Returns (score, breakdown) for per-component logging.
-"""
-from __future__ import annotations
-
-# Substrings whose literal-text appearance in the rollout indicates the model
-# has emitted special tokens as plain text rather than as ID tokens. This is
-# the failure mode we observed in a03_b07: regex-detectable but parser-blind.
-LEAKED_TOKEN_FRAGMENTS = (
- "<|tool_arg|>",
- "<|tool_val|>",
- "<|tool_call_end|>",
-)
-
-
-def _result_keywords(tool_result: str) -> list[str]:
- """Pick distinctive substrings from the tool result that the response should
- plausibly mention if it grounds. Heuristic: words/tokens length ≥ 5
- (avoids stopwords + filler), case preserved.
- """
- return [tok for tok in tool_result.split() if len(tok) >= 5]
-
-
-def mt_aware_reward(generated: str, tool_result: str) -> tuple[float, dict]:
- """Score a single-turn rollout that is conditioned on a prior tool result.
-
- `generated` is the assistant's text after `<|assistant_start|>`, exclusive
- of the prefix (matching `agentic_reward`'s contract).
- """
- breakdown = {
- "mentions_result": 0.0,
- "clean_terminate": 0.0,
- "no_loop": 0.0,
- "format_ok": 0.0,
- }
- keywords = _result_keywords(tool_result)
- if keywords and any(kw in generated for kw in keywords):
- breakdown["mentions_result"] = 1.0
- if "<|assistant_end|>" in generated:
- breakdown["clean_terminate"] = 1.0
- if "<|tool_call_start|>" not in generated:
- breakdown["no_loop"] = 1.0
- if not any(frag in generated for frag in LEAKED_TOKEN_FRAGMENTS):
- breakdown["format_ok"] = 1.0
- score = sum(breakdown.values())
- return score, breakdown
-
-
-# A small mt-conditioned prompt distribution. Each item is a (user, tool_call,
-# tool_result) triple. The training loop renders these into a full agentic
-# conversation context so the rollout starts as the *next* assistant turn.
-#
-# Mirrors the spirit of the agent_eval_mt task suite (read/edit/grep/bash/swe)
-# but uses canned mock results so we don't need a sandbox.
-MT_PROMPTS: list[dict] = [
- # Read tasks
- {
- "user": "show me the contents of README.md",
- "tool_name": "Read",
- "tool_args": {"file_path": "README.md"},
- "tool_result": "# nanoai v0.1\nnanochat is a tiny chat agent that uses tools.",
- },
- {
- "user": "open the configs.py file and tell me what's in it",
- "tool_name": "Read",
- "tool_args": {"file_path": "configs.py"},
- "tool_result": "MODEL_DEPTH = 6\nVOCAB_SIZE = 8000\nMAX_SEQ_LEN = 2048",
- },
- {
- "user": "show me the first lines of run.sh",
- "tool_name": "Read",
- "tool_args": {"file_path": "run.sh"},
- "tool_result": "#!/bin/bash\necho hello\nexec python -m main",
- },
-
- # Bash tasks
- {
- "user": "list the files in the current directory",
- "tool_name": "Bash",
- "tool_args": {"command": "ls"},
- "tool_result": "README.md\napp.py\nconfig.toml\nrun.sh",
- },
- {
- "user": "print the working directory",
- "tool_name": "Bash",
- "tool_args": {"command": "pwd"},
- "tool_result": "/home/vader/projects/nanocode",
- },
- {
- "user": "how many python files are in the data/ directory",
- "tool_name": "Bash",
- "tool_args": {"command": "find data -name '*.py' | wc -l"},
- "tool_result": "4",
- },
-
- # Grep tasks
- {
- "user": "find all uses of 'tokenizer' under scripts/",
- "tool_name": "Grep",
- "tool_args": {"pattern": "tokenizer", "path": "scripts/"},
- "tool_result": "scripts/agentic_sft.py:23: tokenizer = load_tokenizer()\n"
- "scripts/base_train.py:90: tokenizer = build_bpe(vocab=8000)",
- },
- {
- "user": "where is the GPT class defined",
- "tool_name": "Grep",
- "tool_args": {"pattern": "class GPT", "path": "."},
- "tool_result": "./nanoai/gpt.py:341:class GPT(nn.Module):",
- },
- {
- "user": "find TODO comments under scripts",
- "tool_name": "Grep",
- "tool_args": {"pattern": "TODO", "path": "scripts/"},
- "tool_result": "scripts/eval.py:14: # TODO: add streaming\n"
- "scripts/sft.py:88: # TODO: handle long sequences",
- },
-
- # Edit tasks (the model already created the file; next turn confirms)
- {
- "user": "create a file hello.py that prints hi",
- "tool_name": "Edit",
- "tool_args": {"file_path": "hello.py", "new_string": "print('hi')"},
- "tool_result": "Created hello.py with 1 line.",
- },
- {
- "user": "make a new file called notes.txt with content 'TODO: review'",
- "tool_name": "Edit",
- "tool_args": {"file_path": "notes.txt", "new_string": "TODO: review"},
- "tool_result": "Created notes.txt with 1 line.",
- },
- {
- "user": "replace foo with bar in test.py",
- "tool_name": "Edit",
- "tool_args": {"file_path": "test.py", "new_string": "x = bar"},
- "tool_result": "Updated test.py with 1 substitution.",
- },
-]
-
-
-__all__ = [
- "mt_aware_reward",
- "MT_PROMPTS",
- "LEAKED_TOKEN_FRAGMENTS",
-]
-
-
-if __name__ == "__main__":
- # Self-check: hand-craft "good" and "bad" responses for one prompt,
- # confirm the reward discriminates them.
- p = MT_PROMPTS[0]
- good = (
- f"the readme says {p['tool_result'].splitlines()[0]} — looks like a tiny "
- f"chat agent. anything else you want to know?<|assistant_end|>"
- )
- bad_loop = (
- f"sure thing<|tool_call_start|>Read<|tool_arg|>file_path"
- f"<|tool_val|>README.md<|tool_call_end|>"
- )
- bad_format = (
- f"yep<|tool_arg|>file_path<|tool_val|>README.md got it<|assistant_end|>"
- )
- bad_ungrounded = "ok done<|assistant_end|>"
-
- for label, text in [("good", good), ("bad_loop", bad_loop),
- ("bad_format", bad_format), ("bad_ungrounded", bad_ungrounded)]:
- s, br = mt_aware_reward(text, p["tool_result"])
- print(f"{label:18s} score={s:.0f} {br}")
diff --git a/vendor/nanoai/agentic/rl_stability.py b/vendor/nanoai/agentic/rl_stability.py
deleted file mode 100644
index 5b4249e40528deda4c7fa5be80fef40e451d6b6e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/rl_stability.py
+++ /dev/null
@@ -1,288 +0,0 @@
-"""Shared RL-stability instrumentation for the MAI-Thinking-1 absorption arc.
-
-This module collects the *measurement* primitives that the MAI-Thinking-1 report
-(Microsoft, 2026-06-02, "Building a Hill-Climbing Machine") relies on to keep a
-long RL climb stable, plus the diagnostics nanoai already wished it had. None of
-these touch the optimizer — they only observe — so the same code is reused by the
-math single-turn trainer (Stage A), the agentic multi-turn trainer (Stage B), the
-self-distillation bridge, and the BF16-vs-NVFP4 precision probe.
-
-Primitives
-----------
-* `importance_weighted_entropy` — MAI eq.(7): the per-token policy entropy
- estimated *from the current training batch* using the importance ratio. This is
- the signal the adaptive-entropy controller regulates.
-* `EntropyController` — MAI eq.(8): an integral controller that nudges
- the clip *upper* bound `k` toward a target entropy `H*`. Replaces an explicit
- entropy bonus (which the report — and nanoai's own attempts — found weaker).
-* `GradNormSpikeCounter` — counts gradient-norm spikes, the symptom MAI's
- outer ratio clip is meant to remove. Lets an A/B quantify "fewer spikes".
-* `kl_from_init_k3` — Schulman k3 KL estimate with d-clipping. nanoai
- already learned the hard way that a naive logp-diff is NOT a KL anchor and that
- k3 needs d-clipping (see lessons `naive_logp_diff_not_kl_anchor`,
- `k3_needs_d_clipping`); this centralizes the correct form.
-* `logprob_gap_stats` — train-vs-inference per-token logprob gap
- distribution. The core number MAI watches when deciding RL precision (they keep
- RL in BF16 because lower precision widens this gap and destabilizes IS).
-
-Everything is pure torch and runs on CPU, so the `__main__` self-test is part of
-the local (no-CUDA) quality gate.
-"""
-from __future__ import annotations
-
-from dataclasses import dataclass, field
-from typing import Optional
-
-import torch
-
-
-def _masked_select(x: torch.Tensor, valid: Optional[torch.Tensor]) -> torch.Tensor:
- """Flatten `x` to the entries where `valid` is truthy (or all of it)."""
- if valid is None:
- return x.reshape(-1)
- m = valid.reshape(-1).bool()
- return x.reshape(-1)[m]
-
-
-def importance_weighted_entropy(
- logp: torch.Tensor,
- ratio: torch.Tensor,
- valid: Optional[torch.Tensor] = None,
-) -> float:
- """MAI eq.(7): importance-weighted per-token entropy of the current policy.
-
- Ĥ(π_θ) = (1/|T|) Σ_{(i,t)∈T} − log π_θ(y_{i,t}) · r_{i,t}
-
- where `logp` = log π_θ(y_t) under the *training* policy, `ratio` = r_t is the
- importance-sampling ratio π_θ/π_old at that token (so that on a fresh on-policy
- step r≈1 and this reduces to the plain token-entropy estimate), and `valid`
- masks the set T of (response, token) pairs in the batch.
-
- Returns a python float (the batch entropy estimate) so callers can feed it
- straight into `EntropyController.update`.
- """
- lp = _masked_select(logp, valid)
- r = _masked_select(ratio, valid)
- if lp.numel() == 0:
- return 0.0
- # detach: this is a measurement, never a gradient path.
- val = (-lp.detach() * r.detach()).mean()
- return float(val.item())
-
-
-@dataclass
-class EntropyController:
- """Integral controller on the clip upper bound (MAI eq.(8)).
-
- At each training step, given the observed batch entropy Ĥ, update::
-
- k ← clip( k + δ · sign(H* − Ĥ), 0, k_max )
-
- Intuition: entropy too low ⇒ widen the upper clip (raise k) so the policy can
- push alternative tokens more aggressively; entropy high enough ⇒ tighten.
-
- The clip interval used by the PPO/GRPO objective is then::
-
- low = 1 − ε
- high = (1 − ε)^{-1} + k
-
- At the initial k=0 the bounds are multiplicative inverses (symmetric in
- log-ratio space). Note MAI's base ε=0.6 is *much* wider than DAPO's ε≈0.2 —
- the adaptive `k` is what does the regulating, so the base trust region is
- deliberately loose. Always A/B the MAI params against a conservative ε on
- nanoai's micro models.
- """
-
- target_entropy: float = 0.3
- k_max: float = 2.5
- k_step: float = 0.25
- eps_base: float = 0.6
- k: float = 0.0
- # history for logging / plots
- k_history: list[float] = field(default_factory=list)
- entropy_history: list[float] = field(default_factory=list)
-
- def update(self, observed_entropy: float) -> float:
- """Advance the controller one step; returns the new k."""
- sign = 0.0
- diff = self.target_entropy - observed_entropy
- if diff > 0:
- sign = 1.0
- elif diff < 0:
- sign = -1.0
- self.k = float(min(max(self.k + self.k_step * sign, 0.0), self.k_max))
- self.entropy_history.append(float(observed_entropy))
- self.k_history.append(self.k)
- return self.k
-
- def clip_bounds(self) -> tuple[float, float]:
- """Current (low, high) ratio-clip bounds given the controller's k."""
- eps = self.eps_base
- low = 1.0 - eps
- high = 1.0 / (1.0 - eps) + self.k
- return low, high
-
-
-@dataclass
-class GradNormSpikeCounter:
- """Detect & count gradient-norm spikes over a training run.
-
- A step is a "spike" when its grad-norm exceeds both an absolute floor and a
- multiple of the recent running median (median is robust to the spikes
- themselves, unlike a mean). MAI's outer ratio clip is meant to remove these;
- an A/B reports `count` with the clip on vs off.
- """
-
- window: int = 50
- factor: float = 3.0
- min_abs: float = 1.0
- count: int = 0
- _hist: list[float] = field(default_factory=list)
-
- def update(self, grad_norm: float) -> bool:
- g = float(grad_norm)
- is_spike = False
- if len(self._hist) >= max(8, self.window // 4):
- recent = self._hist[-self.window:]
- s = sorted(recent)
- med = s[len(s) // 2]
- thresh = max(self.min_abs, self.factor * med)
- if g > thresh:
- is_spike = True
- self.count += 1
- self._hist.append(g)
- return is_spike
-
-
-def kl_from_init_k3(
- logp_new: torch.Tensor,
- logp_ref: torch.Tensor,
- valid: Optional[torch.Tensor] = None,
- d_clip: float = 3.0,
-) -> float:
- """Schulman k3 estimate of KL(π_θ ‖ π_ref), with d-clipping.
-
- With d = log π_ref − log π_new (so r = π_ref/π_θ = exp(d)), the unbiased,
- always-non-negative k3 estimator is::
-
- KL ≈ E[ exp(d) − d − 1 ]
-
- nanoai lessons baked in here:
- * `naive_logp_diff_not_kl_anchor` — the plain mean(logp_new − logp_ref) is a
- signed quantity that pushes one direction forever; it is NOT a KL. Use k3.
- * `k3_needs_d_clipping` — exp(d) explodes for |d|≫1; clip d to [−d_clip,
- d_clip] before exp so a few off-policy tokens don't blow up the estimate.
-
- Returns a python float (mean over valid tokens).
- """
- d = (logp_ref.detach() - logp_new.detach())
- d = _masked_select(d, valid)
- if d.numel() == 0:
- return 0.0
- d = torch.clamp(d, -abs(d_clip), abs(d_clip))
- k3 = torch.exp(d) - d - 1.0
- return float(k3.mean().item())
-
-
-def logprob_gap_stats(
- logp_a: torch.Tensor,
- logp_b: torch.Tensor,
- valid: Optional[torch.Tensor] = None,
-) -> dict[str, float]:
- """Distribution of the per-token |logp_a − logp_b| gap (e.g. train vs infer).
-
- Returns mean / p50 / p90 / p99 / max. MAI keeps RL in BF16 precisely because
- a lower-precision inference engine widens this gap and compounding it over a
- long rollout destabilizes the importance-sampling correction. Phase 5 uses
- this to test the same claim on nanoai's NVFP4 stack.
- """
- gap = (logp_a.detach() - logp_b.detach()).abs()
- gap = _masked_select(gap, valid)
- if gap.numel() == 0:
- return {"mean": 0.0, "p50": 0.0, "p90": 0.0, "p99": 0.0, "max": 0.0}
- gap_sorted = torch.sort(gap).values
-
- def _pct(p: float) -> float:
- idx = min(gap_sorted.numel() - 1, int(p * gap_sorted.numel()))
- return float(gap_sorted[idx].item())
-
- return {
- "mean": float(gap.mean().item()),
- "p50": _pct(0.50),
- "p90": _pct(0.90),
- "p99": _pct(0.99),
- "max": float(gap_sorted[-1].item()),
- }
-
-
-# ---- Self-test (CPU, part of the no-CUDA local gate) --------------------------
-
-if __name__ == "__main__":
- torch.manual_seed(0)
-
- # 1. importance_weighted_entropy: on a fresh on-policy step (ratio==1) it must
- # equal the plain mean negative-logp; higher with a flatter distribution.
- logp_sharp = torch.log(torch.tensor([[0.9, 0.9, 0.9, 0.9]])) # confident
- logp_flat = torch.log(torch.tensor([[0.25, 0.25, 0.25, 0.25]])) # uniform-ish
- ones = torch.ones_like(logp_sharp)
- h_sharp = importance_weighted_entropy(logp_sharp, ones)
- h_flat = importance_weighted_entropy(logp_flat, ones)
- print(f"entropy sharp={h_sharp:.4f} flat={h_flat:.4f}")
- assert h_flat > h_sharp, "flatter distribution must have higher entropy"
- # ratio scaling: doubling ratio doubles the weighted entropy estimate.
- h_r2 = importance_weighted_entropy(logp_flat, ones * 2.0)
- assert abs(h_r2 - 2 * h_flat) < 1e-5, "entropy must be linear in ratio"
-
- # 2. EntropyController: when observed < target it should RAISE k (widen upper
- # clip); when observed > target it should LOWER k; bounded in [0, k_max].
- ctl = EntropyController(target_entropy=0.3, k_max=2.5, k_step=0.25, eps_base=0.6)
- low0, high0 = ctl.clip_bounds()
- # k=0 ⇒ bounds are multiplicative inverses (symmetric in log space).
- assert abs(low0 * high0 - 1.0) < 1e-9, f"k=0 bounds not inverse: {low0}*{high0}"
- for _ in range(5):
- ctl.update(observed_entropy=0.1) # too low ⇒ raise k
- assert ctl.k > 0.0, "low entropy should raise k"
- k_after_raise = ctl.k
- for _ in range(20):
- ctl.update(observed_entropy=0.9) # too high ⇒ lower k toward 0
- assert ctl.k < k_after_raise, "high entropy should lower k"
- assert 0.0 <= ctl.k <= 2.5, "k must stay in [0, k_max]"
- _, high_now = ctl.clip_bounds()
- print(f"controller k={ctl.k:.3f} bounds=({low0:.3f},{high_now:.3f})")
-
- # 3. GradNormSpikeCounter: a steady baseline then a big spike.
- cnt = GradNormSpikeCounter(window=50, factor=3.0, min_abs=1.0)
- for _ in range(40):
- cnt.update(2.0) # steady baseline
- spiked = cnt.update(50.0) # clear spike
- assert spiked and cnt.count == 1, f"should flag exactly one spike, got {cnt.count}"
- cnt.update(2.1) # back to normal, no new spike
- assert cnt.count == 1, "normal step must not count"
- print(f"spikes count={cnt.count}")
-
- # 4. kl_from_init_k3: identical policies ⇒ 0; divergence ⇒ >0; d-clipping caps
- # a pathological token so the estimate stays finite.
- lp = torch.log(torch.tensor([[0.5, 0.5, 0.5, 0.5]]))
- assert kl_from_init_k3(lp, lp) == 0.0, "same policy ⇒ KL 0"
- lp_ref = torch.log(torch.tensor([[0.5, 0.5, 0.5, 0.5]]))
- lp_new = torch.log(torch.tensor([[0.9, 0.1, 0.5, 0.5]]))
- kl = kl_from_init_k3(lp_new, lp_ref)
- assert kl > 0.0, "divergent policy ⇒ KL > 0"
- # pathological: a token where new policy assigns ~0 prob ⇒ huge d; clipping
- # must keep it finite and bounded by exp(d_clip)-d_clip-1.
- lp_path = torch.log(torch.tensor([[1e-9]]))
- lp_refp = torch.log(torch.tensor([[0.9]]))
- kl_path = kl_from_init_k3(lp_path, lp_refp, d_clip=3.0)
- import math
- assert kl_path <= math.exp(3.0) - 3.0 - 1.0 + 1e-6, "d-clip must bound k3"
- print(f"kl divergent={kl:.4f} clipped_path={kl_path:.4f}")
-
- # 5. logprob_gap_stats: identical ⇒ all zero; known gap ⇒ correct stats.
- a = torch.tensor([[-1.0, -2.0, -3.0, -4.0]])
- assert logprob_gap_stats(a, a)["max"] == 0.0
- b = a + torch.tensor([[0.1, 0.2, 0.3, 0.4]])
- st = logprob_gap_stats(a, b)
- assert abs(st["max"] - 0.4) < 1e-6 and abs(st["mean"] - 0.25) < 1e-6
- print(f"gap {st}")
-
- print("rl_stability self-test OK")
diff --git a/vendor/nanoai/agentic/rollout_metrics.py b/vendor/nanoai/agentic/rollout_metrics.py
deleted file mode 100644
index 5fd3ee6cf510e9bf2782d5dba339cfdb0f3fc3de..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/rollout_metrics.py
+++ /dev/null
@@ -1,106 +0,0 @@
-"""Rollout diversity / exploration metrics (Polaris-style distinct-N).
-
-Shared by the agentic RL trainers (intra-group diversity monitor) and the
-offline difficulty/diagnostic scripts. Pure-Python, no torch / no GPU, so it
-is unit-testable on macOS without CUDA.
-
-Polaris (HKU NLP) uses the *distinct-4-gram ratio* — unique 4-grams over total
-4-grams — as its rollout-diversity signal, and tunes the sampling temperature
-to keep this ratio from collapsing across training stages. For on-policy RL the
-decision-relevant quantity is **intra-group** diversity: if the N rollouts of
-one prompt are near-identical, their rewards have ~zero variance and the group
-contributes no gradient (cf. DAPO dynamic-sampling skipping zero-variance
-groups). A collapsing distinct-N is an early warning of mode collapse.
-
-Reference: experiments/agentic_rl/docs/polaris_followups (see project plan).
-"""
-from __future__ import annotations
-
-from collections import Counter
-from typing import Iterable, Sequence, Hashable
-
-
-def _ngrams(seq: Sequence[Hashable], n: int) -> list[tuple]:
- if n <= 0:
- raise ValueError(f"n must be >= 1, got {n}")
- if len(seq) < n:
- return []
- return [tuple(seq[i:i + n]) for i in range(len(seq) - n + 1)]
-
-
-def distinct_ngram_ratio(sequences: Iterable[Sequence[Hashable]], n: int = 4) -> float:
- """Polaris distinct-N over a *pool* of sequences.
-
- sequences: iterable of sequences (token-id lists, word lists, char lists —
- anything whose elements are hashable). When you pass the N
- rollouts of a single prompt, this is the intra-group diversity.
- n: n-gram order (Polaris uses 4).
-
- Returns unique_ngrams / total_ngrams in [0, 1]. 1.0 = every n-gram occurs
- once (maximal diversity); → 0 = the pool is dominated by repeated n-grams
- (rollouts collapsing onto each other / degenerate loops). Returns 0.0 when
- no sequence is long enough to contain an n-gram (defensive: too-short or
- empty rollouts carry no diversity signal).
- """
- total = 0
- seen: set[tuple] = set()
- for seq in sequences:
- grams = _ngrams(list(seq), n)
- total += len(grams)
- seen.update(grams)
- if total == 0:
- return 0.0
- return len(seen) / total
-
-
-def self_repetition_ratio(seq: Sequence[Hashable], n: int = 4) -> float:
- """Fraction of n-grams *within a single* sequence that are repeats.
-
- 1 - distinct-N of one sequence. High → the trajectory loops on itself
- (the degenerate failure mode strict scoring penalizes). Returns 0.0 for
- sequences too short to form an n-gram.
- """
- grams = _ngrams(list(seq), n)
- if not grams:
- return 0.0
- counts = Counter(grams)
- repeats = sum(c - 1 for c in counts.values())
- return repeats / len(grams)
-
-
-def group_diversity_stats(token_lists: Iterable[Sequence[Hashable]], n: int = 4) -> dict:
- """Convenience bundle for the trainer's per-step diversity log.
-
- Returns {distinct_n, mean_self_repeat, n_seqs, n_nonempty}. distinct_n is
- the intra-group Polaris signal; mean_self_repeat flags per-trajectory loops.
- """
- seqs = [list(s) for s in token_lists]
- nonempty = [s for s in seqs if len(s) >= n]
- self_reps = [self_repetition_ratio(s, n) for s in nonempty]
- return {
- "distinct_n": distinct_ngram_ratio(seqs, n),
- "mean_self_repeat": (sum(self_reps) / len(self_reps)) if self_reps else 0.0,
- "n_seqs": len(seqs),
- "n_nonempty": len(nonempty),
- }
-
-
-if __name__ == "__main__": # lightweight self-test (no GPU needed)
- # All-identical rollouts → minimal intra-group diversity.
- same = [[1, 2, 3, 4, 5]] * 4
- r_same = distinct_ngram_ratio(same, n=4)
- # Each rollout distinct → ratio 1.0.
- diff = [[1, 2, 3, 4, 5], [6, 7, 8, 9, 10], [11, 12, 13, 14, 15], [16, 17, 18, 19, 20]]
- r_diff = distinct_ngram_ratio(diff, n=4)
- # Looping trajectory → high self-repeat.
- loop = [9, 9, 9, 9, 9, 9, 9, 9]
- sr = self_repetition_ratio(loop, n=4)
- print(f"distinct-4 identical group : {r_same:.3f} (expect low, 4 copies of 2 grams → 2/8=0.25)")
- print(f"distinct-4 distinct group : {r_diff:.3f} (expect 1.000)")
- print(f"self-repeat looping seq : {sr:.3f} (expect high, ->1.0)")
- print(f"group_diversity_stats : {group_diversity_stats(diff)}")
- assert r_same < r_diff, "identical group must be less diverse than distinct group"
- assert abs(r_diff - 1.0) < 1e-9, "fully-distinct group must be 1.0"
- assert sr > 0.5, "looping sequence must have high self-repetition"
- assert distinct_ngram_ratio([[1, 2]], n=4) == 0.0, "too-short → 0.0"
- print("OK rollout_metrics self-test passed")
diff --git a/vendor/nanoai/agentic/synth_dataloader.py b/vendor/nanoai/agentic/synth_dataloader.py
deleted file mode 100644
index 8a794762a74f9e2064db3b8e46a8d865dc0ae367..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/synth_dataloader.py
+++ /dev/null
@@ -1,136 +0,0 @@
-"""Pure-SYNTH pretrain dataloader.
-
-Drop-in for `nanoai.dataloader.tokenizing_distributed_data_loader_bos_bestfit`
-that draws documents exclusively from the PleIAs/SYNTH dataset (no web, no
-code mix). Used by `scripts/agentic_pretrain_synth_only.py` for the d6 SYNTH
-vs ClimbMix ablation.
-
-BOS-aligned best-fit packing logic mirrors `code_dataloader.py` exactly; the
-only difference is the document source.
-"""
-
-from __future__ import annotations
-
-import torch
-
-from nanoai.common import get_dist_info
-from nanoai.agentic.data.synth_data import (
- list_synth_parquet_files,
- split_paths,
- synth_doc_batches,
-)
-
-
-def _synth_document_batches(split: str, tokenizer_batch_size: int):
- paths = list_synth_parquet_files()
- if not paths:
- raise FileNotFoundError(
- "No SYNTH parquet shards found. Expected at "
- "$NANOAI_BASE_DIR/code_pretrain_data/synth-pleias/ "
- "(symlink to PleIAs/SYNTH download is fine)."
- )
- paths = split_paths(paths, split)
- _ddp, ddp_rank, _local, ddp_world_size = get_dist_info()
- return synth_doc_batches(paths, ddp_rank, ddp_world_size, tokenizer_batch_size)
-
-
-def tokenizing_synth_only_data_loader(
- tokenizer,
- B: int,
- T: int,
- split: str,
- tokenizer_threads: int = 4,
- tokenizer_batch_size: int = 128,
- device: str = "cuda",
- buffer_size: int = 1000,
- seed: int = 0,
-):
- """BOS-aligned best-fit packing over pure-SYNTH docs.
-
- Same yield shape as nanoai's `tokenizing_distributed_data_loader_bos_bestfit`.
- """
- assert split in ("train", "val")
- row_capacity = T + 1
- bos_token = tokenizer.get_bos_token_id()
-
- batches = _synth_document_batches(split, tokenizer_batch_size)
- doc_buffer: list[list[int]] = []
-
- def refill_buffer():
- doc_batch = next(batches)
- token_lists = tokenizer.encode(
- doc_batch, prepend=bos_token, num_threads=tokenizer_threads
- )
- for tokens in token_lists:
- doc_buffer.append(tokens)
-
- use_cuda = device == "cuda"
- row_buffer = torch.empty((B, row_capacity), dtype=torch.long)
- cpu_buffer = torch.empty(2 * B * T, dtype=torch.long, pin_memory=use_cuda)
- gpu_buffer = torch.empty(2 * B * T, dtype=torch.long, device=device)
- cpu_inputs = cpu_buffer[: B * T].view(B, T)
- cpu_targets = cpu_buffer[B * T :].view(B, T)
- inputs = gpu_buffer[: B * T].view(B, T)
- targets = gpu_buffer[B * T :].view(B, T)
-
- while True:
- for row_idx in range(B):
- pos = 0
- while pos < row_capacity:
- while len(doc_buffer) < buffer_size:
- refill_buffer()
- remaining = row_capacity - pos
-
- best_idx = -1
- best_len = 0
- for i, doc in enumerate(doc_buffer):
- doc_len = len(doc)
- if doc_len <= remaining and doc_len > best_len:
- best_idx = i
- best_len = doc_len
-
- if best_idx >= 0:
- doc = doc_buffer.pop(best_idx)
- doc_len = len(doc)
- row_buffer[row_idx, pos : pos + doc_len] = torch.tensor(doc, dtype=torch.long)
- pos += doc_len
- else:
- shortest_idx = min(range(len(doc_buffer)), key=lambda i: len(doc_buffer[i]))
- doc = doc_buffer.pop(shortest_idx)
- row_buffer[row_idx, pos : pos + remaining] = torch.tensor(
- doc[:remaining], dtype=torch.long
- )
- pos += remaining
-
- cpu_inputs.copy_(row_buffer[:, :-1])
- cpu_targets.copy_(row_buffer[:, 1:])
- gpu_buffer.copy_(cpu_buffer, non_blocking=use_cuda)
- yield inputs, targets
-
-
-def tokenizing_synth_only_data_loader_with_state(
- tokenizer,
- B: int,
- T: int,
- split: str,
- tokenizer_threads: int = 4,
- tokenizer_batch_size: int = 128,
- device: str = "cuda",
- resume_state_dict=None,
- buffer_size: int = 1000,
- seed: int = 0,
-):
- """State-aware variant. Resume state is a stub (single epoch counter) —
- same simplification as `code_dataloader.tokenizing_code_mixed_data_loader_with_state`.
- """
- base = tokenizing_synth_only_data_loader(
- tokenizer, B, T, split,
- tokenizer_threads=tokenizer_threads,
- tokenizer_batch_size=tokenizer_batch_size,
- device=device,
- buffer_size=buffer_size,
- seed=seed,
- )
- state = {"epoch": 1, "pq_idx": 0, "rg_idx": 0}
- for inputs, targets in base:
- yield inputs, targets, state
diff --git a/vendor/nanoai/agentic/teacher.py b/vendor/nanoai/agentic/teacher.py
deleted file mode 100644
index c74c78224d1c02ed3c6923572f7479cb31690d6b..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/teacher.py
+++ /dev/null
@@ -1,673 +0,0 @@
-"""Teacher abstraction for On-Policy Distillation.
-
-The teacher in OPD is FROZEN — no gradient, no optimizer state. Its only job
-is to produce per-token log-probabilities (and optionally full logp distributions)
-for a student rollout, so the student can compute reverse KL = logp_S - logp_T.
-
-This module wraps two kinds of teachers:
-
- - `NanochatTeacher`: a nanoai GPT loaded via `checkpoint_manager.load_model`.
- Shares the student's tokenizer and sequence-length conventions, so no
- re-tokenization is needed. Used in V1 (d12 self-distill) and as the
- primary teacher in V2/V3/V4.
-
- - `HFTeacher` (stub, populated in V2): a HuggingFace Transformers model
- (e.g. Qwen3-1.7B) with a DIFFERENT tokenizer. The byte-span alignment
- helper lives in `nanoai.agentic.cross_tokenizer_kd` and converts
- teacher logp back to student-token positions.
-
-The Teacher object exposes:
-
- teacher.compute_logp(inputs, targets) -> (B, T)
- # per-token logp(targets_t | inputs_<=t) under the teacher.
- # Negative-log-likelihood is returned negated so callers can do
- # rev_kl = logp_student - logp_teacher
- # directly.
-
- teacher.compute_full_logp(inputs) -> (B, T, V)
- # Full log-softmax distribution. Used only for forward-KL / JSD
- # ablations (V1 ablation b). Expensive in memory; gated behind an
- # explicit caller request.
-
-Caching: V2 will gain a `LogpCache` mixin keyed on a hash of the input token
-sequence; for V1 the teacher is invoked every step (no cache).
-"""
-
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import Optional
-
-import torch
-
-
-@dataclass
-class TeacherInfo:
- """Minimal metadata captured at load time. Useful for logging and gating."""
- tag: str
- source: str # "base" | "sft" | "dpo" | "grpo" | "custom_dir" | "hf"
- step: Optional[int]
- depth: Optional[int] # n_layer
- n_embd: Optional[int]
- vocab_size: Optional[int]
- same_tokenizer: bool # True iff teacher's tokenizer matches the student's
-
-
-class Teacher:
- """Interface contract (duck-typed; no abstract base required)."""
-
- info: TeacherInfo
-
- def compute_logp(self, inputs: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
- raise NotImplementedError
-
- def compute_logp_and_validity(
- self, inputs: torch.Tensor, targets: torch.Tensor
- ):
- """Return ``(logp, validity)`` where ``validity[b,t] == 1.0`` iff
- ``logp[b,t]`` is a meaningful teacher logp for the realized target.
-
- The default implementation derives validity from ``targets >= 0``,
- which is correct for same-tokenizer teachers (NanochatTeacher): when
- the target is masked, cross_entropy uses ignore_index=-1 and the
- returned logp is 0.0 anyway — `validity = 0` flags those positions.
-
- Cross-tokenizer teachers MUST override this to also mark
- alignment-failure positions as invalid. Otherwise the caller's
- reverse-KL objective will silently train against a fabricated
- teacher logp=0 wherever the API call or byte-span alignment failed
- (P1-OPD-3 in TODO.md).
- """
- logp = self.compute_logp(inputs, targets)
- validity = (targets >= 0).to(dtype=logp.dtype)
- return logp, validity
-
- def compute_full_logp(self, inputs: torch.Tensor) -> torch.Tensor:
- raise NotImplementedError
-
-
-# ---------------------------------------------------------------------------
-# Nanochat teacher (same tokenizer as student)
-# ---------------------------------------------------------------------------
-
-
-class NanochatTeacher(Teacher):
- """A frozen nanoai GPT used as the OPD teacher.
-
- Construct via `NanochatTeacher.from_source(...)` which delegates to
- `checkpoint_manager.load_model` and freezes the model. The teacher's
- forward MUST be wrapped in `torch.no_grad()` by the caller for correctness;
- we additionally set `requires_grad_(False)` on all parameters so even an
- accidental call inside autograd context produces no graph.
-
- NVFP4 / FP8 conversion happens at the caller's discretion BEFORE the
- teacher is wrapped here — at construction time you decide the precision
- once and for all (typically BF16 for stability; NVFP4 only for the very
- large teachers where BF16 doesn't fit).
- """
-
- def __init__(self, model, tokenizer, info: TeacherInfo):
- self.model = model
- self.tokenizer = tokenizer
- self.info = info
-
- @classmethod
- def from_source(
- cls,
- source: str,
- device: torch.device,
- model_tag: Optional[str] = None,
- step: Optional[int] = None,
- student_tokenizer=None,
- ) -> "NanochatTeacher":
- """Load a frozen nanoai checkpoint by (source, model_tag, step).
-
- source: "base" | "sft" | "dpo" | "grpo" (matches load_model).
- student_tokenizer: optional; if provided we sanity-check that vocab
- sizes match, which is a necessary (not sufficient) condition
- for "same tokenizer".
- """
- from nanoai.checkpoint_manager import load_model
-
- model, tokenizer, meta = load_model(
- source, device, phase="eval", model_tag=model_tag, step=step,
- )
- model.eval()
- for p in model.parameters():
- p.requires_grad_(False)
-
- same_tok = True
- if student_tokenizer is not None:
- same_tok = (
- student_tokenizer.get_vocab_size() == tokenizer.get_vocab_size()
- )
-
- info = TeacherInfo(
- tag=model_tag or "auto",
- source=source,
- step=step if step is not None else meta.get("step"),
- depth=meta["model_config"].get("n_layer"),
- n_embd=meta["model_config"].get("n_embd"),
- vocab_size=meta["model_config"].get("vocab_size"),
- same_tokenizer=same_tok,
- )
- return cls(model=model, tokenizer=tokenizer, info=info)
-
- @classmethod
- def from_dir(
- cls,
- checkpoint_dir: str,
- device: torch.device,
- step: Optional[int] = None,
- student_tokenizer=None,
- ) -> "NanochatTeacher":
- """Load a teacher from an arbitrary directory (when not under NANOAI_BASE_DIR).
-
- Used when the d12 teacher lives in a separate base dir (e.g. the
- d12_fail2pass_v1 ckpt is under /home/vader/.cache/nanochat_agentic_d6/
- while the student trains under /home/vader/.cache/nanochat_agentic/).
- """
- from nanoai.checkpoint_manager import build_model, find_last_step
-
- if step is None:
- step = find_last_step(checkpoint_dir)
- model, tokenizer, meta = build_model(checkpoint_dir, step, device, phase="eval")
- model.eval()
- for p in model.parameters():
- p.requires_grad_(False)
-
- same_tok = True
- if student_tokenizer is not None:
- same_tok = (
- student_tokenizer.get_vocab_size() == tokenizer.get_vocab_size()
- )
-
- info = TeacherInfo(
- tag=checkpoint_dir.rstrip("/").split("/")[-1],
- source="custom_dir",
- step=step,
- depth=meta["model_config"].get("n_layer"),
- n_embd=meta["model_config"].get("n_embd"),
- vocab_size=meta["model_config"].get("vocab_size"),
- same_tokenizer=same_tok,
- )
- return cls(model=model, tokenizer=tokenizer, info=info)
-
- @torch.no_grad()
- def compute_logp(self, inputs: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
- """Per-token log p(targets | context) under the teacher. Returns (B, T).
-
- Mirrors `scripts/agentic_grpo_v2.compute_log_pi`: we use the model's
- own forward(idx, targets, loss_reduction="none") which returns flat
- NLL, then negate and reshape.
-
- Positions with targets == -1 (mask sentinel) yield 0 in the NLL
- (cross_entropy ignore_index=-1) — the caller is responsible for
- also masking these in any reduction.
- """
- nll = self.model(inputs, targets, loss_reduction="none").view_as(inputs)
- return -nll
-
- @torch.no_grad()
- def compute_full_logp(self, inputs: torch.Tensor) -> torch.Tensor:
- """Full log-softmax over the vocab at every position. Returns (B, T, V).
-
- Memory: O(B * T * V). For V=32768, B=8, T=2048 this is 8 * 2048 *
- 32768 * 4 bytes = 2.1 GB in fp32 / 1 GB in bf16. Use only for full-KL
- / JSD ablations on a small validation subset.
- """
- logits = self.model(inputs)
- return logits.log_softmax(dim=-1)
-
- def describe(self) -> str:
- i = self.info
- return (
- f"NanochatTeacher(tag={i.tag} source={i.source} step={i.step} "
- f"depth={i.depth} n_embd={i.n_embd} vocab={i.vocab_size} "
- f"same_tokenizer={i.same_tokenizer})"
- )
-
-
-# ---------------------------------------------------------------------------
-# HuggingFace teacher (stub — fleshed out in V2)
-# ---------------------------------------------------------------------------
-
-
-class HFTeacher(Teacher):
- """Stub. Concrete implementation lives in V2 along with byte-span alignment.
-
- Why a stub now: V1 only needs nanoai self-distill. By having the
- placeholder declared we make the `Teacher` interface concrete and the V1
- trainer can already type-check against `Teacher` (duck-typed), letting V2
- plug in HFTeacher without changing trainer signatures.
- """
-
- def __init__(self, *args, **kwargs):
- raise NotImplementedError(
- "HFTeacher is a stub for V2; implement in nanoai.agentic.teacher.HFTeacher "
- "alongside nanoai.agentic.cross_tokenizer_kd."
- )
-
-
-# ---------------------------------------------------------------------------
-# Qwen vLLM API teacher (V1.5)
-# ---------------------------------------------------------------------------
-
-
-class QwenAPITeacher(Teacher):
- """Teacher that queries a running vLLM OpenAI-compatible API for per-token logp.
-
- Use case: another team's Qwen 27B (or larger) is already serving on the
- box via `vllm.entrypoints.openai.api_server`. We piggyback on it for the
- cost of HTTP round-trips, without any GPU allocation of our own.
-
- How it works:
- 1. Decode student's (nanoai) token ids to a UTF-8 string `text`.
- 2. POST to `/v1/completions` with:
- model=, prompt=text,
- max_tokens=1, temperature=0, prompt_logprobs=1
- vLLM returns `choices[0].prompt_logprobs`, a list[ Dict[token_id, info] ]
- where info has `logprob` for the actual token at that position.
- 3. Re-tokenize `text` with the Qwen tokenizer to get Qwen token byte spans.
- 4. Use byte-span alignment (cross_tokenizer_kd) to sum Qwen per-token
- logp over each nanoai token's byte interval.
-
- Performance:
- - One HTTP roundtrip per batch row (we send rows sequentially).
- - Qwen 27B prefill on a B300: ~1500 tokens/sec, so a 200-token rollout
- adds ~135 ms compute + ~50 ms HTTP = ~200 ms / row.
- - For B=2 rollouts/step, teacher forward ~ 400 ms / step.
- """
-
- REQUEST_TIMEOUT_S = 60.0
-
- def __init__(
- self,
- api_url: str,
- model_name: str,
- student_tokenizer,
- device: torch.device,
- qwen_tokenizer_path: str = "", # kept for API compat; unused
- ):
- try:
- import requests # noqa: F401
- except ImportError as e:
- raise ImportError("QwenAPITeacher requires `requests` installed.") from e
-
- self.api_url = api_url.rstrip("/")
- self.model_name = model_name
- self.student_tokenizer = student_tokenizer
- self.device = device
-
- # We DO NOT load the Qwen tokenizer — vLLM's prompt_logprobs response
- # already includes `decoded_token` per position, which we use directly
- # for byte-span alignment. This avoids needing access to /models/.
-
- self.info = TeacherInfo(
- tag=model_name,
- source="hf",
- step=None,
- depth=None,
- n_embd=None,
- vocab_size=None,
- same_tokenizer=False,
- )
-
- def describe(self) -> str:
- return (
- f"QwenAPITeacher(model={self.model_name} api={self.api_url} "
- f"student_vocab={self.student_tokenizer.get_vocab_size()})"
- )
-
- # ---------------- public API -------------------------------------------
-
- def compute_logp(self, inputs: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
- """Backward-compatible wrapper: returns logp only and DISCARDS validity.
-
- Callers should prefer `compute_logp_and_validity` so they can mask
- out positions where the teacher couldn't score the target (alignment
- / API failure). See P1-OPD-3 in TODO.md.
- """
- logp, _ = self.compute_logp_and_validity(inputs, targets)
- return logp
-
- def compute_logp_and_validity(
- self, inputs: torch.Tensor, targets: torch.Tensor
- ):
- """Per-token (logp, validity) under the Qwen teacher.
-
- inputs: (B, T) long, nanoai token ids
- targets: (B, T) long, nanoai token ids of next tokens (-1 = mask)
-
- Returns:
- logp: (B, T) float tensor. 0.0 wherever validity is 0.
- validity: (B, T) float tensor in {0, 1}. 1 iff:
- - target[b,t] >= 0, AND
- - all HTTP / decode / alignment steps succeeded, AND
- - byte-span alignment produced a non-None logp.
-
- The caller MUST multiply validity into its position-wise loss
- weight — otherwise positions where the teacher silently failed
- are trained against logp_T=0, fabricating a strong reverse-KL
- signal where the teacher actually said nothing.
- """
- B, T = inputs.shape
- result = torch.zeros((B, T), dtype=torch.float32, device=self.device)
- validity = torch.zeros((B, T), dtype=torch.float32, device=self.device)
-
- for b in range(B):
- nano_ids = inputs[b].tolist()
- valid_targets = (targets[b] >= 0)
- if not valid_targets.any():
- continue
- last_valid = int(valid_targets.nonzero(as_tuple=True)[0].max().item())
- nano_ids_eff = nano_ids[: last_valid + 2] # +2 because targets[t] = nano[t+1]
-
- try:
- text = self._decode_nano(nano_ids_eff)
- except Exception as e:
- # Surface what went wrong so this row's silent skip doesn't
- # masquerade as a zero teacher signal (P0 bug encountered when
- # QwenTokenizer student lacked .enc; would yield valid_T=0 across
- # the whole run with no visible error).
- print(f"[QwenAPITeacher] _decode_nano failed for batch {b}: {e}")
- continue
- if not text:
- continue
-
- prompt_token_ids = self._tokenize_qwen(text)
- payload = self._fetch_prompt_logprobs(text)
- if payload is None:
- continue
- choices = payload.get("choices", [])
- if not choices:
- continue
- prompt_logprobs = choices[0].get("prompt_logprobs")
- if not prompt_logprobs:
- continue
-
- try:
- aligned = self._align_row_pure(
- text=text,
- nano_ids_eff=nano_ids_eff,
- prompt_token_ids=prompt_token_ids,
- prompt_logprobs=prompt_logprobs,
- )
- except Exception as e:
- print(f"[QwenAPITeacher] alignment failed for batch {b}: {e}")
- continue
-
- # `aligned[i]` is the teacher logp at nano position i. We want
- # the logp of `targets[t] = nano[t+1]` given context `nano[0..t]`,
- # which equals `aligned[t+1]`.
- for t in range(min(T, len(aligned) - 1)):
- if not valid_targets[t].item():
- continue
- lp = aligned[t + 1]
- if lp is None:
- continue # alignment failure; leave logp=0 + validity=0
- result[b, t] = lp
- validity[b, t] = 1.0
- return result, validity
-
- # ---------------- HTTP helpers ------------------------------------------
-
- def _fetch_prompt_logprobs(self, text: str):
- """Call vLLM `/v1/completions` with prompt_logprobs=1; return payload or None."""
- import requests
- try:
- resp = requests.post(
- f"{self.api_url}/v1/completions",
- headers={"Content-Type": "application/json"},
- json={
- "model": self.model_name,
- "prompt": text,
- "max_tokens": 1,
- "temperature": 0,
- "prompt_logprobs": 1,
- },
- timeout=self.REQUEST_TIMEOUT_S,
- )
- resp.raise_for_status()
- return resp.json()
- except Exception as e:
- print(f"[QwenAPITeacher] /v1/completions failed: {e}")
- return None
-
- def _tokenize_qwen(self, text: str):
- """Call vLLM `/tokenize` and return realized prompt token ids, or None on failure.
-
- Knowing the realized ids lets us pick the correct entry from
- `prompt_logprobs[i]` (which may contain BOTH the top-1 and the
- realized token, with the realized one not always being rank 1).
- Without this, the OPD KD signal is silently corrupted whenever
- Qwen's argmax disagrees with the realized prompt token.
- """
- import requests
- try:
- resp = requests.post(
- f"{self.api_url}/tokenize",
- headers={"Content-Type": "application/json"},
- json={"model": self.model_name, "prompt": text},
- timeout=self.REQUEST_TIMEOUT_S,
- )
- resp.raise_for_status()
- payload = resp.json()
- ids = payload.get("tokens")
- if isinstance(ids, list) and all(isinstance(t, int) for t in ids):
- return ids
- return None
- except Exception as e:
- print(f"[QwenAPITeacher] /tokenize failed (will fall back to rank-based pick): {e}")
- return None
-
- # ---------------- pure alignment helper (testable, no HTTP) ------------
-
- def _align_row_pure(
- self,
- text: str,
- nano_ids_eff,
- prompt_token_ids,
- prompt_logprobs,
- ):
- """Pure alignment: produce per-nano-position teacher logp from vLLM payload.
-
- Bug-fix notes (P1-OPD-1 / P1-OPD-2 in TODO.md):
- - prompt_logprobs[0] is None in vLLM (no prior context). The first
- prompt token STILL occupies bytes; if we drop it from the spans,
- every subsequent byte offset shifts left by len(first_token_bytes)
- and every aligned nano position gets the wrong teacher logp.
- Fix: keep position 0 with logp=None, derive its byte length from
- the prompt minus the concatenated tail of decoded_tokens.
- - When vLLM returns multiple entries in prompt_logprobs[i] (top-1
- plus realized when they differ), picking the smallest rank yields
- the model's argmax, not the realized prompt token. Fix: look up
- the realized token id directly via str(prompt_token_ids[i]); fall
- back to the rank-based pick only when /tokenize is unavailable.
-
- Args:
- text: the exact UTF-8 string sent as `prompt` to vLLM.
- nano_ids_eff: list[int] of length N+1; we align to nano_ids_eff[:-1].
- prompt_token_ids: list[int] from /tokenize, or None on failure.
- prompt_logprobs: list of length M = len(prompt_token_ids) (when
- available); prompt_logprobs[0] is None, [i>=1] is a dict.
-
- Returns:
- list[Optional[float]] of length len(nano_ids_eff) - 1.
- """
- from nanoai.agentic.cross_tokenizer_kd import (
- TokenSpan,
- align_logps,
- compute_nanochat_offsets,
- )
-
- prompt_bytes = text.encode("utf-8")
- n_qwen = len(prompt_logprobs)
-
- # If we have prompt_token_ids, validate length match. A mismatch
- # means /tokenize and /v1/completions tokenized differently.
- ids_authoritative = (
- prompt_token_ids is not None and len(prompt_token_ids) == n_qwen
- )
-
- # Codex adversarial review 2026-05-26 finding #1: when ids are not
- # authoritative (prompt_token_ids missing or mismatched length), the
- # per-position fallback was silent (`pass` → sel_val=None → invalid).
- # That left training with effectively zero teacher signal for the
- # whole row, but the trainer kept stepping. We now surface a loud
- # warning so the caller sees the failure mode AND can trip the
- # trainer-side align_rate guard. The actual position-level fallback
- # remains "mark invalid" (no silent argmax substitution, per #3 fix).
- if not ids_authoritative:
- reason = ("prompt_token_ids=None (/tokenize unavailable)"
- if prompt_token_ids is None
- else f"length mismatch (tokenize={len(prompt_token_ids)} "
- f"vs prompt_logprobs={n_qwen})")
- print(f"[QwenAPITeacher] WARNING: alignment row has no authoritative "
- f"realized-token ids ({reason}); all non-position-0 positions "
- f"will be marked invalid. Trainer-side align_rate guard "
- f"should catch sustained zero-signal.")
-
- # Per-position extraction.
- qwen_ids: list = []
- qwen_decoded: list = [] # may be None for position 0
- qwen_logps: list = []
- for i, entry in enumerate(prompt_logprobs):
- realized_id = prompt_token_ids[i] if ids_authoritative else None
-
- if entry is None:
- # Position 0: vLLM has no prior to score against. Still
- # contributes bytes — we'll fill them in from the prompt
- # bytes that the tail doesn't cover.
- qwen_ids.append(realized_id if realized_id is not None else -1)
- qwen_decoded.append(None)
- qwen_logps.append(None)
- continue
-
- sel_key, sel_val = None, None
- if realized_id is not None:
- # Index by realized token id directly (vLLM's prompt_logprobs
- # entry maps str(id) → {logprob, rank, decoded_token}).
- rid_key = str(realized_id)
- if rid_key in entry:
- sel_key, sel_val = rid_key, entry[rid_key]
- # If realized token not present (e.g. prompt_logprobs=1 returned
- # only argmax which differs from realized), mark this position
- # invalid rather than silently substituting argmax — that yields
- # the teacher's preference for a DIFFERENT token than the one
- # being scored, producing systematically wrong KD signal.
- # (codex adversarial review 2026-05-18, finding #3)
- # else: realized_id is None (/tokenize unavailable or mismatched).
- # We have no way to know which entry corresponds to the realized
- # token. The OLD code picked argmax here, silently corrupting the
- # KD signal whenever argmax ≠ realized. Per codex #3 (per-position)
- # and #1 (row-level loud warning emitted above), we now keep
- # sel_val=None → mark this position invalid and let the
- # trainer-side align_rate guard abort if sustained.
-
- if sel_val is None:
- qwen_ids.append(None)
- qwen_decoded.append(None)
- qwen_logps.append(None)
- continue
-
- try:
- tid = int(sel_key)
- except (TypeError, ValueError):
- tid = -1
- qwen_ids.append(tid)
- qwen_decoded.append(sel_val.get("decoded_token", ""))
- qwen_logps.append(sel_val.get("logprob"))
-
- # Bug 1 fix: derive position-0 byte length from prompt - tail.
- # Sum byte lengths of positions 1..end (positions with decoded text).
- tail_bytes_len = 0
- for d in qwen_decoded[1:]:
- if d is None:
- continue
- tail_bytes_len += len(d.encode("utf-8"))
- first_byte_len = max(0, len(prompt_bytes) - tail_bytes_len)
-
- # Build qwen TokenSpans with correct byte offsets across ALL positions
- # (including position 0). Note: text content of position 0 doesn't
- # matter for alignment — only start/end do.
- qwen_spans = []
- cur = 0
- for i, (tid, dec) in enumerate(zip(qwen_ids, qwen_decoded)):
- if i == 0 and dec is None:
- blen = first_byte_len
- # Best-effort text for diagnostics only; alignment uses bytes.
- sp_text = prompt_bytes[:first_byte_len].decode("utf-8", errors="replace")
- else:
- sp_text = dec if dec is not None else ""
- blen = len(sp_text.encode("utf-8"))
- qwen_spans.append(TokenSpan(
- token_id=tid if tid is not None else -1,
- start=cur,
- end=cur + blen,
- text=sp_text,
- is_special=False,
- ))
- cur += blen
-
- # Nano side: align targets at positions 1..N (logp of nano[i] given
- # nano[0..i-1]). Pass the full nano_ids_eff so the LAST position
- # also gets an aligned logp — without it, the final supervised
- # target token is silently dropped (validity stays 0). The teacher
- # bytes from _decode_nano(nano_ids_eff) cover the full eff, so
- # nano spans must cover the full eff too for byte alignment to
- # be complete. (codex adversarial review 2026-05-26, finding #2)
- nano_spans, _ = compute_nanochat_offsets(
- self.student_tokenizer, list(nano_ids_eff)
- )
- return align_logps(nano_spans, qwen_spans, qwen_logps)
-
- def compute_full_logp(self, inputs: torch.Tensor) -> torch.Tensor:
- raise NotImplementedError(
- "QwenAPITeacher does not expose full vocab logp via the OpenAI API."
- )
-
- # ---------------- internals --------------------------------------------
-
- def _decode_nano(self, token_ids) -> str:
- """Decode student-vocab ids to text, rendering specials as <|name|>.
-
- Qwen path: HF tokenizer's .decode(skip_special_tokens=False) already
- renders specials as their string forms (`<|im_start|>`, ``,
- etc.), so a direct decode is correct.
-
- nanoai-rustbpe path: tiktoken's decode drops specials by default,
- so we walk the id stream and substitute names from enc._special_tokens.
- """
- # Qwen fast path: HF tokenizer handles specials inline.
- if getattr(self.student_tokenizer, "chat_kind", None) == "qwen":
- return self.student_tokenizer.decode(token_ids)
-
- enc = self.student_tokenizer.enc
- special_ids = (
- set(enc._special_tokens.values()) if hasattr(enc, "_special_tokens") else set()
- )
-
- # If no specials, fast path through tokenizer.decode.
- if not any(tid in special_ids for tid in token_ids):
- return self.student_tokenizer.decode(token_ids)
-
- # Slow path: piece together token-by-token, rendering specials as their names.
- parts = []
- pending: list = []
- for tid in token_ids:
- if tid in special_ids:
- if pending:
- parts.append(self.student_tokenizer.decode(pending))
- pending = []
- name = None
- for k, v in enc._special_tokens.items():
- if v == tid:
- name = k
- break
- parts.append(name or f"<|sp_{tid}|>")
- else:
- pending.append(tid)
- if pending:
- parts.append(self.student_tokenizer.decode(pending))
- return "".join(parts)
diff --git a/vendor/nanoai/agentic/tree_rollout.py b/vendor/nanoai/agentic/tree_rollout.py
deleted file mode 100644
index 44c0eb77f8909966f2f62d84a25de49a76ad1322..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/tree_rollout.py
+++ /dev/null
@@ -1,287 +0,0 @@
-"""Tree rollout for agentic Tree-GRPO / TRACE — the missing primitive.
-
-`agentic_tree_grpo.py` imports `tree_rollout` and `leaf_to_multiturn_rollout`
-but they were never implemented. This module provides them, plus the adaptive
-branching hook TRACE needs (arXiv 2606.11119): instead of branching a fixed B
-times at every turn boundary (uniform Tree-GRPO, arXiv 2509.21240), a caller can
-pass `branch_fn(node) -> k` to branch only where the mixed-outcome potential is
-high. `branch_fn = lambda node: B` reproduces uniform Tree-GRPO exactly (asserted
-in tests), so this is a strict superset.
-
-Tree shape (matches the advantage-attribution contract in
-agentic_tree_grpo.pack_tree_leaves_for_grad):
- - root: depth 0, holds the prompt tokens, no turn yet.
- - branching at a node of depth d (< max_tree_depth) samples k continuations,
- each generating ONE more agent turn -> k children of depth d+1. Turn `i` of a
- leaf therefore lives in its depth-(i+1) ancestor (path[i]).
- - a node at depth >= max_tree_depth continues LINEARLY (no branching) to a
- terminal state; those extra turns are owned by the deepest tree node.
-
-Sandbox correctness: branching is stateful — to sample a divergent continuation
-from a prefix we need the sandbox in that prefix's state. The FIRST child of a
-node reuses the parent's live sandbox (already in the parent's state); each
-ADDITIONAL branch gets a FRESH `eval_sandbox.Sandbox` with the ancestor tool
-calls REPLAYED into it. Using eval_sandbox.Sandbox keeps the killpg/RLIMIT
-process-group hygiene (lesson: subprocess timeouts leak firejail grandchildren as
-orphans) — never roll a bespoke sandbox here.
-
-Token behavior is identical to `grpo_multi_rollout.rollout_one` at the per-turn
-level; only the tree bookkeeping is added.
-"""
-from __future__ import annotations
-
-import time
-from dataclasses import dataclass, field
-from typing import Callable, Optional
-
-from .eval_loop import (
- _append_tool_result,
- _assistant_end_id,
- _build_initial_tokens,
- _execute,
- _generate_one_turn,
- _parse_tool_call_dispatch,
- _tool_call_pair,
- TERMINAL_TOOLS,
-)
-from .eval_sandbox import Sandbox
-from .eval_tasks import Task, Transcript
-from .grpo_multi_rollout import MultiTurnRollout, TurnSpan
-
-
-@dataclass
-class TreeNode:
- """One node in a rollout tree. A leaf is a full trajectory (prompt + turns).
-
- `tokens` is the FULL token path from the root to this node (so leaves are
- directly packable). `turn_spans` index into `tokens`. `tool_calls` is the
- ordered (name, args) of executed tool turns on the path — replayed to rebuild
- the sandbox when a sibling branch is sampled.
- """
- depth: int
- parent: Optional["TreeNode"]
- tokens: list[int]
- prompt_len: int
- turn_spans: list[TurnSpan] = field(default_factory=list)
- tool_calls: list[tuple] = field(default_factory=list) # [(name, args), ...]
- children: list["TreeNode"] = field(default_factory=list)
- is_leaf: bool = False
- is_terminal: bool = False
- passed: bool = False
- reason: str = ""
- n_tool_errors: int = 0
- truncated: bool = False
- # filled by the trainer / annotate passes:
- reward: float = 0.0
- _R: float = 0.0
- _A_intra: float = 0.0
-
-
-def _step_one_turn(engine, tokenizer, sandbox: Sandbox, parent: "TreeNode",
- task: Task, *, max_tokens_per_turn: int, temperature: float,
- seed: int) -> "TreeNode":
- """Generate ONE agent turn continuing from `parent.tokens`, execute its tool
- in `sandbox`, and return the resulting child node. `sandbox` MUST already be
- in `parent`'s state (replayed or carried); it is mutated into the child's
- state. Mirrors one iteration of grpo_multi_rollout.rollout_one's loop.
- """
- asst_end = _assistant_end_id(tokenizer)
- tc_start, tc_end = _tool_call_pair(tokenizer)
-
- tokens = list(parent.tokens)
- body_start = len(tokens)
- out_ids = _generate_one_turn(
- engine, tokenizer, tokens, max_tokens_per_turn, temperature, seed)
- tokens.extend(out_ids)
- body_end = len(tokens)
-
- ended_clean = bool(out_ids) and out_ids[-1] == asst_end
- body_ids = out_ids[:-1] if ended_clean else out_ids
- truncated = (not ended_clean) and len(out_ids) >= max_tokens_per_turn
-
- child = TreeNode(
- depth=parent.depth + 1, parent=parent, tokens=tokens,
- prompt_len=parent.prompt_len, turn_spans=list(parent.turn_spans),
- tool_calls=list(parent.tool_calls), n_tool_errors=parent.n_tool_errors,
- truncated=parent.truncated or truncated,
- )
- turn_idx = len(parent.turn_spans)
-
- if tc_start in body_ids and tc_end in body_ids:
- i = body_ids.index(tc_start)
- try:
- j = body_ids.index(tc_end, i + 1)
- except ValueError:
- j = -1
- if j > i:
- tool_body = tokenizer.decode(body_ids[i + 1:j])
- try:
- name, args = _parse_tool_call_dispatch(tokenizer, tool_body)
- except ValueError as e:
- name, args, ok, result_text = "", {}, False, f"tool parse error: {e}"
- else:
- result_text, ok = _execute(sandbox, name, args)
- if not ok:
- child.n_tool_errors += 1
- is_terminal = name in TERMINAL_TOOLS
- child.turn_spans.append(TurnSpan(
- turn_idx=turn_idx, start_idx=body_start, end_idx=body_end,
- has_tool_call=True, tool_name=name, tool_ok=ok, is_terminal=is_terminal))
- child.tool_calls.append((name, args))
- if is_terminal:
- child.is_terminal = True
- else:
- if not ended_clean:
- child.tokens.append(asst_end)
- _append_tool_result(tokenizer, child.tokens, result_text)
- return child
-
- # No tool call -> final answer (terminal).
- child.turn_spans.append(TurnSpan(
- turn_idx=turn_idx, start_idx=body_start, end_idx=body_end,
- has_tool_call=False, tool_name=None, tool_ok=None, is_terminal=True))
- child.is_terminal = True
- return child
-
-
-def _fresh_sandbox_at(parent: "TreeNode", task: Task, bash_timeout: float) -> Sandbox:
- """Open a fresh eval_sandbox.Sandbox and replay the ancestor tool calls so it
- matches `parent`'s state. Caller owns the returned sandbox (context-manage it).
- """
- sb = Sandbox(bash_timeout=bash_timeout)
- sb.__enter__()
- task.setup(sb)
- for (name, args) in parent.tool_calls:
- if name != "":
- _execute(sb, name, args)
- return sb
-
-
-def tree_rollout(
- engine, tokenizer, task: Task, *,
- branch_fn: Callable[["TreeNode"], int],
- max_tree_depth: int,
- max_turns: int,
- max_tokens_per_turn: int = 256,
- temperature: float = 0.8,
- base_seed: int = 42,
- bash_timeout: float = 5.0,
-) -> tuple["TreeNode", list["TreeNode"]]:
- """Build a rollout tree and return (root, leaves).
-
- branch_fn(node) -> k in [1, B]: how many continuations to sample at `node`.
- k <= 1 is a linear continuation (no branch). Branching only happens for nodes
- of depth < max_tree_depth; deeper nodes continue linearly to a terminal state.
- `branch_fn = lambda n: B` reproduces uniform Tree-GRPO (B^D leaves).
- """
- if max_turns is None:
- max_turns = task.max_turns
- else:
- max_turns = max(max_turns, getattr(task, "max_turns", 0) or 0)
-
- prompt_tokens = _build_initial_tokens(
- tokenizer, task.prompt, system_prompt=getattr(task, "system_prompt", None))
- root = TreeNode(depth=0, parent=None, tokens=prompt_tokens, prompt_len=len(prompt_tokens))
- leaves: list[TreeNode] = []
- # Monotonic seed counter so every sampled turn across the tree uses a distinct seed.
- seed_ctr = [0]
-
- def next_seed() -> int:
- seed_ctr[0] += 1
- return base_seed + seed_ctr[0] * 7919
-
- def _close(sb: Sandbox):
- try:
- sb.__exit__(None, None, None)
- except Exception:
- pass
-
- def finalize_leaf(node: "TreeNode", sandbox: Sandbox):
- transcript = _transcript_for(node, tokenizer, task)
- passed, reason = task.check(sandbox, transcript)
- node.passed, node.reason = passed, reason
- node.is_leaf = True
- leaves.append(node)
-
- def expand(node: "TreeNode", sandbox: Sandbox):
- """`sandbox` is live and in `node`'s state. This call OWNS the sandbox and
- closes it exactly once — unless it hands it to child-0's recursive expand
- (which then owns it). Every leaf-producing path closes its sandbox."""
- try:
- if node.is_terminal or len(node.turn_spans) >= max_turns:
- finalize_leaf(node, sandbox)
- _close(sandbox)
- return
- if node.depth >= max_tree_depth:
- # Linear continuation in the SAME sandbox until terminal / max_turns.
- cur = node
- while (not cur.is_terminal) and len(cur.turn_spans) < max_turns:
- nxt = _step_one_turn(
- engine, tokenizer, sandbox, cur, task,
- max_tokens_per_turn=max_tokens_per_turn,
- temperature=temperature, seed=next_seed())
- cur.children.append(nxt) # link the linear chain so tree walks reach the leaf
- cur = nxt
- finalize_leaf(cur, sandbox)
- _close(sandbox)
- return
- k = max(1, int(branch_fn(node)))
- for bi in range(k):
- # child 0 reuses the parent's live sandbox; extra branches get a
- # fresh sandbox with ancestor tool calls replayed into it.
- sb = sandbox if bi == 0 else _fresh_sandbox_at(node, task, bash_timeout)
- try:
- child = _step_one_turn(
- engine, tokenizer, sb, node, task,
- max_tokens_per_turn=max_tokens_per_turn,
- temperature=temperature, seed=next_seed())
- except Exception:
- _close(sb)
- raise
- node.children.append(child)
- expand(child, sb) # child's expand() owns + closes sb
- # parent's sandbox was handed to child 0's expand(); nothing to close here.
- except Exception:
- # On error in the terminal/linear branches the sandbox is still ours.
- if node.is_terminal or node.depth >= max_tree_depth:
- _close(sandbox)
- raise
-
- # Root: open one sandbox in the prompt state; expand() takes ownership.
- sb0 = Sandbox(bash_timeout=bash_timeout)
- sb0.__enter__()
- task.setup(sb0)
- expand(root, sb0)
- return root, leaves
-
-
-def _transcript_for(node: "TreeNode", tokenizer, task: Task) -> Transcript:
- """Reconstruct a minimal Transcript from a leaf's turn_spans/tokens for
- task.check (which inspects assistant text + tool results)."""
- transcript = Transcript(user_prompt=task.prompt)
- for ts in node.turn_spans:
- body = tokenizer.decode(node.tokens[ts.start_idx:ts.end_idx])
- if ts.has_tool_call:
- transcript.turns.append({"role": "assistant", "content": body,
- "tool_call": {"name": ts.tool_name, "args": {}}})
- else:
- transcript.turns.append({"role": "assistant", "content": body})
- return transcript
-
-
-def leaf_to_multiturn_rollout(leaf: "TreeNode", task: Task) -> MultiTurnRollout:
- """Adapt a tree leaf to a MultiTurnRollout so reward_mlr.compute_mlr_reward
- can score it (used by agentic_tree_grpo.reward_for_leaf)."""
- return MultiTurnRollout(
- task_id=task.task_id,
- tokens=leaf.tokens,
- prompt_len=leaf.prompt_len,
- turn_spans=leaf.turn_spans,
- transcript=Transcript(user_prompt=task.prompt),
- passed=leaf.passed,
- reason=leaf.reason,
- n_turns=len(leaf.turn_spans),
- n_tool_errors=leaf.n_tool_errors,
- truncated=leaf.truncated,
- elapsed_s=0.0,
- )
diff --git a/vendor/nanoai/agentic/turn_critic.py b/vendor/nanoai/agentic/turn_critic.py
deleted file mode 100644
index 6467443e04f6bd59b98651d7dd3e777193ea3fb2..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/turn_critic.py
+++ /dev/null
@@ -1,313 +0,0 @@
-"""Turn-Critic Probe — lightweight MLP on LM hidden states at turn boundaries.
-
-Probes V(s_t) at each turn boundary (the position right before the assistant's
-first generated token for that turn). Trained on (hidden_state, discounted_return)
-pairs from logged rollouts. If the probe achieves Pearson r > 0.5, it can be
-used for GAE advantage estimation in a GRPO trainer variant.
-
-The probe is intentionally tiny (~200K params) relative to the policy (~70M for d6)
-so it trains in minutes and does not compete with the policy for GPU memory.
-"""
-from __future__ import annotations
-
-import math
-from dataclasses import dataclass
-from typing import Optional
-
-import torch
-import torch.nn as nn
-import torch.nn.functional as F
-
-
-class TurnCritic(nn.Module):
- """MLP probe: hidden_state (n_embd,) -> scalar V(s_t)."""
-
- def __init__(self, n_embd: int = 384, hidden: int = 256, dropout: float = 0.1):
- super().__init__()
- self.net = nn.Sequential(
- nn.Linear(n_embd, hidden),
- nn.ReLU(),
- nn.Dropout(dropout),
- nn.Linear(hidden, hidden // 4),
- nn.ReLU(),
- nn.Linear(hidden // 4, 1),
- )
-
- def forward(self, h: torch.Tensor) -> torch.Tensor:
- """h: (B, n_embd) -> (B,)"""
- return self.net(h).squeeze(-1)
-
-
-class ScaledTurnCritic(nn.Module):
- """Larger value model with residual connections and layer norm.
-
- ~5-10x more params than TurnCritic. Designed to be trained on MC V(s_t)
- labels (gold-standard) with more data (2000+ rollouts).
- """
-
- def __init__(self, n_embd: int = 384, hidden: int = 512, n_layers: int = 3,
- dropout: float = 0.1):
- super().__init__()
- self.input_proj = nn.Linear(n_embd, hidden)
- self.input_norm = nn.LayerNorm(hidden)
- blocks = []
- for _ in range(n_layers):
- blocks.append(ResBlock(hidden, dropout))
- self.blocks = nn.ModuleList(blocks)
- self.head = nn.Sequential(
- nn.LayerNorm(hidden),
- nn.Linear(hidden, 1),
- )
-
- def forward(self, h: torch.Tensor) -> torch.Tensor:
- """h: (B, n_embd) -> (B,)"""
- x = self.input_norm(F.gelu(self.input_proj(h)))
- for block in self.blocks:
- x = block(x)
- return self.head(x).squeeze(-1)
-
-
-class ResBlock(nn.Module):
- def __init__(self, dim: int, dropout: float = 0.1):
- super().__init__()
- self.net = nn.Sequential(
- nn.Linear(dim, dim * 2),
- nn.GELU(),
- nn.Dropout(dropout),
- nn.Linear(dim * 2, dim),
- nn.Dropout(dropout),
- )
- self.norm = nn.LayerNorm(dim)
-
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- return self.norm(x + self.net(x))
-
-
-def extract_turn_boundary_hidden(
- model: nn.Module,
- token_ids: list[int],
- turn_spans: list[dict],
- device: torch.device,
- layer_idx: int = -1,
-) -> list[torch.Tensor]:
- """Extract hidden states at turn boundary positions via forward hook.
-
- Turn boundary = position right before each turn's first token (start_idx - 1).
- This is where the model has consumed all prior context (user prompt, previous
- turns + tool results) and is about to generate the next assistant turn.
-
- layer_idx: which transformer block to probe. -1 = last block.
-
- Returns list of (n_embd,) tensors, one per turn.
- """
- blocks = model.transformer.h
- target_block = blocks[layer_idx]
-
- captured = {}
-
- def hook_fn(module, inp, out):
- captured["hidden"] = out.detach()
-
- handle = target_block.register_forward_hook(hook_fn)
- try:
- inp = torch.tensor([token_ids], dtype=torch.int32, device=device)
- with torch.no_grad():
- model(inp)
- h = captured["hidden"] # (1, seq_len, n_embd)
- finally:
- handle.remove()
-
- boundary_hidden = []
- for ts in turn_spans:
- start_idx = ts["start_idx"] if isinstance(ts, dict) else ts.start_idx
- pos = max(start_idx - 1, 0)
- if pos < h.size(1):
- boundary_hidden.append(h[0, pos, :].clone())
- else:
- boundary_hidden.append(torch.zeros(h.size(-1), device=device))
- return boundary_hidden
-
-
-@dataclass
-class ProbeDataset:
- """Flat list of (hidden_state, target_return) pairs for probe training."""
- hiddens: torch.Tensor # (N, n_embd)
- targets: torch.Tensor # (N,)
-
- def __len__(self):
- return self.hiddens.size(0)
-
- def batches(self, batch_size: int, shuffle: bool = True):
- N = len(self)
- idx = torch.randperm(N) if shuffle else torch.arange(N)
- for start in range(0, N, batch_size):
- sel = idx[start:start + batch_size]
- yield self.hiddens[sel], self.targets[sel]
-
-
-def discounted_returns(per_turn_rewards: list[float], gamma: float = 0.9) -> list[float]:
- """R_j = sum_{m=j..T} gamma^(m-j) * r_m."""
- T = len(per_turn_rewards)
- out = [0.0] * T
- running = 0.0
- for j in range(T - 1, -1, -1):
- running = per_turn_rewards[j] + gamma * running
- out[j] = running
- return out
-
-
-def train_probe(
- probe: TurnCritic,
- dataset: ProbeDataset,
- *,
- lr: float = 1e-3,
- epochs: int = 50,
- batch_size: int = 128,
- val_frac: float = 0.15,
- device: torch.device = torch.device("cpu"),
- verbose: bool = True,
-) -> dict:
- """Train probe on (hidden, target_return) pairs. Returns quality metrics."""
- N = len(dataset)
- val_n = max(int(N * val_frac), 1)
- perm = torch.randperm(N)
- val_idx, train_idx = perm[:val_n], perm[val_n:]
-
- train_ds = ProbeDataset(
- dataset.hiddens[train_idx].to(device=device, dtype=torch.float32),
- dataset.targets[train_idx].to(device=device, dtype=torch.float32),
- )
- val_ds = ProbeDataset(
- dataset.hiddens[val_idx].to(device=device, dtype=torch.float32),
- dataset.targets[val_idx].to(device=device, dtype=torch.float32),
- )
-
- probe = probe.to(device)
- opt = torch.optim.Adam(probe.parameters(), lr=lr)
-
- best_val_loss = float("inf")
- best_state = None
-
- for epoch in range(epochs):
- probe.train()
- train_loss_sum = 0.0
- train_n = 0
- for h_batch, t_batch in train_ds.batches(batch_size):
- pred = probe(h_batch)
- loss = F.mse_loss(pred, t_batch)
- opt.zero_grad()
- loss.backward()
- opt.step()
- train_loss_sum += loss.item() * h_batch.size(0)
- train_n += h_batch.size(0)
-
- probe.eval()
- with torch.no_grad():
- val_pred = probe(val_ds.hiddens)
- val_loss = F.mse_loss(val_pred, val_ds.targets).item()
-
- if val_loss < best_val_loss:
- best_val_loss = val_loss
- best_state = {k: v.clone() for k, v in probe.state_dict().items()}
-
- if verbose and (epoch + 1) % 10 == 0:
- print(f" probe epoch {epoch+1:3d}/{epochs} "
- f"train_mse={train_loss_sum/max(train_n,1):.4f} val_mse={val_loss:.4f}")
-
- if best_state is not None:
- probe.load_state_dict(best_state)
-
- # Quality metrics on val set
- probe.eval()
- with torch.no_grad():
- val_pred = probe(val_ds.hiddens)
- val_targets = val_ds.targets
-
- metrics = compute_probe_quality(val_pred, val_targets)
- metrics["best_val_mse"] = best_val_loss
- metrics["train_n"] = len(train_ds)
- metrics["val_n"] = len(val_ds)
- return metrics
-
-
-def compute_probe_quality(pred: torch.Tensor, target: torch.Tensor) -> dict:
- """Pearson r, Spearman rho, AUC for binary pass/fail discrimination."""
- pred_np = pred.cpu().float()
- target_np = target.cpu().float()
-
- # Pearson r
- pm = pred_np - pred_np.mean()
- tm = target_np - target_np.mean()
- denom = (pm.norm() * tm.norm()).clamp(min=1e-8)
- pearson_r = (pm * tm).sum() / denom
-
- # Spearman rho (rank correlation with tie-averaging)
- def _rank(x):
- sorted_vals, sorted_idx = x.sort()
- ranks = torch.empty_like(x)
- ranks[sorted_idx] = torch.arange(len(x), dtype=x.dtype)
- # Average ranks for ties
- i = 0
- while i < len(sorted_vals):
- j = i + 1
- while j < len(sorted_vals) and sorted_vals[j] == sorted_vals[i]:
- j += 1
- if j > i + 1:
- tie_rank = ranks[sorted_idx[i:j]].mean()
- for k in range(i, j):
- ranks[sorted_idx[k]] = tie_rank
- i = j
- return ranks
-
- pr = _rank(pred_np)
- tr = _rank(target_np)
- prm = pr - pr.mean()
- trm = tr - tr.mean()
- spearman_rho = (prm * trm).sum() / (prm.norm() * trm.norm()).clamp(min=1e-8)
-
- # AUC for binary pass/fail (target > 0.5 = pass-likely)
- # Uses the Wilcoxon-Mann-Whitney statistic (tie-safe)
- binary = (target_np > 0.5).float()
- n_pos = int(binary.sum().item())
- n_neg = int((1 - binary).sum().item())
- auc = float("nan")
- if n_pos > 0 and n_neg > 0:
- pos_scores = pred_np[binary == 1]
- neg_scores = pred_np[binary == 0]
- # U statistic: fraction of (pos, neg) pairs where pos > neg (+ 0.5 for ties)
- u_sum = 0.0
- for ps in pos_scores:
- u_sum += (ps > neg_scores).float().sum().item()
- u_sum += 0.5 * (ps == neg_scores).float().sum().item()
- auc = u_sum / (n_pos * n_neg)
-
- return {
- "pearson_r": pearson_r.item(),
- "spearman_rho": spearman_rho.item(),
- "auc_pass_fail": auc,
- }
-
-
-def gae_advantages(
- per_turn_rewards: list[float],
- values: list[float],
- terminal_value: float,
- gamma: float = 0.9,
- lam: float = 0.95,
-) -> list[float]:
- """GAE(lambda) advantages from per-turn rewards and probe-estimated V(s_t).
-
- values: [V(s_0), V(s_1), ..., V(s_{T-1})], one per turn.
- terminal_value: V(s_T) = 1.0 if passed, 0.0 if failed.
- Returns: [A(0), A(1), ..., A(T-1)].
- """
- T = len(per_turn_rewards)
- assert len(values) == T
- all_v = list(values) + [terminal_value]
- deltas = [per_turn_rewards[t] + gamma * all_v[t + 1] - all_v[t] for t in range(T)]
- advantages = [0.0] * T
- running = 0.0
- for t in range(T - 1, -1, -1):
- running = deltas[t] + gamma * lam * running
- advantages[t] = running
- return advantages
diff --git a/vendor/nanoai/agentic/turn_weight.py b/vendor/nanoai/agentic/turn_weight.py
deleted file mode 100644
index 89a95b8518c50a9d1fea44f03da0ba10ca51ba01..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/turn_weight.py
+++ /dev/null
@@ -1,143 +0,0 @@
-"""Turn-level loss weighting for SFT/DFT/GFT.
-
-Provides weight functions that assign per-turn importance based on:
- 1. Heuristic: tool_ok, tool complexity, turn position
- 2. OCAW: pre-computed conditional advantage table
- 3. Probe: V(s_t) from trained value probe
-
-Usage in any trainer:
- weights = compute_turn_weights(rollout, method="heuristic")
- for turn_idx, (start, end) in enumerate(turn_spans):
- loss_mask[start:end] *= weights[turn_idx]
-"""
-from __future__ import annotations
-
-from typing import Optional
-
-
-def compute_turn_weights(
- rollout_dict: dict,
- *,
- method: str = "heuristic",
- ocaw_table: Optional[dict] = None,
- normalize: bool = True,
-) -> list[float]:
- """Compute per-turn loss weights for a rollout.
-
- rollout_dict: must have turn_spans with tool_name, tool_ok, is_terminal.
- method: "heuristic", "ocaw", or "uniform" (baseline).
- """
- turns = rollout_dict.get("turn_spans", [])
- if not turns:
- return []
-
- T = len(turns)
-
- if method == "uniform":
- return [1.0] * T
-
- if method == "heuristic":
- return _heuristic_weights(turns)
-
- if method == "ocaw":
- if ocaw_table is None:
- return _heuristic_weights(turns)
- return _ocaw_weights(rollout_dict, ocaw_table)
-
- raise ValueError(f"Unknown method: {method}")
-
-
-def _heuristic_weights(turns: list[dict]) -> list[float]:
- """Simple heuristic: weight by action complexity and outcome.
-
- Intuition:
- - Tool calls that succeed AND require multi-step reasoning (Edit, Bash) are
- the turns the model needs to learn most from
- - Read/Grep are easy (model already masters them at 95%+)
- - __no_tool__ turns are almost always bad in agentic tasks
- - Terminal turns carry the final decision weight
- """
- TOOL_COMPLEXITY = {
- "Edit": 2.0,
- "Write": 2.0,
- "Bash": 1.5,
- "Grep": 0.5,
- "Read": 0.3,
- None: 0.2,
- }
-
- weights = []
- T = len(turns)
- for i, ts in enumerate(turns):
- tool = ts.get("tool_name")
- ok = ts.get("tool_ok")
- is_terminal = ts.get("is_terminal", False)
-
- base = TOOL_COMPLEXITY.get(tool, 1.0)
-
- if ok is False:
- base *= 0.5
-
- if is_terminal:
- base *= 1.5
-
- if i == 0:
- base *= 1.2
-
- weights.append(max(base, 0.1))
-
- total = sum(weights)
- if total > 0:
- weights = [w / total * T for w in weights]
- return weights
-
-
-def _ocaw_weights(rollout_dict: dict, ocaw_table: dict) -> list[float]:
- """OCAW-based weights: |advantage| as importance signal."""
- import re
-
- task_id = rollout_dict.get("task_id", "")
- family = re.sub(r"(?:_v\d+)+$", "", task_id)
- turns = rollout_dict.get("turn_spans", [])
- T = len(turns)
-
- weights = []
- for ts in turns:
- turn_idx = ts.get("turn_idx", 0)
- tool = ts.get("tool_name") or "__no_tool__"
- key = f"{family}|{turn_idx}|{tool}"
- entry = ocaw_table.get(key)
- if entry:
- w = abs(entry["advantage"]) + 0.5
- else:
- w = 0.5
- weights.append(w)
-
- total = sum(weights)
- if total > 0:
- weights = [w / total * T for w in weights]
- return weights
-
-
-def apply_turn_weights_to_loss_mask(
- loss_mask,
- turn_spans: list[dict],
- weights: list[float],
- prefix_len: int = 0,
-):
- """Multiply loss_mask by per-turn weights in-place.
-
- loss_mask: 1D tensor of shape (seq_len,), values 0 or 1.
- turn_spans: list of dicts with start_idx, end_idx.
- weights: per-turn weights from compute_turn_weights().
- prefix_len: offset for token positions (e.g., if turn_spans are relative).
- """
- for ti, ts in enumerate(turn_spans):
- if ti >= len(weights):
- break
- start = ts.get("start_idx", 0) + prefix_len
- end = ts.get("end_idx", start) + prefix_len
- start = max(start - 1, 0)
- end = min(end - 1, loss_mask.size(0))
- if end > start:
- loss_mask[start:end] *= weights[ti]
diff --git a/vendor/nanoai/agentic/value_head.py b/vendor/nanoai/agentic/value_head.py
deleted file mode 100644
index 68c3f1c1c21407f095b0f133fe77b24d73cf9c68..0000000000000000000000000000000000000000
--- a/vendor/nanoai/agentic/value_head.py
+++ /dev/null
@@ -1,141 +0,0 @@
-"""Value Head for actor-critic training (Plan B).
-
-Attaches a lightweight value head to the policy model's final hidden state.
-The value head shares the transformer backbone with the policy (lm_head) and
-only adds a small MLP that maps the last-layer hidden state to a scalar V(s_t).
-
-Two integration patterns:
- 1. Hook-based (non-invasive): register a forward hook on the last transformer
- block, extract hidden states at turn boundaries, feed to the value head.
- Does NOT modify gpt.py. Used during rollout collection.
- 2. Inline (training): modify the forward pass to return both logits and values
- at specified positions. Requires a thin wrapper around GPT.forward().
-
-This module implements pattern 1 (hook-based) to avoid modifying gpt.py.
-"""
-from __future__ import annotations
-
-import torch
-import torch.nn as nn
-import torch.nn.functional as F
-
-
-class ValueHead(nn.Module):
- """MLP value head: final hidden state -> scalar V(s).
-
- Designed to be trained jointly with the policy via:
- loss = policy_loss + vf_coef * value_loss
- where value_loss = MSE(V_pred, V_target).
- """
-
- def __init__(self, n_embd: int = 384, hidden: int = 256):
- super().__init__()
- self.net = nn.Sequential(
- nn.LayerNorm(n_embd),
- nn.Linear(n_embd, hidden),
- nn.GELU(),
- nn.Linear(hidden, hidden // 2),
- nn.GELU(),
- nn.Linear(hidden // 2, 1),
- )
-
- def forward(self, h: torch.Tensor) -> torch.Tensor:
- """h: (..., n_embd) -> (...,)"""
- return self.net(h).squeeze(-1)
-
-
-class PolicyWithValueHead(nn.Module):
- """Wraps a GPT policy model with a value head for actor-critic training.
-
- The value head is attached via a forward hook on the last transformer block.
- During forward passes, hidden states at specified positions are captured and
- fed through the value head to produce V(s) estimates.
-
- This wrapper does NOT modify the policy's forward() method or its parameters.
- The value head has its own parameters that need separate optimizer groups.
- """
-
- def __init__(self, policy: nn.Module, n_embd: int = 384, hidden: int = 256):
- super().__init__()
- self.policy = policy
- self.value_head = ValueHead(n_embd=n_embd, hidden=hidden)
- self._captured_hidden = None
- self._hook_handle = None
-
- def attach_hook(self):
- """Register forward hook on the last transformer block."""
- if self._hook_handle is not None:
- return
- last_block = self.policy.transformer.h[-1]
-
- def hook_fn(module, inp, out):
- self._captured_hidden = out.detach()
-
- self._hook_handle = last_block.register_forward_hook(hook_fn)
-
- def detach_hook(self):
- if self._hook_handle is not None:
- self._hook_handle.remove()
- self._hook_handle = None
- self._captured_hidden = None
-
- def get_values_at_positions(self, positions: list[int]) -> torch.Tensor:
- """Extract value estimates at specified token positions.
-
- Must be called AFTER a forward pass through the policy (which triggers
- the hook and populates self._captured_hidden).
-
- positions: list of token indices to evaluate V(s) at.
- Returns: (len(positions),) tensor of value estimates.
- """
- if self._captured_hidden is None:
- raise RuntimeError("No captured hidden state — run a policy forward pass first")
- h = self._captured_hidden # (1, seq_len, n_embd)
- h_at_pos = h[0, positions, :].float() # (P, n_embd), cast to fp32
- return self.value_head(h_at_pos)
-
- def compute_value_loss(
- self,
- turn_boundary_positions: list[int],
- value_targets: torch.Tensor,
- ) -> torch.Tensor:
- """Compute MSE loss between predicted and target values at turn boundaries.
-
- turn_boundary_positions: token positions (start_idx - 1 for each turn)
- value_targets: (T,) tensor of V(s_t) targets (from MC estimation or GAE)
- """
- v_pred = self.get_values_at_positions(turn_boundary_positions)
- return F.mse_loss(v_pred, value_targets.to(v_pred.device).float())
-
- def value_head_parameters(self):
- """Return only the value head parameters (for separate optimizer group)."""
- return self.value_head.parameters()
-
- def setup_joint_optimizer(self, policy_optimizer, vf_lr: float = 1e-3):
- """Add value head parameters to an existing policy optimizer.
-
- Returns the modified optimizer with an additional param group for the
- value head.
- """
- policy_optimizer.add_param_group({
- "lr": vf_lr,
- "weight_decay": 0.01,
- "params": list(self.value_head_parameters()),
- "kind": "adamw",
- "betas": (0.9, 0.999),
- "eps": 1e-8,
- })
- return policy_optimizer
-
-
-def collect_turn_boundary_positions(turn_spans, prefix_len: int = 0) -> list[int]:
- """Extract token positions for turn boundaries from turn_spans.
-
- Returns positions suitable for get_values_at_positions().
- """
- positions = []
- for ts in turn_spans:
- start_idx = ts["start_idx"] if isinstance(ts, dict) else ts.start_idx
- pos = max(start_idx - 1 + prefix_len, 0)
- positions.append(pos)
- return positions
diff --git a/vendor/nanoai/ane/README.md b/vendor/nanoai/ane/README.md
deleted file mode 100644
index 4c13423951a3cadcbd04050d070ca810b154e8ad..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ane/README.md
+++ /dev/null
@@ -1,26 +0,0 @@
-# nanoai/ane — Apple Neural Engine inference
-
-Export a trained NanoAI model to Core ML and run it on the Apple Neural Engine
-for on-device inference. Inference-only (no training on ANE). Full write-up:
-[docs/apple_ane_inference.md](../../docs/apple_ane_inference.md).
-
-## Modules
-
-- `export.py` — convert a `nanoai` GPT checkpoint to a Core ML `.mlpackage`.
-- `model.py` — `NanochatANEModel`, the ANE-friendly model definition.
-- `package.py` — `.mlpackage` packaging helpers.
-- `runtime.py` — `ANEEngine`, the on-device inference runtime.
-- `__init__.py` exports: `ANEEngine`, `NanochatANEModel`, `compare_logits`.
-
-`compare_logits` is the parity check: ANE output vs the reference PyTorch model.
-
-## Dependencies
-
-`requirements-ane.txt` (`coremltools>=9.0`, `PyYAML>=6.0`). Install on macOS:
-
-```bash
-pip install -r requirements-ane.txt
-```
-
-Load trained models as **dense BF16** (NanoAI's dense-by-default semantics) when
-exporting. Tests: `tests/test_ane_*.py`.
diff --git a/vendor/nanoai/ane/__init__.py b/vendor/nanoai/ane/__init__.py
deleted file mode 100644
index 7df2e993c8108c2ddf7a136fa8e6fa171d258ed6..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ane/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""Apple Neural Engine export and runtime helpers for nanoai.
-
-Imports stay lazy so ``import nanoai.ane.model`` does not require CoreMLTools
-or tokenizer-only dependencies.
-"""
-
-__all__ = ["ANEEngine", "NanochatANEModel", "compare_logits"]
diff --git a/vendor/nanoai/ane/export.py b/vendor/nanoai/ane/export.py
deleted file mode 100644
index b8c498779471b01528d5ba70036e43a51e07b596..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ane/export.py
+++ /dev/null
@@ -1,394 +0,0 @@
-"""Export nanoai checkpoints to a CoreML/ANE package."""
-
-from __future__ import annotations
-
-import shutil
-import subprocess
-import logging
-from pathlib import Path
-from typing import Iterable
-
-import torch
-
-from nanoai.ane.model import (
- NanochatANEKVModel,
- NanochatANEModel,
- NanochatANEStatefulKVModel,
- compare_kv_decode_logits,
- compare_last_logits,
- compare_logits,
-)
-from nanoai.ane.package import normalize_context_buckets, write_metadata
-from nanoai.checkpoint_manager import find_largest_model, find_last_step, load_model
-from nanoai.common import get_base_dir
-
-
-_SOURCE_DIRS = {
- "base": "base_checkpoints",
- "sft": "chatsft_checkpoints",
- "rl": "chatrl_checkpoints",
- "dpo": "dpo_checkpoints",
- "grpo": "grpo_checkpoints",
-}
-
-
-def _copy_tokenizer(out_dir: Path) -> None:
- src = Path(get_base_dir()) / "tokenizer"
- if not src.exists():
- raise FileNotFoundError(f"Tokenizer directory not found: {src}")
- dst = out_dir / "tokenizer"
- if dst.exists():
- shutil.rmtree(dst)
- shutil.copytree(src, dst)
-
-
-def _compile_mlpackage(package_path: Path, out_dir: Path) -> Path:
- compiled_dir = out_dir / (package_path.stem + ".mlmodelc")
- if compiled_dir.exists():
- shutil.rmtree(compiled_dir)
- subprocess.run(
- ["xcrun", "coremlcompiler", "compile", str(package_path), str(out_dir)],
- check=True,
- )
- return compiled_dir
-
-
-def _resolve_checkpoint(source: str, model_tag: str | None, step: int | None) -> tuple[str, int]:
- checkpoints_dir = Path(get_base_dir()) / _SOURCE_DIRS[source]
- resolved_tag = model_tag or find_largest_model(str(checkpoints_dir))
- checkpoint_dir = checkpoints_dir / resolved_tag
- resolved_step = step if step is not None else find_last_step(str(checkpoint_dir))
- return resolved_tag, resolved_step
-
-
-def export_checkpoint(
- *,
- source: str,
- model_tag: str | None,
- step: int | None,
- out_dir: str | Path,
- context_length: int | None = None,
- context_buckets: Iterable[int] | None = None,
- skip_coreml: bool = False,
- compile_model: bool = False,
- parity_prompts: Iterable[str] = (),
- output_mode: str = "last",
- export_kv_trace: bool = False,
- export_kv_coreml: bool = False,
- export_kv_functional_coreml: bool = False,
-) -> dict:
- """Export a nanoai checkpoint package.
-
- ``skip_coreml`` still writes the traced TorchScript model and metadata. This
- keeps tests and non-macOS development useful even when CoreMLTools is absent.
- """
- out = Path(out_dir)
- out.mkdir(parents=True, exist_ok=True)
- if output_mode not in {"full", "last"}:
- raise ValueError("output_mode must be 'full' or 'last'")
- if export_kv_coreml and skip_coreml:
- raise ValueError("export_kv_coreml requires CoreML export; remove --skip-coreml")
- if export_kv_functional_coreml and skip_coreml:
- raise ValueError("export_kv_functional_coreml requires CoreML export; remove --skip-coreml")
- export_kv_trace = export_kv_trace or export_kv_coreml or export_kv_functional_coreml
- resolved_tag, resolved_step = _resolve_checkpoint(source, model_tag, step)
-
- model, tokenizer, meta = load_model(
- source,
- torch.device("cpu"),
- phase="eval",
- model_tag=resolved_tag,
- step=resolved_step,
- )
- model.eval()
- sequence_len = int(meta["model_config"]["sequence_len"])
- context_bucket_list = normalize_context_buckets(context_length, sequence_len, context_buckets)
- context = max(context_bucket_list)
-
- parity = []
- kv_parity = []
- for prompt in parity_prompts:
- ids = tokenizer.encode(prompt, prepend=tokenizer.get_bos_token_id())
- ids = ids[-context:]
- if len(ids) < 1:
- continue
- input_ids = torch.tensor([ids], dtype=torch.long)
- compare_fn = compare_last_logits if output_mode == "last" else compare_logits
- if len(ids) >= 2:
- parity.append({"prompt": prompt, **compare_fn(model, input_ids, context_length=context)})
- if export_kv_trace:
- kv_parity.append({"prompt": prompt, **compare_kv_decode_logits(model, input_ids, context_length=context)})
-
- full_context_exports = []
- for bucket in context_bucket_list:
- ane_model = NanochatANEModel(model, context_length=bucket, output_mode=output_mode).eval()
- example = torch.zeros((1, bucket), dtype=torch.long)
- example_last_pos = torch.tensor([bucket - 1], dtype=torch.long)
- trace_example = (example, example_last_pos) if output_mode == "last" else example
- traced = torch.jit.trace(ane_model, trace_example, strict=False)
- trace_name = (
- "nanochat_ane_trace.pt"
- if len(context_bucket_list) == 1
- else f"nanochat_ane_ctx{bucket}_trace.pt"
- )
- torchscript_path = out / trace_name
- traced.save(str(torchscript_path))
- full_context_exports.append(
- {
- "context_length": bucket,
- "torchscript_path": torchscript_path,
- "traced": traced,
- "example_shape": tuple(example.shape),
- "last_pos_shape": tuple(example_last_pos.shape),
- "coreml_path": None,
- "compiled_path": None,
- }
- )
-
- kv_torchscript_path = None
- kv_stateful_torchscript_path = None
- kv_state_shapes = None
- if export_kv_trace:
- kv_model = NanochatANEKVModel(model, context_length=context).eval()
- k_cache, v_cache, prev_embedding = kv_model.empty_cache(batch_size=1, device=torch.device("cpu"))
- kv_example = (
- torch.zeros((1, 1), dtype=torch.long),
- torch.zeros((1,), dtype=torch.long),
- k_cache,
- v_cache,
- prev_embedding,
- )
- kv_traced = torch.jit.trace(kv_model, kv_example, strict=False)
- kv_torchscript_path = out / "nanochat_ane_kv_trace.pt"
- kv_traced.save(str(kv_torchscript_path))
-
- stateful_kv_model = NanochatANEStatefulKVModel(model, context_length=context).eval()
- kv_state_shapes = {
- "k_cache": tuple(stateful_kv_model.k_cache.shape),
- "v_cache": tuple(stateful_kv_model.v_cache.shape),
- "prev_embedding": tuple(stateful_kv_model.prev_embedding.shape),
- }
- kv_stateful_traced = torch.jit.trace(
- stateful_kv_model,
- (torch.zeros((1, 1), dtype=torch.long), torch.zeros((1,), dtype=torch.long)),
- strict=False,
- check_trace=False,
- )
- kv_stateful_traced.k_cache.zero_()
- kv_stateful_traced.v_cache.zero_()
- kv_stateful_traced.prev_embedding.zero_()
- kv_stateful_torchscript_path = out / "nanochat_ane_kv_stateful_trace.pt"
- kv_stateful_traced.save(str(kv_stateful_torchscript_path))
-
- coreml_path = None
- compiled_path = None
- kv_coreml_path = None
- kv_compiled_path = None
- kv_functional_coreml_path = None
- kv_functional_compiled_path = None
- if not skip_coreml:
- try:
- import coremltools as ct
- import numpy as np
- except ImportError as e:
- raise RuntimeError(
- "CoreML export requires coremltools. Install with `pip install -r requirements-ane.txt`."
- ) from e
- logging.getLogger("coremltools").setLevel(logging.WARNING)
-
- target = getattr(ct.target, "macOS15", None)
-
- def apply_common_convert_options(kwargs: dict) -> dict:
- if target is not None:
- kwargs["minimum_deployment_target"] = target
- if hasattr(ct, "precision"):
- kwargs["compute_precision"] = ct.precision.FLOAT16
- return kwargs
-
- for bucket_export in full_context_exports:
- bucket = bucket_export["context_length"]
- convert_kwargs = {
- "inputs": [ct.TensorType(name="input_ids", shape=bucket_export["example_shape"], dtype=np.int32)],
- "outputs": [ct.TensorType(name="logits")],
- "convert_to": "mlprogram",
- }
- if output_mode == "last":
- convert_kwargs["inputs"].append(
- ct.TensorType(name="last_pos", shape=bucket_export["last_pos_shape"], dtype=np.int32)
- )
- mlmodel = ct.convert(bucket_export["traced"], **apply_common_convert_options(convert_kwargs))
- coreml_name = "nanoai.mlpackage" if len(context_bucket_list) == 1 else f"nanochat_ctx{bucket}.mlpackage"
- coreml_path = out / coreml_name
- if coreml_path.exists():
- shutil.rmtree(coreml_path)
- mlmodel.save(str(coreml_path))
- bucket_export["coreml_path"] = coreml_path
- if compile_model:
- bucket_export["compiled_path"] = _compile_mlpackage(coreml_path, out)
- max_context_export = full_context_exports[-1]
- coreml_path = max_context_export["coreml_path"]
- compiled_path = max_context_export["compiled_path"]
-
- if export_kv_coreml:
- assert kv_stateful_torchscript_path is not None
- assert kv_state_shapes is not None
- kv_stateful_traced = torch.jit.load(str(kv_stateful_torchscript_path))
- kv_convert_kwargs = {
- "inputs": [
- ct.TensorType(name="input_ids", shape=(1, 1), dtype=np.int32),
- ct.TensorType(name="cache_pos", shape=(1,), dtype=np.int32),
- ],
- "outputs": [ct.TensorType(name="logits")],
- "states": [
- ct.StateType(
- wrapped_type=ct.TensorType(shape=kv_state_shapes["k_cache"], dtype=np.float16),
- name="k_cache",
- ),
- ct.StateType(
- wrapped_type=ct.TensorType(shape=kv_state_shapes["v_cache"], dtype=np.float16),
- name="v_cache",
- ),
- ct.StateType(
- wrapped_type=ct.TensorType(shape=kv_state_shapes["prev_embedding"], dtype=np.float16),
- name="prev_embedding",
- ),
- ],
- "convert_to": "mlprogram",
- }
- kv_mlmodel = ct.convert(kv_stateful_traced, **apply_common_convert_options(kv_convert_kwargs))
- kv_coreml_path = out / "nanochat_kv.mlpackage"
- if kv_coreml_path.exists():
- shutil.rmtree(kv_coreml_path)
- kv_mlmodel.save(str(kv_coreml_path))
- if compile_model:
- kv_compiled_path = _compile_mlpackage(kv_coreml_path, out)
-
- if export_kv_functional_coreml:
- assert kv_torchscript_path is not None
- kv_traced = torch.jit.load(str(kv_torchscript_path))
- k_shape = tuple(k_cache.shape)
- v_shape = tuple(v_cache.shape)
- prev_shape = tuple(prev_embedding.shape)
- kv_functional_convert_kwargs = {
- "inputs": [
- ct.TensorType(name="input_ids", shape=(1, 1), dtype=np.int32),
- ct.TensorType(name="cache_pos", shape=(1,), dtype=np.int32),
- ct.TensorType(name="k_cache", shape=k_shape, dtype=np.float16),
- ct.TensorType(name="v_cache", shape=v_shape, dtype=np.float16),
- ct.TensorType(name="prev_embedding", shape=prev_shape, dtype=np.float16),
- ],
- "outputs": [
- ct.TensorType(name="logits"),
- ct.TensorType(name="k_cache_out"),
- ct.TensorType(name="v_cache_out"),
- ct.TensorType(name="prev_embedding_out"),
- ],
- "convert_to": "mlprogram",
- }
- kv_functional_mlmodel = ct.convert(
- kv_traced,
- **apply_common_convert_options(kv_functional_convert_kwargs),
- )
- kv_functional_coreml_path = out / "nanochat_kv_functional.mlpackage"
- if kv_functional_coreml_path.exists():
- shutil.rmtree(kv_functional_coreml_path)
- kv_functional_mlmodel.save(str(kv_functional_coreml_path))
- if compile_model:
- kv_functional_compiled_path = _compile_mlpackage(kv_functional_coreml_path, out)
-
- _copy_tokenizer(out)
-
- metadata = {
- "format": "nanoai-ane-package",
- "format_version": 1,
- "backend": "coreml-full-context",
- "checkpoint": {
- "source": source,
- "model_tag": resolved_tag,
- "step": resolved_step,
- },
- "model_config": dict(meta["model_config"]),
- "runtime": {
- "context_length": context,
- "input_name": "input_ids",
- "position_name": "last_pos" if output_mode == "last" else None,
- "output_name": "logits",
- "output_mode": output_mode,
- "functional_kv_min_context_length": (
- context_bucket_list[-2] + 1
- if export_kv_functional_coreml and len(context_bucket_list) > 1
- else 1
- ),
- "context_buckets": [
- {
- "context_length": bucket_export["context_length"],
- "input_name": "input_ids",
- "position_name": "last_pos" if output_mode == "last" else None,
- "output_name": "logits",
- "output_mode": output_mode,
- "torchscript": bucket_export["torchscript_path"].name,
- "coreml": (
- bucket_export["coreml_path"].name
- if bucket_export["coreml_path"] is not None
- else None
- ),
- "compiled_coreml": (
- bucket_export["compiled_path"].name
- if bucket_export["compiled_path"] is not None
- else None
- ),
- }
- for bucket_export in full_context_exports
- ],
- "sampling_on_host": True,
- "kv_cache": export_kv_coreml or export_kv_functional_coreml,
- "kv_mode": (
- "stateful"
- if export_kv_coreml
- else "functional"
- if export_kv_functional_coreml
- else "none"
- ),
- "kv_cache_trace": export_kv_trace,
- "kv_input_name": "input_ids",
- "kv_position_name": "cache_pos",
- "kv_output_name": "logits",
- "kv_cache_input_names": {
- "k": "k_cache",
- "v": "v_cache",
- "prev_embedding": "prev_embedding",
- },
- "kv_cache_output_names": {
- "k": "k_cache_out",
- "v": "v_cache_out",
- "prev_embedding": "prev_embedding_out",
- },
- "kv_cache_shapes": {
- "k": list(k_cache.shape) if export_kv_trace else None,
- "v": list(v_cache.shape) if export_kv_trace else None,
- "prev_embedding": list(prev_embedding.shape) if export_kv_trace else None,
- },
- },
- "artifacts": {
- "torchscript": full_context_exports[-1]["torchscript_path"].name,
- "kv_torchscript": kv_torchscript_path.name if kv_torchscript_path is not None else None,
- "kv_stateful_torchscript": (
- kv_stateful_torchscript_path.name if kv_stateful_torchscript_path is not None else None
- ),
- "coreml": coreml_path.name if coreml_path is not None else None,
- "compiled_coreml": compiled_path.name if compiled_path is not None else None,
- "kv_coreml": kv_coreml_path.name if kv_coreml_path is not None else None,
- "compiled_kv_coreml": kv_compiled_path.name if kv_compiled_path is not None else None,
- "kv_functional_coreml": (
- kv_functional_coreml_path.name if kv_functional_coreml_path is not None else None
- ),
- "compiled_kv_functional_coreml": (
- kv_functional_compiled_path.name if kv_functional_compiled_path is not None else None
- ),
- "tokenizer": "tokenizer",
- },
- "parity": parity,
- "kv_parity": kv_parity,
- }
- write_metadata(out, metadata)
- return metadata
diff --git a/vendor/nanoai/ane/model.py b/vendor/nanoai/ane/model.py
deleted file mode 100644
index 510a1e117cac5968ff037a715d787d0ef9c180dd..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ane/model.py
+++ /dev/null
@@ -1,446 +0,0 @@
-"""CoreML-friendly nanoai GPT wrapper.
-
-The normal ``nanoai.gpt.GPT`` forward path calls the repo's FlashAttention
-adapter. That is the right path for CUDA/MPS PyTorch, but it is not a good
-conversion surface for CoreML. ``NanochatANEModel`` mirrors the no-KV full
-context inference semantics using plain tensor operations that CoreMLTools can
-trace and lower more predictably.
-"""
-
-from __future__ import annotations
-
-import math
-from typing import Any
-
-import torch
-import torch.nn as nn
-import torch.nn.functional as F
-
-from nanoai.gpt import GPT, apply_rotary_emb
-
-
-def _rms_norm(x: torch.Tensor) -> torch.Tensor:
- eps = torch.finfo(x.dtype).eps
- return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
-
-
-def _repeat_kv(x: torch.Tensor, repeats: int) -> torch.Tensor:
- if repeats == 1:
- return x
- return x.repeat_interleave(repeats, dim=2)
-
-
-def _causal_window_mask(t: int, window: int, device: torch.device) -> torch.Tensor:
- rows = torch.arange(t, device=device).unsqueeze(1)
- cols = torch.arange(t, device=device).unsqueeze(0)
- mask = cols <= rows
- if window >= 0:
- mask = mask & ((rows - cols) <= window)
- return mask
-
-
-class NanochatANEModel(nn.Module):
- """Traceable full-context nanoai model for CoreML export.
-
- This v1 wrapper deliberately does not implement KV-cache decode. Runtime
- generation re-runs the model over the current context window at each token.
- That is inefficient, but it is a concrete Apple Neural Engine baseline and
- keeps the first conversion target faithful to nanoai's custom architecture.
- """
-
- def __init__(self, model: GPT, context_length: int | None = None, output_mode: str = "full"):
- super().__init__()
- if output_mode not in {"full", "last"}:
- raise ValueError("output_mode must be 'full' or 'last'")
- self.model = model
- self.config = model.config
- self.context_length = int(context_length or model.config.sequence_len)
- self.output_mode = output_mode
- if self.context_length > model.cos.size(1):
- raise ValueError(
- f"context_length={self.context_length} exceeds rotary cache "
- f"length {model.cos.size(1)}"
- )
-
- def forward(self, input_ids: torch.Tensor, last_pos: torch.Tensor | None = None) -> torch.Tensor:
- input_ids = input_ids.to(torch.long)
- bsz, seq_len = input_ids.shape
- if not torch.jit.is_tracing() and seq_len > self.context_length:
- raise ValueError(f"input seq_len {seq_len} exceeds context_length {self.context_length}")
-
- m = self.model
- cfg = self.config
- cos_sin = (
- m.cos[:, :seq_len].to(device=input_ids.device),
- m.sin[:, :seq_len].to(device=input_ids.device),
- )
-
- x = m.transformer.wte(input_ids)
- x = _rms_norm(x)
-
- gate = m.smear_lambda.to(x.dtype) * torch.sigmoid(m.smear_gate(x[:, 1:, :24]))
- x = torch.cat([x[:, :1], x[:, 1:] + gate * x[:, :-1]], dim=1)
-
- x0 = x
- backout_layer = cfg.n_layer // 2
- x_backout = None
- for layer_idx, block in enumerate(m.transformer.h):
- x = m.resid_lambdas[layer_idx].to(x.dtype) * x + m.x0_lambdas[layer_idx].to(x.dtype) * x0
- ve_key = m.ve_tid.get(layer_idx)
- ve = m.value_embeds[ve_key](input_ids).to(x.dtype) if ve_key is not None else None
- x = x + self._attention(block.attn, _rms_norm(x), ve, cos_sin, m.window_sizes[layer_idx])
- x = x + block.mlp(_rms_norm(x))
- if layer_idx == backout_layer:
- x_backout = x
-
- if x_backout is not None:
- x = x - m.backout_lambda.to(x.dtype) * x_backout
- x = _rms_norm(x)
-
- if self.output_mode == "last":
- if last_pos is None:
- raise ValueError("last_pos is required when output_mode='last'")
- last_pos = last_pos.to(torch.long).view(-1, 1, 1)
- last_pos = last_pos.expand(bsz, 1, x.size(-1))
- x = torch.gather(x, dim=1, index=last_pos).squeeze(1)
-
- softcap = 15.0
- logits = m.lm_head(x)[..., : cfg.vocab_size]
- logits = logits.float()
- return softcap * torch.tanh(logits / softcap)
-
- def _attention(
- self,
- attn: nn.Module,
- x: torch.Tensor,
- ve: torch.Tensor | None,
- cos_sin: tuple[torch.Tensor, torch.Tensor],
- window_size: tuple[int, int],
- ) -> torch.Tensor:
- bsz, seq_len, _ = x.shape
- head_dim = attn.head_dim
-
- q = attn.c_q(x).view(bsz, seq_len, attn.n_head, head_dim)
- k = attn.c_k(x).view(bsz, seq_len, attn.n_kv_head, head_dim)
- v = attn.c_v(x).view(bsz, seq_len, attn.n_kv_head, head_dim)
-
- if ve is not None:
- ve = ve.view(bsz, seq_len, attn.n_kv_head, head_dim)
- gate = 3 * torch.sigmoid(attn.ve_gate(x[..., : attn.ve_gate_channels]))
- v = v + gate.unsqueeze(-1) * ve
-
- q = apply_rotary_emb(q, *cos_sin)
- k = apply_rotary_emb(k, *cos_sin)
- q = _rms_norm(q) * 1.2
- k = _rms_norm(k) * 1.2
-
- repeats = attn.n_head // attn.n_kv_head
- k = _repeat_kv(k, repeats)
- v = _repeat_kv(v, repeats)
-
- q = q.transpose(1, 2)
- k = k.transpose(1, 2)
- v = v.transpose(1, 2)
- scores = torch.matmul(q, k.transpose(-1, -2)) * (1.0 / math.sqrt(head_dim))
- mask = _causal_window_mask(seq_len, int(window_size[0]), x.device)
- scores = scores.masked_fill(~mask.view(1, 1, seq_len, seq_len), -10000.0)
- weights = F.softmax(scores, dim=-1)
- y = torch.matmul(weights, v)
- y = y.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
- return attn.c_proj(y)
-
-
-class NanochatANEKVModel(nn.Module):
- """Traceable single-token KV-cache decode wrapper for nanoai.
-
- This mirrors ``GPT(..., kv_cache=...)`` semantics without relying on the
- FlashAttention cache kernel. It is a staging target for ANE cache export:
- inputs and outputs are plain tensors, and cache updates are functional so
- they can be traced.
- """
-
- def __init__(self, model: GPT, context_length: int | None = None):
- super().__init__()
- self.model = model
- self.config = model.config
- self.context_length = int(context_length or model.config.sequence_len)
- if self.context_length > model.cos.size(1):
- raise ValueError(
- f"context_length={self.context_length} exceeds rotary cache "
- f"length {model.cos.size(1)}"
- )
-
- def empty_cache(
- self,
- batch_size: int = 1,
- device: torch.device | str | None = None,
- dtype: torch.dtype | None = None,
- ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
- """Return zeroed ``(k_cache, v_cache, prev_embedding)`` tensors."""
- m = self.model
- cfg = self.config
- device = device or m.transformer.wte.weight.device
- dtype = dtype or m.transformer.wte.weight.dtype
- head_dim = cfg.n_embd // cfg.n_head
- cache_shape = (cfg.n_layer, batch_size, self.context_length, cfg.n_kv_head, head_dim)
- k_cache = torch.zeros(cache_shape, device=device, dtype=dtype)
- v_cache = torch.zeros(cache_shape, device=device, dtype=dtype)
- prev_embedding = torch.zeros((batch_size, 1, cfg.n_embd), device=device, dtype=dtype)
- return k_cache, v_cache, prev_embedding
-
- def forward(
- self,
- input_ids: torch.Tensor,
- cache_pos: torch.Tensor,
- k_cache: torch.Tensor,
- v_cache: torch.Tensor,
- prev_embedding: torch.Tensor,
- ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
- input_ids = input_ids.to(torch.long)
- bsz, seq_len = input_ids.shape
- if not torch.jit.is_tracing():
- if seq_len != 1:
- raise ValueError("NanochatANEKVModel expects exactly one token per forward")
- if k_cache.size(2) != self.context_length or v_cache.size(2) != self.context_length:
- raise ValueError("cache tensors do not match context_length")
-
- position = cache_pos.to(torch.long).view(-1)[:1]
- m = self.model
- cfg = self.config
- cos_sin = (
- torch.index_select(m.cos.to(device=input_ids.device), 1, position),
- torch.index_select(m.sin.to(device=input_ids.device), 1, position),
- )
-
- x = m.transformer.wte(input_ids)
- x = _rms_norm(x)
-
- # GPT KV decode stores the normalized current embedding before smear and
- # uses it as the previous-token embedding on the next decode step.
- next_prev_embedding = x
- has_prev = (position > 0).to(dtype=x.dtype).view(1, 1, 1)
- gate = m.smear_lambda.to(x.dtype) * torch.sigmoid(m.smear_gate(x[:, :, :24]))
- x = x + has_prev * gate * prev_embedding.to(x.dtype)
-
- x0 = x
- backout_layer = cfg.n_layer // 2
- x_backout = None
- new_k_layers = []
- new_v_layers = []
- for layer_idx, block in enumerate(m.transformer.h):
- x = m.resid_lambdas[layer_idx].to(x.dtype) * x + m.x0_lambdas[layer_idx].to(x.dtype) * x0
- ve_key = m.ve_tid.get(layer_idx)
- ve = m.value_embeds[ve_key](input_ids).to(x.dtype) if ve_key is not None else None
- attn_out, new_k_layer, new_v_layer = self._attention_decode(
- block.attn,
- _rms_norm(x),
- ve,
- cos_sin,
- m.window_sizes[layer_idx],
- position,
- k_cache[layer_idx],
- v_cache[layer_idx],
- )
- new_k_layers.append(new_k_layer)
- new_v_layers.append(new_v_layer)
- x = x + attn_out
- x = x + block.mlp(_rms_norm(x))
- if layer_idx == backout_layer:
- x_backout = x
-
- if x_backout is not None:
- x = x - m.backout_lambda.to(x.dtype) * x_backout
- x = _rms_norm(x)
-
- softcap = 15.0
- logits = m.lm_head(x).squeeze(1)[..., : cfg.vocab_size]
- logits = logits.float()
- logits = softcap * torch.tanh(logits / softcap)
- return logits, torch.stack(new_k_layers, dim=0), torch.stack(new_v_layers, dim=0), next_prev_embedding
-
- def _attention_decode(
- self,
- attn: nn.Module,
- x: torch.Tensor,
- ve: torch.Tensor | None,
- cos_sin: tuple[torch.Tensor, torch.Tensor],
- window_size: tuple[int, int],
- position: torch.Tensor,
- k_cache_layer: torch.Tensor,
- v_cache_layer: torch.Tensor,
- ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
- bsz, seq_len, _ = x.shape
- head_dim = attn.head_dim
-
- q = attn.c_q(x).view(bsz, seq_len, attn.n_head, head_dim)
- k = attn.c_k(x).view(bsz, seq_len, attn.n_kv_head, head_dim)
- v = attn.c_v(x).view(bsz, seq_len, attn.n_kv_head, head_dim)
-
- if ve is not None:
- ve = ve.view(bsz, seq_len, attn.n_kv_head, head_dim)
- gate = 3 * torch.sigmoid(attn.ve_gate(x[..., : attn.ve_gate_channels]))
- v = v + gate.unsqueeze(-1) * ve
-
- q = apply_rotary_emb(q, *cos_sin)
- k = apply_rotary_emb(k, *cos_sin)
- q = _rms_norm(q) * 1.2
- k = _rms_norm(k) * 1.2
-
- one_hot = F.one_hot(position, num_classes=self.context_length).to(dtype=k.dtype)
- one_hot = one_hot.view(1, self.context_length, 1, 1)
- k_cache_layer = k_cache_layer.to(k.dtype)
- v_cache_layer = v_cache_layer.to(v.dtype)
- new_k_cache_layer = k_cache_layer * (1 - one_hot) + k * one_hot
- new_v_cache_layer = v_cache_layer * (1 - one_hot) + v * one_hot
-
- repeats = attn.n_head // attn.n_kv_head
- k_all = _repeat_kv(new_k_cache_layer, repeats)
- v_all = _repeat_kv(new_v_cache_layer, repeats)
-
- q = q.transpose(1, 2)
- k_all = k_all.transpose(1, 2)
- v_all = v_all.transpose(1, 2)
- scores = torch.matmul(q, k_all.transpose(-1, -2)) * (1.0 / math.sqrt(head_dim))
-
- cols = torch.arange(self.context_length, device=x.device)
- allowed = cols <= position
- window = int(window_size[0])
- if window >= 0 and window < self.context_length:
- allowed = allowed & ((position - cols) <= window)
- scores = scores.masked_fill(~allowed.view(1, 1, 1, self.context_length), -10000.0)
- weights = F.softmax(scores, dim=-1)
- y = torch.matmul(weights, v_all)
- y = y.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
- return attn.c_proj(y), new_k_cache_layer, new_v_cache_layer
-
-
-class NanochatANEStatefulKVModel(nn.Module):
- """Stateful single-token KV decode wrapper for CoreML ``StateType`` export."""
-
- def __init__(
- self,
- model: GPT,
- context_length: int | None = None,
- batch_size: int = 1,
- cache_dtype: torch.dtype | None = None,
- ):
- super().__init__()
- if batch_size != 1:
- raise ValueError("NanochatANEStatefulKVModel currently supports batch_size=1 only")
- self.decoder = NanochatANEKVModel(model, context_length=context_length)
- k_cache, v_cache, prev_embedding = self.decoder.empty_cache(
- batch_size=batch_size,
- dtype=cache_dtype,
- )
- self.register_buffer("k_cache", k_cache.squeeze(1))
- self.register_buffer("v_cache", v_cache.squeeze(1))
- self.register_buffer("prev_embedding", prev_embedding)
-
- def reset_cache(self) -> None:
- self.k_cache.zero_()
- self.v_cache.zero_()
- self.prev_embedding.zero_()
-
- def forward(self, input_ids: torch.Tensor, cache_pos: torch.Tensor) -> torch.Tensor:
- logits, k_cache, v_cache, prev_embedding = self.decoder(
- input_ids,
- cache_pos,
- self.k_cache.unsqueeze(1),
- self.v_cache.unsqueeze(1),
- self.prev_embedding,
- )
- self.k_cache.mul_(0)
- self.k_cache.add_(k_cache.squeeze(1))
- self.v_cache.mul_(0)
- self.v_cache.add_(v_cache.squeeze(1))
- self.prev_embedding.mul_(0)
- self.prev_embedding.add_(prev_embedding)
- return logits
-
-
-@torch.inference_mode()
-def compare_logits(model: GPT, input_ids: torch.Tensor, context_length: int | None = None) -> dict[str, Any]:
- """Compare PyTorch GPT logits with the CoreML-friendly wrapper logits."""
- was_training = model.training
- model.eval()
- wrapper = NanochatANEModel(model, context_length=context_length).eval()
- ref = model(input_ids)
- got = wrapper(input_ids)
- diff = (ref - got).abs()
- top_ref = ref[:, -1, :].argmax(dim=-1)
- top_got = got[:, -1, :].argmax(dim=-1)
- if was_training:
- model.train()
- return {
- "max_abs": float(diff.max().item()),
- "mean_abs": float(diff.mean().item()),
- "last_token_top1_match": bool(torch.equal(top_ref, top_got)),
- }
-
-
-@torch.inference_mode()
-def compare_last_logits(model: GPT, input_ids: torch.Tensor, context_length: int | None = None) -> dict[str, Any]:
- """Compare GPT last-token logits with the last-token ANE wrapper output."""
- was_training = model.training
- model.eval()
- wrapper = NanochatANEModel(model, context_length=context_length, output_mode="last").eval()
- last_pos = torch.full((input_ids.size(0),), input_ids.size(1) - 1, dtype=torch.long, device=input_ids.device)
- ref = model(input_ids)[:, -1, :]
- got = wrapper(input_ids, last_pos)
- diff = (ref - got).abs()
- top_ref = ref.argmax(dim=-1)
- top_got = got.argmax(dim=-1)
- if was_training:
- model.train()
- return {
- "max_abs": float(diff.max().item()),
- "mean_abs": float(diff.mean().item()),
- "last_token_top1_match": bool(torch.equal(top_ref, top_got)),
- }
-
-
-@torch.inference_mode()
-def compare_kv_decode_logits(
- model: GPT,
- input_ids: torch.Tensor,
- context_length: int | None = None,
-) -> dict[str, Any]:
- """Compare single-token KV decode logits with full-context last logits."""
- was_training = model.training
- model.eval()
- context = int(context_length or model.config.sequence_len)
- if input_ids.size(1) > context:
- raise ValueError("input_ids length cannot exceed context_length")
- full = NanochatANEModel(model, context_length=context, output_mode="last").eval()
- kv = NanochatANEKVModel(model, context_length=context).eval()
- k_cache, v_cache, prev_embedding = kv.empty_cache(
- batch_size=input_ids.size(0),
- device=input_ids.device,
- dtype=model.transformer.wte.weight.dtype,
- )
-
- max_abs = 0.0
- mean_abs_total = 0.0
- top1_match = True
- for pos in range(input_ids.size(1)):
- pos_tensor = torch.tensor([pos], dtype=torch.long, device=input_ids.device)
- got, k_cache, v_cache, prev_embedding = kv(
- input_ids[:, pos : pos + 1],
- pos_tensor,
- k_cache,
- v_cache,
- prev_embedding,
- )
- ref = full(input_ids[:, : pos + 1], pos_tensor)
- diff = (ref - got).abs()
- max_abs = max(max_abs, float(diff.max().item()))
- mean_abs_total += float(diff.mean().item())
- top1_match = top1_match and bool(torch.equal(ref.argmax(dim=-1), got.argmax(dim=-1)))
-
- if was_training:
- model.train()
- return {
- "max_abs": max_abs,
- "mean_abs": mean_abs_total / max(input_ids.size(1), 1),
- "last_token_top1_match": top1_match,
- }
diff --git a/vendor/nanoai/ane/package.py b/vendor/nanoai/ane/package.py
deleted file mode 100644
index 2213048432bb48ca7e8b780ae81a85f241822cc2..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ane/package.py
+++ /dev/null
@@ -1,69 +0,0 @@
-"""Packaging helpers for nanoai ANE exports."""
-
-from __future__ import annotations
-
-import json
-from pathlib import Path
-from typing import Any, Iterable
-
-
-def normalize_context_buckets(
- context_length: int | None,
- sequence_len: int,
- context_buckets: Iterable[int] | None,
-) -> list[int]:
- if context_buckets is None:
- buckets = [int(context_length or sequence_len)]
- else:
- buckets = [int(bucket) for bucket in context_buckets]
- if context_length is not None:
- buckets.append(int(context_length))
- if not buckets:
- raise ValueError("context_buckets must not be empty")
- if any(bucket <= 0 for bucket in buckets):
- raise ValueError("context_buckets must be positive")
- if any(bucket > sequence_len for bucket in buckets):
- raise ValueError("context_length cannot exceed checkpoint sequence_len")
- return sorted(set(buckets))
-
-
-def write_metadata(package_dir: str | Path, metadata: dict[str, Any]) -> None:
- """Write metadata as JSON and as YAML-compatible JSON text."""
- out = Path(package_dir)
- out.mkdir(parents=True, exist_ok=True)
- text = json.dumps(metadata, indent=2, sort_keys=True)
- (out / "meta.json").write_text(text + "\n", encoding="utf-8")
- # JSON is valid YAML 1.2 and avoids making PyYAML a hard dependency.
- (out / "meta.yaml").write_text(text + "\n", encoding="utf-8")
-
-
-def load_metadata(package_dir: str | Path) -> dict[str, Any]:
- package = Path(package_dir)
- meta_json = package / "meta.json"
- if meta_json.exists():
- return json.loads(meta_json.read_text(encoding="utf-8"))
- meta_yaml = package / "meta.yaml"
- if not meta_yaml.exists():
- raise FileNotFoundError(f"No meta.json or meta.yaml found in {package}")
- text = meta_yaml.read_text(encoding="utf-8")
- try:
- return json.loads(text)
- except json.JSONDecodeError:
- try:
- import yaml
- except ImportError as e:
- raise RuntimeError("Install PyYAML to read non-JSON meta.yaml files") from e
- return yaml.safe_load(text)
-
-
-def load_tokenizer_from_package(package_dir: str | Path):
- tokenizer_dir = Path(package_dir) / "tokenizer"
- if not tokenizer_dir.exists():
- raise FileNotFoundError(f"Tokenizer directory not found: {tokenizer_dir}")
- if (tokenizer_dir / "tokenizer_config.json").exists() and not (tokenizer_dir / "tokenizer.pkl").exists():
- from nanoai.agentic.tokenizer_qwen import QwenTokenizer
-
- return QwenTokenizer.from_directory(str(tokenizer_dir))
- from nanoai.tokenizer import RustBPETokenizer
-
- return RustBPETokenizer.from_directory(str(tokenizer_dir))
diff --git a/vendor/nanoai/ane/runtime.py b/vendor/nanoai/ane/runtime.py
deleted file mode 100644
index af6c2003bdec06e30771ae43b0cf2f47686af4a4..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ane/runtime.py
+++ /dev/null
@@ -1,443 +0,0 @@
-"""Python runtime for nanoai ANE packages."""
-
-from __future__ import annotations
-
-from collections import deque
-import signal
-import warnings
-from contextlib import contextmanager
-from pathlib import Path
-
-import numpy as np
-import torch
-import torch.nn.functional as F
-
-from nanoai.ane.package import load_metadata, load_tokenizer_from_package
-
-
-@contextmanager
-def _timeout(duration: int, formula: str):
- def timeout_handler(_signum, _frame):
- raise Exception(f"{formula!r}: timed out after {duration} seconds")
-
- signal.signal(signal.SIGALRM, timeout_handler)
- signal.alarm(duration)
- try:
- yield
- finally:
- signal.alarm(0)
-
-
-def _eval_with_timeout(formula: str, max_time: int = 3):
- try:
- with _timeout(max_time, formula):
- with warnings.catch_warnings():
- warnings.simplefilter("ignore", SyntaxWarning)
- return eval(formula, {"__builtins__": {}}, {})
- except Exception:
- return None
-
-
-def _use_calculator(expr: str):
- expr = expr.replace(",", "")
- if all(x in "0123456789*+-/.() " for x in expr):
- if "**" in expr:
- return None
- return _eval_with_timeout(expr)
-
- allowed_chars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789'\"()._ "
- if not all(x in allowed_chars for x in expr):
- return None
- dangerous = ["__", "import", "exec", "eval", "compile", "open", "file", "input", "globals", "locals"]
- if any(pattern in expr.lower() for pattern in dangerous):
- return None
- if ".count(" not in expr:
- return None
- return _eval_with_timeout(expr)
-
-
-class RowState:
- def __init__(self, current_tokens=None):
- self.current_tokens = current_tokens or []
- self.forced_tokens = deque()
- self.in_python_block = False
- self.python_expr_tokens = []
- self.completed = False
- self.kv_state = None
- self.kv_processed = 0
- self.kv_window_start = 0
- self.kv_last_logits = None
- self.kv_k_cache = None
- self.kv_v_cache = None
- self.kv_prev_embedding = None
-
-
-def _sample_logits(logits: torch.Tensor, rng: torch.Generator, temperature: float, top_k: int | None) -> int:
- if temperature == 0.0:
- return int(torch.argmax(logits).item())
- if top_k is not None and top_k > 0:
- k = min(int(top_k), logits.numel())
- vals, idx = torch.topk(logits, k)
- probs = F.softmax(vals / temperature, dim=-1)
- choice = torch.multinomial(probs, num_samples=1, generator=rng)
- return int(idx[choice].item())
- probs = F.softmax(logits / temperature, dim=-1)
- return int(torch.multinomial(probs, num_samples=1, generator=rng).item())
-
-
-def _select_context_bucket(context_models: list[dict], token_count: int) -> dict:
- if not context_models:
- raise RuntimeError("no CoreML context models are loaded")
- target = max(1, int(token_count))
- for entry in context_models:
- if target <= int(entry["context_length"]):
- return entry
- return context_models[-1]
-
-
-class ANEEngine:
- """CoreML generation engine for nanoai ANE packages.
-
- The default path re-runs a fixed context window per token. Packages with KV
- artifacts can instead use either CoreML state or host-managed functional KV
- cache, depending on what CoreML can load on the current machine.
- """
-
- def __init__(self, package_dir: str | Path, compute_units: str = "CPU_AND_NE"):
- self.package_dir = Path(package_dir)
- self.meta = load_metadata(self.package_dir)
- self.tokenizer = load_tokenizer_from_package(self.package_dir)
- self.context_length = int(self.meta["runtime"]["context_length"])
- self.input_name = self.meta["runtime"].get("input_name", "input_ids")
- self.position_name = self.meta["runtime"].get("position_name", "last_pos")
- self.output_name = self.meta["runtime"].get("output_name", "logits")
- self.output_mode = self.meta["runtime"].get("output_mode", "full")
- self.kv_input_name = self.meta["runtime"].get("kv_input_name", "input_ids")
- self.kv_position_name = self.meta["runtime"].get("kv_position_name", "cache_pos")
- self.kv_output_name = self.meta["runtime"].get("kv_output_name", "logits")
- self.kv_mode = self.meta["runtime"].get("kv_mode", "stateful" if self.meta["runtime"].get("kv_cache") else "none")
- self.kv_cache_input_names = self.meta["runtime"].get("kv_cache_input_names", {})
- self.kv_cache_output_names = self.meta["runtime"].get("kv_cache_output_names", {})
- self.kv_cache_shapes = self.meta["runtime"].get("kv_cache_shapes", {})
- self.functional_kv_min_context_length = int(self.meta["runtime"].get("functional_kv_min_context_length", 1))
- self.last_context_length = self.context_length
- self.last_runtime_mode = "full"
-
- coreml_name = self.meta["artifacts"].get("compiled_coreml") or self.meta["artifacts"].get("coreml")
- context_bucket_meta = self.meta["runtime"].get("context_buckets") or []
- if not coreml_name and not context_bucket_meta:
- raise FileNotFoundError("ANE package metadata has no CoreML artifact path")
-
- try:
- import coremltools as ct
- except ImportError as e:
- raise RuntimeError(
- "ANE inference requires coremltools. Install with `pip install -r requirements-ane.txt`."
- ) from e
- cu = getattr(ct.ComputeUnit, compute_units)
-
- def load_coreml(path: Path):
- if path.suffix == ".mlmodelc":
- return ct.models.CompiledMLModel(str(path), compute_units=cu)
- return ct.models.MLModel(str(path), compute_units=cu)
-
- self.context_models = []
- if context_bucket_meta:
- for entry in sorted(context_bucket_meta, key=lambda item: int(item["context_length"])):
- model_name = entry.get("compiled_coreml") or entry.get("coreml")
- if not model_name:
- continue
- model_path = self.package_dir / model_name
- if not model_path.exists():
- raise FileNotFoundError(f"CoreML context bucket model not found: {model_path}")
- self.context_models.append(
- {
- "context_length": int(entry["context_length"]),
- "model": load_coreml(model_path),
- "input_name": entry.get("input_name", self.input_name),
- "position_name": entry.get("position_name", self.position_name),
- "output_name": entry.get("output_name", self.output_name),
- "output_mode": entry.get("output_mode", self.output_mode),
- }
- )
- if not self.context_models:
- if not coreml_name:
- raise FileNotFoundError("ANE package metadata has no runnable CoreML artifact path")
- model_path = self.package_dir / coreml_name
- if not model_path.exists():
- raise FileNotFoundError(f"CoreML model not found: {model_path}")
- self.context_models.append(
- {
- "context_length": self.context_length,
- "model": load_coreml(model_path),
- "input_name": self.input_name,
- "position_name": self.position_name,
- "output_name": self.output_name,
- "output_mode": self.output_mode,
- }
- )
- self.model = self.context_models[-1]["model"]
- self.kv_model = None
- self.kv_functional_model = None
- kv_coreml_name = self.meta["artifacts"].get("compiled_kv_coreml") or self.meta["artifacts"].get("kv_coreml")
- kv_functional_coreml_name = (
- self.meta["artifacts"].get("compiled_kv_functional_coreml")
- or self.meta["artifacts"].get("kv_functional_coreml")
- )
- self.kv_enabled = False
- if self.kv_mode == "stateful" and kv_coreml_name:
- kv_model_path = self.package_dir / kv_coreml_name
- if not kv_model_path.exists():
- raise FileNotFoundError(f"CoreML KV model not found: {kv_model_path}")
- try:
- self.kv_model = load_coreml(kv_model_path)
- self.kv_enabled = True
- except Exception as e:
- warnings.warn(
- f"CoreML KV model could not be loaded; falling back to full-context runtime: {e}",
- RuntimeWarning,
- )
- self.kv_model = None
- self.kv_mode = "none"
- if not self.kv_enabled and kv_functional_coreml_name:
- kv_functional_model_path = self.package_dir / kv_functional_coreml_name
- if not kv_functional_model_path.exists():
- raise FileNotFoundError(f"CoreML functional KV model not found: {kv_functional_model_path}")
- try:
- self.kv_functional_model = load_coreml(kv_functional_model_path)
- self.kv_mode = "functional"
- self.kv_enabled = True
- except Exception as e:
- warnings.warn(
- f"CoreML functional KV model could not be loaded; falling back to full-context runtime: {e}",
- RuntimeWarning,
- )
- self.kv_functional_model = None
- self.kv_mode = "none"
-
- def predict_next_logits(self, tokens: list[int], row_state: RowState | None = None) -> torch.Tensor:
- if self.kv_enabled and row_state is not None:
- self.last_context_length = self.context_length
- if self.kv_mode == "functional":
- if len(row_state.current_tokens) < self.functional_kv_min_context_length:
- return self._predict_next_logits_full_context(row_state.current_tokens)
- self.last_runtime_mode = "functional"
- return self._predict_next_logits_kv_functional(row_state)
- self.last_runtime_mode = "stateful"
- return self._predict_next_logits_kv_stateful(row_state)
- return self._predict_next_logits_full_context(tokens)
-
- def _predict_next_logits_full_context(self, tokens: list[int]) -> torch.Tensor:
- if not tokens:
- raise ValueError("tokens must not be empty")
- bucket = _select_context_bucket(self.context_models, min(len(tokens), self.context_length))
- bucket_context = int(bucket["context_length"])
- self.last_context_length = bucket_context
- self.last_runtime_mode = "full"
- window = tokens[-bucket_context:]
- input_ids = np.zeros((1, bucket_context), dtype=np.int32)
- input_ids[0, : len(window)] = np.asarray(window, dtype=np.int32)
- inputs = {bucket["input_name"]: input_ids}
- if bucket["output_mode"] == "last":
- inputs[bucket["position_name"]] = np.asarray([len(window) - 1], dtype=np.int32)
- outputs = bucket["model"].predict(inputs)
- if bucket["output_name"] in outputs:
- logits = np.asarray(outputs[bucket["output_name"]])
- else:
- logits = np.asarray(next(iter(outputs.values())))
- pos = len(window) - 1
- if logits.ndim == 3:
- logits = logits[0, pos, :]
- elif logits.ndim == 2:
- logits = logits[0, :]
- return torch.from_numpy(logits.astype(np.float32, copy=False))
-
- def _predict_next_logits_kv_stateful(self, state: RowState) -> torch.Tensor:
- if self.kv_model is None:
- raise RuntimeError("KV runtime requested but no KV CoreML model is loaded")
- if not state.current_tokens:
- raise ValueError("tokens must not be empty")
-
- window_start = max(0, len(state.current_tokens) - self.context_length)
- if state.kv_state is None or state.kv_window_start != window_start:
- state.kv_state = self.kv_model.make_state()
- state.kv_window_start = window_start
- state.kv_processed = window_start
- state.kv_last_logits = None
-
- for abs_pos in range(state.kv_processed, len(state.current_tokens)):
- cache_pos = abs_pos - state.kv_window_start
- inputs = {
- self.kv_input_name: np.asarray([[state.current_tokens[abs_pos]]], dtype=np.int32),
- self.kv_position_name: np.asarray([cache_pos], dtype=np.int32),
- }
- outputs = self.kv_model.predict(inputs, state=state.kv_state)
- if self.kv_output_name in outputs:
- logits = np.asarray(outputs[self.kv_output_name])
- else:
- logits = np.asarray(next(iter(outputs.values())))
- if logits.ndim == 2:
- logits = logits[0, :]
- elif logits.ndim == 3:
- logits = logits[0, -1, :]
- state.kv_last_logits = torch.from_numpy(logits.astype(np.float32, copy=False))
- state.kv_processed = abs_pos + 1
-
- if state.kv_last_logits is None:
- raise RuntimeError("KV state has no logits after prefill")
- return state.kv_last_logits
-
- def _predict_next_logits_kv_functional(self, state: RowState) -> torch.Tensor:
- if self.kv_functional_model is None:
- raise RuntimeError("functional KV runtime requested but no KV CoreML model is loaded")
- if not state.current_tokens:
- raise ValueError("tokens must not be empty")
-
- window_start = max(0, len(state.current_tokens) - self.context_length)
- if (
- state.kv_k_cache is None
- or state.kv_v_cache is None
- or state.kv_prev_embedding is None
- or state.kv_window_start != window_start
- ):
- state.kv_window_start = window_start
- state.kv_processed = window_start
- state.kv_last_logits = None
- state.kv_k_cache = np.zeros(tuple(self.kv_cache_shapes["k"]), dtype=np.float16)
- state.kv_v_cache = np.zeros(tuple(self.kv_cache_shapes["v"]), dtype=np.float16)
- state.kv_prev_embedding = np.zeros(tuple(self.kv_cache_shapes["prev_embedding"]), dtype=np.float16)
-
- k_in = self.kv_cache_input_names.get("k", "k_cache")
- v_in = self.kv_cache_input_names.get("v", "v_cache")
- prev_in = self.kv_cache_input_names.get("prev_embedding", "prev_embedding")
- k_out = self.kv_cache_output_names.get("k", "k_cache_out")
- v_out = self.kv_cache_output_names.get("v", "v_cache_out")
- prev_out = self.kv_cache_output_names.get("prev_embedding", "prev_embedding_out")
-
- for abs_pos in range(state.kv_processed, len(state.current_tokens)):
- cache_pos = abs_pos - state.kv_window_start
- inputs = {
- self.kv_input_name: np.asarray([[state.current_tokens[abs_pos]]], dtype=np.int32),
- self.kv_position_name: np.asarray([cache_pos], dtype=np.int32),
- k_in: state.kv_k_cache,
- v_in: state.kv_v_cache,
- prev_in: state.kv_prev_embedding,
- }
- outputs = self.kv_functional_model.predict(inputs)
- if self.kv_output_name in outputs:
- logits = np.asarray(outputs[self.kv_output_name])
- else:
- logits = np.asarray(next(iter(outputs.values())))
- state.kv_k_cache = np.asarray(outputs[k_out], dtype=np.float16)
- state.kv_v_cache = np.asarray(outputs[v_out], dtype=np.float16)
- state.kv_prev_embedding = np.asarray(outputs[prev_out], dtype=np.float16)
- if logits.ndim == 2:
- logits = logits[0, :]
- elif logits.ndim == 3:
- logits = logits[0, -1, :]
- state.kv_last_logits = torch.from_numpy(logits.astype(np.float32, copy=False))
- state.kv_processed = abs_pos + 1
-
- if state.kv_last_logits is None:
- raise RuntimeError("functional KV state has no logits after prefill")
- return state.kv_last_logits
-
- def generate(
- self,
- tokens: list[int],
- num_samples: int = 1,
- max_tokens: int | None = None,
- temperature: float = 1.0,
- top_k: int | None = None,
- seed: int = 42,
- **_ignored,
- ):
- assert isinstance(tokens, list) and isinstance(tokens[0], int), "expecting list of ints"
- rng = torch.Generator(device="cpu")
- rng.manual_seed(seed)
-
- def get_special(name: str) -> int:
- try:
- return self.tokenizer.encode_special(name)
- except (ValueError, KeyError):
- return -1
-
- python_start = get_special("<|python_start|>")
- python_end = get_special("<|python_end|>")
- output_start = get_special("<|output_start|>")
- output_end = get_special("<|output_end|>")
- assistant_end = (
- self.tokenizer.assistant_end_id()
- if hasattr(self.tokenizer, "assistant_end_id")
- else get_special("<|assistant_end|>")
- )
- bos = self.tokenizer.get_bos_token_id()
- row_states = [RowState(tokens.copy()) for _ in range(num_samples)]
-
- generated = 0
- while True:
- if max_tokens is not None and generated >= max_tokens:
- break
- if all(state.completed for state in row_states):
- break
-
- token_column = []
- token_masks = []
- for state in row_states:
- if state.completed:
- token_column.append(assistant_end)
- token_masks.append(0)
- continue
-
- is_forced = len(state.forced_tokens) > 0
- if is_forced:
- next_token = state.forced_tokens.popleft()
- else:
- logits = self.predict_next_logits(state.current_tokens, state)
- next_token = _sample_logits(logits, rng, temperature, top_k)
- token_column.append(next_token)
- token_masks.append(0 if is_forced else 1)
- state.current_tokens.append(next_token)
-
- if next_token == assistant_end or next_token == bos:
- state.completed = True
- if next_token == python_start:
- state.in_python_block = True
- state.python_expr_tokens = []
- elif next_token == python_end and state.in_python_block:
- state.in_python_block = False
- if state.python_expr_tokens:
- expr = self.tokenizer.decode(state.python_expr_tokens)
- result = _use_calculator(expr)
- if result is not None:
- state.forced_tokens.append(output_start)
- state.forced_tokens.extend(self.tokenizer.encode(str(result)))
- state.forced_tokens.append(output_end)
- state.python_expr_tokens = []
- elif state.in_python_block:
- state.python_expr_tokens.append(next_token)
-
- yield token_column, token_masks
- generated += 1
-
- def generate_batch(self, tokens: list[int], num_samples: int = 1, **kwargs):
- assistant_end = (
- self.tokenizer.assistant_end_id()
- if hasattr(self.tokenizer, "assistant_end_id")
- else self.tokenizer.encode_special("<|assistant_end|>")
- )
- bos = self.tokenizer.get_bos_token_id()
- results = [tokens.copy() for _ in range(num_samples)]
- masks = [[0] * len(tokens) for _ in range(num_samples)]
- completed = [False] * num_samples
- for token_column, token_masks in self.generate(tokens, num_samples, **kwargs):
- for i, (token, mask) in enumerate(zip(token_column, token_masks)):
- if not completed[i]:
- if token == assistant_end or token == bos:
- completed[i] = True
- else:
- results[i].append(token)
- masks[i].append(mask)
- if all(completed):
- break
- return results, masks
diff --git a/vendor/nanoai/attribution/__init__.py b/vendor/nanoai/attribution/__init__.py
deleted file mode 100644
index 0b867179ce816f8c42d763033b15ea50181e90f4..0000000000000000000000000000000000000000
--- a/vendor/nanoai/attribution/__init__.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Token-attribution tools for NanoAI models.
-
-Currently exposes FlashTrace (arXiv 2602.01914): efficient, faithful
-multi-token attribution that traces importance through a reasoning chain back
-to the source input. See ``flashtrace.py`` for the algorithm and
-``scripts/flashtrace_smoke.py`` for the needle-in-haystack smoke driver.
-
-This subpackage is a *sibling* to the model code: it never edits ``nanoai/gpt.py``.
-It re-derives the per-layer quantities FlashTrace needs (attention maps,
-residual-stream states, per-head value vectors) via an instrumented eager
-forward that reuses the trained weight modules but computes attention with an
-explicit softmax so the attention map ``alpha`` is available.
-"""
-
-from nanoai.attribution.flashtrace import (
- FlashTrace,
- InstrumentedGPT,
- LayerTrace,
- proximity,
- recovery_rate,
-)
-
-__all__ = [
- "FlashTrace",
- "InstrumentedGPT",
- "LayerTrace",
- "proximity",
- "recovery_rate",
-]
diff --git a/vendor/nanoai/attribution/flashtrace.py b/vendor/nanoai/attribution/flashtrace.py
deleted file mode 100644
index b7d1132a47d358a9825ab14bbca18f094bd935a3..0000000000000000000000000000000000000000
--- a/vendor/nanoai/attribution/flashtrace.py
+++ /dev/null
@@ -1,327 +0,0 @@
-"""FlashTrace: efficient + faithful multi-token attribution for NanoAI GPTs.
-
-Reference: "Towards Long-Horizon Interpretability: Efficient and Faithful
-Multi-Token Attribution for Reasoning LLMs" (Pan et al., arXiv 2602.01914).
-Local notes: ``knowledge/summary_flashtrace_attribution.md``.
-
-What this gives you
--------------------
-Given a context laid out as ``S = I ∘ T ∘ O`` (Input ∘ Reasoning ∘ Output),
-FlashTrace produces a per-input-token importance distribution explaining which
-tokens of ``I`` caused the generated output span ``O`` — tracing *through* the
-reasoning tokens ``T`` instead of letting them absorb all the attribution mass.
-
-Two mechanisms (both from the paper):
-
-1. **Span-wise aggregation** — attribute a whole target *span* in one pass by
- exploiting that the transformed value vector ``v_j`` is independent of the
- target position ``i``. The contribution of source ``j`` to a (weighted) span
- ``S`` factorizes as ``v_j · (Σ_{i∈S} w_i α_{i,j})`` — the expensive D-dim
- vector op runs once per source token regardless of span length M.
- Complexity O(M·N·D) → O(N·(M+D)).
-
-2. **Recursive attribution** — re-attribute the reasoning span ``T`` weighted by
- the importance it received in the previous hop, then discount by the mass
- ``ρ`` that flowed into the reasoning chain. ``K=1`` hop is enough in practice.
-
-Faithfulness vs. the paper
---------------------------
-We compute the *attention contribution term* exactly from the realized forward
-pass (we capture the true per-head value vectors actually fed to attention,
-including the ResFormer value-embedding mixin and GQA expansion), so we do NOT
-rely on the paper's ``v_j ≈ (x ⊙ s) W_V W_O`` linearization for that term — our
-attention term is more faithful than the paper's. The residual term still uses
-the L1 proximity of ALTI/IFR. The MLP sublayer does not move information across
-token positions (it is applied per-position after a norm), so it cannot
-redistribute attribution across *source* tokens; we therefore treat it as a
-per-position pass-through and it does not enter the cross-token attribution.
-This is a deliberate, documented simplification of Algorithm 1's normalization.
-
-This module is CPU-runnable (the instrumented forward computes attention with an
-explicit softmax and never calls the flash-attn / MSA kernels), which lets the
-test suite self-audit on a tiny random model without a GPU.
-"""
-
-from __future__ import annotations
-
-from dataclasses import dataclass
-
-import torch
-import torch.nn.functional as F
-
-from nanoai.gpt import apply_rotary_emb, norm
-
-EPS = 1e-9
-
-
-# -----------------------------------------------------------------------------
-# Captured per-layer quantities
-# -----------------------------------------------------------------------------
-@dataclass
-class LayerTrace:
- """Per-layer activations + attention captured during the instrumented forward.
-
- All tensors are for a single sequence (batch dim squeezed). N = context len.
- """
- x_in: torch.Tensor # (N, D) block input (after resid/x0 mixing)
- x_mid: torch.Tensor # (N, D) post-attention residual state
- x_out: torch.Tensor # (N, D) post-MLP residual state (== block output)
- mlp_out: torch.Tensor # (N, D) MLP contribution at each position
- alpha: torch.Tensor # (H, N, N) attention probs: alpha[h, query_i, key_j]
- v_heads: torch.Tensor # (N, H, hd) per-query-head value vectors (ve-mixed, GQA-expanded)
- wo_heads: torch.Tensor # (H, D, hd) output-proj slice per head: proj_h(v) = wo_heads[h] @ v
-
-
-# -----------------------------------------------------------------------------
-# Proximity (ALTI / IFR L1 metric)
-# -----------------------------------------------------------------------------
-def proximity(contrib: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
- """L1 proximity ``max(0, ||target||_1 - ||target - contrib||_1)``.
-
- Measures how much the magnitude of ``target`` would shrink if ``contrib``
- were removed. ``contrib`` is (N, D) (one row per source), ``target`` is (D,).
- Returns (N,), nonnegative.
- """
- tnorm = target.abs().sum()
- resid = (target.unsqueeze(0) - contrib).abs().sum(dim=-1)
- return torch.clamp(tnorm - resid, min=0.0)
-
-
-# -----------------------------------------------------------------------------
-# Instrumented forward
-# -----------------------------------------------------------------------------
-class InstrumentedGPT:
- """Re-runs a NanoAI ``GPT`` forward, capturing what FlashTrace needs.
-
- Faithfully mirrors ``GPT._forward_impl`` (embedding + norm + smear + per-layer
- resid/x0 mixing + value embeddings + pre-norm attn/MLP + backout), but
- computes attention with an explicit causal softmax so ``alpha`` is exposed.
- Only the dense full-sequence path is supported (kv_cache is None).
- """
-
- def __init__(self, model):
- self.model = model
- self.cfg = model.config
- if any(t != "dense" for t in model.layer_attn_types):
- raise NotImplementedError(
- "FlashTrace instrumented forward supports dense attention only; "
- f"got layer_attn_types={model.layer_attn_types}. MSA/sparse "
- "checkpoints need their own contribution decomposition."
- )
-
- @torch.no_grad()
- def run(self, idx: torch.Tensor) -> list[LayerTrace]:
- """idx: (1, N) or (N,) long tensor. Returns per-layer traces."""
- model, cfg = self.model, self.cfg
- if idx.dim() == 1:
- idx = idx.unsqueeze(0)
- assert idx.size(0) == 1, "instrumented forward is single-sequence (B=1)"
- device = model.transformer.wte.weight.device
- idx = idx.to(device)
- B, T = idx.size()
- H, n_kv = cfg.n_head, cfg.n_kv_head
- hd = cfg.n_embd // cfg.n_head
- group = H // n_kv
- dtype = model.transformer.wte.weight.dtype
-
- cos = model.cos[:, :T]
- sin = model.sin[:, :T]
-
- # --- embedding + norm + smear (mirror _forward_impl, training/B=1 path) ---
- x = model.transformer.wte(idx).to(dtype)
- x = norm(x)
- if T > 1:
- gate = model.smear_lambda.to(x.dtype) * torch.sigmoid(
- model.smear_gate(x[:, 1:, :24])
- )
- x = torch.cat([x[:, :1], x[:, 1:] + gate * x[:, :-1]], dim=1)
- x0 = x
-
- # causal (+ optional left-window) additive mask, shared across heads
- qpos = torch.arange(T, device=device)
- causal = qpos.unsqueeze(1) >= qpos.unsqueeze(0) # (Tq, Tk) lower-tri incl diag
-
- traces: list[LayerTrace] = []
- n_layer = cfg.n_layer
- backout_layer = n_layer // 2
- x_backout = None
- for i, block in enumerate(model.transformer.h):
- attn = block.attn
- x_in = model.resid_lambdas[i] * x + model.x0_lambdas[i] * x0 # (1,T,D)
- ve_key = model.ve_tid.get(i)
- ve = model.value_embeds[ve_key](idx).to(x.dtype) if ve_key is not None else None
-
- xn = norm(x_in)
- q = attn.c_q(xn).view(B, T, H, hd)
- k = attn.c_k(xn).view(B, T, n_kv, hd)
- v = attn.c_v(xn).view(B, T, n_kv, hd)
- if ve is not None:
- ve_r = ve.view(B, T, n_kv, hd)
- vgate = 3 * torch.sigmoid(attn.ve_gate(xn[..., : attn.ve_gate_channels]))
- v = v + vgate.unsqueeze(-1) * ve_r
- q = apply_rotary_emb(q, cos, sin)
- k = apply_rotary_emb(k, cos, sin)
- q, k = norm(q), norm(k) # QK norm
- q = q * 1.2
- k = k * 1.2
-
- # expand GQA kv-heads up to query heads
- k_e = k.repeat_interleave(group, dim=2) # (1,T,H,hd)
- v_e = v.repeat_interleave(group, dim=2) # (1,T,H,hd)
-
- # attention scores + causal/window mask -> probs (1,H,Tq,Tk)
- qh = q[0].permute(1, 0, 2) # (H, Tq, hd)
- kh = k_e[0].permute(1, 0, 2) # (H, Tk, hd)
- scores = torch.matmul(qh, kh.transpose(-2, -1)) / (hd ** 0.5) # (H,Tq,Tk)
- mask = causal
- left = model.window_sizes[i][0]
- if left is not None and left >= 0:
- within = (qpos.unsqueeze(1) - qpos.unsqueeze(0)) <= left
- mask = mask & within
- scores = scores.masked_fill(~mask.unsqueeze(0), float("-inf"))
- alpha = torch.softmax(scores.float(), dim=-1).to(dtype) # (H,Tq,Tk)
-
- vh = v_e[0] # (T, H, hd)
- # attention output: y[i] = Σ_h wo_h @ (Σ_j alpha[h,i,j] vh[j,h])
- ctx = torch.einsum("hij,jhd->ihd", alpha, vh) # (T,H,hd) per-head context
- y = ctx.reshape(T, H * hd).unsqueeze(0)
- attn_out = attn.c_proj(y) # (1,T,D)
- x_mid = x_in + attn_out
-
- mlp_out = block.mlp(norm(x_mid))
- x_out = x_mid + mlp_out
-
- # per-head output-projection slices: wo_heads[h] maps hd-vec -> D
- wo = attn.c_proj.weight.to(dtype) # (D, H*hd)
- wo_heads = wo.view(cfg.n_embd, H, hd).permute(1, 0, 2).contiguous() # (H,D,hd)
-
- traces.append(LayerTrace(
- x_in=x_in[0], x_mid=x_mid[0], x_out=x_out[0], mlp_out=mlp_out[0],
- alpha=alpha, v_heads=vh, wo_heads=wo_heads,
- ))
-
- x = x_out
- if i == backout_layer:
- x_backout = x
-
- # not used by attribution, but lets callers cross-check logits
- if x_backout is not None:
- x = x - model.backout_lambda.to(x.dtype) * x_backout
- self._final_x = norm(x)
- return traces
-
- @torch.no_grad()
- def logits(self) -> torch.Tensor:
- """Reconstruct logits from the last ``run`` for the faithfulness check."""
- softcap = 15
- lg = self.model.lm_head(self._final_x)[..., : self.cfg.vocab_size].float()
- return softcap * torch.tanh(lg / softcap)
-
-
-# -----------------------------------------------------------------------------
-# FlashTrace attribution
-# -----------------------------------------------------------------------------
-class FlashTrace:
- """Span-wise + recursive multi-token attribution over a captured forward."""
-
- def __init__(self, model):
- self.inst = InstrumentedGPT(model)
-
- def span_attribute(self, traces: list[LayerTrace], weights: torch.Tensor) -> torch.Tensor:
- """One hop: attribute the (weighted) target span to all source tokens.
-
- ``weights`` is (N,), zero outside the target span. Returns an (N,)
- nonnegative distribution summing to 1 over all source positions.
- """
- N = weights.shape[0]
- w = weights.to(traces[0].x_mid.dtype)
- E = torch.zeros(N, dtype=torch.float32, device=w.device)
- for lc in traces:
- Y_mid = (w.unsqueeze(1) * lc.x_mid).sum(0) # (D,)
- # pre-aggregated attention weights A^S[h, j] = Σ_i w_i α[h,i,j]
- A_S = torch.einsum("i,hij->hj", w, lc.alpha) # (H, N)
- # projected per-head value vectors: vp[j,h,:] = wo_heads[h] @ v_heads[j,h]
- vp = torch.einsum("hed,jhd->jhe", lc.wo_heads, lc.v_heads) # (N,H,D)
- C = torch.einsum("hj,jhe->je", A_S, vp) # (N,D) attn contribution of j
- e_tok = proximity(C, Y_mid).float() # (N,)
- R = (w.unsqueeze(1) * lc.x_in).sum(0, keepdim=True) # (1,D) residual bypass
- e_res = proximity(R, Y_mid).float().squeeze(0) # scalar
- denom = e_tok.sum() + e_res + EPS
- E = E + e_tok / denom
- return E / (E.sum() + EPS)
-
- @torch.no_grad()
- def attribute(self, idx, input_slice, reasoning_slice, output_slice, n_hops=1):
- """Full FlashTrace over a context laid out I ∘ T ∘ O.
-
- slices are python ``slice`` objects (or index tensors) into the N-length
- context. Returns a dict with the normalized input attribution and the
- per-hop diagnostics.
- """
- traces = self.inst.run(idx)
- N = traces[0].x_mid.shape[0]
- device = traces[0].x_mid.device
-
- def span_mask(s):
- m = torch.zeros(N, device=device)
- m[s] = 1.0
- return m
-
- I_mask = span_mask(input_slice)
- T_mask = span_mask(reasoning_slice)
-
- # hop 0: target = output span O, uniform weights
- hops = [self.span_attribute(traces, span_mask(output_slice))]
- rhos = []
- w_prev = hops[0]
- for _ in range(n_hops):
- rho = (w_prev * T_mask).sum() / (w_prev.sum() + EPS)
- w_next = w_prev * T_mask # restrict previous distribution to reasoning span
- if w_next.sum() <= EPS:
- break
- rhos.append(float(rho))
- hops.append(self.span_attribute(traces, w_next))
- w_prev = hops[-1]
-
- # accumulate input-restricted attribution across hops, discounting by ρ
- A = hops[0].clone()
- accum = 1.0
- for k in range(1, len(hops)):
- accum *= rhos[k - 1]
- A = A + accum * hops[k]
- A_input = A * I_mask
- A_input = A_input / (A_input.sum() + EPS)
-
- # hop-0-only baseline (no recursion), for ablation
- A0 = hops[0] * I_mask
- A0 = A0 / (A0.sum() + EPS)
-
- return {
- "A_input": A_input, # (N,) recursive attribution, support on I
- "A_input_hop0": A0, # (N,) no-recursion baseline
- "hops": hops, # list of (N,) per-hop distributions
- "rhos": rhos, # reasoning-mass flow ratios
- "n_input": int(I_mask.sum().item()),
- }
-
-
-# -----------------------------------------------------------------------------
-# Evaluation helper
-# -----------------------------------------------------------------------------
-def recovery_rate(a_input: torch.Tensor, input_slice, gt_positions, frac=0.1) -> float:
- """Fraction of ground-truth evidence tokens in the top-``frac`` of attribution.
-
- Mirrors the paper's Recovery Rate: rank input tokens by attribution, take the
- top ``floor(frac * |I|)``, report recall of the ground-truth positions.
- """
- idx = torch.arange(a_input.shape[0], device=a_input.device)
- in_pos = idx[input_slice]
- n_input = in_pos.numel()
- k = max(1, int(frac * n_input))
- scores = a_input[in_pos]
- top_local = torch.topk(scores, min(k, n_input)).indices
- top_global = set(in_pos[top_local].tolist())
- gt = set(int(p) for p in gt_positions)
- if not gt:
- return float("nan")
- return len(top_global & gt) / len(gt)
diff --git a/vendor/nanoai/config.py b/vendor/nanoai/config.py
deleted file mode 100644
index d97fe77953deac952c36a567c338a9ae6ede3366..0000000000000000000000000000000000000000
--- a/vendor/nanoai/config.py
+++ /dev/null
@@ -1,413 +0,0 @@
-"""NanoAI central configuration — the one documented home for env vars + tunables.
-
-ADDITIVE, gradual-migration module. It does NOT replace the scattered
-``os.environ`` reads and module-level constants that already exist; those keep
-working untouched. Instead this module:
-
- * gives every environment variable a single typed reader with a default and a
- one-time log (so "what env vars exist and what do they default to?" has one
- answer), and
- * exposes the magic constants that live across the codebase as frozen
- dataclasses, so new code (and a few high-value hotspots) read them from one
- place instead of duplicating literals.
-
-Design constraints (kept deliberately nanoai-minimal — no Hydra/OmegaConf, no
-model factories, no if-else dispatch):
-
- * The module top-level imports ONLY the standard library. It does NOT import
- torch, and it does NOT import ``nanoai.common`` at module top — those are
- deferred into ``paths()`` / ``precision()`` so ``import nanoai.config`` is
- cheap and side-effect free for non-training tools.
- * Path resolution REUSES ``nanoai.common.get_base_dir`` (never reinvented).
- * Frozen dataclasses + ``as_dict()`` returning a fresh dict each call, so a
- consumer cannot mutate the shared singleton.
-
-Functional identifiers are frozen: the package stays ``nanoai``, env vars stay
-``NANOAI_*`` / ``VR_*`` / ``CAPAGENT_*`` / ``GLM_*``, and ``~/.cache/nanochat*``
-checkpoint paths are unchanged. Only the project's presented name is NanoAI.
-
-================================================================================
-Environment variables consolidated here (canonical list)
-================================================================================
-Core / paths / precision:
- NANOAI_BASE_DIR base dir for all caches/checkpoints (default ~/.cache/nanochat)
- NANOAI_DTYPE compute dtype override: bfloat16|float16|float32 (else auto by SM)
- NANOAI_ATTN_BACKEND force attention backend: fa4|fa3|sdpa (else auto)
- NANOAI_MSA_KERNEL MSA block-sparse kernel: sdpa|fa4 (default sdpa)
- NANOAI_PACK_VERIFY enable rollout-packing verification (truthy)
-Distributed / runtime (read by torchrun / scripts; mirrored here for docs):
- RANK / LOCAL_RANK / WORLD_SIZE, CUDA_VISIBLE_DEVICES, OMP_NUM_THREADS,
- PYTORCH_ALLOC_CONF, GFT_NO_COMPILE
-HuggingFace:
- HF_ENDPOINT (default https://huggingface.co), HF_HUB_DISABLE_PROGRESS_BARS
-LLM judge / teacher endpoints:
- VR_LLM_JUDGE_ENABLE / _URL / _MODEL / _KEY / _SAMPLES / _TEMPERATURE / _TIMEOUT
- CAPAGENT_URL / CAPAGENT_ENDPOINT (two historical names) / CAPAGENT_API_KEY
- QWEN_JUDGE_MODEL, GLM_URL, GLM_MODEL, TSHARK_BIN, WANDB_RUN
-
-Magic constants mirrored as dataclasses: COMPUTE_DTYPE (PrecisionConfig),
-GPTConfig defaults (ModelDefaults), reward_mlr.DEFAULT_WEIGHTS (RewardConfig),
-opd_loss seg weights (SegWeightConfig), dense_verifier rewards
-(DenseVerifierConfig), eval_pcap_tasks thresholds (TaskEvalConfig),
-setup_optimizer defaults + init-lr-frac (OptimDefaults), dataset BASE_URL /
-MAX_SHARD + tokenizer/teacher model names (DatasetConfig).
-"""
-from __future__ import annotations
-
-import logging
-import os
-from dataclasses import asdict, dataclass
-from functools import lru_cache
-
-logger = logging.getLogger(__name__)
-
-# ---------------------------------------------------------------------------
-# 1. Environment-variable resolution layer
-# One typed reader per kind, each with a default + one-time INFO log so a
-# non-default value is always visible once in the logs.
-# ---------------------------------------------------------------------------
-_LOGGED: set[str] = set()
-_TRUTHY = ("1", "true", "yes", "on")
-
-
-def _log_once(name: str, value, source: str) -> None:
- if name not in _LOGGED:
- _LOGGED.add(name)
- logger.info("config: %s=%r (%s)", name, value, source)
-
-
-def env_str(name: str, default: str) -> str:
- raw = os.environ.get(name)
- if raw is None:
- return default
- _log_once(name, raw, "env")
- return raw
-
-
-def env_bool(name: str, default: bool) -> bool:
- raw = os.environ.get(name)
- if raw is None:
- return default
- value = raw.strip().lower() in _TRUTHY
- _log_once(name, value, "env")
- return value
-
-
-def env_int(name: str, default: int) -> int:
- raw = os.environ.get(name)
- if raw is None:
- return default
- try:
- value = int(raw)
- except ValueError as e:
- raise ValueError(f"env {name}={raw!r} is not an int") from e
- _log_once(name, value, "env")
- return value
-
-
-def env_float(name: str, default: float) -> float:
- raw = os.environ.get(name)
- if raw is None:
- return default
- try:
- value = float(raw)
- except ValueError as e:
- raise ValueError(f"env {name}={raw!r} is not a float") from e
- _log_once(name, value, "env")
- return value
-
-
-def env_choice(name: str, default: str, allowed: tuple) -> str:
- value = env_str(name, default)
- if value not in allowed:
- raise ValueError(f"env {name}={value!r} not in {allowed}")
- return value
-
-
-# ---------------------------------------------------------------------------
-# 2. Path config — wraps get_base_dir(), derives the standard subdir layout.
-# ---------------------------------------------------------------------------
-@dataclass(frozen=True)
-class PathConfig:
- base_dir: str
- base_checkpoints: str
- base_checkpoints_curated: str
- chatsft_checkpoints: str
- grpo_checkpoints: str
- cvar_diagnostics: str
- code_pretrain_data: str
- base_data_climbmix: str
-
- def checkpoint(self, kind: str, tag: str) -> str:
- """``/_checkpoints/`` — kind in base|chatsft|dpo|grpo|..."""
- return os.path.join(self.base_dir, f"{kind}_checkpoints", tag)
-
- def code_corpus(self, name: str) -> str:
- """``/code_pretrain_data/`` — e.g. the-stack-v2-dedup, synth-pleias."""
- return os.path.join(self.code_pretrain_data, name)
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-@lru_cache(maxsize=1)
-def paths() -> PathConfig:
- # Deferred import keeps module-level config torch-free; get_base_dir reused verbatim.
- from nanoai.common import get_base_dir
-
- base = get_base_dir()
- j = os.path.join
- return PathConfig(
- base_dir=base,
- base_checkpoints=j(base, "base_checkpoints"),
- base_checkpoints_curated=j(base, "base_checkpoints_curated"),
- chatsft_checkpoints=j(base, "chatsft_checkpoints"),
- grpo_checkpoints=j(base, "grpo_checkpoints"),
- cvar_diagnostics=j(base, "cvar_diagnostics"),
- code_pretrain_data=j(base, "code_pretrain_data"),
- base_data_climbmix=j(base, "base_data_climbmix"),
- )
-
-
-# ---------------------------------------------------------------------------
-# 3. Precision config — mirrors the live COMPUTE_DTYPE / ATTN_BACKEND singletons
-# by reading them, so it agrees by construction (no re-derivation drift).
-# ---------------------------------------------------------------------------
-_DTYPE_NAMES = {"torch.bfloat16": "bfloat16", "torch.float16": "float16", "torch.float32": "float32"}
-
-
-@dataclass(frozen=True)
-class PrecisionConfig:
- dtype_name: str # "bfloat16" | "float16" | "float32"
- dtype_reason: str # mirrors COMPUTE_DTYPE_REASON
- attn_backend: str # "fa4" | "fa3" | "sdpa"
- msa_kernel: str # "sdpa" | "fa4"
-
- @property
- def torch_dtype(self):
- import torch
- return {"bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32}[self.dtype_name]
-
-
-@lru_cache(maxsize=1)
-def precision() -> PrecisionConfig:
- # Deferred imports: reading the live singletons guarantees byte-identical agreement.
- from nanoai.common import COMPUTE_DTYPE, COMPUTE_DTYPE_REASON
- from nanoai.flash_attention import ATTN_BACKEND
-
- return PrecisionConfig(
- dtype_name=_DTYPE_NAMES[str(COMPUTE_DTYPE)],
- dtype_reason=COMPUTE_DTYPE_REASON,
- attn_backend=ATTN_BACKEND,
- msa_kernel=env_choice("NANOAI_MSA_KERNEL", "sdpa", ("sdpa", "fa4")),
- )
-
-
-# ---------------------------------------------------------------------------
-# 4. Model defaults — MIRROR of nanoai.gpt.GPTConfig (documentation only).
-# GPTConfig stays the authoritative source; do NOT have gpt.py import this
-# (keeps config torch-free and acyclic).
-# ---------------------------------------------------------------------------
-@dataclass(frozen=True)
-class ModelDefaults:
- sequence_len: int = 2048
- vocab_size: int = 32768
- window_pattern: str = "SSSL"
- msa_block_size: int = 64
- msa_top_k_blocks: int = 16
- msa_local_window: int = 128
- msa_sink_blocks: int = 1
- msa_index_mode: str = "pooled_k"
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-# ---------------------------------------------------------------------------
-# 5. Reward config — reward_mlr.DEFAULT_WEIGHTS (2026-05-17 retune).
-# Terminal MUST dominate per-turn shaping (see CLAUDE.md "MLR shaping weights
-# can overpower terminal"). Encoded as a soft invariant: warn, never assert,
-# so legitimate retuning is not blocked.
-# ---------------------------------------------------------------------------
-@dataclass(frozen=True)
-class RewardConfig:
- term: float = 1.0
- fmt: float = 0.05
- tool: float = 0.1
- rec: float = 0.1
- eff: float = 0.01
-
- def __post_init__(self):
- shaping = abs(self.fmt) + abs(self.tool) + abs(self.rec) + abs(self.eff)
- if shaping > 0 and self.term < 2 * shaping:
- logger.warning(
- "config.reward: term=%.3f < 2*shaping=%.3f — terminal may not dominate "
- "(CLAUDE.md: shaping can overpower terminal). Intentional? otherwise revisit.",
- self.term, 2 * shaping,
- )
-
- def as_dict(self) -> dict:
- return {"term": self.term, "fmt": self.fmt, "tool": self.tool, "rec": self.rec, "eff": self.eff}
-
-
-@dataclass(frozen=True)
-class SegWeightConfig:
- """opd_loss.DEFAULT_SEG_WEIGHTS — SAD segment masking."""
- reason: float = 1.0
- action: float = 1.5
- tool_arg: float = 1.2
- final_answer: float = 2.0
- tool_result: float = 0.0
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-@dataclass(frozen=True)
-class DenseVerifierConfig:
- """dense_verifier.py reward magnitudes (format penalty > dense > shaping > terminal/N)."""
- edit_ok: float = 0.4
- edit_fail: float = -0.5
- read_ok: float = 0.1
- read_fail: float = -0.5
- bash_ok: float = 0.1
- bash_fail: float = -0.2
- grep_hit: float = 0.2
- grep_miss: float = -0.1
- parse_error: float = -0.5
- ngram_repeat_penalty: float = -1.0
- mangled_special_penalty: float = -0.8
- truncated_tool_call_penalty: float = -0.6
- terminal_pass: float = 0.5
- terminal_fail: float = -0.3
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-@dataclass(frozen=True)
-class TaskEvalConfig:
- """Agentic eval/rollout lengths. eval/train length MUST match (CLAUDE.md):
- train at 512x12; the bug-prone eval default is 256x8 (silently truncates)."""
- pass_threshold: float = 0.50
- max_turns: int = 12
- max_tokens_per_turn: int = 512
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-@dataclass(frozen=True)
-class OptimDefaults:
- """Reference only. setup_optimizer attaches per-group lrs scaled by
- 1/sqrt(n_embd/768); do NOT pass a single flat --lr (CLAUDE.md). Use
- init_lr_frac (chat_rl semantics) to scale the per-group lrs uniformly."""
- init_lr_frac: float = 0.05
- # mirror of nanoai.gpt.GPT.setup_optimizer signature defaults
- unembedding_lr: float = 0.004
- embedding_lr: float = 0.2
- matrix_lr: float = 0.02
- scalar_lr: float = 0.5
- weight_decay: float = 0.0
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-# ---------------------------------------------------------------------------
-# 6. Teacher / judge endpoints — VR_* + CAPAGENT_* + GLM_* in one home.
-# Note the two historical CAPAGENT variable names (URL vs ENDPOINT).
-# ---------------------------------------------------------------------------
-@dataclass(frozen=True)
-class TeacherConfig:
- judge_enable: bool
- judge_url: str
- judge_model: str
- judge_api_key: str
- judge_samples: int
- judge_temperature: float
- judge_timeout: float
- capagent_endpoint: str
- capagent_api_key: str
- glm_url: str
- glm_model: str
- qwen_judge_model: str
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-def teacher() -> TeacherConfig:
- # CAPAGENT_URL and CAPAGENT_ENDPOINT are historical aliases; prefer ENDPOINT, fall back to URL.
- capagent = os.environ.get("CAPAGENT_ENDPOINT") or os.environ.get("CAPAGENT_URL") or "http://127.0.0.1:4210"
- return TeacherConfig(
- judge_enable=env_bool("VR_LLM_JUDGE_ENABLE", False),
- judge_url=env_str("VR_LLM_JUDGE_URL", "http://127.0.0.1:4210"),
- judge_model=env_str("VR_LLM_JUDGE_MODEL", "capagent"),
- judge_api_key=env_str("VR_LLM_JUDGE_KEY", "EMPTY"),
- judge_samples=env_int("VR_LLM_JUDGE_SAMPLES", 3),
- judge_temperature=env_float("VR_LLM_JUDGE_TEMPERATURE", 0.6),
- judge_timeout=env_float("VR_LLM_JUDGE_TIMEOUT", 30.0),
- capagent_endpoint=capagent,
- capagent_api_key=env_str("CAPAGENT_API_KEY", "sk-capagent"),
- glm_url=env_str("GLM_URL", "http://172.16.103.81:9000/v1/chat/completions"),
- glm_model=env_str("GLM_MODEL", "GLM-5.1"),
- qwen_judge_model=env_str("QWEN_JUDGE_MODEL", "qwen36-27b"),
- )
-
-
-# ---------------------------------------------------------------------------
-# 7. Dataset config — climbmix source + tokenizer/teacher model names + tools.
-# ---------------------------------------------------------------------------
-@dataclass(frozen=True)
-class DatasetConfig:
- base_url: str = "https://huggingface.co/datasets/karpathy/climbmix-400b-shuffle/resolve/main"
- max_shard: int = 6542
- hf_endpoint: str = "https://huggingface.co"
- qwen_model_default: str = "Qwen/Qwen2.5-0.5B"
- tshark_bin: str = "tshark"
-
- def as_dict(self) -> dict:
- return asdict(self)
-
-
-def dataset() -> DatasetConfig:
- return DatasetConfig(
- hf_endpoint=env_str("HF_ENDPOINT", "https://huggingface.co"),
- tshark_bin=env_str("TSHARK_BIN", "tshark"),
- )
-
-
-# ---------------------------------------------------------------------------
-# Pure-constant singletons (no env / no torch). Built once at import.
-# ---------------------------------------------------------------------------
-_MODEL = ModelDefaults()
-_REWARD = RewardConfig()
-_SEG = SegWeightConfig()
-_DENSE = DenseVerifierConfig()
-_TASK = TaskEvalConfig()
-_OPTIM = OptimDefaults()
-
-
-def model_defaults() -> ModelDefaults:
- return _MODEL
-
-
-def reward() -> RewardConfig:
- return _REWARD
-
-
-def seg_weights() -> SegWeightConfig:
- return _SEG
-
-
-def dense_verifier() -> DenseVerifierConfig:
- return _DENSE
-
-
-def task_eval() -> TaskEvalConfig:
- return _TASK
-
-
-def optim_defaults() -> OptimDefaults:
- return _OPTIM
diff --git a/vendor/nanoai/core_eval.py b/vendor/nanoai/core_eval.py
deleted file mode 100644
index f3c9a9f1f0f7e99abbe83d380913d07a3c0429ad..0000000000000000000000000000000000000000
--- a/vendor/nanoai/core_eval.py
+++ /dev/null
@@ -1,262 +0,0 @@
-"""
-Functions for evaluating the CORE metric, as described in the DCLM paper.
-https://arxiv.org/abs/2406.11794
-
-TODOs:
-- All tasks ~match except for squad. We get 31% reference is 37%. Figure out why.
-"""
-import random
-
-from jinja2 import Template
-import torch
-import torch.distributed as dist
-
-# -----------------------------------------------------------------------------
-# Prompt rendering utilities
-
-def render_prompts_mc(item, continuation_delimiter, fewshot_examples=None):
- """Render complete prompts for a multiple choice question"""
- template_str = """
-{%- for example in fewshot_examples -%}
-{{ example.query }}{{ continuation_delimiter }}{{ example.choices[example.gold] }}
-
-{% endfor -%}
-{{ item.query }}{{ continuation_delimiter }}{{ choice }}""".strip()
- template = Template(template_str)
- fewshot_examples = fewshot_examples or []
- context = {
- 'fewshot_examples': fewshot_examples,
- 'continuation_delimiter': continuation_delimiter,
- 'item': item
- }
- prompts = [template.render(choice=choice, **context) for choice in item['choices']]
- return prompts
-
-
-def render_prompts_schema(item, continuation_delimiter, fewshot_examples=None):
- """Render complete prompts for a schema question"""
- template_str = """
-{%- for example in fewshot_examples -%}
-{{ example.context_options[example.gold] }}{{ continuation_delimiter }}{{ example.continuation }}
-
-{% endfor -%}
-{{ context }}{{ continuation_delimiter }}{{ item.continuation }}""".strip()
- template = Template(template_str)
- fewshot_examples = fewshot_examples or []
- context = {
- 'fewshot_examples': fewshot_examples,
- 'continuation_delimiter': continuation_delimiter,
- 'item': item
- }
- prompts = [template.render(context=context_option, **context)
- for context_option in item['context_options']]
- return prompts
-
-
-def render_prompts_lm(item, continuation_delimiter, fewshot_examples=None):
- """
- Render complete prompt for a language modeling task.
- Notice that we manually trim the context in the template,
- which in some datasets seems to have trailing whitespace (which we don't want).
- """
- template_str = """
-{%- for example in fewshot_examples -%}
-{{ example.context | trim }}{{ continuation_delimiter }}{{ example.continuation }}
-
-{% endfor -%}
-{{ item.context | trim }}{{ continuation_delimiter }}{% if include_continuation %}{{ item.continuation }}{% endif %}""".strip()
- template = Template(template_str)
- fewshot_examples = fewshot_examples or []
- context = {
- 'fewshot_examples': fewshot_examples,
- 'continuation_delimiter': continuation_delimiter,
- 'item': item
- }
- # Return two prompts: without and with the continuation
- prompt_without = template.render(include_continuation=False, **context)
- prompt_with = template.render(include_continuation=True, **context)
- # Due to the way the data seems to be stored, I think I need to strip in the case of LM here.
- # Otherwise we may get trailing whitespaces in prompt_without (which get absorbed into the next
- # token in prompt_with), meaning we don't get a nice and clean prefix in the token space
- # to detect the final continuation. Tokenizers...
- prompt_without = prompt_without.strip()
- return [prompt_without, prompt_with]
-
-
-def find_common_length(token_sequences, direction='left'):
- """
- Find the length of the common prefix or suffix across token sequences
- - direction: 'left' for prefix, 'right' for suffix
- """
- min_len = min(len(seq) for seq in token_sequences)
- indices = {
- 'left': range(min_len),
- 'right': range(-1, -min_len-1, -1)
- }[direction]
- # Find the first position where the token sequences differ
- for i, idx in enumerate(indices):
- token = token_sequences[0][idx]
- if not all(seq[idx] == token for seq in token_sequences):
- return i
- return min_len
-
-
-def stack_sequences(tokens, pad_token_id):
- """Stack up a list of token sequences, pad to longest on the right"""
- bsz, seq_len = len(tokens), max(len(x) for x in tokens)
- input_ids = torch.full((bsz, seq_len), pad_token_id, dtype=torch.long)
- for i, x in enumerate(tokens):
- input_ids[i, :len(x)] = torch.tensor(x, dtype=torch.long)
- return input_ids
-
-
-def batch_sequences_mc(tokenizer, prompts):
- # In multiple choice, contexts are the same but the continuation is different (common prefix)
- tokens = tokenizer(prompts, prepend=tokenizer.get_bos_token_id())
- # figure out the start and end of each continuation
- answer_start_idx = find_common_length(tokens, direction='left')
- start_indices = [answer_start_idx] * len(prompts)
- end_indices = [len(x) for x in tokens]
- return tokens, start_indices, end_indices
-
-
-def batch_sequences_schema(tokenizer, prompts):
- # In schema tasks, contexts vary but continuation is the same (common suffix)
- tokens = tokenizer(prompts, prepend=tokenizer.get_bos_token_id())
- # figure out the start and end of each context
- suffix_length = find_common_length(tokens, direction='right')
- end_indices = [len(x) for x in tokens]
- start_indices = [ei - suffix_length for ei in end_indices]
- return tokens, start_indices, end_indices
-
-
-def batch_sequences_lm(tokenizer, prompts):
- # In LM tasks, we have two prompts: without and with continuation
- tokens = tokenizer(prompts, prepend=tokenizer.get_bos_token_id())
- tokens_without, tokens_with = tokens
- start_idx, end_idx = len(tokens_without), len(tokens_with)
- assert start_idx < end_idx, "prompt without is supposed to be a prefix of prompt with"
- assert tokens_without == tokens_with[:start_idx], "prompt without is supposed to be a prefix of prompt with"
- # we only need the with continuation prompt in the LM task, i.e. batch size of 1
- return [tokens_with], [start_idx], [end_idx]
-
-
-@torch.no_grad()
-def forward_model(model, input_ids):
- """
- Take BxT tensor of token ids, return BxT tensor of losses and argmax predictions.
- The last column of losses is set to nan because we don't have autoregressive targets there.
- """
- batch_size, seq_len = input_ids.size()
- outputs = model(input_ids)
- # Roll the tensor to the left by one position to get the (autoregressive) target ids
- target_ids = torch.roll(input_ids, shifts=-1, dims=1)
- # Calculate cross entropy at all positions
- losses = torch.nn.functional.cross_entropy(
- outputs.view(batch_size * seq_len, -1),
- target_ids.view(batch_size * seq_len),
- reduction='none'
- ).view(batch_size, seq_len)
- # Set the last column to be nan because there is no autoregressive loss there
- losses[:, -1] = float('nan')
- # Get the argmax predictions at each position
- predictions = outputs.argmax(dim=-1)
- return losses, predictions
-
-
-@torch.no_grad()
-def evaluate_example(idx, model, tokenizer, data, device, task_meta):
- """Evaluate a single example, return True if correct, False otherwise"""
- item = data[idx]
- task_type = task_meta['task_type']
- num_fewshot = task_meta['num_fewshot']
- continuation_delimiter = task_meta['continuation_delimiter']
-
- # Sample few-shot examples (excluding current item)
- fewshot_examples = []
- if num_fewshot > 0:
- rng = random.Random(1234 + idx)
- available_indices = [i for i in range(len(data)) if i != idx]
- fewshot_indices = rng.sample(available_indices, num_fewshot)
- fewshot_examples = [data[i] for i in fewshot_indices]
-
- # Render prompts and batch sequences based on task type
- if task_type == 'multiple_choice':
- prompts = render_prompts_mc(item, continuation_delimiter, fewshot_examples)
- tokens, start_idxs, end_idxs = batch_sequences_mc(tokenizer, prompts)
- elif task_type == 'schema':
- prompts = render_prompts_schema(item, continuation_delimiter, fewshot_examples)
- tokens, start_idxs, end_idxs = batch_sequences_schema(tokenizer, prompts)
- elif task_type == 'language_modeling':
- prompts = render_prompts_lm(item, continuation_delimiter, fewshot_examples)
- tokens, start_idxs, end_idxs = batch_sequences_lm(tokenizer, prompts)
- else:
- raise ValueError(f"Unsupported task type: {task_type}")
-
- # Some models can't forward sequences beyond a certain length (e.g. GPT-2)
- # In these cases, we have to truncate sequences to max length and adjust the indices
- if hasattr(model, 'max_seq_len') and model.max_seq_len is not None:
- max_tokens = model.max_seq_len
- new_tokens, new_start_idxs, new_end_idxs = [], [], []
- for t, s, e in zip(tokens, start_idxs, end_idxs):
- if len(t) > max_tokens:
- num_to_crop = len(t) - max_tokens
- new_tokens.append(t[-max_tokens:]) # take the last max_tokens tokens
- new_start_idxs.append(s - num_to_crop) # shift the indices down
- new_end_idxs.append(e - num_to_crop)
- assert s - num_to_crop >= 0, "this should never happen right?"
- assert e - num_to_crop >= 0, "this should never happen right?"
- else:
- new_tokens.append(t) # keep unchanged
- new_start_idxs.append(s)
- new_end_idxs.append(e)
- tokens, start_idxs, end_idxs = new_tokens, new_start_idxs, new_end_idxs
-
- # Stack up all the sequences into a batch
- pad_token_id = tokenizer.get_bos_token_id() # use BOS as pad token is ok
- input_ids = stack_sequences(tokens, pad_token_id)
- input_ids = input_ids.to(device)
-
- # Forward the model, get the autoregressive loss and argmax prediction at each token
- losses, predictions = forward_model(model, input_ids)
-
- # See if the losses/predictions come out correctly
- if task_type == 'language_modeling':
- # language modeling task is currently always batch size 1
- si = start_idxs[0]
- ei = end_idxs[0]
- # predictions[i] predict input_ids[i+1] autoregressively
- predicted_tokens = predictions[0, si-1:ei-1]
- actual_tokens = input_ids[0, si:ei]
- is_correct = torch.all(predicted_tokens == actual_tokens).item()
- elif task_type in ['multiple_choice', 'schema']:
- # For MC/schema: find the option with lowest average loss
- mean_losses = [losses[i, si-1:ei-1].mean().item()
- for i, (si, ei) in enumerate(zip(start_idxs, end_idxs))]
- pred_idx = mean_losses.index(min(mean_losses))
- is_correct = pred_idx == item['gold']
- else:
- raise ValueError(f"Unsupported task type: {task_type}")
-
- return is_correct
-
-
-def evaluate_task(model, tokenizer, data, device, task_meta):
- """
- This function is responsible for evaluating one task across many examples.
- It also handles dispatch to all processes if the script is run with torchrun.
- """
- rank = dist.get_rank() if dist.is_initialized() else 0
- world_size = dist.get_world_size() if dist.is_initialized() else 1
- correct = torch.zeros(len(data), dtype=torch.float32, device=device)
- # stride the examples to each rank
- for idx in range(rank, len(data), world_size):
- is_correct = evaluate_example(idx, model, tokenizer, data, device, task_meta)
- correct[idx] = float(is_correct)
- # sync results across all the processes if running distributed
- if world_size > 1:
- dist.barrier()
- dist.all_reduce(correct, op=dist.ReduceOp.SUM)
- # compute the mean
- mean_correct = correct.mean().item()
- return mean_correct
diff --git a/vendor/nanoai/dataloader.py b/vendor/nanoai/dataloader.py
deleted file mode 100644
index 789fe6895132296a072cb17b269c5ad4d3523c5b..0000000000000000000000000000000000000000
--- a/vendor/nanoai/dataloader.py
+++ /dev/null
@@ -1,166 +0,0 @@
-"""
-Distributed dataloaders for pretraining.
-
-BOS-aligned bestfit:
- - Every row starts with BOS token
- - Documents packed using best-fit algorithm to minimize cropping
- - When no document fits remaining space, crops a document to fill exactly
- - 100% utilization (no padding), ~35% tokens cropped at T=2048
-
-Compared to the original tokenizing_distributed_data_loader:
-BOS-aligned loses ~35% of tokens to cropping, but ensures that
-there are fewer "confusing" tokens in the train/val batches as every token can
-now attend back to the BOS token and sees the full context of the document.
-
-Fallback to the original if you have very limited data AND long documents:
-https://github.com/karpathy/nanoai/blob/3c3a3d7/nanoai/dataloader.py#L78-L117
-"""
-
-import torch
-import pyarrow.parquet as pq
-
-from nanoai.common import get_dist_info
-from nanoai.dataset import list_parquet_files
-
-def _document_batches(split, resume_state_dict, tokenizer_batch_size):
- """
- Infinite iterator over document batches (list of text strings) from parquet files.
-
- Handles DDP sharding and approximate resume. Each yield is (text_batch, (pq_idx, rg_idx, epoch))
- where text_batch is a list of document strings, indices track position for resumption,
- and epoch counts how many times we've cycled through the dataset (starts at 1).
- """
- ddp, ddp_rank, ddp_local_rank, ddp_world_size = get_dist_info()
-
- warn_on_legacy = ddp_rank == 0 and split == "train" # rank 0 on train split will warn on legacy
- parquet_paths = list_parquet_files(warn_on_legacy=warn_on_legacy)
- assert len(parquet_paths) != 0, "No dataset parquet files found, did you run dataset.py?"
- parquet_paths = parquet_paths[:-1] if split == "train" else parquet_paths[-1:]
-
- resume_pq_idx = resume_state_dict["pq_idx"] if resume_state_dict is not None else 0
- resume_rg_idx = resume_state_dict["rg_idx"] if resume_state_dict is not None else None
- resume_epoch = resume_state_dict.get("epoch", 1) if resume_state_dict is not None else 1
- first_pass = True
- pq_idx = resume_pq_idx
- epoch = resume_epoch
-
- while True: # iterate infinitely (multi-epoch)
- pq_idx = resume_pq_idx if first_pass else 0
- while pq_idx < len(parquet_paths):
- filepath = parquet_paths[pq_idx]
- pf = pq.ParquetFile(filepath)
- # Start from resume point if resuming on same file, otherwise from DDP rank
- if first_pass and (resume_rg_idx is not None) and (pq_idx == resume_pq_idx):
- base_idx = resume_rg_idx // ddp_world_size
- base_idx += 1 # advance by 1 so we don't repeat data after resuming
- rg_idx = base_idx * ddp_world_size + ddp_rank
- if rg_idx >= pf.num_row_groups:
- pq_idx += 1
- continue
- resume_rg_idx = None # only do this once
- else:
- rg_idx = ddp_rank
- while rg_idx < pf.num_row_groups:
- rg = pf.read_row_group(rg_idx)
- batch = rg.column('text').to_pylist()
- for i in range(0, len(batch), tokenizer_batch_size):
- yield batch[i:i+tokenizer_batch_size], (pq_idx, rg_idx, epoch)
- rg_idx += ddp_world_size
- pq_idx += 1
- first_pass = False
- epoch += 1
-
-
-def tokenizing_distributed_data_loader_with_state_bos_bestfit(
- tokenizer, B, T, split,
- tokenizer_threads=4, tokenizer_batch_size=128,
- device="cuda", resume_state_dict=None,
- buffer_size=1000
-):
- """
- BOS-aligned dataloader with Best-Fit Cropping.
-
- Reduces token waste compared to simple greedy cropping by searching a buffer
- for documents that fit well, while maintaining 100% utilization (no padding).
-
- Algorithm for each row:
- 1. From buffered docs, pick the LARGEST doc that fits entirely
- 2. Repeat until no doc fits
- 3. When nothing fits, crop a doc to fill remaining space exactly
-
- Key properties:
- - Every row starts with BOS
- - 100% utilization (no padding, every token is trained on)
- - Approximately 35% of all tokens are discarded due to cropping
- """
- assert split in ["train", "val"], "split must be 'train' or 'val'"
-
- row_capacity = T + 1
- batches = _document_batches(split, resume_state_dict, tokenizer_batch_size)
- bos_token = tokenizer.get_bos_token_id()
- doc_buffer = []
- pq_idx, rg_idx, epoch = 0, 0, 1
-
- def refill_buffer():
- nonlocal pq_idx, rg_idx, epoch
- doc_batch, (pq_idx, rg_idx, epoch) = next(batches)
- token_lists = tokenizer.encode(doc_batch, prepend=bos_token, num_threads=tokenizer_threads)
- for tokens in token_lists:
- doc_buffer.append(tokens)
-
- # Pre-allocate buffers once: layout is [inputs (B*T) | targets (B*T)]
- # This gives us contiguous views and a single HtoD transfer
- use_cuda = device == "cuda"
- row_buffer = torch.empty((B, row_capacity), dtype=torch.long) # for building rows without creating Python lists
- cpu_buffer = torch.empty(2 * B * T, dtype=torch.long, pin_memory=use_cuda) # staging area (CPU)
- gpu_buffer = torch.empty(2 * B * T, dtype=torch.long, device=device) # on-device buffer
- cpu_inputs = cpu_buffer[:B * T].view(B, T) # a few views into these buffers just for convenience
- cpu_targets = cpu_buffer[B * T:].view(B, T)
- inputs = gpu_buffer[:B * T].view(B, T)
- targets = gpu_buffer[B * T:].view(B, T)
-
- while True:
- for row_idx in range(B):
- pos = 0
- while pos < row_capacity:
- # Ensure buffer has documents
- while len(doc_buffer) < buffer_size:
- refill_buffer()
-
- remaining = row_capacity - pos
-
- # Find largest doc that fits entirely
- best_idx = -1
- best_len = 0
- for i, doc in enumerate(doc_buffer):
- doc_len = len(doc)
- if doc_len <= remaining and doc_len > best_len:
- best_idx = i
- best_len = doc_len
-
- if best_idx >= 0:
- doc = doc_buffer.pop(best_idx)
- doc_len = len(doc)
- row_buffer[row_idx, pos:pos + doc_len] = torch.tensor(doc, dtype=torch.long)
- pos += doc_len
- else:
- # No doc fits - crop shortest in buffer to fill remaining and minimize waste
- shortest_idx = min(range(len(doc_buffer)), key=lambda i: len(doc_buffer[i]))
- doc = doc_buffer.pop(shortest_idx)
- row_buffer[row_idx, pos:pos + remaining] = torch.tensor(doc[:remaining], dtype=torch.long)
- pos += remaining
-
- # Copy to pinned CPU buffer, then single HtoD transfer
- cpu_inputs.copy_(row_buffer[:, :-1])
- cpu_targets.copy_(row_buffer[:, 1:])
-
- state_dict = {"pq_idx": pq_idx, "rg_idx": rg_idx, "epoch": epoch}
-
- # Single HtoD copy into persistent GPU buffer and yield
- gpu_buffer.copy_(cpu_buffer, non_blocking=use_cuda)
- yield inputs, targets, state_dict
-
-def tokenizing_distributed_data_loader_bos_bestfit(*args, **kwargs):
- """Helper that omits state_dict from yields."""
- for inputs, targets, state_dict in tokenizing_distributed_data_loader_with_state_bos_bestfit(*args, **kwargs):
- yield inputs, targets
diff --git a/vendor/nanoai/dataset.py b/vendor/nanoai/dataset.py
deleted file mode 100644
index 8e9f7af8e7322cc13ce869849d6ca6186c585e93..0000000000000000000000000000000000000000
--- a/vendor/nanoai/dataset.py
+++ /dev/null
@@ -1,160 +0,0 @@
-"""
-The base/pretraining dataset is a set of parquet files.
-This file contains utilities for:
-- iterating over the parquet files and yielding documents from it
-- download the files on demand if they are not on disk
-
-For details of how the dataset was prepared, see `repackage_data_reference.py`.
-"""
-
-import os
-import argparse
-import time
-import requests
-import pyarrow.parquet as pq
-from multiprocessing import Pool
-
-from nanoai.common import get_base_dir
-
-# -----------------------------------------------------------------------------
-# The specifics of the current pretraining dataset
-
-# The URL on the internet where the data is hosted and downloaded from on demand
-BASE_URL = "https://huggingface.co/datasets/karpathy/climbmix-400b-shuffle/resolve/main"
-MAX_SHARD = 6542 # the last datashard is shard_06542.parquet
-index_to_filename = lambda index: f"shard_{index:05d}.parquet" # format of the filenames
-base_dir = get_base_dir()
-DATA_DIR = os.path.join(base_dir, "base_data_climbmix")
-
-# -----------------------------------------------------------------------------
-# These functions are useful utilities to other modules, can/should be imported
-
-def list_parquet_files(data_dir=None, warn_on_legacy=False):
- """ Looks into a data dir and returns full paths to all parquet files. """
- data_dir = DATA_DIR if data_dir is None else data_dir
-
- # Legacy-supporting code due to the upgrade from FinewebEdu-100B to ClimbMix-400B
- # This code will eventually be deleted.
- if not os.path.exists(data_dir):
- if warn_on_legacy:
- print()
- print("=" * 80)
- print(" WARNING: DATASET UPGRADE REQUIRED")
- print("=" * 80)
- print()
- print(f" Could not find: {data_dir}")
- print()
- print(" nanoai recently switched from FinewebEdu-100B to ClimbMix-400B.")
- print(" Everyone who does `git pull` as of March 4, 2026 is expected to see this message.")
- print(" To upgrade to the new ClimbMix-400B dataset, run these two commands:")
- print()
- print(" python -m nanoai.dataset -n 170 # download ~170 shards, enough for GPT-2, adjust as desired")
- print(" python -m scripts.tok_train # re-train tokenizer on new ClimbMix data")
- print()
- print(" For now, falling back to your old FinewebEdu-100B dataset...")
- print("=" * 80)
- print()
- # attempt a fallback to the legacy data directory
- data_dir = os.path.join(base_dir, "base_data")
-
- parquet_files = sorted([
- f for f in os.listdir(data_dir)
- if f.endswith('.parquet') and not f.endswith('.tmp')
- ])
- parquet_paths = [os.path.join(data_dir, f) for f in parquet_files]
- return parquet_paths
-
-def parquets_iter_batched(split, start=0, step=1):
- """
- Iterate through the dataset, in batches of underlying row_groups for efficiency.
- - split can be "train" or "val". the last parquet file will be val.
- - start/step are useful for skipping rows in DDP. e.g. start=rank, step=world_size
- """
- assert split in ["train", "val"], "split must be 'train' or 'val'"
- parquet_paths = list_parquet_files()
- parquet_paths = parquet_paths[:-1] if split == "train" else parquet_paths[-1:]
- for filepath in parquet_paths:
- pf = pq.ParquetFile(filepath)
- for rg_idx in range(start, pf.num_row_groups, step):
- rg = pf.read_row_group(rg_idx)
- texts = rg.column('text').to_pylist()
- yield texts
-
-# -----------------------------------------------------------------------------
-def download_single_file(index):
- """ Downloads a single file index, with some backoff """
-
- # Construct the local filepath for this file and skip if it already exists
- filename = index_to_filename(index)
- filepath = os.path.join(DATA_DIR, filename)
- if os.path.exists(filepath):
- print(f"Skipping {filepath} (already exists)")
- return True
-
- # Construct the remote URL for this file
- url = f"{BASE_URL}/{filename}"
- print(f"Downloading {filename}...")
-
- # Download with retries
- max_attempts = 5
- for attempt in range(1, max_attempts + 1):
- try:
- response = requests.get(url, stream=True, timeout=30)
- response.raise_for_status()
- # Write to temporary file first
- temp_path = filepath + f".tmp"
- with open(temp_path, 'wb') as f:
- for chunk in response.iter_content(chunk_size=1024 * 1024): # 1MB chunks
- if chunk:
- f.write(chunk)
- # Move temp file to final location
- os.rename(temp_path, filepath)
- print(f"Successfully downloaded {filename}")
- return True
-
- except (requests.RequestException, IOError) as e:
- print(f"Attempt {attempt}/{max_attempts} failed for {filename}: {e}")
- # Clean up any partial files
- for path in [filepath + f".tmp", filepath]:
- if os.path.exists(path):
- try:
- os.remove(path)
- except:
- pass
- # Try a few times with exponential backoff: 2^attempt seconds
- if attempt < max_attempts:
- wait_time = 2 ** attempt
- print(f"Waiting {wait_time} seconds before retry...")
- time.sleep(wait_time)
- else:
- print(f"Failed to download {filename} after {max_attempts} attempts")
- return False
-
- return False
-
-
-if __name__ == "__main__":
- parser = argparse.ArgumentParser(description="Download pretraining dataset shards")
- parser.add_argument("-n", "--num-files", type=int, default=-1, help="Number of train shards to download (default: -1), -1 = disable")
- parser.add_argument("-w", "--num-workers", type=int, default=4, help="Number of parallel download workers (default: 4)")
- args = parser.parse_args()
-
- # Prepare the output directory
- os.makedirs(DATA_DIR, exist_ok=True)
-
- # The way this works is that the user specifies the number of train shards to download via the -n flag.
- # In addition to that, the validation shard is *always* downloaded and is pinned to be the last shard.
- num_train_shards = MAX_SHARD if args.num_files == -1 else min(args.num_files, MAX_SHARD)
- ids_to_download = list(range(num_train_shards))
- ids_to_download.append(MAX_SHARD) # always download the validation shard
-
- # Download the shards
- print(f"Downloading {len(ids_to_download)} shards using {args.num_workers} workers...")
- print(f"Target directory: {DATA_DIR}")
- print()
- with Pool(processes=args.num_workers) as pool:
- results = pool.map(download_single_file, ids_to_download)
-
- # Report results
- successful = sum(1 for success in results if success)
- print(f"Done! Downloaded: {successful}/{len(ids_to_download)} shards to {DATA_DIR}")
diff --git a/vendor/nanoai/engine.py b/vendor/nanoai/engine.py
deleted file mode 100644
index 470f0f4ad1bff294d7d18d628ae095eb1af57b73..0000000000000000000000000000000000000000
--- a/vendor/nanoai/engine.py
+++ /dev/null
@@ -1,598 +0,0 @@
-"""
-Engine for efficient inference of our models.
-
-Everything works around token sequences:
-- The user can send token sequences to the engine
-- The engine returns the next token
-
-Notes:
-- The engine knows nothing about tokenization, it's purely token id sequences.
-
-The whole thing is made as efficient as possible.
-"""
-
-import torch
-import torch.nn.functional as F
-import signal
-import warnings
-from contextlib import contextmanager
-from collections import deque
-from nanoai.common import compute_init, autodetect_device_type
-from nanoai.checkpoint_manager import load_model
-
-# -----------------------------------------------------------------------------
-# Calculator tool helpers
-@contextmanager
-def timeout(duration, formula):
- def timeout_handler(signum, frame):
- raise Exception(f"'{formula}': timed out after {duration} seconds")
-
- signal.signal(signal.SIGALRM, timeout_handler)
- signal.alarm(duration)
- yield
- signal.alarm(0)
-
-def eval_with_timeout(formula, max_time=3):
- try:
- with timeout(max_time, formula):
- with warnings.catch_warnings():
- warnings.simplefilter("ignore", SyntaxWarning)
- return eval(formula, {"__builtins__": {}}, {})
- except Exception as e:
- signal.alarm(0)
- # print(f"Warning: Failed to eval {formula}, exception: {e}") # it's ok ignore wrong calculator usage
- return None
-
-def use_calculator(expr):
- """
- Evaluate a Python expression safely.
- Supports both math expressions and string operations like .count()
- """
- # Remove commas from numbers
- expr = expr.replace(",", "")
-
- # Check if it's a pure math expression (old behavior)
- if all([x in "0123456789*+-/.() " for x in expr]):
- if "**" in expr: # disallow power operator
- return None
- return eval_with_timeout(expr)
-
- # Check if it's a string operation we support
- # Allow: strings (single/double quotes), .count(), letters, numbers, spaces, parens
- allowed_chars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789'\"()._ "
- if not all([x in allowed_chars for x in expr]):
- return None
-
- # Disallow dangerous patterns
- dangerous_patterns = ['__', 'import', 'exec', 'eval', 'compile', 'open', 'file',
- 'input', 'raw_input', 'globals', 'locals', 'vars', 'dir',
- 'getattr', 'setattr', 'delattr', 'hasattr']
- expr_lower = expr.lower()
- if any(pattern in expr_lower for pattern in dangerous_patterns):
- return None
-
- # Only allow .count() method for now (can expand later)
- if '.count(' not in expr:
- return None
-
- # Evaluate with timeout
- return eval_with_timeout(expr)
-
-# -----------------------------------------------------------------------------
-class KVCache:
- """
- KV Cache designed for Flash Attention 3's flash_attn_with_kvcache API.
-
- Key differences from FA2-style cache:
- - Tensors are (B, T, H, D) not (B, H, T, D)
- - FA3 updates the cache in-place during flash_attn_with_kvcache
- - Position tracked per batch element via cache_seqlens tensor
- """
-
- def __init__(self, batch_size, num_heads, seq_len, head_dim, num_layers, device, dtype):
- self.batch_size = batch_size
- self.max_seq_len = seq_len
- self.n_layers = num_layers
- self.n_heads = num_heads
- self.head_dim = head_dim
- # Pre-allocate cache tensors: (n_layers, B, T, H, D)
- self.k_cache = torch.zeros(num_layers, batch_size, seq_len, num_heads, head_dim, device=device, dtype=dtype)
- self.v_cache = torch.zeros(num_layers, batch_size, seq_len, num_heads, head_dim, device=device, dtype=dtype)
- # Current sequence length per batch element (FA3 needs int32)
- self.cache_seqlens = torch.zeros(batch_size, dtype=torch.int32, device=device)
- self._pos = 0
- # Previous token's normalized embedding for smear (set by model forward pass)
- self.prev_embedding = None
- # When True, the attention layer routes to flash_attn_with_kvcache_variable_seqlens
- # which honors per-row cache_seqlens. Used by Engine.generate_batched_independent
- # for batched rollout with samples at different KV positions. Defaults to False
- # so the existing uniform-pos path is unchanged for training + Engine.generate.
- self.variable_seqlens = False
-
- def reset(self):
- """Reset cache to empty state."""
- self.cache_seqlens.zero_()
- self._pos = 0
- self.prev_embedding = None
- self._msa_pooled = None # MSA decode incremental pooled-K cache
- self._msa_pooled_pos = None
-
- def get_pos(self):
- """Get current position (assumes all batch elements at same position)."""
- if self._pos is not None:
- return self._pos
- return self.cache_seqlens[0].item()
-
- def get_layer_cache(self, layer_idx):
- """Return (k_cache, v_cache) views for a specific layer."""
- return self.k_cache[layer_idx], self.v_cache[layer_idx]
-
- def advance(self, num_tokens):
- """Advance the cache position by num_tokens (uniform across batch)."""
- self._pos = self.get_pos() + int(num_tokens)
- self.cache_seqlens += num_tokens
-
- def advance_per_row(self, num_tokens_per_row):
- """Advance each row's position by its own count. Used in batched rollout
- where different samples consume different numbers of new tokens per step
- (e.g. prefill phase with variable-length prompts)."""
- old_pos = self._pos
- if isinstance(num_tokens_per_row, torch.Tensor):
- self.cache_seqlens += num_tokens_per_row.to(dtype=torch.int32, device=self.cache_seqlens.device)
- self._pos = None
- else:
- self.cache_seqlens += torch.tensor(num_tokens_per_row, dtype=torch.int32, device=self.cache_seqlens.device)
- if len(set(num_tokens_per_row)) == 1:
- delta = int(num_tokens_per_row[0])
- self._pos = old_pos + delta if old_pos is not None else int(self.cache_seqlens[0].item())
- else:
- self._pos = None
-
- def prefill(self, other):
- """
- Copy cached KV from another cache into this one.
- Used when we do batch=1 prefill and then want to generate multiple samples in parallel.
- """
- assert self.get_pos() == 0, "Cannot prefill a non-empty KV cache"
- assert self.n_layers == other.n_layers and self.n_heads == other.n_heads and self.head_dim == other.head_dim
- assert self.max_seq_len >= other.max_seq_len
- other_pos = other.get_pos()
- self.k_cache[:, :, :other_pos, :, :] = other.k_cache[:, :, :other_pos, :, :]
- self.v_cache[:, :, :other_pos, :, :] = other.v_cache[:, :, :other_pos, :, :]
- self.cache_seqlens.fill_(other_pos)
- self._pos = other_pos
- # Copy smear state: expand batch=1 prev_embedding to num_samples
- if other.prev_embedding is not None:
- self.prev_embedding = other.prev_embedding.expand(self.batch_size, -1, -1).clone()
-
-# -----------------------------------------------------------------------------
-def apply_repetition_penalty(logits, recent_per_row, penalty, exempt_ids=None):
- """HuggingFace-style repetition penalty applied per row.
-
- logits: (B, V) float tensor (returns a NEW tensor when penalty != 1.0).
- recent_per_row: list of iterables of int (one per row); each is the row's
- recently-emitted token ids — used as the penalty index set.
- penalty: float. >1 discourages repetition; positive logits divided by
- penalty, negative logits multiplied (so both pushed toward 0/negative).
- exempt_ids: optional set[int] of token ids to skip penalty on (e.g. enum
- answer tokens in narrow-classification tasks where the correct token
- appears repeatedly in the prompt and should NOT be penalized).
- """
- if penalty == 1.0:
- return logits
- logits = logits.clone()
- for b, recent in enumerate(recent_per_row):
- if not recent:
- continue
- uniq = set(recent)
- if exempt_ids:
- uniq -= exempt_ids
- if not uniq:
- continue
- idx = torch.tensor(list(uniq), device=logits.device, dtype=torch.long)
- vals = logits[b].index_select(0, idx)
- vals = torch.where(vals > 0, vals / penalty, vals * penalty)
- logits[b, idx] = vals
- return logits
-
-@torch.inference_mode()
-def sample_next_token(logits, rng, temperature=1.0, top_k=None):
- """Sample a single next token from given logits of shape (B, vocab_size). Returns (B, 1)."""
- assert temperature >= 0.0, "temperature must be non-negative"
- if temperature == 0.0:
- return torch.argmax(logits, dim=-1, keepdim=True)
- if top_k is not None and top_k > 0:
- k = min(top_k, logits.size(-1))
- vals, idx = torch.topk(logits, k, dim=-1)
- vals = vals / temperature
- probs = F.softmax(vals, dim=-1)
- choice = torch.multinomial(probs, num_samples=1, generator=rng)
- return idx.gather(1, choice)
- else:
- logits = logits / temperature
- probs = F.softmax(logits, dim=-1)
- return torch.multinomial(probs, num_samples=1, generator=rng)
-
-
-def sampled_token_logprob(logits, sampled_ids, temperature=1.0, top_k=None):
- """log prob of each `sampled_ids` token under the SAME distribution that
- `sample_next_token` drew from (post repetition-penalty / temperature / top_k).
-
- This is the behavior-policy logprob needed for truncated importance sampling
- (TIS) in async / off-policy RL (cf. Polar, arxiv 2605.24220, "TIS: Enabled").
- Pass the exact `logits` used at sampling time and the ids that were sampled.
-
- logits: (B, V); sampled_ids: (B,) long. Returns (B,) float (log prob).
- For temperature==0 (argmax, a point mass) returns 0.0 as a sentinel.
- """
- if temperature == 0.0:
- return torch.zeros(logits.size(0), device=logits.device, dtype=torch.float32)
- if top_k is not None and top_k > 0:
- k = min(top_k, logits.size(-1))
- vals, idx = torch.topk(logits, k, dim=-1)
- logp = F.log_softmax(vals / temperature, dim=-1) # (B, k) over the top-k support
- # sampled id is in the top-k support by construction; pick its column.
- match = (idx == sampled_ids.unsqueeze(1))
- return (logp * match).sum(dim=-1)
- logp = F.log_softmax(logits / temperature, dim=-1)
- return logp.gather(1, sampled_ids.unsqueeze(1)).squeeze(1)
-
-# -----------------------------------------------------------------------------
-
-class RowState:
- # Per-row state tracking during generation
- def __init__(self, current_tokens=None):
- self.current_tokens = current_tokens or [] # Current token sequence for this row
- self.forced_tokens = deque() # Queue of tokens to force inject
- self.in_python_block = False # Whether we are inside a python block
- self.python_expr_tokens = [] # Tokens of the current python expression
- self.completed = False # Whether this row has completed generation
-
-class Engine:
-
- def __init__(self, model, tokenizer, *, nvfp4_generation: bool = False):
- self.model = model
- self.tokenizer = tokenizer # needed for tool use
- self.nvfp4_generation = bool(nvfp4_generation)
-
- @contextmanager
- def generation_forward_context(self):
- """Select the numeric path for generation forwards.
-
- Default behavior is dense TE BF16 for NVFP4-trained checkpoints to keep
- legacy eval/generation stable. ``nvfp4_generation=True`` opts into real
- TE NVFP4 generation; Float4Linear pads each GEMM internally so KV-cache
- prefill/decode shapes do not need dummy sequence tokens.
- """
- if getattr(self.model, "_uses_nvfp4", False) and not self.nvfp4_generation:
- from nanoai.nvfp4 import disable_nvfp4_autocast
- with disable_nvfp4_autocast():
- yield
- else:
- yield
-
- def generation_precision_mode(self) -> str:
- if getattr(self.model, "_uses_nvfp4", False) and self.nvfp4_generation:
- return "te_nvfp4"
- return "dense_bf16"
-
- @torch.inference_mode()
- def generate(self, tokens, num_samples=1, max_tokens=None, temperature=1.0, top_k=None, seed=42,
- repetition_penalty=1.0, recent_window=0, exempt_ids=None, return_logprobs=False):
- """Same as generate, but does single prefill and then clones the KV cache.
-
- repetition_penalty / recent_window: when penalty>1 and window>0, each row
- tracks its last `recent_window` GENERATED tokens (not prompt) and the
- logits for those token ids are pushed down before sampling. No-op when
- penalty==1.0 or window==0.
- exempt_ids: optional set[int] of token ids that should NEVER be penalized
- even if they appear in the recent window (useful for narrow-enum tasks).
-
- return_logprobs: when True, each step yields a THIRD element
- `logprob_column` — the behavior-policy log prob of each row's emitted
- token (None for FORCED tokens, which were not sampled). This is the
- UNTEMPERED model log pi(a) (full softmax), matching the training-time
- logp so a TIS ratio exp(logp_now - logp_behavior) is well defined.
- Enables TIS for async/off-policy RL (Polar mechanism). Default False
- keeps the 2-tuple yield contract for all existing callers.
- """
- assert isinstance(tokens, list) and isinstance(tokens[0], int), "expecting list of ints"
- device = self.model.get_device()
- # NOTE: setting the dtype here and in this way is an ugly hack.
- # Currently the repo assumes that cuda -> bfloat16 and everything else -> float32.
- # We need to know the dtype here to call __init__ on KVCache and pre-allocate its tensors.
- # As a quick hack, we're making generate() function inherit and know about this repo-wise assumption.
- # I think there has to be a bigger refactor to deal with device/dtype tracking across the codebase.
- # In particular, the KVCache should allocate its tensors lazily
- dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
- rng = torch.Generator(device=device)
- rng.manual_seed(seed)
-
- # Get the special tokens we need to coordinate the tool use state machine.
- # Qwen tokenizer doesn't expose python_start/end or output_start/end —
- # those drive the in-engine python REPL state machine, which isn't used
- # for the agentic tool-call format anyway. Use sentinels (-1) so the
- # state checks never match a real token id.
- def get_special(s):
- try:
- return self.tokenizer.encode_special(s)
- except (ValueError, KeyError):
- return -1
- python_start = get_special("<|python_start|>")
- python_end = get_special("<|python_end|>")
- output_start = get_special("<|output_start|>")
- output_end = get_special("<|output_end|>")
- # Qwen path: assistant_end_id adapter; nanoai: <|assistant_end|>.
- if hasattr(self.tokenizer, "assistant_end_id"):
- assistant_end = self.tokenizer.assistant_end_id()
- else:
- assistant_end = get_special("<|assistant_end|>")
- bos = self.tokenizer.get_bos_token_id() # if sampled, ends row
-
- # 1) Run a batch 1 prefill of the prompt tokens
- m = self.model.config
- kv_model_kwargs = {"num_heads": m.n_kv_head, "head_dim": m.n_embd // m.n_head, "num_layers": m.n_layer}
- kv_cache_prefill = KVCache(
- batch_size=1,
- seq_len=len(tokens),
- device=device,
- dtype=dtype,
- **kv_model_kwargs,
- )
- ids = torch.tensor([tokens], dtype=torch.long, device=device)
- with self.generation_forward_context():
- logits = self.model.forward(ids, kv_cache=kv_cache_prefill)
- logits = logits[:, -1, :].expand(num_samples, -1) # (num_samples, vocab_size)
-
- # 2) Replicate the KV cache for each sample/row
- kv_length_hint = (len(tokens) + max_tokens) if max_tokens is not None else self.model.config.sequence_len
- kv_cache_decode = KVCache(
- batch_size=num_samples,
- seq_len=kv_length_hint,
- device=device,
- dtype=dtype,
- **kv_model_kwargs,
- )
- kv_cache_decode.prefill(kv_cache_prefill)
- del kv_cache_prefill # no need to keep this memory around
-
- # 3) Initialize states for each sample
- row_states = [RowState(tokens.copy()) for _ in range(num_samples)]
-
- # Per-row recent-token tracking for optional repetition penalty.
- use_rep_penalty = repetition_penalty != 1.0 and recent_window > 0
- recent_per_row = (
- [deque(maxlen=recent_window) for _ in range(num_samples)]
- if use_rep_penalty else None
- )
-
- # 4) Main generation loop
- num_generated = 0
- while True:
- # Stop condition: we've reached max tokens
- if max_tokens is not None and num_generated >= max_tokens:
- break
- # Stop condition: all rows are completed
- if all(state.completed for state in row_states):
- break
-
- # Optional repetition penalty over the last `recent_window` generated tokens.
- if use_rep_penalty:
- logits = apply_repetition_penalty(logits, recent_per_row, repetition_penalty, exempt_ids)
-
- # Sample the next token for each row
- next_ids = sample_next_token(logits, rng, temperature, top_k) # (B, 1)
- sampled_tokens = next_ids[:, 0].tolist()
- # Behavior-policy logprob of the just-sampled token, computed UNTEMPERED
- # (full softmax of the post-penalty logits) — i.e. the model's log pi(a),
- # consistent with the untempered training-time logp (compute_log_pi), so
- # the TIS ratio exp(logp_now - logp_behavior) measures pure policy drift.
- # Temperature / top_k shaped WHICH token was drawn, not the policy we
- # optimize, so they are deliberately excluded here.
- sampled_logprobs = (
- sampled_token_logprob(logits, next_ids[:, 0], temperature=1.0, top_k=None).tolist()
- if return_logprobs else None
- )
-
- # Process each row: choose the next token, update state, optional tool use
- token_column = [] # contains the next token id along each row
- token_masks = [] # contains the mask (was it sampled (1) or forced (0)?) along each row
- logprob_column = [] if return_logprobs else None # behavior logp; None for forced
- for i, state in enumerate(row_states):
- # Select the next token in this row
- is_forced = len(state.forced_tokens) > 0 # are there tokens waiting to be forced in deque?
- token_masks.append(0 if is_forced else 1) # mask is 0 if forced, 1 if sampled
- if return_logprobs:
- logprob_column.append(None if is_forced else sampled_logprobs[i])
- next_token = state.forced_tokens.popleft() if is_forced else sampled_tokens[i]
- token_column.append(next_token)
- # Update the state of this row to include the next token
- state.current_tokens.append(next_token)
- # Track generated tokens for repetition penalty (post-sampling
- # so penalty applies to model's own generations across steps).
- if use_rep_penalty and not state.completed:
- recent_per_row[i].append(next_token)
- # On <|assistant_end|> or <|bos|>, mark the row as completed
- if next_token == assistant_end or next_token == bos:
- state.completed = True
- # Handle tool logic
- if next_token == python_start:
- state.in_python_block = True
- state.python_expr_tokens = []
- elif next_token == python_end and state.in_python_block:
- state.in_python_block = False
- if state.python_expr_tokens:
- expr = self.tokenizer.decode(state.python_expr_tokens)
- result = use_calculator(expr)
- if result is not None:
- result_tokens = self.tokenizer.encode(str(result))
- state.forced_tokens.append(output_start)
- state.forced_tokens.extend(result_tokens)
- state.forced_tokens.append(output_end)
- state.python_expr_tokens = []
- elif state.in_python_block:
- state.python_expr_tokens.append(next_token)
-
- # Yield the token column (3-tuple with behavior logprobs when requested)
- if return_logprobs:
- yield token_column, token_masks, logprob_column
- else:
- yield token_column, token_masks
- num_generated += 1
-
- # Prepare logits for next iteration
- ids = torch.tensor(token_column, dtype=torch.long, device=device).unsqueeze(1)
- with self.generation_forward_context():
- logits = self.model.forward(ids, kv_cache=kv_cache_decode)[:, -1, :] # (B, vocab_size)
-
- def generate_batch(self, tokens, num_samples=1, **kwargs):
- """
- Non-streaming batch generation that just returns the final token sequences.
- Returns a list of token sequences (list of lists of ints).
- Terminal tokens (assistant_end, bos) are not included in the results.
- """
- # Qwen-tokenizer path exposes assistant_end via adapter (= <|im_end|>);
- # nanoai-tokenizer path uses the literal special token name.
- if hasattr(self.tokenizer, "assistant_end_id"):
- assistant_end = self.tokenizer.assistant_end_id()
- else:
- assistant_end = self.tokenizer.encode_special("<|assistant_end|>")
- bos = self.tokenizer.get_bos_token_id()
- results = [tokens.copy() for _ in range(num_samples)]
- masks = [[0] * len(tokens) for _ in range(num_samples)]
- completed = [False] * num_samples
- for token_column, token_masks in self.generate(tokens, num_samples, **kwargs):
- for i, (token, mask) in enumerate(zip(token_column, token_masks)):
- if not completed[i]:
- if token == assistant_end or token == bos:
- completed[i] = True
- else:
- results[i].append(token)
- masks[i].append(mask)
- # Stop if all rows are completed
- if all(completed):
- break
- return results, masks
-
- @torch.inference_mode()
- def generate_batched_independent(
- self,
- prompts: list[list[int]],
- seeds: list[int],
- *,
- max_new_tokens: int,
- stop_tokens: set,
- temperature: float = 1.0,
- top_k: int = None,
- ) -> tuple[list[list[int]], list[bool]]:
- """Per-sample independent generation with byte-equal contract.
-
- Each row b in the batch:
- - starts from its own (potentially different-length) prompt
- - samples with its own torch.Generator seeded by seeds[b]
- - stops at its own first stop token (or max_new_tokens)
-
- Implementation note: bf16 batched matmul (B>1) is NOT bit-equal to
- repeated B=1 matmul due to reduction order differences — even argmax
- flips after a handful of steps on low-confidence outputs. To preserve
- byte-equal vs `generate(num_samples=1, seed=s_b)` per row, this method
- runs each sample through B=1 generate sequentially. The Stage 5 KV
- cache reuse extension (BatchedRolloutSession) is what unlocks the
- speedup by avoiding re-prefill between turns — not batching the GPU
- forward across samples.
-
- Returns:
- (new_tokens_per_sample, hit_stop_flag_per_sample)
- new_tokens_per_sample[b] = generated tokens for row b, INCLUDING
- the trailing stop token if hit (caller can strip if desired).
- """
- B = len(prompts)
- assert len(seeds) == B, f"seeds len {len(seeds)} != prompts len {B}"
- assert B >= 1
- assert max_new_tokens >= 1
- assert temperature >= 0.0
-
- outputs: list[list[int]] = []
- finished_flags: list[bool] = []
- for b in range(B):
- # Use existing generate_batch which already supports stop tokens
- # (assistant_end + bos) and yields per-step. We extract just the
- # new tokens (strip prompt prefix).
- full, _ = self.generate_batch(
- prompts[b],
- num_samples=1,
- max_tokens=max_new_tokens,
- temperature=temperature,
- top_k=top_k,
- seed=int(seeds[b]),
- )
- # generate_batch returns prompt+new_tokens with stops STRIPPED
- new_tokens = full[0][len(prompts[b]):]
- # Detect whether we stopped due to stop-token vs max-tokens by
- # counting: generate_batch breaks at stop, so if len < max_new
- # we hit a stop. (Stop token itself was stripped from output.)
- hit_stop = len(new_tokens) < max_new_tokens
- # Re-append the stop token if we hit it, so callers (rollout_one
- # equivalent) see the assistant_end and can detect terminal turns.
- if hit_stop:
- # Use assistant_end as the canonical stop marker (matches
- # rollout_one's `_generate_one_turn` contract).
- asst_end = self.tokenizer.encode_special("<|assistant_end|>")
- new_tokens = new_tokens + [asst_end]
- outputs.append(new_tokens)
- finished_flags.append(hit_stop)
-
- return outputs, finished_flags
-
-
-if __name__ == "__main__":
- """
- Quick inline test to make sure that the naive/slow model.generate function
- is equivalent to the faster Engine.generate function here.
- """
- import time
- # init compute
- device_type = autodetect_device_type()
- ddp, ddp_rank, ddp_local_rank, ddp_world_size, device = compute_init(device_type)
- # load the model and tokenizer
- model, tokenizer, meta = load_model("base", device, phase="eval")
- bos_token_id = tokenizer.get_bos_token_id()
- # common hyperparameters
- kwargs = dict(max_tokens=64, temperature=0.0)
- # set the starting prompt
- prompt_tokens = tokenizer.encode("The chemical formula of water is", prepend=bos_token_id)
- # generate the reference sequence using the model.generate() function
- generated_tokens = []
- torch.cuda.synchronize()
- t0 = time.time()
- stream = model.generate(prompt_tokens, **kwargs)
- for token in stream:
- generated_tokens.append(token)
- chunk = tokenizer.decode([token])
- print(chunk, end="", flush=True)
- print()
- torch.cuda.synchronize()
- t1 = time.time()
- print(f"Reference time: {t1 - t0:.2f}s")
- reference_ids = generated_tokens
- # generate tokens with Engine
- generated_tokens = []
- engine = Engine(model, tokenizer)
- stream = engine.generate(prompt_tokens, num_samples=1, **kwargs) # note: runs in fp32
- torch.cuda.synchronize()
- t0 = time.time()
- for token_column, token_masks in stream:
- token = token_column[0] # only print out the first row
- generated_tokens.append(token)
- chunk = tokenizer.decode([token])
- print(chunk, end="", flush=True)
- print()
- torch.cuda.synchronize()
- t1 = time.time()
- print(f"Engine time: {t1 - t0:.2f}s")
- # compare the two sequences
- for i in range(len(reference_ids)):
- if reference_ids[i] != generated_tokens[i]:
- print(f"Mismatch at {i}: {reference_ids[i]} != {generated_tokens[i]}")
- break
- print(f"Match: {reference_ids == generated_tokens}")
diff --git a/vendor/nanoai/execution.py b/vendor/nanoai/execution.py
deleted file mode 100644
index 6f50c7492bd2bff65e2380942140d44d11b3640b..0000000000000000000000000000000000000000
--- a/vendor/nanoai/execution.py
+++ /dev/null
@@ -1,349 +0,0 @@
-"""
-Sandboxed execution utilities for running Python code that comes out of an LLM.
-Adapted from OpenAI HumanEval code:
-https://github.com/openai/human-eval/blob/master/human_eval/execution.py
-
-What is covered:
-- Each execution runs in its own process (can be killed if it hangs or crashes)
-- Execution is limited by a timeout to stop infinite loops
-- Memory limits are enforced by default (256MB)
-- stdout and stderr are captured and returned
-- Code runs in a temporary directory that is deleted afterwards
-- Dangerous functions are disabled (examples: os.system, os.kill, shutil.rmtree, subprocess.Popen)
-
-What is not covered:
-- Not a true security sandbox
-- Network access is not blocked (e.g. sockets could be opened)
-- Python's dynamic features (e.g. ctypes) could bypass restrictions
-- No kernel-level isolation (no seccomp, no containers, no virtualization)
-
-Overall this sandbox is good for evaluation of generated code and protects against
-accidental destructive behavior, but it is not safe against malicious adversarial code.
-"""
-
-import contextlib
-import faulthandler
-import io
-import multiprocessing
-import os
-import platform
-import signal
-import tempfile
-from dataclasses import dataclass
-from typing import Optional
-
-# -----------------------------------------------------------------------------
-
-@dataclass
-class ExecutionResult:
- """Result of executing Python code in a sandbox."""
- success: bool
- stdout: str
- stderr: str
- error: Optional[str] = None
- timeout: bool = False
- memory_exceeded: bool = False
-
- def __repr__(self):
- parts = []
- parts.append(f"ExecutionResult(success={self.success}")
- if self.timeout:
- parts.append(", timeout=True")
- if self.memory_exceeded:
- parts.append(", memory_exceeded=True")
- if self.error:
- parts.append(f", error={self.error!r}")
- if self.stdout:
- parts.append(f", stdout={self.stdout!r}")
- if self.stderr:
- parts.append(f", stderr={self.stderr!r}")
- parts.append(")")
- return "".join(parts)
-
-
-@contextlib.contextmanager
-def time_limit(seconds: float):
- def signal_handler(signum, frame):
- raise TimeoutException("Timed out!")
-
- signal.setitimer(signal.ITIMER_REAL, seconds)
- signal.signal(signal.SIGALRM, signal_handler)
- try:
- yield
- finally:
- signal.setitimer(signal.ITIMER_REAL, 0)
-
-
-@contextlib.contextmanager
-def capture_io():
- """Capture stdout and stderr, and disable stdin."""
- stdout_capture = io.StringIO()
- stderr_capture = io.StringIO()
- stdin_block = WriteOnlyStringIO()
- with contextlib.redirect_stdout(stdout_capture):
- with contextlib.redirect_stderr(stderr_capture):
- with redirect_stdin(stdin_block):
- yield stdout_capture, stderr_capture
-
-
-@contextlib.contextmanager
-def create_tempdir():
- with tempfile.TemporaryDirectory() as dirname:
- with chdir(dirname):
- yield dirname
-
-
-class TimeoutException(Exception):
- pass
-
-
-class WriteOnlyStringIO(io.StringIO):
- """StringIO that throws an exception when it's read from"""
-
- def read(self, *args, **kwargs):
- raise IOError
-
- def readline(self, *args, **kwargs):
- raise IOError
-
- def readlines(self, *args, **kwargs):
- raise IOError
-
- def readable(self, *args, **kwargs):
- """Returns True if the IO object can be read."""
- return False
-
-
-class redirect_stdin(contextlib._RedirectStream): # type: ignore
- _stream = "stdin"
-
-
-@contextlib.contextmanager
-def chdir(root):
- if root == ".":
- yield
- return
- cwd = os.getcwd()
- os.chdir(root)
- try:
- yield
- finally:
- os.chdir(cwd)
-
-
-def reliability_guard(maximum_memory_bytes: Optional[int] = None):
- """
- This disables various destructive functions and prevents the generated code
- from interfering with the test (e.g. fork bomb, killing other processes,
- removing filesystem files, etc.)
-
- WARNING
- This function is NOT a security sandbox. Untrusted code, including, model-
- generated code, should not be blindly executed outside of one. See the
- Codex paper for more information about OpenAI's code sandbox, and proceed
- with caution.
- """
-
- if platform.uname().system != "Darwin":
- # These resource limit calls seem to fail on macOS (Darwin), skip?
- import resource
- resource.setrlimit(resource.RLIMIT_AS, (maximum_memory_bytes, maximum_memory_bytes))
- resource.setrlimit(resource.RLIMIT_DATA, (maximum_memory_bytes, maximum_memory_bytes))
- resource.setrlimit(resource.RLIMIT_STACK, (maximum_memory_bytes, maximum_memory_bytes))
-
- faulthandler.disable()
-
- import builtins
-
- builtins.exit = None
- builtins.quit = None
-
- import os
-
- os.environ["OMP_NUM_THREADS"] = "1"
-
- os.kill = None
- os.system = None
- os.putenv = None
- os.remove = None
- os.removedirs = None
- os.rmdir = None
- os.fchdir = None
- os.setuid = None
- os.fork = None
- os.forkpty = None
- os.killpg = None
- os.rename = None
- os.renames = None
- os.truncate = None
- os.replace = None
- os.unlink = None
- os.fchmod = None
- os.fchown = None
- os.chmod = None
- os.chown = None
- os.chroot = None
- os.fchdir = None
- os.lchflags = None
- os.lchmod = None
- os.lchown = None
- os.getcwd = None
- os.chdir = None
-
- import shutil
-
- shutil.rmtree = None
- shutil.move = None
- shutil.chown = None
-
- import subprocess
-
- subprocess.Popen = None # type: ignore
-
- __builtins__["help"] = None
-
- import sys
-
- sys.modules["ipdb"] = None
- sys.modules["joblib"] = None
- sys.modules["resource"] = None
- sys.modules["psutil"] = None
- sys.modules["tkinter"] = None
-
-
-def _unsafe_execute(code: str, timeout: float, maximum_memory_bytes: Optional[int], result_dict):
- """Execute code in a subprocess with safety guards. Results are written to result_dict."""
- with create_tempdir():
-
- # These system calls are needed when cleaning up tempdir.
- import os
- import shutil
-
- rmtree = shutil.rmtree
- rmdir = os.rmdir
- chdir = os.chdir
- unlink = os.unlink
-
- # Disable functionalities that can make destructive changes to the test.
- reliability_guard(maximum_memory_bytes=maximum_memory_bytes)
-
- # Default to failure
- result_dict.update({
- "success": False,
- "stdout": "",
- "stderr": "",
- "timeout": False,
- "memory_exceeded": False,
- "error": None,
- })
-
- try:
- exec_globals = {}
- with capture_io() as (stdout_capture, stderr_capture):
- with time_limit(timeout):
- # WARNING
- # This program exists to execute untrusted model-generated code. Although
- # it is highly unlikely that model-generated code will do something overtly
- # malicious in response to this test suite, model-generated code may act
- # destructively due to a lack of model capability or alignment.
- # Users are strongly encouraged to sandbox this evaluation suite so that it
- # does not perform destructive actions on their host or network. For more
- # information on how OpenAI sandboxes its code, see the accompanying paper.
- # Once you have read this disclaimer and taken appropriate precautions,
- # uncomment the following line and proceed at your own risk:
- exec(code, exec_globals)
-
- result_dict.update({
- "success": True,
- "stdout": stdout_capture.getvalue(),
- "stderr": stderr_capture.getvalue(),
- })
-
- except TimeoutException:
- result_dict.update({
- "timeout": True,
- "error": "Execution timed out",
- })
-
- except MemoryError as e:
- result_dict.update({
- "memory_exceeded": True,
- "error": f"Memory limit exceeded: {e}",
- })
-
- except BaseException as e:
- result_dict.update({
- "error": f"{type(e).__name__}: {e}",
- })
-
- # Needed for cleaning up.
- shutil.rmtree = rmtree
- os.rmdir = rmdir
- os.chdir = chdir
- os.unlink = unlink
-
-
-def execute_code(
- code: str,
- timeout: float = 5.0, # 5 seconds default
- maximum_memory_bytes: Optional[int] = 256 * 1024 * 1024, # 256MB default
-) -> ExecutionResult:
- """
- Execute Python code in a sandboxed environment.
-
- Args:
- code: Python code to execute as a string
- timeout: Maximum execution time in seconds (default: 5.0)
- maximum_memory_bytes: Memory limit in bytes (default: 256MB, None to disable)
-
- Returns:
- ExecutionResult with success status, stdout/stderr, and error information
-
- Example:
- >>> result = execute_code("print('hello world')")
- >>> result.success
- True
- >>> result.stdout
- 'hello world\\n'
- """
-
- manager = multiprocessing.Manager()
- result_dict = manager.dict()
-
- p = multiprocessing.Process(
- target=_unsafe_execute,
- args=(code, timeout, maximum_memory_bytes, result_dict)
- )
- p.start()
- p.join(timeout=timeout + 1)
-
- if p.is_alive():
- p.kill()
- return ExecutionResult(
- success=False,
- stdout="",
- stderr="",
- error="Execution timed out (process killed)",
- timeout=True,
- memory_exceeded=False,
- )
-
- if not result_dict:
- return ExecutionResult(
- success=False,
- stdout="",
- stderr="",
- error="Execution failed (no result returned)",
- timeout=True,
- memory_exceeded=False,
- )
-
- return ExecutionResult(
- success=result_dict["success"],
- stdout=result_dict["stdout"],
- stderr=result_dict["stderr"],
- error=result_dict["error"],
- timeout=result_dict["timeout"],
- memory_exceeded=result_dict["memory_exceeded"],
- )
-
diff --git a/vendor/nanoai/fat_tail.py b/vendor/nanoai/fat_tail.py
deleted file mode 100644
index b53d4c01cbf41b1f0ba528617aa2884bcf43b22e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/fat_tail.py
+++ /dev/null
@@ -1,223 +0,0 @@
-"""Fat-tail-aware statistics for evaluation reporting.
-
-The "mean" hides everything that matters about reliability. This module
-returns the triplet (mean, worst-decile, P99) plus a fail-mode breakdown
-for downstream evals — implementing the report format described in plan
-1-cvar-sft-mean-lively-cerf.md.
-
-Pure functions (numpy + python only). The CLI wrapper lives in
-experiments/fat_tail_cvar/scripts/fat_tail_report.py and consumes the existing eval artifacts:
- - agent_eval.py CSV (per-prompt rows)
- - agent_eval_mt.py CSV/JSONL (per-task rows + optional transcripts)
-
-CVaR_α (Conditional Value-at-Risk at quantile α) = mean of the values
-above the empirical α-quantile. It's the average loss in the *worst*
-(1-α) fraction of cases. For α=0.9 this is the worst 10% mean — the
-"reliability tail."
-
-Note: this module reports loss-flavored statistics (higher = worse).
-For accuracy-flavored signals, convert with `loss = 1 - acc` (binary)
-or `loss = max_reward - reward` (continuous, e.g. agentic 0-4 reward).
-"""
-from __future__ import annotations
-
-from collections import Counter
-from typing import Callable, Iterable, Sequence
-
-import numpy as np
-
-
-def fat_tail_stats(values: Sequence[float], alpha: float = 0.9) -> dict:
- """Compute fat-tail statistics over a 1-D numeric sequence.
-
- Args:
- values: per-example losses (higher = worse). Empty input → all-zero stats.
- alpha: quantile defining the tail boundary. CVaR_α averages over the
- top (1-α) fraction. Default 0.9 → worst-decile mean.
-
- Returns dict with keys:
- n, mean, p50, p90, p99, var_alpha, cvar_alpha, max,
- worst_k_indices (indices into the input, length = ceil((1-alpha)*n)).
- """
- arr = np.asarray(list(values), dtype=np.float64)
- n = arr.size
- if n == 0:
- return {
- "n": 0, "mean": 0.0, "p50": 0.0, "p90": 0.0, "p99": 0.0,
- "var_alpha": 0.0, "cvar_alpha": 0.0, "max": 0.0,
- "worst_k_indices": [],
- }
- # Empirical quantiles (linear interpolation, the numpy default).
- var_alpha = float(np.quantile(arr, alpha))
- p50 = float(np.quantile(arr, 0.5))
- p90 = float(np.quantile(arr, 0.9))
- p99 = float(np.quantile(arr, 0.99))
- # CVaR: mean over points strictly above VaR_α. Fallback to top-k mean
- # if the strict mask is empty (which happens for tiny samples or
- # ties at the boundary).
- above = arr[arr > var_alpha]
- if above.size == 0:
- k = max(1, int(np.ceil((1.0 - alpha) * n)))
- cvar_alpha = float(np.sort(arr)[-k:].mean())
- else:
- cvar_alpha = float(above.mean())
- # Indices of the worst-(1-α) cases.
- k = max(1, int(np.ceil((1.0 - alpha) * n)))
- worst_k_indices = np.argsort(arr)[-k:][::-1].tolist()
- return {
- "n": int(n),
- "mean": float(arr.mean()),
- "p50": p50,
- "p90": p90,
- "p99": p99,
- "var_alpha": var_alpha,
- "cvar_alpha": cvar_alpha,
- "max": float(arr.max()),
- "worst_k_indices": worst_k_indices,
- }
-
-
-def category_breakdown(
- records: Sequence[dict],
- category_fn: Callable[[dict], str],
- value_fn: Callable[[dict], float],
- alpha: float = 0.9,
-) -> dict[str, dict]:
- """Group records by category, run fat_tail_stats per group + overall.
-
- The returned dict has one key per category plus an "__all__" key with the
- aggregate. Each value is the fat_tail_stats dict augmented with a "values"
- list (used by fail_mode for downstream slicing).
- """
- buckets: dict[str, list[tuple[int, float]]] = {}
- all_values: list[tuple[int, float]] = []
- for i, rec in enumerate(records):
- v = float(value_fn(rec))
- cat = str(category_fn(rec))
- buckets.setdefault(cat, []).append((i, v))
- all_values.append((i, v))
-
- def _stats(pairs):
- if not pairs:
- return None
- idxs, vals = zip(*pairs)
- st = fat_tail_stats(vals, alpha=alpha)
- st["values"] = list(vals)
- st["record_indices"] = list(idxs)
- # Translate worst_k_indices (positions in `vals`) back to record indices.
- st["worst_k_record_indices"] = [idxs[j] for j in st["worst_k_indices"]]
- return st
-
- out = {cat: _stats(pairs) for cat, pairs in buckets.items()}
- out["__all__"] = _stats(all_values)
- return out
-
-
-def fail_mode(
- records: Sequence[dict],
- component_keys: Sequence[str],
- bottom_pct: float = 0.1,
- score_fn: Callable[[dict], float] | None = None,
-) -> dict:
- """Tally which binary components most often fired as 0 in the worst-decile.
-
- `component_keys` are dict keys whose values are 0/1 (or False/True). For each
- record in the bottom `bottom_pct` (by score_fn — lower score = worse), we
- count how many times each component was 0. The result is a sorted list of
- (component, count, percentage_of_tail).
-
- If score_fn is None, the components are summed to form a score (matches the
- agentic 0-4 reward semantics).
- """
- if not records:
- return {"tail_n": 0, "modes": []}
- if score_fn is None:
- def score_fn(r): # noqa: E306
- return sum(float(r.get(k, 0)) for k in component_keys)
-
- scores = np.array([score_fn(r) for r in records], dtype=np.float64)
- n = scores.size
- k = max(1, int(np.ceil(bottom_pct * n)))
- # Bottom-k by score (lowest = worst). argsort ascending → first k are tail.
- tail_idx = np.argsort(scores)[:k]
-
- counter: Counter = Counter()
- for i in tail_idx:
- rec = records[int(i)]
- for key in component_keys:
- if not bool(rec.get(key, False)):
- counter[key] += 1
-
- modes = sorted(counter.items(), key=lambda kv: -kv[1])
- return {
- "tail_n": int(k),
- "tail_indices": [int(i) for i in tail_idx],
- "modes": [
- {"component": key, "count": int(cnt), "pct_of_tail": 100.0 * cnt / k}
- for key, cnt in modes
- ],
- }
-
-
-def format_table(rows: list[dict], columns: Sequence[str], floatfmt: str = "{:.4f}") -> str:
- """Render a list-of-dicts as a fixed-width text table (no external deps)."""
- if not rows:
- return "(empty)"
- str_rows = [[str(c) for c in columns]]
- for r in rows:
- line = []
- for c in columns:
- v = r.get(c, "")
- if isinstance(v, float):
- line.append(floatfmt.format(v))
- else:
- line.append(str(v))
- str_rows.append(line)
- widths = [max(len(s[i]) for s in str_rows) for i in range(len(columns))]
- sep = " ".join("-" * w for w in widths)
- out = [" ".join(s.ljust(widths[i]) for i, s in enumerate(str_rows[0])), sep]
- for line in str_rows[1:]:
- out.append(" ".join(s.ljust(widths[i]) for i, s in enumerate(line)))
- return "\n".join(out)
-
-
-# -----------------------------------------------------------------------------
-# Self-check: run `python -m nanoai.fat_tail` to verify the math on a synthetic
-# Pareto-tailed array. Requires no GPU.
-
-def _self_check():
- rng = np.random.default_rng(42)
- # Mix: 90% Gaussian body, 10% heavy-tailed Pareto contamination.
- body = rng.normal(loc=1.0, scale=0.3, size=900)
- tail = rng.pareto(a=1.5, size=100) * 5.0 + 2.0 # heavy right tail
- losses = np.concatenate([body, tail])
- rng.shuffle(losses)
-
- s = fat_tail_stats(losses, alpha=0.9)
- print("fat_tail_stats on 1000-sample mix (Gaussian body + Pareto tail):")
- for k in ("n", "mean", "p50", "p90", "p99", "var_alpha", "cvar_alpha", "max"):
- print(f" {k:12s} = {s[k]:.4f}" if isinstance(s[k], float) else f" {k:12s} = {s[k]}")
- assert s["mean"] < s["cvar_alpha"], "CVaR must exceed mean"
- assert s["p50"] <= s["p90"] <= s["p99"] <= s["max"], "Quantile ordering broken"
- assert s["var_alpha"] <= s["cvar_alpha"], "CVaR must dominate VaR"
- print("\nself-check passed.")
-
- # Fail-mode probe: synthetic 4-component records with known tail bias.
- records = []
- for i in range(40):
- rec = {"a": 1, "b": 1, "c": 1, "d": 1}
- if i < 4: # tail: all components flipped
- rec = {"a": 0, "b": 0, "c": 0, "d": 0}
- records.append(rec)
- fm = fail_mode(records, ["a", "b", "c", "d"], bottom_pct=0.1)
- print(f"\nfail_mode (synthetic): tail_n={fm['tail_n']}")
- for m in fm["modes"]:
- print(f" {m['component']:8s} {m['count']:>3d}/{fm['tail_n']} "
- f"({m['pct_of_tail']:.0f}% of tail)")
- assert fm["tail_n"] == 4
- assert len(fm["modes"]) == 4 and all(m["count"] == 4 for m in fm["modes"])
- print("fail_mode self-check passed.")
-
-
-if __name__ == "__main__":
- _self_check()
diff --git a/vendor/nanoai/fp8.py b/vendor/nanoai/fp8.py
deleted file mode 100644
index e04f36bce53ea18e5d887071233b04d6eb02c26e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/fp8.py
+++ /dev/null
@@ -1,266 +0,0 @@
-"""Minimal FP8 training for nanoai — tensorwise dynamic scaling only.
-
-Drop-in replacement for torchao's Float8Linear (~2000 lines) with ~150 lines.
-We only need the "tensorwise" recipe (one scalar scale per tensor), not the full
-generality of torchao (rowwise scaling, FSDP float8 all-gather, DTensor, tensor
-subclass dispatch tables, etc.)
-
-How FP8 training works
-======================
-A standard Linear layer does one matmul in forward and two in backward:
- forward: output = input @ weight.T
- backward: grad_input = grad_output @ weight
- grad_weight= grad_output.T @ input
-
-FP8 training wraps each of these three matmuls with:
- 1. Compute scale = FP8_MAX / max(|tensor|) for each operand
- 2. Quantize: fp8_tensor = clamp(tensor * scale, -FP8_MAX, FP8_MAX).to(fp8)
- 3. Matmul via torch._scaled_mm (cuBLAS FP8 kernel, ~2x faster than bf16)
- 4. Dequantize: _scaled_mm handles this internally using the inverse scales
-
-The key insight: torch._scaled_mm and the float8 dtypes are PyTorch built-ins.
-torchao is just orchestration around these primitives. We can call them directly.
-
-FP8 dtype choice
-================
-There are two FP8 formats. We use both, following the standard convention:
- - float8_e4m3fn: 4-bit exponent, 3-bit mantissa, range [-448, 448]
- Higher precision (more mantissa bits), used for input and weight.
- - float8_e5m2: 5-bit exponent, 2-bit mantissa, range [-57344, 57344]
- Wider range (more exponent bits), used for gradients which can be large.
-
-torch._scaled_mm layout requirements
-=====================================
-The cuBLAS FP8 kernel requires specific memory layouts:
- - First argument (A): must be row-major (contiguous)
- - Second argument (B): must be column-major (B.t().contiguous().t())
-If B is obtained by transposing a contiguous tensor (e.g. weight.t()), it is
-already column-major — no copy needed. Otherwise we use _to_col_major().
-
-How this differs from torchao's approach
-========================================
-torchao uses a "tensor subclass" architecture: Float8TrainingTensor is a subclass
-of torch.Tensor that bundles FP8 data + scale + metadata. It implements
-__torch_dispatch__ with a dispatch table that intercepts every aten op (mm, t,
-reshape, clone, ...) and handles it in FP8-aware fashion. When you call
- output = input @ weight.T
-the @ operator dispatches to aten.mm, which gets intercepted and routed to
-torch._scaled_mm behind the scenes. This is ~2000 lines of code because you need
-a handler for every tensor operation that might touch an FP8 tensor.
-
-We take a simpler approach: a single autograd.Function (_Float8Matmul) that takes
-full-precision inputs, quantizes to FP8 internally, calls _scaled_mm, and returns
-full-precision outputs. Marked @allow_in_graph so torch.compile treats it as one
-opaque node rather than trying to trace inside.
-
-The trade-off is in how torch.compile sees the two approaches:
- - torchao: compile decomposes the tensor subclass (via __tensor_flatten__) and
- sees every individual op (amax, scale, cast, _scaled_mm) as separate graph
- nodes. Inductor can fuse these with surrounding operations (e.g. fuse the
- amax computation with the preceding layer's activation function).
- - ours: compile sees a single opaque call. It can optimize everything around
- the FP8 linear (attention, norms, etc.) but cannot fuse across the boundary.
-
-Both call the exact same cuBLAS _scaled_mm kernel — the GPU matmul is identical.
-The difference is only in the "glue" ops (amax, scale, cast) which are tiny
-compared to the matmul. In practice this means our version is slightly faster
-(less compilation overhead, no tensor subclass dispatch cost) but can produce
-subtly different floating-point rounding paths under torch.compile, since Inductor
-generates a different graph. Numerics are bitwise identical in eager mode.
-"""
-
-import torch
-import torch.nn as nn
-
-from nanoai.common import COMPUTE_DTYPE
-
-# Avoid division by zero when computing scale from an all-zeros tensor
-EPS = 1e-12
-
-
-@torch.no_grad()
-def _to_fp8(x, fp8_dtype):
- """Dynamically quantize a tensor to FP8 using tensorwise scaling.
-
- "Tensorwise" means one scalar scale for the entire tensor (as opposed to
- "rowwise" which computes a separate scale per row). Tensorwise is faster
- because cuBLAS handles the scaling; rowwise needs the CUTLASS kernel.
-
- Returns (fp8_data, inverse_scale) for use with torch._scaled_mm.
- """
- fp8_max = torch.finfo(fp8_dtype).max
- # Compute the max absolute value across the entire tensor
- amax = x.float().abs().max()
- # Scale maps [0, amax] -> [0, fp8_max]. Use float64 for the division to
- # ensure consistent numerics between torch.compile and eager mode.
- # (torchao does the same upcast — without it, compile/eager can diverge)
- scale = fp8_max / amax.double().clamp(min=EPS)
- scale = scale.float()
- # Quantize: scale into FP8 range, saturate (clamp prevents overflow when
- # casting — PyTorch's default is to wrap, not saturate), then cast to FP8
- x_scaled = x.float() * scale
- x_clamped = x_scaled.clamp(-fp8_max, fp8_max)
- x_fp8 = x_clamped.to(fp8_dtype)
- # _scaled_mm expects the *inverse* of our scale (it multiplies by this to
- # convert FP8 values back to the original range during the matmul)
- inv_scale = scale.reciprocal()
- return x_fp8, inv_scale
-
-
-def _to_col_major(x):
- """Rearrange a 2D tensor's memory to column-major layout.
-
- torch._scaled_mm requires its second operand in column-major layout.
- The trick: transpose -> contiguous (forces a copy in transposed order)
- -> transpose back. The result has the same logical shape but column-major
- strides, e.g. a [M, N] tensor gets strides (1, M) instead of (N, 1).
- """
- return x.t().contiguous().t()
-
-
-# allow_in_graph tells torch.compile to treat this as an opaque operation —
-# dynamo won't try to decompose it into smaller ops. See the module docstring
-# for how this differs from torchao's tensor subclass approach.
-@torch._dynamo.allow_in_graph
-class _Float8Matmul(torch.autograd.Function):
- """Custom autograd for the three FP8 GEMMs of a Linear layer.
-
- The forward quantizes input and weight to FP8 and saves
- the quantized tensors + scales for backward.
- """
-
- @staticmethod
- def forward(ctx, input_2d, weight):
- # Quantize both operands to e4m3 (higher precision format)
- input_fp8, input_inv = _to_fp8(input_2d, torch.float8_e4m3fn)
- weight_fp8, weight_inv = _to_fp8(weight, torch.float8_e4m3fn)
- ctx.save_for_backward(input_fp8, input_inv, weight_fp8, weight_inv)
-
- # output = input @ weight.T
- # input_fp8 is [B, K] contiguous = row-major (good for first arg)
- # weight_fp8 is [N, K] contiguous, so weight_fp8.t() is [K, N] with
- # strides (1, K) = column-major (good for second arg, no copy needed!)
- output = torch._scaled_mm(
- input_fp8,
- weight_fp8.t(),
- scale_a=input_inv,
- scale_b=weight_inv,
- out_dtype=input_2d.dtype,
- # use_fast_accum=True accumulates the dot products in lower precision.
- # Slightly less accurate but measurably faster. Standard practice for
- # the forward pass; we use False in backward for more precise gradients.
- use_fast_accum=True,
- )
- return output
-
- @staticmethod
- def backward(ctx, grad_output):
- in_fp8, in_inv, w_fp8, w_inv = ctx.saved_tensors
-
- # === GEMM 1: grad_input = grad_output @ weight ===
- # Shapes: [B, N] @ [N, K] -> [B, K]
- # Gradients use e5m2 (wider range), weights use e4m3 (higher precision)
- go_fp8, go_inv = _to_fp8(grad_output, torch.float8_e5m2)
- # go_fp8 is [B, N] contiguous = row-major, good for first arg
- # w_fp8 is [N, K] contiguous = row-major, need column-major for second arg
- w_col = _to_col_major(w_fp8)
- grad_input = torch._scaled_mm(
- go_fp8,
- w_col,
- scale_a=go_inv,
- scale_b=w_inv,
- out_dtype=grad_output.dtype,
- use_fast_accum=False,
- )
-
- # === GEMM 2: grad_weight = grad_output.T @ input ===
- # Shapes: [N, B] @ [B, K] -> [N, K]
- # go_fp8 is [B, N] contiguous, we need go.T = [N, B] as first arg.
- # Transposing gives column-major, but first arg needs row-major,
- # so we must call .contiguous() to physically rearrange the memory.
- go_T = go_fp8.t().contiguous() # [N, B] row-major
- in_col = _to_col_major(in_fp8) # [B, K] column-major
- grad_weight = torch._scaled_mm(
- go_T,
- in_col,
- scale_a=go_inv,
- scale_b=in_inv,
- out_dtype=grad_output.dtype,
- use_fast_accum=False,
- )
-
- return grad_input, grad_weight
-
-
-class Float8Linear(nn.Linear):
- """Drop-in nn.Linear replacement that does FP8 compute.
-
- Weights and biases remain in their original precision (e.g. fp32/bf16).
- Only the matmul is performed in FP8 via the _Float8Matmul autograd function.
- """
-
- def forward(self, input):
- # Cast input to COMPUTE_DTYPE (typically bf16) since _scaled_mm expects
- # reduced precision input, and we no longer rely on autocast to do this.
- input = input.to(COMPUTE_DTYPE)
- # _scaled_mm only works on 2D tensors, so flatten batch dimensions
- orig_shape = input.shape
- input_2d = input.reshape(-1, orig_shape[-1])
- output = _Float8Matmul.apply(input_2d, self.weight)
- output = output.reshape(*orig_shape[:-1], output.shape[-1])
- if self.bias is not None:
- output = output + self.bias.to(output.dtype)
- return output
-
- @classmethod
- def from_float(cls, mod):
- """Create Float8Linear from nn.Linear, sharing the same weight and bias.
-
- Uses meta device to avoid allocating a temporary weight tensor — we
- create the module shell on meta (shapes/dtypes only, no memory), then
- point .weight and .bias to the original module's parameters.
- """
- with torch.device("meta"):
- new_mod = cls(mod.in_features, mod.out_features, bias=False)
- new_mod.weight = mod.weight
- new_mod.bias = mod.bias
- return new_mod
-
-
-class Float8LinearConfig:
- """Minimal config matching torchao's API. Only tensorwise recipe is supported."""
-
- @staticmethod
- def from_recipe_name(recipe_name):
- if recipe_name != "tensorwise":
- raise ValueError(
- f"Only 'tensorwise' recipe is supported, got '{recipe_name}'. "
- f"Rowwise/axiswise recipes require the full torchao library."
- )
- return Float8LinearConfig()
-
-
-def convert_to_float8_training(module, *, config=None, module_filter_fn=None):
- """Replace nn.Linear layers with Float8Linear throughout a module.
-
- Walks the module tree in post-order (children before parents) and swaps
- each nn.Linear that passes the optional filter. The new Float8Linear shares
- the original weight and bias tensors — no copies, no extra memory.
-
- Args:
- module: Root module to convert.
- config: Float8LinearConfig (accepted for API compat, only tensorwise supported).
- module_filter_fn: Optional filter(module, fqn) -> bool. Only matching Linears
- are converted. Common use: skip layers with dims not divisible by 16
- (hardware requirement for FP8 matmuls on H100).
- """
- def _convert(mod, prefix=""):
- for name, child in mod.named_children():
- fqn = f"{prefix}.{name}" if prefix else name
- _convert(child, fqn)
- if isinstance(child, nn.Linear) and not isinstance(child, Float8Linear):
- if module_filter_fn is None or module_filter_fn(child, fqn):
- setattr(mod, name, Float8Linear.from_float(child))
-
- _convert(module)
- return module
diff --git a/vendor/nanoai/logo.svg b/vendor/nanoai/logo.svg
deleted file mode 100644
index f0f0571b7837f06cf3be95d5bbe5493313052996..0000000000000000000000000000000000000000
--- a/vendor/nanoai/logo.svg
+++ /dev/null
@@ -1,8 +0,0 @@
-
diff --git a/vendor/nanoai/loss_eval.py b/vendor/nanoai/loss_eval.py
deleted file mode 100644
index 5a556e6c2c8f4409b04d3dba08f487d1fed176a5..0000000000000000000000000000000000000000
--- a/vendor/nanoai/loss_eval.py
+++ /dev/null
@@ -1,65 +0,0 @@
-"""
-A number of functions that help with evaluating a base model.
-"""
-import math
-import torch
-import torch.distributed as dist
-
-@torch.no_grad()
-def evaluate_bpb(model, batches, steps, token_bytes):
- """
- Instead of the naive 'mean loss', this function returns the bits per byte (bpb),
- which is a tokenization vocab size-independent metric, meaning you are still comparing
- apples:apples if you change the vocab size. The way this works is that instead of just
- calculating the average loss as usual, you calculate the sum loss, and independently
- also the sum bytes (of all the target tokens), and divide. This normalizes the loss by
- the number of bytes that the target tokens represent.
-
- The added complexity is so that:
- 1) All "normal" tokens are normalized by the length of the token in bytes
- 2) No special tokens (e.g. <|bos|>) are included in the metric - they are masked out.
- 3) No actively masked tokens (using ignore_index of e.g. -1) are included in the metric.
-
- In addition to evaluate_loss, we need the token_bytes tensor:
- It is a 1D tensor of shape (vocab_size,), indicating the number of bytes for
- each token id, or 0 if the token is to not be counted (e.g. special tokens).
- """
- # record the losses
- total_nats = torch.tensor(0.0, dtype=torch.float32, device=model.get_device())
- total_bytes = torch.tensor(0, dtype=torch.int64, device=model.get_device())
- batch_iter = iter(batches)
- for _ in range(steps):
- x, y = next(batch_iter)
- loss2d = model(x, y, loss_reduction='none') # (B, T)
- loss2d = loss2d.view(-1) # flatten
- y = y.view(-1) # flatten
- if (y.int() < 0).any(): # mps does not currently have kernel for < 0 for int64, only int32
- # slightly more complex code path if some target tokens are ignore_index (e.g. -1)
- # any target token < 0 is to be ignored: do NOT index token_bytes with negatives
- valid = y >= 0
- y_safe = torch.where(valid, y, torch.zeros_like(y))
- # map valid targets to their byte length; ignored targets contribute 0 bytes
- num_bytes2d = torch.where(
- valid,
- token_bytes[y_safe],
- torch.zeros_like(y, dtype=token_bytes.dtype)
- )
- total_nats += (loss2d * (num_bytes2d > 0)).sum()
- total_bytes += num_bytes2d.sum()
- else:
- # fast path: no ignored targets, safe to index directly
- num_bytes2d = token_bytes[y]
- total_nats += (loss2d * (num_bytes2d > 0)).sum()
- total_bytes += num_bytes2d.sum()
- # sum reduce across all ranks
- world_size = dist.get_world_size() if dist.is_initialized() else 1
- if world_size > 1:
- dist.all_reduce(total_nats, op=dist.ReduceOp.SUM)
- dist.all_reduce(total_bytes, op=dist.ReduceOp.SUM)
- # move both to cpu, calculate bpb and return
- total_nats = total_nats.item()
- total_bytes = total_bytes.item()
- if total_bytes == 0:
- return float('inf')
- bpb = total_nats / (math.log(2) * total_bytes)
- return bpb
diff --git a/vendor/nanoai/mlx/README.md b/vendor/nanoai/mlx/README.md
deleted file mode 100644
index dc64d44281bf3c3eda3728de35137d9572de3e3e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/README.md
+++ /dev/null
@@ -1,33 +0,0 @@
-# nanoai.mlx — MLX training & inference (Apple Silicon)
-
-A native [MLX](https://github.com/ml-explore/mlx) port of NanoAI's GPT for
-training and inference on Apple GPUs, with an NVFP4 **simulation** path for
-experimenting with 4-bit precision off-Blackwell. Full write-up:
-[docs/mlx_nvfp4.md](../docs/mlx_nvfp4.md) (and `docs/mlx_nvfp4_port_spike.md`).
-
-## Modules
-
-- `gpt.py` — MLX GPT (mirror of `nanoai/gpt.py`); exports `GPT`, `GPTConfig`, `loss_fn`.
-- `train.py` — MLX training loop.
-- `engine.py` — MLX inference engine.
-- `checkpoint.py` — MLX checkpoint serialization / interop with torch checkpoints.
-- `nvfp4.py` — NVFP4 quantization **simulation** on Apple GPUs
- (`NVFP4_GROUP_SIZE=16`, `NVFP4_BITS=4`, modes `dense|nvfp4|nvfp4-safe`,
- `NVFP4_MIN_LINEAR_DIM=128`).
-
-## Dependencies
-
-The `mlx` optional extra (`mlx>=0.31.2` in `pyproject.toml`):
-
-```bash
-uv sync --extra mlx # or: pip install "mlx>=0.31.2"
-```
-
-## Notes
-
-- This is a separate package (`nanoai.mlx`), not under `nanoai/`, because it
- is a parallel backend rather than a core-library subpackage.
-- The NVFP4 path here is a **simulation** for Apple GPUs; real NVFP4 tensor-core
- execution is the Blackwell/TransformerEngine path under `nanoai/` (NVIDIA).
-- Tests: `tests/test_mlx_*.py`. A d4 MLX speedrun smoke note is in
- `knowledge/tst_mlx_d4_speedrun_20260526.md`.
diff --git a/vendor/nanoai/mlx/__init__.py b/vendor/nanoai/mlx/__init__.py
deleted file mode 100644
index 5bd55fed7bd15deaaa580cca5e34e90e3ff9b0ac..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/__init__.py
+++ /dev/null
@@ -1,11 +0,0 @@
-"""MLX runtime for nanoai with NVFP4 training and inference support."""
-
-__all__ = ["GPT", "GPTConfig", "loss_fn"]
-
-
-def __getattr__(name):
- if name in __all__:
- from nanoai.mlx.gpt import GPT, GPTConfig, loss_fn
-
- return {"GPT": GPT, "GPTConfig": GPTConfig, "loss_fn": loss_fn}[name]
- raise AttributeError(name)
diff --git a/vendor/nanoai/mlx/checkpoint.py b/vendor/nanoai/mlx/checkpoint.py
deleted file mode 100644
index 8979af9272293fa67b5341a70a5de3048db84a59..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/checkpoint.py
+++ /dev/null
@@ -1,305 +0,0 @@
-"""Checkpoint loading/export for the MLX nanoai runtime."""
-
-from __future__ import annotations
-
-import glob
-import json
-import os
-from pathlib import Path
-from typing import Any
-
-import mlx.core as mx
-
-from nanoai.common import get_base_dir
-from nanoai.mlx.gpt import GPT, GPTConfig, Linear
-from nanoai.mlx.nvfp4 import QUANTIZATION_MODES, quantize_weight
-
-
-FORMAT_VERSION = 1
-SOURCE_DIRS = {
- "base": "base_checkpoints",
- "sft": "chatsft_checkpoints",
- "rl": "chatrl_checkpoints",
- "dpo": "dpo_checkpoints",
- "grpo": "grpo_checkpoints",
-}
-
-
-def find_last_step(checkpoint_dir: str | os.PathLike) -> int:
- files = glob.glob(os.path.join(str(checkpoint_dir), "model_*.pt"))
- if not files:
- files = glob.glob(os.path.join(str(checkpoint_dir), "step_*.safetensors"))
- if not files:
- raise FileNotFoundError(f"No model_*.pt or step_*.safetensors found in {checkpoint_dir}")
- return int(max(os.path.basename(f).split("_")[-1].split(".")[0] for f in files if "_meta" not in f))
-
-
-def resolve_pt_checkpoint_dir(source: str, model_tag: str | None) -> str:
- if source not in SOURCE_DIRS:
- raise KeyError(f"unknown source {source!r}; expected one of {sorted(SOURCE_DIRS)}")
- root = os.path.join(get_base_dir(), SOURCE_DIRS[source])
- if model_tag is None:
- tags = [d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]
- if not tags:
- raise FileNotFoundError(f"No checkpoints found in {root}")
- tags.sort(key=lambda d: os.path.getmtime(os.path.join(root, d)), reverse=True)
- model_tag = tags[0]
- return os.path.join(root, model_tag)
-
-
-def _torch_module():
- try:
- import torch
- except Exception as exc: # pragma: no cover - optional dependency
- raise RuntimeError("Loading nanoai .pt checkpoints requires torch") from exc
- return torch
-
-
-def _load_json(path: str | os.PathLike) -> dict[str, Any]:
- with open(path, "r", encoding="utf-8") as f:
- return json.load(f)
-
-
-def _save_json(path: str | os.PathLike, data: dict[str, Any]) -> None:
- with open(path, "w", encoding="utf-8") as f:
- json.dump(data, f, indent=2)
-
-
-def load_pt_state(checkpoint_dir: str | os.PathLike, step: int | None = None) -> tuple[dict[str, Any], dict[str, Any], int]:
- torch = _torch_module()
- if step is None:
- step = find_last_step(checkpoint_dir)
- model_path = os.path.join(str(checkpoint_dir), f"model_{step:06d}.pt")
- meta_path = os.path.join(str(checkpoint_dir), f"meta_{step:06d}.json")
- state = torch.load(model_path, map_location="cpu")
- state = {
- k.removeprefix("_orig_mod."): (v.float() if torch.is_tensor(v) else v)
- for k, v in state.items()
- }
- meta = _load_json(meta_path)
- config = dict(meta["model_config"])
- config.setdefault("window_pattern", "L")
- n_layer = int(config["n_layer"])
- if "resid_lambdas" not in state:
- state["resid_lambdas"] = torch.ones(n_layer)
- if "x0_lambdas" not in state:
- state["x0_lambdas"] = torch.zeros(n_layer)
- return state, meta, step
-
-
-def _to_mx(value):
- if hasattr(value, "detach"):
- return mx.array(value.detach().cpu().float().contiguous().numpy())
- return value
-
-
-def _assign_array(obj, attr: str, value) -> None:
- setattr(obj, attr, _to_mx(value))
-
-
-def _set_linear_weight(linear: Linear, value) -> None:
- linear.packed = False
- linear.weight = _to_mx(value)
-
-
-def load_state_into_model(model: GPT, state: dict[str, Any]) -> None:
- _assign_array(model.wte, "weight", state["transformer.wte.weight"])
- _set_linear_weight(model.lm_head, state["lm_head.weight"])
- _assign_array(model, "resid_lambdas", state["resid_lambdas"])
- _assign_array(model, "x0_lambdas", state["x0_lambdas"])
- _assign_array(model, "smear_lambda", state["smear_lambda"])
- _assign_array(model, "backout_lambda", state["backout_lambda"])
- _set_linear_weight(model.smear_gate, state["smear_gate.weight"])
- for i, block in enumerate(model.blocks):
- prefix = f"transformer.h.{i}"
- _set_linear_weight(block.attn.c_q, state[f"{prefix}.attn.c_q.weight"])
- _set_linear_weight(block.attn.c_k, state[f"{prefix}.attn.c_k.weight"])
- _set_linear_weight(block.attn.c_v, state[f"{prefix}.attn.c_v.weight"])
- _set_linear_weight(block.attn.c_proj, state[f"{prefix}.attn.c_proj.weight"])
- if block.attn.ve_gate is not None:
- _set_linear_weight(block.attn.ve_gate, state[f"{prefix}.attn.ve_gate.weight"])
- _set_linear_weight(block.mlp.c_fc, state[f"{prefix}.mlp.c_fc.weight"])
- _set_linear_weight(block.mlp.c_proj, state[f"{prefix}.mlp.c_proj.weight"])
- for i, ve in model.value_embeds.items():
- _assign_array(ve, "weight", state[f"value_embeds.{i}.weight"])
- mx.eval(model.parameters())
-
-
-def build_model_from_pt(
- checkpoint_dir: str | os.PathLike,
- *,
- step: int | None = None,
- quantization: str = "dense",
-) -> tuple[GPT, dict[str, Any], int]:
- state, meta, resolved_step = load_pt_state(checkpoint_dir, step)
- config = GPTConfig.from_dict(meta["model_config"])
- model = GPT(config, quantization=quantization)
- load_state_into_model(model, state)
- return model, meta, resolved_step
-
-
-def _linear_state(name: str, linear: Linear, *, packed: bool, meta: dict[str, Any], out: dict[str, Any]) -> None:
- weight_key = f"{name}.weight"
- if packed and hasattr(linear, "weight") and linear.nvfp4_eligible and linear.quantization == "nvfp4":
- qweight, scales, qbiases = quantize_weight(linear.weight)
- out[f"{weight_key}.qweight"] = qweight
- out[f"{weight_key}.scales"] = scales
- if qbiases is not None:
- out[f"{weight_key}.qbiases"] = qbiases
- meta["tensors"][weight_key] = {
- "storage": "nvfp4_packed",
- "qweight": f"{weight_key}.qweight",
- "scales": f"{weight_key}.scales",
- "qbiases": f"{weight_key}.qbiases" if qbiases is not None else None,
- }
- else:
- out[weight_key] = linear.weight
- meta["tensors"][weight_key] = {"storage": "dense", "weight": weight_key}
-
-
-def model_state_dict(model: GPT, *, packed_nvfp4: bool = False) -> tuple[dict[str, Any], dict[str, Any]]:
- out: dict[str, Any] = {}
- meta: dict[str, Any] = {"tensors": {}}
- out["transformer.wte.weight"] = model.wte.weight
- meta["tensors"]["transformer.wte.weight"] = {"storage": "dense", "weight": "transformer.wte.weight"}
- _linear_state("lm_head", model.lm_head, packed=packed_nvfp4, meta=meta, out=out)
- out["resid_lambdas"] = model.resid_lambdas
- out["x0_lambdas"] = model.x0_lambdas
- out["smear_lambda"] = model.smear_lambda
- out["backout_lambda"] = model.backout_lambda
- _linear_state("smear_gate", model.smear_gate, packed=packed_nvfp4, meta=meta, out=out)
- for i, block in enumerate(model.blocks):
- prefix = f"transformer.h.{i}"
- _linear_state(f"{prefix}.attn.c_q", block.attn.c_q, packed=packed_nvfp4, meta=meta, out=out)
- _linear_state(f"{prefix}.attn.c_k", block.attn.c_k, packed=packed_nvfp4, meta=meta, out=out)
- _linear_state(f"{prefix}.attn.c_v", block.attn.c_v, packed=packed_nvfp4, meta=meta, out=out)
- _linear_state(f"{prefix}.attn.c_proj", block.attn.c_proj, packed=packed_nvfp4, meta=meta, out=out)
- if block.attn.ve_gate is not None:
- _linear_state(f"{prefix}.attn.ve_gate", block.attn.ve_gate, packed=packed_nvfp4, meta=meta, out=out)
- _linear_state(f"{prefix}.mlp.c_fc", block.mlp.c_fc, packed=packed_nvfp4, meta=meta, out=out)
- _linear_state(f"{prefix}.mlp.c_proj", block.mlp.c_proj, packed=packed_nvfp4, meta=meta, out=out)
- for i, ve in model.value_embeds.items():
- key = f"value_embeds.{i}.weight"
- out[key] = ve.weight
- meta["tensors"][key] = {"storage": "dense", "weight": key}
- return out, meta
-
-
-def save_mlx_checkpoint(
- model: GPT,
- out_dir: str | os.PathLike,
- *,
- step: int,
- source_meta: dict[str, Any] | None = None,
- quantization: str = "dense",
-) -> tuple[str, str]:
- os.makedirs(out_dir, exist_ok=True)
- packed = quantization == "nvfp4_packed"
- weights, tensor_meta = model_state_dict(model, packed_nvfp4=packed)
- weights_path = os.path.join(str(out_dir), f"step_{step:06d}.safetensors")
- meta_path = os.path.join(str(out_dir), f"step_{step:06d}_meta.json")
- mx.save_safetensors(weights_path, weights)
- meta = {
- "format_version": FORMAT_VERSION,
- "step": step,
- "model_config": model.config.asdict(),
- "quantization": quantization,
- "quantization_policy": model.quantization,
- "training_quantization": model.quantization,
- "source_meta": source_meta or {},
- **tensor_meta,
- }
- _save_json(meta_path, meta)
- return weights_path, meta_path
-
-
-def _assign_packed(linear: Linear, qweight, scales, qbiases=None) -> None:
- linear.packed = True
- linear.qweight = qweight
- linear.scales = scales
- linear.qbiases = qbiases
- linear.quantization = "nvfp4"
-
-
-def _load_linear_from_mlx(weights: dict[str, Any], tensor_meta: dict[str, Any], name: str, linear: Linear) -> None:
- key = f"{name}.weight"
- info = tensor_meta[key]
- if info["storage"] == "nvfp4_packed":
- _assign_packed(
- linear,
- weights[info["qweight"]],
- weights[info["scales"]],
- weights.get(info.get("qbiases")) if info.get("qbiases") else None,
- )
- else:
- _set_linear_weight(linear, weights[info["weight"]])
-
-
-def load_mlx_checkpoint(
- checkpoint_dir: str | os.PathLike,
- *,
- step: int | None = None,
- quantization: str | None = None,
-) -> tuple[GPT, dict[str, Any], int]:
- if step is None:
- files = [f for f in glob.glob(os.path.join(str(checkpoint_dir), "step_*.safetensors")) if "_optim" not in f]
- if not files:
- raise FileNotFoundError(f"No step_*.safetensors found in {checkpoint_dir}")
- step = int(max(os.path.basename(f).split("_")[-1].split(".")[0] for f in files))
- weights_path = os.path.join(str(checkpoint_dir), f"step_{step:06d}.safetensors")
- meta_path = os.path.join(str(checkpoint_dir), f"step_{step:06d}_meta.json")
- weights = mx.load(weights_path)
- meta = _load_json(meta_path)
- default_quant = meta.get("quantization_policy") or meta.get("training_quantization")
- model_quant = quantization or (default_quant if default_quant in QUANTIZATION_MODES else "dense")
- model = GPT(GPTConfig.from_dict(meta["model_config"]), quantization=model_quant)
- tensor_meta = meta["tensors"]
- _assign_array(model.wte, "weight", weights["transformer.wte.weight"])
- _load_linear_from_mlx(weights, tensor_meta, "lm_head", model.lm_head)
- _assign_array(model, "resid_lambdas", weights["resid_lambdas"])
- _assign_array(model, "x0_lambdas", weights["x0_lambdas"])
- _assign_array(model, "smear_lambda", weights["smear_lambda"])
- _assign_array(model, "backout_lambda", weights["backout_lambda"])
- _load_linear_from_mlx(weights, tensor_meta, "smear_gate", model.smear_gate)
- for i, block in enumerate(model.blocks):
- prefix = f"transformer.h.{i}"
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.attn.c_q", block.attn.c_q)
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.attn.c_k", block.attn.c_k)
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.attn.c_v", block.attn.c_v)
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.attn.c_proj", block.attn.c_proj)
- if block.attn.ve_gate is not None:
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.attn.ve_gate", block.attn.ve_gate)
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.mlp.c_fc", block.mlp.c_fc)
- _load_linear_from_mlx(weights, tensor_meta, f"{prefix}.mlp.c_proj", block.mlp.c_proj)
- for i, ve in model.value_embeds.items():
- _assign_array(ve, "weight", weights[f"value_embeds.{i}.weight"])
- mx.eval(model.parameters())
- return model, meta, step
-
-
-def load_mlx_model(
- *,
- source: str = "sft",
- model_tag: str | None = None,
- checkpoint_dir: str | os.PathLike | None = None,
- step: int | None = None,
- format: str = "pt",
- quantization: str = "dense",
-):
- if quantization not in QUANTIZATION_MODES:
- raise ValueError(f"unknown quantization {quantization!r}; expected one of {QUANTIZATION_MODES}")
- if format == "pt":
- ckpt_dir = str(checkpoint_dir) if checkpoint_dir is not None else resolve_pt_checkpoint_dir(source, model_tag)
- model, meta, resolved_step = build_model_from_pt(ckpt_dir, step=step, quantization=quantization)
- elif format == "mlx":
- if checkpoint_dir is None:
- root = Path(get_base_dir()) / "mlx_checkpoints"
- tag = model_tag or source
- checkpoint_dir = root / tag
- model, meta, resolved_step = load_mlx_checkpoint(checkpoint_dir, step=step, quantization=quantization)
- else:
- raise ValueError("format must be 'pt' or 'mlx'")
- from nanoai.tokenizer import get_tokenizer
-
- tokenizer = get_tokenizer()
- return model, tokenizer, meta, resolved_step
diff --git a/vendor/nanoai/mlx/engine.py b/vendor/nanoai/mlx/engine.py
deleted file mode 100644
index 05195664162ccd7b63b2a89300de0aaa85fd4fa3..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/engine.py
+++ /dev/null
@@ -1,303 +0,0 @@
-"""Inference engine for MLX nanoai models."""
-
-from __future__ import annotations
-
-import signal
-import warnings
-from collections import deque
-from contextlib import contextmanager
-
-import mlx.core as mx
-import numpy as np
-
-
-@contextmanager
-def timeout(duration, formula):
- def timeout_handler(signum, frame):
- raise Exception(f"'{formula}': timed out after {duration} seconds")
-
- signal.signal(signal.SIGALRM, timeout_handler)
- signal.alarm(duration)
- yield
- signal.alarm(0)
-
-
-def eval_with_timeout(formula, max_time=3):
- try:
- with timeout(max_time, formula):
- with warnings.catch_warnings():
- warnings.simplefilter("ignore", SyntaxWarning)
- return eval(formula, {"__builtins__": {}}, {})
- except Exception:
- signal.alarm(0)
- return None
-
-
-def use_calculator(expr):
- expr = expr.replace(",", "")
- if all(x in "0123456789*+-/.() " for x in expr):
- if "**" in expr:
- return None
- return eval_with_timeout(expr)
- allowed_chars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789'\"()._ "
- if not all(x in allowed_chars for x in expr):
- return None
- dangerous_patterns = [
- "__",
- "import",
- "exec",
- "eval",
- "compile",
- "open",
- "file",
- "input",
- "raw_input",
- "globals",
- "locals",
- "vars",
- "dir",
- "getattr",
- "setattr",
- "delattr",
- "hasattr",
- ]
- if any(pattern in expr.lower() for pattern in dangerous_patterns):
- return None
- if ".count(" not in expr:
- return None
- return eval_with_timeout(expr)
-
-
-class KVCache:
- def __init__(self, n_layers: int, window_sizes=None):
- self.n_layers = n_layers
- self.keys = [None] * n_layers
- self.values = [None] * n_layers
- self.offset = 0
- self.window_sizes = window_sizes
- self.prev_embedding = None
-
- def update(self, layer_idx: int, k, v):
- if self.keys[layer_idx] is None:
- self.keys[layer_idx] = k
- self.values[layer_idx] = v
- else:
- self.keys[layer_idx] = mx.concatenate([self.keys[layer_idx], k], axis=2)
- self.values[layer_idx] = mx.concatenate([self.values[layer_idx], v], axis=2)
-
- if self.window_sizes is not None:
- left = self.window_sizes[layer_idx][0]
- if left >= 0:
- keep = left + 1
- if self.keys[layer_idx].shape[2] > keep:
- self.keys[layer_idx] = self.keys[layer_idx][:, :, -keep:, :]
- self.values[layer_idx] = self.values[layer_idx][:, :, -keep:, :]
-
- if layer_idx == self.n_layers - 1:
- self.offset += k.shape[2]
- return self.keys[layer_idx], self.values[layer_idx]
-
- def reset(self):
- self.keys = [None] * self.n_layers
- self.values = [None] * self.n_layers
- self.offset = 0
- self.prev_embedding = None
-
- def clone_repeated(self, repeats: int) -> "KVCache":
- other = KVCache(self.n_layers, window_sizes=self.window_sizes)
- other.offset = self.offset
- other.keys = [mx.repeat(k, repeats, axis=0) if k is not None else None for k in self.keys]
- other.values = [mx.repeat(v, repeats, axis=0) if v is not None else None for v in self.values]
- if self.prev_embedding is not None:
- other.prev_embedding = mx.repeat(self.prev_embedding, repeats, axis=0)
- return other
-
-
-def _special_id(tokenizer, name: str, default: int = -1) -> int:
- try:
- return tokenizer.encode_special(name)
- except Exception:
- return default
-
-
-def _assistant_end_id(tokenizer) -> int:
- if hasattr(tokenizer, "assistant_end_id"):
- return tokenizer.assistant_end_id()
- return _special_id(tokenizer, "<|assistant_end|>")
-
-
-def apply_repetition_penalty(logits, recent_per_row, penalty: float, exempt_ids=None):
- if penalty == 1.0:
- return logits
- rows = []
- logits_np = np.array(logits)
- exempt_ids = exempt_ids or set()
- for b, recent in enumerate(recent_per_row):
- row = logits_np[b].copy()
- for token in set(recent) - exempt_ids:
- if token < 0 or token >= row.shape[0]:
- continue
- row[token] = row[token] / penalty if row[token] > 0 else row[token] * penalty
- rows.append(row)
- return mx.array(np.stack(rows, axis=0))
-
-
-def sample_next_token(logits, *, temperature: float = 1.0, top_k: int | None = None):
- assert temperature >= 0.0
- if temperature == 0.0:
- return mx.argmax(logits, axis=-1, keepdims=True)
- if top_k is not None and top_k > 0:
- k = min(top_k, logits.shape[-1])
- vals = mx.topk(logits, k=k, axis=-1)
- threshold = vals[:, -1:]
- logits = mx.where(logits >= threshold, logits, mx.array(float("-inf"), dtype=logits.dtype))
- return mx.random.categorical(logits / temperature, axis=-1)[:, None]
-
-
-class RowState:
- def __init__(self, current_tokens=None):
- self.current_tokens = current_tokens or []
- self.forced_tokens = deque()
- self.in_python_block = False
- self.python_expr_tokens = []
- self.completed = False
-
-
-class Engine:
- def __init__(self, model, tokenizer):
- self.model = model
- self.tokenizer = tokenizer
-
- def generate(
- self,
- tokens,
- num_samples: int = 1,
- max_tokens: int | None = None,
- temperature: float = 1.0,
- top_k: int | None = None,
- seed: int = 42,
- repetition_penalty: float = 1.0,
- recent_window: int = 0,
- exempt_ids=None,
- ):
- assert isinstance(tokens, list) and isinstance(tokens[0], int), "expecting list of ints"
- mx.random.seed(seed)
-
- python_start = _special_id(self.tokenizer, "<|python_start|>")
- python_end = _special_id(self.tokenizer, "<|python_end|>")
- output_start = _special_id(self.tokenizer, "<|output_start|>")
- output_end = _special_id(self.tokenizer, "<|output_end|>")
- assistant_end = _assistant_end_id(self.tokenizer)
- bos = self.tokenizer.get_bos_token_id()
-
- kv_cache = KVCache(self.model.config.n_layer, window_sizes=self.model.window_sizes)
- ids = mx.array([tokens], dtype=mx.int32)
- logits = self.model(ids, kv_cache=kv_cache)[:, -1, :]
- mx.eval(logits)
- if num_samples > 1:
- logits = mx.repeat(logits, num_samples, axis=0)
- kv_cache = kv_cache.clone_repeated(num_samples)
-
- row_states = [RowState(tokens.copy()) for _ in range(num_samples)]
- use_rep = repetition_penalty != 1.0 and recent_window > 0
- recent_per_row = [deque(maxlen=recent_window) for _ in range(num_samples)] if use_rep else None
-
- num_generated = 0
- while True:
- if max_tokens is not None and num_generated >= max_tokens:
- break
- if all(state.completed for state in row_states):
- break
-
- step_logits = (
- apply_repetition_penalty(logits, recent_per_row, repetition_penalty, exempt_ids)
- if use_rep
- else logits
- )
- next_ids = sample_next_token(step_logits, temperature=temperature, top_k=top_k)
- sampled_tokens = next_ids[:, 0].tolist()
-
- token_column = []
- token_masks = []
- for i, state in enumerate(row_states):
- is_forced = len(state.forced_tokens) > 0
- token = state.forced_tokens.popleft() if is_forced else int(sampled_tokens[i])
- token_masks.append(0 if is_forced else 1)
- token_column.append(token)
- state.current_tokens.append(token)
- if use_rep and not state.completed:
- recent_per_row[i].append(token)
- if token == assistant_end or token == bos:
- state.completed = True
- if token == python_start:
- state.in_python_block = True
- state.python_expr_tokens = []
- elif token == python_end and state.in_python_block:
- state.in_python_block = False
- if state.python_expr_tokens:
- expr = self.tokenizer.decode(state.python_expr_tokens)
- result = use_calculator(expr)
- if result is not None:
- result_tokens = self.tokenizer.encode(str(result))
- state.forced_tokens.append(output_start)
- state.forced_tokens.extend(result_tokens)
- state.forced_tokens.append(output_end)
- state.python_expr_tokens = []
- elif state.in_python_block:
- state.python_expr_tokens.append(token)
-
- yield token_column, token_masks
- num_generated += 1
-
- ids = mx.array(token_column, dtype=mx.int32)[:, None]
- logits = self.model(ids, kv_cache=kv_cache)[:, -1, :]
- mx.eval(logits)
-
- def generate_batch(self, tokens, num_samples: int = 1, **kwargs):
- assistant_end = _assistant_end_id(self.tokenizer)
- bos = self.tokenizer.get_bos_token_id()
- results = [tokens.copy() for _ in range(num_samples)]
- masks = [[0] * len(tokens) for _ in range(num_samples)]
- completed = [False] * num_samples
- for token_column, token_masks in self.generate(tokens, num_samples=num_samples, **kwargs):
- for i, (token, mask) in enumerate(zip(token_column, token_masks)):
- if completed[i]:
- continue
- if token == assistant_end or token == bos:
- completed[i] = True
- else:
- results[i].append(token)
- masks[i].append(mask)
- if all(completed):
- break
- return results, masks
-
- def generate_batched_independent(
- self,
- prompts: list[list[int]],
- seeds: list[int],
- *,
- max_new_tokens: int,
- stop_tokens: set,
- temperature: float = 1.0,
- top_k: int | None = None,
- ) -> tuple[list[list[int]], list[bool]]:
- outputs = []
- finished = []
- for prompt, seed in zip(prompts, seeds):
- full, _ = self.generate_batch(
- prompt,
- num_samples=1,
- max_tokens=max_new_tokens,
- temperature=temperature,
- top_k=top_k,
- seed=int(seed),
- )
- new_tokens = full[0][len(prompt) :]
- hit_stop = len(new_tokens) < max_new_tokens
- if hit_stop:
- marker = _assistant_end_id(self.tokenizer)
- new_tokens = new_tokens + [marker]
- outputs.append(new_tokens)
- finished.append(hit_stop)
- return outputs, finished
diff --git a/vendor/nanoai/mlx/gpt.py b/vendor/nanoai/mlx/gpt.py
deleted file mode 100644
index 21014efa035e72afc03bcd344774a46aa11d1d0a..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/gpt.py
+++ /dev/null
@@ -1,401 +0,0 @@
-"""nanoai GPT implemented in MLX, including NVFP4 training forward.
-
-The model mirrors ``nanoai/gpt.py`` in this repository. The dense path is the
-numerical reference; the NVFP4 path keeps FP master weights and uses NVFP4
-matmul in the forward pass with a straight-through estimator for training.
-"""
-
-from __future__ import annotations
-
-import math
-from dataclasses import dataclass
-from typing import Iterable
-
-import mlx.core as mx
-import mlx.nn as nn
-from mlx.utils import tree_flatten
-
-from nanoai.mlx.nvfp4 import (
- QUANTIZATION_MODES,
- nvfp4_linear_ste,
- quantize_weight,
- quantized_matmul,
- should_nvfp4_quantize_linear,
- should_nvfp4_quantize_linear_shape,
-)
-
-
-# Single shared config (torch + MLX). from_dict tolerates torch-saved metas that
-# carry msa_* fields (MLX ignores them — no MSA kernel) and drops unknown keys,
-# so a torch checkpoint loads here and an MLX-saved meta loads in torch.
-# Re-exported so `from nanoai.mlx.gpt import GPTConfig` keeps working.
-from nanoai.model_config import GPTConfig, value_embed_table_map # noqa: E402
-
-
-def norm(x):
- return x * mx.rsqrt(mx.mean(mx.square(x), axis=-1, keepdims=True) + mx.finfo(x.dtype).eps)
-
-
-def has_ve(layer_idx: int, n_layer: int) -> bool:
- return layer_idx % 2 == (n_layer - 1) % 2
-
-
-def apply_rotary_emb(x, cos, sin):
- d = x.shape[-1] // 2
- x1, x2 = x[..., :d], x[..., d:]
- y1 = x1 * cos + x2 * sin
- y2 = x1 * (-sin) + x2 * cos
- return mx.concatenate([y1, y2], axis=-1)
-
-
-class Linear(nn.Module):
- """Bias-free Linear with optional NVFP4 forward.
-
- ``quantization="nvfp4"`` uses FP master weights and STE gradients.
- ``packed=True`` stores only a packed NVFP4 weight and is inference-only.
- """
-
- def __init__(
- self,
- in_features: int,
- out_features: int,
- *,
- quantization: str = "dense",
- packed: bool = False,
- ):
- super().__init__()
- self.in_features = in_features
- self.out_features = out_features
- self.quantization = quantization
- self.packed = packed
- if packed:
- self.qweight = mx.zeros((out_features, in_features // 8), dtype=mx.uint32)
- self.scales = mx.zeros((out_features, in_features // 16), dtype=mx.float32)
- self.qbiases = None
- else:
- self.weight = mx.zeros((out_features, in_features), dtype=mx.float32)
-
- @property
- def nvfp4_eligible(self) -> bool:
- shape = (self.out_features, self.in_features)
- return should_nvfp4_quantize_linear_shape(shape)
-
- def pack_from_weight(self) -> None:
- if not self.nvfp4_eligible:
- raise ValueError(f"Linear shape {(self.out_features, self.in_features)} is not NVFP4 eligible")
- self.qweight, self.scales, self.qbiases = quantize_weight(self.weight)
- self.packed = True
-
- def __call__(self, x):
- if self.packed:
- return quantized_matmul(x, self.qweight, self.scales, self.qbiases, transpose=True)
- if self.quantization == "nvfp4" and self.nvfp4_eligible:
- return nvfp4_linear_ste(x, self.weight)
- return mx.matmul(x, mx.transpose(self.weight, (1, 0)))
-
-
-class CausalSelfAttention(nn.Module):
- def __init__(self, config: GPTConfig, layer_idx: int, *, quantization: str = "dense"):
- super().__init__()
- self.layer_idx = layer_idx
- self.n_head = config.n_head
- self.n_kv_head = config.n_kv_head
- self.n_embd = config.n_embd
- self.head_dim = self.n_embd // self.n_head
- assert self.n_embd % self.n_head == 0
- assert self.n_kv_head <= self.n_head and self.n_head % self.n_kv_head == 0
- self.c_q = Linear(self.n_embd, self.n_head * self.head_dim, quantization=quantization)
- self.c_k = Linear(self.n_embd, self.n_kv_head * self.head_dim, quantization=quantization)
- self.c_v = Linear(self.n_embd, self.n_kv_head * self.head_dim, quantization=quantization)
- self.c_proj = Linear(self.n_embd, self.n_embd, quantization=quantization)
- self.ve_gate_channels = 12
- self.ve_gate = (
- Linear(self.ve_gate_channels, self.n_kv_head, quantization="dense")
- if has_ve(layer_idx, config.n_layer)
- else None
- )
-
- def __call__(self, x, ve, cos_sin, window_size, kv_cache):
- B, T, _C = x.shape
- q = self.c_q(x).reshape(B, T, self.n_head, self.head_dim)
- k = self.c_k(x).reshape(B, T, self.n_kv_head, self.head_dim)
- v = self.c_v(x).reshape(B, T, self.n_kv_head, self.head_dim)
-
- if ve is not None:
- ve = ve.reshape(B, T, self.n_kv_head, self.head_dim)
- gate = 3.0 * mx.sigmoid(self.ve_gate(x[..., : self.ve_gate_channels]))
- v = v + mx.expand_dims(gate, axis=-1) * ve
-
- cos, sin = cos_sin
- q = apply_rotary_emb(q, cos, sin)
- k = apply_rotary_emb(k, cos, sin)
- q = norm(q) * 1.2
- k = norm(k) * 1.2
-
- q = q.transpose(0, 2, 1, 3)
- k = k.transpose(0, 2, 1, 3)
- v = v.transpose(0, 2, 1, 3)
- if kv_cache is not None:
- k, v = kv_cache.update(self.layer_idx, k, v)
-
- if self.n_kv_head != self.n_head:
- repeat = self.n_head // self.n_kv_head
- k = mx.repeat(k, repeat, axis=1)
- v = mx.repeat(v, repeat, axis=1)
-
- mask = _attention_mask(T, k.shape[2], kv_cache.offset if kv_cache is not None else 0, window_size)
- y = mx.fast.scaled_dot_product_attention(
- q,
- k,
- v,
- scale=1.0 / math.sqrt(self.head_dim),
- mask=mask,
- )
- y = y.transpose(0, 2, 1, 3).reshape(B, T, -1)
- return self.c_proj(y)
-
-
-class MLP(nn.Module):
- def __init__(self, config: GPTConfig, *, quantization: str = "dense"):
- super().__init__()
- self.c_fc = Linear(config.n_embd, 4 * config.n_embd, quantization=quantization)
- self.c_proj = Linear(4 * config.n_embd, config.n_embd, quantization=quantization)
-
- def __call__(self, x):
- x = self.c_fc(x)
- x = mx.square(mx.maximum(x, 0.0))
- return self.c_proj(x)
-
-
-class Block(nn.Module):
- def __init__(self, config: GPTConfig, layer_idx: int, *, quantization: str = "dense"):
- super().__init__()
- self.attn = CausalSelfAttention(config, layer_idx, quantization=quantization)
- self.mlp = MLP(config, quantization=quantization)
-
- def __call__(self, x, ve, cos_sin, window_size, kv_cache):
- x = x + self.attn(norm(x), ve, cos_sin, window_size, kv_cache)
- x = x + self.mlp(norm(x))
- return x
-
-
-class GPT(nn.Module):
- def __init__(
- self,
- config: GPTConfig,
- *,
- quantization: str = "dense",
- pad_vocab_size_to: int = 64,
- ):
- super().__init__()
- assert quantization in QUANTIZATION_MODES
- self.config = config
- self.quantization = quantization
- self.padded_vocab_size = (
- (config.vocab_size + pad_vocab_size_to - 1) // pad_vocab_size_to
- ) * pad_vocab_size_to
- self.transformer = {
- "wte": nn.Embedding(self.padded_vocab_size, config.n_embd),
- }
- self.blocks = [Block(config, i, quantization="dense") for i in range(config.n_layer)]
- self.lm_head = Linear(config.n_embd, self.padded_vocab_size, quantization="dense")
- self.resid_lambdas = mx.ones((config.n_layer,), dtype=mx.float32)
- self.x0_lambdas = mx.zeros((config.n_layer,), dtype=mx.float32)
- self.smear_gate = Linear(24, 1, quantization="dense")
- self.smear_lambda = mx.zeros((1,), dtype=mx.float32)
- self.backout_lambda = mx.full((1,), 0.2, dtype=mx.float32)
- head_dim = config.n_embd // config.n_head
- kv_dim = config.n_kv_head * head_dim
- # Value-embedding tables, shared per config.ve_n_unique (U-Net). Keyed
- # identically to the torch backend via the shared map so checkpoints load
- # cross-backend. self.ve_tid maps layer_idx -> table key (absent => no VE).
- ve_keys, self.ve_tid = value_embed_table_map(config.n_layer, config.ve_n_unique)
- self.value_embeds = {
- k: nn.Embedding(self.padded_vocab_size, kv_dim) for k in ve_keys
- }
- self.window_sizes = self._compute_window_sizes(config)
- self.rotary_seq_len = config.sequence_len * 10
- self.cos, self.sin = self._precompute_rotary_embeddings(self.rotary_seq_len, head_dim)
- self.set_quantization(quantization)
-
- @property
- def wte(self):
- return self.transformer["wte"]
-
- def init_weights(self):
- n_embd = self.config.n_embd
- s = 3**0.5 * n_embd**-0.5
- self.wte.weight = mx.random.normal(self.wte.weight.shape) * 0.8
- self.lm_head.weight = mx.random.normal(self.lm_head.weight.shape) * 0.001
- for block in self.blocks:
- block.attn.c_q.weight = mx.random.uniform(-s, s, block.attn.c_q.weight.shape)
- block.attn.c_k.weight = mx.random.uniform(-s, s, block.attn.c_k.weight.shape)
- block.attn.c_v.weight = mx.random.uniform(-s, s, block.attn.c_v.weight.shape)
- block.attn.c_proj.weight = mx.zeros_like(block.attn.c_proj.weight)
- block.mlp.c_fc.weight = mx.random.uniform(-s * 0.4, s * 0.4, block.mlp.c_fc.weight.shape)
- block.mlp.c_proj.weight = mx.zeros_like(block.mlp.c_proj.weight)
- if block.attn.ve_gate is not None:
- block.attn.ve_gate.weight = mx.random.uniform(0.0, 0.02, block.attn.ve_gate.weight.shape)
- n_layer = self.config.n_layer
- self.resid_lambdas = mx.array(
- [1.15 - (0.10 * i / max(n_layer - 1, 1)) for i in range(n_layer)],
- dtype=mx.float32,
- )
- self.x0_lambdas = mx.array(
- [0.20 - (0.15 * i / max(n_layer - 1, 1)) for i in range(n_layer)],
- dtype=mx.float32,
- )
- self.smear_gate.weight = mx.random.uniform(0.0, 0.02, self.smear_gate.weight.shape)
- self.smear_lambda = mx.zeros((1,), dtype=mx.float32)
- self.backout_lambda = mx.full((1,), 0.2, dtype=mx.float32)
- for ve in self.value_embeds.values():
- ve.weight = mx.random.uniform(-s, s, ve.weight.shape)
- self.cos, self.sin = self._precompute_rotary_embeddings(
- self.rotary_seq_len,
- self.config.n_embd // self.config.n_head,
- )
-
- def _precompute_rotary_embeddings(self, seq_len: int, head_dim: int, base: int = 100000):
- channel_range = mx.arange(0, head_dim, 2, dtype=mx.float32)
- inv_freq = 1.0 / (base ** (channel_range / head_dim))
- t = mx.arange(seq_len, dtype=mx.float32)
- freqs = t[:, None] * inv_freq[None, :]
- cos = mx.cos(freqs)[None, :, None, :]
- sin = mx.sin(freqs)[None, :, None, :]
- return cos, sin
-
- def _compute_window_sizes(self, config: GPTConfig) -> list[tuple[int, int]]:
- pattern = config.window_pattern.upper()
- assert all(c in "SLMDH" for c in pattern), f"Invalid window_pattern: {pattern}. Use only S, L, M, D, H."
- long_window = config.sequence_len
- short_window = -(-long_window // 4 // 128) * 128
- # MLX has no sparse kernels yet. Sparse letters keep full context and
- # execute through the existing dense attention path.
- mapping = {
- "L": (long_window, 0),
- "S": (short_window, 0),
- "M": (long_window, 0),
- "D": (long_window, 0),
- "H": (long_window, 0),
- }
- windows = [mapping[pattern[i % len(pattern)]] for i in range(config.n_layer)]
- windows[-1] = (long_window, 0)
- return windows
-
- def get_device(self):
- return "mlx"
-
- def num_scaling_params(self) -> dict[str, int]:
- wte = self.wte.weight.size
- value_embeds = sum(ve.weight.size for ve in self.value_embeds.values())
- lm_head = self.lm_head.weight.size if hasattr(self.lm_head, "weight") else 0
- transformer_matrices = 0
- for block in self.blocks:
- for _path, param in tree_flatten(block.parameters()):
- if hasattr(param, "size"):
- transformer_matrices += param.size
- scalars = self.resid_lambdas.size + self.x0_lambdas.size
- scalars += self.smear_gate.weight.size + self.smear_lambda.size + self.backout_lambda.size
- total = wte + value_embeds + lm_head + transformer_matrices + scalars
- return {
- "wte": wte,
- "value_embeds": value_embeds,
- "lm_head": lm_head,
- "transformer_matrices": transformer_matrices,
- "scalars": scalars,
- "total": total,
- }
-
- def iter_linears(self) -> Iterable[tuple[str, Linear]]:
- yield "lm_head", self.lm_head
- yield "smear_gate", self.smear_gate
- for i, block in enumerate(self.blocks):
- prefix = f"transformer.h.{i}"
- yield f"{prefix}.attn.c_q", block.attn.c_q
- yield f"{prefix}.attn.c_k", block.attn.c_k
- yield f"{prefix}.attn.c_v", block.attn.c_v
- yield f"{prefix}.attn.c_proj", block.attn.c_proj
- if block.attn.ve_gate is not None:
- yield f"{prefix}.attn.ve_gate", block.attn.ve_gate
- yield f"{prefix}.mlp.c_fc", block.mlp.c_fc
- yield f"{prefix}.mlp.c_proj", block.mlp.c_proj
-
- def set_quantization(self, quantization: str) -> None:
- assert quantization in QUANTIZATION_MODES
- self.quantization = quantization
- for name, linear in self.iter_linears():
- shape = (linear.out_features, linear.in_features)
- linear.quantization = (
- "nvfp4"
- if should_nvfp4_quantize_linear(name, shape, quantization)
- else "dense"
- )
-
- def __call__(self, idx, targets=None, kv_cache=None, loss_reduction: str = "mean"):
- B, T = idx.shape
- T0 = 0 if kv_cache is None else kv_cache.offset
- assert T0 + T <= self.cos.shape[1], f"rotary cache too short for pos {T0}+{T}"
- cos_sin = self.cos[:, T0 : T0 + T], self.sin[:, T0 : T0 + T]
-
- x = self.wte(idx)
- x = norm(x)
- if kv_cache is None:
- assert T > 1, "training/prefill forward pass should have T > 1"
- gate = self.smear_lambda.astype(x.dtype) * mx.sigmoid(self.smear_gate(x[:, 1:, :24]))
- x = mx.concatenate([x[:, :1], x[:, 1:] + gate * x[:, :-1]], axis=1)
- else:
- prev = kv_cache.prev_embedding
- kv_cache.prev_embedding = x[:, -1:, :]
- if T > 1:
- gate = self.smear_lambda.astype(x.dtype) * mx.sigmoid(self.smear_gate(x[:, 1:, :24]))
- x = mx.concatenate([x[:, :1], x[:, 1:] + gate * x[:, :-1]], axis=1)
- elif prev is not None:
- gate = self.smear_lambda.astype(x.dtype) * mx.sigmoid(self.smear_gate(x[:, :, :24]))
- x = x + gate * prev
-
- x0 = x
- backout_layer = self.config.n_layer // 2
- x_backout = None
- for i, block in enumerate(self.blocks):
- x = self.resid_lambdas[i].astype(x.dtype) * x + self.x0_lambdas[i].astype(x.dtype) * x0
- ve_key = self.ve_tid.get(i)
- ve = self.value_embeds[ve_key](idx).astype(x.dtype) if ve_key is not None else None
- x = block(x, ve, cos_sin, self.window_sizes[i], kv_cache)
- if i == backout_layer:
- x_backout = x
- if x_backout is not None:
- x = x - self.backout_lambda.astype(x.dtype) * x_backout
- x = norm(x)
- logits = self.lm_head(x)[..., : self.config.vocab_size]
- logits = logits.astype(mx.float32)
- logits = 15.0 * mx.tanh(logits / 15.0)
-
- if targets is None:
- return logits
- mask = targets != -1
- targets_safe = mx.where(mask, targets, mx.zeros_like(targets))
- ce = nn.losses.cross_entropy(logits, targets_safe, reduction="none") * mask
- denom = mx.maximum(mx.sum(mask), 1)
- loss = mx.sum(ce) / denom
- if loss_reduction == "sum":
- return mx.sum(ce)
- if loss_reduction == "none":
- return ce
- return loss
-
-
-def _attention_mask(T_query: int, T_key: int, offset: int, window_size: tuple[int, int]):
- if T_query == 1 and T_key > 1:
- return None
- q_pos = mx.arange(offset, offset + T_query)[:, None]
- k_start = offset + T_query - T_key
- k_pos = mx.arange(k_start, k_start + T_key)[None, :]
- allowed = k_pos <= q_pos
- left, _right = window_size
- if left >= 0:
- allowed = allowed & ((q_pos - k_pos) <= left)
- return mx.where(allowed, mx.array(0.0, dtype=mx.float32), mx.array(float("-inf"), dtype=mx.float32))
-
-
-def loss_fn(model: GPT, inputs, targets):
- return model(inputs, targets=targets)
diff --git a/vendor/nanoai/mlx/nvfp4.py b/vendor/nanoai/mlx/nvfp4.py
deleted file mode 100644
index 5cac58810a8725f49b0625d8b7f0c7a87de7788e..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/nvfp4.py
+++ /dev/null
@@ -1,205 +0,0 @@
-"""NVFP4 helpers for MLX.
-
-MLX can execute ``quantized_matmul(..., mode="nvfp4")`` and can propagate
-gradients to the input activation. As of MLX 0.31.2 it does not propagate
-gradients to the packed quantized weight. For training we keep FP master
-weights and use a straight-through estimator (STE): the forward pass executes
-NVFP4 matmul, while the backward pass returns the dense linear weight gradient
-for the master weight.
-"""
-
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import Any
-
-
-NVFP4_GROUP_SIZE = 16
-NVFP4_BITS = 4
-NVFP4_MODE = "nvfp4"
-NVFP4_MIN_LINEAR_DIM = 128
-QUANTIZATION_MODES = ("dense", "nvfp4", "nvfp4-safe")
-
-
-def mlx_available() -> bool:
- try:
- import mlx.core # noqa: F401
- except Exception:
- return False
- return True
-
-
-def require_mlx():
- try:
- import mlx.core as mx
- except Exception as exc: # pragma: no cover - host dependent
- raise RuntimeError(
- "MLX is not available. Install it on Apple Silicon with `pip install mlx`."
- ) from exc
- return mx
-
-
-def assert_mlx_runtime_ready() -> None:
- mx = require_mlx()
- try:
- probe = mx.array([0.0])
- mx.eval(probe)
- except Exception as exc: # pragma: no cover - host dependent
- raise RuntimeError(
- "MLX is importable, but the runtime cannot evaluate arrays. On macOS "
- "this usually means the process cannot access the Metal device."
- ) from exc
-
-
-def can_nvfp4_quantize_shape(shape: tuple[int, ...], group_size: int = NVFP4_GROUP_SIZE) -> bool:
- return len(shape) >= 2 and shape[-1] % group_size == 0
-
-
-def should_nvfp4_quantize_linear_shape(
- shape: tuple[int, ...],
- group_size: int = NVFP4_GROUP_SIZE,
- min_linear_dim: int = NVFP4_MIN_LINEAR_DIM,
-) -> bool:
- """Match nanoai's TransformerEngine linear selection policy."""
-
- if not can_nvfp4_quantize_shape(shape, group_size):
- return False
- out_features, in_features = shape[-2], shape[-1]
- return (
- in_features % group_size == 0
- and out_features % group_size == 0
- and min(in_features, out_features) >= min_linear_dim
- )
-
-
-def should_nvfp4_quantize_linear(
- name: str,
- shape: tuple[int, ...],
- quantization: str,
-) -> bool:
- """Return whether a named Linear should execute through MLX NVFP4.
-
- ``nvfp4`` is the full weight-only MLX path. ``nvfp4-safe`` is a conservative
- inference policy for dense nanoai checkpoints: it quantizes the largest
- output-side matrices that held bounded ChatCORE quality in local ablations
- and leaves attention/value/up-projection matrices dense.
- """
-
- if quantization not in QUANTIZATION_MODES:
- raise ValueError(f"unknown quantization {quantization!r}; expected one of {QUANTIZATION_MODES}")
- if quantization == "dense":
- return False
- if not should_nvfp4_quantize_linear_shape(shape):
- return False
- if quantization == "nvfp4":
- return True
- if quantization == "nvfp4-safe":
- return name == "lm_head" or name.endswith(".mlp.c_proj")
- raise AssertionError(f"unhandled quantization mode {quantization!r}")
-
-
-def _flatten_linear_batch(mx: Any, x, dy):
- in_features = x.shape[-1]
- out_features = dy.shape[-1]
- return (
- x.reshape((-1, in_features)),
- dy.reshape((-1, out_features)),
- )
-
-
-def _make_nvfp4_linear_ste():
- mx = require_mlx()
-
- @mx.custom_function
- def nvfp4_linear_ste(x, weight):
- qweight, scales, *rest = mx.quantize(
- weight,
- group_size=NVFP4_GROUP_SIZE,
- bits=NVFP4_BITS,
- mode=NVFP4_MODE,
- )
- qbiases = rest[0] if rest else None
- return mx.quantized_matmul(
- x,
- qweight,
- scales,
- qbiases,
- transpose=True,
- group_size=NVFP4_GROUP_SIZE,
- bits=NVFP4_BITS,
- mode=NVFP4_MODE,
- )
-
- @nvfp4_linear_ste.vjp
- def nvfp4_linear_ste_vjp(primals, cotangent, output):
- x, weight = primals
- dy = cotangent
- qweight, scales, *rest = mx.quantize(
- weight,
- group_size=NVFP4_GROUP_SIZE,
- bits=NVFP4_BITS,
- mode=NVFP4_MODE,
- )
- qbiases = rest[0] if rest else None
- grad_x = mx.quantized_matmul(
- dy,
- qweight,
- scales,
- qbiases,
- transpose=False,
- group_size=NVFP4_GROUP_SIZE,
- bits=NVFP4_BITS,
- mode=NVFP4_MODE,
- )
- x2, dy2 = _flatten_linear_batch(mx, x, dy)
- grad_weight = mx.matmul(mx.transpose(dy2, (1, 0)), x2)
- return grad_x, grad_weight
-
- return nvfp4_linear_ste
-
-
-_NVFP4_LINEAR_STE = None
-
-
-def nvfp4_linear_ste(x, weight):
- global _NVFP4_LINEAR_STE
- if _NVFP4_LINEAR_STE is None:
- _NVFP4_LINEAR_STE = _make_nvfp4_linear_ste()
- return _NVFP4_LINEAR_STE(x, weight)
-
-
-def quantize_weight(weight, *, group_size: int = NVFP4_GROUP_SIZE):
- mx = require_mlx()
- if not can_nvfp4_quantize_shape(tuple(int(v) for v in weight.shape), group_size):
- raise ValueError(f"shape {tuple(weight.shape)} is not NVFP4 quantizable")
- qweight, scales, *rest = mx.quantize(
- weight,
- group_size=group_size,
- bits=NVFP4_BITS,
- mode=NVFP4_MODE,
- )
- qbiases = rest[0] if rest else None
- return qweight, scales, qbiases
-
-
-def quantized_matmul(x, qweight, scales, qbiases=None, *, transpose: bool = True):
- mx = require_mlx()
- return mx.quantized_matmul(
- x,
- qweight,
- scales,
- qbiases,
- transpose=transpose,
- group_size=NVFP4_GROUP_SIZE,
- bits=NVFP4_BITS,
- mode=NVFP4_MODE,
- )
-
-
-@dataclass(frozen=True)
-class QuantizedTensorInfo:
- name: str
- weight_shape: tuple[int, ...]
- qweight_shape: tuple[int, ...]
- scales_shape: tuple[int, ...]
- has_qbiases: bool
diff --git a/vendor/nanoai/mlx/train.py b/vendor/nanoai/mlx/train.py
deleted file mode 100644
index 905045bff56302a54dfe31d14fdfe394e556d5ac..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx/train.py
+++ /dev/null
@@ -1,184 +0,0 @@
-"""Small MLX training utilities for nanoai.
-
-This intentionally keeps the MLX path independent from the CUDA training stack.
-It supports dense and NVFP4-forward training; the NVFP4 path uses FP master
-weights plus STE gradients from ``nanoai.mlx.nvfp4``.
-"""
-
-from __future__ import annotations
-
-import json
-import math
-import os
-import time
-from dataclasses import asdict
-from typing import Iterable
-
-import mlx.core as mx
-import mlx.nn as nn
-import mlx.optimizers as optim
-import numpy as np
-from mlx.utils import tree_map
-
-from nanoai.mlx.checkpoint import save_mlx_checkpoint
-from nanoai.mlx.gpt import GPT, GPTConfig, loss_fn
-
-
-def build_model(
- *,
- depth: int,
- vocab_size: int,
- max_seq_len: int = 2048,
- aspect_ratio: int = 64,
- head_dim: int = 128,
- window_pattern: str = "SSSL",
- quantization: str = "dense",
-) -> GPT:
- base_dim = depth * aspect_ratio
- model_dim = ((base_dim + head_dim - 1) // head_dim) * head_dim
- num_heads = model_dim // head_dim
- config = GPTConfig(
- sequence_len=max_seq_len,
- vocab_size=vocab_size,
- n_layer=depth,
- n_head=num_heads,
- n_kv_head=num_heads,
- n_embd=model_dim,
- window_pattern=window_pattern,
- )
- model = GPT(config, quantization=quantization)
- model.init_weights()
- mx.eval(model.parameters())
- return model
-
-
-def to_mx_batch(x, y):
- if hasattr(x, "detach"):
- x = x.detach().cpu().numpy()
- if hasattr(y, "detach"):
- y = y.detach().cpu().numpy()
- return mx.array(x, dtype=mx.int32), mx.array(y, dtype=mx.int32)
-
-
-def random_batch(*, batch_size: int, seq_len: int, vocab_size: int, seed: int):
- rng = np.random.default_rng(seed)
- row = rng.integers(0, vocab_size, size=(batch_size, seq_len + 1), dtype=np.int32)
- return mx.array(row[:, :-1]), mx.array(row[:, 1:])
-
-
-def random_batches(*, batch_size: int, seq_len: int, vocab_size: int, seed: int = 0):
- step = 0
- while True:
- yield random_batch(batch_size=batch_size, seq_len=seq_len, vocab_size=vocab_size, seed=seed + step)
- step += 1
-
-
-def base_batches(tokenizer, *, batch_size: int, seq_len: int, split: str = "train"):
- from nanoai.dataloader import tokenizing_distributed_data_loader_with_state_bos_bestfit
-
- loader = tokenizing_distributed_data_loader_with_state_bos_bestfit(
- tokenizer,
- batch_size,
- seq_len,
- split=split,
- device="cpu",
- resume_state_dict=None,
- )
- while True:
- x, y, _state = next(loader)
- yield to_mx_batch(x, y)
-
-
-def sft_jsonl_batches(tokenizer, path: str, *, batch_size: int, seq_len: int, agentic: bool = False):
- rows = []
- with open(path, "r", encoding="utf-8") as f:
- for line in f:
- if line.strip():
- rows.append(json.loads(line))
- if not rows:
- raise ValueError(f"No JSONL rows found in {path}")
-
- if agentic:
- from nanoai.agentic.tokenizer import render_agentic_conversation
-
- def render(row):
- return render_agentic_conversation(tokenizer, row, max_tokens=seq_len + 1)
- else:
- def render(row):
- ids, mask = tokenizer.render_conversation(row, max_tokens=seq_len + 1)
- return ids, mask
-
- pos = 0
- while True:
- x = np.full((batch_size, seq_len), tokenizer.get_bos_token_id(), dtype=np.int32)
- y = np.full((batch_size, seq_len), -1, dtype=np.int32)
- for b in range(batch_size):
- row = rows[pos % len(rows)]
- pos += 1
- ids, mask = render(row)
- if len(ids) < 2:
- continue
- ids = ids[: seq_len + 1]
- mask = mask[: seq_len + 1]
- n = min(len(ids) - 1, seq_len)
- x[b, :n] = np.array(ids[:n], dtype=np.int32)
- targets = np.array(ids[1 : n + 1], dtype=np.int32)
- target_mask = np.array(mask[1 : n + 1], dtype=bool)
- y[b, :n] = np.where(target_mask, targets, -1)
- yield mx.array(x), mx.array(y)
-
-
-def train_loop(
- model: GPT,
- batches: Iterable,
- *,
- num_iterations: int,
- learning_rate: float,
- save_dir: str | None = None,
- save_every: int = -1,
- source_meta: dict | None = None,
-):
- optimizer = optim.AdamW(learning_rate=learning_rate)
- loss_grad_fn = nn.value_and_grad(model, loss_fn)
- smooth_loss = 0.0
- for step in range(num_iterations):
- t0 = time.time()
- x, y = next(batches)
- loss, grads = loss_grad_fn(model, x, y)
- loss_value = float(loss.item())
- if not math.isfinite(loss_value):
- raise RuntimeError(f"non-finite loss at step {step}: {loss_value}")
- optimizer.update(model, grads)
- mx.eval(model.parameters(), optimizer.state)
- dt = time.time() - t0
- smooth_loss = 0.9 * smooth_loss + 0.1 * loss_value
- debiased = smooth_loss / (1.0 - 0.9 ** (step + 1))
- if step < 5 or step == num_iterations - 1 or step % max(1, num_iterations // 10) == 0:
- toks = int(x.size / max(dt, 1e-9))
- print(f"step {step:05d} | loss {debiased:.4f} | {dt*1000:.0f}ms | {toks:,} tok/s")
- if save_dir and save_every > 0 and (step + 1) % save_every == 0:
- save_mlx_checkpoint(
- model,
- save_dir,
- step=step + 1,
- source_meta=source_meta or {"model_config": asdict(model.config)},
- quantization="dense_master",
- )
- if save_dir:
- save_mlx_checkpoint(
- model,
- save_dir,
- step=num_iterations,
- source_meta=source_meta or {"model_config": asdict(model.config)},
- quantization="dense_master",
- )
- return model
-
-
-def one_step_for_test(model: GPT, x, y, *, learning_rate: float = 1e-3):
- optimizer = optim.AdamW(learning_rate=learning_rate)
- loss_grad_fn = nn.value_and_grad(model, loss_fn)
- loss, grads = loss_grad_fn(model, x, y)
- optimizer.update(model, grads)
- mx.eval(model.parameters(), optimizer.state)
- return float(loss.item()), grads
diff --git a/vendor/nanoai/mlx_nvfp4.py b/vendor/nanoai/mlx_nvfp4.py
deleted file mode 100644
index 520723f6719e86af25e945941b2981f5421e07a7..0000000000000000000000000000000000000000
--- a/vendor/nanoai/mlx_nvfp4.py
+++ /dev/null
@@ -1,167 +0,0 @@
-"""Small MLX/NVFP4 probing helpers.
-
-This module is intentionally narrow: it does not try to replace nanoai's
-PyTorch runtime. It checks whether a nanoai Linear weight can be packed with
-MLX's NVFP4 quantizer and used with MLX's quantized matmul.
-"""
-
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import Mapping
-
-import numpy as np
-import torch
-
-
-@dataclass(frozen=True)
-class LinearProbeResult:
- weight_name: str
- weight_shape: tuple[int, ...]
- quantized_weight_shape: tuple[int, ...]
- scale_shape: tuple[int, ...]
- has_qbiases: bool
- input_shape: tuple[int, ...]
- max_abs_error: float
- mean_abs_error: float
- relative_l2_error: float
-
-
-def mlx_available() -> bool:
- try:
- import mlx.core # noqa: F401
- except Exception:
- return False
- return True
-
-
-def require_mlx():
- try:
- import mlx.core as mx
- except Exception as exc: # pragma: no cover - depends on host platform
- raise RuntimeError(
- "MLX is not available. Install it on Apple Silicon with `pip install mlx` "
- "or run this probe in an environment that has MLX with NVFP4 support."
- ) from exc
- return mx
-
-
-def assert_mlx_runtime_ready() -> None:
- mx = require_mlx()
- try:
- probe = mx.array([0.0])
- mx.eval(probe)
- except Exception as exc: # pragma: no cover - depends on host runtime
- raise RuntimeError(
- "MLX is importable, but the runtime cannot evaluate arrays. This often "
- "means the process is running in a headless/sandboxed macOS session "
- "without access to the Metal device."
- ) from exc
-
-
-def candidate_linear_weights(state_dict: Mapping[str, torch.Tensor]) -> list[str]:
- """Return likely nanoai Linear weights that MLX can NVFP4-quantize.
-
- MLX NVFP4 quantization requires a tensor with at least two dimensions and a
- last dimension divisible by group size 16. Embeddings also satisfy this, so
- filter them out for the linear probe.
- """
-
- names: list[str] = []
- for name, tensor in state_dict.items():
- if not name.endswith(".weight"):
- continue
- if "wte" in name or "value_embeds" in name:
- continue
- if not torch.is_tensor(tensor) or tensor.ndim != 2:
- continue
- if tensor.shape[-1] % 16 != 0:
- continue
- names.append(name)
- return sorted(names)
-
-
-def select_linear_weight(
- state_dict: Mapping[str, torch.Tensor],
- requested_name: str | None = None,
-) -> tuple[str, torch.Tensor]:
- if requested_name is not None:
- if requested_name not in state_dict:
- raise KeyError(f"weight {requested_name!r} not found in checkpoint")
- weight = state_dict[requested_name]
- if not torch.is_tensor(weight) or weight.ndim != 2:
- raise ValueError(f"weight {requested_name!r} is not a 2D tensor")
- if weight.shape[-1] % 16 != 0:
- raise ValueError(
- f"weight {requested_name!r} last dimension {weight.shape[-1]} "
- "is not divisible by NVFP4 group size 16"
- )
- return requested_name, weight
-
- candidates = candidate_linear_weights(state_dict)
- if not candidates:
- raise ValueError("no 2D Linear weights with last dimension divisible by 16 found")
- name = candidates[0]
- return name, state_dict[name]
-
-
-def probe_nvfp4_linear(
- weight_name: str,
- weight: torch.Tensor,
- *,
- rows: int = 4,
- seed: int = 0,
- group_size: int = 16,
-) -> LinearProbeResult:
- """Compare Torch dense linear output to MLX NVFP4 quantized matmul output."""
-
- if weight.ndim != 2:
- raise ValueError(f"{weight_name} must be a 2D Linear weight")
- if weight.shape[-1] % group_size != 0:
- raise ValueError(
- f"{weight_name} last dim {weight.shape[-1]} is not divisible by {group_size}"
- )
-
- mx = require_mlx()
- torch.manual_seed(seed)
-
- dense_weight = weight.detach().cpu().float().contiguous()
- x = torch.randn(rows, dense_weight.shape[1], dtype=torch.float32)
- expected = torch.nn.functional.linear(x, dense_weight)
-
- mx_weight = mx.array(dense_weight.numpy())
- mx_input = mx.array(x.numpy())
-
- quantized = mx.quantize(mx_weight, group_size=group_size, bits=4, mode="nvfp4")
- if len(quantized) == 2:
- qweight, scales = quantized
- qbiases = None
- else:
- qweight, scales, qbiases = quantized
-
- got = mx.quantized_matmul(
- mx_input,
- qweight,
- scales,
- qbiases,
- transpose=True,
- group_size=group_size,
- bits=4,
- mode="nvfp4",
- )
- mx.eval(got)
-
- got_torch = torch.from_numpy(np.array(got)).float()
- diff = got_torch - expected
- rel = diff.norm() / expected.norm().clamp_min(1e-12)
- return LinearProbeResult(
- weight_name=weight_name,
- weight_shape=tuple(int(v) for v in dense_weight.shape),
- quantized_weight_shape=tuple(int(v) for v in qweight.shape),
- scale_shape=tuple(int(v) for v in scales.shape),
- has_qbiases=qbiases is not None,
- input_shape=tuple(int(v) for v in x.shape),
- max_abs_error=float(diff.abs().max().item()),
- mean_abs_error=float(diff.abs().mean().item()),
- relative_l2_error=float(rel.item()),
- )
diff --git a/vendor/nanoai/msa/README.md b/vendor/nanoai/msa/README.md
deleted file mode 100644
index 02fe51100f6564709c8e71b54eeef36238072826..0000000000000000000000000000000000000000
--- a/vendor/nanoai/msa/README.md
+++ /dev/null
@@ -1,47 +0,0 @@
-# nanoai/msa — block-sparse attention (MiniMax-M3 style)
-
-A block-sparse attention that selects a small set of far KV blocks per query
-(top-k) plus a sliding local window and sink blocks, for cheaper long-context
-decode. Quality is on par with or better than dense, with extrapolation ≥ full
-attention. Project notes: see the `project_msa` memory and `docs/` (MSA spike).
-
-## Two paths
-
-- **Train-free probe**: drop MSA into a dense-trained model at inference to
- measure the quality/speed trade-off without retraining.
-- **Native training**: train with MSA layers. The `fa4` kernel path is trainable
- (fused forward + native block-sparse backward) and gives a real wall-clock
- speedup that grows with context; the `sdpa` path is correct but overhead-bound.
-
-Which layers are MSA is set by the model's `window_pattern` (e.g. `S`=short
-dense, `L`=long dense, `M`=MSA sparse).
-
-## Exports (`nanoai/msa/__init__.py`)
-
-`msa_block_sparse_attn_func`, `msa_gather_decode`, `update_msa_pooled_cache`,
-`msa_fa4_block_sparse_attn`, `default_block_size`, plus `apply_msa`,
-`num_msa_layers`, `extend_rotary` (`runtime.py`).
-
-## Config knobs (GPTConfig, mirrored in `nanoai.config.model_defaults()`)
-
-`msa_block_size=64`, `msa_top_k_blocks=16`, `msa_local_window=128`,
-`msa_sink_blocks=1`, `msa_index_mode="pooled_k"` (param-free block selection).
-
-Kernel selection via `NANOAI_MSA_KERNEL` (`sdpa` default — flexible/slow; `fa4`
-— fast, requires head_dim==64). `attention.py` holds the block-selection logic,
-`fa4.py` the FA4 CuTeDSL native kernel, `runtime.py` the rotary/cache handling.
-
-The `fa4` path is **trainable** and needs **flash-attn ≥ b11** (earlier betas ran
-the block-sparse backward dense → wrong grads; the b5-only repo-internal
-workaround — a custom autograd Function + a `_tile_size_bwd_sm90` monkey-patch —
-has been removed in favor of the public `flash_attn_func` with
-`block_sparse_tensors` + `block_sparse_tensors_bwd`). The kernel block size is
-arch-dependent (`default_block_size`): Hopper/SM90 → (128, 128), Blackwell/SM100
-(B200 sm_100, B300 sm_103) → (256, 128). Validated fwd + dq/dk/dv vs an SDPA
-reference at rel-err ~3e-3 on B300 (flash_attn_4 b14). Decode uses the
-incremental pooled-K cache (`update_msa_pooled_cache`, O(new_tokens) block
-scoring, bit-exact).
-
-Tests: `tests/test_msa_*.py` (the GPU `fa4` grad tests need flash-attn ≥ b11; the
-model-level `test_msa_gpt.py` builds CPU models, so run it with
-`NANOAI_ATTN_BACKEND=sdpa` to force the dense path onto CPU tensors).
diff --git a/vendor/nanoai/report.py b/vendor/nanoai/report.py
deleted file mode 100644
index c34baa69ae2ae6fc5ae8c9096f30e512be995b43..0000000000000000000000000000000000000000
--- a/vendor/nanoai/report.py
+++ /dev/null
@@ -1,418 +0,0 @@
-"""
-Utilities for generating training report cards. More messy code than usual, will fix.
-"""
-
-import os
-import re
-import shutil
-import subprocess
-import socket
-import datetime
-import platform
-import psutil
-import torch
-
-def run_command(cmd):
- """Run a shell command and return output, or None if it fails."""
- try:
- result = subprocess.run(cmd, shell=True, capture_output=True, text=True, timeout=5)
- # Return stdout if we got output (even if some files in xargs failed)
- if result.stdout.strip():
- return result.stdout.strip()
- if result.returncode == 0:
- return ""
- return None
- except:
- return None
-
-def get_git_info():
- """Get current git commit, branch, and dirty status."""
- info = {}
- info['commit'] = run_command("git rev-parse --short HEAD") or "unknown"
- info['branch'] = run_command("git rev-parse --abbrev-ref HEAD") or "unknown"
-
- # Check if repo is dirty (has uncommitted changes)
- status = run_command("git status --porcelain")
- info['dirty'] = bool(status) if status is not None else False
-
- # Get commit message
- info['message'] = run_command("git log -1 --pretty=%B") or ""
- info['message'] = info['message'].split('\n')[0][:80] # First line, truncated
-
- return info
-
-def get_gpu_info():
- """Get GPU information."""
- if not torch.cuda.is_available():
- return {"available": False}
-
- num_devices = torch.cuda.device_count()
- info = {
- "available": True,
- "count": num_devices,
- "names": [],
- "memory_gb": []
- }
-
- for i in range(num_devices):
- props = torch.cuda.get_device_properties(i)
- info["names"].append(props.name)
- info["memory_gb"].append(props.total_memory / (1024**3))
-
- # Get CUDA version
- info["cuda_version"] = torch.version.cuda or "unknown"
-
- return info
-
-def get_system_info():
- """Get system information."""
- info = {}
-
- # Basic system info
- info['hostname'] = socket.gethostname()
- info['platform'] = platform.system()
- info['python_version'] = platform.python_version()
- info['torch_version'] = torch.__version__
-
- # CPU and memory
- info['cpu_count'] = psutil.cpu_count(logical=False)
- info['cpu_count_logical'] = psutil.cpu_count(logical=True)
- info['memory_gb'] = psutil.virtual_memory().total / (1024**3)
-
- # User and environment
- info['user'] = os.environ.get('USER', 'unknown')
- info['nanochat_base_dir'] = os.environ.get('NANOAI_BASE_DIR', 'out')
- info['working_dir'] = os.getcwd()
-
- return info
-
-def estimate_cost(gpu_info, runtime_hours=None):
- """Estimate training cost based on GPU type and runtime."""
-
- # Rough pricing, from Lambda Cloud
- default_rate = 2.0
- gpu_hourly_rates = {
- "H100": 3.00,
- "A100": 1.79,
- "V100": 0.55,
- }
-
- if not gpu_info.get("available"):
- return None
-
- # Try to identify GPU type from name
- hourly_rate = None
- gpu_name = gpu_info["names"][0] if gpu_info["names"] else "unknown"
- for gpu_type, rate in gpu_hourly_rates.items():
- if gpu_type in gpu_name:
- hourly_rate = rate * gpu_info["count"]
- break
-
- if hourly_rate is None:
- hourly_rate = default_rate * gpu_info["count"] # Default estimate
-
- return {
- "hourly_rate": hourly_rate,
- "gpu_type": gpu_name,
- "estimated_total": hourly_rate * runtime_hours if runtime_hours else None
- }
-
-def generate_header():
- """Generate the header for a training report."""
- timestamp = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
-
- git_info = get_git_info()
- gpu_info = get_gpu_info()
- sys_info = get_system_info()
- cost_info = estimate_cost(gpu_info)
-
- header = f"""# nanoai training report
-
-Generated: {timestamp}
-
-## Environment
-
-### Git Information
-- Branch: {git_info['branch']}
-- Commit: {git_info['commit']} {"(dirty)" if git_info['dirty'] else "(clean)"}
-- Message: {git_info['message']}
-
-### Hardware
-- Platform: {sys_info['platform']}
-- CPUs: {sys_info['cpu_count']} cores ({sys_info['cpu_count_logical']} logical)
-- Memory: {sys_info['memory_gb']:.1f} GB
-"""
-
- if gpu_info.get("available"):
- gpu_names = ", ".join(set(gpu_info["names"]))
- total_vram = sum(gpu_info["memory_gb"])
- header += f"""- GPUs: {gpu_info['count']}x {gpu_names}
-- GPU Memory: {total_vram:.1f} GB total
-- CUDA Version: {gpu_info['cuda_version']}
-"""
- else:
- header += "- GPUs: None available\n"
-
- if cost_info and cost_info["hourly_rate"] > 0:
- header += f"""- Hourly Rate: ${cost_info['hourly_rate']:.2f}/hour\n"""
-
- header += f"""
-### Software
-- Python: {sys_info['python_version']}
-- PyTorch: {sys_info['torch_version']}
-
-"""
-
- # bloat metrics: count lines/chars in git-tracked source files only
- extensions = ['py', 'md', 'rs', 'html', 'toml', 'sh']
- git_patterns = ' '.join(f"'*.{ext}'" for ext in extensions)
- files_output = run_command(f"git ls-files -- {git_patterns}")
- file_list = [f for f in (files_output or '').split('\n') if f]
- num_files = len(file_list)
- num_lines = 0
- num_chars = 0
- if num_files > 0:
- wc_output = run_command(f"git ls-files -- {git_patterns} | xargs wc -lc 2>/dev/null")
- if wc_output:
- total_line = wc_output.strip().split('\n')[-1]
- parts = total_line.split()
- if len(parts) >= 2:
- num_lines = int(parts[0])
- num_chars = int(parts[1])
- num_tokens = num_chars // 4 # assume approximately 4 chars per token
-
- # count dependencies via uv.lock
- uv_lock_lines = 0
- if os.path.exists('uv.lock'):
- with open('uv.lock', 'r', encoding='utf-8') as f:
- uv_lock_lines = len(f.readlines())
-
- header += f"""
-### Bloat
-- Characters: {num_chars:,}
-- Lines: {num_lines:,}
-- Files: {num_files:,}
-- Tokens (approx): {num_tokens:,}
-- Dependencies (uv.lock lines): {uv_lock_lines:,}
-
-"""
- return header
-
-# -----------------------------------------------------------------------------
-
-def slugify(text):
- """Slugify a text string."""
- return text.lower().replace(" ", "-")
-
-# the expected files and their order
-EXPECTED_FILES = [
- "tokenizer-training.md",
- "tokenizer-evaluation.md",
- "base-model-training.md",
- "base-model-loss.md",
- "base-model-evaluation.md",
- "chat-sft.md",
- "chat-evaluation-sft.md",
- "chat-rl.md",
- "chat-evaluation-rl.md",
-]
-# the metrics we're currently interested in
-chat_metrics = ["ARC-Easy", "ARC-Challenge", "MMLU", "GSM8K", "HumanEval", "ChatCORE"]
-
-def extract(section, keys):
- """simple def to extract a single key from a section"""
- if not isinstance(keys, list):
- keys = [keys] # convenience
- out = {}
- for line in section.split("\n"):
- for key in keys:
- if key in line:
- out[key] = line.split(":")[1].strip()
- return out
-
-def extract_timestamp(content, prefix):
- """Extract timestamp from content with given prefix."""
- for line in content.split('\n'):
- if line.startswith(prefix):
- time_str = line.split(":", 1)[1].strip()
- try:
- return datetime.datetime.strptime(time_str, "%Y-%m-%d %H:%M:%S")
- except:
- pass
- return None
-
-class Report:
- """Maintains a bunch of logs, generates a final markdown report."""
-
- def __init__(self, report_dir):
- os.makedirs(report_dir, exist_ok=True)
- self.report_dir = report_dir
-
- def log(self, section, data):
- """Log a section of data to the report."""
- slug = slugify(section)
- file_name = f"{slug}.md"
- file_path = os.path.join(self.report_dir, file_name)
- with open(file_path, "w", encoding="utf-8") as f:
- f.write(f"## {section}\n")
- f.write(f"timestamp: {datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n\n")
- for item in data:
- if not item:
- # skip falsy values like None or empty dict etc.
- continue
- if isinstance(item, str):
- # directly write the string
- f.write(item)
- else:
- # render a dict
- for k, v in item.items():
- if isinstance(v, float):
- vstr = f"{v:.4f}"
- elif isinstance(v, int) and v >= 10000:
- vstr = f"{v:,.0f}"
- else:
- vstr = str(v)
- f.write(f"- {k}: {vstr}\n")
- f.write("\n")
- return file_path
-
- def generate(self):
- """Generate the final report."""
- report_dir = self.report_dir
- report_file = os.path.join(report_dir, "report.md")
- print(f"Generating report to {report_file}")
- final_metrics = {} # the most important final metrics we'll add as table at the end
- start_time = None
- end_time = None
- with open(report_file, "w", encoding="utf-8") as out_file:
- # write the header first
- header_file = os.path.join(report_dir, "header.md")
- if os.path.exists(header_file):
- with open(header_file, "r", encoding="utf-8") as f:
- header_content = f.read()
- out_file.write(header_content)
- start_time = extract_timestamp(header_content, "Run started:")
- # capture bloat data for summary later (the stuff after Bloat header and until \n\n)
- bloat_data = re.search(r"### Bloat\n(.*?)\n\n", header_content, re.DOTALL)
- bloat_data = bloat_data.group(1) if bloat_data else ""
- else:
- start_time = None # will cause us to not write the total wall clock time
- bloat_data = "[bloat data missing]"
- print(f"Warning: {header_file} does not exist. Did you forget to run `nanoai reset`?")
- # process all the individual sections
- for file_name in EXPECTED_FILES:
- section_file = os.path.join(report_dir, file_name)
- if not os.path.exists(section_file):
- print(f"Warning: {section_file} does not exist, skipping")
- continue
- with open(section_file, "r", encoding="utf-8") as in_file:
- section = in_file.read()
- # Extract timestamp from this section (the last section's timestamp will "stick" as end_time)
- if "rl" not in file_name:
- # Skip RL sections for end_time calculation because RL is experimental
- end_time = extract_timestamp(section, "timestamp:")
- # extract the most important metrics from the sections
- if file_name == "base-model-evaluation.md":
- final_metrics["base"] = extract(section, "CORE")
- if file_name == "chat-evaluation-sft.md":
- final_metrics["sft"] = extract(section, chat_metrics)
- if file_name == "chat-evaluation-rl.md":
- final_metrics["rl"] = extract(section, "GSM8K") # RL only evals GSM8K
- # append this section of the report
- out_file.write(section)
- out_file.write("\n")
- # add the final metrics table
- out_file.write("## Summary\n\n")
- # Copy over the bloat metrics from the header
- out_file.write(bloat_data)
- out_file.write("\n\n")
- # Collect all unique metric names
- all_metrics = set()
- for stage_metrics in final_metrics.values():
- all_metrics.update(stage_metrics.keys())
- # Custom ordering: CORE first, ChatCORE last, rest in middle
- all_metrics = sorted(all_metrics, key=lambda x: (x != "CORE", x == "ChatCORE", x))
- # Fixed column widths
- stages = ["base", "sft", "rl"]
- metric_width = 15
- value_width = 8
- # Write table header
- header = f"| {'Metric'.ljust(metric_width)} |"
- for stage in stages:
- header += f" {stage.upper().ljust(value_width)} |"
- out_file.write(header + "\n")
- # Write separator
- separator = f"|{'-' * (metric_width + 2)}|"
- for stage in stages:
- separator += f"{'-' * (value_width + 2)}|"
- out_file.write(separator + "\n")
- # Write table rows
- for metric in all_metrics:
- row = f"| {metric.ljust(metric_width)} |"
- for stage in stages:
- value = final_metrics.get(stage, {}).get(metric, "-")
- row += f" {str(value).ljust(value_width)} |"
- out_file.write(row + "\n")
- out_file.write("\n")
- # Calculate and write total wall clock time
- if start_time and end_time:
- duration = end_time - start_time
- total_seconds = int(duration.total_seconds())
- hours = total_seconds // 3600
- minutes = (total_seconds % 3600) // 60
- out_file.write(f"Total wall clock time: {hours}h{minutes}m\n")
- else:
- out_file.write("Total wall clock time: unknown\n")
- # also cp the report.md file to current directory
- print(f"Copying report.md to current directory for convenience")
- shutil.copy(report_file, "report.md")
- return report_file
-
- def reset(self):
- """Reset the report."""
- # Remove section files
- for file_name in EXPECTED_FILES:
- file_path = os.path.join(self.report_dir, file_name)
- if os.path.exists(file_path):
- os.remove(file_path)
- # Remove report.md if it exists
- report_file = os.path.join(self.report_dir, "report.md")
- if os.path.exists(report_file):
- os.remove(report_file)
- # Generate and write the header section with start timestamp
- header_file = os.path.join(self.report_dir, "header.md")
- header = generate_header()
- start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
- with open(header_file, "w", encoding="utf-8") as f:
- f.write(header)
- f.write(f"Run started: {start_time}\n\n---\n\n")
- print(f"Reset report and wrote header to {header_file}")
-
-# -----------------------------------------------------------------------------
-# nanoai-specific convenience functions
-
-class DummyReport:
- def log(self, *args, **kwargs):
- pass
- def reset(self, *args, **kwargs):
- pass
-
-def get_report():
- # just for convenience, only rank 0 logs to report
- from nanoai.common import get_base_dir, get_dist_info
- ddp, ddp_rank, ddp_local_rank, ddp_world_size = get_dist_info()
- if ddp_rank == 0:
- report_dir = os.path.join(get_base_dir(), "report")
- return Report(report_dir)
- else:
- return DummyReport()
-
-if __name__ == "__main__":
- import argparse
- parser = argparse.ArgumentParser(description="Generate or reset nanoai training reports.")
- parser.add_argument("command", nargs="?", default="generate", choices=["generate", "reset"], help="Operation to perform (default: generate)")
- args = parser.parse_args()
- if args.command == "generate":
- get_report().generate()
- elif args.command == "reset":
- get_report().reset()
diff --git a/vendor/nanoai/ui.html b/vendor/nanoai/ui.html
deleted file mode 100644
index b85532a335cd74410d4b98ef76e1d46ca18985ad..0000000000000000000000000000000000000000
--- a/vendor/nanoai/ui.html
+++ /dev/null
@@ -1,566 +0,0 @@
-
-
-
-
-
- NanoChat
-
-
-
-
-
-
-
-
-
-
-
-
-