Skip to content

Commit 16d3b9f

Browse files
committed
[ET-VK][ops] Add batch support to q8ta convolutions
Add batch-aware direct, pointwise, depthwise, and im2col q8ta convolution dispatch. Batched im2col uses an NCHW scratch tensor and batch-strided shader indexing, with conservative groups=1 high-K/small-spatial routing bounded by a 32 MiB scratch cap. Grouped and large-spatial convolutions stay on the direct path. Authored with Codex. Differential Revision: [D117869783](https://our.internmc.facebook.com/intern/diff/D117869783/) ghstack-source-id: 421281491 Pull-Request: #22254
1 parent 74d024c commit 16d3b9f

12 files changed

Lines changed: 311 additions & 64 deletions

backends/vulkan/runtime/graph/ops/glsl/q8ta_conv2d.glsl

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -77,14 +77,18 @@ void main() {
7777
// Thread mapping
7878
int oc4 = int(gl_GlobalInvocationID.z);
7979
int w4 = int(gl_GlobalInvocationID.x);
80+
const int H = int(outp.sizes[0][1]);
81+
const int hn = int(gl_GlobalInvocationID.y);
82+
const int n = hn / H;
83+
const int h = hn % H;
8084

8185
// Initialize output tensor index (WHCN order)
8286
// Each thread handles 4 adjacent widths starting at base_out_w
8387
TensorIndex4D outp_tidx;
8488
outp_tidx.data[0] = w4 * 4;
85-
outp_tidx.data[1] = int(gl_GlobalInvocationID.y);
89+
outp_tidx.data[1] = h;
8690
outp_tidx.data[2] = oc4 * 4;
87-
outp_tidx.data[3] = 0;
91+
outp_tidx.data[3] = n;
8892

8993
const int W = int(outp.sizes[0][0]);
9094
const int OC = int(outp.sizes[0][2]);
@@ -113,6 +117,7 @@ void main() {
113117
const int inp_w_stride = int(inp.strides[0][0]);
114118
const int inp_h_stride = int(inp.strides[0][1]);
115119
const int inp_c_stride = int(inp.strides[0][2]);
120+
const int inp_n_stride = int(inp.strides[0][3]);
116121
const int w_texel_step = conv2d_params.dilation.x * inp_w_stride;
117122
const int h_texel_step = conv2d_params.dilation.y * inp_h_stride;
118123
const int subtile_w_step = conv2d_params.stride.x * inp_w_stride;
@@ -122,7 +127,7 @@ void main() {
122127
inp_tidx.data[0] = outp_tidx.data[0] * conv2d_params.stride.x - conv2d_params.padding.x;
123128
inp_tidx.data[1] = outp_tidx.data[1] * conv2d_params.stride.y - conv2d_params.padding.y;
124129
inp_tidx.data[2] = ic_group_start;
125-
inp_tidx.data[3] = 0;
130+
inp_tidx.data[3] = n;
126131

127132
int base_inp_texel_idx;
128133
if (get_outer_packed_dim_block_size(inp_layout) == 1) {
@@ -172,7 +177,11 @@ void main() {
172177
// inp_texel_idx = tensor4d_idx_to_texel_idx(inp, inp_tidx, inp_layout);
173178
const int w4 = div_4(inp_tidx.data[0]);
174179
const int inp_c4 = div_4(inp_tidx.data[2]);
175-
inp_texel_idx = (inp_tidx.data[1] * inp_h_stride + w4 * inp_w_stride + inp_c4) * 4 + mod_4(inp_tidx.data[0]);
180+
inp_texel_idx =
181+
(n * inp_n_stride + inp_tidx.data[1] * inp_h_stride +
182+
w4 * inp_w_stride + inp_c4) *
183+
4 +
184+
mod_4(inp_tidx.data[0]);
176185
}
177186
packed_input = t_packed_int8_input[inp_texel_idx];
178187
}

backends/vulkan/runtime/graph/ops/glsl/q8ta_conv2d_dw.glsl

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -71,14 +71,18 @@ ivec4 quantize(const vec4 texel, const float inv_scale, const int zp) {
7171

7272
void main() {
7373
const int c4 = int(gl_GlobalInvocationID.z);
74+
const int H = int(outp.sizes[0][1]);
75+
const int hn = int(gl_GlobalInvocationID.y);
76+
const int n = hn / H;
77+
const int h = hn % H;
7478

7579
// Initialize output tensor index (WHCN order)
7680
// Each thread handles 4 adjacent widths starting at base_out_w
7781
TensorIndex4D outp_tidx;
7882
outp_tidx.data[0] = int(gl_GlobalInvocationID.x) * 4;
79-
outp_tidx.data[1] = int(gl_GlobalInvocationID.y);
83+
outp_tidx.data[1] = h;
8084
outp_tidx.data[2] = c4 * 4;
81-
outp_tidx.data[3] = 0;
85+
outp_tidx.data[3] = n;
8286

8387
const int W = int(outp.sizes[0][0]);
8488
const int C4 = int(div_up_4(outp.sizes[0][2]));
@@ -94,6 +98,7 @@ void main() {
9498
// Get strides for width and height dimensions (in texel space)
9599
const int w_stride = int(inp.strides[0][0]);
96100
const int h_stride = int(inp.strides[0][1]);
101+
const int n_stride = int(inp.strides[0][3]);
97102

98103
// Pre-compute step sizes for efficient indexing
99104
const int w_texel_step = conv2d_params.dilation.x * w_stride;
@@ -106,7 +111,7 @@ void main() {
106111
inp_tidx.data[0] = outp_tidx.data[0] * conv2d_params.stride.x - conv2d_params.padding.x;
107112
inp_tidx.data[1] = outp_tidx.data[1] * conv2d_params.stride.y - conv2d_params.padding.y;
108113
inp_tidx.data[2] = outp_tidx.data[2];
109-
inp_tidx.data[3] = 0; // batch = 0 since N == 1
114+
inp_tidx.data[3] = n;
110115

111116
int base_inp_texel_idx;
112117
if (get_outer_packed_dim_block_size(inp_layout) == 1) {
@@ -152,7 +157,11 @@ void main() {
152157
// inp_texel_idx = base_inp_texel_idx + div_4(w_offset) * w_stride + mod_4(w_offset);
153158
// inp_texel_idx = tensor4d_idx_to_texel_idx(inp, inp_tidx, inp_layout);
154159
const int w4 = div_4(inp_tidx.data[0]);
155-
inp_texel_idx = (inp_tidx.data[1] * h_stride + w4 * w_stride + c4) * 4 + mod_4(inp_tidx.data[0]);
160+
inp_texel_idx =
161+
(n * n_stride + inp_tidx.data[1] * h_stride +
162+
w4 * w_stride + c4) *
163+
4 +
164+
mod_4(inp_tidx.data[0]);
156165
}
157166
const int packed_input = t_packed_int8_input[inp_texel_idx];
158167
input_4c = unpack_int8x4(packed_input);

backends/vulkan/runtime/graph/ops/glsl/q8ta_conv2d_pw.glsl

Lines changed: 16 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -73,14 +73,17 @@ ${layout_declare_spec_const(C, "int", "inp_layout", "CONTIG_LAYOUT_INT")}
7373
int compute_outp_buffer_idx(
7474
const int w_block_idx,
7575
const int h_idx,
76-
const int c_block_idx) {
76+
const int c_block_idx,
77+
const int n_idx) {
7778
if (get_outer_packed_dim_block_size(outp_layout) == 1) {
78-
return h_idx * int(outp.strides[0][1])
79+
return n_idx * int(outp.strides[0][3])
80+
+ h_idx * int(outp.strides[0][1])
7981
+ mul_4(w_block_idx) * int(outp.strides[0][0])
8082
+ c_block_idx * int(outp.strides[0][2]);
8183
} else {
8284
return mul_4(
83-
h_idx * int(outp.strides[0][1])
85+
n_idx * int(outp.strides[0][3])
86+
+ h_idx * int(outp.strides[0][1])
8487
+ w_block_idx * int(outp.strides[0][0])
8588
+ c_block_idx * int(outp.strides[0][2]));
8689
}
@@ -91,20 +94,22 @@ void main() {
9194
// Thread mapping: each thread handles TILE_M widths x TILE_N output channels.
9295
// gl_GlobalInvocationID.x -> output channel blocks.
9396
// gl_GlobalInvocationID.y -> width blocks.
94-
// gl_GlobalInvocationID.z → batch (or height * batch combined)
97+
// gl_GlobalInvocationID.z -> height * batch.
9598
const int oc_block_idx = int(gl_GlobalInvocationID.x) * TILE_N4;
9699
const int ow_block_idx = int(gl_GlobalInvocationID.y) * TILE_M4;
97-
const int oh = int(gl_GlobalInvocationID.z);
98100

99101
// Get output extents in block space (div_up_4 for packed dimensions)
100102
const int W = int(outp.sizes[0][0]);
101103
const int W4 = div_up_4(int(outp.sizes[0][0]));
102104
const int H = int(outp.sizes[0][1]);
103105
const int OC4 = div_up_4(int(outp.sizes[0][2]));
106+
const int hn = int(gl_GlobalInvocationID.z);
107+
const int n = hn / H;
108+
const int oh = hn % H;
104109

105110
// Bounds check in block space
106111
if (ow_block_idx >= W4 ||
107-
oh >= H ||
112+
n >= int(outp.sizes[0][3]) ||
108113
oc_block_idx >= OC4) {
109114
return;
110115
}
@@ -118,6 +123,7 @@ void main() {
118123
const int inp_w_stride = int(inp.strides[0][0]);
119124
const int inp_h_stride = int(inp.strides[0][1]);
120125
const int inp_c_stride = int(inp.strides[0][2]);
126+
const int inp_n_stride = int(inp.strides[0][3]);
121127

122128
// Initialize int32 accumulator
123129
ivec4 out_accum[TILE_M][TILE_N4];
@@ -133,7 +139,8 @@ void main() {
133139
// Compute initial input tile index with group offset
134140
// For grouped im2col, each group's K range starts at group_idx * K4_per_group
135141
// For non-grouped (groups=1), group_idx is always 0 so offset is 0
136-
int input_idx = oh * inp_h_stride
142+
int input_idx = n * inp_n_stride
143+
+ oh * inp_h_stride
137144
+ ow_block_idx * inp_w_stride
138145
+ group_idx * K4_per_group;
139146

@@ -256,7 +263,8 @@ void main() {
256263
const int base_outp_buffer_idx = compute_outp_buffer_idx(
257264
ow_block_idx + m4,
258265
oh,
259-
oc_block_idx + n4);
266+
oc_block_idx + n4,
267+
n);
260268
if (oc_block_idx + n4 < OC4) {
261269
// Store individual ints from the ivec4
262270
const int subtile_w_limit = min(4, W - mul_4(ow_block_idx + m4));

backends/vulkan/runtime/graph/ops/glsl/q8ta_im2col.glsl

Lines changed: 12 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -50,15 +50,19 @@ void main() {
5050
const int im2col_W4 = div_up_4(im2col_sizes.x);
5151
const int im2col_H = im2col_sizes.y;
5252
const int im2col_Z4 = div_up_4(im2col_sizes.z);
53+
const int im2col_N = im2col_sizes.w;
5354

5455
// im2col block index from linear output buffer index
5556
const int c4_idx = out_buf_idx % im2col_Z4;
5657
const int row = out_buf_idx / im2col_Z4;
5758
const int w4_idx = row % im2col_W4;
58-
const int h_idx = row / im2col_W4;
59+
const int hn_idx = row / im2col_W4;
60+
const int h_idx = hn_idx % im2col_H;
61+
const int n_idx = hn_idx / im2col_H;
5962

6063
// out of bounds check
61-
if (w4_idx >= im2col_W4 || h_idx >= im2col_H || c4_idx >= im2col_Z4) {
64+
if (w4_idx >= im2col_W4 || h_idx >= im2col_H ||
65+
c4_idx >= im2col_Z4 || n_idx >= im2col_N) {
6266
return;
6367
}
6468

@@ -108,12 +112,14 @@ void main() {
108112
const int x_mod = mod_4(x);
109113
int scalar_idx;
110114
if (get_outer_packed_dim_block_size(inp_layout) == 1) {
111-
scalar_idx = input_y * int(inp.strides[0][1])
115+
scalar_idx = n_idx * int(inp.strides[0][3])
116+
+ input_y * int(inp.strides[0][1])
112117
+ x * int(inp.strides[0][0])
113118
+ z4 * int(inp.strides[0][2]);
114119
} else {
115120
scalar_idx = mul_4(
116-
input_y * int(inp.strides[0][1])
121+
n_idx * int(inp.strides[0][3])
122+
+ input_y * int(inp.strides[0][1])
117123
+ x4 * int(inp.strides[0][0])
118124
+ z4) + x_mod;
119125
}
@@ -122,7 +128,8 @@ void main() {
122128
}
123129

124130
// store_packed_int8_output_tile (with TILE_M4=1, TILE_N4=1)
125-
const int buffer_idx = h_idx * int(im2col_outp.strides[0][1])
131+
const int buffer_idx = n_idx * int(im2col_outp.strides[0][3])
132+
+ h_idx * int(im2col_outp.strides[0][1])
126133
+ w4_idx * int(im2col_outp.strides[0][0])
127134
+ c4_idx;
128135

backends/vulkan/runtime/graph/ops/impl/Q8taConv2d.cpp

Lines changed: 66 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,52 @@ bool q8ta_conv2d_check_4w4c_packed_dim_info(const api::PackedDimInfo& info) {
4242
info.outer_packed_dim_block_size == 4;
4343
}
4444

45+
namespace {
46+
47+
uint64_t q8ta_conv2d_im2col_scratch_limit(ComputeGraph& graph) {
48+
constexpr uint64_t kMaxBatchedIm2ColScratchBytes = 32ULL * 1024ULL * 1024ULL;
49+
const uint64_t device_scratch_limit =
50+
graph.context()->adapter_ptr()->max_buffer_numel();
51+
return device_scratch_limit < kMaxBatchedIm2ColScratchBytes
52+
? device_scratch_limit
53+
: kMaxBatchedIm2ColScratchBytes;
54+
}
55+
56+
bool should_use_q8ta_conv2d_im2col(
57+
ComputeGraph& graph,
58+
const int64_t batch,
59+
const int64_t groups,
60+
const int64_t in_channels_per_group,
61+
const int64_t flattened_kernel_size,
62+
const int64_t out_height,
63+
const int64_t out_width) {
64+
const bool im2col_eligible = in_channels_per_group % 4 == 0;
65+
if (!im2col_eligible) {
66+
return false;
67+
}
68+
69+
const int64_t spatial_out = out_height * out_width;
70+
if (batch > 1) {
71+
constexpr int64_t kMinFlattenedKernelSize = 1024;
72+
constexpr int64_t kMaxSpatialOutput = 64;
73+
const uint64_t scratch_bytes = static_cast<uint64_t>(batch) *
74+
static_cast<uint64_t>(flattened_kernel_size) *
75+
static_cast<uint64_t>(out_height) *
76+
static_cast<uint64_t>(utils::align_up_4(out_width));
77+
return groups == 1 && flattened_kernel_size >= kMinFlattenedKernelSize &&
78+
spatial_out <= kMaxSpatialOutput &&
79+
scratch_bytes <= q8ta_conv2d_im2col_scratch_limit(graph);
80+
}
81+
82+
if (graph.device_is_mali()) {
83+
return true;
84+
}
85+
86+
return groups == 1 && (in_channels_per_group >= 32 || spatial_out <= 4096);
87+
}
88+
89+
} // namespace
90+
4591
//
4692
// Workgroup size selection functions
4793
//
@@ -70,13 +116,16 @@ GlobalWorkGrid pick_q8ta_conv2d_gwg(
70116
const uint32_t W = graph->size_at<uint32_t>(-1, output);
71117
const uint32_t H = graph->size_at<uint32_t>(-2, output);
72118
const uint32_t C = graph->size_at<uint32_t>(-3, output);
119+
const uint32_t N = graph->size_at<uint32_t>(-4, output);
73120

74121
// Each thread processes 4 adjacent width positions and 4 channels (4Wx4C
75122
// tile)
76123
const uint32_t W4 = utils::div_up_4(W);
77124
const uint32_t C4 = utils::div_up_4(C);
78125

79-
return GlobalWorkGrid({W4, H, C4}, kTiledWorkGrid);
126+
return GlobalWorkGrid(
127+
{W4, utils::safe_downcast<uint32_t>(static_cast<uint64_t>(H) * N), C4},
128+
kTiledWorkGrid);
80129
}
81130

82131
/**
@@ -466,33 +515,33 @@ void q8ta_conv2d_general(
466515

467516
void q8ta_conv2d(ComputeGraph& graph, const std::vector<ValueRef>& args) {
468517
const ValueRef input = args.at(0);
518+
const ValueRef kernel_size_ref = args.at(9);
469519
const ValueRef groups_ref = args.at(13);
470520
const ValueRef output = args.at(15);
471521

472522
const int64_t groups = graph.extract_scalar<int64_t>(groups_ref);
473523
const int64_t in_channels = graph.size_at<int64_t>(-3, input);
474524
const int64_t in_channels_per_group = in_channels / groups;
525+
const int64_t batch = graph.size_at<int64_t>(-4, input);
475526

476527
const int64_t H_out = graph.size_at<int64_t>(-2, output);
477528
const int64_t W_out = graph.size_at<int64_t>(-1, output);
478-
const int64_t spatial_out = H_out * W_out;
479-
480-
// Im2col requires input channels per group to be a multiple of 4
481-
const bool im2col_eligible = in_channels_per_group % 4 == 0;
482-
483-
bool use_im2col = false;
484-
if (graph.device_is_mali()) {
485-
// On Mali, im2col is faster than the general shader across the board.
486-
use_im2col = im2col_eligible;
487-
} else {
488-
// Default: on Adreno and unknown GPU architectures, im2col is only
489-
// beneficial for ungrouped convolutions with sufficient channel depth or
490-
// small spatial output. For grouped convolutions, the general shader is
491-
// more efficient (0.7-0.95x regression measured on Adreno).
492-
use_im2col = im2col_eligible && groups == 1 &&
493-
(in_channels_per_group >= 32 || spatial_out <= 4096);
529+
int64_t flattened_kernel_size;
530+
{
531+
const auto kernel_size = graph.get_int_list(kernel_size_ref);
532+
flattened_kernel_size = utils::align_up_4(
533+
in_channels_per_group * kernel_size->at(0) * kernel_size->at(1));
494534
}
495535

536+
const bool use_im2col = should_use_q8ta_conv2d_im2col(
537+
graph,
538+
batch,
539+
groups,
540+
in_channels_per_group,
541+
flattened_kernel_size,
542+
H_out,
543+
W_out);
544+
496545
if (use_im2col) {
497546
q8ta_conv2d_im2col(graph, args);
498547
} else {

backends/vulkan/runtime/graph/ops/impl/Q8taConv2dDW.cpp

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,13 +34,16 @@ GlobalWorkGrid pick_q8ta_conv2d_dw_gwg(
3434
const uint32_t W = graph->size_at<uint32_t>(-1, output);
3535
const uint32_t H = graph->size_at<uint32_t>(-2, output);
3636
const uint32_t C = graph->size_at<uint32_t>(-3, output);
37+
const uint32_t N = graph->size_at<uint32_t>(-4, output);
3738

3839
// Each thread processes 4 adjacent width positions and 4 channels (4Wx4C
3940
// tile)
4041
const uint32_t W4 = utils::div_up_4(W);
4142
const uint32_t C4 = utils::div_up_4(C);
4243

43-
return GlobalWorkGrid({W4, H, C4}, kTiledWorkGrid);
44+
return GlobalWorkGrid(
45+
{W4, utils::safe_downcast<uint32_t>(static_cast<uint64_t>(H) * N), C4},
46+
kTiledWorkGrid);
4447
}
4548

4649
LocalWorkGroup pick_q8ta_conv2d_dw_lwg(

0 commit comments

Comments
 (0)