use cubecl_common::config::logger::{LogLevel, LoggerConfig};
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct FusionConfig {
#[serde(default)]
pub logger: LoggerConfig<FusionLogLevel>,
#[serde(default)]
pub beam_search: BeamSearchConfig,
#[serde(default)]
pub max_graph_size: Option<usize>,
#[serde(default = "default_growth_patience")]
pub growth_patience: usize,
}
impl Default for FusionConfig {
fn default() -> Self {
Self {
logger: LoggerConfig::default(),
beam_search: BeamSearchConfig::default(),
max_graph_size: None,
growth_patience: default_growth_patience(),
}
}
}
fn default_growth_patience() -> usize {
32
}
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct BeamSearchConfig {
#[serde(default = "default_max_blocks")]
pub max_blocks: usize,
#[serde(default)]
pub max_explorations: Option<usize>,
}
impl Default for BeamSearchConfig {
fn default() -> Self {
Self {
max_blocks: default_max_blocks(),
max_explorations: None,
}
}
}
fn default_max_blocks() -> usize {
5
}
#[derive(
Default,
Clone,
Copy,
Debug,
PartialEq,
Eq,
PartialOrd,
Ord,
serde::Serialize,
serde::Deserialize,
)]
pub enum FusionLogLevel {
#[default]
#[serde(rename = "disabled")]
Disabled,
#[serde(rename = "basic")]
Basic,
#[serde(rename = "medium")]
Medium,
#[serde(rename = "full")]
Full,
}
impl LogLevel for FusionLogLevel {}