Skip to content
Open
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
62 changes: 53 additions & 9 deletions src/operators/dynamic-fully-connected-nc.c
Original file line number Diff line number Diff line change
Expand Up @@ -437,8 +437,17 @@ reshape_dynamic_fully_connected_nc(

const uint32_t kr = ukernel->kr;
const uint32_t sr = ukernel->sr;
const size_t n_stride = round_up(output_channels, nr);
const size_t k_stride = round_up_po2(input_channels, kr * sr);
size_t n_stride;
size_t k_stride;
size_t kr_sr;
if (!xnn_safe_round_up(output_channels, nr, &n_stride) ||
!xnn_safe_mul((size_t)kr, (size_t)sr, &kr_sr) ||
!xnn_safe_round_up_po2(input_channels, kr_sr, &k_stride)) {
xnn_log_error(
"failed to reshape %s operator: GEMM strides overflow size_t",
xnn_operator_type_to_string_v2(dynamic_fully_connected_op));
return xnn_status_out_of_memory;
}
const struct xnn_gemm_config* gemm_config =
dynamic_fully_connected_op->gemm_config;
size_t weights_stride;
Expand All @@ -457,6 +466,12 @@ reshape_dynamic_fully_connected_nc(
return xnn_status_out_of_memory;
}
}
if (weights_stride == SIZE_MAX) {
xnn_log_error(
"failed to reshape %s operator: weights stride overflows size_t",
xnn_operator_type_to_string_v2(dynamic_fully_connected_op));
return xnn_status_out_of_memory;
}
const size_t num_threads = pthreadpool_get_threads_count(threadpool);

// Clear the operator's compute data to avoid accidentally reusing values from
Expand Down Expand Up @@ -547,17 +562,32 @@ reshape_dynamic_fully_connected_nc(
}

// Compute the optimal tile size for this GEMM.
size_t m_stride;
if (!xnn_safe_mul(
input_stride,
(size_t)1 << (packed_lh_config ? packed_lh_config->log2_packed_element_size
: log2_input_element_size),
&m_stride)) {
xnn_log_error(
"failed to reshape %s operator: input stride overflows size_t",
xnn_operator_type_to_string_v2(dynamic_fully_connected_op));
return xnn_status_out_of_memory;
}
const size_t nc = xnn_gemm_best_tile_size(
/*num_groups=*/1, /*m=*/batch_size, /*n=*/output_channels,
/*m_stride=*/input_stride
<< (packed_lh_config ? packed_lh_config->log2_packed_element_size
: log2_input_element_size),
/*m_stride=*/m_stride,
/*n_stride=*/weights_stride,
/*cn_stride=*/1 << log2_output_element_size, mr, nr,
/*num_threads=*/num_threads);

const size_t workspace_offset =
round_up_po2(*workspace_size, XNN_ALLOCATION_ALIGNMENT);
size_t workspace_offset;
if (!xnn_safe_round_up_po2(*workspace_size, XNN_ALLOCATION_ALIGNMENT,
&workspace_offset)) {
xnn_log_error(
"failed to reshape %s operator: workspace offset overflows size_t",
xnn_operator_type_to_string_v2(dynamic_fully_connected_op));
return xnn_status_out_of_memory;
}

