MCPcopy Create free account
hub / github.com/bitsandbytes-foundation/bitsandbytes / Optimizer1State

Class Optimizer1State

bitsandbytes/optim/optimizer.py:593–756  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

591
592
593class Optimizer1State(Optimizer8bit):
594 def __init__(
595 self,
596 optimizer_name,
597 params,
598 lr=1e-3,
599 betas=(0.9, 0.0),
600 eps=1e-8,
601 weight_decay=0.0,
602 optim_bits=32,
603 args=None,
604 min_8bit_size=4096,
605 max_unorm=0.0,
606 skip_zeros=False,
607 is_paged=False,
608 ):
609 """
610 Base 1-state update optimizer class.
611
612 Arguments:
613 optimizer_name (`str`):
614 The name of the optimizer.
615 params (`torch.Tensor`):
616 The input parameters to optimize.
617 lr (`float`, defaults to 1e-3):
618 The learning rate.
619 betas (`tuple`, defaults to (0.9, 0.0)):
620 The beta values for the optimizer.
621 eps (`float`, defaults to 1e-8):
622 The epsilon value for the optimizer.
623 weight_decay (`float`, defaults to 0.0):
624 The weight decay value for the optimizer.
625 optim_bits (`int`, defaults to 32):
626 The number of bits of the optimizer state.
627 args (`object`, defaults to `None`):
628 An object with additional arguments.
629 min_8bit_size (`int`, defaults to 4096):
630 The minimum number of elements of the parameter tensors for 8-bit optimization.
631 max_unorm (`float`, defaults to 0.0):
632 The maximum value to normalize each block with.
633 skip_zeros (`bool`, defaults to `False`):
634 Whether to skip zero values for sparse gradients and models to ensure correct updates.
635 is_paged (`bool`, defaults to `False`):
636 Whether the optimizer is a paged optimizer or not.
637 """
638 if not 0.0 <= lr:
639 raise ValueError(f"Invalid learning rate: {lr}")
640 if not 0.0 <= eps:
641 raise ValueError(f"Invalid epsilon value: {eps}")
642 for i in range(len(betas)):
643 if not 0.0 <= betas[i] < 1.0:
644 raise ValueError(f"Invalid beta parameter at index {i}: {betas[i]}")
645 if not 0.0 <= weight_decay:
646 raise ValueError(f"Invalid weight_decay value: {weight_decay}")
647 defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay)
648 super().__init__(params, defaults, optim_bits, is_paged)
649
650 if args is None:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected