Skip to content

Commit d44924f

Browse files
fix other vendor , flagos , torch func mismatch
1 parent 04f0cd6 commit d44924f

15 files changed

Lines changed: 2178 additions & 1443 deletions

File tree

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

Lines changed: 56 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
import torch
99

10-
from ...ops import TEFLBackendBase, FP8TensorMeta, NVTE_Fused_Attn_Backend
10+
from ...ops import *
1111

1212
from .impl import (
1313
rmsnorm_fwd_fl, rmsnorm_bwd_fl,
@@ -20,7 +20,6 @@
2020
def _check_flagos_available() -> bool:
2121
return True
2222

23-
2423
class FlagOSBackend(TEFLBackendBase):
2524
@staticmethod
2625
def check_available() -> bool:
@@ -29,10 +28,6 @@ def check_available() -> bool:
2928
def is_available(self) -> bool:
3029
return _check_flagos_available()
3130

32-
def get_flash_attention_class(self):
33-
from .attention.dot_product_attention.backends import FlashAttentionFL
34-
return FlashAttentionFL
35-
3631
def get_attention_backend(self, attention_params=None):
3732
from packaging.version import Version as PkgVersion
3833
from ...logger_manager import get_logger
@@ -65,48 +60,50 @@ def get_attention_backend(self, attention_params=None):
6560
available_backends,
6661
)
6762

63+
##### transformer_engine/pytorch/csrc/extensions/pybind.cpp #####
6864
def generic_gemm(
6965
self,
70-
A: torch.Tensor,
71-
transA: bool,
72-
B: torch.Tensor,
73-
transB: bool,
74-
D: torch.Tensor,
66+
A: Any,
67+
transa: bool,
68+
B: Any,
69+
transb: bool,
70+
D: Any,
7571
quantizer: Any,
76-
output_dtype: torch.dtype,
72+
out_dtype: Optional[DType],
7773
bias: Optional[torch.Tensor],
78-
bias_type: Any,
74+
bias_type: DType,
7975
gelu: bool,
8076
gelu_in: Optional[torch.Tensor],
8177
grad: bool,
8278
workspace: torch.Tensor,
83-
workspace_size: int,
79+
workspaceSize: int,
8480
accumulate: bool,
8581
use_split_accumulator: bool,
8682
comm_overlap: Optional[Any] = None,
87-
comm_type: Optional[Any] = None,
83+
comm_type: Optional[CommOverlapType] = None,
8884
extra_output: Optional[torch.Tensor] = None,
8985
bulk_overlap: bool = False,
9086
alpha: float = 1.0,
9187
beta: Optional[float] = None,
92-
) -> Any:
88+
) -> List[Any]:
9389
return generic_gemm_fl(
94-
A, transA, B, transB, D, quantizer, output_dtype,
90+
A, transa, B, transb, D, quantizer, out_dtype,
9591
bias, bias_type, gelu, gelu_in, grad,
96-
workspace, workspace_size, accumulate, use_split_accumulator,
92+
workspace, workspaceSize, accumulate, use_split_accumulator,
9793
comm_overlap=comm_overlap, comm_type=comm_type,
9894
extra_output=extra_output, bulk_overlap=bulk_overlap,
9995
alpha=alpha, beta=beta
10096
)
10197

98+
# Other granular functions
10299
def rmsnorm_fwd(
103100
self,
104101
input: torch.Tensor,
105102
weight: torch.Tensor,
106103
eps: float,
107104
ln_out: Optional[torch.Tensor],
108105
quantizer: Any,
109-
otype: torch.dtype,
106+
otype: DType,
110107
sm_margin: int,
111108
zero_centered_gamma: bool,
112109
) -> Tuple[torch.Tensor, Optional[torch.Tensor], torch.Tensor]:
@@ -115,22 +112,23 @@ def rmsnorm_fwd(
115112
quantizer=quantizer, odtype=otype,
116113
sm_margin=sm_margin, zero_centered_gamma=zero_centered_gamma,
117114
)
118-
119115
def rmsnorm_bwd(
120116
self,
121-
dy: torch.Tensor,
117+
dz: torch.Tensor,
122118
x: torch.Tensor,
123119
rsigma: torch.Tensor,
124120
gamma: torch.Tensor,
125121
sm_margin: int = 0,
126122
zero_centered_gamma: bool = False,
127-
eps: float = 1e-5,
128123
) -> Tuple[torch.Tensor, torch.Tensor]:
129124
return rmsnorm_bwd_fl(
130-
dy=dy, x=x, rsigma=rsigma, gamma=gamma,
131-
sm_margin=sm_margin, zero_centered_gamma=zero_centered_gamma, eps=eps,
125+
dy=dz, x=x, rsigma=rsigma, gamma=gamma,
126+
sm_margin=sm_margin, zero_centered_gamma=zero_centered_gamma
132127
)
128+
def get_fused_attn_backend(self, *args, **kwargs) -> int:
129+
return NVTE_Fused_Attn_Backend.NVTE_No_Backend
133130

