Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 44 additions & 6 deletions scripts/generate_wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -490,6 +490,26 @@ def _find_optional_vector_int64_params(op_name):
)


def _vector_int64_kind(arg, vector_params, optional_vector_params):
"""Classify a vector parameter, preferring its overload-local type."""
spelling = "".join(arg.type.spelling.split())
spelling = spelling.removeprefix("const").removesuffix("&")

if spelling == "std::optional<std::vector<int64_t>>":
return "optional"

if spelling == "std::vector<int64_t>":
return "vector"

if arg.spelling in optional_vector_params:
return "optional"

if arg.spelling in vector_params:
return "vector"

return None


def _find_tensor_params(op_name):
source = _find_base_header(op_name).read_text()

Expand Down Expand Up @@ -545,7 +565,10 @@ def _is_optional(arg):
return "std::optional" in arg.type.spelling

def _is_optional_vector_int64(arg):
return arg.spelling in optional_vector_int64_params
return (
_vector_int64_kind(arg, vector_int64_params, optional_vector_int64_params)
== "optional"
)

def _is_vector_tensor(arg):
if arg.spelling in vector_tensor_params:
Expand All @@ -554,7 +577,10 @@ def _is_vector_tensor(arg):
return "std::vector" in arg.type.spelling and "Tensor" in arg.type.spelling

def _is_vector_int64(arg):
return arg.spelling in vector_int64_params
return (
_vector_int64_kind(arg, vector_int64_params, optional_vector_int64_params)
== "vector"
)

def _is_data_type(arg):
return _is_data_type_spelling(arg.type.spelling)
Expand Down Expand Up @@ -1002,7 +1028,10 @@ def _is_optional_tensor(arg):
return arg.spelling in optional_tensor_params

def _is_optional_vector_int64(arg):
return arg.spelling in optional_vector_int64_params
return (
_vector_int64_kind(arg, vector_int64_params, optional_vector_int64_params)
== "optional"
)

def _is_vector_tensor(arg):
if arg.spelling in vector_tensor_params:
Expand All @@ -1011,7 +1040,10 @@ def _is_vector_tensor(arg):
return "std::vector" in arg.type.spelling and "Tensor" in arg.type.spelling

def _is_vector_int64(arg):
return arg.spelling in vector_int64_params
return (
_vector_int64_kind(arg, vector_int64_params, optional_vector_int64_params)
== "vector"
)

def _is_tensor(arg):
if arg.spelling in optional_non_tensor_params:
Expand Down Expand Up @@ -1239,7 +1271,10 @@ def _is_optional_tensor(arg):
return False

def _is_optional_vector_int64(arg):
return arg.spelling in optional_vector_int64_params
return (
_vector_int64_kind(arg, vector_int64_params, optional_vector_int64_params)
== "optional"
)

def _is_vector_tensor(arg):
if arg.spelling in vector_tensor_params:
Expand All @@ -1250,7 +1285,10 @@ def _is_vector_tensor(arg):
)

def _is_vector_int64(arg):
return arg.spelling in vector_int64_params
return (
_vector_int64_kind(arg, vector_int64_params, optional_vector_int64_params)
== "vector"
)

def _is_tensor(arg):
if arg.spelling in optional_non_tensor_params:
Expand Down
1 change: 0 additions & 1 deletion scripts/torch_ops.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -266,7 +266,6 @@
- max_pool2d_with_indices_backward
- max_pool3d_with_indices
- max_pool3d_with_indices_backward
- max_unpool2d
- max_unpool3d
- maximum
- mean
Expand Down
78 changes: 78 additions & 0 deletions src/base/detail/max_unpool.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
#ifndef INFINI_OPS_BASE_DETAIL_MAX_UNPOOL_H_
#define INFINI_OPS_BASE_DETAIL_MAX_UNPOOL_H_

#include <cassert>
#include <cstddef>
#include <cstdint>
#include <optional>
#include <utility>
#include <vector>

#include "operator.h"

