Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
300 changes: 300 additions & 0 deletions examples/tiny.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,300 @@
//! Run the experimental BTQ1 model without changing `betlang::detect`.
//!
//! cargo run --release --example tiny -- MODEL.bin SOURCE
//! cargo run --release --example tiny -- MODEL.bin --compare CACHE/test.paths.txt

use std::{env, error::Error, fs, io, path::Path, time::Instant};

use rayon::prelude::*;

const BINS: usize = 512;
const HIDDEN: usize = 16;
const CLASSES: usize = 48;
const LABELS: [&str; CLASSES] = [
"asm",
"batch",
"c",
"clojure",
"cmake",
"cobol",
"cpp",
"cs",
"css",
"dart",
"dockerfile",
"elixir",
"erlang",
"gemfile",
"gemspec",
"go",
"gradle",
"groovy",
"haskell",
"html",
"ini",
"java",
"javascript",
"json",
"julia",
"kotlin",
"lisp",
"lua",
"markdown",
"objectivec",
"ocaml",
"perl",
"php",
"powershell",
"python",
"r",
"ruby",
"rust",
"scala",
"shell",
"sql",
"swift",
"toml",
"typescript",
"vba",
"verilog",
"xml",
"yaml",
];

struct Tiny {
kernel: Vec<f32>,
bias: Vec<f32>,
output: Vec<f32>,
output_bias: Vec<f32>,
}

impl Tiny {
fn load(bytes: &[u8]) -> Result<Self, &'static str> {
if bytes.len() != 4752 || bytes[..8] != *b"BTQ1\x01\0\0\0" {
return Err("expected a 4,752-byte BTQ1 model");
}
let floats = |data: &[u8]| {
data.chunks_exact(4)
.map(|chunk| f32::from_le_bytes(chunk.try_into().unwrap()))
.collect::<Vec<_>>()
};
let scales = floats(&bytes[8..16]);
if scales.iter().any(|s| !s.is_finite() || *s <= 0.0) {
return Err("invalid weight scales");
}
let unpack = |data: &[u8], scale: f32| {
data.iter()
.flat_map(|b| [(b & 15) as i8 - 8, (b >> 4) as i8 - 8])
.map(|v| v as f32 * scale)
.collect()
};
let model = Self {
kernel: unpack(&bytes[16..4112], scales[0]),
bias: floats(&bytes[4112..4176]),
output: unpack(&bytes[4176..4560], scales[1]),
output_bias: floats(&bytes[4560..4752]),
};
if model
.bias
.iter()
.chain(&model.output_bias)
.any(|v| !v.is_finite())
{
return Err("non-finite bias");
}
Ok(model)
}

fn logits(&self, source: &[u8]) -> Option<[f32; CLASSES]> {
let inputs = features(source)?;
let mut hidden = [0.0; HIDDEN];
hidden.copy_from_slice(&self.bias);
for (&input, row) in inputs.iter().zip(self.kernel.chunks_exact(HIDDEN)) {
if input != 0.0 {
for (value, weight) in hidden.iter_mut().zip(row) {
*value += input * weight;
}
}
}
let mut logits = [0.0; CLASSES];
logits.copy_from_slice(&self.output_bias);
for (input, row) in hidden.iter().zip(self.output.chunks_exact(CLASSES)) {
for (value, weight) in logits.iter_mut().zip(row) {
*value += input.max(0.0) * weight;
}
}
Some(logits)
}

fn detect(&self, source: &[u8]) -> Option<&'static str> {
let logits = self.logits(source)?;
let index = (0..CLASSES)
.max_by(|&a, &b| logits[a].total_cmp(&logits[b]).then_with(|| b.cmp(&a)))?;
Some(LABELS[index])
}
}

fn hash(bytes: &[u8]) -> u32 {
bytes.iter().fold(2166136261, |h, b| {
(h ^ u32::from(b.to_ascii_lowercase())).wrapping_mul(16777619)
})
}

fn word_start(byte: u8) -> bool {
byte.is_ascii_alphabetic() || byte == b'_' || byte >= 128
}

fn whitespace(byte: &u8) -> bool {
matches!(byte, b' ' | b'\t' | b'\n' | b'\r' | 11 | 12)
}

