mirror of
https://github.com/zhw2590582/ArtPlayer.git
synced 2026-10-08 19:06:15 -08:00
- Created a new documentation file for the danmuku plugin with code examples. - Added a new documentation file for language settings (i18n) with detailed instructions on language imports, default languages, and how to add or modify languages. - Included code snippets for initializing the Artplayer with different language settings.
49176 lines
1.6 MiB
Plaintext
49176 lines
1.6 MiB
Plaintext
/*!
|
|
* artplayer-plugin-danmuku-mask.js v1.1.0
|
|
* Github: https://github.com/zhw2590582/ArtPlayer
|
|
* (c) 2017-2026 Harvey Zhao
|
|
* Released under the MIT License.
|
|
*/
|
|
function _mergeNamespaces(n, m) {
|
|
for (var i = 0; i < m.length; i++) {
|
|
const e = m[i];
|
|
if (typeof e !== "string" && !Array.isArray(e)) {
|
|
for (const k in e) {
|
|
if (k !== "default" && !(k in n)) {
|
|
const d = Object.getOwnPropertyDescriptor(e, k);
|
|
if (d) {
|
|
Object.defineProperty(n, k, d.get ? d : {
|
|
enumerable: true,
|
|
get: () => e[k]
|
|
});
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return Object.freeze(Object.defineProperty(n, Symbol.toStringTag, { value: "Module" }));
|
|
}
|
|
const EPSILON_FLOAT32$1 = 1e-7;
|
|
const EPSILON_FLOAT16$1 = 1e-4;
|
|
class DataStorage {
|
|
constructor(backend2, dataMover) {
|
|
this.backend = backend2;
|
|
this.dataMover = dataMover;
|
|
this.data = /* @__PURE__ */ new WeakMap();
|
|
this.dataIdsCount = 0;
|
|
}
|
|
get(dataId) {
|
|
if (!this.data.has(dataId)) {
|
|
this.dataMover.moveData(this.backend, dataId);
|
|
}
|
|
return this.data.get(dataId);
|
|
}
|
|
set(dataId, value) {
|
|
this.dataIdsCount++;
|
|
this.data.set(dataId, value);
|
|
}
|
|
has(dataId) {
|
|
return this.data.has(dataId);
|
|
}
|
|
delete(dataId) {
|
|
this.dataIdsCount--;
|
|
return this.data.delete(dataId);
|
|
}
|
|
numDataIds() {
|
|
return this.dataIdsCount;
|
|
}
|
|
}
|
|
class KernelBackend {
|
|
refCount(dataId) {
|
|
return notYetImplemented("refCount");
|
|
}
|
|
incRef(dataId) {
|
|
return notYetImplemented("incRef");
|
|
}
|
|
timerAvailable() {
|
|
return true;
|
|
}
|
|
time(f) {
|
|
return notYetImplemented("time");
|
|
}
|
|
read(dataId) {
|
|
return notYetImplemented("read");
|
|
}
|
|
readSync(dataId) {
|
|
return notYetImplemented("readSync");
|
|
}
|
|
readToGPU(dataId, options) {
|
|
return notYetImplemented("readToGPU");
|
|
}
|
|
numDataIds() {
|
|
return notYetImplemented("numDataIds");
|
|
}
|
|
disposeData(dataId, force) {
|
|
return notYetImplemented("disposeData");
|
|
}
|
|
write(values, shape, dtype) {
|
|
return notYetImplemented("write");
|
|
}
|
|
move(dataId, values, shape, dtype, refCount) {
|
|
return notYetImplemented("move");
|
|
}
|
|
createTensorFromGPUData(values, shape, dtype) {
|
|
return notYetImplemented("createTensorFromGPUData");
|
|
}
|
|
memory() {
|
|
return notYetImplemented("memory");
|
|
}
|
|
/** Returns the highest precision for floats in bits (e.g. 16 or 32) */
|
|
floatPrecision() {
|
|
return notYetImplemented("floatPrecision");
|
|
}
|
|
/** Returns the smallest representable number. */
|
|
epsilon() {
|
|
return this.floatPrecision() === 32 ? EPSILON_FLOAT32$1 : EPSILON_FLOAT16$1;
|
|
}
|
|
dispose() {
|
|
return notYetImplemented("dispose");
|
|
}
|
|
}
|
|
function notYetImplemented(kernelName) {
|
|
throw new Error(`'${kernelName}' not yet implemented or not found in the registry. This kernel may not be supported by the tfjs backend you have chosen`);
|
|
}
|
|
function clamp(min2, x, max2) {
|
|
return Math.max(min2, Math.min(x, max2));
|
|
}
|
|
function nearestLargerEven(val) {
|
|
return val % 2 === 0 ? val : val + 1;
|
|
}
|
|
function swap(object, left, right) {
|
|
const temp = object[left];
|
|
object[left] = object[right];
|
|
object[right] = temp;
|
|
}
|
|
function sum$3(arr) {
|
|
let sum2 = 0;
|
|
for (let i = 0; i < arr.length; i++) {
|
|
sum2 += arr[i];
|
|
}
|
|
return sum2;
|
|
}
|
|
function assert(expr, msg) {
|
|
if (!expr) {
|
|
throw new Error(typeof msg === "string" ? msg : msg());
|
|
}
|
|
}
|
|
function assertShapesMatch(shapeA, shapeB, errorMessagePrefix = "") {
|
|
assert(arraysEqual(shapeA, shapeB), () => errorMessagePrefix + ` Shapes ${shapeA} and ${shapeB} must match`);
|
|
}
|
|
function assertNonNull(a) {
|
|
assert(a != null, () => `The input to the tensor constructor must be a non-null value.`);
|
|
}
|
|
function sizeFromShape(shape) {
|
|
if (shape.length === 0) {
|
|
return 1;
|
|
}
|
|
let size = shape[0];
|
|
for (let i = 1; i < shape.length; i++) {
|
|
size *= shape[i];
|
|
}
|
|
return size;
|
|
}
|
|
function arraysEqualWithNull(n1, n2) {
|
|
if (n1 === n2) {
|
|
return true;
|
|
}
|
|
if (n1 == null || n2 == null) {
|
|
return false;
|
|
}
|
|
if (n1.length !== n2.length) {
|
|
return false;
|
|
}
|
|
for (let i = 0; i < n1.length; i++) {
|
|
if (n1[i] !== null && n2[i] !== null && n1[i] !== n2[i]) {
|
|
return false;
|
|
}
|
|
}
|
|
return true;
|
|
}
|
|
function arraysEqual(n1, n2) {
|
|
if (n1 === n2) {
|
|
return true;
|
|
}
|
|
if (n1 == null || n2 == null) {
|
|
return false;
|
|
}
|
|
if (n1.length !== n2.length) {
|
|
return false;
|
|
}
|
|
for (let i = 0; i < n1.length; i++) {
|
|
if (n1[i] !== n2[i]) {
|
|
return false;
|
|
}
|
|
}
|
|
return true;
|
|
}
|
|
function isInt(a) {
|
|
return a % 1 === 0;
|
|
}
|
|
function sizeToSquarishShape(size) {
|
|
const width = Math.ceil(Math.sqrt(size));
|
|
return [width, Math.ceil(size / width)];
|
|
}
|
|
function rightPad(a, size) {
|
|
if (size <= a.length) {
|
|
return a;
|
|
}
|
|
return a + " ".repeat(size - a.length);
|
|
}
|
|
function repeatedTry(checkFn, delayFn = (counter) => 0, maxCounter, scheduleFn) {
|
|
return new Promise((resolve, reject) => {
|
|
let tryCount = 0;
|
|
const tryFn = () => {
|
|
if (checkFn()) {
|
|
resolve();
|
|
return;
|
|
}
|
|
tryCount++;
|
|
const nextBackoff = delayFn(tryCount);
|
|
if (maxCounter != null && tryCount >= maxCounter) {
|
|
reject();
|
|
return;
|
|
}
|
|
if (scheduleFn != null) {
|
|
scheduleFn(tryFn, nextBackoff);
|
|
} else {
|
|
setTimeout(tryFn, nextBackoff);
|
|
}
|
|
};
|
|
tryFn();
|
|
});
|
|
}
|
|
function inferFromImplicitShape(shape, size) {
|
|
let shapeProd = 1;
|
|
let implicitIdx = -1;
|
|
for (let i = 0; i < shape.length; ++i) {
|
|
if (shape[i] >= 0) {
|
|
shapeProd *= shape[i];
|
|
} else if (shape[i] === -1) {
|
|
if (implicitIdx !== -1) {
|
|
throw Error(`Shapes can only have 1 implicit size. Found -1 at dim ${implicitIdx} and dim ${i}`);
|
|
}
|
|
implicitIdx = i;
|
|
} else if (shape[i] < 0) {
|
|
throw Error(`Shapes can not be < 0. Found ${shape[i]} at dim ${i}`);
|
|
}
|
|
}
|
|
if (implicitIdx === -1) {
|
|
if (size > 0 && size !== shapeProd) {
|
|
throw Error(`Size(${size}) must match the product of shape ${shape}`);
|
|
}
|
|
return shape;
|
|
}
|
|
if (shapeProd === 0) {
|
|
throw Error(`Cannot infer the missing size in [${shape}] when there are 0 elements`);
|
|
}
|
|
if (size % shapeProd !== 0) {
|
|
throw Error(`The implicit shape can't be a fractional number. Got ${size} / ${shapeProd}`);
|
|
}
|
|
const newShape = shape.slice();
|
|
newShape[implicitIdx] = size / shapeProd;
|
|
return newShape;
|
|
}
|
|
function parseAxisParam(axis, shape) {
|
|
const rank = shape.length;
|
|
axis = axis == null ? shape.map((s, i) => i) : [].concat(axis);
|
|
assert(axis.every((ax) => ax >= -rank && ax < rank), () => `All values in axis param must be in range [-${rank}, ${rank}) but got axis ${axis}`);
|
|
assert(axis.every((ax) => isInt(ax)), () => `All values in axis param must be integers but got axis ${axis}`);
|
|
return axis.map((a) => a < 0 ? rank + a : a);
|
|
}
|
|
function squeezeShape(shape, axis) {
|
|
const newShape = [];
|
|
const keptDims = [];
|
|
const isEmptyArray = axis != null && Array.isArray(axis) && axis.length === 0;
|
|
const axes = axis == null || isEmptyArray ? null : parseAxisParam(axis, shape).sort();
|
|
let j2 = 0;
|
|
for (let i = 0; i < shape.length; ++i) {
|
|
if (axes != null) {
|
|
if (axes[j2] === i && shape[i] !== 1) {
|
|
throw new Error(`Can't squeeze axis ${i} since its dim '${shape[i]}' is not 1`);
|
|
}
|
|
if ((axes[j2] == null || axes[j2] > i) && shape[i] === 1) {
|
|
newShape.push(shape[i]);
|
|
keptDims.push(i);
|
|
}
|
|
if (axes[j2] <= i) {
|
|
j2++;
|
|
}
|
|
}
|
|
if (shape[i] !== 1) {
|
|
newShape.push(shape[i]);
|
|
keptDims.push(i);
|
|
}
|
|
}
|
|
return { newShape, keptDims };
|
|
}
|
|
function getTypedArrayFromDType(dtype, size) {
|
|
return getArrayFromDType(dtype, size);
|
|
}
|
|
function getArrayFromDType(dtype, size) {
|
|
let values = null;
|
|
if (dtype == null || dtype === "float32") {
|
|
values = new Float32Array(size);
|
|
} else if (dtype === "int32") {
|
|
values = new Int32Array(size);
|
|
} else if (dtype === "bool") {
|
|
values = new Uint8Array(size);
|
|
} else if (dtype === "string") {
|
|
values = new Array(size);
|
|
} else {
|
|
throw new Error(`Unknown data type ${dtype}`);
|
|
}
|
|
return values;
|
|
}
|
|
function checkConversionForErrors(vals, dtype) {
|
|
for (let i = 0; i < vals.length; i++) {
|
|
const num = vals[i];
|
|
if (isNaN(num) || !isFinite(num)) {
|
|
throw Error(`A tensor of type ${dtype} being uploaded contains ${num}.`);
|
|
}
|
|
}
|
|
}
|
|
function isValidDtype(dtype) {
|
|
return dtype === "bool" || dtype === "complex64" || dtype === "float32" || dtype === "int32" || dtype === "string";
|
|
}
|
|
function hasEncodingLoss(oldType, newType) {
|
|
if (newType === "complex64") {
|
|
return false;
|
|
}
|
|
if (newType === "float32" && oldType !== "complex64") {
|
|
return false;
|
|
}
|
|
if (newType === "int32" && oldType !== "float32" && oldType !== "complex64") {
|
|
return false;
|
|
}
|
|
if (newType === "bool" && oldType === "bool") {
|
|
return false;
|
|
}
|
|
return true;
|
|
}
|
|
function bytesPerElement(dtype) {
|
|
if (dtype === "float32" || dtype === "int32") {
|
|
return 4;
|
|
} else if (dtype === "complex64") {
|
|
return 8;
|
|
} else if (dtype === "bool") {
|
|
return 1;
|
|
} else {
|
|
throw new Error(`Unknown dtype ${dtype}`);
|
|
}
|
|
}
|
|
function bytesFromStringArray(arr) {
|
|
if (arr == null) {
|
|
return 0;
|
|
}
|
|
let bytes = 0;
|
|
arr.forEach((x) => bytes += x.length);
|
|
return bytes;
|
|
}
|
|
function isString(value) {
|
|
return typeof value === "string" || value instanceof String;
|
|
}
|
|
function isBoolean(value) {
|
|
return typeof value === "boolean";
|
|
}
|
|
function isNumber(value) {
|
|
return typeof value === "number";
|
|
}
|
|
function inferDtype(values) {
|
|
if (Array.isArray(values)) {
|
|
return inferDtype(values[0]);
|
|
}
|
|
if (values instanceof Float32Array) {
|
|
return "float32";
|
|
} else if (values instanceof Int32Array || values instanceof Uint8Array || values instanceof Uint8ClampedArray) {
|
|
return "int32";
|
|
} else if (isNumber(values)) {
|
|
return "float32";
|
|
} else if (isString(values)) {
|
|
return "string";
|
|
} else if (isBoolean(values)) {
|
|
return "bool";
|
|
}
|
|
return "float32";
|
|
}
|
|
function isFunction(f) {
|
|
return !!(f && f.constructor && f.call && f.apply);
|
|
}
|
|
function nearestDivisor(size, start) {
|
|
for (let i = start; i < size; ++i) {
|
|
if (size % i === 0) {
|
|
return i;
|
|
}
|
|
}
|
|
return size;
|
|
}
|
|
function computeStrides(shape) {
|
|
const rank = shape.length;
|
|
if (rank < 2) {
|
|
return [];
|
|
}
|
|
const strides = new Array(rank - 1);
|
|
strides[rank - 2] = shape[rank - 1];
|
|
for (let i = rank - 3; i >= 0; --i) {
|
|
strides[i] = strides[i + 1] * shape[i + 1];
|
|
}
|
|
return strides;
|
|
}
|
|
function createNestedArray(offset, shape, a, isComplex = false) {
|
|
const ret = new Array();
|
|
if (shape.length === 1) {
|
|
const d = shape[0] * (isComplex ? 2 : 1);
|
|
for (let i = 0; i < d; i++) {
|
|
ret[i] = a[offset + i];
|
|
}
|
|
} else {
|
|
const d = shape[0];
|
|
const rest = shape.slice(1);
|
|
const len = rest.reduce((acc, c) => acc * c) * (isComplex ? 2 : 1);
|
|
for (let i = 0; i < d; i++) {
|
|
ret[i] = createNestedArray(offset + i * len, rest, a, isComplex);
|
|
}
|
|
}
|
|
return ret;
|
|
}
|
|
function toNestedArray(shape, a, isComplex = false) {
|
|
if (shape.length === 0) {
|
|
return a[0];
|
|
}
|
|
const size = shape.reduce((acc, c) => acc * c) * (isComplex ? 2 : 1);
|
|
if (size === 0) {
|
|
return [];
|
|
}
|
|
if (size !== a.length) {
|
|
throw new Error(`[${shape}] does not match the input size ${a.length}${isComplex ? " for a complex tensor" : ""}.`);
|
|
}
|
|
return createNestedArray(0, shape, a, isComplex);
|
|
}
|
|
function convertBackendValuesAndArrayBuffer(data, dtype) {
|
|
if (Array.isArray(data)) {
|
|
return data;
|
|
}
|
|
if (dtype === "float32") {
|
|
return data instanceof Float32Array ? data : new Float32Array(data);
|
|
} else if (dtype === "int32") {
|
|
return data instanceof Int32Array ? data : new Int32Array(data);
|
|
} else if (dtype === "bool" || dtype === "string") {
|
|
return Uint8Array.from(new Int32Array(data));
|
|
} else {
|
|
throw new Error(`Unknown dtype ${dtype}`);
|
|
}
|
|
}
|
|
function makeOnesTypedArray(size, dtype) {
|
|
const array = makeZerosTypedArray(size, dtype);
|
|
for (let i = 0; i < array.length; i++) {
|
|
array[i] = 1;
|
|
}
|
|
return array;
|
|
}
|
|
function makeZerosTypedArray(size, dtype) {
|
|
if (dtype == null || dtype === "float32" || dtype === "complex64") {
|
|
return new Float32Array(size);
|
|
} else if (dtype === "int32") {
|
|
return new Int32Array(size);
|
|
} else if (dtype === "bool") {
|
|
return new Uint8Array(size);
|
|
} else {
|
|
throw new Error(`Unknown data type ${dtype}`);
|
|
}
|
|
}
|
|
function makeZerosNestedTypedArray(shape, dtype) {
|
|
const size = shape.reduce((prev, curr) => prev * curr, 1);
|
|
if (dtype == null || dtype === "float32") {
|
|
return toNestedArray(shape, new Float32Array(size));
|
|
} else if (dtype === "int32") {
|
|
return toNestedArray(shape, new Int32Array(size));
|
|
} else if (dtype === "bool") {
|
|
return toNestedArray(shape, new Uint8Array(size));
|
|
} else {
|
|
throw new Error(`Unknown data type ${dtype}`);
|
|
}
|
|
}
|
|
function assertNonNegativeIntegerDimensions(shape) {
|
|
shape.forEach((dimSize) => {
|
|
assert(Number.isInteger(dimSize) && dimSize >= 0, () => `Tensor must have a shape comprised of positive integers but got shape [${shape}].`);
|
|
});
|
|
}
|
|
function locToIndex(locs, rank, strides) {
|
|
if (rank === 0) {
|
|
return 0;
|
|
} else if (rank === 1) {
|
|
return locs[0];
|
|
}
|
|
let index = locs[locs.length - 1];
|
|
for (let i = 0; i < locs.length - 1; ++i) {
|
|
index += strides[i] * locs[i];
|
|
}
|
|
return index;
|
|
}
|
|
function indexToLoc(index, rank, strides) {
|
|
if (rank === 0) {
|
|
return [];
|
|
} else if (rank === 1) {
|
|
return [index];
|
|
}
|
|
const locs = new Array(rank);
|
|
for (let i = 0; i < locs.length - 1; ++i) {
|
|
locs[i] = Math.floor(index / strides[i]);
|
|
index -= locs[i] * strides[i];
|
|
}
|
|
locs[locs.length - 1] = index;
|
|
return locs;
|
|
}
|
|
function isPromise(object) {
|
|
return object && object.then && typeof object.then === "function";
|
|
}
|
|
const TENSORFLOWJS_FLAGS_PREFIX = "tfjsflags";
|
|
class Environment {
|
|
// tslint:disable-next-line: no-any
|
|
constructor(global2) {
|
|
this.global = global2;
|
|
this.flags = {};
|
|
this.flagRegistry = {};
|
|
this.urlFlags = {};
|
|
this.getQueryParams = getQueryParams;
|
|
this.populateURLFlags();
|
|
}
|
|
setPlatform(platformName, platform) {
|
|
if (this.platform != null) {
|
|
if (!(env().getBool("IS_TEST") || env().getBool("PROD"))) {
|
|
console.warn(`Platform ${this.platformName} has already been set. Overwriting the platform with ${platformName}.`);
|
|
}
|
|
}
|
|
this.platformName = platformName;
|
|
this.platform = platform;
|
|
}
|
|
registerFlag(flagName, evaluationFn, setHook) {
|
|
this.flagRegistry[flagName] = { evaluationFn, setHook };
|
|
if (this.urlFlags[flagName] != null) {
|
|
const flagValue = this.urlFlags[flagName];
|
|
if (!(env().getBool("IS_TEST") || env().getBool("PROD"))) {
|
|
console.warn(`Setting feature override from URL ${flagName}: ${flagValue}.`);
|
|
}
|
|
this.set(flagName, flagValue);
|
|
}
|
|
}
|
|
async getAsync(flagName) {
|
|
if (flagName in this.flags) {
|
|
return this.flags[flagName];
|
|
}
|
|
this.flags[flagName] = await this.evaluateFlag(flagName);
|
|
return this.flags[flagName];
|
|
}
|
|
get(flagName) {
|
|
if (flagName in this.flags) {
|
|
return this.flags[flagName];
|
|
}
|
|
const flagValue = this.evaluateFlag(flagName);
|
|
if (isPromise(flagValue)) {
|
|
throw new Error(`Flag ${flagName} cannot be synchronously evaluated. Please use getAsync() instead.`);
|
|
}
|
|
this.flags[flagName] = flagValue;
|
|
return this.flags[flagName];
|
|
}
|
|
getNumber(flagName) {
|
|
return this.get(flagName);
|
|
}
|
|
getBool(flagName) {
|
|
return this.get(flagName);
|
|
}
|
|
getString(flagName) {
|
|
return this.get(flagName);
|
|
}
|
|
getFlags() {
|
|
return this.flags;
|
|
}
|
|
// For backwards compatibility.
|
|
get features() {
|
|
return this.flags;
|
|
}
|
|
set(flagName, value) {
|
|
if (this.flagRegistry[flagName] == null) {
|
|
throw new Error(`Cannot set flag ${flagName} as it has not been registered.`);
|
|
}
|
|
this.flags[flagName] = value;
|
|
if (this.flagRegistry[flagName].setHook != null) {
|
|
this.flagRegistry[flagName].setHook(value);
|
|
}
|
|
}
|
|
evaluateFlag(flagName) {
|
|
if (this.flagRegistry[flagName] == null) {
|
|
throw new Error(`Cannot evaluate flag '${flagName}': no evaluation function found.`);
|
|
}
|
|
return this.flagRegistry[flagName].evaluationFn();
|
|
}
|
|
setFlags(flags) {
|
|
this.flags = Object.assign({}, flags);
|
|
}
|
|
reset() {
|
|
this.flags = {};
|
|
this.urlFlags = {};
|
|
this.populateURLFlags();
|
|
}
|
|
populateURLFlags() {
|
|
if (typeof this.global === "undefined" || typeof this.global.location === "undefined" || typeof this.global.location.search === "undefined") {
|
|
return;
|
|
}
|
|
const urlParams = this.getQueryParams(this.global.location.search);
|
|
if (TENSORFLOWJS_FLAGS_PREFIX in urlParams) {
|
|
const keyValues = urlParams[TENSORFLOWJS_FLAGS_PREFIX].split(",");
|
|
keyValues.forEach((keyValue) => {
|
|
const [key, value] = keyValue.split(":");
|
|
this.urlFlags[key] = parseValue(key, value);
|
|
});
|
|
}
|
|
}
|
|
}
|
|
function getQueryParams(queryString) {
|
|
const params = {};
|
|
queryString.replace(/[?&]([^=?&]+)(?:=([^&]*))?/g, (s, ...t2) => {
|
|
decodeParam(params, t2[0], t2[1]);
|
|
return t2.join("=");
|
|
});
|
|
return params;
|
|
}
|
|
function decodeParam(params, name, value) {
|
|
params[decodeURIComponent(name)] = decodeURIComponent(value || "");
|
|
}
|
|
function parseValue(flagName, value) {
|
|
const lowerCaseValue = value.toLowerCase();
|
|
if (lowerCaseValue === "true" || lowerCaseValue === "false") {
|
|
return lowerCaseValue === "true";
|
|
} else if (`${+lowerCaseValue}` === lowerCaseValue) {
|
|
return +lowerCaseValue;
|
|
} else {
|
|
return value;
|
|
}
|
|
}
|
|
function env() {
|
|
return ENV$3;
|
|
}
|
|
let ENV$3 = null;
|
|
function setEnvironmentGlobal(environment) {
|
|
ENV$3 = environment;
|
|
}
|
|
let globalNameSpace;
|
|
function getGlobalNamespace() {
|
|
if (globalNameSpace == null) {
|
|
let ns;
|
|
if (typeof window !== "undefined") {
|
|
ns = window;
|
|
} else if (typeof global !== "undefined") {
|
|
ns = global;
|
|
} else if (typeof process !== "undefined") {
|
|
ns = process;
|
|
} else if (typeof self !== "undefined") {
|
|
ns = self;
|
|
} else {
|
|
throw new Error("Could not find a global object");
|
|
}
|
|
globalNameSpace = ns;
|
|
}
|
|
return globalNameSpace;
|
|
}
|
|
function getGlobalMap() {
|
|
const ns = getGlobalNamespace();
|
|
if (ns._tfGlobals == null) {
|
|
ns._tfGlobals = /* @__PURE__ */ new Map();
|
|
}
|
|
return ns._tfGlobals;
|
|
}
|
|
function getGlobal(key, init) {
|
|
const globalMap = getGlobalMap();
|
|
if (globalMap.has(key)) {
|
|
return globalMap.get(key);
|
|
} else {
|
|
const singleton = init();
|
|
globalMap.set(key, singleton);
|
|
return globalMap.get(key);
|
|
}
|
|
}
|
|
const Abs = "Abs";
|
|
const Acos = "Acos";
|
|
const Acosh = "Acosh";
|
|
const Add = "Add";
|
|
const AddN = "AddN";
|
|
const All = "All";
|
|
const Any = "Any";
|
|
const ArgMax = "ArgMax";
|
|
const ArgMin = "ArgMin";
|
|
const Asin = "Asin";
|
|
const Asinh = "Asinh";
|
|
const Atan = "Atan";
|
|
const Atanh = "Atanh";
|
|
const Atan2 = "Atan2";
|
|
const AvgPool = "AvgPool";
|
|
const AvgPoolGrad = "AvgPoolGrad";
|
|
const AvgPool3D = "AvgPool3D";
|
|
const AvgPool3DGrad = "AvgPool3DGrad";
|
|
const BatchMatMul = "BatchMatMul";
|
|
const BatchToSpaceND = "BatchToSpaceND";
|
|
const Bincount = "Bincount";
|
|
const BitwiseAnd = "BitwiseAnd";
|
|
const BroadcastArgs = "BroadcastArgs";
|
|
const Cast = "Cast";
|
|
const Ceil = "Ceil";
|
|
const ClipByValue = "ClipByValue";
|
|
const Complex = "Complex";
|
|
const ComplexAbs = "ComplexAbs";
|
|
const Concat = "Concat";
|
|
const Conv2D = "Conv2D";
|
|
const Conv2DBackpropFilter = "Conv2DBackpropFilter";
|
|
const Conv2DBackpropInput = "Conv2DBackpropInput";
|
|
const Conv3D = "Conv3D";
|
|
const Conv3DBackpropFilterV2 = "Conv3DBackpropFilterV2";
|
|
const Conv3DBackpropInputV2 = "Conv3DBackpropInputV2";
|
|
const Cos = "Cos";
|
|
const Cosh = "Cosh";
|
|
const Cumprod = "Cumprod";
|
|
const Cumsum = "Cumsum";
|
|
const CropAndResize = "CropAndResize";
|
|
const DenseBincount = "DenseBincount";
|
|
const DepthToSpace = "DepthToSpace";
|
|
const DepthwiseConv2dNative = "DepthwiseConv2dNative";
|
|
const DepthwiseConv2dNativeBackpropFilter = "DepthwiseConv2dNativeBackpropFilter";
|
|
const DepthwiseConv2dNativeBackpropInput = "DepthwiseConv2dNativeBackpropInput";
|
|
const Diag = "Diag";
|
|
const Dilation2D = "Dilation2D";
|
|
const Dilation2DBackpropInput = "Dilation2DBackpropInput";
|
|
const Dilation2DBackpropFilter = "Dilation2DBackpropFilter";
|
|
const Draw = "Draw";
|
|
const RealDiv = "RealDiv";
|
|
const Einsum = "Einsum";
|
|
const Elu = "Elu";
|
|
const EluGrad = "EluGrad";
|
|
const Erf = "Erf";
|
|
const Equal = "Equal";
|
|
const Exp = "Exp";
|
|
const ExpandDims = "ExpandDims";
|
|
const Expm1 = "Expm1";
|
|
const FFT = "FFT";
|
|
const Fill = "Fill";
|
|
const FlipLeftRight = "FlipLeftRight";
|
|
const Floor = "Floor";
|
|
const FloorDiv = "FloorDiv";
|
|
const FusedBatchNorm = "FusedBatchNorm";
|
|
const GatherV2 = "GatherV2";
|
|
const GatherNd = "GatherNd";
|
|
const Greater = "Greater";
|
|
const GreaterEqual = "GreaterEqual";
|
|
const Identity = "Identity";
|
|
const IFFT = "IFFT";
|
|
const Imag = "Imag";
|
|
const IsFinite = "IsFinite";
|
|
const IsInf = "IsInf";
|
|
const IsNan = "IsNan";
|
|
const LeakyRelu = "LeakyRelu";
|
|
const Less = "Less";
|
|
const LessEqual = "LessEqual";
|
|
const LinSpace = "LinSpace";
|
|
const Log = "Log";
|
|
const Log1p = "Log1p";
|
|
const LogicalAnd = "LogicalAnd";
|
|
const LogicalNot = "LogicalNot";
|
|
const LogicalOr = "LogicalOr";
|
|
const LRN = "LRN";
|
|
const LRNGrad = "LRNGrad";
|
|
const Max = "Max";
|
|
const Maximum = "Maximum";
|
|
const MaxPool = "MaxPool";
|
|
const MaxPoolGrad = "MaxPoolGrad";
|
|
const MaxPool3D = "MaxPool3D";
|
|
const MaxPool3DGrad = "MaxPool3DGrad";
|
|
const MaxPoolWithArgmax = "MaxPoolWithArgmax";
|
|
const Mean = "Mean";
|
|
const Min = "Min";
|
|
const Minimum = "Minimum";
|
|
const MirrorPad = "MirrorPad";
|
|
const Mod = "Mod";
|
|
const Multinomial = "Multinomial";
|
|
const Multiply = "Multiply";
|
|
const Neg = "Neg";
|
|
const NotEqual = "NotEqual";
|
|
const NonMaxSuppressionV3 = "NonMaxSuppressionV3";
|
|
const NonMaxSuppressionV4 = "NonMaxSuppressionV4";
|
|
const NonMaxSuppressionV5 = "NonMaxSuppressionV5";
|
|
const OnesLike = "OnesLike";
|
|
const OneHot = "OneHot";
|
|
const Pack = "Pack";
|
|
const PadV2 = "PadV2";
|
|
const Pow = "Pow";
|
|
const Prelu = "Prelu";
|
|
const Prod = "Prod";
|
|
const RaggedGather = "RaggedGather";
|
|
const RaggedRange = "RaggedRange";
|
|
const RaggedTensorToTensor = "RaggedTensorToTensor";
|
|
const Range = "Range";
|
|
const Real = "Real";
|
|
const Reciprocal = "Reciprocal";
|
|
const Relu = "Relu";
|
|
const Reshape = "Reshape";
|
|
const ResizeNearestNeighbor = "ResizeNearestNeighbor";
|
|
const ResizeNearestNeighborGrad = "ResizeNearestNeighborGrad";
|
|
const ResizeBilinear = "ResizeBilinear";
|
|
const ResizeBilinearGrad = "ResizeBilinearGrad";
|
|
const Relu6 = "Relu6";
|
|
const Reverse = "Reverse";
|
|
const Round = "Round";
|
|
const Rsqrt = "Rsqrt";
|
|
const ScatterNd = "ScatterNd";
|
|
const TensorScatterUpdate = "TensorScatterUpdate";
|
|
const SearchSorted = "SearchSorted";
|
|
const Select = "Select";
|
|
const Selu = "Selu";
|
|
const Slice = "Slice";
|
|
const Sin = "Sin";
|
|
const Sinh = "Sinh";
|
|
const Sign = "Sign";
|
|
const Sigmoid = "Sigmoid";
|
|
const Softplus = "Softplus";
|
|
const Sqrt = "Sqrt";
|
|
const Sum = "Sum";
|
|
const SpaceToBatchND = "SpaceToBatchND";
|
|
const SplitV = "SplitV";
|
|
const Softmax = "Softmax";
|
|
const SparseFillEmptyRows = "SparseFillEmptyRows";
|
|
const SparseReshape = "SparseReshape";
|
|
const SparseSegmentMean = "SparseSegmentMean";
|
|
const SparseSegmentSum = "SparseSegmentSum";
|
|
const SparseToDense = "SparseToDense";
|
|
const SquaredDifference = "SquaredDifference";
|
|
const Square = "Square";
|
|
const StaticRegexReplace = "StaticRegexReplace";
|
|
const StridedSlice = "StridedSlice";
|
|
const StringNGrams = "StringNGrams";
|
|
const StringSplit = "StringSplit";
|
|
const StringToHashBucketFast = "StringToHashBucketFast";
|
|
const Sub = "Sub";
|
|
const Tan = "Tan";
|
|
const Tanh = "Tanh";
|
|
const Tile = "Tile";
|
|
const TopK = "TopK";
|
|
const Transform = "Transform";
|
|
const Transpose = "Transpose";
|
|
const Unique = "Unique";
|
|
const Unpack = "Unpack";
|
|
const UnsortedSegmentSum = "UnsortedSegmentSum";
|
|
const ZerosLike = "ZerosLike";
|
|
const Step = "Step";
|
|
const FromPixels = "FromPixels";
|
|
const RotateWithOffset = "RotateWithOffset";
|
|
const _FusedMatMul = "_FusedMatMul";
|
|
const FusedConv2D = "FusedConv2D";
|
|
const FusedDepthwiseConv2D = "FusedDepthwiseConv2D";
|
|
function warn(...msg) {
|
|
if (!(env().getBool("IS_TEST") || env().getBool("PROD"))) {
|
|
console.warn(...msg);
|
|
}
|
|
}
|
|
const kernelRegistry = getGlobal("kernelRegistry", () => /* @__PURE__ */ new Map());
|
|
const gradRegistry = getGlobal("gradRegistry", () => /* @__PURE__ */ new Map());
|
|
function getKernel(kernelName, backendName) {
|
|
const key = makeKey(kernelName, backendName);
|
|
return kernelRegistry.get(key);
|
|
}
|
|
function getGradient(kernelName) {
|
|
return gradRegistry.get(kernelName);
|
|
}
|
|
function getKernelsForBackend(backendName) {
|
|
const it2 = kernelRegistry.entries();
|
|
const result = [];
|
|
while (true) {
|
|
const { done, value } = it2.next();
|
|
if (done) {
|
|
break;
|
|
}
|
|
const [key, config] = value;
|
|
const [backend2] = key.split("_");
|
|
if (backend2 === backendName) {
|
|
result.push(config);
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
function registerKernel(config) {
|
|
const { kernelName, backendName } = config;
|
|
const key = makeKey(kernelName, backendName);
|
|
if (kernelRegistry.has(key)) {
|
|
warn(`The kernel '${kernelName}' for backend '${backendName}' is already registered`);
|
|
}
|
|
kernelRegistry.set(key, config);
|
|
}
|
|
function makeKey(kernelName, backendName) {
|
|
return `${backendName}_${kernelName}`;
|
|
}
|
|
function isTypedArrayBrowser(a) {
|
|
return a instanceof Float32Array || a instanceof Int32Array || a instanceof Uint8Array || a instanceof Uint8ClampedArray;
|
|
}
|
|
var commonjsGlobal = typeof globalThis !== "undefined" ? globalThis : typeof window !== "undefined" ? window : typeof global !== "undefined" ? global : typeof self !== "undefined" ? self : {};
|
|
function getDefaultExportFromCjs(x) {
|
|
return x && x.__esModule && Object.prototype.hasOwnProperty.call(x, "default") ? x["default"] : x;
|
|
}
|
|
function getAugmentedNamespace(n) {
|
|
if (Object.prototype.hasOwnProperty.call(n, "__esModule")) return n;
|
|
var f = n.default;
|
|
if (typeof f == "function") {
|
|
var a = function a6() {
|
|
var isInstance = false;
|
|
try {
|
|
isInstance = this instanceof a6;
|
|
} catch {
|
|
}
|
|
if (isInstance) {
|
|
return Reflect.construct(f, arguments, this.constructor);
|
|
}
|
|
return f.apply(this, arguments);
|
|
};
|
|
a.prototype = f.prototype;
|
|
} else a = {};
|
|
Object.defineProperty(a, "__esModule", { value: true });
|
|
Object.keys(n).forEach(function(k) {
|
|
var d = Object.getOwnPropertyDescriptor(n, k);
|
|
Object.defineProperty(a, k, d.get ? d : {
|
|
enumerable: true,
|
|
get: function() {
|
|
return n[k];
|
|
}
|
|
});
|
|
});
|
|
return a;
|
|
}
|
|
var long$1;
|
|
var hasRequiredLong;
|
|
function requireLong() {
|
|
if (hasRequiredLong) return long$1;
|
|
hasRequiredLong = 1;
|
|
long$1 = Long2;
|
|
var wasm = null;
|
|
try {
|
|
wasm = new WebAssembly.Instance(new WebAssembly.Module(new Uint8Array([
|
|
0,
|
|
97,
|
|
115,
|
|
109,
|
|
1,
|
|
0,
|
|
0,
|
|
0,
|
|
1,
|
|
13,
|
|
2,
|
|
96,
|
|
0,
|
|
1,
|
|
127,
|
|
96,
|
|
4,
|
|
127,
|
|
127,
|
|
127,
|
|
127,
|
|
1,
|
|
127,
|
|
3,
|
|
7,
|
|
6,
|
|
0,
|
|
1,
|
|
1,
|
|
1,
|
|
1,
|
|
1,
|
|
6,
|
|
6,
|
|
1,
|
|
127,
|
|
1,
|
|
65,
|
|
0,
|
|
11,
|
|
7,
|
|
50,
|
|
6,
|
|
3,
|
|
109,
|
|
117,
|
|
108,
|
|
0,
|
|
1,
|
|
5,
|
|
100,
|
|
105,
|
|
118,
|
|
95,
|
|
115,
|
|
0,
|
|
2,
|
|
5,
|
|
100,
|
|
105,
|
|
118,
|
|
95,
|
|
117,
|
|
0,
|
|
3,
|
|
5,
|
|
114,
|
|
101,
|
|
109,
|
|
95,
|
|
115,
|
|
0,
|
|
4,
|
|
5,
|
|
114,
|
|
101,
|
|
109,
|
|
95,
|
|
117,
|
|
0,
|
|
5,
|
|
8,
|
|
103,
|
|
101,
|
|
116,
|
|
95,
|
|
104,
|
|
105,
|
|
103,
|
|
104,
|
|
0,
|
|
0,
|
|
10,
|
|
191,
|
|
1,
|
|
6,
|
|
4,
|
|
0,
|
|
35,
|
|
0,
|
|
11,
|
|
36,
|
|
1,
|
|
1,
|
|
126,
|
|
32,
|
|
0,
|
|
173,
|
|
32,
|
|
1,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
32,
|
|
2,
|
|
173,
|
|
32,
|
|
3,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
126,
|
|
34,
|
|
4,
|
|
66,
|
|
32,
|
|
135,
|
|
167,
|
|
36,
|
|
0,
|
|
32,
|
|
4,
|
|
167,
|
|
11,
|
|
36,
|
|
1,
|
|
1,
|
|
126,
|
|
32,
|
|
0,
|
|
173,
|
|
32,
|
|
1,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
32,
|
|
2,
|
|
173,
|
|
32,
|
|
3,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
127,
|
|
34,
|
|
4,
|
|
66,
|
|
32,
|
|
135,
|
|
167,
|
|
36,
|
|
0,
|
|
32,
|
|
4,
|
|
167,
|
|
11,
|
|
36,
|
|
1,
|
|
1,
|
|
126,
|
|
32,
|
|
0,
|
|
173,
|
|
32,
|
|
1,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
32,
|
|
2,
|
|
173,
|
|
32,
|
|
3,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
128,
|
|
34,
|
|
4,
|
|
66,
|
|
32,
|
|
135,
|
|
167,
|
|
36,
|
|
0,
|
|
32,
|
|
4,
|
|
167,
|
|
11,
|
|
36,
|
|
1,
|
|
1,
|
|
126,
|
|
32,
|
|
0,
|
|
173,
|
|
32,
|
|
1,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
32,
|
|
2,
|
|
173,
|
|
32,
|
|
3,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
129,
|
|
34,
|
|
4,
|
|
66,
|
|
32,
|
|
135,
|
|
167,
|
|
36,
|
|
0,
|
|
32,
|
|
4,
|
|
167,
|
|
11,
|
|
36,
|
|
1,
|
|
1,
|
|
126,
|
|
32,
|
|
0,
|
|
173,
|
|
32,
|
|
1,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
32,
|
|
2,
|
|
173,
|
|
32,
|
|
3,
|
|
173,
|
|
66,
|
|
32,
|
|
134,
|
|
132,
|
|
130,
|
|
34,
|
|
4,
|
|
66,
|
|
32,
|
|
135,
|
|
167,
|
|
36,
|
|
0,
|
|
32,
|
|
4,
|
|
167,
|
|
11
|
|
])), {}).exports;
|
|
} catch (e) {
|
|
}
|
|
function Long2(low, high, unsigned) {
|
|
this.low = low | 0;
|
|
this.high = high | 0;
|
|
this.unsigned = !!unsigned;
|
|
}
|
|
Long2.prototype.__isLong__;
|
|
Object.defineProperty(Long2.prototype, "__isLong__", { value: true });
|
|
function isLong(obj) {
|
|
return (obj && obj["__isLong__"]) === true;
|
|
}
|
|
Long2.isLong = isLong;
|
|
var INT_CACHE = {};
|
|
var UINT_CACHE = {};
|
|
function fromInt(value, unsigned) {
|
|
var obj, cachedObj, cache;
|
|
if (unsigned) {
|
|
value >>>= 0;
|
|
if (cache = 0 <= value && value < 256) {
|
|
cachedObj = UINT_CACHE[value];
|
|
if (cachedObj)
|
|
return cachedObj;
|
|
}
|
|
obj = fromBits(value, (value | 0) < 0 ? -1 : 0, true);
|
|
if (cache)
|
|
UINT_CACHE[value] = obj;
|
|
return obj;
|
|
} else {
|
|
value |= 0;
|
|
if (cache = -128 <= value && value < 128) {
|
|
cachedObj = INT_CACHE[value];
|
|
if (cachedObj)
|
|
return cachedObj;
|
|
}
|
|
obj = fromBits(value, value < 0 ? -1 : 0, false);
|
|
if (cache)
|
|
INT_CACHE[value] = obj;
|
|
return obj;
|
|
}
|
|
}
|
|
Long2.fromInt = fromInt;
|
|
function fromNumber(value, unsigned) {
|
|
if (isNaN(value))
|
|
return unsigned ? UZERO : ZERO;
|
|
if (unsigned) {
|
|
if (value < 0)
|
|
return UZERO;
|
|
if (value >= TWO_PWR_64_DBL)
|
|
return MAX_UNSIGNED_VALUE;
|
|
} else {
|
|
if (value <= -TWO_PWR_63_DBL)
|
|
return MIN_VALUE;
|
|
if (value + 1 >= TWO_PWR_63_DBL)
|
|
return MAX_VALUE;
|
|
}
|
|
if (value < 0)
|
|
return fromNumber(-value, unsigned).neg();
|
|
return fromBits(value % TWO_PWR_32_DBL | 0, value / TWO_PWR_32_DBL | 0, unsigned);
|
|
}
|
|
Long2.fromNumber = fromNumber;
|
|
function fromBits(lowBits, highBits, unsigned) {
|
|
return new Long2(lowBits, highBits, unsigned);
|
|
}
|
|
Long2.fromBits = fromBits;
|
|
var pow_dbl = Math.pow;
|
|
function fromString(str, unsigned, radix) {
|
|
if (str.length === 0)
|
|
throw Error("empty string");
|
|
if (str === "NaN" || str === "Infinity" || str === "+Infinity" || str === "-Infinity")
|
|
return ZERO;
|
|
if (typeof unsigned === "number") {
|
|
radix = unsigned, unsigned = false;
|
|
} else {
|
|
unsigned = !!unsigned;
|
|
}
|
|
radix = radix || 10;
|
|
if (radix < 2 || 36 < radix)
|
|
throw RangeError("radix");
|
|
var p2;
|
|
if ((p2 = str.indexOf("-")) > 0)
|
|
throw Error("interior hyphen");
|
|
else if (p2 === 0) {
|
|
return fromString(str.substring(1), unsigned, radix).neg();
|
|
}
|
|
var radixToPower = fromNumber(pow_dbl(radix, 8));
|
|
var result = ZERO;
|
|
for (var i = 0; i < str.length; i += 8) {
|
|
var size = Math.min(8, str.length - i), value = parseInt(str.substring(i, i + size), radix);
|
|
if (size < 8) {
|
|
var power = fromNumber(pow_dbl(radix, size));
|
|
result = result.mul(power).add(fromNumber(value));
|
|
} else {
|
|
result = result.mul(radixToPower);
|
|
result = result.add(fromNumber(value));
|
|
}
|
|
}
|
|
result.unsigned = unsigned;
|
|
return result;
|
|
}
|
|
Long2.fromString = fromString;
|
|
function fromValue(val, unsigned) {
|
|
if (typeof val === "number")
|
|
return fromNumber(val, unsigned);
|
|
if (typeof val === "string")
|
|
return fromString(val, unsigned);
|
|
return fromBits(val.low, val.high, typeof unsigned === "boolean" ? unsigned : val.unsigned);
|
|
}
|
|
Long2.fromValue = fromValue;
|
|
var TWO_PWR_16_DBL = 1 << 16;
|
|
var TWO_PWR_24_DBL = 1 << 24;
|
|
var TWO_PWR_32_DBL = TWO_PWR_16_DBL * TWO_PWR_16_DBL;
|
|
var TWO_PWR_64_DBL = TWO_PWR_32_DBL * TWO_PWR_32_DBL;
|
|
var TWO_PWR_63_DBL = TWO_PWR_64_DBL / 2;
|
|
var TWO_PWR_24 = fromInt(TWO_PWR_24_DBL);
|
|
var ZERO = fromInt(0);
|
|
Long2.ZERO = ZERO;
|
|
var UZERO = fromInt(0, true);
|
|
Long2.UZERO = UZERO;
|
|
var ONE = fromInt(1);
|
|
Long2.ONE = ONE;
|
|
var UONE = fromInt(1, true);
|
|
Long2.UONE = UONE;
|
|
var NEG_ONE = fromInt(-1);
|
|
Long2.NEG_ONE = NEG_ONE;
|
|
var MAX_VALUE = fromBits(4294967295 | 0, 2147483647 | 0, false);
|
|
Long2.MAX_VALUE = MAX_VALUE;
|
|
var MAX_UNSIGNED_VALUE = fromBits(4294967295 | 0, 4294967295 | 0, true);
|
|
Long2.MAX_UNSIGNED_VALUE = MAX_UNSIGNED_VALUE;
|
|
var MIN_VALUE = fromBits(0, 2147483648 | 0, false);
|
|
Long2.MIN_VALUE = MIN_VALUE;
|
|
var LongPrototype = Long2.prototype;
|
|
LongPrototype.toInt = function toInt() {
|
|
return this.unsigned ? this.low >>> 0 : this.low;
|
|
};
|
|
LongPrototype.toNumber = function toNumber() {
|
|
if (this.unsigned)
|
|
return (this.high >>> 0) * TWO_PWR_32_DBL + (this.low >>> 0);
|
|
return this.high * TWO_PWR_32_DBL + (this.low >>> 0);
|
|
};
|
|
LongPrototype.toString = function toString(radix) {
|
|
radix = radix || 10;
|
|
if (radix < 2 || 36 < radix)
|
|
throw RangeError("radix");
|
|
if (this.isZero())
|
|
return "0";
|
|
if (this.isNegative()) {
|
|
if (this.eq(MIN_VALUE)) {
|
|
var radixLong = fromNumber(radix), div2 = this.div(radixLong), rem1 = div2.mul(radixLong).sub(this);
|
|
return div2.toString(radix) + rem1.toInt().toString(radix);
|
|
} else
|
|
return "-" + this.neg().toString(radix);
|
|
}
|
|
var radixToPower = fromNumber(pow_dbl(radix, 6), this.unsigned), rem = this;
|
|
var result = "";
|
|
while (true) {
|
|
var remDiv = rem.div(radixToPower), intval = rem.sub(remDiv.mul(radixToPower)).toInt() >>> 0, digits = intval.toString(radix);
|
|
rem = remDiv;
|
|
if (rem.isZero())
|
|
return digits + result;
|
|
else {
|
|
while (digits.length < 6)
|
|
digits = "0" + digits;
|
|
result = "" + digits + result;
|
|
}
|
|
}
|
|
};
|
|
LongPrototype.getHighBits = function getHighBits() {
|
|
return this.high;
|
|
};
|
|
LongPrototype.getHighBitsUnsigned = function getHighBitsUnsigned() {
|
|
return this.high >>> 0;
|
|
};
|
|
LongPrototype.getLowBits = function getLowBits() {
|
|
return this.low;
|
|
};
|
|
LongPrototype.getLowBitsUnsigned = function getLowBitsUnsigned() {
|
|
return this.low >>> 0;
|
|
};
|
|
LongPrototype.getNumBitsAbs = function getNumBitsAbs() {
|
|
if (this.isNegative())
|
|
return this.eq(MIN_VALUE) ? 64 : this.neg().getNumBitsAbs();
|
|
var val = this.high != 0 ? this.high : this.low;
|
|
for (var bit = 31; bit > 0; bit--)
|
|
if ((val & 1 << bit) != 0)
|
|
break;
|
|
return this.high != 0 ? bit + 33 : bit + 1;
|
|
};
|
|
LongPrototype.isZero = function isZero() {
|
|
return this.high === 0 && this.low === 0;
|
|
};
|
|
LongPrototype.eqz = LongPrototype.isZero;
|
|
LongPrototype.isNegative = function isNegative() {
|
|
return !this.unsigned && this.high < 0;
|
|
};
|
|
LongPrototype.isPositive = function isPositive() {
|
|
return this.unsigned || this.high >= 0;
|
|
};
|
|
LongPrototype.isOdd = function isOdd() {
|
|
return (this.low & 1) === 1;
|
|
};
|
|
LongPrototype.isEven = function isEven2() {
|
|
return (this.low & 1) === 0;
|
|
};
|
|
LongPrototype.equals = function equals(other) {
|
|
if (!isLong(other))
|
|
other = fromValue(other);
|
|
if (this.unsigned !== other.unsigned && this.high >>> 31 === 1 && other.high >>> 31 === 1)
|
|
return false;
|
|
return this.high === other.high && this.low === other.low;
|
|
};
|
|
LongPrototype.eq = LongPrototype.equals;
|
|
LongPrototype.notEquals = function notEquals(other) {
|
|
return !this.eq(
|
|
/* validates */
|
|
other
|
|
);
|
|
};
|
|
LongPrototype.neq = LongPrototype.notEquals;
|
|
LongPrototype.ne = LongPrototype.notEquals;
|
|
LongPrototype.lessThan = function lessThan(other) {
|
|
return this.comp(
|
|
/* validates */
|
|
other
|
|
) < 0;
|
|
};
|
|
LongPrototype.lt = LongPrototype.lessThan;
|
|
LongPrototype.lessThanOrEqual = function lessThanOrEqual(other) {
|
|
return this.comp(
|
|
/* validates */
|
|
other
|
|
) <= 0;
|
|
};
|
|
LongPrototype.lte = LongPrototype.lessThanOrEqual;
|
|
LongPrototype.le = LongPrototype.lessThanOrEqual;
|
|
LongPrototype.greaterThan = function greaterThan(other) {
|
|
return this.comp(
|
|
/* validates */
|
|
other
|
|
) > 0;
|
|
};
|
|
LongPrototype.gt = LongPrototype.greaterThan;
|
|
LongPrototype.greaterThanOrEqual = function greaterThanOrEqual(other) {
|
|
return this.comp(
|
|
/* validates */
|
|
other
|
|
) >= 0;
|
|
};
|
|
LongPrototype.gte = LongPrototype.greaterThanOrEqual;
|
|
LongPrototype.ge = LongPrototype.greaterThanOrEqual;
|
|
LongPrototype.compare = function compare(other) {
|
|
if (!isLong(other))
|
|
other = fromValue(other);
|
|
if (this.eq(other))
|
|
return 0;
|
|
var thisNeg = this.isNegative(), otherNeg = other.isNegative();
|
|
if (thisNeg && !otherNeg)
|
|
return -1;
|
|
if (!thisNeg && otherNeg)
|
|
return 1;
|
|
if (!this.unsigned)
|
|
return this.sub(other).isNegative() ? -1 : 1;
|
|
return other.high >>> 0 > this.high >>> 0 || other.high === this.high && other.low >>> 0 > this.low >>> 0 ? -1 : 1;
|
|
};
|
|
LongPrototype.comp = LongPrototype.compare;
|
|
LongPrototype.negate = function negate() {
|
|
if (!this.unsigned && this.eq(MIN_VALUE))
|
|
return MIN_VALUE;
|
|
return this.not().add(ONE);
|
|
};
|
|
LongPrototype.neg = LongPrototype.negate;
|
|
LongPrototype.add = function add2(addend) {
|
|
if (!isLong(addend))
|
|
addend = fromValue(addend);
|
|
var a48 = this.high >>> 16;
|
|
var a32 = this.high & 65535;
|
|
var a16 = this.low >>> 16;
|
|
var a00 = this.low & 65535;
|
|
var b48 = addend.high >>> 16;
|
|
var b32 = addend.high & 65535;
|
|
var b16 = addend.low >>> 16;
|
|
var b00 = addend.low & 65535;
|
|
var c48 = 0, c32 = 0, c16 = 0, c00 = 0;
|
|
c00 += a00 + b00;
|
|
c16 += c00 >>> 16;
|
|
c00 &= 65535;
|
|
c16 += a16 + b16;
|
|
c32 += c16 >>> 16;
|
|
c16 &= 65535;
|
|
c32 += a32 + b32;
|
|
c48 += c32 >>> 16;
|
|
c32 &= 65535;
|
|
c48 += a48 + b48;
|
|
c48 &= 65535;
|
|
return fromBits(c16 << 16 | c00, c48 << 16 | c32, this.unsigned);
|
|
};
|
|
LongPrototype.subtract = function subtract(subtrahend) {
|
|
if (!isLong(subtrahend))
|
|
subtrahend = fromValue(subtrahend);
|
|
return this.add(subtrahend.neg());
|
|
};
|
|
LongPrototype.sub = LongPrototype.subtract;
|
|
LongPrototype.multiply = function multiply2(multiplier) {
|
|
if (this.isZero())
|
|
return ZERO;
|
|
if (!isLong(multiplier))
|
|
multiplier = fromValue(multiplier);
|
|
if (wasm) {
|
|
var low = wasm.mul(
|
|
this.low,
|
|
this.high,
|
|
multiplier.low,
|
|
multiplier.high
|
|
);
|
|
return fromBits(low, wasm.get_high(), this.unsigned);
|
|
}
|
|
if (multiplier.isZero())
|
|
return ZERO;
|
|
if (this.eq(MIN_VALUE))
|
|
return multiplier.isOdd() ? MIN_VALUE : ZERO;
|
|
if (multiplier.eq(MIN_VALUE))
|
|
return this.isOdd() ? MIN_VALUE : ZERO;
|
|
if (this.isNegative()) {
|
|
if (multiplier.isNegative())
|
|
return this.neg().mul(multiplier.neg());
|
|
else
|
|
return this.neg().mul(multiplier).neg();
|
|
} else if (multiplier.isNegative())
|
|
return this.mul(multiplier.neg()).neg();
|
|
if (this.lt(TWO_PWR_24) && multiplier.lt(TWO_PWR_24))
|
|
return fromNumber(this.toNumber() * multiplier.toNumber(), this.unsigned);
|
|
var a48 = this.high >>> 16;
|
|
var a32 = this.high & 65535;
|
|
var a16 = this.low >>> 16;
|
|
var a00 = this.low & 65535;
|
|
var b48 = multiplier.high >>> 16;
|
|
var b32 = multiplier.high & 65535;
|
|
var b16 = multiplier.low >>> 16;
|
|
var b00 = multiplier.low & 65535;
|
|
var c48 = 0, c32 = 0, c16 = 0, c00 = 0;
|
|
c00 += a00 * b00;
|
|
c16 += c00 >>> 16;
|
|
c00 &= 65535;
|
|
c16 += a16 * b00;
|
|
c32 += c16 >>> 16;
|
|
c16 &= 65535;
|
|
c16 += a00 * b16;
|
|
c32 += c16 >>> 16;
|
|
c16 &= 65535;
|
|
c32 += a32 * b00;
|
|
c48 += c32 >>> 16;
|
|
c32 &= 65535;
|
|
c32 += a16 * b16;
|
|
c48 += c32 >>> 16;
|
|
c32 &= 65535;
|
|
c32 += a00 * b32;
|
|
c48 += c32 >>> 16;
|
|
c32 &= 65535;
|
|
c48 += a48 * b00 + a32 * b16 + a16 * b32 + a00 * b48;
|
|
c48 &= 65535;
|
|
return fromBits(c16 << 16 | c00, c48 << 16 | c32, this.unsigned);
|
|
};
|
|
LongPrototype.mul = LongPrototype.multiply;
|
|
LongPrototype.divide = function divide(divisor) {
|
|
if (!isLong(divisor))
|
|
divisor = fromValue(divisor);
|
|
if (divisor.isZero())
|
|
throw Error("division by zero");
|
|
if (wasm) {
|
|
if (!this.unsigned && this.high === -2147483648 && divisor.low === -1 && divisor.high === -1) {
|
|
return this;
|
|
}
|
|
var low = (this.unsigned ? wasm.div_u : wasm.div_s)(
|
|
this.low,
|
|
this.high,
|
|
divisor.low,
|
|
divisor.high
|
|
);
|
|
return fromBits(low, wasm.get_high(), this.unsigned);
|
|
}
|
|
if (this.isZero())
|
|
return this.unsigned ? UZERO : ZERO;
|
|
var approx, rem, res;
|
|
if (!this.unsigned) {
|
|
if (this.eq(MIN_VALUE)) {
|
|
if (divisor.eq(ONE) || divisor.eq(NEG_ONE))
|
|
return MIN_VALUE;
|
|
else if (divisor.eq(MIN_VALUE))
|
|
return ONE;
|
|
else {
|
|
var halfThis = this.shr(1);
|
|
approx = halfThis.div(divisor).shl(1);
|
|
if (approx.eq(ZERO)) {
|
|
return divisor.isNegative() ? ONE : NEG_ONE;
|
|
} else {
|
|
rem = this.sub(divisor.mul(approx));
|
|
res = approx.add(rem.div(divisor));
|
|
return res;
|
|
}
|
|
}
|
|
} else if (divisor.eq(MIN_VALUE))
|
|
return this.unsigned ? UZERO : ZERO;
|
|
if (this.isNegative()) {
|
|
if (divisor.isNegative())
|
|
return this.neg().div(divisor.neg());
|
|
return this.neg().div(divisor).neg();
|
|
} else if (divisor.isNegative())
|
|
return this.div(divisor.neg()).neg();
|
|
res = ZERO;
|
|
} else {
|
|
if (!divisor.unsigned)
|
|
divisor = divisor.toUnsigned();
|
|
if (divisor.gt(this))
|
|
return UZERO;
|
|
if (divisor.gt(this.shru(1)))
|
|
return UONE;
|
|
res = UZERO;
|
|
}
|
|
rem = this;
|
|
while (rem.gte(divisor)) {
|
|
approx = Math.max(1, Math.floor(rem.toNumber() / divisor.toNumber()));
|
|
var log2 = Math.ceil(Math.log(approx) / Math.LN2), delta = log2 <= 48 ? 1 : pow_dbl(2, log2 - 48), approxRes = fromNumber(approx), approxRem = approxRes.mul(divisor);
|
|
while (approxRem.isNegative() || approxRem.gt(rem)) {
|
|
approx -= delta;
|
|
approxRes = fromNumber(approx, this.unsigned);
|
|
approxRem = approxRes.mul(divisor);
|
|
}
|
|
if (approxRes.isZero())
|
|
approxRes = ONE;
|
|
res = res.add(approxRes);
|
|
rem = rem.sub(approxRem);
|
|
}
|
|
return res;
|
|
};
|
|
LongPrototype.div = LongPrototype.divide;
|
|
LongPrototype.modulo = function modulo(divisor) {
|
|
if (!isLong(divisor))
|
|
divisor = fromValue(divisor);
|
|
if (wasm) {
|
|
var low = (this.unsigned ? wasm.rem_u : wasm.rem_s)(
|
|
this.low,
|
|
this.high,
|
|
divisor.low,
|
|
divisor.high
|
|
);
|
|
return fromBits(low, wasm.get_high(), this.unsigned);
|
|
}
|
|
return this.sub(this.div(divisor).mul(divisor));
|
|
};
|
|
LongPrototype.mod = LongPrototype.modulo;
|
|
LongPrototype.rem = LongPrototype.modulo;
|
|
LongPrototype.not = function not() {
|
|
return fromBits(~this.low, ~this.high, this.unsigned);
|
|
};
|
|
LongPrototype.and = function and(other) {
|
|
if (!isLong(other))
|
|
other = fromValue(other);
|
|
return fromBits(this.low & other.low, this.high & other.high, this.unsigned);
|
|
};
|
|
LongPrototype.or = function or(other) {
|
|
if (!isLong(other))
|
|
other = fromValue(other);
|
|
return fromBits(this.low | other.low, this.high | other.high, this.unsigned);
|
|
};
|
|
LongPrototype.xor = function xor(other) {
|
|
if (!isLong(other))
|
|
other = fromValue(other);
|
|
return fromBits(this.low ^ other.low, this.high ^ other.high, this.unsigned);
|
|
};
|
|
LongPrototype.shiftLeft = function shiftLeft(numBits) {
|
|
if (isLong(numBits))
|
|
numBits = numBits.toInt();
|
|
if ((numBits &= 63) === 0)
|
|
return this;
|
|
else if (numBits < 32)
|
|
return fromBits(this.low << numBits, this.high << numBits | this.low >>> 32 - numBits, this.unsigned);
|
|
else
|
|
return fromBits(0, this.low << numBits - 32, this.unsigned);
|
|
};
|
|
LongPrototype.shl = LongPrototype.shiftLeft;
|
|
LongPrototype.shiftRight = function shiftRight(numBits) {
|
|
if (isLong(numBits))
|
|
numBits = numBits.toInt();
|
|
if ((numBits &= 63) === 0)
|
|
return this;
|
|
else if (numBits < 32)
|
|
return fromBits(this.low >>> numBits | this.high << 32 - numBits, this.high >> numBits, this.unsigned);
|
|
else
|
|
return fromBits(this.high >> numBits - 32, this.high >= 0 ? 0 : -1, this.unsigned);
|
|
};
|
|
LongPrototype.shr = LongPrototype.shiftRight;
|
|
LongPrototype.shiftRightUnsigned = function shiftRightUnsigned(numBits) {
|
|
if (isLong(numBits))
|
|
numBits = numBits.toInt();
|
|
numBits &= 63;
|
|
if (numBits === 0)
|
|
return this;
|
|
else {
|
|
var high = this.high;
|
|
if (numBits < 32) {
|
|
var low = this.low;
|
|
return fromBits(low >>> numBits | high << 32 - numBits, high >>> numBits, this.unsigned);
|
|
} else if (numBits === 32)
|
|
return fromBits(high, 0, this.unsigned);
|
|
else
|
|
return fromBits(high >>> numBits - 32, 0, this.unsigned);
|
|
}
|
|
};
|
|
LongPrototype.shru = LongPrototype.shiftRightUnsigned;
|
|
LongPrototype.shr_u = LongPrototype.shiftRightUnsigned;
|
|
LongPrototype.toSigned = function toSigned() {
|
|
if (!this.unsigned)
|
|
return this;
|
|
return fromBits(this.low, this.high, false);
|
|
};
|
|
LongPrototype.toUnsigned = function toUnsigned() {
|
|
if (this.unsigned)
|
|
return this;
|
|
return fromBits(this.low, this.high, true);
|
|
};
|
|
LongPrototype.toBytes = function toBytes(le2) {
|
|
return le2 ? this.toBytesLE() : this.toBytesBE();
|
|
};
|
|
LongPrototype.toBytesLE = function toBytesLE() {
|
|
var hi = this.high, lo = this.low;
|
|
return [
|
|
lo & 255,
|
|
lo >>> 8 & 255,
|
|
lo >>> 16 & 255,
|
|
lo >>> 24,
|
|
hi & 255,
|
|
hi >>> 8 & 255,
|
|
hi >>> 16 & 255,
|
|
hi >>> 24
|
|
];
|
|
};
|
|
LongPrototype.toBytesBE = function toBytesBE() {
|
|
var hi = this.high, lo = this.low;
|
|
return [
|
|
hi >>> 24,
|
|
hi >>> 16 & 255,
|
|
hi >>> 8 & 255,
|
|
hi & 255,
|
|
lo >>> 24,
|
|
lo >>> 16 & 255,
|
|
lo >>> 8 & 255,
|
|
lo & 255
|
|
];
|
|
};
|
|
Long2.fromBytes = function fromBytes(bytes, unsigned, le2) {
|
|
return le2 ? Long2.fromBytesLE(bytes, unsigned) : Long2.fromBytesBE(bytes, unsigned);
|
|
};
|
|
Long2.fromBytesLE = function fromBytesLE(bytes, unsigned) {
|
|
return new Long2(
|
|
bytes[0] | bytes[1] << 8 | bytes[2] << 16 | bytes[3] << 24,
|
|
bytes[4] | bytes[5] << 8 | bytes[6] << 16 | bytes[7] << 24,
|
|
unsigned
|
|
);
|
|
};
|
|
Long2.fromBytesBE = function fromBytesBE(bytes, unsigned) {
|
|
return new Long2(
|
|
bytes[4] << 24 | bytes[5] << 16 | bytes[6] << 8 | bytes[7],
|
|
bytes[0] << 24 | bytes[1] << 16 | bytes[2] << 8 | bytes[3],
|
|
unsigned
|
|
);
|
|
};
|
|
return long$1;
|
|
}
|
|
var longExports = requireLong();
|
|
const long = /* @__PURE__ */ getDefaultExportFromCjs(longExports);
|
|
const LongExports = /* @__PURE__ */ _mergeNamespaces({
|
|
__proto__: null,
|
|
default: long
|
|
}, [longExports]);
|
|
const Long = (
|
|
// tslint:disable-next-line
|
|
long || LongExports
|
|
);
|
|
function hexToLong(hex) {
|
|
return Long.fromString(hex, true, 16);
|
|
}
|
|
const k0 = hexToLong("c3a5c85c97cb3127");
|
|
const k1 = hexToLong("b492b66fbe98f273");
|
|
const k2 = hexToLong("9ae16a3b2f90404f");
|
|
function shiftMix(val) {
|
|
return val.xor(val.shru(47));
|
|
}
|
|
function fetch$1(s, offset, numBytes) {
|
|
const bytes = s.slice(offset, offset + numBytes);
|
|
return Long.fromBytes(Array.from(bytes), true, true);
|
|
}
|
|
function fetch64(s, offset) {
|
|
return fetch$1(s, offset, 8);
|
|
}
|
|
function fetch32(s, offset) {
|
|
return fetch$1(s, offset, 4);
|
|
}
|
|
function rotate64(val, shift) {
|
|
return shift === 0 ? val : val.shru(shift).or(val.shl(64 - shift));
|
|
}
|
|
function hashLen16(u, v, mul2 = hexToLong("9ddfea08eb382d69")) {
|
|
let a = u.xor(v).mul(mul2);
|
|
a = a.xor(a.shru(47));
|
|
let b = v.xor(a).mul(mul2);
|
|
b = b.xor(b.shru(47));
|
|
b = b.mul(mul2);
|
|
return b;
|
|
}
|
|
function weakHashLen32WithSeeds(w, x, y, z2, a, b) {
|
|
a = a.add(w);
|
|
b = rotate64(b.add(a).add(z2), 21);
|
|
const c = a;
|
|
a = a.add(x);
|
|
a = a.add(y);
|
|
b = b.add(rotate64(a, 44));
|
|
return [a.add(z2), b.add(c)];
|
|
}
|
|
function weakHashLen32WithSeedsStr(s, offset, a, b) {
|
|
return weakHashLen32WithSeeds(fetch64(s, offset), fetch64(s, offset + 8), fetch64(s, offset + 16), fetch64(s, offset + 24), a, b);
|
|
}
|
|
function hashLen0to16(s, len = s.length) {
|
|
if (len >= 8) {
|
|
const mul2 = k2.add(len * 2);
|
|
const a = fetch64(s, 0).add(k2);
|
|
const b = fetch64(s, len - 8);
|
|
const c = rotate64(b, 37).mul(mul2).add(a);
|
|
const d = rotate64(a, 25).add(b).mul(mul2);
|
|
return hashLen16(c, d, mul2);
|
|
}
|
|
if (len >= 4) {
|
|
const mul2 = k2.add(len * 2);
|
|
const a = fetch32(s, 0);
|
|
return hashLen16(a.shl(3).add(len), fetch32(s, len - 4), mul2);
|
|
}
|
|
if (len > 0) {
|
|
const a = s[0];
|
|
const b = s[len >> 1];
|
|
const c = s[len - 1];
|
|
const y = a + (b << 8);
|
|
const z2 = len + (c << 2);
|
|
return shiftMix(k2.mul(y).xor(k0.mul(z2))).mul(k2);
|
|
}
|
|
return k2;
|
|
}
|
|
function hashLen17to32(s, len = s.length) {
|
|
const mul2 = k2.add(len * 2);
|
|
const a = fetch64(s, 0).mul(k1);
|
|
const b = fetch64(s, 8);
|
|
const c = fetch64(s, len - 8).mul(mul2);
|
|
const d = fetch64(s, len - 16).mul(k2);
|
|
return hashLen16(rotate64(a.add(b), 43).add(rotate64(c, 30)).add(d), a.add(rotate64(b.add(k2), 18)).add(c), mul2);
|
|
}
|
|
function hashLen33to64(s, len = s.length) {
|
|
const mul2 = k2.add(len * 2);
|
|
const a = fetch64(s, 0).mul(k2);
|
|
const b = fetch64(s, 8);
|
|
const c = fetch64(s, len - 8).mul(mul2);
|
|
const d = fetch64(s, len - 16).mul(k2);
|
|
const y = rotate64(a.add(b), 43).add(rotate64(c, 30)).add(d);
|
|
const z2 = hashLen16(y, a.add(rotate64(b.add(k2), 18)).add(c), mul2);
|
|
const e = fetch64(s, 16).mul(mul2);
|
|
const f = fetch64(s, 24);
|
|
const g = y.add(fetch64(s, len - 32)).mul(mul2);
|
|
const h = z2.add(fetch64(s, len - 24)).mul(mul2);
|
|
return hashLen16(rotate64(e.add(f), 43).add(rotate64(g, 30)).add(h), e.add(rotate64(f.add(a), 18)).add(g), mul2);
|
|
}
|
|
function fingerPrint64(s, len = s.length) {
|
|
const seed = Long.fromNumber(81, true);
|
|
if (len <= 32) {
|
|
if (len <= 16) {
|
|
return hashLen0to16(s, len);
|
|
} else {
|
|
return hashLen17to32(s, len);
|
|
}
|
|
} else if (len <= 64) {
|
|
return hashLen33to64(s, len);
|
|
}
|
|
let x = seed;
|
|
let y = seed.mul(k1).add(113);
|
|
let z2 = shiftMix(y.mul(k2).add(113)).mul(k2);
|
|
let v = [Long.UZERO, Long.UZERO];
|
|
let w = [Long.UZERO, Long.UZERO];
|
|
x = x.mul(k2).add(fetch64(s, 0));
|
|
let offset = 0;
|
|
const end = (len - 1 >> 6) * 64;
|
|
const last64 = end + (len - 1 & 63) - 63;
|
|
do {
|
|
x = rotate64(x.add(y).add(v[0]).add(fetch64(s, offset + 8)), 37).mul(k1);
|
|
y = rotate64(y.add(v[1]).add(fetch64(s, offset + 48)), 42).mul(k1);
|
|
x = x.xor(w[1]);
|
|
y = y.add(v[0]).add(fetch64(s, offset + 40));
|
|
z2 = rotate64(z2.add(w[0]), 33).mul(k1);
|
|
v = weakHashLen32WithSeedsStr(s, offset, v[1].mul(k1), x.add(w[0]));
|
|
w = weakHashLen32WithSeedsStr(s, offset + 32, z2.add(w[1]), y.add(fetch64(s, offset + 16)));
|
|
[z2, x] = [x, z2];
|
|
offset += 64;
|
|
} while (offset !== end);
|
|
const mul2 = k1.add(z2.and(255).shl(1));
|
|
offset = last64;
|
|
w[0] = w[0].add(len - 1 & 63);
|
|
v[0] = v[0].add(w[0]);
|
|
w[0] = w[0].add(v[0]);
|
|
x = rotate64(x.add(y).add(v[0]).add(fetch64(s, offset + 8)), 37).mul(mul2);
|
|
y = rotate64(y.add(v[1]).add(fetch64(s, offset + 48)), 42).mul(mul2);
|
|
x = x.xor(w[1].mul(9));
|
|
y = y.add(v[0].mul(9).add(fetch64(s, offset + 40)));
|
|
z2 = rotate64(z2.add(w[0]), 33).mul(mul2);
|
|
v = weakHashLen32WithSeedsStr(s, offset, v[1].mul(mul2), x.add(w[0]));
|
|
w = weakHashLen32WithSeedsStr(s, offset + 32, z2.add(w[1]), y.add(fetch64(s, offset + 16)));
|
|
[z2, x] = [x, z2];
|
|
return hashLen16(hashLen16(v[0], w[0], mul2).add(shiftMix(y).mul(k0)).add(z2), hashLen16(v[1], w[1], mul2).add(x), mul2);
|
|
}
|
|
function createScalarValue(value, dtype) {
|
|
if (dtype === "string") {
|
|
return encodeString(value);
|
|
}
|
|
return toTypedArray([value], dtype);
|
|
}
|
|
function noConversionNeeded(a, dtype) {
|
|
return a instanceof Float32Array && dtype === "float32" || a instanceof Int32Array && dtype === "int32" || a instanceof Uint8Array && dtype === "bool";
|
|
}
|
|
function toTypedArray(a, dtype) {
|
|
if (dtype === "string") {
|
|
throw new Error("Cannot convert a string[] to a TypedArray");
|
|
}
|
|
if (Array.isArray(a)) {
|
|
a = flatten(a);
|
|
}
|
|
if (env().getBool("DEBUG")) {
|
|
checkConversionForErrors(a, dtype);
|
|
}
|
|
if (noConversionNeeded(a, dtype)) {
|
|
return a;
|
|
}
|
|
if (dtype == null || dtype === "float32" || dtype === "complex64") {
|
|
return new Float32Array(a);
|
|
} else if (dtype === "int32") {
|
|
return new Int32Array(a);
|
|
} else if (dtype === "bool") {
|
|
const bool = new Uint8Array(a.length);
|
|
for (let i = 0; i < bool.length; ++i) {
|
|
if (Math.round(a[i]) !== 0) {
|
|
bool[i] = 1;
|
|
}
|
|
}
|
|
return bool;
|
|
} else {
|
|
throw new Error(`Unknown data type ${dtype}`);
|
|
}
|
|
}
|
|
function now() {
|
|
return env().platform.now();
|
|
}
|
|
function encodeString(s, encoding = "utf-8") {
|
|
encoding = encoding || "utf-8";
|
|
return env().platform.encode(s, encoding);
|
|
}
|
|
function decodeString(bytes, encoding = "utf-8") {
|
|
encoding = encoding || "utf-8";
|
|
return env().platform.decode(bytes, encoding);
|
|
}
|
|
function isTypedArray(a) {
|
|
if (env().platform.isTypedArray != null) {
|
|
return env().platform.isTypedArray(a);
|
|
} else {
|
|
return isTypedArrayBrowser(a);
|
|
}
|
|
}
|
|
function flatten(arr, result = [], skipTypedArray = false) {
|
|
if (result == null) {
|
|
result = [];
|
|
}
|
|
if (typeof arr === "boolean" || typeof arr === "number" || typeof arr === "string" || isPromise(arr) || arr == null || isTypedArray(arr) && skipTypedArray) {
|
|
result.push(arr);
|
|
} else if (Array.isArray(arr) || isTypedArray(arr)) {
|
|
for (let i = 0; i < arr.length; ++i) {
|
|
flatten(arr[i], result, skipTypedArray);
|
|
}
|
|
} else {
|
|
let maxIndex = -1;
|
|
for (const key of Object.keys(arr)) {
|
|
if (/^([1-9]+[0-9]*|0)$/.test(key)) {
|
|
maxIndex = Math.max(maxIndex, Number(key));
|
|
}
|
|
}
|
|
for (let i = 0; i <= maxIndex; i++) {
|
|
flatten(arr[i], result, skipTypedArray);
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
class Profiler {
|
|
constructor(backendTimer, logger) {
|
|
this.backendTimer = backendTimer;
|
|
this.logger = logger;
|
|
if (logger == null) {
|
|
this.logger = new Logger();
|
|
}
|
|
}
|
|
profileKernel(kernelName, inputs, f) {
|
|
let outputs;
|
|
const holdResultWrapperFn = () => {
|
|
outputs = f();
|
|
};
|
|
let timer;
|
|
const start = now();
|
|
if (this.backendTimer.timerAvailable()) {
|
|
timer = this.backendTimer.time(holdResultWrapperFn);
|
|
} else {
|
|
holdResultWrapperFn();
|
|
for (const output of outputs) {
|
|
output.dataSync();
|
|
}
|
|
timer = Promise.resolve({ kernelMs: now() - start });
|
|
}
|
|
if (env().getBool("CHECK_COMPUTATION_FOR_ERRORS")) {
|
|
for (let i = 0; i < outputs.length; i++) {
|
|
const output = outputs[i];
|
|
output.data().then((tensorVals) => {
|
|
checkComputationForErrors(tensorVals, output.dtype, kernelName);
|
|
});
|
|
}
|
|
}
|
|
const kernelProfile = {
|
|
kernelName,
|
|
outputs,
|
|
inputs,
|
|
timeMs: timer.then((timing) => timing.kernelMs),
|
|
extraInfo: timer.then((timing) => timing.getExtraProfileInfo != null ? timing.getExtraProfileInfo() : "")
|
|
};
|
|
return kernelProfile;
|
|
}
|
|
logKernelProfile(kernelProfile) {
|
|
const { kernelName, outputs, timeMs, inputs, extraInfo } = kernelProfile;
|
|
outputs.forEach((result) => {
|
|
Promise.all([result.data(), timeMs, extraInfo]).then((valueContainer) => {
|
|
this.logger.logKernelProfile(kernelName, result, valueContainer[0], valueContainer[1], inputs, valueContainer[2]);
|
|
});
|
|
});
|
|
}
|
|
}
|
|
function checkComputationForErrors(vals, dtype, kernelName) {
|
|
if (dtype !== "float32") {
|
|
return false;
|
|
}
|
|
for (let i = 0; i < vals.length; i++) {
|
|
const num = vals[i];
|
|
if (isNaN(num) || !isFinite(num)) {
|
|
console.warn(`Found ${num} in the result of '${kernelName}'`);
|
|
return true;
|
|
}
|
|
}
|
|
return false;
|
|
}
|
|
class Logger {
|
|
logKernelProfile(name, result, vals, timeMs, inputs, extraInfo) {
|
|
const time = typeof timeMs === "number" ? rightPad(`${timeMs}ms`, 9) : timeMs["error"];
|
|
const paddedName = rightPad(name, 25);
|
|
const rank = result.rank;
|
|
const size = result.size;
|
|
const shape = rightPad(result.shape.toString(), 14);
|
|
let inputShapesDescription = "";
|
|
for (const name2 in inputs) {
|
|
const input = inputs[name2];
|
|
if (input != null) {
|
|
const inputShape = input.shape || result.shape;
|
|
const inputRank = inputShape.length;
|
|
inputShapesDescription += `${name2}: ${inputRank}D ${inputRank > 0 ? inputShape : ""} `;
|
|
}
|
|
}
|
|
console.log(`%c${paddedName} %c${time} %c${rank}D ${shape} %c${size} %c${inputShapesDescription} %c${extraInfo}`, "font-weight:bold", "color:red", "color:blue", "color: orange", "color: green", "color: steelblue");
|
|
}
|
|
}
|
|
function getFilteredNodesXToY(tape, xs, y) {
|
|
const tensorsFromX = {};
|
|
const nodesFromX = {};
|
|
for (let i = 0; i < xs.length; i++) {
|
|
tensorsFromX[xs[i].id] = true;
|
|
}
|
|
for (let i = 0; i < tape.length; i++) {
|
|
const node = tape[i];
|
|
const nodeInputs = node.inputs;
|
|
for (const inputName in nodeInputs) {
|
|
const input = nodeInputs[inputName];
|
|
let anyInputFromX = false;
|
|
for (let j2 = 0; j2 < xs.length; j2++) {
|
|
if (tensorsFromX[input.id]) {
|
|
node.outputs.forEach((output) => tensorsFromX[output.id] = true);
|
|
anyInputFromX = true;
|
|
nodesFromX[node.id] = true;
|
|
break;
|
|
}
|
|
}
|
|
if (anyInputFromX) {
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
const tensorsLeadToY = {};
|
|
tensorsLeadToY[y.id] = true;
|
|
const nodesToY = {};
|
|
for (let i = tape.length - 1; i >= 0; i--) {
|
|
const node = tape[i];
|
|
const nodeInputs = node.inputs;
|
|
for (let j2 = 0; j2 < node.outputs.length; j2++) {
|
|
if (tensorsLeadToY[node.outputs[j2].id]) {
|
|
for (const inputName in nodeInputs) {
|
|
tensorsLeadToY[nodeInputs[inputName].id] = true;
|
|
nodesToY[node.id] = true;
|
|
}
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
const filteredTape = [];
|
|
for (let i = 0; i < tape.length; i++) {
|
|
const node = tape[i];
|
|
if (nodesFromX[node.id] && nodesToY[node.id]) {
|
|
const prunedInputs = {};
|
|
for (const inputName in node.inputs) {
|
|
const nodeInput = node.inputs[inputName];
|
|
if (tensorsFromX[nodeInput.id]) {
|
|
prunedInputs[inputName] = nodeInput;
|
|
}
|
|
}
|
|
const prunedNode = Object.assign({}, node);
|
|
prunedNode.inputs = prunedInputs;
|
|
prunedNode.outputs = node.outputs;
|
|
filteredTape.push(prunedNode);
|
|
}
|
|
}
|
|
return filteredTape;
|
|
}
|
|
function backpropagateGradients(tensorAccumulatedGradientMap, filteredTape, tidy2, add2) {
|
|
for (let i = filteredTape.length - 1; i >= 0; i--) {
|
|
const node = filteredTape[i];
|
|
const dys = [];
|
|
node.outputs.forEach((o) => {
|
|
const gradTensor = tensorAccumulatedGradientMap[o.id];
|
|
if (gradTensor != null) {
|
|
dys.push(gradTensor);
|
|
} else {
|
|
dys.push(null);
|
|
}
|
|
});
|
|
if (node.gradient == null) {
|
|
throw new Error(`Cannot compute gradient: gradient function not found for ${node.kernelName}.`);
|
|
}
|
|
const inputGradients = node.gradient(dys);
|
|
for (const inputName in node.inputs) {
|
|
if (!(inputName in inputGradients)) {
|
|
throw new Error(`Cannot backprop through input ${inputName}. Available gradients found: ${Object.keys(inputGradients)}.`);
|
|
}
|
|
const dx = tidy2(() => inputGradients[inputName]());
|
|
if (dx.dtype !== "float32") {
|
|
throw new Error(`Error in gradient for op ${node.kernelName}. The gradient of input ${inputName} must have 'float32' dtype, but has '${dx.dtype}'`);
|
|
}
|
|
const x = node.inputs[inputName];
|
|
if (!arraysEqual(dx.shape, x.shape)) {
|
|
throw new Error(`Error in gradient for op ${node.kernelName}. The gradient of input '${inputName}' has shape '${dx.shape}', which does not match the shape of the input '${x.shape}'`);
|
|
}
|
|
if (tensorAccumulatedGradientMap[x.id] == null) {
|
|
tensorAccumulatedGradientMap[x.id] = dx;
|
|
} else {
|
|
const curGradient = tensorAccumulatedGradientMap[x.id];
|
|
tensorAccumulatedGradientMap[x.id] = add2(curGradient, dx);
|
|
curGradient.dispose();
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const FORMAT_LIMIT_NUM_VALS = 20;
|
|
const FORMAT_NUM_FIRST_LAST_VALS = 3;
|
|
const FORMAT_NUM_SIG_DIGITS = 7;
|
|
function tensorToString(vals, shape, dtype, verbose) {
|
|
const strides = computeStrides(shape);
|
|
const padPerCol = computeMaxSizePerColumn(vals, shape, dtype, strides);
|
|
const rank = shape.length;
|
|
const valsLines = subTensorToString(vals, shape, dtype, strides, padPerCol);
|
|
const lines = ["Tensor"];
|
|
if (verbose) {
|
|
lines.push(` dtype: ${dtype}`);
|
|
lines.push(` rank: ${rank}`);
|
|
lines.push(` shape: [${shape}]`);
|
|
lines.push(` values:`);
|
|
}
|
|
lines.push(valsLines.map((l) => " " + l).join("\n"));
|
|
return lines.join("\n");
|
|
}
|
|
function computeMaxSizePerColumn(vals, shape, dtype, strides) {
|
|
const n = sizeFromShape(shape);
|
|
const numCols = strides[strides.length - 1];
|
|
const padPerCol = new Array(numCols).fill(0);
|
|
const rank = shape.length;
|
|
const valuesOrTuples = dtype === "complex64" ? createComplexTuples(vals) : vals;
|
|
if (rank > 1) {
|
|
for (let row = 0; row < n / numCols; row++) {
|
|
const offset = row * numCols;
|
|
for (let j2 = 0; j2 < numCols; j2++) {
|
|
padPerCol[j2] = Math.max(padPerCol[j2], valToString(valuesOrTuples[offset + j2], 0, dtype).length);
|
|
}
|
|
}
|
|
}
|
|
return padPerCol;
|
|
}
|
|
function valToString(val, pad2, dtype) {
|
|
let valStr;
|
|
if (Array.isArray(val)) {
|
|
valStr = `${parseFloat(val[0].toFixed(FORMAT_NUM_SIG_DIGITS))} + ${parseFloat(val[1].toFixed(FORMAT_NUM_SIG_DIGITS))}j`;
|
|
} else if (isString(val)) {
|
|
valStr = `'${val}'`;
|
|
} else if (dtype === "bool") {
|
|
valStr = boolNumToString(val);
|
|
} else {
|
|
valStr = parseFloat(val.toFixed(FORMAT_NUM_SIG_DIGITS)).toString();
|
|
}
|
|
return rightPad(valStr, pad2);
|
|
}
|
|
function boolNumToString(v) {
|
|
return v === 0 ? "false" : "true";
|
|
}
|
|
function subTensorToString(vals, shape, dtype, strides, padPerCol, isLast = true) {
|
|
const storagePerElement = dtype === "complex64" ? 2 : 1;
|
|
const size = shape[0];
|
|
const rank = shape.length;
|
|
if (rank === 0) {
|
|
if (dtype === "complex64") {
|
|
const complexTuple = createComplexTuples(vals);
|
|
return [valToString(complexTuple[0], 0, dtype)];
|
|
}
|
|
if (dtype === "bool") {
|
|
return [boolNumToString(vals[0])];
|
|
}
|
|
return [vals[0].toString()];
|
|
}
|
|
if (rank === 1) {
|
|
if (size > FORMAT_LIMIT_NUM_VALS) {
|
|
const firstValsSize = FORMAT_NUM_FIRST_LAST_VALS * storagePerElement;
|
|
let firstVals = Array.from(vals.slice(0, firstValsSize));
|
|
let lastVals = Array.from(vals.slice((size - FORMAT_NUM_FIRST_LAST_VALS) * storagePerElement, size * storagePerElement));
|
|
if (dtype === "complex64") {
|
|
firstVals = createComplexTuples(firstVals);
|
|
lastVals = createComplexTuples(lastVals);
|
|
}
|
|
return [
|
|
"[" + firstVals.map((x, i) => valToString(x, padPerCol[i], dtype)).join(", ") + ", ..., " + lastVals.map((x, i) => valToString(x, padPerCol[size - FORMAT_NUM_FIRST_LAST_VALS + i], dtype)).join(", ") + "]"
|
|
];
|
|
}
|
|
const displayVals = dtype === "complex64" ? createComplexTuples(vals) : Array.from(vals);
|
|
return [
|
|
"[" + displayVals.map((x, i) => valToString(x, padPerCol[i], dtype)).join(", ") + "]"
|
|
];
|
|
}
|
|
const subshape = shape.slice(1);
|
|
const substrides = strides.slice(1);
|
|
const stride = strides[0] * storagePerElement;
|
|
const lines = [];
|
|
if (size > FORMAT_LIMIT_NUM_VALS) {
|
|
for (let i = 0; i < FORMAT_NUM_FIRST_LAST_VALS; i++) {
|
|
const start = i * stride;
|
|
const end = start + stride;
|
|
lines.push(...subTensorToString(
|
|
vals.slice(start, end),
|
|
subshape,
|
|
dtype,
|
|
substrides,
|
|
padPerCol,
|
|
false
|
|
/* isLast */
|
|
));
|
|
}
|
|
lines.push("...");
|
|
for (let i = size - FORMAT_NUM_FIRST_LAST_VALS; i < size; i++) {
|
|
const start = i * stride;
|
|
const end = start + stride;
|
|
lines.push(...subTensorToString(
|
|
vals.slice(start, end),
|
|
subshape,
|
|
dtype,
|
|
substrides,
|
|
padPerCol,
|
|
i === size - 1
|
|
/* isLast */
|
|
));
|
|
}
|
|
} else {
|
|
for (let i = 0; i < size; i++) {
|
|
const start = i * stride;
|
|
const end = start + stride;
|
|
lines.push(...subTensorToString(
|
|
vals.slice(start, end),
|
|
subshape,
|
|
dtype,
|
|
substrides,
|
|
padPerCol,
|
|
i === size - 1
|
|
/* isLast */
|
|
));
|
|
}
|
|
}
|
|
const sep = rank === 2 ? "," : "";
|
|
lines[0] = "[" + (size > 0 ? lines[0] + sep : "");
|
|
for (let i = 1; i < lines.length - 1; i++) {
|
|
lines[i] = " " + lines[i] + sep;
|
|
}
|
|
let newLineSep = ",\n";
|
|
for (let i = 2; i < rank; i++) {
|
|
newLineSep += "\n";
|
|
}
|
|
lines[lines.length - 1] = " " + lines[lines.length - 1] + "]" + (isLast ? "" : newLineSep);
|
|
return lines;
|
|
}
|
|
function createComplexTuples(vals) {
|
|
const complexTuples = [];
|
|
for (let i = 0; i < vals.length; i += 2) {
|
|
complexTuples.push([vals[i], vals[i + 1]]);
|
|
}
|
|
return complexTuples;
|
|
}
|
|
class TensorBuffer {
|
|
constructor(shape, dtype, values) {
|
|
this.dtype = dtype;
|
|
this.shape = shape.slice();
|
|
this.size = sizeFromShape(shape);
|
|
if (values != null) {
|
|
const n = values.length;
|
|
assert(n === this.size, () => `Length of values '${n}' does not match the size inferred by the shape '${this.size}'.`);
|
|
}
|
|
if (dtype === "complex64") {
|
|
throw new Error(`complex64 dtype TensorBuffers are not supported. Please create a TensorBuffer for the real and imaginary parts separately and call tf.complex(real, imag).`);
|
|
}
|
|
this.values = values || getArrayFromDType(dtype, this.size);
|
|
this.strides = computeStrides(shape);
|
|
}
|
|
/**
|
|
* Sets a value in the buffer at a given location.
|
|
*
|
|
* @param value The value to set.
|
|
* @param locs The location indices.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Creation'}
|
|
*/
|
|
set(value, ...locs) {
|
|
if (locs.length === 0) {
|
|
locs = [0];
|
|
}
|
|
assert(locs.length === this.rank, () => `The number of provided coordinates (${locs.length}) must match the rank (${this.rank})`);
|
|
const index = this.locToIndex(locs);
|
|
this.values[index] = value;
|
|
}
|
|
/**
|
|
* Returns the value in the buffer at the provided location.
|
|
*
|
|
* @param locs The location indices.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Creation'}
|
|
*/
|
|
get(...locs) {
|
|
if (locs.length === 0) {
|
|
locs = [0];
|
|
}
|
|
let i = 0;
|
|
for (const loc of locs) {
|
|
if (loc < 0 || loc >= this.shape[i]) {
|
|
const msg = `Requested out of range element at ${locs}. Buffer shape=${this.shape}`;
|
|
throw new Error(msg);
|
|
}
|
|
i++;
|
|
}
|
|
let index = locs[locs.length - 1];
|
|
for (let i2 = 0; i2 < locs.length - 1; ++i2) {
|
|
index += this.strides[i2] * locs[i2];
|
|
}
|
|
return this.values[index];
|
|
}
|
|
locToIndex(locs) {
|
|
if (this.rank === 0) {
|
|
return 0;
|
|
} else if (this.rank === 1) {
|
|
return locs[0];
|
|
}
|
|
let index = locs[locs.length - 1];
|
|
for (let i = 0; i < locs.length - 1; ++i) {
|
|
index += this.strides[i] * locs[i];
|
|
}
|
|
return index;
|
|
}
|
|
indexToLoc(index) {
|
|
if (this.rank === 0) {
|
|
return [];
|
|
} else if (this.rank === 1) {
|
|
return [index];
|
|
}
|
|
const locs = new Array(this.shape.length);
|
|
for (let i = 0; i < locs.length - 1; ++i) {
|
|
locs[i] = Math.floor(index / this.strides[i]);
|
|
index -= locs[i] * this.strides[i];
|
|
}
|
|
locs[locs.length - 1] = index;
|
|
return locs;
|
|
}
|
|
get rank() {
|
|
return this.shape.length;
|
|
}
|
|
/**
|
|
* Creates an immutable `tf.Tensor` object from the buffer.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Creation'}
|
|
*/
|
|
toTensor() {
|
|
return trackerFn().makeTensor(this.values, this.shape, this.dtype);
|
|
}
|
|
}
|
|
let trackerFn = null;
|
|
let opHandler$1 = null;
|
|
function setTensorTracker(fn) {
|
|
trackerFn = fn;
|
|
}
|
|
function setOpHandler(handler) {
|
|
opHandler$1 = handler;
|
|
}
|
|
class Tensor {
|
|
constructor(shape, dtype, dataId, id) {
|
|
this.kept = false;
|
|
this.isDisposedInternal = false;
|
|
this.shape = shape.slice();
|
|
this.dtype = dtype || "float32";
|
|
this.size = sizeFromShape(shape);
|
|
this.strides = computeStrides(shape);
|
|
this.dataId = dataId;
|
|
this.id = id;
|
|
this.rankType = this.rank < 5 ? this.rank.toString() : "higher";
|
|
}
|
|
get rank() {
|
|
return this.shape.length;
|
|
}
|
|
/**
|
|
* Returns a promise of `tf.TensorBuffer` that holds the underlying data.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
async buffer() {
|
|
const vals = await this.data();
|
|
return opHandler$1.buffer(this.shape, this.dtype, vals);
|
|
}
|
|
/**
|
|
* Returns a `tf.TensorBuffer` that holds the underlying data.
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
bufferSync() {
|
|
return opHandler$1.buffer(this.shape, this.dtype, this.dataSync());
|
|
}
|
|
/**
|
|
* Returns the tensor data as a nested array. The transfer of data is done
|
|
* asynchronously.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
async array() {
|
|
const vals = await this.data();
|
|
return toNestedArray(this.shape, vals, this.dtype === "complex64");
|
|
}
|
|
/**
|
|
* Returns the tensor data as a nested array. The transfer of data is done
|
|
* synchronously.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
arraySync() {
|
|
return toNestedArray(this.shape, this.dataSync(), this.dtype === "complex64");
|
|
}
|
|
/**
|
|
* Asynchronously downloads the values from the `tf.Tensor`. Returns a
|
|
* promise of `TypedArray` that resolves when the computation has finished.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
async data() {
|
|
this.throwIfDisposed();
|
|
const data = trackerFn().read(this.dataId);
|
|
if (this.dtype === "string") {
|
|
const bytes = await data;
|
|
try {
|
|
return bytes.map((b) => decodeString(b));
|
|
} catch (_a) {
|
|
throw new Error("Failed to decode the string bytes into utf-8. To get the original bytes, call tensor.bytes().");
|
|
}
|
|
}
|
|
return data;
|
|
}
|
|
/**
|
|
* Copy the tensor's data to a new GPU resource. Comparing to the `dataSync()`
|
|
* and `data()`, this method prevents data from being downloaded to CPU.
|
|
*
|
|
* For WebGL backend, the data will be stored on a densely packed texture.
|
|
* This means that the texture will use the RGBA channels to store value.
|
|
*
|
|
* For WebGPU backend, the data will be stored on a buffer. There is no
|
|
* parameter, so can not use a user-defined size to create the buffer.
|
|
*
|
|
* @param options:
|
|
* For WebGL,
|
|
* - customTexShape: Optional. If set, will use the user defined
|
|
* texture shape to create the texture.
|
|
*
|
|
* @returns For WebGL backend, a GPUData contains the new texture and
|
|
* its information.
|
|
* {
|
|
* tensorRef: The tensor that is associated with this texture,
|
|
* texture: WebGLTexture,
|
|
* texShape: [number, number] // [height, width]
|
|
* }
|
|
*
|
|
* For WebGPU backend, a GPUData contains the new buffer.
|
|
* {
|
|
* tensorRef: The tensor that is associated with this buffer,
|
|
* buffer: GPUBuffer,
|
|
* }
|
|
*
|
|
* Remember to dispose the GPUData after it is used by
|
|
* `res.tensorRef.dispose()`.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
dataToGPU(options) {
|
|
this.throwIfDisposed();
|
|
return trackerFn().readToGPU(this.dataId, options);
|
|
}
|
|
/**
|
|
* Synchronously downloads the values from the `tf.Tensor`. This blocks the
|
|
* UI thread until the values are ready, which can cause performance issues.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
dataSync() {
|
|
this.throwIfDisposed();
|
|
const data = trackerFn().readSync(this.dataId);
|
|
if (this.dtype === "string") {
|
|
try {
|
|
return data.map((b) => decodeString(b));
|
|
} catch (_a) {
|
|
throw new Error("Failed to decode the string bytes into utf-8. To get the original bytes, call tensor.bytes().");
|
|
}
|
|
}
|
|
return data;
|
|
}
|
|
/** Returns the underlying bytes of the tensor's data. */
|
|
async bytes() {
|
|
this.throwIfDisposed();
|
|
const data = await trackerFn().read(this.dataId);
|
|
if (this.dtype === "string") {
|
|
return data;
|
|
} else {
|
|
return new Uint8Array(data.buffer);
|
|
}
|
|
}
|
|
/**
|
|
* Disposes `tf.Tensor` from memory.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
dispose() {
|
|
if (this.isDisposed) {
|
|
return;
|
|
}
|
|
if (this.kerasMask) {
|
|
this.kerasMask.dispose();
|
|
}
|
|
trackerFn().disposeTensor(this);
|
|
this.isDisposedInternal = true;
|
|
}
|
|
get isDisposed() {
|
|
return this.isDisposedInternal;
|
|
}
|
|
throwIfDisposed() {
|
|
if (this.isDisposed) {
|
|
throw new Error(`Tensor is disposed.`);
|
|
}
|
|
}
|
|
/**
|
|
* Prints the `tf.Tensor`. See `tf.print` for details.
|
|
*
|
|
* @param verbose Whether to print verbose information about the tensor,
|
|
* including dtype and size.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
print(verbose = false) {
|
|
return opHandler$1.print(this, verbose);
|
|
}
|
|
/**
|
|
* Returns a copy of the tensor. See `tf.clone` for details.
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
clone() {
|
|
this.throwIfDisposed();
|
|
return opHandler$1.clone(this);
|
|
}
|
|
/**
|
|
* Returns a human-readable description of the tensor. Useful for logging.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
toString(verbose = false) {
|
|
const vals = this.dataSync();
|
|
return tensorToString(vals, this.shape, this.dtype, verbose);
|
|
}
|
|
cast(dtype) {
|
|
this.throwIfDisposed();
|
|
return opHandler$1.cast(this, dtype);
|
|
}
|
|
variable(trainable = true, name, dtype) {
|
|
this.throwIfDisposed();
|
|
return trackerFn().makeVariable(this, trainable, name, dtype);
|
|
}
|
|
}
|
|
Object.defineProperty(Tensor, Symbol.hasInstance, {
|
|
value: (instance) => {
|
|
return !!instance && instance.data != null && instance.dataSync != null && instance.throwIfDisposed != null;
|
|
}
|
|
});
|
|
function getGlobalTensorClass() {
|
|
return getGlobal("Tensor", () => {
|
|
return Tensor;
|
|
});
|
|
}
|
|
getGlobalTensorClass();
|
|
class Variable extends Tensor {
|
|
constructor(initialValue, trainable, name, tensorId) {
|
|
super(initialValue.shape, initialValue.dtype, initialValue.dataId, tensorId);
|
|
this.trainable = trainable;
|
|
this.name = name;
|
|
}
|
|
/**
|
|
* Assign a new `tf.Tensor` to this variable. The new `tf.Tensor` must have
|
|
* the same shape and dtype as the old `tf.Tensor`.
|
|
*
|
|
* @param newValue New tensor to be assigned to this variable.
|
|
*
|
|
* @doc {heading: 'Tensors', subheading: 'Classes'}
|
|
*/
|
|
assign(newValue) {
|
|
if (newValue.dtype !== this.dtype) {
|
|
throw new Error(`dtype of the new value (${newValue.dtype}) and previous value (${this.dtype}) must match`);
|
|
}
|
|
if (!arraysEqual(newValue.shape, this.shape)) {
|
|
throw new Error(`shape of the new value (${newValue.shape}) and previous value (${this.shape}) must match`);
|
|
}
|
|
trackerFn().disposeTensor(this);
|
|
this.dataId = newValue.dataId;
|
|
trackerFn().incRef(
|
|
this,
|
|
null
|
|
/* backend */
|
|
);
|
|
}
|
|
dispose() {
|
|
trackerFn().disposeVariable(this);
|
|
this.isDisposedInternal = true;
|
|
}
|
|
}
|
|
Object.defineProperty(Variable, Symbol.hasInstance, {
|
|
value: (instance) => {
|
|
return instance instanceof Tensor && instance.assign != null && instance.assign instanceof Function;
|
|
}
|
|
});
|
|
var Rank;
|
|
(function(Rank2) {
|
|
Rank2["R0"] = "R0";
|
|
Rank2["R1"] = "R1";
|
|
Rank2["R2"] = "R2";
|
|
Rank2["R3"] = "R3";
|
|
Rank2["R4"] = "R4";
|
|
Rank2["R5"] = "R5";
|
|
Rank2["R6"] = "R6";
|
|
})(Rank || (Rank = {}));
|
|
var UpcastInt32AndMap;
|
|
(function(UpcastInt32AndMap2) {
|
|
UpcastInt32AndMap2["float32"] = "float32";
|
|
UpcastInt32AndMap2["int32"] = "int32";
|
|
UpcastInt32AndMap2["bool"] = "int32";
|
|
UpcastInt32AndMap2["complex64"] = "complex64";
|
|
})(UpcastInt32AndMap || (UpcastInt32AndMap = {}));
|
|
var UpcastBoolAndMap;
|
|
(function(UpcastBoolAndMap2) {
|
|
UpcastBoolAndMap2["float32"] = "float32";
|
|
UpcastBoolAndMap2["int32"] = "int32";
|
|
UpcastBoolAndMap2["bool"] = "bool";
|
|
UpcastBoolAndMap2["complex64"] = "complex64";
|
|
})(UpcastBoolAndMap || (UpcastBoolAndMap = {}));
|
|
var UpcastFloat32AndMap;
|
|
(function(UpcastFloat32AndMap2) {
|
|
UpcastFloat32AndMap2["float32"] = "float32";
|
|
UpcastFloat32AndMap2["int32"] = "float32";
|
|
UpcastFloat32AndMap2["bool"] = "float32";
|
|
UpcastFloat32AndMap2["complex64"] = "complex64";
|
|
})(UpcastFloat32AndMap || (UpcastFloat32AndMap = {}));
|
|
var UpcastComplex64AndMap;
|
|
(function(UpcastComplex64AndMap2) {
|
|
UpcastComplex64AndMap2["float32"] = "complex64";
|
|
UpcastComplex64AndMap2["int32"] = "complex64";
|
|
UpcastComplex64AndMap2["bool"] = "complex64";
|
|
UpcastComplex64AndMap2["complex64"] = "complex64";
|
|
})(UpcastComplex64AndMap || (UpcastComplex64AndMap = {}));
|
|
const upcastTypeMap = {
|
|
"float32": UpcastFloat32AndMap,
|
|
"int32": UpcastInt32AndMap,
|
|
"bool": UpcastBoolAndMap,
|
|
"complex64": UpcastComplex64AndMap
|
|
};
|
|
function upcastType(typeA, typeB) {
|
|
if (typeA === "string" || typeB === "string") {
|
|
if (typeA === "string" && typeB === "string") {
|
|
return "string";
|
|
}
|
|
throw new Error(`Can not upcast ${typeA} with ${typeB}`);
|
|
}
|
|
return upcastTypeMap[typeA][typeB];
|
|
}
|
|
function sumOutType(type) {
|
|
return upcastType(type, "int32");
|
|
}
|
|
function isWebGLData(values) {
|
|
return values != null && typeof values === "object" && "texture" in values && values.texture instanceof WebGLTexture;
|
|
}
|
|
function isWebGPUData(values) {
|
|
return typeof GPUBuffer !== "undefined" && values != null && typeof values === "object" && "buffer" in values && values.buffer instanceof GPUBuffer;
|
|
}
|
|
function makeTypesMatch(a, b) {
|
|
if (a.dtype === b.dtype) {
|
|
return [a, b];
|
|
}
|
|
const dtype = upcastType(a.dtype, b.dtype);
|
|
return [a.cast(dtype), b.cast(dtype)];
|
|
}
|
|
function assertTypesMatch(a, b) {
|
|
assert(a.dtype === b.dtype, () => `The dtypes of the first(${a.dtype}) and second(${b.dtype}) input must match`);
|
|
}
|
|
function getTensorsInContainer(result) {
|
|
const list = [];
|
|
const seen = /* @__PURE__ */ new Set();
|
|
walkTensorContainer(result, list, seen);
|
|
return list;
|
|
}
|
|
function walkTensorContainer(container, list, seen) {
|
|
if (container == null) {
|
|
return;
|
|
}
|
|
if (container instanceof Tensor) {
|
|
list.push(container);
|
|
return;
|
|
}
|
|
if (!isIterable(container)) {
|
|
return;
|
|
}
|
|
const iterable = container;
|
|
for (const k in iterable) {
|
|
const val = iterable[k];
|
|
if (!seen.has(val)) {
|
|
seen.add(val);
|
|
walkTensorContainer(val, list, seen);
|
|
}
|
|
}
|
|
}
|
|
function isIterable(obj) {
|
|
return Array.isArray(obj) || typeof obj === "object";
|
|
}
|
|
function isRegisteredKernelInvocation(kernelInvocation) {
|
|
return kernelInvocation.kernelName != null;
|
|
}
|
|
class EngineState {
|
|
constructor() {
|
|
this.registeredVariables = {};
|
|
this.nextTapeNodeId = 0;
|
|
this.numBytes = 0;
|
|
this.numTensors = 0;
|
|
this.numStringTensors = 0;
|
|
this.numDataBuffers = 0;
|
|
this.gradientDepth = 0;
|
|
this.kernelDepth = 0;
|
|
this.scopeStack = [];
|
|
this.numDataMovesStack = [];
|
|
this.nextScopeId = 0;
|
|
this.tensorInfo = /* @__PURE__ */ new WeakMap();
|
|
this.profiling = false;
|
|
this.activeProfile = {
|
|
newBytes: 0,
|
|
newTensors: 0,
|
|
peakBytes: 0,
|
|
kernels: [],
|
|
result: null,
|
|
get kernelNames() {
|
|
return Array.from(new Set(this.kernels.map((k) => k.name)));
|
|
}
|
|
};
|
|
}
|
|
dispose() {
|
|
for (const variableName in this.registeredVariables) {
|
|
this.registeredVariables[variableName].dispose();
|
|
}
|
|
}
|
|
}
|
|
class Engine {
|
|
constructor(ENV2) {
|
|
this.ENV = ENV2;
|
|
this.registry = {};
|
|
this.registryFactory = {};
|
|
this.pendingBackendInitId = 0;
|
|
this.state = new EngineState();
|
|
}
|
|
async ready() {
|
|
if (this.pendingBackendInit != null) {
|
|
return this.pendingBackendInit.then(() => {
|
|
});
|
|
}
|
|
if (this.backendInstance != null) {
|
|
return;
|
|
}
|
|
const sortedBackends = this.getSortedBackends();
|
|
for (let i = 0; i < sortedBackends.length; i++) {
|
|
const backendName = sortedBackends[i];
|
|
const success = await this.initializeBackend(backendName).success;
|
|
if (success) {
|
|
await this.setBackend(backendName);
|
|
return;
|
|
}
|
|
}
|
|
throw new Error(`Could not initialize any backends, all backend initializations failed.`);
|
|
}
|
|
get backend() {
|
|
if (this.pendingBackendInit != null) {
|
|
throw new Error(`Backend '${this.backendName}' has not yet been initialized. Make sure to await tf.ready() or await tf.setBackend() before calling other methods`);
|
|
}
|
|
if (this.backendInstance == null) {
|
|
const { name, asyncInit } = this.initializeBackendsAndReturnBest();
|
|
if (asyncInit) {
|
|
throw new Error(`The highest priority backend '${name}' has not yet been initialized. Make sure to await tf.ready() or await tf.setBackend() before calling other methods`);
|
|
}
|
|
this.setBackend(name);
|
|
}
|
|
return this.backendInstance;
|
|
}
|
|
backendNames() {
|
|
return Object.keys(this.registryFactory);
|
|
}
|
|
findBackend(backendName) {
|
|
if (!(backendName in this.registry)) {
|
|
if (backendName in this.registryFactory) {
|
|
const { asyncInit } = this.initializeBackend(backendName);
|
|
if (asyncInit) {
|
|
return null;
|
|
}
|
|
} else {
|
|
return null;
|
|
}
|
|
}
|
|
return this.registry[backendName];
|
|
}
|
|
findBackendFactory(backendName) {
|
|
if (!(backendName in this.registryFactory)) {
|
|
return null;
|
|
}
|
|
return this.registryFactory[backendName].factory;
|
|
}
|
|
registerBackend(backendName, factory, priority = 1) {
|
|
if (backendName in this.registryFactory) {
|
|
warn(`${backendName} backend was already registered. Reusing existing backend factory.`);
|
|
return false;
|
|
}
|
|
this.registryFactory[backendName] = { factory, priority };
|
|
return true;
|
|
}
|
|
async setBackend(backendName) {
|
|
if (this.registryFactory[backendName] == null) {
|
|
throw new Error(`Backend name '${backendName}' not found in registry`);
|
|
}
|
|
this.backendName = backendName;
|
|
if (this.registry[backendName] == null) {
|
|
this.backendInstance = null;
|
|
const { success, asyncInit } = this.initializeBackend(backendName);
|
|
const result = asyncInit ? await success : success;
|
|
if (!result) {
|
|
return false;
|
|
}
|
|
}
|
|
this.backendInstance = this.registry[backendName];
|
|
this.setupRegisteredKernels();
|
|
this.profiler = new Profiler(this.backendInstance);
|
|
return true;
|
|
}
|
|
setupRegisteredKernels() {
|
|
const kernels = getKernelsForBackend(this.backendName);
|
|
kernels.forEach((kernel) => {
|
|
if (kernel.setupFunc != null) {
|
|
kernel.setupFunc(this.backendInstance);
|
|
}
|
|
});
|
|
}
|
|
disposeRegisteredKernels(backendName) {
|
|
const kernels = getKernelsForBackend(backendName);
|
|
kernels.forEach((kernel) => {
|
|
if (kernel.disposeFunc != null) {
|
|
kernel.disposeFunc(this.registry[backendName]);
|
|
}
|
|
});
|
|
}
|
|
/**
|
|
* Initializes a backend by looking up the backend name in the factory
|
|
* registry and calling the factory method. Returns a boolean representing
|
|
* whether the initialization of the backend succeeded. Throws an error if
|
|
* there is no backend in the factory registry.
|
|
*/
|
|
initializeBackend(backendName) {
|
|
const registryFactoryEntry = this.registryFactory[backendName];
|
|
if (registryFactoryEntry == null) {
|
|
throw new Error(`Cannot initialize backend ${backendName}, no registration found.`);
|
|
}
|
|
try {
|
|
const backend2 = registryFactoryEntry.factory();
|
|
if (backend2 && !(backend2 instanceof KernelBackend) && typeof backend2.then === "function") {
|
|
const promiseId = ++this.pendingBackendInitId;
|
|
const success = backend2.then((backendInstance) => {
|
|
if (promiseId < this.pendingBackendInitId) {
|
|
return false;
|
|
}
|
|
this.registry[backendName] = backendInstance;
|
|
this.pendingBackendInit = null;
|
|
return true;
|
|
}).catch((err) => {
|
|
if (promiseId < this.pendingBackendInitId) {
|
|
return false;
|
|
}
|
|
this.pendingBackendInit = null;
|
|
warn(`Initialization of backend ${backendName} failed`);
|
|
warn(err.stack || err.message);
|
|
return false;
|
|
});
|
|
this.pendingBackendInit = success;
|
|
return { success, asyncInit: true };
|
|
} else {
|
|
this.registry[backendName] = backend2;
|
|
return { success: true, asyncInit: false };
|
|
}
|
|
} catch (err) {
|
|
warn(`Initialization of backend ${backendName} failed`);
|
|
warn(err.stack || err.message);
|
|
return { success: false, asyncInit: false };
|
|
}
|
|
}
|
|
removeBackend(backendName) {
|
|
if (!(backendName in this.registryFactory)) {
|
|
throw new Error(`${backendName} backend not found in registry`);
|
|
}
|
|
if (this.backendName === backendName && this.pendingBackendInit != null) {
|
|
this.pendingBackendInitId++;
|
|
}
|
|
if (backendName in this.registry) {
|
|
this.disposeRegisteredKernels(backendName);
|
|
this.registry[backendName].dispose();
|
|
delete this.registry[backendName];
|
|
}
|
|
delete this.registryFactory[backendName];
|
|
if (this.backendName === backendName) {
|
|
this.pendingBackendInit = null;
|
|
this.backendName = null;
|
|
this.backendInstance = null;
|
|
}
|
|
}
|
|
getSortedBackends() {
|
|
if (Object.keys(this.registryFactory).length === 0) {
|
|
throw new Error("No backend found in registry.");
|
|
}
|
|
return Object.keys(this.registryFactory).sort((a, b) => {
|
|
return this.registryFactory[b].priority - this.registryFactory[a].priority;
|
|
});
|
|
}
|
|
initializeBackendsAndReturnBest() {
|
|
const sortedBackends = this.getSortedBackends();
|
|
for (let i = 0; i < sortedBackends.length; i++) {
|
|
const backendName = sortedBackends[i];
|
|
const { success, asyncInit } = this.initializeBackend(backendName);
|
|
if (asyncInit || success) {
|
|
return { name: backendName, asyncInit };
|
|
}
|
|
}
|
|
throw new Error(`Could not initialize any backends, all backend initializations failed.`);
|
|
}
|
|
moveData(backend2, dataId) {
|
|
const info = this.state.tensorInfo.get(dataId);
|
|
const srcBackend = info.backend;
|
|
const values = this.readSync(dataId);
|
|
const refCount = srcBackend.refCount(dataId);
|
|
srcBackend.disposeData(dataId, true);
|
|
info.backend = backend2;
|
|
backend2.move(dataId, values, info.shape, info.dtype, refCount);
|
|
if (this.shouldCheckForMemLeaks()) {
|
|
this.state.numDataMovesStack[this.state.numDataMovesStack.length - 1]++;
|
|
}
|
|
}
|
|
tidy(nameOrFn, fn) {
|
|
let name = null;
|
|
if (fn == null) {
|
|
if (typeof nameOrFn !== "function") {
|
|
throw new Error("Please provide a function to tidy()");
|
|
}
|
|
fn = nameOrFn;
|
|
} else {
|
|
if (typeof nameOrFn !== "string" && !(nameOrFn instanceof String)) {
|
|
throw new Error("When calling with two arguments, the first argument to tidy() must be a string");
|
|
}
|
|
if (typeof fn !== "function") {
|
|
throw new Error("When calling with two arguments, the 2nd argument to tidy() must be a function");
|
|
}
|
|
name = nameOrFn;
|
|
}
|
|
let result;
|
|
return this.scopedRun(() => this.startScope(name), () => this.endScope(result), () => {
|
|
result = fn();
|
|
if (result instanceof Promise) {
|
|
console.error("Cannot return a Promise inside of tidy.");
|
|
}
|
|
return result;
|
|
});
|
|
}
|
|
scopedRun(start, end, f) {
|
|
start();
|
|
try {
|
|
const res = f();
|
|
end();
|
|
return res;
|
|
} catch (ex) {
|
|
end();
|
|
throw ex;
|
|
}
|
|
}
|
|
nextTensorId() {
|
|
return Engine.nextTensorId++;
|
|
}
|
|
nextVariableId() {
|
|
return Engine.nextVariableId++;
|
|
}
|
|
/**
|
|
* This method is called instead of the public-facing tensor.clone() when
|
|
* saving a tensor for backwards pass. It makes sure to add the clone
|
|
* operation to the tape regardless of being called inside a kernel
|
|
* execution.
|
|
*/
|
|
clone(x) {
|
|
const y = ENGINE.runKernel(Identity, { x });
|
|
const inputs = { x };
|
|
const grad = (dy) => ({
|
|
x: () => {
|
|
const dtype = "float32";
|
|
const gradInputs = { x: dy };
|
|
const attrs = { dtype };
|
|
return ENGINE.runKernel(
|
|
Cast,
|
|
gradInputs,
|
|
// tslint:disable-next-line: no-unnecessary-type-assertion
|
|
attrs
|
|
);
|
|
}
|
|
});
|
|
const saved = [];
|
|
this.addTapeNode(this.state.activeScope.name, inputs, [y], grad, saved, {});
|
|
return y;
|
|
}
|
|
/**
|
|
* Execute a kernel with the given name and return the output tensor.
|
|
*
|
|
* @param kernelName The name of the kernel to execute.
|
|
* @param inputs A map of input names to tensors.
|
|
* @param attrs A map of attribute names to their values. An attribute is a
|
|
* primitive (non-tensor) input to the kernel.
|
|
* @param inputsToSave A list of tensors, inputs to save for the backprop
|
|
* computation.
|
|
* @param outputsToSave A list of booleans, specifying which output to save
|
|
* for the backprop computation. These are booleans since the output
|
|
* tensors are not visible to the user.
|
|
*/
|
|
runKernel(kernelName, inputs, attrs) {
|
|
if (this.backendName == null) {
|
|
this.backend;
|
|
}
|
|
const hasKernel = getKernel(kernelName, this.backendName) != null;
|
|
if (!hasKernel) {
|
|
throw new Error(`Kernel '${kernelName}' not registered for backend '${this.backendName}'`);
|
|
}
|
|
return this.runKernelFunc({ kernelName, inputs, attrs });
|
|
}
|
|
shouldCheckForMemLeaks() {
|
|
return this.ENV.getBool("IS_TEST");
|
|
}
|
|
checkKernelForMemLeak(kernelName, numDataIdsBefore, outInfos) {
|
|
const numDataIdsAfter = this.backend.numDataIds();
|
|
let numOutputDataIds = 0;
|
|
outInfos.forEach((info) => {
|
|
numOutputDataIds += info.dtype === "complex64" ? 3 : 1;
|
|
});
|
|
const numMoves = this.state.numDataMovesStack[this.state.numDataMovesStack.length - 1];
|
|
const dataIdsLeaked = numDataIdsAfter - numDataIdsBefore - numOutputDataIds - numMoves;
|
|
if (dataIdsLeaked > 0) {
|
|
throw new Error(`Backend '${this.backendName}' has an internal memory leak (${dataIdsLeaked} data ids) after running '${kernelName}'`);
|
|
}
|
|
}
|
|
/**
|
|
* Internal helper method to execute a kernel Func
|
|
*
|
|
* Use `runKernel` to execute kernels from outside of engine.
|
|
*/
|
|
runKernelFunc(kernelParams) {
|
|
let outputs;
|
|
let saved = [];
|
|
const isTapeOn = this.isTapeOn();
|
|
const startingBytecount = this.state.numBytes;
|
|
const startingNumTensors = this.state.numTensors;
|
|
if (this.shouldCheckForMemLeaks()) {
|
|
this.state.numDataMovesStack.push(0);
|
|
}
|
|
let kernelFunc;
|
|
if (this.backendName == null) {
|
|
this.backend;
|
|
}
|
|
let out;
|
|
const kernelOrScopeName = isRegisteredKernelInvocation(kernelParams) ? kernelParams.kernelName : this.state.activeScope != null ? this.state.activeScope.name : "";
|
|
if (isRegisteredKernelInvocation(kernelParams)) {
|
|
const { kernelName, inputs: inputs2, attrs: attrs2 } = kernelParams;
|
|
if (this.backendName == null) {
|
|
this.backend;
|
|
}
|
|
const kernel = getKernel(kernelName, this.backendName);
|
|
assert(kernel != null, () => `Cannot find registered kernel '${kernelName}' for backend '${this.backendName}'`);
|
|
kernelFunc = () => {
|
|
const numDataIdsBefore = this.backend.numDataIds();
|
|
out = kernel.kernelFunc({ inputs: inputs2, attrs: attrs2, backend: this.backend });
|
|
const outInfos = Array.isArray(out) ? out : [out];
|
|
if (this.shouldCheckForMemLeaks()) {
|
|
this.checkKernelForMemLeak(kernelName, numDataIdsBefore, outInfos);
|
|
}
|
|
const outTensors = outInfos.map((outInfo) => {
|
|
if (outInfo.rank != null) {
|
|
return outInfo;
|
|
}
|
|
return this.makeTensorFromTensorInfo(outInfo);
|
|
});
|
|
if (isTapeOn) {
|
|
const tensorsToSave = this.getTensorsForGradient(kernelName, inputs2, outTensors);
|
|
saved = this.saveTensorsForBackwardMode(tensorsToSave);
|
|
}
|
|
return outTensors;
|
|
};
|
|
} else {
|
|
const { forwardFunc } = kernelParams;
|
|
const saveFunc = (tensors) => {
|
|
if (!isTapeOn) {
|
|
return;
|
|
}
|
|
saved = tensors.map((tensor2) => this.keep(this.clone(tensor2)));
|
|
};
|
|
kernelFunc = () => {
|
|
const numDataIdsBefore = this.backend.numDataIds();
|
|
out = this.tidy(() => forwardFunc(this.backend, saveFunc));
|
|
const outs = Array.isArray(out) ? out : [out];
|
|
if (this.shouldCheckForMemLeaks()) {
|
|
this.checkKernelForMemLeak(kernelOrScopeName, numDataIdsBefore, outs);
|
|
}
|
|
return outs;
|
|
};
|
|
}
|
|
const { inputs, attrs } = kernelParams;
|
|
const backwardsFunc = isRegisteredKernelInvocation(kernelParams) ? null : kernelParams.backwardsFunc;
|
|
let kernelProfile;
|
|
this.scopedRun(
|
|
// Stop recording to a tape when running a kernel.
|
|
() => this.state.kernelDepth++,
|
|
() => this.state.kernelDepth--,
|
|
() => {
|
|
if (!this.ENV.getBool("DEBUG") && !this.state.profiling) {
|
|
outputs = kernelFunc();
|
|
} else {
|
|
kernelProfile = this.profiler.profileKernel(kernelOrScopeName, inputs, () => kernelFunc());
|
|
if (this.ENV.getBool("DEBUG")) {
|
|
this.profiler.logKernelProfile(kernelProfile);
|
|
}
|
|
outputs = kernelProfile.outputs;
|
|
}
|
|
}
|
|
);
|
|
if (isTapeOn) {
|
|
this.addTapeNode(kernelOrScopeName, inputs, outputs, backwardsFunc, saved, attrs);
|
|
}
|
|
if (this.state.profiling) {
|
|
this.state.activeProfile.kernels.push({
|
|
name: kernelOrScopeName,
|
|
bytesAdded: this.state.numBytes - startingBytecount,
|
|
totalBytesSnapshot: this.state.numBytes,
|
|
tensorsAdded: this.state.numTensors - startingNumTensors,
|
|
totalTensorsSnapshot: this.state.numTensors,
|
|
inputShapes: Object.keys(inputs).map((key) => inputs[key] != null ? inputs[key].shape : null),
|
|
outputShapes: outputs.map((item) => item.shape),
|
|
kernelTimeMs: kernelProfile.timeMs,
|
|
extraInfo: kernelProfile.extraInfo
|
|
});
|
|
}
|
|
return Array.isArray(out) ? outputs : outputs[0];
|
|
}
|
|
/**
|
|
* Saves tensors used in forward mode for use in backward mode.
|
|
*
|
|
* @param tensors the list of tensors to save.
|
|
*/
|
|
saveTensorsForBackwardMode(tensors) {
|
|
const saved = tensors.map((tensor2) => this.keep(this.clone(tensor2)));
|
|
return saved;
|
|
}
|
|
/**
|
|
* Returns a list of tensors to save for a given gradient calculation.
|
|
*
|
|
* @param kernelName name of kernel to look up gradient for.
|
|
* @param inputs a map of input tensors.
|
|
* @param outputs an array of output tensors from forward mode of kernel.
|
|
*/
|
|
getTensorsForGradient(kernelName, inputs, outputs) {
|
|
const gradConfig = getGradient(kernelName);
|
|
if (gradConfig != null) {
|
|
const inputsToSave = gradConfig.inputsToSave || [];
|
|
const outputsToSave = gradConfig.outputsToSave || [];
|
|
let inputTensorsToSave;
|
|
if (gradConfig.saveAllInputs) {
|
|
assert(Array.isArray(inputs), () => "saveAllInputs is true, expected inputs to be an array.");
|
|
inputTensorsToSave = Object.keys(inputs).map((key) => inputs[key]);
|
|
} else {
|
|
inputTensorsToSave = inputsToSave.map((inputName) => inputs[inputName]);
|
|
}
|
|
const outputTensorsToSave = outputs.filter((_, i) => outputsToSave[i]);
|
|
return inputTensorsToSave.concat(outputTensorsToSave);
|
|
}
|
|
return [];
|
|
}
|
|
/**
|
|
* Internal method used by public APIs for tensor creation. Makes a new
|
|
* tensor with the provided shape, dtype and values. It always
|
|
* creates a new data id and writes the values to the underlying backend.
|
|
*/
|
|
makeTensor(values, shape, dtype, backend2) {
|
|
if (values == null) {
|
|
throw new Error("Values passed to engine.makeTensor() are null");
|
|
}
|
|
dtype = dtype || "float32";
|
|
backend2 = backend2 || this.backend;
|
|
let backendVals = values;
|
|
if (dtype === "string" && isString(values[0])) {
|
|
backendVals = values.map((d) => encodeString(d));
|
|
}
|
|
const dataId = backend2.write(backendVals, shape, dtype);
|
|
const t2 = new Tensor(shape, dtype, dataId, this.nextTensorId());
|
|
this.trackTensor(t2, backend2);
|
|
if (dtype === "string") {
|
|
const info = this.state.tensorInfo.get(dataId);
|
|
const newBytes = bytesFromStringArray(backendVals);
|
|
this.state.numBytes += newBytes - info.bytes;
|
|
info.bytes = newBytes;
|
|
}
|
|
return t2;
|
|
}
|
|
/**
|
|
* Internal method used by backends. Makes a new tensor
|
|
* that is a wrapper around an existing data id. It doesn't create
|
|
* a new data id, only increments the ref count used in memory tracking.
|
|
* @deprecated
|
|
*/
|
|
makeTensorFromDataId(dataId, shape, dtype, backend2) {
|
|
dtype = dtype || "float32";
|
|
const tensorInfo = { dataId, shape, dtype };
|
|
return this.makeTensorFromTensorInfo(tensorInfo, backend2);
|
|
}
|
|
/**
|
|
* Internal method used by backends. Makes a new tensor that is a wrapper
|
|
* around an existing data id in TensorInfo. It doesn't create a new data id,
|
|
* only increments the ref count used in memory tracking.
|
|
*/
|
|
makeTensorFromTensorInfo(tensorInfo, backend2) {
|
|
const { dataId, shape, dtype } = tensorInfo;
|
|
const t2 = new Tensor(shape, dtype, dataId, this.nextTensorId());
|
|
this.trackTensor(t2, backend2);
|
|
return t2;
|
|
}
|
|
makeVariable(initialValue, trainable = true, name, dtype) {
|
|
name = name || this.nextVariableId().toString();
|
|
if (dtype != null && dtype !== initialValue.dtype) {
|
|
initialValue = initialValue.cast(dtype);
|
|
}
|
|
const v = new Variable(initialValue, trainable, name, this.nextTensorId());
|
|
if (this.state.registeredVariables[v.name] != null) {
|
|
throw new Error(`Variable with name ${v.name} was already registered`);
|
|
}
|
|
this.state.registeredVariables[v.name] = v;
|
|
this.incRef(v, this.backend);
|
|
return v;
|
|
}
|
|
trackTensor(a, backend2) {
|
|
this.state.numTensors++;
|
|
if (a.dtype === "string") {
|
|
this.state.numStringTensors++;
|
|
}
|
|
let bytes = 0;
|
|
if (a.dtype !== "complex64" && a.dtype !== "string") {
|
|
bytes = a.size * bytesPerElement(a.dtype);
|
|
}
|
|
this.state.numBytes += bytes;
|
|
if (!this.state.tensorInfo.has(a.dataId)) {
|
|
this.state.numDataBuffers++;
|
|
this.state.tensorInfo.set(a.dataId, {
|
|
backend: backend2 || this.backend,
|
|
dtype: a.dtype,
|
|
shape: a.shape,
|
|
bytes
|
|
});
|
|
}
|
|
if (!(a instanceof Variable)) {
|
|
this.track(a);
|
|
}
|
|
}
|
|
// Track the tensor by dataId and increase the refCount for the dataId in the
|
|
// backend.
|
|
// TODO(pyu10055): This is currently used by makeVariable method, to increase
|
|
// refCount on the backend for the dataId. It can potentially be replaced with
|
|
// Identity op indead of calling backend directly.
|
|
incRef(a, backend2) {
|
|
this.trackTensor(a, backend2);
|
|
this.backend.incRef(a.dataId);
|
|
}
|
|
removeDataId(dataId, backend2) {
|
|
if (this.state.tensorInfo.has(dataId) && this.state.tensorInfo.get(dataId).backend === backend2) {
|
|
this.state.tensorInfo.delete(dataId);
|
|
this.state.numDataBuffers--;
|
|
}
|
|
}
|
|
disposeTensor(a) {
|
|
if (!this.state.tensorInfo.has(a.dataId)) {
|
|
return;
|
|
}
|
|
const info = this.state.tensorInfo.get(a.dataId);
|
|
this.state.numTensors--;
|
|
if (a.dtype === "string") {
|
|
this.state.numStringTensors--;
|
|
this.state.numBytes -= info.bytes;
|
|
}
|
|
if (a.dtype !== "complex64" && a.dtype !== "string") {
|
|
const bytes = a.size * bytesPerElement(a.dtype);
|
|
this.state.numBytes -= bytes;
|
|
}
|
|
if (info.backend.disposeData(a.dataId)) {
|
|
this.removeDataId(a.dataId, info.backend);
|
|
}
|
|
}
|
|
disposeVariables() {
|
|
for (const varName in this.state.registeredVariables) {
|
|
const v = this.state.registeredVariables[varName];
|
|
this.disposeVariable(v);
|
|
}
|
|
}
|
|
disposeVariable(v) {
|
|
this.disposeTensor(v);
|
|
if (this.state.registeredVariables[v.name] != null) {
|
|
delete this.state.registeredVariables[v.name];
|
|
}
|
|
}
|
|
memory() {
|
|
const info = this.backend.memory();
|
|
info.numTensors = this.state.numTensors;
|
|
info.numDataBuffers = this.state.numDataBuffers;
|
|
info.numBytes = this.state.numBytes;
|
|
if (this.state.numStringTensors > 0) {
|
|
info.unreliable = true;
|
|
if (info.reasons == null) {
|
|
info.reasons = [];
|
|
}
|
|
info.reasons.push("Memory usage by string tensors is approximate (2 bytes per character)");
|
|
}
|
|
return info;
|
|
}
|
|
async profile(query) {
|
|
this.state.profiling = true;
|
|
const startBytes = this.state.numBytes;
|
|
const startNumTensors = this.state.numTensors;
|
|
this.state.activeProfile.kernels = [];
|
|
this.state.activeProfile.result = await query();
|
|
this.state.profiling = false;
|
|
this.state.activeProfile.peakBytes = Math.max(...this.state.activeProfile.kernels.map((d) => d.totalBytesSnapshot));
|
|
this.state.activeProfile.newBytes = this.state.numBytes - startBytes;
|
|
this.state.activeProfile.newTensors = this.state.numTensors - startNumTensors;
|
|
for (const kernel of this.state.activeProfile.kernels) {
|
|
kernel.kernelTimeMs = await kernel.kernelTimeMs;
|
|
kernel.extraInfo = await kernel.extraInfo;
|
|
}
|
|
return this.state.activeProfile;
|
|
}
|
|
isTapeOn() {
|
|
return this.state.gradientDepth > 0 && this.state.kernelDepth === 0;
|
|
}
|
|
addTapeNode(kernelName, inputs, outputs, gradientsFunc, saved, attrs) {
|
|
const tapeNode = { id: this.state.nextTapeNodeId++, kernelName, inputs, outputs, saved };
|
|
const gradConfig = getGradient(kernelName);
|
|
if (gradConfig != null) {
|
|
gradientsFunc = gradConfig.gradFunc;
|
|
}
|
|
if (gradientsFunc != null) {
|
|
tapeNode.gradient = (dys) => {
|
|
dys = dys.map((dy, i) => {
|
|
if (dy == null) {
|
|
const output = outputs[i];
|
|
const vals = makeZerosTypedArray(output.size, output.dtype);
|
|
return this.makeTensor(vals, output.shape, output.dtype);
|
|
}
|
|
return dy;
|
|
});
|
|
return gradientsFunc(dys.length > 1 ? dys : dys[0], saved, attrs);
|
|
};
|
|
}
|
|
this.state.activeTape.push(tapeNode);
|
|
}
|
|
keep(result) {
|
|
result.kept = true;
|
|
return result;
|
|
}
|
|
startTape() {
|
|
if (this.state.gradientDepth === 0) {
|
|
this.state.activeTape = [];
|
|
}
|
|
this.state.gradientDepth++;
|
|
}
|
|
endTape() {
|
|
this.state.gradientDepth--;
|
|
}
|
|
/**
|
|
* Start a scope. Use this with endScope() to achieve the same functionality
|
|
* as scope() without the need for a function closure.
|
|
*/
|
|
startScope(name) {
|
|
const scopeInfo = {
|
|
track: [],
|
|
name: "unnamed scope",
|
|
id: this.state.nextScopeId++
|
|
};
|
|
if (name) {
|
|
scopeInfo.name = name;
|
|
}
|
|
this.state.scopeStack.push(scopeInfo);
|
|
this.state.activeScope = scopeInfo;
|
|
}
|
|
/**
|
|
* End a scope. Use this with startScope() to achieve the same functionality
|
|
* as scope() without the need for a function closure.
|
|
*/
|
|
endScope(result) {
|
|
const tensorsToTrackInParent = getTensorsInContainer(result);
|
|
const tensorsToTrackInParentSet = new Set(tensorsToTrackInParent.map((t2) => t2.id));
|
|
for (let i = 0; i < this.state.activeScope.track.length; i++) {
|
|
const tensor2 = this.state.activeScope.track[i];
|
|
if (!tensor2.kept && !tensorsToTrackInParentSet.has(tensor2.id)) {
|
|
tensor2.dispose();
|
|
}
|
|
}
|
|
const oldScope = this.state.scopeStack.pop();
|
|
this.state.activeScope = this.state.scopeStack.length === 0 ? null : this.state.scopeStack[this.state.scopeStack.length - 1];
|
|
tensorsToTrackInParent.forEach((tensor2) => {
|
|
if (!tensor2.kept && tensor2.scopeId === oldScope.id) {
|
|
this.track(tensor2);
|
|
}
|
|
});
|
|
}
|
|
/**
|
|
* Returns gradients of `f` with respect to each of the `xs`. The gradients
|
|
* returned are of the same length as `xs`, but some might be null if `f`
|
|
* was not a function of that `x`. It also takes optional dy to multiply the
|
|
* gradient, which defaults to `1`.
|
|
*/
|
|
gradients(f, xs, dy, allowNoGradients = false) {
|
|
assert(xs.length > 0, () => "gradients() received an empty list of xs.");
|
|
if (dy != null && dy.dtype !== "float32") {
|
|
throw new Error(`dy must have 'float32' dtype, but has '${dy.dtype}'`);
|
|
}
|
|
const y = this.scopedRun(() => this.startTape(), () => this.endTape(), () => this.tidy("forward", f));
|
|
assert(y instanceof Tensor, () => "The result y returned by f() must be a tensor.");
|
|
const filteredTape = getFilteredNodesXToY(this.state.activeTape, xs, y);
|
|
if (!allowNoGradients && filteredTape.length === 0 && xs.length > 0) {
|
|
throw new Error("Cannot compute gradient of y=f(x) with respect to x. Make sure that the f you passed encloses all operations that lead from x to y.");
|
|
}
|
|
return this.tidy("backward", () => {
|
|
const accumulatedGradientMap = {};
|
|
accumulatedGradientMap[y.id] = dy == null ? ones$1(y.shape) : dy;
|
|
backpropagateGradients(
|
|
accumulatedGradientMap,
|
|
filteredTape,
|
|
// Pass the tidy function to avoid circular dep with `tape.ts`.
|
|
(f2) => this.tidy(f2),
|
|
// Pass an add function to avoide a circular dep with `tape.ts`.
|
|
add$2
|
|
);
|
|
const grads = xs.map((x) => accumulatedGradientMap[x.id]);
|
|
if (this.state.gradientDepth === 0) {
|
|
this.state.activeTape.forEach((node) => {
|
|
for (const tensor2 of node.saved) {
|
|
tensor2.dispose();
|
|
}
|
|
});
|
|
this.state.activeTape = null;
|
|
}
|
|
return { value: y, grads };
|
|
});
|
|
}
|
|
customGrad(f) {
|
|
assert(isFunction(f), () => "The f passed in customGrad(f) must be a function.");
|
|
return (...inputs) => {
|
|
assert(inputs.every((t2) => t2 instanceof Tensor), () => "The args passed in customGrad(f)(x1, x2,...) must all be tensors");
|
|
let res;
|
|
const inputMap = {};
|
|
inputs.forEach((input, i) => {
|
|
inputMap[i] = input;
|
|
});
|
|
const forwardFunc = (_, save) => {
|
|
res = f(...[...inputs, save]);
|
|
assert(res.value instanceof Tensor, () => "The function f passed in customGrad(f) must return an object where `obj.value` is a tensor");
|
|
assert(isFunction(res.gradFunc), () => "The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function.");
|
|
return res.value;
|
|
};
|
|
const backwardsFunc = (dy, saved) => {
|
|
const gradRes = res.gradFunc(dy, saved);
|
|
const grads = Array.isArray(gradRes) ? gradRes : [gradRes];
|
|
assert(grads.length === inputs.length, () => "The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function that returns the same number of tensors as inputs passed to f(...).");
|
|
assert(grads.every((t2) => t2 instanceof Tensor), () => "The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function that returns a list of only tensors.");
|
|
const gradMap = {};
|
|
grads.forEach((grad, i) => {
|
|
gradMap[i] = () => grad;
|
|
});
|
|
return gradMap;
|
|
};
|
|
return this.runKernelFunc({
|
|
forwardFunc,
|
|
backwardsFunc,
|
|
inputs: inputMap
|
|
});
|
|
};
|
|
}
|
|
readSync(dataId) {
|
|
const info = this.state.tensorInfo.get(dataId);
|
|
return info.backend.readSync(dataId);
|
|
}
|
|
read(dataId) {
|
|
const info = this.state.tensorInfo.get(dataId);
|
|
return info.backend.read(dataId);
|
|
}
|
|
readToGPU(dataId, options) {
|
|
const info = this.state.tensorInfo.get(dataId);
|
|
return info.backend.readToGPU(dataId, options);
|
|
}
|
|
async time(query) {
|
|
const start = now();
|
|
const timingInfo = await this.backend.time(query);
|
|
timingInfo.wallMs = now() - start;
|
|
return timingInfo;
|
|
}
|
|
/**
|
|
* Tracks a Tensor in the current scope to be automatically cleaned up
|
|
* when the current scope ends, and returns the value.
|
|
*
|
|
* @param result The Tensor to track in the current scope.
|
|
*/
|
|
track(result) {
|
|
if (this.state.activeScope != null) {
|
|
result.scopeId = this.state.activeScope.id;
|
|
this.state.activeScope.track.push(result);
|
|
}
|
|
return result;
|
|
}
|
|
get registeredVariables() {
|
|
return this.state.registeredVariables;
|
|
}
|
|
/**
|
|
* Resets the engine state. Removes all backends but does not remove
|
|
* registered backend factories.
|
|
*/
|
|
reset() {
|
|
this.pendingBackendInitId++;
|
|
this.state.dispose();
|
|
this.ENV.reset();
|
|
this.state = new EngineState();
|
|
for (const backendName in this.registry) {
|
|
this.disposeRegisteredKernels(backendName);
|
|
this.registry[backendName].dispose();
|
|
delete this.registry[backendName];
|
|
}
|
|
this.backendName = null;
|
|
this.backendInstance = null;
|
|
this.pendingBackendInit = null;
|
|
}
|
|
}
|
|
Engine.nextTensorId = 0;
|
|
Engine.nextVariableId = 0;
|
|
function ones$1(shape) {
|
|
const values = makeOnesTypedArray(sizeFromShape(shape), "float32");
|
|
return ENGINE.makeTensor(values, shape, "float32");
|
|
}
|
|
function getOrMakeEngine() {
|
|
const ns = getGlobalNamespace();
|
|
if (ns._tfengine == null) {
|
|
const environment = new Environment(ns);
|
|
ns._tfengine = new Engine(environment);
|
|
}
|
|
setEnvironmentGlobal(ns._tfengine.ENV);
|
|
setTensorTracker(() => ns._tfengine);
|
|
return ns._tfengine;
|
|
}
|
|
const ENGINE = getOrMakeEngine();
|
|
function add$2(a, b) {
|
|
const inputs = { a, b };
|
|
return ENGINE.runKernel(Add, inputs);
|
|
}
|
|
function _isNavigatorDefined() {
|
|
return typeof navigator !== "undefined" && navigator != null;
|
|
}
|
|
function isMobile(nav) {
|
|
if (nav || _isNavigatorDefined()) {
|
|
if (!nav) {
|
|
nav = navigator;
|
|
}
|
|
if (nav.product === "ReactNative") {
|
|
return true;
|
|
}
|
|
const a = nav.userAgent || nav.vendor || // tslint:disable-next-line:no-any
|
|
(typeof window !== "undefined" ? window.opera : "");
|
|
if (!a) {
|
|
const navAny = nav;
|
|
return navAny.userAgentData && navAny.userAgentData.mobile;
|
|
}
|
|
return /(android|bb\d+|meego).+mobile|avantgo|bada\/|blackberry|blazer|compal|elaine|fennec|hiptop|iemobile|ip(hone|od)|iris|kindle|lge |maemo|midp|mmp|mobile.+firefox|netfront|opera m(ob|in)i|palm( os)?|phone|p(ixi|re)\/|plucker|pocket|psp|series(4|6)0|symbian|treo|up\.(browser|link)|vodafone|wap|windows ce|xda|xiino/i.test(a) || // tslint:disable-next-line:max-line-length
|
|
/1207|6310|6590|3gso|4thp|50[1-6]i|770s|802s|a wa|abac|ac(er|oo|s\-)|ai(ko|rn)|al(av|ca|co)|amoi|an(ex|ny|yw)|aptu|ar(ch|go)|as(te|us)|attw|au(di|\-m|r |s )|avan|be(ck|ll|nq)|bi(lb|rd)|bl(ac|az)|br(e|v)w|bumb|bw\-(n|u)|c55\/|capi|ccwa|cdm\-|cell|chtm|cldc|cmd\-|co(mp|nd)|craw|da(it|ll|ng)|dbte|dc\-s|devi|dica|dmob|do(c|p)o|ds(12|\-d)|el(49|ai)|em(l2|ul)|er(ic|k0)|esl8|ez([4-7]0|os|wa|ze)|fetc|fly(\-|_)|g1 u|g560|gene|gf\-5|g\-mo|go(\.w|od)|gr(ad|un)|haie|hcit|hd\-(m|p|t)|hei\-|hi(pt|ta)|hp( i|ip)|hs\-c|ht(c(\-| |_|a|g|p|s|t)|tp)|hu(aw|tc)|i\-(20|go|ma)|i230|iac( |\-|\/)|ibro|idea|ig01|ikom|im1k|inno|ipaq|iris|ja(t|v)a|jbro|jemu|jigs|kddi|keji|kgt( |\/)|klon|kpt |kwc\-|kyo(c|k)|le(no|xi)|lg( g|\/(k|l|u)|50|54|\-[a-w])|libw|lynx|m1\-w|m3ga|m50\/|ma(te|ui|xo)|mc(01|21|ca)|m\-cr|me(rc|ri)|mi(o8|oa|ts)|mmef|mo(01|02|bi|de|do|t(\-| |o|v)|zz)|mt(50|p1|v )|mwbp|mywa|n10[0-2]|n20[2-3]|n30(0|2)|n50(0|2|5)|n7(0(0|1)|10)|ne((c|m)\-|on|tf|wf|wg|wt)|nok(6|i)|nzph|o2im|op(ti|wv)|oran|owg1|p800|pan(a|d|t)|pdxg|pg(13|\-([1-8]|c))|phil|pire|pl(ay|uc)|pn\-2|po(ck|rt|se)|prox|psio|pt\-g|qa\-a|qc(07|12|21|32|60|\-[2-7]|i\-)|qtek|r380|r600|raks|rim9|ro(ve|zo)|s55\/|sa(ge|ma|mm|ms|ny|va)|sc(01|h\-|oo|p\-)|sdk\/|se(c(\-|0|1)|47|mc|nd|ri)|sgh\-|shar|sie(\-|m)|sk\-0|sl(45|id)|sm(al|ar|b3|it|t5)|so(ft|ny)|sp(01|h\-|v\-|v )|sy(01|mb)|t2(18|50)|t6(00|10|18)|ta(gt|lk)|tcl\-|tdg\-|tel(i|m)|tim\-|t\-mo|to(pl|sh)|ts(70|m\-|m3|m5)|tx\-9|up(\.b|g1|si)|utst|v400|v750|veri|vi(rg|te)|vk(40|5[0-3]|\-v)|vm40|voda|vulc|vx(52|53|60|61|70|80|81|83|85|98)|w3c(\-| )|webc|whit|wi(g |nc|nw)|wmlb|wonu|x700|yas\-|your|zeto|zte\-/i.test(a.substr(0, 4));
|
|
}
|
|
return false;
|
|
}
|
|
function isBrowser() {
|
|
return typeof window !== "undefined" && window.document != null || //@ts-ignore
|
|
typeof WorkerGlobalScope !== "undefined";
|
|
}
|
|
const ENV$2 = env();
|
|
ENV$2.registerFlag("DEBUG", () => false, (debugValue) => {
|
|
if (debugValue) {
|
|
console.warn("Debugging mode is ON. The output of every math call will be downloaded to CPU and checked for NaNs. This significantly impacts performance.");
|
|
}
|
|
});
|
|
ENV$2.registerFlag("IS_BROWSER", () => isBrowser());
|
|
ENV$2.registerFlag("IS_NODE", () => typeof process !== "undefined" && typeof process.versions !== "undefined" && typeof process.versions.node !== "undefined");
|
|
ENV$2.registerFlag("IS_CHROME", () => typeof navigator !== "undefined" && navigator != null && navigator.userAgent != null && /Chrome/.test(navigator.userAgent) && /Google Inc/.test(navigator.vendor));
|
|
ENV$2.registerFlag("IS_SAFARI", () => typeof navigator !== "undefined" && navigator != null && navigator.userAgent != null && /Safari/.test(navigator.userAgent) && /Apple/.test(navigator.vendor));
|
|
ENV$2.registerFlag("PROD", () => false);
|
|
ENV$2.registerFlag("TENSORLIKE_CHECK_SHAPE_CONSISTENCY", () => ENV$2.getBool("DEBUG"));
|
|
ENV$2.registerFlag("DEPRECATION_WARNINGS_ENABLED", () => true);
|
|
ENV$2.registerFlag("IS_TEST", () => false);
|
|
ENV$2.registerFlag("CHECK_COMPUTATION_FOR_ERRORS", () => ENV$2.getBool("DEBUG"));
|
|
ENV$2.registerFlag("WRAP_TO_IMAGEBITMAP", () => false);
|
|
ENV$2.registerFlag("CANVAS2D_WILL_READ_FREQUENTLY_FOR_GPU", () => false);
|
|
ENV$2.registerFlag("USE_SETTIMEOUTCUSTOM", () => false);
|
|
function inferShape(val, dtype) {
|
|
let firstElem = val;
|
|
if (isTypedArray(val)) {
|
|
return dtype === "string" ? [] : [val.length];
|
|
}
|
|
if (isWebGLData(val)) {
|
|
const usedChannels = val.channels || "RGBA";
|
|
return [val.height, val.width * usedChannels.length];
|
|
} else if (isWebGPUData(val)) {
|
|
return [val.buffer.size / (dtype == null ? 4 : bytesPerElement(dtype))];
|
|
}
|
|
if (!Array.isArray(val)) {
|
|
return [];
|
|
}
|
|
const shape = [];
|
|
while (Array.isArray(firstElem) || isTypedArray(firstElem) && dtype !== "string") {
|
|
shape.push(firstElem.length);
|
|
firstElem = firstElem[0];
|
|
}
|
|
if (Array.isArray(val) && env().getBool("TENSORLIKE_CHECK_SHAPE_CONSISTENCY")) {
|
|
deepAssertShapeConsistency(val, shape, []);
|
|
}
|
|
return shape;
|
|
}
|
|
function deepAssertShapeConsistency(val, shape, indices) {
|
|
indices = indices || [];
|
|
if (!Array.isArray(val) && !isTypedArray(val)) {
|
|
assert(shape.length === 0, () => `Element arr[${indices.join("][")}] is a primitive, but should be an array/TypedArray of ${shape[0]} elements`);
|
|
return;
|
|
}
|
|
assert(shape.length > 0, () => `Element arr[${indices.join("][")}] should be a primitive, but is an array of ${val.length} elements`);
|
|
assert(val.length === shape[0], () => `Element arr[${indices.join("][")}] should have ${shape[0]} elements, but has ${val.length} elements`);
|
|
const subShape = shape.slice(1);
|
|
for (let i = 0; i < val.length; ++i) {
|
|
deepAssertShapeConsistency(val[i], subShape, indices.concat(i));
|
|
}
|
|
}
|
|
function assertDtype(expectedDtype, actualDType, argName, functionName) {
|
|
if (expectedDtype === "string_or_numeric") {
|
|
return;
|
|
}
|
|
if (expectedDtype == null) {
|
|
throw new Error(`Expected dtype cannot be null.`);
|
|
}
|
|
if (expectedDtype !== "numeric" && expectedDtype !== actualDType || expectedDtype === "numeric" && actualDType === "string") {
|
|
throw new Error(`Argument '${argName}' passed to '${functionName}' must be ${expectedDtype} tensor, but got ${actualDType} tensor`);
|
|
}
|
|
}
|
|
function convertToTensor(x, argName, functionName, parseAsDtype = "numeric") {
|
|
if (x instanceof getGlobalTensorClass()) {
|
|
assertDtype(parseAsDtype, x.dtype, argName, functionName);
|
|
return x;
|
|
}
|
|
let inferredDtype = inferDtype(x);
|
|
if (inferredDtype !== "string" && ["bool", "int32", "float32"].indexOf(parseAsDtype) >= 0) {
|
|
inferredDtype = parseAsDtype;
|
|
}
|
|
assertDtype(parseAsDtype, inferredDtype, argName, functionName);
|
|
if (x == null || !isTypedArray(x) && !Array.isArray(x) && typeof x !== "number" && typeof x !== "boolean" && typeof x !== "string") {
|
|
const type = x == null ? "null" : x.constructor.name;
|
|
throw new Error(`Argument '${argName}' passed to '${functionName}' must be a Tensor or TensorLike, but got '${type}'`);
|
|
}
|
|
const inferredShape = inferShape(x, inferredDtype);
|
|
if (!isTypedArray(x) && !Array.isArray(x)) {
|
|
x = [x];
|
|
}
|
|
const skipTypedArray = true;
|
|
const values = inferredDtype !== "string" ? toTypedArray(x, inferredDtype) : flatten(x, [], skipTypedArray);
|
|
return ENGINE.makeTensor(values, inferredShape, inferredDtype);
|
|
}
|
|
function convertToTensorArray(arg, argName, functionName, parseAsDtype = "numeric") {
|
|
if (!Array.isArray(arg)) {
|
|
throw new Error(`Argument ${argName} passed to ${functionName} must be a \`Tensor[]\` or \`TensorLike[]\``);
|
|
}
|
|
const tensors = arg;
|
|
return tensors.map((t2, i) => convertToTensor(t2, `${argName}[${i}]`, functionName, parseAsDtype));
|
|
}
|
|
const OP_SCOPE_SUFFIX = "__op";
|
|
function op(f) {
|
|
const keys = Object.keys(f);
|
|
if (keys.length !== 1) {
|
|
throw new Error(`Please provide an object with a single key (operation name) mapping to a function. Got an object with ${keys.length} keys.`);
|
|
}
|
|
let opName = keys[0];
|
|
const fn = f[opName];
|
|
if (opName.endsWith("_")) {
|
|
opName = opName.substring(0, opName.length - 1);
|
|
}
|
|
opName = opName + OP_SCOPE_SUFFIX;
|
|
const f2 = (...args) => {
|
|
ENGINE.startScope(opName);
|
|
try {
|
|
const result = fn(...args);
|
|
if (isPromise(result)) {
|
|
console.error("Cannot return a Promise inside of tidy.");
|
|
}
|
|
ENGINE.endScope(result);
|
|
return result;
|
|
} catch (ex) {
|
|
ENGINE.endScope(null);
|
|
throw ex;
|
|
}
|
|
};
|
|
Object.defineProperty(f2, "name", { value: opName, configurable: true });
|
|
return f2;
|
|
}
|
|
function complex_(real2, imag2) {
|
|
const $real = convertToTensor(real2, "real", "complex");
|
|
const $imag = convertToTensor(imag2, "imag", "complex");
|
|
assertShapesMatch($real.shape, $imag.shape, `real and imag shapes, ${$real.shape} and ${$imag.shape}, must match in call to tf.complex().`);
|
|
const inputs = { real: $real, imag: $imag };
|
|
return ENGINE.runKernel(Complex, inputs);
|
|
}
|
|
const complex$2 = /* @__PURE__ */ op({ complex_ });
|
|
function makeTensor(values, shape, inferredShape, dtype) {
|
|
if (dtype == null) {
|
|
dtype = inferDtype(values);
|
|
} else if (dtype === "complex64") {
|
|
throw new Error(`Cannot construct a complex64 tensor directly. Please use tf.complex(real, imag).`);
|
|
}
|
|
if (isWebGPUData(values) || isWebGLData(values)) {
|
|
if (dtype !== "float32" && dtype !== "int32") {
|
|
throw new Error(`Creating tensor from GPU data only supports 'float32'|'int32' dtype, while the dtype is ${dtype}.`);
|
|
}
|
|
return ENGINE.backend.createTensorFromGPUData(values, shape || inferredShape, dtype);
|
|
}
|
|
if (!isTypedArray(values) && !Array.isArray(values) && typeof values !== "number" && typeof values !== "boolean" && typeof values !== "string") {
|
|
throw new Error("values passed to tensor(values) must be a number/boolean/string or an array of numbers/booleans/strings, or a TypedArray");
|
|
}
|
|
if (shape != null) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
const providedSize = sizeFromShape(shape);
|
|
const inferredSize = sizeFromShape(inferredShape);
|
|
assert(providedSize === inferredSize, () => `Based on the provided shape, [${shape}], the tensor should have ${providedSize} values but has ${inferredSize}`);
|
|
for (let i = 0; i < inferredShape.length; ++i) {
|
|
const inferred = inferredShape[i];
|
|
const flatDimsDontMatch = i === inferredShape.length - 1 ? inferred !== sizeFromShape(shape.slice(i)) : true;
|
|
assert(inferredShape[i] === shape[i] || !flatDimsDontMatch, () => `Error creating a new Tensor. Inferred shape (${inferredShape}) does not match the provided shape (${shape}). `);
|
|
}
|
|
}
|
|
if (!isTypedArray(values) && !Array.isArray(values)) {
|
|
values = [values];
|
|
}
|
|
shape = shape || inferredShape;
|
|
values = dtype !== "string" ? toTypedArray(values, dtype) : flatten(values, [], true);
|
|
return ENGINE.makeTensor(values, shape, dtype);
|
|
}
|
|
function tensor(values, shape, dtype) {
|
|
const inferredShape = inferShape(values, dtype);
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
const DTYPE_VALUE_SIZE_MAP = {
|
|
"float32": 4,
|
|
"float16": 2,
|
|
"int32": 4,
|
|
"uint16": 2,
|
|
"uint8": 1,
|
|
"bool": 1,
|
|
"complex64": 8
|
|
};
|
|
class CompositeArrayBuffer {
|
|
/**
|
|
* Concatenate a number of ArrayBuffers into one.
|
|
*
|
|
* @param buffers An array of ArrayBuffers to concatenate, or a single
|
|
* ArrayBuffer.
|
|
* @returns Result of concatenating `buffers` in order.
|
|
*/
|
|
static join(buffers) {
|
|
return new CompositeArrayBuffer(buffers).slice();
|
|
}
|
|
constructor(buffers) {
|
|
this.shards = [];
|
|
this.previousShardIndex = 0;
|
|
if (buffers == null) {
|
|
return;
|
|
}
|
|
if (!(buffers instanceof Array)) {
|
|
buffers = [buffers];
|
|
}
|
|
buffers = buffers.map((bufferOrTypedArray) => {
|
|
if (isTypedArray(bufferOrTypedArray)) {
|
|
return bufferOrTypedArray.buffer;
|
|
}
|
|
return bufferOrTypedArray;
|
|
});
|
|
if (buffers.length === 0) {
|
|
return;
|
|
}
|
|
this.bufferUniformSize = buffers[0].byteLength;
|
|
let start = 0;
|
|
for (let i = 0; i < buffers.length; i++) {
|
|
const buffer2 = buffers[i];
|
|
if (i !== buffers.length - 1 && buffer2.byteLength !== this.bufferUniformSize) {
|
|
this.bufferUniformSize = void 0;
|
|
}
|
|
const end = start + buffer2.byteLength;
|
|
this.shards.push({ buffer: buffer2, start, end });
|
|
start = end;
|
|
}
|
|
if (this.shards.length === 0) {
|
|
this.byteLength = 0;
|
|
}
|
|
this.byteLength = this.shards[this.shards.length - 1].end;
|
|
}
|
|
slice(start = 0, end = this.byteLength) {
|
|
if (this.shards.length === 0) {
|
|
return new ArrayBuffer(0);
|
|
}
|
|
start = isNaN(Number(start)) ? 0 : start;
|
|
end = isNaN(Number(end)) ? 0 : end;
|
|
start = Math.max(0, start);
|
|
end = Math.min(this.byteLength, end);
|
|
if (end <= start) {
|
|
return new ArrayBuffer(0);
|
|
}
|
|
const startShardIndex = this.findShardForByte(start);
|
|
if (startShardIndex === -1) {
|
|
throw new Error(`Could not find start shard for byte ${start}`);
|
|
}
|
|
const size = end - start;
|
|
const outputBuffer = new ArrayBuffer(size);
|
|
const outputArray = new Uint8Array(outputBuffer);
|
|
let sliced = 0;
|
|
for (let i = startShardIndex; i < this.shards.length; i++) {
|
|
const shard = this.shards[i];
|
|
const globalStart = start + sliced;
|
|
const localStart = globalStart - shard.start;
|
|
const outputStart = sliced;
|
|
const globalEnd = Math.min(end, shard.end);
|
|
const localEnd = globalEnd - shard.start;
|
|
const outputSlice = new Uint8Array(shard.buffer, localStart, localEnd - localStart);
|
|
outputArray.set(outputSlice, outputStart);
|
|
sliced += outputSlice.length;
|
|
if (end < shard.end) {
|
|
break;
|
|
}
|
|
}
|
|
return outputBuffer;
|
|
}
|
|
/**
|
|
* Get the index of the shard that contains the byte at `byteIndex`.
|
|
*/
|
|
findShardForByte(byteIndex) {
|
|
if (this.shards.length === 0 || byteIndex < 0 || byteIndex >= this.byteLength) {
|
|
return -1;
|
|
}
|
|
if (this.bufferUniformSize != null) {
|
|
this.previousShardIndex = Math.floor(byteIndex / this.bufferUniformSize);
|
|
return this.previousShardIndex;
|
|
}
|
|
function check(shard) {
|
|
if (byteIndex < shard.start) {
|
|
return -1;
|
|
}
|
|
if (byteIndex >= shard.end) {
|
|
return 1;
|
|
}
|
|
return 0;
|
|
}
|
|
if (check(this.shards[this.previousShardIndex]) === 0) {
|
|
return this.previousShardIndex;
|
|
}
|
|
const index = search(this.shards, check);
|
|
if (index === -1) {
|
|
return -1;
|
|
}
|
|
this.previousShardIndex = index;
|
|
return this.previousShardIndex;
|
|
}
|
|
}
|
|
function search(sortedArray, compare) {
|
|
let min2 = 0;
|
|
let max2 = sortedArray.length;
|
|
while (min2 <= max2) {
|
|
const middle = Math.floor((max2 - min2) / 2) + min2;
|
|
const side = compare(sortedArray[middle]);
|
|
if (side === 0) {
|
|
return middle;
|
|
} else if (side < 0) {
|
|
max2 = middle;
|
|
} else {
|
|
min2 = middle + 1;
|
|
}
|
|
}
|
|
return -1;
|
|
}
|
|
function engine() {
|
|
return ENGINE;
|
|
}
|
|
function tidy(nameOrFn, fn) {
|
|
return ENGINE.tidy(nameOrFn, fn);
|
|
}
|
|
function dispose(container) {
|
|
const tensors = getTensorsInContainer(container);
|
|
tensors.forEach((tensor2) => tensor2.dispose());
|
|
}
|
|
function keep(result) {
|
|
return ENGINE.keep(result);
|
|
}
|
|
function setBackend(backendName) {
|
|
return ENGINE.setBackend(backendName);
|
|
}
|
|
function getBackend() {
|
|
return ENGINE.backendName;
|
|
}
|
|
function registerBackend(name, factory, priority = 1) {
|
|
return ENGINE.registerBackend(name, factory, priority);
|
|
}
|
|
function backend() {
|
|
return ENGINE.backend;
|
|
}
|
|
const NUM_BYTES_STRING_LENGTH = 4;
|
|
async function encodeWeights(tensors, group) {
|
|
const specs = [];
|
|
const dataPromises = [];
|
|
const names = Array.isArray(tensors) ? tensors.map((tensor2) => tensor2.name) : Object.keys(tensors);
|
|
for (let i = 0; i < names.length; ++i) {
|
|
const name = names[i];
|
|
const t2 = Array.isArray(tensors) ? tensors[i].tensor : tensors[name];
|
|
if (t2.dtype !== "float32" && t2.dtype !== "int32" && t2.dtype !== "bool" && t2.dtype !== "string" && t2.dtype !== "complex64") {
|
|
throw new Error(`Unsupported dtype in weight '${name}': ${t2.dtype}`);
|
|
}
|
|
const spec = { name, shape: t2.shape, dtype: t2.dtype };
|
|
if (t2.dtype === "string") {
|
|
const utf8bytes = new Promise(async (resolve) => {
|
|
const vals = await t2.bytes();
|
|
const totalNumBytes = vals.reduce((p2, c) => p2 + c.length, 0) + NUM_BYTES_STRING_LENGTH * vals.length;
|
|
const bytes = new Uint8Array(totalNumBytes);
|
|
let offset = 0;
|
|
for (let i2 = 0; i2 < vals.length; i2++) {
|
|
const val = vals[i2];
|
|
const bytesOfLength = new Uint8Array(new Uint32Array([val.length]).buffer);
|
|
bytes.set(bytesOfLength, offset);
|
|
offset += NUM_BYTES_STRING_LENGTH;
|
|
bytes.set(val, offset);
|
|
offset += val.length;
|
|
}
|
|
resolve(bytes);
|
|
});
|
|
dataPromises.push(utf8bytes);
|
|
} else {
|
|
dataPromises.push(t2.data());
|
|
}
|
|
if (group != null) {
|
|
spec.group = group;
|
|
}
|
|
specs.push(spec);
|
|
}
|
|
const tensorValues = await Promise.all(dataPromises);
|
|
return { data: concatenateTypedArrays(tensorValues), specs };
|
|
}
|
|
function decodeWeights(weightData, specs) {
|
|
const compositeBuffer = new CompositeArrayBuffer(weightData);
|
|
const out = {};
|
|
let offset = 0;
|
|
for (const spec of specs) {
|
|
const byteLength = getWeightBytelength(spec, (start, end) => {
|
|
return compositeBuffer.slice(offset + start, offset + end);
|
|
});
|
|
out[spec.name] = decodeWeight(spec, compositeBuffer.slice(offset, offset + byteLength));
|
|
offset += byteLength;
|
|
}
|
|
return out;
|
|
}
|
|
function getWeightBytelength(spec, slice2) {
|
|
const size = sizeFromShape(spec.shape);
|
|
let bytesPerValue;
|
|
if ("quantization" in spec) {
|
|
const quantization = spec.quantization;
|
|
bytesPerValue = DTYPE_VALUE_SIZE_MAP[quantization.dtype];
|
|
} else if (spec.dtype === "string") {
|
|
let byteLength = 0;
|
|
for (let i = 0; i < size; i++) {
|
|
byteLength += NUM_BYTES_STRING_LENGTH + new Uint32Array(slice2(byteLength, byteLength + NUM_BYTES_STRING_LENGTH))[0];
|
|
}
|
|
return byteLength;
|
|
} else {
|
|
bytesPerValue = DTYPE_VALUE_SIZE_MAP[spec.dtype];
|
|
}
|
|
return size * bytesPerValue;
|
|
}
|
|
async function getWeightBytelengthAsync(spec, slice2) {
|
|
const size = sizeFromShape(spec.shape);
|
|
let bytesPerValue;
|
|
if ("quantization" in spec) {
|
|
const quantization = spec.quantization;
|
|
bytesPerValue = DTYPE_VALUE_SIZE_MAP[quantization.dtype];
|
|
} else if (spec.dtype === "string") {
|
|
let byteLength = 0;
|
|
for (let i = 0; i < size; i++) {
|
|
byteLength += NUM_BYTES_STRING_LENGTH + new Uint32Array(await slice2(byteLength, byteLength + NUM_BYTES_STRING_LENGTH))[0];
|
|
}
|
|
return byteLength;
|
|
} else {
|
|
bytesPerValue = DTYPE_VALUE_SIZE_MAP[spec.dtype];
|
|
}
|
|
return size * bytesPerValue;
|
|
}
|
|
function decodeWeight(spec, byteBuffer) {
|
|
const name = spec.name;
|
|
const dtype = spec.dtype;
|
|
const shape = spec.shape;
|
|
const size = sizeFromShape(shape);
|
|
let values;
|
|
let offset = 0;
|
|
if ("quantization" in spec) {
|
|
const quantization = spec.quantization;
|
|
if (quantization.dtype === "uint8" || quantization.dtype === "uint16") {
|
|
if (!("min" in quantization && "scale" in quantization)) {
|
|
throw new Error(`Weight ${spec.name} with quantization ${quantization.dtype} doesn't have corresponding metadata min and scale.`);
|
|
}
|
|
} else if (quantization.dtype === "float16") {
|
|
if (dtype !== "float32") {
|
|
throw new Error(`Weight ${spec.name} is quantized with ${quantization.dtype} which only supports weights of type float32 not ${dtype}.`);
|
|
}
|
|
} else {
|
|
throw new Error(`Weight ${spec.name} has unknown quantization dtype ${quantization.dtype}. Supported quantization dtypes are: 'uint8', 'uint16', and 'float16'.`);
|
|
}
|
|
const quantizationSizeFactor = DTYPE_VALUE_SIZE_MAP[quantization.dtype];
|
|
const quantizedArray = quantization.dtype === "uint8" ? new Uint8Array(byteBuffer) : new Uint16Array(byteBuffer);
|
|
if (dtype === "float32") {
|
|
if (quantization.dtype === "uint8" || quantization.dtype === "uint16") {
|
|
values = new Float32Array(quantizedArray.length);
|
|
for (let i = 0; i < quantizedArray.length; i++) {
|
|
const v = quantizedArray[i];
|
|
values[i] = v * quantization.scale + quantization.min;
|
|
}
|
|
} else if (quantization.dtype === "float16") {
|
|
const float16Decode = getFloat16Decoder();
|
|
values = float16Decode(quantizedArray);
|
|
} else {
|
|
throw new Error(`Unsupported quantization type ${quantization.dtype} for weight type float32.`);
|
|
}
|
|
} else if (dtype === "int32") {
|
|
if (quantization.dtype !== "uint8" && quantization.dtype !== "uint16") {
|
|
throw new Error(`Unsupported quantization type ${quantization.dtype} for weight type int32.`);
|
|
}
|
|
values = new Int32Array(quantizedArray.length);
|
|
for (let i = 0; i < quantizedArray.length; i++) {
|
|
const v = quantizedArray[i];
|
|
values[i] = Math.round(v * quantization.scale + quantization.min);
|
|
}
|
|
} else {
|
|
throw new Error(`Unsupported dtype in weight '${name}': ${dtype}`);
|
|
}
|
|
offset += size * quantizationSizeFactor;
|
|
} else if (dtype === "string") {
|
|
const size2 = sizeFromShape(spec.shape);
|
|
values = [];
|
|
for (let i = 0; i < size2; i++) {
|
|
const byteLength = new Uint32Array(byteBuffer.slice(offset, offset + NUM_BYTES_STRING_LENGTH))[0];
|
|
offset += NUM_BYTES_STRING_LENGTH;
|
|
const bytes = new Uint8Array(byteBuffer.slice(offset, offset + byteLength));
|
|
values.push(bytes);
|
|
offset += byteLength;
|
|
}
|
|
} else {
|
|
const dtypeFactor = DTYPE_VALUE_SIZE_MAP[dtype];
|
|
if (dtype === "float32") {
|
|
values = new Float32Array(byteBuffer);
|
|
} else if (dtype === "int32") {
|
|
values = new Int32Array(byteBuffer);
|
|
} else if (dtype === "bool") {
|
|
values = new Uint8Array(byteBuffer);
|
|
} else if (dtype === "complex64") {
|
|
values = new Float32Array(byteBuffer);
|
|
const real2 = new Float32Array(values.length / 2);
|
|
const image2 = new Float32Array(values.length / 2);
|
|
for (let i = 0; i < real2.length; i++) {
|
|
real2[i] = values[i * 2];
|
|
image2[i] = values[i * 2 + 1];
|
|
}
|
|
const realTensor = tensor(real2, shape, "float32");
|
|
const imageTensor = tensor(image2, shape, "float32");
|
|
const complexTensor = complex$2(realTensor, imageTensor);
|
|
realTensor.dispose();
|
|
imageTensor.dispose();
|
|
return complexTensor;
|
|
} else {
|
|
throw new Error(`Unsupported dtype in weight '${name}': ${dtype}`);
|
|
}
|
|
offset += size * dtypeFactor;
|
|
}
|
|
return tensor(values, shape, dtype);
|
|
}
|
|
async function readToLength(reader, initialData, length) {
|
|
let data = new Uint8Array(initialData);
|
|
while (data.byteLength < length) {
|
|
const { done, value } = await reader.read();
|
|
if (done && value == null) {
|
|
const missing = length - data.byteLength;
|
|
throw new Error(`Reader is done but ${missing} bytes are still expected`);
|
|
}
|
|
const newData = new Uint8Array(data.length + value.byteLength);
|
|
newData.set(data, 0);
|
|
newData.set(new Uint8Array(value), data.length);
|
|
data = newData;
|
|
}
|
|
return data.buffer;
|
|
}
|
|
async function decodeWeightsStream(weightStream, specs) {
|
|
const tensors = {};
|
|
const reader = weightStream.getReader();
|
|
let data = new ArrayBuffer(0);
|
|
for (const spec of specs) {
|
|
const byteLength = await getWeightBytelengthAsync(spec, async (start, end) => {
|
|
data = await readToLength(reader, data, end);
|
|
return data.slice(start, end);
|
|
});
|
|
data = await readToLength(reader, data, byteLength);
|
|
const tensorData = data.slice(0, byteLength);
|
|
data = data.slice(byteLength);
|
|
const weightTensor = decodeWeight(spec, tensorData);
|
|
tensors[spec.name] = weightTensor;
|
|
if (getBackend() === "webgpu") {
|
|
const b = backend();
|
|
if ("uploadToGPU" in b && sizeFromShape(weightTensor.shape) >= env().get("WEBGPU_CPU_HANDOFF_SIZE_THRESHOLD")) {
|
|
b.uploadToGPU(weightTensor.dataId);
|
|
}
|
|
}
|
|
}
|
|
return tensors;
|
|
}
|
|
function concatenateTypedArrays(xs) {
|
|
if (xs === null) {
|
|
throw new Error(`Invalid input value: ${JSON.stringify(xs)}`);
|
|
}
|
|
let totalByteLength = 0;
|
|
const normalizedXs = [];
|
|
xs.forEach((x) => {
|
|
totalByteLength += x.byteLength;
|
|
normalizedXs.push(x.byteLength === x.buffer.byteLength ? x : new x.constructor(x));
|
|
if (!(x instanceof Float32Array || x instanceof Int32Array || x instanceof Uint8Array)) {
|
|
throw new Error(`Unsupported TypedArray subtype: ${x.constructor.name}`);
|
|
}
|
|
});
|
|
const y = new Uint8Array(totalByteLength);
|
|
let offset = 0;
|
|
normalizedXs.forEach((x) => {
|
|
y.set(new Uint8Array(x.buffer), offset);
|
|
offset += x.byteLength;
|
|
});
|
|
return y.buffer;
|
|
}
|
|
const useNodeBuffer = typeof Buffer !== "undefined" && (typeof Blob === "undefined" || typeof atob === "undefined" || typeof btoa === "undefined");
|
|
function stringByteLength(str) {
|
|
if (useNodeBuffer) {
|
|
return Buffer.byteLength(str, "utf8");
|
|
}
|
|
return new Blob([str]).size;
|
|
}
|
|
function arrayBufferToBase64String(buffer2) {
|
|
if (useNodeBuffer) {
|
|
return Buffer.from(buffer2).toString("base64");
|
|
}
|
|
const buf = new Uint8Array(buffer2);
|
|
let s = "";
|
|
for (let i = 0, l = buf.length; i < l; i++) {
|
|
s += String.fromCharCode(buf[i]);
|
|
}
|
|
return btoa(s);
|
|
}
|
|
function base64StringToArrayBuffer(str) {
|
|
if (useNodeBuffer) {
|
|
const buf = Buffer.from(str, "base64");
|
|
return buf.buffer.slice(buf.byteOffset, buf.byteOffset + buf.byteLength);
|
|
}
|
|
const s = atob(str);
|
|
const buffer2 = new Uint8Array(s.length);
|
|
for (let i = 0; i < s.length; ++i) {
|
|
buffer2.set([s.charCodeAt(i)], i);
|
|
}
|
|
return buffer2.buffer;
|
|
}
|
|
function concatenateArrayBuffers(buffers) {
|
|
return CompositeArrayBuffer.join(buffers);
|
|
}
|
|
function basename(path) {
|
|
const SEPARATOR = "/";
|
|
path = path.trim();
|
|
while (path.endsWith(SEPARATOR)) {
|
|
path = path.slice(0, path.length - 1);
|
|
}
|
|
const items = path.split(SEPARATOR);
|
|
return items[items.length - 1];
|
|
}
|
|
function getModelJSONForModelArtifacts(artifacts, manifest) {
|
|
const result = {
|
|
modelTopology: artifacts.modelTopology,
|
|
format: artifacts.format,
|
|
generatedBy: artifacts.generatedBy,
|
|
convertedBy: artifacts.convertedBy,
|
|
weightsManifest: manifest
|
|
};
|
|
if (artifacts.signature != null) {
|
|
result.signature = artifacts.signature;
|
|
}
|
|
if (artifacts.userDefinedMetadata != null) {
|
|
result.userDefinedMetadata = artifacts.userDefinedMetadata;
|
|
}
|
|
if (artifacts.modelInitializer != null) {
|
|
result.modelInitializer = artifacts.modelInitializer;
|
|
}
|
|
if (artifacts.initializerSignature != null) {
|
|
result.initializerSignature = artifacts.initializerSignature;
|
|
}
|
|
if (artifacts.trainingConfig != null) {
|
|
result.trainingConfig = artifacts.trainingConfig;
|
|
}
|
|
return result;
|
|
}
|
|
function getModelArtifactsForJSONSync(modelJSON, weightSpecs, weightData) {
|
|
const modelArtifacts = {
|
|
modelTopology: modelJSON.modelTopology,
|
|
format: modelJSON.format,
|
|
generatedBy: modelJSON.generatedBy,
|
|
convertedBy: modelJSON.convertedBy
|
|
};
|
|
if (modelJSON.trainingConfig != null) {
|
|
modelArtifacts.trainingConfig = modelJSON.trainingConfig;
|
|
}
|
|
if (modelJSON.weightsManifest != null) {
|
|
if (!weightSpecs) {
|
|
throw new Error("modelJSON has weightsManifest but weightSpecs is null");
|
|
}
|
|
if (!weightData) {
|
|
throw new Error("modelJSON has weightsManifest but weightData is null");
|
|
}
|
|
modelArtifacts.weightSpecs = weightSpecs;
|
|
modelArtifacts.weightData = weightData;
|
|
}
|
|
if (modelJSON.signature != null) {
|
|
modelArtifacts.signature = modelJSON.signature;
|
|
}
|
|
if (modelJSON.userDefinedMetadata != null) {
|
|
modelArtifacts.userDefinedMetadata = modelJSON.userDefinedMetadata;
|
|
}
|
|
if (modelJSON.modelInitializer != null) {
|
|
modelArtifacts.modelInitializer = modelJSON.modelInitializer;
|
|
}
|
|
if (modelJSON.initializerSignature != null) {
|
|
modelArtifacts.initializerSignature = modelJSON.initializerSignature;
|
|
}
|
|
return modelArtifacts;
|
|
}
|
|
async function getModelArtifactsForJSON(modelJSON, loadWeights2) {
|
|
let weightSpecs;
|
|
let weightData;
|
|
if (modelJSON.weightsManifest != null) {
|
|
[weightSpecs, weightData] = await loadWeights2(modelJSON.weightsManifest);
|
|
}
|
|
return getModelArtifactsForJSONSync(modelJSON, weightSpecs, weightData);
|
|
}
|
|
function getModelArtifactsInfoForJSON(modelArtifacts) {
|
|
if (modelArtifacts.modelTopology instanceof ArrayBuffer) {
|
|
throw new Error("Expected JSON model topology, received ArrayBuffer.");
|
|
}
|
|
return {
|
|
dateSaved: /* @__PURE__ */ new Date(),
|
|
modelTopologyType: "JSON",
|
|
modelTopologyBytes: modelArtifacts.modelTopology == null ? 0 : stringByteLength(JSON.stringify(modelArtifacts.modelTopology)),
|
|
weightSpecsBytes: modelArtifacts.weightSpecs == null ? 0 : stringByteLength(JSON.stringify(modelArtifacts.weightSpecs)),
|
|
weightDataBytes: modelArtifacts.weightData == null ? 0 : new CompositeArrayBuffer(modelArtifacts.weightData).byteLength
|
|
};
|
|
}
|
|
function getWeightSpecs(weightsManifest) {
|
|
const weightSpecs = [];
|
|
for (const entry of weightsManifest) {
|
|
weightSpecs.push(...entry.weights);
|
|
}
|
|
return weightSpecs;
|
|
}
|
|
function computeFloat16MantisaTable() {
|
|
const convertMantissa = (i) => {
|
|
let m = i << 13;
|
|
let e = 0;
|
|
while ((m & 8388608) === 0) {
|
|
e -= 8388608;
|
|
m <<= 1;
|
|
}
|
|
m &= -8388609;
|
|
e += 947912704;
|
|
return m | e;
|
|
};
|
|
const mantisaTable = new Uint32Array(2048);
|
|
mantisaTable[0] = 0;
|
|
for (let i = 1; i < 1024; i++) {
|
|
mantisaTable[i] = convertMantissa(i);
|
|
}
|
|
for (let i = 1024; i < 2048; i++) {
|
|
mantisaTable[i] = 939524096 + (i - 1024 << 13);
|
|
}
|
|
return mantisaTable;
|
|
}
|
|
function computeFloat16ExponentTable() {
|
|
const exponentTable = new Uint32Array(64);
|
|
exponentTable[0] = 0;
|
|
exponentTable[31] = 1199570944;
|
|
exponentTable[32] = 2147483648;
|
|
exponentTable[63] = 3347054592;
|
|
for (let i = 1; i < 31; i++) {
|
|
exponentTable[i] = i << 23;
|
|
}
|
|
for (let i = 33; i < 63; i++) {
|
|
exponentTable[i] = 2147483648 + (i - 32 << 23);
|
|
}
|
|
return exponentTable;
|
|
}
|
|
function computeFloat16OffsetTable() {
|
|
const offsetTable = new Uint32Array(64);
|
|
for (let i = 0; i < 64; i++) {
|
|
offsetTable[i] = 1024;
|
|
}
|
|
offsetTable[0] = offsetTable[32] = 0;
|
|
return offsetTable;
|
|
}
|
|
function getFloat16Decoder() {
|
|
const mantisaTable = computeFloat16MantisaTable();
|
|
const exponentTable = computeFloat16ExponentTable();
|
|
const offsetTable = computeFloat16OffsetTable();
|
|
return (quantizedArray) => {
|
|
const buffer2 = new ArrayBuffer(4 * quantizedArray.length);
|
|
const bufferUint32View = new Uint32Array(buffer2);
|
|
for (let index = 0; index < quantizedArray.length; index++) {
|
|
const float16Bits = quantizedArray[index];
|
|
const float32Bits = mantisaTable[offsetTable[float16Bits >> 10] + (float16Bits & 1023)] + exponentTable[float16Bits >> 10];
|
|
bufferUint32View[index] = float32Bits;
|
|
}
|
|
return new Float32Array(buffer2);
|
|
};
|
|
}
|
|
class IORouterRegistry {
|
|
constructor() {
|
|
this.saveRouters = [];
|
|
this.loadRouters = [];
|
|
}
|
|
static getInstance() {
|
|
if (IORouterRegistry.instance == null) {
|
|
IORouterRegistry.instance = new IORouterRegistry();
|
|
}
|
|
return IORouterRegistry.instance;
|
|
}
|
|
/**
|
|
* Register a save-handler router.
|
|
*
|
|
* @param saveRouter A function that maps a URL-like string onto an instance
|
|
* of `IOHandler` with the `save` method defined or `null`.
|
|
*/
|
|
static registerSaveRouter(saveRouter) {
|
|
IORouterRegistry.getInstance().saveRouters.push(saveRouter);
|
|
}
|
|
/**
|
|
* Register a load-handler router.
|
|
*
|
|
* @param loadRouter A function that maps a URL-like string onto an instance
|
|
* of `IOHandler` with the `load` method defined or `null`.
|
|
*/
|
|
static registerLoadRouter(loadRouter) {
|
|
IORouterRegistry.getInstance().loadRouters.push(loadRouter);
|
|
}
|
|
/**
|
|
* Look up IOHandler for saving, given a URL-like string.
|
|
*
|
|
* @param url
|
|
* @returns If only one match is found, an instance of IOHandler with the
|
|
* `save` method defined. If no match is found, `null`.
|
|
* @throws Error, if more than one match is found.
|
|
*/
|
|
static getSaveHandlers(url) {
|
|
return IORouterRegistry.getHandlers(url, "save");
|
|
}
|
|
/**
|
|
* Look up IOHandler for loading, given a URL-like string.
|
|
*
|
|
* @param url
|
|
* @param loadOptions Optional, custom load options.
|
|
* @returns All valid handlers for `url`, given the currently registered
|
|
* handler routers.
|
|
*/
|
|
static getLoadHandlers(url, loadOptions) {
|
|
return IORouterRegistry.getHandlers(url, "load", loadOptions);
|
|
}
|
|
static getHandlers(url, handlerType, loadOptions) {
|
|
const validHandlers = [];
|
|
const routers = handlerType === "load" ? IORouterRegistry.getInstance().loadRouters : IORouterRegistry.getInstance().saveRouters;
|
|
routers.forEach((router) => {
|
|
const handler = router(url, loadOptions);
|
|
if (handler !== null) {
|
|
validHandlers.push(handler);
|
|
}
|
|
});
|
|
return validHandlers;
|
|
}
|
|
}
|
|
const registerSaveRouter = (loudRouter) => IORouterRegistry.registerSaveRouter(loudRouter);
|
|
const registerLoadRouter = (loudRouter) => IORouterRegistry.registerLoadRouter(loudRouter);
|
|
const getSaveHandlers = (url) => IORouterRegistry.getSaveHandlers(url);
|
|
const getLoadHandlers = (url, loadOptions) => IORouterRegistry.getLoadHandlers(url, loadOptions);
|
|
const DATABASE_NAME = "tensorflowjs";
|
|
const DATABASE_VERSION = 1;
|
|
const MODEL_STORE_NAME = "models_store";
|
|
const INFO_STORE_NAME = "model_info_store";
|
|
function getIndexedDBFactory() {
|
|
if (!env().getBool("IS_BROWSER")) {
|
|
throw new Error("Failed to obtain IndexedDB factory because the current environmentis not a web browser.");
|
|
}
|
|
const theWindow = typeof window === "undefined" ? self : window;
|
|
const factory = theWindow.indexedDB || theWindow.mozIndexedDB || theWindow.webkitIndexedDB || theWindow.msIndexedDB || theWindow.shimIndexedDB;
|
|
if (factory == null) {
|
|
throw new Error("The current browser does not appear to support IndexedDB.");
|
|
}
|
|
return factory;
|
|
}
|
|
function setUpDatabase(openRequest) {
|
|
const db = openRequest.result;
|
|
db.createObjectStore(MODEL_STORE_NAME, { keyPath: "modelPath" });
|
|
db.createObjectStore(INFO_STORE_NAME, { keyPath: "modelPath" });
|
|
}
|
|
class BrowserIndexedDB {
|
|
constructor(modelPath) {
|
|
this.indexedDB = getIndexedDBFactory();
|
|
if (modelPath == null || !modelPath) {
|
|
throw new Error("For IndexedDB, modelPath must not be null, undefined or empty.");
|
|
}
|
|
this.modelPath = modelPath;
|
|
}
|
|
async save(modelArtifacts) {
|
|
if (modelArtifacts.modelTopology instanceof ArrayBuffer) {
|
|
throw new Error("BrowserLocalStorage.save() does not support saving model topology in binary formats yet.");
|
|
}
|
|
return this.databaseAction(this.modelPath, modelArtifacts);
|
|
}
|
|
async load() {
|
|
return this.databaseAction(this.modelPath);
|
|
}
|
|
/**
|
|
* Perform database action to put model artifacts into or read model artifacts
|
|
* from IndexedDB object store.
|
|
*
|
|
* Whether the action is put or get depends on whether `modelArtifacts` is
|
|
* specified. If it is specified, the action will be put; otherwise the action
|
|
* will be get.
|
|
*
|
|
* @param modelPath A unique string path for the model.
|
|
* @param modelArtifacts If specified, it will be the model artifacts to be
|
|
* stored in IndexedDB.
|
|
* @returns A `Promise` of `SaveResult`, if the action is put, or a `Promise`
|
|
* of `ModelArtifacts`, if the action is get.
|
|
*/
|
|
databaseAction(modelPath, modelArtifacts) {
|
|
return new Promise((resolve, reject) => {
|
|
const openRequest = this.indexedDB.open(DATABASE_NAME, DATABASE_VERSION);
|
|
openRequest.onupgradeneeded = () => setUpDatabase(openRequest);
|
|
openRequest.onsuccess = () => {
|
|
const db = openRequest.result;
|
|
if (modelArtifacts == null) {
|
|
const modelTx = db.transaction(MODEL_STORE_NAME, "readonly");
|
|
const modelStore = modelTx.objectStore(MODEL_STORE_NAME);
|
|
const getRequest = modelStore.get(this.modelPath);
|
|
getRequest.onsuccess = () => {
|
|
if (getRequest.result == null) {
|
|
db.close();
|
|
return reject(new Error(`Cannot find model with path '${this.modelPath}' in IndexedDB.`));
|
|
} else {
|
|
resolve(getRequest.result.modelArtifacts);
|
|
}
|
|
};
|
|
getRequest.onerror = (error) => {
|
|
db.close();
|
|
return reject(getRequest.error);
|
|
};
|
|
modelTx.oncomplete = () => db.close();
|
|
} else {
|
|
modelArtifacts.weightData = CompositeArrayBuffer.join(modelArtifacts.weightData);
|
|
const modelArtifactsInfo = getModelArtifactsInfoForJSON(modelArtifacts);
|
|
const infoTx = db.transaction(INFO_STORE_NAME, "readwrite");
|
|
let infoStore = infoTx.objectStore(INFO_STORE_NAME);
|
|
let putInfoRequest;
|
|
try {
|
|
putInfoRequest = infoStore.put({ modelPath: this.modelPath, modelArtifactsInfo });
|
|
} catch (error) {
|
|
return reject(error);
|
|
}
|
|
let modelTx;
|
|
putInfoRequest.onsuccess = () => {
|
|
modelTx = db.transaction(MODEL_STORE_NAME, "readwrite");
|
|
const modelStore = modelTx.objectStore(MODEL_STORE_NAME);
|
|
let putModelRequest;
|
|
try {
|
|
putModelRequest = modelStore.put({
|
|
modelPath: this.modelPath,
|
|
modelArtifacts,
|
|
modelArtifactsInfo
|
|
});
|
|
} catch (error) {
|
|
return reject(error);
|
|
}
|
|
putModelRequest.onsuccess = () => resolve({ modelArtifactsInfo });
|
|
putModelRequest.onerror = (error) => {
|
|
infoStore = infoTx.objectStore(INFO_STORE_NAME);
|
|
const deleteInfoRequest = infoStore.delete(this.modelPath);
|
|
deleteInfoRequest.onsuccess = () => {
|
|
db.close();
|
|
return reject(putModelRequest.error);
|
|
};
|
|
deleteInfoRequest.onerror = (error2) => {
|
|
db.close();
|
|
return reject(putModelRequest.error);
|
|
};
|
|
};
|
|
};
|
|
putInfoRequest.onerror = (error) => {
|
|
db.close();
|
|
return reject(putInfoRequest.error);
|
|
};
|
|
infoTx.oncomplete = () => {
|
|
if (modelTx == null) {
|
|
db.close();
|
|
} else {
|
|
modelTx.oncomplete = () => db.close();
|
|
}
|
|
};
|
|
}
|
|
};
|
|
openRequest.onerror = (error) => reject(openRequest.error);
|
|
});
|
|
}
|
|
}
|
|
BrowserIndexedDB.URL_SCHEME = "indexeddb://";
|
|
const indexedDBRouter = (url) => {
|
|
if (!env().getBool("IS_BROWSER")) {
|
|
return null;
|
|
} else {
|
|
if (!Array.isArray(url) && url.startsWith(BrowserIndexedDB.URL_SCHEME)) {
|
|
return browserIndexedDB(url.slice(BrowserIndexedDB.URL_SCHEME.length));
|
|
} else {
|
|
return null;
|
|
}
|
|
}
|
|
};
|
|
IORouterRegistry.registerSaveRouter(indexedDBRouter);
|
|
IORouterRegistry.registerLoadRouter(indexedDBRouter);
|
|
function browserIndexedDB(modelPath) {
|
|
return new BrowserIndexedDB(modelPath);
|
|
}
|
|
function maybeStripScheme$1(key) {
|
|
return key.startsWith(BrowserIndexedDB.URL_SCHEME) ? key.slice(BrowserIndexedDB.URL_SCHEME.length) : key;
|
|
}
|
|
class BrowserIndexedDBManager {
|
|
constructor() {
|
|
this.indexedDB = getIndexedDBFactory();
|
|
}
|
|
async listModels() {
|
|
return new Promise((resolve, reject) => {
|
|
const openRequest = this.indexedDB.open(DATABASE_NAME, DATABASE_VERSION);
|
|
openRequest.onupgradeneeded = () => setUpDatabase(openRequest);
|
|
openRequest.onsuccess = () => {
|
|
const db = openRequest.result;
|
|
const tx = db.transaction(INFO_STORE_NAME, "readonly");
|
|
const store = tx.objectStore(INFO_STORE_NAME);
|
|
const getAllInfoRequest = store.getAll();
|
|
getAllInfoRequest.onsuccess = () => {
|
|
const out = {};
|
|
for (const item of getAllInfoRequest.result) {
|
|
out[item.modelPath] = item.modelArtifactsInfo;
|
|
}
|
|
resolve(out);
|
|
};
|
|
getAllInfoRequest.onerror = (error) => {
|
|
db.close();
|
|
return reject(getAllInfoRequest.error);
|
|
};
|
|
tx.oncomplete = () => db.close();
|
|
};
|
|
openRequest.onerror = (error) => reject(openRequest.error);
|
|
});
|
|
}
|
|
async removeModel(path) {
|
|
path = maybeStripScheme$1(path);
|
|
return new Promise((resolve, reject) => {
|
|
const openRequest = this.indexedDB.open(DATABASE_NAME, DATABASE_VERSION);
|
|
openRequest.onupgradeneeded = () => setUpDatabase(openRequest);
|
|
openRequest.onsuccess = () => {
|
|
const db = openRequest.result;
|
|
const infoTx = db.transaction(INFO_STORE_NAME, "readwrite");
|
|
const infoStore = infoTx.objectStore(INFO_STORE_NAME);
|
|
const getInfoRequest = infoStore.get(path);
|
|
let modelTx;
|
|
getInfoRequest.onsuccess = () => {
|
|
if (getInfoRequest.result == null) {
|
|
db.close();
|
|
return reject(new Error(`Cannot find model with path '${path}' in IndexedDB.`));
|
|
} else {
|
|
const deleteInfoRequest = infoStore.delete(path);
|
|
const deleteModelData = () => {
|
|
modelTx = db.transaction(MODEL_STORE_NAME, "readwrite");
|
|
const modelStore = modelTx.objectStore(MODEL_STORE_NAME);
|
|
const deleteModelRequest = modelStore.delete(path);
|
|
deleteModelRequest.onsuccess = () => resolve(getInfoRequest.result.modelArtifactsInfo);
|
|
deleteModelRequest.onerror = (error) => reject(getInfoRequest.error);
|
|
};
|
|
deleteInfoRequest.onsuccess = deleteModelData;
|
|
deleteInfoRequest.onerror = (error) => {
|
|
deleteModelData();
|
|
db.close();
|
|
return reject(getInfoRequest.error);
|
|
};
|
|
}
|
|
};
|
|
getInfoRequest.onerror = (error) => {
|
|
db.close();
|
|
return reject(getInfoRequest.error);
|
|
};
|
|
infoTx.oncomplete = () => {
|
|
if (modelTx == null) {
|
|
db.close();
|
|
} else {
|
|
modelTx.oncomplete = () => db.close();
|
|
}
|
|
};
|
|
};
|
|
openRequest.onerror = (error) => reject(openRequest.error);
|
|
});
|
|
}
|
|
}
|
|
const PATH_SEPARATOR = "/";
|
|
const PATH_PREFIX = "tensorflowjs_models";
|
|
const INFO_SUFFIX = "info";
|
|
const MODEL_TOPOLOGY_SUFFIX = "model_topology";
|
|
const WEIGHT_SPECS_SUFFIX = "weight_specs";
|
|
const WEIGHT_DATA_SUFFIX = "weight_data";
|
|
const MODEL_METADATA_SUFFIX = "model_metadata";
|
|
function getModelKeys(path) {
|
|
return {
|
|
info: [PATH_PREFIX, path, INFO_SUFFIX].join(PATH_SEPARATOR),
|
|
topology: [PATH_PREFIX, path, MODEL_TOPOLOGY_SUFFIX].join(PATH_SEPARATOR),
|
|
weightSpecs: [PATH_PREFIX, path, WEIGHT_SPECS_SUFFIX].join(PATH_SEPARATOR),
|
|
weightData: [PATH_PREFIX, path, WEIGHT_DATA_SUFFIX].join(PATH_SEPARATOR),
|
|
modelMetadata: [PATH_PREFIX, path, MODEL_METADATA_SUFFIX].join(PATH_SEPARATOR)
|
|
};
|
|
}
|
|
function removeItems(keys) {
|
|
for (const key of Object.values(keys)) {
|
|
window.localStorage.removeItem(key);
|
|
}
|
|
}
|
|
function getModelPathFromKey(key) {
|
|
const items = key.split(PATH_SEPARATOR);
|
|
if (items.length < 3) {
|
|
throw new Error(`Invalid key format: ${key}`);
|
|
}
|
|
return items.slice(1, items.length - 1).join(PATH_SEPARATOR);
|
|
}
|
|
function maybeStripScheme(key) {
|
|
return key.startsWith(BrowserLocalStorage.URL_SCHEME) ? key.slice(BrowserLocalStorage.URL_SCHEME.length) : key;
|
|
}
|
|
class BrowserLocalStorage {
|
|
constructor(modelPath) {
|
|
if (!env().getBool("IS_BROWSER") || typeof window === "undefined" || typeof window.localStorage === "undefined") {
|
|
throw new Error("The current environment does not support local storage.");
|
|
}
|
|
this.LS = window.localStorage;
|
|
if (modelPath == null || !modelPath) {
|
|
throw new Error("For local storage, modelPath must not be null, undefined or empty.");
|
|
}
|
|
this.modelPath = modelPath;
|
|
this.keys = getModelKeys(this.modelPath);
|
|
}
|
|
/**
|
|
* Save model artifacts to browser local storage.
|
|
*
|
|
* See the documentation to `browserLocalStorage` for details on the saved
|
|
* artifacts.
|
|
*
|
|
* @param modelArtifacts The model artifacts to be stored.
|
|
* @returns An instance of SaveResult.
|
|
*/
|
|
async save(modelArtifacts) {
|
|
if (modelArtifacts.modelTopology instanceof ArrayBuffer) {
|
|
throw new Error("BrowserLocalStorage.save() does not support saving model topology in binary formats yet.");
|
|
} else {
|
|
const topology = JSON.stringify(modelArtifacts.modelTopology);
|
|
const weightSpecs = JSON.stringify(modelArtifacts.weightSpecs);
|
|
const modelArtifactsInfo = getModelArtifactsInfoForJSON(modelArtifacts);
|
|
const weightBuffer = CompositeArrayBuffer.join(modelArtifacts.weightData);
|
|
try {
|
|
this.LS.setItem(this.keys.info, JSON.stringify(modelArtifactsInfo));
|
|
this.LS.setItem(this.keys.topology, topology);
|
|
this.LS.setItem(this.keys.weightSpecs, weightSpecs);
|
|
this.LS.setItem(this.keys.weightData, arrayBufferToBase64String(weightBuffer));
|
|
const metadata = {
|
|
format: modelArtifacts.format,
|
|
generatedBy: modelArtifacts.generatedBy,
|
|
convertedBy: modelArtifacts.convertedBy,
|
|
signature: modelArtifacts.signature != null ? modelArtifacts.signature : void 0,
|
|
userDefinedMetadata: modelArtifacts.userDefinedMetadata != null ? modelArtifacts.userDefinedMetadata : void 0,
|
|
modelInitializer: modelArtifacts.modelInitializer != null ? modelArtifacts.modelInitializer : void 0,
|
|
initializerSignature: modelArtifacts.initializerSignature != null ? modelArtifacts.initializerSignature : void 0,
|
|
trainingConfig: modelArtifacts.trainingConfig != null ? modelArtifacts.trainingConfig : void 0
|
|
};
|
|
this.LS.setItem(this.keys.modelMetadata, JSON.stringify(metadata));
|
|
return { modelArtifactsInfo };
|
|
} catch (err) {
|
|
removeItems(this.keys);
|
|
throw new Error(`Failed to save model '${this.modelPath}' to local storage: size quota being exceeded is a possible cause of this failure: modelTopologyBytes=${modelArtifactsInfo.modelTopologyBytes}, weightSpecsBytes=${modelArtifactsInfo.weightSpecsBytes}, weightDataBytes=${modelArtifactsInfo.weightDataBytes}.`);
|
|
}
|
|
}
|
|
}
|
|
/**
|
|
* Load a model from local storage.
|
|
*
|
|
* See the documentation to `browserLocalStorage` for details on the saved
|
|
* artifacts.
|
|
*
|
|
* @returns The loaded model (if loading succeeds).
|
|
*/
|
|
async load() {
|
|
const info = JSON.parse(this.LS.getItem(this.keys.info));
|
|
if (info == null) {
|
|
throw new Error(`In local storage, there is no model with name '${this.modelPath}'`);
|
|
}
|
|
if (info.modelTopologyType !== "JSON") {
|
|
throw new Error("BrowserLocalStorage does not support loading non-JSON model topology yet.");
|
|
}
|
|
const out = {};
|
|
const topology = JSON.parse(this.LS.getItem(this.keys.topology));
|
|
if (topology == null) {
|
|
throw new Error(`In local storage, the topology of model '${this.modelPath}' is missing.`);
|
|
}
|
|
out.modelTopology = topology;
|
|
const weightSpecs = JSON.parse(this.LS.getItem(this.keys.weightSpecs));
|
|
if (weightSpecs == null) {
|
|
throw new Error(`In local storage, the weight specs of model '${this.modelPath}' are missing.`);
|
|
}
|
|
out.weightSpecs = weightSpecs;
|
|
const metadataString = this.LS.getItem(this.keys.modelMetadata);
|
|
if (metadataString != null) {
|
|
const metadata = JSON.parse(metadataString);
|
|
out.format = metadata.format;
|
|
out.generatedBy = metadata.generatedBy;
|
|
out.convertedBy = metadata.convertedBy;
|
|
if (metadata.signature != null) {
|
|
out.signature = metadata.signature;
|
|
}
|
|
if (metadata.userDefinedMetadata != null) {
|
|
out.userDefinedMetadata = metadata.userDefinedMetadata;
|
|
}
|
|
if (metadata.modelInitializer != null) {
|
|
out.modelInitializer = metadata.modelInitializer;
|
|
}
|
|
if (metadata.initializerSignature != null) {
|
|
out.initializerSignature = metadata.initializerSignature;
|
|
}
|
|
if (metadata.trainingConfig != null) {
|
|
out.trainingConfig = metadata.trainingConfig;
|
|
}
|
|
}
|
|
const weightDataBase64 = this.LS.getItem(this.keys.weightData);
|
|
if (weightDataBase64 == null) {
|
|
throw new Error(`In local storage, the binary weight values of model '${this.modelPath}' are missing.`);
|
|
}
|
|
out.weightData = base64StringToArrayBuffer(weightDataBase64);
|
|
return out;
|
|
}
|
|
}
|
|
BrowserLocalStorage.URL_SCHEME = "localstorage://";
|
|
const localStorageRouter = (url) => {
|
|
if (!env().getBool("IS_BROWSER")) {
|
|
return null;
|
|
} else {
|
|
if (!Array.isArray(url) && url.startsWith(BrowserLocalStorage.URL_SCHEME)) {
|
|
return browserLocalStorage(url.slice(BrowserLocalStorage.URL_SCHEME.length));
|
|
} else {
|
|
return null;
|
|
}
|
|
}
|
|
};
|
|
IORouterRegistry.registerSaveRouter(localStorageRouter);
|
|
IORouterRegistry.registerLoadRouter(localStorageRouter);
|
|
function browserLocalStorage(modelPath) {
|
|
return new BrowserLocalStorage(modelPath);
|
|
}
|
|
class BrowserLocalStorageManager {
|
|
constructor() {
|
|
assert(env().getBool("IS_BROWSER"), () => "Current environment is not a web browser");
|
|
assert(typeof window === "undefined" || typeof window.localStorage !== "undefined", () => "Current browser does not appear to support localStorage");
|
|
this.LS = window.localStorage;
|
|
}
|
|
async listModels() {
|
|
const out = {};
|
|
const prefix = PATH_PREFIX + PATH_SEPARATOR;
|
|
const suffix = PATH_SEPARATOR + INFO_SUFFIX;
|
|
for (let i = 0; i < this.LS.length; ++i) {
|
|
const key = this.LS.key(i);
|
|
if (key.startsWith(prefix) && key.endsWith(suffix)) {
|
|
const modelPath = getModelPathFromKey(key);
|
|
out[modelPath] = JSON.parse(this.LS.getItem(key));
|
|
}
|
|
}
|
|
return out;
|
|
}
|
|
async removeModel(path) {
|
|
path = maybeStripScheme(path);
|
|
const keys = getModelKeys(path);
|
|
if (this.LS.getItem(keys.info) == null) {
|
|
throw new Error(`Cannot find model at path '${path}'`);
|
|
}
|
|
const info = JSON.parse(this.LS.getItem(keys.info));
|
|
removeItems(keys);
|
|
return info;
|
|
}
|
|
}
|
|
const URL_SCHEME_SUFFIX = "://";
|
|
class ModelStoreManagerRegistry {
|
|
constructor() {
|
|
this.managers = {};
|
|
}
|
|
static getInstance() {
|
|
if (ModelStoreManagerRegistry.instance == null) {
|
|
ModelStoreManagerRegistry.instance = new ModelStoreManagerRegistry();
|
|
}
|
|
return ModelStoreManagerRegistry.instance;
|
|
}
|
|
/**
|
|
* Register a save-handler router.
|
|
*
|
|
* @param saveRouter A function that maps a URL-like string onto an instance
|
|
* of `IOHandler` with the `save` method defined or `null`.
|
|
*/
|
|
static registerManager(scheme, manager) {
|
|
assert(scheme != null, () => "scheme must not be undefined or null.");
|
|
if (scheme.endsWith(URL_SCHEME_SUFFIX)) {
|
|
scheme = scheme.slice(0, scheme.indexOf(URL_SCHEME_SUFFIX));
|
|
}
|
|
assert(scheme.length > 0, () => "scheme must not be an empty string.");
|
|
const registry = ModelStoreManagerRegistry.getInstance();
|
|
assert(registry.managers[scheme] == null, () => `A model store manager is already registered for scheme '${scheme}'.`);
|
|
registry.managers[scheme] = manager;
|
|
}
|
|
static getManager(scheme) {
|
|
const manager = ModelStoreManagerRegistry.getInstance().managers[scheme];
|
|
if (manager == null) {
|
|
throw new Error(`Cannot find model manager for scheme '${scheme}'`);
|
|
}
|
|
return manager;
|
|
}
|
|
static getSchemes() {
|
|
return Object.keys(ModelStoreManagerRegistry.getInstance().managers);
|
|
}
|
|
}
|
|
function parseURL(url) {
|
|
if (url.indexOf(URL_SCHEME_SUFFIX) === -1) {
|
|
throw new Error(`The url string provided does not contain a scheme. Supported schemes are: ${ModelStoreManagerRegistry.getSchemes().join(",")}`);
|
|
}
|
|
return {
|
|
scheme: url.split(URL_SCHEME_SUFFIX)[0],
|
|
path: url.split(URL_SCHEME_SUFFIX)[1]
|
|
};
|
|
}
|
|
async function cloneModelInternal(sourceURL, destURL, deleteSource = false) {
|
|
assert(sourceURL !== destURL, () => `Old path and new path are the same: '${sourceURL}'`);
|
|
const loadHandlers = IORouterRegistry.getLoadHandlers(sourceURL);
|
|
assert(loadHandlers.length > 0, () => `Copying failed because no load handler is found for source URL ${sourceURL}.`);
|
|
assert(loadHandlers.length < 2, () => `Copying failed because more than one (${loadHandlers.length}) load handlers for source URL ${sourceURL}.`);
|
|
const loadHandler = loadHandlers[0];
|
|
const saveHandlers = IORouterRegistry.getSaveHandlers(destURL);
|
|
assert(saveHandlers.length > 0, () => `Copying failed because no save handler is found for destination URL ${destURL}.`);
|
|
assert(saveHandlers.length < 2, () => `Copying failed because more than one (${loadHandlers.length}) save handlers for destination URL ${destURL}.`);
|
|
const saveHandler = saveHandlers[0];
|
|
const sourceScheme = parseURL(sourceURL).scheme;
|
|
const sourcePath = parseURL(sourceURL).path;
|
|
const sameMedium = sourceScheme === parseURL(sourceURL).scheme;
|
|
const modelArtifacts = await loadHandler.load();
|
|
if (deleteSource && sameMedium) {
|
|
await ModelStoreManagerRegistry.getManager(sourceScheme).removeModel(sourcePath);
|
|
}
|
|
const saveResult = await saveHandler.save(modelArtifacts);
|
|
if (deleteSource && !sameMedium) {
|
|
await ModelStoreManagerRegistry.getManager(sourceScheme).removeModel(sourcePath);
|
|
}
|
|
return saveResult.modelArtifactsInfo;
|
|
}
|
|
async function listModels() {
|
|
const schemes = ModelStoreManagerRegistry.getSchemes();
|
|
const out = {};
|
|
for (const scheme of schemes) {
|
|
const schemeOut = await ModelStoreManagerRegistry.getManager(scheme).listModels();
|
|
for (const path in schemeOut) {
|
|
const url = scheme + URL_SCHEME_SUFFIX + path;
|
|
out[url] = schemeOut[path];
|
|
}
|
|
}
|
|
return out;
|
|
}
|
|
async function removeModel(url) {
|
|
const schemeAndPath = parseURL(url);
|
|
const manager = ModelStoreManagerRegistry.getManager(schemeAndPath.scheme);
|
|
return manager.removeModel(schemeAndPath.path);
|
|
}
|
|
async function copyModel(sourceURL, destURL) {
|
|
const deleteSource = false;
|
|
return cloneModelInternal(sourceURL, destURL, deleteSource);
|
|
}
|
|
async function moveModel(sourceURL, destURL) {
|
|
const deleteSource = true;
|
|
return cloneModelInternal(sourceURL, destURL, deleteSource);
|
|
}
|
|
class PlatformBrowser {
|
|
constructor() {
|
|
this.messageName = "setTimeoutCustom";
|
|
this.functionRefs = [];
|
|
this.handledMessageCount = 0;
|
|
this.hasEventListener = false;
|
|
}
|
|
fetch(path, init) {
|
|
return fetch(path, init);
|
|
}
|
|
now() {
|
|
return performance.now();
|
|
}
|
|
encode(text, encoding) {
|
|
if (encoding !== "utf-8" && encoding !== "utf8") {
|
|
throw new Error(`Browser's encoder only supports utf-8, but got ${encoding}`);
|
|
}
|
|
if (this.textEncoder == null) {
|
|
this.textEncoder = new TextEncoder();
|
|
}
|
|
return this.textEncoder.encode(text);
|
|
}
|
|
decode(bytes, encoding) {
|
|
return new TextDecoder(encoding).decode(bytes);
|
|
}
|
|
// If the setTimeout nesting level is greater than 5 and timeout is less
|
|
// than 4ms, timeout will be clamped to 4ms, which hurts the perf.
|
|
// Interleaving window.postMessage and setTimeout will trick the browser and
|
|
// avoid the clamp.
|
|
setTimeoutCustom(functionRef, delay) {
|
|
if (typeof window === "undefined" || !env().getBool("USE_SETTIMEOUTCUSTOM")) {
|
|
setTimeout(functionRef, delay);
|
|
return;
|
|
}
|
|
this.functionRefs.push(functionRef);
|
|
setTimeout(() => {
|
|
window.postMessage({ name: this.messageName, index: this.functionRefs.length - 1 }, "*");
|
|
}, delay);
|
|
if (!this.hasEventListener) {
|
|
this.hasEventListener = true;
|
|
window.addEventListener("message", (event) => {
|
|
if (event.source === window && event.data.name === this.messageName) {
|
|
event.stopPropagation();
|
|
const functionRef2 = this.functionRefs[event.data.index];
|
|
functionRef2();
|
|
this.handledMessageCount++;
|
|
if (this.handledMessageCount === this.functionRefs.length) {
|
|
this.functionRefs = [];
|
|
this.handledMessageCount = 0;
|
|
}
|
|
}
|
|
}, true);
|
|
}
|
|
}
|
|
isTypedArray(a) {
|
|
return isTypedArrayBrowser(a);
|
|
}
|
|
}
|
|
if (env().get("IS_BROWSER")) {
|
|
env().setPlatform("browser", new PlatformBrowser());
|
|
try {
|
|
ModelStoreManagerRegistry.registerManager(BrowserLocalStorage.URL_SCHEME, new BrowserLocalStorageManager());
|
|
} catch (err) {
|
|
}
|
|
try {
|
|
ModelStoreManagerRegistry.registerManager(BrowserIndexedDB.URL_SCHEME, new BrowserIndexedDBManager());
|
|
} catch (err) {
|
|
}
|
|
}
|
|
const getNodeFetch = {
|
|
// tslint:disable-next-line:no-require-imports
|
|
importFetch: () => require("node-fetch")
|
|
};
|
|
let systemFetch;
|
|
class PlatformNode {
|
|
constructor() {
|
|
this.util = require("util");
|
|
this.textEncoder = new this.util.TextEncoder();
|
|
}
|
|
fetch(path, requestInits) {
|
|
if (env().global.fetch != null) {
|
|
return env().global.fetch(path, requestInits);
|
|
}
|
|
if (systemFetch == null) {
|
|
systemFetch = getNodeFetch.importFetch();
|
|
}
|
|
return systemFetch(path, requestInits);
|
|
}
|
|
now() {
|
|
const time = process.hrtime();
|
|
return time[0] * 1e3 + time[1] / 1e6;
|
|
}
|
|
encode(text, encoding) {
|
|
if (encoding !== "utf-8" && encoding !== "utf8") {
|
|
throw new Error(`Node built-in encoder only supports utf-8, but got ${encoding}`);
|
|
}
|
|
return this.textEncoder.encode(text);
|
|
}
|
|
decode(bytes, encoding) {
|
|
if (bytes.length === 0) {
|
|
return "";
|
|
}
|
|
return new this.util.TextDecoder(encoding).decode(bytes);
|
|
}
|
|
isTypedArray(a) {
|
|
return this.util.types.isFloat32Array(a) || this.util.types.isInt32Array(a) || this.util.types.isUint8Array(a) || this.util.types.isUint8ClampedArray(a);
|
|
}
|
|
}
|
|
if (env().get("IS_NODE") && !env().get("IS_BROWSER")) {
|
|
env().setPlatform("node", new PlatformNode());
|
|
}
|
|
function buffer(shape, dtype = "float32", values) {
|
|
dtype = dtype || "float32";
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
return new TensorBuffer(shape, dtype, values);
|
|
}
|
|
function cast_(x, dtype) {
|
|
const $x = convertToTensor(x, "x", "cast");
|
|
if (!isValidDtype(dtype)) {
|
|
throw new Error(`Failed to cast to unknown dtype ${dtype}`);
|
|
}
|
|
if (dtype === "string" && $x.dtype !== "string" || dtype !== "string" && $x.dtype === "string") {
|
|
throw new Error("Only strings can be casted to strings");
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { dtype };
|
|
return ENGINE.runKernel(Cast, inputs, attrs);
|
|
}
|
|
const cast$2 = /* @__PURE__ */ op({ cast_ });
|
|
function clone_(x) {
|
|
const $x = convertToTensor(x, "x", "clone", "string_or_numeric");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Identity, inputs);
|
|
}
|
|
const clone = /* @__PURE__ */ op({ clone_ });
|
|
function print(x, verbose = false) {
|
|
console.log(x.toString(verbose));
|
|
}
|
|
getOrMakeEngine();
|
|
const opHandler = {
|
|
buffer,
|
|
cast: cast$2,
|
|
clone,
|
|
print
|
|
};
|
|
setOpHandler(opHandler);
|
|
function add_(a, b) {
|
|
let $a = convertToTensor(a, "a", "add");
|
|
let $b = convertToTensor(b, "b", "add");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Add, inputs);
|
|
}
|
|
const add$1 = /* @__PURE__ */ op({ add_ });
|
|
function floorDiv_(a, b) {
|
|
let $a = convertToTensor(a, "a", "floorDiv");
|
|
let $b = convertToTensor(b, "b", "floorDiv");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(FloorDiv, inputs);
|
|
}
|
|
const floorDiv$2 = /* @__PURE__ */ op({ floorDiv_ });
|
|
function div_(a, b) {
|
|
let $a = convertToTensor(a, "a", "div");
|
|
let $b = convertToTensor(b, "b", "div");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
if ($a.dtype === "int32" && $b.dtype === "int32") {
|
|
return floorDiv$2($a, $b);
|
|
}
|
|
const inputs = { a: $a, b: $b };
|
|
const attrs = {};
|
|
return ENGINE.runKernel(RealDiv, inputs, attrs);
|
|
}
|
|
const div$1 = /* @__PURE__ */ op({ div_ });
|
|
function mul_(a, b) {
|
|
let $a = convertToTensor(a, "a", "mul");
|
|
let $b = convertToTensor(b, "b", "mul");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Multiply, inputs);
|
|
}
|
|
const mul = /* @__PURE__ */ op({ mul_ });
|
|
function abs_(x) {
|
|
const $x = convertToTensor(x, "x", "abs");
|
|
if ($x.dtype === "complex64") {
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(ComplexAbs, inputs);
|
|
} else {
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Abs, inputs);
|
|
}
|
|
}
|
|
const abs$2 = /* @__PURE__ */ op({ abs_ });
|
|
function acos_(x) {
|
|
const $x = convertToTensor(x, "x", "acos");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Acos, inputs);
|
|
}
|
|
const acos$2 = /* @__PURE__ */ op({ acos_ });
|
|
function acosh_(x) {
|
|
const $x = convertToTensor(x, "x", "acosh");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Acosh, inputs);
|
|
}
|
|
const acosh$2 = /* @__PURE__ */ op({ acosh_ });
|
|
function addN_(tensors) {
|
|
assert(Array.isArray(tensors), () => "The argument passed to tf.addN() must be a list of tensors");
|
|
assert(tensors.length >= 1, () => `Must pass at least one tensor to tf.addN(), but got ${tensors.length}`);
|
|
const $tensors = tensors.map((t2, i) => convertToTensor(t2, `tensors${i}`, "addN"));
|
|
const firstTensor = $tensors[0];
|
|
$tensors.forEach((t2) => {
|
|
if (t2.dtype !== firstTensor.dtype) {
|
|
throw new Error("All tensors passed to tf.addN() must have the same dtype");
|
|
}
|
|
});
|
|
$tensors.forEach((t2) => {
|
|
if (!arraysEqual(t2.shape, firstTensor.shape)) {
|
|
throw new Error("All tensors passed to tf.addN() must have the same shape");
|
|
}
|
|
});
|
|
const inputs = $tensors;
|
|
return ENGINE.runKernel(AddN, inputs);
|
|
}
|
|
const addN$2 = /* @__PURE__ */ op({ addN_ });
|
|
function all_(x, axis = null, keepDims = false) {
|
|
const $x = convertToTensor(x, "x", "all", "bool");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, keepDims };
|
|
return ENGINE.runKernel(All, inputs, attrs);
|
|
}
|
|
const all$2 = /* @__PURE__ */ op({ all_ });
|
|
function any_(x, axis = null, keepDims = false) {
|
|
const $x = convertToTensor(x, "x", "any", "bool");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, keepDims };
|
|
return ENGINE.runKernel(Any, inputs, attrs);
|
|
}
|
|
const any$2 = /* @__PURE__ */ op({ any_ });
|
|
function argMax_(x, axis = 0) {
|
|
const $x = convertToTensor(x, "x", "argMax");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis };
|
|
return ENGINE.runKernel(ArgMax, inputs, attrs);
|
|
}
|
|
const argMax$2 = /* @__PURE__ */ op({ argMax_ });
|
|
function argMin_(x, axis = 0) {
|
|
const $x = convertToTensor(x, "x", "argMin");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis };
|
|
return ENGINE.runKernel(ArgMin, inputs, attrs);
|
|
}
|
|
const argMin$2 = /* @__PURE__ */ op({ argMin_ });
|
|
function asin_(x) {
|
|
const $x = convertToTensor(x, "x", "asin");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Asin, inputs);
|
|
}
|
|
const asin$2 = /* @__PURE__ */ op({ asin_ });
|
|
function asinh_(x) {
|
|
const $x = convertToTensor(x, "x", "asinh");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Asinh, inputs);
|
|
}
|
|
const asinh$2 = /* @__PURE__ */ op({ asinh_ });
|
|
function atan_(x) {
|
|
const $x = convertToTensor(x, "x", "atan");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Atan, inputs);
|
|
}
|
|
const atan$2 = /* @__PURE__ */ op({ atan_ });
|
|
function atan2_(a, b) {
|
|
let $a = convertToTensor(a, "a", "atan2");
|
|
let $b = convertToTensor(b, "b", "atan2");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Atan2, inputs);
|
|
}
|
|
const atan2$2 = /* @__PURE__ */ op({ atan2_ });
|
|
function atanh_(x) {
|
|
const $x = convertToTensor(x, "x", "atanh");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Atanh, inputs);
|
|
}
|
|
const atanh$2 = /* @__PURE__ */ op({ atanh_ });
|
|
function computeDilation2DInfo(inputShape, filterShape, strides, pad2, dataFormat = "NHWC", dilations) {
|
|
const inputChannels = inputShape[3];
|
|
const $filterShape = [...filterShape, inputChannels];
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
return computeConv2DInfo(inputShape, $filterShape, strides, dilations, pad2, null, null, $dataFormat);
|
|
}
|
|
function computePool2DInfo(inShape, filterSize, strides, dilations, pad2, roundingMode, dataFormat = "channelsLast") {
|
|
const [filterHeight, filterWidth] = parseTupleParam(filterSize);
|
|
let filterShape;
|
|
if (dataFormat === "channelsLast") {
|
|
filterShape = [filterHeight, filterWidth, inShape[3], inShape[3]];
|
|
} else if (dataFormat === "channelsFirst") {
|
|
filterShape = [filterHeight, filterWidth, inShape[1], inShape[1]];
|
|
} else {
|
|
throw new Error(`Unknown dataFormat ${dataFormat}`);
|
|
}
|
|
return computeConv2DInfo(inShape, filterShape, strides, dilations, pad2, roundingMode, false, dataFormat);
|
|
}
|
|
function computePool3DInfo(inShape, filterSize, strides, dilations, pad2, roundingMode, dataFormat = "NDHWC") {
|
|
const [filterDepth, filterHeight, filterWidth] = parse3TupleParam(filterSize);
|
|
let filterShape;
|
|
let $dataFormat;
|
|
if (dataFormat === "NDHWC") {
|
|
$dataFormat = "channelsLast";
|
|
filterShape = [filterDepth, filterHeight, filterWidth, inShape[4], inShape[4]];
|
|
} else if (dataFormat === "NCDHW") {
|
|
$dataFormat = "channelsFirst";
|
|
filterShape = [filterDepth, filterHeight, filterWidth, inShape[1], inShape[1]];
|
|
} else {
|
|
throw new Error(`Unknown dataFormat ${dataFormat}`);
|
|
}
|
|
return computeConv3DInfo(inShape, filterShape, strides, dilations, pad2, false, $dataFormat, roundingMode);
|
|
}
|
|
function computeConv2DInfo(inShape, filterShape, strides, dilations, pad2, roundingMode, depthwise = false, dataFormat = "channelsLast") {
|
|
let [batchSize, inHeight, inWidth, inChannels] = [-1, -1, -1, -1];
|
|
if (dataFormat === "channelsLast") {
|
|
[batchSize, inHeight, inWidth, inChannels] = inShape;
|
|
} else if (dataFormat === "channelsFirst") {
|
|
[batchSize, inChannels, inHeight, inWidth] = inShape;
|
|
} else {
|
|
throw new Error(`Unknown dataFormat ${dataFormat}`);
|
|
}
|
|
const [filterHeight, filterWidth, , filterChannels] = filterShape;
|
|
const [strideHeight, strideWidth] = parseTupleParam(strides);
|
|
const [dilationHeight, dilationWidth] = parseTupleParam(dilations);
|
|
const effectiveFilterHeight = getEffectiveFilterSize(filterHeight, dilationHeight);
|
|
const effectiveFilterWidth = getEffectiveFilterSize(filterWidth, dilationWidth);
|
|
const { padInfo, outHeight, outWidth } = getPadAndOutInfo(pad2, inHeight, inWidth, strideHeight, strideWidth, effectiveFilterHeight, effectiveFilterWidth, roundingMode, dataFormat);
|
|
const outChannels = depthwise ? filterChannels * inChannels : filterChannels;
|
|
let outShape;
|
|
if (dataFormat === "channelsFirst") {
|
|
outShape = [batchSize, outChannels, outHeight, outWidth];
|
|
} else if (dataFormat === "channelsLast") {
|
|
outShape = [batchSize, outHeight, outWidth, outChannels];
|
|
}
|
|
return {
|
|
batchSize,
|
|
dataFormat,
|
|
inHeight,
|
|
inWidth,
|
|
inChannels,
|
|
outHeight,
|
|
outWidth,
|
|
outChannels,
|
|
padInfo,
|
|
strideHeight,
|
|
strideWidth,
|
|
filterHeight,
|
|
filterWidth,
|
|
effectiveFilterHeight,
|
|
effectiveFilterWidth,
|
|
dilationHeight,
|
|
dilationWidth,
|
|
inShape,
|
|
outShape,
|
|
filterShape
|
|
};
|
|
}
|
|
function computeConv3DInfo(inShape, filterShape, strides, dilations, pad2, depthwise = false, dataFormat = "channelsLast", roundingMode) {
|
|
let [batchSize, inDepth, inHeight, inWidth, inChannels] = [-1, -1, -1, -1, -1];
|
|
if (dataFormat === "channelsLast") {
|
|
[batchSize, inDepth, inHeight, inWidth, inChannels] = inShape;
|
|
} else if (dataFormat === "channelsFirst") {
|
|
[batchSize, inChannels, inDepth, inHeight, inWidth] = inShape;
|
|
} else {
|
|
throw new Error(`Unknown dataFormat ${dataFormat}`);
|
|
}
|
|
const [filterDepth, filterHeight, filterWidth, , filterChannels] = filterShape;
|
|
const [strideDepth, strideHeight, strideWidth] = parse3TupleParam(strides);
|
|
const [dilationDepth, dilationHeight, dilationWidth] = parse3TupleParam(dilations);
|
|
const effectiveFilterDepth = getEffectiveFilterSize(filterDepth, dilationDepth);
|
|
const effectiveFilterHeight = getEffectiveFilterSize(filterHeight, dilationHeight);
|
|
const effectiveFilterWidth = getEffectiveFilterSize(filterWidth, dilationWidth);
|
|
const { padInfo, outDepth, outHeight, outWidth } = get3DPadAndOutInfo(pad2, inDepth, inHeight, inWidth, strideDepth, strideHeight, strideWidth, effectiveFilterDepth, effectiveFilterHeight, effectiveFilterWidth, roundingMode);
|
|
const outChannels = depthwise ? filterChannels * inChannels : filterChannels;
|
|
let outShape;
|
|
if (dataFormat === "channelsFirst") {
|
|
outShape = [batchSize, outChannels, outDepth, outHeight, outWidth];
|
|
} else if (dataFormat === "channelsLast") {
|
|
outShape = [batchSize, outDepth, outHeight, outWidth, outChannels];
|
|
}
|
|
return {
|
|
batchSize,
|
|
dataFormat,
|
|
inDepth,
|
|
inHeight,
|
|
inWidth,
|
|
inChannels,
|
|
outDepth,
|
|
outHeight,
|
|
outWidth,
|
|
outChannels,
|
|
padInfo,
|
|
strideDepth,
|
|
strideHeight,
|
|
strideWidth,
|
|
filterDepth,
|
|
filterHeight,
|
|
filterWidth,
|
|
effectiveFilterDepth,
|
|
effectiveFilterHeight,
|
|
effectiveFilterWidth,
|
|
dilationDepth,
|
|
dilationHeight,
|
|
dilationWidth,
|
|
inShape,
|
|
outShape,
|
|
filterShape
|
|
};
|
|
}
|
|
function computeOutputShape2D(inShape, fieldSize, stride, zeroPad, roundingMode) {
|
|
if (zeroPad == null) {
|
|
zeroPad = computeDefaultPad(inShape, fieldSize, stride);
|
|
}
|
|
const inputRows = inShape[0];
|
|
const inputCols = inShape[1];
|
|
const outputRows = round$3((inputRows - fieldSize + 2 * zeroPad) / stride + 1, roundingMode);
|
|
const outputCols = round$3((inputCols - fieldSize + 2 * zeroPad) / stride + 1, roundingMode);
|
|
return [outputRows, outputCols];
|
|
}
|
|
function computeOutputShape4D(inShape, filterShape, outChannels, strides, zeroPad, roundingMode) {
|
|
if (zeroPad == null) {
|
|
zeroPad = computeDefaultPad(inShape, filterShape[0], strides[0]);
|
|
}
|
|
const outShape = [0, 0, 0, outChannels];
|
|
for (let index = 0; index < 3; index++) {
|
|
if (inShape[index] + 2 * zeroPad >= filterShape[index]) {
|
|
outShape[index] = round$3((inShape[index] - filterShape[index] + 2 * zeroPad) / strides[index] + 1, roundingMode);
|
|
}
|
|
}
|
|
return outShape;
|
|
}
|
|
function computeDefaultPad(inputShape, fieldSize, stride, dilation = 1) {
|
|
const effectiveFieldSize = getEffectiveFilterSize(fieldSize, dilation);
|
|
return Math.floor((inputShape[0] * (stride - 1) - stride + effectiveFieldSize) / 2);
|
|
}
|
|
function parseTupleParam(param) {
|
|
if (typeof param === "number") {
|
|
return [param, param, param];
|
|
}
|
|
if (param.length === 2) {
|
|
return [param[0], param[1], 1];
|
|
}
|
|
return param;
|
|
}
|
|
function parse3TupleParam(param) {
|
|
return typeof param === "number" ? [param, param, param] : param;
|
|
}
|
|
function getEffectiveFilterSize(filterSize, dilation) {
|
|
if (dilation <= 1) {
|
|
return filterSize;
|
|
}
|
|
return filterSize + (filterSize - 1) * (dilation - 1);
|
|
}
|
|
function getPadAndOutInfo(pad2, inHeight, inWidth, strideHeight, strideWidth, filterHeight, filterWidth, roundingMode, dataFormat) {
|
|
let padInfo;
|
|
let outHeight;
|
|
let outWidth;
|
|
if (typeof pad2 === "number") {
|
|
const padType = pad2 === 0 ? "VALID" : "NUMBER";
|
|
padInfo = { top: pad2, bottom: pad2, left: pad2, right: pad2, type: padType };
|
|
const outShape = computeOutputShape2D([inHeight, inWidth], filterHeight, strideHeight, pad2, roundingMode);
|
|
outHeight = outShape[0];
|
|
outWidth = outShape[1];
|
|
} else if (pad2 === "same") {
|
|
outHeight = Math.ceil(inHeight / strideHeight);
|
|
outWidth = Math.ceil(inWidth / strideWidth);
|
|
const padAlongHeight = Math.max(0, (outHeight - 1) * strideHeight + filterHeight - inHeight);
|
|
const padAlongWidth = Math.max(0, (outWidth - 1) * strideWidth + filterWidth - inWidth);
|
|
const top = Math.floor(padAlongHeight / 2);
|
|
const bottom = padAlongHeight - top;
|
|
const left = Math.floor(padAlongWidth / 2);
|
|
const right = padAlongWidth - left;
|
|
padInfo = { top, bottom, left, right, type: "SAME" };
|
|
} else if (pad2 === "valid") {
|
|
padInfo = { top: 0, bottom: 0, left: 0, right: 0, type: "VALID" };
|
|
outHeight = Math.ceil((inHeight - filterHeight + 1) / strideHeight);
|
|
outWidth = Math.ceil((inWidth - filterWidth + 1) / strideWidth);
|
|
} else if (typeof pad2 === "object") {
|
|
const top = dataFormat === "channelsLast" ? pad2[1][0] : pad2[2][0];
|
|
const bottom = dataFormat === "channelsLast" ? pad2[1][1] : pad2[2][1];
|
|
const left = dataFormat === "channelsLast" ? pad2[2][0] : pad2[3][0];
|
|
const right = dataFormat === "channelsLast" ? pad2[2][1] : pad2[3][1];
|
|
const padType = top === 0 && bottom === 0 && left === 0 && right === 0 ? "VALID" : "EXPLICIT";
|
|
padInfo = { top, bottom, left, right, type: padType };
|
|
outHeight = round$3((inHeight - filterHeight + top + bottom) / strideHeight + 1, roundingMode);
|
|
outWidth = round$3((inWidth - filterWidth + left + right) / strideWidth + 1, roundingMode);
|
|
} else {
|
|
throw Error(`Unknown padding parameter: ${pad2}`);
|
|
}
|
|
return { padInfo, outHeight, outWidth };
|
|
}
|
|
function get3DPadAndOutInfo(pad2, inDepth, inHeight, inWidth, strideDepth, strideHeight, strideWidth, filterDepth, filterHeight, filterWidth, roundingMode) {
|
|
let padInfo;
|
|
let outDepth;
|
|
let outHeight;
|
|
let outWidth;
|
|
if (pad2 === "valid") {
|
|
pad2 = 0;
|
|
}
|
|
if (typeof pad2 === "number") {
|
|
const padType = pad2 === 0 ? "VALID" : "NUMBER";
|
|
padInfo = {
|
|
top: pad2,
|
|
bottom: pad2,
|
|
left: pad2,
|
|
right: pad2,
|
|
front: pad2,
|
|
back: pad2,
|
|
type: padType
|
|
};
|
|
const outShape = computeOutputShape4D([inDepth, inHeight, inWidth, 1], [filterDepth, filterHeight, filterWidth], 1, [strideDepth, strideHeight, strideWidth], pad2, roundingMode);
|
|
outDepth = outShape[0];
|
|
outHeight = outShape[1];
|
|
outWidth = outShape[2];
|
|
} else if (pad2 === "same") {
|
|
outDepth = Math.ceil(inDepth / strideDepth);
|
|
outHeight = Math.ceil(inHeight / strideHeight);
|
|
outWidth = Math.ceil(inWidth / strideWidth);
|
|
const padAlongDepth = (outDepth - 1) * strideDepth + filterDepth - inDepth;
|
|
const padAlongHeight = (outHeight - 1) * strideHeight + filterHeight - inHeight;
|
|
const padAlongWidth = (outWidth - 1) * strideWidth + filterWidth - inWidth;
|
|
const front = Math.floor(padAlongDepth / 2);
|
|
const back = padAlongDepth - front;
|
|
const top = Math.floor(padAlongHeight / 2);
|
|
const bottom = padAlongHeight - top;
|
|
const left = Math.floor(padAlongWidth / 2);
|
|
const right = padAlongWidth - left;
|
|
padInfo = { top, bottom, left, right, front, back, type: "SAME" };
|
|
} else {
|
|
throw Error(`Unknown padding parameter: ${pad2}`);
|
|
}
|
|
return { padInfo, outDepth, outHeight, outWidth };
|
|
}
|
|
function round$3(value, roundingMode) {
|
|
if (!roundingMode) {
|
|
return Math.trunc(value);
|
|
}
|
|
switch (roundingMode) {
|
|
case "round":
|
|
return Math.round(value);
|
|
case "ceil":
|
|
return Math.ceil(value);
|
|
case "floor":
|
|
return Math.floor(value);
|
|
default:
|
|
throw new Error(`Unknown roundingMode ${roundingMode}`);
|
|
}
|
|
}
|
|
function tupleValuesAreOne(param) {
|
|
const [dimA, dimB, dimC] = parseTupleParam(param);
|
|
return dimA === 1 && dimB === 1 && dimC === 1;
|
|
}
|
|
function eitherStridesOrDilationsAreOne(strides, dilations) {
|
|
return tupleValuesAreOne(strides) || tupleValuesAreOne(dilations);
|
|
}
|
|
function stridesOrDilationsArePositive(values) {
|
|
return parseTupleParam(values).every((value) => value > 0);
|
|
}
|
|
function convertConv2DDataFormat(dataFormat) {
|
|
if (dataFormat === "NHWC") {
|
|
return "channelsLast";
|
|
} else if (dataFormat === "NCHW") {
|
|
return "channelsFirst";
|
|
} else {
|
|
throw new Error(`Unknown dataFormat ${dataFormat}`);
|
|
}
|
|
}
|
|
function checkPadOnDimRoundingMode(opDesc, pad2, dimRoundingMode) {
|
|
if (dimRoundingMode != null) {
|
|
if (typeof pad2 === "string") {
|
|
throw Error(`Error in ${opDesc}: pad must be an integer when using dimRoundingMode ${dimRoundingMode} but got pad ${pad2}.`);
|
|
} else if (typeof pad2 === "number") {
|
|
assert(isInt(pad2), () => `Error in ${opDesc}: pad must be an integer when using dimRoundingMode ${dimRoundingMode} but got pad ${pad2}.`);
|
|
} else if (typeof pad2 === "object") {
|
|
pad2.forEach((p2) => {
|
|
p2.forEach((v) => {
|
|
assert(isInt(v), () => `Error in ${opDesc}: pad must be an integer when using dimRoundingMode ${dimRoundingMode} but got pad ${v}.`);
|
|
});
|
|
});
|
|
} else {
|
|
throw Error(`Error in ${opDesc}: Unknown padding parameter: ${pad2}`);
|
|
}
|
|
}
|
|
}
|
|
function reshape_(x, shape) {
|
|
const $x = convertToTensor(x, "x", "reshape", "string_or_numeric");
|
|
const inputs = { x: $x };
|
|
const attrs = { shape };
|
|
return ENGINE.runKernel(Reshape, inputs, attrs);
|
|
}
|
|
const reshape$2 = /* @__PURE__ */ op({ reshape_ });
|
|
function avgPool_(x, filterSize, strides, pad2, dimRoundingMode) {
|
|
const $x = convertToTensor(x, "x", "avgPool", "float32");
|
|
const dilations = 1;
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in avgPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in avgPool: x must be rank 4 but got rank ${x4D.rank}.`);
|
|
checkPadOnDimRoundingMode("avgPool", pad2, dimRoundingMode);
|
|
const inputs = { x: x4D };
|
|
const attrs = { filterSize, strides, pad: pad2, dimRoundingMode };
|
|
let res = ENGINE.runKernel(AvgPool, inputs, attrs);
|
|
res = cast$2(res, $x.dtype);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const avgPool$2 = /* @__PURE__ */ op({ avgPool_ });
|
|
function avgPool3d_(x, filterSize, strides, pad2, dimRoundingMode, dataFormat = "NDHWC") {
|
|
const $x = convertToTensor(x, "x", "avgPool3d", "float32");
|
|
let x5D = $x;
|
|
let reshapedTo5D = false;
|
|
if ($x.rank === 4) {
|
|
reshapedTo5D = true;
|
|
x5D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2], $x.shape[3]]);
|
|
}
|
|
assert(x5D.rank === 5, () => `Error in avgPool3d: x must be rank 5 but got rank ${x5D.rank}.`);
|
|
assert(dataFormat === "NDHWC", () => `Error in avgPool3d: Only NDHWC is currently supported, but got dataFormat of ${dataFormat}`);
|
|
assert(typeof strides === "number" && strides > 0 || Array.isArray(strides) && strides[0] > 0 && strides[1] > 0 && strides[2] > 0, () => `Error in avgPool3d: Stride must be > 0, but got '${strides}'`);
|
|
checkPadOnDimRoundingMode("avgPool3d", pad2, dimRoundingMode);
|
|
const inputs = { x: x5D };
|
|
const attrs = { filterSize, strides, pad: pad2, dimRoundingMode, dataFormat };
|
|
let res = ENGINE.runKernel(AvgPool3D, inputs, attrs);
|
|
res = cast$2(res, x5D.dtype);
|
|
if (reshapedTo5D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3], res.shape[4]]);
|
|
}
|
|
return res;
|
|
}
|
|
const avgPool3d = /* @__PURE__ */ op({ avgPool3d_ });
|
|
function concat_(tensors, axis = 0) {
|
|
assert(tensors.length >= 1, () => "Pass at least one tensor to concat");
|
|
const $tensors = convertToTensorArray(tensors, "tensors", "concat", "string_or_numeric");
|
|
if ($tensors[0].dtype === "complex64") {
|
|
$tensors.forEach((tensor2) => {
|
|
if (tensor2.dtype !== "complex64") {
|
|
throw new Error(`Cannot concatenate complex64 tensors with a tensor
|
|
with dtype ${tensor2.dtype}. `);
|
|
}
|
|
});
|
|
}
|
|
if ($tensors.length === 1) {
|
|
return clone($tensors[0]);
|
|
}
|
|
const inputs = $tensors;
|
|
const attr = { axis };
|
|
return ENGINE.runKernel(Concat, inputs, attr);
|
|
}
|
|
const concat$2 = /* @__PURE__ */ op({ concat_ });
|
|
function matMul_(a, b, transposeA = false, transposeB = false) {
|
|
let $a = convertToTensor(a, "a", "matMul");
|
|
let $b = convertToTensor(b, "b", "matMul");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
const attrs = { transposeA, transposeB };
|
|
return ENGINE.runKernel(BatchMatMul, inputs, attrs);
|
|
}
|
|
const matMul$1 = /* @__PURE__ */ op({ matMul_ });
|
|
function sigmoid_(x) {
|
|
const $x = convertToTensor(x, "x", "sigmoid", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Sigmoid, inputs);
|
|
}
|
|
const sigmoid$2 = /* @__PURE__ */ op({ sigmoid_ });
|
|
function slice_(x, begin, size) {
|
|
const $x = convertToTensor(x, "x", "slice", "string_or_numeric");
|
|
if ($x.rank === 0) {
|
|
throw new Error("Slicing scalar is not possible");
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { begin, size };
|
|
return ENGINE.runKernel(Slice, inputs, attrs);
|
|
}
|
|
const slice$2 = /* @__PURE__ */ op({ slice_ });
|
|
function tanh_(x) {
|
|
const $x = convertToTensor(x, "x", "tanh", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Tanh, inputs);
|
|
}
|
|
const tanh$2 = /* @__PURE__ */ op({ tanh_ });
|
|
function basicLSTMCell_(forgetBias, lstmKernel, lstmBias, data, c, h) {
|
|
const $forgetBias = convertToTensor(forgetBias, "forgetBias", "basicLSTMCell");
|
|
const $lstmKernel = convertToTensor(lstmKernel, "lstmKernel", "basicLSTMCell");
|
|
const $lstmBias = convertToTensor(lstmBias, "lstmBias", "basicLSTMCell");
|
|
const $data = convertToTensor(data, "data", "basicLSTMCell");
|
|
const $c = convertToTensor(c, "c", "basicLSTMCell");
|
|
const $h = convertToTensor(h, "h", "basicLSTMCell");
|
|
const combined = concat$2([$data, $h], 1);
|
|
const weighted = matMul$1(combined, $lstmKernel);
|
|
const res = add$1(weighted, $lstmBias);
|
|
const batchSize = res.shape[0];
|
|
const sliceCols = res.shape[1] / 4;
|
|
const sliceSize = [batchSize, sliceCols];
|
|
const i = slice$2(res, [0, 0], sliceSize);
|
|
const j2 = slice$2(res, [0, sliceCols], sliceSize);
|
|
const f = slice$2(res, [0, sliceCols * 2], sliceSize);
|
|
const o = slice$2(res, [0, sliceCols * 3], sliceSize);
|
|
const newC = add$1(mul(sigmoid$2(i), tanh$2(j2)), mul($c, sigmoid$2(add$1($forgetBias, f))));
|
|
const newH = mul(tanh$2(newC), sigmoid$2(o));
|
|
return [newC, newH];
|
|
}
|
|
const basicLSTMCell = /* @__PURE__ */ op({ basicLSTMCell_ });
|
|
function batchToSpaceND_(x, blockShape, crops) {
|
|
const $x = convertToTensor(x, "x", "batchToSpaceND");
|
|
const prod2 = blockShape.reduce((a, b) => a * b);
|
|
assert($x.rank >= 1 + blockShape.length, () => `input rank is ${$x.rank} but should be > than blockShape.length ${blockShape.length}`);
|
|
assert(crops.length === blockShape.length, () => `crops.length is ${crops.length} but should be equal to blockShape.length ${blockShape.length}`);
|
|
assert($x.shape[0] % prod2 === 0, () => `input tensor batch is ${$x.shape[0]} but is not divisible by the product of the elements of blockShape ${blockShape.join(" * ")} === ${prod2}`);
|
|
const inputs = { x: $x };
|
|
const attrs = { blockShape, crops };
|
|
return ENGINE.runKernel(BatchToSpaceND, inputs, attrs);
|
|
}
|
|
const batchToSpaceND$2 = /* @__PURE__ */ op({ batchToSpaceND_ });
|
|
function xAs4D(x) {
|
|
let x4D;
|
|
if (x.rank === 0 || x.rank === 1) {
|
|
x4D = reshape$2(x, [1, 1, 1, x.size]);
|
|
} else if (x.rank === 2) {
|
|
x4D = reshape$2(x, [1, 1, x.shape[0], x.shape[1]]);
|
|
} else if (x.rank === 3) {
|
|
x4D = reshape$2(x, [1, x.shape[0], x.shape[1], x.shape[2]]);
|
|
} else {
|
|
x4D = x;
|
|
}
|
|
return x4D;
|
|
}
|
|
function batchNorm_(x, mean2, variance, offset, scale2, varianceEpsilon) {
|
|
if (varianceEpsilon == null) {
|
|
varianceEpsilon = 1e-3;
|
|
}
|
|
const $x = convertToTensor(x, "x", "batchNorm");
|
|
const $mean = convertToTensor(mean2, "mean", "batchNorm");
|
|
const $variance = convertToTensor(variance, "variance", "batchNorm");
|
|
let $scale;
|
|
if (scale2 != null) {
|
|
$scale = convertToTensor(scale2, "scale", "batchNorm");
|
|
}
|
|
let $offset;
|
|
if (offset != null) {
|
|
$offset = convertToTensor(offset, "offset", "batchNorm");
|
|
}
|
|
assert($mean.rank === $variance.rank, () => "Batch normalization gradient requires mean and variance to have equal ranks.");
|
|
assert($offset == null || $mean.rank === $offset.rank, () => "Batch normalization gradient requires mean and offset to have equal ranks.");
|
|
assert($scale == null || $mean.rank === $scale.rank, () => "Batch normalization gradient requires mean and scale to have equal ranks.");
|
|
const x4D = xAs4D($x);
|
|
const inputs = {
|
|
x: x4D,
|
|
scale: $scale,
|
|
offset: $offset,
|
|
mean: $mean,
|
|
variance: $variance
|
|
};
|
|
const attrs = { varianceEpsilon };
|
|
const res = ENGINE.runKernel(FusedBatchNorm, inputs, attrs);
|
|
return reshape$2(res, $x.shape);
|
|
}
|
|
const batchNorm$2 = /* @__PURE__ */ op({ batchNorm_ });
|
|
function batchNorm2d_(x, mean2, variance, offset, scale2, varianceEpsilon) {
|
|
const $x = convertToTensor(x, "x", "batchNorm");
|
|
const $mean = convertToTensor(mean2, "mean", "batchNorm");
|
|
const $variance = convertToTensor(variance, "variance", "batchNorm");
|
|
let $scale;
|
|
if (scale2 != null) {
|
|
$scale = convertToTensor(scale2, "scale", "batchNorm");
|
|
}
|
|
let $offset;
|
|
if (offset != null) {
|
|
$offset = convertToTensor(offset, "offset", "batchNorm");
|
|
}
|
|
assert($x.rank === 2, () => `Error in batchNorm2D: x must be rank 2 but got rank ${$x.rank}.`);
|
|
assert($mean.rank === 2 || $mean.rank === 1, () => `Error in batchNorm2D: mean must be rank 2 or rank 1 but got rank ${$mean.rank}.`);
|
|
assert($variance.rank === 2 || $variance.rank === 1, () => `Error in batchNorm2D: variance must be rank 2 or rank 1 but got rank ${$variance.rank}.`);
|
|
if ($scale != null) {
|
|
assert($scale.rank === 2 || $scale.rank === 1, () => `Error in batchNorm2D: scale must be rank 2 or rank 1 but got rank ${$scale.rank}.`);
|
|
}
|
|
if ($offset != null) {
|
|
assert($offset.rank === 2 || $offset.rank === 1, () => `Error in batchNorm2D: offset must be rank 2 or rank 1 but got rank ${$offset.rank}.`);
|
|
}
|
|
return batchNorm$2($x, $mean, $variance, $offset, $scale, varianceEpsilon);
|
|
}
|
|
const batchNorm2d = /* @__PURE__ */ op({ batchNorm2d_ });
|
|
function batchNorm3d_(x, mean2, variance, offset, scale2, varianceEpsilon) {
|
|
const $x = convertToTensor(x, "x", "batchNorm");
|
|
const $mean = convertToTensor(mean2, "mean", "batchNorm");
|
|
const $variance = convertToTensor(variance, "variance", "batchNorm");
|
|
let $scale;
|
|
if (scale2 != null) {
|
|
$scale = convertToTensor(scale2, "scale", "batchNorm");
|
|
}
|
|
let $offset;
|
|
if (offset != null) {
|
|
$offset = convertToTensor(offset, "offset", "batchNorm");
|
|
}
|
|
assert($x.rank === 3, () => `Error in batchNorm3D: x must be rank 3 but got rank ${$x.rank}.`);
|
|
assert($mean.rank === 3 || $mean.rank === 1, () => `Error in batchNorm3D: mean must be rank 3 or rank 1 but got rank ${$mean.rank}.`);
|
|
assert($variance.rank === 3 || $variance.rank === 1, () => `Error in batchNorm3D: variance must be rank 3 or rank 1 but got rank ${$variance.rank}.`);
|
|
if ($scale != null) {
|
|
assert($scale.rank === 3 || $scale.rank === 1, () => `Error in batchNorm3D: scale must be rank 3 or rank 1 but got rank ${$scale.rank}.`);
|
|
}
|
|
if ($offset != null) {
|
|
assert($offset.rank === 3 || $offset.rank === 1, () => `Error in batchNorm3D: offset must be rank 3 or rank 1 but got rank ${$offset.rank}.`);
|
|
}
|
|
return batchNorm$2($x, $mean, $variance, $offset, $scale, varianceEpsilon);
|
|
}
|
|
const batchNorm3d = /* @__PURE__ */ op({ batchNorm3d_ });
|
|
function batchNorm4d_(x, mean2, variance, offset, scale2, varianceEpsilon) {
|
|
const $x = convertToTensor(x, "x", "batchNorm");
|
|
const $mean = convertToTensor(mean2, "mean", "batchNorm");
|
|
const $variance = convertToTensor(variance, "variance", "batchNorm");
|
|
let $scale;
|
|
if (scale2 != null) {
|
|
$scale = convertToTensor(scale2, "scale", "batchNorm");
|
|
}
|
|
let $offset;
|
|
if (offset != null) {
|
|
$offset = convertToTensor(offset, "offset", "batchNorm");
|
|
}
|
|
assert($x.rank === 4, () => `Error in batchNorm4D: x must be rank 4 but got rank ${$x.rank}.`);
|
|
assert($mean.rank === 4 || $mean.rank === 1, () => `Error in batchNorm4D: mean must be rank 4 or rank 1 but got rank ${$mean.rank}.`);
|
|
assert($variance.rank === 4 || $variance.rank === 1, () => `Error in batchNorm4D: variance must be rank 4 or rank 1 but got rank ${$variance.rank}.`);
|
|
if ($scale != null) {
|
|
assert($scale.rank === 4 || $scale.rank === 1, () => `Error in batchNorm4D: scale must be rank 4 or rank 1 but got rank ${$scale.rank}.`);
|
|
}
|
|
if ($offset != null) {
|
|
assert($offset.rank === 4 || $offset.rank === 1, () => `Error in batchNorm4D: offset must be rank 4 or rank 1 but got rank ${$offset.rank}.`);
|
|
}
|
|
return batchNorm$2($x, $mean, $variance, $offset, $scale, varianceEpsilon);
|
|
}
|
|
const batchNorm4d = /* @__PURE__ */ op({ batchNorm4d_ });
|
|
function bincount_(x, weights, size) {
|
|
const $x = convertToTensor(x, "x", "bincount");
|
|
const $weights = convertToTensor(weights, "weights", "bincount");
|
|
assert($x.dtype === "int32", () => `Error in bincount: input dtype must be int32, but got ${$x.dtype}`);
|
|
assert(size >= 0, () => `size must be non-negative, but got ${size}.`);
|
|
assert($weights.size === $x.size || $weights.size === 0, () => `Error in bincount: weights must have the same size as input or0-length, but got input shape: ${$x.shape}, weights shape: ${$weights.shape}.`);
|
|
const inputs = { x: $x, weights: $weights };
|
|
const attrs = { size };
|
|
return ENGINE.runKernel(Bincount, inputs, attrs);
|
|
}
|
|
const bincount$2 = /* @__PURE__ */ op({ bincount_ });
|
|
function bitwiseAnd_(x, y) {
|
|
const $x = convertToTensor(x, "x", "bitwiseAnd");
|
|
const $y = convertToTensor(y, "y", "bitwiseAnd");
|
|
if (!arraysEqual($x.shape, $y.shape)) {
|
|
throw new Error(`BitwiseAnd: Tensors must have the same shape. x: ${$x.shape}, y: ${$y.shape}`);
|
|
}
|
|
if ($x.dtype !== "int32" || $y.dtype !== "int32") {
|
|
throw new Error(`BitwiseAnd: Only supports 'int32' values in tensor, found type of x: ${$x.dtype} and type of y: ${$y.dtype}`);
|
|
}
|
|
const inputs = { a: $x, b: $y };
|
|
return ENGINE.runKernel(BitwiseAnd, inputs);
|
|
}
|
|
const bitwiseAnd$2 = /* @__PURE__ */ op({ bitwiseAnd_ });
|
|
function broadcastArgs_(s0, s1) {
|
|
const shape1Input = convertToTensor(s0, "s0", "broadcastArgs", "int32");
|
|
const shape2Input = convertToTensor(s1, "s1", "broadcastArgs", "int32");
|
|
if (shape1Input.rank !== 1) {
|
|
throw new Error(`broadcastArgs(): first input must be a vector (rank=1). Has rank ${shape1Input.rank}`);
|
|
}
|
|
if (shape2Input.rank !== 1) {
|
|
throw new Error(`broadcastArgs(): second input must be a vector (rank=1). Has rank ${shape2Input.rank}`);
|
|
}
|
|
const inputs = { s0: shape1Input, s1: shape2Input };
|
|
return ENGINE.runKernel(BroadcastArgs, inputs);
|
|
}
|
|
const broadcastArgs$2 = /* @__PURE__ */ op({ broadcastArgs_ });
|
|
function broadcastTo_(x, shape) {
|
|
let input = convertToTensor(x, "broadcastTo", "x");
|
|
const xShape = input.shape;
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
if (shape.length < input.rank) {
|
|
throw new Error(`broadcastTo(): shape.length=${shape.length} < input.rank=${input.rank}.`);
|
|
}
|
|
if (shape.length > input.rank) {
|
|
const newShape = input.shape.slice();
|
|
while (newShape.length < shape.length) {
|
|
newShape.unshift(1);
|
|
}
|
|
input = reshape$2(input, newShape);
|
|
}
|
|
const inputShape = input.shape;
|
|
const reps = Array.from(shape);
|
|
for (let i = shape.length - 1; i >= 0; i--) {
|
|
if (inputShape[i] === shape[i]) {
|
|
reps[i] = 1;
|
|
} else if (input.shape[i] !== 1) {
|
|
throw new Error(`broadcastTo(): [${xShape}] cannot be broadcast to [${shape}].`);
|
|
}
|
|
}
|
|
const axes = reps.map((n, i) => n > 1 ? i : -1).filter((i) => i >= 0);
|
|
if (axes.length === 0) {
|
|
return clone(input);
|
|
}
|
|
const inputs = { x: input };
|
|
const attrs = { reps };
|
|
return ENGINE.runKernel(Tile, inputs, attrs);
|
|
}
|
|
const broadcastTo = /* @__PURE__ */ op({ broadcastTo_ });
|
|
function ceil_(x) {
|
|
const $x = convertToTensor(x, "x", "ceil", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Ceil, inputs);
|
|
}
|
|
const ceil$2 = /* @__PURE__ */ op({ ceil_ });
|
|
function fill$2(shape, value, dtype) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
dtype = dtype || inferDtype(value);
|
|
const attrs = { shape, value, dtype };
|
|
return ENGINE.runKernel(Fill, {}, attrs);
|
|
}
|
|
function clipByValue_(x, clipValueMin, clipValueMax) {
|
|
const $x = convertToTensor(x, "x", "clipByValue");
|
|
assert(clipValueMin <= clipValueMax, () => `Error in clip: min (${clipValueMin}) must be less than or equal to max (${clipValueMax}).`);
|
|
if (clipValueMin === clipValueMax) {
|
|
return fill$2($x.shape, clipValueMin, $x.dtype);
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { clipValueMin, clipValueMax };
|
|
return ENGINE.runKernel(ClipByValue, inputs, attrs);
|
|
}
|
|
const clipByValue$2 = /* @__PURE__ */ op({ clipByValue_ });
|
|
function concat1d_(tensors) {
|
|
return concat$2(
|
|
tensors,
|
|
0
|
|
/* axis */
|
|
);
|
|
}
|
|
const concat1d = /* @__PURE__ */ op({ concat1d_ });
|
|
function concat2d_(tensors, axis) {
|
|
return concat$2(tensors, axis);
|
|
}
|
|
const concat2d = /* @__PURE__ */ op({ concat2d_ });
|
|
function concat3d_(tensors, axis) {
|
|
return concat$2(tensors, axis);
|
|
}
|
|
const concat3d = /* @__PURE__ */ op({ concat3d_ });
|
|
function concat4d_(tensors, axis) {
|
|
return concat$2(tensors, axis);
|
|
}
|
|
const concat4d = /* @__PURE__ */ op({ concat4d_ });
|
|
function conv2d_(x, filter, strides, pad2, dataFormat = "NHWC", dilations = [1, 1], dimRoundingMode) {
|
|
const $x = convertToTensor(x, "x", "conv2d", "float32");
|
|
const $filter = convertToTensor(filter, "filter", "conv2d", "float32");
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in conv2d: input must be rank 4, but got rank ${x4D.rank}.`);
|
|
assert($filter.rank === 4, () => `Error in conv2d: filter must be rank 4, but got rank ${$filter.rank}.`);
|
|
checkPadOnDimRoundingMode("conv2d", pad2, dimRoundingMode);
|
|
const inDepth = dataFormat === "NHWC" ? x4D.shape[3] : x4D.shape[1];
|
|
assert(inDepth === $filter.shape[2], () => `Error in conv2d: depth of input (${inDepth}) must match input depth for filter ${$filter.shape[2]}.`);
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in conv2D: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
assert(stridesOrDilationsArePositive(dilations), () => "Error in conv2D: Dilated rates should be larger than 0.");
|
|
assert(stridesOrDilationsArePositive(strides), () => "Error in conv2D: Strides should be larger than 0.");
|
|
const inputs = { x: x4D, filter: $filter };
|
|
const attrs = { strides, pad: pad2, dataFormat, dilations, dimRoundingMode };
|
|
const res = ENGINE.runKernel(Conv2D, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const conv2d$2 = /* @__PURE__ */ op({ conv2d_ });
|
|
function conv1d_(x, filter, stride, pad2, dataFormat = "NWC", dilation = 1, dimRoundingMode) {
|
|
const $x = convertToTensor(x, "x", "conv1d");
|
|
const $filter = convertToTensor(filter, "filter", "conv1d");
|
|
let x3D = $x;
|
|
let reshapedTo3D = false;
|
|
if ($x.rank === 2) {
|
|
reshapedTo3D = true;
|
|
x3D = reshape$2($x, [1, $x.shape[0], $x.shape[1]]);
|
|
}
|
|
assert(x3D.rank === 3, () => `Error in conv1d: input must be rank 3, but got rank ${x3D.rank}.`);
|
|
assert($filter.rank === 3, () => `Error in conv1d: filter must be rank 3, but got rank ${$filter.rank}.`);
|
|
checkPadOnDimRoundingMode("conv1d", pad2, dimRoundingMode);
|
|
assert(x3D.shape[2] === $filter.shape[1], () => `Error in conv1d: depth of input (${x3D.shape[2]}) must match input depth for filter ${$filter.shape[1]}.`);
|
|
assert(eitherStridesOrDilationsAreOne(stride, dilation), () => `Error in conv1D: Either stride or dilation must be 1. Got stride ${stride} and dilation '${dilation}'`);
|
|
assert(stridesOrDilationsArePositive(dilation), () => "Error in conv1D: Dilated rates should be larger than 0.");
|
|
assert(stridesOrDilationsArePositive(stride), () => "Error in conv1D: Stride should be larger than 0.");
|
|
assert(dataFormat === "NWC", () => `Error in conv1d: got dataFormat of ${dataFormat} but only NWC is currently supported.`);
|
|
const filter4D = reshape$2($filter, [1, $filter.shape[0], $filter.shape[1], $filter.shape[2]]);
|
|
const input4D = reshape$2(x3D, [x3D.shape[0], 1, x3D.shape[1], x3D.shape[2]]);
|
|
const strides = [1, stride];
|
|
const dilations = [1, dilation];
|
|
const conv2dDataFormat = "NHWC";
|
|
const res = conv2d$2(input4D, filter4D, strides, pad2, conv2dDataFormat, dilations, dimRoundingMode);
|
|
if (reshapedTo3D) {
|
|
return reshape$2(res, [res.shape[2], res.shape[3]]);
|
|
}
|
|
return reshape$2(res, [res.shape[0], res.shape[2], res.shape[3]]);
|
|
}
|
|
const conv1d = /* @__PURE__ */ op({ conv1d_ });
|
|
function conv2DBackpropInput_(xShape, dy, filter, strides, pad2, dataFormat = "NHWC", dimRoundingMode) {
|
|
assert(xShape.length === dy.rank, () => `Length of inShape (${xShape.length}) and rank of dy (${dy.rank}) must match`);
|
|
let xShape4D = xShape;
|
|
let dy4D = dy;
|
|
let reshapedTo4D = false;
|
|
if (dy.rank === 3) {
|
|
reshapedTo4D = true;
|
|
dy4D = reshape$2(dy, [1, dy.shape[0], dy.shape[1], dy.shape[2]]);
|
|
xShape4D = [1, xShape[0], xShape[1], xShape[2]];
|
|
}
|
|
assert(xShape4D.length === 4, () => `Error in conv2dDerInput: inShape must be length 4, but got length ${xShape4D.length}.`);
|
|
assert(dy4D.rank === 4, () => `Error in conv2dDerInput: dy must be rank 4, but got rank ${dy4D.rank}`);
|
|
assert(filter.rank === 4, () => `Error in conv2dDerInput: filter must be rank 4, but got rank ${filter.rank}`);
|
|
const inDepth = dataFormat === "NHWC" ? xShape4D[3] : xShape4D[1];
|
|
const outDepth = dataFormat === "NHWC" ? dy4D.shape[3] : dy4D.shape[1];
|
|
assert(inDepth === filter.shape[2], () => `Error in conv2dDerInput: depth of input (${inDepth}) must match input depth for filter ${filter.shape[2]}.`);
|
|
assert(outDepth === filter.shape[3], () => `Error in conv2dDerInput: depth of output (${outDepth}) must match output depth for filter ${filter.shape[3]}.`);
|
|
checkPadOnDimRoundingMode("conv2dDerInput", pad2, dimRoundingMode);
|
|
const inputs = { dy: dy4D, filter };
|
|
const attrs = { strides, pad: pad2, dataFormat, dimRoundingMode, inputShape: xShape4D };
|
|
const res = ENGINE.runKernel(Conv2DBackpropInput, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const conv2DBackpropInput$2 = /* @__PURE__ */ op({ conv2DBackpropInput_ });
|
|
function conv2dTranspose_(x, filter, outputShape, strides, pad2, dimRoundingMode) {
|
|
const $x = convertToTensor(x, "x", "conv2dTranspose");
|
|
const $filter = convertToTensor(filter, "filter", "conv2dTranspose");
|
|
return conv2DBackpropInput$2(outputShape, $x, $filter, strides, pad2, "NHWC", dimRoundingMode);
|
|
}
|
|
const conv2dTranspose = /* @__PURE__ */ op({ conv2dTranspose_ });
|
|
function conv3d_(x, filter, strides, pad2, dataFormat = "NDHWC", dilations = [1, 1, 1]) {
|
|
const $x = convertToTensor(x, "x", "conv3d");
|
|
const $filter = convertToTensor(filter, "filter", "conv3d");
|
|
let x5D = $x;
|
|
let reshapedTo5D = false;
|
|
if ($x.rank === 4) {
|
|
reshapedTo5D = true;
|
|
x5D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2], $x.shape[3]]);
|
|
}
|
|
assert(x5D.rank === 5, () => `Error in conv3d: input must be rank 5, but got rank ${x5D.rank}.`);
|
|
assert($filter.rank === 5, () => `Error in conv3d: filter must be rank 5, but got rank ${$filter.rank}.`);
|
|
assert(x5D.shape[4] === $filter.shape[3], () => `Error in conv3d: depth of input (${x5D.shape[4]}) must match input depth for filter ${$filter.shape[3]}.`);
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in conv3D: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
assert(dataFormat === "NDHWC", () => `Error in conv3d: got dataFormat of ${dataFormat} but only NDHWC is currently supported.`);
|
|
assert(stridesOrDilationsArePositive(dilations), () => "Error in conv3D: Dilated rates should be larger than 0.");
|
|
assert(stridesOrDilationsArePositive(strides), () => "Error in conv3D: Strides should be larger than 0.");
|
|
const inputs = { x: x5D, filter: $filter };
|
|
const attrs = { strides, pad: pad2, dataFormat, dilations };
|
|
const res = ENGINE.runKernel(Conv3D, inputs, attrs);
|
|
if (reshapedTo5D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3], res.shape[4]]);
|
|
}
|
|
return res;
|
|
}
|
|
const conv3d = /* @__PURE__ */ op({ conv3d_ });
|
|
function conv3DBackpropInput_(xShape, dy, filter, strides, pad2) {
|
|
assert(xShape.length === dy.rank, () => `Length of inShape (${xShape.length}) and rank of dy (${dy.rank}) must match`);
|
|
let xShape5D = xShape;
|
|
let dy5D = dy;
|
|
let reshapedTo5D = false;
|
|
if (dy.rank === 4) {
|
|
reshapedTo5D = true;
|
|
dy5D = reshape$2(dy, [1, dy.shape[0], dy.shape[1], dy.shape[2], dy.shape[3]]);
|
|
xShape5D = [1, xShape[0], xShape[1], xShape[2], xShape[3]];
|
|
}
|
|
const inDepth = xShape5D[4];
|
|
const outDepth = dy5D.shape[4];
|
|
assert(xShape5D.length === 5, () => `Error in conv3dDerInput: inShape must be length 5, but got length ${xShape5D.length}.`);
|
|
assert(dy5D.rank === 5, () => `Error in conv3dDerInput: dy must be rank 5, but got rank ${dy5D.rank}`);
|
|
assert(filter.rank === 5, () => `Error in conv3dDerInput: filter must be rank 5, but got rank ${filter.rank}`);
|
|
assert(inDepth === filter.shape[3], () => `Error in conv3dDerInput: depth of input (${inDepth}) must match input depth for filter ${filter.shape[3]}.`);
|
|
assert(outDepth === filter.shape[4], () => `Error in conv3dDerInput: depth of output (${outDepth}) must match output depth for filter ${filter.shape[4]}.`);
|
|
const inputs = { dy: dy5D, filter };
|
|
const attrs = { pad: pad2, strides, inputShape: xShape5D };
|
|
const res = ENGINE.runKernel(Conv3DBackpropInputV2, inputs, attrs);
|
|
if (reshapedTo5D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3], res.shape[4]]);
|
|
}
|
|
return res;
|
|
}
|
|
const conv3DBackpropInput$1 = /* @__PURE__ */ op({ conv3DBackpropInput_ });
|
|
function conv3dTranspose_(x, filter, outputShape, strides, pad2) {
|
|
const $x = convertToTensor(x, "x", "conv3dTranspose");
|
|
const $filter = convertToTensor(filter, "filter", "conv3dTranspose");
|
|
return conv3DBackpropInput$1(outputShape, $x, $filter, strides, pad2);
|
|
}
|
|
const conv3dTranspose = /* @__PURE__ */ op({ conv3dTranspose_ });
|
|
function cos_(x) {
|
|
const $x = convertToTensor(x, "x", "cos", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Cos, inputs);
|
|
}
|
|
const cos$2 = /* @__PURE__ */ op({ cos_ });
|
|
function cosh_(x) {
|
|
const $x = convertToTensor(x, "x", "cosh", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Cosh, inputs);
|
|
}
|
|
const cosh$2 = /* @__PURE__ */ op({ cosh_ });
|
|
function cumprod_(x, axis = 0, exclusive = false, reverse2 = false) {
|
|
const $x = convertToTensor(x, "x", "cumprod");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, exclusive, reverse: reverse2 };
|
|
return ENGINE.runKernel(Cumprod, inputs, attrs);
|
|
}
|
|
const cumprod$2 = /* @__PURE__ */ op({ cumprod_ });
|
|
function cumsum_(x, axis = 0, exclusive = false, reverse2 = false) {
|
|
const $x = convertToTensor(x, "x", "cumsum");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, exclusive, reverse: reverse2 };
|
|
return ENGINE.runKernel(Cumsum, inputs, attrs);
|
|
}
|
|
const cumsum$2 = /* @__PURE__ */ op({ cumsum_ });
|
|
function denseBincount_(x, weights, size, binaryOutput = false) {
|
|
const $x = convertToTensor(x, "x", "denseBincount");
|
|
const $weights = convertToTensor(weights, "weights", "denseBincount");
|
|
assert($x.dtype === "int32", () => `Error in denseBincount: input dtype must be int32, but got ${$x.dtype}`);
|
|
assert($x.rank <= 2, () => `Error in denseBincount: input must be at most rank 2, but got rank ${$x.rank}.`);
|
|
assert(size >= 0, () => `size must be non-negative, but got ${size}.`);
|
|
assert($weights.size === $x.size || $weights.size === 0, () => `Error in denseBincount: weights must have the same shape as x or 0-length, but got x shape: ${$x.shape}, weights shape: ${$weights.shape}.`);
|
|
const inputs = { x: $x, weights: $weights };
|
|
const attrs = { size, binaryOutput };
|
|
return ENGINE.runKernel(DenseBincount, inputs, attrs);
|
|
}
|
|
const denseBincount$2 = /* @__PURE__ */ op({ denseBincount_ });
|
|
function depthToSpace_(x, blockSize, dataFormat = "NHWC") {
|
|
const $x = convertToTensor(x, "x", "depthToSpace", "float32");
|
|
const inputHeight = dataFormat === "NHWC" ? $x.shape[1] : $x.shape[2];
|
|
const inputWidth = dataFormat === "NHWC" ? $x.shape[2] : $x.shape[3];
|
|
const inputDepth = dataFormat === "NHWC" ? $x.shape[3] : $x.shape[1];
|
|
assert(blockSize > 1, () => `blockSize should be > 1 for depthToSpace, but was: ${blockSize}`);
|
|
assert(inputHeight * blockSize >= 0, () => `Negative dimension size caused by overflow when multiplying
|
|
${inputHeight} and ${blockSize} for depthToSpace with input shape
|
|
${$x.shape}`);
|
|
assert(inputWidth * blockSize >= 0, () => `Negative dimension size caused by overflow when multiplying
|
|
${inputWidth} and ${blockSize} for depthToSpace with input shape
|
|
${$x.shape}`);
|
|
assert(inputDepth % (blockSize * blockSize) === 0, () => `Dimension size must be evenly divisible by ${blockSize * blockSize} but is ${inputDepth} for depthToSpace with input shape ${$x.shape}`);
|
|
const inputs = { x: $x };
|
|
const attrs = { blockSize, dataFormat };
|
|
return ENGINE.runKernel(DepthToSpace, inputs, attrs);
|
|
}
|
|
const depthToSpace$2 = /* @__PURE__ */ op({ depthToSpace_ });
|
|
function depthwiseConv2d_(x, filter, strides, pad2, dataFormat = "NHWC", dilations = [1, 1], dimRoundingMode) {
|
|
const $x = convertToTensor(x, "x", "depthwiseConv2d", "float32");
|
|
const $filter = convertToTensor(filter, "filter", "depthwiseConv2d", "float32");
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in depthwiseConv2d: input must be rank 4, but got rank ${x4D.rank}.`);
|
|
assert($filter.rank === 4, () => `Error in depthwiseConv2d: filter must be rank 4, but got rank ${$filter.rank}.`);
|
|
const inChannels = dataFormat === "NHWC" ? x4D.shape[3] : x4D.shape[1];
|
|
assert(inChannels === $filter.shape[2], () => `Error in depthwiseConv2d: number of input channels (${inChannels}) must match the inChannels dimension in filter ${$filter.shape[2]}.`);
|
|
checkPadOnDimRoundingMode("depthwiseConv2d", pad2, dimRoundingMode);
|
|
const inputs = { x: x4D, filter: $filter };
|
|
const attrs = { strides, pad: pad2, dataFormat, dilations, dimRoundingMode };
|
|
const res = ENGINE.runKernel(DepthwiseConv2dNative, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const depthwiseConv2d$1 = /* @__PURE__ */ op({ depthwiseConv2d_ });
|
|
function diag_(x) {
|
|
const $x = convertToTensor(x, "x", "diag");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Diag, inputs);
|
|
}
|
|
const diag$2 = /* @__PURE__ */ op({ diag_ });
|
|
function dilation2d_(x, filter, strides, pad2, dilations = [1, 1], dataFormat = "NHWC") {
|
|
const $x = convertToTensor(x, "x", "dilation2d");
|
|
const $filter = convertToTensor(filter, "filter", "dilation2d");
|
|
assert($x.rank === 3 || $x.rank === 4, () => `Error in dilation2d: input must be rank 3 or 4, but got rank ${$x.rank}.`);
|
|
assert($filter.rank === 3, () => `Error in dilation2d: filter must be rank 3, but got rank ${$filter.rank}.`);
|
|
assert(dataFormat === "NHWC", () => `Error in dilation2d: Only NHWC is currently supported, but got dataFormat of ${dataFormat}`);
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
reshapedTo4D = true;
|
|
}
|
|
assert(x4D.shape[3] === $filter.shape[2], () => `Error in dilation2d: input and filter must have the same depth: ${x4D.shape[3]} vs ${$filter.shape[2]}`);
|
|
const inputs = { x: x4D, filter: $filter };
|
|
const attrs = { strides, pad: pad2, dilations };
|
|
const res = ENGINE.runKernel(Dilation2D, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const dilation2d = /* @__PURE__ */ op({ dilation2d_ });
|
|
function getBroadcastDims$1(inShape, outShape) {
|
|
const inRank = inShape.length;
|
|
const dims = [];
|
|
for (let i = 0; i < inRank; i++) {
|
|
const dim = inRank - 1 - i;
|
|
const a = inShape[dim] || 1;
|
|
const b = outShape[outShape.length - 1 - i] || 1;
|
|
if (b > 1 && a === 1) {
|
|
dims.unshift(dim);
|
|
}
|
|
}
|
|
return dims;
|
|
}
|
|
function getReductionAxes(inShape, outShape) {
|
|
const result = [];
|
|
for (let i = 0; i < outShape.length; i++) {
|
|
const inDim = inShape[inShape.length - i - 1];
|
|
const outAxis = outShape.length - i - 1;
|
|
const outDim = outShape[outAxis];
|
|
if (inDim == null || inDim === 1 && outDim > 1) {
|
|
result.unshift(outAxis);
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
function assertAndGetBroadcastShape(shapeA, shapeB) {
|
|
const l = Math.max(shapeA.length, shapeB.length);
|
|
const result = new Array(l);
|
|
for (let i = 0; i < l; i++) {
|
|
let a = shapeA[shapeA.length - i - 1];
|
|
if (a == null) {
|
|
a = 1;
|
|
}
|
|
let b = shapeB[shapeB.length - i - 1];
|
|
if (b == null) {
|
|
b = 1;
|
|
}
|
|
if (a === 1) {
|
|
result[l - i - 1] = b;
|
|
} else if (b === 1) {
|
|
result[l - i - 1] = a;
|
|
} else if (a !== b) {
|
|
const errMsg = `Operands could not be broadcast together with shapes ${shapeA} and ${shapeB}.`;
|
|
throw Error(errMsg);
|
|
} else {
|
|
result[l - i - 1] = a;
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
function equal_(a, b) {
|
|
let $a = convertToTensor(a, "a", "equal", "string_or_numeric");
|
|
let $b = convertToTensor(b, "b", "equal", "string_or_numeric");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Equal, inputs);
|
|
}
|
|
const equal$2 = /* @__PURE__ */ op({ equal_ });
|
|
function where_(condition, a, b) {
|
|
const $a = convertToTensor(a, "a", "where");
|
|
const $b = convertToTensor(b, "b", "where");
|
|
const $condition = convertToTensor(condition, "condition", "where", "bool");
|
|
const broadcastShape = assertAndGetBroadcastShape(assertAndGetBroadcastShape($condition.shape, $a.shape), $b.shape);
|
|
const $broadcastedCondition = broadcastTo($condition, broadcastShape);
|
|
const $broadcastedA = broadcastTo($a, broadcastShape);
|
|
const $broadcastedB = broadcastTo($b, broadcastShape);
|
|
const inputs = {
|
|
condition: $broadcastedCondition,
|
|
t: $broadcastedA,
|
|
e: $broadcastedB
|
|
};
|
|
return ENGINE.runKernel(Select, inputs);
|
|
}
|
|
const where = /* @__PURE__ */ op({ where_ });
|
|
function zerosLike_(x) {
|
|
const $x = convertToTensor(x, "x", "zerosLike");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(ZerosLike, inputs);
|
|
}
|
|
const zerosLike$2 = /* @__PURE__ */ op({ zerosLike_ });
|
|
function divNoNan_(a, b) {
|
|
let $a = convertToTensor(a, "a", "div");
|
|
let $b = convertToTensor(b, "b", "div");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const divResult = div$1($a, $b);
|
|
const zeros2 = zerosLike$2(divResult);
|
|
const bEqualsZero = equal$2($b, zeros2);
|
|
return where(bEqualsZero, zeros2, divResult);
|
|
}
|
|
const divNoNan = /* @__PURE__ */ op({ divNoNan_ });
|
|
function dot_(t1, t2) {
|
|
const $t1 = convertToTensor(t1, "t1", "dot");
|
|
const $t2 = convertToTensor(t2, "t2", "dot");
|
|
assert(($t1.rank === 1 || $t1.rank === 2) && ($t2.rank === 1 || $t2.rank === 2), () => `Error in dot: inputs must all be rank 1 or 2, but got ranks ${$t1.rank} and ${$t2.rank}.`);
|
|
const t1Inner = $t1.rank === 1 ? $t1.size : $t1.shape[1];
|
|
const t2Inner = $t2.rank === 1 ? $t2.size : $t2.shape[0];
|
|
assert(t1Inner === t2Inner, () => `Error in dot: inner dimensions of inputs must match, but got ${t1Inner} and ${t2Inner}.`);
|
|
if ($t1.rank === 1 && $t2.rank === 1) {
|
|
const t12D = reshape$2($t1, [1, -1]);
|
|
const t22D = reshape$2($t2, [-1, 1]);
|
|
const t1t2 = matMul$1(t12D, t22D);
|
|
return reshape$2(t1t2, []);
|
|
} else if ($t1.rank === 1 && $t2.rank === 2) {
|
|
const t12D = reshape$2($t1, [1, -1]);
|
|
const t22D = reshape$2($t2, [$t2.shape[0], $t2.shape[1]]);
|
|
const t1t2 = matMul$1(t12D, t22D);
|
|
return reshape$2(t1t2, [t1t2.size]);
|
|
} else if ($t1.rank === 2 && $t2.rank === 1) {
|
|
const t22D = reshape$2($t2, [-1, 1]);
|
|
const t1t2 = matMul$1($t1, t22D);
|
|
return reshape$2(t1t2, [t1t2.size]);
|
|
} else {
|
|
const t22D = reshape$2($t2, [$t2.shape[0], $t2.shape[1]]);
|
|
const t1t2 = matMul$1($t1, t22D);
|
|
return t1t2;
|
|
}
|
|
}
|
|
const dot = /* @__PURE__ */ op({ dot_ });
|
|
function einsum_(equation, ...tensors) {
|
|
const $tensors = tensors.map((t2, i) => convertToTensor(t2, `tensors${i}`, "einsum"));
|
|
const attrs = { equation };
|
|
return ENGINE.runKernel(Einsum, $tensors, attrs);
|
|
}
|
|
const einsum$2 = /* @__PURE__ */ op({ einsum_ });
|
|
function elu_(x) {
|
|
const $x = convertToTensor(x, "x", "elu", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Elu, inputs);
|
|
}
|
|
const elu$2 = /* @__PURE__ */ op({ elu_ });
|
|
function ensureShape_(x, shape) {
|
|
const $x = convertToTensor(x, "x", "ensureShape", "string_or_numeric");
|
|
if (!arraysEqualWithNull($x.shape, shape)) {
|
|
throw new Error(`EnsureShape: Shape of tensor ${$x.shape} is not compatible with expected shape ${shape}`);
|
|
}
|
|
return x;
|
|
}
|
|
const ensureShape = /* @__PURE__ */ op({ ensureShape_ });
|
|
function erf_(x) {
|
|
let $x = convertToTensor(x, "x", "erf");
|
|
assert($x.dtype === "int32" || $x.dtype === "float32", () => "Input dtype must be `int32` or `float32`.");
|
|
if ($x.dtype === "int32") {
|
|
$x = cast$2($x, "float32");
|
|
}
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Erf, inputs);
|
|
}
|
|
const erf$2 = /* @__PURE__ */ op({ erf_ });
|
|
function axesAreInnerMostDims(axes, rank) {
|
|
for (let i = 0; i < axes.length; ++i) {
|
|
if (axes[axes.length - i - 1] !== rank - 1 - i) {
|
|
return false;
|
|
}
|
|
}
|
|
return true;
|
|
}
|
|
function combineLocations(outputLoc, reduceLoc, axes) {
|
|
const rank = outputLoc.length + reduceLoc.length;
|
|
const loc = [];
|
|
let outIdx = 0;
|
|
let reduceIdx = 0;
|
|
for (let dim = 0; dim < rank; dim++) {
|
|
if (axes.indexOf(dim) === -1) {
|
|
loc.push(outputLoc[outIdx++]);
|
|
} else {
|
|
loc.push(reduceLoc[reduceIdx++]);
|
|
}
|
|
}
|
|
return loc;
|
|
}
|
|
function computeOutAndReduceShapes(aShape, axes) {
|
|
const outShape = [];
|
|
const rank = aShape.length;
|
|
for (let dim = 0; dim < rank; dim++) {
|
|
if (axes.indexOf(dim) === -1) {
|
|
outShape.push(aShape[dim]);
|
|
}
|
|
}
|
|
const reduceShape = axes.map((dim) => aShape[dim]);
|
|
return [outShape, reduceShape];
|
|
}
|
|
function expandShapeToKeepDim(shape, axes) {
|
|
const reduceSubShape = axes.map((x) => 1);
|
|
return combineLocations(shape, reduceSubShape, axes);
|
|
}
|
|
function assertAxesAreInnerMostDims(msg, axes, rank) {
|
|
assert(axesAreInnerMostDims(axes, rank), () => `${msg} supports only inner-most axes for now. Got axes ${axes} and rank-${rank} input.`);
|
|
}
|
|
function getAxesPermutation(axes, rank) {
|
|
if (axesAreInnerMostDims(axes, rank)) {
|
|
return null;
|
|
}
|
|
const result = [];
|
|
for (let i = 0; i < rank; ++i) {
|
|
if (axes.indexOf(i) === -1) {
|
|
result.push(i);
|
|
}
|
|
}
|
|
axes.forEach((axis) => result.push(axis));
|
|
return result;
|
|
}
|
|
function getUndoAxesPermutation(axes) {
|
|
return axes.map((axis, i) => [i, axis]).sort((a, b) => a[1] - b[1]).map((x) => x[0]);
|
|
}
|
|
function getInnerMostAxes(numAxes, rank) {
|
|
const res = [];
|
|
for (let i = rank - numAxes; i < rank; ++i) {
|
|
res.push(i);
|
|
}
|
|
return res;
|
|
}
|
|
function max_(x, axis = null, keepDims = false) {
|
|
const $x = convertToTensor(x, "x", "max");
|
|
const inputs = { x: $x };
|
|
const attrs = { reductionIndices: axis, keepDims };
|
|
return ENGINE.runKernel(Max, inputs, attrs);
|
|
}
|
|
const max$2 = /* @__PURE__ */ op({ max_ });
|
|
function min_(x, axis = null, keepDims = false) {
|
|
const $x = convertToTensor(x, "x", "min");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, keepDims };
|
|
return ENGINE.runKernel(Min, inputs, attrs);
|
|
}
|
|
const min$2 = /* @__PURE__ */ op({ min_ });
|
|
function pow_(base, exp2) {
|
|
let $base = convertToTensor(base, "base", "pow");
|
|
let $exp = convertToTensor(exp2, "exp", "pow");
|
|
[$base, $exp] = makeTypesMatch($base, $exp);
|
|
const inputs = { a: $base, b: $exp };
|
|
return ENGINE.runKernel(Pow, inputs);
|
|
}
|
|
const pow$2 = /* @__PURE__ */ op({ pow_ });
|
|
function scalar(value, dtype) {
|
|
if ((isTypedArray(value) && dtype !== "string" || Array.isArray(value)) && dtype !== "complex64") {
|
|
throw new Error("Error creating a new Scalar: value must be a primitive (number|boolean|string)");
|
|
}
|
|
if (dtype === "string" && isTypedArray(value) && !(value instanceof Uint8Array)) {
|
|
throw new Error("When making a scalar from encoded string, the value must be `Uint8Array`.");
|
|
}
|
|
const shape = [];
|
|
const inferredShape = [];
|
|
return makeTensor(value, shape, inferredShape, dtype);
|
|
}
|
|
function sqrt_(x) {
|
|
const $x = convertToTensor(x, "x", "sqrt", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Sqrt, inputs);
|
|
}
|
|
const sqrt$2 = /* @__PURE__ */ op({ sqrt_ });
|
|
function square_(x) {
|
|
const $x = convertToTensor(x, "x", "square");
|
|
const attrs = {};
|
|
return ENGINE.runKernel("Square", { x: $x }, attrs);
|
|
}
|
|
const square$1 = /* @__PURE__ */ op({ square_ });
|
|
function sum_(x, axis = null, keepDims = false) {
|
|
let $x = convertToTensor(x, "x", "sum");
|
|
if ($x.dtype === "bool") {
|
|
$x = cast$2($x, "int32");
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, keepDims };
|
|
return ENGINE.runKernel(Sum, inputs, attrs);
|
|
}
|
|
const sum$2 = /* @__PURE__ */ op({ sum_ });
|
|
function norm_(x, ord = "euclidean", axis = null, keepDims = false) {
|
|
x = convertToTensor(x, "x", "norm");
|
|
const norm2 = normImpl(x, ord, axis);
|
|
let keepDimsShape = norm2.shape;
|
|
if (keepDims) {
|
|
const axes = parseAxisParam(axis, x.shape);
|
|
keepDimsShape = expandShapeToKeepDim(norm2.shape, axes);
|
|
}
|
|
return reshape$2(norm2, keepDimsShape);
|
|
}
|
|
function normImpl(x, p2, axis = null) {
|
|
if (x.rank === 0) {
|
|
return abs$2(x);
|
|
}
|
|
if (x.rank !== 1 && axis === null) {
|
|
return normImpl(reshape$2(x, [-1]), p2, axis);
|
|
}
|
|
if (x.rank === 1 || typeof axis === "number" || Array.isArray(axis) && axis.length === 1) {
|
|
if (p2 === 1) {
|
|
return sum$2(abs$2(x), axis);
|
|
}
|
|
if (p2 === Infinity) {
|
|
return max$2(abs$2(x), axis);
|
|
}
|
|
if (p2 === -Infinity) {
|
|
return min$2(abs$2(x), axis);
|
|
}
|
|
if (p2 === "euclidean" || p2 === 2) {
|
|
return sqrt$2(sum$2(pow$2(abs$2(x), scalar(2, "int32")), axis));
|
|
}
|
|
throw new Error(`Error in norm: invalid ord value: ${p2}`);
|
|
}
|
|
if (Array.isArray(axis) && axis.length === 2) {
|
|
if (p2 === 1) {
|
|
return max$2(sum$2(abs$2(x), axis[0]), axis[1] - 1);
|
|
}
|
|
if (p2 === Infinity) {
|
|
return max$2(sum$2(abs$2(x), axis[1]), axis[0]);
|
|
}
|
|
if (p2 === -Infinity) {
|
|
return min$2(sum$2(abs$2(x), axis[1]), axis[0]);
|
|
}
|
|
if (p2 === "fro" || p2 === "euclidean") {
|
|
return sqrt$2(sum$2(square$1(x), axis));
|
|
}
|
|
throw new Error(`Error in norm: invalid ord value: ${p2}`);
|
|
}
|
|
throw new Error(`Error in norm: invalid axis: ${axis}`);
|
|
}
|
|
const norm = /* @__PURE__ */ op({ norm_ });
|
|
function euclideanNorm_(x, axis = null, keepDims = false) {
|
|
return norm(x, "euclidean", axis, keepDims);
|
|
}
|
|
const euclideanNorm = /* @__PURE__ */ op({ euclideanNorm_ });
|
|
function exp_(x) {
|
|
const $x = convertToTensor(x, "x", "exp");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Exp, inputs);
|
|
}
|
|
const exp$2 = /* @__PURE__ */ op({ exp_ });
|
|
function expandDims_(x, axis = 0) {
|
|
const $x = convertToTensor(x, "x", "expandDims", "string_or_numeric");
|
|
assert(axis <= $x.rank, () => "Axis must be <= rank of the tensor");
|
|
const inputs = { input: $x };
|
|
const attrs = { dim: axis };
|
|
return ENGINE.runKernel(ExpandDims, inputs, attrs);
|
|
}
|
|
const expandDims$2 = /* @__PURE__ */ op({ expandDims_ });
|
|
function expm1_(x) {
|
|
const $x = convertToTensor(x, "x", "expm1");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Expm1, inputs);
|
|
}
|
|
const expm1$2 = /* @__PURE__ */ op({ expm1_ });
|
|
function tile_(x, reps) {
|
|
const $x = convertToTensor(x, "x", "tile", "string_or_numeric");
|
|
assert($x.rank === reps.length, () => `Error in transpose: rank of input ${$x.rank} must match length of reps ${reps}.`);
|
|
const inputs = { x: $x };
|
|
const attrs = { reps };
|
|
return ENGINE.runKernel(Tile, inputs, attrs);
|
|
}
|
|
const tile$2 = /* @__PURE__ */ op({ tile_ });
|
|
function eye_(numRows, numColumns, batchShape, dtype = "float32") {
|
|
if (numColumns == null) {
|
|
numColumns = numRows;
|
|
}
|
|
const buff = buffer([numRows, numColumns], dtype);
|
|
const n = numRows <= numColumns ? numRows : numColumns;
|
|
for (let i = 0; i < n; ++i) {
|
|
buff.set(1, i, i);
|
|
}
|
|
const out = reshape$2(buff.toTensor(), [numRows, numColumns]);
|
|
if (batchShape == null) {
|
|
return out;
|
|
} else {
|
|
if (batchShape.length === 1) {
|
|
return tile$2(expandDims$2(out, 0), [batchShape[0], 1, 1]);
|
|
} else if (batchShape.length === 2) {
|
|
return tile$2(expandDims$2(expandDims$2(out, 0), 0), [batchShape[0], batchShape[1], 1, 1]);
|
|
} else if (batchShape.length === 3) {
|
|
return tile$2(expandDims$2(expandDims$2(expandDims$2(out, 0), 0), 0), [
|
|
batchShape[0],
|
|
batchShape[1],
|
|
batchShape[2],
|
|
1,
|
|
1
|
|
]);
|
|
} else {
|
|
throw new Error(`eye() currently supports only 1D and 2D batchShapes, but received ${batchShape.length}D.`);
|
|
}
|
|
}
|
|
}
|
|
const eye = /* @__PURE__ */ op({ eye_ });
|
|
function floor_(x) {
|
|
const $x = convertToTensor(x, "x", "floor", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Floor, inputs);
|
|
}
|
|
const floor$2 = /* @__PURE__ */ op({ floor_ });
|
|
function gather_(x, indices, axis = 0, batchDims = 0) {
|
|
const $x = convertToTensor(x, "x", "gather");
|
|
const $indices = convertToTensor(indices, "indices", "gather", "int32");
|
|
const inputs = { x: $x, indices: $indices };
|
|
const attrs = { axis, batchDims };
|
|
return ENGINE.runKernel(GatherV2, inputs, attrs);
|
|
}
|
|
const gather = /* @__PURE__ */ op({ gather_ });
|
|
function greater_(a, b) {
|
|
let $a = convertToTensor(a, "a", "greater", "string_or_numeric");
|
|
let $b = convertToTensor(b, "b", "greater", "string_or_numeric");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Greater, inputs);
|
|
}
|
|
const greater$2 = /* @__PURE__ */ op({ greater_ });
|
|
function greaterEqual_(a, b) {
|
|
let $a = convertToTensor(a, "a", "greaterEqual", "string_or_numeric");
|
|
let $b = convertToTensor(b, "b", "greaterEqual", "string_or_numeric");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(GreaterEqual, inputs);
|
|
}
|
|
const greaterEqual$2 = /* @__PURE__ */ op({ greaterEqual_ });
|
|
function imag_(input) {
|
|
const $input = convertToTensor(input, "input", "imag");
|
|
const inputs = { input: $input };
|
|
return ENGINE.runKernel(Imag, inputs);
|
|
}
|
|
const imag$2 = /* @__PURE__ */ op({ imag_ });
|
|
function isFinite_(x) {
|
|
const $x = convertToTensor(x, "x", "isFinite");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(IsFinite, inputs);
|
|
}
|
|
const isFinite$3 = /* @__PURE__ */ op({ isFinite_ });
|
|
function isInf_(x) {
|
|
const $x = convertToTensor(x, "x", "isInf");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(IsInf, inputs);
|
|
}
|
|
const isInf$2 = /* @__PURE__ */ op({ isInf_ });
|
|
function isNaN_(x) {
|
|
const $x = convertToTensor(x, "x", "isNaN");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(IsNan, inputs);
|
|
}
|
|
const isNaN$3 = /* @__PURE__ */ op({ isNaN_ });
|
|
function leakyRelu_(x, alpha = 0.2) {
|
|
const $x = convertToTensor(x, "x", "leakyRelu");
|
|
const inputs = { x: $x };
|
|
const attrs = { alpha };
|
|
return ENGINE.runKernel(LeakyRelu, inputs, attrs);
|
|
}
|
|
const leakyRelu$2 = /* @__PURE__ */ op({ leakyRelu_ });
|
|
function less_(a, b) {
|
|
let $a = convertToTensor(a, "a", "less", "string_or_numeric");
|
|
let $b = convertToTensor(b, "b", "less", "string_or_numeric");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Less, inputs);
|
|
}
|
|
const less$2 = /* @__PURE__ */ op({ less_ });
|
|
function lessEqual_(a, b) {
|
|
let $a = convertToTensor(a, "a", "lessEqual", "string_or_numeric");
|
|
let $b = convertToTensor(b, "b", "lessEqual", "string_or_numeric");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(LessEqual, inputs);
|
|
}
|
|
const lessEqual$2 = /* @__PURE__ */ op({ lessEqual_ });
|
|
function linspace(start, stop, num) {
|
|
if (num <= 0) {
|
|
throw new Error("The number of values should be positive.");
|
|
}
|
|
const attrs = { start, stop, num };
|
|
return ENGINE.runKernel(LinSpace, {}, attrs);
|
|
}
|
|
function localResponseNormalization_(x, depthRadius = 5, bias = 1, alpha = 1, beta = 0.5) {
|
|
const $x = convertToTensor(x, "x", "localResponseNormalization");
|
|
assert($x.rank === 4 || $x.rank === 3, () => `Error in localResponseNormalization: x must be rank 3 or 4 but got
|
|
rank ${$x.rank}.`);
|
|
assert(isInt(depthRadius), () => `Error in localResponseNormalization: depthRadius must be an integer but got depthRadius ${depthRadius}.`);
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
const inputs = { x: x4D };
|
|
const attrs = { depthRadius, bias, alpha, beta };
|
|
const res = ENGINE.runKernel(LRN, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
} else {
|
|
return res;
|
|
}
|
|
}
|
|
const localResponseNormalization = /* @__PURE__ */ op({ localResponseNormalization_ });
|
|
function log_(x) {
|
|
const $x = convertToTensor(x, "x", "log", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Log, inputs);
|
|
}
|
|
const log$2 = /* @__PURE__ */ op({ log_ });
|
|
function log1p_(x) {
|
|
const $x = convertToTensor(x, "x", "log1p");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Log1p, inputs);
|
|
}
|
|
const log1p$2 = /* @__PURE__ */ op({ log1p_ });
|
|
function variableGrads(f, varList) {
|
|
assert(isFunction(f), () => "The f passed in variableGrads(f) must be a function");
|
|
assert(varList == null || Array.isArray(varList) && varList.every((v) => v instanceof Variable), () => "The varList passed in variableGrads(f, varList) must be an array of variables");
|
|
const specifiedVarList = varList != null;
|
|
if (!specifiedVarList) {
|
|
varList = [];
|
|
for (const varName in ENGINE.registeredVariables) {
|
|
varList.push(ENGINE.registeredVariables[varName]);
|
|
}
|
|
}
|
|
const specifiedNonTrainable = specifiedVarList ? varList.filter((variable2) => !variable2.trainable) : null;
|
|
const originalVarCount = varList.length;
|
|
varList = varList.filter((variable2) => variable2.trainable);
|
|
assert(varList.length > 0, () => `variableGrads() expects at least one of the input variables to be trainable, but none of the ${originalVarCount} variables is trainable.`);
|
|
const allowNoGradients = true;
|
|
const { value, grads } = ENGINE.gradients(f, varList, null, allowNoGradients);
|
|
assert(grads.some((g) => g != null), () => "Cannot find a connection between any variable and the result of the loss function y=f(x). Please make sure the operations that use variables are inside the function f passed to minimize().");
|
|
assert(value.rank === 0, () => `The f passed in variableGrads(f) must return a scalar, but it returned a rank-${value.rank} tensor`);
|
|
const namedGrads = {};
|
|
varList.forEach((v, i) => {
|
|
if (grads[i] != null) {
|
|
namedGrads[v.name] = grads[i];
|
|
}
|
|
});
|
|
if (specifiedNonTrainable != null) {
|
|
specifiedNonTrainable.forEach((v) => namedGrads[v.name] = null);
|
|
}
|
|
return { value, grads: namedGrads };
|
|
}
|
|
function customGrad(f) {
|
|
return ENGINE.customGrad(f);
|
|
}
|
|
function neg_(x) {
|
|
const $x = convertToTensor(x, "x", "neg");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Neg, inputs);
|
|
}
|
|
const neg$2 = /* @__PURE__ */ op({ neg_ });
|
|
function softplus_(x) {
|
|
const $x = convertToTensor(x, "x", "softplus");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Softplus, inputs);
|
|
}
|
|
const softplus$2 = /* @__PURE__ */ op({ softplus_ });
|
|
function logSigmoid_(x) {
|
|
const $x = convertToTensor(x, "x", "logSigmoid");
|
|
const customOp = customGrad((x2) => {
|
|
const value = neg$2(softplus$2(neg$2(x2)));
|
|
const gradFunc = (dy) => {
|
|
const derX = mul(dy, sigmoid$2(neg$2(x2)));
|
|
return derX;
|
|
};
|
|
return { value, gradFunc };
|
|
});
|
|
return customOp($x);
|
|
}
|
|
const logSigmoid = /* @__PURE__ */ op({ logSigmoid_ });
|
|
function sub_(a, b) {
|
|
let $a = convertToTensor(a, "a", "sub");
|
|
let $b = convertToTensor(b, "b", "sub");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Sub, inputs);
|
|
}
|
|
const sub$2 = /* @__PURE__ */ op({ sub_ });
|
|
function logSoftmax_(logits, axis = -1) {
|
|
const $logits = convertToTensor(logits, "logits", "logSoftmax");
|
|
if (axis === -1) {
|
|
axis = $logits.rank - 1;
|
|
}
|
|
if (axis !== $logits.rank - 1) {
|
|
throw Error(`Log Softmax along a non-last dimension is not yet supported. Logits was rank ${$logits.rank} and axis was ${axis}`);
|
|
}
|
|
const customOp = customGrad((logits2, save) => {
|
|
const keepDims = true;
|
|
const xMax = max$2(logits2, axis, true);
|
|
const shifted = sub$2(logits2, xMax);
|
|
const value = sub$2(cast$2(shifted, "float32"), log$2(sum$2(exp$2(shifted), axis, keepDims)));
|
|
save([value]);
|
|
const gradFunc = (dy, saved) => {
|
|
const [value2] = saved;
|
|
const keepDims2 = true;
|
|
const softmax2 = exp$2(value2);
|
|
return sub$2(dy, mul(sum$2(dy, axis, keepDims2), softmax2));
|
|
};
|
|
return { value, gradFunc };
|
|
});
|
|
return customOp($logits);
|
|
}
|
|
const logSoftmax = /* @__PURE__ */ op({ logSoftmax_ });
|
|
function logSumExp_(x, axis = null, keepDims = false) {
|
|
const $x = convertToTensor(x, "x", "logSumExp");
|
|
const axes = parseAxisParam(axis, $x.shape);
|
|
const xMax = max$2(
|
|
$x,
|
|
axes,
|
|
true
|
|
/* keepDims */
|
|
);
|
|
const a = sub$2($x, xMax);
|
|
const b = exp$2(a);
|
|
const c = sum$2(b, axes);
|
|
const d = log$2(c);
|
|
const res = add$1(reshape$2(xMax, d.shape), d);
|
|
if (keepDims) {
|
|
const newShape = expandShapeToKeepDim(res.shape, axes);
|
|
return reshape$2(res, newShape);
|
|
}
|
|
return res;
|
|
}
|
|
const logSumExp = /* @__PURE__ */ op({ logSumExp_ });
|
|
function logicalAnd_(a, b) {
|
|
const $a = convertToTensor(a, "a", "logicalAnd", "bool");
|
|
const $b = convertToTensor(b, "b", "logicalAnd", "bool");
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(LogicalAnd, inputs);
|
|
}
|
|
const logicalAnd$2 = /* @__PURE__ */ op({ logicalAnd_ });
|
|
function logicalNot_(x) {
|
|
const $x = convertToTensor(x, "x", "logicalNot", "bool");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(LogicalNot, inputs);
|
|
}
|
|
const logicalNot$2 = /* @__PURE__ */ op({ logicalNot_ });
|
|
function logicalOr_(a, b) {
|
|
const $a = convertToTensor(a, "a", "logicalOr", "bool");
|
|
const $b = convertToTensor(b, "b", "logicalOr", "bool");
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(LogicalOr, inputs);
|
|
}
|
|
const logicalOr$2 = /* @__PURE__ */ op({ logicalOr_ });
|
|
function logicalXor_(a, b) {
|
|
const $a = convertToTensor(a, "a", "logicalXor", "bool");
|
|
const $b = convertToTensor(b, "b", "logicalXor", "bool");
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
return logicalAnd$2(logicalOr$2(a, b), logicalNot$2(logicalAnd$2(a, b)));
|
|
}
|
|
const logicalXor = /* @__PURE__ */ op({ logicalXor_ });
|
|
const INT32_MAX$1 = 2147483648;
|
|
function searchSorted_(sortedSequence, values, side = "left") {
|
|
const $sortedSequence = convertToTensor(sortedSequence, "sortedSequence", "searchSorted");
|
|
const $values = convertToTensor(values, "values", "searchSorted");
|
|
const sequenceSize = $sortedSequence.shape[$sortedSequence.shape.length - 1];
|
|
const valuesSize = $values.shape[$values.shape.length - 1];
|
|
const $sortedSequence2D = reshape$2($sortedSequence, [-1, sequenceSize]);
|
|
const $values2D = reshape$2($values, [-1, valuesSize]);
|
|
if ($sortedSequence2D.rank < 2) {
|
|
throw new Error(`Sorted input argument must be at least 2-dimensional`);
|
|
}
|
|
if ($sortedSequence2D.shape[0] !== $values2D.shape[0]) {
|
|
throw new Error(`Leading dimension of 'sortedSequence' and 'values' must match.`);
|
|
}
|
|
if (sizeFromShape($values2D.shape) >= INT32_MAX$1) {
|
|
throw new Error(`values tensor size must less than ${INT32_MAX$1}`);
|
|
}
|
|
if ($sortedSequence2D.shape[1] >= INT32_MAX$1) {
|
|
throw new Error(`trailing dim_size must less than ${INT32_MAX$1} for int32 output type, was ${$sortedSequence2D.shape[1]}`);
|
|
}
|
|
const inputs = {
|
|
sortedSequence: $sortedSequence2D,
|
|
values: $values2D
|
|
};
|
|
const attrs = { side };
|
|
return ENGINE.runKernel(SearchSorted, inputs, attrs);
|
|
}
|
|
const searchSorted$2 = /* @__PURE__ */ op({ searchSorted_ });
|
|
function lowerBound$1(sortedSequence, values) {
|
|
return searchSorted$2(sortedSequence, values, "left");
|
|
}
|
|
function maxPool_(x, filterSize, strides, pad2, dimRoundingMode) {
|
|
const $x = convertToTensor(x, "x", "maxPool");
|
|
const dilations = 1;
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in maxPool: input must be rank 4 but got rank ${x4D.rank}.`);
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in maxPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
checkPadOnDimRoundingMode("maxPool", pad2, dimRoundingMode);
|
|
const inputs = { x: x4D };
|
|
const attrs = { filterSize, strides, pad: pad2, dimRoundingMode };
|
|
const res = ENGINE.runKernel(MaxPool, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const maxPool$2 = /* @__PURE__ */ op({ maxPool_ });
|
|
function maxPool3d_(x, filterSize = [1, 1, 1], strides, pad2, dimRoundingMode, dataFormat = "NDHWC") {
|
|
const $x = convertToTensor(x, "x", "maxPool3d");
|
|
let x5D = $x;
|
|
let reshapedTo5D = false;
|
|
if ($x.rank === 4) {
|
|
reshapedTo5D = true;
|
|
x5D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2], $x.shape[3]]);
|
|
}
|
|
assert(x5D.rank === 5, () => `Error in maxPool3d: x must be rank 5 but got rank ${x5D.rank}.`);
|
|
assert(dataFormat === "NDHWC", () => `Error in maxPool3d: Only NDHWC is currently supported, but got dataFormat of ${dataFormat}`);
|
|
checkPadOnDimRoundingMode("maxPool3d", pad2, dimRoundingMode);
|
|
const inputs = { x: x5D };
|
|
const attrs = { filterSize, strides, pad: pad2, dimRoundingMode, dataFormat };
|
|
const res = ENGINE.runKernel(MaxPool3D, inputs, attrs);
|
|
if (reshapedTo5D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3], res.shape[4]]);
|
|
}
|
|
return res;
|
|
}
|
|
const maxPool3d$1 = /* @__PURE__ */ op({ maxPool3d_ });
|
|
function maxPoolWithArgmax_(x, filterSize, strides, pad2, includeBatchInIndex = false) {
|
|
const $x = convertToTensor(x, "x", "maxPoolWithArgmax");
|
|
const inputs = { x: $x };
|
|
const attrs = { filterSize, strides, pad: pad2, includeBatchInIndex };
|
|
const result = ENGINE.runKernel(MaxPoolWithArgmax, inputs, attrs);
|
|
return { result: result[0], indexes: result[1] };
|
|
}
|
|
const maxPoolWithArgmax = /* @__PURE__ */ op({ maxPoolWithArgmax_ });
|
|
function maximum_(a, b) {
|
|
let $a = convertToTensor(a, "a", "maximum");
|
|
let $b = convertToTensor(b, "b", "maximum");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
if ($a.dtype === "bool") {
|
|
$a = cast$2($a, "int32");
|
|
$b = cast$2($b, "int32");
|
|
}
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Maximum, inputs);
|
|
}
|
|
const maximum$2 = /* @__PURE__ */ op({ maximum_ });
|
|
function mean_(x, axis = null, keepDims = false) {
|
|
const $x = convertToTensor(x, "x", "mean");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, keepDims };
|
|
return ENGINE.runKernel(Mean, inputs, attrs);
|
|
}
|
|
const mean$1 = /* @__PURE__ */ op({ mean_ });
|
|
function zeros$1(shape, dtype = "float32") {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
if (dtype === "complex64") {
|
|
const real2 = zeros$1(shape, "float32");
|
|
const imag2 = zeros$1(shape, "float32");
|
|
return complex$2(real2, imag2);
|
|
}
|
|
const values = makeZerosTypedArray(sizeFromShape(shape), dtype);
|
|
return ENGINE.makeTensor(values, shape, dtype);
|
|
}
|
|
function ones(shape, dtype = "float32") {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
if (dtype === "complex64") {
|
|
const real2 = ones(shape, "float32");
|
|
const imag2 = zeros$1(shape, "float32");
|
|
return complex$2(real2, imag2);
|
|
}
|
|
const values = makeOnesTypedArray(sizeFromShape(shape), dtype);
|
|
return ENGINE.makeTensor(values, shape, dtype);
|
|
}
|
|
function meshgrid(x, y, { indexing = "xy" } = {}) {
|
|
if (indexing !== "xy" && indexing !== "ij") {
|
|
throw new TypeError(`${indexing} is not a valid third argument to meshgrid`);
|
|
}
|
|
if (x === void 0) {
|
|
return [];
|
|
}
|
|
let $x = convertToTensor(x, "x", "meshgrid", x instanceof Tensor ? x.dtype : "float32");
|
|
if (y === void 0) {
|
|
return [$x];
|
|
}
|
|
let $y = convertToTensor(y, "y", "meshgrid", y instanceof Tensor ? y.dtype : "float32");
|
|
const w = sizeFromShape($x.shape);
|
|
const h = sizeFromShape($y.shape);
|
|
if (indexing === "xy") {
|
|
$x = reshape$2($x, [1, -1]);
|
|
$y = reshape$2($y, [-1, 1]);
|
|
return [
|
|
matMul$1(ones([h, 1], $x.dtype), $x),
|
|
matMul$1($y, ones([1, w], $y.dtype))
|
|
];
|
|
}
|
|
$x = reshape$2($x, [-1, 1]);
|
|
$y = reshape$2($y, [1, -1]);
|
|
return [
|
|
matMul$1($x, ones([1, h], $x.dtype)),
|
|
matMul$1(ones([w, 1], $y.dtype), $y)
|
|
];
|
|
}
|
|
function minimum_(a, b) {
|
|
let $a = convertToTensor(a, "a", "minimum");
|
|
let $b = convertToTensor(b, "b", "minimum");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
if ($a.dtype === "bool") {
|
|
$a = cast$2($a, "int32");
|
|
$b = cast$2($b, "int32");
|
|
}
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Minimum, inputs);
|
|
}
|
|
const minimum$2 = /* @__PURE__ */ op({ minimum_ });
|
|
function mirrorPad_(x, paddings, mode) {
|
|
assert(mode === "reflect" || mode === "symmetric", () => `Invalid mode. Mode must be either reflect or symmetric. Got ${mode}.`);
|
|
const $x = convertToTensor(x, "x", "mirrorPad");
|
|
if ($x.rank === 0) {
|
|
throw new Error("mirrorPad(scalar) is not defined. Pass non-scalar to mirrorPad");
|
|
}
|
|
assert(paddings.length === $x.rank, () => `Padding doesn't match input. Must be ${$x.rank}. Got ${paddings.length}.`);
|
|
const shapeOffset = mode === "reflect" ? 1 : 0;
|
|
for (let i = 0; i < $x.rank; i++) {
|
|
assert(paddings[i].length === 2, () => `Invalid number of paddings. Must be length of 2 each.`);
|
|
assert(paddings[i][0] >= 0 && paddings[i][0] <= $x.shape[i] - shapeOffset && paddings[i][1] >= 0 && paddings[i][1] <= $x.shape[i] - shapeOffset, () => `Padding in dimension ${i} cannot be greater than or equal to ${$x.shape[i] - shapeOffset} or less than 0 for input of shape ${$x.shape}`);
|
|
}
|
|
const attrs = { paddings, mode };
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(MirrorPad, inputs, attrs);
|
|
}
|
|
const mirrorPad$1 = /* @__PURE__ */ op({ mirrorPad_ });
|
|
function mod_(a, b) {
|
|
let $a = convertToTensor(a, "a", "mod");
|
|
let $b = convertToTensor(b, "b", "mod");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(Mod, inputs);
|
|
}
|
|
const mod$2 = /* @__PURE__ */ op({ mod_ });
|
|
function moments_(x, axis = null, keepDims = false) {
|
|
x = convertToTensor(x, "x", "moments");
|
|
const axes = parseAxisParam(axis, x.shape);
|
|
const xMean = mean$1(x, axes, keepDims);
|
|
let keepDimsShape = xMean.shape;
|
|
if (!keepDims) {
|
|
keepDimsShape = expandShapeToKeepDim(xMean.shape, axes);
|
|
}
|
|
const devSquared = square$1(sub$2(cast$2(x, "float32"), reshape$2(xMean, keepDimsShape)));
|
|
const variance = mean$1(devSquared, axes, keepDims);
|
|
return { mean: xMean, variance };
|
|
}
|
|
const moments = /* @__PURE__ */ op({ moments_ });
|
|
function multiRNNCell_(lstmCells, data, c, h) {
|
|
const $data = convertToTensor(data, "data", "multiRNNCell");
|
|
const $c = convertToTensorArray(c, "c", "multiRNNCell");
|
|
const $h = convertToTensorArray(h, "h", "multiRNNCell");
|
|
let input = $data;
|
|
const newStates = [];
|
|
for (let i = 0; i < lstmCells.length; i++) {
|
|
const output = lstmCells[i](input, $c[i], $h[i]);
|
|
newStates.push(output[0]);
|
|
newStates.push(output[1]);
|
|
input = output[1];
|
|
}
|
|
const newC = [];
|
|
const newH = [];
|
|
for (let i = 0; i < newStates.length; i += 2) {
|
|
newC.push(newStates[i]);
|
|
newH.push(newStates[i + 1]);
|
|
}
|
|
return [newC, newH];
|
|
}
|
|
const multiRNNCell = /* @__PURE__ */ op({ multiRNNCell_ });
|
|
function multinomial_(logits, numSamples, seed, normalized = false) {
|
|
const $logits = convertToTensor(logits, "logits", "multinomial");
|
|
const numOutcomes = $logits.size;
|
|
const origRank = $logits.rank;
|
|
if (numOutcomes < 2) {
|
|
throw new Error(`Error in multinomial: you need at least 2 outcomes, but got ${numOutcomes}.`);
|
|
}
|
|
if (origRank > 2) {
|
|
throw new Error(`Rank of probabilities must be 1 or 2, but is ${origRank}`);
|
|
}
|
|
seed = seed || Math.random();
|
|
const logits2D = origRank === 1 ? reshape$2($logits, [1, -1]) : $logits;
|
|
const inputs = { logits: logits2D };
|
|
const attrs = { numSamples, seed, normalized };
|
|
const res = ENGINE.runKernel(Multinomial, inputs, attrs);
|
|
return origRank === 1 ? reshape$2(res, [res.size]) : res;
|
|
}
|
|
const multinomial$2 = /* @__PURE__ */ op({ multinomial_ });
|
|
function notEqual_(a, b) {
|
|
let $a = convertToTensor(a, "a", "notEqual", "string_or_numeric");
|
|
let $b = convertToTensor(b, "b", "notEqual", "string_or_numeric");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
return ENGINE.runKernel(NotEqual, inputs);
|
|
}
|
|
const notEqual$2 = /* @__PURE__ */ op({ notEqual_ });
|
|
function oneHot_(indices, depth, onValue = 1, offValue = 0, dtype = "int32") {
|
|
if (depth < 2) {
|
|
throw new Error(`Error in oneHot: depth must be >=2, but it is ${depth}`);
|
|
}
|
|
const $indices = convertToTensor(indices, "indices", "oneHot", "int32");
|
|
const inputs = { indices: $indices };
|
|
const attrs = { dtype, depth, onValue, offValue };
|
|
return ENGINE.runKernel(OneHot, inputs, attrs);
|
|
}
|
|
const oneHot$2 = /* @__PURE__ */ op({ oneHot_ });
|
|
function onesLike_(x) {
|
|
const $x = convertToTensor(x, "x", "onesLike");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(OnesLike, inputs);
|
|
}
|
|
const onesLike$2 = /* @__PURE__ */ op({ onesLike_ });
|
|
function outerProduct_(v1, v2) {
|
|
const $v1 = convertToTensor(v1, "v1", "outerProduct");
|
|
const $v2 = convertToTensor(v2, "v2", "outerProduct");
|
|
assert($v1.rank === 1 && $v2.rank === 1, () => `Error in outerProduct: inputs must be rank 1, but got ranks ${$v1.rank} and ${$v2.rank}.`);
|
|
const v12D = reshape$2($v1, [-1, 1]);
|
|
const v22D = reshape$2($v2, [1, -1]);
|
|
return matMul$1(v12D, v22D);
|
|
}
|
|
const outerProduct = /* @__PURE__ */ op({ outerProduct_ });
|
|
function pad_(x, paddings, constantValue = 0) {
|
|
const $x = convertToTensor(x, "x", "pad");
|
|
if ($x.rank === 0) {
|
|
throw new Error("pad(scalar) is not defined. Pass non-scalar to pad");
|
|
}
|
|
const attrs = { paddings, constantValue };
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(PadV2, inputs, attrs);
|
|
}
|
|
const pad = /* @__PURE__ */ op({ pad_ });
|
|
function pad1d_(x, paddings, constantValue = 0) {
|
|
assert(paddings.length === 2, () => "Invalid number of paddings. Must be length of 2.");
|
|
return pad(x, [paddings], constantValue);
|
|
}
|
|
const pad1d = /* @__PURE__ */ op({ pad1d_ });
|
|
function pad2d_(x, paddings, constantValue = 0) {
|
|
assert(paddings.length === 2 && paddings[0].length === 2 && paddings[1].length === 2, () => "Invalid number of paddings. Must be length of 2 each.");
|
|
return pad(x, paddings, constantValue);
|
|
}
|
|
const pad2d = /* @__PURE__ */ op({ pad2d_ });
|
|
function pad3d_(x, paddings, constantValue = 0) {
|
|
assert(paddings.length === 3 && paddings[0].length === 2 && paddings[1].length === 2 && paddings[2].length === 2, () => "Invalid number of paddings. Must be length of 2 each.");
|
|
return pad(x, paddings, constantValue);
|
|
}
|
|
const pad3d = /* @__PURE__ */ op({ pad3d_ });
|
|
function pad4d_(x, paddings, constantValue = 0) {
|
|
assert(paddings.length === 4 && paddings[0].length === 2 && paddings[1].length === 2 && paddings[2].length === 2 && paddings[3].length === 2, () => "Invalid number of paddings. Must be length of 2 each.");
|
|
return pad(x, paddings, constantValue);
|
|
}
|
|
const pad4d = /* @__PURE__ */ op({ pad4d_ });
|
|
function spaceToBatchND_(x, blockShape, paddings) {
|
|
const $x = convertToTensor(x, "x", "spaceToBatchND");
|
|
assert($x.rank >= 1 + blockShape.length, () => `input rank ${$x.rank} should be > than [blockShape] ${blockShape.length}`);
|
|
assert(paddings.length === blockShape.length, () => `paddings.shape[0] ${paddings.length} must be equal to [blockShape] ${blockShape.length}`);
|
|
assert($x.shape.reduce((a, b, i) => {
|
|
if (i > 0 && i <= blockShape.length) {
|
|
return a && (b + paddings[i - 1][0] + paddings[i - 1][1]) % blockShape[i - 1] === 0;
|
|
}
|
|
return a;
|
|
}, true), () => `input spatial dimensions ${$x.shape.slice(1)} with paddings ${paddings.toString()} must be divisible by blockShapes ${blockShape.toString()}`);
|
|
const inputs = { x: $x };
|
|
const attrs = { blockShape, paddings };
|
|
return ENGINE.runKernel(SpaceToBatchND, inputs, attrs);
|
|
}
|
|
const spaceToBatchND$2 = /* @__PURE__ */ op({ spaceToBatchND_ });
|
|
function pool_(input, windowShape, poolingType, pad2, dilations, strides, dimRoundingMode) {
|
|
if (dilations == null) {
|
|
dilations = [1, 1];
|
|
}
|
|
if (strides == null) {
|
|
strides = 1;
|
|
}
|
|
if (pad2 === 0) {
|
|
pad2 = "valid";
|
|
}
|
|
const $x = convertToTensor(input, "x", "maxPool");
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in pool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computePool2DInfo(x4D.shape, windowShape, strides, dilations, pad2);
|
|
const dilation = [convInfo.dilationHeight, convInfo.dilationWidth];
|
|
let basePadding;
|
|
if (pad2 === "same") {
|
|
basePadding = withSpaceToBatchBasePaddings([convInfo.filterHeight, convInfo.filterWidth], dilation);
|
|
} else {
|
|
basePadding = [[0, 0], [0, 0]];
|
|
}
|
|
const isDilationOne = dilation[0] === 1 && dilation[1] === 1;
|
|
const [adjustedPadding, adjustedCrops] = requiredSpaceToBatchPaddings([convInfo.inHeight, convInfo.inWidth], dilation, basePadding);
|
|
const convertedPad = isDilationOne ? pad2 : "valid";
|
|
const convertedX = isDilationOne ? x4D : spaceToBatchND$2(x4D, dilation, adjustedPadding);
|
|
const forwardOp = poolingType === "avg" ? () => avgPool$2(convertedX, windowShape, strides, convertedPad, dimRoundingMode) : () => maxPool$2(convertedX, windowShape, strides, convertedPad, dimRoundingMode);
|
|
const y = forwardOp();
|
|
const res = isDilationOne ? y : batchToSpaceND$2(y, dilation, adjustedCrops);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
function requiredSpaceToBatchPaddings(inputShape, blockShape, basePadding) {
|
|
const padStart = basePadding.map((b) => b[0]);
|
|
const origPadEnd = basePadding.map((b) => b[1]);
|
|
const fullInputShape = inputShape.concat(padStart, origPadEnd);
|
|
const padEndExtra = blockShape.map((b, i) => (b - fullInputShape[i] % b) % b);
|
|
const padEnd = origPadEnd.map((s, i) => s + padEndExtra[i]);
|
|
const paddings = blockShape.map((_, i) => [padStart[i], padEnd[i]]);
|
|
const crops = blockShape.map((_, i) => [0, padEndExtra[i]]);
|
|
return [paddings, crops];
|
|
}
|
|
function withSpaceToBatchBasePaddings(filterShape, dilation) {
|
|
const dilatedFilterShape = filterShape.map((s, i) => {
|
|
return s + (s - 1) * (dilation[i] - 1);
|
|
});
|
|
const padExtraShape = dilatedFilterShape.map((s) => s - 1);
|
|
const padExtraStart = padExtraShape.map((s) => Math.floor(s / 2));
|
|
const padExtraEnd = padExtraShape.map((s, i) => s - padExtraStart[i]);
|
|
return padExtraShape.map((_, i) => {
|
|
return [padExtraStart[i], padExtraEnd[i]];
|
|
});
|
|
}
|
|
const pool$1 = /* @__PURE__ */ op({ pool_ });
|
|
function prelu_(x, alpha) {
|
|
const $x = convertToTensor(x, "x", "prelu");
|
|
const $alpha = convertToTensor(alpha, "alpha", "prelu");
|
|
const inputs = { x: $x, alpha: $alpha };
|
|
return ENGINE.runKernel(Prelu, inputs);
|
|
}
|
|
const prelu$2 = /* @__PURE__ */ op({ prelu_ });
|
|
function prod_(x, axis = null, keepDims = false) {
|
|
let $x = convertToTensor(x, "x", "prod");
|
|
if ($x.dtype === "bool") {
|
|
$x = cast$2($x, "int32");
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { axis, keepDims };
|
|
return ENGINE.runKernel(Prod, inputs, attrs);
|
|
}
|
|
const prod$2 = /* @__PURE__ */ op({ prod_ });
|
|
function raggedGather_(paramsNestedSplits, paramsDenseValues, indices, outputRaggedRank) {
|
|
const $paramsNestedSplits = paramsNestedSplits.map((t2, i) => convertToTensor(t2, `tensors${i}`, "raggedGather", "int32"));
|
|
const $paramsDenseValues = convertToTensor(paramsDenseValues, "paramsDenseValues", "raggedGather");
|
|
const $indices = convertToTensor(indices, "indices", "raggedGather", "int32");
|
|
const inputs = {
|
|
paramsNestedSplits: $paramsNestedSplits,
|
|
paramsDenseValues: $paramsDenseValues,
|
|
indices: $indices
|
|
};
|
|
const attrs = { outputRaggedRank };
|
|
const result = ENGINE.runKernel(RaggedGather, inputs, attrs);
|
|
return {
|
|
outputNestedSplits: result.slice(0, result.length - 1),
|
|
outputDenseValues: result[result.length - 1]
|
|
};
|
|
}
|
|
const raggedGather$2 = /* @__PURE__ */ op({ raggedGather_ });
|
|
function raggedRange_(starts, limits, deltas) {
|
|
const $starts = convertToTensor(starts, "starts", "raggedRange");
|
|
const $limits = convertToTensor(limits, "limits", "raggedRange", $starts.dtype);
|
|
const $deltas = convertToTensor(deltas, "deltas", "raggedRange", $starts.dtype);
|
|
const inputs = {
|
|
starts: $starts,
|
|
limits: $limits,
|
|
deltas: $deltas
|
|
};
|
|
const result = ENGINE.runKernel(RaggedRange, inputs);
|
|
return {
|
|
rtNestedSplits: result[0],
|
|
rtDenseValues: result[1]
|
|
};
|
|
}
|
|
const raggedRange$2 = /* @__PURE__ */ op({ raggedRange_ });
|
|
function raggedTensorToTensor_(shape, values, defaultValue, rowPartitionTensors, rowPartitionTypes) {
|
|
const $shape = convertToTensor(shape, "shape", "raggedTensorToTensor", "int32");
|
|
const $values = convertToTensor(values, "values", "raggedTensorToTensor");
|
|
const $defaultValue = convertToTensor(defaultValue, "defaultValue", "raggedTensorToTensor", $values.dtype);
|
|
const $rowPartitionTensors = rowPartitionTensors.map((t2, i) => convertToTensor(t2, `tensors${i}`, "raggedTensorToTensor", "int32"));
|
|
const inputs = {
|
|
shape: $shape,
|
|
values: $values,
|
|
defaultValue: $defaultValue,
|
|
rowPartitionTensors: $rowPartitionTensors
|
|
};
|
|
const attrs = { rowPartitionTypes };
|
|
return ENGINE.runKernel(RaggedTensorToTensor, inputs, attrs);
|
|
}
|
|
const raggedTensorToTensor$2 = /* @__PURE__ */ op({ raggedTensorToTensor_ });
|
|
function rand_(shape, randFunction, dtype) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
const size = sizeFromShape(shape);
|
|
let values = null;
|
|
if (dtype == null || dtype === "float32") {
|
|
values = new Float32Array(size);
|
|
} else if (dtype === "int32") {
|
|
values = new Int32Array(size);
|
|
} else if (dtype === "bool") {
|
|
values = new Uint8Array(size);
|
|
} else {
|
|
throw new Error(`Unknown data type ${dtype}`);
|
|
}
|
|
for (let i = 0; i < size; i++) {
|
|
values[i] = randFunction();
|
|
}
|
|
return ENGINE.makeTensor(values, shape, dtype);
|
|
}
|
|
const rand = /* @__PURE__ */ op({ rand_ });
|
|
var alea$1 = { exports: {} };
|
|
var alea = alea$1.exports;
|
|
var hasRequiredAlea;
|
|
function requireAlea() {
|
|
if (hasRequiredAlea) return alea$1.exports;
|
|
hasRequiredAlea = 1;
|
|
(function(module) {
|
|
(function(global2, module2, define) {
|
|
function Alea(seed) {
|
|
var me2 = this, mash = Mash();
|
|
me2.next = function() {
|
|
var t2 = 2091639 * me2.s0 + me2.c * 23283064365386963e-26;
|
|
me2.s0 = me2.s1;
|
|
me2.s1 = me2.s2;
|
|
return me2.s2 = t2 - (me2.c = t2 | 0);
|
|
};
|
|
me2.c = 1;
|
|
me2.s0 = mash(" ");
|
|
me2.s1 = mash(" ");
|
|
me2.s2 = mash(" ");
|
|
me2.s0 -= mash(seed);
|
|
if (me2.s0 < 0) {
|
|
me2.s0 += 1;
|
|
}
|
|
me2.s1 -= mash(seed);
|
|
if (me2.s1 < 0) {
|
|
me2.s1 += 1;
|
|
}
|
|
me2.s2 -= mash(seed);
|
|
if (me2.s2 < 0) {
|
|
me2.s2 += 1;
|
|
}
|
|
mash = null;
|
|
}
|
|
function copy(f, t2) {
|
|
t2.c = f.c;
|
|
t2.s0 = f.s0;
|
|
t2.s1 = f.s1;
|
|
t2.s2 = f.s2;
|
|
return t2;
|
|
}
|
|
function impl(seed, opts) {
|
|
var xg = new Alea(seed), state = opts && opts.state, prng = xg.next;
|
|
prng.int32 = function() {
|
|
return xg.next() * 4294967296 | 0;
|
|
};
|
|
prng.double = function() {
|
|
return prng() + (prng() * 2097152 | 0) * 11102230246251565e-32;
|
|
};
|
|
prng.quick = prng;
|
|
if (state) {
|
|
if (typeof state == "object") copy(state, xg);
|
|
prng.state = function() {
|
|
return copy(xg, {});
|
|
};
|
|
}
|
|
return prng;
|
|
}
|
|
function Mash() {
|
|
var n = 4022871197;
|
|
var mash = function(data) {
|
|
data = String(data);
|
|
for (var i = 0; i < data.length; i++) {
|
|
n += data.charCodeAt(i);
|
|
var h = 0.02519603282416938 * n;
|
|
n = h >>> 0;
|
|
h -= n;
|
|
h *= n;
|
|
n = h >>> 0;
|
|
h -= n;
|
|
n += h * 4294967296;
|
|
}
|
|
return (n >>> 0) * 23283064365386963e-26;
|
|
};
|
|
return mash;
|
|
}
|
|
if (module2 && module2.exports) {
|
|
module2.exports = impl;
|
|
} else {
|
|
this.alea = impl;
|
|
}
|
|
})(
|
|
alea,
|
|
module
|
|
);
|
|
})(alea$1);
|
|
return alea$1.exports;
|
|
}
|
|
var xor128$1 = { exports: {} };
|
|
var xor128 = xor128$1.exports;
|
|
var hasRequiredXor128;
|
|
function requireXor128() {
|
|
if (hasRequiredXor128) return xor128$1.exports;
|
|
hasRequiredXor128 = 1;
|
|
(function(module) {
|
|
(function(global2, module2, define) {
|
|
function XorGen(seed) {
|
|
var me2 = this, strseed = "";
|
|
me2.x = 0;
|
|
me2.y = 0;
|
|
me2.z = 0;
|
|
me2.w = 0;
|
|
me2.next = function() {
|
|
var t2 = me2.x ^ me2.x << 11;
|
|
me2.x = me2.y;
|
|
me2.y = me2.z;
|
|
me2.z = me2.w;
|
|
return me2.w ^= me2.w >>> 19 ^ t2 ^ t2 >>> 8;
|
|
};
|
|
if (seed === (seed | 0)) {
|
|
me2.x = seed;
|
|
} else {
|
|
strseed += seed;
|
|
}
|
|
for (var k = 0; k < strseed.length + 64; k++) {
|
|
me2.x ^= strseed.charCodeAt(k) | 0;
|
|
me2.next();
|
|
}
|
|
}
|
|
function copy(f, t2) {
|
|
t2.x = f.x;
|
|
t2.y = f.y;
|
|
t2.z = f.z;
|
|
t2.w = f.w;
|
|
return t2;
|
|
}
|
|
function impl(seed, opts) {
|
|
var xg = new XorGen(seed), state = opts && opts.state, prng = function() {
|
|
return (xg.next() >>> 0) / 4294967296;
|
|
};
|
|
prng.double = function() {
|
|
do {
|
|
var top = xg.next() >>> 11, bot = (xg.next() >>> 0) / 4294967296, result = (top + bot) / (1 << 21);
|
|
} while (result === 0);
|
|
return result;
|
|
};
|
|
prng.int32 = xg.next;
|
|
prng.quick = prng;
|
|
if (state) {
|
|
if (typeof state == "object") copy(state, xg);
|
|
prng.state = function() {
|
|
return copy(xg, {});
|
|
};
|
|
}
|
|
return prng;
|
|
}
|
|
if (module2 && module2.exports) {
|
|
module2.exports = impl;
|
|
} else {
|
|
this.xor128 = impl;
|
|
}
|
|
})(
|
|
xor128,
|
|
module
|
|
);
|
|
})(xor128$1);
|
|
return xor128$1.exports;
|
|
}
|
|
var xorwow$1 = { exports: {} };
|
|
var xorwow = xorwow$1.exports;
|
|
var hasRequiredXorwow;
|
|
function requireXorwow() {
|
|
if (hasRequiredXorwow) return xorwow$1.exports;
|
|
hasRequiredXorwow = 1;
|
|
(function(module) {
|
|
(function(global2, module2, define) {
|
|
function XorGen(seed) {
|
|
var me2 = this, strseed = "";
|
|
me2.next = function() {
|
|
var t2 = me2.x ^ me2.x >>> 2;
|
|
me2.x = me2.y;
|
|
me2.y = me2.z;
|
|
me2.z = me2.w;
|
|
me2.w = me2.v;
|
|
return (me2.d = me2.d + 362437 | 0) + (me2.v = me2.v ^ me2.v << 4 ^ (t2 ^ t2 << 1)) | 0;
|
|
};
|
|
me2.x = 0;
|
|
me2.y = 0;
|
|
me2.z = 0;
|
|
me2.w = 0;
|
|
me2.v = 0;
|
|
if (seed === (seed | 0)) {
|
|
me2.x = seed;
|
|
} else {
|
|
strseed += seed;
|
|
}
|
|
for (var k = 0; k < strseed.length + 64; k++) {
|
|
me2.x ^= strseed.charCodeAt(k) | 0;
|
|
if (k == strseed.length) {
|
|
me2.d = me2.x << 10 ^ me2.x >>> 4;
|
|
}
|
|
me2.next();
|
|
}
|
|
}
|
|
function copy(f, t2) {
|
|
t2.x = f.x;
|
|
t2.y = f.y;
|
|
t2.z = f.z;
|
|
t2.w = f.w;
|
|
t2.v = f.v;
|
|
t2.d = f.d;
|
|
return t2;
|
|
}
|
|
function impl(seed, opts) {
|
|
var xg = new XorGen(seed), state = opts && opts.state, prng = function() {
|
|
return (xg.next() >>> 0) / 4294967296;
|
|
};
|
|
prng.double = function() {
|
|
do {
|
|
var top = xg.next() >>> 11, bot = (xg.next() >>> 0) / 4294967296, result = (top + bot) / (1 << 21);
|
|
} while (result === 0);
|
|
return result;
|
|
};
|
|
prng.int32 = xg.next;
|
|
prng.quick = prng;
|
|
if (state) {
|
|
if (typeof state == "object") copy(state, xg);
|
|
prng.state = function() {
|
|
return copy(xg, {});
|
|
};
|
|
}
|
|
return prng;
|
|
}
|
|
if (module2 && module2.exports) {
|
|
module2.exports = impl;
|
|
} else {
|
|
this.xorwow = impl;
|
|
}
|
|
})(
|
|
xorwow,
|
|
module
|
|
);
|
|
})(xorwow$1);
|
|
return xorwow$1.exports;
|
|
}
|
|
var xorshift7$1 = { exports: {} };
|
|
var xorshift7 = xorshift7$1.exports;
|
|
var hasRequiredXorshift7;
|
|
function requireXorshift7() {
|
|
if (hasRequiredXorshift7) return xorshift7$1.exports;
|
|
hasRequiredXorshift7 = 1;
|
|
(function(module) {
|
|
(function(global2, module2, define) {
|
|
function XorGen(seed) {
|
|
var me2 = this;
|
|
me2.next = function() {
|
|
var X2 = me2.x, i = me2.i, t2, v;
|
|
t2 = X2[i];
|
|
t2 ^= t2 >>> 7;
|
|
v = t2 ^ t2 << 24;
|
|
t2 = X2[i + 1 & 7];
|
|
v ^= t2 ^ t2 >>> 10;
|
|
t2 = X2[i + 3 & 7];
|
|
v ^= t2 ^ t2 >>> 3;
|
|
t2 = X2[i + 4 & 7];
|
|
v ^= t2 ^ t2 << 7;
|
|
t2 = X2[i + 7 & 7];
|
|
t2 = t2 ^ t2 << 13;
|
|
v ^= t2 ^ t2 << 9;
|
|
X2[i] = v;
|
|
me2.i = i + 1 & 7;
|
|
return v;
|
|
};
|
|
function init(me3, seed2) {
|
|
var j2, X2 = [];
|
|
if (seed2 === (seed2 | 0)) {
|
|
X2[0] = seed2;
|
|
} else {
|
|
seed2 = "" + seed2;
|
|
for (j2 = 0; j2 < seed2.length; ++j2) {
|
|
X2[j2 & 7] = X2[j2 & 7] << 15 ^ seed2.charCodeAt(j2) + X2[j2 + 1 & 7] << 13;
|
|
}
|
|
}
|
|
while (X2.length < 8) X2.push(0);
|
|
for (j2 = 0; j2 < 8 && X2[j2] === 0; ++j2) ;
|
|
if (j2 == 8) X2[7] = -1;
|
|
else X2[j2];
|
|
me3.x = X2;
|
|
me3.i = 0;
|
|
for (j2 = 256; j2 > 0; --j2) {
|
|
me3.next();
|
|
}
|
|
}
|
|
init(me2, seed);
|
|
}
|
|
function copy(f, t2) {
|
|
t2.x = f.x.slice();
|
|
t2.i = f.i;
|
|
return t2;
|
|
}
|
|
function impl(seed, opts) {
|
|
if (seed == null) seed = +/* @__PURE__ */ new Date();
|
|
var xg = new XorGen(seed), state = opts && opts.state, prng = function() {
|
|
return (xg.next() >>> 0) / 4294967296;
|
|
};
|
|
prng.double = function() {
|
|
do {
|
|
var top = xg.next() >>> 11, bot = (xg.next() >>> 0) / 4294967296, result = (top + bot) / (1 << 21);
|
|
} while (result === 0);
|
|
return result;
|
|
};
|
|
prng.int32 = xg.next;
|
|
prng.quick = prng;
|
|
if (state) {
|
|
if (state.x) copy(state, xg);
|
|
prng.state = function() {
|
|
return copy(xg, {});
|
|
};
|
|
}
|
|
return prng;
|
|
}
|
|
if (module2 && module2.exports) {
|
|
module2.exports = impl;
|
|
} else {
|
|
this.xorshift7 = impl;
|
|
}
|
|
})(
|
|
xorshift7,
|
|
module
|
|
);
|
|
})(xorshift7$1);
|
|
return xorshift7$1.exports;
|
|
}
|
|
var xor4096$1 = { exports: {} };
|
|
var xor4096 = xor4096$1.exports;
|
|
var hasRequiredXor4096;
|
|
function requireXor4096() {
|
|
if (hasRequiredXor4096) return xor4096$1.exports;
|
|
hasRequiredXor4096 = 1;
|
|
(function(module) {
|
|
(function(global2, module2, define) {
|
|
function XorGen(seed) {
|
|
var me2 = this;
|
|
me2.next = function() {
|
|
var w = me2.w, X2 = me2.X, i = me2.i, t2, v;
|
|
me2.w = w = w + 1640531527 | 0;
|
|
v = X2[i + 34 & 127];
|
|
t2 = X2[i = i + 1 & 127];
|
|
v ^= v << 13;
|
|
t2 ^= t2 << 17;
|
|
v ^= v >>> 15;
|
|
t2 ^= t2 >>> 12;
|
|
v = X2[i] = v ^ t2;
|
|
me2.i = i;
|
|
return v + (w ^ w >>> 16) | 0;
|
|
};
|
|
function init(me3, seed2) {
|
|
var t2, v, i, j2, w, X2 = [], limit = 128;
|
|
if (seed2 === (seed2 | 0)) {
|
|
v = seed2;
|
|
seed2 = null;
|
|
} else {
|
|
seed2 = seed2 + "\0";
|
|
v = 0;
|
|
limit = Math.max(limit, seed2.length);
|
|
}
|
|
for (i = 0, j2 = -32; j2 < limit; ++j2) {
|
|
if (seed2) v ^= seed2.charCodeAt((j2 + 32) % seed2.length);
|
|
if (j2 === 0) w = v;
|
|
v ^= v << 10;
|
|
v ^= v >>> 15;
|
|
v ^= v << 4;
|
|
v ^= v >>> 13;
|
|
if (j2 >= 0) {
|
|
w = w + 1640531527 | 0;
|
|
t2 = X2[j2 & 127] ^= v + w;
|
|
i = 0 == t2 ? i + 1 : 0;
|
|
}
|
|
}
|
|
if (i >= 128) {
|
|
X2[(seed2 && seed2.length || 0) & 127] = -1;
|
|
}
|
|
i = 127;
|
|
for (j2 = 4 * 128; j2 > 0; --j2) {
|
|
v = X2[i + 34 & 127];
|
|
t2 = X2[i = i + 1 & 127];
|
|
v ^= v << 13;
|
|
t2 ^= t2 << 17;
|
|
v ^= v >>> 15;
|
|
t2 ^= t2 >>> 12;
|
|
X2[i] = v ^ t2;
|
|
}
|
|
me3.w = w;
|
|
me3.X = X2;
|
|
me3.i = i;
|
|
}
|
|
init(me2, seed);
|
|
}
|
|
function copy(f, t2) {
|
|
t2.i = f.i;
|
|
t2.w = f.w;
|
|
t2.X = f.X.slice();
|
|
return t2;
|
|
}
|
|
function impl(seed, opts) {
|
|
if (seed == null) seed = +/* @__PURE__ */ new Date();
|
|
var xg = new XorGen(seed), state = opts && opts.state, prng = function() {
|
|
return (xg.next() >>> 0) / 4294967296;
|
|
};
|
|
prng.double = function() {
|
|
do {
|
|
var top = xg.next() >>> 11, bot = (xg.next() >>> 0) / 4294967296, result = (top + bot) / (1 << 21);
|
|
} while (result === 0);
|
|
return result;
|
|
};
|
|
prng.int32 = xg.next;
|
|
prng.quick = prng;
|
|
if (state) {
|
|
if (state.X) copy(state, xg);
|
|
prng.state = function() {
|
|
return copy(xg, {});
|
|
};
|
|
}
|
|
return prng;
|
|
}
|
|
if (module2 && module2.exports) {
|
|
module2.exports = impl;
|
|
} else {
|
|
this.xor4096 = impl;
|
|
}
|
|
})(
|
|
xor4096,
|
|
// window object or global
|
|
module
|
|
);
|
|
})(xor4096$1);
|
|
return xor4096$1.exports;
|
|
}
|
|
var tychei$1 = { exports: {} };
|
|
var tychei = tychei$1.exports;
|
|
var hasRequiredTychei;
|
|
function requireTychei() {
|
|
if (hasRequiredTychei) return tychei$1.exports;
|
|
hasRequiredTychei = 1;
|
|
(function(module) {
|
|
(function(global2, module2, define) {
|
|
function XorGen(seed) {
|
|
var me2 = this, strseed = "";
|
|
me2.next = function() {
|
|
var b = me2.b, c = me2.c, d = me2.d, a = me2.a;
|
|
b = b << 25 ^ b >>> 7 ^ c;
|
|
c = c - d | 0;
|
|
d = d << 24 ^ d >>> 8 ^ a;
|
|
a = a - b | 0;
|
|
me2.b = b = b << 20 ^ b >>> 12 ^ c;
|
|
me2.c = c = c - d | 0;
|
|
me2.d = d << 16 ^ c >>> 16 ^ a;
|
|
return me2.a = a - b | 0;
|
|
};
|
|
me2.a = 0;
|
|
me2.b = 0;
|
|
me2.c = 2654435769 | 0;
|
|
me2.d = 1367130551;
|
|
if (seed === Math.floor(seed)) {
|
|
me2.a = seed / 4294967296 | 0;
|
|
me2.b = seed | 0;
|
|
} else {
|
|
strseed += seed;
|
|
}
|
|
for (var k = 0; k < strseed.length + 20; k++) {
|
|
me2.b ^= strseed.charCodeAt(k) | 0;
|
|
me2.next();
|
|
}
|
|
}
|
|
function copy(f, t2) {
|
|
t2.a = f.a;
|
|
t2.b = f.b;
|
|
t2.c = f.c;
|
|
t2.d = f.d;
|
|
return t2;
|
|
}
|
|
function impl(seed, opts) {
|
|
var xg = new XorGen(seed), state = opts && opts.state, prng = function() {
|
|
return (xg.next() >>> 0) / 4294967296;
|
|
};
|
|
prng.double = function() {
|
|
do {
|
|
var top = xg.next() >>> 11, bot = (xg.next() >>> 0) / 4294967296, result = (top + bot) / (1 << 21);
|
|
} while (result === 0);
|
|
return result;
|
|
};
|
|
prng.int32 = xg.next;
|
|
prng.quick = prng;
|
|
if (state) {
|
|
if (typeof state == "object") copy(state, xg);
|
|
prng.state = function() {
|
|
return copy(xg, {});
|
|
};
|
|
}
|
|
return prng;
|
|
}
|
|
if (module2 && module2.exports) {
|
|
module2.exports = impl;
|
|
} else {
|
|
this.tychei = impl;
|
|
}
|
|
})(
|
|
tychei,
|
|
module
|
|
);
|
|
})(tychei$1);
|
|
return tychei$1.exports;
|
|
}
|
|
var seedrandom$2 = { exports: {} };
|
|
const __viteBrowserExternal = {};
|
|
const __viteBrowserExternal$1 = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
default: __viteBrowserExternal
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const require$$0 = /* @__PURE__ */ getAugmentedNamespace(__viteBrowserExternal$1);
|
|
var seedrandom$1 = seedrandom$2.exports;
|
|
var hasRequiredSeedrandom$1;
|
|
function requireSeedrandom$1() {
|
|
if (hasRequiredSeedrandom$1) return seedrandom$2.exports;
|
|
hasRequiredSeedrandom$1 = 1;
|
|
(function(module) {
|
|
(function(global2, pool2, math) {
|
|
var width = 256, chunks = 6, digits = 52, rngname = "random", startdenom = math.pow(width, chunks), significance = math.pow(2, digits), overflow = significance * 2, mask = width - 1, nodecrypto;
|
|
function seedrandom2(seed, options, callback) {
|
|
var key = [];
|
|
options = options == true ? { entropy: true } : options || {};
|
|
var shortseed = mixkey(flatten2(
|
|
options.entropy ? [seed, tostring(pool2)] : seed == null ? autoseed() : seed,
|
|
3
|
|
), key);
|
|
var arc4 = new ARC4(key);
|
|
var prng = function() {
|
|
var n = arc4.g(chunks), d = startdenom, x = 0;
|
|
while (n < significance) {
|
|
n = (n + x) * width;
|
|
d *= width;
|
|
x = arc4.g(1);
|
|
}
|
|
while (n >= overflow) {
|
|
n /= 2;
|
|
d /= 2;
|
|
x >>>= 1;
|
|
}
|
|
return (n + x) / d;
|
|
};
|
|
prng.int32 = function() {
|
|
return arc4.g(4) | 0;
|
|
};
|
|
prng.quick = function() {
|
|
return arc4.g(4) / 4294967296;
|
|
};
|
|
prng.double = prng;
|
|
mixkey(tostring(arc4.S), pool2);
|
|
return (options.pass || callback || function(prng2, seed2, is_math_call, state) {
|
|
if (state) {
|
|
if (state.S) {
|
|
copy(state, arc4);
|
|
}
|
|
prng2.state = function() {
|
|
return copy(arc4, {});
|
|
};
|
|
}
|
|
if (is_math_call) {
|
|
math[rngname] = prng2;
|
|
return seed2;
|
|
} else return prng2;
|
|
})(
|
|
prng,
|
|
shortseed,
|
|
"global" in options ? options.global : this == math,
|
|
options.state
|
|
);
|
|
}
|
|
function ARC4(key) {
|
|
var t2, keylen = key.length, me2 = this, i = 0, j2 = me2.i = me2.j = 0, s = me2.S = [];
|
|
if (!keylen) {
|
|
key = [keylen++];
|
|
}
|
|
while (i < width) {
|
|
s[i] = i++;
|
|
}
|
|
for (i = 0; i < width; i++) {
|
|
s[i] = s[j2 = mask & j2 + key[i % keylen] + (t2 = s[i])];
|
|
s[j2] = t2;
|
|
}
|
|
(me2.g = function(count) {
|
|
var t3, r = 0, i2 = me2.i, j3 = me2.j, s2 = me2.S;
|
|
while (count--) {
|
|
t3 = s2[i2 = mask & i2 + 1];
|
|
r = r * width + s2[mask & (s2[i2] = s2[j3 = mask & j3 + t3]) + (s2[j3] = t3)];
|
|
}
|
|
me2.i = i2;
|
|
me2.j = j3;
|
|
return r;
|
|
})(width);
|
|
}
|
|
function copy(f, t2) {
|
|
t2.i = f.i;
|
|
t2.j = f.j;
|
|
t2.S = f.S.slice();
|
|
return t2;
|
|
}
|
|
function flatten2(obj, depth) {
|
|
var result = [], typ = typeof obj, prop;
|
|
if (depth && typ == "object") {
|
|
for (prop in obj) {
|
|
try {
|
|
result.push(flatten2(obj[prop], depth - 1));
|
|
} catch (e) {
|
|
}
|
|
}
|
|
}
|
|
return result.length ? result : typ == "string" ? obj : obj + "\0";
|
|
}
|
|
function mixkey(seed, key) {
|
|
var stringseed = seed + "", smear, j2 = 0;
|
|
while (j2 < stringseed.length) {
|
|
key[mask & j2] = mask & (smear ^= key[mask & j2] * 19) + stringseed.charCodeAt(j2++);
|
|
}
|
|
return tostring(key);
|
|
}
|
|
function autoseed() {
|
|
try {
|
|
var out;
|
|
if (nodecrypto && (out = nodecrypto.randomBytes)) {
|
|
out = out(width);
|
|
} else {
|
|
out = new Uint8Array(width);
|
|
(global2.crypto || global2.msCrypto).getRandomValues(out);
|
|
}
|
|
return tostring(out);
|
|
} catch (e) {
|
|
var browser = global2.navigator, plugins = browser && browser.plugins;
|
|
return [+/* @__PURE__ */ new Date(), global2, plugins, global2.screen, tostring(pool2)];
|
|
}
|
|
}
|
|
function tostring(a) {
|
|
return String.fromCharCode.apply(0, a);
|
|
}
|
|
mixkey(math.random(), pool2);
|
|
if (module.exports) {
|
|
module.exports = seedrandom2;
|
|
try {
|
|
nodecrypto = require$$0;
|
|
} catch (ex) {
|
|
}
|
|
} else {
|
|
math["seed" + rngname] = seedrandom2;
|
|
}
|
|
})(
|
|
// global: `self` in browsers (including strict mode and web workers),
|
|
// otherwise `this` in Node and other environments
|
|
typeof self !== "undefined" ? self : seedrandom$1,
|
|
[],
|
|
// pool: entropy pool starts empty
|
|
Math
|
|
// math: package containing random, pow, and seedrandom
|
|
);
|
|
})(seedrandom$2);
|
|
return seedrandom$2.exports;
|
|
}
|
|
var seedrandom;
|
|
var hasRequiredSeedrandom;
|
|
function requireSeedrandom() {
|
|
if (hasRequiredSeedrandom) return seedrandom;
|
|
hasRequiredSeedrandom = 1;
|
|
var alea2 = requireAlea();
|
|
var xor1282 = requireXor128();
|
|
var xorwow2 = requireXorwow();
|
|
var xorshift72 = requireXorshift7();
|
|
var xor40962 = requireXor4096();
|
|
var tychei2 = requireTychei();
|
|
var sr = requireSeedrandom$1();
|
|
sr.alea = alea2;
|
|
sr.xor128 = xor1282;
|
|
sr.xorwow = xorwow2;
|
|
sr.xorshift7 = xorshift72;
|
|
sr.xor4096 = xor40962;
|
|
sr.tychei = tychei2;
|
|
seedrandom = sr;
|
|
return seedrandom;
|
|
}
|
|
var seedrandomExports = requireSeedrandom();
|
|
class MPRandGauss {
|
|
constructor(mean2, stdDeviation, dtype, truncated, seed) {
|
|
this.mean = mean2;
|
|
this.stdDev = stdDeviation;
|
|
this.dtype = dtype;
|
|
this.nextVal = NaN;
|
|
this.truncated = truncated;
|
|
if (this.truncated) {
|
|
this.upper = this.mean + this.stdDev * 2;
|
|
this.lower = this.mean - this.stdDev * 2;
|
|
}
|
|
const seedValue = seed ? seed : Math.random();
|
|
this.random = seedrandomExports.alea(seedValue.toString());
|
|
}
|
|
/** Returns next sample from a Gaussian distribution. */
|
|
nextValue() {
|
|
if (!isNaN(this.nextVal)) {
|
|
const value = this.nextVal;
|
|
this.nextVal = NaN;
|
|
return value;
|
|
}
|
|
let resultX, resultY;
|
|
let isValid = false;
|
|
while (!isValid) {
|
|
let v1, v2, s;
|
|
do {
|
|
v1 = 2 * this.random() - 1;
|
|
v2 = 2 * this.random() - 1;
|
|
s = v1 * v1 + v2 * v2;
|
|
} while (s >= 1 || s === 0);
|
|
const mul2 = Math.sqrt(-2 * Math.log(s) / s);
|
|
resultX = this.mean + this.stdDev * v1 * mul2;
|
|
resultY = this.mean + this.stdDev * v2 * mul2;
|
|
if (!this.truncated || this.isValidTruncated(resultX)) {
|
|
isValid = true;
|
|
}
|
|
}
|
|
if (!this.truncated || this.isValidTruncated(resultY)) {
|
|
this.nextVal = this.convertValue(resultY);
|
|
}
|
|
return this.convertValue(resultX);
|
|
}
|
|
/** Handles proper rounding for non-floating-point numbers. */
|
|
convertValue(value) {
|
|
if (this.dtype == null || this.dtype === "float32") {
|
|
return value;
|
|
}
|
|
return Math.round(value);
|
|
}
|
|
/** Returns true if less than 2-standard-deviations from the mean. */
|
|
isValidTruncated(value) {
|
|
return value <= this.upper && value >= this.lower;
|
|
}
|
|
}
|
|
class RandGamma {
|
|
constructor(alpha, beta, dtype, seed) {
|
|
this.alpha = alpha;
|
|
this.beta = 1 / beta;
|
|
this.dtype = dtype;
|
|
const seedValue = seed ? seed : Math.random();
|
|
this.randu = seedrandomExports.alea(seedValue.toString());
|
|
this.randn = new MPRandGauss(0, 1, dtype, false, this.randu());
|
|
if (alpha < 1) {
|
|
this.d = alpha + 2 / 3;
|
|
} else {
|
|
this.d = alpha - 1 / 3;
|
|
}
|
|
this.c = 1 / Math.sqrt(9 * this.d);
|
|
}
|
|
/** Returns next sample from a gamma distribution. */
|
|
nextValue() {
|
|
let x2, v0, v1, x, u, v;
|
|
while (true) {
|
|
do {
|
|
x = this.randn.nextValue();
|
|
v = 1 + this.c * x;
|
|
} while (v <= 0);
|
|
v *= v * v;
|
|
x2 = x * x;
|
|
v0 = 1 - 0.331 * x2 * x2;
|
|
v1 = 0.5 * x2 + this.d * (1 - v + Math.log(v));
|
|
u = this.randu();
|
|
if (u < v0 || Math.log(u) < v1) {
|
|
break;
|
|
}
|
|
}
|
|
v = 1 / this.beta * this.d * v;
|
|
if (this.alpha < 1) {
|
|
v *= Math.pow(this.randu(), 1 / this.alpha);
|
|
}
|
|
return this.convertValue(v);
|
|
}
|
|
/** Handles proper rounding for non-floating-point numbers. */
|
|
convertValue(value) {
|
|
if (this.dtype === "float32") {
|
|
return value;
|
|
}
|
|
return Math.round(value);
|
|
}
|
|
}
|
|
class UniformRandom {
|
|
constructor(min2 = 0, max2 = 1, dtype, seed) {
|
|
this.canReturnFloat = () => this.dtype == null || this.dtype === "float32";
|
|
this.min = min2;
|
|
this.range = max2 - min2;
|
|
this.dtype = dtype;
|
|
if (seed == null) {
|
|
seed = Math.random();
|
|
}
|
|
if (typeof seed === "number") {
|
|
seed = seed.toString();
|
|
}
|
|
if (!this.canReturnFloat() && this.range <= 1) {
|
|
throw new Error(`The difference between ${min2} - ${max2} <= 1 and dtype is not float`);
|
|
}
|
|
this.random = seedrandomExports.alea(seed);
|
|
}
|
|
convertValue(value) {
|
|
if (this.canReturnFloat()) {
|
|
return value;
|
|
}
|
|
return Math.round(value);
|
|
}
|
|
nextValue() {
|
|
return this.convertValue(this.min + this.range * this.random());
|
|
}
|
|
}
|
|
function randomGamma_(shape, alpha, beta = 1, dtype = "float32", seed) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
if (beta == null) {
|
|
beta = 1;
|
|
}
|
|
if (dtype == null) {
|
|
dtype = "float32";
|
|
}
|
|
if (dtype !== "float32" && dtype !== "int32") {
|
|
throw new Error(`Unsupported data type ${dtype}`);
|
|
}
|
|
const rgamma = new RandGamma(alpha, beta, dtype, seed);
|
|
const res = buffer(shape, dtype);
|
|
for (let i = 0; i < res.values.length; i++) {
|
|
res.values[i] = rgamma.nextValue();
|
|
}
|
|
return res.toTensor();
|
|
}
|
|
const randomGamma = /* @__PURE__ */ op({ randomGamma_ });
|
|
function randomNormal_(shape, mean2 = 0, stdDev = 1, dtype, seed) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
if (dtype != null && dtype === "bool") {
|
|
throw new Error(`Unsupported data type ${dtype}`);
|
|
}
|
|
const randGauss = new MPRandGauss(mean2, stdDev, dtype, false, seed);
|
|
const res = buffer(shape, dtype);
|
|
for (let i = 0; i < res.values.length; i++) {
|
|
res.values[i] = randGauss.nextValue();
|
|
}
|
|
return res.toTensor();
|
|
}
|
|
const randomNormal = /* @__PURE__ */ op({ randomNormal_ });
|
|
function randomStandardNormal_(shape, dtype, seed) {
|
|
if (dtype != null && dtype === "bool") {
|
|
throw new Error(`Unsupported data type ${dtype}`);
|
|
}
|
|
return randomNormal(shape, 0, 1, dtype, seed);
|
|
}
|
|
const randomStandardNormal = /* @__PURE__ */ op({ randomStandardNormal_ });
|
|
function randomUniform_(shape, minval = 0, maxval = 1, dtype = "float32", seed) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
const res = buffer(shape, dtype);
|
|
const random = new UniformRandom(minval, maxval, null, seed);
|
|
for (let i = 0; i < res.values.length; i++) {
|
|
res.values[i] = random.nextValue();
|
|
}
|
|
return res.toTensor();
|
|
}
|
|
const randomUniform = /* @__PURE__ */ op({ randomUniform_ });
|
|
function randomUniformInt_(shape, minval, maxval, seed) {
|
|
return randomUniform(shape, minval, maxval, "int32", seed);
|
|
}
|
|
const randomUniformInt = /* @__PURE__ */ op({ randomUniformInt_ });
|
|
function range$2(start, stop, step2 = 1, dtype = "float32") {
|
|
if (step2 === 0) {
|
|
throw new Error("Cannot have a step of zero");
|
|
}
|
|
const attrs = { start, stop, step: step2, dtype };
|
|
return ENGINE.runKernel(Range, {}, attrs);
|
|
}
|
|
function real_(input) {
|
|
const $input = convertToTensor(input, "input", "real");
|
|
const inputs = { input: $input };
|
|
return ENGINE.runKernel(Real, inputs);
|
|
}
|
|
const real$2 = /* @__PURE__ */ op({ real_ });
|
|
function reciprocal_(x) {
|
|
const $x = convertToTensor(x, "x", "reciprocal");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Reciprocal, inputs);
|
|
}
|
|
const reciprocal$2 = /* @__PURE__ */ op({ reciprocal_ });
|
|
function relu_(x) {
|
|
const $x = convertToTensor(x, "x", "relu");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Relu, inputs);
|
|
}
|
|
const relu$2 = /* @__PURE__ */ op({ relu_ });
|
|
function relu6_(x) {
|
|
const $x = convertToTensor(x, "x", "relu6");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Relu6, inputs);
|
|
}
|
|
const relu6$2 = /* @__PURE__ */ op({ relu6_ });
|
|
function reverse_(x, axis) {
|
|
const $x = convertToTensor(x, "x", "reverse");
|
|
const inputs = { x: $x };
|
|
const attrs = { dims: axis };
|
|
return ENGINE.runKernel(Reverse, inputs, attrs);
|
|
}
|
|
const reverse$2 = /* @__PURE__ */ op({ reverse_ });
|
|
function reverse1d_(x) {
|
|
const $x = convertToTensor(x, "x", "reverse");
|
|
assert($x.rank === 1, () => `Error in reverse1D: x must be rank 1 but got rank ${$x.rank}.`);
|
|
return reverse$2($x, 0);
|
|
}
|
|
const reverse1d = /* @__PURE__ */ op({ reverse1d_ });
|
|
function reverse2d_(x, axis) {
|
|
const $x = convertToTensor(x, "x", "reverse");
|
|
assert($x.rank === 2, () => `Error in reverse2D: x must be rank 2 but got rank ${$x.rank}.`);
|
|
return reverse$2($x, axis);
|
|
}
|
|
const reverse2d = /* @__PURE__ */ op({ reverse2d_ });
|
|
function reverse3d_(x, axis) {
|
|
const $x = convertToTensor(x, "x", "reverse");
|
|
assert($x.rank === 3, () => `Error in reverse3D: x must be rank 3 but got rank ${$x.rank}.`);
|
|
return reverse$2($x, axis);
|
|
}
|
|
const reverse3d = /* @__PURE__ */ op({ reverse3d_ });
|
|
function reverse4d_(x, axis) {
|
|
const $x = convertToTensor(x, "x", "reverse");
|
|
assert($x.rank === 4, () => `Error in reverse4D: x must be rank 4 but got rank ${$x.rank}.`);
|
|
return reverse$2($x, axis);
|
|
}
|
|
const reverse4d = /* @__PURE__ */ op({ reverse4d_ });
|
|
function round_(x) {
|
|
const $x = convertToTensor(x, "x", "round");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Round, inputs);
|
|
}
|
|
const round$2 = /* @__PURE__ */ op({ round_ });
|
|
function rsqrt_(x) {
|
|
const $x = convertToTensor(x, "x", "rsqrt", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Rsqrt, inputs);
|
|
}
|
|
const rsqrt$2 = /* @__PURE__ */ op({ rsqrt_ });
|
|
function selu_(x) {
|
|
const $x = convertToTensor(x, "x", "selu");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Selu, inputs);
|
|
}
|
|
const selu$2 = /* @__PURE__ */ op({ selu_ });
|
|
function separableConv2d_(x, depthwiseFilter, pointwiseFilter, strides, pad2, dilation = [1, 1], dataFormat = "NHWC") {
|
|
const $x = convertToTensor(x, "x", "separableConv2d");
|
|
const $depthwiseFilter = convertToTensor(depthwiseFilter, "depthwiseFilter", "separableConv2d");
|
|
const $pointwiseFilter = convertToTensor(pointwiseFilter, "pointwiseFilter", "separableConv2d");
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
if (dataFormat === "NCHW") {
|
|
throw new Error("separableConv2d currently does not support dataFormat NCHW; only NHWC is supported");
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in separableConv2d: input must be rank 4, but got rank ${x4D.rank}.`);
|
|
assert($depthwiseFilter.rank === 4, () => `Error in separableConv2d: depthwise filter must be rank 4, but got rank ${$depthwiseFilter.rank}.`);
|
|
assert($pointwiseFilter.rank === 4, () => `Error in separableConv2d: pointwise filter must be rank 4, but got rank ${$depthwiseFilter.rank}.`);
|
|
assert($pointwiseFilter.shape[0] === 1, () => `Error in separableConv2d: the first dimension of pointwise filter must be 1, but got ${$pointwiseFilter.shape[0]}.`);
|
|
assert($pointwiseFilter.shape[1] === 1, () => `Error in separableConv2d: the second dimension of pointwise filter must be 1, but got ${$pointwiseFilter.shape[1]}.`);
|
|
const inChannels = $depthwiseFilter.shape[2];
|
|
const channelMultiplier = $depthwiseFilter.shape[3];
|
|
assert($pointwiseFilter.shape[2] === inChannels * channelMultiplier, () => `Error in separableConv2d: the third dimension of pointwise filter must be ${inChannels * channelMultiplier}, but got ${$pointwiseFilter.shape[2]}.`);
|
|
const depthwise = depthwiseConv2d$1(x4D, $depthwiseFilter, strides, pad2, dataFormat, dilation);
|
|
const pointwiseStride = 1;
|
|
const res = conv2d$2(depthwise, $pointwiseFilter, pointwiseStride, "valid", dataFormat);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const separableConv2d = /* @__PURE__ */ op({ separableConv2d_ });
|
|
async function setdiff1dAsync_(x, y) {
|
|
const $x = convertToTensor(x, "x", "setdiff1d");
|
|
const $y = convertToTensor(y, "y", "setdiff1d");
|
|
assert($x.dtype === $y.dtype, () => `x and y should have the same dtype, but got x (${$x.dtype}) and y (${$y.dtype}).`);
|
|
assert($x.rank === 1, () => `x should be 1D tensor, but got x (${$x.shape}).`);
|
|
assert($y.rank === 1, () => `y should be 1D tensor, but got y (${$y.shape}).`);
|
|
const xVals = await $x.data();
|
|
const yVals = await $y.data();
|
|
const ySet = new Set(yVals);
|
|
let outputSize = 0;
|
|
for (let i = 0; i < xVals.length; i++) {
|
|
if (!ySet.has(xVals[i])) {
|
|
outputSize++;
|
|
}
|
|
}
|
|
const buffer2 = new TensorBuffer([outputSize], $x.dtype);
|
|
const indices = new TensorBuffer([outputSize], "int32");
|
|
for (let i = 0, p2 = 0; i < xVals.length; i++) {
|
|
if (!ySet.has(xVals[i])) {
|
|
buffer2.values[p2] = xVals[i];
|
|
indices.values[p2] = i;
|
|
p2++;
|
|
}
|
|
}
|
|
return [buffer2.toTensor(), indices.toTensor()];
|
|
}
|
|
const setdiff1dAsync = setdiff1dAsync_;
|
|
function sign_(x) {
|
|
const $x = convertToTensor(x, "x", "sign");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Sign, inputs);
|
|
}
|
|
const sign$2 = /* @__PURE__ */ op({ sign_ });
|
|
function sin_(x) {
|
|
const $x = convertToTensor(x, "x", "sin", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Sin, inputs);
|
|
}
|
|
const sin$2 = /* @__PURE__ */ op({ sin_ });
|
|
function sinh_(x) {
|
|
const $x = convertToTensor(x, "x", "sinh");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Sinh, inputs);
|
|
}
|
|
const sinh$2 = /* @__PURE__ */ op({ sinh_ });
|
|
function slice1d_(x, begin, size) {
|
|
const $x = convertToTensor(x, "x", "slice1d");
|
|
assert($x.rank === 1, () => `slice1d expects a rank-1 tensor, but got a rank-${$x.rank} tensor`);
|
|
return slice$2($x, [begin], [size]);
|
|
}
|
|
const slice1d = /* @__PURE__ */ op({ slice1d_ });
|
|
function slice2d_(x, begin, size) {
|
|
const $x = convertToTensor(x, "x", "slice2d");
|
|
assert($x.rank === 2, () => `slice2d expects a rank-2 tensor, but got a rank-${$x.rank} tensor`);
|
|
return slice$2($x, begin, size);
|
|
}
|
|
const slice2d = /* @__PURE__ */ op({ slice2d_ });
|
|
function slice3d_(x, begin, size) {
|
|
const $x = convertToTensor(x, "x", "slice3d");
|
|
assert($x.rank === 3, () => `slice3d expects a rank-3 tensor, but got a rank-${$x.rank} tensor`);
|
|
return slice$2($x, begin, size);
|
|
}
|
|
const slice3d = /* @__PURE__ */ op({ slice3d_ });
|
|
function slice4d_(x, begin, size) {
|
|
const $x = convertToTensor(x, "x", "slice4d");
|
|
assert($x.rank === 4, () => `slice4d expects a rank-4 tensor, but got a rank-${$x.rank} tensor`);
|
|
return slice$2($x, begin, size);
|
|
}
|
|
const slice4d = /* @__PURE__ */ op({ slice4d_ });
|
|
function softmax_(logits, dim = -1) {
|
|
const $logits = convertToTensor(logits, "logits", "softmax", "float32");
|
|
if (dim === -1) {
|
|
dim = $logits.rank - 1;
|
|
}
|
|
if (dim !== $logits.rank - 1) {
|
|
throw Error(`Softmax along a non-last dimension is not yet supported. Logits was rank ${$logits.rank} and dim was ${dim}`);
|
|
}
|
|
const inputs = { logits: $logits };
|
|
const attrs = { dim };
|
|
return ENGINE.runKernel(Softmax, inputs, attrs);
|
|
}
|
|
const softmax$2 = /* @__PURE__ */ op({ softmax_ });
|
|
function fft_(input) {
|
|
assert(input.dtype === "complex64", () => `The dtype for tf.spectral.fft() must be complex64 but got ${input.dtype}.`);
|
|
const inputs = { input };
|
|
return ENGINE.runKernel(FFT, inputs);
|
|
}
|
|
const fft$2 = /* @__PURE__ */ op({ fft_ });
|
|
function ifft_(input) {
|
|
assert(input.dtype === "complex64", () => `The dtype for tf.spectral.ifft() must be complex64 but got ${input.dtype}.`);
|
|
const inputs = { input };
|
|
return ENGINE.runKernel(IFFT, inputs);
|
|
}
|
|
const ifft$2 = /* @__PURE__ */ op({ ifft_ });
|
|
function irfft_(input) {
|
|
const innerDimensionSize = input.shape[input.shape.length - 1];
|
|
const batch = input.size / innerDimensionSize;
|
|
let ret;
|
|
if (innerDimensionSize <= 2) {
|
|
const complexInput = reshape$2(input, [batch, innerDimensionSize]);
|
|
ret = ifft$2(complexInput);
|
|
} else {
|
|
const outputShape = [batch, 2 * (innerDimensionSize - 1)];
|
|
const realInput = reshape$2(real$2(input), [batch, innerDimensionSize]);
|
|
const imagInput = reshape$2(imag$2(input), [batch, innerDimensionSize]);
|
|
const realConjugate = reverse$2(slice$2(realInput, [0, 1], [batch, innerDimensionSize - 2]), 1);
|
|
const imagConjugate = mul(reverse$2(slice$2(imagInput, [0, 1], [batch, innerDimensionSize - 2]), 1), scalar(-1));
|
|
const r = concat$2([realInput, realConjugate], 1);
|
|
const i = concat$2([imagInput, imagConjugate], 1);
|
|
const complexInput = reshape$2(complex$2(r, i), [outputShape[0], outputShape[1]]);
|
|
ret = ifft$2(complexInput);
|
|
}
|
|
ret = real$2(ret);
|
|
if (input.rank === 3 && input.shape[0] !== 0) {
|
|
const temp = ret;
|
|
const batch2 = input.shape[0];
|
|
ret = reshape$2(ret, [batch2, ret.shape[0] / batch2, ret.shape[1]]);
|
|
temp.dispose();
|
|
}
|
|
return ret;
|
|
}
|
|
const irfft = /* @__PURE__ */ op({ irfft_ });
|
|
function split_(x, numOrSizeSplits, axis = 0) {
|
|
const $x = convertToTensor(x, "x", "split");
|
|
const inputs = { x: $x };
|
|
const attr = { numOrSizeSplits, axis };
|
|
return ENGINE.runKernel(SplitV, inputs, attr);
|
|
}
|
|
const split$2 = /* @__PURE__ */ op({ split_ });
|
|
function rfft_(input, fftLength) {
|
|
assert(input.dtype === "float32", () => `The dtype for rfft() must be real value but got ${input.dtype}`);
|
|
let innerDimensionSize = input.shape[input.shape.length - 1];
|
|
const batch = input.size / innerDimensionSize;
|
|
let adjustedInput;
|
|
if (fftLength != null && fftLength < innerDimensionSize) {
|
|
const begin = input.shape.map((v) => 0);
|
|
const size = input.shape.map((v) => v);
|
|
size[input.shape.length - 1] = fftLength;
|
|
adjustedInput = slice$2(input, begin, size);
|
|
innerDimensionSize = fftLength;
|
|
} else if (fftLength != null && fftLength > innerDimensionSize) {
|
|
const zerosShape = input.shape.map((v) => v);
|
|
zerosShape[input.shape.length - 1] = fftLength - innerDimensionSize;
|
|
adjustedInput = concat$2([input, zeros$1(zerosShape)], input.shape.length - 1);
|
|
innerDimensionSize = fftLength;
|
|
} else {
|
|
adjustedInput = input;
|
|
}
|
|
const zerosInput = zerosLike$2(adjustedInput);
|
|
const complexInput = reshape$2(complex$2(adjustedInput, zerosInput), [batch, innerDimensionSize]);
|
|
const ret = fft$2(complexInput);
|
|
const half = Math.floor(innerDimensionSize / 2) + 1;
|
|
const realValues = real$2(ret);
|
|
const imagValues = imag$2(ret);
|
|
const realComplexConjugate = split$2(realValues, [half, innerDimensionSize - half], realValues.shape.length - 1);
|
|
const imagComplexConjugate = split$2(imagValues, [half, innerDimensionSize - half], imagValues.shape.length - 1);
|
|
const outputShape = adjustedInput.shape.slice();
|
|
outputShape[adjustedInput.shape.length - 1] = half;
|
|
return reshape$2(complex$2(realComplexConjugate[0], imagComplexConjugate[0]), outputShape);
|
|
}
|
|
const rfft = /* @__PURE__ */ op({ rfft_ });
|
|
function squaredDifference_(a, b) {
|
|
let $a = convertToTensor(a, "a", "squaredDifference");
|
|
let $b = convertToTensor(b, "b", "squaredDifference");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
assertAndGetBroadcastShape($a.shape, $b.shape);
|
|
const inputs = { a: $a, b: $b };
|
|
const attrs = {};
|
|
return ENGINE.runKernel(SquaredDifference, inputs, attrs);
|
|
}
|
|
const squaredDifference$2 = /* @__PURE__ */ op({ squaredDifference_ });
|
|
function squeeze_(x, axis) {
|
|
const $x = convertToTensor(x, "x", "squeeze", "string_or_numeric");
|
|
return reshape$2($x, squeezeShape($x.shape, axis).newShape);
|
|
}
|
|
const squeeze = /* @__PURE__ */ op({ squeeze_ });
|
|
function stack_(tensors, axis = 0) {
|
|
const $tensors = convertToTensorArray(tensors, "tensors", "stack", "string_or_numeric");
|
|
assert($tensors.length >= 1, () => "Pass at least one tensor to tf.stack");
|
|
if ($tensors.length > 0) {
|
|
assert(axis <= $tensors[0].rank, () => "Axis must be <= rank of the tensor");
|
|
}
|
|
const inputs = $tensors;
|
|
const attrs = { axis };
|
|
return ENGINE.runKernel(Pack, inputs, attrs);
|
|
}
|
|
const stack = /* @__PURE__ */ op({ stack_ });
|
|
function step_(x, alpha = 0) {
|
|
const $x = convertToTensor(x, "x", "step");
|
|
const inputs = { x: $x };
|
|
const attrs = { alpha };
|
|
return ENGINE.runKernel(Step, inputs, attrs);
|
|
}
|
|
const step$2 = /* @__PURE__ */ op({ step_ });
|
|
function stridedSlice_(x, begin, end, strides, beginMask = 0, endMask = 0, ellipsisMask = 0, newAxisMask = 0, shrinkAxisMask = 0) {
|
|
const $x = convertToTensor(x, "x", "stridedSlice", "string_or_numeric");
|
|
const inputs = { x: $x };
|
|
const attrs = {
|
|
begin,
|
|
end,
|
|
strides,
|
|
beginMask,
|
|
endMask,
|
|
ellipsisMask,
|
|
newAxisMask,
|
|
shrinkAxisMask
|
|
};
|
|
return ENGINE.runKernel(StridedSlice, inputs, attrs);
|
|
}
|
|
const stridedSlice$2 = /* @__PURE__ */ op({ stridedSlice_ });
|
|
function tan_(x) {
|
|
const $x = convertToTensor(x, "x", "tan", "float32");
|
|
const inputs = { x: $x };
|
|
return ENGINE.runKernel(Tan, inputs);
|
|
}
|
|
const tan$2 = /* @__PURE__ */ op({ tan_ });
|
|
function tensor1d(values, dtype) {
|
|
assertNonNull(values);
|
|
const inferredShape = inferShape(values, dtype);
|
|
if (inferredShape.length !== 1) {
|
|
throw new Error("tensor1d() requires values to be a flat/TypedArray");
|
|
}
|
|
const shape = null;
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
function tensor2d(values, shape, dtype) {
|
|
assertNonNull(values);
|
|
if (shape != null && shape.length !== 2) {
|
|
throw new Error("tensor2d() requires shape to have two numbers");
|
|
}
|
|
const inferredShape = inferShape(values, dtype);
|
|
if (inferredShape.length !== 2 && inferredShape.length !== 1) {
|
|
throw new Error("tensor2d() requires values to be number[][] or flat/TypedArray");
|
|
}
|
|
if (inferredShape.length === 1 && shape == null) {
|
|
throw new Error("tensor2d() requires shape to be provided when `values` are a flat/TypedArray");
|
|
}
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
function tensor3d(values, shape, dtype) {
|
|
assertNonNull(values);
|
|
if (shape != null && shape.length !== 3) {
|
|
throw new Error("tensor3d() requires shape to have three numbers");
|
|
}
|
|
const inferredShape = inferShape(values, dtype);
|
|
if (inferredShape.length !== 3 && inferredShape.length !== 1) {
|
|
throw new Error("tensor3d() requires values to be number[][][] or flat/TypedArray");
|
|
}
|
|
if (inferredShape.length === 1 && shape == null) {
|
|
throw new Error("tensor3d() requires shape to be provided when `values` are a flat array");
|
|
}
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
function tensor4d(values, shape, dtype) {
|
|
assertNonNull(values);
|
|
if (shape != null && shape.length !== 4) {
|
|
throw new Error("tensor4d() requires shape to have four numbers");
|
|
}
|
|
const inferredShape = inferShape(values, dtype);
|
|
if (inferredShape.length !== 4 && inferredShape.length !== 1) {
|
|
throw new Error("tensor4d() requires values to be number[][][][] or flat/TypedArray");
|
|
}
|
|
if (inferredShape.length === 1 && shape == null) {
|
|
throw new Error("tensor4d() requires shape to be provided when `values` are a flat array");
|
|
}
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
function tensor5d(values, shape, dtype) {
|
|
assertNonNull(values);
|
|
if (shape != null && shape.length !== 5) {
|
|
throw new Error("tensor5d() requires shape to have five numbers");
|
|
}
|
|
const inferredShape = inferShape(values, dtype);
|
|
if (inferredShape.length !== 5 && inferredShape.length !== 1) {
|
|
throw new Error("tensor5d() requires values to be number[][][][][] or flat/TypedArray");
|
|
}
|
|
if (inferredShape.length === 1 && shape == null) {
|
|
throw new Error("tensor5d() requires shape to be provided when `values` are a flat array");
|
|
}
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
function tensor6d(values, shape, dtype) {
|
|
assertNonNull(values);
|
|
if (shape != null && shape.length !== 6) {
|
|
throw new Error("tensor6d() requires shape to have six numbers");
|
|
}
|
|
const inferredShape = inferShape(values, dtype);
|
|
if (inferredShape.length !== 6 && inferredShape.length !== 1) {
|
|
throw new Error("tensor6d() requires values to be number[][][][][][] or flat/TypedArray");
|
|
}
|
|
if (inferredShape.length === 1 && shape == null) {
|
|
throw new Error("tensor6d() requires shape to be provided when `values` are a flat array");
|
|
}
|
|
shape = shape || inferredShape;
|
|
return makeTensor(values, shape, inferredShape, dtype);
|
|
}
|
|
function validateUpdateShape(shape, indices, updates) {
|
|
const sliceDim = indices.rank > 1 ? indices.shape[indices.rank - 1] : 1;
|
|
const batchDim = indices.rank > 1 ? indices.rank - 1 : 1;
|
|
const shapeError = `Must have updates.shape = indices.shape[:batchDim] + shape[sliceDim:], got updates.shape: ${updates.shape}, indices.shape: ${indices.shape}, shape: ${shape}, sliceDim: ${sliceDim}, and batchDim: ${batchDim}.`;
|
|
if (updates.rank < batchDim) {
|
|
throw new Error(shapeError + ` update.rank < ${batchDim}. `);
|
|
}
|
|
if (shape.length < sliceDim + (updates.rank - batchDim)) {
|
|
throw new Error(shapeError + ` Output shape length < ${sliceDim + (updates.rank - batchDim)}`);
|
|
}
|
|
if (updates.rank !== batchDim + shape.length - sliceDim) {
|
|
throw new Error(shapeError + ` update.rank != ${batchDim + shape.length - sliceDim}`);
|
|
}
|
|
for (let d = 0; d < batchDim; ++d) {
|
|
if (updates.shape[d] !== indices.shape[d]) {
|
|
throw new Error(shapeError + ` updates.shape[${d}] (${updates.shape[d]}) != indices.shape[${d}] (${indices.shape[d]}).`);
|
|
}
|
|
}
|
|
for (let d = 0; d < updates.rank - batchDim; ++d) {
|
|
if (updates.shape[d + batchDim] !== shape[d + sliceDim]) {
|
|
throw new Error(shapeError + ` updates.shape[${d + batchDim}] (${updates.shape[d + batchDim]}) != shape[${d + batchDim}] (${shape[d + batchDim]})`);
|
|
}
|
|
}
|
|
}
|
|
function validateInput$1(updates, indices, shape) {
|
|
if (indices.rank < 1) {
|
|
throw new Error(`tf.scatterND() expects the indices to be rank 1 or higher, but the rank was ${indices.rank}.`);
|
|
}
|
|
if (updates.rank < 1) {
|
|
throw new Error(`tf.scatterND() expects the updates to be rank 1 or higher, but the rank was ${updates.rank}.`);
|
|
}
|
|
if (indices.dtype !== "int32") {
|
|
throw new Error(`The dtype of 'indices' should be int32, but got dtype: ${indices.dtype}`);
|
|
}
|
|
if (shape.length < 1) {
|
|
throw new Error(`Output rank must be greater or equal to 1, but got shape: ${shape}`);
|
|
}
|
|
if (shape.length === 0) {
|
|
if (indices.size === 0) {
|
|
throw new Error(`Indices specified for empty output. indices shape: ${indices.shape}`);
|
|
}
|
|
if (updates.size === 0) {
|
|
throw new Error(`Updates specified for empty output. updates shape: ${updates.shape}`);
|
|
}
|
|
}
|
|
validateUpdateShape(shape, indices, updates);
|
|
}
|
|
function calculateShapes(updates, indices, shape) {
|
|
const indicesRank = indices.shape.length;
|
|
const sliceRank = indicesRank > 1 ? indices.shape[indicesRank - 1] : 1;
|
|
const totalNd = shape.length;
|
|
let sliceSize = 1;
|
|
for (let i = sliceRank; i < totalNd; ++i) {
|
|
sliceSize *= shape[i];
|
|
}
|
|
const safeSliceDim = sliceRank < 1 ? 1 : sliceRank;
|
|
const numUpdates = sizeFromShape(indices.shape) / safeSliceDim;
|
|
const strides = [...computeStrides(shape.slice(0, sliceRank)), 1];
|
|
const outputSize = sizeFromShape(shape);
|
|
return { sliceRank, numUpdates, sliceSize, strides, outputSize };
|
|
}
|
|
function tensorScatterUpdate_(tensor2, indices, updates) {
|
|
const $tensor = convertToTensor(tensor2, "tensor", "tensorScatterupdate");
|
|
const $indices = convertToTensor(indices, "indices", "tensorScatterupdate", "int32");
|
|
const $updates = convertToTensor(updates, "updates", "tensorScatterupdate");
|
|
validateInput$1($updates, $indices, $tensor.shape);
|
|
if ($tensor.dtype !== $updates.dtype) {
|
|
throw new Error(`tensor and updates must have the same dtype, instead they are ${$tensor.dtype} and ${$updates.dtype}.`);
|
|
}
|
|
const inputs = {
|
|
tensor: $tensor,
|
|
indices: $indices,
|
|
updates: $updates
|
|
};
|
|
const attrs = {};
|
|
return ENGINE.runKernel(TensorScatterUpdate, inputs, attrs);
|
|
}
|
|
const tensorScatterUpdate$2 = op({ tensorScatterUpdate_ });
|
|
function topk_(x, k = 1, sorted = true) {
|
|
const $x = convertToTensor(x, "x", "topk");
|
|
if ($x.rank === 0) {
|
|
throw new Error("topk() expects the input to be of rank 1 or higher");
|
|
}
|
|
const lastDim = $x.shape[$x.shape.length - 1];
|
|
if (k < 0) {
|
|
throw new Error(`'k' passed to topk() must be >= 0 but got ${k}`);
|
|
}
|
|
if (k > lastDim) {
|
|
throw new Error(`'k' passed to topk() must be <= the last dimension (${lastDim}) but got ${k}`);
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { k, sorted };
|
|
const [values, indices] = ENGINE.runKernel(TopK, inputs, attrs);
|
|
return { values, indices };
|
|
}
|
|
const topk = /* @__PURE__ */ op({ topk_ });
|
|
function truncatedNormal_(shape, mean2 = 0, stdDev = 1, dtype, seed) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
if (dtype != null && dtype === "bool") {
|
|
throw new Error(`Unsupported data type $ { dtype }`);
|
|
}
|
|
const randGauss = new MPRandGauss(mean2, stdDev, dtype, true, seed);
|
|
const res = buffer(shape, dtype);
|
|
for (let i = 0; i < res.values.length; i++) {
|
|
res.values[i] = randGauss.nextValue();
|
|
}
|
|
return res.toTensor();
|
|
}
|
|
const truncatedNormal = /* @__PURE__ */ op({ truncatedNormal_ });
|
|
function unique_(x, axis = 0) {
|
|
const $x = convertToTensor(x, "x", "unique", "string_or_numeric");
|
|
assert($x.rank > 0, () => "The input tensor must be at least 1D");
|
|
const inputs = { x: $x };
|
|
const attrs = { axis };
|
|
const [values, indices] = ENGINE.runKernel(Unique, inputs, attrs);
|
|
return { values, indices };
|
|
}
|
|
const unique$2 = /* @__PURE__ */ op({ unique_ });
|
|
function unsortedSegmentSum_(x, segmentIds, numSegments) {
|
|
const $x = convertToTensor(x, "x", "unsortedSegmentSum");
|
|
const $segmentIds = convertToTensor(segmentIds, "segmentIds", "unsortedSegmentSum", "int32");
|
|
assert(isInt(numSegments), () => "numSegments must be of dtype int");
|
|
const inputs = { x: $x, segmentIds: $segmentIds };
|
|
const attrs = { numSegments };
|
|
return ENGINE.runKernel(UnsortedSegmentSum, inputs, attrs);
|
|
}
|
|
const unsortedSegmentSum$2 = /* @__PURE__ */ op({ unsortedSegmentSum_ });
|
|
function unstack_(x, axis = 0) {
|
|
const $x = convertToTensor(x, "x", "unstack", "string_or_numeric");
|
|
assert(axis >= -$x.shape.length && axis < $x.shape.length, () => `Axis = ${axis} is not in [-${$x.shape.length}, ${$x.shape.length})`);
|
|
const inputs = { value: $x };
|
|
const attrs = { axis };
|
|
return ENGINE.runKernel(Unpack, inputs, attrs);
|
|
}
|
|
const unstack = /* @__PURE__ */ op({ unstack_ });
|
|
function upperBound$1(sortedSequence, values) {
|
|
return searchSorted$2(sortedSequence, values, "right");
|
|
}
|
|
function variable(initialValue, trainable = true, name, dtype) {
|
|
return ENGINE.makeVariable(initialValue, trainable, name, dtype);
|
|
}
|
|
function whereImpl$2(condShape, condVals) {
|
|
const indices = [];
|
|
for (let i = 0; i < condVals.length; i++) {
|
|
if (condVals[i]) {
|
|
indices.push(i);
|
|
}
|
|
}
|
|
const inBuffer = buffer(condShape, "int32");
|
|
const out = buffer([indices.length, condShape.length], "int32");
|
|
for (let i = 0; i < indices.length; i++) {
|
|
const loc = inBuffer.indexToLoc(indices[i]);
|
|
const offset = i * condShape.length;
|
|
out.values.set(loc, offset);
|
|
}
|
|
return out.toTensor();
|
|
}
|
|
async function whereAsync_(condition) {
|
|
const $condition = convertToTensor(condition, "condition", "whereAsync", "bool");
|
|
const vals = await $condition.data();
|
|
const res = whereImpl$2($condition.shape, vals);
|
|
if (condition !== $condition) {
|
|
$condition.dispose();
|
|
}
|
|
return res;
|
|
}
|
|
const whereAsync = whereAsync_;
|
|
async function booleanMaskAsync_(tensor2, mask, axis) {
|
|
const $tensor = convertToTensor(tensor2, "tensor", "boolMask");
|
|
const $mask = convertToTensor(mask, "mask", "boolMask", "bool");
|
|
const axisFrom = axis == null ? 0 : axis;
|
|
const maskDim = $mask.rank;
|
|
const tensorShape = $tensor.shape;
|
|
assert(maskDim > 0, () => "mask cannot be scalar");
|
|
assertShapesMatch(tensorShape.slice(axisFrom, axisFrom + maskDim), $mask.shape, `mask's shape must match the first K dimensions of tensor's shape,`);
|
|
let leadingSize = 1;
|
|
for (let i = axisFrom; i < axisFrom + maskDim; i++) {
|
|
leadingSize *= tensorShape[i];
|
|
}
|
|
const targetTensorShape = tensorShape.slice(0, axisFrom).concat([leadingSize], tensorShape.slice(axisFrom + maskDim));
|
|
const reshapedTensor = reshape$2($tensor, targetTensorShape);
|
|
const reshapedMask = reshape$2($mask, [-1]);
|
|
const positivePositions = await whereAsync(reshapedMask);
|
|
const indices = squeeze(positivePositions, [1]);
|
|
const res = gather(reshapedTensor, indices, axisFrom);
|
|
if (tensor2 !== $tensor) {
|
|
$tensor.dispose();
|
|
}
|
|
if (mask !== $mask) {
|
|
$mask.dispose();
|
|
}
|
|
indices.dispose();
|
|
reshapedTensor.dispose();
|
|
reshapedMask.dispose();
|
|
positivePositions.dispose();
|
|
return res;
|
|
}
|
|
const booleanMaskAsync = booleanMaskAsync_;
|
|
function transpose_(x, perm, conjugate) {
|
|
const $x = convertToTensor(x, "x", "transpose");
|
|
if (perm == null) {
|
|
perm = $x.shape.map((s, i) => i).reverse();
|
|
}
|
|
assert($x.rank === perm.length, () => `Error in transpose: rank of input ${$x.rank} must match length of perm ${perm}.`);
|
|
perm.forEach((axis) => {
|
|
assert(axis >= 0 && axis < $x.rank, () => `All entries in 'perm' must be between 0 and ${$x.rank - 1} but got ${perm}`);
|
|
});
|
|
if ($x.rank <= 1) {
|
|
return $x.clone();
|
|
}
|
|
const inputs = { x: $x };
|
|
const attrs = { perm };
|
|
if ($x.dtype === "complex64") {
|
|
return tidy(() => {
|
|
let $real = real$2($x);
|
|
let $imag = imag$2($x);
|
|
$real = ENGINE.runKernel(Transpose, { x: $real }, attrs);
|
|
$imag = ENGINE.runKernel(Transpose, { x: $imag }, attrs);
|
|
if (conjugate) {
|
|
$imag = neg$2($imag);
|
|
}
|
|
return complex$2($real, $imag);
|
|
});
|
|
}
|
|
return ENGINE.runKernel(Transpose, inputs, attrs);
|
|
}
|
|
const transpose$2 = /* @__PURE__ */ op({ transpose_ });
|
|
function movingAverage_(v, x, decay, step2, zeroDebias = true) {
|
|
const $v = convertToTensor(v, "v", "movingAverage");
|
|
const $x = convertToTensor(x, "x", "movingAverage");
|
|
const $decay = convertToTensor(decay, "decay", "movingAverage");
|
|
assertTypesMatch($v, $x);
|
|
assert(arraysEqual($v.shape, $x.shape), () => "Shape mismatch in v and x");
|
|
const one = scalar(1);
|
|
const oneMinusDecay = sub$2(one, $decay);
|
|
let update = mul(sub$2($x, $v), oneMinusDecay);
|
|
if (zeroDebias) {
|
|
assert(step2 != null, () => "When using zeroDebias: true, step is required.");
|
|
const $step = convertToTensor(step2, "step", "movingAverage");
|
|
update = div$1(update, sub$2(one, pow$2($decay, $step)));
|
|
}
|
|
return add$1($v, update);
|
|
}
|
|
const movingAverage = /* @__PURE__ */ op({ movingAverage_ });
|
|
function scatterND_(indices, updates, shape) {
|
|
assertNonNegativeIntegerDimensions(shape);
|
|
const $indices = convertToTensor(indices, "indices", "scatterND", "int32");
|
|
const $updates = convertToTensor(updates, "updates", "scatterND");
|
|
validateInput$1($updates, $indices, shape);
|
|
const inputs = { indices: $indices, updates: $updates };
|
|
const attrs = { shape };
|
|
return ENGINE.runKernel(ScatterNd, inputs, attrs);
|
|
}
|
|
const scatterND = /* @__PURE__ */ op({ scatterND_ });
|
|
function validateInput(sparseIndices, sparseValues, outputShape, defaultValues) {
|
|
if (sparseIndices.dtype !== "int32") {
|
|
throw new Error(`tf.sparseToDense() expects the indices to be int32 type, but the dtype was ${sparseIndices.dtype}.`);
|
|
}
|
|
if (sparseIndices.rank > 2) {
|
|
throw new Error(`sparseIndices should be a scalar, vector, or matrix, but got shape ${sparseIndices.shape}.`);
|
|
}
|
|
const numElems = sparseIndices.rank > 0 ? sparseIndices.shape[0] : 1;
|
|
const numDims = sparseIndices.rank > 1 ? sparseIndices.shape[1] : 1;
|
|
if (outputShape.length !== numDims) {
|
|
throw new Error(`outputShape has incorrect number of elements:, ${outputShape.length}, should be: ${numDims}.`);
|
|
}
|
|
const numValues = sparseValues.size;
|
|
if (!(sparseValues.rank === 0 || sparseValues.rank === 1 && numValues === numElems)) {
|
|
throw new Error(`sparseValues has incorrect shape ${sparseValues.shape}, should be [] or [${numElems}]`);
|
|
}
|
|
if (sparseValues.dtype !== defaultValues.dtype) {
|
|
throw new Error("sparseValues.dtype must match defaultValues.dtype");
|
|
}
|
|
}
|
|
function sparseToDense_(sparseIndices, sparseValues, outputShape, defaultValue = 0) {
|
|
assertNonNegativeIntegerDimensions(outputShape);
|
|
const $sparseIndices = convertToTensor(sparseIndices, "sparseIndices", "sparseToDense", "int32");
|
|
const $sparseValues = convertToTensor(sparseValues, "sparseValues", "sparseToDense", "string_or_numeric");
|
|
const $defaultValue = convertToTensor(defaultValue, "defaultValue", "sparseToDense", $sparseValues.dtype);
|
|
validateInput($sparseIndices, $sparseValues, outputShape, $defaultValue);
|
|
const inputs = {
|
|
sparseIndices: $sparseIndices,
|
|
sparseValues: $sparseValues,
|
|
defaultValue: $defaultValue
|
|
};
|
|
const attrs = { outputShape };
|
|
return ENGINE.runKernel(SparseToDense, inputs, attrs);
|
|
}
|
|
const sparseToDense$2 = /* @__PURE__ */ op({ sparseToDense_ });
|
|
function gatherND_(x, indices) {
|
|
const $indices = convertToTensor(indices, "indices", "gatherND", "int32");
|
|
const $x = convertToTensor(x, "x", "gatherND", "string_or_numeric");
|
|
const inputs = { params: $x, indices: $indices };
|
|
return ENGINE.runKernel(GatherNd, inputs);
|
|
}
|
|
const gatherND = /* @__PURE__ */ op({ gatherND_ });
|
|
function getNoiseShape(x, noiseShape) {
|
|
if (noiseShape == null) {
|
|
return x.shape.slice();
|
|
}
|
|
if (arraysEqual(x.shape, noiseShape)) {
|
|
return noiseShape;
|
|
}
|
|
if (x.shape.length === noiseShape.length) {
|
|
const newDimension = [];
|
|
for (let i = 0; i < x.shape.length; i++) {
|
|
if (noiseShape[i] == null && x.shape[i] != null) {
|
|
newDimension.push(x.shape[i]);
|
|
} else {
|
|
newDimension.push(noiseShape[i]);
|
|
}
|
|
}
|
|
return newDimension;
|
|
}
|
|
return noiseShape;
|
|
}
|
|
function dropout_(x, rate, noiseShape, seed) {
|
|
const $x = convertToTensor(x, "x", "dropout");
|
|
assert($x.dtype === "float32", () => `x has to be a floating point tensor since it's going to be scaled, but got a ${$x.dtype} tensor instead.`);
|
|
assert(rate >= 0 && rate < 1, () => `rate must be a float in the range [0, 1), but got ${rate}.`);
|
|
if (rate === 0) {
|
|
return x instanceof Tensor ? $x.clone() : $x;
|
|
}
|
|
const $noiseShape = getNoiseShape($x, noiseShape);
|
|
const keepProb = 1 - rate;
|
|
const multiplier = div$1(floor$2(add$1(randomUniform($noiseShape, 0, 1, "float32", seed), keepProb)), keepProb);
|
|
return mul($x, multiplier);
|
|
}
|
|
const dropout = /* @__PURE__ */ op({ dropout_ });
|
|
function enclosingPowerOfTwo(value) {
|
|
return Math.floor(Math.pow(2, Math.ceil(Math.log(value) / Math.log(2))));
|
|
}
|
|
function cosineWindow(windowLength, a, b) {
|
|
const even = 1 - windowLength % 2;
|
|
const newValues = new Float32Array(windowLength);
|
|
for (let i = 0; i < windowLength; ++i) {
|
|
const cosArg = 2 * Math.PI * i / (windowLength + even - 1);
|
|
newValues[i] = a - b * Math.cos(cosArg);
|
|
}
|
|
return tensor1d(newValues, "float32");
|
|
}
|
|
async function inTopKAsync_(predictions, targets, k = 1) {
|
|
const $predictions = convertToTensor(predictions, "predictions", "inTopK");
|
|
const $targets = convertToTensor(targets, "targets", "inTopK");
|
|
assert($predictions.rank > 1, () => `inTopK() expects the predictions to be of rank 2 or higher, but got ${$predictions.rank}`);
|
|
assert($predictions.rank - 1 === $targets.rank, () => `predictions rank should be 1 larger than targets rank, but got predictions rank ${$predictions.rank} and targets rank ${$targets.rank}`);
|
|
assertShapesMatch($predictions.shape.slice(0, $predictions.shape.length - 1), $targets.shape, `predictions's shape should be align with the targets' shape, except the last dimension.`);
|
|
const lastDim = $predictions.shape[$predictions.shape.length - 1];
|
|
assert(k > 0 && k <= lastDim, () => `'k' passed to inTopK() must be > 0 && <= the predictions last dimension (${lastDim}), but got ${k}`);
|
|
const predictionsVals = await $predictions.data();
|
|
const targetsVals = await $targets.data();
|
|
const [batch, size] = [predictionsVals.length / lastDim, lastDim];
|
|
const precision = getTypedArrayFromDType("bool", batch);
|
|
for (let b = 0; b < batch; b++) {
|
|
const offset = b * size;
|
|
const vals = predictionsVals.subarray(offset, offset + size);
|
|
const valAndInd = [];
|
|
for (let i = 0; i < vals.length; i++) {
|
|
valAndInd.push({ value: vals[i], index: i });
|
|
}
|
|
valAndInd.sort((a, b2) => b2.value - a.value);
|
|
precision[b] = 0;
|
|
for (let i = 0; i < k; i++) {
|
|
if (valAndInd[i].index === targetsVals[b]) {
|
|
precision[b] = 1;
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
if (predictions !== $predictions) {
|
|
$predictions.dispose();
|
|
}
|
|
if (targets !== $targets) {
|
|
$targets.dispose();
|
|
}
|
|
return tensor(precision, $targets.shape, "bool");
|
|
}
|
|
const inTopKAsync = inTopKAsync_;
|
|
function conv2DBackpropFilter_(x, dy, filterShape, strides, pad2, dataFormat = "NHWC", dimRoundingMode) {
|
|
let x4D = x;
|
|
if (x.rank === 3) {
|
|
x4D = reshape$2(x, [1, x.shape[0], x.shape[1], x.shape[2]]);
|
|
}
|
|
let dy4D = dy;
|
|
if (dy4D.rank === 3) {
|
|
dy4D = reshape$2(dy, [1, dy.shape[0], dy.shape[1], dy.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in conv2dDerFilter: input must be rank 4, but got shape ${x4D.shape}.`);
|
|
assert(dy4D.rank === 4, () => `Error in conv2dDerFilter: dy must be rank 4, but got shape ${dy4D.shape}.`);
|
|
assert(filterShape.length === 4, () => `Error in conv2dDerFilter: filterShape must be length 4, but got ${filterShape}.`);
|
|
const inDepth = dataFormat === "NHWC" ? x4D.shape[3] : x4D.shape[1];
|
|
const outDepth = dataFormat === "NHWC" ? dy4D.shape[3] : dy4D.shape[1];
|
|
assert(inDepth === filterShape[2], () => `Error in conv2dDerFilter: depth of input ${inDepth}) must match input depth in filter (${filterShape[2]}.`);
|
|
assert(outDepth === filterShape[3], () => `Error in conv2dDerFilter: depth of dy (${outDepth}) must match output depth for filter (${filterShape[3]}).`);
|
|
checkPadOnDimRoundingMode("conv2dDerFilter", pad2, dimRoundingMode);
|
|
const inputs = { x: x4D, dy: dy4D };
|
|
const attrs = { strides, pad: pad2, dataFormat, dimRoundingMode, filterShape };
|
|
return ENGINE.runKernel(Conv2DBackpropFilter, inputs, attrs);
|
|
}
|
|
const conv2DBackpropFilter$2 = /* @__PURE__ */ op({ conv2DBackpropFilter_ });
|
|
function getFusedDyActivation(dy, y, activation) {
|
|
if (activation == null || activation === "linear") {
|
|
return dy;
|
|
}
|
|
if (activation === "relu") {
|
|
return mul(dy, step$2(y));
|
|
}
|
|
throw new Error(`Cannot compute gradient for fused activation ${activation}.`);
|
|
}
|
|
function getFusedBiasGradient(bias, dyActivation) {
|
|
let res = dyActivation;
|
|
const reduceAxes = getReductionAxes(bias.shape, dyActivation.shape);
|
|
if (reduceAxes.length > 0) {
|
|
res = sum$2(res, reduceAxes);
|
|
}
|
|
return reshape$2(res, bias.shape);
|
|
}
|
|
function applyActivation$1(x, activation, preluActivationWeights, leakyreluAlpha) {
|
|
if (activation === "linear") {
|
|
return x;
|
|
} else if (activation === "relu") {
|
|
return relu$2(x);
|
|
} else if (activation === "elu") {
|
|
return elu$2(x);
|
|
} else if (activation === "relu6") {
|
|
return relu6$2(x);
|
|
} else if (activation === "prelu") {
|
|
return prelu$2(x, preluActivationWeights);
|
|
} else if (activation === "leakyrelu") {
|
|
return leakyRelu$2(x, leakyreluAlpha);
|
|
} else if (activation === "sigmoid") {
|
|
return sigmoid$2(x);
|
|
}
|
|
throw new Error(`Unknown fused activation ${activation}.`);
|
|
}
|
|
const shouldFuse = (gradientDepth, activation) => {
|
|
const gradientMode = gradientDepth > 0;
|
|
return !gradientMode || activation === "linear";
|
|
};
|
|
function fusedConv2d_({ x, filter, strides, pad: pad2, dataFormat = "NHWC", dilations = [1, 1], dimRoundingMode, bias, activation = "linear", preluActivationWeights, leakyreluAlpha }) {
|
|
activation = activation || "linear";
|
|
if (shouldFuse(ENGINE.state.gradientDepth, activation) === false) {
|
|
assert(dataFormat === "NHWC", () => `Error in fused conv2d: got dataFormat of ${dataFormat} but only NHWC is currently supported for the case of gradient depth is 0 and the activation is not linear.`);
|
|
let result = conv2d$2(x, filter, strides, pad2, dataFormat, dilations, dimRoundingMode);
|
|
if (bias != null) {
|
|
result = add$1(result, bias);
|
|
}
|
|
return applyActivation$1(result, activation, preluActivationWeights, leakyreluAlpha);
|
|
}
|
|
const $x = convertToTensor(x, "x", "conv2d", "float32");
|
|
const $filter = convertToTensor(filter, "filter", "conv2d", "float32");
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in fused conv2d: input must be rank 4, but got rank ${x4D.rank}.`);
|
|
assert($filter.rank === 4, () => `Error in fused conv2d: filter must be rank 4, but got rank ${$filter.rank}.`);
|
|
checkPadOnDimRoundingMode("fused conv2d", pad2, dimRoundingMode);
|
|
const inputChannels = dataFormat === "NHWC" ? x4D.shape[3] : x4D.shape[1];
|
|
assert($filter.shape[2] === inputChannels, () => `Error in conv2d: depth of input (${inputChannels}) must match input depth for filter ${$filter.shape[2]}.`);
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in conv2D: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computeConv2DInfo(x4D.shape, $filter.shape, strides, dilations, pad2, dimRoundingMode);
|
|
let $bias;
|
|
if (bias != null) {
|
|
$bias = convertToTensor(bias, "bias", "fused conv2d");
|
|
[$bias] = makeTypesMatch($bias, $x);
|
|
if (dataFormat === "NHWC") {
|
|
assertAndGetBroadcastShape(convInfo.outShape, $bias.shape);
|
|
} else {
|
|
assert($bias.shape.length <= 1, () => `Error in fused conv2d: only supports scalar or 1-D Tensor bias for NCHW format but got the bias of rank-${$bias.shape.length}.`);
|
|
assert($bias.shape.length === 0 || $bias.shape[0] === convInfo.outChannels || $bias.shape[0] === 1, () => `Error in fused conv2d: bias shape (${$bias.shape}) is not compatible with the number of output channels (${convInfo.outChannels})`);
|
|
}
|
|
}
|
|
let $preluActivationWeights;
|
|
if (preluActivationWeights != null) {
|
|
const alphaShape = preluActivationWeights.shape;
|
|
assert(alphaShape.length <= 1 || alphaShape.length === 3, () => `Error in fused conv2d: only supports scalar, 1-D Tensor or 3-D Tensor PReLU activation weights but got a tensor of rank-${alphaShape.length}.`);
|
|
if (alphaShape.length === 1) {
|
|
assert(alphaShape[0] === 1 || alphaShape[0] === convInfo.outChannels, () => `Error in fused conv2d: PReLU activation weights (${alphaShape}) is not compatible with the number of output channels (${convInfo.outChannels}).`);
|
|
} else if (alphaShape.length === 3) {
|
|
try {
|
|
assertAndGetBroadcastShape(alphaShape, convInfo.outShape);
|
|
} catch (e) {
|
|
const errMsg = `Error in fused conv2d: PReLU activation weights (${alphaShape}) is not compatible with the output shape of the conv2d (${convInfo.outShape}).`;
|
|
throw Error(errMsg);
|
|
}
|
|
}
|
|
$preluActivationWeights = convertToTensor(preluActivationWeights, "prelu weights", "fused conv2d");
|
|
}
|
|
const grad = (dy, saved) => {
|
|
assert(dataFormat === "NHWC", () => `Error in gradient of fused conv2D: got dataFormat of ${dataFormat} but only NHWC is currently supported.`);
|
|
const [$filter2, x4D2, y, $bias2] = saved;
|
|
const dyActivation = getFusedDyActivation(dy, y, activation);
|
|
assert(tupleValuesAreOne(dilations), () => `Error in gradient of fused conv2D: dilation rates greater than 1 are not yet supported in gradients. Got dilations '${dilations}'`);
|
|
const xDer = conv2DBackpropInput$2(x4D2.shape, dyActivation, $filter2, strides, pad2);
|
|
const filterDer = conv2DBackpropFilter$2(x4D2, dyActivation, $filter2.shape, strides, pad2);
|
|
const der = [xDer, filterDer];
|
|
if ($bias2 != null) {
|
|
const biasDer = getFusedBiasGradient($bias2, dyActivation);
|
|
der.push(biasDer);
|
|
}
|
|
return der;
|
|
};
|
|
const inputs = {
|
|
x: x4D,
|
|
filter: $filter,
|
|
bias: $bias,
|
|
preluActivationWeights: $preluActivationWeights
|
|
};
|
|
const attrs = {
|
|
strides,
|
|
pad: pad2,
|
|
dataFormat,
|
|
dilations,
|
|
dimRoundingMode,
|
|
activation,
|
|
leakyreluAlpha
|
|
};
|
|
if (bias == null) {
|
|
const customOp = customGrad((x4D2, filter2, save) => {
|
|
let res = (
|
|
// tslint:disable-next-line: no-unnecessary-type-assertion
|
|
ENGINE.runKernel(FusedConv2D, inputs, attrs)
|
|
);
|
|
save([filter2, x4D2, res]);
|
|
if (reshapedTo4D) {
|
|
res = reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return { value: res, gradFunc: grad };
|
|
});
|
|
return customOp(x4D, $filter);
|
|
} else {
|
|
const customOpWithBias = customGrad((x4D2, filter2, bias2, save) => {
|
|
let res = ENGINE.runKernel(FusedConv2D, inputs, attrs);
|
|
save([filter2, x4D2, res, bias2]);
|
|
if (reshapedTo4D) {
|
|
res = reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return { value: res, gradFunc: grad };
|
|
});
|
|
return customOpWithBias(x4D, $filter, $bias);
|
|
}
|
|
}
|
|
const conv2d$1 = /* @__PURE__ */ op({ fusedConv2d_ });
|
|
function depthwiseConv2dNativeBackpropFilter_(x, dy, filterShape, strides, pad2, dilations = [1, 1], dimRoundingMode) {
|
|
let x4D = x;
|
|
if (x.rank === 3) {
|
|
x4D = reshape$2(x, [1, x.shape[0], x.shape[1], x.shape[2]]);
|
|
}
|
|
let dy4D = dy;
|
|
if (dy4D.rank === 3) {
|
|
dy4D = reshape$2(dy, [1, dy.shape[0], dy.shape[1], dy.shape[2]]);
|
|
}
|
|
const inputs = { x: x4D, dy: dy4D };
|
|
const attrs = { strides, pad: pad2, dimRoundingMode, dilations, filterShape };
|
|
return ENGINE.runKernel(DepthwiseConv2dNativeBackpropFilter, inputs, attrs);
|
|
}
|
|
const depthwiseConv2dNativeBackpropFilter$2 = op({ depthwiseConv2dNativeBackpropFilter_ });
|
|
function depthwiseConv2dNativeBackpropInput_(xShape, dy, filter, strides, pad2, dilations = [1, 1], dimRoundingMode) {
|
|
let dy4D = dy;
|
|
let reshapedTo4D = false;
|
|
if (dy.rank === 3) {
|
|
reshapedTo4D = true;
|
|
dy4D = reshape$2(dy, [1, dy.shape[0], dy.shape[1], dy.shape[2]]);
|
|
}
|
|
const inputs = { dy: dy4D, filter };
|
|
const attrs = { strides, pad: pad2, dimRoundingMode, dilations, inputShape: xShape };
|
|
const res = (
|
|
// tslint:disable-next-line: no-unnecessary-type-assertion
|
|
ENGINE.runKernel(DepthwiseConv2dNativeBackpropInput, inputs, attrs)
|
|
);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const depthwiseConv2dNativeBackpropInput$2 = op({ depthwiseConv2dNativeBackpropInput_ });
|
|
function fusedDepthwiseConv2d_({ x, filter, strides, pad: pad2, dataFormat = "NHWC", dilations = [1, 1], dimRoundingMode, bias, activation = "linear", preluActivationWeights, leakyreluAlpha }) {
|
|
if (shouldFuse(ENGINE.state.gradientDepth, activation) === false) {
|
|
let result = depthwiseConv2d$1(x, filter, strides, pad2, dataFormat, dilations, dimRoundingMode);
|
|
if (bias != null) {
|
|
result = add$1(result, bias);
|
|
}
|
|
return applyActivation$1(result, activation, preluActivationWeights, leakyreluAlpha);
|
|
}
|
|
const $x = convertToTensor(x, "x", "depthwiseConv2d", "float32");
|
|
const $filter = convertToTensor(filter, "filter", "depthwiseConv2d", "float32");
|
|
let x4D = $x;
|
|
let reshapedTo4D = false;
|
|
if ($x.rank === 3) {
|
|
reshapedTo4D = true;
|
|
x4D = reshape$2($x, [1, $x.shape[0], $x.shape[1], $x.shape[2]]);
|
|
}
|
|
assert(x4D.rank === 4, () => `Error in fused depthwiseConv2d: input must be rank 4, but got rank ${x4D.rank}.`);
|
|
assert($filter.rank === 4, () => `Error in fused depthwiseConv2d: filter must be rank 4, but got rank ${$filter.rank}.`);
|
|
assert(x4D.shape[3] === $filter.shape[2], () => `Error in fused depthwiseConv2d: number of input channels (${x4D.shape[3]}) must match the inChannels dimension in filter ${$filter.shape[2]}.`);
|
|
if (dilations == null) {
|
|
dilations = [1, 1];
|
|
}
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in fused depthwiseConv2d: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
checkPadOnDimRoundingMode("fused depthwiseConv2d", pad2, dimRoundingMode);
|
|
const convInfo = computeConv2DInfo(
|
|
x4D.shape,
|
|
$filter.shape,
|
|
strides,
|
|
dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
let $bias;
|
|
if (bias != null) {
|
|
$bias = convertToTensor(bias, "bias", "fused conv2d");
|
|
[$bias] = makeTypesMatch($bias, $x);
|
|
assertAndGetBroadcastShape(convInfo.outShape, $bias.shape);
|
|
}
|
|
let $preluActivationWeights;
|
|
if (preluActivationWeights != null) {
|
|
$preluActivationWeights = convertToTensor(preluActivationWeights, "prelu weights", "fused depthwiseConv2d");
|
|
}
|
|
const grad = (dy, saved) => {
|
|
assert(tupleValuesAreOne(dilations), () => `Error in gradient of fused depthwiseConv2d: dilation rates greater than 1 are not yet supported. Got dilations '${dilations}'`);
|
|
const [$filter2, x4D2, y, bias2] = saved;
|
|
const dyActivation = getFusedDyActivation(dy, y, activation);
|
|
const xDer = depthwiseConv2dNativeBackpropInput$2(x4D2.shape, dyActivation, $filter2, strides, pad2, dilations, dimRoundingMode);
|
|
const filterDer = depthwiseConv2dNativeBackpropFilter$2(x4D2, dyActivation, $filter2.shape, strides, pad2, dilations, dimRoundingMode);
|
|
if (bias2 != null) {
|
|
const biasDer = getFusedBiasGradient($bias, dyActivation);
|
|
return [xDer, filterDer, biasDer];
|
|
}
|
|
return [xDer, filterDer];
|
|
};
|
|
const inputs = {
|
|
x: x4D,
|
|
filter: $filter,
|
|
bias: $bias,
|
|
preluActivationWeights: $preluActivationWeights
|
|
};
|
|
const attrs = {
|
|
strides,
|
|
pad: pad2,
|
|
dataFormat,
|
|
dilations,
|
|
dimRoundingMode,
|
|
activation,
|
|
leakyreluAlpha
|
|
};
|
|
if (bias == null) {
|
|
const customOp = customGrad((x4D2, filter2, save) => {
|
|
let res = ENGINE.runKernel(FusedDepthwiseConv2D, inputs, attrs);
|
|
save([filter2, x4D2, res]);
|
|
if (reshapedTo4D) {
|
|
res = reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return { value: res, gradFunc: grad };
|
|
});
|
|
return customOp(x4D, $filter);
|
|
} else {
|
|
const customOpWithBias = customGrad((x4D2, filter2, bias2, save) => {
|
|
let res = ENGINE.runKernel(FusedDepthwiseConv2D, inputs, attrs);
|
|
save([filter2, x4D2, res, bias2]);
|
|
if (reshapedTo4D) {
|
|
res = reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return { value: res, gradFunc: grad };
|
|
});
|
|
return customOpWithBias(x4D, $filter, $bias);
|
|
}
|
|
}
|
|
const depthwiseConv2d = /* @__PURE__ */ op({ fusedDepthwiseConv2d_ });
|
|
function fusedMatMul_({ a, b, transposeA = false, transposeB = false, bias, activation = "linear", preluActivationWeights, leakyreluAlpha = 0.2 }) {
|
|
if (shouldFuse(ENGINE.state.gradientDepth, activation) === false) {
|
|
let result = matMul$1(a, b, transposeA, transposeB);
|
|
if (bias != null) {
|
|
result = add$1(result, bias);
|
|
}
|
|
return applyActivation$1(result, activation, preluActivationWeights, leakyreluAlpha);
|
|
}
|
|
let $a = convertToTensor(a, "a", "fused matMul");
|
|
let $b = convertToTensor(b, "b", "fused matMul");
|
|
[$a, $b] = makeTypesMatch($a, $b);
|
|
const innerShapeA = transposeA ? $a.shape[$a.rank - 2] : $a.shape[$a.rank - 1];
|
|
const innerShapeB = transposeB ? $b.shape[$b.rank - 1] : $b.shape[$b.rank - 2];
|
|
const outerShapeA = transposeA ? $a.shape[$a.rank - 1] : $a.shape[$a.rank - 2];
|
|
const outerShapeB = transposeB ? $b.shape[$b.rank - 2] : $b.shape[$b.rank - 1];
|
|
const outerDimsA = $a.shape.slice(0, -2);
|
|
const outerDimsB = $b.shape.slice(0, -2);
|
|
const batchDimA = sizeFromShape(outerDimsA);
|
|
const batchDimB = sizeFromShape(outerDimsB);
|
|
assert(innerShapeA === innerShapeB, () => `Error in fused matMul: inner shapes (${innerShapeA}) and (${innerShapeB}) of Tensors with shapes ${$a.shape} and ${$b.shape} and transposeA=${transposeA} and transposeB=${transposeB} must match.`);
|
|
const outShapeOuterDims = assertAndGetBroadcastShape($a.shape.slice(0, -2), $b.shape.slice(0, -2));
|
|
const outShape = outShapeOuterDims.concat([outerShapeA, outerShapeB]);
|
|
const a3D = transposeA ? reshape$2($a, [batchDimA, innerShapeA, outerShapeA]) : reshape$2($a, [batchDimA, outerShapeA, innerShapeA]);
|
|
const b3D = transposeB ? reshape$2($b, [batchDimB, outerShapeB, innerShapeB]) : reshape$2($b, [batchDimB, innerShapeB, outerShapeB]);
|
|
let $bias;
|
|
if (bias != null) {
|
|
$bias = convertToTensor(bias, "bias", "fused matMul");
|
|
[$bias] = makeTypesMatch($bias, $a);
|
|
assertAndGetBroadcastShape(outShape, $bias.shape);
|
|
}
|
|
let $preluActivationWeights;
|
|
if (preluActivationWeights != null) {
|
|
$preluActivationWeights = convertToTensor(preluActivationWeights, "prelu weights", "fused matMul");
|
|
}
|
|
const grad = (dy, saved) => {
|
|
const [a3D2, b3D2, y, $bias2] = saved;
|
|
const dyActivation = getFusedDyActivation(reshape$2(dy, y.shape), y, activation);
|
|
let aDer;
|
|
let bDer;
|
|
if (!transposeA && !transposeB) {
|
|
aDer = matMul$1(dyActivation, b3D2, false, true);
|
|
bDer = matMul$1(a3D2, dyActivation, true, false);
|
|
} else if (!transposeA && transposeB) {
|
|
aDer = matMul$1(dyActivation, b3D2, false, false);
|
|
bDer = matMul$1(dyActivation, a3D2, true, false);
|
|
} else if (transposeA && !transposeB) {
|
|
aDer = matMul$1(b3D2, dyActivation, false, true);
|
|
bDer = matMul$1(a3D2, dyActivation, false, false);
|
|
} else {
|
|
aDer = matMul$1(b3D2, dyActivation, true, true);
|
|
bDer = matMul$1(dyActivation, a3D2, true, true);
|
|
}
|
|
if (bias != null) {
|
|
const biasDer = getFusedBiasGradient($bias2, dyActivation);
|
|
return [aDer, bDer, biasDer];
|
|
} else {
|
|
return [aDer, bDer];
|
|
}
|
|
};
|
|
const inputs = {
|
|
a: a3D,
|
|
b: b3D,
|
|
bias: $bias,
|
|
preluActivationWeights: $preluActivationWeights
|
|
};
|
|
const attrs = { transposeA, transposeB, activation, leakyreluAlpha };
|
|
if (bias == null) {
|
|
const customOp = customGrad((a3D2, b3D2, save) => {
|
|
const res = (
|
|
// tslint:disable-next-line: no-unnecessary-type-assertion
|
|
ENGINE.runKernel(_FusedMatMul, inputs, attrs)
|
|
);
|
|
save([a3D2, b3D2, res]);
|
|
return { value: reshape$2(res, outShape), gradFunc: grad };
|
|
});
|
|
return customOp(a3D, b3D);
|
|
} else {
|
|
const customOpWithBias = customGrad((a3D2, b3D2, $bias2, save) => {
|
|
const res = (
|
|
// tslint:disable-next-line: no-unnecessary-type-assertion
|
|
ENGINE.runKernel(_FusedMatMul, inputs, attrs)
|
|
);
|
|
save([a3D2, b3D2, res, $bias2]);
|
|
return { value: reshape$2(res, outShape), gradFunc: grad };
|
|
});
|
|
return customOpWithBias(a3D, b3D, $bias);
|
|
}
|
|
}
|
|
const matMul = /* @__PURE__ */ op({ fusedMatMul_ });
|
|
const fused_ops = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
conv2d: conv2d$1,
|
|
depthwiseConv2d,
|
|
matMul
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
function hammingWindow_(windowLength) {
|
|
return cosineWindow(windowLength, 0.54, 0.46);
|
|
}
|
|
const hammingWindow = /* @__PURE__ */ op({ hammingWindow_ });
|
|
function hannWindow_(windowLength) {
|
|
return cosineWindow(windowLength, 0.5, 0.5);
|
|
}
|
|
const hannWindow = /* @__PURE__ */ op({ hannWindow_ });
|
|
function frame_(signal2, frameLength, frameStep, padEnd = false, padValue = 0) {
|
|
let start = 0;
|
|
const output = [];
|
|
while (start + frameLength <= signal2.size) {
|
|
output.push(slice$2(signal2, start, frameLength));
|
|
start += frameStep;
|
|
}
|
|
if (padEnd) {
|
|
while (start < signal2.size) {
|
|
const padLen = start + frameLength - signal2.size;
|
|
const pad2 = concat$2([
|
|
slice$2(signal2, start, frameLength - padLen),
|
|
fill$2([padLen], padValue)
|
|
]);
|
|
output.push(pad2);
|
|
start += frameStep;
|
|
}
|
|
}
|
|
if (output.length === 0) {
|
|
return tensor2d([], [0, frameLength]);
|
|
}
|
|
return reshape$2(concat$2(output), [output.length, frameLength]);
|
|
}
|
|
const frame = /* @__PURE__ */ op({ frame_ });
|
|
function stft_(signal2, frameLength, frameStep, fftLength, windowFn = hannWindow) {
|
|
if (fftLength == null) {
|
|
fftLength = enclosingPowerOfTwo(frameLength);
|
|
}
|
|
const framedSignal = frame(signal2, frameLength, frameStep);
|
|
const windowedSignal = mul(framedSignal, windowFn(frameLength));
|
|
return rfft(windowedSignal, fftLength);
|
|
}
|
|
const stft = /* @__PURE__ */ op({ stft_ });
|
|
function cropAndResize_(image2, boxes, boxInd, cropSize, method = "bilinear", extrapolationValue = 0) {
|
|
const $image = convertToTensor(image2, "image", "cropAndResize");
|
|
const $boxes = convertToTensor(boxes, "boxes", "cropAndResize", "float32");
|
|
const $boxInd = convertToTensor(boxInd, "boxInd", "cropAndResize", "int32");
|
|
const numBoxes = $boxes.shape[0];
|
|
assert($image.rank === 4, () => `Error in cropAndResize: image must be rank 4,but got rank ${$image.rank}.`);
|
|
assert($boxes.rank === 2 && $boxes.shape[1] === 4, () => `Error in cropAndResize: boxes must be have size [${numBoxes},4] but had shape ${$boxes.shape}.`);
|
|
assert($boxInd.rank === 1 && $boxInd.shape[0] === numBoxes, () => `Error in cropAndResize: boxInd must be have size [${numBoxes}] but had shape ${$boxes.shape}.`);
|
|
assert(cropSize.length === 2, () => `Error in cropAndResize: cropSize must be of length 2, but got length ${cropSize.length}.`);
|
|
assert(cropSize[0] >= 1 && cropSize[1] >= 1, () => `cropSize must be atleast [1,1], but was ${cropSize}`);
|
|
assert(method === "bilinear" || method === "nearest", () => `method must be bilinear or nearest, but was ${method}`);
|
|
const inputs = { image: $image, boxes: $boxes, boxInd: $boxInd };
|
|
const attrs = { method, extrapolationValue, cropSize };
|
|
const res = ENGINE.runKernel(CropAndResize, inputs, attrs);
|
|
return res;
|
|
}
|
|
const cropAndResize$2 = /* @__PURE__ */ op({ cropAndResize_ });
|
|
function flipLeftRight_(image2) {
|
|
const $image = convertToTensor(image2, "image", "flipLeftRight", "float32");
|
|
assert($image.rank === 4, () => `Error in flipLeftRight: image must be rank 4,but got rank ${$image.rank}.`);
|
|
const inputs = { image: $image };
|
|
const res = ENGINE.runKernel(FlipLeftRight, inputs, {});
|
|
return res;
|
|
}
|
|
const flipLeftRight = /* @__PURE__ */ op({ flipLeftRight_ });
|
|
function grayscaleToRGB_(image2) {
|
|
const $image = convertToTensor(image2, "image", "grayscaleToRGB");
|
|
const lastDimsIdx = $image.rank - 1;
|
|
const lastDims = $image.shape[lastDimsIdx];
|
|
assert($image.rank >= 2, () => `Error in grayscaleToRGB: images must be at least rank 2, but got rank ${$image.rank}.`);
|
|
assert(lastDims === 1, () => `Error in grayscaleToRGB: last dimension of a grayscale image should be size 1, but got size ${lastDims}.`);
|
|
const reps = new Array($image.rank);
|
|
reps.fill(1, 0, lastDimsIdx);
|
|
reps[lastDimsIdx] = 3;
|
|
return tile$2($image, reps);
|
|
}
|
|
const grayscaleToRGB = /* @__PURE__ */ op({ grayscaleToRGB_ });
|
|
function rgbToGrayscale_(image2) {
|
|
const $image = convertToTensor(image2, "image", "RGBToGrayscale");
|
|
const lastDimsIdx = $image.rank - 1;
|
|
const lastDims = $image.shape[lastDimsIdx];
|
|
assert($image.rank >= 2, () => `Error in RGBToGrayscale: images must be at least rank 2, but got rank ${$image.rank}.`);
|
|
assert(lastDims === 3, () => `Error in RGBToGrayscale: last dimension of an RGB image should be size 3, but got size ${lastDims}.`);
|
|
const origDtype = $image.dtype;
|
|
const fltImage = cast$2($image, "float32");
|
|
const rgbWeights = tensor1d([0.2989, 0.587, 0.114]);
|
|
let grayFloat;
|
|
switch ($image.rank) {
|
|
case 2:
|
|
grayFloat = einsum$2("ij,j->i", fltImage, rgbWeights);
|
|
break;
|
|
case 3:
|
|
grayFloat = einsum$2("ijk,k->ij", fltImage, rgbWeights);
|
|
break;
|
|
case 4:
|
|
grayFloat = einsum$2("ijkl,l->ijk", fltImage, rgbWeights);
|
|
break;
|
|
case 5:
|
|
grayFloat = einsum$2("ijklm,m->ijkl", fltImage, rgbWeights);
|
|
break;
|
|
case 6:
|
|
grayFloat = einsum$2("ijklmn,n->ijklm", fltImage, rgbWeights);
|
|
break;
|
|
default:
|
|
throw new Error("Not a valid tensor rank.");
|
|
}
|
|
grayFloat = expandDims$2(grayFloat, -1);
|
|
return cast$2(grayFloat, origDtype);
|
|
}
|
|
const rgbToGrayscale = /* @__PURE__ */ op({ rgbToGrayscale_ });
|
|
function rotateWithOffset_(image2, radians, fillValue = 0, center = 0.5) {
|
|
const $image = convertToTensor(image2, "image", "rotateWithOffset", "float32");
|
|
assert($image.rank === 4, () => `Error in rotateWithOffset: image must be rank 4,but got rank ${$image.rank}.`);
|
|
const inputs = { image: $image };
|
|
const attrs = { radians, fillValue, center };
|
|
const res = ENGINE.runKernel(RotateWithOffset, inputs, attrs);
|
|
return res;
|
|
}
|
|
const rotateWithOffset = /* @__PURE__ */ op({ rotateWithOffset_ });
|
|
function nonMaxSuppSanityCheck(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma) {
|
|
if (iouThreshold == null) {
|
|
iouThreshold = 0.5;
|
|
}
|
|
if (scoreThreshold == null) {
|
|
scoreThreshold = Number.NEGATIVE_INFINITY;
|
|
}
|
|
if (softNmsSigma == null) {
|
|
softNmsSigma = 0;
|
|
}
|
|
const numBoxes = boxes.shape[0];
|
|
maxOutputSize = Math.min(maxOutputSize, numBoxes);
|
|
assert(0 <= iouThreshold && iouThreshold <= 1, () => `iouThreshold must be in [0, 1], but was '${iouThreshold}'`);
|
|
assert(boxes.rank === 2, () => `boxes must be a 2D tensor, but was of rank '${boxes.rank}'`);
|
|
assert(boxes.shape[1] === 4, () => `boxes must have 4 columns, but 2nd dimension was ${boxes.shape[1]}`);
|
|
assert(scores.rank === 1, () => "scores must be a 1D tensor");
|
|
assert(scores.shape[0] === numBoxes, () => `scores has incompatible shape with boxes. Expected ${numBoxes}, but was ${scores.shape[0]}`);
|
|
assert(0 <= softNmsSigma && softNmsSigma <= 1, () => `softNmsSigma must be in [0, 1], but was '${softNmsSigma}'`);
|
|
return { maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma };
|
|
}
|
|
function nonMaxSuppression_(boxes, scores, maxOutputSize, iouThreshold = 0.5, scoreThreshold = Number.NEGATIVE_INFINITY) {
|
|
const $boxes = convertToTensor(boxes, "boxes", "nonMaxSuppression", "float32");
|
|
const $scores = convertToTensor(scores, "scores", "nonMaxSuppression", "float32");
|
|
const inputs = nonMaxSuppSanityCheck($boxes, $scores, maxOutputSize, iouThreshold, scoreThreshold);
|
|
maxOutputSize = inputs.maxOutputSize;
|
|
iouThreshold = inputs.iouThreshold;
|
|
scoreThreshold = inputs.scoreThreshold;
|
|
const attrs = { maxOutputSize, iouThreshold, scoreThreshold };
|
|
return ENGINE.runKernel(NonMaxSuppressionV3, { boxes: $boxes, scores: $scores }, attrs);
|
|
}
|
|
const nonMaxSuppression = /* @__PURE__ */ op({ nonMaxSuppression_ });
|
|
function binaryInsert(arr, element, comparator) {
|
|
const index = binarySearch(arr, element, comparator);
|
|
const insertionPoint = index < 0 ? -(index + 1) : index;
|
|
arr.splice(insertionPoint, 0, element);
|
|
}
|
|
function binarySearch(arr, target, comparator) {
|
|
return binarySearch_(arr, target, comparator || defaultComparator);
|
|
}
|
|
function defaultComparator(a, b) {
|
|
return a > b ? 1 : a < b ? -1 : 0;
|
|
}
|
|
function binarySearch_(arr, target, comparator) {
|
|
let left = 0;
|
|
let right = arr.length;
|
|
let middle = 0;
|
|
let found = false;
|
|
while (left < right) {
|
|
middle = left + (right - left >>> 1);
|
|
const compareResult = comparator(target, arr[middle]);
|
|
if (compareResult > 0) {
|
|
left = middle + 1;
|
|
} else {
|
|
right = middle;
|
|
found = !compareResult;
|
|
}
|
|
}
|
|
return found ? left : -left - 1;
|
|
}
|
|
function nonMaxSuppressionV3Impl$2(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold) {
|
|
return nonMaxSuppressionImpl_(
|
|
boxes,
|
|
scores,
|
|
maxOutputSize,
|
|
iouThreshold,
|
|
scoreThreshold,
|
|
0
|
|
/* softNmsSigma */
|
|
);
|
|
}
|
|
function nonMaxSuppressionV4Impl$2(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, padToMaxOutputSize) {
|
|
return nonMaxSuppressionImpl_(
|
|
boxes,
|
|
scores,
|
|
maxOutputSize,
|
|
iouThreshold,
|
|
scoreThreshold,
|
|
0,
|
|
false,
|
|
padToMaxOutputSize,
|
|
true
|
|
/* returnValidOutputs */
|
|
);
|
|
}
|
|
function nonMaxSuppressionV5Impl$2(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma) {
|
|
return nonMaxSuppressionImpl_(
|
|
boxes,
|
|
scores,
|
|
maxOutputSize,
|
|
iouThreshold,
|
|
scoreThreshold,
|
|
softNmsSigma,
|
|
true
|
|
/* returnScoresTensor */
|
|
);
|
|
}
|
|
function nonMaxSuppressionImpl_(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma, returnScoresTensor = false, padToMaxOutputSize = false, returnValidOutputs = false) {
|
|
const candidates = [];
|
|
for (let i = 0; i < scores.length; i++) {
|
|
if (scores[i] > scoreThreshold) {
|
|
candidates.push({ score: scores[i], boxIndex: i, suppressBeginIndex: 0 });
|
|
}
|
|
}
|
|
candidates.sort(ascendingComparator);
|
|
const scale2 = softNmsSigma > 0 ? -0.5 / softNmsSigma : 0;
|
|
const selectedIndices = [];
|
|
const selectedScores = [];
|
|
while (selectedIndices.length < maxOutputSize && candidates.length > 0) {
|
|
const candidate = candidates.pop();
|
|
const { score: originalScore, boxIndex, suppressBeginIndex } = candidate;
|
|
if (originalScore < scoreThreshold) {
|
|
break;
|
|
}
|
|
let ignoreCandidate = false;
|
|
for (let j2 = selectedIndices.length - 1; j2 >= suppressBeginIndex; --j2) {
|
|
const iou = intersectionOverUnion(boxes, boxIndex, selectedIndices[j2]);
|
|
if (iou >= iouThreshold) {
|
|
ignoreCandidate = true;
|
|
break;
|
|
}
|
|
candidate.score = candidate.score * suppressWeight(iouThreshold, scale2, iou);
|
|
if (candidate.score <= scoreThreshold) {
|
|
break;
|
|
}
|
|
}
|
|
candidate.suppressBeginIndex = selectedIndices.length;
|
|
if (!ignoreCandidate) {
|
|
if (candidate.score === originalScore) {
|
|
selectedIndices.push(boxIndex);
|
|
selectedScores.push(candidate.score);
|
|
} else if (candidate.score > scoreThreshold) {
|
|
binaryInsert(candidates, candidate, ascendingComparator);
|
|
}
|
|
}
|
|
}
|
|
const validOutputs = selectedIndices.length;
|
|
const elemsToPad = maxOutputSize - validOutputs;
|
|
if (padToMaxOutputSize && elemsToPad > 0) {
|
|
selectedIndices.push(...new Array(elemsToPad).fill(0));
|
|
selectedScores.push(...new Array(elemsToPad).fill(0));
|
|
}
|
|
const result = { selectedIndices };
|
|
if (returnScoresTensor) {
|
|
result["selectedScores"] = selectedScores;
|
|
}
|
|
if (returnValidOutputs) {
|
|
result["validOutputs"] = validOutputs;
|
|
}
|
|
return result;
|
|
}
|
|
function intersectionOverUnion(boxes, i, j2) {
|
|
const iCoord = boxes.subarray(i * 4, i * 4 + 4);
|
|
const jCoord = boxes.subarray(j2 * 4, j2 * 4 + 4);
|
|
const yminI = Math.min(iCoord[0], iCoord[2]);
|
|
const xminI = Math.min(iCoord[1], iCoord[3]);
|
|
const ymaxI = Math.max(iCoord[0], iCoord[2]);
|
|
const xmaxI = Math.max(iCoord[1], iCoord[3]);
|
|
const yminJ = Math.min(jCoord[0], jCoord[2]);
|
|
const xminJ = Math.min(jCoord[1], jCoord[3]);
|
|
const ymaxJ = Math.max(jCoord[0], jCoord[2]);
|
|
const xmaxJ = Math.max(jCoord[1], jCoord[3]);
|
|
const areaI = (ymaxI - yminI) * (xmaxI - xminI);
|
|
const areaJ = (ymaxJ - yminJ) * (xmaxJ - xminJ);
|
|
if (areaI <= 0 || areaJ <= 0) {
|
|
return 0;
|
|
}
|
|
const intersectionYmin = Math.max(yminI, yminJ);
|
|
const intersectionXmin = Math.max(xminI, xminJ);
|
|
const intersectionYmax = Math.min(ymaxI, ymaxJ);
|
|
const intersectionXmax = Math.min(xmaxI, xmaxJ);
|
|
const intersectionArea = Math.max(intersectionYmax - intersectionYmin, 0) * Math.max(intersectionXmax - intersectionXmin, 0);
|
|
return intersectionArea / (areaI + areaJ - intersectionArea);
|
|
}
|
|
function suppressWeight(iouThreshold, scale2, iou) {
|
|
const weight = Math.exp(scale2 * iou * iou);
|
|
return iou <= iouThreshold ? weight : 0;
|
|
}
|
|
function ascendingComparator(c1, c2) {
|
|
return c1.score - c2.score || c1.score === c2.score && c2.boxIndex - c1.boxIndex;
|
|
}
|
|
async function nonMaxSuppressionAsync_(boxes, scores, maxOutputSize, iouThreshold = 0.5, scoreThreshold = Number.NEGATIVE_INFINITY) {
|
|
const $boxes = convertToTensor(boxes, "boxes", "nonMaxSuppressionAsync");
|
|
const $scores = convertToTensor(scores, "scores", "nonMaxSuppressionAsync");
|
|
const inputs = nonMaxSuppSanityCheck($boxes, $scores, maxOutputSize, iouThreshold, scoreThreshold);
|
|
maxOutputSize = inputs.maxOutputSize;
|
|
iouThreshold = inputs.iouThreshold;
|
|
scoreThreshold = inputs.scoreThreshold;
|
|
const boxesAndScores = await Promise.all([$boxes.data(), $scores.data()]);
|
|
const boxesVals = boxesAndScores[0];
|
|
const scoresVals = boxesAndScores[1];
|
|
const { selectedIndices } = nonMaxSuppressionV3Impl$2(boxesVals, scoresVals, maxOutputSize, iouThreshold, scoreThreshold);
|
|
if ($boxes !== boxes) {
|
|
$boxes.dispose();
|
|
}
|
|
if ($scores !== scores) {
|
|
$scores.dispose();
|
|
}
|
|
return tensor1d(selectedIndices, "int32");
|
|
}
|
|
const nonMaxSuppressionAsync = nonMaxSuppressionAsync_;
|
|
function nonMaxSuppressionWithScore_(boxes, scores, maxOutputSize, iouThreshold = 0.5, scoreThreshold = Number.NEGATIVE_INFINITY, softNmsSigma = 0) {
|
|
const $boxes = convertToTensor(boxes, "boxes", "nonMaxSuppression");
|
|
const $scores = convertToTensor(scores, "scores", "nonMaxSuppression");
|
|
const params = nonMaxSuppSanityCheck($boxes, $scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma);
|
|
maxOutputSize = params.maxOutputSize;
|
|
iouThreshold = params.iouThreshold;
|
|
scoreThreshold = params.scoreThreshold;
|
|
softNmsSigma = params.softNmsSigma;
|
|
const inputs = { boxes: $boxes, scores: $scores };
|
|
const attrs = { maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma };
|
|
const result = ENGINE.runKernel(NonMaxSuppressionV5, inputs, attrs);
|
|
return { selectedIndices: result[0], selectedScores: result[1] };
|
|
}
|
|
const nonMaxSuppressionWithScore = /* @__PURE__ */ op({ nonMaxSuppressionWithScore_ });
|
|
async function nonMaxSuppressionWithScoreAsync_(boxes, scores, maxOutputSize, iouThreshold = 0.5, scoreThreshold = Number.NEGATIVE_INFINITY, softNmsSigma = 0) {
|
|
const $boxes = convertToTensor(boxes, "boxes", "nonMaxSuppressionAsync");
|
|
const $scores = convertToTensor(scores, "scores", "nonMaxSuppressionAsync");
|
|
const params = nonMaxSuppSanityCheck($boxes, $scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma);
|
|
maxOutputSize = params.maxOutputSize;
|
|
iouThreshold = params.iouThreshold;
|
|
scoreThreshold = params.scoreThreshold;
|
|
softNmsSigma = params.softNmsSigma;
|
|
const boxesAndScores = await Promise.all([$boxes.data(), $scores.data()]);
|
|
const boxesVals = boxesAndScores[0];
|
|
const scoresVals = boxesAndScores[1];
|
|
const { selectedIndices, selectedScores } = nonMaxSuppressionV5Impl$2(boxesVals, scoresVals, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma);
|
|
if ($boxes !== boxes) {
|
|
$boxes.dispose();
|
|
}
|
|
if ($scores !== scores) {
|
|
$scores.dispose();
|
|
}
|
|
return {
|
|
selectedIndices: tensor1d(selectedIndices, "int32"),
|
|
selectedScores: tensor1d(selectedScores)
|
|
};
|
|
}
|
|
const nonMaxSuppressionWithScoreAsync = nonMaxSuppressionWithScoreAsync_;
|
|
function nonMaxSuppressionPadded_(boxes, scores, maxOutputSize, iouThreshold = 0.5, scoreThreshold = Number.NEGATIVE_INFINITY, padToMaxOutputSize = false) {
|
|
const $boxes = convertToTensor(boxes, "boxes", "nonMaxSuppression");
|
|
const $scores = convertToTensor(scores, "scores", "nonMaxSuppression");
|
|
const params = nonMaxSuppSanityCheck(
|
|
$boxes,
|
|
$scores,
|
|
maxOutputSize,
|
|
iouThreshold,
|
|
scoreThreshold,
|
|
null
|
|
/* softNmsSigma */
|
|
);
|
|
const $maxOutputSize = params.maxOutputSize;
|
|
const $iouThreshold = params.iouThreshold;
|
|
const $scoreThreshold = params.scoreThreshold;
|
|
const inputs = { boxes: $boxes, scores: $scores };
|
|
const attrs = {
|
|
maxOutputSize: $maxOutputSize,
|
|
iouThreshold: $iouThreshold,
|
|
scoreThreshold: $scoreThreshold,
|
|
padToMaxOutputSize
|
|
};
|
|
const result = ENGINE.runKernel(NonMaxSuppressionV4, inputs, attrs);
|
|
return { selectedIndices: result[0], validOutputs: result[1] };
|
|
}
|
|
const nonMaxSuppressionPadded = /* @__PURE__ */ op({ nonMaxSuppressionPadded_ });
|
|
async function nonMaxSuppressionPaddedAsync_(boxes, scores, maxOutputSize, iouThreshold = 0.5, scoreThreshold = Number.NEGATIVE_INFINITY, padToMaxOutputSize = false) {
|
|
const $boxes = convertToTensor(boxes, "boxes", "nonMaxSuppressionAsync");
|
|
const $scores = convertToTensor(scores, "scores", "nonMaxSuppressionAsync");
|
|
const params = nonMaxSuppSanityCheck(
|
|
$boxes,
|
|
$scores,
|
|
maxOutputSize,
|
|
iouThreshold,
|
|
scoreThreshold,
|
|
null
|
|
/* softNmsSigma */
|
|
);
|
|
const $maxOutputSize = params.maxOutputSize;
|
|
const $iouThreshold = params.iouThreshold;
|
|
const $scoreThreshold = params.scoreThreshold;
|
|
const [boxesVals, scoresVals] = await Promise.all([$boxes.data(), $scores.data()]);
|
|
const { selectedIndices, validOutputs } = nonMaxSuppressionV4Impl$2(boxesVals, scoresVals, $maxOutputSize, $iouThreshold, $scoreThreshold, padToMaxOutputSize);
|
|
if ($boxes !== boxes) {
|
|
$boxes.dispose();
|
|
}
|
|
if ($scores !== scores) {
|
|
$scores.dispose();
|
|
}
|
|
return {
|
|
selectedIndices: tensor1d(selectedIndices, "int32"),
|
|
validOutputs: scalar(validOutputs, "int32")
|
|
};
|
|
}
|
|
const nonMaxSuppressionPaddedAsync = nonMaxSuppressionPaddedAsync_;
|
|
function resizeBilinear_(images, size, alignCorners = false, halfPixelCenters = false) {
|
|
const $images = convertToTensor(images, "images", "resizeBilinear");
|
|
assert($images.rank === 3 || $images.rank === 4, () => `Error in resizeBilinear: x must be rank 3 or 4, but got rank ${$images.rank}.`);
|
|
assert(size.length === 2, () => `Error in resizeBilinear: new shape must 2D, but got shape ${size}.`);
|
|
assert(halfPixelCenters === false || alignCorners === false, () => `Error in resizeBilinear: If halfPixelCenters is true, alignCorners must be false.`);
|
|
let batchImages = $images;
|
|
let reshapedTo4D = false;
|
|
if ($images.rank === 3) {
|
|
reshapedTo4D = true;
|
|
batchImages = reshape$2($images, [1, $images.shape[0], $images.shape[1], $images.shape[2]]);
|
|
}
|
|
const inputs = { images: batchImages };
|
|
const attrs = { alignCorners, halfPixelCenters, size };
|
|
const res = ENGINE.runKernel(ResizeBilinear, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const resizeBilinear$2 = /* @__PURE__ */ op({ resizeBilinear_ });
|
|
function resizeNearestNeighbor_(images, size, alignCorners = false, halfPixelCenters = false) {
|
|
const $images = convertToTensor(images, "images", "resizeNearestNeighbor");
|
|
assert($images.rank === 3 || $images.rank === 4, () => `Error in resizeNearestNeighbor: x must be rank 3 or 4, but got rank ${$images.rank}.`);
|
|
assert(size.length === 2, () => `Error in resizeNearestNeighbor: new shape must 2D, but got shape ${size}.`);
|
|
assert($images.dtype === "float32" || $images.dtype === "int32", () => "`images` must have `int32` or `float32` as dtype");
|
|
assert(halfPixelCenters === false || alignCorners === false, () => `Error in resizeNearestNeighbor: If halfPixelCenters is true, alignCorners must be false.`);
|
|
let batchImages = $images;
|
|
let reshapedTo4D = false;
|
|
if ($images.rank === 3) {
|
|
reshapedTo4D = true;
|
|
batchImages = reshape$2($images, [1, $images.shape[0], $images.shape[1], $images.shape[2]]);
|
|
}
|
|
const inputs = { images: batchImages };
|
|
const attrs = { alignCorners, halfPixelCenters, size };
|
|
const res = ENGINE.runKernel(ResizeNearestNeighbor, inputs, attrs);
|
|
if (reshapedTo4D) {
|
|
return reshape$2(res, [res.shape[1], res.shape[2], res.shape[3]]);
|
|
}
|
|
return res;
|
|
}
|
|
const resizeNearestNeighbor$2 = /* @__PURE__ */ op({ resizeNearestNeighbor_ });
|
|
function threshold_(image2, method = "binary", inverted = false, threshValue = 0.5) {
|
|
const $image = convertToTensor(image2, "image", "threshold");
|
|
const RED_INTENCITY_COEF = 0.2989;
|
|
const GREEN_INTENCITY_COEF = 0.587;
|
|
const BLUE_INTENCITY_COEF = 0.114;
|
|
const totalPixelsInImage = $image.shape[0] * $image.shape[1];
|
|
let $threshold = mul(tensor1d([threshValue]), 255);
|
|
let r, g, b, grayscale;
|
|
assert($image.rank === 3, () => `Error in threshold: image must be rank 3,but got rank ${$image.rank}.`);
|
|
assert($image.shape[2] === 3 || $image.shape[2] === 1, () => `Error in threshold: image color channel must be equal to 3 or 1but got ${$image.shape[2]}.`);
|
|
assert($image.dtype === "int32" || $image.dtype === "float32", () => `Error in dtype: image dtype must be int32 or float32,but got dtype ${$image.dtype}.`);
|
|
assert(method === "otsu" || method === "binary", () => `Method must be binary or otsu, but was ${method}`);
|
|
if ($image.shape[2] === 3) {
|
|
[r, g, b] = split$2($image, [1, 1, 1], -1);
|
|
const $r = mul(r, RED_INTENCITY_COEF);
|
|
const $g = mul(g, GREEN_INTENCITY_COEF);
|
|
const $b = mul(b, BLUE_INTENCITY_COEF);
|
|
grayscale = add$1(add$1($r, $g), $b);
|
|
} else {
|
|
grayscale = image2;
|
|
}
|
|
if (method === "otsu") {
|
|
const $histogram = bincount$2(cast$2(round$2(grayscale), "int32"), tensor([]), 256);
|
|
$threshold = otsu($histogram, totalPixelsInImage);
|
|
}
|
|
const invCondition = inverted ? lessEqual$2(grayscale, $threshold) : greater$2(grayscale, $threshold);
|
|
const result = cast$2(mul(invCondition, 255), "int32");
|
|
return result;
|
|
}
|
|
function otsu(histogram, total) {
|
|
let bestThresh = tensor1d([-1]);
|
|
let bestInBetVar = tensor1d([0]);
|
|
let cInBetVar = tensor1d([0]);
|
|
let classFirst, classSecond, meanFirst, meanSec, weightForeground, weightBack;
|
|
for (let index = 0; index < histogram.size - 1; index++) {
|
|
classFirst = slice$2(histogram, 0, index + 1);
|
|
classSecond = slice$2(histogram, index + 1);
|
|
weightForeground = div$1(sum$2(classFirst), total);
|
|
weightBack = div$1(sum$2(classSecond), total);
|
|
const meanFirstDivA = sum$2(mul(classFirst, range$2(0, classFirst.size)));
|
|
meanFirst = div$1(meanFirstDivA, sum$2(classFirst));
|
|
const meanSecFill = fill$2(classSecond.shape, classFirst.size);
|
|
const meanSecAdd = add$1(range$2(0, classSecond.size), meanSecFill);
|
|
const meanSecMul = mul(classSecond, meanSecAdd);
|
|
meanSec = div$1(sum$2(meanSecMul), sum$2(classSecond));
|
|
const cInBetVarSubA = sub$2(meanFirst, meanSec);
|
|
const cInBetVarSubB = sub$2(meanFirst, meanSec);
|
|
const cInBetVarMul = mul(weightForeground, weightBack);
|
|
cInBetVar = mul(mul(cInBetVarMul, cInBetVarSubA), cInBetVarSubB);
|
|
const condition = greater$2(cInBetVar, bestInBetVar);
|
|
bestInBetVar = where(condition, cInBetVar, bestInBetVar);
|
|
bestThresh = where(condition, tensor1d([index]), bestThresh);
|
|
}
|
|
return bestThresh;
|
|
}
|
|
const threshold$1 = /* @__PURE__ */ op({ threshold_ });
|
|
function transform_(image2, transforms, interpolation = "nearest", fillMode = "constant", fillValue = 0, outputShape) {
|
|
const $image = convertToTensor(image2, "image", "transform", "float32");
|
|
const $transforms = convertToTensor(transforms, "transforms", "transform", "float32");
|
|
assert($image.rank === 4, () => `Error in transform: image must be rank 4,but got rank ${$image.rank}.`);
|
|
assert($transforms.rank === 2 && ($transforms.shape[0] === $image.shape[0] || $transforms.shape[0] === 1) && $transforms.shape[1] === 8, () => `Error in transform: Input transform should be batch x 8 or 1 x 8`);
|
|
assert(outputShape == null || outputShape.length === 2, () => `Error in transform: outputShape must be [height, width] or null, but got ${outputShape}.`);
|
|
const inputs = { image: $image, transforms: $transforms };
|
|
const attrs = { interpolation, fillMode, fillValue, outputShape };
|
|
return ENGINE.runKernel(Transform, inputs, attrs);
|
|
}
|
|
const transform$2 = /* @__PURE__ */ op({ transform_ });
|
|
function bandPart_(a, numLower, numUpper) {
|
|
const $a = convertToTensor(a, "a", "bandPart");
|
|
assert($a.rank >= 2, () => `bandPart(): Rank must be at least 2, got ${$a.rank}.`);
|
|
const shape = $a.shape;
|
|
const [M, N2] = $a.shape.slice(-2);
|
|
let $numLower;
|
|
let $numUpper;
|
|
if (typeof numLower === "number") {
|
|
assert(numLower % 1 === 0, () => `bandPart(): numLower must be an integer, got ${numLower}.`);
|
|
assert(numLower <= M, () => `bandPart(): numLower (${numLower}) must not be greater than the number of rows (${M}).`);
|
|
$numLower = convertToTensor(numLower < 0 ? M : numLower, "numLower", "bandPart");
|
|
} else {
|
|
assert(numLower.dtype === "int32", () => `bandPart(): numLower's dtype must be an int32.`);
|
|
$numLower = where(less$2(numLower, 0), M, minimum$2(numLower, M));
|
|
}
|
|
if (typeof numUpper === "number") {
|
|
assert(numUpper % 1 === 0, () => `bandPart(): numUpper must be an integer, got ${numUpper}.`);
|
|
assert(numUpper <= N2, () => `bandPart(): numUpper (${numUpper}) must not be greater than the number of columns (${N2}).`);
|
|
$numUpper = convertToTensor(numUpper < 0 ? N2 : numUpper, "numUpper", "bandPart");
|
|
} else {
|
|
assert(numUpper.dtype === "int32", () => `bandPart(): numUpper's dtype must be an int32.`);
|
|
$numUpper = where(less$2(numUpper, 0), N2, minimum$2(numUpper, N2));
|
|
}
|
|
const i = reshape$2(range$2(0, M, 1, "int32"), [-1, 1]);
|
|
const j2 = range$2(0, N2, 1, "int32");
|
|
const ij = sub$2(i, j2);
|
|
const inBand = logicalAnd$2(lessEqual$2(ij, $numLower), greaterEqual$2(ij, neg$2($numUpper)));
|
|
const zero = zeros$1([M, N2], $a.dtype);
|
|
return reshape$2(stack(unstack(reshape$2($a, [-1, M, N2])).map((mat) => where(inBand, mat, zero))), shape);
|
|
}
|
|
const bandPart = /* @__PURE__ */ op({ bandPart_ });
|
|
function gramSchmidt_(xs) {
|
|
let inputIsTensor2D;
|
|
if (Array.isArray(xs)) {
|
|
inputIsTensor2D = false;
|
|
assert(xs != null && xs.length > 0, () => "Gram-Schmidt process: input must not be null, undefined, or empty");
|
|
const dim = xs[0].shape[0];
|
|
for (let i = 1; i < xs.length; ++i) {
|
|
assert(xs[i].shape[0] === dim, () => `Gram-Schmidt: Non-unique lengths found in the input vectors: (${xs[i].shape[0]} vs. ${dim})`);
|
|
}
|
|
} else {
|
|
inputIsTensor2D = true;
|
|
xs = split$2(xs, xs.shape[0], 0).map((x) => squeeze(x, [0]));
|
|
}
|
|
assert(xs.length <= xs[0].shape[0], () => `Gram-Schmidt: Number of vectors (${xs.length}) exceeds number of dimensions (${xs[0].shape[0]}).`);
|
|
const ys = [];
|
|
const xs1d = xs;
|
|
for (let i = 0; i < xs.length; ++i) {
|
|
ys.push(ENGINE.tidy(() => {
|
|
let x = xs1d[i];
|
|
if (i > 0) {
|
|
for (let j2 = 0; j2 < i; ++j2) {
|
|
const proj = mul(sum$2(mul(ys[j2], x)), ys[j2]);
|
|
x = sub$2(x, proj);
|
|
}
|
|
}
|
|
return div$1(x, norm(x, "euclidean"));
|
|
}));
|
|
}
|
|
if (inputIsTensor2D) {
|
|
return stack(ys, 0);
|
|
} else {
|
|
return ys;
|
|
}
|
|
}
|
|
const gramSchmidt = /* @__PURE__ */ op({ gramSchmidt_ });
|
|
function qr_(x, fullMatrices = false) {
|
|
assert(x.rank >= 2, () => `qr() requires input tensor to have a rank >= 2, but got rank ${x.rank}`);
|
|
if (x.rank === 2) {
|
|
return qr2d(x, fullMatrices);
|
|
} else {
|
|
const outerDimsProd = x.shape.slice(0, x.shape.length - 2).reduce((value, prev) => value * prev);
|
|
const x2ds = unstack(reshape$2(x, [
|
|
outerDimsProd,
|
|
x.shape[x.shape.length - 2],
|
|
x.shape[x.shape.length - 1]
|
|
]), 0);
|
|
const q2ds = [];
|
|
const r2ds = [];
|
|
x2ds.forEach((x2d) => {
|
|
const [q2d, r2d] = qr2d(x2d, fullMatrices);
|
|
q2ds.push(q2d);
|
|
r2ds.push(r2d);
|
|
});
|
|
const q2 = reshape$2(stack(q2ds, 0), x.shape);
|
|
const r = reshape$2(stack(r2ds, 0), x.shape);
|
|
return [q2, r];
|
|
}
|
|
}
|
|
function qr2d(x, fullMatrices = false) {
|
|
return ENGINE.tidy(() => {
|
|
assert(x.shape.length === 2, () => `qr2d() requires a 2D Tensor, but got a ${x.shape.length}D Tensor.`);
|
|
const m = x.shape[0];
|
|
const n = x.shape[1];
|
|
let q2 = eye(m);
|
|
let r = clone(x);
|
|
const one2D = tensor2d([[1]], [1, 1]);
|
|
let w = clone(one2D);
|
|
const iters = m >= n ? n : m;
|
|
for (let j2 = 0; j2 < iters; ++j2) {
|
|
const rTemp = r;
|
|
const wTemp = w;
|
|
const qTemp = q2;
|
|
[w, r, q2] = ENGINE.tidy(() => {
|
|
const rjEnd1 = slice$2(r, [j2, j2], [m - j2, 1]);
|
|
const normX = norm(rjEnd1);
|
|
const rjj = slice$2(r, [j2, j2], [1, 1]);
|
|
const s = where(greater$2(rjj, 0), tensor2d([[-1]]), tensor2d([[1]]));
|
|
const u1 = sub$2(rjj, mul(s, normX));
|
|
const wPre = div$1(rjEnd1, u1);
|
|
if (wPre.shape[0] === 1) {
|
|
w = clone(one2D);
|
|
} else {
|
|
w = concat$2([
|
|
one2D,
|
|
slice$2(wPre, [1, 0], [wPre.shape[0] - 1, wPre.shape[1]])
|
|
], 0);
|
|
}
|
|
const tau = neg$2(div$1(matMul$1(s, u1), normX));
|
|
const rjEndAll = slice$2(r, [j2, 0], [m - j2, n]);
|
|
const tauTimesW = mul(tau, w);
|
|
const wT = transpose$2(w);
|
|
if (j2 === 0) {
|
|
r = sub$2(rjEndAll, matMul$1(tauTimesW, matMul$1(wT, rjEndAll)));
|
|
} else {
|
|
const rTimesTau = sub$2(rjEndAll, matMul$1(tauTimesW, matMul$1(wT, rjEndAll)));
|
|
r = concat$2([slice$2(r, [0, 0], [j2, n]), rTimesTau], 0);
|
|
}
|
|
const tawTimesWT = transpose$2(tauTimesW);
|
|
const qAllJEnd = slice$2(q2, [0, j2], [m, q2.shape[1] - j2]);
|
|
if (j2 === 0) {
|
|
q2 = sub$2(qAllJEnd, matMul$1(matMul$1(qAllJEnd, w), tawTimesWT));
|
|
} else {
|
|
const qTimesTau = sub$2(qAllJEnd, matMul$1(matMul$1(qAllJEnd, w), tawTimesWT));
|
|
q2 = concat$2([slice$2(q2, [0, 0], [m, j2]), qTimesTau], 1);
|
|
}
|
|
return [w, r, q2];
|
|
});
|
|
dispose([rTemp, wTemp, qTemp]);
|
|
}
|
|
if (!fullMatrices && m > n) {
|
|
q2 = slice$2(q2, [0, 0], [m, n]);
|
|
r = slice$2(r, [0, 0], [n, n]);
|
|
}
|
|
return [q2, r];
|
|
});
|
|
}
|
|
const qr = /* @__PURE__ */ op({ qr_ });
|
|
var Reduction;
|
|
(function(Reduction2) {
|
|
Reduction2[Reduction2["NONE"] = 0] = "NONE";
|
|
Reduction2[Reduction2["MEAN"] = 1] = "MEAN";
|
|
Reduction2[Reduction2["SUM"] = 2] = "SUM";
|
|
Reduction2[Reduction2["SUM_BY_NONZERO_WEIGHTS"] = 3] = "SUM_BY_NONZERO_WEIGHTS";
|
|
})(Reduction || (Reduction = {}));
|
|
function computeWeightedLoss_(losses2, weights, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
const $losses = convertToTensor(losses2, "losses", "computeWeightedLoss");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "computeWeightedLoss");
|
|
}
|
|
const weightedLoss = $weights == null ? $losses : mul($losses, $weights);
|
|
if (reduction2 === Reduction.NONE) {
|
|
return weightedLoss;
|
|
}
|
|
if (reduction2 === Reduction.SUM) {
|
|
return sum$2(weightedLoss);
|
|
}
|
|
if (reduction2 === Reduction.MEAN) {
|
|
if ($weights == null) {
|
|
return mean$1(weightedLoss);
|
|
} else {
|
|
const broadcastFactor = $losses.size / $weights.size;
|
|
const result = div$1(sum$2(weightedLoss), sum$2($weights));
|
|
return broadcastFactor > 1 ? div$1(result, scalar(broadcastFactor)) : result;
|
|
}
|
|
}
|
|
if (reduction2 === Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
if ($weights == null) {
|
|
return div$1(sum$2(weightedLoss), scalar($losses.size));
|
|
} else {
|
|
const broadcastedWeights = mul($weights, ones($losses.shape));
|
|
const numNonZeros = cast$2(sum$2(notEqual$2(broadcastedWeights, scalar(0))), "float32");
|
|
return div$1(sum$2(weightedLoss), numNonZeros);
|
|
}
|
|
}
|
|
throw Error(`Unknown reduction: ${reduction2}`);
|
|
}
|
|
const computeWeightedLoss = /* @__PURE__ */ op({ computeWeightedLoss_ });
|
|
function absoluteDifference_(labels, predictions, weights, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
const $labels = convertToTensor(labels, "labels", "absoluteDifference");
|
|
const $predictions = convertToTensor(predictions, "predictions", "absoluteDifference");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "absoluteDifference");
|
|
}
|
|
assertShapesMatch($labels.shape, $predictions.shape, "Error in absoluteDifference: ");
|
|
const losses2 = abs$2(sub$2($labels, $predictions));
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const absoluteDifference = /* @__PURE__ */ op({ absoluteDifference_ });
|
|
function cosineDistance_(labels, predictions, axis, weights, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
const $labels = convertToTensor(labels, "labels", "cosineDistance");
|
|
const $predictions = convertToTensor(predictions, "predictions", "cosineDistance");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "cosineDistance");
|
|
}
|
|
assertShapesMatch($labels.shape, $predictions.shape, "Error in cosineDistance: ");
|
|
const one = scalar(1);
|
|
const losses2 = sub$2(one, sum$2(mul($labels, $predictions), axis, true));
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const cosineDistance = /* @__PURE__ */ op({ cosineDistance_ });
|
|
function hingeLoss_(labels, predictions, weights, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
let $labels = convertToTensor(labels, "labels", "hingeLoss");
|
|
const $predictions = convertToTensor(predictions, "predictions", "hingeLoss");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "hingeLoss");
|
|
}
|
|
assertShapesMatch($labels.shape, $predictions.shape, "Error in hingeLoss: ");
|
|
const one = scalar(1);
|
|
$labels = sub$2(mul(scalar(2), $labels), one);
|
|
const losses2 = relu$2(sub$2(one, mul($labels, $predictions)));
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const hingeLoss = /* @__PURE__ */ op({ hingeLoss_ });
|
|
function huberLoss_(labels, predictions, weights, delta = 1, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
const $labels = convertToTensor(labels, "labels", "huberLoss");
|
|
const $predictions = convertToTensor(predictions, "predictions", "huberLoss");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "huberLoss");
|
|
}
|
|
assertShapesMatch($labels.shape, $predictions.shape, "Error in huberLoss: ");
|
|
const deltaScalar = scalar(delta);
|
|
const error = abs$2(sub$2($predictions, $labels));
|
|
const quadratic = minimum$2(error, deltaScalar);
|
|
const linear = sub$2(error, quadratic);
|
|
const losses2 = add$1(mul(scalar(0.5), square$1(quadratic)), mul(deltaScalar, linear));
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const huberLoss = /* @__PURE__ */ op({ huberLoss_ });
|
|
function logLoss_(labels, predictions, weights, epsilon2 = 1e-7, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
const $labels = convertToTensor(labels, "labels", "logLoss");
|
|
const $predictions = convertToTensor(predictions, "predictions", "logLoss");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "logLoss");
|
|
}
|
|
assertShapesMatch($labels.shape, $predictions.shape, "Error in logLoss: ");
|
|
const one = scalar(1);
|
|
const epsilonScalar = scalar(epsilon2);
|
|
const l1 = neg$2(mul($labels, log$2(add$1($predictions, epsilonScalar))));
|
|
const l2 = mul(sub$2(one, $labels), log$2(add$1(sub$2(one, $predictions), epsilonScalar)));
|
|
const losses2 = sub$2(l1, l2);
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const logLoss = /* @__PURE__ */ op({ logLoss_ });
|
|
function meanSquaredError_(labels, predictions, weights, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
const $labels = convertToTensor(labels, "labels", "meanSquaredError");
|
|
const $predictions = convertToTensor(predictions, "predictions", "meanSquaredError");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "meanSquaredError");
|
|
}
|
|
assertShapesMatch($labels.shape, $predictions.shape, "Error in meanSquaredError: ");
|
|
const losses2 = squaredDifference$2($labels, $predictions);
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const meanSquaredError = /* @__PURE__ */ op({ meanSquaredError_ });
|
|
function sigmoidCrossEntropyWithLogits_(labels, logits) {
|
|
const $labels = convertToTensor(labels, "labels", "sigmoidCrossEntropyWithLogits");
|
|
const $logits = convertToTensor(logits, "logits", "sigmoidCrossEntropyWithLogits");
|
|
assertShapesMatch($labels.shape, $logits.shape, "Error in sigmoidCrossEntropyWithLogits: ");
|
|
const maxOutput = relu$2($logits);
|
|
const outputXTarget = mul($logits, $labels);
|
|
const sigmoidOutput = log1p$2(exp$2(neg$2(abs$2($logits))));
|
|
return add$1(sub$2(maxOutput, outputXTarget), sigmoidOutput);
|
|
}
|
|
function sigmoidCrossEntropy_(multiClassLabels, logits, weights, labelSmoothing = 0, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
let $multiClassLabels = convertToTensor(multiClassLabels, "multiClassLabels", "sigmoidCrossEntropy");
|
|
const $logits = convertToTensor(logits, "logits", "sigmoidCrossEntropy");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "sigmoidCrossEntropy");
|
|
}
|
|
assertShapesMatch($multiClassLabels.shape, $logits.shape, "Error in sigmoidCrossEntropy: ");
|
|
if (labelSmoothing > 0) {
|
|
const labelSmoothingScalar = scalar(labelSmoothing);
|
|
const one = scalar(1);
|
|
const half = scalar(0.5);
|
|
$multiClassLabels = add$1(mul($multiClassLabels, sub$2(one, labelSmoothingScalar)), mul(half, labelSmoothingScalar));
|
|
}
|
|
const losses2 = sigmoidCrossEntropyWithLogits_($multiClassLabels, $logits);
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const sigmoidCrossEntropy = /* @__PURE__ */ op({ sigmoidCrossEntropy_ });
|
|
function softmaxCrossEntropyWithLogits_(labels, logits, dim = -1) {
|
|
if (dim === -1) {
|
|
dim = logits.rank - 1;
|
|
}
|
|
if (dim !== logits.rank - 1) {
|
|
throw Error(`Softmax cross entropy along a non-last dimension is not yet supported. Labels / logits was rank ${logits.rank} and dim was ${dim}`);
|
|
}
|
|
const customOp = customGrad((labels2, logits2, save) => {
|
|
const keepDims = true;
|
|
const lse = logSumExp(logits2, [dim], keepDims);
|
|
const logResult = sub$2(cast$2(logits2, "float32"), lse);
|
|
save([labels2, logResult]);
|
|
const costVector = neg$2(mul(logResult, labels2));
|
|
const value = sum$2(costVector, [dim]);
|
|
const gradFunc = (dy, saved) => {
|
|
const [labels3, logResult2] = saved;
|
|
const dyShape = expandShapeToKeepDim(dy.shape, [dim]);
|
|
return [
|
|
mul(reshape$2(dy, dyShape), sub$2(cast$2(labels3, "float32"), exp$2(logResult2))),
|
|
mul(reshape$2(dy, dyShape), sub$2(exp$2(logResult2), cast$2(labels3, "float32")))
|
|
];
|
|
};
|
|
return { value, gradFunc };
|
|
});
|
|
return customOp(labels, logits);
|
|
}
|
|
function softmaxCrossEntropy_(onehotLabels, logits, weights, labelSmoothing = 0, reduction2 = Reduction.SUM_BY_NONZERO_WEIGHTS) {
|
|
let $onehotLabels = convertToTensor(onehotLabels, "onehotLabels", "softmaxCrossEntropy");
|
|
const $logits = convertToTensor(logits, "logits", "softmaxCrossEntropy");
|
|
let $weights = null;
|
|
if (weights != null) {
|
|
$weights = convertToTensor(weights, "weights", "softmaxCrossEntropy");
|
|
}
|
|
assertShapesMatch($onehotLabels.shape, $logits.shape, "Error in softmaxCrossEntropy: ");
|
|
if (labelSmoothing > 0) {
|
|
const labelSmoothingScalar = scalar(labelSmoothing);
|
|
const one = scalar(1);
|
|
const numClasses = scalar($onehotLabels.shape[1]);
|
|
$onehotLabels = add$1(mul($onehotLabels, sub$2(one, labelSmoothingScalar)), div$1(labelSmoothingScalar, numClasses));
|
|
}
|
|
const losses2 = softmaxCrossEntropyWithLogits_($onehotLabels, $logits);
|
|
return computeWeightedLoss(losses2, $weights, reduction2);
|
|
}
|
|
const softmaxCrossEntropy = /* @__PURE__ */ op({ softmaxCrossEntropy_ });
|
|
function sparseFillEmptyRows_(indices, values, denseShape, defaultValue) {
|
|
const $indices = convertToTensor(indices, "indices", "sparseFillEmptyRows", "int32");
|
|
const $values = convertToTensor(values, "values", "sparseFillEmptyRows");
|
|
const $denseShape = convertToTensor(denseShape, "denseShape", "sparseFillEmptyRows", "int32");
|
|
const $defaultValue = convertToTensor(defaultValue, "defaultValue", "sparseFillEmptyRows", $values.dtype);
|
|
if ($indices.rank !== 2) {
|
|
throw new Error(`Indices should be Tensor2D but received shape
|
|
${$indices.shape}`);
|
|
}
|
|
if ($values.rank !== 1) {
|
|
throw new Error(`Values should be Tensor1D but received shape ${$values.shape}`);
|
|
}
|
|
if ($denseShape.rank !== 1) {
|
|
throw new Error(`Dense shape should be Tensor1D but received shape ${$denseShape.shape}`);
|
|
}
|
|
if ($defaultValue.rank !== 0) {
|
|
throw new Error(`Default value should be a scalar but received shape ${$defaultValue.shape}`);
|
|
}
|
|
const inputs = {
|
|
indices: $indices,
|
|
values: $values,
|
|
denseShape: $denseShape,
|
|
defaultValue: $defaultValue
|
|
};
|
|
const result = ENGINE.runKernel(SparseFillEmptyRows, inputs);
|
|
return {
|
|
outputIndices: result[0],
|
|
outputValues: result[1],
|
|
emptyRowIndicator: result[2],
|
|
reverseIndexMap: result[3]
|
|
};
|
|
}
|
|
const sparseFillEmptyRows$2 = /* @__PURE__ */ op({ sparseFillEmptyRows_ });
|
|
function sparseReshape_(inputIndices, inputShape, newShape) {
|
|
const $inputIndices = convertToTensor(inputIndices, "inputIndices", "sparseReshape", "int32");
|
|
const $inputShape = convertToTensor(inputShape, "inputShape", "sparseReshape", "int32");
|
|
const $newShape = convertToTensor(newShape, "newShape", "sparseReshape", "int32");
|
|
if ($inputIndices.rank !== 2) {
|
|
throw new Error(`Input indices should be Tensor2D but received shape
|
|
${$inputIndices.shape}`);
|
|
}
|
|
if ($inputShape.rank !== 1) {
|
|
throw new Error(`Input shape should be Tensor1D but received shape ${$inputShape.shape}`);
|
|
}
|
|
if ($newShape.rank !== 1) {
|
|
throw new Error(`New shape should be Tensor1D but received shape ${$newShape.shape}`);
|
|
}
|
|
const inputs = {
|
|
inputIndices: $inputIndices,
|
|
inputShape: $inputShape,
|
|
newShape: $newShape
|
|
};
|
|
const result = ENGINE.runKernel(SparseReshape, inputs);
|
|
return { outputIndices: result[0], outputShape: result[1] };
|
|
}
|
|
const sparseReshape$2 = /* @__PURE__ */ op({ sparseReshape_ });
|
|
function sparseSegmentMean_(data, indices, segmentIds) {
|
|
const $data = convertToTensor(data, "data", "sparseSegmentMean");
|
|
const $indices = convertToTensor(indices, "indices", "sparseSegmentMean", "int32");
|
|
const $segmentIds = convertToTensor(segmentIds, "segmentIds", "sparseSegmentMean", "int32");
|
|
if ($data.rank < 1) {
|
|
throw new Error(`Data should be at least 1 dimensional but received scalar`);
|
|
}
|
|
if ($indices.rank !== 1) {
|
|
throw new Error(`Indices should be Tensor1D but received shape
|
|
${$indices.shape}`);
|
|
}
|
|
if ($segmentIds.rank !== 1) {
|
|
throw new Error(`Segment ids should be Tensor1D but received shape
|
|
${$segmentIds.shape}`);
|
|
}
|
|
const inputs = {
|
|
data: $data,
|
|
indices: $indices,
|
|
segmentIds: $segmentIds
|
|
};
|
|
return ENGINE.runKernel(SparseSegmentMean, inputs);
|
|
}
|
|
const sparseSegmentMean$2 = /* @__PURE__ */ op({ sparseSegmentMean_ });
|
|
function sparseSegmentSum_(data, indices, segmentIds) {
|
|
const $data = convertToTensor(data, "data", "sparseSegmentSum");
|
|
const $indices = convertToTensor(indices, "indices", "sparseSegmentSum", "int32");
|
|
const $segmentIds = convertToTensor(segmentIds, "segmentIds", "sparseSegmentSum", "int32");
|
|
if ($data.rank < 1) {
|
|
throw new Error(`Data should be at least 1 dimensional but received scalar`);
|
|
}
|
|
if ($indices.rank !== 1) {
|
|
throw new Error(`Indices should be Tensor1D but received shape
|
|
${$indices.shape}`);
|
|
}
|
|
if ($segmentIds.rank !== 1) {
|
|
throw new Error(`Segment ids should be Tensor1D but received shape
|
|
${$segmentIds.shape}`);
|
|
}
|
|
const inputs = {
|
|
data: $data,
|
|
indices: $indices,
|
|
segmentIds: $segmentIds
|
|
};
|
|
return ENGINE.runKernel(SparseSegmentSum, inputs);
|
|
}
|
|
const sparseSegmentSum$2 = /* @__PURE__ */ op({ sparseSegmentSum_ });
|
|
function stringNGrams_(data, dataSplits, separator, nGramWidths, leftPad, rightPad2, padWidth, preserveShortSequences) {
|
|
const $data = convertToTensor(data, "data", "stringNGrams", "string");
|
|
if ($data.dtype !== "string") {
|
|
throw new Error("Data must be of datatype string");
|
|
}
|
|
if ($data.shape.length !== 1) {
|
|
throw new Error(`Data must be a vector, saw: ${$data.shape}`);
|
|
}
|
|
const $dataSplits = convertToTensor(dataSplits, "dataSplits", "stringNGrams");
|
|
if ($dataSplits.dtype !== "int32") {
|
|
throw new Error("Data splits must be of datatype int32");
|
|
}
|
|
const attrs = {
|
|
separator,
|
|
nGramWidths,
|
|
leftPad,
|
|
rightPad: rightPad2,
|
|
padWidth,
|
|
preserveShortSequences
|
|
};
|
|
const inputs = { data: $data, dataSplits: $dataSplits };
|
|
const result = ENGINE.runKernel(StringNGrams, inputs, attrs);
|
|
return { nGrams: result[0], nGramsSplits: result[1] };
|
|
}
|
|
const stringNGrams$2 = /* @__PURE__ */ op({ stringNGrams_ });
|
|
function stringSplit_(input, delimiter, skipEmpty = true) {
|
|
const $input = convertToTensor(input, "input", "stringSplit", "string");
|
|
const $delimiter = convertToTensor(delimiter, "delimiter", "stringSplit", "string");
|
|
if ($input.rank !== 1) {
|
|
throw new Error(`Input should be Tensor1D but received shape ${$input.shape}`);
|
|
}
|
|
if ($delimiter.rank !== 0) {
|
|
throw new Error(`Delimiter should be a scalar but received shape ${$delimiter.shape}`);
|
|
}
|
|
const attrs = { skipEmpty };
|
|
const inputs = { input: $input, delimiter: $delimiter };
|
|
const result = ENGINE.runKernel(StringSplit, inputs, attrs);
|
|
return { indices: result[0], values: result[1], shape: result[2] };
|
|
}
|
|
const stringSplit$2 = /* @__PURE__ */ op({ stringSplit_ });
|
|
function stringToHashBucketFast_(input, numBuckets) {
|
|
const $input = convertToTensor(input, "input", "stringToHashBucketFast", "string");
|
|
const attrs = { numBuckets };
|
|
if (numBuckets <= 0) {
|
|
throw new Error(`Number of buckets must be at least 1`);
|
|
}
|
|
const inputs = { input: $input };
|
|
return ENGINE.runKernel(StringToHashBucketFast, inputs, attrs);
|
|
}
|
|
const stringToHashBucketFast$2 = /* @__PURE__ */ op({ stringToHashBucketFast_ });
|
|
function staticRegexReplace_(input, pattern, rewrite, replaceGlobal = true) {
|
|
const $input = convertToTensor(input, "input", "staticRegexReplace", "string");
|
|
const attrs = { pattern, rewrite, replaceGlobal };
|
|
return ENGINE.runKernel(StaticRegexReplace, { x: $input }, attrs);
|
|
}
|
|
const staticRegexReplace$2 = /* @__PURE__ */ op({ staticRegexReplace_ });
|
|
const spectral$1 = {
|
|
fft: fft$2,
|
|
ifft: ifft$2,
|
|
rfft,
|
|
irfft
|
|
};
|
|
const signal = {
|
|
hammingWindow,
|
|
hannWindow,
|
|
frame,
|
|
stft
|
|
};
|
|
const image$1 = {
|
|
flipLeftRight,
|
|
grayscaleToRGB,
|
|
resizeNearestNeighbor: resizeNearestNeighbor$2,
|
|
resizeBilinear: resizeBilinear$2,
|
|
rgbToGrayscale,
|
|
rotateWithOffset,
|
|
cropAndResize: cropAndResize$2,
|
|
nonMaxSuppression,
|
|
nonMaxSuppressionAsync,
|
|
nonMaxSuppressionWithScore,
|
|
nonMaxSuppressionWithScoreAsync,
|
|
nonMaxSuppressionPadded,
|
|
nonMaxSuppressionPaddedAsync,
|
|
threshold: threshold$1,
|
|
transform: transform$2
|
|
};
|
|
const linalg = {
|
|
bandPart,
|
|
gramSchmidt,
|
|
qr
|
|
};
|
|
const losses = {
|
|
absoluteDifference,
|
|
computeWeightedLoss,
|
|
cosineDistance,
|
|
hingeLoss,
|
|
huberLoss,
|
|
logLoss,
|
|
meanSquaredError,
|
|
sigmoidCrossEntropy,
|
|
softmaxCrossEntropy
|
|
};
|
|
const sparse$1 = {
|
|
sparseFillEmptyRows: sparseFillEmptyRows$2,
|
|
sparseReshape: sparseReshape$2,
|
|
sparseSegmentMean: sparseSegmentMean$2,
|
|
sparseSegmentSum: sparseSegmentSum$2
|
|
};
|
|
const string$1 = {
|
|
stringNGrams: stringNGrams$2,
|
|
stringSplit: stringSplit$2,
|
|
stringToHashBucketFast: stringToHashBucketFast$2,
|
|
staticRegexReplace: staticRegexReplace$2
|
|
};
|
|
const GLOBAL_CUSTOM_OBJECT = /* @__PURE__ */ new Map();
|
|
const GLOBAL_CUSTOM_NAMES = /* @__PURE__ */ new Map();
|
|
class Serializable {
|
|
/**
|
|
* Return the class name for this class to use in serialization contexts.
|
|
*
|
|
* Generally speaking this will be the same thing that constructor.name
|
|
* would have returned. However, the class name needs to be robust
|
|
* against minification for serialization/deserialization to work properly.
|
|
*
|
|
* There's also places such as initializers.VarianceScaling, where
|
|
* implementation details between different languages led to different
|
|
* class hierarchies and a non-leaf node is used for serialization purposes.
|
|
*/
|
|
getClassName() {
|
|
return this.constructor.className;
|
|
}
|
|
/**
|
|
* Creates an instance of T from a ConfigDict.
|
|
*
|
|
* This works for most descendants of serializable. A few need to
|
|
* provide special handling.
|
|
* @param cls A Constructor for the class to instantiate.
|
|
* @param config The Configuration for the object.
|
|
*/
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config);
|
|
}
|
|
}
|
|
class SerializationMap {
|
|
constructor() {
|
|
this.classNameMap = {};
|
|
}
|
|
/**
|
|
* Returns the singleton instance of the map.
|
|
*/
|
|
static getMap() {
|
|
if (SerializationMap.instance == null) {
|
|
SerializationMap.instance = new SerializationMap();
|
|
}
|
|
return SerializationMap.instance;
|
|
}
|
|
/**
|
|
* Registers the class as serializable.
|
|
*/
|
|
static register(cls) {
|
|
SerializationMap.getMap().classNameMap[cls.className] = [cls, cls.fromConfig];
|
|
}
|
|
}
|
|
function registerClass(cls, pkg, name) {
|
|
assert(cls.className != null, () => `Class being registered does not have the static className property defined.`);
|
|
assert(typeof cls.className === "string", () => `className is required to be a string, but got type ` + typeof cls.className);
|
|
assert(cls.className.length > 0, () => `Class being registered has an empty-string as its className, which is disallowed.`);
|
|
if (typeof pkg === "undefined") {
|
|
pkg = "Custom";
|
|
}
|
|
if (typeof name === "undefined") {
|
|
name = cls.className;
|
|
}
|
|
const className = name;
|
|
const registerName = pkg + ">" + className;
|
|
SerializationMap.register(cls);
|
|
GLOBAL_CUSTOM_OBJECT.set(registerName, cls);
|
|
GLOBAL_CUSTOM_NAMES.set(cls, registerName);
|
|
return cls;
|
|
}
|
|
class Optimizer extends Serializable {
|
|
/**
|
|
* Executes `f()` and minimizes the scalar output of `f()` by computing
|
|
* gradients of y with respect to the list of trainable variables provided by
|
|
* `varList`. If no list is provided, it defaults to all trainable variables.
|
|
*
|
|
* @param f The function to execute and whose output to minimize.
|
|
* @param returnCost Whether to return the scalar cost value produced by
|
|
* executing `f()`.
|
|
* @param varList An optional list of variables to update. If specified, only
|
|
* the trainable variables in varList will be updated by minimize. Defaults to
|
|
* all trainable variables.
|
|
*
|
|
* @doc {heading: 'Training', subheading: 'Optimizers'}
|
|
*/
|
|
minimize(f, returnCost = false, varList) {
|
|
const { value, grads } = this.computeGradients(f, varList);
|
|
if (varList != null) {
|
|
const gradArray = varList.map((v) => ({ name: v.name, tensor: grads[v.name] }));
|
|
this.applyGradients(gradArray);
|
|
} else {
|
|
this.applyGradients(grads);
|
|
}
|
|
dispose(grads);
|
|
if (returnCost) {
|
|
return value;
|
|
} else {
|
|
value.dispose();
|
|
return null;
|
|
}
|
|
}
|
|
/**
|
|
* The number of iterations that this optimizer instance has been invoked for.
|
|
*/
|
|
get iterations() {
|
|
if (this.iterations_ == null) {
|
|
this.iterations_ = 0;
|
|
}
|
|
return this.iterations_;
|
|
}
|
|
incrementIterations() {
|
|
this.iterations_ = this.iterations + 1;
|
|
}
|
|
/**
|
|
* Executes f() and computes the gradient of the scalar output of f() with
|
|
* respect to the list of trainable variables provided by `varList`. If no
|
|
* list is provided, it defaults to all trainable variables.
|
|
*
|
|
* @param f The function to execute and whose output to use for computing
|
|
* gradients with respect to variables.
|
|
* @param varList An optional list of variables to compute gradients with
|
|
* respect to. If specified, only the trainable variables in varList will have
|
|
* gradients computed with respect to. Defaults to all trainable variables.
|
|
*
|
|
* @doc {heading: 'Training', subheading: 'Optimizers'}
|
|
*/
|
|
computeGradients(f, varList) {
|
|
return variableGrads(f, varList);
|
|
}
|
|
/**
|
|
* Dispose the variables (if any) owned by this optimizer instance.
|
|
*/
|
|
dispose() {
|
|
if (this.iterations_ != null) {
|
|
dispose(this.iterations_);
|
|
}
|
|
}
|
|
async saveIterations() {
|
|
if (this.iterations_ == null) {
|
|
this.iterations_ = 0;
|
|
}
|
|
return {
|
|
name: "iter",
|
|
// TODO(cais): Use 'int64' type when available.
|
|
tensor: scalar(this.iterations_, "int32")
|
|
};
|
|
}
|
|
async getWeights() {
|
|
throw new Error("getWeights() is not implemented for this optimizer yet.");
|
|
}
|
|
async setWeights(weightValues) {
|
|
throw new Error(`setWeights() is not implemented for this optimizer class ${this.getClassName()}`);
|
|
}
|
|
/**
|
|
* Extract the first element of the weight values and set it
|
|
* as the iterations counter variable of this instance of optimizer.
|
|
*
|
|
* @param weightValues
|
|
* @returns Weight values with the first element consumed and excluded.
|
|
*/
|
|
async extractIterations(weightValues) {
|
|
this.iterations_ = (await weightValues[0].tensor.data())[0];
|
|
return weightValues.slice(1);
|
|
}
|
|
}
|
|
Object.defineProperty(Optimizer, Symbol.hasInstance, {
|
|
value: (instance) => {
|
|
return instance.minimize != null && instance.computeGradients != null && instance.applyGradients != null;
|
|
}
|
|
});
|
|
class AdadeltaOptimizer extends Optimizer {
|
|
/** @nocollapse */
|
|
static get className() {
|
|
return "Adadelta";
|
|
}
|
|
constructor(learningRate, rho, epsilon2 = null) {
|
|
super();
|
|
this.learningRate = learningRate;
|
|
this.rho = rho;
|
|
this.epsilon = epsilon2;
|
|
this.accumulatedGrads = [];
|
|
this.accumulatedUpdates = [];
|
|
if (epsilon2 == null) {
|
|
this.epsilon = ENGINE.backend.epsilon();
|
|
}
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const variableNames = Array.isArray(variableGradients) ? variableGradients.map((item) => item.name) : Object.keys(variableGradients);
|
|
variableNames.forEach((name, i) => {
|
|
const value = ENGINE.registeredVariables[name];
|
|
const trainable = false;
|
|
if (this.accumulatedGrads[i] == null) {
|
|
this.accumulatedGrads[i] = {
|
|
originalName: `${name}/accum_grad`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
if (this.accumulatedUpdates[i] == null) {
|
|
this.accumulatedUpdates[i] = {
|
|
originalName: `${name}/accum_var`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
const accumulatedGrad = this.accumulatedGrads[i].variable;
|
|
const accumulatedUpdate = this.accumulatedUpdates[i].variable;
|
|
tidy(() => {
|
|
const newAccumulatedGrad = add$1(mul(accumulatedGrad, this.rho), mul(square$1(gradient), 1 - this.rho));
|
|
const updates = mul(div$1(sqrt$2(add$1(accumulatedUpdate, this.epsilon)), sqrt$2(add$1(accumulatedGrad, this.epsilon))), gradient);
|
|
const newAccumulatedUpdate = add$1(mul(accumulatedUpdate, this.rho), mul(square$1(updates), 1 - this.rho));
|
|
accumulatedGrad.assign(newAccumulatedGrad);
|
|
accumulatedUpdate.assign(newAccumulatedUpdate);
|
|
const newValue = add$1(mul(updates, -this.learningRate), value);
|
|
value.assign(newValue);
|
|
});
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
dispose() {
|
|
if (this.accumulatedUpdates != null) {
|
|
dispose(this.accumulatedGrads.map((v) => v.variable));
|
|
dispose(this.accumulatedUpdates.map((v) => v.variable));
|
|
}
|
|
}
|
|
async getWeights() {
|
|
const variables = [...this.accumulatedGrads, ...this.accumulatedUpdates];
|
|
return [await this.saveIterations()].concat(variables.map((v) => ({ name: v.originalName, tensor: v.variable })));
|
|
}
|
|
async setWeights(weightValues) {
|
|
weightValues = await this.extractIterations(weightValues);
|
|
const variableCount = weightValues.length / 2;
|
|
const trainable = false;
|
|
this.accumulatedGrads = weightValues.slice(0, variableCount).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
this.accumulatedUpdates = weightValues.slice(variableCount, variableCount * 2).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
}
|
|
getConfig() {
|
|
return {
|
|
"learningRate": this.learningRate,
|
|
"rho": this.rho,
|
|
"epsilon": this.epsilon
|
|
};
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"], config["rho"], config["epsilon"]);
|
|
}
|
|
}
|
|
class AdagradOptimizer extends Optimizer {
|
|
/** @nocollapse */
|
|
static get className() {
|
|
return "Adagrad";
|
|
}
|
|
constructor(learningRate, initialAccumulatorValue = 0.1) {
|
|
super();
|
|
this.learningRate = learningRate;
|
|
this.initialAccumulatorValue = initialAccumulatorValue;
|
|
this.accumulatedGrads = [];
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const variableNames = Array.isArray(variableGradients) ? variableGradients.map((item) => item.name) : Object.keys(variableGradients);
|
|
variableNames.forEach((name, i) => {
|
|
const value = ENGINE.registeredVariables[name];
|
|
if (this.accumulatedGrads[i] == null) {
|
|
const trainable = false;
|
|
this.accumulatedGrads[i] = {
|
|
originalName: `${name}/accumulator`,
|
|
variable: tidy(() => fill$2(value.shape, this.initialAccumulatorValue).variable(trainable))
|
|
};
|
|
}
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
const accumulatedGrad = this.accumulatedGrads[i].variable;
|
|
tidy(() => {
|
|
const newAccumulatedGrad = add$1(accumulatedGrad, square$1(gradient));
|
|
accumulatedGrad.assign(newAccumulatedGrad);
|
|
const newValue = add$1(mul(div$1(gradient, sqrt$2(add$1(newAccumulatedGrad, ENGINE.backend.epsilon()))), -this.learningRate), value);
|
|
value.assign(newValue);
|
|
});
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
dispose() {
|
|
if (this.accumulatedGrads != null) {
|
|
dispose(this.accumulatedGrads.map((v) => v.variable));
|
|
}
|
|
}
|
|
async getWeights() {
|
|
return [await this.saveIterations()].concat(this.accumulatedGrads.map((v) => ({ name: v.originalName, tensor: v.variable })));
|
|
}
|
|
async setWeights(weightValues) {
|
|
weightValues = await this.extractIterations(weightValues);
|
|
const trainable = false;
|
|
this.accumulatedGrads = weightValues.map((v) => ({ originalName: v.name, variable: v.tensor.variable(trainable) }));
|
|
}
|
|
getConfig() {
|
|
return {
|
|
"learningRate": this.learningRate,
|
|
"initialAccumulatorValue": this.initialAccumulatorValue
|
|
};
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"], config["initialAccumulatorValue"]);
|
|
}
|
|
}
|
|
class AdamOptimizer extends Optimizer {
|
|
/** @nocollapse */
|
|
static get className() {
|
|
return "Adam";
|
|
}
|
|
constructor(learningRate, beta1, beta2, epsilon2 = null) {
|
|
super();
|
|
this.learningRate = learningRate;
|
|
this.beta1 = beta1;
|
|
this.beta2 = beta2;
|
|
this.epsilon = epsilon2;
|
|
this.accumulatedFirstMoment = [];
|
|
this.accumulatedSecondMoment = [];
|
|
tidy(() => {
|
|
this.accBeta1 = scalar(beta1).variable();
|
|
this.accBeta2 = scalar(beta2).variable();
|
|
});
|
|
if (epsilon2 == null) {
|
|
this.epsilon = ENGINE.backend.epsilon();
|
|
}
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const varNames = Array.isArray(variableGradients) ? variableGradients.map((v) => v.name) : Object.keys(variableGradients);
|
|
tidy(() => {
|
|
const oneMinusAccBeta1 = sub$2(1, this.accBeta1);
|
|
const oneMinusAccBeta2 = sub$2(1, this.accBeta2);
|
|
varNames.forEach((name, i) => {
|
|
const value = ENGINE.registeredVariables[name];
|
|
const trainable = false;
|
|
if (this.accumulatedFirstMoment[i] == null) {
|
|
this.accumulatedFirstMoment[i] = {
|
|
originalName: `${name}/m`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
if (this.accumulatedSecondMoment[i] == null) {
|
|
this.accumulatedSecondMoment[i] = {
|
|
originalName: `${name}/v`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
const firstMoment = this.accumulatedFirstMoment[i].variable;
|
|
const secondMoment = this.accumulatedSecondMoment[i].variable;
|
|
const newFirstMoment = add$1(mul(firstMoment, this.beta1), mul(gradient, 1 - this.beta1));
|
|
const newSecondMoment = add$1(mul(secondMoment, this.beta2), mul(square$1(gradient), 1 - this.beta2));
|
|
const biasCorrectedFirstMoment = div$1(newFirstMoment, oneMinusAccBeta1);
|
|
const biasCorrectedSecondMoment = div$1(newSecondMoment, oneMinusAccBeta2);
|
|
firstMoment.assign(newFirstMoment);
|
|
secondMoment.assign(newSecondMoment);
|
|
const newValue = add$1(mul(div$1(biasCorrectedFirstMoment, add$1(sqrt$2(biasCorrectedSecondMoment), this.epsilon)), -this.learningRate), value);
|
|
value.assign(newValue);
|
|
});
|
|
this.accBeta1.assign(mul(this.accBeta1, this.beta1));
|
|
this.accBeta2.assign(mul(this.accBeta2, this.beta2));
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
dispose() {
|
|
this.accBeta1.dispose();
|
|
this.accBeta2.dispose();
|
|
if (this.accumulatedFirstMoment != null) {
|
|
dispose(this.accumulatedFirstMoment.map((v) => v.variable));
|
|
}
|
|
if (this.accumulatedSecondMoment != null) {
|
|
dispose(this.accumulatedSecondMoment.map((v) => v.variable));
|
|
}
|
|
}
|
|
async getWeights() {
|
|
const variables = [...this.accumulatedFirstMoment, ...this.accumulatedSecondMoment];
|
|
return [await this.saveIterations()].concat(variables.map((v) => ({ name: v.originalName, tensor: v.variable })));
|
|
}
|
|
async setWeights(weightValues) {
|
|
weightValues = await this.extractIterations(weightValues);
|
|
tidy(() => {
|
|
this.accBeta1.assign(pow$2(this.beta1, this.iterations_ + 1));
|
|
this.accBeta2.assign(pow$2(this.beta2, this.iterations_ + 1));
|
|
});
|
|
const variableCount = weightValues.length / 2;
|
|
const trainable = false;
|
|
this.accumulatedFirstMoment = weightValues.slice(0, variableCount).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
this.accumulatedSecondMoment = weightValues.slice(variableCount, variableCount * 2).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
}
|
|
getConfig() {
|
|
return {
|
|
"learningRate": this.learningRate,
|
|
"beta1": this.beta1,
|
|
"beta2": this.beta2,
|
|
"epsilon": this.epsilon
|
|
};
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"], config["beta1"], config["beta2"], config["epsilon"]);
|
|
}
|
|
}
|
|
class AdamaxOptimizer extends Optimizer {
|
|
/** @nocollapse */
|
|
static get className() {
|
|
return "Adamax";
|
|
}
|
|
constructor(learningRate, beta1, beta2, epsilon2 = null, decay = 0) {
|
|
super();
|
|
this.learningRate = learningRate;
|
|
this.beta1 = beta1;
|
|
this.beta2 = beta2;
|
|
this.epsilon = epsilon2;
|
|
this.decay = decay;
|
|
this.accumulatedFirstMoment = [];
|
|
this.accumulatedWeightedInfNorm = [];
|
|
tidy(() => {
|
|
this.iteration = scalar(0).variable();
|
|
this.accBeta1 = scalar(beta1).variable();
|
|
});
|
|
if (epsilon2 == null) {
|
|
this.epsilon = ENGINE.backend.epsilon();
|
|
}
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const variableNames = Array.isArray(variableGradients) ? variableGradients.map((item) => item.name) : Object.keys(variableGradients);
|
|
tidy(() => {
|
|
const oneMinusAccBeta1 = sub$2(1, this.accBeta1);
|
|
const lr = div$1(-this.learningRate, add$1(mul(this.iteration, this.decay), 1));
|
|
variableNames.forEach((name, i) => {
|
|
const value = ENGINE.registeredVariables[name];
|
|
const trainable = false;
|
|
if (this.accumulatedFirstMoment[i] == null) {
|
|
this.accumulatedFirstMoment[i] = {
|
|
originalName: `${name}/m`,
|
|
variable: zerosLike$2(value).variable(trainable)
|
|
};
|
|
}
|
|
if (this.accumulatedWeightedInfNorm[i] == null) {
|
|
this.accumulatedWeightedInfNorm[i] = {
|
|
originalName: `${name}/v`,
|
|
variable: zerosLike$2(value).variable(trainable)
|
|
};
|
|
}
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
const firstMoment = this.accumulatedFirstMoment[i].variable;
|
|
const weightedInfNorm = this.accumulatedWeightedInfNorm[i].variable;
|
|
const newFirstMoment = add$1(mul(firstMoment, this.beta1), mul(gradient, 1 - this.beta1));
|
|
const ut0 = mul(weightedInfNorm, this.beta2);
|
|
const ut1 = abs$2(gradient);
|
|
const newWeightedInfNorm = maximum$2(ut0, ut1);
|
|
firstMoment.assign(newFirstMoment);
|
|
weightedInfNorm.assign(newWeightedInfNorm);
|
|
const newValue = add$1(mul(div$1(lr, oneMinusAccBeta1), div$1(newFirstMoment, add$1(newWeightedInfNorm, this.epsilon))), value);
|
|
value.assign(newValue);
|
|
});
|
|
this.iteration.assign(add$1(this.iteration, 1));
|
|
this.accBeta1.assign(mul(this.accBeta1, this.beta1));
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
dispose() {
|
|
this.accBeta1.dispose();
|
|
this.iteration.dispose();
|
|
if (this.accumulatedFirstMoment != null) {
|
|
dispose(this.accumulatedFirstMoment.map((v) => v.variable));
|
|
}
|
|
if (this.accumulatedWeightedInfNorm != null) {
|
|
dispose(this.accumulatedWeightedInfNorm.map((v) => v.variable));
|
|
}
|
|
}
|
|
async getWeights() {
|
|
throw new Error("getWeights() is not implemented for Adamax yet.");
|
|
}
|
|
async setWeights(weightValues) {
|
|
throw new Error("setWeights() is not implemented for Adamax yet.");
|
|
}
|
|
getConfig() {
|
|
return {
|
|
"learningRate": this.learningRate,
|
|
"beta1": this.beta1,
|
|
"beta2": this.beta2,
|
|
"epsilon": this.epsilon,
|
|
"decay": this.decay
|
|
};
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"], config["beta1"], config["beta2"], config["epsilon"], config["decay"]);
|
|
}
|
|
}
|
|
class SGDOptimizer extends Optimizer {
|
|
/** @nocollapse */
|
|
static get className() {
|
|
return "SGD";
|
|
}
|
|
constructor(learningRate) {
|
|
super();
|
|
this.learningRate = learningRate;
|
|
this.setLearningRate(learningRate);
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const varNames = Array.isArray(variableGradients) ? variableGradients.map((v) => v.name) : Object.keys(variableGradients);
|
|
varNames.forEach((name, i) => {
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
const value = ENGINE.registeredVariables[name];
|
|
tidy(() => {
|
|
const newValue = add$1(mul(this.c, gradient), value);
|
|
value.assign(newValue);
|
|
});
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
/**
|
|
* Sets the learning rate of the optimizer.
|
|
*/
|
|
setLearningRate(learningRate) {
|
|
this.learningRate = learningRate;
|
|
if (this.c != null) {
|
|
this.c.dispose();
|
|
}
|
|
this.c = keep(scalar(-learningRate));
|
|
}
|
|
dispose() {
|
|
this.c.dispose();
|
|
}
|
|
async getWeights() {
|
|
return [await this.saveIterations()];
|
|
}
|
|
async setWeights(weightValues) {
|
|
weightValues = await this.extractIterations(weightValues);
|
|
if (weightValues.length !== 0) {
|
|
throw new Error("SGD optimizer does not have settable weights.");
|
|
}
|
|
}
|
|
getConfig() {
|
|
return { "learningRate": this.learningRate };
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"]);
|
|
}
|
|
}
|
|
class MomentumOptimizer extends SGDOptimizer {
|
|
/** @nocollapse */
|
|
// Name matters for Python compatibility.
|
|
static get className() {
|
|
return "Momentum";
|
|
}
|
|
constructor(learningRate, momentum, useNesterov = false) {
|
|
super(learningRate);
|
|
this.learningRate = learningRate;
|
|
this.momentum = momentum;
|
|
this.useNesterov = useNesterov;
|
|
this.accumulations = [];
|
|
this.m = scalar(this.momentum);
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const variableNames = Array.isArray(variableGradients) ? variableGradients.map((item) => item.name) : Object.keys(variableGradients);
|
|
variableNames.forEach((name, i) => {
|
|
const value = ENGINE.registeredVariables[name];
|
|
if (this.accumulations[i] == null) {
|
|
const trainable = false;
|
|
this.accumulations[i] = {
|
|
originalName: `${name}/momentum`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
const accumulation = this.accumulations[i].variable;
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
tidy(() => {
|
|
let newValue;
|
|
const newAccumulation = add$1(mul(this.m, accumulation), gradient);
|
|
if (this.useNesterov) {
|
|
newValue = add$1(mul(this.c, add$1(gradient, mul(newAccumulation, this.m))), value);
|
|
} else {
|
|
newValue = add$1(mul(this.c, newAccumulation), value);
|
|
}
|
|
accumulation.assign(newAccumulation);
|
|
value.assign(newValue);
|
|
});
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
dispose() {
|
|
this.m.dispose();
|
|
if (this.accumulations != null) {
|
|
dispose(this.accumulations.map((v) => v.variable));
|
|
}
|
|
}
|
|
/**
|
|
* Sets the momentum of the optimizer.
|
|
*
|
|
* @param momentum
|
|
*/
|
|
setMomentum(momentum) {
|
|
this.momentum = momentum;
|
|
}
|
|
async getWeights() {
|
|
return [await this.saveIterations()].concat(this.accumulations.map((v) => ({ name: v.originalName, tensor: v.variable })));
|
|
}
|
|
async setWeights(weightValues) {
|
|
weightValues = await this.extractIterations(weightValues);
|
|
const trainable = false;
|
|
this.accumulations = weightValues.map((v) => ({ originalName: v.name, variable: v.tensor.variable(trainable) }));
|
|
}
|
|
getConfig() {
|
|
return {
|
|
"learningRate": this.learningRate,
|
|
"momentum": this.momentum,
|
|
"useNesterov": this.useNesterov
|
|
};
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"], config["momentum"], config["useNesterov"]);
|
|
}
|
|
}
|
|
class RMSPropOptimizer extends Optimizer {
|
|
/** @nocollapse */
|
|
static get className() {
|
|
return "RMSProp";
|
|
}
|
|
constructor(learningRate, decay = 0.9, momentum = 0, epsilon2 = null, centered = false) {
|
|
super();
|
|
this.learningRate = learningRate;
|
|
this.decay = decay;
|
|
this.momentum = momentum;
|
|
this.epsilon = epsilon2;
|
|
this.accumulatedMeanSquares = [];
|
|
this.accumulatedMoments = [];
|
|
this.accumulatedMeanGrads = [];
|
|
this.centered = centered;
|
|
if (epsilon2 == null) {
|
|
this.epsilon = ENGINE.backend.epsilon();
|
|
}
|
|
if (learningRate == null) {
|
|
throw new Error(`learningRate for RMSPropOptimizer must be defined.`);
|
|
}
|
|
}
|
|
applyGradients(variableGradients) {
|
|
const variableNames = Array.isArray(variableGradients) ? variableGradients.map((item) => item.name) : Object.keys(variableGradients);
|
|
variableNames.forEach((name, i) => {
|
|
const value = ENGINE.registeredVariables[name];
|
|
const trainable = false;
|
|
if (this.accumulatedMeanSquares[i] == null) {
|
|
this.accumulatedMeanSquares[i] = {
|
|
originalName: `${name}/rms`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
if (this.accumulatedMoments[i] == null) {
|
|
this.accumulatedMoments[i] = {
|
|
originalName: `${name}/momentum`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
if (this.accumulatedMeanGrads[i] == null && this.centered) {
|
|
this.accumulatedMeanGrads[i] = {
|
|
originalName: `${name}/mg`,
|
|
variable: tidy(() => zerosLike$2(value).variable(trainable))
|
|
};
|
|
}
|
|
const gradient = Array.isArray(variableGradients) ? variableGradients[i].tensor : variableGradients[name];
|
|
if (gradient == null) {
|
|
return;
|
|
}
|
|
const accumulatedMeanSquare = this.accumulatedMeanSquares[i].variable;
|
|
const accumulatedMoments = this.accumulatedMoments[i].variable;
|
|
tidy(() => {
|
|
const newAccumulatedMeanSquare = add$1(mul(accumulatedMeanSquare, this.decay), mul(square$1(gradient), 1 - this.decay));
|
|
if (this.centered) {
|
|
const accumulatedMeanGrad = this.accumulatedMeanGrads[i].variable;
|
|
const newAccumulatedMeanGrad = add$1(mul(accumulatedMeanGrad, this.decay), mul(gradient, 1 - this.decay));
|
|
const gradContribution = div$1(mul(gradient, this.learningRate), sqrt$2(sub$2(newAccumulatedMeanSquare, add$1(square$1(newAccumulatedMeanGrad), this.epsilon))));
|
|
const newAccumulatedMoments = add$1(mul(accumulatedMoments, this.momentum), gradContribution);
|
|
accumulatedMeanSquare.assign(newAccumulatedMeanSquare);
|
|
accumulatedMeanGrad.assign(newAccumulatedMeanGrad);
|
|
accumulatedMoments.assign(newAccumulatedMoments);
|
|
const newValue = sub$2(value, newAccumulatedMoments);
|
|
value.assign(newValue);
|
|
} else {
|
|
const newAccumulatedMeanSquare2 = add$1(mul(accumulatedMeanSquare, this.decay), mul(square$1(gradient), 1 - this.decay));
|
|
const newAccumulatedMoments = add$1(mul(accumulatedMoments, this.momentum), div$1(mul(gradient, this.learningRate), sqrt$2(add$1(newAccumulatedMeanSquare2, this.epsilon))));
|
|
accumulatedMeanSquare.assign(newAccumulatedMeanSquare2);
|
|
accumulatedMoments.assign(newAccumulatedMoments);
|
|
const newValue = sub$2(value, newAccumulatedMoments);
|
|
value.assign(newValue);
|
|
}
|
|
});
|
|
});
|
|
this.incrementIterations();
|
|
}
|
|
dispose() {
|
|
if (this.accumulatedMeanSquares != null) {
|
|
dispose(this.accumulatedMeanSquares.map((v) => v.variable));
|
|
}
|
|
if (this.accumulatedMeanGrads != null && this.centered) {
|
|
dispose(this.accumulatedMeanGrads.map((v) => v.variable));
|
|
}
|
|
if (this.accumulatedMoments != null) {
|
|
dispose(this.accumulatedMoments.map((v) => v.variable));
|
|
}
|
|
}
|
|
async getWeights() {
|
|
const variables = [...this.accumulatedMeanSquares, ...this.accumulatedMoments];
|
|
if (this.centered) {
|
|
variables.push(...this.accumulatedMeanGrads);
|
|
}
|
|
return [await this.saveIterations()].concat(variables.map((v) => ({ name: v.originalName, tensor: v.variable })));
|
|
}
|
|
async setWeights(weightValues) {
|
|
weightValues = await this.extractIterations(weightValues);
|
|
const variableCount = this.centered ? weightValues.length / 3 : weightValues.length / 2;
|
|
const trainable = false;
|
|
this.accumulatedMeanSquares = weightValues.slice(0, variableCount).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
this.accumulatedMoments = weightValues.slice(variableCount, variableCount * 2).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
if (this.centered) {
|
|
this.accumulatedMeanGrads = weightValues.slice(variableCount * 2, variableCount * 3).map((v) => ({
|
|
originalName: v.name,
|
|
variable: v.tensor.variable(trainable)
|
|
}));
|
|
}
|
|
}
|
|
getConfig() {
|
|
return {
|
|
"learningRate": this.learningRate,
|
|
"decay": this.decay,
|
|
"momentum": this.momentum,
|
|
"epsilon": this.epsilon,
|
|
"centered": this.centered
|
|
};
|
|
}
|
|
/** @nocollapse */
|
|
static fromConfig(cls, config) {
|
|
return new cls(config["learningRate"], config["decay"], config["momentum"], config["epsilon"], config["centered"]);
|
|
}
|
|
}
|
|
const OPTIMIZERS = [
|
|
AdadeltaOptimizer,
|
|
AdagradOptimizer,
|
|
AdamOptimizer,
|
|
AdamaxOptimizer,
|
|
MomentumOptimizer,
|
|
RMSPropOptimizer,
|
|
SGDOptimizer
|
|
];
|
|
function registerOptimizers() {
|
|
for (const optimizer of OPTIMIZERS) {
|
|
registerClass(optimizer);
|
|
}
|
|
}
|
|
const DEFAULT_FILE_NAME_PREFIX = "model";
|
|
const DEFAULT_JSON_EXTENSION_NAME = ".json";
|
|
const DEFAULT_WEIGHT_DATA_EXTENSION_NAME = ".weights.bin";
|
|
function defer(f) {
|
|
return new Promise((resolve) => setTimeout(resolve)).then(f);
|
|
}
|
|
class BrowserDownloads {
|
|
constructor(fileNamePrefix) {
|
|
if (!env().getBool("IS_BROWSER")) {
|
|
throw new Error("browserDownloads() cannot proceed because the current environment is not a browser.");
|
|
}
|
|
if (fileNamePrefix.startsWith(BrowserDownloads.URL_SCHEME)) {
|
|
fileNamePrefix = fileNamePrefix.slice(BrowserDownloads.URL_SCHEME.length);
|
|
}
|
|
if (fileNamePrefix == null || fileNamePrefix.length === 0) {
|
|
fileNamePrefix = DEFAULT_FILE_NAME_PREFIX;
|
|
}
|
|
this.modelJsonFileName = fileNamePrefix + DEFAULT_JSON_EXTENSION_NAME;
|
|
this.weightDataFileName = fileNamePrefix + DEFAULT_WEIGHT_DATA_EXTENSION_NAME;
|
|
}
|
|
async save(modelArtifacts) {
|
|
if (typeof document === "undefined") {
|
|
throw new Error("Browser downloads are not supported in this environment since `document` is not present");
|
|
}
|
|
const weightBuffer = CompositeArrayBuffer.join(modelArtifacts.weightData);
|
|
const weightsURL = window.URL.createObjectURL(new Blob([weightBuffer], { type: "application/octet-stream" }));
|
|
if (modelArtifacts.modelTopology instanceof ArrayBuffer) {
|
|
throw new Error("BrowserDownloads.save() does not support saving model topology in binary formats yet.");
|
|
} else {
|
|
const weightsManifest = [{
|
|
paths: ["./" + this.weightDataFileName],
|
|
weights: modelArtifacts.weightSpecs
|
|
}];
|
|
const modelJSON = getModelJSONForModelArtifacts(modelArtifacts, weightsManifest);
|
|
const modelJsonURL = window.URL.createObjectURL(new Blob([JSON.stringify(modelJSON)], { type: "application/json" }));
|
|
const jsonAnchor = this.modelJsonAnchor == null ? document.createElement("a") : this.modelJsonAnchor;
|
|
jsonAnchor.download = this.modelJsonFileName;
|
|
jsonAnchor.href = modelJsonURL;
|
|
await defer(() => jsonAnchor.dispatchEvent(new MouseEvent("click")));
|
|
if (modelArtifacts.weightData != null) {
|
|
const weightDataAnchor = this.weightDataAnchor == null ? document.createElement("a") : this.weightDataAnchor;
|
|
weightDataAnchor.download = this.weightDataFileName;
|
|
weightDataAnchor.href = weightsURL;
|
|
await defer(() => weightDataAnchor.dispatchEvent(new MouseEvent("click")));
|
|
}
|
|
return { modelArtifactsInfo: getModelArtifactsInfoForJSON(modelArtifacts) };
|
|
}
|
|
}
|
|
}
|
|
BrowserDownloads.URL_SCHEME = "downloads://";
|
|
class BrowserFiles {
|
|
constructor(files) {
|
|
if (files == null || files.length < 1) {
|
|
throw new Error(`When calling browserFiles, at least 1 file is required, but received ${files}`);
|
|
}
|
|
this.jsonFile = files[0];
|
|
this.weightsFiles = files.slice(1);
|
|
}
|
|
async load() {
|
|
return new Promise((resolve, reject) => {
|
|
const jsonReader = new FileReader();
|
|
jsonReader.onload = (event) => {
|
|
const modelJSON = JSON.parse(event.target.result);
|
|
const modelTopology = modelJSON.modelTopology;
|
|
if (modelTopology == null) {
|
|
reject(new Error(`modelTopology field is missing from file ${this.jsonFile.name}`));
|
|
return;
|
|
}
|
|
const weightsManifest = modelJSON.weightsManifest;
|
|
if (weightsManifest == null) {
|
|
reject(new Error(`weightManifest field is missing from file ${this.jsonFile.name}`));
|
|
return;
|
|
}
|
|
if (this.weightsFiles.length === 0) {
|
|
resolve({ modelTopology });
|
|
return;
|
|
}
|
|
const modelArtifactsPromise = getModelArtifactsForJSON(modelJSON, (weightsManifest2) => this.loadWeights(weightsManifest2));
|
|
resolve(modelArtifactsPromise);
|
|
};
|
|
jsonReader.onerror = (error) => reject(`Failed to read model topology and weights manifest JSON from file '${this.jsonFile.name}'. BrowserFiles supports loading Keras-style tf.Model artifacts only.`);
|
|
jsonReader.readAsText(this.jsonFile);
|
|
});
|
|
}
|
|
loadWeights(weightsManifest) {
|
|
const weightSpecs = [];
|
|
const paths = [];
|
|
for (const entry of weightsManifest) {
|
|
weightSpecs.push(...entry.weights);
|
|
paths.push(...entry.paths);
|
|
}
|
|
const pathToFile = this.checkManifestAndWeightFiles(weightsManifest);
|
|
const promises = paths.map((path) => this.loadWeightsFile(path, pathToFile[path]));
|
|
return Promise.all(promises).then((buffers) => [weightSpecs, buffers]);
|
|
}
|
|
loadWeightsFile(path, file) {
|
|
return new Promise((resolve, reject) => {
|
|
const weightFileReader = new FileReader();
|
|
weightFileReader.onload = (event) => {
|
|
const weightData = event.target.result;
|
|
resolve(weightData);
|
|
};
|
|
weightFileReader.onerror = (error) => reject(`Failed to weights data from file of path '${path}'.`);
|
|
weightFileReader.readAsArrayBuffer(file);
|
|
});
|
|
}
|
|
/**
|
|
* Check the compatibility between weights manifest and weight files.
|
|
*/
|
|
checkManifestAndWeightFiles(manifest) {
|
|
const basenames = [];
|
|
const fileNames = this.weightsFiles.map((file) => basename(file.name));
|
|
const pathToFile = {};
|
|
for (const group of manifest) {
|
|
group.paths.forEach((path) => {
|
|
const pathBasename = basename(path);
|
|
if (basenames.indexOf(pathBasename) !== -1) {
|
|
throw new Error(`Duplicate file basename found in weights manifest: '${pathBasename}'`);
|
|
}
|
|
basenames.push(pathBasename);
|
|
if (fileNames.indexOf(pathBasename) === -1) {
|
|
throw new Error(`Weight file with basename '${pathBasename}' is not provided.`);
|
|
} else {
|
|
pathToFile[path] = this.weightsFiles[fileNames.indexOf(pathBasename)];
|
|
}
|
|
});
|
|
}
|
|
if (basenames.length !== this.weightsFiles.length) {
|
|
throw new Error(`Mismatch in the number of files in weights manifest (${basenames.length}) and the number of weight files provided (${this.weightsFiles.length}).`);
|
|
}
|
|
return pathToFile;
|
|
}
|
|
}
|
|
const browserDownloadsRouter = (url) => {
|
|
if (!env().getBool("IS_BROWSER")) {
|
|
return null;
|
|
} else {
|
|
if (!Array.isArray(url) && url.startsWith(BrowserDownloads.URL_SCHEME)) {
|
|
return browserDownloads(url.slice(BrowserDownloads.URL_SCHEME.length));
|
|
} else {
|
|
return null;
|
|
}
|
|
}
|
|
};
|
|
IORouterRegistry.registerSaveRouter(browserDownloadsRouter);
|
|
function browserDownloads(fileNamePrefix = "model") {
|
|
return new BrowserDownloads(fileNamePrefix);
|
|
}
|
|
function browserFiles(files) {
|
|
return new BrowserFiles(files);
|
|
}
|
|
function monitorPromisesProgress(promises, onProgress, startFraction, endFraction) {
|
|
checkPromises(promises);
|
|
startFraction = startFraction == null ? 0 : startFraction;
|
|
endFraction = endFraction == null ? 1 : endFraction;
|
|
checkFraction(startFraction, endFraction);
|
|
let resolvedPromise = 0;
|
|
const registerMonitor = (promise) => {
|
|
promise.then((value) => {
|
|
const fraction = startFraction + ++resolvedPromise / promises.length * (endFraction - startFraction);
|
|
onProgress(fraction);
|
|
return value;
|
|
});
|
|
return promise;
|
|
};
|
|
function checkPromises(promises2) {
|
|
assert(promises2 != null && Array.isArray(promises2) && promises2.length > 0, () => "promises must be a none empty array");
|
|
}
|
|
function checkFraction(startFraction2, endFraction2) {
|
|
assert(startFraction2 >= 0 && startFraction2 <= 1, () => `Progress fraction must be in range [0, 1], but got startFraction ${startFraction2}`);
|
|
assert(endFraction2 >= 0 && endFraction2 <= 1, () => `Progress fraction must be in range [0, 1], but got endFraction ${endFraction2}`);
|
|
assert(endFraction2 >= startFraction2, () => `startFraction must be no more than endFraction, but got startFraction ${startFraction2} and endFraction ${endFraction2}`);
|
|
}
|
|
return Promise.all(promises.map(registerMonitor));
|
|
}
|
|
async function loadWeightsAsArrayBuffer(fetchURLs, loadOptions) {
|
|
if (loadOptions == null) {
|
|
loadOptions = {};
|
|
}
|
|
const fetchFunc = loadOptions.fetchFunc == null ? env().platform.fetch : loadOptions.fetchFunc;
|
|
const requests = fetchURLs.map((fetchURL) => fetchFunc(fetchURL, loadOptions.requestInit, { isBinary: true }));
|
|
const fetchStartFraction = 0;
|
|
const fetchEndFraction = 0.5;
|
|
const responses = loadOptions.onProgress == null ? await Promise.all(requests) : await monitorPromisesProgress(requests, loadOptions.onProgress, fetchStartFraction, fetchEndFraction);
|
|
const bufferPromises = responses.map((response) => response.arrayBuffer());
|
|
const bufferStartFraction = 0.5;
|
|
const bufferEndFraction = 1;
|
|
const buffers = loadOptions.onProgress == null ? await Promise.all(bufferPromises) : await monitorPromisesProgress(bufferPromises, loadOptions.onProgress, bufferStartFraction, bufferEndFraction);
|
|
return buffers;
|
|
}
|
|
function streamWeights(fetchURLs, loadOptions) {
|
|
var _a;
|
|
const fetchFunc = loadOptions.fetchFunc == null ? env().platform.fetch : loadOptions.fetchFunc;
|
|
let fetchIndex = 0;
|
|
let chunkReader;
|
|
(_a = loadOptions.onProgress) === null || _a === void 0 ? void 0 : _a.call(loadOptions, 0);
|
|
return new ReadableStream({
|
|
pull: async (controller) => {
|
|
var _a2;
|
|
while (fetchIndex < fetchURLs.length) {
|
|
if (!chunkReader) {
|
|
const body = (await fetchFunc(fetchURLs[fetchIndex], loadOptions.requestInit, { isBinary: true })).body;
|
|
chunkReader = body.getReader();
|
|
}
|
|
const { done, value } = await chunkReader.read();
|
|
if (done) {
|
|
fetchIndex++;
|
|
chunkReader = void 0;
|
|
(_a2 = loadOptions.onProgress) === null || _a2 === void 0 ? void 0 : _a2.call(loadOptions, fetchIndex / fetchURLs.length);
|
|
continue;
|
|
}
|
|
controller.enqueue(value);
|
|
return;
|
|
}
|
|
controller.close();
|
|
}
|
|
});
|
|
}
|
|
async function loadWeights(manifest, filePathPrefix = "", weightNames, requestInit) {
|
|
const fetchWeights = (fetchUrls) => loadWeightsAsArrayBuffer(fetchUrls, { requestInit });
|
|
const loadWeights2 = weightsLoaderFactory(fetchWeights);
|
|
return loadWeights2(manifest, filePathPrefix, weightNames);
|
|
}
|
|
function weightsLoaderFactory(fetchWeightsFunction) {
|
|
return async (manifest, filePathPrefix = "", weightNames) => {
|
|
const groupIndicesToFetchMap = manifest.map(() => false);
|
|
const groupWeightsToFetch = {};
|
|
const weightsFound = weightNames != null ? weightNames.map(() => false) : [];
|
|
const allManifestWeightNames = [];
|
|
manifest.forEach((manifestGroupConfig, groupIndex) => {
|
|
let groupOffset = 0;
|
|
manifestGroupConfig.weights.forEach((weightsEntry) => {
|
|
const rawDtype = "quantization" in weightsEntry ? weightsEntry.quantization.dtype : weightsEntry.dtype;
|
|
const weightsBytes = DTYPE_VALUE_SIZE_MAP[rawDtype] * sizeFromShape(weightsEntry.shape);
|
|
const enqueueWeightsForFetchingFn = () => {
|
|
groupIndicesToFetchMap[groupIndex] = true;
|
|
if (groupWeightsToFetch[groupIndex] == null) {
|
|
groupWeightsToFetch[groupIndex] = [];
|
|
}
|
|
groupWeightsToFetch[groupIndex].push({
|
|
manifestEntry: weightsEntry,
|
|
groupOffset,
|
|
sizeBytes: weightsBytes
|
|
});
|
|
};
|
|
if (weightNames != null) {
|
|
weightNames.forEach((weightName, weightIndex) => {
|
|
if (weightName === weightsEntry.name) {
|
|
enqueueWeightsForFetchingFn();
|
|
weightsFound[weightIndex] = true;
|
|
}
|
|
});
|
|
} else {
|
|
enqueueWeightsForFetchingFn();
|
|
}
|
|
allManifestWeightNames.push(weightsEntry.name);
|
|
groupOffset += weightsBytes;
|
|
});
|
|
});
|
|
if (!weightsFound.every((found) => found)) {
|
|
const weightsNotFound = weightNames.filter((_, i) => !weightsFound[i]);
|
|
throw new Error(`Could not find weights in manifest with names: ${weightsNotFound.join(", ")}.
|
|
Manifest JSON has weights with names: ${allManifestWeightNames.join(", ")}.`);
|
|
}
|
|
const groupIndicesToFetch = groupIndicesToFetchMap.reduce((accumulator, shouldFetch, i) => {
|
|
if (shouldFetch) {
|
|
accumulator.push(i);
|
|
}
|
|
return accumulator;
|
|
}, []);
|
|
const fetchUrls = [];
|
|
groupIndicesToFetch.forEach((i) => {
|
|
manifest[i].paths.forEach((filepath) => {
|
|
const fetchUrl = filePathPrefix + (!filePathPrefix.endsWith("/") ? "/" : "") + filepath;
|
|
fetchUrls.push(fetchUrl);
|
|
});
|
|
});
|
|
const buffers = await fetchWeightsFunction(fetchUrls);
|
|
const weightsTensorMap = {};
|
|
let bufferIndexOffset = 0;
|
|
groupIndicesToFetch.forEach((i) => {
|
|
const numBuffers = manifest[i].paths.length;
|
|
const weightsBuffer = new CompositeArrayBuffer(buffers.slice(bufferIndexOffset, bufferIndexOffset + numBuffers));
|
|
const weightsEntries = groupWeightsToFetch[i];
|
|
weightsEntries.forEach((weightsEntry) => {
|
|
const byteBuffer = weightsBuffer.slice(weightsEntry.groupOffset, weightsEntry.groupOffset + weightsEntry.sizeBytes);
|
|
const nameToTensorMap = decodeWeights(byteBuffer, [weightsEntry.manifestEntry]);
|
|
for (const name in nameToTensorMap) {
|
|
weightsTensorMap[name] = nameToTensorMap[name];
|
|
}
|
|
});
|
|
bufferIndexOffset += numBuffers;
|
|
});
|
|
return weightsTensorMap;
|
|
};
|
|
}
|
|
const OCTET_STREAM_MIME_TYPE = "application/octet-stream";
|
|
const JSON_TYPE = "application/json";
|
|
class HTTPRequest {
|
|
constructor(path, loadOptions) {
|
|
this.DEFAULT_METHOD = "POST";
|
|
if (loadOptions == null) {
|
|
loadOptions = {};
|
|
}
|
|
this.weightPathPrefix = loadOptions.weightPathPrefix;
|
|
this.weightUrlConverter = loadOptions.weightUrlConverter;
|
|
if (loadOptions.fetchFunc != null) {
|
|
assert(typeof loadOptions.fetchFunc === "function", () => "Must pass a function that matches the signature of `fetch` (see https://developer.mozilla.org/en-US/docs/Web/API/Fetch_API)");
|
|
this.fetch = loadOptions.fetchFunc;
|
|
} else {
|
|
this.fetch = env().platform.fetch;
|
|
}
|
|
assert(path != null && path.length > 0, () => "URL path for http must not be null, undefined or empty.");
|
|
if (Array.isArray(path)) {
|
|
assert(path.length === 2, () => `URL paths for http must have a length of 2, (actual length is ${path.length}).`);
|
|
}
|
|
this.path = path;
|
|
if (loadOptions.requestInit != null && loadOptions.requestInit.body != null) {
|
|
throw new Error("requestInit is expected to have no pre-existing body, but has one.");
|
|
}
|
|
this.requestInit = loadOptions.requestInit || {};
|
|
this.loadOptions = loadOptions;
|
|
}
|
|
async save(modelArtifacts) {
|
|
if (modelArtifacts.modelTopology instanceof ArrayBuffer) {
|
|
throw new Error("BrowserHTTPRequest.save() does not support saving model topology in binary formats yet.");
|
|
}
|
|
const init = Object.assign({ method: this.DEFAULT_METHOD }, this.requestInit);
|
|
init.body = new FormData();
|
|
const weightsManifest = [{
|
|
paths: ["./model.weights.bin"],
|
|
weights: modelArtifacts.weightSpecs
|
|
}];
|
|
const modelTopologyAndWeightManifest = getModelJSONForModelArtifacts(modelArtifacts, weightsManifest);
|
|
init.body.append("model.json", new Blob([JSON.stringify(modelTopologyAndWeightManifest)], { type: JSON_TYPE }), "model.json");
|
|
if (modelArtifacts.weightData != null) {
|
|
const weightBuffer = CompositeArrayBuffer.join(modelArtifacts.weightData);
|
|
init.body.append("model.weights.bin", new Blob([weightBuffer], { type: OCTET_STREAM_MIME_TYPE }), "model.weights.bin");
|
|
}
|
|
const response = await this.fetch(this.path, init);
|
|
if (response.ok) {
|
|
return {
|
|
modelArtifactsInfo: getModelArtifactsInfoForJSON(modelArtifacts),
|
|
responses: [response]
|
|
};
|
|
} else {
|
|
throw new Error(`BrowserHTTPRequest.save() failed due to HTTP response status ${response.status}.`);
|
|
}
|
|
}
|
|
async loadModelJSON() {
|
|
const modelConfigRequest = await this.fetch(this.path, this.requestInit);
|
|
if (!modelConfigRequest.ok) {
|
|
throw new Error(`Request to ${this.path} failed with status code ${modelConfigRequest.status}. Please verify this URL points to the model JSON of the model to load.`);
|
|
}
|
|
let modelJSON;
|
|
try {
|
|
modelJSON = await modelConfigRequest.json();
|
|
} catch (e) {
|
|
let message = `Failed to parse model JSON of response from ${this.path}.`;
|
|
if (this.path.endsWith(".pb")) {
|
|
message += " Your path contains a .pb file extension. Support for .pb models have been removed in TensorFlow.js 1.0 in favor of .json models. You can re-convert your Python TensorFlow model using the TensorFlow.js 1.0 conversion scripts or you can convert your.pb models with the 'pb2json'NPM script in the tensorflow/tfjs-converter repository.";
|
|
} else {
|
|
message += " Please make sure the server is serving valid JSON for this request.";
|
|
}
|
|
throw new Error(message);
|
|
}
|
|
const modelTopology = modelJSON.modelTopology;
|
|
const weightsManifest = modelJSON.weightsManifest;
|
|
if (modelTopology == null && weightsManifest == null) {
|
|
throw new Error(`The JSON from HTTP path ${this.path} contains neither model topology or manifest for weights.`);
|
|
}
|
|
return modelJSON;
|
|
}
|
|
/**
|
|
* Load model artifacts via HTTP request(s).
|
|
*
|
|
* See the documentation to `tf.io.http` for details on the saved
|
|
* artifacts.
|
|
*
|
|
* @returns The loaded model artifacts (if loading succeeds).
|
|
*/
|
|
async load() {
|
|
if (this.loadOptions.streamWeights) {
|
|
return this.loadStream();
|
|
}
|
|
const modelJSON = await this.loadModelJSON();
|
|
return getModelArtifactsForJSON(modelJSON, (weightsManifest) => this.loadWeights(weightsManifest));
|
|
}
|
|
async loadStream() {
|
|
const modelJSON = await this.loadModelJSON();
|
|
const fetchURLs = await this.getWeightUrls(modelJSON.weightsManifest);
|
|
const weightSpecs = getWeightSpecs(modelJSON.weightsManifest);
|
|
const stream = () => streamWeights(fetchURLs, this.loadOptions);
|
|
return Object.assign(Object.assign({}, modelJSON), { weightSpecs, getWeightStream: stream });
|
|
}
|
|
async getWeightUrls(weightsManifest) {
|
|
const weightPath = Array.isArray(this.path) ? this.path[1] : this.path;
|
|
const [prefix, suffix] = parseUrl(weightPath);
|
|
const pathPrefix = this.weightPathPrefix || prefix;
|
|
const fetchURLs = [];
|
|
const urlPromises = [];
|
|
for (const weightsGroup of weightsManifest) {
|
|
for (const path of weightsGroup.paths) {
|
|
if (this.weightUrlConverter != null) {
|
|
urlPromises.push(this.weightUrlConverter(path));
|
|
} else {
|
|
fetchURLs.push(pathPrefix + path + suffix);
|
|
}
|
|
}
|
|
}
|
|
if (this.weightUrlConverter) {
|
|
fetchURLs.push(...await Promise.all(urlPromises));
|
|
}
|
|
return fetchURLs;
|
|
}
|
|
async loadWeights(weightsManifest) {
|
|
const fetchURLs = await this.getWeightUrls(weightsManifest);
|
|
const weightSpecs = getWeightSpecs(weightsManifest);
|
|
const buffers = await loadWeightsAsArrayBuffer(fetchURLs, this.loadOptions);
|
|
return [weightSpecs, buffers];
|
|
}
|
|
}
|
|
HTTPRequest.URL_SCHEME_REGEX = /^https?:\/\//;
|
|
function parseUrl(url) {
|
|
const lastSlash = url.lastIndexOf("/");
|
|
const lastSearchParam = url.lastIndexOf("?");
|
|
const prefix = url.substring(0, lastSlash);
|
|
const suffix = lastSearchParam > lastSlash ? url.substring(lastSearchParam) : "";
|
|
return [prefix + "/", suffix];
|
|
}
|
|
function isHTTPScheme(url) {
|
|
return url.match(HTTPRequest.URL_SCHEME_REGEX) != null;
|
|
}
|
|
const httpRouter = (url, loadOptions) => {
|
|
if (typeof fetch === "undefined" && (loadOptions == null || loadOptions.fetchFunc == null)) {
|
|
return null;
|
|
} else {
|
|
let isHTTP = true;
|
|
if (Array.isArray(url)) {
|
|
isHTTP = url.every((urlItem) => isHTTPScheme(urlItem));
|
|
} else {
|
|
isHTTP = isHTTPScheme(url);
|
|
}
|
|
if (isHTTP) {
|
|
return http(url, loadOptions);
|
|
}
|
|
}
|
|
return null;
|
|
};
|
|
IORouterRegistry.registerSaveRouter(httpRouter);
|
|
IORouterRegistry.registerLoadRouter(httpRouter);
|
|
function http(path, loadOptions) {
|
|
return new HTTPRequest(path, loadOptions);
|
|
}
|
|
function browserHTTPRequest(path, loadOptions) {
|
|
return http(path, loadOptions);
|
|
}
|
|
class PassthroughLoader {
|
|
constructor(modelArtifacts) {
|
|
this.modelArtifacts = modelArtifacts;
|
|
}
|
|
load() {
|
|
return this.modelArtifacts;
|
|
}
|
|
}
|
|
class PassthroughSaver {
|
|
constructor(saveHandler) {
|
|
this.saveHandler = saveHandler;
|
|
}
|
|
save(modelArtifacts) {
|
|
return this.saveHandler(modelArtifacts);
|
|
}
|
|
}
|
|
class PassthroughAsync {
|
|
constructor(handler) {
|
|
if (handler.load) {
|
|
this.load = () => Promise.resolve(handler.load());
|
|
}
|
|
if (handler.save) {
|
|
this.save = (modelArtifacts) => Promise.resolve(handler.save(modelArtifacts));
|
|
}
|
|
}
|
|
}
|
|
function fromMemory(modelArtifacts, weightSpecs, weightData, trainingConfig) {
|
|
const args = arguments;
|
|
return new PassthroughAsync(fromMemorySync(...args));
|
|
}
|
|
function fromMemorySync(modelArtifacts, weightSpecs, weightData, trainingConfig) {
|
|
if (arguments.length === 1) {
|
|
const isModelArtifacts = modelArtifacts.modelTopology != null || modelArtifacts.weightSpecs != null;
|
|
if (isModelArtifacts) {
|
|
return new PassthroughLoader(modelArtifacts);
|
|
} else {
|
|
console.warn("Please call tf.io.fromMemory() with only one argument. The argument should be of type ModelArtifacts. The multi-argument signature of tf.io.fromMemory() has been deprecated and will be removed in a future release.");
|
|
return new PassthroughLoader({ modelTopology: modelArtifacts });
|
|
}
|
|
} else {
|
|
console.warn("Please call tf.io.fromMemory() with only one argument. The argument should be of type ModelArtifacts. The multi-argument signature of tf.io.fromMemory() has been deprecated and will be removed in a future release.");
|
|
return new PassthroughLoader({
|
|
modelTopology: modelArtifacts,
|
|
weightSpecs,
|
|
weightData,
|
|
trainingConfig
|
|
});
|
|
}
|
|
}
|
|
function withSaveHandler(saveHandler) {
|
|
return new PassthroughSaver(saveHandler);
|
|
}
|
|
function withSaveHandlerSync(saveHandler) {
|
|
return new PassthroughSaver(saveHandler);
|
|
}
|
|
const io = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
CompositeArrayBuffer,
|
|
browserFiles,
|
|
browserHTTPRequest,
|
|
concatenateArrayBuffers,
|
|
copyModel,
|
|
decodeWeights,
|
|
decodeWeightsStream,
|
|
encodeWeights,
|
|
fromMemory,
|
|
fromMemorySync,
|
|
getLoadHandlers,
|
|
getModelArtifactsForJSON,
|
|
getModelArtifactsForJSONSync,
|
|
getModelArtifactsInfoForJSON,
|
|
getSaveHandlers,
|
|
getWeightSpecs,
|
|
http,
|
|
isHTTPScheme,
|
|
listModels,
|
|
loadWeights,
|
|
moveModel,
|
|
registerLoadRouter,
|
|
registerSaveRouter,
|
|
removeModel,
|
|
weightsLoaderFactory,
|
|
withSaveHandler,
|
|
withSaveHandlerSync
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
let fromPixels2DContext$1;
|
|
let hasToPixelsWarned = false;
|
|
function fromPixels_(pixels, numChannels = 3) {
|
|
if (numChannels > 4) {
|
|
throw new Error("Cannot construct Tensor with more than 4 channels from pixels.");
|
|
}
|
|
if (pixels == null) {
|
|
throw new Error("pixels passed to tf.browser.fromPixels() can not be null");
|
|
}
|
|
let isPixelData = false;
|
|
let isImageData = false;
|
|
let isVideo = false;
|
|
let isImage = false;
|
|
let isCanvasLike = false;
|
|
let isImageBitmap = false;
|
|
if (pixels.data instanceof Uint8Array) {
|
|
isPixelData = true;
|
|
} else if (typeof ImageData !== "undefined" && pixels instanceof ImageData) {
|
|
isImageData = true;
|
|
} else if (typeof HTMLVideoElement !== "undefined" && pixels instanceof HTMLVideoElement) {
|
|
isVideo = true;
|
|
} else if (typeof HTMLImageElement !== "undefined" && pixels instanceof HTMLImageElement) {
|
|
isImage = true;
|
|
} else if (pixels.getContext != null) {
|
|
isCanvasLike = true;
|
|
} else if (typeof ImageBitmap !== "undefined" && pixels instanceof ImageBitmap) {
|
|
isImageBitmap = true;
|
|
} else {
|
|
throw new Error(`pixels passed to tf.browser.fromPixels() must be either an HTMLVideoElement, HTMLImageElement, HTMLCanvasElement, ImageData in browser, or OffscreenCanvas, ImageData in webworker or {data: Uint32Array, width: number, height: number}, but was ${pixels.constructor.name}`);
|
|
}
|
|
const kernel = getKernel(FromPixels, ENGINE.backendName);
|
|
if (kernel != null) {
|
|
const inputs = { pixels };
|
|
const attrs = { numChannels };
|
|
return ENGINE.runKernel(FromPixels, inputs, attrs);
|
|
}
|
|
const [width, height] = isVideo ? [
|
|
pixels.videoWidth,
|
|
pixels.videoHeight
|
|
] : [pixels.width, pixels.height];
|
|
let vals;
|
|
if (isCanvasLike) {
|
|
vals = // tslint:disable-next-line:no-any
|
|
pixels.getContext("2d").getImageData(0, 0, width, height).data;
|
|
} else if (isImageData || isPixelData) {
|
|
vals = pixels.data;
|
|
} else if (isImage || isVideo || isImageBitmap) {
|
|
if (fromPixels2DContext$1 == null) {
|
|
if (typeof document === "undefined") {
|
|
if (typeof OffscreenCanvas !== "undefined" && typeof OffscreenCanvasRenderingContext2D !== "undefined") {
|
|
fromPixels2DContext$1 = new OffscreenCanvas(1, 1).getContext("2d");
|
|
} else {
|
|
throw new Error("Cannot parse input in current context. Reason: OffscreenCanvas Context2D rendering is not supported.");
|
|
}
|
|
} else {
|
|
fromPixels2DContext$1 = document.createElement("canvas").getContext("2d", { willReadFrequently: true });
|
|
}
|
|
}
|
|
fromPixels2DContext$1.canvas.width = width;
|
|
fromPixels2DContext$1.canvas.height = height;
|
|
fromPixels2DContext$1.drawImage(pixels, 0, 0, width, height);
|
|
vals = fromPixels2DContext$1.getImageData(0, 0, width, height).data;
|
|
}
|
|
let values;
|
|
if (numChannels === 4) {
|
|
values = new Int32Array(vals);
|
|
} else {
|
|
const numPixels = width * height;
|
|
values = new Int32Array(numPixels * numChannels);
|
|
for (let i = 0; i < numPixels; i++) {
|
|
for (let channel = 0; channel < numChannels; ++channel) {
|
|
values[i * numChannels + channel] = vals[i * 4 + channel];
|
|
}
|
|
}
|
|
}
|
|
const outShape = [height, width, numChannels];
|
|
return tensor3d(values, outShape, "int32");
|
|
}
|
|
function validateImgTensor(img) {
|
|
if (img.rank !== 2 && img.rank !== 3) {
|
|
throw new Error(`toPixels only supports rank 2 or 3 tensors, got rank ${img.rank}.`);
|
|
}
|
|
const depth = img.rank === 2 ? 1 : img.shape[2];
|
|
if (depth > 4 || depth === 2) {
|
|
throw new Error(`toPixels only supports depth of size 1, 3 or 4 but got ${depth}`);
|
|
}
|
|
if (img.dtype !== "float32" && img.dtype !== "int32") {
|
|
throw new Error(`Unsupported type for toPixels: ${img.dtype}. Please use float32 or int32 tensors.`);
|
|
}
|
|
}
|
|
async function toPixels(img, canvas) {
|
|
let $img = convertToTensor(img, "img", "toPixels");
|
|
if (!(img instanceof Tensor)) {
|
|
const originalImgTensor = $img;
|
|
$img = cast$2(originalImgTensor, "int32");
|
|
originalImgTensor.dispose();
|
|
}
|
|
validateImgTensor($img);
|
|
const [height, width] = $img.shape.slice(0, 2);
|
|
const depth = $img.rank === 2 ? 1 : $img.shape[2];
|
|
const data = await $img.data();
|
|
const multiplier = $img.dtype === "float32" ? 255 : 1;
|
|
const bytes = new Uint8ClampedArray(width * height * 4);
|
|
for (let i = 0; i < height * width; ++i) {
|
|
const rgba = [0, 0, 0, 255];
|
|
for (let d = 0; d < depth; d++) {
|
|
const value = data[i * depth + d];
|
|
if ($img.dtype === "float32") {
|
|
if (value < 0 || value > 1) {
|
|
throw new Error(`Tensor values for a float32 Tensor must be in the range [0 - 1] but encountered ${value}.`);
|
|
}
|
|
} else if ($img.dtype === "int32") {
|
|
if (value < 0 || value > 255) {
|
|
throw new Error(`Tensor values for a int32 Tensor must be in the range [0 - 255] but encountered ${value}.`);
|
|
}
|
|
}
|
|
if (depth === 1) {
|
|
rgba[0] = value * multiplier;
|
|
rgba[1] = value * multiplier;
|
|
rgba[2] = value * multiplier;
|
|
} else {
|
|
rgba[d] = value * multiplier;
|
|
}
|
|
}
|
|
const j2 = i * 4;
|
|
bytes[j2 + 0] = Math.round(rgba[0]);
|
|
bytes[j2 + 1] = Math.round(rgba[1]);
|
|
bytes[j2 + 2] = Math.round(rgba[2]);
|
|
bytes[j2 + 3] = Math.round(rgba[3]);
|
|
}
|
|
if (canvas != null) {
|
|
if (!hasToPixelsWarned) {
|
|
const kernel = getKernel(Draw, ENGINE.backendName);
|
|
if (kernel != null) {
|
|
console.warn("tf.browser.toPixels is not efficient to draw tensor on canvas. Please try tf.browser.draw instead.");
|
|
hasToPixelsWarned = true;
|
|
}
|
|
}
|
|
canvas.width = width;
|
|
canvas.height = height;
|
|
const ctx = canvas.getContext("2d");
|
|
const imageData = new ImageData(bytes, width, height);
|
|
ctx.putImageData(imageData, 0, 0);
|
|
}
|
|
if ($img !== img) {
|
|
$img.dispose();
|
|
}
|
|
return bytes;
|
|
}
|
|
const fromPixels$1 = /* @__PURE__ */ op({ fromPixels_ });
|
|
function prepareAndValidate(tensor2, indices) {
|
|
const tensorRank = tensor2.shape.length;
|
|
const indicesRank = indices.shape.length;
|
|
if (tensorRank < 1) {
|
|
throw new Error(`tf.gatherND() expects the input to be rank 1 or higher, but the rank was ${tensorRank}.`);
|
|
}
|
|
if (indicesRank < 1) {
|
|
throw new Error(`tf.gatherND() expects the indices to be rank 1 or higher, but the rank was ${indicesRank}.`);
|
|
}
|
|
if (indices.dtype !== "int32") {
|
|
throw new Error(`tf.gatherND() expects the indices to be int32 type, but the dtype was ${indices.dtype}.`);
|
|
}
|
|
if (indices.shape[indicesRank - 1] > tensorRank) {
|
|
throw new Error(`index innermost dimension length must be <= tensor rank; saw: ${indices.shape[indicesRank - 1]} vs. ${tensorRank}`);
|
|
}
|
|
if (sizeFromShape(tensor2.shape) === 0) {
|
|
throw new Error(`Requested more than 0 entries, but input is empty. Input shape: ${tensor2.shape}.`);
|
|
}
|
|
const indicesShape = indices.shape;
|
|
const sliceRank = indicesShape[indicesShape.length - 1];
|
|
let nResult = 1;
|
|
for (let i = 0; i < indicesShape.length - 1; ++i) {
|
|
nResult *= indicesShape[i];
|
|
}
|
|
const inputShape = tensor2.shape;
|
|
const resultShape = indicesShape.slice();
|
|
resultShape.pop();
|
|
let sliceSize = 1;
|
|
for (let i = sliceRank; i < tensorRank; ++i) {
|
|
sliceSize *= inputShape[i];
|
|
resultShape.push(inputShape[i]);
|
|
}
|
|
const strides = [
|
|
...computeStrides(tensor2.shape).map((stride) => stride / sliceSize),
|
|
1
|
|
].slice(0, sliceRank);
|
|
return [resultShape, nResult, sliceSize, strides];
|
|
}
|
|
const NEW_AXIS = -2;
|
|
const SHRINK_AXIS = -1;
|
|
function assertParamsValid(input, begin, size) {
|
|
const inputRank = input.shape.length;
|
|
assert(inputRank === begin.length, () => `Error in slice${inputRank}D: Length of begin ${begin} must match the rank of the array (${inputRank}).`);
|
|
assert(inputRank === size.length, () => `Error in slice${inputRank}D: Length of size ${size} must match the rank of the array (${inputRank}).`);
|
|
for (let i = 0; i < inputRank; ++i) {
|
|
assert(begin[i] + size[i] <= input.shape[i], () => `Error in slice${inputRank}D: begin[${i}] + size[${i}] (${begin[i] + size[i]}) would overflow input.shape[${i}] (${input.shape[i]})`);
|
|
}
|
|
}
|
|
function computeOutShape$2(begin, end, strides) {
|
|
const size = [];
|
|
for (let axis = 0; axis < begin.length; axis++) {
|
|
size[axis] = Math.ceil((end[axis] - begin[axis]) / strides[axis]);
|
|
}
|
|
return size;
|
|
}
|
|
function isSliceContinous(shape, begin, size) {
|
|
let firstNonOneAxis = size.length;
|
|
for (let i = 0; i < size.length; i++) {
|
|
if (size[i] > 1) {
|
|
firstNonOneAxis = i;
|
|
break;
|
|
}
|
|
}
|
|
for (let i = firstNonOneAxis + 1; i < size.length; i++) {
|
|
if (begin[i] > 0 || size[i] !== shape[i]) {
|
|
return false;
|
|
}
|
|
}
|
|
return true;
|
|
}
|
|
function computeFlatOffset(begin, strides) {
|
|
let flatOffset = begin.length > 0 ? begin[begin.length - 1] : 1;
|
|
for (let i = 0; i < begin.length - 1; i++) {
|
|
flatOffset += begin[i] * strides[i];
|
|
}
|
|
return flatOffset;
|
|
}
|
|
function parseSliceParams(x, begin, size) {
|
|
let begin_;
|
|
const xRank = x.shape.length;
|
|
if (typeof begin === "number") {
|
|
begin_ = [begin, ...new Array(xRank - 1).fill(0)];
|
|
} else if (begin.length < xRank) {
|
|
begin_ = begin.concat(new Array(xRank - begin.length).fill(0));
|
|
} else {
|
|
begin_ = begin.slice();
|
|
}
|
|
begin_.forEach((d) => {
|
|
assert(d !== -1, () => "slice() does not support negative begin indexing.");
|
|
});
|
|
let size_;
|
|
if (size == null) {
|
|
size_ = new Array(xRank).fill(-1);
|
|
} else if (typeof size === "number") {
|
|
size_ = [size, ...new Array(xRank - 1).fill(-1)];
|
|
} else if (size.length < xRank) {
|
|
size_ = size.concat(new Array(xRank - size.length).fill(-1));
|
|
} else {
|
|
size_ = size;
|
|
}
|
|
size_ = size_.map((d, i) => {
|
|
if (d >= 0) {
|
|
return d;
|
|
} else {
|
|
assert(d === -1, () => `Negative size values should be exactly -1 but got ${d} for the slice() size at index ${i}.`);
|
|
return x.shape[i] - begin_[i];
|
|
}
|
|
});
|
|
return [begin_, size_];
|
|
}
|
|
function sliceInfo(xShape, begin, end, strides, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask) {
|
|
let stridesNonNull;
|
|
if (strides == null) {
|
|
stridesNonNull = new Array(begin.length);
|
|
stridesNonNull.fill(1);
|
|
} else {
|
|
stridesNonNull = strides;
|
|
}
|
|
if (ellipsisMask != null && (ellipsisMask & ellipsisMask - 1) !== 0) {
|
|
throw new Error("Multiple ellipses in slice is not allowed.");
|
|
}
|
|
let ellipsisSeen = false;
|
|
const sparseSpec = {
|
|
dims: stridesNonNull.length,
|
|
numAddAxisAfterEllipsis: 0,
|
|
begin: begin.slice(),
|
|
end: end.slice(),
|
|
strides: stridesNonNull.slice(),
|
|
beginMask,
|
|
endMask,
|
|
ellipsisMask,
|
|
newAxisMask,
|
|
shrinkAxisMask
|
|
};
|
|
for (let i = 0; i < sparseSpec.dims; i++) {
|
|
if (ellipsisSeen && (1 << i & newAxisMask) !== 0) {
|
|
sparseSpec.numAddAxisAfterEllipsis++;
|
|
}
|
|
if (1 << i & ellipsisMask) {
|
|
ellipsisSeen = true;
|
|
}
|
|
}
|
|
if (!ellipsisSeen) {
|
|
sparseSpec.ellipsisMask |= 1 << sparseSpec.dims;
|
|
sparseSpec.dims++;
|
|
}
|
|
const denseSpec = {
|
|
dims: xShape.length,
|
|
beginMask: 0,
|
|
endMask: 0,
|
|
beginValid: false,
|
|
endValid: false
|
|
};
|
|
buildDenseSpec(sparseSpec, denseSpec);
|
|
let isIdentity = true;
|
|
let sliceDim0 = true;
|
|
let isSimpleSlice = true;
|
|
const processingShape = [];
|
|
const finalShape = [];
|
|
for (let i = 0; i < xShape.length; ++i) {
|
|
if (denseSpec.strides[i] === 0) {
|
|
throw Error(`strides[${i}] must be non-zero`);
|
|
}
|
|
const shrinkI = !!(denseSpec.shrinkAxisMask & 1 << i);
|
|
const dimI = xShape[i];
|
|
if (dimI === -1) {
|
|
processingShape.push(shrinkI ? 1 : -1);
|
|
continue;
|
|
}
|
|
const masks = [denseSpec.beginMask & 1 << i, denseSpec.endMask & 1 << i];
|
|
const validRange = [
|
|
denseSpec.strides[i] > 0 ? 0 : -1,
|
|
denseSpec.strides[i] > 0 ? dimI : dimI - 1
|
|
];
|
|
if (shrinkI && denseSpec.strides[i] <= 0) {
|
|
throw Error("only stride 1 allowed on non-range indexing.");
|
|
}
|
|
isSimpleSlice = isSimpleSlice && denseSpec.strides[i] === 1;
|
|
const beginAndEndMasked = !!(denseSpec.beginMask & 1 << i && denseSpec.endMask & 1 << i);
|
|
if (denseSpec.beginValid && denseSpec.endValid) {
|
|
if (shrinkI) {
|
|
const xFwd = denseSpec.begin[i] < 0 ? dimI + denseSpec.begin[i] : denseSpec.begin[i];
|
|
denseSpec.begin[i] = xFwd;
|
|
denseSpec.end[i] = denseSpec.begin[i] + 1;
|
|
if (xFwd < 0 || xFwd >= dimI) {
|
|
throw Error(`slice index ${denseSpec.begin[i]} of dimension ${i} out of bounds.`);
|
|
}
|
|
} else {
|
|
denseSpec.begin[i] = canonical(denseSpec.begin[i], 0, denseSpec.strides[i], dimI, masks, validRange);
|
|
denseSpec.end[i] = canonical(denseSpec.end[i], 1, denseSpec.strides[i], dimI, masks, validRange);
|
|
}
|
|
const takeAllInDimension = denseSpec.strides[i] === 1 && denseSpec.begin[i] === 0 && denseSpec.end[i] === dimI;
|
|
isIdentity = isIdentity && takeAllInDimension;
|
|
sliceDim0 = sliceDim0 && (i === 0 && denseSpec.strides[i] === 1 || takeAllInDimension);
|
|
} else {
|
|
isIdentity = isIdentity && (denseSpec.strides[i] === 1 && beginAndEndMasked);
|
|
sliceDim0 = sliceDim0 && (i === 0 && denseSpec.strides[i] === 1 || beginAndEndMasked);
|
|
}
|
|
let intervalLength;
|
|
let knownInterval = false;
|
|
if (denseSpec.beginValid && denseSpec.endValid) {
|
|
intervalLength = denseSpec.end[i] - denseSpec.begin[i];
|
|
knownInterval = true;
|
|
} else if (shrinkI) {
|
|
intervalLength = 1;
|
|
knownInterval = true;
|
|
} else if (beginAndEndMasked) {
|
|
if (dimI >= 0) {
|
|
if (denseSpec.strides[i] < 0) {
|
|
intervalLength = -dimI;
|
|
} else {
|
|
intervalLength = dimI;
|
|
}
|
|
knownInterval = true;
|
|
}
|
|
}
|
|
if (knownInterval) {
|
|
let sizeI;
|
|
if (intervalLength === 0 || intervalLength < 0 !== denseSpec.strides[i] < 0) {
|
|
sizeI = 0;
|
|
} else {
|
|
sizeI = Math.trunc(intervalLength / denseSpec.strides[i]) + (intervalLength % denseSpec.strides[i] !== 0 ? 1 : 0);
|
|
}
|
|
processingShape.push(sizeI);
|
|
} else {
|
|
processingShape.push(-1);
|
|
}
|
|
}
|
|
for (let denseDim = 0; denseDim < denseSpec.finalShapeGatherIndices.length; ++denseDim) {
|
|
const gatherIndex = denseSpec.finalShapeGatherIndices[denseDim];
|
|
if (gatherIndex >= 0) {
|
|
finalShape.push(processingShape[gatherIndex]);
|
|
} else if (gatherIndex === NEW_AXIS) {
|
|
finalShape.push(1);
|
|
}
|
|
}
|
|
const finalShapeSparse = finalShape.filter((dim, i) => denseSpec.finalShapeGatherIndices[i] !== NEW_AXIS);
|
|
return {
|
|
finalShapeSparse,
|
|
finalShape,
|
|
isIdentity,
|
|
sliceDim0,
|
|
isSimpleSlice,
|
|
begin: denseSpec.begin,
|
|
end: denseSpec.end,
|
|
strides: denseSpec.strides
|
|
};
|
|
}
|
|
function buildDenseSpec(sparse2, dense) {
|
|
dense.beginMask = 0;
|
|
dense.endMask = 0;
|
|
dense.shrinkAxisMask = 0;
|
|
let fullIndex = 0;
|
|
dense.beginValid = sparse2.begin != null;
|
|
dense.endValid = sparse2.end != null;
|
|
dense.begin = new Array(dense.dims);
|
|
dense.end = new Array(dense.dims);
|
|
dense.strides = new Array(dense.dims);
|
|
dense.finalShapeGatherIndices = [];
|
|
dense.finalShapeGatherIndicesSparse = [];
|
|
dense.inputShapeGatherIndicesSparse = new Array(dense.dims);
|
|
for (let i = 0; i < sparse2.dims; i++) {
|
|
if (1 << i & sparse2.ellipsisMask) {
|
|
const nextIndex = Math.min(dense.dims - (sparse2.dims - i) + 1 + sparse2.numAddAxisAfterEllipsis, dense.dims);
|
|
for (; fullIndex < nextIndex; fullIndex++) {
|
|
dense.begin[fullIndex] = 0;
|
|
dense.end[fullIndex] = 0;
|
|
dense.strides[fullIndex] = 1;
|
|
dense.beginMask |= 1 << fullIndex;
|
|
dense.endMask |= 1 << fullIndex;
|
|
dense.finalShapeGatherIndices.push(fullIndex);
|
|
dense.finalShapeGatherIndicesSparse.push(-1);
|
|
dense.inputShapeGatherIndicesSparse[fullIndex] = i;
|
|
}
|
|
} else if (1 << i & sparse2.newAxisMask) {
|
|
dense.finalShapeGatherIndices.push(NEW_AXIS);
|
|
dense.finalShapeGatherIndicesSparse.push(-1);
|
|
} else {
|
|
if (fullIndex === dense.begin.length) {
|
|
throw Error(`Index out of range using input dim ${fullIndex}; input has only ${dense.dims} dims, ${dense.begin.length}.`);
|
|
}
|
|
if (sparse2.begin != null) {
|
|
dense.begin[fullIndex] = sparse2.begin[i];
|
|
}
|
|
if (sparse2.end != null) {
|
|
dense.end[fullIndex] = sparse2.end[i];
|
|
}
|
|
dense.strides[fullIndex] = sparse2.strides[i];
|
|
if (sparse2.beginMask & 1 << i) {
|
|
dense.beginMask |= 1 << fullIndex;
|
|
}
|
|
if (sparse2.endMask & 1 << i) {
|
|
dense.endMask |= 1 << fullIndex;
|
|
}
|
|
if (sparse2.shrinkAxisMask & 1 << i) {
|
|
dense.finalShapeGatherIndices.push(SHRINK_AXIS);
|
|
dense.finalShapeGatherIndicesSparse.push(-1);
|
|
dense.shrinkAxisMask |= 1 << fullIndex;
|
|
} else {
|
|
dense.finalShapeGatherIndices.push(fullIndex);
|
|
dense.finalShapeGatherIndicesSparse.push(i);
|
|
}
|
|
dense.inputShapeGatherIndicesSparse[fullIndex] = i;
|
|
fullIndex++;
|
|
}
|
|
}
|
|
}
|
|
function canonical(x, c, strideI, dimI, masks, validRange) {
|
|
if (masks[c]) {
|
|
return strideI > 0 ? validRange[c] : validRange[c + 1 & 1];
|
|
} else {
|
|
const xFwd = x < 0 ? dimI + x : x;
|
|
return xFwd < validRange[0] ? validRange[0] : xFwd > validRange[1] ? validRange[1] : xFwd;
|
|
}
|
|
}
|
|
const delayCallback = (() => {
|
|
if (typeof requestAnimationFrame !== "undefined") {
|
|
return requestAnimationFrame;
|
|
} else if (typeof setImmediate !== "undefined") {
|
|
return setImmediate;
|
|
}
|
|
return (f) => f();
|
|
})();
|
|
function nextFrame() {
|
|
return new Promise((resolve) => delayCallback(() => resolve()));
|
|
}
|
|
function assertParamsConsistent(shapes, axis) {
|
|
const rank = shapes[0].length;
|
|
shapes.forEach((shape, i) => {
|
|
assert(shape.length === rank, () => `Error in concat${rank}D: rank of tensors[${i}] must be the same as the rank of the rest (${rank})`);
|
|
});
|
|
assert(axis >= 0 && axis < rank, () => `Error in concat${rank}D: axis must be between 0 and ${rank - 1}.`);
|
|
const firstShape = shapes[0];
|
|
shapes.forEach((shape, i) => {
|
|
for (let r = 0; r < rank; r++) {
|
|
assert(r === axis || shape[r] === firstShape[r], () => `Error in concat${rank}D: Shape of tensors[${i}] (${shape}) does not match the shape of the rest (${firstShape}) along the non-concatenated axis ${i}.`);
|
|
}
|
|
});
|
|
}
|
|
function computeOutShape$1(shapes, axis) {
|
|
const outputShape = shapes[0].slice();
|
|
for (let i = 1; i < shapes.length; i++) {
|
|
outputShape[axis] += shapes[i][axis];
|
|
}
|
|
return outputShape;
|
|
}
|
|
var RowPartitionType$1;
|
|
(function(RowPartitionType2) {
|
|
RowPartitionType2[RowPartitionType2["FIRST_DIM_SIZE"] = 0] = "FIRST_DIM_SIZE";
|
|
RowPartitionType2[RowPartitionType2["VALUE_ROWIDS"] = 1] = "VALUE_ROWIDS";
|
|
RowPartitionType2[RowPartitionType2["ROW_LENGTHS"] = 2] = "ROW_LENGTHS";
|
|
RowPartitionType2[RowPartitionType2["ROW_SPLITS"] = 3] = "ROW_SPLITS";
|
|
RowPartitionType2[RowPartitionType2["ROW_LIMITS"] = 4] = "ROW_LIMITS";
|
|
RowPartitionType2[RowPartitionType2["ROW_STARTS"] = 5] = "ROW_STARTS";
|
|
})(RowPartitionType$1 || (RowPartitionType$1 = {}));
|
|
function combineRaggedTensorToTensorShapes(raggedRank, shape, valueShape) {
|
|
let outputShape = new Array();
|
|
if (valueShape == null && shape == null) {
|
|
return outputShape;
|
|
}
|
|
if (shape == null) {
|
|
while (outputShape.length < raggedRank + valueShape.length) {
|
|
outputShape.push(-1);
|
|
}
|
|
} else {
|
|
outputShape = shape.slice();
|
|
}
|
|
if (valueShape == null) {
|
|
return outputShape;
|
|
}
|
|
if (raggedRank + valueShape.length !== outputShape.length) {
|
|
throw new Error(`rt input.shape and shape=${shape} are incompatible: rt input.rank = ${raggedRank + valueShape.length}, but shape.rank = ${outputShape.length}`);
|
|
}
|
|
for (let i = 1; i < valueShape.length; ++i) {
|
|
const valueDim = valueShape[i];
|
|
const outputShapeDimIndex = outputShape[outputShape.length - valueShape.length + i];
|
|
const outputShapeDim = outputShape[outputShapeDimIndex];
|
|
if (valueDim >= 0) {
|
|
if (outputShapeDim >= 0) {
|
|
if (outputShapeDim !== valueDim) {
|
|
throw new Error(`rt input.shape and shape=${shape} are incompatible: rt input.shape[${i + raggedRank}] = ${valueDim} but shape[${i + raggedRank}] = ${outputShapeDim}`);
|
|
}
|
|
} else {
|
|
outputShape[outputShapeDimIndex] = valueDim;
|
|
}
|
|
}
|
|
}
|
|
return outputShape;
|
|
}
|
|
function getRowPartitionTypesHelper(rowPartitionTypeStrings) {
|
|
const stringToType = {
|
|
"FIRST_DIM_SIZE": RowPartitionType$1.FIRST_DIM_SIZE,
|
|
"VALUE_ROWIDS": RowPartitionType$1.VALUE_ROWIDS,
|
|
"ROW_LENGTHS": RowPartitionType$1.ROW_LENGTHS,
|
|
"ROW_SPLITS": RowPartitionType$1.ROW_SPLITS,
|
|
"ROW_LIMITS": RowPartitionType$1.ROW_LIMITS,
|
|
"ROW_STARTS": RowPartitionType$1.ROW_STARTS
|
|
};
|
|
const result = [];
|
|
for (const typeStr of rowPartitionTypeStrings) {
|
|
if (typeStr in stringToType) {
|
|
result.push(stringToType[typeStr]);
|
|
} else {
|
|
break;
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
function getRaggedRank(rowPartitionTypes) {
|
|
if (rowPartitionTypes.length === 0) {
|
|
return 0;
|
|
}
|
|
if (rowPartitionTypes[0] === RowPartitionType$1.FIRST_DIM_SIZE) {
|
|
return rowPartitionTypes.length - 1;
|
|
}
|
|
return rowPartitionTypes.length;
|
|
}
|
|
function validateDefaultValueShape(defaultValueShape, valueShape) {
|
|
if (defaultValueShape == null || valueShape == null) {
|
|
return;
|
|
}
|
|
const defaultNDims = defaultValueShape.length;
|
|
const valuesNDims = valueShape.length;
|
|
if (defaultNDims >= valuesNDims) {
|
|
throw new Error(`defaultValue.shape=${defaultValueShape} and ragged tensor flatValues.shape=${valueShape}, are incompatible: defaultValue.rank = ${defaultNDims} must be less than ragged tensor input flatValues.rank = ${valuesNDims})`);
|
|
}
|
|
for (let i = 0; i < Math.min(defaultNDims, valuesNDims - 1); ++i) {
|
|
const defaultDim = defaultValueShape[i];
|
|
const valueDim = valueShape[i + 1];
|
|
if (defaultDim >= 0 && valueDim >= 0 && defaultDim !== 1 && defaultDim !== valueDim) {
|
|
throw new Error(`defaultValue.shape=${defaultValueShape}, and ragged tensor input flatValues.shape=${valueShape} are incompatible: defaultValue.shape[${i - defaultValueShape.length}] = ${defaultDim} but ragged tensor input.flatValues.shape[${i - defaultValueShape.length}] = ${valueDim}`);
|
|
}
|
|
}
|
|
}
|
|
const PARALLELIZE_THRESHOLD = 30;
|
|
function computeOptimalWindowSize(inSize) {
|
|
if (inSize <= PARALLELIZE_THRESHOLD) {
|
|
return inSize;
|
|
}
|
|
return nearestDivisor(inSize, Math.floor(Math.sqrt(inSize)));
|
|
}
|
|
function getImageCenter(center, imageHeight, imageWidth) {
|
|
const centerX = imageWidth * (typeof center === "number" ? center : center[0]);
|
|
const centerY = imageHeight * (typeof center === "number" ? center : center[1]);
|
|
return [centerX, centerY];
|
|
}
|
|
function getReshaped(inputShape, blockShape, prod2, batchToSpace = true) {
|
|
let reshaped = [];
|
|
if (batchToSpace) {
|
|
reshaped = reshaped.concat(blockShape.slice(0));
|
|
reshaped.push(inputShape[0] / prod2);
|
|
reshaped = reshaped.concat(inputShape.slice(1));
|
|
} else {
|
|
reshaped = reshaped.concat(inputShape[0]);
|
|
const spatialLength = blockShape.length;
|
|
for (let i = 0; i < spatialLength; ++i) {
|
|
reshaped = reshaped.concat([inputShape[i + 1] / blockShape[i], blockShape[i]]);
|
|
}
|
|
reshaped = reshaped.concat(inputShape.slice(spatialLength + 1));
|
|
}
|
|
return reshaped;
|
|
}
|
|
function getPermuted(reshapedRank, blockShapeRank, batchToSpace = true) {
|
|
const permuted = [];
|
|
if (batchToSpace) {
|
|
permuted.push(blockShapeRank);
|
|
for (let i = blockShapeRank + 1; i < reshapedRank; ++i) {
|
|
if (i <= 2 * blockShapeRank) {
|
|
permuted.push(i);
|
|
permuted.push(i - (blockShapeRank + 1));
|
|
} else {
|
|
permuted.push(i);
|
|
}
|
|
}
|
|
} else {
|
|
const permutedBeforeBatch = [];
|
|
const permutedAfterBatch = [];
|
|
for (let i = 1; i < reshapedRank; ++i) {
|
|
if (i >= blockShapeRank * 2 + 1 || i % 2 === 1) {
|
|
permutedAfterBatch.push(i);
|
|
} else {
|
|
permutedBeforeBatch.push(i);
|
|
}
|
|
}
|
|
permuted.push(...permutedBeforeBatch);
|
|
permuted.push(0);
|
|
permuted.push(...permutedAfterBatch);
|
|
}
|
|
return permuted;
|
|
}
|
|
function getReshapedPermuted(inputShape, blockShape, prod2, batchToSpace = true) {
|
|
const reshapedPermuted = [];
|
|
if (batchToSpace) {
|
|
reshapedPermuted.push(inputShape[0] / prod2);
|
|
} else {
|
|
reshapedPermuted.push(inputShape[0] * prod2);
|
|
}
|
|
for (let i = 1; i < inputShape.length; ++i) {
|
|
if (i <= blockShape.length) {
|
|
if (batchToSpace) {
|
|
reshapedPermuted.push(blockShape[i - 1] * inputShape[i]);
|
|
} else {
|
|
reshapedPermuted.push(inputShape[i] / blockShape[i - 1]);
|
|
}
|
|
} else {
|
|
reshapedPermuted.push(inputShape[i]);
|
|
}
|
|
}
|
|
return reshapedPermuted;
|
|
}
|
|
function getSliceBeginCoords(crops, blockShape) {
|
|
const sliceBeginCoords = [0];
|
|
for (let i = 0; i < blockShape; ++i) {
|
|
sliceBeginCoords.push(crops[i][0]);
|
|
}
|
|
return sliceBeginCoords;
|
|
}
|
|
function getSliceSize(uncroppedShape, crops, blockShape) {
|
|
const sliceSize = uncroppedShape.slice(0, 1);
|
|
for (let i = 0; i < blockShape; ++i) {
|
|
sliceSize.push(uncroppedShape[i + 1] - crops[i][0] - crops[i][1]);
|
|
}
|
|
return sliceSize;
|
|
}
|
|
const SELU_SCALEALPHA = 1.7580993408473768;
|
|
const SELU_SCALE = 1.0507009873554805;
|
|
const ERF_P = 0.3275911;
|
|
const ERF_A1 = 0.254829592;
|
|
const ERF_A2 = -0.284496736;
|
|
const ERF_A3 = 1.421413741;
|
|
const ERF_A4 = -1.453152027;
|
|
const ERF_A5 = 1.061405429;
|
|
function mergeRealAndImagArrays(real2, imag2) {
|
|
if (real2.length !== imag2.length) {
|
|
throw new Error(`Cannot merge real and imag arrays of different lengths. real:${real2.length}, imag: ${imag2.length}.`);
|
|
}
|
|
const result = new Float32Array(real2.length * 2);
|
|
for (let i = 0; i < result.length; i += 2) {
|
|
result[i] = real2[i / 2];
|
|
result[i + 1] = imag2[i / 2];
|
|
}
|
|
return result;
|
|
}
|
|
function splitRealAndImagArrays(complex2) {
|
|
const real2 = new Float32Array(complex2.length / 2);
|
|
const imag2 = new Float32Array(complex2.length / 2);
|
|
for (let i = 0; i < complex2.length; i += 2) {
|
|
real2[i / 2] = complex2[i];
|
|
imag2[i / 2] = complex2[i + 1];
|
|
}
|
|
return { real: real2, imag: imag2 };
|
|
}
|
|
function complexWithEvenIndex(complex2) {
|
|
const len = Math.ceil(complex2.length / 4);
|
|
const real2 = new Float32Array(len);
|
|
const imag2 = new Float32Array(len);
|
|
for (let i = 0; i < complex2.length; i += 4) {
|
|
real2[Math.floor(i / 4)] = complex2[i];
|
|
imag2[Math.floor(i / 4)] = complex2[i + 1];
|
|
}
|
|
return { real: real2, imag: imag2 };
|
|
}
|
|
function complexWithOddIndex(complex2) {
|
|
const len = Math.floor(complex2.length / 4);
|
|
const real2 = new Float32Array(len);
|
|
const imag2 = new Float32Array(len);
|
|
for (let i = 2; i < complex2.length; i += 4) {
|
|
real2[Math.floor(i / 4)] = complex2[i];
|
|
imag2[Math.floor(i / 4)] = complex2[i + 1];
|
|
}
|
|
return { real: real2, imag: imag2 };
|
|
}
|
|
function getComplexWithIndex(complex2, index) {
|
|
const real2 = complex2[index * 2];
|
|
const imag2 = complex2[index * 2 + 1];
|
|
return { real: real2, imag: imag2 };
|
|
}
|
|
function assignToTypedArray(data, real2, imag2, index) {
|
|
data[index * 2] = real2;
|
|
data[index * 2 + 1] = imag2;
|
|
}
|
|
function exponents(n, inverse) {
|
|
const real2 = new Float32Array(n / 2);
|
|
const imag2 = new Float32Array(n / 2);
|
|
for (let i = 0; i < Math.ceil(n / 2); i++) {
|
|
const x = (inverse ? 2 : -2) * Math.PI * (i / n);
|
|
real2[i] = Math.cos(x);
|
|
imag2[i] = Math.sin(x);
|
|
}
|
|
return { real: real2, imag: imag2 };
|
|
}
|
|
function exponent(k, n, inverse) {
|
|
const x = (inverse ? 2 : -2) * Math.PI * (k / n);
|
|
const real2 = Math.cos(x);
|
|
const imag2 = Math.sin(x);
|
|
return { real: real2, imag: imag2 };
|
|
}
|
|
const ARROW = "->";
|
|
const ARROW_REGEX = /->/g;
|
|
const COMMA = ",";
|
|
const ELLIPSIS = "...";
|
|
function decodeEinsumEquation(equation, numTensors) {
|
|
equation = equation.replace(/\s/g, "");
|
|
const numArrows = (equation.length - equation.replace(ARROW_REGEX, "").length) / ARROW.length;
|
|
if (numArrows < 1) {
|
|
throw new Error("Equations without an arrow are not supported.");
|
|
} else if (numArrows > 1) {
|
|
throw new Error(`Equation must contain exactly one arrow ("${ARROW}").`);
|
|
}
|
|
const [inputString, outputString] = equation.split(ARROW);
|
|
assert(inputString.indexOf(ELLIPSIS) === -1, () => `The ellipsis notation ("${ELLIPSIS}") is not supported yet.`);
|
|
const inputTerms = inputString.split(COMMA);
|
|
const numInputs = inputTerms.length;
|
|
if (numTensors !== numInputs) {
|
|
throw new Error(`Expected ${numInputs} input tensors, received ${numTensors}`);
|
|
}
|
|
if (numInputs > 2) {
|
|
throw new Error("Support for more than 2 input tensors is not implemented yet.");
|
|
}
|
|
const allDims = [];
|
|
for (let i = 0; i < outputString.length; ++i) {
|
|
const dimName = outputString[i];
|
|
if (!inputTerms.some((inputTerm) => inputTerm.indexOf(dimName) !== -1)) {
|
|
throw new Error(`Output subscripts contain the label ${dimName} not present in the input subscripts.`);
|
|
}
|
|
if (allDims.indexOf(dimName) === -1) {
|
|
allDims.push(dimName);
|
|
}
|
|
}
|
|
for (let i = 0; i < inputString.length; ++i) {
|
|
const dimName = inputString[i];
|
|
if (allDims.indexOf(dimName) === -1 && dimName !== COMMA) {
|
|
allDims.push(dimName);
|
|
}
|
|
}
|
|
const idDims = new Array(inputTerms.length);
|
|
for (let i = 0; i < numInputs; ++i) {
|
|
if (new Set(inputTerms[i].split("")).size !== inputTerms[i].length) {
|
|
throw new Error(`Found duplicate axes in input component ${inputTerms[i]}. Support for duplicate axes in input is not implemented yet.`);
|
|
}
|
|
idDims[i] = [];
|
|
for (let j2 = 0; j2 < inputTerms[i].length; ++j2) {
|
|
idDims[i].push(allDims.indexOf(inputTerms[i][j2]));
|
|
}
|
|
}
|
|
const numDims = allDims.length;
|
|
const numOutDims = outputString.length;
|
|
const summedDims = [];
|
|
for (let i = numOutDims; i < numDims; ++i) {
|
|
summedDims.push(i);
|
|
}
|
|
return { allDims, summedDims, idDims };
|
|
}
|
|
function getEinsumPermutation(nDims, idDims) {
|
|
let permutationIndices = new Array(nDims);
|
|
permutationIndices.fill(-1);
|
|
for (let i = 0; i < idDims.length; ++i) {
|
|
permutationIndices[idDims[i]] = i;
|
|
}
|
|
const expandDims2 = [];
|
|
for (let i = 0; i < nDims; ++i) {
|
|
if (permutationIndices[i] === -1) {
|
|
expandDims2.push(i);
|
|
}
|
|
}
|
|
permutationIndices = permutationIndices.filter((d) => d !== -1);
|
|
return { permutationIndices, expandDims: expandDims2 };
|
|
}
|
|
function checkEinsumDimSizes(nDims, idDims, tensors) {
|
|
const dimSizes = new Array(nDims);
|
|
for (let i = 0; i < tensors.length; ++i) {
|
|
const shape = tensors[i].shape;
|
|
for (let j2 = 0; j2 < idDims[i].length; ++j2) {
|
|
if (dimSizes[idDims[i][j2]] === void 0) {
|
|
dimSizes[idDims[i][j2]] = shape[j2];
|
|
} else {
|
|
assert(dimSizes[idDims[i][j2]] === shape[j2], () => `Expected dimension ${dimSizes[idDims[i][j2]]} at axis ${j2} of input shaped ${JSON.stringify(shape)}, but got dimension ${shape[j2]}`);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
function getEinsumComputePath(summedDims, idDims) {
|
|
const path = summedDims;
|
|
const steps = [];
|
|
let nSteps = 0;
|
|
if (summedDims.length === 0) {
|
|
path.push(-1);
|
|
}
|
|
nSteps = summedDims.length + 1;
|
|
for (let i = 0; i < nSteps; ++i) {
|
|
steps.push([]);
|
|
}
|
|
const computedTermIndices = [];
|
|
for (let i = 0; i < path.length; ++i) {
|
|
const summedDim = path[i];
|
|
const termIndices = findTermsWithDim(idDims, summedDim);
|
|
for (const termIndex of termIndices) {
|
|
if (computedTermIndices.indexOf(termIndex) === -1) {
|
|
steps[i].push(termIndex);
|
|
computedTermIndices.push(termIndex);
|
|
}
|
|
}
|
|
}
|
|
return { path, steps };
|
|
}
|
|
function isIdentityPermutation(perm) {
|
|
return perm.every((dim, index) => dim === index);
|
|
}
|
|
function findTermsWithDim(idDims, dim) {
|
|
const termIndices = [];
|
|
for (let i = 0; i < idDims.length; ++i) {
|
|
if (idDims[i].length === 0 || idDims[i].indexOf(dim) !== -1 || dim === -1) {
|
|
termIndices.push(i);
|
|
}
|
|
}
|
|
return termIndices;
|
|
}
|
|
function prepareSplitSize(x, numOrSizeSplits, axis = 0) {
|
|
let splitSizes = [];
|
|
if (typeof numOrSizeSplits === "number") {
|
|
assert(x.shape[axis] % numOrSizeSplits === 0, () => "Number of splits must evenly divide the axis.");
|
|
splitSizes = new Array(numOrSizeSplits).fill(x.shape[axis] / numOrSizeSplits);
|
|
} else {
|
|
const numOfNegs = numOrSizeSplits.reduce((count, value) => {
|
|
if (value === -1) {
|
|
count += 1;
|
|
}
|
|
return count;
|
|
}, 0);
|
|
assert(numOfNegs <= 1, () => "There should be only one negative value in split array.");
|
|
const negIndex = numOrSizeSplits.indexOf(-1);
|
|
if (negIndex !== -1) {
|
|
const total = numOrSizeSplits.reduce((a, b) => b > 0 ? a + b : a);
|
|
numOrSizeSplits[negIndex] = x.shape[axis] - total;
|
|
}
|
|
assert(x.shape[axis] === numOrSizeSplits.reduce((a, b) => a + b), () => "The sum of sizes must match the size of the axis dimension.");
|
|
splitSizes = numOrSizeSplits;
|
|
}
|
|
return splitSizes;
|
|
}
|
|
function getSparseFillEmptyRowsIndicesDenseShapeMismatch(indicesLength) {
|
|
return `Received SparseTensor with denseShape[0] = 0 but
|
|
indices.shape[0] = ${indicesLength}`;
|
|
}
|
|
function getSparseFillEmptyRowsNegativeIndexErrorMessage(index, value) {
|
|
return `indices(${index}, 0) is invalid: ${value} < 0`;
|
|
}
|
|
function getSparseFillEmptyRowsOutOfRangeIndexErrorMessage(index, value, limit) {
|
|
return `indices(${index}, 0) is invalid: ${value} >= ${limit}`;
|
|
}
|
|
function getSparseReshapeMultipleNegativeOneOutputDimErrorMessage(dim1, dim2) {
|
|
return `only one output dimension may be -1, not both ${dim1} and ${dim2}`;
|
|
}
|
|
function getSparseReshapeNegativeOutputDimErrorMessage(dim, value) {
|
|
return `size ${dim} must be non-negative, not ${value}`;
|
|
}
|
|
function getSparseReshapeEmptyTensorZeroOutputDimErrorMessage() {
|
|
return "reshape cannot infer the missing input size for an empty tensor unless all specified input sizes are non-zero";
|
|
}
|
|
function getSparseReshapeInputOutputMultipleErrorMessage(inputShape, outputShape) {
|
|
const inputSize = sizeFromShape(inputShape);
|
|
const outputSize = sizeFromShape(outputShape);
|
|
return `Input to reshape is a SparseTensor with ${inputSize}
|
|
dense values, but the requested shape requires a multiple of ${outputSize}. inputShape=${inputShape} outputShape= ${outputShape}`;
|
|
}
|
|
function getSparseReshapeInputOutputMismatchErrorMessage(inputShape, outputShape) {
|
|
const inputSize = sizeFromShape(inputShape);
|
|
const outputSize = sizeFromShape(outputShape);
|
|
return `Input to reshape is a tensor with ${inputSize} dense values, but the requested shape has ${outputSize}. inputShape=${inputShape} outputShape=${outputShape}`;
|
|
}
|
|
function getSparseSegmentReductionNegativeSegmentIdsErrorMessage() {
|
|
return `segment ids must be >= 0`;
|
|
}
|
|
function getSparseSegmentReductionNonIncreasingSegmentIdsErrorMessage() {
|
|
return `segment ids are not increasing`;
|
|
}
|
|
function getSparseSegmentReductionSegmentIdOutOfRangeErrorMessage(segmentId, outputRows) {
|
|
return `Segment id ${segmentId} out of range [0, ${outputRows}), possibly because segmentIds input is not sorted.`;
|
|
}
|
|
function getSparseSegmentReductionIndicesOutOfRangeErrorMessage(index, indexValue, inputRows) {
|
|
return `Bad: indices[${index}] == ${indexValue} out of range [0, ${inputRows})`;
|
|
}
|
|
function segOpComputeOptimalWindowSize(inSize, numSegments) {
|
|
let done = false;
|
|
let res;
|
|
if (inSize <= PARALLELIZE_THRESHOLD) {
|
|
res = inSize;
|
|
done = true;
|
|
} else {
|
|
res = nearestDivisor(inSize, Math.floor(Math.sqrt(inSize)));
|
|
}
|
|
while (!done) {
|
|
if (res > numSegments || res === inSize) {
|
|
done = true;
|
|
} else {
|
|
res = nearestDivisor(inSize, res + 1);
|
|
}
|
|
}
|
|
return res;
|
|
}
|
|
function computeOutShape(aShape, axis, numSegments) {
|
|
const outShape = [];
|
|
const rank = aShape.length;
|
|
for (let dim = 0; dim < rank; dim++) {
|
|
if (dim !== axis) {
|
|
outShape.push(aShape[dim]);
|
|
} else {
|
|
outShape.push(numSegments);
|
|
}
|
|
}
|
|
return outShape;
|
|
}
|
|
function collectGatherOpShapeInfo(x, indices, axis, batchDims) {
|
|
const indicesRank = indices.shape.length;
|
|
const xRank = x.shape.length;
|
|
if (batchDims !== 0) {
|
|
if (batchDims < -indicesRank || batchDims > indicesRank) {
|
|
throw new Error(`Expect batchDims in the range of [-${indicesRank}, ${indicesRank}], but got ${batchDims}`);
|
|
}
|
|
}
|
|
if (batchDims < 0) {
|
|
batchDims += indicesRank;
|
|
}
|
|
if (batchDims > xRank) {
|
|
throw new Error(`batchDims (${batchDims}) must be less than rank(x) (
|
|
${xRank}).`);
|
|
}
|
|
if (axis < batchDims) {
|
|
throw new Error(`batchDims (${batchDims}) must be less than or equal to axis (${axis}).`);
|
|
}
|
|
for (let i = 0; i < batchDims; ++i) {
|
|
if (x.shape[i] !== indices.shape[i]) {
|
|
throw new Error(`x.shape[${i}]: ${x.shape[i]} should be equal to indices.shape[${i}]: ${indices.shape[i]}.`);
|
|
}
|
|
}
|
|
const dimSize = x.shape[axis];
|
|
const outputShape = [];
|
|
let batchSize = 1;
|
|
let outerSize = 1;
|
|
let sliceSize = 1;
|
|
for (let i = 0; i < batchDims; ++i) {
|
|
outputShape.push(x.shape[i]);
|
|
batchSize *= x.shape[i];
|
|
}
|
|
for (let i = batchDims; i < axis; i++) {
|
|
outputShape.push(x.shape[i]);
|
|
outerSize *= x.shape[i];
|
|
}
|
|
for (let i = batchDims; i < indicesRank; i++) {
|
|
outputShape.push(indices.shape[i]);
|
|
}
|
|
for (let i = axis + 1; i < xRank; i++) {
|
|
outputShape.push(x.shape[i]);
|
|
sliceSize *= x.shape[i];
|
|
}
|
|
return { batchSize, sliceSize, outerSize, dimSize, outputShape };
|
|
}
|
|
function fromUint8ToStringArray(vals) {
|
|
try {
|
|
return vals.map((val) => decodeString(val));
|
|
} catch (err) {
|
|
throw new Error(`Failed to decode encoded string bytes into utf-8, error: ${err}`);
|
|
}
|
|
}
|
|
function fromStringArrayToUint8(strings) {
|
|
return strings.map((s) => encodeString(s));
|
|
}
|
|
const backend_util = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
ERF_A1,
|
|
ERF_A2,
|
|
ERF_A3,
|
|
ERF_A4,
|
|
ERF_A5,
|
|
ERF_P,
|
|
PARALLELIZE_THRESHOLD,
|
|
get RowPartitionType() {
|
|
return RowPartitionType$1;
|
|
},
|
|
SELU_SCALE,
|
|
SELU_SCALEALPHA,
|
|
applyActivation: applyActivation$1,
|
|
assertAndGetBroadcastShape,
|
|
assertAxesAreInnerMostDims,
|
|
assertParamsConsistent,
|
|
assignToTypedArray,
|
|
axesAreInnerMostDims,
|
|
calculateShapes,
|
|
checkEinsumDimSizes,
|
|
checkPadOnDimRoundingMode,
|
|
combineLocations,
|
|
combineRaggedTensorToTensorShapes,
|
|
complexWithEvenIndex,
|
|
complexWithOddIndex,
|
|
computeConv2DInfo,
|
|
computeConv3DInfo,
|
|
computeDefaultPad,
|
|
computeDilation2DInfo,
|
|
computeOptimalWindowSize,
|
|
computeOutAndReduceShapes,
|
|
computeOutShape: computeOutShape$1,
|
|
computePool2DInfo,
|
|
computePool3DInfo,
|
|
convertConv2DDataFormat,
|
|
decodeEinsumEquation,
|
|
eitherStridesOrDilationsAreOne,
|
|
expandShapeToKeepDim,
|
|
exponent,
|
|
exponents,
|
|
fromStringArrayToUint8,
|
|
fromUint8ToStringArray,
|
|
getAxesPermutation,
|
|
getBroadcastDims: getBroadcastDims$1,
|
|
getComplexWithIndex,
|
|
getEinsumComputePath,
|
|
getEinsumPermutation,
|
|
getFusedBiasGradient,
|
|
getFusedDyActivation,
|
|
getImageCenter,
|
|
getInnerMostAxes,
|
|
getPermuted,
|
|
getRaggedRank,
|
|
getReductionAxes,
|
|
getReshaped,
|
|
getReshapedPermuted,
|
|
getRowPartitionTypesHelper,
|
|
getSliceBeginCoords,
|
|
getSliceSize,
|
|
getSparseFillEmptyRowsIndicesDenseShapeMismatch,
|
|
getSparseFillEmptyRowsNegativeIndexErrorMessage,
|
|
getSparseFillEmptyRowsOutOfRangeIndexErrorMessage,
|
|
getSparseReshapeEmptyTensorZeroOutputDimErrorMessage,
|
|
getSparseReshapeInputOutputMismatchErrorMessage,
|
|
getSparseReshapeInputOutputMultipleErrorMessage,
|
|
getSparseReshapeMultipleNegativeOneOutputDimErrorMessage,
|
|
getSparseReshapeNegativeOutputDimErrorMessage,
|
|
getSparseSegmentReductionIndicesOutOfRangeErrorMessage,
|
|
getSparseSegmentReductionNegativeSegmentIdsErrorMessage,
|
|
getSparseSegmentReductionNonIncreasingSegmentIdsErrorMessage,
|
|
getSparseSegmentReductionSegmentIdOutOfRangeErrorMessage,
|
|
getUndoAxesPermutation,
|
|
isIdentityPermutation,
|
|
mergeRealAndImagArrays,
|
|
prepareAndValidate,
|
|
prepareSplitSize,
|
|
shouldFuse,
|
|
splitRealAndImagArrays,
|
|
stridesOrDilationsArePositive,
|
|
tupleValuesAreOne,
|
|
upcastType,
|
|
validateDefaultValueShape,
|
|
validateInput: validateInput$1,
|
|
validateUpdateShape,
|
|
warn
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
registerOptimizers();
|
|
const t = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
Abs,
|
|
Acos,
|
|
Acosh,
|
|
AdadeltaOptimizer,
|
|
AdagradOptimizer,
|
|
AdamOptimizer,
|
|
AdamaxOptimizer,
|
|
Add,
|
|
AddN,
|
|
All,
|
|
Any,
|
|
ArgMax,
|
|
ArgMin,
|
|
Asin,
|
|
Asinh,
|
|
Atan,
|
|
Atan2,
|
|
Atanh,
|
|
AvgPool,
|
|
AvgPool3D,
|
|
AvgPool3DGrad,
|
|
AvgPoolGrad,
|
|
BatchMatMul,
|
|
BatchToSpaceND,
|
|
Bincount,
|
|
BitwiseAnd,
|
|
BroadcastArgs,
|
|
Cast,
|
|
Ceil,
|
|
ClipByValue,
|
|
Complex,
|
|
ComplexAbs,
|
|
Concat,
|
|
Conv2D,
|
|
Conv2DBackpropFilter,
|
|
Conv2DBackpropInput,
|
|
Conv3D,
|
|
Conv3DBackpropFilterV2,
|
|
Conv3DBackpropInputV2,
|
|
Cos,
|
|
Cosh,
|
|
CropAndResize,
|
|
Cumprod,
|
|
Cumsum,
|
|
DataStorage,
|
|
DenseBincount,
|
|
DepthToSpace,
|
|
DepthwiseConv2dNative,
|
|
DepthwiseConv2dNativeBackpropFilter,
|
|
DepthwiseConv2dNativeBackpropInput,
|
|
Diag,
|
|
Dilation2D,
|
|
Dilation2DBackpropFilter,
|
|
Dilation2DBackpropInput,
|
|
Draw,
|
|
get ENV() {
|
|
return ENV$3;
|
|
},
|
|
Einsum,
|
|
Elu,
|
|
EluGrad,
|
|
Environment,
|
|
Equal,
|
|
Erf,
|
|
Exp,
|
|
ExpandDims,
|
|
Expm1,
|
|
FFT,
|
|
Fill,
|
|
FlipLeftRight,
|
|
Floor,
|
|
FloorDiv,
|
|
FromPixels,
|
|
FusedBatchNorm,
|
|
FusedConv2D,
|
|
FusedDepthwiseConv2D,
|
|
GatherNd,
|
|
GatherV2,
|
|
Greater,
|
|
GreaterEqual,
|
|
IFFT,
|
|
Identity,
|
|
Imag,
|
|
IsFinite,
|
|
IsInf,
|
|
IsNan,
|
|
KernelBackend,
|
|
LRN,
|
|
LRNGrad,
|
|
LeakyRelu,
|
|
Less,
|
|
LessEqual,
|
|
LinSpace,
|
|
Log,
|
|
Log1p,
|
|
LogicalAnd,
|
|
LogicalNot,
|
|
LogicalOr,
|
|
Max,
|
|
MaxPool,
|
|
MaxPool3D,
|
|
MaxPool3DGrad,
|
|
MaxPoolGrad,
|
|
MaxPoolWithArgmax,
|
|
Maximum,
|
|
Mean,
|
|
Min,
|
|
Minimum,
|
|
MirrorPad,
|
|
Mod,
|
|
MomentumOptimizer,
|
|
Multinomial,
|
|
Multiply,
|
|
Neg,
|
|
NonMaxSuppressionV3,
|
|
NonMaxSuppressionV4,
|
|
NonMaxSuppressionV5,
|
|
NotEqual,
|
|
OP_SCOPE_SUFFIX,
|
|
OneHot,
|
|
OnesLike,
|
|
Optimizer,
|
|
Pack,
|
|
PadV2,
|
|
Pow,
|
|
Prelu,
|
|
Prod,
|
|
RMSPropOptimizer,
|
|
RaggedGather,
|
|
RaggedRange,
|
|
RaggedTensorToTensor,
|
|
Range,
|
|
get Rank() {
|
|
return Rank;
|
|
},
|
|
Real,
|
|
RealDiv,
|
|
Reciprocal,
|
|
get Reduction() {
|
|
return Reduction;
|
|
},
|
|
Relu,
|
|
Relu6,
|
|
Reshape,
|
|
ResizeBilinear,
|
|
ResizeBilinearGrad,
|
|
ResizeNearestNeighbor,
|
|
ResizeNearestNeighborGrad,
|
|
Reverse,
|
|
RotateWithOffset,
|
|
Round,
|
|
Rsqrt,
|
|
SGDOptimizer,
|
|
ScatterNd,
|
|
SearchSorted,
|
|
Select,
|
|
Selu,
|
|
Sigmoid,
|
|
Sign,
|
|
Sin,
|
|
Sinh,
|
|
Slice,
|
|
Softmax,
|
|
Softplus,
|
|
SpaceToBatchND,
|
|
SparseFillEmptyRows,
|
|
SparseReshape,
|
|
SparseSegmentMean,
|
|
SparseSegmentSum,
|
|
SparseToDense,
|
|
SplitV,
|
|
Sqrt,
|
|
Square,
|
|
SquaredDifference,
|
|
StaticRegexReplace,
|
|
Step,
|
|
StridedSlice,
|
|
StringNGrams,
|
|
StringSplit,
|
|
StringToHashBucketFast,
|
|
Sub,
|
|
Sum,
|
|
Tan,
|
|
Tanh,
|
|
Tensor,
|
|
TensorBuffer,
|
|
TensorScatterUpdate,
|
|
Tile,
|
|
TopK,
|
|
Transform,
|
|
Transpose,
|
|
Unique,
|
|
Unpack,
|
|
UnsortedSegmentSum,
|
|
Variable,
|
|
ZerosLike,
|
|
_FusedMatMul,
|
|
abs: abs$2,
|
|
acos: acos$2,
|
|
acosh: acosh$2,
|
|
add: add$1,
|
|
addN: addN$2,
|
|
all: all$2,
|
|
any: any$2,
|
|
argMax: argMax$2,
|
|
argMin: argMin$2,
|
|
asin: asin$2,
|
|
asinh: asinh$2,
|
|
atan: atan$2,
|
|
atan2: atan2$2,
|
|
atanh: atanh$2,
|
|
avgPool: avgPool$2,
|
|
avgPool3d,
|
|
backend,
|
|
backend_util,
|
|
basicLSTMCell,
|
|
batchNorm: batchNorm$2,
|
|
batchNorm2d,
|
|
batchNorm3d,
|
|
batchNorm4d,
|
|
batchToSpaceND: batchToSpaceND$2,
|
|
bincount: bincount$2,
|
|
bitwiseAnd: bitwiseAnd$2,
|
|
booleanMaskAsync,
|
|
broadcastArgs: broadcastArgs$2,
|
|
broadcastTo,
|
|
buffer,
|
|
cast: cast$2,
|
|
ceil: ceil$2,
|
|
clipByValue: clipByValue$2,
|
|
clone,
|
|
complex: complex$2,
|
|
concat: concat$2,
|
|
concat1d,
|
|
concat2d,
|
|
concat3d,
|
|
concat4d,
|
|
conv1d,
|
|
conv2d: conv2d$2,
|
|
conv2dTranspose,
|
|
conv3d,
|
|
conv3dTranspose,
|
|
cos: cos$2,
|
|
cosh: cosh$2,
|
|
cosineWindow,
|
|
cumprod: cumprod$2,
|
|
cumsum: cumsum$2,
|
|
customGrad,
|
|
denseBincount: denseBincount$2,
|
|
depthToSpace: depthToSpace$2,
|
|
depthwiseConv2d: depthwiseConv2d$1,
|
|
diag: diag$2,
|
|
dilation2d,
|
|
dispose,
|
|
div: div$1,
|
|
divNoNan,
|
|
dot,
|
|
dropout,
|
|
einsum: einsum$2,
|
|
elu: elu$2,
|
|
enclosingPowerOfTwo,
|
|
engine,
|
|
ensureShape,
|
|
env,
|
|
equal: equal$2,
|
|
erf: erf$2,
|
|
euclideanNorm,
|
|
exp: exp$2,
|
|
expandDims: expandDims$2,
|
|
expm1: expm1$2,
|
|
eye,
|
|
fft: fft$2,
|
|
fill: fill$2,
|
|
floor: floor$2,
|
|
floorDiv: floorDiv$2,
|
|
fused: fused_ops,
|
|
gather,
|
|
gatherND,
|
|
getBackend,
|
|
getGradient,
|
|
getKernel,
|
|
getKernelsForBackend,
|
|
greater: greater$2,
|
|
greaterEqual: greaterEqual$2,
|
|
ifft: ifft$2,
|
|
imag: imag$2,
|
|
image: image$1,
|
|
inTopKAsync,
|
|
io,
|
|
irfft,
|
|
isFinite: isFinite$3,
|
|
isInf: isInf$2,
|
|
isNaN: isNaN$3,
|
|
keep,
|
|
leakyRelu: leakyRelu$2,
|
|
less: less$2,
|
|
lessEqual: lessEqual$2,
|
|
linalg,
|
|
linspace,
|
|
localResponseNormalization,
|
|
log: log$2,
|
|
log1p: log1p$2,
|
|
logSigmoid,
|
|
logSoftmax,
|
|
logSumExp,
|
|
logicalAnd: logicalAnd$2,
|
|
logicalNot: logicalNot$2,
|
|
logicalOr: logicalOr$2,
|
|
logicalXor,
|
|
losses,
|
|
lowerBound: lowerBound$1,
|
|
matMul: matMul$1,
|
|
max: max$2,
|
|
maxPool: maxPool$2,
|
|
maxPool3d: maxPool3d$1,
|
|
maxPoolWithArgmax,
|
|
maximum: maximum$2,
|
|
mean: mean$1,
|
|
meshgrid,
|
|
min: min$2,
|
|
minimum: minimum$2,
|
|
mirrorPad: mirrorPad$1,
|
|
mod: mod$2,
|
|
moments,
|
|
movingAverage,
|
|
mul,
|
|
multiRNNCell,
|
|
multinomial: multinomial$2,
|
|
neg: neg$2,
|
|
nextFrame,
|
|
norm,
|
|
notEqual: notEqual$2,
|
|
oneHot: oneHot$2,
|
|
ones,
|
|
onesLike: onesLike$2,
|
|
op,
|
|
outerProduct,
|
|
pad,
|
|
pad1d,
|
|
pad2d,
|
|
pad3d,
|
|
pad4d,
|
|
pool: pool$1,
|
|
pow: pow$2,
|
|
prelu: prelu$2,
|
|
print,
|
|
prod: prod$2,
|
|
raggedGather: raggedGather$2,
|
|
raggedRange: raggedRange$2,
|
|
raggedTensorToTensor: raggedTensorToTensor$2,
|
|
rand,
|
|
randomGamma,
|
|
randomNormal,
|
|
randomStandardNormal,
|
|
randomUniform,
|
|
randomUniformInt,
|
|
range: range$2,
|
|
real: real$2,
|
|
reciprocal: reciprocal$2,
|
|
registerBackend,
|
|
registerKernel,
|
|
relu: relu$2,
|
|
relu6: relu6$2,
|
|
reshape: reshape$2,
|
|
reverse: reverse$2,
|
|
reverse1d,
|
|
reverse2d,
|
|
reverse3d,
|
|
reverse4d,
|
|
rfft,
|
|
round: round$2,
|
|
rsqrt: rsqrt$2,
|
|
scalar,
|
|
scatterND,
|
|
searchSorted: searchSorted$2,
|
|
selu: selu$2,
|
|
separableConv2d,
|
|
setBackend,
|
|
setdiff1dAsync,
|
|
sigmoid: sigmoid$2,
|
|
sign: sign$2,
|
|
signal,
|
|
sin: sin$2,
|
|
sinh: sinh$2,
|
|
slice: slice$2,
|
|
slice1d,
|
|
slice2d,
|
|
slice3d,
|
|
slice4d,
|
|
softmax: softmax$2,
|
|
softplus: softplus$2,
|
|
spaceToBatchND: spaceToBatchND$2,
|
|
sparse: sparse$1,
|
|
sparseToDense: sparseToDense$2,
|
|
spectral: spectral$1,
|
|
split: split$2,
|
|
sqrt: sqrt$2,
|
|
square: square$1,
|
|
squaredDifference: squaredDifference$2,
|
|
squeeze,
|
|
stack,
|
|
step: step$2,
|
|
stridedSlice: stridedSlice$2,
|
|
string: string$1,
|
|
sub: sub$2,
|
|
sum: sum$2,
|
|
sumOutType,
|
|
tan: tan$2,
|
|
tanh: tanh$2,
|
|
tensor,
|
|
tensor1d,
|
|
tensor2d,
|
|
tensor3d,
|
|
tensor4d,
|
|
tensor5d,
|
|
tensor6d,
|
|
tensorScatterUpdate: tensorScatterUpdate$2,
|
|
tidy,
|
|
tile: tile$2,
|
|
topk,
|
|
transpose: transpose$2,
|
|
truncatedNormal,
|
|
unique: unique$2,
|
|
unsortedSegmentSum: unsortedSegmentSum$2,
|
|
unstack,
|
|
upcastType,
|
|
upperBound: upperBound$1,
|
|
variable,
|
|
variableGrads,
|
|
where,
|
|
whereAsync,
|
|
zeros: zeros$1,
|
|
zerosLike: zerosLike$2
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const ENV$1 = env();
|
|
ENV$1.registerFlag("KEEP_INTERMEDIATE_TENSORS", () => false, (debugValue) => {
|
|
if (debugValue) {
|
|
console.warn("Keep intermediate tensors is ON. This will print the values of all intermediate tensors during model inference. Not all models support this mode. For details, check e2e/benchmarks/ model_config.js. This significantly impacts performance.");
|
|
}
|
|
});
|
|
var DataType;
|
|
(function(DataType2) {
|
|
DataType2[DataType2["DT_INVALID"] = 0] = "DT_INVALID";
|
|
DataType2[DataType2["DT_FLOAT"] = 1] = "DT_FLOAT";
|
|
DataType2[DataType2["DT_DOUBLE"] = 2] = "DT_DOUBLE";
|
|
DataType2[DataType2["DT_INT32"] = 3] = "DT_INT32";
|
|
DataType2[DataType2["DT_UINT8"] = 4] = "DT_UINT8";
|
|
DataType2[DataType2["DT_INT16"] = 5] = "DT_INT16";
|
|
DataType2[DataType2["DT_INT8"] = 6] = "DT_INT8";
|
|
DataType2[DataType2["DT_STRING"] = 7] = "DT_STRING";
|
|
DataType2[DataType2["DT_COMPLEX64"] = 8] = "DT_COMPLEX64";
|
|
DataType2[DataType2["DT_INT64"] = 9] = "DT_INT64";
|
|
DataType2[DataType2["DT_BOOL"] = 10] = "DT_BOOL";
|
|
DataType2[DataType2["DT_QINT8"] = 11] = "DT_QINT8";
|
|
DataType2[DataType2["DT_QUINT8"] = 12] = "DT_QUINT8";
|
|
DataType2[DataType2["DT_QINT32"] = 13] = "DT_QINT32";
|
|
DataType2[DataType2["DT_BFLOAT16"] = 14] = "DT_BFLOAT16";
|
|
DataType2[DataType2["DT_QINT16"] = 15] = "DT_QINT16";
|
|
DataType2[DataType2["DT_QUINT16"] = 16] = "DT_QUINT16";
|
|
DataType2[DataType2["DT_UINT16"] = 17] = "DT_UINT16";
|
|
DataType2[DataType2["DT_COMPLEX128"] = 18] = "DT_COMPLEX128";
|
|
DataType2[DataType2["DT_HALF"] = 19] = "DT_HALF";
|
|
DataType2[DataType2["DT_RESOURCE"] = 20] = "DT_RESOURCE";
|
|
DataType2[DataType2["DT_VARIANT"] = 21] = "DT_VARIANT";
|
|
DataType2[DataType2["DT_UINT32"] = 22] = "DT_UINT32";
|
|
DataType2[DataType2["DT_UINT64"] = 23] = "DT_UINT64";
|
|
DataType2[DataType2["DT_FLOAT_REF"] = 101] = "DT_FLOAT_REF";
|
|
DataType2[DataType2["DT_DOUBLE_REF"] = 102] = "DT_DOUBLE_REF";
|
|
DataType2[DataType2["DT_INT32_REF"] = 103] = "DT_INT32_REF";
|
|
DataType2[DataType2["DT_UINT8_REF"] = 104] = "DT_UINT8_REF";
|
|
DataType2[DataType2["DT_INT16_REF"] = 105] = "DT_INT16_REF";
|
|
DataType2[DataType2["DT_INT8_REF"] = 106] = "DT_INT8_REF";
|
|
DataType2[DataType2["DT_STRING_REF"] = 107] = "DT_STRING_REF";
|
|
DataType2[DataType2["DT_COMPLEX64_REF"] = 108] = "DT_COMPLEX64_REF";
|
|
DataType2[DataType2["DT_INT64_REF"] = 109] = "DT_INT64_REF";
|
|
DataType2[DataType2["DT_BOOL_REF"] = 110] = "DT_BOOL_REF";
|
|
DataType2[DataType2["DT_QINT8_REF"] = 111] = "DT_QINT8_REF";
|
|
DataType2[DataType2["DT_QUINT8_REF"] = 112] = "DT_QUINT8_REF";
|
|
DataType2[DataType2["DT_QINT32_REF"] = 113] = "DT_QINT32_REF";
|
|
DataType2[DataType2["DT_BFLOAT16_REF"] = 114] = "DT_BFLOAT16_REF";
|
|
DataType2[DataType2["DT_QINT16_REF"] = 115] = "DT_QINT16_REF";
|
|
DataType2[DataType2["DT_QUINT16_REF"] = 116] = "DT_QUINT16_REF";
|
|
DataType2[DataType2["DT_UINT16_REF"] = 117] = "DT_UINT16_REF";
|
|
DataType2[DataType2["DT_COMPLEX128_REF"] = 118] = "DT_COMPLEX128_REF";
|
|
DataType2[DataType2["DT_HALF_REF"] = 119] = "DT_HALF_REF";
|
|
DataType2[DataType2["DT_RESOURCE_REF"] = 120] = "DT_RESOURCE_REF";
|
|
DataType2[DataType2["DT_VARIANT_REF"] = 121] = "DT_VARIANT_REF";
|
|
DataType2[DataType2["DT_UINT32_REF"] = 122] = "DT_UINT32_REF";
|
|
DataType2[DataType2["DT_UINT64_REF"] = 123] = "DT_UINT64_REF";
|
|
})(DataType || (DataType = {}));
|
|
var SaverDef;
|
|
(function(SaverDef2) {
|
|
(function(CheckpointFormatVersion) {
|
|
CheckpointFormatVersion[CheckpointFormatVersion["LEGACY"] = 0] = "LEGACY";
|
|
CheckpointFormatVersion[CheckpointFormatVersion["V1"] = 1] = "V1";
|
|
CheckpointFormatVersion[CheckpointFormatVersion["V2"] = 2] = "V2";
|
|
})(SaverDef2.CheckpointFormatVersion || (SaverDef2.CheckpointFormatVersion = {}));
|
|
})(SaverDef || (SaverDef = {}));
|
|
const CUSTOM_OPS = {};
|
|
function getRegisteredOp(name) {
|
|
return CUSTOM_OPS[name];
|
|
}
|
|
function getParamValue(paramName, node, tensorMap, context, resourceManager) {
|
|
const inputParam = node.inputParams[paramName];
|
|
if (inputParam && inputParam.inputIndexStart !== void 0) {
|
|
const start = inputParam.inputIndexStart;
|
|
const end = inputParam.inputIndexEnd === 0 ? void 0 : inputParam.inputIndexEnd === void 0 ? start + 1 : inputParam.inputIndexEnd;
|
|
const shiftedStart = start < 0 ? node.inputNames.length + start : start;
|
|
if (inputParam.type === "tensor") {
|
|
return getTensor(node.inputNames[shiftedStart], tensorMap, context, resourceManager);
|
|
}
|
|
if (inputParam.type === "tensors") {
|
|
const inputs = node.inputs.slice(start, end);
|
|
const inputNames = node.inputNames.slice(start, end).filter((_name, index) => {
|
|
var _a;
|
|
return ((_a = inputs[index]) === null || _a === void 0 ? void 0 : _a.op) !== "NoOp";
|
|
});
|
|
return inputNames.map((name) => getTensor(name, tensorMap, context, resourceManager));
|
|
}
|
|
const tensor2 = getTensor(node.inputNames[shiftedStart], tensorMap, context, resourceManager);
|
|
const data = tensor2.dataSync();
|
|
return inputParam.type === "number" ? data[0] : toNestedArray(tensor2.shape, data);
|
|
}
|
|
const attrParam = node.attrParams[paramName];
|
|
return attrParam && attrParam.value;
|
|
}
|
|
function getTensor(name, tensorsMap, context, resourceManager) {
|
|
const [nodeName, index] = parseNodeName(name, context);
|
|
if (resourceManager != null) {
|
|
const tensor2 = resourceManager.getHashTableHandleByName(nodeName);
|
|
if (tensor2 != null) {
|
|
return tensor2;
|
|
}
|
|
}
|
|
const contextId = context.currentContextIds.find((contextId2) => {
|
|
return !!tensorsMap[getNodeNameWithContextId(nodeName, contextId2)];
|
|
});
|
|
return contextId !== void 0 ? tensorsMap[getNodeNameWithContextId(nodeName, contextId)][index] : void 0;
|
|
}
|
|
function getTensorsForCurrentContext(name, tensorsMap, context) {
|
|
return tensorsMap[getNodeNameWithContextId(name, context.currentContextId)];
|
|
}
|
|
function getNodeNameAndIndex(inputName, context) {
|
|
const [nodeName, index, outputName] = parseNodeName(inputName, context);
|
|
return [
|
|
getNodeNameWithContextId(nodeName, context && context.currentContextId),
|
|
index,
|
|
outputName
|
|
];
|
|
}
|
|
function getNodeNameWithContextId(name, contextId) {
|
|
return !!contextId ? `${name}-${contextId}` : name;
|
|
}
|
|
function parseNodeName(name, context) {
|
|
if (name === "") {
|
|
return ["", 0, void 0];
|
|
}
|
|
const isCacheEnabled = context != null && context.parseNodeNameCache != null;
|
|
if (isCacheEnabled) {
|
|
const cachedResult = context.parseNodeNameCache.get(name);
|
|
if (cachedResult != null) {
|
|
return cachedResult;
|
|
}
|
|
}
|
|
const parts = name.split(":");
|
|
let result;
|
|
if (parts.length === 1) {
|
|
result = [name, 0, void 0];
|
|
} else {
|
|
const nodeName = parts[0];
|
|
const outputName = parts.length === 3 ? parts[1] : void 0;
|
|
const index = Number(parts[parts.length - 1]);
|
|
result = [nodeName, index, outputName];
|
|
}
|
|
if (isCacheEnabled) {
|
|
context.parseNodeNameCache.set(name, result);
|
|
}
|
|
return result;
|
|
}
|
|
function getPadding(node, tensorMap, context) {
|
|
let pad2 = getParamValue("pad", node, tensorMap, context);
|
|
if (pad2 === "explicit") {
|
|
pad2 = getParamValue("explicitPaddings", node, tensorMap, context);
|
|
const explicitPadding = [[0, 0], [0, 0], [0, 0], [0, 0]];
|
|
for (let i = 0; i < 4; i++) {
|
|
explicitPadding[i][0] = pad2[i * 2];
|
|
explicitPadding[i][1] = pad2[i * 2 + 1];
|
|
}
|
|
return explicitPadding;
|
|
}
|
|
return pad2;
|
|
}
|
|
function cloneTensor(tensor2) {
|
|
return tensor2.kept ? tensor2 : clone(tensor2);
|
|
}
|
|
const json$i = [
|
|
{
|
|
"tfOpName": "Add",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "AddV2",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "AddN",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "tensors",
|
|
"type": "tensors"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "BiasAdd",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sub",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "RealDiv",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Div",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "DivNoNan",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FloorDiv",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Mul",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Maximum",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Minimum",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Pow",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SquaredDifference",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Mod",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FloorMod",
|
|
"category": "arithmetic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const arithmetic = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$i
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$h = [
|
|
{
|
|
"tfOpName": "Abs",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Acos",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Asin",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Atan",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Atan2",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "y",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Ceil",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ClipByValue",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "clipValueMin",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "clipValueMax",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Complex",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "real",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "imag",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ComplexAbs",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Cos",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Cosh",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Elu",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Exp",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Floor",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Log",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Imag",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "Tout",
|
|
"name": "outputType",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Neg",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Real",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "Tout",
|
|
"name": "outputType",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Prelu",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "alpha",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Relu",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Relu6",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Selu",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sigmoid",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sin",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sinh",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sqrt",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Rsqrt",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Square",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Tan",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Tanh",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sign",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Round",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Expm1",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Log1p",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Reciprocal",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Softplus",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Asinh",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Acosh",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Atanh",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Erf",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LeakyRelu",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "alpha",
|
|
"name": "alpha",
|
|
"type": "number",
|
|
"defaultValue": 0.2
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "IsNan",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "IsFinite",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "IsInf",
|
|
"category": "basic_math",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const basicMath = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$h
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$g = [
|
|
{
|
|
"tfOpName": "EmptyTensorList",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "maxNumElements",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LoopCond",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "pred",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Switch",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "data",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "pred",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Merge",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "tensors",
|
|
"type": "tensors"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Enter",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "frame_name",
|
|
"name": "frameName",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "is_constant",
|
|
"name": "isConstant",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Exit",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "NextIteration",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "size",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "element_shape",
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"tfName": "dynamic_size",
|
|
"name": "dynamicSize",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "clear_after_read",
|
|
"name": "clearAfterRead",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "identical_element_shapes",
|
|
"name": "identicalElementShapes",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "tensor_array_name",
|
|
"name": "name",
|
|
"type": "string"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayWriteV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "index",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayReadV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "index",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayGatherV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "element_shape",
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayScatterV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayConcatV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "element_shape_except0",
|
|
"name": "elementShapeExcept0",
|
|
"type": "shape",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArraySplitV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "lengths",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArraySizeV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "flowIn",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorArrayCloseV3",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorArrayId",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "StatelessIf",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "cond",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "then_branch",
|
|
"name": "thenBranch",
|
|
"type": "func"
|
|
},
|
|
{
|
|
"tfName": "else_branch",
|
|
"name": "elseBranch",
|
|
"type": "func"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "If",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "cond",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "then_branch",
|
|
"name": "thenBranch",
|
|
"type": "func"
|
|
},
|
|
{
|
|
"tfName": "else_branch",
|
|
"name": "elseBranch",
|
|
"type": "func"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "StatelessWhile",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "cond",
|
|
"name": "cond",
|
|
"type": "func"
|
|
},
|
|
{
|
|
"tfName": "body",
|
|
"name": "body",
|
|
"type": "func"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "While",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "cond",
|
|
"name": "cond",
|
|
"type": "func"
|
|
},
|
|
{
|
|
"tfName": "body",
|
|
"name": "body",
|
|
"type": "func"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListScatter",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListScatterV2",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "numElements",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListGather",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListGetItem",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "index",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListSetItem",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "index",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListReserve",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "numElements",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListFromTensor",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListStack",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "num_elements",
|
|
"name": "numElements",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListSplit",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "lengths",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListConcat",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_shape",
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListConcatV2",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_shape",
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListPopBack",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "elementShape",
|
|
"type": "shape"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListPushBack",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "element_dtype",
|
|
"name": "elementDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListLength",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorListResize",
|
|
"category": "control",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensorListId",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "size",
|
|
"type": "number"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const control = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$g
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$f = [
|
|
{
|
|
"tfOpName": "AvgPool",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "ksize",
|
|
"name": "kernelSize",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "MaxPool",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "ksize",
|
|
"name": "kernelSize",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": [],
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "MaxPoolWithArgmax",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "ksize",
|
|
"name": "kernelSize",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "include_batch_in_index",
|
|
"name": "includeBatchInIndex",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "AvgPool3D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "ksize",
|
|
"name": "kernelSize",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "MaxPool3D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "ksize",
|
|
"name": "kernelSize",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Conv1D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "stride",
|
|
"name": "stride",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NWC"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "dilation",
|
|
"name": "dilation",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Conv2D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "useCudnnOnGpu",
|
|
"name": "useCudnnOnGpu",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NHWC"
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "_FusedConv2D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "num_args",
|
|
"name": "numArgs",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "use_cudnn_on_gpu",
|
|
"name": "useCudnnOnGpu",
|
|
"type": "bool",
|
|
"defaultValue": true
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NHWC"
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]",
|
|
"defaultValue": [
|
|
1,
|
|
1,
|
|
1,
|
|
1
|
|
]
|
|
},
|
|
{
|
|
"tfName": "fused_ops",
|
|
"name": "fusedOps",
|
|
"type": "string[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "epsilon",
|
|
"name": "epsilon",
|
|
"type": "number",
|
|
"defaultValue": 1e-4
|
|
},
|
|
{
|
|
"tfName": "leakyrelu_alpha",
|
|
"name": "leakyreluAlpha",
|
|
"type": "number",
|
|
"defaultValue": 0.2
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Conv2DBackpropInput",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 2,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 0,
|
|
"name": "outputShape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "DepthwiseConv2d",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "input",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NHWC"
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "DepthwiseConv2dNative",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "input",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NHWC"
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FusedDepthwiseConv2dNative",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "num_args",
|
|
"name": "numArgs",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NHWC"
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]",
|
|
"defaultValue": [
|
|
1,
|
|
1,
|
|
1,
|
|
1
|
|
]
|
|
},
|
|
{
|
|
"tfName": "fused_ops",
|
|
"name": "fusedOps",
|
|
"type": "string[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "explicit_paddings",
|
|
"name": "explicitPaddings",
|
|
"type": "number[]",
|
|
"defaultValue": []
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Conv3D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"defaultValue": "NHWC"
|
|
},
|
|
{
|
|
"tfName": "dilations",
|
|
"name": "dilations",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Dilation2D",
|
|
"category": "convolution",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "filter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "strides",
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "rates",
|
|
"name": "dilations",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "padding",
|
|
"name": "pad",
|
|
"type": "string"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const convolution = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$f
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$e = [
|
|
{
|
|
"tfOpName": "Fill",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "value",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LinSpace",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "start",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "stop",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "num",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "OneHot",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "depth",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "onValue",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "offValue",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "axis",
|
|
"name": "axis",
|
|
"type": "number",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Ones",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "OnesLike",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "RandomStandardNormal",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "seed",
|
|
"name": "seed",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "seed2",
|
|
"name": "seed2",
|
|
"type": "number",
|
|
"defaultValue": 0,
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "T",
|
|
"type": "number",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "RandomUniform",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "minval",
|
|
"name": "minval",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "maxval",
|
|
"name": "maxval",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
},
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "seed",
|
|
"name": "seed",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "seed2",
|
|
"name": "seed2",
|
|
"type": "number",
|
|
"defaultValue": 0,
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "T",
|
|
"type": "number",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "RandomUniformInt",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "minval",
|
|
"name": "minval",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "maxval",
|
|
"name": "maxval",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "seed",
|
|
"name": "seed",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "seed2",
|
|
"name": "seed2",
|
|
"type": "number",
|
|
"defaultValue": 0,
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Range",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "start",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "stop",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "step",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "Tidx",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TruncatedNormal",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "means",
|
|
"name": "mean",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "stddev",
|
|
"name": "stdDev",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
},
|
|
{
|
|
"tfName": "seed",
|
|
"name": "seed",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "seed2",
|
|
"name": "seed2",
|
|
"type": "number",
|
|
"defaultValue": 0,
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "T",
|
|
"type": "number",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Zeros",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ZerosLike",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Multinomial",
|
|
"category": "creation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "logits",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "numSamples",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "seed",
|
|
"name": "seed",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "seed2",
|
|
"name": "seed2",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "output_dtype",
|
|
"name": "output_dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const creation = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$e
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$d = [
|
|
{
|
|
"tfOpName": "NonMaxSuppressionV2",
|
|
"category": "dynamic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "boxes",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scores",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "maxOutputSize",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "iouThreshold",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "NonMaxSuppressionV3",
|
|
"category": "dynamic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "boxes",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scores",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "maxOutputSize",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "iouThreshold",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 4,
|
|
"name": "scoreThreshold",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "NonMaxSuppressionV4",
|
|
"category": "dynamic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "boxes",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scores",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "maxOutputSize",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "iouThreshold",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 4,
|
|
"name": "scoreThreshold",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "T_threshold",
|
|
"name": "threshold",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "pad_to_max_output_size",
|
|
"name": "padToMaxOutputSize",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "NonMaxSuppressionV5",
|
|
"category": "dynamic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "boxes",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scores",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "maxOutputSize",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "iouThreshold",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 4,
|
|
"name": "scoreThreshold",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 5,
|
|
"name": "softNmsSigma",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Where",
|
|
"category": "dynamic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "condition",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ListDiff",
|
|
"category": "dynamic",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "y",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const dynamic = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$d
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$c = [
|
|
{
|
|
"tfOpName": "LowerBound",
|
|
"category": "evaluation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "sortedSequence",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TopKV2",
|
|
"category": "evaluation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "k",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "sorted",
|
|
"name": "sorted",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "UpperBound",
|
|
"category": "evaluation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "sortedSequence",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Unique",
|
|
"category": "evaluation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "UniqueV2",
|
|
"category": "evaluation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const evaluation = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$c
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$b = [
|
|
{
|
|
"tfOpName": "PlaceholderWithDefault",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "default",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "shape",
|
|
"name": "shape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Placeholder",
|
|
"category": "graph",
|
|
"attrs": [
|
|
{
|
|
"tfName": "shape",
|
|
"name": "shape",
|
|
"type": "shape"
|
|
},
|
|
{
|
|
"tfName": "dtype",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Const",
|
|
"category": "graph"
|
|
},
|
|
{
|
|
"tfOpName": "Identity",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "IdentityN",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "x",
|
|
"type": "tensors"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Snapshot",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Rank",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Size",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Shape",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ShapeN",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "x",
|
|
"type": "tensors"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Print",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "data",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "message",
|
|
"name": "message",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "first_n",
|
|
"name": "firstN",
|
|
"type": "number",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "summarize",
|
|
"name": "summarize",
|
|
"type": "number",
|
|
"defaultValue": 3
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "NoOp",
|
|
"category": "graph",
|
|
"inputs": []
|
|
},
|
|
{
|
|
"tfOpName": "StopGradient",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FakeQuantWithMinMaxVars",
|
|
"category": "graph",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "min",
|
|
"name": "min",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "max",
|
|
"name": "max",
|
|
"type": "number"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const graph = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$b
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$a = [
|
|
{
|
|
"tfOpName": "HashTable",
|
|
"category": "hash_table",
|
|
"inputs": [],
|
|
"attrs": [
|
|
{
|
|
"tfName": "shared_name",
|
|
"name": "sharedName",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "use_node_name_sharing",
|
|
"name": "useNodeNameSharing",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "key_dtype",
|
|
"name": "keyDType",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "value_dtype",
|
|
"name": "valueDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "HashTableV2",
|
|
"category": "hash_table",
|
|
"inputs": [],
|
|
"attrs": [
|
|
{
|
|
"tfName": "shared_name",
|
|
"name": "sharedName",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "use_node_name_sharing",
|
|
"name": "useNodeNameSharing",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "key_dtype",
|
|
"name": "keyDType",
|
|
"type": "dtype"
|
|
},
|
|
{
|
|
"tfName": "value_dtype",
|
|
"name": "valueDType",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LookupTableImport",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "keys",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "Tin",
|
|
"name": "tIn",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "Tout",
|
|
"name": "tOut",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LookupTableImportV2",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "keys",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "Tin",
|
|
"name": "tIn",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "Tout",
|
|
"name": "tOut",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LookupTableFind",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "keys",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "defaultValue",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "Tin",
|
|
"name": "tIn",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "Tout",
|
|
"name": "tOut",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LookupTableFindV2",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "keys",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "defaultValue",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "Tin",
|
|
"name": "tIn",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "Tout",
|
|
"name": "tOut",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LookupTableSize",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LookupTableSizeV2",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "InitializeTable",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "keys",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "InitializeTableV2",
|
|
"category": "hash_table",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tableHandle",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "keys",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const hashTable = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$a
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$9 = [
|
|
{
|
|
"tfOpName": "ResizeBilinear",
|
|
"category": "image",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "images",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "size",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "align_corners",
|
|
"name": "alignCorners",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "half_pixel_centers",
|
|
"name": "halfPixelCenters",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ResizeNearestNeighbor",
|
|
"category": "image",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "images",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "size",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "align_corners",
|
|
"name": "alignCorners",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "half_pixel_centers",
|
|
"name": "halfPixelCenters",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "CropAndResize",
|
|
"category": "image",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "image",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "boxes",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "boxInd",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "cropSize",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "method",
|
|
"name": "method",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "extrapolation_value",
|
|
"name": "extrapolationValue",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ImageProjectiveTransformV3",
|
|
"category": "image",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "images",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "transforms",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "outputShape",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "fillValue",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "interpolation",
|
|
"name": "interpolation",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "fill_mode",
|
|
"name": "fillMode",
|
|
"type": "string"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const image = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$9
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$8 = [
|
|
{
|
|
"tfOpName": "Equal",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "NotEqual",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Greater",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "GreaterEqual",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Less",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LessEqual",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LogicalAnd",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LogicalNot",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LogicalOr",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Select",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "condition",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SelectV2",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "condition",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "BitwiseAnd",
|
|
"category": "logical",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "y",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const logical = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$8
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$7 = [
|
|
{
|
|
"tfOpName": "_FusedMatMul",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"end": 0,
|
|
"name": "args",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "num_args",
|
|
"name": "numArgs",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "fused_ops",
|
|
"name": "fusedOps",
|
|
"type": "string[]",
|
|
"defaultValue": []
|
|
},
|
|
{
|
|
"tfName": "epsilon",
|
|
"name": "epsilon",
|
|
"type": "number",
|
|
"defaultValue": 1e-4
|
|
},
|
|
{
|
|
"tfName": "transpose_a",
|
|
"name": "transposeA",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "transpose_b",
|
|
"name": "transposeB",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "leakyrelu_alpha",
|
|
"name": "leakyreluAlpha",
|
|
"type": "number",
|
|
"defaultValue": 0.2
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "MatMul",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "transpose_a",
|
|
"name": "transposeA",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "transpose_b",
|
|
"name": "transposeB",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "BatchMatMul",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "adj_x",
|
|
"name": "transposeA",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "adj_y",
|
|
"name": "transposeB",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "BatchMatMulV2",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "b",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "adj_x",
|
|
"name": "transposeA",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "adj_y",
|
|
"name": "transposeB",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Transpose",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "perm",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Einsum",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "tensors",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "equation",
|
|
"name": "equation",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "N",
|
|
"name": "n",
|
|
"type": "number",
|
|
"defaultValue": 2
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "MatrixBandPart",
|
|
"category": "matrices",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "a",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "numLower",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "numUpper",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const matrices = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$7
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$6 = [
|
|
{
|
|
"tfOpName": "EuclideanNorm",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool",
|
|
"defaultValue": false
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FusedBatchNorm",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scale",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "offset",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "mean",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 4,
|
|
"name": "variance",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "epsilon",
|
|
"name": "epsilon",
|
|
"type": "number",
|
|
"defaultValue": 1e-3
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FusedBatchNormV2",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scale",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "offset",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "mean",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 4,
|
|
"name": "variance",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "epsilon",
|
|
"name": "epsilon",
|
|
"type": "number",
|
|
"defaultValue": 1e-3
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "FusedBatchNormV3",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "scale",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "offset",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "mean",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 4,
|
|
"name": "variance",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "epsilon",
|
|
"name": "epsilon",
|
|
"type": "number",
|
|
"defaultValue": 1e-3
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LRN",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "depth_radius",
|
|
"name": "radius",
|
|
"type": "number",
|
|
"defaultValue": 5
|
|
},
|
|
{
|
|
"tfName": "bias",
|
|
"name": "bias",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
},
|
|
{
|
|
"tfName": "alpha",
|
|
"name": "alpha",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
},
|
|
{
|
|
"tfName": "beta",
|
|
"name": "beta",
|
|
"type": "number",
|
|
"defaultValue": 0.5
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Softmax",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "LogSoftmax",
|
|
"category": "normalization",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const normalization = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$6
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$5 = [
|
|
{
|
|
"tfOpName": "Bincount",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "size",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "weights",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "DenseBincount",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "size",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "weights",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "binary_output",
|
|
"name": "binaryOutput",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Max",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Mean",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Min",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Sum",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "All",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Any",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ArgMax",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ArgMin",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Prod",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "keep_dims",
|
|
"name": "keepDims",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Cumprod",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "exclusive",
|
|
"name": "exclusive",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "reverse",
|
|
"name": "reverse",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Cumsum",
|
|
"category": "reduction",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "exclusive",
|
|
"name": "exclusive",
|
|
"type": "bool"
|
|
},
|
|
{
|
|
"tfName": "reverse",
|
|
"name": "reverse",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const reduction = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$5
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$4 = [
|
|
{
|
|
"tfOpName": "ConcatV2",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": -1,
|
|
"name": "tensors",
|
|
"type": "tensors"
|
|
},
|
|
{
|
|
"start": -1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "N",
|
|
"name": "n",
|
|
"type": "number",
|
|
"defaultValue": 2
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Concat",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 1,
|
|
"end": 0,
|
|
"name": "tensors",
|
|
"type": "tensors"
|
|
},
|
|
{
|
|
"start": 0,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "N",
|
|
"name": "n",
|
|
"type": "number",
|
|
"defaultValue": 2
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "GatherV2",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "axis",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "batch_dims",
|
|
"name": "batchDims",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Gather",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "validate_indices",
|
|
"name": "validateIndices",
|
|
"type": "bool",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Reverse",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "dims",
|
|
"type": "bool[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ReverseV2",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Slice",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "begin",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "size",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "StridedSlice",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "begin",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "end",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "strides",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "begin_mask",
|
|
"name": "beginMask",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "end_mask",
|
|
"name": "endMask",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "new_axis_mask",
|
|
"name": "newAxisMask",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "ellipsis_mask",
|
|
"name": "ellipsisMask",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "shrink_axis_mask",
|
|
"name": "shrinkAxisMask",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Pack",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"end": 0,
|
|
"name": "tensors",
|
|
"type": "tensors"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "axis",
|
|
"name": "axis",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Unpack",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "axis",
|
|
"name": "axis",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"tfName": "num",
|
|
"name": "num",
|
|
"type": "number",
|
|
"defaultValue": 0,
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Tile",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "reps",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Split",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "axis",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "num_split",
|
|
"name": "numOrSizeSplits",
|
|
"type": "number",
|
|
"defaultValue": 1
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SplitV",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "numOrSizeSplits",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "axis",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ScatterNd",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "GatherNd",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SparseToDense",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "sparseIndices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "outputShape",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "sparseValues",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "defaultValue",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "validate_indices",
|
|
"name": "validateIndices",
|
|
"type": "bool",
|
|
"defaultValue": false,
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "TensorScatterUpdate",
|
|
"category": "slice_join",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "tensor",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const sliceJoin = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$4
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$3 = [
|
|
{
|
|
"tfOpName": "SparseFillEmptyRows",
|
|
"category": "sparse",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "values",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "denseShape",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 3,
|
|
"name": "defaultValue",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SparseReshape",
|
|
"category": "sparse",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "inputIndices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "inputShape",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "newShape",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "T",
|
|
"name": "dtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SparseSegmentMean",
|
|
"category": "sparse",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "data",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "segmentIds",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SparseSegmentSum",
|
|
"category": "sparse",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "data",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "indices",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "segmentIds",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const sparse = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$3
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$2 = [
|
|
{
|
|
"tfOpName": "FFT",
|
|
"category": "spectral",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "IFFT",
|
|
"category": "spectral",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "RFFT",
|
|
"category": "spectral",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "fft_length",
|
|
"type": "number",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "IRFFT",
|
|
"category": "spectral",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "fft_length",
|
|
"type": "number",
|
|
"notSupported": true
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const spectral = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$2
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json$1 = [
|
|
{
|
|
"tfOpName": "StaticRegexReplace",
|
|
"category": "string",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "input",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "pattern",
|
|
"name": "pattern",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "rewrite",
|
|
"name": "rewrite",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "replace_global",
|
|
"name": "replaceGlobal",
|
|
"type": "bool"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "StringNGrams",
|
|
"category": "string",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "data",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "dataSplits",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "separator",
|
|
"name": "separator",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "ngram_widths",
|
|
"name": "nGramWidths",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"tfName": "left_pad",
|
|
"name": "leftPad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "right_pad",
|
|
"name": "rightPad",
|
|
"type": "string"
|
|
},
|
|
{
|
|
"tfName": "pad_width",
|
|
"name": "padWidth",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "preserve_short_sequences",
|
|
"name": "preserveShortSequences",
|
|
"type": "bool"
|
|
}
|
|
],
|
|
"outputs": [
|
|
"ngrams",
|
|
"ngrams_splits"
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "StringSplit",
|
|
"category": "string",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "input",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "delimiter",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "skip_empty",
|
|
"name": "skipEmpty",
|
|
"type": "bool"
|
|
}
|
|
],
|
|
"outputs": [
|
|
"indices",
|
|
"values",
|
|
"shape"
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "StringToHashBucketFast",
|
|
"category": "string",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "input",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "num_buckets",
|
|
"name": "numBuckets",
|
|
"type": "number"
|
|
}
|
|
]
|
|
}
|
|
];
|
|
const string = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json: json$1
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const json = [
|
|
{
|
|
"tfOpName": "Cast",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "SrcT",
|
|
"name": "sdtype",
|
|
"type": "dtype",
|
|
"notSupported": true
|
|
},
|
|
{
|
|
"tfName": "DstT",
|
|
"name": "dtype",
|
|
"type": "dtype"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "ExpandDims",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "axis",
|
|
"type": "number"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "MirrorPad",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "padding",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "mode",
|
|
"name": "mode",
|
|
"type": "string"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Pad",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "padding",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "constant_value",
|
|
"name": "constantValue",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "PadV2",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "padding",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "constantValue",
|
|
"type": "number",
|
|
"defaultValue": 0
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Reshape",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "EnsureShape",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "Squeeze",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "axis",
|
|
"tfDeprecatedName": "squeeze_dims",
|
|
"name": "axis",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "SpaceToBatchND",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "blockShape",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "paddings",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "BatchToSpaceND",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "blockShape",
|
|
"type": "number[]"
|
|
},
|
|
{
|
|
"start": 2,
|
|
"name": "crops",
|
|
"type": "number[]"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "DepthToSpace",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": [
|
|
{
|
|
"tfName": "block_size",
|
|
"name": "blockSize",
|
|
"type": "number"
|
|
},
|
|
{
|
|
"tfName": "data_format",
|
|
"name": "dataFormat",
|
|
"type": "string"
|
|
}
|
|
]
|
|
},
|
|
{
|
|
"tfOpName": "BroadcastTo",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "x",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "shape",
|
|
"type": "number[]"
|
|
}
|
|
],
|
|
"attrs": []
|
|
},
|
|
{
|
|
"tfOpName": "BroadcastArgs",
|
|
"category": "transformation",
|
|
"inputs": [
|
|
{
|
|
"start": 0,
|
|
"name": "s0",
|
|
"type": "tensor"
|
|
},
|
|
{
|
|
"start": 1,
|
|
"name": "s1",
|
|
"type": "tensor"
|
|
}
|
|
],
|
|
"attrs": []
|
|
}
|
|
];
|
|
const transformation = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
json
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
class OperationMapper {
|
|
// Singleton instance for the mapper
|
|
static get Instance() {
|
|
return this._instance || (this._instance = new this());
|
|
}
|
|
// Loads the op mapping from the JSON file.
|
|
constructor() {
|
|
const ops = [
|
|
arithmetic,
|
|
basicMath,
|
|
control,
|
|
convolution,
|
|
creation,
|
|
dynamic,
|
|
evaluation,
|
|
graph,
|
|
hashTable,
|
|
image,
|
|
logical,
|
|
matrices,
|
|
normalization,
|
|
reduction,
|
|
sliceJoin,
|
|
sparse,
|
|
spectral,
|
|
string,
|
|
transformation
|
|
];
|
|
const mappersJson = [].concat(...ops.map((op2) => op2.json));
|
|
this.opMappers = mappersJson.reduce((map, mapper) => {
|
|
map[mapper.tfOpName] = mapper;
|
|
return map;
|
|
}, {});
|
|
}
|
|
// Converts the model inference graph from Tensorflow GraphDef to local
|
|
// representation for TensorFlow.js API
|
|
transformGraph(graph2, signature = {}) {
|
|
const tfNodes = graph2.node;
|
|
const placeholders = [];
|
|
const weights = [];
|
|
const initNodes = [];
|
|
const nodes = tfNodes.reduce((map, node) => {
|
|
map[node.name] = this.mapNode(node);
|
|
if (node.op.startsWith("Placeholder")) {
|
|
placeholders.push(map[node.name]);
|
|
} else if (node.op === "Const") {
|
|
weights.push(map[node.name]);
|
|
} else if (node.input == null || node.input.length === 0) {
|
|
initNodes.push(map[node.name]);
|
|
}
|
|
return map;
|
|
}, {});
|
|
let inputs = [];
|
|
const outputs = [];
|
|
let inputNodeNameToKey = {};
|
|
let outputNodeNameToKey = {};
|
|
if (signature != null) {
|
|
inputNodeNameToKey = this.mapSignatureEntries(signature.inputs);
|
|
outputNodeNameToKey = this.mapSignatureEntries(signature.outputs);
|
|
}
|
|
const allNodes = Object.keys(nodes);
|
|
allNodes.forEach((key) => {
|
|
const node = nodes[key];
|
|
node.inputNames.forEach((name, index) => {
|
|
const [nodeName, , outputName] = getNodeNameAndIndex(name);
|
|
const inputNode = nodes[nodeName];
|
|
if (inputNode.outputs != null) {
|
|
const outputIndex = inputNode.outputs.indexOf(outputName);
|
|
if (outputIndex !== -1) {
|
|
const inputName = `${nodeName}:${outputIndex}`;
|
|
node.inputNames[index] = inputName;
|
|
}
|
|
}
|
|
node.inputs.push(inputNode);
|
|
inputNode.children.push(node);
|
|
});
|
|
});
|
|
if (Object.keys(outputNodeNameToKey).length === 0) {
|
|
allNodes.forEach((key) => {
|
|
const node = nodes[key];
|
|
if (node.children.length === 0) {
|
|
outputs.push(node);
|
|
}
|
|
});
|
|
} else {
|
|
Object.keys(outputNodeNameToKey).forEach((name) => {
|
|
const [nodeName] = getNodeNameAndIndex(name);
|
|
const node = nodes[nodeName];
|
|
if (node != null) {
|
|
node.signatureKey = outputNodeNameToKey[name];
|
|
outputs.push(node);
|
|
}
|
|
});
|
|
}
|
|
if (Object.keys(inputNodeNameToKey).length > 0) {
|
|
Object.keys(inputNodeNameToKey).forEach((name) => {
|
|
const [nodeName] = getNodeNameAndIndex(name);
|
|
const node = nodes[nodeName];
|
|
if (node) {
|
|
node.signatureKey = inputNodeNameToKey[name];
|
|
inputs.push(node);
|
|
}
|
|
});
|
|
} else {
|
|
inputs = placeholders;
|
|
}
|
|
let functions = {};
|
|
if (graph2.library != null && graph2.library.function != null) {
|
|
functions = graph2.library.function.reduce((functions2, func) => {
|
|
functions2[func.signature.name] = this.mapFunction(func);
|
|
return functions2;
|
|
}, {});
|
|
}
|
|
const result = { nodes, inputs, outputs, weights, placeholders, signature, functions };
|
|
if (initNodes.length > 0) {
|
|
result.initNodes = initNodes;
|
|
}
|
|
return result;
|
|
}
|
|
mapSignatureEntries(entries) {
|
|
return Object.keys(entries || {}).reduce((prev, curr) => {
|
|
prev[entries[curr].name] = curr;
|
|
return prev;
|
|
}, {});
|
|
}
|
|
mapNode(node) {
|
|
const mapper = getRegisteredOp(node.op) || this.opMappers[node.op] || {};
|
|
if (node.attr == null) {
|
|
node.attr = {};
|
|
}
|
|
const newNode = {
|
|
name: node.name,
|
|
op: node.op,
|
|
category: mapper.category,
|
|
inputNames: (node.input || []).map((input) => input.startsWith("^") ? input.slice(1) : input),
|
|
inputs: [],
|
|
children: [],
|
|
inputParams: {},
|
|
attrParams: {},
|
|
rawAttrs: node.attr,
|
|
outputs: mapper.outputs
|
|
};
|
|
if (mapper.inputs != null) {
|
|
newNode.inputParams = mapper.inputs.reduce((map, param) => {
|
|
map[param.name] = {
|
|
type: param.type,
|
|
inputIndexStart: param.start,
|
|
inputIndexEnd: param.end
|
|
};
|
|
return map;
|
|
}, {});
|
|
}
|
|
if (mapper.attrs != null) {
|
|
newNode.attrParams = mapper.attrs.reduce((map, param) => {
|
|
const type = param.type;
|
|
let value = void 0;
|
|
switch (param.type) {
|
|
case "string":
|
|
value = getStringParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getStringParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "string[]":
|
|
value = getStringArrayParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getStringArrayParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "number":
|
|
value = getNumberParam(node.attr, param.tfName, param.defaultValue || 0);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getNumberParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "number[]":
|
|
value = getNumericArrayParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getNumericArrayParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "bool":
|
|
value = getBoolParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getBoolParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "bool[]":
|
|
value = getBoolArrayParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getBoolArrayParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "shape":
|
|
value = getTensorShapeParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getTensorShapeParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "shape[]":
|
|
value = getTensorShapeArrayParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getTensorShapeArrayParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "dtype":
|
|
value = getDtypeParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getDtypeParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "dtype[]":
|
|
value = getDtypeArrayParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getDtypeArrayParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "func":
|
|
value = getFuncParam(node.attr, param.tfName, param.defaultValue);
|
|
if (value === void 0 && !!param.tfDeprecatedName) {
|
|
value = getFuncParam(node.attr, param.tfDeprecatedName, param.defaultValue);
|
|
}
|
|
break;
|
|
case "tensor":
|
|
case "tensors":
|
|
break;
|
|
default:
|
|
throw new Error(`Unsupported param type: ${param.type} for op: ${node.op}`);
|
|
}
|
|
map[param.name] = { value, type };
|
|
return map;
|
|
}, {});
|
|
}
|
|
return newNode;
|
|
}
|
|
// map the TFunctionDef to TFJS graph object
|
|
mapFunction(functionDef) {
|
|
const tfNodes = functionDef.nodeDef;
|
|
const placeholders = [];
|
|
const weights = [];
|
|
let nodes = {};
|
|
if (tfNodes != null) {
|
|
nodes = tfNodes.reduce((map, node) => {
|
|
map[node.name] = this.mapNode(node);
|
|
if (node.op === "Const") {
|
|
weights.push(map[node.name]);
|
|
}
|
|
return map;
|
|
}, {});
|
|
}
|
|
const inputs = [];
|
|
const outputs = [];
|
|
functionDef.signature.inputArg.forEach((arg) => {
|
|
const [nodeName] = getNodeNameAndIndex(arg.name);
|
|
const node = {
|
|
name: nodeName,
|
|
op: "Placeholder",
|
|
inputs: [],
|
|
inputNames: [],
|
|
category: "graph",
|
|
inputParams: {},
|
|
attrParams: { dtype: { value: parseDtypeParam(arg.type), type: "dtype" } },
|
|
children: []
|
|
};
|
|
node.signatureKey = arg.name;
|
|
inputs.push(node);
|
|
nodes[nodeName] = node;
|
|
});
|
|
const allNodes = Object.keys(nodes);
|
|
allNodes.forEach((key) => {
|
|
const node = nodes[key];
|
|
node.inputNames.forEach((name, index) => {
|
|
const [nodeName, , outputName] = getNodeNameAndIndex(name);
|
|
const inputNode = nodes[nodeName];
|
|
if (inputNode.outputs != null) {
|
|
const outputIndex = inputNode.outputs.indexOf(outputName);
|
|
if (outputIndex !== -1) {
|
|
const inputName = `${nodeName}:${outputIndex}`;
|
|
node.inputNames[index] = inputName;
|
|
}
|
|
}
|
|
node.inputs.push(inputNode);
|
|
inputNode.children.push(node);
|
|
});
|
|
});
|
|
const returnNodeMap = functionDef.ret;
|
|
functionDef.signature.outputArg.forEach((output) => {
|
|
const [nodeName, index] = getNodeNameAndIndex(returnNodeMap[output.name]);
|
|
const node = nodes[nodeName];
|
|
if (node != null) {
|
|
node.defaultOutput = index;
|
|
outputs.push(node);
|
|
}
|
|
});
|
|
const signature = this.mapArgsToSignature(functionDef);
|
|
return { nodes, inputs, outputs, weights, placeholders, signature };
|
|
}
|
|
mapArgsToSignature(functionDef) {
|
|
return {
|
|
methodName: functionDef.signature.name,
|
|
inputs: functionDef.signature.inputArg.reduce((map, arg) => {
|
|
map[arg.name] = this.mapArgToTensorInfo(arg);
|
|
return map;
|
|
}, {}),
|
|
outputs: functionDef.signature.outputArg.reduce((map, arg) => {
|
|
map[arg.name] = this.mapArgToTensorInfo(arg, functionDef.ret);
|
|
return map;
|
|
}, {})
|
|
};
|
|
}
|
|
mapArgToTensorInfo(arg, nameMap) {
|
|
let name = arg.name;
|
|
if (nameMap != null) {
|
|
name = nameMap[name];
|
|
}
|
|
return { name, dtype: arg.type };
|
|
}
|
|
}
|
|
function decodeBase64(text) {
|
|
const global2 = env().global;
|
|
if (typeof global2.atob !== "undefined") {
|
|
return global2.atob(text);
|
|
} else if (typeof Buffer !== "undefined") {
|
|
return new Buffer(text, "base64").toString();
|
|
} else {
|
|
throw new Error("Unable to decode base64 in this environment. Missing built-in atob() or Buffer()");
|
|
}
|
|
}
|
|
function parseStringParam(s, keepCase) {
|
|
const value = Array.isArray(s) ? String.fromCharCode.apply(null, s) : decodeBase64(s);
|
|
return keepCase ? value : value.toLowerCase();
|
|
}
|
|
function getStringParam(attrs, name, def, keepCase = false) {
|
|
const param = attrs[name];
|
|
if (param != null) {
|
|
return parseStringParam(param.s, keepCase);
|
|
}
|
|
return def;
|
|
}
|
|
function getBoolParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
return param ? param.b : def;
|
|
}
|
|
function getNumberParam(attrs, name, def) {
|
|
const param = attrs[name] || {};
|
|
const value = param["i"] != null ? param["i"] : param["f"] != null ? param["f"] : def;
|
|
return typeof value === "number" ? value : parseInt(value, 10);
|
|
}
|
|
function parseDtypeParam(value) {
|
|
if (typeof value === "string") {
|
|
value = DataType[value];
|
|
}
|
|
switch (value) {
|
|
case DataType.DT_FLOAT:
|
|
case DataType.DT_HALF:
|
|
return "float32";
|
|
case DataType.DT_INT32:
|
|
case DataType.DT_INT64:
|
|
case DataType.DT_INT8:
|
|
case DataType.DT_UINT8:
|
|
return "int32";
|
|
case DataType.DT_BOOL:
|
|
return "bool";
|
|
case DataType.DT_DOUBLE:
|
|
return "float32";
|
|
case DataType.DT_STRING:
|
|
return "string";
|
|
case DataType.DT_COMPLEX64:
|
|
case DataType.DT_COMPLEX128:
|
|
return "complex64";
|
|
default:
|
|
return null;
|
|
}
|
|
}
|
|
function getFuncParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param && param.func) {
|
|
return param.func.name;
|
|
}
|
|
return def;
|
|
}
|
|
function getDtypeParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param && param.type) {
|
|
return parseDtypeParam(param.type);
|
|
}
|
|
return def;
|
|
}
|
|
function getDtypeArrayParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param && param.list && param.list.type) {
|
|
return param.list.type.map((v) => parseDtypeParam(v));
|
|
}
|
|
return def;
|
|
}
|
|
function parseTensorShapeParam(shape) {
|
|
if (shape.unknownRank) {
|
|
return void 0;
|
|
}
|
|
if (shape.dim != null) {
|
|
return shape.dim.map((dim) => typeof dim.size === "number" ? dim.size : parseInt(dim.size, 10));
|
|
}
|
|
return [];
|
|
}
|
|
function getTensorShapeParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param && param.shape) {
|
|
return parseTensorShapeParam(param.shape);
|
|
}
|
|
return def;
|
|
}
|
|
function getNumericArrayParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param) {
|
|
return ((param.list.f && param.list.f.length ? param.list.f : param.list.i) || []).map((v) => typeof v === "number" ? v : parseInt(v, 10));
|
|
}
|
|
return def;
|
|
}
|
|
function getStringArrayParam(attrs, name, def, keepCase = false) {
|
|
const param = attrs[name];
|
|
if (param && param.list && param.list.s) {
|
|
return param.list.s.map((v) => {
|
|
return parseStringParam(v, keepCase);
|
|
});
|
|
}
|
|
return def;
|
|
}
|
|
function getTensorShapeArrayParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param && param.list && param.list.shape) {
|
|
return param.list.shape.map((v) => {
|
|
return parseTensorShapeParam(v);
|
|
});
|
|
}
|
|
return def;
|
|
}
|
|
function getBoolArrayParam(attrs, name, def) {
|
|
const param = attrs[name];
|
|
if (param && param.list && param.list.b) {
|
|
return param.list.b;
|
|
}
|
|
return def;
|
|
}
|
|
class NodeValueImpl {
|
|
constructor(node, tensorMap, context) {
|
|
this.node = node;
|
|
this.tensorMap = tensorMap;
|
|
this.context = context;
|
|
this.inputs = [];
|
|
this.attrs = {};
|
|
this.inputs = node.inputNames.map((name) => this.getInput(name));
|
|
if (node.rawAttrs != null) {
|
|
this.attrs = Object.keys(node.rawAttrs).reduce((attrs, key) => {
|
|
attrs[key] = this.getAttr(key);
|
|
return attrs;
|
|
}, {});
|
|
}
|
|
}
|
|
/**
|
|
* Return the value of the attribute or input param.
|
|
* @param name String: name of attribute or input param.
|
|
*/
|
|
getInput(name) {
|
|
return getTensor(name, this.tensorMap, this.context);
|
|
}
|
|
/**
|
|
* Return the value of the attribute or input param.
|
|
* @param name String: name of attribute or input param.
|
|
*/
|
|
getAttr(name, defaultValue) {
|
|
const value = this.node.rawAttrs[name];
|
|
if (value.tensor != null) {
|
|
return getTensor(name, this.tensorMap, this.context);
|
|
}
|
|
if (value.i != null || value.f != null) {
|
|
return getNumberParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.s != null) {
|
|
return getStringParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.b != null) {
|
|
return getBoolParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.shape != null) {
|
|
return getTensorShapeParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.type != null) {
|
|
return getDtypeParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.list != null) {
|
|
if (value.list.i != null || value.list.f != null) {
|
|
return getNumericArrayParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.list.s != null) {
|
|
return getStringArrayParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.list.shape != null) {
|
|
return getTensorShapeArrayParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.list.b != null) {
|
|
return getBoolArrayParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
if (value.list.type != null) {
|
|
return getDtypeArrayParam(this.node.rawAttrs, name, defaultValue);
|
|
}
|
|
}
|
|
return defaultValue;
|
|
}
|
|
}
|
|
const tfOps = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
OP_SCOPE_SUFFIX,
|
|
abs: abs$2,
|
|
acos: acos$2,
|
|
acosh: acosh$2,
|
|
add: add$1,
|
|
addN: addN$2,
|
|
all: all$2,
|
|
any: any$2,
|
|
argMax: argMax$2,
|
|
argMin: argMin$2,
|
|
asin: asin$2,
|
|
asinh: asinh$2,
|
|
atan: atan$2,
|
|
atan2: atan2$2,
|
|
atanh: atanh$2,
|
|
avgPool: avgPool$2,
|
|
avgPool3d,
|
|
basicLSTMCell,
|
|
batchNorm: batchNorm$2,
|
|
batchNorm2d,
|
|
batchNorm3d,
|
|
batchNorm4d,
|
|
batchToSpaceND: batchToSpaceND$2,
|
|
bincount: bincount$2,
|
|
bitwiseAnd: bitwiseAnd$2,
|
|
booleanMaskAsync,
|
|
broadcastArgs: broadcastArgs$2,
|
|
broadcastTo,
|
|
buffer,
|
|
cast: cast$2,
|
|
ceil: ceil$2,
|
|
clipByValue: clipByValue$2,
|
|
clone,
|
|
complex: complex$2,
|
|
concat: concat$2,
|
|
concat1d,
|
|
concat2d,
|
|
concat3d,
|
|
concat4d,
|
|
conv1d,
|
|
conv2d: conv2d$2,
|
|
conv2dTranspose,
|
|
conv3d,
|
|
conv3dTranspose,
|
|
cos: cos$2,
|
|
cosh: cosh$2,
|
|
cosineWindow,
|
|
cumprod: cumprod$2,
|
|
cumsum: cumsum$2,
|
|
denseBincount: denseBincount$2,
|
|
depthToSpace: depthToSpace$2,
|
|
depthwiseConv2d: depthwiseConv2d$1,
|
|
diag: diag$2,
|
|
dilation2d,
|
|
div: div$1,
|
|
divNoNan,
|
|
dot,
|
|
dropout,
|
|
einsum: einsum$2,
|
|
elu: elu$2,
|
|
enclosingPowerOfTwo,
|
|
ensureShape,
|
|
equal: equal$2,
|
|
erf: erf$2,
|
|
euclideanNorm,
|
|
exp: exp$2,
|
|
expandDims: expandDims$2,
|
|
expm1: expm1$2,
|
|
eye,
|
|
fft: fft$2,
|
|
fill: fill$2,
|
|
floor: floor$2,
|
|
floorDiv: floorDiv$2,
|
|
fused: fused_ops,
|
|
gather,
|
|
gatherND,
|
|
greater: greater$2,
|
|
greaterEqual: greaterEqual$2,
|
|
ifft: ifft$2,
|
|
imag: imag$2,
|
|
image: image$1,
|
|
inTopKAsync,
|
|
irfft,
|
|
isFinite: isFinite$3,
|
|
isInf: isInf$2,
|
|
isNaN: isNaN$3,
|
|
leakyRelu: leakyRelu$2,
|
|
less: less$2,
|
|
lessEqual: lessEqual$2,
|
|
linalg,
|
|
linspace,
|
|
localResponseNormalization,
|
|
log: log$2,
|
|
log1p: log1p$2,
|
|
logSigmoid,
|
|
logSoftmax,
|
|
logSumExp,
|
|
logicalAnd: logicalAnd$2,
|
|
logicalNot: logicalNot$2,
|
|
logicalOr: logicalOr$2,
|
|
logicalXor,
|
|
losses,
|
|
lowerBound: lowerBound$1,
|
|
matMul: matMul$1,
|
|
max: max$2,
|
|
maxPool: maxPool$2,
|
|
maxPool3d: maxPool3d$1,
|
|
maxPoolWithArgmax,
|
|
maximum: maximum$2,
|
|
mean: mean$1,
|
|
meshgrid,
|
|
min: min$2,
|
|
minimum: minimum$2,
|
|
mirrorPad: mirrorPad$1,
|
|
mod: mod$2,
|
|
moments,
|
|
movingAverage,
|
|
mul,
|
|
multiRNNCell,
|
|
multinomial: multinomial$2,
|
|
neg: neg$2,
|
|
norm,
|
|
notEqual: notEqual$2,
|
|
oneHot: oneHot$2,
|
|
ones,
|
|
onesLike: onesLike$2,
|
|
op,
|
|
outerProduct,
|
|
pad,
|
|
pad1d,
|
|
pad2d,
|
|
pad3d,
|
|
pad4d,
|
|
pool: pool$1,
|
|
pow: pow$2,
|
|
prelu: prelu$2,
|
|
print,
|
|
prod: prod$2,
|
|
raggedGather: raggedGather$2,
|
|
raggedRange: raggedRange$2,
|
|
raggedTensorToTensor: raggedTensorToTensor$2,
|
|
rand,
|
|
randomGamma,
|
|
randomNormal,
|
|
randomStandardNormal,
|
|
randomUniform,
|
|
randomUniformInt,
|
|
range: range$2,
|
|
real: real$2,
|
|
reciprocal: reciprocal$2,
|
|
relu: relu$2,
|
|
relu6: relu6$2,
|
|
reshape: reshape$2,
|
|
reverse: reverse$2,
|
|
reverse1d,
|
|
reverse2d,
|
|
reverse3d,
|
|
reverse4d,
|
|
rfft,
|
|
round: round$2,
|
|
rsqrt: rsqrt$2,
|
|
scalar,
|
|
scatterND,
|
|
searchSorted: searchSorted$2,
|
|
selu: selu$2,
|
|
separableConv2d,
|
|
setdiff1dAsync,
|
|
sigmoid: sigmoid$2,
|
|
sign: sign$2,
|
|
signal,
|
|
sin: sin$2,
|
|
sinh: sinh$2,
|
|
slice: slice$2,
|
|
slice1d,
|
|
slice2d,
|
|
slice3d,
|
|
slice4d,
|
|
softmax: softmax$2,
|
|
softplus: softplus$2,
|
|
spaceToBatchND: spaceToBatchND$2,
|
|
sparse: sparse$1,
|
|
sparseToDense: sparseToDense$2,
|
|
spectral: spectral$1,
|
|
split: split$2,
|
|
sqrt: sqrt$2,
|
|
square: square$1,
|
|
squaredDifference: squaredDifference$2,
|
|
squeeze,
|
|
stack,
|
|
step: step$2,
|
|
stridedSlice: stridedSlice$2,
|
|
string: string$1,
|
|
sub: sub$2,
|
|
sum: sum$2,
|
|
tan: tan$2,
|
|
tanh: tanh$2,
|
|
tensor,
|
|
tensor1d,
|
|
tensor2d,
|
|
tensor3d,
|
|
tensor4d,
|
|
tensor5d,
|
|
tensor6d,
|
|
tensorScatterUpdate: tensorScatterUpdate$2,
|
|
tile: tile$2,
|
|
topk,
|
|
transpose: transpose$2,
|
|
truncatedNormal,
|
|
unique: unique$2,
|
|
unsortedSegmentSum: unsortedSegmentSum$2,
|
|
unstack,
|
|
upperBound: upperBound$1,
|
|
variable,
|
|
where,
|
|
whereAsync,
|
|
zeros: zeros$1,
|
|
zerosLike: zerosLike$2
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const executeOp$k = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "BiasAdd":
|
|
case "AddV2":
|
|
case "Add": {
|
|
return [ops.add(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "AddN": {
|
|
return [ops.addN(getParamValue("tensors", node, tensorMap, context))];
|
|
}
|
|
case "FloorMod":
|
|
case "Mod":
|
|
return [ops.mod(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
case "Mul":
|
|
return [ops.mul(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
case "RealDiv":
|
|
case "Div": {
|
|
return [ops.div(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "DivNoNan": {
|
|
return [ops.divNoNan(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "FloorDiv": {
|
|
return [ops.floorDiv(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Sub": {
|
|
return [ops.sub(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Minimum": {
|
|
return [ops.minimum(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Maximum": {
|
|
return [ops.maximum(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Pow": {
|
|
return [ops.pow(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "SquaredDifference": {
|
|
return [ops.squaredDifference(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$j = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Abs":
|
|
case "ComplexAbs":
|
|
return [ops.abs(getParamValue("x", node, tensorMap, context))];
|
|
case "Acos":
|
|
return [ops.acos(getParamValue("x", node, tensorMap, context))];
|
|
case "Acosh":
|
|
return [ops.acosh(getParamValue("x", node, tensorMap, context))];
|
|
case "Asin":
|
|
return [ops.asin(getParamValue("x", node, tensorMap, context))];
|
|
case "Asinh":
|
|
return [ops.asinh(getParamValue("x", node, tensorMap, context))];
|
|
case "Atan":
|
|
return [ops.atan(getParamValue("x", node, tensorMap, context))];
|
|
case "Atan2":
|
|
return [ops.atan2(getParamValue("x", node, tensorMap, context), getParamValue("y", node, tensorMap, context))];
|
|
case "Atanh":
|
|
return [ops.atanh(getParamValue("x", node, tensorMap, context))];
|
|
case "Ceil":
|
|
return [ops.ceil(getParamValue("x", node, tensorMap, context))];
|
|
case "Complex":
|
|
return [ops.complex(getParamValue("real", node, tensorMap, context), getParamValue("imag", node, tensorMap, context))];
|
|
case "Cos":
|
|
return [ops.cos(getParamValue("x", node, tensorMap, context))];
|
|
case "Cosh":
|
|
return [ops.cosh(getParamValue("x", node, tensorMap, context))];
|
|
case "Elu":
|
|
return [ops.elu(getParamValue("x", node, tensorMap, context))];
|
|
case "Erf":
|
|
return [ops.erf(getParamValue("x", node, tensorMap, context))];
|
|
case "Exp":
|
|
return [ops.exp(getParamValue("x", node, tensorMap, context))];
|
|
case "Expm1": {
|
|
return [ops.expm1(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Floor":
|
|
return [ops.floor(getParamValue("x", node, tensorMap, context))];
|
|
case "Log":
|
|
return [ops.log(getParamValue("x", node, tensorMap, context))];
|
|
case "Log1p": {
|
|
return [ops.log1p(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Imag":
|
|
return [ops.imag(getParamValue("x", node, tensorMap, context))];
|
|
case "Neg":
|
|
return [ops.neg(getParamValue("x", node, tensorMap, context))];
|
|
case "Reciprocal": {
|
|
return [ops.reciprocal(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Real":
|
|
return [ops.real(getParamValue("x", node, tensorMap, context))];
|
|
case "Relu":
|
|
return [ops.relu(getParamValue("x", node, tensorMap, context))];
|
|
case "Round": {
|
|
return [ops.round(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Selu":
|
|
return [ops.selu(getParamValue("x", node, tensorMap, context))];
|
|
case "Sigmoid":
|
|
return [ops.sigmoid(getParamValue("x", node, tensorMap, context))];
|
|
case "Sin":
|
|
return [ops.sin(getParamValue("x", node, tensorMap, context))];
|
|
case "Sign": {
|
|
return [ops.sign(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Sinh": {
|
|
return [ops.sinh(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Softplus": {
|
|
return [ops.softplus(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Sqrt": {
|
|
return [ops.sqrt(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Square": {
|
|
return [ops.square(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Tanh": {
|
|
return [ops.tanh(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "Tan":
|
|
return [ops.tan(getParamValue("x", node, tensorMap, context))];
|
|
case "ClipByValue":
|
|
return [ops.clipByValue(getParamValue("x", node, tensorMap, context), getParamValue("clipValueMin", node, tensorMap, context), getParamValue("clipValueMax", node, tensorMap, context))];
|
|
case "Relu6":
|
|
return [ops.relu6(getParamValue("x", node, tensorMap, context))];
|
|
case "Rsqrt":
|
|
return [ops.rsqrt(getTensor(node.inputNames[0], tensorMap, context))];
|
|
case "LeakyRelu":
|
|
return [ops.leakyRelu(getParamValue("x", node, tensorMap, context), getParamValue("alpha", node, tensorMap, context))];
|
|
case "Prelu":
|
|
return [ops.prelu(getParamValue("x", node, tensorMap, context), getParamValue("alpha", node, tensorMap, context))];
|
|
case "IsNan":
|
|
return [ops.isNaN(getTensor(node.inputNames[0], tensorMap, context))];
|
|
case "IsInf":
|
|
return [ops.isInf(getTensor(node.inputNames[0], tensorMap, context))];
|
|
case "IsFinite":
|
|
return [ops.isFinite(getTensor(node.inputNames[0], tensorMap, context))];
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
function assertShapesMatchAllowUndefinedSize(shapeA, shapeB, errorMessagePrefix = "") {
|
|
if (typeof shapeA === "number" || typeof shapeB === "number") {
|
|
return;
|
|
}
|
|
assert(shapeA.length === shapeB.length, () => errorMessagePrefix + ` Shapes ${shapeA} and ${shapeB} must match`);
|
|
for (let i = 0; i < shapeA.length; i++) {
|
|
const dim0 = shapeA[i];
|
|
const dim1 = shapeB[i];
|
|
assert(dim0 < 0 || dim1 < 0 || dim0 === dim1, () => errorMessagePrefix + ` Shapes ${shapeA} and ${shapeB} must match`);
|
|
}
|
|
}
|
|
function fullDefinedShape(elementShape) {
|
|
if (typeof elementShape === "number" || elementShape.some((dim) => dim < 0)) {
|
|
return false;
|
|
}
|
|
return true;
|
|
}
|
|
function inferElementShape(listElementShape, tensors, elementShape) {
|
|
let partialShape = mergeElementShape(listElementShape, elementShape);
|
|
const notfullDefinedShape = !fullDefinedShape(partialShape);
|
|
if (notfullDefinedShape && tensors.length === 0) {
|
|
throw new Error(`Tried to calculate elements of an empty list with non-fully-defined elementShape: ${partialShape}`);
|
|
}
|
|
if (notfullDefinedShape) {
|
|
tensors.forEach((tensor2) => {
|
|
partialShape = mergeElementShape(tensor2.shape, partialShape);
|
|
});
|
|
}
|
|
if (!fullDefinedShape(partialShape)) {
|
|
throw new Error(`Non-fully-defined elementShape: ${partialShape}`);
|
|
}
|
|
return partialShape;
|
|
}
|
|
function mergeElementShape(elementShapeA, elementShapeB) {
|
|
if (typeof elementShapeA === "number") {
|
|
return elementShapeB;
|
|
}
|
|
if (typeof elementShapeB === "number") {
|
|
return elementShapeA;
|
|
}
|
|
if (elementShapeA.length !== elementShapeB.length) {
|
|
throw new Error(`Incompatible ranks during merge: ${elementShapeA} vs. ${elementShapeB}`);
|
|
}
|
|
const result = [];
|
|
for (let i = 0; i < elementShapeA.length; ++i) {
|
|
const dim0 = elementShapeA[i];
|
|
const dim1 = elementShapeB[i];
|
|
if (dim0 >= 0 && dim1 >= 0 && dim0 !== dim1) {
|
|
throw new Error(`Incompatible shape during merge: ${elementShapeA} vs. ${elementShapeB}`);
|
|
}
|
|
result[i] = dim0 >= 0 ? dim0 : dim1;
|
|
}
|
|
return result;
|
|
}
|
|
class TensorArray {
|
|
constructor(name, dtype, maxSize, elementShape, identicalElementShapes, dynamicSize, clearAfterRead) {
|
|
this.name = name;
|
|
this.dtype = dtype;
|
|
this.maxSize = maxSize;
|
|
this.elementShape = elementShape;
|
|
this.identicalElementShapes = identicalElementShapes;
|
|
this.dynamicSize = dynamicSize;
|
|
this.clearAfterRead = clearAfterRead;
|
|
this.tensors = [];
|
|
this.closed_ = false;
|
|
this.idTensor = scalar(0);
|
|
keep(this.idTensor);
|
|
}
|
|
get id() {
|
|
return this.idTensor.id;
|
|
}
|
|
get closed() {
|
|
return this.closed_;
|
|
}
|
|
/**
|
|
* Dispose the tensors and idTensor and mark the TensoryArray as closed.
|
|
*/
|
|
clearAndClose(keepIds) {
|
|
this.tensors.forEach((tensor2) => {
|
|
if (keepIds == null || !keepIds.has(tensor2.tensor.id)) {
|
|
tensor2.tensor.dispose();
|
|
}
|
|
});
|
|
this.tensors = [];
|
|
this.closed_ = true;
|
|
this.idTensor.dispose();
|
|
}
|
|
size() {
|
|
return this.tensors.length;
|
|
}
|
|
/**
|
|
* Read the value at location index in the TensorArray.
|
|
* @param index Number the index to read from.
|
|
*/
|
|
read(index) {
|
|
if (this.closed_) {
|
|
throw new Error(`TensorArray ${this.name} has already been closed.`);
|
|
}
|
|
if (index < 0 || index >= this.size()) {
|
|
throw new Error(`Tried to read from index ${index}, but array size is: ${this.size()}`);
|
|
}
|
|
const tensorWithState = this.tensors[index];
|
|
if (tensorWithState.cleared) {
|
|
throw new Error(`TensorArray ${this.name}: Could not read index ${index} twice because it was cleared after a previous read (perhaps try setting clear_after_read = false?).`);
|
|
}
|
|
if (this.clearAfterRead) {
|
|
tensorWithState.cleared = true;
|
|
}
|
|
tensorWithState.read = true;
|
|
return tensorWithState.tensor;
|
|
}
|
|
/**
|
|
* Helper method to read multiple tensors from the specified indices.
|
|
*/
|
|
readMany(indices) {
|
|
return indices.map((index) => this.read(index));
|
|
}
|
|
/**
|
|
* Write value into the index of the TensorArray.
|
|
* @param index number the index to write to.
|
|
* @param tensor
|
|
*/
|
|
write(index, tensor2) {
|
|
if (this.closed_) {
|
|
throw new Error(`TensorArray ${this.name} has already been closed.`);
|
|
}
|
|
if (index < 0 || !this.dynamicSize && index >= this.maxSize) {
|
|
throw new Error(`Tried to write to index ${index}, but array is not resizeable and size is: ${this.maxSize}`);
|
|
}
|
|
const t2 = this.tensors[index] || {};
|
|
if (tensor2.dtype !== this.dtype) {
|
|
throw new Error(`TensorArray ${this.name}: Could not write to TensorArray index ${index},
|
|
because the value dtype is ${tensor2.dtype}, but TensorArray dtype is ${this.dtype}.`);
|
|
}
|
|
if (this.size() === 0 && (this.elementShape == null || this.elementShape.length === 0)) {
|
|
this.elementShape = tensor2.shape;
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(this.elementShape, tensor2.shape, `TensorArray ${this.name}: Could not write to TensorArray index ${index}.`);
|
|
if (t2.read) {
|
|
throw new Error(`TensorArray ${this.name}: Could not write to TensorArray index ${index}, because it has already been read.`);
|
|
}
|
|
if (t2.written) {
|
|
throw new Error(`TensorArray ${this.name}: Could not write to TensorArray index ${index}, because it has already been written.`);
|
|
}
|
|
t2.tensor = tensor2;
|
|
keep(tensor2);
|
|
t2.written = true;
|
|
this.tensors[index] = t2;
|
|
}
|
|
/**
|
|
* Helper method to write multiple tensors to the specified indices.
|
|
*/
|
|
writeMany(indices, tensors) {
|
|
if (indices.length !== tensors.length) {
|
|
throw new Error(`TensorArray ${this.name}: could not write multiple tensors,because the index size: ${indices.length} is not the same as tensors size: ${tensors.length}.`);
|
|
}
|
|
indices.forEach((i, index) => this.write(i, tensors[index]));
|
|
}
|
|
/**
|
|
* Return selected values in the TensorArray as a packed Tensor. All of
|
|
* selected values must have been written and their shapes must all match.
|
|
* @param [indices] number[] Optional. Taking values in [0, max_value). If the
|
|
* TensorArray is not dynamic, max_value=size(). If not specified returns
|
|
* all tensors in the original order.
|
|
* @param [dtype]
|
|
*/
|
|
gather(indices, dtype) {
|
|
if (!!dtype && dtype !== this.dtype) {
|
|
throw new Error(`TensorArray dtype is ${this.dtype} but gather requested dtype ${dtype}`);
|
|
}
|
|
if (!indices) {
|
|
indices = [];
|
|
for (let i = 0; i < this.size(); i++) {
|
|
indices.push(i);
|
|
}
|
|
} else {
|
|
indices = indices.slice(0, this.size());
|
|
}
|
|
if (indices.length === 0) {
|
|
return tensor([], [0].concat(this.elementShape));
|
|
}
|
|
const tensors = this.readMany(indices);
|
|
assertShapesMatchAllowUndefinedSize(this.elementShape, tensors[0].shape, "TensorArray shape mismatch: ");
|
|
return stack(tensors, 0);
|
|
}
|
|
/**
|
|
* Return the values in the TensorArray as a concatenated Tensor.
|
|
*/
|
|
concat(dtype) {
|
|
if (!!dtype && dtype !== this.dtype) {
|
|
throw new Error(`TensorArray dtype is ${this.dtype} but concat requested dtype ${dtype}`);
|
|
}
|
|
if (this.size() === 0) {
|
|
return tensor([], [0].concat(this.elementShape));
|
|
}
|
|
const indices = [];
|
|
for (let i = 0; i < this.size(); i++) {
|
|
indices.push(i);
|
|
}
|
|
const tensors = this.readMany(indices);
|
|
assertShapesMatchAllowUndefinedSize(this.elementShape, tensors[0].shape, `TensorArray shape mismatch: tensor array shape (${this.elementShape}) vs first tensor shape (${tensors[0].shape})`);
|
|
return concat$2(tensors, 0);
|
|
}
|
|
/**
|
|
* Scatter the values of a Tensor in specific indices of a TensorArray.
|
|
* @param indices number[] values in [0, max_value). If the
|
|
* TensorArray is not dynamic, max_value=size().
|
|
* @param tensor Tensor input tensor.
|
|
*/
|
|
scatter(indices, tensor2) {
|
|
if (tensor2.dtype !== this.dtype) {
|
|
throw new Error(`TensorArray dtype is ${this.dtype} but tensor has dtype ${tensor2.dtype}`);
|
|
}
|
|
if (indices.length !== tensor2.shape[0]) {
|
|
throw new Error(`Expected len(indices) == tensor.shape[0], but saw: ${indices.length} vs. ${tensor2.shape[0]}`);
|
|
}
|
|
const maxIndex = Math.max(...indices);
|
|
if (!this.dynamicSize && maxIndex >= this.maxSize) {
|
|
throw new Error(`Max index must be < array size (${maxIndex} vs. ${this.maxSize})`);
|
|
}
|
|
this.writeMany(indices, unstack(tensor2, 0));
|
|
}
|
|
/**
|
|
* Split the values of a Tensor into the TensorArray.
|
|
* @param length number[] with the lengths to use when splitting value along
|
|
* its first dimension.
|
|
* @param tensor Tensor, the tensor to split.
|
|
*/
|
|
split(length, tensor2) {
|
|
if (tensor2.dtype !== this.dtype) {
|
|
throw new Error(`TensorArray dtype is ${this.dtype} but tensor has dtype ${tensor2.dtype}`);
|
|
}
|
|
let totalLength = 0;
|
|
const cumulativeLengths = length.map((len) => {
|
|
totalLength += len;
|
|
return totalLength;
|
|
});
|
|
if (totalLength !== tensor2.shape[0]) {
|
|
throw new Error(`Expected sum of lengths to be equal to
|
|
tensor.shape[0], but sum of lengths is
|
|
${totalLength}, and tensor's shape is: ${tensor2.shape}`);
|
|
}
|
|
if (!this.dynamicSize && length.length !== this.maxSize) {
|
|
throw new Error(`TensorArray's size is not equal to the size of lengths (${this.maxSize} vs. ${length.length}), and the TensorArray is not marked as dynamically resizeable`);
|
|
}
|
|
const elementPerRow = totalLength === 0 ? 0 : tensor2.size / totalLength;
|
|
const tensors = [];
|
|
tidy(() => {
|
|
tensor2 = reshape$2(tensor2, [1, totalLength, elementPerRow]);
|
|
for (let i = 0; i < length.length; ++i) {
|
|
const previousLength = i === 0 ? 0 : cumulativeLengths[i - 1];
|
|
const indices2 = [0, previousLength, 0];
|
|
const sizes = [1, length[i], elementPerRow];
|
|
tensors[i] = reshape$2(slice$2(tensor2, indices2, sizes), this.elementShape);
|
|
}
|
|
return tensors;
|
|
});
|
|
const indices = [];
|
|
for (let i = 0; i < length.length; i++) {
|
|
indices[i] = i;
|
|
}
|
|
this.writeMany(indices, tensors);
|
|
}
|
|
}
|
|
class TensorList {
|
|
get id() {
|
|
return this.idTensor.id;
|
|
}
|
|
/**
|
|
*
|
|
* @param tensors list of tensors
|
|
* @param elementShape shape of each tensor, this can be a single number (any
|
|
* shape is allowed) or partial shape (dim = -1).
|
|
* @param elementDtype data type of each tensor
|
|
* @param maxNumElements The maximum allowed size of `tensors`. Defaults to -1
|
|
* meaning that the size of `tensors` is unbounded.
|
|
*/
|
|
constructor(tensors, elementShape, elementDtype, maxNumElements = -1) {
|
|
this.tensors = tensors;
|
|
this.elementShape = elementShape;
|
|
this.elementDtype = elementDtype;
|
|
if (tensors != null) {
|
|
tensors.forEach((tensor2) => {
|
|
if (elementDtype !== tensor2.dtype) {
|
|
throw new Error(`Invalid data types; op elements ${elementDtype}, but list elements ${tensor2.dtype}`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(elementShape, tensor2.shape, "TensorList shape mismatch: ");
|
|
keep(tensor2);
|
|
});
|
|
}
|
|
this.idTensor = scalar(0);
|
|
this.maxNumElements = maxNumElements;
|
|
keep(this.idTensor);
|
|
}
|
|
/**
|
|
* Get a new TensorList containing a copy of the underlying tensor container.
|
|
*/
|
|
copy() {
|
|
return new TensorList([...this.tensors], this.elementShape, this.elementDtype);
|
|
}
|
|
/**
|
|
* Dispose the tensors and idTensor and clear the tensor list.
|
|
*/
|
|
clearAndClose(keepIds) {
|
|
this.tensors.forEach((tensor2) => {
|
|
if (keepIds == null || !keepIds.has(tensor2.id)) {
|
|
tensor2.dispose();
|
|
}
|
|
});
|
|
this.tensors.length = 0;
|
|
this.idTensor.dispose();
|
|
}
|
|
/**
|
|
* The size of the tensors in the tensor list.
|
|
*/
|
|
size() {
|
|
return this.tensors.length;
|
|
}
|
|
/**
|
|
* Return a tensor that stacks a list of rank-R tf.Tensors into one rank-(R+1)
|
|
* tf.Tensor.
|
|
* @param elementShape shape of each tensor
|
|
* @param elementDtype data type of each tensor
|
|
* @param numElements the number of elements to stack
|
|
*/
|
|
stack(elementShape, elementDtype, numElements = -1) {
|
|
if (elementDtype !== this.elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${elementDtype}, but list elements ${this.elementDtype}`);
|
|
}
|
|
if (numElements !== -1 && this.tensors.length !== numElements) {
|
|
throw new Error(`Operation expected a list with ${numElements} elements but got a list with ${this.tensors.length} elements.`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(elementShape, this.elementShape, "TensorList shape mismatch: ");
|
|
const outputElementShape = inferElementShape(this.elementShape, this.tensors, elementShape);
|
|
return tidy(() => {
|
|
const reshapedTensors = this.tensors.map((tensor2) => reshape$2(tensor2, outputElementShape));
|
|
return stack(reshapedTensors, 0);
|
|
});
|
|
}
|
|
/**
|
|
* Pop a tensor from the end of the list.
|
|
* @param elementShape shape of the tensor
|
|
* @param elementDtype data type of the tensor
|
|
*/
|
|
popBack(elementShape, elementDtype) {
|
|
if (elementDtype !== this.elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${elementDtype}, but list elements ${this.elementDtype}`);
|
|
}
|
|
if (this.size() === 0) {
|
|
throw new Error("Trying to pop from an empty list.");
|
|
}
|
|
const outputElementShape = inferElementShape(this.elementShape, this.tensors, elementShape);
|
|
const tensor2 = this.tensors.pop();
|
|
tensor2.kept = false;
|
|
assertShapesMatchAllowUndefinedSize(tensor2.shape, elementShape, "TensorList shape mismatch: ");
|
|
return reshape$2(tensor2, outputElementShape);
|
|
}
|
|
/**
|
|
* Push a tensor to the end of the list.
|
|
* @param tensor Tensor to be pushed.
|
|
*/
|
|
pushBack(tensor2) {
|
|
if (tensor2.dtype !== this.elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${tensor2.dtype}, but list elements ${this.elementDtype}`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(tensor2.shape, this.elementShape, "TensorList shape mismatch: ");
|
|
if (this.maxNumElements === this.size()) {
|
|
throw new Error(`Trying to push element into a full list.`);
|
|
}
|
|
keep(tensor2);
|
|
this.tensors.push(tensor2);
|
|
}
|
|
/**
|
|
* Update the size of the list.
|
|
* @param size the new size of the list.
|
|
*/
|
|
resize(size) {
|
|
if (size < 0) {
|
|
throw new Error(`TensorListResize expects size to be non-negative. Got: ${size}`);
|
|
}
|
|
if (this.maxNumElements !== -1 && size > this.maxNumElements) {
|
|
throw new Error(`TensorListResize input size ${size} is greater maxNumElement ${this.maxNumElements}.`);
|
|
}
|
|
const destTensorList = new TensorList([], this.elementShape, this.elementDtype, this.maxNumElements);
|
|
destTensorList.tensors.length = size;
|
|
for (let i = 0; i < Math.min(this.tensors.length, size); ++i) {
|
|
destTensorList.tensors[i] = this.tensors[i];
|
|
}
|
|
return destTensorList;
|
|
}
|
|
/**
|
|
* Retrieve the element at the provided index
|
|
* @param elementShape shape of the tensor
|
|
* @param elementDtype dtype of the tensor
|
|
* @param elementIndex index of the tensor
|
|
*/
|
|
getItem(elementIndex, elementShape, elementDtype) {
|
|
if (elementDtype !== this.elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${elementDtype}, but list elements ${this.elementDtype}`);
|
|
}
|
|
if (elementIndex < 0 || elementIndex > this.tensors.length) {
|
|
throw new Error(`Trying to access element ${elementIndex} in a list with ${this.tensors.length} elements.`);
|
|
}
|
|
if (this.tensors[elementIndex] == null) {
|
|
throw new Error(`element at index ${elementIndex} is null.`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(this.tensors[elementIndex].shape, elementShape, "TensorList shape mismatch: ");
|
|
const outputElementShape = inferElementShape(this.elementShape, this.tensors, elementShape);
|
|
return reshape$2(this.tensors[elementIndex], outputElementShape);
|
|
}
|
|
/**
|
|
* Set the tensor at the index
|
|
* @param elementIndex index of the tensor
|
|
* @param tensor the tensor to be inserted into the list
|
|
*/
|
|
setItem(elementIndex, tensor2) {
|
|
if (tensor2.dtype !== this.elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${tensor2.dtype}, but list elements ${this.elementDtype}`);
|
|
}
|
|
if (elementIndex < 0 || this.maxNumElements !== -1 && elementIndex >= this.maxNumElements) {
|
|
throw new Error(`Trying to set element ${elementIndex} in a list with max ${this.maxNumElements} elements.`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(this.elementShape, tensor2.shape, "TensorList shape mismatch: ");
|
|
keep(tensor2);
|
|
if (this.tensors[elementIndex] != null) {
|
|
this.tensors[elementIndex].kept = false;
|
|
}
|
|
this.tensors[elementIndex] = tensor2;
|
|
}
|
|
/**
|
|
* Return selected values in the TensorList as a stacked Tensor. All of
|
|
* selected values must have been written and their shapes must all match.
|
|
* @param indices indices of tensors to gather
|
|
* @param elementDtype output tensor dtype
|
|
* @param elementShape output tensor element shape
|
|
*/
|
|
gather(indices, elementDtype, elementShape) {
|
|
if (elementDtype !== this.elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${elementDtype}, but list elements ${this.elementDtype}`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(this.elementShape, elementShape, "TensorList shape mismatch: ");
|
|
indices = indices.slice(0, this.size());
|
|
const outputElementShape = inferElementShape(this.elementShape, this.tensors, elementShape);
|
|
if (indices.length === 0) {
|
|
return tensor([], [0].concat(outputElementShape));
|
|
}
|
|
return tidy(() => {
|
|
const tensors = indices.map((i) => reshape$2(this.tensors[i], outputElementShape));
|
|
return stack(tensors, 0);
|
|
});
|
|
}
|
|
/**
|
|
* Return the values in the TensorList as a concatenated Tensor.
|
|
* @param elementDtype output tensor dtype
|
|
* @param elementShape output tensor element shape
|
|
*/
|
|
concat(elementDtype, elementShape) {
|
|
if (!!elementDtype && elementDtype !== this.elementDtype) {
|
|
throw new Error(`TensorList dtype is ${this.elementDtype} but concat requested dtype ${elementDtype}`);
|
|
}
|
|
assertShapesMatchAllowUndefinedSize(this.elementShape, elementShape, "TensorList shape mismatch: ");
|
|
const outputElementShape = inferElementShape(this.elementShape, this.tensors, elementShape);
|
|
if (this.size() === 0) {
|
|
return tensor([], [0].concat(outputElementShape));
|
|
}
|
|
return tidy(() => {
|
|
const tensors = this.tensors.map((t2) => reshape$2(t2, outputElementShape));
|
|
return concat$2(tensors, 0);
|
|
});
|
|
}
|
|
}
|
|
function fromTensor(tensor2, elementShape, elementDtype) {
|
|
const dtype = tensor2.dtype;
|
|
if (tensor2.shape.length < 1) {
|
|
throw new Error(`Tensor must be at least a vector, but saw shape: ${tensor2.shape}`);
|
|
}
|
|
if (tensor2.dtype !== elementDtype) {
|
|
throw new Error(`Invalid data types; op elements ${tensor2.dtype}, but list elements ${elementDtype}`);
|
|
}
|
|
const tensorElementShape = tensor2.shape.slice(1);
|
|
assertShapesMatchAllowUndefinedSize(tensorElementShape, elementShape, "TensorList shape mismatch: ");
|
|
const tensorList = unstack(tensor2);
|
|
return new TensorList(tensorList, elementShape, dtype);
|
|
}
|
|
function reserve(elementShape, elementDtype, numElements, maxNumElements) {
|
|
return new TensorList([], elementShape, elementDtype, maxNumElements);
|
|
}
|
|
function scatter(tensor2, indices, elementShape, numElements) {
|
|
if (indices.length !== tensor2.shape[0]) {
|
|
throw new Error(`Expected len(indices) == tensor.shape[0], but saw: ${indices.length} vs. ${tensor2.shape[0]}`);
|
|
}
|
|
const maxIndex = Math.max(...indices);
|
|
if (numElements != null && numElements !== -1 && maxIndex >= numElements) {
|
|
throw new Error(`Max index must be < array size (${maxIndex} vs. ${numElements})`);
|
|
}
|
|
const list = new TensorList([], elementShape, tensor2.dtype, numElements);
|
|
const tensors = unstack(tensor2, 0);
|
|
indices.forEach((value, index) => {
|
|
list.setItem(value, tensors[index]);
|
|
});
|
|
return list;
|
|
}
|
|
function split$1(tensor2, length, elementShape) {
|
|
let totalLength = 0;
|
|
const cumulativeLengths = length.map((len) => {
|
|
totalLength += len;
|
|
return totalLength;
|
|
});
|
|
if (totalLength !== tensor2.shape[0]) {
|
|
throw new Error(`Expected sum of lengths to be equal to
|
|
tensor.shape[0], but sum of lengths is
|
|
${totalLength}, and tensor's shape is: ${tensor2.shape}`);
|
|
}
|
|
const shapeWithoutFirstDim = tensor2.shape.slice(1);
|
|
const outputElementShape = mergeElementShape(shapeWithoutFirstDim, elementShape);
|
|
const elementPerRow = totalLength === 0 ? 0 : tensor2.size / totalLength;
|
|
const tensors = tidy(() => {
|
|
const tensors2 = [];
|
|
tensor2 = reshape$2(tensor2, [1, totalLength, elementPerRow]);
|
|
for (let i = 0; i < length.length; ++i) {
|
|
const previousLength = i === 0 ? 0 : cumulativeLengths[i - 1];
|
|
const indices = [0, previousLength, 0];
|
|
const sizes = [1, length[i], elementPerRow];
|
|
tensors2[i] = reshape$2(slice$2(tensor2, indices, sizes), outputElementShape);
|
|
}
|
|
tensor2.dispose();
|
|
return tensors2;
|
|
});
|
|
const list = new TensorList([], elementShape, tensor2.dtype, length.length);
|
|
for (let i = 0; i < tensors.length; i++) {
|
|
list.setItem(i, tensors[i]);
|
|
}
|
|
return list;
|
|
}
|
|
const executeOp$i = async (node, tensorMap, context) => {
|
|
switch (node.op) {
|
|
case "If":
|
|
case "StatelessIf": {
|
|
const thenFunc = getParamValue("thenBranch", node, tensorMap, context);
|
|
const elseFunc = getParamValue("elseBranch", node, tensorMap, context);
|
|
const cond = getParamValue("cond", node, tensorMap, context);
|
|
const args = getParamValue("args", node, tensorMap, context);
|
|
const condValue = await cond.data();
|
|
if (condValue[0]) {
|
|
return context.functionMap[thenFunc].executeFunctionAsync(args, context.tensorArrayMap, context.tensorListMap);
|
|
} else {
|
|
return context.functionMap[elseFunc].executeFunctionAsync(args, context.tensorArrayMap, context.tensorListMap);
|
|
}
|
|
}
|
|
case "While":
|
|
case "StatelessWhile": {
|
|
const bodyFunc = getParamValue("body", node, tensorMap, context);
|
|
const condFunc = getParamValue("cond", node, tensorMap, context);
|
|
const args = getParamValue("args", node, tensorMap, context);
|
|
const condResult = await context.functionMap[condFunc].executeFunctionAsync(args, context.tensorArrayMap, context.tensorListMap);
|
|
const argIds = args.map((tensor2) => tensor2.id);
|
|
let condValue = await condResult[0].data();
|
|
condResult.forEach((tensor2) => {
|
|
if (!tensor2.kept && argIds.indexOf(tensor2.id) === -1) {
|
|
tensor2.dispose();
|
|
}
|
|
});
|
|
let result = args;
|
|
while (condValue[0]) {
|
|
const origResult = result;
|
|
result = await context.functionMap[bodyFunc].executeFunctionAsync(result, context.tensorArrayMap, context.tensorListMap);
|
|
const resultIds = result.map((tensor2) => tensor2.id);
|
|
origResult.forEach((tensor2) => {
|
|
if (!tensor2.kept && argIds.indexOf(tensor2.id) === -1 && resultIds.indexOf(tensor2.id) === -1) {
|
|
tensor2.dispose();
|
|
}
|
|
});
|
|
const condResult2 = await context.functionMap[condFunc].executeFunctionAsync(result, context.tensorArrayMap, context.tensorListMap);
|
|
condValue = await condResult2[0].data();
|
|
condResult2.forEach((tensor2) => {
|
|
if (!tensor2.kept && argIds.indexOf(tensor2.id) === -1 && resultIds.indexOf(tensor2.id) === -1) {
|
|
tensor2.dispose();
|
|
}
|
|
});
|
|
}
|
|
return result;
|
|
}
|
|
case "LoopCond": {
|
|
const pred = getParamValue("pred", node, tensorMap, context);
|
|
return [cloneTensor(pred)];
|
|
}
|
|
case "Switch": {
|
|
const pred = getParamValue("pred", node, tensorMap, context);
|
|
let data = getParamValue("data", node, tensorMap, context);
|
|
if (!data.kept) {
|
|
data = cloneTensor(data);
|
|
}
|
|
return (await pred.data())[0] ? [void 0, data] : [data, void 0];
|
|
}
|
|
case "Merge": {
|
|
const inputName = node.inputNames.find((name) => getTensor(name, tensorMap, context) !== void 0);
|
|
if (inputName) {
|
|
const data = getTensor(inputName, tensorMap, context);
|
|
return [cloneTensor(data)];
|
|
}
|
|
return void 0;
|
|
}
|
|
case "Enter": {
|
|
const frameId = getParamValue("frameName", node, tensorMap, context);
|
|
const data = getParamValue("tensor", node, tensorMap, context);
|
|
context.enterFrame(frameId);
|
|
return [cloneTensor(data)];
|
|
}
|
|
case "Exit": {
|
|
const data = getParamValue("tensor", node, tensorMap, context);
|
|
context.exitFrame();
|
|
return [cloneTensor(data)];
|
|
}
|
|
case "NextIteration": {
|
|
const data = getParamValue("tensor", node, tensorMap, context);
|
|
context.nextIteration();
|
|
return [cloneTensor(data)];
|
|
}
|
|
case "TensorArrayV3": {
|
|
const size = getParamValue("size", node, tensorMap, context);
|
|
const dtype = getParamValue("dtype", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const dynamicSize = getParamValue("dynamicSize", node, tensorMap, context);
|
|
const clearAfterRead = getParamValue("clearAfterRead", node, tensorMap, context);
|
|
const identicalElementShapes = getParamValue("identicalElementShapes", node, tensorMap, context);
|
|
const name = getParamValue("name", node, tensorMap, context);
|
|
const tensorArray = new TensorArray(name, dtype, size, elementShape, identicalElementShapes, dynamicSize, clearAfterRead);
|
|
context.addTensorArray(tensorArray);
|
|
return [tensorArray.idTensor, scalar(1)];
|
|
}
|
|
case "TensorArrayWriteV3": {
|
|
const id = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const index = getParamValue("index", node, tensorMap, context);
|
|
const writeTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const writeTensorArray = context.getTensorArray(id.id);
|
|
writeTensorArray.write(index, writeTensor);
|
|
return [writeTensorArray.idTensor];
|
|
}
|
|
case "TensorArrayReadV3": {
|
|
const readId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const readIndex = getParamValue("index", node, tensorMap, context);
|
|
const readTensorArray = context.getTensorArray(readId.id);
|
|
return [readTensorArray.read(readIndex)];
|
|
}
|
|
case "TensorArrayGatherV3": {
|
|
const gatherId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const gatherIndices = getParamValue("indices", node, tensorMap, context);
|
|
const gatherDtype = getParamValue("dtype", node, tensorMap, context);
|
|
const gatherTensorArray = context.getTensorArray(gatherId.id);
|
|
return [gatherTensorArray.gather(gatherIndices, gatherDtype)];
|
|
}
|
|
case "TensorArrayScatterV3": {
|
|
const scatterId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const scatterIndices = getParamValue("indices", node, tensorMap, context);
|
|
const scatterTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const scatterTensorArray = context.getTensorArray(scatterId.id);
|
|
scatterTensorArray.scatter(scatterIndices, scatterTensor);
|
|
return [scatterTensorArray.idTensor];
|
|
}
|
|
case "TensorArrayConcatV3": {
|
|
const concatId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const concatTensorArray = context.getTensorArray(concatId.id);
|
|
const concatDtype = getParamValue("dtype", node, tensorMap, context);
|
|
return [concatTensorArray.concat(concatDtype)];
|
|
}
|
|
case "TensorArraySplitV3": {
|
|
const splitId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const splitTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const lengths = getParamValue("lengths", node, tensorMap, context);
|
|
const splitTensorArray = context.getTensorArray(splitId.id);
|
|
splitTensorArray.split(lengths, splitTensor);
|
|
return [splitTensorArray.idTensor];
|
|
}
|
|
case "TensorArraySizeV3": {
|
|
const sizeId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const sizeTensorArray = context.getTensorArray(sizeId.id);
|
|
return [scalar(sizeTensorArray.size(), "int32")];
|
|
}
|
|
case "TensorArrayCloseV3": {
|
|
const closeId = getParamValue("tensorArrayId", node, tensorMap, context);
|
|
const closeTensorArray = context.getTensorArray(closeId.id);
|
|
closeTensorArray.clearAndClose();
|
|
return [closeTensorArray.idTensor];
|
|
}
|
|
case "TensorListSetItem": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const index = getParamValue("index", node, tensorMap, context);
|
|
const writeTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(idTensor.id);
|
|
tensorList.setItem(index, writeTensor);
|
|
return [tensorList.idTensor];
|
|
}
|
|
case "TensorListGetItem": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const readIndex = getParamValue("index", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const elementDType = getParamValue("elementDType", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(idTensor.id);
|
|
return [tensorList.getItem(readIndex, elementShape, elementDType)];
|
|
}
|
|
case "TensorListScatterV2":
|
|
case "TensorListScatter": {
|
|
const scatterIndices = getParamValue("indices", node, tensorMap, context);
|
|
const scatterTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const numElements = getParamValue("numElements", node, tensorMap, context);
|
|
const tensorList = scatter(scatterTensor, scatterIndices, elementShape, numElements);
|
|
context.addTensorList(tensorList);
|
|
return [tensorList.idTensor];
|
|
}
|
|
case "TensorListReserve":
|
|
case "EmptyTensorList": {
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const elementDtype = getParamValue("elementDType", node, tensorMap, context);
|
|
let numElementsParam;
|
|
if (node.op === "TensorListReserve") {
|
|
numElementsParam = "numElements";
|
|
} else {
|
|
numElementsParam = "maxNumElements";
|
|
}
|
|
const numElements = getParamValue(numElementsParam, node, tensorMap, context);
|
|
const maxNumElements = node.op === "TensorListReserve" ? -1 : numElements;
|
|
const tensorList = reserve(elementShape, elementDtype, numElements, maxNumElements);
|
|
context.addTensorList(tensorList);
|
|
return [tensorList.idTensor];
|
|
}
|
|
case "TensorListGather": {
|
|
const gatherId = getParamValue("tensorListId", node, tensorMap, context);
|
|
const gatherIndices = getParamValue("indices", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const elementDtype = getParamValue("elementDType", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(gatherId.id);
|
|
return [tensorList.gather(gatherIndices, elementDtype, elementShape)];
|
|
}
|
|
case "TensorListStack": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const elementDtype = getParamValue("elementDType", node, tensorMap, context);
|
|
const numElements = getParamValue("numElements", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(idTensor.id);
|
|
return [tensorList.stack(elementShape, elementDtype, numElements)];
|
|
}
|
|
case "TensorListFromTensor": {
|
|
const tensor2 = getParamValue("tensor", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const elementDtype = getParamValue("elementDType", node, tensorMap, context);
|
|
const tensorList = fromTensor(tensor2, elementShape, elementDtype);
|
|
context.addTensorList(tensorList);
|
|
return [tensorList.idTensor];
|
|
}
|
|
case "TensorListConcat":
|
|
case "TensorListConcatV2": {
|
|
const concatId = getParamValue("tensorListId", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(concatId.id);
|
|
const concatDtype = getParamValue("dtype", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
return [tensorList.concat(concatDtype, elementShape)];
|
|
}
|
|
case "TensorListPushBack": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const writeTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(idTensor.id);
|
|
tensorList.pushBack(writeTensor);
|
|
return [tensorList.idTensor];
|
|
}
|
|
case "TensorListPopBack": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const elementDType = getParamValue("elementDType", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(idTensor.id);
|
|
return [tensorList.popBack(elementShape, elementDType)];
|
|
}
|
|
case "TensorListSplit": {
|
|
const splitTensor = getParamValue("tensor", node, tensorMap, context);
|
|
const elementShape = getParamValue("elementShape", node, tensorMap, context);
|
|
const lengths = getParamValue("lengths", node, tensorMap, context);
|
|
const tensorList = split$1(splitTensor, lengths, elementShape);
|
|
context.addTensorList(tensorList);
|
|
return [tensorList.idTensor];
|
|
}
|
|
case "TensorListLength": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const tensorList = context.getTensorList(idTensor.id);
|
|
return [scalar(tensorList.size(), "int32")];
|
|
}
|
|
case "TensorListResize": {
|
|
const idTensor = getParamValue("tensorListId", node, tensorMap, context);
|
|
const size = getParamValue("size", node, tensorMap, context);
|
|
const srcTensorList = context.getTensorList(idTensor.id);
|
|
const destTensorList = srcTensorList.resize(size);
|
|
context.addTensorList(destTensorList);
|
|
return [destTensorList.idTensor];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
function fusedConvAndDepthWiseParams(node, tensorMap, context) {
|
|
const [extraOp, activationFunc] = getParamValue("fusedOps", node, tensorMap, context);
|
|
const isBiasAdd = extraOp === "biasadd";
|
|
const noBiasAdd = !isBiasAdd;
|
|
const isPrelu = activationFunc === "prelu";
|
|
const isBatchNorm = extraOp === "fusedbatchnorm";
|
|
const numArgs = getParamValue("numArgs", node, tensorMap, context);
|
|
if (isBiasAdd) {
|
|
if (isPrelu && numArgs !== 2) {
|
|
throw new Error("FusedConv2d and DepthwiseConv2d with BiasAdd and Prelu must have two extra arguments: bias and alpha.");
|
|
}
|
|
if (!isPrelu && isBiasAdd && numArgs !== 1) {
|
|
throw new Error("FusedConv2d and DepthwiseConv2d with BiasAdd must have one extra argument: bias.");
|
|
}
|
|
}
|
|
if (isBatchNorm) {
|
|
throw new Error("FusedConv2d and DepthwiseConv2d with FusedBatchNorm is not supported");
|
|
}
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getPadding(node, tensorMap, context);
|
|
const dataFormat = getParamValue("dataFormat", node, tensorMap, context).toUpperCase();
|
|
const dilations = getParamValue("dilations", node, tensorMap, context);
|
|
let [biasArg, preluArg] = getParamValue("args", node, tensorMap, context);
|
|
if (noBiasAdd) {
|
|
preluArg = biasArg;
|
|
biasArg = void 0;
|
|
}
|
|
const leakyreluAlpha = getParamValue("leakyreluAlpha", node, tensorMap, context);
|
|
return {
|
|
stride,
|
|
pad: pad2,
|
|
dataFormat,
|
|
dilations,
|
|
biasArg,
|
|
preluArg,
|
|
activationFunc,
|
|
leakyreluAlpha
|
|
};
|
|
}
|
|
const executeOp$h = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Conv1D": {
|
|
const stride = getParamValue("stride", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const dataFormat = getParamValue("dataFormat", node, tensorMap, context).toUpperCase();
|
|
const dilation = getParamValue("dilation", node, tensorMap, context);
|
|
return [ops.conv1d(getParamValue("x", node, tensorMap, context), getParamValue("filter", node, tensorMap, context), stride, pad2, dataFormat, dilation)];
|
|
}
|
|
case "Conv2D": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getPadding(node, tensorMap, context);
|
|
const dataFormat = getParamValue("dataFormat", node, tensorMap, context).toUpperCase();
|
|
const dilations = getParamValue("dilations", node, tensorMap, context);
|
|
return [ops.conv2d(getParamValue("x", node, tensorMap, context), getParamValue("filter", node, tensorMap, context), [stride[1], stride[2]], pad2, dataFormat, [dilations[1], dilations[2]])];
|
|
}
|
|
case "_FusedConv2D": {
|
|
const { stride, pad: pad2, dataFormat, dilations, biasArg, preluArg, activationFunc, leakyreluAlpha } = fusedConvAndDepthWiseParams(node, tensorMap, context);
|
|
return [ops.fused.conv2d({
|
|
x: getParamValue("x", node, tensorMap, context),
|
|
filter: getParamValue("filter", node, tensorMap, context),
|
|
strides: [stride[1], stride[2]],
|
|
pad: pad2,
|
|
dataFormat,
|
|
dilations: [dilations[1], dilations[2]],
|
|
bias: biasArg,
|
|
activation: activationFunc,
|
|
preluActivationWeights: preluArg,
|
|
leakyreluAlpha
|
|
})];
|
|
}
|
|
case "FusedDepthwiseConv2dNative": {
|
|
const { stride, pad: pad2, dataFormat, dilations, biasArg, preluArg, activationFunc, leakyreluAlpha } = fusedConvAndDepthWiseParams(node, tensorMap, context);
|
|
return [ops.fused.depthwiseConv2d({
|
|
x: getParamValue("x", node, tensorMap, context),
|
|
filter: getParamValue("filter", node, tensorMap, context),
|
|
strides: [stride[1], stride[2]],
|
|
pad: pad2,
|
|
dataFormat,
|
|
dilations: [dilations[1], dilations[2]],
|
|
bias: biasArg,
|
|
activation: activationFunc,
|
|
preluActivationWeights: preluArg,
|
|
leakyreluAlpha
|
|
})];
|
|
}
|
|
case "Conv2DBackpropInput":
|
|
case "Conv2dTranspose": {
|
|
const shape = getParamValue("outputShape", node, tensorMap, context);
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getPadding(node, tensorMap, context);
|
|
return [ops.conv2dTranspose(getParamValue("x", node, tensorMap, context), getParamValue("filter", node, tensorMap, context), shape, [stride[1], stride[2]], pad2)];
|
|
}
|
|
case "DepthwiseConv2dNative":
|
|
case "DepthwiseConv2d": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getPadding(node, tensorMap, context);
|
|
const dilations = getParamValue("dilations", node, tensorMap, context);
|
|
const dataFormat = getParamValue("dataFormat", node, tensorMap, context).toUpperCase();
|
|
return [ops.depthwiseConv2d(getParamValue("input", node, tensorMap, context), getParamValue("filter", node, tensorMap, context), [stride[1], stride[2]], pad2, dataFormat, [dilations[1], dilations[2]])];
|
|
}
|
|
case "Conv3D": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const dataFormat = getParamValue("dataFormat", node, tensorMap, context).toUpperCase();
|
|
const dilations = getParamValue("dilations", node, tensorMap, context);
|
|
return [ops.conv3d(getParamValue("x", node, tensorMap, context), getParamValue("filter", node, tensorMap, context), [stride[1], stride[2], stride[3]], pad2, dataFormat, [dilations[1], dilations[2], dilations[3]])];
|
|
}
|
|
case "AvgPool": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const kernelSize = getParamValue("kernelSize", node, tensorMap, context);
|
|
return [ops.avgPool(getParamValue("x", node, tensorMap, context), [kernelSize[1], kernelSize[2]], [stride[1], stride[2]], pad2)];
|
|
}
|
|
case "MaxPool": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const kernelSize = getParamValue("kernelSize", node, tensorMap, context);
|
|
return [ops.maxPool(getParamValue("x", node, tensorMap, context), [kernelSize[1], kernelSize[2]], [stride[1], stride[2]], pad2)];
|
|
}
|
|
case "MaxPoolWithArgmax": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const kernelSize = getParamValue("kernelSize", node, tensorMap, context);
|
|
const includeBatchInIndex = getParamValue("includeBatchInIndex", node, tensorMap, context);
|
|
const { result, indexes } = ops.maxPoolWithArgmax(getParamValue("x", node, tensorMap, context), [kernelSize[1], kernelSize[2]], [stride[1], stride[2]], pad2, includeBatchInIndex);
|
|
return [result, indexes];
|
|
}
|
|
case "AvgPool3D": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const kernelSize = getParamValue("kernelSize", node, tensorMap, context);
|
|
return [ops.avgPool3d(getParamValue("x", node, tensorMap, context), [kernelSize[1], kernelSize[2], kernelSize[3]], [stride[1], stride[2], stride[3]], pad2)];
|
|
}
|
|
case "MaxPool3D": {
|
|
const stride = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const kernelSize = getParamValue("kernelSize", node, tensorMap, context);
|
|
return [ops.maxPool3d(getParamValue("x", node, tensorMap, context), [kernelSize[1], kernelSize[2], kernelSize[3]], [stride[1], stride[2], stride[3]], pad2)];
|
|
}
|
|
case "Dilation2D": {
|
|
const strides = getParamValue("strides", node, tensorMap, context);
|
|
const pad2 = getParamValue("pad", node, tensorMap, context);
|
|
const dilations = getParamValue("dilations", node, tensorMap, context);
|
|
const strideHeight = strides[1];
|
|
const strideWidth = strides[2];
|
|
const dilationHeight = dilations[1];
|
|
const dilationWidth = dilations[2];
|
|
return [ops.dilation2d(
|
|
getParamValue("x", node, tensorMap, context),
|
|
getParamValue("filter", node, tensorMap, context),
|
|
[strideHeight, strideWidth],
|
|
pad2,
|
|
[dilationHeight, dilationWidth],
|
|
"NHWC"
|
|
/* dataFormat */
|
|
)];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$g = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Fill": {
|
|
const shape = getParamValue("shape", node, tensorMap, context);
|
|
const dtype = getParamValue("dtype", node, tensorMap, context);
|
|
const value = getParamValue("value", node, tensorMap, context);
|
|
return [ops.fill(shape, value, dtype)];
|
|
}
|
|
case "LinSpace": {
|
|
const start = getParamValue("start", node, tensorMap, context);
|
|
const stop = getParamValue("stop", node, tensorMap, context);
|
|
const num = getParamValue("num", node, tensorMap, context);
|
|
return [ops.linspace(start, stop, num)];
|
|
}
|
|
case "Multinomial": {
|
|
const logits = getParamValue("logits", node, tensorMap, context);
|
|
const numSamples = getParamValue("numSamples", node, tensorMap, context);
|
|
const seed = getParamValue("seed", node, tensorMap, context);
|
|
return [ops.multinomial(logits, numSamples, seed)];
|
|
}
|
|
case "OneHot": {
|
|
const indices = getParamValue("indices", node, tensorMap, context);
|
|
const depth = getParamValue("depth", node, tensorMap, context);
|
|
const onValue = getParamValue("onValue", node, tensorMap, context);
|
|
const offValue = getParamValue("offValue", node, tensorMap, context);
|
|
const dtype = getParamValue("dtype", node, tensorMap, context);
|
|
return [ops.oneHot(indices, depth, onValue, offValue, dtype)];
|
|
}
|
|
case "Ones": {
|
|
return [ops.ones(getParamValue("shape", node, tensorMap, context), getParamValue("dtype", node, tensorMap, context))];
|
|
}
|
|
case "OnesLike": {
|
|
return [ops.onesLike(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "RandomStandardNormal": {
|
|
return [ops.randomStandardNormal(getParamValue("shape", node, tensorMap, context), getParamValue("dtype", node, tensorMap, context), getParamValue("seed", node, tensorMap, context))];
|
|
}
|
|
case "RandomUniform": {
|
|
return [ops.randomUniform(
|
|
// tslint:disable-next-line:no-any
|
|
getParamValue("shape", node, tensorMap, context),
|
|
getParamValue("minval", node, tensorMap, context),
|
|
getParamValue("maxval", node, tensorMap, context),
|
|
getParamValue("dtype", node, tensorMap, context)
|
|
)];
|
|
}
|
|
case "RandomUniformInt": {
|
|
return [ops.randomUniformInt(getParamValue("shape", node, tensorMap, context), getParamValue("minval", node, tensorMap, context), getParamValue("maxval", node, tensorMap, context), getParamValue("seed", node, tensorMap, context))];
|
|
}
|
|
case "Range": {
|
|
const start = getParamValue("start", node, tensorMap, context);
|
|
const stop = getParamValue("stop", node, tensorMap, context);
|
|
const step2 = getParamValue("step", node, tensorMap, context);
|
|
return [ops.range(start, stop, step2, getParamValue("dtype", node, tensorMap, context))];
|
|
}
|
|
case "TruncatedNormal": {
|
|
const shape = getParamValue("shape", node, tensorMap, context);
|
|
const mean2 = getParamValue("mean", node, tensorMap, context);
|
|
const stdDev = getParamValue("stdDev", node, tensorMap, context);
|
|
const seed = getParamValue("seed", node, tensorMap, context);
|
|
return [ops.truncatedNormal(shape, mean2, stdDev, getParamValue("dtype", node, tensorMap, context), seed)];
|
|
}
|
|
case "Zeros": {
|
|
return [ops.zeros(getParamValue("shape", node, tensorMap, context), getParamValue("dtype", node, tensorMap, context))];
|
|
}
|
|
case "ZerosLike": {
|
|
return [ops.zerosLike(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
function nmsParams(node, tensorMap, context) {
|
|
const boxes = getParamValue("boxes", node, tensorMap, context);
|
|
const scores = getParamValue("scores", node, tensorMap, context);
|
|
const maxOutputSize = getParamValue("maxOutputSize", node, tensorMap, context);
|
|
const iouThreshold = getParamValue("iouThreshold", node, tensorMap, context);
|
|
const scoreThreshold = getParamValue("scoreThreshold", node, tensorMap, context);
|
|
const softNmsSigma = getParamValue("softNmsSigma", node, tensorMap, context);
|
|
return {
|
|
boxes,
|
|
scores,
|
|
maxOutputSize,
|
|
iouThreshold,
|
|
scoreThreshold,
|
|
softNmsSigma
|
|
};
|
|
}
|
|
const executeOp$f = async (node, tensorMap, context, resourceManager, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "NonMaxSuppressionV5": {
|
|
const { boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma } = nmsParams(node, tensorMap, context);
|
|
const result = await ops.image.nonMaxSuppressionWithScoreAsync(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma);
|
|
return [result.selectedIndices, result.selectedScores];
|
|
}
|
|
case "NonMaxSuppressionV4": {
|
|
const { boxes, scores, maxOutputSize, iouThreshold, scoreThreshold } = nmsParams(node, tensorMap, context);
|
|
const padToMaxOutputSize = getParamValue("padToMaxOutputSize", node, tensorMap, context);
|
|
const result = await ops.image.nonMaxSuppressionPaddedAsync(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold, padToMaxOutputSize);
|
|
return [result.selectedIndices, result.validOutputs];
|
|
}
|
|
case "NonMaxSuppressionV3":
|
|
case "NonMaxSuppressionV2": {
|
|
const { boxes, scores, maxOutputSize, iouThreshold, scoreThreshold } = nmsParams(node, tensorMap, context);
|
|
return [await ops.image.nonMaxSuppressionAsync(boxes, scores, maxOutputSize, iouThreshold, scoreThreshold)];
|
|
}
|
|
case "Where": {
|
|
const condition = ops.cast(getParamValue("condition", node, tensorMap, context), "bool");
|
|
const result = [await ops.whereAsync(condition)];
|
|
condition.dispose();
|
|
return result;
|
|
}
|
|
case "ListDiff": {
|
|
return ops.setdiff1dAsync(getParamValue("x", node, tensorMap, context), getParamValue("y", node, tensorMap, context));
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$e = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "LowerBound": {
|
|
const sortedSequence = getParamValue("sortedSequence", node, tensorMap, context);
|
|
const values = getParamValue("values", node, tensorMap, context);
|
|
return [ops.lowerBound(sortedSequence, values)];
|
|
}
|
|
case "TopKV2": {
|
|
const x = getParamValue("x", node, tensorMap, context);
|
|
const k = getParamValue("k", node, tensorMap, context);
|
|
const sorted = getParamValue("sorted", node, tensorMap, context);
|
|
const result = ops.topk(x, k, sorted);
|
|
return [result.values, result.indices];
|
|
}
|
|
case "UpperBound": {
|
|
const sortedSequence = getParamValue("sortedSequence", node, tensorMap, context);
|
|
const values = getParamValue("values", node, tensorMap, context);
|
|
return [ops.upperBound(sortedSequence, values)];
|
|
}
|
|
case "Unique": {
|
|
const x = getParamValue("x", node, tensorMap, context);
|
|
const result = ops.unique(x);
|
|
return [result.values, result.indices];
|
|
}
|
|
case "UniqueV2": {
|
|
const x = getParamValue("x", node, tensorMap, context);
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const result = ops.unique(x, axis);
|
|
return [result.values, result.indices];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$d = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Const": {
|
|
return tensorMap[node.name];
|
|
}
|
|
case "PlaceholderWithDefault":
|
|
const def = getParamValue("default", node, tensorMap, context);
|
|
return [getTensor(node.name, tensorMap, context) || def];
|
|
case "Placeholder":
|
|
return [getTensor(node.name, tensorMap, context)];
|
|
case "Identity":
|
|
case "StopGradient":
|
|
case "FakeQuantWithMinMaxVars": {
|
|
const data2 = getParamValue("x", node, tensorMap, context);
|
|
return [cloneTensor(data2)];
|
|
}
|
|
case "IdentityN":
|
|
return getParamValue("x", node, tensorMap, context).map((t2) => cloneTensor(t2));
|
|
case "Snapshot":
|
|
const snapshot = getParamValue("x", node, tensorMap, context);
|
|
return [cloneTensor(snapshot)];
|
|
case "Shape":
|
|
return [ops.tensor1d(getParamValue("x", node, tensorMap, context).shape, "int32")];
|
|
case "ShapeN":
|
|
return getParamValue("x", node, tensorMap, context).map((t2) => ops.tensor1d(t2.shape));
|
|
case "Size":
|
|
return [ops.scalar(getParamValue("x", node, tensorMap, context).size, "int32")];
|
|
case "Rank":
|
|
return [ops.scalar(getParamValue("x", node, tensorMap, context).rank, "int32")];
|
|
case "NoOp":
|
|
return [ops.scalar(1)];
|
|
case "Print":
|
|
const input = getParamValue("x", node, tensorMap, context);
|
|
const data = getParamValue("data", node, tensorMap, context);
|
|
const message = getParamValue("message", node, tensorMap, context);
|
|
const summarize = getParamValue("summarize", node, tensorMap, context);
|
|
console.warn("The graph has a tf.print() operation,usually used for debugging, which slows down performance.");
|
|
console.log(message);
|
|
for (let i = 0; i < data.length; i++) {
|
|
console.log(Array.prototype.slice.call(data[i].dataSync()).slice(0, summarize));
|
|
}
|
|
return [input];
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
class HashTable {
|
|
get id() {
|
|
return this.handle.id;
|
|
}
|
|
/**
|
|
* Constructor of HashTable. Creates a hash table.
|
|
*
|
|
* @param keyDType `dtype` of the table keys.
|
|
* @param valueDType `dtype` of the table values.
|
|
*/
|
|
constructor(keyDType, valueDType) {
|
|
this.keyDType = keyDType;
|
|
this.valueDType = valueDType;
|
|
this.handle = scalar(0);
|
|
this.tensorMap = /* @__PURE__ */ new Map();
|
|
keep(this.handle);
|
|
}
|
|
/**
|
|
* Dispose the tensors and handle and clear the hashtable.
|
|
*/
|
|
clearAndClose() {
|
|
this.tensorMap.forEach((value) => value.dispose());
|
|
this.tensorMap.clear();
|
|
this.handle.dispose();
|
|
}
|
|
/**
|
|
* The number of items in the hash table.
|
|
*/
|
|
size() {
|
|
return this.tensorMap.size;
|
|
}
|
|
/**
|
|
* The number of items in the hash table as a rank-0 tensor.
|
|
*/
|
|
tensorSize() {
|
|
return scalar(this.size(), "int32");
|
|
}
|
|
/**
|
|
* Replaces the contents of the table with the specified keys and values.
|
|
* @param keys Keys to store in the hashtable.
|
|
* @param values Values to store in the hashtable.
|
|
*/
|
|
async import(keys, values) {
|
|
this.checkKeyAndValueTensor(keys, values);
|
|
const $keys = await keys.data();
|
|
this.tensorMap.forEach((value) => value.dispose());
|
|
this.tensorMap.clear();
|
|
return tidy(() => {
|
|
const $values = unstack(values);
|
|
const keysLength = $keys.length;
|
|
const valuesLength = $values.length;
|
|
assert(keysLength === valuesLength, () => `The number of elements doesn't match, keys has ${keysLength} elements, the values has ${valuesLength} elements.`);
|
|
for (let i = 0; i < keysLength; i++) {
|
|
const key = $keys[i];
|
|
const value = $values[i];
|
|
keep(value);
|
|
this.tensorMap.set(key, value);
|
|
}
|
|
return this.handle;
|
|
});
|
|
}
|
|
/**
|
|
* Looks up keys in a hash table, outputs the corresponding values.
|
|
*
|
|
* Performs batch lookups, for every element in the key tensor, `find`
|
|
* stacks the corresponding value into the return tensor.
|
|
*
|
|
* If an element is not present in the table, the given `defaultValue` is
|
|
* used.
|
|
*
|
|
* @param keys Keys to look up. Must have the same type as the keys of the
|
|
* table.
|
|
* @param defaultValue The scalar `defaultValue` is the value output for keys
|
|
* not present in the table. It must also be of the same type as the
|
|
* table values.
|
|
*/
|
|
async find(keys, defaultValue) {
|
|
this.checkKeyAndValueTensor(keys, defaultValue);
|
|
const $keys = await keys.data();
|
|
return tidy(() => {
|
|
const result = [];
|
|
for (let i = 0; i < $keys.length; i++) {
|
|
const key = $keys[i];
|
|
const value = this.findWithDefault(key, defaultValue);
|
|
result.push(value);
|
|
}
|
|
return stack(result);
|
|
});
|
|
}
|
|
// tslint:disable-next-line: no-any
|
|
findWithDefault(key, defaultValue) {
|
|
const result = this.tensorMap.get(key);
|
|
return result != null ? result : defaultValue;
|
|
}
|
|
checkKeyAndValueTensor(key, value) {
|
|
if (key.dtype !== this.keyDType) {
|
|
throw new Error(`Expect key dtype ${this.keyDType}, but got ${key.dtype}`);
|
|
}
|
|
if (value.dtype !== this.valueDType) {
|
|
throw new Error(`Expect value dtype ${this.valueDType}, but got ${value.dtype}`);
|
|
}
|
|
}
|
|
}
|
|
const executeOp$c = async (node, tensorMap, context, resourceManager) => {
|
|
switch (node.op) {
|
|
case "HashTable":
|
|
case "HashTableV2": {
|
|
const existingTableHandle = resourceManager.getHashTableHandleByName(node.name);
|
|
if (existingTableHandle != null) {
|
|
return [existingTableHandle];
|
|
} else {
|
|
const keyDType = getParamValue("keyDType", node, tensorMap, context);
|
|
const valueDType = getParamValue("valueDType", node, tensorMap, context);
|
|
const hashTable2 = new HashTable(keyDType, valueDType);
|
|
resourceManager.addHashTable(node.name, hashTable2);
|
|
return [hashTable2.handle];
|
|
}
|
|
}
|
|
case "InitializeTable":
|
|
case "InitializeTableV2":
|
|
case "LookupTableImport":
|
|
case "LookupTableImportV2": {
|
|
const handle = getParamValue("tableHandle", node, tensorMap, context, resourceManager);
|
|
const keys = getParamValue("keys", node, tensorMap, context);
|
|
const values = getParamValue("values", node, tensorMap, context);
|
|
const hashTable2 = resourceManager.getHashTableById(handle.id);
|
|
return [await hashTable2.import(keys, values)];
|
|
}
|
|
case "LookupTableFind":
|
|
case "LookupTableFindV2": {
|
|
const handle = getParamValue("tableHandle", node, tensorMap, context, resourceManager);
|
|
const keys = getParamValue("keys", node, tensorMap, context);
|
|
const defaultValue = getParamValue("defaultValue", node, tensorMap, context);
|
|
const hashTable2 = resourceManager.getHashTableById(handle.id);
|
|
return [await hashTable2.find(keys, defaultValue)];
|
|
}
|
|
case "LookupTableSize":
|
|
case "LookupTableSizeV2": {
|
|
const handle = getParamValue("tableHandle", node, tensorMap, context, resourceManager);
|
|
const hashTable2 = resourceManager.getHashTableById(handle.id);
|
|
return [hashTable2.tensorSize()];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$b = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "ResizeBilinear": {
|
|
const images = getParamValue("images", node, tensorMap, context);
|
|
const size = getParamValue("size", node, tensorMap, context);
|
|
const alignCorners = getParamValue("alignCorners", node, tensorMap, context);
|
|
const halfPixelCenters = getParamValue("halfPixelCenters", node, tensorMap, context);
|
|
return [ops.image.resizeBilinear(images, [size[0], size[1]], alignCorners, halfPixelCenters)];
|
|
}
|
|
case "ResizeNearestNeighbor": {
|
|
const images = getParamValue("images", node, tensorMap, context);
|
|
const size = getParamValue("size", node, tensorMap, context);
|
|
const alignCorners = getParamValue("alignCorners", node, tensorMap, context);
|
|
const halfPixelCenters = getParamValue("halfPixelCenters", node, tensorMap, context);
|
|
return [ops.image.resizeNearestNeighbor(images, [size[0], size[1]], alignCorners, halfPixelCenters)];
|
|
}
|
|
case "CropAndResize": {
|
|
const image2 = getParamValue("image", node, tensorMap, context);
|
|
const boxes = getParamValue("boxes", node, tensorMap, context);
|
|
const boxInd = getParamValue("boxInd", node, tensorMap, context);
|
|
const cropSize = getParamValue("cropSize", node, tensorMap, context);
|
|
const method = getParamValue("method", node, tensorMap, context);
|
|
const extrapolationValue = getParamValue("extrapolationValue", node, tensorMap, context);
|
|
return [ops.image.cropAndResize(image2, boxes, boxInd, cropSize, method, extrapolationValue)];
|
|
}
|
|
case "ImageProjectiveTransformV3": {
|
|
const images = getParamValue("images", node, tensorMap, context);
|
|
const transforms = getParamValue("transforms", node, tensorMap, context);
|
|
const outputShape = getParamValue("outputShape", node, tensorMap, context);
|
|
const fillValue = getParamValue("fillValue", node, tensorMap, context);
|
|
const interpolation = getParamValue("interpolation", node, tensorMap, context);
|
|
const fillMode = getParamValue("fillMode", node, tensorMap, context);
|
|
return [ops.image.transform(images, transforms, interpolation.toLowerCase(), fillMode.toLowerCase(), fillValue, outputShape)];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$a = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Equal": {
|
|
return [ops.equal(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "NotEqual": {
|
|
return [ops.notEqual(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Greater": {
|
|
return [ops.greater(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "GreaterEqual": {
|
|
return [ops.greaterEqual(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Less": {
|
|
return [ops.less(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "LessEqual": {
|
|
return [ops.lessEqual(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "LogicalAnd": {
|
|
return [ops.logicalAnd(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "LogicalNot": {
|
|
return [ops.logicalNot(getParamValue("a", node, tensorMap, context))];
|
|
}
|
|
case "LogicalOr": {
|
|
return [ops.logicalOr(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "Select":
|
|
case "SelectV2": {
|
|
return [ops.where(getParamValue("condition", node, tensorMap, context), getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
case "BitwiseAnd": {
|
|
return [ops.bitwiseAnd(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$9 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "BatchMatMul":
|
|
case "BatchMatMulV2":
|
|
case "MatMul":
|
|
return [ops.matMul(getParamValue("a", node, tensorMap, context), getParamValue("b", node, tensorMap, context), getParamValue("transposeA", node, tensorMap, context), getParamValue("transposeB", node, tensorMap, context))];
|
|
case "Einsum":
|
|
return [ops.einsum(getParamValue("equation", node, tensorMap, context), ...getParamValue("tensors", node, tensorMap, context))];
|
|
case "Transpose":
|
|
return [ops.transpose(getParamValue("x", node, tensorMap, context), getParamValue("perm", node, tensorMap, context))];
|
|
case "_FusedMatMul":
|
|
const [extraOp, activationFunc] = getParamValue("fusedOps", node, tensorMap, context);
|
|
const isBiasAdd = extraOp === "biasadd";
|
|
const isPrelu = activationFunc === "prelu";
|
|
const numArgs = getParamValue("numArgs", node, tensorMap, context);
|
|
const leakyreluAlpha = getParamValue("leakyreluAlpha", node, tensorMap, context);
|
|
if (isBiasAdd) {
|
|
if (isPrelu && numArgs !== 2) {
|
|
throw new Error("Fused MatMul with BiasAdd and Prelu must have two extra arguments: bias and alpha.");
|
|
}
|
|
if (!isPrelu && numArgs !== 1) {
|
|
throw new Error("Fused MatMul with BiasAdd must have one extra argument: bias.");
|
|
}
|
|
}
|
|
const [biasArg, preluArg] = getParamValue("args", node, tensorMap, context);
|
|
return [ops.fused.matMul({
|
|
a: getParamValue("a", node, tensorMap, context),
|
|
b: getParamValue("b", node, tensorMap, context),
|
|
transposeA: getParamValue("transposeA", node, tensorMap, context),
|
|
transposeB: getParamValue("transposeB", node, tensorMap, context),
|
|
bias: biasArg,
|
|
activation: activationFunc,
|
|
preluActivationWeights: preluArg,
|
|
leakyreluAlpha
|
|
})];
|
|
case "MatrixBandPart":
|
|
return [ops.linalg.bandPart(getParamValue("a", node, tensorMap, context), getParamValue("numLower", node, tensorMap, context), getParamValue("numUpper", node, tensorMap, context))];
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$8 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "EuclideanNorm":
|
|
return [ops.euclideanNorm(getParamValue("x", node, tensorMap, context), getParamValue("axis", node, tensorMap, context), getParamValue("keepDims", node, tensorMap, context))];
|
|
case "FusedBatchNorm":
|
|
case "FusedBatchNormV2": {
|
|
return [ops.batchNorm(getParamValue("x", node, tensorMap, context), getParamValue("mean", node, tensorMap, context), getParamValue("variance", node, tensorMap, context), getParamValue("offset", node, tensorMap, context), getParamValue("scale", node, tensorMap, context), getParamValue("epsilon", node, tensorMap, context))];
|
|
}
|
|
case "FusedBatchNormV3": {
|
|
return [ops.batchNorm(getParamValue("x", node, tensorMap, context), getParamValue("mean", node, tensorMap, context), getParamValue("variance", node, tensorMap, context), getParamValue("offset", node, tensorMap, context), getParamValue("scale", node, tensorMap, context), getParamValue("epsilon", node, tensorMap, context))];
|
|
}
|
|
case "LRN": {
|
|
return [ops.localResponseNormalization(getParamValue("x", node, tensorMap, context), getParamValue("radius", node, tensorMap, context), getParamValue("bias", node, tensorMap, context), getParamValue("alpha", node, tensorMap, context), getParamValue("beta", node, tensorMap, context))];
|
|
}
|
|
case "Softmax": {
|
|
return [ops.softmax(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "LogSoftmax": {
|
|
return [ops.logSoftmax(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$7 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "RaggedGather": {
|
|
const { outputNestedSplits, outputDenseValues } = ops.raggedGather(getParamValue("paramsNestedSplits", node, tensorMap, context), getParamValue("paramsDenseValues", node, tensorMap, context), getParamValue("indices", node, tensorMap, context), getParamValue("outputRaggedRank", node, tensorMap, context));
|
|
return outputNestedSplits.concat(outputDenseValues);
|
|
}
|
|
case "RaggedRange": {
|
|
const { rtNestedSplits, rtDenseValues } = ops.raggedRange(getParamValue("starts", node, tensorMap, context), getParamValue("limits", node, tensorMap, context), getParamValue("splits", node, tensorMap, context));
|
|
return [rtNestedSplits, rtDenseValues];
|
|
}
|
|
case "RaggedTensorToTensor": {
|
|
return [ops.raggedTensorToTensor(getParamValue("shape", node, tensorMap, context), getParamValue("values", node, tensorMap, context), getParamValue("defaultValue", node, tensorMap, context), getParamValue("rowPartitionTensors", node, tensorMap, context), getParamValue("rowPartitionTypes", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$6 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Max": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.max(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "Mean": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.mean(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "Min": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.min(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "Sum": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.sum(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "All": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.all(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "Any": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.any(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "ArgMax": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
return [ops.argMax(getParamValue("x", node, tensorMap, context), axis)];
|
|
}
|
|
case "ArgMin": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
return [ops.argMin(getParamValue("x", node, tensorMap, context), axis)];
|
|
}
|
|
case "Prod": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const keepDims = getParamValue("keepDims", node, tensorMap, context);
|
|
return [ops.prod(getParamValue("x", node, tensorMap, context), axis, keepDims)];
|
|
}
|
|
case "Cumprod": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const exclusive = getParamValue("exclusive", node, tensorMap, context);
|
|
const reverse2 = getParamValue("reverse", node, tensorMap, context);
|
|
return [ops.cumprod(getParamValue("x", node, tensorMap, context), axis, exclusive, reverse2)];
|
|
}
|
|
case "Cumsum": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const exclusive = getParamValue("exclusive", node, tensorMap, context);
|
|
const reverse2 = getParamValue("reverse", node, tensorMap, context);
|
|
return [ops.cumsum(getParamValue("x", node, tensorMap, context), axis, exclusive, reverse2)];
|
|
}
|
|
case "Bincount":
|
|
const x = getParamValue("x", node, tensorMap, context);
|
|
const weights = getParamValue("weights", node, tensorMap, context);
|
|
const size = getParamValue("size", node, tensorMap, context);
|
|
return [ops.bincount(x, weights, size)];
|
|
case "DenseBincount": {
|
|
const x2 = getParamValue("x", node, tensorMap, context);
|
|
const weights2 = getParamValue("weights", node, tensorMap, context);
|
|
const size2 = getParamValue("size", node, tensorMap, context);
|
|
const binaryOutput = getParamValue("binaryOutput", node, tensorMap, context);
|
|
return [ops.denseBincount(x2, weights2, size2, binaryOutput)];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$5 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "ConcatV2":
|
|
case "Concat": {
|
|
const n = getParamValue("n", node, tensorMap, context);
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
let inputs = getParamValue("tensors", node, tensorMap, context);
|
|
inputs = inputs.slice(0, n);
|
|
return [ops.concat(inputs, axis)];
|
|
}
|
|
case "Gather": {
|
|
const input = getParamValue("x", node, tensorMap, context);
|
|
const indices = getParamValue("indices", node, tensorMap, context);
|
|
return [ops.gather(input, ops.cast(indices, "int32"), 0)];
|
|
}
|
|
case "GatherV2": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const batchDims = getParamValue("batchDims", node, tensorMap, context);
|
|
const input = getParamValue("x", node, tensorMap, context);
|
|
const indices = getParamValue("indices", node, tensorMap, context);
|
|
return [ops.gather(input, ops.cast(indices, "int32"), axis, batchDims)];
|
|
}
|
|
case "Reverse": {
|
|
const dims = getParamValue("dims", node, tensorMap, context);
|
|
const axis = [];
|
|
for (let i = 0; i < dims.length; i++) {
|
|
if (dims[i]) {
|
|
axis.push(i);
|
|
}
|
|
}
|
|
const input = getParamValue("x", node, tensorMap, context);
|
|
return [ops.reverse(input, axis)];
|
|
}
|
|
case "ReverseV2": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const input = getParamValue("x", node, tensorMap, context);
|
|
return [ops.reverse(input, axis)];
|
|
}
|
|
case "Slice": {
|
|
const begin = getParamValue("begin", node, tensorMap, context);
|
|
const size = getParamValue("size", node, tensorMap, context);
|
|
return [ops.slice(getParamValue("x", node, tensorMap, context), begin, size)];
|
|
}
|
|
case "StridedSlice": {
|
|
const begin = getParamValue("begin", node, tensorMap, context);
|
|
const end = getParamValue("end", node, tensorMap, context);
|
|
const strides = getParamValue("strides", node, tensorMap, context);
|
|
const beginMask = getParamValue("beginMask", node, tensorMap, context);
|
|
const endMask = getParamValue("endMask", node, tensorMap, context);
|
|
const ellipsisMask = getParamValue("ellipsisMask", node, tensorMap, context);
|
|
const newAxisMask = getParamValue("newAxisMask", node, tensorMap, context);
|
|
const shrinkAxisMask = getParamValue("shrinkAxisMask", node, tensorMap, context);
|
|
const tensor2 = getParamValue("x", node, tensorMap, context);
|
|
return [ops.stridedSlice(tensor2, begin, end, strides, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask)];
|
|
}
|
|
case "Pack": {
|
|
return tidy(() => {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const tensors = getParamValue("tensors", node, tensorMap, context);
|
|
const shape = tensors[0].shape;
|
|
const squeezedShape = ops.squeeze(tensors[0]).shape;
|
|
const mapped = tensors.map((tensor2) => {
|
|
const sameShape = arraysEqual(tensor2.shape, shape);
|
|
if (!sameShape && !arraysEqual(ops.squeeze(tensor2).shape, squeezedShape)) {
|
|
throw new Error("the input tensors shape does not match");
|
|
}
|
|
return sameShape ? tensor2 : ops.reshape(tensor2, shape);
|
|
});
|
|
return [ops.stack(mapped, axis)];
|
|
});
|
|
}
|
|
case "Unpack": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const tensor2 = getParamValue("tensor", node, tensorMap, context);
|
|
return ops.unstack(tensor2, axis);
|
|
}
|
|
case "Tile": {
|
|
const reps = getParamValue("reps", node, tensorMap, context);
|
|
return [ops.tile(getParamValue("x", node, tensorMap, context), reps)];
|
|
}
|
|
case "Split":
|
|
case "SplitV": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
const numOrSizeSplits = getParamValue("numOrSizeSplits", node, tensorMap, context);
|
|
const tensor2 = getParamValue("x", node, tensorMap, context);
|
|
return ops.split(tensor2, numOrSizeSplits, axis);
|
|
}
|
|
case "ScatterNd": {
|
|
const indices = getParamValue("indices", node, tensorMap, context);
|
|
const values = getParamValue("values", node, tensorMap, context);
|
|
const shape = getParamValue("shape", node, tensorMap, context);
|
|
return [ops.scatterND(indices, values, shape)];
|
|
}
|
|
case "GatherNd": {
|
|
const x = getParamValue("x", node, tensorMap, context);
|
|
const indices = getParamValue("indices", node, tensorMap, context);
|
|
return [ops.gatherND(x, indices)];
|
|
}
|
|
case "SparseToDense": {
|
|
const indices = getParamValue("sparseIndices", node, tensorMap, context);
|
|
const shape = getParamValue("outputShape", node, tensorMap, context);
|
|
const sparseValues = getParamValue("sparseValues", node, tensorMap, context);
|
|
const defaultValue = getParamValue("defaultValue", node, tensorMap, context);
|
|
return [ops.sparseToDense(indices, sparseValues, shape, sparseValues.dtype === defaultValue.dtype ? defaultValue : ops.cast(defaultValue, sparseValues.dtype))];
|
|
}
|
|
case "TensorScatterUpdate": {
|
|
const indices = getParamValue("indices", node, tensorMap, context);
|
|
const values = getParamValue("values", node, tensorMap, context);
|
|
const tensor2 = getParamValue("tensor", node, tensorMap, context);
|
|
return [ops.tensorScatterUpdate(tensor2, indices, values)];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$4 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "SparseFillEmptyRows": {
|
|
const { outputIndices, outputValues, emptyRowIndicator, reverseIndexMap } = ops.sparse.sparseFillEmptyRows(getParamValue("indices", node, tensorMap, context), getParamValue("values", node, tensorMap, context), getParamValue("denseShape", node, tensorMap, context), getParamValue("defaultValue", node, tensorMap, context));
|
|
return [
|
|
outputIndices,
|
|
outputValues,
|
|
emptyRowIndicator,
|
|
reverseIndexMap
|
|
];
|
|
}
|
|
case "SparseReshape": {
|
|
const { outputIndices, outputShape } = ops.sparse.sparseReshape(getParamValue("inputIndices", node, tensorMap, context), getParamValue("inputShape", node, tensorMap, context), getParamValue("newShape", node, tensorMap, context));
|
|
return [outputIndices, outputShape];
|
|
}
|
|
case "SparseSegmentMean": {
|
|
const outputData = ops.sparse.sparseSegmentMean(getParamValue("data", node, tensorMap, context), getParamValue("indices", node, tensorMap, context), getParamValue("segmentIds", node, tensorMap, context));
|
|
return [outputData];
|
|
}
|
|
case "SparseSegmentSum": {
|
|
const outputData = ops.sparse.sparseSegmentSum(getParamValue("data", node, tensorMap, context), getParamValue("indices", node, tensorMap, context), getParamValue("segmentIds", node, tensorMap, context));
|
|
return [outputData];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$3 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "FFT": {
|
|
return [ops.fft(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "IFFT": {
|
|
return [ops.ifft(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "RFFT": {
|
|
return [ops.rfft(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
case "IRFFT": {
|
|
return [ops.irfft(getParamValue("x", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$2 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "StaticRegexReplace": {
|
|
return [ops.string.staticRegexReplace(getParamValue("input", node, tensorMap, context), getParamValue("pattern", node, tensorMap, context), getParamValue("rewrite", node, tensorMap, context), getParamValue("replaceGlobal", node, tensorMap, context))];
|
|
}
|
|
case "StringNGrams": {
|
|
const { nGrams, nGramsSplits } = ops.string.stringNGrams(getParamValue("data", node, tensorMap, context), getParamValue("dataSplits", node, tensorMap, context), getParamValue("separator", node, tensorMap, context), getParamValue("nGramWidths", node, tensorMap, context), getParamValue("leftPad", node, tensorMap, context), getParamValue("rightPad", node, tensorMap, context), getParamValue("padWidth", node, tensorMap, context), getParamValue("preserveShortSequences", node, tensorMap, context));
|
|
return [nGrams, nGramsSplits];
|
|
}
|
|
case "StringSplit": {
|
|
const { indices, values, shape } = ops.string.stringSplit(getParamValue("input", node, tensorMap, context), getParamValue("delimiter", node, tensorMap, context), getParamValue("skipEmpty", node, tensorMap, context));
|
|
return [indices, values, shape];
|
|
}
|
|
case "StringToHashBucketFast": {
|
|
const output = ops.string.stringToHashBucketFast(getParamValue("input", node, tensorMap, context), getParamValue("numBuckets", node, tensorMap, context));
|
|
return [output];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
const executeOp$1 = (node, tensorMap, context, ops = tfOps) => {
|
|
switch (node.op) {
|
|
case "Cast": {
|
|
return [ops.cast(getParamValue("x", node, tensorMap, context), getParamValue("dtype", node, tensorMap, context))];
|
|
}
|
|
case "ExpandDims": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
return [ops.expandDims(getParamValue("x", node, tensorMap, context), axis)];
|
|
}
|
|
case "Squeeze": {
|
|
const axis = getParamValue("axis", node, tensorMap, context);
|
|
return [ops.squeeze(getParamValue("x", node, tensorMap, context), axis)];
|
|
}
|
|
case "Reshape": {
|
|
return [ops.reshape(getParamValue("x", node, tensorMap, context), getParamValue("shape", node, tensorMap, context))];
|
|
}
|
|
case "EnsureShape": {
|
|
return [ops.ensureShape(getParamValue("x", node, tensorMap, context), getParamValue("shape", node, tensorMap, context))];
|
|
}
|
|
case "MirrorPad": {
|
|
return [ops.mirrorPad(getParamValue("x", node, tensorMap, context), getParamValue("padding", node, tensorMap, context), getParamValue("mode", node, tensorMap, context))];
|
|
}
|
|
case "PadV2":
|
|
case "Pad": {
|
|
return [ops.pad(getParamValue("x", node, tensorMap, context), getParamValue("padding", node, tensorMap, context), getParamValue("constantValue", node, tensorMap, context))];
|
|
}
|
|
case "SpaceToBatchND": {
|
|
const blockShape = getParamValue("blockShape", node, tensorMap, context);
|
|
const paddings = getParamValue("paddings", node, tensorMap, context);
|
|
return [ops.spaceToBatchND(getParamValue("x", node, tensorMap, context), blockShape, paddings)];
|
|
}
|
|
case "BatchToSpaceND": {
|
|
const blockShape = getParamValue("blockShape", node, tensorMap, context);
|
|
const crops = getParamValue("crops", node, tensorMap, context);
|
|
return [ops.batchToSpaceND(getParamValue("x", node, tensorMap, context), blockShape, crops)];
|
|
}
|
|
case "DepthToSpace": {
|
|
const blockSize = getParamValue("blockSize", node, tensorMap, context);
|
|
const dataFormat = getParamValue("dataFormat", node, tensorMap, context).toUpperCase();
|
|
return [ops.depthToSpace(getParamValue("x", node, tensorMap, context), blockSize, dataFormat)];
|
|
}
|
|
case "BroadcastTo": {
|
|
return [ops.broadcastTo(getParamValue("x", node, tensorMap, context), getParamValue("shape", node, tensorMap, context))];
|
|
}
|
|
case "BroadcastArgs": {
|
|
return [ops.broadcastArgs(getParamValue("s0", node, tensorMap, context), getParamValue("s1", node, tensorMap, context))];
|
|
}
|
|
default:
|
|
throw TypeError(`Node type ${node.op} is not implemented`);
|
|
}
|
|
};
|
|
function executeOp(node, tensorMap, context, resourceManager, tidy$1 = tidy) {
|
|
const value = ((node2, tensorMap2, context2) => {
|
|
switch (node2.category) {
|
|
case "arithmetic":
|
|
return tidy$1(() => executeOp$k(node2, tensorMap2, context2));
|
|
case "basic_math":
|
|
return tidy$1(() => executeOp$j(node2, tensorMap2, context2));
|
|
case "control":
|
|
return executeOp$i(node2, tensorMap2, context2);
|
|
case "convolution":
|
|
return tidy$1(() => executeOp$h(node2, tensorMap2, context2));
|
|
case "creation":
|
|
return tidy$1(() => executeOp$g(node2, tensorMap2, context2));
|
|
case "dynamic":
|
|
return executeOp$f(node2, tensorMap2, context2);
|
|
case "evaluation":
|
|
return tidy$1(() => executeOp$e(node2, tensorMap2, context2));
|
|
case "image":
|
|
return tidy$1(() => executeOp$b(node2, tensorMap2, context2));
|
|
case "graph":
|
|
return tidy$1(() => executeOp$d(node2, tensorMap2, context2));
|
|
case "logical":
|
|
return tidy$1(() => executeOp$a(node2, tensorMap2, context2));
|
|
case "matrices":
|
|
return tidy$1(() => executeOp$9(node2, tensorMap2, context2));
|
|
case "normalization":
|
|
return tidy$1(() => executeOp$8(node2, tensorMap2, context2));
|
|
case "ragged":
|
|
return tidy$1(() => executeOp$7(node2, tensorMap2, context2));
|
|
case "reduction":
|
|
return tidy$1(() => executeOp$6(node2, tensorMap2, context2));
|
|
case "slice_join":
|
|
return tidy$1(() => executeOp$5(node2, tensorMap2, context2));
|
|
case "sparse":
|
|
return tidy$1(() => executeOp$4(node2, tensorMap2, context2));
|
|
case "spectral":
|
|
return tidy$1(() => executeOp$3(node2, tensorMap2, context2));
|
|
case "string":
|
|
return tidy$1(() => executeOp$2(node2, tensorMap2, context2));
|
|
case "transformation":
|
|
return tidy$1(() => executeOp$1(node2, tensorMap2, context2));
|
|
case "hash_table":
|
|
return executeOp$c(node2, tensorMap2, context2, resourceManager);
|
|
case "custom":
|
|
const opMapper = getRegisteredOp(node2.op);
|
|
if (opMapper && opMapper.customExecutor) {
|
|
return opMapper.customExecutor(new NodeValueImpl(node2, tensorMap2, context2));
|
|
} else {
|
|
throw TypeError(`Custom op ${node2.op} is not registered.`);
|
|
}
|
|
default:
|
|
throw TypeError(`Unknown op '${node2.op}'. File an issue at https://github.com/tensorflow/tfjs/issues so we can add it, or register a custom execution with tf.registerOp()`);
|
|
}
|
|
})(node, tensorMap, context);
|
|
if (isPromise(value)) {
|
|
return value.then((data) => [].concat(data));
|
|
}
|
|
return [].concat(value);
|
|
}
|
|
class ExecutionContext {
|
|
constructor(weightMap = {}, tensorArrayMap = {}, tensorListMap = {}, functionMap = {}, parseNodeNameCache) {
|
|
this.weightMap = weightMap;
|
|
this.tensorArrayMap = tensorArrayMap;
|
|
this.tensorListMap = tensorListMap;
|
|
this.functionMap = functionMap;
|
|
this.parseNodeNameCache = parseNodeNameCache;
|
|
this.rootContext = { id: 0, frameName: "", iterationId: 0 };
|
|
this.contexts = [this.rootContext];
|
|
this.lastId = 0;
|
|
this.generateCurrentContextIds();
|
|
}
|
|
newFrame(id, frameName) {
|
|
return { id, frameName, iterationId: 0 };
|
|
}
|
|
/**
|
|
* Set the current context
|
|
* @param contexts: ExecutionContextInfo[] the current path of execution
|
|
* frames
|
|
*/
|
|
set currentContext(contexts2) {
|
|
if (this.contexts !== contexts2) {
|
|
this.contexts = contexts2;
|
|
this.generateCurrentContextIds();
|
|
}
|
|
}
|
|
get currentContext() {
|
|
return this.contexts;
|
|
}
|
|
/**
|
|
* Returns the current context in string format.
|
|
*/
|
|
get currentContextId() {
|
|
return this._currentContextIds[0];
|
|
}
|
|
/**
|
|
* Returns the current context and all parent contexts in string format.
|
|
* This allow access to the nodes in the current and parent frames.
|
|
*/
|
|
get currentContextIds() {
|
|
return this._currentContextIds;
|
|
}
|
|
generateCurrentContextIds() {
|
|
const names = [];
|
|
for (let i = 0; i < this.contexts.length - 1; i++) {
|
|
const contexts2 = this.contexts.slice(0, this.contexts.length - i);
|
|
names.push(this.contextIdforContexts(contexts2));
|
|
}
|
|
names.push("");
|
|
this._currentContextIds = names;
|
|
}
|
|
contextIdforContexts(contexts2) {
|
|
return contexts2 ? contexts2.map((context) => context.id === 0 && context.iterationId === 0 ? "" : `${context.frameName}-${context.iterationId}`).join("/") : "";
|
|
}
|
|
/**
|
|
* Enter a new frame, a new context is pushed on the current context list.
|
|
* @param frameId new frame id
|
|
*/
|
|
enterFrame(frameId) {
|
|
if (this.contexts) {
|
|
this.lastId++;
|
|
this.contexts = this.contexts.slice();
|
|
this.contexts.push(this.newFrame(this.lastId, frameId));
|
|
this._currentContextIds.unshift(this.contextIdforContexts(this.contexts));
|
|
}
|
|
}
|
|
/**
|
|
* Exit the current frame, the last context is removed from the current
|
|
* context list.
|
|
*/
|
|
exitFrame() {
|
|
if (this.contexts && this.contexts.length > 1) {
|
|
this.contexts = this.contexts.slice();
|
|
this.contexts.splice(-1);
|
|
this.currentContextIds.shift();
|
|
} else {
|
|
throw new Error("Cannot exit frame, the context is empty");
|
|
}
|
|
}
|
|
/**
|
|
* Enter the next iteration of a loop, the iteration id of last context is
|
|
* increased.
|
|
*/
|
|
nextIteration() {
|
|
if (this.contexts && this.contexts.length > 0) {
|
|
this.contexts = this.contexts.slice();
|
|
this.lastId++;
|
|
const context = Object.assign({}, this.contexts[this.contexts.length - 1]);
|
|
context.iterationId += 1;
|
|
context.id = this.lastId;
|
|
this.contexts.splice(-1, 1, context);
|
|
this._currentContextIds.splice(0, 1, this.contextIdforContexts(this.contexts));
|
|
} else {
|
|
throw new Error("Cannot increase frame iteration, the context is empty");
|
|
}
|
|
}
|
|
getWeight(name) {
|
|
return this.weightMap[name];
|
|
}
|
|
addTensorArray(tensorArray) {
|
|
this.tensorArrayMap[tensorArray.id] = tensorArray;
|
|
}
|
|
getTensorArray(id) {
|
|
return this.tensorArrayMap[id];
|
|
}
|
|
addTensorList(tensorList) {
|
|
this.tensorListMap[tensorList.id] = tensorList;
|
|
}
|
|
getTensorList(id) {
|
|
return this.tensorListMap[id];
|
|
}
|
|
dispose(keepIds) {
|
|
for (const key in this.tensorArrayMap) {
|
|
this.tensorArrayMap[key].clearAndClose(keepIds);
|
|
}
|
|
for (const key in this.tensorListMap) {
|
|
this.tensorListMap[key].clearAndClose(keepIds);
|
|
}
|
|
}
|
|
}
|
|
function getExecutionSubgraph(inputs, outputs, weightMap, initNodes) {
|
|
const usedNodes = /* @__PURE__ */ new Set();
|
|
const missingInputs = [];
|
|
let dynamicNode = null;
|
|
let syncInputs = null;
|
|
const seen = /* @__PURE__ */ new Set();
|
|
const inputNodeNames = new Set(Object.keys(inputs).map((name) => parseNodeName(name)[0]));
|
|
initNodes = initNodes || [];
|
|
const initNodeNames = new Set(initNodes.map((node) => parseNodeName(node.name)[0]));
|
|
const frontier = [...outputs];
|
|
while (frontier.length > 0) {
|
|
const node = frontier.pop();
|
|
if (isControlFlow(node) || isDynamicShape(node) || isHashTable(node)) {
|
|
if (dynamicNode == null) {
|
|
dynamicNode = node;
|
|
syncInputs = dynamicNode.children.map((child) => child.name).filter((name) => usedNodes.has(name));
|
|
}
|
|
}
|
|
usedNodes.add(node.name);
|
|
if (weightMap[node.name] != null) {
|
|
continue;
|
|
}
|
|
if (inputNodeNames.has(node.name)) {
|
|
continue;
|
|
}
|
|
if (initNodeNames.has(node.name)) {
|
|
continue;
|
|
}
|
|
if (node.inputs.length === 0) {
|
|
missingInputs.push(node.name);
|
|
continue;
|
|
}
|
|
node.inputs.forEach((input) => {
|
|
if (seen.has(input.name)) {
|
|
return;
|
|
}
|
|
seen.add(input.name);
|
|
frontier.push(input);
|
|
});
|
|
}
|
|
return { inputs, outputs, usedNodes, missingInputs, dynamicNode, syncInputs };
|
|
}
|
|
function getNodesInTopologicalOrder(graph2, executionInfo) {
|
|
const { usedNodes, inputs } = executionInfo;
|
|
const inputNodes = Object.keys(inputs).map((name) => parseNodeName(name)[0]).map((name) => graph2.nodes[name]);
|
|
const initNodes = graph2.initNodes || [];
|
|
const isUsed = (node) => usedNodes.has(typeof node === "string" ? node : node.name);
|
|
function unique2(nodes) {
|
|
return [...new Map(nodes.map((node) => [node.name, node])).values()];
|
|
}
|
|
const predefinedNodes = unique2([
|
|
...inputNodes,
|
|
...graph2.weights,
|
|
...initNodes
|
|
]).filter(isUsed);
|
|
const allNodes = unique2([
|
|
...predefinedNodes,
|
|
...Object.values(graph2.nodes)
|
|
]).filter(isUsed);
|
|
const nameToNode = new Map(allNodes.map((node) => [node.name, node]));
|
|
const inCounts = {};
|
|
for (const node of allNodes) {
|
|
inCounts[node.name] = inCounts[node.name] || 0;
|
|
for (const child of node.children) {
|
|
if (!isUsed(child)) {
|
|
inCounts[child.name] = Number.POSITIVE_INFINITY;
|
|
}
|
|
inCounts[child.name] = (inCounts[child.name] || 0) + 1;
|
|
}
|
|
}
|
|
const frontier = Object.entries(inCounts).filter(([, inCount]) => inCount === 0).map(([name]) => name);
|
|
const orderedNodeNames = [...frontier];
|
|
while (frontier.length > 0) {
|
|
const nodeName = frontier.pop();
|
|
const node = nameToNode.get(nodeName);
|
|
for (const child of node.children.filter(isUsed)) {
|
|
if (--inCounts[child.name] === 0) {
|
|
orderedNodeNames.push(child.name);
|
|
frontier.push(child.name);
|
|
}
|
|
}
|
|
}
|
|
const orderedNodes = orderedNodeNames.map((name) => nameToNode.get(name));
|
|
const filteredOrderedNodes = filterPredefinedReachableNodes(orderedNodes, predefinedNodes);
|
|
validateNodesExecutionOrder(filteredOrderedNodes, predefinedNodes);
|
|
return filteredOrderedNodes;
|
|
}
|
|
function filterPredefinedReachableNodes(orderedNodes, predefinedNodes) {
|
|
const nameToNode = new Map(orderedNodes.map((node) => [node.name, node]));
|
|
const stack2 = predefinedNodes.map((node) => node.name);
|
|
const predefinedReachableNodeNames = new Set(stack2);
|
|
while (stack2.length > 0) {
|
|
const nodeName = stack2.pop();
|
|
const node = nameToNode.get(nodeName);
|
|
for (const child of node.children) {
|
|
if (!nameToNode.has(child.name) || predefinedReachableNodeNames.has(child.name)) {
|
|
continue;
|
|
}
|
|
predefinedReachableNodeNames.add(child.name);
|
|
stack2.push(child.name);
|
|
}
|
|
}
|
|
const filteredOrderedNodes = orderedNodes.filter((node) => predefinedReachableNodeNames.has(node.name));
|
|
return filteredOrderedNodes;
|
|
}
|
|
class NodesExecutionOrderError extends Error {
|
|
constructor(message) {
|
|
super(`NodesExecutionOrderError: ${message}`);
|
|
}
|
|
}
|
|
function validateNodesExecutionOrder(orderedNodes, predefinedNodes) {
|
|
const nodeNameToOrder = new Map(orderedNodes.map((node, order) => [node.name, order]));
|
|
const predefinedNodeNames = new Set(predefinedNodes.map((node) => node.name));
|
|
const isPredefined = (node) => predefinedNodeNames.has(typeof node === "string" ? node : node.name);
|
|
const willBeExecutedNodeNames = new Set(orderedNodes.map((node) => node.name));
|
|
const willBeExecuted = (node) => willBeExecutedNodeNames.has(typeof node === "string" ? node : node.name);
|
|
for (const node of orderedNodes) {
|
|
for (const child of node.children.filter(willBeExecuted)) {
|
|
if (!nodeNameToOrder.has(child.name)) {
|
|
throw new NodesExecutionOrderError(`Child ${child.name} of node ${node.name} is unreachable.`);
|
|
}
|
|
if (nodeNameToOrder.get(node.name) > nodeNameToOrder.get(child.name)) {
|
|
throw new NodesExecutionOrderError(`Node ${node.name} is scheduled to run after its child ${child.name}.`);
|
|
}
|
|
}
|
|
if (!isPredefined(node)) {
|
|
for (const input of node.inputs) {
|
|
if (!nodeNameToOrder.has(input.name)) {
|
|
throw new NodesExecutionOrderError(`Input ${input.name} of node ${node.name} is unreachable.`);
|
|
}
|
|
if (nodeNameToOrder.get(input.name) > nodeNameToOrder.get(node.name)) {
|
|
throw new NodesExecutionOrderError(`Node ${node.name} is scheduled to run before its input ${input.name}.`);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
function getNodeLiveUntilMap(orderedNodes) {
|
|
const nodeNameToOrder = new Map(orderedNodes.map((node, order) => [node.name, order]));
|
|
const INF_LIFE = Number.MAX_SAFE_INTEGER;
|
|
const selfLifespans = orderedNodes.map((node, nodeOrder) => isControlFlow(node) ? INF_LIFE : nodeOrder);
|
|
const getSelfLifeSpan = (node) => {
|
|
const selfLife = selfLifespans[nodeNameToOrder.get(node.name)];
|
|
if (selfLife == null) {
|
|
return -1;
|
|
}
|
|
return selfLife;
|
|
};
|
|
const liveUntilOrders = orderedNodes.map((node, nodeOrder) => {
|
|
return node.children.map(getSelfLifeSpan).reduce((a, b) => Math.max(a, b), selfLifespans[nodeOrder]);
|
|
});
|
|
const liveUntilMap = /* @__PURE__ */ new Map();
|
|
for (let nodeOrder = 0; nodeOrder < orderedNodes.length; ++nodeOrder) {
|
|
const liveUntilOrder = liveUntilOrders[nodeOrder];
|
|
if (liveUntilOrder === INF_LIFE) {
|
|
continue;
|
|
}
|
|
const node = orderedNodes[nodeOrder];
|
|
const liveUntilNode = orderedNodes[liveUntilOrder];
|
|
if (!liveUntilMap.has(liveUntilNode.name)) {
|
|
liveUntilMap.set(liveUntilNode.name, []);
|
|
}
|
|
liveUntilMap.get(liveUntilNode.name).push(node);
|
|
}
|
|
return liveUntilMap;
|
|
}
|
|
const CONTROL_FLOW_OPS = /* @__PURE__ */ new Set([
|
|
"Switch",
|
|
"Merge",
|
|
"Enter",
|
|
"Exit",
|
|
"NextIteration",
|
|
"StatelessIf",
|
|
"StatelessWhile",
|
|
"if",
|
|
"While"
|
|
]);
|
|
const DYNAMIC_SHAPE_OPS = /* @__PURE__ */ new Set([
|
|
"NonMaxSuppressionV2",
|
|
"NonMaxSuppressionV3",
|
|
"NonMaxSuppressionV5",
|
|
"Where"
|
|
]);
|
|
const HASH_TABLE_OPS = /* @__PURE__ */ new Set([
|
|
"HashTable",
|
|
"HashTableV2",
|
|
"LookupTableImport",
|
|
"LookupTableImportV2",
|
|
"LookupTableFind",
|
|
"LookupTableFindV2",
|
|
"LookupTableSize",
|
|
"LookupTableSizeV2"
|
|
]);
|
|
function isControlFlow(node) {
|
|
return CONTROL_FLOW_OPS.has(node.op);
|
|
}
|
|
function isDynamicShape(node) {
|
|
return DYNAMIC_SHAPE_OPS.has(node.op);
|
|
}
|
|
function isHashTable(node) {
|
|
return HASH_TABLE_OPS.has(node.op);
|
|
}
|
|
class GraphExecutor {
|
|
get weightIds() {
|
|
return this.parent ? this.parent.weightIds : this._weightIds;
|
|
}
|
|
get functionExecutorMap() {
|
|
return this.parent ? this.parent.functionExecutorMap : this._functionExecutorMap;
|
|
}
|
|
get weightMap() {
|
|
return this.parent ? this.parent.weightMap : this._weightMap;
|
|
}
|
|
set weightMap(weightMap) {
|
|
const weightIds = Object.keys(weightMap).map((key) => weightMap[key].map((tensor2) => tensor2.id));
|
|
this._weightIds = [].concat(...weightIds);
|
|
this._weightMap = weightMap;
|
|
}
|
|
/**
|
|
* Set `ResourceManager` shared by executors of a model.
|
|
* @param resourceManager: `ResourceManager` of the `GraphModel`.
|
|
*/
|
|
set resourceManager(resourceManager) {
|
|
this._resourceManager = resourceManager;
|
|
}
|
|
get inputs() {
|
|
return this._inputs.map((node) => {
|
|
return {
|
|
name: node.name,
|
|
shape: node.attrParams["shape"] ? node.attrParams["shape"].value : void 0,
|
|
dtype: node.attrParams["dtype"] ? node.attrParams["dtype"].value : void 0
|
|
};
|
|
});
|
|
}
|
|
get outputs() {
|
|
return this._outputs.map((node) => {
|
|
return {
|
|
name: node.name,
|
|
shape: node.attrParams["shape"] ? node.attrParams["shape"].value : void 0,
|
|
dtype: node.attrParams["dtype"] ? node.attrParams["dtype"].value : void 0
|
|
};
|
|
});
|
|
}
|
|
get inputNodes() {
|
|
return this._inputs.map((node) => node.signatureKey || node.name);
|
|
}
|
|
get outputNodes() {
|
|
return this._outputs.map((node) => {
|
|
const name = node.signatureKey || node.name;
|
|
return node.defaultOutput ? `${name}:${node.defaultOutput}` : name;
|
|
});
|
|
}
|
|
get functions() {
|
|
return Object.keys(this._functions).reduce((map, key) => {
|
|
map[key] = this._functions[key].signature;
|
|
return map;
|
|
}, {});
|
|
}
|
|
/**
|
|
*
|
|
* @param graph Graph the model or function graph to be executed.
|
|
* @param parent When building function exector you need to set the parent
|
|
* executor. Since the weights and function executor maps are set at parant
|
|
* level, that function executor can access the function maps and weight maps
|
|
* through the parent.
|
|
*/
|
|
constructor(graph2, parent) {
|
|
this.graph = graph2;
|
|
this.parent = parent;
|
|
this.compiledMap = /* @__PURE__ */ new Map();
|
|
this.parseNodeNameCache = /* @__PURE__ */ new Map();
|
|
this._weightMap = {};
|
|
this.SEPARATOR = ",";
|
|
this._functions = {};
|
|
this._functionExecutorMap = {};
|
|
this.keepIntermediateTensors = false;
|
|
this._outputs = graph2.outputs;
|
|
this._inputs = graph2.inputs;
|
|
this._initNodes = graph2.initNodes;
|
|
this._signature = graph2.signature;
|
|
this._functions = graph2.functions;
|
|
if (graph2.functions != null) {
|
|
Object.keys(graph2.functions).forEach((name) => {
|
|
this._functionExecutorMap[name] = new GraphExecutor(graph2.functions[name], this);
|
|
});
|
|
}
|
|
}
|
|
getCompilationKey(inputs, outputs) {
|
|
const sortedInputs = inputs.map((node) => node.name).sort();
|
|
const sortedOutputs = outputs.map((node) => node.name).sort();
|
|
return sortedInputs.join(this.SEPARATOR) + "--" + sortedOutputs.join(this.SEPARATOR);
|
|
}
|
|
/**
|
|
* Compiles the inference graph and returns the minimal set of nodes that are
|
|
* required for execution, in the correct execution order.
|
|
* @returns {Object} compilation The compile result.
|
|
* @returns {Node[]} compilation.orderedNodes Nodes in the correct execution
|
|
* order.
|
|
* @returns {Map<string, Node[]>} compilation.nodeLiveUntilMap A map from node
|
|
* to disposable nodes after its execution. That is, for a node `x`,
|
|
* `nodeLiveUntilMap[x]` indicates all nodes whose intermediate
|
|
* tensors should be disposed after `x` is executed.
|
|
*/
|
|
compile(inputs, outputs) {
|
|
const executionInfo = getExecutionSubgraph(inputs, outputs, this.weightMap, this._initNodes);
|
|
const { missingInputs, dynamicNode, syncInputs } = executionInfo;
|
|
if (dynamicNode != null) {
|
|
throw new Error(`This execution contains the node '${dynamicNode.name}', which has the dynamic op '${dynamicNode.op}'. Please use model.executeAsync() instead. Alternatively, to avoid the dynamic ops, specify the inputs [${syncInputs}]`);
|
|
}
|
|
if (missingInputs.length > 0) {
|
|
const outNames = outputs.map((n) => n.name);
|
|
const inNames = Object.keys(inputs);
|
|
throw new Error(`Cannot compute the outputs [${outNames}] from the provided inputs [${inNames}]. Missing the following inputs: [${missingInputs}]`);
|
|
}
|
|
const orderedNodes = getNodesInTopologicalOrder(this.graph, executionInfo);
|
|
const nodeLiveUntilMap = getNodeLiveUntilMap(orderedNodes);
|
|
return { orderedNodes, nodeLiveUntilMap };
|
|
}
|
|
cloneAndKeepTensor(tensor2) {
|
|
if (tensor2 == null) {
|
|
return null;
|
|
}
|
|
const clone2 = tensor2.clone();
|
|
keep(clone2);
|
|
return clone2;
|
|
}
|
|
cloneTensorList(tensors) {
|
|
if (!tensors) {
|
|
return null;
|
|
}
|
|
const clonedTensor = tensors.map((tensor2) => {
|
|
return this.cloneAndKeepTensor(tensor2);
|
|
});
|
|
return clonedTensor;
|
|
}
|
|
cloneTensorMap(tensorsMap) {
|
|
return Object.fromEntries(Object.entries(tensorsMap).map(([name, tensorsList]) => {
|
|
return [name, this.cloneTensorList(tensorsList)];
|
|
}));
|
|
}
|
|
/**
|
|
* Executes the inference for given input tensors.
|
|
* @param inputs Tensor map for the model inputs, keyed by the input node
|
|
* names.
|
|
* @param outputs Optional. output node name from the Tensorflow model, if
|
|
* no outputs are specified, the default outputs of the model would be used.
|
|
* You can inspect intermediate nodes of the model by adding them to the
|
|
* outputs array.
|
|
*/
|
|
execute(inputs, outputs) {
|
|
this.disposeIntermediateTensors();
|
|
inputs = this.mapInputs(inputs);
|
|
const names = Object.keys(inputs).sort();
|
|
this.checkInputs(inputs);
|
|
this.checkInputShapeAndType(inputs);
|
|
outputs = this.mapOutputs(outputs);
|
|
this.checkOutputs(outputs);
|
|
const inputNodes = names.map((name) => this.graph.nodes[parseNodeName(name)[0]]);
|
|
const outputNodeNames = outputs.map((name) => parseNodeName(name)[0]);
|
|
const outputNodeNameSet = new Set(outputNodeNames);
|
|
let outputNodes = outputNodeNames.map((name) => this.graph.nodes[name]);
|
|
if (outputNodes.length === 0) {
|
|
outputNodes = this._outputs;
|
|
}
|
|
const compilationKey = this.getCompilationKey(inputNodes, outputNodes);
|
|
let compilation = this.compiledMap.get(compilationKey);
|
|
if (compilation == null) {
|
|
compilation = this.compile(inputs, outputNodes);
|
|
this.compiledMap.set(compilationKey, compilation);
|
|
}
|
|
try {
|
|
this.keepIntermediateTensors = env().getBool("KEEP_INTERMEDIATE_TENSORS");
|
|
} catch (e) {
|
|
this.keepIntermediateTensors = false;
|
|
console.warn(e.message);
|
|
}
|
|
const tensorArrayMap = {};
|
|
const tensorListMap = {};
|
|
return tidy(() => {
|
|
const context = new ExecutionContext(this.weightMap, tensorArrayMap, tensorListMap, this.functionExecutorMap, this.parseNodeNameCache);
|
|
const tensorsMap = Object.assign({}, this.weightMap);
|
|
if (this.keepIntermediateTensors) {
|
|
this.clonedTensorsMap = this.cloneTensorMap(this.weightMap);
|
|
}
|
|
Object.keys(inputs).forEach((name) => {
|
|
const [nodeName, index] = parseNodeName(name, context);
|
|
const tensors = [];
|
|
tensors[index] = inputs[name];
|
|
tensorsMap[nodeName] = tensors;
|
|
if (this.keepIntermediateTensors) {
|
|
this.clonedTensorsMap[nodeName] = this.cloneTensorList(tensors);
|
|
}
|
|
});
|
|
const tensorsToKeep = this.getFrozenTensorIds(tensorsMap);
|
|
const { orderedNodes, nodeLiveUntilMap } = compilation;
|
|
for (const node of orderedNodes) {
|
|
if (tensorsMap[node.name]) {
|
|
continue;
|
|
}
|
|
const tensors = executeOp(node, tensorsMap, context, this._resourceManager);
|
|
if (isPromise(tensors)) {
|
|
throw new Error(`The execution of the op '${node.op}' returned a promise. Please use model.executeAsync() instead.`);
|
|
}
|
|
tensorsMap[node.name] = tensors;
|
|
if (this.keepIntermediateTensors) {
|
|
this.clonedTensorsMap[node.name] = this.cloneTensorList(tensors);
|
|
}
|
|
this.checkTensorForDisposalWithNodeLiveUntilInfo(node, tensorsMap, context, tensorsToKeep, outputNodeNameSet, nodeLiveUntilMap.get(node.name));
|
|
}
|
|
if (this.parent == null) {
|
|
context.dispose(tensorsToKeep);
|
|
}
|
|
return outputs.map((name) => getTensor(name, tensorsMap, context));
|
|
});
|
|
}
|
|
getFrozenTensorIds(tensorMap) {
|
|
const ids = [].concat.apply([], Object.keys(tensorMap).map((key) => tensorMap[key]).map((tensors) => tensors.map((tensor2) => tensor2.id)));
|
|
return new Set(ids);
|
|
}
|
|
checkTensorForDisposal(nodeName, node, tensorMap, context, tensorsToKeep, outputNodeNameSet, intermediateTensorConsumerCount) {
|
|
if (isControlFlow(node) || outputNodeNameSet.has(nodeName)) {
|
|
return;
|
|
}
|
|
for (const tensor2 of tensorMap[nodeName]) {
|
|
if (tensor2 == null) {
|
|
continue;
|
|
}
|
|
intermediateTensorConsumerCount[tensor2.id] = (intermediateTensorConsumerCount[tensor2.id] || 0) + node.children.length;
|
|
}
|
|
for (const input of node.inputs) {
|
|
if (isControlFlow(input)) {
|
|
continue;
|
|
}
|
|
const tensors = getTensorsForCurrentContext(input.name, tensorMap, context);
|
|
if (tensors == null) {
|
|
continue;
|
|
}
|
|
for (const tensor2 of tensors) {
|
|
if (!tensor2 || tensor2.kept || tensorsToKeep.has(tensor2.id)) {
|
|
continue;
|
|
}
|
|
const count = intermediateTensorConsumerCount[tensor2.id];
|
|
if (count === 1) {
|
|
tensor2.dispose();
|
|
delete intermediateTensorConsumerCount[tensor2.id];
|
|
} else if (count != null) {
|
|
intermediateTensorConsumerCount[tensor2.id]--;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
checkTensorForDisposalWithNodeLiveUntilInfo(node, tensorMap, context, tensorsToKeep, outputNodeNameSet, liveUntilNodes) {
|
|
function isNonDisposableNode(node2) {
|
|
return isControlFlow(node2) || outputNodeNameSet.has(node2.name);
|
|
}
|
|
if (isControlFlow(node) || liveUntilNodes == null) {
|
|
return;
|
|
}
|
|
for (const nodeToDispose of liveUntilNodes) {
|
|
if (isNonDisposableNode(nodeToDispose)) {
|
|
continue;
|
|
}
|
|
const tensors = getTensorsForCurrentContext(nodeToDispose.name, tensorMap, context);
|
|
for (const tensor2 of tensors) {
|
|
if (!tensor2 || tensor2.kept || tensorsToKeep.has(tensor2.id)) {
|
|
continue;
|
|
}
|
|
tensor2.dispose();
|
|
}
|
|
}
|
|
}
|
|
/**
|
|
* Executes the inference for given input tensors in Async fashion.
|
|
* @param inputs Tensor map for the model inputs, keyed by the input node
|
|
* names.
|
|
* @param outputs output node name from the Tensorflow model, if no outputs
|
|
* are specified, the default outputs of the model would be used. You can
|
|
* inspect intermediate nodes of the model by adding them to the outputs
|
|
* array.
|
|
*/
|
|
async executeAsync(inputs, outputs) {
|
|
return this._executeAsync(inputs, outputs);
|
|
}
|
|
disposeIntermediateTensors() {
|
|
if (!this.clonedTensorsMap) {
|
|
return;
|
|
}
|
|
Object.values(this.clonedTensorsMap).forEach((tensorsList) => {
|
|
for (const tensor2 of tensorsList) {
|
|
if (tensor2 && !tensor2.isDisposed) {
|
|
tensor2.dispose();
|
|
}
|
|
}
|
|
});
|
|
this.clonedTensorsMap = null;
|
|
}
|
|
getIntermediateTensors() {
|
|
return this.clonedTensorsMap;
|
|
}
|
|
/**
|
|
* Executes the inference for given input tensors in Async fashion.
|
|
* @param inputs Tensor map for the model inputs, keyed by the input node
|
|
* names.
|
|
* @param outputs Optional. output node name from the Tensorflow model,
|
|
* if no outputs are specified, the default outputs of the model would be
|
|
* used. You can inspect intermediate nodes of the model by adding them to
|
|
* the outputs array.
|
|
* @param isFunctionExecution Optional. Flag for executing a function.
|
|
* @param tensorArrayMap Optional, global TensorArray map by id. Used for
|
|
* function execution.
|
|
* @param tensorArrayMap Optional global TensorList map by id. Used for
|
|
* function execution.
|
|
*/
|
|
async _executeAsync(inputs, outputs, isFunctionExecution = false, tensorArrayMap = {}, tensorListMap = {}) {
|
|
this.disposeIntermediateTensors();
|
|
if (!isFunctionExecution) {
|
|
inputs = this.mapInputs(inputs);
|
|
this.checkInputs(inputs);
|
|
this.checkInputShapeAndType(inputs);
|
|
outputs = this.mapOutputs(outputs);
|
|
this.checkOutputs(outputs);
|
|
}
|
|
try {
|
|
this.keepIntermediateTensors = env().getBool("KEEP_INTERMEDIATE_TENSORS");
|
|
} catch (e) {
|
|
this.keepIntermediateTensors = false;
|
|
console.warn(e.message);
|
|
}
|
|
const context = new ExecutionContext(this.weightMap, tensorArrayMap, tensorListMap, this.functionExecutorMap, this.parseNodeNameCache);
|
|
if (this.keepIntermediateTensors) {
|
|
this.clonedTensorsMap = this.cloneTensorMap(this.weightMap);
|
|
}
|
|
const tensorsMap = await this.executeWithControlFlow(inputs, context, outputs, isFunctionExecution);
|
|
const results = outputs.map((name) => getTensor(name, tensorsMap, context));
|
|
const outputIds = results.map((t2) => t2.id);
|
|
const inputIds = Object.keys(inputs).map((name) => inputs[name].id);
|
|
const keepIds = /* @__PURE__ */ new Set([...outputIds, ...inputIds, ...this.weightIds]);
|
|
Object.values(tensorsMap).forEach((tensorsList) => {
|
|
tensorsList.forEach((tensor2) => {
|
|
if (tensor2 && !tensor2.isDisposed && !keepIds.has(tensor2.id)) {
|
|
tensor2.dispose();
|
|
}
|
|
});
|
|
});
|
|
if (this.parent == null) {
|
|
context.dispose(keepIds);
|
|
}
|
|
return results;
|
|
}
|
|
async executeFunctionAsync(inputs, tensorArrayMap, tensorListMap) {
|
|
const mappedInputs = inputs.reduce((map, tensor2, index) => {
|
|
map[this.inputs[index].name] = tensor2;
|
|
return map;
|
|
}, {});
|
|
return this._executeAsync(mappedInputs, this.outputNodes, true, tensorArrayMap, tensorListMap);
|
|
}
|
|
/**
|
|
* When there are control flow nodes in the graph, the graph execution use
|
|
* ExecutionContext to keep track of the frames and loop iterators.
|
|
* @param inputs placeholder tensors for the graph.
|
|
* @param context the execution context object for current execution.
|
|
* @param outputNames Optional. output node name from the Tensorflow model,
|
|
* if no outputs are specified, the default outputs of the model would be
|
|
* used. You can inspect intermediate nodes of the model by adding them to
|
|
* the outputs array.
|
|
* @param isFunctionExecution Flag for executing a function.
|
|
*/
|
|
async executeWithControlFlow(inputs, context, outputNames, isFunctionExecution) {
|
|
const names = Object.keys(inputs);
|
|
const inputNodes = names.map((name) => this.graph.nodes[parseNodeName(name)[0]]);
|
|
const outputNodeNames = outputNames.map((name) => parseNodeName(name)[0]);
|
|
const outputNodeNameSet = new Set(outputNodeNames);
|
|
let outputNodes = outputNodeNames.map((name) => this.graph.nodes[name]);
|
|
if (outputNodes.length === 0) {
|
|
outputNodes = this._outputs;
|
|
}
|
|
const { usedNodes, missingInputs, dynamicNode, syncInputs } = getExecutionSubgraph(inputs, outputNodes, this.weightMap, this._initNodes);
|
|
const stack2 = [
|
|
...inputNodes,
|
|
...this.graph.weights,
|
|
...this._initNodes || []
|
|
].map((node) => {
|
|
return { node, contexts: context.currentContext };
|
|
});
|
|
const tensorsMap = Object.assign({}, this.weightMap);
|
|
Object.keys(inputs).forEach((name) => {
|
|
const [nodeName, index] = parseNodeName(name);
|
|
const tensors = [];
|
|
tensors[index] = inputs[name];
|
|
tensorsMap[nodeName] = tensors;
|
|
});
|
|
const intermediateTensorConsumerCount = {};
|
|
const tensorsToKeep = this.getFrozenTensorIds(tensorsMap);
|
|
const added = {};
|
|
while (stack2.length > 0) {
|
|
const promises = this.processStack(inputNodes, stack2, context, tensorsMap, added, tensorsToKeep, outputNodeNameSet, intermediateTensorConsumerCount, usedNodes);
|
|
await Promise.all(promises);
|
|
}
|
|
if (dynamicNode == null && !isFunctionExecution) {
|
|
console.warn(`This model execution did not contain any nodes with control flow or dynamic output shapes. You can use model.execute() instead.`);
|
|
}
|
|
const missingOutputs = outputNodes.filter((node) => !isControlFlow(node) && !getTensor(node.name, tensorsMap, context)).map((node) => node.name);
|
|
if (missingOutputs.length > 0) {
|
|
let alternativeMsg = "";
|
|
if (dynamicNode != null) {
|
|
alternativeMsg = `Alternatively, to avoid the dynamic ops, use model.execute() and specify the inputs [${syncInputs}]`;
|
|
}
|
|
throw new Error(`Cannot compute the outputs [${missingOutputs}] from the provided inputs [${names}]. Consider providing the following inputs: [${missingInputs}]. ${alternativeMsg}`);
|
|
}
|
|
return tensorsMap;
|
|
}
|
|
processStack(inputNodes, stack2, context, tensorMap, added, tensorsToKeep, outputNodeNameSet, intermediateTensorConsumerCount, usedNodes) {
|
|
const promises = [];
|
|
while (stack2.length > 0) {
|
|
const item = stack2.pop();
|
|
context.currentContext = item.contexts;
|
|
let nodeName = "";
|
|
if (item.node.op === "Enter" && getParamValue("isConstant", item.node, tensorMap, context)) {
|
|
[nodeName] = getNodeNameAndIndex(item.node.name, context);
|
|
}
|
|
if (tensorMap[item.node.name] == null) {
|
|
const tensors = executeOp(item.node, tensorMap, context, this._resourceManager);
|
|
if (!nodeName) {
|
|
[nodeName] = getNodeNameAndIndex(item.node.name, context);
|
|
}
|
|
const currentContext = context.currentContext;
|
|
if (isPromise(tensors)) {
|
|
promises.push(tensors.then((t2) => {
|
|
tensorMap[nodeName] = t2;
|
|
if (this.keepIntermediateTensors) {
|
|
this.clonedTensorsMap[nodeName] = this.cloneTensorList(t2);
|
|
}
|
|
context.currentContext = currentContext;
|
|
this.checkTensorForDisposal(nodeName, item.node, tensorMap, context, tensorsToKeep, outputNodeNameSet, intermediateTensorConsumerCount);
|
|
this.processChildNodes(item.node, stack2, context, tensorMap, added, usedNodes);
|
|
return t2;
|
|
}));
|
|
} else {
|
|
tensorMap[nodeName] = tensors;
|
|
if (this.keepIntermediateTensors) {
|
|
this.clonedTensorsMap[nodeName] = this.cloneTensorList(tensors);
|
|
}
|
|
this.checkTensorForDisposal(nodeName, item.node, tensorMap, context, tensorsToKeep, outputNodeNameSet, intermediateTensorConsumerCount);
|
|
this.processChildNodes(item.node, stack2, context, tensorMap, added, usedNodes);
|
|
}
|
|
} else {
|
|
this.processChildNodes(item.node, stack2, context, tensorMap, added, usedNodes);
|
|
}
|
|
}
|
|
return promises;
|
|
}
|
|
processChildNodes(node, stack2, context, tensorMap, added, usedNodes) {
|
|
node.children.forEach((childNode) => {
|
|
const [nodeName] = getNodeNameAndIndex(childNode.name, context);
|
|
if (added[nodeName] || !usedNodes.has(childNode.name)) {
|
|
return;
|
|
}
|
|
if (childNode.op === "Merge") {
|
|
if (childNode.inputNames.some((name) => {
|
|
return !!getTensor(name, tensorMap, context);
|
|
})) {
|
|
added[nodeName] = true;
|
|
stack2.push({ contexts: context.currentContext, node: childNode });
|
|
}
|
|
} else if (childNode.inputNames.every((name) => {
|
|
return !!getTensor(name, tensorMap, context);
|
|
})) {
|
|
added[nodeName] = true;
|
|
stack2.push({ contexts: context.currentContext, node: childNode });
|
|
}
|
|
});
|
|
}
|
|
/**
|
|
* Releases the memory used by the weight tensors.
|
|
*/
|
|
dispose() {
|
|
Object.keys(this.weightMap).forEach((key) => this.weightMap[key].forEach((tensor2) => tensor2.dispose()));
|
|
}
|
|
checkInputShapeAndType(inputs) {
|
|
Object.keys(inputs).forEach((name) => {
|
|
const input = inputs[name];
|
|
const [nodeName] = parseNodeName(name);
|
|
const node = this.graph.nodes[nodeName];
|
|
if (node.attrParams["shape"] && node.attrParams["shape"].value) {
|
|
const shape = node.attrParams["shape"].value;
|
|
const match = shape.length === input.shape.length && input.shape.every((dim, index) => shape[index] === -1 || shape[index] === dim);
|
|
assert(match, () => `The shape of dict['${node.name}'] provided in model.execute(dict) must be [${shape}], but was [${input.shape}]`);
|
|
}
|
|
if (node.attrParams["dtype"] && node.attrParams["dtype"].value) {
|
|
assert(input.dtype === node.attrParams["dtype"].value, () => `The dtype of dict['${node.name}'] provided in model.execute(dict) must be ${node.attrParams["dtype"].value}, but was ${input.dtype}`);
|
|
}
|
|
});
|
|
}
|
|
mapInputs(inputs) {
|
|
var _a, _b;
|
|
const result = {};
|
|
for (const inputName in inputs) {
|
|
const tensor2 = (_b = (_a = this._signature) === null || _a === void 0 ? void 0 : _a.inputs) === null || _b === void 0 ? void 0 : _b[inputName];
|
|
if (tensor2 != null) {
|
|
result[tensor2.name] = inputs[inputName];
|
|
} else {
|
|
result[inputName] = inputs[inputName];
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
checkInputs(inputs) {
|
|
const notInGraph = Object.keys(inputs).filter((name) => {
|
|
const [nodeName] = parseNodeName(name);
|
|
return this.graph.nodes[nodeName] == null;
|
|
});
|
|
if (notInGraph.length > 0) {
|
|
throw new Error(`The dict provided in model.execute(dict) has keys: [${notInGraph}] that are not part of graph`);
|
|
}
|
|
}
|
|
mapOutputs(outputs) {
|
|
return outputs.map((name) => {
|
|
var _a, _b;
|
|
const tensor2 = (_b = (_a = this._signature) === null || _a === void 0 ? void 0 : _a.outputs) === null || _b === void 0 ? void 0 : _b[name];
|
|
if (tensor2 != null) {
|
|
return tensor2.name;
|
|
}
|
|
return name;
|
|
}, {});
|
|
}
|
|
checkOutputs(outputs) {
|
|
outputs.forEach((name) => {
|
|
const [normalizedName] = parseNodeName(name);
|
|
if (!this.graph.nodes[normalizedName]) {
|
|
throw new Error(`The output '${name}' is not found in the graph`);
|
|
}
|
|
});
|
|
}
|
|
}
|
|
class ResourceManager {
|
|
constructor(hashTableNameToHandle = {}, hashTableMap = {}) {
|
|
this.hashTableNameToHandle = hashTableNameToHandle;
|
|
this.hashTableMap = hashTableMap;
|
|
}
|
|
/**
|
|
* Register a `HashTable` in the resource manager.
|
|
*
|
|
* The `HashTable` can be retrieved by `resourceManager.getHashTableById`,
|
|
* where id is the table handle tensor's id.
|
|
*
|
|
* @param name Op node name that creates the `HashTable`.
|
|
* @param hashTable The `HashTable` to be added to resource manager.
|
|
*/
|
|
addHashTable(name, hashTable2) {
|
|
this.hashTableNameToHandle[name] = hashTable2.handle;
|
|
this.hashTableMap[hashTable2.id] = hashTable2;
|
|
}
|
|
/**
|
|
* Get the table handle by node name.
|
|
* @param name Op node name that creates the `HashTable`. This name is also
|
|
* used in the inputs list of lookup and import `HashTable` ops.
|
|
*/
|
|
getHashTableHandleByName(name) {
|
|
return this.hashTableNameToHandle[name];
|
|
}
|
|
/**
|
|
* Get the actual `HashTable` by its handle tensor's id.
|
|
* @param id The id of the handle tensor.
|
|
*/
|
|
getHashTableById(id) {
|
|
return this.hashTableMap[id];
|
|
}
|
|
/**
|
|
* Dispose `ResourceManager`, including its hashTables and tensors in them.
|
|
*/
|
|
dispose() {
|
|
for (const key in this.hashTableMap) {
|
|
this.hashTableMap[key].clearAndClose();
|
|
delete this.hashTableMap[key];
|
|
}
|
|
for (const name in this.hashTableNameToHandle) {
|
|
this.hashTableNameToHandle[name].dispose();
|
|
delete this.hashTableNameToHandle[name];
|
|
}
|
|
}
|
|
}
|
|
const TFHUB_SEARCH_PARAM = "?tfjs-format=file";
|
|
const DEFAULT_MODEL_NAME = "model.json";
|
|
class GraphModel {
|
|
// Returns the version information for the tensorflow model GraphDef.
|
|
get modelVersion() {
|
|
return this.version;
|
|
}
|
|
get inputNodes() {
|
|
return this.executor.inputNodes;
|
|
}
|
|
get outputNodes() {
|
|
return this.executor.outputNodes;
|
|
}
|
|
get inputs() {
|
|
return this.executor.inputs;
|
|
}
|
|
get outputs() {
|
|
return this.executor.outputs;
|
|
}
|
|
get weights() {
|
|
return this.executor.weightMap;
|
|
}
|
|
get metadata() {
|
|
return this.artifacts.userDefinedMetadata;
|
|
}
|
|
get modelSignature() {
|
|
return this.signature;
|
|
}
|
|
get modelStructuredOutputKeys() {
|
|
return this.structuredOutputKeys;
|
|
}
|
|
/**
|
|
* @param modelUrl url for the model, or an `io.IOHandler`.
|
|
* @param weightManifestUrl url for the weight file generated by
|
|
* scripts/convert.py script.
|
|
* @param requestOption options for Request, which allows to send credentials
|
|
* and custom headers.
|
|
* @param onProgress Optional, progress callback function, fired periodically
|
|
* before the load is completed.
|
|
*/
|
|
constructor(modelUrl, loadOptions = {}, tfio = io) {
|
|
this.modelUrl = modelUrl;
|
|
this.loadOptions = loadOptions;
|
|
this.version = "n/a";
|
|
this.io = tfio;
|
|
if (loadOptions == null) {
|
|
this.loadOptions = {};
|
|
}
|
|
this.resourceManager = new ResourceManager();
|
|
}
|
|
findIOHandler() {
|
|
const path = this.modelUrl;
|
|
if (path.load != null) {
|
|
this.handler = path;
|
|
} else if (this.loadOptions.requestInit != null) {
|
|
this.handler = this.io.browserHTTPRequest(path, this.loadOptions);
|
|
} else {
|
|
const handlers = this.io.getLoadHandlers(path, this.loadOptions);
|
|
if (handlers.length === 0) {
|
|
handlers.push(this.io.browserHTTPRequest(path, this.loadOptions));
|
|
} else if (handlers.length > 1) {
|
|
throw new Error(`Found more than one (${handlers.length}) load handlers for URL '${[path]}'`);
|
|
}
|
|
this.handler = handlers[0];
|
|
}
|
|
}
|
|
/**
|
|
* Loads the model and weight files, construct the in memory weight map and
|
|
* compile the inference graph.
|
|
*/
|
|
load() {
|
|
this.findIOHandler();
|
|
if (this.handler.load == null) {
|
|
throw new Error("Cannot proceed with model loading because the IOHandler provided does not have the `load` method implemented.");
|
|
}
|
|
const loadResult = this.handler.load();
|
|
if (isPromise(loadResult)) {
|
|
return loadResult.then((artifacts) => {
|
|
if (artifacts.getWeightStream == null) {
|
|
return this.loadSync(artifacts);
|
|
}
|
|
return this.loadStreaming(artifacts);
|
|
});
|
|
}
|
|
return this.loadSync(loadResult);
|
|
}
|
|
/**
|
|
* Synchronously construct the in memory weight map and
|
|
* compile the inference graph.
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes', ignoreCI: true}
|
|
*/
|
|
loadSync(artifacts) {
|
|
const weightMap = this.io.decodeWeights(artifacts.weightData, artifacts.weightSpecs);
|
|
return this.loadWithWeightMap(artifacts, weightMap);
|
|
}
|
|
async loadStreaming(artifacts) {
|
|
if (artifacts.getWeightStream == null) {
|
|
throw new Error("Model artifacts missing streamWeights function");
|
|
}
|
|
const weightMap = await decodeWeightsStream(artifacts.getWeightStream(), artifacts.weightSpecs);
|
|
return this.loadWithWeightMap(artifacts, weightMap);
|
|
}
|
|
loadWithWeightMap(artifacts, weightMap) {
|
|
this.artifacts = artifacts;
|
|
const graph2 = this.artifacts.modelTopology;
|
|
let signature = this.artifacts.signature;
|
|
if (this.artifacts.userDefinedMetadata != null) {
|
|
const metadata = this.artifacts.userDefinedMetadata;
|
|
if (metadata.signature != null) {
|
|
signature = metadata.signature;
|
|
}
|
|
if (metadata.structuredOutputKeys != null) {
|
|
this.structuredOutputKeys = metadata.structuredOutputKeys;
|
|
}
|
|
}
|
|
this.signature = signature;
|
|
this.version = `${graph2.versions.producer}.${graph2.versions.minConsumer}`;
|
|
this.executor = new GraphExecutor(OperationMapper.Instance.transformGraph(graph2, this.signature));
|
|
this.executor.weightMap = this.convertTensorMapToTensorsMap(weightMap);
|
|
this.executor.resourceManager = this.resourceManager;
|
|
if (artifacts.modelInitializer != null && artifacts.modelInitializer.node != null) {
|
|
const initializer = OperationMapper.Instance.transformGraph(artifacts.modelInitializer);
|
|
this.initializer = new GraphExecutor(initializer);
|
|
this.initializer.weightMap = this.executor.weightMap;
|
|
this.initializer.resourceManager = this.resourceManager;
|
|
this.initializerSignature = artifacts.initializerSignature;
|
|
}
|
|
return true;
|
|
}
|
|
/**
|
|
* Save the configuration and/or weights of the GraphModel.
|
|
*
|
|
* An `IOHandler` is an object that has a `save` method of the proper
|
|
* signature defined. The `save` method manages the storing or
|
|
* transmission of serialized data ("artifacts") that represent the
|
|
* model's topology and weights onto or via a specific medium, such as
|
|
* file downloads, local storage, IndexedDB in the web browser and HTTP
|
|
* requests to a server. TensorFlow.js provides `IOHandler`
|
|
* implementations for a number of frequently used saving mediums, such as
|
|
* `tf.io.browserDownloads` and `tf.io.browserLocalStorage`. See `tf.io`
|
|
* for more details.
|
|
*
|
|
* This method also allows you to refer to certain types of `IOHandler`s
|
|
* as URL-like string shortcuts, such as 'localstorage://' and
|
|
* 'indexeddb://'.
|
|
*
|
|
* Example 1: Save `model`'s topology and weights to browser [local
|
|
* storage](https://developer.mozilla.org/en-US/docs/Web/API/Window/localStorage);
|
|
* then load it back.
|
|
*
|
|
* ```js
|
|
* const modelUrl =
|
|
* 'https://storage.googleapis.com/tfjs-models/savedmodel/mobilenet_v2_1.0_224/model.json';
|
|
* const model = await tf.loadGraphModel(modelUrl);
|
|
* const zeros = tf.zeros([1, 224, 224, 3]);
|
|
* model.predict(zeros).print();
|
|
*
|
|
* const saveResults = await model.save('localstorage://my-model-1');
|
|
*
|
|
* const loadedModel = await tf.loadGraphModel('localstorage://my-model-1');
|
|
* console.log('Prediction from loaded model:');
|
|
* model.predict(zeros).print();
|
|
* ```
|
|
*
|
|
* @param handlerOrURL An instance of `IOHandler` or a URL-like,
|
|
* scheme-based string shortcut for `IOHandler`.
|
|
* @param config Options for saving the model.
|
|
* @returns A `Promise` of `SaveResult`, which summarizes the result of
|
|
* the saving, such as byte sizes of the saved artifacts for the model's
|
|
* topology and weight values.
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes', ignoreCI: true}
|
|
*/
|
|
async save(handlerOrURL, config) {
|
|
if (typeof handlerOrURL === "string") {
|
|
const handlers = this.io.getSaveHandlers(handlerOrURL);
|
|
if (handlers.length === 0) {
|
|
throw new Error(`Cannot find any save handlers for URL '${handlerOrURL}'`);
|
|
} else if (handlers.length > 1) {
|
|
throw new Error(`Found more than one (${handlers.length}) save handlers for URL '${handlerOrURL}'`);
|
|
}
|
|
handlerOrURL = handlers[0];
|
|
}
|
|
if (handlerOrURL.save == null) {
|
|
throw new Error("GraphModel.save() cannot proceed because the IOHandler provided does not have the `save` attribute defined.");
|
|
}
|
|
return handlerOrURL.save(this.artifacts);
|
|
}
|
|
addStructuredOutputNames(outputTensors) {
|
|
if (this.structuredOutputKeys) {
|
|
const outputTensorsArray = outputTensors instanceof Tensor ? [outputTensors] : outputTensors;
|
|
const outputTensorMap = {};
|
|
outputTensorsArray.forEach((outputTensor, i) => outputTensorMap[this.structuredOutputKeys[i]] = outputTensor);
|
|
return outputTensorMap;
|
|
}
|
|
return outputTensors;
|
|
}
|
|
/**
|
|
* Execute the inference for the input tensors.
|
|
*
|
|
* @param input The input tensors, when there is single input for the model,
|
|
* inputs param should be a `tf.Tensor`. For models with multiple inputs,
|
|
* inputs params should be in either `tf.Tensor`[] if the input order is
|
|
* fixed, or otherwise NamedTensorMap format.
|
|
*
|
|
* For model with multiple inputs, we recommend you use NamedTensorMap as the
|
|
* input type, if you use `tf.Tensor`[], the order of the array needs to
|
|
* follow the
|
|
* order of inputNodes array. @see {@link GraphModel.inputNodes}
|
|
*
|
|
* You can also feed any intermediate nodes using the NamedTensorMap as the
|
|
* input type. For example, given the graph
|
|
* InputNode => Intermediate => OutputNode,
|
|
* you can execute the subgraph Intermediate => OutputNode by calling
|
|
* model.execute('IntermediateNode' : tf.tensor(...));
|
|
*
|
|
* This is useful for models that uses tf.dynamic_rnn, where the intermediate
|
|
* state needs to be fed manually.
|
|
*
|
|
* For batch inference execution, the tensors for each input need to be
|
|
* concatenated together. For example with mobilenet, the required input shape
|
|
* is [1, 244, 244, 3], which represents the [batch, height, width, channel].
|
|
* If we are provide a batched data of 100 images, the input tensor should be
|
|
* in the shape of [100, 244, 244, 3].
|
|
*
|
|
* @param config Prediction configuration for specifying the batch size.
|
|
* Currently the batch size option is ignored for graph model.
|
|
*
|
|
* @returns Inference result tensors. If the model is converted and it
|
|
* originally had structured_outputs in tensorflow, then a NamedTensorMap
|
|
* will be returned matching the structured_outputs. If no structured_outputs
|
|
* are present, the output will be single `tf.Tensor` if the model has single
|
|
* output node, otherwise Tensor[].
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
predict(inputs, config) {
|
|
const outputTensors = this.execute(inputs, this.outputNodes);
|
|
return this.addStructuredOutputNames(outputTensors);
|
|
}
|
|
/**
|
|
* Execute the inference for the input tensors in async fashion, use this
|
|
* method when your model contains control flow ops.
|
|
*
|
|
* @param input The input tensors, when there is single input for the model,
|
|
* inputs param should be a `tf.Tensor`. For models with mutliple inputs,
|
|
* inputs params should be in either `tf.Tensor`[] if the input order is
|
|
* fixed, or otherwise NamedTensorMap format.
|
|
*
|
|
* For model with multiple inputs, we recommend you use NamedTensorMap as the
|
|
* input type, if you use `tf.Tensor`[], the order of the array needs to
|
|
* follow the
|
|
* order of inputNodes array. @see {@link GraphModel.inputNodes}
|
|
*
|
|
* You can also feed any intermediate nodes using the NamedTensorMap as the
|
|
* input type. For example, given the graph
|
|
* InputNode => Intermediate => OutputNode,
|
|
* you can execute the subgraph Intermediate => OutputNode by calling
|
|
* model.execute('IntermediateNode' : tf.tensor(...));
|
|
*
|
|
* This is useful for models that uses tf.dynamic_rnn, where the intermediate
|
|
* state needs to be fed manually.
|
|
*
|
|
* For batch inference execution, the tensors for each input need to be
|
|
* concatenated together. For example with mobilenet, the required input shape
|
|
* is [1, 244, 244, 3], which represents the [batch, height, width, channel].
|
|
* If we are provide a batched data of 100 images, the input tensor should be
|
|
* in the shape of [100, 244, 244, 3].
|
|
*
|
|
* @param config Prediction configuration for specifying the batch size.
|
|
* Currently the batch size option is ignored for graph model.
|
|
*
|
|
* @returns A Promise of inference result tensors. If the model is converted
|
|
* and it originally had structured_outputs in tensorflow, then a
|
|
* NamedTensorMap will be returned matching the structured_outputs. If no
|
|
* structured_outputs are present, the output will be single `tf.Tensor` if
|
|
* the model has single output node, otherwise Tensor[].
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
async predictAsync(inputs, config) {
|
|
const outputTensors = await this.executeAsync(inputs, this.outputNodes);
|
|
return this.addStructuredOutputNames(outputTensors);
|
|
}
|
|
normalizeInputs(inputs) {
|
|
var _a;
|
|
if (!(inputs instanceof Tensor) && !Array.isArray(inputs)) {
|
|
const signatureInputs = (_a = this.signature) === null || _a === void 0 ? void 0 : _a.inputs;
|
|
if (signatureInputs != null) {
|
|
for (const input in signatureInputs) {
|
|
const tensor2 = signatureInputs[input];
|
|
if (tensor2.resourceId != null) {
|
|
inputs[input] = this.resourceIdToCapturedInput[tensor2.resourceId];
|
|
}
|
|
}
|
|
}
|
|
return inputs;
|
|
}
|
|
inputs = Array.isArray(inputs) ? inputs : [inputs];
|
|
const numCapturedInputs = Object.keys(this.resourceIdToCapturedInput).length;
|
|
if (inputs.length + numCapturedInputs !== this.inputNodes.length) {
|
|
throw new Error(`Input tensor count mismatch, the graph model has ${this.inputNodes.length - numCapturedInputs} non-resource placeholders, while there are ${inputs.length} input tensors provided.`);
|
|
}
|
|
let inputIndex = 0;
|
|
return this.inputNodes.reduce((map, inputName) => {
|
|
var _a2, _b, _c;
|
|
const resourceId = (_c = (_b = (_a2 = this.signature) === null || _a2 === void 0 ? void 0 : _a2.inputs) === null || _b === void 0 ? void 0 : _b[inputName]) === null || _c === void 0 ? void 0 : _c.resourceId;
|
|
if (resourceId != null) {
|
|
map[inputName] = this.resourceIdToCapturedInput[resourceId];
|
|
} else {
|
|
map[inputName] = inputs[inputIndex++];
|
|
}
|
|
return map;
|
|
}, {});
|
|
}
|
|
normalizeOutputs(outputs) {
|
|
outputs = outputs || this.outputNodes;
|
|
return !Array.isArray(outputs) ? [outputs] : outputs;
|
|
}
|
|
executeInitializerGraph() {
|
|
if (this.initializer == null) {
|
|
return [];
|
|
}
|
|
if (this.initializerSignature == null) {
|
|
return this.initializer.execute({}, []);
|
|
} else {
|
|
return this.initializer.execute({}, Object.keys(this.initializerSignature.outputs));
|
|
}
|
|
}
|
|
async executeInitializerGraphAsync() {
|
|
if (this.initializer == null) {
|
|
return [];
|
|
}
|
|
if (this.initializerSignature == null) {
|
|
return this.initializer.executeAsync({}, []);
|
|
} else {
|
|
return this.initializer.executeAsync({}, Object.keys(this.initializerSignature.outputs));
|
|
}
|
|
}
|
|
setResourceIdToCapturedInput(outputs) {
|
|
this.resourceIdToCapturedInput = {};
|
|
if (this.initializerSignature) {
|
|
const signatureOutputs = this.initializerSignature.outputs;
|
|
const outputNames = Object.keys(signatureOutputs);
|
|
for (let i = 0; i < outputNames.length; i++) {
|
|
const outputName = outputNames[i];
|
|
const tensorInfo = signatureOutputs[outputName];
|
|
this.resourceIdToCapturedInput[tensorInfo.resourceId] = outputs[i];
|
|
}
|
|
}
|
|
}
|
|
/**
|
|
* Executes inference for the model for given input tensors.
|
|
* @param inputs tensor, tensor array or tensor map of the inputs for the
|
|
* model, keyed by the input node names.
|
|
* @param outputs output node name from the TensorFlow model, if no
|
|
* outputs are specified, the default outputs of the model would be used.
|
|
* You can inspect intermediate nodes of the model by adding them to the
|
|
* outputs array.
|
|
*
|
|
* @returns A single tensor if provided with a single output or no outputs
|
|
* are provided and there is only one default output, otherwise return a
|
|
* tensor array. The order of the tensor array is the same as the outputs
|
|
* if provided, otherwise the order of outputNodes attribute of the model.
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
execute(inputs, outputs) {
|
|
if (this.resourceIdToCapturedInput == null) {
|
|
this.setResourceIdToCapturedInput(this.executeInitializerGraph());
|
|
}
|
|
inputs = this.normalizeInputs(inputs);
|
|
outputs = this.normalizeOutputs(outputs);
|
|
const result = this.executor.execute(inputs, outputs);
|
|
return result.length > 1 ? result : result[0];
|
|
}
|
|
/**
|
|
* Executes inference for the model for given input tensors in async
|
|
* fashion, use this method when your model contains control flow ops.
|
|
* @param inputs tensor, tensor array or tensor map of the inputs for the
|
|
* model, keyed by the input node names.
|
|
* @param outputs output node name from the TensorFlow model, if no outputs
|
|
* are specified, the default outputs of the model would be used. You can
|
|
* inspect intermediate nodes of the model by adding them to the outputs
|
|
* array.
|
|
*
|
|
* @returns A Promise of single tensor if provided with a single output or
|
|
* no outputs are provided and there is only one default output, otherwise
|
|
* return a tensor map.
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
async executeAsync(inputs, outputs) {
|
|
if (this.resourceIdToCapturedInput == null) {
|
|
this.setResourceIdToCapturedInput(await this.executeInitializerGraphAsync());
|
|
}
|
|
inputs = this.normalizeInputs(inputs);
|
|
outputs = this.normalizeOutputs(outputs);
|
|
const result = await this.executor.executeAsync(inputs, outputs);
|
|
return result.length > 1 ? result : result[0];
|
|
}
|
|
/**
|
|
* Get intermediate tensors for model debugging mode (flag
|
|
* KEEP_INTERMEDIATE_TENSORS is true).
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
getIntermediateTensors() {
|
|
return this.executor.getIntermediateTensors();
|
|
}
|
|
/**
|
|
* Dispose intermediate tensors for model debugging mode (flag
|
|
* KEEP_INTERMEDIATE_TENSORS is true).
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
disposeIntermediateTensors() {
|
|
this.executor.disposeIntermediateTensors();
|
|
}
|
|
convertTensorMapToTensorsMap(map) {
|
|
return Object.keys(map).reduce((newMap, key) => {
|
|
newMap[key] = [map[key]];
|
|
return newMap;
|
|
}, {});
|
|
}
|
|
/**
|
|
* Releases the memory used by the weight tensors and resourceManager.
|
|
*
|
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
*/
|
|
dispose() {
|
|
this.executor.dispose();
|
|
if (this.initializer) {
|
|
this.initializer.dispose();
|
|
if (this.resourceIdToCapturedInput) {
|
|
dispose(this.resourceIdToCapturedInput);
|
|
}
|
|
}
|
|
this.resourceManager.dispose();
|
|
}
|
|
}
|
|
async function loadGraphModel(modelUrl, options = {}, tfio = io) {
|
|
if (modelUrl == null) {
|
|
throw new Error("modelUrl in loadGraphModel() cannot be null. Please provide a url or an IOHandler that loads the model");
|
|
}
|
|
if (options == null) {
|
|
options = {};
|
|
}
|
|
if (options.fromTFHub && typeof modelUrl === "string") {
|
|
modelUrl = getTFHubUrl(modelUrl);
|
|
}
|
|
const model = new GraphModel(modelUrl, options, tfio);
|
|
await model.load();
|
|
return model;
|
|
}
|
|
function getTFHubUrl(modelUrl) {
|
|
if (!modelUrl.endsWith("/")) {
|
|
modelUrl = modelUrl + "/";
|
|
}
|
|
return `${modelUrl}${DEFAULT_MODEL_NAME}${TFHUB_SEARCH_PARAM}`;
|
|
}
|
|
var selfie_segmentation = {};
|
|
var hasRequiredSelfie_segmentation;
|
|
function requireSelfie_segmentation() {
|
|
if (hasRequiredSelfie_segmentation) return selfie_segmentation;
|
|
hasRequiredSelfie_segmentation = 1;
|
|
(function() {
|
|
var x;
|
|
function aa(a) {
|
|
var b = 0;
|
|
return function() {
|
|
return b < a.length ? { done: false, value: a[b++] } : { done: true };
|
|
};
|
|
}
|
|
var ba = "function" == typeof Object.defineProperties ? Object.defineProperty : function(a, b, c) {
|
|
if (a == Array.prototype || a == Object.prototype) return a;
|
|
a[b] = c.value;
|
|
return a;
|
|
};
|
|
function ca(a) {
|
|
a = ["object" == typeof globalThis && globalThis, a, "object" == typeof window && window, "object" == typeof self && self, "object" == typeof commonjsGlobal && commonjsGlobal];
|
|
for (var b = 0; b < a.length; ++b) {
|
|
var c = a[b];
|
|
if (c && c.Math == Math) return c;
|
|
}
|
|
throw Error("Cannot find global object");
|
|
}
|
|
var y = ca(this);
|
|
function z2(a, b) {
|
|
if (b) a: {
|
|
var c = y;
|
|
a = a.split(".");
|
|
for (var d = 0; d < a.length - 1; d++) {
|
|
var e = a[d];
|
|
if (!(e in c)) break a;
|
|
c = c[e];
|
|
}
|
|
a = a[a.length - 1];
|
|
d = c[a];
|
|
b = b(d);
|
|
b != d && null != b && ba(c, a, { configurable: true, writable: true, value: b });
|
|
}
|
|
}
|
|
z2("Symbol", function(a) {
|
|
function b(g) {
|
|
if (this instanceof b) throw new TypeError("Symbol is not a constructor");
|
|
return new c(d + (g || "") + "_" + e++, g);
|
|
}
|
|
function c(g, f) {
|
|
this.h = g;
|
|
ba(this, "description", { configurable: true, writable: true, value: f });
|
|
}
|
|
if (a) return a;
|
|
c.prototype.toString = function() {
|
|
return this.h;
|
|
};
|
|
var d = "jscomp_symbol_" + (1e9 * Math.random() >>> 0) + "_", e = 0;
|
|
return b;
|
|
});
|
|
z2("Symbol.iterator", function(a) {
|
|
if (a) return a;
|
|
a = /* @__PURE__ */ Symbol("Symbol.iterator");
|
|
for (var b = "Array Int8Array Uint8Array Uint8ClampedArray Int16Array Uint16Array Int32Array Uint32Array Float32Array Float64Array".split(" "), c = 0; c < b.length; c++) {
|
|
var d = y[b[c]];
|
|
"function" === typeof d && "function" != typeof d.prototype[a] && ba(d.prototype, a, { configurable: true, writable: true, value: function() {
|
|
return da(aa(this));
|
|
} });
|
|
}
|
|
return a;
|
|
});
|
|
function da(a) {
|
|
a = { next: a };
|
|
a[Symbol.iterator] = function() {
|
|
return this;
|
|
};
|
|
return a;
|
|
}
|
|
function A(a) {
|
|
var b = "undefined" != typeof Symbol && Symbol.iterator && a[Symbol.iterator];
|
|
return b ? b.call(a) : { next: aa(a) };
|
|
}
|
|
function ea(a) {
|
|
if (!(a instanceof Array)) {
|
|
a = A(a);
|
|
for (var b, c = []; !(b = a.next()).done; ) c.push(b.value);
|
|
a = c;
|
|
}
|
|
return a;
|
|
}
|
|
var fa = "function" == typeof Object.assign ? Object.assign : function(a, b) {
|
|
for (var c = 1; c < arguments.length; c++) {
|
|
var d = arguments[c];
|
|
if (d) for (var e in d) Object.prototype.hasOwnProperty.call(d, e) && (a[e] = d[e]);
|
|
}
|
|
return a;
|
|
};
|
|
z2("Object.assign", function(a) {
|
|
return a || fa;
|
|
});
|
|
var ha = "function" == typeof Object.create ? Object.create : function(a) {
|
|
function b() {
|
|
}
|
|
b.prototype = a;
|
|
return new b();
|
|
}, ia;
|
|
if ("function" == typeof Object.setPrototypeOf) ia = Object.setPrototypeOf;
|
|
else {
|
|
var ja;
|
|
a: {
|
|
var ka = { a: true }, la = {};
|
|
try {
|
|
la.__proto__ = ka;
|
|
ja = la.a;
|
|
break a;
|
|
} catch (a) {
|
|
}
|
|
ja = false;
|
|
}
|
|
ia = ja ? function(a, b) {
|
|
a.__proto__ = b;
|
|
if (a.__proto__ !== b) throw new TypeError(a + " is not extensible");
|
|
return a;
|
|
} : null;
|
|
}
|
|
var ma = ia;
|
|
function na(a, b) {
|
|
a.prototype = ha(b.prototype);
|
|
a.prototype.constructor = a;
|
|
if (ma) ma(a, b);
|
|
else for (var c in b) if ("prototype" != c) if (Object.defineProperties) {
|
|
var d = Object.getOwnPropertyDescriptor(b, c);
|
|
d && Object.defineProperty(a, c, d);
|
|
} else a[c] = b[c];
|
|
a.za = b.prototype;
|
|
}
|
|
function oa() {
|
|
this.m = false;
|
|
this.j = null;
|
|
this.i = void 0;
|
|
this.h = 1;
|
|
this.v = this.s = 0;
|
|
this.l = null;
|
|
}
|
|
function pa(a) {
|
|
if (a.m) throw new TypeError("Generator is already running");
|
|
a.m = true;
|
|
}
|
|
oa.prototype.u = function(a) {
|
|
this.i = a;
|
|
};
|
|
function qa(a, b) {
|
|
a.l = { ma: b, na: true };
|
|
a.h = a.s || a.v;
|
|
}
|
|
oa.prototype.return = function(a) {
|
|
this.l = { return: a };
|
|
this.h = this.v;
|
|
};
|
|
function D2(a, b, c) {
|
|
a.h = c;
|
|
return { value: b };
|
|
}
|
|
function ra(a) {
|
|
this.h = new oa();
|
|
this.i = a;
|
|
}
|
|
function sa(a, b) {
|
|
pa(a.h);
|
|
var c = a.h.j;
|
|
if (c) return ta(a, "return" in c ? c["return"] : function(d) {
|
|
return { value: d, done: true };
|
|
}, b, a.h.return);
|
|
a.h.return(b);
|
|
return ua(a);
|
|
}
|
|
function ta(a, b, c, d) {
|
|
try {
|
|
var e = b.call(a.h.j, c);
|
|
if (!(e instanceof Object)) throw new TypeError("Iterator result " + e + " is not an object");
|
|
if (!e.done) return a.h.m = false, e;
|
|
var g = e.value;
|
|
} catch (f) {
|
|
return a.h.j = null, qa(a.h, f), ua(a);
|
|
}
|
|
a.h.j = null;
|
|
d.call(a.h, g);
|
|
return ua(a);
|
|
}
|
|
function ua(a) {
|
|
for (; a.h.h; ) try {
|
|
var b = a.i(a.h);
|
|
if (b) return a.h.m = false, { value: b.value, done: false };
|
|
} catch (c) {
|
|
a.h.i = void 0, qa(a.h, c);
|
|
}
|
|
a.h.m = false;
|
|
if (a.h.l) {
|
|
b = a.h.l;
|
|
a.h.l = null;
|
|
if (b.na) throw b.ma;
|
|
return { value: b.return, done: true };
|
|
}
|
|
return { value: void 0, done: true };
|
|
}
|
|
function va(a) {
|
|
this.next = function(b) {
|
|
pa(a.h);
|
|
a.h.j ? b = ta(a, a.h.j.next, b, a.h.u) : (a.h.u(b), b = ua(a));
|
|
return b;
|
|
};
|
|
this.throw = function(b) {
|
|
pa(a.h);
|
|
a.h.j ? b = ta(a, a.h.j["throw"], b, a.h.u) : (qa(a.h, b), b = ua(a));
|
|
return b;
|
|
};
|
|
this.return = function(b) {
|
|
return sa(a, b);
|
|
};
|
|
this[Symbol.iterator] = function() {
|
|
return this;
|
|
};
|
|
}
|
|
function wa(a) {
|
|
function b(d) {
|
|
return a.next(d);
|
|
}
|
|
function c(d) {
|
|
return a.throw(d);
|
|
}
|
|
return new Promise(function(d, e) {
|
|
function g(f) {
|
|
f.done ? d(f.value) : Promise.resolve(f.value).then(b, c).then(g, e);
|
|
}
|
|
g(a.next());
|
|
});
|
|
}
|
|
function E(a) {
|
|
return wa(new va(new ra(a)));
|
|
}
|
|
z2("Promise", function(a) {
|
|
function b(f) {
|
|
this.i = 0;
|
|
this.j = void 0;
|
|
this.h = [];
|
|
this.u = false;
|
|
var h = this.l();
|
|
try {
|
|
f(h.resolve, h.reject);
|
|
} catch (k) {
|
|
h.reject(k);
|
|
}
|
|
}
|
|
function c() {
|
|
this.h = null;
|
|
}
|
|
function d(f) {
|
|
return f instanceof b ? f : new b(function(h) {
|
|
h(f);
|
|
});
|
|
}
|
|
if (a) return a;
|
|
c.prototype.i = function(f) {
|
|
if (null == this.h) {
|
|
this.h = [];
|
|
var h = this;
|
|
this.j(function() {
|
|
h.m();
|
|
});
|
|
}
|
|
this.h.push(f);
|
|
};
|
|
var e = y.setTimeout;
|
|
c.prototype.j = function(f) {
|
|
e(f, 0);
|
|
};
|
|
c.prototype.m = function() {
|
|
for (; this.h && this.h.length; ) {
|
|
var f = this.h;
|
|
this.h = [];
|
|
for (var h = 0; h < f.length; ++h) {
|
|
var k = f[h];
|
|
f[h] = null;
|
|
try {
|
|
k();
|
|
} catch (l) {
|
|
this.l(l);
|
|
}
|
|
}
|
|
}
|
|
this.h = null;
|
|
};
|
|
c.prototype.l = function(f) {
|
|
this.j(function() {
|
|
throw f;
|
|
});
|
|
};
|
|
b.prototype.l = function() {
|
|
function f(l) {
|
|
return function(m) {
|
|
k || (k = true, l.call(h, m));
|
|
};
|
|
}
|
|
var h = this, k = false;
|
|
return { resolve: f(this.I), reject: f(this.m) };
|
|
};
|
|
b.prototype.I = function(f) {
|
|
if (f === this) this.m(new TypeError("A Promise cannot resolve to itself"));
|
|
else if (f instanceof b) this.L(f);
|
|
else {
|
|
a: switch (typeof f) {
|
|
case "object":
|
|
var h = null != f;
|
|
break a;
|
|
case "function":
|
|
h = true;
|
|
break a;
|
|
default:
|
|
h = false;
|
|
}
|
|
h ? this.F(f) : this.s(f);
|
|
}
|
|
};
|
|
b.prototype.F = function(f) {
|
|
var h = void 0;
|
|
try {
|
|
h = f.then;
|
|
} catch (k) {
|
|
this.m(k);
|
|
return;
|
|
}
|
|
"function" == typeof h ? this.M(h, f) : this.s(f);
|
|
};
|
|
b.prototype.m = function(f) {
|
|
this.v(2, f);
|
|
};
|
|
b.prototype.s = function(f) {
|
|
this.v(1, f);
|
|
};
|
|
b.prototype.v = function(f, h) {
|
|
if (0 != this.i) throw Error("Cannot settle(" + f + ", " + h + "): Promise already settled in state" + this.i);
|
|
this.i = f;
|
|
this.j = h;
|
|
2 === this.i && this.K();
|
|
this.H();
|
|
};
|
|
b.prototype.K = function() {
|
|
var f = this;
|
|
e(function() {
|
|
if (f.D()) {
|
|
var h = y.console;
|
|
"undefined" !== typeof h && h.error(f.j);
|
|
}
|
|
}, 1);
|
|
};
|
|
b.prototype.D = function() {
|
|
if (this.u) return false;
|
|
var f = y.CustomEvent, h = y.Event, k = y.dispatchEvent;
|
|
if ("undefined" === typeof k) return true;
|
|
"function" === typeof f ? f = new f("unhandledrejection", { cancelable: true }) : "function" === typeof h ? f = new h("unhandledrejection", { cancelable: true }) : (f = y.document.createEvent("CustomEvent"), f.initCustomEvent("unhandledrejection", false, true, f));
|
|
f.promise = this;
|
|
f.reason = this.j;
|
|
return k(f);
|
|
};
|
|
b.prototype.H = function() {
|
|
if (null != this.h) {
|
|
for (var f = 0; f < this.h.length; ++f) g.i(this.h[f]);
|
|
this.h = null;
|
|
}
|
|
};
|
|
var g = new c();
|
|
b.prototype.L = function(f) {
|
|
var h = this.l();
|
|
f.T(h.resolve, h.reject);
|
|
};
|
|
b.prototype.M = function(f, h) {
|
|
var k = this.l();
|
|
try {
|
|
f.call(h, k.resolve, k.reject);
|
|
} catch (l) {
|
|
k.reject(l);
|
|
}
|
|
};
|
|
b.prototype.then = function(f, h) {
|
|
function k(p2, n) {
|
|
return "function" == typeof p2 ? function(q2) {
|
|
try {
|
|
l(p2(q2));
|
|
} catch (t2) {
|
|
m(t2);
|
|
}
|
|
} : n;
|
|
}
|
|
var l, m, r = new b(function(p2, n) {
|
|
l = p2;
|
|
m = n;
|
|
});
|
|
this.T(k(f, l), k(h, m));
|
|
return r;
|
|
};
|
|
b.prototype.catch = function(f) {
|
|
return this.then(void 0, f);
|
|
};
|
|
b.prototype.T = function(f, h) {
|
|
function k() {
|
|
switch (l.i) {
|
|
case 1:
|
|
f(l.j);
|
|
break;
|
|
case 2:
|
|
h(l.j);
|
|
break;
|
|
default:
|
|
throw Error("Unexpected state: " + l.i);
|
|
}
|
|
}
|
|
var l = this;
|
|
null == this.h ? g.i(k) : this.h.push(k);
|
|
this.u = true;
|
|
};
|
|
b.resolve = d;
|
|
b.reject = function(f) {
|
|
return new b(function(h, k) {
|
|
k(f);
|
|
});
|
|
};
|
|
b.race = function(f) {
|
|
return new b(function(h, k) {
|
|
for (var l = A(f), m = l.next(); !m.done; m = l.next()) d(m.value).T(h, k);
|
|
});
|
|
};
|
|
b.all = function(f) {
|
|
var h = A(f), k = h.next();
|
|
return k.done ? d([]) : new b(function(l, m) {
|
|
function r(q2) {
|
|
return function(t2) {
|
|
p2[q2] = t2;
|
|
n--;
|
|
0 == n && l(p2);
|
|
};
|
|
}
|
|
var p2 = [], n = 0;
|
|
do
|
|
p2.push(void 0), n++, d(k.value).T(r(p2.length - 1), m), k = h.next();
|
|
while (!k.done);
|
|
});
|
|
};
|
|
return b;
|
|
});
|
|
function xa(a, b) {
|
|
a instanceof String && (a += "");
|
|
var c = 0, d = false, e = { next: function() {
|
|
if (!d && c < a.length) {
|
|
var g = c++;
|
|
return { value: b(g, a[g]), done: false };
|
|
}
|
|
d = true;
|
|
return { done: true, value: void 0 };
|
|
} };
|
|
e[Symbol.iterator] = function() {
|
|
return e;
|
|
};
|
|
return e;
|
|
}
|
|
z2("Array.prototype.keys", function(a) {
|
|
return a ? a : function() {
|
|
return xa(this, function(b) {
|
|
return b;
|
|
});
|
|
};
|
|
});
|
|
z2("Array.prototype.fill", function(a) {
|
|
return a ? a : function(b, c, d) {
|
|
var e = this.length || 0;
|
|
0 > c && (c = Math.max(0, e + c));
|
|
if (null == d || d > e) d = e;
|
|
d = Number(d);
|
|
0 > d && (d = Math.max(0, e + d));
|
|
for (c = Number(c || 0); c < d; c++) this[c] = b;
|
|
return this;
|
|
};
|
|
});
|
|
function F2(a) {
|
|
return a ? a : Array.prototype.fill;
|
|
}
|
|
z2("Int8Array.prototype.fill", F2);
|
|
z2("Uint8Array.prototype.fill", F2);
|
|
z2("Uint8ClampedArray.prototype.fill", F2);
|
|
z2("Int16Array.prototype.fill", F2);
|
|
z2("Uint16Array.prototype.fill", F2);
|
|
z2("Int32Array.prototype.fill", F2);
|
|
z2("Uint32Array.prototype.fill", F2);
|
|
z2("Float32Array.prototype.fill", F2);
|
|
z2("Float64Array.prototype.fill", F2);
|
|
z2("Object.is", function(a) {
|
|
return a ? a : function(b, c) {
|
|
return b === c ? 0 !== b || 1 / b === 1 / c : b !== b && c !== c;
|
|
};
|
|
});
|
|
z2("Array.prototype.includes", function(a) {
|
|
return a ? a : function(b, c) {
|
|
var d = this;
|
|
d instanceof String && (d = String(d));
|
|
var e = d.length;
|
|
c = c || 0;
|
|
for (0 > c && (c = Math.max(c + e, 0)); c < e; c++) {
|
|
var g = d[c];
|
|
if (g === b || Object.is(g, b)) return true;
|
|
}
|
|
return false;
|
|
};
|
|
});
|
|
z2("String.prototype.includes", function(a) {
|
|
return a ? a : function(b, c) {
|
|
if (null == this) throw new TypeError("The 'this' value for String.prototype.includes must not be null or undefined");
|
|
if (b instanceof RegExp) throw new TypeError("First argument to String.prototype.includes must not be a regular expression");
|
|
return -1 !== this.indexOf(b, c || 0);
|
|
};
|
|
});
|
|
var ya = this || self;
|
|
function Aa(a, b) {
|
|
a = a.split(".");
|
|
var c = ya;
|
|
a[0] in c || "undefined" == typeof c.execScript || c.execScript("var " + a[0]);
|
|
for (var d; a.length && (d = a.shift()); ) a.length || void 0 === b ? c[d] && c[d] !== Object.prototype[d] ? c = c[d] : c = c[d] = {} : c[d] = b;
|
|
}
|
|
function Ba(a) {
|
|
var b;
|
|
a: {
|
|
if (b = ya.navigator) {
|
|
if (b = b.userAgent) break a;
|
|
}
|
|
b = "";
|
|
}
|
|
return -1 != b.indexOf(a);
|
|
}
|
|
var Ca = Array.prototype.map ? function(a, b) {
|
|
return Array.prototype.map.call(a, b, void 0);
|
|
} : function(a, b) {
|
|
for (var c = a.length, d = Array(c), e = "string" === typeof a ? a.split("") : a, g = 0; g < c; g++) g in e && (d[g] = b.call(void 0, e[g], g, a));
|
|
return d;
|
|
};
|
|
var Da = {}, Ea = null;
|
|
function Fa(a) {
|
|
var b = a.length, c = 3 * b / 4;
|
|
c % 3 ? c = Math.floor(c) : -1 != "=.".indexOf(a[b - 1]) && (c = -1 != "=.".indexOf(a[b - 2]) ? c - 2 : c - 1);
|
|
var d = new Uint8Array(c), e = 0;
|
|
Ga(a, function(g) {
|
|
d[e++] = g;
|
|
});
|
|
return e !== c ? d.subarray(0, e) : d;
|
|
}
|
|
function Ga(a, b) {
|
|
function c(k) {
|
|
for (; d < a.length; ) {
|
|
var l = a.charAt(d++), m = Ea[l];
|
|
if (null != m) return m;
|
|
if (!/^[\s\xa0]*$/.test(l)) throw Error("Unknown base64 encoding at char: " + l);
|
|
}
|
|
return k;
|
|
}
|
|
Ha();
|
|
for (var d = 0; ; ) {
|
|
var e = c(-1), g = c(0), f = c(64), h = c(64);
|
|
if (64 === h && -1 === e) break;
|
|
b(e << 2 | g >> 4);
|
|
64 != f && (b(g << 4 & 240 | f >> 2), 64 != h && b(f << 6 & 192 | h));
|
|
}
|
|
}
|
|
function Ha() {
|
|
if (!Ea) {
|
|
Ea = {};
|
|
for (var a = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789".split(""), b = ["+/=", "+/", "-_=", "-_.", "-_"], c = 0; 5 > c; c++) {
|
|
var d = a.concat(b[c].split(""));
|
|
Da[c] = d;
|
|
for (var e = 0; e < d.length; e++) {
|
|
var g = d[e];
|
|
void 0 === Ea[g] && (Ea[g] = e);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
var Ia = "undefined" !== typeof Uint8Array, Ja = !(Ba("Trident") || Ba("MSIE")) && "function" === typeof ya.btoa;
|
|
function Ka(a) {
|
|
if (!Ja) {
|
|
var b;
|
|
void 0 === b && (b = 0);
|
|
Ha();
|
|
b = Da[b];
|
|
for (var c = Array(Math.floor(a.length / 3)), d = b[64] || "", e = 0, g = 0; e < a.length - 2; e += 3) {
|
|
var f = a[e], h = a[e + 1], k = a[e + 2], l = b[f >> 2];
|
|
f = b[(f & 3) << 4 | h >> 4];
|
|
h = b[(h & 15) << 2 | k >> 6];
|
|
k = b[k & 63];
|
|
c[g++] = l + f + h + k;
|
|
}
|
|
l = 0;
|
|
k = d;
|
|
switch (a.length - e) {
|
|
case 2:
|
|
l = a[e + 1], k = b[(l & 15) << 2] || d;
|
|
case 1:
|
|
a = a[e], c[g] = b[a >> 2] + b[(a & 3) << 4 | l >> 4] + k + d;
|
|
}
|
|
return c.join("");
|
|
}
|
|
for (b = ""; 10240 < a.length; ) b += String.fromCharCode.apply(null, a.subarray(0, 10240)), a = a.subarray(10240);
|
|
b += String.fromCharCode.apply(
|
|
null,
|
|
a
|
|
);
|
|
return btoa(b);
|
|
}
|
|
var La = RegExp("[-_.]", "g");
|
|
function Ma(a) {
|
|
switch (a) {
|
|
case "-":
|
|
return "+";
|
|
case "_":
|
|
return "/";
|
|
case ".":
|
|
return "=";
|
|
default:
|
|
return "";
|
|
}
|
|
}
|
|
function Na(a) {
|
|
if (!Ja) return Fa(a);
|
|
La.test(a) && (a = a.replace(La, Ma));
|
|
a = atob(a);
|
|
for (var b = new Uint8Array(a.length), c = 0; c < a.length; c++) b[c] = a.charCodeAt(c);
|
|
return b;
|
|
}
|
|
var Oa;
|
|
function Pa() {
|
|
return Oa || (Oa = new Uint8Array(0));
|
|
}
|
|
var Qa = {};
|
|
var Ra = "function" === typeof Uint8Array.prototype.slice, G2 = 0, H = 0;
|
|
function Sa(a) {
|
|
var b = 0 > a;
|
|
a = Math.abs(a);
|
|
var c = a >>> 0;
|
|
a = Math.floor((a - c) / 4294967296);
|
|
b && (c = A(Ta(c, a)), b = c.next().value, a = c.next().value, c = b);
|
|
G2 = c >>> 0;
|
|
H = a >>> 0;
|
|
}
|
|
var Ua = "function" === typeof BigInt;
|
|
function Ta(a, b) {
|
|
b = ~b;
|
|
a ? a = ~a + 1 : b += 1;
|
|
return [a, b];
|
|
}
|
|
function Va(a, b) {
|
|
this.i = a >>> 0;
|
|
this.h = b >>> 0;
|
|
}
|
|
function Wa(a) {
|
|
if (!a) return Xa || (Xa = new Va(0, 0));
|
|
if (!/^-?\d+$/.test(a)) return null;
|
|
if (16 > a.length) Sa(Number(a));
|
|
else if (Ua) a = BigInt(a), G2 = Number(a & BigInt(4294967295)) >>> 0, H = Number(a >> BigInt(32) & BigInt(4294967295));
|
|
else {
|
|
var b = +("-" === a[0]);
|
|
H = G2 = 0;
|
|
for (var c = a.length, d = b, e = (c - b) % 6 + b; e <= c; d = e, e += 6) d = Number(a.slice(d, e)), H *= 1e6, G2 = 1e6 * G2 + d, 4294967296 <= G2 && (H += G2 / 4294967296 | 0, G2 %= 4294967296);
|
|
b && (b = A(Ta(G2, H)), a = b.next().value, b = b.next().value, G2 = a, H = b);
|
|
}
|
|
return new Va(G2, H);
|
|
}
|
|
var Xa;
|
|
function Ya(a, b) {
|
|
return Error("Invalid wire type: " + a + " (at position " + b + ")");
|
|
}
|
|
function Za() {
|
|
return Error("Failed to read varint, encoding is invalid.");
|
|
}
|
|
function $a(a, b) {
|
|
return Error("Tried to read past the end of the data " + b + " > " + a);
|
|
}
|
|
function K() {
|
|
throw Error("Invalid UTF8");
|
|
}
|
|
function ab(a, b) {
|
|
b = String.fromCharCode.apply(null, b);
|
|
return null == a ? b : a + b;
|
|
}
|
|
var bb = void 0, cb, db = "undefined" !== typeof TextDecoder, eb, fb = "undefined" !== typeof TextEncoder;
|
|
var gb;
|
|
function hb(a) {
|
|
if (a !== Qa) throw Error("illegal external caller");
|
|
}
|
|
function ib(a, b) {
|
|
hb(b);
|
|
this.V = a;
|
|
if (null != a && 0 === a.length) throw Error("ByteString should be constructed with non-empty values");
|
|
}
|
|
function jb() {
|
|
return gb || (gb = new ib(null, Qa));
|
|
}
|
|
function kb(a) {
|
|
hb(Qa);
|
|
var b = a.V;
|
|
b = null == b || Ia && null != b && b instanceof Uint8Array ? b : "string" === typeof b ? Na(b) : null;
|
|
return null == b ? b : a.V = b;
|
|
}
|
|
function lb(a) {
|
|
if ("string" === typeof a) return { buffer: Na(a), C: false };
|
|
if (Array.isArray(a)) return { buffer: new Uint8Array(a), C: false };
|
|
if (a.constructor === Uint8Array) return { buffer: a, C: false };
|
|
if (a.constructor === ArrayBuffer) return { buffer: new Uint8Array(a), C: false };
|
|
if (a.constructor === ib) return { buffer: kb(a) || Pa(), C: true };
|
|
if (a instanceof Uint8Array) return { buffer: new Uint8Array(a.buffer, a.byteOffset, a.byteLength), C: false };
|
|
throw Error("Type not convertible to a Uint8Array, expected a Uint8Array, an ArrayBuffer, a base64 encoded string, a ByteString or an Array of numbers");
|
|
}
|
|
function mb(a, b) {
|
|
this.i = null;
|
|
this.m = false;
|
|
this.h = this.j = this.l = 0;
|
|
nb(this, a, b);
|
|
}
|
|
function nb(a, b, c) {
|
|
c = void 0 === c ? {} : c;
|
|
a.S = void 0 === c.S ? false : c.S;
|
|
b && (b = lb(b), a.i = b.buffer, a.m = b.C, a.l = 0, a.j = a.i.length, a.h = a.l);
|
|
}
|
|
mb.prototype.reset = function() {
|
|
this.h = this.l;
|
|
};
|
|
function L2(a, b) {
|
|
a.h = b;
|
|
if (b > a.j) throw $a(a.j, b);
|
|
}
|
|
function ob(a) {
|
|
var b = a.i, c = a.h, d = b[c++], e = d & 127;
|
|
if (d & 128 && (d = b[c++], e |= (d & 127) << 7, d & 128 && (d = b[c++], e |= (d & 127) << 14, d & 128 && (d = b[c++], e |= (d & 127) << 21, d & 128 && (d = b[c++], e |= d << 28, d & 128 && b[c++] & 128 && b[c++] & 128 && b[c++] & 128 && b[c++] & 128 && b[c++] & 128))))) throw Za();
|
|
L2(a, c);
|
|
return e;
|
|
}
|
|
function pb(a, b) {
|
|
if (0 > b) throw Error("Tried to read a negative byte length: " + b);
|
|
var c = a.h, d = c + b;
|
|
if (d > a.j) throw $a(b, a.j - c);
|
|
a.h = d;
|
|
return c;
|
|
}
|
|
var qb = [];
|
|
function rb() {
|
|
this.h = [];
|
|
}
|
|
rb.prototype.length = function() {
|
|
return this.h.length;
|
|
};
|
|
rb.prototype.end = function() {
|
|
var a = this.h;
|
|
this.h = [];
|
|
return a;
|
|
};
|
|
function sb(a, b, c) {
|
|
for (; 0 < c || 127 < b; ) a.h.push(b & 127 | 128), b = (b >>> 7 | c << 25) >>> 0, c >>>= 7;
|
|
a.h.push(b);
|
|
}
|
|
function M(a, b) {
|
|
for (; 127 < b; ) a.h.push(b & 127 | 128), b >>>= 7;
|
|
a.h.push(b);
|
|
}
|
|
function tb(a, b) {
|
|
if (qb.length) {
|
|
var c = qb.pop();
|
|
nb(c, a, b);
|
|
a = c;
|
|
} else a = new mb(a, b);
|
|
this.h = a;
|
|
this.j = this.h.h;
|
|
this.i = this.l = -1;
|
|
this.setOptions(b);
|
|
}
|
|
tb.prototype.setOptions = function(a) {
|
|
a = void 0 === a ? {} : a;
|
|
this.ca = void 0 === a.ca ? false : a.ca;
|
|
};
|
|
tb.prototype.reset = function() {
|
|
this.h.reset();
|
|
this.j = this.h.h;
|
|
this.i = this.l = -1;
|
|
};
|
|
function ub(a) {
|
|
var b = a.h;
|
|
if (b.h == b.j) return false;
|
|
a.j = a.h.h;
|
|
var c = ob(a.h) >>> 0;
|
|
b = c >>> 3;
|
|
c &= 7;
|
|
if (!(0 <= c && 5 >= c)) throw Ya(c, a.j);
|
|
if (1 > b) throw Error("Invalid field number: " + b + " (at position " + a.j + ")");
|
|
a.l = b;
|
|
a.i = c;
|
|
return true;
|
|
}
|
|
function vb(a) {
|
|
switch (a.i) {
|
|
case 0:
|
|
if (0 != a.i) vb(a);
|
|
else a: {
|
|
a = a.h;
|
|
for (var b = a.h, c = b + 10, d = a.i; b < c; ) if (0 === (d[b++] & 128)) {
|
|
L2(a, b);
|
|
break a;
|
|
}
|
|
throw Za();
|
|
}
|
|
break;
|
|
case 1:
|
|
a = a.h;
|
|
L2(a, a.h + 8);
|
|
break;
|
|
case 2:
|
|
2 != a.i ? vb(a) : (b = ob(a.h) >>> 0, a = a.h, L2(a, a.h + b));
|
|
break;
|
|
case 5:
|
|
a = a.h;
|
|
L2(a, a.h + 4);
|
|
break;
|
|
case 3:
|
|
b = a.l;
|
|
do {
|
|
if (!ub(a)) throw Error("Unmatched start-group tag: stream EOF");
|
|
if (4 == a.i) {
|
|
if (a.l != b) throw Error("Unmatched end-group tag");
|
|
break;
|
|
}
|
|
vb(a);
|
|
} while (1);
|
|
break;
|
|
default:
|
|
throw Ya(a.i, a.j);
|
|
}
|
|
}
|
|
var wb = [];
|
|
function xb() {
|
|
this.j = [];
|
|
this.i = 0;
|
|
this.h = new rb();
|
|
}
|
|
function N2(a, b) {
|
|
0 !== b.length && (a.j.push(b), a.i += b.length);
|
|
}
|
|
function yb(a, b) {
|
|
if (b = b.R) {
|
|
N2(a, a.h.end());
|
|
for (var c = 0; c < b.length; c++) N2(a, kb(b[c]) || Pa());
|
|
}
|
|
}
|
|
var O = "function" === typeof Symbol && "symbol" === typeof /* @__PURE__ */ Symbol() ? /* @__PURE__ */ Symbol() : void 0;
|
|
function P(a, b) {
|
|
if (O) return a[O] |= b;
|
|
if (void 0 !== a.A) return a.A |= b;
|
|
Object.defineProperties(a, { A: { value: b, configurable: true, writable: true, enumerable: false } });
|
|
return b;
|
|
}
|
|
function zb(a, b) {
|
|
O ? a[O] && (a[O] &= ~b) : void 0 !== a.A && (a.A &= ~b);
|
|
}
|
|
function Q2(a) {
|
|
var b;
|
|
O ? b = a[O] : b = a.A;
|
|
return null == b ? 0 : b;
|
|
}
|
|
function R2(a, b) {
|
|
O ? a[O] = b : void 0 !== a.A ? a.A = b : Object.defineProperties(a, { A: { value: b, configurable: true, writable: true, enumerable: false } });
|
|
}
|
|
function Ab(a) {
|
|
P(a, 1);
|
|
return a;
|
|
}
|
|
function Bb(a, b) {
|
|
R2(b, (a | 0) & -51);
|
|
}
|
|
function Cb(a, b) {
|
|
R2(b, (a | 18) & -41);
|
|
}
|
|
var Db = {};
|
|
function Eb(a) {
|
|
return null !== a && "object" === typeof a && !Array.isArray(a) && a.constructor === Object;
|
|
}
|
|
var Fb, Gb = [];
|
|
R2(Gb, 23);
|
|
Fb = Object.freeze(Gb);
|
|
function Hb(a) {
|
|
if (Q2(a.o) & 2) throw Error("Cannot mutate an immutable Message");
|
|
}
|
|
function Ib(a) {
|
|
var b = a.length;
|
|
(b = b ? a[b - 1] : void 0) && Eb(b) ? b.g = 1 : (b = {}, a.push((b.g = 1, b)));
|
|
}
|
|
function Jb(a) {
|
|
var b = a.i + a.G;
|
|
return a.B || (a.B = a.o[b] = {});
|
|
}
|
|
function S(a, b) {
|
|
return -1 === b ? null : b >= a.i ? a.B ? a.B[b] : void 0 : a.o[b + a.G];
|
|
}
|
|
function U(a, b, c, d) {
|
|
Hb(a);
|
|
Kb(a, b, c, d);
|
|
}
|
|
function Kb(a, b, c, d) {
|
|
a.j && (a.j = void 0);
|
|
b >= a.i || d ? Jb(a)[b] = c : (a.o[b + a.G] = c, (a = a.B) && b in a && delete a[b]);
|
|
}
|
|
function Lb(a, b, c, d) {
|
|
var e = S(a, b);
|
|
Array.isArray(e) || (e = Fb);
|
|
var g = Q2(e);
|
|
g & 1 || Ab(e);
|
|
if (d) g & 2 || P(e, 2), c & 1 || Object.freeze(e);
|
|
else {
|
|
d = !(c & 2);
|
|
var f = g & 2;
|
|
c & 1 || !f ? d && g & 16 && !f && zb(e, 16) : (e = Ab(Array.prototype.slice.call(e)), Kb(a, b, e));
|
|
}
|
|
return e;
|
|
}
|
|
function Mb(a, b) {
|
|
var c = S(a, b);
|
|
var d = null == c ? c : "number" === typeof c || "NaN" === c || "Infinity" === c || "-Infinity" === c ? Number(c) : void 0;
|
|
null != d && d !== c && Kb(a, b, d);
|
|
return d;
|
|
}
|
|
function Nb(a, b, c, d, e) {
|
|
a.h || (a.h = {});
|
|
var g = a.h[c], f = Lb(a, c, 3, e);
|
|
if (!g) {
|
|
var h = f;
|
|
g = [];
|
|
var k = !!(Q2(a.o) & 16);
|
|
f = !!(Q2(h) & 2);
|
|
var l = h;
|
|
!e && f && (h = Array.prototype.slice.call(h));
|
|
for (var m = f, r = 0; r < h.length; r++) {
|
|
var p2 = h[r];
|
|
var n = b, q2 = false;
|
|
q2 = void 0 === q2 ? false : q2;
|
|
p2 = Array.isArray(p2) ? new n(p2) : q2 ? new n() : void 0;
|
|
if (void 0 !== p2) {
|
|
n = p2.o;
|
|
var t2 = q2 = Q2(n);
|
|
f && (t2 |= 2);
|
|
k && (t2 |= 16);
|
|
t2 != q2 && R2(n, t2);
|
|
n = t2;
|
|
m = m || !!(2 & n);
|
|
g.push(p2);
|
|
}
|
|
}
|
|
a.h[c] = g;
|
|
k = Q2(h);
|
|
b = k | 33;
|
|
b = m ? b & -9 : b | 8;
|
|
k != b && (m = h, Object.isFrozen(m) && (m = Array.prototype.slice.call(m)), R2(m, b), h = m);
|
|
l !== h && Kb(
|
|
a,
|
|
c,
|
|
h
|
|
);
|
|
(e || d && f) && P(g, 2);
|
|
d && Object.freeze(g);
|
|
return g;
|
|
}
|
|
e || (e = Object.isFrozen(g), d && !e ? Object.freeze(g) : !d && e && (g = Array.prototype.slice.call(g), a.h[c] = g));
|
|
return g;
|
|
}
|
|
function Ob(a, b, c) {
|
|
var d = !!(Q2(a.o) & 2);
|
|
b = Nb(a, b, c, d, d);
|
|
a = Lb(a, c, 3, d);
|
|
if (!(d || Q2(a) & 8)) {
|
|
for (d = 0; d < b.length; d++) {
|
|
c = b[d];
|
|
if (Q2(c.o) & 2) {
|
|
var e = Pb(c, false);
|
|
e.j = c;
|
|
} else e = c;
|
|
c !== e && (b[d] = e, a[d] = e.o);
|
|
}
|
|
P(a, 8);
|
|
}
|
|
return b;
|
|
}
|
|
function V2(a, b, c) {
|
|
if (null != c && "number" !== typeof c) throw Error("Value of float/double field must be a number|null|undefined, found " + typeof c + ": " + c);
|
|
U(a, b, c);
|
|
}
|
|
function Qb(a, b, c, d, e) {
|
|
Hb(a);
|
|
var g = Nb(a, c, b, false, false);
|
|
c = null != d ? d : new c();
|
|
a = Lb(a, b, 2, false);
|
|
void 0 != e ? (g.splice(e, 0, c), a.splice(e, 0, c.o)) : (g.push(c), a.push(c.o));
|
|
c.C() && zb(a, 8);
|
|
return c;
|
|
}
|
|
function Rb(a, b) {
|
|
return null == a ? b : a;
|
|
}
|
|
function W2(a, b, c) {
|
|
c = void 0 === c ? 0 : c;
|
|
return Rb(Mb(a, b), c);
|
|
}
|
|
var Sb;
|
|
function Tb(a) {
|
|
switch (typeof a) {
|
|
case "number":
|
|
return isFinite(a) ? a : String(a);
|
|
case "object":
|
|
if (a) if (Array.isArray(a)) {
|
|
if (0 !== (Q2(a) & 128)) return a = Array.prototype.slice.call(a), Ib(a), a;
|
|
} else {
|
|
if (Ia && null != a && a instanceof Uint8Array) return Ka(a);
|
|
if (a instanceof ib) {
|
|
var b = a.V;
|
|
return null == b ? "" : "string" === typeof b ? b : a.V = Ka(b);
|
|
}
|
|
}
|
|
}
|
|
return a;
|
|
}
|
|
function Ub(a, b, c, d) {
|
|
if (null != a) {
|
|
if (Array.isArray(a)) a = Vb(a, b, c, void 0 !== d);
|
|
else if (Eb(a)) {
|
|
var e = {}, g;
|
|
for (g in a) e[g] = Ub(a[g], b, c, d);
|
|
a = e;
|
|
} else a = b(a, d);
|
|
return a;
|
|
}
|
|
}
|
|
function Vb(a, b, c, d) {
|
|
var e = Q2(a);
|
|
d = d ? !!(e & 16) : void 0;
|
|
a = Array.prototype.slice.call(a);
|
|
for (var g = 0; g < a.length; g++) a[g] = Ub(a[g], b, c, d);
|
|
c(e, a);
|
|
return a;
|
|
}
|
|
function Wb(a) {
|
|
return a.ja === Db ? a.toJSON() : Tb(a);
|
|
}
|
|
function Xb(a, b) {
|
|
a & 128 && Ib(b);
|
|
}
|
|
function Yb(a, b, c) {
|
|
c = void 0 === c ? Cb : c;
|
|
if (null != a) {
|
|
if (Ia && a instanceof Uint8Array) return a.length ? new ib(new Uint8Array(a), Qa) : jb();
|
|
if (Array.isArray(a)) {
|
|
var d = Q2(a);
|
|
if (d & 2) return a;
|
|
if (b && !(d & 32) && (d & 16 || 0 === d)) return R2(a, d | 2), a;
|
|
a = Vb(a, Yb, d & 4 ? Cb : c, true);
|
|
b = Q2(a);
|
|
b & 4 && b & 2 && Object.freeze(a);
|
|
return a;
|
|
}
|
|
return a.ja === Db ? Zb(a) : a;
|
|
}
|
|
}
|
|
function $b(a, b, c, d, e, g, f) {
|
|
if (a = a.h && a.h[c]) {
|
|
d = Q2(a);
|
|
d & 2 ? d = a : (g = Ca(a, Zb), Cb(d, g), Object.freeze(g), d = g);
|
|
Hb(b);
|
|
f = null == d ? Fb : Ab([]);
|
|
if (null != d) {
|
|
g = !!d.length;
|
|
for (a = 0; a < d.length; a++) {
|
|
var h = d[a];
|
|
g = g && !(Q2(h.o) & 2);
|
|
f[a] = h.o;
|
|
}
|
|
g = (g ? 8 : 0) | 1;
|
|
a = Q2(f);
|
|
(a & g) !== g && (Object.isFrozen(f) && (f = Array.prototype.slice.call(f)), R2(f, a | g));
|
|
b.h || (b.h = {});
|
|
b.h[c] = d;
|
|
} else b.h && (b.h[c] = void 0);
|
|
Kb(b, c, f, e);
|
|
} else U(b, c, Yb(d, g, f), e);
|
|
}
|
|
function Zb(a) {
|
|
if (Q2(a.o) & 2) return a;
|
|
a = Pb(a, true);
|
|
P(a.o, 2);
|
|
return a;
|
|
}
|
|
function Pb(a, b) {
|
|
var c = a.o, d = [];
|
|
P(d, 16);
|
|
var e = a.constructor.h;
|
|
e && d.push(e);
|
|
e = a.B;
|
|
if (e) {
|
|
d.length = c.length;
|
|
d.fill(void 0, d.length, c.length);
|
|
var g = {};
|
|
d[d.length - 1] = g;
|
|
}
|
|
0 !== (Q2(c) & 128) && Ib(d);
|
|
b = b || a.C() ? Cb : Bb;
|
|
g = a.constructor;
|
|
Sb = d;
|
|
d = new g(d);
|
|
Sb = void 0;
|
|
a.R && (d.R = a.R.slice());
|
|
g = !!(Q2(c) & 16);
|
|
for (var f = e ? c.length - 1 : c.length, h = 0; h < f; h++) $b(a, d, h - a.G, c[h], false, g, b);
|
|
if (e) for (var k in e) $b(a, d, +k, e[k], true, g, b);
|
|
return d;
|
|
}
|
|
function X2(a, b, c) {
|
|
null == a && (a = Sb);
|
|
Sb = void 0;
|
|
var d = this.constructor.i || 0, e = 0 < d, g = this.constructor.h, f = false;
|
|
if (null == a) {
|
|
a = g ? [g] : [];
|
|
var h = 48;
|
|
var k = true;
|
|
e && (d = 0, h |= 128);
|
|
R2(a, h);
|
|
} else {
|
|
if (!Array.isArray(a)) throw Error();
|
|
if (g && g !== a[0]) throw Error();
|
|
var l = h = P(a, 0);
|
|
if (k = 0 !== (16 & l)) (f = 0 !== (32 & l)) || (l |= 32);
|
|
if (e) if (128 & l) d = 0;
|
|
else {
|
|
if (0 < a.length) {
|
|
var m = a[a.length - 1];
|
|
if (Eb(m) && "g" in m) {
|
|
d = 0;
|
|
l |= 128;
|
|
delete m.g;
|
|
var r = true, p2;
|
|
for (p2 in m) {
|
|
r = false;
|
|
break;
|
|
}
|
|
r && a.pop();
|
|
}
|
|
}
|
|
}
|
|
else if (128 & l) throw Error();
|
|
h !== l && R2(a, l);
|
|
}
|
|
this.G = (g ? 0 : -1) - d;
|
|
this.h = void 0;
|
|
this.o = a;
|
|
a: {
|
|
g = this.o.length;
|
|
d = g - 1;
|
|
if (g && (g = this.o[d], Eb(g))) {
|
|
this.B = g;
|
|
this.i = d - this.G;
|
|
break a;
|
|
}
|
|
void 0 !== b && -1 < b ? (this.i = Math.max(b, d + 1 - this.G), this.B = void 0) : this.i = Number.MAX_VALUE;
|
|
}
|
|
if (!e && this.B && "g" in this.B) throw Error('Unexpected "g" flag in sparse object of message that is not a group type.');
|
|
if (c) {
|
|
b = k && !f && true;
|
|
e = this.i;
|
|
var n;
|
|
for (k = 0; k < c.length; k++) f = c[k], f < e ? (f += this.G, (d = a[f]) ? ac(d, b) : a[f] = Fb) : (n || (n = Jb(this)), (d = n[f]) ? ac(d, b) : n[f] = Fb);
|
|
}
|
|
}
|
|
X2.prototype.toJSON = function() {
|
|
return Vb(this.o, Wb, Xb);
|
|
};
|
|
X2.prototype.C = function() {
|
|
return !!(Q2(this.o) & 2);
|
|
};
|
|
function ac(a, b) {
|
|
if (Array.isArray(a)) {
|
|
var c = Q2(a), d = 1;
|
|
!b || c & 2 || (d |= 16);
|
|
(c & d) !== d && R2(a, c | d);
|
|
}
|
|
}
|
|
X2.prototype.ja = Db;
|
|
X2.prototype.toString = function() {
|
|
return this.o.toString();
|
|
};
|
|
function bc(a, b, c) {
|
|
if (c) {
|
|
var d = {}, e;
|
|
for (e in c) {
|
|
var g = c[e], f = g.ra;
|
|
f || (d.J = g.xa || g.oa.W, g.ia ? (d.aa = cc(g.ia), f = /* @__PURE__ */ (function(h) {
|
|
return function(k, l, m) {
|
|
return h.J(k, l, m, h.aa);
|
|
};
|
|
})(d)) : g.ka ? (d.Z = dc(g.da.P, g.ka), f = /* @__PURE__ */ (function(h) {
|
|
return function(k, l, m) {
|
|
return h.J(k, l, m, h.Z);
|
|
};
|
|
})(d)) : f = d.J, g.ra = f);
|
|
f(b, a, g.da);
|
|
d = { J: d.J, aa: d.aa, Z: d.Z };
|
|
}
|
|
}
|
|
yb(b, a);
|
|
}
|
|
var ec = /* @__PURE__ */ Symbol();
|
|
function fc(a, b, c) {
|
|
return a[ec] || (a[ec] = function(d, e) {
|
|
return b(d, e, c);
|
|
});
|
|
}
|
|
function gc(a) {
|
|
var b = a[ec];
|
|
if (!b) {
|
|
var c = hc(a);
|
|
b = function(d, e) {
|
|
return ic(d, e, c);
|
|
};
|
|
a[ec] = b;
|
|
}
|
|
return b;
|
|
}
|
|
function jc(a) {
|
|
var b = a.ia;
|
|
if (b) return gc(b);
|
|
if (b = a.wa) return fc(a.da.P, b, a.ka);
|
|
}
|
|
function kc(a) {
|
|
var b = jc(a), c = a.da, d = a.oa.U;
|
|
return b ? function(e, g) {
|
|
return d(e, g, c, b);
|
|
} : function(e, g) {
|
|
return d(e, g, c);
|
|
};
|
|
}
|
|
function lc(a, b) {
|
|
var c = a[b];
|
|
"function" == typeof c && 0 === c.length && (c = c(), a[b] = c);
|
|
return Array.isArray(c) && (mc in c || nc in c || 0 < c.length && "function" == typeof c[0]) ? c : void 0;
|
|
}
|
|
function oc(a, b, c, d, e, g) {
|
|
b.P = a[0];
|
|
var f = 1;
|
|
if (a.length > f && "number" !== typeof a[f]) {
|
|
var h = a[f++];
|
|
c(b, h);
|
|
}
|
|
for (; f < a.length; ) {
|
|
c = a[f++];
|
|
for (var k = f + 1; k < a.length && "number" !== typeof a[k]; ) k++;
|
|
h = a[f++];
|
|
k -= f;
|
|
switch (k) {
|
|
case 0:
|
|
d(b, c, h);
|
|
break;
|
|
case 1:
|
|
(k = lc(a, f)) ? (f++, e(b, c, h, k)) : d(b, c, h, a[f++]);
|
|
break;
|
|
case 2:
|
|
k = f++;
|
|
k = lc(a, k);
|
|
e(b, c, h, k, a[f++]);
|
|
break;
|
|
case 3:
|
|
g(b, c, h, a[f++], a[f++], a[f++]);
|
|
break;
|
|
case 4:
|
|
g(b, c, h, a[f++], a[f++], a[f++], a[f++]);
|
|
break;
|
|
default:
|
|
throw Error("unexpected number of binary field arguments: " + k);
|
|
}
|
|
}
|
|
return b;
|
|
}
|
|
var pc = /* @__PURE__ */ Symbol();
|
|
function cc(a) {
|
|
var b = a[pc];
|
|
if (!b) {
|
|
var c = qc(a);
|
|
b = function(d, e) {
|
|
return rc(d, e, c);
|
|
};
|
|
a[pc] = b;
|
|
}
|
|
return b;
|
|
}
|
|
function dc(a, b) {
|
|
var c = a[pc];
|
|
c || (c = function(d, e) {
|
|
return bc(d, e, b);
|
|
}, a[pc] = c);
|
|
return c;
|
|
}
|
|
var nc = /* @__PURE__ */ Symbol();
|
|
function sc(a, b) {
|
|
a.push(b);
|
|
}
|
|
function tc(a, b, c) {
|
|
a.push(b, c.W);
|
|
}
|
|
function uc(a, b, c, d) {
|
|
var e = cc(d), g = qc(d).P, f = c.W;
|
|
a.push(b, function(h, k, l) {
|
|
return f(h, k, l, g, e);
|
|
});
|
|
}
|
|
function vc(a, b, c, d, e, g) {
|
|
var f = dc(d, g), h = c.W;
|
|
a.push(b, function(k, l, m) {
|
|
return h(k, l, m, d, f);
|
|
});
|
|
}
|
|
function qc(a) {
|
|
var b = a[nc];
|
|
if (b) return b;
|
|
b = oc(a, a[nc] = [], sc, tc, uc, vc);
|
|
mc in a && nc in a && (a.length = 0);
|
|
return b;
|
|
}
|
|
var mc = /* @__PURE__ */ Symbol();
|
|
function wc(a, b) {
|
|
a[0] = b;
|
|
}
|
|
function xc(a, b, c, d) {
|
|
var e = c.U;
|
|
a[b] = d ? function(g, f, h) {
|
|
return e(g, f, h, d);
|
|
} : e;
|
|
}
|
|
function yc(a, b, c, d, e) {
|
|
var g = c.U, f = gc(d), h = hc(d).P;
|
|
a[b] = function(k, l, m) {
|
|
return g(k, l, m, h, f, e);
|
|
};
|
|
}
|
|
function zc(a, b, c, d, e, g, f) {
|
|
var h = c.U, k = fc(d, e, g);
|
|
a[b] = function(l, m, r) {
|
|
return h(l, m, r, d, k, f);
|
|
};
|
|
}
|
|
function hc(a) {
|
|
var b = a[mc];
|
|
if (b) return b;
|
|
b = oc(a, a[mc] = {}, wc, xc, yc, zc);
|
|
mc in a && nc in a && (a.length = 0);
|
|
return b;
|
|
}
|
|
function ic(a, b, c) {
|
|
for (; ub(b) && 4 != b.i; ) {
|
|
var d = b.l, e = c[d];
|
|
if (!e) {
|
|
var g = c[0];
|
|
g && (g = g[d]) && (e = c[d] = kc(g));
|
|
}
|
|
if (!e || !e(b, a, d)) {
|
|
e = b;
|
|
d = a;
|
|
g = e.j;
|
|
vb(e);
|
|
var f = e;
|
|
if (!f.ca) {
|
|
e = f.h.h - g;
|
|
f.h.h = g;
|
|
f = f.h;
|
|
if (0 == e) e = jb();
|
|
else {
|
|
g = pb(f, e);
|
|
if (f.S && f.m) e = f.i.subarray(g, g + e);
|
|
else {
|
|
f = f.i;
|
|
var h = g;
|
|
e = g + e;
|
|
e = h === e ? Pa() : Ra ? f.slice(h, e) : new Uint8Array(f.subarray(h, e));
|
|
}
|
|
e = 0 == e.length ? jb() : new ib(e, Qa);
|
|
}
|
|
(g = d.R) ? g.push(e) : d.R = [e];
|
|
}
|
|
}
|
|
}
|
|
return a;
|
|
}
|
|
function rc(a, b, c) {
|
|
for (var d = c.length, e = 1 == d % 2, g = e ? 1 : 0; g < d; g += 2) (0, c[g + 1])(b, a, c[g]);
|
|
bc(a, b, e ? c[0] : void 0);
|
|
}
|
|
function Ac(a, b) {
|
|
return { U: a, W: b };
|
|
}
|
|
var Y2 = Ac(function(a, b, c) {
|
|
if (5 !== a.i) return false;
|
|
a = a.h;
|
|
var d = a.i, e = a.h, g = d[e];
|
|
var f = d[e + 1];
|
|
var h = d[e + 2];
|
|
d = d[e + 3];
|
|
L2(a, a.h + 4);
|
|
f = (g << 0 | f << 8 | h << 16 | d << 24) >>> 0;
|
|
a = 2 * (f >> 31) + 1;
|
|
g = f >>> 23 & 255;
|
|
f &= 8388607;
|
|
U(b, c, 255 == g ? f ? NaN : Infinity * a : 0 == g ? a * Math.pow(2, -149) * f : a * Math.pow(2, g - 150) * (f + Math.pow(2, 23)));
|
|
return true;
|
|
}, function(a, b, c) {
|
|
b = Mb(b, c);
|
|
if (null != b) {
|
|
M(a.h, 8 * c + 5);
|
|
a = a.h;
|
|
var d = +b;
|
|
0 === d ? 0 < 1 / d ? G2 = H = 0 : (H = 0, G2 = 2147483648) : isNaN(d) ? (H = 0, G2 = 2147483647) : (d = (c = 0 > d ? -2147483648 : 0) ? -d : d, 34028234663852886e22 < d ? (H = 0, G2 = (c | 2139095040) >>> 0) : 11754943508222875e-54 > d ? (d = Math.round(d / Math.pow(2, -149)), H = 0, G2 = (c | d) >>> 0) : (b = Math.floor(Math.log(d) / Math.LN2), d *= Math.pow(2, -b), d = Math.round(8388608 * d), 16777216 <= d && ++b, H = 0, G2 = (c | b + 127 << 23 | d & 8388607) >>> 0));
|
|
c = G2;
|
|
a.h.push(c >>> 0 & 255);
|
|
a.h.push(c >>> 8 & 255);
|
|
a.h.push(c >>> 16 & 255);
|
|
a.h.push(c >>> 24 & 255);
|
|
}
|
|
}), Bc = Ac(function(a, b, c) {
|
|
if (0 !== a.i) return false;
|
|
var d = a.h, e = 0, g = a = 0, f = d.i, h = d.h;
|
|
do {
|
|
var k = f[h++];
|
|
e |= (k & 127) << g;
|
|
g += 7;
|
|
} while (32 > g && k & 128);
|
|
32 < g && (a |= (k & 127) >> 4);
|
|
for (g = 3; 32 > g && k & 128; g += 7) k = f[h++], a |= (k & 127) << g;
|
|
L2(
|
|
d,
|
|
h
|
|
);
|
|
if (128 > k) {
|
|
d = e >>> 0;
|
|
k = a >>> 0;
|
|
if (a = k & 2147483648) d = ~d + 1 >>> 0, k = ~k >>> 0, 0 == d && (k = k + 1 >>> 0);
|
|
d = 4294967296 * k + (d >>> 0);
|
|
} else throw Za();
|
|
U(b, c, a ? -d : d);
|
|
return true;
|
|
}, function(a, b, c) {
|
|
b = S(b, c);
|
|
null != b && ("string" === typeof b && Wa(b), null != b && (M(a.h, 8 * c), "number" === typeof b ? (a = a.h, Sa(b), sb(a, G2, H)) : (c = Wa(b), sb(a.h, c.i, c.h))));
|
|
}), Cc = Ac(function(a, b, c) {
|
|
if (0 !== a.i) return false;
|
|
U(b, c, ob(a.h));
|
|
return true;
|
|
}, function(a, b, c) {
|
|
b = S(b, c);
|
|
if (null != b && null != b) if (M(a.h, 8 * c), a = a.h, c = b, 0 <= c) M(a, c);
|
|
else {
|
|
for (b = 0; 9 > b; b++) a.h.push(c & 127 | 128), c >>= 7;
|
|
a.h.push(1);
|
|
}
|
|
}), Dc = Ac(function(a, b, c) {
|
|
if (2 !== a.i) return false;
|
|
var d = ob(a.h) >>> 0;
|
|
a = a.h;
|
|
var e = pb(a, d);
|
|
a = a.i;
|
|
if (db) {
|
|
var g = a, f;
|
|
(f = cb) || (f = cb = new TextDecoder("utf-8", { fatal: true }));
|
|
a = e + d;
|
|
g = 0 === e && a === g.length ? g : g.subarray(e, a);
|
|
try {
|
|
var h = f.decode(g);
|
|
} catch (r) {
|
|
if (void 0 === bb) {
|
|
try {
|
|
f.decode(new Uint8Array([128]));
|
|
} catch (p2) {
|
|
}
|
|
try {
|
|
f.decode(new Uint8Array([97])), bb = true;
|
|
} catch (p2) {
|
|
bb = false;
|
|
}
|
|
}
|
|
!bb && (cb = void 0);
|
|
throw r;
|
|
}
|
|
} else {
|
|
h = e;
|
|
d = h + d;
|
|
e = [];
|
|
for (var k = null, l, m; h < d; ) l = a[h++], 128 > l ? e.push(l) : 224 > l ? h >= d ? K() : (m = a[h++], 194 > l || 128 !== (m & 192) ? (h--, K()) : e.push((l & 31) << 6 | m & 63)) : 240 > l ? h >= d - 1 ? K() : (m = a[h++], 128 !== (m & 192) || 224 === l && 160 > m || 237 === l && 160 <= m || 128 !== ((g = a[h++]) & 192) ? (h--, K()) : e.push((l & 15) << 12 | (m & 63) << 6 | g & 63)) : 244 >= l ? h >= d - 2 ? K() : (m = a[h++], 128 !== (m & 192) || 0 !== (l << 28) + (m - 144) >> 30 || 128 !== ((g = a[h++]) & 192) || 128 !== ((f = a[h++]) & 192) ? (h--, K()) : (l = (l & 7) << 18 | (m & 63) << 12 | (g & 63) << 6 | f & 63, l -= 65536, e.push((l >> 10 & 1023) + 55296, (l & 1023) + 56320))) : K(), 8192 <= e.length && (k = ab(k, e), e.length = 0);
|
|
h = ab(k, e);
|
|
}
|
|
U(b, c, h);
|
|
return true;
|
|
}, function(a, b, c) {
|
|
b = S(b, c);
|
|
if (null != b) {
|
|
var d = false;
|
|
d = void 0 === d ? false : d;
|
|
if (fb) {
|
|
if (d && /(?:[^\uD800-\uDBFF]|^)[\uDC00-\uDFFF]|[\uD800-\uDBFF](?![\uDC00-\uDFFF])/.test(b)) throw Error("Found an unpaired surrogate");
|
|
b = (eb || (eb = new TextEncoder())).encode(b);
|
|
} else {
|
|
for (var e = 0, g = new Uint8Array(3 * b.length), f = 0; f < b.length; f++) {
|
|
var h = b.charCodeAt(f);
|
|
if (128 > h) g[e++] = h;
|
|
else {
|
|
if (2048 > h) g[e++] = h >> 6 | 192;
|
|
else {
|
|
if (55296 <= h && 57343 >= h) {
|
|
if (56319 >= h && f < b.length) {
|
|
var k = b.charCodeAt(++f);
|
|
if (56320 <= k && 57343 >= k) {
|
|
h = 1024 * (h - 55296) + k - 56320 + 65536;
|
|
g[e++] = h >> 18 | 240;
|
|
g[e++] = h >> 12 & 63 | 128;
|
|
g[e++] = h >> 6 & 63 | 128;
|
|
g[e++] = h & 63 | 128;
|
|
continue;
|
|
} else f--;
|
|
}
|
|
if (d) throw Error("Found an unpaired surrogate");
|
|
h = 65533;
|
|
}
|
|
g[e++] = h >> 12 | 224;
|
|
g[e++] = h >> 6 & 63 | 128;
|
|
}
|
|
g[e++] = h & 63 | 128;
|
|
}
|
|
}
|
|
b = e === g.length ? g : g.subarray(0, e);
|
|
}
|
|
M(a.h, 8 * c + 2);
|
|
M(a.h, b.length);
|
|
N2(a, a.h.end());
|
|
N2(a, b);
|
|
}
|
|
}), Ec = Ac(function(a, b, c, d, e) {
|
|
if (2 !== a.i) return false;
|
|
b = Qb(b, c, d);
|
|
c = a.h.j;
|
|
d = ob(a.h) >>> 0;
|
|
var g = a.h.h + d, f = g - c;
|
|
0 >= f && (a.h.j = g, e(b, a, void 0, void 0, void 0), f = g - a.h.h);
|
|
if (f) throw Error("Message parsing ended unexpectedly. Expected to read " + (d + " bytes, instead read " + (d - f) + " bytes, either the data ended unexpectedly or the message misreported its own length"));
|
|
a.h.h = g;
|
|
a.h.j = c;
|
|
return true;
|
|
}, function(a, b, c, d, e) {
|
|
b = Ob(b, d, c);
|
|
if (null != b) for (d = 0; d < b.length; d++) {
|
|
var g = a;
|
|
M(g.h, 8 * c + 2);
|
|
var f = g.h.end();
|
|
N2(g, f);
|
|
f.push(g.i);
|
|
g = f;
|
|
e(b[d], a);
|
|
f = a;
|
|
var h = g.pop();
|
|
for (h = f.i + f.h.length() - h; 127 < h; ) g.push(h & 127 | 128), h >>>= 7, f.i++;
|
|
g.push(h);
|
|
f.i++;
|
|
}
|
|
});
|
|
function Fc(a) {
|
|
return function(b, c) {
|
|
a: {
|
|
if (wb.length) {
|
|
var d = wb.pop();
|
|
d.setOptions(c);
|
|
nb(d.h, b, c);
|
|
b = d;
|
|
} else b = new tb(b, c);
|
|
try {
|
|
var e = hc(a);
|
|
var g = ic(new e.P(), b, e);
|
|
break a;
|
|
} finally {
|
|
e = b.h, e.i = null, e.m = false, e.l = 0, e.j = 0, e.h = 0, e.S = false, b.l = -1, b.i = -1, 100 > wb.length && wb.push(b);
|
|
}
|
|
g = void 0;
|
|
}
|
|
return g;
|
|
};
|
|
}
|
|
function Gc(a) {
|
|
return function() {
|
|
var b = new xb();
|
|
rc(this, b, qc(a));
|
|
N2(b, b.h.end());
|
|
for (var c = new Uint8Array(b.i), d = b.j, e = d.length, g = 0, f = 0; f < e; f++) {
|
|
var h = d[f];
|
|
c.set(h, g);
|
|
g += h.length;
|
|
}
|
|
b.j = [c];
|
|
return c;
|
|
};
|
|
}
|
|
function Z2(a) {
|
|
X2.call(this, a);
|
|
}
|
|
na(Z2, X2);
|
|
var Hc = [Z2, 1, Cc, 2, Y2, 3, Dc, 4, Dc];
|
|
Z2.prototype.l = Gc(Hc);
|
|
function Ic(a) {
|
|
X2.call(this, a, -1, Jc);
|
|
}
|
|
na(Ic, X2);
|
|
Ic.prototype.addClassification = function(a, b) {
|
|
Qb(this, 1, Z2, a, b);
|
|
return this;
|
|
};
|
|
var Jc = [1], Kc = Fc([Ic, 1, Ec, Hc]);
|
|
function Lc(a) {
|
|
X2.call(this, a);
|
|
}
|
|
na(Lc, X2);
|
|
var Mc = [Lc, 1, Y2, 2, Y2, 3, Y2, 4, Y2, 5, Y2];
|
|
Lc.prototype.l = Gc(Mc);
|
|
function Nc(a) {
|
|
X2.call(this, a, -1, Oc);
|
|
}
|
|
na(Nc, X2);
|
|
var Oc = [1], Pc = Fc([Nc, 1, Ec, Mc]);
|
|
function Qc(a) {
|
|
X2.call(this, a);
|
|
}
|
|
na(Qc, X2);
|
|
var Rc = [Qc, 1, Y2, 2, Y2, 3, Y2, 4, Y2, 5, Y2, 6, Bc], Sc = Fc(Rc);
|
|
Qc.prototype.l = Gc(Rc);
|
|
function Tc(a, b, c) {
|
|
c = a.createShader(0 === c ? a.VERTEX_SHADER : a.FRAGMENT_SHADER);
|
|
a.shaderSource(c, b);
|
|
a.compileShader(c);
|
|
if (!a.getShaderParameter(c, a.COMPILE_STATUS)) throw Error("Could not compile WebGL shader.\n\n" + a.getShaderInfoLog(c));
|
|
return c;
|
|
}
|
|
function Uc(a) {
|
|
return Ob(a, Z2, 1).map(function(b) {
|
|
var c = S(b, 1);
|
|
return { index: null == c ? 0 : c, qa: W2(b, 2), label: null != S(b, 3) ? Rb(S(b, 3), "") : void 0, displayName: null != S(b, 4) ? Rb(S(b, 4), "") : void 0 };
|
|
});
|
|
}
|
|
function Vc(a) {
|
|
return { x: W2(a, 1), y: W2(a, 2), z: W2(a, 3), visibility: null != Mb(a, 4) ? W2(a, 4) : void 0 };
|
|
}
|
|
function Wc(a, b) {
|
|
this.i = a;
|
|
this.h = b;
|
|
this.m = 0;
|
|
}
|
|
function Xc(a, b, c) {
|
|
Yc(a, b);
|
|
if ("function" === typeof a.h.canvas.transferToImageBitmap) return Promise.resolve(a.h.canvas.transferToImageBitmap());
|
|
if (c) return Promise.resolve(a.h.canvas);
|
|
if ("function" === typeof createImageBitmap) return createImageBitmap(a.h.canvas);
|
|
void 0 === a.j && (a.j = document.createElement("canvas"));
|
|
return new Promise(function(d) {
|
|
a.j.height = a.h.canvas.height;
|
|
a.j.width = a.h.canvas.width;
|
|
a.j.getContext("2d", {}).drawImage(a.h.canvas, 0, 0, a.h.canvas.width, a.h.canvas.height);
|
|
d(a.j);
|
|
});
|
|
}
|
|
function Yc(a, b) {
|
|
var c = a.h;
|
|
if (void 0 === a.s) {
|
|
var d = Tc(c, "\n attribute vec2 aVertex;\n attribute vec2 aTex;\n varying vec2 vTex;\n void main(void) {\n gl_Position = vec4(aVertex, 0.0, 1.0);\n vTex = aTex;\n }", 0), e = Tc(c, "\n precision mediump float;\n varying vec2 vTex;\n uniform sampler2D sampler0;\n void main(){\n gl_FragColor = texture2D(sampler0, vTex);\n }", 1), g = c.createProgram();
|
|
c.attachShader(g, d);
|
|
c.attachShader(g, e);
|
|
c.linkProgram(g);
|
|
if (!c.getProgramParameter(g, c.LINK_STATUS)) throw Error("Could not compile WebGL program.\n\n" + c.getProgramInfoLog(g));
|
|
d = a.s = g;
|
|
c.useProgram(d);
|
|
e = c.getUniformLocation(d, "sampler0");
|
|
a.l = { O: c.getAttribLocation(d, "aVertex"), N: c.getAttribLocation(d, "aTex"), ya: e };
|
|
a.v = c.createBuffer();
|
|
c.bindBuffer(c.ARRAY_BUFFER, a.v);
|
|
c.enableVertexAttribArray(a.l.O);
|
|
c.vertexAttribPointer(a.l.O, 2, c.FLOAT, false, 0, 0);
|
|
c.bufferData(c.ARRAY_BUFFER, new Float32Array([-1, -1, -1, 1, 1, 1, 1, -1]), c.STATIC_DRAW);
|
|
c.bindBuffer(c.ARRAY_BUFFER, null);
|
|
a.u = c.createBuffer();
|
|
c.bindBuffer(c.ARRAY_BUFFER, a.u);
|
|
c.enableVertexAttribArray(a.l.N);
|
|
c.vertexAttribPointer(
|
|
a.l.N,
|
|
2,
|
|
c.FLOAT,
|
|
false,
|
|
0,
|
|
0
|
|
);
|
|
c.bufferData(c.ARRAY_BUFFER, new Float32Array([0, 1, 0, 0, 1, 0, 1, 1]), c.STATIC_DRAW);
|
|
c.bindBuffer(c.ARRAY_BUFFER, null);
|
|
c.uniform1i(e, 0);
|
|
}
|
|
d = a.l;
|
|
c.useProgram(a.s);
|
|
c.canvas.width = b.width;
|
|
c.canvas.height = b.height;
|
|
c.viewport(0, 0, b.width, b.height);
|
|
c.activeTexture(c.TEXTURE0);
|
|
a.i.bindTexture2d(b.glName);
|
|
c.enableVertexAttribArray(d.O);
|
|
c.bindBuffer(c.ARRAY_BUFFER, a.v);
|
|
c.vertexAttribPointer(d.O, 2, c.FLOAT, false, 0, 0);
|
|
c.enableVertexAttribArray(d.N);
|
|
c.bindBuffer(c.ARRAY_BUFFER, a.u);
|
|
c.vertexAttribPointer(
|
|
d.N,
|
|
2,
|
|
c.FLOAT,
|
|
false,
|
|
0,
|
|
0
|
|
);
|
|
c.bindFramebuffer(c.DRAW_FRAMEBUFFER ? c.DRAW_FRAMEBUFFER : c.FRAMEBUFFER, null);
|
|
c.clearColor(0, 0, 0, 0);
|
|
c.clear(c.COLOR_BUFFER_BIT);
|
|
c.colorMask(true, true, true, true);
|
|
c.drawArrays(c.TRIANGLE_FAN, 0, 4);
|
|
c.disableVertexAttribArray(d.O);
|
|
c.disableVertexAttribArray(d.N);
|
|
c.bindBuffer(c.ARRAY_BUFFER, null);
|
|
a.i.bindTexture2d(0);
|
|
}
|
|
function Zc(a) {
|
|
this.h = a;
|
|
}
|
|
var $c = new Uint8Array([0, 97, 115, 109, 1, 0, 0, 0, 1, 4, 1, 96, 0, 0, 3, 2, 1, 0, 10, 9, 1, 7, 0, 65, 0, 253, 15, 26, 11]);
|
|
function ad(a, b) {
|
|
return b + a;
|
|
}
|
|
function bd(a, b) {
|
|
window[a] = b;
|
|
}
|
|
function cd(a) {
|
|
var b = document.createElement("script");
|
|
b.setAttribute("src", a);
|
|
b.setAttribute("crossorigin", "anonymous");
|
|
return new Promise(function(c) {
|
|
b.addEventListener("load", function() {
|
|
c();
|
|
}, false);
|
|
b.addEventListener("error", function() {
|
|
c();
|
|
}, false);
|
|
document.body.appendChild(b);
|
|
});
|
|
}
|
|
function dd() {
|
|
return E(function(a) {
|
|
switch (a.h) {
|
|
case 1:
|
|
return a.s = 2, D2(a, WebAssembly.instantiate($c), 4);
|
|
case 4:
|
|
a.h = 3;
|
|
a.s = 0;
|
|
break;
|
|
case 2:
|
|
return a.s = 0, a.l = null, a.return(false);
|
|
case 3:
|
|
return a.return(true);
|
|
}
|
|
});
|
|
}
|
|
function ed(a) {
|
|
this.h = a;
|
|
this.listeners = {};
|
|
this.l = {};
|
|
this.L = {};
|
|
this.s = {};
|
|
this.v = {};
|
|
this.M = this.u = this.ga = true;
|
|
this.I = Promise.resolve();
|
|
this.fa = "";
|
|
this.D = {};
|
|
this.locateFile = a && a.locateFile || ad;
|
|
if ("object" === typeof window) var b = window.location.pathname.toString().substring(0, window.location.pathname.toString().lastIndexOf("/")) + "/";
|
|
else if ("undefined" !== typeof location) b = location.pathname.toString().substring(0, location.pathname.toString().lastIndexOf("/")) + "/";
|
|
else throw Error("solutions can only be loaded on a web page or in a web worker");
|
|
this.ha = b;
|
|
if (a.options) {
|
|
b = A(Object.keys(a.options));
|
|
for (var c = b.next(); !c.done; c = b.next()) {
|
|
c = c.value;
|
|
var d = a.options[c].default;
|
|
void 0 !== d && (this.l[c] = "function" === typeof d ? d() : d);
|
|
}
|
|
}
|
|
}
|
|
x = ed.prototype;
|
|
x.close = function() {
|
|
this.j && this.j.delete();
|
|
return Promise.resolve();
|
|
};
|
|
function fd(a) {
|
|
var b, c, d, e, g, f, h, k, l, m, r;
|
|
return E(function(p2) {
|
|
switch (p2.h) {
|
|
case 1:
|
|
if (!a.ga) return p2.return();
|
|
b = void 0 === a.h.files ? [] : "function" === typeof a.h.files ? a.h.files(a.l) : a.h.files;
|
|
return D2(p2, dd(), 2);
|
|
case 2:
|
|
c = p2.i;
|
|
if ("object" === typeof window) return bd("createMediapipeSolutionsWasm", { locateFile: a.locateFile }), bd("createMediapipeSolutionsPackedAssets", { locateFile: a.locateFile }), f = b.filter(function(n) {
|
|
return void 0 !== n.data;
|
|
}), h = b.filter(function(n) {
|
|
return void 0 === n.data;
|
|
}), k = Promise.all(f.map(function(n) {
|
|
var q2 = gd(a, n.url);
|
|
if (void 0 !== n.path) {
|
|
var t2 = n.path;
|
|
q2 = q2.then(function(w) {
|
|
a.overrideFile(t2, w);
|
|
return Promise.resolve(w);
|
|
});
|
|
}
|
|
return q2;
|
|
})), l = Promise.all(h.map(function(n) {
|
|
return void 0 === n.simd || n.simd && c || !n.simd && !c ? cd(a.locateFile(n.url, a.ha)) : Promise.resolve();
|
|
})).then(function() {
|
|
var n, q2, t2;
|
|
return E(function(w) {
|
|
if (1 == w.h) return n = window.createMediapipeSolutionsWasm, q2 = window.createMediapipeSolutionsPackedAssets, t2 = a, D2(w, n(q2), 2);
|
|
t2.i = w.i;
|
|
w.h = 0;
|
|
});
|
|
}), m = (function() {
|
|
return E(function(n) {
|
|
a.h.graph && a.h.graph.url ? n = D2(
|
|
n,
|
|
gd(a, a.h.graph.url),
|
|
0
|
|
) : (n.h = 0, n = void 0);
|
|
return n;
|
|
});
|
|
})(), D2(p2, Promise.all([l, k, m]), 7);
|
|
if ("function" !== typeof importScripts) throw Error("solutions can only be loaded on a web page or in a web worker");
|
|
d = b.filter(function(n) {
|
|
return void 0 === n.simd || n.simd && c || !n.simd && !c;
|
|
}).map(function(n) {
|
|
return a.locateFile(n.url, a.ha);
|
|
});
|
|
importScripts.apply(null, ea(d));
|
|
e = a;
|
|
return D2(p2, createMediapipeSolutionsWasm(Module), 6);
|
|
case 6:
|
|
e.i = p2.i;
|
|
a.m = new OffscreenCanvas(1, 1);
|
|
a.i.canvas = a.m;
|
|
g = a.i.GL.createContext(a.m, {
|
|
antialias: false,
|
|
alpha: false,
|
|
va: "undefined" !== typeof WebGL2RenderingContext ? 2 : 1
|
|
});
|
|
a.i.GL.makeContextCurrent(g);
|
|
p2.h = 4;
|
|
break;
|
|
case 7:
|
|
a.m = document.createElement("canvas");
|
|
r = a.m.getContext("webgl2", {});
|
|
if (!r && (r = a.m.getContext("webgl", {}), !r)) return alert("Failed to create WebGL canvas context when passing video frame."), p2.return();
|
|
a.K = r;
|
|
a.i.canvas = a.m;
|
|
a.i.createContext(a.m, true, true, {});
|
|
case 4:
|
|
a.j = new a.i.SolutionWasm(), a.ga = false, p2.h = 0;
|
|
}
|
|
});
|
|
}
|
|
function hd(a) {
|
|
var b, c, d, e, g, f, h, k;
|
|
return E(function(l) {
|
|
if (1 == l.h) {
|
|
if (a.h.graph && a.h.graph.url && a.fa === a.h.graph.url) return l.return();
|
|
a.u = true;
|
|
if (!a.h.graph || !a.h.graph.url) {
|
|
l.h = 2;
|
|
return;
|
|
}
|
|
a.fa = a.h.graph.url;
|
|
return D2(l, gd(a, a.h.graph.url), 3);
|
|
}
|
|
2 != l.h && (b = l.i, a.j.loadGraph(b));
|
|
c = A(Object.keys(a.D));
|
|
for (d = c.next(); !d.done; d = c.next()) e = d.value, a.j.overrideFile(e, a.D[e]);
|
|
a.D = {};
|
|
if (a.h.listeners) for (g = A(a.h.listeners), f = g.next(); !f.done; f = g.next()) h = f.value, id(a, h);
|
|
k = a.l;
|
|
a.l = {};
|
|
a.setOptions(k);
|
|
l.h = 0;
|
|
});
|
|
}
|
|
x.reset = function() {
|
|
var a = this;
|
|
return E(function(b) {
|
|
a.j && (a.j.reset(), a.s = {}, a.v = {});
|
|
b.h = 0;
|
|
});
|
|
};
|
|
x.setOptions = function(a, b) {
|
|
var c = this;
|
|
if (b = b || this.h.options) {
|
|
for (var d = [], e = [], g = {}, f = A(Object.keys(a)), h = f.next(); !h.done; g = { X: g.X, Y: g.Y }, h = f.next()) if (h = h.value, !(h in this.l && this.l[h] === a[h])) {
|
|
this.l[h] = a[h];
|
|
var k = b[h];
|
|
void 0 !== k && (k.onChange && (g.X = k.onChange, g.Y = a[h], d.push(/* @__PURE__ */ (function(l) {
|
|
return function() {
|
|
var m;
|
|
return E(function(r) {
|
|
if (1 == r.h) return D2(r, l.X(l.Y), 2);
|
|
m = r.i;
|
|
true === m && (c.u = true);
|
|
r.h = 0;
|
|
});
|
|
};
|
|
})(g))), k.graphOptionXref && (h = Object.assign(
|
|
{},
|
|
{ calculatorName: "", calculatorIndex: 0 },
|
|
k.graphOptionXref,
|
|
{ valueNumber: 1 === k.type ? a[h] : 0, valueBoolean: 0 === k.type ? a[h] : false, valueString: 2 === k.type ? a[h] : "" }
|
|
), e.push(h)));
|
|
}
|
|
if (0 !== d.length || 0 !== e.length) this.u = true, this.H = (void 0 === this.H ? [] : this.H).concat(e), this.F = (void 0 === this.F ? [] : this.F).concat(d);
|
|
}
|
|
};
|
|
function jd(a) {
|
|
var b, c, d, e, g, f, h;
|
|
return E(function(k) {
|
|
switch (k.h) {
|
|
case 1:
|
|
if (!a.u) return k.return();
|
|
if (!a.F) {
|
|
k.h = 2;
|
|
break;
|
|
}
|
|
b = A(a.F);
|
|
c = b.next();
|
|
case 3:
|
|
if (c.done) {
|
|
k.h = 5;
|
|
break;
|
|
}
|
|
d = c.value;
|
|
return D2(k, d(), 4);
|
|
case 4:
|
|
c = b.next();
|
|
k.h = 3;
|
|
break;
|
|
case 5:
|
|
a.F = void 0;
|
|
case 2:
|
|
if (a.H) {
|
|
e = new a.i.GraphOptionChangeRequestList();
|
|
g = A(a.H);
|
|
for (f = g.next(); !f.done; f = g.next()) h = f.value, e.push_back(h);
|
|
a.j.changeOptions(e);
|
|
e.delete();
|
|
a.H = void 0;
|
|
}
|
|
a.u = false;
|
|
k.h = 0;
|
|
}
|
|
});
|
|
}
|
|
x.initialize = function() {
|
|
var a = this;
|
|
return E(function(b) {
|
|
return 1 == b.h ? D2(b, fd(a), 2) : 3 != b.h ? D2(b, hd(a), 3) : D2(b, jd(a), 0);
|
|
});
|
|
};
|
|
function gd(a, b) {
|
|
var c, d;
|
|
return E(function(e) {
|
|
if (b in a.L) return e.return(a.L[b]);
|
|
c = a.locateFile(b, "");
|
|
d = fetch(c).then(function(g) {
|
|
return g.arrayBuffer();
|
|
});
|
|
a.L[b] = d;
|
|
return e.return(d);
|
|
});
|
|
}
|
|
x.overrideFile = function(a, b) {
|
|
this.j ? this.j.overrideFile(a, b) : this.D[a] = b;
|
|
};
|
|
x.clearOverriddenFiles = function() {
|
|
this.D = {};
|
|
this.j && this.j.clearOverriddenFiles();
|
|
};
|
|
x.send = function(a, b) {
|
|
var c = this, d, e, g, f, h, k, l, m, r;
|
|
return E(function(p2) {
|
|
switch (p2.h) {
|
|
case 1:
|
|
if (!c.h.inputs) return p2.return();
|
|
d = 1e3 * (void 0 === b || null === b ? performance.now() : b);
|
|
return D2(p2, c.I, 2);
|
|
case 2:
|
|
return D2(p2, c.initialize(), 3);
|
|
case 3:
|
|
e = new c.i.PacketDataList();
|
|
g = A(Object.keys(a));
|
|
for (f = g.next(); !f.done; f = g.next()) if (h = f.value, k = c.h.inputs[h]) {
|
|
a: {
|
|
var n = a[h];
|
|
switch (k.type) {
|
|
case "video":
|
|
var q2 = c.s[k.stream];
|
|
q2 || (q2 = new Wc(c.i, c.K), c.s[k.stream] = q2);
|
|
0 === q2.m && (q2.m = q2.i.createTexture());
|
|
if ("undefined" !== typeof HTMLVideoElement && n instanceof HTMLVideoElement) {
|
|
var t2 = n.videoWidth;
|
|
var w = n.videoHeight;
|
|
} else "undefined" !== typeof HTMLImageElement && n instanceof HTMLImageElement ? (t2 = n.naturalWidth, w = n.naturalHeight) : (t2 = n.width, w = n.height);
|
|
w = { glName: q2.m, width: t2, height: w };
|
|
t2 = q2.h;
|
|
t2.canvas.width = w.width;
|
|
t2.canvas.height = w.height;
|
|
t2.activeTexture(t2.TEXTURE0);
|
|
q2.i.bindTexture2d(q2.m);
|
|
t2.texImage2D(t2.TEXTURE_2D, 0, t2.RGBA, t2.RGBA, t2.UNSIGNED_BYTE, n);
|
|
q2.i.bindTexture2d(0);
|
|
q2 = w;
|
|
break a;
|
|
case "detections":
|
|
q2 = c.s[k.stream];
|
|
q2 || (q2 = new Zc(c.i), c.s[k.stream] = q2);
|
|
q2.data || (q2.data = new q2.h.DetectionListData());
|
|
q2.data.reset(n.length);
|
|
for (w = 0; w < n.length; ++w) {
|
|
t2 = n[w];
|
|
var v = q2.data, B2 = v.setBoundingBox, J2 = w;
|
|
var I = t2.la;
|
|
var u = new Qc();
|
|
V2(u, 1, I.sa);
|
|
V2(u, 2, I.ta);
|
|
V2(u, 3, I.height);
|
|
V2(u, 4, I.width);
|
|
V2(u, 5, I.rotation);
|
|
U(u, 6, I.pa);
|
|
I = u.l();
|
|
B2.call(v, J2, I);
|
|
if (t2.ea) for (v = 0; v < t2.ea.length; ++v) {
|
|
u = t2.ea[v];
|
|
B2 = q2.data;
|
|
J2 = B2.addNormalizedLandmark;
|
|
I = w;
|
|
u = Object.assign({}, u, { visibility: u.visibility ? u.visibility : 0 });
|
|
var C2 = new Lc();
|
|
V2(C2, 1, u.x);
|
|
V2(C2, 2, u.y);
|
|
V2(C2, 3, u.z);
|
|
u.visibility && V2(C2, 4, u.visibility);
|
|
u = C2.l();
|
|
J2.call(
|
|
B2,
|
|
I,
|
|
u
|
|
);
|
|
}
|
|
if (t2.ba) for (v = 0; v < t2.ba.length; ++v) B2 = q2.data, J2 = B2.addClassification, I = w, u = t2.ba[v], C2 = new Z2(), V2(C2, 2, u.qa), u.index && U(C2, 1, u.index), u.label && U(C2, 3, u.label), u.displayName && U(C2, 4, u.displayName), u = C2.l(), J2.call(B2, I, u);
|
|
}
|
|
q2 = q2.data;
|
|
break a;
|
|
default:
|
|
q2 = {};
|
|
}
|
|
}
|
|
l = q2;
|
|
m = k.stream;
|
|
switch (k.type) {
|
|
case "video":
|
|
e.pushTexture2d(Object.assign({}, l, { stream: m, timestamp: d }));
|
|
break;
|
|
case "detections":
|
|
r = l;
|
|
r.stream = m;
|
|
r.timestamp = d;
|
|
e.pushDetectionList(r);
|
|
break;
|
|
default:
|
|
throw Error("Unknown input config type: '" + k.type + "'");
|
|
}
|
|
}
|
|
c.j.send(e);
|
|
return D2(p2, c.I, 4);
|
|
case 4:
|
|
e.delete(), p2.h = 0;
|
|
}
|
|
});
|
|
};
|
|
function kd(a, b, c) {
|
|
var d, e, g, f, h, k, l, m, r, p2, n, q2, t2, w;
|
|
return E(function(v) {
|
|
switch (v.h) {
|
|
case 1:
|
|
if (!c) return v.return(b);
|
|
d = {};
|
|
e = 0;
|
|
g = A(Object.keys(c));
|
|
for (f = g.next(); !f.done; f = g.next()) h = f.value, k = c[h], "string" !== typeof k && "texture" === k.type && void 0 !== b[k.stream] && ++e;
|
|
1 < e && (a.M = false);
|
|
l = A(Object.keys(c));
|
|
f = l.next();
|
|
case 2:
|
|
if (f.done) {
|
|
v.h = 4;
|
|
break;
|
|
}
|
|
m = f.value;
|
|
r = c[m];
|
|
if ("string" === typeof r) return t2 = d, w = m, D2(v, ld(a, m, b[r]), 14);
|
|
p2 = b[r.stream];
|
|
if ("detection_list" === r.type) {
|
|
if (p2) {
|
|
var B2 = p2.getRectList();
|
|
for (var J2 = p2.getLandmarksList(), I = p2.getClassificationsList(), u = [], C2 = 0; C2 < B2.size(); ++C2) {
|
|
var T = Sc(B2.get(C2)), od = W2(T, 1), pd = W2(T, 2), qd = W2(T, 3), rd = W2(T, 4), sd = W2(T, 5, 0), za = void 0;
|
|
za = void 0 === za ? 0 : za;
|
|
T = { la: { sa: od, ta: pd, height: qd, width: rd, rotation: sd, pa: Rb(S(T, 6), za) }, ea: Ob(Pc(J2.get(C2)), Lc, 1).map(Vc), ba: Uc(Kc(I.get(C2))) };
|
|
u.push(T);
|
|
}
|
|
B2 = u;
|
|
} else B2 = [];
|
|
d[m] = B2;
|
|
v.h = 7;
|
|
break;
|
|
}
|
|
if ("proto_list" === r.type) {
|
|
if (p2) {
|
|
B2 = Array(p2.size());
|
|
for (J2 = 0; J2 < p2.size(); J2++) B2[J2] = p2.get(J2);
|
|
p2.delete();
|
|
} else B2 = [];
|
|
d[m] = B2;
|
|
v.h = 7;
|
|
break;
|
|
}
|
|
if (void 0 === p2) {
|
|
v.h = 3;
|
|
break;
|
|
}
|
|
if ("float_list" === r.type) {
|
|
d[m] = p2;
|
|
v.h = 7;
|
|
break;
|
|
}
|
|
if ("proto" === r.type) {
|
|
d[m] = p2;
|
|
v.h = 7;
|
|
break;
|
|
}
|
|
if ("texture" !== r.type) throw Error("Unknown output config type: '" + r.type + "'");
|
|
n = a.v[m];
|
|
n || (n = new Wc(a.i, a.K), a.v[m] = n);
|
|
return D2(v, Xc(n, p2, a.M), 13);
|
|
case 13:
|
|
q2 = v.i, d[m] = q2;
|
|
case 7:
|
|
r.transform && d[m] && (d[m] = r.transform(d[m]));
|
|
v.h = 3;
|
|
break;
|
|
case 14:
|
|
t2[w] = v.i;
|
|
case 3:
|
|
f = l.next();
|
|
v.h = 2;
|
|
break;
|
|
case 4:
|
|
return v.return(d);
|
|
}
|
|
});
|
|
}
|
|
function ld(a, b, c) {
|
|
var d;
|
|
return E(function(e) {
|
|
return "number" === typeof c || c instanceof Uint8Array || c instanceof a.i.Uint8BlobList ? e.return(c) : c instanceof a.i.Texture2dDataOut ? (d = a.v[b], d || (d = new Wc(a.i, a.K), a.v[b] = d), e.return(Xc(d, c, a.M))) : e.return(void 0);
|
|
});
|
|
}
|
|
function id(a, b) {
|
|
for (var c = b.name || "$", d = [].concat(ea(b.wants)), e = new a.i.StringList(), g = A(b.wants), f = g.next(); !f.done; f = g.next()) e.push_back(f.value);
|
|
g = a.i.PacketListener.implement({ onResults: function(h) {
|
|
for (var k = {}, l = 0; l < b.wants.length; ++l) k[d[l]] = h.get(l);
|
|
var m = a.listeners[c];
|
|
m && (a.I = kd(a, k, b.outs).then(function(r) {
|
|
r = m(r);
|
|
for (var p2 = 0; p2 < b.wants.length; ++p2) {
|
|
var n = k[d[p2]];
|
|
"object" === typeof n && n.hasOwnProperty && n.hasOwnProperty("delete") && n.delete();
|
|
}
|
|
r && (a.I = r);
|
|
}));
|
|
} });
|
|
a.j.attachMultiListener(e, g);
|
|
e.delete();
|
|
}
|
|
x.onResults = function(a, b) {
|
|
this.listeners[b || "$"] = a;
|
|
};
|
|
Aa("Solution", ed);
|
|
Aa("OptionType", { BOOL: 0, NUMBER: 1, ua: 2, 0: "BOOL", 1: "NUMBER", 2: "STRING" });
|
|
function md(a) {
|
|
void 0 === a && (a = 0);
|
|
switch (a) {
|
|
case 1:
|
|
return "selfie_segmentation_landscape.tflite";
|
|
default:
|
|
return "selfie_segmentation.tflite";
|
|
}
|
|
}
|
|
function nd(a) {
|
|
var b = this;
|
|
a = a || {};
|
|
this.h = new ed({ locateFile: a.locateFile, files: function(c) {
|
|
return [{ simd: true, url: "selfie_segmentation_solution_simd_wasm_bin.js" }, { simd: false, url: "selfie_segmentation_solution_wasm_bin.js" }, { data: true, url: md(c.modelSelection) }];
|
|
}, graph: { url: "selfie_segmentation.binarypb" }, listeners: [{ wants: ["segmentation_mask", "image_transformed"], outs: { image: { type: "texture", stream: "image_transformed" }, segmentationMask: { type: "texture", stream: "segmentation_mask" } } }], inputs: { image: {
|
|
type: "video",
|
|
stream: "input_frames_gpu"
|
|
} }, options: { useCpuInference: { type: 0, graphOptionXref: { calculatorType: "InferenceCalculator", fieldName: "use_cpu_inference" }, default: "object" !== typeof window || void 0 === window.navigator ? false : "iPad Simulator;iPhone Simulator;iPod Simulator;iPad;iPhone;iPod".split(";").includes(navigator.platform) || navigator.userAgent.includes("Mac") && "ontouchend" in document }, selfieMode: { type: 0, graphOptionXref: { calculatorType: "GlScalerCalculator", calculatorIndex: 1, fieldName: "flip_horizontal" } }, modelSelection: {
|
|
type: 1,
|
|
graphOptionXref: { calculatorType: "ConstantSidePacketCalculator", calculatorName: "ConstantSidePacketCalculatorModelSelection", fieldName: "int_value" },
|
|
onChange: function(c) {
|
|
var d, e, g;
|
|
return E(function(f) {
|
|
if (1 == f.h) return d = md(c), e = "third_party/mediapipe/modules/selfie_segmentation/" + d, D2(f, gd(b.h, d), 2);
|
|
g = f.i;
|
|
b.h.overrideFile(e, g);
|
|
return f.return(true);
|
|
});
|
|
}
|
|
} } });
|
|
}
|
|
x = nd.prototype;
|
|
x.close = function() {
|
|
this.h.close();
|
|
return Promise.resolve();
|
|
};
|
|
x.onResults = function(a) {
|
|
this.h.onResults(a);
|
|
};
|
|
x.initialize = function() {
|
|
var a = this;
|
|
return E(function(b) {
|
|
return D2(b, a.h.initialize(), 0);
|
|
});
|
|
};
|
|
x.reset = function() {
|
|
this.h.reset();
|
|
};
|
|
x.send = function(a) {
|
|
var b = this;
|
|
return E(function(c) {
|
|
return D2(c, b.h.send(a), 0);
|
|
});
|
|
};
|
|
x.setOptions = function(a) {
|
|
this.h.setOptions(a);
|
|
};
|
|
Aa("SelfieSegmentation", nd);
|
|
Aa("VERSION", "0.1.1675465747");
|
|
}).call(selfie_segmentation);
|
|
return selfie_segmentation;
|
|
}
|
|
var selfie_segmentationExports = /* @__PURE__ */ requireSelfie_segmentation();
|
|
var R = function(t2, e) {
|
|
return R = Object.setPrototypeOf || { __proto__: [] } instanceof Array && function(t3, e2) {
|
|
t3.__proto__ = e2;
|
|
} || function(t3, e2) {
|
|
for (var n in e2) e2.hasOwnProperty(n) && (t3[n] = e2[n]);
|
|
}, R(t2, e);
|
|
};
|
|
function F(t2, e) {
|
|
function n() {
|
|
this.constructor = t2;
|
|
}
|
|
R(t2, e), t2.prototype = null === e ? Object.create(e) : (n.prototype = e.prototype, new n());
|
|
}
|
|
var C = function() {
|
|
return C = Object.assign || function(t2) {
|
|
for (var e, n = 1, r = arguments.length; n < r; n++) for (var i in e = arguments[n]) Object.prototype.hasOwnProperty.call(e, i) && (t2[i] = e[i]);
|
|
return t2;
|
|
}, C.apply(this, arguments);
|
|
};
|
|
function D(t2, e, n, r) {
|
|
return new (n || (n = Promise))((function(i, o) {
|
|
function a(t3) {
|
|
try {
|
|
u(r.next(t3));
|
|
} catch (t4) {
|
|
o(t4);
|
|
}
|
|
}
|
|
function s(t3) {
|
|
try {
|
|
u(r.throw(t3));
|
|
} catch (t4) {
|
|
o(t4);
|
|
}
|
|
}
|
|
function u(t3) {
|
|
var e2;
|
|
t3.done ? i(t3.value) : (e2 = t3.value, e2 instanceof n ? e2 : new n((function(t4) {
|
|
t4(e2);
|
|
}))).then(a, s);
|
|
}
|
|
u((r = r.apply(t2, [])).next());
|
|
}));
|
|
}
|
|
function B(t2, e) {
|
|
var n, r, i, o, a = { label: 0, sent: function() {
|
|
if (1 & i[0]) throw i[1];
|
|
return i[1];
|
|
}, trys: [], ops: [] };
|
|
return o = { next: s(0), throw: s(1), return: s(2) }, "function" == typeof Symbol && (o[Symbol.iterator] = function() {
|
|
return this;
|
|
}), o;
|
|
function s(o2) {
|
|
return function(s2) {
|
|
return (function(o3) {
|
|
if (n) throw new TypeError("Generator is already executing.");
|
|
for (; a; ) try {
|
|
if (n = 1, r && (i = 2 & o3[0] ? r.return : o3[0] ? r.throw || ((i = r.return) && i.call(r), 0) : r.next) && !(i = i.call(r, o3[1])).done) return i;
|
|
switch (r = 0, i && (o3 = [2 & o3[0], i.value]), o3[0]) {
|
|
case 0:
|
|
case 1:
|
|
i = o3;
|
|
break;
|
|
case 4:
|
|
return a.label++, { value: o3[1], done: false };
|
|
case 5:
|
|
a.label++, r = o3[1], o3 = [0];
|
|
continue;
|
|
case 7:
|
|
o3 = a.ops.pop(), a.trys.pop();
|
|
continue;
|
|
default:
|
|
if (!(i = a.trys, (i = i.length > 0 && i[i.length - 1]) || 6 !== o3[0] && 2 !== o3[0])) {
|
|
a = 0;
|
|
continue;
|
|
}
|
|
if (3 === o3[0] && (!i || o3[1] > i[0] && o3[1] < i[3])) {
|
|
a.label = o3[1];
|
|
break;
|
|
}
|
|
if (6 === o3[0] && a.label < i[1]) {
|
|
a.label = i[1], i = o3;
|
|
break;
|
|
}
|
|
if (i && a.label < i[2]) {
|
|
a.label = i[2], a.ops.push(o3);
|
|
break;
|
|
}
|
|
i[2] && a.ops.pop(), a.trys.pop();
|
|
continue;
|
|
}
|
|
o3 = e.call(t2, a);
|
|
} catch (t3) {
|
|
o3 = [6, t3], r = 0;
|
|
} finally {
|
|
n = i = 0;
|
|
}
|
|
if (5 & o3[0]) throw o3[1];
|
|
return { value: o3[0] ? o3[1] : void 0, done: true };
|
|
})([o2, s2]);
|
|
};
|
|
}
|
|
}
|
|
function L(t2) {
|
|
return t2 instanceof SVGAnimatedLength ? t2.baseVal.value : t2;
|
|
}
|
|
function z(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var r, i;
|
|
return B(this, (function(o) {
|
|
switch (o.label) {
|
|
case 0:
|
|
return r = document.createElement("canvas"), t2 instanceof Tensor ? [4, toPixels(t2, r)] : [3, 2];
|
|
case 1:
|
|
return o.sent(), [3, 3];
|
|
case 2:
|
|
r.width = L(t2.width), r.height = L(t2.height), i = r.getContext("2d"), t2 instanceof ImageData ? i.putImageData(t2, 0, 0) : i.drawImage(t2, 0, 0), o.label = 3;
|
|
case 3:
|
|
return [2, r];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function W(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var r, i, o, a, s, u;
|
|
return B(this, (function(f) {
|
|
switch (f.label) {
|
|
case 0:
|
|
return t2 instanceof Tensor ? (r = t2.shape.slice(0, 2), i = r[0], o = r[1], a = ImageData.bind, [4, toPixels(t2)]) : [3, 2];
|
|
case 1:
|
|
return [2, new (a.apply(ImageData, [void 0, f.sent(), o, i]))()];
|
|
case 2:
|
|
return s = document.createElement("canvas"), u = s.getContext("2d"), s.width = L(t2.width), s.height = L(t2.height), u.drawImage(t2, 0, 0), [2, u.getImageData(0, 0, s.width, s.height)];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function j(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var e, r;
|
|
return B(this, (function(i) {
|
|
switch (i.label) {
|
|
case 0:
|
|
return t2 instanceof SVGImageElement || t2 instanceof OffscreenCanvas ? [4, z(t2)] : [3, 2];
|
|
case 1:
|
|
return r = i.sent(), [3, 3];
|
|
case 2:
|
|
r = t2, i.label = 3;
|
|
case 3:
|
|
return e = r, [2, fromPixels$1(e, 4)];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function V(t2) {
|
|
if (t2 < 0 || t2 >= 256) throw new Error("Mask value must be in range [0, 255] but got " + t2);
|
|
if (!Number.isInteger(t2)) throw new Error("Mask value must be an integer but got " + t2);
|
|
}
|
|
function q(t2) {
|
|
var e = t2.shape[2], n = argMax$2(t2, 2), r = reshape$2(n, [-1]);
|
|
return oneHot$2(r, e);
|
|
}
|
|
function N(t2, e) {
|
|
return tidy((function() {
|
|
return cast$2(greater$2(t2, scalar(e)), "int32");
|
|
}));
|
|
}
|
|
function Q(t2, e) {
|
|
var n = e.shape, o = n[0], s = n[1], m = n[2];
|
|
return tidy((function() {
|
|
var n2 = q(e), r = expandDims$2(range$2(0, m, 1, "int32"), 1), g = cast$2(matMul$1(n2, r), "int32"), v = reshape$2(g, [o, s]), w = add$1(v, scalar(1, "int32"));
|
|
return sub$2((function(t3, e2) {
|
|
return mul(t3, e2);
|
|
})(w, t2), scalar(1, "int32"));
|
|
}));
|
|
}
|
|
var X = (function() {
|
|
function t2(t3, e) {
|
|
this.model = t3, this.outputStride = e;
|
|
var n = this.model.inputs[0].shape;
|
|
assert(-1 === n[1] && -1 === n[2], (function() {
|
|
return "Input shape [" + n[1] + ", " + n[2] + "] must both be equal to or -1";
|
|
}));
|
|
}
|
|
return t2.prototype.predict = function(t3) {
|
|
var e = this;
|
|
return tidy((function() {
|
|
var n = e.preprocessInput(cast$2(t3, "float32")), r = expandDims$2(n, 0), o = e.model.predict(r).map((function(t4) {
|
|
return squeeze(t4, [0]);
|
|
})), a = e.nameOutputResults(o);
|
|
return { heatmapScores: sigmoid$2(a.heatmap), offsets: a.offsets, displacementFwd: a.displacementFwd, displacementBwd: a.displacementBwd, segmentation: a.segmentation, partHeatmaps: a.partHeatmaps, longOffsets: a.longOffsets, partOffsets: a.partOffsets };
|
|
}));
|
|
}, t2.prototype.dispose = function() {
|
|
this.model.dispose();
|
|
}, t2;
|
|
})(), Y = (function(t2) {
|
|
function e() {
|
|
return null !== t2 && t2.apply(this, arguments) || this;
|
|
}
|
|
return F(e, t2), e.prototype.preprocessInput = function(t3) {
|
|
return tidy((function() {
|
|
return sub$2(div$1(t3, 127.5), 1);
|
|
}));
|
|
}, e.prototype.nameOutputResults = function(t3) {
|
|
return { offsets: t3[0], segmentation: t3[1], partHeatmaps: t3[2], longOffsets: t3[3], heatmap: t3[4], displacementFwd: t3[5], displacementBwd: t3[6], partOffsets: t3[7] };
|
|
}, e;
|
|
})(X), G = ["nose", "leftEye", "rightEye", "leftEar", "rightEar", "leftShoulder", "rightShoulder", "leftElbow", "rightElbow", "leftWrist", "rightWrist", "leftHip", "rightHip", "leftKnee", "rightKnee", "leftAnkle", "rightAnkle"], $ = G.length, J = G.reduce((function(t2, e, n) {
|
|
return t2[e] = n, t2;
|
|
}), {});
|
|
[["leftHip", "leftShoulder"], ["leftElbow", "leftShoulder"], ["leftElbow", "leftWrist"], ["leftHip", "leftKnee"], ["leftKnee", "leftAnkle"], ["rightHip", "rightShoulder"], ["rightElbow", "rightShoulder"], ["rightElbow", "rightWrist"], ["rightHip", "rightKnee"], ["rightKnee", "rightAnkle"], ["leftShoulder", "rightShoulder"], ["leftHip", "rightHip"]].map((function(t2) {
|
|
var e = t2[0], n = t2[1];
|
|
return [J[e], J[n]];
|
|
}));
|
|
function Z(t2, e, n) {
|
|
var r = t2[0], i = t2[1], o = e[0], a = e[1], s = n.top, u = n.bottom;
|
|
return [a / (n.left + n.right + i), o / (s + u + r)];
|
|
}
|
|
function tt(t2, e, n, r) {
|
|
return { y: r.get(t2, e, n), x: r.get(t2, e, n + $) };
|
|
}
|
|
function et(t2, e, n) {
|
|
var r = tt(t2.heatmapY, t2.heatmapX, t2.id, n), i = r.y, o = r.x;
|
|
return { x: t2.heatmapX * e + o, y: t2.heatmapY * e + i };
|
|
}
|
|
function nt(t2, e, n) {
|
|
return t2 < e ? e : t2 > n ? n : t2;
|
|
}
|
|
function rt(t2, e) {
|
|
return { x: t2.x + e.x, y: t2.y + e.y };
|
|
}
|
|
function it(t2, e, n) {
|
|
void 0 === n && (n = 0.3);
|
|
for (var r = 0, i = 0, o = 0; o < t2.length; o++) e.keypoints[o].score > n && (i += 1, r += Math.pow(t2[o].x - e.keypoints[o].position.x, 2) + Math.pow(t2[o].y - e.keypoints[o].position.y, 2));
|
|
return 0 === i ? r = 1 / 0 : r /= i, r;
|
|
}
|
|
function ot(t2, e, n, r, i, o, a) {
|
|
for (var s = a[0], u = a[1], f = n(t2), l = f.y * r + f.x, c = i[$ * (2 * l) + e], d = i[$ * (2 * l + 1) + e], h = t2.y + c, p2 = t2.x + d, m = 0; m < o; m++) {
|
|
h = Math.min(h, s - 1);
|
|
var g = n({ x: p2 = Math.min(p2, u - 1), y: h }), v = g.y * r + g.x;
|
|
h += c = i[$ * (2 * v) + e], p2 += d = i[$ * (2 * v + 1) + e];
|
|
}
|
|
return { x: p2, y: h };
|
|
}
|
|
function at(t2, e, n, r, i, o, a, s, u, f) {
|
|
for (var l = i[0], c = i[1], d = o[0], h = o[1], p2 = s[0], m = s[1], g = [], v = function(t3) {
|
|
return (function(t4, e2, n2, r2) {
|
|
var i2 = e2[0], o2 = e2[1], a6 = n2[0], s2 = n2[1], u2 = Math.round(((i2 + t4.y + 1) * s2 - 1) / r2);
|
|
return { x: Math.round(((o2 + t4.x + 1) * a6 - 1) / r2), y: u2 };
|
|
})(t3, [l, c], [d, h], u);
|
|
}, w = 0; w < r; w++) {
|
|
var y = ot(t2, w, v, a, e, f, [p2, m]);
|
|
g.push(y);
|
|
}
|
|
for (var b = -1, S = 1 / 0, x = 0; x < n.length; x++) {
|
|
var M = it(g, n[x]);
|
|
M < S && (b = x, S = M);
|
|
}
|
|
return b;
|
|
}
|
|
function st(t2, e) {
|
|
var n = t2[0], r = t2[1];
|
|
return [Math.round((r - 1) / e + 1), Math.round((n - 1) / e + 1)];
|
|
}
|
|
function ut(t2, e, n, r, i, o, a, s, u, f, l) {
|
|
for (var d = a[0], h = a[1], p2 = t2.shape, m = p2[0], g = p2[1], v = e.shape.slice(0, 2), w = v[0], y = v[1], x = reshape$2(e, [w, y, 2, $]), M = new Float32Array(l * $ * 3).fill(0), k = 0; k < n.length; k++) for (var E = k * $ * 3, T = n[k], P = 0; P < $; P++) {
|
|
var I = T.keypoints[P], O = E + 3 * P;
|
|
M[O] = I.score, M[O + 1] = I.position.y, M[O + 2] = I.position.x;
|
|
}
|
|
var _ = Z([r, i], [d, h], s), H = _[0], A = _[1], R2 = tensor(M, [l, $, 3]), F2 = s.top, C2 = s.left, D2 = { variableNames: ["segmentation", "longOffsets", "poses"], outputShape: [m, g], userCode: "\n int convertToPositionInOutput(int pos, int pad, float scale, int stride) {\n return round(((float(pos + pad) + 1.0) * scale - 1.0) / float(stride));\n }\n\n float convertToPositionInOutputFloat(\n int pos, int pad, float scale, int stride) {\n return ((float(pos + pad) + 1.0) * scale - 1.0) / float(stride);\n }\n\n float dist(float x1, float y1, float x2, float y2) {\n return pow(x1 - x2, 2.0) + pow(y1 - y2, 2.0);\n }\n\n float sampleLongOffsets(float h, float w, int d, int k) {\n float fh = fract(h);\n float fw = fract(w);\n int clH = int(ceil(h));\n int clW = int(ceil(w));\n int flH = int(floor(h));\n int flW = int(floor(w));\n float o11 = getLongOffsets(flH, flW, d, k);\n float o12 = getLongOffsets(flH, clW, d, k);\n float o21 = getLongOffsets(clH, flW, d, k);\n float o22 = getLongOffsets(clH, clW, d, k);\n float o1 = mix(o11, o12, fw);\n float o2 = mix(o21, o22, fw);\n return mix(o1, o2, fh);\n }\n\n int findNearestPose(int h, int w) {\n float prob = getSegmentation(h, w);\n if (prob < 1.0) {\n return -1;\n }\n\n // Done(Tyler): convert from output space h/w to strided space.\n float stridedH = convertToPositionInOutputFloat(\n h, " + F2 + ", " + A + ", " + o + ");\n float stridedW = convertToPositionInOutputFloat(\n w, " + C2 + ", " + H + ", " + o + ");\n\n float minDist = 1000000.0;\n int iMin = -1;\n for (int i = 0; i < " + l + "; i++) {\n float curDistSum = 0.0;\n int numKpt = 0;\n for (int k = 0; k < " + $ + "; k++) {\n float dy = sampleLongOffsets(stridedH, stridedW, 0, k);\n float dx = sampleLongOffsets(stridedH, stridedW, 1, k);\n\n float y = float(h) + dy;\n float x = float(w) + dx;\n\n for (int s = 0; s < " + u + "; s++) {\n int yRounded = round(min(y, float(" + (r - 1) + ")));\n int xRounded = round(min(x, float(" + (i - 1) + ")));\n\n float yStrided = convertToPositionInOutputFloat(\n yRounded, " + F2 + ", " + A + ", " + o + ");\n float xStrided = convertToPositionInOutputFloat(\n xRounded, " + C2 + ", " + H + ", " + o + ");\n\n float dy = sampleLongOffsets(yStrided, xStrided, 0, k);\n float dx = sampleLongOffsets(yStrided, xStrided, 1, k);\n\n y = y + dy;\n x = x + dx;\n }\n\n float poseScore = getPoses(i, k, 0);\n float poseY = getPoses(i, k, 1);\n float poseX = getPoses(i, k, 2);\n if (poseScore > " + f + ") {\n numKpt = numKpt + 1;\n curDistSum = curDistSum + dist(x, y, poseX, poseY);\n }\n }\n if (numKpt > 0 && curDistSum / float(numKpt) < minDist) {\n minDist = curDistSum / float(numKpt);\n iMin = i;\n }\n }\n return iMin;\n }\n\n void main() {\n ivec2 coords = getOutputCoords();\n int nearestPose = findNearestPose(coords[0], coords[1]);\n setOutput(float(nearestPose));\n }\n " };
|
|
return backend().compileAndRun(D2, [t2, x, R2]);
|
|
}
|
|
function ft() {
|
|
return "webgl" === getBackend();
|
|
}
|
|
function lt(t2, e, n, o, s, u, f, l, c, d, h, p2) {
|
|
var m = f[0], g = f[1];
|
|
return void 0 === c && (c = 0.2), void 0 === d && (d = 8), void 0 === h && (h = 0.3), void 0 === p2 && (p2 = 10), D(this, void 0, void 0, (function() {
|
|
var f2, v, w, y, b;
|
|
return B(this, (function(S) {
|
|
switch (S.label) {
|
|
case 0:
|
|
return f2 = n.filter((function(t3) {
|
|
return t3.score >= c;
|
|
})), ft() ? (w = tidy((function() {
|
|
var n2 = ut(t2, e, f2, o, s, u, [m, g], l, d, h, p2), c2 = engine().makeTensorFromDataId(n2.dataId, n2.shape, n2.dtype);
|
|
return f2.map((function(t3, e2) {
|
|
return (function(t4, e3) {
|
|
return tidy((function() {
|
|
return cast$2(equal$2(t4, scalar(e3)), "int32");
|
|
}));
|
|
})(c2, e2);
|
|
}));
|
|
})), [4, Promise.all(w.map((function(t3) {
|
|
return t3.data();
|
|
})))]) : [3, 2];
|
|
case 1:
|
|
return v = S.sent(), w.forEach((function(t3) {
|
|
return t3.dispose();
|
|
})), [3, 5];
|
|
case 2:
|
|
return [4, t2.data()];
|
|
case 3:
|
|
return y = S.sent(), [4, e.data()];
|
|
case 4:
|
|
b = S.sent(), v = (function(t3, e2, n2, r, i, o2, a, s2, u2, f3) {
|
|
var l2 = a[0], c2 = a[1];
|
|
void 0 === f3 && (f3 = 5);
|
|
for (var d2 = n2.map((function(t4) {
|
|
return new Uint8Array(r * i).fill(0);
|
|
})), h2 = s2.top, p3 = s2.left, m2 = Z([r, i], [l2, c2], s2), g2 = m2[0], v2 = m2[1], w2 = st([l2, c2], o2)[0], y2 = 0; y2 < r; y2 += 1) for (var b2 = 0; b2 < i; b2 += 1) {
|
|
var S2 = y2 * i + b2;
|
|
if (1 === t3[S2]) {
|
|
var x = at({ x: b2, y: y2 }, e2, n2, f3, [h2, p3], [g2, v2], w2, [r, i], o2, u2);
|
|
x >= 0 && (d2[x][S2] = 1);
|
|
}
|
|
}
|
|
return d2;
|
|
})(y, b, f2, o, s, u, [m, g], l, d), S.label = 5;
|
|
case 5:
|
|
return [2, v.map((function(t3, e2) {
|
|
return { data: t3, pose: f2[e2], width: s, height: o };
|
|
}))];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function ct(t2, e, n, o, s, u, f, l, c, m, g, v, w) {
|
|
var y = l[0], b = l[1];
|
|
return void 0 === m && (m = 0.2), void 0 === g && (g = 8), void 0 === v && (v = 0.3), void 0 === w && (w = 10), D(this, void 0, void 0, (function() {
|
|
var l2, S, x, E, T, P;
|
|
return B(this, (function(I) {
|
|
switch (I.label) {
|
|
case 0:
|
|
return l2 = o.filter((function(t3) {
|
|
return t3.score >= m;
|
|
})), ft() ? (x = tidy((function() {
|
|
var o2 = ut(t2, e, l2, s, u, f, [y, b], c, g, v, w), m2 = engine().makeTensorFromDataId(o2.dataId, o2.shape, o2.dtype);
|
|
return l2.map((function(t3, e2) {
|
|
return (function(t4, e3, n2) {
|
|
return tidy((function() {
|
|
return sub$2(mul(cast$2(equal$2(t4, scalar(n2)), "int32"), add$1(e3, 1)), 1);
|
|
}));
|
|
})(m2, n, e2);
|
|
}));
|
|
})), [4, Promise.all(x.map((function(t3) {
|
|
return t3.data();
|
|
})))]) : [3, 2];
|
|
case 1:
|
|
return S = I.sent(), x.forEach((function(t3) {
|
|
return t3.dispose();
|
|
})), [3, 6];
|
|
case 2:
|
|
return [4, t2.data()];
|
|
case 3:
|
|
return E = I.sent(), [4, e.data()];
|
|
case 4:
|
|
return T = I.sent(), [4, n.data()];
|
|
case 5:
|
|
P = I.sent(), S = (function(t3, e2, n2, r, i, o2, a, s2, u2, f2, l3) {
|
|
var c2 = s2[0], d = s2[1];
|
|
void 0 === l3 && (l3 = 5);
|
|
for (var h = r.map((function(t4) {
|
|
return new Int32Array(i * o2).fill(-1);
|
|
})), p2 = u2.top, m2 = u2.left, g2 = Z([i, o2], [c2, d], u2), v2 = g2[0], w2 = g2[1], y2 = st([c2, d], a)[0], b2 = 0; b2 < i; b2 += 1) for (var S2 = 0; S2 < o2; S2 += 1) {
|
|
var x2 = b2 * o2 + S2;
|
|
if (1 === t3[x2]) {
|
|
var M = at({ x: S2, y: b2 }, e2, r, l3, [p2, m2], [v2, w2], y2, [i, o2], a, f2);
|
|
M >= 0 && (h[M][x2] = n2[x2]);
|
|
}
|
|
}
|
|
return h;
|
|
})(E, T, P, l2, s, u, f, [y, b], c, g), I.label = 6;
|
|
case 6:
|
|
return [2, S.map((function(t3, e2) {
|
|
return { pose: l2[e2], data: t3, height: s, width: u };
|
|
}))];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function dt(t2) {
|
|
return Math.floor(t2 / 2);
|
|
}
|
|
var ht = (function() {
|
|
function t2(t3, e) {
|
|
this.priorityQueue = new Array(t3), this.numberOfElements = -1, this.getElementValue = e;
|
|
}
|
|
return t2.prototype.enqueue = function(t3) {
|
|
this.priorityQueue[++this.numberOfElements] = t3, this.swim(this.numberOfElements);
|
|
}, t2.prototype.dequeue = function() {
|
|
var t3 = this.priorityQueue[0];
|
|
return this.exchange(0, this.numberOfElements--), this.sink(0), this.priorityQueue[this.numberOfElements + 1] = null, t3;
|
|
}, t2.prototype.empty = function() {
|
|
return -1 === this.numberOfElements;
|
|
}, t2.prototype.size = function() {
|
|
return this.numberOfElements + 1;
|
|
}, t2.prototype.all = function() {
|
|
return this.priorityQueue.slice(0, this.numberOfElements + 1);
|
|
}, t2.prototype.max = function() {
|
|
return this.priorityQueue[0];
|
|
}, t2.prototype.swim = function(t3) {
|
|
for (; t3 > 0 && this.less(dt(t3), t3); ) this.exchange(t3, dt(t3)), t3 = dt(t3);
|
|
}, t2.prototype.sink = function(t3) {
|
|
for (; 2 * t3 <= this.numberOfElements; ) {
|
|
var e = 2 * t3;
|
|
if (e < this.numberOfElements && this.less(e, e + 1) && e++, !this.less(t3, e)) break;
|
|
this.exchange(t3, e), t3 = e;
|
|
}
|
|
}, t2.prototype.getValueAt = function(t3) {
|
|
return this.getElementValue(this.priorityQueue[t3]);
|
|
}, t2.prototype.less = function(t3, e) {
|
|
return this.getValueAt(t3) < this.getValueAt(e);
|
|
}, t2.prototype.exchange = function(t3, e) {
|
|
var n = this.priorityQueue[t3];
|
|
this.priorityQueue[t3] = this.priorityQueue[e], this.priorityQueue[e] = n;
|
|
}, t2;
|
|
})();
|
|
function pt(t2, e, n, r, i, o) {
|
|
for (var a = o.shape, s = a[0], u = a[1], f = true, l = Math.max(n - i, 0), c = Math.min(n + i + 1, s), d = l; d < c; ++d) {
|
|
for (var h = Math.max(r - i, 0), p2 = Math.min(r + i + 1, u), m = h; m < p2; ++m) if (o.get(d, m, t2) > e) {
|
|
f = false;
|
|
break;
|
|
}
|
|
if (!f) break;
|
|
}
|
|
return f;
|
|
}
|
|
var mt = [["nose", "leftEye"], ["leftEye", "leftEar"], ["nose", "rightEye"], ["rightEye", "rightEar"], ["nose", "leftShoulder"], ["leftShoulder", "leftElbow"], ["leftElbow", "leftWrist"], ["leftShoulder", "leftHip"], ["leftHip", "leftKnee"], ["leftKnee", "leftAnkle"], ["nose", "rightShoulder"], ["rightShoulder", "rightElbow"], ["rightElbow", "rightWrist"], ["rightShoulder", "rightHip"], ["rightHip", "rightKnee"], ["rightKnee", "rightAnkle"]].map((function(t2) {
|
|
var e = t2[0], n = t2[1];
|
|
return [J[e], J[n]];
|
|
})), gt = mt.map((function(t2) {
|
|
return t2[1];
|
|
})), vt = mt.map((function(t2) {
|
|
return t2[0];
|
|
}));
|
|
function wt(t2, e, n, r) {
|
|
return { y: nt(Math.round(t2.y / e), 0, n - 1), x: nt(Math.round(t2.x / e), 0, r - 1) };
|
|
}
|
|
function yt(t2, e, n, r, i, o, a, s) {
|
|
void 0 === s && (s = 2);
|
|
for (var u = r.shape, f = u[0], l = u[1], c = (function(t3, e2, n2) {
|
|
var r2 = n2.shape[2] / 2;
|
|
return { y: n2.get(e2.y, e2.x, t3), x: n2.get(e2.y, e2.x, r2 + t3) };
|
|
})(t2, wt(e.position, o, f, l), a), d = rt(e.position, c), h = 0; h < s; h++) {
|
|
var p2 = wt(d, o, f, l), m = tt(p2.y, p2.x, n, i);
|
|
d = rt({ x: p2.x * o, y: p2.y * o }, { x: m.x, y: m.y });
|
|
}
|
|
var g = wt(d, o, f, l), v = r.get(g.y, g.x, n);
|
|
return { position: d, part: G[n], score: v };
|
|
}
|
|
function bt(t2, e, n, r, i, o) {
|
|
var a = e.shape[2], s = gt.length, u = new Array(a), f = t2.part, l = t2.score, c = et(f, r, n);
|
|
u[f.id] = { score: l, part: G[f.id], position: c };
|
|
for (var d = s - 1; d >= 0; --d) {
|
|
var h = gt[d], p2 = vt[d];
|
|
u[h] && !u[p2] && (u[p2] = yt(d, u[h], p2, e, n, r, o));
|
|
}
|
|
for (d = 0; d < s; ++d) {
|
|
h = vt[d], p2 = gt[d];
|
|
u[h] && !u[p2] && (u[p2] = yt(d, u[h], p2, e, n, r, i));
|
|
}
|
|
return u;
|
|
}
|
|
function St(t2, e, n, r) {
|
|
var i = n.x, o = n.y;
|
|
return t2.some((function(t3) {
|
|
var n2, a, s, u, f, l, c = t3.keypoints[r].position;
|
|
return n2 = o, a = i, s = c.y, u = c.x, (f = s - n2) * f + (l = u - a) * l <= e;
|
|
}));
|
|
}
|
|
function xt(t2, e, n) {
|
|
var r = n.reduce((function(n2, r2, i) {
|
|
var o = r2.position, a = r2.score;
|
|
return St(t2, e, o, i) || (n2 += a), n2;
|
|
}), 0);
|
|
return r / n.length;
|
|
}
|
|
function Mt(t2, e, n, r, i, o, a, s) {
|
|
void 0 === a && (a = 0.5), void 0 === s && (s = 20);
|
|
for (var u = [], f = (function(t3, e2, n2) {
|
|
for (var r2 = n2.shape, i2 = r2[0], o2 = r2[1], a6 = r2[2], s2 = new ht(i2 * o2 * a6, (function(t4) {
|
|
return t4.score;
|
|
})), u2 = 0; u2 < i2; ++u2) for (var f2 = 0; f2 < o2; ++f2) for (var l2 = 0; l2 < a6; ++l2) {
|
|
var c2 = n2.get(u2, f2, l2);
|
|
c2 < t3 || pt(l2, c2, u2, f2, e2, n2) && s2.enqueue({ score: c2, part: { heatmapY: u2, heatmapX: f2, id: l2 } });
|
|
}
|
|
return s2;
|
|
})(a, 1, t2), l = s * s; u.length < o && !f.empty(); ) {
|
|
var c = f.dequeue();
|
|
if (!St(u, l, et(c.part, i, e), c.part.id)) {
|
|
var d = bt(c, t2, e, i, n, r), h = xt(u, l, d);
|
|
u.push({ keypoints: d, score: h });
|
|
}
|
|
}
|
|
return u;
|
|
}
|
|
var kt, Et = [-123.15, -115.9, -103.06], Tt = (function(t2) {
|
|
function e() {
|
|
return null !== t2 && t2.apply(this, arguments) || this;
|
|
}
|
|
return F(e, t2), e.prototype.preprocessInput = function(t3) {
|
|
return add$1(t3, Et);
|
|
}, e.prototype.nameOutputResults = function(t3) {
|
|
var e2 = t3[0], n = t3[1], r = t3[2], i = t3[3], o = t3[4], a = t3[5];
|
|
return { offsets: o, segmentation: t3[6], partHeatmaps: a, longOffsets: i, heatmap: r, displacementFwd: n, displacementBwd: e2, partOffsets: t3[7] };
|
|
}, e;
|
|
})(X), Pt = "https://storage.googleapis.com/tfjs-models/savedmodel/bodypix/resnet50/", It = "https://storage.googleapis.com/tfjs-models/savedmodel/bodypix/mobilenet/";
|
|
function Ot(t2) {
|
|
if ("undefined" != typeof HTMLCanvasElement && t2 instanceof HTMLCanvasElement || "undefined" != typeof OffscreenCanvas && t2 instanceof OffscreenCanvas || "undefined" != typeof HTMLImageElement && t2 instanceof HTMLImageElement) return (function(t3) {
|
|
if ("offsetHeight" in t3 && 0 !== t3.offsetHeight && "offsetWidth" in t3 && 0 !== t3.offsetWidth) return [t3.offsetHeight, t3.offsetWidth];
|
|
if (null != t3.height && null != t3.width) return [t3.height, t3.width];
|
|
throw new Error("HTMLImageElement must have height and width attributes set.");
|
|
})(t2);
|
|
if ("undefined" != typeof ImageData && t2 instanceof ImageData) return [t2.height, t2.width];
|
|
if ("undefined" != typeof HTMLVideoElement && t2 instanceof HTMLVideoElement) return (function(t3) {
|
|
return t3.hasAttribute("height") && t3.hasAttribute("width") ? [t3.height, t3.width] : [t3.videoHeight, t3.videoWidth];
|
|
})(t2);
|
|
if (t2 instanceof Tensor) return [t2.shape[0], t2.shape[1]];
|
|
throw new Error("error: Unknown input type: " + t2 + ".");
|
|
}
|
|
function _t(t2, e) {
|
|
return (function(t3, e2) {
|
|
return (t3 - 1) % e2 == 0;
|
|
})(t2, e) ? t2 : Math.floor(t2 / e) * e + 1;
|
|
}
|
|
var Ht = { low: "low", medium: "medium", high: "high", full: "full" }, At = ((kt = {})[Ht.low] = 0.25, kt[Ht.medium] = 0.5, kt[Ht.high] = 0.75, kt[Ht.full] = 1, kt);
|
|
function Rt(t2, e, n) {
|
|
var r = n[0], i = n[1], o = (function(t3) {
|
|
if ("string" == typeof t3) {
|
|
var e2 = At[t3];
|
|
return assert("number" == typeof e2, (function() {
|
|
return "string value of inputResolution must be one of " + Object.values(Ht).join(",") + " but was " + t3 + ".";
|
|
})), e2;
|
|
}
|
|
return assert("number" == typeof t3 && t3 <= 2 && t3 >= 0.1, (function() {
|
|
return "inputResolution must be a string or number between 0.1 and 2, but was " + t3;
|
|
})), t3;
|
|
})(t2);
|
|
return [_t(r * o, e), _t(i * o, e)];
|
|
}
|
|
function Ft(t2, e, n, i, o) {
|
|
var a = e[0], s = e[1], f = n[0], l = n[1], c = i[0], d = c[0], h = c[1], p2 = i[1], m = p2[0], w = p2[1];
|
|
return tidy((function() {
|
|
var e2 = image$1.resizeBilinear(t2, [f, l], true);
|
|
return e2 = sigmoid$2(e2), (function(t3, e3, n2) {
|
|
var i2 = e3[0], o2 = e3[1], a6 = n2[0], s2 = a6[0], f2 = a6[1], l2 = n2[1], c2 = l2[0], d2 = l2[1];
|
|
return tidy((function() {
|
|
var e4 = expandDims$2(t3);
|
|
return squeeze(image$1.cropAndResize(e4, [[s2 / (i2 + s2 + f2 - 1), c2 / (o2 + c2 + d2 - 1), (s2 + i2 - 1) / (i2 + s2 + f2 - 1), (c2 + o2 - 1) / (o2 + c2 + d2 - 1)]], [0], [i2, o2]), [0]);
|
|
}));
|
|
})(e2, [a, s], [[d, h], [m, w]]);
|
|
}));
|
|
}
|
|
function Ct(t2, i) {
|
|
var o = i[0], a = i[1], s = Ot(t2), u = s[0], f = s[1], l = a / o, c = [0, 0, 0, 0], d = c[0], h = c[1], p2 = c[2], m = c[3];
|
|
f / u < l ? (d = 0, h = 0, p2 = Math.round(0.5 * (l * u - f)), m = Math.round(0.5 * (l * u - f))) : (d = Math.round(0.5 * (1 / l * f - u)), h = Math.round(0.5 * (1 / l * f - u)), p2 = 0, m = 0);
|
|
var g = tidy((function() {
|
|
var r = (function(t3) {
|
|
return t3 instanceof Tensor ? t3 : fromPixels$1(t3);
|
|
})(t2);
|
|
return r = pad3d(r, [[d, h], [p2, m], [0, 0]]), image$1.resizeBilinear(r, [o, a]);
|
|
}));
|
|
return { resized: g, padding: { top: d, left: p2, right: m, bottom: h } };
|
|
}
|
|
function Dt(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(e) {
|
|
return [2, Promise.all(t2.map((function(t3) {
|
|
return t3.buffer();
|
|
})))];
|
|
}));
|
|
}));
|
|
}
|
|
function Bt(t2, e, n, r, i) {
|
|
var o = e[0], a = e[1], s = n[0], u = n[1], f = (function(t3, e2, n2, r2, i2) {
|
|
return void 0 === r2 && (r2 = 0), void 0 === i2 && (i2 = 0), 1 === n2 && 1 === e2 && 0 === r2 && 0 === i2 ? t3 : t3.map((function(t4) {
|
|
return (function(t5, e3, n3, r3, i3) {
|
|
return void 0 === r3 && (r3 = 0), void 0 === i3 && (i3 = 0), { score: t5.score, keypoints: t5.keypoints.map((function(t6) {
|
|
var o2 = t6.score, a6 = t6.part, s2 = t6.position;
|
|
return { score: o2, part: a6, position: { x: s2.x * n3 + i3, y: s2.y * e3 + r3 } };
|
|
})) };
|
|
})(t4, e2, n2, r2, i2);
|
|
}));
|
|
})(t2, (o + r.top + r.bottom) / s, (a + r.left + r.right) / u, -r.top, -r.left);
|
|
return i ? (function(t3, e2) {
|
|
return e2 <= 0 ? t3 : t3.map((function(t4) {
|
|
return (function(t5, e3) {
|
|
return { score: t5.score, keypoints: t5.keypoints.map((function(t6) {
|
|
var n2 = t6.score, r2 = t6.part, i2 = t6.position;
|
|
return { score: n2, part: r2, position: { x: e3 - 1 - i2.x, y: i2.y } };
|
|
})) };
|
|
})(t4, e2);
|
|
}));
|
|
})(f, a) : f;
|
|
}
|
|
var Lt = { architecture: "MobileNetV1", outputStride: 16, quantBytes: 4, multiplier: 0.75 }, zt = ["MobileNetV1", "ResNet50"], Wt = { MobileNetV1: [8, 16, 32], ResNet50: [32, 16] }, jt = { MobileNetV1: [0.5, 0.75, 1], ResNet50: [1] }, Vt = [1, 2, 4];
|
|
var Kt = { flipHorizontal: false, internalResolution: "medium", segmentationThreshold: 0.7, maxDetections: 10, scoreThreshold: 0.4, nmsRadius: 20 }, Ut = { flipHorizontal: false, internalResolution: "medium", segmentationThreshold: 0.7, maxDetections: 10, scoreThreshold: 0.4, nmsRadius: 20, minKeypointScore: 0.3, refineSteps: 10 };
|
|
function qt(t2) {
|
|
var e = t2.segmentationThreshold, n = t2.maxDetections, r = t2.scoreThreshold, i = t2.nmsRadius;
|
|
if (e < 0 || e > 1) throw new Error("segmentationThreshold " + e + ". Should be in range [0.0, 1.0]");
|
|
if (n <= 0) throw new Error("Invalid maxDetections " + n + ". Should be > 0");
|
|
if (r < 0 || r > 1) throw new Error("Invalid scoreThreshold " + r + ". Should be in range [0.0, 1.0]");
|
|
if (i <= 0) throw new Error("Invalid nmsRadius " + i + ".");
|
|
}
|
|
function Nt(t2) {
|
|
var e = t2.segmentationThreshold, n = t2.maxDetections, r = t2.scoreThreshold, i = t2.nmsRadius, o = t2.minKeypointScore, a = t2.refineSteps;
|
|
if (e < 0 || e > 1) throw new Error("segmentationThreshold " + e + ". Should be in range [0.0, 1.0]");
|
|
if (n <= 0) throw new Error("Invalid maxDetections " + n + ". Should be > 0");
|
|
if (r < 0 || r > 1) throw new Error("Invalid scoreThreshold " + r + ". Should be in range [0.0, 1.0]");
|
|
if (i <= 0) throw new Error("Invalid nmsRadius " + i + ".");
|
|
if (o < 0 || o > 1) throw new Error("Invalid minKeypointScore " + o + ".Should be in range [0.0, 1.0]");
|
|
if (a <= 0 || a > 20) throw new Error("Invalid refineSteps " + a + ".Should be in range [1, 20]");
|
|
}
|
|
var Qt = (function() {
|
|
function t2(t3) {
|
|
this.baseModel = t3;
|
|
}
|
|
return t2.prototype.predictForPersonSegmentation = function(t3) {
|
|
var e = this.baseModel.predict(t3);
|
|
return { segmentLogits: e.segmentation, heatmapScores: e.heatmapScores, offsets: e.offsets, displacementFwd: e.displacementFwd, displacementBwd: e.displacementBwd };
|
|
}, t2.prototype.predictForPersonSegmentationAndPart = function(t3) {
|
|
var e = this.baseModel.predict(t3);
|
|
return { segmentLogits: e.segmentation, partHeatmapLogits: e.partHeatmaps, heatmapScores: e.heatmapScores, offsets: e.offsets, displacementFwd: e.displacementFwd, displacementBwd: e.displacementBwd };
|
|
}, t2.prototype.predictForMultiPersonInstanceSegmentationAndPart = function(t3) {
|
|
var e = this.baseModel.predict(t3);
|
|
return { segmentLogits: e.segmentation, longOffsets: e.longOffsets, heatmapScores: e.heatmapScores, offsets: e.offsets, displacementFwd: e.displacementFwd, displacementBwd: e.displacementBwd, partHeatmaps: e.partHeatmaps };
|
|
}, t2.prototype.segmentPersonActivation = function(t3, e, n) {
|
|
var i = this;
|
|
void 0 === n && (n = 0.5);
|
|
var o = Ot(t3), a = o[0], s = o[1], u = Rt(e, this.baseModel.outputStride, [a, s]), f = Ct(t3, u), l = f.resized, c = f.padding, d = tidy((function() {
|
|
var t4 = i.predictForPersonSegmentation(l), e2 = t4.segmentLogits, r = t4.heatmapScores, o2 = t4.offsets, u2 = t4.displacementFwd, f2 = t4.displacementBwd, d2 = l.shape, h2 = d2[0], p3 = d2[1], m2 = Ft(e2, [a, s], [h2, p3], [[c.top, c.bottom], [c.left, c.right]]);
|
|
return { segmentation: N(squeeze(m2), n), heatmapScores: r, offsets: o2, displacementFwd: u2, displacementBwd: f2 };
|
|
})), h = d.segmentation, p2 = d.heatmapScores, m = d.offsets, v = d.displacementFwd, w = d.displacementBwd;
|
|
return l.dispose(), { segmentation: h, heatmapScores: p2, offsets: m, displacementFwd: v, displacementBwd: w, padding: c, internalResolutionHeightAndWidth: u };
|
|
}, t2.prototype.segmentPerson = function(t3, e) {
|
|
return void 0 === e && (e = Kt), D(this, void 0, void 0, (function() {
|
|
var n, r, i, o, a, s, u, f, l, c, d, h, p2, m, g, v, w, y;
|
|
return B(this, (function(b) {
|
|
switch (b.label) {
|
|
case 0:
|
|
return qt(e = C(C({}, Kt), e)), n = this.segmentPersonActivation(t3, e.internalResolution, e.segmentationThreshold), r = n.segmentation, i = n.heatmapScores, o = n.offsets, a = n.displacementFwd, s = n.displacementBwd, u = n.padding, f = n.internalResolutionHeightAndWidth, l = r.shape, c = l[0], d = l[1], [4, r.data()];
|
|
case 1:
|
|
return h = b.sent(), r.dispose(), [4, Dt([i, o, a, s])];
|
|
case 2:
|
|
return p2 = b.sent(), m = p2[0], g = p2[1], v = p2[2], w = p2[3], y = Bt(y = Mt(m, g, v, w, this.baseModel.outputStride, e.maxDetections, e.scoreThreshold, e.nmsRadius), [c, d], f, u, false), i.dispose(), o.dispose(), a.dispose(), s.dispose(), [2, { height: c, width: d, data: h, allPoses: y }];
|
|
}
|
|
}));
|
|
}));
|
|
}, t2.prototype.segmentMultiPerson = function(t3, e) {
|
|
return void 0 === e && (e = Ut), D(this, void 0, void 0, (function() {
|
|
var n, i, o, a, s, u, f, l, c, d, h, p2, m, v, w, y, b, S, x, M, k, E = this;
|
|
return B(this, (function(T) {
|
|
switch (T.label) {
|
|
case 0:
|
|
return Nt(e = C(C({}, Ut), e)), n = Ot(t3), i = n[0], o = n[1], a = Rt(e.internalResolution, this.baseModel.outputStride, [i, o]), s = Ct(t3, a), u = s.resized, f = s.padding, l = tidy((function() {
|
|
var t4, n2 = E.predictForMultiPersonInstanceSegmentationAndPart(u), r = n2.segmentLogits, s2 = n2.longOffsets, l2 = n2.heatmapScores, c2 = n2.offsets, d2 = n2.displacementFwd, h2 = n2.displacementBwd, p3 = Ft(r, [i, o], a, [[f.top, f.bottom], [f.left, f.right]]);
|
|
return t4 = s2, { segmentation: N(squeeze(p3), e.segmentationThreshold), longOffsets: t4, heatmapScoresRaw: l2, offsetsRaw: c2, displacementFwdRaw: d2, displacementBwdRaw: h2 };
|
|
})), c = l.segmentation, d = l.longOffsets, h = l.heatmapScoresRaw, p2 = l.offsetsRaw, m = l.displacementFwdRaw, v = l.displacementBwdRaw, [4, Dt([h, p2, m, v])];
|
|
case 1:
|
|
return w = T.sent(), y = w[0], b = w[1], S = w[2], x = w[3], M = Bt(M = Mt(y, b, S, x, this.baseModel.outputStride, e.maxDetections, e.scoreThreshold, e.nmsRadius), [i, o], a, f, false), [4, lt(c, d, M, i, o, this.baseModel.outputStride, a, f, e.scoreThreshold, e.refineSteps, e.minKeypointScore, e.maxDetections)];
|
|
case 2:
|
|
return k = T.sent(), u.dispose(), c.dispose(), d.dispose(), h.dispose(), p2.dispose(), m.dispose(), v.dispose(), [2, k];
|
|
}
|
|
}));
|
|
}));
|
|
}, t2.prototype.segmentPersonPartsActivation = function(t3, e, n) {
|
|
var i = this;
|
|
void 0 === n && (n = 0.5);
|
|
var o = Ot(t3), a = o[0], s = o[1], u = Rt(e, this.baseModel.outputStride, [a, s]), f = Ct(t3, u), l = f.resized, c = f.padding, d = tidy((function() {
|
|
var t4 = i.predictForPersonSegmentationAndPart(l), e2 = t4.segmentLogits, r = t4.partHeatmapLogits, o2 = t4.heatmapScores, u2 = t4.offsets, f2 = t4.displacementFwd, d2 = t4.displacementBwd, h2 = l.shape, p3 = h2[0], m2 = h2[1], v2 = Ft(e2, [a, s], [p3, m2], [[c.top, c.bottom], [c.left, c.right]]), w2 = Ft(r, [a, s], [p3, m2], [[c.top, c.bottom], [c.left, c.right]]);
|
|
return { partSegmentation: Q(N(squeeze(v2), n), w2), heatmapScores: o2, offsets: u2, displacementFwd: f2, displacementBwd: d2 };
|
|
})), h = d.partSegmentation, p2 = d.heatmapScores, m = d.offsets, v = d.displacementFwd, w = d.displacementBwd;
|
|
return l.dispose(), { partSegmentation: h, heatmapScores: p2, offsets: m, displacementFwd: v, displacementBwd: w, padding: c, internalResolutionHeightAndWidth: u };
|
|
}, t2.prototype.segmentPersonParts = function(t3, e) {
|
|
return void 0 === e && (e = Kt), D(this, void 0, void 0, (function() {
|
|
var n, r, i, o, a, s, u, f, l, c, d, h, p2, m, g, v, w, y;
|
|
return B(this, (function(b) {
|
|
switch (b.label) {
|
|
case 0:
|
|
return qt(e = C(C({}, Kt), e)), n = this.segmentPersonPartsActivation(t3, e.internalResolution, e.segmentationThreshold), r = n.partSegmentation, i = n.heatmapScores, o = n.offsets, a = n.displacementFwd, s = n.displacementBwd, u = n.padding, f = n.internalResolutionHeightAndWidth, l = r.shape, c = l[0], d = l[1], [4, r.data()];
|
|
case 1:
|
|
return h = b.sent(), r.dispose(), [4, Dt([i, o, a, s])];
|
|
case 2:
|
|
return p2 = b.sent(), m = p2[0], g = p2[1], v = p2[2], w = p2[3], y = Bt(y = Mt(m, g, v, w, this.baseModel.outputStride, e.maxDetections, e.scoreThreshold, e.nmsRadius), [c, d], f, u, false), i.dispose(), o.dispose(), a.dispose(), s.dispose(), [2, { height: c, width: d, data: h, allPoses: y }];
|
|
}
|
|
}));
|
|
}));
|
|
}, t2.prototype.segmentMultiPersonParts = function(t3, e) {
|
|
return void 0 === e && (e = Ut), D(this, void 0, void 0, (function() {
|
|
var n, o, a, s, d, h, p2, m, v, w, y, b, S, x, M, k, E, T, P, I, O, _, H = this;
|
|
return B(this, (function(A) {
|
|
switch (A.label) {
|
|
case 0:
|
|
return Nt(e = C(C({}, Ut), e)), n = Ot(t3), o = n[0], a = n[1], s = Rt(e.internalResolution, this.baseModel.outputStride, [o, a]), d = Ct(t3, s), h = d.resized, p2 = d.padding, m = tidy((function() {
|
|
var t4 = H.predictForMultiPersonInstanceSegmentationAndPart(h), n2 = t4.segmentLogits, d2 = t4.longOffsets, m2 = t4.heatmapScores, v2 = t4.offsets, w2 = t4.displacementFwd, y2 = t4.displacementBwd, b2 = t4.partHeatmaps, S2 = Ft(n2, [o, a], s, [[p2.top, p2.bottom], [p2.left, p2.right]]), x2 = Ft(b2, [o, a], s, [[p2.top, p2.bottom], [p2.left, p2.right]]), M2 = d2, k3 = N(squeeze(S2), e.segmentationThreshold), E2 = (function(t5) {
|
|
var e2 = t5.shape, n3 = e2[0], o2 = e2[1], a6 = e2[2];
|
|
return tidy((function() {
|
|
var e3 = q(t5), r = expandDims$2(range$2(0, a6, 1, "int32"), 1), s2 = cast$2(matMul$1(e3, r), "int32");
|
|
return reshape$2(s2, [n3, o2]);
|
|
}));
|
|
})(x2);
|
|
return { segmentation: k3, longOffsets: M2, heatmapScoresRaw: m2, offsetsRaw: v2, displacementFwdRaw: w2, displacementBwdRaw: y2, partSegmentation: E2 };
|
|
})), v = m.segmentation, w = m.longOffsets, y = m.heatmapScoresRaw, b = m.offsetsRaw, S = m.displacementFwdRaw, x = m.displacementBwdRaw, M = m.partSegmentation, [4, Dt([y, b, S, x])];
|
|
case 1:
|
|
return k = A.sent(), E = k[0], T = k[1], P = k[2], I = k[3], O = Bt(O = Mt(E, T, P, I, this.baseModel.outputStride, e.maxDetections, e.scoreThreshold, e.nmsRadius), [o, a], s, p2, false), [4, ct(v, w, M, O, o, a, this.baseModel.outputStride, s, p2, e.scoreThreshold, e.refineSteps, e.minKeypointScore, e.maxDetections)];
|
|
case 2:
|
|
return _ = A.sent(), h.dispose(), v.dispose(), w.dispose(), y.dispose(), b.dispose(), S.dispose(), x.dispose(), M.dispose(), [2, _];
|
|
}
|
|
}));
|
|
}));
|
|
}, t2.prototype.dispose = function() {
|
|
this.baseModel.dispose();
|
|
}, t2;
|
|
})();
|
|
function Xt(e) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var n, r, i, o, a, s;
|
|
return B(this, (function(u) {
|
|
switch (u.label) {
|
|
case 0:
|
|
if (n = e.outputStride, r = e.quantBytes, i = e.multiplier, null == t) throw new Error("Cannot find TensorFlow.js. If you are using a <script> tag, please also include @tensorflow/tfjs on the page before using this\n model.");
|
|
return o = (function(t2, e2, n2) {
|
|
var r2 = { 1: "100", 0.75: "075", 0.5: "050" }, i2 = "model-stride" + t2 + ".json";
|
|
return 4 === n2 ? It + "float/" + r2[e2] + "/" + i2 : It + "quant" + n2 + "/" + r2[e2] + "/" + i2;
|
|
})(n, i, r), [4, loadGraphModel(e.modelUrl || o)];
|
|
case 1:
|
|
return a = u.sent(), s = new Y(a, n), [2, new Qt(s)];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function Yt(e) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var n, r, i, o, a;
|
|
return B(this, (function(s) {
|
|
switch (s.label) {
|
|
case 0:
|
|
if (n = e.outputStride, r = e.quantBytes, null == t) throw new Error("Cannot find TensorFlow.js. If you are using a <script> tag, please also include @tensorflow/tfjs on the page before using this\n model.");
|
|
return i = (function(t2, e2) {
|
|
var n2 = "model-stride" + t2 + ".json";
|
|
return 4 === e2 ? Pt + "float/" + n2 : Pt + "quant" + e2 + "/" + n2;
|
|
})(n, r), [4, loadGraphModel(e.modelUrl || i)];
|
|
case 1:
|
|
return o = s.sent(), a = new Tt(o, n), [2, new Qt(a)];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function Gt(t2) {
|
|
return void 0 === t2 && (t2 = Lt), D(this, void 0, void 0, (function() {
|
|
return B(this, (function(e) {
|
|
return "ResNet50" === (t2 = (function(t3) {
|
|
if (null == (t3 = t3 || Lt).architecture && (t3.architecture = "MobileNetV1"), zt.indexOf(t3.architecture) < 0) throw new Error("Invalid architecture " + t3.architecture + ". Should be one of " + zt);
|
|
if (null == t3.outputStride && (t3.outputStride = 16), Wt[t3.architecture].indexOf(t3.outputStride) < 0) throw new Error("Invalid outputStride " + t3.outputStride + ". Should be one of " + Wt[t3.architecture] + " for architecture " + t3.architecture + ".");
|
|
if (null == t3.multiplier && (t3.multiplier = 1), jt[t3.architecture].indexOf(t3.multiplier) < 0) throw new Error("Invalid multiplier " + t3.multiplier + ". Should be one of " + jt[t3.architecture] + " for architecture " + t3.architecture + ".");
|
|
if (null == t3.quantBytes && (t3.quantBytes = 4), Vt.indexOf(t3.quantBytes) < 0) throw new Error("Invalid quantBytes " + t3.quantBytes + ". Should be one of " + Vt + " for architecture " + t3.architecture + ".");
|
|
return t3;
|
|
})(t2)).architecture ? [2, Yt(t2)] : "MobileNetV1" === t2.architecture ? [2, Xt(t2)] : [2, null];
|
|
}));
|
|
}));
|
|
}
|
|
var $t = ["left_face", "right_face", "left_upper_arm_front", "left_upper_arm_back", "right_upper_arm_front", "right_upper_arm_back", "left_lower_arm_front", "left_lower_arm_back", "right_lower_arm_front", "right_lower_arm_back", "left_hand", "right_hand", "torso_front", "torso_back", "left_upper_leg_front", "left_upper_leg_back", "right_upper_leg_front", "right_upper_leg_back", "left_lower_leg_front", "left_lower_leg_back", "right_lower_leg_front", "right_lower_leg_back", "left_feet", "right_feet"], Jt = (function() {
|
|
function t2(t3) {
|
|
this.mask = t3;
|
|
}
|
|
return t2.prototype.toCanvasImageSource = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, z(this.mask)];
|
|
}));
|
|
}));
|
|
}, t2.prototype.toImageData = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, this.mask];
|
|
}));
|
|
}));
|
|
}, t2.prototype.toTensor = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, j(this.mask)];
|
|
}));
|
|
}));
|
|
}, t2.prototype.getUnderlyingType = function() {
|
|
return "imagedata";
|
|
}, t2;
|
|
})();
|
|
function Zt(t2) {
|
|
if (V(t2), 255 !== t2) throw new Error("Foreground id must be 255 but got " + t2);
|
|
return "person";
|
|
}
|
|
function te(t2) {
|
|
if (V(t2), t2 >= $t.length) throw new Error("Invalid body part value " + t2);
|
|
return $t[t2];
|
|
}
|
|
var ee = (function() {
|
|
function t2(t3) {
|
|
this.bodyPixModel = t3;
|
|
}
|
|
return t2.prototype.segmentPeople = function(t3, e) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var n, r, i, o;
|
|
return B(this, (function(a) {
|
|
switch (a.label) {
|
|
case 0:
|
|
return t3 instanceof ImageBitmap && ((n = document.createElement("canvas")).getContext("2d").drawImage(t3, 0, 0), t3 = n), e.segmentBodyParts ? e.multiSegmentation ? [4, this.bodyPixModel.segmentMultiPersonParts(t3, e)] : [3, 2] : [3, 5];
|
|
case 1:
|
|
return i = a.sent(), [3, 4];
|
|
case 2:
|
|
return [4, this.bodyPixModel.segmentPersonParts(t3, e)];
|
|
case 3:
|
|
i = [a.sent()], a.label = 4;
|
|
case 4:
|
|
return r = i.map((function(t4) {
|
|
var e2 = t4.data, n2 = t4.width, r2 = t4.height, i2 = new Uint8ClampedArray(n2 * r2 * 4).fill(0);
|
|
return e2.forEach((function(t5, e3) {
|
|
-1 === t5 ? (i2[4 * e3] = $t.length, i2[4 * e3 + 3] = 0) : (i2[4 * e3] = t5, i2[4 * e3 + 3] = 255);
|
|
})), { maskValueToLabel: te, mask: new Jt(new ImageData(i2, n2, r2)) };
|
|
})), [3, 10];
|
|
case 5:
|
|
return e.multiSegmentation ? [4, this.bodyPixModel.segmentMultiPerson(t3, e)] : [3, 7];
|
|
case 6:
|
|
return o = a.sent(), [3, 9];
|
|
case 7:
|
|
return [4, this.bodyPixModel.segmentPerson(t3, e)];
|
|
case 8:
|
|
o = [a.sent()], a.label = 9;
|
|
case 9:
|
|
r = o.map((function(t4) {
|
|
var e2 = t4.data, n2 = t4.width, r2 = t4.height, i2 = new Uint8ClampedArray(n2 * r2 * 4).fill(0);
|
|
return e2.forEach((function(t5, e3) {
|
|
0 === t5 ? (i2[4 * e3] = 0, i2[4 * e3 + 3] = 0) : (i2[4 * e3] = 255, i2[4 * e3 + 3] = 255);
|
|
})), { maskValueToLabel: Zt, mask: new Jt(new ImageData(i2, n2, r2)) };
|
|
})), a.label = 10;
|
|
case 10:
|
|
return [2, r];
|
|
}
|
|
}));
|
|
}));
|
|
}, t2.prototype.dispose = function() {
|
|
this.bodyPixModel.dispose();
|
|
}, t2.prototype.reset = function() {
|
|
}, t2;
|
|
})();
|
|
function ne(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(e) {
|
|
return [2, Gt(t2).then((function(t3) {
|
|
return new ee(t3);
|
|
}))];
|
|
}));
|
|
}));
|
|
}
|
|
var re = { runtime: "mediapipe", modelType: "general" };
|
|
var ie = (function() {
|
|
function t2(t3) {
|
|
this.mask = t3;
|
|
}
|
|
return t2.prototype.toCanvasImageSource = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, this.mask];
|
|
}));
|
|
}));
|
|
}, t2.prototype.toImageData = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, W(this.mask)];
|
|
}));
|
|
}));
|
|
}, t2.prototype.toTensor = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, j(this.mask)];
|
|
}));
|
|
}));
|
|
}, t2.prototype.getUnderlyingType = function() {
|
|
return "canvasimagesource";
|
|
}, t2;
|
|
})();
|
|
function oe(t2) {
|
|
return V(t2), "person";
|
|
}
|
|
var ae = (function() {
|
|
function t2(t3) {
|
|
var e, n = this;
|
|
this.selfieMode = false;
|
|
var r;
|
|
if (this.selfieSegmentationSolution = new selfie_segmentationExports.SelfieSegmentation({ locateFile: null !== (e = t3.locateFile) && void 0 !== e ? e : function(e2, n2) {
|
|
return t3.solutionPath ? t3.solutionPath.replace(/\/+$/, "") + "/" + e2 : n2 + "/" + e2;
|
|
} }), "landscape" === t3.modelType) r = 1;
|
|
else r = 0;
|
|
this.selfieSegmentationSolution.setOptions({ modelSelection: r, selfieMode: this.selfieMode }), this.selfieSegmentationSolution.onResults((function(t4) {
|
|
n.segmentation = [{ maskValueToLabel: oe, mask: new ie(t4.segmentationMask) }];
|
|
}));
|
|
}
|
|
return t2.prototype.segmentPeople = function(t3, r) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var i, o;
|
|
return B(this, (function(a) {
|
|
switch (a.label) {
|
|
case 0:
|
|
return r && r.flipHorizontal && r.flipHorizontal !== this.selfieMode && (this.selfieMode = r.flipHorizontal, this.selfieSegmentationSolution.setOptions({ selfieMode: this.selfieMode })), t3 instanceof Tensor ? (o = ImageData.bind, [4, toPixels(t3)]) : [3, 2];
|
|
case 1:
|
|
return i = new (o.apply(ImageData, [void 0, a.sent(), t3.shape[1], t3.shape[0]]))(), [3, 3];
|
|
case 2:
|
|
i = t3, a.label = 3;
|
|
case 3:
|
|
return t3 = i, [4, this.selfieSegmentationSolution.send({ image: t3 })];
|
|
case 4:
|
|
return a.sent(), [2, this.segmentation];
|
|
}
|
|
}));
|
|
}));
|
|
}, t2.prototype.dispose = function() {
|
|
this.selfieSegmentationSolution.close();
|
|
}, t2.prototype.reset = function() {
|
|
this.selfieSegmentationSolution.reset(), this.segmentation = null, this.selfieMode = false;
|
|
}, t2.prototype.initialize = function() {
|
|
return this.selfieSegmentationSolution.initialize();
|
|
}, t2;
|
|
})();
|
|
function se(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var e, n;
|
|
return B(this, (function(r) {
|
|
switch (r.label) {
|
|
case 0:
|
|
return e = (function(t3) {
|
|
if (null == t3) return C({}, re);
|
|
var e2 = C({}, t3);
|
|
return e2.runtime = "mediapipe", null == e2.modelType && (e2.modelType = re.modelType), e2;
|
|
})(t2), [4, (n = new ae(e)).initialize()];
|
|
case 1:
|
|
return r.sent(), [2, n];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function ue(t2, e, n, r) {
|
|
var i = t2.width, o = t2.height, a = 1, s = Math.cos(t2.rotation), u = Math.sin(t2.rotation), f = t2.xCenter, l = t2.yCenter, c = 1 / e, d = 1 / n, h = new Array(16);
|
|
return h[0] = i * s * a * c, h[1] = -o * u * c, h[2] = 0, h[3] = (-0.5 * i * s * a + 0.5 * o * u + f) * c, h[4] = i * u * a * d, h[5] = o * s * d, h[6] = 0, h[7] = (-0.5 * o * s - 0.5 * i * u * a + l) * d, h[8] = 0, h[9] = 0, h[10] = i * c, h[11] = 0, h[12] = 0, h[13] = 0, h[14] = 0, h[15] = 1, (function(t3) {
|
|
if (16 !== t3.length) throw new Error("Array length must be 16 but got " + t3.length);
|
|
return [[t3[0], t3[1], t3[2], t3[3]], [t3[4], t3[5], t3[6], t3[7]], [t3[8], t3[9], t3[10], t3[11]], [t3[12], t3[13], t3[14], t3[15]]];
|
|
})(h);
|
|
}
|
|
function fe(t2) {
|
|
return t2 instanceof Tensor ? { height: t2.shape[0], width: t2.shape[1] } : { height: t2.height, width: t2.width };
|
|
}
|
|
function le(t2, e) {
|
|
assert(0 !== t2.width, (function() {
|
|
return e + " width cannot be 0.";
|
|
})), assert(0 !== t2.height, (function() {
|
|
return e + " height cannot be 0.";
|
|
}));
|
|
}
|
|
function ce(t2, e) {
|
|
var n = (function(t3, e2, n2, r) {
|
|
var i = e2 - t3, o = r - n2;
|
|
var a = o / i;
|
|
return { scale: a, offset: n2 - t3 * a };
|
|
})(0, 255, e[0], e[1]);
|
|
return tidy((function() {
|
|
return add$1(mul(t2, n.scale), n.offset);
|
|
}));
|
|
}
|
|
function de(t2, o, a) {
|
|
var s = o.outputTensorSize, c = o.outputTensorFloatRange, d = fe(t2), h = (function(t3, e) {
|
|
return { xCenter: 0.5 * t3.width, yCenter: 0.5 * t3.height, width: t3.width, height: t3.height, rotation: 0 };
|
|
})(d), p2 = /* @__PURE__ */ (function(t3, e, n) {
|
|
return { top: 0, left: 0, right: 0, bottom: 0 };
|
|
})(), m = ue(h, d.width, d.height), g = tidy((function() {
|
|
var r, o2 = (r = t2) instanceof Tensor ? r : fromPixels$1(r), a6 = tensor2d((function(t3, e, n) {
|
|
return le(n, "inputResolution"), [1 / n.width * t3[0][0] * e.width, 1 / n.height * t3[0][1] * e.width, t3[0][3] * e.width, 1 / n.width * t3[1][0] * e.height, 1 / n.height * t3[1][1] * e.height, t3[1][3] * e.height, 0, 0];
|
|
})(m, d, s), [1, 8]), f = "constant", h2 = image$1.transform(expandDims$2(cast$2(o2, "float32")), a6, "bilinear", f, 0, [s.height, s.width]);
|
|
return null != c ? ce(h2, c) : h2;
|
|
}));
|
|
return { imageTensor: g, padding: p2, transformationMatrix: m };
|
|
}
|
|
function he(t2, e, n) {
|
|
return tidy((function() {
|
|
var r = squeeze(t2, [0]), i = r.shape[2];
|
|
if (1 === i) {
|
|
var o = r;
|
|
switch (e.activation) {
|
|
case "none":
|
|
break;
|
|
case "sigmoid":
|
|
o = sigmoid$2(o);
|
|
break;
|
|
case "softmax":
|
|
throw new Error("Softmax activation requires two channels.");
|
|
default:
|
|
throw new Error("Activation not supported (" + e.activation + ")");
|
|
}
|
|
var a = n ? image$1.resizeBilinear(o, [n.height, n.width]) : o;
|
|
return squeeze(a, [2]);
|
|
}
|
|
throw new Error("Unsupported number of tensor channels " + i);
|
|
}));
|
|
}
|
|
var pe = { runtime: "tfjs", modelType: "general", modelUrl: "https://tfhub.dev/mediapipe/tfjs-model/selfie_segmentation/general/1" }, me = { flipHorizontal: false }, ge = { outputTensorSize: { width: 256, height: 256 }, outputTensorFloatRange: [0, 1] }, ve = { outputTensorSize: { width: 256, height: 144 }, outputTensorFloatRange: [0, 1] }, we = { activation: "none" };
|
|
var ye = (function() {
|
|
function t2(t3) {
|
|
this.mask = t3;
|
|
}
|
|
return t2.prototype.toCanvasImageSource = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, z(this.mask)];
|
|
}));
|
|
}));
|
|
}, t2.prototype.toImageData = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, W(this.mask)];
|
|
}));
|
|
}));
|
|
}, t2.prototype.toTensor = function() {
|
|
return D(this, void 0, void 0, (function() {
|
|
return B(this, (function(t3) {
|
|
return [2, this.mask];
|
|
}));
|
|
}));
|
|
}, t2.prototype.getUnderlyingType = function() {
|
|
return "tensor";
|
|
}, t2;
|
|
})();
|
|
function be(t2) {
|
|
return V(t2), "person";
|
|
}
|
|
var Se, xe = (function() {
|
|
function t2(t3, e) {
|
|
this.modelType = t3, this.model = e;
|
|
}
|
|
return t2.prototype.segmentPeople = function(t3, e) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var n, i = this;
|
|
return B(this, (function(o) {
|
|
return e = (function(t4) {
|
|
if (null == t4) return C({}, me);
|
|
var e2 = C({}, t4);
|
|
return null == e2.flipHorizontal && (e2.flipHorizontal = me.flipHorizontal), e2;
|
|
})(e), null == t3 ? (this.reset(), [2, []]) : (n = tidy((function() {
|
|
var e2 = de(t3, "general" === i.modelType ? ge : ve).imageTensor, n2 = slice$2(i.model.predict(e2), [0, 0, 0, 1], -1), r = fe(t3), o2 = he(n2, we, r), a = expandDims$2(o2, 2), s = pad(a, [[0, 0], [0, 0], [0, 1]]);
|
|
return mirrorPad$1(s, [[0, 0], [0, 0], [0, 2]], "symmetric");
|
|
})), [2, [{ maskValueToLabel: be, mask: new ye(n) }]]);
|
|
}));
|
|
}));
|
|
}, t2.prototype.dispose = function() {
|
|
this.model.dispose();
|
|
}, t2.prototype.reset = function() {
|
|
}, t2;
|
|
})();
|
|
function Me(t2) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var e, n, r;
|
|
return B(this, (function(i) {
|
|
switch (i.label) {
|
|
case 0:
|
|
return e = (function(t3) {
|
|
if (null == t3) return C({}, pe);
|
|
var e2 = C({}, t3);
|
|
if (e2.runtime = "tfjs", null == e2.modelType && (e2.modelType = pe.modelType), "general" !== e2.modelType && "landscape" !== e2.modelType) throw new Error("Model type must be one of general or landscape, but got " + e2.modelType);
|
|
null == e2.modelUrl && ("general" === e2.modelType ? e2.modelUrl = "https://tfhub.dev/mediapipe/tfjs-model/selfie_segmentation/general/1" : e2.modelUrl = "https://tfhub.dev/mediapipe/tfjs-model/selfie_segmentation/landscape/1");
|
|
return e2;
|
|
})(t2), n = "string" == typeof e.modelUrl && e.modelUrl.indexOf("https://tfhub.dev") > -1, [4, loadGraphModel(e.modelUrl, { fromTFHub: n })];
|
|
case 1:
|
|
return r = i.sent(), [2, new xe(e.modelType, r)];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function ke(t2, e) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var n, r;
|
|
return B(this, (function(i) {
|
|
switch (t2) {
|
|
case Se.MediaPipeSelfieSegmentation:
|
|
if (n = void 0, null != (r = e)) {
|
|
if ("tfjs" === r.runtime) return [2, Me(r)];
|
|
if ("mediapipe" === r.runtime) return [2, se(r)];
|
|
n = r.runtime;
|
|
}
|
|
throw new Error("Expect modelConfig.runtime to be either 'tfjs' or 'mediapipe', but got " + n);
|
|
case Se.BodyPix:
|
|
return [2, ne(r = e)];
|
|
default:
|
|
throw new Error(t2 + " is not a supported model name.");
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
!(function(t2) {
|
|
t2.BodyPix = "BodyPix", t2.MediaPipeSelfieSegmentation = "MediaPipeSelfieSegmentation";
|
|
})(Se || (Se = {}));
|
|
var Te = "blurred-mask", Pe = "mask", Oe = "draw-image", _e = {};
|
|
function He(t2, e, n, r) {
|
|
var i = t2.width, o = t2.height, a = e.width, s = e.height;
|
|
if (i !== a || o !== s) throw new Error("error: dimensions must match. " + n + " has dimensions " + i + "x" + o + ", " + r + " has dimensions " + a + "x" + s);
|
|
}
|
|
function Ae(t2) {
|
|
if ("undefined" != typeof HTMLCanvasElement && t2 instanceof HTMLCanvasElement || "undefined" != typeof OffscreenCanvas && t2 instanceof OffscreenCanvas || "undefined" != typeof HTMLImageElement && t2 instanceof HTMLImageElement) return (function(t3) {
|
|
if ("offsetHeight" in t3 && 0 !== t3.offsetHeight && "offsetWidth" in t3 && 0 !== t3.offsetWidth) return [t3.offsetHeight, t3.offsetWidth];
|
|
if (null != t3.height && null != t3.width) return [t3.height, t3.width];
|
|
throw new Error("HTMLImageElement must have height and width attributes set.");
|
|
})(t2);
|
|
if ("undefined" != typeof ImageData && t2 instanceof ImageData) return [t2.height, t2.width];
|
|
if ("undefined" != typeof HTMLVideoElement && t2 instanceof HTMLVideoElement) return (function(t3) {
|
|
return t3.hasAttribute("height") && t3.hasAttribute("width") ? [t3.height, t3.width] : [t3.videoHeight, t3.videoWidth];
|
|
})(t2);
|
|
if (t2 instanceof Tensor) return [t2.shape[0], t2.shape[1]];
|
|
throw new Error("error: Unknown input type: " + t2 + ".");
|
|
}
|
|
function Re(t2) {
|
|
return _e[t2] || (_e[t2] = (function() {
|
|
if ("undefined" != typeof document) return document.createElement("canvas");
|
|
if ("undefined" != typeof OffscreenCanvas) return new OffscreenCanvas(0, 0);
|
|
throw new Error("Cannot create a canvas in this context");
|
|
})()), _e[t2];
|
|
}
|
|
function Fe(t2, e) {
|
|
var n = Re(e);
|
|
return (function(t3, e2) {
|
|
e2.width = t3.width, e2.height = t3.height, e2.getContext("2d").putImageData(t3, 0, 0);
|
|
})(t2, n), n;
|
|
}
|
|
function Ce(t2, r, i, o, a, s) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var u, f, l, c;
|
|
return B(this, (function(d) {
|
|
switch (d.label) {
|
|
case 0:
|
|
return r instanceof Tensor ? [4, toPixels(r)] : [3, 2];
|
|
case 1:
|
|
u = d.sent(), f = Ae(r), l = f[0], c = f[1], r = new ImageData(u, c, l), d.label = 2;
|
|
case 2:
|
|
return r instanceof ImageData && (r = Fe(r, Oe)), null == a || null == s ? t2.drawImage(r, i, o) : t2.drawImage(r, i, o, a, s), [2];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function De(t2, e) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var n, r, i;
|
|
return B(this, (function(o) {
|
|
switch (o.label) {
|
|
case 0:
|
|
return n = Ae(t2), r = n[0], i = n[1], e.width = i, e.height = r, [4, Ce(e.getContext("2d"), t2, 0, 0, i, r)];
|
|
case 1:
|
|
return o.sent(), [2];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function Be(t2) {
|
|
var e = t2.getContext("2d");
|
|
e.scale(-1, 1), e.translate(-t2.width, 0);
|
|
}
|
|
function ze(t2, e, n) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var r, i, o, a, s, u, f, l;
|
|
return B(this, (function(c) {
|
|
switch (c.label) {
|
|
case 0:
|
|
for (r = t2.getContext("2d"), i = 0, o = 5, a = 1 / (2 * Math.PI * o * o), s = n < 3 ? 1 : 2, f = -n; f <= n; f += s) for (l = -n; l <= n; l += s) u = a * Math.exp(-(l * l + f * f) / (2 * o * o)), i += u;
|
|
f = -n, c.label = 1;
|
|
case 1:
|
|
if (!(f <= n)) return [3, 6];
|
|
l = -n, c.label = 2;
|
|
case 2:
|
|
return l <= n ? (r.globalAlpha = a * Math.exp(-(l * l + f * f) / (2 * o * o)) / i * n, [4, Ce(r, e, l, f)]) : [3, 5];
|
|
case 3:
|
|
c.sent(), c.label = 4;
|
|
case 4:
|
|
return l += s, [3, 2];
|
|
case 5:
|
|
return f += s, [3, 1];
|
|
case 6:
|
|
return r.globalAlpha = 1, [2];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function We(t2, e, n) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var r, i, o, a;
|
|
return B(this, (function(s) {
|
|
switch (s.label) {
|
|
case 0:
|
|
return r = Ae(t2), i = r[0], o = r[1], a = n.getContext("2d"), n.width = o, n.height = i, a.clearRect(0, 0, o, i), a.save(), /^((?!chrome|android).)*safari/i.test(navigator.userAgent) ? [4, ze(n, t2, e)] : [3, 2];
|
|
case 1:
|
|
return s.sent(), [3, 4];
|
|
case 2:
|
|
return a.filter = "blur(" + e + "px)", [4, Ce(a, t2, 0, 0, o, i)];
|
|
case 3:
|
|
s.sent(), s.label = 4;
|
|
case 4:
|
|
return a.restore(), [2];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function je(t2, e, n) {
|
|
return D(this, void 0, void 0, (function() {
|
|
var r;
|
|
return B(this, (function(i) {
|
|
switch (i.label) {
|
|
case 0:
|
|
return r = Re(n), 0 !== e ? [3, 2] : [4, De(t2, r)];
|
|
case 1:
|
|
return i.sent(), [3, 4];
|
|
case 2:
|
|
return [4, We(t2, e, r)];
|
|
case 3:
|
|
i.sent(), i.label = 4;
|
|
case 4:
|
|
return [2, r];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function Ve(t2, e, n, r, i, o) {
|
|
void 0 === o && (o = { r: 0, g: 255, b: 255, a: 255 });
|
|
for (var a = -i; a <= i; a++) for (var s = -i; s <= i; s++) if (0 !== a && 0 !== s) {
|
|
var u = (e + a) * r + (n + s);
|
|
t2[4 * u + 0] = o.r, t2[4 * u + 1] = o.g, t2[4 * u + 2] = o.b, t2[4 * u + 3] = o.a;
|
|
}
|
|
}
|
|
function Ke(t2, e, n, r, i, o, a) {
|
|
void 0 === a && (a = 1);
|
|
for (var s = 0, u = -a; u <= a; u++) for (var f = -a; f <= a; f++) if (0 !== u && 0 !== f) {
|
|
var l = (e + u) * r + (n + f);
|
|
(!i[t2[4 * l]] || t2[4 * l + 3] < o) && (s += 1);
|
|
}
|
|
return s > 0;
|
|
}
|
|
function Ue(t2, e, n, r, i, o) {
|
|
return void 0 === e && (e = { r: 0, g: 0, b: 0, a: 0 }), void 0 === n && (n = { r: 0, g: 0, b: 0, a: 255 }), void 0 === r && (r = false), void 0 === i && (i = 0.5), void 0 === o && (o = Array.from(Array(256).keys())), D(this, void 0, void 0, (function() {
|
|
var a, s, u, f, l, c, d, h, p2, m, g, v, w, y;
|
|
return B(this, (function(b) {
|
|
switch (b.label) {
|
|
case 0:
|
|
return 0 === (a = Array.isArray(t2) ? t2 : [t2]).length ? [2, null] : [4, Promise.all(a.map((function(t3) {
|
|
return t3.mask.toImageData();
|
|
})))];
|
|
case 1:
|
|
for (s = b.sent(), u = s[0], f = u.width, l = u.height, c = new Uint8ClampedArray(f * l * 4), d = Math.round(255 * i), h = new Array(256).fill(false), o.forEach((function(t3) {
|
|
return h[t3] = true;
|
|
})), p2 = 0; p2 < l; p2++) for (m = 0; m < f; m++) for (c[4 * (g = p2 * f + m) + 0] = n.r, c[4 * g + 1] = n.g, c[4 * g + 2] = n.b, c[4 * g + 3] = n.a, v = 0, w = s; v < w.length; v++) y = w[v], h[y.data[4 * g]] && y.data[4 * g + 3] >= d && (c[4 * g] = e.r, c[4 * g + 1] = e.g, c[4 * g + 2] = e.b, c[4 * g + 3] = e.a, r && p2 - 1 >= 0 && p2 + 1 < l && m - 1 >= 0 && m + 1 < f && Ke(y.data, p2, m, f, h, d) && Ve(c, p2, m, f, 1));
|
|
return [2, new ImageData(c, f, l)];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
function Ne(t2, e, n, r, i, o) {
|
|
return void 0 === r && (r = 0.7), void 0 === i && (i = 0), void 0 === o && (o = false), D(this, void 0, void 0, (function() {
|
|
var a, s, u, f, l;
|
|
return B(this, (function(c) {
|
|
switch (c.label) {
|
|
case 0:
|
|
return a = Ae(e), s = a[0], u = a[1], t2.width = u, t2.height = s, (f = t2.getContext("2d")).save(), o && Be(t2), [4, Ce(f, e, 0, 0)];
|
|
case 1:
|
|
return c.sent(), f.globalAlpha = r, n ? (He({ width: u, height: s }, n, "image", "mask"), [4, je(Fe(n, Pe), i, Te)]) : [3, 3];
|
|
case 2:
|
|
l = c.sent(), f.drawImage(l, 0, 0, u, s), c.label = 3;
|
|
case 3:
|
|
return f.restore(), [2];
|
|
}
|
|
}));
|
|
}));
|
|
}
|
|
const contexts = {};
|
|
const WEBGL_ATTRIBUTES = {
|
|
alpha: false,
|
|
antialias: false,
|
|
premultipliedAlpha: false,
|
|
preserveDrawingBuffer: false,
|
|
depth: false,
|
|
stencil: false,
|
|
failIfMajorPerformanceCaveat: true
|
|
};
|
|
function setWebGLContext(webGLVersion, gl) {
|
|
contexts[webGLVersion] = gl;
|
|
}
|
|
function getWebGLContext(webGLVersion, customCanvas) {
|
|
if (!(webGLVersion in contexts) || customCanvas != null) {
|
|
const newCtx = getWebGLRenderingContext(webGLVersion, customCanvas);
|
|
if (newCtx !== null) {
|
|
contexts[webGLVersion] = newCtx;
|
|
} else {
|
|
console.log("Could not get context for WebGL version", webGLVersion);
|
|
return null;
|
|
}
|
|
}
|
|
const gl = contexts[webGLVersion];
|
|
if (gl == null || gl.isContextLost()) {
|
|
delete contexts[webGLVersion];
|
|
return getWebGLContext(webGLVersion);
|
|
}
|
|
gl.disable(gl.DEPTH_TEST);
|
|
gl.disable(gl.STENCIL_TEST);
|
|
gl.disable(gl.BLEND);
|
|
gl.disable(gl.DITHER);
|
|
gl.disable(gl.POLYGON_OFFSET_FILL);
|
|
gl.disable(gl.SAMPLE_COVERAGE);
|
|
gl.enable(gl.SCISSOR_TEST);
|
|
gl.enable(gl.CULL_FACE);
|
|
gl.cullFace(gl.BACK);
|
|
return contexts[webGLVersion];
|
|
}
|
|
function createCanvas(webGLVersion) {
|
|
if (!env().getBool("IS_SAFARI") && typeof OffscreenCanvas !== "undefined" && webGLVersion === 2) {
|
|
return new OffscreenCanvas(300, 150);
|
|
} else if (typeof document !== "undefined") {
|
|
return document.createElement("canvas");
|
|
} else {
|
|
throw new Error("Cannot create a canvas in this context");
|
|
}
|
|
}
|
|
function getWebGLRenderingContext(webGLVersion, customCanvas) {
|
|
if (webGLVersion !== 1 && webGLVersion !== 2) {
|
|
throw new Error("Cannot get WebGL rendering context, WebGL is disabled.");
|
|
}
|
|
const canvas = customCanvas == null ? createCanvas(webGLVersion) : customCanvas;
|
|
canvas.addEventListener("webglcontextlost", (ev) => {
|
|
ev.preventDefault();
|
|
delete contexts[webGLVersion];
|
|
}, false);
|
|
if (env().getBool("SOFTWARE_WEBGL_ENABLED")) {
|
|
WEBGL_ATTRIBUTES.failIfMajorPerformanceCaveat = false;
|
|
}
|
|
if (webGLVersion === 1) {
|
|
return (
|
|
// tslint:disable-next-line
|
|
canvas.getContext("webgl", WEBGL_ATTRIBUTES) || canvas.getContext("experimental-webgl", WEBGL_ATTRIBUTES)
|
|
);
|
|
}
|
|
return canvas.getContext("webgl2", WEBGL_ATTRIBUTES);
|
|
}
|
|
var PackingScheme;
|
|
(function(PackingScheme2) {
|
|
PackingScheme2[PackingScheme2["DENSE"] = 0] = "DENSE";
|
|
PackingScheme2[PackingScheme2["SHARED_BATCH"] = 1] = "SHARED_BATCH";
|
|
})(PackingScheme || (PackingScheme = {}));
|
|
var TextureUsage;
|
|
(function(TextureUsage2) {
|
|
TextureUsage2[TextureUsage2["RENDER"] = 0] = "RENDER";
|
|
TextureUsage2[TextureUsage2["UPLOAD"] = 1] = "UPLOAD";
|
|
TextureUsage2[TextureUsage2["PIXELS"] = 2] = "PIXELS";
|
|
TextureUsage2[TextureUsage2["DOWNLOAD"] = 3] = "DOWNLOAD";
|
|
})(TextureUsage || (TextureUsage = {}));
|
|
var PhysicalTextureType;
|
|
(function(PhysicalTextureType2) {
|
|
PhysicalTextureType2[PhysicalTextureType2["UNPACKED_FLOAT16"] = 0] = "UNPACKED_FLOAT16";
|
|
PhysicalTextureType2[PhysicalTextureType2["UNPACKED_FLOAT32"] = 1] = "UNPACKED_FLOAT32";
|
|
PhysicalTextureType2[PhysicalTextureType2["PACKED_4X1_UNSIGNED_BYTE"] = 2] = "PACKED_4X1_UNSIGNED_BYTE";
|
|
PhysicalTextureType2[PhysicalTextureType2["PACKED_2X2_FLOAT32"] = 3] = "PACKED_2X2_FLOAT32";
|
|
PhysicalTextureType2[PhysicalTextureType2["PACKED_2X2_FLOAT16"] = 4] = "PACKED_2X2_FLOAT16";
|
|
})(PhysicalTextureType || (PhysicalTextureType = {}));
|
|
function getUnpackedMatrixTextureShapeWidthHeight(rows, columns) {
|
|
return [columns, rows];
|
|
}
|
|
function getUnpackedArraySizeFromMatrixSize(matrixSize, channelsPerTexture) {
|
|
return matrixSize * channelsPerTexture;
|
|
}
|
|
function getDenseTexShape(shape) {
|
|
const size = sizeFromShape(shape);
|
|
const texelsNeeded = Math.ceil(size / 4);
|
|
return sizeToSquarishShape(texelsNeeded);
|
|
}
|
|
function getPackedMatrixTextureShapeWidthHeight(rows, columns) {
|
|
return [
|
|
Math.max(1, Math.ceil(columns / 2)),
|
|
Math.max(1, Math.ceil(rows / 2))
|
|
];
|
|
}
|
|
function getPackedRGBAArraySizeFromMatrixShape(rows, columns) {
|
|
const [w, h] = getPackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
return w * h * 4;
|
|
}
|
|
function getTextureConfig(gl, textureHalfFloatExtension) {
|
|
const glany = gl;
|
|
let internalFormatFloat;
|
|
let internalFormatHalfFloat;
|
|
let internalFormatPackedHalfFloat;
|
|
let internalFormatPackedFloat;
|
|
let textureFormatFloat;
|
|
let downloadTextureFormat;
|
|
let downloadUnpackNumChannels;
|
|
let defaultNumChannels;
|
|
let textureTypeHalfFloat;
|
|
let textureTypeFloat;
|
|
if (env().getNumber("WEBGL_VERSION") === 2) {
|
|
internalFormatFloat = glany.R32F;
|
|
internalFormatHalfFloat = glany.R16F;
|
|
internalFormatPackedHalfFloat = glany.RGBA16F;
|
|
internalFormatPackedFloat = glany.RGBA32F;
|
|
textureFormatFloat = glany.RED;
|
|
downloadUnpackNumChannels = 4;
|
|
defaultNumChannels = 1;
|
|
textureTypeHalfFloat = glany.HALF_FLOAT;
|
|
textureTypeFloat = glany.FLOAT;
|
|
downloadTextureFormat = glany.RGBA8;
|
|
} else {
|
|
internalFormatFloat = gl.RGBA;
|
|
internalFormatHalfFloat = gl.RGBA;
|
|
internalFormatPackedHalfFloat = gl.RGBA;
|
|
internalFormatPackedFloat = glany.RGBA;
|
|
textureFormatFloat = gl.RGBA;
|
|
downloadUnpackNumChannels = 4;
|
|
defaultNumChannels = 4;
|
|
textureTypeHalfFloat = textureHalfFloatExtension != null ? textureHalfFloatExtension.HALF_FLOAT_OES : null;
|
|
textureTypeFloat = gl.FLOAT;
|
|
downloadTextureFormat = gl.RGBA;
|
|
}
|
|
return {
|
|
internalFormatFloat,
|
|
internalFormatHalfFloat,
|
|
internalFormatPackedHalfFloat,
|
|
internalFormatPackedFloat,
|
|
textureFormatFloat,
|
|
downloadTextureFormat,
|
|
downloadUnpackNumChannels,
|
|
defaultNumChannels,
|
|
textureTypeHalfFloat,
|
|
textureTypeFloat
|
|
};
|
|
}
|
|
function callAndCheck(gl, func) {
|
|
const returnValue = func();
|
|
if (env().getBool("DEBUG")) {
|
|
checkWebGLError(gl);
|
|
}
|
|
return returnValue;
|
|
}
|
|
function checkWebGLError(gl) {
|
|
const error = gl.getError();
|
|
if (error !== gl.NO_ERROR) {
|
|
throw new Error("WebGL Error: " + getWebGLErrorMessage(gl, error));
|
|
}
|
|
}
|
|
const MIN_FLOAT16 = 596e-10;
|
|
const MAX_FLOAT16 = 65504;
|
|
function canBeRepresented(num) {
|
|
if (env().getBool("WEBGL_RENDER_FLOAT32_ENABLED") || num === 0 || MIN_FLOAT16 < Math.abs(num) && Math.abs(num) < MAX_FLOAT16) {
|
|
return true;
|
|
}
|
|
return false;
|
|
}
|
|
function getWebGLErrorMessage(gl, status) {
|
|
switch (status) {
|
|
case gl.NO_ERROR:
|
|
return "NO_ERROR";
|
|
case gl.INVALID_ENUM:
|
|
return "INVALID_ENUM";
|
|
case gl.INVALID_VALUE:
|
|
return "INVALID_VALUE";
|
|
case gl.INVALID_OPERATION:
|
|
return "INVALID_OPERATION";
|
|
case gl.INVALID_FRAMEBUFFER_OPERATION:
|
|
return "INVALID_FRAMEBUFFER_OPERATION";
|
|
case gl.OUT_OF_MEMORY:
|
|
return "OUT_OF_MEMORY";
|
|
case gl.CONTEXT_LOST_WEBGL:
|
|
return "CONTEXT_LOST_WEBGL";
|
|
default:
|
|
return `Unknown error code ${status}`;
|
|
}
|
|
}
|
|
function getExtensionOrThrow(gl, extensionName) {
|
|
return throwIfNull(gl, () => gl.getExtension(extensionName), 'Extension "' + extensionName + '" not supported on this browser.');
|
|
}
|
|
function createVertexShader$1(gl, vertexShaderSource) {
|
|
const vertexShader = throwIfNull(gl, () => gl.createShader(gl.VERTEX_SHADER), "Unable to create vertex WebGLShader.");
|
|
callAndCheck(gl, () => gl.shaderSource(vertexShader, vertexShaderSource));
|
|
callAndCheck(gl, () => gl.compileShader(vertexShader));
|
|
if (gl.getShaderParameter(vertexShader, gl.COMPILE_STATUS) === false) {
|
|
console.log(gl.getShaderInfoLog(vertexShader));
|
|
throw new Error("Failed to compile vertex shader.");
|
|
}
|
|
return vertexShader;
|
|
}
|
|
function createFragmentShader(gl, fragmentShaderSource) {
|
|
const fragmentShader = throwIfNull(gl, () => gl.createShader(gl.FRAGMENT_SHADER), "Unable to create fragment WebGLShader.");
|
|
callAndCheck(gl, () => gl.shaderSource(fragmentShader, fragmentShaderSource));
|
|
callAndCheck(gl, () => gl.compileShader(fragmentShader));
|
|
if (env().get("ENGINE_COMPILE_ONLY")) {
|
|
return fragmentShader;
|
|
}
|
|
if (gl.getShaderParameter(fragmentShader, gl.COMPILE_STATUS) === false) {
|
|
logShaderSourceAndInfoLog(fragmentShaderSource, gl.getShaderInfoLog(fragmentShader));
|
|
throw new Error("Failed to compile fragment shader.");
|
|
}
|
|
return fragmentShader;
|
|
}
|
|
const lineNumberRegex = /ERROR: [0-9]+:([0-9]+):/g;
|
|
function logShaderSourceAndInfoLog(shaderSource, shaderInfoLog) {
|
|
const lineNumberRegexResult = lineNumberRegex.exec(shaderInfoLog);
|
|
if (lineNumberRegexResult == null) {
|
|
console.log(`Couldn't parse line number in error: ${shaderInfoLog}`);
|
|
console.log(shaderSource);
|
|
return;
|
|
}
|
|
const lineNumber = +lineNumberRegexResult[1];
|
|
const shaderLines = shaderSource.split("\n");
|
|
const pad2 = shaderLines.length.toString().length + 2;
|
|
const linesWithLineNumbers = shaderLines.map((line, lineNumber2) => rightPad((lineNumber2 + 1).toString(), pad2) + line);
|
|
let maxLineLength = 0;
|
|
for (let i = 0; i < linesWithLineNumbers.length; i++) {
|
|
maxLineLength = Math.max(linesWithLineNumbers[i].length, maxLineLength);
|
|
}
|
|
const beforeErrorLines = linesWithLineNumbers.slice(0, lineNumber - 1);
|
|
const errorLine = linesWithLineNumbers.slice(lineNumber - 1, lineNumber);
|
|
const afterErrorLines = linesWithLineNumbers.slice(lineNumber);
|
|
console.log(beforeErrorLines.join("\n"));
|
|
console.log(shaderInfoLog.split("\n")[0]);
|
|
console.log(`%c ${rightPad(errorLine[0], maxLineLength)}`, "border:1px solid red; background-color:#e3d2d2; color:#a61717");
|
|
console.log(afterErrorLines.join("\n"));
|
|
}
|
|
function createProgram(gl) {
|
|
return throwIfNull(gl, () => gl.createProgram(), "Unable to create WebGLProgram.");
|
|
}
|
|
function linkProgram(gl, program) {
|
|
callAndCheck(gl, () => gl.linkProgram(program));
|
|
if (env().get("ENGINE_COMPILE_ONLY")) {
|
|
return;
|
|
}
|
|
if (gl.getProgramParameter(program, gl.LINK_STATUS) === false) {
|
|
console.log(gl.getProgramInfoLog(program));
|
|
throw new Error("Failed to link vertex and fragment shaders.");
|
|
}
|
|
}
|
|
function validateProgram(gl, program) {
|
|
callAndCheck(gl, () => gl.validateProgram(program));
|
|
if (gl.getProgramParameter(program, gl.VALIDATE_STATUS) === false) {
|
|
console.log(gl.getProgramInfoLog(program));
|
|
throw new Error("Shader program validation failed.");
|
|
}
|
|
}
|
|
function createStaticVertexBuffer(gl, data) {
|
|
const buffer2 = throwIfNull(gl, () => gl.createBuffer(), "Unable to create WebGLBuffer");
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ARRAY_BUFFER, buffer2));
|
|
callAndCheck(gl, () => gl.bufferData(gl.ARRAY_BUFFER, data, gl.STATIC_DRAW));
|
|
return buffer2;
|
|
}
|
|
function createStaticIndexBuffer(gl, data) {
|
|
const buffer2 = throwIfNull(gl, () => gl.createBuffer(), "Unable to create WebGLBuffer");
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ELEMENT_ARRAY_BUFFER, buffer2));
|
|
callAndCheck(gl, () => gl.bufferData(gl.ELEMENT_ARRAY_BUFFER, data, gl.STATIC_DRAW));
|
|
return buffer2;
|
|
}
|
|
function createTexture(gl) {
|
|
return throwIfNull(gl, () => gl.createTexture(), "Unable to create WebGLTexture.");
|
|
}
|
|
function validateTextureSize(width, height) {
|
|
const maxTextureSize = env().getNumber("WEBGL_MAX_TEXTURE_SIZE");
|
|
if (width <= 0 || height <= 0) {
|
|
const requested = `[${width}x${height}]`;
|
|
throw new Error("Requested texture size " + requested + " is invalid.");
|
|
}
|
|
if (width > maxTextureSize || height > maxTextureSize) {
|
|
const requested = `[${width}x${height}]`;
|
|
const max2 = `[${maxTextureSize}x${maxTextureSize}]`;
|
|
throw new Error("Requested texture size " + requested + " greater than WebGL maximum on this browser / GPU " + max2 + ".");
|
|
}
|
|
}
|
|
function createFramebuffer(gl) {
|
|
return throwIfNull(gl, () => gl.createFramebuffer(), "Unable to create WebGLFramebuffer.");
|
|
}
|
|
function bindVertexBufferToProgramAttribute(gl, program, attribute, buffer2, arrayEntriesPerItem, itemStrideInBytes, itemOffsetInBytes) {
|
|
const loc = gl.getAttribLocation(program, attribute);
|
|
if (loc === -1) {
|
|
return false;
|
|
}
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ARRAY_BUFFER, buffer2));
|
|
callAndCheck(gl, () => gl.vertexAttribPointer(loc, arrayEntriesPerItem, gl.FLOAT, false, itemStrideInBytes, itemOffsetInBytes));
|
|
callAndCheck(gl, () => gl.enableVertexAttribArray(loc));
|
|
return true;
|
|
}
|
|
function bindTextureUnit(gl, texture, textureUnit) {
|
|
validateTextureUnit(gl, textureUnit);
|
|
callAndCheck(gl, () => gl.activeTexture(gl.TEXTURE0 + textureUnit));
|
|
callAndCheck(gl, () => gl.bindTexture(gl.TEXTURE_2D, texture));
|
|
}
|
|
function getProgramUniformLocationOrThrow(gl, program, uniformName) {
|
|
return throwIfNull(gl, () => gl.getUniformLocation(program, uniformName), 'uniform "' + uniformName + '" not present in program.');
|
|
}
|
|
function getProgramUniformLocation(gl, program, uniformName) {
|
|
return gl.getUniformLocation(program, uniformName);
|
|
}
|
|
function bindTextureToProgramUniformSampler(gl, texture, uniformSamplerLocation, textureUnit) {
|
|
callAndCheck(gl, () => bindTextureUnit(gl, texture, textureUnit));
|
|
callAndCheck(gl, () => gl.uniform1i(uniformSamplerLocation, textureUnit));
|
|
}
|
|
function bindColorTextureToFramebuffer(gl, texture, framebuffer) {
|
|
callAndCheck(gl, () => gl.bindFramebuffer(gl.FRAMEBUFFER, framebuffer));
|
|
callAndCheck(gl, () => gl.framebufferTexture2D(gl.FRAMEBUFFER, gl.COLOR_ATTACHMENT0, gl.TEXTURE_2D, texture, 0));
|
|
}
|
|
function unbindColorTextureFromFramebuffer(gl, framebuffer) {
|
|
callAndCheck(gl, () => gl.bindFramebuffer(gl.FRAMEBUFFER, framebuffer));
|
|
callAndCheck(gl, () => gl.framebufferTexture2D(gl.FRAMEBUFFER, gl.COLOR_ATTACHMENT0, gl.TEXTURE_2D, null, 0));
|
|
}
|
|
function validateFramebuffer(gl) {
|
|
const status = gl.checkFramebufferStatus(gl.FRAMEBUFFER);
|
|
if (status !== gl.FRAMEBUFFER_COMPLETE) {
|
|
throw new Error("Error binding framebuffer: " + getFramebufferErrorMessage(gl, status));
|
|
}
|
|
}
|
|
function getFramebufferErrorMessage(gl, status) {
|
|
switch (status) {
|
|
case gl.FRAMEBUFFER_INCOMPLETE_ATTACHMENT:
|
|
return "FRAMEBUFFER_INCOMPLETE_ATTACHMENT";
|
|
case gl.FRAMEBUFFER_INCOMPLETE_MISSING_ATTACHMENT:
|
|
return "FRAMEBUFFER_INCOMPLETE_MISSING_ATTACHMENT";
|
|
case gl.FRAMEBUFFER_INCOMPLETE_DIMENSIONS:
|
|
return "FRAMEBUFFER_INCOMPLETE_DIMENSIONS";
|
|
case gl.FRAMEBUFFER_UNSUPPORTED:
|
|
return "FRAMEBUFFER_UNSUPPORTED";
|
|
default:
|
|
return `unknown error ${status}`;
|
|
}
|
|
}
|
|
function throwIfNull(gl, returnTOrNull, failureMessage) {
|
|
const tOrNull = callAndCheck(gl, () => returnTOrNull());
|
|
if (tOrNull == null) {
|
|
throw new Error(failureMessage);
|
|
}
|
|
return tOrNull;
|
|
}
|
|
function validateTextureUnit(gl, textureUnit) {
|
|
const maxTextureUnit = gl.MAX_COMBINED_TEXTURE_IMAGE_UNITS - 1;
|
|
const glTextureUnit = textureUnit + gl.TEXTURE0;
|
|
if (glTextureUnit < gl.TEXTURE0 || glTextureUnit > maxTextureUnit) {
|
|
const textureUnitRange = `[gl.TEXTURE0, gl.TEXTURE${maxTextureUnit}]`;
|
|
throw new Error(`textureUnit must be in ${textureUnitRange}.`);
|
|
}
|
|
}
|
|
function getBatchDim(shape, dimsToSkip = 2) {
|
|
return sizeFromShape(shape.slice(0, shape.length - dimsToSkip));
|
|
}
|
|
function getRowsCols(shape) {
|
|
if (shape.length === 0) {
|
|
throw Error("Cannot get rows and columns of an empty shape array.");
|
|
}
|
|
return [
|
|
shape.length > 1 ? shape[shape.length - 2] : 1,
|
|
shape[shape.length - 1]
|
|
];
|
|
}
|
|
function getShapeAs3D(shape) {
|
|
let shapeAs3D = [1, 1, 1];
|
|
const isScalar = shape.length === 0 || shape.length === 1 && shape[0] === 1;
|
|
if (!isScalar) {
|
|
shapeAs3D = [getBatchDim(shape), ...getRowsCols(shape)];
|
|
}
|
|
return shapeAs3D;
|
|
}
|
|
function getTextureShapeFromLogicalShape(logShape, isPacked = false) {
|
|
let maxTexSize = env().getNumber("WEBGL_MAX_TEXTURE_SIZE");
|
|
let maxSizeForNarrowTex = env().getNumber("WEBGL_MAX_SIZE_FOR_NARROW_TEXTURE");
|
|
if (maxSizeForNarrowTex === Infinity && env().getBool("WEBGL_AUTO_SQUARIFY_NARROW_TEXTURE_SHAPE")) {
|
|
maxSizeForNarrowTex = maxTexSize / 2;
|
|
}
|
|
if (isPacked) {
|
|
maxTexSize = maxTexSize * 2;
|
|
maxSizeForNarrowTex = maxSizeForNarrowTex * 2;
|
|
logShape = logShape.map((d, i) => i >= logShape.length - 2 ? nearestLargerEven(logShape[i]) : logShape[i]);
|
|
if (logShape.length === 1) {
|
|
logShape = [2, logShape[0]];
|
|
}
|
|
}
|
|
if (logShape.length !== 2) {
|
|
const squeezeResult = squeezeShape(logShape);
|
|
logShape = squeezeResult.newShape;
|
|
}
|
|
let size = sizeFromShape(logShape);
|
|
let textureShape = null;
|
|
if (logShape.length <= 1 && size <= maxTexSize) {
|
|
textureShape = [1, size];
|
|
} else if (logShape.length === 2 && logShape[0] <= maxTexSize && logShape[1] <= maxTexSize) {
|
|
textureShape = logShape;
|
|
} else if (logShape.length === 3 && logShape[0] * logShape[1] <= maxTexSize && logShape[2] <= maxTexSize) {
|
|
textureShape = [logShape[0] * logShape[1], logShape[2]];
|
|
} else if (logShape.length === 3 && logShape[0] <= maxTexSize && logShape[1] * logShape[2] <= maxTexSize) {
|
|
textureShape = [logShape[0], logShape[1] * logShape[2]];
|
|
} else if (logShape.length === 4 && logShape[0] * logShape[1] * logShape[2] <= maxTexSize && logShape[3] <= maxTexSize) {
|
|
textureShape = [logShape[0] * logShape[1] * logShape[2], logShape[3]];
|
|
} else if (logShape.length === 4 && logShape[0] <= maxTexSize && logShape[1] * logShape[2] * logShape[3] <= maxTexSize) {
|
|
textureShape = [logShape[0], logShape[1] * logShape[2] * logShape[3]];
|
|
}
|
|
const isLongNarrowTex = textureShape != null && Math.max(...textureShape) > maxSizeForNarrowTex && Math.min(...textureShape) <= (isPacked ? 2 : 1) && Math.min(...textureShape) > 0;
|
|
if (textureShape == null || isLongNarrowTex) {
|
|
if (isPacked) {
|
|
const batchDim = getBatchDim(logShape);
|
|
let rows = 2, cols = 2;
|
|
if (logShape.length) {
|
|
[rows, cols] = getRowsCols(logShape);
|
|
}
|
|
size = batchDim * (rows / 2) * (cols / 2);
|
|
textureShape = sizeToSquarishShape(size).map((d) => d * 2);
|
|
} else {
|
|
textureShape = sizeToSquarishShape(size);
|
|
}
|
|
}
|
|
return textureShape;
|
|
}
|
|
function isEven(n) {
|
|
return n % 2 === 0;
|
|
}
|
|
function isReshapeFree(shape1, shape2) {
|
|
shape1 = shape1.slice(-2);
|
|
shape2 = shape2.slice(-2);
|
|
if (arraysEqual(shape1, shape2)) {
|
|
return true;
|
|
}
|
|
if (!shape1.length || !shape2.length) {
|
|
return true;
|
|
}
|
|
if (shape1[0] === 0 || shape1[1] === 0 || shape2[0] === 0 || shape2[1] === 0) {
|
|
return true;
|
|
}
|
|
if (shape1.length !== shape2.length) {
|
|
const shape1Cols = shape1[shape1.length - 1];
|
|
const shape2Cols = shape2[shape2.length - 1];
|
|
if (shape1Cols === shape2Cols) {
|
|
return true;
|
|
}
|
|
if (isEven(shape1Cols) && isEven(shape2Cols) && (shape1[0] === 1 || shape2[0] === 1)) {
|
|
return true;
|
|
}
|
|
}
|
|
return shape1[1] === shape2[1] && isEven(shape1[0]) && isEven(shape2[0]);
|
|
}
|
|
let MAX_TEXTURE_SIZE;
|
|
let MAX_TEXTURES_IN_SHADER;
|
|
function getWebGLMaxTextureSize(webGLVersion) {
|
|
if (MAX_TEXTURE_SIZE == null) {
|
|
const gl = getWebGLContext(webGLVersion);
|
|
MAX_TEXTURE_SIZE = gl.getParameter(gl.MAX_TEXTURE_SIZE);
|
|
}
|
|
return MAX_TEXTURE_SIZE;
|
|
}
|
|
function getMaxTexturesInShader(webGLVersion) {
|
|
if (MAX_TEXTURES_IN_SHADER == null) {
|
|
const gl = getWebGLContext(webGLVersion);
|
|
MAX_TEXTURES_IN_SHADER = gl.getParameter(gl.MAX_TEXTURE_IMAGE_UNITS);
|
|
}
|
|
return Math.min(16, MAX_TEXTURES_IN_SHADER);
|
|
}
|
|
function getWebGLDisjointQueryTimerVersion(webGLVersion) {
|
|
if (webGLVersion === 0) {
|
|
return 0;
|
|
}
|
|
let queryTimerVersion;
|
|
const gl = getWebGLContext(webGLVersion);
|
|
if (hasExtension(gl, "EXT_disjoint_timer_query_webgl2") && webGLVersion === 2) {
|
|
queryTimerVersion = 2;
|
|
} else if (hasExtension(gl, "EXT_disjoint_timer_query")) {
|
|
queryTimerVersion = 1;
|
|
} else {
|
|
queryTimerVersion = 0;
|
|
}
|
|
return queryTimerVersion;
|
|
}
|
|
function hasExtension(gl, extensionName) {
|
|
const ext = gl.getExtension(extensionName);
|
|
return ext != null;
|
|
}
|
|
function isWebGLVersionEnabled(webGLVersion) {
|
|
try {
|
|
const gl = getWebGLContext(webGLVersion);
|
|
if (gl != null) {
|
|
return true;
|
|
}
|
|
} catch (e) {
|
|
console.log("Error when getting WebGL context: ", e);
|
|
return false;
|
|
}
|
|
return false;
|
|
}
|
|
function isCapableOfRenderingToFloatTexture(webGLVersion) {
|
|
if (webGLVersion === 0) {
|
|
return false;
|
|
}
|
|
const gl = getWebGLContext(webGLVersion);
|
|
if (webGLVersion === 1) {
|
|
if (!hasExtension(gl, "OES_texture_float")) {
|
|
return false;
|
|
}
|
|
} else {
|
|
if (!hasExtension(gl, "EXT_color_buffer_float")) {
|
|
return false;
|
|
}
|
|
}
|
|
const isFrameBufferComplete = createFloatTextureAndBindToFramebuffer(gl);
|
|
return isFrameBufferComplete;
|
|
}
|
|
function isDownloadFloatTextureEnabled(webGLVersion) {
|
|
if (webGLVersion === 0) {
|
|
return false;
|
|
}
|
|
const gl = getWebGLContext(webGLVersion);
|
|
if (webGLVersion === 1) {
|
|
if (!hasExtension(gl, "OES_texture_float")) {
|
|
return false;
|
|
}
|
|
if (!hasExtension(gl, "WEBGL_color_buffer_float")) {
|
|
return false;
|
|
}
|
|
} else {
|
|
if (hasExtension(gl, "EXT_color_buffer_float")) {
|
|
return createFloatTextureAndBindToFramebuffer(gl);
|
|
}
|
|
const COLOR_BUFFER_HALF_FLOAT = "EXT_color_buffer_half_float";
|
|
if (hasExtension(gl, COLOR_BUFFER_HALF_FLOAT)) {
|
|
const textureHalfFloatExtension = gl.getExtension(COLOR_BUFFER_HALF_FLOAT);
|
|
return createHalfFloatTextureAndBindToFramebuffer(gl, textureHalfFloatExtension);
|
|
}
|
|
return false;
|
|
}
|
|
const isFrameBufferComplete = createFloatTextureAndBindToFramebuffer(gl);
|
|
return isFrameBufferComplete;
|
|
}
|
|
function createFloatTextureAndBindToFramebuffer(gl) {
|
|
const texConfig = getTextureConfig(gl);
|
|
const texture = gl.createTexture();
|
|
gl.bindTexture(gl.TEXTURE_2D, texture);
|
|
const width = 1;
|
|
const height = 1;
|
|
gl.texImage2D(gl.TEXTURE_2D, 0, texConfig.internalFormatFloat, width, height, 0, texConfig.textureFormatFloat, texConfig.textureTypeFloat, null);
|
|
const frameBuffer = gl.createFramebuffer();
|
|
gl.bindFramebuffer(gl.FRAMEBUFFER, frameBuffer);
|
|
gl.framebufferTexture2D(gl.FRAMEBUFFER, gl.COLOR_ATTACHMENT0, gl.TEXTURE_2D, texture, 0);
|
|
const isFrameBufferComplete = gl.checkFramebufferStatus(gl.FRAMEBUFFER) === gl.FRAMEBUFFER_COMPLETE;
|
|
gl.bindTexture(gl.TEXTURE_2D, null);
|
|
gl.bindFramebuffer(gl.FRAMEBUFFER, null);
|
|
gl.deleteTexture(texture);
|
|
gl.deleteFramebuffer(frameBuffer);
|
|
return isFrameBufferComplete;
|
|
}
|
|
function createHalfFloatTextureAndBindToFramebuffer(gl, textureHalfFloatExtension) {
|
|
const texConfig = getTextureConfig(gl, textureHalfFloatExtension);
|
|
const texture = gl.createTexture();
|
|
gl.bindTexture(gl.TEXTURE_2D, texture);
|
|
const width = 1;
|
|
const height = 1;
|
|
gl.texImage2D(gl.TEXTURE_2D, 0, texConfig.internalFormatHalfFloat, width, height, 0, texConfig.textureFormatFloat, texConfig.textureTypeHalfFloat, null);
|
|
const frameBuffer = gl.createFramebuffer();
|
|
gl.bindFramebuffer(gl.FRAMEBUFFER, frameBuffer);
|
|
gl.framebufferTexture2D(gl.FRAMEBUFFER, gl.COLOR_ATTACHMENT0, gl.TEXTURE_2D, texture, 0);
|
|
const isFrameBufferComplete = gl.checkFramebufferStatus(gl.FRAMEBUFFER) === gl.FRAMEBUFFER_COMPLETE;
|
|
gl.bindTexture(gl.TEXTURE_2D, null);
|
|
gl.bindFramebuffer(gl.FRAMEBUFFER, null);
|
|
gl.deleteTexture(texture);
|
|
gl.deleteFramebuffer(frameBuffer);
|
|
return isFrameBufferComplete;
|
|
}
|
|
function isWebGLFenceEnabled(webGLVersion) {
|
|
if (webGLVersion !== 2) {
|
|
return false;
|
|
}
|
|
const gl = getWebGLContext(webGLVersion);
|
|
const isEnabled = gl.fenceSync != null;
|
|
return isEnabled;
|
|
}
|
|
function assertNotComplex$1(tensor2, opName) {
|
|
if (!Array.isArray(tensor2)) {
|
|
tensor2 = [tensor2];
|
|
}
|
|
tensor2.forEach((t2) => {
|
|
if (t2 != null) {
|
|
assert(t2.dtype !== "complex64", () => `${opName} does not support complex64 tensors in the WebGL backend.`);
|
|
}
|
|
});
|
|
}
|
|
const ENV = env();
|
|
ENV.registerFlag("HAS_WEBGL", () => ENV.getNumber("WEBGL_VERSION") > 0);
|
|
ENV.registerFlag("WEBGL_VERSION", () => {
|
|
if (isWebGLVersionEnabled(2)) {
|
|
return 2;
|
|
} else if (isWebGLVersionEnabled(1)) {
|
|
return 1;
|
|
}
|
|
return 0;
|
|
});
|
|
ENV.registerFlag("WEBGL_CHECK_NUMERICAL_PROBLEMS", () => false);
|
|
ENV.registerFlag("WEBGL_BUFFER_SUPPORTED", () => ENV.get("WEBGL_VERSION") === 2);
|
|
ENV.registerFlag("WEBGL_CPU_FORWARD", () => true);
|
|
ENV.registerFlag("WEBGL_FORCE_F16_TEXTURES", () => false);
|
|
ENV.registerFlag("WEBGL_PACK", () => ENV.getBool("HAS_WEBGL"));
|
|
ENV.registerFlag("WEBGL_PACK_NORMALIZATION", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_CLIP", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_DEPTHWISECONV", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_BINARY_OPERATIONS", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_UNARY_OPERATIONS", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_ARRAY_OPERATIONS", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_IMAGE_OPERATIONS", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_REDUCE", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_LAZILY_UNPACK", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_CONV_IM2COL", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_PACK_CONV2DTRANSPOSE", () => ENV.getBool("WEBGL_PACK"));
|
|
ENV.registerFlag("WEBGL_MAX_TEXTURE_SIZE", () => getWebGLMaxTextureSize(ENV.getNumber("WEBGL_VERSION")));
|
|
ENV.registerFlag("WEBGL_MAX_TEXTURES_IN_SHADER", () => getMaxTexturesInShader(ENV.getNumber("WEBGL_VERSION")));
|
|
ENV.registerFlag("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION", () => {
|
|
const webGLVersion = ENV.getNumber("WEBGL_VERSION");
|
|
if (webGLVersion === 0) {
|
|
return 0;
|
|
}
|
|
return getWebGLDisjointQueryTimerVersion(webGLVersion);
|
|
});
|
|
ENV.registerFlag("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_RELIABLE", () => ENV.getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION") > 0 && !isMobile());
|
|
ENV.registerFlag("WEBGL_RENDER_FLOAT32_CAPABLE", () => isCapableOfRenderingToFloatTexture(ENV.getNumber("WEBGL_VERSION")));
|
|
ENV.registerFlag("WEBGL_RENDER_FLOAT32_ENABLED", () => {
|
|
return ENV.getBool("WEBGL_FORCE_F16_TEXTURES") ? false : ENV.getBool("WEBGL_RENDER_FLOAT32_CAPABLE");
|
|
});
|
|
ENV.registerFlag("WEBGL_DOWNLOAD_FLOAT_ENABLED", () => isDownloadFloatTextureEnabled(ENV.getNumber("WEBGL_VERSION")));
|
|
ENV.registerFlag("WEBGL_FENCE_API_ENABLED", () => isWebGLFenceEnabled(ENV.getNumber("WEBGL_VERSION")));
|
|
ENV.registerFlag("WEBGL_SIZE_UPLOAD_UNIFORM", () => {
|
|
const useUniforms = ENV.getBool("WEBGL_RENDER_FLOAT32_ENABLED");
|
|
return useUniforms ? 4 : 0;
|
|
});
|
|
ENV.registerFlag("WEBGL_DELETE_TEXTURE_THRESHOLD", () => {
|
|
return -1;
|
|
}, (threshold2) => {
|
|
if (!(typeof threshold2 === "number")) {
|
|
throw new Error(`WEBGL_DELETE_TEXTURE_THRESHOLD must be a number but got ${threshold2}.`);
|
|
}
|
|
if (threshold2 < 0 && threshold2 !== -1) {
|
|
throw new Error(`WEBGL_DELETE_TEXTURE_THRESHOLD must be -1 (indicating never delete) or at least 0, but got ${threshold2}.`);
|
|
}
|
|
});
|
|
ENV.registerFlag("WEBGL_FLUSH_THRESHOLD", () => {
|
|
return isMobile() ? 1 : -1;
|
|
}, (threshold2) => {
|
|
if (!(typeof threshold2 === "number")) {
|
|
throw new Error(`WEBGL_FLUSH_THRESHOLD must be a number but got ${threshold2}.`);
|
|
}
|
|
if (threshold2 < 0 && threshold2 !== -1) {
|
|
throw new Error(`WEBGL_FLUSH_THRESHOLD must be -1 (indicating never manual flush) or at least 0, but got ${threshold2}.`);
|
|
}
|
|
});
|
|
ENV.registerFlag("CPU_HANDOFF_SIZE_THRESHOLD", () => 128);
|
|
ENV.registerFlag("WEBGL_USE_SHAPES_UNIFORMS", () => false);
|
|
ENV.registerFlag("TOPK_LAST_DIM_CPU_HANDOFF_SIZE_THRESHOLD", () => 1e5);
|
|
ENV.registerFlag("TOPK_K_CPU_HANDOFF_THRESHOLD", () => 128);
|
|
ENV.registerFlag("WEBGL_EXP_CONV", () => false);
|
|
ENV.registerFlag("SOFTWARE_WEBGL_ENABLED", () => ENV.getBool("IS_TEST"));
|
|
ENV.registerFlag("WEBGL_MAX_SIZE_FOR_NARROW_TEXTURE", () => Infinity);
|
|
ENV.registerFlag("WEBGL_AUTO_SQUARIFY_NARROW_TEXTURE_SHAPE", () => false);
|
|
ENV.registerFlag("WEBGL2_ISNAN_CUSTOM", () => false);
|
|
ENV.registerFlag("ENGINE_COMPILE_ONLY", () => false);
|
|
function getGlslDifferences() {
|
|
let version;
|
|
let attribute;
|
|
let varyingVs;
|
|
let varyingFs;
|
|
let texture2D;
|
|
let output;
|
|
let defineOutput;
|
|
let defineSpecialNaN;
|
|
let defineSpecialInf;
|
|
let defineRound;
|
|
if (env().getNumber("WEBGL_VERSION") === 2) {
|
|
version = "#version 300 es";
|
|
attribute = "in";
|
|
varyingVs = "out";
|
|
varyingFs = "in";
|
|
texture2D = "texture";
|
|
output = "outputColor";
|
|
defineOutput = "out vec4 outputColor;";
|
|
defineSpecialNaN = env().getBool("WEBGL2_ISNAN_CUSTOM") ? `
|
|
bool isnan_custom(float val) {
|
|
uint floatToUint = floatBitsToUint(val);
|
|
return (floatToUint & 0x7fffffffu) > 0x7f800000u;
|
|
}
|
|
|
|
bvec4 isnan_custom(vec4 val) {
|
|
return bvec4(isnan_custom(val.x),
|
|
isnan_custom(val.y), isnan_custom(val.z), isnan_custom(val.w));
|
|
}
|
|
|
|
#define isnan(value) isnan_custom(value)
|
|
` : "";
|
|
defineSpecialInf = ``;
|
|
defineRound = `
|
|
#define round(value) newRound(value)
|
|
int newRound(float value) {
|
|
return int(floor(value + 0.5));
|
|
}
|
|
|
|
ivec4 newRound(vec4 value) {
|
|
return ivec4(floor(value + vec4(0.5)));
|
|
}
|
|
`;
|
|
} else {
|
|
version = "";
|
|
attribute = "attribute";
|
|
varyingVs = "varying";
|
|
varyingFs = "varying";
|
|
texture2D = "texture2D";
|
|
output = "gl_FragColor";
|
|
defineOutput = "";
|
|
defineSpecialNaN = `
|
|
#define isnan(value) isnan_custom(value)
|
|
bool isnan_custom(float val) {
|
|
return (val > 0. || val < 1. || val == 0.) ? false : true;
|
|
}
|
|
bvec4 isnan_custom(vec4 val) {
|
|
return bvec4(isnan(val.x), isnan(val.y), isnan(val.z), isnan(val.w));
|
|
}
|
|
`;
|
|
defineSpecialInf = `
|
|
uniform float INFINITY;
|
|
|
|
bool isinf(float val) {
|
|
return abs(val) == INFINITY;
|
|
}
|
|
bvec4 isinf(vec4 val) {
|
|
return equal(abs(val), vec4(INFINITY));
|
|
}
|
|
`;
|
|
defineRound = `
|
|
int round(float value) {
|
|
return int(floor(value + 0.5));
|
|
}
|
|
|
|
ivec4 round(vec4 value) {
|
|
return ivec4(floor(value + vec4(0.5)));
|
|
}
|
|
`;
|
|
}
|
|
return {
|
|
version,
|
|
attribute,
|
|
varyingVs,
|
|
varyingFs,
|
|
texture2D,
|
|
output,
|
|
defineOutput,
|
|
defineSpecialNaN,
|
|
defineSpecialInf,
|
|
defineRound
|
|
};
|
|
}
|
|
function getLogicalCoordinatesFromFlatIndex(coords2, shape, index = "index") {
|
|
const strides = computeStrides(shape);
|
|
return strides.map((stride, i) => {
|
|
const line1 = `int ${coords2[i]} = ${index} / ${stride}`;
|
|
const line2 = i === strides.length - 1 ? `int ${coords2[i + 1]} = ${index} - ${coords2[i]} * ${stride}` : `index -= ${coords2[i]} * ${stride}`;
|
|
return `${line1}; ${line2};`;
|
|
}).join("");
|
|
}
|
|
function getOutputLogicalCoordinatesFromFlatIndexByUniform(coords2, shape, index = "index") {
|
|
const strides = computeStrides(shape);
|
|
return strides.map((_, i) => {
|
|
const line1 = `int ${coords2[i]} = ${index} / outShapeStrides[${i}]`;
|
|
const line2 = i === strides.length - 1 ? `int ${coords2[i + 1]} = ${index} - ${coords2[i]} * outShapeStrides[${i}]` : `index -= ${coords2[i]} * outShapeStrides[${i}]`;
|
|
return `${line1}; ${line2};`;
|
|
}).join("");
|
|
}
|
|
function symbolicallyComputeStrides(indicesArr, variableName) {
|
|
const numCoords = indicesArr.length;
|
|
const shape = indicesArr.map((d) => `${variableName}[${d}]`);
|
|
const strides = new Array(numCoords - 1);
|
|
strides[numCoords - 2] = shape[numCoords - 1];
|
|
for (let i = numCoords - 3; i >= 0; --i) {
|
|
strides[i] = `(${strides[i + 1]} * ${shape[i + 1]})`;
|
|
}
|
|
return strides;
|
|
}
|
|
function getLogicalCoordinatesFromFlatIndexByUniform(coords2, variableName, index = "index") {
|
|
const indicesArray = coords2.map((_, i) => i);
|
|
const strides = symbolicallyComputeStrides(indicesArray, variableName);
|
|
return strides.map((_, i) => {
|
|
const line1 = `int ${coords2[i]} = ${index} / ${strides[i]}`;
|
|
const line2 = i === strides.length - 1 ? `int ${coords2[i + 1]} = ${index} - ${coords2[i]} * ${strides[i]}` : `index -= ${coords2[i]} * ${strides[i]}`;
|
|
return `${line1}; ${line2};`;
|
|
}).join("");
|
|
}
|
|
function getFlatIndexFrom3D(shape) {
|
|
const strides = computeStrides(shape).map((d) => d.toString());
|
|
return `
|
|
int getFlatIndex(ivec3 coords) {
|
|
return coords.x * ${strides[0]} + coords.y * ${strides[1]} + coords.z;
|
|
}
|
|
`;
|
|
}
|
|
function getFlatIndexFrom3DOutput() {
|
|
return `
|
|
int getFlatIndex(ivec3 coords) {
|
|
return coords.x * outShapeStrides[0] + coords.y * outShapeStrides[1] + coords.z;
|
|
}
|
|
`;
|
|
}
|
|
const ENCODE_FLOAT_SNIPPET = `
|
|
const float FLOAT_MAX = 1.70141184e38;
|
|
const float FLOAT_MIN = 1.17549435e-38;
|
|
|
|
lowp vec4 encode_float(highp float v) {
|
|
if (isnan(v)) {
|
|
return vec4(255, 255, 255, 255);
|
|
}
|
|
|
|
highp float av = abs(v);
|
|
|
|
if(av < FLOAT_MIN) {
|
|
return vec4(0.0, 0.0, 0.0, 0.0);
|
|
} else if(v > FLOAT_MAX) {
|
|
return vec4(0.0, 0.0, 128.0, 127.0) / 255.0;
|
|
} else if(v < -FLOAT_MAX) {
|
|
return vec4(0.0, 0.0, 128.0, 255.0) / 255.0;
|
|
}
|
|
|
|
highp vec4 c = vec4(0,0,0,0);
|
|
|
|
highp float e = floor(log2(av));
|
|
highp float m = exp2(fract(log2(av))) - 1.0;
|
|
|
|
c[2] = floor(128.0 * m);
|
|
m -= c[2] / 128.0;
|
|
c[1] = floor(32768.0 * m);
|
|
m -= c[1] / 32768.0;
|
|
c[0] = floor(8388608.0 * m);
|
|
|
|
highp float ebias = e + 127.0;
|
|
c[3] = floor(ebias / 2.0);
|
|
ebias -= c[3] * 2.0;
|
|
c[2] += floor(ebias) * 128.0;
|
|
|
|
c[3] += 128.0 * step(0.0, -v);
|
|
|
|
return c / 255.0;
|
|
}
|
|
`;
|
|
const { getBroadcastDims } = backend_util;
|
|
function makeShader(inputsInfo, outputShape, program) {
|
|
const prefixSnippets = [];
|
|
inputsInfo.forEach((x) => {
|
|
const size = sizeFromShape(x.shapeInfo.logicalShape);
|
|
if (x.shapeInfo.isUniform) {
|
|
prefixSnippets.push(`uniform float ${x.name}${size > 1 ? `[${size}]` : ""};`);
|
|
} else {
|
|
prefixSnippets.push(`uniform sampler2D ${x.name};`);
|
|
prefixSnippets.push(`uniform int offset${x.name};`);
|
|
}
|
|
if (program.enableShapeUniforms) {
|
|
const { uniformShape } = getUniformInfoFromShape(program.packedInputs, x.shapeInfo.logicalShape, x.shapeInfo.texShape);
|
|
switch (uniformShape.length) {
|
|
case 1:
|
|
prefixSnippets.push(`uniform int ${x.name}Shape;`);
|
|
break;
|
|
case 2:
|
|
prefixSnippets.push(`uniform ivec2 ${x.name}Shape;`);
|
|
break;
|
|
case 3:
|
|
prefixSnippets.push(`uniform ivec3 ${x.name}Shape;`);
|
|
break;
|
|
case 4:
|
|
prefixSnippets.push(`uniform ivec4 ${x.name}Shape;`);
|
|
break;
|
|
}
|
|
prefixSnippets.push(`uniform ivec2 ${x.name}TexShape;`);
|
|
}
|
|
});
|
|
if (program.enableShapeUniforms) {
|
|
switch (outputShape.logicalShape.length) {
|
|
case 1:
|
|
prefixSnippets.push(`uniform int outShape;`);
|
|
break;
|
|
case 2:
|
|
prefixSnippets.push(`uniform ivec2 outShape;`);
|
|
prefixSnippets.push(`uniform int outShapeStrides;`);
|
|
break;
|
|
case 3:
|
|
prefixSnippets.push(`uniform ivec3 outShape;`);
|
|
prefixSnippets.push(`uniform ivec2 outShapeStrides;`);
|
|
break;
|
|
case 4:
|
|
prefixSnippets.push(`uniform ivec4 outShape;`);
|
|
prefixSnippets.push(`uniform ivec3 outShapeStrides;`);
|
|
break;
|
|
}
|
|
prefixSnippets.push(`uniform ivec2 outTexShape;`);
|
|
}
|
|
if (program.customUniforms) {
|
|
program.customUniforms.forEach((d) => {
|
|
prefixSnippets.push(`uniform ${d.type} ${d.name}${d.arrayIndex ? `[${d.arrayIndex}]` : ""};`);
|
|
});
|
|
}
|
|
const inputPrefixSnippet = prefixSnippets.join("\n");
|
|
const inputSamplingSnippet = inputsInfo.map((x) => getInputSamplingSnippet(x, outputShape, program.packedInputs, program.enableShapeUniforms)).join("\n");
|
|
const outTexShape = outputShape.texShape;
|
|
const glsl = getGlslDifferences();
|
|
const floatTextureSampleSnippet = getFloatTextureSampleSnippet(glsl);
|
|
let outputSamplingSnippet;
|
|
let floatTextureSetOutputSnippet;
|
|
let shaderPrefix = getShaderPrefix(glsl);
|
|
if (outputShape.isPacked) {
|
|
outputSamplingSnippet = getPackedOutputSamplingSnippet(outputShape.logicalShape, outTexShape, program.enableShapeUniforms);
|
|
floatTextureSetOutputSnippet = getFloatTextureSetRGBASnippet(glsl);
|
|
} else {
|
|
outputSamplingSnippet = getOutputSamplingSnippet(outputShape.logicalShape, outTexShape, program.enableShapeUniforms);
|
|
floatTextureSetOutputSnippet = getFloatTextureSetRSnippet(glsl);
|
|
}
|
|
if (program.packedInputs) {
|
|
shaderPrefix += SHADER_PACKED_PREFIX;
|
|
}
|
|
const source = [
|
|
shaderPrefix,
|
|
floatTextureSampleSnippet,
|
|
floatTextureSetOutputSnippet,
|
|
inputPrefixSnippet,
|
|
outputSamplingSnippet,
|
|
inputSamplingSnippet,
|
|
program.userCode
|
|
].join("\n");
|
|
return source;
|
|
}
|
|
function getSamplerFromInInfo(inInfo, enableShapeUniforms = false) {
|
|
const shape = inInfo.shapeInfo.logicalShape;
|
|
switch (shape.length) {
|
|
case 0:
|
|
return getSamplerScalar(inInfo, enableShapeUniforms);
|
|
case 1:
|
|
return getSampler1D(inInfo, enableShapeUniforms);
|
|
case 2:
|
|
return getSampler2D(inInfo, enableShapeUniforms);
|
|
case 3:
|
|
return getSampler3D(inInfo, enableShapeUniforms);
|
|
case 4:
|
|
return getSampler4D(inInfo, enableShapeUniforms);
|
|
case 5:
|
|
return getSampler5D(inInfo);
|
|
case 6:
|
|
return getSampler6D(inInfo);
|
|
default:
|
|
throw new Error(`${shape.length}-D input sampling is not yet supported`);
|
|
}
|
|
}
|
|
function getPackedSamplerFromInInfo(inInfo, enableShapeUniforms) {
|
|
const shape = inInfo.shapeInfo.logicalShape;
|
|
switch (shape.length) {
|
|
case 0:
|
|
return getPackedSamplerScalar(inInfo);
|
|
case 1:
|
|
return getPackedSampler1D(inInfo, enableShapeUniforms);
|
|
case 2:
|
|
return getPackedSampler2D(inInfo, enableShapeUniforms);
|
|
case 3:
|
|
return getPackedSampler3D(inInfo, enableShapeUniforms);
|
|
default:
|
|
return getPackedSamplerND(inInfo, enableShapeUniforms);
|
|
}
|
|
}
|
|
function getInputSamplingSnippet(inInfo, outShapeInfo, usesPackedTextures = false, enableShapeUniforms) {
|
|
let res = "";
|
|
if (usesPackedTextures) {
|
|
res += getPackedSamplerFromInInfo(inInfo, enableShapeUniforms);
|
|
} else {
|
|
res += getSamplerFromInInfo(inInfo, enableShapeUniforms);
|
|
}
|
|
const inShape = inInfo.shapeInfo.logicalShape;
|
|
const outShape = outShapeInfo.logicalShape;
|
|
if (inShape.length <= outShape.length) {
|
|
if (usesPackedTextures) {
|
|
res += getPackedSamplerAtOutputCoords(inInfo, outShapeInfo);
|
|
} else {
|
|
res += getSamplerAtOutputCoords(inInfo, outShapeInfo);
|
|
}
|
|
}
|
|
return res;
|
|
}
|
|
function getPackedOutputSamplingSnippet(outShape, outTexShape, enableShapeUniforms) {
|
|
switch (outShape.length) {
|
|
case 0:
|
|
return getOutputScalarCoords();
|
|
case 1:
|
|
return getOutputPacked1DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
case 2:
|
|
return getOutputPacked2DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
case 3:
|
|
return getOutputPacked3DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
default:
|
|
return getOutputPackedNDCoords(outShape, outTexShape, enableShapeUniforms);
|
|
}
|
|
}
|
|
function getOutputSamplingSnippet(outShape, outTexShape, enableShapeUniforms) {
|
|
switch (outShape.length) {
|
|
case 0:
|
|
return getOutputScalarCoords();
|
|
case 1:
|
|
return getOutput1DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
case 2:
|
|
return getOutput2DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
case 3:
|
|
return getOutput3DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
case 4:
|
|
return getOutput4DCoords(outShape, outTexShape, enableShapeUniforms);
|
|
case 5:
|
|
return getOutput5DCoords(outShape, outTexShape);
|
|
case 6:
|
|
return getOutput6DCoords(outShape, outTexShape);
|
|
default:
|
|
throw new Error(`${outShape.length}-D output sampling is not yet supported`);
|
|
}
|
|
}
|
|
function getFloatTextureSampleSnippet(glsl) {
|
|
return `
|
|
float sampleTexture(sampler2D textureSampler, vec2 uv) {
|
|
return ${glsl.texture2D}(textureSampler, uv).r;
|
|
}
|
|
`;
|
|
}
|
|
function getFloatTextureSetRSnippet(glsl) {
|
|
return `
|
|
void setOutput(float val) {
|
|
${glsl.output} = vec4(val, 0, 0, 0);
|
|
}
|
|
`;
|
|
}
|
|
function getFloatTextureSetRGBASnippet(glsl) {
|
|
return `
|
|
void setOutput(vec4 val) {
|
|
${glsl.output} = val;
|
|
}
|
|
`;
|
|
}
|
|
function getShaderPrefix(glsl) {
|
|
const SHADER_PREFIX = `${glsl.version}
|
|
precision highp float;
|
|
precision highp int;
|
|
precision highp sampler2D;
|
|
${glsl.varyingFs} vec2 resultUV;
|
|
${glsl.defineOutput}
|
|
const vec2 halfCR = vec2(0.5, 0.5);
|
|
|
|
struct ivec5
|
|
{
|
|
int x;
|
|
int y;
|
|
int z;
|
|
int w;
|
|
int u;
|
|
};
|
|
|
|
struct ivec6
|
|
{
|
|
int x;
|
|
int y;
|
|
int z;
|
|
int w;
|
|
int u;
|
|
int v;
|
|
};
|
|
|
|
uniform float NAN;
|
|
${glsl.defineSpecialNaN}
|
|
${glsl.defineSpecialInf}
|
|
${glsl.defineRound}
|
|
|
|
int imod(int x, int y) {
|
|
return x - y * (x / y);
|
|
}
|
|
|
|
int idiv(int a, int b, float sign) {
|
|
int res = a / b;
|
|
int mod = imod(a, b);
|
|
if (sign < 0. && mod != 0) {
|
|
res -= 1;
|
|
}
|
|
return res;
|
|
}
|
|
|
|
//Based on the work of Dave Hoskins
|
|
//https://www.shadertoy.com/view/4djSRW
|
|
#define HASHSCALE1 443.8975
|
|
float random(float seed){
|
|
vec2 p = resultUV * seed;
|
|
vec3 p3 = fract(vec3(p.xyx) * HASHSCALE1);
|
|
p3 += dot(p3, p3.yzx + 19.19);
|
|
return fract((p3.x + p3.y) * p3.z);
|
|
}
|
|
|
|
${SAMPLE_1D_SNIPPET}
|
|
${SAMPLE_2D_SNIPPET}
|
|
${SAMPLE_3D_SNIPPET}
|
|
`;
|
|
return SHADER_PREFIX;
|
|
}
|
|
const SAMPLE_1D_SNIPPET = `
|
|
vec2 uvFromFlat(int texNumR, int texNumC, int index) {
|
|
int texR = index / texNumC;
|
|
int texC = index - texR * texNumC;
|
|
return (vec2(texC, texR) + halfCR) / vec2(texNumC, texNumR);
|
|
}
|
|
vec2 packedUVfrom1D(int texNumR, int texNumC, int index) {
|
|
int texelIndex = index / 2;
|
|
int texR = texelIndex / texNumC;
|
|
int texC = texelIndex - texR * texNumC;
|
|
return (vec2(texC, texR) + halfCR) / vec2(texNumC, texNumR);
|
|
}
|
|
`;
|
|
const SAMPLE_2D_SNIPPET = `
|
|
vec2 packedUVfrom2D(int texelsInLogicalRow, int texNumR,
|
|
int texNumC, int row, int col) {
|
|
int texelIndex = (row / 2) * texelsInLogicalRow + (col / 2);
|
|
int texR = texelIndex / texNumC;
|
|
int texC = texelIndex - texR * texNumC;
|
|
return (vec2(texC, texR) + halfCR) / vec2(texNumC, texNumR);
|
|
}
|
|
`;
|
|
const SAMPLE_3D_SNIPPET = `
|
|
vec2 packedUVfrom3D(int texNumR, int texNumC,
|
|
int texelsInBatch, int texelsInLogicalRow, int b,
|
|
int row, int col) {
|
|
int index = b * texelsInBatch + (row / 2) * texelsInLogicalRow + (col / 2);
|
|
int texR = index / texNumC;
|
|
int texC = index - texR * texNumC;
|
|
return (vec2(texC, texR) + halfCR) / vec2(texNumC, texNumR);
|
|
}
|
|
`;
|
|
const SHADER_PACKED_PREFIX = `
|
|
float getChannel(vec4 frag, vec2 innerDims) {
|
|
vec2 modCoord = mod(innerDims, 2.);
|
|
return modCoord.x == 0. ?
|
|
(modCoord.y == 0. ? frag.r : frag.g) :
|
|
(modCoord.y == 0. ? frag.b : frag.a);
|
|
}
|
|
float getChannel(vec4 frag, int dim) {
|
|
float modCoord = mod(float(dim), 2.);
|
|
return modCoord == 0. ? frag.r : frag.g;
|
|
}
|
|
`;
|
|
function getOutputScalarCoords() {
|
|
return `
|
|
int getOutputCoords() {
|
|
return 0;
|
|
}
|
|
`;
|
|
}
|
|
function getOutputPacked1DCoords(shape, texShape, enableShapeUniforms) {
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
if (packedTexShape[0] === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
int getOutputCoords() {
|
|
return 2 * int(resultUV.x * ceil(float(outTexShape[1]) / 2.0));
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
int getOutputCoords() {
|
|
return 2 * int(resultUV.x * ${packedTexShape[1]}.0);
|
|
}
|
|
`;
|
|
}
|
|
if (packedTexShape[1] === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
int getOutputCoords() {
|
|
return 2 * int(resultUV.y * ceil(float(outTexShape[0]) / 2.0));
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
int getOutputCoords() {
|
|
return 2 * int(resultUV.y * ${packedTexShape[0]}.0);
|
|
}
|
|
`;
|
|
}
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
int getOutputCoords() {
|
|
ivec2 packedTexShape = ivec2(ceil(float(outTexShape[0]) / 2.0), ceil(float(outTexShape[1]) / 2.0));
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(packedTexShape[0], packedTexShape[1]));
|
|
return 2 * (resTexRC.x * packedTexShape[1] + resTexRC.y);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
int getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${packedTexShape[0]}, ${packedTexShape[1]}));
|
|
return 2 * (resTexRC.x * ${packedTexShape[1]} + resTexRC.y);
|
|
}
|
|
`;
|
|
}
|
|
function getOutput1DCoords(shape, texShape, enableShapeUniforms) {
|
|
if (texShape[0] === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
int getOutputCoords() {
|
|
return int(resultUV.x * float(outTexShape[1]));
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
int getOutputCoords() {
|
|
return int(resultUV.x * ${texShape[1]}.0);
|
|
}
|
|
`;
|
|
}
|
|
if (texShape[1] === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
int getOutputCoords() {
|
|
return int(resultUV.y * float(outTexShape[0]));
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
int getOutputCoords() {
|
|
return int(resultUV.y * ${texShape[0]}.0);
|
|
}
|
|
`;
|
|
}
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
int getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(outTexShape[0], outTexShape[1]));
|
|
return resTexRC.x * outTexShape[1] + resTexRC.y;
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
int getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
return resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
}
|
|
`;
|
|
}
|
|
function getOutputPacked3DCoords(shape, texShape, enableShapeUniforms) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec3 getOutputCoords() {
|
|
ivec2 packedTexShape = ivec2(ceil(float(outTexShape[0]) / 2.0), ceil(float(outTexShape[1]) / 2.0));
|
|
int texelsInLogicalRow = int(ceil(float(outShape[2]) / 2.0));
|
|
int texelsInBatch = texelsInLogicalRow * int(ceil(float(outShape[1]) / 2.0));
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(packedTexShape[0], packedTexShape[1]));
|
|
int index = resTexRC.x * packedTexShape[1] + resTexRC.y;
|
|
|
|
int b = index / texelsInBatch;
|
|
index -= b * texelsInBatch;
|
|
|
|
int r = 2 * (index / texelsInLogicalRow);
|
|
int c = imod(index, texelsInLogicalRow) * 2;
|
|
|
|
return ivec3(b, r, c);
|
|
}
|
|
`;
|
|
}
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
const texelsInLogicalRow = Math.ceil(shape[2] / 2);
|
|
const texelsInBatch = texelsInLogicalRow * Math.ceil(shape[1] / 2);
|
|
return `
|
|
ivec3 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${packedTexShape[0]}, ${packedTexShape[1]}));
|
|
int index = resTexRC.x * ${packedTexShape[1]} + resTexRC.y;
|
|
|
|
int b = index / ${texelsInBatch};
|
|
index -= b * ${texelsInBatch};
|
|
|
|
int r = 2 * (index / ${texelsInLogicalRow});
|
|
int c = imod(index, ${texelsInLogicalRow}) * 2;
|
|
|
|
return ivec3(b, r, c);
|
|
}
|
|
`;
|
|
}
|
|
function getOutput3DCoords(shape, texShape, enableShapeUniforms) {
|
|
if (enableShapeUniforms) {
|
|
const coordsFromIndexSnippet2 = getOutputLogicalCoordinatesFromFlatIndexByUniform(["r", "c", "d"], shape);
|
|
return `
|
|
ivec3 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(outTexShape[0], outTexShape[1]));
|
|
int index = resTexRC.x * outTexShape[1] + resTexRC.y;
|
|
${coordsFromIndexSnippet2}
|
|
return ivec3(r, c, d);
|
|
}
|
|
`;
|
|
}
|
|
const coordsFromIndexSnippet = getLogicalCoordinatesFromFlatIndex(["r", "c", "d"], shape);
|
|
return `
|
|
ivec3 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
${coordsFromIndexSnippet}
|
|
return ivec3(r, c, d);
|
|
}
|
|
`;
|
|
}
|
|
function getOutputPackedNDCoords(shape, texShape, enableShapeUniforms) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec4 getOutputCoords() {
|
|
ivec2 packedTexShape = ivec2(ceil(float(outTexShape[0]) / 2.0), ceil(float(outTexShape[1]) / 2.0));
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(packedTexShape[0], packedTexShape[1]));
|
|
int index = resTexRC.x * packedTexShape[1] + resTexRC.y;
|
|
|
|
int texelsInLogicalRow = int(ceil(float(outShape[3]) / 2.0));
|
|
int texelsInBatch = texelsInLogicalRow * int(ceil(float(outShape[2]) / 2.0));
|
|
int texelsInBatchN = texelsInBatch * outShape[1];
|
|
|
|
int b2 = index / texelsInBatchN;
|
|
index -= b2 * texelsInBatchN;
|
|
|
|
int b = index / texelsInBatch;
|
|
index -= b * texelsInBatch;
|
|
|
|
int r = 2 * (index / texelsInLogicalRow);
|
|
int c = imod(index, texelsInLogicalRow) * 2;
|
|
|
|
return ivec4(b2, b, r, c);
|
|
}
|
|
`;
|
|
}
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
const texelsInLogicalRow = Math.ceil(shape[shape.length - 1] / 2);
|
|
const texelsInBatch = texelsInLogicalRow * Math.ceil(shape[shape.length - 2] / 2);
|
|
let texelsInBatchN = texelsInBatch;
|
|
let batches = ``;
|
|
let coords2 = "b, r, c";
|
|
for (let b = 2; b < shape.length - 1; b++) {
|
|
texelsInBatchN *= shape[shape.length - b - 1];
|
|
batches = `
|
|
int b${b} = index / ${texelsInBatchN};
|
|
index -= b${b} * ${texelsInBatchN};
|
|
` + batches;
|
|
coords2 = `b${b}, ` + coords2;
|
|
}
|
|
return `
|
|
ivec${shape.length} getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${packedTexShape[0]}, ${packedTexShape[1]}));
|
|
int index = resTexRC.x * ${packedTexShape[1]} + resTexRC.y;
|
|
|
|
${batches}
|
|
|
|
int b = index / ${texelsInBatch};
|
|
index -= b * ${texelsInBatch};
|
|
|
|
int r = 2 * (index / ${texelsInLogicalRow});
|
|
int c = imod(index, ${texelsInLogicalRow}) * 2;
|
|
|
|
return ivec${shape.length}(${coords2});
|
|
}
|
|
`;
|
|
}
|
|
function getOutput4DCoords(shape, texShape, enableShapeUniforms) {
|
|
if (enableShapeUniforms) {
|
|
const coordsFromIndexSnippet2 = getOutputLogicalCoordinatesFromFlatIndexByUniform(["r", "c", "d", "d2"], shape);
|
|
return `
|
|
ivec4 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(outTexShape[0], outTexShape[1]));
|
|
int index = resTexRC.x * outTexShape[1] + resTexRC.y;
|
|
${coordsFromIndexSnippet2}
|
|
return ivec4(r, c, d, d2);
|
|
}
|
|
`;
|
|
}
|
|
const coordsFromIndexSnippet = getLogicalCoordinatesFromFlatIndex(["r", "c", "d", "d2"], shape);
|
|
return `
|
|
ivec4 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
${coordsFromIndexSnippet}
|
|
return ivec4(r, c, d, d2);
|
|
}
|
|
`;
|
|
}
|
|
function getOutput5DCoords(shape, texShape) {
|
|
const coordsFromIndexSnippet = getLogicalCoordinatesFromFlatIndex(["r", "c", "d", "d2", "d3"], shape);
|
|
return `
|
|
ivec5 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx * vec2(${texShape[0]},
|
|
${texShape[1]}));
|
|
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
|
|
${coordsFromIndexSnippet}
|
|
|
|
ivec5 outShape = ivec5(r, c, d, d2, d3);
|
|
return outShape;
|
|
}
|
|
`;
|
|
}
|
|
function getOutput6DCoords(shape, texShape) {
|
|
const coordsFromIndexSnippet = getLogicalCoordinatesFromFlatIndex(["r", "c", "d", "d2", "d3", "d4"], shape);
|
|
return `
|
|
ivec6 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
|
|
${coordsFromIndexSnippet}
|
|
|
|
ivec6 result = ivec6(r, c, d, d2, d3, d4);
|
|
return result;
|
|
}
|
|
`;
|
|
}
|
|
function getOutputPacked2DCoords(shape, texShape, enableShapeUniforms) {
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
if (arraysEqual(shape, texShape)) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 packedTexShape = ivec2(ceil(float(outTexShape[0]) / 2.0), ceil(float(outTexShape[1]) / 2.0));
|
|
return 2 * ivec2(resultUV.yx * vec2(packedTexShape[0], packedTexShape[1]));
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
return 2 * ivec2(resultUV.yx * vec2(${packedTexShape[0]}, ${packedTexShape[1]}));
|
|
}
|
|
`;
|
|
}
|
|
const texelsInLogicalRow = Math.ceil(shape[1] / 2);
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 packedTexShape = ivec2(ceil(float(outTexShape[0]) / 2.0), ceil(float(outTexShape[1]) / 2.0));
|
|
int texelsInLogicalRow = int(ceil(float(outShape[1]) / 2.0));
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(packedTexShape[0], packedTexShape[1]));
|
|
|
|
int index = resTexRC.x * packedTexShape[1] + resTexRC.y;
|
|
int r = 2 * (index / texelsInLogicalRow);
|
|
int c = imod(index, texelsInLogicalRow) * 2;
|
|
|
|
return ivec2(r, c);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${packedTexShape[0]}, ${packedTexShape[1]}));
|
|
|
|
int index = resTexRC.x * ${packedTexShape[1]} + resTexRC.y;
|
|
int r = 2 * (index / ${texelsInLogicalRow});
|
|
int c = imod(index, ${texelsInLogicalRow}) * 2;
|
|
|
|
return ivec2(r, c);
|
|
}
|
|
`;
|
|
}
|
|
function getOutput2DCoords(shape, texShape, enableShapeUniforms) {
|
|
if (arraysEqual(shape, texShape)) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
return ivec2(resultUV.yx * vec2(outTexShape[0], outTexShape[1]));
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
return ivec2(resultUV.yx * vec2(${texShape[0]}, ${texShape[1]}));
|
|
}
|
|
`;
|
|
}
|
|
if (shape[1] === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(outTexShape[0], outTexShape[1]));
|
|
int index = resTexRC.x * outTexShape[1] + resTexRC.y;
|
|
return ivec2(index, 0);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
return ivec2(index, 0);
|
|
}
|
|
`;
|
|
}
|
|
if (shape[0] === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(outTexShape[0], outTexShape[1]));
|
|
int index = resTexRC.x * outTexShape[1] + resTexRC.y;
|
|
return ivec2(0, index);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
return ivec2(0, index);
|
|
}
|
|
`;
|
|
}
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(outTexShape[0], outTexShape[1]));
|
|
int index = resTexRC.x * outTexShape[1] + resTexRC.y;
|
|
int r = index / outShape[1];
|
|
int c = index - r * outShape[1];
|
|
return ivec2(r, c);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
ivec2 getOutputCoords() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx *
|
|
vec2(${texShape[0]}, ${texShape[1]}));
|
|
int index = resTexRC.x * ${texShape[1]} + resTexRC.y;
|
|
int r = index / ${shape[1]};
|
|
int c = index - r * ${shape[1]};
|
|
return ivec2(r, c);
|
|
}
|
|
`;
|
|
}
|
|
function getFlatOffsetUniformName(texName) {
|
|
return `offset${texName}`;
|
|
}
|
|
function getPackedSamplerScalar(inputInfo) {
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const glsl = getGlslDifferences();
|
|
return `
|
|
vec4 ${funcName}() {
|
|
return ${glsl.texture2D}(${texName}, halfCR);
|
|
}
|
|
`;
|
|
}
|
|
function getSamplerScalar(inputInfo, enableShapeUniforms) {
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `float ${funcName}() {return ${texName};}`;
|
|
}
|
|
const [texNumR, texNumC] = inputInfo.shapeInfo.texShape;
|
|
if (texNumR === 1 && texNumC === 1) {
|
|
return `
|
|
float ${funcName}() {
|
|
return sampleTexture(${texName}, halfCR);
|
|
}
|
|
`;
|
|
}
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}() {
|
|
vec2 uv = uvFromFlat(${texName}TexShape[0], ${texName}TexShape[1], ${offset});
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const [tNumR, tNumC] = inputInfo.shapeInfo.texShape;
|
|
return `
|
|
float ${funcName}() {
|
|
vec2 uv = uvFromFlat(${tNumR}, ${tNumC}, ${offset});
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getPackedSampler1D(inputInfo, enableShapeUniforms) {
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const glsl = getGlslDifferences();
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
vec4 ${funcName}(int index) {
|
|
ivec2 packedTexShape = ivec2(ceil(float(${texName}TexShape[0]) / 2.0), ceil(float(${texName}TexShape[1]) / 2.0));
|
|
vec2 uv = packedUVfrom1D(
|
|
packedTexShape[0], packedTexShape[1], index);
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
return `
|
|
vec4 ${funcName}(int index) {
|
|
vec2 uv = packedUVfrom1D(
|
|
${packedTexShape[0]}, ${packedTexShape[1]}, index);
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getSampler1D(inputInfo, enableShapeUniforms) {
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `
|
|
float ${funcName}(int index) {
|
|
${getUniformSampler(inputInfo)}
|
|
}
|
|
`;
|
|
}
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const tNumR = texShape[0];
|
|
const tNumC = texShape[1];
|
|
if (tNumC === 1 && tNumR === 1) {
|
|
return `
|
|
float ${funcName}(int index) {
|
|
return sampleTexture(${texName}, halfCR);
|
|
}
|
|
`;
|
|
}
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
if (tNumC === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int index) {
|
|
vec2 uv = vec2(0.5, (float(index + ${offset}) + 0.5) / float(${texName}TexShape[0]));
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int index) {
|
|
vec2 uv = vec2(0.5, (float(index + ${offset}) + 0.5) / ${tNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (tNumR === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int index) {
|
|
vec2 uv = vec2((float(index + ${offset}) + 0.5) / float(${texName}TexShape[1]), 0.5);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int index) {
|
|
vec2 uv = vec2((float(index + ${offset}) + 0.5) / ${tNumC}.0, 0.5);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int index) {
|
|
vec2 uv = uvFromFlat(${texName}TexShape[0], ${texName}TexShape[1], index + ${offset});
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int index) {
|
|
vec2 uv = uvFromFlat(${tNumR}, ${tNumC}, index + ${offset});
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getPackedSampler2D(inputInfo, enableShapeUniforms) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const texNumR = texShape[0];
|
|
const texNumC = texShape[1];
|
|
const glsl = getGlslDifferences();
|
|
if (texShape != null && arraysEqual(shape, texShape)) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
vec4 ${funcName}(int row, int col) {
|
|
vec2 uv = (vec2(col, row) + halfCR) / vec2(${texName}TexShape[1], ${texName}TexShape[0]);
|
|
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
vec4 ${funcName}(int row, int col) {
|
|
vec2 uv = (vec2(col, row) + halfCR) / vec2(${texNumC}.0, ${texNumR}.0);
|
|
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
vec4 ${funcName}(int row, int col) {
|
|
ivec2 packedTexShape = ivec2(ceil(float(${texName}TexShape[0]) / 2.0), ceil(float(${texName}TexShape[1]) / 2.0));
|
|
int valuesPerRow = int(ceil(float(${texName}Shape[1]) / 2.0));
|
|
vec2 uv = packedUVfrom2D(valuesPerRow, packedTexShape[0], packedTexShape[1], row, col);
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
const valuesPerRow = Math.ceil(shape[1] / 2);
|
|
return `
|
|
vec4 ${funcName}(int row, int col) {
|
|
vec2 uv = packedUVfrom2D(${valuesPerRow}, ${packedTexShape[0]}, ${packedTexShape[1]}, row, col);
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getSampler2D(inputInfo, enableShapeUniforms) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
if (texShape != null && arraysEqual(shape, texShape)) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
vec2 uv = (vec2(col, row) + halfCR) / vec2(${texName}TexShape[1], ${texName}TexShape[0]);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const texNumR2 = texShape[0];
|
|
const texNumC2 = texShape[1];
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
vec2 uv = (vec2(col, row) + halfCR) / vec2(${texNumC2}.0, ${texNumR2}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const { newShape, keptDims } = squeezeShape(shape);
|
|
const squeezedShape = newShape;
|
|
if (squeezedShape.length < shape.length) {
|
|
const newInputInfo = squeezeInputInfo(inputInfo, squeezedShape);
|
|
const params = ["row", "col"];
|
|
return `
|
|
${getSamplerFromInInfo(newInputInfo, enableShapeUniforms)}
|
|
float ${funcName}(int row, int col) {
|
|
return ${funcName}(${getSqueezedParams(params, keptDims)});
|
|
}
|
|
`;
|
|
}
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
int index = round(dot(vec2(row, col), vec2(${shape[1]}, 1)));
|
|
${getUniformSampler(inputInfo)}
|
|
}
|
|
`;
|
|
}
|
|
const texNumR = texShape[0];
|
|
const texNumC = texShape[1];
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
if (texNumC === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
float index = dot(vec3(row, col, ${offset}), vec3(${texName}Shape[1], 1, 1));
|
|
vec2 uv = vec2(0.5, (index + 0.5) / float(${texName}TexShape[0]));
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
float index = dot(vec3(row, col, ${offset}), vec3(${shape[1]}, 1, 1));
|
|
vec2 uv = vec2(0.5, (index + 0.5) / ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (texNumR === 1) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
float index = dot(vec3(row, col, ${offset}), vec3(${texName}Shape[1], 1, 1));
|
|
vec2 uv = vec2((index + 0.5) / float(${texName}TexShape[1]), 0.5);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
float index = dot(vec3(row, col, ${offset}), vec3(${shape[1]}, 1, 1));
|
|
vec2 uv = vec2((index + 0.5) / ${texNumC}.0, 0.5);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int index = row * ${texName}Shape[1] + col + ${offset};
|
|
vec2 uv = uvFromFlat(${texName}TexShape[0], ${texName}TexShape[1], index);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int index = row * ${shape[1]} + col + ${offset};
|
|
vec2 uv = uvFromFlat(${texNumR}, ${texNumC}, index);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getPackedSampler3D(inputInfo, enableShapeUniforms) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
if (shape[0] === 1) {
|
|
const squeezedShape = shape.slice(1);
|
|
const keptDims = [1, 2];
|
|
const newInputInfo = squeezeInputInfo(inputInfo, squeezedShape);
|
|
const params = ["b", "row", "col"];
|
|
return `
|
|
${getPackedSamplerFromInInfo(newInputInfo, enableShapeUniforms)}
|
|
vec4 ${funcName}(int b, int row, int col) {
|
|
return ${funcName}(${getSqueezedParams(params, keptDims)});
|
|
}
|
|
`;
|
|
}
|
|
const glsl = getGlslDifferences();
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
vec4 ${funcName}(int b, int row, int col) {
|
|
ivec2 packedTexShape = ivec2(ceil(float(${texName}TexShape[0]) / 2.0), ceil(float(${texName}TexShape[1]) / 2.0));
|
|
int valuesPerRow = int(ceil(float(${texName}Shape[2]) / 2.0));
|
|
int texelsInBatch = valuesPerRow * int(ceil(float(${texName}Shape[1]) / 2.0));
|
|
vec2 uv = packedUVfrom3D(
|
|
packedTexShape[0], packedTexShape[1], texelsInBatch, valuesPerRow, b, row, col);
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const texNumR = packedTexShape[0];
|
|
const texNumC = packedTexShape[1];
|
|
const valuesPerRow = Math.ceil(shape[2] / 2);
|
|
const texelsInBatch = valuesPerRow * Math.ceil(shape[1] / 2);
|
|
return `
|
|
vec4 ${funcName}(int b, int row, int col) {
|
|
vec2 uv = packedUVfrom3D(
|
|
${texNumR}, ${texNumC}, ${texelsInBatch}, ${valuesPerRow}, b, row, col);
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getSampler3D(inputInfo, enableShapeUniforms) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const stride0 = shape[1] * shape[2];
|
|
const stride1 = shape[2];
|
|
const { newShape, keptDims } = squeezeShape(shape);
|
|
const squeezedShape = newShape;
|
|
if (squeezedShape.length < shape.length) {
|
|
const newInputInfo = squeezeInputInfo(inputInfo, squeezedShape);
|
|
const params = ["row", "col", "depth"];
|
|
return `
|
|
${getSamplerFromInInfo(newInputInfo, enableShapeUniforms)}
|
|
float ${funcName}(int row, int col, int depth) {
|
|
return ${funcName}(${getSqueezedParams(params, keptDims)});
|
|
}
|
|
`;
|
|
}
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
int index = round(dot(vec3(row, col, depth),
|
|
vec3(${stride0}, ${stride1}, 1)));
|
|
${getUniformSampler(inputInfo)}
|
|
}
|
|
`;
|
|
}
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const texNumR = texShape[0];
|
|
const texNumC = texShape[1];
|
|
const flatOffset = inputInfo.shapeInfo.flatOffset;
|
|
if (texNumC === stride0 && flatOffset == null) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
int stride1 = ${texName}Shape[2];
|
|
float texR = float(row);
|
|
float texC = dot(vec2(col, depth), vec2(stride1, 1));
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texName}TexShape[1], ${texName}TexShape[0]);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
float texR = float(row);
|
|
float texC = dot(vec2(col, depth), vec2(${stride1}, 1));
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (texNumC === stride1 && flatOffset == null) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
float texR = dot(vec2(row, col), vec2(${texName}Shape[1], 1));
|
|
float texC = float(depth);
|
|
vec2 uv = (vec2(texC, texR) + halfCR) / vec2(${texName}TexShape[1], ${texName}TexShape[0]);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
float texR = dot(vec2(row, col), vec2(${shape[1]}, 1));
|
|
float texC = float(depth);
|
|
vec2 uv = (vec2(texC, texR) + halfCR) / vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int stride0 = ${texName}Shape[1] * ${texName}Shape[2];
|
|
int stride1 = ${texName}Shape[2];
|
|
int index = row * stride0 + col * stride1 + depth + ${offset};
|
|
vec2 uv = uvFromFlat(${texName}TexShape[0], ${texName}TexShape[1], index);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col, int depth) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int index = row * ${stride0} + col * ${stride1} + depth + ${offset};
|
|
vec2 uv = uvFromFlat(${texNumR}, ${texNumC}, index);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getPackedSamplerND(inputInfo, enableShapeUniforms) {
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const glsl = getGlslDifferences();
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
vec4 ${funcName}(int b2, int b, int row, int col) {
|
|
int valuesPerRow = int(ceil(float(${texName}Shape[3]) / 2.0));
|
|
int texelsInBatch = valuesPerRow * int(ceil(float(${texName}Shape[2]) / 2.0));
|
|
int index = b * texelsInBatch + (row / 2) * valuesPerRow + (col / 2);
|
|
texelsInBatch *= ${texName}Shape[1];
|
|
index = b2 * texelsInBatch + index;
|
|
ivec2 packedTexShape = ivec2(ceil(float(${texName}TexShape[0]) / 2.0), ceil(float(${texName}TexShape[1]) / 2.0));
|
|
int texR = index / packedTexShape[1];
|
|
int texC = index - texR * packedTexShape[1];
|
|
vec2 uv = (vec2(texC, texR) + halfCR) / vec2(packedTexShape[1], packedTexShape[0]); return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const rank = shape.length;
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const packedTexShape = [Math.ceil(texShape[0] / 2), Math.ceil(texShape[1] / 2)];
|
|
const texNumR = packedTexShape[0];
|
|
const texNumC = packedTexShape[1];
|
|
const valuesPerRow = Math.ceil(shape[rank - 1] / 2);
|
|
let texelsInBatch = valuesPerRow * Math.ceil(shape[rank - 2] / 2);
|
|
let params = `int b, int row, int col`;
|
|
let index = `b * ${texelsInBatch} + (row / 2) * ${valuesPerRow} + (col / 2)`;
|
|
for (let b = 2; b < rank - 1; b++) {
|
|
params = `int b${b}, ` + params;
|
|
texelsInBatch *= shape[rank - b - 1];
|
|
index = `b${b} * ${texelsInBatch} + ` + index;
|
|
}
|
|
return `
|
|
vec4 ${funcName}(${params}) {
|
|
int index = ${index};
|
|
int texR = index / ${texNumC};
|
|
int texC = index - texR * ${texNumC};
|
|
vec2 uv = (vec2(texC, texR) + halfCR) / vec2(${texNumC}, ${texNumR});
|
|
return ${glsl.texture2D}(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getSampler4D(inputInfo, enableShapeUniforms) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const stride2 = shape[3];
|
|
const stride1 = shape[2] * stride2;
|
|
const stride0 = shape[1] * stride1;
|
|
const { newShape, keptDims } = squeezeShape(shape);
|
|
if (newShape.length < shape.length) {
|
|
const newInputInfo = squeezeInputInfo(inputInfo, newShape);
|
|
const params = ["row", "col", "depth", "depth2"];
|
|
return `
|
|
${getSamplerFromInInfo(newInputInfo, enableShapeUniforms)}
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
return ${funcName}(${getSqueezedParams(params, keptDims)});
|
|
}
|
|
`;
|
|
}
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
int index = round(dot(vec4(row, col, depth, depth2),
|
|
vec4(${stride0}, ${stride1}, ${stride2}, 1)));
|
|
${getUniformSampler(inputInfo)}
|
|
}
|
|
`;
|
|
}
|
|
const flatOffset = inputInfo.shapeInfo.flatOffset;
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const texNumR = texShape[0];
|
|
const texNumC = texShape[1];
|
|
const stride2Str = `int stride2 = ${texName}Shape[3];`;
|
|
const stride1Str = `int stride1 = ${texName}Shape[2] * stride2;`;
|
|
const stride0Str = `int stride0 = ${texName}Shape[1] * stride1;`;
|
|
if (texNumC === stride0 && flatOffset == null) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
${stride2Str}
|
|
${stride1Str}
|
|
float texR = float(row);
|
|
float texC =
|
|
dot(vec3(col, depth, depth2),
|
|
vec3(stride1, stride2, 1));
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texName}TexShape[1], ${texName}TexShape[0]);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
float texR = float(row);
|
|
float texC =
|
|
dot(vec3(col, depth, depth2),
|
|
vec3(${stride1}, ${stride2}, 1));
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (texNumC === stride2 && flatOffset == null) {
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
float texR = dot(vec3(row, col, depth),
|
|
vec3(${texName}Shape[1] * ${texName}Shape[2], ${texName}Shape[2], 1));
|
|
float texC = float(depth2);
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texName}TexShape[1], ${texName}TexShape[0]);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
float texR = dot(vec3(row, col, depth),
|
|
vec3(${shape[1] * shape[2]}, ${shape[2]}, 1));
|
|
float texC = float(depth2);
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
if (enableShapeUniforms) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
${stride2Str}
|
|
${stride1Str}
|
|
${stride0Str}
|
|
int index = row * stride0 + col * stride1 +
|
|
depth * stride2 + depth2;
|
|
vec2 uv = uvFromFlat(${texName}TexShape[0], ${texName}TexShape[1], index + ${offset});
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int index = row * ${stride0} + col * ${stride1} +
|
|
depth * ${stride2} + depth2;
|
|
vec2 uv = uvFromFlat(${texNumR}, ${texNumC}, index + ${offset});
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getSampler5D(inputInfo) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const stride3 = shape[4];
|
|
const stride2 = shape[3] * stride3;
|
|
const stride1 = shape[2] * stride2;
|
|
const stride0 = shape[1] * stride1;
|
|
const { newShape, keptDims } = squeezeShape(shape);
|
|
if (newShape.length < shape.length) {
|
|
const newInputInfo = squeezeInputInfo(inputInfo, newShape);
|
|
const params = ["row", "col", "depth", "depth2", "depth3"];
|
|
return `
|
|
${getSamplerFromInInfo(newInputInfo)}
|
|
float ${funcName}(int row, int col, int depth, int depth2, int depth3) {
|
|
return ${funcName}(${getSqueezedParams(params, keptDims)});
|
|
}
|
|
`;
|
|
}
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2, int depth3) {
|
|
float index = dot(
|
|
vec4(row, col, depth, depth2),
|
|
vec4(${stride0}, ${stride1}, ${stride2}, ${stride3})) +
|
|
depth3;
|
|
${getUniformSampler(inputInfo)}
|
|
}
|
|
`;
|
|
}
|
|
const flatOffset = inputInfo.shapeInfo.flatOffset;
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const texNumR = texShape[0];
|
|
const texNumC = texShape[1];
|
|
if (texNumC === stride0 && flatOffset == null) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2, int depth3) {
|
|
int texR = row;
|
|
float texC = dot(vec4(col, depth, depth2, depth3),
|
|
vec4(${stride1}, ${stride2}, ${stride3}, 1));
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (texNumC === stride3 && flatOffset == null) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2, int depth3) {
|
|
float texR = dot(
|
|
vec4(row, col, depth, depth2),
|
|
vec4(${shape[1] * shape[2] * shape[3]},
|
|
${shape[2] * shape[3]}, ${shape[3]}, 1));
|
|
int texC = depth3;
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
return `
|
|
float ${funcName}(int row, int col, int depth, int depth2, int depth3) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int index = row * ${stride0} + col * ${stride1} + depth * ${stride2} +
|
|
depth2 * ${stride3} + depth3 + ${offset};
|
|
vec2 uv = uvFromFlat(${texNumR}, ${texNumC}, index);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getSampler6D(inputInfo) {
|
|
const shape = inputInfo.shapeInfo.logicalShape;
|
|
const texName = inputInfo.name;
|
|
const funcName = "get" + texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const { newShape, keptDims } = squeezeShape(shape);
|
|
if (newShape.length < shape.length) {
|
|
const newInputInfo = squeezeInputInfo(inputInfo, newShape);
|
|
const params = ["row", "col", "depth", "depth2", "depth3", "depth4"];
|
|
return `
|
|
${getSamplerFromInInfo(newInputInfo)}
|
|
float ${funcName}(int row, int col, int depth,
|
|
int depth2, int depth3, int depth4) {
|
|
return ${funcName}(${getSqueezedParams(params, keptDims)});
|
|
}
|
|
`;
|
|
}
|
|
const stride4 = shape[5];
|
|
const stride3 = shape[4] * stride4;
|
|
const stride2 = shape[3] * stride3;
|
|
const stride1 = shape[2] * stride2;
|
|
const stride0 = shape[1] * stride1;
|
|
if (inputInfo.shapeInfo.isUniform) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth,
|
|
int depth2, int depth3, int depth4) {
|
|
int index = round(dot(
|
|
vec4(row, col, depth, depth2),
|
|
vec4(${stride0}, ${stride1}, ${stride2}, ${stride3})) +
|
|
dot(
|
|
vec2(depth3, depth4),
|
|
vec2(${stride4}, 1)));
|
|
${getUniformSampler(inputInfo)}
|
|
}
|
|
`;
|
|
}
|
|
const flatOffset = inputInfo.shapeInfo.flatOffset;
|
|
const texShape = inputInfo.shapeInfo.texShape;
|
|
const texNumR = texShape[0];
|
|
const texNumC = texShape[1];
|
|
if (texNumC === stride0 && flatOffset == null) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth,
|
|
int depth2, int depth3, int depth4) {
|
|
int texR = row;
|
|
float texC = dot(vec4(col, depth, depth2, depth3),
|
|
vec4(${stride1}, ${stride2}, ${stride3}, ${stride4})) +
|
|
float(depth4);
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
if (texNumC === stride4 && flatOffset == null) {
|
|
return `
|
|
float ${funcName}(int row, int col, int depth,
|
|
int depth2, int depth3, int depth4) {
|
|
float texR = dot(vec4(row, col, depth, depth2),
|
|
vec4(${shape[1] * shape[2] * shape[3] * shape[4]},
|
|
${shape[2] * shape[3] * shape[4]},
|
|
${shape[3] * shape[4]},
|
|
${shape[4]})) + float(depth3);
|
|
int texC = depth4;
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${texNumC}.0, ${texNumR}.0);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
const offset = getFlatOffsetUniformName(texName);
|
|
return `
|
|
float ${funcName}(int row, int col, int depth,
|
|
int depth2, int depth3, int depth4) {
|
|
// Explicitly use integer operations as dot() only works on floats.
|
|
int index = row * ${stride0} + col * ${stride1} + depth * ${stride2} +
|
|
depth2 * ${stride3} + depth3 * ${stride4} + depth4 + ${offset};
|
|
vec2 uv = uvFromFlat(${texNumR}, ${texNumC}, index);
|
|
return sampleTexture(${texName}, uv);
|
|
}
|
|
`;
|
|
}
|
|
function getUniformSampler(inputInfo) {
|
|
const texName = inputInfo.name;
|
|
const inSize = sizeFromShape(inputInfo.shapeInfo.logicalShape);
|
|
if (inSize < 2) {
|
|
return `return ${texName};`;
|
|
}
|
|
return `
|
|
for (int i = 0; i < ${inSize}; i++) {
|
|
if (i == index) {
|
|
return ${texName}[i];
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
function getPackedSamplerAtOutputCoords(inputInfo, outShapeInfo) {
|
|
const texName = inputInfo.name;
|
|
const texFuncSnippet = texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const funcName = "get" + texFuncSnippet + "AtOutCoords";
|
|
const inRank = inputInfo.shapeInfo.logicalShape.length;
|
|
const outRank = outShapeInfo.logicalShape.length;
|
|
const broadcastDims = getBroadcastDims(inputInfo.shapeInfo.logicalShape, outShapeInfo.logicalShape);
|
|
const type = getCoordsDataType(outRank);
|
|
const rankDiff = outRank - inRank;
|
|
let coordsSnippet;
|
|
const fields = ["x", "y", "z", "w", "u", "v"];
|
|
if (inRank === 0) {
|
|
coordsSnippet = "";
|
|
} else if (outRank < 2 && broadcastDims.length >= 1) {
|
|
coordsSnippet = "coords = 0;";
|
|
} else {
|
|
coordsSnippet = broadcastDims.map((d) => `coords.${fields[d + rankDiff]} = 0;`).join("\n");
|
|
}
|
|
let unpackedCoordsSnippet = "";
|
|
if (outRank < 2 && inRank > 0) {
|
|
unpackedCoordsSnippet = "coords";
|
|
} else {
|
|
unpackedCoordsSnippet = inputInfo.shapeInfo.logicalShape.map((s, i) => `coords.${fields[i + rankDiff]}`).join(", ");
|
|
}
|
|
let output = `return outputValue;`;
|
|
const inSize = sizeFromShape(inputInfo.shapeInfo.logicalShape);
|
|
const isInputScalar = inSize === 1;
|
|
const outSize = sizeFromShape(outShapeInfo.logicalShape);
|
|
const isOutputScalar = outSize === 1;
|
|
if (inRank === 1 && !isInputScalar && !isOutputScalar) {
|
|
output = `
|
|
return vec4(outputValue.xy, outputValue.xy);
|
|
`;
|
|
} else if (isInputScalar && !isOutputScalar) {
|
|
if (outRank === 1) {
|
|
output = `
|
|
return vec4(outputValue.x, outputValue.x, 0., 0.);
|
|
`;
|
|
} else {
|
|
output = `
|
|
return vec4(outputValue.x);
|
|
`;
|
|
}
|
|
} else if (broadcastDims.length) {
|
|
const rows = inRank - 2;
|
|
const cols = inRank - 1;
|
|
if (broadcastDims.indexOf(rows) > -1 && broadcastDims.indexOf(cols) > -1) {
|
|
output = `return vec4(outputValue.x);`;
|
|
} else if (broadcastDims.indexOf(rows) > -1) {
|
|
output = `return vec4(outputValue.x, outputValue.y, outputValue.x, outputValue.y);`;
|
|
} else if (broadcastDims.indexOf(cols) > -1) {
|
|
output = `return vec4(outputValue.xx, outputValue.zz);`;
|
|
}
|
|
}
|
|
return `
|
|
vec4 ${funcName}() {
|
|
${type} coords = getOutputCoords();
|
|
${coordsSnippet}
|
|
vec4 outputValue = get${texFuncSnippet}(${unpackedCoordsSnippet});
|
|
${output}
|
|
}
|
|
`;
|
|
}
|
|
function getSamplerAtOutputCoords(inputInfo, outShapeInfo) {
|
|
const texName = inputInfo.name;
|
|
const texFuncSnippet = texName.charAt(0).toUpperCase() + texName.slice(1);
|
|
const funcName = "get" + texFuncSnippet + "AtOutCoords";
|
|
const outTexShape = outShapeInfo.texShape;
|
|
const inTexShape = inputInfo.shapeInfo.texShape;
|
|
const inRank = inputInfo.shapeInfo.logicalShape.length;
|
|
const outRank = outShapeInfo.logicalShape.length;
|
|
if (!inputInfo.shapeInfo.isUniform && inRank === outRank && inputInfo.shapeInfo.flatOffset == null && arraysEqual(inTexShape, outTexShape)) {
|
|
return `
|
|
float ${funcName}() {
|
|
return sampleTexture(${texName}, resultUV);
|
|
}
|
|
`;
|
|
}
|
|
const type = getCoordsDataType(outRank);
|
|
const broadcastDims = getBroadcastDims(inputInfo.shapeInfo.logicalShape, outShapeInfo.logicalShape);
|
|
const rankDiff = outRank - inRank;
|
|
let coordsSnippet;
|
|
const fields = ["x", "y", "z", "w", "u", "v"];
|
|
if (inRank === 0) {
|
|
coordsSnippet = "";
|
|
} else if (outRank < 2 && broadcastDims.length >= 1) {
|
|
coordsSnippet = "coords = 0;";
|
|
} else {
|
|
coordsSnippet = broadcastDims.map((d) => `coords.${fields[d + rankDiff]} = 0;`).join("\n");
|
|
}
|
|
let unpackedCoordsSnippet = "";
|
|
if (outRank < 2 && inRank > 0) {
|
|
unpackedCoordsSnippet = "coords";
|
|
} else {
|
|
unpackedCoordsSnippet = inputInfo.shapeInfo.logicalShape.map((s, i) => `coords.${fields[i + rankDiff]}`).join(", ");
|
|
}
|
|
return `
|
|
float ${funcName}() {
|
|
${type} coords = getOutputCoords();
|
|
${coordsSnippet}
|
|
return get${texFuncSnippet}(${unpackedCoordsSnippet});
|
|
}
|
|
`;
|
|
}
|
|
function getCoordsDataType(rank) {
|
|
if (rank <= 1) {
|
|
return "int";
|
|
} else if (rank === 2) {
|
|
return "ivec2";
|
|
} else if (rank === 3) {
|
|
return "ivec3";
|
|
} else if (rank === 4) {
|
|
return "ivec4";
|
|
} else if (rank === 5) {
|
|
return "ivec5";
|
|
} else if (rank === 6) {
|
|
return "ivec6";
|
|
} else {
|
|
throw Error(`GPU for rank ${rank} is not yet supported`);
|
|
}
|
|
}
|
|
function getUniformInfoFromShape(isPacked, shape, texShape) {
|
|
const { newShape, keptDims } = squeezeShape(shape);
|
|
const rank = shape.length;
|
|
const useSqueezePackedShape = isPacked && rank === 3 && shape[0] === 1;
|
|
const squeezeShape$1 = useSqueezePackedShape ? shape.slice(1) : newShape;
|
|
const useSqueezeShape = !isPacked && rank > 1 && !arraysEqual(shape, texShape) && newShape.length < rank || useSqueezePackedShape;
|
|
const uniformShape = useSqueezeShape ? squeezeShape$1 : shape;
|
|
return { useSqueezeShape, uniformShape, keptDims };
|
|
}
|
|
function squeezeInputInfo(inInfo, squeezedShape) {
|
|
const newInputInfo = JSON.parse(JSON.stringify(inInfo));
|
|
newInputInfo.shapeInfo.logicalShape = squeezedShape;
|
|
return newInputInfo;
|
|
}
|
|
function getSqueezedParams(params, keptDims) {
|
|
return keptDims.map((d) => params[d]).join(", ");
|
|
}
|
|
function compileProgram(gpgpu, program, inputs, output) {
|
|
const inputInfos = inputs.map((input, i) => {
|
|
const shapeInfo = {
|
|
logicalShape: input.shape,
|
|
texShape: input.isUniform ? null : input.texData.texShape,
|
|
isUniform: input.isUniform,
|
|
isPacked: input.isUniform ? false : input.texData.isPacked,
|
|
flatOffset: null
|
|
};
|
|
if (input.texData != null && input.texData.slice != null && input.texData.slice.flatOffset > 0) {
|
|
shapeInfo.flatOffset = input.texData.slice.flatOffset;
|
|
}
|
|
return { name: program.variableNames[i], shapeInfo };
|
|
});
|
|
const inShapeInfos = inputInfos.map((x) => x.shapeInfo);
|
|
const outShapeInfo = {
|
|
logicalShape: output.shape,
|
|
texShape: output.texData.texShape,
|
|
isUniform: false,
|
|
isPacked: output.texData.isPacked,
|
|
flatOffset: null
|
|
};
|
|
const source = makeShader(inputInfos, outShapeInfo, program);
|
|
const fragmentShader = createFragmentShader(gpgpu.gl, source);
|
|
const webGLProgram = gpgpu.createProgram(fragmentShader);
|
|
if (!env().get("ENGINE_COMPILE_ONLY")) {
|
|
gpgpu.buildVao(webGLProgram);
|
|
return Object.assign({
|
|
program,
|
|
fragmentShader,
|
|
source,
|
|
webGLProgram,
|
|
inShapeInfos,
|
|
outShapeInfo
|
|
}, getUniformLocations(gpgpu, program, webGLProgram));
|
|
} else {
|
|
return {
|
|
program,
|
|
fragmentShader,
|
|
source,
|
|
webGLProgram,
|
|
inShapeInfos,
|
|
outShapeInfo,
|
|
variablesLocations: null,
|
|
customUniformLocations: null,
|
|
infLoc: null,
|
|
nanLoc: null,
|
|
outShapeLocation: null,
|
|
outShapeStridesLocation: null,
|
|
outTexShapeLocation: null
|
|
};
|
|
}
|
|
}
|
|
function getUniformLocations(gpgpu, program, webGLProgram) {
|
|
const variablesLocations = [];
|
|
const customUniformLocations = [];
|
|
let outShapeLocation;
|
|
let outTexShapeLocation;
|
|
let outShapeStridesLocation;
|
|
let infLoc = null;
|
|
let nanLoc = null;
|
|
nanLoc = gpgpu.getUniformLocation(webGLProgram, "NAN", false);
|
|
if (env().getNumber("WEBGL_VERSION") === 1) {
|
|
infLoc = gpgpu.getUniformLocation(webGLProgram, "INFINITY", false);
|
|
}
|
|
const shouldThrow = false;
|
|
for (const varName of program.variableNames) {
|
|
const varLocs = {
|
|
name: varName,
|
|
uniform: gpgpu.getUniformLocation(webGLProgram, varName, shouldThrow),
|
|
offset: gpgpu.getUniformLocation(webGLProgram, `offset${varName}`, shouldThrow)
|
|
};
|
|
if (program.enableShapeUniforms) {
|
|
varLocs.shape = gpgpu.getUniformLocation(webGLProgram, `${varName}Shape`, shouldThrow);
|
|
varLocs.texShape = gpgpu.getUniformLocation(webGLProgram, `${varName}TexShape`, shouldThrow);
|
|
}
|
|
variablesLocations.push(varLocs);
|
|
}
|
|
if (program.enableShapeUniforms) {
|
|
outShapeLocation = gpgpu.getUniformLocation(webGLProgram, "outShape", shouldThrow);
|
|
outShapeStridesLocation = gpgpu.getUniformLocation(webGLProgram, "outShapeStrides", shouldThrow);
|
|
outTexShapeLocation = gpgpu.getUniformLocation(webGLProgram, "outTexShape", shouldThrow);
|
|
}
|
|
if (program.customUniforms) {
|
|
for (const d of program.customUniforms) {
|
|
customUniformLocations.push(gpgpu.getUniformLocation(webGLProgram, d.name, shouldThrow));
|
|
}
|
|
}
|
|
return {
|
|
variablesLocations,
|
|
customUniformLocations,
|
|
infLoc,
|
|
nanLoc,
|
|
outShapeLocation,
|
|
outShapeStridesLocation,
|
|
outTexShapeLocation
|
|
};
|
|
}
|
|
function validateBinaryAndProgram(shapeInfos, inputs) {
|
|
if (shapeInfos.length !== inputs.length) {
|
|
throw Error(`Binary was compiled with ${shapeInfos.length} inputs, but was executed with ${inputs.length} inputs`);
|
|
}
|
|
shapeInfos.forEach((s, i) => {
|
|
const shapeA = s.logicalShape;
|
|
const input = inputs[i];
|
|
const shapeB = input.shape;
|
|
if (!arraysEqual(shapeA, shapeB)) {
|
|
throw Error(`Binary was compiled with different shapes than the current args. Shapes ${shapeA} and ${shapeB} must match`);
|
|
}
|
|
if (s.isUniform && input.isUniform) {
|
|
return;
|
|
}
|
|
const texShapeA = s.texShape;
|
|
const texShapeB = input.isUniform ? null : input.texData.texShape;
|
|
if (!arraysEqual(texShapeA, texShapeB)) {
|
|
throw Error(`Binary was compiled with different texture shapes than the current args. Shape ${texShapeA} and ${texShapeB} must match`);
|
|
}
|
|
});
|
|
}
|
|
function runProgram(gpgpu, binary, inputs, output, customUniformValues) {
|
|
if (!binary.program.enableShapeUniforms) {
|
|
validateBinaryAndProgram(binary.inShapeInfos, inputs);
|
|
validateBinaryAndProgram([binary.outShapeInfo], [output]);
|
|
}
|
|
const outTex = output.texData.texture;
|
|
const outTexShape = output.texData.texShape;
|
|
if (output.texData.isPacked) {
|
|
gpgpu.setOutputPackedMatrixTexture(outTex.texture, outTexShape[0], outTexShape[1]);
|
|
} else {
|
|
gpgpu.setOutputMatrixTexture(outTex.texture, outTexShape[0], outTexShape[1]);
|
|
}
|
|
gpgpu.setProgram(binary.webGLProgram);
|
|
gpgpu.bindVertexArray(binary.webGLProgram.vao);
|
|
if (env().getNumber("WEBGL_VERSION") === 1) {
|
|
if (binary.infLoc !== null) {
|
|
gpgpu.gl.uniform1f(binary.infLoc, Infinity);
|
|
}
|
|
}
|
|
if (binary.nanLoc !== null) {
|
|
gpgpu.gl.uniform1f(binary.nanLoc, NaN);
|
|
}
|
|
for (let i = 0; i < inputs.length; ++i) {
|
|
const input = inputs[i];
|
|
const { uniform: varLoc, offset: varOffsetLoc, shape: varShapeLoc, texShape: varTexShapeLoc } = binary.variablesLocations[i];
|
|
if (varShapeLoc) {
|
|
const { uniformShape } = getUniformInfoFromShape(binary.program.packedInputs, input.shape, input.texData.texShape);
|
|
switch (uniformShape.length) {
|
|
case 1:
|
|
gpgpu.gl.uniform1iv(varShapeLoc, new Int32Array(uniformShape));
|
|
break;
|
|
case 2:
|
|
gpgpu.gl.uniform2iv(varShapeLoc, new Int32Array(uniformShape));
|
|
break;
|
|
case 3:
|
|
gpgpu.gl.uniform3iv(varShapeLoc, new Int32Array(uniformShape));
|
|
break;
|
|
case 4:
|
|
gpgpu.gl.uniform4iv(varShapeLoc, new Int32Array(uniformShape));
|
|
break;
|
|
}
|
|
}
|
|
if (varTexShapeLoc) {
|
|
gpgpu.gl.uniform2i(varTexShapeLoc, input.texData.texShape[0], input.texData.texShape[1]);
|
|
}
|
|
if (varLoc == null) {
|
|
continue;
|
|
}
|
|
if (input.isUniform) {
|
|
if (sizeFromShape(input.shape) < 2) {
|
|
gpgpu.gl.uniform1f(varLoc, input.uniformValues[0]);
|
|
} else {
|
|
let vals = input.uniformValues;
|
|
if (!(vals instanceof Float32Array)) {
|
|
vals = new Float32Array(vals);
|
|
}
|
|
gpgpu.gl.uniform1fv(varLoc, vals);
|
|
}
|
|
continue;
|
|
}
|
|
if (input.texData.slice != null && varOffsetLoc != null) {
|
|
gpgpu.gl.uniform1i(varOffsetLoc, input.texData.slice.flatOffset);
|
|
}
|
|
gpgpu.setInputMatrixTexture(input.texData.texture.texture, varLoc, i);
|
|
}
|
|
const outShapeLoc = binary.outShapeLocation;
|
|
if (outShapeLoc) {
|
|
switch (output.shape.length) {
|
|
case 1:
|
|
gpgpu.gl.uniform1iv(outShapeLoc, new Int32Array(output.shape));
|
|
break;
|
|
case 2:
|
|
gpgpu.gl.uniform2iv(outShapeLoc, new Int32Array(output.shape));
|
|
break;
|
|
case 3:
|
|
gpgpu.gl.uniform3iv(outShapeLoc, new Int32Array(output.shape));
|
|
break;
|
|
case 4:
|
|
gpgpu.gl.uniform4iv(outShapeLoc, new Int32Array(output.shape));
|
|
break;
|
|
}
|
|
}
|
|
if (binary.outShapeStridesLocation) {
|
|
const strides = computeStrides(output.shape);
|
|
switch (output.shape.length) {
|
|
case 2:
|
|
gpgpu.gl.uniform1iv(binary.outShapeStridesLocation, new Int32Array(strides));
|
|
break;
|
|
case 3:
|
|
gpgpu.gl.uniform2iv(binary.outShapeStridesLocation, new Int32Array(strides));
|
|
break;
|
|
case 4:
|
|
gpgpu.gl.uniform3iv(binary.outShapeStridesLocation, new Int32Array(strides));
|
|
break;
|
|
}
|
|
}
|
|
if (binary.outTexShapeLocation) {
|
|
gpgpu.gl.uniform2i(binary.outTexShapeLocation, output.texData.texShape[0], output.texData.texShape[1]);
|
|
}
|
|
if (binary.program.customUniforms && customUniformValues) {
|
|
for (let i = 0; i < binary.program.customUniforms.length; ++i) {
|
|
const d = binary.program.customUniforms[i];
|
|
const customLoc = binary.customUniformLocations[i];
|
|
const customValue = customUniformValues[i];
|
|
if (d.type === "float") {
|
|
gpgpu.gl.uniform1fv(customLoc, customValue);
|
|
} else if (d.type === "vec2") {
|
|
gpgpu.gl.uniform2fv(customLoc, customValue);
|
|
} else if (d.type === "vec3") {
|
|
gpgpu.gl.uniform3fv(customLoc, customValue);
|
|
} else if (d.type === "vec4") {
|
|
gpgpu.gl.uniform4fv(customLoc, customValue);
|
|
} else if (d.type === "int") {
|
|
gpgpu.gl.uniform1iv(customLoc, customValue);
|
|
} else if (d.type === "ivec2") {
|
|
gpgpu.gl.uniform2iv(customLoc, customValue);
|
|
} else if (d.type === "ivec3") {
|
|
gpgpu.gl.uniform3iv(customLoc, customValue);
|
|
} else if (d.type === "ivec4") {
|
|
gpgpu.gl.uniform4iv(customLoc, customValue);
|
|
} else {
|
|
throw Error(`uniform type ${d.type} is not supported yet.`);
|
|
}
|
|
}
|
|
}
|
|
gpgpu.executeProgram();
|
|
}
|
|
function makeShaderKey(program, inputs, output) {
|
|
let keyInputs = "";
|
|
inputs.concat(output).forEach((x) => {
|
|
const hasOffset = x.texData != null && x.texData.slice != null && x.texData.slice.flatOffset > 0;
|
|
if (program.enableShapeUniforms && !x.isUniform) {
|
|
const xTexShape = x.texData.texShape;
|
|
const { useSqueezeShape, uniformShape, keptDims } = getUniformInfoFromShape(program.packedInputs, x.shape, xTexShape);
|
|
let rank1 = "", rank2 = "", rank34 = "";
|
|
if (uniformShape.length === 1 && program.packedInputs) {
|
|
const packedTexShape = [Math.ceil(xTexShape[0] / 2), Math.ceil(xTexShape[1] / 2)];
|
|
rank1 = `${packedTexShape[0] > 1}_${packedTexShape[1] > 1}`;
|
|
} else if (uniformShape.length === 2 && !program.packedInputs) {
|
|
rank2 = `${uniformShape[0] > 1}_${uniformShape[1] > 1}`;
|
|
} else if (uniformShape.length > 2 && !program.packedInputs) {
|
|
const strides = computeStrides(uniformShape);
|
|
rank34 = `${strides[0] === xTexShape[1]}_${strides[strides.length - 1] === xTexShape[1]}`;
|
|
}
|
|
const xRank = x.shape.length;
|
|
const isLogicalShapTexShapeEqual = uniformShape.length === 2 && arraysEqual(x.shape, xTexShape);
|
|
const isScalar = sizeFromShape(x.shape) === 1;
|
|
const broadcastDims = getBroadcastDims$1(x.shape, output.shape);
|
|
const isInOutTexShapeEqual = !program.packedInputs && xRank === output.shape.length && arraysEqual(xTexShape, output.texData.texShape);
|
|
const isTexShapeGreaterThanOne = program.packedInputs || uniformShape.length > 2 ? "" : `${xTexShape[0] > 1}_${xTexShape[1] > 1}`;
|
|
keyInputs += `${xRank}_${isInOutTexShapeEqual}_${useSqueezeShape ? keptDims : ""}_${uniformShape.length}_${isScalar}_${broadcastDims}_${isLogicalShapTexShapeEqual}_${rank1}_${rank2}_${rank34}_${isTexShapeGreaterThanOne}_${hasOffset}`;
|
|
} else {
|
|
const texShape = x.isUniform ? "uniform" : x.texData.texShape;
|
|
keyInputs += `${x.shape}_${texShape}_${hasOffset}`;
|
|
}
|
|
});
|
|
const keyUserCode = program.userCode;
|
|
let key = program.constructor.name;
|
|
key += "_" + keyInputs + "_" + keyUserCode + `${env().getNumber("WEBGL_VERSION")}`;
|
|
return key;
|
|
}
|
|
function useShapeUniforms(rank) {
|
|
return env().getBool("WEBGL_USE_SHAPES_UNIFORMS") && rank <= 4;
|
|
}
|
|
class DecodeMatrixProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = false;
|
|
this.packedOutput = true;
|
|
this.outPackingScheme = PackingScheme.DENSE;
|
|
this.customUniforms = [{ name: "texShape", type: "ivec2" }];
|
|
const glsl = getGlslDifferences();
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
this.userCode = `
|
|
ivec3 outCoordsFromFlatIndex(int index) {
|
|
${this.enableShapeUniforms ? getOutputLogicalCoordinatesFromFlatIndexByUniform(["r", "c", "d"], outputShape) : getLogicalCoordinatesFromFlatIndex(["r", "c", "d"], outputShape)}
|
|
return ivec3(r, c, d);
|
|
}
|
|
|
|
void main() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx * vec2(texShape[0], texShape[1]));
|
|
int index = 4 * (resTexRC.x * texShape[1] + resTexRC.y);
|
|
|
|
vec4 result = vec4(0.);
|
|
|
|
for (int i=0; i<4; i++) {
|
|
int flatIndex = index + i;
|
|
ivec3 rc = outCoordsFromFlatIndex(flatIndex);
|
|
result[i] = getA(rc.x, rc.y, rc.z);
|
|
}
|
|
|
|
${glsl.output} = result;
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class DecodeMatrixPackedProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outPackingScheme = PackingScheme.DENSE;
|
|
this.customUniforms = [{ name: "texShape", type: "ivec2" }];
|
|
const glsl = getGlslDifferences();
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
this.userCode = `
|
|
ivec3 outCoordsFromFlatIndex(int index) {
|
|
${this.enableShapeUniforms ? getOutputLogicalCoordinatesFromFlatIndexByUniform(["r", "c", "d"], outputShape) : getLogicalCoordinatesFromFlatIndex(["r", "c", "d"], outputShape)}
|
|
return ivec3(r, c, d);
|
|
}
|
|
|
|
void main() {
|
|
ivec2 resTexRC = ivec2(resultUV.yx * vec2(texShape[0], texShape[1]));
|
|
int index = 4 * (resTexRC.x * texShape[1] + resTexRC.y);
|
|
|
|
vec4 result = vec4(0.);
|
|
|
|
for (int i=0; i<4; i++) {
|
|
int flatIndex = index + i;
|
|
ivec3 rc = outCoordsFromFlatIndex(flatIndex);
|
|
result[i] = getChannel(getA(rc.x, rc.y, rc.z), vec2(rc.y, rc.z));
|
|
}
|
|
|
|
${glsl.output} = result;
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class EncodeFloatProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.outTexUsage = TextureUsage.DOWNLOAD;
|
|
const glsl = getGlslDifferences();
|
|
this.outputShape = outputShape;
|
|
this.userCode = `
|
|
${ENCODE_FLOAT_SNIPPET}
|
|
|
|
void main() {
|
|
float x = getAAtOutCoords();
|
|
${glsl.output} = encode_float(x);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class EncodeFloatPackedProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = false;
|
|
this.outTexUsage = TextureUsage.DOWNLOAD;
|
|
const glsl = getGlslDifferences();
|
|
this.outputShape = outputShape;
|
|
this.userCode = `
|
|
${ENCODE_FLOAT_SNIPPET}
|
|
|
|
void main() {
|
|
ivec3 coords = getOutputCoords();
|
|
float x = getChannel(getAAtOutCoords(), vec2(coords.y, coords.z));
|
|
${glsl.output} = encode_float(x);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const CHANNEL_CHAR_TO_INDEX_MAP = {
|
|
"R": 0,
|
|
"G": 1,
|
|
"B": 2,
|
|
"A": 3
|
|
};
|
|
class EncodeMatrixProgram {
|
|
constructor(outputShape, inputIsUnsignedByte = false, usedChannels = "RGBA") {
|
|
this.variableNames = ["A"];
|
|
this.customUniforms = [{ name: "texShape", type: "ivec2" }];
|
|
const glsl = getGlslDifferences();
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
let output = `result`;
|
|
if (inputIsUnsignedByte) {
|
|
output = `floor(result * 255. + 0.5)`;
|
|
}
|
|
let mainLoop = "";
|
|
for (let usedChannelIndex = 0; usedChannelIndex < usedChannels.length; usedChannelIndex++) {
|
|
const curChannel = usedChannels[usedChannelIndex];
|
|
mainLoop += `
|
|
if(offset == ${usedChannelIndex}) {
|
|
result = values[${CHANNEL_CHAR_TO_INDEX_MAP[curChannel]}];
|
|
}`;
|
|
}
|
|
this.userCode = `
|
|
${this.enableShapeUniforms ? getFlatIndexFrom3DOutput() : getFlatIndexFrom3D(outputShape)}
|
|
|
|
void main() {
|
|
ivec3 coords = getOutputCoords();
|
|
int flatIndex = getFlatIndex(coords);
|
|
float result = 0.;
|
|
int offset = imod(flatIndex, ${usedChannels.length});
|
|
|
|
flatIndex = idiv(flatIndex, ${usedChannels.length}, 1.);
|
|
|
|
int r = flatIndex / texShape[1];
|
|
if (r < texShape[0]) {
|
|
int c = imod(flatIndex, texShape[1]);
|
|
vec2 uv = (vec2(c, r) + halfCR) / vec2(texShape[1], texShape[0]);
|
|
vec4 values = ${glsl.texture2D}(A, uv);
|
|
${mainLoop}
|
|
}
|
|
${glsl.output} = vec4(${output}, 0., 0., 0.);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class EncodeMatrixPackedProgram {
|
|
constructor(outputShape, inputIsUnsignedByte = false) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = false;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [{ name: "texShape", type: "ivec2" }];
|
|
const glsl = getGlslDifferences();
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
let mainLoop = "";
|
|
let output = "result";
|
|
if (inputIsUnsignedByte) {
|
|
output = "floor(result * 255. + 0.5)";
|
|
}
|
|
for (let row = 0; row <= 1; row++) {
|
|
for (let col = 0; col <= 1; col++) {
|
|
const channel = row * 2 + col;
|
|
mainLoop += `
|
|
localCoords = coords;
|
|
if(localCoords[2] + ${col} < ${this.enableShapeUniforms ? "outShape[2]" : `${outputShape[2]}`}) {
|
|
localCoords[2] += ${col};
|
|
if (localCoords[1] + ${row} < ${this.enableShapeUniforms ? "outShape[1]" : `${outputShape[1]}`}) {
|
|
localCoords[1] += ${row};
|
|
|
|
flatIndex = getFlatIndex(localCoords);
|
|
offset = imod(flatIndex, 4);
|
|
|
|
flatIndex = idiv(flatIndex, 4, 1.);
|
|
|
|
int r = flatIndex / texShape[1];
|
|
int c = imod(flatIndex, texShape[1]);
|
|
vec2 uv = (vec2(c, r) + halfCR) / vec2(texShape[1], texShape[0]);
|
|
values = ${glsl.texture2D}(A, uv);
|
|
|
|
if (offset == 0) {
|
|
result[${channel}] = values[0];
|
|
} else if (offset == 1) {
|
|
result[${channel}] = values[1];
|
|
} else if (offset == 2) {
|
|
result[${channel}] = values[2];
|
|
} else {
|
|
result[${channel}] = values[3];
|
|
}
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
this.userCode = `
|
|
${this.enableShapeUniforms ? getFlatIndexFrom3DOutput() : getFlatIndexFrom3D(outputShape)}
|
|
|
|
void main() {
|
|
ivec3 coords = getOutputCoords();
|
|
|
|
vec4 result = vec4(0.);
|
|
int flatIndex, r, c, offset;
|
|
ivec3 localCoords;
|
|
vec2 uv;
|
|
vec4 values;
|
|
|
|
${mainLoop}
|
|
|
|
${glsl.output} = ${output};
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function createVertexShader(gl) {
|
|
const glsl = getGlslDifferences();
|
|
const vertexShaderSource = `${glsl.version}
|
|
precision highp float;
|
|
${glsl.attribute} vec3 clipSpacePos;
|
|
${glsl.attribute} vec2 uv;
|
|
${glsl.varyingVs} vec2 resultUV;
|
|
|
|
void main() {
|
|
gl_Position = vec4(clipSpacePos, 1);
|
|
resultUV = uv;
|
|
}`;
|
|
return createVertexShader$1(gl, vertexShaderSource);
|
|
}
|
|
function createVertexBuffer(gl) {
|
|
const vertexArray = new Float32Array([-1, 1, 0, 0, 1, -1, -1, 0, 0, 0, 1, 1, 0, 1, 1, 1, -1, 0, 1, 0]);
|
|
return createStaticVertexBuffer(gl, vertexArray);
|
|
}
|
|
function createIndexBuffer(gl) {
|
|
const triangleVertexIndices = new Uint16Array([0, 1, 2, 2, 1, 3]);
|
|
return createStaticIndexBuffer(gl, triangleVertexIndices);
|
|
}
|
|
function createAndConfigureTexture(gl, width, height, internalFormat, textureFormat, textureType) {
|
|
validateTextureSize(width, height);
|
|
const texture = createTexture(gl);
|
|
const tex2d = gl.TEXTURE_2D;
|
|
callAndCheck(gl, () => gl.bindTexture(tex2d, texture));
|
|
callAndCheck(gl, () => gl.texParameteri(tex2d, gl.TEXTURE_WRAP_S, gl.CLAMP_TO_EDGE));
|
|
callAndCheck(gl, () => gl.texParameteri(tex2d, gl.TEXTURE_WRAP_T, gl.CLAMP_TO_EDGE));
|
|
callAndCheck(gl, () => gl.texParameteri(tex2d, gl.TEXTURE_MIN_FILTER, gl.NEAREST));
|
|
callAndCheck(gl, () => gl.texParameteri(tex2d, gl.TEXTURE_MAG_FILTER, gl.NEAREST));
|
|
if (env().getNumber("WEBGL_VERSION") === 1) {
|
|
callAndCheck(gl, () => gl.texImage2D(tex2d, 0, internalFormat, width, height, 0, textureFormat, textureType, null));
|
|
} else {
|
|
callAndCheck(gl, () => gl.texStorage2D(tex2d, 1, internalFormat, width, height));
|
|
}
|
|
callAndCheck(gl, () => gl.bindTexture(gl.TEXTURE_2D, null));
|
|
return { texture, texShape: [height, width] };
|
|
}
|
|
function getInternalFormatForFloat32MatrixTexture(textureConfig) {
|
|
return textureConfig.internalFormatFloat;
|
|
}
|
|
function createFloat32MatrixTexture(gl, rows, columns, textureConfig) {
|
|
const [width, height] = getUnpackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
return createAndConfigureTexture(gl, width, height, getInternalFormatForFloat32MatrixTexture(textureConfig), textureConfig.textureFormatFloat, gl.FLOAT);
|
|
}
|
|
function getInternalFormatForFloat16MatrixTexture(textureConfig) {
|
|
return textureConfig.internalFormatHalfFloat;
|
|
}
|
|
function createFloat16MatrixTexture(gl, rows, columns, textureConfig) {
|
|
const [width, height] = getUnpackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
return createAndConfigureTexture(gl, width, height, getInternalFormatForFloat16MatrixTexture(textureConfig), textureConfig.textureFormatFloat, textureConfig.textureTypeHalfFloat);
|
|
}
|
|
function getInternalFormatForUnsignedBytesMatrixTexture(textureConfig) {
|
|
return textureConfig.downloadTextureFormat;
|
|
}
|
|
function createUnsignedBytesMatrixTexture(gl, rows, columns, textureConfig) {
|
|
const [width, height] = getUnpackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
return createAndConfigureTexture(gl, width, height, getInternalFormatForUnsignedBytesMatrixTexture(textureConfig), gl.RGBA, gl.UNSIGNED_BYTE);
|
|
}
|
|
function getInternalFormatForPackedMatrixTexture(textureConfig) {
|
|
return textureConfig.internalFormatPackedFloat;
|
|
}
|
|
function createPackedMatrixTexture(gl, rows, columns, textureConfig) {
|
|
const [width, height] = getPackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
return createAndConfigureTexture(gl, width, height, getInternalFormatForPackedMatrixTexture(textureConfig), gl.RGBA, gl.FLOAT);
|
|
}
|
|
function getInternalFormatForFloat16PackedMatrixTexture(textureConfig) {
|
|
return textureConfig.internalFormatPackedHalfFloat;
|
|
}
|
|
function createFloat16PackedMatrixTexture(gl, rows, columns, textureConfig) {
|
|
const [width, height] = getPackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
return createAndConfigureTexture(gl, width, height, getInternalFormatForFloat16PackedMatrixTexture(textureConfig), gl.RGBA, textureConfig.textureTypeHalfFloat);
|
|
}
|
|
function bindVertexProgramAttributeStreams(gl, program, vertexBuffer) {
|
|
const posOffset = 0;
|
|
const uvOffset = 3 * 4;
|
|
const stride = 3 * 4 + 2 * 4;
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ARRAY_BUFFER, vertexBuffer));
|
|
const success = bindVertexBufferToProgramAttribute(gl, program, "clipSpacePos", vertexBuffer, 3, stride, posOffset);
|
|
return success && bindVertexBufferToProgramAttribute(gl, program, "uv", vertexBuffer, 2, stride, uvOffset);
|
|
}
|
|
function uploadDenseMatrixToTexture(gl, texture, width, height, data, textureConfig) {
|
|
callAndCheck(gl, () => gl.bindTexture(gl.TEXTURE_2D, texture));
|
|
let dataForUpload, texelDataType, internalFormat;
|
|
if (data instanceof Uint8Array) {
|
|
dataForUpload = new Uint8Array(width * height * 4);
|
|
texelDataType = gl.UNSIGNED_BYTE;
|
|
internalFormat = gl.RGBA;
|
|
} else {
|
|
dataForUpload = new Float32Array(width * height * 4);
|
|
texelDataType = gl.FLOAT;
|
|
internalFormat = textureConfig.internalFormatPackedFloat;
|
|
}
|
|
dataForUpload.set(data);
|
|
if (env().getNumber("WEBGL_VERSION") === 2) {
|
|
callAndCheck(gl, () => gl.texSubImage2D(gl.TEXTURE_2D, 0, 0, 0, width, height, gl.RGBA, texelDataType, dataForUpload));
|
|
} else {
|
|
callAndCheck(gl, () => gl.texImage2D(gl.TEXTURE_2D, 0, internalFormat, width, height, 0, gl.RGBA, texelDataType, dataForUpload));
|
|
}
|
|
callAndCheck(gl, () => gl.bindTexture(gl.TEXTURE_2D, null));
|
|
}
|
|
function uploadPixelDataToTexture(gl, texture, pixels) {
|
|
callAndCheck(gl, () => gl.bindTexture(gl.TEXTURE_2D, texture));
|
|
if (pixels.data instanceof Uint8Array) {
|
|
if (env().getNumber("WEBGL_VERSION") === 2) {
|
|
callAndCheck(gl, () => gl.texSubImage2D(gl.TEXTURE_2D, 0, 0, 0, pixels.width, pixels.height, gl.RGBA, gl.UNSIGNED_BYTE, pixels.data));
|
|
} else {
|
|
callAndCheck(gl, () => gl.texImage2D(gl.TEXTURE_2D, 0, gl.RGBA, pixels.width, pixels.height, 0, gl.RGBA, gl.UNSIGNED_BYTE, pixels.data));
|
|
}
|
|
} else {
|
|
if (env().getNumber("WEBGL_VERSION") === 2) {
|
|
callAndCheck(gl, () => gl.texSubImage2D(gl.TEXTURE_2D, 0, 0, 0, gl.RGBA, gl.UNSIGNED_BYTE, pixels));
|
|
} else {
|
|
callAndCheck(gl, () => gl.texImage2D(gl.TEXTURE_2D, 0, gl.RGBA, gl.RGBA, gl.UNSIGNED_BYTE, pixels));
|
|
}
|
|
}
|
|
callAndCheck(gl, () => gl.bindTexture(gl.TEXTURE_2D, null));
|
|
}
|
|
function createBufferFromOutputTexture(gl2, rows, columns, textureConfig) {
|
|
const buffer2 = gl2.createBuffer();
|
|
callAndCheck(gl2, () => gl2.bindBuffer(gl2.PIXEL_PACK_BUFFER, buffer2));
|
|
const bytesPerFloat = 4;
|
|
const valuesPerTexel = 4;
|
|
const bufferSizeBytes = bytesPerFloat * valuesPerTexel * rows * columns;
|
|
callAndCheck(gl2, () => gl2.bufferData(gl2.PIXEL_PACK_BUFFER, bufferSizeBytes, gl2.STREAM_READ));
|
|
callAndCheck(gl2, () => gl2.readPixels(0, 0, columns, rows, gl2.RGBA, gl2.FLOAT, 0));
|
|
callAndCheck(gl2, () => gl2.bindBuffer(gl2.PIXEL_PACK_BUFFER, null));
|
|
return buffer2;
|
|
}
|
|
function downloadFloat32MatrixFromBuffer(gl, buffer2, size) {
|
|
const gl2 = gl;
|
|
const downloadTarget = new Float32Array(size);
|
|
gl2.bindBuffer(gl2.PIXEL_PACK_BUFFER, buffer2);
|
|
gl2.getBufferSubData(gl2.PIXEL_PACK_BUFFER, 0, downloadTarget);
|
|
gl2.bindBuffer(gl2.PIXEL_PACK_BUFFER, null);
|
|
return downloadTarget;
|
|
}
|
|
function downloadByteEncodedFloatMatrixFromOutputTexture(gl, rows, columns, textureConfig) {
|
|
const [w, h] = getUnpackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
const numChannels = 4;
|
|
const downloadTarget = new Uint8Array(getUnpackedArraySizeFromMatrixSize(rows * columns, numChannels));
|
|
callAndCheck(gl, () => gl.readPixels(0, 0, w, h, textureConfig.downloadTextureFormat, gl.UNSIGNED_BYTE, downloadTarget));
|
|
return new Float32Array(downloadTarget.buffer);
|
|
}
|
|
function downloadPackedMatrixFromBuffer(gl, buffer2, batch, rows, cols, physicalRows, physicalCols, textureConfig) {
|
|
const gl2 = gl;
|
|
const downloadTarget = new Float32Array(getPackedRGBAArraySizeFromMatrixShape(physicalRows, physicalCols));
|
|
gl2.bindBuffer(gl2.PIXEL_PACK_BUFFER, buffer2);
|
|
gl2.getBufferSubData(gl2.PIXEL_PACK_BUFFER, 0, downloadTarget);
|
|
gl2.bindBuffer(gl2.PIXEL_PACK_BUFFER, null);
|
|
return downloadTarget;
|
|
}
|
|
function downloadMatrixFromPackedOutputTexture(gl, physicalRows, physicalCols) {
|
|
const packedRGBA = new Float32Array(physicalRows * physicalCols * 4);
|
|
callAndCheck(gl, () => gl.readPixels(0, 0, physicalCols, physicalRows, gl.RGBA, gl.FLOAT, packedRGBA));
|
|
return packedRGBA;
|
|
}
|
|
class GPGPUContext {
|
|
constructor(gl) {
|
|
this.outputTexture = null;
|
|
this.program = null;
|
|
this.disposed = false;
|
|
this.itemsToPoll = [];
|
|
const glVersion = env().getNumber("WEBGL_VERSION");
|
|
if (gl != null) {
|
|
this.gl = gl;
|
|
setWebGLContext(glVersion, gl);
|
|
} else {
|
|
this.gl = getWebGLContext(glVersion);
|
|
}
|
|
gl = this.gl;
|
|
if (env().getNumber("WEBGL_VERSION") === 2) {
|
|
const gl2 = gl;
|
|
this.createVertexArray = () => {
|
|
return callAndCheck(gl2, () => gl2.createVertexArray());
|
|
};
|
|
this.bindVertexArray = (vao) => {
|
|
return callAndCheck(gl2, () => gl2.bindVertexArray(vao));
|
|
};
|
|
this.deleteVertexArray = (vao) => {
|
|
return callAndCheck(gl2, () => gl2.deleteVertexArray(vao));
|
|
};
|
|
this.getVertexArray = () => {
|
|
return callAndCheck(gl2, () => gl2.getParameter(gl2.VERTEX_ARRAY_BINDING));
|
|
};
|
|
} else if (gl != null) {
|
|
const ext = gl.getExtension("OES_vertex_array_object");
|
|
if (ext == null) {
|
|
throw new Error("All WebGL1 implementations are expected to offer OES_vertex_array_object.");
|
|
}
|
|
this.createVertexArray = () => {
|
|
return callAndCheck(gl, () => ext.createVertexArrayOES());
|
|
};
|
|
this.bindVertexArray = (vao) => {
|
|
return callAndCheck(gl, () => ext.bindVertexArrayOES(vao));
|
|
};
|
|
this.deleteVertexArray = (vao) => {
|
|
return callAndCheck(gl, () => ext.deleteVertexArrayOES(vao));
|
|
};
|
|
this.getVertexArray = () => {
|
|
return callAndCheck(gl, () => gl.getParameter(ext.VERTEX_ARRAY_BINDING_OES));
|
|
};
|
|
}
|
|
let COLOR_BUFFER_FLOAT = "WEBGL_color_buffer_float";
|
|
const COLOR_BUFFER_HALF_FLOAT = "EXT_color_buffer_half_float";
|
|
this.parallelCompilationExtension = this.gl.getExtension("KHR_parallel_shader_compile");
|
|
if (env().getNumber("WEBGL_VERSION") === 1) {
|
|
const TEXTURE_FLOAT = "OES_texture_float";
|
|
const TEXTURE_HALF_FLOAT = "OES_texture_half_float";
|
|
this.textureFloatExtension = getExtensionOrThrow(this.gl, TEXTURE_FLOAT);
|
|
if (hasExtension(this.gl, TEXTURE_HALF_FLOAT)) {
|
|
this.textureHalfFloatExtension = getExtensionOrThrow(this.gl, TEXTURE_HALF_FLOAT);
|
|
} else if (env().get("WEBGL_FORCE_F16_TEXTURES")) {
|
|
throw new Error("GL context does not support half float textures, yet the environment flag WEBGL_FORCE_F16_TEXTURES is set to true.");
|
|
}
|
|
this.colorBufferFloatExtension = this.gl.getExtension(COLOR_BUFFER_FLOAT);
|
|
if (hasExtension(this.gl, COLOR_BUFFER_HALF_FLOAT)) {
|
|
this.colorBufferHalfFloatExtension = getExtensionOrThrow(this.gl, COLOR_BUFFER_HALF_FLOAT);
|
|
} else if (env().get("WEBGL_FORCE_F16_TEXTURES")) {
|
|
throw new Error("GL context does not support color renderable half floats, yet the environment flag WEBGL_FORCE_F16_TEXTURES is set to true.");
|
|
}
|
|
} else {
|
|
COLOR_BUFFER_FLOAT = "EXT_color_buffer_float";
|
|
if (hasExtension(this.gl, COLOR_BUFFER_FLOAT)) {
|
|
this.colorBufferFloatExtension = this.gl.getExtension(COLOR_BUFFER_FLOAT);
|
|
} else if (hasExtension(this.gl, COLOR_BUFFER_HALF_FLOAT)) {
|
|
this.colorBufferHalfFloatExtension = this.gl.getExtension(COLOR_BUFFER_HALF_FLOAT);
|
|
} else {
|
|
throw new Error("GL context does not support color renderable floats");
|
|
}
|
|
}
|
|
this.vertexBuffer = createVertexBuffer(this.gl);
|
|
this.indexBuffer = createIndexBuffer(this.gl);
|
|
this.framebuffer = createFramebuffer(this.gl);
|
|
this.textureConfig = getTextureConfig(this.gl, this.textureHalfFloatExtension);
|
|
}
|
|
get debug() {
|
|
return env().getBool("DEBUG");
|
|
}
|
|
dispose() {
|
|
if (this.disposed) {
|
|
return;
|
|
}
|
|
if (this.program != null) {
|
|
console.warn("Disposing a GPGPUContext that still has a bound WebGLProgram. This is probably a resource leak, delete the program with GPGPUContext.deleteProgram before disposing.");
|
|
}
|
|
if (this.outputTexture != null) {
|
|
console.warn("Disposing a GPGPUContext that still has a bound output matrix texture. This is probably a resource leak, delete the output matrix texture with GPGPUContext.deleteMatrixTexture before disposing.");
|
|
}
|
|
const gl = this.gl;
|
|
callAndCheck(gl, () => gl.finish());
|
|
callAndCheck(gl, () => gl.bindFramebuffer(gl.FRAMEBUFFER, null));
|
|
callAndCheck(gl, () => gl.deleteFramebuffer(this.framebuffer));
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ARRAY_BUFFER, null));
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ELEMENT_ARRAY_BUFFER, null));
|
|
callAndCheck(gl, () => gl.deleteBuffer(this.indexBuffer));
|
|
this.disposed = true;
|
|
}
|
|
createFloat32MatrixTexture(rows, columns) {
|
|
this.throwIfDisposed();
|
|
return createFloat32MatrixTexture(this.gl, rows, columns, this.textureConfig);
|
|
}
|
|
createFloat16MatrixTexture(rows, columns) {
|
|
this.throwIfDisposed();
|
|
return createFloat16MatrixTexture(this.gl, rows, columns, this.textureConfig);
|
|
}
|
|
createUnsignedBytesMatrixTexture(rows, columns) {
|
|
this.throwIfDisposed();
|
|
return createUnsignedBytesMatrixTexture(this.gl, rows, columns, this.textureConfig);
|
|
}
|
|
uploadPixelDataToTexture(texture, pixels) {
|
|
this.throwIfDisposed();
|
|
uploadPixelDataToTexture(this.gl, texture, pixels);
|
|
}
|
|
uploadDenseMatrixToTexture(texture, width, height, data) {
|
|
this.throwIfDisposed();
|
|
uploadDenseMatrixToTexture(this.gl, texture, width, height, data, this.textureConfig);
|
|
}
|
|
createFloat16PackedMatrixTexture(rows, columns) {
|
|
this.throwIfDisposed();
|
|
return createFloat16PackedMatrixTexture(this.gl, rows, columns, this.textureConfig);
|
|
}
|
|
createPackedMatrixTexture(rows, columns) {
|
|
this.throwIfDisposed();
|
|
return createPackedMatrixTexture(this.gl, rows, columns, this.textureConfig);
|
|
}
|
|
deleteMatrixTexture(texture) {
|
|
this.throwIfDisposed();
|
|
if (this.outputTexture === texture) {
|
|
unbindColorTextureFromFramebuffer(this.gl, this.framebuffer);
|
|
this.outputTexture = null;
|
|
}
|
|
callAndCheck(this.gl, () => this.gl.deleteTexture(texture));
|
|
}
|
|
downloadByteEncodedFloatMatrixFromOutputTexture(texture, rows, columns) {
|
|
return this.downloadMatrixDriver(texture, () => downloadByteEncodedFloatMatrixFromOutputTexture(this.gl, rows, columns, this.textureConfig));
|
|
}
|
|
downloadPackedMatrixFromBuffer(buffer2, batch, rows, columns, physicalRows, physicalCols) {
|
|
return downloadPackedMatrixFromBuffer(this.gl, buffer2, batch, rows, columns, physicalRows, physicalCols, this.textureConfig);
|
|
}
|
|
downloadFloat32MatrixFromBuffer(buffer2, size) {
|
|
return downloadFloat32MatrixFromBuffer(this.gl, buffer2, size);
|
|
}
|
|
createBufferFromTexture(texture, rows, columns) {
|
|
this.bindTextureToFrameBuffer(texture);
|
|
const result = createBufferFromOutputTexture(this.gl, rows, columns, this.textureConfig);
|
|
this.unbindTextureToFrameBuffer();
|
|
return result;
|
|
}
|
|
createAndWaitForFence() {
|
|
const fenceContext = this.createFence(this.gl);
|
|
return this.pollFence(fenceContext);
|
|
}
|
|
createFence(gl) {
|
|
let query;
|
|
let isFencePassed;
|
|
if (env().getBool("WEBGL_FENCE_API_ENABLED")) {
|
|
const gl2 = gl;
|
|
const sync = gl2.fenceSync(gl2.SYNC_GPU_COMMANDS_COMPLETE, 0);
|
|
gl.flush();
|
|
isFencePassed = () => {
|
|
const status = gl2.clientWaitSync(sync, 0, 0);
|
|
return status === gl2.ALREADY_SIGNALED || status === gl2.CONDITION_SATISFIED;
|
|
};
|
|
query = sync;
|
|
} else if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION") > 0) {
|
|
query = this.beginQuery();
|
|
this.endQuery();
|
|
isFencePassed = () => this.isQueryAvailable(query, env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION"));
|
|
} else {
|
|
isFencePassed = () => true;
|
|
}
|
|
return { query, isFencePassed };
|
|
}
|
|
downloadMatrixFromPackedTexture(texture, physicalRows, physicalCols) {
|
|
return this.downloadMatrixDriver(texture, () => downloadMatrixFromPackedOutputTexture(this.gl, physicalRows, physicalCols));
|
|
}
|
|
createProgram(fragmentShader) {
|
|
this.throwIfDisposed();
|
|
const gl = this.gl;
|
|
if (this.vertexShader == null) {
|
|
this.vertexShader = createVertexShader(gl);
|
|
}
|
|
const program = createProgram(gl);
|
|
callAndCheck(gl, () => gl.attachShader(program, this.vertexShader));
|
|
callAndCheck(gl, () => gl.attachShader(program, fragmentShader));
|
|
linkProgram(gl, program);
|
|
const program2 = Object.assign(program, { vao: this.createVertexArray() });
|
|
if (this.debug) {
|
|
validateProgram(gl, program2);
|
|
}
|
|
return program2;
|
|
}
|
|
buildVao(program) {
|
|
this.setProgram(program);
|
|
this.bindVertexArray(program.vao);
|
|
const gl = this.gl;
|
|
callAndCheck(gl, () => gl.bindBuffer(gl.ELEMENT_ARRAY_BUFFER, this.indexBuffer));
|
|
bindVertexProgramAttributeStreams(gl, program, this.vertexBuffer);
|
|
}
|
|
deleteProgram(program) {
|
|
this.throwIfDisposed();
|
|
if (program === this.program) {
|
|
this.program = null;
|
|
}
|
|
if (program != null) {
|
|
callAndCheck(this.gl, () => this.gl.deleteProgram(program));
|
|
this.deleteVertexArray(program.vao);
|
|
}
|
|
}
|
|
setProgram(program) {
|
|
this.throwIfDisposed();
|
|
this.program = program;
|
|
if (this.program != null) {
|
|
if (this.debug) {
|
|
validateProgram(this.gl, this.program);
|
|
}
|
|
}
|
|
callAndCheck(this.gl, () => this.gl.useProgram(program));
|
|
}
|
|
getUniformLocation(program, uniformName, shouldThrow = true) {
|
|
this.throwIfDisposed();
|
|
if (shouldThrow) {
|
|
return getProgramUniformLocationOrThrow(this.gl, program, uniformName);
|
|
} else {
|
|
return getProgramUniformLocation(this.gl, program, uniformName);
|
|
}
|
|
}
|
|
getAttributeLocation(program, attribute) {
|
|
this.throwIfDisposed();
|
|
return callAndCheck(this.gl, () => this.gl.getAttribLocation(program, attribute));
|
|
}
|
|
getUniformLocationNoThrow(program, uniformName) {
|
|
this.throwIfDisposed();
|
|
return this.gl.getUniformLocation(program, uniformName);
|
|
}
|
|
setInputMatrixTexture(inputMatrixTexture, uniformLocation, textureUnit) {
|
|
this.throwIfDisposed();
|
|
this.throwIfNoProgram();
|
|
bindTextureToProgramUniformSampler(this.gl, inputMatrixTexture, uniformLocation, textureUnit);
|
|
}
|
|
setOutputMatrixTexture(outputMatrixTexture, rows, columns) {
|
|
this.setOutputMatrixTextureDriver(outputMatrixTexture, columns, rows);
|
|
}
|
|
setOutputPackedMatrixTexture(outputPackedMatrixTexture, rows, columns) {
|
|
this.throwIfDisposed();
|
|
const [width, height] = getPackedMatrixTextureShapeWidthHeight(rows, columns);
|
|
this.setOutputMatrixTextureDriver(outputPackedMatrixTexture, width, height);
|
|
}
|
|
setOutputMatrixWriteRegion(startRow, numRows, startColumn, numColumns) {
|
|
this.setOutputMatrixWriteRegionDriver(startColumn, startRow, numColumns, numRows);
|
|
}
|
|
setOutputPackedMatrixWriteRegion(startRow, numRows, startColumn, numColumns) {
|
|
throw new Error("setOutputPackedMatrixWriteRegion not implemented.");
|
|
}
|
|
debugValidate() {
|
|
if (this.program != null) {
|
|
validateProgram(this.gl, this.program);
|
|
}
|
|
validateFramebuffer(this.gl);
|
|
}
|
|
executeProgram() {
|
|
this.throwIfDisposed();
|
|
this.throwIfNoProgram();
|
|
const gl = this.gl;
|
|
if (this.debug) {
|
|
const boundVao = this.getVertexArray();
|
|
console.assert(boundVao === this.program.vao, "VAO changed between setProgram and executeProgram!");
|
|
this.debugValidate();
|
|
}
|
|
callAndCheck(gl, () => gl.drawElements(gl.TRIANGLES, 6, gl.UNSIGNED_SHORT, 0));
|
|
}
|
|
blockUntilAllProgramsCompleted() {
|
|
this.throwIfDisposed();
|
|
callAndCheck(this.gl, () => this.gl.finish());
|
|
}
|
|
getQueryTimerExtension() {
|
|
if (this.disjointQueryTimerExtension == null) {
|
|
this.disjointQueryTimerExtension = getExtensionOrThrow(this.gl, env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION") === 2 ? "EXT_disjoint_timer_query_webgl2" : "EXT_disjoint_timer_query");
|
|
}
|
|
return this.disjointQueryTimerExtension;
|
|
}
|
|
getQueryTimerExtensionWebGL2() {
|
|
return this.getQueryTimerExtension();
|
|
}
|
|
getQueryTimerExtensionWebGL1() {
|
|
return this.getQueryTimerExtension();
|
|
}
|
|
beginQuery() {
|
|
if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION") === 2) {
|
|
const gl2 = this.gl;
|
|
const ext2 = this.getQueryTimerExtensionWebGL2();
|
|
const query2 = gl2.createQuery();
|
|
gl2.beginQuery(ext2.TIME_ELAPSED_EXT, query2);
|
|
return query2;
|
|
}
|
|
const ext = this.getQueryTimerExtensionWebGL1();
|
|
const query = ext.createQueryEXT();
|
|
ext.beginQueryEXT(ext.TIME_ELAPSED_EXT, query);
|
|
return query;
|
|
}
|
|
endQuery() {
|
|
if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION") === 2) {
|
|
const gl2 = this.gl;
|
|
const ext2 = this.getQueryTimerExtensionWebGL2();
|
|
gl2.endQuery(ext2.TIME_ELAPSED_EXT);
|
|
return;
|
|
}
|
|
const ext = this.getQueryTimerExtensionWebGL1();
|
|
ext.endQueryEXT(ext.TIME_ELAPSED_EXT);
|
|
}
|
|
async waitForQueryAndGetTime(query) {
|
|
await repeatedTry(() => this.disposed || // while testing contexts are created / disposed
|
|
// in rapid succession, so without this check we
|
|
// may poll for the query timer indefinitely
|
|
this.isQueryAvailable(query, env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION")));
|
|
return this.getQueryTime(query, env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_VERSION"));
|
|
}
|
|
getQueryTime(query, queryTimerVersion) {
|
|
if (queryTimerVersion === 0) {
|
|
return null;
|
|
}
|
|
if (queryTimerVersion === 2) {
|
|
const gl2 = this.gl;
|
|
const timeElapsedNanos = gl2.getQueryParameter(query, gl2.QUERY_RESULT);
|
|
return timeElapsedNanos / 1e6;
|
|
} else {
|
|
const ext = this.getQueryTimerExtensionWebGL1();
|
|
const timeElapsedNanos = ext.getQueryObjectEXT(query, ext.QUERY_RESULT_EXT);
|
|
return timeElapsedNanos / 1e6;
|
|
}
|
|
}
|
|
isQueryAvailable(query, queryTimerVersion) {
|
|
if (queryTimerVersion === 0) {
|
|
return true;
|
|
}
|
|
if (queryTimerVersion === 2) {
|
|
const gl2 = this.gl;
|
|
const ext = this.getQueryTimerExtensionWebGL2();
|
|
const available = gl2.getQueryParameter(query, gl2.QUERY_RESULT_AVAILABLE);
|
|
if (this.disjoint == null) {
|
|
this.disjoint = this.gl.getParameter(ext.GPU_DISJOINT_EXT);
|
|
}
|
|
return available && !this.disjoint;
|
|
} else {
|
|
const ext = this.getQueryTimerExtensionWebGL1();
|
|
const available = ext.getQueryObjectEXT(query, ext.QUERY_RESULT_AVAILABLE_EXT);
|
|
if (this.disjoint == null) {
|
|
this.disjoint = this.gl.getParameter(ext.GPU_DISJOINT_EXT);
|
|
}
|
|
return available && !this.disjoint;
|
|
}
|
|
}
|
|
pollFence(fenceContext) {
|
|
return new Promise((resolve) => {
|
|
this.addItemToPoll(() => fenceContext.isFencePassed(), () => resolve());
|
|
});
|
|
}
|
|
pollItems() {
|
|
const index = linearSearchLastTrue(this.itemsToPoll.map((x) => x.isDoneFn));
|
|
for (let i = 0; i <= index; ++i) {
|
|
const { resolveFn } = this.itemsToPoll[i];
|
|
resolveFn();
|
|
}
|
|
this.itemsToPoll = this.itemsToPoll.slice(index + 1);
|
|
}
|
|
addItemToPoll(isDoneFn, resolveFn) {
|
|
this.itemsToPoll.push({ isDoneFn, resolveFn });
|
|
if (this.itemsToPoll.length > 1) {
|
|
return;
|
|
}
|
|
let scheduleFn = void 0;
|
|
if ("setTimeoutCustom" in env().platform) {
|
|
scheduleFn = env().platform.setTimeoutCustom.bind(env().platform);
|
|
}
|
|
repeatedTry(() => {
|
|
this.pollItems();
|
|
return this.itemsToPoll.length === 0;
|
|
}, () => 0, null, scheduleFn);
|
|
}
|
|
bindTextureToFrameBuffer(texture) {
|
|
this.throwIfDisposed();
|
|
bindColorTextureToFramebuffer(this.gl, texture, this.framebuffer);
|
|
if (this.debug) {
|
|
validateFramebuffer(this.gl);
|
|
}
|
|
}
|
|
unbindTextureToFrameBuffer() {
|
|
if (this.outputTexture != null) {
|
|
bindColorTextureToFramebuffer(this.gl, this.outputTexture, this.framebuffer);
|
|
if (this.debug) {
|
|
validateFramebuffer(this.gl);
|
|
}
|
|
} else {
|
|
unbindColorTextureFromFramebuffer(this.gl, this.framebuffer);
|
|
}
|
|
}
|
|
downloadMatrixDriver(texture, downloadAndDecode) {
|
|
this.bindTextureToFrameBuffer(texture);
|
|
const result = downloadAndDecode();
|
|
this.unbindTextureToFrameBuffer();
|
|
return result;
|
|
}
|
|
setOutputMatrixTextureDriver(outputMatrixTextureMaybePacked, width, height) {
|
|
this.throwIfDisposed();
|
|
const gl = this.gl;
|
|
bindColorTextureToFramebuffer(gl, outputMatrixTextureMaybePacked, this.framebuffer);
|
|
if (this.debug) {
|
|
validateFramebuffer(gl);
|
|
}
|
|
this.outputTexture = outputMatrixTextureMaybePacked;
|
|
callAndCheck(gl, () => gl.viewport(0, 0, width, height));
|
|
callAndCheck(gl, () => gl.scissor(0, 0, width, height));
|
|
}
|
|
setOutputMatrixWriteRegionDriver(x, y, width, height) {
|
|
this.throwIfDisposed();
|
|
callAndCheck(this.gl, () => this.gl.scissor(x, y, width, height));
|
|
}
|
|
throwIfDisposed() {
|
|
if (this.disposed) {
|
|
throw new Error("Attempted to use disposed GPGPUContext.");
|
|
}
|
|
}
|
|
throwIfNoProgram() {
|
|
if (this.program == null) {
|
|
throw new Error("No GPU program is currently set.");
|
|
}
|
|
}
|
|
}
|
|
function linearSearchLastTrue(arr) {
|
|
let i = 0;
|
|
for (; i < arr.length; ++i) {
|
|
const isDone = arr[i]();
|
|
if (!isDone) {
|
|
break;
|
|
}
|
|
}
|
|
return i - 1;
|
|
}
|
|
function assertNotComplex(tensor2, opName) {
|
|
if (!Array.isArray(tensor2)) {
|
|
tensor2 = [tensor2];
|
|
}
|
|
tensor2.forEach((t2) => {
|
|
if (t2 != null) {
|
|
assert(t2.dtype !== "complex64", () => `${opName} does not support complex64 tensors in the CPU backend.`);
|
|
}
|
|
});
|
|
}
|
|
function simpleAbsImpl(vals) {
|
|
const resultValues = new Float32Array(vals.length);
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
resultValues[i] = Math.abs(vals[i]);
|
|
}
|
|
return resultValues;
|
|
}
|
|
const abs$1 = (args) => {
|
|
const { x } = args.inputs;
|
|
const cpuBackend = args.backend;
|
|
assertNotComplex(x, "abs");
|
|
let resultValues = new Float32Array(sizeFromShape(x.shape));
|
|
const values = cpuBackend.data.get(x.dataId).values;
|
|
resultValues = simpleAbsImpl(values);
|
|
return cpuBackend.makeOutput(resultValues, x.shape, x.dtype);
|
|
};
|
|
const absConfig$1 = {
|
|
kernelName: Abs,
|
|
backendName: "cpu",
|
|
kernelFunc: abs$1
|
|
};
|
|
function createSimpleBinaryKernelImpl(op2) {
|
|
return (aShape, bShape, aVals, bVals, dtype) => {
|
|
const newShape = assertAndGetBroadcastShape(aShape, bShape);
|
|
const resultRank = newShape.length;
|
|
const resultStrides = computeStrides(newShape);
|
|
const resultSize = sizeFromShape(newShape);
|
|
const result = getTypedArrayFromDType(dtype, resultSize);
|
|
const aRank = aShape.length;
|
|
const bRank = bShape.length;
|
|
const aStrides = computeStrides(aShape);
|
|
const bStrides = computeStrides(bShape);
|
|
const aBroadcastDims = getBroadcastDims$1(aShape, newShape);
|
|
const bBroadcastDims = getBroadcastDims$1(bShape, newShape);
|
|
if (aBroadcastDims.length + bBroadcastDims.length === 0) {
|
|
for (let i = 0; i < result.length; ++i) {
|
|
result[i] = op2(aVals[i % aVals.length], bVals[i % bVals.length]);
|
|
}
|
|
} else {
|
|
for (let i = 0; i < result.length; ++i) {
|
|
const loc = indexToLoc(i, resultRank, resultStrides);
|
|
const aLoc = loc.slice(-aRank);
|
|
aBroadcastDims.forEach((d) => aLoc[d] = 0);
|
|
const aIndex = locToIndex(aLoc, aRank, aStrides);
|
|
const bLoc = loc.slice(-bRank);
|
|
bBroadcastDims.forEach((d) => bLoc[d] = 0);
|
|
const bIndex = locToIndex(bLoc, bRank, bStrides);
|
|
result[i] = op2(aVals[aIndex], bVals[bIndex]);
|
|
}
|
|
}
|
|
return [result, newShape];
|
|
};
|
|
}
|
|
function complex$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { real: real2, imag: imag2 } = inputs;
|
|
const realVals = backend2.data.get(real2.dataId).values;
|
|
const imagVals = backend2.data.get(imag2.dataId).values;
|
|
const complexInfo = backend2.makeTensorInfo(real2.shape, "complex64");
|
|
const complex2 = backend2.data.get(complexInfo.dataId);
|
|
complex2.complexTensorInfos = {
|
|
real: backend2.makeTensorInfo(real2.shape, "float32", realVals),
|
|
imag: backend2.makeTensorInfo(imag2.shape, "float32", imagVals)
|
|
};
|
|
return complexInfo;
|
|
}
|
|
const complexConfig$1 = {
|
|
kernelName: Complex,
|
|
backendName: "cpu",
|
|
kernelFunc: complex$1
|
|
};
|
|
function zeros(backend2, shape, dtype = "float32") {
|
|
if (dtype === "complex64") {
|
|
const real2 = zeros(backend2, shape, "float32");
|
|
const imag2 = zeros(backend2, shape, "float32");
|
|
return complex$1({ inputs: { real: real2, imag: imag2 }, backend: backend2 });
|
|
}
|
|
const values = makeZerosTypedArray(sizeFromShape(shape), dtype);
|
|
return backend2.makeTensorInfo(shape, dtype, values);
|
|
}
|
|
function identity$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
backend2.incRef(x.dataId);
|
|
return { dataId: x.dataId, shape: x.shape, dtype: x.dtype };
|
|
}
|
|
const identityConfig$1 = {
|
|
kernelName: Identity,
|
|
backendName: "cpu",
|
|
kernelFunc: identity$1
|
|
};
|
|
function real$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
const real2 = backend2.data.get(input.dataId).complexTensorInfos.real;
|
|
const realVal = backend2.data.get(real2.dataId).values;
|
|
return backend2.makeTensorInfo(real2.shape, real2.dtype, realVal);
|
|
}
|
|
const realConfig$1 = {
|
|
kernelName: Real,
|
|
backendName: "cpu",
|
|
kernelFunc: real$1
|
|
};
|
|
function castImpl(values, shape, inputType, dtype) {
|
|
if (dtype === "int32") {
|
|
const resultValues = Int32Array.from(values);
|
|
return [shape, "int32", resultValues];
|
|
}
|
|
if (dtype === "bool") {
|
|
const zero = toTypedArray([0], inputType);
|
|
const [resultData, resultShape] = createSimpleBinaryKernelImpl((a, b) => a !== b ? 1 : 0)(shape, [], values, zero, "bool");
|
|
return [resultShape, "bool", resultData];
|
|
}
|
|
throw new Error(`Error in Cast: failed to cast ${inputType} to ${dtype}`);
|
|
}
|
|
function cast$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { dtype } = attrs;
|
|
if (dtype === "complex64") {
|
|
if (x.dtype === "complex64") {
|
|
return identity$1({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const zerosTensorInfo = zeros(backend2, x.shape, x.dtype);
|
|
const floatX = cast$1({ inputs: { x }, backend: backend2, attrs: { dtype: "float32" } });
|
|
const result = complex$1({ inputs: { real: floatX, imag: zerosTensorInfo }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(zerosTensorInfo);
|
|
backend2.disposeIntermediateTensorInfo(floatX);
|
|
return result;
|
|
}
|
|
if (x.dtype === "complex64") {
|
|
const realPart = real$1({ inputs: { input: x }, backend: backend2 });
|
|
const result = cast$1({ inputs: { x: realPart }, backend: backend2, attrs: { dtype } });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
return result;
|
|
}
|
|
if (!hasEncodingLoss(x.dtype, dtype)) {
|
|
const result = identity$1({ inputs: { x }, backend: backend2 });
|
|
return { dataId: result.dataId, shape: result.shape, dtype };
|
|
}
|
|
const values = backend2.data.get(x.dataId).values;
|
|
const [resultShape, resultType, resultData] = castImpl(values, x.shape, x.dtype, dtype);
|
|
return backend2.makeTensorInfo(resultShape, resultType, resultData);
|
|
}
|
|
const castConfig$1 = {
|
|
kernelName: Cast,
|
|
backendName: "cpu",
|
|
kernelFunc: cast$1
|
|
};
|
|
function binaryKernelFunc$1(name, simpleImpl, complexImpl, dtype) {
|
|
if (complexImpl == null) {
|
|
return ({ inputs, backend: backend2 }) => {
|
|
const { a, b } = inputs;
|
|
const cpuBackend = backend2;
|
|
assertNotComplex([a, b], name);
|
|
const aVals = cpuBackend.data.get(a.dataId).values;
|
|
const bVals = cpuBackend.data.get(b.dataId).values;
|
|
const decodedAVals = a.dtype === "string" ? (
|
|
// tslint:disable-next-line: no-any
|
|
fromUint8ToStringArray(aVals)
|
|
) : aVals;
|
|
const decodedBVals = a.dtype === "string" ? (
|
|
// tslint:disable-next-line: no-any
|
|
fromUint8ToStringArray(bVals)
|
|
) : bVals;
|
|
const $dtype = dtype || a.dtype;
|
|
const [resultData, resultShape] = simpleImpl(a.shape, b.shape, decodedAVals, decodedBVals, $dtype);
|
|
return cpuBackend.makeTensorInfo(resultShape, $dtype, resultData);
|
|
};
|
|
}
|
|
return ({ inputs, backend: backend2 }) => {
|
|
const { a, b } = inputs;
|
|
const cpuBackend = backend2;
|
|
if (a.dtype === "complex64" || b.dtype === "complex64") {
|
|
const $aComplex = cast$1({ inputs: { x: a }, backend: cpuBackend, attrs: { dtype: "complex64" } });
|
|
const $aComplexVals = cpuBackend.data.get($aComplex.dataId);
|
|
const aReal = $aComplexVals.complexTensorInfos.real;
|
|
const aImag = $aComplexVals.complexTensorInfos.imag;
|
|
const aRealVals = cpuBackend.data.get(aReal.dataId).values;
|
|
const aImagVals = cpuBackend.data.get(aImag.dataId).values;
|
|
const $bComplex = cast$1({ inputs: { x: b }, backend: cpuBackend, attrs: { dtype: "complex64" } });
|
|
const $bComplexVals = cpuBackend.data.get($bComplex.dataId);
|
|
const bReal = $bComplexVals.complexTensorInfos.real;
|
|
const bImag = $bComplexVals.complexTensorInfos.imag;
|
|
const bRealVals = cpuBackend.data.get(bReal.dataId).values;
|
|
const bImagVals = cpuBackend.data.get(bImag.dataId).values;
|
|
const [resultRealData, resultImagData, resultShape] = complexImpl(a.shape, b.shape, aRealVals, aImagVals, bRealVals, bImagVals);
|
|
const resultReal = cpuBackend.makeTensorInfo(resultShape, "float32", resultRealData);
|
|
const resultImag = cpuBackend.makeTensorInfo(resultShape, "float32", resultImagData);
|
|
const result = complex$1({ inputs: { real: resultReal, imag: resultImag }, backend: cpuBackend });
|
|
cpuBackend.disposeIntermediateTensorInfo($aComplex);
|
|
cpuBackend.disposeIntermediateTensorInfo($bComplex);
|
|
cpuBackend.disposeIntermediateTensorInfo(resultReal);
|
|
cpuBackend.disposeIntermediateTensorInfo(resultImag);
|
|
return result;
|
|
} else {
|
|
const aVals = cpuBackend.data.get(a.dataId).values;
|
|
const bVals = cpuBackend.data.get(b.dataId).values;
|
|
const $dtype = dtype || a.dtype;
|
|
const [resultData, resultShape] = simpleImpl(a.shape, b.shape, aVals, bVals, $dtype);
|
|
return cpuBackend.makeTensorInfo(resultShape, $dtype, resultData);
|
|
}
|
|
};
|
|
}
|
|
function createComplexBinaryKernelImpl(op2) {
|
|
return (aShape, bShape, aRealVals, aImagVals, bRealVals, bImagVals) => {
|
|
const resultShape = assertAndGetBroadcastShape(aShape, bShape);
|
|
const resultSize = sizeFromShape(resultShape);
|
|
const resultRank = resultShape.length;
|
|
const resultStrides = computeStrides(resultShape);
|
|
const resultRealVals = getTypedArrayFromDType("float32", resultSize);
|
|
const resultImagVals = getTypedArrayFromDType("float32", resultSize);
|
|
const aBroadcastDims = getBroadcastDims$1(aShape, resultShape);
|
|
const bBroadcastDims = getBroadcastDims$1(bShape, resultShape);
|
|
const aVals = mergeRealAndImagArrays(aRealVals, aImagVals);
|
|
const bVals = mergeRealAndImagArrays(bRealVals, bImagVals);
|
|
const aRank = aShape.length;
|
|
const aStrides = computeStrides(aShape);
|
|
const bRank = bShape.length;
|
|
const bStrides = computeStrides(bShape);
|
|
if (aBroadcastDims.length + bBroadcastDims.length === 0) {
|
|
for (let i = 0; i < resultRealVals.length; i++) {
|
|
const aIdx = i % aVals.length;
|
|
const bIdx = i % bVals.length;
|
|
const result = op2(aVals[aIdx * 2], aVals[aIdx * 2 + 1], bVals[bIdx * 2], bVals[bIdx * 2 + 1]);
|
|
resultRealVals[i] = result.real;
|
|
resultImagVals[i] = result.imag;
|
|
}
|
|
} else {
|
|
for (let i = 0; i < resultRealVals.length; i++) {
|
|
const loc = indexToLoc(i, resultRank, resultStrides);
|
|
const aLoc = loc.slice(-aRank);
|
|
aBroadcastDims.forEach((d) => aLoc[d] = 0);
|
|
const aIndex = locToIndex(aLoc, aRank, aStrides);
|
|
const bLoc = loc.slice(-bRank);
|
|
bBroadcastDims.forEach((d) => bLoc[d] = 0);
|
|
const bIndex = locToIndex(bLoc, bRank, bStrides);
|
|
const opResult = op2(aVals[aIndex * 2], aVals[aIndex * 2 + 1], bVals[bIndex * 2], bVals[bIndex * 2 + 1]);
|
|
resultRealVals[i] = opResult.real;
|
|
resultImagVals[i] = opResult.imag;
|
|
}
|
|
}
|
|
return [resultRealVals, resultImagVals, resultShape];
|
|
};
|
|
}
|
|
const addImpl = createSimpleBinaryKernelImpl(((a, b) => a + b));
|
|
const addComplexImpl = createComplexBinaryKernelImpl(((aReal, aImag, bReal, bImag) => {
|
|
return { real: aReal + bReal, imag: aImag + bImag };
|
|
}));
|
|
const add = binaryKernelFunc$1(Add, addImpl, addComplexImpl);
|
|
const addConfig$1 = {
|
|
kernelName: Add,
|
|
backendName: "cpu",
|
|
kernelFunc: add
|
|
};
|
|
function bincountImpl(xVals, weightsVals, weightsDtype, weightsShape, size) {
|
|
const weightsSize = sizeFromShape(weightsShape);
|
|
const outVals = makeZerosTypedArray(size, weightsDtype);
|
|
for (let i = 0; i < xVals.length; i++) {
|
|
const value = xVals[i];
|
|
if (value < 0) {
|
|
throw new Error("Input x must be non-negative!");
|
|
}
|
|
if (value >= size) {
|
|
continue;
|
|
}
|
|
if (weightsSize > 0) {
|
|
outVals[value] += weightsVals[i];
|
|
} else {
|
|
outVals[value] += 1;
|
|
}
|
|
}
|
|
return outVals;
|
|
}
|
|
function bincountReduceImpl(xBuf, weightsBuf, size, binaryOutput = false) {
|
|
const numRows = xBuf.shape[0];
|
|
const numCols = xBuf.shape[1];
|
|
const outBuf = buffer([numRows, size], weightsBuf.dtype);
|
|
for (let i = 0; i < numRows; i++) {
|
|
for (let j2 = 0; j2 < numCols; j2++) {
|
|
const value = xBuf.get(i, j2);
|
|
if (value < 0) {
|
|
throw new Error("Input x must be non-negative!");
|
|
}
|
|
if (value >= size) {
|
|
continue;
|
|
}
|
|
if (binaryOutput) {
|
|
outBuf.set(1, i, value);
|
|
} else {
|
|
if (weightsBuf.size > 0) {
|
|
outBuf.set(outBuf.get(i, value) + weightsBuf.get(i, j2), i, value);
|
|
} else {
|
|
outBuf.set(outBuf.get(i, value) + 1, i, value);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return outBuf;
|
|
}
|
|
const bitwiseAndImpl = createSimpleBinaryKernelImpl(((a, b) => a & b));
|
|
const bitwiseAnd$1 = binaryKernelFunc$1(BitwiseAnd, bitwiseAndImpl);
|
|
const bitwiseAndConfig$1 = {
|
|
kernelName: BitwiseAnd,
|
|
backendName: "cpu",
|
|
kernelFunc: bitwiseAnd$1
|
|
};
|
|
function createSimpleUnaryImpl(op2) {
|
|
return (values, dtype, attrs) => {
|
|
const newValues = getArrayFromDType(dtype, values.length);
|
|
for (let i = 0; i < values.length; ++i) {
|
|
newValues[i] = op2(values[i], attrs);
|
|
}
|
|
return newValues;
|
|
};
|
|
}
|
|
function unaryKernelFunc$1(name, op2, dtype) {
|
|
const impl = createSimpleUnaryImpl(op2);
|
|
return unaryKernelFuncFromImpl(name, impl, dtype);
|
|
}
|
|
function unaryKernelFuncFromImpl(name, unaryImpl, dtype) {
|
|
return ({ inputs, attrs, backend: backend2 }) => {
|
|
const { x } = inputs;
|
|
assertNotComplex(x, name);
|
|
const cpuBackend = backend2;
|
|
const values = cpuBackend.data.get(x.dataId).values;
|
|
let decoded;
|
|
if (x.dtype === "string") {
|
|
if (!Array.isArray(values)) {
|
|
throw new Error("String tensor's value was not an instance of Array");
|
|
}
|
|
decoded = fromUint8ToStringArray(values);
|
|
} else {
|
|
decoded = values;
|
|
}
|
|
const $dtype = dtype || x.dtype;
|
|
const newValues = unaryImpl(decoded, $dtype, attrs);
|
|
return cpuBackend.makeTensorInfo(x.shape, $dtype, newValues);
|
|
};
|
|
}
|
|
const ceilImpl = createSimpleUnaryImpl((xi) => Math.ceil(xi));
|
|
const ceil$1 = unaryKernelFuncFromImpl(Ceil, ceilImpl);
|
|
const ceilConfig$1 = {
|
|
kernelName: Ceil,
|
|
backendName: "cpu",
|
|
kernelFunc: ceil$1
|
|
};
|
|
function concatImpl$1(inputs, outShape, dtype, simplyConcat) {
|
|
const outVals = getArrayFromDType(dtype, sizeFromShape(outShape));
|
|
if (simplyConcat && dtype !== "string") {
|
|
let offset = 0;
|
|
inputs.forEach((input) => {
|
|
const size = sizeFromShape(input.shape);
|
|
outVals.set(input.vals, offset);
|
|
offset += size;
|
|
});
|
|
} else {
|
|
let colOffset = 0;
|
|
inputs.forEach((input) => {
|
|
const decodedData = dtype === "string" ? fromUint8ToStringArray(input.vals) : input.vals;
|
|
let tIdx = 0;
|
|
for (let row = 0; row < input.shape[0]; ++row) {
|
|
const resIdx = row * outShape[1] + colOffset;
|
|
for (let col = 0; col < input.shape[1]; ++col) {
|
|
outVals[resIdx + col] = decodedData[tIdx++];
|
|
}
|
|
}
|
|
colOffset += input.shape[1];
|
|
});
|
|
}
|
|
return outVals;
|
|
}
|
|
const equalImpl = createSimpleBinaryKernelImpl((a, b) => a === b ? 1 : 0);
|
|
const equal$1 = binaryKernelFunc$1(Equal, equalImpl, null, "bool");
|
|
const equalConfig$1 = {
|
|
kernelName: Equal,
|
|
backendName: "cpu",
|
|
kernelFunc: equal$1
|
|
};
|
|
const expImpl = createSimpleUnaryImpl((xi) => Math.exp(xi));
|
|
const exp$1 = unaryKernelFuncFromImpl(Exp, expImpl, "float32");
|
|
const expConfig$1 = {
|
|
kernelName: Exp,
|
|
backendName: "cpu",
|
|
kernelFunc: exp$1
|
|
};
|
|
const expm1Impl = createSimpleUnaryImpl((xi) => Math.expm1(xi));
|
|
const expm1$1 = unaryKernelFuncFromImpl(Expm1, expm1Impl);
|
|
const expm1Config$1 = {
|
|
kernelName: Expm1,
|
|
backendName: "cpu",
|
|
kernelFunc: expm1$1
|
|
};
|
|
const floorImpl = createSimpleUnaryImpl((xi) => Math.floor(xi));
|
|
const floor$1 = unaryKernelFuncFromImpl(Floor, floorImpl);
|
|
const floorConfig$1 = {
|
|
kernelName: Floor,
|
|
backendName: "cpu",
|
|
kernelFunc: floor$1
|
|
};
|
|
const floorDivImpl = createSimpleBinaryKernelImpl((a, b) => Math.floor(a / b));
|
|
const floorDiv$1 = binaryKernelFunc$1(FloorDiv, floorDivImpl, null, "int32");
|
|
const floorDivConfig$1 = {
|
|
kernelName: FloorDiv,
|
|
backendName: "cpu",
|
|
kernelFunc: floorDiv$1
|
|
};
|
|
function gatherNdImpl(indicesData, paramsBuf, dtype, numSlices, sliceRank, sliceSize, strides, paramsShape, paramsSize) {
|
|
const outBuf = buffer([numSlices, sliceSize], dtype);
|
|
for (let i = 0; i < numSlices; i++) {
|
|
const index = [];
|
|
let flattenIndex = 0;
|
|
for (let j2 = 0; j2 < sliceRank; j2++) {
|
|
const dim = indicesData[i * sliceRank + j2];
|
|
flattenIndex += dim * strides[j2];
|
|
index.push(dim);
|
|
}
|
|
if (flattenIndex < 0 || flattenIndex >= paramsSize / sliceSize) {
|
|
throw new Error(`Invalid indices: ${index} does not index into ${paramsShape}`);
|
|
}
|
|
for (let k = 0; k < sliceSize; k++) {
|
|
outBuf.values[i * sliceSize + k] = paramsBuf.get(...paramsBuf.indexToLoc(flattenIndex * sliceSize + k));
|
|
}
|
|
}
|
|
return outBuf;
|
|
}
|
|
function gatherV2Impl(xBuf, indicesBuf, flattenOutputShape) {
|
|
const outBuf = buffer(flattenOutputShape, xBuf.dtype);
|
|
for (let i = 0; i < outBuf.size; ++i) {
|
|
const newLoc = outBuf.indexToLoc(i);
|
|
const originalLoc = newLoc.slice();
|
|
const batchIdx = originalLoc[0];
|
|
const indicesIdx = originalLoc[2];
|
|
const indicesIndex = indicesBuf.locToIndex([batchIdx, indicesIdx]);
|
|
originalLoc[2] = indicesBuf.values[indicesIndex];
|
|
const originalIndex = xBuf.locToIndex(originalLoc);
|
|
if (0 <= originalIndex && originalIndex < xBuf.values.length) {
|
|
outBuf.values[i] = xBuf.values[originalIndex];
|
|
}
|
|
}
|
|
return outBuf;
|
|
}
|
|
const greaterImpl = createSimpleBinaryKernelImpl((a, b) => a > b ? 1 : 0);
|
|
const greater$1 = binaryKernelFunc$1(Greater, greaterImpl, null, "bool");
|
|
const greaterConfig$1 = {
|
|
kernelName: Greater,
|
|
backendName: "cpu",
|
|
kernelFunc: greater$1
|
|
};
|
|
const greaterEqualImpl = createSimpleBinaryKernelImpl((a, b) => a >= b ? 1 : 0);
|
|
const greaterEqual$1 = binaryKernelFunc$1(GreaterEqual, greaterEqualImpl, null, "bool");
|
|
const greaterEqualConfig$1 = {
|
|
kernelName: GreaterEqual,
|
|
backendName: "cpu",
|
|
kernelFunc: greaterEqual$1
|
|
};
|
|
const lessImpl = createSimpleBinaryKernelImpl((a, b) => a < b ? 1 : 0);
|
|
const less$1 = binaryKernelFunc$1(Less, lessImpl, null, "bool");
|
|
const lessConfig$1 = {
|
|
kernelName: Less,
|
|
backendName: "cpu",
|
|
kernelFunc: less$1
|
|
};
|
|
const lessEqualImpl = createSimpleBinaryKernelImpl((a, b) => a <= b ? 1 : 0);
|
|
const lessEqual$1 = binaryKernelFunc$1(LessEqual, lessEqualImpl, null, "bool");
|
|
const lessEqualConfig$1 = {
|
|
kernelName: LessEqual,
|
|
backendName: "cpu",
|
|
kernelFunc: lessEqual$1
|
|
};
|
|
function linSpaceImpl(start, stop, num) {
|
|
const step2 = (stop - start) / (num - 1);
|
|
const values = makeZerosTypedArray(num, "float32");
|
|
values[0] = start;
|
|
for (let i = 1; i < values.length; i++) {
|
|
values[i] = values[i - 1] + step2;
|
|
}
|
|
return values;
|
|
}
|
|
const logImpl = createSimpleUnaryImpl((xi) => Math.log(xi));
|
|
const log$1 = unaryKernelFuncFromImpl(Log, logImpl);
|
|
const logConfig$1 = {
|
|
kernelName: Log,
|
|
backendName: "cpu",
|
|
kernelFunc: log$1
|
|
};
|
|
function maxImpl$1(aVals, reduceSize, outShape, dtype) {
|
|
const vals = getTypedArrayFromDType(dtype, sizeFromShape(outShape));
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let max2 = aVals[offset];
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
const value = aVals[offset + j2];
|
|
if (Number.isNaN(value) || value > max2) {
|
|
max2 = value;
|
|
}
|
|
}
|
|
vals[i] = max2;
|
|
}
|
|
return vals;
|
|
}
|
|
const maximumImpl = createSimpleBinaryKernelImpl(((aValue, bValue) => Math.max(aValue, bValue)));
|
|
const maximum$1 = binaryKernelFunc$1(Maximum, maximumImpl);
|
|
const maximumConfig$1 = {
|
|
kernelName: Maximum,
|
|
backendName: "cpu",
|
|
kernelFunc: maximum$1
|
|
};
|
|
const minimumImpl = createSimpleBinaryKernelImpl(((aValue, bValue) => Math.min(aValue, bValue)));
|
|
const minimum$1 = binaryKernelFunc$1(Minimum, minimumImpl);
|
|
const minimumConfig$1 = {
|
|
kernelName: Minimum,
|
|
backendName: "cpu",
|
|
kernelFunc: minimum$1
|
|
};
|
|
const multiplyImpl = createSimpleBinaryKernelImpl(((aValue, bValue) => aValue * bValue));
|
|
const multiplyComplexImpl = createComplexBinaryKernelImpl(((aReal, aImag, bReal, bImag) => {
|
|
return {
|
|
real: aReal * bReal - aImag * bImag,
|
|
imag: aReal * bImag + aImag * bReal
|
|
};
|
|
}));
|
|
const multiply$1 = binaryKernelFunc$1(Multiply, multiplyImpl, multiplyComplexImpl);
|
|
const multiplyConfig$1 = {
|
|
kernelName: Multiply,
|
|
backendName: "cpu",
|
|
kernelFunc: multiply$1
|
|
};
|
|
function negImpl(xVals, xShape, xDtype) {
|
|
const minusOne = createScalarValue(-1, xDtype);
|
|
return multiplyImpl([], xShape, minusOne, xVals, xDtype);
|
|
}
|
|
function neg$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
assertNotComplex(x, "neg");
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const [res, newShape] = negImpl(xVals, x.shape, x.dtype);
|
|
return backend2.makeTensorInfo(newShape, x.dtype, res);
|
|
}
|
|
const negConfig$1 = {
|
|
kernelName: Neg,
|
|
backendName: "cpu",
|
|
kernelFunc: neg$1
|
|
};
|
|
const notEqualImpl = createSimpleBinaryKernelImpl(((a, b) => a !== b ? 1 : 0));
|
|
const notEqual$1 = binaryKernelFunc$1(NotEqual, notEqualImpl, null, "bool");
|
|
const notEqualConfig$1 = {
|
|
kernelName: NotEqual,
|
|
backendName: "cpu",
|
|
kernelFunc: notEqual$1
|
|
};
|
|
function transposeImpl$1(xVals, xShape, dtype, perm, newShape) {
|
|
const xRank = xShape.length;
|
|
const xSize = sizeFromShape(xShape);
|
|
const xStrides = computeStrides(xShape);
|
|
const newStrides = computeStrides(newShape);
|
|
const result = getTypedArrayFromDType(dtype, sizeFromShape(newShape));
|
|
for (let i = 0; i < xSize; ++i) {
|
|
const loc = indexToLoc(i, xRank, xStrides);
|
|
const newLoc = new Array(loc.length);
|
|
for (let i2 = 0; i2 < newLoc.length; i2++) {
|
|
newLoc[i2] = loc[perm[i2]];
|
|
}
|
|
const newIndex = locToIndex(newLoc, xRank, newStrides);
|
|
result[newIndex] = xVals[i];
|
|
}
|
|
return result;
|
|
}
|
|
function transpose$1(args) {
|
|
const { inputs, attrs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
const { perm } = attrs;
|
|
assertNotComplex(x, "transpose");
|
|
const xRank = x.shape.length;
|
|
const newShape = new Array(xRank);
|
|
for (let i = 0; i < newShape.length; i++) {
|
|
newShape[i] = x.shape[perm[i]];
|
|
}
|
|
const values = backend2.data.get(x.dataId).values;
|
|
const result = transposeImpl$1(values, x.shape, x.dtype, perm, newShape);
|
|
const dataId = backend2.write(result, newShape, x.dtype);
|
|
return { dataId, shape: newShape, dtype: x.dtype };
|
|
}
|
|
const transposeConfig$1 = {
|
|
kernelName: Transpose,
|
|
backendName: "cpu",
|
|
kernelFunc: transpose$1
|
|
};
|
|
function prodImpl(xShape, xDtype, xVals, reductionAxes) {
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(xShape, reductionAxes);
|
|
const outDtype = upcastType(xDtype, "int32");
|
|
const outVals = makeZerosTypedArray(sizeFromShape(outShape), outDtype);
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
for (let i = 0; i < outVals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let prod2 = 1;
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
prod2 *= xVals[offset + j2];
|
|
}
|
|
outVals[i] = prod2;
|
|
}
|
|
return { outVals, outShape, outDtype };
|
|
}
|
|
function prod$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
assertNotComplex(x, "prod");
|
|
const xRank = x.shape.length;
|
|
const axes = parseAxisParam(axis, x.shape);
|
|
const permutation = getAxesPermutation(axes, xRank);
|
|
let reductionAxes = axes;
|
|
let permutedX = x;
|
|
const intermediateTensorInfos = [];
|
|
if (permutation != null) {
|
|
permutedX = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutation } });
|
|
intermediateTensorInfos.push(permutedX);
|
|
reductionAxes = getInnerMostAxes(reductionAxes.length, xRank);
|
|
}
|
|
const xVals = backend2.data.get(permutedX.dataId).values;
|
|
const { outVals, outShape, outDtype } = prodImpl(permutedX.shape, permutedX.dtype, xVals, reductionAxes);
|
|
let resultShape = outShape;
|
|
if (keepDims) {
|
|
resultShape = expandShapeToKeepDim(outShape, axes);
|
|
}
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return backend2.makeTensorInfo(resultShape, outDtype, outVals);
|
|
}
|
|
const prodConfig$1 = {
|
|
kernelName: Prod,
|
|
backendName: "cpu",
|
|
kernelFunc: prod$1
|
|
};
|
|
function validateIndices(indices, indicesShape, numParams) {
|
|
indices.forEach((index, i) => {
|
|
if (index < 0 || index >= numParams) {
|
|
const locString = indexToLoc(i, indicesShape.length, computeStrides(indicesShape)).join(",");
|
|
throw new Error(`indices[${locString}] = ${index} is not in [0, ${numParams})`);
|
|
}
|
|
});
|
|
}
|
|
function validateSplits(paramsNestedSplits, numParamsDenseValues) {
|
|
for (let dim = 0; dim < paramsNestedSplits.length; ++dim) {
|
|
const splits = paramsNestedSplits[dim];
|
|
const lastSplit = dim === paramsNestedSplits.length - 1 ? numParamsDenseValues : paramsNestedSplits[dim + 1].length;
|
|
if (splits.length === 0) {
|
|
throw new Error("Ragged splits may not be empty");
|
|
}
|
|
if (splits[0] < 0) {
|
|
throw new Error("Ragged splits must be non-negative");
|
|
}
|
|
if (splits[splits.length - 1] > lastSplit) {
|
|
throw new Error("Ragged splits must not point past values");
|
|
}
|
|
for (let i = 1; i < splits.length; ++i) {
|
|
if (splits[i - 1] > splits[i]) {
|
|
throw new Error("Ragged splits must be sorted in ascending order");
|
|
}
|
|
}
|
|
}
|
|
}
|
|
function makeSplits(indices, indicesShape, paramsNestedSplits, numParamsDenseValues) {
|
|
const valueSlices = [];
|
|
let numValues = 0;
|
|
const numSplits = indicesShape.length - 1 + paramsNestedSplits.length;
|
|
const outSplits = new Array(numSplits).fill(null).map(() => [0]);
|
|
validateSplits(paramsNestedSplits, numParamsDenseValues);
|
|
let nrows = 1;
|
|
for (let dim = 0; dim < indicesShape.length - 1; ++dim) {
|
|
nrows *= indicesShape[dim];
|
|
const rowLength = indicesShape[dim + 1];
|
|
for (let i = 1; i < nrows + 1; ++i) {
|
|
outSplits[dim].push(i * rowLength);
|
|
}
|
|
}
|
|
for (let i = 0; i < indices.length; ++i) {
|
|
let start = indices[i];
|
|
let limit = indices[i] + 1;
|
|
for (let dim = 0; dim < paramsNestedSplits.length; ++dim) {
|
|
const splits = paramsNestedSplits[dim];
|
|
const outDim = dim + indicesShape.length - 1;
|
|
if (outDim >= 0) {
|
|
const outSplitsOutDim = outSplits[outDim];
|
|
const delta = outSplitsOutDim[outSplitsOutDim.length - 1] - splits[start];
|
|
for (let j2 = start; j2 < limit; ++j2) {
|
|
outSplits[outDim].push(splits[j2 + 1] + delta);
|
|
}
|
|
}
|
|
start = splits[start];
|
|
limit = splits[limit];
|
|
}
|
|
if (limit !== start) {
|
|
valueSlices.push([start, limit]);
|
|
numValues += limit - start;
|
|
}
|
|
}
|
|
return { outSplits, valueSlices, numValues };
|
|
}
|
|
function getSplits(outSplits) {
|
|
const splitsOut = [];
|
|
for (let i = 0; i < outSplits.length; ++i) {
|
|
const numSplits = outSplits[i].length;
|
|
const splits = getArrayFromDType("int32", numSplits);
|
|
splitsOut.push(splits);
|
|
outSplits[i].forEach((value, j2) => splits[j2] = value);
|
|
}
|
|
return splitsOut;
|
|
}
|
|
function computeFlatOuterDims(orig, numOutDims) {
|
|
const outDims = orig.slice(0, numOutDims);
|
|
while (outDims.length < numOutDims) {
|
|
outDims.push(1);
|
|
}
|
|
for (let inDim = numOutDims; inDim < orig.length; inDim++) {
|
|
outDims[numOutDims - 1] *= orig[inDim];
|
|
}
|
|
return outDims;
|
|
}
|
|
function writeValueSlices(paramsDenseValues, paramsDenseValuesShape, valueSlices, valueSize, values, valuesShape) {
|
|
const denseM = computeFlatOuterDims(paramsDenseValuesShape, 2)[1];
|
|
const valuesM = computeFlatOuterDims(valuesShape, 2)[1];
|
|
let outPos = 0;
|
|
for (const slice2 of valueSlices) {
|
|
for (let i = slice2[0]; i < slice2[1]; ++i) {
|
|
for (let j2 = 0; j2 < valueSize; ++j2) {
|
|
values[outPos * valuesM + j2] = paramsDenseValues[i * denseM + j2];
|
|
}
|
|
++outPos;
|
|
}
|
|
}
|
|
}
|
|
function getValues(paramsDenseValues, paramsDenseValuesShape, paramsDenseValuesDType, valueSlices, numValues) {
|
|
const valuesShape = paramsDenseValuesShape.slice();
|
|
valuesShape[0] = numValues;
|
|
const valuesOut = getArrayFromDType(paramsDenseValuesDType, sizeFromShape(valuesShape));
|
|
const numElements = paramsDenseValues.length;
|
|
const valueSize = numElements === 0 ? 0 : numElements / paramsDenseValuesShape[0];
|
|
writeValueSlices(paramsDenseValues, paramsDenseValuesShape, valueSlices, valueSize, valuesOut, valuesShape);
|
|
return [valuesOut, valuesShape];
|
|
}
|
|
function raggedGatherImpl(paramsNestedSplits, paramsNestedSplitsShapes, paramsDenseValues, paramsDenseValuesShape, paramsDenseValuesDType, indices, indicesShape, outputRaggedRank) {
|
|
if (paramsNestedSplits.length === 0) {
|
|
throw new Error("paramsNestedSplits must be non empty");
|
|
}
|
|
if (paramsNestedSplitsShapes[0].length === 0) {
|
|
throw new Error("Split tensors must not be scalars");
|
|
}
|
|
const numParams = paramsNestedSplitsShapes[0][0] - 1;
|
|
validateIndices(indices, indicesShape, numParams);
|
|
if (paramsDenseValuesShape.length === 0) {
|
|
throw new Error("params.rank must be nonzero");
|
|
}
|
|
const numParamsDenseValues = paramsDenseValuesShape[0];
|
|
const { outSplits, valueSlices, numValues } = makeSplits(indices, indicesShape, paramsNestedSplits, numParamsDenseValues);
|
|
const outputNestedSplits = getSplits(outSplits);
|
|
const outputDenseValues = getValues(paramsDenseValues, paramsDenseValuesShape, paramsDenseValuesDType, valueSlices, numValues);
|
|
return [outputNestedSplits, outputDenseValues[0], outputDenseValues[1]];
|
|
}
|
|
const INT32_MAX = 2147483647;
|
|
function raggedRangeImpl(starts, startsShape, startsDType, limits, limitsShape, deltas, deltasShape) {
|
|
if (startsShape.length > 1) {
|
|
throw new Error("starts must be a scalar or vector");
|
|
}
|
|
if (limitsShape.length > 1) {
|
|
throw new Error("limits must be a scalar or vector");
|
|
}
|
|
if (deltasShape.length > 1) {
|
|
throw new Error("deltas must be a scalar or vector");
|
|
}
|
|
const broadcastStarts = startsShape.length === 0;
|
|
const broadcastLimits = limitsShape.length === 0;
|
|
const broadcastDeltas = deltasShape.length === 0;
|
|
const inSizes = [];
|
|
if (!broadcastStarts) {
|
|
inSizes.push(startsShape[0]);
|
|
}
|
|
if (!broadcastLimits) {
|
|
inSizes.push(limitsShape[0]);
|
|
}
|
|
if (!broadcastDeltas) {
|
|
inSizes.push(deltasShape[0]);
|
|
}
|
|
for (let i = 1; i < inSizes.length; ++i) {
|
|
if (inSizes[i] !== inSizes[i - 1]) {
|
|
throw new Error("starts, limits, and deltas must have the same shape");
|
|
}
|
|
}
|
|
const nRows = inSizes.length === 0 ? 1 : inSizes[0];
|
|
const rtNestedSplits = getArrayFromDType("int32", nRows + 1);
|
|
rtNestedSplits[0] = 0;
|
|
for (let row = 0; row < nRows; ++row) {
|
|
const start = broadcastStarts ? starts[0] : starts[row];
|
|
const limit = broadcastLimits ? limits[0] : limits[row];
|
|
const delta = broadcastDeltas ? deltas[0] : deltas[row];
|
|
if (delta === 0) {
|
|
throw new Error("Requires delta != 0");
|
|
}
|
|
let size;
|
|
if (delta > 0 && limit < start || delta < 0 && limit > start) {
|
|
size = 0;
|
|
} else {
|
|
size = Math.ceil(Math.abs((limit - start) / delta));
|
|
if (size > INT32_MAX) {
|
|
throw new Error(`Requires ((limit - start) / delta) <= ${INT32_MAX}`);
|
|
}
|
|
}
|
|
rtNestedSplits[row + 1] = rtNestedSplits[row] + size;
|
|
}
|
|
const nVals = rtNestedSplits[nRows];
|
|
const rtDenseValues = getArrayFromDType(startsDType, nVals);
|
|
let valueIndex = 0;
|
|
for (let row = 0; row < nRows; ++row) {
|
|
const rowSize = rtNestedSplits[row + 1] - rtNestedSplits[row];
|
|
let value = broadcastStarts ? starts[0] : starts[row];
|
|
const delta = broadcastDeltas ? deltas[0] : deltas[row];
|
|
for (let i = 0; i < rowSize; ++i) {
|
|
rtDenseValues[valueIndex++] = value;
|
|
value += delta;
|
|
}
|
|
}
|
|
return [rtNestedSplits, rtDenseValues];
|
|
}
|
|
var RowPartitionType = RowPartitionType$1;
|
|
class RaggedTensorToTensorOp {
|
|
constructor(shape, shapeShape, values, valuesShape, valuesDType, defaultValue, defaultValueShape, rowPartitionValues, rowPartitionValuesShapes, rowPartitionTypeStrings) {
|
|
this.shape = shape;
|
|
this.shapeShape = shapeShape;
|
|
this.values = values;
|
|
this.valuesShape = valuesShape;
|
|
this.valuesDType = valuesDType;
|
|
this.defaultValue = defaultValue;
|
|
this.defaultValueShape = defaultValueShape;
|
|
this.rowPartitionValues = rowPartitionValues;
|
|
this.rowPartitionValuesShapes = rowPartitionValuesShapes;
|
|
this.rowPartitionTypes = getRowPartitionTypesHelper(rowPartitionTypeStrings);
|
|
this.raggedRank = getRaggedRank(this.rowPartitionTypes);
|
|
}
|
|
getRowPartitionTypeByDimension(dimension) {
|
|
if (this.rowPartitionTypes[0] === RowPartitionType.FIRST_DIM_SIZE) {
|
|
return this.rowPartitionTypes[dimension + 1];
|
|
} else {
|
|
return this.rowPartitionTypes[dimension];
|
|
}
|
|
}
|
|
// Returns the relationship between dimension and dimension + 1.
|
|
getRowPartitionTensor(dimension) {
|
|
if (this.rowPartitionTypes[0] === RowPartitionType.FIRST_DIM_SIZE) {
|
|
return this.rowPartitionValues[dimension + 1];
|
|
} else {
|
|
return this.rowPartitionValues[dimension];
|
|
}
|
|
}
|
|
getMaxWidth(dimension) {
|
|
const rowPartitionTensor = this.getRowPartitionTensor(dimension - 1);
|
|
switch (this.getRowPartitionTypeByDimension(dimension - 1)) {
|
|
case RowPartitionType.VALUE_ROWIDS:
|
|
return RaggedTensorToTensorOp.getMaxWidthValueRowID(rowPartitionTensor);
|
|
case RowPartitionType.ROW_SPLITS:
|
|
return RaggedTensorToTensorOp.getMaxWidthRowSplit(rowPartitionTensor);
|
|
default:
|
|
throw new Error(`Cannot handle partition type ${RowPartitionType[this.getRowPartitionTypeByDimension(dimension - 1)]}`);
|
|
}
|
|
}
|
|
static getMaxWidthRowSplit(rowSplit) {
|
|
const tensorLength = rowSplit.length;
|
|
if (tensorLength === 0 || tensorLength === 1) {
|
|
return 0;
|
|
}
|
|
let maxWidth = 0;
|
|
for (let i = 0; i < tensorLength - 1; ++i) {
|
|
const currentWidth = rowSplit[i + 1] - rowSplit[i];
|
|
if (currentWidth > maxWidth) {
|
|
maxWidth = currentWidth;
|
|
}
|
|
}
|
|
return maxWidth;
|
|
}
|
|
static getMaxWidthValueRowID(valueRowIds) {
|
|
const indexLength = valueRowIds.length;
|
|
if (indexLength === 0) {
|
|
return 0;
|
|
}
|
|
let firstEqualIndex = 0;
|
|
let firstEqualIndexValue = valueRowIds[0];
|
|
let maxWidth = 0;
|
|
for (let i = 1; i < indexLength; ++i) {
|
|
const value = valueRowIds[i];
|
|
if (value !== firstEqualIndexValue) {
|
|
firstEqualIndexValue = value;
|
|
maxWidth = Math.max(i - firstEqualIndex, maxWidth);
|
|
firstEqualIndex = i;
|
|
}
|
|
}
|
|
return Math.max(indexLength - firstEqualIndex, maxWidth);
|
|
}
|
|
tensorShapeFromTensor(t2, tShape, isPartial = true) {
|
|
if (tShape.length === 0) {
|
|
if (t2[0] === -1) {
|
|
return [];
|
|
}
|
|
throw new Error(`The only valid scalar shape tensor is the fully unknown shape specified as -1.`);
|
|
}
|
|
return makeShape(t2, isPartial);
|
|
}
|
|
calculateOutputSize(firstDim) {
|
|
const valueShape = this.valuesShape;
|
|
const defaultValueShape = this.defaultValueShape;
|
|
validateDefaultValueShape(defaultValueShape, valueShape);
|
|
const shape = this.tensorShapeFromTensor(this.shape, this.shapeShape);
|
|
const outputShape = combineRaggedTensorToTensorShapes(this.raggedRank, shape, valueShape);
|
|
const result = outputShape;
|
|
if (result[0] < 0) {
|
|
result[0] = firstDim;
|
|
}
|
|
for (let i = 1; i <= this.raggedRank; ++i) {
|
|
if (result[i] < 0) {
|
|
result[i] = this.getMaxWidth(i);
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
/**
|
|
* The outputIndex represents the index in the output tensor
|
|
* where the first element of a particular dimension would be written.
|
|
* If it is -1, it indicates that the index is out of scope.
|
|
* Example, given firstDimension = 10, firstDimensionOutput = 6,
|
|
* and outputIndexMultiplier = 100:
|
|
* result = [0 100 200 300 400 500 -1 -1 -1 -1]
|
|
* If firstDimensionOutput = 11 instead, then:
|
|
* result = [0 100 200 300 400 500 600 700 800 900]
|
|
*/
|
|
calculateFirstParentOutputIndex(firstDimension, outputIndexMultiplier, firstDimensionOutput) {
|
|
const minDimension = Math.min(firstDimension, firstDimensionOutput);
|
|
const result = [];
|
|
let currentOutputIndex = 0;
|
|
for (let i = 0; i < minDimension; ++i, currentOutputIndex += outputIndexMultiplier) {
|
|
result.push(currentOutputIndex);
|
|
}
|
|
for (let i = minDimension; i < firstDimension; ++i) {
|
|
result.push(-1);
|
|
}
|
|
assert(result.length === firstDimension, () => "Final length of result must be equal to firstDimension.");
|
|
return result;
|
|
}
|
|
calculateOutputIndexRowSplit(rowSplit, parentOutputIndex, outputIndexMultiplier, outputSize) {
|
|
const rowSplitSize = rowSplit.length;
|
|
const result = [];
|
|
for (let i = 0; i < rowSplitSize - 1; ++i) {
|
|
const rowLength = rowSplit[i + 1] - rowSplit[i];
|
|
let realLength = Math.min(outputSize, rowLength);
|
|
let parentOutputIndexCurrent = parentOutputIndex[i];
|
|
if (parentOutputIndexCurrent === -1) {
|
|
realLength = 0;
|
|
}
|
|
for (let j2 = 0; j2 < realLength; ++j2) {
|
|
result.push(parentOutputIndexCurrent);
|
|
parentOutputIndexCurrent += outputIndexMultiplier;
|
|
}
|
|
for (let j2 = 0; j2 < rowLength - realLength; ++j2) {
|
|
result.push(-1);
|
|
}
|
|
}
|
|
if (rowSplitSize > 0 && result.length !== rowSplit[rowSplitSize - 1]) {
|
|
throw new Error("Invalid row split size.");
|
|
}
|
|
return result;
|
|
}
|
|
// Calculate the output index of the first element of a list.
|
|
// The parentOutputIndex is the same computation for the previous list.
|
|
// -1 indicates an element or list that is out of range.
|
|
// The outputIndexMultiplier is the number of output indices one moves
|
|
// forward for each column.
|
|
// E.g., given:
|
|
// valueRowIds:[0 1 2 2 2 3 5 5 6]
|
|
// parentOutputIndex:[1000 1100 2000 2100 -1 3000 4000]
|
|
// outputIndexMultiplier: 10
|
|
// outputSize: 2
|
|
// You get:
|
|
// result = [1000 1100 2000 2010 -1 2100 -1 -1 3000]
|
|
// result[0] = parentOutputIndex[valueRowIds[0]]
|
|
// result[1] = parentOutputIndex[valueRowIds[1]]
|
|
// result[2] = parentOutputIndex[valueRowIds[2]]
|
|
// result[3] = parentOutputIndex[valueRowIds[2] + 10]
|
|
// result[4] = -1 because it is the third element the size is 2.
|
|
// result[5] = parentOutputIndex[valueRowIds[3]]
|
|
// result[6] = -1 because parentOutputIndex[valueRowIds[6]] == -1
|
|
// result[7] = -1 because parentOutputIndex[valueRowIds[6]] == -1
|
|
// result[8] = parentOutputIndex[valueRowIds[7]]
|
|
calculateOutputIndexValueRowID(valueRowIds, parentOutputIndex, outputIndexMultiplier, outputSize) {
|
|
const indexSize = valueRowIds.length;
|
|
const result = [];
|
|
if (indexSize === 0) {
|
|
return [];
|
|
}
|
|
let currentOutputColumn = 0;
|
|
let currentValueRowId = valueRowIds[0];
|
|
if (currentValueRowId >= parentOutputIndex.length) {
|
|
throw new Error(`Got currentValueRowId=${currentValueRowId}, which is not less than ${parentOutputIndex.length}`);
|
|
}
|
|
let currentOutputIndex = parentOutputIndex[currentValueRowId];
|
|
result.push(currentOutputIndex);
|
|
for (let i = 1; i < indexSize; ++i) {
|
|
const nextValueRowId = valueRowIds[i];
|
|
if (nextValueRowId === currentValueRowId) {
|
|
if (currentOutputIndex >= 0) {
|
|
++currentOutputColumn;
|
|
if (currentOutputColumn < outputSize) {
|
|
currentOutputIndex += outputIndexMultiplier;
|
|
} else {
|
|
currentOutputIndex = -1;
|
|
}
|
|
}
|
|
} else {
|
|
currentOutputColumn = 0;
|
|
currentValueRowId = nextValueRowId;
|
|
if (nextValueRowId >= parentOutputIndex.length) {
|
|
throw new Error(`Got nextValueRowId=${nextValueRowId} which is not less than ${parentOutputIndex.length}`);
|
|
}
|
|
currentOutputIndex = parentOutputIndex[nextValueRowId];
|
|
}
|
|
result.push(currentOutputIndex);
|
|
}
|
|
if (result.length !== valueRowIds.length) {
|
|
throw new Error("Invalid row ids.");
|
|
}
|
|
return result;
|
|
}
|
|
calculateOutputIndex(dimension, parentOutputIndex, outputIndexMultiplier, outputSize) {
|
|
const rowPartitionTensor = this.getRowPartitionTensor(dimension);
|
|
const partitionType = this.getRowPartitionTypeByDimension(dimension);
|
|
switch (partitionType) {
|
|
case RowPartitionType.VALUE_ROWIDS:
|
|
return this.calculateOutputIndexValueRowID(rowPartitionTensor, parentOutputIndex, outputIndexMultiplier, outputSize);
|
|
case RowPartitionType.ROW_SPLITS:
|
|
if (rowPartitionTensor.length - 1 > parentOutputIndex.length) {
|
|
throw new Error(`Row partition size is greater than output size: ${rowPartitionTensor.length - 1} > ${parentOutputIndex.length}`);
|
|
}
|
|
return this.calculateOutputIndexRowSplit(rowPartitionTensor, parentOutputIndex, outputIndexMultiplier, outputSize);
|
|
default:
|
|
throw new Error(`Unsupported partition type: ${RowPartitionType[partitionType]}`);
|
|
}
|
|
}
|
|
getFirstDimensionSize() {
|
|
const firstPartitionTensor = this.rowPartitionValues[0];
|
|
if (this.rowPartitionTypes.length === 0) {
|
|
throw new Error("No row_partition_types given.");
|
|
}
|
|
const firstPartitionType = this.rowPartitionTypes[0];
|
|
switch (firstPartitionType) {
|
|
case RowPartitionType.FIRST_DIM_SIZE:
|
|
return firstPartitionTensor[0];
|
|
case RowPartitionType.VALUE_ROWIDS:
|
|
throw new Error("Cannot handle VALUE_ROWIDS in first dimension.");
|
|
case RowPartitionType.ROW_SPLITS:
|
|
return this.rowPartitionValuesShapes[0][0] - 1;
|
|
default:
|
|
throw new Error(`Cannot handle type ${RowPartitionType[firstPartitionType]}`);
|
|
}
|
|
}
|
|
compute() {
|
|
const firstPartitionTensor = this.rowPartitionValues[0];
|
|
if (firstPartitionTensor.length <= 0) {
|
|
throw new Error("Invalid first partition input. Tensor requires at least one element.");
|
|
}
|
|
const firstDimension = this.getFirstDimensionSize();
|
|
const outputSize = this.calculateOutputSize(firstDimension);
|
|
const multiplier = new Array(this.raggedRank + 1);
|
|
multiplier[multiplier.length - 1] = 1;
|
|
for (let i = multiplier.length - 2; i >= 0; --i) {
|
|
multiplier[i] = multiplier[i + 1] * outputSize[i + 1];
|
|
}
|
|
const outputShape = makeShape(outputSize, false);
|
|
const outputTensor = getArrayFromDType(this.valuesDType, sizeFromShape(outputShape));
|
|
const fullSize = multiplier[0] * outputSize[0];
|
|
if (fullSize > 0) {
|
|
let outputIndex = this.calculateFirstParentOutputIndex(firstDimension, multiplier[0], outputSize[0]);
|
|
for (let i = 1; i <= this.raggedRank; ++i) {
|
|
const newOutputIndex = this.calculateOutputIndex(i - 1, outputIndex, multiplier[i], outputSize[i]);
|
|
outputIndex = newOutputIndex;
|
|
}
|
|
this.setOutput(this.raggedRank, outputIndex, outputTensor, outputShape);
|
|
}
|
|
return [outputShape, outputTensor];
|
|
}
|
|
setOutput(raggedRank, outputIndex, outputTensor, outputShape) {
|
|
if (outputTensor.length === 0) {
|
|
return;
|
|
}
|
|
const valuesBase = this.values;
|
|
const outputBase = outputTensor;
|
|
let elementShape = outputShape.slice();
|
|
elementShape = elementShape.slice(raggedRank + 1);
|
|
const valueElementSize = sizeFromShape(elementShape);
|
|
const outputIndexSize = outputIndex.length;
|
|
let defaultValue = this.defaultValue;
|
|
if (defaultValue.length !== valueElementSize && defaultValue.length !== 1) {
|
|
const srcShape = this.defaultValueShape;
|
|
tidy(() => {
|
|
const defaultValueTensor = reshape$2(defaultValue, srcShape);
|
|
const bCastDefault = broadcastTo(defaultValueTensor, elementShape);
|
|
defaultValue = bCastDefault.dataSync();
|
|
});
|
|
}
|
|
let srcStart = 0;
|
|
let dstStart = 0;
|
|
let dstEnd = 0;
|
|
for (let srcI = 0; srcI <= outputIndexSize; ++srcI) {
|
|
let dstI = srcI < outputIndexSize ? outputIndex[srcI] : -1;
|
|
if (dstI === dstEnd) {
|
|
++dstEnd;
|
|
continue;
|
|
}
|
|
if (dstStart < dstEnd) {
|
|
const src = valuesBase.subarray(srcStart * valueElementSize);
|
|
const dst = outputBase.subarray(dstStart * valueElementSize);
|
|
const nVals = (dstEnd - dstStart) * valueElementSize;
|
|
copyArray(dst, src, nVals);
|
|
}
|
|
if (srcI >= outputIndexSize) {
|
|
const outputSize = outputTensor.length;
|
|
dstI = Math.floor(outputSize / valueElementSize);
|
|
}
|
|
if (dstI > dstEnd) {
|
|
if (this.defaultValue.length === 1) {
|
|
outputBase.subarray(dstEnd * valueElementSize, dstI * valueElementSize).fill(this.defaultValue[0]);
|
|
dstEnd = dstI;
|
|
} else {
|
|
while (dstI > dstEnd) {
|
|
const dst = outputBase.slice(dstEnd * valueElementSize);
|
|
copyArray(dst, defaultValue, valueElementSize);
|
|
++dstEnd;
|
|
}
|
|
}
|
|
}
|
|
if (dstI < 0) {
|
|
srcStart = srcI + 1;
|
|
dstStart = dstEnd;
|
|
} else {
|
|
srcStart = srcI;
|
|
dstStart = dstEnd;
|
|
dstEnd = dstStart + 1;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
function copyArray(dst, src, size) {
|
|
for (let i = 0; i < size; i++) {
|
|
dst[i] = src[i];
|
|
}
|
|
}
|
|
function makeShape(shape, isPartial) {
|
|
const out = [];
|
|
for (let dim of shape) {
|
|
if (dim < 0) {
|
|
if (!isPartial) {
|
|
throw new Error(`Dimension ${dim} must be >= 0`);
|
|
}
|
|
if (dim < -1) {
|
|
throw new Error(`Dimension ${dim} must be >= -1`);
|
|
}
|
|
dim = -1;
|
|
}
|
|
out.push(dim);
|
|
}
|
|
return out;
|
|
}
|
|
function raggedTensorToTensorImpl(shape, shapesShape, values, valuesShape, valuesDType, defaultValue, defaultValueShape, rowPartitionValues, rowPartitionValuesShapes, rowPartitionTypes) {
|
|
return new RaggedTensorToTensorOp(shape, shapesShape, values, valuesShape, valuesDType, defaultValue, defaultValueShape, rowPartitionValues, rowPartitionValuesShapes, rowPartitionTypes).compute();
|
|
}
|
|
function rangeImpl(start, stop, step2, dtype) {
|
|
const sameStartStop = start === stop;
|
|
const increasingRangeNegativeStep = start < stop && step2 < 0;
|
|
const decreasingRangePositiveStep = stop < start && step2 > 1;
|
|
if (sameStartStop || increasingRangeNegativeStep || decreasingRangePositiveStep) {
|
|
return makeZerosTypedArray(0, dtype);
|
|
}
|
|
const numElements = Math.abs(Math.ceil((stop - start) / step2));
|
|
const values = makeZerosTypedArray(numElements, dtype);
|
|
if (stop < start && step2 === 1) {
|
|
step2 = -1;
|
|
}
|
|
values[0] = start;
|
|
for (let i = 1; i < values.length; i++) {
|
|
values[i] = values[i - 1] + step2;
|
|
}
|
|
return values;
|
|
}
|
|
const rsqrtImpl = createSimpleUnaryImpl((xi) => 1 / Math.sqrt(xi));
|
|
const rsqrt$1 = unaryKernelFuncFromImpl(Rsqrt, rsqrtImpl);
|
|
const rsqrtConfig$1 = {
|
|
kernelName: Rsqrt,
|
|
backendName: "cpu",
|
|
kernelFunc: rsqrt$1
|
|
};
|
|
function scatterImpl(indices, updates, shape, outputSize, sliceSize, numUpdates, sliceRank, strides, defaultValue, sumDupeIndices) {
|
|
const flattenShape = [outputSize / sliceSize, sliceSize];
|
|
const indicesData = indices.values;
|
|
const updatesData = updates.values;
|
|
if (outputSize === 0) {
|
|
return buffer(shape, updates.dtype);
|
|
}
|
|
const outBuf = defaultValue instanceof TensorBuffer ? defaultValue : buffer(flattenShape, updates.dtype);
|
|
if (typeof defaultValue === "string") {
|
|
outBuf.values.fill(defaultValue);
|
|
} else if (typeof defaultValue === "number") {
|
|
outBuf.values.fill(defaultValue);
|
|
} else if (typeof defaultValue === "boolean") {
|
|
outBuf.values.fill(+defaultValue);
|
|
}
|
|
for (let i = 0; i < numUpdates; i++) {
|
|
const index = [];
|
|
let flattenIndex = 0;
|
|
for (let j2 = 0; j2 < sliceRank; j2++) {
|
|
const dim = indicesData[i * sliceRank + j2];
|
|
index.push(dim);
|
|
flattenIndex += dim * strides[j2];
|
|
}
|
|
if (flattenIndex < 0 || flattenIndex >= outputSize / sliceSize) {
|
|
throw new Error(`Invalid indices: ${index} does not index into ${shape}`);
|
|
}
|
|
for (let k = 0; k < sliceSize; k++) {
|
|
if (sumDupeIndices) {
|
|
outBuf.values[flattenIndex * sliceSize + k] += updatesData[i * sliceSize + k];
|
|
} else {
|
|
outBuf.values[flattenIndex * sliceSize + k] = updates.rank === 0 ? updatesData[0] : updatesData[i * sliceSize + k];
|
|
}
|
|
}
|
|
}
|
|
return outBuf;
|
|
}
|
|
const sigmoidImpl = createSimpleUnaryImpl((xi) => 1 / (1 + Math.exp(-xi)));
|
|
const sigmoid$1 = unaryKernelFunc$1(Sigmoid, (xi) => 1 / (1 + Math.exp(-xi)));
|
|
const sigmoidConfig$1 = {
|
|
kernelName: Sigmoid,
|
|
backendName: "cpu",
|
|
kernelFunc: sigmoid$1
|
|
};
|
|
function sliceImpl(vals, begin, size, shape, dtype) {
|
|
const isContinous = isSliceContinous(shape, begin, size);
|
|
const length = sizeFromShape(size);
|
|
const xStrides = computeStrides(shape);
|
|
if (isContinous) {
|
|
const flatOffset = computeFlatOffset(begin, xStrides);
|
|
if (dtype === "string") {
|
|
return vals.slice(flatOffset, flatOffset + length);
|
|
}
|
|
return vals.subarray(flatOffset, flatOffset + length);
|
|
}
|
|
const decodedData = dtype === "string" ? fromUint8ToStringArray(vals) : vals;
|
|
const inBuf = buffer(shape, dtype, decodedData);
|
|
const outBuf = buffer(size, dtype);
|
|
for (let i = 0; i < outBuf.size; ++i) {
|
|
const outLoc = outBuf.indexToLoc(i);
|
|
const inLoc = outLoc.map((idx, j2) => idx + begin[j2]);
|
|
outBuf.set(inBuf.get(...inLoc), ...outLoc);
|
|
}
|
|
if (dtype === "string") {
|
|
return fromStringArrayToUint8(outBuf.values);
|
|
}
|
|
return outBuf.values;
|
|
}
|
|
function slice$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { begin, size } = attrs;
|
|
assertNotComplex(x, "slice");
|
|
const [$begin, $size] = parseSliceParams(x, begin, size);
|
|
assertParamsValid(x, $begin, $size);
|
|
const vals = backend2.data.get(x.dataId).values;
|
|
const outVals = sliceImpl(vals, $begin, $size, x.shape, x.dtype);
|
|
return backend2.makeTensorInfo($size, x.dtype, outVals);
|
|
}
|
|
const sliceConfig$1 = {
|
|
kernelName: Slice,
|
|
backendName: "cpu",
|
|
kernelFunc: slice$1
|
|
};
|
|
function sparseFillEmptyRowsImpl(indices, indicesShape, indicesDType, values, valuesDType, denseShape, defaultValue) {
|
|
const indicesCount = indicesShape[0];
|
|
const denseRows = denseShape[0];
|
|
const emptyRowIndicator = new Array(denseRows);
|
|
const reverseIndexMap = new Array(indicesCount);
|
|
const rank = indicesShape[1];
|
|
if (denseRows === 0) {
|
|
if (indicesCount !== 0) {
|
|
throw new Error(getSparseFillEmptyRowsIndicesDenseShapeMismatch(indicesCount));
|
|
}
|
|
const outputIndices = getArrayFromDType(indicesDType, 0);
|
|
const outputValues = getArrayFromDType(valuesDType, 0);
|
|
return [
|
|
outputIndices,
|
|
[0, rank],
|
|
outputValues,
|
|
emptyRowIndicator,
|
|
reverseIndexMap
|
|
];
|
|
}
|
|
let rowsAreOrdered = true;
|
|
let lastIndicesRow = 0;
|
|
const csrOffset = new Array(denseRows).fill(0);
|
|
for (let i = 0; i < indicesCount; ++i) {
|
|
const row = indices[i * rank];
|
|
if (row < 0) {
|
|
throw new Error(getSparseFillEmptyRowsNegativeIndexErrorMessage(i, row));
|
|
}
|
|
if (row >= denseRows) {
|
|
throw new Error(getSparseFillEmptyRowsOutOfRangeIndexErrorMessage(i, row, denseRows));
|
|
}
|
|
++csrOffset[row];
|
|
rowsAreOrdered = rowsAreOrdered && row >= lastIndicesRow;
|
|
lastIndicesRow = row;
|
|
}
|
|
let allRowsFull = true;
|
|
for (let row = 0; row < denseRows; ++row) {
|
|
const rowEmpty = csrOffset[row] === 0;
|
|
emptyRowIndicator[row] = rowEmpty;
|
|
allRowsFull = allRowsFull && !rowEmpty;
|
|
csrOffset[row] = Math.max(csrOffset[row], 1);
|
|
if (row > 0) {
|
|
csrOffset[row] += csrOffset[row - 1];
|
|
}
|
|
}
|
|
if (allRowsFull && rowsAreOrdered) {
|
|
const outputIndices = indices;
|
|
const outputValues = values;
|
|
for (let i = 0; i < indicesCount; ++i) {
|
|
reverseIndexMap[i] = i;
|
|
}
|
|
return [
|
|
outputIndices,
|
|
[indicesCount, rank],
|
|
outputValues,
|
|
emptyRowIndicator,
|
|
reverseIndexMap
|
|
];
|
|
} else {
|
|
const fullIndicesCount = csrOffset[denseRows - 1];
|
|
const outputIndices = getArrayFromDType(indicesDType, fullIndicesCount * rank);
|
|
const outputValues = getArrayFromDType(valuesDType, fullIndicesCount);
|
|
const filledCount = new Array(denseRows).fill(0);
|
|
for (let i = 0; i < indicesCount; ++i) {
|
|
const row = indices[i * rank];
|
|
const offset = filledCount[row];
|
|
const outputI = (row === 0 ? 0 : csrOffset[row - 1]) + offset;
|
|
filledCount[row]++;
|
|
for (let j2 = 0; j2 < rank; ++j2) {
|
|
outputIndices[outputI * rank + j2] = indices[i * rank + j2];
|
|
}
|
|
outputValues[outputI] = values[i];
|
|
reverseIndexMap[i] = outputI;
|
|
}
|
|
for (let row = 0; row < denseRows; ++row) {
|
|
const rowCount = filledCount[row];
|
|
if (rowCount === 0) {
|
|
const startingIndex = row === 0 ? 0 : csrOffset[row - 1];
|
|
outputIndices[startingIndex * rank + 0] = row;
|
|
for (let col = 1; col < rank; ++col) {
|
|
outputIndices[startingIndex * rank + col] = 0;
|
|
}
|
|
outputValues[startingIndex] = defaultValue;
|
|
}
|
|
}
|
|
return [
|
|
outputIndices,
|
|
[fullIndicesCount, rank],
|
|
outputValues,
|
|
emptyRowIndicator,
|
|
reverseIndexMap
|
|
];
|
|
}
|
|
}
|
|
function sparseReshapeImpl(inputIndices, inputIndicesShape, inputDType, inputShape, targetShape) {
|
|
const denseSize = sizeFromShape(inputShape);
|
|
const nnz = inputIndicesShape[0];
|
|
const outputRank = targetShape.length;
|
|
const outputShape = [];
|
|
let product = 1;
|
|
let unknownIndex = -1;
|
|
for (let d = 0; d < outputRank; ++d) {
|
|
const size = targetShape[d];
|
|
if (size === -1) {
|
|
if (unknownIndex !== -1) {
|
|
throw new Error(getSparseReshapeMultipleNegativeOneOutputDimErrorMessage(unknownIndex, d));
|
|
}
|
|
unknownIndex = d;
|
|
outputShape.push(1);
|
|
} else {
|
|
if (size < 0) {
|
|
throw new Error(getSparseReshapeNegativeOutputDimErrorMessage(d, size));
|
|
}
|
|
product *= size;
|
|
outputShape.push(size);
|
|
}
|
|
}
|
|
if (unknownIndex !== -1) {
|
|
if (product <= 0) {
|
|
throw new Error(getSparseReshapeEmptyTensorZeroOutputDimErrorMessage());
|
|
}
|
|
const missing = Math.trunc(denseSize / product);
|
|
if (product * missing !== denseSize) {
|
|
throw new Error(getSparseReshapeInputOutputMultipleErrorMessage(inputShape, outputShape));
|
|
}
|
|
outputShape[unknownIndex] = missing;
|
|
}
|
|
const outputSize = sizeFromShape(outputShape);
|
|
if (outputSize !== denseSize) {
|
|
throw new Error(getSparseReshapeInputOutputMismatchErrorMessage(inputShape, outputShape));
|
|
}
|
|
const inputRank = inputShape.length;
|
|
const inputStrides = [];
|
|
if (inputRank > 0) {
|
|
inputStrides[inputRank - 1] = 1;
|
|
for (let d = inputRank - 2; d >= 0; --d) {
|
|
inputStrides[d] = inputStrides[d + 1] * inputShape[d + 1];
|
|
}
|
|
}
|
|
const outputStrides = [];
|
|
if (outputRank > 0) {
|
|
outputStrides[outputRank - 1] = 1;
|
|
for (let d = outputRank - 2; d >= 0; --d) {
|
|
outputStrides[d] = outputStrides[d + 1] * outputShape[d + 1];
|
|
}
|
|
}
|
|
const newIndices = getArrayFromDType(inputDType, nnz * outputRank);
|
|
for (let i = 0; i < nnz; ++i) {
|
|
let id = 0;
|
|
for (let j2 = 0; j2 < inputRank; ++j2) {
|
|
id += inputIndices[i * inputRank + j2] * inputStrides[j2];
|
|
}
|
|
for (let j2 = 0; j2 < outputRank; ++j2) {
|
|
newIndices[i * outputRank + j2] = Math.trunc(id / outputStrides[j2]);
|
|
id %= outputStrides[j2];
|
|
}
|
|
}
|
|
return [newIndices, [nnz, outputRank], outputShape];
|
|
}
|
|
function sparseSegmentReductionImpl(input, inputShape, inputDType, indices, segmentIds, isMean = false, defaultValue = 0) {
|
|
const numIndices = indices.length;
|
|
const inputFlat = [inputShape[0], input.length / inputShape[0]];
|
|
const numCol = inputFlat[1];
|
|
const lastSegmentIdPlusOne = numIndices > 0 ? segmentIds[numIndices - 1] + 1 : 0;
|
|
const outputRows = lastSegmentIdPlusOne;
|
|
if (outputRows < 0) {
|
|
throw new Error(getSparseSegmentReductionNegativeSegmentIdsErrorMessage());
|
|
}
|
|
const outputShape = inputShape.slice();
|
|
outputShape[0] = outputRows;
|
|
const outputLength = outputShape.reduce((product, value) => product * value, 1);
|
|
const output = getArrayFromDType(inputDType, outputLength);
|
|
if (numIndices === 0) {
|
|
if (outputRows > 0) {
|
|
output.fill(defaultValue);
|
|
}
|
|
return [output, outputShape];
|
|
}
|
|
if (outputRows <= 0) {
|
|
throw new Error(getSparseSegmentReductionNegativeSegmentIdsErrorMessage());
|
|
}
|
|
let start = 0, end = 1;
|
|
let uninitializedIndex = 0;
|
|
let outIndex = segmentIds[start];
|
|
while (true) {
|
|
let nextIndex = 0;
|
|
if (end < numIndices) {
|
|
nextIndex = segmentIds[end];
|
|
if (outIndex === nextIndex) {
|
|
++end;
|
|
continue;
|
|
}
|
|
if (outIndex >= nextIndex) {
|
|
throw new Error(getSparseSegmentReductionNonIncreasingSegmentIdsErrorMessage());
|
|
}
|
|
}
|
|
if (outIndex < 0 || outIndex >= outputRows) {
|
|
throw new Error(getSparseSegmentReductionSegmentIdOutOfRangeErrorMessage(outIndex, outputRows));
|
|
}
|
|
if (outIndex > uninitializedIndex) {
|
|
output.fill(defaultValue, uninitializedIndex * numCol, outIndex * numCol);
|
|
}
|
|
for (let i = start; i < end; ++i) {
|
|
const index = indices[i];
|
|
if (index < 0 || index >= inputFlat[0]) {
|
|
throw new Error(getSparseSegmentReductionIndicesOutOfRangeErrorMessage(i, indices[i], inputFlat[0]));
|
|
}
|
|
for (let j2 = 0; j2 < numCol; j2++) {
|
|
output[outIndex * numCol + j2] += input[index * numCol + j2];
|
|
}
|
|
}
|
|
if (isMean) {
|
|
for (let j2 = 0; j2 < numCol; j2++) {
|
|
output[outIndex * numCol + j2] /= end - start;
|
|
}
|
|
}
|
|
start = end;
|
|
++end;
|
|
uninitializedIndex = outIndex + 1;
|
|
outIndex = nextIndex;
|
|
if (end > numIndices) {
|
|
break;
|
|
}
|
|
}
|
|
if (uninitializedIndex < outputRows) {
|
|
output.fill(defaultValue, uninitializedIndex * numCol, outputRows * numCol);
|
|
}
|
|
return [output, outputShape];
|
|
}
|
|
const sqrtImpl = createSimpleUnaryImpl((xi) => Math.sqrt(xi));
|
|
const sqrt$1 = unaryKernelFunc$1(Sqrt, (xi) => Math.sqrt(xi));
|
|
const sqrtConfig$1 = {
|
|
kernelName: Sqrt,
|
|
backendName: "cpu",
|
|
kernelFunc: sqrt$1
|
|
};
|
|
const squaredDifferenceImpl = createSimpleBinaryKernelImpl(((a, b) => {
|
|
const diff = a - b;
|
|
return diff * diff;
|
|
}));
|
|
const squaredDifference$1 = binaryKernelFunc$1(SquaredDifference, squaredDifferenceImpl);
|
|
const squaredDifferenceConfig$1 = {
|
|
kernelName: SquaredDifference,
|
|
backendName: "cpu",
|
|
kernelFunc: squaredDifference$1
|
|
};
|
|
const staticRegexReplaceImpl = createSimpleUnaryImpl((x, attrs) => {
|
|
const { pattern, replaceGlobal, rewrite } = attrs;
|
|
return x.replace(new RegExp(pattern, replaceGlobal ? "g" : ""), rewrite);
|
|
});
|
|
const staticRegexReplace$1 = unaryKernelFuncFromImpl(StaticRegexReplace, staticRegexReplaceImpl);
|
|
const staticRegexReplaceConfig$1 = {
|
|
kernelName: StaticRegexReplace,
|
|
backendName: "cpu",
|
|
kernelFunc: staticRegexReplace$1
|
|
};
|
|
function stridedSliceImpl(outShape, xBuf, strides, begin) {
|
|
const outBuf = buffer(outShape, xBuf.dtype);
|
|
for (let i = 0; i < outBuf.size; i++) {
|
|
const loc = outBuf.indexToLoc(i);
|
|
const newLoc = new Array(loc.length);
|
|
for (let j2 = 0; j2 < newLoc.length; j2++) {
|
|
newLoc[j2] = loc[j2] * strides[j2] + begin[j2];
|
|
}
|
|
outBuf.set(xBuf.get(...newLoc), ...loc);
|
|
}
|
|
return outBuf;
|
|
}
|
|
class StringNGramsOp {
|
|
constructor(separator, nGramWidths, leftPad, rightPad2, padWidth, preserveShortSequences) {
|
|
this.separator = encodeString(separator);
|
|
this.nGramWidths = nGramWidths;
|
|
this.leftPad = encodeString(leftPad);
|
|
this.rightPad = encodeString(rightPad2);
|
|
this.padWidth = padWidth;
|
|
this.preserveShort = preserveShortSequences;
|
|
}
|
|
getPadWidth(nGramWidth) {
|
|
return Math.min(this.padWidth < 0 ? nGramWidth - 1 : this.padWidth, nGramWidth - 1);
|
|
}
|
|
getNumNGrams(length, nGramWidth) {
|
|
const padWidth = this.getPadWidth(nGramWidth);
|
|
return Math.max(0, length + 2 * padWidth - nGramWidth + 1);
|
|
}
|
|
createNGrams(data, splitIndex, output, outputStartIndex, numNGrams, nGramWidth) {
|
|
for (let nGramIndex = 0; nGramIndex < numNGrams; ++nGramIndex) {
|
|
const padWidth = this.getPadWidth(nGramWidth);
|
|
const leftPadding = Math.max(0, padWidth - nGramIndex);
|
|
const rightPadding = Math.max(0, padWidth - (numNGrams - (nGramIndex + 1)));
|
|
const numTokens = nGramWidth - (leftPadding + rightPadding);
|
|
const dataStartIndex = splitIndex + (leftPadding > 0 ? 0 : nGramIndex - padWidth);
|
|
let nGramSize = 0;
|
|
nGramSize += leftPadding * this.leftPad.length;
|
|
for (let n = 0; n < numTokens; ++n) {
|
|
nGramSize += data[dataStartIndex + n].length;
|
|
}
|
|
nGramSize += rightPadding * this.rightPad.length;
|
|
const numSeparators = leftPadding + rightPadding + numTokens - 1;
|
|
nGramSize += numSeparators * this.separator.length;
|
|
output[outputStartIndex + nGramIndex] = new Uint8Array(nGramSize);
|
|
const nGram = output[outputStartIndex + nGramIndex];
|
|
let nextNGramIndex = 0;
|
|
const appendToNGram = (str) => str.forEach((value) => nGram[nextNGramIndex++] = value);
|
|
for (let n = 0; n < leftPadding; ++n) {
|
|
appendToNGram(this.leftPad);
|
|
appendToNGram(this.separator);
|
|
}
|
|
for (let n = 0; n < numTokens - 1; ++n) {
|
|
appendToNGram(data[dataStartIndex + n]);
|
|
appendToNGram(this.separator);
|
|
}
|
|
if (numTokens > 0) {
|
|
appendToNGram(data[dataStartIndex + numTokens - 1]);
|
|
for (let n = 0; n < rightPadding; ++n) {
|
|
appendToNGram(this.separator);
|
|
appendToNGram(this.rightPad);
|
|
}
|
|
} else {
|
|
for (let n = 0; n < rightPadding - 1; ++n) {
|
|
appendToNGram(this.rightPad);
|
|
appendToNGram(this.separator);
|
|
}
|
|
appendToNGram(this.rightPad);
|
|
}
|
|
}
|
|
}
|
|
// Data and splits together form the definition of the ragged tensor,
|
|
// where data is 1 dimensional and contains the values of the tensor
|
|
// and splits denotes the indices at which each row starts.
|
|
compute(data, splits) {
|
|
const inputDataSize = data.length;
|
|
const splitsSize = splits.length;
|
|
if (splitsSize > 0) {
|
|
let prevSplit = splits[0];
|
|
if (prevSplit !== 0) {
|
|
throw new Error(`First split value must be 0, got ${prevSplit}`);
|
|
}
|
|
for (let i = 1; i < splitsSize; ++i) {
|
|
let validSplits = splits[i] >= prevSplit;
|
|
validSplits = validSplits && splits[i] <= inputDataSize;
|
|
if (!validSplits) {
|
|
throw new Error(`Invalid split value ${splits[i]}, must be in [${prevSplit}, ${inputDataSize}]`);
|
|
}
|
|
prevSplit = splits[i];
|
|
}
|
|
if (prevSplit !== inputDataSize) {
|
|
throw new Error(`Last split value must be data size. Expected ${inputDataSize}, got ${prevSplit}`);
|
|
}
|
|
}
|
|
const numBatchItems = splitsSize - 1;
|
|
const nGramsSplits = getArrayFromDType("int32", splitsSize);
|
|
if (inputDataSize === 0 || splitsSize === 0) {
|
|
const empty = new Array(inputDataSize);
|
|
for (let i = 0; i <= numBatchItems; ++i) {
|
|
nGramsSplits[i] = 0;
|
|
}
|
|
return [empty, nGramsSplits];
|
|
}
|
|
nGramsSplits[0] = 0;
|
|
for (let i = 1; i <= numBatchItems; ++i) {
|
|
const length = splits[i] - splits[i - 1];
|
|
let numNGrams = 0;
|
|
this.nGramWidths.forEach((nGramWidth) => {
|
|
numNGrams += this.getNumNGrams(length, nGramWidth);
|
|
});
|
|
if (this.preserveShort && length > 0 && numNGrams === 0) {
|
|
numNGrams = 1;
|
|
}
|
|
nGramsSplits[i] = nGramsSplits[i - 1] + numNGrams;
|
|
}
|
|
const nGrams = new Array(nGramsSplits[numBatchItems]);
|
|
for (let i = 0; i < numBatchItems; ++i) {
|
|
const splitIndex = splits[i];
|
|
let outputStartIdx = nGramsSplits[i];
|
|
this.nGramWidths.forEach((nGramWidth) => {
|
|
const length = splits[i + 1] - splits[i];
|
|
const numNGrams = this.getNumNGrams(length, nGramWidth);
|
|
this.createNGrams(data, splitIndex, nGrams, outputStartIdx, numNGrams, nGramWidth);
|
|
outputStartIdx += numNGrams;
|
|
});
|
|
if (this.preserveShort && outputStartIdx === nGramsSplits[i]) {
|
|
const dataLength = splits[i + 1] - splits[i];
|
|
if (dataLength === 0) {
|
|
continue;
|
|
}
|
|
const nGramWidth = dataLength + 2 * this.padWidth;
|
|
const numNGrams = 1;
|
|
this.createNGrams(data, splitIndex, nGrams, outputStartIdx, numNGrams, nGramWidth);
|
|
}
|
|
}
|
|
return [nGrams, nGramsSplits];
|
|
}
|
|
}
|
|
function stringNGramsImpl(data, dataSplits, separator, nGramWidths, leftPad, rightPad2, padWidth, preserveShortSequences) {
|
|
return new StringNGramsOp(separator, nGramWidths, leftPad, rightPad2, padWidth, preserveShortSequences).compute(data, dataSplits);
|
|
}
|
|
function split(str, delimiters, skipEmpty, result) {
|
|
if (!str.length) {
|
|
return;
|
|
}
|
|
if (delimiters.length === 0) {
|
|
for (let i = 0; i < str.length; ++i) {
|
|
result.push(str.subarray(i, i + 1));
|
|
}
|
|
return;
|
|
}
|
|
if (delimiters.length === 1) {
|
|
const delimiter = delimiters[0];
|
|
let f = str.indexOf(delimiter);
|
|
while (f !== -1) {
|
|
const token = str.subarray(0, f);
|
|
if (!skipEmpty || token.length !== 0) {
|
|
result.push(token);
|
|
}
|
|
str = str.subarray(f + 1);
|
|
f = str.indexOf(delimiter);
|
|
}
|
|
if (!skipEmpty || str.length !== 0) {
|
|
result.push(str);
|
|
}
|
|
return;
|
|
}
|
|
let tokenStart = 0;
|
|
for (let i = 0; i < str.length + 1; i++) {
|
|
if (i === str.length || delimiters.indexOf(str[i]) !== -1) {
|
|
const token = str.subarray(tokenStart, i);
|
|
if (!skipEmpty || token.length !== 0) {
|
|
result.push(token);
|
|
}
|
|
tokenStart = i + 1;
|
|
}
|
|
}
|
|
}
|
|
function stringSplitImpl(input, delimiter, skipEmpty) {
|
|
const batchSize = input.length;
|
|
const tokens = [];
|
|
let outputSize = 0;
|
|
let maxNumEntries = 0;
|
|
const numIndices = new Array(batchSize);
|
|
for (let i = 0; i < batchSize; ++i) {
|
|
const prevTokensLength = tokens.length;
|
|
split(input[i], delimiter, skipEmpty, tokens);
|
|
const nEntries = tokens.length - prevTokensLength;
|
|
numIndices[i] = nEntries;
|
|
outputSize += nEntries;
|
|
maxNumEntries = Math.max(maxNumEntries, nEntries);
|
|
}
|
|
const indices = getArrayFromDType("int32", outputSize * 2);
|
|
const values = new Array(outputSize);
|
|
const shape = [batchSize, maxNumEntries];
|
|
let c = 0;
|
|
for (let i = 0; i < batchSize; ++i) {
|
|
for (let j2 = 0; j2 < numIndices[i]; ++j2) {
|
|
indices[c * 2] = i;
|
|
indices[c * 2 + 1] = j2;
|
|
values[c] = tokens[c];
|
|
++c;
|
|
}
|
|
}
|
|
return [indices, values, shape];
|
|
}
|
|
function stringToHashBucketFastImpl(input, numBuckets) {
|
|
const output = getArrayFromDType("int32", input.length);
|
|
for (let i = 0; i < input.length; ++i) {
|
|
output[i] = fingerPrint64(input[i]).modulo(numBuckets).getLowBitsUnsigned();
|
|
}
|
|
return output;
|
|
}
|
|
const subImpl = createSimpleBinaryKernelImpl(((aValue, bValue) => aValue - bValue));
|
|
const subComplexImpl = createComplexBinaryKernelImpl(((aReal, aImag, bReal, bImag) => {
|
|
return { real: aReal - bReal, imag: aImag - bImag };
|
|
}));
|
|
const sub$1 = binaryKernelFunc$1(Sub, subImpl, subComplexImpl);
|
|
const subConfig$1 = {
|
|
kernelName: Sub,
|
|
backendName: "cpu",
|
|
kernelFunc: sub$1
|
|
};
|
|
function tileImpl(xBuf, reps) {
|
|
const newShape = new Array(xBuf.rank);
|
|
for (let i = 0; i < newShape.length; i++) {
|
|
newShape[i] = xBuf.shape[i] * reps[i];
|
|
}
|
|
const result = buffer(newShape, xBuf.dtype);
|
|
for (let i = 0; i < result.values.length; ++i) {
|
|
const newLoc = result.indexToLoc(i);
|
|
const originalLoc = new Array(xBuf.rank);
|
|
for (let j2 = 0; j2 < originalLoc.length; j2++) {
|
|
originalLoc[j2] = newLoc[j2] % xBuf.shape[j2];
|
|
}
|
|
const originalIndex = xBuf.locToIndex(originalLoc);
|
|
result.values[i] = xBuf.values[originalIndex];
|
|
}
|
|
return result;
|
|
}
|
|
const comparePair = (a, b) => {
|
|
const valueDiff = b.value - a.value;
|
|
return valueDiff === 0 ? a.index - b.index : valueDiff;
|
|
};
|
|
function select$2(array, k, left = 0, right = array.length - 1) {
|
|
while (right > left) {
|
|
if (right - left > 600) {
|
|
const n = right - left + 1;
|
|
const i2 = k - left + 1;
|
|
const z2 = Math.log(n);
|
|
const s = 0.5 * Math.exp(2 * z2 / 3);
|
|
const sd = 0.5 * Math.sqrt(z2 * s * (n - s) / n) * Math.sign(i2 - n / 2);
|
|
const newLeft = Math.max(left, Math.floor(k - i2 * s / n + sd));
|
|
const newRight = Math.min(right, Math.floor(k + (n - i2) * s / n + sd));
|
|
select$2(array, k, newLeft, newRight);
|
|
}
|
|
const t2 = array[k];
|
|
let i = left;
|
|
let j2 = right;
|
|
swap(array, left, k);
|
|
if (comparePair(array[right], t2) > 0) {
|
|
swap(array, left, right);
|
|
}
|
|
while (i < j2) {
|
|
swap(array, i, j2);
|
|
i++;
|
|
j2--;
|
|
while (comparePair(array[i], t2) < 0) {
|
|
i = i + 1;
|
|
}
|
|
while (comparePair(array[j2], t2) > 0) {
|
|
j2 = j2 - 1;
|
|
}
|
|
}
|
|
if (comparePair(array[left], t2) === 0) {
|
|
swap(array, left, j2);
|
|
} else {
|
|
j2 = j2 + 1;
|
|
swap(array, j2, right);
|
|
}
|
|
if (j2 <= k) {
|
|
left = j2 + 1;
|
|
}
|
|
if (k <= j2) {
|
|
right = j2 - 1;
|
|
}
|
|
}
|
|
}
|
|
function topKImpl(x, xShape, xDtype, k, sorted) {
|
|
const lastDim = xShape[xShape.length - 1];
|
|
const [batch, size] = [x.length / lastDim, lastDim];
|
|
const allTopKVals = getTypedArrayFromDType(xDtype, batch * k);
|
|
const allTopKIndices = getTypedArrayFromDType("int32", batch * k);
|
|
for (let b = 0; b < batch; b++) {
|
|
const offset = b * size;
|
|
const vals = x.subarray(offset, offset + size);
|
|
let valAndInd = new Array(vals.length);
|
|
vals.forEach((value, index) => valAndInd[index] = { value, index });
|
|
if (k < valAndInd.length) {
|
|
select$2(valAndInd, k);
|
|
valAndInd = valAndInd.slice(0, k);
|
|
}
|
|
if (sorted) {
|
|
valAndInd.sort(comparePair);
|
|
}
|
|
const outOffset = b * k;
|
|
const topKVals = allTopKVals.subarray(outOffset, outOffset + k);
|
|
const topKIndices = allTopKIndices.subarray(outOffset, outOffset + k);
|
|
for (let i = 0; i < k; i++) {
|
|
topKVals[i] = valAndInd[i].value;
|
|
topKIndices[i] = valAndInd[i].index;
|
|
}
|
|
}
|
|
const outputShape = xShape.slice();
|
|
outputShape[outputShape.length - 1] = k;
|
|
return [
|
|
buffer(outputShape, xDtype, allTopKVals),
|
|
buffer(outputShape, "int32", allTopKIndices)
|
|
];
|
|
}
|
|
function uniqueImpl(values, axis, shape, dtype) {
|
|
const $axis = parseAxisParam(axis, shape)[0];
|
|
const newShape = [1, shape[0], 1];
|
|
for (let i = 0; i < $axis; i++) {
|
|
newShape[0] *= shape[i];
|
|
}
|
|
newShape[1] = shape[$axis];
|
|
for (let i = $axis + 1; i < shape.length; i++) {
|
|
newShape[2] *= shape[i];
|
|
}
|
|
const uniqueElements = /* @__PURE__ */ new Map();
|
|
const indices = new Int32Array(shape[$axis]);
|
|
const inputBuffer = new TensorBuffer(newShape, dtype, values);
|
|
const uniqueIndices = [];
|
|
const is1DTensor = newShape[0] === 1 && newShape[2] === 1;
|
|
for (let i = 0; i < shape[$axis]; i++) {
|
|
let element;
|
|
if (is1DTensor) {
|
|
element = values[i].toString();
|
|
} else {
|
|
const axisValues = [];
|
|
for (let m = 0; m < newShape[0]; m++) {
|
|
for (let n = 0; n < newShape[2]; n++) {
|
|
axisValues.push(inputBuffer.get(m, i, n));
|
|
}
|
|
}
|
|
element = axisValues.join(",");
|
|
}
|
|
const existingIndex = uniqueElements.get(element);
|
|
if (existingIndex != null) {
|
|
indices[i] = existingIndex;
|
|
} else {
|
|
const uniqueIndex = uniqueElements.size;
|
|
uniqueElements.set(element, uniqueIndex);
|
|
indices[i] = uniqueIndex;
|
|
uniqueIndices.push(i);
|
|
}
|
|
}
|
|
const outputTmpShape = newShape.slice();
|
|
outputTmpShape[1] = uniqueElements.size;
|
|
const outputBuffer = new TensorBuffer(outputTmpShape, dtype);
|
|
uniqueIndices.forEach((uniqueElementIndex, i) => {
|
|
for (let m = 0; m < newShape[0]; m++) {
|
|
for (let n = 0; n < newShape[2]; n++) {
|
|
outputBuffer.set(inputBuffer.get(m, uniqueElementIndex, n), m, i, n);
|
|
}
|
|
}
|
|
});
|
|
const outputShape = shape.slice();
|
|
outputShape[$axis] = outputTmpShape[1];
|
|
return {
|
|
outputValues: outputBuffer.values,
|
|
outputShape,
|
|
indices
|
|
};
|
|
}
|
|
const shared = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
__proto__: null,
|
|
addImpl,
|
|
bincountImpl,
|
|
bincountReduceImpl,
|
|
bitwiseAndImpl,
|
|
castImpl,
|
|
ceilImpl,
|
|
concatImpl: concatImpl$1,
|
|
equalImpl,
|
|
expImpl,
|
|
expm1Impl,
|
|
floorDivImpl,
|
|
floorImpl,
|
|
gatherNdImpl,
|
|
gatherV2Impl,
|
|
greaterEqualImpl,
|
|
greaterImpl,
|
|
lessEqualImpl,
|
|
lessImpl,
|
|
linSpaceImpl,
|
|
logImpl,
|
|
maxImpl: maxImpl$1,
|
|
maximumImpl,
|
|
minimumImpl,
|
|
multiplyImpl,
|
|
negImpl,
|
|
notEqualImpl,
|
|
prodImpl,
|
|
raggedGatherImpl,
|
|
raggedRangeImpl,
|
|
raggedTensorToTensorImpl,
|
|
rangeImpl,
|
|
rsqrtImpl,
|
|
scatterImpl,
|
|
sigmoidImpl,
|
|
simpleAbsImpl,
|
|
sliceImpl,
|
|
sparseFillEmptyRowsImpl,
|
|
sparseReshapeImpl,
|
|
sparseSegmentReductionImpl,
|
|
sqrtImpl,
|
|
squaredDifferenceImpl,
|
|
staticRegexReplaceImpl,
|
|
stridedSliceImpl,
|
|
stringNGramsImpl,
|
|
stringSplitImpl,
|
|
stringToHashBucketFastImpl,
|
|
subImpl,
|
|
tileImpl,
|
|
topKImpl,
|
|
transposeImpl: transposeImpl$1,
|
|
uniqueImpl
|
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
const { addImpl: addImplCPU, bincountImpl: bincountImplCPU, bincountReduceImpl: bincountReduceImplCPU, bitwiseAndImpl: bitwiseAndImplCPU, castImpl: castImplCPU, ceilImpl: ceilImplCPU, concatImpl: concatImplCPU, equalImpl: equalImplCPU, expImpl: expImplCPU, expm1Impl: expm1ImplCPU, floorImpl: floorImplCPU, gatherNdImpl: gatherNdImplCPU, gatherV2Impl: gatherV2ImplCPU, greaterImpl: greaterImplCPU, greaterEqualImpl: greaterEqualImplCPU, lessImpl: lessImplCPU, lessEqualImpl: lessEqualImplCPU, linSpaceImpl: linSpaceImplCPU, logImpl: logImplCPU, maxImpl: maxImplCPU, maximumImpl: maximumImplCPU, minimumImpl: minimumImplCPU, multiplyImpl: multiplyImplCPU, negImpl: negImplCPU, notEqualImpl: notEqualImplCPU, prodImpl: prodImplCPU, raggedGatherImpl: raggedGatherImplCPU, raggedRangeImpl: raggedRangeImplCPU, raggedTensorToTensorImpl: raggedTensorToTensorImplCPU, rangeImpl: rangeImplCPU, rsqrtImpl: rsqrtImplCPU, scatterImpl: scatterImplCPU, sigmoidImpl: sigmoidImplCPU, simpleAbsImpl: simpleAbsImplCPU, sliceImpl: sliceImplCPU, sparseFillEmptyRowsImpl: sparseFillEmptyRowsImplCPU, sparseReshapeImpl: sparseReshapeImplCPU, sparseSegmentReductionImpl: sparseSegmentReductionImplCPU, sqrtImpl: sqrtImplCPU, staticRegexReplaceImpl: staticRegexReplaceImplCPU, stridedSliceImpl: stridedSliceImplCPU, stringNGramsImpl: stringNGramsImplCPU, stringSplitImpl: stringSplitImplCPU, stringToHashBucketFastImpl: stringToHashBucketFastImplCPU, subImpl: subImplCPU, tileImpl: tileImplCPU, topKImpl: topKImplCPU, transposeImpl: transposeImplCPU, uniqueImpl: uniqueImplCPU } = shared;
|
|
function getVecChannels(name, rank) {
|
|
return ["x", "y", "z", "w", "u", "v"].slice(0, rank).map((d) => `${name}.${d}`);
|
|
}
|
|
function getChannels(name, rank) {
|
|
if (rank === 1) {
|
|
return [name];
|
|
}
|
|
return getVecChannels(name, rank);
|
|
}
|
|
function getSourceCoords$2(rank, dims) {
|
|
if (rank === 1) {
|
|
return "rc";
|
|
}
|
|
let coords2 = "";
|
|
for (let i = 0; i < rank; i++) {
|
|
coords2 += dims[i];
|
|
if (i < rank - 1) {
|
|
coords2 += ",";
|
|
}
|
|
}
|
|
return coords2;
|
|
}
|
|
class PackProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = false;
|
|
this.packedOutput = true;
|
|
this.outputShape = outputShape;
|
|
this.rank = outputShape.length;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
if (this.rank === 0) {
|
|
this.userCode = `
|
|
void main() {
|
|
setOutput(vec4(getA(), 0., 0., 0.));
|
|
}
|
|
`;
|
|
} else {
|
|
const channels = getChannels("rc", this.rank);
|
|
const dtype = getCoordsDataType(this.rank);
|
|
const outOfBoundsCondition = this.getOutOfBoundsCondition(channels);
|
|
const setup = this.getSetup(channels);
|
|
const output = this.getOutput(channels);
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} rc = getOutputCoords();
|
|
|
|
if(${outOfBoundsCondition}) {
|
|
setOutput(vec4(0));
|
|
} else {
|
|
${setup}
|
|
|
|
setOutput(vec4(${output}));
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
getSourceCoordsArr(dims) {
|
|
const coords2 = [];
|
|
for (let row = 0; row <= 1; row++) {
|
|
for (let col = 0; col <= 1; col++) {
|
|
let coord = `${row === 0 ? "r" : "rp1"}, ${col === 0 ? "c" : "cp1"}`;
|
|
for (let d = 2; d < this.rank; d++) {
|
|
coord = `${dims[dims.length - 1 - d]},` + coord;
|
|
}
|
|
coords2.push(coord);
|
|
}
|
|
}
|
|
return coords2;
|
|
}
|
|
getOutOfBoundsCondition(dims) {
|
|
if (this.rank === 1) {
|
|
return `rc > ${this.enableShapeUniforms ? "outShape" : this.outputShape[0]}`;
|
|
}
|
|
let cond = "";
|
|
for (let i = this.rank - 2; i < this.rank; i++) {
|
|
cond += `${dims[i]} >= ${this.enableShapeUniforms ? `outShape[${i}]` : this.outputShape[i]}`;
|
|
if (i < this.rank - 1) {
|
|
cond += "||";
|
|
}
|
|
}
|
|
return cond;
|
|
}
|
|
getSetup(dims) {
|
|
if (this.rank === 1) {
|
|
return "";
|
|
}
|
|
const innerDims = dims.slice(-2);
|
|
const col = this.enableShapeUniforms ? `outShape[${this.rank} - 1]` : this.outputShape[this.rank - 1];
|
|
const row = this.enableShapeUniforms ? `outShape[${this.rank} - 2]` : this.outputShape[this.rank - 2];
|
|
return `
|
|
int r = ${innerDims[0]};
|
|
int c = ${innerDims[1]};
|
|
int rp1 = r + 1;
|
|
int cp1 = c + 1;
|
|
|
|
bool cEdge = cp1 >= ${col};
|
|
bool rEdge = rp1 >= ${row};
|
|
`;
|
|
}
|
|
getOutput(dims) {
|
|
const sourceCoords = this.getSourceCoordsArr(dims);
|
|
if (this.rank === 1) {
|
|
const outShape = this.enableShapeUniforms ? "outShape" : this.outputShape[0];
|
|
return `getA(rc), (rc + 1 >= ${outShape} ? 0. : getA(rc + 1)), 0, 0`;
|
|
}
|
|
return `getA(${sourceCoords[0]}),
|
|
cEdge ? 0. : getA(${sourceCoords[1]}),
|
|
rEdge ? 0. : getA(${sourceCoords[2]}),
|
|
rEdge || cEdge ? 0. : getA(${sourceCoords[3]})`;
|
|
}
|
|
}
|
|
class ReshapePackedProgram {
|
|
constructor(outputShape, inputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [{ name: "inputShape", type: "ivec3" }];
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
let mainLoop = ``;
|
|
for (let i = 0; i < 4; i++) {
|
|
let thisRC = `thisRC = rc;`;
|
|
if (i % 2 === 1) {
|
|
thisRC += `thisRC.z += 1;`;
|
|
}
|
|
if (i > 1) {
|
|
thisRC += `thisRC.y += 1;`;
|
|
}
|
|
mainLoop += `
|
|
${thisRC}
|
|
${i > 0 ? `if(thisRC.y < rows && thisRC.z < cols){` : ""}
|
|
int flatIndex = getFlatIndex(thisRC);
|
|
|
|
ivec3 inputRC = inputCoordsFromReshapedOutCoords(flatIndex);
|
|
vec2 inputRCInnerDims = vec2(float(inputRC.y),float(inputRC.z));
|
|
|
|
result[${i}] =
|
|
getChannel(getA(inputRC.x, inputRC.y, inputRC.z), inputRCInnerDims);
|
|
${i > 0 ? "}" : ""}
|
|
`;
|
|
}
|
|
this.userCode = `
|
|
${getReshapedInputCoords(inputShape, this.enableShapeUniforms)}
|
|
${this.enableShapeUniforms ? getFlatIndexFrom3DOutput() : getFlatIndexFrom3D(outputShape)}
|
|
|
|
void main() {
|
|
ivec3 rc = getOutputCoords();
|
|
|
|
vec4 result = vec4(0.);
|
|
|
|
ivec3 thisRC;
|
|
int rows = ${this.enableShapeUniforms ? "outShape[1]" : outputShape[1]};
|
|
int cols = ${this.enableShapeUniforms ? "outShape[2]" : outputShape[2]};
|
|
|
|
${mainLoop}
|
|
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getReshapedInputCoords(shape, enableShapeUniforms) {
|
|
const coordsFromIndexSnippet = enableShapeUniforms ? getLogicalCoordinatesFromFlatIndexByUniform(["r", "c", "d"], "inputShape") : getLogicalCoordinatesFromFlatIndex(["r", "c", "d"], shape);
|
|
return `
|
|
ivec3 inputCoordsFromReshapedOutCoords(int index) {
|
|
${coordsFromIndexSnippet}
|
|
return ivec3(r, c, d);
|
|
}
|
|
`;
|
|
}
|
|
class TextureManager {
|
|
constructor(gpgpu) {
|
|
this.gpgpu = gpgpu;
|
|
this.numUsedTextures = 0;
|
|
this.numFreeTextures = 0;
|
|
this._numBytesAllocated = 0;
|
|
this._numBytesFree = 0;
|
|
this.freeTextures = {};
|
|
this.usedTextures = {};
|
|
this.logEnabled = false;
|
|
}
|
|
acquireTexture(shapeRC, usage, isPacked) {
|
|
const physicalTexType = getPhysicalFromLogicalTextureType(usage, isPacked);
|
|
const shapeKey = getKeyFromTextureShape(shapeRC, physicalTexType, isPacked);
|
|
if (!(shapeKey in this.freeTextures)) {
|
|
this.freeTextures[shapeKey] = [];
|
|
}
|
|
if (!(shapeKey in this.usedTextures)) {
|
|
this.usedTextures[shapeKey] = [];
|
|
}
|
|
const texBytes = computeBytes(shapeRC, physicalTexType, this.gpgpu.gl, this.gpgpu.textureConfig, isPacked);
|
|
if (this.freeTextures[shapeKey].length > 0) {
|
|
this.numFreeTextures--;
|
|
this.numUsedTextures++;
|
|
this._numBytesFree -= texBytes;
|
|
this.log();
|
|
const newTexture2 = this.freeTextures[shapeKey].pop();
|
|
this.usedTextures[shapeKey].push(newTexture2);
|
|
return newTexture2;
|
|
}
|
|
let newTexture;
|
|
if (physicalTexType === PhysicalTextureType.PACKED_2X2_FLOAT32) {
|
|
newTexture = this.gpgpu.createPackedMatrixTexture(shapeRC[0], shapeRC[1]);
|
|
} else if (physicalTexType === PhysicalTextureType.PACKED_2X2_FLOAT16) {
|
|
newTexture = this.gpgpu.createFloat16PackedMatrixTexture(shapeRC[0], shapeRC[1]);
|
|
} else if (physicalTexType === PhysicalTextureType.UNPACKED_FLOAT32) {
|
|
newTexture = this.gpgpu.createFloat32MatrixTexture(shapeRC[0], shapeRC[1]);
|
|
} else if (physicalTexType === PhysicalTextureType.UNPACKED_FLOAT16) {
|
|
newTexture = this.gpgpu.createFloat16MatrixTexture(shapeRC[0], shapeRC[1]);
|
|
} else if (physicalTexType === PhysicalTextureType.PACKED_4X1_UNSIGNED_BYTE) {
|
|
newTexture = this.gpgpu.createUnsignedBytesMatrixTexture(shapeRC[0], shapeRC[1]);
|
|
}
|
|
this.usedTextures[shapeKey].push(newTexture);
|
|
this.numUsedTextures++;
|
|
this._numBytesAllocated += texBytes;
|
|
this.log();
|
|
return newTexture;
|
|
}
|
|
releaseTexture(texture, shape, logicalTexType, isPacked) {
|
|
if (this.freeTextures == null) {
|
|
return;
|
|
}
|
|
const physicalTexType = getPhysicalFromLogicalTextureType(logicalTexType, isPacked);
|
|
const shapeKey = getKeyFromTextureShape(shape, physicalTexType, isPacked);
|
|
if (!(shapeKey in this.freeTextures)) {
|
|
this.freeTextures[shapeKey] = [];
|
|
}
|
|
const texBytes = computeBytes(shape, physicalTexType, this.gpgpu.gl, this.gpgpu.textureConfig, isPacked);
|
|
const deleteTexThreshold = env().getNumber("WEBGL_DELETE_TEXTURE_THRESHOLD");
|
|
if (deleteTexThreshold !== -1 && this._numBytesAllocated > deleteTexThreshold) {
|
|
this.gpgpu.deleteMatrixTexture(texture.texture);
|
|
this._numBytesAllocated -= texBytes;
|
|
} else {
|
|
this.freeTextures[shapeKey].push(texture);
|
|
this.numFreeTextures++;
|
|
this._numBytesFree += texBytes;
|
|
}
|
|
this.numUsedTextures--;
|
|
const texList = this.usedTextures[shapeKey];
|
|
const texIndex = texList && texList.indexOf(texture);
|
|
if (texIndex == null || texIndex < 0) {
|
|
throw new Error("Cannot release a texture that was never provided by this texture manager");
|
|
}
|
|
texList[texIndex] = texList[texList.length - 1];
|
|
texList.pop();
|
|
this.log();
|
|
}
|
|
log() {
|
|
if (!this.logEnabled) {
|
|
return;
|
|
}
|
|
const total = this.numFreeTextures + this.numUsedTextures;
|
|
console.log("Free/Used", `${this.numFreeTextures} / ${this.numUsedTextures}`, `(${total})`);
|
|
const freeRatio = this._numBytesFree / this._numBytesAllocated;
|
|
console.log(`Bytes allocated: ${this._numBytesAllocated}`);
|
|
console.log(`Bytes unused: ${this._numBytesFree} (${Math.round(100 * freeRatio)}%)`);
|
|
}
|
|
get numBytesAllocated() {
|
|
return this._numBytesAllocated;
|
|
}
|
|
get numBytesFree() {
|
|
return this._numBytesFree;
|
|
}
|
|
getNumUsedTextures() {
|
|
return this.numUsedTextures;
|
|
}
|
|
getNumFreeTextures() {
|
|
return this.numFreeTextures;
|
|
}
|
|
dispose() {
|
|
if (this.freeTextures == null) {
|
|
return;
|
|
}
|
|
for (const texShape in this.freeTextures) {
|
|
this.freeTextures[texShape].forEach((tex) => {
|
|
this.gpgpu.deleteMatrixTexture(tex.texture);
|
|
});
|
|
}
|
|
for (const texShape in this.usedTextures) {
|
|
this.usedTextures[texShape].forEach((tex) => {
|
|
this.gpgpu.deleteMatrixTexture(tex.texture);
|
|
});
|
|
}
|
|
this.freeTextures = null;
|
|
this.usedTextures = null;
|
|
this.numUsedTextures = 0;
|
|
this.numFreeTextures = 0;
|
|
this._numBytesAllocated = 0;
|
|
this._numBytesFree = 0;
|
|
}
|
|
}
|
|
function numBytesForInternalFormat(gl, internalFormat) {
|
|
const glany = gl;
|
|
if (internalFormat === glany.R32F) {
|
|
return 4;
|
|
} else if (internalFormat === glany.R16F) {
|
|
return 2;
|
|
} else if (internalFormat === glany.RGBA32F) {
|
|
return 16;
|
|
} else if (internalFormat === gl.RGBA) {
|
|
return 16;
|
|
} else if (internalFormat === glany.RGBA16F) {
|
|
return 8;
|
|
} else if (internalFormat === glany.RGBA8) {
|
|
return 4;
|
|
}
|
|
throw new Error(`Unknown internal format ${internalFormat}`);
|
|
}
|
|
function computeBytes(shape, physicalTexType, gl, textureConfig, isPacked) {
|
|
const internalFormat = internalFormatForPhysicalTexType(physicalTexType, textureConfig);
|
|
let numElements;
|
|
if (isPacked) {
|
|
const [packedWidth, packedHeight] = getPackedMatrixTextureShapeWidthHeight(shape[0], shape[1]);
|
|
numElements = packedWidth * packedHeight;
|
|
} else {
|
|
const [width, height] = getUnpackedMatrixTextureShapeWidthHeight(shape[0], shape[1]);
|
|
numElements = width * height;
|
|
}
|
|
const bytesPerElement2 = numBytesForInternalFormat(gl, internalFormat);
|
|
return numElements * bytesPerElement2;
|
|
}
|
|
function internalFormatForPhysicalTexType(physicalTexType, textureConfig) {
|
|
switch (physicalTexType) {
|
|
case PhysicalTextureType.PACKED_2X2_FLOAT32:
|
|
return getInternalFormatForPackedMatrixTexture(textureConfig);
|
|
case PhysicalTextureType.PACKED_2X2_FLOAT16:
|
|
return getInternalFormatForFloat16PackedMatrixTexture(textureConfig);
|
|
case PhysicalTextureType.UNPACKED_FLOAT32:
|
|
return getInternalFormatForFloat32MatrixTexture(textureConfig);
|
|
case PhysicalTextureType.UNPACKED_FLOAT16:
|
|
return getInternalFormatForFloat16MatrixTexture(textureConfig);
|
|
case PhysicalTextureType.PACKED_4X1_UNSIGNED_BYTE:
|
|
return getInternalFormatForUnsignedBytesMatrixTexture(textureConfig);
|
|
default:
|
|
throw new Error(`Unknown physical texture type ${physicalTexType}`);
|
|
}
|
|
}
|
|
function getPhysicalTextureForRendering(isPacked) {
|
|
if (env().getBool("WEBGL_RENDER_FLOAT32_ENABLED")) {
|
|
if (isPacked) {
|
|
return PhysicalTextureType.PACKED_2X2_FLOAT32;
|
|
}
|
|
return PhysicalTextureType.UNPACKED_FLOAT32;
|
|
}
|
|
if (isPacked) {
|
|
return PhysicalTextureType.PACKED_2X2_FLOAT16;
|
|
}
|
|
return PhysicalTextureType.UNPACKED_FLOAT16;
|
|
}
|
|
function getPhysicalFromLogicalTextureType(logicalTexType, isPacked) {
|
|
if (logicalTexType === TextureUsage.UPLOAD) {
|
|
return PhysicalTextureType.PACKED_2X2_FLOAT32;
|
|
} else if (logicalTexType === TextureUsage.RENDER || logicalTexType == null) {
|
|
return getPhysicalTextureForRendering(isPacked);
|
|
} else if (logicalTexType === TextureUsage.DOWNLOAD || logicalTexType === TextureUsage.PIXELS) {
|
|
return PhysicalTextureType.PACKED_4X1_UNSIGNED_BYTE;
|
|
}
|
|
throw new Error(`Unknown logical texture type ${logicalTexType}`);
|
|
}
|
|
function getKeyFromTextureShape(shapeRowsCol, physicalTexType, isPacked) {
|
|
return `${shapeRowsCol[0]}_${shapeRowsCol[1]}_${physicalTexType}_${isPacked}`;
|
|
}
|
|
class UnaryOpProgram {
|
|
constructor(aShape, opSnippet) {
|
|
this.variableNames = ["A"];
|
|
this.outputShape = aShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
this.userCode = `
|
|
float unaryOperation(float x) {
|
|
${opSnippet}
|
|
}
|
|
|
|
void main() {
|
|
float x = getAAtOutCoords();
|
|
float y = unaryOperation(x);
|
|
|
|
setOutput(y);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const CHECK_NAN_SNIPPET$1 = `if (isnan(x)) return x;`;
|
|
const LINEAR$1 = `return x;`;
|
|
const ABS$1 = `return abs(x);`;
|
|
const ELU$2 = `return (x >= 0.0) ? x : (exp(x) - 1.0);`;
|
|
const RELU$2 = CHECK_NAN_SNIPPET$1 + `
|
|
return (x < 0.0) ? 0.0 : x;
|
|
`;
|
|
const RELU6$2 = CHECK_NAN_SNIPPET$1 + `
|
|
return (x < 0.0) ? 0.0 : min(6.0, x);
|
|
`;
|
|
const CLONE = "return x;";
|
|
const SIGMOID$2 = `return 1.0 / (1.0 + exp(-1.0 * x));`;
|
|
const LINEAR = `return x;`;
|
|
const ELU$1 = `
|
|
vec4 result;
|
|
|
|
result.r = (x.r >= 0.0) ? x.r : (exp(x.r) - 1.0);
|
|
result.g = (x.g >= 0.0) ? x.g : (exp(x.g) - 1.0);
|
|
result.b = (x.b >= 0.0) ? x.b : (exp(x.b) - 1.0);
|
|
result.a = (x.a >= 0.0) ? x.a : (exp(x.a) - 1.0);
|
|
|
|
return result;
|
|
`;
|
|
const RELU$1 = `
|
|
vec4 result = x * vec4(greaterThanEqual(x, vec4(0.0)));
|
|
bvec4 isNaN = isnan(x);
|
|
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
const RELU6$1 = `
|
|
vec4 result = min(x, vec4(6.)) * vec4(greaterThanEqual(x, vec4(0.0)));
|
|
bvec4 isNaN = isnan(x);
|
|
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
const SIGMOID$1 = `return 1.0 / (1.0 + exp(-1.0 * x));`;
|
|
class UnaryOpPackedProgram {
|
|
constructor(aShape, opSnippet) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = aShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
this.userCode = `
|
|
vec4 unaryOperation(vec4 x) {
|
|
${opSnippet}
|
|
}
|
|
|
|
void main() {
|
|
vec4 x = getAAtOutCoords();
|
|
vec4 y = unaryOperation(x);
|
|
|
|
setOutput(y);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class UnpackProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = false;
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const rank = outputShape.length;
|
|
const channels = getChannels("rc", rank);
|
|
const dtype = getCoordsDataType(rank);
|
|
const sourceCoords = getSourceCoords$2(rank, channels);
|
|
const innerDims = channels.slice(-2);
|
|
const coords2 = rank <= 1 ? "rc" : `vec2(${innerDims.join(",")})`;
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} rc = getOutputCoords();
|
|
vec4 packedInput = getA(${sourceCoords});
|
|
|
|
setOutput(getChannel(packedInput, ${coords2}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const whereImpl$1 = whereImpl$2;
|
|
const EPSILON_FLOAT32 = 1e-7;
|
|
const EPSILON_FLOAT16 = 1e-4;
|
|
const binaryCaches = {};
|
|
function getBinaryCache(webGLVersion) {
|
|
if (webGLVersion in binaryCaches) {
|
|
return binaryCaches[webGLVersion];
|
|
}
|
|
binaryCaches[webGLVersion] = {};
|
|
return binaryCaches[webGLVersion];
|
|
}
|
|
const CPU_HANDOFF_SIZE_THRESHOLD = env().getNumber("CPU_HANDOFF_SIZE_THRESHOLD");
|
|
const BEFORE_PAGING_CONSTANT = 600;
|
|
function numMBBeforeWarning() {
|
|
if (env().global.screen == null) {
|
|
return 1024;
|
|
}
|
|
return env().global.screen.height * env().global.screen.width * window.devicePixelRatio * BEFORE_PAGING_CONSTANT / 1024 / 1024;
|
|
}
|
|
class MathBackendWebGL extends KernelBackend {
|
|
nextDataId() {
|
|
return MathBackendWebGL.nextDataId++;
|
|
}
|
|
constructor(gpuResource) {
|
|
super();
|
|
this.pendingRead = /* @__PURE__ */ new WeakMap();
|
|
this.pendingDisposal = /* @__PURE__ */ new WeakSet();
|
|
this.dataRefCount = /* @__PURE__ */ new WeakMap();
|
|
this.numBytesInGPU = 0;
|
|
this.uploadWaitMs = 0;
|
|
this.downloadWaitMs = 0;
|
|
this.lastGlFlushTime = 0;
|
|
this.warnedAboutMemory = false;
|
|
this.pendingDeletes = 0;
|
|
this.disposed = false;
|
|
if (!env().getBool("HAS_WEBGL")) {
|
|
throw new Error("WebGL is not supported on this device");
|
|
}
|
|
let newGPGPU;
|
|
if (gpuResource != null) {
|
|
if (gpuResource instanceof GPGPUContext) {
|
|
newGPGPU = gpuResource;
|
|
} else {
|
|
const gl = getWebGLContext(env().getNumber("WEBGL_VERSION"), gpuResource);
|
|
newGPGPU = new GPGPUContext(gl);
|
|
}
|
|
this.binaryCache = {};
|
|
this.gpgpuCreatedLocally = false;
|
|
} else {
|
|
const gl = getWebGLContext(env().getNumber("WEBGL_VERSION"));
|
|
newGPGPU = new GPGPUContext(gl);
|
|
this.binaryCache = getBinaryCache(env().getNumber("WEBGL_VERSION"));
|
|
this.gpgpuCreatedLocally = true;
|
|
}
|
|
this.gpgpu = newGPGPU;
|
|
this.canvas = this.gpgpu.gl.canvas;
|
|
this.textureManager = new TextureManager(this.gpgpu);
|
|
this.numMBBeforeWarning = numMBBeforeWarning();
|
|
this.texData = new DataStorage(this, engine());
|
|
}
|
|
numDataIds() {
|
|
return this.texData.numDataIds() - this.pendingDeletes;
|
|
}
|
|
// Writes a new entry to the data store with a WebGL texture, and registers it
|
|
// to the texture manager.
|
|
writeTexture(texture, shape, dtype, texHeight, texWidth, channels) {
|
|
const input = this.makeTensorInfo(shape, dtype);
|
|
const inData = this.texData.get(input.dataId);
|
|
inData.isPacked = false;
|
|
inData.texture = { texture, texShape: [texHeight, texWidth] };
|
|
inData.texShape = [texHeight, texWidth];
|
|
const shapeAs3D = getShapeAs3D(shape);
|
|
const program = new EncodeMatrixProgram(shapeAs3D, false, channels);
|
|
const output = this.runWebGLProgram(program, [input], dtype, [[texHeight, texWidth]]);
|
|
output.shape = shape;
|
|
inData.texture = null;
|
|
this.disposeIntermediateTensorInfo(input);
|
|
return output.dataId;
|
|
}
|
|
write(values, shape, dtype) {
|
|
if (env().getBool("WEBGL_CHECK_NUMERICAL_PROBLEMS") || env().getBool("DEBUG")) {
|
|
this.checkNumericalProblems(values);
|
|
}
|
|
if (dtype === "complex64" && values != null) {
|
|
throw new Error(`Cannot write to a complex64 dtype. Please use tf.complex(real, imag).`);
|
|
}
|
|
const dataId = { id: this.nextDataId() };
|
|
this.texData.set(dataId, { shape, dtype, values, usage: TextureUsage.UPLOAD, refCount: 1 });
|
|
return dataId;
|
|
}
|
|
/** Return refCount of a `TensorData`. */
|
|
refCount(dataId) {
|
|
if (this.texData.has(dataId)) {
|
|
const tensorData = this.texData.get(dataId);
|
|
return tensorData.refCount;
|
|
}
|
|
return 0;
|
|
}
|
|
/** Increase refCount of a `TextureData`. */
|
|
incRef(dataId) {
|
|
const texData = this.texData.get(dataId);
|
|
texData.refCount++;
|
|
}
|
|
/** Decrease refCount of a `TextureData`. */
|
|
decRef(dataId) {
|
|
if (this.texData.has(dataId)) {
|
|
const texData = this.texData.get(dataId);
|
|
texData.refCount--;
|
|
}
|
|
}
|
|
move(dataId, values, shape, dtype, refCount) {
|
|
if (env().getBool("DEBUG")) {
|
|
this.checkNumericalProblems(values);
|
|
}
|
|
if (dtype === "complex64") {
|
|
throw new Error(`Cannot write to a complex64 dtype. Please use tf.complex(real, imag).`);
|
|
}
|
|
this.texData.set(dataId, { shape, dtype, values, usage: TextureUsage.UPLOAD, refCount });
|
|
}
|
|
disposeIntermediateTensorInfo(tensorInfo) {
|
|
this.disposeData(tensorInfo.dataId);
|
|
}
|
|
readSync(dataId) {
|
|
const texData = this.texData.get(dataId);
|
|
const { values, dtype, complexTensorInfos, slice: slice2, shape, isPacked } = texData;
|
|
if (slice2 != null) {
|
|
let program;
|
|
if (isPacked) {
|
|
program = new UnaryOpPackedProgram(shape, CLONE);
|
|
} else {
|
|
program = new UnaryOpProgram(shape, CLONE);
|
|
}
|
|
const res = this.runWebGLProgram(program, [{ dataId, shape, dtype }], dtype);
|
|
const data = this.readSync(res.dataId);
|
|
this.disposeIntermediateTensorInfo(res);
|
|
return data;
|
|
}
|
|
if (values != null) {
|
|
return this.convertAndCacheOnCPU(dataId);
|
|
}
|
|
if (dtype === "string") {
|
|
return values;
|
|
}
|
|
const shouldTimeProgram = this.activeTimers != null;
|
|
let start;
|
|
if (shouldTimeProgram) {
|
|
start = now();
|
|
}
|
|
let result;
|
|
if (dtype === "complex64") {
|
|
const realValues = this.readSync(complexTensorInfos.real.dataId);
|
|
const imagValues = this.readSync(complexTensorInfos.imag.dataId);
|
|
result = mergeRealAndImagArrays(realValues, imagValues);
|
|
} else {
|
|
result = this.getValuesFromTexture(dataId);
|
|
}
|
|
if (shouldTimeProgram) {
|
|
this.downloadWaitMs += now() - start;
|
|
}
|
|
return this.convertAndCacheOnCPU(dataId, result);
|
|
}
|
|
async read(dataId) {
|
|
if (this.pendingRead.has(dataId)) {
|
|
const subscribers2 = this.pendingRead.get(dataId);
|
|
return new Promise((resolve) => subscribers2.push(resolve));
|
|
}
|
|
const texData = this.texData.get(dataId);
|
|
const { values, shape, slice: slice2, dtype, complexTensorInfos, isPacked } = texData;
|
|
if (slice2 != null) {
|
|
let program;
|
|
if (isPacked) {
|
|
program = new UnaryOpPackedProgram(shape, CLONE);
|
|
} else {
|
|
program = new UnaryOpProgram(shape, CLONE);
|
|
}
|
|
const res = this.runWebGLProgram(program, [{ dataId, shape, dtype }], dtype);
|
|
const data = this.read(res.dataId);
|
|
this.disposeIntermediateTensorInfo(res);
|
|
return data;
|
|
}
|
|
if (values != null) {
|
|
return this.convertAndCacheOnCPU(dataId);
|
|
}
|
|
if (env().getBool("DEBUG")) {
|
|
if (!env().getBool("WEBGL_DOWNLOAD_FLOAT_ENABLED") && env().getNumber("WEBGL_VERSION") === 2) {
|
|
throw new Error(`tensor.data() with WEBGL_DOWNLOAD_FLOAT_ENABLED=false and WEBGL_VERSION=2 not yet supported.`);
|
|
}
|
|
}
|
|
let buffer2 = null;
|
|
let tmpDownloadTarget;
|
|
if (dtype !== "complex64" && env().get("WEBGL_BUFFER_SUPPORTED")) {
|
|
tmpDownloadTarget = this.decode(dataId);
|
|
const tmpData = this.texData.get(tmpDownloadTarget.dataId);
|
|
buffer2 = this.gpgpu.createBufferFromTexture(tmpData.texture.texture, ...getDenseTexShape(shape));
|
|
}
|
|
this.pendingRead.set(dataId, []);
|
|
if (dtype !== "complex64") {
|
|
await this.gpgpu.createAndWaitForFence();
|
|
}
|
|
let vals;
|
|
if (dtype === "complex64") {
|
|
const ps = await Promise.all([
|
|
this.read(complexTensorInfos.real.dataId),
|
|
this.read(complexTensorInfos.imag.dataId)
|
|
]);
|
|
const realValues = ps[0];
|
|
const imagValues = ps[1];
|
|
vals = mergeRealAndImagArrays(realValues, imagValues);
|
|
} else if (buffer2 == null) {
|
|
vals = this.getValuesFromTexture(dataId);
|
|
} else {
|
|
const size = sizeFromShape(shape);
|
|
vals = this.gpgpu.downloadFloat32MatrixFromBuffer(buffer2, size);
|
|
}
|
|
if (tmpDownloadTarget != null) {
|
|
this.disposeIntermediateTensorInfo(tmpDownloadTarget);
|
|
}
|
|
if (buffer2 != null) {
|
|
const gl = this.gpgpu.gl;
|
|
callAndCheck(gl, () => gl.deleteBuffer(buffer2));
|
|
}
|
|
const dTypeVals = this.convertAndCacheOnCPU(dataId, vals);
|
|
const subscribers = this.pendingRead.get(dataId);
|
|
this.pendingRead.delete(dataId);
|
|
subscribers.forEach((resolve) => resolve(dTypeVals));
|
|
if (this.pendingDisposal.has(dataId)) {
|
|
this.pendingDisposal.delete(dataId);
|
|
if (this.disposeData(dataId)) {
|
|
engine().removeDataId(dataId, this);
|
|
}
|
|
this.pendingDeletes--;
|
|
}
|
|
return dTypeVals;
|
|
}
|
|
/**
|
|
* Read tensor to a new texture that is densely packed for ease of use.
|
|
* @param dataId The source tensor.
|
|
* @param options
|
|
* customTexShape: Optional. If set, will use the user defined texture
|
|
* shape to create the texture.
|
|
*/
|
|
readToGPU(dataId, options = {}) {
|
|
const texData = this.texData.get(dataId);
|
|
const { values, shape, slice: slice2, dtype, isPacked, texture } = texData;
|
|
if (dtype === "complex64") {
|
|
throw new Error("Does not support reading texture for complex64 dtype.");
|
|
}
|
|
if (slice2 != null) {
|
|
let program;
|
|
if (isPacked) {
|
|
program = new UnaryOpPackedProgram(shape, CLONE);
|
|
} else {
|
|
program = new UnaryOpProgram(shape, CLONE);
|
|
}
|
|
const res = this.runWebGLProgram(program, [{ dataId, shape, dtype }], dtype);
|
|
const gpuResouorce = this.readToGPU(res, options);
|
|
this.disposeIntermediateTensorInfo(res);
|
|
return gpuResouorce;
|
|
}
|
|
if (texture == null) {
|
|
if (values != null) {
|
|
throw new Error("Data is not on GPU but on CPU.");
|
|
} else {
|
|
throw new Error("There is no data on GPU or CPU.");
|
|
}
|
|
}
|
|
const tmpTarget = this.decode(dataId, options.customTexShape);
|
|
const tensorRef = engine().makeTensorFromTensorInfo(tmpTarget);
|
|
const tmpData = this.texData.get(tmpTarget.dataId);
|
|
return Object.assign({ tensorRef }, tmpData.texture);
|
|
}
|
|
bufferSync(t2) {
|
|
const data = this.readSync(t2.dataId);
|
|
if (t2.dtype === "string") {
|
|
try {
|
|
const strings = data.map((d) => decodeString(d));
|
|
return buffer(t2.shape, t2.dtype, strings);
|
|
} catch (_a) {
|
|
throw new Error("Failed to decode encoded string bytes into utf-8");
|
|
}
|
|
}
|
|
return buffer(t2.shape, t2.dtype, data);
|
|
}
|
|
checkNumericalProblems(values) {
|
|
if (values == null) {
|
|
return;
|
|
}
|
|
for (let i = 0; i < values.length; i++) {
|
|
const num = values[i];
|
|
if (!canBeRepresented(num)) {
|
|
if (env().getBool("WEBGL_RENDER_FLOAT32_CAPABLE")) {
|
|
throw Error(`The value ${num} cannot be represented with your current settings. Consider enabling float32 rendering: 'tf.env().set('WEBGL_RENDER_FLOAT32_ENABLED', true);'`);
|
|
}
|
|
throw Error(`The value ${num} cannot be represented on this device.`);
|
|
}
|
|
}
|
|
}
|
|
getValuesFromTexture(dataId) {
|
|
const { shape, dtype, isPacked } = this.texData.get(dataId);
|
|
const size = sizeFromShape(shape);
|
|
if (env().getBool("WEBGL_DOWNLOAD_FLOAT_ENABLED")) {
|
|
const tmpTarget = this.decode(dataId);
|
|
const tmpData2 = this.texData.get(tmpTarget.dataId);
|
|
const vals2 = this.gpgpu.downloadMatrixFromPackedTexture(tmpData2.texture.texture, ...getDenseTexShape(shape)).subarray(0, size);
|
|
this.disposeIntermediateTensorInfo(tmpTarget);
|
|
return vals2;
|
|
}
|
|
const shouldUsePackedProgram = env().getBool("WEBGL_PACK") && isPacked === true;
|
|
const outputShape = shouldUsePackedProgram ? getShapeAs3D(shape) : shape;
|
|
const program = shouldUsePackedProgram ? new EncodeFloatPackedProgram(outputShape) : new EncodeFloatProgram(outputShape);
|
|
const output = this.runWebGLProgram(program, [{ shape: outputShape, dtype, dataId }], "float32");
|
|
const tmpData = this.texData.get(output.dataId);
|
|
const vals = this.gpgpu.downloadByteEncodedFloatMatrixFromOutputTexture(tmpData.texture.texture, tmpData.texShape[0], tmpData.texShape[1]).subarray(0, size);
|
|
this.disposeIntermediateTensorInfo(output);
|
|
return vals;
|
|
}
|
|
timerAvailable() {
|
|
return env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_RELIABLE") > 0;
|
|
}
|
|
time(f) {
|
|
const oldActiveTimers = this.activeTimers;
|
|
const newActiveTimers = [];
|
|
let outerMostTime = false;
|
|
if (this.programTimersStack == null) {
|
|
this.programTimersStack = newActiveTimers;
|
|
outerMostTime = true;
|
|
} else {
|
|
this.activeTimers.push(newActiveTimers);
|
|
}
|
|
this.activeTimers = newActiveTimers;
|
|
f();
|
|
const flattenedActiveTimerQueries = flatten(this.activeTimers.map((d) => d.query)).filter((d) => d != null);
|
|
const flattenedActiveTimerNames = flatten(this.activeTimers.map((d) => d.name)).filter((d) => d != null);
|
|
this.activeTimers = oldActiveTimers;
|
|
if (outerMostTime) {
|
|
this.programTimersStack = null;
|
|
}
|
|
const res = {
|
|
uploadWaitMs: this.uploadWaitMs,
|
|
downloadWaitMs: this.downloadWaitMs,
|
|
kernelMs: null,
|
|
wallMs: null
|
|
// will be filled by the engine
|
|
};
|
|
return (async () => {
|
|
if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_RELIABLE") > 0) {
|
|
const kernelMs = await Promise.all(flattenedActiveTimerQueries);
|
|
res["kernelMs"] = sum$3(kernelMs);
|
|
res["getExtraProfileInfo"] = () => kernelMs.map((d, i) => ({ name: flattenedActiveTimerNames[i], ms: d })).map((d) => `${d.name}: ${d.ms}`).join(", ");
|
|
} else {
|
|
res["kernelMs"] = {
|
|
error: "WebGL query timers are not supported in this environment."
|
|
};
|
|
}
|
|
this.uploadWaitMs = 0;
|
|
this.downloadWaitMs = 0;
|
|
return res;
|
|
})();
|
|
}
|
|
memory() {
|
|
return {
|
|
unreliable: false,
|
|
numBytesInGPU: this.numBytesInGPU,
|
|
numBytesInGPUAllocated: this.textureManager.numBytesAllocated,
|
|
numBytesInGPUFree: this.textureManager.numBytesFree
|
|
};
|
|
}
|
|
startTimer() {
|
|
if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_RELIABLE") > 0) {
|
|
return this.gpgpu.beginQuery();
|
|
}
|
|
return { startMs: now(), endMs: null };
|
|
}
|
|
endTimer(query) {
|
|
if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_RELIABLE") > 0) {
|
|
this.gpgpu.endQuery();
|
|
return query;
|
|
}
|
|
query.endMs = now();
|
|
return query;
|
|
}
|
|
async getQueryTime(query) {
|
|
if (env().getNumber("WEBGL_DISJOINT_QUERY_TIMER_EXTENSION_RELIABLE") > 0) {
|
|
return this.gpgpu.waitForQueryAndGetTime(query);
|
|
}
|
|
const timerQuery = query;
|
|
return timerQuery.endMs - timerQuery.startMs;
|
|
}
|
|
/**
|
|
* Decrease the RefCount on the dataId and dispose the memory if the dataId
|
|
* has 0 refCount. If there are pending read on the data, the disposal would
|
|
* added to the pending delete queue. Return true if the dataId is removed
|
|
* from backend or the backend does not contain the dataId, false if the
|
|
* dataId is not removed. Memory may or may not be released even when dataId
|
|
* is removed, which also depends on dataRefCount, see `releaseGPU`.
|
|
* @param dataId
|
|
* @oaram force Optional, remove the data regardless of refCount
|
|
*/
|
|
disposeData(dataId, force = false) {
|
|
if (this.pendingDisposal.has(dataId)) {
|
|
return false;
|
|
}
|
|
if (!this.texData.has(dataId)) {
|
|
return true;
|
|
}
|
|
if (force) {
|
|
this.texData.get(dataId).refCount = 0;
|
|
} else {
|
|
this.texData.get(dataId).refCount--;
|
|
}
|
|
if (!force && this.texData.get(dataId).refCount > 0) {
|
|
return false;
|
|
}
|
|
if (this.pendingRead.has(dataId)) {
|
|
this.pendingDisposal.add(dataId);
|
|
this.pendingDeletes++;
|
|
return false;
|
|
}
|
|
this.releaseGPUData(dataId);
|
|
const { complexTensorInfos } = this.texData.get(dataId);
|
|
if (complexTensorInfos != null) {
|
|
this.disposeData(complexTensorInfos.real.dataId, force);
|
|
this.disposeData(complexTensorInfos.imag.dataId, force);
|
|
}
|
|
this.texData.delete(dataId);
|
|
return true;
|
|
}
|
|
releaseGPUData(dataId) {
|
|
const { texture, dtype, texShape, usage, isPacked, slice: slice2 } = this.texData.get(dataId);
|
|
const key = slice2 && slice2.origDataId || dataId;
|
|
const refCount = this.dataRefCount.get(key);
|
|
if (refCount > 1) {
|
|
this.dataRefCount.set(key, refCount - 1);
|
|
} else {
|
|
this.dataRefCount.delete(key);
|
|
if (texture != null) {
|
|
this.numBytesInGPU -= this.computeBytes(texShape, dtype);
|
|
this.textureManager.releaseTexture(texture, texShape, usage, isPacked);
|
|
}
|
|
}
|
|
const texData = this.texData.get(dataId);
|
|
texData.texture = null;
|
|
texData.texShape = null;
|
|
texData.isPacked = false;
|
|
texData.slice = null;
|
|
}
|
|
getTexture(dataId) {
|
|
this.uploadToGPU(dataId);
|
|
return this.texData.get(dataId).texture.texture;
|
|
}
|
|
/**
|
|
* Returns internal information for the specific data bucket. Used in unit
|
|
* tests.
|
|
*/
|
|
getDataInfo(dataId) {
|
|
return this.texData.get(dataId);
|
|
}
|
|
/*
|
|
Tests whether all the inputs to an op are small and on the CPU. This heuristic
|
|
determines when it would be faster to execute a kernel on the CPU. WebGL
|
|
kernels opt into running this check and forwarding when appropriate.
|
|
TODO(https://github.com/tensorflow/tfjs/issues/872): Develop a more
|
|
sustainable strategy for optimizing backend execution of ops.
|
|
*/
|
|
shouldExecuteOnCPU(inputs, sizeThreshold = CPU_HANDOFF_SIZE_THRESHOLD) {
|
|
return env().getBool("WEBGL_CPU_FORWARD") && inputs.every((input) => this.texData.get(input.dataId).texture == null && sizeFromShape(input.shape) < sizeThreshold);
|
|
}
|
|
getGPGPUContext() {
|
|
return this.gpgpu;
|
|
}
|
|
where(condition) {
|
|
warn("tf.where() in webgl locks the UI thread. Call tf.whereAsync() instead");
|
|
const condVals = condition.dataSync();
|
|
return whereImpl$1(condition.shape, condVals);
|
|
}
|
|
packedUnaryOp(x, op2, dtype) {
|
|
const program = new UnaryOpPackedProgram(x.shape, op2);
|
|
const outInfo = this.compileAndRun(program, [x], dtype);
|
|
return engine().makeTensorFromTensorInfo(outInfo);
|
|
}
|
|
// TODO(msoulanille) remove this once the backend has been modularized
|
|
// a copy is needed here to break a circular dependency.
|
|
// Also remove the op from unary_op.
|
|
abs(x) {
|
|
if (this.shouldExecuteOnCPU([x]) && x.dtype !== "complex64") {
|
|
const outValues = simpleAbsImplCPU(this.texData.get(x.dataId).values);
|
|
return this.makeOutput(x.shape, x.dtype, outValues);
|
|
}
|
|
if (env().getBool("WEBGL_PACK_UNARY_OPERATIONS")) {
|
|
return this.packedUnaryOp(x, ABS$1, x.dtype);
|
|
}
|
|
const program = new UnaryOpProgram(x.shape, ABS$1);
|
|
const outInfo = this.compileAndRun(program, [x]);
|
|
return engine().makeTensorFromTensorInfo(outInfo);
|
|
}
|
|
makeTensorInfo(shape, dtype, values) {
|
|
let dataId;
|
|
if (dtype === "string" && values != null && values.length > 0 && isString(values[0])) {
|
|
const encodedValues = values.map((d) => encodeString(d));
|
|
dataId = this.write(encodedValues, shape, dtype);
|
|
} else {
|
|
dataId = this.write(values, shape, dtype);
|
|
}
|
|
this.texData.get(dataId).usage = null;
|
|
return { dataId, shape, dtype };
|
|
}
|
|
makeOutput(shape, dtype, values) {
|
|
return engine().makeTensorFromTensorInfo(this.makeTensorInfo(shape, dtype, values), this);
|
|
}
|
|
unpackTensor(input) {
|
|
const program = new UnpackProgram(input.shape);
|
|
return this.runWebGLProgram(program, [input], input.dtype);
|
|
}
|
|
packTensor(input) {
|
|
const program = new PackProgram(input.shape);
|
|
const preventEagerUnpackingOutput = true;
|
|
return this.runWebGLProgram(program, [input], input.dtype, null, preventEagerUnpackingOutput);
|
|
}
|
|
packedReshape(input, afterShape) {
|
|
const input3DShape = [
|
|
getBatchDim(input.shape),
|
|
...getRowsCols(input.shape)
|
|
];
|
|
const input3D = {
|
|
dtype: input.dtype,
|
|
shape: input3DShape,
|
|
dataId: input.dataId
|
|
};
|
|
const afterShapeAs3D = [
|
|
getBatchDim(afterShape),
|
|
...getRowsCols(afterShape)
|
|
];
|
|
const program = new ReshapePackedProgram(afterShapeAs3D, input3DShape);
|
|
const preventEagerUnpackingOfOutput = true;
|
|
const customValues = [input3DShape];
|
|
const output = this.runWebGLProgram(program, [input3D], input.dtype, customValues, preventEagerUnpackingOfOutput);
|
|
return { dataId: output.dataId, shape: afterShape, dtype: output.dtype };
|
|
}
|
|
decode(dataId, customTexShape) {
|
|
const texData = this.texData.get(dataId);
|
|
const { isPacked, shape, dtype } = texData;
|
|
if (customTexShape != null) {
|
|
const size = sizeFromShape(shape);
|
|
const texSize = customTexShape[0] * customTexShape[1] * 4;
|
|
assert(size <= texSize, () => "customTexShape is too small. Row * Column * 4 should be equal or larger than the size of the tensor data.");
|
|
}
|
|
const shapeAs3D = getShapeAs3D(shape);
|
|
let program;
|
|
if (isPacked) {
|
|
program = new DecodeMatrixPackedProgram(shapeAs3D);
|
|
} else {
|
|
program = new DecodeMatrixProgram(shapeAs3D);
|
|
}
|
|
const preventEagerUnpackingOfOutput = true;
|
|
const customValues = [customTexShape != null ? customTexShape : getDenseTexShape(shapeAs3D)];
|
|
const out = this.runWebGLProgram(program, [{ shape: shapeAs3D, dtype, dataId }], dtype, customValues, preventEagerUnpackingOfOutput, customTexShape);
|
|
return { dtype, shape, dataId: out.dataId };
|
|
}
|
|
runWebGLProgram(program, inputs, outputDtype, customUniformValues, preventEagerUnpackingOfOutput = false, customTexShape) {
|
|
const output = this.makeTensorInfo(program.outputShape, outputDtype);
|
|
const outData = this.texData.get(output.dataId);
|
|
if (program.packedOutput) {
|
|
outData.isPacked = true;
|
|
}
|
|
if (program.outPackingScheme === PackingScheme.DENSE) {
|
|
const texelShape = customTexShape != null ? customTexShape : getDenseTexShape(program.outputShape);
|
|
outData.texShape = texelShape.map((d) => d * 2);
|
|
}
|
|
if (program.outTexUsage != null) {
|
|
outData.usage = program.outTexUsage;
|
|
}
|
|
if (sizeFromShape(output.shape) === 0) {
|
|
outData.values = getTypedArrayFromDType(output.dtype, 0);
|
|
return output;
|
|
}
|
|
const dataToDispose = [];
|
|
const inputsData = inputs.map((input) => {
|
|
if (input.dtype === "complex64") {
|
|
throw new Error(`GPGPUProgram does not support complex64 input. For complex64 dtypes, please separate the program into real and imaginary parts.`);
|
|
}
|
|
let texData = this.texData.get(input.dataId);
|
|
if (texData.texture == null) {
|
|
if (!program.packedInputs && sizeFromShape(input.shape) <= env().getNumber("WEBGL_SIZE_UPLOAD_UNIFORM")) {
|
|
return {
|
|
shape: input.shape,
|
|
texData: null,
|
|
isUniform: true,
|
|
uniformValues: texData.values
|
|
};
|
|
}
|
|
if (program.packedInputs) {
|
|
texData.isPacked = true;
|
|
texData.shape = input.shape;
|
|
}
|
|
}
|
|
this.uploadToGPU(input.dataId);
|
|
if (!!texData.isPacked !== !!program.packedInputs) {
|
|
input = texData.isPacked ? this.unpackTensor(input) : this.packTensor(input);
|
|
dataToDispose.push(input);
|
|
texData = this.texData.get(input.dataId);
|
|
} else if (texData.isPacked && !isReshapeFree(texData.shape, input.shape)) {
|
|
const savedInput = input;
|
|
const targetShape = input.shape;
|
|
input.shape = texData.shape;
|
|
input = this.packedReshape(input, targetShape);
|
|
dataToDispose.push(input);
|
|
texData = this.texData.get(input.dataId);
|
|
savedInput.shape = targetShape;
|
|
}
|
|
return { shape: input.shape, texData, isUniform: false };
|
|
});
|
|
this.uploadToGPU(output.dataId);
|
|
const outputData = { shape: output.shape, texData: outData, isUniform: false };
|
|
const key = makeShaderKey(program, inputsData, outputData);
|
|
const binary = this.getAndSaveBinary(key, () => {
|
|
return compileProgram(this.gpgpu, program, inputsData, outputData);
|
|
});
|
|
const shouldTimeProgram = this.activeTimers != null;
|
|
let query;
|
|
if (shouldTimeProgram) {
|
|
query = this.startTimer();
|
|
}
|
|
if (!env().get("ENGINE_COMPILE_ONLY")) {
|
|
runProgram(this.gpgpu, binary, inputsData, outputData, customUniformValues);
|
|
}
|
|
dataToDispose.forEach((info) => this.disposeIntermediateTensorInfo(info));
|
|
if (shouldTimeProgram) {
|
|
query = this.endTimer(query);
|
|
this.activeTimers.push({ name: program.constructor.name, query: this.getQueryTime(query) });
|
|
}
|
|
const glFlushThreshold = env().getNumber("WEBGL_FLUSH_THRESHOLD");
|
|
if (glFlushThreshold > 0) {
|
|
const time = now();
|
|
if (time - this.lastGlFlushTime > glFlushThreshold) {
|
|
this.gpgpu.gl.flush();
|
|
this.lastGlFlushTime = time;
|
|
}
|
|
}
|
|
if (!env().getBool("WEBGL_LAZILY_UNPACK") && outData.isPacked && preventEagerUnpackingOfOutput === false) {
|
|
const unpacked = this.unpackTensor(output);
|
|
this.disposeIntermediateTensorInfo(output);
|
|
return unpacked;
|
|
}
|
|
return output;
|
|
}
|
|
compileAndRun(program, inputs, outputDtype, customUniformValues, preventEagerUnpackingOfOutput = false) {
|
|
outputDtype = outputDtype || inputs[0].dtype;
|
|
const outInfo = this.runWebGLProgram(program, inputs, outputDtype, customUniformValues, preventEagerUnpackingOfOutput);
|
|
return outInfo;
|
|
}
|
|
getAndSaveBinary(key, getBinary) {
|
|
if (!(key in this.binaryCache)) {
|
|
this.binaryCache[key] = getBinary();
|
|
}
|
|
return this.binaryCache[key];
|
|
}
|
|
getTextureManager() {
|
|
return this.textureManager;
|
|
}
|
|
dispose() {
|
|
if (this.disposed) {
|
|
return;
|
|
}
|
|
if (!env().getBool("IS_TEST")) {
|
|
const allKeys = Object.keys(this.binaryCache);
|
|
allKeys.forEach((key) => {
|
|
this.gpgpu.deleteProgram(this.binaryCache[key].webGLProgram);
|
|
delete this.binaryCache[key];
|
|
});
|
|
}
|
|
this.textureManager.dispose();
|
|
if (this.canvas != null && (typeof HTMLCanvasElement !== "undefined" && this.canvas instanceof HTMLCanvasElement)) {
|
|
this.canvas.remove();
|
|
} else {
|
|
this.canvas = null;
|
|
}
|
|
if (this.gpgpuCreatedLocally) {
|
|
this.gpgpu.program = null;
|
|
this.gpgpu.dispose();
|
|
}
|
|
this.disposed = true;
|
|
}
|
|
floatPrecision() {
|
|
if (this.floatPrecisionValue == null) {
|
|
this.floatPrecisionValue = tidy(() => {
|
|
if (!env().get("WEBGL_RENDER_FLOAT32_ENABLED")) {
|
|
const debugFlag = env().getBool("DEBUG");
|
|
env().set("DEBUG", false);
|
|
const underflowCheckValue = this.abs(scalar(1e-8)).dataSync()[0];
|
|
env().set("DEBUG", debugFlag);
|
|
if (underflowCheckValue > 0) {
|
|
return 32;
|
|
}
|
|
}
|
|
return 16;
|
|
});
|
|
}
|
|
return this.floatPrecisionValue;
|
|
}
|
|
/** Returns the smallest representable number. */
|
|
epsilon() {
|
|
return this.floatPrecision() === 32 ? EPSILON_FLOAT32 : EPSILON_FLOAT16;
|
|
}
|
|
uploadToGPU(dataId) {
|
|
const texData = this.texData.get(dataId);
|
|
const { shape, dtype, values, texture, usage, isPacked } = texData;
|
|
if (texture != null) {
|
|
return;
|
|
}
|
|
const shouldTimeProgram = this.activeTimers != null;
|
|
let start;
|
|
if (shouldTimeProgram) {
|
|
start = now();
|
|
}
|
|
let texShape = texData.texShape;
|
|
if (texShape == null) {
|
|
texShape = getTextureShapeFromLogicalShape(shape, isPacked);
|
|
texData.texShape = texShape;
|
|
}
|
|
if (values != null) {
|
|
const shapeAs3D = getShapeAs3D(shape);
|
|
let program;
|
|
let width = texShape[1], height = texShape[0];
|
|
const isByteArray = values instanceof Uint8Array || values instanceof Uint8ClampedArray;
|
|
if (isPacked || !isByteArray) {
|
|
[width, height] = getPackedMatrixTextureShapeWidthHeight(texShape[0], texShape[1]);
|
|
}
|
|
if (isPacked) {
|
|
program = new EncodeMatrixPackedProgram(shapeAs3D, isByteArray);
|
|
} else {
|
|
program = new EncodeMatrixProgram(shapeAs3D, isByteArray);
|
|
}
|
|
const tempDenseInputTexShape = isByteArray ? [height, width] : texShape;
|
|
const tempDenseInputHandle = this.makeTensorInfo(tempDenseInputTexShape, dtype);
|
|
const tempDenseInputTexData = this.texData.get(tempDenseInputHandle.dataId);
|
|
if (isByteArray) {
|
|
tempDenseInputTexData.usage = TextureUsage.PIXELS;
|
|
} else {
|
|
tempDenseInputTexData.usage = TextureUsage.UPLOAD;
|
|
}
|
|
tempDenseInputTexData.texShape = tempDenseInputTexShape;
|
|
this.gpgpu.uploadDenseMatrixToTexture(this.getTexture(tempDenseInputHandle.dataId), width, height, values);
|
|
const customValues = [[height, width]];
|
|
const preventEagerUnpacking = true;
|
|
const encodedOutputTarget = this.runWebGLProgram(program, [tempDenseInputHandle], dtype, customValues, preventEagerUnpacking);
|
|
const outputTexData = this.texData.get(encodedOutputTarget.dataId);
|
|
texData.texShape = outputTexData.texShape;
|
|
texData.isPacked = outputTexData.isPacked;
|
|
texData.usage = outputTexData.usage;
|
|
if (!env().get("ENGINE_COMPILE_ONLY")) {
|
|
texData.texture = outputTexData.texture;
|
|
texData.values = null;
|
|
this.texData.delete(encodedOutputTarget.dataId);
|
|
} else {
|
|
this.disposeData(encodedOutputTarget.dataId);
|
|
}
|
|
this.disposeIntermediateTensorInfo(tempDenseInputHandle);
|
|
if (shouldTimeProgram) {
|
|
this.uploadWaitMs += now() - start;
|
|
}
|
|
} else {
|
|
const newTexture = this.acquireTexture(texShape, usage, dtype, isPacked);
|
|
texData.texture = newTexture;
|
|
}
|
|
}
|
|
convertAndCacheOnCPU(dataId, float32Values) {
|
|
const texData = this.texData.get(dataId);
|
|
const { dtype } = texData;
|
|
if (float32Values != null) {
|
|
texData.values = float32ToTypedArray(float32Values, dtype);
|
|
}
|
|
return texData.values;
|
|
}
|
|
acquireTexture(texShape, texType, dtype, isPacked) {
|
|
this.numBytesInGPU += this.computeBytes(texShape, dtype);
|
|
if (!this.warnedAboutMemory && this.numBytesInGPU > this.numMBBeforeWarning * 1024 * 1024) {
|
|
const mb = (this.numBytesInGPU / 1024 / 1024).toFixed(2);
|
|
this.warnedAboutMemory = true;
|
|
console.warn(`High memory usage in GPU: ${mb} MB, most likely due to a memory leak`);
|
|
}
|
|
return this.textureManager.acquireTexture(texShape, texType, isPacked);
|
|
}
|
|
computeBytes(shape, dtype) {
|
|
return shape[0] * shape[1] * bytesPerElement(dtype);
|
|
}
|
|
checkCompileCompletion() {
|
|
for (const [, binary] of Object.entries(this.binaryCache)) {
|
|
this.checkCompletion_(binary);
|
|
}
|
|
}
|
|
async checkCompileCompletionAsync() {
|
|
const ps = [];
|
|
if (this.gpgpu.parallelCompilationExtension) {
|
|
for (const [, binary] of Object.entries(this.binaryCache)) {
|
|
ps.push(this.checkCompletionAsync_(binary));
|
|
}
|
|
return Promise.all(ps);
|
|
} else {
|
|
for (const [, binary] of Object.entries(this.binaryCache)) {
|
|
const p2 = new Promise((resolve) => {
|
|
try {
|
|
this.checkCompletion_(binary);
|
|
resolve(true);
|
|
} catch (error) {
|
|
throw error;
|
|
}
|
|
});
|
|
ps.push(p2);
|
|
}
|
|
return Promise.all(ps);
|
|
}
|
|
}
|
|
async checkCompletionAsync_(binary) {
|
|
if (this.gpgpu.gl.getProgramParameter(binary.webGLProgram, this.gpgpu.parallelCompilationExtension.COMPLETION_STATUS_KHR)) {
|
|
return this.checkCompletion_(binary);
|
|
} else {
|
|
await nextFrame();
|
|
return this.checkCompletionAsync_(binary);
|
|
}
|
|
}
|
|
checkCompletion_(binary) {
|
|
if (this.gpgpu.gl.getProgramParameter(binary.webGLProgram, this.gpgpu.gl.LINK_STATUS) === false) {
|
|
console.log(this.gpgpu.gl.getProgramInfoLog(binary.webGLProgram));
|
|
if (this.gpgpu.gl.getShaderParameter(binary.fragmentShader, this.gpgpu.gl.COMPILE_STATUS) === false) {
|
|
logShaderSourceAndInfoLog(binary.source, this.gpgpu.gl.getShaderInfoLog(binary.fragmentShader));
|
|
throw new Error("Failed to compile fragment shader.");
|
|
}
|
|
throw new Error("Failed to link vertex and fragment shaders.");
|
|
}
|
|
return true;
|
|
}
|
|
getUniformLocations() {
|
|
for (const binary of Object.values(this.binaryCache)) {
|
|
this.gpgpu.buildVao(binary.webGLProgram);
|
|
const { variablesLocations, customUniformLocations, infLoc, nanLoc, outShapeLocation, outShapeStridesLocation, outTexShapeLocation } = getUniformLocations(this.gpgpu, binary.program, binary.webGLProgram);
|
|
binary.variablesLocations = variablesLocations;
|
|
binary.customUniformLocations = customUniformLocations;
|
|
binary.infLoc = infLoc;
|
|
binary.nanLoc = nanLoc;
|
|
binary.outShapeLocation = outShapeLocation;
|
|
binary.outShapeStridesLocation = outShapeStridesLocation;
|
|
binary.outTexShapeLocation = outTexShapeLocation;
|
|
}
|
|
}
|
|
/**
|
|
* Create a TF.js tensor out of an existing WebGL texture. A new texture will
|
|
* be created.
|
|
*/
|
|
createTensorFromGPUData(values, shape, dtype) {
|
|
values.channels = values.channels || "RGBA";
|
|
const { texture, height, width, channels } = values;
|
|
const backend2 = engine().backend;
|
|
if (!backend2.gpgpu.gl.isTexture(texture)) {
|
|
throw new Error(`The texture is invalid. Also, please make sure the texture and the TFJS WebGL backend are using the same canvas. If you want to use your own custom canvas, you have to create and use the custom TFJS WebGL backend created from the canvas through 'new tf.MathBackendWebGL(customCanvas)'.`);
|
|
}
|
|
const dataId = backend2.writeTexture(texture, shape, dtype, height, width, channels);
|
|
return engine().makeTensorFromDataId(dataId, shape, dtype, backend2);
|
|
}
|
|
}
|
|
MathBackendWebGL.nextDataId = 0;
|
|
function float32ToTypedArray(a, dtype) {
|
|
if (dtype === "float32" || dtype === "complex64") {
|
|
return a;
|
|
} else if (dtype === "int32" || dtype === "bool") {
|
|
const result = dtype === "int32" ? new Int32Array(a.length) : new Uint8Array(a.length);
|
|
for (let i = 0; i < result.length; ++i) {
|
|
result[i] = Math.round(a[i]);
|
|
}
|
|
return result;
|
|
} else {
|
|
throw new Error(`Unknown dtype ${dtype}`);
|
|
}
|
|
}
|
|
if (isBrowser()) {
|
|
registerBackend(
|
|
"webgl",
|
|
() => new MathBackendWebGL(),
|
|
2
|
|
/* priority */
|
|
);
|
|
}
|
|
const CHECK_NAN_SNIPPET = `
|
|
if (isnan(a)) return a;
|
|
if (isnan(b)) return b;
|
|
`;
|
|
class BinaryOpProgram {
|
|
constructor(op2, aShape, bShape) {
|
|
this.variableNames = ["A", "B"];
|
|
this.outputShape = assertAndGetBroadcastShape(aShape, bShape);
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
this.userCode = `
|
|
float binaryOperation(float a, float b) {
|
|
${op2}
|
|
}
|
|
|
|
void main() {
|
|
float a = getAAtOutCoords();
|
|
float b = getBAtOutCoords();
|
|
setOutput(binaryOperation(a, b));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const CHECK_NAN_SNIPPET_PACKED = `
|
|
result.r = isNaN.r ? NAN : result.r;
|
|
result.g = isNaN.g ? NAN : result.g;
|
|
result.b = isNaN.b ? NAN : result.b;
|
|
result.a = isNaN.a ? NAN : result.a;
|
|
`;
|
|
class BinaryOpPackedProgram {
|
|
constructor(op2, aShape, bShape, checkOutOfBounds = false) {
|
|
this.variableNames = ["A", "B"];
|
|
this.supportsBroadcasting = true;
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = assertAndGetBroadcastShape(aShape, bShape);
|
|
const rank = this.outputShape.length;
|
|
this.enableShapeUniforms = useShapeUniforms(rank);
|
|
let checkOutOfBoundsString = "";
|
|
if (checkOutOfBounds) {
|
|
if (rank === 0 || sizeFromShape(this.outputShape) === 1) {
|
|
checkOutOfBoundsString = `
|
|
result.y = 0.;
|
|
result.z = 0.;
|
|
result.w = 0.;
|
|
`;
|
|
} else {
|
|
const dtype = getCoordsDataType(rank);
|
|
checkOutOfBoundsString = `
|
|
${dtype} coords = getOutputCoords();
|
|
`;
|
|
if (rank === 1) {
|
|
if (this.enableShapeUniforms) {
|
|
checkOutOfBoundsString += `
|
|
result.y = (coords + 1) >= outShape ? 0. : result.y;
|
|
result.z = 0.;
|
|
result.w = 0.;
|
|
`;
|
|
} else {
|
|
checkOutOfBoundsString += `
|
|
result.y = (coords + 1) >= ${this.outputShape[0]} ? 0. : result.y;
|
|
result.z = 0.;
|
|
result.w = 0.;
|
|
`;
|
|
}
|
|
} else {
|
|
const channels = getChannels("coords", rank);
|
|
if (this.enableShapeUniforms) {
|
|
checkOutOfBoundsString += `
|
|
bool nextRowOutOfBounds =
|
|
(${channels[rank - 2]} + 1) >= outShape[${rank} - 2];
|
|
bool nextColOutOfBounds =
|
|
(${channels[rank - 1]} + 1) >= outShape[${rank} - 1];
|
|
result.y = nextColOutOfBounds ? 0. : result.y;
|
|
result.z = nextRowOutOfBounds ? 0. : result.z;
|
|
result.w = nextColOutOfBounds || nextRowOutOfBounds ? 0. : result.w;
|
|
`;
|
|
} else {
|
|
checkOutOfBoundsString += `
|
|
bool nextRowOutOfBounds =
|
|
(${channels[rank - 2]} + 1) >= ${this.outputShape[rank - 2]};
|
|
bool nextColOutOfBounds =
|
|
(${channels[rank - 1]} + 1) >= ${this.outputShape[rank - 1]};
|
|
result.y = nextColOutOfBounds ? 0. : result.y;
|
|
result.z = nextRowOutOfBounds ? 0. : result.z;
|
|
result.w = nextColOutOfBounds || nextRowOutOfBounds ? 0. : result.w;
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
this.userCode = `
|
|
vec4 binaryOperation(vec4 a, vec4 b) {
|
|
${op2}
|
|
}
|
|
|
|
void main() {
|
|
vec4 a = getAAtOutCoords();
|
|
vec4 b = getBAtOutCoords();
|
|
|
|
vec4 result = binaryOperation(a, b);
|
|
${checkOutOfBoundsString}
|
|
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function identity(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
backend2.incRef(x.dataId);
|
|
return { dataId: x.dataId, shape: x.shape, dtype: x.dtype };
|
|
}
|
|
const identityConfig = {
|
|
kernelName: Identity,
|
|
backendName: "webgl",
|
|
kernelFunc: identity
|
|
};
|
|
function complex(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { real: real2, imag: imag2 } = inputs;
|
|
const complexInfo = backend2.makeTensorInfo(real2.shape, "complex64");
|
|
const complex2 = backend2.texData.get(complexInfo.dataId);
|
|
const realTensorInfo = identity({ inputs: { x: real2 }, backend: backend2 });
|
|
const imagTensorInfo = identity({ inputs: { x: imag2 }, backend: backend2 });
|
|
complex2.complexTensorInfos = { real: realTensorInfo, imag: imagTensorInfo };
|
|
return complexInfo;
|
|
}
|
|
const complexConfig = {
|
|
kernelName: Complex,
|
|
backendName: "webgl",
|
|
kernelFunc: complex
|
|
};
|
|
const LEAKYRELU = `return (a < 0.) ? b * a : a;`;
|
|
const LEAKYRELU_PACKED = `
|
|
vec4 aLessThanZero = vec4(lessThan(a, vec4(0.)));
|
|
return (aLessThanZero * (b * a)) + ((vec4(1.0) - aLessThanZero) * a);
|
|
`;
|
|
function leakyRelu$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { alpha } = attrs;
|
|
const $alpha = backend2.makeTensorInfo([], "float32", createScalarValue(alpha, "float32"));
|
|
const program = env().getBool("WEBGL_PACK_BINARY_OPERATIONS") ? new BinaryOpPackedProgram(LEAKYRELU_PACKED, x.shape, $alpha.shape) : new BinaryOpProgram(LEAKYRELU, x.shape, $alpha.shape);
|
|
const result = backend2.runWebGLProgram(program, [x, $alpha], "float32");
|
|
backend2.disposeIntermediateTensorInfo($alpha);
|
|
return result;
|
|
}
|
|
const leakyReluConfig$1 = {
|
|
kernelName: LeakyRelu,
|
|
backendName: "webgl",
|
|
kernelFunc: leakyRelu$1
|
|
};
|
|
const PRELU = `return (a < 0.) ? b * a : a;`;
|
|
const PRELU_PACKED = `
|
|
vec4 aLessThanZero = vec4(lessThan(a, vec4(0.)));
|
|
return (aLessThanZero * (b * a)) + ((vec4(1.0) - aLessThanZero) * a);
|
|
`;
|
|
function prelu$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x, alpha } = inputs;
|
|
const program = env().getBool("WEBGL_PACK_BINARY_OPERATIONS") ? new BinaryOpPackedProgram(PRELU_PACKED, x.shape, alpha.shape) : new BinaryOpProgram(PRELU, x.shape, alpha.shape);
|
|
return backend2.runWebGLProgram(program, [x, alpha], "float32");
|
|
}
|
|
const preluConfig$1 = {
|
|
kernelName: Prelu,
|
|
backendName: "webgl",
|
|
kernelFunc: prelu$1
|
|
};
|
|
const CHECK_NAN_SNIPPET_UNARY = `if (isnan(x)) return x;`;
|
|
function unaryKernelFunc({ opSnippet, packedOpSnippet, cpuKernelImpl, dtype }) {
|
|
return ({ inputs, backend: backend2 }) => {
|
|
const { x } = inputs;
|
|
const webglBackend = backend2;
|
|
const $dtype = dtype || x.dtype;
|
|
if (webglBackend.shouldExecuteOnCPU([x]) && cpuKernelImpl != null) {
|
|
const xData = webglBackend.texData.get(x.dataId);
|
|
const outValues = cpuKernelImpl(xData.values, $dtype);
|
|
return webglBackend.makeTensorInfo(x.shape, $dtype, outValues);
|
|
}
|
|
const shouldUsePackedProgram = env().getBool("WEBGL_PACK_UNARY_OPERATIONS") && packedOpSnippet != null;
|
|
let program;
|
|
if (shouldUsePackedProgram) {
|
|
program = new UnaryOpPackedProgram(x.shape, packedOpSnippet);
|
|
} else {
|
|
program = new UnaryOpProgram(x.shape, opSnippet);
|
|
}
|
|
return webglBackend.runWebGLProgram(program, [x], $dtype);
|
|
};
|
|
}
|
|
function binaryKernelFunc({ opSnippet, packedOpSnippet, checkOutOfBounds = false, supportsComplex = false, cpuKernelImpl, dtype }) {
|
|
return ({ inputs, backend: backend2 }) => {
|
|
const { a, b } = inputs;
|
|
const webglBackend = backend2;
|
|
if (supportsComplex && a.dtype === "complex64") {
|
|
const aData = webglBackend.texData.get(a.dataId);
|
|
const bData = webglBackend.texData.get(b.dataId);
|
|
const [real2, imag2] = [
|
|
[aData.complexTensorInfos.real, bData.complexTensorInfos.real],
|
|
[aData.complexTensorInfos.imag, bData.complexTensorInfos.imag]
|
|
].map((complexParts) => {
|
|
const [aPart, bPart] = complexParts;
|
|
const aHandle = {
|
|
dataId: aPart.dataId,
|
|
dtype: aPart.dtype,
|
|
shape: a.shape
|
|
};
|
|
const bHandle = {
|
|
dataId: bPart.dataId,
|
|
dtype: bPart.dtype,
|
|
shape: b.shape
|
|
};
|
|
const program2 = new BinaryOpProgram(opSnippet, a.shape, b.shape);
|
|
return webglBackend.runWebGLProgram(program2, [aHandle, bHandle], upcastType(aPart.dtype, bPart.dtype));
|
|
});
|
|
const complexOutput = complex({ inputs: { real: real2, imag: imag2 }, backend: webglBackend });
|
|
webglBackend.disposeIntermediateTensorInfo(real2);
|
|
webglBackend.disposeIntermediateTensorInfo(imag2);
|
|
return complexOutput;
|
|
}
|
|
const $dtype = dtype || upcastType(a.dtype, b.dtype);
|
|
if ((a.dtype === "string" || b.dtype === "string" || webglBackend.shouldExecuteOnCPU([a, b])) && cpuKernelImpl != null) {
|
|
const aVals = webglBackend.texData.get(a.dataId).values;
|
|
const bVals = webglBackend.texData.get(b.dataId).values;
|
|
const decodedAVals = a.dtype === "string" ? (
|
|
// tslint:disable-next-line: no-any
|
|
fromUint8ToStringArray(aVals)
|
|
) : aVals;
|
|
const decodedBVals = a.dtype === "string" ? (
|
|
// tslint:disable-next-line: no-any
|
|
fromUint8ToStringArray(bVals)
|
|
) : bVals;
|
|
const [outValues, outShape] = cpuKernelImpl(a.shape, b.shape, decodedAVals, decodedBVals, $dtype);
|
|
const out = webglBackend.makeTensorInfo(outShape, $dtype);
|
|
const outData = webglBackend.texData.get(out.dataId);
|
|
outData.values = outValues;
|
|
return out;
|
|
}
|
|
const shouldUsePackedProgram = env().getBool("WEBGL_PACK_BINARY_OPERATIONS") && packedOpSnippet != null;
|
|
let program;
|
|
if (shouldUsePackedProgram) {
|
|
program = new BinaryOpPackedProgram(packedOpSnippet, a.shape, b.shape, checkOutOfBounds);
|
|
} else {
|
|
program = new BinaryOpProgram(opSnippet, a.shape, b.shape);
|
|
}
|
|
return webglBackend.runWebGLProgram(program, [a, b], $dtype);
|
|
};
|
|
}
|
|
function mapActivationToShaderProgram(activation, packed = false) {
|
|
if (activation === "linear") {
|
|
if (packed) {
|
|
return LINEAR;
|
|
}
|
|
return LINEAR$1;
|
|
} else if (activation === "relu") {
|
|
if (packed) {
|
|
return RELU$1;
|
|
}
|
|
return RELU$2;
|
|
} else if (activation === "elu") {
|
|
if (packed) {
|
|
return ELU$1;
|
|
}
|
|
return ELU$2;
|
|
} else if (activation === "relu6") {
|
|
if (packed) {
|
|
return RELU6$1;
|
|
}
|
|
return RELU6$2;
|
|
} else if (activation === "prelu") {
|
|
if (packed) {
|
|
return PRELU_PACKED;
|
|
}
|
|
return PRELU;
|
|
} else if (activation === "leakyrelu") {
|
|
if (packed) {
|
|
return LEAKYRELU_PACKED;
|
|
}
|
|
return LEAKYRELU;
|
|
} else if (activation === "sigmoid") {
|
|
if (packed) {
|
|
return SIGMOID$1;
|
|
}
|
|
return SIGMOID$2;
|
|
}
|
|
throw new Error(`Activation ${activation} has not been implemented for the WebGL backend.`);
|
|
}
|
|
class MatMulPackedProgram {
|
|
constructor(aShape, bShape, outputShape, transposeA = false, transposeB = false, addBias = false, activation = null, hasPreluActivation = false, hasLeakyreluActivation = false) {
|
|
this.variableNames = ["matrixA", "matrixB"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const sharedDim = transposeA ? aShape[1] : aShape[2];
|
|
const sharedDimensionPacked = Math.ceil(sharedDim / 2);
|
|
const aSample = transposeA ? "i * 2, rc.y" : "rc.y, i * 2";
|
|
const bSample = transposeB ? "rc.z, i * 2" : "i * 2, rc.z";
|
|
const aSwizzle = transposeA ? ["a.xxyy", "a.zzww"] : ["a.xxzz", "a.yyww"];
|
|
const bSwizzle = transposeB ? ["b.xzxz", "b.ywyw"] : ["b.xyxy", "b.zwzw"];
|
|
let activationSnippet = "", applyActivationSnippet = "";
|
|
if (activation) {
|
|
if (hasPreluActivation) {
|
|
activationSnippet = `vec4 activation(vec4 a) {
|
|
vec4 b = getPreluActivationWeightsAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else if (hasLeakyreluActivation) {
|
|
activationSnippet = `vec4 activation(vec4 a) {
|
|
vec4 b = getLeakyreluAlphaAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else {
|
|
activationSnippet = `vec4 activation(vec4 x) {
|
|
${activation}
|
|
}`;
|
|
}
|
|
applyActivationSnippet = `result = activation(result);`;
|
|
}
|
|
const addBiasSnippet = addBias ? "result += getBiasAtOutCoords();" : "";
|
|
if (addBias) {
|
|
this.variableNames.push("bias");
|
|
}
|
|
if (hasPreluActivation) {
|
|
this.variableNames.push("preluActivationWeights");
|
|
}
|
|
if (hasLeakyreluActivation) {
|
|
this.variableNames.push("leakyreluAlpha");
|
|
}
|
|
let batchASnippet = "rc.x";
|
|
let batchBSnippet = "rc.x";
|
|
if (aShape[0] < bShape[0]) {
|
|
batchASnippet = `imod(rc.x, ${aShape[0]})`;
|
|
} else if (bShape[0] < aShape[0]) {
|
|
batchBSnippet = `imod(rc.x, ${bShape[0]})`;
|
|
}
|
|
this.userCode = `
|
|
${activationSnippet}
|
|
// Don't use uniform for sharedDimensionPacked for performance.
|
|
const float sharedDimension = ${sharedDimensionPacked}.0;
|
|
|
|
vec4 dot2x2ARowBCol(ivec3 rc) {
|
|
vec4 result = vec4(0);
|
|
int batchA = ${batchASnippet};
|
|
int batchB = ${batchBSnippet};
|
|
for (int i = 0; i < ${sharedDimensionPacked}; i++) {
|
|
vec4 a = getMatrixA(batchA, ${aSample});
|
|
vec4 b = getMatrixB(batchB, ${bSample});
|
|
|
|
// These swizzled products need to be separately added.
|
|
// See: https://github.com/tensorflow/tfjs/issues/1735
|
|
result += (${aSwizzle[0]} * ${bSwizzle[0]});
|
|
result += (${aSwizzle[1]} * ${bSwizzle[1]});
|
|
}
|
|
return result;
|
|
}
|
|
|
|
void main() {
|
|
ivec3 rc = getOutputCoords();
|
|
vec4 result = dot2x2ARowBCol(rc);
|
|
|
|
${addBiasSnippet}
|
|
|
|
${applyActivationSnippet}
|
|
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const COMPLEX_MULTIPLY = {
|
|
REAL: "return areal * breal - aimag * bimag;",
|
|
IMAG: "return areal * bimag + aimag * breal;"
|
|
};
|
|
class BinaryOpComplexProgram {
|
|
constructor(op2, aShape, bShape) {
|
|
this.variableNames = ["AReal", "AImag", "BReal", "BImag"];
|
|
this.outputShape = assertAndGetBroadcastShape(aShape, bShape);
|
|
this.userCode = `
|
|
float binaryOpComplex(
|
|
float areal, float aimag, float breal, float bimag) {
|
|
${op2}
|
|
}
|
|
|
|
void main() {
|
|
float areal = getARealAtOutCoords();
|
|
float aimag = getAImagAtOutCoords();
|
|
float breal = getBRealAtOutCoords();
|
|
float bimag = getBImagAtOutCoords();
|
|
setOutput(binaryOpComplex(areal, aimag, breal, bimag));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const MUL = "return a * b;";
|
|
function multiply(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { a, b } = inputs;
|
|
const dtype = upcastType(a.dtype, b.dtype);
|
|
if (a.dtype === "complex64") {
|
|
const aData = backend2.texData.get(a.dataId);
|
|
const bData = backend2.texData.get(b.dataId);
|
|
const realProgram = new BinaryOpComplexProgram(COMPLEX_MULTIPLY.REAL, a.shape, b.shape);
|
|
const imagProgram = new BinaryOpComplexProgram(COMPLEX_MULTIPLY.IMAG, a.shape, b.shape);
|
|
const inputs2 = [
|
|
{
|
|
dataId: aData.complexTensorInfos.real.dataId,
|
|
dtype: aData.complexTensorInfos.real.dtype,
|
|
shape: a.shape
|
|
},
|
|
{
|
|
dataId: aData.complexTensorInfos.imag.dataId,
|
|
dtype: aData.complexTensorInfos.imag.dtype,
|
|
shape: a.shape
|
|
},
|
|
{
|
|
dataId: bData.complexTensorInfos.real.dataId,
|
|
dtype: bData.complexTensorInfos.real.dtype,
|
|
shape: b.shape
|
|
},
|
|
{
|
|
dataId: bData.complexTensorInfos.imag.dataId,
|
|
dtype: bData.complexTensorInfos.imag.dtype,
|
|
shape: b.shape
|
|
}
|
|
];
|
|
const realPart = backend2.runWebGLProgram(realProgram, inputs2, "float32");
|
|
const imagPart = backend2.runWebGLProgram(imagProgram, inputs2, "float32");
|
|
const complexOutput = complex({ inputs: { real: realPart, imag: imagPart }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
backend2.disposeIntermediateTensorInfo(imagPart);
|
|
return complexOutput;
|
|
}
|
|
if (backend2.shouldExecuteOnCPU([a, b])) {
|
|
const aData = backend2.texData.get(a.dataId);
|
|
const bData = backend2.texData.get(b.dataId);
|
|
const [outValues, outShape] = multiplyImplCPU(a.shape, b.shape, aData.values, bData.values, dtype);
|
|
const out = backend2.makeTensorInfo(outShape, dtype);
|
|
const outData = backend2.texData.get(out.dataId);
|
|
outData.values = outValues;
|
|
return out;
|
|
}
|
|
let program;
|
|
if (env().getBool("WEBGL_PACK_BINARY_OPERATIONS")) {
|
|
program = new BinaryOpPackedProgram(MUL, a.shape, b.shape);
|
|
} else {
|
|
program = new BinaryOpProgram(MUL, a.shape, b.shape);
|
|
}
|
|
return backend2.runWebGLProgram(program, [a, b], dtype);
|
|
}
|
|
const multiplyConfig = {
|
|
kernelName: Multiply,
|
|
backendName: "webgl",
|
|
kernelFunc: multiply
|
|
};
|
|
function packedReshape(input, afterShape, backend2) {
|
|
const input3DShape = [
|
|
getBatchDim(input.shape),
|
|
...getRowsCols(input.shape)
|
|
];
|
|
const input3D = {
|
|
dtype: input.dtype,
|
|
shape: input3DShape,
|
|
dataId: input.dataId
|
|
};
|
|
const afterShapeAs3D = [
|
|
getBatchDim(afterShape),
|
|
...getRowsCols(afterShape)
|
|
];
|
|
const program = new ReshapePackedProgram(afterShapeAs3D, input3DShape);
|
|
const preventEagerUnpackingOfOutput = true;
|
|
const customValues = [input3DShape];
|
|
const output = backend2.runWebGLProgram(program, [input3D], input.dtype, customValues, preventEagerUnpackingOfOutput);
|
|
return { dataId: output.dataId, shape: afterShape, dtype: output.dtype };
|
|
}
|
|
function reshape$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { shape } = attrs;
|
|
const webglBackend = backend2;
|
|
const xSize = sizeFromShape(x.shape);
|
|
const $shape = inferFromImplicitShape(shape, xSize);
|
|
const $xSize = sizeFromShape($shape);
|
|
assert(xSize === $xSize, () => `The new shape (${$shape}) has ${$xSize} elements and the old shape (${x.shape}) has ${xSize} elements. The new shape and old shape must have the same number of elements.`);
|
|
const xTexData = webglBackend.texData.get(x.dataId);
|
|
if (xTexData.isPacked && !isReshapeFree(x.shape, $shape) && !(xTexData.texture !== null && isReshapeFree(xTexData.shape, $shape))) {
|
|
return packedReshape(x, $shape, webglBackend);
|
|
}
|
|
webglBackend.incRef(x.dataId);
|
|
return { dataId: x.dataId, shape: $shape, dtype: x.dtype };
|
|
}
|
|
const reshapeConfig$1 = {
|
|
kernelName: Reshape,
|
|
backendName: "webgl",
|
|
kernelFunc: reshape$1
|
|
};
|
|
class MeanProgram {
|
|
constructor(reduceInfo, divisor) {
|
|
this.variableNames = ["x"];
|
|
const { windowSize, batchSize, inSize, outSize } = reduceInfo;
|
|
this.outputShape = [batchSize, outSize];
|
|
const windowSizeNearestVec4 = Math.floor(windowSize / 4) * 4;
|
|
const windowSizeVec4Remainder = windowSize % 4;
|
|
let updateSnippet = `sumValue += dot(values, ones);`;
|
|
if (divisor != null) {
|
|
const denominator = 1 / divisor;
|
|
updateSnippet = `sumValue += dot(values * ${isInt(denominator) ? denominator.toPrecision(2) : denominator}, ones);`;
|
|
}
|
|
let checkOutOfBounds = "";
|
|
if (inSize % windowSize > 0) {
|
|
checkOutOfBounds = `
|
|
if (inIdx < 0 || inIdx >= ${inSize}) {
|
|
return 0.0;
|
|
}
|
|
`;
|
|
}
|
|
this.userCode = `
|
|
const vec4 ones = vec4(1.0, 1.0, 1.0, 1.0);
|
|
|
|
float getValue(int batch, int inIdx) {
|
|
${checkOutOfBounds}
|
|
return getX(batch, inIdx);
|
|
}
|
|
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int outIdx = coords[1];
|
|
int inOffset = outIdx * ${windowSize};
|
|
|
|
float sumValue = 0.0;
|
|
|
|
for (int i = 0; i < ${windowSizeNearestVec4}; i += 4) {
|
|
int inIdx = inOffset + i;
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
getValue(batch, inIdx + 2),
|
|
getValue(batch, inIdx + 3)
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
|
|
int inIdx = inOffset + ${windowSizeNearestVec4};
|
|
if (${windowSizeVec4Remainder === 1}) {
|
|
vec4 values = vec4(getValue(batch, inIdx), 0.0, 0.0, 0.0);
|
|
|
|
${updateSnippet}
|
|
} else if (${windowSizeVec4Remainder === 2}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1), 0.0, 0.0);
|
|
|
|
${updateSnippet}
|
|
} else if (${windowSizeVec4Remainder === 3}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
getValue(batch, inIdx + 2), 0.0);
|
|
|
|
${updateSnippet}
|
|
}
|
|
setOutput(sumValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ReduceProgram {
|
|
constructor(reduceInfo, reduceType) {
|
|
this.variableNames = ["x"];
|
|
const { windowSize, batchSize, inSize, outSize } = reduceInfo;
|
|
this.outputShape = [batchSize, outSize];
|
|
let initializationValue = "0.0";
|
|
let compareOp = ``;
|
|
if (reduceType === "prod") {
|
|
initializationValue = "1.0";
|
|
} else if (reduceType === "min") {
|
|
initializationValue = "1.0 / 1e-20";
|
|
compareOp = `min`;
|
|
} else if (reduceType === "max") {
|
|
initializationValue = "-1.0 / 1e-20";
|
|
compareOp = `max`;
|
|
}
|
|
let returnValue = `${reduceType}(${reduceType}(${reduceType}(minMaxValue[0], minMaxValue[1]), minMaxValue[2]), minMaxValue[3])`;
|
|
if (reduceType === "sum") {
|
|
returnValue = `sumValue`;
|
|
} else if (reduceType === "prod") {
|
|
returnValue = `prodValue`;
|
|
} else if (reduceType === "all") {
|
|
returnValue = `allValue`;
|
|
} else if (reduceType === "any") {
|
|
returnValue = `anyValue`;
|
|
}
|
|
const windowSizeNearestVec4 = Math.floor(windowSize / 4) * 4;
|
|
const windowSizeVec4Remainder = windowSize % 4;
|
|
let updateSnippet = `
|
|
if (${reduceType === "sum"}) {
|
|
sumValue += dot(values, ones);
|
|
} else if (${reduceType === "prod"}) {
|
|
vec2 tmp = vec2(values[0], values[1]) * vec2(values[2], values[3]);
|
|
prodValue *= tmp[0] * tmp[1];
|
|
} else {
|
|
minMaxValue = ${compareOp}(values, minMaxValue);
|
|
if (${reduceType === "min"} || ${reduceType === "max"}) {
|
|
minMaxValue = ${compareOp}(values, minMaxValue);
|
|
bvec4 isNaN = isnan(values);
|
|
if (isNaN.r || isNaN.g || isNaN.b || isNaN.a) {
|
|
minMaxValue = vec4(NAN);
|
|
}
|
|
}
|
|
}
|
|
`;
|
|
let vecType = `vec4`;
|
|
if (reduceType === "all") {
|
|
initializationValue = "1.0";
|
|
updateSnippet = `
|
|
bool reducedAllValue = all(values);
|
|
float floatedReducedAllValue = float(reducedAllValue);
|
|
allValue = float(allValue >= 1.0 && floatedReducedAllValue >= 1.0);
|
|
`;
|
|
vecType = `bvec4`;
|
|
} else if (reduceType === "any") {
|
|
initializationValue = "0.0";
|
|
updateSnippet = `
|
|
bool reducedAnyValue = any(values);
|
|
float floatedReducedAnyValue = float(reducedAnyValue);
|
|
anyValue = float(anyValue >= 1.0 || floatedReducedAnyValue >= 1.0);
|
|
`;
|
|
vecType = `bvec4`;
|
|
}
|
|
let checkOutOfBounds = "";
|
|
if (inSize % windowSize > 0) {
|
|
checkOutOfBounds = `
|
|
if (inIdx < 0 || inIdx >= ${inSize}) {
|
|
return initializationValue;
|
|
}
|
|
`;
|
|
}
|
|
this.userCode = `
|
|
const float initializationValue = ${initializationValue};
|
|
const vec4 ones = vec4(1.0, 1.0, 1.0, 1.0);
|
|
|
|
float getValue(int batch, int inIdx) {
|
|
${checkOutOfBounds}
|
|
return getX(batch, inIdx);
|
|
}
|
|
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int outIdx = coords[1];
|
|
int inOffset = outIdx * ${windowSize};
|
|
|
|
vec4 minMaxValue = vec4(${initializationValue});
|
|
float prodValue = 1.0;
|
|
float sumValue = 0.0;
|
|
float allValue = 1.0;
|
|
float anyValue = 0.0;
|
|
|
|
for (int i = 0; i < ${windowSizeNearestVec4}; i += 4) {
|
|
int inIdx = inOffset + i;
|
|
${vecType} values = ${vecType}(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
getValue(batch, inIdx + 2),
|
|
getValue(batch, inIdx + 3)
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
|
|
int inIdx = inOffset + ${windowSizeNearestVec4};
|
|
if (${windowSizeVec4Remainder === 1}) {
|
|
${vecType} values = ${vecType}(
|
|
getValue(batch, inIdx),
|
|
initializationValue,
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${windowSizeVec4Remainder === 2}) {
|
|
${vecType} values = ${vecType}(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${windowSizeVec4Remainder === 3}) {
|
|
${vecType} values = ${vecType}(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
getValue(batch, inIdx + 2),
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
setOutput(${returnValue});
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getReductionStages(inShape) {
|
|
const stages = [];
|
|
while (stages.length === 0 || stages[stages.length - 1].outSize !== 1) {
|
|
const outSize = stages.length ? stages[stages.length - 1].outSize : inShape[1];
|
|
const windowSize = computeOptimalWindowSize(outSize);
|
|
stages.push({
|
|
inSize: outSize,
|
|
windowSize,
|
|
outSize: Math.ceil(outSize / windowSize)
|
|
});
|
|
}
|
|
return stages;
|
|
}
|
|
function reduce(x, dtype, reductionType, backend2) {
|
|
const reductionStages = getReductionStages(x.shape);
|
|
let result = x;
|
|
for (let i = 0; i < reductionStages.length; i++) {
|
|
const { inSize, windowSize, outSize } = reductionStages[i];
|
|
let program;
|
|
let previousResult;
|
|
if (reductionType === "mean") {
|
|
program = i === 0 ? new MeanProgram({ windowSize, inSize, batchSize: x.shape[0], outSize }, inSize) : new MeanProgram({ windowSize, inSize, batchSize: x.shape[0], outSize });
|
|
} else {
|
|
program = new ReduceProgram({ windowSize, inSize, batchSize: x.shape[0], outSize }, reductionType);
|
|
}
|
|
previousResult = result;
|
|
result = backend2.runWebGLProgram(program, [result], dtype);
|
|
if (previousResult.dataId !== x.dataId) {
|
|
backend2.disposeIntermediateTensorInfo(previousResult);
|
|
}
|
|
}
|
|
return result;
|
|
}
|
|
class TransposeProgram {
|
|
constructor(aShape, newDim) {
|
|
this.variableNames = ["A"];
|
|
const outputShape = new Array(aShape.length);
|
|
for (let i = 0; i < outputShape.length; i++) {
|
|
outputShape[i] = aShape[newDim[i]];
|
|
}
|
|
this.outputShape = outputShape;
|
|
this.rank = outputShape.length;
|
|
const dtype = getCoordsDataType(this.rank);
|
|
const switched = getSwitchedCoords(newDim);
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} resRC = getOutputCoords();
|
|
setOutput(getA(${switched}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getSwitchedCoords(newDim) {
|
|
const rank = newDim.length;
|
|
if (rank > 6) {
|
|
throw Error(`Transpose for rank ${rank} is not yet supported`);
|
|
}
|
|
const originalOrder = ["resRC.x", "resRC.y", "resRC.z", "resRC.w", "resRC.u", "resRC.v"];
|
|
const switchedCoords = new Array(rank);
|
|
for (let i = 0; i < newDim.length; i++) {
|
|
switchedCoords[newDim[i]] = originalOrder[i];
|
|
}
|
|
return switchedCoords.join();
|
|
}
|
|
class TransposePackedProgram {
|
|
constructor(aShape, newDim) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
const outputShape = new Array(aShape.length);
|
|
for (let i = 0; i < outputShape.length; i++) {
|
|
outputShape[i] = aShape[newDim[i]];
|
|
}
|
|
this.outputShape = outputShape;
|
|
this.rank = outputShape.length;
|
|
if (this.rank > 6) {
|
|
throw Error(`Packed transpose for rank ${this.rank} is not yet supported.`);
|
|
}
|
|
const dtype = getCoordsDataType(this.rank);
|
|
const outputOrder = getVecChannels("rc", this.rank);
|
|
const switchedOrder = new Array(this.rank);
|
|
for (let i = 0; i < newDim.length; i++) {
|
|
switchedOrder[newDim[i]] = outputOrder[i];
|
|
}
|
|
const innerDims = `vec2(${switchedOrder.slice(-2).join()})`;
|
|
const nextColumn = `++${outputOrder[this.rank - 1]} < ${outputShape[this.rank - 1]}`;
|
|
const getc = `getChannel(getA(${switchedOrder.join()}), ${innerDims})`;
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} rc = getOutputCoords();
|
|
vec4 result = vec4(0.);
|
|
result[0] = ${getc};
|
|
if(${nextColumn}) {
|
|
result[1] = ${getc};
|
|
}
|
|
--${outputOrder[this.rank - 1]};
|
|
if(++${outputOrder[this.rank - 2]} < ${outputShape[this.rank - 2]}) {
|
|
result[2] = ${getc};
|
|
if(${nextColumn}) {
|
|
result[3] = ${getc};
|
|
}
|
|
}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function transposeImpl(x, perm, backend2) {
|
|
const program = env().getBool("WEBGL_PACK_ARRAY_OPERATIONS") ? new TransposePackedProgram(x.shape, perm) : new TransposeProgram(x.shape, perm);
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
function sumImpl(x, axis, keepDims, backend2) {
|
|
const reductionIndices = axis;
|
|
const xRank = x.shape.length;
|
|
const origAxes = parseAxisParam(reductionIndices, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
const sumInputIsTransposed = permutedAxes != null;
|
|
let sumInput = x;
|
|
if (sumInputIsTransposed) {
|
|
sumInput = transposeImpl(x, permutedAxes, backend2);
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
}
|
|
assertAxesAreInnerMostDims("sum", axes, xRank);
|
|
const [sumOutShape, reduceShape] = computeOutAndReduceShapes(sumInput.shape, axes);
|
|
let outShape = sumOutShape;
|
|
if (keepDims) {
|
|
outShape = expandShapeToKeepDim(sumOutShape, origAxes);
|
|
}
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const xSize = sizeFromShape(x.shape);
|
|
const batchSize = xSize / inSize;
|
|
const reshapedInput = reshape$1({ inputs: { x: sumInput }, attrs: { shape: [batchSize, inSize] }, backend: backend2 });
|
|
const outType = sumOutType(x.dtype);
|
|
const reduced = reduce(reshapedInput, outType, "sum", backend2);
|
|
const out = reshape$1({ inputs: { x: reduced }, attrs: { shape: outShape }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(reshapedInput);
|
|
backend2.disposeIntermediateTensorInfo(reduced);
|
|
if (sumInputIsTransposed) {
|
|
backend2.disposeIntermediateTensorInfo(sumInput);
|
|
}
|
|
return out;
|
|
}
|
|
function sum$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
return sumImpl(x, axis, keepDims, backend2);
|
|
}
|
|
const sumConfig$1 = {
|
|
kernelName: Sum,
|
|
backendName: "webgl",
|
|
kernelFunc: sum$1
|
|
};
|
|
function transpose(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { perm } = attrs;
|
|
const webglBackend = backend2;
|
|
const xRank = x.shape.length;
|
|
const newShape = new Array(xRank);
|
|
for (let i = 0; i < newShape.length; i++) {
|
|
newShape[i] = x.shape[perm[i]];
|
|
}
|
|
let out;
|
|
if (webglBackend.shouldExecuteOnCPU([x])) {
|
|
const xTexData = webglBackend.texData.get(x.dataId);
|
|
const values = xTexData.values;
|
|
const outValues = transposeImplCPU(values, x.shape, x.dtype, perm, newShape);
|
|
out = webglBackend.makeTensorInfo(newShape, x.dtype);
|
|
const outData = webglBackend.texData.get(out.dataId);
|
|
outData.values = outValues;
|
|
} else {
|
|
out = transposeImpl(x, perm, webglBackend);
|
|
}
|
|
return out;
|
|
}
|
|
const transposeConfig = {
|
|
kernelName: Transpose,
|
|
backendName: "webgl",
|
|
kernelFunc: transpose
|
|
};
|
|
const MATMUL_SHARED_DIM_THRESHOLD = 1e3;
|
|
function batchMatMulImpl({ a, b, transposeA, transposeB, backend: backend2, bias = null, preluActivationWeights = null, leakyreluAlpha = 0, activation = null }) {
|
|
const aRank = a.shape.length;
|
|
const bRank = b.shape.length;
|
|
const innerShapeA = transposeA ? a.shape[aRank - 2] : a.shape[aRank - 1];
|
|
const innerShapeB = transposeB ? b.shape[bRank - 1] : b.shape[bRank - 2];
|
|
const outerShapeA = transposeA ? a.shape[aRank - 1] : a.shape[aRank - 2];
|
|
const outerShapeB = transposeB ? b.shape[bRank - 2] : b.shape[bRank - 1];
|
|
const outerDimsA = a.shape.slice(0, -2);
|
|
const outerDimsB = b.shape.slice(0, -2);
|
|
const batchDimA = sizeFromShape(outerDimsA);
|
|
const batchDimB = sizeFromShape(outerDimsB);
|
|
const outShapeOuterDims = assertAndGetBroadcastShape(a.shape.slice(0, -2), b.shape.slice(0, -2));
|
|
const outShape = outShapeOuterDims.concat([outerShapeA, outerShapeB]);
|
|
assert(innerShapeA === innerShapeB, () => `Error in matMul: inner shapes (${innerShapeA}) and (${innerShapeB}) of Tensors with shapes ${a.shape} and ${b.shape} and transposeA=${transposeA} and transposeB=${transposeB} must match.`);
|
|
const a3dShape = transposeA ? [batchDimA, innerShapeA, outerShapeA] : [batchDimA, outerShapeA, innerShapeA];
|
|
const b3dShape = transposeB ? [batchDimB, outerShapeB, innerShapeB] : [batchDimB, innerShapeB, outerShapeB];
|
|
const a3d = reshape$1({ inputs: { x: a }, backend: backend2, attrs: { shape: a3dShape } });
|
|
const b3d = reshape$1({ inputs: { x: b }, backend: backend2, attrs: { shape: b3dShape } });
|
|
const intermediates = [a3d, b3d];
|
|
const batchDim = Math.max(batchDimA, batchDimB);
|
|
const sharedDim = transposeA ? a3d.shape[1] : a3d.shape[2];
|
|
const hasBias = bias != null;
|
|
const hasPreluActivationWeights = preluActivationWeights != null;
|
|
const hasLeakyreluAlpha = activation === "leakyrelu";
|
|
const fusedActivation = activation != null ? mapActivationToShaderProgram(activation, true) : null;
|
|
const containsFusedOps = hasBias || hasPreluActivationWeights || hasLeakyreluAlpha || fusedActivation != null;
|
|
let out;
|
|
if ((outerShapeA === 1 || outerShapeB === 1) && sharedDim > MATMUL_SHARED_DIM_THRESHOLD && containsFusedOps === false) {
|
|
let aVec = a3d;
|
|
let bVec = b3d;
|
|
if (transposeA) {
|
|
aVec = transpose({ inputs: { x: a3d }, backend: backend2, attrs: { perm: [0, 2, 1] } });
|
|
intermediates.push(aVec);
|
|
}
|
|
if (transposeB) {
|
|
bVec = transpose({ inputs: { x: b3d }, backend: backend2, attrs: { perm: [0, 2, 1] } });
|
|
intermediates.push(bVec);
|
|
}
|
|
const shouldReshapeA = outerShapeB !== 1;
|
|
const shouldReshapeB = outerShapeB === 1;
|
|
let aVec3d = aVec;
|
|
if (shouldReshapeA) {
|
|
aVec3d = reshape$1({
|
|
inputs: { x: aVec },
|
|
backend: backend2,
|
|
attrs: { shape: [batchDim, sharedDim, 1] }
|
|
});
|
|
intermediates.push(aVec3d);
|
|
}
|
|
const axis = outerShapeB === 1 ? 2 : 1;
|
|
let bVec3d = bVec;
|
|
if (shouldReshapeB) {
|
|
bVec3d = reshape$1({
|
|
inputs: { x: bVec },
|
|
backend: backend2,
|
|
attrs: { shape: [batchDim, 1, sharedDim] }
|
|
});
|
|
intermediates.push(bVec3d);
|
|
}
|
|
const product = multiply({ inputs: { a: aVec3d, b: bVec3d }, backend: backend2 });
|
|
out = sum$1({ inputs: { x: product }, backend: backend2, attrs: { axis, keepDims: true } });
|
|
intermediates.push(product);
|
|
} else {
|
|
const dtype = upcastType(a.dtype, b.dtype);
|
|
const program = new MatMulPackedProgram(a3dShape, b3dShape, [batchDim, outerShapeA, outerShapeB], transposeA, transposeB, hasBias, fusedActivation, hasPreluActivationWeights, hasLeakyreluAlpha);
|
|
const inputs = [a3d, b3d];
|
|
if (bias != null) {
|
|
inputs.push(bias);
|
|
}
|
|
if (hasPreluActivationWeights) {
|
|
inputs.push(preluActivationWeights);
|
|
}
|
|
if (hasLeakyreluAlpha) {
|
|
const $leakyreluAlpha = backend2.makeTensorInfo([], "float32", createScalarValue(leakyreluAlpha, "float32"));
|
|
inputs.push($leakyreluAlpha);
|
|
intermediates.push($leakyreluAlpha);
|
|
}
|
|
out = backend2.runWebGLProgram(program, inputs, dtype);
|
|
}
|
|
const outReshaped = reshape$1({ inputs: { x: out }, backend: backend2, attrs: { shape: outShape } });
|
|
intermediates.push(out);
|
|
for (const i of intermediates) {
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
}
|
|
return outReshaped;
|
|
}
|
|
function _fusedMatMul$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { a, b, bias, preluActivationWeights } = inputs;
|
|
const { transposeA, transposeB, activation, leakyreluAlpha } = attrs;
|
|
return batchMatMulImpl({
|
|
a,
|
|
b,
|
|
transposeA,
|
|
transposeB,
|
|
backend: backend2,
|
|
bias,
|
|
preluActivationWeights,
|
|
leakyreluAlpha,
|
|
activation
|
|
});
|
|
}
|
|
const _fusedMatMulConfig$1 = {
|
|
kernelName: _FusedMatMul,
|
|
backendName: "webgl",
|
|
kernelFunc: _fusedMatMul$1
|
|
};
|
|
const ABS = `return abs(x);`;
|
|
function abs(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
if (backend2.shouldExecuteOnCPU([x]) && x.dtype !== "complex64") {
|
|
const xData = backend2.texData.get(x.dataId);
|
|
const outValues = simpleAbsImplCPU(xData.values);
|
|
return backend2.makeTensorInfo(x.shape, x.dtype, outValues);
|
|
}
|
|
let program;
|
|
if (env().getBool("WEBGL_PACK_UNARY_OPERATIONS")) {
|
|
program = new UnaryOpPackedProgram(x.shape, ABS);
|
|
} else {
|
|
program = new UnaryOpProgram(x.shape, ABS);
|
|
}
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
const absConfig = {
|
|
kernelName: Abs,
|
|
backendName: "webgl",
|
|
kernelFunc: abs
|
|
};
|
|
const ACOS = CHECK_NAN_SNIPPET$1 + `
|
|
if (abs(x) > 1.) {
|
|
return NAN;
|
|
}
|
|
return acos(x);
|
|
`;
|
|
const acos$1 = unaryKernelFunc({ opSnippet: ACOS });
|
|
const acosConfig$1 = {
|
|
kernelName: Acos,
|
|
backendName: "webgl",
|
|
kernelFunc: acos$1
|
|
};
|
|
const ACOSH = CHECK_NAN_SNIPPET$1 + `
|
|
if (x < 1.0) return NAN;
|
|
return log(x + sqrt(x * x - 1.0));`;
|
|
const acosh$1 = unaryKernelFunc({ opSnippet: ACOSH });
|
|
const acoshConfig$1 = {
|
|
kernelName: Acosh,
|
|
backendName: "webgl",
|
|
kernelFunc: acosh$1
|
|
};
|
|
const ADD = "return a + b;";
|
|
const addKernelFunc = binaryKernelFunc({
|
|
opSnippet: ADD,
|
|
packedOpSnippet: ADD,
|
|
supportsComplex: true,
|
|
cpuKernelImpl: addImplCPU
|
|
});
|
|
const addConfig = {
|
|
kernelName: Add,
|
|
backendName: "webgl",
|
|
kernelFunc: addKernelFunc
|
|
};
|
|
class AddNProgram {
|
|
constructor(outputShape, shapes) {
|
|
this.outputShape = [];
|
|
this.outputShape = outputShape;
|
|
this.variableNames = shapes.map((_, i) => `T${i}`);
|
|
const snippets = [];
|
|
this.variableNames.forEach((variable2) => {
|
|
snippets.push(`float v${variable2} = get${variable2}AtOutCoords();`);
|
|
});
|
|
const operation = this.variableNames.map((variable2) => {
|
|
return `v${variable2}`;
|
|
}).join(" + ");
|
|
this.userCode = `
|
|
void main() {
|
|
${snippets.join("\n ")}
|
|
|
|
float result = ${operation};
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class AddNPackedProgram {
|
|
constructor(outputShape, shapes) {
|
|
this.outputShape = [];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = outputShape;
|
|
this.variableNames = shapes.map((_, i) => `T${i}`);
|
|
const snippets = [];
|
|
this.variableNames.forEach((variable2) => {
|
|
snippets.push(`vec4 v${variable2} = get${variable2}AtOutCoords();`);
|
|
});
|
|
const operation = this.variableNames.map((variable2) => {
|
|
return `v${variable2}`;
|
|
}).join(" + ");
|
|
this.userCode = `
|
|
void main() {
|
|
${snippets.join("\n ")}
|
|
|
|
vec4 result = ${operation};
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function addN$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const tensors = inputs;
|
|
if (tensors.length === 1) {
|
|
return identity({ inputs: { x: tensors[0] }, backend: backend2 });
|
|
}
|
|
if (tensors.length > env().getNumber("WEBGL_MAX_TEXTURES_IN_SHADER")) {
|
|
const midIndex = Math.floor(tensors.length / 2);
|
|
const leftSide = addN$1({ inputs: tensors.slice(0, midIndex), backend: backend2 });
|
|
const rightSide = addN$1({ inputs: tensors.slice(midIndex), backend: backend2 });
|
|
return addN$1({ inputs: [leftSide, rightSide], backend: backend2 });
|
|
}
|
|
const dtype = tensors.map((t2) => t2.dtype).reduce((d1, d2) => upcastType(d1, d2));
|
|
const shapes = tensors.map((t2) => t2.shape);
|
|
const usePackedOp = env().getBool("WEBGL_PACK");
|
|
const program = usePackedOp ? new AddNPackedProgram(tensors[0].shape, shapes) : new AddNProgram(tensors[0].shape, shapes);
|
|
return backend2.runWebGLProgram(program, tensors, dtype);
|
|
}
|
|
const addNConfig$1 = {
|
|
kernelName: AddN,
|
|
backendName: "webgl",
|
|
kernelFunc: addN$1
|
|
};
|
|
function all$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
const xRank = x.shape.length;
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
let permutedX = x;
|
|
if (permutedAxes != null) {
|
|
permutedX = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
}
|
|
assertAxesAreInnerMostDims("all", axes, xRank);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(permutedX.shape, axes);
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const a2D = reshape$1({ inputs: { x: permutedX }, backend: backend2, attrs: { shape: [-1, inSize] } });
|
|
const reduced = reduce(a2D, a2D.dtype, "all", backend2);
|
|
let res;
|
|
if (keepDims) {
|
|
const newShape = expandShapeToKeepDim(outShape, origAxes);
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: newShape } });
|
|
} else {
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: outShape } });
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(a2D);
|
|
backend2.disposeIntermediateTensorInfo(reduced);
|
|
if (permutedAxes != null) {
|
|
backend2.disposeIntermediateTensorInfo(permutedX);
|
|
}
|
|
return res;
|
|
}
|
|
const allConfig$1 = {
|
|
kernelName: All,
|
|
backendName: "webgl",
|
|
kernelFunc: all$1
|
|
};
|
|
function any$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
const xRank = x.shape.length;
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
let permutedX = x;
|
|
if (permutedAxes != null) {
|
|
permutedX = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
}
|
|
assertAxesAreInnerMostDims("any", axes, xRank);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(permutedX.shape, axes);
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const a2D = reshape$1({ inputs: { x: permutedX }, backend: backend2, attrs: { shape: [-1, inSize] } });
|
|
const reduced = reduce(a2D, a2D.dtype, "any", backend2);
|
|
let res;
|
|
if (keepDims) {
|
|
const newShape = expandShapeToKeepDim(outShape, origAxes);
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: newShape } });
|
|
} else {
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: outShape } });
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(a2D);
|
|
backend2.disposeIntermediateTensorInfo(reduced);
|
|
if (permutedAxes != null) {
|
|
backend2.disposeIntermediateTensorInfo(permutedX);
|
|
}
|
|
return res;
|
|
}
|
|
const anyConfig$1 = {
|
|
kernelName: Any,
|
|
backendName: "webgl",
|
|
kernelFunc: any$1
|
|
};
|
|
class ArgMinMaxProgram {
|
|
constructor(reduceInfo, op2, firstPass) {
|
|
this.variableNames = ["A"];
|
|
const { windowSize, batchSize, outSize } = reduceInfo;
|
|
if (!firstPass) {
|
|
this.variableNames.push("bestIndicesA");
|
|
}
|
|
this.outputShape = [batchSize, outSize];
|
|
const compOp = op2 === "max" ? ">" : "<";
|
|
const indexSnippet = firstPass ? "inOffset + i;" : "round(getBestIndicesA(batch, inOffset + i));";
|
|
this.userCode = `
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int outIdx = coords[1];
|
|
int inOffset = outIdx * ${windowSize};
|
|
|
|
int bestIndex = inOffset;
|
|
float bestValue = getA(batch, bestIndex);
|
|
|
|
for (int i = 0; i < ${windowSize}; i++) {
|
|
int inIdx = ${indexSnippet};
|
|
float candidate = getA(batch, inIdx);
|
|
if (candidate ${compOp} bestValue) {
|
|
bestValue = candidate;
|
|
bestIndex = inIdx;
|
|
}
|
|
}
|
|
setOutput(float(bestIndex));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ArgMinMaxPackedProgram {
|
|
constructor(shape, windowSize, op2, firstPass) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
assert(shape.length > 2, () => `Packed arg${op2.charAt(0).toUpperCase() + op2.slice(1)} supports only inputs with rank above 2.`);
|
|
const inSize = shape[shape.length - 1];
|
|
const outSize = Math.ceil(inSize / windowSize);
|
|
this.outputShape = shape.slice(0, -1);
|
|
if (outSize > 1) {
|
|
this.outputShape.push(outSize);
|
|
}
|
|
if (!firstPass) {
|
|
this.variableNames.push("bestIndicesA");
|
|
}
|
|
const outShape = this.outputShape;
|
|
const rank = outShape.length;
|
|
const dtype = getCoordsDataType(rank);
|
|
const coords2 = getChannels("coords", rank);
|
|
let sourceLocSetup;
|
|
let sourceRank;
|
|
if (outSize === 1) {
|
|
sourceRank = rank + 1;
|
|
const sourceLocDType = getCoordsDataType(sourceRank);
|
|
sourceLocSetup = `
|
|
${sourceLocDType} sourceLocR = ${sourceLocDType}(${coords2.join()}, 0);
|
|
++${coords2[rank - 1]};
|
|
${sourceLocDType} sourceLocG = ${sourceLocDType}(${coords2.join()}, 0);
|
|
++${coords2[rank - 2]};
|
|
${sourceLocDType} sourceLocA = ${sourceLocDType}(${coords2.join()}, 0);
|
|
--${coords2[rank - 1]};
|
|
${sourceLocDType} sourceLocB = ${sourceLocDType}(${coords2.join()}, 0);
|
|
--${coords2[rank - 2]};`;
|
|
} else {
|
|
sourceRank = rank;
|
|
sourceLocSetup = `
|
|
${dtype} sourceLocR = coords;
|
|
++${coords2[rank - 1]};
|
|
${dtype} sourceLocG = coords;
|
|
++${coords2[rank - 2]};
|
|
${dtype} sourceLocA = coords;
|
|
--${coords2[rank - 1]};
|
|
${dtype} sourceLocB = coords;
|
|
--${coords2[rank - 2]};`;
|
|
}
|
|
const channels = ["x", "y", "z", "w", "u", "v"].slice(0, sourceRank);
|
|
const inChannel = "." + channels[sourceRank - 1];
|
|
const intChannels = channels.map((x) => "int " + x);
|
|
const srcRCoords = getChannels("sourceLocR", sourceRank - 1).concat("inIdx.r");
|
|
const srcGCoords = getChannels("sourceLocG", sourceRank - 1).concat("inIdx.g");
|
|
const srcBCoords = getChannels("sourceLocB", sourceRank - 1).concat("inIdx.b");
|
|
const srcACoords = getChannels("sourceLocA", sourceRank - 1).concat("inIdx.a");
|
|
const compOp = op2 === "max" ? "greaterThan" : "lessThan";
|
|
const fetchCandidateIdx = firstPass ? "" : `
|
|
inIdx = round(vec4(getBestIndicesAChannel(${srcRCoords.join()}),
|
|
getBestIndicesAChannel(${srcGCoords.join()}),
|
|
getBestIndicesAChannel(${srcBCoords.join()}),
|
|
getBestIndicesAChannel(${srcACoords.join()})));`;
|
|
const fetchValue = `vec4(
|
|
getAChannel(${srcRCoords.join()}),
|
|
hasNextCol ? getAChannel(${srcGCoords.join()}) : 0.,
|
|
hasNextRow ? getAChannel(${srcBCoords.join()}) : 0.,
|
|
hasNextRow && hasNextCol ? getAChannel(${srcACoords.join()}) : 0.)`;
|
|
const getBestIndicesAChannelSnippet = firstPass ? "" : `
|
|
float getBestIndicesAChannel(${intChannels.join()}) {
|
|
return getChannel(getBestIndicesA(${channels.join()}),
|
|
vec2(${channels.slice(-2).join()}));
|
|
}`;
|
|
this.userCode = `
|
|
float getAChannel(${intChannels.join()}) {
|
|
return getChannel(getA(${channels.join()}),
|
|
vec2(${channels.slice(-2).join()}));
|
|
}
|
|
${getBestIndicesAChannelSnippet}
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
bool hasNextCol = ${coords2[rank - 1]} < ${outShape[rank - 1] - 1};
|
|
bool hasNextRow = ${coords2[rank - 2]} < ${outShape[rank - 2] - 1};
|
|
${sourceLocSetup}
|
|
ivec4 srcIdx = ivec4(sourceLocR${inChannel}, sourceLocG${inChannel},
|
|
sourceLocB${inChannel}, sourceLocA${inChannel}) * ${windowSize};
|
|
ivec4 inIdx = srcIdx;
|
|
vec4 bestIndex = vec4(inIdx);
|
|
vec4 bestValue = ${fetchValue};
|
|
|
|
for (int i = 0; i < ${windowSize}; i++) {
|
|
inIdx = srcIdx;
|
|
${fetchCandidateIdx}
|
|
vec4 candidate = ${fetchValue};
|
|
bvec4 nan = isnan(candidate);
|
|
bvec4 replace = bvec4(
|
|
vec4(${compOp}(candidate, bestValue)) * (vec4(1.0) - vec4(nan)));
|
|
|
|
bestValue = vec4(replace.x ? candidate.x : bestValue.x,
|
|
replace.y ? candidate.y : bestValue.y,
|
|
replace.z ? candidate.z : bestValue.z,
|
|
replace.w ? candidate.w : bestValue.w);
|
|
bestIndex = mix(bestIndex, vec4(inIdx), vec4(replace));
|
|
srcIdx++;
|
|
}
|
|
setOutput(bestIndex);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function argReduce(backend2, x, reduceType, bestIndicesA = null) {
|
|
let batchSize = x.shape[0];
|
|
let inSize = x.shape[1];
|
|
if (bestIndicesA != null) {
|
|
batchSize = bestIndicesA.shape[0];
|
|
inSize = bestIndicesA.shape[1];
|
|
}
|
|
const windowSize = computeOptimalWindowSize(inSize);
|
|
const reduceInfo = { windowSize, inSize, batchSize, outSize: Math.ceil(inSize / windowSize) };
|
|
const program = new ArgMinMaxProgram(reduceInfo, reduceType, bestIndicesA == null);
|
|
const inputs = [x];
|
|
if (bestIndicesA != null) {
|
|
inputs.push(bestIndicesA);
|
|
}
|
|
const output = backend2.runWebGLProgram(program, inputs, "int32");
|
|
if (output.shape[1] === 1) {
|
|
return output;
|
|
}
|
|
const result = argReduce(backend2, x, reduceType, output);
|
|
backend2.disposeIntermediateTensorInfo(output);
|
|
return result;
|
|
}
|
|
function argReducePacked(backend2, x, reduceType, bestIndicesA = null) {
|
|
const inShape = bestIndicesA != null ? bestIndicesA.shape : x.shape;
|
|
const inSize = inShape[inShape.length - 1];
|
|
const windowSize = computeOptimalWindowSize(inSize);
|
|
const program = new ArgMinMaxPackedProgram(inShape, windowSize, reduceType, bestIndicesA == null);
|
|
const inputs = bestIndicesA == null ? [x] : [x, bestIndicesA];
|
|
const output = backend2.runWebGLProgram(program, inputs, "int32");
|
|
if (output.shape.length === x.shape.length) {
|
|
const result = argReducePacked(backend2, x, reduceType, output);
|
|
backend2.disposeIntermediateTensorInfo(output);
|
|
return result;
|
|
}
|
|
return output;
|
|
}
|
|
function argMinMaxReduce(backend2, x, axis, reduceType) {
|
|
const axes = [axis];
|
|
assertAxesAreInnerMostDims("arg" + reduceType.charAt(0).toUpperCase() + reduceType.slice(1), axes, x.shape.length);
|
|
if (!env().getBool("WEBGL_PACK_REDUCE") || x.shape.length <= 2) {
|
|
const intermediateTensorInfos = [];
|
|
const xtexData = backend2.texData.get(x.dataId);
|
|
const xIsPacked = xtexData !== null && xtexData.isPacked;
|
|
let xUnPacked = x;
|
|
if (xIsPacked) {
|
|
xUnPacked = backend2.unpackTensor(x);
|
|
intermediateTensorInfos.push(xUnPacked);
|
|
}
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(xUnPacked.shape, axes);
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const a2D = reshape$1({ inputs: { x: xUnPacked }, backend: backend2, attrs: { shape: [-1, inSize] } });
|
|
intermediateTensorInfos.push(a2D);
|
|
const reduced = argReduce(backend2, a2D, reduceType);
|
|
intermediateTensorInfos.push(reduced);
|
|
const reshaped = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: outShape } });
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return reshaped;
|
|
}
|
|
return argReducePacked(backend2, x, reduceType);
|
|
}
|
|
function argMax$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis } = attrs;
|
|
let axes = parseAxisParam(axis, x.shape);
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
const intermediateTensorInfos = [];
|
|
if (permutedAxes != null) {
|
|
$x = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
intermediateTensorInfos.push($x);
|
|
axes = getInnerMostAxes(axes.length, $x.shape.length);
|
|
}
|
|
assertAxesAreInnerMostDims("argMax", [axes[0]], $x.shape.length);
|
|
const out = argMinMaxReduce(backend2, $x, axes[0], "max");
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return out;
|
|
}
|
|
const argMaxConfig$1 = {
|
|
kernelName: ArgMax,
|
|
backendName: "webgl",
|
|
kernelFunc: argMax$1
|
|
};
|
|
function argMin$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis } = attrs;
|
|
let axes = parseAxisParam(axis, x.shape);
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
const intermediateTensorInfos = [];
|
|
if (permutedAxes != null) {
|
|
$x = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
intermediateTensorInfos.push($x);
|
|
axes = getInnerMostAxes(axes.length, $x.shape.length);
|
|
}
|
|
assertAxesAreInnerMostDims("argMin", [axes[0]], $x.shape.length);
|
|
const out = argMinMaxReduce(backend2, $x, axes[0], "min");
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return out;
|
|
}
|
|
const argMinConfig$1 = {
|
|
kernelName: ArgMin,
|
|
backendName: "webgl",
|
|
kernelFunc: argMin$1
|
|
};
|
|
const ASIN = CHECK_NAN_SNIPPET$1 + `
|
|
if (abs(x) > 1.) {
|
|
return NAN;
|
|
}
|
|
return asin(x);
|
|
`;
|
|
const asin$1 = unaryKernelFunc({ opSnippet: ASIN });
|
|
const asinConfig$1 = {
|
|
kernelName: Asin,
|
|
backendName: "webgl",
|
|
kernelFunc: asin$1
|
|
};
|
|
const ASINH = CHECK_NAN_SNIPPET$1 + `return log(x + sqrt(x * x + 1.0));`;
|
|
const asinh$1 = unaryKernelFunc({ opSnippet: ASINH });
|
|
const asinhConfig$1 = {
|
|
kernelName: Asinh,
|
|
backendName: "webgl",
|
|
kernelFunc: asinh$1
|
|
};
|
|
const ATAN = CHECK_NAN_SNIPPET$1 + `
|
|
return atan(x);
|
|
`;
|
|
const atan$1 = unaryKernelFunc({ opSnippet: ATAN });
|
|
const atanConfig$1 = {
|
|
kernelName: Atan,
|
|
backendName: "webgl",
|
|
kernelFunc: atan$1
|
|
};
|
|
const ATAN2 = CHECK_NAN_SNIPPET + `
|
|
return atan(a, b);
|
|
`;
|
|
const ATAN2_PACKED = `
|
|
vec4 result = atan(a, b);
|
|
bvec4 isNaNA = isnan(a);
|
|
bvec4 isNaNB = isnan(b);
|
|
bvec4 isNaN = bvec4(isNaNA.x || isNaNB.x, isNaNA.y || isNaNB.y, isNaNA.z || isNaNB.z, isNaNA.w || isNaNB.w);
|
|
` + CHECK_NAN_SNIPPET_PACKED + `
|
|
return result;
|
|
`;
|
|
const atan2$1 = binaryKernelFunc({ opSnippet: ATAN2, packedOpSnippet: ATAN2_PACKED });
|
|
const atan2Config$1 = {
|
|
kernelName: Atan2,
|
|
backendName: "webgl",
|
|
kernelFunc: atan2$1
|
|
};
|
|
const ATANH = CHECK_NAN_SNIPPET$1 + `
|
|
if ((x < -1.0) || (x > 1.0)) return NAN;
|
|
return (log(1.0 + x) - log(1.0 - x)) / 2.0;`;
|
|
const atanh$1 = unaryKernelFunc({ opSnippet: ATANH });
|
|
const atanhConfig$1 = {
|
|
kernelName: Atanh,
|
|
backendName: "webgl",
|
|
kernelFunc: atanh$1
|
|
};
|
|
class Pool2DProgram {
|
|
constructor(convInfo, poolType, computePositions, flattenPositions = false, includeBatchInIndex = false) {
|
|
this.variableNames = ["x"];
|
|
if (poolType === "avg" && computePositions) {
|
|
throw new Error("Cannot compute positions for average pool.");
|
|
}
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
this.outputShape = convInfo.outShape;
|
|
const isAvgPool = poolType === "avg";
|
|
const batchFlattenPositionStr = `((batch * ${convInfo.inHeight} + xR) * ${convInfo.inWidth} + xC) * ${convInfo.inChannels} + d`;
|
|
const flattenPositionStr = `(xR * ${convInfo.inWidth} + xC) * ${convInfo.inChannels} + d`;
|
|
let initializationValue = "0.0";
|
|
if (!isAvgPool) {
|
|
initializationValue = "-1.0 / 1e-20";
|
|
}
|
|
if (computePositions) {
|
|
const compareOp2 = ">=";
|
|
this.userCode = `
|
|
const ivec2 strides = ivec2(${strideHeight}, ${strideWidth});
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int d = coords[3];
|
|
|
|
ivec2 xRCCorner = coords.yz * strides - pads;
|
|
int xRCorner = xRCCorner.x;
|
|
int xCCorner = xRCCorner.y;
|
|
|
|
// max/min x(?, ?, d) to get y(yR, yC, d).
|
|
// ? = to be determined
|
|
float minMaxValue = 0.0;
|
|
float minMaxValueFound = 0.0;
|
|
int minMaxPosition = 0;
|
|
float avgValue = 0.0;
|
|
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
int xR = xRCorner + wR;
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${effectiveFilterWidth};
|
|
wC += ${dilationWidth}) {
|
|
int xC = xCCorner + wC;
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
float value = getX(batch, xR, xC, d);
|
|
|
|
// If a min / max value has already been found, use it. If not,
|
|
// use the current value.
|
|
float currMinMaxValue = mix(
|
|
value, minMaxValue, minMaxValueFound);
|
|
if (value ${compareOp2} currMinMaxValue) {
|
|
minMaxValue = value;
|
|
minMaxValueFound = 1.0;
|
|
minMaxPosition = ${flattenPositions ? includeBatchInIndex ? batchFlattenPositionStr : flattenPositionStr : `wR * ${effectiveFilterWidth} + wC`};
|
|
}
|
|
}
|
|
}
|
|
setOutput(float(minMaxPosition));
|
|
}
|
|
`;
|
|
return;
|
|
}
|
|
const compareOp = "max";
|
|
let returnValue = `${poolType}(${poolType}(${poolType}(minMaxValue[0], minMaxValue[1]), minMaxValue[2]), minMaxValue[3])`;
|
|
if (poolType === "avg") {
|
|
returnValue = `avgValue / max(count, 1.0)`;
|
|
}
|
|
const filterWidthNearestVec4 = Math.floor(filterWidth / 4) * 4;
|
|
const filterWidthVec4Remainder = filterWidth % 4;
|
|
const updateSnippet = `
|
|
if (${isAvgPool}) {
|
|
avgValue += dot(values, ones);
|
|
} else {
|
|
minMaxValue = ${compareOp}(values, minMaxValue);
|
|
}
|
|
`;
|
|
this.userCode = `
|
|
const ivec2 strides = ivec2(${strideHeight}, ${strideWidth});
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
const float initializationValue = ${initializationValue};
|
|
const vec4 ones = vec4(1.0, 1.0, 1.0, 1.0);
|
|
|
|
float count = 0.0;
|
|
|
|
float getValue(int batch, int xR, int xC, int d) {
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
return initializationValue;
|
|
}
|
|
count += 1.0;
|
|
return getX(batch, xR, xC, d);
|
|
}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int d = coords[3];
|
|
|
|
ivec2 xRCCorner = coords.yz * strides - pads;
|
|
int xRCorner = xRCCorner.x;
|
|
int xCCorner = xRCCorner.y;
|
|
|
|
// max/min x(?, ?, d) to get y(yR, yC, d).
|
|
// ? = to be determined
|
|
vec4 minMaxValue = vec4(${initializationValue});
|
|
float avgValue = 0.0;
|
|
count = 0.0;
|
|
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
int xR = xRCorner + wR;
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${filterWidthNearestVec4}; wC += 4) {
|
|
int xC = xCCorner + wC * ${dilationWidth};
|
|
|
|
vec4 values = vec4(
|
|
getValue(batch, xR, xC, d),
|
|
getValue(batch, xR, xC + ${dilationWidth}, d),
|
|
getValue(batch, xR, xC + 2 * ${dilationWidth}, d),
|
|
getValue(batch, xR, xC + 3 * ${dilationWidth}, d)
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
|
|
int xC = xCCorner + ${filterWidthNearestVec4};
|
|
if (${filterWidthVec4Remainder === 1}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, xR, xC, d),
|
|
initializationValue,
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${filterWidthVec4Remainder === 2}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, xR, xC, d),
|
|
getValue(batch, xR, xC + ${dilationWidth}, d),
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${filterWidthVec4Remainder === 3}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, xR, xC, d),
|
|
getValue(batch, xR, xC + ${dilationWidth}, d),
|
|
getValue(batch, xR, xC + 2 * ${dilationWidth}, d),
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
}
|
|
setOutput(${returnValue});
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Pool3DProgram {
|
|
constructor(convInfo, poolType, computePositions, flattenPositions = false, includeBatchInIndex = false) {
|
|
this.variableNames = ["x"];
|
|
if (poolType === "avg" && computePositions) {
|
|
throw new Error("Cannot compute positions for average pool.");
|
|
}
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = convInfo.padInfo.front;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
this.outputShape = convInfo.outShape;
|
|
const isAvgPool = poolType === "avg";
|
|
let initializationValue = "0.0";
|
|
if (!isAvgPool) {
|
|
initializationValue = "-1.0 / 1e-20";
|
|
}
|
|
if (computePositions) {
|
|
const compareOp2 = ">=";
|
|
this.userCode = `
|
|
const ivec3 strides =
|
|
ivec3(${strideDepth}, ${strideHeight}, ${strideWidth});
|
|
const ivec3 pads = ivec3(${padFront}, ${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int ch = coords.u;
|
|
|
|
ivec3 xCorner = ivec3(coords.y, coords.z, coords.w) * strides - pads;
|
|
int xDCorner = xCorner.x;
|
|
int xRCorner = xCorner.y;
|
|
int xCCorner = xCorner.z;
|
|
|
|
// max/min x(?, ?, ?, ch) to get y(yD, yR, yC, ch).
|
|
// ? = to be determined
|
|
float minMaxValue = 0.0;
|
|
float minMaxValueFound = 0.0;
|
|
int minMaxPosition = 0;
|
|
|
|
for (int wD = 0; wD < ${effectiveFilterDepth};
|
|
wD += ${dilationDepth}) {
|
|
int xD = xDCorner + wD;
|
|
|
|
if (xD < 0 || xD >= ${convInfo.inDepth}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
int xR = xRCorner + wR;
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${effectiveFilterWidth};
|
|
wC += ${dilationWidth}) {
|
|
int xC = xCCorner + wC;
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
float value = getX(batch, xD, xR, xC, ch);
|
|
|
|
// If a min / max value has already been found, use it. If not,
|
|
// use the current value.
|
|
float currMinMaxValue = mix(
|
|
value, minMaxValue, minMaxValueFound);
|
|
if (value ${compareOp2} currMinMaxValue) {
|
|
minMaxValue = value;
|
|
minMaxValueFound = 1.0;
|
|
minMaxPosition = ${flattenPositions ? includeBatchInIndex ? `(((batch * ${convInfo.inDepth} + xD) * ${convInfo.inHeight} + xR) * ${convInfo.inWidth} + xC) * ${convInfo.inChannels} + ch` : `((xD * ${convInfo.inHeight} + xR) * ${convInfo.inWidth} + xC) * ${convInfo.inChannels} + ch` : `wD * ${effectiveFilterHeight} * ${effectiveFilterWidth} +
|
|
wR * ${effectiveFilterWidth} + wC`};
|
|
}
|
|
}
|
|
}
|
|
}
|
|
setOutput(float(minMaxPosition));
|
|
}
|
|
`;
|
|
return;
|
|
}
|
|
const compareOp = "max";
|
|
let returnValue = `${poolType}(${poolType}(${poolType}(minMaxValue[0], minMaxValue[1]), minMaxValue[2]), minMaxValue[3])`;
|
|
if (poolType === "avg") {
|
|
returnValue = `avgValue / max(count, 1.0)`;
|
|
}
|
|
const filterWidthNearestVec4 = Math.floor(filterWidth / 4) * 4;
|
|
const filterWidthVec4Remainder = filterWidth % 4;
|
|
const updateSnippet = `
|
|
if (${isAvgPool}) {
|
|
avgValue += dot(values, ones);
|
|
} else {
|
|
minMaxValue = ${compareOp}(values, minMaxValue);
|
|
}
|
|
`;
|
|
this.userCode = `
|
|
const ivec3 strides =
|
|
ivec3(${strideDepth}, ${strideHeight}, ${strideWidth});
|
|
const ivec3 pads = ivec3(${padFront}, ${padTop}, ${padLeft});
|
|
const float initializationValue = ${initializationValue};
|
|
const vec4 ones = vec4(1.0, 1.0, 1.0, 1.0);
|
|
|
|
float count = 0.0;
|
|
|
|
float getValue(int batch, int xD, int xR, int xC, int ch) {
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
return initializationValue;
|
|
}
|
|
count += 1.0;
|
|
return getX(batch, xD, xR, xC, ch);
|
|
}
|
|
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int ch = coords.u;
|
|
|
|
ivec3 xCorner = ivec3(coords.y, coords.z, coords.w) * strides - pads;
|
|
int xDCorner = xCorner.x;
|
|
int xRCorner = xCorner.y;
|
|
int xCCorner = xCorner.z;
|
|
|
|
// max/min x(?, ?, ?, d) to get y(yD, yR, yC, ch).
|
|
// ? = to be determined
|
|
vec4 minMaxValue = vec4(${initializationValue});
|
|
float avgValue = 0.0;
|
|
count = 0.0;
|
|
|
|
for (int wD = 0; wD < ${effectiveFilterDepth};
|
|
wD += ${dilationDepth}) {
|
|
int xD = xDCorner + wD;
|
|
|
|
if (xD < 0 || xD >= ${convInfo.inDepth}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
int xR = xRCorner + wR;
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${filterWidthNearestVec4}; wC += 4) {
|
|
int xC = xCCorner + wC * ${dilationWidth};
|
|
|
|
vec4 values = vec4(
|
|
getValue(batch, xD, xR, xC, ch),
|
|
getValue(batch, xD, xR, xC + ${dilationWidth}, ch),
|
|
getValue(batch, xD, xR, xC + 2 * ${dilationWidth}, ch),
|
|
getValue(batch, xD, xR, xC + 3 * ${dilationWidth}, ch)
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
|
|
int xC = xCCorner + ${filterWidthNearestVec4};
|
|
if (${filterWidthVec4Remainder === 1}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, xD, xR, xC, ch),
|
|
initializationValue,
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${filterWidthVec4Remainder === 2}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, xD, xR, xC, ch),
|
|
getValue(batch, xD, xR, xC + ${dilationWidth}, ch),
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${filterWidthVec4Remainder === 3}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, xD, xR, xC, ch),
|
|
getValue(batch, xD, xR, xC + ${dilationWidth}, ch),
|
|
getValue(batch, xD, xR, xC + 2 * ${dilationWidth}, ch),
|
|
initializationValue
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
}
|
|
}
|
|
setOutput(${returnValue});
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function avgPool$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
assertNotComplex$1(x, "avgPool");
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const dilations = 1;
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in avgPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode);
|
|
if (convInfo.filterWidth === 1 && convInfo.filterHeight === 1 && arraysEqual(convInfo.inShape, convInfo.outShape)) {
|
|
return identity({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const avgPoolProgram = new Pool2DProgram(convInfo, "avg", false);
|
|
return backend2.runWebGLProgram(avgPoolProgram, [x], "float32");
|
|
}
|
|
const avgPoolConfig$1 = {
|
|
kernelName: AvgPool,
|
|
backendName: "webgl",
|
|
kernelFunc: avgPool$1
|
|
};
|
|
function avgPool3D$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode, dataFormat } = attrs;
|
|
const dilations = [1, 1, 1];
|
|
const convInfo = computePool3DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode, dataFormat);
|
|
const avgPoolProgram = new Pool3DProgram(convInfo, "avg", false);
|
|
return backend2.runWebGLProgram(avgPoolProgram, [x], "float32");
|
|
}
|
|
const avgPool3DConfig$1 = {
|
|
kernelName: AvgPool3D,
|
|
backendName: "webgl",
|
|
kernelFunc: avgPool3D$1
|
|
};
|
|
class AvgPool2DBackpropProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy"];
|
|
this.outputShape = convInfo.inShape;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const avgMultiplier = 1 / (filterHeight * filterWidth);
|
|
this.userCode = `
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
const float avgMultiplier = float(${avgMultiplier});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
|
|
ivec2 dyRCCorner = coords.yz - pads;
|
|
int dyRCorner = dyRCCorner.x;
|
|
int dyCCorner = dyRCCorner.y;
|
|
|
|
// Convolve dy(?, ?, d) with pos mask(:, :, d) to get dx(xR, xC, d).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 || fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
for (int wC = 0; wC < ${effectiveFilterWidth};
|
|
wC+= ${dilationWidth}) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
float dyValue = getDy(b, idyR, idyC, d);
|
|
|
|
dotProd += dyValue * avgMultiplier;
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class AvgPool3DBackpropProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy"];
|
|
this.outputShape = convInfo.inShape;
|
|
const filterDepth = convInfo.filterDepth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = effectiveFilterDepth - 1 - convInfo.padInfo.front;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const avgMultiplier = 1 / (filterDepth * filterHeight * filterWidth);
|
|
this.userCode = `
|
|
const ivec3 pads = ivec3(${padFront}, ${padTop}, ${padLeft});
|
|
const float avgMultiplier = float(${avgMultiplier});
|
|
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int ch = coords.u;
|
|
|
|
ivec3 dyCorner = ivec3(coords.y, coords.z, coords.w) - pads;
|
|
int dyDCorner = dyCorner.x;
|
|
int dyRCorner = dyCorner.y;
|
|
int dyCCorner = dyCorner.z;
|
|
|
|
// Convolve dy(?, ?, ?, d) with pos mask(:, :, :, ch) to get
|
|
// dx(xD, xR, xC, ch).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
|
|
for (int wD = 0; wD < ${effectiveFilterDepth};
|
|
wD += ${dilationDepth}) {
|
|
float dyD = float(dyDCorner + wD) / ${strideDepth}.0;
|
|
|
|
if (dyD < 0.0 || dyD >= ${convInfo.outDepth}.0 || fract(dyD) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyD = int(dyD);
|
|
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 ||
|
|
fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
for (int wC = 0; wC < ${effectiveFilterWidth};
|
|
wC += ${dilationWidth}) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
float dyValue = getDy(batch, idyD, idyR, idyC, ch);
|
|
|
|
dotProd += dyValue * avgMultiplier;
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function avgPool3DGrad$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input } = inputs;
|
|
const x = input;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const dilations = [1, 1, 1];
|
|
const convInfo = computePool3DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode);
|
|
const avgPoolBackpropProgram = new AvgPool3DBackpropProgram(convInfo);
|
|
return backend2.runWebGLProgram(avgPoolBackpropProgram, [dy], x.dtype);
|
|
}
|
|
const avgPool3DGradConfig$1 = {
|
|
kernelName: AvgPool3DGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: avgPool3DGrad$1
|
|
};
|
|
function avgPoolGrad$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input } = inputs;
|
|
const x = input;
|
|
assertNotComplex$1([dy, input], "avgPoolGrad");
|
|
const { filterSize, strides, pad: pad2 } = attrs;
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, 1, pad2);
|
|
const avgPoolBackpropProgram = new AvgPool2DBackpropProgram(convInfo);
|
|
return backend2.runWebGLProgram(avgPoolBackpropProgram, [dy], x.dtype);
|
|
}
|
|
const avgPoolGradConfig$1 = {
|
|
kernelName: AvgPoolGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: avgPoolGrad$1
|
|
};
|
|
function batchMatMul$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { a, b } = inputs;
|
|
const { transposeA, transposeB } = attrs;
|
|
return batchMatMulImpl({ a, b, transposeA, transposeB, backend: backend2 });
|
|
}
|
|
const batchMatMulConfig$1 = {
|
|
kernelName: BatchMatMul,
|
|
backendName: "webgl",
|
|
kernelFunc: batchMatMul$1
|
|
};
|
|
class BatchNormProgram {
|
|
constructor(xShape, meanShape, varianceShape, offsetShape, scaleShape, varianceEpsilon) {
|
|
this.outputShape = [];
|
|
this.variableNames = ["x", "mean", "variance"];
|
|
assertAndGetBroadcastShape(xShape, meanShape);
|
|
assertAndGetBroadcastShape(xShape, varianceShape);
|
|
let offsetSnippet = "0.0";
|
|
if (offsetShape != null) {
|
|
assertAndGetBroadcastShape(xShape, offsetShape);
|
|
this.variableNames.push("offset");
|
|
offsetSnippet = "getOffsetAtOutCoords()";
|
|
}
|
|
let scaleSnippet = "1.0";
|
|
if (scaleShape != null) {
|
|
assertAndGetBroadcastShape(xShape, scaleShape);
|
|
this.variableNames.push("scale");
|
|
scaleSnippet = "getScaleAtOutCoords()";
|
|
}
|
|
this.outputShape = xShape;
|
|
this.userCode = `
|
|
void main() {
|
|
float x = getXAtOutCoords();
|
|
float mean = getMeanAtOutCoords();
|
|
float variance = getVarianceAtOutCoords();
|
|
float offset = ${offsetSnippet};
|
|
float scale = ${scaleSnippet};
|
|
float inv = scale * inversesqrt(variance + float(${varianceEpsilon}));
|
|
setOutput(dot(vec3(x, -mean, offset), vec3(inv, inv, 1)));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class BatchNormPackedProgram {
|
|
constructor(xShape, meanShape, varianceShape, offsetShape, scaleShape, varianceEpsilon) {
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.variableNames = ["x", "mean", "variance"];
|
|
assertAndGetBroadcastShape(xShape, meanShape);
|
|
assertAndGetBroadcastShape(xShape, varianceShape);
|
|
let offsetSnippet = "vec4(0.0)";
|
|
if (offsetShape != null) {
|
|
assertAndGetBroadcastShape(xShape, offsetShape);
|
|
this.variableNames.push("offset");
|
|
offsetSnippet = "getOffsetAtOutCoords()";
|
|
}
|
|
let scaleSnippet = "vec4(1.0)";
|
|
if (scaleShape != null) {
|
|
assertAndGetBroadcastShape(xShape, scaleShape);
|
|
this.variableNames.push("scale");
|
|
scaleSnippet = "getScaleAtOutCoords()";
|
|
}
|
|
this.outputShape = xShape;
|
|
this.userCode = `
|
|
void main() {
|
|
vec4 offset = ${offsetSnippet};
|
|
vec4 scale = ${scaleSnippet};
|
|
|
|
vec4 x = getXAtOutCoords();
|
|
vec4 mean = getMeanAtOutCoords();
|
|
vec4 variance = getVarianceAtOutCoords();
|
|
|
|
vec4 inv = scale * inversesqrt(variance + vec4(${varianceEpsilon}));
|
|
|
|
setOutput((x - mean) * inv + offset);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const batchNorm$1 = ({ inputs, backend: backend2, attrs }) => {
|
|
const { x, mean: mean2, variance, offset, scale: scale2 } = inputs;
|
|
assert(mean2.shape.length === variance.shape.length, () => "Batch normalization gradient requires mean and variance to have equal ranks.");
|
|
assert(offset == null || mean2.shape.length === offset.shape.length, () => "Batch normalization gradient requires mean and offset to have equal ranks.");
|
|
assert(scale2 == null || mean2.shape.length === scale2.shape.length, () => "Batch normalization gradient requires mean and scale to have equal ranks.");
|
|
let { varianceEpsilon } = attrs;
|
|
if (varianceEpsilon == null) {
|
|
varianceEpsilon = 1e-3;
|
|
}
|
|
const finalInputs = [x, mean2, variance];
|
|
let offsetShape = null;
|
|
if (offset != null) {
|
|
offsetShape = offset.shape;
|
|
finalInputs.push(offset);
|
|
}
|
|
let scaleShape = null;
|
|
if (scale2 != null) {
|
|
scaleShape = scale2.shape;
|
|
finalInputs.push(scale2);
|
|
}
|
|
const program = env().getBool("WEBGL_PACK_NORMALIZATION") ? new BatchNormPackedProgram(x.shape, mean2.shape, variance.shape, offsetShape, scaleShape, varianceEpsilon) : new BatchNormProgram(x.shape, mean2.shape, variance.shape, offsetShape, scaleShape, varianceEpsilon);
|
|
const output = backend2.runWebGLProgram(program, finalInputs, finalInputs[0].dtype);
|
|
return output;
|
|
};
|
|
const batchNormConfig$1 = {
|
|
kernelName: FusedBatchNorm,
|
|
backendName: "webgl",
|
|
kernelFunc: batchNorm$1
|
|
};
|
|
class SliceProgram {
|
|
constructor(destSize) {
|
|
this.variableNames = ["source"];
|
|
this.outputShape = destSize;
|
|
this.rank = destSize.length;
|
|
const dtype = getCoordsDataType(this.rank);
|
|
this.customUniforms = [{ name: "start", arrayIndex: this.rank, type: "int" }];
|
|
const sourceCoords = getCoords$1(this.rank);
|
|
let body;
|
|
const coordSum = destSize.map((_, i) => {
|
|
return `sourceLoc.${coords[i]} = start[${i}] + coords.${coords[i]};`;
|
|
});
|
|
body = `
|
|
${dtype} sourceLoc;
|
|
${dtype} coords = getOutputCoords();
|
|
${coordSum.join("\n")}
|
|
`;
|
|
this.userCode = `
|
|
void main() {
|
|
${body}
|
|
setOutput(getSource(${sourceCoords}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const coords = ["x", "y", "z", "w", "u", "v"];
|
|
function getCoords$1(rank) {
|
|
if (rank === 1) {
|
|
return "sourceLoc";
|
|
} else if (rank <= 6) {
|
|
return coords.slice(0, rank).map((x) => "sourceLoc." + x).join(",");
|
|
} else {
|
|
throw Error(`Slicing for rank ${rank} is not yet supported`);
|
|
}
|
|
}
|
|
class SlicePackedProgram {
|
|
constructor(destSize) {
|
|
this.variableNames = ["source"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = destSize;
|
|
this.rank = destSize.length;
|
|
this.customUniforms = [{ name: "start", arrayIndex: this.rank, type: "int" }];
|
|
const dtype = getCoordsDataType(this.rank);
|
|
const coords2 = getChannels("coords", this.rank);
|
|
const sourceLoc = getChannels("sourceLoc", this.rank);
|
|
const innerDims = this.rank === 1 ? "sourceLoc" : `vec2(${sourceLoc.slice(-2).join()})`;
|
|
const getChannel = `getChannel(getSource(${sourceLoc.join()}), ${innerDims})`;
|
|
const upperRow = `
|
|
result.x = ${getChannel};
|
|
if (++${coords2[this.rank - 1]} < ${destSize[this.rank - 1]}) {
|
|
++${sourceLoc[this.rank - 1]};
|
|
result.y = ${getChannel};
|
|
--${sourceLoc[this.rank - 1]};
|
|
}
|
|
`;
|
|
const lowerRow = this.rank === 1 ? "" : `
|
|
--${coords2[this.rank - 1]};
|
|
if (++${coords2[this.rank - 2]} < ${destSize[this.rank - 2]}) {
|
|
++${sourceLoc[this.rank - 2]};
|
|
result.z = ${getChannel};
|
|
if (++${coords2[this.rank - 1]} < ${destSize[this.rank - 1]}) {
|
|
++${sourceLoc[this.rank - 1]};
|
|
result.w = ${getChannel};
|
|
}
|
|
}
|
|
`;
|
|
const sourceLocSetup = this.rank <= 4 ? `sourceLoc = coords +
|
|
${dtype}(${destSize.map((_, i) => `start[${i}]`).join()});` : destSize.map((_, i) => `${sourceLoc[i]} = ${coords2[i]} + start[${i}];`).join("\n");
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
${dtype} sourceLoc;
|
|
${sourceLocSetup}
|
|
vec4 result = vec4(0.);
|
|
${upperRow}
|
|
${lowerRow}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function shallowSlice(x, begin, size, backend2) {
|
|
const xTexData = backend2.texData.get(x.dataId);
|
|
const t2 = backend2.makeTensorInfo(size, x.dtype);
|
|
const newTexData = backend2.texData.get(t2.dataId);
|
|
Object.assign(newTexData, xTexData);
|
|
newTexData.refCount = 1;
|
|
newTexData.shape = size;
|
|
newTexData.dtype = x.dtype;
|
|
let flatOffset = computeFlatOffset(begin, computeStrides(x.shape));
|
|
if (xTexData.slice) {
|
|
flatOffset += xTexData.slice.flatOffset;
|
|
}
|
|
newTexData.slice = {
|
|
flatOffset,
|
|
// Point to the original dataId, which is used to do ref counting.
|
|
origDataId: xTexData.slice && xTexData.slice.origDataId || x.dataId
|
|
};
|
|
const refCount = backend2.dataRefCount.get(newTexData.slice.origDataId) || 1;
|
|
backend2.dataRefCount.set(newTexData.slice.origDataId, refCount + 1);
|
|
return t2;
|
|
}
|
|
function slice(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { begin, size } = attrs;
|
|
const [$begin, $size] = parseSliceParams(x, begin, size);
|
|
assertParamsValid(x, $begin, $size);
|
|
if (sizeFromShape($size) === 0) {
|
|
return backend2.makeTensorInfo($size, x.dtype, []);
|
|
}
|
|
if (backend2.shouldExecuteOnCPU([x]) || x.dtype === "string") {
|
|
const xTexData = backend2.texData.get(x.dataId);
|
|
const outValues = sliceImplCPU(xTexData.values, $begin, $size, x.shape, x.dtype);
|
|
return backend2.makeTensorInfo($size, x.dtype, outValues);
|
|
}
|
|
const { isPacked } = backend2.texData.get(x.dataId);
|
|
const isContinous = isSliceContinous(x.shape, $begin, $size);
|
|
if (isPacked || !isContinous) {
|
|
const program = env().getBool("WEBGL_PACK_ARRAY_OPERATIONS") ? new SlicePackedProgram($size) : new SliceProgram($size);
|
|
const customValues = [$begin];
|
|
return backend2.runWebGLProgram(program, [x], x.dtype, customValues);
|
|
}
|
|
backend2.uploadToGPU(x.dataId);
|
|
return shallowSlice(x, $begin, $size, backend2);
|
|
}
|
|
const sliceConfig = {
|
|
kernelName: Slice,
|
|
backendName: "webgl",
|
|
kernelFunc: slice
|
|
};
|
|
const batchToSpaceND$1 = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { blockShape, crops } = attrs;
|
|
assert(x.shape.length <= 4, () => "batchToSpaceND for rank > 4 with a WebGL backend not implemented yet");
|
|
const prod2 = blockShape.reduce((a, b) => a * b);
|
|
const reshaped = getReshaped(x.shape, blockShape, prod2);
|
|
const permuted = getPermuted(reshaped.length, blockShape.length);
|
|
const reshapedPermuted = getReshapedPermuted(x.shape, blockShape, prod2);
|
|
const sliceBeginCoords = getSliceBeginCoords(crops, blockShape.length);
|
|
const sliceSize = getSliceSize(reshapedPermuted, crops, blockShape.length);
|
|
const toDispose = [];
|
|
const reshapedIntermediate = reshape$1({ inputs: { x }, backend: backend2, attrs: { shape: reshaped } });
|
|
const transposedIntermediate = transpose({ inputs: { x: reshapedIntermediate }, backend: backend2, attrs: { perm: permuted } });
|
|
const reshapedIntermediate2 = reshape$1({
|
|
inputs: { x: transposedIntermediate },
|
|
backend: backend2,
|
|
attrs: { shape: reshapedPermuted }
|
|
});
|
|
const sliced = slice({
|
|
inputs: { x: reshapedIntermediate2 },
|
|
backend: backend2,
|
|
attrs: { begin: sliceBeginCoords, size: sliceSize }
|
|
});
|
|
toDispose.push(reshapedIntermediate);
|
|
toDispose.push(transposedIntermediate);
|
|
toDispose.push(reshapedIntermediate2);
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return sliced;
|
|
};
|
|
const batchToSpaceNDConfig$1 = {
|
|
kernelName: BatchToSpaceND,
|
|
backendName: "webgl",
|
|
kernelFunc: batchToSpaceND$1
|
|
};
|
|
function bincount$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, weights } = inputs;
|
|
const { size } = attrs;
|
|
const xVals = backend2.readSync(x.dataId);
|
|
const weightsVals = backend2.readSync(weights.dataId);
|
|
const outVals = bincountImplCPU(xVals, weightsVals, weights.dtype, weights.shape, size);
|
|
return backend2.makeTensorInfo([size], weights.dtype, outVals);
|
|
}
|
|
const bincountConfig$1 = {
|
|
kernelName: Bincount,
|
|
backendName: "webgl",
|
|
kernelFunc: bincount$1
|
|
};
|
|
const BITWISEAND = `
|
|
int r = int(a.r) & int(b.r);
|
|
int g = int(a.g) & int(b.g);
|
|
int rb = int(a.b) & int(b.b);
|
|
int ra = int(a.a) & int(b.a);
|
|
return vec4(r, g, rb, ra);
|
|
`;
|
|
const BITWISEAND_UNPACKED = `
|
|
return float(int(a.r) & int(b.r));
|
|
`;
|
|
function bitwiseAnd(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { a, b } = inputs;
|
|
const shouldUsePackedProgram = env().getBool("WEBGL_PACK_BINARY_OPERATIONS");
|
|
const versionNumber = env().getNumber("WEBGL_VERSION");
|
|
if (backend2.shouldExecuteOnCPU([a, b]) || versionNumber === 1) {
|
|
const aVals = backend2.texData.get(a.dataId).values;
|
|
const bVals = backend2.texData.get(b.dataId).values;
|
|
const [outValues, outShape] = bitwiseAndImplCPU(a.shape, b.shape, aVals, bVals, a.dtype);
|
|
const out = backend2.makeTensorInfo(outShape, a.dtype);
|
|
const outData = backend2.texData.get(out.dataId);
|
|
outData.values = outValues;
|
|
return out;
|
|
}
|
|
let program;
|
|
if (shouldUsePackedProgram) {
|
|
program = new BinaryOpPackedProgram(BITWISEAND, a.shape, b.shape, false);
|
|
} else {
|
|
program = new BinaryOpProgram(BITWISEAND_UNPACKED, a.shape, b.shape);
|
|
}
|
|
return backend2.runWebGLProgram(program, [a, b], a.dtype);
|
|
}
|
|
const bitwiseAndConfig = {
|
|
kernelName: BitwiseAnd,
|
|
backendName: "webgl",
|
|
kernelFunc: bitwiseAnd
|
|
};
|
|
function broadcastArgs$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { s0, s1 } = inputs;
|
|
const s0Vals = backend2.readSync(s0.dataId);
|
|
const s1Vals = backend2.readSync(s1.dataId);
|
|
const broadcastShape = assertAndGetBroadcastShape(Array.from(s0Vals), Array.from(s1Vals));
|
|
return backend2.makeTensorInfo([broadcastShape.length], "int32", Int32Array.from(broadcastShape));
|
|
}
|
|
const broadcastArgsConfig$1 = {
|
|
kernelName: BroadcastArgs,
|
|
backendName: "webgl",
|
|
kernelFunc: broadcastArgs$1
|
|
};
|
|
const NOT_EQUAL = `return float(a != b);`;
|
|
const notEqual = binaryKernelFunc({ opSnippet: NOT_EQUAL, cpuKernelImpl: notEqualImplCPU, dtype: "bool" });
|
|
const notEqualConfig = {
|
|
kernelName: NotEqual,
|
|
backendName: "webgl",
|
|
kernelFunc: notEqual
|
|
};
|
|
function real(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
const inputData = backend2.texData.get(input.dataId);
|
|
return identity({ inputs: { x: inputData.complexTensorInfos.real }, backend: backend2 });
|
|
}
|
|
const realConfig = {
|
|
kernelName: Real,
|
|
backendName: "webgl",
|
|
kernelFunc: real
|
|
};
|
|
const TO_INT = `return float(int(x));`;
|
|
function int(input, backend2) {
|
|
const program = new UnaryOpProgram(input.shape, TO_INT);
|
|
const output = backend2.runWebGLProgram(program, [input], "int32");
|
|
return { dataId: output.dataId, shape: output.shape, dtype: output.dtype };
|
|
}
|
|
function cast(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { dtype } = attrs;
|
|
if (dtype === "complex64") {
|
|
if (x.dtype === "complex64") {
|
|
return identity({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const zerosTensor = zeros$1(x.shape);
|
|
const floatX = cast({ inputs: { x }, backend: backend2, attrs: { dtype: "float32" } });
|
|
const result = complex({ inputs: { real: floatX, imag: zerosTensor }, backend: backend2 });
|
|
zerosTensor.dispose();
|
|
backend2.disposeIntermediateTensorInfo(floatX);
|
|
return result;
|
|
}
|
|
if (x.dtype === "complex64") {
|
|
const realPart = real({ inputs: { input: x }, backend: backend2 });
|
|
const result = cast({ inputs: { x: realPart }, backend: backend2, attrs: { dtype } });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
return result;
|
|
}
|
|
if (!hasEncodingLoss(x.dtype, dtype)) {
|
|
const result = identity({ inputs: { x }, backend: backend2 });
|
|
return { dataId: result.dataId, shape: result.shape, dtype };
|
|
}
|
|
if (backend2.shouldExecuteOnCPU([x])) {
|
|
const values = backend2.texData.get(x.dataId).values;
|
|
const [resultShape, resultType, resultData] = castImplCPU(values, x.shape, x.dtype, dtype);
|
|
return backend2.makeTensorInfo(resultShape, resultType, resultData);
|
|
}
|
|
if (dtype === "int32") {
|
|
return int(x, backend2);
|
|
}
|
|
if (dtype === "bool") {
|
|
const zerosTensorInfo = backend2.makeTensorInfo([], "bool", getTypedArrayFromDType("bool", 1));
|
|
const binaryInputs = { a: x, b: zerosTensorInfo };
|
|
const result = notEqual({ inputs: binaryInputs, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(zerosTensorInfo);
|
|
return result;
|
|
}
|
|
throw new Error(`Error in Cast: failed to cast ${x.dtype} to ${dtype}`);
|
|
}
|
|
const castConfig = {
|
|
kernelName: Cast,
|
|
backendName: "webgl",
|
|
kernelFunc: cast
|
|
};
|
|
const CEIL = `return ceil(x);`;
|
|
const ceil = unaryKernelFunc({ opSnippet: CEIL, packedOpSnippet: CEIL, cpuKernelImpl: ceilImplCPU });
|
|
const ceilConfig = {
|
|
kernelName: Ceil,
|
|
backendName: "webgl",
|
|
kernelFunc: ceil
|
|
};
|
|
class ClipProgram {
|
|
constructor(aShape) {
|
|
this.variableNames = ["A"];
|
|
this.customUniforms = [
|
|
{ name: "minVal", type: "float" },
|
|
{ name: "maxVal", type: "float" }
|
|
];
|
|
this.outputShape = aShape;
|
|
this.userCode = `
|
|
|
|
void main() {
|
|
float value = getAAtOutCoords();
|
|
if (isnan(value)) {
|
|
setOutput(value);
|
|
return;
|
|
}
|
|
|
|
setOutput(clamp(value, minVal, maxVal));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ClipPackedProgram {
|
|
constructor(aShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [
|
|
{ name: "minVal", type: "float" },
|
|
{ name: "maxVal", type: "float" }
|
|
];
|
|
this.outputShape = aShape;
|
|
this.userCode = `
|
|
void main() {
|
|
vec4 value = getAAtOutCoords();
|
|
|
|
if (any(isnan(value))) {
|
|
setOutput(value);
|
|
return;
|
|
}
|
|
|
|
setOutput(clamp(value, vec4(minVal), vec4(maxVal)));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function clipByValue$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { clipValueMin, clipValueMax } = attrs;
|
|
let program;
|
|
if (env().getBool("WEBGL_PACK_CLIP")) {
|
|
program = new ClipPackedProgram(x.shape);
|
|
} else {
|
|
program = new ClipProgram(x.shape);
|
|
}
|
|
const customValues = [[clipValueMin], [clipValueMax]];
|
|
return backend2.runWebGLProgram(program, [x], x.dtype, customValues);
|
|
}
|
|
const clipByValueConfig$1 = {
|
|
kernelName: ClipByValue,
|
|
backendName: "webgl",
|
|
kernelFunc: clipByValue$1
|
|
};
|
|
class ComplexAbsProgram {
|
|
constructor(shape) {
|
|
this.variableNames = ["real", "imag"];
|
|
this.outputShape = shape;
|
|
this.userCode = `
|
|
void main() {
|
|
float re = abs(getRealAtOutCoords());
|
|
float im = abs(getImagAtOutCoords());
|
|
float mx = max(re, im);
|
|
|
|
// sadly the length function in glsl is not underflow-safe
|
|
// (at least not on Intel GPUs). So the safe solution is
|
|
// to ensure underflow-safety in all cases.
|
|
setOutput(
|
|
mx == 0.0 ? 0.0 : mx * length(vec2(1, min(re, im)/mx))
|
|
);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function makeComplexComponentTensorInfo(complexTensor, complexPart) {
|
|
return {
|
|
dataId: complexPart.dataId,
|
|
dtype: complexPart.dtype,
|
|
shape: complexTensor.shape
|
|
};
|
|
}
|
|
function complexAbs$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
const xData = backend2.texData.get(x.dataId);
|
|
const program = new ComplexAbsProgram(x.shape);
|
|
const programInputs = [
|
|
makeComplexComponentTensorInfo(x, xData.complexTensorInfos.real),
|
|
makeComplexComponentTensorInfo(x, xData.complexTensorInfos.imag)
|
|
];
|
|
return backend2.runWebGLProgram(program, programInputs, programInputs[0].dtype);
|
|
}
|
|
const complexAbsConfig$1 = {
|
|
kernelName: ComplexAbs,
|
|
backendName: "webgl",
|
|
kernelFunc: complexAbs$1
|
|
};
|
|
class ConcatProgram {
|
|
// Concats 2d tensors along axis=1. See comments in MathBackendWebGL.concat().
|
|
constructor(shapes) {
|
|
this.outputShape = [];
|
|
this.outputShape = computeOutShape$1(
|
|
shapes,
|
|
1
|
|
/* axis */
|
|
);
|
|
this.variableNames = shapes.map((_, i) => `T${i}`);
|
|
const offsets = new Array(shapes.length - 1);
|
|
offsets[0] = shapes[0][1];
|
|
for (let i = 1; i < offsets.length; i++) {
|
|
offsets[i] = offsets[i - 1] + shapes[i][1];
|
|
}
|
|
const snippets = [`if (yC < ${offsets[0]}) setOutput(getT0(yR, yC));`];
|
|
for (let i = 1; i < offsets.length; i++) {
|
|
const shift = offsets[i - 1];
|
|
snippets.push(`else if (yC < ${offsets[i]}) setOutput(getT${i}(yR, yC-${shift}));`);
|
|
}
|
|
const lastIndex = offsets.length;
|
|
const lastShift = offsets[offsets.length - 1];
|
|
snippets.push(`else setOutput(getT${lastIndex}(yR, yC-${lastShift}));`);
|
|
this.userCode = `
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int yR = coords.x;
|
|
int yC = coords.y;
|
|
|
|
${snippets.join("\n ")}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ConcatPackedProgram {
|
|
constructor(shapes, axis) {
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = [];
|
|
this.outputShape = computeOutShape$1(shapes, axis);
|
|
const shape = this.outputShape;
|
|
const rank = shape.length;
|
|
const dtype = getCoordsDataType(rank);
|
|
const coords2 = getChannels("coords", rank);
|
|
const channels = ["x", "y", "z", "w", "u", "v"].slice(0, rank);
|
|
this.variableNames = shapes.map((_, i) => `T${i}`);
|
|
const offsets = new Array(shapes.length - 1);
|
|
offsets[0] = shapes[0][axis];
|
|
for (let i = 1; i < offsets.length; i++) {
|
|
offsets[i] = offsets[i - 1] + shapes[i][axis];
|
|
}
|
|
const channel = channels[axis];
|
|
const lastChannels = channels.slice(-2);
|
|
const allChannels = channels.join();
|
|
let getValueSnippet = `if (${channel} < ${offsets[0]}) {
|
|
return getChannel(
|
|
getT0(${allChannels}), vec2(${lastChannels.join()}));
|
|
}`;
|
|
for (let i = 1; i < offsets.length; i++) {
|
|
const shift2 = offsets[i - 1];
|
|
getValueSnippet += `
|
|
if (${channel} < ${offsets[i]} && ${channel} >= ${offsets[i - 1]}) {
|
|
return getChannel(
|
|
getT${i}(${shiftedChannels(channels, channel, shift2)}),
|
|
vec2(${shiftedChannels(lastChannels, channel, shift2)}));
|
|
}`;
|
|
}
|
|
const lastIndex = offsets.length;
|
|
const shift = offsets[offsets.length - 1];
|
|
getValueSnippet += `
|
|
return getChannel(
|
|
getT${lastIndex}(${shiftedChannels(channels, channel, shift)}),
|
|
vec2(${shiftedChannels(lastChannels, channel, shift)}));`;
|
|
this.userCode = `
|
|
float getValue(${channels.map((x) => "int " + x)}) {
|
|
${getValueSnippet}
|
|
}
|
|
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
vec4 result = vec4(getValue(${coords2}), 0., 0., 0.);
|
|
|
|
${coords2[rank - 1]} = ${coords2[rank - 1]} + 1;
|
|
if (${coords2[rank - 1]} < ${shape[rank - 1]}) {
|
|
result.g = getValue(${coords2});
|
|
}
|
|
|
|
${coords2[rank - 2]} = ${coords2[rank - 2]} + 1;
|
|
if (${coords2[rank - 2]} < ${shape[rank - 2]}) {
|
|
result.a = getValue(${coords2});
|
|
}
|
|
|
|
${coords2[rank - 1]} = ${coords2[rank - 1]} - 1;
|
|
if (${coords2[rank - 2]} < ${shape[rank - 2]} &&
|
|
${coords2[rank - 1]} < ${shape[rank - 1]}) {
|
|
result.b = getValue(${coords2});
|
|
}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function shiftedChannels(channels, channel, shift) {
|
|
const channelIdx = channels.indexOf(channel);
|
|
const res = channels.map((c, idx) => {
|
|
if (idx === channelIdx) {
|
|
return `${c} - ${shift}`;
|
|
} else {
|
|
return c;
|
|
}
|
|
});
|
|
return res.join();
|
|
}
|
|
function imag$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
const inputData = backend2.texData.get(input.dataId);
|
|
return identity({ inputs: { x: inputData.complexTensorInfos.imag }, backend: backend2 });
|
|
}
|
|
const imagConfig$1 = {
|
|
kernelName: Imag,
|
|
backendName: "webgl",
|
|
kernelFunc: imag$1
|
|
};
|
|
function concatImpl(inputs, axis, backend2) {
|
|
const dtype = inputs[0].dtype;
|
|
if (dtype === "complex64") {
|
|
const reals = inputs.map((t2) => real({ inputs: { input: t2 }, backend: backend2 }));
|
|
const imags = inputs.map((t2) => imag$1({ inputs: { input: t2 }, backend: backend2 }));
|
|
const realConcated = concatImpl(reals, axis, backend2);
|
|
const imagConcated = concatImpl(imags, axis, backend2);
|
|
const result2 = complex({ inputs: { real: realConcated, imag: imagConcated }, backend: backend2 });
|
|
reals.forEach((r) => backend2.disposeIntermediateTensorInfo(r));
|
|
imags.forEach((i) => backend2.disposeIntermediateTensorInfo(i));
|
|
backend2.disposeIntermediateTensorInfo(realConcated);
|
|
backend2.disposeIntermediateTensorInfo(imagConcated);
|
|
return result2;
|
|
}
|
|
let runOnCpu = backend2.shouldExecuteOnCPU(inputs);
|
|
if (dtype === "string") {
|
|
runOnCpu = true;
|
|
}
|
|
if (runOnCpu) {
|
|
const tensors2D2 = inputs.map((t2) => {
|
|
const innerSize = sizeFromShape(t2.shape.slice(axis));
|
|
const shape = [-1, innerSize];
|
|
return reshape$1({ inputs: { x: t2 }, backend: backend2, attrs: { shape } });
|
|
});
|
|
const inputsValShapes = tensors2D2.map((t2) => {
|
|
return { vals: backend2.readSync(t2.dataId), shape: t2.shape };
|
|
});
|
|
const outShape2 = computeOutShape$1(
|
|
tensors2D2.map((t2) => t2.shape),
|
|
1
|
|
/* axis */
|
|
);
|
|
const simplyConcat = tensors2D2[0].shape[0] === 1;
|
|
const outVals = concatImplCPU(inputsValShapes, outShape2, dtype, simplyConcat);
|
|
const finalOutShape = computeOutShape$1(inputs.map((t2) => t2.shape), axis);
|
|
const outInfo = backend2.makeTensorInfo(finalOutShape, dtype, outVals);
|
|
tensors2D2.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return outInfo;
|
|
}
|
|
const $inputs = inputs.filter((t2) => sizeFromShape(t2.shape) > 0);
|
|
const shouldPack = env().getBool("WEBGL_PACK_ARRAY_OPERATIONS") && $inputs[0].shape.length > 1;
|
|
if ($inputs.length === 1) {
|
|
const program2 = shouldPack ? new UnaryOpProgram(inputs[0].shape, CLONE) : new UnaryOpPackedProgram(inputs[0].shape, CLONE);
|
|
return backend2.runWebGLProgram(program2, inputs, dtype);
|
|
}
|
|
const maxTexturesInShader = env().getNumber("WEBGL_MAX_TEXTURES_IN_SHADER");
|
|
if ($inputs.length > maxTexturesInShader) {
|
|
const reducedInputs = [];
|
|
for (let i = 0; i < $inputs.length; i += maxTexturesInShader) {
|
|
const subArray = $inputs.slice(i, i + maxTexturesInShader);
|
|
reducedInputs.push(concatImpl(subArray, axis, backend2));
|
|
}
|
|
const result2 = concatImpl(reducedInputs, axis, backend2);
|
|
for (const i of reducedInputs) {
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
}
|
|
return result2;
|
|
}
|
|
if (shouldPack) {
|
|
const program2 = new ConcatPackedProgram($inputs.map((t2) => t2.shape), axis);
|
|
return backend2.runWebGLProgram(program2, $inputs, dtype);
|
|
}
|
|
const { tensors2D, outShape } = computeTensors2D($inputs, axis, backend2);
|
|
const program = new ConcatProgram(tensors2D.map((t2) => t2.shape));
|
|
const result = backend2.runWebGLProgram(program, tensors2D, dtype);
|
|
tensors2D.forEach((r) => backend2.disposeIntermediateTensorInfo(r));
|
|
const reshapedResult = reshape$1({ inputs: { x: result }, attrs: { shape: outShape }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return reshapedResult;
|
|
}
|
|
function computeTensors2D(inputs, axis, backend2) {
|
|
const outShape = computeOutShape$1(inputs.map((t2) => t2.shape), axis);
|
|
const tensors2D = inputs.map((x) => reshape$1({
|
|
inputs: { x },
|
|
attrs: { shape: [-1, sizeFromShape(x.shape.slice(axis))] },
|
|
backend: backend2
|
|
}));
|
|
return { tensors2D, outShape };
|
|
}
|
|
function concat$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { axis } = attrs;
|
|
const $axis = parseAxisParam(axis, inputs[0].shape)[0];
|
|
const shapes = inputs.map((t2) => t2.shape);
|
|
assertParamsConsistent(shapes, $axis);
|
|
const outShape = computeOutShape$1(inputs.map((t2) => t2.shape), $axis);
|
|
if (sizeFromShape(outShape) === 0) {
|
|
return backend2.makeTensorInfo(outShape, inputs[0].dtype, []);
|
|
}
|
|
const $inputs = inputs.filter((t2) => sizeFromShape(t2.shape) > 0);
|
|
if ($inputs.length === 1) {
|
|
return identity({ inputs: { x: $inputs[0] }, backend: backend2 });
|
|
}
|
|
return concatImpl($inputs, $axis, backend2);
|
|
}
|
|
const concatConfig$1 = {
|
|
kernelName: Concat,
|
|
backendName: "webgl",
|
|
kernelFunc: concat$1
|
|
};
|
|
class Conv2DProgram {
|
|
constructor(convInfo, addBias = false, activation = null, hasPreluActivationWeights = false, hasLeakyreluAlpha = false) {
|
|
this.variableNames = ["x", "W"];
|
|
this.outputShape = convInfo.outShape;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const inputDepthNearestVec4 = Math.floor(convInfo.inChannels / 4) * 4;
|
|
const inputDepthVec4Remainder = convInfo.inChannels % 4;
|
|
const isChannelsLast = convInfo.dataFormat === "channelsLast";
|
|
const rowDim = isChannelsLast ? 1 : 2;
|
|
const colDim = isChannelsLast ? 2 : 3;
|
|
const channelDim = isChannelsLast ? 3 : 1;
|
|
let activationSnippet = "", applyActivationSnippet = "";
|
|
if (activation) {
|
|
if (hasPreluActivationWeights) {
|
|
activationSnippet = `float activation(float a) {
|
|
float b = getPreluActivationWeightsAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else if (hasLeakyreluAlpha) {
|
|
activationSnippet = `float activation(float a) {
|
|
float b = getLeakyreluAlphaAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else {
|
|
activationSnippet = `
|
|
float activation(float x) {
|
|
${activation}
|
|
}
|
|
`;
|
|
}
|
|
applyActivationSnippet = `result = activation(result);`;
|
|
}
|
|
const addBiasSnippet = addBias ? "result += getBiasAtOutCoords();" : "";
|
|
if (addBias) {
|
|
this.variableNames.push("bias");
|
|
}
|
|
if (hasPreluActivationWeights) {
|
|
this.variableNames.push("preluActivationWeights");
|
|
}
|
|
if (hasLeakyreluAlpha) {
|
|
this.variableNames.push("leakyreluAlpha");
|
|
}
|
|
this.userCode = `
|
|
${activationSnippet}
|
|
|
|
const ivec2 strides = ivec2(${strideHeight}, ${strideWidth});
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int d2 = coords[${channelDim}];
|
|
|
|
ivec2 xRCCorner =
|
|
ivec2(coords[${rowDim}], coords[${colDim}]) * strides - pads;
|
|
int xRCorner = xRCCorner.x;
|
|
int xCCorner = xRCCorner.y;
|
|
|
|
// Convolve x(?, ?, d1) with w(:, :, d1, d2) to get y(yR, yC, d2).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
int xR = xRCorner + wR * ${dilationHeight};
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
int xC = xCCorner + wC * ${dilationWidth};
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
for (int d1 = 0; d1 < ${inputDepthNearestVec4}; d1 += 4) {
|
|
vec4 wValues = vec4(
|
|
getW(wR, wC, d1, d2),
|
|
getW(wR, wC, d1 + 1, d2),
|
|
getW(wR, wC, d1 + 2, d2),
|
|
getW(wR, wC, d1 + 3, d2)
|
|
);
|
|
|
|
if (${isChannelsLast}) {
|
|
vec4 xValues = vec4(
|
|
getX(batch, xR, xC, d1),
|
|
getX(batch, xR, xC, d1 + 1),
|
|
getX(batch, xR, xC, d1 + 2),
|
|
getX(batch, xR, xC, d1 + 3)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
} else {
|
|
vec4 xValues = vec4(
|
|
getX(batch, d1, xR, xC),
|
|
getX(batch, d1 + 1, xR, xC),
|
|
getX(batch, d1 + 2, xR, xC),
|
|
getX(batch, d1 + 3, xR, xC)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
}
|
|
}
|
|
|
|
if (${inputDepthVec4Remainder === 1}) {
|
|
|
|
if (${isChannelsLast}) {
|
|
dotProd +=
|
|
getX(batch, xR, xC, ${inputDepthNearestVec4}) *
|
|
getW(wR, wC, ${inputDepthNearestVec4}, d2);
|
|
} else {
|
|
dotProd +=
|
|
getX(batch, ${inputDepthNearestVec4}, xR, xC) *
|
|
getW(wR, wC, ${inputDepthNearestVec4}, d2);
|
|
}
|
|
|
|
} else if (${inputDepthVec4Remainder === 2}) {
|
|
vec2 wValues = vec2(
|
|
getW(wR, wC, ${inputDepthNearestVec4}, d2),
|
|
getW(wR, wC, ${inputDepthNearestVec4} + 1, d2)
|
|
);
|
|
|
|
if (${isChannelsLast}) {
|
|
vec2 xValues = vec2(
|
|
getX(batch, xR, xC, ${inputDepthNearestVec4}),
|
|
getX(batch, xR, xC, ${inputDepthNearestVec4} + 1)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
} else {
|
|
vec2 xValues = vec2(
|
|
getX(batch, ${inputDepthNearestVec4}, xR, xC),
|
|
getX(batch, ${inputDepthNearestVec4} + 1, xR, xC)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
}
|
|
|
|
} else if (${inputDepthVec4Remainder === 3}) {
|
|
vec3 wValues = vec3(
|
|
getW(wR, wC, ${inputDepthNearestVec4}, d2),
|
|
getW(wR, wC, ${inputDepthNearestVec4} + 1, d2),
|
|
getW(wR, wC, ${inputDepthNearestVec4} + 2, d2)
|
|
);
|
|
|
|
if (${isChannelsLast}) {
|
|
vec3 xValues = vec3(
|
|
getX(batch, xR, xC, ${inputDepthNearestVec4}),
|
|
getX(batch, xR, xC, ${inputDepthNearestVec4} + 1),
|
|
getX(batch, xR, xC, ${inputDepthNearestVec4} + 2)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
} else {
|
|
vec3 xValues = vec3(
|
|
getX(batch, ${inputDepthNearestVec4}, xR, xC),
|
|
getX(batch, ${inputDepthNearestVec4} + 1, xR, xC),
|
|
getX(batch, ${inputDepthNearestVec4} + 2, xR, xC)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
}
|
|
|
|
}
|
|
}
|
|
}
|
|
|
|
float result = dotProd;
|
|
${addBiasSnippet}
|
|
${applyActivationSnippet}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Conv3DProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["x", "W"];
|
|
this.outputShape = convInfo.outShape;
|
|
const padFront = convInfo.padInfo.front;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const filterDepth = convInfo.filterDepth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const inputDepthNearestVec4 = Math.floor(convInfo.inChannels / 4) * 4;
|
|
const inputDepthVec4Remainder = convInfo.inChannels % 4;
|
|
this.userCode = `
|
|
const ivec3 strides = ivec3(${strideDepth}, ${strideHeight}, ${strideWidth});
|
|
const ivec3 pads = ivec3(${padFront}, ${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int d2 = coords.u;
|
|
|
|
ivec3 xFRCCorner = ivec3(coords.y, coords.z, coords.w) * strides - pads;
|
|
int xFCorner = xFRCCorner.x;
|
|
int xRCorner = xFRCCorner.y;
|
|
int xCCorner = xFRCCorner.z;
|
|
|
|
// Convolve x(?, ?, ?, d1) with w(:, :, :, d1, d2) to get
|
|
// y(yF, yR, yC, d2). ? = to be determined. : = across all
|
|
// values in that axis.
|
|
float dotProd = 0.0;
|
|
for (int wF = 0; wF < ${filterDepth}; wF++) {
|
|
int xF = xFCorner + wF * ${dilationDepth};
|
|
|
|
if (xF < 0 || xF >= ${convInfo.inDepth}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
int xR = xRCorner + wR * ${dilationHeight};
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
int xC = xCCorner + wC * ${dilationWidth};
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
for (int d1 = 0; d1 < ${inputDepthNearestVec4}; d1 += 4) {
|
|
vec4 xValues = vec4(
|
|
getX(batch, xF, xR, xC, d1),
|
|
getX(batch, xF, xR, xC, d1 + 1),
|
|
getX(batch, xF, xR, xC, d1 + 2),
|
|
getX(batch, xF, xR, xC, d1 + 3)
|
|
);
|
|
vec4 wValues = vec4(
|
|
getW(wF, wR, wC, d1, d2),
|
|
getW(wF, wR, wC, d1 + 1, d2),
|
|
getW(wF, wR, wC, d1 + 2, d2),
|
|
getW(wF, wR, wC, d1 + 3, d2)
|
|
);
|
|
|
|
dotProd += dot(xValues, wValues);
|
|
}
|
|
|
|
if (${inputDepthVec4Remainder === 1}) {
|
|
dotProd +=
|
|
getX(batch, xF, xR, xC, ${inputDepthNearestVec4}) *
|
|
getW(wF, wR, wC, ${inputDepthNearestVec4}, d2);
|
|
} else if (${inputDepthVec4Remainder === 2}) {
|
|
vec2 xValues = vec2(
|
|
getX(batch, xF, xR, xC, ${inputDepthNearestVec4}),
|
|
getX(batch, xF, xR, xC, ${inputDepthNearestVec4} + 1)
|
|
);
|
|
vec2 wValues = vec2(
|
|
getW(wF, wR, wC, ${inputDepthNearestVec4}, d2),
|
|
getW(wF, wR, wC, ${inputDepthNearestVec4} + 1, d2)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
} else if (${inputDepthVec4Remainder === 3}) {
|
|
vec3 xValues = vec3(
|
|
getX(batch, xF, xR, xC, ${inputDepthNearestVec4}),
|
|
getX(batch, xF, xR, xC, ${inputDepthNearestVec4} + 1),
|
|
getX(batch, xF, xR, xC, ${inputDepthNearestVec4} + 2)
|
|
);
|
|
vec3 wValues = vec3(
|
|
getW(wF, wR, wC, ${inputDepthNearestVec4}, d2),
|
|
getW(wF, wR, wC, ${inputDepthNearestVec4} + 1, d2),
|
|
getW(wF, wR, wC, ${inputDepthNearestVec4} + 2, d2)
|
|
);
|
|
dotProd += dot(xValues, wValues);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Conv2DPackedProgram {
|
|
constructor(convInfo, addBias = false, activation = null, hasPreluActivation = false, hasLeakyReluAlpha = false) {
|
|
this.variableNames = ["x", "W"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [
|
|
{ name: "pads", type: "ivec2" },
|
|
{ name: "strides", type: "ivec2" },
|
|
{ name: "dilations", type: "ivec2" },
|
|
{ name: "inDims", type: "ivec2" }
|
|
];
|
|
this.outputShape = convInfo.outShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const padLeft = convInfo.padInfo.left;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const texelsAcross = filterWidth;
|
|
let mainLoop = `
|
|
int xR; int xC; int xCOffset;
|
|
vec4 wTexel; vec4 previous; vec4 final;`;
|
|
for (let c = 0; c < filterWidth; c++) {
|
|
mainLoop += `
|
|
vec4 xTexelC${c * 2};
|
|
int xTexelC${c * 2}Ready;
|
|
vec4 xTexelC${c * 2 + 1};
|
|
int xTexelC${c * 2 + 1}Ready;
|
|
vec4 xC${c};`;
|
|
}
|
|
mainLoop += `
|
|
for (int r = 0; r < ${filterHeight}; r++) {
|
|
for (int d1 = 0; d1 < ${convInfo.inChannels}; d1 += 2) {
|
|
`;
|
|
for (let c = 0; c < filterWidth; c++) {
|
|
mainLoop += `
|
|
xTexelC${c * 2} = vec4(0.0);
|
|
xTexelC${c * 2}Ready = 0;
|
|
xTexelC${c * 2 + 1} = vec4(0.0);
|
|
xTexelC${c * 2 + 1}Ready = 0;
|
|
xC${c} = vec4(0.0);`;
|
|
}
|
|
mainLoop += `
|
|
xR = xRCorner + r * dilations[0];
|
|
if (xR >=0 && xR < inDims[0]) {
|
|
`;
|
|
for (let texelC = 0; texelC < (texelsAcross + 1) / 2; texelC++) {
|
|
const colIndex = texelC * 2;
|
|
mainLoop += `
|
|
xC = xCCorner + ${colIndex * dilationWidth};
|
|
`;
|
|
if (strideWidth === 1) {
|
|
if (colIndex < filterWidth) {
|
|
if (padLeft % 2 === 1) {
|
|
mainLoop += `
|
|
xCOffset = xC + 1;
|
|
if (xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xCOffset, d1);
|
|
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
`;
|
|
if (dilationWidth === 1 && colIndex > 0) {
|
|
mainLoop += `
|
|
xC${colIndex} = vec4(xTexelC${colIndex - 2}.zw, xTexelC${colIndex}.xy);
|
|
`;
|
|
} else {
|
|
mainLoop += `
|
|
xCOffset = xC + 1 - 2;
|
|
|
|
if (xCOffset >= 0 && xCOffset < inDims[1]) {
|
|
previous = getX(batch, xR, xCOffset, d1);
|
|
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
previous.zw = vec2(0.0);
|
|
}
|
|
|
|
xC${colIndex} = vec4(previous.zw, xTexelC${colIndex}.xy);
|
|
} else {
|
|
xC${colIndex} = vec4(0.0, 0.0, xTexelC${colIndex}.xy);
|
|
}
|
|
`;
|
|
}
|
|
} else {
|
|
mainLoop += `
|
|
if (xC >= 0 && xC < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xC, d1);
|
|
if (xC + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex} = xTexelC${colIndex};
|
|
`;
|
|
}
|
|
if (colIndex + 1 < filterWidth) {
|
|
const nextTexelOffset = padLeft % 2 === 0 ? nearestLargerEven(dilationWidth) : dilationWidth;
|
|
if (dilationWidth % 2 === 0 && padLeft % 2 === 1 || dilationWidth % 2 !== 0 && padLeft % 2 !== 1) {
|
|
mainLoop += `
|
|
xCOffset = xC + imod(pads[1], 2) + ${nextTexelOffset};
|
|
|
|
if (xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xCOffset, d1);
|
|
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
`;
|
|
if (dilationWidth > 1) {
|
|
mainLoop += `
|
|
xCOffset -= 2;
|
|
if (xCOffset >= 0 && xCOffset < inDims[1]) {
|
|
previous = getX(batch, xR, xCOffset, d1);
|
|
xC${colIndex + 1} = vec4(previous.zw, xTexelC${colIndex + 1}.xy);
|
|
} else {
|
|
xC${colIndex + 1} = vec4(0.0, 0.0, xTexelC${colIndex + 1}.xy);
|
|
}
|
|
`;
|
|
} else {
|
|
mainLoop += `
|
|
xC${colIndex + 1} = vec4(xTexelC${colIndex}.zw, xTexelC${colIndex + 1}.xy);
|
|
`;
|
|
}
|
|
} else {
|
|
if (nextTexelOffset === 1) {
|
|
mainLoop += `
|
|
xC${colIndex + 1} = xTexelC${colIndex};
|
|
`;
|
|
} else {
|
|
mainLoop += `
|
|
xCOffset = xC + ${nextTexelOffset};
|
|
|
|
if (xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xCOffset, d1);
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex + 1} = xTexelC${colIndex + 1};
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
} else {
|
|
if (colIndex < filterWidth) {
|
|
if (padLeft % 2 === 1) {
|
|
mainLoop += `
|
|
xCOffset = xC + 1 - strides[1];
|
|
if(xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xCOffset, d1);
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
|
|
if(xC + 1 >= 0 && xC + 1 < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xC + 1, d1);
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xC + 2 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex} = vec4(xTexelC${colIndex}.zw, xTexelC${colIndex + 1}.zw);
|
|
`;
|
|
if (colIndex + 1 < filterWidth) {
|
|
mainLoop += `
|
|
final = vec4(0.0);
|
|
xCOffset = xC + 1 + strides[1];
|
|
if(xCOffset >= 0 && xCOffset < inDims[1]) {
|
|
final = getX(batch, xR, xCOffset, d1);
|
|
}
|
|
xC${colIndex + 1} = vec4(xTexelC${colIndex + 1}.xy, final.xy);
|
|
`;
|
|
}
|
|
} else {
|
|
mainLoop += `
|
|
if(xC >= 0 && xC < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xC, d1);
|
|
if (xC + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
|
|
xCOffset = xC + strides[1];
|
|
if(xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xCOffset, d1);
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex} = vec4(
|
|
xTexelC${colIndex}.xy, xTexelC${colIndex + 1}.xy);
|
|
`;
|
|
if (colIndex + 1 < filterWidth) {
|
|
mainLoop += `
|
|
xC${colIndex + 1} = vec4(xTexelC${colIndex}.zw, xTexelC${colIndex + 1}.zw);
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if (colIndex < filterWidth) {
|
|
mainLoop += `
|
|
wTexel = getW(r, ${colIndex}, d1, d2);
|
|
dotProd += xC${colIndex}.xxzz * vec4(wTexel.xy, wTexel.xy);
|
|
if(d1 + 1 < ${convInfo.inChannels}) {
|
|
dotProd += xC${colIndex}.yyww * vec4(wTexel.zw, wTexel.zw);
|
|
}
|
|
`;
|
|
if (colIndex + 1 < filterWidth) {
|
|
mainLoop += `
|
|
wTexel = getW(r, ${colIndex + 1}, d1, d2);
|
|
dotProd += xC${colIndex + 1}.xxzz * vec4(wTexel.xy, wTexel.xy);
|
|
if(d1 + 1 < ${convInfo.inChannels}) {
|
|
dotProd += xC${colIndex + 1}.yyww * vec4(wTexel.zw, wTexel.zw);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
mainLoop += `
|
|
}
|
|
`;
|
|
mainLoop += `
|
|
}
|
|
`;
|
|
mainLoop += `
|
|
}
|
|
`;
|
|
let activationSnippet = "", applyActivationSnippet = "";
|
|
if (activation) {
|
|
if (hasPreluActivation) {
|
|
activationSnippet = `vec4 activation(vec4 a) {
|
|
vec4 b = getPreluActivationWeightsAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else if (hasLeakyReluAlpha) {
|
|
activationSnippet = `vec4 activation(vec4 a) {
|
|
vec4 b = getLeakyreluAlphaAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else {
|
|
activationSnippet = `vec4 activation(vec4 x) {
|
|
${activation}
|
|
}`;
|
|
}
|
|
applyActivationSnippet = `result = activation(result);`;
|
|
}
|
|
const addBiasSnippet = addBias ? "result += getBiasAtOutCoords();" : "";
|
|
if (addBias) {
|
|
this.variableNames.push("bias");
|
|
}
|
|
if (hasPreluActivation) {
|
|
this.variableNames.push("preluActivationWeights");
|
|
}
|
|
if (hasLeakyReluAlpha) {
|
|
this.variableNames.push("leakyreluAlpha");
|
|
}
|
|
this.userCode = `
|
|
${activationSnippet}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
ivec2 xRCCorner = coords.yz * strides - pads;
|
|
int d2 = coords.w;
|
|
int xRCorner = xRCCorner.x;
|
|
int xCCorner = xRCCorner.y;
|
|
|
|
//intialize dotProd with a small epsilon seems to reduce GPU accuracy loss.
|
|
vec4 dotProd = vec4(0.000000000000001);
|
|
|
|
${mainLoop}
|
|
|
|
vec4 result = dotProd - vec4(0.000000000000001);
|
|
${addBiasSnippet}
|
|
${applyActivationSnippet}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Im2ColPackedProgram {
|
|
constructor(outputShape, convInfo) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [
|
|
{ name: "inputShape", type: "ivec4" },
|
|
{ name: "pad", type: "ivec2" },
|
|
{ name: "stride", type: "ivec2" },
|
|
{ name: "dilation", type: "ivec2" },
|
|
{ name: "inChannels", type: "int" },
|
|
{ name: "itemsPerBlockRow", type: "int" },
|
|
{ name: "outWidth", type: "int" }
|
|
];
|
|
this.outputShape = outputShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const { dataFormat } = convInfo;
|
|
const glsl = getGlslDifferences();
|
|
const isChannelsLast = dataFormat === "channelsLast";
|
|
const rowDim = isChannelsLast ? 1 : 2;
|
|
const colDim = isChannelsLast ? 2 : 3;
|
|
const boundsCheckingSnippet = this.enableShapeUniforms ? "if(blockIndex < outShape[2] && pos < outShape[1]) {" : `if(blockIndex < ${outputShape[2]} && pos < ${outputShape[1]}) {`;
|
|
let unrolled = ``;
|
|
for (let row = 0; row <= 1; row++) {
|
|
for (let col = 0; col <= 1; col++) {
|
|
unrolled += `
|
|
blockIndex = rc.z + ${col};
|
|
pos = rc.y + ${row};
|
|
|
|
${boundsCheckingSnippet}
|
|
offsetY = int(blockIndex / outWidth) * stride[0] - pad[0];
|
|
d0 = offsetY + dilation[0] * (pos / itemsPerBlockRow);
|
|
|
|
if(d0 < inputShape[${rowDim}] && d0 >= 0) {
|
|
// Use custom imod instead mod. On Intel GPU, mod may generate
|
|
// unexpected value.
|
|
// https://github.com/tensorflow/tfjs/issues/5447
|
|
offsetX = imod(blockIndex, outWidth) * stride[1] - pad[1];
|
|
d1 = offsetX + dilation[1] * (imod(pos, itemsPerBlockRow) /
|
|
inChannels);
|
|
|
|
if(d1 < inputShape[${colDim}] && d1 >= 0) {
|
|
|
|
ch = imod(pos, inChannels);
|
|
|
|
if (${isChannelsLast}) {
|
|
innerDims = vec2(d1, ch);
|
|
result[${row * 2 + col}] = getChannel(
|
|
getA(rc.x, d0, int(innerDims.x),
|
|
int(innerDims.y)), innerDims);
|
|
} else {
|
|
innerDims = vec2(d0, d1);
|
|
result[${row * 2 + col}] = getChannel(
|
|
getA(rc.x, ch, int(innerDims.x),
|
|
int(innerDims.y)), innerDims);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
this.userCode = `
|
|
void main() {
|
|
ivec3 rc = getOutputCoords();
|
|
|
|
vec4 result = vec4(0);
|
|
|
|
int blockIndex, pos, offsetY, d0, offsetX, d1, ch;
|
|
vec2 innerDims;
|
|
|
|
${unrolled}
|
|
|
|
${glsl.output} = result;
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getShapeForBatchMatMul(shape, isChannelsLast) {
|
|
const length = shape.length;
|
|
if (length >= 3) {
|
|
return isChannelsLast ? [
|
|
...shape.slice(0, -3),
|
|
shape[length - 3] * shape[length - 2],
|
|
shape[length - 1]
|
|
/* channel */
|
|
] : [
|
|
...shape.slice(0, -3),
|
|
shape[length - 3],
|
|
shape[length - 2] * shape[length - 1]
|
|
/* height * width */
|
|
];
|
|
} else if (!isChannelsLast && length === 1 && shape[0] > 1) {
|
|
return [shape[0], 1];
|
|
} else {
|
|
return null;
|
|
}
|
|
}
|
|
function conv2dByMatMul({ x, filter, convInfo, backend: backend2, bias = null, preluActivationWeights = null, leakyreluAlpha = 0, activation = null }) {
|
|
const xShape = x.shape;
|
|
const xTexData = backend2.texData.get(x.dataId);
|
|
const sharedMatMulDim = convInfo.inChannels;
|
|
const outerShapeX = xShape[0] * xShape[1] * xShape[2];
|
|
const outerShapeFilter = convInfo.outChannels;
|
|
const isChannelsLast = convInfo.dataFormat === "channelsLast";
|
|
const transposeA = false;
|
|
const transposeB = false;
|
|
let out;
|
|
const intermediates = [];
|
|
if (preluActivationWeights != null) {
|
|
const targetShape = getShapeForBatchMatMul(preluActivationWeights.shape, isChannelsLast);
|
|
if (targetShape != null) {
|
|
preluActivationWeights = reshape$1({
|
|
inputs: { x: preluActivationWeights },
|
|
backend: backend2,
|
|
attrs: { shape: targetShape }
|
|
});
|
|
intermediates.push(preluActivationWeights);
|
|
}
|
|
}
|
|
if (bias != null) {
|
|
const targetShape = getShapeForBatchMatMul(bias.shape, isChannelsLast);
|
|
if (targetShape != null) {
|
|
bias = reshape$1({ inputs: { x: bias }, backend: backend2, attrs: { shape: targetShape } });
|
|
intermediates.push(bias);
|
|
}
|
|
}
|
|
const batchMatMulWillBeUnpacked = (outerShapeX === 1 || outerShapeFilter === 1) && sharedMatMulDim > MATMUL_SHARED_DIM_THRESHOLD;
|
|
const canOptimize = !batchMatMulWillBeUnpacked && xTexData.isPacked && isChannelsLast && xTexData.texture != null && xShape[2] % 2 !== 0 && arraysEqual(xTexData.shape.slice(-3), xShape.slice(-3));
|
|
if (canOptimize) {
|
|
const targetShape = xShape[0] * xShape[1] * (xShape[2] + 1);
|
|
const xReshaped = {
|
|
dataId: x.dataId,
|
|
shape: [1, targetShape, convInfo.inChannels],
|
|
dtype: x.dtype
|
|
};
|
|
const originalXTexDataShape = xTexData.shape;
|
|
xTexData.shape = xTexData.shape.slice();
|
|
xTexData.shape[xTexData.shape.length - 2]++;
|
|
assert(isReshapeFree(xTexData.shape, xReshaped.shape), () => `packed reshape ${xTexData.shape} to ${xReshaped.shape} isn't free`);
|
|
const filterReshaped = reshape$1({
|
|
inputs: { x: filter },
|
|
backend: backend2,
|
|
attrs: { shape: [1, convInfo.inChannels, convInfo.outChannels] }
|
|
});
|
|
intermediates.push(filterReshaped);
|
|
const pointwiseConv = batchMatMulImpl({
|
|
a: xReshaped,
|
|
b: filterReshaped,
|
|
backend: backend2,
|
|
transposeA,
|
|
transposeB,
|
|
bias,
|
|
activation,
|
|
preluActivationWeights,
|
|
leakyreluAlpha
|
|
});
|
|
const pointwiseConvTexData = backend2.texData.get(pointwiseConv.dataId);
|
|
assert(pointwiseConvTexData.isPacked, () => "batchMatMul result is expected to be packed");
|
|
xTexData.shape = originalXTexDataShape;
|
|
pointwiseConvTexData.shape = convInfo.outShape;
|
|
out = identity({ inputs: { x: pointwiseConv }, backend: backend2 });
|
|
out.shape = convInfo.outShape;
|
|
intermediates.push(pointwiseConv);
|
|
} else {
|
|
const numCols = convInfo.outHeight * convInfo.outWidth;
|
|
const xReshaped = reshape$1({
|
|
inputs: { x },
|
|
backend: backend2,
|
|
attrs: {
|
|
shape: isChannelsLast ? [convInfo.batchSize, numCols, convInfo.inChannels] : [convInfo.batchSize, convInfo.inChannels, numCols]
|
|
}
|
|
});
|
|
const filterReshaped = reshape$1({
|
|
inputs: { x: filter },
|
|
backend: backend2,
|
|
attrs: { shape: [1, convInfo.inChannels, convInfo.outChannels] }
|
|
});
|
|
const result = batchMatMulImpl({
|
|
a: isChannelsLast ? xReshaped : filterReshaped,
|
|
b: isChannelsLast ? filterReshaped : xReshaped,
|
|
transposeA: !isChannelsLast,
|
|
transposeB,
|
|
backend: backend2,
|
|
bias,
|
|
activation,
|
|
preluActivationWeights,
|
|
leakyreluAlpha
|
|
});
|
|
out = reshape$1({ inputs: { x: result }, backend: backend2, attrs: { shape: convInfo.outShape } });
|
|
intermediates.push(xReshaped);
|
|
intermediates.push(filterReshaped);
|
|
intermediates.push(result);
|
|
}
|
|
for (const i of intermediates) {
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
}
|
|
return out;
|
|
}
|
|
function conv2dWithIm2Row({ x, filter, convInfo, backend: backend2, bias = null, preluActivationWeights = null, leakyreluAlpha = 0, activation = null }) {
|
|
const { filterWidth, filterHeight, inChannels, outWidth, outHeight, dataFormat } = convInfo;
|
|
const isChannelsLast = dataFormat === "channelsLast";
|
|
const sharedDim = filterWidth * filterHeight * inChannels;
|
|
const numCols = outHeight * outWidth;
|
|
const x2ColShape = [convInfo.batchSize, sharedDim, numCols];
|
|
const transposeA = true;
|
|
const transposeB = false;
|
|
const intermediates = [];
|
|
if (preluActivationWeights != null) {
|
|
const targetShape = getShapeForBatchMatMul(preluActivationWeights.shape, isChannelsLast);
|
|
if (targetShape != null) {
|
|
preluActivationWeights = reshape$1({
|
|
inputs: { x: preluActivationWeights },
|
|
backend: backend2,
|
|
attrs: { shape: targetShape }
|
|
});
|
|
intermediates.push(preluActivationWeights);
|
|
}
|
|
}
|
|
if (bias != null) {
|
|
const targetShape = getShapeForBatchMatMul(bias.shape, isChannelsLast);
|
|
if (targetShape != null) {
|
|
bias = reshape$1({ inputs: { x: bias }, backend: backend2, attrs: { shape: targetShape } });
|
|
intermediates.push(bias);
|
|
}
|
|
}
|
|
const w2Row = reshape$1({
|
|
inputs: { x: filter },
|
|
backend: backend2,
|
|
attrs: { shape: [1, sharedDim, sizeFromShape(filter.shape) / sharedDim] }
|
|
});
|
|
intermediates.push(w2Row);
|
|
const im2ColProgram = new Im2ColPackedProgram(x2ColShape, convInfo);
|
|
const customValues = [
|
|
x.shape,
|
|
[convInfo.padInfo.top, convInfo.padInfo.left],
|
|
[convInfo.strideHeight, convInfo.strideWidth],
|
|
[convInfo.dilationHeight, convInfo.dilationWidth],
|
|
[convInfo.inChannels],
|
|
[convInfo.filterWidth * convInfo.inChannels],
|
|
[convInfo.outWidth]
|
|
];
|
|
const im2Col = backend2.runWebGLProgram(im2ColProgram, [x], "float32", customValues);
|
|
const im2ColReshaped = reshape$1({ inputs: { x: im2Col }, backend: backend2, attrs: { shape: x2ColShape } });
|
|
intermediates.push(im2Col);
|
|
intermediates.push(im2ColReshaped);
|
|
const hasBias = bias != null;
|
|
const hasPreluActivationWeights = preluActivationWeights != null;
|
|
const hasLeakyreluAlpha = activation === "leakyrelu";
|
|
const fusedActivation = activation ? mapActivationToShaderProgram(activation, true) : null;
|
|
const matmulProgram = new MatMulPackedProgram(isChannelsLast ? im2ColReshaped.shape : w2Row.shape, isChannelsLast ? w2Row.shape : im2ColReshaped.shape, isChannelsLast ? [convInfo.batchSize, numCols, convInfo.outChannels] : [convInfo.batchSize, convInfo.outChannels, numCols], transposeA, transposeB, hasBias, fusedActivation, hasPreluActivationWeights, hasLeakyreluAlpha);
|
|
const inputs = isChannelsLast ? [im2ColReshaped, w2Row] : [w2Row, im2ColReshaped];
|
|
if (bias) {
|
|
inputs.push(bias);
|
|
}
|
|
if (hasPreluActivationWeights) {
|
|
inputs.push(preluActivationWeights);
|
|
}
|
|
if (hasLeakyreluAlpha) {
|
|
const $leakyreluAlpha = backend2.makeTensorInfo([], "float32", createScalarValue(leakyreluAlpha, "float32"));
|
|
inputs.push($leakyreluAlpha);
|
|
intermediates.push($leakyreluAlpha);
|
|
}
|
|
const product = backend2.runWebGLProgram(matmulProgram, inputs, "float32");
|
|
const out = reshape$1({ inputs: { x: product }, backend: backend2, attrs: { shape: convInfo.outShape } });
|
|
intermediates.push(product);
|
|
for (const i of intermediates) {
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
}
|
|
return out;
|
|
}
|
|
function conv2d(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dilations, dimRoundingMode } = attrs;
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(x.shape, filter.shape, strides, dilations, pad2, dimRoundingMode, false, $dataFormat);
|
|
let out;
|
|
if (convInfo.filterHeight === 1 && convInfo.filterWidth === 1 && convInfo.dilationHeight === 1 && convInfo.dilationWidth === 1 && convInfo.strideHeight === 1 && convInfo.strideWidth === 1 && (convInfo.padInfo.type === "SAME" || convInfo.padInfo.type === "VALID")) {
|
|
out = conv2dByMatMul({ x, filter, convInfo, backend: backend2 });
|
|
} else if (convInfo.strideWidth <= 2 && $dataFormat === "channelsLast" && env().getBool("WEBGL_EXP_CONV")) {
|
|
const program = new Conv2DPackedProgram(convInfo);
|
|
const customValues = [
|
|
[convInfo.padInfo.top, convInfo.padInfo.left],
|
|
[convInfo.strideHeight, convInfo.strideWidth],
|
|
[convInfo.dilationHeight, convInfo.dilationWidth],
|
|
[convInfo.inHeight, convInfo.inWidth]
|
|
];
|
|
out = backend2.runWebGLProgram(program, [x, filter], "float32", customValues);
|
|
} else if (env().getBool("WEBGL_CONV_IM2COL")) {
|
|
out = conv2dWithIm2Row({ x, filter, convInfo, backend: backend2 });
|
|
} else {
|
|
const program = new Conv2DProgram(convInfo);
|
|
out = backend2.runWebGLProgram(program, [x, filter], "float32");
|
|
}
|
|
const outReshaped = reshape$1({ inputs: { x: out }, backend: backend2, attrs: { shape: convInfo.outShape } });
|
|
backend2.disposeIntermediateTensorInfo(out);
|
|
return outReshaped;
|
|
}
|
|
const conv2DConfig$1 = {
|
|
kernelName: Conv2D,
|
|
backendName: "webgl",
|
|
kernelFunc: conv2d
|
|
};
|
|
class Conv2DDerFilterProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["x", "dy"];
|
|
this.outputShape = convInfo.filterShape;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const isChannelsLast = convInfo.dataFormat === "channelsLast";
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int wR = coords.x;
|
|
int wC = coords.y;
|
|
int d1 = coords.z;
|
|
int d2 = coords.w;
|
|
|
|
// Convolve x(?, ?, d1) with dy(:, :, d2) to get dw(wR, wC, d1, d2).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
|
|
for (int b = 0; b < ${convInfo.batchSize}; b++) {
|
|
for (int yR = 0; yR < ${convInfo.outHeight}; yR++) {
|
|
int xR = wR + yR * ${strideHeight} - ${padTop};
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int yC = 0; yC < ${convInfo.outWidth}; yC++) {
|
|
int xC = wC + yC * ${strideWidth} - ${padLeft};
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
${isChannelsLast ? `float dyValue = getDy(b, yR, yC, d2);
|
|
float xValue = getX(b, xR, xC, d1);
|
|
dotProd += (xValue * dyValue);` : `float dyValue = getDy(b, d2, yR, yC);
|
|
float xValue = getX(b, d1, xR, xC);
|
|
dotProd += (xValue * dyValue);`}
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Conv2DDerInputProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy", "W"];
|
|
this.outputShape = convInfo.inShape;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const isChannelsLast = convInfo.dataFormat === "channelsLast";
|
|
const padTop = filterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = filterWidth - 1 - convInfo.padInfo.left;
|
|
const rowDim = isChannelsLast ? 1 : 2;
|
|
const colDim = isChannelsLast ? 2 : 3;
|
|
const channelDim = isChannelsLast ? 3 : 1;
|
|
this.userCode = `
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int d1 = coords[${channelDim}];
|
|
|
|
ivec2 dyCorner = ivec2(coords[${rowDim}], coords[${colDim}]) - pads;
|
|
int dyRCorner = dyCorner.x;
|
|
int dyCCorner = dyCorner.y;
|
|
|
|
// Convolve dy(?, ?, d2) with w(:, :, d1, d2) to compute dx(xR, xC, d1).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 || fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
int wRPerm = ${filterHeight} - 1 - wR;
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
int wCPerm = ${filterWidth} - 1 - wC;
|
|
|
|
for (int d2 = 0; d2 < ${convInfo.outChannels}; d2++) {
|
|
|
|
if (${isChannelsLast}) {
|
|
float xValue = getDy(batch, idyR, idyC, d2);
|
|
float wValue = getW(wRPerm, wCPerm, d1, d2);
|
|
dotProd += xValue * wValue;
|
|
} else {
|
|
float xValue = getDy(batch, d2, idyR, idyC);
|
|
float wValue = getW(wRPerm, wCPerm, d1, d2);
|
|
dotProd += xValue * wValue;
|
|
}
|
|
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Conv3DDerFilterProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["x", "dy"];
|
|
this.outputShape = convInfo.filterShape;
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const padFront = convInfo.padInfo.front;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int wF = coords.x;
|
|
int wR = coords.y;
|
|
int wC = coords.z;
|
|
int d1 = coords.w;
|
|
int d2 = coords.u;
|
|
|
|
float dotProd = 0.0;
|
|
|
|
for (int b = 0; b < ${convInfo.batchSize}; b++) {
|
|
for (int yF = 0; yF < ${convInfo.outDepth}; yF++) {
|
|
int xF = wF + yF * ${strideDepth} - ${padFront};
|
|
|
|
if (xF < 0 || xF >= ${convInfo.inDepth}) {
|
|
continue;
|
|
}
|
|
|
|
for (int yR = 0; yR < ${convInfo.outHeight}; yR++) {
|
|
int xR = wR + yR * ${strideHeight} - ${padTop};
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int yC = 0; yC < ${convInfo.outWidth}; yC++) {
|
|
int xC = wC + yC * ${strideWidth} - ${padLeft};
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
float dyValue = getDy(b, yF, yR, yC, d2);
|
|
float xValue = getX(b, xF, xR, xC, d1);
|
|
dotProd += (xValue * dyValue);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class Conv3DDerInputProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy", "W"];
|
|
this.outputShape = convInfo.inShape;
|
|
const filterDepth = convInfo.filterDepth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const padFront = filterDepth - 1 - convInfo.padInfo.front;
|
|
const padTop = filterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = filterWidth - 1 - convInfo.padInfo.left;
|
|
this.userCode = `
|
|
const ivec3 pads = ivec3(${padFront}, ${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int d1 = coords.u;
|
|
|
|
|
|
ivec3 dyCorner = ivec3(coords.y, coords.z, coords.w) - pads;
|
|
int dyFCorner = dyCorner.x;
|
|
int dyRCorner = dyCorner.y;
|
|
int dyCCorner = dyCorner.z;
|
|
|
|
float dotProd = 0.0;
|
|
for (int wF = 0; wF < ${filterDepth}; wF++) {
|
|
float dyF = float(dyFCorner + wF) / ${strideDepth}.0;
|
|
|
|
if (dyF < 0.0 || dyF >= ${convInfo.outDepth}.0 || fract(dyF) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyF = int(dyF);
|
|
|
|
int wFPerm = ${filterDepth} - 1 - wF;
|
|
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 ||
|
|
fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
int wRPerm = ${filterHeight} - 1 - wR;
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
int wCPerm = ${filterWidth} - 1 - wC;
|
|
|
|
for (int d2 = 0; d2 < ${convInfo.outChannels}; d2++) {
|
|
float xValue = getDy(batch, idyF, idyR, idyC, d2);
|
|
float wValue = getW(wFPerm, wRPerm, wCPerm, d1, d2);
|
|
dotProd += xValue * wValue;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function conv2DBackpropFilter$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, dy } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dimRoundingMode, filterShape } = attrs;
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(x.shape, filterShape, strides, 1, pad2, dimRoundingMode, false, $dataFormat);
|
|
const program = new Conv2DDerFilterProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [x, dy], "float32");
|
|
}
|
|
const conv2DBackpropFilterConfig$1 = {
|
|
kernelName: Conv2DBackpropFilter,
|
|
backendName: "webgl",
|
|
kernelFunc: conv2DBackpropFilter$1
|
|
};
|
|
class Conv2DDerInputPackedProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy", "W"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [
|
|
{ name: "strides", type: "vec2" }
|
|
];
|
|
this.outputShape = convInfo.inShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const padTop = filterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = filterWidth - 1 - convInfo.padInfo.left;
|
|
this.userCode = `
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int d1 = coords[3];
|
|
|
|
ivec2 dyCorner = ivec2(coords[1], coords[2]) - pads;
|
|
int dyRCorner = dyCorner.x;
|
|
int dyCCorner = dyCorner.y;
|
|
|
|
vec4 result = vec4(0.);
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
float dyR = float(dyRCorner + wR) / strides[0];
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 || fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
int wRPerm = ${filterHeight} - 1 - wR;
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
int wCPerm = ${filterWidth} - 1 - wC;
|
|
|
|
float dyC = float(dyCCorner + wC) / strides[1];
|
|
bool idyCVal = (dyC >= 0.0) && (dyC < ${convInfo.outWidth}.0)
|
|
&& (fract(dyC) == 0.0);
|
|
int idyC = int(dyC);
|
|
|
|
float dyC2 = float(dyCCorner + wC + 1) / strides[1];
|
|
bool idyCVal2 = (dyC2 >= 0.0) && (dyC2 < ${convInfo.outWidth}.0)
|
|
&& (fract(dyC2) == 0.0);
|
|
int idyC2 = int(dyC2);
|
|
|
|
if (idyCVal && idyCVal2) {
|
|
for (int d2 = 0; d2 < ${convInfo.outChannels}; d2 += 2) {
|
|
vec4 wValue = getW(wRPerm, wCPerm, d1, d2);
|
|
vec4 dySample = getDy(batch, idyR, idyC, d2);
|
|
vec4 dySample2 = (idyC / 2 == idyC2 / 2) ?
|
|
dySample : getDy(batch, idyR, idyC2, d2);
|
|
|
|
vec2 dyValue = mod(float(idyC), 2.) == 0. ?
|
|
dySample.xy : dySample.zw;
|
|
result.xy += vec2(dot(dyValue, wValue.xy),
|
|
dot(dyValue, wValue.zw));
|
|
|
|
dyValue = mod(float(idyC2), 2.) == 0. ?
|
|
dySample2.xy : dySample2.zw;
|
|
result.zw += vec2(dot(dyValue, wValue.xy),
|
|
dot(dyValue, wValue.zw));
|
|
}
|
|
} else if (idyCVal) {
|
|
for (int d2 = 0; d2 < ${convInfo.outChannels}; d2 += 2) {
|
|
vec4 wValue = getW(wRPerm, wCPerm, d1, d2);
|
|
vec4 dySample = getDy(batch, idyR, idyC, d2);
|
|
vec2 dyValue = mod(float(idyC), 2.) == 0. ?
|
|
dySample.xy : dySample.zw;
|
|
result.xy += vec2(dot(dyValue, wValue.xy),
|
|
dot(dyValue, wValue.zw));
|
|
}
|
|
} else if (idyCVal2) {
|
|
for (int d2 = 0; d2 < ${convInfo.outChannels}; d2 += 2) {
|
|
vec4 wValue = getW(wRPerm, wCPerm, d1, d2);
|
|
vec4 dySample = getDy(batch, idyR, idyC2, d2);
|
|
vec2 dyValue = mod(float(idyC2), 2.) == 0. ?
|
|
dySample.xy : dySample.zw;
|
|
result.zw += vec2(dot(dyValue, wValue.xy),
|
|
dot(dyValue, wValue.zw));
|
|
}
|
|
}
|
|
}
|
|
}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function conv2DBackpropInput$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, filter } = inputs;
|
|
const { inputShape, strides, pad: pad2, dataFormat, dimRoundingMode } = attrs;
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(inputShape, filter.shape, strides, 1, pad2, dimRoundingMode, false, $dataFormat);
|
|
if (env().getBool("WEBGL_PACK_CONV2DTRANSPOSE") && $dataFormat === "channelsLast") {
|
|
const customValues = [
|
|
[convInfo.strideHeight, convInfo.strideWidth]
|
|
];
|
|
const program = new Conv2DDerInputPackedProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [dy, filter], "float32", customValues);
|
|
} else {
|
|
const program = new Conv2DDerInputProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [dy, filter], "float32");
|
|
}
|
|
}
|
|
const conv2DBackpropInputConfig$1 = {
|
|
kernelName: Conv2DBackpropInput,
|
|
backendName: "webgl",
|
|
kernelFunc: conv2DBackpropInput$1
|
|
};
|
|
function conv3D$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dilations } = attrs;
|
|
const convInfo = computeConv3DInfo(x.shape, filter.shape, strides, dilations, pad2);
|
|
const program = new Conv3DProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [x, filter], "float32");
|
|
}
|
|
const conv3DConfig$1 = {
|
|
kernelName: Conv3D,
|
|
backendName: "webgl",
|
|
kernelFunc: conv3D$1
|
|
};
|
|
function conv3DBackpropFilterV2$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, dy } = inputs;
|
|
const { strides, pad: pad2, filterShape } = attrs;
|
|
const convInfo = computeConv3DInfo(x.shape, filterShape, strides, 1, pad2);
|
|
const program = new Conv3DDerFilterProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [x, dy], "float32");
|
|
}
|
|
const conv3DBackpropFilterV2Config$1 = {
|
|
kernelName: Conv3DBackpropFilterV2,
|
|
backendName: "webgl",
|
|
kernelFunc: conv3DBackpropFilterV2$1
|
|
};
|
|
function conv3DBackpropInput(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, filter } = inputs;
|
|
const { pad: pad2, strides, inputShape } = attrs;
|
|
const convInfo = computeConv3DInfo(inputShape, filter.shape, strides, 1, pad2);
|
|
const program = new Conv3DDerInputProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [dy, filter], "float32");
|
|
}
|
|
const conv3DBackpropInputConfig = {
|
|
kernelName: Conv3DBackpropInputV2,
|
|
backendName: "webgl",
|
|
kernelFunc: conv3DBackpropInput
|
|
};
|
|
const COS = CHECK_NAN_SNIPPET_UNARY + `
|
|
return cos(x);
|
|
`;
|
|
const COS_PACKED = `
|
|
vec4 result = cos(x);
|
|
bvec4 isNaN = isnan(x);
|
|
${CHECK_NAN_SNIPPET_PACKED}
|
|
return result;
|
|
`;
|
|
const cos$1 = unaryKernelFunc({ opSnippet: COS, packedOpSnippet: COS_PACKED });
|
|
const cosConfig$1 = {
|
|
kernelName: Cos,
|
|
backendName: "webgl",
|
|
kernelFunc: cos$1
|
|
};
|
|
const COSH = `
|
|
float e2x = exp(-x);
|
|
return (e2x + 1.0 / e2x) / 2.0;
|
|
`;
|
|
const cosh$1 = unaryKernelFunc({ opSnippet: COSH });
|
|
const coshConfig$1 = {
|
|
kernelName: Cosh,
|
|
backendName: "webgl",
|
|
kernelFunc: cosh$1
|
|
};
|
|
class CropAndResizeProgram {
|
|
constructor(imageShape, boxShape, cropSize, method, extrapolationValue) {
|
|
this.variableNames = ["Image", "Boxes", "BoxInd"];
|
|
this.outputShape = [];
|
|
const [batch, imageHeight, imageWidth, depth] = imageShape;
|
|
const [numBoxes] = boxShape;
|
|
const [cropHeight, cropWidth] = cropSize;
|
|
this.outputShape = [numBoxes, cropHeight, cropWidth, depth];
|
|
const methodId = method === "bilinear" ? 1 : 0;
|
|
const [inputHeightFloat, inputWidthFloat] = [`${imageHeight - 1}.0`, `${imageWidth - 1}.0`];
|
|
const [heightRatio, heightScale, inY] = cropHeight > 1 ? [
|
|
`${(imageHeight - 1) / (cropHeight - 1)}`,
|
|
"(y2-y1) * height_ratio",
|
|
`y1*${inputHeightFloat} + float(y)*(height_scale)`
|
|
] : [
|
|
"0.0",
|
|
"0.0",
|
|
`0.5 * (y1+y2) * ${inputHeightFloat}`
|
|
];
|
|
const [widthRatio, widthScale, inX] = cropWidth > 1 ? [
|
|
`${(imageWidth - 1) / (cropWidth - 1)}`,
|
|
"(x2-x1) * width_ratio",
|
|
`x1*${inputWidthFloat} + float(x)*(width_scale)`
|
|
] : [
|
|
"0.0",
|
|
"0.0",
|
|
`0.5 * (x1+x2) * ${inputWidthFloat}`
|
|
];
|
|
this.userCode = `
|
|
const float height_ratio = float(${heightRatio});
|
|
const float width_ratio = float(${widthRatio});
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int y = coords[1];
|
|
int x = coords[2];
|
|
int d = coords[3];
|
|
|
|
// get box vals
|
|
float y1 = getBoxes(b,0);
|
|
float x1 = getBoxes(b,1);
|
|
float y2 = getBoxes(b,2);
|
|
float x2 = getBoxes(b,3);
|
|
|
|
// get image in batch index
|
|
int bInd = round(getBoxInd(b));
|
|
if(bInd < 0 || bInd >= ${batch}) {
|
|
return;
|
|
}
|
|
|
|
float height_scale = ${heightScale};
|
|
float width_scale = ${widthScale};
|
|
|
|
float in_y = ${inY};
|
|
if( in_y < 0.0 || in_y > ${inputHeightFloat} ) {
|
|
setOutput(float(${extrapolationValue}));
|
|
return;
|
|
}
|
|
float in_x = ${inX};
|
|
if( in_x < 0.0 || in_x > ${inputWidthFloat} ) {
|
|
setOutput(float(${extrapolationValue}));
|
|
return;
|
|
}
|
|
|
|
vec2 sourceFracIndexCR = vec2(in_x,in_y);
|
|
if(${methodId} == 1) {
|
|
// Compute the four integer indices.
|
|
ivec2 sourceFloorCR = ivec2(sourceFracIndexCR);
|
|
ivec2 sourceCeilCR = ivec2(ceil(sourceFracIndexCR));
|
|
|
|
float topLeft = getImage(b, sourceFloorCR.y, sourceFloorCR.x, d);
|
|
float bottomLeft = getImage(b, sourceCeilCR.y, sourceFloorCR.x, d);
|
|
float topRight = getImage(b, sourceFloorCR.y, sourceCeilCR.x, d);
|
|
float bottomRight = getImage(b, sourceCeilCR.y, sourceCeilCR.x, d);
|
|
|
|
vec2 fracCR = sourceFracIndexCR - vec2(sourceFloorCR);
|
|
|
|
float top = topLeft + (topRight - topLeft) * fracCR.x;
|
|
float bottom = bottomLeft + (bottomRight - bottomLeft) * fracCR.x;
|
|
float newValue = top + (bottom - top) * fracCR.y;
|
|
setOutput(newValue);
|
|
} else {
|
|
// Compute the coordinators of nearest neighbor point.
|
|
ivec2 sourceNearestCR = ivec2(floor(
|
|
sourceFracIndexCR + vec2(0.5,0.5)));
|
|
float newValue = getImage(b, sourceNearestCR.y, sourceNearestCR.x, d);
|
|
setOutput(newValue);
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const cropAndResize$1 = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { image: image2, boxes, boxInd } = inputs;
|
|
const { cropSize, method, extrapolationValue } = attrs;
|
|
const program = new CropAndResizeProgram(image2.shape, boxes.shape, cropSize, method, extrapolationValue);
|
|
return backend2.runWebGLProgram(program, [image2, boxes, boxInd], "float32");
|
|
};
|
|
const cropAndResizeConfig$1 = {
|
|
kernelName: CropAndResize,
|
|
backendName: "webgl",
|
|
kernelFunc: cropAndResize$1
|
|
};
|
|
var CumOpType;
|
|
(function(CumOpType2) {
|
|
CumOpType2["Prod"] = "*";
|
|
CumOpType2["Sum"] = "+";
|
|
})(CumOpType || (CumOpType = {}));
|
|
class CumProgram {
|
|
constructor(op2, outputShape, exclusive, reverse2) {
|
|
this.op = op2;
|
|
this.outputShape = outputShape;
|
|
this.variableNames = ["x"];
|
|
this.customUniforms = [{ name: "index", type: "float" }];
|
|
const rank = this.outputShape.length;
|
|
const initVal = this.op === CumOpType.Prod ? "1.0" : "0.0";
|
|
const val = exclusive ? initVal : `getX(${getCoords(rank, "coords", this.op)})`;
|
|
const length = this.outputShape[this.outputShape.length - 1];
|
|
let condition = "";
|
|
let idxString = "";
|
|
if (exclusive) {
|
|
condition = reverse2 ? `end != ${length - 1}` : "end != 0";
|
|
idxString = reverse2 ? "end + 1" : "end - 1";
|
|
} else {
|
|
condition = reverse2 ? `end + pow2 < ${length}` : "end >= pow2";
|
|
idxString = reverse2 ? "end + pow2" : "end - pow2";
|
|
}
|
|
this.userCode = `
|
|
void main() {
|
|
${getCoordsDataType(rank)} coords = getOutputCoords();
|
|
int end = ${getFinalCoord(rank, "coords", this.op)};
|
|
float val = ${val};
|
|
int pow2 = int(pow(2.0, index));
|
|
if (${condition}) {
|
|
int idx = ${idxString};
|
|
${getFinalCoord(rank, "coords", this.op)} = idx;
|
|
val ${this.op}= getX(${getCoords(rank, "coords", this.op)});
|
|
}
|
|
setOutput(val);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getCoords(rank, name, op2) {
|
|
if (rank === 1) {
|
|
return `${name}`;
|
|
} else if (rank === 2) {
|
|
return `${name}.x, ${name}.y`;
|
|
} else if (rank === 3) {
|
|
return `${name}.x, ${name}.y, ${name}.z`;
|
|
} else if (rank === 4) {
|
|
return `${name}.x, ${name}.y, ${name}.z, ${name}.w`;
|
|
} else {
|
|
throw new Error(`Cumulative ${op2} for rank ${rank} is not yet supported`);
|
|
}
|
|
}
|
|
function getFinalCoord(rank, name, op2) {
|
|
if (rank === 1) {
|
|
return `${name}`;
|
|
} else if (rank === 2) {
|
|
return `${name}.y`;
|
|
} else if (rank === 3) {
|
|
return `${name}.z`;
|
|
} else if (rank === 4) {
|
|
return `${name}.w`;
|
|
} else {
|
|
throw new Error(`Cumulative ${op2} for rank ${rank} is not yet supported`);
|
|
}
|
|
}
|
|
function cumImpl(op2, x, backend2, axis, exclusive, reverse2) {
|
|
const xRank = x.shape.length;
|
|
const permutation = getAxesPermutation([axis], xRank);
|
|
let permutedX = x;
|
|
if (permutation != null) {
|
|
permutedX = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutation } });
|
|
}
|
|
const permutedAxis = getInnerMostAxes(1, xRank)[0];
|
|
if (permutedAxis !== xRank - 1) {
|
|
throw new Error(`WebGL cumprod shader expects an inner-most axis=${x.shape.length - 1} but got axis=${axis}`);
|
|
}
|
|
const size = permutedX.shape[permutedAxis];
|
|
let result = identity({ inputs: { x: permutedX }, backend: backend2 });
|
|
for (let i = 0; i <= Math.ceil(Math.log2(size)) - 1; i++) {
|
|
const program = new CumProgram(op2, permutedX.shape, false, reverse2);
|
|
const customValues = [[i]];
|
|
const prevResult = result;
|
|
result = backend2.runWebGLProgram(program, [result], result.dtype, customValues);
|
|
backend2.disposeIntermediateTensorInfo(prevResult);
|
|
}
|
|
if (exclusive) {
|
|
const program = new CumProgram(op2, permutedX.shape, exclusive, reverse2);
|
|
const prevResult = result;
|
|
result = backend2.runWebGLProgram(program, [result], result.dtype);
|
|
backend2.disposeIntermediateTensorInfo(prevResult);
|
|
}
|
|
if (permutation != null) {
|
|
const reversePermutation = getUndoAxesPermutation(permutation);
|
|
const reverseTransposedResult = transpose({ inputs: { x: result }, backend: backend2, attrs: { perm: reversePermutation } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
backend2.disposeIntermediateTensorInfo(permutedX);
|
|
return reverseTransposedResult;
|
|
}
|
|
return result;
|
|
}
|
|
function cumprod$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, exclusive, reverse: reverse2 } = attrs;
|
|
return cumImpl(CumOpType.Prod, x, backend2, axis, exclusive, reverse2);
|
|
}
|
|
const cumprodConfig$1 = {
|
|
kernelName: Cumprod,
|
|
backendName: "webgl",
|
|
kernelFunc: cumprod$1
|
|
};
|
|
function cumsum$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, exclusive, reverse: reverse2 } = attrs;
|
|
return cumImpl(CumOpType.Sum, x, backend2, axis, exclusive, reverse2);
|
|
}
|
|
const cumsumConfig$1 = {
|
|
kernelName: Cumsum,
|
|
backendName: "webgl",
|
|
kernelFunc: cumsum$1
|
|
};
|
|
function denseBincount$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, weights } = inputs;
|
|
const { size, binaryOutput } = attrs;
|
|
if (x.shape.length === 1) {
|
|
const xVals = backend2.readSync(x.dataId);
|
|
const weightsVals = backend2.readSync(weights.dataId);
|
|
const outVals = bincountImplCPU(xVals, weightsVals, weights.dtype, weights.shape, size);
|
|
return backend2.makeTensorInfo([size], weights.dtype, outVals);
|
|
} else if (x.shape.length === 2) {
|
|
const xBuf = backend2.bufferSync(x);
|
|
const weightsBuf = backend2.bufferSync(weights);
|
|
const outBuf = bincountReduceImplCPU(xBuf, weightsBuf, size, binaryOutput);
|
|
return backend2.makeTensorInfo(outBuf.shape, weights.dtype, outBuf.values);
|
|
}
|
|
throw new Error(`Error in denseBincount: input must be at most rank 2, but got rank${x.shape.length}.`);
|
|
}
|
|
const denseBincountConfig$1 = {
|
|
kernelName: DenseBincount,
|
|
backendName: "webgl",
|
|
kernelFunc: denseBincount$1
|
|
};
|
|
class DepthToSpaceProgram {
|
|
constructor(outputShape, blockSize, dataFormat) {
|
|
this.variableNames = ["x"];
|
|
this.outputShape = [];
|
|
this.outputShape = outputShape;
|
|
this.blockSize = blockSize;
|
|
this.dataFormat = dataFormat;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int h = ${this.getHeightCoordString()};
|
|
int w = ${this.getWidthCoordString()};
|
|
int d = ${this.getDepthCoordString()};
|
|
|
|
int in_h = h / ${blockSize};
|
|
int offset_h = imod(h, ${blockSize});
|
|
int in_w = w / ${blockSize};
|
|
int offset_w = imod(w, ${blockSize});
|
|
int offset_d = (offset_h * ${blockSize} + offset_w) *
|
|
${this.getOutputDepthSize()};
|
|
int in_d = d + offset_d;
|
|
|
|
float result = ${this.getInputSamplingString()};
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
getHeightCoordString() {
|
|
if (this.dataFormat === "NHWC") {
|
|
return `coords[1]`;
|
|
} else {
|
|
return `coords[2]`;
|
|
}
|
|
}
|
|
getWidthCoordString() {
|
|
if (this.dataFormat === "NHWC") {
|
|
return `coords[2]`;
|
|
} else {
|
|
return `coords[3]`;
|
|
}
|
|
}
|
|
getDepthCoordString() {
|
|
if (this.dataFormat === "NHWC") {
|
|
return `coords[3]`;
|
|
} else {
|
|
return `coords[1]`;
|
|
}
|
|
}
|
|
getOutputDepthSize() {
|
|
if (this.dataFormat === "NHWC") {
|
|
return this.outputShape[3];
|
|
} else {
|
|
return this.outputShape[1];
|
|
}
|
|
}
|
|
getInputSamplingString() {
|
|
if (this.dataFormat === "NHWC") {
|
|
return `getX(b, in_h, in_w, in_d)`;
|
|
} else {
|
|
return `getX(b, in_d, in_h, in_w)`;
|
|
}
|
|
}
|
|
}
|
|
function depthToSpace$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { blockSize, dataFormat } = attrs;
|
|
const batchSize = x.shape[0];
|
|
const inputHeight = dataFormat === "NHWC" ? x.shape[1] : x.shape[2];
|
|
const inputWidth = dataFormat === "NHWC" ? x.shape[2] : x.shape[3];
|
|
const inputDepth = dataFormat === "NHWC" ? x.shape[3] : x.shape[1];
|
|
const outputHeight = inputHeight * blockSize;
|
|
const outputWidth = inputWidth * blockSize;
|
|
const outputDepth = inputDepth / (blockSize * blockSize);
|
|
const outputShape = dataFormat === "NHWC" ? [batchSize, outputHeight, outputWidth, outputDepth] : [batchSize, outputDepth, outputHeight, outputWidth];
|
|
const program = new DepthToSpaceProgram(outputShape, blockSize, dataFormat);
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
const depthToSpaceConfig$1 = {
|
|
kernelName: DepthToSpace,
|
|
backendName: "webgl",
|
|
kernelFunc: depthToSpace$1
|
|
};
|
|
class DepthwiseConv2DProgram {
|
|
constructor(convInfo, addBias = false, activation = null, hasPreluActivation = false, hasLeakyReluAlpha = false) {
|
|
this.variableNames = ["x", "W"];
|
|
this.customUniforms = [
|
|
{ name: "pads", type: "ivec2" },
|
|
{ name: "strides", type: "ivec2" },
|
|
{ name: "dilations", type: "ivec2" },
|
|
{ name: "inDims", type: "ivec2" }
|
|
];
|
|
this.outputShape = convInfo.outShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const channelMul = convInfo.outChannels / convInfo.inChannels;
|
|
let activationSnippet = "", applyActivationSnippet = "";
|
|
if (activation) {
|
|
if (hasPreluActivation) {
|
|
activationSnippet = `float activation(float a) {
|
|
float b = getPreluActivationWeightsAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else if (hasLeakyReluAlpha) {
|
|
activationSnippet = `float activation(float a) {
|
|
float b = getLeakyreluAlphaAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else {
|
|
activationSnippet = `
|
|
float activation(float x) {
|
|
${activation}
|
|
}
|
|
`;
|
|
}
|
|
applyActivationSnippet = `result = activation(result);`;
|
|
}
|
|
const addBiasSnippet = addBias ? "result += getBiasAtOutCoords();" : "";
|
|
if (addBias) {
|
|
this.variableNames.push("bias");
|
|
}
|
|
if (hasPreluActivation) {
|
|
this.variableNames.push("preluActivationWeights");
|
|
}
|
|
if (hasLeakyReluAlpha) {
|
|
this.variableNames.push("leakyreluAlpha");
|
|
}
|
|
this.userCode = `
|
|
${activationSnippet}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
ivec2 xRCCorner = coords.yz * strides - pads;
|
|
int d2 = coords.w;
|
|
int d1 = d2 / ${channelMul};
|
|
int q = d2 - d1 * ${channelMul};
|
|
|
|
int xRCorner = xRCCorner.x;
|
|
int xCCorner = xRCCorner.y;
|
|
|
|
// Convolve x(?, ?, d1) with w(:, :, d1, q) to get y(yR, yC, d2).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
// TO DO(dsmilkov): Flatten the two for loops and vec4 the operations.
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
int xR = xRCorner + wR * dilations[0];
|
|
|
|
if (xR < 0 || xR >= inDims[0]) {
|
|
continue;
|
|
}
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
int xC = xCCorner + wC * dilations[1];
|
|
|
|
if (xC < 0 || xC >= inDims[1]) {
|
|
continue;
|
|
}
|
|
|
|
float xVal = getX(batch, xR, xC, d1);
|
|
float wVal = getW(wR, wC, d1, q);
|
|
dotProd += xVal * wVal;
|
|
}
|
|
}
|
|
|
|
float result = dotProd;
|
|
${addBiasSnippet}
|
|
${applyActivationSnippet}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class DepthwiseConvPacked2DProgram {
|
|
constructor(convInfo, addBias = false, activation = null, hasPreluActivation = false, hasLeakyReluAlpha = false) {
|
|
this.variableNames = ["x", "W"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [
|
|
{ name: "pads", type: "ivec2" },
|
|
{ name: "strides", type: "ivec2" },
|
|
{ name: "dilations", type: "ivec2" },
|
|
{ name: "inDims", type: "ivec2" }
|
|
];
|
|
this.outputShape = convInfo.outShape;
|
|
this.enableShapeUniforms = useShapeUniforms(this.outputShape.length);
|
|
const channelMul = convInfo.outChannels / convInfo.inChannels;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const texelsAcross = filterWidth;
|
|
let mainLoop = `
|
|
int xR; int xC; int xCOffset;
|
|
vec4 wTexel; vec4 previous; vec4 final;`;
|
|
for (let c = 0; c < filterWidth; c++) {
|
|
mainLoop += `
|
|
vec4 xTexelC${c * 2};
|
|
int xTexelC${c * 2}Ready;
|
|
vec4 xTexelC${c * 2 + 1};
|
|
int xTexelC${c * 2 + 1}Ready;
|
|
vec4 xC${c};`;
|
|
}
|
|
mainLoop += `
|
|
for (int r = 0; r < ${filterHeight}; r++) {
|
|
`;
|
|
for (let c = 0; c < filterWidth; c++) {
|
|
mainLoop += `
|
|
xTexelC${c * 2} = vec4(0.0);
|
|
xTexelC${c * 2}Ready = 0;
|
|
xTexelC${c * 2 + 1} = vec4(0.0);
|
|
xTexelC${c * 2 + 1}Ready = 0;
|
|
xC${c} = vec4(0.0);`;
|
|
}
|
|
mainLoop += `
|
|
xR = xRCorner + r * dilations[0];
|
|
if (xR >=0 && xR < inDims[0]) {
|
|
`;
|
|
for (let texelC = 0; texelC < (texelsAcross + 1) / 2; texelC++) {
|
|
const colIndex = texelC * 2;
|
|
mainLoop += `
|
|
xC = xCCorner + ${colIndex * dilationWidth};
|
|
`;
|
|
if (strideWidth === 1) {
|
|
if (colIndex < filterWidth) {
|
|
if (padLeft % 2 === 1) {
|
|
mainLoop += `
|
|
xCOffset = xC + 1;
|
|
if (xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xCOffset, d1);
|
|
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
`;
|
|
if (dilationWidth === 1 && colIndex > 0) {
|
|
mainLoop += `
|
|
xC${colIndex} = vec4(xTexelC${colIndex - 2}.zw, xTexelC${colIndex}.xy);
|
|
`;
|
|
} else {
|
|
mainLoop += `
|
|
xCOffset = xC + 1 - 2;
|
|
|
|
if (xCOffset >= 0 && xCOffset < inDims[1]) {
|
|
previous = getX(batch, xR, xCOffset, d1);
|
|
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
previous.zw = vec2(0.0);
|
|
}
|
|
|
|
xC${colIndex} = vec4(previous.zw, xTexelC${colIndex}.xy);
|
|
} else {
|
|
xC${colIndex} = vec4(0.0, 0.0, xTexelC${colIndex}.xy);
|
|
}
|
|
`;
|
|
}
|
|
} else {
|
|
mainLoop += `
|
|
if (xC >= 0 && xC < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xC, d1);
|
|
if (xC + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex} = xTexelC${colIndex};
|
|
`;
|
|
}
|
|
if (colIndex + 1 < filterWidth) {
|
|
const nextTexelOffset = padLeft % 2 === 0 ? nearestLargerEven(dilationWidth) : dilationWidth;
|
|
if (dilationWidth % 2 === 0 && padLeft % 2 === 1 || dilationWidth % 2 !== 0 && padLeft % 2 !== 1) {
|
|
mainLoop += `
|
|
xCOffset = xC + imod(pads[1], 2) + ${nextTexelOffset};
|
|
|
|
if (xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xCOffset, d1);
|
|
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
`;
|
|
if (dilationWidth > 1) {
|
|
mainLoop += `
|
|
xCOffset -= 2;
|
|
if (xCOffset >= 0 && xCOffset < inDims[1]) {
|
|
previous = getX(batch, xR, xCOffset, d1);
|
|
xC${colIndex + 1} = vec4(previous.zw, xTexelC${colIndex + 1}.xy);
|
|
} else {
|
|
xC${colIndex + 1} = vec4(0.0, 0.0, xTexelC${colIndex + 1}.xy);
|
|
}
|
|
`;
|
|
} else {
|
|
mainLoop += `
|
|
xC${colIndex + 1} = vec4(xTexelC${colIndex}.zw, xTexelC${colIndex + 1}.xy);
|
|
`;
|
|
}
|
|
} else {
|
|
if (nextTexelOffset === 1) {
|
|
mainLoop += `
|
|
xC${colIndex + 1} = xTexelC${colIndex};
|
|
`;
|
|
} else {
|
|
mainLoop += `
|
|
xCOffset = xC + ${nextTexelOffset};
|
|
|
|
if (xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xCOffset, d1);
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex + 1} = xTexelC${colIndex + 1};
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
} else {
|
|
if (colIndex < filterWidth) {
|
|
if (padLeft % 2 === 1) {
|
|
mainLoop += `
|
|
xCOffset = xC + 1 - strides[1];
|
|
if(xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xCOffset, d1);
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
|
|
if(xC + 1 >= 0 && xC + 1 < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xC + 1, d1);
|
|
// Need to manually clear unused channels in case
|
|
// we're reading from recycled texture.
|
|
if (xC + 2 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex} = vec4(xTexelC${colIndex}.zw, xTexelC${colIndex + 1}.zw);
|
|
`;
|
|
if (colIndex + 1 < filterWidth) {
|
|
mainLoop += `
|
|
final = vec4(0.0);
|
|
xCOffset = xC + 1 + strides[1];
|
|
if(xCOffset >= 0 && xCOffset < inDims[1]) {
|
|
final = getX(batch, xR, xCOffset, d1);
|
|
}
|
|
xC${colIndex + 1} = vec4(xTexelC${colIndex + 1}.xy, final.xy);
|
|
`;
|
|
}
|
|
} else {
|
|
mainLoop += `
|
|
if(xC >= 0 && xC < inDims[1] && xTexelC${colIndex}Ready == 0) {
|
|
xTexelC${colIndex} = getX(batch, xR, xC, d1);
|
|
if (xC + 1 >= inDims[1]) {
|
|
xTexelC${colIndex}.zw = vec2(0.0);
|
|
}
|
|
xTexelC${colIndex}Ready = 1;
|
|
}
|
|
|
|
xCOffset = xC + strides[1];
|
|
if(xCOffset >= 0 && xCOffset < inDims[1] && xTexelC${colIndex + 1}Ready == 0) {
|
|
xTexelC${colIndex + 1} = getX(batch, xR, xCOffset, d1);
|
|
if (xCOffset + 1 >= inDims[1]) {
|
|
xTexelC${colIndex + 1}.zw = vec2(0.);
|
|
}
|
|
xTexelC${colIndex + 1}Ready = 1;
|
|
}
|
|
|
|
xC${colIndex} = vec4(
|
|
xTexelC${colIndex}.xy, xTexelC${colIndex + 1}.xy);
|
|
`;
|
|
if (colIndex + 1 < filterWidth) {
|
|
mainLoop += `
|
|
xC${colIndex + 1} = vec4(xTexelC${colIndex}.zw, xTexelC${colIndex + 1}.zw);
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if (colIndex < filterWidth) {
|
|
mainLoop += `
|
|
wTexel = getW(r, ${colIndex}, d1, q);
|
|
dotProd += xC${colIndex} * vec4(wTexel.xz, wTexel.xz);
|
|
`;
|
|
if (colIndex + 1 < filterWidth) {
|
|
mainLoop += `
|
|
wTexel = getW(r, ${colIndex + 1}, d1, q);
|
|
dotProd += xC${colIndex + 1} * vec4(wTexel.xz, wTexel.xz);
|
|
`;
|
|
}
|
|
}
|
|
}
|
|
mainLoop += `
|
|
}
|
|
`;
|
|
mainLoop += `
|
|
}
|
|
`;
|
|
let activationSnippet = "", applyActivationSnippet = "";
|
|
if (activation) {
|
|
if (hasPreluActivation) {
|
|
activationSnippet = `vec4 activation(vec4 a) {
|
|
vec4 b = getPreluActivationWeightsAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else if (hasLeakyReluAlpha) {
|
|
activationSnippet = `vec4 activation(vec4 a) {
|
|
vec4 b = getLeakyreluAlphaAtOutCoords();
|
|
${activation}
|
|
}`;
|
|
} else {
|
|
activationSnippet = `vec4 activation(vec4 x) {
|
|
${activation}
|
|
}`;
|
|
}
|
|
applyActivationSnippet = `result = activation(result);`;
|
|
}
|
|
const addBiasSnippet = addBias ? "result += getBiasAtOutCoords();" : "";
|
|
if (addBias) {
|
|
this.variableNames.push("bias");
|
|
}
|
|
if (hasPreluActivation) {
|
|
this.variableNames.push("preluActivationWeights");
|
|
}
|
|
if (hasLeakyReluAlpha) {
|
|
this.variableNames.push("leakyreluAlpha");
|
|
}
|
|
this.userCode = `
|
|
${activationSnippet}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
ivec2 xRCCorner = coords.yz * strides - pads;
|
|
int d2 = coords.w;
|
|
int d1 = d2 / ${channelMul};
|
|
int q = d2 - d1 * ${channelMul};
|
|
int xRCorner = xRCCorner.x;
|
|
int xCCorner = xRCCorner.y;
|
|
|
|
//intialize dotProd with a small epsilon seems to reduce GPU accuracy loss.
|
|
vec4 dotProd = vec4(0.000000000000001);
|
|
|
|
${mainLoop}
|
|
|
|
vec4 result = dotProd - vec4(0.000000000000001);
|
|
${addBiasSnippet}
|
|
${applyActivationSnippet}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function depthwiseConv2dNative$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dilations, dimRoundingMode } = attrs;
|
|
let $dilations = dilations;
|
|
if ($dilations == null) {
|
|
$dilations = [1, 1];
|
|
}
|
|
assert(eitherStridesOrDilationsAreOne(strides, $dilations), () => `Error in depthwiseConv2d: Either strides or dilations must be 1. Got strides ${strides} and dilations '${$dilations}'`);
|
|
const convInfo = computeConv2DInfo(
|
|
x.shape,
|
|
filter.shape,
|
|
strides,
|
|
$dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
let program;
|
|
if (env().getBool("WEBGL_PACK_DEPTHWISECONV") && convInfo.strideWidth <= 2 && convInfo.outChannels / convInfo.inChannels === 1) {
|
|
program = new DepthwiseConvPacked2DProgram(convInfo);
|
|
} else {
|
|
program = new DepthwiseConv2DProgram(convInfo);
|
|
}
|
|
const customValues = [
|
|
[convInfo.padInfo.top, convInfo.padInfo.left],
|
|
[convInfo.strideHeight, convInfo.strideWidth],
|
|
[convInfo.dilationHeight, convInfo.dilationWidth],
|
|
[convInfo.inHeight, convInfo.inWidth]
|
|
];
|
|
return backend2.runWebGLProgram(program, [x, filter], "float32", customValues);
|
|
}
|
|
const depthwiseConv2dNativeConfig$1 = {
|
|
kernelName: DepthwiseConv2dNative,
|
|
backendName: "webgl",
|
|
kernelFunc: depthwiseConv2dNative$1
|
|
};
|
|
class DepthwiseConv2DDerFilterProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["x", "dy"];
|
|
this.outputShape = convInfo.filterShape;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const channelMul = convInfo.outChannels / convInfo.inChannels;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int wR = coords.x;
|
|
int wC = coords.y;
|
|
int d1 = coords.z;
|
|
int dm = coords.w;
|
|
int d2 = d1 * ${channelMul} + dm;
|
|
|
|
float dotProd = 0.0;
|
|
|
|
// TO DO: Vec4 over the batch size
|
|
for (int b = 0; b < ${convInfo.batchSize}; b++) {
|
|
for (int yR = 0; yR < ${convInfo.outHeight}; yR++) {
|
|
int xR = wR + yR * ${strideHeight} - ${padTop};
|
|
|
|
if (xR < 0 || xR >= ${convInfo.inHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int yC = 0; yC < ${convInfo.outWidth}; yC++) {
|
|
int xC = wC + yC * ${strideWidth} - ${padLeft};
|
|
|
|
if (xC < 0 || xC >= ${convInfo.inWidth}) {
|
|
continue;
|
|
}
|
|
|
|
float dyValue = getDy(b, yR, yC, d2);
|
|
float xValue = getX(b, xR, xC, d1);
|
|
dotProd += (xValue * dyValue);
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class DepthwiseConv2DDerInputProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy", "W"];
|
|
this.outputShape = convInfo.inShape;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const padTop = filterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = filterWidth - 1 - convInfo.padInfo.left;
|
|
const channelMul = convInfo.outChannels / convInfo.inChannels;
|
|
this.userCode = `
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int d1 = coords[3];
|
|
ivec2 dyCorner = coords.yz - pads;
|
|
int dyRCorner = dyCorner.x;
|
|
int dyCCorner = dyCorner.y;
|
|
|
|
float dotProd = 0.0;
|
|
|
|
for (int wR = 0; wR < ${filterHeight}; wR++) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 || fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
int wRPerm = ${filterHeight} - 1 - wR;
|
|
|
|
for (int wC = 0; wC < ${filterWidth}; wC++) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
int wCPerm = ${filterWidth} - 1 - wC;
|
|
|
|
// TO DO: Vec4 over the channelMul
|
|
for (int dm = 0; dm < ${channelMul}; dm++) {
|
|
int d2 = d1 * ${channelMul} + dm;
|
|
float xValue = getDy(batch, idyR, idyC, d2);
|
|
float wValue = getW(wRPerm, wCPerm, d1, dm);
|
|
dotProd += xValue * wValue;
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function depthwiseConv2dNativeBackpropFilter$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, dy } = inputs;
|
|
const { strides, dilations, pad: pad2, dimRoundingMode, filterShape } = attrs;
|
|
const convInfo = computeConv2DInfo(
|
|
x.shape,
|
|
filterShape,
|
|
strides,
|
|
dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
const program = new DepthwiseConv2DDerFilterProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [x, dy], "float32");
|
|
}
|
|
const depthwiseConv2dNativeBackpropFilterConfig$1 = {
|
|
kernelName: DepthwiseConv2dNativeBackpropFilter,
|
|
backendName: "webgl",
|
|
kernelFunc: depthwiseConv2dNativeBackpropFilter$1
|
|
};
|
|
function depthwiseConv2dNativeBackpropInput$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, filter } = inputs;
|
|
const { strides, dilations, pad: pad2, dimRoundingMode, inputShape } = attrs;
|
|
const convInfo = computeConv2DInfo(
|
|
inputShape,
|
|
filter.shape,
|
|
strides,
|
|
dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
const program = new DepthwiseConv2DDerInputProgram(convInfo);
|
|
return backend2.runWebGLProgram(program, [dy, filter], "float32");
|
|
}
|
|
const depthwiseConv2dNativeBackpropInputConfig$1 = {
|
|
kernelName: DepthwiseConv2dNativeBackpropInput,
|
|
backendName: "webgl",
|
|
kernelFunc: depthwiseConv2dNativeBackpropInput$1
|
|
};
|
|
class DiagProgram {
|
|
constructor(size) {
|
|
this.variableNames = ["X"];
|
|
this.outputShape = [size, size];
|
|
this.userCode = `
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
float val = coords[0] == coords[1] ? getX(coords[0]) : 0.0;
|
|
setOutput(val);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function diag$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
const outShape = [...x.shape, ...x.shape];
|
|
const xSize = sizeFromShape(x.shape);
|
|
const flat = reshape$1({ inputs: { x }, backend: backend2, attrs: { shape: [xSize] } });
|
|
const program = new DiagProgram(xSize);
|
|
const res = backend2.runWebGLProgram(program, [flat], flat.dtype);
|
|
const out = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape: outShape } });
|
|
backend2.disposeIntermediateTensorInfo(flat);
|
|
backend2.disposeIntermediateTensorInfo(res);
|
|
return out;
|
|
}
|
|
const diagConfig$1 = {
|
|
kernelName: Diag,
|
|
backendName: "webgl",
|
|
kernelFunc: diag$1
|
|
};
|
|
class Dilation2DProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["x", "W"];
|
|
this.outputShape = convInfo.outShape;
|
|
const { inHeight, inWidth, padInfo, strideHeight, strideWidth, filterHeight, filterWidth, dilationHeight, dilationWidth } = convInfo;
|
|
const { top: padTop, left: padLeft } = padInfo;
|
|
this.userCode = `
|
|
const ivec2 strides = ivec2(${strideHeight}, ${strideWidth});
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
const float neg_infinity = -3.4e38;
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int d1 = coords.w;
|
|
ivec2 outTopLeftCorner =
|
|
coords.yz * strides - pads;
|
|
int hBeg = outTopLeftCorner.x;
|
|
int wBeg = outTopLeftCorner.y;
|
|
|
|
float curVal = neg_infinity;
|
|
for (int h = 0; h < ${filterHeight}; h++) {
|
|
int hIn = hBeg + h * ${dilationHeight};
|
|
|
|
if (hIn >= 0 && hIn < ${inHeight}) {
|
|
for (int w = 0; w < ${filterWidth}; w++) {
|
|
int wIn = wBeg + w * ${dilationWidth};
|
|
|
|
if (wIn >= 0 && wIn < ${inWidth}) {
|
|
float xVal = getX(batch, hIn, wIn, d1);
|
|
float wVal = getW(h, w, d1);
|
|
|
|
float val = xVal + wVal;
|
|
if (val > curVal) {
|
|
curVal = val;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
float result = curVal;
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function dilation2D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dilations } = attrs;
|
|
const convInfo = computeDilation2DInfo(x.shape, filter.shape, strides, pad2, "NHWC", dilations);
|
|
let out;
|
|
const program = new Dilation2DProgram(convInfo);
|
|
out = backend2.runWebGLProgram(program, [x, filter], "float32");
|
|
const outReshaped = reshape$1({ inputs: { x: out }, backend: backend2, attrs: { shape: convInfo.outShape } });
|
|
backend2.disposeIntermediateTensorInfo(out);
|
|
return outReshaped;
|
|
}
|
|
const dilation2DConfig$1 = {
|
|
kernelName: Dilation2D,
|
|
backendName: "webgl",
|
|
kernelFunc: dilation2D
|
|
};
|
|
function einsum$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { equation } = attrs;
|
|
const tensors = inputs;
|
|
const { allDims, summedDims, idDims } = decodeEinsumEquation(equation, tensors.length);
|
|
checkEinsumDimSizes(allDims.length, idDims, tensors);
|
|
const { path, steps } = getEinsumComputePath(summedDims, idDims);
|
|
const nSteps = steps.length;
|
|
let out = null;
|
|
let numDimsRemaining = allDims.length;
|
|
const tensorsToDispose = [];
|
|
for (let i = 0; i < nSteps; ++i) {
|
|
for (const idTerm of steps[i]) {
|
|
const { permutationIndices: perm, expandDims: dimsToExpand } = getEinsumPermutation(numDimsRemaining, idDims[idTerm]);
|
|
let x;
|
|
if (isIdentityPermutation(perm)) {
|
|
x = tensors[idTerm];
|
|
} else {
|
|
x = transpose({ inputs: { x: tensors[idTerm] }, backend: backend2, attrs: { perm } });
|
|
tensorsToDispose.push(x);
|
|
}
|
|
const targetShape = x.shape.slice();
|
|
for (let k = 0; k < dimsToExpand.length; ++k) {
|
|
targetShape.splice(dimsToExpand[k], 0, 1);
|
|
}
|
|
if (!arraysEqual(x.shape, targetShape)) {
|
|
x = reshape$1({ inputs: { x }, backend: backend2, attrs: { shape: targetShape } });
|
|
tensorsToDispose.push(x);
|
|
}
|
|
if (out === null) {
|
|
out = x;
|
|
} else {
|
|
out = multiply({ inputs: { a: x, b: out }, backend: backend2 });
|
|
tensorsToDispose.push(out);
|
|
}
|
|
}
|
|
if (i < nSteps - 1) {
|
|
if (path[i] >= 0) {
|
|
out = sum$1({
|
|
inputs: { x: out },
|
|
backend: backend2,
|
|
attrs: {
|
|
axis: path[i] - (allDims.length - numDimsRemaining),
|
|
keepDims: false
|
|
}
|
|
});
|
|
tensorsToDispose.push(out);
|
|
}
|
|
numDimsRemaining--;
|
|
}
|
|
}
|
|
for (const tensorInfo of tensorsToDispose) {
|
|
if (tensorInfo === out) {
|
|
continue;
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(tensorInfo);
|
|
}
|
|
return out;
|
|
}
|
|
const einsumConfig$1 = {
|
|
kernelName: Einsum,
|
|
backendName: "webgl",
|
|
kernelFunc: einsum$1
|
|
};
|
|
const ELU = `return (x >= 0.0) ? x : (exp(x) - 1.0);`;
|
|
const ELU_PACKED = `
|
|
vec4 result;
|
|
|
|
result.r = (x.r >= 0.0) ? x.r : (exp(x.r) - 1.0);
|
|
result.g = (x.g >= 0.0) ? x.g : (exp(x.g) - 1.0);
|
|
result.b = (x.b >= 0.0) ? x.b : (exp(x.b) - 1.0);
|
|
result.a = (x.a >= 0.0) ? x.a : (exp(x.a) - 1.0);
|
|
|
|
return result;
|
|
`;
|
|
const elu$1 = unaryKernelFunc({ opSnippet: ELU, packedOpSnippet: ELU_PACKED });
|
|
const eluConfig$1 = {
|
|
kernelName: Elu,
|
|
backendName: "webgl",
|
|
kernelFunc: elu$1
|
|
};
|
|
const ELU_DER = `return (b >= 0.0) ? a : a * (b + 1.0);`;
|
|
const ELU_DER_PACKED = `
|
|
vec4 bGTEZero = vec4(greaterThanEqual(b, vec4(0.)));
|
|
return (bGTEZero * a) + ((vec4(1.0) - bGTEZero) * (a * (b + vec4(1.0))));
|
|
`;
|
|
const eluGrad$1 = (args) => {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { dy, y } = inputs;
|
|
const program = env().getBool("WEBGL_PACK_BINARY_OPERATIONS") ? new BinaryOpPackedProgram(ELU_DER_PACKED, dy.shape, y.shape) : new BinaryOpProgram(ELU_DER, dy.shape, y.shape);
|
|
return backend2.runWebGLProgram(program, [dy, y], dy.dtype);
|
|
};
|
|
const eluGradConfig$1 = {
|
|
kernelName: EluGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: eluGrad$1
|
|
};
|
|
const PACKED_EQUAL = `
|
|
return vec4(equal(a, b));
|
|
`;
|
|
const EQUAL = `return float(a == b);`;
|
|
const equal = binaryKernelFunc({
|
|
opSnippet: EQUAL,
|
|
packedOpSnippet: PACKED_EQUAL,
|
|
dtype: "bool",
|
|
cpuKernelImpl: equalImplCPU
|
|
});
|
|
const equalConfig = {
|
|
kernelName: Equal,
|
|
backendName: "webgl",
|
|
kernelFunc: equal
|
|
};
|
|
const ERF = `
|
|
// Error function is calculated approximately with elementary function.
|
|
// See "Handbook of Mathematical Functions with Formulas,
|
|
// Graphs, and Mathematical Tables", Abramowitz and Stegun.
|
|
float p = ${ERF_P};
|
|
float a1 = ${ERF_A1};
|
|
float a2 = ${ERF_A2};
|
|
float a3 = ${ERF_A3};
|
|
float a4 = ${ERF_A4};
|
|
float a5 = ${ERF_A5};
|
|
|
|
float sign = sign(x);
|
|
x = abs(x);
|
|
float t = 1.0 / (1.0 + p * x);
|
|
return sign * (1.0 - (((((a5*t + a4)*t) + a3)*t + a2)*t + a1)*t*exp(-x*x));
|
|
`;
|
|
const erf$1 = unaryKernelFunc({ opSnippet: ERF });
|
|
const erfConfig$1 = {
|
|
kernelName: Erf,
|
|
backendName: "webgl",
|
|
kernelFunc: erf$1
|
|
};
|
|
const EXP = CHECK_NAN_SNIPPET_UNARY + `
|
|
return exp(x);
|
|
`;
|
|
const EXP_PACKED = `
|
|
vec4 result = exp(x);
|
|
bvec4 isNaN = isnan(x);
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
const exp = unaryKernelFunc({
|
|
opSnippet: EXP,
|
|
packedOpSnippet: EXP_PACKED,
|
|
cpuKernelImpl: expImplCPU,
|
|
dtype: "float32"
|
|
});
|
|
const expConfig = {
|
|
kernelName: Exp,
|
|
backendName: "webgl",
|
|
kernelFunc: exp
|
|
};
|
|
function expandDims$1(args) {
|
|
const { inputs, attrs, backend: backend2 } = args;
|
|
const { dim } = attrs;
|
|
const { input } = inputs;
|
|
const inputRank = input.shape.length;
|
|
const newShape = input.shape.slice();
|
|
let $dim = dim;
|
|
if (dim < 0) {
|
|
assert(-(inputRank + 1) <= dim, () => `Axis must be in the interval [${-(inputRank + 1)}, ${inputRank}]`);
|
|
$dim = inputRank + dim + 1;
|
|
}
|
|
newShape.splice($dim, 0, 1);
|
|
return reshape$1({ inputs: { x: input }, backend: backend2, attrs: { shape: newShape } });
|
|
}
|
|
const expandDimsConfig$1 = {
|
|
kernelName: ExpandDims,
|
|
backendName: "webgl",
|
|
kernelFunc: expandDims$1
|
|
};
|
|
const EXPM1 = `return exp(x) - 1.0;`;
|
|
const expm1 = unaryKernelFunc({ opSnippet: EXPM1, packedOpSnippet: EXPM1, cpuKernelImpl: expm1ImplCPU });
|
|
const expm1Config = {
|
|
kernelName: Expm1,
|
|
backendName: "webgl",
|
|
kernelFunc: expm1
|
|
};
|
|
class FFTProgram {
|
|
constructor(component, inputShape, inverse) {
|
|
this.variableNames = ["real", "imag"];
|
|
const innerDim = inputShape[1];
|
|
this.outputShape = inputShape;
|
|
const exponentMultiplierSnippet = inverse ? `2.0 * ${Math.PI}` : `-2.0 * ${Math.PI}`;
|
|
const resultDenominator = inverse ? `${innerDim}.0` : "1.0";
|
|
let opString;
|
|
if (component === "real") {
|
|
opString = "return real * expR - imag * expI;";
|
|
} else if (component === "imag") {
|
|
opString = "return real * expI + imag * expR;";
|
|
} else {
|
|
throw new Error(`FFT component must be either "real" or "imag", got ${component}.`);
|
|
}
|
|
this.userCode = `
|
|
const float exponentMultiplier = ${exponentMultiplierSnippet};
|
|
|
|
float unaryOpComplex(float real, float expR, float imag, float expI) {
|
|
${opString}
|
|
}
|
|
|
|
float mulMatDFT(int batch, int index) {
|
|
float indexRatio = float(index) / float(${innerDim});
|
|
float exponentMultiplierTimesIndexRatio =
|
|
exponentMultiplier * indexRatio;
|
|
|
|
float result = 0.0;
|
|
|
|
for (int i = 0; i < ${innerDim}; i++) {
|
|
// x = (-2|2 * PI / N) * index * i;
|
|
float x = exponentMultiplierTimesIndexRatio * float(i);
|
|
float expR = cos(x);
|
|
float expI = sin(x);
|
|
float real = getReal(batch, i);
|
|
float imag = getImag(batch, i);
|
|
|
|
result +=
|
|
unaryOpComplex(real, expR, imag, expI) / ${resultDenominator};
|
|
}
|
|
|
|
return result;
|
|
}
|
|
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
setOutput(mulMatDFT(coords[0], coords[1]));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function fftImpl$1(x, inverse, backend2) {
|
|
const xData = backend2.texData.get(x.dataId);
|
|
const inputSize = sizeFromShape(x.shape);
|
|
const innerDimensionSize = x.shape[x.shape.length - 1];
|
|
const batch = inputSize / innerDimensionSize;
|
|
const input2D = reshape$1({ inputs: { x }, backend: backend2, attrs: { shape: [batch, innerDimensionSize] } });
|
|
const xShape = input2D.shape;
|
|
const realProgram = new FFTProgram("real", xShape, inverse);
|
|
const imagProgram = new FFTProgram("imag", xShape, inverse);
|
|
const inputs = [
|
|
{
|
|
dataId: xData.complexTensorInfos.real.dataId,
|
|
dtype: xData.complexTensorInfos.real.dtype,
|
|
shape: xShape
|
|
},
|
|
{
|
|
dataId: xData.complexTensorInfos.imag.dataId,
|
|
dtype: xData.complexTensorInfos.imag.dtype,
|
|
shape: xShape
|
|
}
|
|
];
|
|
const realPart = backend2.runWebGLProgram(realProgram, inputs, "float32");
|
|
const imagPart = backend2.runWebGLProgram(imagProgram, inputs, "float32");
|
|
const complexOutput = complex({ inputs: { real: realPart, imag: imagPart }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
backend2.disposeIntermediateTensorInfo(imagPart);
|
|
const complexOutputReshaped = reshape$1({ inputs: { x: complexOutput }, backend: backend2, attrs: { shape: x.shape } });
|
|
backend2.disposeIntermediateTensorInfo(input2D);
|
|
backend2.disposeIntermediateTensorInfo(complexOutput);
|
|
return complexOutputReshaped;
|
|
}
|
|
function fft$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
return fftImpl$1(input, false, backend2);
|
|
}
|
|
const fftConfig$1 = {
|
|
kernelName: FFT,
|
|
backendName: "webgl",
|
|
kernelFunc: fft$1
|
|
};
|
|
class FillProgram {
|
|
constructor(shape, value) {
|
|
this.outputShape = [];
|
|
this.customUniforms = [{ name: "value", type: "float" }];
|
|
this.variableNames = ["x"];
|
|
this.outputShape = shape;
|
|
this.userCode = `
|
|
void main() {
|
|
// Input can be obtained from uniform value.
|
|
setOutput(value);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function fill$1(args) {
|
|
const { backend: backend2, attrs } = args;
|
|
const { shape, value } = attrs;
|
|
let { dtype } = attrs;
|
|
dtype = dtype || inferDtype(value);
|
|
if (dtype === "string") {
|
|
const values = getArrayFromDType(dtype, sizeFromShape(shape));
|
|
values.fill(value);
|
|
return backend2.makeTensorInfo(shape, dtype, values);
|
|
} else {
|
|
const program = new FillProgram(shape, value);
|
|
const customValues = [[value]];
|
|
return backend2.runWebGLProgram(program, [], dtype, customValues);
|
|
}
|
|
}
|
|
const fillConfig$1 = {
|
|
kernelName: Fill,
|
|
backendName: "webgl",
|
|
kernelFunc: fill$1
|
|
};
|
|
class FlipLeftRightProgram {
|
|
constructor(imageShape) {
|
|
this.variableNames = ["Image"];
|
|
this.outputShape = [];
|
|
const imageWidth = imageShape[2];
|
|
this.outputShape = imageShape;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int x = coords[2];
|
|
|
|
int coordX = ${imageWidth} - x - 1;
|
|
float outputValue;
|
|
if(coordX >= 0 && coordX < ${imageWidth}) {
|
|
outputValue = getImage(coords[0], coords[1], coordX, coords[3]);
|
|
} else {
|
|
outputValue = getImage(coords[0], coords[1], coords[2], coords[3]);
|
|
}
|
|
setOutput(outputValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const flipLeftRightConfig$1 = {
|
|
kernelName: FlipLeftRight,
|
|
backendName: "webgl",
|
|
kernelFunc: ({ inputs, backend: backend2 }) => {
|
|
const { image: image2 } = inputs;
|
|
const webglBackend = backend2;
|
|
const program = new FlipLeftRightProgram(image2.shape);
|
|
const output = webglBackend.runWebGLProgram(program, [image2], image2.dtype);
|
|
return output;
|
|
}
|
|
};
|
|
const FLOOR = `return floor(x);`;
|
|
const floor = unaryKernelFunc({ opSnippet: FLOOR, packedOpSnippet: FLOOR, cpuKernelImpl: floorImplCPU });
|
|
const floorConfig = {
|
|
kernelName: Floor,
|
|
backendName: "webgl",
|
|
kernelFunc: floor
|
|
};
|
|
const INT_DIV = `
|
|
float s = sign(a) * sign(b);
|
|
int ia = round(a);
|
|
int ib = round(b);
|
|
if (ib != 0) {
|
|
// Windows (D3D) wants guaranteed non-zero int division at compile-time.
|
|
return float(idiv(ia, ib, s));
|
|
} else {
|
|
return NAN;
|
|
}
|
|
`;
|
|
const INT_DIV_PACKED = `
|
|
ivec4 ia = round(a);
|
|
ivec4 ib = round(b);
|
|
bvec4 cond = notEqual(ib, ivec4(0));
|
|
ivec4 result = ivec4(0);
|
|
vec4 s = sign(a) * sign(b);
|
|
|
|
// Windows (D3D) wants guaranteed non-zero int division at compile-time.
|
|
if (cond[0]) {
|
|
result[0] = idiv(ia[0], ib[0], s[0]);
|
|
}
|
|
if (cond[1]) {
|
|
result[1] = idiv(ia[1], ib[1], s[1]);
|
|
}
|
|
if (cond[2]) {
|
|
result[2] = idiv(ia[2], ib[2], s[2]);
|
|
}
|
|
if (cond[3]) {
|
|
result[3] = idiv(ia[3], ib[3], s[3]);
|
|
}
|
|
return vec4(result);
|
|
`;
|
|
const floorDiv = binaryKernelFunc({ opSnippet: INT_DIV, packedOpSnippet: INT_DIV_PACKED, dtype: "int32" });
|
|
const floorDivConfig = {
|
|
kernelName: FloorDiv,
|
|
backendName: "webgl",
|
|
kernelFunc: floorDiv
|
|
};
|
|
class FromPixelsProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
const glsl = getGlslDifferences();
|
|
const [height, width] = outputShape;
|
|
this.outputShape = outputShape;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec3 coords = getOutputCoords();
|
|
int texR = coords[0];
|
|
int texC = coords[1];
|
|
int depth = coords[2];
|
|
vec2 uv = (vec2(texC, texR) + halfCR) / vec2(${width}.0, ${height}.0);
|
|
|
|
vec4 values = ${glsl.texture2D}(A, uv);
|
|
float value;
|
|
if (depth == 0) {
|
|
value = values.r;
|
|
} else if (depth == 1) {
|
|
value = values.g;
|
|
} else if (depth == 2) {
|
|
value = values.b;
|
|
} else if (depth == 3) {
|
|
value = values.a;
|
|
}
|
|
|
|
setOutput(floor(value * 255.0 + 0.5));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class FromPixelsPackedProgram {
|
|
constructor(outputShape) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = false;
|
|
this.packedOutput = true;
|
|
const glsl = getGlslDifferences();
|
|
const [height, width] = outputShape;
|
|
this.outputShape = outputShape;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec3 coords = getOutputCoords();
|
|
int texR = coords[0];
|
|
int texC = coords[1];
|
|
int depth = coords[2];
|
|
|
|
vec4 result = vec4(0.);
|
|
|
|
for(int row=0; row<=1; row++) {
|
|
for(int col=0; col<=1; col++) {
|
|
texC = coords[1] + row;
|
|
depth = coords[2] + col;
|
|
|
|
vec2 uv = (vec2(texC, texR) + halfCR) /
|
|
vec2(${width}.0, ${height}.0);
|
|
vec4 values = ${glsl.texture2D}(A, uv);
|
|
float value;
|
|
if (depth == 0) {
|
|
value = values.r;
|
|
} else if (depth == 1) {
|
|
value = values.g;
|
|
} else if (depth == 2) {
|
|
value = values.b;
|
|
} else if (depth == 3) {
|
|
value = values.a;
|
|
}
|
|
|
|
result[row * 2 + col] = floor(value * 255.0 + 0.5);
|
|
}
|
|
}
|
|
|
|
${glsl.output} = result;
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const fromPixelsConfig = {
|
|
kernelName: FromPixels,
|
|
backendName: "webgl",
|
|
kernelFunc: fromPixels
|
|
};
|
|
let fromPixels2DContext;
|
|
let willReadFrequently = env().getBool("CANVAS2D_WILL_READ_FREQUENTLY_FOR_GPU");
|
|
function fromPixels(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
let { pixels } = inputs;
|
|
const { numChannels } = attrs;
|
|
const isVideo = typeof HTMLVideoElement !== "undefined" && pixels instanceof HTMLVideoElement;
|
|
const isImage = typeof HTMLImageElement !== "undefined" && pixels instanceof HTMLImageElement;
|
|
const [width, height] = isVideo ? [
|
|
pixels.videoWidth,
|
|
pixels.videoHeight
|
|
] : [pixels.width, pixels.height];
|
|
const texShape = [height, width];
|
|
const outShape = [height, width, numChannels];
|
|
if (isImage || isVideo) {
|
|
const newWillReadFrequently = env().getBool("CANVAS2D_WILL_READ_FREQUENTLY_FOR_GPU");
|
|
if (fromPixels2DContext == null || newWillReadFrequently !== willReadFrequently) {
|
|
willReadFrequently = newWillReadFrequently;
|
|
fromPixels2DContext = document.createElement("canvas").getContext("2d", { willReadFrequently });
|
|
}
|
|
fromPixels2DContext.canvas.width = width;
|
|
fromPixels2DContext.canvas.height = height;
|
|
fromPixels2DContext.drawImage(pixels, 0, 0, width, height);
|
|
pixels = fromPixels2DContext.canvas;
|
|
}
|
|
const tempPixelHandle = backend2.makeTensorInfo(texShape, "int32");
|
|
backend2.texData.get(tempPixelHandle.dataId).usage = TextureUsage.PIXELS;
|
|
backend2.gpgpu.uploadPixelDataToTexture(backend2.getTexture(tempPixelHandle.dataId), pixels);
|
|
const program = env().getBool("WEBGL_PACK") ? new FromPixelsPackedProgram(outShape) : new FromPixelsProgram(outShape);
|
|
const res = backend2.runWebGLProgram(program, [tempPixelHandle], "int32");
|
|
backend2.disposeData(tempPixelHandle.dataId);
|
|
return res;
|
|
}
|
|
function fusedConv2d(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter, bias, preluActivationWeights } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dilations, dimRoundingMode, activation, leakyreluAlpha } = attrs;
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(x.shape, filter.shape, strides, dilations, pad2, dimRoundingMode, false, $dataFormat);
|
|
let out;
|
|
const intermediates = [];
|
|
const hasBias = bias != null;
|
|
const hasPreluActivationWeights = preluActivationWeights != null;
|
|
const hasLeakyreluAlpha = activation === "leakyrelu";
|
|
const prepareInputs = () => {
|
|
const inputs2 = [x, filter];
|
|
const alignInputWithDataFormat = (input, dataFormat2) => {
|
|
if (dataFormat2 === "NCHW" && input.shape.length === 1 && input.shape[0] !== 1) {
|
|
const alignedInput = reshape$1({
|
|
inputs: { x: input },
|
|
backend: backend2,
|
|
attrs: { shape: [input.shape[0], 1, 1] }
|
|
});
|
|
intermediates.push(alignedInput);
|
|
return alignedInput;
|
|
}
|
|
return input;
|
|
};
|
|
if (hasBias) {
|
|
inputs2.push(alignInputWithDataFormat(bias, dataFormat));
|
|
}
|
|
if (hasPreluActivationWeights) {
|
|
inputs2.push(alignInputWithDataFormat(preluActivationWeights, dataFormat));
|
|
}
|
|
if (hasLeakyreluAlpha) {
|
|
const $leakyreluAlpha = backend2.makeTensorInfo([], "float32", createScalarValue(leakyreluAlpha, "float32"));
|
|
inputs2.push($leakyreluAlpha);
|
|
intermediates.push($leakyreluAlpha);
|
|
}
|
|
return inputs2;
|
|
};
|
|
if (convInfo.filterHeight === 1 && convInfo.filterWidth === 1 && convInfo.dilationHeight === 1 && convInfo.dilationWidth === 1 && convInfo.strideHeight === 1 && convInfo.strideWidth === 1 && (convInfo.padInfo.type === "SAME" || convInfo.padInfo.type === "VALID")) {
|
|
out = conv2dByMatMul({
|
|
x,
|
|
filter,
|
|
convInfo,
|
|
backend: backend2,
|
|
bias,
|
|
activation,
|
|
preluActivationWeights,
|
|
leakyreluAlpha
|
|
});
|
|
} else if (convInfo.strideWidth <= 2 && $dataFormat === "channelsLast" && env().getBool("WEBGL_EXP_CONV")) {
|
|
const fusedActivation = activation ? mapActivationToShaderProgram(activation, true) : null;
|
|
const program = new Conv2DPackedProgram(convInfo, hasBias, fusedActivation, hasPreluActivationWeights, hasLeakyreluAlpha);
|
|
const customValues = [
|
|
[convInfo.padInfo.top, convInfo.padInfo.left],
|
|
[convInfo.strideHeight, convInfo.strideWidth],
|
|
[convInfo.dilationHeight, convInfo.dilationWidth],
|
|
[convInfo.inHeight, convInfo.inWidth]
|
|
];
|
|
const inputs2 = prepareInputs();
|
|
out = backend2.runWebGLProgram(program, inputs2, "float32", customValues);
|
|
} else if (env().getBool("WEBGL_CONV_IM2COL")) {
|
|
out = conv2dWithIm2Row({
|
|
x,
|
|
filter,
|
|
convInfo,
|
|
backend: backend2,
|
|
bias,
|
|
activation,
|
|
preluActivationWeights,
|
|
leakyreluAlpha
|
|
});
|
|
} else {
|
|
const fusedActivation = activation ? mapActivationToShaderProgram(activation, false) : null;
|
|
const program = new Conv2DProgram(convInfo, hasBias, fusedActivation, hasPreluActivationWeights, hasLeakyreluAlpha);
|
|
const inputs2 = prepareInputs();
|
|
out = backend2.runWebGLProgram(program, inputs2, "float32");
|
|
}
|
|
const outReshaped = reshape$1({ inputs: { x: out }, backend: backend2, attrs: { shape: convInfo.outShape } });
|
|
intermediates.push(out);
|
|
intermediates.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return outReshaped;
|
|
}
|
|
const fusedConv2DConfig$1 = {
|
|
kernelName: FusedConv2D,
|
|
backendName: "webgl",
|
|
kernelFunc: fusedConv2d
|
|
};
|
|
function fusedDepthwiseConv2D$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter, bias, preluActivationWeights } = inputs;
|
|
const { strides, pad: pad2, dilations, dimRoundingMode, activation, leakyreluAlpha } = attrs;
|
|
const intermediates = [];
|
|
let $dilations = dilations;
|
|
if ($dilations == null) {
|
|
$dilations = [1, 1];
|
|
}
|
|
assert(eitherStridesOrDilationsAreOne(strides, $dilations), () => `Error in depthwiseConv2d: Either strides or dilations must be 1. Got strides ${strides} and dilations '${$dilations}'`);
|
|
const convInfo = computeConv2DInfo(
|
|
x.shape,
|
|
filter.shape,
|
|
strides,
|
|
$dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
const shouldPackDepthwiseConv = env().getBool("WEBGL_PACK_DEPTHWISECONV") && convInfo.strideWidth <= 2 && convInfo.outChannels / convInfo.inChannels === 1;
|
|
const fusedActivation = activation ? mapActivationToShaderProgram(activation, shouldPackDepthwiseConv) : null;
|
|
const programInputs = [x, filter];
|
|
const hasBias = bias != null;
|
|
const hasPreluActivationWeights = preluActivationWeights != null;
|
|
const hasLeakyreluAlpha = activation === "leakyrelu";
|
|
if (hasBias) {
|
|
programInputs.push(bias);
|
|
}
|
|
if (hasPreluActivationWeights) {
|
|
programInputs.push(preluActivationWeights);
|
|
}
|
|
if (hasLeakyreluAlpha) {
|
|
const $leakyreluAlpha = backend2.makeTensorInfo([], "float32", createScalarValue(leakyreluAlpha, "float32"));
|
|
programInputs.push($leakyreluAlpha);
|
|
intermediates.push($leakyreluAlpha);
|
|
}
|
|
let program;
|
|
if (shouldPackDepthwiseConv) {
|
|
program = new DepthwiseConvPacked2DProgram(convInfo, hasBias, fusedActivation, hasPreluActivationWeights, hasLeakyreluAlpha);
|
|
} else {
|
|
program = new DepthwiseConv2DProgram(convInfo, hasBias, fusedActivation, hasPreluActivationWeights, hasLeakyreluAlpha);
|
|
}
|
|
const customValues = [
|
|
[convInfo.padInfo.top, convInfo.padInfo.left],
|
|
[convInfo.strideHeight, convInfo.strideWidth],
|
|
[convInfo.dilationHeight, convInfo.dilationWidth],
|
|
[convInfo.inHeight, convInfo.inWidth]
|
|
];
|
|
const result = backend2.runWebGLProgram(program, programInputs, "float32", customValues);
|
|
intermediates.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
}
|
|
const fusedDepthwiseConv2DConfig$1 = {
|
|
kernelName: FusedDepthwiseConv2D,
|
|
backendName: "webgl",
|
|
kernelFunc: fusedDepthwiseConv2D$1
|
|
};
|
|
class GatherNDProgram {
|
|
constructor(sliceDim, strides, shape, paramsShape) {
|
|
this.sliceDim = sliceDim;
|
|
this.strides = strides;
|
|
this.paramsShape = paramsShape;
|
|
this.variableNames = ["x", "indices"];
|
|
this.outputShape = shape;
|
|
const dtype = getCoordsDataType(shape.length);
|
|
let mainLoop = `
|
|
int index;`;
|
|
for (let j2 = 0; j2 < this.sliceDim; j2++) {
|
|
mainLoop += `
|
|
index = round(getIndices(coords[0], ${j2}));
|
|
out_of_bounds = out_of_bounds || index < 0;
|
|
out_of_bounds = out_of_bounds || index >= ${this.paramsShape[j2]};
|
|
flattenIndex += index * ${this.strides[j2]};`;
|
|
}
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
int flattenIndex = 0;
|
|
bool out_of_bounds = false;
|
|
|
|
${mainLoop}
|
|
|
|
setOutput(out_of_bounds ? 0.0 : getX(flattenIndex, coords[1]));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function gatherNd$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { params, indices } = inputs;
|
|
const indicesShape = indices.shape;
|
|
const sliceRank = indicesShape[indicesShape.length - 1];
|
|
const paramsSize = sizeFromShape(params.shape);
|
|
const [resultShape, numSlices, sliceSize, strides] = prepareAndValidate(params, indices);
|
|
const flattenIndices = reshape$1({ inputs: { x: indices }, backend: backend2, attrs: { shape: [numSlices, sliceRank] } });
|
|
const flattenX = reshape$1({
|
|
inputs: { x: params },
|
|
backend: backend2,
|
|
attrs: { shape: [sizeFromShape(params.shape) / sliceSize, sliceSize] }
|
|
});
|
|
if (backend2.shouldExecuteOnCPU([params, indices]) || params.dtype === "string") {
|
|
const indicesData = backend2.readSync(indices.dataId);
|
|
const paramsBuf = backend2.bufferSync(params);
|
|
const outValue = gatherNdImplCPU(indicesData, paramsBuf, params.dtype, numSlices, sliceRank, sliceSize, strides, params.shape, paramsSize);
|
|
return backend2.makeTensorInfo(resultShape, params.dtype, outValue.values);
|
|
}
|
|
const program = new GatherNDProgram(sliceRank, strides, [numSlices, sliceSize], params.shape);
|
|
const res = backend2.runWebGLProgram(program, [flattenX, flattenIndices], flattenX.dtype);
|
|
const reshaped = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape: resultShape } });
|
|
backend2.disposeIntermediateTensorInfo(flattenIndices);
|
|
backend2.disposeIntermediateTensorInfo(flattenX);
|
|
backend2.disposeIntermediateTensorInfo(res);
|
|
return reshaped;
|
|
}
|
|
const gatherNdConfig$1 = {
|
|
kernelName: GatherNd,
|
|
backendName: "webgl",
|
|
kernelFunc: gatherNd$1
|
|
};
|
|
class GatherProgram {
|
|
constructor(aShape, outputShape) {
|
|
this.variableNames = ["A", "indices"];
|
|
this.outputShape = outputShape;
|
|
this.rank = outputShape.length;
|
|
const dtype = getCoordsDataType(this.rank);
|
|
const sourceCoords = getSourceCoords$1(aShape);
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} resRC = getOutputCoords();
|
|
int index = int(getIndices(resRC.x, resRC.z));
|
|
float inBounds = (index >= 0) && (index < ${aShape[2]}) ? 1.0 : 0.0;
|
|
setOutput(inBounds * getA(${sourceCoords}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getSourceCoords$1(aShape, axis) {
|
|
const currentCoords = ["resRC.x", "resRC.y", "resRC.z", "resRC.w"];
|
|
const sourceCoords = [];
|
|
for (let i = 0; i < aShape.length; i++) {
|
|
if (i === 2) {
|
|
sourceCoords.push("index");
|
|
} else {
|
|
sourceCoords.push(`${currentCoords[i]}`);
|
|
}
|
|
}
|
|
return sourceCoords.join();
|
|
}
|
|
function gatherV2$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, indices } = inputs;
|
|
const { axis, batchDims } = attrs;
|
|
const parsedAxis = parseAxisParam(axis, x.shape)[0];
|
|
if (env().get("DEBUG")) {
|
|
const indicesVals = backend2.readSync(indices.dataId);
|
|
const axisDim = x.shape[parsedAxis];
|
|
for (let i = 0; i < indicesVals.length; ++i) {
|
|
const index = indicesVals[i];
|
|
assert(index <= axisDim - 1 && index >= 0, () => `GatherV2: the index value ${index} is not in [0, ${axisDim - 1}]`);
|
|
}
|
|
}
|
|
const shapeInfo = collectGatherOpShapeInfo(x, indices, parsedAxis, batchDims);
|
|
const indicesSize = sizeFromShape(indices.shape);
|
|
const toDispose = [];
|
|
const flattenX = reshape$1({
|
|
inputs: { x },
|
|
backend: backend2,
|
|
attrs: {
|
|
shape: [
|
|
shapeInfo.batchSize,
|
|
shapeInfo.outerSize,
|
|
shapeInfo.dimSize,
|
|
shapeInfo.sliceSize
|
|
]
|
|
}
|
|
});
|
|
const flattenIndex = reshape$1({
|
|
inputs: { x: indices },
|
|
backend: backend2,
|
|
attrs: { shape: [shapeInfo.batchSize, indicesSize / shapeInfo.batchSize] }
|
|
});
|
|
toDispose.push(flattenX);
|
|
toDispose.push(flattenIndex);
|
|
const flattenOutputShape = [
|
|
shapeInfo.batchSize,
|
|
shapeInfo.outerSize,
|
|
indicesSize / shapeInfo.batchSize,
|
|
shapeInfo.sliceSize
|
|
];
|
|
if (backend2.shouldExecuteOnCPU([x, indices]) || x.dtype === "string") {
|
|
const indicesBuf = backend2.bufferSync(flattenIndex);
|
|
const xBuf = backend2.bufferSync(flattenX);
|
|
const outBuf = gatherV2ImplCPU(xBuf, indicesBuf, flattenOutputShape);
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return backend2.makeTensorInfo(shapeInfo.outputShape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const program = new GatherProgram(flattenX.shape, flattenOutputShape);
|
|
const res = backend2.runWebGLProgram(program, [flattenX, flattenIndex], flattenX.dtype);
|
|
toDispose.push(res);
|
|
const reshaped = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape: shapeInfo.outputShape } });
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return reshaped;
|
|
}
|
|
const gatherV2Config$1 = {
|
|
kernelName: GatherV2,
|
|
backendName: "webgl",
|
|
kernelFunc: gatherV2$1
|
|
};
|
|
const GREATER = `return float(a > b);`;
|
|
const GREATER_PACKED = `
|
|
return vec4(greaterThan(a, b));
|
|
`;
|
|
const greater = binaryKernelFunc({
|
|
opSnippet: GREATER,
|
|
packedOpSnippet: GREATER_PACKED,
|
|
cpuKernelImpl: greaterImplCPU,
|
|
dtype: "bool"
|
|
});
|
|
const greaterConfig = {
|
|
kernelName: Greater,
|
|
backendName: "webgl",
|
|
kernelFunc: greater
|
|
};
|
|
const GREATER_EQUAL = `return float(a >= b);`;
|
|
const GREATER_EQUAL_PACKED = `
|
|
return vec4(greaterThanEqual(a, b));
|
|
`;
|
|
const greaterEqual = binaryKernelFunc({
|
|
opSnippet: GREATER_EQUAL,
|
|
packedOpSnippet: GREATER_EQUAL_PACKED,
|
|
dtype: "bool",
|
|
cpuKernelImpl: greaterEqualImplCPU
|
|
});
|
|
const greaterEqualConfig = {
|
|
kernelName: GreaterEqual,
|
|
backendName: "webgl",
|
|
kernelFunc: greaterEqual
|
|
};
|
|
function ifft$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
return fftImpl$1(input, true, backend2);
|
|
}
|
|
const ifftConfig$1 = {
|
|
kernelName: IFFT,
|
|
backendName: "webgl",
|
|
kernelFunc: ifft$1
|
|
};
|
|
const IS_FINITE = `return float(!isnan(x) && !isinf(x));`;
|
|
const isFinite$2 = unaryKernelFunc({ opSnippet: IS_FINITE, dtype: "bool" });
|
|
const isFiniteConfig$1 = {
|
|
kernelName: IsFinite,
|
|
backendName: "webgl",
|
|
kernelFunc: isFinite$2
|
|
};
|
|
const IS_INF = `return float(isinf(x));`;
|
|
const isInf$1 = unaryKernelFunc({ opSnippet: IS_INF, dtype: "bool" });
|
|
const isInfConfig$1 = {
|
|
kernelName: IsInf,
|
|
backendName: "webgl",
|
|
kernelFunc: isInf$1
|
|
};
|
|
const IS_NAN = `return float(isnan(x));`;
|
|
const isNaN$2 = unaryKernelFunc({ opSnippet: IS_NAN, dtype: "bool" });
|
|
const isNaNConfig$1 = {
|
|
kernelName: IsNan,
|
|
backendName: "webgl",
|
|
kernelFunc: isNaN$2
|
|
};
|
|
const LESS = `return float(a < b);`;
|
|
const LESS_PACKED = `
|
|
return vec4(lessThan(a, b));
|
|
`;
|
|
const less = binaryKernelFunc({
|
|
opSnippet: LESS,
|
|
packedOpSnippet: LESS_PACKED,
|
|
cpuKernelImpl: lessImplCPU,
|
|
dtype: "bool"
|
|
});
|
|
const lessConfig = {
|
|
kernelName: Less,
|
|
backendName: "webgl",
|
|
kernelFunc: less
|
|
};
|
|
const LESS_EQUAL = `return float(a <= b);`;
|
|
const LESS_EQUAL_PACKED = `
|
|
return vec4(lessThanEqual(a, b));
|
|
`;
|
|
const lessEqual = binaryKernelFunc({
|
|
opSnippet: LESS_EQUAL,
|
|
packedOpSnippet: LESS_EQUAL_PACKED,
|
|
cpuKernelImpl: lessEqualImplCPU,
|
|
dtype: "bool"
|
|
});
|
|
const lessEqualConfig = {
|
|
kernelName: LessEqual,
|
|
backendName: "webgl",
|
|
kernelFunc: lessEqual
|
|
};
|
|
function linSpace$1(args) {
|
|
const { backend: backend2, attrs } = args;
|
|
const { start, stop, num } = attrs;
|
|
const outVals = linSpaceImplCPU(start, stop, num);
|
|
return backend2.makeTensorInfo([outVals.length], "float32", outVals);
|
|
}
|
|
const linSpaceConfig$1 = {
|
|
kernelName: LinSpace,
|
|
backendName: "webgl",
|
|
kernelFunc: linSpace$1
|
|
};
|
|
const LOG = CHECK_NAN_SNIPPET_UNARY + `
|
|
return x < 0.0 ? 0./0. : log(x);
|
|
`;
|
|
const LOG_PACKED = `
|
|
vec4 result = log(x);
|
|
bvec4 isNaN = isnan(x);
|
|
result.r = isNaN.r ? x.r : (x.r < 0.0 ? 0./0. : result.r);
|
|
result.g = isNaN.g ? x.g : (x.g < 0.0 ? 0./0. : result.g);
|
|
result.b = isNaN.b ? x.b : (x.b < 0.0 ? 0./0. : result.b);
|
|
result.a = isNaN.a ? x.a : (x.a < 0.0 ? 0./0. : result.a);
|
|
return result;
|
|
`;
|
|
const log = unaryKernelFunc({ opSnippet: LOG, packedOpSnippet: LOG_PACKED, cpuKernelImpl: logImplCPU });
|
|
const logConfig = {
|
|
kernelName: Log,
|
|
backendName: "webgl",
|
|
kernelFunc: log
|
|
};
|
|
const LOG1P = CHECK_NAN_SNIPPET_UNARY + `
|
|
return log(1.0 + x);
|
|
`;
|
|
const log1p$1 = unaryKernelFunc({ opSnippet: LOG1P });
|
|
const log1pConfig$1 = {
|
|
kernelName: Log1p,
|
|
backendName: "webgl",
|
|
kernelFunc: log1p$1
|
|
};
|
|
const LOGICAL_AND = `return float(a >= 1.0 && b >= 1.0);`;
|
|
const LOGICAL_AND_PACKED = `
|
|
return vec4(
|
|
vec4(greaterThanEqual(a, vec4(1.0))) *
|
|
vec4(greaterThanEqual(b, vec4(1.0))));
|
|
`;
|
|
const logicalAnd$1 = binaryKernelFunc({
|
|
opSnippet: LOGICAL_AND,
|
|
packedOpSnippet: LOGICAL_AND_PACKED,
|
|
dtype: "bool"
|
|
});
|
|
const logicalAndConfig$1 = {
|
|
kernelName: LogicalAnd,
|
|
backendName: "webgl",
|
|
kernelFunc: logicalAnd$1
|
|
};
|
|
const LOGICAL_NOT = `return float(!(x >= 1.0));`;
|
|
const logicalNot$1 = unaryKernelFunc({ opSnippet: LOGICAL_NOT });
|
|
const logicalNotConfig$1 = {
|
|
kernelName: LogicalNot,
|
|
backendName: "webgl",
|
|
kernelFunc: logicalNot$1
|
|
};
|
|
const LOGICAL_OR = `return float(a >= 1.0 || b >= 1.0);`;
|
|
const LOGICAL_OR_PACKED = `
|
|
return min(
|
|
vec4(greaterThanEqual(a, vec4(1.0))) +
|
|
vec4(greaterThanEqual(b, vec4(1.0))),
|
|
vec4(1.0));
|
|
`;
|
|
const logicalOr$1 = binaryKernelFunc({ opSnippet: LOGICAL_OR, packedOpSnippet: LOGICAL_OR_PACKED, dtype: "bool" });
|
|
const logicalOrConfig$1 = {
|
|
kernelName: LogicalOr,
|
|
backendName: "webgl",
|
|
kernelFunc: logicalOr$1
|
|
};
|
|
class LRNProgram {
|
|
constructor(xShape, radius, bias, alpha, beta) {
|
|
this.variableNames = ["x"];
|
|
this.outputShape = [];
|
|
const rad = radius;
|
|
const maxD = xShape[3] - 1;
|
|
this.outputShape = xShape;
|
|
let powOperator;
|
|
const basis = `float(${bias}) + float(${alpha}) * sum`;
|
|
if (beta === 0.5) {
|
|
powOperator = `inversesqrt(${basis})`;
|
|
} else if (beta === 1) {
|
|
powOperator = `1.0/(${basis})`;
|
|
} else {
|
|
powOperator = `exp(log(${basis}) * float(-${beta}));`;
|
|
}
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int r = coords[1];
|
|
int c = coords[2];
|
|
int d = coords[3];
|
|
float x = getX(b, r, c, d);
|
|
float sum = 0.0;
|
|
for (int j = -${rad}; j <= ${rad}; j++) {
|
|
int idx = d + j;
|
|
if (idx >= 0 && idx <= ${maxD}) {
|
|
float z = getX(b, r, c, idx);
|
|
sum += z * z;
|
|
}
|
|
}
|
|
float val = x * ${powOperator};
|
|
setOutput(val);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class LRNPackedProgram {
|
|
constructor(xShape, radius, bias, alpha, beta) {
|
|
this.variableNames = ["x"];
|
|
this.outputShape = [];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
const rad = radius;
|
|
const maxD = xShape[3] - 1;
|
|
this.outputShape = xShape;
|
|
let powOperator;
|
|
const basis = `float(${bias}) + float(${alpha}) * sum`;
|
|
if (beta === 0.5) {
|
|
powOperator = `inversesqrt(${basis})`;
|
|
} else if (beta === 1) {
|
|
powOperator = `1.0/(${basis})`;
|
|
} else {
|
|
powOperator = `exp(log(${basis}) * float(-${beta}));`;
|
|
}
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords.x;
|
|
int r = coords.y;
|
|
int c = coords.z;
|
|
int d = coords.w;
|
|
|
|
bool hasNextCol = d < ${this.outputShape[3]};
|
|
bool hasNextRow = c < ${this.outputShape[2]};
|
|
|
|
vec4 sum = vec4(0.);
|
|
vec4 xFragAtOutputCoords = getX(b, r, c, d);
|
|
|
|
vec4 xAtOutputCoords = vec4(
|
|
getChannel(xFragAtOutputCoords, vec2(c, d)),
|
|
hasNextCol ?
|
|
getChannel(xFragAtOutputCoords, vec2(c, d + 1)) : 0.0,
|
|
hasNextRow ?
|
|
getChannel(xFragAtOutputCoords , vec2(c + 1, d)) : 0.0,
|
|
(hasNextRow && hasNextCol) ?
|
|
getChannel(xFragAtOutputCoords, vec2(c + 1, d + 1)) : 0.0
|
|
);
|
|
|
|
int firstChannel = d - ${rad};
|
|
vec2 cache = vec2(0.);
|
|
if(firstChannel >= 0){
|
|
vec4 firstChannelFrag = getX(b, r, c, firstChannel);
|
|
cache.x = getChannel(firstChannelFrag, vec2(c, firstChannel));
|
|
if(hasNextRow){
|
|
cache.y = getChannel(firstChannelFrag, vec2(c + 1, firstChannel));
|
|
}
|
|
}
|
|
|
|
ivec2 depth = ivec2(d, d + 1);
|
|
for (int j = - ${rad}; j <= ${rad}; j++) {
|
|
ivec2 idx = depth + j;
|
|
bvec2 aboveLowerBound = greaterThanEqual(idx, ivec2(0));
|
|
bvec2 belowUpperBound = lessThanEqual(idx, ivec2(${maxD}));
|
|
|
|
bool depthInRange = aboveLowerBound.x && belowUpperBound.x;
|
|
bool depthPlusOneInRange = aboveLowerBound.y && belowUpperBound.y;
|
|
|
|
if(depthInRange || depthPlusOneInRange){
|
|
vec4 z = vec4(0.);
|
|
vec4 xFragAtCurrentDepth;
|
|
z.xz = cache.xy;
|
|
if(depthPlusOneInRange && hasNextCol){
|
|
xFragAtCurrentDepth = idx.y != d ?
|
|
getX(b, r, c, idx.y) : xFragAtOutputCoords;
|
|
z.y = getChannel(xFragAtCurrentDepth, vec2(c, idx.y));
|
|
if(hasNextRow){
|
|
z.w = getChannel(xFragAtCurrentDepth, vec2(c + 1, idx.y));
|
|
}
|
|
}
|
|
cache.xy = z.yw;
|
|
sum += z * z;
|
|
}
|
|
}
|
|
vec4 result = xAtOutputCoords * ${powOperator};
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const lrn = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { depthRadius, bias, alpha, beta } = attrs;
|
|
const program = env().getBool("WEBGL_PACK_NORMALIZATION") ? new LRNPackedProgram(x.shape, depthRadius, bias, alpha, beta) : new LRNProgram(x.shape, depthRadius, bias, alpha, beta);
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
};
|
|
const LRNConfig$1 = {
|
|
kernelName: LRN,
|
|
backendName: "webgl",
|
|
kernelFunc: lrn
|
|
};
|
|
class LRNGradProgram {
|
|
constructor(inputShape, depthRadius, bias, alpha, beta) {
|
|
this.variableNames = ["inputImage", "outputImage", "dy"];
|
|
this.outputShape = [];
|
|
this.outputShape = inputShape;
|
|
this.depth = inputShape[3];
|
|
this.depthRadius = depthRadius;
|
|
this.bias = bias;
|
|
this.alpha = alpha;
|
|
this.beta = beta;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int r = coords[1];
|
|
int c = coords[2];
|
|
|
|
float result = 0.0;
|
|
for (int d = 0; d < ${this.depth}; ++d) {
|
|
int depthBegin = int(max(0.0, float(d - ${depthRadius})));
|
|
int depthEnd = int(min(float(${this.depth}),
|
|
float(d + ${depthRadius} + 1)));
|
|
|
|
const int MIN_DEPTH_BEGIN = 0;
|
|
const int MAX_DEPTH_END = ${this.depth};
|
|
|
|
float norm = 0.0;
|
|
for (int k = MIN_DEPTH_BEGIN; k < MAX_DEPTH_END; ++k) {
|
|
if (k < depthBegin){
|
|
continue;
|
|
}
|
|
else if (k >= depthBegin && k < depthEnd) {
|
|
norm += getInputImage(b, r, c, k) * getInputImage(b, r, c, k);
|
|
}
|
|
else {
|
|
break;
|
|
}
|
|
}
|
|
|
|
norm = float(${alpha}) * norm + float(${bias});
|
|
|
|
for(int k = MIN_DEPTH_BEGIN; k < MAX_DEPTH_END; ++k){
|
|
if (k < depthBegin){
|
|
continue;
|
|
}
|
|
else if (k >= depthBegin && k < depthEnd){
|
|
float dyi = -2.0 * float(${alpha})
|
|
* float(${beta})
|
|
* getInputImage(b, r, c, k) * getOutputImage(b, r, c, d)
|
|
/ norm;
|
|
if (k == d) {
|
|
dyi += pow(norm, -1.0 * ${beta});
|
|
}
|
|
if (k == coords[3]) {
|
|
dyi *= getDy(b, r, c, d);
|
|
result += dyi;
|
|
}
|
|
}
|
|
else {
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const lrnGrad = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, y, dy } = inputs;
|
|
const { depthRadius, bias, alpha, beta } = attrs;
|
|
const program = new LRNGradProgram(x.shape, depthRadius, bias, alpha, beta);
|
|
return backend2.runWebGLProgram(program, [x, y, dy], x.dtype);
|
|
};
|
|
const LRNGradConfig$1 = {
|
|
kernelName: LRNGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: lrnGrad
|
|
};
|
|
function maxImpl(x, reduceShape, outShape, backend2) {
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const xSize = sizeFromShape(x.shape);
|
|
const batchSize = xSize / inSize;
|
|
const reshapedInput = reshape$1({ inputs: { x }, attrs: { shape: [batchSize, inSize] }, backend: backend2 });
|
|
const reduced = reduce(reshapedInput, x.dtype, "max", backend2);
|
|
const reshapedOutput = reshape$1({ inputs: { x: reduced }, attrs: { shape: outShape }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(reshapedInput);
|
|
backend2.disposeIntermediateTensorInfo(reduced);
|
|
return reshapedOutput;
|
|
}
|
|
function max$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { reductionIndices, keepDims } = attrs;
|
|
const xRank = x.shape.length;
|
|
const origAxes = parseAxisParam(reductionIndices, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
const maxInputIsTransposed = permutedAxes != null;
|
|
const shouldExecuteOnCPU = backend2.shouldExecuteOnCPU([x]);
|
|
let maxInput = x;
|
|
if (maxInputIsTransposed) {
|
|
if (shouldExecuteOnCPU) {
|
|
const xTexData = backend2.texData.get(maxInput.dataId);
|
|
const values = xTexData.values;
|
|
const newShape = new Array(xRank);
|
|
for (let i = 0; i < newShape.length; i++) {
|
|
newShape[i] = x.shape[permutedAxes[i]];
|
|
}
|
|
const maxInputValues = transposeImplCPU(values, x.shape, x.dtype, permutedAxes, newShape);
|
|
maxInput = backend2.makeTensorInfo(newShape, x.dtype);
|
|
const maxInputData = backend2.texData.get(maxInput.dataId);
|
|
maxInputData.values = maxInputValues;
|
|
} else {
|
|
maxInput = transposeImpl(x, permutedAxes, backend2);
|
|
}
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
}
|
|
assertAxesAreInnerMostDims("max", axes, xRank);
|
|
const [maxOutShape, reduceShape] = computeOutAndReduceShapes(maxInput.shape, axes);
|
|
let outShape = maxOutShape;
|
|
if (keepDims) {
|
|
outShape = expandShapeToKeepDim(maxOutShape, origAxes);
|
|
}
|
|
let out;
|
|
if (shouldExecuteOnCPU) {
|
|
const xTexData = backend2.texData.get(maxInput.dataId);
|
|
const values = xTexData.values;
|
|
const outValues = maxImplCPU(values, sizeFromShape(reduceShape), outShape, x.dtype);
|
|
out = backend2.makeTensorInfo(outShape, x.dtype);
|
|
const outData = backend2.texData.get(out.dataId);
|
|
outData.values = outValues;
|
|
} else {
|
|
out = maxImpl(maxInput, reduceShape, outShape, backend2);
|
|
}
|
|
if (maxInputIsTransposed) {
|
|
backend2.disposeIntermediateTensorInfo(maxInput);
|
|
}
|
|
return out;
|
|
}
|
|
const maxConfig$1 = {
|
|
kernelName: Max,
|
|
backendName: "webgl",
|
|
kernelFunc: max$1
|
|
};
|
|
const MAXIMUM = CHECK_NAN_SNIPPET + `
|
|
return max(a, b);
|
|
`;
|
|
const MAXIMUM_PACKED = `
|
|
vec4 result = vec4(max(a, b));
|
|
bvec4 isNaNA = isnan(a);
|
|
bvec4 isNaNB = isnan(b);
|
|
bvec4 isNaN = bvec4(isNaNA.x || isNaNB.x, isNaNA.y || isNaNB.y, isNaNA.z || isNaNB.z, isNaNA.w || isNaNB.w);
|
|
` + CHECK_NAN_SNIPPET_PACKED + `
|
|
return result;
|
|
`;
|
|
const maximum = binaryKernelFunc({
|
|
opSnippet: MAXIMUM,
|
|
packedOpSnippet: MAXIMUM_PACKED,
|
|
cpuKernelImpl: maximumImplCPU
|
|
});
|
|
const maximumConfig = {
|
|
kernelName: Maximum,
|
|
backendName: "webgl",
|
|
kernelFunc: maximum
|
|
};
|
|
function maxPool$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
assertNotComplex$1(x, "maxPool");
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const dilations = 1;
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in maxPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode);
|
|
if (convInfo.filterWidth === 1 && convInfo.filterHeight === 1 && arraysEqual(convInfo.inShape, convInfo.outShape)) {
|
|
return identity({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const maxPoolProgram = new Pool2DProgram(convInfo, "max", false);
|
|
return backend2.runWebGLProgram(maxPoolProgram, [x], x.dtype);
|
|
}
|
|
const maxPoolConfig$1 = {
|
|
kernelName: MaxPool,
|
|
backendName: "webgl",
|
|
kernelFunc: maxPool$1
|
|
};
|
|
function maxPool3d(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { filterSize, strides, pad: pad2, dataFormat, dimRoundingMode } = attrs;
|
|
const dilations = [1, 1, 1];
|
|
const convInfo = computePool3DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode, dataFormat);
|
|
const maxPoolProgram = new Pool3DProgram(convInfo, "max", false);
|
|
return backend2.runWebGLProgram(maxPoolProgram, [x], x.dtype);
|
|
}
|
|
const maxPool3DConfig$1 = {
|
|
kernelName: MaxPool3D,
|
|
backendName: "webgl",
|
|
kernelFunc: maxPool3d
|
|
};
|
|
class MaxPool2DBackpropProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy", "maxPos"];
|
|
this.outputShape = convInfo.inShape;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const lastIndex = effectiveFilterHeight * effectiveFilterWidth - 1;
|
|
this.userCode = `
|
|
const ivec2 pads = ivec2(${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
|
|
ivec2 dyRCCorner = coords.yz - pads;
|
|
int dyRCorner = dyRCCorner.x;
|
|
int dyCCorner = dyRCCorner.y;
|
|
|
|
// Convolve dy(?, ?, d) with pos mask(:, :, d) to get dx(xR, xC, d).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 || fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
for (int wC = 0; wC < ${effectiveFilterWidth}; wC++) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
float dyValue = getDy(b, idyR, idyC, d);
|
|
int maxPosValue = ${lastIndex} - int(getMaxPos(b, idyR, idyC, d));
|
|
|
|
// Get the current value, check it against the value from the
|
|
// position matrix.
|
|
int curPosValue = wR * ${effectiveFilterWidth} + wC;
|
|
float mask = float(maxPosValue == curPosValue ? 1.0 : 0.0);
|
|
|
|
dotProd += dyValue * mask;
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class MaxPool3DBackpropProgram {
|
|
constructor(convInfo) {
|
|
this.variableNames = ["dy", "maxPos"];
|
|
this.outputShape = convInfo.inShape;
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = effectiveFilterDepth - 1 - convInfo.padInfo.front;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const lastIndex = effectiveFilterDepth * effectiveFilterHeight * effectiveFilterWidth - 1;
|
|
this.userCode = `
|
|
const ivec3 pads = ivec3(${padFront}, ${padTop}, ${padLeft});
|
|
|
|
void main() {
|
|
ivec5 coords = getOutputCoords();
|
|
int batch = coords.x;
|
|
int ch = coords.u;
|
|
|
|
ivec3 dyCorner = ivec3(coords.y, coords.z, coords.w) - pads;
|
|
int dyDCorner = dyCorner.x;
|
|
int dyRCorner = dyCorner.y;
|
|
int dyCCorner = dyCorner.z;
|
|
|
|
// Convolve dy(?, ?, ?, ch) with pos mask(:, :, :, d) to get
|
|
// dx(xD, xR, xC, ch).
|
|
// ? = to be determined. : = across all values in that axis.
|
|
float dotProd = 0.0;
|
|
|
|
for (int wD = 0; wD < ${effectiveFilterDepth};
|
|
wD += ${dilationDepth}) {
|
|
float dyD = float(dyDCorner + wD) / ${strideDepth}.0;
|
|
|
|
if (dyD < 0.0 || dyD >= ${convInfo.outDepth}.0 || fract(dyD) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyD = int(dyD);
|
|
|
|
for (int wR = 0; wR < ${effectiveFilterHeight};
|
|
wR += ${dilationHeight}) {
|
|
float dyR = float(dyRCorner + wR) / ${strideHeight}.0;
|
|
|
|
if (dyR < 0.0 || dyR >= ${convInfo.outHeight}.0 ||
|
|
fract(dyR) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyR = int(dyR);
|
|
|
|
for (int wC = 0; wC < ${effectiveFilterWidth};
|
|
wC += ${dilationWidth}) {
|
|
float dyC = float(dyCCorner + wC) / ${strideWidth}.0;
|
|
|
|
if (dyC < 0.0 || dyC >= ${convInfo.outWidth}.0 ||
|
|
fract(dyC) > 0.0) {
|
|
continue;
|
|
}
|
|
int idyC = int(dyC);
|
|
|
|
float dyValue = getDy(batch, idyD, idyR, idyC, ch);
|
|
int maxPosValue = ${lastIndex} -
|
|
int(getMaxPos(batch, idyD, idyR, idyC, ch));
|
|
|
|
// Get the current value, check it against the value from the
|
|
// position matrix.
|
|
int curPosValue =
|
|
wD * ${effectiveFilterHeight} * ${effectiveFilterWidth} +
|
|
wR * ${effectiveFilterWidth} + wC;
|
|
float mask = float(maxPosValue == curPosValue ? 1.0 : 0.0);
|
|
|
|
dotProd += dyValue * mask;
|
|
}
|
|
}
|
|
}
|
|
setOutput(dotProd);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function maxPool3DGrad$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input } = inputs;
|
|
const x = input;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const dilations = [1, 1, 1];
|
|
const convInfo = computePool3DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode);
|
|
const maxPool3dPositionsProgram = new Pool3DProgram(
|
|
convInfo,
|
|
"max",
|
|
true
|
|
/* get positions */
|
|
);
|
|
const maxPool3dPositions2 = backend2.runWebGLProgram(maxPool3dPositionsProgram, [x], x.dtype);
|
|
const maxPoolBackpropProgram = new MaxPool3DBackpropProgram(convInfo);
|
|
const result = backend2.runWebGLProgram(maxPoolBackpropProgram, [dy, maxPool3dPositions2], x.dtype);
|
|
backend2.disposeIntermediateTensorInfo(maxPool3dPositions2);
|
|
return result;
|
|
}
|
|
const maxPool3DGradConfig$1 = {
|
|
kernelName: MaxPool3DGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: maxPool3DGrad$1
|
|
};
|
|
function maxPoolGrad$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input, output } = inputs;
|
|
const x = input;
|
|
assertNotComplex$1([input, output], "maxPoolGrad");
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, 1, pad2, dimRoundingMode);
|
|
const getPositions = true;
|
|
const maxPoolPositionsProgram = new Pool2DProgram(convInfo, "max", getPositions);
|
|
const maxPoolPositions2 = backend2.runWebGLProgram(maxPoolPositionsProgram, [x], x.dtype);
|
|
const maxPoolBackPropProgram = new MaxPool2DBackpropProgram(convInfo);
|
|
const result = backend2.runWebGLProgram(maxPoolBackPropProgram, [dy, maxPoolPositions2], x.dtype);
|
|
backend2.disposeIntermediateTensorInfo(maxPoolPositions2);
|
|
return result;
|
|
}
|
|
const maxPoolGradConfig$1 = {
|
|
kernelName: MaxPoolGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: maxPoolGrad$1
|
|
};
|
|
function maxPoolWithArgmaxImpl$1(x, includeBatchInIndex, convInfo, backend2) {
|
|
let program = new Pool2DProgram(convInfo, "max", false);
|
|
const poolOutput = backend2.runWebGLProgram(program, [x], "float32");
|
|
program = new Pool2DProgram(convInfo, "max", true, true, includeBatchInIndex);
|
|
const indexOutput = backend2.runWebGLProgram(program, [x], "float32");
|
|
return [poolOutput, indexOutput];
|
|
}
|
|
const maxPoolWithArgmaxConfig$1 = {
|
|
kernelName: MaxPoolWithArgmax,
|
|
backendName: "webgl",
|
|
kernelFunc: ({ inputs, attrs, backend: backend2 }) => {
|
|
const { x } = inputs;
|
|
const { filterSize, strides, pad: pad2, includeBatchInIndex } = attrs;
|
|
const webglBackend = backend2;
|
|
assert(x.shape.length === 4, () => `Error in maxPool: input must be rank 4 but got rank ${x.shape.length}.`);
|
|
const dilations = [1, 1];
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in maxPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, dilations, pad2);
|
|
const [result, indexes] = maxPoolWithArgmaxImpl$1(x, includeBatchInIndex, convInfo, webglBackend);
|
|
return [result, indexes];
|
|
}
|
|
};
|
|
function meanImpl(x, reduceShape, outShape, backend2) {
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const xSize = sizeFromShape(x.shape);
|
|
const batchSize = xSize / inSize;
|
|
const reshapedInput = reshape$1({ inputs: { x }, attrs: { shape: [batchSize, inSize] }, backend: backend2 });
|
|
const reduced = reduce(reshapedInput, "float32", "mean", backend2);
|
|
const reshapedOutput = reshape$1({ inputs: { x: reduced }, attrs: { shape: outShape }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(reshapedInput);
|
|
backend2.disposeIntermediateTensorInfo(reduced);
|
|
return reshapedOutput;
|
|
}
|
|
const meanConfig$1 = {
|
|
kernelName: Mean,
|
|
backendName: "webgl",
|
|
kernelFunc: ({ inputs, attrs, backend: backend2 }) => {
|
|
const { x } = inputs;
|
|
const { keepDims, axis } = attrs;
|
|
const webglBackend = backend2;
|
|
const xRank = x.shape.length;
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
const meanInputIsTransposed = permutedAxes != null;
|
|
const shouldExecuteOnCPU = webglBackend.shouldExecuteOnCPU([x]);
|
|
const intermediates = [];
|
|
let meanInput = x;
|
|
if (meanInputIsTransposed) {
|
|
if (shouldExecuteOnCPU) {
|
|
const xTexData = webglBackend.texData.get(meanInput.dataId);
|
|
const values = xTexData.values;
|
|
const newShape = new Array(xRank);
|
|
for (let i = 0; i < newShape.length; i++) {
|
|
newShape[i] = x.shape[permutedAxes[i]];
|
|
}
|
|
const meanInputValues = transposeImplCPU(values, x.shape, x.dtype, permutedAxes, newShape);
|
|
meanInput = webglBackend.makeTensorInfo(newShape, x.dtype);
|
|
const meanInputData = webglBackend.texData.get(meanInput.dataId);
|
|
meanInputData.values = meanInputValues;
|
|
} else {
|
|
meanInput = transposeImpl(x, permutedAxes, webglBackend);
|
|
}
|
|
intermediates.push(meanInput);
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
}
|
|
assertAxesAreInnerMostDims("sum", axes, xRank);
|
|
const [meanOutShape, reduceShape] = computeOutAndReduceShapes(meanInput.shape, axes);
|
|
let outShape = meanOutShape;
|
|
if (keepDims) {
|
|
outShape = expandShapeToKeepDim(meanOutShape, origAxes);
|
|
}
|
|
const out = meanImpl(meanInput, reduceShape, outShape, webglBackend);
|
|
for (const i of intermediates) {
|
|
webglBackend.disposeIntermediateTensorInfo(i);
|
|
}
|
|
return out;
|
|
}
|
|
};
|
|
function min$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
const xRank = x.shape.length;
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
let permutedX = x;
|
|
if (permutedAxes != null) {
|
|
permutedX = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, x.shape.length);
|
|
}
|
|
assertAxesAreInnerMostDims("min", axes, xRank);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(permutedX.shape, axes);
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const a2D = reshape$1({ inputs: { x: permutedX }, backend: backend2, attrs: { shape: [-1, inSize] } });
|
|
const reduced = reduce(a2D, a2D.dtype, "min", backend2);
|
|
let res;
|
|
if (keepDims) {
|
|
const newShape = expandShapeToKeepDim(outShape, origAxes);
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: newShape } });
|
|
} else {
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: outShape } });
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(a2D);
|
|
backend2.disposeIntermediateTensorInfo(reduced);
|
|
if (permutedAxes != null) {
|
|
backend2.disposeIntermediateTensorInfo(permutedX);
|
|
}
|
|
return res;
|
|
}
|
|
const minConfig$1 = {
|
|
kernelName: Min,
|
|
backendName: "webgl",
|
|
kernelFunc: min$1
|
|
};
|
|
const MINIMUM = CHECK_NAN_SNIPPET + `
|
|
return min(a, b);
|
|
`;
|
|
const MINIMUM_PACKED = `
|
|
vec4 result = vec4(min(a, b));
|
|
bvec4 isNaNA = isnan(a);
|
|
bvec4 isNaNB = isnan(b);
|
|
bvec4 isNaN = bvec4(isNaNA.x || isNaNB.x, isNaNA.y || isNaNB.y, isNaNA.z || isNaNB.z, isNaNA.w || isNaNB.w);
|
|
` + CHECK_NAN_SNIPPET_PACKED + `
|
|
return result;
|
|
`;
|
|
const minimum = binaryKernelFunc({
|
|
opSnippet: MINIMUM,
|
|
packedOpSnippet: MINIMUM_PACKED,
|
|
cpuKernelImpl: minimumImplCPU
|
|
});
|
|
const minimumConfig = {
|
|
kernelName: Minimum,
|
|
backendName: "webgl",
|
|
kernelFunc: minimum
|
|
};
|
|
class MirrorPadProgram {
|
|
constructor(xShape, paddings, mode) {
|
|
this.variableNames = ["x"];
|
|
this.outputShape = paddings.map(
|
|
(p2, i) => p2[0] + xShape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
const rank = xShape.length;
|
|
const dtype = getCoordsDataType(rank);
|
|
const start = paddings.map((p2) => p2[0]).join(",");
|
|
const end = paddings.map((p2, i) => p2[0] + xShape[i]).join(",");
|
|
const unpackedCoords = ["coords[0]", "coords[1]", "coords[2]", "coords[3]"].slice(0, rank);
|
|
const offset = mode === "reflect" ? 0 : 1;
|
|
if (rank === 1) {
|
|
this.userCode = `
|
|
int start = ${start};
|
|
int end = ${end};
|
|
|
|
void main() {
|
|
int outC = getOutputCoords();
|
|
if (outC < start) {
|
|
outC = start * 2 - outC - ${offset};
|
|
} else if(outC >= end) {
|
|
outC = (end - 1) * 2 - outC + ${offset};
|
|
}
|
|
setOutput(getX(outC - start));
|
|
}
|
|
`;
|
|
return;
|
|
}
|
|
this.userCode = `
|
|
${dtype} start = ${dtype}(${start});
|
|
${dtype} end = ${dtype}(${end});
|
|
|
|
void main() {
|
|
${dtype} outC = getOutputCoords();
|
|
for (int i = 0; i < ${rank}; i++) {
|
|
if (outC[i] < start[i]) {
|
|
outC[i] = start[i] * 2 - outC[i] - ${offset};
|
|
} else if(outC[i] >= end[i]) {
|
|
outC[i] = (end[i] - 1) * 2 - outC[i] + ${offset};
|
|
}
|
|
}
|
|
${dtype} coords = outC - start;
|
|
setOutput(getX(${unpackedCoords}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class MirrorPadPackedProgram {
|
|
constructor(xShape, paddings, mode) {
|
|
this.variableNames = ["x"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = paddings.map(
|
|
(p2, i) => p2[0] + xShape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
const rank = xShape.length;
|
|
const dtype = getCoordsDataType(rank);
|
|
const start = paddings.map((p2) => p2[0]).join(",");
|
|
const end = paddings.map((p2, i) => p2[0] + xShape[i]).join(",");
|
|
const coords2 = getChannels("rc", rank);
|
|
const source = getChannels("source", rank);
|
|
const cLimit = `${coords2[rank - 1]} < ${this.outputShape[rank - 1]}`;
|
|
const innerDims = rank === 1 ? "source" : `vec2(${source.slice(-2).join()})`;
|
|
const offset = mode === "reflect" ? 0 : 1;
|
|
let mainLoop = "";
|
|
if (rank === 1) {
|
|
const padSetup = `
|
|
${dtype} source = rc;
|
|
if (source < start) {
|
|
source = start * 2 - source - ${offset};
|
|
} else if (source >= end) {
|
|
source = (end - 1) * 2 - source + ${offset};
|
|
}
|
|
source -= start;
|
|
`;
|
|
mainLoop = `
|
|
${dtype} rc = outputLoc;
|
|
${padSetup}
|
|
result[0] = getChannel(getX(${source.join()}), ${innerDims});
|
|
${coords2[rank - 1]} += 1;
|
|
if(${cLimit}) {
|
|
${padSetup}
|
|
result[1] = getChannel(getX(${source.join()}), ${innerDims});
|
|
}
|
|
`;
|
|
} else {
|
|
const padSetup = `
|
|
${dtype} source = rc;
|
|
${dtype} lt = ${dtype}(lessThan(source, start));
|
|
${dtype} gte = ${dtype}(greaterThanEqual(source, end));
|
|
${dtype} orig = 1 - (lt + gte);
|
|
source = orig * source +
|
|
lt * (start * 2 - source - ${offset}) +
|
|
gte * ((end - 1) * 2 - source + ${offset});
|
|
source -= start;
|
|
`;
|
|
mainLoop = `
|
|
${dtype} rc = outputLoc;
|
|
${padSetup}
|
|
result[0] = getChannel(getX(${source.join()}), ${innerDims});
|
|
${coords2[rank - 1]} += 1;
|
|
if(${cLimit}) {
|
|
${padSetup}
|
|
result[1] = getChannel(getX(${source.join()}), ${innerDims});
|
|
}
|
|
rc = outputLoc;
|
|
${coords2[rank - 2]} += 1;
|
|
if(${coords2[rank - 2]} < ${this.outputShape[rank - 2]}) {
|
|
${padSetup}
|
|
result[2] = getChannel(getX(${source.join()}), ${innerDims});
|
|
${coords2[rank - 1]} += 1;
|
|
if(${cLimit}) {
|
|
${padSetup}
|
|
result[3] = getChannel(getX(${source.join()}), ${innerDims});
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
this.userCode = `
|
|
const ${dtype} start = ${dtype}(${start});
|
|
const ${dtype} end = ${dtype}(${end});
|
|
|
|
void main() {
|
|
${dtype} outputLoc = getOutputCoords();
|
|
vec4 result = vec4(0.);
|
|
${mainLoop}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const mirrorPadKernelFunc = ({ inputs, backend: backend2, attrs }) => {
|
|
const { x } = inputs;
|
|
const { paddings, mode } = attrs;
|
|
const program = env().getBool("WEBGL_PACK_ARRAY_OPERATIONS") ? new MirrorPadPackedProgram(x.shape, paddings, mode) : new MirrorPadProgram(x.shape, paddings, mode);
|
|
const output = backend2.runWebGLProgram(program, [x], x.dtype);
|
|
return output;
|
|
};
|
|
const mirrorPadConfig$1 = {
|
|
kernelName: MirrorPad,
|
|
backendName: "webgl",
|
|
kernelFunc: mirrorPadKernelFunc
|
|
};
|
|
const MOD = `if (b == 0.0) return NAN;
|
|
return mod(a, b);`;
|
|
const MOD_PACKED = `
|
|
vec4 result = mod(a, b);
|
|
bvec4 isNaN = equal(b, vec4(0.0));
|
|
` + CHECK_NAN_SNIPPET_PACKED + `
|
|
return result;
|
|
`;
|
|
const mod$1 = binaryKernelFunc({
|
|
opSnippet: MOD,
|
|
packedOpSnippet: MOD_PACKED
|
|
});
|
|
const modConfig$1 = {
|
|
kernelName: Mod,
|
|
backendName: "webgl",
|
|
kernelFunc: mod$1
|
|
};
|
|
class MultinomialProgram {
|
|
constructor(batchSize, numOutcomes, numSamples) {
|
|
this.variableNames = ["probs"];
|
|
this.customUniforms = [{ name: "seed", type: "float" }];
|
|
this.outputShape = [batchSize, numSamples];
|
|
this.userCode = `
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
|
|
float r = random(seed);
|
|
float cdf = 0.0;
|
|
|
|
for (int i = 0; i < ${numOutcomes - 1}; i++) {
|
|
cdf += getProbs(batch, i);
|
|
|
|
if (r < cdf) {
|
|
setOutput(float(i));
|
|
return;
|
|
}
|
|
}
|
|
|
|
// If no other event happened, last event happened.
|
|
setOutput(float(${numOutcomes - 1}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const DIV = `
|
|
if (a == b) {
|
|
return 1.0;
|
|
};
|
|
return a / b;`;
|
|
const DIV_PACKED = `
|
|
// vec4 one = vec4(equal(a, b));
|
|
// return one + (vec4(1.0) - one) * a / b;
|
|
vec4 result = a / b;
|
|
if(a.x == b.x) {
|
|
result.x = 1.;
|
|
}
|
|
if(a.y == b.y) {
|
|
result.y = 1.;
|
|
}
|
|
if(a.z == b.z) {
|
|
result.z = 1.;
|
|
}
|
|
if(a.w == b.w) {
|
|
result.w = 1.;
|
|
}
|
|
|
|
return result;
|
|
`;
|
|
const realDiv = binaryKernelFunc({ opSnippet: DIV, packedOpSnippet: DIV_PACKED, checkOutOfBounds: true });
|
|
const realDivConfig$1 = {
|
|
kernelName: RealDiv,
|
|
backendName: "webgl",
|
|
kernelFunc: realDiv
|
|
};
|
|
const SUB = "return a - b;";
|
|
const sub = binaryKernelFunc({
|
|
opSnippet: SUB,
|
|
packedOpSnippet: SUB,
|
|
supportsComplex: true,
|
|
cpuKernelImpl: subImplCPU
|
|
});
|
|
const subConfig = {
|
|
kernelName: Sub,
|
|
backendName: "webgl",
|
|
kernelFunc: sub
|
|
};
|
|
function softmax$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { logits } = inputs;
|
|
const { dim } = attrs;
|
|
const axes = parseAxisParam([dim], logits.shape);
|
|
const maxLogit = max$1({
|
|
inputs: { x: logits },
|
|
backend: backend2,
|
|
attrs: { reductionIndices: axes, keepDims: false }
|
|
});
|
|
const expandedShape = expandShapeToKeepDim(maxLogit.shape, axes);
|
|
const maxLogitsReshaped = reshape$1({ inputs: { x: maxLogit }, backend: backend2, attrs: { shape: expandedShape } });
|
|
const a = sub({ inputs: { a: logits, b: maxLogitsReshaped }, backend: backend2 });
|
|
const b = exp({ inputs: { x: a }, backend: backend2 });
|
|
const sumExp = sum$1({ inputs: { x: b }, backend: backend2, attrs: { axis: axes, keepDims: false } });
|
|
const sumExpReshaped = reshape$1({ inputs: { x: sumExp }, backend: backend2, attrs: { shape: expandedShape } });
|
|
const res = realDiv({ inputs: { a: b, b: sumExpReshaped }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(maxLogit);
|
|
backend2.disposeIntermediateTensorInfo(maxLogitsReshaped);
|
|
backend2.disposeIntermediateTensorInfo(a);
|
|
backend2.disposeIntermediateTensorInfo(b);
|
|
backend2.disposeIntermediateTensorInfo(sumExp);
|
|
backend2.disposeIntermediateTensorInfo(sumExpReshaped);
|
|
return res;
|
|
}
|
|
const softmaxConfig$1 = {
|
|
kernelName: Softmax,
|
|
backendName: "webgl",
|
|
kernelFunc: softmax$1
|
|
};
|
|
function multinomial$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { logits } = inputs;
|
|
const { numSamples, seed, normalized } = attrs;
|
|
const probs = normalized ? logits : softmax$1({ inputs: { logits }, backend: backend2, attrs: { dim: logits.shape.length - 1 } });
|
|
const batchSize = probs.shape[0];
|
|
const numOutcomes = probs.shape[1];
|
|
const program = new MultinomialProgram(batchSize, numOutcomes, numSamples);
|
|
const customValues = [[seed]];
|
|
const res = backend2.runWebGLProgram(program, [probs], "int32", customValues);
|
|
if (!normalized) {
|
|
backend2.disposeIntermediateTensorInfo(probs);
|
|
}
|
|
return res;
|
|
}
|
|
const multinomialConfig$1 = {
|
|
kernelName: Multinomial,
|
|
backendName: "webgl",
|
|
kernelFunc: multinomial$1
|
|
};
|
|
const NEG = CHECK_NAN_SNIPPET$1 + `
|
|
return -x;
|
|
`;
|
|
const NEG_PACKED = `
|
|
vec4 result = -x;
|
|
bvec4 isNaN = isnan(x);
|
|
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
function neg(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
if (backend2.shouldExecuteOnCPU([x])) {
|
|
const xData = backend2.texData.get(x.dataId);
|
|
const [outValues, newShape] = negImplCPU(xData.values, x.shape, x.dtype);
|
|
return backend2.makeTensorInfo(newShape, x.dtype, outValues);
|
|
}
|
|
let program;
|
|
if (env().getBool("WEBGL_PACK_UNARY_OPERATIONS")) {
|
|
program = new UnaryOpPackedProgram(x.shape, NEG_PACKED);
|
|
} else {
|
|
program = new UnaryOpProgram(x.shape, NEG);
|
|
}
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
const negConfig = {
|
|
kernelName: Neg,
|
|
backendName: "webgl",
|
|
kernelFunc: neg
|
|
};
|
|
const nonMaxSuppressionV3Impl$1 = nonMaxSuppressionV3Impl$2;
|
|
function nonMaxSuppressionV3$1(args) {
|
|
warn("tf.nonMaxSuppression() in webgl locks the UI thread. Call tf.nonMaxSuppressionAsync() instead");
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { boxes, scores } = inputs;
|
|
const { maxOutputSize, iouThreshold, scoreThreshold } = attrs;
|
|
const boxesVals = backend2.readSync(boxes.dataId);
|
|
const scoresVals = backend2.readSync(scores.dataId);
|
|
const { selectedIndices } = nonMaxSuppressionV3Impl$1(boxesVals, scoresVals, maxOutputSize, iouThreshold, scoreThreshold);
|
|
return backend2.makeTensorInfo([selectedIndices.length], "int32", new Int32Array(selectedIndices));
|
|
}
|
|
const nonMaxSuppressionV3Config$1 = {
|
|
kernelName: NonMaxSuppressionV3,
|
|
backendName: "webgl",
|
|
kernelFunc: nonMaxSuppressionV3$1
|
|
};
|
|
const nonMaxSuppressionV4Impl$1 = nonMaxSuppressionV4Impl$2;
|
|
function nonMaxSuppressionV4$1(args) {
|
|
warn("tf.nonMaxSuppression() in webgl locks the UI thread. Call tf.nonMaxSuppressionAsync() instead");
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { boxes, scores } = inputs;
|
|
const { maxOutputSize, iouThreshold, scoreThreshold, padToMaxOutputSize } = attrs;
|
|
const boxesVals = backend2.readSync(boxes.dataId);
|
|
const scoresVals = backend2.readSync(scores.dataId);
|
|
const { selectedIndices, validOutputs } = nonMaxSuppressionV4Impl$1(boxesVals, scoresVals, maxOutputSize, iouThreshold, scoreThreshold, padToMaxOutputSize);
|
|
return [
|
|
backend2.makeTensorInfo([selectedIndices.length], "int32", new Int32Array(selectedIndices)),
|
|
backend2.makeTensorInfo([], "int32", new Int32Array([validOutputs]))
|
|
];
|
|
}
|
|
const nonMaxSuppressionV4Config$1 = {
|
|
kernelName: NonMaxSuppressionV4,
|
|
backendName: "webgl",
|
|
kernelFunc: nonMaxSuppressionV4$1
|
|
};
|
|
const nonMaxSuppressionV5Impl$1 = nonMaxSuppressionV5Impl$2;
|
|
function nonMaxSuppressionV5$1(args) {
|
|
warn("tf.nonMaxSuppression() in webgl locks the UI thread. Call tf.nonMaxSuppressionAsync() instead");
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { boxes, scores } = inputs;
|
|
const { maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma } = attrs;
|
|
const boxesVals = backend2.readSync(boxes.dataId);
|
|
const scoresVals = backend2.readSync(scores.dataId);
|
|
const maxOutputSizeVal = maxOutputSize;
|
|
const iouThresholdVal = iouThreshold;
|
|
const scoreThresholdVal = scoreThreshold;
|
|
const softNmsSigmaVal = softNmsSigma;
|
|
const { selectedIndices, selectedScores } = nonMaxSuppressionV5Impl$1(boxesVals, scoresVals, maxOutputSizeVal, iouThresholdVal, scoreThresholdVal, softNmsSigmaVal);
|
|
return [
|
|
backend2.makeTensorInfo([selectedIndices.length], "int32", new Int32Array(selectedIndices)),
|
|
backend2.makeTensorInfo([selectedScores.length], "float32", new Float32Array(selectedScores))
|
|
];
|
|
}
|
|
const nonMaxSuppressionV5Config$1 = {
|
|
kernelName: NonMaxSuppressionV5,
|
|
backendName: "webgl",
|
|
kernelFunc: nonMaxSuppressionV5$1
|
|
};
|
|
class OneHotProgram {
|
|
constructor(numIndices, depth, onValue, offValue) {
|
|
this.variableNames = ["indices"];
|
|
this.outputShape = [numIndices, depth];
|
|
this.userCode = `
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int index = round(getIndices(coords.x));
|
|
setOutput(mix(float(${offValue}), float(${onValue}),
|
|
float(index == coords.y)));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const oneHot$1 = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { indices } = inputs;
|
|
const { dtype, depth, onValue, offValue } = attrs;
|
|
const indicesSize = sizeFromShape(indices.shape);
|
|
const program = new OneHotProgram(indicesSize, depth, onValue, offValue);
|
|
const reshaped = reshape$1({ inputs: { x: indices }, backend: backend2, attrs: { shape: [indicesSize] } });
|
|
const result = backend2.runWebGLProgram(program, [reshaped], dtype);
|
|
backend2.disposeIntermediateTensorInfo(reshaped);
|
|
const outShape = [...indices.shape, depth];
|
|
const out = reshape$1({ inputs: { x: result }, backend: backend2, attrs: { shape: outShape } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return out;
|
|
};
|
|
const oneHotConfig$1 = {
|
|
kernelName: OneHot,
|
|
backendName: "webgl",
|
|
kernelFunc: oneHot$1
|
|
};
|
|
function zerosLike$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
if (x.dtype === "complex64") {
|
|
const realPart = real({ inputs: { input: x }, backend: backend2 });
|
|
const r = zerosLike$1({ inputs: { x: realPart }, backend: backend2 });
|
|
const imagPart = imag$1({ inputs: { input: x }, backend: backend2 });
|
|
const i = zerosLike$1({ inputs: { x: imagPart }, backend: backend2 });
|
|
const result = complex({ inputs: { real: r, imag: i }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
backend2.disposeIntermediateTensorInfo(r);
|
|
backend2.disposeIntermediateTensorInfo(imagPart);
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
return result;
|
|
} else {
|
|
return fill$1({
|
|
attrs: {
|
|
shape: x.shape,
|
|
dtype: x.dtype,
|
|
value: x.dtype === "string" ? "" : 0
|
|
},
|
|
backend: backend2
|
|
});
|
|
}
|
|
}
|
|
const zerosLikeConfig$1 = {
|
|
kernelName: ZerosLike,
|
|
backendName: "webgl",
|
|
kernelFunc: zerosLike$1
|
|
};
|
|
function onesLike$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
if (x.dtype === "string") {
|
|
throw new Error("onesLike is not supported under string dtype");
|
|
} else if (x.dtype === "complex64") {
|
|
const realPart = real({ inputs: { input: x }, backend: backend2 });
|
|
const r = onesLike$1({ inputs: { x: realPart }, backend: backend2 });
|
|
const imagPart = imag$1({ inputs: { input: x }, backend: backend2 });
|
|
const i = zerosLike$1({ inputs: { x: imagPart }, backend: backend2 });
|
|
const result = complex({ inputs: { real: r, imag: i }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
backend2.disposeIntermediateTensorInfo(r);
|
|
backend2.disposeIntermediateTensorInfo(imagPart);
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
return result;
|
|
} else {
|
|
return fill$1({ attrs: { shape: x.shape, dtype: x.dtype, value: 1 }, backend: backend2 });
|
|
}
|
|
}
|
|
const onesLikeConfig$1 = {
|
|
kernelName: OnesLike,
|
|
backendName: "webgl",
|
|
kernelFunc: onesLike$1
|
|
};
|
|
function pack$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { axis } = attrs;
|
|
if (inputs.length === 1) {
|
|
return expandDims$1({ inputs: { input: inputs[0] }, backend: backend2, attrs: { dim: axis } });
|
|
}
|
|
const shape = inputs[0].shape;
|
|
const dtype = inputs[0].dtype;
|
|
inputs.forEach((t2) => {
|
|
assertShapesMatch(shape, t2.shape, "All tensors passed to stack must have matching shapes");
|
|
assert(dtype === t2.dtype, () => "All tensors passed to stack must have matching dtypes");
|
|
});
|
|
const intermediateTensorInfos = [];
|
|
const expandedTensors = inputs.map((t2) => {
|
|
const expandedT = expandDims$1({ inputs: { input: t2 }, backend: backend2, attrs: { dim: axis } });
|
|
intermediateTensorInfos.push(expandedT);
|
|
return expandedT;
|
|
});
|
|
const result = concat$1({ inputs: expandedTensors, backend: backend2, attrs: { axis } });
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
}
|
|
const packConfig$1 = {
|
|
kernelName: Pack,
|
|
backendName: "webgl",
|
|
kernelFunc: pack$1
|
|
};
|
|
class PadProgram {
|
|
constructor(xShape, paddings, constantValue) {
|
|
this.variableNames = ["x"];
|
|
this.customUniforms = [{ name: "value", type: "float" }];
|
|
this.outputShape = paddings.map(
|
|
(p2, i) => p2[0] + xShape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
const rank = xShape.length;
|
|
const type = getCoordsDataType(rank);
|
|
const start = paddings.map((p2) => p2[0]).join(",");
|
|
const end = paddings.map((p2, i) => p2[0] + xShape[i]).join(",");
|
|
const unpackedCoords = ["coords[0]", "coords[1]", "coords[2]", "coords[3]"].slice(0, rank);
|
|
if (rank === 1) {
|
|
this.userCode = `
|
|
int start = ${start};
|
|
int end = ${end};
|
|
|
|
void main() {
|
|
int outC = getOutputCoords();
|
|
if (outC < start || outC >= end) {
|
|
setOutput(value);
|
|
} else {
|
|
setOutput(getX(outC - start));
|
|
}
|
|
}
|
|
`;
|
|
return;
|
|
}
|
|
this.userCode = `
|
|
${type} start = ${type}(${start});
|
|
${type} end = ${type}(${end});
|
|
|
|
void main() {
|
|
${type} outC = getOutputCoords();
|
|
if (any(lessThan(outC, start)) || any(greaterThanEqual(outC, end))) {
|
|
setOutput(value);
|
|
} else {
|
|
${type} coords = outC - start;
|
|
setOutput(getX(${unpackedCoords}));
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class PadPackedProgram {
|
|
constructor(xShape, paddings, constantValue) {
|
|
this.variableNames = ["x"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.customUniforms = [{ name: "value", type: "float" }];
|
|
this.outputShape = paddings.map(
|
|
(p2, i) => p2[0] + xShape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
const rank = xShape.length;
|
|
const dtype = getCoordsDataType(rank);
|
|
const start = paddings.map((p2) => p2[0]).join(",");
|
|
const end = paddings.map((p2, i) => p2[0] + xShape[i]).join(",");
|
|
const coords2 = getChannels("rc", rank);
|
|
const source = getChannels("source", rank);
|
|
const cLimit = `${coords2[rank - 1]} < ${this.outputShape[rank - 1]}`;
|
|
const innerDims = rank === 1 ? "source" : `vec2(${source.slice(-2).join()})`;
|
|
const componentSetup = [
|
|
`${dtype} rc = outputLoc;`,
|
|
`${coords2[rank - 1]} += 1;
|
|
if(${cLimit}) {
|
|
`,
|
|
rank === 1 ? "" : `}
|
|
rc = outputLoc;
|
|
${coords2[rank - 2]} += 1;
|
|
if(${coords2[rank - 2]} < ${this.outputShape[rank - 2]}) {`,
|
|
rank === 1 ? "" : ` ${coords2[rank - 1]} += 1;
|
|
if(${cLimit}) {`
|
|
];
|
|
const paddingArea = rank === 1 ? "rc < start || rc >= end" : "any(lessThan(rc, start)) || any(greaterThanEqual(rc, end))";
|
|
let mainLoop = "";
|
|
for (let i = 0, j2 = rank === 1 ? 2 : 4; i < j2; i++) {
|
|
mainLoop += `
|
|
${componentSetup[i]}
|
|
if (${paddingArea}) {
|
|
result[${i}] = float(value);
|
|
} else {
|
|
${dtype} source = rc - start;
|
|
result[${i}] = getChannel(getX(${source.join()}), ${innerDims});
|
|
}
|
|
`;
|
|
}
|
|
mainLoop += rank === 1 ? `} ` : `}}`;
|
|
this.userCode = `
|
|
const ${dtype} start = ${dtype}(${start});
|
|
const ${dtype} end = ${dtype}(${end});
|
|
|
|
void main() {
|
|
${dtype} outputLoc = getOutputCoords();
|
|
vec4 result = vec4(0.);
|
|
${mainLoop}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const padV2$1 = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { paddings, constantValue } = attrs;
|
|
if (sizeFromShape(x.shape) === 0) {
|
|
const outputShape = paddings.map(
|
|
(p2, i) => p2[0] + x.shape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
return fill$1({
|
|
backend: backend2,
|
|
attrs: { shape: outputShape, value: constantValue, dtype: x.dtype }
|
|
});
|
|
}
|
|
const program = env().getBool("WEBGL_PACK_ARRAY_OPERATIONS") ? new PadPackedProgram(x.shape, paddings, constantValue) : new PadProgram(x.shape, paddings, constantValue);
|
|
const customValues = [[constantValue]];
|
|
return backend2.runWebGLProgram(program, [x], x.dtype, customValues);
|
|
};
|
|
const padV2Config$1 = {
|
|
kernelName: PadV2,
|
|
backendName: "webgl",
|
|
kernelFunc: padV2$1
|
|
};
|
|
const POW = `
|
|
if(a < 0.0 && floor(b) < b){
|
|
return NAN;
|
|
}
|
|
if (b == 0.0) {
|
|
return 1.0;
|
|
}
|
|
return (round(mod(b, 2.0)) != 1) ?
|
|
pow(abs(a), b) : sign(a) * pow(abs(a), b);
|
|
`;
|
|
const POW_PACKED = `
|
|
// isModRound1 has 1 for components with round(mod(b, 2.0)) == 1, 0 otherwise.
|
|
vec4 isModRound1 = vec4(equal(round(mod(b, 2.0)), ivec4(1)));
|
|
vec4 multiplier = sign(a) * isModRound1 + (vec4(1.0) - isModRound1);
|
|
vec4 result = multiplier * pow(abs(a), b);
|
|
|
|
// Ensure that a^0 = 1, including 0^0 = 1 as this correspond to TF and JS
|
|
bvec4 isExpZero = equal(b, vec4(0.0));
|
|
result.r = isExpZero.r ? 1.0 : result.r;
|
|
result.g = isExpZero.g ? 1.0 : result.g;
|
|
result.b = isExpZero.b ? 1.0 : result.b;
|
|
result.a = isExpZero.a ? 1.0 : result.a;
|
|
|
|
bvec4 isNaN1 = lessThan(a, vec4(0.0));
|
|
bvec4 isNaN2 = lessThan(floor(b), b);
|
|
bvec4 isNaN = bvec4(isNaN1.x && isNaN2.x, isNaN1.y && isNaN2.y, isNaN1.z && isNaN2.z, isNaN1.w && isNaN2.w);
|
|
` + CHECK_NAN_SNIPPET_PACKED + `
|
|
return result;
|
|
`;
|
|
const pow$1 = binaryKernelFunc({ opSnippet: POW, packedOpSnippet: POW_PACKED });
|
|
const powConfig$1 = {
|
|
kernelName: Pow,
|
|
backendName: "webgl",
|
|
kernelFunc: pow$1
|
|
};
|
|
function prod(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
const xRank = x.shape.length;
|
|
const toDispose = [];
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
let permutedX = x;
|
|
if (permutedAxes != null) {
|
|
permutedX = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
toDispose.push(permutedX);
|
|
}
|
|
assertAxesAreInnerMostDims("prod", axes, xRank);
|
|
let res;
|
|
if (backend2.shouldExecuteOnCPU([permutedX])) {
|
|
const xVals = backend2.texData.get(permutedX.dataId).values;
|
|
const { outVals, outShape, outDtype } = prodImplCPU(permutedX.shape, permutedX.dtype, xVals, axes);
|
|
res = backend2.makeTensorInfo(outShape, outDtype, outVals);
|
|
} else {
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(permutedX.shape, axes);
|
|
const inSize = sizeFromShape(reduceShape);
|
|
const a2D = reshape$1({ inputs: { x: permutedX }, backend: backend2, attrs: { shape: [-1, inSize] } });
|
|
const outputDType = sumOutType(x.dtype);
|
|
const reduced = reduce(a2D, outputDType, "prod", backend2);
|
|
res = reshape$1({ inputs: { x: reduced }, backend: backend2, attrs: { shape: outShape } });
|
|
toDispose.push(a2D);
|
|
toDispose.push(reduced);
|
|
}
|
|
if (keepDims) {
|
|
toDispose.push(res);
|
|
const newShape = expandShapeToKeepDim(res.shape, origAxes);
|
|
res = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape: newShape } });
|
|
}
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return res;
|
|
}
|
|
const prodConfig = {
|
|
kernelName: Prod,
|
|
backendName: "webgl",
|
|
kernelFunc: prod
|
|
};
|
|
function raggedGather$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { paramsNestedSplits, paramsDenseValues, indices } = inputs;
|
|
const { outputRaggedRank } = attrs;
|
|
const $paramsNestedSplits = paramsNestedSplits.map((t2) => backend2.readSync(t2.dataId));
|
|
const $paramsNestedSplitsShapes = paramsNestedSplits.map((t2) => t2.shape);
|
|
const $paramsDenseValues = backend2.readSync(paramsDenseValues.dataId);
|
|
const $indices = backend2.readSync(indices.dataId);
|
|
const [outputNestedSplits, outputDenseValues, outputDenseValuesShape] = raggedGatherImplCPU($paramsNestedSplits, $paramsNestedSplitsShapes, $paramsDenseValues, paramsDenseValues.shape, paramsDenseValues.dtype, $indices, indices.shape, outputRaggedRank);
|
|
const outputNestedSplitsTensors = outputNestedSplits.map((splits) => backend2.makeTensorInfo([splits.length], "int32", splits));
|
|
const outputDenseValuesTensor = backend2.makeTensorInfo(outputDenseValuesShape, paramsDenseValues.dtype, outputDenseValues);
|
|
return outputNestedSplitsTensors.concat([outputDenseValuesTensor]);
|
|
}
|
|
const raggedGatherConfig$1 = {
|
|
kernelName: RaggedGather,
|
|
backendName: "webgl",
|
|
kernelFunc: raggedGather$1
|
|
};
|
|
function raggedRange$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { starts, limits, deltas } = inputs;
|
|
const $starts = backend2.readSync(starts.dataId);
|
|
const $limits = backend2.readSync(limits.dataId);
|
|
const $deltas = backend2.readSync(deltas.dataId);
|
|
const [rtNestedSplitsData, rtDenseValuesData] = raggedRangeImplCPU($starts, starts.shape, starts.dtype, $limits, limits.shape, $deltas, deltas.shape);
|
|
const rtNestedSplits = backend2.makeTensorInfo([rtNestedSplitsData.length], "int32", rtNestedSplitsData);
|
|
const rtDenseValues = backend2.makeTensorInfo([rtDenseValuesData.length], starts.dtype, rtDenseValuesData);
|
|
return [rtNestedSplits, rtDenseValues];
|
|
}
|
|
const raggedRangeConfig$1 = {
|
|
kernelName: RaggedRange,
|
|
backendName: "webgl",
|
|
kernelFunc: raggedRange$1
|
|
};
|
|
function raggedTensorToTensor$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { shape, values, defaultValue, rowPartitionTensors } = inputs;
|
|
const { rowPartitionTypes } = attrs;
|
|
const $shape = backend2.readSync(shape.dataId);
|
|
const $values = backend2.readSync(values.dataId);
|
|
const $defaultValue = backend2.readSync(defaultValue.dataId);
|
|
const $rowPartitionValues = rowPartitionTensors.map((t2) => backend2.readSync(t2.dataId));
|
|
const rowPartitionValuesShapes = rowPartitionTensors.map((t2) => t2.shape);
|
|
const [outputShape, output] = raggedTensorToTensorImplCPU($shape, shape.shape, $values, values.shape, values.dtype, $defaultValue, defaultValue.shape, $rowPartitionValues, rowPartitionValuesShapes, rowPartitionTypes);
|
|
return backend2.makeTensorInfo(outputShape, values.dtype, output);
|
|
}
|
|
const raggedTensorToTensorConfig$1 = {
|
|
kernelName: RaggedTensorToTensor,
|
|
backendName: "webgl",
|
|
kernelFunc: raggedTensorToTensor$1
|
|
};
|
|
const range$1 = (args) => {
|
|
const { backend: backend2, attrs } = args;
|
|
const { start, stop, step: step2, dtype } = attrs;
|
|
const values = rangeImplCPU(start, stop, step2, dtype);
|
|
return backend2.makeTensorInfo([values.length], dtype, values);
|
|
};
|
|
const rangeConfig$1 = {
|
|
kernelName: Range,
|
|
backendName: "webgl",
|
|
kernelFunc: range$1
|
|
};
|
|
const RECIPROCAL = `return 1.0 / x;`;
|
|
const reciprocal$1 = unaryKernelFunc({ opSnippet: RECIPROCAL });
|
|
const reciprocalConfig$1 = {
|
|
kernelName: Reciprocal,
|
|
backendName: "webgl",
|
|
kernelFunc: reciprocal$1
|
|
};
|
|
const RELU = CHECK_NAN_SNIPPET$1 + `
|
|
return (x < 0.0) ? 0.0 : x;
|
|
`;
|
|
const RELU_PACKED = `
|
|
vec4 result = x * vec4(greaterThanEqual(x, vec4(0.0)));
|
|
bvec4 isNaN = isnan(x);
|
|
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
const relu$1 = unaryKernelFunc({ opSnippet: RELU, packedOpSnippet: RELU_PACKED });
|
|
const reluConfig$1 = {
|
|
kernelName: Relu,
|
|
backendName: "webgl",
|
|
kernelFunc: relu$1
|
|
};
|
|
const RELU6 = CHECK_NAN_SNIPPET$1 + `
|
|
return (x < 0.0) ? 0.0 : min(6.0, x);
|
|
`;
|
|
const RELU6_PACKED = `
|
|
vec4 result = min(x, vec4(6.)) * vec4(greaterThanEqual(x, vec4(0.0)));
|
|
bvec4 isNaN = isnan(x);
|
|
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
const relu6$1 = unaryKernelFunc({ opSnippet: RELU6, packedOpSnippet: RELU6_PACKED });
|
|
const relu6Config$1 = {
|
|
kernelName: Relu6,
|
|
backendName: "webgl",
|
|
kernelFunc: relu6$1
|
|
};
|
|
class ResizeBilinearProgram {
|
|
constructor(inputShape, newHeight, newWidth, alignCorners, halfPixelCenters) {
|
|
this.variableNames = ["A"];
|
|
this.outputShape = [];
|
|
const [batch, oldHeight, oldWidth, depth] = inputShape;
|
|
this.outputShape = [batch, newHeight, newWidth, depth];
|
|
const effectiveInSize = [
|
|
alignCorners && newHeight > 1 ? oldHeight - 1 : oldHeight,
|
|
alignCorners && newWidth > 1 ? oldWidth - 1 : oldWidth
|
|
];
|
|
const effectiveOutSize = [
|
|
alignCorners && newHeight > 1 ? newHeight - 1 : newHeight,
|
|
alignCorners && newWidth > 1 ? newWidth - 1 : newWidth
|
|
];
|
|
let sourceFracIndexRC;
|
|
if (halfPixelCenters) {
|
|
sourceFracIndexRC = `(vec2(yRC) + vec2(0.5)) * effectiveInputOverOutputRatioRC - vec2(0.5)`;
|
|
} else {
|
|
sourceFracIndexRC = `vec2(yRC) * effectiveInputOverOutputRatioRC`;
|
|
}
|
|
this.userCode = `
|
|
const vec2 effectiveInputOverOutputRatioRC = vec2(
|
|
${effectiveInSize[0] / effectiveOutSize[0]},
|
|
${effectiveInSize[1] / effectiveOutSize[1]});
|
|
const vec2 inputShapeRC = vec2(${oldHeight}.0, ${oldWidth}.0);
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
ivec2 yRC = coords.yz;
|
|
|
|
// Fractional source index.
|
|
vec2 sourceFracIndexRC = ${sourceFracIndexRC};
|
|
|
|
// Compute the four integer indices.
|
|
ivec2 sourceFloorRC = ivec2(max(sourceFracIndexRC, vec2(0.0)));
|
|
ivec2 sourceCeilRC = ivec2(
|
|
min(inputShapeRC - 1.0, ceil(sourceFracIndexRC)));
|
|
|
|
float topLeft = getA(b, sourceFloorRC.x, sourceFloorRC.y, d);
|
|
float bottomLeft = getA(b, sourceCeilRC.x, sourceFloorRC.y, d);
|
|
float topRight = getA(b, sourceFloorRC.x, sourceCeilRC.y, d);
|
|
float bottomRight = getA(b, sourceCeilRC.x, sourceCeilRC.y, d);
|
|
|
|
vec2 fracRC = sourceFracIndexRC - vec2(sourceFloorRC);
|
|
|
|
float top = topLeft + (topRight - topLeft) * fracRC.y;
|
|
float bottom = bottomLeft + (bottomRight - bottomLeft) * fracRC.y;
|
|
float newValue = top + (bottom - top) * fracRC.x;
|
|
|
|
setOutput(newValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ResizeBilinearPackedProgram {
|
|
constructor(inputShape, newHeight, newWidth, alignCorners, halfPixelCenters) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = [];
|
|
const [batch, oldHeight, oldWidth, depth] = inputShape;
|
|
this.outputShape = [batch, newHeight, newWidth, depth];
|
|
const effectiveInSize = [
|
|
alignCorners && newHeight > 1 ? oldHeight - 1 : oldHeight,
|
|
alignCorners && newWidth > 1 ? oldWidth - 1 : oldWidth
|
|
];
|
|
const effectiveOutSize = [
|
|
alignCorners && newHeight > 1 ? newHeight - 1 : newHeight,
|
|
alignCorners && newWidth > 1 ? newWidth - 1 : newWidth
|
|
];
|
|
let sourceFracIndexRC;
|
|
if (halfPixelCenters) {
|
|
sourceFracIndexRC = `(vec3(yRC) + vec3(0.5)) * effectiveInputOverOutputRatioRC - vec3(0.5)`;
|
|
} else {
|
|
sourceFracIndexRC = `vec3(yRC) * effectiveInputOverOutputRatioRC`;
|
|
}
|
|
this.userCode = `
|
|
const vec3 effectiveInputOverOutputRatioRC = vec3(
|
|
${effectiveInSize[0] / effectiveOutSize[0]},
|
|
${effectiveInSize[1] / effectiveOutSize[1]},
|
|
${effectiveInSize[1] / effectiveOutSize[1]});
|
|
const vec3 inputShapeRC = vec3(${oldHeight}.0, ${oldWidth}.0,
|
|
${oldWidth}.0);
|
|
|
|
float getAValue(int b, int r, int c, int d) {
|
|
return getChannel(getA(b, r, c, d), vec2(c, d));
|
|
}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
// Calculate values for next column in yRC.z.
|
|
ivec3 yRC = coords.yzz + ivec3(0, 0, 1);
|
|
|
|
// Fractional source index.
|
|
vec3 sourceFracIndexRC = ${sourceFracIndexRC};
|
|
|
|
// Compute the four integer indices.
|
|
ivec3 sourceFloorRC = ivec3(max(sourceFracIndexRC, vec3(0.0)));
|
|
ivec3 sourceCeilRC = ivec3(
|
|
min(inputShapeRC - 1.0, ceil(sourceFracIndexRC)));
|
|
|
|
// Should we calculate next column and row elements in 2x2 packed cell.
|
|
bool hasNextCol = d < ${depth - 1};
|
|
bool hasNextRow = coords.z < ${newWidth - 1};
|
|
|
|
// In parallel, construct four corners for all four components in
|
|
// packed 2x2 cell.
|
|
vec4 topLeft = vec4(
|
|
getAValue(b, sourceFloorRC.x, sourceFloorRC.y, d),
|
|
hasNextCol ? getAValue(b, sourceFloorRC.x, sourceFloorRC.y, d + 1)
|
|
: 0.0,
|
|
hasNextRow ? getAValue(b, sourceFloorRC.x, sourceFloorRC.z, d)
|
|
: 0.0,
|
|
(hasNextRow && hasNextCol) ?
|
|
getAValue(b, sourceFloorRC.x, sourceFloorRC.z, d + 1) : 0.0);
|
|
|
|
vec4 bottomLeft = vec4(
|
|
getAValue(b, sourceCeilRC.x, sourceFloorRC.y, d),
|
|
hasNextCol ? getAValue(b, sourceCeilRC.x, sourceFloorRC.y, d + 1)
|
|
: 0.0,
|
|
hasNextRow ? getAValue(b, sourceCeilRC.x, sourceFloorRC.z, d)
|
|
: 0.0,
|
|
(hasNextRow && hasNextCol) ?
|
|
getAValue(b, sourceCeilRC.x, sourceFloorRC.z, d + 1) : 0.0);
|
|
|
|
vec4 topRight = vec4(
|
|
getAValue(b, sourceFloorRC.x, sourceCeilRC.y, d),
|
|
hasNextCol ? getAValue(b, sourceFloorRC.x, sourceCeilRC.y, d + 1)
|
|
: 0.0,
|
|
hasNextRow ? getAValue(b, sourceFloorRC.x, sourceCeilRC.z, d)
|
|
: 0.0,
|
|
(hasNextRow && hasNextCol) ?
|
|
getAValue(b, sourceFloorRC.x, sourceCeilRC.z, d + 1) : 0.0);
|
|
|
|
vec4 bottomRight = vec4(
|
|
getAValue(b, sourceCeilRC.x, sourceCeilRC.y, d),
|
|
hasNextCol ? getAValue(b, sourceCeilRC.x, sourceCeilRC.y, d + 1)
|
|
: 0.0,
|
|
hasNextRow ? getAValue(b, sourceCeilRC.x, sourceCeilRC.z, d)
|
|
: 0.0,
|
|
(hasNextRow && hasNextCol) ?
|
|
getAValue(b, sourceCeilRC.x, sourceCeilRC.z, d + 1) : 0.0);
|
|
|
|
vec3 fracRC = sourceFracIndexRC - vec3(sourceFloorRC);
|
|
|
|
vec4 top = mix(topLeft, topRight, fracRC.yyzz);
|
|
vec4 bottom = mix(bottomLeft, bottomRight, fracRC.yyzz);
|
|
vec4 newValue = mix(top, bottom, fracRC.x);
|
|
|
|
setOutput(newValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function resizeBilinear$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images } = inputs;
|
|
const { alignCorners, halfPixelCenters, size } = attrs;
|
|
const [newHeight, newWidth] = size;
|
|
const program = env().getBool("WEBGL_PACK_IMAGE_OPERATIONS") ? new ResizeBilinearPackedProgram(images.shape, newHeight, newWidth, alignCorners, halfPixelCenters) : new ResizeBilinearProgram(images.shape, newHeight, newWidth, alignCorners, halfPixelCenters);
|
|
return backend2.runWebGLProgram(program, [images], "float32");
|
|
}
|
|
const resizeBilinearConfig$1 = {
|
|
kernelName: ResizeBilinear,
|
|
backendName: "webgl",
|
|
kernelFunc: resizeBilinear$1
|
|
};
|
|
class ResizeBilinearBackpropProgram {
|
|
constructor(dyShape, inputShape, alignCorners) {
|
|
this.variableNames = ["dy"];
|
|
this.outputShape = [];
|
|
this.outputShape = inputShape;
|
|
const [, xHeight, xWidth] = inputShape;
|
|
const [, yHeight, yWidth] = dyShape;
|
|
const effectiveXSize = [
|
|
alignCorners && yHeight > 1 ? xHeight - 1 : xHeight,
|
|
alignCorners && yWidth > 1 ? xWidth - 1 : xWidth
|
|
];
|
|
const effectiveYSize = [
|
|
alignCorners && yHeight > 1 ? yHeight - 1 : yHeight,
|
|
alignCorners && yWidth > 1 ? yWidth - 1 : yWidth
|
|
];
|
|
const heightScale = effectiveXSize[0] / effectiveYSize[0];
|
|
const widthScale = effectiveXSize[1] / effectiveYSize[1];
|
|
const invHeightScale = 1 / heightScale;
|
|
const invWidthScale = 1 / widthScale;
|
|
const winHeight = Math.ceil(invHeightScale) * 2 + 2;
|
|
const winWidth = Math.ceil(invWidthScale) * 2 + 2;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
int r = coords[1];
|
|
int c = coords[2];
|
|
|
|
float accumulator = 0.0;
|
|
|
|
const float heightScale = float(${heightScale});
|
|
const float widthScale = float(${widthScale});
|
|
|
|
const float invHeightScale = float(${invHeightScale});
|
|
const float invWidthScale = float(${invWidthScale});
|
|
|
|
const int winHeight = int(${winHeight});
|
|
const int winWidth = int(${winWidth});
|
|
|
|
// Compute bounds for where in dy we will look
|
|
float startRLerp = floor(float(r) * invHeightScale);
|
|
int startDyR = int(startRLerp - float(winHeight / 2));
|
|
|
|
float startCLerp = floor(float(c) * invWidthScale);
|
|
int startDyC = int(startCLerp - float(winWidth / 2));
|
|
|
|
// Loop over dy
|
|
for (int dyROffset = 0; dyROffset < winHeight; dyROffset++) {
|
|
int dyR = dyROffset + startDyR;
|
|
|
|
// Guard against the window exceeding the bounds of dy
|
|
if (dyR < 0 || dyR >= ${yHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int dyCOffset = 0; dyCOffset < winWidth; dyCOffset++) {
|
|
int dyC = dyCOffset + startDyC;
|
|
|
|
// Guard against the window exceeding the bounds of dy
|
|
if (dyC < 0 || dyC >= ${yWidth}) {
|
|
continue;
|
|
}
|
|
|
|
float dxR = float(dyR) * heightScale;
|
|
int topDxRIndex = int(floor(dxR));
|
|
int bottomDxRIndex = int(min(ceil(dxR), ${xHeight - 1}.0));
|
|
float dxRLerp = dxR - float(topDxRIndex);
|
|
float inverseDxRLerp = 1.0 - dxRLerp;
|
|
|
|
float dxC = float(dyC) * widthScale;
|
|
int leftDxCIndex = int(floor(dxC));
|
|
int rightDxCIndex = int(min(ceil(dxC), ${xWidth - 1}.0));
|
|
float dxCLerp = dxC - float(leftDxCIndex);
|
|
float inverseDxCLerp = 1.0 - dxCLerp;
|
|
|
|
if (r == topDxRIndex && c == leftDxCIndex) {
|
|
// topLeft
|
|
accumulator +=
|
|
getDy(b, dyR, dyC, d) * inverseDxRLerp * inverseDxCLerp;
|
|
}
|
|
|
|
if (r == topDxRIndex && c == rightDxCIndex) {
|
|
// topRight
|
|
accumulator += getDy(b, dyR, dyC, d) * inverseDxRLerp * dxCLerp;
|
|
}
|
|
|
|
if (r == bottomDxRIndex && c == leftDxCIndex) {
|
|
// bottomLeft
|
|
accumulator += getDy(b, dyR, dyC, d) * dxRLerp * inverseDxCLerp;
|
|
}
|
|
|
|
if (r == bottomDxRIndex && c == rightDxCIndex) {
|
|
// bottomRight
|
|
accumulator += getDy(b, dyR, dyC, d) * dxRLerp * dxCLerp;
|
|
}
|
|
}
|
|
}
|
|
// End loop over dy
|
|
|
|
setOutput(accumulator);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function resizeBilinearGrad$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images, dy } = inputs;
|
|
const { alignCorners } = attrs;
|
|
const program = new ResizeBilinearBackpropProgram(dy.shape, images.shape, alignCorners);
|
|
return backend2.runWebGLProgram(program, [dy], dy.dtype);
|
|
}
|
|
const resizeBilinearGradConfig$1 = {
|
|
kernelName: ResizeBilinearGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: resizeBilinearGrad$1
|
|
};
|
|
class ResizeNearestNeighborProgram {
|
|
constructor(inputShape, newHeight, newWidth, alignCorners, halfPixelCenters) {
|
|
this.variableNames = ["A"];
|
|
this.outputShape = [];
|
|
const [batch, oldHeight, oldWidth, depth] = inputShape;
|
|
this.outputShape = [batch, newHeight, newWidth, depth];
|
|
const effectiveInSize = [
|
|
alignCorners && newHeight > 1 ? oldHeight - 1 : oldHeight,
|
|
alignCorners && newWidth > 1 ? oldWidth - 1 : oldWidth
|
|
];
|
|
const effectiveOutSize = [
|
|
alignCorners && newHeight > 1 ? newHeight - 1 : newHeight,
|
|
alignCorners && newWidth > 1 ? newWidth - 1 : newWidth
|
|
];
|
|
const roundBase = alignCorners ? "0.5" : "0.0";
|
|
let sourceFracIndexRC;
|
|
if (halfPixelCenters) {
|
|
sourceFracIndexRC = `max((vec2(yRC) + vec2(0.5)) * effectiveInputOverOutputRatioRC, vec2(0.0))`;
|
|
} else {
|
|
sourceFracIndexRC = `vec2(yRC) * effectiveInputOverOutputRatioRC`;
|
|
}
|
|
this.userCode = `
|
|
const vec2 effectiveInputOverOutputRatioRC = vec2(
|
|
${effectiveInSize[0] / effectiveOutSize[0]},
|
|
${effectiveInSize[1] / effectiveOutSize[1]});
|
|
const vec2 inputShapeRC = vec2(${oldHeight}.0, ${oldWidth}.0);
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
ivec2 yRC = coords.yz;
|
|
|
|
// Fractional source index.
|
|
vec2 sourceFracIndexRC = ${sourceFracIndexRC};
|
|
|
|
// Compute the coordinators of nearest neighbor point.
|
|
ivec2 sourceNearestRC = ivec2(
|
|
min(inputShapeRC - 1.0, floor(sourceFracIndexRC + ${roundBase})));
|
|
float newValue = getA(b, sourceNearestRC.x, sourceNearestRC.y, d);
|
|
|
|
setOutput(newValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ResizeNearestNeighborPackedProgram {
|
|
constructor(inputShape, newHeight, newWidth, alignCorners, halfPixelCenters) {
|
|
this.variableNames = ["A"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = [];
|
|
const [batch, oldHeight, oldWidth, depth] = inputShape;
|
|
this.outputShape = [batch, newHeight, newWidth, depth];
|
|
const effectiveInSize = [
|
|
alignCorners && newHeight > 1 ? oldHeight - 1 : oldHeight,
|
|
alignCorners && newWidth > 1 ? oldWidth - 1 : oldWidth
|
|
];
|
|
const effectiveOutSize = [
|
|
alignCorners && newHeight > 1 ? newHeight - 1 : newHeight,
|
|
alignCorners && newWidth > 1 ? newWidth - 1 : newWidth
|
|
];
|
|
const roundBase = alignCorners ? "0.5" : "0.0";
|
|
let sourceFracIndexRC;
|
|
if (halfPixelCenters) {
|
|
sourceFracIndexRC = `max((vec3(yRC) + vec3(0.5)) * effectiveInputOverOutputRatioRC, vec3(0.0))`;
|
|
} else {
|
|
sourceFracIndexRC = `vec3(yRC) * effectiveInputOverOutputRatioRC`;
|
|
}
|
|
this.userCode = `
|
|
const vec3 effectiveInputOverOutputRatioRC = vec3(
|
|
${effectiveInSize[0] / effectiveOutSize[0]},
|
|
${effectiveInSize[1] / effectiveOutSize[1]},
|
|
${effectiveInSize[1] / effectiveOutSize[1]});
|
|
const vec3 inputShapeRC = vec3(${oldHeight}.0, ${oldWidth}.0,
|
|
${oldWidth}.0);
|
|
|
|
float getAValue(int b, int r, int c, int d) {
|
|
return getChannel(getA(b, r, c, d), vec2(c, d));
|
|
}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
// Calculate values for next column in yRC.z.
|
|
ivec3 yRC = coords.yzz + ivec3(0, 0, 1);
|
|
|
|
// Fractional source index.
|
|
vec3 sourceFracIndexRC = ${sourceFracIndexRC};
|
|
|
|
// Compute the coordinators of nearest neighbor point.
|
|
ivec3 sourceNearestRC = ivec3(
|
|
min(inputShapeRC - 1.0, floor(sourceFracIndexRC + ${roundBase})));
|
|
|
|
// Should we calculate next column and row elements in 2x2 packed cell.
|
|
bool hasNextCol = d < ${depth - 1};
|
|
bool hasNextRow = coords.z < ${newWidth - 1};
|
|
|
|
vec4 newValue = vec4(
|
|
getAValue(b, sourceNearestRC.x, sourceNearestRC.y, d),
|
|
hasNextCol ? getAValue(b, sourceNearestRC.x, sourceNearestRC.y, d + 1)
|
|
: 0.0,
|
|
hasNextRow ? getAValue(b, sourceNearestRC.x, sourceNearestRC.z, d)
|
|
: 0.0,
|
|
(hasNextRow && hasNextCol) ?
|
|
getAValue(b, sourceNearestRC.x, sourceNearestRC.z, d + 1) : 0.0);
|
|
|
|
setOutput(newValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function resizeNearestNeighbor$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images } = inputs;
|
|
const { alignCorners, halfPixelCenters, size } = attrs;
|
|
const [newHeight, newWidth] = size;
|
|
const program = env().getBool("WEBGL_PACK_IMAGE_OPERATIONS") ? new ResizeNearestNeighborPackedProgram(images.shape, newHeight, newWidth, alignCorners, halfPixelCenters) : new ResizeNearestNeighborProgram(images.shape, newHeight, newWidth, alignCorners, halfPixelCenters);
|
|
return backend2.runWebGLProgram(program, [images], images.dtype);
|
|
}
|
|
const resizeNearestNeighborConfig$1 = {
|
|
kernelName: ResizeNearestNeighbor,
|
|
backendName: "webgl",
|
|
kernelFunc: resizeNearestNeighbor$1
|
|
};
|
|
class ResizeNearestNeigborBackpropProgram {
|
|
constructor(dyShape, inputShape, alignCorners) {
|
|
this.variableNames = ["dy"];
|
|
this.outputShape = [];
|
|
this.outputShape = inputShape;
|
|
const [, xHeight, xWidth] = inputShape;
|
|
const [, yHeight, yWidth] = dyShape;
|
|
const effectiveXSize = [
|
|
alignCorners && yHeight > 1 ? xHeight - 1 : xHeight,
|
|
alignCorners && yWidth > 1 ? xWidth - 1 : xWidth
|
|
];
|
|
const effectiveYSize = [
|
|
alignCorners && yHeight > 1 ? yHeight - 1 : yHeight,
|
|
alignCorners && yWidth > 1 ? yWidth - 1 : yWidth
|
|
];
|
|
const heightScale = effectiveXSize[0] / effectiveYSize[0];
|
|
const widthScale = effectiveXSize[1] / effectiveYSize[1];
|
|
const invHeightScale = 1 / heightScale;
|
|
const invWidthScale = 1 / widthScale;
|
|
const winHeight = Math.ceil(invHeightScale) * 2 + 2;
|
|
const winWidth = Math.ceil(invWidthScale) * 2 + 2;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int b = coords[0];
|
|
int d = coords[3];
|
|
int r = coords[1];
|
|
int c = coords[2];
|
|
|
|
float accumulator = 0.0;
|
|
|
|
const float heightScale = float(${heightScale});
|
|
const float widthScale = float(${widthScale});
|
|
|
|
const float invHeightScale = float(${invHeightScale});
|
|
const float invWidthScale = float(${invWidthScale});
|
|
|
|
const int winHeight = int(${winHeight});
|
|
const int winWidth = int(${winWidth});
|
|
|
|
// Compute bounds for where in dy we will look
|
|
float startRLerp = floor(float(r) * invHeightScale);
|
|
int startDyR = int(floor(startRLerp - float(winHeight / 2)));
|
|
|
|
float startCLerp = floor(float(c) * invWidthScale);
|
|
int startDyC = int(floor(startCLerp - float(winWidth / 2)));
|
|
|
|
// Loop over dy
|
|
for (int dyROffset = 0; dyROffset < winHeight; dyROffset++) {
|
|
int dyR = dyROffset + startDyR;
|
|
|
|
// Guard against the window exceeding the bounds of dy
|
|
if (dyR < 0 || dyR >= ${yHeight}) {
|
|
continue;
|
|
}
|
|
|
|
for (int dyCOffset = 0; dyCOffset < winWidth; dyCOffset++) {
|
|
int dyC = dyCOffset + startDyC;
|
|
|
|
// Guard against the window exceeding the bounds of dy
|
|
if (dyC < 0 || dyC >= ${yWidth}) {
|
|
continue;
|
|
}
|
|
|
|
float sourceFracRow =
|
|
float(${effectiveXSize[0]}) *
|
|
(float(dyR) / float(${effectiveYSize[0]}));
|
|
|
|
float sourceFracCol =
|
|
float(${effectiveXSize[1]}) *
|
|
(float(dyC) / float(${effectiveYSize[1]}));
|
|
|
|
int sourceNearestRow = int(min(
|
|
float(int(${xHeight}) - 1),
|
|
${alignCorners} ? float(round(sourceFracRow)) :
|
|
float(floor(sourceFracRow))));
|
|
|
|
int sourceNearestCol = int(min(
|
|
float(int(${xWidth}) - 1),
|
|
${alignCorners} ? float(round(sourceFracCol)) :
|
|
float(floor(sourceFracCol))));
|
|
|
|
if (r == sourceNearestRow && c == sourceNearestCol) {
|
|
accumulator += getDy(b, dyR, dyC, d);
|
|
}
|
|
}
|
|
}
|
|
// End loop over dy
|
|
|
|
setOutput(accumulator);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function resizeNearestNeighborGrad$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images, dy } = inputs;
|
|
const { alignCorners } = attrs;
|
|
const program = new ResizeNearestNeigborBackpropProgram(dy.shape, images.shape, alignCorners);
|
|
return backend2.runWebGLProgram(program, [dy], dy.dtype);
|
|
}
|
|
const resizeNearestNeighborGradConfig$1 = {
|
|
kernelName: ResizeNearestNeighborGrad,
|
|
backendName: "webgl",
|
|
kernelFunc: resizeNearestNeighborGrad$1
|
|
};
|
|
class ReverseProgram {
|
|
constructor(xShape, axis) {
|
|
this.variableNames = ["x"];
|
|
const rank = xShape.length;
|
|
if (rank > 4) {
|
|
throw new Error(`WebGL backend: Reverse of rank-${rank} tensor is not yet supported`);
|
|
}
|
|
this.outputShape = xShape;
|
|
if (rank === 1) {
|
|
this.userCode = `
|
|
void main() {
|
|
int coord = getOutputCoords();
|
|
setOutput(getX(${xShape[0]} - coord - 1));
|
|
}
|
|
`;
|
|
return;
|
|
}
|
|
const getInCoord = (i) => {
|
|
if (axis.indexOf(i) !== -1 && xShape[i] !== 1) {
|
|
return `${xShape[i]} - coords[${i}] - 1`;
|
|
}
|
|
return `coords[${i}]`;
|
|
};
|
|
const inCoords = xShape.map((_, i) => getInCoord(i)).join(",");
|
|
const type = getCoordsDataType(rank);
|
|
this.userCode = `
|
|
void main() {
|
|
${type} coords = getOutputCoords();
|
|
setOutput(getX(${inCoords}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ReversePackedProgram {
|
|
constructor(xShape, axis) {
|
|
this.variableNames = ["x"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
const rank = xShape.length;
|
|
if (rank > 4) {
|
|
throw new Error(`WebGL backend: Reverse of rank-${rank} tensor is not yet supported`);
|
|
}
|
|
this.outputShape = xShape;
|
|
const channels = getChannels("rc", rank);
|
|
const nextColumn = `${channels[rank - 1]} + 1 < ${this.outputShape[rank - 1]}`;
|
|
const nextRow = `${channels[rank - 2]} + 1 < ${this.outputShape[rank - 2]}`;
|
|
const type = getCoordsDataType(rank);
|
|
if (rank === 1) {
|
|
this.userCode = `
|
|
void main(){
|
|
int rc = getOutputCoords();
|
|
vec4 result = vec4(0.);
|
|
result.r = getChannel(getX(${xShape[0]} - rc - 1),
|
|
${xShape[0]} - rc - 1);
|
|
if(${nextColumn}){
|
|
result.g = getChannel(getX(${xShape[0]} - (rc + 1) - 1),
|
|
${xShape[0]} - (rc + 1) - 1);
|
|
}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
} else {
|
|
this.userCode = `
|
|
void main() {
|
|
${type} rc = getOutputCoords();
|
|
vec4 result = vec4(0.);
|
|
result.r = ${getR(channels.slice())};
|
|
if(${nextColumn}){
|
|
result.g = ${getG(channels.slice())};
|
|
}
|
|
if(${nextRow}) {
|
|
result.b = ${getB(channels.slice())};
|
|
if(${nextColumn}) {
|
|
result.a = ${getA(channels.slice())};
|
|
}
|
|
}
|
|
setOutput(result);
|
|
}
|
|
`;
|
|
}
|
|
function getR(channels2) {
|
|
return getChannel(channels2);
|
|
}
|
|
function getG(channels2) {
|
|
channels2[rank - 1] = "(" + channels2[rank - 1] + ` + 1)`;
|
|
return getChannel(channels2);
|
|
}
|
|
function getB(channels2) {
|
|
channels2[rank - 2] = "(" + channels2[rank - 2] + ` + 1)`;
|
|
return getChannel(channels2);
|
|
}
|
|
function getA(channels2) {
|
|
channels2[rank - 1] = "(" + channels2[rank - 1] + ` + 1)`;
|
|
channels2[rank - 2] = "(" + channels2[rank - 2] + ` + 1)`;
|
|
return getChannel(channels2);
|
|
}
|
|
function getChannel(channels2) {
|
|
const inCoordsArray = xShape.map((_, i) => getInCoord(i, channels2));
|
|
const inCoords = inCoordsArray.join(",");
|
|
const innerDims = inCoordsArray.slice(-2).join(",");
|
|
return `getChannel(getX(${inCoords}), vec2(${innerDims}))`;
|
|
}
|
|
function getInCoord(i, channels1) {
|
|
if (axis.indexOf(i) !== -1 && xShape[i] !== 1) {
|
|
return `${xShape[i]} - ${channels1[i]} - 1`;
|
|
} else {
|
|
return `${channels1[i]}`;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
function reverse$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { dims } = attrs;
|
|
const xRank = x.shape.length;
|
|
const $dims = parseAxisParam(dims, x.shape);
|
|
if (xRank === 0) {
|
|
return identity({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const program = env().getBool("WEBGL_PACK_ARRAY_OPERATIONS") ? new ReversePackedProgram(x.shape, $dims) : new ReverseProgram(x.shape, $dims);
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
const reverseConfig$1 = {
|
|
kernelName: Reverse,
|
|
backendName: "webgl",
|
|
kernelFunc: reverse$1
|
|
};
|
|
class RotateProgram {
|
|
constructor(imageShape, fillValue) {
|
|
this.variableNames = ["Image"];
|
|
this.outputShape = [];
|
|
this.customUniforms = [{ name: "params", type: "vec4" }];
|
|
const imageHeight = imageShape[1];
|
|
const imageWidth = imageShape[2];
|
|
this.outputShape = imageShape;
|
|
let fillSnippet = "";
|
|
if (typeof fillValue === "number") {
|
|
fillSnippet = `float outputValue = ${fillValue.toFixed(2)};`;
|
|
} else {
|
|
fillSnippet = `
|
|
vec3 fill = vec3(${fillValue.join(",")});
|
|
float outputValue = fill[coords[3]];`;
|
|
}
|
|
this.userCode = `
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
int x = coords[2];
|
|
int y = coords[1];
|
|
float coordXFloat = (float(x) - params[0]) * params[3] -
|
|
(float(y) - params[1]) * params[2];
|
|
float coordYFloat = (float(x) - params[0]) * params[2] +
|
|
(float(y) - params[1]) * params[3];
|
|
int coordX = int(round(coordXFloat + params[0]));
|
|
int coordY = int(round(coordYFloat + params[1]));
|
|
${fillSnippet}
|
|
if(coordX >= 0 && coordX < ${imageWidth} && coordY >= 0 && coordY < ${imageHeight}) {
|
|
outputValue = getImage(coords[0], coordY, coordX, coords[3]);
|
|
}
|
|
setOutput(outputValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
const rotateWithOffsetConfig$1 = {
|
|
kernelName: RotateWithOffset,
|
|
backendName: "webgl",
|
|
kernelFunc: ({ inputs, attrs, backend: backend2 }) => {
|
|
const { image: image2 } = inputs;
|
|
const { radians, fillValue, center } = attrs;
|
|
const webglBackend = backend2;
|
|
const program = new RotateProgram(image2.shape, fillValue);
|
|
const [centerX, centerY] = getImageCenter(center, image2.shape[1], image2.shape[2]);
|
|
const customValues = [[centerX, centerY, Math.sin(radians), Math.cos(radians)]];
|
|
const output = webglBackend.runWebGLProgram(program, [image2], image2.dtype, customValues);
|
|
return output;
|
|
}
|
|
};
|
|
const ROUND = `
|
|
// OpenGL ES does not support round function.
|
|
// The algorithm is based on banker's rounding.
|
|
float base = floor(x);
|
|
if ((x - base) < 0.5) {
|
|
return floor(x);
|
|
} else if ((x - base) > 0.5) {
|
|
return ceil(x);
|
|
} else {
|
|
if (mod(base, 2.0) == 0.0) {
|
|
return base;
|
|
} else {
|
|
return base + 1.0;
|
|
}
|
|
}
|
|
`;
|
|
const round$1 = unaryKernelFunc({ opSnippet: ROUND });
|
|
const roundConfig$1 = {
|
|
kernelName: Round,
|
|
backendName: "webgl",
|
|
kernelFunc: round$1
|
|
};
|
|
const RSQRT = `return inversesqrt(x);`;
|
|
const rsqrt = unaryKernelFunc({ opSnippet: RSQRT, cpuKernelImpl: rsqrtImplCPU });
|
|
const rsqrtConfig = {
|
|
kernelName: Rsqrt,
|
|
backendName: "webgl",
|
|
kernelFunc: rsqrt
|
|
};
|
|
class ScatterProgram {
|
|
constructor(updateSize, sliceDim, indicesRank, updatesRank, strides, shape, summingDupeIndex = true, defaultIsTensor = false) {
|
|
this.variableNames = ["updates", "indices", "defaultValue"];
|
|
this.outputShape = shape;
|
|
const stridesType = getCoordsDataType(strides.length);
|
|
const dtype = getCoordsDataType(shape.length);
|
|
let indicesString = "";
|
|
if (indicesRank === 1) {
|
|
indicesString = "i";
|
|
} else if (indicesRank === 2) {
|
|
indicesString = "i, j";
|
|
}
|
|
const indicesSnippet = `getIndices(${indicesString})`;
|
|
let updatesString = "";
|
|
if (updatesRank === 1) {
|
|
updatesString = "i";
|
|
} else if (updatesRank === 2) {
|
|
updatesString = "i, coords[1]";
|
|
}
|
|
const updatesSnippet = `getUpdates(${updatesString})`;
|
|
let defaultValuesString = "";
|
|
if (defaultIsTensor) {
|
|
defaultValuesString = "coords[0], coords[1]";
|
|
}
|
|
const defaultValueSnippet = `getDefaultValue(${defaultValuesString})`;
|
|
const strideString = sliceDim > 1 ? "strides[j]" : "strides";
|
|
this.userCode = `
|
|
${stridesType} strides = ${stridesType}(${strides});
|
|
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
float sum = 0.0;
|
|
bool found = false;
|
|
for (int i = 0; i < ${updateSize}; i++) {
|
|
int flattenedIndex = 0;
|
|
for (int j = 0; j < ${sliceDim}; j++) {
|
|
int index = round(${indicesSnippet});
|
|
flattenedIndex += index * ${strideString};
|
|
}
|
|
if (flattenedIndex == coords[0]) {
|
|
sum += ${updatesSnippet};
|
|
found = true;
|
|
}
|
|
}
|
|
setOutput(mix(${defaultValueSnippet}, sum, float(found)));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class ScatterPackedProgram {
|
|
constructor(updateSize, sliceDim, indicesRank, updatesRank, strides, shape, summingDupeIndex = true, defaultIsTensor = false) {
|
|
this.variableNames = ["updates", "indices", "defaultValue"];
|
|
this.packedInputs = true;
|
|
this.packedOutput = true;
|
|
this.outputShape = shape;
|
|
const stridesType = getCoordsDataType(strides.length);
|
|
const dtype = getCoordsDataType(shape.length);
|
|
let indicesString = "";
|
|
if (indicesRank === 1) {
|
|
indicesString = "i";
|
|
} else if (indicesRank === 2) {
|
|
indicesString = "i, j";
|
|
}
|
|
const indicesSnippet = `getIndices(${indicesString})`;
|
|
let updatesString = "";
|
|
if (updatesRank === 1) {
|
|
updatesString = "i";
|
|
} else if (updatesRank === 2) {
|
|
updatesString = "i, coords[1]";
|
|
}
|
|
const updatesSnippet = `getUpdates(${updatesString})`;
|
|
let defaultValuesString = "";
|
|
if (defaultIsTensor) {
|
|
defaultValuesString = "coords[0], coords[1]";
|
|
}
|
|
const defaultValueSnippet = `getDefaultValue(${defaultValuesString})`;
|
|
const strideString = sliceDim > 1 ? "strides[j]" : "strides";
|
|
const strideString2 = sliceDim > 1 ? "strides[j + 1]" : "strides";
|
|
this.userCode = `
|
|
${stridesType} strides = ${stridesType}(${strides});
|
|
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
vec4 sum = vec4(0.);
|
|
vec4 found = vec4(0.);
|
|
for (int i = 0; i < ${updateSize}; i+=2) {
|
|
ivec2 flattenedIndex = ivec2(0);
|
|
for (int j = 0; j < ${sliceDim}; j+=2) {
|
|
ivec4 index = round(${indicesSnippet});
|
|
flattenedIndex += index.xz * ${strideString};
|
|
if (j + 1 < ${sliceDim}) {
|
|
flattenedIndex += index.yw * ${strideString2};
|
|
}
|
|
}
|
|
if (flattenedIndex[0] == coords[0] || flattenedIndex[1] == coords[0] ||
|
|
flattenedIndex[0] == coords[0] + 1 || flattenedIndex[1] == coords[0] + 1) {
|
|
vec4 updVals = ${updatesSnippet};
|
|
if (flattenedIndex[0] == coords[0]) {
|
|
sum.xy += updVals.xy;
|
|
found.xy = vec2(1.);
|
|
} else if (flattenedIndex[0] == coords[0] + 1) {
|
|
sum.zw += updVals.xy;
|
|
found.zw = vec2(1.);
|
|
}
|
|
if (flattenedIndex[1] == coords[0]) {
|
|
sum.xy += updVals.zw;
|
|
found.xy = vec2(1.);
|
|
} else if (flattenedIndex[1] == coords[0] + 1) {
|
|
sum.zw += updVals.zw;
|
|
found.zw = vec2(1.);
|
|
}
|
|
}
|
|
}
|
|
setOutput(mix(${defaultValueSnippet}, sum, found));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function scatterNd$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { indices, updates } = inputs;
|
|
const { shape } = attrs;
|
|
const { sliceRank, numUpdates, sliceSize, strides, outputSize } = calculateShapes(updates, indices, shape);
|
|
const flattenShape = [outputSize / sliceSize, sliceSize];
|
|
if (outputSize === 0) {
|
|
return backend2.makeTensorInfo(shape, indices.dtype);
|
|
}
|
|
const flattenIndices = reshape$1({ inputs: { x: indices }, backend: backend2, attrs: { shape: [numUpdates, sliceRank] } });
|
|
const flattenX = reshape$1({ inputs: { x: updates }, backend: backend2, attrs: { shape: [numUpdates, sliceSize] } });
|
|
const defaultValue = backend2.makeTensorInfo([], "float32", new Float32Array([0]));
|
|
let program;
|
|
if (env().getBool("WEBGL_PACK")) {
|
|
program = new ScatterPackedProgram(numUpdates, sliceRank, flattenIndices.shape.length, flattenX.shape.length, strides, flattenShape);
|
|
} else {
|
|
program = new ScatterProgram(numUpdates, sliceRank, flattenIndices.shape.length, flattenX.shape.length, strides, flattenShape);
|
|
}
|
|
const res = backend2.runWebGLProgram(program, [flattenX, flattenIndices, defaultValue], flattenX.dtype);
|
|
const reshaped = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape } });
|
|
backend2.disposeIntermediateTensorInfo(flattenIndices);
|
|
backend2.disposeIntermediateTensorInfo(flattenX);
|
|
backend2.disposeIntermediateTensorInfo(res);
|
|
backend2.disposeIntermediateTensorInfo(defaultValue);
|
|
return reshaped;
|
|
}
|
|
const scatterNdConfig$1 = {
|
|
kernelName: ScatterNd,
|
|
backendName: "webgl",
|
|
kernelFunc: scatterNd$1
|
|
};
|
|
class SearchSortedProgram {
|
|
constructor(batchSize, numInputs, numValues, side) {
|
|
this.variableNames = ["sortedSequence", "values"];
|
|
this.customUniforms = [{ name: "numInputs", type: "int" }];
|
|
this.outputShape = [batchSize, numValues];
|
|
const webGL2LoopHead = "while (left < right) {";
|
|
const webGL1LoopHead = `for (int i = 0; i < ${Math.ceil(Math.log2(numInputs + 1))}; ++i) { if (left >= right) break;`;
|
|
const loopHead = env().getNumber("WEBGL_VERSION") === 2 ? webGL2LoopHead : webGL1LoopHead;
|
|
const boundComparator = side === "left" ? "<" : "<=";
|
|
this.userCode = `
|
|
int findBound(int batch, float value) {
|
|
int left = 0;
|
|
int right = numInputs;
|
|
int mid;
|
|
${loopHead}
|
|
mid = (left + right) / 2;
|
|
if (getSortedSequence(batch, mid) ${boundComparator} value) {
|
|
left = mid + 1;
|
|
} else {
|
|
right = mid;
|
|
}
|
|
}
|
|
return right;
|
|
}
|
|
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int valueIndex = coords[1];
|
|
|
|
float value = getValues(batch, valueIndex);
|
|
|
|
setOutput(float(findBound(batch, value)));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function searchSorted$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { sortedSequence, values } = inputs;
|
|
const { side } = attrs;
|
|
const program = new SearchSortedProgram(sortedSequence.shape[0], sortedSequence.shape[1], values.shape[1], side);
|
|
const customValues = [[sortedSequence.shape[1]]];
|
|
return backend2.runWebGLProgram(program, [sortedSequence, values], "int32", customValues);
|
|
}
|
|
const searchSortedConfig$1 = {
|
|
kernelName: SearchSorted,
|
|
backendName: "webgl",
|
|
kernelFunc: searchSorted$1
|
|
};
|
|
class SelectProgram {
|
|
constructor(cRank, shape, rank) {
|
|
this.variableNames = ["c", "a", "b"];
|
|
this.outputShape = shape;
|
|
let cCoords;
|
|
let abCoords;
|
|
if (rank > 4) {
|
|
throw Error(`Where for rank ${rank} is not yet supported`);
|
|
}
|
|
if (rank === 1) {
|
|
abCoords = `resRC`;
|
|
cCoords = `resRC`;
|
|
} else {
|
|
const currentCoords = ["resRC.x", "resRC.y", "resRC.z", "resRC.w"];
|
|
const cCoordVars = [];
|
|
const abCoordVars = [];
|
|
for (let i = 0; i < shape.length; i++) {
|
|
abCoordVars.push(`${currentCoords[i]}`);
|
|
if (i < cRank) {
|
|
cCoordVars.push(`${currentCoords[i]}`);
|
|
}
|
|
}
|
|
cCoords = cCoordVars.join();
|
|
abCoords = abCoordVars.join();
|
|
}
|
|
const dtype = getCoordsDataType(rank);
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} resRC = getOutputCoords();
|
|
float cVal = getC(${cCoords});
|
|
if (cVal >= 1.0) {
|
|
setOutput(getA(${abCoords}));
|
|
} else {
|
|
setOutput(getB(${abCoords}));
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function select$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { condition, t: t2, e } = inputs;
|
|
const program = new SelectProgram(condition.shape.length, t2.shape, t2.shape.length);
|
|
return backend2.runWebGLProgram(program, [condition, t2, e], upcastType(t2.dtype, e.dtype));
|
|
}
|
|
const selectConfig$1 = {
|
|
kernelName: Select,
|
|
backendName: "webgl",
|
|
kernelFunc: select$1
|
|
};
|
|
const SELU = `
|
|
// Stable and Attracting Fixed Point (0, 1) for Normalized Weights.
|
|
// see: https://arxiv.org/abs/1706.02515
|
|
float scaleAlpha = ${SELU_SCALEALPHA};
|
|
float scale = ${SELU_SCALE};
|
|
return (x >= 0.0) ? scale * x : scaleAlpha * (exp(x) - 1.0);
|
|
`;
|
|
const selu$1 = unaryKernelFunc({ opSnippet: SELU });
|
|
const seluConfig$1 = {
|
|
kernelName: Selu,
|
|
backendName: "webgl",
|
|
kernelFunc: selu$1
|
|
};
|
|
const SIGMOID = CHECK_NAN_SNIPPET_UNARY + `
|
|
return 1.0 / (1.0 + exp(-1.0 * x));
|
|
`;
|
|
const SIGMOID_PACKED = `
|
|
vec4 result = 1.0 / (1.0 + exp(-1.0 * x));
|
|
bvec4 isNaN = isnan(x);
|
|
|
|
result.r = isNaN.r ? x.r : result.r;
|
|
result.g = isNaN.g ? x.g : result.g;
|
|
result.b = isNaN.b ? x.b : result.b;
|
|
result.a = isNaN.a ? x.a : result.a;
|
|
|
|
return result;
|
|
`;
|
|
const sigmoid = unaryKernelFunc({
|
|
opSnippet: SIGMOID,
|
|
packedOpSnippet: SIGMOID_PACKED,
|
|
cpuKernelImpl: sigmoidImplCPU
|
|
});
|
|
const sigmoidConfig = {
|
|
kernelName: Sigmoid,
|
|
backendName: "webgl",
|
|
kernelFunc: sigmoid
|
|
};
|
|
const SIGN = `
|
|
if (isnan(x)) { return 0.0; }
|
|
return sign(x);
|
|
`;
|
|
const sign$1 = unaryKernelFunc({ opSnippet: SIGN });
|
|
const signConfig$1 = {
|
|
kernelName: Sign,
|
|
backendName: "webgl",
|
|
kernelFunc: sign$1
|
|
};
|
|
const SIN = CHECK_NAN_SNIPPET_UNARY + `
|
|
return sin(x);
|
|
`;
|
|
const SIN_PACKED = `
|
|
vec4 result = sin(x);
|
|
bvec4 isNaN = isnan(x);
|
|
${CHECK_NAN_SNIPPET_PACKED}
|
|
return result;
|
|
`;
|
|
const sin$1 = unaryKernelFunc({ opSnippet: SIN, packedOpSnippet: SIN_PACKED });
|
|
const sinConfig$1 = {
|
|
kernelName: Sin,
|
|
backendName: "webgl",
|
|
kernelFunc: sin$1
|
|
};
|
|
const SINH = `
|
|
float e2x = exp(x);
|
|
return (e2x - 1.0 / e2x) / 2.0;
|
|
`;
|
|
const sinh$1 = unaryKernelFunc({ opSnippet: SINH });
|
|
const sinhConfig$1 = {
|
|
kernelName: Sinh,
|
|
backendName: "webgl",
|
|
kernelFunc: sinh$1
|
|
};
|
|
const SOFTPLUS = `
|
|
float epsilon = 1.1920928955078125e-7;
|
|
float threshold = log(epsilon) + 2.0;
|
|
|
|
bool too_large = x > -threshold;
|
|
bool too_small = x < threshold;
|
|
|
|
float result;
|
|
float exp_x = exp(x);
|
|
|
|
if (too_large){
|
|
result = x;
|
|
}
|
|
else if (too_small){
|
|
result = exp_x;
|
|
}
|
|
else{
|
|
result = log(exp_x + 1.0);
|
|
}
|
|
return result;
|
|
`;
|
|
const softplus$1 = unaryKernelFunc({ opSnippet: SOFTPLUS });
|
|
const softplusConfig$1 = {
|
|
kernelName: Softplus,
|
|
backendName: "webgl",
|
|
kernelFunc: softplus$1
|
|
};
|
|
const spaceToBatchND$1 = (args) => {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { blockShape, paddings } = attrs;
|
|
assert(x.shape.length <= 4, () => "spaceToBatchND for rank > 4 with a WebGL backend not implemented yet");
|
|
const prod2 = blockShape.reduce((a, b) => a * b);
|
|
const completePaddings = [[0, 0]];
|
|
completePaddings.push(...paddings);
|
|
for (let i = 1 + blockShape.length; i < x.shape.length; ++i) {
|
|
completePaddings.push([0, 0]);
|
|
}
|
|
const toDispose = [];
|
|
const paddedX = padV2$1({
|
|
inputs: { x },
|
|
backend: backend2,
|
|
attrs: { paddings: completePaddings, constantValue: 0 }
|
|
});
|
|
const reshapedPaddedShape = getReshaped(paddedX.shape, blockShape, prod2, false);
|
|
const permutedReshapedPaddedPermutation = getPermuted(reshapedPaddedShape.length, blockShape.length, false);
|
|
const flattenShape = getReshapedPermuted(paddedX.shape, blockShape, prod2, false);
|
|
const reshapedPaddedX = reshape$1({ inputs: { x: paddedX }, backend: backend2, attrs: { shape: reshapedPaddedShape } });
|
|
const paddedXT = transpose({
|
|
inputs: { x: reshapedPaddedX },
|
|
backend: backend2,
|
|
attrs: { perm: permutedReshapedPaddedPermutation }
|
|
});
|
|
const result = reshape$1({ inputs: { x: paddedXT }, backend: backend2, attrs: { shape: flattenShape } });
|
|
toDispose.push(paddedX);
|
|
toDispose.push(reshapedPaddedX);
|
|
toDispose.push(paddedXT);
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
};
|
|
const spaceToBatchNDConfig$1 = {
|
|
kernelName: SpaceToBatchND,
|
|
backendName: "webgl",
|
|
kernelFunc: spaceToBatchND$1
|
|
};
|
|
function sparseFillEmptyRows$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { indices, values, denseShape, defaultValue } = inputs;
|
|
if (denseShape.shape.length !== 1) {
|
|
throw new Error(`Dense shape must be a vector, saw:
|
|
${denseShape.shape}`);
|
|
}
|
|
if (indices.shape.length !== 2) {
|
|
throw new Error(`Indices must be a matrix, saw:
|
|
${indices.shape}`);
|
|
}
|
|
if (values.shape.length !== 1) {
|
|
throw new Error(`Values must be a vector, saw:
|
|
${values.shape}`);
|
|
}
|
|
if (defaultValue.shape.length !== 0) {
|
|
throw new Error(`Default value must be a scalar, saw:
|
|
${defaultValue.shape}`);
|
|
}
|
|
const $indices = backend2.readSync(indices.dataId);
|
|
const $values = backend2.readSync(values.dataId);
|
|
const $denseShape = backend2.readSync(denseShape.dataId);
|
|
const $defaultValue = backend2.readSync(defaultValue.dataId)[0];
|
|
const [outputIndices, outputIndicesShape, outputValues, emptyRowIndicator, reverseIndexMap] = sparseFillEmptyRowsImplCPU($indices, indices.shape, indices.dtype, $values, values.dtype, $denseShape, $defaultValue);
|
|
return [
|
|
backend2.makeTensorInfo(outputIndicesShape, indices.dtype, outputIndices),
|
|
backend2.makeTensorInfo([outputIndicesShape[0]], values.dtype, outputValues),
|
|
backend2.makeTensorInfo([emptyRowIndicator.length], "bool", new Uint8Array(emptyRowIndicator.map((value) => Number(value)))),
|
|
backend2.makeTensorInfo([reverseIndexMap.length], indices.dtype, new Int32Array(reverseIndexMap))
|
|
];
|
|
}
|
|
const sparseFillEmptyRowsConfig$1 = {
|
|
kernelName: SparseFillEmptyRows,
|
|
backendName: "webgl",
|
|
kernelFunc: sparseFillEmptyRows$1
|
|
};
|
|
function sparseReshape$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { inputIndices, inputShape, newShape } = inputs;
|
|
if (inputIndices.shape.length !== 2) {
|
|
throw new Error(`Input indices should be a matrix but received shape ${inputIndices.shape}`);
|
|
}
|
|
if (inputShape.shape.length !== 1) {
|
|
throw new Error(`Input shape should be a vector but received shape ${inputShape.shape}`);
|
|
}
|
|
if (newShape.shape.length !== 1) {
|
|
throw new Error(`Target shape should be a vector but received shape ${newShape.shape}`);
|
|
}
|
|
const $inputShape = Array.from(backend2.readSync(inputShape.dataId));
|
|
const $inputIndices = backend2.readSync(inputIndices.dataId);
|
|
const targetShape = Array.from(backend2.readSync(newShape.dataId));
|
|
const [newIndices, indicesShape, outputShape] = sparseReshapeImplCPU($inputIndices, inputIndices.shape, inputIndices.dtype, $inputShape, targetShape);
|
|
return [
|
|
backend2.makeTensorInfo(indicesShape, inputIndices.dtype, newIndices),
|
|
backend2.makeTensorInfo([outputShape.length], newShape.dtype, new Int32Array(outputShape))
|
|
];
|
|
}
|
|
const sparseReshapeConfig$1 = {
|
|
kernelName: SparseReshape,
|
|
backendName: "webgl",
|
|
kernelFunc: sparseReshape$1
|
|
};
|
|
function sparseSegmentMean$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { data, indices, segmentIds } = inputs;
|
|
if (data.shape.length < 1) {
|
|
throw new Error(`Data should be at least 1 dimensional but received scalar`);
|
|
}
|
|
if (indices.shape.length !== 1) {
|
|
throw new Error(`Indices should be a vector but received shape
|
|
${indices.shape}`);
|
|
}
|
|
if (segmentIds.shape.length !== 1) {
|
|
throw new Error(`Segment ids should be a vector but received shape
|
|
${segmentIds.shape}`);
|
|
}
|
|
const $data = backend2.readSync(data.dataId);
|
|
const $indices = backend2.readSync(indices.dataId);
|
|
const $segmentIds = backend2.readSync(segmentIds.dataId);
|
|
const [outputData, outputDataShape] = sparseSegmentReductionImplCPU($data, data.shape, data.dtype, $indices, $segmentIds, true);
|
|
return backend2.makeTensorInfo(outputDataShape, data.dtype, outputData);
|
|
}
|
|
const sparseSegmentMeanConfig$1 = {
|
|
kernelName: SparseSegmentMean,
|
|
backendName: "webgl",
|
|
kernelFunc: sparseSegmentMean$1
|
|
};
|
|
function sparseSegmentSum$1(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { data, indices, segmentIds } = inputs;
|
|
if (data.shape.length < 1) {
|
|
throw new Error(`Data should be at least 1 dimensional but received scalar`);
|
|
}
|
|
if (indices.shape.length !== 1) {
|
|
throw new Error(`Indices should be a vector but received shape
|
|
${indices.shape}`);
|
|
}
|
|
if (segmentIds.shape.length !== 1) {
|
|
throw new Error(`Segment ids should be a vector but received shape
|
|
${segmentIds.shape}`);
|
|
}
|
|
const $data = backend2.readSync(data.dataId);
|
|
const $indices = backend2.readSync(indices.dataId);
|
|
const $segmentIds = backend2.readSync(segmentIds.dataId);
|
|
const [outputData, outputDataShape] = sparseSegmentReductionImplCPU($data, data.shape, data.dtype, $indices, $segmentIds);
|
|
return backend2.makeTensorInfo(outputDataShape, data.dtype, outputData);
|
|
}
|
|
const sparseSegmentSumConfig$1 = {
|
|
kernelName: SparseSegmentSum,
|
|
backendName: "webgl",
|
|
kernelFunc: sparseSegmentSum$1
|
|
};
|
|
function sparseToDense$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { sparseIndices, sparseValues, defaultValue } = inputs;
|
|
const { outputShape } = attrs;
|
|
const { sliceRank, numUpdates, sliceSize, strides, outputSize } = calculateShapes(sparseValues, sparseIndices, outputShape);
|
|
const sumDupeIndices = false;
|
|
if (sparseValues.dtype === "string") {
|
|
const indicesBuf = backend2.bufferSync(sparseIndices);
|
|
const updatesBuf = backend2.bufferSync(sparseValues);
|
|
const $defaultValue = decodeString(backend2.readSync(defaultValue.dataId)[0]);
|
|
const outBuf = scatterImplCPU(indicesBuf, updatesBuf, outputShape, outputSize, sliceSize, numUpdates, sliceRank, strides, $defaultValue, sumDupeIndices);
|
|
return backend2.makeTensorInfo(outputShape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const program = new ScatterProgram(numUpdates, sliceRank, sparseIndices.shape.length, sparseValues.shape.length, strides, [outputSize, 1], sumDupeIndices);
|
|
const res = backend2.runWebGLProgram(program, [sparseValues, sparseIndices, defaultValue], sparseValues.dtype);
|
|
const reshaped = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape: outputShape } });
|
|
backend2.disposeIntermediateTensorInfo(res);
|
|
return reshaped;
|
|
}
|
|
const sparseToDenseConfig$1 = {
|
|
kernelName: SparseToDense,
|
|
backendName: "webgl",
|
|
kernelFunc: sparseToDense$1
|
|
};
|
|
function splitV$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { numOrSizeSplits, axis } = attrs;
|
|
const $axis = parseAxisParam(axis, x.shape)[0];
|
|
const splitSizes = prepareSplitSize(x, numOrSizeSplits, $axis);
|
|
const xRank = x.shape.length;
|
|
const begin = new Array(xRank).fill(0);
|
|
const size = x.shape.slice();
|
|
return splitSizes.map((s) => {
|
|
const sliceSize = [...size];
|
|
sliceSize[$axis] = s;
|
|
const sliceT = slice({ inputs: { x }, backend: backend2, attrs: { begin, size: sliceSize } });
|
|
begin[$axis] += s;
|
|
return sliceT;
|
|
});
|
|
}
|
|
const splitVConfig$1 = {
|
|
kernelName: SplitV,
|
|
backendName: "webgl",
|
|
kernelFunc: splitV$1
|
|
};
|
|
const SQRT = `return sqrt(x);`;
|
|
const sqrt = unaryKernelFunc({ opSnippet: SQRT, packedOpSnippet: SQRT, cpuKernelImpl: sqrtImplCPU });
|
|
const sqrtConfig = {
|
|
kernelName: Sqrt,
|
|
backendName: "webgl",
|
|
kernelFunc: sqrt
|
|
};
|
|
const SQUARE = `return x * x;`;
|
|
const square = unaryKernelFunc({ opSnippet: SQUARE });
|
|
const squareConfig$1 = {
|
|
kernelName: Square,
|
|
backendName: "webgl",
|
|
kernelFunc: square
|
|
};
|
|
const SQUARED_DIFFERENCE = "return (a - b) * (a - b);";
|
|
const squaredDifference = binaryKernelFunc({ opSnippet: SQUARED_DIFFERENCE, packedOpSnippet: SQUARED_DIFFERENCE });
|
|
const squaredDifferenceConfig = {
|
|
kernelName: SquaredDifference,
|
|
backendName: "webgl",
|
|
kernelFunc: squaredDifference
|
|
};
|
|
function staticRegexReplace(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
if (x.dtype !== "string") {
|
|
throw new Error("Input must be of datatype string");
|
|
}
|
|
const $x = backend2.readSync(x.dataId);
|
|
const stringInput = fromUint8ToStringArray($x);
|
|
const output = staticRegexReplaceImplCPU(stringInput, "string", attrs);
|
|
return backend2.makeTensorInfo(x.shape, "string", output);
|
|
}
|
|
const staticRegexReplaceConfig = {
|
|
kernelName: StaticRegexReplace,
|
|
backendName: "webgl",
|
|
kernelFunc: staticRegexReplace
|
|
};
|
|
function step$1({ inputs, attrs, backend: backend2 }) {
|
|
const { x } = inputs;
|
|
const opSnippet = CHECK_NAN_SNIPPET$1 + `
|
|
return x > 0.0 ? 1.0 : float(${attrs.alpha});
|
|
`;
|
|
const program = new UnaryOpProgram(x.shape, opSnippet);
|
|
return backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
const stepConfig$1 = {
|
|
kernelName: Step,
|
|
backendName: "webgl",
|
|
kernelFunc: step$1
|
|
};
|
|
class StridedSliceProgram {
|
|
constructor(begin, strides, size) {
|
|
this.variableNames = ["x"];
|
|
this.outputShape = size;
|
|
const rank = size.length;
|
|
const inputDtype = getCoordsDataType(size.length);
|
|
const dtype = getCoordsDataType(size.length);
|
|
let newCoords = "";
|
|
if (rank === 1) {
|
|
newCoords = "coords * strides + begin";
|
|
} else {
|
|
let outputAxis = 0;
|
|
newCoords = size.map((_, i) => {
|
|
outputAxis++;
|
|
return size.length === 1 ? `coords * strides[${i}] + begin[${i}]` : `coords[${outputAxis - 1}] * strides[${i}] + begin[${i}]`;
|
|
}).join(",");
|
|
}
|
|
this.userCode = `
|
|
${inputDtype} begin = ${inputDtype}(${begin});
|
|
${inputDtype} strides = ${inputDtype}(${strides});
|
|
|
|
void main() {
|
|
${dtype} coords = getOutputCoords();
|
|
setOutput(getX(${newCoords}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function stridedSlice$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { begin, end, strides, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask } = attrs;
|
|
const { finalShapeSparse, finalShape, isIdentity, sliceDim0, isSimpleSlice, begin: $begin, end: $end, strides: $strides } = sliceInfo(x.shape, begin, end, strides, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask);
|
|
let result;
|
|
if (isIdentity) {
|
|
result = reshape$1({ inputs: { x }, backend: backend2, attrs: { shape: finalShape } });
|
|
} else if (sliceDim0 || isSimpleSlice) {
|
|
assert(x.shape.length >= 1, () => `Input must have rank at least 1, got: ${x.shape.length}`);
|
|
const size = computeOutShape$2($begin, $end, $strides);
|
|
const sliced = slice({ inputs: { x }, backend: backend2, attrs: { begin: $begin, size } });
|
|
result = reshape$1({ inputs: { x: sliced }, backend: backend2, attrs: { shape: finalShape } });
|
|
backend2.disposeIntermediateTensorInfo(sliced);
|
|
} else {
|
|
const shouldExecuteOnCPU = backend2.shouldExecuteOnCPU([x]);
|
|
if (shouldExecuteOnCPU) {
|
|
const values = backend2.readSync(x.dataId);
|
|
const xBuf = buffer(x.shape, x.dtype, values);
|
|
const resultValues = stridedSliceImplCPU(finalShapeSparse, xBuf, $strides, $begin);
|
|
result = backend2.makeTensorInfo(finalShape, x.dtype, resultValues.values);
|
|
} else {
|
|
const program = new StridedSliceProgram($begin, $strides, finalShapeSparse);
|
|
result = backend2.runWebGLProgram(program, [x], x.dtype);
|
|
}
|
|
}
|
|
const resultReshaped = reshape$1({ inputs: { x: result }, backend: backend2, attrs: { shape: finalShape } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return resultReshaped;
|
|
}
|
|
const stridedSliceConfig$1 = {
|
|
kernelName: StridedSlice,
|
|
backendName: "webgl",
|
|
kernelFunc: stridedSlice$1
|
|
};
|
|
function stringNGrams$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { separator, nGramWidths, leftPad, rightPad: rightPad2, padWidth, preserveShortSequences } = attrs;
|
|
const { data, dataSplits } = inputs;
|
|
const $data = backend2.readSync(data.dataId);
|
|
const $dataSplits = backend2.readSync(dataSplits.dataId);
|
|
const [nGrams, nGramsSplits] = stringNGramsImplCPU($data, $dataSplits, separator, nGramWidths, leftPad, rightPad2, padWidth, preserveShortSequences);
|
|
return [
|
|
backend2.makeTensorInfo([nGrams.length], "string", nGrams),
|
|
backend2.makeTensorInfo(dataSplits.shape, "int32", nGramsSplits)
|
|
];
|
|
}
|
|
const stringNGramsConfig$1 = {
|
|
kernelName: StringNGrams,
|
|
backendName: "webgl",
|
|
kernelFunc: stringNGrams$1
|
|
};
|
|
function stringSplit$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { skipEmpty } = attrs;
|
|
const { input, delimiter } = inputs;
|
|
if (input.dtype !== "string") {
|
|
throw new Error("Input must be of datatype string");
|
|
}
|
|
if (input.shape.length !== 1) {
|
|
throw new Error(`Input must be a vector, got shape: ${input.shape}`);
|
|
}
|
|
if (delimiter.shape.length !== 0) {
|
|
throw new Error(`Delimiter must be a scalar, got shape: ${delimiter.shape}`);
|
|
}
|
|
const $input = backend2.readSync(input.dataId);
|
|
const $delimiter = backend2.readSync(delimiter.dataId)[0];
|
|
const [indices, values, shape] = stringSplitImplCPU($input, $delimiter, skipEmpty);
|
|
const outputSize = values.length;
|
|
return [
|
|
backend2.makeTensorInfo([outputSize, 2], "int32", indices),
|
|
backend2.makeTensorInfo([outputSize], "string", values),
|
|
backend2.makeTensorInfo([2], "int32", new Int32Array(shape))
|
|
];
|
|
}
|
|
const stringSplitConfig$1 = {
|
|
kernelName: StringSplit,
|
|
backendName: "webgl",
|
|
kernelFunc: stringSplit$1
|
|
};
|
|
function stringToHashBucketFast$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { numBuckets } = attrs;
|
|
const { input } = inputs;
|
|
if (input.dtype !== "string") {
|
|
throw new Error("Input must be of datatype string");
|
|
}
|
|
if (numBuckets <= 0) {
|
|
throw new Error(`Number of buckets must be at least 1`);
|
|
}
|
|
const $input = backend2.readSync(input.dataId);
|
|
const output = stringToHashBucketFastImplCPU($input, numBuckets);
|
|
return backend2.makeTensorInfo(input.shape, "int32", output);
|
|
}
|
|
const stringToHashBucketFastConfig$1 = {
|
|
kernelName: StringToHashBucketFast,
|
|
backendName: "webgl",
|
|
kernelFunc: stringToHashBucketFast$1
|
|
};
|
|
const TAN = `return tan(x);`;
|
|
const tan$1 = unaryKernelFunc({ opSnippet: TAN });
|
|
const tanConfig$1 = {
|
|
kernelName: Tan,
|
|
backendName: "webgl",
|
|
kernelFunc: tan$1
|
|
};
|
|
const TANH = `
|
|
float e2x = exp(-2.0 * abs(x));
|
|
return sign(x) * (1.0 - e2x) / (1.0 + e2x);
|
|
`;
|
|
const tanh$1 = unaryKernelFunc({ opSnippet: TANH });
|
|
const tanhConfig$1 = {
|
|
kernelName: Tanh,
|
|
backendName: "webgl",
|
|
kernelFunc: tanh$1
|
|
};
|
|
function tensorScatterUpdate$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { tensor: tensor2, indices, updates } = inputs;
|
|
const { sliceRank, numUpdates, sliceSize, strides, outputSize } = calculateShapes(updates, indices, tensor2.shape);
|
|
const flattenShape = [outputSize / sliceSize, sliceSize];
|
|
if (outputSize === 0) {
|
|
return backend2.makeTensorInfo(tensor2.shape, indices.dtype);
|
|
}
|
|
const flattenIndices = reshape$1({ inputs: { x: indices }, backend: backend2, attrs: { shape: [numUpdates, sliceRank] } });
|
|
const flattenX = reshape$1({ inputs: { x: updates }, backend: backend2, attrs: { shape: [numUpdates, sliceSize] } });
|
|
const flattenTensor = reshape$1({ inputs: { x: tensor2 }, backend: backend2, attrs: { shape: flattenShape } });
|
|
const program = new ScatterProgram(numUpdates, sliceRank, flattenIndices.shape.length, flattenX.shape.length, strides, flattenShape, false, true);
|
|
const res = backend2.runWebGLProgram(program, [flattenX, flattenIndices, flattenTensor], flattenTensor.dtype);
|
|
const reshaped = reshape$1({ inputs: { x: res }, backend: backend2, attrs: { shape: tensor2.shape } });
|
|
backend2.disposeIntermediateTensorInfo(flattenIndices);
|
|
backend2.disposeIntermediateTensorInfo(flattenX);
|
|
backend2.disposeIntermediateTensorInfo(flattenTensor);
|
|
backend2.disposeIntermediateTensorInfo(res);
|
|
return reshaped;
|
|
}
|
|
const tensorScatterUpdateConfig$1 = {
|
|
kernelName: TensorScatterUpdate,
|
|
backendName: "webgl",
|
|
kernelFunc: tensorScatterUpdate$1
|
|
};
|
|
class TileProgram {
|
|
constructor(aShape, reps) {
|
|
this.variableNames = ["A"];
|
|
const outputShape = new Array(aShape.length);
|
|
for (let i = 0; i < outputShape.length; i++) {
|
|
outputShape[i] = aShape[i] * reps[i];
|
|
}
|
|
this.outputShape = outputShape;
|
|
this.rank = outputShape.length;
|
|
const dtype = getCoordsDataType(this.rank);
|
|
const sourceCoords = getSourceCoords(aShape);
|
|
this.userCode = `
|
|
void main() {
|
|
${dtype} resRC = getOutputCoords();
|
|
setOutput(getA(${sourceCoords}));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function getSourceCoords(aShape) {
|
|
const rank = aShape.length;
|
|
if (rank > 5) {
|
|
throw Error(`Tile for rank ${rank} is not yet supported`);
|
|
}
|
|
if (rank === 1) {
|
|
return `imod(resRC, ${aShape[0]})`;
|
|
}
|
|
const currentCoords = ["resRC.x", "resRC.y", "resRC.z", "resRC.w", "resRC.u"];
|
|
const sourceCoords = [];
|
|
for (let i = 0; i < aShape.length; i++) {
|
|
sourceCoords.push(`imod(${currentCoords[i]}, ${aShape[i]})`);
|
|
}
|
|
return sourceCoords.join();
|
|
}
|
|
function tile$1(params) {
|
|
const { inputs, backend: backend2, attrs } = params;
|
|
const { x } = inputs;
|
|
const { reps } = attrs;
|
|
if (x.dtype === "string" || x.shape.length > 5) {
|
|
const data = backend2.readSync(x.dataId);
|
|
const value = x.dtype === "string" ? data.map((d) => decodeString(d)) : data;
|
|
const buf = buffer(x.shape, x.dtype, value);
|
|
const outBuf = tileImplCPU(buf, reps);
|
|
return backend2.makeTensorInfo(outBuf.shape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const program = new TileProgram(x.shape, reps);
|
|
const output = backend2.runWebGLProgram(program, [x], x.dtype);
|
|
return output;
|
|
}
|
|
const tileConfig$1 = {
|
|
kernelName: Tile,
|
|
backendName: "webgl",
|
|
kernelFunc: tile$1
|
|
};
|
|
class SwapProgram {
|
|
/**
|
|
* @param shape desired output shape (can be larger than input shape, output
|
|
* will be padded with -Infinity)
|
|
*/
|
|
constructor(shape) {
|
|
this.variableNames = ["x", "indices"];
|
|
this.customUniforms = [
|
|
{ name: "n", type: "int" },
|
|
{ name: "firstPass", type: "int" },
|
|
{ name: "negativeInf", type: "float" },
|
|
{ name: "dir", type: "int" },
|
|
{ name: "inc", type: "int" }
|
|
];
|
|
this.outputShape = shape;
|
|
this.userCode = `
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int elemIdx = coords[1];
|
|
|
|
// We compare elements pair-wise within a group of size 2 * inc.
|
|
// The comparing rule for each group alternates between ascending
|
|
// and descending. Within each group, we compare each pair at
|
|
// positions i and i+inc. To decide whether an element at position i
|
|
// is x0 or x1, we mod it by 2 * inc, if the result is smaller than
|
|
// inc, it is in the first half of the group, we denote it as x0,
|
|
// otherwise we denote it as x1.
|
|
// For example, as shown in the Bitonic top K paper referenced above,
|
|
// Figure5(a) shows that element[1] is in the
|
|
// second half of the group when group size is 2, but it is in the
|
|
// first half of the group when group size is 4.
|
|
|
|
bool isFirstInPair = imod(elemIdx, 2 * inc) < inc;
|
|
int i = isFirstInPair ? elemIdx : elemIdx - inc;
|
|
|
|
int i0 = firstPass == 1 ? i : int(getIndices(batch, i));
|
|
int i1 = firstPass == 1 ? i + inc : int(getIndices(batch, i + inc));
|
|
float x0 = i0 < n ? getX(batch, i0) : negativeInf;
|
|
float x1 = i1 < n ? getX(batch, i1) : negativeInf;
|
|
|
|
// Denotes which direction indices are in (ascending or descending).
|
|
bool reverse = imod(elemIdx, 2 * dir) >= dir;
|
|
bool isGreater = x0 > x1 || (x0 == x1 && i1 > i0);
|
|
if (reverse == isGreater) { // Elements in opposite order of direction
|
|
int iTemp = i0;
|
|
i0 = i1;
|
|
i1 = iTemp;
|
|
}
|
|
if (isFirstInPair) {
|
|
setOutput(float(i0));
|
|
} else {
|
|
setOutput(float(i1));
|
|
}
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
class MergeProgram {
|
|
/**
|
|
* @param shape desired output shape (must be half of the input size)
|
|
*/
|
|
constructor(shape) {
|
|
this.variableNames = ["x", "indices"];
|
|
this.customUniforms = [
|
|
{ name: "n", type: "int" },
|
|
{ name: "firstPass", type: "int" },
|
|
{ name: "k", type: "int" }
|
|
];
|
|
this.outputShape = shape;
|
|
this.userCode = `
|
|
void main() {
|
|
// Takes max of indices (0, k), (1, k + 1), (2, k + 2) ...
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int elemIdx = coords[1];
|
|
|
|
// The output size is half of the previous size.
|
|
// If the previous sequence is | | | | _ _ _ _ | | | | _ _ _ _ (k=4),
|
|
// we only need to output the indices at positions |, the indices at
|
|
// positions _ can be thrown away, see Figure5(b) After Phase 2
|
|
// (Merge phase) in the Bitonic Top K paper referenced above.
|
|
// For example, the paper shows we only need to output the orange bars.
|
|
// The output sequence should look like this | | | | | | | |.
|
|
// Because the sequence is halved, to map the output index back
|
|
// to the previous sequence to find the corresponding value,
|
|
// we need to double the index. When we double the index,
|
|
// we basically interpolate a position, so 2i looks like
|
|
// | _ | _ | _ | _ | _ | _ | _. We move the | to the first k position
|
|
// of each 2k positions by - elemIdx % k. E.g. for output at
|
|
// index 4,5,6,7, we want to get the corresponding element at
|
|
// original index 8,9,10,11, for output at index 8,9,10,11,
|
|
// we want to get the corresponding element at original index
|
|
// 16,17,18,19, so on and so forth.
|
|
|
|
int i = elemIdx < k ? elemIdx : (elemIdx * 2 - imod(elemIdx, k));
|
|
int i0 = firstPass == 1 ? i : int(getIndices(batch, i));
|
|
int i1 = firstPass == 1 ? i + k : int(getIndices(batch, i + k));
|
|
|
|
float x0 = getX(batch, i0);
|
|
float x1 = i1 < n ? getX(batch, i1) : x0;
|
|
|
|
setOutput(x0 >= x1 ? float(i0) : float(i1));
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function disposeIntermediateTensorInfoOrNull(backend2, tensorInfo) {
|
|
if (tensorInfo !== null) {
|
|
backend2.disposeIntermediateTensorInfo(tensorInfo);
|
|
}
|
|
}
|
|
function roundUpToPow2(num) {
|
|
let pow2 = 1;
|
|
while (pow2 < num) {
|
|
pow2 *= 2;
|
|
}
|
|
return pow2;
|
|
}
|
|
function topK$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { k, sorted } = attrs;
|
|
const TOPK_LAST_DIM_CPU_HANDOFF_SIZE_THRESHOLD = env().getNumber("TOPK_LAST_DIM_CPU_HANDOFF_SIZE_THRESHOLD");
|
|
const TOPK_K_CPU_HANDOFF_THRESHOLD = env().getNumber("TOPK_K_CPU_HANDOFF_THRESHOLD");
|
|
const xShape = x.shape;
|
|
const lastDim = xShape[xShape.length - 1];
|
|
if (backend2.shouldExecuteOnCPU([x]) || lastDim < TOPK_LAST_DIM_CPU_HANDOFF_SIZE_THRESHOLD || k > TOPK_K_CPU_HANDOFF_THRESHOLD) {
|
|
const xVals = backend2.readSync(x.dataId);
|
|
const [allTopKVals, allTopKIndices] = topKImplCPU(xVals, xShape, x.dtype, k, sorted);
|
|
return [
|
|
backend2.makeTensorInfo(allTopKVals.shape, allTopKVals.dtype, allTopKVals.values),
|
|
backend2.makeTensorInfo(allTopKIndices.shape, allTopKIndices.dtype, allTopKIndices.values)
|
|
];
|
|
}
|
|
if (k === 0) {
|
|
xShape[xShape.length - 1] = 0;
|
|
return [
|
|
backend2.makeTensorInfo(xShape, x.dtype, []),
|
|
backend2.makeTensorInfo(xShape, "int32", [])
|
|
];
|
|
}
|
|
if (lastDim === 1) {
|
|
return [
|
|
x,
|
|
fill$1({ attrs: { shape: xShape, dtype: "int32", value: 0 }, backend: backend2 })
|
|
];
|
|
}
|
|
const xtexData = backend2.texData.get(x.dataId);
|
|
const xIsPacked = xtexData !== null && xtexData.isPacked;
|
|
const xUnPacked = xIsPacked ? backend2.unpackTensor(x) : x;
|
|
const xSize = sizeFromShape(xShape);
|
|
const batch = xSize / lastDim;
|
|
const x2D = reshape$1({ inputs: { x: xUnPacked }, attrs: { shape: [batch, lastDim] }, backend: backend2 });
|
|
if (xIsPacked) {
|
|
disposeIntermediateTensorInfoOrNull(backend2, xUnPacked);
|
|
}
|
|
const kPow2 = roundUpToPow2(k);
|
|
const lastDimPow2 = roundUpToPow2(lastDim);
|
|
let indices = null;
|
|
const getInputs = () => indices === null ? [x2D, x2D] : [x2D, indices];
|
|
const runSwap = (dir, inc, shape) => {
|
|
const inputs2 = getInputs();
|
|
const program = new SwapProgram(shape);
|
|
const fistPass = indices === null ? 1 : 0;
|
|
const customValues = [[lastDim], [fistPass], [Number.NEGATIVE_INFINITY], [dir], [inc]];
|
|
const prevIndices2 = indices;
|
|
indices = backend2.runWebGLProgram(program, inputs2, "int32", customValues);
|
|
disposeIntermediateTensorInfoOrNull(backend2, prevIndices2);
|
|
};
|
|
for (let len = 1; len < kPow2; len *= 2) {
|
|
const dir = len * 2;
|
|
for (let inc = len; inc >= 1; inc /= 2) {
|
|
runSwap(dir, inc, [batch, lastDimPow2]);
|
|
}
|
|
}
|
|
for (let indicesSize = lastDimPow2; indicesSize > kPow2; indicesSize /= 2) {
|
|
const inputs2 = getInputs();
|
|
const mergeProgram = new MergeProgram([batch, indicesSize / 2]);
|
|
const firstPass = indices === null ? 1 : 0;
|
|
const customValues = [[lastDim], [firstPass], [kPow2]];
|
|
const prevIndices2 = indices;
|
|
indices = backend2.runWebGLProgram(mergeProgram, inputs2, "int32", customValues);
|
|
disposeIntermediateTensorInfoOrNull(backend2, prevIndices2);
|
|
const len = kPow2 / 2;
|
|
const dir = len * 2;
|
|
for (let inc = len; inc >= 1; inc /= 2) {
|
|
runSwap(dir, inc, indices.shape);
|
|
}
|
|
}
|
|
let prevIndices = indices;
|
|
indices = slice({ inputs: { x: indices }, backend: backend2, attrs: { begin: 0, size: [batch, k] } });
|
|
disposeIntermediateTensorInfoOrNull(backend2, prevIndices);
|
|
let values = gatherV2$1({ inputs: { x: x2D, indices }, backend: backend2, attrs: { axis: 1, batchDims: 1 } });
|
|
disposeIntermediateTensorInfoOrNull(backend2, x2D);
|
|
const newShape = xShape.slice(0, -1);
|
|
newShape.push(k);
|
|
prevIndices = indices;
|
|
indices = reshape$1({ inputs: { x: indices }, attrs: { shape: newShape }, backend: backend2 });
|
|
disposeIntermediateTensorInfoOrNull(backend2, prevIndices);
|
|
const prevValues = values;
|
|
values = reshape$1({ inputs: { x: values }, attrs: { shape: newShape }, backend: backend2 });
|
|
disposeIntermediateTensorInfoOrNull(backend2, prevValues);
|
|
return [values, indices];
|
|
}
|
|
const topKConfig$1 = {
|
|
kernelName: TopK,
|
|
backendName: "webgl",
|
|
kernelFunc: topK$1
|
|
};
|
|
class TransformProgram {
|
|
constructor(imageHeight, imageWidth, interpolation, fillMode, fillValue, outShape) {
|
|
this.variableNames = ["Image", "Transforms"];
|
|
this.outputShape = outShape;
|
|
const interpolationModeId = interpolation === "nearest" ? 1 : 2;
|
|
let fillModeId;
|
|
switch (fillMode) {
|
|
case "constant":
|
|
fillModeId = 1;
|
|
break;
|
|
case "reflect":
|
|
fillModeId = 2;
|
|
break;
|
|
case "wrap":
|
|
fillModeId = 3;
|
|
break;
|
|
case "nearest":
|
|
fillModeId = 4;
|
|
break;
|
|
default:
|
|
fillModeId = 1;
|
|
break;
|
|
}
|
|
this.userCode = `
|
|
float mapCoord(float outCoord, float len) {
|
|
float inCoord = outCoord;
|
|
if(${fillModeId} == 2) {
|
|
if (inCoord < 0.0) {
|
|
if (len <= 1.0) {
|
|
inCoord = 0.0;
|
|
} else {
|
|
float sz2 = 2.0 * len;
|
|
if (inCoord < sz2) {
|
|
inCoord = sz2 * float(int(float(-inCoord / sz2))) +
|
|
inCoord;
|
|
}
|
|
inCoord = inCoord < -len ? inCoord + sz2 : -inCoord - 1.0;
|
|
}
|
|
} else if (inCoord > len - 1.0) {
|
|
if (len <= 1.0) {
|
|
inCoord = 0.0;
|
|
} else {
|
|
float sz2 = 2.0 * len;
|
|
inCoord -= sz2 * float(int(float(inCoord / sz2)));
|
|
if (inCoord >= len) {
|
|
inCoord = sz2 - inCoord - 1.0;
|
|
}
|
|
}
|
|
}
|
|
return clamp(inCoord, 0.0, len - 1.0);
|
|
} else if (${fillModeId} == 3) {
|
|
if (inCoord < 0.0) {
|
|
if (len <= 1.0) {
|
|
inCoord = 0.0;
|
|
} else {
|
|
float sz = len - 1.0;
|
|
inCoord += len * (float(int(float(-inCoord / sz))) + 1.0);
|
|
}
|
|
} else if (inCoord > len - 1.0) {
|
|
if (len <= 1.0) {
|
|
inCoord = 0.0;
|
|
} else {
|
|
float sz = len - 1.0;
|
|
inCoord -= len * float(int(float(inCoord / sz)));
|
|
}
|
|
}
|
|
return clamp(inCoord, 0.0, len - 1.0);
|
|
} else if (${fillModeId} == 4) {
|
|
return clamp(outCoord, 0.0, len - 1.0);
|
|
} else {
|
|
return outCoord;
|
|
}
|
|
}
|
|
|
|
float readWithFillValue(int batch, int coordY, int coordX,
|
|
int channel) {
|
|
float outputValue;
|
|
if (0 <= coordY && coordY < ${imageHeight} && 0 <= coordX && coordX < ${imageWidth}) {
|
|
outputValue = getImage(batch, coordY, coordX, channel);
|
|
} else {
|
|
outputValue = float(${fillValue});
|
|
}
|
|
return outputValue;
|
|
}
|
|
|
|
void main() {
|
|
ivec4 coords = getOutputCoords();
|
|
float outputValue;
|
|
int batch = coords[0];
|
|
int x = coords[2];
|
|
int y = coords[1];
|
|
int channel = coords[3];
|
|
float xf = float(x);
|
|
float yf = float(y);
|
|
float a1 = getTransforms(batch, 0);
|
|
float a2 = getTransforms(batch, 1);
|
|
float a3 = getTransforms(batch, 2);
|
|
float b1 = getTransforms(batch, 3);
|
|
float b2 = getTransforms(batch, 4);
|
|
float b3 = getTransforms(batch, 5);
|
|
float c1 = getTransforms(batch, 6);
|
|
float c2 = getTransforms(batch, 7);
|
|
float projection = c1 * xf + c2 * yf + 1.0;
|
|
if (projection == 0.0) {
|
|
outputValue = float(${fillValue});
|
|
} else {
|
|
float inX = (a1 * xf + a2 * yf + a3) / projection;
|
|
float inY = (b1 * xf + b2 * yf + b3) / projection;
|
|
float mapX = mapCoord(inX, float(${imageWidth}));
|
|
float mapY = mapCoord(inY, float(${imageHeight}));
|
|
|
|
if (${interpolationModeId} == 1) {
|
|
int coordY = int(round(mapY));
|
|
int coordX = int(round(mapX));
|
|
outputValue = readWithFillValue(batch, coordY, coordX,
|
|
channel);
|
|
} else {
|
|
float yFloor = floor(mapY);
|
|
float xFloor = floor(mapX);
|
|
float yCeil = yFloor + 1.0;
|
|
float xCeil = xFloor + 1.0;
|
|
float valueYFloor = (xCeil - mapX) *
|
|
readWithFillValue(batch, int(yFloor), int(xFloor), channel) +
|
|
(mapX - xFloor) *
|
|
readWithFillValue(batch, int(yFloor), int(xCeil), channel);
|
|
float valueYCeil = (xCeil - mapX) *
|
|
readWithFillValue(batch, int(yCeil), int(xFloor), channel) +
|
|
(mapX - xFloor) *
|
|
readWithFillValue(batch, int(yCeil), int(xCeil), channel);
|
|
outputValue = (yCeil - mapY) * valueYFloor +
|
|
(mapY - yFloor) * valueYCeil;
|
|
}
|
|
}
|
|
setOutput(outputValue);
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function transform$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { image: image2, transforms } = inputs;
|
|
const { interpolation, fillMode, fillValue, outputShape } = attrs;
|
|
const [batch, imageHeight, imageWidth, numChannels] = image2.shape;
|
|
const [outHeight, outWidth] = outputShape != null ? outputShape : [imageHeight, imageWidth];
|
|
const outShape = [
|
|
batch,
|
|
outHeight,
|
|
outWidth,
|
|
numChannels
|
|
];
|
|
const program = new TransformProgram(imageHeight, imageWidth, interpolation, fillMode, fillValue, outShape);
|
|
return backend2.runWebGLProgram(program, [image2, transforms], "float32");
|
|
}
|
|
const transformConfig$1 = {
|
|
kernelName: Transform,
|
|
backendName: "webgl",
|
|
kernelFunc: transform$1
|
|
};
|
|
function unique$1(args) {
|
|
const { inputs, attrs, backend: backend2 } = args;
|
|
const { axis } = attrs;
|
|
const { x } = inputs;
|
|
assertNotComplex$1(x, "unique");
|
|
console.warn("WARNING: ", "UI might be locked temporarily as data is being downloaded");
|
|
const values = backend2.readSync(x.dataId);
|
|
const { outputValues, outputShape, indices } = uniqueImplCPU(values, axis, x.shape, x.dtype);
|
|
return [
|
|
backend2.makeTensorInfo(outputShape, x.dtype, outputValues),
|
|
backend2.makeTensorInfo([indices.length], "int32", indices)
|
|
];
|
|
}
|
|
const uniqueConfig$1 = {
|
|
kernelName: Unique,
|
|
backendName: "webgl",
|
|
kernelFunc: unique$1
|
|
};
|
|
function unpack$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { value } = inputs;
|
|
let { axis } = attrs;
|
|
if (axis < 0) {
|
|
axis += value.shape.length;
|
|
}
|
|
const x = value;
|
|
const xRank = x.shape.length;
|
|
const num = value.shape[axis];
|
|
const outShape = new Array(xRank - 1);
|
|
let outIndex = 0;
|
|
for (let i = 0; i < xRank; i++) {
|
|
if (i !== axis) {
|
|
outShape[outIndex++] = x.shape[i];
|
|
}
|
|
}
|
|
const toDispose = [];
|
|
const begin = new Array(xRank).fill(0);
|
|
const size = x.shape.slice();
|
|
size[axis] = 1;
|
|
const res = new Array(num);
|
|
for (let i = 0; i < res.length; i++) {
|
|
begin[axis] = i;
|
|
const sliced = slice({ inputs: { x }, backend: backend2, attrs: { begin, size } });
|
|
const reshaped = reshape$1({ inputs: { x: sliced }, backend: backend2, attrs: { shape: outShape } });
|
|
res[i] = reshaped;
|
|
toDispose.push(sliced);
|
|
}
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return res;
|
|
}
|
|
const unpackConfig$1 = {
|
|
kernelName: Unpack,
|
|
backendName: "webgl",
|
|
kernelFunc: unpack$1
|
|
};
|
|
class SegmentOpProgram {
|
|
constructor(segOpInfo, segOpType) {
|
|
this.variableNames = ["x", "segmentIds"];
|
|
const windowSize = segOpInfo.windowSize;
|
|
const batchSize = segOpInfo.batchSize;
|
|
const inSize = segOpInfo.inSize;
|
|
const numSegments = segOpInfo.numSegments;
|
|
const outSize = numSegments * Math.ceil(inSize / windowSize);
|
|
this.outputShape = [batchSize, outSize];
|
|
const initializationValue = "0.0";
|
|
const returnValue = `sumValue`;
|
|
const windowSizeNearestVec4 = Math.floor(windowSize / 4) * 4;
|
|
const windowSizeVec4Remainder = windowSize % 4;
|
|
const updateSnippet = `
|
|
sumValue += dot(values, segFilter);
|
|
`;
|
|
let checkValueOutOfBounds = "";
|
|
if (inSize % windowSize > 0) {
|
|
checkValueOutOfBounds = `
|
|
if (inIdx < 0 || inIdx >= ${inSize}) {
|
|
return initializationValue;
|
|
}
|
|
`;
|
|
}
|
|
let checkSegmentIdOutOfBounds = "";
|
|
if (inSize % windowSize > 0) {
|
|
checkSegmentIdOutOfBounds = `
|
|
if (inIdx < 0 || inIdx >= ${inSize}) {
|
|
return -1.0;
|
|
}
|
|
`;
|
|
}
|
|
this.userCode = `
|
|
const float initializationValue = ${initializationValue};
|
|
|
|
float getValue(int batch, int inIdx) {
|
|
${checkValueOutOfBounds}
|
|
return getX(batch, inIdx);
|
|
}
|
|
|
|
float getSegmentIdAtIndex(int inIdx) {
|
|
${checkSegmentIdOutOfBounds}
|
|
return getSegmentIds(inIdx);
|
|
}
|
|
|
|
void main() {
|
|
ivec2 coords = getOutputCoords();
|
|
int batch = coords[0];
|
|
int outIdx = coords[1];
|
|
int inOffset = int(floor(float(outIdx) / float(
|
|
${numSegments})) * float(${windowSize}));
|
|
int currentSeg = int(mod(float(outIdx), float(${numSegments})));
|
|
|
|
float sumValue = 0.0;
|
|
|
|
for (int i = 0; i < ${windowSizeNearestVec4}; i += 4) {
|
|
int inIdx = inOffset + i;
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
getValue(batch, inIdx + 2),
|
|
getValue(batch, inIdx + 3)
|
|
);
|
|
|
|
vec4 segFilter = vec4(
|
|
int(getSegmentIdAtIndex(inIdx)) == currentSeg ? 1 : 0,
|
|
int(getSegmentIdAtIndex(inIdx + 1)) == currentSeg ? 1 : 0,
|
|
int(getSegmentIdAtIndex(inIdx + 2)) == currentSeg ? 1 : 0,
|
|
int(getSegmentIdAtIndex(inIdx + 3)) == currentSeg ? 1 : 0
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
|
|
int inIdx = inOffset + ${windowSizeNearestVec4};
|
|
if (${windowSizeVec4Remainder === 1}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
initializationValue,
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
int inIdxSeg = int(getSegmentIdAtIndex(inIdx));
|
|
|
|
vec4 segFilter = vec4(
|
|
int(getSegmentIdAtIndex(inIdx)) == currentSeg ? 1 : 0,
|
|
0,
|
|
0,
|
|
0
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${windowSizeVec4Remainder === 2}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
initializationValue,
|
|
initializationValue
|
|
);
|
|
|
|
vec4 segFilter = vec4(
|
|
int(getSegmentIdAtIndex(inIdx)) == currentSeg ? 1 : 0,
|
|
int(getSegmentIdAtIndex(inIdx + 1)) == currentSeg ? 1 : 0,
|
|
0,
|
|
0
|
|
);
|
|
|
|
${updateSnippet}
|
|
} else if (${windowSizeVec4Remainder === 3}) {
|
|
vec4 values = vec4(
|
|
getValue(batch, inIdx),
|
|
getValue(batch, inIdx + 1),
|
|
getValue(batch, inIdx + 2),
|
|
initializationValue
|
|
);
|
|
|
|
vec4 segFilter = vec4(
|
|
int(getSegmentIdAtIndex(inIdx)) == currentSeg ? 1 : 0,
|
|
int(getSegmentIdAtIndex(inIdx + 1)) == currentSeg ? 1 : 0,
|
|
int(getSegmentIdAtIndex(inIdx + 2)) == currentSeg ? 1 : 0,
|
|
0
|
|
);
|
|
|
|
${updateSnippet}
|
|
}
|
|
setOutput(${returnValue});
|
|
}
|
|
`;
|
|
}
|
|
}
|
|
function unsortedSegmentSum$1(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, segmentIds } = inputs;
|
|
const { numSegments } = attrs;
|
|
const xRank = x.shape.length;
|
|
const toDispose = [];
|
|
let axis = 0;
|
|
const permutation = getAxesPermutation([axis], xRank);
|
|
let permutedX = x;
|
|
if (permutation != null) {
|
|
permutedX = transpose({ inputs: { x }, backend: backend2, attrs: { perm: permutation } });
|
|
toDispose.push(permutedX);
|
|
axis = getInnerMostAxes(1, xRank)[0];
|
|
}
|
|
const outShape = computeOutShape(permutedX.shape, axis, numSegments);
|
|
const inSize = sizeFromShape([permutedX.shape[axis]]);
|
|
const a2D = reshape$1({ inputs: { x: permutedX }, backend: backend2, attrs: { shape: [-1, inSize] } });
|
|
toDispose.push(a2D);
|
|
const outputDType = sumOutType(x.dtype);
|
|
const segOpCompute = (x2, segOpType, segmentIds2, dtype, numSegments2) => {
|
|
const batchSize = x2.shape[0];
|
|
const inSize2 = x2.shape[1];
|
|
const windowSize = segOpComputeOptimalWindowSize(inSize2, numSegments2);
|
|
const segOpInfo = { windowSize, inSize: inSize2, batchSize, numSegments: numSegments2 };
|
|
const program = new SegmentOpProgram(segOpInfo, segOpType);
|
|
const output = backend2.compileAndRun(program, [x2, segmentIds2], dtype);
|
|
toDispose.push(output);
|
|
if (output.shape[1] === numSegments2) {
|
|
return output;
|
|
}
|
|
const rangeInfo = range$1({
|
|
backend: backend2,
|
|
attrs: { start: 0, stop: numSegments2, step: 1, dtype: "float32" }
|
|
});
|
|
const tileInfo = tile$1({
|
|
inputs: { x: rangeInfo },
|
|
backend: backend2,
|
|
attrs: { reps: [inSize2 / windowSize] }
|
|
});
|
|
toDispose.push(rangeInfo);
|
|
toDispose.push(tileInfo);
|
|
const result2 = segOpCompute(output, segOpType, tileInfo, dtype, numSegments2);
|
|
return result2;
|
|
};
|
|
const segOpResult = segOpCompute(a2D, "unsortedSegmentSum", segmentIds, outputDType, numSegments);
|
|
const reshaped = reshape$1({ inputs: { x: segOpResult }, backend: backend2, attrs: { shape: outShape } });
|
|
let result = reshaped;
|
|
if (permutation != null) {
|
|
toDispose.push(reshaped);
|
|
const perm = getUndoAxesPermutation(permutation);
|
|
result = transpose({ inputs: { x: result }, backend: backend2, attrs: { perm } });
|
|
}
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
}
|
|
const unsortedSegmentSumConfig$1 = {
|
|
kernelName: UnsortedSegmentSum,
|
|
backendName: "webgl",
|
|
kernelFunc: unsortedSegmentSum$1
|
|
};
|
|
const kernelConfigs$1 = [
|
|
_fusedMatMulConfig$1,
|
|
absConfig,
|
|
acosConfig$1,
|
|
acoshConfig$1,
|
|
addConfig,
|
|
addNConfig$1,
|
|
allConfig$1,
|
|
anyConfig$1,
|
|
argMaxConfig$1,
|
|
argMinConfig$1,
|
|
asinConfig$1,
|
|
asinhConfig$1,
|
|
atanConfig$1,
|
|
atan2Config$1,
|
|
atanhConfig$1,
|
|
avgPoolConfig$1,
|
|
avgPool3DConfig$1,
|
|
avgPool3DGradConfig$1,
|
|
avgPoolGradConfig$1,
|
|
batchMatMulConfig$1,
|
|
batchNormConfig$1,
|
|
batchToSpaceNDConfig$1,
|
|
bincountConfig$1,
|
|
bitwiseAndConfig,
|
|
broadcastArgsConfig$1,
|
|
castConfig,
|
|
ceilConfig,
|
|
clipByValueConfig$1,
|
|
complexConfig,
|
|
complexAbsConfig$1,
|
|
concatConfig$1,
|
|
conv2DConfig$1,
|
|
conv2DBackpropFilterConfig$1,
|
|
conv2DBackpropInputConfig$1,
|
|
conv3DConfig$1,
|
|
conv3DBackpropFilterV2Config$1,
|
|
conv3DBackpropInputConfig,
|
|
cosConfig$1,
|
|
coshConfig$1,
|
|
cropAndResizeConfig$1,
|
|
cumprodConfig$1,
|
|
cumsumConfig$1,
|
|
denseBincountConfig$1,
|
|
depthToSpaceConfig$1,
|
|
depthwiseConv2dNativeConfig$1,
|
|
depthwiseConv2dNativeBackpropFilterConfig$1,
|
|
depthwiseConv2dNativeBackpropInputConfig$1,
|
|
diagConfig$1,
|
|
dilation2DConfig$1,
|
|
einsumConfig$1,
|
|
eluConfig$1,
|
|
eluGradConfig$1,
|
|
equalConfig,
|
|
erfConfig$1,
|
|
expConfig,
|
|
expandDimsConfig$1,
|
|
expm1Config,
|
|
fftConfig$1,
|
|
fillConfig$1,
|
|
flipLeftRightConfig$1,
|
|
floorConfig,
|
|
floorDivConfig,
|
|
fromPixelsConfig,
|
|
fusedConv2DConfig$1,
|
|
fusedDepthwiseConv2DConfig$1,
|
|
gatherNdConfig$1,
|
|
gatherV2Config$1,
|
|
greaterConfig,
|
|
greaterEqualConfig,
|
|
identityConfig,
|
|
ifftConfig$1,
|
|
imagConfig$1,
|
|
isFiniteConfig$1,
|
|
isInfConfig$1,
|
|
isNaNConfig$1,
|
|
leakyReluConfig$1,
|
|
lessConfig,
|
|
lessEqualConfig,
|
|
linSpaceConfig$1,
|
|
logConfig,
|
|
log1pConfig$1,
|
|
logicalAndConfig$1,
|
|
logicalNotConfig$1,
|
|
logicalOrConfig$1,
|
|
LRNConfig$1,
|
|
LRNGradConfig$1,
|
|
maxConfig$1,
|
|
maximumConfig,
|
|
maxPoolConfig$1,
|
|
maxPool3DConfig$1,
|
|
maxPool3DGradConfig$1,
|
|
maxPoolGradConfig$1,
|
|
maxPoolWithArgmaxConfig$1,
|
|
meanConfig$1,
|
|
minConfig$1,
|
|
minimumConfig,
|
|
mirrorPadConfig$1,
|
|
modConfig$1,
|
|
multinomialConfig$1,
|
|
multiplyConfig,
|
|
negConfig,
|
|
nonMaxSuppressionV3Config$1,
|
|
nonMaxSuppressionV4Config$1,
|
|
nonMaxSuppressionV5Config$1,
|
|
notEqualConfig,
|
|
oneHotConfig$1,
|
|
onesLikeConfig$1,
|
|
packConfig$1,
|
|
padV2Config$1,
|
|
powConfig$1,
|
|
preluConfig$1,
|
|
prodConfig,
|
|
raggedGatherConfig$1,
|
|
raggedRangeConfig$1,
|
|
raggedTensorToTensorConfig$1,
|
|
rangeConfig$1,
|
|
realConfig,
|
|
realDivConfig$1,
|
|
reciprocalConfig$1,
|
|
reluConfig$1,
|
|
relu6Config$1,
|
|
reshapeConfig$1,
|
|
resizeBilinearConfig$1,
|
|
resizeBilinearGradConfig$1,
|
|
resizeNearestNeighborConfig$1,
|
|
resizeNearestNeighborGradConfig$1,
|
|
reverseConfig$1,
|
|
rotateWithOffsetConfig$1,
|
|
roundConfig$1,
|
|
rsqrtConfig,
|
|
scatterNdConfig$1,
|
|
searchSortedConfig$1,
|
|
selectConfig$1,
|
|
seluConfig$1,
|
|
sigmoidConfig,
|
|
signConfig$1,
|
|
sinConfig$1,
|
|
sinhConfig$1,
|
|
sliceConfig,
|
|
softmaxConfig$1,
|
|
softplusConfig$1,
|
|
spaceToBatchNDConfig$1,
|
|
sparseFillEmptyRowsConfig$1,
|
|
sparseReshapeConfig$1,
|
|
sparseSegmentMeanConfig$1,
|
|
sparseSegmentSumConfig$1,
|
|
sparseToDenseConfig$1,
|
|
splitVConfig$1,
|
|
sqrtConfig,
|
|
squareConfig$1,
|
|
squaredDifferenceConfig,
|
|
staticRegexReplaceConfig,
|
|
stepConfig$1,
|
|
stridedSliceConfig$1,
|
|
stringNGramsConfig$1,
|
|
stringSplitConfig$1,
|
|
stringToHashBucketFastConfig$1,
|
|
subConfig,
|
|
sumConfig$1,
|
|
tanConfig$1,
|
|
tanhConfig$1,
|
|
tensorScatterUpdateConfig$1,
|
|
tileConfig$1,
|
|
topKConfig$1,
|
|
transformConfig$1,
|
|
transposeConfig,
|
|
uniqueConfig$1,
|
|
unpackConfig$1,
|
|
unsortedSegmentSumConfig$1,
|
|
zerosLikeConfig$1
|
|
];
|
|
for (const kernelConfig of kernelConfigs$1) {
|
|
registerKernel(kernelConfig);
|
|
}
|
|
const whereImpl = whereImpl$2;
|
|
class MathBackendCPU extends KernelBackend {
|
|
nextDataId() {
|
|
return MathBackendCPU.nextDataId++;
|
|
}
|
|
constructor() {
|
|
super();
|
|
this.blockSize = 48;
|
|
this.firstUse = true;
|
|
this.data = new DataStorage(this, engine());
|
|
}
|
|
write(values, shape, dtype) {
|
|
if (this.firstUse) {
|
|
this.firstUse = false;
|
|
if (env().get("IS_NODE")) {
|
|
warn("\n============================\nHi, looks like you are running TensorFlow.js in Node.js. To speed things up dramatically, install our node backend, visit https://github.com/tensorflow/tfjs-node for more details. \n============================");
|
|
}
|
|
}
|
|
const dataId = { id: this.nextDataId() };
|
|
this.data.set(dataId, { values, dtype, refCount: 1 });
|
|
return dataId;
|
|
}
|
|
/**
|
|
* Create a data bucket in cpu backend.
|
|
* @param shape Shape of the `TensorInfo`.
|
|
* @param dtype DType of the `TensorInfo`.
|
|
* @param values The value of the `TensorInfo` stored as a flattened array.
|
|
*/
|
|
makeTensorInfo(shape, dtype, values) {
|
|
let outId;
|
|
if (dtype === "string" && values != null && values.length > 0 && isString(values[0])) {
|
|
const encodedValues = values.map((d) => encodeString(d));
|
|
outId = this.write(encodedValues, shape, dtype);
|
|
} else {
|
|
outId = this.write(values, shape, dtype);
|
|
}
|
|
return { dataId: outId, shape, dtype };
|
|
}
|
|
/** Return refCount of a `TensorData`. */
|
|
refCount(dataId) {
|
|
if (this.data.has(dataId)) {
|
|
const tensorData = this.data.get(dataId);
|
|
return tensorData.refCount;
|
|
}
|
|
return 0;
|
|
}
|
|
/** Increase refCount of a `TensorData`. */
|
|
incRef(dataId) {
|
|
const tensorData = this.data.get(dataId);
|
|
tensorData.refCount++;
|
|
}
|
|
/** Decrease refCount of a `TensorData`. */
|
|
decRef(dataId) {
|
|
if (this.data.has(dataId)) {
|
|
const tensorData = this.data.get(dataId);
|
|
tensorData.refCount--;
|
|
}
|
|
}
|
|
move(dataId, values, shape, dtype, refCount) {
|
|
this.data.set(dataId, { values, dtype, refCount });
|
|
}
|
|
numDataIds() {
|
|
return this.data.numDataIds();
|
|
}
|
|
async read(dataId) {
|
|
return this.readSync(dataId);
|
|
}
|
|
readSync(dataId) {
|
|
const { dtype, complexTensorInfos } = this.data.get(dataId);
|
|
if (dtype === "complex64") {
|
|
const realValues = this.readSync(complexTensorInfos.real.dataId);
|
|
const imagValues = this.readSync(complexTensorInfos.imag.dataId);
|
|
return mergeRealAndImagArrays(realValues, imagValues);
|
|
}
|
|
return convertBackendValuesAndArrayBuffer(this.data.get(dataId).values, dtype);
|
|
}
|
|
bufferSync(t2) {
|
|
const data = this.readSync(t2.dataId);
|
|
if (t2.dtype === "string") {
|
|
try {
|
|
const strings = data.map((d) => decodeString(d));
|
|
return buffer(t2.shape, t2.dtype, strings);
|
|
} catch (_a) {
|
|
throw new Error("Failed to decode encoded string bytes into utf-8");
|
|
}
|
|
}
|
|
return buffer(t2.shape, t2.dtype, data);
|
|
}
|
|
makeOutput(values, shape, dtype) {
|
|
return engine().makeTensorFromTensorInfo(this.makeTensorInfo(shape, dtype, values), this);
|
|
}
|
|
/**
|
|
* Dispose the memory if the dataId has 0 refCount. Return true if the memory
|
|
* is released or memory is not managed in this backend, false if memory is
|
|
* not cleared.
|
|
* @param dataId
|
|
* @oaram force Optional, remove the data regardless of refCount
|
|
*/
|
|
disposeData(dataId, force = false) {
|
|
if (this.data.has(dataId)) {
|
|
this.data.get(dataId).refCount--;
|
|
if (!force && this.data.get(dataId).refCount > 0) {
|
|
return false;
|
|
}
|
|
const { complexTensorInfos } = this.data.get(dataId);
|
|
if (complexTensorInfos != null) {
|
|
this.disposeData(complexTensorInfos.real.dataId, true);
|
|
this.disposeData(complexTensorInfos.imag.dataId, true);
|
|
}
|
|
this.data.delete(dataId);
|
|
}
|
|
return true;
|
|
}
|
|
disposeIntermediateTensorInfo(tensorInfo) {
|
|
this.disposeData(tensorInfo.dataId);
|
|
}
|
|
async time(f) {
|
|
const start = now();
|
|
f();
|
|
const kernelMs = now() - start;
|
|
return { kernelMs };
|
|
}
|
|
memory() {
|
|
return {
|
|
// Unreliable due to automatic gc. The numbers above are cumulative.
|
|
unreliable: true,
|
|
reasons: ["The reported memory is an upper bound. Due to automatic garbage collection, the true allocated memory may be less."]
|
|
};
|
|
}
|
|
where(condition) {
|
|
assertNotComplex([condition], "where");
|
|
const condVals = this.readSync(condition.dataId);
|
|
return whereImpl(condition.shape, condVals);
|
|
}
|
|
dispose() {
|
|
}
|
|
floatPrecision() {
|
|
return 32;
|
|
}
|
|
/** Returns the smallest representable number. */
|
|
epsilon() {
|
|
return super.epsilon();
|
|
}
|
|
}
|
|
MathBackendCPU.nextDataId = 0;
|
|
registerBackend(
|
|
"cpu",
|
|
() => new MathBackendCPU(),
|
|
1
|
|
/* priority */
|
|
);
|
|
const elu = unaryKernelFunc$1(Elu, (xi) => xi >= 0 ? xi : Math.exp(xi) - 1);
|
|
const eluConfig = {
|
|
kernelName: Elu,
|
|
backendName: "cpu",
|
|
kernelFunc: elu
|
|
};
|
|
function leakyRelu(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { alpha } = attrs;
|
|
assertNotComplex([x], "leakyRelu");
|
|
const xSize = sizeFromShape(x.shape);
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const outVals = getTypedArrayFromDType("float32", xSize);
|
|
for (let i = 0; i < xVals.length; i++) {
|
|
outVals[i] = xVals[i] < 0 ? alpha * xVals[i] : xVals[i];
|
|
}
|
|
return backend2.makeTensorInfo(x.shape, "float32", outVals);
|
|
}
|
|
const leakyReluConfig = {
|
|
kernelName: LeakyRelu,
|
|
backendName: "cpu",
|
|
kernelFunc: leakyRelu
|
|
};
|
|
const preluImpl = createSimpleBinaryKernelImpl((xValue, aValue) => xValue < 0 ? aValue * xValue : xValue);
|
|
function prelu(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x, alpha } = inputs;
|
|
assertNotComplex([x, alpha], "prelu");
|
|
const aVals = backend2.data.get(x.dataId).values;
|
|
const bVals = backend2.data.get(alpha.dataId).values;
|
|
const [resultData, resultShape] = preluImpl(x.shape, alpha.shape, aVals, bVals, "float32");
|
|
return backend2.makeTensorInfo(resultShape, "float32", resultData);
|
|
}
|
|
const preluConfig = {
|
|
kernelName: Prelu,
|
|
backendName: "cpu",
|
|
kernelFunc: prelu
|
|
};
|
|
const relu = unaryKernelFunc$1(Relu, (xi) => Math.max(0, xi));
|
|
const reluConfig = {
|
|
kernelName: Relu,
|
|
backendName: "cpu",
|
|
kernelFunc: relu
|
|
};
|
|
const relu6 = unaryKernelFunc$1(Relu6, (xi) => Math.min(Math.max(0, xi), 6));
|
|
const relu6Config = {
|
|
kernelName: Relu6,
|
|
backendName: "cpu",
|
|
kernelFunc: relu6
|
|
};
|
|
function applyActivation(backend2, x, activation, preluActivationWeights, leakyreluAlpha) {
|
|
if (activation === "linear") {
|
|
return identity$1({ inputs: { x }, backend: backend2 });
|
|
} else if (activation === "relu") {
|
|
return relu({ inputs: { x }, backend: backend2 });
|
|
} else if (activation === "elu") {
|
|
return elu({ inputs: { x }, backend: backend2 });
|
|
} else if (activation === "relu6") {
|
|
return relu6({ inputs: { x }, backend: backend2 });
|
|
} else if (activation === "prelu") {
|
|
return prelu({ inputs: { x, alpha: preluActivationWeights }, backend: backend2 });
|
|
} else if (activation === "leakyrelu") {
|
|
return leakyRelu({ inputs: { x }, backend: backend2, attrs: { alpha: leakyreluAlpha } });
|
|
} else if (activation === "sigmoid") {
|
|
return sigmoid$1({ inputs: { x }, backend: backend2 });
|
|
}
|
|
throw new Error(`Activation ${activation} has not been implemented for the CPU backend.`);
|
|
}
|
|
function reshape(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { shape } = attrs;
|
|
const xSize = sizeFromShape(x.shape);
|
|
const $shape = inferFromImplicitShape(shape, xSize);
|
|
const $xSize = sizeFromShape($shape);
|
|
assert(xSize === $xSize, () => `The new shape (${$shape}) has ${$xSize} elements and the old shape (${x.shape}) has ${xSize} elements. The new shape and old shape must have the same number of elements.`);
|
|
backend2.incRef(x.dataId);
|
|
const xData = backend2.data.get(x.dataId);
|
|
if (xData.complexTensorInfos != null) {
|
|
const real2 = xData.complexTensorInfos.real;
|
|
const imag2 = xData.complexTensorInfos.imag;
|
|
real2.shape = $shape;
|
|
imag2.shape = $shape;
|
|
}
|
|
return { dataId: x.dataId, shape: $shape, dtype: x.dtype };
|
|
}
|
|
const reshapeConfig = {
|
|
kernelName: Reshape,
|
|
backendName: "cpu",
|
|
kernelFunc: reshape
|
|
};
|
|
function batchMatMul(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { a, b } = inputs;
|
|
const { transposeA, transposeB } = attrs;
|
|
assertNotComplex([a, b], "matMul");
|
|
const aRank = a.shape.length;
|
|
const bRank = b.shape.length;
|
|
const innerShapeA = transposeA ? a.shape[aRank - 2] : a.shape[aRank - 1];
|
|
const innerShapeB = transposeB ? b.shape[bRank - 1] : b.shape[bRank - 2];
|
|
const outerShapeA = transposeA ? a.shape[aRank - 1] : a.shape[aRank - 2];
|
|
const outerShapeB = transposeB ? b.shape[bRank - 2] : b.shape[bRank - 1];
|
|
const outerDimsA = a.shape.slice(0, -2);
|
|
const outerDimsB = b.shape.slice(0, -2);
|
|
const batchDimA = sizeFromShape(outerDimsA);
|
|
const batchDimB = sizeFromShape(outerDimsB);
|
|
const outShapeOuterDims = assertAndGetBroadcastShape(a.shape.slice(0, -2), b.shape.slice(0, -2));
|
|
const outShape = outShapeOuterDims.concat([outerShapeA, outerShapeB]);
|
|
assert(innerShapeA === innerShapeB, () => `Error in matMul: inner shapes (${innerShapeA}) and (${innerShapeB}) of Tensors with shapes ${a.shape} and ${b.shape} and transposeA=${transposeA} and transposeB=${transposeB} must match.`);
|
|
const a3dShape = transposeA ? [batchDimA, innerShapeA, outerShapeA] : [batchDimA, outerShapeA, innerShapeA];
|
|
const b3dShape = transposeB ? [batchDimB, outerShapeB, innerShapeB] : [batchDimB, innerShapeB, outerShapeB];
|
|
const a3d = reshape({ inputs: { x: a }, backend: backend2, attrs: { shape: a3dShape } });
|
|
const b3d = reshape({ inputs: { x: b }, backend: backend2, attrs: { shape: b3dShape } });
|
|
const sharedDim = transposeA ? a3d.shape[1] : a3d.shape[2];
|
|
const leftDim = transposeA ? a3d.shape[2] : a3d.shape[1];
|
|
const rightDim = transposeB ? b3d.shape[1] : b3d.shape[2];
|
|
const batchDim = Math.max(batchDimA, batchDimB);
|
|
const a3dValues = backend2.data.get(a3d.dataId).values;
|
|
const b3dValues = backend2.data.get(b3d.dataId).values;
|
|
const a3dStrides = computeStrides(a3d.shape);
|
|
const b3dStrides = computeStrides(b3d.shape);
|
|
const [aBatch, aOuterStep, aInnerStep] = transposeA ? [a3dStrides[0], 1, a3dStrides[1]] : [a3dStrides[0], a3dStrides[1], 1];
|
|
const [bInnerStep, bOuterStep, bBatch] = transposeB ? [1, b3dStrides[1], b3dStrides[0]] : [b3dStrides[1], 1, b3dStrides[0]];
|
|
const size = leftDim * rightDim;
|
|
const result = buffer([batchDim, leftDim, rightDim], a3d.dtype);
|
|
const resVals = result.values;
|
|
const blockSize = backend2.blockSize;
|
|
for (let bi = 0; bi < batchDim; bi++) {
|
|
const batchIndexA = bi % batchDimA;
|
|
const batchIndexB = bi % batchDimB;
|
|
for (let i0 = 0; i0 < leftDim; i0 += blockSize) {
|
|
const iBlock = Math.min(i0 + blockSize, leftDim);
|
|
for (let j0 = 0; j0 < rightDim; j0 += blockSize) {
|
|
const jBlock = Math.min(j0 + blockSize, rightDim);
|
|
for (let k02 = 0; k02 < sharedDim; k02 += blockSize) {
|
|
const kBlock = Math.min(k02 + blockSize, sharedDim);
|
|
for (let i = i0; i < iBlock; i++) {
|
|
for (let j2 = j0; j2 < jBlock; j2++) {
|
|
let sum2 = 0;
|
|
for (let k = k02; k < kBlock; k++) {
|
|
const aVal = (
|
|
// tslint:disable-next-line: max-line-length
|
|
a3dValues[batchIndexA * aBatch + i * aOuterStep + k * aInnerStep]
|
|
);
|
|
const bVal = (
|
|
// tslint:disable-next-line: max-line-length
|
|
b3dValues[k * bInnerStep + j2 * bOuterStep + batchIndexB * bBatch]
|
|
);
|
|
sum2 += aVal * bVal;
|
|
}
|
|
resVals[bi * size + (i * rightDim + j2)] += sum2;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(a3d);
|
|
backend2.disposeIntermediateTensorInfo(b3d);
|
|
return backend2.makeTensorInfo(outShape, result.dtype, result.values);
|
|
}
|
|
const batchMatMulConfig = {
|
|
kernelName: BatchMatMul,
|
|
backendName: "cpu",
|
|
kernelFunc: batchMatMul
|
|
};
|
|
function _fusedMatMul(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { a, b, bias, preluActivationWeights } = inputs;
|
|
const { transposeA, transposeB, activation, leakyreluAlpha } = attrs;
|
|
let current;
|
|
let addRes;
|
|
let activationRes;
|
|
const intermediates = [];
|
|
const matMulRes = batchMatMul({ inputs: { a, b }, attrs: { transposeA, transposeB }, backend: backend2 });
|
|
current = matMulRes;
|
|
if (bias) {
|
|
addRes = add({ inputs: { a: current, b: bias }, backend: backend2 });
|
|
intermediates.push(current);
|
|
current = addRes;
|
|
}
|
|
if (activation) {
|
|
activationRes = applyActivation(backend2, current, activation, preluActivationWeights, leakyreluAlpha);
|
|
intermediates.push(current);
|
|
current = activationRes;
|
|
}
|
|
for (const i of intermediates) {
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
}
|
|
return current;
|
|
}
|
|
const _fusedMatMulConfig = {
|
|
kernelName: _FusedMatMul,
|
|
backendName: "cpu",
|
|
kernelFunc: _fusedMatMul
|
|
};
|
|
const acos = unaryKernelFunc$1(Acos, (xi) => Math.acos(xi));
|
|
const acosConfig = {
|
|
kernelName: Acos,
|
|
backendName: "cpu",
|
|
kernelFunc: acos
|
|
};
|
|
const acosh = unaryKernelFunc$1(Acosh, (xi) => Math.acosh(xi));
|
|
const acoshConfig = {
|
|
kernelName: Acosh,
|
|
backendName: "cpu",
|
|
kernelFunc: acosh
|
|
};
|
|
function addN(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const tensors = inputs;
|
|
assertNotComplex(inputs, "addN");
|
|
const vals = tensors.map((t2) => backend2.data.get(t2.dataId).values);
|
|
const outBuf = buffer(tensors[0].shape, tensors[0].dtype);
|
|
const outVals = outBuf.values;
|
|
for (let i = 0; i < tensors.length; i++) {
|
|
const currVals = vals[i];
|
|
for (let j2 = 0; j2 < outVals.length; j2++) {
|
|
outVals[j2] += currVals[j2];
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(outBuf.shape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const addNConfig = {
|
|
kernelName: AddN,
|
|
backendName: "cpu",
|
|
kernelFunc: addN
|
|
};
|
|
function all(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
assertNotComplex(x, "all");
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
if (permutedAxes != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, x.shape.length);
|
|
}
|
|
assertAxesAreInnerMostDims("all", axes, $x.shape.length);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes($x.shape, axes);
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const vals = makeZerosTypedArray(sizeFromShape(outShape), $x.dtype);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let all2 = aVals[offset];
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
const value = aVals[offset + j2];
|
|
all2 = all2 && value;
|
|
}
|
|
vals[i] = all2;
|
|
}
|
|
if (permutedAxes != null) {
|
|
backend2.disposeIntermediateTensorInfo($x);
|
|
}
|
|
const result = backend2.makeTensorInfo(outShape, $x.dtype, vals);
|
|
if (keepDims) {
|
|
const expandedShape = expandShapeToKeepDim(outShape, origAxes);
|
|
const reshapedResult = reshape({ inputs: { x: result }, backend: backend2, attrs: { shape: expandedShape } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return reshapedResult;
|
|
}
|
|
return result;
|
|
}
|
|
const allConfig = {
|
|
kernelName: All,
|
|
backendName: "cpu",
|
|
kernelFunc: all
|
|
};
|
|
function any(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
assertNotComplex(x, "any");
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
if (permutedAxes != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, x.shape.length);
|
|
}
|
|
assertAxesAreInnerMostDims("any", axes, $x.shape.length);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes($x.shape, axes);
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const vals = makeZerosTypedArray(sizeFromShape(outShape), $x.dtype);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let anyVal = aVals[offset];
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
const value = aVals[offset + j2];
|
|
anyVal = anyVal || value;
|
|
}
|
|
vals[i] = anyVal;
|
|
}
|
|
if (permutedAxes != null) {
|
|
backend2.disposeIntermediateTensorInfo($x);
|
|
}
|
|
const result = backend2.makeTensorInfo(outShape, $x.dtype, vals);
|
|
if (keepDims) {
|
|
const expandedShape = expandShapeToKeepDim(outShape, origAxes);
|
|
const reshapedResult = reshape({ inputs: { x: result }, backend: backend2, attrs: { shape: expandedShape } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return reshapedResult;
|
|
}
|
|
return result;
|
|
}
|
|
const anyConfig = {
|
|
kernelName: Any,
|
|
backendName: "cpu",
|
|
kernelFunc: any
|
|
};
|
|
function argMax(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis } = attrs;
|
|
assertNotComplex(x, "argMax");
|
|
let axes = parseAxisParam(axis, x.shape);
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
const intermediateTensorInfos = [];
|
|
if (permutedAxes != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
intermediateTensorInfos.push($x);
|
|
axes = getInnerMostAxes(axes.length, $x.shape.length);
|
|
}
|
|
axes = [axes[0]];
|
|
assertAxesAreInnerMostDims("argMax", axes, $x.shape.length);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes($x.shape, axes);
|
|
const outSize = sizeFromShape(outShape);
|
|
const vals = makeZerosTypedArray(outSize, "int32");
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let max2 = aVals[offset];
|
|
let maxIndex = 0;
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
const value = aVals[offset + j2];
|
|
if (value > max2) {
|
|
max2 = value;
|
|
maxIndex = j2;
|
|
}
|
|
}
|
|
vals[i] = maxIndex;
|
|
}
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return backend2.makeTensorInfo(outShape, "int32", vals);
|
|
}
|
|
const argMaxConfig = {
|
|
kernelName: ArgMax,
|
|
backendName: "cpu",
|
|
kernelFunc: argMax
|
|
};
|
|
function argMin(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis } = attrs;
|
|
assertNotComplex(x, "argMin");
|
|
let axes = parseAxisParam(axis, x.shape);
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
const intermediateTensorInfos = [];
|
|
if (permutedAxes != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
intermediateTensorInfos.push($x);
|
|
axes = getInnerMostAxes(axes.length, $x.shape.length);
|
|
}
|
|
axes = [axes[0]];
|
|
assertAxesAreInnerMostDims("argMin", axes, $x.shape.length);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes($x.shape, axes);
|
|
const outSize = sizeFromShape(outShape);
|
|
const vals = makeZerosTypedArray(outSize, "int32");
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let min2 = aVals[offset];
|
|
let minIndex = 0;
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
const value = aVals[offset + j2];
|
|
if (value < min2) {
|
|
min2 = value;
|
|
minIndex = j2;
|
|
}
|
|
}
|
|
vals[i] = minIndex;
|
|
}
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return backend2.makeTensorInfo(outShape, "int32", vals);
|
|
}
|
|
const argMinConfig = {
|
|
kernelName: ArgMin,
|
|
backendName: "cpu",
|
|
kernelFunc: argMin
|
|
};
|
|
const asin = unaryKernelFunc$1(Asin, (xi) => Math.asin(xi));
|
|
const asinConfig = {
|
|
kernelName: Asin,
|
|
backendName: "cpu",
|
|
kernelFunc: asin
|
|
};
|
|
const asinh = unaryKernelFunc$1(Asinh, (xi) => Math.asinh(xi));
|
|
const asinhConfig = {
|
|
kernelName: Asinh,
|
|
backendName: "cpu",
|
|
kernelFunc: asinh
|
|
};
|
|
const atan = unaryKernelFunc$1(Atan, (xi) => Math.atan(xi));
|
|
const atanConfig = {
|
|
kernelName: Atan,
|
|
backendName: "cpu",
|
|
kernelFunc: atan
|
|
};
|
|
const atan2Impl = createSimpleBinaryKernelImpl((aValue, bValue) => Math.atan2(aValue, bValue));
|
|
const atan2 = binaryKernelFunc$1(Atan2, atan2Impl);
|
|
const atan2Config = {
|
|
kernelName: Atan2,
|
|
backendName: "cpu",
|
|
kernelFunc: atan2
|
|
};
|
|
const atanh = unaryKernelFunc$1(Atanh, (xi) => Math.atanh(xi));
|
|
const atanhConfig = {
|
|
kernelName: Atanh,
|
|
backendName: "cpu",
|
|
kernelFunc: atanh
|
|
};
|
|
function pool(xValues, xShape, dtype, strides, convInfo, poolType) {
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const initialValue = poolType === "max" ? Number.NEGATIVE_INFINITY : Number.POSITIVE_INFINITY;
|
|
const output = buffer(convInfo.outShape, dtype);
|
|
const outputVals = output.values;
|
|
const outputBatchStrides = convInfo.outShape[1] * convInfo.outShape[2] * convInfo.outShape[3];
|
|
const outputRowStrides = convInfo.outShape[2] * convInfo.outShape[3];
|
|
const outputColStrides = convInfo.outShape[3];
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
const outputBatchOffset = b * outputBatchStrides;
|
|
const inputBatchOffset = b * strides[0];
|
|
for (let d = 0; d < convInfo.inChannels; ++d) {
|
|
for (let yR = 0; yR < convInfo.outHeight; ++yR) {
|
|
const xRCorner = yR * strideHeight - padTop;
|
|
const xRMin = Math.max(0, xRCorner);
|
|
const xRMax = Math.min(convInfo.inHeight, effectiveFilterHeight + xRCorner);
|
|
const outputRowOffset = outputBatchOffset + yR * outputRowStrides;
|
|
for (let yC = 0; yC < convInfo.outWidth; ++yC) {
|
|
const xCCorner = yC * strideWidth - padLeft;
|
|
const xCMin = Math.max(0, xCCorner);
|
|
const xCMax = Math.min(convInfo.inWidth, effectiveFilterWidth + xCCorner);
|
|
let minMaxValue = initialValue;
|
|
let avgValue = 0;
|
|
let count = 0;
|
|
for (let xR = xRMin; xR < xRMax; xR += dilationHeight) {
|
|
const xROffset = inputBatchOffset + xR * strides[1];
|
|
for (let xC = xCMin; xC < xCMax; xC += dilationWidth) {
|
|
const xCOffset = xROffset + xC * strides[2];
|
|
const pixel = xValues[xCOffset + d];
|
|
if (poolType === "max" && pixel > minMaxValue) {
|
|
minMaxValue = pixel;
|
|
} else if (poolType === "avg") {
|
|
avgValue += pixel;
|
|
count++;
|
|
}
|
|
}
|
|
if (isNaN(minMaxValue)) {
|
|
break;
|
|
}
|
|
}
|
|
const outputOffset = outputRowOffset + yC * outputColStrides + d;
|
|
outputVals[outputOffset] = poolType === "avg" ? avgValue / count : minMaxValue;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return output;
|
|
}
|
|
function maxPoolPositions(xValues, xShape, dtype, convInfo, flattenPositions = false, includeBatchInIndex = false) {
|
|
const maxPositions = buffer(convInfo.outShape, "int32");
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const xBuf = buffer(xShape, dtype, xValues);
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
for (let d = 0; d < convInfo.inChannels; ++d) {
|
|
for (let yR = 0; yR < convInfo.outHeight; ++yR) {
|
|
const xRCorner = yR * strideHeight - padTop;
|
|
let xRMin = xRCorner;
|
|
while (xRMin < 0) {
|
|
xRMin += dilationHeight;
|
|
}
|
|
const xRMax = Math.min(convInfo.inHeight, effectiveFilterHeight + xRCorner);
|
|
for (let yC = 0; yC < convInfo.outWidth; ++yC) {
|
|
const xCCorner = yC * strideWidth - padLeft;
|
|
let xCMin = xCCorner;
|
|
while (xCMin < 0) {
|
|
xCMin += dilationWidth;
|
|
}
|
|
const xCMax = Math.min(convInfo.inWidth, effectiveFilterWidth + xCCorner);
|
|
let maxValue = Number.NEGATIVE_INFINITY;
|
|
let maxPosition = -1;
|
|
for (let xR = xRMin; xR < xRMax; xR += dilationHeight) {
|
|
const wR = xR - xRCorner;
|
|
for (let xC = xCMin; xC < xCMax; xC += dilationWidth) {
|
|
const wC = xC - xCCorner;
|
|
const pixel = xBuf.get(b, xR, xC, d);
|
|
if (pixel > maxValue) {
|
|
maxValue = pixel;
|
|
if (flattenPositions) {
|
|
maxPosition = includeBatchInIndex ? ((b * convInfo.inHeight + xR) * convInfo.inWidth + xC) * convInfo.inChannels + d : (xR * convInfo.inWidth + xC) * convInfo.inChannels + d;
|
|
} else {
|
|
maxPosition = wR * effectiveFilterWidth + wC;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
maxPositions.set(maxPosition, b, yR, yC, d);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return maxPositions;
|
|
}
|
|
function pool3d(xValues, xShape, dtype, strides, convInfo, poolType) {
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = convInfo.padInfo.front;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const initialValue = poolType === "max" ? Number.NEGATIVE_INFINITY : Number.POSITIVE_INFINITY;
|
|
const output = buffer(convInfo.outShape, dtype);
|
|
const outputVals = output.values;
|
|
const outputBatchStrides = convInfo.outShape[1] * convInfo.outShape[2] * convInfo.outShape[3] * convInfo.outShape[4];
|
|
const outputDepthStrides = convInfo.outShape[2] * convInfo.outShape[3] * convInfo.outShape[4];
|
|
const outputRowStrides = convInfo.outShape[3] * convInfo.outShape[4];
|
|
const outputColStrides = convInfo.outShape[4];
|
|
for (let batch = 0; batch < convInfo.batchSize; ++batch) {
|
|
const outputBatchOffset = batch * outputBatchStrides;
|
|
const inputBatchOffset = batch * strides[0];
|
|
for (let channel = 0; channel < convInfo.inChannels; ++channel) {
|
|
for (let yDepth = 0; yDepth < convInfo.outDepth; ++yDepth) {
|
|
const xDepthCorner = yDepth * strideDepth - padFront;
|
|
let xDepthMin = xDepthCorner;
|
|
while (xDepthMin < 0) {
|
|
xDepthMin += dilationDepth;
|
|
}
|
|
const xDepthMax = Math.min(convInfo.inDepth, effectiveFilterDepth + xDepthCorner);
|
|
const outputDepthOffset = outputBatchOffset + yDepth * outputDepthStrides;
|
|
for (let yRow = 0; yRow < convInfo.outHeight; ++yRow) {
|
|
const xRowCorner = yRow * strideHeight - padTop;
|
|
let xRowMin = xRowCorner;
|
|
while (xRowMin < 0) {
|
|
xRowMin += dilationHeight;
|
|
}
|
|
const xRowMax = Math.min(convInfo.inHeight, effectiveFilterHeight + xRowCorner);
|
|
const outputRowOffset = outputDepthOffset + yRow * outputRowStrides;
|
|
for (let yCol = 0; yCol < convInfo.outWidth; ++yCol) {
|
|
const xColCorner = yCol * strideWidth - padLeft;
|
|
let xColMin = xColCorner;
|
|
while (xColMin < 0) {
|
|
xColMin += dilationWidth;
|
|
}
|
|
const xColMax = Math.min(convInfo.inWidth, effectiveFilterWidth + xColCorner);
|
|
const outputColOffset = outputRowOffset + yCol * outputColStrides;
|
|
let minMaxValue = initialValue;
|
|
let avgValue = 0;
|
|
let count = 0;
|
|
for (let xDepth = xDepthMin; xDepth < xDepthMax; xDepth += dilationDepth) {
|
|
const xDepthOffset = inputBatchOffset + xDepth * strides[1];
|
|
for (let xRow = xRowMin; xRow < xRowMax; xRow += dilationHeight) {
|
|
const xRowOffset = xDepthOffset + xRow * strides[2];
|
|
for (let xCol = xColMin; xCol < xColMax; xCol += dilationWidth) {
|
|
const xColOffset = xRowOffset + xCol * strides[3];
|
|
const pixel = xValues[xColOffset + channel];
|
|
if (poolType === "max" && pixel > minMaxValue) {
|
|
minMaxValue = pixel;
|
|
} else if (poolType === "avg") {
|
|
avgValue += pixel;
|
|
count++;
|
|
}
|
|
if (isNaN(minMaxValue)) {
|
|
break;
|
|
}
|
|
}
|
|
if (isNaN(minMaxValue)) {
|
|
break;
|
|
}
|
|
}
|
|
if (isNaN(minMaxValue)) {
|
|
break;
|
|
}
|
|
}
|
|
const outputOffset = outputColOffset + channel;
|
|
outputVals[outputOffset] = poolType === "avg" ? avgValue / Math.max(count, 1) : minMaxValue;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return output;
|
|
}
|
|
function maxPool3dPositions(xBuf, convInfo) {
|
|
const maxPositions = buffer(convInfo.outShape, "int32");
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = convInfo.padInfo.front;
|
|
const padTop = convInfo.padInfo.top;
|
|
const padLeft = convInfo.padInfo.left;
|
|
for (let batch = 0; batch < convInfo.batchSize; ++batch) {
|
|
for (let channel = 0; channel < convInfo.inChannels; ++channel) {
|
|
for (let yDepth = 0; yDepth < convInfo.outDepth; ++yDepth) {
|
|
const xDepthCorner = yDepth * strideDepth - padFront;
|
|
let xDepthMin = xDepthCorner;
|
|
while (xDepthMin < 0) {
|
|
xDepthMin += dilationDepth;
|
|
}
|
|
const xDepthMax = Math.min(convInfo.inDepth, effectiveFilterDepth + xDepthCorner);
|
|
for (let yRow = 0; yRow < convInfo.outHeight; ++yRow) {
|
|
const xRowCorner = yRow * strideHeight - padTop;
|
|
let xRowMin = xRowCorner;
|
|
while (xRowMin < 0) {
|
|
xRowMin += dilationHeight;
|
|
}
|
|
const xRowMax = Math.min(convInfo.inHeight, effectiveFilterHeight + xRowCorner);
|
|
for (let yCol = 0; yCol < convInfo.outWidth; ++yCol) {
|
|
const xColCorner = yCol * strideWidth - padLeft;
|
|
let xColMin = xColCorner;
|
|
while (xColMin < 0) {
|
|
xColMin += dilationWidth;
|
|
}
|
|
const xColMax = Math.min(convInfo.inWidth, effectiveFilterWidth + xColCorner);
|
|
let maxValue = Number.NEGATIVE_INFINITY;
|
|
let maxPosition = -1;
|
|
for (let xDepth = xDepthMin; xDepth < xDepthMax; xDepth += dilationDepth) {
|
|
const wDepth = xDepth - xDepthCorner;
|
|
for (let xRow = xRowMin; xRow < xRowMax; xRow += dilationHeight) {
|
|
const wRow = xRow - xRowCorner;
|
|
for (let xCol = xColMin; xCol < xColMax; xCol += dilationWidth) {
|
|
const wCol = xCol - xColCorner;
|
|
const pixel = xBuf.get(batch, xDepth, xRow, xCol, channel);
|
|
if (pixel >= maxValue) {
|
|
maxValue = pixel;
|
|
maxPosition = wDepth * effectiveFilterHeight * effectiveFilterWidth + wRow * effectiveFilterHeight + wCol;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
maxPositions.set(maxPosition, batch, yDepth, yRow, yCol, channel);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return maxPositions;
|
|
}
|
|
function avgPool(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
assertNotComplex(x, "avgPool");
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const dilations = 1;
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in avgPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode);
|
|
let res;
|
|
if (convInfo.filterWidth === 1 && convInfo.filterHeight === 1 && arraysEqual(convInfo.inShape, convInfo.outShape)) {
|
|
res = identity$1({ inputs: { x }, backend: backend2 });
|
|
} else {
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const strides2 = computeStrides(x.shape);
|
|
const buffer2 = pool(xValues, x.shape, x.dtype, strides2, convInfo, "avg");
|
|
res = backend2.makeTensorInfo(convInfo.outShape, x.dtype, buffer2.values);
|
|
}
|
|
return res;
|
|
}
|
|
const avgPoolConfig = {
|
|
kernelName: AvgPool,
|
|
backendName: "cpu",
|
|
kernelFunc: avgPool
|
|
};
|
|
function avgPool3D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode, dataFormat } = attrs;
|
|
assertNotComplex(x, "avgPool3d");
|
|
const convInfo = computePool3DInfo(x.shape, filterSize, strides, 1, pad2, dimRoundingMode, dataFormat);
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const outBuf = pool3d(xValues, x.shape, x.dtype, computeStrides(x.shape), convInfo, "avg");
|
|
return backend2.makeTensorInfo(outBuf.shape, "float32", outBuf.values);
|
|
}
|
|
const avgPool3DConfig = {
|
|
kernelName: AvgPool3D,
|
|
backendName: "cpu",
|
|
kernelFunc: avgPool3D
|
|
};
|
|
function avgPool3DGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input } = inputs;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
assertNotComplex([dy, input], "avgPool3DGrad");
|
|
const convInfo = computePool3DInfo(input.shape, filterSize, strides, 1, pad2, dimRoundingMode);
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const filterDepth = convInfo.filterDepth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = effectiveFilterDepth - 1 - convInfo.padInfo.front;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const dx = buffer(input.shape, "float32");
|
|
const avgMultiplier = 1 / (filterDepth * filterHeight * filterWidth);
|
|
const dyBuf = backend2.bufferSync(dy);
|
|
for (let batch = 0; batch < convInfo.batchSize; ++batch) {
|
|
for (let channel = 0; channel < convInfo.inChannels; ++channel) {
|
|
for (let dxDepth = 0; dxDepth < convInfo.inDepth; ++dxDepth) {
|
|
for (let dxRow = 0; dxRow < convInfo.inHeight; ++dxRow) {
|
|
for (let dxCol = 0; dxCol < convInfo.inWidth; ++dxCol) {
|
|
const dyDepthCorner = dxDepth - padFront;
|
|
const dyRowCorner = dxRow - padTop;
|
|
const dyColCorner = dxCol - padLeft;
|
|
let dotProd = 0;
|
|
for (let wDepth = 0; wDepth < effectiveFilterDepth; wDepth += dilationDepth) {
|
|
const dyDepth = (dyDepthCorner + wDepth) / strideDepth;
|
|
if (dyDepth < 0 || dyDepth >= convInfo.outDepth || Math.floor(dyDepth) !== dyDepth) {
|
|
continue;
|
|
}
|
|
for (let wRow = 0; wRow < effectiveFilterHeight; wRow += dilationHeight) {
|
|
const dyRow = (dyRowCorner + wRow) / strideHeight;
|
|
if (dyRow < 0 || dyRow >= convInfo.outHeight || Math.floor(dyRow) !== dyRow) {
|
|
continue;
|
|
}
|
|
for (let wCol = 0; wCol < effectiveFilterWidth; wCol += dilationWidth) {
|
|
const dyCol = (dyColCorner + wCol) / strideWidth;
|
|
if (dyCol < 0 || dyCol >= convInfo.outWidth || Math.floor(dyCol) !== dyCol) {
|
|
continue;
|
|
}
|
|
const pixel = dyBuf.get(batch, dyDepth, dyRow, dyCol, channel);
|
|
dotProd += pixel;
|
|
}
|
|
}
|
|
}
|
|
dx.set(dotProd * avgMultiplier, batch, dxDepth, dxRow, dxCol, channel);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const avgPool3DGradConfig = {
|
|
kernelName: AvgPool3DGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: avgPool3DGrad
|
|
};
|
|
function avgPoolGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input } = inputs;
|
|
const x = input;
|
|
assertNotComplex([dy, input], "avgPoolGrad");
|
|
const { filterSize, strides, pad: pad2 } = attrs;
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, 1, pad2);
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const dx = buffer(x.shape, "float32");
|
|
const avgMultiplier = 1 / (filterHeight * filterWidth);
|
|
const dyData = backend2.data.get(dy.dataId).values;
|
|
const dyBuf = buffer(dy.shape, "float32", dyData);
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
for (let d = 0; d < convInfo.inChannels; ++d) {
|
|
for (let dxR = 0; dxR < convInfo.inHeight; ++dxR) {
|
|
for (let dxC = 0; dxC < convInfo.inWidth; ++dxC) {
|
|
const dyRCorner = dxR - padTop;
|
|
const dyCCorner = dxC - padLeft;
|
|
let dotProd = 0;
|
|
for (let wR = 0; wR < effectiveFilterHeight; wR += dilationHeight) {
|
|
const dyR = (dyRCorner + wR) / strideHeight;
|
|
if (dyR < 0 || dyR >= convInfo.outHeight || Math.floor(dyR) !== dyR) {
|
|
continue;
|
|
}
|
|
for (let wC = 0; wC < effectiveFilterWidth; wC += dilationWidth) {
|
|
const dyC = (dyCCorner + wC) / strideWidth;
|
|
if (dyC < 0 || dyC >= convInfo.outWidth || Math.floor(dyC) !== dyC) {
|
|
continue;
|
|
}
|
|
const pixel = dyBuf.get(b, dyR, dyC, d);
|
|
dotProd += pixel;
|
|
}
|
|
}
|
|
dx.set(dotProd * avgMultiplier, b, dxR, dxC, d);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const avgPoolGradConfig = {
|
|
kernelName: AvgPoolGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: avgPoolGrad
|
|
};
|
|
function batchNorm(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, scale: scale2, offset, mean: mean2, variance } = inputs;
|
|
assert(mean2.shape.length === variance.shape.length, () => "Batch normalization gradient requires mean and variance to have equal ranks.");
|
|
assert(offset == null || mean2.shape.length === offset.shape.length, () => "Batch normalization gradient requires mean and offset to have equal ranks.");
|
|
assert(scale2 == null || mean2.shape.length === scale2.shape.length, () => "Batch normalization gradient requires mean and scale to have equal ranks.");
|
|
assertNotComplex([x, mean2, variance, scale2, offset], "batchNorm");
|
|
let { varianceEpsilon } = attrs;
|
|
if (varianceEpsilon == null) {
|
|
varianceEpsilon = 1e-3;
|
|
}
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const mVals = backend2.data.get(mean2.dataId).values;
|
|
const varVals = backend2.data.get(variance.dataId).values;
|
|
const sVals = scale2 ? backend2.data.get(scale2.dataId).values : new Float32Array([1]);
|
|
const offVals = offset ? backend2.data.get(offset.dataId).values : new Float32Array([0]);
|
|
const outVals = new Float32Array(xVals.length);
|
|
const offValsLength = offVals.length;
|
|
const sValsLength = sVals.length;
|
|
const varValsLength = varVals.length;
|
|
const mValsLength = mVals.length;
|
|
let offi = 0;
|
|
let mi = 0;
|
|
let si = 0;
|
|
let vi = 0;
|
|
for (let i = 0; i < xVals.length; ++i) {
|
|
outVals[i] = offVals[offi++] + (xVals[i] - mVals[mi++]) * sVals[si++] / Math.sqrt(varVals[vi++] + varianceEpsilon);
|
|
if (offi >= offValsLength) {
|
|
offi = 0;
|
|
}
|
|
if (mi >= mValsLength) {
|
|
mi = 0;
|
|
}
|
|
if (si >= sValsLength) {
|
|
si = 0;
|
|
}
|
|
if (vi >= varValsLength) {
|
|
vi = 0;
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(x.shape, x.dtype, outVals);
|
|
}
|
|
const batchNormConfig = {
|
|
kernelName: FusedBatchNorm,
|
|
backendName: "cpu",
|
|
kernelFunc: batchNorm
|
|
};
|
|
function batchToSpaceND(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { blockShape, crops } = attrs;
|
|
assertNotComplex([x], "batchToSpaceND");
|
|
const prod2 = blockShape.reduce((a, b) => a * b);
|
|
const reshaped = getReshaped(x.shape, blockShape, prod2);
|
|
const permuted = getPermuted(reshaped.length, blockShape.length);
|
|
const reshapedPermuted = getReshapedPermuted(x.shape, blockShape, prod2);
|
|
const sliceBeginCoords = getSliceBeginCoords(crops, blockShape.length);
|
|
const sliceSize = getSliceSize(reshapedPermuted, crops, blockShape.length);
|
|
const xReshaped = reshape({ inputs: { x }, backend: backend2, attrs: { shape: reshaped } });
|
|
const xTransposed = transpose$1({ inputs: { x: xReshaped }, backend: backend2, attrs: { perm: permuted } });
|
|
const xTransposedReshaped = reshape({ inputs: { x: xTransposed }, backend: backend2, attrs: { shape: reshapedPermuted } });
|
|
const result = slice$1({
|
|
inputs: { x: xTransposedReshaped },
|
|
backend: backend2,
|
|
attrs: { begin: sliceBeginCoords, size: sliceSize }
|
|
});
|
|
backend2.disposeIntermediateTensorInfo(xReshaped);
|
|
backend2.disposeIntermediateTensorInfo(xTransposed);
|
|
backend2.disposeIntermediateTensorInfo(xTransposedReshaped);
|
|
return result;
|
|
}
|
|
const batchToSpaceNDConfig = {
|
|
kernelName: BatchToSpaceND,
|
|
backendName: "cpu",
|
|
kernelFunc: batchToSpaceND
|
|
};
|
|
function bincount(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, weights } = inputs;
|
|
const { size } = attrs;
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const weightsVals = backend2.data.get(weights.dataId).values;
|
|
const outVals = bincountImpl(xVals, weightsVals, weights.dtype, weights.shape, size);
|
|
return backend2.makeTensorInfo([size], weights.dtype, outVals);
|
|
}
|
|
const bincountConfig = {
|
|
kernelName: Bincount,
|
|
backendName: "cpu",
|
|
kernelFunc: bincount
|
|
};
|
|
function broadcastArgs(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { s0, s1 } = inputs;
|
|
const s0Vals = backend2.data.get(s0.dataId).values;
|
|
const s1Vals = backend2.data.get(s1.dataId).values;
|
|
const broadcastShape = assertAndGetBroadcastShape(Array.from(s0Vals), Array.from(s1Vals));
|
|
return backend2.makeTensorInfo([broadcastShape.length], "int32", Int32Array.from(broadcastShape));
|
|
}
|
|
const broadcastArgsConfig = {
|
|
kernelName: BroadcastArgs,
|
|
backendName: "cpu",
|
|
kernelFunc: broadcastArgs
|
|
};
|
|
const clipByValue = unaryKernelFunc$1(ClipByValue, (xi, attrs) => {
|
|
const clipAttrs = attrs;
|
|
if (xi > clipAttrs.clipValueMax) {
|
|
return clipAttrs.clipValueMax;
|
|
}
|
|
return xi < clipAttrs.clipValueMin ? clipAttrs.clipValueMin : xi;
|
|
});
|
|
const clipByValueConfig = {
|
|
kernelName: ClipByValue,
|
|
backendName: "cpu",
|
|
kernelFunc: clipByValue
|
|
};
|
|
const complexAbs = (args) => {
|
|
const { x } = args.inputs;
|
|
const cpuBackend = args.backend;
|
|
const resultValues = new Float32Array(sizeFromShape(x.shape));
|
|
const complexVals = cpuBackend.data.get(x.dataId);
|
|
const real2 = complexVals.complexTensorInfos.real;
|
|
const imag2 = complexVals.complexTensorInfos.imag;
|
|
const realVals = cpuBackend.data.get(real2.dataId).values;
|
|
const imagVals = cpuBackend.data.get(imag2.dataId).values;
|
|
for (let i = 0; i < realVals.length; i++) {
|
|
const real3 = realVals[i];
|
|
const imag3 = imagVals[i];
|
|
resultValues[i] = Math.hypot(real3, imag3);
|
|
}
|
|
return cpuBackend.makeOutput(resultValues, x.shape, "float32");
|
|
};
|
|
const complexAbsConfig = {
|
|
kernelName: ComplexAbs,
|
|
backendName: "cpu",
|
|
kernelFunc: complexAbs
|
|
};
|
|
function imag(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
const imag2 = backend2.data.get(input.dataId).complexTensorInfos.imag;
|
|
const imagVal = backend2.data.get(imag2.dataId).values;
|
|
return backend2.makeTensorInfo(imag2.shape, imag2.dtype, imagVal);
|
|
}
|
|
const imagConfig = {
|
|
kernelName: Imag,
|
|
backendName: "cpu",
|
|
kernelFunc: imag
|
|
};
|
|
function concat(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { axis } = attrs;
|
|
const $axis = parseAxisParam(axis, inputs[0].shape)[0];
|
|
const shapes = inputs.map((t2) => t2.shape);
|
|
assertParamsConsistent(shapes, $axis);
|
|
let outShape = computeOutShape$1(inputs.map((t2) => t2.shape), $axis);
|
|
if (sizeFromShape(outShape) === 0) {
|
|
return backend2.makeTensorInfo(outShape, inputs[0].dtype, []);
|
|
}
|
|
const $inputs = inputs.filter((t2) => sizeFromShape(t2.shape) > 0);
|
|
if ($inputs.length === 1) {
|
|
return identity$1({ inputs: { x: $inputs[0] }, backend: backend2 });
|
|
}
|
|
if ($inputs[0].dtype === "complex64") {
|
|
const reals = $inputs.map((t2) => real$1({ inputs: { input: t2 }, backend: backend2 }));
|
|
const imags = $inputs.map((t2) => imag({ inputs: { input: t2 }, backend: backend2 }));
|
|
const realConcated = concat({ inputs: reals, backend: backend2, attrs: { axis: $axis } });
|
|
const imagConcated = concat({ inputs: imags, backend: backend2, attrs: { axis: $axis } });
|
|
const result = complex$1({ inputs: { real: realConcated, imag: imagConcated }, backend: backend2 });
|
|
reals.forEach((r) => backend2.disposeIntermediateTensorInfo(r));
|
|
imags.forEach((i) => backend2.disposeIntermediateTensorInfo(i));
|
|
backend2.disposeIntermediateTensorInfo(realConcated);
|
|
backend2.disposeIntermediateTensorInfo(imagConcated);
|
|
return result;
|
|
}
|
|
const inputs2D = $inputs.map((t2) => {
|
|
const innerSize = sizeFromShape(t2.shape.slice($axis));
|
|
const shape = [-1, innerSize];
|
|
return reshape({ inputs: { x: t2 }, backend: backend2, attrs: { shape } });
|
|
});
|
|
const inputsValShapes = inputs2D.map((t2) => {
|
|
return { vals: backend2.data.get(t2.dataId).values, shape: t2.shape };
|
|
});
|
|
outShape = computeOutShape$1(
|
|
inputs2D.map((t2) => t2.shape),
|
|
1
|
|
/* axis */
|
|
);
|
|
const simplyConcat = inputs2D[0].shape[0] === 1;
|
|
const outVals = concatImpl$1(inputsValShapes, outShape, inputs[0].dtype, simplyConcat);
|
|
const finalOutShape = computeOutShape$1($inputs.map((t2) => t2.shape), $axis);
|
|
const outInfo = backend2.makeTensorInfo(finalOutShape, inputs[0].dtype, outVals);
|
|
inputs2D.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return outInfo;
|
|
}
|
|
const concatConfig = {
|
|
kernelName: Concat,
|
|
backendName: "cpu",
|
|
kernelFunc: concat
|
|
};
|
|
function conv2D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dilations, dimRoundingMode } = attrs;
|
|
assertNotComplex([x, filter], "conv2d");
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(x.shape, filter.shape, strides, dilations, pad2, dimRoundingMode, false, $dataFormat);
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const padLeft = convInfo.padInfo.left;
|
|
const padTop = convInfo.padInfo.top;
|
|
const isChannelsLast = convInfo.dataFormat === "channelsLast";
|
|
const y = new TensorBuffer(convInfo.outShape, x.dtype);
|
|
const xStrides = computeStrides(x.shape);
|
|
const filterStrides = computeStrides(filter.shape);
|
|
const xBatchStride = xStrides[0];
|
|
const xRowStride = isChannelsLast ? xStrides[1] : xStrides[2];
|
|
const xColStride = isChannelsLast ? xStrides[2] : 1;
|
|
const xChannelStride = isChannelsLast ? 1 : xStrides[1];
|
|
const yBatchStride = y.strides[0];
|
|
const yRowStride = isChannelsLast ? y.strides[1] : y.strides[2];
|
|
const yColStride = isChannelsLast ? y.strides[2] : 1;
|
|
const yChannelStride = isChannelsLast ? 1 : y.strides[1];
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const wVals = backend2.data.get(filter.dataId).values;
|
|
const yVals = y.values;
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
const xOffset1 = b * xBatchStride;
|
|
const yOffset1 = b * yBatchStride;
|
|
for (let yR = 0; yR < convInfo.outHeight; ++yR) {
|
|
const yOffset2 = yOffset1 + yR * yRowStride;
|
|
const xRCorner = yR * convInfo.strideHeight - padTop;
|
|
for (let wR = 0; wR < filterHeight; ++wR) {
|
|
const xR = xRCorner + wR * dilationHeight;
|
|
if (xR < 0 || xR >= convInfo.inHeight) {
|
|
continue;
|
|
}
|
|
const wOffset1 = wR * filterStrides[0];
|
|
const xOffset2 = xOffset1 + xR * xRowStride;
|
|
for (let yC = 0; yC < convInfo.outWidth; ++yC) {
|
|
const yOffset3 = yOffset2 + yC * yColStride;
|
|
const xCCorner = yC * convInfo.strideWidth - padLeft;
|
|
for (let wC = 0; wC < filterWidth; ++wC) {
|
|
const xC = xCCorner + wC * dilationWidth;
|
|
if (xC < 0 || xC >= convInfo.inWidth) {
|
|
continue;
|
|
}
|
|
const wOffset2 = wOffset1 + wC * filterStrides[1];
|
|
const xOffset3 = xOffset2 + xC * xColStride;
|
|
let wOffset3 = wOffset2;
|
|
for (let d1 = 0; d1 < convInfo.inChannels; ++d1) {
|
|
const xVal = xVals[xOffset3 + d1 * xChannelStride];
|
|
for (let d2 = 0; d2 < convInfo.outChannels; ++d2) {
|
|
yVals[yOffset3 + d2 * yChannelStride] += xVal * wVals[wOffset3 + d2];
|
|
}
|
|
wOffset3 += convInfo.outChannels;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(y.shape, y.dtype, yVals);
|
|
}
|
|
const conv2DConfig = {
|
|
kernelName: Conv2D,
|
|
backendName: "cpu",
|
|
kernelFunc: conv2D
|
|
};
|
|
function conv2DBackpropFilter(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, dy } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dimRoundingMode, filterShape } = attrs;
|
|
assertNotComplex([x, dy], "conv2dBackpropFilter");
|
|
const $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(x.shape, filterShape, strides, 1, pad2, dimRoundingMode, false, $dataFormat);
|
|
const { strideHeight, strideWidth, filterHeight, filterWidth } = convInfo;
|
|
const isChannelsLast = convInfo.dataFormat === "channelsLast";
|
|
const dW = new TensorBuffer(convInfo.filterShape, "float32");
|
|
const leftPad = convInfo.padInfo.left;
|
|
const topPad = convInfo.padInfo.top;
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const dyVals = backend2.data.get(dy.dataId).values;
|
|
const xBuf = new TensorBuffer(x.shape, x.dtype, xVals);
|
|
const dyBuf = new TensorBuffer(dy.shape, dy.dtype, dyVals);
|
|
for (let wR = 0; wR < filterHeight; ++wR) {
|
|
const yRMin = Math.max(0, Math.ceil((topPad - wR) / strideHeight));
|
|
const yRMax = Math.min(convInfo.outHeight, (convInfo.inHeight + topPad - wR) / strideHeight);
|
|
for (let wC = 0; wC < filterWidth; ++wC) {
|
|
const yCMin = Math.max(0, Math.ceil((leftPad - wC) / strideWidth));
|
|
const yCMax = Math.min(convInfo.outWidth, (convInfo.inWidth + leftPad - wC) / strideWidth);
|
|
for (let d1 = 0; d1 < convInfo.inChannels; ++d1) {
|
|
for (let d2 = 0; d2 < convInfo.outChannels; ++d2) {
|
|
let dotProd = 0;
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
for (let yR = yRMin; yR < yRMax; ++yR) {
|
|
const xR = wR + yR * strideHeight - topPad;
|
|
for (let yC = yCMin; yC < yCMax; ++yC) {
|
|
const xC = wC + yC * strideWidth - leftPad;
|
|
if (isChannelsLast) {
|
|
dotProd += xBuf.get(b, xR, xC, d1) * dyBuf.get(b, yR, yC, d2);
|
|
} else {
|
|
dotProd += xBuf.get(b, d1, xR, xC) * dyBuf.get(b, d2, yR, yC);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
dW.set(dotProd, wR, wC, d1, d2);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dW.shape, dW.dtype, dW.values);
|
|
}
|
|
const conv2DBackpropFilterConfig = {
|
|
kernelName: Conv2DBackpropFilter,
|
|
backendName: "cpu",
|
|
kernelFunc: conv2DBackpropFilter
|
|
};
|
|
function conv2DBackpropInput(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, filter } = inputs;
|
|
const { inputShape, strides, pad: pad2, dataFormat, dimRoundingMode } = attrs;
|
|
assertNotComplex([dy, filter], "conv2dBackpropInput");
|
|
const filterStrides = computeStrides(filter.shape);
|
|
const dyStrides = computeStrides(dy.shape);
|
|
let $dataFormat = convertConv2DDataFormat(dataFormat);
|
|
const convInfo = computeConv2DInfo(inputShape, filter.shape, strides, 1, pad2, dimRoundingMode, false, $dataFormat);
|
|
const dx = new TensorBuffer(convInfo.inShape, "float32");
|
|
const dxValues = dx.values;
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
const fltValues = backend2.data.get(filter.dataId).values;
|
|
const [fltS0, fltS1, fltS2] = filterStrides;
|
|
const { batchSize, filterHeight, filterWidth, inChannels, inHeight, inWidth, outChannels, outHeight, outWidth, strideHeight, strideWidth } = convInfo;
|
|
$dataFormat = convInfo.dataFormat;
|
|
const topPad = filterHeight - 1 - convInfo.padInfo.top;
|
|
const leftPad = filterWidth - 1 - convInfo.padInfo.left;
|
|
const isChannelsLast = $dataFormat === "channelsLast";
|
|
const xBatchStride = dx.strides[0];
|
|
const xRowStride = isChannelsLast ? dx.strides[1] : dx.strides[2];
|
|
const xColStride = isChannelsLast ? dx.strides[2] : 1;
|
|
const xChannelStride = isChannelsLast ? 1 : dx.strides[1];
|
|
const yBatchStride = dyStrides[0];
|
|
const yRowStride = isChannelsLast ? dyStrides[1] : dyStrides[2];
|
|
const yColStride = isChannelsLast ? dyStrides[2] : 1;
|
|
const yChannelStride = isChannelsLast ? 1 : dyStrides[1];
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let d1 = 0; d1 < inChannels; ++d1) {
|
|
for (let xR = 0; xR < inHeight; ++xR) {
|
|
const xRCorner = xR - topPad;
|
|
const xRMin = Math.max(0, Math.ceil(xRCorner / strideHeight));
|
|
const yRMax = Math.min(outHeight, (filterHeight + xRCorner) / strideHeight);
|
|
for (let xC = 0; xC < inWidth; ++xC) {
|
|
const xCCorner = xC - leftPad;
|
|
const xCMin = Math.max(0, Math.ceil(xCCorner / strideWidth));
|
|
const yCMax = Math.min(outWidth, (filterWidth + xCCorner) / strideWidth);
|
|
let dotProd = 0;
|
|
for (let yR = xRMin; yR < yRMax; ++yR) {
|
|
const wR = yR * strideHeight - xRCorner;
|
|
for (let yC = xCMin; yC < yCMax; ++yC) {
|
|
const wC = yC * strideWidth - xCCorner;
|
|
const dyOffset = yBatchStride * b + yRowStride * yR + yColStride * yC;
|
|
const fltOffset = fltS0 * (filterHeight - 1 - wR) + fltS1 * (filterWidth - 1 - wC) + fltS2 * d1;
|
|
for (let d2 = 0; d2 < outChannels; ++d2) {
|
|
const pixel = dyValues[dyOffset + yChannelStride * d2];
|
|
const weight = fltValues[fltOffset + d2];
|
|
dotProd += pixel * weight;
|
|
}
|
|
}
|
|
}
|
|
const dxOffset = xBatchStride * b + xRowStride * xR + xColStride * xC + xChannelStride * d1;
|
|
dxValues[dxOffset] = dotProd;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const conv2DBackpropInputConfig = {
|
|
kernelName: Conv2DBackpropInput,
|
|
backendName: "cpu",
|
|
kernelFunc: conv2DBackpropInput
|
|
};
|
|
function conv3D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dilations } = attrs;
|
|
assertNotComplex([x, filter], "conv3d");
|
|
const convInfo = computeConv3DInfo(x.shape, filter.shape, strides, dilations, pad2);
|
|
const { filterDepth, filterHeight, filterWidth, dilationDepth, dilationHeight, dilationWidth, padInfo } = convInfo;
|
|
const padFront = padInfo.front;
|
|
const padLeft = padInfo.left;
|
|
const padTop = padInfo.top;
|
|
const y = new TensorBuffer(convInfo.outShape, x.dtype);
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const wVals = backend2.data.get(filter.dataId).values;
|
|
const yVals = y.values;
|
|
const xStrides = computeStrides(x.shape);
|
|
const filterStrides = computeStrides(filter.shape);
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
const xOffset1 = b * xStrides[0];
|
|
const yOffset1 = b * y.strides[0];
|
|
for (let yF = 0; yF < convInfo.outDepth; ++yF) {
|
|
const yOffset2 = yOffset1 + yF * y.strides[1];
|
|
const xFCorner = yF * convInfo.strideDepth - padFront;
|
|
for (let wF = 0; wF < filterDepth; ++wF) {
|
|
const xF = xFCorner + wF * dilationDepth;
|
|
if (xF < 0 || xF >= convInfo.inDepth) {
|
|
continue;
|
|
}
|
|
const wOffset1 = wF * filterStrides[0];
|
|
const xOffset2 = xOffset1 + xF * xStrides[1];
|
|
for (let yR = 0; yR < convInfo.outHeight; ++yR) {
|
|
const yOffset3 = yOffset2 + yR * y.strides[2];
|
|
const xRCorner = yR * convInfo.strideHeight - padTop;
|
|
for (let wR = 0; wR < filterHeight; ++wR) {
|
|
const xR = xRCorner + wR * dilationHeight;
|
|
if (xR < 0 || xR >= convInfo.inHeight) {
|
|
continue;
|
|
}
|
|
const wOffset2 = wOffset1 + wR * filterStrides[1];
|
|
const xOffset3 = xOffset2 + xR * xStrides[2];
|
|
for (let yC = 0; yC < convInfo.outWidth; ++yC) {
|
|
const yOffset4 = yOffset3 + yC * convInfo.outChannels;
|
|
const xCCorner = yC * convInfo.strideWidth - padLeft;
|
|
for (let wC = 0; wC < filterWidth; ++wC) {
|
|
const xC = xCCorner + wC * dilationWidth;
|
|
if (xC < 0 || xC >= convInfo.inWidth) {
|
|
continue;
|
|
}
|
|
const wOffset3 = wOffset2 + wC * filterStrides[2];
|
|
const xOffset4 = xOffset3 + xC * convInfo.inChannels;
|
|
let wOffset4 = wOffset3;
|
|
for (let d1 = 0; d1 < convInfo.inChannels; ++d1) {
|
|
const xVal = xVals[xOffset4 + d1];
|
|
for (let d2 = 0; d2 < convInfo.outChannels; ++d2) {
|
|
yVals[yOffset4 + d2] += xVal * wVals[wOffset4 + d2];
|
|
}
|
|
wOffset4 += convInfo.outChannels;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(y.shape, y.dtype, y.values);
|
|
}
|
|
const conv3DConfig = {
|
|
kernelName: Conv3D,
|
|
backendName: "cpu",
|
|
kernelFunc: conv3D
|
|
};
|
|
function conv3DBackpropFilterV2(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, dy } = inputs;
|
|
const { strides, pad: pad2, filterShape } = attrs;
|
|
assertNotComplex([x, dy], "conv3dBackpropFilterV2");
|
|
const xStrides = computeStrides(x.shape);
|
|
const dyStrides = computeStrides(dy.shape);
|
|
const convInfo = computeConv3DInfo(x.shape, filterShape, strides, 1, pad2);
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const filterDepth = convInfo.filterDepth;
|
|
const filterHeight = convInfo.filterHeight;
|
|
const filterWidth = convInfo.filterWidth;
|
|
const dw = new TensorBuffer(convInfo.filterShape, "float32");
|
|
const dwValues = dw.values;
|
|
const [dwS0, dwS1, dwS2, dwS3] = dw.strides;
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
const [dyS0, dyS1, dyS2, dyS3] = dyStrides;
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const [xS0, xS1, xS2, xS3] = xStrides;
|
|
const frontPad = convInfo.padInfo.front;
|
|
const leftPad = convInfo.padInfo.left;
|
|
const topPad = convInfo.padInfo.top;
|
|
for (let wF = 0; wF < filterDepth; ++wF) {
|
|
const yFMin = Math.max(0, Math.ceil((frontPad - wF) / strideDepth));
|
|
const yFMax = Math.min(convInfo.outDepth, (convInfo.inDepth + frontPad - wF) / strideDepth);
|
|
const wOffset1 = wF * dwS0;
|
|
for (let wR = 0; wR < filterHeight; ++wR) {
|
|
const yRMin = Math.max(0, Math.ceil((topPad - wR) / strideHeight));
|
|
const yRMax = Math.min(convInfo.outHeight, (convInfo.inHeight + topPad - wR) / strideHeight);
|
|
const wOffset2 = wR * dwS1 + wOffset1;
|
|
for (let wC = 0; wC < filterWidth; ++wC) {
|
|
const yCMin = Math.max(0, Math.ceil((leftPad - wC) / strideWidth));
|
|
const yCMax = Math.min(convInfo.outWidth, (convInfo.inWidth + leftPad - wC) / strideWidth);
|
|
const wOffset3 = wC * dwS2 + wOffset2;
|
|
for (let d1 = 0; d1 < convInfo.inChannels; ++d1) {
|
|
const wOffset4 = d1 * dwS3 + wOffset3;
|
|
for (let d2 = 0; d2 < convInfo.outChannels; ++d2) {
|
|
let dotProd = 0;
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
const xOffset1 = b * xS0;
|
|
const yOffset1 = b * dyS0;
|
|
for (let yF = yFMin; yF < yFMax; ++yF) {
|
|
const xF = wF + yF * strideDepth - frontPad;
|
|
const xOffset2 = xF * xS1 + xOffset1;
|
|
const yOffset2 = yF * dyS1 + yOffset1;
|
|
for (let yR = yRMin; yR < yRMax; ++yR) {
|
|
const xR = wR + yR * strideHeight - topPad;
|
|
const xOffset3 = xR * xS2 + xOffset2;
|
|
const yOffset3 = yR * dyS2 + yOffset2;
|
|
for (let yC = yCMin; yC < yCMax; ++yC) {
|
|
const xC = wC + yC * strideWidth - leftPad;
|
|
const xOffset4 = xC * xS3 + xOffset3;
|
|
const yOffset4 = yC * dyS3 + yOffset3;
|
|
dotProd += xValues[xOffset4 + d1] * dyValues[yOffset4 + d2];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
dwValues[wOffset4 + d2] = dotProd;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dw.shape, dw.dtype, dw.values);
|
|
}
|
|
const conv3DBackpropFilterV2Config = {
|
|
kernelName: Conv3DBackpropFilterV2,
|
|
backendName: "cpu",
|
|
kernelFunc: conv3DBackpropFilterV2
|
|
};
|
|
function conv3DBackpropInputV2(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, filter } = inputs;
|
|
const { pad: pad2, strides, inputShape } = attrs;
|
|
assertNotComplex([dy], "conv3dBackpropInputV2");
|
|
const dyStrides = computeStrides(dy.shape);
|
|
const filterStrides = computeStrides(filter.shape);
|
|
const convInfo = computeConv3DInfo(inputShape, filter.shape, strides, 1, pad2);
|
|
const dx = new TensorBuffer(convInfo.inShape, "float32");
|
|
const dxValues = dx.values;
|
|
const [dxS0, dxS1, dxS2, dxS3] = dx.strides;
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
const [dyS0, dyS1, dyS2, dyS3] = dyStrides;
|
|
const fltValues = backend2.data.get(filter.dataId).values;
|
|
const [fltS0, fltS1, fltS2, fltS3] = filterStrides;
|
|
const { batchSize, filterDepth, filterHeight, filterWidth, inChannels, inDepth, inHeight, inWidth, outChannels, outDepth, outHeight, outWidth, strideDepth, strideHeight, strideWidth } = convInfo;
|
|
const frontPad = filterDepth - 1 - convInfo.padInfo.front;
|
|
const topPad = filterHeight - 1 - convInfo.padInfo.top;
|
|
const leftPad = filterWidth - 1 - convInfo.padInfo.left;
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let d1 = 0; d1 < inChannels; ++d1) {
|
|
for (let xF = 0; xF < inDepth; ++xF) {
|
|
const xFCorner = xF - frontPad;
|
|
const xFMin = Math.max(0, Math.ceil(xFCorner / strideDepth));
|
|
const yFMax = Math.min(outDepth, (filterDepth + xFCorner) / strideDepth);
|
|
for (let xR = 0; xR < inHeight; ++xR) {
|
|
const xRCorner = xR - topPad;
|
|
const xRMin = Math.max(0, Math.ceil(xRCorner / strideHeight));
|
|
const yRMax = Math.min(outHeight, (filterHeight + xRCorner) / strideHeight);
|
|
for (let xC = 0; xC < inWidth; ++xC) {
|
|
const xCCorner = xC - leftPad;
|
|
const xCMin = Math.max(0, Math.ceil(xCCorner / strideWidth));
|
|
const yCMax = Math.min(outWidth, (filterWidth + xCCorner) / strideWidth);
|
|
let dotProd = 0;
|
|
for (let yF = xFMin; yF < yFMax; ++yF) {
|
|
const wF = yF * strideDepth - xFCorner;
|
|
for (let yR = xRMin; yR < yRMax; ++yR) {
|
|
const wR = yR * strideHeight - xRCorner;
|
|
for (let yC = xCMin; yC < yCMax; ++yC) {
|
|
const wC = yC * strideWidth - xCCorner;
|
|
const dyOffset = dyS0 * b + dyS1 * yF + dyS2 * yR + dyS3 * yC;
|
|
const fltOffset = fltS0 * (filterDepth - 1 - wF) + fltS1 * (filterHeight - 1 - wR) + fltS2 * (filterWidth - 1 - wC) + fltS3 * d1;
|
|
for (let d2 = 0; d2 < outChannels; ++d2) {
|
|
const pixel = dyValues[dyOffset + d2];
|
|
const weight = fltValues[fltOffset + d2];
|
|
dotProd += pixel * weight;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
dxValues[dxS0 * b + dxS1 * xF + dxS2 * xR + dxS3 * xC + d1] = dotProd;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const conv3DBackpropInputV2Config = {
|
|
kernelName: Conv3DBackpropInputV2,
|
|
backendName: "cpu",
|
|
kernelFunc: conv3DBackpropInputV2
|
|
};
|
|
const cos = unaryKernelFunc$1(Cos, (xi) => Math.cos(xi));
|
|
const cosConfig = {
|
|
kernelName: Cos,
|
|
backendName: "cpu",
|
|
kernelFunc: cos
|
|
};
|
|
const cosh = unaryKernelFunc$1(Cosh, (xi) => Math.cosh(xi));
|
|
const coshConfig = {
|
|
kernelName: Cosh,
|
|
backendName: "cpu",
|
|
kernelFunc: cosh
|
|
};
|
|
function cropAndResize(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { image: image2, boxes, boxInd } = inputs;
|
|
const { cropSize, method, extrapolationValue } = attrs;
|
|
const [batch, imageHeight, imageWidth, numChannels] = image2.shape;
|
|
const numBoxes = boxes.shape[0];
|
|
const [cropHeight, cropWidth] = cropSize;
|
|
const output = buffer([numBoxes, cropHeight, cropWidth, numChannels], "float32");
|
|
const boxVals = backend2.data.get(boxes.dataId).values;
|
|
const boxIndVals = backend2.data.get(boxInd.dataId).values;
|
|
const imageVals = backend2.data.get(image2.dataId).values;
|
|
const inStride = computeStrides(image2.shape);
|
|
const outStride = computeStrides(output.shape);
|
|
for (let b = 0; b < numBoxes; b++) {
|
|
const startInd = b * 4;
|
|
const y1 = boxVals[startInd];
|
|
const x1 = boxVals[startInd + 1];
|
|
const y2 = boxVals[startInd + 2];
|
|
const x2 = boxVals[startInd + 3];
|
|
const bInd = boxIndVals[b];
|
|
if (bInd >= batch) {
|
|
continue;
|
|
}
|
|
const heightScale = cropHeight > 1 ? (y2 - y1) * (imageHeight - 1) / (cropHeight - 1) : 0;
|
|
const widthScale = cropWidth > 1 ? (x2 - x1) * (imageWidth - 1) / (cropWidth - 1) : 0;
|
|
for (let y = 0; y < cropHeight; y++) {
|
|
const yInd = cropHeight > 1 ? y1 * (imageHeight - 1) + y * heightScale : 0.5 * (y1 + y2) * (imageHeight - 1);
|
|
if (yInd < 0 || yInd > imageHeight - 1) {
|
|
for (let x = 0; x < cropWidth; x++) {
|
|
for (let c = 0; c < numChannels; c++) {
|
|
const ind = c + x * outStride[2] + y * outStride[1] + b * outStride[0];
|
|
output.values[ind] = extrapolationValue;
|
|
}
|
|
}
|
|
continue;
|
|
}
|
|
if (method === "bilinear") {
|
|
const topInd = Math.floor(yInd);
|
|
const bottomInd = Math.ceil(yInd);
|
|
const yLerp = yInd - topInd;
|
|
for (let x = 0; x < cropWidth; x++) {
|
|
const xInd = cropWidth > 1 ? x1 * (imageWidth - 1) + x * widthScale : 0.5 * (x1 + x2) * (imageWidth - 1);
|
|
if (xInd < 0 || xInd > imageWidth - 1) {
|
|
for (let c = 0; c < numChannels; c++) {
|
|
const ind = c + x * outStride[2] + y * outStride[1] + b * outStride[0];
|
|
output.values[ind] = extrapolationValue;
|
|
}
|
|
continue;
|
|
}
|
|
const leftInd = Math.floor(xInd);
|
|
const rightInd = Math.ceil(xInd);
|
|
const xLerp = xInd - leftInd;
|
|
for (let c = 0; c < numChannels; c++) {
|
|
let ind = c + leftInd * inStride[2] + topInd * inStride[1] + bInd * inStride[0];
|
|
const topLeft = imageVals[ind];
|
|
ind = c + rightInd * inStride[2] + topInd * inStride[1] + bInd * inStride[0];
|
|
const topRight = imageVals[ind];
|
|
ind = c + leftInd * inStride[2] + bottomInd * inStride[1] + bInd * inStride[0];
|
|
const bottomLeft = imageVals[ind];
|
|
ind = c + rightInd * inStride[2] + bottomInd * inStride[1] + bInd * inStride[0];
|
|
const bottomRight = imageVals[ind];
|
|
const top = topLeft + (topRight - topLeft) * xLerp;
|
|
const bottom = bottomLeft + (bottomRight - bottomLeft) * xLerp;
|
|
ind = c + x * outStride[2] + y * outStride[1] + b * outStride[0];
|
|
output.values[ind] = top + (bottom - top) * yLerp;
|
|
}
|
|
}
|
|
} else {
|
|
for (let x = 0; x < cropWidth; ++x) {
|
|
const xInd = cropWidth > 1 ? x1 * (imageWidth - 1) + x * widthScale : 0.5 * (x1 + x2) * (imageWidth - 1);
|
|
if (xInd < 0 || xInd > imageWidth - 1) {
|
|
for (let c = 0; c < numChannels; c++) {
|
|
const ind = c + x * outStride[2] + y * outStride[1] + b * outStride[0];
|
|
output.values[ind] = extrapolationValue;
|
|
}
|
|
continue;
|
|
}
|
|
const closestX = Math.round(xInd);
|
|
const closestY = Math.round(yInd);
|
|
for (let c = 0; c < numChannels; c++) {
|
|
const inInd = c + closestX * inStride[2] + closestY * inStride[1] + bInd * inStride[0];
|
|
const outInd = c + x * outStride[2] + y * outStride[1] + b * outStride[0];
|
|
output.values[outInd] = imageVals[inInd];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(output.shape, output.dtype, output.values);
|
|
}
|
|
const cropAndResizeConfig = {
|
|
kernelName: CropAndResize,
|
|
backendName: "cpu",
|
|
kernelFunc: cropAndResize
|
|
};
|
|
function cumprod(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, exclusive, reverse: reverse2 } = attrs;
|
|
assertNotComplex(x, "cumprod");
|
|
const permutation = getAxesPermutation([axis], x.shape.length);
|
|
let $x = x;
|
|
if (permutation != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutation } });
|
|
}
|
|
const permutedAxis = getInnerMostAxes(1, x.shape.length)[0];
|
|
if (permutedAxis !== $x.shape.length - 1) {
|
|
throw new Error(`backend.cumprod in CPU expects an inner-most axis=${$x.shape.length - 1} but got axis=${permutedAxis}`);
|
|
}
|
|
const resultDtype = upcastType($x.dtype, "int32");
|
|
const vals = makeOnesTypedArray(sizeFromShape($x.shape), resultDtype);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
const finalDim = $x.shape[$x.shape.length - 1];
|
|
const indexAdjuster = reverse2 ? (i, j2) => i + finalDim - j2 - 1 : (i, j2) => i + j2;
|
|
for (let i = 0; i < aVals.length; i += finalDim) {
|
|
for (let j2 = 0; j2 < finalDim; j2++) {
|
|
const idx = indexAdjuster(i, j2);
|
|
if (j2 === 0) {
|
|
vals[idx] = exclusive ? 1 : aVals[idx];
|
|
} else {
|
|
const prevIdx = indexAdjuster(i, j2 - 1);
|
|
vals[idx] = exclusive ? aVals[prevIdx] * vals[prevIdx] : aVals[idx] * vals[prevIdx];
|
|
}
|
|
}
|
|
}
|
|
const result = backend2.makeTensorInfo($x.shape, resultDtype, vals);
|
|
if (permutation != null) {
|
|
const reversePermutation = getUndoAxesPermutation(permutation);
|
|
const reverseTransposedResult = transpose$1({ inputs: { x: result }, backend: backend2, attrs: { perm: reversePermutation } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
backend2.disposeIntermediateTensorInfo($x);
|
|
return reverseTransposedResult;
|
|
}
|
|
return result;
|
|
}
|
|
const cumprodConfig = {
|
|
kernelName: Cumprod,
|
|
backendName: "cpu",
|
|
kernelFunc: cumprod
|
|
};
|
|
function cumsum(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, exclusive, reverse: reverse2 } = attrs;
|
|
assertNotComplex(x, "cumsum");
|
|
const permutation = getAxesPermutation([axis], x.shape.length);
|
|
let $x = x;
|
|
if (permutation != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutation } });
|
|
}
|
|
const permutedAxis = getInnerMostAxes(1, x.shape.length)[0];
|
|
if (permutedAxis !== $x.shape.length - 1) {
|
|
throw new Error(`backend.cumsum in CPU expects an inner-most axis=${$x.shape.length - 1} but got axis=${permutedAxis}`);
|
|
}
|
|
const resultDtype = upcastType($x.dtype, "int32");
|
|
const vals = makeZerosTypedArray(sizeFromShape($x.shape), resultDtype);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
const finalDim = $x.shape[$x.shape.length - 1];
|
|
const indexAdjuster = reverse2 ? (i, j2) => i + finalDim - j2 - 1 : (i, j2) => i + j2;
|
|
for (let i = 0; i < aVals.length; i += finalDim) {
|
|
for (let j2 = 0; j2 < finalDim; j2++) {
|
|
const idx = indexAdjuster(i, j2);
|
|
if (j2 === 0) {
|
|
vals[idx] = exclusive ? 0 : aVals[idx];
|
|
} else {
|
|
const prevIdx = indexAdjuster(i, j2 - 1);
|
|
vals[idx] = exclusive ? aVals[prevIdx] + vals[prevIdx] : aVals[idx] + vals[prevIdx];
|
|
}
|
|
}
|
|
}
|
|
const result = backend2.makeTensorInfo($x.shape, resultDtype, vals);
|
|
if (permutation != null) {
|
|
const reversePermutation = getUndoAxesPermutation(permutation);
|
|
const reverseTransposedResult = transpose$1({ inputs: { x: result }, backend: backend2, attrs: { perm: reversePermutation } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
backend2.disposeIntermediateTensorInfo($x);
|
|
return reverseTransposedResult;
|
|
}
|
|
return result;
|
|
}
|
|
const cumsumConfig = {
|
|
kernelName: Cumsum,
|
|
backendName: "cpu",
|
|
kernelFunc: cumsum
|
|
};
|
|
function denseBincount(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, weights } = inputs;
|
|
const { size, binaryOutput } = attrs;
|
|
if (x.shape.length === 1) {
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const weightsVals = backend2.data.get(weights.dataId).values;
|
|
const outVals = bincountImpl(xVals, weightsVals, weights.dtype, weights.shape, size);
|
|
return backend2.makeTensorInfo([size], weights.dtype, outVals);
|
|
} else if (x.shape.length === 2) {
|
|
const xBuf = backend2.bufferSync(x);
|
|
const weightsBuf = backend2.bufferSync(weights);
|
|
const outBuf = bincountReduceImpl(xBuf, weightsBuf, size, binaryOutput);
|
|
return backend2.makeTensorInfo(outBuf.shape, weights.dtype, outBuf.values);
|
|
}
|
|
throw new Error(`Error in denseBincount: input must be at most rank 2, but got rank${x.shape.length}.`);
|
|
}
|
|
const denseBincountConfig = {
|
|
kernelName: DenseBincount,
|
|
backendName: "cpu",
|
|
kernelFunc: denseBincount
|
|
};
|
|
function depthToSpace(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { blockSize, dataFormat } = attrs;
|
|
assert(dataFormat === "NHWC", () => `Only NHWC dataFormat supported on CPU for depthToSpace. Got ${dataFormat}`);
|
|
const batchSize = x.shape[0];
|
|
const inputHeight = x.shape[1];
|
|
const inputWidth = x.shape[2];
|
|
const inputDepth = x.shape[3];
|
|
const outputHeight = inputHeight * blockSize;
|
|
const outputWidth = inputWidth * blockSize;
|
|
const outputDepth = inputDepth / (blockSize * blockSize);
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const result = new Float32Array(batchSize * outputHeight * outputWidth * outputDepth);
|
|
let outputIdx = 0;
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let h = 0; h < outputHeight; ++h) {
|
|
const inH = Math.floor(h / blockSize);
|
|
const offsetH = h % blockSize;
|
|
for (let w = 0; w < outputWidth; ++w) {
|
|
const inW = Math.floor(w / blockSize);
|
|
const offsetW = w % blockSize;
|
|
const offsetD = (offsetH * blockSize + offsetW) * outputDepth;
|
|
for (let d = 0; d < outputDepth; ++d) {
|
|
const inD = d + offsetD;
|
|
const inputIdx = inD + inputDepth * (inW + inputWidth * (inH + inputHeight * b));
|
|
result[outputIdx++] = xValues[inputIdx];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo([batchSize, outputHeight, outputWidth, outputDepth], x.dtype, result);
|
|
}
|
|
const depthToSpaceConfig = {
|
|
kernelName: DepthToSpace,
|
|
backendName: "cpu",
|
|
kernelFunc: depthToSpace
|
|
};
|
|
function depthwiseConv2dNative(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dilations, dimRoundingMode } = attrs;
|
|
assertNotComplex([x, filter], "depthwiseConv2DNative");
|
|
const xStrides = computeStrides(x.shape);
|
|
const filterStrides = computeStrides(filter.shape);
|
|
let $dilations = dilations;
|
|
if ($dilations == null) {
|
|
$dilations = [1, 1];
|
|
}
|
|
assert(eitherStridesOrDilationsAreOne(strides, $dilations), () => `Error in depthwiseConv2d: Either strides or dilations must be 1. Got strides ${strides} and dilations '${$dilations}'`);
|
|
const convInfo = computeConv2DInfo(
|
|
x.shape,
|
|
filter.shape,
|
|
strides,
|
|
$dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
const { filterHeight, filterWidth, dilationHeight, dilationWidth, padInfo } = convInfo;
|
|
const padLeft = padInfo.left;
|
|
const padTop = padInfo.top;
|
|
const chMul = convInfo.outChannels / convInfo.inChannels;
|
|
const y = new TensorBuffer(convInfo.outShape, x.dtype);
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const wVals = backend2.data.get(filter.dataId).values;
|
|
const yVals = y.values;
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
const xOffset1 = b * xStrides[0];
|
|
const yOffset1 = b * y.strides[0];
|
|
for (let yR = 0; yR < convInfo.outHeight; ++yR) {
|
|
const yOffset2 = yOffset1 + yR * y.strides[1];
|
|
const xRCorner = yR * convInfo.strideHeight - padTop;
|
|
for (let wR = 0; wR < filterHeight; ++wR) {
|
|
const xR = xRCorner + wR * dilationHeight;
|
|
if (xR < 0 || xR >= convInfo.inHeight) {
|
|
continue;
|
|
}
|
|
const wOffset1 = wR * filterStrides[0];
|
|
const xOffset2 = xOffset1 + xR * xStrides[1];
|
|
for (let yC = 0; yC < convInfo.outWidth; ++yC) {
|
|
const yOffset3 = yOffset2 + yC * y.strides[2];
|
|
const xCCorner = yC * convInfo.strideWidth - padLeft;
|
|
for (let wC = 0; wC < filterWidth; ++wC) {
|
|
const xC = xCCorner + wC * dilationWidth;
|
|
if (xC < 0 || xC >= convInfo.inWidth) {
|
|
continue;
|
|
}
|
|
const wOffset2 = wOffset1 + wC * filterStrides[1];
|
|
const xOffset3 = xOffset2 + xC * convInfo.inChannels;
|
|
let yOffset4 = yOffset3;
|
|
let wOffset3 = wOffset2;
|
|
for (let d1 = 0; d1 < convInfo.inChannels; ++d1) {
|
|
const xVal = xVals[xOffset3 + d1];
|
|
for (let q2 = 0; q2 < chMul; ++q2) {
|
|
yVals[yOffset4 + q2] += xVal * wVals[wOffset3 + q2];
|
|
}
|
|
yOffset4 += chMul;
|
|
wOffset3 += chMul;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(y.shape, y.dtype, y.values);
|
|
}
|
|
const depthwiseConv2dNativeConfig = {
|
|
kernelName: DepthwiseConv2dNative,
|
|
backendName: "cpu",
|
|
kernelFunc: depthwiseConv2dNative
|
|
};
|
|
function depthwiseConv2dNativeBackpropFilter(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, dy } = inputs;
|
|
const { strides, dilations, pad: pad2, dimRoundingMode, filterShape } = attrs;
|
|
assertNotComplex([x, dy], "depthwiseConv2dNativeBackpropFilter");
|
|
const convInfo = computeConv2DInfo(
|
|
x.shape,
|
|
filterShape,
|
|
strides,
|
|
dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
const { strideHeight, strideWidth, filterHeight, filterWidth } = convInfo;
|
|
const dW = new TensorBuffer(convInfo.filterShape, "float32");
|
|
const leftPad = convInfo.padInfo.left;
|
|
const topPad = convInfo.padInfo.top;
|
|
const chMul = convInfo.outChannels / convInfo.inChannels;
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const xBuf = new TensorBuffer(x.shape, x.dtype, xVals);
|
|
const dyVals = backend2.data.get(dy.dataId).values;
|
|
const dyBuf = new TensorBuffer(dy.shape, dy.dtype, dyVals);
|
|
for (let wR = 0; wR < filterHeight; ++wR) {
|
|
const yRMin = Math.max(0, Math.ceil((topPad - wR) / strideHeight));
|
|
const yRMax = Math.min(convInfo.outHeight, (convInfo.inHeight + topPad - wR) / strideHeight);
|
|
for (let wC = 0; wC < filterWidth; ++wC) {
|
|
const yCMin = Math.max(0, Math.ceil((leftPad - wC) / strideWidth));
|
|
const yCMax = Math.min(convInfo.outWidth, (convInfo.inWidth + leftPad - wC) / strideWidth);
|
|
for (let d2 = 0; d2 < convInfo.outChannels; ++d2) {
|
|
const d1 = Math.trunc(d2 / chMul);
|
|
const dm = d2 % chMul;
|
|
let dotProd = 0;
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
for (let yR = yRMin; yR < yRMax; ++yR) {
|
|
const xR = wR + yR * strideHeight - topPad;
|
|
for (let yC = yCMin; yC < yCMax; ++yC) {
|
|
const xC = wC + yC * strideWidth - leftPad;
|
|
dotProd += xBuf.get(b, xR, xC, d1) * dyBuf.get(b, yR, yC, d2);
|
|
}
|
|
}
|
|
}
|
|
dW.set(dotProd, wR, wC, d1, dm);
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dW.shape, dW.dtype, dW.values);
|
|
}
|
|
const depthwiseConv2dNativeBackpropFilterConfig = {
|
|
kernelName: DepthwiseConv2dNativeBackpropFilter,
|
|
backendName: "cpu",
|
|
kernelFunc: depthwiseConv2dNativeBackpropFilter
|
|
};
|
|
function depthwiseConv2dNativeBackpropInput(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, filter } = inputs;
|
|
const { strides, dilations, pad: pad2, dimRoundingMode, inputShape } = attrs;
|
|
assertNotComplex([dy, filter], "depthwiseConv2DNativeBackpropInput");
|
|
const dyStrides = computeStrides(dy.shape);
|
|
const filterStrides = computeStrides(filter.shape);
|
|
const convInfo = computeConv2DInfo(
|
|
inputShape,
|
|
filter.shape,
|
|
strides,
|
|
dilations,
|
|
pad2,
|
|
dimRoundingMode,
|
|
true
|
|
/* depthwise */
|
|
);
|
|
const dx = new TensorBuffer(convInfo.inShape, "float32");
|
|
const dxValues = dx.values;
|
|
const [dxS0, dxS1, dxS2] = dx.strides;
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
const [dyS0, dyS1, dyS2] = dyStrides;
|
|
const fltValues = backend2.data.get(filter.dataId).values;
|
|
const [fltS0, fltS1, fltS2] = filterStrides;
|
|
const { batchSize, filterHeight, filterWidth, inChannels, inHeight, inWidth, outChannels, outHeight, outWidth, strideHeight, strideWidth } = convInfo;
|
|
const topPad = filterHeight - 1 - convInfo.padInfo.top;
|
|
const leftPad = filterWidth - 1 - convInfo.padInfo.left;
|
|
const chMul = outChannels / inChannels;
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let d1 = 0; d1 < inChannels; ++d1) {
|
|
for (let xR = 0; xR < inHeight; ++xR) {
|
|
const xRCorner = xR - topPad;
|
|
const xRMin = Math.max(0, Math.ceil(xRCorner / strideHeight));
|
|
const yRMax = Math.min(outHeight, (filterHeight + xRCorner) / strideHeight);
|
|
for (let xC = 0; xC < inWidth; ++xC) {
|
|
const xCCorner = xC - leftPad;
|
|
const xCMin = Math.max(0, Math.ceil(xCCorner / strideWidth));
|
|
const yCMax = Math.min(outWidth, (filterWidth + xCCorner) / strideWidth);
|
|
let dotProd = 0;
|
|
for (let yR = xRMin; yR < yRMax; ++yR) {
|
|
const wR = yR * strideHeight - xRCorner;
|
|
for (let yC = xCMin; yC < yCMax; ++yC) {
|
|
const wC = yC * strideWidth - xCCorner;
|
|
const dyOffset = dyS0 * b + dyS1 * yR + dyS2 * yC;
|
|
const fltOffset = fltS0 * (filterHeight - 1 - wR) + fltS1 * (filterWidth - 1 - wC) + fltS2 * d1;
|
|
for (let dm = 0; dm < chMul; ++dm) {
|
|
const d2 = d1 * chMul + dm;
|
|
const pixel = dyValues[dyOffset + d2];
|
|
const weight = fltValues[fltOffset + dm];
|
|
dotProd += pixel * weight;
|
|
}
|
|
}
|
|
}
|
|
dxValues[dxS0 * b + dxS1 * xR + dxS2 * xC + d1] = dotProd;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const depthwiseConv2dNativeBackpropInputConfig = {
|
|
kernelName: DepthwiseConv2dNativeBackpropInput,
|
|
backendName: "cpu",
|
|
kernelFunc: depthwiseConv2dNativeBackpropInput
|
|
};
|
|
function diag(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
const xSize = sizeFromShape(x.shape);
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const outBuf = buffer([xSize, xSize], x.dtype);
|
|
const vals = outBuf.values;
|
|
for (let i = 0; i < xVals.length; i++) {
|
|
vals[i * xSize + i] = xVals[i];
|
|
}
|
|
const outShape = [...x.shape, ...x.shape];
|
|
return backend2.makeTensorInfo(outShape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const diagConfig = {
|
|
kernelName: Diag,
|
|
backendName: "cpu",
|
|
kernelFunc: diag
|
|
};
|
|
const dilation2DConfig = {
|
|
kernelName: Dilation2D,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, backend: backend2, attrs }) => {
|
|
const { x, filter } = inputs;
|
|
const { strides, pad: pad2, dilations } = attrs;
|
|
const cpuBackend = backend2;
|
|
const xVals = cpuBackend.data.get(x.dataId).values;
|
|
const xRank = x.shape.length;
|
|
const filterVals = cpuBackend.data.get(filter.dataId).values;
|
|
const filterRank = filter.shape.length;
|
|
const { batchSize, inHeight, inWidth, inChannels, outHeight, outWidth, padInfo, strideHeight, strideWidth, filterHeight, filterWidth, dilationHeight, dilationWidth, outShape } = computeDilation2DInfo(x.shape, filter.shape, strides, pad2, "NHWC", dilations);
|
|
const outSize = sizeFromShape(outShape);
|
|
const outRank = outShape.length;
|
|
const outputVals = getArrayFromDType(x.dtype, outSize);
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let hOut = 0; hOut < outHeight; ++hOut) {
|
|
const hBeg = hOut * strideHeight - padInfo.top;
|
|
for (let wOut = 0; wOut < outWidth; ++wOut) {
|
|
const wBeg = wOut * strideWidth - padInfo.left;
|
|
for (let d = 0; d < inChannels; ++d) {
|
|
let curVal = Number.MIN_SAFE_INTEGER;
|
|
for (let h = 0; h < filterHeight; ++h) {
|
|
const hIn = hBeg + h * dilationHeight;
|
|
if (hIn >= 0 && hIn < inHeight) {
|
|
for (let w = 0; w < filterWidth; ++w) {
|
|
const wIn = wBeg + w * dilationWidth;
|
|
if (wIn >= 0 && wIn < inWidth) {
|
|
const xIndex = locToIndex([b, hIn, wIn, d], xRank, computeStrides(x.shape));
|
|
const filterIndex = locToIndex([h, w, d], filterRank, computeStrides(filter.shape));
|
|
const val = xVals[xIndex] + filterVals[filterIndex];
|
|
if (val > curVal) {
|
|
curVal = val;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const outputIndex = locToIndex([b, hOut, wOut, d], outRank, computeStrides(outShape));
|
|
outputVals[outputIndex] = curVal;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const dataId = cpuBackend.write(toTypedArray(outputVals, x.dtype), outShape, x.dtype);
|
|
return { dataId, shape: outShape, dtype: x.dtype };
|
|
}
|
|
};
|
|
const dilation2DBackpropFilterConfig = {
|
|
kernelName: Dilation2DBackpropFilter,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, backend: backend2, attrs }) => {
|
|
const { x, filter, dy } = inputs;
|
|
const { strides, pad: pad2, dilations } = attrs;
|
|
const cpuBackend = backend2;
|
|
const $x = toNestedArray(x.shape, cpuBackend.data.get(x.dataId).values);
|
|
const $filter = toNestedArray(filter.shape, cpuBackend.data.get(filter.dataId).values);
|
|
const { batchSize, inHeight, inWidth, inChannels, outHeight, outWidth, padInfo, strideHeight, strideWidth, filterHeight, filterWidth, dilationHeight, dilationWidth, outShape } = computeDilation2DInfo(x.shape, filter.shape, strides, pad2, "NHWC", dilations);
|
|
assert(dy.rank === outShape.length, () => `Error in ${Dilation2DBackpropFilter}, dy must have the same rank as output ${outShape.length}, but got ${dy.rank}`);
|
|
const $dy = toNestedArray(outShape, cpuBackend.data.get(dy.dataId).values);
|
|
const gradients = makeZerosNestedTypedArray(filter.shape, filter.dtype);
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let hOut = 0; hOut < outHeight; ++hOut) {
|
|
const hBeg = hOut * strideHeight - padInfo.top;
|
|
for (let wOut = 0; wOut < outWidth; ++wOut) {
|
|
const wBeg = wOut * strideWidth - padInfo.left;
|
|
for (let d = 0; d < inChannels; ++d) {
|
|
let curVal = Number.MIN_SAFE_INTEGER;
|
|
let hMax = 0;
|
|
let wMax = 0;
|
|
for (let h = 0; h < filterHeight; ++h) {
|
|
const hIn = hBeg + h * dilationHeight;
|
|
if (hIn >= 0 && hIn < inHeight) {
|
|
for (let w = 0; w < filterWidth; ++w) {
|
|
const wIn = wBeg + w * dilationWidth;
|
|
if (wIn >= 0 && wIn < inWidth) {
|
|
const val = $x[b][hIn][wIn][d] + $filter[h][w][d];
|
|
if (val > curVal) {
|
|
curVal = val;
|
|
hMax = h;
|
|
wMax = w;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
gradients[hMax][wMax][d] += $dy[b][hOut][wOut][d];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const dataId = cpuBackend.write(toTypedArray(gradients, x.dtype), filter.shape, filter.dtype);
|
|
return { dataId, shape: filter.shape, dtype: filter.dtype };
|
|
}
|
|
};
|
|
const dilation2DBackpropInputConfig = {
|
|
kernelName: Dilation2DBackpropInput,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, backend: backend2, attrs }) => {
|
|
const { x, filter, dy } = inputs;
|
|
const { strides, pad: pad2, dilations } = attrs;
|
|
const cpuBackend = backend2;
|
|
const $x = toNestedArray(x.shape, cpuBackend.data.get(x.dataId).values);
|
|
const $filter = toNestedArray(filter.shape, cpuBackend.data.get(filter.dataId).values);
|
|
const { batchSize, inHeight, inWidth, inChannels, outHeight, outWidth, padInfo, strideHeight, strideWidth, filterHeight, filterWidth, dilationHeight, dilationWidth, outShape } = computeDilation2DInfo(x.shape, filter.shape, strides, pad2, "NHWC", dilations);
|
|
assert(dy.rank === outShape.length, () => `Error in ${Dilation2DBackpropInput}, dy must have the same rank as output ${outShape.length}, but got ${dy.rank}`);
|
|
const $dy = toNestedArray(outShape, cpuBackend.data.get(dy.dataId).values);
|
|
const gradients = makeZerosNestedTypedArray(x.shape, x.dtype);
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
for (let hOut = 0; hOut < outHeight; ++hOut) {
|
|
const hBeg = hOut * strideHeight - padInfo.top;
|
|
for (let wOut = 0; wOut < outWidth; ++wOut) {
|
|
const wBeg = wOut * strideWidth - padInfo.left;
|
|
for (let d = 0; d < inChannels; ++d) {
|
|
let curVal = Number.MIN_SAFE_INTEGER;
|
|
let hInMax = hBeg < 0 ? 0 : hBeg;
|
|
let wInMax = wBeg < 0 ? 0 : wBeg;
|
|
for (let h = 0; h < filterHeight; ++h) {
|
|
const hIn = hBeg + h * dilationHeight;
|
|
if (hIn >= 0 && hIn < inHeight) {
|
|
for (let w = 0; w < filterWidth; ++w) {
|
|
const wIn = wBeg + w * dilationWidth;
|
|
if (wIn >= 0 && wIn < inWidth) {
|
|
const val = $x[b][hIn][wIn][d] + $filter[h][w][d];
|
|
if (val > curVal) {
|
|
curVal = val;
|
|
hInMax = hIn;
|
|
wInMax = wIn;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
gradients[b][hInMax][wInMax][d] += $dy[b][hOut][wOut][d];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const dataId = cpuBackend.write(toTypedArray(gradients, x.dtype), x.shape, x.dtype);
|
|
return { dataId, shape: x.shape, dtype: x.dtype };
|
|
}
|
|
};
|
|
function draw(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { image: image2 } = inputs;
|
|
const { canvas, options } = attrs;
|
|
const { contextOptions, imageOptions } = options || {};
|
|
const alpha = (imageOptions === null || imageOptions === void 0 ? void 0 : imageOptions.alpha) || 1;
|
|
const contextType = (contextOptions === null || contextOptions === void 0 ? void 0 : contextOptions.contextType) || "2d";
|
|
if (contextType !== "2d") {
|
|
throw new Error(`Context type ${contextOptions.contextType} is not supported by the CPU backend.`);
|
|
}
|
|
const ctx = canvas.getContext(contextType, (contextOptions === null || contextOptions === void 0 ? void 0 : contextOptions.contextAttributes) || {});
|
|
if (ctx == null) {
|
|
throw new Error(`Could not get the context with ${contextType} type.`);
|
|
}
|
|
const [height, width] = image2.shape.slice(0, 2);
|
|
const depth = image2.shape.length === 2 ? 1 : image2.shape[2];
|
|
const data = backend2.data.get(image2.dataId).values;
|
|
const multiplier = image2.dtype === "float32" ? 255 : 1;
|
|
const bytes = new Uint8ClampedArray(width * height * 4);
|
|
for (let i = 0; i < height * width; ++i) {
|
|
const rgba = [0, 0, 0, 255 * alpha];
|
|
for (let d = 0; d < depth; d++) {
|
|
const value = data[i * depth + d];
|
|
if (image2.dtype === "float32") {
|
|
if (value < 0 || value > 1) {
|
|
throw new Error(`Tensor values for a float32 Tensor must be in the range [0 - 1] but encountered ${value}.`);
|
|
}
|
|
} else if (image2.dtype === "int32") {
|
|
if (value < 0 || value > 255) {
|
|
throw new Error(`Tensor values for a int32 Tensor must be in the range [0 - 255] but encountered ${value}.`);
|
|
}
|
|
}
|
|
if (depth === 1) {
|
|
rgba[0] = value * multiplier;
|
|
rgba[1] = value * multiplier;
|
|
rgba[2] = value * multiplier;
|
|
} else {
|
|
rgba[d] = value * multiplier;
|
|
}
|
|
}
|
|
const j2 = i * 4;
|
|
bytes[j2 + 0] = Math.round(rgba[0]);
|
|
bytes[j2 + 1] = Math.round(rgba[1]);
|
|
bytes[j2 + 2] = Math.round(rgba[2]);
|
|
bytes[j2 + 3] = Math.round(rgba[3]);
|
|
}
|
|
canvas.width = width;
|
|
canvas.height = height;
|
|
const imageData = new ImageData(bytes, width, height);
|
|
ctx.putImageData(imageData, 0, 0);
|
|
return image2;
|
|
}
|
|
const drawConfig = {
|
|
kernelName: Draw,
|
|
backendName: "cpu",
|
|
kernelFunc: draw
|
|
};
|
|
function sum(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
assertNotComplex(x, "sum");
|
|
let $x;
|
|
if (x.dtype === "bool") {
|
|
$x = cast$1({ inputs: { x }, backend: backend2, attrs: { dtype: "int32" } });
|
|
} else {
|
|
$x = identity$1({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const xRank = $x.shape.length;
|
|
const axes = parseAxisParam(axis, $x.shape);
|
|
const permutation = getAxesPermutation(axes, xRank);
|
|
let reductionAxes = axes;
|
|
let permutedX = $x;
|
|
if (permutation != null) {
|
|
permutedX = transpose$1({ inputs: { x: $x }, backend: backend2, attrs: { perm: permutation } });
|
|
reductionAxes = getInnerMostAxes(reductionAxes.length, xRank);
|
|
}
|
|
assertAxesAreInnerMostDims("sum", reductionAxes, permutedX.shape.length);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes(permutedX.shape, reductionAxes);
|
|
const resultDtype = upcastType(permutedX.dtype, "int32");
|
|
let result = zeros(backend2, outShape, resultDtype);
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const vals = backend2.data.get(result.dataId).values;
|
|
const aVals = backend2.data.get(permutedX.dataId).values;
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let sum2 = 0;
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
sum2 += aVals[offset + j2];
|
|
}
|
|
vals[i] = sum2;
|
|
}
|
|
if (keepDims) {
|
|
const newShape = expandShapeToKeepDim(result.shape, axes);
|
|
const oldResult = result;
|
|
result = reshape({ inputs: { x: result }, backend: backend2, attrs: { shape: newShape } });
|
|
backend2.disposeIntermediateTensorInfo(oldResult);
|
|
}
|
|
backend2.disposeIntermediateTensorInfo($x);
|
|
if (permutation != null) {
|
|
backend2.disposeIntermediateTensorInfo(permutedX);
|
|
}
|
|
return result;
|
|
}
|
|
const sumConfig = {
|
|
kernelName: Sum,
|
|
backendName: "cpu",
|
|
kernelFunc: sum
|
|
};
|
|
function einsum(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { equation } = attrs;
|
|
const tensors = inputs;
|
|
const { allDims, summedDims, idDims } = decodeEinsumEquation(equation, tensors.length);
|
|
checkEinsumDimSizes(allDims.length, idDims, tensors);
|
|
const { path, steps } = getEinsumComputePath(summedDims, idDims);
|
|
const nSteps = steps.length;
|
|
let out = null;
|
|
let numDimsRemaining = allDims.length;
|
|
const tensorsToDispose = [];
|
|
for (let i = 0; i < nSteps; ++i) {
|
|
for (const idTerm of steps[i]) {
|
|
const { permutationIndices: perm, expandDims: dimsToExpand } = getEinsumPermutation(numDimsRemaining, idDims[idTerm]);
|
|
let x;
|
|
if (isIdentityPermutation(perm)) {
|
|
x = tensors[idTerm];
|
|
} else {
|
|
x = transpose$1({ inputs: { x: tensors[idTerm] }, backend: backend2, attrs: { perm } });
|
|
tensorsToDispose.push(x);
|
|
}
|
|
const targetShape = x.shape.slice();
|
|
for (let k = 0; k < dimsToExpand.length; ++k) {
|
|
targetShape.splice(dimsToExpand[k], 0, 1);
|
|
}
|
|
if (!arraysEqual(x.shape, targetShape)) {
|
|
x = reshape({ inputs: { x }, backend: backend2, attrs: { shape: targetShape } });
|
|
tensorsToDispose.push(x);
|
|
}
|
|
if (out === null) {
|
|
out = x;
|
|
} else {
|
|
out = multiply$1({ inputs: { a: x, b: out }, backend: backend2 });
|
|
tensorsToDispose.push(out);
|
|
}
|
|
}
|
|
if (i < nSteps - 1) {
|
|
if (path[i] >= 0) {
|
|
out = sum({
|
|
inputs: { x: out },
|
|
backend: backend2,
|
|
attrs: {
|
|
axis: path[i] - (allDims.length - numDimsRemaining),
|
|
keepDims: false
|
|
}
|
|
});
|
|
tensorsToDispose.push(out);
|
|
}
|
|
numDimsRemaining--;
|
|
}
|
|
}
|
|
for (const tensorInfo of tensorsToDispose) {
|
|
if (tensorInfo === out) {
|
|
continue;
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(tensorInfo);
|
|
}
|
|
return out;
|
|
}
|
|
const einsumConfig = {
|
|
kernelName: Einsum,
|
|
backendName: "cpu",
|
|
kernelFunc: einsum
|
|
};
|
|
function eluGrad(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { dy, y } = inputs;
|
|
assertNotComplex([dy, y], "eluGrad");
|
|
const resultValues = new Float32Array(sizeFromShape(y.shape));
|
|
const values = backend2.data.get(y.dataId).values;
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
for (let i = 0; i < values.length; ++i) {
|
|
const v = values[i];
|
|
if (v >= 0) {
|
|
resultValues[i] = dyValues[i];
|
|
} else {
|
|
resultValues[i] = dyValues[i] * (v + 1);
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(y.shape, "float32", resultValues);
|
|
}
|
|
const eluGradConfig = {
|
|
kernelName: EluGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: eluGrad
|
|
};
|
|
const p = ERF_P;
|
|
const a1 = ERF_A1;
|
|
const a2 = ERF_A2;
|
|
const a3 = ERF_A3;
|
|
const a4 = ERF_A4;
|
|
const a5 = ERF_A5;
|
|
const erf = unaryKernelFunc$1(Erf, (xi) => {
|
|
const sign2 = Math.sign(xi);
|
|
const v = Math.abs(xi);
|
|
const t2 = 1 / (1 + p * v);
|
|
return sign2 * (1 - ((((a5 * t2 + a4) * t2 + a3) * t2 + a2) * t2 + a1) * t2 * Math.exp(-v * v));
|
|
});
|
|
const erfConfig = {
|
|
kernelName: Erf,
|
|
backendName: "cpu",
|
|
kernelFunc: erf
|
|
};
|
|
function expandDims(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { input } = inputs;
|
|
const { dim } = attrs;
|
|
const inputRank = input.shape.length;
|
|
const newShape = input.shape.slice();
|
|
let $dim = dim;
|
|
if (dim < 0) {
|
|
assert(-(inputRank + 1) <= dim, () => `Axis must be in the interval [${-(inputRank + 1)}, ${inputRank}]`);
|
|
$dim = inputRank + dim + 1;
|
|
}
|
|
newShape.splice($dim, 0, 1);
|
|
return reshape({ inputs: { x: input }, backend: backend2, attrs: { shape: newShape } });
|
|
}
|
|
const expandDimsConfig = {
|
|
kernelName: ExpandDims,
|
|
backendName: "cpu",
|
|
kernelFunc: expandDims
|
|
};
|
|
const realDivImpl = createSimpleBinaryKernelImpl((a, b) => a / b);
|
|
const div = binaryKernelFunc$1(RealDiv, realDivImpl);
|
|
const realDivConfig = {
|
|
kernelName: RealDiv,
|
|
backendName: "cpu",
|
|
kernelFunc: div
|
|
};
|
|
function fftBatch(input, inverse, cpuBackend) {
|
|
const inputShape = input.shape;
|
|
const batch = inputShape[0];
|
|
const innerDim = inputShape[1];
|
|
const inputVals = cpuBackend.data.get(input.dataId);
|
|
const real2D = inputVals.complexTensorInfos.real;
|
|
const imag2D = inputVals.complexTensorInfos.imag;
|
|
const resultShape = [batch, innerDim];
|
|
const resultSize = sizeFromShape(resultShape);
|
|
const resultReal = getTypedArrayFromDType("float32", resultSize);
|
|
const resultImag = getTypedArrayFromDType("float32", resultSize);
|
|
for (let b = 0; b < batch; b++) {
|
|
const r = slice$1({
|
|
inputs: { x: real2D },
|
|
backend: cpuBackend,
|
|
attrs: { begin: [b, 0], size: [1, innerDim] }
|
|
});
|
|
const i = slice$1({
|
|
inputs: { x: imag2D },
|
|
backend: cpuBackend,
|
|
attrs: { begin: [b, 0], size: [1, innerDim] }
|
|
});
|
|
const input2 = complex$1({ inputs: { real: r, imag: i }, backend: cpuBackend });
|
|
const { real: real2, imag: imag2 } = fftImpl(input2, inverse, cpuBackend);
|
|
const res = mergeRealAndImagArrays(real2, imag2);
|
|
for (let d = 0; d < innerDim; d++) {
|
|
const c = getComplexWithIndex(res, d);
|
|
resultReal[b * innerDim + d] = c.real;
|
|
resultImag[b * innerDim + d] = c.imag;
|
|
}
|
|
cpuBackend.disposeIntermediateTensorInfo(r);
|
|
cpuBackend.disposeIntermediateTensorInfo(i);
|
|
cpuBackend.disposeIntermediateTensorInfo(input2);
|
|
}
|
|
const $realInfo = cpuBackend.makeTensorInfo(resultShape, "float32", resultReal);
|
|
const $imagInfo = cpuBackend.makeTensorInfo(resultShape, "float32", resultImag);
|
|
const result = complex$1({ inputs: { real: $realInfo, imag: $imagInfo }, backend: cpuBackend });
|
|
cpuBackend.disposeIntermediateTensorInfo($realInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($imagInfo);
|
|
return result;
|
|
}
|
|
function fftImpl(input, inverse, cpuBackend) {
|
|
const inputSize = sizeFromShape(input.shape);
|
|
const inputVals = cpuBackend.data.get(input.dataId);
|
|
const realVals = cpuBackend.data.get(inputVals.complexTensorInfos.real.dataId).values;
|
|
const imagVals = cpuBackend.data.get(inputVals.complexTensorInfos.imag.dataId).values;
|
|
if (isExponentOf2(inputSize)) {
|
|
const result = fftRadix2(realVals, imagVals, inputSize, inverse, cpuBackend);
|
|
const resultShape = [input.shape[0], input.shape[1]];
|
|
if (inverse) {
|
|
const realInfo = cpuBackend.makeTensorInfo(resultShape, "float32", result.real);
|
|
const imagInfo = cpuBackend.makeTensorInfo(resultShape, "float32", result.imag);
|
|
const sizeInfo = cpuBackend.makeTensorInfo([], "float32", createScalarValue(inputSize, "float32"));
|
|
const sizeInfoCopy = identity$1({ inputs: { x: sizeInfo }, backend: cpuBackend });
|
|
const divRealInfo = realDivConfig.kernelFunc({ inputs: { a: realInfo, b: sizeInfo }, backend: cpuBackend });
|
|
const divImagInfo = realDivConfig.kernelFunc({ inputs: { a: imagInfo, b: sizeInfoCopy }, backend: cpuBackend });
|
|
const divRealVals = cpuBackend.data.get(divRealInfo.dataId).values;
|
|
const divImagVals = cpuBackend.data.get(divImagInfo.dataId).values;
|
|
cpuBackend.disposeIntermediateTensorInfo(realInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(imagInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(sizeInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(sizeInfoCopy);
|
|
cpuBackend.disposeIntermediateTensorInfo(divRealInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(divImagInfo);
|
|
return { real: divRealVals, imag: divImagVals };
|
|
}
|
|
return result;
|
|
} else {
|
|
const data = mergeRealAndImagArrays(realVals, imagVals);
|
|
const rawOutput = fourierTransformByMatmul(data, inputSize, inverse);
|
|
return splitRealAndImagArrays(rawOutput);
|
|
}
|
|
}
|
|
function isExponentOf2(size) {
|
|
return (size & size - 1) === 0;
|
|
}
|
|
function fftRadix2(realVals, imagVals, size, inverse, cpuBackend) {
|
|
if (size === 1) {
|
|
return { real: realVals, imag: imagVals };
|
|
}
|
|
const data = mergeRealAndImagArrays(realVals, imagVals);
|
|
const half = size / 2;
|
|
const evenComplex = complexWithEvenIndex(data);
|
|
const evenRealVals = evenComplex.real;
|
|
const evenImagVals = evenComplex.imag;
|
|
const evenShape = [evenRealVals.length];
|
|
const evenRealInfo = cpuBackend.makeTensorInfo(evenShape, "float32", evenRealVals);
|
|
const evenImagInfo = cpuBackend.makeTensorInfo(evenShape, "float32", evenImagVals);
|
|
const evenTensorInfo = complex$1({ inputs: { real: evenRealInfo, imag: evenImagInfo }, backend: cpuBackend });
|
|
const oddComplex = complexWithOddIndex(data);
|
|
const oddRealVals = oddComplex.real;
|
|
const oddImagVals = oddComplex.imag;
|
|
const oddShape = [oddRealVals.length];
|
|
const oddRealInfo = cpuBackend.makeTensorInfo(oddShape, "float32", oddRealVals);
|
|
const oddImagInfo = cpuBackend.makeTensorInfo(oddShape, "float32", oddImagVals);
|
|
const oddTensorInfo = complex$1({ inputs: { real: oddRealInfo, imag: oddImagInfo }, backend: cpuBackend });
|
|
const $evenComplex = fftRadix2(evenRealVals, evenImagVals, half, inverse, cpuBackend);
|
|
const $evenRealVals = $evenComplex.real;
|
|
const $evenImagVals = $evenComplex.imag;
|
|
const $evenShape = [$evenRealVals.length];
|
|
const $evenRealInfo = cpuBackend.makeTensorInfo($evenShape, "float32", $evenRealVals);
|
|
const $evenImagInfo = cpuBackend.makeTensorInfo($evenShape, "float32", $evenImagVals);
|
|
const $evenTensorInfo = complex$1({
|
|
inputs: { real: $evenRealInfo, imag: $evenImagInfo },
|
|
backend: cpuBackend
|
|
});
|
|
const $oddComplex = fftRadix2(oddRealVals, oddImagVals, half, inverse, cpuBackend);
|
|
const $oddRealVals = $oddComplex.real;
|
|
const $oddImagVals = $oddComplex.imag;
|
|
const $oddShape = [$oddRealVals.length];
|
|
const $oddRealInfo = cpuBackend.makeTensorInfo($oddShape, "float32", $oddRealVals);
|
|
const $oddImagInfo = cpuBackend.makeTensorInfo($oddShape, "float32", $oddImagVals);
|
|
const $oddTensorInfo = complex$1({ inputs: { real: $oddRealInfo, imag: $oddImagInfo }, backend: cpuBackend });
|
|
const e = exponents(size, inverse);
|
|
const eShape = [e.real.length];
|
|
const eRealInfo = cpuBackend.makeTensorInfo(eShape, "float32", e.real);
|
|
const eImagInfo = cpuBackend.makeTensorInfo(eShape, "float32", e.imag);
|
|
const complexInfo = complex$1({ inputs: { real: eRealInfo, imag: eImagInfo }, backend: cpuBackend });
|
|
const exponentInfo = multiply$1({ inputs: { a: complexInfo, b: $oddTensorInfo }, backend: cpuBackend });
|
|
const addPart = add({
|
|
inputs: { a: $evenTensorInfo, b: exponentInfo },
|
|
backend: cpuBackend
|
|
});
|
|
const subPart = sub$1({
|
|
inputs: { a: $evenTensorInfo, b: exponentInfo },
|
|
backend: cpuBackend
|
|
});
|
|
const addPartReal = real$1({ inputs: { input: addPart }, backend: cpuBackend });
|
|
const subPartReal = real$1({ inputs: { input: subPart }, backend: cpuBackend });
|
|
const addPartImag = imag({ inputs: { input: addPart }, backend: cpuBackend });
|
|
const subPartImag = imag({ inputs: { input: subPart }, backend: cpuBackend });
|
|
const $real = concat({
|
|
inputs: [addPartReal, subPartReal],
|
|
backend: cpuBackend,
|
|
attrs: { axis: 0 }
|
|
});
|
|
const $imag = concat({
|
|
inputs: [addPartImag, subPartImag],
|
|
backend: cpuBackend,
|
|
attrs: { axis: 0 }
|
|
});
|
|
const $realVals = cpuBackend.data.get($real.dataId).values;
|
|
const $imagVals = cpuBackend.data.get($imag.dataId).values;
|
|
cpuBackend.disposeIntermediateTensorInfo(evenRealInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(evenImagInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(evenTensorInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(oddRealInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(oddImagInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(oddTensorInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($evenRealInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($evenImagInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($evenTensorInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($oddRealInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($oddImagInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo($oddTensorInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(eRealInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(eImagInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(complexInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(exponentInfo);
|
|
cpuBackend.disposeIntermediateTensorInfo(addPart);
|
|
cpuBackend.disposeIntermediateTensorInfo(subPart);
|
|
cpuBackend.disposeIntermediateTensorInfo(addPartReal);
|
|
cpuBackend.disposeIntermediateTensorInfo(addPartImag);
|
|
cpuBackend.disposeIntermediateTensorInfo(subPartReal);
|
|
cpuBackend.disposeIntermediateTensorInfo(subPartImag);
|
|
cpuBackend.disposeIntermediateTensorInfo($real);
|
|
cpuBackend.disposeIntermediateTensorInfo($imag);
|
|
return { real: $realVals, imag: $imagVals };
|
|
}
|
|
function fourierTransformByMatmul(data, size, inverse) {
|
|
const ret = new Float32Array(size * 2);
|
|
for (let r = 0; r < size; r++) {
|
|
let real2 = 0;
|
|
let imag2 = 0;
|
|
for (let c = 0; c < size; c++) {
|
|
const e = exponent(r * c, size, inverse);
|
|
const term = getComplexWithIndex(data, c);
|
|
real2 += term.real * e.real - term.imag * e.imag;
|
|
imag2 += term.real * e.imag + term.imag * e.real;
|
|
}
|
|
if (inverse) {
|
|
real2 /= size;
|
|
imag2 /= size;
|
|
}
|
|
assignToTypedArray(ret, real2, imag2, r);
|
|
}
|
|
return ret;
|
|
}
|
|
function fft(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
const inputSize = sizeFromShape(input.shape);
|
|
const innerDimensionSize = input.shape[input.shape.length - 1];
|
|
const batch = inputSize / innerDimensionSize;
|
|
const input2D = reshape({
|
|
inputs: { x: input },
|
|
backend: backend2,
|
|
attrs: { shape: [batch, innerDimensionSize] }
|
|
});
|
|
const result = fftBatch(input2D, false, backend2);
|
|
const resultReshaped = reshape({ inputs: { x: result }, backend: backend2, attrs: { shape: input.shape } });
|
|
backend2.disposeIntermediateTensorInfo(input2D);
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return resultReshaped;
|
|
}
|
|
const fftConfig = {
|
|
kernelName: FFT,
|
|
backendName: "cpu",
|
|
kernelFunc: fft
|
|
};
|
|
function fill(args) {
|
|
const { backend: backend2, attrs } = args;
|
|
const { shape, value, dtype } = attrs;
|
|
const $dtype = dtype || inferDtype(value);
|
|
const values = getArrayFromDType($dtype, sizeFromShape(shape));
|
|
fillValues(values, value, $dtype);
|
|
return backend2.makeTensorInfo(shape, $dtype, values);
|
|
}
|
|
const fillConfig = {
|
|
kernelName: Fill,
|
|
backendName: "cpu",
|
|
kernelFunc: fill
|
|
};
|
|
function fillValues(values, value, dtype) {
|
|
if (dtype === "string") {
|
|
values.fill(value);
|
|
} else {
|
|
values.fill(value);
|
|
}
|
|
}
|
|
const flipLeftRightConfig = {
|
|
kernelName: FlipLeftRight,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, attrs, backend: backend2 }) => {
|
|
const { image: image2 } = inputs;
|
|
const cpuBackend = backend2;
|
|
const output = getTypedArrayFromDType(image2.dtype, sizeFromShape(image2.shape));
|
|
const [batch, imageHeight, imageWidth, numChannels] = image2.shape;
|
|
const imageVals = cpuBackend.data.get(image2.dataId).values;
|
|
for (let batchIdx = 0; batchIdx < batch; batchIdx++) {
|
|
const batchOffset = batchIdx * imageWidth * imageHeight * numChannels;
|
|
for (let row = 0; row < imageHeight; row++) {
|
|
const rowOffset = row * (imageWidth * numChannels);
|
|
for (let col = 0; col < imageWidth; col++) {
|
|
const colOffset = col * numChannels;
|
|
for (let channel = 0; channel < numChannels; channel++) {
|
|
const coordX = Math.round(imageWidth - col - 1);
|
|
const outIdx = batchOffset + rowOffset + colOffset + channel;
|
|
let outputValue = imageVals[outIdx];
|
|
if (coordX >= 0 && coordX < imageWidth) {
|
|
const rotatedColOffset = coordX * numChannels;
|
|
const imageIdx = batchOffset + rowOffset + rotatedColOffset + channel;
|
|
outputValue = imageVals[imageIdx];
|
|
}
|
|
output[outIdx] = outputValue;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const dataId = cpuBackend.write(output, image2.shape, image2.dtype);
|
|
return { dataId, shape: image2.shape, dtype: image2.dtype };
|
|
}
|
|
};
|
|
function fusedConv2D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter, bias, preluActivationWeights } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dilations, dimRoundingMode, activation, leakyreluAlpha } = attrs;
|
|
let result = conv2D({
|
|
inputs: { x, filter },
|
|
backend: backend2,
|
|
attrs: { strides, pad: pad2, dataFormat, dilations, dimRoundingMode }
|
|
});
|
|
if (bias) {
|
|
const resultOld = result;
|
|
if (dataFormat === "NCHW" && bias.shape.length === 1 && bias.shape[0] !== 1) {
|
|
const reshapedBias = reshape({ inputs: { x: bias }, backend: backend2, attrs: { shape: [bias.shape[0], 1, 1] } });
|
|
result = add({ inputs: { a: result, b: reshapedBias }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(reshapedBias);
|
|
} else {
|
|
result = add({ inputs: { a: result, b: bias }, backend: backend2 });
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(resultOld);
|
|
}
|
|
if (activation) {
|
|
const resultOld = result;
|
|
if (dataFormat === "NCHW" && activation === "prelu" && preluActivationWeights.shape.length === 1 && preluActivationWeights.shape[0] !== 1) {
|
|
const reshapedAlpha = reshape({
|
|
inputs: { x: preluActivationWeights },
|
|
backend: backend2,
|
|
attrs: { shape: [preluActivationWeights.shape[0], 1, 1] }
|
|
});
|
|
result = applyActivation(backend2, result, activation, reshapedAlpha, leakyreluAlpha);
|
|
backend2.disposeIntermediateTensorInfo(reshapedAlpha);
|
|
} else {
|
|
result = applyActivation(backend2, result, activation, preluActivationWeights, leakyreluAlpha);
|
|
}
|
|
backend2.disposeIntermediateTensorInfo(resultOld);
|
|
}
|
|
return result;
|
|
}
|
|
const fusedConv2DConfig = {
|
|
kernelName: FusedConv2D,
|
|
backendName: "cpu",
|
|
kernelFunc: fusedConv2D
|
|
};
|
|
function fusedDepthwiseConv2D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, filter, bias, preluActivationWeights } = inputs;
|
|
const { strides, pad: pad2, dataFormat, dilations, dimRoundingMode, activation, leakyreluAlpha } = attrs;
|
|
let result = depthwiseConv2dNative({
|
|
inputs: { x, filter },
|
|
backend: backend2,
|
|
attrs: { strides, pad: pad2, dataFormat, dilations, dimRoundingMode }
|
|
});
|
|
if (bias) {
|
|
const oldResult = result;
|
|
result = add({ inputs: { a: result, b: bias }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(oldResult);
|
|
}
|
|
if (activation) {
|
|
const oldResult = result;
|
|
result = applyActivation(backend2, result, activation, preluActivationWeights, leakyreluAlpha);
|
|
backend2.disposeIntermediateTensorInfo(oldResult);
|
|
}
|
|
return result;
|
|
}
|
|
const fusedDepthwiseConv2DConfig = {
|
|
kernelName: FusedDepthwiseConv2D,
|
|
backendName: "cpu",
|
|
kernelFunc: fusedDepthwiseConv2D
|
|
};
|
|
function gatherNd(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { params, indices } = inputs;
|
|
const paramsSize = sizeFromShape(params.shape);
|
|
const indicesShape = indices.shape;
|
|
const sliceRank = indicesShape[indicesShape.length - 1];
|
|
const [resultShape, numSlices, sliceSize, strides] = prepareAndValidate(params, indices);
|
|
if (numSlices === 0) {
|
|
return backend2.makeTensorInfo(resultShape, params.dtype, []);
|
|
}
|
|
const indicesData = backend2.data.get(indices.dataId).values;
|
|
const paramsBuf = backend2.bufferSync(params);
|
|
const outBuf = gatherNdImpl(indicesData, paramsBuf, params.dtype, numSlices, sliceRank, sliceSize, strides, params.shape, paramsSize);
|
|
return backend2.makeTensorInfo(resultShape, params.dtype, outBuf.values);
|
|
}
|
|
const gatherNdConfig = {
|
|
kernelName: GatherNd,
|
|
backendName: "cpu",
|
|
kernelFunc: gatherNd
|
|
};
|
|
function gatherV2(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, indices } = inputs;
|
|
const { axis, batchDims } = attrs;
|
|
assertNotComplex([x, indices], "gatherV2");
|
|
const parsedAxis = parseAxisParam(axis, x.shape)[0];
|
|
const indicesVals = backend2.data.get(indices.dataId).values;
|
|
const axisDim = x.shape[parsedAxis];
|
|
for (let i = 0; i < indicesVals.length; ++i) {
|
|
const index = indicesVals[i];
|
|
assert(index <= axisDim - 1 && index >= 0, () => `GatherV2: the index value ${index} is not in [0, ${axisDim - 1}]`);
|
|
}
|
|
let $batchDims = batchDims;
|
|
if (batchDims == null) {
|
|
$batchDims = 0;
|
|
}
|
|
const indicesSize = sizeFromShape(indices.shape);
|
|
const shapeInfo = collectGatherOpShapeInfo(x, indices, parsedAxis, $batchDims);
|
|
const flattenX = reshape({
|
|
inputs: { x },
|
|
backend: backend2,
|
|
attrs: {
|
|
shape: [
|
|
shapeInfo.batchSize,
|
|
shapeInfo.outerSize,
|
|
shapeInfo.dimSize,
|
|
shapeInfo.sliceSize
|
|
]
|
|
}
|
|
});
|
|
const flattenIndex = reshape({
|
|
inputs: { x: indices },
|
|
backend: backend2,
|
|
attrs: { shape: [shapeInfo.batchSize, indicesSize / shapeInfo.batchSize] }
|
|
});
|
|
const flattenOutputShape = [
|
|
shapeInfo.batchSize,
|
|
shapeInfo.outerSize,
|
|
indicesSize / shapeInfo.batchSize,
|
|
shapeInfo.sliceSize
|
|
];
|
|
const indicesBuf = backend2.bufferSync(flattenIndex);
|
|
const xBuf = backend2.bufferSync(flattenX);
|
|
const outBuf = gatherV2Impl(xBuf, indicesBuf, flattenOutputShape);
|
|
backend2.disposeIntermediateTensorInfo(flattenX);
|
|
backend2.disposeIntermediateTensorInfo(flattenIndex);
|
|
return backend2.makeTensorInfo(shapeInfo.outputShape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const gatherV2Config = {
|
|
kernelName: GatherV2,
|
|
backendName: "cpu",
|
|
kernelFunc: gatherV2
|
|
};
|
|
function ifft(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { input } = inputs;
|
|
const inputSize = sizeFromShape(input.shape);
|
|
const innerDimensionSize = input.shape[input.shape.length - 1];
|
|
const batch = inputSize / innerDimensionSize;
|
|
const input2D = reshape({
|
|
inputs: { x: input },
|
|
backend: backend2,
|
|
attrs: { shape: [batch, innerDimensionSize] }
|
|
});
|
|
const result = fftBatch(input2D, true, backend2);
|
|
const resultReshaped = reshape({ inputs: { x: result }, backend: backend2, attrs: { shape: input.shape } });
|
|
backend2.disposeIntermediateTensorInfo(input2D);
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return resultReshaped;
|
|
}
|
|
const ifftConfig = {
|
|
kernelName: IFFT,
|
|
backendName: "cpu",
|
|
kernelFunc: ifft
|
|
};
|
|
const isFinite$1 = unaryKernelFunc$1(IsFinite, (xi) => Number.isFinite(xi) ? 1 : 0, "bool");
|
|
const isFiniteConfig = {
|
|
kernelName: IsFinite,
|
|
backendName: "cpu",
|
|
kernelFunc: isFinite$1
|
|
};
|
|
const isInf = unaryKernelFunc$1(IsInf, (xi) => Math.abs(xi) === Infinity ? 1 : 0, "bool");
|
|
const isInfConfig = {
|
|
kernelName: IsInf,
|
|
backendName: "cpu",
|
|
kernelFunc: isInf
|
|
};
|
|
const isNaN$1 = unaryKernelFunc$1(IsNan, (xi) => Number.isNaN(xi) ? 1 : 0, "bool");
|
|
const isNaNConfig = {
|
|
kernelName: IsNan,
|
|
backendName: "cpu",
|
|
kernelFunc: isNaN$1
|
|
};
|
|
function linSpace(args) {
|
|
const { backend: backend2, attrs } = args;
|
|
const { start, stop, num } = attrs;
|
|
const outVals = linSpaceImpl(start, stop, num);
|
|
return backend2.makeTensorInfo([outVals.length], "float32", outVals);
|
|
}
|
|
const linSpaceConfig = {
|
|
kernelName: LinSpace,
|
|
backendName: "cpu",
|
|
kernelFunc: linSpace
|
|
};
|
|
const log1p = unaryKernelFunc$1(Log1p, (xi) => Math.log1p(xi));
|
|
const log1pConfig = {
|
|
kernelName: Log1p,
|
|
backendName: "cpu",
|
|
kernelFunc: log1p
|
|
};
|
|
const logicalAndImpl = createSimpleBinaryKernelImpl((a, b) => a && b);
|
|
const logicalAnd = binaryKernelFunc$1(LogicalAnd, logicalAndImpl, null, "bool");
|
|
const logicalAndConfig = {
|
|
kernelName: LogicalAnd,
|
|
backendName: "cpu",
|
|
kernelFunc: logicalAnd
|
|
};
|
|
const logicalNot = unaryKernelFunc$1(LogicalNot, (xi) => xi ? 0 : 1, "bool");
|
|
const logicalNotConfig = {
|
|
kernelName: LogicalNot,
|
|
backendName: "cpu",
|
|
kernelFunc: logicalNot
|
|
};
|
|
const logicalOrImpl = createSimpleBinaryKernelImpl((a, b) => a || b);
|
|
const logicalOr = binaryKernelFunc$1(LogicalOr, logicalOrImpl, null, "bool");
|
|
const logicalOrConfig = {
|
|
kernelName: LogicalOr,
|
|
backendName: "cpu",
|
|
kernelFunc: logicalOr
|
|
};
|
|
function lRN(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { depthRadius, bias, alpha, beta } = attrs;
|
|
assertNotComplex(x, "LRN");
|
|
const channels = x.shape[3];
|
|
const maxD = channels - 1;
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const size = sizeFromShape(x.shape);
|
|
const result = new Float32Array(size);
|
|
function sumAcrossChannels(offset) {
|
|
const currentChannel = offset % channels;
|
|
let beginSumOffset = offset - currentChannel + Math.max(0, currentChannel - depthRadius);
|
|
const endSumOffset = offset - currentChannel + Math.min(currentChannel + depthRadius, maxD);
|
|
let sum2 = 0;
|
|
for (; beginSumOffset <= endSumOffset; beginSumOffset++) {
|
|
const z2 = xValues[beginSumOffset];
|
|
sum2 += z2 * z2;
|
|
}
|
|
return sum2;
|
|
}
|
|
for (let offset = 0; offset < size; offset++) {
|
|
const sum2 = sumAcrossChannels(offset);
|
|
const val = xValues[offset] * Math.pow(bias + alpha * sum2, -beta);
|
|
result[offset] = val;
|
|
}
|
|
return backend2.makeTensorInfo(x.shape, x.dtype, result);
|
|
}
|
|
const LRNConfig = {
|
|
kernelName: LRN,
|
|
backendName: "cpu",
|
|
kernelFunc: lRN
|
|
};
|
|
function lRNGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, y, dy } = inputs;
|
|
const { depthRadius, bias, alpha, beta } = attrs;
|
|
assertNotComplex(dy, "LRNGrad");
|
|
const dySize = sizeFromShape(dy.shape);
|
|
const channels = dy.shape[3];
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const yValues = backend2.data.get(y.dataId).values;
|
|
const result = new Float32Array(dySize);
|
|
const size = dySize;
|
|
for (let offset = 0; offset < size; offset++) {
|
|
const currentChannel = offset % channels;
|
|
const depthBegin = offset - currentChannel + Math.max(0, currentChannel - depthRadius);
|
|
const depthEnd = offset - currentChannel + Math.min(channels, currentChannel + depthRadius + 1);
|
|
let norm2 = 0;
|
|
for (let k = depthBegin; k < depthEnd; k++) {
|
|
norm2 += Math.pow(xValues[k], 2);
|
|
}
|
|
norm2 = alpha * norm2 + bias;
|
|
for (let k = depthBegin; k < depthEnd; k++) {
|
|
let dyi = -2 * alpha * beta * xValues[k] * yValues[offset] / norm2;
|
|
if (offset === k) {
|
|
dyi += Math.pow(norm2, -beta);
|
|
}
|
|
dyi *= dyValues[offset];
|
|
result[k] += dyi;
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dy.shape, x.dtype, result);
|
|
}
|
|
const LRNGradConfig = {
|
|
kernelName: LRNGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: lRNGrad
|
|
};
|
|
function max(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { reductionIndices, keepDims } = attrs;
|
|
const cpuBackend = backend2;
|
|
let xShape = x.shape;
|
|
const xRank = xShape.length;
|
|
const origAxes = parseAxisParam(reductionIndices, xShape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, xRank);
|
|
let xVals = cpuBackend.data.get(x.dataId).values;
|
|
if (permutedAxes != null) {
|
|
const newShape = new Array(xRank);
|
|
for (let i = 0; i < newShape.length; i++) {
|
|
newShape[i] = xShape[permutedAxes[i]];
|
|
}
|
|
xVals = transposeImpl$1(xVals, xShape, x.dtype, permutedAxes, newShape);
|
|
axes = getInnerMostAxes(axes.length, xRank);
|
|
xShape = newShape;
|
|
}
|
|
assertNotComplex(x, "max");
|
|
assertAxesAreInnerMostDims("max", axes, xRank);
|
|
const [maxOutShape, reduceShape] = computeOutAndReduceShapes(xShape, axes);
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const result = maxImpl$1(xVals, reduceSize, maxOutShape, x.dtype);
|
|
const dataId = cpuBackend.write(result, maxOutShape, x.dtype);
|
|
let outShape = maxOutShape;
|
|
if (keepDims) {
|
|
const newShape = expandShapeToKeepDim(maxOutShape, origAxes);
|
|
outShape = newShape;
|
|
}
|
|
return { dataId, shape: outShape, dtype: x.dtype };
|
|
}
|
|
const maxConfig = {
|
|
kernelName: Max,
|
|
backendName: "cpu",
|
|
kernelFunc: max
|
|
};
|
|
function maxPool(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
assertNotComplex(x, "maxPool");
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const dilations = 1;
|
|
assert(eitherStridesOrDilationsAreOne(strides, dilations), () => `Error in maxPool: Either strides or dilations must be 1. Got strides ${strides} and dilations '${dilations}'`);
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, dilations, pad2, dimRoundingMode);
|
|
let res;
|
|
if (convInfo.filterWidth === 1 && convInfo.filterHeight === 1 && arraysEqual(convInfo.inShape, convInfo.outShape)) {
|
|
res = identity$1({ inputs: { x }, backend: backend2 });
|
|
} else {
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const strides2 = computeStrides(x.shape);
|
|
const buffer2 = pool(xValues, x.shape, x.dtype, strides2, convInfo, "max");
|
|
res = backend2.makeTensorInfo(convInfo.outShape, x.dtype, buffer2.values);
|
|
}
|
|
return res;
|
|
}
|
|
const maxPoolConfig = {
|
|
kernelName: MaxPool,
|
|
backendName: "cpu",
|
|
kernelFunc: maxPool
|
|
};
|
|
function maxPool3D(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode, dataFormat } = attrs;
|
|
assertNotComplex(x, "maxPool3d");
|
|
const convInfo = computePool3DInfo(x.shape, filterSize, strides, 1, pad2, dimRoundingMode, dataFormat);
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const outBuf = pool3d(xValues, x.shape, x.dtype, computeStrides(x.shape), convInfo, "max");
|
|
return backend2.makeTensorInfo(outBuf.shape, "float32", outBuf.values);
|
|
}
|
|
const maxPool3DConfig = {
|
|
kernelName: MaxPool3D,
|
|
backendName: "cpu",
|
|
kernelFunc: maxPool3D
|
|
};
|
|
function maxPool3DGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input } = inputs;
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
assertNotComplex([dy, input], "maxPool3DGrad");
|
|
const convInfo = computePool3DInfo(input.shape, filterSize, strides, 1, pad2, dimRoundingMode);
|
|
const inputBuf = backend2.bufferSync(input);
|
|
const maxPosBuf = maxPool3dPositions(inputBuf, convInfo);
|
|
const strideDepth = convInfo.strideDepth;
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationDepth = convInfo.dilationDepth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterDepth = convInfo.effectiveFilterDepth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padFront = effectiveFilterDepth - 1 - convInfo.padInfo.front;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const dx = buffer(input.shape, "float32");
|
|
const dyBuf = backend2.bufferSync(dy);
|
|
for (let batch = 0; batch < convInfo.batchSize; ++batch) {
|
|
for (let channel = 0; channel < convInfo.inChannels; ++channel) {
|
|
for (let dxDepth = 0; dxDepth < convInfo.inDepth; ++dxDepth) {
|
|
for (let dxRow = 0; dxRow < convInfo.inHeight; ++dxRow) {
|
|
for (let dxCol = 0; dxCol < convInfo.inWidth; ++dxCol) {
|
|
const dyDepthCorner = dxDepth - padFront;
|
|
const dyRowCorner = dxRow - padTop;
|
|
const dyColCorner = dxCol - padLeft;
|
|
let dotProd = 0;
|
|
for (let wDepth = 0; wDepth < effectiveFilterDepth; wDepth += dilationDepth) {
|
|
const dyDepth = (dyDepthCorner + wDepth) / strideDepth;
|
|
if (dyDepth < 0 || dyDepth >= convInfo.outDepth || Math.floor(dyDepth) !== dyDepth) {
|
|
continue;
|
|
}
|
|
for (let wRow = 0; wRow < effectiveFilterHeight; wRow += dilationHeight) {
|
|
const dyRow = (dyRowCorner + wRow) / strideHeight;
|
|
if (dyRow < 0 || dyRow >= convInfo.outHeight || Math.floor(dyRow) !== dyRow) {
|
|
continue;
|
|
}
|
|
for (let wCol = 0; wCol < effectiveFilterWidth; wCol += dilationWidth) {
|
|
const dyCol = (dyColCorner + wCol) / strideWidth;
|
|
if (dyCol < 0 || dyCol >= convInfo.outWidth || Math.floor(dyCol) !== dyCol) {
|
|
continue;
|
|
}
|
|
const maxPos = effectiveFilterDepth * effectiveFilterHeight * effectiveFilterWidth - 1 - maxPosBuf.get(batch, dyDepth, dyRow, dyCol, channel);
|
|
const curPos = wDepth * effectiveFilterHeight * effectiveFilterWidth + wRow * effectiveFilterWidth + wCol;
|
|
const mask = maxPos === curPos ? 1 : 0;
|
|
if (mask === 0) {
|
|
continue;
|
|
}
|
|
const pixel = dyBuf.get(batch, dyDepth, dyRow, dyCol, channel);
|
|
dotProd += pixel * mask;
|
|
}
|
|
}
|
|
}
|
|
dx.set(dotProd, batch, dxDepth, dxRow, dxCol, channel);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const maxPool3DGradConfig = {
|
|
kernelName: MaxPool3DGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: maxPool3DGrad
|
|
};
|
|
function maxPoolGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { dy, input, output } = inputs;
|
|
const x = input;
|
|
assertNotComplex([input, output], "maxPoolGrad");
|
|
const { filterSize, strides, pad: pad2, dimRoundingMode } = attrs;
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, 1, pad2, dimRoundingMode);
|
|
const xValues = backend2.data.get(x.dataId).values;
|
|
const maxPosBuf = buffer(convInfo.outShape, x.dtype, maxPoolPositions(xValues, x.shape, x.dtype, convInfo).values);
|
|
const strideHeight = convInfo.strideHeight;
|
|
const strideWidth = convInfo.strideWidth;
|
|
const dilationHeight = convInfo.dilationHeight;
|
|
const dilationWidth = convInfo.dilationWidth;
|
|
const effectiveFilterHeight = convInfo.effectiveFilterHeight;
|
|
const effectiveFilterWidth = convInfo.effectiveFilterWidth;
|
|
const padLeft = effectiveFilterWidth - 1 - convInfo.padInfo.left;
|
|
const padTop = effectiveFilterHeight - 1 - convInfo.padInfo.top;
|
|
const dx = buffer(x.shape, "float32");
|
|
const dyData = backend2.data.get(dy.dataId).values;
|
|
const dyBuf = buffer(dy.shape, "float32", dyData);
|
|
for (let b = 0; b < convInfo.batchSize; ++b) {
|
|
for (let d = 0; d < convInfo.inChannels; ++d) {
|
|
for (let dxR = 0; dxR < convInfo.inHeight; ++dxR) {
|
|
for (let dxC = 0; dxC < convInfo.inWidth; ++dxC) {
|
|
const dyRCorner = dxR - padTop;
|
|
const dyCCorner = dxC - padLeft;
|
|
let dotProd = 0;
|
|
for (let wR = 0; wR < effectiveFilterHeight; wR += dilationHeight) {
|
|
const dyR = (dyRCorner + wR) / strideHeight;
|
|
if (dyR < 0 || dyR >= convInfo.outHeight || Math.floor(dyR) !== dyR) {
|
|
continue;
|
|
}
|
|
for (let wC = 0; wC < effectiveFilterWidth; wC += dilationWidth) {
|
|
const dyC = (dyCCorner + wC) / strideWidth;
|
|
if (dyC < 0 || dyC >= convInfo.outWidth || Math.floor(dyC) !== dyC) {
|
|
continue;
|
|
}
|
|
const maxPos = effectiveFilterHeight * effectiveFilterWidth - 1 - maxPosBuf.get(b, dyR, dyC, d);
|
|
const curPos = wR * effectiveFilterWidth + wC;
|
|
const mask = maxPos === curPos ? 1 : 0;
|
|
if (mask === 0) {
|
|
continue;
|
|
}
|
|
const pixel = dyBuf.get(b, dyR, dyC, d);
|
|
dotProd += pixel * mask;
|
|
}
|
|
}
|
|
dx.set(dotProd, b, dxR, dxC, d);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(dx.shape, dx.dtype, dx.values);
|
|
}
|
|
const maxPoolGradConfig = {
|
|
kernelName: MaxPoolGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: maxPoolGrad
|
|
};
|
|
function maxPoolWithArgmaxImpl(xValues, xShape, dtype, includeBatchInIndex, convInfo) {
|
|
const strides = computeStrides(xShape);
|
|
const maxPools = pool(xValues, xShape, dtype, strides, convInfo, "max");
|
|
const maxPositions = maxPoolPositions(xValues, xShape, dtype, convInfo, true, includeBatchInIndex);
|
|
return [maxPools.values, maxPositions.values];
|
|
}
|
|
const maxPoolWithArgmaxConfig = {
|
|
kernelName: MaxPoolWithArgmax,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, attrs, backend: backend2 }) => {
|
|
const { x } = inputs;
|
|
const { filterSize, strides, pad: pad2, includeBatchInIndex } = attrs;
|
|
const cpuBackend = backend2;
|
|
assertNotComplex(x, "MaxPoolWithArgmax");
|
|
const values = cpuBackend.data.get(x.dataId).values;
|
|
const convInfo = computePool2DInfo(x.shape, filterSize, strides, [1, 1], pad2);
|
|
const [pooled, indexes] = maxPoolWithArgmaxImpl(values, x.shape, x.dtype, includeBatchInIndex, convInfo);
|
|
const pooledDataId = cpuBackend.write(pooled, convInfo.outShape, x.dtype);
|
|
const indexesDataId = cpuBackend.write(indexes, convInfo.outShape, x.dtype);
|
|
return [
|
|
{ dataId: pooledDataId, shape: convInfo.outShape, dtype: x.dtype },
|
|
{ dataId: indexesDataId, shape: convInfo.outShape, dtype: "int32" }
|
|
];
|
|
}
|
|
};
|
|
function mean(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
const axes = parseAxisParam(axis, x.shape);
|
|
const shapes = computeOutAndReduceShapes(x.shape, axes);
|
|
const reduceShape = shapes[1];
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const toDispose = [];
|
|
const reduceSizeScalar = backend2.makeTensorInfo([], "float32", new Float32Array([reduceSize]));
|
|
toDispose.push(reduceSizeScalar);
|
|
const $x = cast$1({ inputs: { x }, backend: backend2, attrs: { dtype: "float32" } });
|
|
toDispose.push($x);
|
|
const res = div({ inputs: { a: $x, b: reduceSizeScalar }, backend: backend2 });
|
|
toDispose.push(res);
|
|
const result = sum({ inputs: { x: res }, backend: backend2, attrs: { axis, keepDims } });
|
|
toDispose.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
}
|
|
const meanConfig = {
|
|
kernelName: Mean,
|
|
backendName: "cpu",
|
|
kernelFunc: mean
|
|
};
|
|
function min(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { axis, keepDims } = attrs;
|
|
assertNotComplex(x, "min");
|
|
const origAxes = parseAxisParam(axis, x.shape);
|
|
let axes = origAxes;
|
|
const permutedAxes = getAxesPermutation(axes, x.shape.length);
|
|
let $x = x;
|
|
if (permutedAxes != null) {
|
|
$x = transpose$1({ inputs: { x }, backend: backend2, attrs: { perm: permutedAxes } });
|
|
axes = getInnerMostAxes(axes.length, x.shape.length);
|
|
}
|
|
assertAxesAreInnerMostDims("min", axes, $x.shape.length);
|
|
const [outShape, reduceShape] = computeOutAndReduceShapes($x.shape, axes);
|
|
const reduceSize = sizeFromShape(reduceShape);
|
|
const vals = makeZerosTypedArray(sizeFromShape(outShape), $x.dtype);
|
|
const aVals = backend2.data.get($x.dataId).values;
|
|
for (let i = 0; i < vals.length; ++i) {
|
|
const offset = i * reduceSize;
|
|
let min2 = aVals[offset];
|
|
for (let j2 = 0; j2 < reduceSize; ++j2) {
|
|
const value = aVals[offset + j2];
|
|
if (Number.isNaN(value) || value < min2) {
|
|
min2 = value;
|
|
}
|
|
}
|
|
vals[i] = min2;
|
|
}
|
|
if (permutedAxes != null) {
|
|
backend2.disposeIntermediateTensorInfo($x);
|
|
}
|
|
const result = backend2.makeTensorInfo(outShape, $x.dtype, vals);
|
|
if (keepDims) {
|
|
const expandedShape = expandShapeToKeepDim(outShape, origAxes);
|
|
const reshapedResult = reshape({ inputs: { x: result }, backend: backend2, attrs: { shape: expandedShape } });
|
|
backend2.disposeIntermediateTensorInfo(result);
|
|
return reshapedResult;
|
|
}
|
|
return result;
|
|
}
|
|
const minConfig = {
|
|
kernelName: Min,
|
|
backendName: "cpu",
|
|
kernelFunc: min
|
|
};
|
|
function mirrorPad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { paddings, mode } = attrs;
|
|
assertNotComplex(x, "mirrorPad");
|
|
const outShape = paddings.map(
|
|
(p2, i) => p2[0] + x.shape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
const start = paddings.map((p2) => p2[0]);
|
|
const end = paddings.map((p2, i) => p2[0] + x.shape[i]);
|
|
const offset = mode === "reflect" ? 0 : 1;
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const xRank = x.shape.length;
|
|
const xStrides = computeStrides(x.shape);
|
|
const resultSize = sizeFromShape(outShape);
|
|
const resultRank = outShape.length;
|
|
const resultStrides = computeStrides(outShape);
|
|
const resVals = getTypedArrayFromDType(x.dtype, resultSize);
|
|
for (let i = 0; i < resultSize; i++) {
|
|
let coords2 = indexToLoc(i, resultRank, resultStrides);
|
|
for (let i2 = 0; i2 < resultRank; i2++) {
|
|
if (coords2[i2] < start[i2]) {
|
|
coords2[i2] = start[i2] * 2 - coords2[i2] - offset;
|
|
} else if (coords2[i2] >= end[i2]) {
|
|
coords2[i2] = (end[i2] - 1) * 2 - coords2[i2] + offset;
|
|
}
|
|
}
|
|
coords2 = coords2.map((c, i2) => c - start[i2]);
|
|
const inIndex = locToIndex(coords2, xRank, xStrides);
|
|
resVals[i] = xVals[inIndex];
|
|
}
|
|
const outId = backend2.write(resVals, outShape, x.dtype);
|
|
return { dataId: outId, shape: outShape, dtype: x.dtype };
|
|
}
|
|
const mirrorPadConfig = {
|
|
kernelName: MirrorPad,
|
|
backendName: "cpu",
|
|
kernelFunc: mirrorPad
|
|
};
|
|
const modImpl = createSimpleBinaryKernelImpl(((aValue, bValue) => {
|
|
const rem = aValue % bValue;
|
|
if (aValue < 0 && bValue < 0 || aValue >= 0 && bValue >= 0) {
|
|
return rem;
|
|
} else {
|
|
return (rem + bValue) % bValue;
|
|
}
|
|
}));
|
|
const mod = binaryKernelFunc$1(Mod, modImpl);
|
|
const modConfig = {
|
|
kernelName: Mod,
|
|
backendName: "cpu",
|
|
kernelFunc: mod
|
|
};
|
|
function softmax(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { logits } = inputs;
|
|
const { dim } = attrs;
|
|
const logitsRank = logits.shape.length;
|
|
let $dim = dim;
|
|
if ($dim === -1) {
|
|
$dim = logitsRank - 1;
|
|
}
|
|
if ($dim !== logitsRank - 1) {
|
|
throw Error(`Softmax along a non-last dimension is not yet supported. Logits was rank ${logitsRank} and dim was ${$dim}`);
|
|
}
|
|
const axes = parseAxisParam([$dim], logits.shape);
|
|
const maxLogit = max({
|
|
inputs: { x: logits },
|
|
backend: backend2,
|
|
attrs: { reductionIndices: axes, keepDims: false }
|
|
});
|
|
const expandedShape = expandShapeToKeepDim(maxLogit.shape, axes);
|
|
const maxLogitReshaped = reshape({ inputs: { x: maxLogit }, backend: backend2, attrs: { shape: expandedShape } });
|
|
const a = sub$1({ inputs: { a: logits, b: maxLogitReshaped }, backend: backend2 });
|
|
const b = exp$1({ inputs: { x: a }, backend: backend2 });
|
|
const sumExp = sum({ inputs: { x: b }, backend: backend2, attrs: { axis: axes, keepDims: false } });
|
|
const sumReshaped = reshape({ inputs: { x: sumExp }, backend: backend2, attrs: { shape: expandedShape } });
|
|
const result = div({ inputs: { a: b, b: sumReshaped }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(maxLogit);
|
|
backend2.disposeIntermediateTensorInfo(maxLogitReshaped);
|
|
backend2.disposeIntermediateTensorInfo(a);
|
|
backend2.disposeIntermediateTensorInfo(b);
|
|
backend2.disposeIntermediateTensorInfo(sumExp);
|
|
backend2.disposeIntermediateTensorInfo(sumReshaped);
|
|
return result;
|
|
}
|
|
const softmaxConfig = {
|
|
kernelName: Softmax,
|
|
backendName: "cpu",
|
|
kernelFunc: softmax
|
|
};
|
|
function multinomial(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { logits } = inputs;
|
|
const { numSamples, seed, normalized } = attrs;
|
|
assertNotComplex(logits, "multinomial");
|
|
const probabilities = normalized ? logits : softmax({ inputs: { logits }, backend: backend2, attrs: { dim: -1 } });
|
|
const batchSize = probabilities.shape[0];
|
|
const numEvents = probabilities.shape[1];
|
|
const probVals = backend2.data.get(probabilities.dataId).values;
|
|
const resShape = [batchSize, numSamples];
|
|
const resVals = makeZerosTypedArray(sizeFromShape(resShape), "int32");
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
const offset = b * numEvents;
|
|
const cdf = new Float32Array(numEvents - 1);
|
|
cdf[0] = probVals[offset];
|
|
for (let event = 1; event < cdf.length; ++event) {
|
|
cdf[event] = cdf[event - 1] + probVals[offset + event];
|
|
}
|
|
const random = seedrandomExports.alea(seed.toString());
|
|
const outOffset = b * numSamples;
|
|
for (let sampleId = 0; sampleId < numSamples; ++sampleId) {
|
|
const r = random();
|
|
resVals[outOffset + sampleId] = cdf.length;
|
|
for (let event = 0; event < cdf.length; event++) {
|
|
if (r < cdf[event]) {
|
|
resVals[outOffset + sampleId] = event;
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if (!normalized) {
|
|
backend2.disposeIntermediateTensorInfo(probabilities);
|
|
}
|
|
return backend2.makeTensorInfo(resShape, "int32", resVals);
|
|
}
|
|
const multinomialConfig = {
|
|
kernelName: Multinomial,
|
|
backendName: "cpu",
|
|
kernelFunc: multinomial
|
|
};
|
|
const nonMaxSuppressionV3Impl = nonMaxSuppressionV3Impl$2;
|
|
function nonMaxSuppressionV3(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { boxes, scores } = inputs;
|
|
const { maxOutputSize, iouThreshold, scoreThreshold } = attrs;
|
|
assertNotComplex(boxes, "NonMaxSuppression");
|
|
const boxesVals = backend2.data.get(boxes.dataId).values;
|
|
const scoresVals = backend2.data.get(scores.dataId).values;
|
|
const { selectedIndices } = nonMaxSuppressionV3Impl(boxesVals, scoresVals, maxOutputSize, iouThreshold, scoreThreshold);
|
|
return backend2.makeTensorInfo([selectedIndices.length], "int32", new Int32Array(selectedIndices));
|
|
}
|
|
const nonMaxSuppressionV3Config = {
|
|
kernelName: NonMaxSuppressionV3,
|
|
backendName: "cpu",
|
|
kernelFunc: nonMaxSuppressionV3
|
|
};
|
|
const nonMaxSuppressionV4Impl = nonMaxSuppressionV4Impl$2;
|
|
function nonMaxSuppressionV4(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { boxes, scores } = inputs;
|
|
const { maxOutputSize, iouThreshold, scoreThreshold, padToMaxOutputSize } = attrs;
|
|
assertNotComplex(boxes, "NonMaxSuppressionPadded");
|
|
const boxesVals = backend2.data.get(boxes.dataId).values;
|
|
const scoresVals = backend2.data.get(scores.dataId).values;
|
|
const { selectedIndices, validOutputs } = nonMaxSuppressionV4Impl(boxesVals, scoresVals, maxOutputSize, iouThreshold, scoreThreshold, padToMaxOutputSize);
|
|
return [
|
|
backend2.makeTensorInfo([selectedIndices.length], "int32", new Int32Array(selectedIndices)),
|
|
backend2.makeTensorInfo([], "int32", new Int32Array([validOutputs]))
|
|
];
|
|
}
|
|
const nonMaxSuppressionV4Config = {
|
|
kernelName: NonMaxSuppressionV4,
|
|
backendName: "cpu",
|
|
kernelFunc: nonMaxSuppressionV4
|
|
};
|
|
const nonMaxSuppressionV5Impl = nonMaxSuppressionV5Impl$2;
|
|
function nonMaxSuppressionV5(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { boxes, scores } = inputs;
|
|
const { maxOutputSize, iouThreshold, scoreThreshold, softNmsSigma } = attrs;
|
|
assertNotComplex(boxes, "NonMaxSuppressionWithScore");
|
|
const boxesVals = backend2.data.get(boxes.dataId).values;
|
|
const scoresVals = backend2.data.get(scores.dataId).values;
|
|
const maxOutputSizeVal = maxOutputSize;
|
|
const iouThresholdVal = iouThreshold;
|
|
const scoreThresholdVal = scoreThreshold;
|
|
const softNmsSigmaVal = softNmsSigma;
|
|
const { selectedIndices, selectedScores } = nonMaxSuppressionV5Impl(boxesVals, scoresVals, maxOutputSizeVal, iouThresholdVal, scoreThresholdVal, softNmsSigmaVal);
|
|
return [
|
|
backend2.makeTensorInfo([selectedIndices.length], "int32", new Int32Array(selectedIndices)),
|
|
backend2.makeTensorInfo([selectedScores.length], "float32", new Float32Array(selectedScores))
|
|
];
|
|
}
|
|
const nonMaxSuppressionV5Config = {
|
|
kernelName: NonMaxSuppressionV5,
|
|
backendName: "cpu",
|
|
kernelFunc: nonMaxSuppressionV5
|
|
};
|
|
function oneHot(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { indices } = inputs;
|
|
const { dtype, depth, onValue, offValue } = attrs;
|
|
assertNotComplex(indices, "oneHot");
|
|
const indicesSize = sizeFromShape(indices.shape);
|
|
const res = new Float32Array(indicesSize * depth);
|
|
res.fill(offValue);
|
|
const indicesVal = backend2.data.get(indices.dataId).values;
|
|
for (let event = 0; event < indicesSize; ++event) {
|
|
if (indicesVal[event] >= 0 && indicesVal[event] < depth) {
|
|
res[event * depth + indicesVal[event]] = onValue;
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo([...indices.shape, depth], dtype, res);
|
|
}
|
|
const oneHotConfig = {
|
|
kernelName: OneHot,
|
|
backendName: "cpu",
|
|
kernelFunc: oneHot
|
|
};
|
|
function zerosLike(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
if (x.dtype === "string") {
|
|
throw new Error("zerosLike is not supported for string tensors");
|
|
} else if (x.dtype === "complex64") {
|
|
const realPart = real$1({ inputs: { input: x }, backend: backend2 });
|
|
const r = zerosLike({ inputs: { x: realPart }, backend: backend2 });
|
|
const imagPart = imag({ inputs: { input: x }, backend: backend2 });
|
|
const i = zerosLike({ inputs: { x: imagPart }, backend: backend2 });
|
|
const result = complex$1({ inputs: { real: r, imag: i }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
backend2.disposeIntermediateTensorInfo(r);
|
|
backend2.disposeIntermediateTensorInfo(imagPart);
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
return result;
|
|
} else {
|
|
return fill({ backend: backend2, attrs: { shape: x.shape, value: 0, dtype: x.dtype } });
|
|
}
|
|
}
|
|
const zerosLikeConfig = {
|
|
kernelName: ZerosLike,
|
|
backendName: "cpu",
|
|
kernelFunc: zerosLike
|
|
};
|
|
function onesLike(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { x } = inputs;
|
|
if (x.dtype === "string") {
|
|
throw new Error("onesLike is not supported for string tensors");
|
|
} else if (x.dtype === "complex64") {
|
|
const realPart = real$1({ inputs: { input: x }, backend: backend2 });
|
|
const r = onesLike({ inputs: { x: realPart }, backend: backend2 });
|
|
const imagPart = imag({ inputs: { input: x }, backend: backend2 });
|
|
const i = zerosLike({ inputs: { x: imagPart }, backend: backend2 });
|
|
const result = complex$1({ inputs: { real: r, imag: i }, backend: backend2 });
|
|
backend2.disposeIntermediateTensorInfo(realPart);
|
|
backend2.disposeIntermediateTensorInfo(r);
|
|
backend2.disposeIntermediateTensorInfo(imagPart);
|
|
backend2.disposeIntermediateTensorInfo(i);
|
|
return result;
|
|
} else {
|
|
return fill({ backend: backend2, attrs: { shape: x.shape, value: 1, dtype: x.dtype } });
|
|
}
|
|
}
|
|
const onesLikeConfig = {
|
|
kernelName: OnesLike,
|
|
backendName: "cpu",
|
|
kernelFunc: onesLike
|
|
};
|
|
function pack(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { axis } = attrs;
|
|
if (inputs.length === 1) {
|
|
return expandDims({ inputs: { input: inputs[0] }, backend: backend2, attrs: { dim: axis } });
|
|
}
|
|
const shape = inputs[0].shape;
|
|
const dtype = inputs[0].dtype;
|
|
inputs.forEach((t2) => {
|
|
assertShapesMatch(shape, t2.shape, "All tensors passed to stack must have matching shapes");
|
|
assert(dtype === t2.dtype, () => "All tensors passed to stack must have matching dtypes");
|
|
});
|
|
const intermediateTensorInfos = [];
|
|
const expandedTensors = inputs.map((t2) => {
|
|
const expandedT = expandDims({ inputs: { input: t2 }, backend: backend2, attrs: { dim: axis } });
|
|
intermediateTensorInfos.push(expandedT);
|
|
return expandedT;
|
|
});
|
|
const result = concat({ inputs: expandedTensors, backend: backend2, attrs: { axis } });
|
|
intermediateTensorInfos.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
}
|
|
const packConfig = {
|
|
kernelName: Pack,
|
|
backendName: "cpu",
|
|
kernelFunc: pack
|
|
};
|
|
function padV2(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { paddings, constantValue } = attrs;
|
|
assertNotComplex(x, "pad");
|
|
const outShape = paddings.map(
|
|
(p2, i) => p2[0] + x.shape[i] + p2[1]
|
|
/* afterPad */
|
|
);
|
|
const start = paddings.map((p2) => p2[0]);
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const xSize = sizeFromShape(x.shape);
|
|
const xRank = x.shape.length;
|
|
const xStrides = computeStrides(x.shape);
|
|
const resultSize = sizeFromShape(outShape);
|
|
const resultRank = outShape.length;
|
|
const resultStrides = computeStrides(outShape);
|
|
const resVals = getTypedArrayFromDType(x.dtype, resultSize);
|
|
if (constantValue !== 0) {
|
|
resVals.fill(constantValue);
|
|
}
|
|
for (let i = 0; i < xSize; i++) {
|
|
const coords2 = indexToLoc(i, xRank, xStrides);
|
|
const outCoords = coords2.map((c, i2) => c + start[i2]);
|
|
const outIndex = locToIndex(outCoords, resultRank, resultStrides);
|
|
resVals[outIndex] = xVals[i];
|
|
}
|
|
const outId = backend2.write(resVals, outShape, x.dtype);
|
|
return { dataId: outId, shape: outShape, dtype: x.dtype };
|
|
}
|
|
const padV2Config = {
|
|
kernelName: PadV2,
|
|
backendName: "cpu",
|
|
kernelFunc: padV2
|
|
};
|
|
const powImpl = createSimpleBinaryKernelImpl((a, b) => Math.pow(a, b));
|
|
const pow = binaryKernelFunc$1(Pow, powImpl);
|
|
const powConfig = {
|
|
kernelName: Pow,
|
|
backendName: "cpu",
|
|
kernelFunc: pow
|
|
};
|
|
function raggedGather(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { paramsNestedSplits, paramsDenseValues, indices } = inputs;
|
|
const { outputRaggedRank } = attrs;
|
|
const $paramsNestedSplits = paramsNestedSplits.map((t2) => backend2.data.get(t2.dataId).values);
|
|
const $paramsNestedSplitsShapes = paramsNestedSplits.map((t2) => t2.shape);
|
|
const $paramsDenseValues = backend2.data.get(paramsDenseValues.dataId).values;
|
|
const $indices = backend2.data.get(indices.dataId).values;
|
|
const [outputNestedSplits, outputDenseValues, outputDenseValuesShape] = raggedGatherImpl($paramsNestedSplits, $paramsNestedSplitsShapes, $paramsDenseValues, paramsDenseValues.shape, paramsDenseValues.dtype, $indices, indices.shape);
|
|
const outputNestedSplitsTensors = outputNestedSplits.map((splits) => backend2.makeTensorInfo([splits.length], "int32", splits));
|
|
const outputDenseValuesTensor = backend2.makeTensorInfo(outputDenseValuesShape, paramsDenseValues.dtype, outputDenseValues);
|
|
return outputNestedSplitsTensors.concat([outputDenseValuesTensor]);
|
|
}
|
|
const raggedGatherConfig = {
|
|
kernelName: RaggedGather,
|
|
backendName: "cpu",
|
|
kernelFunc: raggedGather
|
|
};
|
|
function raggedRange(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { starts, limits, deltas } = inputs;
|
|
const $starts = backend2.data.get(starts.dataId).values;
|
|
const $limits = backend2.data.get(limits.dataId).values;
|
|
const $deltas = backend2.data.get(deltas.dataId).values;
|
|
const [rtNestedSplitsData, rtDenseValuesData] = raggedRangeImpl($starts, starts.shape, starts.dtype, $limits, limits.shape, $deltas, deltas.shape);
|
|
const rtNestedSplits = backend2.makeTensorInfo([rtNestedSplitsData.length], "int32", rtNestedSplitsData);
|
|
const rtDenseValues = backend2.makeTensorInfo([rtDenseValuesData.length], starts.dtype, rtDenseValuesData);
|
|
return [rtNestedSplits, rtDenseValues];
|
|
}
|
|
const raggedRangeConfig = {
|
|
kernelName: RaggedRange,
|
|
backendName: "cpu",
|
|
kernelFunc: raggedRange
|
|
};
|
|
function raggedTensorToTensor(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { shape, values, defaultValue, rowPartitionTensors } = inputs;
|
|
const { rowPartitionTypes } = attrs;
|
|
const $shape = backend2.data.get(shape.dataId).values;
|
|
const $values = backend2.data.get(values.dataId).values;
|
|
const $defaultValue = backend2.data.get(defaultValue.dataId).values;
|
|
const $rowPartitionValues = rowPartitionTensors.map((t2) => backend2.data.get(t2.dataId).values);
|
|
const rowPartitionValuesShapes = rowPartitionTensors.map((t2) => t2.shape);
|
|
const [outputShape, output] = raggedTensorToTensorImpl($shape, shape.shape, $values, values.shape, values.dtype, $defaultValue, defaultValue.shape, $rowPartitionValues, rowPartitionValuesShapes, rowPartitionTypes);
|
|
return backend2.makeTensorInfo(outputShape, values.dtype, output);
|
|
}
|
|
const raggedTensorToTensorConfig = {
|
|
kernelName: RaggedTensorToTensor,
|
|
backendName: "cpu",
|
|
kernelFunc: raggedTensorToTensor
|
|
};
|
|
function range(args) {
|
|
const { backend: backend2, attrs } = args;
|
|
const { start, stop, dtype, step: step2 } = attrs;
|
|
const values = rangeImpl(start, stop, step2, dtype);
|
|
return backend2.makeTensorInfo([values.length], dtype, values);
|
|
}
|
|
const rangeConfig = {
|
|
kernelName: Range,
|
|
backendName: "cpu",
|
|
kernelFunc: range
|
|
};
|
|
const reciprocal = unaryKernelFunc$1(Reciprocal, (xi) => 1 / xi);
|
|
const reciprocalConfig = {
|
|
kernelName: Reciprocal,
|
|
backendName: "cpu",
|
|
kernelFunc: reciprocal
|
|
};
|
|
function resizeBilinear(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images } = inputs;
|
|
const { alignCorners, halfPixelCenters, size } = attrs;
|
|
assertNotComplex(images, "resizeBilinear");
|
|
const imagesStrides = computeStrides(images.shape);
|
|
const [newHeight, newWidth] = size;
|
|
const [batch, oldHeight, oldWidth, numChannels] = images.shape;
|
|
const xValues = backend2.data.get(images.dataId).values;
|
|
const result = new Float32Array(sizeFromShape([batch, newHeight, newWidth, numChannels]));
|
|
const effectiveInputSize = [
|
|
alignCorners && newHeight > 1 ? oldHeight - 1 : oldHeight,
|
|
alignCorners && newWidth > 1 ? oldWidth - 1 : oldWidth
|
|
];
|
|
const effectiveOutputSize = [
|
|
alignCorners && newHeight > 1 ? newHeight - 1 : newHeight,
|
|
alignCorners && newWidth > 1 ? newWidth - 1 : newWidth
|
|
];
|
|
let outputIdx = 0;
|
|
const effectiveRowSizeRatio = effectiveInputSize[0] / effectiveOutputSize[0];
|
|
const effectiveColSizeRatio = effectiveInputSize[1] / effectiveOutputSize[1];
|
|
for (let b = 0; b < batch; b++) {
|
|
for (let r = 0; r < newHeight; r++) {
|
|
let sourceFracRow;
|
|
if (halfPixelCenters) {
|
|
sourceFracRow = effectiveRowSizeRatio * (r + 0.5) - 0.5;
|
|
} else {
|
|
sourceFracRow = effectiveRowSizeRatio * r;
|
|
}
|
|
const sourceRowFloor = Math.max(0, Math.floor(sourceFracRow));
|
|
const rowFrac = sourceFracRow - sourceRowFloor;
|
|
const sourceRowCeil = Math.min(oldHeight - 1, Math.ceil(sourceFracRow));
|
|
const topRowOffset = b * imagesStrides[0] + sourceRowFloor * imagesStrides[1];
|
|
const botRowOffset = b * imagesStrides[0] + sourceRowCeil * imagesStrides[1];
|
|
for (let c = 0; c < newWidth; c++) {
|
|
let sourceFracCol;
|
|
if (halfPixelCenters) {
|
|
sourceFracCol = effectiveColSizeRatio * (c + 0.5) - 0.5;
|
|
} else {
|
|
sourceFracCol = effectiveColSizeRatio * c;
|
|
}
|
|
const sourceColFloor = Math.max(0, Math.floor(sourceFracCol));
|
|
const colFrac = sourceFracCol - sourceColFloor;
|
|
const sourceColCeil = Math.min(oldWidth - 1, Math.ceil(sourceFracCol));
|
|
const topLeftOffest = topRowOffset + sourceColFloor * imagesStrides[2];
|
|
const botLeftOffset = botRowOffset + sourceColFloor * imagesStrides[2];
|
|
const topRightOffset = topRowOffset + sourceColCeil * imagesStrides[2];
|
|
const botRightOffest = botRowOffset + sourceColCeil * imagesStrides[2];
|
|
for (let d = 0; d < numChannels; d++) {
|
|
const topLeft = xValues[topLeftOffest + d];
|
|
const bottomLeft = xValues[botLeftOffset + d];
|
|
const topRight = xValues[topRightOffset + d];
|
|
const bottomRight = xValues[botRightOffest + d];
|
|
const top = topLeft + (topRight - topLeft) * colFrac;
|
|
const bottom = bottomLeft + (bottomRight - bottomLeft) * colFrac;
|
|
const newValue = top + (bottom - top) * rowFrac;
|
|
result[outputIdx++] = newValue;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo([batch, newHeight, newWidth, numChannels], "float32", result);
|
|
}
|
|
const resizeBilinearConfig = {
|
|
kernelName: ResizeBilinear,
|
|
backendName: "cpu",
|
|
kernelFunc: resizeBilinear
|
|
};
|
|
function resizeBilinearGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images, dy } = inputs;
|
|
const { alignCorners } = attrs;
|
|
assertNotComplex([dy, images], "resizeBilinearGrad");
|
|
const imagesStrides = computeStrides(images.shape);
|
|
const [batch, xHeight, xWidth, depth] = images.shape;
|
|
const [, yHeight, yWidth] = dy.shape;
|
|
const output = new Float32Array(batch * xHeight * xWidth * depth);
|
|
const effectiveXSize = [
|
|
alignCorners && yHeight > 1 ? xHeight - 1 : xHeight,
|
|
alignCorners && yWidth > 1 ? xWidth - 1 : xWidth
|
|
];
|
|
const effectiveYSize = [
|
|
alignCorners && yHeight > 1 ? yHeight - 1 : yHeight,
|
|
alignCorners && yWidth > 1 ? yWidth - 1 : yWidth
|
|
];
|
|
const heightScale = effectiveXSize[0] / effectiveYSize[0];
|
|
const widthScale = effectiveXSize[1] / effectiveYSize[1];
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
let offset = 0;
|
|
for (let b = 0; b < batch; b++) {
|
|
const bOffset = b * imagesStrides[0];
|
|
for (let r = 0; r < yHeight; r++) {
|
|
const dxR = r * heightScale;
|
|
const topDxRIndex = Math.floor(dxR);
|
|
const bottomDxRIndex = Math.min(Math.ceil(dxR), xHeight - 1);
|
|
const topDxROffset = bOffset + topDxRIndex * imagesStrides[1];
|
|
const bottomDxROffset = bOffset + bottomDxRIndex * imagesStrides[1];
|
|
const dxRLerp = dxR - topDxRIndex;
|
|
const inverseDxRLerp = 1 - dxRLerp;
|
|
for (let c = 0; c < yWidth; c++) {
|
|
const dxC = c * widthScale;
|
|
const leftDxCIndex = Math.floor(dxC);
|
|
const rightDxCIndex = Math.min(Math.ceil(dxC), xWidth - 1);
|
|
const dxCLerp = dxC - leftDxCIndex;
|
|
const inverseDxCLerp = 1 - dxCLerp;
|
|
const topLeftRCOffset = topDxROffset + leftDxCIndex * imagesStrides[2];
|
|
const topRightRCOffset = topDxROffset + rightDxCIndex * imagesStrides[2];
|
|
const bottomLeftRCOffset = bottomDxROffset + leftDxCIndex * imagesStrides[2];
|
|
const bottomRightRCOffset = bottomDxROffset + rightDxCIndex * imagesStrides[2];
|
|
const inverseDxRLerpTimesInverseDxCLerp = inverseDxRLerp * inverseDxCLerp;
|
|
const inverseDxRLerpTimesDxCLerp = inverseDxRLerp * dxCLerp;
|
|
const dxRLerpTimesInverseDxCLerp = dxRLerp * inverseDxCLerp;
|
|
const dxRLerpTimesDxCLerp = dxRLerp * dxCLerp;
|
|
for (let d = 0; d < depth; d++) {
|
|
const dyVal = dyValues[offset++];
|
|
output[topLeftRCOffset + d] += dyVal * inverseDxRLerpTimesInverseDxCLerp;
|
|
output[topRightRCOffset + d] += dyVal * inverseDxRLerpTimesDxCLerp;
|
|
output[bottomLeftRCOffset + d] += dyVal * dxRLerpTimesInverseDxCLerp;
|
|
output[bottomRightRCOffset + d] += dyVal * dxRLerpTimesDxCLerp;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo([batch, xWidth, xHeight, depth], "float32", output);
|
|
}
|
|
const resizeBilinearGradConfig = {
|
|
kernelName: ResizeBilinearGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: resizeBilinearGrad
|
|
};
|
|
function resizeNearestNeighbor(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images } = inputs;
|
|
const { alignCorners, halfPixelCenters, size } = attrs;
|
|
assertNotComplex(images, "resizeNearestNeighbor");
|
|
const imagesStrides = computeStrides(images.shape);
|
|
const [newHeight, newWidth] = size;
|
|
const [batch, oldHeight, oldWidth, numChannels] = images.shape;
|
|
const xValues = backend2.data.get(images.dataId).values;
|
|
const output = new Float32Array(batch * newHeight * newWidth * numChannels);
|
|
const effectiveInputSize = [
|
|
alignCorners && newHeight > 1 ? oldHeight - 1 : oldHeight,
|
|
alignCorners && newWidth > 1 ? oldWidth - 1 : oldWidth
|
|
];
|
|
const effectiveOutputSize = [
|
|
alignCorners && newHeight > 1 ? newHeight - 1 : newHeight,
|
|
alignCorners && newWidth > 1 ? newWidth - 1 : newWidth
|
|
];
|
|
const effectiveRowSizeRatio = effectiveInputSize[0] / effectiveOutputSize[0];
|
|
const effectiveColSizeRatio = effectiveInputSize[1] / effectiveOutputSize[1];
|
|
let outputOffset = 0;
|
|
for (let b = 0; b < batch; b++) {
|
|
const batchOffset = b * imagesStrides[0];
|
|
for (let r = 0; r < newHeight; r++) {
|
|
const sourceFracRow = halfPixelCenters ? effectiveRowSizeRatio * (r + 0.5) : effectiveRowSizeRatio * r;
|
|
let sourceNearestRow = Math.min(oldHeight - 1, alignCorners ? Math.round(sourceFracRow) : Math.floor(sourceFracRow));
|
|
if (halfPixelCenters) {
|
|
sourceNearestRow = Math.max(0, sourceNearestRow);
|
|
}
|
|
const rowOffset = batchOffset + sourceNearestRow * imagesStrides[1];
|
|
for (let c = 0; c < newWidth; c++) {
|
|
const sourceFracCol = halfPixelCenters ? effectiveColSizeRatio * (c + 0.5) : effectiveColSizeRatio * c;
|
|
let sourceNearestCol = Math.min(oldWidth - 1, alignCorners ? Math.round(sourceFracCol) : Math.floor(sourceFracCol));
|
|
if (halfPixelCenters) {
|
|
sourceNearestCol = Math.max(0, sourceNearestCol);
|
|
}
|
|
const colOffset = rowOffset + sourceNearestCol * imagesStrides[2];
|
|
for (let d = 0; d < numChannels; d++) {
|
|
const newVal = xValues[colOffset + d];
|
|
output[outputOffset++] = newVal;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo([batch, newHeight, newWidth, numChannels], images.dtype, output);
|
|
}
|
|
const resizeNearestNeighborConfig = {
|
|
kernelName: ResizeNearestNeighbor,
|
|
backendName: "cpu",
|
|
kernelFunc: resizeNearestNeighbor
|
|
};
|
|
function resizeNearestNeighborGrad(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { images, dy } = inputs;
|
|
const { alignCorners } = attrs;
|
|
assertNotComplex([dy, images], "resizeNearestNeighborGrad");
|
|
const imagesStrides = computeStrides(images.shape);
|
|
const dyStrides = computeStrides(dy.shape);
|
|
const [batch, xHeight, xWidth, depth] = images.shape;
|
|
const [, yHeight, yWidth] = dy.shape;
|
|
const output = new Float32Array(batch * xHeight * xWidth * depth);
|
|
const dyValues = backend2.data.get(dy.dataId).values;
|
|
const effectiveXSize = [
|
|
alignCorners && yHeight > 1 ? xHeight - 1 : xHeight,
|
|
alignCorners && yWidth > 1 ? xWidth - 1 : xWidth
|
|
];
|
|
const effectiveYSize = [
|
|
alignCorners && yHeight > 1 ? yHeight - 1 : yHeight,
|
|
alignCorners && yWidth > 1 ? yWidth - 1 : yWidth
|
|
];
|
|
const heightScale = effectiveXSize[0] / effectiveYSize[0];
|
|
const widthScale = effectiveXSize[1] / effectiveYSize[1];
|
|
const invHeightScale = 1 / heightScale;
|
|
const invWidthScale = 1 / widthScale;
|
|
const winHeight = Math.ceil(invHeightScale) * 2 + 2;
|
|
const winWidth = Math.ceil(invWidthScale) * 2 + 2;
|
|
for (let b = 0; b < batch; b++) {
|
|
const batchOffset = b * imagesStrides[0];
|
|
for (let r = 0; r < xHeight; r++) {
|
|
const rowOffset = batchOffset + r * imagesStrides[1];
|
|
const startRLerp = Math.floor(r * invHeightScale);
|
|
const startDyR = Math.floor(startRLerp - winHeight / 2);
|
|
for (let c = 0; c < xWidth; c++) {
|
|
const colOffset = rowOffset + c * imagesStrides[2];
|
|
const startCLerp = Math.floor(c * invWidthScale);
|
|
const startDyC = Math.floor(startCLerp - winWidth / 2);
|
|
for (let d = 0; d < depth; d++) {
|
|
let accum = 0;
|
|
for (let dyRIndex = 0; dyRIndex < winHeight; dyRIndex++) {
|
|
const dyR = dyRIndex + startDyR;
|
|
if (dyR < 0 || dyR >= yHeight) {
|
|
continue;
|
|
}
|
|
const dyROffset = batchOffset + dyR * dyStrides[1];
|
|
const sourceFracRow = dyR * heightScale;
|
|
const sourceNearestRow = Math.min(xHeight - 1, alignCorners ? Math.round(sourceFracRow) : Math.floor(sourceFracRow));
|
|
if (r !== sourceNearestRow) {
|
|
continue;
|
|
}
|
|
for (let dyCIndex = 0; dyCIndex < winWidth; dyCIndex++) {
|
|
const dyC = dyCIndex + startDyC;
|
|
if (dyC < 0 || dyC >= yWidth) {
|
|
continue;
|
|
}
|
|
const dyCOffset = dyROffset + dyC * dyStrides[2];
|
|
const sourceFracCol = dyC * widthScale;
|
|
const sourceNearestCol = Math.min(xWidth - 1, alignCorners ? Math.round(sourceFracCol) : Math.floor(sourceFracCol));
|
|
if (c === sourceNearestCol) {
|
|
accum += dyValues[dyCOffset + d];
|
|
}
|
|
}
|
|
}
|
|
output[colOffset + d] = accum;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(images.shape, images.dtype, output);
|
|
}
|
|
const resizeNearestNeighborGradConfig = {
|
|
kernelName: ResizeNearestNeighborGrad,
|
|
backendName: "cpu",
|
|
kernelFunc: resizeNearestNeighborGrad
|
|
};
|
|
function reverse(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { dims } = attrs;
|
|
assertNotComplex(x, "reverse");
|
|
const xRank = x.shape.length;
|
|
const $dims = parseAxisParam(dims, x.shape);
|
|
if (xRank === 0) {
|
|
return identity$1({ inputs: { x }, backend: backend2 });
|
|
}
|
|
const outBuf = new TensorBuffer(x.shape, x.dtype);
|
|
const xBuf = backend2.bufferSync(x);
|
|
for (let i = 0; i < outBuf.size; i++) {
|
|
const outLoc = outBuf.indexToLoc(i);
|
|
const inLoc = outLoc.slice();
|
|
$dims.forEach((d) => inLoc[d] = x.shape[d] - 1 - inLoc[d]);
|
|
outBuf.set(xBuf.get(...inLoc), ...outLoc);
|
|
}
|
|
return backend2.makeTensorInfo(outBuf.shape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const reverseConfig = {
|
|
kernelName: Reverse,
|
|
backendName: "cpu",
|
|
kernelFunc: reverse
|
|
};
|
|
const rotateWithOffsetConfig = {
|
|
kernelName: RotateWithOffset,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, attrs, backend: backend2 }) => {
|
|
const { image: image2 } = inputs;
|
|
const { radians, fillValue, center } = attrs;
|
|
const cpuBackend = backend2;
|
|
const output = getTypedArrayFromDType(image2.dtype, sizeFromShape(image2.shape));
|
|
const [batch, imageHeight, imageWidth, numChannels] = image2.shape;
|
|
const [centerX, centerY] = getImageCenter(center, imageHeight, imageWidth);
|
|
const fullOpacityValue = 255;
|
|
const sinFactor = Math.sin(radians);
|
|
const cosFactor = Math.cos(radians);
|
|
const imageVals = cpuBackend.data.get(image2.dataId).values;
|
|
for (let batchIdx = 0; batchIdx < batch; batchIdx++) {
|
|
const batchOffset = batchIdx * imageWidth * imageHeight * numChannels;
|
|
for (let row = 0; row < imageHeight; row++) {
|
|
const rowOffset = row * (imageWidth * numChannels);
|
|
for (let col = 0; col < imageWidth; col++) {
|
|
const colOffset = col * numChannels;
|
|
for (let channel = 0; channel < numChannels; channel++) {
|
|
const coords2 = [batch, row, col, channel];
|
|
const x = coords2[2];
|
|
const y = coords2[1];
|
|
let coordX = (x - centerX) * cosFactor - (y - centerY) * sinFactor;
|
|
let coordY = (x - centerX) * sinFactor + (y - centerY) * cosFactor;
|
|
coordX = Math.round(coordX + centerX);
|
|
coordY = Math.round(coordY + centerY);
|
|
let outputValue = fillValue;
|
|
if (typeof fillValue !== "number") {
|
|
if (channel === 3) {
|
|
outputValue = fullOpacityValue;
|
|
} else {
|
|
outputValue = fillValue[channel];
|
|
}
|
|
}
|
|
if (coordX >= 0 && coordX < imageWidth && coordY >= 0 && coordY < imageHeight) {
|
|
const rotatedRowOffset = coordY * (imageWidth * numChannels);
|
|
const rotatedColOffset = coordX * numChannels;
|
|
const imageIdx = batchOffset + rotatedRowOffset + rotatedColOffset + channel;
|
|
outputValue = imageVals[imageIdx];
|
|
}
|
|
const outIdx = batchOffset + rowOffset + colOffset + channel;
|
|
output[outIdx] = outputValue;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
const dataId = cpuBackend.write(output, image2.shape, image2.dtype);
|
|
return { dataId, shape: image2.shape, dtype: image2.dtype };
|
|
}
|
|
};
|
|
const round = unaryKernelFunc$1(Round, (xi) => {
|
|
const base = Math.floor(xi);
|
|
if (xi - base < 0.5) {
|
|
return Math.floor(xi);
|
|
} else if (xi - base > 0.5) {
|
|
return Math.ceil(xi);
|
|
} else {
|
|
if (base % 2 === 0) {
|
|
return base;
|
|
} else {
|
|
return base + 1;
|
|
}
|
|
}
|
|
});
|
|
const roundConfig = {
|
|
kernelName: Round,
|
|
backendName: "cpu",
|
|
kernelFunc: round
|
|
};
|
|
function scatterNd(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { indices, updates } = inputs;
|
|
const { shape } = attrs;
|
|
const { sliceRank, numUpdates, sliceSize, strides, outputSize } = calculateShapes(updates, indices, shape);
|
|
const sumDupeIndices = true;
|
|
const indicesBuf = backend2.bufferSync(indices);
|
|
const updatesBuf = backend2.bufferSync(updates);
|
|
const outBuf = scatterImpl(indicesBuf, updatesBuf, shape, outputSize, sliceSize, numUpdates, sliceRank, strides, 0, sumDupeIndices);
|
|
return backend2.makeTensorInfo(shape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const scatterNdConfig = {
|
|
kernelName: ScatterNd,
|
|
backendName: "cpu",
|
|
kernelFunc: scatterNd
|
|
};
|
|
function lowerBound(array, value) {
|
|
let left = 0;
|
|
let right = array.length;
|
|
let mid = 0;
|
|
while (left < right) {
|
|
mid = Math.floor((left + right) / 2);
|
|
if (array[mid] < value) {
|
|
left = mid + 1;
|
|
} else {
|
|
right = mid;
|
|
}
|
|
}
|
|
return right;
|
|
}
|
|
function upperBound(array, value) {
|
|
let left = 0;
|
|
let right = array.length;
|
|
let mid = 0;
|
|
while (left < right) {
|
|
mid = Math.floor((left + right) / 2);
|
|
if (array[mid] <= value) {
|
|
left = mid + 1;
|
|
} else {
|
|
right = mid;
|
|
}
|
|
}
|
|
return right;
|
|
}
|
|
function searchSortedImpl(sortedInputs, values, batchSize, numInputs, numValues, side) {
|
|
const output = getArrayFromDType("int32", batchSize * numValues);
|
|
for (let b = 0; b < batchSize; ++b) {
|
|
const sortedInputsSlice = sortedInputs.slice(b * numInputs, (b + 1) * numInputs);
|
|
const outputOffset = b * numValues;
|
|
for (let i = 0; i < numValues; ++i) {
|
|
output[outputOffset + i] = side === "left" ? lowerBound(sortedInputsSlice, values[i + outputOffset]) : upperBound(sortedInputsSlice, values[i + outputOffset]);
|
|
}
|
|
}
|
|
return output;
|
|
}
|
|
function searchSorted(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { sortedSequence, values } = inputs;
|
|
const { side } = attrs;
|
|
const $sortedSequence = backend2.data.get(sortedSequence.dataId).values;
|
|
const $values = backend2.data.get(values.dataId).values;
|
|
const output = searchSortedImpl($sortedSequence, $values, sortedSequence.shape[0], sortedSequence.shape[1], values.shape[1], side);
|
|
return backend2.makeTensorInfo(values.shape, "int32", output);
|
|
}
|
|
const searchSortedConfig = {
|
|
kernelName: SearchSorted,
|
|
backendName: "cpu",
|
|
kernelFunc: searchSorted
|
|
};
|
|
function select(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { condition, t: t2, e } = inputs;
|
|
assertNotComplex([condition, t2, e], "select");
|
|
const conditionRank = condition.shape.length;
|
|
const values = backend2.data.get(condition.dataId).values;
|
|
const tValues = backend2.data.get(t2.dataId).values;
|
|
const eValues = backend2.data.get(e.dataId).values;
|
|
const resultDtype = upcastType(t2.dtype, e.dtype);
|
|
const newValues = makeZerosTypedArray(sizeFromShape(t2.shape), resultDtype);
|
|
let index = 0;
|
|
const offset = conditionRank === 0 || conditionRank > 1 || t2.shape.length === 1 ? 1 : sizeFromShape(t2.shape.slice(1));
|
|
for (let i = 0; i < values.length; i++) {
|
|
for (let j2 = 0; j2 < offset; j2++) {
|
|
if (values[i] === 1) {
|
|
newValues[index++] = tValues[i];
|
|
} else {
|
|
newValues[index++] = eValues[i];
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(t2.shape, resultDtype, newValues);
|
|
}
|
|
const selectConfig = {
|
|
kernelName: Select,
|
|
backendName: "cpu",
|
|
kernelFunc: select
|
|
};
|
|
const scaleAlpha = SELU_SCALEALPHA;
|
|
const scale = SELU_SCALE;
|
|
const selu = unaryKernelFunc$1(Selu, (xi) => {
|
|
if (xi >= 0) {
|
|
return scale * xi;
|
|
} else {
|
|
return scaleAlpha * (Math.exp(xi) - 1);
|
|
}
|
|
});
|
|
const seluConfig = {
|
|
kernelName: Selu,
|
|
backendName: "cpu",
|
|
kernelFunc: selu
|
|
};
|
|
const sign = unaryKernelFunc$1(Sign, (xi) => {
|
|
if (xi < 0) {
|
|
return -1;
|
|
} else if (xi > 0) {
|
|
return 1;
|
|
} else {
|
|
return 0;
|
|
}
|
|
});
|
|
const signConfig = {
|
|
kernelName: Sign,
|
|
backendName: "cpu",
|
|
kernelFunc: sign
|
|
};
|
|
const sin = unaryKernelFunc$1(Sin, (xi) => Math.sin(xi));
|
|
const sinConfig = {
|
|
kernelName: Sin,
|
|
backendName: "cpu",
|
|
kernelFunc: sin
|
|
};
|
|
const sinh = unaryKernelFunc$1(Sinh, (xi) => Math.sinh(xi));
|
|
const sinhConfig = {
|
|
kernelName: Sinh,
|
|
backendName: "cpu",
|
|
kernelFunc: sinh
|
|
};
|
|
const epsilon = 11920928955078125e-23;
|
|
const threshold = Math.log(epsilon) + 2;
|
|
const softplus = unaryKernelFunc$1(Softplus, (xi) => {
|
|
const tooLarge = xi > -threshold;
|
|
const tooSmall = xi < threshold;
|
|
const expX = Math.exp(xi);
|
|
let result;
|
|
if (tooSmall) {
|
|
result = expX;
|
|
} else if (tooLarge) {
|
|
result = xi;
|
|
} else {
|
|
result = Math.log(1 + expX);
|
|
}
|
|
return result;
|
|
});
|
|
const softplusConfig = {
|
|
kernelName: Softplus,
|
|
backendName: "cpu",
|
|
kernelFunc: softplus
|
|
};
|
|
function spaceToBatchND(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { blockShape, paddings } = attrs;
|
|
assertNotComplex([x], "spaceToBatchND");
|
|
const prod2 = sizeFromShape(blockShape);
|
|
const completePaddings = [[0, 0]];
|
|
completePaddings.push(...paddings);
|
|
for (let i = 1 + blockShape.length; i < x.shape.length; ++i) {
|
|
completePaddings.push([0, 0]);
|
|
}
|
|
const paddedX = padV2Config.kernelFunc({
|
|
inputs: { x },
|
|
backend: backend2,
|
|
attrs: { paddings: completePaddings, constantValue: 0 }
|
|
});
|
|
const reshapedPaddedShape = getReshaped(paddedX.shape, blockShape, prod2, false);
|
|
const permutedReshapedPaddedPermutation = getPermuted(reshapedPaddedShape.length, blockShape.length, false);
|
|
const flattenShape = getReshapedPermuted(paddedX.shape, blockShape, prod2, false);
|
|
const reshapeInputs = { x: paddedX };
|
|
const reshapeAttrs = { shape: reshapedPaddedShape };
|
|
const paddedXReshaped = reshape({ inputs: reshapeInputs, backend: backend2, attrs: reshapeAttrs });
|
|
const transposeInputs = { x: paddedXReshaped };
|
|
const transposeAttrs = { perm: permutedReshapedPaddedPermutation };
|
|
const paddedXT = transpose$1({ inputs: transposeInputs, backend: backend2, attrs: transposeAttrs });
|
|
const resultReshapeInputs = { x: paddedXT };
|
|
const resultReshapeAttrs = { shape: flattenShape };
|
|
const result = reshape({ inputs: resultReshapeInputs, backend: backend2, attrs: resultReshapeAttrs });
|
|
backend2.disposeIntermediateTensorInfo(paddedX);
|
|
backend2.disposeIntermediateTensorInfo(paddedXReshaped);
|
|
backend2.disposeIntermediateTensorInfo(paddedXT);
|
|
return result;
|
|
}
|
|
const spaceToBatchNDConfig = {
|
|
kernelName: SpaceToBatchND,
|
|
backendName: "cpu",
|
|
kernelFunc: spaceToBatchND
|
|
};
|
|
function sparseFillEmptyRows(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { indices, values, denseShape, defaultValue } = inputs;
|
|
if (denseShape.shape.length !== 1) {
|
|
throw new Error(`Dense shape must be a vector, saw:
|
|
${denseShape.shape}`);
|
|
}
|
|
if (indices.shape.length !== 2) {
|
|
throw new Error(`Indices must be a matrix, saw:
|
|
${indices.shape}`);
|
|
}
|
|
if (values.shape.length !== 1) {
|
|
throw new Error(`Values must be a vector, saw:
|
|
${values.shape}`);
|
|
}
|
|
if (defaultValue.shape.length !== 0) {
|
|
throw new Error(`Default value must be a scalar, saw:
|
|
${defaultValue.shape}`);
|
|
}
|
|
const $indices = backend2.data.get(indices.dataId).values;
|
|
const $values = backend2.data.get(values.dataId).values;
|
|
const $denseShape = backend2.data.get(denseShape.dataId).values;
|
|
const $defaultValue = backend2.data.get(defaultValue.dataId).values[0];
|
|
const [outputIndices, outputIndicesShape, outputValues, emptyRowIndicator, reverseIndexMap] = sparseFillEmptyRowsImpl($indices, indices.shape, indices.dtype, $values, values.dtype, $denseShape, $defaultValue);
|
|
return [
|
|
backend2.makeTensorInfo(outputIndicesShape, indices.dtype, outputIndices),
|
|
backend2.makeTensorInfo([outputIndicesShape[0]], values.dtype, outputValues),
|
|
backend2.makeTensorInfo([emptyRowIndicator.length], "bool", new Uint8Array(emptyRowIndicator.map((value) => Number(value)))),
|
|
backend2.makeTensorInfo([reverseIndexMap.length], indices.dtype, new Int32Array(reverseIndexMap))
|
|
];
|
|
}
|
|
const sparseFillEmptyRowsConfig = {
|
|
kernelName: SparseFillEmptyRows,
|
|
backendName: "cpu",
|
|
kernelFunc: sparseFillEmptyRows
|
|
};
|
|
function sparseReshape(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { inputIndices, inputShape, newShape } = inputs;
|
|
if (inputIndices.shape.length !== 2) {
|
|
throw new Error(`Input indices should be a matrix but received shape
|
|
${inputIndices.shape}`);
|
|
}
|
|
if (inputShape.shape.length !== 1) {
|
|
throw new Error(`Input shape should be a vector but received shape
|
|
${inputShape.shape}`);
|
|
}
|
|
if (newShape.shape.length !== 1) {
|
|
throw new Error(`Target shape should be a vector but received shape ${newShape.shape}`);
|
|
}
|
|
const $inputShape = Array.from(backend2.data.get(inputShape.dataId).values);
|
|
const $inputIndices = backend2.data.get(inputIndices.dataId).values;
|
|
const targetShape = Array.from(backend2.data.get(newShape.dataId).values);
|
|
const [newIndices, indicesShape, outputShape] = sparseReshapeImpl($inputIndices, inputIndices.shape, inputIndices.dtype, $inputShape, targetShape);
|
|
return [
|
|
backend2.makeTensorInfo(indicesShape, inputIndices.dtype, newIndices),
|
|
backend2.makeTensorInfo([outputShape.length], newShape.dtype, new Int32Array(outputShape))
|
|
];
|
|
}
|
|
const sparseReshapeConfig = {
|
|
kernelName: SparseReshape,
|
|
backendName: "cpu",
|
|
kernelFunc: sparseReshape
|
|
};
|
|
function sparseSegmentMean(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { data, indices, segmentIds } = inputs;
|
|
if (data.shape.length < 1) {
|
|
throw new Error(`Data should be at least 1 dimensional but received scalar`);
|
|
}
|
|
if (indices.shape.length !== 1) {
|
|
throw new Error(`Indices should be a vector but received shape
|
|
${indices.shape}`);
|
|
}
|
|
if (segmentIds.shape.length !== 1) {
|
|
throw new Error(`Segment ids should be a vector but received shape
|
|
${segmentIds.shape}`);
|
|
}
|
|
if (indices.shape[0] !== segmentIds.shape[0]) {
|
|
throw new Error(`segmentIds and indices should have same size.`);
|
|
}
|
|
const $data = backend2.data.get(data.dataId).values;
|
|
const $indices = backend2.data.get(indices.dataId).values;
|
|
const $segmentIds = backend2.data.get(segmentIds.dataId).values;
|
|
const [outputData, outputDataShape] = sparseSegmentReductionImpl($data, data.shape, data.dtype, $indices, $segmentIds, true);
|
|
return backend2.makeTensorInfo(outputDataShape, data.dtype, outputData);
|
|
}
|
|
const sparseSegmentMeanConfig = {
|
|
kernelName: SparseSegmentMean,
|
|
backendName: "cpu",
|
|
kernelFunc: sparseSegmentMean
|
|
};
|
|
function sparseSegmentSum(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { data, indices, segmentIds } = inputs;
|
|
if (data.shape.length < 1) {
|
|
throw new Error(`Data should be at least 1 dimensional but received scalar`);
|
|
}
|
|
if (indices.shape.length !== 1) {
|
|
throw new Error(`Indices should be a vector but received shape
|
|
${indices.shape}`);
|
|
}
|
|
if (segmentIds.shape.length !== 1) {
|
|
throw new Error(`Segment ids should be a vector but received shape
|
|
${segmentIds.shape}`);
|
|
}
|
|
if (indices.shape[0] !== segmentIds.shape[0]) {
|
|
throw new Error(`segmentIds and indices should have same size.`);
|
|
}
|
|
const $data = backend2.data.get(data.dataId).values;
|
|
const $indices = backend2.data.get(indices.dataId).values;
|
|
const $segmentIds = backend2.data.get(segmentIds.dataId).values;
|
|
const [outputData, outputDataShape] = sparseSegmentReductionImpl($data, data.shape, data.dtype, $indices, $segmentIds);
|
|
return backend2.makeTensorInfo(outputDataShape, data.dtype, outputData);
|
|
}
|
|
const sparseSegmentSumConfig = {
|
|
kernelName: SparseSegmentSum,
|
|
backendName: "cpu",
|
|
kernelFunc: sparseSegmentSum
|
|
};
|
|
function sparseToDense(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { sparseIndices, sparseValues, defaultValue } = inputs;
|
|
const { outputShape } = attrs;
|
|
const { sliceRank, numUpdates, sliceSize, strides, outputSize } = calculateShapes(sparseValues, sparseIndices, outputShape);
|
|
const sumDupeIndices = false;
|
|
const indicesBuf = backend2.bufferSync(sparseIndices);
|
|
let outBuf;
|
|
switch (sparseValues.dtype) {
|
|
case "bool": {
|
|
const updatesBuf = backend2.bufferSync(sparseValues);
|
|
const $defaultValue = Boolean(backend2.data.get(defaultValue.dataId).values[0]);
|
|
outBuf = scatterImpl(indicesBuf, updatesBuf, outputShape, outputSize, sliceSize, numUpdates, sliceRank, strides, $defaultValue, sumDupeIndices);
|
|
break;
|
|
}
|
|
case "float32": {
|
|
const updatesBuf = backend2.bufferSync(sparseValues);
|
|
const $defaultValue = backend2.data.get(defaultValue.dataId).values[0];
|
|
outBuf = scatterImpl(indicesBuf, updatesBuf, outputShape, outputSize, sliceSize, numUpdates, sliceRank, strides, $defaultValue, sumDupeIndices);
|
|
break;
|
|
}
|
|
case "int32": {
|
|
const updatesBuf = backend2.bufferSync(sparseValues);
|
|
const $defaultValue = backend2.data.get(defaultValue.dataId).values[0];
|
|
outBuf = scatterImpl(indicesBuf, updatesBuf, outputShape, outputSize, sliceSize, numUpdates, sliceRank, strides, $defaultValue, sumDupeIndices);
|
|
break;
|
|
}
|
|
case "string": {
|
|
const updatesBuf = backend2.bufferSync(sparseValues);
|
|
const $defaultValue = decodeString(backend2.data.get(defaultValue.dataId).values[0]);
|
|
outBuf = scatterImpl(indicesBuf, updatesBuf, outputShape, outputSize, sliceSize, numUpdates, sliceRank, strides, $defaultValue, sumDupeIndices);
|
|
break;
|
|
}
|
|
default:
|
|
throw new Error(`Unsupported type ${sparseValues.dtype}`);
|
|
}
|
|
return backend2.makeTensorInfo(outputShape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const sparseToDenseConfig = {
|
|
kernelName: SparseToDense,
|
|
backendName: "cpu",
|
|
kernelFunc: sparseToDense
|
|
};
|
|
function splitV(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { numOrSizeSplits, axis } = attrs;
|
|
const $axis = parseAxisParam(axis, x.shape)[0];
|
|
const splitSizes = prepareSplitSize(x, numOrSizeSplits, $axis);
|
|
const begin = new Array(x.shape.length).fill(0);
|
|
const size = x.shape.slice();
|
|
return splitSizes.map((s) => {
|
|
const sliceSize = [...size];
|
|
sliceSize[$axis] = s;
|
|
const sliceT = slice$1({ inputs: { x }, backend: backend2, attrs: { begin, size: sliceSize } });
|
|
begin[$axis] += s;
|
|
return sliceT;
|
|
});
|
|
}
|
|
const splitVConfig = {
|
|
kernelName: SplitV,
|
|
backendName: "cpu",
|
|
kernelFunc: splitV
|
|
};
|
|
const squareConfig = {
|
|
kernelName: Square,
|
|
backendName: "cpu",
|
|
kernelFunc: ({ inputs, backend: backend2 }) => {
|
|
const { x } = inputs;
|
|
const cpuBackend = backend2;
|
|
assertNotComplex(x, "square");
|
|
const values = cpuBackend.data.get(x.dataId).values;
|
|
const newValues = new Float32Array(values.length);
|
|
for (let i = 0; i < values.length; ++i) {
|
|
const value = values[i];
|
|
newValues[i] = value * value;
|
|
}
|
|
const dataId = cpuBackend.write(newValues, x.shape, x.dtype);
|
|
return { dataId, shape: x.shape, dtype: x.dtype };
|
|
}
|
|
};
|
|
const step = unaryKernelFunc$1(Step, (xi, attrs) => {
|
|
const stepAttrs = attrs;
|
|
if (isNaN(xi)) {
|
|
return NaN;
|
|
} else {
|
|
return xi > 0 ? 1 : stepAttrs.alpha;
|
|
}
|
|
});
|
|
const stepConfig = {
|
|
kernelName: Step,
|
|
backendName: "cpu",
|
|
kernelFunc: step
|
|
};
|
|
function stridedSlice(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { begin, end, strides, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask } = attrs;
|
|
assertNotComplex(x, "stridedSlice");
|
|
const { finalShapeSparse, finalShape, isIdentity, sliceDim0, isSimpleSlice, begin: $begin, end: $end, strides: $strides } = sliceInfo(x.shape, begin, end, strides, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask);
|
|
let result;
|
|
if (isIdentity) {
|
|
result = reshape({ inputs: { x }, backend: backend2, attrs: { shape: finalShape } });
|
|
} else if (sliceDim0 || isSimpleSlice) {
|
|
assert(x.shape.length >= 1, () => `Input must have rank at least 1, got: ${x.shape.length}`);
|
|
const size = computeOutShape$2($begin, $end, $strides);
|
|
const sliced = slice$1({ inputs: { x }, backend: backend2, attrs: { begin: $begin, size } });
|
|
result = reshape({ inputs: { x: sliced }, backend: backend2, attrs: { shape: finalShape } });
|
|
backend2.disposeIntermediateTensorInfo(sliced);
|
|
} else {
|
|
const xBuf = backend2.bufferSync(x);
|
|
const outBuf = stridedSliceImpl(finalShapeSparse, xBuf, $strides, $begin);
|
|
result = backend2.makeTensorInfo(finalShape, outBuf.dtype, outBuf.values);
|
|
}
|
|
return result;
|
|
}
|
|
const stridedSliceConfig = {
|
|
kernelName: StridedSlice,
|
|
backendName: "cpu",
|
|
kernelFunc: stridedSlice
|
|
};
|
|
function stringNGrams(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { separator, nGramWidths, leftPad, rightPad: rightPad2, padWidth, preserveShortSequences } = attrs;
|
|
const { data, dataSplits } = inputs;
|
|
const $data = backend2.data.get(data.dataId).values;
|
|
const $dataSplits = backend2.data.get(dataSplits.dataId).values;
|
|
const [nGrams, nGramsSplits] = stringNGramsImpl($data, $dataSplits, separator, nGramWidths, leftPad, rightPad2, padWidth, preserveShortSequences);
|
|
return [
|
|
backend2.makeTensorInfo([nGrams.length], "string", nGrams),
|
|
backend2.makeTensorInfo(dataSplits.shape, "int32", nGramsSplits)
|
|
];
|
|
}
|
|
const stringNGramsConfig = {
|
|
kernelName: StringNGrams,
|
|
backendName: "cpu",
|
|
kernelFunc: stringNGrams
|
|
};
|
|
function stringSplit(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { skipEmpty } = attrs;
|
|
const { input, delimiter } = inputs;
|
|
if (input.dtype !== "string") {
|
|
throw new Error("Input must be of datatype string");
|
|
}
|
|
if (input.shape.length !== 1) {
|
|
throw new Error(`Input must be a vector, got shape: ${input.shape}`);
|
|
}
|
|
if (delimiter.shape.length !== 0) {
|
|
throw new Error(`Delimiter must be a scalar, got shape: ${delimiter.shape}`);
|
|
}
|
|
const $input = backend2.data.get(input.dataId).values;
|
|
const $delimiter = backend2.data.get(delimiter.dataId).values[0];
|
|
const [indices, values, shape] = stringSplitImpl($input, $delimiter, skipEmpty);
|
|
const outputSize = values.length;
|
|
return [
|
|
backend2.makeTensorInfo([outputSize, 2], "int32", indices),
|
|
backend2.makeTensorInfo([outputSize], "string", values),
|
|
backend2.makeTensorInfo([2], "int32", new Int32Array(shape))
|
|
];
|
|
}
|
|
const stringSplitConfig = {
|
|
kernelName: StringSplit,
|
|
backendName: "cpu",
|
|
kernelFunc: stringSplit
|
|
};
|
|
function stringToHashBucketFast(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { numBuckets } = attrs;
|
|
const { input } = inputs;
|
|
if (input.dtype !== "string") {
|
|
throw new Error("Input must be of datatype string");
|
|
}
|
|
if (numBuckets <= 0) {
|
|
throw new Error(`Number of buckets must be at least 1`);
|
|
}
|
|
const $input = backend2.data.get(input.dataId).values;
|
|
const output = stringToHashBucketFastImpl($input, numBuckets);
|
|
return backend2.makeTensorInfo(input.shape, "int32", output);
|
|
}
|
|
const stringToHashBucketFastConfig = {
|
|
kernelName: StringToHashBucketFast,
|
|
backendName: "cpu",
|
|
kernelFunc: stringToHashBucketFast
|
|
};
|
|
const tan = unaryKernelFunc$1(Tan, (xi) => Math.tan(xi));
|
|
const tanConfig = {
|
|
kernelName: Tan,
|
|
backendName: "cpu",
|
|
kernelFunc: tan
|
|
};
|
|
const tanh = unaryKernelFunc$1(Tanh, (xi) => Math.tanh(xi));
|
|
const tanhConfig = {
|
|
kernelName: Tanh,
|
|
backendName: "cpu",
|
|
kernelFunc: tanh
|
|
};
|
|
function tensorScatterUpdate(args) {
|
|
const { inputs, backend: backend2 } = args;
|
|
const { tensor: tensor2, indices, updates } = inputs;
|
|
const { sliceRank, numUpdates, sliceSize, strides, outputSize } = calculateShapes(updates, indices, tensor2.shape);
|
|
const sumDupeIndices = false;
|
|
const indicesBuf = backend2.bufferSync(indices);
|
|
const updatesBuf = backend2.bufferSync(updates);
|
|
const tensorBuf = backend2.bufferSync(tensor2);
|
|
const outBuf = scatterImpl(indicesBuf, updatesBuf, tensor2.shape, outputSize, sliceSize, numUpdates, sliceRank, strides, tensorBuf, sumDupeIndices);
|
|
return backend2.makeTensorInfo(tensor2.shape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const tensorScatterUpdateConfig = {
|
|
kernelName: TensorScatterUpdate,
|
|
backendName: "cpu",
|
|
kernelFunc: tensorScatterUpdate
|
|
};
|
|
function tile(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { reps } = attrs;
|
|
assertNotComplex(x, "tile");
|
|
const outBuf = tileImpl(backend2.bufferSync(x), reps);
|
|
return backend2.makeTensorInfo(outBuf.shape, outBuf.dtype, outBuf.values);
|
|
}
|
|
const tileConfig = {
|
|
kernelName: Tile,
|
|
backendName: "cpu",
|
|
kernelFunc: tile
|
|
};
|
|
function topK(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x } = inputs;
|
|
const { k, sorted } = attrs;
|
|
assertNotComplex(x, "topk");
|
|
const xVals = backend2.data.get(x.dataId).values;
|
|
const [allTopKVals, allTopKIndices] = topKImpl(xVals, x.shape, x.dtype, k, sorted);
|
|
return [
|
|
backend2.makeTensorInfo(allTopKVals.shape, allTopKVals.dtype, allTopKVals.values),
|
|
backend2.makeTensorInfo(allTopKIndices.shape, allTopKIndices.dtype, allTopKIndices.values)
|
|
];
|
|
}
|
|
const topKConfig = {
|
|
kernelName: TopK,
|
|
backendName: "cpu",
|
|
kernelFunc: topK
|
|
};
|
|
function transform(args) {
|
|
const { inputs, attrs, backend: backend2 } = args;
|
|
const { image: image2, transforms } = inputs;
|
|
const { interpolation, fillMode, fillValue, outputShape } = attrs;
|
|
const [batch, imageHeight, imageWidth, numChannels] = image2.shape;
|
|
const [outHeight, outWidth] = outputShape != null ? outputShape : [imageHeight, imageWidth];
|
|
const outShape = [batch, outHeight, outWidth, numChannels];
|
|
const inStrides = computeStrides(image2.shape);
|
|
const batchInStride = inStrides[0];
|
|
const rowInStride = inStrides[1];
|
|
const colInStride = inStrides[2];
|
|
const outStrides = computeStrides(outShape);
|
|
const batchOutStride = outStrides[0];
|
|
const rowOutStride = outStrides[1];
|
|
const colOutStride = outStrides[2];
|
|
const outVals = getTypedArrayFromDType(image2.dtype, sizeFromShape(outShape));
|
|
outVals.fill(fillValue);
|
|
const imageVals = backend2.data.get(image2.dataId).values;
|
|
const transformVals = backend2.data.get(transforms.dataId).values;
|
|
for (let b = 0; b < batch; ++b) {
|
|
const transform2 = transforms.shape[0] === 1 ? transformVals : transformVals.subarray(b * 8, b * 8 + 8);
|
|
for (let outY = 0; outY < outHeight; ++outY) {
|
|
for (let outX = 0; outX < outWidth; ++outX) {
|
|
for (let channel = 0; channel < numChannels; ++channel) {
|
|
let val;
|
|
const projection = transform2[6] * outX + transform2[7] * outY + 1;
|
|
if (projection === 0) {
|
|
continue;
|
|
}
|
|
const inX = (transform2[0] * outX + transform2[1] * outY + transform2[2]) / projection;
|
|
const inY = (transform2[3] * outX + transform2[4] * outY + transform2[5]) / projection;
|
|
const x = mapCoord(inX, imageWidth, fillMode);
|
|
const y = mapCoord(inY, imageHeight, fillMode);
|
|
switch (interpolation) {
|
|
case "nearest":
|
|
val = nearestInterpolation(imageVals, imageHeight, imageWidth, batchInStride, rowInStride, colInStride, b, y, x, channel, fillValue);
|
|
break;
|
|
case "bilinear":
|
|
val = bilinearInterpolation(imageVals, imageHeight, imageWidth, batchInStride, rowInStride, colInStride, b, y, x, channel, fillValue);
|
|
break;
|
|
default:
|
|
throw new Error(`Error in Transform: Expect 'nearest' or 'bilinear', but got ${interpolation}`);
|
|
}
|
|
const ind = b * batchOutStride + outY * rowOutStride + outX * colOutStride + channel;
|
|
outVals[ind] = val;
|
|
}
|
|
}
|
|
}
|
|
return backend2.makeTensorInfo(outShape, image2.dtype, outVals);
|
|
}
|
|
const dataId = backend2.write(outVals, outShape, image2.dtype);
|
|
return { dataId, shape: image2.shape, dtype: image2.dtype };
|
|
}
|
|
const transformConfig = {
|
|
kernelName: Transform,
|
|
backendName: "cpu",
|
|
kernelFunc: transform
|
|
};
|
|
function mapCoord(outCoord, len, mode) {
|
|
switch (mode) {
|
|
case "reflect":
|
|
return mapCoordReflect(outCoord, len);
|
|
case "wrap":
|
|
return mapCoordWrap(outCoord, len);
|
|
case "nearest":
|
|
return mapCoordNearest(outCoord, len);
|
|
case "constant":
|
|
default:
|
|
return mapCoordConstant(outCoord);
|
|
}
|
|
}
|
|
function mapCoordReflect(outCoord, len) {
|
|
let inCoord = outCoord;
|
|
if (inCoord < 0) {
|
|
if (len <= 1) {
|
|
inCoord = 0;
|
|
} else {
|
|
const sz2 = 2 * len;
|
|
if (inCoord < sz2) {
|
|
inCoord = sz2 * Math.trunc(-inCoord / sz2) + inCoord;
|
|
}
|
|
inCoord = inCoord < -len ? inCoord + sz2 : -inCoord - 1;
|
|
}
|
|
} else if (inCoord > len - 1) {
|
|
if (len <= 1) {
|
|
inCoord = 0;
|
|
} else {
|
|
const sz2 = 2 * len;
|
|
inCoord -= sz2 * Math.trunc(inCoord / sz2);
|
|
if (inCoord >= len) {
|
|
inCoord = sz2 - inCoord - 1;
|
|
}
|
|
}
|
|
}
|
|
return clamp(0, inCoord, len - 1);
|
|
}
|
|
function mapCoordWrap(outCoord, len) {
|
|
let inCoord = outCoord;
|
|
if (inCoord < 0) {
|
|
if (len <= 1) {
|
|
inCoord = 0;
|
|
} else {
|
|
const sz = len - 1;
|
|
inCoord += len * (Math.trunc(-inCoord / sz) + 1);
|
|
}
|
|
} else if (inCoord > len - 1) {
|
|
if (len <= 1) {
|
|
inCoord = 0;
|
|
} else {
|
|
const sz = len - 1;
|
|
inCoord -= len * Math.trunc(inCoord / sz);
|
|
}
|
|
}
|
|
return clamp(0, inCoord, len - 1);
|
|
}
|
|
function mapCoordConstant(outCoord, len) {
|
|
return outCoord;
|
|
}
|
|
function mapCoordNearest(outCoord, len) {
|
|
return clamp(0, outCoord, len - 1);
|
|
}
|
|
function readWithFillValue(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, y, x, channel, fillValue) {
|
|
const ind = batch * batchStride + y * rowStride + x * colStride + channel;
|
|
if (0 <= y && y < imageHeight && 0 <= x && x < imageWidth) {
|
|
return imageVals[ind];
|
|
} else {
|
|
return fillValue;
|
|
}
|
|
}
|
|
function nearestInterpolation(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, y, x, channel, fillValue) {
|
|
const $y = Math.round(y);
|
|
const $x = Math.round(x);
|
|
return readWithFillValue(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, $y, $x, channel, fillValue);
|
|
}
|
|
function bilinearInterpolation(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, y, x, channel, fillValue) {
|
|
const yFloor = Math.floor(y);
|
|
const xFloor = Math.floor(x);
|
|
const yCeil = yFloor + 1;
|
|
const xCeil = xFloor + 1;
|
|
const valueYFloor = (xCeil - x) * readWithFillValue(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, yFloor, xFloor, channel, fillValue) + (x - xFloor) * readWithFillValue(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, yFloor, xCeil, channel, fillValue);
|
|
const valueYCeil = (xCeil - x) * readWithFillValue(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, yCeil, xFloor, channel, fillValue) + (x - xFloor) * readWithFillValue(imageVals, imageHeight, imageWidth, batchStride, rowStride, colStride, batch, yCeil, xCeil, channel, fillValue);
|
|
return (yCeil - y) * valueYFloor + (y - yFloor) * valueYCeil;
|
|
}
|
|
function unique(args) {
|
|
const { inputs, attrs, backend: backend2 } = args;
|
|
const { axis } = attrs;
|
|
const { x } = inputs;
|
|
assertNotComplex(x, "unique");
|
|
const values = backend2.data.get(x.dataId).values;
|
|
const { outputValues, outputShape, indices } = uniqueImpl(values, axis, x.shape, x.dtype);
|
|
return [
|
|
backend2.makeTensorInfo(outputShape, x.dtype, outputValues),
|
|
backend2.makeTensorInfo([indices.length], "int32", indices)
|
|
];
|
|
}
|
|
const uniqueConfig = {
|
|
kernelName: Unique,
|
|
backendName: "cpu",
|
|
kernelFunc: unique
|
|
};
|
|
function unpack(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { value } = inputs;
|
|
let { axis } = attrs;
|
|
if (axis < 0) {
|
|
axis += value.shape.length;
|
|
}
|
|
const valueRank = value.shape.length;
|
|
const num = value.shape[axis];
|
|
const outShape = new Array(valueRank - 1);
|
|
let outIndex = 0;
|
|
for (let i = 0; i < valueRank; i++) {
|
|
if (i !== axis) {
|
|
outShape[outIndex++] = value.shape[i];
|
|
}
|
|
}
|
|
const begin = new Array(valueRank).fill(0);
|
|
const size = value.shape.slice();
|
|
size[axis] = 1;
|
|
const res = new Array(num);
|
|
for (let i = 0; i < res.length; i++) {
|
|
begin[axis] = i;
|
|
const tempRes = slice$1({ inputs: { x: value }, backend: backend2, attrs: { begin, size } });
|
|
res[i] = reshape({ inputs: { x: tempRes }, backend: backend2, attrs: { shape: outShape } });
|
|
backend2.disposeIntermediateTensorInfo(tempRes);
|
|
}
|
|
return res;
|
|
}
|
|
const unpackConfig = {
|
|
kernelName: Unpack,
|
|
backendName: "cpu",
|
|
kernelFunc: unpack
|
|
};
|
|
function unsortedSegmentSum(args) {
|
|
const { inputs, backend: backend2, attrs } = args;
|
|
const { x, segmentIds } = inputs;
|
|
const { numSegments } = attrs;
|
|
assertNotComplex(x, "unsortedSegmentSum");
|
|
const xRank = x.shape.length;
|
|
const segmentIdsRank = segmentIds.shape.length;
|
|
const res = [];
|
|
const intermediates = [];
|
|
const numIters = xRank - segmentIdsRank;
|
|
let $segmentIds = segmentIds;
|
|
for (let i = 0; i < numIters; ++i) {
|
|
const expanded = expandDims({ inputs: { input: $segmentIds }, backend: backend2, attrs: { dim: i + 1 } });
|
|
$segmentIds = expanded;
|
|
intermediates.push(expanded);
|
|
}
|
|
for (let i = 0; i < numSegments; ++i) {
|
|
const scalarValue = createScalarValue(i, "int32");
|
|
const segmentId = backend2.makeTensorInfo([], "int32", scalarValue);
|
|
const mask = equal$1({ inputs: { a: segmentId, b: $segmentIds }, backend: backend2 });
|
|
const maskCasted = cast$1({ inputs: { x: mask }, backend: backend2, attrs: { dtype: "float32" } });
|
|
const mul2 = multiply$1({ inputs: { a: maskCasted, b: x }, backend: backend2 });
|
|
const sumTensorInfo = sum({ inputs: { x: mul2 }, backend: backend2, attrs: { axis: 0, keepDims: false } });
|
|
res.push(sumTensorInfo);
|
|
intermediates.push(segmentId);
|
|
intermediates.push(mask);
|
|
intermediates.push(maskCasted);
|
|
intermediates.push(mul2);
|
|
intermediates.push(sumTensorInfo);
|
|
}
|
|
const result = pack({ inputs: res, backend: backend2, attrs: { axis: 0 } });
|
|
intermediates.forEach((t2) => backend2.disposeIntermediateTensorInfo(t2));
|
|
return result;
|
|
}
|
|
const unsortedSegmentSumConfig = {
|
|
kernelName: UnsortedSegmentSum,
|
|
backendName: "cpu",
|
|
kernelFunc: unsortedSegmentSum
|
|
};
|
|
const kernelConfigs = [
|
|
_fusedMatMulConfig,
|
|
absConfig$1,
|
|
acosConfig,
|
|
acoshConfig,
|
|
addConfig$1,
|
|
addNConfig,
|
|
allConfig,
|
|
anyConfig,
|
|
argMaxConfig,
|
|
argMinConfig,
|
|
asinConfig,
|
|
asinhConfig,
|
|
atanConfig,
|
|
atan2Config,
|
|
atanhConfig,
|
|
avgPoolConfig,
|
|
avgPool3DConfig,
|
|
avgPool3DGradConfig,
|
|
avgPoolGradConfig,
|
|
batchMatMulConfig,
|
|
batchNormConfig,
|
|
batchToSpaceNDConfig,
|
|
bincountConfig,
|
|
bitwiseAndConfig$1,
|
|
broadcastArgsConfig,
|
|
castConfig$1,
|
|
ceilConfig$1,
|
|
clipByValueConfig,
|
|
complexConfig$1,
|
|
complexAbsConfig,
|
|
concatConfig,
|
|
conv2DConfig,
|
|
conv2DBackpropFilterConfig,
|
|
conv2DBackpropInputConfig,
|
|
conv3DConfig,
|
|
conv3DBackpropFilterV2Config,
|
|
conv3DBackpropInputV2Config,
|
|
cosConfig,
|
|
coshConfig,
|
|
cropAndResizeConfig,
|
|
cumprodConfig,
|
|
cumsumConfig,
|
|
denseBincountConfig,
|
|
depthToSpaceConfig,
|
|
depthwiseConv2dNativeConfig,
|
|
depthwiseConv2dNativeBackpropFilterConfig,
|
|
depthwiseConv2dNativeBackpropInputConfig,
|
|
diagConfig,
|
|
dilation2DConfig,
|
|
dilation2DBackpropFilterConfig,
|
|
dilation2DBackpropInputConfig,
|
|
drawConfig,
|
|
einsumConfig,
|
|
eluConfig,
|
|
eluGradConfig,
|
|
equalConfig$1,
|
|
erfConfig,
|
|
expConfig$1,
|
|
expandDimsConfig,
|
|
expm1Config$1,
|
|
fftConfig,
|
|
fillConfig,
|
|
flipLeftRightConfig,
|
|
floorConfig$1,
|
|
floorDivConfig$1,
|
|
fusedConv2DConfig,
|
|
fusedDepthwiseConv2DConfig,
|
|
gatherNdConfig,
|
|
gatherV2Config,
|
|
greaterConfig$1,
|
|
greaterEqualConfig$1,
|
|
identityConfig$1,
|
|
ifftConfig,
|
|
imagConfig,
|
|
isFiniteConfig,
|
|
isInfConfig,
|
|
isNaNConfig,
|
|
leakyReluConfig,
|
|
lessConfig$1,
|
|
lessEqualConfig$1,
|
|
linSpaceConfig,
|
|
logConfig$1,
|
|
log1pConfig,
|
|
logicalAndConfig,
|
|
logicalNotConfig,
|
|
logicalOrConfig,
|
|
LRNConfig,
|
|
LRNGradConfig,
|
|
maxConfig,
|
|
maximumConfig$1,
|
|
maxPoolConfig,
|
|
maxPool3DConfig,
|
|
maxPool3DGradConfig,
|
|
maxPoolGradConfig,
|
|
maxPoolWithArgmaxConfig,
|
|
meanConfig,
|
|
minConfig,
|
|
minimumConfig$1,
|
|
mirrorPadConfig,
|
|
modConfig,
|
|
multinomialConfig,
|
|
multiplyConfig$1,
|
|
negConfig$1,
|
|
nonMaxSuppressionV3Config,
|
|
nonMaxSuppressionV4Config,
|
|
nonMaxSuppressionV5Config,
|
|
notEqualConfig$1,
|
|
oneHotConfig,
|
|
onesLikeConfig,
|
|
packConfig,
|
|
padV2Config,
|
|
powConfig,
|
|
preluConfig,
|
|
prodConfig$1,
|
|
raggedGatherConfig,
|
|
raggedRangeConfig,
|
|
raggedTensorToTensorConfig,
|
|
rangeConfig,
|
|
realConfig$1,
|
|
realDivConfig,
|
|
reciprocalConfig,
|
|
reluConfig,
|
|
relu6Config,
|
|
reshapeConfig,
|
|
resizeBilinearConfig,
|
|
resizeBilinearGradConfig,
|
|
resizeNearestNeighborConfig,
|
|
resizeNearestNeighborGradConfig,
|
|
reverseConfig,
|
|
rotateWithOffsetConfig,
|
|
roundConfig,
|
|
rsqrtConfig$1,
|
|
scatterNdConfig,
|
|
searchSortedConfig,
|
|
selectConfig,
|
|
seluConfig,
|
|
sigmoidConfig$1,
|
|
signConfig,
|
|
sinConfig,
|
|
sinhConfig,
|
|
sliceConfig$1,
|
|
softmaxConfig,
|
|
softplusConfig,
|
|
spaceToBatchNDConfig,
|
|
sparseFillEmptyRowsConfig,
|
|
sparseReshapeConfig,
|
|
sparseSegmentMeanConfig,
|
|
sparseSegmentSumConfig,
|
|
sparseToDenseConfig,
|
|
splitVConfig,
|
|
sqrtConfig$1,
|
|
squareConfig,
|
|
squaredDifferenceConfig$1,
|
|
staticRegexReplaceConfig$1,
|
|
stepConfig,
|
|
stridedSliceConfig,
|
|
stringNGramsConfig,
|
|
stringSplitConfig,
|
|
stringToHashBucketFastConfig,
|
|
subConfig$1,
|
|
sumConfig,
|
|
tanConfig,
|
|
tanhConfig,
|
|
tensorScatterUpdateConfig,
|
|
tileConfig,
|
|
topKConfig,
|
|
transformConfig,
|
|
transposeConfig$1,
|
|
uniqueConfig,
|
|
unpackConfig,
|
|
unsortedSegmentSumConfig,
|
|
zerosLikeConfig
|
|
];
|
|
for (const kernelConfig of kernelConfigs) {
|
|
registerKernel(kernelConfig);
|
|
}
|
|
function artplayerPluginDanmukuMask(option = {}) {
|
|
return (art) => {
|
|
const {
|
|
template: { $video, $danmuku }
|
|
} = art;
|
|
let segmenter = null;
|
|
let canvas = null;
|
|
let ctx = null;
|
|
let animationFrameId = null;
|
|
let isInitialized = false;
|
|
const config = {
|
|
solutionPath: option.solutionPath || "https://cdn.jsdelivr.net/npm/@mediapipe/selfie_segmentation",
|
|
modelSelection: option.modelSelection || 1,
|
|
smoothSegmentation: option.smoothSegmentation !== void 0 ? option.smoothSegmentation : true,
|
|
minDetectionConfidence: option.minDetectionConfidence || 0.5,
|
|
minTrackingConfidence: option.minTrackingConfidence || 0.5,
|
|
selfieMode: option.selfieMode || false,
|
|
drawContour: option.drawContour || false,
|
|
foregroundThreshold: option.foregroundThreshold || 0.5,
|
|
opacity: option.opacity || 1,
|
|
maskBlurAmount: option.maskBlurAmount || 3
|
|
};
|
|
async function initTensorFlow() {
|
|
try {
|
|
await setBackend("webgl");
|
|
} catch (error) {
|
|
console.warn("WebGL backend not available, falling back to CPU", error.message);
|
|
await setBackend("cpu");
|
|
}
|
|
}
|
|
async function initSegmenter() {
|
|
await initTensorFlow();
|
|
const model = Se.MediaPipeSelfieSegmentation;
|
|
const segmenterConfig = {
|
|
runtime: "mediapipe",
|
|
modelType: "general",
|
|
solutionPath: config.solutionPath,
|
|
modelSelection: config.modelSelection,
|
|
smoothSegmentation: config.smoothSegmentation,
|
|
minDetectionConfidence: config.minDetectionConfidence,
|
|
minTrackingConfidence: config.minTrackingConfidence,
|
|
selfieMode: config.selfieMode
|
|
};
|
|
try {
|
|
segmenter = await ke(model, segmenterConfig);
|
|
isInitialized = true;
|
|
} catch (error) {
|
|
console.error("Error initializing segmenter:", error);
|
|
isInitialized = false;
|
|
}
|
|
}
|
|
function setupDanmukuStyle() {
|
|
Object.assign($danmuku.style, {
|
|
maskMode: "alpha",
|
|
maskSize: "contain",
|
|
maskRepeat: "no-repeat",
|
|
backgroundSize: "contain",
|
|
backgroundRepeat: "no-repeat"
|
|
});
|
|
}
|
|
function createCanvas2() {
|
|
canvas = document.createElement("canvas");
|
|
ctx = canvas.getContext("2d");
|
|
}
|
|
function makeWhiteTransparent(imageData) {
|
|
const data = imageData.data;
|
|
for (let i = 0; i < data.length; i += 4) {
|
|
if (data[i] > 250 && data[i + 1] > 250 && data[i + 2] > 250) {
|
|
data[i + 3] = 0;
|
|
}
|
|
}
|
|
return imageData;
|
|
}
|
|
async function segmentBody() {
|
|
if (!isInitialized || $video.paused || $video.ended) {
|
|
animationFrameId = requestAnimationFrame(segmentBody);
|
|
return;
|
|
}
|
|
try {
|
|
canvas.width = $video.videoWidth;
|
|
canvas.height = $video.videoHeight;
|
|
const segmentation = await segmenter.segmentPeople($video);
|
|
if (!segmentation || segmentation.length === 0) {
|
|
animationFrameId = requestAnimationFrame(segmentBody);
|
|
return;
|
|
}
|
|
const foregroundColor = { r: 255, g: 255, b: 255, a: 255 };
|
|
const backgroundColor = { r: 0, g: 0, b: 0, a: 255 };
|
|
const mask = await Ue(
|
|
segmentation,
|
|
foregroundColor,
|
|
backgroundColor,
|
|
config.drawContour,
|
|
config.foregroundThreshold
|
|
);
|
|
await Ne(canvas, $video, mask, config.opacity, config.maskBlurAmount);
|
|
const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height);
|
|
ctx.putImageData(makeWhiteTransparent(imageData), 0, 0);
|
|
$danmuku.style.maskImage = `url(${canvas.toDataURL()})`;
|
|
} catch (error) {
|
|
console.error("Error in segmentBody:", error);
|
|
}
|
|
animationFrameId = requestAnimationFrame(segmentBody);
|
|
}
|
|
async function startSegmentation() {
|
|
if (!isInitialized) {
|
|
await initSegmenter();
|
|
}
|
|
if (!canvas) {
|
|
createCanvas2();
|
|
}
|
|
setupDanmukuStyle();
|
|
segmentBody();
|
|
}
|
|
function stopSegmentation() {
|
|
$danmuku.style.maskImage = "none";
|
|
if (animationFrameId) {
|
|
cancelAnimationFrame(animationFrameId);
|
|
animationFrameId = null;
|
|
}
|
|
}
|
|
art.on("ready", startSegmentation);
|
|
art.on("destroy", stopSegmentation);
|
|
return {
|
|
name: "artplayerPluginDanmukuMask",
|
|
start: startSegmentation,
|
|
stop: stopSegmentation
|
|
};
|
|
};
|
|
}
|
|
export {
|
|
artplayerPluginDanmukuMask as default
|
|
};
|