use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use sha2::{Digest, Sha256};
use trueno_rand::Philox4x32;
use crate::autograd::Tensor;
use crate::nn::transformer::AttentionDropoutMasks;
const DOMAIN_TAG: &[u8] = b"apr-setfit-dropout-v1\0";
const TWO_POW_64_F64: f64 = 18_446_744_073_709_551_616.0;
const TWO_POW_64_U128: u128 = 1_u128 << 64;
#[derive(Debug, Clone, PartialEq)]
pub enum DropoutRngError {
RateNotFinite {
observed: f32,
},
RateNegative {
observed: f32,
},
RateAtOrAboveOne {
observed: f32,
},
RateScaleNotFinite {
observed: f32,
scale: f32,
},
ForwardOrdinalOverflow {
observed: u64,
},
}
impl std::fmt::Display for DropoutRngError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::RateNotFinite { observed } => write!(
f,
"DropoutRngError::RateNotFinite(dropout rate {observed} is not finite)"
),
Self::RateNegative { observed } => write!(
f,
"DropoutRngError::RateNegative(dropout rate {observed} is negative)"
),
Self::RateAtOrAboveOne { observed } => write!(
f,
"DropoutRngError::RateAtOrAboveOne(dropout rate {observed} would drop every element)"
),
Self::RateScaleNotFinite { observed, scale } => write!(
f,
"DropoutRngError::RateScaleNotFinite(dropout rate {observed} is below 1.0 but its \
inverted-dropout scale 1/(1-p) is {scale})"
),
Self::ForwardOrdinalOverflow { observed } => write!(
f,
"DropoutRngError::ForwardOrdinalOverflow(forward ordinal {observed} does not fit \
the u32 counter lane; the limit is {} exclusive)",
u32::MAX
),
}
}
}
impl std::error::Error for DropoutRngError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DomainKey([u32; 2]);
impl DomainKey {
#[must_use]
pub fn lanes(self) -> [u32; 2] {
self.0
}
}
#[must_use]
pub fn derive_key(root_seed: u64, site: &str) -> DomainKey {
let mut hasher = Sha256::new();
hasher.update(DOMAIN_TAG);
hasher.update(root_seed.to_le_bytes());
hasher.update(site.as_bytes());
let digest: [u8; 32] = hasher.finalize().into();
let lane0 = u32::from_le_bytes([digest[0], digest[1], digest[2], digest[3]]);
let lane1 = u32::from_le_bytes([digest[4], digest[5], digest[6], digest[7]]);
DomainKey([lane0, lane1])
}
#[must_use]
pub fn draw(key: &DomainKey, forward_ordinal: u32, element: u64) -> [u32; 4] {
let counter = [element as u32, (element >> 32) as u32, forward_ordinal, 0];
Philox4x32::generate_at(key.0, counter)
}
#[must_use]
fn assemble64(lanes: [u32; 4]) -> u64 {
(u64::from(lanes[1]) << 32) | u64::from(lanes[0])
}
#[must_use]
pub fn keep_threshold(p: f64) -> u128 {
if !p.is_finite() || p <= 0.0 {
return 0;
}
let scaled = (p * TWO_POW_64_F64).floor();
let raw = scaled as u128;
raw.min(TWO_POW_64_U128)
}
pub fn validate_rate(p: f32) -> Result<f32, DropoutRngError> {
if !p.is_finite() {
return Err(DropoutRngError::RateNotFinite { observed: p });
}
if p < 0.0 {
return Err(DropoutRngError::RateNegative { observed: p });
}
if p >= 1.0 {
return Err(DropoutRngError::RateAtOrAboveOne { observed: p });
}
let scale = 1.0 / (1.0 - p);
if !scale.is_finite() {
return Err(DropoutRngError::RateScaleNotFinite { observed: p, scale });
}
Ok(scale)
}
pub fn forward_ordinal(step: u64, branch: u32) -> Result<u32, DropoutRngError> {
let doubled = step
.checked_mul(2)
.and_then(|s| s.checked_add(u64::from(branch)));
match doubled {
Some(ordinal) => checked_forward_ordinal(ordinal),
None => Err(DropoutRngError::ForwardOrdinalOverflow { observed: u64::MAX }),
}
}
pub fn checked_forward_ordinal(ordinal: u64) -> Result<u32, DropoutRngError> {
match u32::try_from(ordinal) {
Ok(narrowed) if narrowed < u32::MAX => Ok(narrowed),
_ => Err(DropoutRngError::ForwardOrdinalOverflow { observed: ordinal }),
}
}
#[derive(Debug)]
pub struct SiteDropout {
site: String,
key: DomainKey,
p: f32,
scale: f32,
threshold: u128,
training: AtomicBool,
forward_ordinal: AtomicU32,
}
impl SiteDropout {
pub fn new(root_seed: u64, site: &str, p: f32) -> Result<Self, DropoutRngError> {
let scale = validate_rate(p)?;
Ok(Self {
site: site.to_string(),
key: derive_key(root_seed, site),
p,
scale,
threshold: keep_threshold(f64::from(p)),
training: AtomicBool::new(true),
forward_ordinal: AtomicU32::new(0),
})
}
#[must_use]
pub fn site(&self) -> &str {
&self.site
}
#[must_use]
pub fn key(&self) -> DomainKey {
self.key
}
#[must_use]
pub fn probability(&self) -> f32 {
self.p
}
#[must_use]
pub fn scale(&self) -> f32 {
self.scale
}
#[must_use]
pub fn threshold(&self) -> u128 {
self.threshold
}
#[must_use]
pub fn training(&self) -> bool {
self.training.load(Ordering::Relaxed)
}
pub fn set_training(&self, training: bool) {
self.training.store(training, Ordering::Relaxed);
}
#[must_use]
pub fn current_forward_ordinal(&self) -> u32 {
self.forward_ordinal.load(Ordering::Relaxed)
}
pub fn set_forward_ordinal(&self, ordinal: u64) -> Result<(), DropoutRngError> {
let narrowed = checked_forward_ordinal(ordinal)?;
self.forward_ordinal.store(narrowed, Ordering::Relaxed);
Ok(())
}
#[must_use]
pub fn mask_element(&self, forward_ordinal: u32, i: u64) -> f32 {
let x = assemble64(draw(&self.key, forward_ordinal, i));
if u128::from(x) >= self.threshold {
self.scale
} else {
0.0
}
}
#[must_use]
pub fn mask_at(&self, forward_ordinal: u32, len: usize) -> Vec<f32> {
(0..len)
.map(|i| self.mask_element(forward_ordinal, i as u64))
.collect()
}
#[must_use]
pub fn mask(&self, len: usize) -> Vec<f32> {
self.mask_at(self.current_forward_ordinal(), len)
}
#[must_use]
pub fn forward(&self, input: &Tensor) -> Tensor {
if !self.training() || self.p == 0.0 {
return input.clone();
}
let mask_data = self.mask(input.data().len());
let mask = Tensor::from_vec(mask_data, input.shape());
input.mul(&mask)
}
}
impl AttentionDropoutMasks for SiteDropout {
fn attention_dropout_mask(&self, len: usize) -> Vec<f32> {
self.mask(len)
}
}
#[cfg(test)]
mod dropout_rng_tests {
use super::{
assemble64, checked_forward_ordinal, derive_key, draw, forward_ordinal, keep_threshold,
validate_rate, DropoutRngError, SiteDropout,
};
use crate::autograd::Tensor;
const GOLDEN_SITE: &str = "embeddings.dropout";
const GOLDEN_SEED: u64 = 13;
const KEY_13_EMBEDDINGS: [u32; 2] = [0xa5ba_f332, 0x7e65_61f9];
const KEY_13_LAYER0_ATTN: [u32; 2] = [0xd0cf_7225, 0x1ade_02c9];
const KEY_14_EMBEDDINGS: [u32; 2] = [0x3a40_a784, 0x1322_f43f];
const BLOCK_ORD0_ELEM0: [u32; 4] = [3_836_206_948, 4_227_518_855, 1_470_901_809, 3_372_378_841];
const ASSEMBLED_ORD0_ELEM0: u64 = 18_157_055_229_284_573_028;
const BLOCK_ORD7_ELEM3: [u32; 4] = [2_242_403_448, 2_398_921_103, 4_169_238_919, 3_718_170_778];
const BLOCK_ORD1_HIGH_ELEM: [u32; 4] =
[1_286_292_542, 2_418_118_001, 4_129_126_676, 3_075_353_215];
const THRESHOLD_P0: u128 = 0;
const THRESHOLD_P0_1_F64: u128 = 1_844_674_407_370_955_264;
const THRESHOLD_P0_5: u128 = 9_223_372_036_854_775_808;
const THRESHOLD_DROPOUT_P: u128 = 1_844_674_434_858_745_856;
const THRESHOLD_P1: u128 = 1_u128 << 64;
const THRESHOLD_P1E_5: u128 = 184_467_440_737_095;
fn site(p: f32) -> SiteDropout {
SiteDropout::new(GOLDEN_SEED, GOLDEN_SITE, p).expect("test rates are valid")
}
#[test]
fn dropout_rng_accessors_report_the_constructed_values() {
let s = site(0.1);
assert_eq!(s.site(), GOLDEN_SITE);
assert_eq!(s.probability(), 0.1);
assert_eq!(s.scale(), validate_rate(0.1).expect("0.1 is a valid rate"));
assert_eq!(s.threshold(), THRESHOLD_DROPOUT_P);
assert_eq!(s.threshold(), keep_threshold(f64::from(0.1_f32)));
s.set_forward_ordinal(7).expect("7 fits u32");
assert_eq!(s.current_forward_ordinal(), 7);
}
#[test]
fn dropout_rng_attention_mask_trait_forwarder_reaches_the_real_mask() {
use crate::nn::transformer::AttentionDropoutMasks;
let s = site(0.5);
s.set_forward_ordinal(3).expect("3 fits u32");
let via_trait: &dyn AttentionDropoutMasks = &s;
let len = 64_usize;
let mask = via_trait.attention_dropout_mask(len);
assert_eq!(mask.len(), len, "a fixed 0- or 1-element vec is not a mask");
assert_eq!(
mask,
s.mask(len),
"the forwarder must return THE mask, not a lookalike"
);
assert!(
mask.contains(&0.0),
"no element dropped at p=0.5 across 64 draws"
);
assert!(mask.contains(&2.0), "no element kept-and-scaled at p=0.5");
}
#[test]
fn dropout_rng_byte_encoding_golden_is_frozen() {
assert_eq!(
derive_key(GOLDEN_SEED, GOLDEN_SITE).lanes(),
KEY_13_EMBEDDINGS
);
assert_eq!(
derive_key(GOLDEN_SEED, "encoder.layer.0.attention.self.dropout").lanes(),
KEY_13_LAYER0_ATTN,
"site separation"
);
assert_eq!(
derive_key(14, GOLDEN_SITE).lanes(),
KEY_14_EMBEDDINGS,
"seed separation"
);
let key = derive_key(GOLDEN_SEED, GOLDEN_SITE);
assert_eq!(draw(&key, 0, 0), BLOCK_ORD0_ELEM0, "counter layout");
assert_eq!(
draw(&key, 7, 3),
BLOCK_ORD7_ELEM3,
"ordinal + element lanes"
);
assert_eq!(
draw(&key, 1, 12_345_678_901),
BLOCK_ORD1_HIGH_ELEM,
"the HIGH element word must reach counter lane 1"
);
assert_eq!(assemble64(BLOCK_ORD0_ELEM0), ASSEMBLED_ORD0_ELEM0);
}
#[test]
fn dropout_rng_tag_is_this_modules_own_and_not_phase_twos() {
let ours = derive_key(GOLDEN_SEED, GOLDEN_SITE).lanes();
let phase_two_tag_would_give = {
use sha2::{Digest, Sha256};
let mut h = Sha256::new();
h.update(b"apr-contrastive-v1\0");
h.update(GOLDEN_SEED.to_le_bytes());
h.update(GOLDEN_SITE.as_bytes());
let d: [u8; 32] = h.finalize().into();
[
u32::from_le_bytes([d[0], d[1], d[2], d[3]]),
u32::from_le_bytes([d[4], d[5], d[6], d[7]]),
]
};
assert_ne!(
ours, phase_two_tag_would_give,
"this module derived the SAME key Phase 2's tag would — the domain tag \
is not separating the phases"
);
}
#[test]
fn dropout_rng_threshold_goldens_are_frozen() {
assert_eq!(keep_threshold(0.0), THRESHOLD_P0);
assert_eq!(keep_threshold(0.1), THRESHOLD_P0_1_F64);
assert_eq!(keep_threshold(0.5), THRESHOLD_P0_5);
assert_eq!(keep_threshold(f64::from(0.1_f32)), THRESHOLD_DROPOUT_P);
assert_ne!(
THRESHOLD_P0_1_F64, THRESHOLD_DROPOUT_P,
"the f32 and f64 spellings of 0.1 must NOT collapse to one threshold — \
if they do, the widening step was silently dropped"
);
assert_eq!(keep_threshold(1e-5), THRESHOLD_P1E_5);
assert_ne!(
keep_threshold(1e-5),
THRESHOLD_P1E_5 + 1,
"the threshold is rounding AWAY from zero"
);
assert_eq!(keep_threshold(1.0), THRESHOLD_P1);
assert!(THRESHOLD_P1 > u128::from(u64::MAX));
assert_eq!(keep_threshold(f64::NAN), 0);
assert_eq!(keep_threshold(-0.5), 0);
assert_eq!(keep_threshold(f64::INFINITY), 0);
}
#[test]
fn dropout_rng_threshold_is_the_drop_boundary_not_the_keep_boundary() {
let s = SiteDropout::new(GOLDEN_SEED, GOLDEN_SITE, 0.9).expect("0.9 is valid");
let mask = s.mask_at(0, 20_000);
let dropped = mask.iter().filter(|v| **v == 0.0).count();
assert!(
(17_400..=18_600).contains(&dropped),
"dropped {dropped} of 20000 at p = 0.9; expected ~18000, so the keep \
comparison is inverted or the threshold is wrong"
);
}
#[test]
fn dropout_rng_mask_element_is_a_pure_function_of_its_index() {
let s = site(0.3);
for (ordinal, len) in [(0_u32, 64_usize), (1, 64), (7, 256), (4_242, 33)] {
let sequential = s.mask_at(ordinal, len);
for i in [0_usize, 1, 5, len / 2, len - 1] {
assert_eq!(
s.mask_element(ordinal, i as u64),
sequential[i],
"element {i} at ordinal {ordinal} differs when computed directly"
);
}
let shuffled: Vec<f32> = [len - 1, 0, len / 2]
.iter()
.map(|i| s.mask_element(ordinal, *i as u64))
.collect();
assert_eq!(
shuffled,
vec![sequential[len - 1], sequential[0], sequential[len / 2]]
);
}
}
#[test]
fn dropout_rng_replay_is_bitwise_and_every_coordinate_separates() {
let base = SiteDropout::new(7, "embeddings.dropout", 0.25).expect("valid");
let a = base.mask_at(4, 512);
let b = base.mask_at(4, 512);
for (i, (x, y)) in a.iter().zip(b.iter()).enumerate() {
assert_eq!(x.to_bits(), y.to_bits(), "element {i}: replay is not exact");
}
let twin = SiteDropout::new(7, "embeddings.dropout", 0.25).expect("valid");
assert_eq!(a, twin.mask_at(4, 512));
let other_seed = SiteDropout::new(8, "embeddings.dropout", 0.25).expect("valid");
let other_site =
SiteDropout::new(7, "encoder.layer.0.output.dropout", 0.25).expect("valid");
assert_ne!(a, other_seed.mask_at(4, 512), "root seed does not separate");
assert_ne!(a, other_site.mask_at(4, 512), "site does not separate");
assert_ne!(a, base.mask_at(5, 512), "forward ordinal does not separate");
}
#[test]
fn dropout_rng_branches_of_one_step_draw_independent_masks() {
const LEN: usize = 512;
for (p, step) in [(0.1_f32, 3_u64), (0.5, 3), (0.1, 0), (0.5, 17)] {
let s = SiteDropout::new(GOLDEN_SEED, GOLDEN_SITE, p).expect("valid rate");
let branch_a = forward_ordinal(step, 0).expect("2*step fits");
let branch_b = forward_ordinal(step, 1).expect("2*step+1 fits");
assert_eq!(u64::from(branch_a), 2 * step, "branch A is 2*step");
assert_eq!(u64::from(branch_b), 2 * step + 1, "branch B is 2*step + 1");
let mask_a = s.mask_at(branch_a, LEN);
let mask_b = s.mask_at(branch_b, LEN);
let hamming = mask_a
.iter()
.zip(mask_b.iter())
.filter(|(x, y)| x.to_bits() != y.to_bits())
.count();
let q = 2.0 * f64::from(p) * (1.0 - f64::from(p));
let n = LEN as f64;
let mean = n * q;
let sd = (n * q * (1.0 - q)).sqrt();
let lo = (mean - 4.0 * sd).max(1.0);
let hi = mean + 4.0 * sd;
assert!(
(lo..=hi).contains(&(hamming as f64)),
"p = {p}, step = {step}: Hamming distance {hamming} is outside \
[{lo:.1}, {hi:.1}] (mean {mean:.1}, sd {sd:.2}). A distance of 0 \
means the two siamese branches collapsed onto ONE stream and are \
sharing a mask — D-15's whole point"
);
}
}
#[test]
fn dropout_rng_forward_ordinal_is_monotone_in_step_and_branch() {
let mut previous = None;
for step in 0_u64..8 {
for branch in 0_u32..2 {
let ordinal = forward_ordinal(step, branch).expect("small ordinals fit");
if let Some(prev) = previous {
assert!(ordinal > prev, "2*{step}+{branch} is not monotone");
}
previous = Some(ordinal);
}
}
}
#[test]
fn dropout_rng_eval_mode_is_the_identity_and_consumes_nothing() {
let s = site(0.5);
let x = Tensor::new(&[0.25_f32, -1.5, 3.0, 7.75], &[4]);
s.set_training(false);
let y = s.forward(&x);
for (i, (a, b)) in x.data().iter().zip(y.data().iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"element {i}: eval is not identity"
);
}
assert_eq!(s.current_forward_ordinal(), 0);
s.set_training(true);
let z = s.forward(&x);
assert!(
x.data()
.iter()
.zip(z.data().iter())
.any(|(a, b)| a.to_bits() != b.to_bits()),
"train mode was also the identity — no element was dropped or scaled"
);
}
#[test]
fn dropout_rng_rate_zero_keeps_everything_unscaled() {
let s = site(0.0);
assert_eq!(s.scale(), 1.0);
assert_eq!(s.threshold(), 0);
assert!(
s.mask_at(0, 1_000).iter().all(|v| *v == 1.0),
"p = 0 dropped or rescaled an element"
);
let x = Tensor::new(&[1.0_f32, 2.0, 3.0], &[3]);
let y = s.forward(&x);
assert_eq!(x.data(), y.data());
}
#[test]
fn dropout_rng_rate_one_is_a_typed_error_naming_the_value() {
let err = validate_rate(1.0).expect_err("p = 1.0 must be rejected");
assert_eq!(err, DropoutRngError::RateAtOrAboveOne { observed: 1.0 });
assert!(err.to_string().contains('1'), "got {err}");
assert!(SiteDropout::new(1, "s", 1.0).is_err());
}
#[test]
fn dropout_rng_rate_just_below_one_is_a_typed_error_naming_the_value() {
let p: f32 = 1.0 - 1e-40;
assert_eq!(p.to_bits(), 1.0_f32.to_bits(), "1.0 - 1e-40 is exactly 1.0");
let err = validate_rate(p).expect_err("must be rejected");
assert_eq!(err, DropoutRngError::RateAtOrAboveOne { observed: p });
let text = err.to_string();
assert!(
text.contains('1'),
"the message must name the value: {text}"
);
}
#[test]
fn dropout_rng_rate_scale_guard_is_unreachable_for_f32() {
let worst = f32::from_bits(0x3F7F_FFFF);
assert!(worst < 1.0);
assert_eq!(1.0 - worst, 2.0_f32.powi(-24));
let scale = validate_rate(worst).expect("the worst case is still valid");
assert_eq!(
scale, 16_777_216.0,
"1/(1-p) at the boundary is exactly 2^24"
);
assert!(scale.is_finite());
for bits in [0x3F7F_FFFEu32, 0x3F7F_FFF0, 0x3F7F_FF00, 0x3F00_0000] {
let p = f32::from_bits(bits);
let s = validate_rate(p).expect("still below one");
assert!(s.is_finite() && s <= 16_777_216.0, "p = {p} gave scale {s}");
}
}
#[test]
fn dropout_rng_rate_rejects_nan_infinity_and_negatives_by_name() {
match validate_rate(f32::NAN).expect_err("NaN") {
DropoutRngError::RateNotFinite { observed } => assert!(observed.is_nan()),
other => panic!("NaN must be RateNotFinite, got {other}"),
}
assert_eq!(
validate_rate(f32::INFINITY).expect_err("inf"),
DropoutRngError::RateNotFinite {
observed: f32::INFINITY
}
);
let err = validate_rate(-0.25).expect_err("negative");
assert_eq!(err, DropoutRngError::RateNegative { observed: -0.25 });
assert!(err.to_string().contains("-0.25"), "got {err}");
}
#[test]
fn dropout_rng_forward_ordinal_overflow_is_a_typed_error_naming_the_value() {
let limit = u64::from(u32::MAX);
assert!(checked_forward_ordinal(limit - 1).is_ok());
for observed in [limit, limit + 1, u64::MAX] {
let err = checked_forward_ordinal(observed).expect_err("must be rejected");
assert_eq!(err, DropoutRngError::ForwardOrdinalOverflow { observed });
assert!(
err.to_string().contains(&observed.to_string()),
"the message must name the observed ordinal: {err}"
);
}
let s = site(0.1);
assert!(s.set_forward_ordinal(limit).is_err());
assert_eq!(
s.current_forward_ordinal(),
0,
"a rejected ordinal must leave the site where it was"
);
assert!(
forward_ordinal(u64::MAX / 2, 1).is_err(),
"2*step must be checked before it can wrap into a SMALL ordinal"
);
}
}