fn window(source: &[u8]) -> Option<Vec<u8>> {
let begin = source[..source.len().min(4096)].trim_ascii_start();
if begin.len() < 8 {
return None;
}
if begin.len() < 1024 {
return Some(begin.to_vec());
}
let end = source[source.len().saturating_sub(4096)..].trim_ascii_end();
let mut window = begin[..1024].to_vec();
if end.len() >= 1024 {
window.extend_from_slice(&end[end.len() - 1024..]);
}
Some(window)
}

fn features(source: &[u8]) -> Option<[f32; BINS]> {
let bytes = window(source)?;
let mut counts = [0.0f32; BINS];
let mut pos = 0;
let mut previous = 0u32;
while pos < bytes.len() {
let start = pos;
let byte = bytes[pos];
pos += 1;
let value = if word_start(byte) {
while pos < bytes.len() && (word_start(bytes[pos]) || bytes[pos].is_ascii_digit()) {
pos += 1;
}
hash(&bytes[start..pos])
} else if byte.is_ascii_digit() {
while pos < bytes.len() && bytes[pos].is_ascii_digit() {
pos += 1;
}
hash(b"0")
} else if byte == b'\n' || !whitespace(&byte) {
hash(&[byte])
} else {
continue;
};
counts[value as usize % BINS] += 1.0;
if previous != 0 {
let pair = previous.wrapping_mul(16777619) ^ value;
counts[pair as usize % BINS] += 1.0;
}
previous = value;
}
for count in &mut counts {
*count = count.ln_1p();
}
let norm = counts.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
for value in &mut counts {
*value /= norm;
}
}
Some(counts)
}

fn compare(model: &Tiny, root: &Path) -> Result<(), Box<dyn Error>> {
let mut paths = Vec::new();
for line in fs::read_to_string(root)?.lines() {
let path = Path::new(line).to_path_buf();
let slug = path
.parent()
.and_then(Path::file_name)
.and_then(|s| s.to_str());
let label = LABELS
.iter()
.find(|&&label| Some(label) == slug)
.ok_or("unknown label")?;
paths.push((*label, path));
}
paths.sort();
if paths.is_empty() {
return Err("no labeled files found".into());
}
let results: Vec<_> = paths
.par_iter()
.map(|(truth, path)| {
let source = fs::read(path)?;
let tiny = model.detect(&source).unwrap_or("unknown");
let base = betlang::detect(&source)
.language()
.map(|l| l.slug())
.unwrap_or("unknown");
Ok((*truth, tiny, base, path))
})
.collect::<Result<_, io::Error>>()?;
println!("truth\ttiny\tbaseline\tpath");
for (truth, tiny, base, path) in results {
println!("{truth}\t{tiny}\t{base}\t{}", path.display());
}
Ok(())
}

fn main() -> Result<(), Box<dyn Error>> {
let args: Vec<_> = env::args().collect();
if args.len() < 3 {
return Err(
"usage: tiny MODEL.bin SOURCE | MODEL.bin --compare CACHE/test.paths.txt".into(),
);
}
let model = Tiny::load(&fs::read(&args[1])?)?;
if args[2] == "--compare" {
let root = args.get(3).ok_or("--compare requires a paths manifest")?;
return compare(&model, Path::new(root));
}
let source = fs::read(&args[2])?;
if args.get(3).is_some_and(|arg| arg == "--features") {
println!("{:?}", features(&source));
return Ok(());
}
let start = Instant::now();
let logits = model.logits(&source);
let elapsed = start.elapsed();
println!("{}", model.detect(&source).unwrap_or("unknown"));
println!("{logits:?}");
eprintln!("inference_us={:.2}", elapsed.as_secs_f64() * 1e6);
Ok(())
}

#[cfg(test)]
mod tests {
use super::*;

#[test]
fn labels_match_production_order() {
for (index, label) in LABELS.iter().enumerate() {
assert_eq!(label.parse::<betlang::Language>().unwrap() as usize, index);
}
}

#[test]
fn rejects_invalid_models() {
for size in [0, 8, 4751, 4752, 4753] {
assert!(Tiny::load(&vec![0; size]).is_err());
}
}

#[test]
fn rejects_short_inputs_and_normalizes_features() {
assert!(features(b" \n\t").is_none());
assert!(features(b"short").is_none());
for source in [b"fn main() { 123; }\n".as_slice(), &[255u8; 2048]] {
let x = features(source).unwrap();
assert!((x.iter().map(|v| v * v).sum::<f32>() - 1.0).abs() < 1e-5);
}
}
}
94 changes: 94 additions & 0 deletions scripts/TINY_STUDENT.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
# 5 KB lexical student experiment

