Skip to content

Commit 47e8ee7

Browse files
Refactor optimizer implementations and improve multi_tensor ops (#36)
## Summary Refactor and improve the FlagOS optimizer and multi_tensor implementations to better match CUDA behavior and improve code quality. ## Changes ### `fused_adam.py` (FlagOS backend) - Remove unused `inv_scale` and `out_dtype` parameters from `multi_tensor_adam_fl` - `multi_tensor_adam_param_remainder_fl`: rewrite FP32 master weight reconstruction using bit manipulation (int16 high/low bits), matching the CUDA implementation exactly ### `multi_tensor.py` (FlagOS backend) - `multi_tensor_l2_norm_fl`: add proper type hints, noop_flag check, inf/nan detection, and replace raw `**` / `+` operators with `flag_gems.mul` / `flag_gems.add` - `multi_tensor_scale_fl`: add type hints, noop_flag check, inf/nan detection, and replace `src * scale` with `flag_gems.mul(src, scale)` ### `optimizer.py` (reference backend) - Update `multi_tensor_l2norm_torch` and `multi_tensor_adam_torch` to match new signatures and CUDA behavior (L2 vs AdamW mode split) - Rewrite `multi_tensor_adam_param_remainder_torch` with bit manipulation matching CUDA - Rename `eps` → `epsilon` for consistency ### `optimizers/__init__.py` - Export `multi_tensor_scale` and `multi_tensor_l2norm` ### Misc - Fix missing newline at end of files
1 parent f808816 commit 47e8ee7

4 files changed

Lines changed: 255 additions & 144 deletions

File tree

transformer_engine/plugin/core/backends/flagos/impl/fused_adam.py

Lines changed: 57 additions & 69 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
#
33
# See LICENSE for license information.
44

5-
from typing import Optional, List
5+
from typing import List
66
import torch
77
import 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

8275
def 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)

transformer_engine/plugin/core/backends/flagos/impl/multi_tensor.py

Lines changed: 51 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -2,25 +2,67 @@
22
#
33
# See LICENSE for license information.
44

5+
from typing import List, Tuple
56
import torch
6-
from torch.distributed._tensor import DTensor
77
import flag_gems
88

99

10-
def multi_tensor_l2_norm_fl(chunk_size, noop_flag, tensor_lists, per_tensor, *args):
10+
def multi_tensor_l2_norm_fl(
11+
_chunk_size: int,
12+
noop_flag: torch.Tensor,
13+
tensor_lists: List[List[torch.Tensor]],
14+
per_tensor: bool = False,
15+
) -> Tuple[torch.Tensor, torch.Tensor]:
16+
"""
17+
Compute L2 norm of tensors using flag_gems.
18+
19+
Returns:
20+
Tuple of (total_norm, per_tensor_norms_or_dummy)
21+
- total_norm: The combined L2 norm of all tensors
22+
- per_tensor_norms_or_dummy: Per-tensor norms stacked if per_tensor=True, else dummy tensor
23+
"""
24+
device = tensor_lists[0][0].device if tensor_lists and tensor_lists[0] else "cpu"
25+
26+
if noop_flag.item() != 0:
27+
return torch.tensor(0.0, device=device), torch.tensor(0.0, device=device)
1128

1229
tensors = tensor_lists[0]
1330

31+
# Compute per-tensor norms
32+
per_tensor_norms = []
33+
total_norm_sq = torch.tensor(0.0, device=device)
34+
35+
for tensor in tensors:
36+
t_float = tensor.float()
37+
norm_sq = flag_gems.sum(flag_gems.mul(t_float, t_float))
38+
# Check for inf/nan (matches CUDA behavior)
39+
if not torch.isfinite(norm_sq):
40+
noop_flag.fill_(1)
41+
total_norm_sq = flag_gems.add(total_norm_sq, norm_sq)
42+
if per_tensor:
43+
per_tensor_norms.append(flag_gems.sqrt(norm_sq))
44+
45+
total_norm = flag_gems.sqrt(total_norm_sq)
46+
1447
if per_tensor:
15-
norms = [torch.norm(t.float(), p=2) for t in tensors]
16-
return norms, None
48+
per_tensor_result = torch.stack(per_tensor_norms)
1749
else:
18-
total_norm_sq = sum(flag_gems.sum(flag_gems.pow_func(t.float(), 2)) for t in tensors)
19-
total_norm = flag_gems.sqrt(total_norm_sq)
20-
return total_norm, None
50+
per_tensor_result = torch.tensor(0.0, device=device)
51+
52+
return total_norm, per_tensor_result
2153

2254

23-
def multi_tensor_scale_fl(chunk_size, noop_flag, tensor_lists, scale):
55+
def multi_tensor_scale_fl(
56+
_chunk_size: int,
57+
noop_flag: torch.Tensor,
58+
tensor_lists: List[List[torch.Tensor]],
59+
scale: float,
60+
) -> None:
61+
if noop_flag.item() != 0:
62+
return
2463

2564
for src, dst in zip(tensor_lists[0], tensor_lists[1]):
26-
flag_gems.copy_(dst, src * scale)
65+
# Check for inf/nan (matches CUDA behavior for AMP gradient scaling)
66+
if not torch.isfinite(src).all():
67+
noop_flag.fill_(1)
68+
flag_gems.copy_(dst, flag_gems.mul(src, scale))

0 commit comments

Comments
 (0)