use crate::common::spsc::RtStatusFlags;
use log::{debug, info};
use serde::{Deserialize, Serialize};
use std::sync::atomic::Ordering;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[repr(u8)]
pub enum SlimOverride {
#[default]
Auto = 0,
ForceFull = 1,
ForceLite = 2,
}
impl SlimOverride {
pub fn from_f32(val: f32) -> Self {
match val.round() as i32 {
1 => Self::ForceFull,
2 => Self::ForceLite,
_ => Self::Auto,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[repr(u8)]
pub enum AdaptiveComputeMode {
#[default]
Off = 0,
Conservative = 1,
Aggressive = 2,
}
impl AdaptiveComputeMode {
pub fn from_f32(val: f32) -> Self {
match val.round() as i32 {
1 => Self::Conservative,
2 => Self::Aggressive,
_ => Self::Off,
}
}
pub fn to_f32(self) -> f32 {
self as u8 as f32
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum AdaptiveState {
Full = 0,
Reduced = 1,
Minimal = 2,
}
#[derive(Debug, Clone, Copy, PartialEq)]
enum CrossfadePhase {
Idle,
Active,
}
pub struct AdaptiveCompute {
state: AdaptiveState,
prev_state: AdaptiveState,
overload_counter: u32,
recovery_counter: u32,
crossfade: CrossfadePhase,
mode: AdaptiveComputeMode,
crossfade_total: usize,
crossfade_elapsed: usize,
slim_override: SlimOverride,
wavenet_full_ch: Option<usize>,
wavenet_slim_ch_current: Option<usize>,
}
const DEGRADE_CONSECUTIVE: u32 = 3;
const RECOVER_CONSECUTIVE: u32 = 5;
const THRESHOLD_FULL_TO_REDUCED_CONSERVATIVE: f32 = 0.70;
const THRESHOLD_REDUCED_TO_MINIMAL_CONSERVATIVE: f32 = 0.85;
const THRESHOLD_FULL_TO_REDUCED_AGGRESSIVE: f32 = 0.55;
const THRESHOLD_REDUCED_TO_MINIMAL_AGGRESSIVE: f32 = 0.70;
const CROSSFADE_DURATION_MS: f32 = 32.0;
impl AdaptiveCompute {
#[cold]
pub fn new(mode: AdaptiveComputeMode) -> Self {
debug!("[Adaptive] FSM initialized: mode={:?}, state=Full", mode);
Self {
state: AdaptiveState::Full,
prev_state: AdaptiveState::Full,
overload_counter: 0,
recovery_counter: 0,
crossfade: CrossfadePhase::Idle,
mode,
crossfade_total: 0,
crossfade_elapsed: 0,
slim_override: SlimOverride::Auto,
wavenet_full_ch: None,
wavenet_slim_ch_current: None,
}
}
#[inline]
pub fn set_mode(&mut self, mode: AdaptiveComputeMode, rt_status: &RtStatusFlags) {
info!("[Adaptive] Mode changed to {:?}", mode);
self.mode = mode;
if mode == AdaptiveComputeMode::Off {
self.state = AdaptiveState::Full;
self.prev_state = AdaptiveState::Full;
self.overload_counter = 0;
self.recovery_counter = 0;
self.crossfade = CrossfadePhase::Idle;
self.crossfade_total = 0;
self.crossfade_elapsed = 0;
rt_status.clear_flag(crate::common::spsc::RT_STATUS_DEGRADE_REDUCED);
rt_status.clear_flag(crate::common::spsc::RT_STATUS_DEGRADE_MINIMAL);
}
}
pub fn mode(&self) -> AdaptiveComputeMode {
self.mode
}
#[inline]
pub fn set_slim_override(&mut self, ov: SlimOverride) {
info!("[Adaptive] Slim override: {:?}", ov);
self.slim_override = ov;
}
pub fn slim_override(&self) -> SlimOverride {
self.slim_override
}
pub fn state(&self) -> AdaptiveState {
match self.slim_override {
SlimOverride::ForceFull => AdaptiveState::Full,
SlimOverride::ForceLite => AdaptiveState::Reduced,
SlimOverride::Auto => self.state,
}
}
#[inline(always)]
pub fn crossfade_multiplier(&mut self, _sample_rate: u32, frame_advance: usize) -> f32 {
match self.crossfade {
CrossfadePhase::Idle => {
if self.state == AdaptiveState::Full {
0.0
} else {
1.0
}
}
CrossfadePhase::Active => {
let total = self.crossfade_total;
let progress = if total == 0 {
1.0
} else {
(self.crossfade_elapsed as f32 / total as f32).min(1.0)
};
if self.crossfade_elapsed + frame_advance >= total {
self.crossfade = CrossfadePhase::Idle;
self.crossfade_elapsed = 0;
1.0
} else {
self.crossfade_elapsed += frame_advance;
progress
}
}
}
}
#[inline(always)]
pub fn is_crossfading(&self) -> bool {
!matches!(self.crossfade, CrossfadePhase::Idle)
}
pub fn update(
&mut self,
latency_us: u64,
budget_us: u64,
sample_rate: u32,
rt_status: &RtStatusFlags,
) {
if self.mode == AdaptiveComputeMode::Off
|| budget_us == 0
|| self.slim_override != SlimOverride::Auto
{
return;
}
let ratio = latency_us as f32 / budget_us as f32;
let (full_to_reduced, reduced_to_minimal) = match self.mode {
AdaptiveComputeMode::Conservative => (
THRESHOLD_FULL_TO_REDUCED_CONSERVATIVE,
THRESHOLD_REDUCED_TO_MINIMAL_CONSERVATIVE,
),
AdaptiveComputeMode::Aggressive => (
THRESHOLD_FULL_TO_REDUCED_AGGRESSIVE,
THRESHOLD_REDUCED_TO_MINIMAL_AGGRESSIVE,
),
AdaptiveComputeMode::Off => return,
};
let recovery_reduced = full_to_reduced * 0.5;
let recovery_minimal = reduced_to_minimal * 0.5;
match self.state {
AdaptiveState::Full => {
if ratio > full_to_reduced {
self.overload_counter = self.overload_counter.saturating_add(1);
self.recovery_counter = 0;
if self.overload_counter >= DEGRADE_CONSECUTIVE {
self.transition_to(AdaptiveState::Reduced, sample_rate, rt_status);
}
} else {
self.overload_counter = 0;
}
}
AdaptiveState::Reduced => {
if ratio > reduced_to_minimal {
self.overload_counter = self.overload_counter.saturating_add(1);
self.recovery_counter = 0;
if self.overload_counter >= DEGRADE_CONSECUTIVE {
self.transition_to(AdaptiveState::Minimal, sample_rate, rt_status);
}
} else if ratio < recovery_reduced {
self.recovery_counter = self.recovery_counter.saturating_add(1);
self.overload_counter = 0;
if self.recovery_counter >= RECOVER_CONSECUTIVE {
self.transition_to(AdaptiveState::Full, sample_rate, rt_status);
self.recovery_counter = 0;
}
} else {
self.overload_counter = 0;
self.recovery_counter = 0;
}
}
AdaptiveState::Minimal => {
if ratio < recovery_minimal {
self.recovery_counter = self.recovery_counter.saturating_add(1);
if self.recovery_counter >= RECOVER_CONSECUTIVE {
self.transition_to(AdaptiveState::Reduced, sample_rate, rt_status);
self.recovery_counter = 0;
}
} else {
self.recovery_counter = 0;
}
}
}
}
fn transition_to(
&mut self,
new_state: AdaptiveState,
sample_rate: u32,
rt_status: &RtStatusFlags,
) {
if self.state == new_state {
return;
}
let current_progress = match self.crossfade {
CrossfadePhase::Active if self.crossfade_total > 0 => {
(self.crossfade_elapsed as f32 / self.crossfade_total as f32).min(1.0)
}
_ => 0.0,
};
if !matches!(self.crossfade, CrossfadePhase::Active) {
self.prev_state = self.state;
}
self.state = new_state;
self.overload_counter = 0;
self.recovery_counter = 0;
let crossfade_samples =
(CROSSFADE_DURATION_MS / 1000.0 * sample_rate as f32).round() as usize;
let new_total = crossfade_samples.max(1);
self.crossfade = CrossfadePhase::Active;
self.crossfade_total = new_total;
self.crossfade_elapsed = (current_progress * new_total as f32).round() as usize;
rt_status
.degrade_transitions_total
.fetch_add(1, Ordering::Relaxed);
match new_state {
AdaptiveState::Full => {
rt_status.clear_flag(crate::common::spsc::RT_STATUS_DEGRADE_REDUCED);
rt_status.clear_flag(crate::common::spsc::RT_STATUS_DEGRADE_MINIMAL);
}
AdaptiveState::Reduced => {
rt_status.set_flag(crate::common::spsc::RT_STATUS_DEGRADE_REDUCED);
rt_status.clear_flag(crate::common::spsc::RT_STATUS_DEGRADE_MINIMAL);
}
AdaptiveState::Minimal => {
rt_status.set_flag(crate::common::spsc::RT_STATUS_DEGRADE_REDUCED);
rt_status.set_flag(crate::common::spsc::RT_STATUS_DEGRADE_MINIMAL);
}
}
}
pub fn wavenet_skip_fraction(&self) -> f32 {
match self.state() {
AdaptiveState::Full => 0.0,
AdaptiveState::Reduced => 0.25,
AdaptiveState::Minimal => 0.50,
}
}
#[inline(always)]
pub fn wavenet_effective_layers(&self, total_layers: usize) -> usize {
let skip = (total_layers as f32 * self.wavenet_skip_fraction()).round() as usize;
total_layers.saturating_sub(skip).max(1)
}
pub fn lstm_effective_layers(&self, total_layers: usize) -> usize {
match self.state() {
AdaptiveState::Full => total_layers,
AdaptiveState::Reduced => total_layers.min(1), AdaptiveState::Minimal => 0, }
}
pub fn slimmable_size(&self) -> f32 {
match self.slim_override {
SlimOverride::ForceFull => 1.0,
SlimOverride::ForceLite => 0.25,
SlimOverride::Auto => match self.state {
AdaptiveState::Full => 1.0,
AdaptiveState::Reduced => 0.25,
AdaptiveState::Minimal => 0.0,
},
}
}
pub fn wavenet_slimmable_ch_target(&self) -> Option<usize> {
let original = self.wavenet_full_ch?;
let fraction = match self.slim_override {
SlimOverride::ForceFull => 1.0,
SlimOverride::ForceLite => 0.75,
SlimOverride::Auto => match self.state {
AdaptiveState::Full => 1.0,
AdaptiveState::Reduced => 0.75,
AdaptiveState::Minimal => 0.50,
},
};
let target = (original as f32 * fraction).round() as usize;
Some(target.max(4).min(original))
}
pub fn set_wavenet_full_ch(&mut self, ch: usize) {
info!("[Adaptive] WaveNet full channels set to {}", ch);
self.wavenet_full_ch = Some(ch);
self.wavenet_slim_ch_current = Some(ch);
}
pub fn take_slimmable_rebuild(&mut self) -> Option<usize> {
let target = self.wavenet_slimmable_ch_target()?;
let current = self.wavenet_slim_ch_current?;
if current != target && target >= 4 {
self.wavenet_slim_ch_current = Some(target);
Some(target)
} else {
None
}
}
#[inline(always)]
pub fn prev_state(&self) -> AdaptiveState {
self.prev_state
}
#[inline(always)]
pub fn current_crossfade_multiplier(&self) -> f32 {
match self.crossfade {
CrossfadePhase::Idle => {
if self.state() == AdaptiveState::Full {
0.0
} else {
1.0
}
}
CrossfadePhase::Active => {
let total = self.crossfade_total;
if total == 0 {
1.0
} else {
(self.crossfade_elapsed as f32 / total as f32).min(1.0)
}
}
}
}
#[inline(always)]
pub fn wavenet_skip_fraction_for_state(&self, state: AdaptiveState) -> f32 {
match state {
AdaptiveState::Full => 0.0,
AdaptiveState::Reduced => 0.25,
AdaptiveState::Minimal => 0.50,
}
}
#[inline(always)]
pub fn wavenet_effective_layers_for_state(
&self,
state: AdaptiveState,
total_layers: usize,
) -> usize {
let state_to_use = match self.slim_override {
SlimOverride::ForceFull => AdaptiveState::Full,
SlimOverride::ForceLite => AdaptiveState::Reduced,
SlimOverride::Auto => state,
};
let skip = (total_layers as f32 * self.wavenet_skip_fraction_for_state(state_to_use))
.round() as usize;
total_layers.saturating_sub(skip).max(1)
}
}
#[cfg(test)]
#[path = "adaptive_test.rs"]
mod adaptive_test;