/*
 * SPDX-FileCopyrightText: Copyright 2025 the Regents of the University of California, Nerfstudio Team and contributors. All rights reserved.
 * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
 * SPDX-License-Identifier: Apache-2.0
 *
 * 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.
 */

#pragma once

#include <cstdint>
#include <tuple>

#include <ATen/core/Tensor.h>

#include "Cameras.h"
#include "Common.h"
#include "Lidars.h"
#include "ExternalDistortion.h"
#include "Cameras.cuh"
#include "TorchUtils.h"

namespace at
{
class Tensor;
}

namespace gsplat
{
struct Projection2DGSFusedResult
{
    at::Tensor radii;
    at::Tensor means2d;
    at::Tensor depths;
    at::Tensor ray_transforms;
    at::Tensor normals;
};

Projection2DGSFusedResult projection_2dgs_fused(
    const at::Tensor &means,
    const at::Tensor &quats,
    const at::Tensor &scales,
    const at::Tensor &viewmats,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double eps2d,
    double near_plane,
    double far_plane,
    double radius_clip
);

struct Projection2DGSPackedResult
{
    at::Tensor batch_ids;
    at::Tensor camera_ids;
    at::Tensor gaussian_ids;
    at::Tensor indptr;
    at::Tensor radii;
    at::Tensor means2d;
    at::Tensor depths;
    at::Tensor ray_transforms;
    at::Tensor normals;
};

Projection2DGSPackedResult projection_2dgs_packed(
    const at::Tensor &means,
    const at::Tensor &quats,
    const at::Tensor &scales,
    const at::Tensor &viewmats,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double near_plane,
    double far_plane,
    double radius_clip,
    bool sparse_grad
);

struct ProjectionEWA3DGSFusedFwdResult
{
    at::Tensor radii;
    at::Tensor means2d;
    at::Tensor depths;
    at::Tensor conics;
    at::optional<at::Tensor> compensations;
};

using ProjectionEWA3DGSFusedResult = ProjectionEWA3DGSFusedFwdResult;

ProjectionEWA3DGSFusedFwdResult projection_ewa_3dgs_fused_fwd(
    const at::Tensor &means,
    const at::optional<at::Tensor> &covars,
    const at::optional<at::Tensor> &quats,
    const at::optional<at::Tensor> &scales,
    const at::optional<at::Tensor> &opacities,
    const at::Tensor &viewmats,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double eps2d,
    double near_plane,
    double far_plane,
    double radius_clip,
    bool calc_compensations,
    CameraModelType camera_model
);
ProjectionEWA3DGSFusedResult projection_ewa_3dgs_fused(
    const at::Tensor &means,
    const at::optional<at::Tensor> &covars,
    const at::optional<at::Tensor> &quats,
    const at::optional<at::Tensor> &scales,
    const at::optional<at::Tensor> &opacities,
    const at::Tensor &viewmats,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double eps2d,
    double near_plane,
    double far_plane,
    double radius_clip,
    bool calc_compensations,
    CameraModelType camera_model
);

struct ProjectionEWA3DGSPackedFwdResult
{
    at::Tensor batch_ids;
    at::Tensor camera_ids;
    at::Tensor gaussian_ids;
    at::Tensor indptr;
    at::Tensor radii;
    at::Tensor means2d;
    at::Tensor depths;
    at::Tensor conics;
    at::optional<at::Tensor> compensations;
};

using ProjectionEWA3DGSPackedResult = ProjectionEWA3DGSPackedFwdResult;

ProjectionEWA3DGSPackedFwdResult projection_ewa_3dgs_packed_fwd(
    const at::Tensor &means,
    const at::optional<at::Tensor> &covars,
    const at::optional<at::Tensor> &quats,
    const at::optional<at::Tensor> &scales,
    const at::optional<at::Tensor> &opacities,
    const at::Tensor &viewmats,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double eps2d,
    double near_plane,
    double far_plane,
    double radius_clip,
    bool sparse_grad,
    bool calc_compensations,
    CameraModelType camera_model
);
ProjectionEWA3DGSPackedResult projection_ewa_3dgs_packed(
    const at::Tensor &means,
    const at::optional<at::Tensor> &covars,
    const at::optional<at::Tensor> &quats,
    const at::optional<at::Tensor> &scales,
    const at::optional<at::Tensor> &opacities,
    const at::Tensor &viewmats,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double eps2d,
    double near_plane,
    double far_plane,
    double radius_clip,
    bool sparse_grad,
    bool calc_compensations,
    CameraModelType camera_model
);

struct ProjectionUT3DGSFusedResult
{
    at::Tensor radii;
    at::Tensor means2d;
    at::Tensor depths;
    at::Tensor conics;
    at::optional<at::Tensor> compensations;
};

template<>
struct TorchArgDef<ProjectionUT3DGSFusedResult>
{
    static auto to(const ProjectionUT3DGSFusedResult &r)
    {
        return to_torch_args(r.radii, r.means2d, r.depths, r.conics, r.compensations);
    }

