oidn-web 0.3.3 → 0.3.5
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/dist/oidn.js +510 -510
- package/dist/oidn.umd.cjs +12 -12
- package/lib/process.js +35 -35
- package/lib/process.js.map +1 -1
- package/package.json +3 -3
- package/src/process.ts +35 -35
package/dist/oidn.js
CHANGED
|
@@ -110,7 +110,7 @@ function xf(n) {
|
|
|
110
110
|
* =============================================================================
|
|
111
111
|
*/
|
|
112
112
|
const Sf = 1e-7, vf = 1e-4;
|
|
113
|
-
class
|
|
113
|
+
class If {
|
|
114
114
|
constructor(t, e) {
|
|
115
115
|
this.backend = t, this.dataMover = e, this.data = /* @__PURE__ */ new WeakMap(), this.dataIdsCount = 0;
|
|
116
116
|
}
|
|
@@ -201,7 +201,7 @@ function Bt(n) {
|
|
|
201
201
|
* limitations under the License.
|
|
202
202
|
* =============================================================================
|
|
203
203
|
*/
|
|
204
|
-
function
|
|
204
|
+
function $f(n) {
|
|
205
205
|
let t = n.length, e = 0;
|
|
206
206
|
for (; t > 0; )
|
|
207
207
|
e = Math.random() * t | 0, t--, En(n, t, e);
|
|
@@ -583,7 +583,7 @@ function ai(n, t) {
|
|
|
583
583
|
return e.set(n, s), e.get(n);
|
|
584
584
|
}
|
|
585
585
|
}
|
|
586
|
-
const Ff = "Abs", Vl = "Add", zf = "All", Uf = "ArgMax", Wf = "AvgPool", Gf = "AvgPool3D", Vf = "BatchMatMul", qf = "Bincount", ql = "Cast", jf = "ClipByValue", Hf = "Complex", Kf = "ComplexAbs", jl = "Concat", Yf = "Conv2D", Xf = "Conv2DBackpropFilter", Jf = "Conv2DBackpropInput", Zf = "Conv3D", Qf = "Conv3DBackpropInputV2", td = "CropAndResize", ed = "DepthwiseConv2dNative", nd = "RealDiv", sd = "Einsum", rd = "Elu", od = "Erf", id = "Equal", ad = "Exp", ld = "ExpandDims", ud = "Fill", cd = "FlipLeftRight", hd = "Floor", fd = "FloorDiv", dd = "GatherV2", pd = "Greater", md = "GreaterEqual", li = "Identity", gd = "Imag", bd = "LeakyRelu", yd = "Less", wd = "LessEqual", xd = "Log", Sd = "Log1p", vd = "LogicalAnd",
|
|
586
|
+
const Ff = "Abs", Vl = "Add", zf = "All", Uf = "ArgMax", Wf = "AvgPool", Gf = "AvgPool3D", Vf = "BatchMatMul", qf = "Bincount", ql = "Cast", jf = "ClipByValue", Hf = "Complex", Kf = "ComplexAbs", jl = "Concat", Yf = "Conv2D", Xf = "Conv2DBackpropFilter", Jf = "Conv2DBackpropInput", Zf = "Conv3D", Qf = "Conv3DBackpropInputV2", td = "CropAndResize", ed = "DepthwiseConv2dNative", nd = "RealDiv", sd = "Einsum", rd = "Elu", od = "Erf", id = "Equal", ad = "Exp", ld = "ExpandDims", ud = "Fill", cd = "FlipLeftRight", hd = "Floor", fd = "FloorDiv", dd = "GatherV2", pd = "Greater", md = "GreaterEqual", li = "Identity", gd = "Imag", bd = "LeakyRelu", yd = "Less", wd = "LessEqual", xd = "Log", Sd = "Log1p", vd = "LogicalAnd", Id = "Max", $d = "Maximum", Hl = "MaxPool", Ad = "MaxPool3D", Ed = "Mean", _d = "Min", Cd = "Minimum", kd = "MirrorPad", Td = "Multiply", Nd = "Neg", Dd = "NonMaxSuppressionV3", Pd = "NonMaxSuppressionV4", Rd = "NonMaxSuppressionV5", Ld = "OnesLike", Od = "OneHot", Md = "Pack", Kl = "PadV2", Bd = "Pow", Fd = "Prelu", zd = "Range", Ud = "Real", Wd = "Relu", Gd = "Reshape", Yl = "ResizeNearestNeighbor", Vd = "ResizeBilinear", qd = "Relu6", jd = "Round", Hd = "Select", Kd = "Selu", Xl = "Slice", Yd = "Sigmoid", Xd = "Softplus", Jd = "Sqrt", Zd = "Sum", Qd = "SplitV", tp = "Softmax", ep = "Sub", np = "Tanh", Jl = "Tile", sp = "Transform", io = "Transpose", rp = "Unpack", op = "ZerosLike", ip = "Step", ap = "RotateWithOffset", No = "FusedConv2D";
|
|
587
587
|
/**
|
|
588
588
|
* @license
|
|
589
589
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
@@ -1394,7 +1394,7 @@ function au(n, t, e) {
|
|
|
1394
1394
|
function rt(n, t) {
|
|
1395
1395
|
return au(n, t, 8);
|
|
1396
1396
|
}
|
|
1397
|
-
function
|
|
1397
|
+
function Ia(n, t) {
|
|
1398
1398
|
return au(n, t, 4);
|
|
1399
1399
|
}
|
|
1400
1400
|
function bt(n, t) {
|
|
@@ -1420,8 +1420,8 @@ function gp(n, t = n.length) {
|
|
|
1420
1420
|
return We(o, i, e);
|
|
1421
1421
|
}
|
|
1422
1422
|
if (t >= 4) {
|
|
1423
|
-
const e = kt.add(t * 2), s =
|
|
1424
|
-
return We(s.shl(3).add(t),
|
|
1423
|
+
const e = kt.add(t * 2), s = Ia(n, 0);
|
|
1424
|
+
return We(s.shl(3).add(t), Ia(n, t - 4), e);
|
|
1425
1425
|
}
|
|
1426
1426
|
if (t > 0) {
|
|
1427
1427
|
const e = n[0], s = n[t >> 1], r = n[t - 1], o = e + (s << 8), i = t + (r << 2);
|
|
@@ -1537,7 +1537,7 @@ function fn(n, t = [], e = !1) {
|
|
|
1537
1537
|
*/
|
|
1538
1538
|
class vp {
|
|
1539
1539
|
constructor(t, e) {
|
|
1540
|
-
this.backendTimer = t, this.logger = e, e == null && (this.logger = new
|
|
1540
|
+
this.backendTimer = t, this.logger = e, e == null && (this.logger = new $p());
|
|
1541
1541
|
}
|
|
1542
1542
|
profileKernel(t, e, s) {
|
|
1543
1543
|
let r;
|
|
@@ -1558,7 +1558,7 @@ class vp {
|
|
|
1558
1558
|
for (let u = 0; u < r.length; u++) {
|
|
1559
1559
|
const c = r[u];
|
|
1560
1560
|
c.data().then((h) => {
|
|
1561
|
-
|
|
1561
|
+
Ip(h, c.dtype, t);
|
|
1562
1562
|
});
|
|
1563
1563
|
}
|
|
1564
1564
|
return {
|
|
@@ -1578,7 +1578,7 @@ class vp {
|
|
|
1578
1578
|
});
|
|
1579
1579
|
}
|
|
1580
1580
|
}
|
|
1581
|
-
function
|
|
1581
|
+
function Ip(n, t, e) {
|
|
1582
1582
|
if (t !== "float32")
|
|
1583
1583
|
return !1;
|
|
1584
1584
|
for (let s = 0; s < n.length; s++) {
|
|
@@ -1588,7 +1588,7 @@ function $p(n, t, e) {
|
|
|
1588
1588
|
}
|
|
1589
1589
|
return !1;
|
|
1590
1590
|
}
|
|
1591
|
-
class
|
|
1591
|
+
class $p {
|
|
1592
1592
|
logKernelProfile(t, e, s, r, o, i) {
|
|
1593
1593
|
const a = typeof r == "number" ? tr(`${r}ms`, 9) : r.error, l = tr(t, 25), u = e.rank, c = e.size, h = tr(e.shape.toString(), 14);
|
|
1594
1594
|
let f = "";
|
|
@@ -1706,7 +1706,7 @@ function Ep(n, t, e, s) {
|
|
|
1706
1706
|
* limitations under the License.
|
|
1707
1707
|
* =============================================================================
|
|
1708
1708
|
*/
|
|
1709
|
-
const
|
|
1709
|
+
const $a = 20, ss = 3, ao = 7;
|
|
1710
1710
|
function _p(n, t, e, s) {
|
|
1711
1711
|
const r = Kt(t), o = Cp(n, t, e, r), i = t.length, a = er(n, t, e, r, o), l = ["Tensor"];
|
|
1712
1712
|
return s && (l.push(` dtype: ${e}`), l.push(` rank: ${i}`), l.push(` shape: [${t}]`), l.push(" values:")), l.push(a.map((u) => " " + u).join(`
|
|
@@ -1740,7 +1740,7 @@ function er(n, t, e, s, r, o = !0) {
|
|
|
1740
1740
|
return e === "bool" ? [lu(n[0])] : [n[0].toString()];
|
|
1741
1741
|
}
|
|
1742
1742
|
if (l === 1) {
|
|
1743
|
-
if (a >
|
|
1743
|
+
if (a > $a) {
|
|
1744
1744
|
const m = ss * i;
|
|
1745
1745
|
let b = Array.from(n.slice(0, m)), y = Array.from(n.slice((a - ss) * i, a * i));
|
|
1746
1746
|
return e === "complex64" && (b = us(b), y = us(y)), [
|
|
@@ -1752,7 +1752,7 @@ function er(n, t, e, s, r, o = !0) {
|
|
|
1752
1752
|
];
|
|
1753
1753
|
}
|
|
1754
1754
|
const u = t.slice(1), c = s.slice(1), h = s[0] * i, f = [];
|
|
1755
|
-
if (a >
|
|
1755
|
+
if (a > $a) {
|
|
1756
1756
|
for (let g = 0; g < ss; g++) {
|
|
1757
1757
|
const m = g * h, b = m + h;
|
|
1758
1758
|
f.push(...er(
|
|
@@ -2329,7 +2329,7 @@ class Mn {
|
|
|
2329
2329
|
/**
|
|
2330
2330
|
* Initializes a backend by looking up the backend name in the factory
|
|
2331
2331
|
* registry and calling the factory method. Returns a boolean representing
|
|
2332
|
-
* whether the initialization of the backend
|
|
2332
|
+
* whether the initialization of the backend suceeded. Throws an error if
|
|
2333
2333
|
* there is no backend in the factory registry.
|
|
2334
2334
|
*/
|
|
2335
2335
|
initializeBackend(t) {
|
|
@@ -2844,7 +2844,7 @@ function _a(n, t, e, s) {
|
|
|
2844
2844
|
throw new Error(`Argument '${e}' passed to '${s}' must be ${n} tensor, but got ${t} tensor`);
|
|
2845
2845
|
}
|
|
2846
2846
|
}
|
|
2847
|
-
function
|
|
2847
|
+
function $(n, t, e, s = "numeric") {
|
|
2848
2848
|
if (n instanceof uu())
|
|
2849
2849
|
return _a(s, n.dtype, t, e), n;
|
|
2850
2850
|
let r = Cs(n);
|
|
@@ -2860,7 +2860,7 @@ function I(n, t, e, s = "numeric") {
|
|
|
2860
2860
|
function gu(n, t, e, s = "numeric") {
|
|
2861
2861
|
if (!Array.isArray(n))
|
|
2862
2862
|
throw new Error(`Argument ${t} passed to ${e} must be a \`Tensor[]\` or \`TensorLike[]\``);
|
|
2863
|
-
return n.map((o, i) =>
|
|
2863
|
+
return n.map((o, i) => $(o, `${t}[${i}]`, e, s));
|
|
2864
2864
|
}
|
|
2865
2865
|
/**
|
|
2866
2866
|
* @license
|
|
@@ -2996,7 +2996,7 @@ function C(n) {
|
|
|
2996
2996
|
* =============================================================================
|
|
2997
2997
|
*/
|
|
2998
2998
|
function Mp(n, t, e = 0) {
|
|
2999
|
-
const s =
|
|
2999
|
+
const s = $(n, "x", "pad");
|
|
3000
3000
|
if (s.rank === 0)
|
|
3001
3001
|
throw new Error("pad(scalar) is not defined. Pass non-scalar to pad");
|
|
3002
3002
|
const r = { paddings: t, constantValue: e }, o = { x: s };
|
|
@@ -3024,7 +3024,7 @@ const zp = /* @__PURE__ */ C({ pad4d_: Fp });
|
|
|
3024
3024
|
* =============================================================================
|
|
3025
3025
|
*/
|
|
3026
3026
|
function Up(n, t, e) {
|
|
3027
|
-
const s =
|
|
3027
|
+
const s = $(n, "x", "slice", "string_or_numeric");
|
|
3028
3028
|
if (s.rank === 0)
|
|
3029
3029
|
throw new Error("Slicing scalar is not possible");
|
|
3030
3030
|
const r = { x: s }, o = { begin: t, size: e };
|
|
@@ -3048,7 +3048,7 @@ const At = /* @__PURE__ */ C({ slice_: Up });
|
|
|
3048
3048
|
* =============================================================================
|
|
3049
3049
|
*/
|
|
3050
3050
|
function Wp(n, t, e) {
|
|
3051
|
-
const s =
|
|
3051
|
+
const s = $(n, "x", "slice4d");
|
|
3052
3052
|
return w(s.rank === 4, () => `slice4d expects a rank-4 tensor, but got a rank-${s.rank} tensor`), At(s, t, e);
|
|
3053
3053
|
}
|
|
3054
3054
|
const ws = /* @__PURE__ */ C({ slice4d_: Wp });
|
|
@@ -3069,7 +3069,7 @@ const ws = /* @__PURE__ */ C({ slice4d_: Wp });
|
|
|
3069
3069
|
* =============================================================================
|
|
3070
3070
|
*/
|
|
3071
3071
|
function Gp(n) {
|
|
3072
|
-
const e = { x:
|
|
3072
|
+
const e = { x: $(n, "x", "clone", "string_or_numeric") };
|
|
3073
3073
|
return A.runKernel(li, e);
|
|
3074
3074
|
}
|
|
3075
3075
|
const sn = /* @__PURE__ */ C({ clone_: Gp });
|
|
@@ -3175,7 +3175,7 @@ Pt.registerFlag("USE_SETTIMEOUTCUSTOM", () => !1);
|
|
|
3175
3175
|
* =============================================================================
|
|
3176
3176
|
*/
|
|
3177
3177
|
function Kp(n, t) {
|
|
3178
|
-
const e =
|
|
3178
|
+
const e = $(n, "real", "complex"), s = $(t, "imag", "complex");
|
|
3179
3179
|
Ef(e.shape, s.shape, `real and imag shapes, ${e.shape} and ${s.shape}, must match in call to tf.complex().`);
|
|
3180
3180
|
const r = { real: e, imag: s };
|
|
3181
3181
|
return A.runKernel(Hf, r);
|
|
@@ -3722,9 +3722,9 @@ class pn {
|
|
|
3722
3722
|
}
|
|
3723
3723
|
}
|
|
3724
3724
|
pn.URL_SCHEME = "localstorage://";
|
|
3725
|
-
const
|
|
3726
|
-
Ct.registerSaveRouter(
|
|
3727
|
-
Ct.registerLoadRouter(
|
|
3725
|
+
const Iu = (n) => V().getBool("IS_BROWSER") && !Array.isArray(n) && n.startsWith(pn.URL_SCHEME) ? fm(n.slice(pn.URL_SCHEME.length)) : null;
|
|
3726
|
+
Ct.registerSaveRouter(Iu);
|
|
3727
|
+
Ct.registerLoadRouter(Iu);
|
|
3728
3728
|
function fm(n) {
|
|
3729
3729
|
return new pn(n);
|
|
3730
3730
|
}
|
|
@@ -3946,7 +3946,7 @@ function vt(n, t = "float32", e) {
|
|
|
3946
3946
|
* =============================================================================
|
|
3947
3947
|
*/
|
|
3948
3948
|
function bm(n, t) {
|
|
3949
|
-
const e =
|
|
3949
|
+
const e = $(n, "x", "cast");
|
|
3950
3950
|
if (!Tf(t))
|
|
3951
3951
|
throw new Error(`Failed to cast to unknown dtype ${t}`);
|
|
3952
3952
|
if (t === "string" && e.dtype !== "string" || t !== "string" && e.dtype === "string")
|
|
@@ -4015,7 +4015,7 @@ Tp(wm);
|
|
|
4015
4015
|
* =============================================================================
|
|
4016
4016
|
*/
|
|
4017
4017
|
function xm(n, t) {
|
|
4018
|
-
let e =
|
|
4018
|
+
let e = $(n, "a", "add"), s = $(t, "b", "add");
|
|
4019
4019
|
[e, s] = Rt(e, s);
|
|
4020
4020
|
const r = { a: e, b: s };
|
|
4021
4021
|
return A.runKernel(Vl, r);
|
|
@@ -4038,7 +4038,7 @@ const M = /* @__PURE__ */ C({ add_: xm });
|
|
|
4038
4038
|
* =============================================================================
|
|
4039
4039
|
*/
|
|
4040
4040
|
function Sm(n, t) {
|
|
4041
|
-
let e =
|
|
4041
|
+
let e = $(n, "a", "floorDiv"), s = $(t, "b", "floorDiv");
|
|
4042
4042
|
[e, s] = Rt(e, s);
|
|
4043
4043
|
const r = { a: e, b: s };
|
|
4044
4044
|
return A.runKernel(fd, r);
|
|
@@ -4060,14 +4060,14 @@ const vm = /* @__PURE__ */ C({ floorDiv_: Sm });
|
|
|
4060
4060
|
* limitations under the License.
|
|
4061
4061
|
* =============================================================================
|
|
4062
4062
|
*/
|
|
4063
|
-
function
|
|
4064
|
-
let e =
|
|
4063
|
+
function Im(n, t) {
|
|
4064
|
+
let e = $(n, "a", "div"), s = $(t, "b", "div");
|
|
4065
4065
|
if ([e, s] = Rt(e, s), e.dtype === "int32" && s.dtype === "int32")
|
|
4066
4066
|
return vm(e, s);
|
|
4067
4067
|
const r = { a: e, b: s }, o = {};
|
|
4068
4068
|
return A.runKernel(nd, r, o);
|
|
4069
4069
|
}
|
|
4070
|
-
const Y = /* @__PURE__ */ C({ div_:
|
|
4070
|
+
const Y = /* @__PURE__ */ C({ div_: Im });
|
|
4071
4071
|
/**
|
|
4072
4072
|
* @license
|
|
4073
4073
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
@@ -4084,13 +4084,13 @@ const Y = /* @__PURE__ */ C({ div_: $m });
|
|
|
4084
4084
|
* limitations under the License.
|
|
4085
4085
|
* =============================================================================
|
|
4086
4086
|
*/
|
|
4087
|
-
function
|
|
4088
|
-
let e =
|
|
4087
|
+
function $m(n, t) {
|
|
4088
|
+
let e = $(n, "a", "mul"), s = $(t, "b", "mul");
|
|
4089
4089
|
[e, s] = Rt(e, s);
|
|
4090
4090
|
const r = { a: e, b: s };
|
|
4091
4091
|
return A.runKernel(Td, r);
|
|
4092
4092
|
}
|
|
4093
|
-
const N = /* @__PURE__ */ C({ mul_:
|
|
4093
|
+
const N = /* @__PURE__ */ C({ mul_: $m });
|
|
4094
4094
|
/**
|
|
4095
4095
|
* @license
|
|
4096
4096
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
@@ -4108,7 +4108,7 @@ const N = /* @__PURE__ */ C({ mul_: Im });
|
|
|
4108
4108
|
* =============================================================================
|
|
4109
4109
|
*/
|
|
4110
4110
|
function Am(n) {
|
|
4111
|
-
const t =
|
|
4111
|
+
const t = $(n, "x", "abs");
|
|
4112
4112
|
if (t.dtype === "complex64") {
|
|
4113
4113
|
const e = { x: t };
|
|
4114
4114
|
return A.runKernel(Kf, e);
|
|
@@ -4135,7 +4135,7 @@ const Nt = /* @__PURE__ */ C({ abs_: Am });
|
|
|
4135
4135
|
* =============================================================================
|
|
4136
4136
|
*/
|
|
4137
4137
|
function Em(n, t = null, e = !1) {
|
|
4138
|
-
const r = { x:
|
|
4138
|
+
const r = { x: $(n, "x", "all", "bool") }, o = { axis: t, keepDims: e };
|
|
4139
4139
|
return A.runKernel(zf, r, o);
|
|
4140
4140
|
}
|
|
4141
4141
|
const _m = /* @__PURE__ */ C({ all_: Em });
|
|
@@ -4156,7 +4156,7 @@ const _m = /* @__PURE__ */ C({ all_: Em });
|
|
|
4156
4156
|
* =============================================================================
|
|
4157
4157
|
*/
|
|
4158
4158
|
function Cm(n, t = 0) {
|
|
4159
|
-
const s = { x:
|
|
4159
|
+
const s = { x: $(n, "x", "argMax") }, r = { axis: t };
|
|
4160
4160
|
return A.runKernel(Uf, s, r);
|
|
4161
4161
|
}
|
|
4162
4162
|
const hr = /* @__PURE__ */ C({ argMax_: Cm });
|
|
@@ -4318,7 +4318,7 @@ function Te(n, t, e) {
|
|
|
4318
4318
|
* =============================================================================
|
|
4319
4319
|
*/
|
|
4320
4320
|
function Rm(n, t) {
|
|
4321
|
-
const s = { x:
|
|
4321
|
+
const s = { x: $(n, "x", "reshape", "string_or_numeric") }, r = { shape: t };
|
|
4322
4322
|
return A.runKernel(Gd, s, r);
|
|
4323
4323
|
}
|
|
4324
4324
|
const L = /* @__PURE__ */ C({ reshape_: Rm });
|
|
@@ -4339,7 +4339,7 @@ const L = /* @__PURE__ */ C({ reshape_: Rm });
|
|
|
4339
4339
|
* =============================================================================
|
|
4340
4340
|
*/
|
|
4341
4341
|
function Lm(n, t, e, s, r) {
|
|
4342
|
-
const o =
|
|
4342
|
+
const o = $(n, "x", "avgPool", "float32"), i = 1;
|
|
4343
4343
|
w(Vn(e, i), () => `Error in avgPool: Either strides or dilations must be 1. Got strides ${e} and dilations '${i}'`);
|
|
4344
4344
|
let a = o, l = !1;
|
|
4345
4345
|
o.rank === 3 && (l = !0, a = L(o, [1, o.shape[0], o.shape[1], o.shape[2]])), w(a.rank === 4, () => `Error in avgPool: x must be rank 4 but got rank ${a.rank}.`), Te("avgPool", s, r);
|
|
@@ -4365,7 +4365,7 @@ const Om = /* @__PURE__ */ C({ avgPool_: Lm });
|
|
|
4365
4365
|
* =============================================================================
|
|
4366
4366
|
*/
|
|
4367
4367
|
function Mm(n, t, e, s, r, o = "NDHWC") {
|
|
4368
|
-
const i =
|
|
4368
|
+
const i = $(n, "x", "avgPool3d", "float32");
|
|
4369
4369
|
let a = i, l = !1;
|
|
4370
4370
|
i.rank === 4 && (l = !0, a = L(i, [1, i.shape[0], i.shape[1], i.shape[2], i.shape[3]])), w(a.rank === 5, () => `Error in avgPool3d: x must be rank 5 but got rank ${a.rank}.`), w(o === "NDHWC", () => `Error in avgPool3d: Only NDHWC is currently supported, but got dataFormat of ${o}`), w(typeof e == "number" && e > 0 || Array.isArray(e) && e[0] > 0 && e[1] > 0 && e[2] > 0, () => `Error in avgPool3d: Stride must be > 0, but got '${e}'`), Te("avgPool3d", s, r);
|
|
4371
4371
|
const u = { x: a }, c = { filterSize: t, strides: e, pad: s, dimRoundingMode: r, dataFormat: o };
|
|
@@ -4390,7 +4390,7 @@ const Bm = /* @__PURE__ */ C({ avgPool3d_: Mm });
|
|
|
4390
4390
|
* =============================================================================
|
|
4391
4391
|
*/
|
|
4392
4392
|
function Fm(n, t, e = !1, s = !1) {
|
|
4393
|
-
let r =
|
|
4393
|
+
let r = $(n, "a", "matMul"), o = $(t, "b", "matMul");
|
|
4394
4394
|
[r, o] = Rt(r, o);
|
|
4395
4395
|
const i = { a: r, b: o }, a = { transposeA: e, transposeB: s };
|
|
4396
4396
|
return A.runKernel(Vf, i, a);
|
|
@@ -4413,7 +4413,7 @@ const we = /* @__PURE__ */ C({ matMul_: Fm });
|
|
|
4413
4413
|
* =============================================================================
|
|
4414
4414
|
*/
|
|
4415
4415
|
function zm(n) {
|
|
4416
|
-
const e = { x:
|
|
4416
|
+
const e = { x: $(n, "x", "sigmoid", "float32") };
|
|
4417
4417
|
return A.runKernel(Yd, e);
|
|
4418
4418
|
}
|
|
4419
4419
|
const pi = /* @__PURE__ */ C({ sigmoid_: zm });
|
|
@@ -4434,7 +4434,7 @@ const pi = /* @__PURE__ */ C({ sigmoid_: zm });
|
|
|
4434
4434
|
* =============================================================================
|
|
4435
4435
|
*/
|
|
4436
4436
|
function Um(n) {
|
|
4437
|
-
const e = { x:
|
|
4437
|
+
const e = { x: $(n, "x", "tanh", "float32") };
|
|
4438
4438
|
return A.runKernel(np, e);
|
|
4439
4439
|
}
|
|
4440
4440
|
const mi = /* @__PURE__ */ C({ tanh_: Um });
|
|
@@ -4455,7 +4455,7 @@ const mi = /* @__PURE__ */ C({ tanh_: Um });
|
|
|
4455
4455
|
* =============================================================================
|
|
4456
4456
|
*/
|
|
4457
4457
|
function Wm(n, t, e) {
|
|
4458
|
-
const s =
|
|
4458
|
+
const s = $(n, "x", "bincount"), r = $(t, "weights", "bincount");
|
|
4459
4459
|
w(s.dtype === "int32", () => `Error in bincount: input dtype must be int32, but got ${s.dtype}`), w(e >= 0, () => `size must be non-negative, but got ${e}.`), w(r.size === s.size || r.size === 0, () => `Error in bincount: weights must have the same size as input or0-length, but got input shape: ${s.shape}, weights shape: ${r.shape}.`);
|
|
4460
4460
|
const o = { x: s, weights: r }, i = { size: e };
|
|
4461
4461
|
return A.runKernel(qf, o, i);
|
|
@@ -4478,7 +4478,7 @@ const Gm = /* @__PURE__ */ C({ bincount_: Wm });
|
|
|
4478
4478
|
* =============================================================================
|
|
4479
4479
|
*/
|
|
4480
4480
|
function Vm(n, t) {
|
|
4481
|
-
let e =
|
|
4481
|
+
let e = $(n, "broadcastTo", "x");
|
|
4482
4482
|
const s = e.shape;
|
|
4483
4483
|
if (Me(t), t.length < e.rank)
|
|
4484
4484
|
throw new Error(`broadcastTo(): shape.length=${t.length} < input.rank=${e.rank}.`);
|
|
@@ -4538,7 +4538,7 @@ function Vr(n, t, e) {
|
|
|
4538
4538
|
* =============================================================================
|
|
4539
4539
|
*/
|
|
4540
4540
|
function qm(n, t, e) {
|
|
4541
|
-
const s =
|
|
4541
|
+
const s = $(n, "x", "clipByValue");
|
|
4542
4542
|
if (w(t <= e, () => `Error in clip: min (${t}) must be less than or equal to max (${e}).`), t === e)
|
|
4543
4543
|
return Vr(s.shape, t, s.dtype);
|
|
4544
4544
|
const r = { x: s }, o = { clipValueMin: t, clipValueMax: e };
|
|
@@ -4562,7 +4562,7 @@ const pe = /* @__PURE__ */ C({ clipByValue_: qm });
|
|
|
4562
4562
|
* =============================================================================
|
|
4563
4563
|
*/
|
|
4564
4564
|
function jm(n, t, e, s, r = "NHWC", o = [1, 1], i) {
|
|
4565
|
-
const a =
|
|
4565
|
+
const a = $(n, "x", "conv2d", "float32"), l = $(t, "filter", "conv2d", "float32");
|
|
4566
4566
|
let u = a, c = !1;
|
|
4567
4567
|
a.rank === 3 && (c = !0, u = L(a, [1, a.shape[0], a.shape[1], a.shape[2]])), w(u.rank === 4, () => `Error in conv2d: input must be rank 4, but got rank ${u.rank}.`), w(l.rank === 4, () => `Error in conv2d: filter must be rank 4, but got rank ${l.rank}.`), Te("conv2d", s, i);
|
|
4568
4568
|
const h = r === "NHWC" ? u.shape[3] : u.shape[1];
|
|
@@ -4572,7 +4572,7 @@ function jm(n, t, e, s, r = "NHWC", o = [1, 1], i) {
|
|
|
4572
4572
|
}
|
|
4573
4573
|
const gi = /* @__PURE__ */ C({ conv2d_: jm });
|
|
4574
4574
|
function Hm(n, t, e, s, r = "NWC", o = 1, i) {
|
|
4575
|
-
const a =
|
|
4575
|
+
const a = $(n, "x", "conv1d"), l = $(t, "filter", "conv1d");
|
|
4576
4576
|
let u = a, c = !1;
|
|
4577
4577
|
a.rank === 2 && (c = !0, u = L(a, [1, a.shape[0], a.shape[1]])), w(u.rank === 3, () => `Error in conv1d: input must be rank 3, but got rank ${u.rank}.`), w(l.rank === 3, () => `Error in conv1d: filter must be rank 3, but got rank ${l.rank}.`), Te("conv1d", s, i), w(u.shape[2] === l.shape[1], () => `Error in conv1d: depth of input (${u.shape[2]}) must match input depth for filter ${l.shape[1]}.`), w(Vn(e, o), () => `Error in conv1D: Either stride or dilation must be 1. Got stride ${e} and dilation '${o}'`), w(Bn(o), () => "Error in conv1D: Dilated rates should be larger than 0."), w(Bn(e), () => "Error in conv1D: Stride should be larger than 0."), w(r === "NWC", () => `Error in conv1d: got dataFormat of ${r} but only NWC is currently supported.`);
|
|
4578
4578
|
const h = L(l, [1, l.shape[0], l.shape[1], l.shape[2]]), f = L(u, [u.shape[0], 1, u.shape[1], u.shape[2]]), m = gi(f, h, [1, e], s, "NHWC", [1, o], i);
|
|
@@ -4604,10 +4604,10 @@ function Ym(n, t, e, s, r, o = "NHWC", i) {
|
|
|
4604
4604
|
const f = { dy: l, filter: e }, d = { strides: s, pad: r, dataFormat: o, dimRoundingMode: i, inputShape: a }, p = A.runKernel(Jf, f, d);
|
|
4605
4605
|
return u ? L(p, [p.shape[1], p.shape[2], p.shape[3]]) : p;
|
|
4606
4606
|
}
|
|
4607
|
-
const
|
|
4607
|
+
const $u = /* @__PURE__ */ C({ conv2DBackpropInput_: Ym });
|
|
4608
4608
|
function Xm(n, t, e, s, r, o) {
|
|
4609
|
-
const i =
|
|
4610
|
-
return
|
|
4609
|
+
const i = $(n, "x", "conv2dTranspose"), a = $(t, "filter", "conv2dTranspose");
|
|
4610
|
+
return $u(e, i, a, s, r, "NHWC", o);
|
|
4611
4611
|
}
|
|
4612
4612
|
const Jm = /* @__PURE__ */ C({ conv2dTranspose_: Xm });
|
|
4613
4613
|
/**
|
|
@@ -4627,7 +4627,7 @@ const Jm = /* @__PURE__ */ C({ conv2dTranspose_: Xm });
|
|
|
4627
4627
|
* =============================================================================
|
|
4628
4628
|
*/
|
|
4629
4629
|
function Zm(n, t, e, s, r = "NDHWC", o = [1, 1, 1]) {
|
|
4630
|
-
const i =
|
|
4630
|
+
const i = $(n, "x", "conv3d"), a = $(t, "filter", "conv3d");
|
|
4631
4631
|
let l = i, u = !1;
|
|
4632
4632
|
i.rank === 4 && (u = !0, l = L(i, [1, i.shape[0], i.shape[1], i.shape[2], i.shape[3]])), w(l.rank === 5, () => `Error in conv3d: input must be rank 5, but got rank ${l.rank}.`), w(a.rank === 5, () => `Error in conv3d: filter must be rank 5, but got rank ${a.rank}.`), w(l.shape[4] === a.shape[3], () => `Error in conv3d: depth of input (${l.shape[4]}) must match input depth for filter ${a.shape[3]}.`), w(Vn(e, o), () => `Error in conv3D: Either strides or dilations must be 1. Got strides ${e} and dilations '${o}'`), w(r === "NDHWC", () => `Error in conv3d: got dataFormat of ${r} but only NDHWC is currently supported.`), w(Bn(o), () => "Error in conv3D: Dilated rates should be larger than 0."), w(Bn(e), () => "Error in conv3D: Strides should be larger than 0.");
|
|
4633
4633
|
const c = { x: l, filter: a }, h = { strides: e, pad: s, dataFormat: r, dilations: o }, f = A.runKernel(Zf, c, h);
|
|
@@ -4661,7 +4661,7 @@ function tg(n, t, e, s, r) {
|
|
|
4661
4661
|
}
|
|
4662
4662
|
const eg = /* @__PURE__ */ C({ conv3DBackpropInput_: tg });
|
|
4663
4663
|
function ng(n, t, e, s, r) {
|
|
4664
|
-
const o =
|
|
4664
|
+
const o = $(n, "x", "conv3dTranspose"), i = $(t, "filter", "conv3dTranspose");
|
|
4665
4665
|
return eg(e, o, i, s, r);
|
|
4666
4666
|
}
|
|
4667
4667
|
const sg = /* @__PURE__ */ C({ conv3dTranspose_: ng });
|
|
@@ -4682,7 +4682,7 @@ const sg = /* @__PURE__ */ C({ conv3dTranspose_: ng });
|
|
|
4682
4682
|
* =============================================================================
|
|
4683
4683
|
*/
|
|
4684
4684
|
function rg(n, t, e, s, r = "NHWC", o = [1, 1], i) {
|
|
4685
|
-
const a =
|
|
4685
|
+
const a = $(n, "x", "depthwiseConv2d", "float32"), l = $(t, "filter", "depthwiseConv2d", "float32");
|
|
4686
4686
|
let u = a, c = !1;
|
|
4687
4687
|
a.rank === 3 && (c = !0, u = L(a, [1, a.shape[0], a.shape[1], a.shape[2]])), w(u.rank === 4, () => `Error in depthwiseConv2d: input must be rank 4, but got rank ${u.rank}.`), w(l.rank === 4, () => `Error in depthwiseConv2d: filter must be rank 4, but got rank ${l.rank}.`);
|
|
4688
4688
|
const h = r === "NHWC" ? u.shape[3] : u.shape[1];
|
|
@@ -4758,7 +4758,7 @@ function Wt(n, t) {
|
|
|
4758
4758
|
* =============================================================================
|
|
4759
4759
|
*/
|
|
4760
4760
|
function ag(n, t) {
|
|
4761
|
-
let e =
|
|
4761
|
+
let e = $(n, "a", "equal", "string_or_numeric"), s = $(t, "b", "equal", "string_or_numeric");
|
|
4762
4762
|
[e, s] = Rt(e, s), Wt(e.shape, s.shape);
|
|
4763
4763
|
const r = { a: e, b: s };
|
|
4764
4764
|
return A.runKernel(id, r);
|
|
@@ -4781,7 +4781,7 @@ const mn = /* @__PURE__ */ C({ equal_: ag });
|
|
|
4781
4781
|
* =============================================================================
|
|
4782
4782
|
*/
|
|
4783
4783
|
function lg(n, t, e) {
|
|
4784
|
-
const s =
|
|
4784
|
+
const s = $(t, "a", "where"), r = $(e, "b", "where"), o = $(n, "condition", "where", "bool"), i = Wt(Wt(o.shape, s.shape), r.shape), a = sr(o, i), l = sr(s, i), u = sr(r, i), c = {
|
|
4785
4785
|
condition: a,
|
|
4786
4786
|
t: l,
|
|
4787
4787
|
e: u
|
|
@@ -4806,7 +4806,7 @@ const on = /* @__PURE__ */ C({ where_: lg });
|
|
|
4806
4806
|
* =============================================================================
|
|
4807
4807
|
*/
|
|
4808
4808
|
function ug(n) {
|
|
4809
|
-
const e = { x:
|
|
4809
|
+
const e = { x: $(n, "x", "zerosLike") };
|
|
4810
4810
|
return A.runKernel(op, e);
|
|
4811
4811
|
}
|
|
4812
4812
|
const _e = /* @__PURE__ */ C({ zerosLike_: ug });
|
|
@@ -4827,7 +4827,7 @@ const _e = /* @__PURE__ */ C({ zerosLike_: ug });
|
|
|
4827
4827
|
* =============================================================================
|
|
4828
4828
|
*/
|
|
4829
4829
|
function cg(n, ...t) {
|
|
4830
|
-
const e = t.map((r, o) =>
|
|
4830
|
+
const e = t.map((r, o) => $(r, `tensors${o}`, "einsum")), s = { equation: n };
|
|
4831
4831
|
return A.runKernel(sd, e, s);
|
|
4832
4832
|
}
|
|
4833
4833
|
const rs = /* @__PURE__ */ C({ einsum_: cg });
|
|
@@ -4848,7 +4848,7 @@ const rs = /* @__PURE__ */ C({ einsum_: cg });
|
|
|
4848
4848
|
* =============================================================================
|
|
4849
4849
|
*/
|
|
4850
4850
|
function hg(n) {
|
|
4851
|
-
const e = { x:
|
|
4851
|
+
const e = { x: $(n, "x", "elu", "float32") };
|
|
4852
4852
|
return A.runKernel(rd, e);
|
|
4853
4853
|
}
|
|
4854
4854
|
const Au = /* @__PURE__ */ C({ elu_: hg });
|
|
@@ -4869,7 +4869,7 @@ const Au = /* @__PURE__ */ C({ elu_: hg });
|
|
|
4869
4869
|
* =============================================================================
|
|
4870
4870
|
*/
|
|
4871
4871
|
function fg(n) {
|
|
4872
|
-
let t =
|
|
4872
|
+
let t = $(n, "x", "erf");
|
|
4873
4873
|
w(t.dtype === "int32" || t.dtype === "float32", () => "Input dtype must be `int32` or `float32`."), t.dtype === "int32" && (t = ot(t, "float32"));
|
|
4874
4874
|
const e = { x: t };
|
|
4875
4875
|
return A.runKernel(od, e);
|
|
@@ -4949,8 +4949,8 @@ function bg(n, t) {
|
|
|
4949
4949
|
* =============================================================================
|
|
4950
4950
|
*/
|
|
4951
4951
|
function yg(n, t = null, e = !1) {
|
|
4952
|
-
const r = { x:
|
|
4953
|
-
return A.runKernel(
|
|
4952
|
+
const r = { x: $(n, "x", "max") }, o = { reductionIndices: t, keepDims: e };
|
|
4953
|
+
return A.runKernel(Id, r, o);
|
|
4954
4954
|
}
|
|
4955
4955
|
const Ge = /* @__PURE__ */ C({ max_: yg });
|
|
4956
4956
|
/**
|
|
@@ -4970,7 +4970,7 @@ const Ge = /* @__PURE__ */ C({ max_: yg });
|
|
|
4970
4970
|
* =============================================================================
|
|
4971
4971
|
*/
|
|
4972
4972
|
function wg(n, t = null, e = !1) {
|
|
4973
|
-
const r = { x:
|
|
4973
|
+
const r = { x: $(n, "x", "min") }, o = { axis: t, keepDims: e };
|
|
4974
4974
|
return A.runKernel(_d, r, o);
|
|
4975
4975
|
}
|
|
4976
4976
|
const Pa = /* @__PURE__ */ C({ min_: wg });
|
|
@@ -4991,7 +4991,7 @@ const Pa = /* @__PURE__ */ C({ min_: wg });
|
|
|
4991
4991
|
* =============================================================================
|
|
4992
4992
|
*/
|
|
4993
4993
|
function xg(n, t) {
|
|
4994
|
-
let e =
|
|
4994
|
+
let e = $(n, "base", "pow"), s = $(t, "exp", "pow");
|
|
4995
4995
|
[e, s] = Rt(e, s);
|
|
4996
4996
|
const r = { a: e, b: s };
|
|
4997
4997
|
return A.runKernel(Bd, r);
|
|
@@ -5037,7 +5037,7 @@ function Yt(n, t) {
|
|
|
5037
5037
|
* =============================================================================
|
|
5038
5038
|
*/
|
|
5039
5039
|
function Sg(n) {
|
|
5040
|
-
const e = { x:
|
|
5040
|
+
const e = { x: $(n, "x", "sqrt", "float32") };
|
|
5041
5041
|
return A.runKernel(Jd, e);
|
|
5042
5042
|
}
|
|
5043
5043
|
const me = /* @__PURE__ */ C({ sqrt_: Sg });
|
|
@@ -5058,7 +5058,7 @@ const me = /* @__PURE__ */ C({ sqrt_: Sg });
|
|
|
5058
5058
|
* =============================================================================
|
|
5059
5059
|
*/
|
|
5060
5060
|
function vg(n) {
|
|
5061
|
-
const t =
|
|
5061
|
+
const t = $(n, "x", "square"), e = {};
|
|
5062
5062
|
return A.runKernel("Square", { x: t }, e);
|
|
5063
5063
|
}
|
|
5064
5064
|
const Ve = /* @__PURE__ */ C({ square_: vg });
|
|
@@ -5078,13 +5078,13 @@ const Ve = /* @__PURE__ */ C({ square_: vg });
|
|
|
5078
5078
|
* limitations under the License.
|
|
5079
5079
|
* =============================================================================
|
|
5080
5080
|
*/
|
|
5081
|
-
function
|
|
5082
|
-
let s =
|
|
5081
|
+
function Ig(n, t = null, e = !1) {
|
|
5082
|
+
let s = $(n, "x", "sum");
|
|
5083
5083
|
s.dtype === "bool" && (s = ot(s, "int32"));
|
|
5084
5084
|
const r = { x: s }, o = { axis: t, keepDims: e };
|
|
5085
5085
|
return A.runKernel(Zd, r, o);
|
|
5086
5086
|
}
|
|
5087
|
-
const et = /* @__PURE__ */ C({ sum_:
|
|
5087
|
+
const et = /* @__PURE__ */ C({ sum_: Ig });
|
|
5088
5088
|
/**
|
|
5089
5089
|
* @license
|
|
5090
5090
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
@@ -5101,8 +5101,8 @@ const et = /* @__PURE__ */ C({ sum_: $g });
|
|
|
5101
5101
|
* limitations under the License.
|
|
5102
5102
|
* =============================================================================
|
|
5103
5103
|
*/
|
|
5104
|
-
function
|
|
5105
|
-
n =
|
|
5104
|
+
function $g(n, t = "euclidean", e = null, s = !1) {
|
|
5105
|
+
n = $(n, "x", "norm");
|
|
5106
5106
|
const r = Cu(n, t, e);
|
|
5107
5107
|
let o = r.shape;
|
|
5108
5108
|
if (s) {
|
|
@@ -5140,7 +5140,7 @@ function Cu(n, t, e = null) {
|
|
|
5140
5140
|
}
|
|
5141
5141
|
throw new Error(`Error in norm: invalid axis: ${e}`);
|
|
5142
5142
|
}
|
|
5143
|
-
const ku = /* @__PURE__ */ C({ norm_:
|
|
5143
|
+
const ku = /* @__PURE__ */ C({ norm_: $g });
|
|
5144
5144
|
/**
|
|
5145
5145
|
* @license
|
|
5146
5146
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
@@ -5158,7 +5158,7 @@ const ku = /* @__PURE__ */ C({ norm_: Ig });
|
|
|
5158
5158
|
* =============================================================================
|
|
5159
5159
|
*/
|
|
5160
5160
|
function Ag(n) {
|
|
5161
|
-
const e = { x:
|
|
5161
|
+
const e = { x: $(n, "x", "exp") };
|
|
5162
5162
|
return A.runKernel(ad, e);
|
|
5163
5163
|
}
|
|
5164
5164
|
const Go = /* @__PURE__ */ C({ exp_: Ag });
|
|
@@ -5179,12 +5179,12 @@ const Go = /* @__PURE__ */ C({ exp_: Ag });
|
|
|
5179
5179
|
* =============================================================================
|
|
5180
5180
|
*/
|
|
5181
5181
|
function Eg(n, t = 0) {
|
|
5182
|
-
const e =
|
|
5182
|
+
const e = $(n, "x", "expandDims", "string_or_numeric");
|
|
5183
5183
|
w(t <= e.rank, () => "Axis must be <= rank of the tensor");
|
|
5184
5184
|
const s = { input: e }, r = { dim: t };
|
|
5185
5185
|
return A.runKernel(ld, s, r);
|
|
5186
5186
|
}
|
|
5187
|
-
const
|
|
5187
|
+
const Ie = /* @__PURE__ */ C({ expandDims_: Eg });
|
|
5188
5188
|
/**
|
|
5189
5189
|
* @license
|
|
5190
5190
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
@@ -5202,7 +5202,7 @@ const $e = /* @__PURE__ */ C({ expandDims_: Eg });
|
|
|
5202
5202
|
* =============================================================================
|
|
5203
5203
|
*/
|
|
5204
5204
|
function _g(n, t) {
|
|
5205
|
-
const e =
|
|
5205
|
+
const e = $(n, "x", "tile", "string_or_numeric");
|
|
5206
5206
|
w(e.rank === t.length, () => `Error in transpose: rank of input ${e.rank} must match length of reps ${t}.`);
|
|
5207
5207
|
const s = { x: e }, r = { reps: t };
|
|
5208
5208
|
return A.runKernel(Jl, s, r);
|
|
@@ -5233,11 +5233,11 @@ function Cg(n, t, e, s = "float32") {
|
|
|
5233
5233
|
if (e == null)
|
|
5234
5234
|
return i;
|
|
5235
5235
|
if (e.length === 1)
|
|
5236
|
-
return rr(
|
|
5236
|
+
return rr(Ie(i, 0), [e[0], 1, 1]);
|
|
5237
5237
|
if (e.length === 2)
|
|
5238
|
-
return rr(
|
|
5238
|
+
return rr(Ie(Ie(i, 0), 0), [e[0], e[1], 1, 1]);
|
|
5239
5239
|
if (e.length === 3)
|
|
5240
|
-
return rr(
|
|
5240
|
+
return rr(Ie(Ie(Ie(i, 0), 0), 0), [
|
|
5241
5241
|
e[0],
|
|
5242
5242
|
e[1],
|
|
5243
5243
|
e[2],
|
|
@@ -5264,7 +5264,7 @@ const Tu = /* @__PURE__ */ C({ eye_: Cg });
|
|
|
5264
5264
|
* =============================================================================
|
|
5265
5265
|
*/
|
|
5266
5266
|
function kg(n) {
|
|
5267
|
-
const e = { x:
|
|
5267
|
+
const e = { x: $(n, "x", "floor", "float32") };
|
|
5268
5268
|
return A.runKernel(hd, e);
|
|
5269
5269
|
}
|
|
5270
5270
|
const Tg = /* @__PURE__ */ C({ floor_: kg });
|
|
@@ -5285,7 +5285,7 @@ const Tg = /* @__PURE__ */ C({ floor_: kg });
|
|
|
5285
5285
|
* =============================================================================
|
|
5286
5286
|
*/
|
|
5287
5287
|
function Ng(n, t, e = 0, s = 0) {
|
|
5288
|
-
const r =
|
|
5288
|
+
const r = $(n, "x", "gather"), o = $(t, "indices", "gather", "int32"), i = { x: r, indices: o }, a = { axis: e, batchDims: s };
|
|
5289
5289
|
return A.runKernel(dd, i, a);
|
|
5290
5290
|
}
|
|
5291
5291
|
const Dg = /* @__PURE__ */ C({ gather_: Ng });
|
|
@@ -5306,7 +5306,7 @@ const Dg = /* @__PURE__ */ C({ gather_: Ng });
|
|
|
5306
5306
|
* =============================================================================
|
|
5307
5307
|
*/
|
|
5308
5308
|
function Pg(n, t) {
|
|
5309
|
-
let e =
|
|
5309
|
+
let e = $(n, "a", "greater", "string_or_numeric"), s = $(t, "b", "greater", "string_or_numeric");
|
|
5310
5310
|
[e, s] = Rt(e, s), Wt(e.shape, s.shape);
|
|
5311
5311
|
const r = { a: e, b: s };
|
|
5312
5312
|
return A.runKernel(pd, r);
|
|
@@ -5329,7 +5329,7 @@ const ks = /* @__PURE__ */ C({ greater_: Pg });
|
|
|
5329
5329
|
* =============================================================================
|
|
5330
5330
|
*/
|
|
5331
5331
|
function Rg(n, t) {
|
|
5332
|
-
let e =
|
|
5332
|
+
let e = $(n, "a", "greaterEqual", "string_or_numeric"), s = $(t, "b", "greaterEqual", "string_or_numeric");
|
|
5333
5333
|
[e, s] = Rt(e, s), Wt(e.shape, s.shape);
|
|
5334
5334
|
const r = { a: e, b: s };
|
|
5335
5335
|
return A.runKernel(md, r);
|
|
@@ -5352,7 +5352,7 @@ const Lg = /* @__PURE__ */ C({ greaterEqual_: Rg });
|
|
|
5352
5352
|
* =============================================================================
|
|
5353
5353
|
*/
|
|
5354
5354
|
function Og(n) {
|
|
5355
|
-
const e = { input:
|
|
5355
|
+
const e = { input: $(n, "input", "imag") };
|
|
5356
5356
|
return A.runKernel(gd, e);
|
|
5357
5357
|
}
|
|
5358
5358
|
const Mg = /* @__PURE__ */ C({ imag_: Og });
|
|
@@ -5373,7 +5373,7 @@ const Mg = /* @__PURE__ */ C({ imag_: Og });
|
|
|
5373
5373
|
* =============================================================================
|
|
5374
5374
|
*/
|
|
5375
5375
|
function Bg(n, t = 0.2) {
|
|
5376
|
-
const s = { x:
|
|
5376
|
+
const s = { x: $(n, "x", "leakyRelu") }, r = { alpha: t };
|
|
5377
5377
|
return A.runKernel(bd, s, r);
|
|
5378
5378
|
}
|
|
5379
5379
|
const Fg = /* @__PURE__ */ C({ leakyRelu_: Bg });
|
|
@@ -5394,7 +5394,7 @@ const Fg = /* @__PURE__ */ C({ leakyRelu_: Bg });
|
|
|
5394
5394
|
* =============================================================================
|
|
5395
5395
|
*/
|
|
5396
5396
|
function zg(n, t) {
|
|
5397
|
-
let e =
|
|
5397
|
+
let e = $(n, "a", "less", "string_or_numeric"), s = $(t, "b", "less", "string_or_numeric");
|
|
5398
5398
|
[e, s] = Rt(e, s), Wt(e.shape, s.shape);
|
|
5399
5399
|
const r = { a: e, b: s };
|
|
5400
5400
|
return A.runKernel(yd, r);
|
|
@@ -5417,7 +5417,7 @@ const Ra = /* @__PURE__ */ C({ less_: zg });
|
|
|
5417
5417
|
* =============================================================================
|
|
5418
5418
|
*/
|
|
5419
5419
|
function Ug(n, t) {
|
|
5420
|
-
let e =
|
|
5420
|
+
let e = $(n, "a", "lessEqual", "string_or_numeric"), s = $(t, "b", "lessEqual", "string_or_numeric");
|
|
5421
5421
|
[e, s] = Rt(e, s), Wt(e.shape, s.shape);
|
|
5422
5422
|
const r = { a: e, b: s };
|
|
5423
5423
|
return A.runKernel(wd, r);
|
|
@@ -5440,7 +5440,7 @@ const Nu = /* @__PURE__ */ C({ lessEqual_: Ug });
|
|
|
5440
5440
|
* =============================================================================
|
|
5441
5441
|
*/
|
|
5442
5442
|
function Wg(n) {
|
|
5443
|
-
const e = { x:
|
|
5443
|
+
const e = { x: $(n, "x", "log", "float32") };
|
|
5444
5444
|
return A.runKernel(xd, e);
|
|
5445
5445
|
}
|
|
5446
5446
|
const gn = /* @__PURE__ */ C({ log_: Wg });
|
|
@@ -5461,7 +5461,7 @@ const gn = /* @__PURE__ */ C({ log_: Wg });
|
|
|
5461
5461
|
* =============================================================================
|
|
5462
5462
|
*/
|
|
5463
5463
|
function Gg(n) {
|
|
5464
|
-
const e = { x:
|
|
5464
|
+
const e = { x: $(n, "x", "log1p") };
|
|
5465
5465
|
return A.runKernel(Sd, e);
|
|
5466
5466
|
}
|
|
5467
5467
|
const Vg = /* @__PURE__ */ C({ log1p_: Gg });
|
|
@@ -5518,7 +5518,7 @@ function Vo(n) {
|
|
|
5518
5518
|
* =============================================================================
|
|
5519
5519
|
*/
|
|
5520
5520
|
function jg(n) {
|
|
5521
|
-
const e = { x:
|
|
5521
|
+
const e = { x: $(n, "x", "neg") };
|
|
5522
5522
|
return A.runKernel(Nd, e);
|
|
5523
5523
|
}
|
|
5524
5524
|
const qn = /* @__PURE__ */ C({ neg_: jg });
|
|
@@ -5539,7 +5539,7 @@ const qn = /* @__PURE__ */ C({ neg_: jg });
|
|
|
5539
5539
|
* =============================================================================
|
|
5540
5540
|
*/
|
|
5541
5541
|
function Hg(n) {
|
|
5542
|
-
const e = { x:
|
|
5542
|
+
const e = { x: $(n, "x", "softplus") };
|
|
5543
5543
|
return A.runKernel(Xd, e);
|
|
5544
5544
|
}
|
|
5545
5545
|
const yi = /* @__PURE__ */ C({ softplus_: Hg });
|
|
@@ -5560,7 +5560,7 @@ const yi = /* @__PURE__ */ C({ softplus_: Hg });
|
|
|
5560
5560
|
* =============================================================================
|
|
5561
5561
|
*/
|
|
5562
5562
|
function Kg(n, t) {
|
|
5563
|
-
let e =
|
|
5563
|
+
let e = $(n, "a", "sub"), s = $(t, "b", "sub");
|
|
5564
5564
|
[e, s] = Rt(e, s);
|
|
5565
5565
|
const r = { a: e, b: s };
|
|
5566
5566
|
return A.runKernel(ep, r);
|
|
@@ -5583,7 +5583,7 @@ const Z = /* @__PURE__ */ C({ sub_: Kg });
|
|
|
5583
5583
|
* =============================================================================
|
|
5584
5584
|
*/
|
|
5585
5585
|
function Yg(n, t = -1) {
|
|
5586
|
-
const e =
|
|
5586
|
+
const e = $(n, "logits", "logSoftmax");
|
|
5587
5587
|
if (t === -1 && (t = e.rank - 1), t !== e.rank - 1)
|
|
5588
5588
|
throw Error(`Log Softmax along a non-last dimension is not yet supported. Logits was rank ${e.rank} and axis was ${t}`);
|
|
5589
5589
|
return Vo((r, o) => {
|
|
@@ -5612,7 +5612,7 @@ const Xg = /* @__PURE__ */ C({ logSoftmax_: Yg });
|
|
|
5612
5612
|
* =============================================================================
|
|
5613
5613
|
*/
|
|
5614
5614
|
function Jg(n, t) {
|
|
5615
|
-
const e =
|
|
5615
|
+
const e = $(n, "a", "logicalAnd", "bool"), s = $(t, "b", "logicalAnd", "bool");
|
|
5616
5616
|
Wt(e.shape, s.shape);
|
|
5617
5617
|
const r = { a: e, b: s };
|
|
5618
5618
|
return A.runKernel(vd, r);
|
|
@@ -5635,7 +5635,7 @@ const qr = /* @__PURE__ */ C({ logicalAnd_: Jg });
|
|
|
5635
5635
|
* =============================================================================
|
|
5636
5636
|
*/
|
|
5637
5637
|
function Zg(n, t, e, s, r) {
|
|
5638
|
-
const o =
|
|
5638
|
+
const o = $(n, "x", "maxPool"), i = 1;
|
|
5639
5639
|
let a = o, l = !1;
|
|
5640
5640
|
o.rank === 3 && (l = !0, a = L(o, [1, o.shape[0], o.shape[1], o.shape[2]])), w(a.rank === 4, () => `Error in maxPool: input must be rank 4 but got rank ${a.rank}.`), w(Vn(e, i), () => `Error in maxPool: Either strides or dilations must be 1. Got strides ${e} and dilations '${i}'`), Te("maxPool", s, r);
|
|
5641
5641
|
const u = { x: a }, c = { filterSize: t, strides: e, pad: s, dimRoundingMode: r }, h = A.runKernel(Hl, u, c);
|
|
@@ -5659,7 +5659,7 @@ const Qg = /* @__PURE__ */ C({ maxPool_: Zg });
|
|
|
5659
5659
|
* =============================================================================
|
|
5660
5660
|
*/
|
|
5661
5661
|
function t0(n, t = [1, 1, 1], e, s, r, o = "NDHWC") {
|
|
5662
|
-
const i =
|
|
5662
|
+
const i = $(n, "x", "maxPool3d");
|
|
5663
5663
|
let a = i, l = !1;
|
|
5664
5664
|
i.rank === 4 && (l = !0, a = L(i, [1, i.shape[0], i.shape[1], i.shape[2], i.shape[3]])), w(a.rank === 5, () => `Error in maxPool3d: x must be rank 5 but got rank ${a.rank}.`), w(o === "NDHWC", () => `Error in maxPool3d: Only NDHWC is currently supported, but got dataFormat of ${o}`), Te("maxPool3d", s, r);
|
|
5665
5665
|
const u = { x: a }, c = { filterSize: t, strides: e, pad: s, dimRoundingMode: r, dataFormat: o }, h = A.runKernel(Ad, u, c);
|
|
@@ -5683,10 +5683,10 @@ const e0 = /* @__PURE__ */ C({ maxPool3d_: t0 });
|
|
|
5683
5683
|
* =============================================================================
|
|
5684
5684
|
*/
|
|
5685
5685
|
function n0(n, t) {
|
|
5686
|
-
let e =
|
|
5686
|
+
let e = $(n, "a", "maximum"), s = $(t, "b", "maximum");
|
|
5687
5687
|
[e, s] = Rt(e, s), e.dtype === "bool" && (e = ot(e, "int32"), s = ot(s, "int32")), Wt(e.shape, s.shape);
|
|
5688
5688
|
const r = { a: e, b: s };
|
|
5689
|
-
return A.runKernel(
|
|
5689
|
+
return A.runKernel($d, r);
|
|
5690
5690
|
}
|
|
5691
5691
|
const jn = /* @__PURE__ */ C({ maximum_: n0 });
|
|
5692
5692
|
/**
|
|
@@ -5706,7 +5706,7 @@ const jn = /* @__PURE__ */ C({ maximum_: n0 });
|
|
|
5706
5706
|
* =============================================================================
|
|
5707
5707
|
*/
|
|
5708
5708
|
function s0(n, t = null, e = !1) {
|
|
5709
|
-
const r = { x:
|
|
5709
|
+
const r = { x: $(n, "x", "mean") }, o = { axis: t, keepDims: e };
|
|
5710
5710
|
return A.runKernel(Ed, r, o);
|
|
5711
5711
|
}
|
|
5712
5712
|
const St = /* @__PURE__ */ C({ mean_: s0 });
|
|
@@ -5775,7 +5775,7 @@ function wi(n, t = "float32") {
|
|
|
5775
5775
|
* =============================================================================
|
|
5776
5776
|
*/
|
|
5777
5777
|
function r0(n, t) {
|
|
5778
|
-
let e =
|
|
5778
|
+
let e = $(n, "a", "minimum"), s = $(t, "b", "minimum");
|
|
5779
5779
|
[e, s] = Rt(e, s), e.dtype === "bool" && (e = ot(e, "int32"), s = ot(s, "int32")), Wt(e.shape, s.shape);
|
|
5780
5780
|
const r = { a: e, b: s };
|
|
5781
5781
|
return A.runKernel(Cd, r);
|
|
@@ -5800,7 +5800,7 @@ const mr = /* @__PURE__ */ C({ minimum_: r0 });
|
|
|
5800
5800
|
function o0(n, t, e = 1, s = 0, r = "int32") {
|
|
5801
5801
|
if (t < 2)
|
|
5802
5802
|
throw new Error(`Error in oneHot: depth must be >=2, but it is ${t}`);
|
|
5803
|
-
const i = { indices:
|
|
5803
|
+
const i = { indices: $(n, "indices", "oneHot", "int32") }, a = { dtype: r, depth: t, onValue: e, offValue: s };
|
|
5804
5804
|
return A.runKernel(Od, i, a);
|
|
5805
5805
|
}
|
|
5806
5806
|
const i0 = /* @__PURE__ */ C({ oneHot_: o0 });
|
|
@@ -5821,7 +5821,7 @@ const i0 = /* @__PURE__ */ C({ oneHot_: o0 });
|
|
|
5821
5821
|
* =============================================================================
|
|
5822
5822
|
*/
|
|
5823
5823
|
function a0(n) {
|
|
5824
|
-
const e = { x:
|
|
5824
|
+
const e = { x: $(n, "x", "onesLike") };
|
|
5825
5825
|
return A.runKernel(Ld, e);
|
|
5826
5826
|
}
|
|
5827
5827
|
const Du = /* @__PURE__ */ C({ onesLike_: a0 });
|
|
@@ -5842,7 +5842,7 @@ const Du = /* @__PURE__ */ C({ onesLike_: a0 });
|
|
|
5842
5842
|
* =============================================================================
|
|
5843
5843
|
*/
|
|
5844
5844
|
function l0(n, t) {
|
|
5845
|
-
const e =
|
|
5845
|
+
const e = $(n, "x", "prelu"), s = $(t, "alpha", "prelu"), r = { x: e, alpha: s };
|
|
5846
5846
|
return A.runKernel(Fd, r);
|
|
5847
5847
|
}
|
|
5848
5848
|
const u0 = /* @__PURE__ */ C({ prelu_: l0 });
|
|
@@ -5958,8 +5958,8 @@ vi.exports;
|
|
|
5958
5958
|
n
|
|
5959
5959
|
);
|
|
5960
5960
|
})(vi);
|
|
5961
|
-
var f0 = vi.exports,
|
|
5962
|
-
|
|
5961
|
+
var f0 = vi.exports, Ii = { exports: {} };
|
|
5962
|
+
Ii.exports;
|
|
5963
5963
|
(function(n) {
|
|
5964
5964
|
(function(t, e, s) {
|
|
5965
5965
|
function r(a) {
|
|
@@ -6004,9 +6004,9 @@ $i.exports;
|
|
|
6004
6004
|
wn,
|
|
6005
6005
|
n
|
|
6006
6006
|
);
|
|
6007
|
-
})(
|
|
6008
|
-
var d0 =
|
|
6009
|
-
|
|
6007
|
+
})(Ii);
|
|
6008
|
+
var d0 = Ii.exports, $i = { exports: {} };
|
|
6009
|
+
$i.exports;
|
|
6010
6010
|
(function(n) {
|
|
6011
6011
|
(function(t, e, s) {
|
|
6012
6012
|
function r(a) {
|
|
@@ -6048,8 +6048,8 @@ Ii.exports;
|
|
|
6048
6048
|
// window object or global
|
|
6049
6049
|
n
|
|
6050
6050
|
);
|
|
6051
|
-
})(
|
|
6052
|
-
var p0 =
|
|
6051
|
+
})($i);
|
|
6052
|
+
var p0 = $i.exports, Ai = { exports: {} };
|
|
6053
6053
|
Ai.exports;
|
|
6054
6054
|
(function(n) {
|
|
6055
6055
|
(function(t, e, s) {
|
|
@@ -6180,12 +6180,12 @@ const g0 = {}, b0 = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineP
|
|
|
6180
6180
|
// math: package containing random, pow, and seedrandom
|
|
6181
6181
|
);
|
|
6182
6182
|
})(Pu);
|
|
6183
|
-
var w0 = Pu.exports, x0 = c0, S0 = h0, v0 = f0,
|
|
6183
|
+
var w0 = Pu.exports, x0 = c0, S0 = h0, v0 = f0, I0 = d0, $0 = p0, A0 = m0, Sn = w0;
|
|
6184
6184
|
Sn.alea = x0;
|
|
6185
6185
|
Sn.xor128 = S0;
|
|
6186
6186
|
Sn.xorwow = v0;
|
|
6187
|
-
Sn.xorshift7 =
|
|
6188
|
-
Sn.xor4096 =
|
|
6187
|
+
Sn.xorshift7 = I0;
|
|
6188
|
+
Sn.xor4096 = $0;
|
|
6189
6189
|
Sn.tychei = A0;
|
|
6190
6190
|
var Ru = Sn;
|
|
6191
6191
|
/**
|
|
@@ -6337,7 +6337,7 @@ function gr(n, t, e = 1, s = "float32") {
|
|
|
6337
6337
|
* =============================================================================
|
|
6338
6338
|
*/
|
|
6339
6339
|
function T0(n) {
|
|
6340
|
-
const e = { input:
|
|
6340
|
+
const e = { input: $(n, "input", "real") };
|
|
6341
6341
|
return A.runKernel(Ud, e);
|
|
6342
6342
|
}
|
|
6343
6343
|
const N0 = /* @__PURE__ */ C({ real_: T0 });
|
|
@@ -6358,7 +6358,7 @@ const N0 = /* @__PURE__ */ C({ real_: T0 });
|
|
|
6358
6358
|
* =============================================================================
|
|
6359
6359
|
*/
|
|
6360
6360
|
function D0(n) {
|
|
6361
|
-
const e = { x:
|
|
6361
|
+
const e = { x: $(n, "x", "relu") };
|
|
6362
6362
|
return A.runKernel(Wd, e);
|
|
6363
6363
|
}
|
|
6364
6364
|
const Ts = /* @__PURE__ */ C({ relu_: D0 });
|
|
@@ -6379,7 +6379,7 @@ const Ts = /* @__PURE__ */ C({ relu_: D0 });
|
|
|
6379
6379
|
* =============================================================================
|
|
6380
6380
|
*/
|
|
6381
6381
|
function P0(n) {
|
|
6382
|
-
const e = { x:
|
|
6382
|
+
const e = { x: $(n, "x", "relu6") };
|
|
6383
6383
|
return A.runKernel(qd, e);
|
|
6384
6384
|
}
|
|
6385
6385
|
const R0 = /* @__PURE__ */ C({ relu6_: P0 });
|
|
@@ -6400,7 +6400,7 @@ const R0 = /* @__PURE__ */ C({ relu6_: P0 });
|
|
|
6400
6400
|
* =============================================================================
|
|
6401
6401
|
*/
|
|
6402
6402
|
function L0(n) {
|
|
6403
|
-
const e = { x:
|
|
6403
|
+
const e = { x: $(n, "x", "round") };
|
|
6404
6404
|
return A.runKernel(jd, e);
|
|
6405
6405
|
}
|
|
6406
6406
|
const O0 = /* @__PURE__ */ C({ round_: L0 });
|
|
@@ -6421,12 +6421,12 @@ const O0 = /* @__PURE__ */ C({ round_: L0 });
|
|
|
6421
6421
|
* =============================================================================
|
|
6422
6422
|
*/
|
|
6423
6423
|
function M0(n) {
|
|
6424
|
-
const e = { x:
|
|
6424
|
+
const e = { x: $(n, "x", "selu") };
|
|
6425
6425
|
return A.runKernel(Kd, e);
|
|
6426
6426
|
}
|
|
6427
6427
|
const B0 = /* @__PURE__ */ C({ selu_: M0 });
|
|
6428
6428
|
function F0(n, t, e, s, r, o = [1, 1], i = "NHWC") {
|
|
6429
|
-
const a =
|
|
6429
|
+
const a = $(n, "x", "separableConv2d"), l = $(t, "depthwiseFilter", "separableConv2d"), u = $(e, "pointwiseFilter", "separableConv2d");
|
|
6430
6430
|
let c = a, h = !1;
|
|
6431
6431
|
if (a.rank === 3 && (h = !0, c = L(a, [1, a.shape[0], a.shape[1], a.shape[2]])), i === "NCHW")
|
|
6432
6432
|
throw new Error("separableConv2d currently does not support dataFormat NCHW; only NHWC is supported");
|
|
@@ -6454,7 +6454,7 @@ const z0 = /* @__PURE__ */ C({ separableConv2d_: F0 });
|
|
|
6454
6454
|
* =============================================================================
|
|
6455
6455
|
*/
|
|
6456
6456
|
function U0(n, t, e) {
|
|
6457
|
-
const s =
|
|
6457
|
+
const s = $(n, "x", "slice1d");
|
|
6458
6458
|
return w(s.rank === 1, () => `slice1d expects a rank-1 tensor, but got a rank-${s.rank} tensor`), At(s, [t], [e]);
|
|
6459
6459
|
}
|
|
6460
6460
|
const Ei = /* @__PURE__ */ C({ slice1d_: U0 });
|
|
@@ -6475,7 +6475,7 @@ const Ei = /* @__PURE__ */ C({ slice1d_: U0 });
|
|
|
6475
6475
|
* =============================================================================
|
|
6476
6476
|
*/
|
|
6477
6477
|
function W0(n, t, e) {
|
|
6478
|
-
const s =
|
|
6478
|
+
const s = $(n, "x", "slice2d");
|
|
6479
6479
|
return w(s.rank === 2, () => `slice2d expects a rank-2 tensor, but got a rank-${s.rank} tensor`), At(s, t, e);
|
|
6480
6480
|
}
|
|
6481
6481
|
const Mu = /* @__PURE__ */ C({ slice2d_: W0 });
|
|
@@ -6496,7 +6496,7 @@ const Mu = /* @__PURE__ */ C({ slice2d_: W0 });
|
|
|
6496
6496
|
* =============================================================================
|
|
6497
6497
|
*/
|
|
6498
6498
|
function G0(n, t, e) {
|
|
6499
|
-
const s =
|
|
6499
|
+
const s = $(n, "x", "slice3d");
|
|
6500
6500
|
return w(s.rank === 3, () => `slice3d expects a rank-3 tensor, but got a rank-${s.rank} tensor`), At(s, t, e);
|
|
6501
6501
|
}
|
|
6502
6502
|
const _i = /* @__PURE__ */ C({ slice3d_: G0 });
|
|
@@ -6517,7 +6517,7 @@ const _i = /* @__PURE__ */ C({ slice3d_: G0 });
|
|
|
6517
6517
|
* =============================================================================
|
|
6518
6518
|
*/
|
|
6519
6519
|
function V0(n, t = -1) {
|
|
6520
|
-
const e =
|
|
6520
|
+
const e = $(n, "logits", "softmax", "float32");
|
|
6521
6521
|
if (t === -1 && (t = e.rank - 1), t !== e.rank - 1)
|
|
6522
6522
|
throw Error(`Softmax along a non-last dimension is not yet supported. Logits was rank ${e.rank} and dim was ${t}`);
|
|
6523
6523
|
const s = { logits: e }, r = { dim: t };
|
|
@@ -6541,7 +6541,7 @@ const Bu = /* @__PURE__ */ C({ softmax_: V0 });
|
|
|
6541
6541
|
* =============================================================================
|
|
6542
6542
|
*/
|
|
6543
6543
|
function q0(n, t, e = 0) {
|
|
6544
|
-
const r = { x:
|
|
6544
|
+
const r = { x: $(n, "x", "split") }, o = { numOrSizeSplits: t, axis: e };
|
|
6545
6545
|
return A.runKernel(Qd, r, o);
|
|
6546
6546
|
}
|
|
6547
6547
|
const Fu = /* @__PURE__ */ C({ split_: q0 });
|
|
@@ -6562,7 +6562,7 @@ const Fu = /* @__PURE__ */ C({ split_: q0 });
|
|
|
6562
6562
|
* =============================================================================
|
|
6563
6563
|
*/
|
|
6564
6564
|
function j0(n, t) {
|
|
6565
|
-
const e =
|
|
6565
|
+
const e = $(n, "x", "squeeze", "string_or_numeric");
|
|
6566
6566
|
return L(e, Cf(e.shape, t).newShape);
|
|
6567
6567
|
}
|
|
6568
6568
|
const jr = /* @__PURE__ */ C({ squeeze_: j0 });
|
|
@@ -6606,7 +6606,7 @@ const br = /* @__PURE__ */ C({ stack_: H0 });
|
|
|
6606
6606
|
* =============================================================================
|
|
6607
6607
|
*/
|
|
6608
6608
|
function K0(n, t = 0) {
|
|
6609
|
-
const s = { x:
|
|
6609
|
+
const s = { x: $(n, "x", "step") }, r = { alpha: t };
|
|
6610
6610
|
return A.runKernel(ip, s, r);
|
|
6611
6611
|
}
|
|
6612
6612
|
const Y0 = /* @__PURE__ */ C({ step_: K0 });
|
|
@@ -6678,7 +6678,7 @@ const zu = /* @__PURE__ */ C({ truncatedNormal_: X0 });
|
|
|
6678
6678
|
* =============================================================================
|
|
6679
6679
|
*/
|
|
6680
6680
|
function J0(n, t = 0) {
|
|
6681
|
-
const e =
|
|
6681
|
+
const e = $(n, "x", "unstack", "string_or_numeric");
|
|
6682
6682
|
w(t >= -e.shape.length && t < e.shape.length, () => `Axis = ${t} is not in [-${e.shape.length}, ${e.shape.length})`);
|
|
6683
6683
|
const s = { value: e }, r = { axis: t };
|
|
6684
6684
|
return A.runKernel(rp, s, r);
|
|
@@ -6720,7 +6720,7 @@ function Z0(n, t = !0, e, s) {
|
|
|
6720
6720
|
* =============================================================================
|
|
6721
6721
|
*/
|
|
6722
6722
|
function Q0(n, t, e) {
|
|
6723
|
-
const s =
|
|
6723
|
+
const s = $(n, "x", "transpose");
|
|
6724
6724
|
if (t == null && (t = s.shape.map((i, a) => a).reverse()), w(s.rank === t.length, () => `Error in transpose: rank of input ${s.rank} must match length of perm ${t}.`), t.forEach((i) => {
|
|
6725
6725
|
w(i >= 0 && i < s.rank, () => `All entries in 'perm' must be between 0 and ${s.rank - 1} but got ${t}`);
|
|
6726
6726
|
}), s.rank <= 1)
|
|
@@ -6827,14 +6827,14 @@ function ib({ x: n, filter: t, strides: e, pad: s, dataFormat: r = "NHWC", dilat
|
|
|
6827
6827
|
let E = gi(n, t, e, s, r, o, i);
|
|
6828
6828
|
return a != null && (E = M(E, a)), rb(E, l, u, c);
|
|
6829
6829
|
}
|
|
6830
|
-
const h =
|
|
6830
|
+
const h = $(n, "x", "conv2d", "float32"), f = $(t, "filter", "conv2d", "float32");
|
|
6831
6831
|
let d = h, p = !1;
|
|
6832
6832
|
h.rank === 3 && (p = !0, d = L(h, [1, h.shape[0], h.shape[1], h.shape[2]])), w(d.rank === 4, () => `Error in fused conv2d: input must be rank 4, but got rank ${d.rank}.`), w(f.rank === 4, () => `Error in fused conv2d: filter must be rank 4, but got rank ${f.rank}.`), Te("fused conv2d", s, i);
|
|
6833
6833
|
const g = r === "NHWC" ? d.shape[3] : d.shape[1];
|
|
6834
6834
|
w(f.shape[2] === g, () => `Error in conv2d: depth of input (${g}) must match input depth for filter ${f.shape[2]}.`), w(Vn(e, o), () => `Error in conv2D: Either strides or dilations must be 1. Got strides ${e} and dilations '${o}'`);
|
|
6835
6835
|
const m = di(d.shape, f.shape, e, o, s, i);
|
|
6836
6836
|
let b;
|
|
6837
|
-
a != null && (b =
|
|
6837
|
+
a != null && (b = $(a, "bias", "fused conv2d"), [b] = Rt(b, h), r === "NHWC" ? Wt(m.outShape, b.shape) : (w(b.shape.length <= 1, () => `Error in fused conv2d: only supports scalar or 1-D Tensor bias for NCHW format but got the bias of rank-${b.shape.length}.`), w(b.shape.length === 0 || b.shape[0] === m.outChannels || b.shape[0] === 1, () => `Error in fused conv2d: bias shape (${b.shape}) is not compatible with the number of output channels (${m.outChannels})`)));
|
|
6838
6838
|
let y;
|
|
6839
6839
|
if (u != null) {
|
|
6840
6840
|
const E = u.shape;
|
|
@@ -6847,13 +6847,13 @@ function ib({ x: n, filter: t, strides: e, pad: s, dataFormat: r = "NHWC", dilat
|
|
|
6847
6847
|
const k = `Error in fused conv2d: PReLU activation weights (${E}) is not compatible with the output shape of the conv2d (${m.outShape}).`;
|
|
6848
6848
|
throw Error(k);
|
|
6849
6849
|
}
|
|
6850
|
-
y =
|
|
6850
|
+
y = $(u, "prelu weights", "fused conv2d");
|
|
6851
6851
|
}
|
|
6852
6852
|
const S = (E, D) => {
|
|
6853
6853
|
w(r === "NHWC", () => `Error in gradient of fused conv2D: got dataFormat of ${r} but only NHWC is currently supported.`);
|
|
6854
6854
|
const [k, T, R, B] = D, H = nb(E, R, l);
|
|
6855
6855
|
w(Wo(o), () => `Error in gradient of fused conv2D: dilation rates greater than 1 are not yet supported in gradients. Got dilations '${o}'`);
|
|
6856
|
-
const X =
|
|
6856
|
+
const X = $u(T.shape, H, k, e, s), W = eb(T, H, k.shape, e, s), U = [X, W];
|
|
6857
6857
|
if (B != null) {
|
|
6858
6858
|
const j = sb(B, H);
|
|
6859
6859
|
U.push(j);
|
|
@@ -6902,7 +6902,7 @@ const ab = /* @__PURE__ */ C({ fusedConv2d_: ib });
|
|
|
6902
6902
|
* =============================================================================
|
|
6903
6903
|
*/
|
|
6904
6904
|
function lb(n, t, e, s, r = "bilinear", o = 0) {
|
|
6905
|
-
const i =
|
|
6905
|
+
const i = $(n, "image", "cropAndResize"), a = $(t, "boxes", "cropAndResize", "float32"), l = $(e, "boxInd", "cropAndResize", "int32"), u = a.shape[0];
|
|
6906
6906
|
w(i.rank === 4, () => `Error in cropAndResize: image must be rank 4,but got rank ${i.rank}.`), w(a.rank === 2 && a.shape[1] === 4, () => `Error in cropAndResize: boxes must be have size [${u},4] but had shape ${a.shape}.`), w(l.rank === 1 && l.shape[0] === u, () => `Error in cropAndResize: boxInd must be have size [${u}] but had shape ${a.shape}.`), w(s.length === 2, () => `Error in cropAndResize: cropSize must be of length 2, but got length ${s.length}.`), w(s[0] >= 1 && s[1] >= 1, () => `cropSize must be atleast [1,1], but was ${s}`), w(r === "bilinear" || r === "nearest", () => `method must be bilinear or nearest, but was ${r}`);
|
|
6907
6907
|
const c = { image: i, boxes: a, boxInd: l }, h = { method: r, extrapolationValue: o, cropSize: s };
|
|
6908
6908
|
return A.runKernel(td, c, h);
|
|
@@ -6925,7 +6925,7 @@ const ub = /* @__PURE__ */ C({ cropAndResize_: lb });
|
|
|
6925
6925
|
* =============================================================================
|
|
6926
6926
|
*/
|
|
6927
6927
|
function cb(n) {
|
|
6928
|
-
const t =
|
|
6928
|
+
const t = $(n, "image", "flipLeftRight", "float32");
|
|
6929
6929
|
w(t.rank === 4, () => `Error in flipLeftRight: image must be rank 4,but got rank ${t.rank}.`);
|
|
6930
6930
|
const e = { image: t };
|
|
6931
6931
|
return A.runKernel(cd, e, {});
|
|
@@ -6948,7 +6948,7 @@ const hb = /* @__PURE__ */ C({ flipLeftRight_: cb });
|
|
|
6948
6948
|
* =============================================================================
|
|
6949
6949
|
*/
|
|
6950
6950
|
function fb(n) {
|
|
6951
|
-
const t =
|
|
6951
|
+
const t = $(n, "image", "grayscaleToRGB"), e = t.rank - 1, s = t.shape[e];
|
|
6952
6952
|
w(t.rank >= 2, () => `Error in grayscaleToRGB: images must be at least rank 2, but got rank ${t.rank}.`), w(s === 1, () => `Error in grayscaleToRGB: last dimension of a grayscale image should be size 1, but got size ${s}.`);
|
|
6953
6953
|
const r = new Array(t.rank);
|
|
6954
6954
|
return r.fill(1, 0, e), r[e] = 3, rr(t, r);
|
|
@@ -6971,7 +6971,7 @@ const db = /* @__PURE__ */ C({ grayscaleToRGB_: fb });
|
|
|
6971
6971
|
* =============================================================================
|
|
6972
6972
|
*/
|
|
6973
6973
|
function pb(n) {
|
|
6974
|
-
const t =
|
|
6974
|
+
const t = $(n, "image", "RGBToGrayscale"), e = t.rank - 1, s = t.shape[e];
|
|
6975
6975
|
w(t.rank >= 2, () => `Error in RGBToGrayscale: images must be at least rank 2, but got rank ${t.rank}.`), w(s === 3, () => `Error in RGBToGrayscale: last dimension of an RGB image should be size 3, but got size ${s}.`);
|
|
6976
6976
|
const r = t.dtype, o = ot(t, "float32"), i = Dt([0.2989, 0.587, 0.114]);
|
|
6977
6977
|
let a;
|
|
@@ -6994,7 +6994,7 @@ function pb(n) {
|
|
|
6994
6994
|
default:
|
|
6995
6995
|
throw new Error("Not a valid tensor rank.");
|
|
6996
6996
|
}
|
|
6997
|
-
return a =
|
|
6997
|
+
return a = Ie(a, -1), ot(a, r);
|
|
6998
6998
|
}
|
|
6999
6999
|
const mb = /* @__PURE__ */ C({ rgbToGrayscale_: pb });
|
|
7000
7000
|
/**
|
|
@@ -7014,7 +7014,7 @@ const mb = /* @__PURE__ */ C({ rgbToGrayscale_: pb });
|
|
|
7014
7014
|
* =============================================================================
|
|
7015
7015
|
*/
|
|
7016
7016
|
function gb(n, t, e = 0, s = 0.5) {
|
|
7017
|
-
const r =
|
|
7017
|
+
const r = $(n, "image", "rotateWithOffset", "float32");
|
|
7018
7018
|
w(r.rank === 4, () => `Error in rotateWithOffset: image must be rank 4,but got rank ${r.rank}.`);
|
|
7019
7019
|
const o = { image: r }, i = { radians: t, fillValue: e, center: s };
|
|
7020
7020
|
return A.runKernel(ap, o, i);
|
|
@@ -7058,7 +7058,7 @@ function Hn(n, t, e, s, r, o) {
|
|
|
7058
7058
|
* =============================================================================
|
|
7059
7059
|
*/
|
|
7060
7060
|
function yb(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY) {
|
|
7061
|
-
const o =
|
|
7061
|
+
const o = $(n, "boxes", "nonMaxSuppression", "float32"), i = $(t, "scores", "nonMaxSuppression", "float32"), a = Hn(o, i, e, s, r);
|
|
7062
7062
|
e = a.maxOutputSize, s = a.iouThreshold, r = a.scoreThreshold;
|
|
7063
7063
|
const l = { maxOutputSize: e, iouThreshold: s, scoreThreshold: r };
|
|
7064
7064
|
return A.runKernel(Dd, { boxes: o, scores: i }, l);
|
|
@@ -7112,7 +7112,7 @@ function vb(n, t, e) {
|
|
|
7112
7112
|
* limitations under the License.
|
|
7113
7113
|
* =============================================================================
|
|
7114
7114
|
*/
|
|
7115
|
-
function
|
|
7115
|
+
function Ib(n, t, e, s, r) {
|
|
7116
7116
|
return Ci(
|
|
7117
7117
|
n,
|
|
7118
7118
|
t,
|
|
@@ -7123,7 +7123,7 @@ function $b(n, t, e, s, r) {
|
|
|
7123
7123
|
/* softNmsSigma */
|
|
7124
7124
|
);
|
|
7125
7125
|
}
|
|
7126
|
-
function
|
|
7126
|
+
function $b(n, t, e, s, r, o) {
|
|
7127
7127
|
return Ci(
|
|
7128
7128
|
n,
|
|
7129
7129
|
t,
|
|
@@ -7207,9 +7207,9 @@ function La(n, t) {
|
|
|
7207
7207
|
* =============================================================================
|
|
7208
7208
|
*/
|
|
7209
7209
|
async function Cb(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY) {
|
|
7210
|
-
const o =
|
|
7210
|
+
const o = $(n, "boxes", "nonMaxSuppressionAsync"), i = $(t, "scores", "nonMaxSuppressionAsync"), a = Hn(o, i, e, s, r);
|
|
7211
7211
|
e = a.maxOutputSize, s = a.iouThreshold, r = a.scoreThreshold;
|
|
7212
|
-
const l = await Promise.all([o.data(), i.data()]), u = l[0], c = l[1], { selectedIndices: h } =
|
|
7212
|
+
const l = await Promise.all([o.data(), i.data()]), u = l[0], c = l[1], { selectedIndices: h } = Ib(u, c, e, s, r);
|
|
7213
7213
|
return o !== n && o.dispose(), i !== t && i.dispose(), Dt(h, "int32");
|
|
7214
7214
|
}
|
|
7215
7215
|
const kb = Cb;
|
|
@@ -7230,7 +7230,7 @@ const kb = Cb;
|
|
|
7230
7230
|
* =============================================================================
|
|
7231
7231
|
*/
|
|
7232
7232
|
function Tb(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY, o = 0) {
|
|
7233
|
-
const i =
|
|
7233
|
+
const i = $(n, "boxes", "nonMaxSuppression"), a = $(t, "scores", "nonMaxSuppression"), l = Hn(i, a, e, s, r, o);
|
|
7234
7234
|
e = l.maxOutputSize, s = l.iouThreshold, r = l.scoreThreshold, o = l.softNmsSigma;
|
|
7235
7235
|
const u = { boxes: i, scores: a }, c = { maxOutputSize: e, iouThreshold: s, scoreThreshold: r, softNmsSigma: o }, h = A.runKernel(Rd, u, c);
|
|
7236
7236
|
return { selectedIndices: h[0], selectedScores: h[1] };
|
|
@@ -7253,7 +7253,7 @@ const Nb = /* @__PURE__ */ C({ nonMaxSuppressionWithScore_: Tb });
|
|
|
7253
7253
|
* =============================================================================
|
|
7254
7254
|
*/
|
|
7255
7255
|
async function Db(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY, o = 0) {
|
|
7256
|
-
const i =
|
|
7256
|
+
const i = $(n, "boxes", "nonMaxSuppressionAsync"), a = $(t, "scores", "nonMaxSuppressionAsync"), l = Hn(i, a, e, s, r, o);
|
|
7257
7257
|
e = l.maxOutputSize, s = l.iouThreshold, r = l.scoreThreshold, o = l.softNmsSigma;
|
|
7258
7258
|
const u = await Promise.all([i.data(), a.data()]), c = u[0], h = u[1], { selectedIndices: f, selectedScores: d } = Ab(c, h, e, s, r, o);
|
|
7259
7259
|
return i !== n && i.dispose(), a !== t && a.dispose(), {
|
|
@@ -7279,7 +7279,7 @@ const Pb = Db;
|
|
|
7279
7279
|
* =============================================================================
|
|
7280
7280
|
*/
|
|
7281
7281
|
function Rb(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY, o = !1) {
|
|
7282
|
-
const i =
|
|
7282
|
+
const i = $(n, "boxes", "nonMaxSuppression"), a = $(t, "scores", "nonMaxSuppression"), l = Hn(
|
|
7283
7283
|
i,
|
|
7284
7284
|
a,
|
|
7285
7285
|
e,
|
|
@@ -7313,7 +7313,7 @@ const Lb = /* @__PURE__ */ C({ nonMaxSuppressionPadded_: Rb });
|
|
|
7313
7313
|
* =============================================================================
|
|
7314
7314
|
*/
|
|
7315
7315
|
async function Ob(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY, o = !1) {
|
|
7316
|
-
const i =
|
|
7316
|
+
const i = $(n, "boxes", "nonMaxSuppressionAsync"), a = $(t, "scores", "nonMaxSuppressionAsync"), l = Hn(
|
|
7317
7317
|
i,
|
|
7318
7318
|
a,
|
|
7319
7319
|
e,
|
|
@@ -7321,7 +7321,7 @@ async function Ob(n, t, e, s = 0.5, r = Number.NEGATIVE_INFINITY, o = !1) {
|
|
|
7321
7321
|
r,
|
|
7322
7322
|
null
|
|
7323
7323
|
/* softNmsSigma */
|
|
7324
|
-
), u = l.maxOutputSize, c = l.iouThreshold, h = l.scoreThreshold, [f, d] = await Promise.all([i.data(), a.data()]), { selectedIndices: p, validOutputs: g } =
|
|
7324
|
+
), u = l.maxOutputSize, c = l.iouThreshold, h = l.scoreThreshold, [f, d] = await Promise.all([i.data(), a.data()]), { selectedIndices: p, validOutputs: g } = $b(f, d, u, c, h, o);
|
|
7325
7325
|
return i !== n && i.dispose(), a !== t && a.dispose(), {
|
|
7326
7326
|
selectedIndices: Dt(p, "int32"),
|
|
7327
7327
|
validOutputs: Yt(g, "int32")
|
|
@@ -7345,7 +7345,7 @@ const Mb = Ob;
|
|
|
7345
7345
|
* =============================================================================
|
|
7346
7346
|
*/
|
|
7347
7347
|
function Bb(n, t, e = !1, s = !1) {
|
|
7348
|
-
const r =
|
|
7348
|
+
const r = $(n, "images", "resizeBilinear");
|
|
7349
7349
|
w(r.rank === 3 || r.rank === 4, () => `Error in resizeBilinear: x must be rank 3 or 4, but got rank ${r.rank}.`), w(t.length === 2, () => `Error in resizeBilinear: new shape must 2D, but got shape ${t}.`), w(s === !1 || e === !1, () => "Error in resizeBilinear: If halfPixelCenters is true, alignCorners must be false.");
|
|
7350
7350
|
let o = r, i = !1;
|
|
7351
7351
|
r.rank === 3 && (i = !0, o = L(r, [1, r.shape[0], r.shape[1], r.shape[2]]));
|
|
@@ -7370,7 +7370,7 @@ const Fb = /* @__PURE__ */ C({ resizeBilinear_: Bb });
|
|
|
7370
7370
|
* =============================================================================
|
|
7371
7371
|
*/
|
|
7372
7372
|
function zb(n, t, e = !1, s = !1) {
|
|
7373
|
-
const r =
|
|
7373
|
+
const r = $(n, "images", "resizeNearestNeighbor");
|
|
7374
7374
|
w(r.rank === 3 || r.rank === 4, () => `Error in resizeNearestNeighbor: x must be rank 3 or 4, but got rank ${r.rank}.`), w(t.length === 2, () => `Error in resizeNearestNeighbor: new shape must 2D, but got shape ${t}.`), w(r.dtype === "float32" || r.dtype === "int32", () => "`images` must have `int32` or `float32` as dtype"), w(s === !1 || e === !1, () => "Error in resizeNearestNeighbor: If halfPixelCenters is true, alignCorners must be false.");
|
|
7375
7375
|
let o = r, i = !1;
|
|
7376
7376
|
r.rank === 3 && (i = !0, o = L(r, [1, r.shape[0], r.shape[1], r.shape[2]]));
|
|
@@ -7395,7 +7395,7 @@ const Ub = /* @__PURE__ */ C({ resizeNearestNeighbor_: zb });
|
|
|
7395
7395
|
* =============================================================================
|
|
7396
7396
|
*/
|
|
7397
7397
|
function Wb(n, t = "binary", e = !1, s = 0.5) {
|
|
7398
|
-
const r =
|
|
7398
|
+
const r = $(n, "image", "threshold"), o = 0.2989, i = 0.587, a = 0.114, l = r.shape[0] * r.shape[1];
|
|
7399
7399
|
let u = N(Dt([s]), 255), c, h, f, d;
|
|
7400
7400
|
if (w(r.rank === 3, () => `Error in threshold: image must be rank 3,but got rank ${r.rank}.`), w(r.shape[2] === 3 || r.shape[2] === 1, () => `Error in threshold: image color channel must be equal to 3 or 1but got ${r.shape[2]}.`), w(r.dtype === "int32" || r.dtype === "float32", () => `Error in dtype: image dtype must be int32 or float32,but got dtype ${r.dtype}.`), w(t === "otsu" || t === "binary", () => `Method must be binary or otsu, but was ${t}`), r.shape[2] === 3) {
|
|
7401
7401
|
[c, h, f] = Fu(r, [1, 1, 1], -1);
|
|
@@ -7443,7 +7443,7 @@ const Vb = /* @__PURE__ */ C({ threshold_: Wb });
|
|
|
7443
7443
|
* =============================================================================
|
|
7444
7444
|
*/
|
|
7445
7445
|
function qb(n, t, e = "nearest", s = "constant", r = 0, o) {
|
|
7446
|
-
const i =
|
|
7446
|
+
const i = $(n, "image", "transform", "float32"), a = $(t, "transforms", "transform", "float32");
|
|
7447
7447
|
w(i.rank === 4, () => `Error in transform: image must be rank 4,but got rank ${i.rank}.`), w(a.rank === 2 && (a.shape[0] === i.shape[0] || a.shape[0] === 1) && a.shape[1] === 8, () => "Error in transform: Input transform should be batch x 8 or 1 x 8"), w(o == null || o.length === 2, () => `Error in transform: outputShape must be [height, width] or null, but got ${o}.`);
|
|
7448
7448
|
const l = { image: i, transforms: a }, u = { interpolation: e, fillMode: s, fillValue: r, outputShape: o };
|
|
7449
7449
|
return A.runKernel(sp, l, u);
|
|
@@ -7466,11 +7466,11 @@ const jb = /* @__PURE__ */ C({ transform_: qb });
|
|
|
7466
7466
|
* =============================================================================
|
|
7467
7467
|
*/
|
|
7468
7468
|
function Hb(n, t, e) {
|
|
7469
|
-
const s =
|
|
7469
|
+
const s = $(n, "a", "bandPart");
|
|
7470
7470
|
w(s.rank >= 2, () => `bandPart(): Rank must be at least 2, got ${s.rank}.`);
|
|
7471
7471
|
const r = s.shape, [o, i] = s.shape.slice(-2);
|
|
7472
7472
|
let a, l;
|
|
7473
|
-
typeof t == "number" ? (w(t % 1 === 0, () => `bandPart(): numLower must be an integer, got ${t}.`), w(t <= o, () => `bandPart(): numLower (${t}) must not be greater than the number of rows (${o}).`), a =
|
|
7473
|
+
typeof t == "number" ? (w(t % 1 === 0, () => `bandPart(): numLower must be an integer, got ${t}.`), w(t <= o, () => `bandPart(): numLower (${t}) must not be greater than the number of rows (${o}).`), a = $(t < 0 ? o : t, "numLower", "bandPart")) : (w(t.dtype === "int32", () => "bandPart(): numLower's dtype must be an int32."), a = on(Ra(t, 0), o, mr(t, o))), typeof e == "number" ? (w(e % 1 === 0, () => `bandPart(): numUpper must be an integer, got ${e}.`), w(e <= i, () => `bandPart(): numUpper (${e}) must not be greater than the number of columns (${i}).`), l = $(e < 0 ? i : e, "numUpper", "bandPart")) : (w(e.dtype === "int32", () => "bandPart(): numUpper's dtype must be an int32."), l = on(Ra(e, 0), i, mr(e, i)));
|
|
7474
7474
|
const u = L(gr(0, o, 1, "int32"), [-1, 1]), c = gr(0, i, 1, "int32"), h = Z(u, c), f = qr(Nu(h, a), Lg(h, qn(l))), d = Fn([o, i], s.dtype);
|
|
7475
7475
|
return L(br(Uu(L(s, [-1, o, i])).map((p) => on(f, p, d))), r);
|
|
7476
7476
|
}
|
|
@@ -8564,7 +8564,7 @@ class ly {
|
|
|
8564
8564
|
* limitations under the License.
|
|
8565
8565
|
* =============================================================================
|
|
8566
8566
|
*/
|
|
8567
|
-
const
|
|
8567
|
+
const In = ly;
|
|
8568
8568
|
/**
|
|
8569
8569
|
* @license
|
|
8570
8570
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
@@ -8634,10 +8634,10 @@ function Ss(n, t) {
|
|
|
8634
8634
|
* limitations under the License.
|
|
8635
8635
|
* =============================================================================
|
|
8636
8636
|
*/
|
|
8637
|
-
var
|
|
8637
|
+
var $e;
|
|
8638
8638
|
(function(n) {
|
|
8639
8639
|
n[n.FIRST_DIM_SIZE = 0] = "FIRST_DIM_SIZE", n[n.VALUE_ROWIDS = 1] = "VALUE_ROWIDS", n[n.ROW_LENGTHS = 2] = "ROW_LENGTHS", n[n.ROW_SPLITS = 3] = "ROW_SPLITS", n[n.ROW_LIMITS = 4] = "ROW_LIMITS", n[n.ROW_STARTS = 5] = "ROW_STARTS";
|
|
8640
|
-
})(
|
|
8640
|
+
})($e || ($e = {}));
|
|
8641
8641
|
function fy(n, t, e) {
|
|
8642
8642
|
let s = new Array();
|
|
8643
8643
|
if (e == null && t == null)
|
|
@@ -8664,12 +8664,12 @@ function fy(n, t, e) {
|
|
|
8664
8664
|
}
|
|
8665
8665
|
function dy(n) {
|
|
8666
8666
|
const t = {
|
|
8667
|
-
FIRST_DIM_SIZE:
|
|
8668
|
-
VALUE_ROWIDS:
|
|
8669
|
-
ROW_LENGTHS:
|
|
8670
|
-
ROW_SPLITS:
|
|
8671
|
-
ROW_LIMITS:
|
|
8672
|
-
ROW_STARTS:
|
|
8667
|
+
FIRST_DIM_SIZE: $e.FIRST_DIM_SIZE,
|
|
8668
|
+
VALUE_ROWIDS: $e.VALUE_ROWIDS,
|
|
8669
|
+
ROW_LENGTHS: $e.ROW_LENGTHS,
|
|
8670
|
+
ROW_SPLITS: $e.ROW_SPLITS,
|
|
8671
|
+
ROW_LIMITS: $e.ROW_LIMITS,
|
|
8672
|
+
ROW_STARTS: $e.ROW_STARTS
|
|
8673
8673
|
}, e = [];
|
|
8674
8674
|
for (const s of n)
|
|
8675
8675
|
if (s in t)
|
|
@@ -8679,7 +8679,7 @@ function dy(n) {
|
|
|
8679
8679
|
return e;
|
|
8680
8680
|
}
|
|
8681
8681
|
function py(n) {
|
|
8682
|
-
return n.length === 0 ? 0 : n[0] ===
|
|
8682
|
+
return n.length === 0 ? 0 : n[0] === $e.FIRST_DIM_SIZE ? n.length - 1 : n.length;
|
|
8683
8683
|
}
|
|
8684
8684
|
function my(n, t) {
|
|
8685
8685
|
if (n == null || t == null)
|
|
@@ -8726,7 +8726,7 @@ const gy = 1.7580993408473768, by = 1.0507009873554805;
|
|
|
8726
8726
|
* limitations under the License.
|
|
8727
8727
|
* =============================================================================
|
|
8728
8728
|
*/
|
|
8729
|
-
const yy = 0.3275911, wy = 0.254829592, xy = -0.284496736, Sy = 1.421413741, vy = -1.453152027,
|
|
8729
|
+
const yy = 0.3275911, wy = 0.254829592, xy = -0.284496736, Sy = 1.421413741, vy = -1.453152027, Iy = 1.061405429;
|
|
8730
8730
|
/**
|
|
8731
8731
|
* @license
|
|
8732
8732
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
@@ -8767,7 +8767,7 @@ function Ma(n, t) {
|
|
|
8767
8767
|
* limitations under the License.
|
|
8768
8768
|
* =============================================================================
|
|
8769
8769
|
*/
|
|
8770
|
-
function
|
|
8770
|
+
function $y(n) {
|
|
8771
8771
|
return `Received SparseTensor with denseShape[0] = 0 but
|
|
8772
8772
|
indices.shape[0] = ${n}`;
|
|
8773
8773
|
}
|
|
@@ -8911,9 +8911,9 @@ class qe extends Error {
|
|
|
8911
8911
|
super(t), Object.setPrototypeOf(this, qe.prototype);
|
|
8912
8912
|
}
|
|
8913
8913
|
}
|
|
8914
|
-
class
|
|
8914
|
+
class I extends Error {
|
|
8915
8915
|
constructor(t) {
|
|
8916
|
-
super(t), Object.setPrototypeOf(this,
|
|
8916
|
+
super(t), Object.setPrototypeOf(this, I.prototype);
|
|
8917
8917
|
}
|
|
8918
8918
|
}
|
|
8919
8919
|
class J extends Error {
|
|
@@ -8997,19 +8997,19 @@ function Ns(n, t = {}, e = {}, s = "object", r = !1) {
|
|
|
8997
8997
|
else if (o in Jt)
|
|
8998
8998
|
i = Jt[o];
|
|
8999
8999
|
else if (i = t[o], i == null)
|
|
9000
|
-
throw new
|
|
9000
|
+
throw new I(`Unknown ${s}: ${n}. This may be due to one of the following reasons:
|
|
9001
9001
|
1. The ${s} is defined in Python, in which case it needs to be ported to TensorFlow.js or your JavaScript code.
|
|
9002
9002
|
2. The custom ${s} is defined in JavaScript, but is not registered properly with tf.serialization.registerClass().`);
|
|
9003
9003
|
return i;
|
|
9004
9004
|
} else {
|
|
9005
9005
|
const o = n;
|
|
9006
9006
|
if (o.className == null || o.config == null)
|
|
9007
|
-
throw new
|
|
9007
|
+
throw new I(`${s}: Improper config format: ${JSON.stringify(o)}.
|
|
9008
9008
|
'className' and 'config' must set.`);
|
|
9009
9009
|
const i = o.className;
|
|
9010
9010
|
let a, l;
|
|
9011
9011
|
if (i in e ? [a, l] = e[i] : i in Jt ? [a, l] = Jt.className : i in t && ([a, l] = t[i]), a == null)
|
|
9012
|
-
throw new
|
|
9012
|
+
throw new I(`Unknown ${s}: ${i}. This may be due to one of the following reasons:
|
|
9013
9013
|
1. The ${s} is defined in Python, in which case it needs to be ported to TensorFlow.js or your JavaScript code.
|
|
9014
9014
|
2. The custom ${s} is defined in JavaScript, but is not registered properly with tf.serialization.registerClass().`);
|
|
9015
9015
|
if (l != null) {
|
|
@@ -9051,7 +9051,7 @@ function an(n) {
|
|
|
9051
9051
|
}
|
|
9052
9052
|
function Uy(n) {
|
|
9053
9053
|
if (n == null)
|
|
9054
|
-
throw new
|
|
9054
|
+
throw new I(`Invalid value in obj: ${JSON.stringify(n)}`);
|
|
9055
9055
|
for (const t in n)
|
|
9056
9056
|
if (n.hasOwnProperty(t))
|
|
9057
9057
|
return !1;
|
|
@@ -9059,7 +9059,7 @@ function Uy(n) {
|
|
|
9059
9059
|
}
|
|
9060
9060
|
function Yn(n, t, e) {
|
|
9061
9061
|
if (e != null && n.indexOf(e) < 0)
|
|
9062
|
-
throw new
|
|
9062
|
+
throw new I(`${e} is not a valid ${t}. Valid values are ${n} or null/undefined.`);
|
|
9063
9063
|
}
|
|
9064
9064
|
function Di(n, t, e = 0, s = 1 / 0) {
|
|
9065
9065
|
return Ae(e >= 0), Ae(s >= e), Array.isArray(n) && n.length >= e && n.length <= s && n.every((r) => typeof r === t);
|
|
@@ -9089,7 +9089,7 @@ function Gy(n) {
|
|
|
9089
9089
|
* https://opensource.org/licenses/MIT.
|
|
9090
9090
|
* =============================================================================
|
|
9091
9091
|
*/
|
|
9092
|
-
const
|
|
9092
|
+
const $n = /* @__PURE__ */ new Map();
|
|
9093
9093
|
function gt(n) {
|
|
9094
9094
|
Yn(Oy, "DataFormat", n);
|
|
9095
9095
|
}
|
|
@@ -9123,11 +9123,11 @@ function Ju(n) {
|
|
|
9123
9123
|
function Zu(n) {
|
|
9124
9124
|
if (!Qu(n))
|
|
9125
9125
|
throw new Error("Not a valid tensor name: '" + n + "'");
|
|
9126
|
-
|
|
9127
|
-
const t =
|
|
9128
|
-
if (
|
|
9126
|
+
$n.has(n) || $n.set(n, 0);
|
|
9127
|
+
const t = $n.get(n);
|
|
9128
|
+
if ($n.set(n, $n.get(n) + 1), t > 0) {
|
|
9129
9129
|
const e = `${n}_${t}`;
|
|
9130
|
-
return
|
|
9130
|
+
return $n.set(e, 1), e;
|
|
9131
9131
|
} else
|
|
9132
9132
|
return n;
|
|
9133
9133
|
}
|
|
@@ -9166,7 +9166,7 @@ function tc(n) {
|
|
|
9166
9166
|
}
|
|
9167
9167
|
function wr(n, t) {
|
|
9168
9168
|
if (t < n)
|
|
9169
|
-
throw new
|
|
9169
|
+
throw new I(`end (${t}) < begin (${n}) is forbidden.`);
|
|
9170
9170
|
const e = [];
|
|
9171
9171
|
for (let s = n; s < t; ++s)
|
|
9172
9172
|
e.push(s);
|
|
@@ -9237,7 +9237,7 @@ function ln(n, t, e) {
|
|
|
9237
9237
|
n.shape[5]
|
|
9238
9238
|
]);
|
|
9239
9239
|
default:
|
|
9240
|
-
throw new
|
|
9240
|
+
throw new I(`sliceAlongFirstAxis() received an unsupported tensor rank: ${n.rank}`);
|
|
9241
9241
|
}
|
|
9242
9242
|
});
|
|
9243
9243
|
}
|
|
@@ -9253,7 +9253,7 @@ function po(n, t, e) {
|
|
|
9253
9253
|
case 4:
|
|
9254
9254
|
return ws(n, [0, 0, 0, t], [n.shape[0], n.shape[1], n.shape[2], e]);
|
|
9255
9255
|
default:
|
|
9256
|
-
throw new
|
|
9256
|
+
throw new I(`sliceAlongLastAxis() received an unsupported tensor rank: ${n.rank}`);
|
|
9257
9257
|
}
|
|
9258
9258
|
});
|
|
9259
9259
|
}
|
|
@@ -9269,7 +9269,7 @@ function Gs(n, t, e, s) {
|
|
|
9269
9269
|
case 2:
|
|
9270
9270
|
return po(n, t, e);
|
|
9271
9271
|
default:
|
|
9272
|
-
throw new
|
|
9272
|
+
throw new I(`The axis is not within the rank of the tensor ${s}`);
|
|
9273
9273
|
}
|
|
9274
9274
|
case 3:
|
|
9275
9275
|
switch (s) {
|
|
@@ -9280,7 +9280,7 @@ function Gs(n, t, e, s) {
|
|
|
9280
9280
|
case 3:
|
|
9281
9281
|
return po(n, t, e);
|
|
9282
9282
|
default:
|
|
9283
|
-
throw new
|
|
9283
|
+
throw new I(`The axis is not within the rank of the tensor ${s}`);
|
|
9284
9284
|
}
|
|
9285
9285
|
case 4:
|
|
9286
9286
|
switch (s) {
|
|
@@ -9293,10 +9293,10 @@ function Gs(n, t, e, s) {
|
|
|
9293
9293
|
case 4:
|
|
9294
9294
|
return po(n, t, e);
|
|
9295
9295
|
default:
|
|
9296
|
-
throw new
|
|
9296
|
+
throw new I(`The axis is not within the rank of the tensor ${s}`);
|
|
9297
9297
|
}
|
|
9298
9298
|
default:
|
|
9299
|
-
throw new
|
|
9299
|
+
throw new I(`sliceAlongLastAxis() received an unsupported tensor rank: ${n.rank}`);
|
|
9300
9300
|
}
|
|
9301
9301
|
});
|
|
9302
9302
|
}
|
|
@@ -9316,7 +9316,7 @@ function Ds(n) {
|
|
|
9316
9316
|
function Jy(n, t, e) {
|
|
9317
9317
|
const s = t.shape;
|
|
9318
9318
|
if (t.rank !== 1 && t.rank !== n)
|
|
9319
|
-
throw new
|
|
9319
|
+
throw new I(`Unexpected bias dimensions: ${t.rank}; expected it to be 1 or ${n}`);
|
|
9320
9320
|
if (n === 5) {
|
|
9321
9321
|
if (e === "channelsFirst")
|
|
9322
9322
|
return s.length === 1 ? L(t, [1, s[0], 1, 1, 1]) : L(t, [1, s[3], s[0], s[1], s[2]]);
|
|
@@ -9334,7 +9334,7 @@ function Jy(n, t, e) {
|
|
|
9334
9334
|
return s.length === 1 ? L(t, [1, 1, s[0]]) : L(t, [1].concat(s));
|
|
9335
9335
|
} else if (n < 3)
|
|
9336
9336
|
return t;
|
|
9337
|
-
throw new
|
|
9337
|
+
throw new I(`Unsupported input rank by biasAdd: ${t.rank}`);
|
|
9338
9338
|
}
|
|
9339
9339
|
function Ps(n, t, e) {
|
|
9340
9340
|
return _(() => (e == null && (e = Xn()), gt(e), M(n, Jy(n.rank, t, e))));
|
|
@@ -9659,7 +9659,7 @@ function Sr(n) {
|
|
|
9659
9659
|
* =============================================================================
|
|
9660
9660
|
*/
|
|
9661
9661
|
let sw = 0;
|
|
9662
|
-
function
|
|
9662
|
+
function Ic() {
|
|
9663
9663
|
return sw++;
|
|
9664
9664
|
}
|
|
9665
9665
|
const Vs = {};
|
|
@@ -9699,13 +9699,13 @@ class Ne extends Kn {
|
|
|
9699
9699
|
return {};
|
|
9700
9700
|
}
|
|
9701
9701
|
}
|
|
9702
|
-
class
|
|
9702
|
+
class $c extends Ne {
|
|
9703
9703
|
apply(t, e) {
|
|
9704
9704
|
return Fn(t, e);
|
|
9705
9705
|
}
|
|
9706
9706
|
}
|
|
9707
|
-
|
|
9708
|
-
O(
|
|
9707
|
+
$c.className = "Zeros";
|
|
9708
|
+
O($c);
|
|
9709
9709
|
class Ac extends Ne {
|
|
9710
9710
|
apply(t, e) {
|
|
9711
9711
|
return wi(t, e);
|
|
@@ -9716,9 +9716,9 @@ O(Ac);
|
|
|
9716
9716
|
class Ec extends Ne {
|
|
9717
9717
|
constructor(t) {
|
|
9718
9718
|
if (super(), typeof t != "object")
|
|
9719
|
-
throw new
|
|
9719
|
+
throw new I(`Expected argument of type ConstantConfig but got ${t}`);
|
|
9720
9720
|
if (t.value === void 0)
|
|
9721
|
-
throw new
|
|
9721
|
+
throw new I(`config must have value set but got ${t}`);
|
|
9722
9722
|
this.value = t.value;
|
|
9723
9723
|
}
|
|
9724
9724
|
apply(t, e) {
|
|
@@ -9782,7 +9782,7 @@ class Tc extends Ne {
|
|
|
9782
9782
|
apply(t, e) {
|
|
9783
9783
|
return _(() => {
|
|
9784
9784
|
if (t.length !== 2 || t[0] !== t[1])
|
|
9785
|
-
throw new
|
|
9785
|
+
throw new I("Identity matrix initializer can only be used for 2D square matrices.");
|
|
9786
9786
|
return N(this.gain, Tu(t[0]));
|
|
9787
9787
|
});
|
|
9788
9788
|
}
|
|
@@ -9817,7 +9817,7 @@ class Gt extends Ne {
|
|
|
9817
9817
|
*/
|
|
9818
9818
|
constructor(t) {
|
|
9819
9819
|
if (super(), t.scale < 0)
|
|
9820
|
-
throw new
|
|
9820
|
+
throw new I(`scale must be a positive float. Got: ${t.scale}`);
|
|
9821
9821
|
this.scale = t.scale == null ? 1 : t.scale, this.mode = t.mode == null ? "fanIn" : t.mode, iw(this.mode), this.distribution = t.distribution == null ? "normal" : t.distribution, aw(this.distribution), this.seed = t.seed;
|
|
9822
9822
|
}
|
|
9823
9823
|
apply(t, e) {
|
|
@@ -10029,14 +10029,14 @@ function vs(n) {
|
|
|
10029
10029
|
* https://opensource.org/licenses/MIT.
|
|
10030
10030
|
* =============================================================================
|
|
10031
10031
|
*/
|
|
10032
|
-
function
|
|
10032
|
+
function Ir(n) {
|
|
10033
10033
|
return n.length === 0 ? [] : Array.isArray(n[0]) ? n : [n];
|
|
10034
10034
|
}
|
|
10035
10035
|
function Vt(n) {
|
|
10036
10036
|
let t;
|
|
10037
10037
|
if (Array.isArray(n)) {
|
|
10038
10038
|
if (n.length !== 1)
|
|
10039
|
-
throw new
|
|
10039
|
+
throw new I(`Expected Tensor length to be 1; got ${n.length}`);
|
|
10040
10040
|
t = n[0];
|
|
10041
10041
|
} else
|
|
10042
10042
|
t = n;
|
|
@@ -10046,7 +10046,7 @@ function ge(n) {
|
|
|
10046
10046
|
if (Array.isArray(n) && Array.isArray(n[0])) {
|
|
10047
10047
|
if (n.length === 1)
|
|
10048
10048
|
return n = n, n[0];
|
|
10049
|
-
throw new
|
|
10049
|
+
throw new I(`Expected exactly 1 Shape; got ${n.length}`);
|
|
10050
10050
|
} else
|
|
10051
10051
|
return n;
|
|
10052
10052
|
}
|
|
@@ -10059,7 +10059,7 @@ function ge(n) {
|
|
|
10059
10059
|
* https://opensource.org/licenses/MIT.
|
|
10060
10060
|
* =============================================================================
|
|
10061
10061
|
*/
|
|
10062
|
-
function
|
|
10062
|
+
function $r(n) {
|
|
10063
10063
|
let t = 0;
|
|
10064
10064
|
for (const e of n)
|
|
10065
10065
|
e.shape.length === 0 ? t += 1 : t += e.shape.reduce((s, r) => s * r);
|
|
@@ -10091,7 +10091,7 @@ class uw {
|
|
|
10091
10091
|
* @throws ValueError if `name` is `null` or `undefined`.
|
|
10092
10092
|
*/
|
|
10093
10093
|
constructor(t, e = "float32", s = qa, r = !0, o = null) {
|
|
10094
|
-
this.dtype = e ?? "float32", this.shape = t.shape, this.id =
|
|
10094
|
+
this.dtype = e ?? "float32", this.shape = t.shape, this.id = Ic(), s = s ?? qa, this.originalName = Ju(s), this.name = Zu(this.originalName), this.trainable_ = r, this.constraint = o, this.val = Z0(t, this.trainable_, this.name, this.dtype);
|
|
10095
10095
|
}
|
|
10096
10096
|
/**
|
|
10097
10097
|
* Get a snapshot of the Variable's value.
|
|
@@ -10171,7 +10171,7 @@ class bn {
|
|
|
10171
10171
|
* returned by apply().
|
|
10172
10172
|
*/
|
|
10173
10173
|
constructor(t, e, s, r, o, i, a) {
|
|
10174
|
-
this.dtype = t, this.shape = e, this.sourceLayer = s, this.inputs = r, this.callArgs = o, this.outputTensorIndex = a, this.id =
|
|
10174
|
+
this.dtype = t, this.shape = e, this.sourceLayer = s, this.inputs = r, this.callArgs = o, this.outputTensorIndex = a, this.id = Ic(), i != null && (this.originalName = Ju(i), this.name = Zu(this.originalName)), this.rank = e.length;
|
|
10175
10175
|
}
|
|
10176
10176
|
}
|
|
10177
10177
|
let hw = 0;
|
|
@@ -10240,7 +10240,7 @@ class ye extends Kn {
|
|
|
10240
10240
|
if (this.inboundNodes.length === 0)
|
|
10241
10241
|
throw new qe(`The layer has never been called and thus has no defined ${e}.`);
|
|
10242
10242
|
if (this.inboundNodes.length <= t)
|
|
10243
|
-
throw new
|
|
10243
|
+
throw new I(`Asked to get ${e} at node ${t}, but the layer has only ${this.inboundNodes.length} inbound nodes.`);
|
|
10244
10244
|
return this.inboundNodes[t];
|
|
10245
10245
|
}
|
|
10246
10246
|
/**
|
|
@@ -10381,33 +10381,33 @@ class ye extends Kn {
|
|
|
10381
10381
|
return;
|
|
10382
10382
|
const s = st(this.inputSpec);
|
|
10383
10383
|
if (e.length !== s.length)
|
|
10384
|
-
throw new
|
|
10384
|
+
throw new I(`Layer ${this.name} expects ${s.length} inputs, but it received ${e.length} input tensors. Input received: ${t}`);
|
|
10385
10385
|
for (let r = 0; r < e.length; r++) {
|
|
10386
10386
|
const o = e[r], i = s[r];
|
|
10387
10387
|
if (i == null)
|
|
10388
10388
|
continue;
|
|
10389
10389
|
const a = o.rank;
|
|
10390
10390
|
if (i.ndim != null && a !== i.ndim)
|
|
10391
|
-
throw new
|
|
10391
|
+
throw new I(`Input ${r} is incompatible with layer ${this.name}: expected ndim=${i.ndim}, found ndim=${a}`);
|
|
10392
10392
|
if (i.maxNDim != null && a > i.maxNDim)
|
|
10393
|
-
throw new
|
|
10393
|
+
throw new I(`Input ${r} is incompatible with layer ${this.name}: expected max_ndim=${i.maxNDim}, found ndim=${a}`);
|
|
10394
10394
|
if (i.minNDim != null && a < i.minNDim)
|
|
10395
|
-
throw new
|
|
10395
|
+
throw new I(`Input ${r} is incompatible with layer ${this.name}: expected min_ndim=${i.minNDim}, found ndim=${a}.`);
|
|
10396
10396
|
if (i.dtype != null && o.dtype !== i.dtype)
|
|
10397
|
-
throw new
|
|
10397
|
+
throw new I(`Input ${r} is incompatible with layer ${this.name} : expected dtype=${i.dtype}, found dtype=${o.dtype}.`);
|
|
10398
10398
|
if (i.axes) {
|
|
10399
10399
|
const l = o.shape;
|
|
10400
10400
|
for (const u in i.axes) {
|
|
10401
10401
|
const c = Number(u), h = i.axes[u], f = c >= 0 ? l[c] : l[l.length + c];
|
|
10402
10402
|
if (h != null && [h, null].indexOf(f) === -1)
|
|
10403
|
-
throw new
|
|
10403
|
+
throw new I(`Input ${r} is incompatible with layer ${this.name}: expected axis ${c} of input shape to have value ${h} but got shape ${l}.`);
|
|
10404
10404
|
}
|
|
10405
10405
|
}
|
|
10406
10406
|
if (i.shape != null)
|
|
10407
10407
|
for (let l = 0; l < i.shape.length; ++l) {
|
|
10408
10408
|
const u = i.shape[l], c = o.shape[l];
|
|
10409
10409
|
if (u != null && c != null && u !== c)
|
|
10410
|
-
throw new
|
|
10410
|
+
throw new I(`Input ${r} is incompatible with layer ${this.name}: expected shape=${i.shape}, found shape=${o.shape}.`);
|
|
10411
10411
|
}
|
|
10412
10412
|
}
|
|
10413
10413
|
}
|
|
@@ -10513,7 +10513,7 @@ class ye extends Kn {
|
|
|
10513
10513
|
e = e || {}, this.assertNotDisposed();
|
|
10514
10514
|
const s = st(t), r = mw(t), o = gw(t);
|
|
10515
10515
|
if (r === o)
|
|
10516
|
-
throw new
|
|
10516
|
+
throw new I("Arguments to apply() must be all SymbolicTensors or all Tensors");
|
|
10517
10517
|
return or(this.name, () => {
|
|
10518
10518
|
if (!this.built) {
|
|
10519
10519
|
this.assertInputCompatibility(t);
|
|
@@ -10598,7 +10598,7 @@ class ye extends Kn {
|
|
|
10598
10598
|
countParams() {
|
|
10599
10599
|
if (!this.built)
|
|
10600
10600
|
throw new qe(`You tried to call countParams() on ${this.name}, but the layer is not built yet. Build it first by calling build(batchInputShape).`);
|
|
10601
|
-
return
|
|
10601
|
+
return $r(this.weights);
|
|
10602
10602
|
}
|
|
10603
10603
|
/**
|
|
10604
10604
|
* Creates the layer weights.
|
|
@@ -10641,14 +10641,14 @@ class ye extends Kn {
|
|
|
10641
10641
|
_(() => {
|
|
10642
10642
|
const e = this.weights;
|
|
10643
10643
|
if (e.length !== t.length)
|
|
10644
|
-
throw new
|
|
10644
|
+
throw new I(`You called setWeights(weights) on layer "${this.name}" with a weight list of length ${t.length}, but the layer was expecting ${e.length} weights. Provided weights: ${t}...`);
|
|
10645
10645
|
if (e.length === 0)
|
|
10646
10646
|
return;
|
|
10647
10647
|
const s = [], r = ja(e);
|
|
10648
10648
|
for (let o = 0; o < r.length; ++o) {
|
|
10649
10649
|
const i = r[o], a = e[o], l = t[o];
|
|
10650
10650
|
if (!oe(i.shape, l.shape))
|
|
10651
|
-
throw new
|
|
10651
|
+
throw new I(`Layer weight shape ${i.shape} not compatible with provided weight shape ${l.shape}`);
|
|
10652
10652
|
s.push([a, l]);
|
|
10653
10653
|
}
|
|
10654
10654
|
Dc(s);
|
|
@@ -10671,7 +10671,7 @@ class ye extends Kn {
|
|
|
10671
10671
|
*/
|
|
10672
10672
|
addWeight(t, e, s, r, o, i, a, l) {
|
|
10673
10673
|
if (this._addedWeightNames.indexOf(t) !== -1)
|
|
10674
|
-
throw new
|
|
10674
|
+
throw new I(`Duplicate weight name ${t} for layer ${this.name}`);
|
|
10675
10675
|
this._addedWeightNames.push(t), s == null && (s = "float32"), this.fastWeightInitDuringBuild && (r = l != null ? l() : vs("zeros"));
|
|
10676
10676
|
const u = r.apply(e, s), c = new uw(u, s, t, i, a);
|
|
10677
10677
|
return u.dispose(), o != null && this.addLoss(() => o.apply(c.read())), i == null && (i = !0), i ? this._trainableWeights.push(c) : this._nonTrainableWeights.push(c), c;
|
|
@@ -10760,7 +10760,7 @@ class ye extends Kn {
|
|
|
10760
10760
|
*/
|
|
10761
10761
|
addInboundNode(t, e, s, r, o, i, a = null) {
|
|
10762
10762
|
const l = st(t);
|
|
10763
|
-
e = st(e), s = st(s), r = st(r), o =
|
|
10763
|
+
e = st(e), s = st(s), r = st(r), o = Ir(o), i = Ir(i);
|
|
10764
10764
|
const u = [], c = [], h = [];
|
|
10765
10765
|
for (const f of l)
|
|
10766
10766
|
u.push(f.sourceLayer), c.push(f.nodeIndex), h.push(f.tensorIndex);
|
|
@@ -10926,13 +10926,13 @@ O(Rc);
|
|
|
10926
10926
|
const Ha = {
|
|
10927
10927
|
l1l2: "L1L2"
|
|
10928
10928
|
};
|
|
10929
|
-
function
|
|
10929
|
+
function Is(n) {
|
|
10930
10930
|
return Ni(n);
|
|
10931
10931
|
}
|
|
10932
10932
|
function Ka(n, t = {}) {
|
|
10933
10933
|
return Ns(n, te.getMap().classNameMap, t, "regularizer");
|
|
10934
10934
|
}
|
|
10935
|
-
function
|
|
10935
|
+
function $s(n) {
|
|
10936
10936
|
if (n == null)
|
|
10937
10937
|
return null;
|
|
10938
10938
|
if (typeof n == "string") {
|
|
@@ -10953,11 +10953,11 @@ function go(n, t, e) {
|
|
|
10953
10953
|
if (typeof n == "number")
|
|
10954
10954
|
return yr(n, t);
|
|
10955
10955
|
if (n.length !== t)
|
|
10956
|
-
throw new
|
|
10956
|
+
throw new I(`The ${e} argument must be an integer or tuple of ${t} integers. Received: ${n.length} elements.`);
|
|
10957
10957
|
for (let s = 0; s < t; ++s) {
|
|
10958
10958
|
const r = n[s];
|
|
10959
10959
|
if (!Hy(r))
|
|
10960
|
-
throw new
|
|
10960
|
+
throw new I(`The ${e} argument must be an integer or tuple of ${t} integers. Received: ${JSON.stringify(n)} including a non-integer number ${r}`);
|
|
10961
10961
|
}
|
|
10962
10962
|
return n;
|
|
10963
10963
|
}
|
|
@@ -10976,7 +10976,7 @@ function Ee(n, t, e, s) {
|
|
|
10976
10976
|
else if (s === "same")
|
|
10977
10977
|
n = n * t;
|
|
10978
10978
|
else
|
|
10979
|
-
throw new
|
|
10979
|
+
throw new I(`Unsupport padding mode: ${s}.`);
|
|
10980
10980
|
return n;
|
|
10981
10981
|
}
|
|
10982
10982
|
/**
|
|
@@ -10997,11 +10997,11 @@ function Oc(n, t) {
|
|
|
10997
10997
|
function yw(n, t, e, s = 1, r = "valid", o, i = 1) {
|
|
10998
10998
|
return _(() => {
|
|
10999
10999
|
if (o == null && (o = Xn()), gt(o), n.shape.length !== 3)
|
|
11000
|
-
throw new
|
|
11000
|
+
throw new I(`The input of a conv1dWithBias operation should be 3, but is ${n.shape.length} instead.`);
|
|
11001
11001
|
if (t.shape.length !== 3)
|
|
11002
|
-
throw new
|
|
11002
|
+
throw new I(`The kernel for a conv1dWithBias operation should be 3, but is ${t.shape.length} instead`);
|
|
11003
11003
|
if (e != null && e.shape.length !== 1)
|
|
11004
|
-
throw new
|
|
11004
|
+
throw new I(`The bias for a conv1dWithBias operation should be 1, but is ${e.shape.length} instead`);
|
|
11005
11005
|
if (o === "channelsFirst" && (n = pt(n, [0, 2, 1])), r === "causal")
|
|
11006
11006
|
throw new J("The support for CAUSAL padding mode in conv1dWithBias is not implemented yet.");
|
|
11007
11007
|
let a = Km(n, t, s, r === "same" ? "same" : "valid", "NWC", i);
|
|
@@ -11011,9 +11011,9 @@ function yw(n, t, e, s = 1, r = "valid", o, i = 1) {
|
|
|
11011
11011
|
function Ya(n, t, e, s = [1, 1], r = "valid", o, i, a = null) {
|
|
11012
11012
|
return _(() => {
|
|
11013
11013
|
if (o == null && (o = Xn()), gt(o), n.rank !== 3 && n.rank !== 4)
|
|
11014
|
-
throw new
|
|
11014
|
+
throw new I(`conv2dWithBiasActivation expects input to be of rank 3 or 4, but received ${n.rank}.`);
|
|
11015
11015
|
if (t.rank !== 3 && t.rank !== 4)
|
|
11016
|
-
throw new
|
|
11016
|
+
throw new I(`conv2dWithBiasActivation expects kernel to be of rank 3 or 4, but received ${n.rank}.`);
|
|
11017
11017
|
let l = Lc(n, o);
|
|
11018
11018
|
if (r === "causal")
|
|
11019
11019
|
throw new J("The support for CAUSAL padding mode in conv1dWithBias is not implemented yet.");
|
|
@@ -11032,9 +11032,9 @@ function Ya(n, t, e, s = [1, 1], r = "valid", o, i, a = null) {
|
|
|
11032
11032
|
function ww(n, t, e, s = [1, 1, 1], r = "valid", o, i) {
|
|
11033
11033
|
return _(() => {
|
|
11034
11034
|
if (o == null && (o = Xn()), gt(o), n.rank !== 4 && n.rank !== 5)
|
|
11035
|
-
throw new
|
|
11035
|
+
throw new I(`conv3dWithBias expects input to be of rank 4 or 5, but received ${n.rank}.`);
|
|
11036
11036
|
if (t.rank !== 4 && t.rank !== 5)
|
|
11037
|
-
throw new
|
|
11037
|
+
throw new I(`conv3dWithBias expects kernel to be of rank 4 or 5, but received ${n.rank}.`);
|
|
11038
11038
|
let a = Oc(n, o);
|
|
11039
11039
|
if (r === "causal")
|
|
11040
11040
|
throw new J("The support for CAUSAL padding mode in conv3dWithBias is not implemented yet.");
|
|
@@ -11045,23 +11045,23 @@ class Gi extends ye {
|
|
|
11045
11045
|
constructor(t, e) {
|
|
11046
11046
|
if (super(e), this.bias = null, this.DEFAULT_KERNEL_INITIALIZER = "glorotNormal", this.DEFAULT_BIAS_INITIALIZER = "zeros", Gi.verifyArgs(e), this.rank = t, Oe(this.rank, "rank"), this.rank !== 1 && this.rank !== 2 && this.rank !== 3)
|
|
11047
11047
|
throw new J(`Convolution layer for rank other than 1, 2, or 3 (${this.rank}) is not implemented yet.`);
|
|
11048
|
-
if (this.kernelSize = go(e.kernelSize, t, "kernelSize"), this.strides = go(e.strides == null ? 1 : e.strides, t, "strides"), this.padding = e.padding == null ? "valid" : e.padding, ie(this.padding), this.dataFormat = e.dataFormat == null ? "channelsLast" : e.dataFormat, gt(this.dataFormat), this.activation = nw(e.activation), this.useBias = e.useBias == null ? !0 : e.useBias, this.biasInitializer = vs(e.biasInitializer || this.DEFAULT_BIAS_INITIALIZER), this.biasConstraint = Sr(e.biasConstraint), this.biasRegularizer =
|
|
11049
|
-
throw new
|
|
11048
|
+
if (this.kernelSize = go(e.kernelSize, t, "kernelSize"), this.strides = go(e.strides == null ? 1 : e.strides, t, "strides"), this.padding = e.padding == null ? "valid" : e.padding, ie(this.padding), this.dataFormat = e.dataFormat == null ? "channelsLast" : e.dataFormat, gt(this.dataFormat), this.activation = nw(e.activation), this.useBias = e.useBias == null ? !0 : e.useBias, this.biasInitializer = vs(e.biasInitializer || this.DEFAULT_BIAS_INITIALIZER), this.biasConstraint = Sr(e.biasConstraint), this.biasRegularizer = $s(e.biasRegularizer), this.activityRegularizer = $s(e.activityRegularizer), this.dilationRate = go(e.dilationRate == null ? 1 : e.dilationRate, t, "dilationRate"), this.rank === 1 && Array.isArray(this.dilationRate) && this.dilationRate.length !== 1)
|
|
11049
|
+
throw new I(`dilationRate must be a number or an array of a single number for 1D convolution, but received ${JSON.stringify(this.dilationRate)}`);
|
|
11050
11050
|
if (this.rank === 2) {
|
|
11051
11051
|
if (typeof this.dilationRate == "number")
|
|
11052
11052
|
this.dilationRate = [this.dilationRate, this.dilationRate];
|
|
11053
11053
|
else if (this.dilationRate.length !== 2)
|
|
11054
|
-
throw new
|
|
11054
|
+
throw new I(`dilationRate must be a number or array of two numbers for 2D convolution, but received ${JSON.stringify(this.dilationRate)}`);
|
|
11055
11055
|
} else if (this.rank === 3) {
|
|
11056
11056
|
if (typeof this.dilationRate == "number")
|
|
11057
11057
|
this.dilationRate = [this.dilationRate, this.dilationRate, this.dilationRate];
|
|
11058
11058
|
else if (this.dilationRate.length !== 3)
|
|
11059
|
-
throw new
|
|
11059
|
+
throw new I(`dilationRate must be a number or array of three numbers for 3D convolution, but received ${JSON.stringify(this.dilationRate)}`);
|
|
11060
11060
|
}
|
|
11061
11061
|
}
|
|
11062
11062
|
static verifyArgs(t) {
|
|
11063
11063
|
if (Ae("kernelSize" in t, "required key 'kernelSize' not in config"), typeof t.kernelSize != "number" && !Di(t.kernelSize, "number", 1, 3))
|
|
11064
|
-
throw new
|
|
11064
|
+
throw new I(`BaseConv expects config.kernelSize to be number or number[] with length 1, 2, or 3, but received ${JSON.stringify(t.kernelSize)}.`);
|
|
11065
11065
|
}
|
|
11066
11066
|
getConfig() {
|
|
11067
11067
|
const t = {
|
|
@@ -11073,8 +11073,8 @@ class Gi extends ye {
|
|
|
11073
11073
|
activation: ew(this.activation),
|
|
11074
11074
|
useBias: this.useBias,
|
|
11075
11075
|
biasInitializer: vr(this.biasInitializer),
|
|
11076
|
-
biasRegularizer:
|
|
11077
|
-
activityRegularizer:
|
|
11076
|
+
biasRegularizer: Is(this.biasRegularizer),
|
|
11077
|
+
activityRegularizer: Is(this.activityRegularizer),
|
|
11078
11078
|
biasConstraint: xr(this.biasConstraint)
|
|
11079
11079
|
}, e = super.getConfig();
|
|
11080
11080
|
return Object.assign(t, e), t;
|
|
@@ -11082,13 +11082,13 @@ class Gi extends ye {
|
|
|
11082
11082
|
}
|
|
11083
11083
|
class Jn extends Gi {
|
|
11084
11084
|
constructor(t, e) {
|
|
11085
|
-
super(t, e), this.kernel = null, Jn.verifyArgs(e), this.filters = e.filters, Oe(this.filters, "filters"), this.kernelInitializer = vs(e.kernelInitializer || this.DEFAULT_KERNEL_INITIALIZER), this.kernelConstraint = Sr(e.kernelConstraint), this.kernelRegularizer =
|
|
11085
|
+
super(t, e), this.kernel = null, Jn.verifyArgs(e), this.filters = e.filters, Oe(this.filters, "filters"), this.kernelInitializer = vs(e.kernelInitializer || this.DEFAULT_KERNEL_INITIALIZER), this.kernelConstraint = Sr(e.kernelConstraint), this.kernelRegularizer = $s(e.kernelRegularizer);
|
|
11086
11086
|
}
|
|
11087
11087
|
build(t) {
|
|
11088
11088
|
t = ge(t);
|
|
11089
11089
|
const e = this.dataFormat === "channelsFirst" ? 1 : t.length - 1;
|
|
11090
11090
|
if (t[e] == null)
|
|
11091
|
-
throw new
|
|
11091
|
+
throw new I(`The channel dimension of the input should be defined. Found ${t[e]}`);
|
|
11092
11092
|
const s = t[e], r = this.kernelSize.concat([s, this.filters]);
|
|
11093
11093
|
this.kernel = this.addWeight("kernel", r, null, this.kernelInitializer, this.kernelRegularizer, !0, this.kernelConstraint), this.useBias && (this.bias = this.addWeight("bias", [this.filters], null, this.biasInitializer, this.biasRegularizer, !0, this.biasConstraint)), this.inputSpec = [{ ndim: this.rank + 2, axes: { [e]: s } }], this.built = !0;
|
|
11094
11094
|
}
|
|
@@ -11127,14 +11127,14 @@ class Jn extends Gi {
|
|
|
11127
11127
|
const t = {
|
|
11128
11128
|
filters: this.filters,
|
|
11129
11129
|
kernelInitializer: vr(this.kernelInitializer),
|
|
11130
|
-
kernelRegularizer:
|
|
11130
|
+
kernelRegularizer: Is(this.kernelRegularizer),
|
|
11131
11131
|
kernelConstraint: xr(this.kernelConstraint)
|
|
11132
11132
|
}, e = super.getConfig();
|
|
11133
11133
|
return Object.assign(t, e), t;
|
|
11134
11134
|
}
|
|
11135
11135
|
static verifyArgs(t) {
|
|
11136
11136
|
if (!("filters" in t) || typeof t.filters != "number" || t.filters < 1)
|
|
11137
|
-
throw new
|
|
11137
|
+
throw new I(`Convolution layer expected config.filters to be a 'number' > 0 but got ${JSON.stringify(t.filters)}`);
|
|
11138
11138
|
}
|
|
11139
11139
|
}
|
|
11140
11140
|
class Zn extends Jn {
|
|
@@ -11147,7 +11147,7 @@ class Zn extends Jn {
|
|
|
11147
11147
|
}
|
|
11148
11148
|
static verifyArgs(t) {
|
|
11149
11149
|
if (typeof t.kernelSize != "number" && !Di(t.kernelSize, "number", 1, 2))
|
|
11150
|
-
throw new
|
|
11150
|
+
throw new I(`Conv2D expects config.kernelSize to be number or number[] with length 1 or 2, but received ${JSON.stringify(t.kernelSize)}.`);
|
|
11151
11151
|
}
|
|
11152
11152
|
}
|
|
11153
11153
|
Zn.className = "Conv2D";
|
|
@@ -11162,7 +11162,7 @@ class Ls extends Jn {
|
|
|
11162
11162
|
}
|
|
11163
11163
|
static verifyArgs(t) {
|
|
11164
11164
|
if (typeof t.kernelSize != "number" && !(Array.isArray(t.kernelSize) && (t.kernelSize.length === 1 || t.kernelSize.length === 3)))
|
|
11165
|
-
throw new
|
|
11165
|
+
throw new I(`Conv3D expects config.kernelSize to be number or [number, number, number], but received ${JSON.stringify(t.kernelSize)}.`);
|
|
11166
11166
|
}
|
|
11167
11167
|
}
|
|
11168
11168
|
Ls.className = "Conv3D";
|
|
@@ -11170,14 +11170,14 @@ O(Ls);
|
|
|
11170
11170
|
class Mc extends Zn {
|
|
11171
11171
|
constructor(t) {
|
|
11172
11172
|
if (super(t), this.inputSpec = [new Ce({ ndim: 4 })], this.padding !== "same" && this.padding !== "valid")
|
|
11173
|
-
throw new
|
|
11173
|
+
throw new I(`Conv2DTranspose currently supports only padding modes 'same' and 'valid', but received padding mode ${this.padding}`);
|
|
11174
11174
|
}
|
|
11175
11175
|
build(t) {
|
|
11176
11176
|
if (t = ge(t), t.length !== 4)
|
|
11177
|
-
throw new
|
|
11177
|
+
throw new I("Input should have rank 4; Received input shape: " + JSON.stringify(t));
|
|
11178
11178
|
const e = this.dataFormat === "channelsFirst" ? 1 : t.length - 1;
|
|
11179
11179
|
if (t[e] == null)
|
|
11180
|
-
throw new
|
|
11180
|
+
throw new I("The channel dimension of the inputs should be defined. Found `None`.");
|
|
11181
11181
|
const s = t[e], r = this.kernelSize.concat([this.filters, s]);
|
|
11182
11182
|
this.kernel = this.addWeight("kernel", r, "float32", this.kernelInitializer, this.kernelRegularizer, !0, this.kernelConstraint), this.useBias && (this.bias = this.addWeight("bias", [this.filters], "float32", this.biasInitializer, this.biasRegularizer, !0, this.biasConstraint)), this.inputSpec = [new Ce({ ndim: 4, axes: { [e]: s } })], this.built = !0;
|
|
11183
11183
|
}
|
|
@@ -11185,7 +11185,7 @@ class Mc extends Zn {
|
|
|
11185
11185
|
return _(() => {
|
|
11186
11186
|
let s = Vt(t);
|
|
11187
11187
|
if (s.shape.length !== 4)
|
|
11188
|
-
throw new
|
|
11188
|
+
throw new I(`Conv2DTranspose.call() expects input tensor to be rank-4, but received a tensor of rank-${s.shape.length}`);
|
|
11189
11189
|
const r = s.shape, o = r[0];
|
|
11190
11190
|
let i, a;
|
|
11191
11191
|
this.dataFormat === "channelsFirst" ? (i = 2, a = 3) : (i = 1, a = 2);
|
|
@@ -11213,14 +11213,14 @@ O(Mc);
|
|
|
11213
11213
|
class Bc extends Ls {
|
|
11214
11214
|
constructor(t) {
|
|
11215
11215
|
if (super(t), this.inputSpec = [new Ce({ ndim: 5 })], this.padding !== "same" && this.padding !== "valid")
|
|
11216
|
-
throw new
|
|
11216
|
+
throw new I(`Conv3DTranspose currently supports only padding modes 'same' and 'valid', but received padding mode ${this.padding}`);
|
|
11217
11217
|
}
|
|
11218
11218
|
build(t) {
|
|
11219
11219
|
if (t = ge(t), t.length !== 5)
|
|
11220
|
-
throw new
|
|
11220
|
+
throw new I("Input should have rank 5; Received input shape: " + JSON.stringify(t));
|
|
11221
11221
|
const e = this.dataFormat === "channelsFirst" ? 1 : t.length - 1;
|
|
11222
11222
|
if (t[e] == null)
|
|
11223
|
-
throw new
|
|
11223
|
+
throw new I("The channel dimension of the inputs should be defined. Found `None`.");
|
|
11224
11224
|
const s = t[e], r = this.kernelSize.concat([this.filters, s]);
|
|
11225
11225
|
this.kernel = this.addWeight("kernel", r, "float32", this.kernelInitializer, this.kernelRegularizer, !0, this.kernelConstraint), this.useBias && (this.bias = this.addWeight("bias", [this.filters], "float32", this.biasInitializer, this.biasRegularizer, !0, this.biasConstraint)), this.inputSpec = [new Ce({ ndim: 5, axes: { [e]: s } })], this.built = !0;
|
|
11226
11226
|
}
|
|
@@ -11228,7 +11228,7 @@ class Bc extends Ls {
|
|
|
11228
11228
|
return _(() => {
|
|
11229
11229
|
let s = Vt(t);
|
|
11230
11230
|
if (s.shape.length !== 5)
|
|
11231
|
-
throw new
|
|
11231
|
+
throw new I(`Conv3DTranspose.call() expects input tensor to be rank-4, but received a tensor of rank-${s.shape.length}`);
|
|
11232
11232
|
const r = s.shape, o = r[0];
|
|
11233
11233
|
let i, a, l;
|
|
11234
11234
|
this.dataFormat === "channelsFirst" ? (l = 2, i = 3, a = 4) : (l = 1, i = 2, a = 3);
|
|
@@ -11256,19 +11256,19 @@ O(Bc);
|
|
|
11256
11256
|
class Fc extends Jn {
|
|
11257
11257
|
constructor(t, e) {
|
|
11258
11258
|
if (super(t, e), this.DEFAULT_DEPTHWISE_INITIALIZER = "glorotUniform", this.DEFAULT_POINTWISE_INITIALIZER = "glorotUniform", this.depthwiseKernel = null, this.pointwiseKernel = null, e.filters == null)
|
|
11259
|
-
throw new
|
|
11259
|
+
throw new I("The `filters` configuration field is required by SeparableConv, but is unspecified.");
|
|
11260
11260
|
if (e.kernelInitializer != null || e.kernelRegularizer != null || e.kernelConstraint != null)
|
|
11261
|
-
throw new
|
|
11261
|
+
throw new I("Fields kernelInitializer, kernelRegularizer and kernelConstraint are invalid for SeparableConv2D. Use depthwiseInitializer, depthwiseRegularizer, depthwiseConstraint, pointwiseInitializer, pointwiseRegularizer and pointwiseConstraint instead.");
|
|
11262
11262
|
if (e.padding != null && e.padding !== "same" && e.padding !== "valid")
|
|
11263
|
-
throw new
|
|
11264
|
-
this.depthMultiplier = e.depthMultiplier == null ? 1 : e.depthMultiplier, this.depthwiseInitializer = vs(e.depthwiseInitializer || this.DEFAULT_DEPTHWISE_INITIALIZER), this.depthwiseRegularizer =
|
|
11263
|
+
throw new I(`SeparableConv${this.rank}D supports only padding modes: 'same' and 'valid', but received ${JSON.stringify(e.padding)}`);
|
|
11264
|
+
this.depthMultiplier = e.depthMultiplier == null ? 1 : e.depthMultiplier, this.depthwiseInitializer = vs(e.depthwiseInitializer || this.DEFAULT_DEPTHWISE_INITIALIZER), this.depthwiseRegularizer = $s(e.depthwiseRegularizer), this.depthwiseConstraint = Sr(e.depthwiseConstraint), this.pointwiseInitializer = vs(e.depthwiseInitializer || this.DEFAULT_POINTWISE_INITIALIZER), this.pointwiseRegularizer = $s(e.pointwiseRegularizer), this.pointwiseConstraint = Sr(e.pointwiseConstraint);
|
|
11265
11265
|
}
|
|
11266
11266
|
build(t) {
|
|
11267
11267
|
if (t = ge(t), t.length < this.rank + 2)
|
|
11268
|
-
throw new
|
|
11268
|
+
throw new I(`Inputs to SeparableConv${this.rank}D should have rank ${this.rank + 2}, but received input shape: ${JSON.stringify(t)}`);
|
|
11269
11269
|
const e = this.dataFormat === "channelsFirst" ? 1 : t.length - 1;
|
|
11270
11270
|
if (t[e] == null || t[e] < 0)
|
|
11271
|
-
throw new
|
|
11271
|
+
throw new I(`The channel dimension of the inputs should be defined, but found ${JSON.stringify(t[e])}`);
|
|
11272
11272
|
const s = t[e], r = this.kernelSize.concat([s, this.depthMultiplier]), o = [];
|
|
11273
11273
|
for (let a = 0; a < this.rank; ++a)
|
|
11274
11274
|
o.push(1);
|
|
@@ -11287,7 +11287,7 @@ class Fc extends Jn {
|
|
|
11287
11287
|
}
|
|
11288
11288
|
getConfig() {
|
|
11289
11289
|
const t = super.getConfig();
|
|
11290
|
-
return delete t.rank, delete t.kernelInitializer, delete t.kernelRegularizer, delete t.kernelConstraint, t.depthwiseInitializer = vr(this.depthwiseInitializer), t.pointwiseInitializer = vr(this.pointwiseInitializer), t.depthwiseRegularizer =
|
|
11290
|
+
return delete t.rank, delete t.kernelInitializer, delete t.kernelRegularizer, delete t.kernelConstraint, t.depthwiseInitializer = vr(this.depthwiseInitializer), t.pointwiseInitializer = vr(this.pointwiseInitializer), t.depthwiseRegularizer = Is(this.depthwiseRegularizer), t.pointwiseRegularizer = Is(this.pointwiseRegularizer), t.depthwiseConstraint = xr(this.depthwiseConstraint), t.pointwiseConstraint = xr(this.pointwiseConstraint), t;
|
|
11291
11291
|
}
|
|
11292
11292
|
}
|
|
11293
11293
|
Fc.className = "SeparableConv";
|
|
@@ -11308,7 +11308,7 @@ class Hr extends Jn {
|
|
|
11308
11308
|
}
|
|
11309
11309
|
static verifyArgs(t) {
|
|
11310
11310
|
if (typeof t.kernelSize != "number" && !Di(t.kernelSize, "number", 1, 1))
|
|
11311
|
-
throw new
|
|
11311
|
+
throw new I(`Conv1D expects config.kernelSize to be number or number[] with length 1, but received ${JSON.stringify(t.kernelSize)}.`);
|
|
11312
11312
|
}
|
|
11313
11313
|
}
|
|
11314
11314
|
Hr.className = "Conv1D";
|
|
@@ -11433,7 +11433,7 @@ class Gc extends ye {
|
|
|
11433
11433
|
else if (Array.isArray(t.poolSize) && t.poolSize.length === 1 && typeof t.poolSize[0] == "number")
|
|
11434
11434
|
this.poolSize = t.poolSize;
|
|
11435
11435
|
else
|
|
11436
|
-
throw new
|
|
11436
|
+
throw new I(`poolSize for 1D convolutional layer must be a number or an Array of a single number, but received ${JSON.stringify(t.poolSize)}`);
|
|
11437
11437
|
if (Oe(this.poolSize, "poolSize"), t.strides == null)
|
|
11438
11438
|
this.strides = this.poolSize;
|
|
11439
11439
|
else if (typeof t.strides == "number")
|
|
@@ -11441,7 +11441,7 @@ class Gc extends ye {
|
|
|
11441
11441
|
else if (Array.isArray(t.strides) && t.strides.length === 1 && typeof t.strides[0] == "number")
|
|
11442
11442
|
this.strides = t.strides;
|
|
11443
11443
|
else
|
|
11444
|
-
throw new
|
|
11444
|
+
throw new I(`strides for 1D convolutional layer must be a number or an Array of a single number, but received ${JSON.stringify(t.strides)}`);
|
|
11445
11445
|
Oe(this.strides, "strides"), this.padding = t.padding == null ? "valid" : t.padding, ie(this.padding), this.inputSpec = [new Ce({ ndim: 3 })];
|
|
11446
11446
|
}
|
|
11447
11447
|
computeOutputShape(t) {
|
|
@@ -11491,7 +11491,7 @@ class jc extends ye {
|
|
|
11491
11491
|
this.strides = this.poolSize;
|
|
11492
11492
|
else if (Array.isArray(t.strides)) {
|
|
11493
11493
|
if (t.strides.length !== 2)
|
|
11494
|
-
throw new
|
|
11494
|
+
throw new I(`If the strides property of a 2D pooling layer is an Array, it is expected to have a length of 2, but received length ${t.strides.length}.`);
|
|
11495
11495
|
this.strides = t.strides;
|
|
11496
11496
|
} else
|
|
11497
11497
|
this.strides = [t.strides, t.strides];
|
|
@@ -11541,7 +11541,7 @@ class Kc extends ye {
|
|
|
11541
11541
|
this.strides = this.poolSize;
|
|
11542
11542
|
else if (Array.isArray(t.strides)) {
|
|
11543
11543
|
if (t.strides.length !== 3)
|
|
11544
|
-
throw new
|
|
11544
|
+
throw new I(`If the strides property of a 3D pooling layer is an Array, it is expected to have a length of 3, but received length ${t.strides.length}.`);
|
|
11545
11545
|
this.strides = t.strides;
|
|
11546
11546
|
} else
|
|
11547
11547
|
this.strides = [t.strides, t.strides, t.strides];
|
|
@@ -11703,13 +11703,13 @@ function vw(n, t) {
|
|
|
11703
11703
|
return St(e, -1);
|
|
11704
11704
|
});
|
|
11705
11705
|
}
|
|
11706
|
-
function
|
|
11706
|
+
function Iw(n, t) {
|
|
11707
11707
|
return _(() => {
|
|
11708
11708
|
const e = et(N(n, t), -1), s = Ge(N(Z(1, n), t), -1);
|
|
11709
11709
|
return jn(0, M(1, Z(s, e)));
|
|
11710
11710
|
});
|
|
11711
11711
|
}
|
|
11712
|
-
function
|
|
11712
|
+
function $w(n, t) {
|
|
11713
11713
|
return _(() => {
|
|
11714
11714
|
const e = Math.log(2), s = Z(t, n), r = Z(M(s, yi(N(-2, s))), e);
|
|
11715
11715
|
return St(r, -1);
|
|
@@ -11736,7 +11736,7 @@ function Er(n, t, e = !1) {
|
|
|
11736
11736
|
}
|
|
11737
11737
|
function Aw(n, t) {
|
|
11738
11738
|
if (!oe(n.shape, t.shape))
|
|
11739
|
-
throw new
|
|
11739
|
+
throw new I(`logits and labels must have the same shape, but got shapes ${JSON.stringify(n.shape)} and ${JSON.stringify(t.shape)}`);
|
|
11740
11740
|
return _(() => {
|
|
11741
11741
|
const e = Ts(t), s = qn(Nt(t));
|
|
11742
11742
|
return M(Z(e, N(t, n)), Vg(Go(s)));
|
|
@@ -11773,8 +11773,8 @@ const _r = {
|
|
|
11773
11773
|
meanSquaredLogarithmicError: xw,
|
|
11774
11774
|
squaredHinge: Sw,
|
|
11775
11775
|
hinge: vw,
|
|
11776
|
-
categoricalHinge:
|
|
11777
|
-
logcosh:
|
|
11776
|
+
categoricalHinge: Iw,
|
|
11777
|
+
logcosh: $w,
|
|
11778
11778
|
categoricalCrossentropy: As,
|
|
11779
11779
|
sparseCategoricalCrossentropy: Er,
|
|
11780
11780
|
binaryCrossentropy: Xr,
|
|
@@ -11787,7 +11787,7 @@ function bo(n) {
|
|
|
11787
11787
|
if (n in _r)
|
|
11788
11788
|
return _r[n];
|
|
11789
11789
|
let t = `Unknown loss ${n}`;
|
|
11790
|
-
throw n.toLowerCase().includes("softmaxcrossentropy") && (t = `Unknown loss ${n}. Use "categoricalCrossentropy" as the string name for tf.losses.softmaxCrossEntropy`), new
|
|
11790
|
+
throw n.toLowerCase().includes("softmaxcrossentropy") && (t = `Unknown loss ${n}. Use "categoricalCrossentropy" as the string name for tf.losses.softmaxCrossEntropy`), new I(t);
|
|
11791
11791
|
} else
|
|
11792
11792
|
return n;
|
|
11793
11793
|
}
|
|
@@ -11839,7 +11839,7 @@ class vn extends ye {
|
|
|
11839
11839
|
s.push(o);
|
|
11840
11840
|
else {
|
|
11841
11841
|
if (o !== i)
|
|
11842
|
-
throw new
|
|
11842
|
+
throw new I("Operands could not be broadcast together with shapes " + JSON.stringify(t) + " " + JSON.stringify(e));
|
|
11843
11843
|
s.push(o);
|
|
11844
11844
|
}
|
|
11845
11845
|
}
|
|
@@ -11847,12 +11847,12 @@ class vn extends ye {
|
|
|
11847
11847
|
}
|
|
11848
11848
|
build(t) {
|
|
11849
11849
|
if (Array.isArray(t) && !Array.isArray(t[0]) && (t = [ge(t)]), t = t, t.length < 2)
|
|
11850
|
-
throw new
|
|
11850
|
+
throw new I(`A merge layer should be called on an Array of at least 2 inputs. Got ${t.length} input(s).`);
|
|
11851
11851
|
let e = [];
|
|
11852
11852
|
for (const o of t)
|
|
11853
11853
|
o != null && o[0] !== null && e.push(o[0]);
|
|
11854
11854
|
if (e = an(e), e.length > 1)
|
|
11855
|
-
throw new
|
|
11855
|
+
throw new I(`Can not merge tensors with different batch sizes. Got tensors with shapes: ${JSON.stringify(t)}.`);
|
|
11856
11856
|
let s = t[0] == null ? null : t[0].slice(1);
|
|
11857
11857
|
for (let o = 1; o < t.length; ++o) {
|
|
11858
11858
|
const i = t[o] == null ? null : t[o].slice(1);
|
|
@@ -11923,14 +11923,14 @@ class vn extends ye {
|
|
|
11923
11923
|
if (e == null)
|
|
11924
11924
|
return null;
|
|
11925
11925
|
if (!Array.isArray(e))
|
|
11926
|
-
throw new
|
|
11926
|
+
throw new I("`mask` should be an Array");
|
|
11927
11927
|
if (!Array.isArray(t))
|
|
11928
|
-
throw new
|
|
11928
|
+
throw new I("`inputs` should be an Array");
|
|
11929
11929
|
if (e.length !== t.length)
|
|
11930
|
-
throw new
|
|
11930
|
+
throw new I(`The Array 'inputs' and 'mask' are expected to have the same length, but have different lengths (${t.length} vs ${e.length})`);
|
|
11931
11931
|
if (e.every((r) => r == null))
|
|
11932
11932
|
return null;
|
|
11933
|
-
e = e.map((r) => r == null ? r :
|
|
11933
|
+
e = e.map((r) => r == null ? r : Ie(r, 0));
|
|
11934
11934
|
let s = e[0];
|
|
11935
11935
|
for (let r = 1; r < e.length - 1; ++r)
|
|
11936
11936
|
s = qr(s, e[r]);
|
|
@@ -12019,7 +12019,7 @@ class Ki extends vn {
|
|
|
12019
12019
|
}
|
|
12020
12020
|
build(t) {
|
|
12021
12021
|
if (!(Array.isArray(t) && Array.isArray(t[0])) || t.length === 1)
|
|
12022
|
-
throw new
|
|
12022
|
+
throw new I("A `Concatenate` layer should be called on a list of at least 2 inputs");
|
|
12023
12023
|
t = t;
|
|
12024
12024
|
let e = !0;
|
|
12025
12025
|
for (const r of t)
|
|
@@ -12042,14 +12042,14 @@ class Ki extends vn {
|
|
|
12042
12042
|
i || s.push(o);
|
|
12043
12043
|
}
|
|
12044
12044
|
if (s.length > 1)
|
|
12045
|
-
throw new
|
|
12045
|
+
throw new I("A `Concatenate` layer requires inputs with matching shapes except for the concat axis. Got input shapes: " + JSON.stringify(t));
|
|
12046
12046
|
}
|
|
12047
12047
|
mergeFunction(t) {
|
|
12048
12048
|
return _(() => Yy(t, this.axis));
|
|
12049
12049
|
}
|
|
12050
12050
|
computeOutputShape(t) {
|
|
12051
12051
|
if (!(Array.isArray(t) && Array.isArray(t[0])))
|
|
12052
|
-
throw new
|
|
12052
|
+
throw new I("A `Concatenate` layer should be called on a list of inputs.");
|
|
12053
12053
|
const e = t, s = e[0].slice(), r = this.axis < 0 ? s.length + this.axis : this.axis;
|
|
12054
12054
|
for (const o of e.slice(1)) {
|
|
12055
12055
|
if (s[r] == null || o[r] == null) {
|
|
@@ -12064,11 +12064,11 @@ class Ki extends vn {
|
|
|
12064
12064
|
if (e == null)
|
|
12065
12065
|
return null;
|
|
12066
12066
|
if (!Array.isArray(e))
|
|
12067
|
-
throw new
|
|
12067
|
+
throw new I("`mask` should be an array for Concatenate");
|
|
12068
12068
|
if (!Array.isArray(t))
|
|
12069
|
-
throw new
|
|
12069
|
+
throw new I("`inputs` should be an array for Concatenate");
|
|
12070
12070
|
if (e.length !== t.length)
|
|
12071
|
-
throw new
|
|
12071
|
+
throw new I(`Mismatch in the length of mask (${e.length}) and the legnth of inputs (${t.length})`);
|
|
12072
12072
|
return _(() => {
|
|
12073
12073
|
let s = !0;
|
|
12074
12074
|
if (e.forEach((i) => {
|
|
@@ -12080,7 +12080,7 @@ class Ki extends vn {
|
|
|
12080
12080
|
return null;
|
|
12081
12081
|
const r = [];
|
|
12082
12082
|
for (let i = 0; i < t.length; ++i)
|
|
12083
|
-
e[i] == null ? r.push(ot(Du(t[i]), "bool")) : e[i].rank < t[i].rank ? r.push(
|
|
12083
|
+
e[i] == null ? r.push(ot(Du(t[i]), "bool")) : e[i].rank < t[i].rank ? r.push(Ie(e[i], -1)) : r.push(e[i]);
|
|
12084
12084
|
const o = rn(r, this.axis);
|
|
12085
12085
|
return _m(o, -1, !1);
|
|
12086
12086
|
});
|
|
@@ -12138,7 +12138,7 @@ function Cw(n, t, e) {
|
|
|
12138
12138
|
u.push(c);
|
|
12139
12139
|
a = jr(a, u);
|
|
12140
12140
|
}
|
|
12141
|
-
return a.shape.length === 1 && (a =
|
|
12141
|
+
return a.shape.length === 1 && (a = Ie(a, 1)), a;
|
|
12142
12142
|
});
|
|
12143
12143
|
}
|
|
12144
12144
|
class uh extends vn {
|
|
@@ -12152,11 +12152,11 @@ class uh extends vn {
|
|
|
12152
12152
|
throw new J("Dot layer does not support tensors of 4D or higher rank yet.");
|
|
12153
12153
|
const r = this.interpretAxes(e, s);
|
|
12154
12154
|
if (e[r[0]] !== s[r[1]])
|
|
12155
|
-
throw new
|
|
12155
|
+
throw new I(`Dimension incompatibility: ${e[r[0]]} !== ${s[r[1]]}`);
|
|
12156
12156
|
}
|
|
12157
12157
|
mergeFunction(t) {
|
|
12158
12158
|
if (t.length !== 2)
|
|
12159
|
-
throw new
|
|
12159
|
+
throw new I(`A \`Dot\` layer must be called on exactly 2 inputs, but received ${t.length} input(s).`);
|
|
12160
12160
|
let e = t[0], s = t[1], r;
|
|
12161
12161
|
return Array.isArray(this.axes) ? r = this.axes.map((o, i) => os(o, t[i].shape.length)) : r = [
|
|
12162
12162
|
os(this.axes, e.shape.length),
|
|
@@ -12473,7 +12473,7 @@ class Qt {
|
|
|
12473
12473
|
for (const e in Qt.constructors)
|
|
12474
12474
|
Qt.constructors[+e].forEach((r) => {
|
|
12475
12475
|
if (r === t)
|
|
12476
|
-
throw new
|
|
12476
|
+
throw new I("Duplicate callback constructor.");
|
|
12477
12477
|
});
|
|
12478
12478
|
}
|
|
12479
12479
|
/**
|
|
@@ -12585,7 +12585,7 @@ function jw(n) {
|
|
|
12585
12585
|
return Cr[n];
|
|
12586
12586
|
if (typeof n != "string" && n != null)
|
|
12587
12587
|
return n;
|
|
12588
|
-
throw new
|
|
12588
|
+
throw new I(`Unknown metric ${n}`);
|
|
12589
12589
|
}
|
|
12590
12590
|
function qs(n) {
|
|
12591
12591
|
if (Ae(n !== null, `Unknown LossOrMetricFn ${n}`), typeof n == "string")
|
|
@@ -12618,16 +12618,16 @@ function qs(n) {
|
|
|
12618
12618
|
*/
|
|
12619
12619
|
function Hw(n) {
|
|
12620
12620
|
const t = {
|
|
12621
|
-
Adagrad: () =>
|
|
12622
|
-
Adadelta: () =>
|
|
12623
|
-
Adam: () =>
|
|
12624
|
-
Adamax: () =>
|
|
12625
|
-
RMSProp: () =>
|
|
12626
|
-
SGD: () =>
|
|
12621
|
+
Adagrad: () => In.adagrad(0.01),
|
|
12622
|
+
Adadelta: () => In.adadelta(1, 0.95, mt()),
|
|
12623
|
+
Adam: () => In.adam(1e-3, 0.9, 0.999, mt()),
|
|
12624
|
+
Adamax: () => In.adamax(2e-3, 0.9, 0.999, mt(), 0),
|
|
12625
|
+
RMSProp: () => In.rmsprop(1e-3, 0.9, 0, mt()),
|
|
12626
|
+
SGD: () => In.sgd(0.01)
|
|
12627
12627
|
};
|
|
12628
12628
|
if (t.adagrad = t.Adagrad, t.adadelta = t.Adadelta, t.adam = t.Adam, t.adamax = t.Adamax, t.rmsprop = t.RMSProp, t.sgd = t.SGD, n in t)
|
|
12629
12629
|
return t[n]();
|
|
12630
|
-
throw new
|
|
12630
|
+
throw new I(`Unknown Optimizer ${n}`);
|
|
12631
12631
|
}
|
|
12632
12632
|
/**
|
|
12633
12633
|
* @license
|
|
@@ -12692,12 +12692,12 @@ function Kw(n, t, e, s = console.log) {
|
|
|
12692
12692
|
for (let c = 0; c < a.length; ++c)
|
|
12693
12693
|
r ? Jw(a[c], e, s) : Zw(a[c], e, i, s), s((c === a.length - 1 ? "=" : "_").repeat(t));
|
|
12694
12694
|
n.checkTrainableWeightsConsistency();
|
|
12695
|
-
const l = Yw(n), u =
|
|
12695
|
+
const l = Yw(n), u = $r(n.nonTrainableWeights);
|
|
12696
12696
|
s(`Total params: ${l + u}`), s(`Trainable params: ${l}`), s(`Non-trainable params: ${u}`), s("_".repeat(t));
|
|
12697
12697
|
}
|
|
12698
12698
|
function Yw(n) {
|
|
12699
12699
|
let t;
|
|
12700
|
-
return n.collectedTrainableWeights != null ? t =
|
|
12700
|
+
return n.collectedTrainableWeights != null ? t = $r(n.collectedTrainableWeights) : t = $r(n.trainableWeights), t;
|
|
12701
12701
|
}
|
|
12702
12702
|
function Xw(n) {
|
|
12703
12703
|
let t = !0;
|
|
@@ -12846,7 +12846,7 @@ function Ko(n, t) {
|
|
|
12846
12846
|
}
|
|
12847
12847
|
}
|
|
12848
12848
|
/** @license See the LICENSE file. */
|
|
12849
|
-
const wh = "4.
|
|
12849
|
+
const wh = "4.20.0";
|
|
12850
12850
|
/**
|
|
12851
12851
|
* @license
|
|
12852
12852
|
* Copyright 2022 Google LLC
|
|
@@ -12916,14 +12916,14 @@ class Os extends ye {
|
|
|
12916
12916
|
dtype: t.dtype,
|
|
12917
12917
|
name: t.name != null ? t.name : Li("input").toString()
|
|
12918
12918
|
}), t.batchSize == null && (t.batchSize = null), t.sparse == null && (t.sparse = !1), this.trainable = !1, this.built = !0, this.sparse = t.sparse, t.inputShape != null && t.batchInputShape != null)
|
|
12919
|
-
throw new
|
|
12919
|
+
throw new I("Only provide the inputShape OR batchInputShape argument to inputLayer, not both at the same time.");
|
|
12920
12920
|
let e = t.batchInputShape;
|
|
12921
12921
|
if (e == null) {
|
|
12922
12922
|
if (t.inputShape == null)
|
|
12923
|
-
throw new
|
|
12923
|
+
throw new I("An InputLayer should be passed either a `batchInputShape` or an `inputShape`.");
|
|
12924
12924
|
e = [t.batchSize].concat(t.inputShape);
|
|
12925
12925
|
} else if (t.batchSize != null)
|
|
12926
|
-
throw new
|
|
12926
|
+
throw new I("Cannot specify batchSize if batchInputShape is specified when creating an InputLayer.");
|
|
12927
12927
|
const s = t.dtype || "float32";
|
|
12928
12928
|
this.batchInputShape = e, this.dtype = s, this.inputSpec = [{ shape: e }];
|
|
12929
12929
|
const r = new bn(this.dtype, this.batchInputShape, this, [], {}, this.name);
|
|
@@ -12941,7 +12941,7 @@ class Os extends ye {
|
|
|
12941
12941
|
});
|
|
12942
12942
|
}
|
|
12943
12943
|
apply(t, e) {
|
|
12944
|
-
throw new
|
|
12944
|
+
throw new I(`Cannot pass any input to an InputLayer's apply() method. InputLayer name: ${this.name}`);
|
|
12945
12945
|
}
|
|
12946
12946
|
dispose() {
|
|
12947
12947
|
return { refCountAfterDispose: this._refCount, numDisposedVariables: 0 };
|
|
@@ -12961,7 +12961,7 @@ function Qw(n) {
|
|
|
12961
12961
|
if (n.batchShape == null && n.shape == null)
|
|
12962
12962
|
throw new Error("Please provide to Input either a `shape` or a `batchShape` argument. Note that `shape` does not include the batch dimension.");
|
|
12963
12963
|
if (n.batchShape != null && n.shape != null)
|
|
12964
|
-
throw new
|
|
12964
|
+
throw new I("Please provide either a `shape` or `batchShape` argument to Input, but not both.");
|
|
12965
12965
|
let t = n.batchShape;
|
|
12966
12966
|
n.shape != null && t == null && (t = [null].concat(n.shape));
|
|
12967
12967
|
let e = n.dtype;
|
|
@@ -12987,7 +12987,7 @@ function t1(n, t) {
|
|
|
12987
12987
|
try {
|
|
12988
12988
|
return ot(t, n.dtype);
|
|
12989
12989
|
} catch {
|
|
12990
|
-
throw new
|
|
12990
|
+
throw new I(`The dtype of the feed (${t.dtype}) can not be cast to the dtype of the key '${n.name}' (${n.dtype}).`);
|
|
12991
12991
|
}
|
|
12992
12992
|
}
|
|
12993
12993
|
class Ue {
|
|
@@ -13021,7 +13021,7 @@ class Ue {
|
|
|
13021
13021
|
if (this.id2Value[t.id] == null)
|
|
13022
13022
|
this.id2Value[t.id] = t1(t, e), this.name2Id[t.name] = t.id, s != null && (this.id2Mask[t.id] = s);
|
|
13023
13023
|
else
|
|
13024
|
-
throw new
|
|
13024
|
+
throw new I(`Duplicate key: name=${t.name}, id=${t.id}`);
|
|
13025
13025
|
return this;
|
|
13026
13026
|
}
|
|
13027
13027
|
/**
|
|
@@ -13055,12 +13055,12 @@ class Ue {
|
|
|
13055
13055
|
getValue(t) {
|
|
13056
13056
|
if (t instanceof bn) {
|
|
13057
13057
|
if (this.id2Value[t.id] == null)
|
|
13058
|
-
throw new
|
|
13058
|
+
throw new I(`Nonexistent key: ${t.name}`);
|
|
13059
13059
|
return this.id2Value[t.id];
|
|
13060
13060
|
} else {
|
|
13061
13061
|
const e = this.name2Id[t];
|
|
13062
13062
|
if (e == null)
|
|
13063
|
-
throw new
|
|
13063
|
+
throw new I(`Feed dict has no SymbolicTensor name: ${t}`);
|
|
13064
13064
|
return this.id2Value[e];
|
|
13065
13065
|
}
|
|
13066
13066
|
}
|
|
@@ -13074,12 +13074,12 @@ class Ue {
|
|
|
13074
13074
|
getMask(t) {
|
|
13075
13075
|
if (t instanceof bn) {
|
|
13076
13076
|
if (this.id2Value[t.id] == null)
|
|
13077
|
-
throw new
|
|
13077
|
+
throw new I(`Nonexistent key: ${t.name}`);
|
|
13078
13078
|
return this.id2Mask[t.id];
|
|
13079
13079
|
} else {
|
|
13080
13080
|
const e = this.name2Id[t];
|
|
13081
13081
|
if (e == null)
|
|
13082
|
-
throw new
|
|
13082
|
+
throw new I(`Feed dict has no SymbolicTensor name: ${t}`);
|
|
13083
13083
|
return this.id2Mask[e];
|
|
13084
13084
|
}
|
|
13085
13085
|
}
|
|
@@ -13213,7 +13213,7 @@ class fe extends ye {
|
|
|
13213
13213
|
this.name = Li(y);
|
|
13214
13214
|
}
|
|
13215
13215
|
if (this.supportsMasking = !1, this.trainable_ = !0, Array.isArray(t.inputs) ? this.inputs = t.inputs.slice() : this.inputs = [t.inputs], Array.isArray(t.outputs) ? this.outputs = t.outputs.slice() : this.outputs = [t.outputs], an(this.inputs).length !== this.inputs.length)
|
|
13216
|
-
throw new
|
|
13216
|
+
throw new I(`The list of inputs passed to the model is redundant. All inputs should only appear once. Found: ${this.inputs.map((y) => y.name)}`);
|
|
13217
13217
|
an(this.outputs).length !== this.outputs.length && console.warn(`The list of outputs passed to the model is redundant. All outputs should only appear once. Found: ${this.outputs.map((y) => y.name)}`), this.inputLayers = [], this.inputLayersNodeIndices = [], this.inputLayersTensorIndices = [], this.outputLayers = [], this.outputLayersNodeIndices = [], this.outputLayersTensorIndices = [], this.layers = [], this.internalContainerRefs = [];
|
|
13218
13218
|
for (const y of this.outputs) {
|
|
13219
13219
|
const S = y.sourceLayer, x = y.nodeIndex, v = y.tensorIndex;
|
|
@@ -13369,7 +13369,7 @@ class fe extends ye {
|
|
|
13369
13369
|
}
|
|
13370
13370
|
get trainableWeights() {
|
|
13371
13371
|
if (this._trainableWeights.length > 0)
|
|
13372
|
-
throw new
|
|
13372
|
+
throw new I("Container instance unexpectedly contains _trainableWeights.The trainable weights of a Container are a union of the trainable weights of its consituent Layers. Its own _trainableWeights must remain an empty Array.");
|
|
13373
13373
|
if (!this.trainable)
|
|
13374
13374
|
return [];
|
|
13375
13375
|
let t = [];
|
|
@@ -13416,7 +13416,7 @@ class fe extends ye {
|
|
|
13416
13416
|
for (const [l, u] of a.weights.entries()) {
|
|
13417
13417
|
const c = o ? `${u.name.split("/").slice(0, -1).join("/") + "/"}${l}` : u.originalName;
|
|
13418
13418
|
if (s[c] != null)
|
|
13419
|
-
throw new
|
|
13419
|
+
throw new I(`Duplicate weight name: ${c}`);
|
|
13420
13420
|
s[c] = u, r++;
|
|
13421
13421
|
}
|
|
13422
13422
|
const i = [];
|
|
@@ -13429,7 +13429,7 @@ class fe extends ye {
|
|
|
13429
13429
|
if (s[l] != null)
|
|
13430
13430
|
i.push([s[l], t[a]]);
|
|
13431
13431
|
else if (e)
|
|
13432
|
-
throw new
|
|
13432
|
+
throw new I(`Provided weight data has no target variable: ${a}`);
|
|
13433
13433
|
delete s[l];
|
|
13434
13434
|
}
|
|
13435
13435
|
if (e) {
|
|
@@ -13437,7 +13437,7 @@ class fe extends ye {
|
|
|
13437
13437
|
for (const l in s)
|
|
13438
13438
|
a.push(l);
|
|
13439
13439
|
if (a.length > 0)
|
|
13440
|
-
throw new
|
|
13440
|
+
throw new I(`${a.length} of ${r} weights are not set: ${a}`);
|
|
13441
13441
|
}
|
|
13442
13442
|
Dc(i);
|
|
13443
13443
|
}
|
|
@@ -13519,9 +13519,9 @@ class fe extends ye {
|
|
|
13519
13519
|
* free dimensions, instead of an integer.
|
|
13520
13520
|
*/
|
|
13521
13521
|
computeOutputShape(t) {
|
|
13522
|
-
const e =
|
|
13522
|
+
const e = Ir(t);
|
|
13523
13523
|
if (e.length !== this.inputLayers.length)
|
|
13524
|
-
throw new
|
|
13524
|
+
throw new I(`Invalid inputShape argument ${t}: model has ${this.inputLayers.length} tensor inputs.`);
|
|
13525
13525
|
const s = {};
|
|
13526
13526
|
for (let a = 0; a < e.length; a++) {
|
|
13527
13527
|
const l = this.inputLayers[a], u = e[a], c = l.name + "_0_0";
|
|
@@ -13540,7 +13540,7 @@ class fe extends ye {
|
|
|
13540
13540
|
const m = u.inboundLayers[g], b = u.nodeIndices[g], y = u.tensorIndices[g], S = `${m.name}_${b}_${y}`, x = s[S];
|
|
13541
13541
|
h.push(x);
|
|
13542
13542
|
}
|
|
13543
|
-
const f = c.computeOutputShape(zt(h)), d =
|
|
13543
|
+
const f = c.computeOutputShape(zt(h)), d = Ir(f), p = c.inboundNodes.indexOf(u);
|
|
13544
13544
|
for (let g = 0; g < d.length; g++) {
|
|
13545
13545
|
const m = `${c.name}_${p}_${g}`;
|
|
13546
13546
|
s[m] = d[g];
|
|
@@ -13630,17 +13630,17 @@ class fe extends ye {
|
|
|
13630
13630
|
if (e != null)
|
|
13631
13631
|
return this.findLayer(e);
|
|
13632
13632
|
if (t == null)
|
|
13633
|
-
throw new
|
|
13633
|
+
throw new I("Provide either a layer name or layer index");
|
|
13634
13634
|
if (typeof t == "number")
|
|
13635
13635
|
return this.findLayer(t);
|
|
13636
13636
|
for (const s of this.layers)
|
|
13637
13637
|
if (s.name === t)
|
|
13638
13638
|
return s;
|
|
13639
|
-
throw new
|
|
13639
|
+
throw new I(`No such layer: ${t}`);
|
|
13640
13640
|
}
|
|
13641
13641
|
findLayer(t) {
|
|
13642
13642
|
if (this.layers.length <= t)
|
|
13643
|
-
throw new
|
|
13643
|
+
throw new I(`Was asked to retrieve layer at index ${t}, but model only has ${this.layers.length} layer(s).`);
|
|
13644
13644
|
return this.layers[t];
|
|
13645
13645
|
}
|
|
13646
13646
|
/**
|
|
@@ -13752,7 +13752,7 @@ class fe extends ye {
|
|
|
13752
13752
|
const b = m.name, y = dh(m, e.customObjects != null ? e.customObjects : {});
|
|
13753
13753
|
y.setFastWeightInitDuringBuild(r), o[b] = y, m.inboundNodes.forEach((x) => {
|
|
13754
13754
|
if (!(x instanceof Array))
|
|
13755
|
-
throw new
|
|
13755
|
+
throw new I(`Corrupted configuration, expected array for nodeData: ${x}`);
|
|
13756
13756
|
a(y, x);
|
|
13757
13757
|
});
|
|
13758
13758
|
}
|
|
@@ -13793,7 +13793,7 @@ class fe extends ye {
|
|
|
13793
13793
|
*/
|
|
13794
13794
|
get stateful() {
|
|
13795
13795
|
if (this._stateful)
|
|
13796
|
-
throw new
|
|
13796
|
+
throw new I("Container instance unexpectedly has _stateful = true. The statefulness of a Container is determined by the Layers it contains. Its _stateful property must remain the default false.");
|
|
13797
13797
|
for (const t of this.layers)
|
|
13798
13798
|
if (t.stateful)
|
|
13799
13799
|
return !0;
|
|
@@ -13880,7 +13880,7 @@ function i1(n, t) {
|
|
|
13880
13880
|
* =============================================================================
|
|
13881
13881
|
*/
|
|
13882
13882
|
const a1 = 32;
|
|
13883
|
-
function
|
|
13883
|
+
function Ih(n, t) {
|
|
13884
13884
|
let e, s;
|
|
13885
13885
|
const r = t;
|
|
13886
13886
|
e = r.xs, s = r.ys, w(e != null && s != null, () => `A Dataset iterator for fitDataset() is expected to generate objects of the form \`{xs: xVal, ys: yVal}\`, where the two values may be \`tf.Tensor\`, an array of Tensors, or a map of string to Tensor. The provided Dataset instead generates ${t}`);
|
|
@@ -13901,7 +13901,7 @@ function nl(n, t, e) {
|
|
|
13901
13901
|
const s = [];
|
|
13902
13902
|
for (const r of t) {
|
|
13903
13903
|
if (e[r] == null)
|
|
13904
|
-
throw new
|
|
13904
|
+
throw new I(`The feature data generated by the dataset lacks the required ${n} key '${r}'.`);
|
|
13905
13905
|
s.push(e[r]);
|
|
13906
13906
|
}
|
|
13907
13907
|
return s;
|
|
@@ -13959,7 +13959,7 @@ async function u1(n, t, e) {
|
|
|
13959
13959
|
break;
|
|
13960
13960
|
}
|
|
13961
13961
|
if (S.value != null) {
|
|
13962
|
-
const { xs: x, ys: v } =
|
|
13962
|
+
const { xs: x, ys: v } = Ih(n, S.value), E = {};
|
|
13963
13963
|
E.batch = y, E.size = x[0].shape[0], await f.onBatchBegin(y, E);
|
|
13964
13964
|
const D = [];
|
|
13965
13965
|
if (e.classWeight != null) {
|
|
@@ -14021,7 +14021,7 @@ async function f1(n, t, e) {
|
|
|
14021
14021
|
const u = await i.next();
|
|
14022
14022
|
if (o = _(() => {
|
|
14023
14023
|
if (u.value) {
|
|
14024
|
-
const { xs: c, ys: h } =
|
|
14024
|
+
const { xs: c, ys: h } = Ih(n, u.value), f = c.concat(h), d = _(() => r(f));
|
|
14025
14025
|
if (ut(f), l === 0)
|
|
14026
14026
|
for (let g = 0; g < d.length; ++g)
|
|
14027
14027
|
o.push(Yt(0));
|
|
@@ -14069,7 +14069,7 @@ function wo(n, t) {
|
|
|
14069
14069
|
r = s + t, r >= n && (r = n), e.push([s, r]), s = r;
|
|
14070
14070
|
return e;
|
|
14071
14071
|
}
|
|
14072
|
-
function
|
|
14072
|
+
function $h(n) {
|
|
14073
14073
|
const t = [];
|
|
14074
14074
|
n instanceof Et && (n = [n]);
|
|
14075
14075
|
for (let e = 0; e < n.length; ++e) {
|
|
@@ -14146,7 +14146,7 @@ function ol(n, t, e, s = !0, r = "") {
|
|
|
14146
14146
|
} else
|
|
14147
14147
|
i = !0;
|
|
14148
14148
|
if (i)
|
|
14149
|
-
throw new
|
|
14149
|
+
throw new I(`Error when checking model ${r} expected no data, but got ${n}`);
|
|
14150
14150
|
}
|
|
14151
14151
|
return [];
|
|
14152
14152
|
}
|
|
@@ -14157,31 +14157,31 @@ function ol(n, t, e, s = !0, r = "") {
|
|
|
14157
14157
|
n = n, o = [];
|
|
14158
14158
|
for (const i of t) {
|
|
14159
14159
|
if (n[i] == null)
|
|
14160
|
-
throw new
|
|
14160
|
+
throw new I(`No data provided for "${i}". Need data for each key in: ${t}`);
|
|
14161
14161
|
o.push(n[i]);
|
|
14162
14162
|
}
|
|
14163
14163
|
} else if (Xo(n)) {
|
|
14164
14164
|
if (n = n, n.length !== t.length)
|
|
14165
|
-
throw new
|
|
14165
|
+
throw new I(`Error when checking model ${r}: the Array of Tensors that you are passing to your model is not the size the model expected. Expected to see ${t.length} Tensor(s), but instead got the following list of Tensor(s): ${n}`);
|
|
14166
14166
|
o = n;
|
|
14167
14167
|
} else {
|
|
14168
14168
|
if (n = n, t.length > 1)
|
|
14169
|
-
throw new
|
|
14169
|
+
throw new I(`The model ${r} expects ${t.length} Tensor(s), but only received one Tensor. Found: Tensor with shape ${n.shape}`);
|
|
14170
14170
|
o = [n];
|
|
14171
14171
|
}
|
|
14172
|
-
if (o =
|
|
14172
|
+
if (o = $h(o), e != null)
|
|
14173
14173
|
for (let i = 0; i < t.length; ++i) {
|
|
14174
14174
|
if (e[i] == null)
|
|
14175
14175
|
continue;
|
|
14176
14176
|
const a = o[i];
|
|
14177
14177
|
if (a.shape.length !== e[i].length)
|
|
14178
|
-
throw new
|
|
14178
|
+
throw new I(`Error when checking ${r}: expected ${t[i]} to have ${e[i].length} dimension(s). but got array with shape ${a.shape}`);
|
|
14179
14179
|
for (let l = 0; l < e[i].length; ++l) {
|
|
14180
14180
|
if (l === 0 && !s)
|
|
14181
14181
|
continue;
|
|
14182
14182
|
const u = a.shape[l], c = e[i][l];
|
|
14183
14183
|
if (c != null && c >= 0 && u !== c)
|
|
14184
|
-
throw new
|
|
14184
|
+
throw new I(`${r} expected a batch of elements where each example has shape [${e[i].slice(1, e[i].length)}] (i.e.,tensor shape [*,${e[i].slice(1, e[i].length)}]) but the ${r} received an input with ${a.shape[0]} examples, each with shape [${a.shape.slice(1, a.shape.length)}] (tensor shape [${a.shape}])`);
|
|
14185
14185
|
}
|
|
14186
14186
|
}
|
|
14187
14187
|
return o;
|
|
@@ -14191,11 +14191,11 @@ function p1(n, t, e) {
|
|
|
14191
14191
|
s.sort();
|
|
14192
14192
|
const r = an(t.map((o) => o.shape[0]));
|
|
14193
14193
|
if (r.sort(), s.length > 1)
|
|
14194
|
-
throw new
|
|
14194
|
+
throw new I(`All input Tensors (x) should have the same number of samples. Got array shapes: ${JSON.stringify(n.map((o) => o.shape))}`);
|
|
14195
14195
|
if (r.length > 1)
|
|
14196
|
-
throw new
|
|
14196
|
+
throw new I(`All target Tensors (y) should have the same number of samples. Got array shapes: ${JSON.stringify(t.map((o) => o.shape))}`);
|
|
14197
14197
|
if (s.length > 0 && r.length > 0 && !oe(s, r))
|
|
14198
|
-
throw new
|
|
14198
|
+
throw new I(`Input Tensors should have the same number of samples as target Tensors. Found ${s[0]} input sample(s) and ${r[0]} target sample(s).`);
|
|
14199
14199
|
}
|
|
14200
14200
|
function m1(n, t, e) {
|
|
14201
14201
|
const s = [
|
|
@@ -14207,13 +14207,13 @@ function m1(n, t, e) {
|
|
|
14207
14207
|
const o = n[r], i = t[r], a = e[r];
|
|
14208
14208
|
if (i != null) {
|
|
14209
14209
|
if (i === As && o.shape[o.shape.length - 1] === 1)
|
|
14210
|
-
throw new
|
|
14210
|
+
throw new I(`You are passing a target array of shape ${o.shape} while using a loss 'categorical_crossentropy'. 'categorical_crossentropy'expects targets to be binary matrices (1s and 0s) of shape [samples, classes].`);
|
|
14211
14211
|
if (s.indexOf(i) !== -1) {
|
|
14212
14212
|
const l = o.shape.slice(1), u = a.slice(1);
|
|
14213
14213
|
for (let c = 0; c < l.length; ++c) {
|
|
14214
14214
|
const h = l[c], f = u[c];
|
|
14215
14215
|
if (f != null && h !== f)
|
|
14216
|
-
throw new
|
|
14216
|
+
throw new I(`A target Tensor with shape ${o.shape} was passed for an output of shape ${a}, while using a loss function that expects targets to have the same shape as the output.`);
|
|
14217
14217
|
}
|
|
14218
14218
|
}
|
|
14219
14219
|
}
|
|
@@ -14223,11 +14223,11 @@ function il(n, t, e, s = !0, r = "") {
|
|
|
14223
14223
|
let o;
|
|
14224
14224
|
if (Array.isArray(n)) {
|
|
14225
14225
|
if (n.length !== t.length)
|
|
14226
|
-
throw new
|
|
14226
|
+
throw new I(`Error when checking model ${r}: the Array of Tensors that you are passing to your model is not the size the the model expected. Expected to see ${t.length} Tensor(s), but instead got ${n.length} Tensors(s).`);
|
|
14227
14227
|
o = n;
|
|
14228
14228
|
} else {
|
|
14229
14229
|
if (t.length > 1)
|
|
14230
|
-
throw new
|
|
14230
|
+
throw new I(`The model expects ${t.length} ${r} Tensors, but only received one Tensor. Found: array with shape ${JSON.stringify(n.shape)}.`);
|
|
14231
14231
|
o = [n];
|
|
14232
14232
|
}
|
|
14233
14233
|
if (e != null)
|
|
@@ -14236,13 +14236,13 @@ function il(n, t, e, s = !0, r = "") {
|
|
|
14236
14236
|
continue;
|
|
14237
14237
|
const a = o[i];
|
|
14238
14238
|
if (a.shape.length !== e[i].length)
|
|
14239
|
-
throw new
|
|
14239
|
+
throw new I(`Error when checking ${r}: expected ${t[i]} to have ${e[i].length} dimension(s), but got array with shape ${JSON.stringify(a.shape)}`);
|
|
14240
14240
|
for (let l = 0; l < e[i].length; ++l) {
|
|
14241
14241
|
if (l === 0 && !s)
|
|
14242
14242
|
continue;
|
|
14243
14243
|
const u = a.shape[l], c = e[i][l];
|
|
14244
14244
|
if (c != null && c !== u)
|
|
14245
|
-
throw new
|
|
14245
|
+
throw new I(`Error when checking ${r}: expected ${t[i]} to have shape ${JSON.stringify(e[i])} but got array with shape ${JSON.stringify(a.shape)}.`);
|
|
14246
14246
|
}
|
|
14247
14247
|
}
|
|
14248
14248
|
}
|
|
@@ -14309,7 +14309,7 @@ class Jr extends fe {
|
|
|
14309
14309
|
*/
|
|
14310
14310
|
summary(t, e, s = console.log) {
|
|
14311
14311
|
if (!this.built)
|
|
14312
|
-
throw new
|
|
14312
|
+
throw new I("This model has never been called, thus its weights have not been created yet. So no summary can be displayed. Build the model first (e.g., by calling it on some test data).");
|
|
14313
14313
|
Kw(this, t, e, s);
|
|
14314
14314
|
}
|
|
14315
14315
|
/**
|
|
@@ -14327,7 +14327,7 @@ class Jr extends fe {
|
|
|
14327
14327
|
this.optimizer_ = Hw(t.optimizer), this.isOptimizerOwned = !0;
|
|
14328
14328
|
else {
|
|
14329
14329
|
if (!(t.optimizer instanceof Ke))
|
|
14330
|
-
throw new
|
|
14330
|
+
throw new I("User-defined optimizer must be an instance of tf.Optimizer.");
|
|
14331
14331
|
this.optimizer_ = t.optimizer, this.isOptimizerOwned = !1;
|
|
14332
14332
|
}
|
|
14333
14333
|
let e = [];
|
|
@@ -14335,12 +14335,12 @@ class Jr extends fe {
|
|
|
14335
14335
|
t.loss = t.loss;
|
|
14336
14336
|
for (const i in t.loss)
|
|
14337
14337
|
if (this.outputNames.indexOf(i) === -1)
|
|
14338
|
-
throw new
|
|
14338
|
+
throw new I(`Unknown entry in loss dictionary: "${i}". Only expected the following keys: ${this.outputNames}`);
|
|
14339
14339
|
for (const i of this.outputNames)
|
|
14340
14340
|
t.loss[i] == null && console.warn(`Output "${i}" is missing from loss dictionary. We assume this was done on purpose, and we will not be expecting data to be passed to ${i} during training`), e.push(bo(t.loss[i]));
|
|
14341
14341
|
} else if (Array.isArray(t.loss)) {
|
|
14342
14342
|
if (t.loss.length !== this.outputs.length)
|
|
14343
|
-
throw new
|
|
14343
|
+
throw new I(`When passing an Array as loss, it should have one entry per model output. The model has ${this.outputs.length} output(s), but you passed loss=${t.loss}.`);
|
|
14344
14344
|
e = t.loss.map((a) => bo(a));
|
|
14345
14345
|
} else {
|
|
14346
14346
|
const i = bo(t.loss);
|
|
@@ -14485,11 +14485,11 @@ class Jr extends fe {
|
|
|
14485
14485
|
let o;
|
|
14486
14486
|
if (s != null) {
|
|
14487
14487
|
if (o = null, e != null)
|
|
14488
|
-
throw new
|
|
14488
|
+
throw new I(`If ${r} is set, batchSize must be null or undefined.Got batchSize = ${e}`);
|
|
14489
14489
|
} else if (t != null)
|
|
14490
14490
|
Array.isArray(t) ? o = t[0].shape[0] : o = t.shape[0];
|
|
14491
14491
|
else
|
|
14492
|
-
throw new
|
|
14492
|
+
throw new I(`Either the input data should have a defined shape, or ${r} shoud be specified.`);
|
|
14493
14493
|
return o;
|
|
14494
14494
|
}
|
|
14495
14495
|
/**
|
|
@@ -14501,18 +14501,18 @@ class Jr extends fe {
|
|
|
14501
14501
|
*/
|
|
14502
14502
|
execute(t, e) {
|
|
14503
14503
|
if (Array.isArray(e) && e.length === 0)
|
|
14504
|
-
throw new
|
|
14504
|
+
throw new I("`outputs` is an empty Array, which is not allowed.");
|
|
14505
14505
|
const s = Array.isArray(e), r = s ? e : [e], o = this.retrieveSymbolicTensors(r), i = new Ue();
|
|
14506
14506
|
if (t instanceof Et && (t = [t]), Array.isArray(t)) {
|
|
14507
14507
|
if (t.length !== this.inputs.length)
|
|
14508
|
-
throw new
|
|
14508
|
+
throw new I(`The number of inputs provided (${t.length}) does not match the number of inputs of this model (${this.inputs.length}).`);
|
|
14509
14509
|
for (let l = 0; l < this.inputs.length; ++l)
|
|
14510
14510
|
i.add(this.inputs[l], t[l]);
|
|
14511
14511
|
} else
|
|
14512
14512
|
for (const l of this.inputs) {
|
|
14513
14513
|
const u = t[l.name];
|
|
14514
14514
|
if (u == null)
|
|
14515
|
-
throw new
|
|
14515
|
+
throw new I(`No value is provided for the model's input ${l.name}`);
|
|
14516
14516
|
i.add(l, u);
|
|
14517
14517
|
}
|
|
14518
14518
|
const a = cs(o, i);
|
|
@@ -14538,7 +14538,7 @@ class Jr extends fe {
|
|
|
14538
14538
|
const r = [];
|
|
14539
14539
|
throw e.forEach((o, i) => {
|
|
14540
14540
|
o == null && r.push(t[i]);
|
|
14541
|
-
}), new
|
|
14541
|
+
}), new I(`Cannot find SymbolicTensors for output name(s): ${JSON.stringify(r)}`);
|
|
14542
14542
|
}
|
|
14543
14543
|
return e;
|
|
14544
14544
|
}
|
|
@@ -14603,7 +14603,7 @@ class Jr extends fe {
|
|
|
14603
14603
|
* @doc {heading: 'Models', subheading: 'Classes'}
|
|
14604
14604
|
*/
|
|
14605
14605
|
predict(t, e = {}) {
|
|
14606
|
-
const s =
|
|
14606
|
+
const s = $h(t);
|
|
14607
14607
|
il(s, this.inputNames, this.feedInputShapes, !1);
|
|
14608
14608
|
try {
|
|
14609
14609
|
const r = e.batchSize == null ? 32 : e.batchSize;
|
|
@@ -14641,7 +14641,7 @@ class Jr extends fe {
|
|
|
14641
14641
|
this.feedLossFns[i] === Er ? o.push(a.slice(0, a.length - 1).concat([1])) : o.push(a);
|
|
14642
14642
|
}
|
|
14643
14643
|
if (t = ol(t, this.feedInputNames, this.feedInputShapes, !1, "input"), e = ol(e, this.feedOutputNames, o, !1, "target"), p1(t, e), m1(e, this.feedLossFns, this.feedOutputShapes), this.stateful && r != null && r > 0 && t[0].shape[0] % r !== 0)
|
|
14644
|
-
throw new
|
|
14644
|
+
throw new I(`In a stateful network, you should only pass inputs with a number of samples that is divisible by the batch size ${r}. Found: ${t[0].shape[0]} sample(s).`);
|
|
14645
14645
|
return [t, e];
|
|
14646
14646
|
}
|
|
14647
14647
|
async standardizeUserData(t, e, s, r, o = !0, i) {
|
|
@@ -14820,7 +14820,7 @@ class Jr extends fe {
|
|
|
14820
14820
|
if (s.validationData != null && s.validationData.length > 0) {
|
|
14821
14821
|
if (m = !0, s.validationData.length === 2)
|
|
14822
14822
|
l = s.validationData[0], u = s.validationData[1];
|
|
14823
|
-
else throw s.validationData.length === 3 ? new J("validationData including sample weights is not supported yet.") : new
|
|
14823
|
+
else throw s.validationData.length === 3 ? new J("validationData including sample weights is not supported yet.") : new I(`When passing validation data, it must contain 2 (valX, valY) or 3 (valX, valY, valSampleWeight) items; ${s.validationData} is invalid.`);
|
|
14824
14824
|
const R = await this.standardizeUserData(
|
|
14825
14825
|
l,
|
|
14826
14826
|
u,
|
|
@@ -14878,7 +14878,7 @@ class Jr extends fe {
|
|
|
14878
14878
|
r == null && (r = 32), o == null && (o = 1), c == null && (c = !0), f == null && (f = 0);
|
|
14879
14879
|
let g = !1;
|
|
14880
14880
|
if (l != null && u != null && (g = !0), p != null && (g = !0, d == null))
|
|
14881
|
-
throw new
|
|
14881
|
+
throw new I("Can only use `validationSteps` when doing step-wise training, i.e., `stepsPerEpoch` must be set.");
|
|
14882
14882
|
const m = this.checkNumSamples(e, r, d, "steps_per_epoch");
|
|
14883
14883
|
let b;
|
|
14884
14884
|
m != null && (b = wr(0, m)), i == null && (i = 1);
|
|
@@ -14892,7 +14892,7 @@ class Jr extends fe {
|
|
|
14892
14892
|
{
|
|
14893
14893
|
if (c === "batch")
|
|
14894
14894
|
throw new J("batch shuffling is not implemneted yet");
|
|
14895
|
-
c &&
|
|
14895
|
+
c && $f(b);
|
|
14896
14896
|
const E = Dt(b), D = wo(m, r);
|
|
14897
14897
|
for (let k = 0; k < D.length; ++k) {
|
|
14898
14898
|
const T = {};
|
|
@@ -15199,13 +15199,13 @@ class Jr extends fe {
|
|
|
15199
15199
|
if (typeof t == "string") {
|
|
15200
15200
|
const u = nm(t);
|
|
15201
15201
|
if (u.length === 0)
|
|
15202
|
-
throw new
|
|
15202
|
+
throw new I(`Cannot find any save handlers for URL '${t}'`);
|
|
15203
15203
|
if (u.length > 1)
|
|
15204
|
-
throw new
|
|
15204
|
+
throw new I(`Found more than one (${u.length}) save handlers for URL '${t}'`);
|
|
15205
15205
|
t = u[0];
|
|
15206
15206
|
}
|
|
15207
15207
|
if (t.save == null)
|
|
15208
|
-
throw new
|
|
15208
|
+
throw new I("LayersModel.save() cannot proceed because the IOHandler provided does not have the `save` attribute defined.");
|
|
15209
15209
|
const s = await Ta(this.getNamedWeights(e)), a = {
|
|
15210
15210
|
modelTopology: this.toJSON(null, !1),
|
|
15211
15211
|
format: b1,
|
|
@@ -15274,8 +15274,8 @@ const {
|
|
|
15274
15274
|
ownKeys: _h,
|
|
15275
15275
|
set: hl,
|
|
15276
15276
|
setPrototypeOf: Ch
|
|
15277
|
-
} = Reflect,
|
|
15278
|
-
EPSILON:
|
|
15277
|
+
} = Reflect, I1 = Proxy, {
|
|
15278
|
+
EPSILON: $1,
|
|
15279
15279
|
MAX_SAFE_INTEGER: fl,
|
|
15280
15280
|
isFinite: kh,
|
|
15281
15281
|
isNaN: Un
|
|
@@ -15314,27 +15314,27 @@ const {
|
|
|
15314
15314
|
), Xi = Qr[ke], D1 = ct(Xi), {
|
|
15315
15315
|
abs: P1,
|
|
15316
15316
|
trunc: Dh
|
|
15317
|
-
} = Math, to = ArrayBuffer, R1 = to.isView, Ph = to.prototype, L1 = ct(Ph.slice), O1 = Qn(Ph, "byteLength"), ei = typeof SharedArrayBuffer < "u" ? SharedArrayBuffer : null, M1 = ei && Qn(ei.prototype, "byteLength"), Ji = Ms(Uint8Array), B1 = Ji.from,
|
|
15318
|
-
|
|
15317
|
+
} = Math, to = ArrayBuffer, R1 = to.isView, Ph = to.prototype, L1 = ct(Ph.slice), O1 = Qn(Ph, "byteLength"), ei = typeof SharedArrayBuffer < "u" ? SharedArrayBuffer : null, M1 = ei && Qn(ei.prototype, "byteLength"), Ji = Ms(Uint8Array), B1 = Ji.from, It = Ji.prototype, F1 = It[ke], z1 = ct(It.keys), U1 = ct(
|
|
15318
|
+
It.values
|
|
15319
15319
|
), W1 = ct(
|
|
15320
|
-
|
|
15321
|
-
), G1 = ct(
|
|
15322
|
-
|
|
15323
|
-
), V1 = ct(
|
|
15324
|
-
|
|
15325
|
-
), ml = ct(
|
|
15326
|
-
|
|
15320
|
+
It.entries
|
|
15321
|
+
), G1 = ct(It.set), pl = ct(
|
|
15322
|
+
It.reverse
|
|
15323
|
+
), V1 = ct(It.fill), q1 = ct(
|
|
15324
|
+
It.copyWithin
|
|
15325
|
+
), ml = ct(It.sort), as = ct(It.slice), j1 = ct(
|
|
15326
|
+
It.subarray
|
|
15327
15327
|
), xt = Qn(
|
|
15328
|
-
|
|
15328
|
+
It,
|
|
15329
15329
|
"buffer"
|
|
15330
15330
|
), Xe = Qn(
|
|
15331
|
-
|
|
15331
|
+
It,
|
|
15332
15332
|
"byteOffset"
|
|
15333
15333
|
), tt = Qn(
|
|
15334
|
-
|
|
15334
|
+
It,
|
|
15335
15335
|
"length"
|
|
15336
15336
|
), Rh = Qn(
|
|
15337
|
-
|
|
15337
|
+
It,
|
|
15338
15338
|
Yi
|
|
15339
15339
|
), H1 = Uint8Array, Ht = Uint16Array, gl = (...n) => Ft(B1, Ht, n), Zi = Uint32Array, K1 = Float32Array, yn = Ms([][ke]()), eo = ct(yn.next), Y1 = ct(function* () {
|
|
15340
15340
|
}().next), X1 = Ms(yn), ft = TypeError, vo = RangeError, Lh = WeakSet, Oh = Lh.prototype, J1 = ct(Oh.add), Z1 = ct(Oh.has), no = WeakMap, Qi = no.prototype, Tr = ct(Qi.get), Q1 = ct(Qi.has), ta = ct(Qi.set), Mh = new no(), tx = Zr(null, {
|
|
@@ -15436,7 +15436,7 @@ function rx(n) {
|
|
|
15436
15436
|
throw ft(Eh);
|
|
15437
15437
|
return Qo(e, Pr);
|
|
15438
15438
|
}
|
|
15439
|
-
const si = 1 /
|
|
15439
|
+
const si = 1 / $1;
|
|
15440
15440
|
function ox(n) {
|
|
15441
15441
|
return n + si - si;
|
|
15442
15442
|
}
|
|
@@ -15489,7 +15489,7 @@ function De(n) {
|
|
|
15489
15489
|
const t = +n;
|
|
15490
15490
|
return Un(t) || t === 0 ? 0 : Dh(t);
|
|
15491
15491
|
}
|
|
15492
|
-
function
|
|
15492
|
+
function Io(n) {
|
|
15493
15493
|
const t = De(n);
|
|
15494
15494
|
return t < 0 ? 0 : t < fl ? t : fl;
|
|
15495
15495
|
}
|
|
@@ -15589,10 +15589,10 @@ function vl(n) {
|
|
|
15589
15589
|
return e;
|
|
15590
15590
|
}
|
|
15591
15591
|
const Hh = new Lh();
|
|
15592
|
-
for (const n of _h(
|
|
15592
|
+
for (const n of _h(It)) {
|
|
15593
15593
|
if (n === Yi)
|
|
15594
15594
|
continue;
|
|
15595
|
-
const t = zn(
|
|
15595
|
+
const t = zn(It, n);
|
|
15596
15596
|
Le(t, "get") && typeof t.get == "function" && J1(Hh, t.get);
|
|
15597
15597
|
}
|
|
15598
15598
|
const ux = _1(
|
|
@@ -15641,7 +15641,7 @@ class ht {
|
|
|
15641
15641
|
throw ft(ll);
|
|
15642
15642
|
l != null ? wl(t) ? (i = t, a = t.length) : (i = [.../** @type {Iterable<unknown>} */
|
|
15643
15643
|
t], a = i.length) : (i = /** @type {ArrayLike<unknown>} */
|
|
15644
|
-
t, a =
|
|
15644
|
+
t, a = Io(i.length)), r = hs(Ht, [a], new.target);
|
|
15645
15645
|
}
|
|
15646
15646
|
for (let l = 0; l < a; ++l)
|
|
15647
15647
|
r[l] = xe(i[l]);
|
|
@@ -15649,7 +15649,7 @@ class ht {
|
|
|
15649
15649
|
r = hs(Ht, arguments, new.target);
|
|
15650
15650
|
const o = (
|
|
15651
15651
|
/** @type {any} */
|
|
15652
|
-
new
|
|
15652
|
+
new I1(r, ux)
|
|
15653
15653
|
);
|
|
15654
15654
|
return ta(Rr, o, r), o;
|
|
15655
15655
|
}
|
|
@@ -15702,7 +15702,7 @@ class ht {
|
|
|
15702
15702
|
throw ft(
|
|
15703
15703
|
Jo
|
|
15704
15704
|
);
|
|
15705
|
-
r = Wn(t), o =
|
|
15705
|
+
r = Wn(t), o = Io(r.length);
|
|
15706
15706
|
}
|
|
15707
15707
|
const a = new s(o);
|
|
15708
15708
|
if (e.length === 0)
|
|
@@ -15970,7 +15970,7 @@ class ht {
|
|
|
15970
15970
|
if (bs(l))
|
|
15971
15971
|
throw ft(gs);
|
|
15972
15972
|
}
|
|
15973
|
-
const o = tt(s), i = Wn(t), a =
|
|
15973
|
+
const o = tt(s), i = Wn(t), a = Io(i.length);
|
|
15974
15974
|
if (r === 1 / 0 || a + r > o)
|
|
15975
15975
|
throw vo(xo);
|
|
15976
15976
|
for (let l = 0; l < a; ++l)
|
|
@@ -16162,7 +16162,7 @@ Bs(Lr, ke, {
|
|
|
16162
16162
|
writable: !0,
|
|
16163
16163
|
configurable: !0
|
|
16164
16164
|
});
|
|
16165
|
-
Ch(Lr,
|
|
16165
|
+
Ch(Lr, It);
|
|
16166
16166
|
function cx(n, t) {
|
|
16167
16167
|
return n.channels === t.channels;
|
|
16168
16168
|
}
|
|
@@ -16378,7 +16378,7 @@ function hx(n) {
|
|
|
16378
16378
|
return n <= Xh ? n = n / sa : n <= Jh ? n = Math.pow((n - ia) / ra, 1 / oa) : n = Math.exp((n - ua) / aa) - la, n;
|
|
16379
16379
|
}
|
|
16380
16380
|
const fx = 65504, Qh = Zh(fx), tf = 1 / Qh, ef = Qh;
|
|
16381
|
-
class
|
|
16381
|
+
class $o {
|
|
16382
16382
|
constructor(t, e, s, r) {
|
|
16383
16383
|
this.x = t, this.y = e, this.width = s, this.height = r;
|
|
16384
16384
|
}
|
|
@@ -16424,7 +16424,7 @@ function mx({
|
|
|
16424
16424
|
}
|
|
16425
16425
|
return s;
|
|
16426
16426
|
}
|
|
16427
|
-
const
|
|
16427
|
+
const Il = `
|
|
16428
16428
|
const a = ${sa};
|
|
16429
16429
|
const b = ${ra};
|
|
16430
16430
|
const c = ${oa};
|
|
@@ -16456,18 +16456,18 @@ class gx {
|
|
|
16456
16456
|
},
|
|
16457
16457
|
{
|
|
16458
16458
|
label: "inputSize",
|
|
16459
|
-
type: "
|
|
16460
|
-
data: new
|
|
16459
|
+
type: "vec2i",
|
|
16460
|
+
data: new Int32Array(2)
|
|
16461
16461
|
},
|
|
16462
16462
|
{
|
|
16463
16463
|
label: "outputSize",
|
|
16464
|
-
type: "
|
|
16465
|
-
data: new
|
|
16464
|
+
type: "vec2i",
|
|
16465
|
+
data: new Int32Array(2)
|
|
16466
16466
|
},
|
|
16467
16467
|
{
|
|
16468
16468
|
label: "inputOffset",
|
|
16469
|
-
type: "
|
|
16470
|
-
data: new
|
|
16469
|
+
type: "vec2i",
|
|
16470
|
+
data: new Int32Array(2)
|
|
16471
16471
|
}
|
|
16472
16472
|
];
|
|
16473
16473
|
this._inputPassAux = new Js("inputPassAux", this._device, {
|
|
@@ -16493,28 +16493,28 @@ class gx {
|
|
|
16493
16493
|
},
|
|
16494
16494
|
{
|
|
16495
16495
|
label: "inputSize",
|
|
16496
|
-
type: "
|
|
16497
|
-
data: new
|
|
16496
|
+
type: "vec2i",
|
|
16497
|
+
data: new Int32Array(2)
|
|
16498
16498
|
},
|
|
16499
16499
|
{
|
|
16500
16500
|
label: "outputSize",
|
|
16501
|
-
type: "
|
|
16502
|
-
data: new
|
|
16501
|
+
type: "vec2i",
|
|
16502
|
+
data: new Int32Array(2)
|
|
16503
16503
|
},
|
|
16504
16504
|
{
|
|
16505
16505
|
label: "imageSize",
|
|
16506
|
-
type: "
|
|
16507
|
-
data: new
|
|
16506
|
+
type: "vec2i",
|
|
16507
|
+
data: new Int32Array(2)
|
|
16508
16508
|
},
|
|
16509
16509
|
{
|
|
16510
16510
|
label: "inputOffset",
|
|
16511
|
-
type: "
|
|
16512
|
-
data: new
|
|
16511
|
+
type: "vec2i",
|
|
16512
|
+
data: new Int32Array(2)
|
|
16513
16513
|
},
|
|
16514
16514
|
{
|
|
16515
16515
|
label: "outputOffset",
|
|
16516
|
-
type: "
|
|
16517
|
-
data: new
|
|
16516
|
+
type: "vec2i",
|
|
16517
|
+
data: new Int32Array(2)
|
|
16518
16518
|
}
|
|
16519
16519
|
],
|
|
16520
16520
|
csDefine: "",
|
|
@@ -16526,8 +16526,8 @@ class gx {
|
|
|
16526
16526
|
uniforms: [
|
|
16527
16527
|
{
|
|
16528
16528
|
label: "size",
|
|
16529
|
-
type: "
|
|
16530
|
-
data: new
|
|
16529
|
+
type: "vec2i",
|
|
16530
|
+
data: new Int32Array(2)
|
|
16531
16531
|
}
|
|
16532
16532
|
],
|
|
16533
16533
|
csMain: (
|
|
@@ -16554,7 +16554,7 @@ out_color[outIdx] = textureLoad(in_color, globalId.xy, 0);
|
|
|
16554
16554
|
const s = this._isHDR, r = (
|
|
16555
16555
|
/* wgsl */
|
|
16556
16556
|
`
|
|
16557
|
-
${
|
|
16557
|
+
${Il}
|
|
16558
16558
|
fn PUForward(y: f32) -> f32 {
|
|
16559
16559
|
if (y <= y0) {
|
|
16560
16560
|
return a * y;
|
|
@@ -16571,12 +16571,12 @@ fn PUForward(y: f32) -> f32 {
|
|
|
16571
16571
|
const i = (
|
|
16572
16572
|
/* wgsl */
|
|
16573
16573
|
`
|
|
16574
|
-
let x =
|
|
16575
|
-
let y =
|
|
16576
|
-
let inIdx =
|
|
16574
|
+
let x = i32(globalId.x);
|
|
16575
|
+
let y = i32(globalId.y);
|
|
16576
|
+
let inIdx = (y + inputOffset.y) * inputSize.x + (x + inputOffset.x);
|
|
16577
16577
|
let col = ${o("color")};
|
|
16578
16578
|
|
|
16579
|
-
let outIdx =
|
|
16579
|
+
let outIdx = y * outputSize.x + x;
|
|
16580
16580
|
|
|
16581
16581
|
if (${e}) {
|
|
16582
16582
|
// Denoise the inversed alpha. Or the anti aliased edge will be too dark after denoised
|
|
@@ -16614,7 +16614,7 @@ ${i}
|
|
|
16614
16614
|
csDefine: (
|
|
16615
16615
|
/* wgsl */
|
|
16616
16616
|
`
|
|
16617
|
-
${
|
|
16617
|
+
${Il}
|
|
16618
16618
|
fn PUInverse(y: f32) -> f32 {
|
|
16619
16619
|
if (y <= x0) {
|
|
16620
16620
|
return y / a;
|
|
@@ -16629,13 +16629,13 @@ fn PUInverse(y: f32) -> f32 {
|
|
|
16629
16629
|
csMain: (
|
|
16630
16630
|
/* wgsl */
|
|
16631
16631
|
`
|
|
16632
|
-
let x =
|
|
16633
|
-
let y =
|
|
16632
|
+
let x = i32(globalId.x);
|
|
16633
|
+
let y = i32(globalId.y);
|
|
16634
16634
|
if (x >= outputSize.x || y >= outputSize.y) {
|
|
16635
16635
|
return;
|
|
16636
16636
|
}
|
|
16637
|
-
let inIdx =
|
|
16638
|
-
let outIdx =
|
|
16637
|
+
let inIdx = (y + inputOffset.y) * inputSize.x + x + inputOffset.x;
|
|
16638
|
+
let outIdx = (y + outputOffset.y) * imageSize.x + x + outputOffset.x;
|
|
16639
16639
|
let col = in_color[inIdx];
|
|
16640
16640
|
let raw = ${t ? "textureLoad(in_raw, globalId.xy + vec2u(outputOffset), 0)" : "in_raw[outIdx]"};
|
|
16641
16641
|
|
|
@@ -16657,19 +16657,19 @@ else {
|
|
|
16657
16657
|
});
|
|
16658
16658
|
}
|
|
16659
16659
|
setImageSize(t, e) {
|
|
16660
|
-
this._inputPassAux.setUniform("inputSize", new
|
|
16660
|
+
this._inputPassAux.setUniform("inputSize", new Int32Array([t, e])), this._inputPassColor.setUniform("inputSize", new Int32Array([t, e])), this._outputPass.setUniform("imageSize", new Int32Array([t, e])), this._outputPass.setSize(t, e), this._copyPass.setSize(t, e), this._copyPass.setUniform("size", new Int32Array([t, e]));
|
|
16661
16661
|
}
|
|
16662
16662
|
setInputTile(t) {
|
|
16663
|
-
const e = new
|
|
16663
|
+
const e = new Int32Array([t.width, t.height]);
|
|
16664
16664
|
[this._inputPassAux, this._inputPassColor].forEach((s) => {
|
|
16665
|
-
s.setUniform("inputOffset", new
|
|
16665
|
+
s.setUniform("inputOffset", new Int32Array([t.x, t.y])), s.setUniform("outputSize", e), s.setSize(e[0], e[1]);
|
|
16666
16666
|
}), this._outputPass.setUniform("inputSize", e);
|
|
16667
16667
|
}
|
|
16668
16668
|
setOutputTile(t, e) {
|
|
16669
|
-
const s = this._outputPass, r = new
|
|
16670
|
-
s.setUniform("outputSize", r), s.setUniform("inputOffset", new
|
|
16669
|
+
const s = this._outputPass, r = new Int32Array([t.width, t.height]), o = t.x - e.x, i = t.y - e.y;
|
|
16670
|
+
s.setUniform("outputSize", r), s.setUniform("inputOffset", new Int32Array([o, i])), s.setUniform(
|
|
16671
16671
|
"outputOffset",
|
|
16672
|
-
new
|
|
16672
|
+
new Int32Array([t.x, t.y])
|
|
16673
16673
|
), s.setExecuteSize(r[0], r[1]);
|
|
16674
16674
|
}
|
|
16675
16675
|
forward(t, e, s, r) {
|
|
@@ -16724,7 +16724,7 @@ else {
|
|
|
16724
16724
|
this._outputPass.dispose(), this._inputPassAux.dispose();
|
|
16725
16725
|
}
|
|
16726
16726
|
}
|
|
16727
|
-
function
|
|
16727
|
+
function $l(n, t) {
|
|
16728
16728
|
const e = n.buffer;
|
|
16729
16729
|
if (t === "Float32")
|
|
16730
16730
|
return new Float32Array(n.buffer);
|
|
@@ -16796,7 +16796,7 @@ class xx {
|
|
|
16796
16796
|
const h = u.desc.dims;
|
|
16797
16797
|
a = nr(
|
|
16798
16798
|
bx(
|
|
16799
|
-
|
|
16799
|
+
$l(u.data, u.desc.dataType),
|
|
16800
16800
|
h
|
|
16801
16801
|
),
|
|
16802
16802
|
[h[2], h[3], h[1], h[0]],
|
|
@@ -16806,7 +16806,7 @@ class xx {
|
|
|
16806
16806
|
if (!l) {
|
|
16807
16807
|
const h = this._hostTensors.get(t + ".bias");
|
|
16808
16808
|
l = Dt(
|
|
16809
|
-
|
|
16809
|
+
$l(h.data, h.desc.dataType),
|
|
16810
16810
|
"float32"
|
|
16811
16811
|
), i.set(o, l);
|
|
16812
16812
|
}
|
|
@@ -16940,7 +16940,7 @@ class xx {
|
|
|
16940
16940
|
g = Math.max(m - d.width, 0);
|
|
16941
16941
|
let b = o > 0 ? o * p.height - f : 0, y = Math.min(b + d.height, a);
|
|
16942
16942
|
b = Math.max(y - d.height, 0);
|
|
16943
|
-
const S = d.width, x = d.height, v = new
|
|
16943
|
+
const S = d.width, x = d.height, v = new $o(g, b, S, x);
|
|
16944
16944
|
let E, D = 1;
|
|
16945
16945
|
const k = this._device;
|
|
16946
16946
|
let T = this._dataProcessGPU;
|
|
@@ -16990,7 +16990,7 @@ class xx {
|
|
|
16990
16990
|
E = Ot(U);
|
|
16991
16991
|
}
|
|
16992
16992
|
let R;
|
|
16993
|
-
const B = this._tfModel.predict(E), H = Math.min(p.width, i), X = Math.min(p.height, a), W = new
|
|
16993
|
+
const B = this._tfModel.predict(E), H = Math.min(p.width, i), X = Math.min(p.height, a), W = new $o(r * H, o * X, H, X);
|
|
16994
16994
|
if (W.width = Math.min(W.width, i - W.x), W.height = Math.min(W.height, a - W.y), t instanceof Float32Array) {
|
|
16995
16995
|
let U = B.dataSync();
|
|
16996
16996
|
l && (U = mx({
|
|
@@ -17086,7 +17086,7 @@ class xx {
|
|
|
17086
17086
|
D,
|
|
17087
17087
|
// Is undefined if using webgpu buffer
|
|
17088
17088
|
b,
|
|
17089
|
-
new
|
|
17089
|
+
new $o(x * h, v * f, h, f),
|
|
17090
17090
|
x + v * p,
|
|
17091
17091
|
p * d
|
|
17092
17092
|
), x + 1 < p || v + 1 < d ? requestAnimationFrame(() => {
|
|
@@ -17235,7 +17235,7 @@ function El(n, t) {
|
|
|
17235
17235
|
* limitations under the License.
|
|
17236
17236
|
* =============================================================================
|
|
17237
17237
|
*/
|
|
17238
|
-
class
|
|
17238
|
+
class Ix {
|
|
17239
17239
|
constructor(t) {
|
|
17240
17240
|
this.device = t, this.numUsedTextures = 0, this.numFreeTextures = 0, this.freeTextures = /* @__PURE__ */ new Map(), this.usedTextures = /* @__PURE__ */ new Map(), this.numBytesUsed = 0, this.numBytesAllocated = 0;
|
|
17241
17241
|
}
|
|
@@ -17308,7 +17308,7 @@ function Cl(n) {
|
|
|
17308
17308
|
* limitations under the License.
|
|
17309
17309
|
* =============================================================================
|
|
17310
17310
|
*/
|
|
17311
|
-
function
|
|
17311
|
+
function $x(n, t) {
|
|
17312
17312
|
if (Math.max(...n) > 5)
|
|
17313
17313
|
throw new Error("Cannot symbolically compute strides for rank > 6 tensor.");
|
|
17314
17314
|
const e = n.length, s = "xyzwuv", r = n.map((i) => `${t}.${s[i]}`), o = new Array(e - 1);
|
|
@@ -17706,7 +17706,7 @@ function Rx(n, t) {
|
|
|
17706
17706
|
if (d.length === 1)
|
|
17707
17707
|
a += `let d${d[0]} = i32(globalId[${f}]);`;
|
|
17708
17708
|
else {
|
|
17709
|
-
const p =
|
|
17709
|
+
const p = $x(d, "uniforms.outShape");
|
|
17710
17710
|
a += `var index${f} = i32(globalId[${f}]);`;
|
|
17711
17711
|
for (let g = 0; g < p.length; g++)
|
|
17712
17712
|
a += `let d${d[g]} = index${f} / ${p[g]};`, g === p.length - 1 ? a += `let d${d[g + 1]} = index${f} - d${d[g]} * ${p[g]};` : a += `index${f} = index${f} - d${d[g]} * ${p[g]};`;
|
|
@@ -17849,7 +17849,7 @@ const hn = (n) => {
|
|
|
17849
17849
|
t *= n[e];
|
|
17850
17850
|
return t;
|
|
17851
17851
|
};
|
|
17852
|
-
function
|
|
17852
|
+
function $t(n, t, e = [1, 1, 1], s = [1, 1, 1]) {
|
|
17853
17853
|
const [r, o, i] = [
|
|
17854
17854
|
Math.ceil(hn(n.x.map((a) => t[a])) / (e[0] * s[0])),
|
|
17855
17855
|
n.y ? Math.ceil(hn(n.y.map((a) => t[a])) / (e[1] * s[1])) : 1,
|
|
@@ -17921,7 +17921,7 @@ class Fs extends Bl {
|
|
|
17921
17921
|
constructor(t, e) {
|
|
17922
17922
|
if (super(), this.commandQueueOwnedIds = /* @__PURE__ */ new WeakSet(), this.dispatchCountInPass = 0, this.disposed = !1, this.downloadWaitMs = 0, this.tensorDataPendingDisposal = [], this.queryResolveBuffer = null, this.querySet = null, this.querySetCount = 2, this.stagingPendingDisposal = [], this.uniformPendingDisposal = [], this.uploadWaitMs = 0, this.hasReadSyncWarned = !1, this.hasTimestampQueryWarned = !1, !rf())
|
|
17923
17923
|
throw new Error("WebGPU is not supported on this device");
|
|
17924
|
-
this.pipelineCache = {}, this.device = t, this.queue = t.queue, this.commandEncoder = null, this.computePassEncoder = null, this.adapterInfo = new Sx(e), this.supportTimestampQuery = this.device.features.has("timestamp-query"), this.thresholdToIncreaseWorkgroups = this.adapterInfo.intelGPUGeneration >= 12 ? 16 : 8, this.bufferManager = new vx(this.device), this.textureManager = new
|
|
17924
|
+
this.pipelineCache = {}, this.device = t, this.queue = t.queue, this.commandEncoder = null, this.computePassEncoder = null, this.adapterInfo = new Sx(e), this.supportTimestampQuery = this.device.features.has("timestamp-query"), this.thresholdToIncreaseWorkgroups = this.adapterInfo.intelGPUGeneration >= 12 ? 16 : 8, this.bufferManager = new vx(this.device), this.textureManager = new Ix(this.device), this.tensorMap = new If(this, uo()), V().getBool("WEBGPU_USE_PROFILE_TOOL") && (this.dummyCanvas = document.createElement("canvas"), this.dummyCanvas.width = 1, this.dummyCanvas.height = 1, this.dummyContext = this.dummyCanvas.getContext("webgpu"), this.dummyContext.configure({
|
|
17925
17925
|
device: t,
|
|
17926
17926
|
format: "bgra8unorm"
|
|
17927
17927
|
}), document.body.appendChild(this.dummyCanvas));
|
|
@@ -18350,7 +18350,7 @@ rf() && Xp(
|
|
|
18350
18350
|
maxComputeWorkgroupSizeX: r.maxComputeWorkgroupSizeX,
|
|
18351
18351
|
maxComputeInvocationsPerWorkgroup: r.maxComputeInvocationsPerWorkgroup
|
|
18352
18352
|
};
|
|
18353
|
-
const o = await t.requestDevice(e), i =
|
|
18353
|
+
const o = await t.requestDevice(e), i = await t.requestAdapterInfo();
|
|
18354
18354
|
return new Fs(o, i);
|
|
18355
18355
|
},
|
|
18356
18356
|
3
|
|
@@ -18377,7 +18377,7 @@ class Gx {
|
|
|
18377
18377
|
this.uniforms = "", this.variableNames = ["x"], this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = e.map(
|
|
18378
18378
|
(r, o) => r[0] + t[o] + r[1]
|
|
18379
18379
|
/* afterPad */
|
|
18380
|
-
), this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
18380
|
+
), this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.xShape = t, e.map((r, o) => {
|
|
18381
18381
|
this.uniforms += ` pad${o} : vec2<i32>,`;
|
|
18382
18382
|
}), this.offset = s === "reflect" ? 0 : 1, this.shaderKey = `mirrorPad_${s}`;
|
|
18383
18383
|
}
|
|
@@ -18486,7 +18486,7 @@ class Hx {
|
|
|
18486
18486
|
this.variableNames = ["x"], this.uniforms = "constantValue : f32,", this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = e.map(
|
|
18487
18487
|
(s, r) => s[0] + t[r] + s[1]
|
|
18488
18488
|
/* afterPad */
|
|
18489
|
-
), this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
18489
|
+
), this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), e.map((s, r) => {
|
|
18490
18490
|
this.uniforms += ` pad${r} : vec2<i32>,`;
|
|
18491
18491
|
}), this.xShape = t, this.shaderKey = "pad";
|
|
18492
18492
|
}
|
|
@@ -18519,7 +18519,7 @@ class Hx {
|
|
|
18519
18519
|
*/
|
|
18520
18520
|
class Kx {
|
|
18521
18521
|
constructor(t) {
|
|
18522
|
-
this.variableNames = [], this.outputShape = [], this.uniforms = "value : f32,", this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = t, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
18522
|
+
this.variableNames = [], this.outputShape = [], this.uniforms = "value : f32,", this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = t, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.shaderKey = "fill";
|
|
18523
18523
|
}
|
|
18524
18524
|
getUserCode() {
|
|
18525
18525
|
return `
|
|
@@ -19215,7 +19215,7 @@ const vS = Lt((n, t) => n !== t ? 1 : 0);
|
|
|
19215
19215
|
* limitations under the License.
|
|
19216
19216
|
* =============================================================================
|
|
19217
19217
|
*/
|
|
19218
|
-
function
|
|
19218
|
+
function IS(n, t, e, s, r) {
|
|
19219
19219
|
const o = t.length, i = z(t), a = Kt(t), l = Kt(r), u = Rn(e, z(r));
|
|
19220
19220
|
for (let c = 0; c < i; ++c) {
|
|
19221
19221
|
const h = oi(c, o, a), f = new Array(h.length);
|
|
@@ -19242,7 +19242,7 @@ function $S(n, t, e, s, r) {
|
|
|
19242
19242
|
* limitations under the License.
|
|
19243
19243
|
* =============================================================================
|
|
19244
19244
|
*/
|
|
19245
|
-
function
|
|
19245
|
+
function $S(n, t, e, s) {
|
|
19246
19246
|
const [r, o] = bi(n, s), i = ci(t, "int32"), a = je(z(r), i), l = z(o);
|
|
19247
19247
|
for (let u = 0; u < a.length; ++u) {
|
|
19248
19248
|
const c = u * l;
|
|
@@ -19430,7 +19430,7 @@ function DS(n, t, e, s, r, o, i) {
|
|
|
19430
19430
|
* limitations under the License.
|
|
19431
19431
|
* =============================================================================
|
|
19432
19432
|
*/
|
|
19433
|
-
var Zt =
|
|
19433
|
+
var Zt = $e;
|
|
19434
19434
|
class Mr {
|
|
19435
19435
|
constructor(t, e, s, r, o, i, a, l, u, c) {
|
|
19436
19436
|
this.shape = t, this.shapeShape = e, this.values = s, this.valuesShape = r, this.valuesDType = o, this.defaultValue = i, this.defaultValueShape = a, this.rowPartitionValues = l, this.rowPartitionValuesShapes = u, this.rowPartitionTypes = dy(c), this.raggedRank = py(this.rowPartitionTypes);
|
|
@@ -19824,7 +19824,7 @@ function FS(n, t, e, s, r, o, i) {
|
|
|
19824
19824
|
const a = t[0], l = o[0], u = new Array(l), c = new Array(a), h = t[1];
|
|
19825
19825
|
if (l === 0) {
|
|
19826
19826
|
if (a !== 0)
|
|
19827
|
-
throw new Error(
|
|
19827
|
+
throw new Error($y(a));
|
|
19828
19828
|
const m = yt(e, 0), b = yt(r, 0);
|
|
19829
19829
|
return [
|
|
19830
19830
|
m,
|
|
@@ -20479,7 +20479,7 @@ const e2 = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
|
20479
20479
|
multiplyImpl: af,
|
|
20480
20480
|
negImpl: SS,
|
|
20481
20481
|
notEqualImpl: vS,
|
|
20482
|
-
prodImpl:
|
|
20482
|
+
prodImpl: $S,
|
|
20483
20483
|
raggedGatherImpl: NS,
|
|
20484
20484
|
raggedRangeImpl: DS,
|
|
20485
20485
|
raggedTensorToTensorImpl: PS,
|
|
@@ -20502,7 +20502,7 @@ const e2 = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
|
20502
20502
|
subImpl: JS,
|
|
20503
20503
|
tileImpl: ZS,
|
|
20504
20504
|
topKImpl: QS,
|
|
20505
|
-
transposeImpl:
|
|
20505
|
+
transposeImpl: IS,
|
|
20506
20506
|
uniqueImpl: t2
|
|
20507
20507
|
}, Symbol.toStringTag, { value: "Module" }));
|
|
20508
20508
|
/**
|
|
@@ -20521,7 +20521,7 @@ const e2 = /* @__PURE__ */ Object.freeze(/* @__PURE__ */ Object.defineProperty({
|
|
|
20521
20521
|
* limitations under the License.
|
|
20522
20522
|
* =============================================================================
|
|
20523
20523
|
*/
|
|
20524
|
-
const { addImpl:
|
|
20524
|
+
const { addImpl: EI, castImpl: _I, ceilImpl: CI, concatImpl: n2, equalImpl: kI, expImpl: TI, expm1Impl: NI, floorImpl: DI, floorDivImpl: PI, gatherNdImpl: RI, gatherV2Impl: LI, greaterEqualImpl: OI, greaterImpl: MI, lessEqualImpl: BI, lessImpl: FI, logImpl: zI, maxImpl: s2, maximumImpl: UI, minimumImpl: WI, multiplyImpl: GI, negImpl: VI, notEqualImpl: qI, prodImpl: r2, rangeImpl: jI, rsqrtImpl: HI, scatterImpl: KI, simpleAbsImpl: YI, sliceImpl: o2, stridedSliceImpl: XI, stringNGramsImpl: JI, subImpl: ZI, tileImpl: QI, topKImpl: t$, transposeImpl: i2, uniqueImpl: e$ } = e2;
|
|
20525
20525
|
/**
|
|
20526
20526
|
* @license
|
|
20527
20527
|
* Copyright 2019 Google LLC. All Rights Reserved.
|
|
@@ -20540,7 +20540,7 @@ const { addImpl: E$, castImpl: _$, ceilImpl: C$, concatImpl: n2, equalImpl: k$,
|
|
|
20540
20540
|
*/
|
|
20541
20541
|
class a2 {
|
|
20542
20542
|
constructor(t, e) {
|
|
20543
|
-
this.variableNames = ["source"], this.workPerThread = 1, this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = e, this.rank = e.length, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
20543
|
+
this.variableNames = ["source"], this.workPerThread = 1, this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = e, this.rank = e.length, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize, [this.workPerThread, 1, 1]), this.start = t, this.uniforms = `start : ${Tt(t.length)}, `, this.shaderKey = "slice";
|
|
20544
20544
|
}
|
|
20545
20545
|
getUserCode() {
|
|
20546
20546
|
const t = Tt(this.rank), e = l2(this.rank);
|
|
@@ -20644,7 +20644,7 @@ const h2 = "let resultTemp = a + b;", f2 = "let resultTemp = atan2(a, b);", d2 =
|
|
|
20644
20644
|
let zero = sign(a) * 0 + 0;
|
|
20645
20645
|
let one = sign(b) * 0 + 1;
|
|
20646
20646
|
let resultTemp = select(zero, one, a <= b);
|
|
20647
|
-
`,
|
|
20647
|
+
`, I2 = "return f32(a >= 1.0 && b >= 1.0);", $2 = `return (vec4<f32>(a >= vec4<f32>(1.0)) *
|
|
20648
20648
|
vec4<f32>(b >= vec4<f32>(1.0)));`, A2 = "return f32(a >= 1.0 || b >= 1.0);", E2 = `return min(vec4<f32>(a >= vec4<f32>(1.0)) +
|
|
20649
20649
|
vec4<f32>(b >= vec4<f32>(1.0)), vec4<f32>(1.0));`, _2 = "let resultTemp = max(a, b);", C2 = "let resultTemp = min(a, b);", k2 = `
|
|
20650
20650
|
let isNaN = b == 0.;
|
|
@@ -20782,7 +20782,7 @@ function z2(n, t) {
|
|
|
20782
20782
|
e = v2;
|
|
20783
20783
|
break;
|
|
20784
20784
|
case it.LOGICAL_AND:
|
|
20785
|
-
return t ?
|
|
20785
|
+
return t ? $2 : I2;
|
|
20786
20786
|
case it.LOGICAL_OR:
|
|
20787
20787
|
return t ? E2 : A2;
|
|
20788
20788
|
case it.MUL:
|
|
@@ -20880,7 +20880,7 @@ const U2 = "return abs(a);", W2 = `
|
|
|
20880
20880
|
let a2 = ${xy};
|
|
20881
20881
|
let a3 = ${Sy};
|
|
20882
20882
|
let a4 = ${vy};
|
|
20883
|
-
let a5 = ${
|
|
20883
|
+
let a5 = ${Iy};
|
|
20884
20884
|
|
|
20885
20885
|
let sign = sign(a);
|
|
20886
20886
|
let absA = abs(a);
|
|
@@ -20901,7 +20901,7 @@ const U2 = "return abs(a);", W2 = `
|
|
|
20901
20901
|
} else {
|
|
20902
20902
|
return ${gy} * (exp(a) - 1.0);
|
|
20903
20903
|
}
|
|
20904
|
-
`, Sv = "return 1.0 / (1.0 + exp(-1.0 * a));", vv = "return sign(a);",
|
|
20904
|
+
`, Sv = "return 1.0 / (1.0 + exp(-1.0 * a));", vv = "return sign(a);", Iv = "return sin(a);", $v = `
|
|
20905
20905
|
let e2x = exp(a);
|
|
20906
20906
|
return (e2x - 1.0 / e2x) / 2.0;
|
|
20907
20907
|
`, Av = `
|
|
@@ -20996,9 +20996,9 @@ function An(n, t) {
|
|
|
20996
20996
|
case F.SIGN:
|
|
20997
20997
|
return vv;
|
|
20998
20998
|
case F.SIN:
|
|
20999
|
-
return $v;
|
|
21000
|
-
case F.SINH:
|
|
21001
20999
|
return Iv;
|
|
21000
|
+
case F.SINH:
|
|
21001
|
+
return $v;
|
|
21002
21002
|
case F.SOFTPLUS:
|
|
21003
21003
|
return Av;
|
|
21004
21004
|
case F.SQRT:
|
|
@@ -21420,7 +21420,7 @@ class Mv {
|
|
|
21420
21420
|
const f = Bx(e[1], u, e[2], s);
|
|
21421
21421
|
this.workgroupSize = f.workgroupSize, this.elementsPerThread = f.elementsPerThread;
|
|
21422
21422
|
}
|
|
21423
|
-
this.dispatch =
|
|
21423
|
+
this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize, this.elementsPerThread);
|
|
21424
21424
|
const c = o != null, h = a != null;
|
|
21425
21425
|
c && this.variableNames.push("bias"), h && this.variableNames.push("preluActivationWeights"), this.sequentialAccessByThreads = l, this.transposeA = s, this.transposeB = r, this.addBias = c, this.activation = i, this.hasPreluActivationWeights = h, [this.fitAOuter, this.fitBOuter, this.fitInner] = this.getShapeFit(e[1], e[2], u), this.shaderKey = `matMulPacked_${this.elementsPerThread}_${s}_${r}_${this.activation}_${this.fitAOuter}_${this.fitBOuter}_${this.fitInner}_${this.isVec4}_${this.isVectorA}_${this.sequentialAccessByThreads}`;
|
|
21426
21426
|
}
|
|
@@ -21544,7 +21544,7 @@ function Bv(n, t, e, s, r = !1, o = null, i = !1, a = 4, l = 4, u = 4) {
|
|
|
21544
21544
|
}
|
|
21545
21545
|
class Fv {
|
|
21546
21546
|
constructor(t, e, s, r, o = !1, i = null, a = !1, l = !1) {
|
|
21547
|
-
this.variableNames = ["x", "W"], this.uniforms = "filterDims : vec2<i32>, pads : vec2<i32>, strides : vec2<i32>, dilations : vec2<i32>, dimAOuter : i32, dimBOuter : i32, dimInner : i32,", this.outputShape = t.outShape, this.isChannelsLast = t.dataFormat === "channelsLast", this.isVec4 = ((t.inChannels % 4 === 0 || t.inChannels % 3 === 0) && this.isChannelsLast || t.outWidth % 4 === 0 && !this.isChannelsLast) && t.outChannels % 4 === 0, this.dispatchLayout = this.isChannelsLast ? { x: [3], y: [1, 2], z: [0] } : { x: [2, 3], y: [1], z: [0] }, this.workgroupSize = Fx(this.dispatchLayout, this.outputShape, this.isVec4), this.elementsPerThread = zx(this.dispatchLayout, this.outputShape, this.isVec4), this.dispatch =
|
|
21547
|
+
this.variableNames = ["x", "W"], this.uniforms = "filterDims : vec2<i32>, pads : vec2<i32>, strides : vec2<i32>, dilations : vec2<i32>, dimAOuter : i32, dimBOuter : i32, dimInner : i32,", this.outputShape = t.outShape, this.isChannelsLast = t.dataFormat === "channelsLast", this.isVec4 = ((t.inChannels % 4 === 0 || t.inChannels % 3 === 0) && this.isChannelsLast || t.outWidth % 4 === 0 && !this.isChannelsLast) && t.outChannels % 4 === 0, this.dispatchLayout = this.isChannelsLast ? { x: [3], y: [1, 2], z: [0] } : { x: [2, 3], y: [1], z: [0] }, this.workgroupSize = Fx(this.dispatchLayout, this.outputShape, this.isVec4), this.elementsPerThread = zx(this.dispatchLayout, this.outputShape, this.isVec4), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize, this.elementsPerThread), this.isVec4 ? (this.outputComponent = 4, this.isChannelsLast && t.inChannels % 4 !== 0 ? (this.innerElementSize = 3, this.variableComponents = [1, 4]) : (this.innerElementSize = 4, this.variableComponents = [4, 4]), o && (this.variableNames.push("bias"), this.variableComponents.push(4)), a && (this.variableNames.push("preluActivationWeights"), this.variableComponents.push(4))) : (this.innerElementSize = this.elementsPerThread[0], o && this.variableNames.push("bias"), a && this.variableNames.push("preluActivationWeights")), this.sequentialAccessByThreads = l, this.addBias = o, this.activation = i, this.hasPreluActivationWeights = a, this.tileAOuter = this.workgroupSize[1] * this.elementsPerThread[1], this.tileBOuter = this.workgroupSize[0] * this.elementsPerThread[0], this.tileInner = Math.max(this.workgroupSize[0] * this.innerElementSize, this.workgroupSize[1]), this.fitAOuter = e % this.tileAOuter === 0, this.fitBOuter = s % this.tileBOuter === 0, this.fitInner = r % this.tileInner === 0, this.shaderKey = `conv2DMM_${this.elementsPerThread}_${this.activation}}_${this.fitAOuter}_${this.fitBOuter}_${this.fitInner}_${this.isVec4}_${this.innerElementSize}_${this.isChannelsLast}_${this.sequentialAccessByThreads}`;
|
|
21548
21548
|
}
|
|
21549
21549
|
getUserCode() {
|
|
21550
21550
|
const t = this.isVec4 ? ha(this.elementsPerThread, this.workgroupSize, !this.isChannelsLast, this.tileInner) : fa(this.elementsPerThread, this.workgroupSize, !this.isChannelsLast, this.tileInner, !1, null, this.sequentialAccessByThreads), e = this.isVec4 ? [this.innerElementSize, 4, 4] : [1, 1, 1];
|
|
@@ -21572,7 +21572,7 @@ class Fv {
|
|
|
21572
21572
|
*/
|
|
21573
21573
|
class zv {
|
|
21574
21574
|
constructor(t, e = !1, s = null, r = !1) {
|
|
21575
|
-
this.variableNames = ["x", "W"], this.uniforms = "filterDims: vec2<i32>, pads: vec2<i32>, strides: vec2<i32>, dilations: vec2<i32>,", this.workgroupSize = [4, 4, 8], this.outputShape = t.outShape, this.isChannelsLast = t.dataFormat === "channelsLast", this.dispatchLayout = this.isChannelsLast ? { x: [2], y: [1], z: [0, 3] } : { x: [3], y: [2], z: [0, 1] }, this.dispatch =
|
|
21575
|
+
this.variableNames = ["x", "W"], this.uniforms = "filterDims: vec2<i32>, pads: vec2<i32>, strides: vec2<i32>, dilations: vec2<i32>,", this.workgroupSize = [4, 4, 8], this.outputShape = t.outShape, this.isChannelsLast = t.dataFormat === "channelsLast", this.dispatchLayout = this.isChannelsLast ? { x: [2], y: [1], z: [0, 3] } : { x: [3], y: [2], z: [0, 1] }, this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.addBias = e, this.activation = s, this.hasPreluActivationWeights = r, e && this.variableNames.push("bias"), r && this.variableNames.push("preluActivationWeights"), this.shaderKey = `conv2dnaive_${this.activation}_${this.isChannelsLast}`;
|
|
21576
21576
|
}
|
|
21577
21577
|
getUserCode() {
|
|
21578
21578
|
return `
|
|
@@ -21643,7 +21643,7 @@ class zv {
|
|
|
21643
21643
|
class Uv {
|
|
21644
21644
|
constructor(t, e) {
|
|
21645
21645
|
this.variableNames = ["x"], this.uniforms = `pads : vec2<i32>, strides : vec2<i32>, dilations : vec2<i32>, outWidth : i32, itemsPerBlockRow : i32,
|
|
21646
|
-
inChannels : i32,`, this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = t, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
21646
|
+
inChannels : i32,`, this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = t, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.isChannelsLast = e, this.shaderKey = `im2col_${this.isChannelsLast}`;
|
|
21647
21647
|
}
|
|
21648
21648
|
getUserCode() {
|
|
21649
21649
|
const t = this.isChannelsLast ? 1 : 2, e = this.isChannelsLast ? 2 : 3, s = this.isChannelsLast ? "coords[1]" : "coords[2]", r = this.isChannelsLast ? "coords[2]" : "coords[1]", o = this.isChannelsLast ? "getX(batch, xRow, xCol, ch)" : "getX(batch, ch, xRow, xCol)";
|
|
@@ -21727,7 +21727,7 @@ function Wv(n) {
|
|
|
21727
21727
|
}
|
|
21728
21728
|
class Gv {
|
|
21729
21729
|
constructor(t, e = !1, s = !1, r = null, o = null, i = null) {
|
|
21730
|
-
this.variableNames = ["A", "B"], this.uniforms = "dimAOuter : i32, dimBOuter : i32, dimInner : i32,", this.workgroupSize = [256, 1, 1], this.outputShape = t, this.dispatchLayout = { x: [], y: [1, 2], z: [0] }, this.dispatch =
|
|
21730
|
+
this.variableNames = ["A", "B"], this.uniforms = "dimAOuter : i32, dimBOuter : i32, dimInner : i32,", this.workgroupSize = [256, 1, 1], this.outputShape = t, this.dispatchLayout = { x: [], y: [1, 2], z: [0] }, this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize);
|
|
21731
21731
|
const a = r != null, l = i != null;
|
|
21732
21732
|
a && this.variableNames.push("bias"), l && this.variableNames.push("preluActivationWeights"), this.transposeA = e, this.transposeB = s, this.addBias = a, this.activation = o, this.hasPreluActivationWeights = l, this.shaderKey = `matMulReduce_${this.activation}_${e}_${s}`;
|
|
21733
21733
|
}
|
|
@@ -21851,7 +21851,7 @@ class jv {
|
|
|
21851
21851
|
constructor(t, e, s = !1, r = !1) {
|
|
21852
21852
|
this.variableNames = ["A", "B"], this.uniforms = "dimAOuter : i32, dimBOuter : i32, dimInner : i32,", this.workgroupSize = [8, 8, 1], this.atomic = !0, this.splitedDimInner = 128, w(t[0] === 1, () => "MatMulSplitKProgram only supports batch = 1."), this.outputShape = t, this.dispatchLayout = { x: [2], y: [1], z: [0, 3] };
|
|
21853
21853
|
const o = (s && this.outputShape[1] % 4 === 0 || !s && e % 4 === 0) && this.outputShape[2] % 4 === 0;
|
|
21854
|
-
this.elementsPerThread = [4, 4, this.splitedDimInner], this.outputComponent = o ? 4 : 1, o || (this.outputShape[1] < 16 && (this.elementsPerThread[1] = 1), this.outputShape[2] < 16 && (this.elementsPerThread[0] = 1)), this.dispatch =
|
|
21854
|
+
this.elementsPerThread = [4, 4, this.splitedDimInner], this.outputComponent = o ? 4 : 1, o || (this.outputShape[1] < 16 && (this.elementsPerThread[1] = 1), this.outputShape[2] < 16 && (this.elementsPerThread[0] = 1)), this.dispatch = $t(this.dispatchLayout, [
|
|
21855
21855
|
this.outputShape[0],
|
|
21856
21856
|
this.outputShape[1],
|
|
21857
21857
|
this.outputShape[2],
|
|
@@ -21879,7 +21879,7 @@ class jv {
|
|
|
21879
21879
|
}
|
|
21880
21880
|
class Hv {
|
|
21881
21881
|
constructor(t, e = null, s = null, r = null) {
|
|
21882
|
-
this.uniforms = "", this.variableNames = ["x"], this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = t, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
21882
|
+
this.uniforms = "", this.variableNames = ["x"], this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = t, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.addBias = e != null, this.hasPreluActivationWeights = r != null, this.activation = s, this.addBias && this.variableNames.push("bias"), this.hasPreluActivationWeights && this.variableNames.push("preluActivationWeights"), this.shaderKey = `biasActivation_${s}`;
|
|
21883
21883
|
}
|
|
21884
21884
|
getUserCode() {
|
|
21885
21885
|
return `
|
|
@@ -22218,7 +22218,7 @@ const Zv = {
|
|
|
22218
22218
|
*/
|
|
22219
22219
|
class Qv {
|
|
22220
22220
|
constructor(t) {
|
|
22221
|
-
this.variableNames = ["x"], this.uniforms = "strides : vec2<i32>,", this.workgroupSize = [256, 1, 1], this.size = !0, this.outputShape = t.outShape, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
22221
|
+
this.variableNames = ["x"], this.uniforms = "strides : vec2<i32>,", this.workgroupSize = [256, 1, 1], this.size = !0, this.outputShape = t.outShape, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.shaderKey = "poolWithFilterSizeEqualsOne";
|
|
22222
22222
|
}
|
|
22223
22223
|
getUserCode() {
|
|
22224
22224
|
return `
|
|
@@ -22255,11 +22255,11 @@ class Qv {
|
|
|
22255
22255
|
* limitations under the License.
|
|
22256
22256
|
* =============================================================================
|
|
22257
22257
|
*/
|
|
22258
|
-
class
|
|
22258
|
+
class tI {
|
|
22259
22259
|
constructor(t, e, s = !1, r = !1, o = !1) {
|
|
22260
22260
|
if (this.variableNames = ["x"], this.uniforms = "strides : vec2<i32>, pads : vec2<i32>, dilations : vec2<i32>, convDims : vec2<i32>, filterDims : vec2<i32>,", this.workgroupSize = [128, 1, 1], this.size = !0, e === "avg" && s)
|
|
22261
22261
|
throw new Error("Cannot compute positions for average pool.");
|
|
22262
|
-
this.outputShape = t.outShape, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
22262
|
+
this.outputShape = t.outShape, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.poolType = e, this.computePositions = s, this.flattenPositions = r, this.includeBatchIndex = o, this.shaderKey = `pool2D_${e}_${s}_${r}_${o}`;
|
|
22263
22263
|
}
|
|
22264
22264
|
getUserCode() {
|
|
22265
22265
|
let t;
|
|
@@ -22325,13 +22325,13 @@ class t$ {
|
|
|
22325
22325
|
* limitations under the License.
|
|
22326
22326
|
* =============================================================================
|
|
22327
22327
|
*/
|
|
22328
|
-
class
|
|
22328
|
+
class eI {
|
|
22329
22329
|
constructor(t, e) {
|
|
22330
22330
|
this.variableNames = ["A"], this.workgroupSize = [16, 16, 1];
|
|
22331
22331
|
const s = new Array(t.length);
|
|
22332
22332
|
for (let r = 0; r < s.length; r++)
|
|
22333
22333
|
s[r] = t[e[r]];
|
|
22334
|
-
this.outputShape = s, this.dispatchLayout = { x: [0], y: [1] }, this.dispatch =
|
|
22334
|
+
this.outputShape = s, this.dispatchLayout = { x: [0], y: [1] }, this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize, [1, 1, 1]), this.shaderKey = "transposeShared";
|
|
22335
22335
|
}
|
|
22336
22336
|
getUserCode() {
|
|
22337
22337
|
w(this.workgroupSize[0] === this.workgroupSize[1], () => `Must be a square tile, current tile shape is ${this.workgroupSize[0]} x ${this.workgroupSize[1]}`);
|
|
@@ -22374,16 +22374,16 @@ class e$ {
|
|
|
22374
22374
|
* limitations under the License.
|
|
22375
22375
|
* =============================================================================
|
|
22376
22376
|
*/
|
|
22377
|
-
class
|
|
22377
|
+
class nI {
|
|
22378
22378
|
constructor(t, e) {
|
|
22379
22379
|
this.variableNames = ["A"], this.workPerThread = 1, this.workgroupSize = [64, 1, 1], this.size = !0;
|
|
22380
22380
|
const s = new Array(t.length);
|
|
22381
22381
|
for (let r = 0; r < s.length; r++)
|
|
22382
22382
|
s[r] = t[e[r]];
|
|
22383
|
-
this.outputShape = s, this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
22383
|
+
this.outputShape = s, this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize, [this.workPerThread, 1, 1]), this.newDim = e, this.shaderKey = `transpose_${e}`;
|
|
22384
22384
|
}
|
|
22385
22385
|
getUserCode() {
|
|
22386
|
-
const t = Tt(this.outputShape.length), e =
|
|
22386
|
+
const t = Tt(this.outputShape.length), e = sI(this.newDim);
|
|
22387
22387
|
return `
|
|
22388
22388
|
${wt("index")} {
|
|
22389
22389
|
for(var i = 0; i < ${this.workPerThread}; i = i + 1) {
|
|
@@ -22398,7 +22398,7 @@ class n$ {
|
|
|
22398
22398
|
`;
|
|
22399
22399
|
}
|
|
22400
22400
|
}
|
|
22401
|
-
function
|
|
22401
|
+
function sI(n) {
|
|
22402
22402
|
const t = n.length;
|
|
22403
22403
|
if (t > 6)
|
|
22404
22404
|
throw Error(`Transpose for rank ${t} is not yet supported`);
|
|
@@ -22423,7 +22423,7 @@ function s$(n) {
|
|
|
22423
22423
|
* limitations under the License.
|
|
22424
22424
|
* =============================================================================
|
|
22425
22425
|
*/
|
|
22426
|
-
function
|
|
22426
|
+
function rI(n) {
|
|
22427
22427
|
const { inputs: t, backend: e, attrs: s } = n, { x: r } = t, { perm: o } = s, i = e, a = r.shape.length, l = new Array(a);
|
|
22428
22428
|
for (let c = 0; c < l.length; c++)
|
|
22429
22429
|
l[c] = r.shape[o[c]];
|
|
@@ -22432,10 +22432,10 @@ function r$(n) {
|
|
|
22432
22432
|
return e.makeTensorInfo(l, r.dtype, f);
|
|
22433
22433
|
}
|
|
22434
22434
|
if (r.shape.length === 2 && oe(o, [1, 0])) {
|
|
22435
|
-
const c = new
|
|
22435
|
+
const c = new eI(r.shape, o);
|
|
22436
22436
|
return i.runWebGPUProgram(c, [r], r.dtype);
|
|
22437
22437
|
}
|
|
22438
|
-
const u = new
|
|
22438
|
+
const u = new nI(r.shape, o);
|
|
22439
22439
|
return i.runWebGPUProgram(u, [r], r.dtype);
|
|
22440
22440
|
}
|
|
22441
22441
|
/**
|
|
@@ -22454,11 +22454,11 @@ function r$(n) {
|
|
|
22454
22454
|
* limitations under the License.
|
|
22455
22455
|
* =============================================================================
|
|
22456
22456
|
*/
|
|
22457
|
-
class
|
|
22457
|
+
class oI {
|
|
22458
22458
|
constructor(t, e, s) {
|
|
22459
22459
|
this.variableNames = ["x"], this.uniforms = "reduceSize : i32,", this.size = !0, this.inputShape = [t.batchSize, t.inSize];
|
|
22460
22460
|
const [r] = bi(this.inputShape, [1]);
|
|
22461
|
-
this.outputShape = r.length === 0 ? [1] : r, t.inSize >= 32768 && s >= 512 ? this.workgroupSize = [512, 1, 1] : t.inSize >= 4096 ? this.workgroupSize = [256, 1, 1] : this.workgroupSize = [64, 1, 1], this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
22461
|
+
this.outputShape = r.length === 0 ? [1] : r, t.inSize >= 32768 && s >= 512 ? this.workgroupSize = [512, 1, 1] : t.inSize >= 4096 ? this.workgroupSize = [256, 1, 1] : this.workgroupSize = [64, 1, 1], this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, [1, 1, 1]), this.reduceType = e, this.shaderKey = `reduce_${e}`;
|
|
22462
22462
|
}
|
|
22463
22463
|
getUserCode() {
|
|
22464
22464
|
let t = "", e = "0.0";
|
|
@@ -22535,17 +22535,17 @@ class o$ {
|
|
|
22535
22535
|
* limitations under the License.
|
|
22536
22536
|
* =============================================================================
|
|
22537
22537
|
*/
|
|
22538
|
-
const
|
|
22538
|
+
const iI = {
|
|
22539
22539
|
mean: "float32",
|
|
22540
22540
|
all: "bool",
|
|
22541
22541
|
any: "bool"
|
|
22542
22542
|
};
|
|
22543
|
-
function
|
|
22543
|
+
function aI(n, t, e, s, r) {
|
|
22544
22544
|
const o = n.shape.length, i = [], a = _s(t, n.shape);
|
|
22545
22545
|
let l = a;
|
|
22546
22546
|
const u = gg(l, o);
|
|
22547
22547
|
let c = n;
|
|
22548
|
-
u != null && (c =
|
|
22548
|
+
u != null && (c = rI({ inputs: { x: n }, attrs: { perm: u }, backend: r }), l = bg(l.length, o), i.push(c)), mg(s, l, o);
|
|
22549
22549
|
const [h, f] = bi(c.shape, l);
|
|
22550
22550
|
let d = h;
|
|
22551
22551
|
e && (d = _u(h, a));
|
|
@@ -22565,9 +22565,9 @@ function a$(n, t, e, s, r) {
|
|
|
22565
22565
|
throw new Error(`${s} CPU implementation is not yet supported.`);
|
|
22566
22566
|
}
|
|
22567
22567
|
} else {
|
|
22568
|
-
const g = z(f), b = z(c.shape) / g, y = { windowSize: g, inSize: g, batchSize: b, outSize: 1 }, S =
|
|
22568
|
+
const g = z(f), b = z(c.shape) / g, y = { windowSize: g, inSize: g, batchSize: b, outSize: 1 }, S = iI[s] || Dp(n.dtype), x = [
|
|
22569
22569
|
{ type: "int32", data: [g] }
|
|
22570
|
-
], v = new
|
|
22570
|
+
], v = new oI(y, s, r.device.limits.maxComputeWorkgroupSizeX), E = r.runWebGPUProgram(v, [c], S, x);
|
|
22571
22571
|
i.push(E), p = dt({ inputs: { x: E }, attrs: { shape: d }, backend: r });
|
|
22572
22572
|
}
|
|
22573
22573
|
return i.forEach((g) => r.disposeData(g.dataId)), p;
|
|
@@ -22588,9 +22588,9 @@ function a$(n, t, e, s, r) {
|
|
|
22588
22588
|
* limitations under the License.
|
|
22589
22589
|
* =============================================================================
|
|
22590
22590
|
*/
|
|
22591
|
-
function
|
|
22591
|
+
function lI(n) {
|
|
22592
22592
|
const { inputs: t, backend: e, attrs: s } = n, { x: r } = t, { reductionIndices: o, keepDims: i } = s;
|
|
22593
|
-
return
|
|
22593
|
+
return aI(r, o, i, "max", e);
|
|
22594
22594
|
}
|
|
22595
22595
|
/**
|
|
22596
22596
|
* @license
|
|
@@ -22608,7 +22608,7 @@ function l$(n) {
|
|
|
22608
22608
|
* limitations under the License.
|
|
22609
22609
|
* =============================================================================
|
|
22610
22610
|
*/
|
|
22611
|
-
function
|
|
22611
|
+
function uI(n, t, e, s) {
|
|
22612
22612
|
if (t.filterWidth === 1 && t.filterHeight === 1 && oe(t.inShape, t.outShape))
|
|
22613
22613
|
return He({ inputs: { x: n }, backend: s });
|
|
22614
22614
|
if (t.filterWidth === t.inWidth && t.filterHeight === t.inHeight && t.batchSize === 1 && t.padInfo.type === "VALID") {
|
|
@@ -22624,7 +22624,7 @@ function u$(n, t, e, s) {
|
|
|
22624
22624
|
}
|
|
22625
22625
|
});
|
|
22626
22626
|
let l;
|
|
22627
|
-
w(e === "max", () => `Invalid pool type ${e}`), l =
|
|
22627
|
+
w(e === "max", () => `Invalid pool type ${e}`), l = lI({
|
|
22628
22628
|
inputs: { x: a },
|
|
22629
22629
|
backend: s,
|
|
22630
22630
|
attrs: { reductionIndices: 0, keepDims: !1 }
|
|
@@ -22634,7 +22634,7 @@ function u$(n, t, e, s) {
|
|
|
22634
22634
|
}
|
|
22635
22635
|
let r;
|
|
22636
22636
|
const o = [{ type: "int32", data: [t.strideHeight, t.strideWidth] }];
|
|
22637
|
-
return t.filterHeight === 1 && t.filterWidth === 1 ? r = new Qv(t) : (w(e === "max", () => `Invalid pool type ${e}`), r = new
|
|
22637
|
+
return t.filterHeight === 1 && t.filterWidth === 1 ? r = new Qv(t) : (w(e === "max", () => `Invalid pool type ${e}`), r = new tI(t, "max"), o.push({ type: "int32", data: [t.padInfo.top, t.padInfo.left] }, {
|
|
22638
22638
|
type: "int32",
|
|
22639
22639
|
data: [t.dilationHeight, t.dilationWidth]
|
|
22640
22640
|
}, { type: "int32", data: [t.inHeight, t.inWidth] }, {
|
|
@@ -22658,14 +22658,14 @@ function u$(n, t, e, s) {
|
|
|
22658
22658
|
* limitations under the License.
|
|
22659
22659
|
* =============================================================================
|
|
22660
22660
|
*/
|
|
22661
|
-
function
|
|
22661
|
+
function cI(n) {
|
|
22662
22662
|
const { inputs: t, backend: e, attrs: s } = n, { x: r } = t, { filterSize: o, strides: i, pad: a, dimRoundingMode: l } = s, c = km(r.shape, o, i, 1, a, l);
|
|
22663
|
-
return
|
|
22663
|
+
return uI(r, c, "max", e);
|
|
22664
22664
|
}
|
|
22665
|
-
const
|
|
22665
|
+
const hI = {
|
|
22666
22666
|
kernelName: Hl,
|
|
22667
22667
|
backendName: "webgpu",
|
|
22668
|
-
kernelFunc:
|
|
22668
|
+
kernelFunc: cI
|
|
22669
22669
|
};
|
|
22670
22670
|
/**
|
|
22671
22671
|
* @license
|
|
@@ -22683,9 +22683,9 @@ const h$ = {
|
|
|
22683
22683
|
* limitations under the License.
|
|
22684
22684
|
* =============================================================================
|
|
22685
22685
|
*/
|
|
22686
|
-
class
|
|
22686
|
+
class fI {
|
|
22687
22687
|
constructor(t, e, s, r) {
|
|
22688
|
-
this.variableNames = ["x"], this.uniforms = "adjustHeightWidth : vec2<f32>, roundBase : f32,", this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = [t[0], e, s, t[3]], this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
22688
|
+
this.variableNames = ["x"], this.uniforms = "adjustHeightWidth : vec2<f32>, roundBase : f32,", this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = [t[0], e, s, t[3]], this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize), this.halfPixelCenters = r, this.shaderKey = `resizeNearest_${r}`;
|
|
22689
22689
|
}
|
|
22690
22690
|
getUserCode() {
|
|
22691
22691
|
let t;
|
|
@@ -22739,17 +22739,17 @@ class f$ {
|
|
|
22739
22739
|
* limitations under the License.
|
|
22740
22740
|
* =============================================================================
|
|
22741
22741
|
*/
|
|
22742
|
-
function
|
|
22742
|
+
function dI(n) {
|
|
22743
22743
|
const { inputs: t, backend: e, attrs: s } = n, { images: r } = t, { alignCorners: o, halfPixelCenters: i, size: a } = s, [l, u] = a, c = o && l > 1 ? 1 : 0, h = o && u > 1 ? 1 : 0, d = [
|
|
22744
22744
|
{ type: "float32", data: [c, h] },
|
|
22745
22745
|
{ type: "float32", data: [o ? 0.5 : 0] }
|
|
22746
|
-
], p = new
|
|
22746
|
+
], p = new fI(r.shape, l, u, i);
|
|
22747
22747
|
return e.runWebGPUProgram(p, [r], r.dtype, d);
|
|
22748
22748
|
}
|
|
22749
|
-
const
|
|
22749
|
+
const pI = {
|
|
22750
22750
|
kernelName: Yl,
|
|
22751
22751
|
backendName: "webgpu",
|
|
22752
|
-
kernelFunc:
|
|
22752
|
+
kernelFunc: dI
|
|
22753
22753
|
};
|
|
22754
22754
|
/**
|
|
22755
22755
|
* @license
|
|
@@ -22767,13 +22767,13 @@ const p$ = {
|
|
|
22767
22767
|
* limitations under the License.
|
|
22768
22768
|
* =============================================================================
|
|
22769
22769
|
*/
|
|
22770
|
-
class
|
|
22770
|
+
class mI {
|
|
22771
22771
|
constructor(t) {
|
|
22772
22772
|
this.uniforms = "", this.workPerThread = 1, this.workgroupSize = [64, 1, 1], this.size = !0, this.outputShape = Ss(
|
|
22773
22773
|
t,
|
|
22774
22774
|
1
|
|
22775
22775
|
/* axis */
|
|
22776
|
-
), this.variableNames = t.map((e, s) => `T${s}`), this.dispatchLayout = ae(this.outputShape), this.dispatch =
|
|
22776
|
+
), this.variableNames = t.map((e, s) => `T${s}`), this.dispatchLayout = ae(this.outputShape), this.dispatch = $t(this.dispatchLayout, this.outputShape, this.workgroupSize, [this.workPerThread, 1, 1]), this.offsetLength = t.length - 1;
|
|
22777
22777
|
for (let e = 0; e < this.offsetLength; e++)
|
|
22778
22778
|
this.uniforms += `offset${e} : i32,`;
|
|
22779
22779
|
this.shaderKey = "concat";
|
|
@@ -22821,7 +22821,7 @@ class m$ {
|
|
|
22821
22821
|
* limitations under the License.
|
|
22822
22822
|
* =============================================================================
|
|
22823
22823
|
*/
|
|
22824
|
-
function
|
|
22824
|
+
function gI(n) {
|
|
22825
22825
|
const { inputs: t, backend: e } = n, { real: s, imag: r } = t, o = e.makeTensorInfo(s.shape, "complex64"), i = e.tensorMap.get(o.dataId), a = He({ inputs: { x: s }, backend: e }), l = He({ inputs: { x: r }, backend: e });
|
|
22826
22826
|
return i.complexTensorInfos = { real: a, imag: l }, o;
|
|
22827
22827
|
}
|
|
@@ -22841,7 +22841,7 @@ function g$(n) {
|
|
|
22841
22841
|
* limitations under the License.
|
|
22842
22842
|
* =============================================================================
|
|
22843
22843
|
*/
|
|
22844
|
-
function
|
|
22844
|
+
function bI(n) {
|
|
22845
22845
|
const { inputs: t, backend: e } = n, { input: s } = t, r = e.tensorMap.get(s.dataId);
|
|
22846
22846
|
return He({ inputs: { x: r.complexTensorInfos.imag }, backend: e });
|
|
22847
22847
|
}
|
|
@@ -22861,7 +22861,7 @@ function b$(n) {
|
|
|
22861
22861
|
* limitations under the License.
|
|
22862
22862
|
* =============================================================================
|
|
22863
22863
|
*/
|
|
22864
|
-
function
|
|
22864
|
+
function yI(n) {
|
|
22865
22865
|
const { inputs: t, backend: e } = n, { input: s } = t, r = e.tensorMap.get(s.dataId);
|
|
22866
22866
|
return He({ inputs: { x: r.complexTensorInfos.real }, backend: e });
|
|
22867
22867
|
}
|
|
@@ -22884,7 +22884,7 @@ function y$(n) {
|
|
|
22884
22884
|
function ds(n, t, e) {
|
|
22885
22885
|
const s = n[0].dtype;
|
|
22886
22886
|
if (s === "complex64") {
|
|
22887
|
-
const p = n.map((S) =>
|
|
22887
|
+
const p = n.map((S) => yI({ inputs: { input: S }, backend: e })), g = n.map((S) => bI({ inputs: { input: S }, backend: e })), m = ds(p, t, e), b = ds(g, t, e), y = gI({ inputs: { real: m, imag: b }, backend: e });
|
|
22888
22888
|
return p.forEach((S) => e.disposeData(S.dataId)), g.forEach((S) => e.disposeData(S.dataId)), e.disposeData(m.dataId), e.disposeData(b.dataId), y;
|
|
22889
22889
|
}
|
|
22890
22890
|
let r = e.shouldExecuteOnCPU(n);
|
|
@@ -22911,7 +22911,7 @@ function ds(n, t, e) {
|
|
|
22911
22911
|
e.disposeData(m.dataId);
|
|
22912
22912
|
return g;
|
|
22913
22913
|
}
|
|
22914
|
-
const { tensors2D: i, outShape: a } =
|
|
22914
|
+
const { tensors2D: i, outShape: a } = wI(n, t, e), l = i.map((p) => p.shape), u = new mI(l), c = [], h = new Array(l.length - 1);
|
|
22915
22915
|
if (h.length > 0) {
|
|
22916
22916
|
h[0] = l[0][1], c.push({ type: "int32", data: [h[0]] });
|
|
22917
22917
|
for (let p = 1; p < h.length; p++)
|
|
@@ -22922,7 +22922,7 @@ function ds(n, t, e) {
|
|
|
22922
22922
|
const d = dt({ inputs: { x: f }, backend: e, attrs: { shape: a } });
|
|
22923
22923
|
return e.disposeData(f.dataId), d;
|
|
22924
22924
|
}
|
|
22925
|
-
function
|
|
22925
|
+
function wI(n, t, e) {
|
|
22926
22926
|
const s = Ss(n.map((o) => o.shape), t);
|
|
22927
22927
|
return { tensors2D: n.map((o) => dt({
|
|
22928
22928
|
inputs: { x: o },
|
|
@@ -22951,7 +22951,7 @@ function w$(n, t, e) {
|
|
|
22951
22951
|
* limitations under the License.
|
|
22952
22952
|
* =============================================================================
|
|
22953
22953
|
*/
|
|
22954
|
-
function
|
|
22954
|
+
function xI(n) {
|
|
22955
22955
|
const { inputs: t, backend: e, attrs: s } = n, { axis: r } = s, o = _s(r, t[0].shape)[0], i = t.map((u) => u.shape);
|
|
22956
22956
|
hy(i, o);
|
|
22957
22957
|
const a = Ss(t.map((u) => u.shape), o);
|
|
@@ -22960,26 +22960,26 @@ function x$(n) {
|
|
|
22960
22960
|
const l = t.filter((u) => z(u.shape) > 0);
|
|
22961
22961
|
return l.length === 1 ? He({ inputs: { x: l[0] }, backend: e }) : ds(l, o, e);
|
|
22962
22962
|
}
|
|
22963
|
-
const
|
|
22963
|
+
const SI = {
|
|
22964
22964
|
kernelName: jl,
|
|
22965
22965
|
backendName: "webgpu",
|
|
22966
|
-
kernelFunc:
|
|
22967
|
-
},
|
|
22966
|
+
kernelFunc: xI
|
|
22967
|
+
}, vI = [
|
|
22968
22968
|
Vx,
|
|
22969
22969
|
Xx,
|
|
22970
22970
|
c2,
|
|
22971
22971
|
Zv,
|
|
22972
|
-
|
|
22973
|
-
|
|
22974
|
-
|
|
22972
|
+
hI,
|
|
22973
|
+
pI,
|
|
22974
|
+
SI,
|
|
22975
22975
|
qx
|
|
22976
22976
|
];
|
|
22977
|
-
for (const n of
|
|
22977
|
+
for (const n of vI)
|
|
22978
22978
|
up({
|
|
22979
22979
|
...n,
|
|
22980
22980
|
backendName: "webgpu-oidn"
|
|
22981
22981
|
});
|
|
22982
|
-
async function
|
|
22982
|
+
async function II() {
|
|
22983
22983
|
var n;
|
|
22984
22984
|
try {
|
|
22985
22985
|
const t = {
|
|
@@ -23008,19 +23008,19 @@ async function hf(n, t) {
|
|
|
23008
23008
|
let e = A.findBackend("webgpu-oidn");
|
|
23009
23009
|
return e != null || (e = new Fs(n, t), A.registerBackend("webgpu-oidn", () => e), await A.setBackend("webgpu-oidn")), e;
|
|
23010
23010
|
}
|
|
23011
|
-
async function I
|
|
23011
|
+
async function $I(n, t, e) {
|
|
23012
23012
|
const s = await (t ? hf(
|
|
23013
23013
|
t.device,
|
|
23014
23014
|
t.adapterInfo
|
|
23015
|
-
) :
|
|
23015
|
+
) : II()), r = xf(n);
|
|
23016
23016
|
return new xx(r, s, e);
|
|
23017
23017
|
}
|
|
23018
|
-
async function
|
|
23019
|
-
return fetch(n).then((s) => s.arrayBuffer()).then((s) => I
|
|
23018
|
+
async function n$(n, t, e) {
|
|
23019
|
+
return fetch(n).then((s) => s.arrayBuffer()).then((s) => $I(s, t, e));
|
|
23020
23020
|
}
|
|
23021
23021
|
export {
|
|
23022
23022
|
xx as UNet,
|
|
23023
|
-
I
|
|
23024
|
-
|
|
23023
|
+
$I as initUNetFromBuffer,
|
|
23024
|
+
n$ as initUNetFromURL,
|
|
23025
23025
|
xf as parseTZA
|
|
23026
23026
|
};
|