Text Generation
Transformers
Safetensors
PyTorch
English
ac_swiglu
custom_code
causal-lm
channelmix-swiglu
channel-mixing
qana
zeros-qana-5m
generalist
4k-tokenizer
fromziro
Instructions to use fromziro/ZeroS-Qana-5M with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use fromziro/ZeroS-Qana-5M with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="fromziro/ZeroS-Qana-5M", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("fromziro/ZeroS-Qana-5M", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use fromziro/ZeroS-Qana-5M with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "fromziro/ZeroS-Qana-5M" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "fromziro/ZeroS-Qana-5M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/fromziro/ZeroS-Qana-5M
- SGLang
How to use fromziro/ZeroS-Qana-5M with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "fromziro/ZeroS-Qana-5M" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "fromziro/ZeroS-Qana-5M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "fromziro/ZeroS-Qana-5M" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "fromziro/ZeroS-Qana-5M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use fromziro/ZeroS-Qana-5M with Docker Model Runner:
docker model run hf.co/fromziro/ZeroS-Qana-5M
Step 1000: Int 4.59, Avg 34.55%, BPB 1.6713
Browse files- README.md +186 -0
- benchmark_results.json +202 -0
- config.json +39 -0
- generation_config.json +8 -0
- model.safetensors +3 -0
- modeling_ac_swiglu.py +415 -0
- tokenizer.json +0 -0
- tokenizer_config.json +14 -0
- training_config.json +64 -0
README.md
ADDED
|
@@ -0,0 +1,186 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
library_name: transformers
|
| 3 |
+
pipeline_tag: text-generation
|
| 4 |
+
tags:
|
| 5 |
+
- custom_code
|
| 6 |
+
- causal-lm
|
| 7 |
+
- text-generation
|
| 8 |
+
- pytorch
|
| 9 |
+
- ac-swiglu
|
| 10 |
+
- inclusive-mixing
|
| 11 |
+
- small-language-model
|
| 12 |
+
- generalist
|
| 13 |
+
- 4k-tokenizer
|
| 14 |
+
datasets:
|
| 15 |
+
- HuggingFaceFW/fineweb_edu_100BT-shuffled
|
| 16 |
+
- HuggingFaceTB/smollm-corpus
|
| 17 |
+
- epfml/FineWeb-HQ
|
| 18 |
+
- HuggingFaceTB/finemath
|
| 19 |
+
---
|
| 20 |
+
|
| 21 |
+
# Inclusive AC-SwiGLU Mini
|
| 22 |
+
|
| 23 |
+
Evaluated training checkpoint from a 4.94M-parameter
|
| 24 |
+
inclusive attention-mixed SwiGLU
|
| 25 |
+
generalist language model using the third-party
|
| 26 |
+
AxiomicLabs GPT-S 4,096-token tokenizer. It has no place embeddings, role
|
| 27 |
+
embeddings, or inference-time equation detection. It was recorded at step
|
| 28 |
+
1,000 with WikiText normalized BPB 1.6713.
|
| 29 |
+
|
| 30 |
+
## Loading
|
| 31 |
+
|
| 32 |
+
This is a custom Transformers architecture. `trust_remote_code=True` is
|
| 33 |
+
required because stock Hugging Face model classes do not implement AC-SwiGLU or this
|
| 34 |
+
model's exact rotary convention.
|
| 35 |
+
|
| 36 |
+
```bash
|
| 37 |
+
pip install "torch>=2.5" "transformers>=4.50" safetensors
|
| 38 |
+
```
|
| 39 |
+
|
| 40 |
+
```python
|
| 41 |
+
import torch
|
| 42 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 43 |
+
|
| 44 |
+
repo = "User01110/attention-contraction-mini"
|
| 45 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 46 |
+
dtype = torch.bfloat16 if device == "cuda" else torch.float32
|
| 47 |
+
tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
|
| 48 |
+
model = AutoModelForCausalLM.from_pretrained(
|
| 49 |
+
repo,
|
| 50 |
+
trust_remote_code=True,
|
| 51 |
+
torch_dtype=dtype,
|
| 52 |
+
).to(device).eval()
|
| 53 |
+
|
| 54 |
+
inputs = tokenizer("The process of photosynthesis", return_tensors="pt").to(device)
|
| 55 |
+
output = model.generate(**inputs, max_new_tokens=64)
|
| 56 |
+
print(tokenizer.decode(output[0], skip_special_tokens=True))
|
| 57 |
+
```
|
| 58 |
+
|
| 59 |
+
Checkpoint tensors are stored in bfloat16. Pass `torch_dtype=torch.float32`
|
| 60 |
+
when an FP32 runtime is required; every stored BF16 value widens exactly to
|
| 61 |
+
FP32, though the pre-export FP32 master-weight mantissa cannot be reconstructed.
|
| 62 |
+
The model intentionally sets `use_cache=False`: generation is standard
|
| 63 |
+
Transformers generation, but the visible context is recomputed for every new
|
| 64 |
+
token because this experimental architecture does not implement a KV cache.
|
| 65 |
+
|
| 66 |
+
## Architecture
|
| 67 |
+
|
| 68 |
+
- Parameters: 4,943,712, with tied input/output embeddings
|
| 69 |
+
- Weights: native bfloat16 safetensors (`model.safetensors`); no `.bin` weights
|
| 70 |
+
- Runtime: native SDPA, a fused SwiGLU input projection, and small dense native
|
| 71 |
+
PyTorch expanded-space attention-mixing matmuls
|
| 72 |
+
- Tokenizer: `AxiomicLabs/GPT-S-5M` at revision
|
| 73 |
+
`275b9c3ca78736bf6aeb154c7e2d5f5764fe9035`
|
| 74 |
+
- Vocabulary: 4,096 third-party tokens
|
| 75 |
+
- Parameter allocation: 884,736 tied embedding parameters and
|
| 76 |
+
4,058,976 non-embedding parameters
|
| 77 |
+
- Supported/exported context: 1,024 tokens
|
| 78 |
+
- Training block length: 1,024 tokens
|
| 79 |
+
- Standalone prompt tokenization automatically prepends the native BOS token
|
| 80 |
+
- Width/layers: 216 / 10
|
| 81 |
+
- Token-attention heads: 6 query, 2 KV
|
| 82 |
+
- Native SwiGLU expands each token to 432 channels with
|
| 83 |
+
one fused projection, applies `up * SiLU(gate)`, and retains an ordinary
|
| 84 |
+
zero-initialized linear contraction into the residual stream
|
| 85 |
+
- Inclusive AC mixing uses all
|
| 86 |
+
18 expanded chunks as queries, keys,
|
| 87 |
+
and values, each of 24 channels split over 3 heads
|
| 88 |
+
- Every head computes an exact dense
|
| 89 |
+
18x18
|
| 90 |
+
softmax mixing graph; no edge is removed or approximated
|
| 91 |
+
- The inclusive operator is `V_mix = [I + s(A(x) - A0)]V`; `V_mix` then passes
|
| 92 |
+
through the ordinary SwiGLU down projection exactly once
|
| 93 |
+
- There is no parallel AC output branch. At neutral mixing `A(x) = A0`, the
|
| 94 |
+
architecture is exactly native SwiGLU with `V_mix = V`
|
| 95 |
+
- The learned key scale and zero-initialized down projection make the complete
|
| 96 |
+
AC-SwiGLU residual branch exactly zero at initialization without a route gate
|
| 97 |
+
- The difference of two row-stochastic mixtures has zero row sum and bounded
|
| 98 |
+
L1 norm, providing a stable internal deformation of the expanded features
|
| 99 |
+
- Contiguous-half RoPE without scaling
|
| 100 |
+
- No task-specific model features or inference-time benchmark handling
|
| 101 |
+
|
| 102 |
+
## Intended use and limitations
|
| 103 |
+
|
| 104 |
+
This is a small base language model released for architecture research,
|
| 105 |
+
representation analysis, and controlled comparisons. It is not instruction
|
| 106 |
+
tuned and should not be treated as a factual authority or used for consequential
|
| 107 |
+
decisions. Its 4.94M-parameter scale, English-weighted web mixture,
|
| 108 |
+
1,024-token context, and cache-free generation materially limit
|
| 109 |
+
capability and throughput. The usual web-corpus biases and inaccuracies remain.
|
| 110 |
+
|
| 111 |
+
## Tokenizer provenance
|
| 112 |
+
|
| 113 |
+
The tokenizer and its 4,096-token vocabulary were **not
|
| 114 |
+
created or owned by the AC-SwiGLU model author**. They are reused from the public
|
| 115 |
+
[`AxiomicLabs/GPT-S-5M`](https://huggingface.co/AxiomicLabs/GPT-S-5M)
|
| 116 |
+
repository at the exact revision listed above, whose repository metadata
|
| 117 |
+
identifies Axiomic Labs as the publisher and Apache-2.0 as the license. No
|
| 118 |
+
claim of tokenizer ownership beyond that public attribution is made here. The
|
| 119 |
+
exported copy preserves its vocabulary and tokenization pipeline; AC-SwiGLU only
|
| 120 |
+
configures its existing `<bos>` token to be prepended automatically and raises
|
| 121 |
+
the exported model/tokenizer context limit to 1,024 tokens.
|
| 122 |
+
|
| 123 |
+
## Optimization
|
| 124 |
+
|
| 125 |
+
- Training budget: 20,971,520,000 tokens over 40,000 updates
|
| 126 |
+
- Effective batch: 524,288 tokens per update
|
| 127 |
+
- Native microbatch: 512 sequences x 1,024 tokens,
|
| 128 |
+
accumulated 1 times
|
| 129 |
+
- Runtime implementation: explicit BF16 residual/AC-SwiGLU activations, one
|
| 130 |
+
fused expansion GEMM, dense 18x18 inclusive mixing matmuls, and full-model
|
| 131 |
+
`torch.compile` enabled by default
|
| 132 |
+
- Runtime: 1 resumable 40,000-update sessions
|
| 133 |
+
with exact full-precision optimizer/RNG/data-stream recovery checkpoints
|
| 134 |
+
- Learning rate: 1,000-update linear warmup, flat at 2.5e-03 through update 30,000, then cosine to zero on update 40,000; original-scaling Muon follows the same
|
| 135 |
+
multiplier from its 3.0e-02 peak
|
| 136 |
+
- The final configured update has exactly zero learning rate
|
| 137 |
+
- Official PyTorch Muon with original matrix-shape scaling for hidden matrices;
|
| 138 |
+
AdamW for embeddings and remaining parameters
|
| 139 |
+
|
| 140 |
+
## Training mixture
|
| 141 |
+
|
| 142 |
+
- FineWeb-Edu 100BT shuffled: 55%
|
| 143 |
+
- Cosmopedia v2: 25%
|
| 144 |
+
- FineWeb-HQ: 10%
|
| 145 |
+
- FineMath 4+: 10%
|
| 146 |
+
|
| 147 |
+
FineWeb-Edu supplies the primary educational web text, Cosmopedia supplies
|
| 148 |
+
synthetic textbook-style coverage, FineWeb-HQ contributes model-filtered,
|
| 149 |
+
knowledge-rich general web text, and FineMath-4+ supplies mathematical
|
| 150 |
+
reasoning as ordinary causal-language-model text. The mixture remains fixed
|
| 151 |
+
for the complete run.
|
| 152 |
+
There are no benchmark labels or benchmark-specific preprocessing. Dataset
|
| 153 |
+
revisions are pinned and the same fixed mixture is maintained across both
|
| 154 |
+
runtime sessions.
|
| 155 |
+
|
| 156 |
+
## Zero-shot evaluation at step 1,000
|
| 157 |
+
|
| 158 |
+
The four lm-eval tasks use normalized accuracy when supplied by lm-eval
|
| 159 |
+
0.4.12, with native bfloat16 weights and float32
|
| 160 |
+
likelihood softmax. ArithMark-3 uses its official primary `acc_norm` metric:
|
| 161 |
+
mean continuation-token log likelihood, with context and continuation tokenized
|
| 162 |
+
separately after one native BOS prefix. Autocast is not used for evaluation. The four
|
| 163 |
+
lm-eval benchmark contexts use the tokenizer's native BOS behavior.
|
| 164 |
+
|
| 165 |
+
**Open SLM Intelligence Index: 4.59**
|
| 166 |
+
**Open SLM Average: 34.55%**
|
| 167 |
+
|
| 168 |
+
| Benchmark | Accuracy |
|
| 169 |
+
|---|---:|
|
| 170 |
+
| HellaSwag | 27.49% |
|
| 171 |
+
| ARC-Easy | 31.14% |
|
| 172 |
+
| ARC-Challenge | 20.05% |
|
| 173 |
+
| PIQA | 54.19% |
|
| 174 |
+
| ArithMark-3 (`acc_norm`) | 29.70% |
|
| 175 |
+
|
| 176 |
+
The Average weights HellaSwag, combined ARC (the mean of ARC-Easy and
|
| 177 |
+
ARC-Challenge), and PIQA at 1.0, and ArithMark-3 at
|
| 178 |
+
0.75. The Intelligence Index first maps random
|
| 179 |
+
chance to 0 and perfect accuracy to 100, then applies those same weights,
|
| 180 |
+
matching Open SLM Leaderboard revision
|
| 181 |
+
`2fbaaa164009c6c0d201ddf245878f5942cc8628`.
|
| 182 |
+
|
| 183 |
+
WikiText-103 validation at this step: loss 3.6964, perplexity
|
| 184 |
+
40.30, normalized BPB 1.6713 over 359,037
|
| 185 |
+
scored tokens and 1,145,591 normalized UTF-8 bytes, using one initial BOS,
|
| 186 |
+
1,024-token windows, and a 512-token stride.
|
benchmark_results.json
ADDED
|
@@ -0,0 +1,202 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"step": 1000,
|
| 3 |
+
"parameters": 4943712,
|
| 4 |
+
"lm_eval_version": "0.4.12",
|
| 5 |
+
"evaluation_dtype": "bfloat16",
|
| 6 |
+
"softmax_dtype": "float32",
|
| 7 |
+
"evaluation_autocast": false,
|
| 8 |
+
"num_fewshot": 0,
|
| 9 |
+
"model_context": 1024,
|
| 10 |
+
"training_context": 1024,
|
| 11 |
+
"lm_eval_context": 1024,
|
| 12 |
+
"lm_eval_bos_prefix": true,
|
| 13 |
+
"arithmark_3.0_bos_prefix": true,
|
| 14 |
+
"benchmark_accuracy": {
|
| 15 |
+
"arc_easy": 0.3114478114478115,
|
| 16 |
+
"arc_challenge": 0.20051194539249148,
|
| 17 |
+
"hellaswag": 0.2749452300338578,
|
| 18 |
+
"piqa": 0.5418933623503809
|
| 19 |
+
},
|
| 20 |
+
"lm_eval_metric_policy": "acc_norm when available, otherwise acc",
|
| 21 |
+
"intelligence_index": 4.587205403718424,
|
| 22 |
+
"average_accuracy_percent": 34.548492554783735,
|
| 23 |
+
"open_slm_scores": {
|
| 24 |
+
"leaderboard": "AxiomicLabs/Open_SLM_Leaderboard",
|
| 25 |
+
"leaderboard_revision": "2fbaaa164009c6c0d201ddf245878f5942cc8628",
|
| 26 |
+
"intelligence_index": 4.587205403718424,
|
| 27 |
+
"average_accuracy_percent": 34.548492554783735,
|
| 28 |
+
"components_percent": {
|
| 29 |
+
"hellaswag": 27.49452300338578,
|
| 30 |
+
"combined_arc": 25.59798784201515,
|
| 31 |
+
"piqa": 54.18933623503809,
|
| 32 |
+
"arithmark_3.0": 29.7
|
| 33 |
+
},
|
| 34 |
+
"chance_normalized_components": {
|
| 35 |
+
"hellaswag": 3.3260306711810395,
|
| 36 |
+
"combined_arc": 0.7973171226868677,
|
| 37 |
+
"piqa": 8.378672470076182,
|
| 38 |
+
"arithmark_3.0": 6.266666666666666
|
| 39 |
+
},
|
| 40 |
+
"average_accuracy_weights": {
|
| 41 |
+
"hellaswag": 1.0,
|
| 42 |
+
"combined_arc": 1.0,
|
| 43 |
+
"piqa": 1.0,
|
| 44 |
+
"arithmark_3.0": 0.75
|
| 45 |
+
},
|
| 46 |
+
"intelligence_index_weights": {
|
| 47 |
+
"hellaswag": 1.0,
|
| 48 |
+
"combined_arc": 1.0,
|
| 49 |
+
"piqa": 1.0,
|
| 50 |
+
"arithmark_3.0": 0.75
|
| 51 |
+
}
|
| 52 |
+
},
|
| 53 |
+
"checkpoint_policy": "every completed evaluation is committed regardless of whether validation BPB is best; rolling full-precision model, optimizer, RNG, stream cursor, and prefetched batches are stored separately",
|
| 54 |
+
"best_tracking_metric": "validation.normalized_bpb@1024 (lower is better)",
|
| 55 |
+
"best_validation": {
|
| 56 |
+
"normalized_bpb": 1.6713139089025868,
|
| 57 |
+
"step": 1000
|
| 58 |
+
},
|
| 59 |
+
"arithmark_3.0": {
|
| 60 |
+
"dataset": "AxiomicLabs/Arithmark-3.0",
|
| 61 |
+
"dataset_revision": "6f6e59dd9b7e2c63455f7af7f838f9ecc3d0a746",
|
| 62 |
+
"primary_metric": "acc_norm",
|
| 63 |
+
"metric_unit": "fraction",
|
| 64 |
+
"acc": 0.314,
|
| 65 |
+
"acc_norm": 0.297,
|
| 66 |
+
"raw_correct": 314,
|
| 67 |
+
"normalized_correct": 297,
|
| 68 |
+
"total": 1000,
|
| 69 |
+
"max_context": 1024,
|
| 70 |
+
"by_category": {
|
| 71 |
+
"elementary_school_math_continuation::addition::grades_1_2::easy": {
|
| 72 |
+
"acc": 0.21875,
|
| 73 |
+
"acc_norm": 0.2109375,
|
| 74 |
+
"raw_correct": 28,
|
| 75 |
+
"normalized_correct": 27,
|
| 76 |
+
"total": 128
|
| 77 |
+
},
|
| 78 |
+
"elementary_school_math_continuation::comparison::grades_2_3::medium": {
|
| 79 |
+
"acc": 0.1590909090909091,
|
| 80 |
+
"acc_norm": 0.13636363636363635,
|
| 81 |
+
"raw_correct": 7,
|
| 82 |
+
"normalized_correct": 6,
|
| 83 |
+
"total": 44
|
| 84 |
+
},
|
| 85 |
+
"elementary_school_math_continuation::comparison_difference::grades_2_3::medium": {
|
| 86 |
+
"acc": 0.25,
|
| 87 |
+
"acc_norm": 0.25,
|
| 88 |
+
"raw_correct": 12,
|
| 89 |
+
"normalized_correct": 12,
|
| 90 |
+
"total": 48
|
| 91 |
+
},
|
| 92 |
+
"elementary_school_math_continuation::data::grades_2_3::easy": {
|
| 93 |
+
"acc": 0.16279069767441862,
|
| 94 |
+
"acc_norm": 0.16279069767441862,
|
| 95 |
+
"raw_correct": 7,
|
| 96 |
+
"normalized_correct": 7,
|
| 97 |
+
"total": 43
|
| 98 |
+
},
|
| 99 |
+
"elementary_school_math_continuation::division::grades_3_4::medium": {
|
| 100 |
+
"acc": 0.2962962962962963,
|
| 101 |
+
"acc_norm": 0.2962962962962963,
|
| 102 |
+
"raw_correct": 16,
|
| 103 |
+
"normalized_correct": 16,
|
| 104 |
+
"total": 54
|
| 105 |
+
},
|
| 106 |
+
"elementary_school_math_continuation::fractions_counting::grades_3_4::medium": {
|
| 107 |
+
"acc": 0.22,
|
| 108 |
+
"acc_norm": 0.22,
|
| 109 |
+
"raw_correct": 11,
|
| 110 |
+
"normalized_correct": 11,
|
| 111 |
+
"total": 50
|
| 112 |
+
},
|
| 113 |
+
"elementary_school_math_continuation::geometry_area::grades_4_5::medium": {
|
| 114 |
+
"acc": 0.4423076923076923,
|
| 115 |
+
"acc_norm": 0.34615384615384615,
|
| 116 |
+
"raw_correct": 23,
|
| 117 |
+
"normalized_correct": 18,
|
| 118 |
+
"total": 52
|
| 119 |
+
},
|
| 120 |
+
"elementary_school_math_continuation::geometry_perimeter::grades_4_5::medium": {
|
| 121 |
+
"acc": 0.6222222222222222,
|
| 122 |
+
"acc_norm": 0.5777777777777777,
|
| 123 |
+
"raw_correct": 28,
|
| 124 |
+
"normalized_correct": 26,
|
| 125 |
+
"total": 45
|
| 126 |
+
},
|
| 127 |
+
"elementary_school_math_continuation::measurement::grades_2_3::easy": {
|
| 128 |
+
"acc": 0.3157894736842105,
|
| 129 |
+
"acc_norm": 0.2631578947368421,
|
| 130 |
+
"raw_correct": 24,
|
| 131 |
+
"normalized_correct": 20,
|
| 132 |
+
"total": 76
|
| 133 |
+
},
|
| 134 |
+
"elementary_school_math_continuation::money::grades_3_4::medium": {
|
| 135 |
+
"acc": 0.265625,
|
| 136 |
+
"acc_norm": 0.25,
|
| 137 |
+
"raw_correct": 17,
|
| 138 |
+
"normalized_correct": 16,
|
| 139 |
+
"total": 64
|
| 140 |
+
},
|
| 141 |
+
"elementary_school_math_continuation::multiplication::grades_3_4::medium": {
|
| 142 |
+
"acc": 0.35135135135135137,
|
| 143 |
+
"acc_norm": 0.32432432432432434,
|
| 144 |
+
"raw_correct": 26,
|
| 145 |
+
"normalized_correct": 24,
|
| 146 |
+
"total": 74
|
| 147 |
+
},
|
| 148 |
+
"elementary_school_math_continuation::patterns::grades_3_4::medium": {
|
| 149 |
+
"acc": 0.3018867924528302,
|
| 150 |
+
"acc_norm": 0.2830188679245283,
|
| 151 |
+
"raw_correct": 16,
|
| 152 |
+
"normalized_correct": 15,
|
| 153 |
+
"total": 53
|
| 154 |
+
},
|
| 155 |
+
"elementary_school_math_continuation::subtraction::grades_1_2::easy": {
|
| 156 |
+
"acc": 0.21367521367521367,
|
| 157 |
+
"acc_norm": 0.2222222222222222,
|
| 158 |
+
"raw_correct": 25,
|
| 159 |
+
"normalized_correct": 26,
|
| 160 |
+
"total": 117
|
| 161 |
+
},
|
| 162 |
+
"elementary_school_math_continuation::time::grades_2_3::easy": {
|
| 163 |
+
"acc": 0.9454545454545454,
|
| 164 |
+
"acc_norm": 0.8909090909090909,
|
| 165 |
+
"raw_correct": 52,
|
| 166 |
+
"normalized_correct": 49,
|
| 167 |
+
"total": 55
|
| 168 |
+
},
|
| 169 |
+
"elementary_school_math_continuation::two_step_add_subtract::grades_2_3::medium": {
|
| 170 |
+
"acc": 0.1956521739130435,
|
| 171 |
+
"acc_norm": 0.2391304347826087,
|
| 172 |
+
"raw_correct": 9,
|
| 173 |
+
"normalized_correct": 11,
|
| 174 |
+
"total": 46
|
| 175 |
+
},
|
| 176 |
+
"elementary_school_math_continuation::two_step_addition::grades_2_3::medium": {
|
| 177 |
+
"acc": 0.10526315789473684,
|
| 178 |
+
"acc_norm": 0.10526315789473684,
|
| 179 |
+
"raw_correct": 2,
|
| 180 |
+
"normalized_correct": 2,
|
| 181 |
+
"total": 19
|
| 182 |
+
},
|
| 183 |
+
"elementary_school_math_continuation::two_step_subtraction::grades_2_3::medium": {
|
| 184 |
+
"acc": 0.34375,
|
| 185 |
+
"acc_norm": 0.34375,
|
| 186 |
+
"raw_correct": 11,
|
| 187 |
+
"normalized_correct": 11,
|
| 188 |
+
"total": 32
|
| 189 |
+
}
|
| 190 |
+
},
|
| 191 |
+
"choice_batch_size": 64
|
| 192 |
+
},
|
| 193 |
+
"validation": {
|
| 194 |
+
"loss": 3.696356708225175,
|
| 195 |
+
"perplexity": 40.300211143236396,
|
| 196 |
+
"normalized_bpb": 1.6713139089025868,
|
| 197 |
+
"tokens": 359037,
|
| 198 |
+
"normalized_utf8_bytes": 1145591,
|
| 199 |
+
"window": 1024,
|
| 200 |
+
"stride": 512
|
| 201 |
+
}
|
| 202 |
+
}
|
config.json
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"ACSwiGLUForCausalLM"
|
| 4 |
+
],
|
| 5 |
+
"auto_map": {
|
| 6 |
+
"AutoConfig": "modeling_ac_swiglu.ACSwiGLUConfig",
|
| 7 |
+
"AutoModelForCausalLM": "modeling_ac_swiglu.ACSwiGLUForCausalLM"
|
| 8 |
+
},
|
| 9 |
+
"model_type": "ac_swiglu",
|
| 10 |
+
"vocab_size": 4096,
|
| 11 |
+
"seq_len": 1024,
|
| 12 |
+
"max_position_embeddings": 1024,
|
| 13 |
+
"n_positions": 1024,
|
| 14 |
+
"n_ctx": 1024,
|
| 15 |
+
"d_model": 216,
|
| 16 |
+
"n_layers": 10,
|
| 17 |
+
"n_heads": 6,
|
| 18 |
+
"n_kv_heads": 2,
|
| 19 |
+
"chunk": 24,
|
| 20 |
+
"ac_heads": 3,
|
| 21 |
+
"expand": 2,
|
| 22 |
+
"ac_mix_scale": 1.0,
|
| 23 |
+
"mix_mode": "inclusive_centered_expanded_self_attention",
|
| 24 |
+
"ac_conditioning": "swiglu_hidden_queries_keys_values",
|
| 25 |
+
"ac_mixing_backend": "dense_18_by_18_softmax",
|
| 26 |
+
"tokenizer_source": "AxiomicLabs/GPT-S-5M",
|
| 27 |
+
"tokenizer_revision": "275b9c3ca78736bf6aeb154c7e2d5f5764fe9035",
|
| 28 |
+
"bos_token_id": 1,
|
| 29 |
+
"eos_token_id": 2,
|
| 30 |
+
"pad_token_id": 2,
|
| 31 |
+
"tie_word_embeddings": true,
|
| 32 |
+
"is_decoder": true,
|
| 33 |
+
"is_encoder_decoder": false,
|
| 34 |
+
"use_cache": false,
|
| 35 |
+
"safe_serialization": true,
|
| 36 |
+
"minimum_torch_version": "2.5",
|
| 37 |
+
"torch_dtype": "bfloat16",
|
| 38 |
+
"transformers_version": "5.14.1"
|
| 39 |
+
}
|
generation_config.json
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token_id": 1,
|
| 3 |
+
"eos_token_id": 2,
|
| 4 |
+
"pad_token_id": 2,
|
| 5 |
+
"do_sample": false,
|
| 6 |
+
"use_cache": false,
|
| 7 |
+
"repetition_penalty": 1.2
|
| 8 |
+
}
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:f4541226faf35b590a94cf426b0de2e387df5bf1a5d718bbb7951a5c57923de0
|
| 3 |
+
size 9973768
|
modeling_ac_swiglu.py
ADDED
|
@@ -0,0 +1,415 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
|
| 2 |
+
import math
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn as nn
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
try:
|
| 8 |
+
from transformers import GenerationMixin
|
| 9 |
+
except ImportError:
|
| 10 |
+
from transformers.generation import GenerationMixin
|
| 11 |
+
from transformers import PretrainedConfig, PreTrainedModel
|
| 12 |
+
from transformers.modeling_outputs import CausalLMOutputWithPast
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class ACSwiGLUConfig(PretrainedConfig):
|
| 16 |
+
model_type = "ac_swiglu"
|
| 17 |
+
attribute_map = {
|
| 18 |
+
"hidden_size": "d_model",
|
| 19 |
+
"num_hidden_layers": "n_layers",
|
| 20 |
+
"num_attention_heads": "n_heads",
|
| 21 |
+
"num_key_value_heads": "n_kv_heads",
|
| 22 |
+
}
|
| 23 |
+
|
| 24 |
+
def __init__(
|
| 25 |
+
self,
|
| 26 |
+
vocab_size=4096,
|
| 27 |
+
seq_len=1024,
|
| 28 |
+
d_model=216,
|
| 29 |
+
n_layers=10,
|
| 30 |
+
n_heads=6,
|
| 31 |
+
n_kv_heads=2,
|
| 32 |
+
chunk=24,
|
| 33 |
+
ac_heads=3,
|
| 34 |
+
expand=2,
|
| 35 |
+
ac_mix_scale=1.0,
|
| 36 |
+
max_position_embeddings=None,
|
| 37 |
+
n_positions=None,
|
| 38 |
+
n_ctx=None,
|
| 39 |
+
tie_word_embeddings=True,
|
| 40 |
+
**kwargs,
|
| 41 |
+
):
|
| 42 |
+
if min(vocab_size, seq_len, d_model, n_layers) <= 0:
|
| 43 |
+
raise ValueError("Vocabulary, context, width, and depth must be positive.")
|
| 44 |
+
if min(n_heads, n_kv_heads, chunk, ac_heads, expand) <= 0:
|
| 45 |
+
raise ValueError(
|
| 46 |
+
"Attention, AC-SwiGLU, and expansion dimensions must be positive."
|
| 47 |
+
)
|
| 48 |
+
if d_model % n_heads or n_heads % n_kv_heads:
|
| 49 |
+
raise ValueError(
|
| 50 |
+
"d_model % n_heads and n_heads % n_kv_heads must be zero."
|
| 51 |
+
)
|
| 52 |
+
if (d_model // n_heads) % 2:
|
| 53 |
+
raise ValueError("The token-attention head dimension must be even for RoPE.")
|
| 54 |
+
if d_model % chunk or (d_model * expand) % chunk or chunk % ac_heads:
|
| 55 |
+
raise ValueError(
|
| 56 |
+
"AC-SwiGLU requires d_model and d_model * expand divisible by "
|
| 57 |
+
"chunk, and chunk divisible by ac_heads."
|
| 58 |
+
)
|
| 59 |
+
if not math.isfinite(ac_mix_scale) or ac_mix_scale <= 0.0:
|
| 60 |
+
raise ValueError("ac_mix_scale must be finite and positive.")
|
| 61 |
+
kwargs.setdefault("is_decoder", True)
|
| 62 |
+
kwargs.setdefault("is_encoder_decoder", False)
|
| 63 |
+
kwargs.setdefault("tie_word_embeddings", tie_word_embeddings)
|
| 64 |
+
# AC-SwiGLU exports intentionally recompute the full visible context during
|
| 65 |
+
# generation; this architecture does not expose a KV cache.
|
| 66 |
+
kwargs.setdefault("use_cache", False)
|
| 67 |
+
super().__init__(**kwargs)
|
| 68 |
+
self.vocab_size = vocab_size
|
| 69 |
+
self.seq_len = seq_len
|
| 70 |
+
self.max_position_embeddings = max_position_embeddings or seq_len
|
| 71 |
+
self.n_positions = n_positions or self.max_position_embeddings
|
| 72 |
+
self.n_ctx = n_ctx or self.max_position_embeddings
|
| 73 |
+
self.d_model = d_model
|
| 74 |
+
self.n_layers = n_layers
|
| 75 |
+
self.n_heads = n_heads
|
| 76 |
+
self.n_kv_heads = n_kv_heads
|
| 77 |
+
self.chunk = chunk
|
| 78 |
+
self.ac_heads = ac_heads
|
| 79 |
+
self.expand = expand
|
| 80 |
+
self.ac_mix_scale = ac_mix_scale
|
| 81 |
+
self.head_dim = d_model // n_heads
|
| 82 |
+
self.num_key_value_groups = n_heads // n_kv_heads
|
| 83 |
+
self.is_decoder = True
|
| 84 |
+
self.use_cache = False
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
class RMSNorm(nn.Module):
|
| 88 |
+
def __init__(self, d):
|
| 89 |
+
super().__init__()
|
| 90 |
+
self.w = nn.Parameter(torch.ones(d))
|
| 91 |
+
|
| 92 |
+
def forward(self, x):
|
| 93 |
+
return F.rms_norm(
|
| 94 |
+
x, (x.shape[-1],), self.w.to(dtype=x.dtype), eps=1e-6
|
| 95 |
+
)
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def rope_cache(seq_len, head_dim, device, base=10000.0):
|
| 99 |
+
inv = 1.0 / (
|
| 100 |
+
base ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim)
|
| 101 |
+
)
|
| 102 |
+
t = torch.arange(seq_len, device=device).float()
|
| 103 |
+
freqs = torch.outer(t, inv)
|
| 104 |
+
return torch.cos(freqs), torch.sin(freqs)
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def apply_rope(x, cos, sin):
|
| 108 |
+
x1, x2 = x.chunk(2, dim=-1)
|
| 109 |
+
return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
class TokenAttention(nn.Module):
|
| 113 |
+
def __init__(self, d, n_heads, n_kv_heads):
|
| 114 |
+
super().__init__()
|
| 115 |
+
if d % n_heads or n_heads % n_kv_heads:
|
| 116 |
+
raise ValueError("d_model % n_heads and n_heads % n_kv_heads must be zero.")
|
| 117 |
+
self.h, self.kv_h, self.hd = n_heads, n_kv_heads, d // n_heads
|
| 118 |
+
version = torch.__version__.split("+", 1)[0].split(".")
|
| 119 |
+
self.native_gqa = tuple(map(int, version[:2])) >= (2, 5)
|
| 120 |
+
self.q = nn.Linear(d, d, bias=False)
|
| 121 |
+
self.k = nn.Linear(d, n_kv_heads * self.hd, bias=False)
|
| 122 |
+
self.v = nn.Linear(d, n_kv_heads * self.hd, bias=False)
|
| 123 |
+
self.o = nn.Linear(d, d, bias=False)
|
| 124 |
+
self.qn, self.kn = RMSNorm(self.hd), RMSNorm(self.hd)
|
| 125 |
+
|
| 126 |
+
def forward(self, x, cos, sin):
|
| 127 |
+
B, T, d = x.shape
|
| 128 |
+
q = self.q(x).view(B, T, self.h, self.hd).transpose(1, 2)
|
| 129 |
+
k = self.k(x).view(B, T, self.kv_h, self.hd).transpose(1, 2)
|
| 130 |
+
v = self.v(x).view(B, T, self.kv_h, self.hd).transpose(1, 2)
|
| 131 |
+
q, k = self.qn(q), self.kn(k)
|
| 132 |
+
q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)
|
| 133 |
+
if self.h != self.kv_h and not self.native_gqa:
|
| 134 |
+
repeats = self.h // self.kv_h
|
| 135 |
+
k = k.repeat_interleave(repeats, dim=1)
|
| 136 |
+
v = v.repeat_interleave(repeats, dim=1)
|
| 137 |
+
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
|
| 138 |
+
else:
|
| 139 |
+
y = F.scaled_dot_product_attention(
|
| 140 |
+
q,
|
| 141 |
+
k,
|
| 142 |
+
v,
|
| 143 |
+
is_causal=True,
|
| 144 |
+
enable_gqa=self.h != self.kv_h,
|
| 145 |
+
)
|
| 146 |
+
return self.o(y.transpose(1, 2).reshape(B, T, d))
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
class ACSwiGLU(nn.Module):
|
| 150 |
+
# Fused SwiGLU with centered attention inside its expanded representation.
|
| 151 |
+
|
| 152 |
+
def __init__(self, d, chunk=24, heads=3, expand=2, mix_scale=1.0):
|
| 153 |
+
super().__init__()
|
| 154 |
+
expanded = d * expand
|
| 155 |
+
if min(d, chunk, heads, expand) <= 0:
|
| 156 |
+
raise ValueError("AC-SwiGLU dimensions and expansion must be positive.")
|
| 157 |
+
if d % chunk or expanded % chunk or chunk % heads:
|
| 158 |
+
raise ValueError(
|
| 159 |
+
"AC-SwiGLU requires d and d * expand divisible by chunk, "
|
| 160 |
+
"and chunk divisible by heads."
|
| 161 |
+
)
|
| 162 |
+
if not math.isfinite(mix_scale) or mix_scale <= 0.0:
|
| 163 |
+
raise ValueError("ac_mix_scale must be finite and positive.")
|
| 164 |
+
self.d, self.m, self.c, self.h = d, expanded, chunk, heads
|
| 165 |
+
self.n = expanded // chunk
|
| 166 |
+
self.hd, self.expand = chunk // heads, expand
|
| 167 |
+
self.mix_scale = float(mix_scale)
|
| 168 |
+
self.in_proj = nn.Linear(d, 2 * expanded, bias=False)
|
| 169 |
+
self.down = nn.Linear(expanded, d, bias=False)
|
| 170 |
+
self.q_scale = nn.Parameter(torch.ones(chunk))
|
| 171 |
+
self.k_scale = nn.Parameter(torch.zeros(chunk))
|
| 172 |
+
self.bias = nn.Parameter(torch.zeros(heads, self.n, self.n))
|
| 173 |
+
nn.init.zeros_(self.down.weight)
|
| 174 |
+
|
| 175 |
+
def _mix_expanded(self, x):
|
| 176 |
+
B, T, _ = x.shape
|
| 177 |
+
batch_tokens = B * T
|
| 178 |
+
up, gate = self.in_proj(x).chunk(2, dim=-1)
|
| 179 |
+
hidden = up * F.silu(gate)
|
| 180 |
+
values = hidden.reshape(
|
| 181 |
+
batch_tokens, self.n, self.h, self.hd
|
| 182 |
+
).transpose(1, 2)
|
| 183 |
+
normalized = F.rms_norm(values, (self.hd,), eps=1e-6)
|
| 184 |
+
q = normalized * self.q_scale.view(1, self.h, 1, self.hd).to(
|
| 185 |
+
normalized.dtype
|
| 186 |
+
)
|
| 187 |
+
k = normalized * self.k_scale.view(1, self.h, 1, self.hd).to(
|
| 188 |
+
normalized.dtype
|
| 189 |
+
)
|
| 190 |
+
logits = (q @ k.transpose(-2, -1)) * (self.hd ** -0.5)
|
| 191 |
+
logits = logits + self.bias.unsqueeze(0).to(logits.dtype)
|
| 192 |
+
attn = F.softmax(logits, dim=-1, dtype=torch.float32).to(values.dtype)
|
| 193 |
+
reference = F.softmax(
|
| 194 |
+
self.bias.float(), dim=-1
|
| 195 |
+
).to(values.dtype).unsqueeze(0)
|
| 196 |
+
delta_attn = attn - reference
|
| 197 |
+
mixed = values + self.mix_scale * (delta_attn @ values)
|
| 198 |
+
return hidden, values, mixed, logits, attn, delta_attn
|
| 199 |
+
|
| 200 |
+
def forward(self, x):
|
| 201 |
+
B, T, _ = x.shape
|
| 202 |
+
_, _, mixed, _, _, _ = self._mix_expanded(x)
|
| 203 |
+
mixed = mixed.transpose(1, 2).reshape(B, T, self.m)
|
| 204 |
+
return self.down(mixed)
|
| 205 |
+
|
| 206 |
+
@torch.no_grad()
|
| 207 |
+
def diagnostic_stats(self, x):
|
| 208 |
+
probe_count = x.shape[0]
|
| 209 |
+
B, T, _ = x.shape
|
| 210 |
+
hidden, values, mixed, logits, attn, delta_attn = self._mix_expanded(x)
|
| 211 |
+
mixed_hidden = mixed.transpose(1, 2).reshape(B, T, self.m)
|
| 212 |
+
base_output = self.down(hidden)
|
| 213 |
+
output = self.down(mixed_hidden)
|
| 214 |
+
flat_output = output.reshape(-1, self.d)
|
| 215 |
+
token_effect = (
|
| 216 |
+
flat_output.float().norm(dim=-1)
|
| 217 |
+
/ x.reshape(-1, self.d).float().norm(dim=-1).clamp_min(1e-12)
|
| 218 |
+
)
|
| 219 |
+
interaction_effect = (
|
| 220 |
+
(output - base_output).reshape(-1, self.d).float().norm(dim=-1)
|
| 221 |
+
/ flat_output.float().norm(dim=-1).clamp_min(1e-12)
|
| 222 |
+
)
|
| 223 |
+
probs = attn.float().clamp_min(1e-9)
|
| 224 |
+
entropy = -(probs * probs.log()).sum(dim=-1) / math.log(self.n)
|
| 225 |
+
identity_coefficient = (
|
| 226 |
+
1.0
|
| 227 |
+
+ self.mix_scale
|
| 228 |
+
* delta_attn.float().diagonal(dim1=-2, dim2=-1)
|
| 229 |
+
).mean()
|
| 230 |
+
if self.h > 1:
|
| 231 |
+
flattened = probs.transpose(0, 1).reshape(self.h, -1)
|
| 232 |
+
normalized = F.normalize(flattened, dim=-1)
|
| 233 |
+
similarity = normalized @ normalized.mT
|
| 234 |
+
head_similarity = (
|
| 235 |
+
similarity.sum() - similarity.diagonal().sum()
|
| 236 |
+
) / (self.h * (self.h - 1))
|
| 237 |
+
else:
|
| 238 |
+
head_similarity = probs.new_zeros(())
|
| 239 |
+
mix_delta = mixed - values
|
| 240 |
+
retention_cosine = F.cosine_similarity(
|
| 241 |
+
mixed.float().reshape(B * T, -1),
|
| 242 |
+
values.float().reshape(B * T, -1),
|
| 243 |
+
dim=-1,
|
| 244 |
+
).mean()
|
| 245 |
+
return {
|
| 246 |
+
"token_effect": token_effect,
|
| 247 |
+
"interaction_effect": interaction_effect,
|
| 248 |
+
"probe_effect": interaction_effect.reshape(probe_count, -1).mean(dim=-1),
|
| 249 |
+
"entropy": entropy.mean().item(),
|
| 250 |
+
"identity_coefficient": identity_coefficient.item(),
|
| 251 |
+
"dominant_mass": probs.max(dim=-1).values.mean().item(),
|
| 252 |
+
"head_similarity": head_similarity.item(),
|
| 253 |
+
"delta_rms": delta_attn.float().square().mean().sqrt().item(),
|
| 254 |
+
"mix_ratio": (
|
| 255 |
+
mix_delta.float().norm()
|
| 256 |
+
/ values.float().norm().clamp_min(1e-9)
|
| 257 |
+
).item(),
|
| 258 |
+
"retention_cosine": retention_cosine.item(),
|
| 259 |
+
"mix_gain": (
|
| 260 |
+
mixed.float().norm()
|
| 261 |
+
/ values.float().norm().clamp_min(1e-9)
|
| 262 |
+
).item(),
|
| 263 |
+
"key_scale_abs_mean": self.k_scale.float().abs().mean().item(),
|
| 264 |
+
"logit_max": logits.float().abs().max().item(),
|
| 265 |
+
}
|
| 266 |
+
|
| 267 |
+
|
| 268 |
+
class Block(nn.Module):
|
| 269 |
+
def __init__(self, config):
|
| 270 |
+
super().__init__()
|
| 271 |
+
d = config.d_model
|
| 272 |
+
self.n1, self.n2 = RMSNorm(d), RMSNorm(d)
|
| 273 |
+
self.attn = TokenAttention(d, config.n_heads, config.n_kv_heads)
|
| 274 |
+
self.mix = ACSwiGLU(
|
| 275 |
+
d,
|
| 276 |
+
config.chunk,
|
| 277 |
+
config.ac_heads,
|
| 278 |
+
config.expand,
|
| 279 |
+
config.ac_mix_scale,
|
| 280 |
+
)
|
| 281 |
+
|
| 282 |
+
def forward(self, x, cos, sin):
|
| 283 |
+
x = x + self.attn(self.n1(x), cos, sin)
|
| 284 |
+
x = x + self.mix(self.n2(x))
|
| 285 |
+
return x
|
| 286 |
+
|
| 287 |
+
|
| 288 |
+
class ACSwiGLUModel(nn.Module):
|
| 289 |
+
def __init__(self, config):
|
| 290 |
+
super().__init__()
|
| 291 |
+
d = config.d_model
|
| 292 |
+
self.config = config
|
| 293 |
+
self.emb = nn.Embedding(config.vocab_size, d)
|
| 294 |
+
self.blocks = nn.ModuleList(Block(config) for _ in range(config.n_layers))
|
| 295 |
+
self.norm = RMSNorm(d)
|
| 296 |
+
hd = d // config.n_heads
|
| 297 |
+
cos, sin = rope_cache(config.seq_len, hd, "cpu")
|
| 298 |
+
self.register_buffer("cos", cos)
|
| 299 |
+
self.register_buffer("sin", sin)
|
| 300 |
+
|
| 301 |
+
def forward_hidden(self, idx):
|
| 302 |
+
if idx.size(1) > self.config.seq_len:
|
| 303 |
+
idx = idx[:, -self.config.seq_len :]
|
| 304 |
+
x = self.emb(idx)
|
| 305 |
+
device_type = x.device.type
|
| 306 |
+
compute_dtype = (
|
| 307 |
+
torch.get_autocast_dtype(device_type)
|
| 308 |
+
if torch.is_autocast_enabled(device_type)
|
| 309 |
+
else x.dtype
|
| 310 |
+
)
|
| 311 |
+
x = x.to(dtype=compute_dtype)
|
| 312 |
+
cos = self.cos[: idx.size(1)].to(device=idx.device, dtype=compute_dtype)
|
| 313 |
+
sin = self.sin[: idx.size(1)].to(device=idx.device, dtype=compute_dtype)
|
| 314 |
+
for b in self.blocks:
|
| 315 |
+
x = b(x, cos, sin)
|
| 316 |
+
return self.norm(x)
|
| 317 |
+
|
| 318 |
+
|
| 319 |
+
class ACSwiGLUForCausalLM(PreTrainedModel, GenerationMixin):
|
| 320 |
+
config_class = ACSwiGLUConfig
|
| 321 |
+
base_model_prefix = "model"
|
| 322 |
+
_no_split_modules = ["Block"]
|
| 323 |
+
_tied_weights_keys = {"head.weight": "model.emb.weight"}
|
| 324 |
+
# Transformers 4.x reads the expanded map directly; Transformers 5.x
|
| 325 |
+
# replaces it during post_init(). Keeping both forms makes the same remote
|
| 326 |
+
# code load cleanly across that boundary.
|
| 327 |
+
all_tied_weights_keys = {"head.weight": "model.emb.weight"}
|
| 328 |
+
|
| 329 |
+
def __init__(self, config):
|
| 330 |
+
super().__init__(config)
|
| 331 |
+
self.model = ACSwiGLUModel(config)
|
| 332 |
+
self.head = nn.Linear(config.d_model, config.vocab_size, bias=False)
|
| 333 |
+
# Modern Transformers creates loader metadata and performs configured
|
| 334 |
+
# tying in post_init(); omitting it leaves all_tied_weights_keys absent.
|
| 335 |
+
self.post_init()
|
| 336 |
+
self.head.weight = self.model.emb.weight
|
| 337 |
+
|
| 338 |
+
def get_input_embeddings(self):
|
| 339 |
+
return self.model.emb
|
| 340 |
+
|
| 341 |
+
def set_input_embeddings(self, value):
|
| 342 |
+
self.model.emb = value
|
| 343 |
+
|
| 344 |
+
def get_output_embeddings(self):
|
| 345 |
+
return self.head
|
| 346 |
+
|
| 347 |
+
def set_output_embeddings(self, new_embeddings):
|
| 348 |
+
self.head = new_embeddings
|
| 349 |
+
|
| 350 |
+
def raw_logits(self, idx):
|
| 351 |
+
return self.head(self.model.forward_hidden(idx))
|
| 352 |
+
|
| 353 |
+
def logits(self, idx):
|
| 354 |
+
return self.raw_logits(idx)
|
| 355 |
+
|
| 356 |
+
def _masked_logits(self, input_ids, attention_mask):
|
| 357 |
+
if attention_mask is None or bool(attention_mask.all()):
|
| 358 |
+
return self.logits(input_ids)
|
| 359 |
+
|
| 360 |
+
B, T = input_ids.shape
|
| 361 |
+
out = None
|
| 362 |
+
for i in range(B):
|
| 363 |
+
keep = attention_mask[i].bool().nonzero(as_tuple=False).flatten()
|
| 364 |
+
if keep.numel() == 0:
|
| 365 |
+
keep = torch.tensor([T - 1], device=input_ids.device)
|
| 366 |
+
trimmed = input_ids[i, keep].unsqueeze(0)
|
| 367 |
+
logits_i = self.logits(trimmed)
|
| 368 |
+
if out is None:
|
| 369 |
+
out = logits_i.new_zeros(B, T, logits_i.size(-1))
|
| 370 |
+
out[i, keep, :] = logits_i[0, -keep.numel() :, :]
|
| 371 |
+
return out
|
| 372 |
+
|
| 373 |
+
def forward(
|
| 374 |
+
self,
|
| 375 |
+
input_ids=None,
|
| 376 |
+
attention_mask=None,
|
| 377 |
+
labels=None,
|
| 378 |
+
use_cache=False,
|
| 379 |
+
past_key_values=None,
|
| 380 |
+
**kwargs,
|
| 381 |
+
):
|
| 382 |
+
return_dict = kwargs.pop(
|
| 383 |
+
"return_dict", getattr(self.config, "use_return_dict", True)
|
| 384 |
+
)
|
| 385 |
+
if input_ids is None:
|
| 386 |
+
raise ValueError("input_ids must be provided.")
|
| 387 |
+
if input_ids.size(1) > self.config.seq_len:
|
| 388 |
+
input_ids = input_ids[:, -self.config.seq_len :]
|
| 389 |
+
if attention_mask is not None:
|
| 390 |
+
attention_mask = attention_mask[:, -self.config.seq_len :]
|
| 391 |
+
if labels is not None:
|
| 392 |
+
labels = labels[:, -self.config.seq_len :]
|
| 393 |
+
logits = self._masked_logits(input_ids, attention_mask)
|
| 394 |
+
loss = None
|
| 395 |
+
if labels is not None:
|
| 396 |
+
shift_logits = logits[:, :-1, :].contiguous()
|
| 397 |
+
shift_labels = labels[:, 1:].contiguous()
|
| 398 |
+
loss = F.cross_entropy(
|
| 399 |
+
shift_logits.view(-1, shift_logits.size(-1)).float(),
|
| 400 |
+
shift_labels.view(-1),
|
| 401 |
+
ignore_index=-100,
|
| 402 |
+
)
|
| 403 |
+
if not return_dict:
|
| 404 |
+
return (loss, logits) if loss is not None else (logits,)
|
| 405 |
+
return CausalLMOutputWithPast(loss=loss, logits=logits, past_key_values=None)
|
| 406 |
+
|
| 407 |
+
def prepare_inputs_for_generation(self, input_ids, **kwargs):
|
| 408 |
+
attention_mask = kwargs.get("attention_mask")
|
| 409 |
+
result = {
|
| 410 |
+
"input_ids": input_ids[:, -self.config.seq_len :],
|
| 411 |
+
"use_cache": False,
|
| 412 |
+
}
|
| 413 |
+
if attention_mask is not None:
|
| 414 |
+
result["attention_mask"] = attention_mask[:, -self.config.seq_len :]
|
| 415 |
+
return result
|
tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"add_prefix_space": true,
|
| 3 |
+
"backend": "tokenizers",
|
| 4 |
+
"bos_token": "<bos>",
|
| 5 |
+
"clean_up_tokenization_spaces": false,
|
| 6 |
+
"eos_token": "<eos>",
|
| 7 |
+
"is_local": false,
|
| 8 |
+
"local_files_only": false,
|
| 9 |
+
"model_max_length": 1024,
|
| 10 |
+
"pad_token": "<eos>",
|
| 11 |
+
"tokenizer_class": "TokenizersBackend",
|
| 12 |
+
"unk_token": "<unk>",
|
| 13 |
+
"vocab_size": 4096
|
| 14 |
+
}
|
training_config.json
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"steps": 40000,
|
| 3 |
+
"session_steps": 40000,
|
| 4 |
+
"session_count": 1,
|
| 5 |
+
"batch_size": 512,
|
| 6 |
+
"grad_accum": 1,
|
| 7 |
+
"seq_len": 1024,
|
| 8 |
+
"model_context": 1024,
|
| 9 |
+
"d_model": 216,
|
| 10 |
+
"n_layers": 10,
|
| 11 |
+
"n_heads": 6,
|
| 12 |
+
"n_kv_heads": 2,
|
| 13 |
+
"chunk": 24,
|
| 14 |
+
"ac_heads": 3,
|
| 15 |
+
"expand": 2,
|
| 16 |
+
"ac_mix_scale": 1.0,
|
| 17 |
+
"lr": 0.0025,
|
| 18 |
+
"linear_decay_start": 30000,
|
| 19 |
+
"cosine_decay_start": 30000,
|
| 20 |
+
"mid_lr": 0.0025,
|
| 21 |
+
"muon_lr": 0.03,
|
| 22 |
+
"muon_momentum": 0.95,
|
| 23 |
+
"muon_ns_steps": 5,
|
| 24 |
+
"muon_adjust_lr_fn": null,
|
| 25 |
+
"warmup": 1000,
|
| 26 |
+
"wd": 0.01,
|
| 27 |
+
"grad_clip": 1.0,
|
| 28 |
+
"log_every": 10,
|
| 29 |
+
"diag_every": 1000,
|
| 30 |
+
"eval_every": 10000,
|
| 31 |
+
"val_batch_size": 32,
|
| 32 |
+
"val_batches": 0,
|
| 33 |
+
"val_context": 1024,
|
| 34 |
+
"val_stride": 512,
|
| 35 |
+
"infer_tokens": 512,
|
| 36 |
+
"infer_repeat_penalty": 1.2,
|
| 37 |
+
"infer_prompt": "The process of photosynthesis",
|
| 38 |
+
"lm_eval_tasks": "arc_easy,arc_challenge,hellaswag,piqa",
|
| 39 |
+
"lm_eval_batch_size": "32",
|
| 40 |
+
"lm_eval_device": "cuda",
|
| 41 |
+
"lm_eval_dtype": "bfloat16",
|
| 42 |
+
"lm_eval_softmax_dtype": "float32",
|
| 43 |
+
"lm_eval_expected_version": "0.4.12",
|
| 44 |
+
"lm_eval_retries": 3,
|
| 45 |
+
"lm_eval_export_dir": "AC_SwiGLU_Inclusive_40k_lm_eval_hf",
|
| 46 |
+
"lm_eval_output_dir": "lm_eval_results_AC_SwiGLU_Inclusive_40k",
|
| 47 |
+
"arithmark3_choice_batch_size": 64,
|
| 48 |
+
"arithmark3_force_download": false,
|
| 49 |
+
"recipe_version": "ACSwiGLUInclusive_d216_l10_q6_kv2_chunk24_h3_expand2_centered18x18_mix1_train1024_ctx1024_v4k_gpts4944k_fullcompile_adam2p5e3_hold30k_cos0_40k_muon3e2_w1k_clip1_20p97B_b512_a1_tpu524k_session1x40k_static_fwe55_cos25_fwhq10_math10_interleaved_pinneddata_buf1k_prefetch16_eval1024s512_b32_arithb64_eval10k_diag1k_arithbos_scorew75_v1",
|
| 50 |
+
"hf_repo_id": "User01110/attention-contraction-mini",
|
| 51 |
+
"hf_repo_private": false,
|
| 52 |
+
"hub_upload_retries": 3,
|
| 53 |
+
"resume_branch": "resume-latest",
|
| 54 |
+
"resume_model_file": "resume_model.safetensors",
|
| 55 |
+
"resume_state_file": "resume_state.pt",
|
| 56 |
+
"resume_manifest_file": "resume_manifest.json",
|
| 57 |
+
"tokenizer_name": "AxiomicLabs/GPT-S-5M",
|
| 58 |
+
"tokenizer_revision": "275b9c3ca78736bf6aeb154c7e2d5f5764fe9035",
|
| 59 |
+
"data_seed": 1337,
|
| 60 |
+
"shuffle_buffer": 1024,
|
| 61 |
+
"tokenize_batch_size": 64,
|
| 62 |
+
"prefetch_batches": 16,
|
| 63 |
+
"compile": true
|
| 64 |
+
}
|