    template<class TT>
    static ProjectionUT3DGSFusedResult from(TT &&t)
    {
        return {
            .radii         = std::get<0>(std::forward<TT>(t)),
            .means2d       = std::get<1>(std::forward<TT>(t)),
            .depths        = std::get<2>(std::forward<TT>(t)),
            .conics        = std::get<3>(std::forward<TT>(t)),
            .compensations = std::get<4>(std::forward<TT>(t)),
        };
    }
};

ProjectionUT3DGSFusedResult projection_ut_3dgs_fused(
    const at::Tensor &means,
    const at::Tensor &quats,
    const at::Tensor &scales,
    const at::optional<at::Tensor> &opacities,
    const at::Tensor &viewmats0,
    const at::optional<at::Tensor> &viewmats1,
    const at::Tensor &Ks,
    int64_t image_width,
    int64_t image_height,
    double eps2d,
    double near_plane,
    double far_plane,
    double radius_clip,
    bool calc_compensations,
    CameraModelType camera_model,
    bool global_z_order,
    const at::optional<c10::intrusive_ptr<UnscentedTransformParameters>> &ut_params,
    ShutterType rs_type,
    const at::optional<at::Tensor> &radial_coeffs,
    const at::optional<at::Tensor> &tangential_coeffs,
    const at::optional<at::Tensor> &thin_prism_coeffs,
    const at::optional<c10::intrusive_ptr<FThetaCameraDistortionParameters>> &ftheta_coeffs,
    const at::optional<c10::intrusive_ptr<RowOffsetStructuredSpinningLidarModelParametersExt>> &lidar_coeffs,
    const at::optional<c10::intrusive_ptr<extdist::BivariateWindshieldModelParameters>> &external_distortion_params
);

void launch_projection_ewa_simple_fwd_kernel(
    // inputs
    const at::Tensor means,  // [..., C, N, 3]
    const at::Tensor covars, // [..., C, N, 3, 3]
    const at::Tensor Ks,     // [..., C, 3, 3]
    const uint32_t width,
    const uint32_t height,
    const CameraModelType camera_model,
    // outputs
    at::Tensor means2d, // [..., C, N, 2]
    at::Tensor covars2d // [..., C, N, 2, 2]
);
void launch_projection_ewa_simple_bwd_kernel(
    // inputs
    const at::Tensor means,  // [..., C, N, 3]
    const at::Tensor covars, // [..., C, N, 3, 3]
    const at::Tensor Ks,     // [..., C, 3, 3]
    const uint32_t width,
    const uint32_t height,
    const CameraModelType camera_model,
    const at::Tensor v_means2d,  // [..., C, N, 2]
    const at::Tensor v_covars2d, // [..., C, N, 2, 2]
    // outputs
    at::Tensor v_means, // [..., C, N, 3]
    at::Tensor v_covars // [..., C, N, 3, 3]
);

void launch_projection_ewa_3dgs_fused_fwd_kernel(
    // inputs
    const at::Tensor means,                   // [..., N, 3]
    const at::optional<at::Tensor> covars,    // [..., N, 6] optional
    const at::optional<at::Tensor> quats,     // [..., N, 4] optional
    const at::optional<at::Tensor> scales,    // [..., N, 3] optional
    const at::optional<at::Tensor> opacities, // [..., N] optional
    const at::Tensor viewmats,                // [..., C, 4, 4]
    const at::Tensor Ks,                      // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const float near_plane,
    const float far_plane,
    const float radius_clip,
    const CameraModelType camera_model,
    // outputs
    at::Tensor radii,                      // [..., C, N, 2]
    at::Tensor means2d,                    // [..., C, N, 2]
    at::Tensor depths,                     // [..., C, N]
    at::Tensor conics,                     // [..., C, N, 3]
    at::optional<at::Tensor> compensations // [..., C, N] optional
);

void launch_projection_ewa_3dgs_fused_fwd_kernels(
    // inputs
    const at::Tensor means,                   // [..., N, 3]
    const at::optional<at::Tensor> covars,    // [..., N, 6] optional
    const at::optional<at::Tensor> quats,     // [..., N, 4] optional
    const at::optional<at::Tensor> scales,    // [..., N, 3] optional
    const at::optional<at::Tensor> opacities, // [..., N] optional
    const at::Tensor viewmats,                // [..., C, 4, 4]
    const at::Tensor Ks,                      // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const float near_plane,
    const float far_plane,
    const float radius_clip,
    const CameraModelType camera_model,
    // outputs
    at::Tensor radii,                      // [..., C, N, 2]
    at::Tensor means2d,                    // [..., C, N, 2]
    at::Tensor depths,                     // [..., C, N]
    at::Tensor conics,                     // [..., C, N, 3]
    at::optional<at::Tensor> compensations // [..., C, N] optional
);

void launch_projection_ewa_3dgs_fused_bwd_kernel(
    // inputs
    // fwd inputs
    const at::Tensor means,                // [..., N, 3]
    const at::optional<at::Tensor> covars, // [..., N, 6] optional
    const at::optional<at::Tensor> quats,  // [..., N, 4] optional
    const at::optional<at::Tensor> scales, // [..., N, 3] optional
    const at::Tensor viewmats,             // [..., C, 4, 4]
    const at::Tensor Ks,                   // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const CameraModelType camera_model,
    // fwd outputs
    const at::Tensor radii,                       // [..., C, N, 2]
    const at::Tensor conics,                      // [..., C, N, 3]
    const at::optional<at::Tensor> compensations, // [..., C, N] optional
    // grad outputs
    const at::Tensor v_means2d,                     // [..., C, N, 2]
    const at::Tensor v_depths,                      // [..., C, N]
    const at::Tensor v_conics,                      // [..., C, N, 3]
    const at::optional<at::Tensor> v_compensations, // [..., C, N] optional
    const bool viewmats_requires_grad,
    // outputs
    at::Tensor v_means,   // [..., N, 3]
    at::Tensor v_covars,  // [..., N, 3, 3]
    at::Tensor v_quats,   // [..., N, 4]
    at::Tensor v_scales,  // [..., N, 3]
    at::Tensor v_viewmats // [..., C, 4, 4]
);

void launch_projection_ewa_3dgs_fused_bwd_kernels(
    // inputs
    // fwd inputs
    const at::Tensor means,                // [..., N, 3]
    const at::optional<at::Tensor> covars, // [..., N, 6] optional
    const at::optional<at::Tensor> quats,  // [..., N, 4] optional
    const at::optional<at::Tensor> scales, // [..., N, 3] optional
    const at::Tensor viewmats,             // [..., C, 4, 4]
    const at::Tensor Ks,                   // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const CameraModelType camera_model,
    // fwd outputs
    const at::Tensor radii,                       // [..., C, N, 2]
    const at::Tensor conics,                      // [..., C, N, 3]
    const at::optional<at::Tensor> compensations, // [..., C, N] optional
    // grad outputs
    const at::Tensor v_means2d,                     // [..., C, N, 2]
    const at::Tensor v_depths,                      // [..., C, N]
    const at::Tensor v_conics,                      // [..., C, N, 3]
    const at::optional<at::Tensor> v_compensations, // [..., C, N] optional
    const bool viewmats_requires_grad,
    // outputs
    at::Tensor v_means,   // [..., N, 3]
    at::Tensor v_covars,  // [..., N, 3, 3]
    at::Tensor v_quats,   // [..., N, 4]
    at::Tensor v_scales,  // [..., N, 3]
    at::Tensor v_viewmats // [..., C, 4, 4]
);

void launch_projection_ewa_3dgs_packed_fwd_kernel(
    // inputs
    const at::Tensor means,                   // [..., N, 3]
    const at::optional<at::Tensor> covars,    // [..., N, 6] optional
    const at::optional<at::Tensor> quats,     // [..., N, 4] optional
    const at::optional<at::Tensor> scales,    // [..., N, 3] optional
    const at::optional<at::Tensor> opacities, // [..., N] optional
    const at::Tensor viewmats,                // [..., C, 4, 4]
    const at::Tensor Ks,                      // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const float near_plane,
    const float far_plane,
    const float radius_clip,
    const at::optional<at::Tensor> block_accum, // [B * C * blocks_per_row] packing helper
    const CameraModelType camera_model,
    // outputs
    at::optional<at::Tensor> block_cnts,   // [B * C * blocks_per_row] packing helper
    at::optional<at::Tensor> indptr,       // [B * C + 1]
    at::optional<at::Tensor> batch_ids,    // [nnz]
    at::optional<at::Tensor> camera_ids,   // [nnz]
    at::optional<at::Tensor> gaussian_ids, // [nnz]
    at::optional<at::Tensor> radii,        // [nnz, 2]
    at::optional<at::Tensor> means2d,      // [nnz, 2]
    at::optional<at::Tensor> depths,       // [nnz]
    at::optional<at::Tensor> conics,       // [nnz, 3]
    at::optional<at::Tensor> compensations // [nnz] optional
);
void launch_projection_ewa_3dgs_packed_bwd_kernel(
    // fwd inputs
    const at::Tensor means,                // [..., N, 3]
    const at::optional<at::Tensor> covars, // [..., N, 6]
    const at::optional<at::Tensor> quats,  // [..., N, 4]
    const at::optional<at::Tensor> scales, // [..., N, 3]
    const at::Tensor viewmats,             // [..., C, 4, 4]
    const at::Tensor Ks,                   // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const CameraModelType camera_model,
    // fwd outputs
    const at::Tensor batch_ids,                   // [nnz]
    const at::Tensor camera_ids,                  // [nnz]
    const at::Tensor gaussian_ids,                // [nnz]
    const at::Tensor conics,                      // [nnz, 3]
    const at::optional<at::Tensor> compensations, // [nnz] optional
    // grad outputs
    const at::Tensor v_means2d,                     // [nnz, 2]
    const at::Tensor v_depths,                      // [nnz]
    const at::Tensor v_conics,                      // [nnz, 3]
    const at::optional<at::Tensor> v_compensations, // [nnz] optional
    const bool sparse_grad,
    // grad inputs
    at::Tensor v_means,                 // [..., N, 3] or [nnz, 3]
    at::optional<at::Tensor> v_covars,  // [..., N, 6] or [nnz, 6] Optional
    at::optional<at::Tensor> v_quats,   // [..., N, 4] or [nnz, 4] Optional
    at::optional<at::Tensor> v_scales,  // [..., N, 3] or [nnz, 3] Optional
    at::optional<at::Tensor> v_viewmats // [..., C, 4, 4] Optional
);

void launch_projection_2dgs_fused_fwd_kernel(
    // inputs
    const at::Tensor means,    // [..., N, 3]
    const at::Tensor quats,    // [..., N, 4]
    const at::Tensor scales,   // [..., N, 3]
    const at::Tensor viewmats, // [..., C, 4, 4]
    const at::Tensor Ks,       // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float near_plane,
    const float far_plane,
    const float radius_clip,
    // outputs
    at::Tensor radii,          // [..., C, N, 2]
    at::Tensor means2d,        // [..., C, N, 2]
    at::Tensor depths,         // [..., C, N]
    at::Tensor ray_transforms, // [..., C, N, 3, 3]
    at::Tensor normals         // [..., C, N, 3]
);
void launch_projection_2dgs_fused_bwd_kernel(
    // fwd inputs
    const at::Tensor means,    // [..., N, 3]
    const at::Tensor quats,    // [..., N, 4]
    const at::Tensor scales,   // [..., N, 3]
    const at::Tensor viewmats, // [..., C, 4, 4]
    const at::Tensor Ks,       // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    // fwd outputs
    const at::Tensor radii,          // [..., C, N, 2]
    const at::Tensor ray_transforms, // [..., C, N, 3, 3]
    // grad outputs
    const at::Tensor v_means2d,        // [..., C, N, 2]
    const at::Tensor v_depths,         // [..., C, N]
    const at::Tensor v_normals,        // [..., C, N, 3]
    const at::Tensor v_ray_transforms, // [..., C, N, 3, 3]
    const bool viewmats_requires_grad,
    // outputs
    at::Tensor v_means,   // [..., N, 3]
    at::Tensor v_quats,   // [..., N, 4]
    at::Tensor v_scales,  // [..., N, 3]
    at::Tensor v_viewmats // [..., C, 4, 4]
);

void launch_projection_2dgs_packed_fwd_kernel(
    // inputs
    const at::Tensor means,    // [..., N, 3]
    const at::Tensor quats,    // [..., N, 4]
    const at::Tensor scales,   // [..., N, 3]
    const at::Tensor viewmats, // [..., C, 4, 4]
    const at::Tensor Ks,       // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float near_plane,
    const float far_plane,
    const float radius_clip,
    const at::optional<at::Tensor> block_accum, // [B * C * blocks_per_row] packing helper
    // outputs
    at::optional<at::Tensor> block_cnts,     // [B * C * blocks_per_row] packing helper
    at::optional<at::Tensor> indptr,         // [B * C + 1]
    at::optional<at::Tensor> batch_ids,      // [nnz]
    at::optional<at::Tensor> camera_ids,     // [nnz]
    at::optional<at::Tensor> gaussian_ids,   // [nnz]
    at::optional<at::Tensor> radii,          // [nnz, 2]
    at::optional<at::Tensor> means2d,        // [nnz, 2]
    at::optional<at::Tensor> depths,         // [nnz]
    at::optional<at::Tensor> ray_transforms, // [nnz, 3, 3]
    at::optional<at::Tensor> normals         // [nnz]
);
void launch_projection_2dgs_packed_bwd_kernel(
    // fwd inputs
    const at::Tensor means,    // [..., N, 3]
    const at::Tensor quats,    // [..., N, 4]
    const at::Tensor scales,   // [..., N, 3]
    const at::Tensor viewmats, // [..., C, 4, 4]
    const at::Tensor Ks,       // [..., C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    // fwd outputs
    const at::Tensor batch_ids,      // [nnz]
    const at::Tensor camera_ids,     // [nnz]
    const at::Tensor gaussian_ids,   // [nnz]
    const at::Tensor ray_transforms, // [nnz, 3, 3]
    // grad outputs
    const at::Tensor v_means2d,        // [nnz, 2]
    const at::Tensor v_depths,         // [nnz]
    const at::Tensor v_ray_transforms, // [nnz, 3, 3]
    const at::Tensor v_normals,        // [nnz, 3]
    const bool sparse_grad,
    // grad inputs
    at::Tensor v_means,                 // [..., N, 3] or [nnz, 3]
    at::Tensor v_quats,                 // [..., N, 4] or [nnz, 4]
    at::Tensor v_scales,                // [..., N, 3] or [nnz, 3]
    at::optional<at::Tensor> v_viewmats // [..., C, 4, 4] Optional
);

void launch_projection_ut_3dgs_fused_kernel(
    // inputs
    const at::Tensor means,                   // [N, 3]
    const at::Tensor quats,                   // [N, 4]
    const at::Tensor scales,                  // [N, 3]
    const at::optional<at::Tensor> opacities, // [N] optional
    const at::Tensor viewmats0,               // [C, 4, 4]
    const at::optional<at::Tensor> viewmats1, // [C, 4, 4] optional for rolling shutter
    const at::Tensor Ks,                      // [C, 3, 3]
    const uint32_t image_width,
    const uint32_t image_height,
    const float eps2d,
    const float near_plane,
    const float far_plane,
    const float radius_clip,
    const CameraModelType camera_model,
    const bool global_z_order,
    // uncented transform
    const at::optional<c10::intrusive_ptr<UnscentedTransformParameters>> &ut_params,
    ShutterType rs_type,
    const at::optional<at::Tensor> radial_coeffs,     // [C, 6] or [C, 4] optional
    const at::optional<at::Tensor> tangential_coeffs, // [C, 2] optional
    const at::optional<at::Tensor> thin_prism_coeffs, // [C, 4] optional
    const at::optional<c10::intrusive_ptr<FThetaCameraDistortionParameters>>
        &ftheta_coeffs, // shared parameters for all cameras
    const at::optional<c10::intrusive_ptr<RowOffsetStructuredSpinningLidarModelParametersExt>> &lidar_coeffs,
    const at::optional<c10::intrusive_ptr<extdist::BivariateWindshieldModelParameters>>
        &external_distortion_params, // external distortion parameters
    // outputs
    at::Tensor radii,                      // [C, N, 2]
    at::Tensor means2d,                    // [C, N, 2]
    at::Tensor depths,                     // [C, N]
    at::Tensor conics,                     // [C, N, 3]
    at::optional<at::Tensor> compensations // [C, N] optional
);
} // namespace gsplat