namespace infini::ops::max_unpool_detail {

template <std::size_t SpatialDimensions>
inline std::pair<std::vector<int64_t>, std::vector<int64_t>> ResolveGeometry(
const Tensor input, const std::vector<int64_t>& kernel_size,
const std::optional<std::vector<int64_t>>& stride,
const std::vector<int64_t>& padding,
const std::optional<std::vector<int64_t>>& output_size) {
assert((input.ndim() == SpatialDimensions + 1 ||
input.ndim() == SpatialDimensions + 2) &&
"`MaxUnpool` input rank must include the spatial dimensions");
assert(kernel_size.size() == SpatialDimensions &&
"`MaxUnpool` `kernel_size` has the wrong length");
assert(padding.size() == SpatialDimensions &&
"`MaxUnpool` `padding` has the wrong length");

auto resolved_stride = stride.value_or(kernel_size);
assert(resolved_stride.size() == SpatialDimensions &&
"`MaxUnpool` `stride` has the wrong length");

std::vector<int64_t> default_size(SpatialDimensions);

for (std::size_t dim = 0; dim < SpatialDimensions; ++dim) {
assert(kernel_size[dim] > 0 &&
"`MaxUnpool` requires positive `kernel_size` values");
assert(resolved_stride[dim] > 0 &&
"`MaxUnpool` requires positive `stride` values");
assert(padding[dim] >= 0 &&
"`MaxUnpool` requires non-negative `padding` values");

const auto input_size = static_cast<int64_t>(
input.size(input.ndim() - SpatialDimensions + dim));
default_size[dim] = (input_size - 1) * resolved_stride[dim] +
kernel_size[dim] - 2 * padding[dim];
}

auto resolved_output_size = output_size.value_or(default_size);

if (output_size.has_value()) {
if (resolved_output_size.size() == SpatialDimensions + 2) {
resolved_output_size.erase(resolved_output_size.begin(),
resolved_output_size.begin() + 2);
}

assert(resolved_output_size.size() == SpatialDimensions &&
"`MaxUnpool` `output_size` has the wrong length");

for (std::size_t dim = 0; dim < SpatialDimensions; ++dim) {
const auto min_size = default_size[dim] - resolved_stride[dim];
const auto max_size = default_size[dim] + resolved_stride[dim];
assert((min_size < resolved_output_size[dim] &&
resolved_output_size[dim] < max_size) &&
"`MaxUnpool` `output_size` is outside the valid range");
}
}

for (const auto value : resolved_output_size) {
assert(value >= 0 && "`MaxUnpool` requires non-negative output dimensions");
}

return {std::move(resolved_output_size), std::move(resolved_stride)};
}

} // namespace infini::ops::max_unpool_detail

#endif
36 changes: 28 additions & 8 deletions src/base/max_unpool2d.h
Original file line number Diff line number Diff line change
@@ -1,16 +1,21 @@
#ifndef INFINI_OPS_BASE_MAX_UNPOOL2D_H_
#define INFINI_OPS_BASE_MAX_UNPOOL2D_H_

#include <cstdint>
#include <optional>
#include <vector>

#include "operator.h"
#include "detail/max_unpool.h"

