bbkdevops's picture
download
raw
19.9 kB
from __future__ import annotations
from datetime import datetime, timezone
import hashlib
import json
import math
from pathlib import Path
from typing import Any
import torch
from model.native_axiom_regenesis import AxiomReGenesisConfig, TinyMindAxiomReGenesis, config_to_dict
def _render(row: dict[str, Any]) -> str:
messages = row.get("messages")
if isinstance(messages, list):
parts: list[str] = []
for message in messages:
if not isinstance(message, dict):
continue
role = str(message.get("role", "user")).upper()
content = str(message.get("content", "")).strip()
if content:
parts.append(f"{role}: {content}")
if parts:
return "\n".join(parts)
for key in ("text", "content", "response", "answer"):
value = row.get(key)
if isinstance(value, str) and value.strip():
return value.strip()
return json.dumps(row, ensure_ascii=False, sort_keys=True)
def _prompt_answer(row: dict[str, Any]) -> tuple[str, str]:
messages = row.get("messages")
if isinstance(messages, list):
prefix: list[str] = []
answer = ""
for message in messages:
if not isinstance(message, dict):
continue
role = str(message.get("role", "user")).upper()
content = str(message.get("content", "")).strip()
if not content:
continue
if role == "ASSISTANT":
answer = content
break
prefix.append(f"{role}: {content}")
if answer:
return "\n".join(prefix) + "\nASSISTANT:", answer
text = _render(row)
return "USER:", text
def _load_rows(path: str | Path, limit_records: int | None = None) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
with Path(path).open("r", encoding="utf-8") as handle:
for line in handle:
line = line.strip()
if not line:
continue
rows.append(json.loads(line))
if limit_records is not None and len(rows) >= limit_records:
break
return rows
def _sha(text: str) -> str:
return hashlib.sha256(text.encode("utf-8", errors="replace")).hexdigest()
def _row_digest(row: dict[str, Any]) -> str:
metadata = row.get("metadata") if isinstance(row.get("metadata"), dict) else {}
governor = row.get("quality_governor") if isinstance(row.get("quality_governor"), dict) else {}
provenance = (
governor.get("semantic_sha256")
or metadata.get("fingerprint_sha256")
or metadata.get("semantic_sha256")
or ""
)
return _sha(f"{provenance}\n{_render(row)}")
def _char_to_id(ch: str, vocab_size: int) -> int:
code = ord(ch)
if ch == "\n":
return min(vocab_size - 1, 99)
if ch == "\t":
return min(vocab_size - 1, 100)
if 32 <= code <= 126:
return min(vocab_size - 1, 4 + (code - 32))
if 0x0E00 <= code <= 0x0E7F:
return min(vocab_size - 1, 128 + (code - 0x0E00))
return 3
def _id_to_char(idx: int) -> str:
if idx == 99:
return "\n"
if idx == 100:
return "\t"
if 4 <= idx <= 98:
return chr(32 + (idx - 4))
if 128 <= idx <= 255:
return chr(0x0E00 + (idx - 128))
return ""
def _encode_text(text: str, vocab_size: int, tokenizer_mode: str) -> list[int]:
if tokenizer_mode == "char_v1":
return [_char_to_id(ch, vocab_size) for ch in text]
return [4 + (byte % max(1, vocab_size - 4)) for byte in text.encode("utf-8", errors="replace")]
def _encode(text: str, seq_len: int, vocab_size: int, tokenizer_mode: str = "byte") -> torch.Tensor:
ids = [1]
ids.extend(_encode_text(text, vocab_size, tokenizer_mode))
ids.append(2)
ids = ids[:seq_len]
if len(ids) < seq_len:
ids.extend([0] * (seq_len - len(ids)))
return torch.tensor(ids, dtype=torch.long)
def _encode_example(row: dict[str, Any], seq_len: int, vocab_size: int, tokenizer_mode: str = "byte") -> tuple[torch.Tensor, torch.Tensor]:
prompt, answer = _prompt_answer(row)
prompt_ids = [1]
prompt_ids.extend(_encode_text(prompt, vocab_size, tokenizer_mode))
answer_ids = _encode_text(answer, vocab_size, tokenizer_mode)
answer_ids.append(2)
if len(prompt_ids) + len(answer_ids) > seq_len:
answer_budget = min(len(answer_ids), max(1, seq_len // 2))
prompt_budget = max(1, seq_len - answer_budget)
prompt_ids = prompt_ids[:prompt_budget]
answer_ids = answer_ids[:answer_budget]
ids = (prompt_ids + answer_ids)[:seq_len]
labels = [-100] * min(len(prompt_ids), len(ids))
remaining = len(ids) - len(labels)
if remaining > 0:
labels.extend(ids[-remaining:])
if not any(label >= 0 for label in labels):
labels[-1] = ids[-1]
if len(ids) < seq_len:
pad = seq_len - len(ids)
ids.extend([0] * pad)
labels.extend([-100] * pad)
return torch.tensor(ids, dtype=torch.long), torch.tensor(labels, dtype=torch.long)
def _decode(ids: torch.Tensor, vocab_size: int, tokenizer_mode: str = "byte") -> str:
if tokenizer_mode == "char_v1":
chars = []
for value in ids.detach().cpu().tolist():
if value in (0, 1, 2):
continue
chars.append(_id_to_char(int(value)))
return "".join(chars)
bytes_out = []
for value in ids.detach().cpu().tolist():
if value in (0, 1, 2, 3):
continue
bytes_out.append(((int(value) - 4) % max(1, vocab_size - 4)) % 256)
return bytes(bytes_out).decode("utf-8", errors="replace")
def _batch(rows: list[dict[str, Any]], seq_len: int, vocab_size: int) -> torch.Tensor:
return torch.stack([_encode(_render(row), seq_len, vocab_size) for row in rows], dim=0)
def _batch_examples(rows: list[dict[str, Any]], seq_len: int, vocab_size: int, tokenizer_mode: str) -> tuple[torch.Tensor, torch.Tensor]:
encoded = [_encode_example(row, seq_len, vocab_size, tokenizer_mode) for row in rows]
return torch.stack([item[0] for item in encoded], dim=0), torch.stack([item[1] for item in encoded], dim=0)
def _retrieved_from_ids(ids: torch.Tensor, top_k: int, chunk_tokens: int = 16) -> torch.Tensor:
batch, seq_len = ids.shape
chunks = torch.zeros(batch, top_k, chunk_tokens, dtype=ids.dtype, device=ids.device)
if seq_len == 0:
return chunks
for k in range(top_k):
start = min(max(0, seq_len - chunk_tokens), (seq_len * k) // max(1, top_k))
piece = ids[:, start : start + chunk_tokens]
chunks[:, k, : piece.shape[1]] = piece
return chunks
def _split_rows(rows: list[dict[str, Any]], eval_records: int, seed: int) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]:
dedup: dict[str, dict[str, Any]] = {}
for row in rows:
dedup.setdefault(_row_digest(row), row)
groups: dict[str, list[tuple[str, dict[str, Any]]]] = {}
for digest, row in dedup.items():
groups.setdefault(_sha(_render(row)), []).append((digest, row))
ordered_groups = sorted(groups.items(), key=lambda item: hashlib.sha256(f"{seed}:{item[0]}".encode()).hexdigest())
eval_items: list[tuple[str, dict[str, Any]]] = []
train_items: list[tuple[str, dict[str, Any]]] = []
for _, group in ordered_groups:
if len(eval_items) < eval_records and len(ordered_groups) > 1:
eval_items.extend(group)
else:
train_items.extend(group)
if not train_items and len(eval_items) > 1:
train_items = eval_items[1:]
eval_items = eval_items[:1]
train_hashes = {digest for digest, _ in train_items}
eval_hashes = {digest for digest, _ in eval_items}
train_content_hashes = {_sha(_render(row)) for _, row in train_items}
eval_content_hashes = {_sha(_render(row)) for _, row in eval_items}
return [row for _, row in train_items], [row for _, row in eval_items], {
"input_records": len(rows),
"unique_records": len(dedup),
"unique_content_groups": len(groups),
"dropped_duplicates": len(rows) - len(dedup),
"train_eval_hash_overlap": len(train_hashes & eval_hashes),
"train_eval_content_overlap": len(train_content_hashes & eval_content_hashes),
}
@torch.no_grad()
def _eval_loss(model: TinyMindAxiomReGenesis, ids: torch.Tensor, labels: torch.Tensor, batch_size: int = 4) -> float:
model.eval()
losses: list[float] = []
for start in range(0, ids.shape[0], batch_size):
batch = ids[start : start + batch_size]
batch_labels = labels[start : start + batch_size]
retrieved = _retrieved_from_ids(batch, model.cfg.regen_top_k)
out = model(batch, labels=batch_labels, retrieved_tokens=retrieved)
loss = out["loss"]
if loss is not None:
losses.append(float(loss.detach().cpu()))
return float(sum(losses) / max(1, len(losses)))
@torch.no_grad()
def _memory_smoke(model: TinyMindAxiomReGenesis, seq_len: int, vocab_size: int, device: torch.device) -> dict[str, Any]:
lengths = sorted(set(max(8, min(seq_len, value)) for value in (32, 64, 128, seq_len)))
shapes = []
reports = []
for length in lengths:
ids = torch.randint(4, vocab_size, (1, length), device=device)
retrieved = _retrieved_from_ids(ids, model.cfg.regen_top_k)
out = model(ids, retrieved_tokens=retrieved, return_report=True)
states = out["states"]
shapes.append([list(state.shape) for state in states])
reports.append(out["report"])
state_shapes_constant = all(shape == shapes[0] for shape in shapes)
return {
"sequence_lengths": lengths,
"state_shapes": shapes,
"state_shapes_constant_by_context_length": state_shapes_constant,
"kv_tokens_stored": 0,
"sample_layer_report": reports[-1]["layer_reports"][0] if reports and reports[-1]["layer_reports"] else {},
}
@torch.no_grad()
def _probe_generation(model: TinyMindAxiomReGenesis, seq_len: int, vocab_size: int, device: torch.device) -> list[dict[str, Any]]:
prompts = [
{"axis": "thai", "prompt": "USER: อธิบาย data governance แบบสั้นและตรงประเด็น\nASSISTANT:"},
{"axis": "code", "prompt": "USER: write python function add(a,b)\nASSISTANT:"},
{"axis": "math", "prompt": "USER: solve 2+3 and explain\nASSISTANT:"},
{"axis": "tool", "prompt": "USER: return JSON with name and arguments\nASSISTANT:"},
]
outputs = []
for item in prompts:
ids = _encode(item["prompt"], min(seq_len, 96), vocab_size, model.cfg.tokenizer_mode).unsqueeze(0).to(device)
retrieved = _retrieved_from_ids(ids, model.cfg.regen_top_k)
generated = model.generate(ids, max_new_tokens=24, retrieved_tokens=retrieved)
tail = generated[:, ids.shape[1] :]
text = _decode(tail[0], vocab_size, model.cfg.tokenizer_mode)
outputs.append(
{
"axis": item["axis"],
"prompt": item["prompt"],
"raw_token_count": int(tail.numel()),
"decoded": text,
"non_empty": bool(text.strip()),
"unique_char_ratio": len(set(text)) / max(1, len(text)),
}
)
return outputs
def _fixed_answer_collapse(probes: list[dict[str, Any]]) -> dict[str, Any]:
outputs = [str(probe.get("decoded", "")) for probe in probes]
unique = len(set(outputs))
return {
"unique_outputs": unique,
"probe_count": len(outputs),
"fixed_answer_collapse_detected": len(outputs) > 1 and unique <= 1,
}
def run_native_axiom_regenesis_train(
out_dir: str | Path,
*,
dataset: str | Path,
max_steps: int = 16,
eval_records: int = 16,
limit_records: int | None = None,
dim: int = 128,
layers: int = 3,
lanes: int = 8,
seq_len: int = 128,
vocab_size: int = 512,
virtual_dim: int = 20_480,
basis_rank: int = 32,
facets: int = 8,
tokenizer_mode: str = "byte",
learning_rate: float = 3e-4,
train_batch_size: int = 1,
seed: int = 20260528,
device: str | None = None,
resume_checkpoint: str | Path | None = None,
) -> dict[str, Any]:
torch.manual_seed(seed)
out = Path(out_dir)
out.mkdir(parents=True, exist_ok=True)
rows = _load_rows(dataset, limit_records=limit_records)
if len(rows) < 6:
raise ValueError("native AxiomReGenesis train requires at least 6 rows")
train_rows, eval_rows, split_report = _split_rows(rows, eval_records=eval_records, seed=seed)
run_device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
resume_payload: dict[str, Any] | None = None
if resume_checkpoint is not None:
resume_payload = torch.load(resume_checkpoint, map_location=run_device)
cfg_payload = resume_payload.get("config")
if not isinstance(cfg_payload, dict):
raise ValueError(f"resume checkpoint missing config: {resume_checkpoint}")
cfg = AxiomReGenesisConfig(**cfg_payload)
tokenizer_mode = cfg.tokenizer_mode
seq_len = min(seq_len, cfg.max_seq_len)
vocab_size = cfg.vocab_size
else:
cfg = AxiomReGenesisConfig(
vocab_size=vocab_size,
tokenizer_mode=tokenizer_mode,
dim=dim,
n_layers=layers,
lanes=lanes,
max_seq_len=seq_len,
local_window=min(64, seq_len),
memory_slots=max(4, min(16, lanes)),
memory_rank=max(8, min(64, dim // 4)),
regen_top_k=4,
regen_rank=4,
axiom_effective_dim=virtual_dim,
axiom_basis_rank=min(basis_rank, max(4, dim)),
axiom_facets=facets,
dropout=0.0,
residual_alpha=layers ** -0.5,
)
model = TinyMindAxiomReGenesis(cfg).to(run_device)
if resume_payload is not None:
state_dict = resume_payload.get("state_dict")
if not isinstance(state_dict, dict):
raise ValueError(f"resume checkpoint missing state_dict: {resume_checkpoint}")
model.load_state_dict(state_dict)
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=0.01)
train_ids, train_labels = _batch_examples(train_rows, seq_len, vocab_size, cfg.tokenizer_mode)
eval_ids, eval_labels = _batch_examples(eval_rows, seq_len, vocab_size, cfg.tokenizer_mode)
train_ids = train_ids.to(run_device)
train_labels = train_labels.to(run_device)
eval_ids = eval_ids.to(run_device)
eval_labels = eval_labels.to(run_device)
pre_loss = _eval_loss(model, eval_ids, eval_labels)
losses: list[float] = []
steps = max(1, int(max_steps))
batch_size = max(1, int(train_batch_size))
generator = torch.Generator(device=run_device)
generator.manual_seed(seed)
model.train()
for step in range(steps):
if batch_size == 1:
index = int(torch.randint(0, train_ids.shape[0], (1,), generator=generator, device=run_device).item())
sample = train_ids[index : index + 1]
sample_labels = train_labels[index : index + 1]
else:
indices = torch.randint(0, train_ids.shape[0], (batch_size,), generator=generator, device=run_device)
sample = train_ids.index_select(0, indices)
sample_labels = train_labels.index_select(0, indices)
retrieved = _retrieved_from_ids(sample, cfg.regen_top_k)
optimizer.zero_grad(set_to_none=True)
out_dict = model(sample, labels=sample_labels, retrieved_tokens=retrieved)
loss = out_dict["loss"]
if loss is None or not torch.isfinite(loss):
raise RuntimeError(f"non-finite native train loss at step {step}")
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
losses.append(float(loss.detach().cpu()))
post_loss = _eval_loss(model, eval_ids, eval_labels)
probes = _probe_generation(model, seq_len, vocab_size, run_device)
collapse = _fixed_answer_collapse(probes)
memory = _memory_smoke(model, seq_len, vocab_size, run_device)
checkpoint_path = out / "checkpoint.pt"
torch.save(
{
"schema": "tinymind.native_axiom_regenesis.checkpoint.v1",
"config": config_to_dict(cfg),
"state_dict": model.state_dict(),
},
checkpoint_path,
)
report = {
"schema": "tinymind.native_axiom_regenesis_train.v1",
"created_at": datetime.now(timezone.utc).isoformat(),
"dataset": str(dataset),
"summary": {
"model_name": cfg.architecture_name,
"device": str(run_device),
"records_loaded": len(rows),
"train_records": len(train_rows),
"eval_records": len(eval_rows),
"dim": cfg.dim,
"layers": cfg.n_layers,
"lanes": cfg.lanes,
"seq_len": seq_len,
"vocab_size": cfg.vocab_size,
"virtual_dim": cfg.axiom_effective_dim,
"train_batch_size": batch_size,
"loss_mode": "assistant_completion_only",
"tokenizer_mode": cfg.tokenizer_mode,
"parameter_count": model.parameter_count,
"resume_checkpoint": str(resume_checkpoint) if resume_checkpoint is not None else None,
},
"split": split_report,
"metrics": {
"pre_eval_loss": pre_loss,
"post_eval_loss": post_loss,
"pre_eval_perplexity": math.exp(pre_loss) if math.isfinite(pre_loss) and pre_loss < 50 else float("inf"),
"post_eval_perplexity": math.exp(post_loss) if math.isfinite(post_loss) and post_loss < 50 else float("inf"),
"eval_loss_delta": post_loss - pre_loss,
"eval_loss_improved": post_loss < pre_loss,
"train_loss_first": losses[0],
"train_loss_last": losses[-1],
"train_loss_min": min(losses),
"train_loss_mean": sum(losses) / max(1, len(losses)),
"train_steps_completed": len(losses),
"train_loss_finite": all(math.isfinite(value) for value in losses),
},
"memory_context": memory,
"generation_probes": probes,
"collapse_check": collapse,
"artifacts": {
"checkpoint_path": str(checkpoint_path),
},
"claim_gate": {
"native_architecture_independent": True,
"native_checkpoint_infer_ready": checkpoint_path.exists() and all(probe["non_empty"] for probe in probes),
"native_training_proven": len(losses) == steps and all(math.isfinite(value) for value in losses),
"eval_credible": split_report["train_eval_hash_overlap"] == 0
and split_report["train_eval_content_overlap"] == 0
and math.isfinite(pre_loss)
and math.isfinite(post_loss),
"memory_bounded": memory["state_shapes_constant_by_context_length"] and memory["kv_tokens_stored"] == 0,
"promotion_allowed": False,
"world_best_claim_allowed": False,
"reason": "Native smoke proves train/infer/memory mechanics only; promotion requires baseline probe wins and external evidence.",
},
}
report_path = out / "native_axiom_regenesis_train_report.json"
report["json_path"] = str(report_path)
runtime_metadata = model.export_runtime_metadata(out, training_report=report)
report["artifacts"]["runtime_metadata_path"] = runtime_metadata["json_path"]
report_path.write_text(json.dumps(report, ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8")
return report

Xet Storage Details

Size:
19.9 kB
·
Xet hash:
b3c3e06d35801978ba02971c28e9f6d11c3f9519c223cb3220cc963c1cef4bc2

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.