📄 contexts.rs
/home/palash/git/iron_learn/src/examples/contexts.rs
Language: rs • Lines: 138
//! # Application Context Module
//!
//! Provides global application state management and capability detection.
//!
//! Manages the single global application context containing:
//! - Application metadata (name, version)
//! - Training hyperparameters (learning rate, epochs)
//! - Data configuration (path to data file)
//! - GPU capability flags and CUDA context handle
//!
//! The context is initialized once at application startup and remains immutable
//! throughout the program lifetime using `OnceLock` for thread-safe access.

use std::sync::OnceLock;

use crate::nn::DistributionType;

/// Global singleton instance of application context
///
/// Thread-safe access to the immutable application state initialized at startup.
/// Use `GLOBAL_CONTEXT.get()` to access the context after initialization.
pub static GLOBAL_CONTEXT: OnceLock<AppContext> = OnceLock::new();
use crate::examples::types::ExampleMode;
use crate::examples::types::IronLearnArgs;

/// Global application context with training configuration and GPU capabilities
///
/// # Fields
///
/// * `app_name` - Application name/title
/// * `version` - Application version number
/// * `data_path` - Path to JSON data file
/// * `learning_rate` - Gradient descent learning rate
/// * `epochs` - Number of training iterations
/// * `gpu_enabled` - Flag indicating whether GPU acceleration is available
/// * `context` - CUDA context handle (None if GPU is not available)
/// * `lr_adjust` - Flag indicating whether learning rate adjustment is enabled
/// * `hidden_layer_length` - Number of hidden layers in the neural network
/// * `weights_path` - Path to model weights file, the parameters file
/// * `monitor_interval` - The interval on how many epochs, the monitor should be called for internal state monitoring
/// * `sleep_time` - Sometimes the workload is high and CPU/GPU gets exhausted and machine generates a lot of heat, you may choose to let it cool down
/// * `name` - Name of the model. A similarly named directory must exist from execution path
/// * `restore` - Restore a network from model file.
/// * `distribution` - Weight initialization distribution
///

#[derive(Debug)]
pub struct AppContext {
    pub app_name: &'static str,
    pub version: u32,
    pub data_path: String,
    pub learning_rate: f64,
    pub epochs: u32,
    pub gpu_enabled: bool,
    pub lr_adjust: bool,
    pub hidden_layer_length: u32,
    pub monitor_interval: usize,
    pub sleep_time: u64,
    pub name: String,
    pub restore: bool,
    pub distribution: DistributionType,
    pub example_mode: ExampleMode,
    pub predict_only: bool,
    pub resize: u32,
    pub temparature: f64,
    pub repeat: bool,
    pub n_gram_seed: String,
    pub n_gram_size: u8,
    pub weights_path: String,
}

/// Initialize the global application context
///
/// Must be called exactly once at application startup. Subsequent calls will fail silently.
/// This function captures all training configuration and GPU state for access throughout
/// the application lifetime.
///
/// # Arguments
///
/// * `app_name` - Application name/title
/// * `version` - Application version number
/// * `data_path` - Path to JSON data file
/// * `learning_rate` - Gradient descent learning rate
/// * `epochs` - Number of training iterations
/// * `gpu_enabled` - Flag indicating whether GPU acceleration is available
/// * `context` - CUDA context handle (None if GPU is not available)
/// * `lr_adjust` - Flag indicating whether learning rate adjustment is enabled
/// * `hidden_layer_length` - Number of hidden layers in the neural network
/// * `weights_path` - Path to model weights file, the parameters file
/// * `monitor_interval` - The interval on how many epochs, the monitor should be called for internal state monitoring
/// * `sleep_time` - Sometimes the workload is high and CPU/GPU gets exhausted and machine generates a lot of heat, you may choose to let it cool down
/// * `name` - Name of the model. A similarly named directory must exist from execution path
/// * `restore` - Restore a network from model file.
/// * `distribution` - Weight initialization distribution
///
pub fn init_context(app_name: &'static str, version: u32, gpu_enabled: bool, args: IronLearnArgs) {
    let distribution = match args.distribution.as_str().to_uppercase().as_str() {
        "NORMAL" => DistributionType::Normal,
        "XAVIER" => DistributionType::Xavier,
        "UNIFORM" => DistributionType::Uniform,
        "HE" => DistributionType::He,
        _ => DistributionType::Normal,
    };

    let restore = if args.predict_only {
        true
    } else {
        args.restore
    };

    let ctx = AppContext {
        app_name,
        version,
        gpu_enabled,
        distribution,
        data_path: args.data_file,
        learning_rate: args.lr,
        epochs: args.epochs,
        lr_adjust: args.adjust_lr,
        hidden_layer_length: args.internal_layers,
        monitor_interval: args.monitor_interval,
        sleep_time: args.sleep_time,
        name: args.name.clone(),
        restore,
        example_mode: args.mode,
        predict_only: args.predict_only,
        resize: args.reproduce,
        temparature: args.temparature,
        repeat: args.repeat,
        n_gram_seed: args.n_gram_seed,
        n_gram_size: args.n_gram_size,
        weights_path: "model_outputs/".to_owned() + &args.name.to_owned() + "/model.json",
    };
    match GLOBAL_CONTEXT.set(ctx) {
        Ok(_) => (),
        Err(_) => eprintln!("AppContext has already been initialized!"),
    }
}