This applies the small hashed lexical representation idea from
[gpu-lexer](https://gpu-lexer.vercel.app) to file-level language detection.
It is a separate experiment; `betlang::detect` retains its production model.

The published gpu-lexer 0.0.1 runtime uses token kinds, lengths, first/last
characters, two word hashes, shape flags and neighboring punctuation features
as additive embeddings. It combines local depthwise filtering with recurrent
and hierarchical file context. Its advertised 27.5 KB is a **minified,
Brotli-compressed browser bundle**, not an uncompressed weights budget.
This experiment borrows compact hashing and local/global lexical context,
not the WebGPU implementation or its trained weights.

## Architecture and byte budget

The detector uses the same beginning/end byte-window policy as betlang.
It splits that window into identifiers, digit runs, punctuation and newlines.
Identifiers are case-folded; digit runs become `0`. Token and adjacent-token
hashes accumulate in 512 shared buckets. `log1p` counts followed by L2
normalization provide file context. A 16-unit ReLU projection maps these
features to the existing 48 output labels.

| BTQ1 component | Raw bytes |
|---|---:|
| Magic/version | 8 |
| Two f32 quantization scales | 8 |
| 512 × 16 int4 projection | 4,096 |
| 16 f32 hidden biases | 64 |
| 16 × 48 int4 output | 384 |
| 48 f32 output biases | 192 |
| **Total** | **4,752** |

The exporter asserts the exact layout and the 5,000-byte budget.
This is model size only: it excludes inference code, runtime allocations
and executable overhead. Quantization-aware training and export use the
existing trainer's symmetric int4 quantizer.

This has much less context capacity than gpu-lexer or betlang's convolutional
model. A size result is not evidence that it can replace the production model.
In particular, confidence calibration and short ambiguous inputs need their
own validation before exposing this as a production detector.

## Run

Python 3.10 or newer:

```sh
python3 -m venv .venv
.venv/bin/pip install tensorflow-cpu==2.20.0 onnxruntime==1.22.1 numpy==2.2.6
export TF_NUM_INTEROP_THREADS=1 TF_NUM_INTRAOP_THREADS=2

# Dataset layout: files/{train,valid,test}/{production-label}/*
.venv/bin/python scripts/train_tiny_student.py \
--prepare --dataset /path/to/files --cache /path/to/cache \
--output /path/to/tiny-q4.bin
.venv/bin/python scripts/train_tiny_student.py \
--cache /path/to/cache --output /path/to/tiny-q4.bin

cargo run --release --example tiny -- /path/to/tiny-q4.bin source.rs
cargo run --release --example tiny -- /path/to/tiny-q4.bin \
--compare /path/to/cache/test.paths.txt > comparison.tsv
```

Prepare the corpus with repository-disjoint train/validation/test splits
(see `build_finetune_corpus.py`). The feature preparation additionally removes
identical model windows across and within splits, keeping the first occurrence
in train, validation, test order. Re-run preparation whenever the corpus or
tokenizer changes. The paths manifest is the exact evaluation population.
Checkpoint selection uses validation accuracy; test metrics come from the
reloaded packed file. The adjacent JSON report includes all 48 class recalls,
missing test labels, epoch history and the artifact hash. Never compare its
accuracy directly with the model card's score on a different corpus.

The Rust example loads weights once, provides independent native inference,
and can compare both models on the same manifest. The Python inference module
`tiny_student.py` needs only NumPy.

## Checks

```sh
cargo build --release --example tiny
BETLANG_TINY_EXAMPLE="$PWD/target/release/examples/tiny" \
.venv/bin/python scripts/test_tiny_student.py
cargo test --all-targets
cargo fmt --check
cargo clippy --all-targets -- -D warnings
python3 -m flake8 scripts/{tiny_student,train_tiny_student,test_tiny_student}.py \
--select E9,F63,F7,F82
```

The Python tests compare QAT against the reloaded int4 model and compare both
features and logits with Rust on UTF-8, arbitrary bytes, whitespace and window
boundaries. They also exercise malformed files and non-finite parameters.
Loading