optimizer

package
v1.60.0 Latest Latest
Warning

This package is not in the latest version of its module.

Go to latest
Published: Aug 23, 2026 License: Apache-2.0 Imports: 9 Imported by: 0

Documentation

Overview

Package optimizer provides various optimization algorithms for neural networks.

Package optimizer provides neural network optimizers including AdamW and SGD.

Stability: beta

Package optimizer provides various optimization algorithms for neural networks.

Package optimizer provides various optimization algorithms for neural networks.

Index

Constants

This section is empty.

Variables

This section is empty.

Functions

func AdamWUpdateF64 added in v1.41.0

func AdamWUpdateF64(params []*float64, grads []float64, state *AdamWStateF64, lr, beta1, beta2, epsilon, weightDecay, t float64)

AdamWUpdateF64 applies one AdamW parameter update step on raw float64 slices. t is the global step count (1-indexed for bias correction).

func ClipGradientsF64 added in v1.41.0

func ClipGradientsF64(grads []float64, maxNorm float64)

ClipGradientsF64 applies global norm gradient clipping in-place. If the L2 norm of grads exceeds maxNorm, all gradients are scaled down proportionally. A maxNorm of 0 disables clipping.

Types

type AdamW

type AdamW[T tensor.Numeric] struct {
	// contains filtered or unexported fields
}

AdamW implements the AdamW optimizer.

When T is float32 or a sub-float32 precision (float16, float8) AND the engine is a CPU engine, the second-moment accumulator (v) is held as a float64 sidecar instead of T and the per-element "sqrt(v) + epsilon" and "m / (sqrt(v) + epsilon)" computations run in float64. This removes the underflow cliff where v drifts into denormals and sqrt(v) + eps collapses to eps, producing runaway update magnitudes. Storage for param.Value and param.Gradient is unchanged. On GPU engines the original T-only path is preserved; mixed-precision GPU kernels are a follow-up.

func NewAdamW

func NewAdamW[T tensor.Numeric](engine compute.Engine[T], learningRate, beta1, beta2, epsilon, weightDecay T) *AdamW[T]

NewAdamW creates a new AdamW optimizer.

For reduced-precision element types (bfloat16, float16, float8) prefer NewAdamWFromFloat64: passing hyperparameters as T rounds beta2 = 0.999 to 1.0 and epsilon to 0 in bfloat16, which disables learning. NewAdamW recovers the float64 sidecars from the T inputs, so it is lossless for float32/float64 callers but inherits whatever precision the caller already lost when it constructed the T-typed beta2/epsilon.

func NewAdamWFromFloat64 added in v1.52.0

func NewAdamWFromFloat64[T tensor.Numeric](
	engine compute.Engine[T],
	learningRate, beta1, beta2, epsilon, weightDecay float64,
) *AdamW[T]

NewAdamWFromFloat64 creates a new AdamW optimizer with the hyperparameters given in full float64 precision, which is the correct constructor for reduced-precision element types (T = bfloat16/float16/float8).

The T-typed fields (used only by the all-T stepEngine path and by SetLR) are derived once via ops.FromFloat64, but the float64 fields retain the exact values, and the mixed-precision update -- the path every sub-f32 T takes -- reads only the float64 fields. This keeps beta2 = 0.999 and epsilon = 1e-7 from collapsing to 1.0 / 0 in bfloat16.

func (*AdamW[T]) SetLR added in v1.8.0

func (a *AdamW[T]) SetLR(lr T)

SetLR sets the learning rate. This is typically called by a scheduler.

Both the T-typed field (stepEngine path) and the float64 sidecar (mixed- precision path) are updated so the new rate takes effect regardless of T.

func (*AdamW[T]) SetLRFloat64 added in v1.52.0

func (a *AdamW[T]) SetLRFloat64(lr float64)

SetLRFloat64 sets the learning rate in full float64 precision. Prefer this over SetLR for reduced-precision T so a small rate (e.g. 1e-4) is not rounded when stored in the T-typed field.

func (*AdamW[T]) SetMaxGradNorm added in v1.11.0

func (a *AdamW[T]) SetMaxGradNorm(maxGradNorm float64)

SetMaxGradNorm sets the maximum gradient norm for gradient clipping. If maxGradNorm <= 0, gradient clipping is disabled.

func (*AdamW[T]) Step

func (a *AdamW[T]) Step(ctx context.Context, params []*graph.Parameter[T]) error

Step updates the parameters based on their gradients.

type AdamW8bit added in v1.5.0

type AdamW8bit[T tensor.Numeric] struct {
	// contains filtered or unexported fields
}

AdamW8bit implements the AdamW optimizer with block-wise INT8 quantization for first and second moment estimates. Parameters remain in full precision. This reduces optimizer state memory by ~4x compared to FP32 AdamW.

func NewAdamW8bit added in v1.5.0

func NewAdamW8bit[T tensor.Numeric](engine compute.Engine[T], lr, beta1, beta2, eps, wd float32) *AdamW8bit[T]

NewAdamW8bit creates a new 8-bit AdamW optimizer.

