/* Copyright (c) 2023 PaddlePaddle Authors. All Rights Reserved. Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the License. You may obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0 Unless required by applicable law or agreed to in writing, software distributed under the License is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions and limitations under the License. */ #include "paddle/phi/infermeta/spmd_rules/rules.h" /** * Design Notes: * * 1. SPMD info is the special meta info of DistTensor, so we put Spmd infer * functions in `infermeta` directory. * * 2. Since the infer functions of Spmd forward and backward are closely related * and need to be registered together, we manage them together in one file. * * 3. SPMD rules are much smaller than infermeta function, and we manage files * in operator units. * * 4. The previous registration used some compile-time regular matching methods, * which was less flexible, and the registration of SPMD rules here is declare * directly in the header file */ namespace phi::distributed { // matmul rule PD_REGISTER_SPMD_RULE(matmul, PD_INFER_SPMD(MatmulInferSpmd), PD_INFER_SPMD(MatmulInferSpmdReverse)); PD_REGISTER_SPMD_RULE(matmul_v2, // static mode PD_INFER_SPMD(MatmulInferSpmd), PD_INFER_SPMD(MatmulInferSpmdReverse)); PD_REGISTER_SPMD_RULE(bmm, PD_INFER_SPMD(BmmInferSpmd), PD_INFER_SPMD(BmmGradInferSpmd)); PD_REGISTER_SPMD_RULE(elementwise_unary, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(elementwise_binary, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); // default data parallel rule PD_REGISTER_SPMD_RULE(default_data_parallel, PD_INFER_SPMD(DefaultDataParallelInferSpmd), PD_INFER_SPMD(DefaultDataParallelInferSpmdReverse)); PD_REGISTER_SPMD_RULE(default_, PD_INFER_SPMD(DefaultDataParallelInferSpmd), PD_INFER_SPMD(DefaultDataParallelInferSpmdReverse)); // fused rope PD_REGISTER_SPMD_RULE(fused_rotary_position_embedding, PD_INFER_SPMD(FusedRopeInferSpmd), PD_INFER_SPMD(FusedRopeInferSpmdReverse)); // replicated rule /* for unittest */ PD_REGISTER_SPMD_RULE(replicated, PD_INFER_SPMD(ReplicatedInferSpmd), PD_INFER_SPMD(ReplicatedInferSpmdReverse)); // unsqueeze rule PD_REGISTER_SPMD_RULE(unsqueeze, PD_INFER_SPMD(UnsqueezeInferSpmd), PD_INFER_SPMD(UnsqueezeInferSpmdReverse)); PD_REGISTER_SPMD_RULE(unsqueeze2, PD_INFER_SPMD(UnsqueezeInferSpmd), PD_INFER_SPMD(UnsqueezeInferSpmdReverse)); // elementwise unary rule PD_REGISTER_SPMD_RULE(abs, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(assign, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(hardswish, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(mish, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(relu6, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(swish, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(acos, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(acosh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(asin, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(asinh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(atan, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(atanh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(bernoulli, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(bitwise_not, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(ceil, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(celu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(clip, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(conj, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(conv2d, PD_INFER_SPMD(Conv2dInferSpmdBase), PD_INFER_SPMD(Conv2dInferSpmdReverseBase)); PD_REGISTER_SPMD_RULE(cos, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(cosh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(digamma, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(elu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(erf, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(erfinv, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(exp, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(expm1, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(fill, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(floor, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(hardshrink, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(hardsigmoid, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(hardtanh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(label_smooth, PD_INFER_SPMD(LabelSmoothInferSpmd), PD_INFER_SPMD(LabelSmoothGradInferSpmd)); PD_REGISTER_SPMD_RULE(leaky_relu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(lgamma, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(log, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(log10, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(log1p, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(log2, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(logical_not, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(logit, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(logsigmoid, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(poisson, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(pow, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(reciprocal, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(relu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(round, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(rsqrt, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(scale, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(selu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(sigmoid, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(sign, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(silu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(sin, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(sinh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(softplus, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(softshrink, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(softsign, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(sqrt, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(square, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(stanh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(tan, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(tanh, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(tanh_shrink, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(thresholded_relu, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(trunc, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(dropout, PD_INFER_SPMD(ElementwiseUnaryInferSpmd), PD_INFER_SPMD(ElementwiseUnaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(fused_dropout_add, PD_INFER_SPMD(FusedDropoutAddSpmdBase), PD_INFER_SPMD(FusedDropoutAddSpmdReverseBase)); // elementwise binary rule PD_REGISTER_SPMD_RULE(add, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(elementwise_add, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(divide, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(elementwise_div, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(elementwise_pow, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(floor_divide, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(trunc_divide, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(fmin, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(heaviside, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(maximum, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(minimum, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(multiply, PD_INFER_SPMD(ElementwiseBinaryWithPartialInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(elementwise_mul, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(remainder, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(subtract, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(bitwise_and, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(bitwise_or, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(bitwise_xor, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(fmax, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(logical_and, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(logical_or, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(logical_xor, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(not_equal, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(greater_than, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(less_than, PD_INFER_SPMD(ElementwiseBinaryInferSpmd), PD_INFER_SPMD(ElementwiseBinaryInferSpmdReverse)); PD_REGISTER_SPMD_RULE(swiglu, PD_INFER_SPMD(SwiGLUInferSpmd), PD_INFER_SPMD(SwiGLUInferSpmdReverse)); // TODO(pkuzyc): add multiary elementwise rule // reduction rule PD_REGISTER_SPMD_RULE(reduce_base, PD_INFER_SPMD(ReductionInferSpmdBase), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(all, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(amax, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(amin, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(any, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(frobenius_norm, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(max, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(reduce_max, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(min, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(prod, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(sum, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(reduce_sum, // static PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); PD_REGISTER_SPMD_RULE(squared_l2_norm, PD_INFER_SPMD(ReductionInferSpmd), PD_INFER_SPMD(ReductionInferSpmdReverse)); // mean_all PD_REGISTER_SPMD_RULE(mean_all, PD_INFER_SPMD(MeanAllInferSpmd), PD_INFER_SPMD(MeanAllGradInferSpmd)); // batch_norm PD_REGISTER_SPMD_RULE(batch_norm, PD_INFER_SPMD(BatchNormInferSpmdStatic)); // layer_norm PD_REGISTER_SPMD_RULE(layer_norm, PD_INFER_SPMD(LayerNormInferSpmd), PD_INFER_SPMD(LayerNormInferSpmdReverse)); // instance_norm PD_REGISTER_SPMD_RULE(instance_norm, PD_INFER_SPMD(InstanceNormInferSpmd), PD_INFER_SPMD(InstanceNormGradInferSpmd)); // fused_rms_norm // NOTE(ZHIQIU): Temporally register fused_rms_norm rule, // this is not for rms_norm kernel, but for the custom kernel // 'fused_rms_norm' in PaddleNLP. // It will be no longer needed when the PIR-AutoParallel project // is finished. PD_REGISTER_SPMD_RULE(fused_rms_norm, PD_INFER_SPMD(RmsNormInferSpmd), PD_INFER_SPMD(RmsNormInferSpmdReverse)); // index_put PD_REGISTER_SPMD_RULE(index_put, PD_INFER_SPMD(IndexPutInferSpmd), PD_INFER_SPMD(IndexPutGradInferSpmd)); PD_REGISTER_SPMD_RULE(flash_attention, PD_INFER_SPMD(FlashAttInferSpmdStatic), PD_INFER_SPMD(FlashAttInferSpmdReverse)); // reshape rule PD_REGISTER_SPMD_RULE(reshape, PD_INFER_SPMD(ReshapeInferSpmd), PD_INFER_SPMD(ReshapeInferSpmdReverse)); PD_REGISTER_SPMD_RULE(reshape2, PD_INFER_SPMD(ReshapeInferSpmd), PD_INFER_SPMD(ReshapeInferSpmdReverse)); // squeeze rule PD_REGISTER_SPMD_RULE(squeeze, PD_INFER_SPMD(SqueezeInferSpmd), PD_INFER_SPMD(SqueezeInferSpmdReverse)); // flatten rule PD_REGISTER_SPMD_RULE(flatten, PD_INFER_SPMD(FlattenInferSpmd), PD_INFER_SPMD(FlattenInferSpmdReverse)); // embedding rule PD_REGISTER_SPMD_RULE(embedding, PD_INFER_SPMD(EmbeddingInferSpmd), PD_INFER_SPMD(EmbeddingInferSpmdReverse)); PD_REGISTER_SPMD_RULE(c_embedding, PD_INFER_SPMD(CEmbeddingInferSpmd), PD_INFER_SPMD(CEmbeddingGradInferSpmd)); PD_REGISTER_SPMD_RULE(lookup_table_v2, PD_INFER_SPMD(EmbeddingInferSpmd), PD_INFER_SPMD(EmbeddingInferSpmdReverse)); // split rule PD_REGISTER_SPMD_RULE(split, PD_INFER_SPMD(SplitInferSpmd), PD_INFER_SPMD(SplitInferSpmdReverse)); PD_REGISTER_SPMD_RULE(split_with_num, PD_INFER_SPMD(SplitWithNumInferSpmd), PD_INFER_SPMD(SplitWithNumInferSpmdReverse)); // slice rule PD_REGISTER_SPMD_RULE(slice, PD_INFER_SPMD(SliceInferSpmd), PD_INFER_SPMD(SliceInferSpmdReverse)); PD_REGISTER_SPMD_RULE(strided_slice, PD_INFER_SPMD(StridedSliceInferSpmd), PD_INFER_SPMD(StridedSliceGradInferSpmd)); PD_REGISTER_SPMD_RULE(concat, PD_INFER_SPMD(ConcatInferSpmd), PD_INFER_SPMD(ConcatInferSpmdReverse)); PD_REGISTER_SPMD_RULE(stack, PD_INFER_SPMD(StackInferSpmd), PD_INFER_SPMD(StackInferSpmdReverse)); // transpose rule PD_REGISTER_SPMD_RULE(transpose, PD_INFER_SPMD(TransposeInferSpmd), PD_INFER_SPMD(TransposeInferSpmdReverse)); PD_REGISTER_SPMD_RULE(transpose2, PD_INFER_SPMD(TransposeInferSpmd), PD_INFER_SPMD(TransposeInferSpmdReverse)); // softmax rule PD_REGISTER_SPMD_RULE(softmax, PD_INFER_SPMD(SoftmaxInferSpmd), PD_INFER_SPMD(SoftmaxInferSpmdReverse)); PD_REGISTER_SPMD_RULE(log_softmax, PD_INFER_SPMD(SoftmaxInferSpmd), PD_INFER_SPMD(SoftmaxInferSpmdReverse)); PD_REGISTER_SPMD_RULE(where, PD_INFER_SPMD(WhereInferSpmd), PD_INFER_SPMD(WhereInferSpmdReverse)); PD_REGISTER_SPMD_RULE(triu, PD_INFER_SPMD(TriuInferSpmd), PD_INFER_SPMD(TriuInferSpmdReverse)); PD_REGISTER_SPMD_RULE(tril_triu, PD_INFER_SPMD(TrilTriuInferSpmd), PD_INFER_SPMD(TrilTriuInferSpmdReverse)); PD_REGISTER_SPMD_RULE(tile, PD_INFER_SPMD(TileInferSpmd), PD_INFER_SPMD(TileInferSpmdReverse)); // cross_entropy_with_softmax PD_REGISTER_SPMD_RULE(cross_entropy_with_softmax, PD_INFER_SPMD(CrossEntropyWithSoftmaxInferSpmdStatic), PD_INFER_SPMD(CrossEntropyWithSoftmaxInferSpmdReverse)); PD_REGISTER_SPMD_RULE(softmax_with_cross_entropy, PD_INFER_SPMD(CrossEntropyWithSoftmaxInferSpmdStatic), PD_INFER_SPMD(CrossEntropyWithSoftmaxInferSpmdReverse)); PD_REGISTER_SPMD_RULE(c_softmax_with_cross_entropy, PD_INFER_SPMD(CSoftmaxWithCrossEntropyInferSpmd)); PD_REGISTER_SPMD_RULE( c_softmax_with_multi_label_cross_entropy, PD_INFER_SPMD(CSoftmaxWithMultiLabelCrossEntropyInferSpmd)); // fused_linear_param_grad_add got no reverse infer spmd rule PD_REGISTER_SPMD_RULE( fused_linear_param_grad_add, PD_INFER_SPMD(FusedLinearParamGradAddInferSpmd), PD_INFER_SPMD(FusedLinearParamGradAddInferSpmdFakeReverse)); PD_REGISTER_SPMD_RULE(expand_as, PD_INFER_SPMD(ExpandAsInferSpmd), PD_INFER_SPMD(ExpandAsGradInferSpmd)); PD_REGISTER_SPMD_RULE(expand_as_v2, PD_INFER_SPMD(ExpandAsInferSpmd), PD_INFER_SPMD(ExpandAsGradInferSpmd)); // scatter PD_REGISTER_SPMD_RULE(scatter, PD_INFER_SPMD(ScatterInferSpmd), PD_INFER_SPMD(ScatterInferSpmdReverse)); // scatter_nd_add PD_REGISTER_SPMD_RULE(scatter_nd_add, PD_INFER_SPMD(ScatterNdAddInferSpmd), PD_INFER_SPMD(ScatterNdAddInferSpmdReverse)); // gather PD_REGISTER_SPMD_RULE(gather, PD_INFER_SPMD(GatherInferSpmdBase), PD_INFER_SPMD(GatherInferSpmdReverseBase)); PD_REGISTER_SPMD_RULE(gather_nd, PD_INFER_SPMD(GatherNdInferSpmd), PD_INFER_SPMD(GatherNdInferSpmdReverse)); // gelu PD_REGISTER_SPMD_RULE(gelu, PD_INFER_SPMD(GeluInferSpmd), PD_INFER_SPMD(GeluGradInferSpmd)); // one_hot PD_REGISTER_SPMD_RULE(one_hot, PD_INFER_SPMD(OneHotInferSpmd), PD_INFER_SPMD(OneHotInferSpmdReverse)); PD_REGISTER_SPMD_RULE(cumsum, PD_INFER_SPMD(CumSumInferSpmd), PD_INFER_SPMD(CumSumInferSpmdReverse)); // unique PD_REGISTER_SPMD_RULE(unique, PD_INFER_SPMD(UniqueInferSpmd)); // argmin PD_REGISTER_SPMD_RULE(argmin, PD_INFER_SPMD(ArgMinInferSpmdBase), PD_INFER_SPMD(ArgMinInferSpmdReverseBase)); // argmax PD_REGISTER_SPMD_RULE(argmax, PD_INFER_SPMD(ArgMaxInferSpmdBase), PD_INFER_SPMD(ArgMaxInferSpmdReverseBase)); // topk PD_REGISTER_SPMD_RULE(topk, PD_INFER_SPMD(TopkInferSpmd), PD_INFER_SPMD(TopkGradInferSpmd)); // unbind PD_REGISTER_SPMD_RULE(unbind, PD_INFER_SPMD(UnbindInferSpmd), PD_INFER_SPMD(UnbindInferSpmdReverse)); // logsumexp PD_REGISTER_SPMD_RULE(logsumexp, PD_INFER_SPMD(LogSumExpInferSpmd), PD_INFER_SPMD(LogSumExpInferSpmdReverse)); // p_norm PD_REGISTER_SPMD_RULE(p_norm, PD_INFER_SPMD(PNormInferSpmd), PD_INFER_SPMD(PNormInferSpmdReverse)); // pad PD_REGISTER_SPMD_RULE(pad, PD_INFER_SPMD(PadInferSpmd), PD_INFER_SPMD(PadGradInferSpmd)); // group_norm PD_REGISTER_SPMD_RULE(group_norm, PD_INFER_SPMD(GroupNormInferSpmdBase)); // nonzero PD_REGISTER_SPMD_RULE(nonzero, PD_INFER_SPMD(NonZeroInferSpmd), PD_INFER_SPMD(NonZeroInferSpmdReverse)); // add_n PD_REGISTER_SPMD_RULE(add_n, PD_INFER_SPMD(AddNInferSpmd)); // roll PD_REGISTER_SPMD_RULE(roll, PD_INFER_SPMD(RollInferSpmd), PD_INFER_SPMD(RollGradInferSpmd)); // cummax PD_REGISTER_SPMD_RULE(cummax, PD_INFER_SPMD(CummaxInferSpmd), PD_INFER_SPMD(CummaxGradInferSpmd)); // cummin PD_REGISTER_SPMD_RULE(cummin, PD_INFER_SPMD(CumminInferSpmd), PD_INFER_SPMD(CumminGradInferSpmd)); // argsort PD_REGISTER_SPMD_RULE(argsort, PD_INFER_SPMD(ArgSortInferSpmd), PD_INFER_SPMD(ArgSortGradInferSpmd)); // index_select PD_REGISTER_SPMD_RULE(index_select, PD_INFER_SPMD(IndexSelectInferSpmd), PD_INFER_SPMD(IndexSelectGradInferSpmd)); // put_along_axis PD_REGISTER_SPMD_RULE(put_along_axis, PD_INFER_SPMD(PutAlongAxisInferSpmd), PD_INFER_SPMD(PutAlongAxisGradInferSpmd)); // roi_align PD_REGISTER_SPMD_RULE(roi_align, PD_INFER_SPMD(RoiAlignInferSpmd), PD_INFER_SPMD(RoiAlignGradInferSpmd)); // fused gemm epilogue PD_REGISTER_SPMD_RULE(fused_gemm_epilogue, PD_INFER_SPMD(FusedGemmEpilogueInferSpmdBase)); // linear_v2 PD_REGISTER_SPMD_RULE(linear_v2, PD_INFER_SPMD(LinearV2InferSpmdBase)); // take_along_axis PD_REGISTER_SPMD_RULE(take_along_axis, PD_INFER_SPMD(TakeAlongAxisInferSpmd), PD_INFER_SPMD(TakeAlongAxisGradInferSpmd)); // conv3d PD_REGISTER_SPMD_RULE(conv3d, PD_INFER_SPMD(Conv3dInferSpmd), PD_INFER_SPMD(Conv3dGradInferSpmd)); // depthwise_conv2d PD_REGISTER_SPMD_RULE(depthwise_conv2d, PD_INFER_SPMD(DepthwiseConv2dInferSpmd), PD_INFER_SPMD(DepthwiseConv2dGradInferSpmd)); // conv2d_transpose PD_REGISTER_SPMD_RULE(conv2d_transpose, PD_INFER_SPMD(Conv2dTransposeInferSpmd), PD_INFER_SPMD(Conv2dTransposeGradInferSpmd)); // einsum PD_REGISTER_SPMD_RULE(einsum, PD_INFER_SPMD(EinsumInferSpmd), PD_INFER_SPMD(EinsumGradInferSpmd)); // moe_gate_dispatch PD_REGISTER_SPMD_RULE(moe_gate_dispatch, PD_INFER_SPMD(MoEGateDispatchInferSpmd), PD_INFER_SPMD(MoEGateDispatchGradInferSpmd)); // moe_combine PD_REGISTER_SPMD_RULE(moe_combine, PD_INFER_SPMD(MoECombineInferSpmd), PD_INFER_SPMD(MoECombineGradInferSpmd)); } // namespace phi::distributed