mod rules;
mod support;
pub use rules::association_rules;
use std::collections::BTreeSet;
use support::count_to_f64;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AssociationError {
EmptyInput,
NoItems,
RaggedRows,
InvalidSupport,
InvalidThreshold,
}
impl std::fmt::Display for AssociationError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::EmptyInput => write!(f, "transaction matrix has no rows"),
Self::NoItems => write!(f, "transaction matrix has no item columns"),
Self::RaggedRows => write!(f, "transaction matrix rows have differing lengths"),
Self::InvalidSupport => write!(f, "minimum support must lie in (0, 1]"),
Self::InvalidThreshold => write!(f, "minimum metric threshold must be finite"),
}
}
}
impl std::error::Error for AssociationError {}
#[derive(Debug, Clone, PartialEq)]
pub struct FrequentItemset {
items: Vec<usize>,
support: f64,
}
impl FrequentItemset {
#[must_use]
pub fn items(&self) -> &[usize] {
&self.items
}
#[must_use]
pub const fn support(&self) -> f64 {
self.support
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct AssociationRule {
antecedent: Vec<usize>,
consequent: Vec<usize>,
support: f64,
confidence: f64,
lift: f64,
leverage: f64,
conviction: f64,
}
impl AssociationRule {
#[must_use]
pub fn antecedent(&self) -> &[usize] {
&self.antecedent
}
#[must_use]
pub fn consequent(&self) -> &[usize] {
&self.consequent
}
#[must_use]
pub const fn support(&self) -> f64 {
self.support
}
#[must_use]
pub const fn confidence(&self) -> f64 {
self.confidence
}
#[must_use]
pub const fn lift(&self) -> f64 {
self.lift
}
#[must_use]
pub const fn leverage(&self) -> f64 {
self.leverage
}
#[must_use]
pub const fn conviction(&self) -> f64 {
self.conviction
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RuleMetric {
Confidence,
Lift,
Support,
Leverage,
Conviction,
}
fn validate_matrix(transactions: &[Vec<bool>]) -> Result<(usize, usize), AssociationError> {
let n_transactions = transactions.len();
if n_transactions == 0 {
return Err(AssociationError::EmptyInput);
}
let n_items = transactions.first().map_or(0, Vec::len);
if n_items == 0 {
return Err(AssociationError::NoItems);
}
if transactions.iter().any(|row| row.len() != n_items) {
return Err(AssociationError::RaggedRows);
}
Ok((n_transactions, n_items))
}
fn cover_count(transactions: &[Vec<bool>], items: &[usize]) -> usize {
transactions
.iter()
.filter(|row| items.iter().all(|&i| row.get(i).copied().unwrap_or(false)))
.count()
}
pub fn apriori(
transactions: &[Vec<bool>],
min_support: f64,
) -> Result<Vec<FrequentItemset>, AssociationError> {
let (n_transactions, n_items) = validate_matrix(transactions)?;
if !(min_support > 0.0 && min_support <= 1.0) {
return Err(AssociationError::InvalidSupport);
}
let n = count_to_f64(n_transactions);
let support_of = |count: usize| count_to_f64(count) / n;
let mut all: Vec<FrequentItemset> = Vec::new();
let mut current: Vec<Vec<usize>> = Vec::new();
for item in 0..n_items {
let count = cover_count(transactions, &[item]);
let support = support_of(count);
if support >= min_support {
current.push(vec![item]);
all.push(FrequentItemset {
items: vec![item],
support,
});
}
}
while !current.is_empty() {
let candidates = join_and_prune(¤t);
let mut next: Vec<Vec<usize>> = Vec::new();
for cand in candidates {
let count = cover_count(transactions, &cand);
let support = support_of(count);
if support >= min_support {
next.push(cand.clone());
all.push(FrequentItemset {
items: cand,
support,
});
}
}
current = next;
}
sort_itemsets(&mut all);
Ok(all)
}
fn join_and_prune(frequent: &[Vec<usize>]) -> Vec<Vec<usize>> {
let frequent_set: BTreeSet<&Vec<usize>> = frequent.iter().collect();
let mut candidates: BTreeSet<Vec<usize>> = BTreeSet::new();
for (ai, a) in frequent.iter().enumerate() {
for b in frequent.iter().skip(ai + 1) {
let k = a.len();
if a.get(..k - 1) != b.get(..k - 1) {
continue;
}
let (Some(&last_a), Some(&last_b)) = (a.last(), b.last()) else {
continue;
};
if last_a >= last_b {
continue;
}
let mut candidate = a.clone();
candidate.push(last_b);
if all_subsets_frequent(&candidate, &frequent_set) {
candidates.insert(candidate);
}
}
}
candidates.into_iter().collect()
}
fn all_subsets_frequent(candidate: &[usize], frequent: &BTreeSet<&Vec<usize>>) -> bool {
(0..candidate.len()).all(|drop| {
let subset: Vec<usize> = candidate
.iter()
.enumerate()
.filter_map(|(idx, &item)| (idx != drop).then_some(item))
.collect();
frequent.contains(&subset)
})
}
fn sort_itemsets(itemsets: &mut [FrequentItemset]) {
itemsets.sort_by(|a, b| {
a.items
.len()
.cmp(&b.items.len())
.then(a.items.cmp(&b.items))
});
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;