///| Exact single-path speculative verification. A rejected draft token is
///| replaced by a sample from max(target - draft, 0), preserving the target
///|
/// distribution when draft and target distributions are valid.
pub enum VerifyError {
TargetLengthMismatch
RandomLengthMismatch
InvalidTargetDistribution(Int)
ProbabilityFailure(Int)
InvalidProposal
} derive(Eq, Debug)
///|
pub struct VerifyResult {
emitted : Array[Int]
accepted_count : Int
rejected_at : Int?
}
///|
pub fn VerifyResult::accepted_all(
self : VerifyResult,
proposal : DraftProposal,
) -> Bool {
self.accepted_count == proposal.length()
}
///|
fn valid_distribution(values : Array[Double]) -> Bool {
distribution_is_valid(values)
}
///|
/// Validate all rows, including those after an early rejection.
fn validate_verification_inputs(
proposal : DraftProposal,
targets : Array[Array[Double]],
uniforms : Array[Double],
) -> Result[Unit, VerifyError] {
match validate_proposal(proposal) {
Err(_) => return Err(InvalidProposal)
Ok(_) => ()
}
if targets.length() != proposal.length() {
return Err(TargetLengthMismatch)
}
if uniforms.length() != proposal.length() {
return Err(RandomLengthMismatch)
}
for i in 0.. Result[VerifyResult, VerifyError] {
let length = proposal.length()
match
validate_verification_inputs(
proposal, target_distributions, accept_uniforms,
) {
Err(error) => return Err(error)
Ok(_) => ()
}
if target_distributions.length() != length {
return Err(TargetLengthMismatch)
}
if accept_uniforms.length() != length || fallback_uniforms.length() != length {
return Err(RandomLengthMismatch)
}
for i in 0.. value
Err(_) => return Err(ProbabilityFailure(index))
}
if accept_uniforms[index] < acceptance {
emitted.push(draft.token)
accepted = accepted + 1
} else {
let residual = match residual_distribution(target, draft.distribution) {
Ok(value) => value
Err(_) => return Err(ProbabilityFailure(index))
}
let replacement = match
sample_categorical(residual, fallback_uniforms[index]) {
Ok(value) => value
Err(_) => return Err(ProbabilityFailure(index))
}
emitted.push(replacement)
return Ok({ emitted, accepted_count: accepted, rejected_at: Some(index) })
}
}
Ok({ emitted, accepted_count: accepted, rejected_at: None })
}