forked from tinygrad/tinygrad
313 lines
11 KiB
Python
313 lines
11 KiB
Python
import pyopencl as cl
|
|
import numpy as np
|
|
from ..tensor import Function
|
|
from ..llops.gpu import GPUBuffer, clbuild, buffer_new, unary_op, binary_op, reduce_op, perm_axis, inner_slice
|
|
|
|
i32 = np.int32
|
|
|
|
# ************* unary ops *************
|
|
|
|
class UnaryOp(Function):
|
|
def forward(ctx, input):
|
|
ctx.save_for_backward(input)
|
|
return unary_op(ctx, ctx.fop, input)
|
|
|
|
def backward(ctx, grad_output):
|
|
input, = ctx.saved_tensors
|
|
return binary_op(ctx, ctx.bop, input, grad_output)
|
|
|
|
class ReLU(UnaryOp):
|
|
fop = 'max(a, (float)0.)'
|
|
bop = 'b * (a >= 0)'
|
|
|
|
class Log(UnaryOp):
|
|
fop = 'log(a)'
|
|
bop = 'b / a'
|
|
|
|
class Exp(UnaryOp):
|
|
fop = 'exp(a)'
|
|
bop = 'b * exp(a)'
|
|
|
|
# ************* reduce ops *************
|
|
|
|
class Sum(Function):
|
|
def forward(ctx, input, axis=None):
|
|
ctx.save_for_backward(input.shape)
|
|
return reduce_op(ctx, "out += a", "out", input, axis=axis)
|
|
|
|
def backward(ctx, grad_output):
|
|
shape_input, = ctx.saved_tensors
|
|
# NOTE: the b buffer_new isn't used, since this is just for broadcast
|
|
return binary_op(ctx, 'a', grad_output, buffer_new(ctx, shape_input))
|
|
|
|
class Max(Function):
|
|
def forward(ctx, input, axis=None):
|
|
ret = reduce_op(ctx, "out = max(a,out)", "out", input, axis=axis, start="-INFINITY")
|
|
ctx.save_for_backward(input, axis, ret)
|
|
return ret
|
|
|
|
def backward(ctx, grad_output):
|
|
input, axis, ret = ctx.saved_tensors
|
|
ret2 = binary_op(ctx, "1.0*(a==b)", input, ret)
|
|
div = reduce_op(ctx, "out += a", "out+1e-10", ret2, axis=axis)
|
|
ret3 = binary_op(ctx, "a/b", ret2, div)
|
|
return binary_op(ctx, 'a*b', ret3, grad_output)
|
|
|
|
# ************* binary ops *************
|
|
|
|
def unbroadcast(ctx, out, in_sh):
|
|
sum_axis = [i for i in range(len(in_sh)) if in_sh[i]==1 and out.shape[i]>1] if in_sh != (1,) else None
|
|
return reduce_op(ctx, "out += a", "out", out, sum_axis)
|
|
|
|
class Add(Function):
|
|
def forward(ctx, x, y):
|
|
ctx.save_for_backward(x.shape, y.shape)
|
|
return binary_op(ctx, 'a+b', x, y)
|
|
|
|
def backward(ctx, grad_output):
|
|
grad_x, grad_y = grad_output, grad_output
|
|
shape_x, shape_y = ctx.saved_tensors
|
|
return unbroadcast(ctx, grad_x, shape_x), unbroadcast(ctx, grad_y, shape_y)
|
|
|
|
class Sub(Function):
|
|
def forward(ctx, x, y):
|
|
ctx.save_for_backward(x.shape, y.shape)
|
|
return binary_op(ctx, 'a-b', x, y)
|
|
|
|
def backward(ctx, grad_output):
|
|
grad_x, grad_y = grad_output, unary_op(ctx, '-a', grad_output)
|
|
shape_x, shape_y = ctx.saved_tensors
|
|
return unbroadcast(ctx, grad_x, shape_x), unbroadcast(ctx, grad_y, shape_y)
|
|
|
|
class Mul(Function):
|
|
def forward(ctx, x, y):
|
|
ctx.save_for_backward(x, y)
|
|
return binary_op(ctx, 'a*b', x, y)
|
|
|
|
def backward(ctx, grad_output):
|
|
x,y = ctx.saved_tensors
|
|
grad_x = binary_op(ctx, 'a*b', y, grad_output)
|
|
grad_y = binary_op(ctx, 'a*b', x, grad_output)
|
|
return unbroadcast(ctx, grad_x, x.shape), unbroadcast(ctx, grad_y, y.shape)
|
|
|
|
class Pow(Function):
|
|
def forward(ctx, x, y):
|
|
ctx.save_for_backward(x, y)
|
|
return binary_op(ctx, 'pow(a,b)', x, y)
|
|
|
|
def backward(ctx, grad_output):
|
|
x,y = ctx.saved_tensors
|
|
grad_x_inter = binary_op(ctx, 'b * (pow((float)a, (float)(b-1.0)))', x, y)
|
|
grad_x = binary_op(ctx, 'a*b', grad_output, grad_x_inter)
|
|
grad_y_inter = binary_op(ctx, 'pow(a, (float)b) * log(a);', x, y)
|
|
grad_y = binary_op(ctx, 'a*b', grad_output, grad_y_inter)
|
|
return unbroadcast(ctx, grad_x, x.shape), unbroadcast(ctx, grad_y, y.shape)
|
|
|
|
# ************* movement ops *************
|
|
|
|
class Reshape(Function):
|
|
def forward(ctx, x, shape):
|
|
ctx.save_for_backward(x.shape)
|
|
shape = tuple(-np.prod(x.shape) // np.prod(shape) if s == -1 else s for s in shape)
|
|
r = GPUBuffer(shape, hostbuf=x) # NOTE: this is not a copy
|
|
assert np.prod(x.shape) == np.prod(r.shape)
|
|
return r
|
|
|
|
def backward(ctx, grad_output):
|
|
in_shape, = ctx.saved_tensors
|
|
return GPUBuffer(in_shape, hostbuf=grad_output)
|
|
|
|
|
|
class Transpose(Function):
|
|
def forward(ctx, x, order=(1,0)):
|
|
ctx.save_for_backward(order)
|
|
return perm_axis(ctx, x, order)
|
|
|
|
def backward(ctx, grad_output):
|
|
return perm_axis(ctx, grad_output, np.argsort(ctx.order))
|
|
|
|
class Slice(Function):
|
|
def forward(ctx, x, arg=None):
|
|
ctx.save_for_backward(x.shape)
|
|
return inner_slice(ctx, x, arg)
|
|
|
|
def backward(ctx, grad_output):
|
|
shape, = ctx.saved_tensors
|
|
narg = [(0-p[0], grad_output.shape[i]+(shape[i]-p[1])) for i,p in enumerate(ctx.arg)]
|
|
return inner_slice(ctx, grad_output, narg)
|
|
|
|
# ************* processing ops *************
|
|
|
|
class Matmul(Function):
|
|
def forward(ctx, input, weight):
|
|
assert input.shape[-1] == weight.shape[-2]
|
|
cnt = np.prod(input.shape[0:-2]) if len(input.shape) > 2 else 1
|
|
isize, msize, osize = i32(input.shape[-2]), i32(input.shape[-1]), i32(weight.shape[-1])
|
|
ret = buffer_new(ctx, list(input.shape[0:-2])+[isize, osize])
|
|
|
|
matmul = clbuild("matmul", """
|
|
__kernel void matmul(
|
|
__global const float *input, __global const float *weight, __global float *res,
|
|
int isize, int is0, int is1, int msize, int ws0, int ws1, int osize
|
|
) {
|
|
int stride = get_global_id(2);
|
|
|
|
int X = get_global_id(0); // isize
|
|
int Y = get_global_id(1); // osize
|
|
|
|
float ret = 0.0;
|
|
for (int x = 0; x < msize; x++) {
|
|
ret += input[X * is0 + x * is1 + isize*msize*stride] *
|
|
weight[Y * ws0 + x * ws1 + msize*osize*stride];
|
|
}
|
|
|
|
res[X * osize + Y + isize*osize*stride] = ret;
|
|
}""")
|
|
ctx.save_for_backward(input, weight, matmul, cnt)
|
|
|
|
# (isize,msize) x (msize,osize) = (isize,osize)
|
|
matmul([isize, osize, cnt], None,
|
|
input.cl, weight.cl, ret.cl, isize,
|
|
msize, i32(1), msize, i32(1), osize, osize)
|
|
return ret
|
|
|
|
def backward(ctx, grad_output):
|
|
input, weight, matmul, cnt = ctx.saved_tensors
|
|
isize, msize, osize = i32(input.shape[-2]), i32(input.shape[-1]), i32(weight.shape[-1])
|
|
|
|
grad_input = buffer_new(ctx, input.shape)
|
|
grad_weight = buffer_new(ctx, weight.shape)
|
|
|
|
# (isize,osize) x (msize,osize) = (isize,msize)
|
|
matmul([isize, msize, cnt], None,
|
|
grad_output.cl, weight.cl, grad_input.cl, isize,
|
|
osize, i32(1), osize, osize, i32(1), msize)
|
|
|
|
# (isize,msize) x (isize,osize) = (msize,osize)
|
|
matmul([msize, osize, cnt], None,
|
|
input.cl, grad_output.cl, grad_weight.cl, msize,
|
|
i32(1), msize, isize, i32(1), osize, osize)
|
|
|
|
return grad_input, grad_weight
|
|
|
|
class Conv2D(Function):
|
|
def forward(ctx, x, w, stride=1, groups=1):
|
|
if isinstance(ctx.stride, int): ctx.stride = (ctx.stride, ctx.stride)
|
|
cout,cin,H,W = w.shape
|
|
ys,xs = ctx.stride
|
|
bs,cin_,iy,ix = x.shape
|
|
oy,ox = (iy-(H-ys))//ys, (ix-(W-xs))//xs
|
|
if cin*ctx.groups != cin_: raise Exception(f"Input Tensor shape {x.shape} does not match the shape of the weights {w.shape}. ({cin*ctx.groups} vs. {cin_})")
|
|
assert cout % ctx.groups == 0
|
|
rcout = cout//ctx.groups
|
|
|
|
ctx.save_for_backward(x,w)
|
|
|
|
# output buffer
|
|
ret = buffer_new(ctx, (bs, cout, oy, ox))
|
|
|
|
# input = (bs, groups, cin, iy, ix)
|
|
# weight = (groups, rcout, cin, H, W)
|
|
# output = (bs, groups, rcout, oy, ox)
|
|
|
|
conv = clbuild("conv", """
|
|
__kernel void conv(__global const float *input, __global const float *weight, __global float *output,
|
|
int H, int W, int groups, int rcout, int cin, int oy, int ox, int iy, int ix, int ys, int xs) {
|
|
|
|
int B = get_global_id(0)/(groups*rcout); // range 0-bs
|
|
int g = (get_global_id(0)/rcout)%groups;
|
|
int c = get_global_id(0) % rcout;
|
|
|
|
int Y = get_global_id(1); // range 0-oy
|
|
int X = get_global_id(2); // range 0-ox
|
|
int IY = Y*ys;
|
|
int IX = X*xs;
|
|
|
|
float acc = 0.0;
|
|
for (int ci = 0; ci < cin; ci++) {
|
|
for (int y = IY; y < IY+H; y++) {
|
|
for (int x = IX; x < IX+W; x++) {
|
|
acc += input[B*groups*cin*iy*ix + g*cin*iy*ix + ci*iy*ix + y*ix + x] * \
|
|
weight[g*rcout*cin*H*W + c*cin*H*W + ci*H*W + (y-IY)*W + (x-IX)];
|
|
}
|
|
}
|
|
}
|
|
output[B*groups*rcout*oy*ox + g*rcout*oy*ox + c*oy*ox + Y*ox + X] = acc;
|
|
}""")
|
|
|
|
conv([bs*groups*rcout, oy, ox], None,
|
|
x.cl, w.cl, ret.cl,
|
|
i32(H), i32(W), i32(groups), i32(rcout), i32(cin),
|
|
i32(oy), i32(ox), i32(iy), i32(ix), i32(ys), i32(xs)
|
|
)
|
|
return ret
|
|
|
|
def backward(ctx, grad_output):
|
|
bs,_,oy,ox = grad_output.shape
|
|
x, w = ctx.saved_tensors
|
|
cout,cin,H,W = w.shape
|
|
ys,xs = ctx.stride
|
|
bs,cin_,iy,ix = x.shape
|
|
oy,ox = (iy-(H-ys))//ys, (ix-(W-xs))//xs
|
|
assert cin*ctx.groups == cin_
|
|
assert cout % ctx.groups == 0
|
|
rcout = cout//ctx.groups
|
|
|
|
dx = buffer_new(ctx, (bs, cin_, iy, ix), zero=True)
|
|
dw = buffer_new(ctx, (cout, cin, H, W))
|
|
|
|
# tensx = (bs, groups*cin, iy, ix)
|
|
# tensw = (groups*rcout, cin, H, W)
|
|
# ggg = (bs, groups*rout, oy, ox)
|
|
|
|
convw = clbuild("convw", """
|
|
__kernel void convw(__global const float *tensx, __global const float *ggg, __global float *dw,
|
|
int H, int W, int groups, int rcout, int cin, int oy, int ox, int iy, int ix, int ys, int xs, int bs) {
|
|
|
|
int g = get_global_id(0)/(rcout*cin) ; // range 0-groups
|
|
int c = (get_global_id(0)/(cin)) %rcout; // range 0-rcout
|
|
int ci = get_global_id(0) % cin; // range 0-cin
|
|
int y = get_global_id(1); // range 0-H
|
|
int x = get_global_id(2); // range 0-W
|
|
|
|
float acc = 0.0;
|
|
for (int Y = 0; Y < oy; Y++) {
|
|
for (int X = 0; X < ox; X++) {
|
|
for (int B = 0; B < bs; B++) {
|
|
acc += ggg[B*groups*rcout*oy*ox + +g*rcout*oy*ox + c*oy*ox + Y*ox + X] * \
|
|
tensx[B*groups*cin*iy*ix + g*cin*iy*ix + ci*iy*ix + (Y*ys+y)*ix + X*xs+x];
|
|
}
|
|
}
|
|
}
|
|
dw[get_global_id(0)*H*W + y*W + x] = acc;
|
|
}""")
|
|
convx = clbuild("convx", """
|
|
__kernel void convx(__global const float *tensw, __global const float *ggg, __global float *dx,
|
|
int H, int W, int groups, int rcout, int cin, int oy, int ox, int iy, int ix, int ys, int xs, int bs) {
|
|
|
|
int B = get_global_id(0);
|
|
int g = get_global_id(1);
|
|
int ci = get_global_id(2);
|
|
|
|
for (int Y = 0; Y < oy; Y++) {
|
|
for (int X = 0; X < ox; X++) {
|
|
for (int y = 0; y < H; y++) {
|
|
for (int x = 0; x < W; x++) {
|
|
float acc = 0.0;
|
|
for (int c = 0; c < rcout; c++) {
|
|
acc += ggg[B*groups*rcout*oy*ox + g*rcout*oy*ox + c*oy*ox + Y*ox + X] * \
|
|
tensw[g*rcout*cin*H*W + c*cin*H*W + ci*H*W + y*W + x];
|
|
}
|
|
dx[B*groups*cin*iy*ix + g*cin*iy*ix + ci*iy*ix + (Y*ys+y)*ix + X*xs+x] += acc;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
""")
|
|
|
|
conv_args = i32(H), i32(W), i32(ctx.groups), i32(rcout), i32(cin), i32(oy), i32(ox), i32(iy), i32(ix), i32(ys), i32(xs), i32(bs)
|
|
convw([ctx.groups*rcout*cin, H, W], None, x.cl, grad_output.cl, dw.cl, *conv_args)
|
|
convx([bs, ctx.groups, cin], None, w.cl, grad_output.cl, dx.cl, *conv_args)
|
|
return dx, dw
|