use cubecl_common::config::RuntimeConfig;
use cubecl_common::stub::Arc;
use super::autodiff::AutodiffConfig;
use super::fusion::FusionConfig;
use super::remote::RemoteConfig;
static BURN_GLOBAL_CONFIG: spin::Mutex<Option<Arc<BurnConfig>>> = spin::Mutex::new(None);
#[derive(Default, Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct BurnConfig {
#[serde(default)]
fusion: FusionConfig,
#[serde(default)]
autodiff: AutodiffConfig,
#[serde(default)]
remote: RemoteConfig,
}
impl BurnConfig {
pub fn fusion(&self) -> &FusionConfig {
&self.fusion
}
pub fn autodiff(&self) -> &AutodiffConfig {
&self.autodiff
}
pub fn remote(&self) -> &RemoteConfig {
&self.remote
}
}
impl RuntimeConfig for BurnConfig {
fn storage() -> &'static spin::Mutex<Option<Arc<Self>>> {
&BURN_GLOBAL_CONFIG
}
fn file_names() -> &'static [&'static str] {
&["burn.toml", "Burn.toml"]
}
#[cfg(all(
feature = "std",
any(
target_os = "windows",
target_os = "linux",
target_os = "macos",
target_os = "android"
)
))]
fn override_from_env(mut self) -> Self {
use super::fusion::FusionLogLevel;
use super::remote::RemoteLogLevel;
if let Ok(val) = std::env::var("BURN_FUSION_LOG") {
let level = match val.to_ascii_lowercase().as_str() {
"disabled" | "off" | "0" => FusionLogLevel::Disabled,
"basic" => FusionLogLevel::Basic,
"medium" => FusionLogLevel::Medium,
"full" | "1" => FusionLogLevel::Full,
_ => self.fusion.logger.level,
};
self.fusion.logger.level = level;
if level != FusionLogLevel::Disabled {
self.fusion.logger.stderr = true;
}
}
if let Ok(val) = std::env::var("BURN_FUSION_MAX_EXPLORATIONS")
&& let Ok(n) = val.parse::<usize>()
{
self.fusion.beam_search.max_explorations = Some(n);
}
if let Ok(val) = std::env::var("BURN_REMOTE_LOG") {
let level = match val.to_ascii_lowercase().as_str() {
"disabled" | "off" | "0" => RemoteLogLevel::Disabled,
"basic" | "1" => RemoteLogLevel::Basic,
"full" | "2" => RemoteLogLevel::Full,
_ => self.remote.logger.level,
};
self.remote.logger.level = level;
if level != RemoteLogLevel::Disabled {
self.remote.logger.stderr = true;
}
}
self
}
}