@@ -251,19 +251,6 @@ inline std::vector<int64_t> process_input_output_shape(
251251 return {shape.num_output_rows , columns};
252252}
253253
254- inline std::vector<int64_t > process_input_leading_shape (
255- const ProcessInputKernelData &data, const Tensor &inputs, const ProcessInputShape &shape) {
256- if (data.layout == 0 ) {
257- std::vector<int64_t > leading_shape;
258- leading_shape.reserve (inputs.dim () - 1 );
259- for (int64_t dimension = 0 ; dimension + 1 < inputs.dim (); dimension++)
260- leading_shape.push_back (inputs.size (dimension));
261- return leading_shape;
262- }
263- if (data.layout == 3 ) return {shape.num_experts , shape.max_tokens_per_expert };
264- return {shape.num_output_rows };
265- }
266-
267254inline Tensor prepare_process_input_output (
268255 const ProcessInputKernelData &data,
269256 const Tensor &inputs,
@@ -278,7 +265,7 @@ inline Tensor prepare_process_input_output(
278265 ASSERT_CHECK (!outputs.has_value () || outputs->data_ptr () == inputs.data_ptr (), " inplace output must alias inputs" );
279266 outputs = inputs;
280267 }
281- if (! outputs.has_value ()) return torch_empty (expected_shape, dtype, inputs. device () );
268+ ASSERT_CHECK ( outputs.has_value (), " outputs must be allocated by prepare_process_input " );
282269 check_process_input_tensor (*outputs, " outputs" , inputs.get_device (), dtype);
283270 ASSERT_CHECK (outputs->dim () == static_cast <int64_t >(expected_shape.size ()), " invalid output rank" );
284271 for (size_t dimension = 0 ; dimension < expected_shape.size (); dimension++)
@@ -302,18 +289,7 @@ inline std::optional<Tensor> prepare_process_input_group_scales(
302289 if (data.scale_layout == 2 ) elements = shape.group_scale_stride * CEIL_DIV (groups, 4 ) * 4 ;
303290 ScalarType dtype = dtype_id_to_tensor_dtype (data.group_scale_dtype_id );
304291 bool allow_byte = data.group_scale_dtype_id == 20080800 ;
305- if (!scales.has_value ()) {
306- std::vector<int64_t > scale_shape;
307- if (data.scale_layout == 0 ) {
308- scale_shape = process_input_leading_shape (data, inputs, shape);
309- scale_shape.push_back (groups);
310- } else if (data.scale_layout == 1 ) {
311- scale_shape = {groups, shape.group_scale_stride };
312- } else {
313- scale_shape = {CEIL_DIV (groups, 4 ), shape.group_scale_stride , 4 };
314- }
315- return torch_empty (scale_shape, dtype, inputs.device ());
316- }
292+ ASSERT_CHECK (scales.has_value (), " group_scales must be allocated by prepare_process_input" );
317293 check_process_input_tensor (*scales, " group_scales" , inputs.get_device (), dtype, allow_byte);
318294 ASSERT_CHECK (scales->numel () == elements, " invalid group_scales size" );
319295 return scales;
@@ -331,11 +307,7 @@ inline std::optional<Tensor> prepare_process_input_token_scales(
331307 return std::nullopt ;
332308 }
333309 int64_t elements = static_scale ? shape.num_experts : shape.num_output_rows ;
334- if (!scales.has_value ()) {
335- ASSERT_CHECK (dynamic_scale, " static token_scales is required" );
336- auto scale_shape = process_input_leading_shape (data, inputs, shape);
337- return torch_empty (scale_shape, ScalarType::Float, inputs.device ());
338- }
310+ ASSERT_CHECK (scales.has_value (), " token_scales must be allocated by prepare_process_input" );
339311 check_process_input_tensor (*scales, " token_scales" , inputs.get_device (), ScalarType::Float);
340312 ASSERT_CHECK (scales->numel () == elements, " invalid token_scales size" );
341313 return scales;
@@ -457,7 +429,7 @@ inline void launch_process_input_finalizer(
457429 check_curesult (cuLaunchKernelEx (&config, data.func , kernel_args, nullptr ), " cuLaunchKernelEx" );
458430}
459431
460- inline std::tuple<Tensor, std::optional<Tensor>, std::optional<Tensor>> launch_process_input (
432+ inline void launch_process_input_impl (
461433 Tensor configs_tensor,
462434 Tensor inputs,
463435 std::optional<Tensor> outputs,
@@ -524,5 +496,29 @@ inline std::tuple<Tensor, std::optional<Tensor>, std::optional<Tensor>> launch_p
524496 secondary, intermediate, *group_scales, *token_scales, expert_layout, indices, shape);
525497 }
526498 }
527- return std::make_tuple (prepared_output, group_scales, token_scales);
499+ }
500+
501+ inline void launch_process_input (
502+ Tensor configs_tensor,
503+ Tensor inputs,
504+ Tensor outputs,
505+ std::optional<Tensor> group_scales,
506+ std::optional<Tensor> token_scales,
507+ std::optional<Tensor> expert_layout,
508+ std::optional<Tensor> indices) {
509+ if (!inputs.is_cuda ()) return ;
510+ launch_process_input_impl (
511+ configs_tensor, inputs, outputs, group_scales, token_scales,
512+ expert_layout, indices, false );
513+ }
514+
515+ inline void launch_process_input_inplace (
516+ Tensor configs_tensor,
517+ Tensor inputs,
518+ std::optional<Tensor> expert_layout,
519+ std::optional<Tensor> indices) {
520+ if (!inputs.is_cuda ()) return ;
521+ launch_process_input_impl (
522+ configs_tensor, inputs, inputs, std::nullopt , std::nullopt ,
523+ expert_layout, indices, true );
528524}
0 commit comments