func (*AdamW8bit[T]) Step added in v1.5.0

func (a *AdamW8bit[T]) Step(ctx context.Context, params []*graph.Parameter[T]) error

Step updates parameters based on their gradients. Moment estimates are stored in INT8 and dequantized for computation, then re-quantized after update.

type AdamWStateF64 added in v1.41.0

type AdamWStateF64 struct {
	M []float64 // first moment
	V []float64 // second moment
}

AdamWStateF64 holds first and second moment estimates for scalar AdamW updates.

func NewAdamWStateF64 added in v1.41.0

func NewAdamWStateF64(nParams int) *AdamWStateF64

NewAdamWStateF64 creates a new scalar AdamW state for nParams parameters.

type EMA added in v0.2.1

type EMA[T tensor.Numeric] struct {
	// contains filtered or unexported fields
}

EMA wraps an Optimizer with Exponential Moving Average weight averaging. After each inner optimizer step, it updates shadow weights:

shadow = decay * shadow + (1-decay) * param.Value

Call SwapShadow before validation to use averaged weights, then SwapBack to restore training weights.

func NewEMA added in v0.2.1

func NewEMA[T tensor.Numeric](inner Optimizer[T], engine compute.Engine[T], decay T) *EMA[T]

NewEMA creates a new EMA wrapper around the given optimizer.

func (*EMA[T]) Step added in v0.2.1

func (e *EMA[T]) Step(ctx context.Context, params []*graph.Parameter[T]) error

Step runs the inner optimizer step and then updates shadow weights.

func (*EMA[T]) SwapBack added in v0.2.1

func (e *EMA[T]) SwapBack(ctx context.Context, params []*graph.Parameter[T]) error

SwapBack is a semantic alias for SwapShadow — the swap operation is symmetric.

func (*EMA[T]) SwapShadow added in v0.2.1

func (e *EMA[T]) SwapShadow(ctx context.Context, params []*graph.Parameter[T]) error

SwapShadow swaps param.Value with shadow weights for validation/checkpointing.

type Int8State added in v1.5.0

type Int8State struct {
	// contains filtered or unexported fields
}

Int8State holds a block-wise INT8-quantized representation of a float32 slice. Each block of blockSize elements shares a single scale factor, reducing memory from 4 bytes/element to ~1 byte/element (+ negligible scale overhead).

type Optimizer

type Optimizer[T tensor.Numeric] interface {
	Step(ctx context.Context, params []*graph.Parameter[T]) error
}

Optimizer defines the interface for optimization algorithms.

type SGD

type SGD[T tensor.Numeric] struct {
	// contains filtered or unexported fields
}

SGD implements the stochastic gradient descent optimizer.

func NewSGD

func NewSGD[T tensor.Numeric](engine compute.Engine[T], ops numeric.Arithmetic[T], learningRate float32) *SGD[T]

NewSGD creates a new SGD optimizer.

func (*SGD[T]) Clip

func (s *SGD[T]) Clip(ctx context.Context, params []*graph.Parameter[T], threshold float32)

Clip clips the gradients of the parameters by a threshold.

func (*SGD[T]) SetLR added in v1.8.0

func (s *SGD[T]) SetLR(lr T)

SetLR sets the learning rate. This is typically called by a scheduler.

func (*SGD[T]) Step

func (s *SGD[T]) Step(ctx context.Context, params []*graph.Parameter[T]) error

Step updates the parameters based on their gradients.

type SWA added in v0.2.1

type SWA[T tensor.Numeric] struct {
	// contains filtered or unexported fields
}

SWA wraps an Optimizer with Stochastic Weight Averaging. Unlike EMA which averages every step, SWA averages at epoch boundaries. Call UpdateAverage at the end of each epoch (after startEpoch). Call SwapWeights before validation to use averaged weights.

func NewSWA added in v0.2.1

func NewSWA[T tensor.Numeric](inner Optimizer[T], engine compute.Engine[T], startEpoch int) *SWA[T]

NewSWA creates a new SWA wrapper around the given optimizer.

func (*SWA[T]) NAveraged added in v0.2.1

func (s *SWA[T]) NAveraged() int

NAveraged returns the number of checkpoints averaged so far.

func (*SWA[T]) Step added in v0.2.1

func (s *SWA[T]) Step(ctx context.Context, params []*graph.Parameter[T]) error

Step delegates to the inner optimizer.

func (*SWA[T]) SwapWeights added in v0.2.1

func (s *SWA[T]) SwapWeights(ctx context.Context, params []*graph.Parameter[T]) error

SwapWeights swaps live params with averaged params.

func (*SWA[T]) UpdateAverage added in v0.2.1

func (s *SWA[T]) UpdateAverage(ctx context.Context, params []*graph.Parameter[T], epoch int) error

UpdateAverage updates the running average of parameters. Should be called at the end of each epoch. Only averages when epoch >= startEpoch. Formula: avg = avg + (param - avg) / (n + 1)

Jump to

Keyboard shortcuts

? : This menu
/ : Search site
f or F : Jump to
y or Y : Canonical URL