Fix small matrices matmuls, imp3 working but slow

This commit is contained in:
Corentin 2021-06-28 19:05:39 +09:00
parent a3f7793c17
commit 3962ee70af
5 changed files with 77 additions and 219 deletions

View file

@ -5,7 +5,7 @@ import numpy as np
class MatMulOp:
def __init__(self, manager: kp.Manager, tile_size: int = -1, thread_work_ratio: int = 8):
def __init__(self, manager: kp.Manager, tile_size: int = -1, thread_work_ratio: int = 16):
self.mgr = manager
props = self.mgr.get_device_properties()
@ -14,9 +14,9 @@ class MatMulOp:
if tile_size < 0:
tile_size = 1
local_size_y = tile_size // thread_work_ratio
while (4 * tile_size * local_size_y <= max_workgroup_invocation
while (4 * tile_size * tile_size <= max_workgroup_invocation
and 2 * tile_size <= max_workgroup_size[0]
and 2 * local_size_y <= max_workgroup_size[1]):
and 2 * tile_size <= max_workgroup_size[1]):
tile_size *= 2
local_size_y = tile_size // thread_work_ratio
else:
@ -32,10 +32,10 @@ class MatMulOp:
self.local_size_x = tile_size
self.local_size_y = tile_size // thread_work_ratio
self.shader = f'''
self.shader = '''
#version 450
layout (local_size_x = {tile_size}, local_size_y = {self.local_size_y}) in;
layout (local_size_x = {tile_size}, local_size_y = {local_size_y}) in;
layout (set = 0, binding = 0) readonly buffer buf_in_tensor_1 {{ float in_tensor_1[]; }};
layout (set = 0, binding = 1) readonly buffer buf_in_tensor_2 {{ float in_tensor_2[]; }};
@ -51,14 +51,13 @@ void main()
uint row = gl_LocalInvocationID.x;
uint col = gl_LocalInvocationID.y;
uint globalRow = {tile_size} * gl_WorkGroupID.x + row;
uint globalCol = {tile_size} * gl_WorkGroupID.y + row;
uint globalCol = {tile_size} * gl_WorkGroupID.y + col;
uint tensor_size = uint(tensor_size_f);
float acc[{thread_work_ratio}];
for(uint w = 0u; w < {thread_work_ratio}; w++)
acc[w] = 0.0;
/*
uint numTiles = tensor_size / {tile_size};
for(uint t = 0u; t < numTiles; t++)
{{
@ -66,10 +65,10 @@ void main()
{{
uint tiledRow = {tile_size} * t + row;
uint tiledCol = {tile_size} * t + col;
sub_tensor_1[col + t * {self.local_size_y}][row] = in_tensor_1[
(tiledCol + w * {self.local_size_y}) * tensor_size + globalRow];
sub_tensor_2[col + t * {self.local_size_y}][row] = in_tensor_2[
(globalCol + w * {self.local_size_y})* tensor_size + tiledRow];
sub_tensor_1[col + w * {local_size_y}][row] = in_tensor_1[
(tiledCol + w * {local_size_y}) * tensor_size + globalRow];
sub_tensor_2[col + w * {local_size_y}][row] = in_tensor_2[
(globalCol + w * {local_size_y})* tensor_size + tiledRow];
}}
memoryBarrierShared();
@ -77,17 +76,15 @@ void main()
for(uint k = 0u; k < {tile_size}; k++)
for(uint w = 0u; w < {thread_work_ratio}; w++)
acc[w] += sub_tensor_1[k][row] * sub_tensor_2[col + w * {self.local_size_y}][k];
acc[w] += sub_tensor_1[k][row] * sub_tensor_2[col + w * {local_size_y}][k];
barrier();
}}*/
for(uint w = 0u; w < {thread_work_ratio}; w++)
{{
//out_tensor[(globalCol + w * {self.local_size_y}) * tensor_size + globalRow] = acc[w];
out_tensor[(globalCol + w * {self.local_size_y}) * tensor_size + globalRow] = w;
}}
for(uint w = 0u; w < {thread_work_ratio}; w++)
out_tensor[(globalCol + w * {local_size_y}) * tensor_size + globalRow] = acc[w];
}}'''
self.compiled_shader = kp.Shader.compile_source(self.shader)
self.compiled_shader = kp.Shader.compile_source(self.shader.format(
tile_size=tile_size, thread_work_ratio=thread_work_ratio, local_size_y=local_size_y))
self.tensor_shape: tuple[int, int] = (0, 0)
self.params: list[kp.Tensor] = []
self.algo = None
@ -99,9 +96,12 @@ void main()
if self.algo is None or self.tensor_shape != tensor_shape or self.params != params:
self.tensor_shape = tensor_shape
self.params = params
# workgroup = (tensor_shape[0] // self.local_size_x, tensor_shape[1] // self.local_size_y, 1)
workgroup = (2, 2, 1)
print(tensor_shape, self.local_size_x, self.local_size_y, workgroup)
tile_size = min(self.tensor_shape[0], self.tile_size)
thread_work_ratio = min(self.tensor_shape[1] // self.tile_size, self.thread_work_ratio)
local_size_y = tile_size // thread_work_ratio
self.compiled_shader = kp.Shader.compile_source(self.shader.format(
tile_size=tile_size, thread_work_ratio=thread_work_ratio, local_size_y=local_size_y))
workgroup = (tensor_shape[0] // self.local_size_x, tensor_shape[1] // self.local_size_y, 1)
self.algo = self.mgr.algorithm(
params, # params
self.compiled_shader, # spirv
@ -110,7 +110,7 @@ void main()
[]) # push_consts
(self.mgr.sequence()
.record(kp.OpTensorSyncDevice(self.params))
.record(kp.OpTensorSyncDevice([tensor_in_1, tensor_in_2]))
.record(kp.OpAlgoDispatch(self.algo))
.record(kp.OpTensorSyncLocal([tensor_out]))
.eval())
@ -133,7 +133,7 @@ def main():
matmul_op(tensor_shape, tensor_in_1, tensor_in_2, tensor_out)
experiment_count = 8
experiment_count = 2
start_time = time.time()
for _ in range(experiment_count):
matmul_op(tensor_shape, tensor_in_1, tensor_in_2, tensor_out)