1616
1717
1818def _run_gpt (
19+ dtype = "fp32" ,
1920 device_throughput = constants .DEFAULT_DEVICE_THROUGHPUT ,
2021 dram_bandwidth = constants .DEFAULT_DRAM_BANDWIDTH ,
2122 kernel_launch_overhead = constants .DEFAULT_KERNEL_LAUNCH_OVERHEAD ,
@@ -38,6 +39,7 @@ def _run_gpt(
3839 topology ,
3940 ) = get_transformed_function_and_input_data (
4041 MODEL_PATH ,
42+ dtype ,
4143 device_throughput ,
4244 dram_bandwidth ,
4345 kernel_launch_overhead ,
@@ -77,6 +79,7 @@ def _run_gpt(
7779
7880def _test (
7981 original_outputs ,
82+ dtype ,
8083 dp_degree = 1 ,
8184 hp_degree = 1 ,
8285 pp_degree = 1 ,
@@ -86,6 +89,7 @@ def _test(
8689
8790 # Test with real weights
8891 transformed_outputs = _run_gpt (
92+ dtype = dtype ,
8993 dp_degree = dp_degree ,
9094 hp_degree = hp_degree ,
9195 pp_degree = pp_degree ,
@@ -95,13 +99,18 @@ def _test(
9599 assert len (transformed_outputs ) == dp_degree * hp_degree
96100 for i in range (len (transformed_outputs )):
97101 np .testing .assert_array_almost_equal (
98- original_outputs [0 ].val , transformed_outputs [i ].val , decimal = 2
102+ original_outputs [0 ].val ,
103+ transformed_outputs [i ].val ,
104+ decimal = (2 if dtype == "fp32" else 1 ),
99105 )
100106
101107
102108@pytest .fixture (scope = "session" )
103109def original_outputs ():
104- return _run_gpt ()
110+ return {
111+ "fp16" : _run_gpt (dtype = "fp16" , use_pytorch_backend = True ),
112+ "fp32" : _run_gpt (dtype = "fp32" , use_pytorch_backend = True ),
113+ }
105114
106115
107116@pytest .mark .parametrize (
@@ -110,7 +119,8 @@ def original_outputs():
110119)
111120def test_reference_execution (original_outputs , dp_degree , hp_degree , pp_degree ):
112121 _test (
113- original_outputs ,
122+ original_outputs ["fp32" ],
123+ dtype = "fp32" ,
114124 dp_degree = dp_degree ,
115125 hp_degree = hp_degree ,
116126 pp_degree = pp_degree ,
@@ -119,12 +129,23 @@ def test_reference_execution(original_outputs, dp_degree, hp_degree, pp_degree):
119129
120130
121131@pytest .mark .parametrize (
122- ("dp_degree" , "hp_degree" , "pp_degree" ),
123- list (itertools .product ([1 , 2 ], [1 , 2 ], [1 , 2 ])),
132+ ("dtype" , "dp_degree" , "hp_degree" , "pp_degree" ),
133+ list (
134+ itertools .product (
135+ ["fp16" , "fp32" ] if torch .cuda .is_available () else ["fp32" ],
136+ [1 , 2 ],
137+ [1 , 2 ],
138+ [1 , 2 ],
139+ )
140+ ),
124141)
125- def test_pytorch_backend (original_outputs , dp_degree , hp_degree , pp_degree ):
142+ def test_pytorch_backend (original_outputs , dtype , dp_degree , hp_degree , pp_degree ):
143+ world_size = dp_degree * hp_degree * pp_degree
144+ if dtype == "fp16" and world_size > torch .cuda .device_count ():
145+ pytest .skip ("Not enough GPUs available" )
126146 _test (
127- original_outputs ,
147+ original_outputs [dtype ],
148+ dtype ,
128149 dp_degree = dp_degree ,
129150 hp_degree = hp_degree ,
130151 pp_degree = pp_degree ,
@@ -134,14 +155,24 @@ def test_pytorch_backend(original_outputs, dp_degree, hp_degree, pp_degree):
134155
135156
136157@pytest .mark .parametrize (
137- ("dp_degree" , "hp_degree" , "pp_degree" ),
138- list (itertools .product ([1 , 2 ], [1 , 2 ], [1 , 2 ])),
158+ ("dtype" , " dp_degree" , "hp_degree" , "pp_degree" ),
159+ list (itertools .product (["fp16" , "fp32" ], [ 1 , 2 ], [1 , 2 ], [1 , 2 ])),
139160)
140- def test_mixed_simulation (dp_degree , hp_degree , pp_degree ):
161+ def test_mixed_simulation (dtype , dp_degree , hp_degree , pp_degree ):
141162 _run_gpt (
163+ dtype = dtype ,
142164 dp_degree = dp_degree ,
143165 hp_degree = hp_degree ,
144166 pp_degree = pp_degree ,
145167 num_microbatches = pp_degree ,
146168 use_real_weights = False ,
147169 )
170+
171+ if __name__ == "__main__" :
172+ original_outputs = {
173+ "fp16" : _run_gpt (dtype = "fp16" , use_pytorch_backend = True ),
174+ "fp32" : _run_gpt (dtype = "fp32" , use_pytorch_backend = True ),
175+ }
176+ for dtype , dp_degree , hp_degree , pp_degree in list (itertools .product (["fp16" , "fp32" ], [1 , 2 ], [1 , 2 ], [1 , 2 ])):
177+ print (f"dtype={ dtype } , dp_degree={ dp_degree } , hp_degree={ hp_degree } , pp_degree={ pp_degree } " )
178+ test_pytorch_backend (original_outputs , dtype , dp_degree , hp_degree , pp_degree )
0 commit comments