namespace infini::ops {

class MaxUnpool2d : public Operator<MaxUnpool2d> {
public:
MaxUnpool2d(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size, Tensor out)
const std::vector<int64_t> kernel_size,
const std::optional<std::vector<int64_t>> stride,
const std::vector<int64_t> padding,
const std::optional<std::vector<int64_t>> output_size, Tensor out)
: input_shape_{input.shape()},
input_strides_{input.strides()},
input_type_{input.dtype()},
Expand All @@ -20,14 +25,29 @@ class MaxUnpool2d : public Operator<MaxUnpool2d> {
out_shape_{out.shape()},
out_strides_{out.strides()},
out_type_{out.dtype()},
output_size_{output_size},
device_index_{out.device().index()} {}

virtual void operator()(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size,
Tensor out) const = 0;
output_size_{},
device_index_{out.device().index()} {
auto geometry = max_unpool_detail::ResolveGeometry<2>(
input, kernel_size, stride, padding, output_size);
output_size_ = std::move(geometry.first);
}

void operator()(const Tensor input, const Tensor indices,
const std::vector<int64_t> kernel_size,
const std::optional<std::vector<int64_t>> stride,
const std::vector<int64_t> padding,
const std::optional<std::vector<int64_t>> output_size,
Tensor out) const {
auto geometry = max_unpool_detail::ResolveGeometry<2>(
input, kernel_size, stride, padding, output_size);
Run(input, indices, std::move(geometry.first), out);
}

protected:
virtual void Run(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size,
Tensor out) const = 0;

Tensor::Shape input_shape_;

Tensor::Strides input_strides_;
Expand Down
54 changes: 48 additions & 6 deletions src/base/max_unpool3d.h
Original file line number Diff line number Diff line change
@@ -1,14 +1,43 @@
#ifndef INFINI_OPS_BASE_MAX_UNPOOL3D_H_
#define INFINI_OPS_BASE_MAX_UNPOOL3D_H_

#include <cstdint>
#include <optional>
#include <vector>

#include "operator.h"
#include "detail/max_unpool.h"

namespace infini::ops {

class MaxUnpool3d : public Operator<MaxUnpool3d> {
public:
MaxUnpool3d(const Tensor input, const Tensor indices,
const std::vector<int64_t> kernel_size,
const std::optional<std::vector<int64_t>> stride,
const std::vector<int64_t> padding,
const std::optional<std::vector<int64_t>> output_size, Tensor out)
: input_shape_{input.shape()},
input_strides_{input.strides()},
input_type_{input.dtype()},
indices_shape_{indices.shape()},
indices_strides_{indices.strides()},
indices_type_{indices.dtype()},
out_shape_{out.shape()},
out_strides_{out.strides()},
out_type_{out.dtype()},
output_size_{},
stride_{},
padding_{padding},
device_index_{out.device().index()} {
auto geometry = max_unpool_detail::ResolveGeometry<3>(
input, kernel_size, stride, padding, output_size);
output_size_ = std::move(geometry.first);
stride_ = std::move(geometry.second);
}

/// \deprecated Use the overload that accepts `kernel_size`, `stride`,
/// `padding`, and `output_size` instead.
[[deprecated("Use the `kernel_size` overload instead.")]]
MaxUnpool3d(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size,
const std::vector<int64_t> stride,
Expand All @@ -27,11 +56,24 @@ class MaxUnpool3d : public Operator<MaxUnpool3d> {
padding_{padding},
device_index_{out.device().index()} {}

virtual void operator()(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size,
const std::vector<int64_t> stride,
const std::vector<int64_t> padding,
Tensor out) const = 0;
void operator()(const Tensor input, const Tensor indices,
const std::vector<int64_t> kernel_size,
const std::optional<std::vector<int64_t>> stride,
const std::vector<int64_t> padding,
const std::optional<std::vector<int64_t>> output_size,
Tensor out) const {
const auto geometry = max_unpool_detail::ResolveGeometry<3>(
input, kernel_size, stride, padding, output_size);
(*this)(input, indices, geometry.first, geometry.second, padding, out);
}

/// \deprecated Use the overload that accepts `kernel_size`, `stride`,
/// `padding`, and `output_size` instead.
[[deprecated("Use the `kernel_size` overload instead.")]] virtual void
operator()(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size,
const std::vector<int64_t> stride,
const std::vector<int64_t> padding, Tensor out) const = 0;

protected:
Tensor::Shape input_shape_;
Expand Down
50 changes: 50 additions & 0 deletions src/torch/ops/max_unpool2d/max_unpool2d.cc
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
#include "torch/ops/max_unpool2d/max_unpool2d.h"

#include "torch/tensor_.h"

namespace infini::ops {

template <Device::Type kDev>
void Operator<MaxUnpool2d, kDev, 8>::Run(const Tensor input,
const Tensor indices,
const std::vector<int64_t> output_size,
Tensor out) const {
const auto device_index = out.device().index();
auto at_input =
ToAtenTensor<kDev>(const_cast<void*>(input.data()), input_shape_,
input_strides_, input_type_, device_index);
auto at_indices =
ToAtenTensor<kDev>(const_cast<void*>(indices.data()), indices_shape_,
indices_strides_, indices_type_, device_index);
auto at_out = ToAtenTensor<kDev>(out.data(), out_shape_, out_strides_,
out_type_, device_index);

at::max_unpool2d_out(at_out, at_input, at_indices, output_size);
}

#ifdef WITH_CPU
template class Operator<MaxUnpool2d, Device::Type::kCpu, 8>;
#endif
#ifdef WITH_NVIDIA
template class Operator<MaxUnpool2d, Device::Type::kNvidia, 8>;
#endif
#ifdef WITH_CAMBRICON
template class Operator<MaxUnpool2d, Device::Type::kCambricon, 8>;
#endif
#ifdef WITH_ASCEND
template class Operator<MaxUnpool2d, Device::Type::kAscend, 8>;
#endif
#ifdef WITH_METAX
template class Operator<MaxUnpool2d, Device::Type::kMetax, 8>;
#endif
#ifdef WITH_MOORE
template class Operator<MaxUnpool2d, Device::Type::kMoore, 8>;
#endif
#ifdef WITH_ILUVATAR
template class Operator<MaxUnpool2d, Device::Type::kIluvatar, 8>;
#endif
#ifdef WITH_HYGON
template class Operator<MaxUnpool2d, Device::Type::kHygon, 8>;
#endif

} // namespace infini::ops
20 changes: 20 additions & 0 deletions src/torch/ops/max_unpool2d/max_unpool2d.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
#ifndef INFINI_OPS_TORCH_MAX_UNPOOL2D_H_
#define INFINI_OPS_TORCH_MAX_UNPOOL2D_H_

#include "base/max_unpool2d.h"

namespace infini::ops {

template <Device::Type kDev>
class Operator<MaxUnpool2d, kDev, 8> : public MaxUnpool2d {
public:
using MaxUnpool2d::MaxUnpool2d;

protected:
void Run(const Tensor input, const Tensor indices,
const std::vector<int64_t> output_size, Tensor out) const override;
};

} // namespace infini::ops

#endif
Loading
Loading