/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */

/* Generated By CANNBot */

/*!
 * \file sparse_apply_ftrl_infershape.cpp
 * \brief InferShape implementation for sparse_apply_ftrl
 *
 * Output shape rules:
 *   var_out.shape = var.shape
 *   accum_out.shape = var.shape
 *   linear_out.shape = var.shape
 * Output dtype rules:
 *   All outputs have same dtype as input var (float32)
 *
 * Constraints validated:
 *   R-01: var/accum/linear must have the same shape
 *   R-02: grad.shape[0] == indices.shape[0], grad.shape[1:] == var.shape[1:]
 *   R-03: indices must be 1-D
 *   R-04: lr/l1/l2/lr_power must be scalar (0-D) or 1-D tensor with shape size 1
 *   R-07: var must be >= 2-D
 */

#include "register/op_impl_registry.h"
#include "log/log.h"
#include "graph/utils/type_utils.h"

using namespace ge;

namespace ops {

static constexpr int32_t IDX_VAR = 0;
static constexpr int32_t IDX_ACCUM = 1;
static constexpr int32_t IDX_LINEAR = 2;
static constexpr int32_t IDX_GRAD = 3;
static constexpr int32_t IDX_INDICES = 4;
static constexpr int32_t IDX_LR = 5;
static constexpr int32_t IDX_L1 = 6;
static constexpr int32_t IDX_L2 = 7;
static constexpr int32_t IDX_LR_POWER = 8;
static constexpr size_t MIN_VAR_DIM_NUM = 2;

static ge::graphStatus InferShapeSparseApplyFtrl(gert::InferShapeContext* context)
{
    const gert::Shape* varShape = context->GetInputShape(IDX_VAR);
    OP_CHECK_NULL_WITH_CONTEXT(context, varShape);
    const gert::Shape* accumShape = context->GetInputShape(IDX_ACCUM);
    OP_CHECK_NULL_WITH_CONTEXT(context, accumShape);
    const gert::Shape* linearShape = context->GetInputShape(IDX_LINEAR);
    OP_CHECK_NULL_WITH_CONTEXT(context, linearShape);
    const gert::Shape* gradShape = context->GetInputShape(IDX_GRAD);
    OP_CHECK_NULL_WITH_CONTEXT(context, gradShape);
    const gert::Shape* indicesShape = context->GetInputShape(IDX_INDICES);
    OP_CHECK_NULL_WITH_CONTEXT(context, indicesShape);
    const gert::Shape* lrShape = context->GetInputShape(IDX_LR);
    OP_CHECK_NULL_WITH_CONTEXT(context, lrShape);
    const gert::Shape* l1Shape = context->GetInputShape(IDX_L1);
    OP_CHECK_NULL_WITH_CONTEXT(context, l1Shape);
    const gert::Shape* l2Shape = context->GetInputShape(IDX_L2);
    OP_CHECK_NULL_WITH_CONTEXT(context, l2Shape);
    const gert::Shape* lrPowerShape = context->GetInputShape(IDX_LR_POWER);
    OP_CHECK_NULL_WITH_CONTEXT(context, lrPowerShape);

    auto nodeName = context->GetNodeName();

    // R-07: var must be >= 2-D
    auto varDimNum = varShape->GetDimNum();
    OP_CHECK_IF(varDimNum < MIN_VAR_DIM_NUM,
                OP_LOGE_FOR_INVALID_SHAPEDIM(nodeName, "var", std::to_string(varDimNum).c_str(), ">= 2"),
                return GRAPH_FAILED);

    // R-03: indices must be 1-D
    OP_CHECK_IF(
        indicesShape->GetDimNum() != 1,
        OP_LOGE_FOR_INVALID_SHAPEDIM(nodeName, "indices", std::to_string(indicesShape->GetDimNum()).c_str(), "1"),
        return GRAPH_FAILED);

    // R-01: var/accum/linear must have the same rank
    OP_CHECK_IF(accumShape->GetDimNum() != varDimNum || linearShape->GetDimNum() != varDimNum,
                OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
                    nodeName, "var, accum, linear",
                    (Ops::Base::ToString(*varShape) + ", " + Ops::Base::ToString(*accumShape) + ", " +
                     Ops::Base::ToString(*linearShape))
                        .c_str(),
                    "var, accum and linear must have the same rank"),
                return GRAPH_FAILED);

