Skip to main content

AdamW4bit

Struct AdamW4bit 

Source
pub struct AdamW4bit { /* private fields */ }
Expand description

4-bit AdamW: Adam4bit with decoupled weight decay.

Identical quantized state (NF4 momentum and variance) and identical adaptive step; the only difference is that λ·w is subtracted from the parameter directly rather than folded into the gradient, which is what makes AdamW’s decay independent of the gradient magnitude.

Implementations§

Source§

impl AdamW4bit

Source

pub fn new( learning_rate: f32, beta1: f32, beta2: f32, epsilon: f32, weight_decay: f32, ) -> Self

Creates a 4-bit AdamW optimizer.

Source

pub fn with_quantization_config( optimizer_config: Adam4bitOptimizerConfig, quantization_config: AdvancedQuantizationConfig, ) -> Self

Creates a 4-bit AdamW optimizer with a custom quantization configuration.

Source

pub fn memory_savings(&self) -> f32

Memory saved relative to full-precision AdamW, measured from the live buffers.

Trait Implementations§

Source§

impl Debug for AdamW4bit

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result

Formats the value using the given formatter. Read more
Source§

impl Optimizer for AdamW4bit

Source§

fn update(&mut self, parameter: &mut Tensor, grad: &Tensor) -> Result<()>

Updates a parameter based on its gradient. Read more
Source§

fn zero_grad(&mut self)

Clears all accumulated gradients. Read more
Source§

fn step(&mut self)

Performs a single optimization step. Read more
Source§

fn get_lr(&self) -> f32

Gets the current learning rate. Read more
Source§

fn set_lr(&mut self, lr: f32)

Sets a new learning rate. Read more
Source§

fn accumulate_grad( &mut self, parameter: &mut Tensor, grad: &Tensor, ) -> Result<(), TrustformersError>

Accumulates gradients for gradient accumulation. Read more
Source§

fn apply_accumulated_grads( &mut self, accumulation_steps: usize, ) -> Result<(), TrustformersError>

Applies accumulated gradients after gradient accumulation. Read more
Source§

impl StatefulOptimizer for AdamW4bit

Source§

type Config = <Adam4bit as StatefulOptimizer>::Config

The configuration type for this optimizer.
Source§

type State = <Adam4bit as StatefulOptimizer>::State

The state type used by this optimizer.
Source§

fn config(&self) -> &Self::Config

Gets a reference to the optimizer’s configuration.
Source§

fn state(&self) -> &Self::State

Gets a reference to the optimizer’s internal state.
Source§

fn state_mut(&mut self) -> &mut Self::State

Gets a mutable reference to the optimizer’s internal state.
Source§

fn state_dict(&self) -> Result<HashMap<String, Tensor>>

Saves the optimizer state to a dictionary for checkpointing.
Source§

fn load_state_dict(&mut self, state: HashMap<String, Tensor>) -> Result<()>

Loads optimizer state from a dictionary during checkpoint restoration.
Source§

fn memory_usage(&self) -> StateMemoryStats

Gets memory usage statistics for this optimizer.
Source§

fn reset_state(&mut self)

Resets the optimizer state (useful for training restarts).
Source§

fn num_parameters(&self) -> usize

Returns the number of parameters being optimized.
Source§

fn save_state(&self, path: &Path) -> Result<()>

Saves the optimizer state to path. Read more
Source§

fn load_state(&mut self, path: &Path) -> Result<()>

Loads the optimizer state written by Self::save_state. Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T> Instrument for T

Source§

fn instrument(self, span: Span) -> Instrumented<Self>

Instruments this type with the provided Span, returning an Instrumented wrapper. Read more
Source§

fn in_current_span(self) -> Instrumented<Self>

Instruments this type with the current Span, returning an Instrumented wrapper. Read more
Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> IntoEither for T

Source§

fn into_either(self, into_left: bool) -> Either<Self, Self>

Converts self into a Left variant of Either<Self, Self> if into_left is true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
Source§

fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
where F: FnOnce(&Self) -> bool,

Converts self into a Left variant of Either<Self, Self> if into_left(&self) returns true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
Source§

impl<T> Pointable for T

Source§

const ALIGN: usize

The alignment of pointer.
Source§

type Init = T

The type for initializers.
Source§

unsafe fn init(init: <T as Pointable>::Init) -> usize

Initializes a with the given initializer. Read more
Source§

unsafe fn deref<'a>(ptr: usize) -> &'a T

Dereferences the given pointer. Read more
Source§

unsafe fn deref_mut<'a>(ptr: usize) -> &'a mut T

Mutably dereferences the given pointer. Read more
Source§

unsafe fn drop(ptr: usize)

Drops the object pointed to by the given pointer. Read more
Source§

impl<T> Read<Exclusive, BecauseExclusive> for T
where T: ?Sized,

Source§

impl<T> Same for T

Source§

type Output = T

Should always be Self
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, <T as TryFrom<U>>::Error>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.
Source§

impl<V, T> VZip<V> for T
where V: MultiLane<T>,

Source§

fn vzip(self) -> V

Source§

impl<T> WithSubscriber for T

Source§

fn with_subscriber<S>(self, subscriber: S) -> WithDispatch<Self>
where S: Into<Dispatch>,

Attaches the provided Subscriber to this type, returning a WithDispatch wrapper. Read more
Source§

fn with_current_subscriber(self) -> WithDispatch<Self>

Attaches the current default Subscriber to this type, returning a WithDispatch wrapper. Read more