#[must_use]
#[cfg_attr(not(feature = "gpu"), allow(unused_variables))]
pub fn compress_gpu(input: &[u8]) -> Option<f64> {
#[cfg(feature = "gpu")]
{
cuda::compress(input)
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub struct CompressPhases {
pub chunks: usize,
pub prior_bytes: usize,
pub host_prep_us: f64,
pub upload_input_us: f64,
pub alloc_us: f64,
pub upload_prior_us: f64,
pub kernel_and_download_us: f64,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct ProjectionStats {
pub model: usize,
pub nodes: usize,
pub placed: usize,
pub mass_kept: f64,
}
#[must_use]
#[cfg_attr(not(feature = "gpu"), allow(unused_variables))]
pub fn project_device_prior(blob: &[u8], bits: usize, cap: u64) -> Option<(Vec<u16>, Vec<ProjectionStats>)> {
#[cfg(feature = "gpu")]
{
Some(cuda::project_prior(blob, bits, cap))
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
#[must_use]
#[cfg_attr(not(feature = "gpu"), allow(unused_variables))]
pub fn compress_gpu_with_prior(input: &[u8], prior: &[u16], bits: usize, predict: bool) -> Option<f64> {
#[cfg(feature = "gpu")]
{
cuda::compress_with(input, prior, bits, predict)
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
pub struct GpuPrior {
#[cfg(feature = "gpu")]
held: cuda::HeldPrior,
}
impl GpuPrior {
#[must_use]
#[cfg_attr(not(feature = "gpu"), allow(unused_variables))]
pub fn upload(table: &[u16], bits: usize, predict: bool) -> Option<Self> {
#[cfg(feature = "gpu")]
{
cuda::hold_prior(table, bits, predict).map(|held| Self { held })
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
#[must_use]
pub fn default_prior() -> Option<Self> {
#[cfg(feature = "gpu")]
{
cuda::hold_default_prior().map(|held| Self { held })
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
#[must_use]
#[cfg_attr(not(feature = "gpu"), allow(unused_variables))]
pub fn compress(&self, input: &[u8]) -> Option<f64> {
#[cfg(feature = "gpu")]
{
cuda::compress_held(input, &self.held)
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
}
#[must_use]
#[cfg_attr(not(feature = "gpu"), allow(unused_variables))]
pub fn compress_gpu_phases(input: &[u8]) -> Option<(f64, CompressPhases)> {
#[cfg(feature = "gpu")]
{
cuda::compress_phases(input)
}
#[cfg(not(feature = "gpu"))]
{
None
}
}
#[cfg(feature = "gpu")]
#[allow(unsafe_code)]
mod cuda {
use std::sync::{Arc, OnceLock};
use cudarc::driver::{CudaContext, CudaFunction, LaunchConfig, PushKernelArg};
use cudarc::nvrtc::Ptx;
use crate::gpu::cuda::load_or_cpu;
const COMPRESS_PTX: &str = include_str!(env!("TREX_COMPRESS_PTX"));
const NCTX: usize = 11;
const MATCH_SIZE: usize = 1 << 14;
const CHUNK: usize = 8192;
const NEMB: usize = 2;
const EMB_SIZE: usize = 1 << 14;
const EMB_DIM: usize = 16;
const PRIOR_CORPUS: &[u8] = include_bytes!("../../_corpus/prior_en.txt");
const PRIOR_BITS: usize = 24;
const PRIOR_CAP: u64 = 8;
const MODEL_ORDER: [usize; NCTX] = [0, 1, 2, 3, 4, 6, 8, 4, 3, 6, 0];
const MODEL_KIND: [u8; NCTX] = [0, 0, 0, 0, 0, 0, 0, 1, 2, 2, 3];
const MIXK: u64 = 0x9E37_79B9_7F4A_7C15;
const FNV_OFFSET: u64 = 14695981039346656037;
const FNV_PRIME: u64 = 1099511628211;
fn fold_byte(b: u8) -> u8 {
if b.is_ascii_uppercase() { b + 32 } else { b }
}
fn shape_byte(b: u8) -> u8 {
if b.is_ascii_digit() {
b'D'
} else if matches!(b, b'a' | b'e' | b'i' | b'o' | b'u' | b'A' | b'E' | b'I' | b'O' | b'U') {
b'V'
} else if b.is_ascii_alphabetic() {
b'C'
} else {
b'.'
}
}
fn model_hash(data: &[u8], t: usize, k: usize) -> u64 {
let mut h = FNV_OFFSET ^ (k as u64 + 1);
if MODEL_KIND[k] == 3 {
let mut s = t;
while s > 0 && data[s - 1].is_ascii_alphanumeric() && t - s < 32 {
s -= 1;
}
for &raw in &data[s..t] {
h = (h ^ u64::from(fold_byte(raw))).wrapping_mul(FNV_PRIME);
}
} else {
let lo = t.saturating_sub(MODEL_ORDER[k]);
for &raw in &data[lo..t] {
let by = match MODEL_KIND[k] {
1 => fold_byte(raw),
2 => shape_byte(raw),
_ => raw,
};
h = (h ^ u64::from(by)).wrapping_mul(FNV_PRIME);
}
}
h
}
fn corpus_prior(data: &[u8], bits: usize) -> Vec<u16> {
let psize = 1usize << bits;
let pmask = psize - 1;
let mut prior = vec![0u16; NCTX * psize * 3];
for t in 0..data.len() {
let byte = data[t];
let mut base = [0u64; NCTX];
for (k, b) in base.iter_mut().enumerate() {
*b = model_hash(data, t, k);
}
let mut node = 1u64;
for bit in (0..8).rev() {
let a = (byte >> bit) & 1;
for (k, &bk) in base.iter().enumerate() {
let key = bk.wrapping_mul(MIXK).wrapping_add(node);
let idx = (key as usize) & pmask;
let chk = (key >> bits) as u16;
let off = (k * psize + idx) * 3;
if prior[off] != chk {
prior[off] = chk;
prior[off + 1] = 0;
prior[off + 2] = 0;
}
if a == 1 { prior[off + 2] += 1 } else { prior[off + 1] += 1 }
if prior[off + 1] + prior[off + 2] > 1024 {
prior[off + 1] = prior[off + 1].div_ceil(2);
prior[off + 2] = prior[off + 2].div_ceil(2);
}
}
node = (node << 1) | u64::from(a);
}
}
prior
}
fn default_prior() -> (&'static [u16], usize) {
static P: OnceLock<Vec<u16>> = OnceLock::new();
let table = P.get_or_init(|| project_prior(crate::seam::baked_model_blob(), PRIOR_BITS, PRIOR_CAP).0);
(table, PRIOR_BITS)
}
fn projection_source(k: usize) -> Option<usize> {
match (MODEL_KIND[k], MODEL_ORDER[k]) {
(3, _) => None,
(_, 0) => Some(1),
(_, o) => Some(o),
}
}
type BlobRows = std::collections::BTreeMap<Vec<u8>, std::collections::BTreeMap<u8, u32>>;
fn push_row_nodes(k: usize, ctx: &[u8], followers: impl Iterator<Item = (u8, u64)>, nodes: &mut Vec<(u64, u64, u64)>) {
let mut h = FNV_OFFSET ^ (k as u64 + 1);
for &b in ctx {
h = (h ^ u64::from(b)).wrapping_mul(FNV_PRIME);
}
let mut n0 = [0u64; 256];
let mut n1 = [0u64; 256];
for (byte, c) in followers {
let mut node = 1usize;
for bit in (0..8).rev() {
let a = (byte >> bit) & 1;
if a == 1 { n1[node] += c } else { n0[node] += c }
node = (node << 1) | usize::from(a);
}
}
for (node, (&zeros, &ones)) in n0.iter().zip(&n1).enumerate().skip(1) {
if zeros + ones > 0 {
nodes.push((h.wrapping_mul(MIXK).wrapping_add(node as u64), zeros, ones));
}
}
}
fn projected_nodes(source: &BlobRows, k: usize) -> Vec<(u64, u64, u64)> {
let width = MODEL_ORDER[k];
let mut nodes = Vec::new();
if MODEL_KIND[k] == 0 && projection_source(k) == Some(width) {
for (ctx, followers) in source {
push_row_nodes(k, ctx, followers.iter().map(|(&b, &c)| (b, u64::from(c))), &mut nodes);
}
return nodes;
}
let mut grouped: std::collections::BTreeMap<Vec<u8>, std::collections::BTreeMap<u8, u64>> =
std::collections::BTreeMap::new();
for (ctx, followers) in source {
let key: Vec<u8> = ctx[ctx.len() - width..]
.iter()
.map(|&b| match MODEL_KIND[k] {
1 => fold_byte(b),
2 => shape_byte(b),
_ => b,
})
.collect();
let sums = grouped.entry(key).or_default();
for (&b, &c) in followers {
*sums.entry(b).or_insert(0) += u64::from(c);
}
}
for (ctx, sums) in &grouped {
push_row_nodes(k, ctx, sums.iter().map(|(&b, &c)| (b, c)), &mut nodes);
}
nodes
}
fn place_nodes(nodes: &mut [(u64, u64, u64)], block: &mut [u16], bits: usize, cap: u64) -> (usize, u64) {
assert!(cap >= 2, "a slot cap below 2 cannot hold one count each way, and halving never reaches it");
let psize = 1usize << bits;
let pmask = psize - 1;
assert_eq!(block.len(), psize * 3, "a block at {bits} bits holds {} u16s", psize * 3);
nodes.sort_unstable_by(|a, b| (b.1 + b.2).cmp(&(a.1 + a.2)).then(a.0.cmp(&b.0)));
block.fill(0);
let mut taken = vec![false; psize];
let mut placed = 0usize;
let mut kept_mass = 0u64;
for &(key, mut n0, mut n1) in nodes.iter() {
let idx = (key as usize) & pmask;
if taken[idx] {
continue;
}
taken[idx] = true;
placed += 1;
kept_mass += n0 + n1;
while n0 + n1 > cap {
n0 = n0.div_ceil(2);
n1 = n1.div_ceil(2);
}
let off = idx * 3;
block[off] = (key >> bits) as u16;
block[off + 1] = u16::try_from(n0).expect("a capped count fits a slot");
block[off + 2] = u16::try_from(n1).expect("a capped count fits a slot");
}
(placed, kept_mass)
}
pub(super) fn project_prior(blob: &[u8], bits: usize, cap: u64) -> (Vec<u16>, Vec<super::ProjectionStats>) {
assert!((1..=28).contains(&bits), "a device prior needs 1 to 28 bits, not {bits}");
assert!(cap >= 2, "a slot cap below 2 cannot hold one count each way, and halving never reaches it");
let rows = crate::seam::decode_baked(blob);
let psize = 1usize << bits;
let mut prior = corpus_prior(PRIOR_CORPUS, bits);
struct ModelJob<'a> {
k: usize,
source: &'a BlobRows,
block: &'a mut [u16],
stats: Option<super::ProjectionStats>,
}
let mut jobs: Vec<ModelJob<'_>> = Vec::with_capacity(NCTX);
for (k, block) in prior.chunks_mut(psize * 3).enumerate() {
let Some(order) = projection_source(k) else {
continue;
};
let Some(oi) = crate::seam::MODEL_ORDERS.iter().position(|&o| o == order) else {
panic!("device model {k} reads order {order}, which the blob format does not carry");
};
let Some(source) = rows.get(oi) else {
panic!("the blob decoded {} orders and has no order {order}", rows.len());
};
jobs.push(ModelJob { k, source, block, stats: None });
}
let plan = flynnel::JobPlan::new(0, jobs.len() as u32).with_leaf_shape(flynnel::LeafShape::Gather);
flynnel::sched::par_iter::for_each_chunk_indexed_min_leaf(&plan, &mut jobs, 1, |_, slice| {
for job in slice {
let mut nodes = projected_nodes(job.source, job.k);
let total_mass: u64 = nodes.iter().map(|n| n.1 + n.2).sum();
let (placed, kept_mass) = place_nodes(&mut nodes, job.block, bits, cap);
job.stats = Some(super::ProjectionStats {
model: job.k,
nodes: nodes.len(),
placed,
mass_kept: if total_mass == 0 { 1.0 } else { kept_mass as f64 / total_mass as f64 },
});
}
});
let stats = jobs.into_iter().map(|job| job.stats.expect("every model's job ran")).collect();
(prior, stats)
}
struct GpuCompress {
ctx: Arc<CudaContext>,
func: CudaFunction,
}
fn gpu_compress() -> Option<&'static GpuCompress> {
static GC: OnceLock<Option<GpuCompress>> = OnceLock::new();
GC.get_or_init(|| load_or_cpu("compress", load_compress_kernel)).as_ref()
}
fn load_compress_kernel() -> Result<GpuCompress, cudarc::driver::DriverError> {
let ctx = CudaContext::new(0)?;
let module = ctx.load_module(Ptx::from_src(COMPRESS_PTX))?;
let func = module.load_function("trex_compress")?;
Ok(GpuCompress { ctx, func })
}
pub(super) fn compress(input: &[u8]) -> Option<f64> {
match held_default_prior() {
Some(held) => compress_held(input, held),
None => compress_with(input, &[0u16; 3], 0, false),
}
}
pub(super) fn compress_with(input: &[u8], prior: &[u16], prior_bits: usize, predict: bool) -> Option<f64> {
let held = hold_prior(prior, prior_bits, predict)?;
compress_held(input, &held)
}
pub(super) struct HeldPrior {
d_prior: cudarc::driver::CudaSlice<u16>,
bits: usize,
predict: bool,
}
pub(super) fn hold_prior(prior: &[u16], prior_bits: usize, predict: bool) -> Option<HeldPrior> {
let want = if prior_bits == 0 { prior.len() } else { NCTX * (1usize << prior_bits) * 3 };
assert!(
!prior.is_empty() && prior.len() == want,
"a prior at {prior_bits} bits holds {want} u16s, not {}",
prior.len()
);
let g = gpu_compress()?;
match g.ctx.default_stream().clone_htod(prior) {
Ok(d_prior) => Some(HeldPrior { d_prior, bits: prior_bits, predict }),
Err(e) => {
eprintln!("trex gpu compress: holding the {}-u16 prior failed: {e:?}", prior.len());
None
}
}
}
fn held_default_prior() -> Option<&'static HeldPrior> {
static HELD: OnceLock<Option<HeldPrior>> = OnceLock::new();
HELD.get_or_init(hold_default_prior).as_ref()
}
pub(super) fn hold_default_prior() -> Option<HeldPrior> {
let (table, bits) = default_prior();
hold_prior(table, bits, bits > 0)
}
pub(super) fn compress_held(input: &[u8], held: &HeldPrior) -> Option<f64> {
let g = gpu_compress()?;
let n = input.len();
if n == 0 {
return Some(0.0);
}
let chunk: usize = env_or("TREX_GPU_CHUNK", (n / 4096).clamp(CHUNK, 65536));
assert!(chunk > 0, "TREX_GPU_CHUNK must be positive");
let nc = n.div_ceil(chunk);
let starts: Vec<i32> = (0..nc).map(|i| (i * chunk) as i32).collect();
let ends: Vec<i32> = (0..nc).map(|i| ((i + 1) * chunk).min(n) as i32).collect();
let mut ctx_bits = (chunk.max(1).ilog2() as usize + 4).clamp(14, 20);
while ctx_bits > 14 && nc * NCTX * (1usize << ctx_bits) * 6 > (2usize << 30) {
ctx_bits -= 1;
}
let nc_i = nc as i32;
let ctx_bits_i = ctx_bits as i32;
let prior_bits_i = held.bits as i32;
let prior_predict_i = i32::from(held.predict);
let overlap: i32 = env_or("TREX_GPU_OVERLAP", (chunk / 2) as i32);
let block = 128u32;
let cfg = LaunchConfig {
grid_dim: (nc.div_ceil(block as usize) as u32, 1, 1),
block_dim: (block, 1, 1),
shared_mem_bytes: 0,
};
let stream = g.ctx.default_stream();
let ran = (|| {
let d_input = stream.clone_htod(input)?;
let d_start = stream.clone_htod(&starts)?;
let d_end = stream.clone_htod(&ends)?;
let d_tables = stream.alloc_zeros::<u16>(nc * NCTX * (1 << ctx_bits) * 3)?;
let d_mtables = stream.alloc_zeros::<u32>(nc * MATCH_SIZE)?;
let d_emb = stream.alloc_zeros::<i32>(nc * NEMB * EMB_SIZE * EMB_DIM)?;
let d_wnode = stream.alloc_zeros::<i32>(nc * NEMB * 256 * EMB_DIM)?;
let d_weights = stream.alloc_zeros::<i32>(nc * 256 * (NCTX + 1 + NEMB + 1))?;
let d_out = stream.alloc_zeros::<f64>(nc)?;
let mut builder = stream.launch_builder(&g.func);
builder.arg(&d_input);
builder.arg(&d_start);
builder.arg(&d_end);
builder.arg(&nc_i);
builder.arg(&held.d_prior);
builder.arg(&prior_bits_i);
builder.arg(&prior_predict_i);
builder.arg(&overlap);
builder.arg(&ctx_bits_i);
builder.arg(&d_tables);
builder.arg(&d_mtables);
builder.arg(&d_emb);
builder.arg(&d_wnode);
builder.arg(&d_weights);
builder.arg(&d_out);
unsafe { builder.launch(cfg)? };
let out: Vec<f64> = stream.clone_dtoh(&d_out)?;
stream.synchronize()?;
Ok::<_, cudarc::driver::DriverError>(out)
})();
match ran {
Ok(out) => Some(out.iter().sum()),
Err(e) => {
eprintln!("trex gpu compress: coding {n} bytes in {nc} chunks failed: {e:?}");
None
}
}
}
fn env_or<T: std::str::FromStr>(name: &str, default: T) -> T
where
T::Err: std::fmt::Display,
{
match std::env::var_os(name) {
None => default,
Some(raw) => match raw.to_str() {
None => panic!("{name} is set but is not valid Unicode"),
Some(s) => match s.parse::<T>() {
Ok(v) => v,
Err(e) => panic!("{name}={s:?} does not parse: {e}"),
},
},
}
}
pub(super) fn compress_phases(input: &[u8]) -> Option<(f64, super::CompressPhases)> {
use std::time::Instant;
let g = gpu_compress()?;
let mut ph = super::CompressPhases::default();
let n = input.len();
if n == 0 {
return Some((0.0, ph));
}
let t = Instant::now();
let chunk: usize = env_or("TREX_GPU_CHUNK", (n / 4096).clamp(CHUNK, 65536));
assert!(chunk > 0, "TREX_GPU_CHUNK must be positive");
let nc = n.div_ceil(chunk);
let starts: Vec<i32> = (0..nc).map(|i| (i * chunk) as i32).collect();
let ends: Vec<i32> = (0..nc).map(|i| ((i + 1) * chunk).min(n) as i32).collect();
let (prior, prior_bits) = default_prior();
ph.host_prep_us = t.elapsed().as_secs_f64() * 1e6;
ph.chunks = nc;
ph.prior_bytes = std::mem::size_of_val(prior);
let stream = g.ctx.default_stream();
let report = |what: &str, e: &cudarc::driver::DriverError| {
eprintln!("trex gpu compress probe: {what} for {n} bytes failed: {e:?}");
};
let t = Instant::now();
let staged = (|| {
let d_input = stream.clone_htod(input)?;
let d_start = stream.clone_htod(&starts)?;
let d_end = stream.clone_htod(&ends)?;
stream.synchronize()?;
Ok::<_, cudarc::driver::DriverError>((d_input, d_start, d_end))
})();
let (d_input, d_start, d_end) = match staged {
Ok(v) => v,
Err(e) => {
report("uploading the input", &e);
return None;
}
};
ph.upload_input_us = t.elapsed().as_secs_f64() * 1e6;
let mut ctx_bits = (chunk.ilog2() as usize + 4).clamp(14, 20);
while ctx_bits > 14 && nc * NCTX * (1usize << ctx_bits) * 6 > (2usize << 30) {
ctx_bits -= 1;
}
let t = Instant::now();
let allocated = (|| {
let d_tables = stream.alloc_zeros::<u16>(nc * NCTX * (1 << ctx_bits) * 3)?;
let d_mtables = stream.alloc_zeros::<u32>(nc * MATCH_SIZE)?;
let d_emb = stream.alloc_zeros::<i32>(nc * NEMB * EMB_SIZE * EMB_DIM)?;
let d_wnode = stream.alloc_zeros::<i32>(nc * NEMB * 256 * EMB_DIM)?;
let d_weights = stream.alloc_zeros::<i32>(nc * 256 * (NCTX + 1 + NEMB + 1))?;
let d_out = stream.alloc_zeros::<f64>(nc)?;
stream.synchronize()?;
Ok::<_, cudarc::driver::DriverError>((d_tables, d_mtables, d_emb, d_wnode, d_weights, d_out))
})();
let (d_tables, d_mtables, d_emb, d_wnode, d_weights, d_out) = match allocated {
Ok(v) => v,
Err(e) => {
report("allocating the per-thread tables", &e);
return None;
}
};
ph.alloc_us = t.elapsed().as_secs_f64() * 1e6;
let t = Instant::now();
let uploaded = (|| {
let d = stream.clone_htod(prior)?;
stream.synchronize()?;
Ok::<_, cudarc::driver::DriverError>(d)
})();
let d_prior = match uploaded {
Ok(d) => d,
Err(e) => {
report("uploading the prior", &e);
return None;
}
};
ph.upload_prior_us = t.elapsed().as_secs_f64() * 1e6;
let nc_i = nc as i32;
let prior_bits_i = prior_bits as i32;
let prior_predict_i = i32::from(prior_bits > 0);
let overlap: i32 = env_or("TREX_GPU_OVERLAP", (chunk / 2) as i32);
let block = 128u32;
let cfg = LaunchConfig {
grid_dim: (nc.div_ceil(block as usize) as u32, 1, 1),
block_dim: (block, 1, 1),
shared_mem_bytes: 0,
};
let ctx_bits_i = ctx_bits as i32;
let t = Instant::now();
let ran = (|| {
let mut builder = stream.launch_builder(&g.func);
builder.arg(&d_input);
builder.arg(&d_start);
builder.arg(&d_end);
builder.arg(&nc_i);
builder.arg(&d_prior);
builder.arg(&prior_bits_i);
builder.arg(&prior_predict_i);
builder.arg(&overlap);
builder.arg(&ctx_bits_i);
builder.arg(&d_tables);
builder.arg(&d_mtables);
builder.arg(&d_emb);
builder.arg(&d_wnode);
builder.arg(&d_weights);
builder.arg(&d_out);
unsafe { builder.launch(cfg)? };
let out: Vec<f64> = stream.clone_dtoh(&d_out)?;
stream.synchronize()?;
Ok::<_, cudarc::driver::DriverError>(out)
})();
let out = match ran {
Ok(o) => o,
Err(e) => {
report("running the kernel", &e);
return None;
}
};
ph.kernel_and_download_us = t.elapsed().as_secs_f64() * 1e6;
Some((out.iter().sum(), ph))
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use super::*;
const SAMPLE: &[u8] = b"The cat sat. the Cat ran 42 times; THE CAT sat again, and 7 cats sat by 42 mats.";
#[test]
fn projected_nodes_are_the_counts_the_kernel_hashing_gives_the_corpus() {
let corpus = SAMPLE.repeat(3);
let rows = crate::seam::byte_ngram_train(&corpus);
let mut projected_models = 0;
for k in 0..NCTX {
let Some(order) = projection_source(k) else {
continue;
};
let oi = crate::seam::MODEL_ORDERS
.iter()
.position(|&o| o == order)
.expect("every source order is one the trainer counts");
let mut want: BTreeMap<u64, (u64, u64)> = BTreeMap::new();
for t in order..corpus.len() {
let base = model_hash(&corpus, t, k);
let mut node = 1u64;
for bit in (0..8).rev() {
let a = (corpus[t] >> bit) & 1;
let e = want.entry(base.wrapping_mul(MIXK).wrapping_add(node)).or_insert((0, 0));
if a == 1 { e.1 += 1 } else { e.0 += 1 }
node = (node << 1) | u64::from(a);
}
}
let mut got: BTreeMap<u64, (u64, u64)> = BTreeMap::new();
for (key, n0, n1) in projected_nodes(&rows[oi], k) {
assert!(got.insert(key, (n0, n1)).is_none(), "model {k} projected key {key:#x} twice");
}
assert!(!want.is_empty(), "model {k} counted nothing on the sample");
assert_eq!(got, want, "model {k} projected counts the corpus does not give it");
projected_models += 1;
}
assert_eq!(projected_models, NCTX - 1, "every model but the word stem projects");
}
#[test]
fn a_heavier_node_takes_the_slot_and_its_counts_halve_to_the_cap() {
let bits = 4;
let mut block = vec![7u16; 16 * 3];
let mut nodes = vec![(0x103u64, 6u64, 4u64), (0x203, 1500, 500), (0x305, 0, 7)];
let (placed, kept) = place_nodes(&mut nodes, &mut block, bits, 1024);
assert_eq!((placed, kept), (2, 2007));
assert_eq!(&block[3 * 3..4 * 3], &[0x20, 750, 250], "the 2000-count node wins slot 3, halved once");
assert_eq!(&block[5 * 3..6 * 3], &[0x30, 0, 7], "a node under the cap keeps its counts");
let untouched: usize = (0..16).filter(|&i| i != 3 && i != 5).map(|i| usize::from(block[i * 3 + 1])).sum();
assert_eq!(untouched, 0, "every other slot is cleared");
}
#[test]
fn the_smallest_cap_halves_a_two_sided_slot_to_one_count_each_way() {
let mut block = vec![0u16; 16 * 3];
let mut nodes = vec![(0x203u64, 1500u64, 500u64), (0x305, 3, 0)];
let (placed, _) = place_nodes(&mut nodes, &mut block, 4, 2);
assert_eq!(placed, 2);
assert_eq!(&block[3 * 3..4 * 3], &[0x20, 1, 1], "halving rounds up, so both sides keep one count");
assert_eq!(&block[5 * 3..6 * 3], &[0x30, 2, 0], "a one-sided slot halves to the cap");
}
}
}