Skip to content
Merged
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
24 changes: 15 additions & 9 deletions mlx/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4344,16 +4344,21 @@ array conv_transpose_general(
std::vector<int> padding_lo(padding.size());
std::vector<int> padding_hi(padding.size());
for (int i = 0; i < padding.size(); ++i) {
int wt_size = 1 + dilation[i] * (weight.shape(1 + i) - 1);
padding_lo[i] = wt_size - padding[i] - 1;
int64_t wt_size =
1 + static_cast<int64_t>(dilation[i]) * (weight.shape(1 + i) - 1);
padding_lo[i] = safe_cast(wt_size - padding[i] - 1, "conv");

int conv_output_shape = (input.shape(i + 1) - 1) * stride[i] -
2 * padding[i] + dilation[i] * (weight.shape(i + 1) - 1) + 1;
int64_t conv_output_shape =
static_cast<int64_t>(input.shape(i + 1) - 1) * stride[i] -
2 * static_cast<int64_t>(padding[i]) +
static_cast<int64_t>(dilation[i]) * (weight.shape(i + 1) - 1) + 1;

int in_size = 1 + (conv_output_shape - 1);
int out_size = 1 + stride[i] * (input.shape(1 + i) - 1);
padding_hi[i] = in_size - out_size + padding[i] +
output_padding[i]; // Adjust with output_padding
int64_t in_size = 1 + (conv_output_shape - 1);
int64_t out_size =
1 + static_cast<int64_t>(stride[i]) * (input.shape(1 + i) - 1);
// Adjust with output_padding
padding_hi[i] =
safe_cast(in_size - out_size + padding[i] + output_padding[i], "conv");
}

auto ndim = stride.size();
Expand Down Expand Up @@ -4504,7 +4509,8 @@ array conv_general(

for (int i = 0; i < spatial_dims; i++) {
if (padding_lo[i] < 0) {
starts[i + 1] -= padding_lo[i];
starts[i + 1] = safe_cast(
starts[i + 1] - static_cast<int64_t>(padding_lo[i]), "conv");
padding_lo[i] = 0;
}

Expand Down
56 changes: 36 additions & 20 deletions mlx/primitives.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1257,9 +1257,11 @@ array conv_weight_backward_patches(

// padded shape
for (int i = 1; i < in.ndim() - 1; i++) {
in_padded_shape[i] += padding_lo[i - 1] + padding_hi[i - 1];
padding_ends[i] += padding_lo[i - 1];
padding_starts[i] += padding_lo[i - 1];
int64_t lo = padding_lo[i - 1];
int64_t hi = padding_hi[i - 1];
in_padded_shape[i] = safe_cast(in_padded_shape[i] + lo + hi, "conv");
padding_ends[i] = safe_cast(padding_ends[i] + lo, "conv");
padding_starts[i] = safe_cast(padding_starts[i] + lo, "conv");
}

// padded strides (contiguous)
Expand Down Expand Up @@ -1315,12 +1317,18 @@ array conv_weight_backward_patches(
namespace {

// Conv helpers
inline int conv_out_axis_size(int in_dim, int wt_dim, int stride, int padding) {
// Computed in 64 bits so extreme but in-range int32 parameters do not overflow.
inline int64_t conv_out_axis_size(
int64_t in_dim,
int64_t wt_dim,
int64_t stride,
int64_t padding) {
return ((in_dim + padding - wt_dim) / stride) + 1;
}

// Conv helpers
inline int dilate_size(int dim, int dil) {
// Computed in 64 bits so extreme but in-range int32 parameters do not overflow.
inline int64_t dilate_size(int64_t dim, int64_t dil) {
return 1 + dil * (dim - 1);
}

Expand Down Expand Up @@ -1399,13 +1407,16 @@ Shape Convolution::conv_out_shape(
throw std::invalid_argument(msg.str());
}

int kd = dilate_size(wt_shape[i], kernel_dilation[i - 1]);
int id = dilate_size(in_shape[i], input_dilation[i - 1]);
int64_t kd = dilate_size(wt_shape[i], kernel_dilation[i - 1]);
int64_t id = dilate_size(in_shape[i], input_dilation[i - 1]);

out_shape[i] = conv_out_axis_size(
id, kd, strides[i - 1], pads_lo[i - 1] + pads_hi[i - 1]);
int64_t out_size = conv_out_axis_size(
id,
kd,
strides[i - 1],
static_cast<int64_t>(pads_lo[i - 1]) + pads_hi[i - 1]);

if (out_shape[i] <= 0) {
if (out_size <= 0) {
std::ostringstream msg;
msg << "[conv] Spatial dimensions of input after padding"
<< " cannot be smaller than weight spatial dimensions."
Expand All @@ -1414,6 +1425,8 @@ Shape Convolution::conv_out_shape(
<< ", and weight of shape " << wt_shape << ".";
throw std::invalid_argument(msg.str());
}

out_shape[i] = safe_cast(out_size, "conv");
}
out_shape[i] = O;

Expand Down Expand Up @@ -1457,12 +1470,12 @@ std::vector<array> Convolution::vjp(
std::vector<int> padding_hi = padding_hi_;

for (int i = 0; i < padding_lo.size(); ++i) {
int wt_size = 1 + kernel_dilation_[i] * (wt.shape(1 + i) - 1);
padding_lo[i] = wt_size - padding_lo_[i] - 1;
int64_t wt_size = dilate_size(wt.shape(1 + i), kernel_dilation_[i]);
padding_lo[i] = safe_cast(wt_size - padding_lo_[i] - 1, "conv");

int in_size = 1 + input_dilation_[i] * (in.shape(1 + i) - 1);
int out_size = 1 + kernel_strides_[i] * (cotan.shape(1 + i) - 1);
padding_hi[i] = in_size - out_size + padding_hi_[i];
int64_t in_size = dilate_size(in.shape(1 + i), input_dilation_[i]);
int64_t out_size = dilate_size(cotan.shape(1 + i), kernel_strides_[i]);
padding_hi[i] = safe_cast(in_size - out_size + padding_hi_[i], "conv");
}

// Check for negative padding
Expand Down Expand Up @@ -1494,7 +1507,8 @@ std::vector<array> Convolution::vjp(

for (int i = 0; i < grad.ndim() - 2; i++) {
if (padding_lo[i] < 0) {
starts[i + 1] -= padding_lo[i];
starts[i + 1] = safe_cast(
starts[i + 1] - static_cast<int64_t>(padding_lo[i]), "conv");
}
if (padding_hi[i] < 0) {
stops[i + 1] += padding_hi[i];
Expand Down Expand Up @@ -1522,10 +1536,12 @@ std::vector<array> Convolution::vjp(
auto padding_hi = padding_lo_;

for (int i = 0; i < padding_hi.size(); ++i) {
int in_size = 1 + input_dilation_[i] * (in.shape(1 + i) - 1);
int out_size = 1 + kernel_strides_[i] * (cotan.shape(1 + i) - 1);
int wt_size = 1 + kernel_dilation_[i] * (wt.shape(1 + i) - 1);
padding_hi[i] = out_size - in_size + wt_size - padding_hi[i] - 1;
int64_t in_size = dilate_size(in.shape(1 + i), input_dilation_[i]);
int64_t out_size =
dilate_size(cotan.shape(1 + i), kernel_strides_[i]);
int64_t wt_size = dilate_size(wt.shape(1 + i), kernel_dilation_[i]);
padding_hi[i] = safe_cast(
out_size - in_size + wt_size - padding_hi[i] - 1, "conv");
}

auto cotan_trans = swapaxes(cotan, 0, -1, stream());
Expand Down
72 changes: 72 additions & 0 deletions tests/ops_tests.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4342,6 +4342,78 @@ TEST_CASE("test conv_transpose3d with output_padding") {
CHECK(array_equal(out, expected).item<bool>());
}

TEST_CASE("test conv shape overflow") {
// Conv shape arithmetic must not overflow (signed-int UB) for large but
// otherwise valid int32 parameters; out-of-range results are rejected
// gracefully. https://github.com/ml-explore/mlx/issues/3611
const int imax = 2147483647;
const int imin = -2147483647 - 1;
auto in = zeros({1, 8, 8, 1});
auto wt = zeros({1, 3, 3, 1});

// A kernel dilated past the input reports the spatial-size error.
CHECK_THROWS_AS(
conv_general(in, wt, {1, 1}, {0, 0}, {0, 0}, {imax, imax}, {1, 1}),
std::invalid_argument);

// Padding sums, input dilation, and negating a padding of INT_MIN raise.
CHECK_THROWS_AS(
conv_general(in, wt, {1, 1}, {imax, imax}, {imax, imax}, {1, 1}, {1, 1}),
std::overflow_error);
CHECK_THROWS_AS(
conv_general(in, wt, {1, 1}, {imax, 0}, {0, 0}, {1, 1}, {1, 1}),
std::overflow_error);
CHECK_THROWS_AS(
conv_general(in, wt, {1, 1}, {0, 0}, {0, 0}, {1, 1}, {imax, imax}),
std::overflow_error);
CHECK_THROWS_AS(
conv_general(in, wt, {1, 1}, {imin, imin}, {0, 0}, {1, 1}, {1, 1}),
std::overflow_error);

// The transposed padding setup runs before conv_general validates it.
auto in_t = zeros({1, 4, 4, 1});
CHECK_THROWS_AS(
conv_transpose2d(in_t, wt, {1, 1}, {0, 0}, {imax, imax}, {0, 0}),
std::overflow_error);
CHECK_THROWS_AS(
conv_transpose2d(in_t, wt, {1, 1}, {imin, imin}, {1, 1}, {0, 0}),
std::overflow_error);

// The dilated input and kernel are both near 4e9 and cancel in the forward
// output, so only the gradient's own recompute goes out of range.
auto in_g = zeros({1, 3, 1, 1});
auto wt_g = zeros({1, 200000, 1, 1});
auto conv_g = [](const std::vector<array>& primals) {
return std::vector<array>{conv_general(
primals[0],
primals[1],
{1, 1},
{0, 0},
{0, 0},
{20000, 1},
{2000000000, 1})};
};
auto cotan = ones(conv_g({in_g, wt_g})[0].shape());
CHECK_THROWS_AS(vjp(conv_g, {in_g, wt_g}, {cotan}), std::overflow_error);

// The weight gradient pads without dividing by the stride.
auto in_w = zeros({1, 8, 8, 1});
auto conv_w = [&in_w, imax](const std::vector<array>& primals) {
return std::vector<array>{conv_general(
in_w, primals[0], {imax, 1}, {imax, 0}, {imax, 0}, {1, 1}, {1, 1})};
};
auto cotan_w = ones(conv_w({wt})[0].shape());
CHECK_THROWS_AS(vjp(conv_w, {wt}, {cotan_w}), std::overflow_error);

// In-range parameters still give the same shapes.
CHECK_EQ(
conv_general(in, wt, {1, 1}, {1, 1}, {1, 1}, {2, 2}, {1, 1}).shape(),
Shape{1, 6, 6, 1});
CHECK_EQ(
conv_transpose2d(in_t, wt, {2, 2}, {1, 1}, {1, 1}, {1, 1}).shape(),
Shape{1, 8, 8, 1});
}

TEST_CASE("test fp8 conversion") {
for (auto t : {float32, float16, bfloat16}) {
array in({-1.125, -1.0, 0.0, 1.0, 1.125, 4.5, 448.0}, t);
Expand Down
Loading