22#
33# See LICENSE for license information.
44
5- from typing import Optional , List
5+ from typing import List
66import torch
77import flag_gems
88
@@ -19,8 +19,6 @@ def multi_tensor_adam_fl(
1919 mode : int ,
2020 bias_correction : int ,
2121 weight_decay : float ,
22- inv_scale : Optional [float ] = 1.0 ,
23- out_dtype : Optional [torch .dtype ] = None ,
2422) -> None :
2523
2624 num_lists = len (tensor_lists )
@@ -50,9 +48,6 @@ def multi_tensor_adam_fl(
5048 if not g .is_contiguous ():
5149 g = g .contiguous ()
5250
53- if inv_scale is not None and inv_scale != 1.0 :
54- g = flag_gems .mul (g , inv_scale )
55-
5651 m = flag_gems .add_ (flag_gems .mul_ (m , beta1 ), g , alpha = 1 - beta1 )
5752 v = flag_gems .add_ (
5853 flag_gems .mul_ (v , beta2 ), flag_gems .mul_ (flag_gems .mul_ (g , g ), 1 - beta2 )
@@ -75,8 +70,6 @@ def multi_tensor_adam_fl(
7570
7671 if p_master is not None :
7772 flag_gems .copy_ (p_master , p )
78- out_dtype = p_master .dtype if out_dtype is None else out_dtype
79- p .data = p .data .to (out_dtype )
8073
8174
8275def multi_tensor_adam_param_remainder_fl (
@@ -91,27 +84,9 @@ def multi_tensor_adam_param_remainder_fl(
9184 mode : int ,
9285 bias_correction : int ,
9386 weight_decay : float ,
94- inv_scale : Optional [float ] = 1.0 ,
9587) -> None :
9688 """
9789 Adam optimizer with parameter remainders for BF16 precision (FlagOS implementation).
98-
99- This variant stores BF16 parameters + int16 remainders to reconstruct FP32 master weights.
100- Used when you have BF16 params and need FP32 master params without storing full FP32 copies.
101-
102- Args:
103- chunk_size: Chunk size for processing (unused in this implementation)
104- noop_flag: If non-zero, skip computation
105- tensor_lists: [grads, params (bf16), exp_avgs (fp32), exp_avg_sqs (fp32), param_remainders (int16)]
106- lr: Learning rate
107- beta1: First moment decay rate
108- beta2: Second moment decay rate
109- eps: Epsilon for numerical stability
110- step: Current optimization step
111- mode: 0 = L2 regularization, 1 = AdamW (decoupled weight decay)
112- bias_correction: Whether to apply bias correction (1 = yes, 0 = no)
113- weight_decay: Weight decay coefficient
114- inv_scale: Inverse gradient scale for mixed precision training
11590 """
11691 if noop_flag .item () != 0 :
11792 return
@@ -135,65 +110,78 @@ def multi_tensor_adam_param_remainder_fl(
135110
136111 for i in range (num_tensors ):
137112 g = tensor_lists [0 ][i ]
138- p = tensor_lists [1 ][i ] # BF16 parameter
113+ p = tensor_lists [1 ][i ] # int16 parameter (high 16 bits of FP32)
139114 m = tensor_lists [2 ][i ] # FP32 first moment
140115 v = tensor_lists [3 ][i ] # FP32 second moment
141- p_remainder = tensor_lists [4 ][i ] # int16 remainder
116+ p_remainder = tensor_lists [4 ][i ] # int16 remainder (low 16 bits of FP32)
142117
143118 if not g .is_contiguous ():
144119 g = g .contiguous ()
145120
146- # Apply gradient unscaling if needed
147- if inv_scale is not None and inv_scale != 1.0 :
148- g = flag_gems .mul (g , inv_scale )
121+ # Convert gradient to float
122+ g_float = g .float ()
149123
150- # Reconstruct FP32 master weight from BF16 param + int16 remainder
151- # The remainder represents the lower 16 bits lost in BF16 conversion
152- param_fp32 = p .float ()
153- param_master = flag_gems .add (param_fp32 , flag_gems .mul (p_remainder .float (), 2.0 ** - 16 ))
124+ # Reconstruct FP32 master weight from int16 param + int16 remainder using bit manipulation
125+ # This matches the CUDA implementation exactly:
126+ # 1. If p_remainder < 0, decrement p (undo rounding)
127+ # 2. Combine high 16 bits (p) and low 16 bits (p_remainder) into FP32
128+ # Note: Use PyTorch native ops for bit manipulation (int16/int32 operations)
154129
155- # Compute gradient with weight decay (if L2 mode)
156- grad_with_decay = g .float ()
157- if not is_adamw : # L2 regularization mode
158- grad_with_decay = flag_gems .add (
159- grad_with_decay , flag_gems .mul (param_master , weight_decay )
160- )
130+ local_p = p .view (torch .int16 ).clone ()
131+ local_p_rem = p_remainder .clone ()
161132
162- # Update moments
163- m = flag_gems .add_ (flag_gems .mul_ (m , beta1 ), grad_with_decay , alpha = 1 - beta1 )
164- v = flag_gems .add_ (
165- flag_gems .mul_ (v , beta2 ),
166- flag_gems .mul_ (flag_gems .mul_ (grad_with_decay , grad_with_decay ), 1 - beta2 ),
167- )
133+ # Undo rounding: if remainder < 0, decrement p
134+ local_p = torch .where (local_p_rem < 0 , local_p - 1 , local_p )
135+
136+ # Combine into FP32 using bit shift operations
137+ # local_p is high 16 bits, local_p_rem is low 16 bits
138+ high_bits = local_p .to (torch .int32 ) << 16
139+ low_bits = local_p_rem .to (torch .int32 ) & 0xFFFF # Mask off sign extension
140+ param_int32 = high_bits | low_bits
141+ param_master = param_int32 .view (torch .float32 )
142+
143+ # L2 mode: add weight decay to gradient before updating moments
144+ if not is_adamw and weight_decay != 0 :
145+ g_float = flag_gems .add (g_float , param_master , alpha = weight_decay )
146+
147+ # Update first moment: m = beta1 * m + (1 - beta1) * g
148+ flag_gems .add_ (flag_gems .mul_ (m , beta1 ), g_float , alpha = 1 - beta1 )
149+
150+ # Update second moment: v = beta2 * v + (1 - beta2) * g^2
151+ flag_gems .add_ (flag_gems .mul_ (v , beta2 ), flag_gems .mul (g_float , g_float ), alpha = 1 - beta2 )
168152
169153 # Apply bias correction
170- m_corr = m . clone ( )
171- v_corr = v . clone ( )
172- if bias_correction == 1 :
173- m_corr = flag_gems . true_divide ( m_corr , bias_correction1 )
174- v_corr = flag_gems .true_divide ( v_corr , bias_correction2 )
154+ m_corr = flag_gems . true_divide ( m , bias_correction1 )
155+ v_corr = flag_gems . true_divide ( v , bias_correction2 )
156+
157+ # Compute denominator: sqrt(v_corr) + eps
158+ denom = flag_gems .add ( flag_gems . sqrt ( v_corr ), eps )
175159
176160 # Compute update
177- update = flag_gems .true_divide (m_corr , flag_gems . add ( flag_gems . sqrt ( v_corr ), eps ) )
161+ update = flag_gems .true_divide (m_corr , denom )
178162
179- # Apply weight decay (if AdamW mode)
180- if is_adamw :
181- param_master = flag_gems .mul_ ( param_master , 1 - lr * weight_decay )
163+ # AdamW mode: add decoupled weight decay to update
164+ if is_adamw and weight_decay != 0 :
165+ update = flag_gems .add ( update , param_master , alpha = weight_decay )
182166
183- # Update master weight
184- param_master = flag_gems .add_ (param_master , update , alpha = - lr )
167+ # Update master weight: p = p - lr * update
168+ param_master = flag_gems .sub (param_master , flag_gems . mul ( update , lr ) )
185169
186- # Split back into BF16 param + int16 remainder
187- # Convert to BF16 (this is the rounded version)
188- param_bf16 = param_master .to (dtype = p .dtype )
170+ # Split FP32 back into int16 param + int16 remainder using bit manipulation
171+ # This matches the CUDA implementation exactly:
172+ # 1. Extract high 16 bits as p
173+ # 2. Extract low 16 bits as p_remainder
174+ # 3. If p_remainder < 0, increment p (round up)
175+ # Note: Use PyTorch native ops for bit manipulation (int32 operations)
189176
190- # Compute remainder: difference between FP32 master and BF16 representation
191- # Scale and quantize to int16 range
192- remainder_fp32 = flag_gems .mul (flag_gems .sub (param_master , param_bf16 .float ()), 2.0 ** 16 )
193- remainder_int16 = flag_gems .clamp (torch .round (remainder_fp32 ), - 32768 , 32767 ).to (
194- dtype = torch .int16
195- )
177+ param_int32 = param_master .view (torch .int32 )
178+ # Extract low 16 bits (remainder) and high 16 bits (param)
179+ new_p_rem = (param_int32 & 0xFFFF ).to (torch .int16 )
180+ new_p = ((param_int32 >> 16 ) & 0xFFFF ).to (torch .int16 )
181+
182+ # Round up: if remainder < 0, increment p
183+ new_p = torch .where (new_p_rem < 0 , new_p + 1 , new_p )
196184
197185 # Write back
198- flag_gems .copy_ (p , param_bf16 )
199- flag_gems .copy_ (p_remainder , remainder_int16 )
186+ flag_gems .copy_ (p , new_p . view ( torch . bfloat16 ) )
187+ flag_gems .copy_ (p_remainder , new_p_rem )
0 commit comments