    // R-01: var/accum/linear must have the same shape (each dim)
    for (size_t i = 0; i < varDimNum; i++) {
        int64_t varDim = varShape->GetDim(i);
        OP_CHECK_IF(accumShape->GetDim(i) != varDim || linearShape->GetDim(i) != varDim,
                    OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
                        nodeName, "var, accum, linear",
                        (Ops::Base::ToString(*varShape) + ", " + Ops::Base::ToString(*accumShape) + ", " +
                         Ops::Base::ToString(*linearShape))
                            .c_str(),
                        "var, accum and linear must have the same shape"),
                    return GRAPH_FAILED);
    }

    // R-02: grad must have at least 1 dimension
    OP_CHECK_IF(gradShape->GetDimNum() < 1,
                OP_LOGE_FOR_INVALID_SHAPEDIM(nodeName, "grad", std::to_string(gradShape->GetDimNum()).c_str(), ">= 1"),
                return GRAPH_FAILED);

    // R-02: grad.shape[0] == indices.shape[0]
    int64_t numIndices = indicesShape->GetDim(0);
    OP_CHECK_IF(gradShape->GetDim(0) != numIndices,
                OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
                    nodeName, "grad, indices",
                    (Ops::Base::ToString(*gradShape) + ", " + Ops::Base::ToString(*indicesShape)).c_str(),
                    "grad.shape[0] must equal indices.shape[0]"),
                return GRAPH_FAILED);

    // R-02: grad and var must have the same rank
    OP_CHECK_IF(
        gradShape->GetDimNum() != varDimNum,
        OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
            nodeName, "grad, var", (Ops::Base::ToString(*gradShape) + ", " + Ops::Base::ToString(*varShape)).c_str(),
            "grad and var must have the same rank"),
        return GRAPH_FAILED);

    // R-02: grad.shape[1:] == var.shape[1:]
    for (size_t i = 1; i < varDimNum; i++) {
        OP_CHECK_IF(gradShape->GetDim(i) != varShape->GetDim(i),
                    OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
                        nodeName, "grad, var",
                        (Ops::Base::ToString(*gradShape) + ", " + Ops::Base::ToString(*varShape)).c_str(),
                        "grad.shape[1:] must equal var.shape[1:]"),
                    return GRAPH_FAILED);
    }

    // R-04: lr/l1/l2/lr_power must be scalar (0-D) or 1-D tensor with shape size 1
    const std::vector<std::pair<int32_t, const char*>> scalarInputs = {
        {IDX_LR, "lr"}, {IDX_L1, "l1"}, {IDX_L2, "l2"}, {IDX_LR_POWER, "lr_power"}};
    for (const auto& item : scalarInputs) {
        const gert::Shape* scalarShape = context->GetInputShape(item.first);
        OP_CHECK_NULL_WITH_CONTEXT(context, scalarShape);
        auto scalarDimNum = scalarShape->GetDimNum();
        auto scalarShapeSize = scalarShape->GetShapeSize();
        OP_CHECK_IF(
            scalarDimNum != 0 && scalarShapeSize != 1,
            OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(nodeName, item.second, Ops::Base::ToString(*scalarShape).c_str(),
                                                  "The param must be a scalar(0D) or have shape size 1"),
            return GRAPH_FAILED);
    }

    // Output 0: var_out (shape = var.shape)
    gert::Shape* varOutShape = context->GetOutputShape(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, varOutShape);
    varOutShape->SetDimNum(varDimNum);
    for (size_t i = 0; i < varDimNum; i++) {
        varOutShape->SetDim(i, varShape->GetDim(i));
    }

    // Output 1: accum_out (shape = var.shape)
    gert::Shape* accumOutShape = context->GetOutputShape(1);
    OP_CHECK_NULL_WITH_CONTEXT(context, accumOutShape);
    accumOutShape->SetDimNum(varDimNum);
    for (size_t i = 0; i < varDimNum; i++) {
        accumOutShape->SetDim(i, varShape->GetDim(i));
    }

    // Output 2: linear_out (shape = var.shape)
    gert::Shape* linearOutShape = context->GetOutputShape(2);
    OP_CHECK_NULL_WITH_CONTEXT(context, linearOutShape);
    linearOutShape->SetDimNum(varDimNum);
    for (size_t i = 0; i < varDimNum; i++) {
        linearOutShape->SetDim(i, varShape->GetDim(i));
    }

    return GRAPH_SUCCESS;
}

IMPL_OP_INFERSHAPE(SparseApplyFtrl).InferShape(InferShapeSparseApplyFtrl);
} // namespace ops