131+
# multi-tensor functions
134132
def multi_tensor_scale(
135133
self,
136134
chunk_size: int,
@@ -139,73 +137,66 @@ def multi_tensor_scale(
139137
scale: float,
140138
) -> None:
141139
return multi_tensor_scale_fl(chunk_size, noop_flag, tensor_lists, scale)
142-
143140
def multi_tensor_l2norm(
144141
self,
145142
chunk_size: int,
146143
noop_flag: torch.Tensor,
147144
tensor_lists: List[List[torch.Tensor]],
148-
per_tensor: bool = False,
149-
) -> Union[torch.Tensor, List[torch.Tensor]]:
150-
result, _ = multi_tensor_l2_norm_fl(chunk_size, noop_flag, tensor_lists, per_tensor)
151-
return result
152-
145+
per_tensor: Optional[bool] = False,
146+
) -> Tuple[torch.Tensor, torch.Tensor]:
147+
return multi_tensor_l2_norm_fl(chunk_size, noop_flag, tensor_lists, per_tensor)
153148
def multi_tensor_adam(
154149
self,
155-
chunk_size: int = None,
156-
noop_flag: torch.Tensor = None,
157-
tensor_lists: List[List[torch.Tensor]] = None,
158-
lr: float = None,
159-
beta1: float = None,
160-
beta2: float = None,
161-
eps: float = None,
162-
step: int = None,
163-
mode: int = None,
164-
bias_correction: int = None,
165-
weight_decay: float = None,
166-
):
150+
chunk_size: int,
151+
noop_flag: torch.Tensor,
152+
tensor_lists: List[List[torch.Tensor]],
153+
lr: float,
154+
beta1: float,
155+
beta2: float,
156+
epsilon: float,
157+
step: int,
158+
mode: int,
159+
bias_correction: int,
160+
weight_decay: float,
161+
) -> None:
167162
if chunk_size is None:
168163
return multi_tensor_adam_fl
169164
return multi_tensor_adam_fl(
170165
chunk_size=chunk_size, noop_flag=noop_flag, tensor_lists=tensor_lists,
171-
lr=lr, beta1=beta1, beta2=beta2, eps=eps,
166+
lr=lr, beta1=beta1, beta2=beta2, eps=epsilon,
172167
step=step, mode=mode, bias_correction=bias_correction, weight_decay=weight_decay,
173168
)
174-
175169
def multi_tensor_adam_param_remainder(
176170
self,
177-
chunk_size: int = None,
178-
noop_flag: torch.Tensor = None,
179-
tensor_lists: List[List[torch.Tensor]] = None,
180-
lr: float = None,
181-
beta1: float = None,
182-
beta2: float = None,
183-
eps: float = None,
184-
step: int = None,
185-
mode: int = None,
186-
bias_correction: int = None,
187-
weight_decay: float = None,
188-
):
171+
chunk_size: int,
172+
noop_flag: torch.Tensor,
173+
tensor_lists: List[List[torch.Tensor]],
174+
lr: float,
175+
beta1: float,
176+
beta2: float,
177+
epsilon: float,
178+
step: int,
179+
mode: int,
180+
bias_correction: int,
181+
weight_decay: float,
182+
) -> None:
189183
if chunk_size is None:
190184
return multi_tensor_adam_param_remainder_fl
191185
return multi_tensor_adam_param_remainder_fl(
192186
chunk_size=chunk_size, noop_flag=noop_flag, tensor_lists=tensor_lists,
193-
lr=lr, beta1=beta1, beta2=beta2, eps=eps,
187+
lr=lr, beta1=beta1, beta2=beta2, eps=epsilon,
194188
step=step, mode=mode, bias_correction=bias_correction, weight_decay=weight_decay,
195189
)
196190

191+
# Misc
197192
def get_cublasLt_version(self) -> int:
198193
return 110000
199-
200194
def get_cudnn_version(self) -> int:
201195
return 90000
202-
203196
def get_num_cublas_streams(self) -> int:
204197
return 0
205198

