UNPKG

@shumai/shumai

Version:

A fast, network-connected, differentiable tensor library for TypeScript (and JavaScript). Built with bun + flashlight for software engineers and researchers alike.

324 lines (307 loc) 11.9 kB
import type { Tensor } from './tensor' import * as base from './tensor' import * as ops from './tensor_ops' const sm = { ...base, ...ops } type ArgType = Tensor | number | number[] | BigInt64Array | boolean export interface GradContext { forward_inputs: [Tensor, ...ArgType[]] forward_output: Tensor backward_input: Tensor // the associated gradient of forward_output backward_output_index: number // // index of the associated forward input to be differentiated } function recoverShape(tensor: Tensor, originalShape: number[], lostAxes: number[]) { const shapeForBroadcast = [...originalShape] for (let axis of lostAxes) { if (axis < 0) { axis += originalShape.length } shapeForBroadcast[axis] = 1 } const tensorForBroadcast = tensor.reshape(shapeForBroadcast) return tensorForBroadcast.add(sm.full(originalShape, 0)) } function possiblyReduce(grad_out: Tensor, ctx: GradContext) { const input = <Tensor>ctx.forward_inputs[ctx.backward_output_index] const new_shape = input.shape if (ctx.backward_input.shape.length != input.shape.length) { for (let i = 0; i < ctx.backward_input.shape.length - input.shape.length; ++i) { new_shape.unshift(1) } } const reduction_axes = [] for (let i = 0; i < new_shape.length; ++i) { if (new_shape[i] === 1 && ctx.backward_input.shape[i] !== 1) { reduction_axes.push(i) } } if (reduction_axes.length) { return grad_out.sum(reduction_axes, true).reshape(input.shape) } return grad_out.reshape(input.shape) } const impls = { absolute: (ctx: GradContext) => { return ctx.backward_input.mul( sm.where( ctx.forward_inputs[0].greaterThan(sm.scalar(0)), sm.full(ctx.forward_inputs[0].shape, 1), sm.full(ctx.forward_inputs[0].shape, -1) ) ) }, add: (ctx: GradContext) => { return possiblyReduce(ctx.backward_input, ctx) }, amax: (ctx: GradContext) => { const inShape = (<Tensor>ctx.forward_inputs[0]).shape let axes = <number[]>ctx.forward_inputs[1] if (axes === undefined || axes.length === 0) { axes = inShape.map((x, i) => i) // All axes } return ctx.backward_input.mul( ctx.forward_inputs[0].eq(ctx.forward_output).astype(ctx.backward_input.dtype) ) }, conv2d: (ctx: GradContext) => { const [x, w, sx, sy, px, py, dx, dy, g] = < [Tensor, Tensor, number, number, number, number, number, number, number] >ctx.forward_inputs try { if (ctx.backward_output_index == 0) { return sm.conv2dBackwardData(ctx.backward_input, x, w, sx, sy, px, py, dx, dy, g) } else if (ctx.backward_output_index == 1) { return sm.conv2dBackwardFilter(ctx.backward_input, x, w, sx, sy, px, py, dx, dy, g) } } catch (e) { console.warn("Couldn't use native conv2d backward, falling back...") } if (dx !== 1 || dy !== 1) { throw new Error( `cannot differentiate convolution with dilation (${dx}, ${dy}), please file an issue.` ) } else if (sx !== sy) { throw new Error( `cannot differentiate convolution with stride (${sx}, ${sy}), please file an issue.` ) } else if (px !== py) { throw new Error( `cannot differentiate convolution with padding (${px}, ${py}), please file an issue.` ) } /* eslint-disable @typescript-eslint/no-unused-vars */ const batch = x.shape[0] const channel_out = ctx.backward_input.shape[1] const channel_in = x.shape[1] /* eslint-enable @typescript-eslint/no-unused-vars */ const k = w.shape[2] if (ctx.backward_output_index === 0) { const padding = k - 1 - (sx - 1) - px let dxgrad = ctx.backward_input if (sx > 1) { dxgrad = dxgrad.reshape(dxgrad.shape.concat([1, 1])) dxgrad = dxgrad .pad([ [0, 0], [0, 0], [sx - 1, 0], [sx - 1, 0] ]) .transpose([3, 4]) const shape = ctx.backward_input.shape dxgrad = dxgrad.reshape(shape.slice(0, 2).concat(shape.slice(2).map((x) => x * 2))) dxgrad = dxgrad.pad([ [0, 0], [0, 0], [0, sx - 1], [0, sx - 1] ]) } let dxw = w.flip(2).flip(3) if (g > 1) { dxw = dxw.reshape([g, dxw.shape[0] / g].concat(dxw.shape.slice(1))) dxw = dxw.transpose([0, 2]) dxw = dxw.reshape([dxw.shape[0] * dxw.shape[1]].concat(dxw.shape.slice(2))) } else { dxw = dxw.transpose([1, 0, 2, 3]) } return sm.conv2d(dxgrad, dxw, 1, 1, padding, padding, 1, 1, g) } const dwgrad = ctx.backward_input.transpose([0, 1]) let dwx = x.transpose([0, 1]) if (g > 1) { dwx = x.reshape([x.shape[0], g, x.shape[1] / g].concat(x.shape.slice(2))) dwx = dwx.transpose([0, 2]) dwx = dwx.reshape([dwx.shape[0], dwx.shape[1] * dwx.shape[2]].concat(dwx.shape.slice(3))) } return sm.conv2d(dwx, dwgrad, 1, 1, px, px, sx, sx, g) }, div: (ctx: GradContext) => { const recip = sm.scalar(1).div(<Tensor>ctx.forward_inputs[1]) const go = ctx.backward_input.mul(recip) if (ctx.backward_output_index === 0) { return possiblyReduce(go, ctx) } else if (ctx.backward_output_index === 1) { return possiblyReduce(go.negative().mul(recip), ctx) } }, sqrt: (ctx: GradContext): Tensor => { return ctx.backward_input.div(ctx.forward_output.mul(sm.scalar(2))) }, exp: (ctx: GradContext) => { return sm.exp(ctx.forward_inputs[0]).mul(ctx.backward_input) }, log: (ctx: GradContext) => { return ctx.backward_input.mul(sm.scalar(1).div(ctx.forward_inputs[0])) }, matmul: (ctx: GradContext) => { if (ctx.backward_output_index === 0) { const y = <Tensor>ctx.forward_inputs[1] if (ctx.backward_input.shape.length === 1 && y.shape.length === 1) { // backward_input and y are 1D column vectors const expandedGradIn = ctx.backward_input.reshape([ctx.backward_input.shape[0], 1]) const expandedY = y.reshape([y.shape[0], 1]) return expandedGradIn.matmul(expandedY.T()) } return ctx.backward_input.matmul(y.T()) // this is 1D if backward_input is a 1D row vector } else if (ctx.backward_output_index === 1) { const x = <Tensor>ctx.forward_inputs[0] if (ctx.backward_input.shape.length === 1 && x.shape.length === 1) { // backward_input and x are 1D row vectors const expandedGradIn = ctx.backward_input.reshape([1, ctx.backward_input.shape[0]]) const expandedX = x.reshape([1, x.shape[0]]) return expandedX.T().matmul(expandedGradIn) } return x.T().matmul(ctx.backward_input) // this is 1D if backward_input is a 1D column vector } else { throw new Error(`Invalid GradContext argument`) } }, maximum: (ctx: GradContext) => { const a_idx = ctx.backward_output_index const b_idx = <0 | 1>(1 - ctx.backward_output_index) const mask = (<Tensor>ctx.forward_inputs[a_idx]).greaterThan(<Tensor>ctx.forward_inputs[b_idx]) return mask.mul(ctx.backward_input) }, mean: (ctx: GradContext) => { const inShape = (<Tensor>ctx.forward_inputs[0]).shape let axes = <number[]>ctx.forward_inputs[1] if (axes === undefined || axes.length === 0) { axes = inShape.map((x, i) => i) // All axes } let num = 1 for (const axis of axes) { num *= inShape[axis] } return recoverShape(ctx.backward_input.div(sm.scalar(num)), inShape, axes) }, var: (ctx: GradContext) => { const input = <Tensor>ctx.forward_inputs[0] const inShape = input.shape let axes = <number[]>ctx.forward_inputs[1] if (axes.length === 0) { axes = inShape.map((x, i) => i) } const bias = <boolean>ctx.forward_inputs[2] let num = 1 for (const axis of axes) { num *= inShape[axis] } if (bias) { num -= 1 } const expandedGradIn = recoverShape(ctx.backward_input, inShape, axes) const expandedMean = recoverShape(input.mean(axes), inShape, axes) return expandedGradIn.mul(sm.scalar(2 / num)).mul(input.sub(expandedMean)) }, mul: (ctx: GradContext) => { const backward_output_index = <0 | 1>(1 - ctx.backward_output_index) return possiblyReduce( (<Tensor>ctx.forward_inputs[backward_output_index]).mul(ctx.backward_input), ctx ) }, greaterThanEqual: (ctx: GradContext) => { return ctx.backward_input.mul(ctx.forward_output).mul(sm.scalar(1).sub(ctx.forward_output)) }, logicalNot: (ctx: GradContext) => { return ctx.backward_input.mul(ctx.forward_output).mul(sm.scalar(1).sub(ctx.forward_output)) }, sigmoid: (ctx: GradContext) => { return ctx.backward_input.mul(ctx.forward_output).mul(sm.scalar(1).sub(ctx.forward_output)) }, clip: (ctx: GradContext) => { const result = <Tensor>ctx.forward_output const low = <Tensor>ctx.forward_inputs[1] const high = <Tensor>ctx.forward_inputs[2] const lowMask = result.greaterThan(low) const highMask = result.lessThan(high) const lowHighMask = lowMask.bitwiseAnd(highMask) const gradMask = sm.where(lowHighMask, ctx.forward_output, sm.full(ctx.forward_output.shape, 0)) return gradMask }, erf: (ctx: GradContext) => { const input = <Tensor>ctx.forward_inputs[0] return ctx.backward_input .mul(sm.scalar(2)) .div(sm.sqrt(sm.scalar(Math.PI))) .mul(sm.exp(input.mul(input).mul(sm.scalar(-1)))) }, minimum: (ctx: GradContext) => { const input = <Tensor>ctx.forward_inputs[0] const rhsVal = <Tensor>ctx.forward_inputs[1] const mask = input.lessThan(rhsVal).astype(ctx.backward_input.dtype) return mask.mul(ctx.backward_input) }, sub: (ctx: GradContext) => { if (ctx.backward_output_index) { return possiblyReduce(ctx.backward_input.negative(), ctx) } return possiblyReduce(ctx.backward_input, ctx) }, sum: (ctx: GradContext) => { const inShape = ctx.forward_inputs[0].shape let axes = <number[]>ctx.forward_inputs[1] if (axes.length === 0) { axes = inShape.map((x, i) => i) // All axes } return recoverShape(ctx.backward_input, inShape, axes) }, tanh: (ctx: GradContext) => { return sm.scalar(1).sub(ctx.forward_output.mul(ctx.forward_output)).mul(ctx.backward_input) }, concatenate: (ctx: GradContext): Tensor => { const axis = <number>ctx.forward_inputs[ctx.forward_inputs.length - 1] const { backward_output_index } = ctx const prevTensors = <Tensor[]>ctx.forward_inputs.slice(0, backward_output_index) const start = prevTensors.reduce((r, t) => r + t.shape[axis], 0) const end = start + (<Tensor>ctx.forward_inputs[backward_output_index]).shape[axis] const range = ctx.forward_output.shape.map(() => ':') range[axis] = start + ':' + end return ctx.backward_input.index(range) }, transpose: (ctx: GradContext): Tensor => { const forwardAxes = <number[]>ctx.forward_inputs[1] const reverseAxes = [...forwardAxes] for (let i = 0; i < forwardAxes.length; i++) { reverseAxes[forwardAxes[i]] = i } // If forwardAxes === [], reverseAxes === [] return ctx.backward_input.transpose(reverseAxes) }, reshape: (ctx: GradContext): Tensor => { const inShape = (<Tensor>ctx.forward_inputs[0]).shape return ctx.backward_input.reshape(inShape) }, where: (ctx: GradContext): Tensor => { const zeros = sm.full(ctx.backward_input.shape, 0) if (ctx.backward_output_index === 0) { throw new Error(`Gradient cannot be propagated to the cond Tensor`) } else if (ctx.backward_output_index === 1) { return sm.where(ctx.forward_inputs[0], ctx.backward_input, zeros) } else if (ctx.backward_output_index === 2) { return sm.where(ctx.forward_inputs[0], zeros, ctx.backward_input) } else { throw new Error(`Invalid Grad argument`) } } } Object.assign(sm.gradient_functions, impls)