| 591 | |
| 592 | |
| 593 | class 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: |
nothing calls this directly
no outgoing calls
no test coverage detected