77
88import torch
99
10- from ...ops import TEFLBackendBase , FP8TensorMeta , NVTE_Fused_Attn_Backend
10+ from ...ops import *
1111
1212from .impl import (
1313 rmsnorm_fwd_fl , rmsnorm_bwd_fl ,
2020def _check_flagos_available () -> bool :
2121 return True
2222
23-
2423class 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
0 commit comments