// If we are packing the LHS, provide a per-thread workspace to do so inline.
memset(&gemm_context->pack_lh, 0, sizeof(struct pack_lh_context));
Expand All @@ -567,6 +597,21 @@ reshape_dynamic_fully_connected_nc(
assert(workspace_size);
const size_t per_thread_workspace_size = packed_lh_config->size_fn(
mr, /*k=*/input_channels, mr_packed, kr, sr);
if (per_thread_workspace_size == SIZE_MAX) {
xnn_log_error(
"failed to reshape %s operator: packed LHS size overflows size_t",
xnn_operator_type_to_string_v2(dynamic_fully_connected_op));
return xnn_status_out_of_memory;
}
size_t rounded_batch_size;
size_t thread_rows;
if (!xnn_safe_round_up(batch_size, mr, &rounded_batch_size) ||
!xnn_safe_mul(num_threads, (size_t)mr, &thread_rows)) {
xnn_log_error(
"failed to reshape %s operator: batch tiling overflows size_t",
xnn_operator_type_to_string_v2(dynamic_fully_connected_op));
return xnn_status_out_of_memory;
}

// If `xnn_gemm_best_tile_size` suggests an `nc` that is smaller than `n`,
// i.e. it suggests splitting along `output_channels`, then it's probably
Expand All @@ -590,8 +635,7 @@ reshape_dynamic_fully_connected_nc(
"it is a no-op for GEMV.",
xnn_operator_type_to_string(dynamic_fully_connected_op->type),
batch_size, output_channels, input_channels);
} else if (!should_inline_lhs_packing ||
num_threads * mr > round_up(batch_size, mr)) {
} else if (!should_inline_lhs_packing || thread_rows > rounded_batch_size) {
xnn_log_debug(
"Pre-packing LHS of %s with m=%zu, n=%zu, and k=%zu despite "
"request to inline because %s.",
Expand Down
59 changes: 41 additions & 18 deletions src/operators/fully-connected-nc.c
Original file line number Diff line number Diff line change
Expand Up @@ -227,32 +227,49 @@ static XNN_NO_SANITIZE_FUNCTION enum xnn_status create_fully_connected_nc(
const uint32_t sr = UINT32_C(1) << gemm_config->log2_sr;
const uint32_t planes = gemm_config->planes;

const size_t n_stride = round_up(output_channels, nr);

size_t k_stride = round_up_po2(input_channels, kr * sr);
size_t n_stride;
size_t k_stride;
size_t kr_sr;
if (!xnn_safe_round_up(output_channels, nr, &n_stride) ||
!xnn_safe_mul((size_t)kr, (size_t)sr, &kr_sr) ||
!xnn_safe_round_up_po2(input_channels, kr_sr, &k_stride)) {
xnn_log_error(
"failed to create %s operator: GEMM strides overflow size_t",
xnn_operator_type_to_string(operator_type));
goto error;
}

if (filter_is_crumb) {
if (planes != 4) {
xnn_log_error("planes is %u but expected to be 4 for 2 bit", planes);
goto error;
}
k_stride = round_up_po2(input_channels, kr * sr * planes);

// If filter is 2-bit, quarter k_stride (since we will scale k_stride by
// log2_filter_element_size, and we pass 0 for qc2).
k_stride = round_up_po2(k_stride, 4) >> 2;
size_t k_block;
if (!xnn_safe_mul(kr_sr, (size_t)planes, &k_block) ||
!xnn_safe_round_up_po2(input_channels, k_block, &k_stride) ||
!xnn_safe_round_up_po2(k_stride, 4, &k_stride)) {
xnn_log_error(
"failed to create %s operator: GEMM strides overflow size_t",
xnn_operator_type_to_string(operator_type));
goto error;
}
k_stride >>= 2;
} else if (filter_is_nibble) {
input_channels = round_up_po2(input_channels, planes);

if (planes < 1 || planes > 2) {
xnn_log_error("planes is %u but expected to be 1 or 2 for 4 bit", planes);
goto error;
}
k_stride = round_up_po2(input_channels, kr * sr * planes);

// If filter is 4-bit, half k_stride (since we will scale k_stride by
// log2_filter_element_size, and we pass 0 for qc4).
k_stride = round_up_po2(k_stride, 2) >> 1;
size_t k_block;
if (!xnn_safe_round_up_po2(input_channels, planes, &input_channels) ||
!xnn_safe_mul(kr_sr, (size_t)planes, &k_block) ||
!xnn_safe_round_up_po2(input_channels, k_block, &k_stride) ||
!xnn_safe_round_up_po2(k_stride, 2, &k_stride)) {
xnn_log_error(
"failed to create %s operator: GEMM strides overflow size_t",
xnn_operator_type_to_string(operator_type));
goto error;
}
k_stride >>= 1;
}

size_t block_scale_bytes = 0;
Expand Down Expand Up @@ -285,6 +302,12 @@ static XNN_NO_SANITIZE_FUNCTION enum xnn_status create_fully_connected_nc(
goto error;
}
}
if (weights_stride == SIZE_MAX) {
xnn_log_error(
"failed to create %s operator: weights stride overflows size_t",
xnn_operator_type_to_string(operator_type));
goto error;
}
size_t packed_weights_size = 0;
if (!xnn_safe_mul(n_stride, weights_stride, &packed_weights_size)) {
xnn_log_error(
Expand All @@ -294,9 +317,9 @@ static XNN_NO_SANITIZE_FUNCTION enum xnn_status create_fully_connected_nc(
goto error;
}
fully_connected_op->weights_stride = weights_stride;
size_t aligned_total_weights_size =
round_up_po2(packed_weights_size, XNN_ALLOCATION_ALIGNMENT);
if (aligned_total_weights_size < packed_weights_size) {
size_t aligned_total_weights_size;
if (!xnn_safe_round_up_po2(packed_weights_size, XNN_ALLOCATION_ALIGNMENT,
&aligned_total_weights_size)) {
xnn_log_error(
"failed to create %s operator: aligned total weights size overflows "
"size_t",
Expand Down
17 changes: 17 additions & 0 deletions src/xnnpack/math.h
Original file line number Diff line number Diff line change
Expand Up @@ -685,6 +685,23 @@ XNN_INLINE static bool xnn_safe_add(size_t a, size_t b, size_t* result) {
#endif
}

XNN_INLINE static bool xnn_safe_round_up(size_t n, size_t q, size_t* result) {
if (q == 0) {
return false;
}
const size_t quotient = n / q + (n % q != 0);
return xnn_safe_mul(quotient, q, result);
}

XNN_INLINE static bool xnn_safe_round_up_po2(size_t n, size_t q,
size_t* result) {
if (!is_po2(q) || n > SIZE_MAX - (q - 1)) {
return false;
}
*result = round_down_po2(n + q - 1, q);
return true;
}

#ifdef __cplusplus
} // extern "C"
#endif
Expand Down
41 changes: 41 additions & 0 deletions test/operators/dynamic-fully-connected-nc.cc
Original file line number Diff line number Diff line change
Expand Up @@ -327,3 +327,44 @@ TEST(DYNAMIC_FULLY_CONNECTED_NC_F32, overflow_batch_stride) {
/*threadpool=*/nullptr));
}

TEST(DYNAMIC_FULLY_CONNECTED_NC_F32, overflow_n_stride) {
ASSERT_EQ(xnn_status_success, xnn_initialize(/*allocator=*/nullptr));
xnn_operator_t op = nullptr;
ASSERT_EQ(xnn_status_success,
xnn_create_dynamic_fully_connected_nc_f32(
-std::numeric_limits<float>::infinity(),
+std::numeric_limits<float>::infinity(),
/*flags=*/0, &op));
std::unique_ptr<xnn_operator, decltype(&xnn_delete_operator)> auto_op(
op, xnn_delete_operator);

size_t workspace_size = 0;
EXPECT_EQ(
xnn_status_out_of_memory,
xnn_reshape_dynamic_fully_connected_nc_f32(
op, /*batch_size=*/1, /*input_channels=*/1,
/*output_channels=*/(SIZE_MAX / 2) + 1, /*input_stride=*/1,
/*output_stride=*/(SIZE_MAX / 2) + 1, &workspace_size,
/*threadpool=*/nullptr));
}

TEST(DYNAMIC_FULLY_CONNECTED_NC_F32, overflow_k_stride) {
ASSERT_EQ(xnn_status_success, xnn_initialize(/*allocator=*/nullptr));
xnn_operator_t op = nullptr;
ASSERT_EQ(xnn_status_success,
xnn_create_dynamic_fully_connected_nc_f32(
-std::numeric_limits<float>::infinity(),
+std::numeric_limits<float>::infinity(),
/*flags=*/0, &op));
std::unique_ptr<xnn_operator, decltype(&xnn_delete_operator)> auto_op(
op, xnn_delete_operator);

size_t workspace_size = 0;
EXPECT_EQ(
xnn_status_out_of_memory,
xnn_reshape_dynamic_fully_connected_nc_f32(
op, /*batch_size=*/1, /*input_channels=*/(SIZE_MAX / 2) + 1,
/*output_channels=*/1, /*input_stride=*/(SIZE_MAX / 2) + 1,
/*output_stride=*/1, &workspace_size, /*threadpool=*/nullptr));
}

32 changes: 32 additions & 0 deletions test/operators/fully-connected-nc.cc
Original file line number Diff line number Diff line change
Expand Up @@ -3790,6 +3790,38 @@ TEST(FULLY_CONNECTED_NC_F32, overflow_batch_stride) {
op, /*batch_size=*/SIZE_MAX / 50, /*threadpool=*/nullptr));
}

TEST(FULLY_CONNECTED_NC_F32, overflow_n_stride) {
ASSERT_EQ(xnn_status_success, xnn_initialize(/*allocator=*/nullptr));
xnn_operator_t op = nullptr;
const std::vector<float> kernel(1);
const std::vector<float> bias(1);
EXPECT_EQ(
xnn_status_out_of_memory,
xnn_create_fully_connected_nc_f32(
/*input_channels=*/1, /*output_channels=*/(SIZE_MAX / 2) + 1,
/*input_stride=*/1, /*output_stride=*/(SIZE_MAX / 2) + 1,
kernel.data(), bias.data(),
-std::numeric_limits<float>::infinity(),
std::numeric_limits<float>::infinity(),
/*flags=*/0, /*weights_cache=*/nullptr, &op));
}

TEST(FULLY_CONNECTED_NC_F32, overflow_k_stride) {
ASSERT_EQ(xnn_status_success, xnn_initialize(/*allocator=*/nullptr));
xnn_operator_t op = nullptr;
const std::vector<float> kernel(1);
const std::vector<float> bias(1);
EXPECT_EQ(
xnn_status_out_of_memory,
xnn_create_fully_connected_nc_f32(
/*input_channels=*/(SIZE_MAX / 2) + 1, /*output_channels=*/1,
/*input_stride=*/(SIZE_MAX / 2) + 1, /*output_stride=*/1,
kernel.data(), bias.data(),
-std::numeric_limits<float>::infinity(),
std::numeric_limits<float>::infinity(),
/*flags=*/0, /*weights_cache=*/nullptr, &op));
}

TEST(FULLY_CONNECTED_NC_F32, overflow_create_weights_stride) {
ASSERT_EQ(xnn_status_success, xnn_initialize(/*allocator=*/nullptr));
xnn_operator_t op = nullptr;
Expand Down
Loading