206-
def get_fused_attn_backend(self, *args, **kwargs) -> int:
207-
return NVTE_Fused_Attn_Backend.NVTE_No_Backend
208-
209-
def create_fp8_tensor_meta(self) -> FP8TensorMeta:
210-
return FP8TensorMeta()
211-
199+
############## class func #################################
200+
def get_flash_attention_class(self):
201+
from .attention.dot_product_attention.backends import FlashAttentionFL
202+
return FlashAttentionFL

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

Lines changed: 43 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -2,25 +2,60 @@
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+
norm_sq = flag_gems.sum(tensor.float() ** 2)
37+
total_norm_sq = total_norm_sq + norm_sq
38+
if per_tensor:
39+
per_tensor_norms.append(flag_gems.sqrt(norm_sq))
40+
41+
total_norm = flag_gems.sqrt(total_norm_sq)
42+
1443
if per_tensor:
15-
norms = [torch.norm(t.float(), p=2) for t in tensors]
16-
return norms, None
44+
per_tensor_result = torch.stack(per_tensor_norms)
1745
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
46+
per_tensor_result = torch.tensor(0.0, device=device)
47+
48+
return total_norm, per_tensor_result
2149

2250

23-
def multi_tensor_scale_fl(chunk_size, noop_flag, tensor_lists, scale):
51+
def multi_tensor_scale_fl(
52+
_chunk_size: int,
53+
noop_flag: torch.Tensor,
54+
tensor_lists: List[List[torch.Tensor]],
55+
scale: float,
56+
) -> None:
57+
if noop_flag.item() != 0:
58+
return
2459

2560
for src, dst in zip(tensor_lists[0], tensor_lists[1]):
2661
flag_gems.copy_(dst, src * scale)

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,7 @@ def rmsnorm_bwd_fl(
4242
gamma,
4343
sm_margin,
4444
zero_centered_gamma,
45-
eps,
45+
eps=1e-5,
4646
):
4747
# When zero_centered_gamma is True, forward uses (1 + gamma) as weight
4848
# So backward needs to use (1 + gamma) for computing dx

transformer_engine/plugin/core/backends/reference/impl/normalization.py

Lines changed: 27 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,37 @@
55
from typing import Any, Optional, Tuple
66
import torch
77
import torch.nn.functional as F
8+
from ....ops import DType
89

910
__all__ = [
1011
"layernorm_fwd_torch",
1112
"layernorm_bwd_torch",
1213
]
1314

15+
# Mapping from DType enum to torch.dtype
16+
_DTYPE_TO_TORCH_DTYPE = {
17+
DType.kByte: torch.uint8,
18+
DType.kInt16: torch.int16,
19+
DType.kInt32: torch.int32,
20+
DType.kInt64: torch.int64,
21+
DType.kFloat32: torch.float32,
22+
DType.kFloat16: torch.float16,
23+
DType.kBFloat16: torch.bfloat16,
24+
DType.kFloat8E4M3: torch.float8_e4m3fn,
25+
DType.kFloat8E5M2: torch.float8_e5m2,
26+
}
27+
28+
def _to_torch_dtype(dtype):
29+
"""Convert DType enum to torch.dtype."""
30+
if dtype is None:
31+
return None
32+
if isinstance(dtype, torch.dtype):
33+
return dtype
34+
if isinstance(dtype, (int, DType)):
35+
dtype_enum = DType(dtype)
36+
if dtype_enum in _DTYPE_TO_TORCH_DTYPE:
37+
return _DTYPE_TO_TORCH_DTYPE[dtype_enum]
38+
raise ValueError(f"Unsupported dtype: {dtype}")
1439

1540
def layernorm_fwd_torch(
1641
input: torch.Tensor,
@@ -19,10 +44,11 @@ def layernorm_fwd_torch(
1944
eps: float,
2045
ln_out: Optional[torch.Tensor],
2146
quantizer: Any,
22-
odtype: torch.dtype,
47+
odtype: DType,
2348
sm_margin: int,
2449
zero_centered_gamma: bool,
2550
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
51+
odtype = _to_torch_dtype(odtype)
2652
mean = input.mean(dim=-1, keepdim=True)
2753
var = input.var(dim=-1, keepdim=True, unbiased=False)
2854
rsigma = torch.rsqrt(var + eps)
@@ -45,7 +71,6 @@ def layernorm_fwd_torch(
4571

4672
return output, mean, rsigma
4773

48-
4974
def layernorm_bwd_torch(
5075
dy: torch.Tensor,
5176
x: torch.Tensor,

0 commit comments

Comments
 (0)