🛠️ Additive-Rand-Transformer:极简算术 TinyGPT 实战手册与工具箱

极小自回归模型(CPU / Colab GPU)的训练、推理与评测全流程

Open In Colab License: MIT PyTorch Hugging Face


📖 项目背景与已实现能力

本项目提供了一套完整的端到端框架,用于在动态生成的任意多位加减法任务上训练、评估和量化紧凑型自回归 Transformer(TinyGPT,词表大小 = 16,参数量 $\le$ 926K)。

核心组件与实现:

  1. 动态数据生成器:生成 1–4 位非负整数加减法表达式,支持竖式思维链(CoT)进位/借位草稿纸展开格式(data.py)。
  2. 模型架构:支持标准因果自注意力(Causal Self-Attention)、动态稀疏注意力(DSA)、ALiBi 相对位置编码、混合专家(MoE)以及低秩适配(LoRA)的可配置 Decoder-Only 架构(model.py)。
  3. 参数扫描矩阵:138 项独立实验配置,全面覆盖深度($L=1\dots10$)、通道宽度($d=32\dots512$)、训练步数($500\dots8000$)与 Dropout 影响。
  4. 训练后量化工具:PyTorch 动态 INT8 量化及模拟 INT4 量化流水线,评测体积压缩比与推理吞吐(quantize.py)。
  5. 强化学习链路:支持自博弈(Self-Play)强化学习与基于组相对优势的 GRPO 训练,具备新颖性奖励追踪机制(rl_selfplay.py, grpo.py)。
  6. 数据与实验总库:完整的 58 列超大宽表全景实验矩阵存储于 EXPERIMENTS_ALL.xlsx

🧭 DIY 快速导航


🚀 DIY 第一步:Google Colab 零配置极速上手

无需本地 GPU 算力,直接利用 Google Colab 免费算力完成训练、评估与量化:

  1. 下载 Colab_Run_Additive_Transformer.ipynb
  2. 访问 Google Colab -> 点击 上传 (Upload) -> 选择该 .ipynb 文件。
  3. 点击 代码执行程序 (Runtime) -> **全部运行 (Run All)**(或按顺序单步执行)。

💻 DIY 第二步:本地安装与环境准备

# 1. 克隆本仓库
git clone https://huggingface.co/Hana-ame/additive-rand-transformer
cd additive-rand-transformer

# 2. 安装依赖(Python >= 3.8, PyTorch >= 2.0)
pip install torch openpyxl huggingface_hub pandas matplotlib
pip install -e .

🎮 DIY 第三步:交互式推理(REPL 与单题求解)

选项 A:使用交互式脚本

# 单题快速求解
./use_model.sh -s "1234 + 5678"
./use_model.sh -s "9999 - 4321"

# 启动交互式 REPL 会话
./use_model.sh -i

选项 B:Python API 调用

import torch
from huggingface_hub import hf_hub_download
from additive_rand_transformer.model import TinyGPT, TinyGPTConfig
from additive_rand_transformer.data import BOS, EOS, SP, PLUS, MINUS, _int_to_tokens, decode

# 1. 下载官方预训练权重
ckpt_path = hf_hub_download(
    repo_id="Hana-ame/additive-rand-transformer",
    filename="checkpoints/l4_d128_cot_bias05_final.pt",
    local_dir="."
)

# 2. 加载模型
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
cfg = TinyGPTConfig(**{k: v for k, v in ck["config"].items() if k in TinyGPTConfig.__dataclass_fields__})
model = TinyGPT(cfg)
model.load_state_dict(ck["model"])
model.eval()

# 3. 逐列竖式 CoT 生成求解
def solve(a: int, b: int, op_char: str = "+"):
    op = PLUS if op_char == "+" else MINUS
    prompt = [BOS] + _int_to_tokens(a) + [SP, op, SP] + _int_to_tokens(b) + [SP]
    ids = list(prompt)
    with torch.no_grad():
        for _ in range(80):
            x = torch.tensor([ids], dtype=torch.long)
            logits, _ = model(x, None)
            nxt = int(logits[0, -1].argmax())
            ids.append(nxt)
            if nxt == EOS:
                break
    print(f"算式: {a} {op_char} {b}")
    print(f"推理输出: {decode(ids)}\n")

solve(1234, 5678, "+")
solve(523, 194, "-")

🏋️ DIY 第四步:使用 config.json 自定义训练模型

1. 定义 config.json

{
  "layers": 4,
  "d": 128,
  "heads": 4,
  "steps": 4000,
  "batch_size": 32,
  "lr": 3e-4,
  "wd": 0.1,
  "warmup": 200,
  "datasource": {
    "type": "cot",
    "max_digits": 4,
    "bias": 0.5,
    "max_spaces": 3,
    "single": true
  }
}

2. 启动训练

# 基于配置文件训练
python -m additive_rand_transformer.train --config config.json

# 50 步快速冒烟测试(验证环境)
python -m additive_rand_transformer.train --quick

⚡ DIY 第五步:训练后动态 INT8 量化评测

对所有 Linear 全连接层应用 PyTorch 动态 INT8 量化,并测量内存占用与推理吞吐:

python -m additive_rand_transformer.quantize --checkpoint checkpoints/l4_d128_cot_bias05_final.pt
  • 量化实测收益:体积压缩 3.8x(1.7MB $\to$ 0.45MB),推理吞吐加速 1.4x,1–4 位加减法准确率 100% 保持(零退化)

🤖 DIY 第六步:自博弈(Self-Play)强化学习训练

启动自监督自动出题与求解自博弈训练:

python -m additive_rand_transformer.rl_selfplay \
    --checkpoint checkpoints/l4_d128_cot_bias05_final.pt \
    --reward_mode both \
    --memory_bonus 1.5 \
    --min_digits_reward 2 \
    --runs_dir runs/selfplay

📦 预训练权重 Checkpoint 清单

Checkpoint 名称 架构配置 参数量 说明与特性
l4_d128_cot_bias05_final.pt 4L · 128d · CoT Bias 0.5 926,464 主力交付模型,多位进位借位闭环
l4_d128_sft_nobias_final.pt 4L · 128d · SFT 均匀采样 926,464 操作数均匀采样基准模型
l4_d128_grpo_final.pt 4L · 128d · GRPO RL (150步) 926,464 基于组相对优势微调的强化学习模型
l4_d128_reinforce_conservative.pt 4L · 128d · REINFORCE (150步) 926,464 策略梯度微调的对照模型
l4_d128_selfplay_anchor.pt 4L · 128d · Self-Play 锚定 926,464 包含自适应出题难度的自博弈模型
l2_d64_attn_dsa.pt 2L · 64d · 动态稀疏注意力 166,656 采用 Top-8 数据依赖动态稀疏注意力的 2 层模型
l2_d64_attn_causal.pt 2L · 64d · 因果自注意力基准 166,656 2 层标准自注意力对照基准模型

📊 实验明细 Excel 工作簿

全部 138 项实验配置与评测指标均归档在: 👉 EXPERIMENTS_ALL.xlsx


📄 开源协议

本项目采用 MIT 协议 开源。

Downloads last month

-

Downloads are not tracked for this model. How to track
Video Preview
loading