oidn-web 0.1.0 → 0.1.1
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.mjs +602 -605
- package/dist/oidn.umd.js +137 -137
- package/lib/UNet.js +5 -2
- package/lib/UNet.js.map +1 -1
- package/package.json +1 -1
- package/src/UNet.ts +14 -14
- package/weights/rt_hdr_calb_cnrm.tza +0 -0
package/dist/oidn.umd.js
CHANGED
|
@@ -1,4 +1,4 @@
|
|
|
1
|
-
(function(Nt,
|
|
1
|
+
(function(Nt,Ht){typeof exports=="object"&&typeof module<"u"?Ht(exports):typeof define=="function"&&define.amd?define(["exports"],Ht):(Nt=typeof globalThis<"u"?globalThis:Nt||self,Ht(Nt.oidn={}))})(this,function(Nt){"use strict";var xv=Object.defineProperty;var Sv=(Nt,Ht,Je)=>Ht in Nt?xv(Nt,Ht,{enumerable:!0,configurable:!0,writable:!0,value:Je}):Nt[Ht]=Je;var K=(Nt,Ht,Je)=>(Sv(Nt,typeof Ht!="symbol"?Ht+"":Ht,Je),Je);function Ht(n,t){for(var e=0;e<t.length;e++){const s=t[e];if(typeof s!="string"&&!Array.isArray(s)){for(const r in s)if(r!=="default"&&!(r in n)){const o=Object.getOwnPropertyDescriptor(s,r);o&&Object.defineProperty(n,r,o.get?o:{enumerable:!0,get:()=>s[r]})}}}return Object.freeze(Object.defineProperty(n,Symbol.toStringTag,{value:"Module"}))}class Je{constructor(){K(this,"dims",[]);K(this,"paddedDims",[]);K(this,"layout","x");K(this,"dataType","Float32")}getByteSize(){let t=1;for(const e of this.paddedDims)t*=e;return this.dataType==="Float32"?t*=4:this.dataType==="Float16"&&(t*=2),t}}class wf{constructor(t,e){this.desc=t,this.data=e}}class xf{constructor(t){K(this,"offset",0);this._view=t}read(t){const e=this._view,s=this.offset;switch(this.offset+=t,t){case 1:return e.getUint8(s);case 2:return e.getUint16(s,!0);case 4:return e.getUint32(s,!0);case 8:return Number(e.getBigUint64(s,!0));default:throw new Error("unsupported read size")}}}function ga(n){const t=new Uint8Array(n),e=new xf(new DataView(n));if(e.read(2)!==16855)throw new Error("invalid or corrupted weights blob");const r=e.read(1);if(e.read(1),r!==2)throw new Error("unsupported weights blob version");const o=e.read(8);e.offset=o;const i=e.read(4),a=new Map;for(let l=0;l<i;++l){const u=new Je,c=e.read(2),h=new TextDecoder().decode(t.subarray(e.offset,e.offset+c));e.offset+=c;const f=e.read(1);for(let b=0;b<f;++b)u.dims.push(e.read(4));u.paddedDims=[...u.dims],new TextDecoder().decode(t.subarray(e.offset,e.offset+f))==="oihw"&&(u.layout="oihw"),e.offset+=f;const p=String.fromCharCode(e.read(1));if(p==="f")u.dataType="Float32";else if(p==="h")u.dataType="Float16";else throw new Error("invalid tensor data type");const g=e.read(8),m=t.slice(g,g+u.getByteSize());a.set(h,new wf(u,m))}return a}/**
|
|
2
2
|
* @license
|
|
3
3
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
4
4
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -28,7 +28,7 @@
|
|
|
28
28
|
* See the License for the specific language governing permissions and
|
|
29
29
|
* limitations under the License.
|
|
30
30
|
* =============================================================================
|
|
31
|
-
*/function If(n){let t=n.length,e=0;for(;t>0;)e=Math.random()*t|0,t--,En(n,t,e)}function En(n,t,e){const s=n[t];n[t]=n[e],n[e]=s}function Af(n){let t=0;for(let e=0;e<n.length;e++)t+=n[e];return t}function w(n,t){if(!n)throw new Error(typeof t=="string"?t:t())}function Ef(n,t,e=""){w(Zt(n,t),()=>e+` Shapes ${n} and ${t} must match`)}function ya(n){w(n!=null,()=>"The input to the tensor constructor must be a non-null value.")}function z(n){if(n.length===0)return 1;let t=n[0];for(let e=1;e<n.length;e++)t*=n[e];return t}function Zt(n,t){if(n===t)return!0;if(n==null||t==null||n.length!==t.length)return!1;for(let e=0;e<n.length;e++)if(n[e]!==t[e])return!1;return!0}function oo(n){return n%1===0}function
|
|
31
|
+
*/function If(n){let t=n.length,e=0;for(;t>0;)e=Math.random()*t|0,t--,En(n,t,e)}function En(n,t,e){const s=n[t];n[t]=n[e],n[e]=s}function Af(n){let t=0;for(let e=0;e<n.length;e++)t+=n[e];return t}function w(n,t){if(!n)throw new Error(typeof t=="string"?t:t())}function Ef(n,t,e=""){w(Zt(n,t),()=>e+` Shapes ${n} and ${t} must match`)}function ya(n){w(n!=null,()=>"The input to the tensor constructor must be a non-null value.")}function z(n){if(n.length===0)return 1;let t=n[0];for(let e=1;e<n.length;e++)t*=n[e];return t}function Zt(n,t){if(n===t)return!0;if(n==null||t==null||n.length!==t.length)return!1;for(let e=0;e<n.length;e++)if(n[e]!==t[e])return!1;return!0}function oo(n){return n%1===0}function Ws(n,t){return t<=n.length?n:n+" ".repeat(t-n.length)}function Cf(n,t){let e=1,s=-1;for(let o=0;o<n.length;++o)if(n[o]>=0)e*=n[o];else if(n[o]===-1){if(s!==-1)throw Error(`Shapes can only have 1 implicit size. Found -1 at dim ${s} and dim ${o}`);s=o}else if(n[o]<0)throw Error(`Shapes can not be < 0. Found ${n[o]} at dim ${o}`);if(s===-1){if(t>0&&t!==e)throw Error(`Size(${t}) must match the product of shape ${n}`);return n}if(e===0)throw Error(`Cannot infer the missing size in [${n}] when there are 0 elements`);if(t%e!==0)throw Error(`The implicit shape can't be a fractional number. Got ${t} / ${e}`);const r=n.slice();return r[s]=t/e,r}function is(n,t){const e=t.length;return n=n==null?t.map((s,r)=>r):[].concat(n),w(n.every(s=>s>=-e&&s<e),()=>`All values in axis param must be in range [-${e}, ${e}) but got axis ${n}`),w(n.every(s=>oo(s)),()=>`All values in axis param must be integers but got axis ${n}`),n.map(s=>s<0?e+s:s)}function kf(n,t){const e=[],s=[],r=t!=null&&Array.isArray(t)&&t.length===0,o=t==null||r?null:is(t,n).sort();let i=0;for(let a=0;a<n.length;++a){if(o!=null){if(o[i]===a&&n[a]!==1)throw new Error(`Can't squeeze axis ${a} since its dim '${n[a]}' is not 1`);(o[i]==null||o[i]>a)&&n[a]===1&&(e.push(n[a]),s.push(a)),o[i]<=a&&i++}n[a]!==1&&(e.push(n[a]),s.push(a))}return{newShape:e,keptDims:s}}function Cn(n,t){return bt(n,t)}function bt(n,t){let e=null;if(n==null||n==="float32")e=new Float32Array(t);else if(n==="int32")e=new Int32Array(t);else if(n==="bool")e=new Uint8Array(t);else if(n==="string")e=new Array(t);else throw new Error(`Unknown data type ${n}`);return e}function _f(n,t){for(let e=0;e<n.length;e++){const s=n[e];if(isNaN(s)||!isFinite(s))throw Error(`A tensor of type ${t} being uploaded contains ${s}.`)}}function Tf(n){return n==="bool"||n==="complex64"||n==="float32"||n==="int32"||n==="string"}function io(n){if(n==="float32"||n==="int32")return 4;if(n==="complex64")return 8;if(n==="bool")return 1;throw new Error(`Unknown dtype ${n}`)}function Nf(n){if(n==null)return 0;let t=0;return n.forEach(e=>t+=e.length),t}function Gs(n){return typeof n=="string"||n instanceof String}function Df(n){return typeof n=="boolean"}function ao(n){return typeof n=="number"}function as(n){return Array.isArray(n)?as(n[0]):n instanceof Float32Array?"float32":n instanceof Int32Array||n instanceof Uint8Array||n instanceof Uint8ClampedArray?"int32":ao(n)?"float32":Gs(n)?"string":Df(n)?"bool":"float32"}function lo(n){return!!(n&&n.constructor&&n.call&&n.apply)}function Kt(n){const t=n.length;if(t<2)return[];const e=new Array(t-1);e[t-2]=n[t-1];for(let s=t-3;s>=0;--s)e[s]=e[s+1]*n[s+1];return e}function wa(n,t,e,s=!1){const r=new Array;if(t.length===1){const o=t[0]*(s?2:1);for(let i=0;i<o;i++)r[i]=e[n+i]}else{const o=t[0],i=t.slice(1),a=i.reduce((l,u)=>l*u)*(s?2:1);for(let l=0;l<o;l++)r[l]=wa(n+l*a,i,e,s)}return r}function xa(n,t,e=!1){if(n.length===0)return t[0];const s=n.reduce((r,o)=>r*o)*(e?2:1);if(s===0)return[];if(s!==t.length)throw new Error(`[${n}] does not match the input size ${t.length}${e?" for a complex tensor":""}.`);return wa(0,n,t,e)}function uo(n,t){if(Array.isArray(n))return n;if(t==="float32")return n instanceof Float32Array?n:new Float32Array(n);if(t==="int32")return n instanceof Int32Array?n:new Int32Array(n);if(t==="bool"||t==="string")return Uint8Array.from(new Int32Array(n));throw new Error(`Unknown dtype ${t}`)}function Sa(n,t){const e=ze(n,t);for(let s=0;s<e.length;s++)e[s]=1;return e}function ze(n,t){if(t==null||t==="float32"||t==="complex64")return new Float32Array(n);if(t==="int32")return new Int32Array(n);if(t==="bool")return new Uint8Array(n);throw new Error(`Unknown data type ${t}`)}function Re(n){n.forEach(t=>{w(Number.isInteger(t)&&t>=0,()=>`Tensor must have a shape comprised of positive integers but got shape [${n}].`)})}function co(n,t,e){if(t===0)return 0;if(t===1)return n[0];let s=n[n.length-1];for(let r=0;r<n.length-1;++r)s+=e[r]*n[r];return s}function ho(n,t,e){if(t===0)return[];if(t===1)return[n];const s=new Array(t);for(let r=0;r<s.length-1;++r)s[r]=Math.floor(n/e[r]),n-=s[r]*e[r];return s[s.length-1]=n,s}function fo(n){return n&&n.then&&typeof n.then=="function"}/**
|
|
32
32
|
* @license
|
|
33
33
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
34
34
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -88,7 +88,7 @@
|
|
|
88
88
|
* See the License for the specific language governing permissions and
|
|
89
89
|
* limitations under the License.
|
|
90
90
|
* =============================================================================
|
|
91
|
-
*/const
|
|
91
|
+
*/const Vs=mo("kernelRegistry",()=>new Map),ap=mo("gradRegistry",()=>new Map);function Pa(n,t){const e=Oa(n,t);return Vs.get(e)}function La(n){return ap.get(n)}function Ma(n){const t=Vs.entries(),e=[];for(;;){const{done:s,value:r}=t.next();if(s)break;const[o,i]=r,[a]=o.split("_");a===n&&e.push(i)}return e}function lp(n){const{kernelName:t,backendName:e}=n,s=Oa(t,e);Vs.has(s)&&kn(`The kernel '${t}' for backend '${e}' is already registered`),Vs.set(s,n)}function Oa(n,t){return`${t}_${n}`}/**
|
|
92
92
|
* @license
|
|
93
93
|
* Copyright 2023 Google LLC.
|
|
94
94
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -103,7 +103,7 @@
|
|
|
103
103
|
* See the License for the specific language governing permissions and
|
|
104
104
|
* limitations under the License.
|
|
105
105
|
* =============================================================================
|
|
106
|
-
*/function Ba(n){return n instanceof Float32Array||n instanceof Int32Array||n instanceof Uint8Array||n instanceof Uint8ClampedArray}var Je=typeof globalThis<"u"?globalThis:typeof window<"u"?window:typeof global<"u"?global:typeof self<"u"?self:{};function up(n){return n&&n.__esModule&&Object.prototype.hasOwnProperty.call(n,"default")?n.default:n}function cp(n){if(n.__esModule)return n;var t=n.default;if(typeof t=="function"){var e=function s(){return this instanceof s?Reflect.construct(t,arguments,this.constructor):t.apply(this,arguments)};e.prototype=t.prototype}else e={};return Object.defineProperty(e,"__esModule",{value:!0}),Object.keys(n).forEach(function(s){var r=Object.getOwnPropertyDescriptor(n,s);Object.defineProperty(e,s,r.get?r:{enumerable:!0,get:function(){return n[s]}})}),e}var Fa=it,Qt=null;try{Qt=new WebAssembly.Instance(new WebAssembly.Module(new Uint8Array([0,97,115,109,1,0,0,0,1,13,2,96,0,1,127,96,4,127,127,127,127,1,127,3,7,6,0,1,1,1,1,1,6,6,1,127,1,65,0,11,7,50,6,3,109,117,108,0,1,5,100,105,118,95,115,0,2,5,100,105,118,95,117,0,3,5,114,101,109,95,115,0,4,5,114,101,109,95,117,0,5,8,103,101,116,95,104,105,103,104,0,0,10,191,1,6,4,0,35,0,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,126,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,127,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,128,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,129,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,130,34,4,66,32,135,167,36,0,32,4,167,11])),{}).exports}catch{}function it(n,t,e){this.low=n|0,this.high=t|0,this.unsigned=!!e}it.prototype.__isLong__,Object.defineProperty(it.prototype,"__isLong__",{value:!0});function Bt(n){return(n&&n.__isLong__)===!0}it.isLong=Bt;var za={},Ua={};function Ze(n,t){var e,s,r;return t?(n>>>=0,(r=0<=n&&n<256)&&(s=Ua[n],s)?s:(e=at(n,(n|0)<0?-1:0,!0),r&&(Ua[n]=e),e)):(n|=0,(r=-128<=n&&n<128)&&(s=za[n],s)?s:(e=at(n,n<0?-1:0,!1),r&&(za[n]=e),e))}it.fromInt=Ze;function te(n,t){if(isNaN(n))return t?Qe:ee;if(t){if(n<0)return Qe;if(n>=Ga)return Ka}else{if(n<=-Va)return Ft;if(n+1>=Va)return Ha}return n<0?te(-n,t).neg():at(n%_n|0,n/_n|0,t)}it.fromNumber=te;function at(n,t,e){return new it(n,t,e)}it.fromBits=at;var Vs=Math.pow;function wo(n,t,e){if(n.length===0)throw Error("empty string");if(n==="NaN"||n==="Infinity"||n==="+Infinity"||n==="-Infinity")return ee;if(typeof t=="number"?(e=t,t=!1):t=!!t,e=e||10,e<2||36<e)throw RangeError("radix");var s;if((s=n.indexOf("-"))>0)throw Error("interior hyphen");if(s===0)return wo(n.substring(1),t,e).neg();for(var r=te(Vs(e,8)),o=ee,i=0;i<n.length;i+=8){var a=Math.min(8,n.length-i),l=parseInt(n.substring(i,i+a),e);if(a<8){var u=te(Vs(e,a));o=o.mul(u).add(te(l))}else o=o.mul(r),o=o.add(te(l))}return o.unsigned=t,o}it.fromString=wo;function ue(n,t){return typeof n=="number"?te(n,t):typeof n=="string"?wo(n,t):at(n.low,n.high,typeof t=="boolean"?t:n.unsigned)}it.fromValue=ue;var Wa=65536,hp=1<<24,_n=Wa*Wa,Ga=_n*_n,Va=Ga/2,qa=Ze(hp),ee=Ze(0);it.ZERO=ee;var Qe=Ze(0,!0);it.UZERO=Qe;var Tn=Ze(1);it.ONE=Tn;var ja=Ze(1,!0);it.UONE=ja;var xo=Ze(-1);it.NEG_ONE=xo;var Ha=at(-1,2147483647,!1);it.MAX_VALUE=Ha;var Ka=at(-1,-1,!0);it.MAX_UNSIGNED_VALUE=Ka;var Ft=at(0,-2147483648,!1);it.MIN_VALUE=Ft;var R=it.prototype;R.toInt=function(){return this.unsigned?this.low>>>0:this.low},R.toNumber=function(){return this.unsigned?(this.high>>>0)*_n+(this.low>>>0):this.high*_n+(this.low>>>0)},R.toString=function(t){if(t=t||10,t<2||36<t)throw RangeError("radix");if(this.isZero())return"0";if(this.isNegative())if(this.eq(Ft)){var e=te(t),s=this.div(e),r=s.mul(e).sub(this);return s.toString(t)+r.toInt().toString(t)}else return"-"+this.neg().toString(t);for(var o=te(Vs(t,6),this.unsigned),i=this,a="";;){var l=i.div(o),u=i.sub(l.mul(o)).toInt()>>>0,c=u.toString(t);if(i=l,i.isZero())return c+a;for(;c.length<6;)c="0"+c;a=""+c+a}},R.getHighBits=function(){return this.high},R.getHighBitsUnsigned=function(){return this.high>>>0},R.getLowBits=function(){return this.low},R.getLowBitsUnsigned=function(){return this.low>>>0},R.getNumBitsAbs=function(){if(this.isNegative())return this.eq(Ft)?64:this.neg().getNumBitsAbs();for(var t=this.high!=0?this.high:this.low,e=31;e>0&&!(t&1<<e);e--);return this.high!=0?e+33:e+1},R.isZero=function(){return this.high===0&&this.low===0},R.eqz=R.isZero,R.isNegative=function(){return!this.unsigned&&this.high<0},R.isPositive=function(){return this.unsigned||this.high>=0},R.isOdd=function(){return(this.low&1)===1},R.isEven=function(){return(this.low&1)===0},R.equals=function(t){return Bt(t)||(t=ue(t)),this.unsigned!==t.unsigned&&this.high>>>31===1&&t.high>>>31===1?!1:this.high===t.high&&this.low===t.low},R.eq=R.equals,R.notEquals=function(t){return!this.eq(t)},R.neq=R.notEquals,R.ne=R.notEquals,R.lessThan=function(t){return this.comp(t)<0},R.lt=R.lessThan,R.lessThanOrEqual=function(t){return this.comp(t)<=0},R.lte=R.lessThanOrEqual,R.le=R.lessThanOrEqual,R.greaterThan=function(t){return this.comp(t)>0},R.gt=R.greaterThan,R.greaterThanOrEqual=function(t){return this.comp(t)>=0},R.gte=R.greaterThanOrEqual,R.ge=R.greaterThanOrEqual,R.compare=function(t){if(Bt(t)||(t=ue(t)),this.eq(t))return 0;var e=this.isNegative(),s=t.isNegative();return e&&!s?-1:!e&&s?1:this.unsigned?t.high>>>0>this.high>>>0||t.high===this.high&&t.low>>>0>this.low>>>0?-1:1:this.sub(t).isNegative()?-1:1},R.comp=R.compare,R.negate=function(){return!this.unsigned&&this.eq(Ft)?Ft:this.not().add(Tn)},R.neg=R.negate,R.add=function(t){Bt(t)||(t=ue(t));var e=this.high>>>16,s=this.high&65535,r=this.low>>>16,o=this.low&65535,i=t.high>>>16,a=t.high&65535,l=t.low>>>16,u=t.low&65535,c=0,h=0,f=0,d=0;return d+=o+u,f+=d>>>16,d&=65535,f+=r+l,h+=f>>>16,f&=65535,h+=s+a,c+=h>>>16,h&=65535,c+=e+i,c&=65535,at(f<<16|d,c<<16|h,this.unsigned)},R.subtract=function(t){return Bt(t)||(t=ue(t)),this.add(t.neg())},R.sub=R.subtract,R.multiply=function(t){if(this.isZero())return ee;if(Bt(t)||(t=ue(t)),Qt){var e=Qt.mul(this.low,this.high,t.low,t.high);return at(e,Qt.get_high(),this.unsigned)}if(t.isZero())return ee;if(this.eq(Ft))return t.isOdd()?Ft:ee;if(t.eq(Ft))return this.isOdd()?Ft:ee;if(this.isNegative())return t.isNegative()?this.neg().mul(t.neg()):this.neg().mul(t).neg();if(t.isNegative())return this.mul(t.neg()).neg();if(this.lt(qa)&&t.lt(qa))return te(this.toNumber()*t.toNumber(),this.unsigned);var s=this.high>>>16,r=this.high&65535,o=this.low>>>16,i=this.low&65535,a=t.high>>>16,l=t.high&65535,u=t.low>>>16,c=t.low&65535,h=0,f=0,d=0,p=0;return p+=i*c,d+=p>>>16,p&=65535,d+=o*c,f+=d>>>16,d&=65535,d+=i*u,f+=d>>>16,d&=65535,f+=r*c,h+=f>>>16,f&=65535,f+=o*u,h+=f>>>16,f&=65535,f+=i*l,h+=f>>>16,f&=65535,h+=s*c+r*u+o*l+i*a,h&=65535,at(d<<16|p,h<<16|f,this.unsigned)},R.mul=R.multiply,R.divide=function(t){if(Bt(t)||(t=ue(t)),t.isZero())throw Error("division by zero");if(Qt){if(!this.unsigned&&this.high===-2147483648&&t.low===-1&&t.high===-1)return this;var e=(this.unsigned?Qt.div_u:Qt.div_s)(this.low,this.high,t.low,t.high);return at(e,Qt.get_high(),this.unsigned)}if(this.isZero())return this.unsigned?Qe:ee;var s,r,o;if(this.unsigned){if(t.unsigned||(t=t.toUnsigned()),t.gt(this))return Qe;if(t.gt(this.shru(1)))return ja;o=Qe}else{if(this.eq(Ft)){if(t.eq(Tn)||t.eq(xo))return Ft;if(t.eq(Ft))return Tn;var i=this.shr(1);return s=i.div(t).shl(1),s.eq(ee)?t.isNegative()?Tn:xo:(r=this.sub(t.mul(s)),o=s.add(r.div(t)),o)}else if(t.eq(Ft))return this.unsigned?Qe:ee;if(this.isNegative())return t.isNegative()?this.neg().div(t.neg()):this.neg().div(t).neg();if(t.isNegative())return this.div(t.neg()).neg();o=ee}for(r=this;r.gte(t);){s=Math.max(1,Math.floor(r.toNumber()/t.toNumber()));for(var a=Math.ceil(Math.log(s)/Math.LN2),l=a<=48?1:Vs(2,a-48),u=te(s),c=u.mul(t);c.isNegative()||c.gt(r);)s-=l,u=te(s,this.unsigned),c=u.mul(t);u.isZero()&&(u=Tn),o=o.add(u),r=r.sub(c)}return o},R.div=R.divide,R.modulo=function(t){if(Bt(t)||(t=ue(t)),Qt){var e=(this.unsigned?Qt.rem_u:Qt.rem_s)(this.low,this.high,t.low,t.high);return at(e,Qt.get_high(),this.unsigned)}return this.sub(this.div(t).mul(t))},R.mod=R.modulo,R.rem=R.modulo,R.not=function(){return at(~this.low,~this.high,this.unsigned)},R.and=function(t){return Bt(t)||(t=ue(t)),at(this.low&t.low,this.high&t.high,this.unsigned)},R.or=function(t){return Bt(t)||(t=ue(t)),at(this.low|t.low,this.high|t.high,this.unsigned)},R.xor=function(t){return Bt(t)||(t=ue(t)),at(this.low^t.low,this.high^t.high,this.unsigned)},R.shiftLeft=function(t){return Bt(t)&&(t=t.toInt()),(t&=63)===0?this:t<32?at(this.low<<t,this.high<<t|this.low>>>32-t,this.unsigned):at(0,this.low<<t-32,this.unsigned)},R.shl=R.shiftLeft,R.shiftRight=function(t){return Bt(t)&&(t=t.toInt()),(t&=63)===0?this:t<32?at(this.low>>>t|this.high<<32-t,this.high>>t,this.unsigned):at(this.high>>t-32,this.high>=0?0:-1,this.unsigned)},R.shr=R.shiftRight,R.shiftRightUnsigned=function(t){if(Bt(t)&&(t=t.toInt()),t&=63,t===0)return this;var e=this.high;if(t<32){var s=this.low;return at(s>>>t|e<<32-t,e>>>t,this.unsigned)}else return t===32?at(e,0,this.unsigned):at(e>>>t-32,0,this.unsigned)},R.shru=R.shiftRightUnsigned,R.shr_u=R.shiftRightUnsigned,R.toSigned=function(){return this.unsigned?at(this.low,this.high,!1):this},R.toUnsigned=function(){return this.unsigned?this:at(this.low,this.high,!0)},R.toBytes=function(t){return t?this.toBytesLE():this.toBytesBE()},R.toBytesLE=function(){var t=this.high,e=this.low;return[e&255,e>>>8&255,e>>>16&255,e>>>24,t&255,t>>>8&255,t>>>16&255,t>>>24]},R.toBytesBE=function(){var t=this.high,e=this.low;return[t>>>24,t>>>16&255,t>>>8&255,t&255,e>>>24,e>>>16&255,e>>>8&255,e&255]},it.fromBytes=function(t,e,s){return s?it.fromBytesLE(t,e):it.fromBytesBE(t,e)},it.fromBytesLE=function(t,e){return new it(t[0]|t[1]<<8|t[2]<<16|t[3]<<24,t[4]|t[5]<<8|t[6]<<16|t[7]<<24,e)},it.fromBytesBE=function(t,e){return new it(t[4]<<24|t[5]<<16|t[6]<<8|t[7],t[0]<<24|t[1]<<16|t[2]<<8|t[3],e)};const Ya=up(Fa),fp=qt({__proto__:null,default:Ya},[Fa]);/**
|
|
106
|
+
*/function Ba(n){return n instanceof Float32Array||n instanceof Int32Array||n instanceof Uint8Array||n instanceof Uint8ClampedArray}var Ze=typeof globalThis<"u"?globalThis:typeof window<"u"?window:typeof global<"u"?global:typeof self<"u"?self:{};function up(n){return n&&n.__esModule&&Object.prototype.hasOwnProperty.call(n,"default")?n.default:n}function cp(n){if(n.__esModule)return n;var t=n.default;if(typeof t=="function"){var e=function s(){return this instanceof s?Reflect.construct(t,arguments,this.constructor):t.apply(this,arguments)};e.prototype=t.prototype}else e={};return Object.defineProperty(e,"__esModule",{value:!0}),Object.keys(n).forEach(function(s){var r=Object.getOwnPropertyDescriptor(n,s);Object.defineProperty(e,s,r.get?r:{enumerable:!0,get:function(){return n[s]}})}),e}var Fa=it,Qt=null;try{Qt=new WebAssembly.Instance(new WebAssembly.Module(new Uint8Array([0,97,115,109,1,0,0,0,1,13,2,96,0,1,127,96,4,127,127,127,127,1,127,3,7,6,0,1,1,1,1,1,6,6,1,127,1,65,0,11,7,50,6,3,109,117,108,0,1,5,100,105,118,95,115,0,2,5,100,105,118,95,117,0,3,5,114,101,109,95,115,0,4,5,114,101,109,95,117,0,5,8,103,101,116,95,104,105,103,104,0,0,10,191,1,6,4,0,35,0,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,126,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,127,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,128,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,129,34,4,66,32,135,167,36,0,32,4,167,11,36,1,1,126,32,0,173,32,1,173,66,32,134,132,32,2,173,32,3,173,66,32,134,132,130,34,4,66,32,135,167,36,0,32,4,167,11])),{}).exports}catch{}function it(n,t,e){this.low=n|0,this.high=t|0,this.unsigned=!!e}it.prototype.__isLong__,Object.defineProperty(it.prototype,"__isLong__",{value:!0});function Bt(n){return(n&&n.__isLong__)===!0}it.isLong=Bt;var za={},Ua={};function Qe(n,t){var e,s,r;return t?(n>>>=0,(r=0<=n&&n<256)&&(s=Ua[n],s)?s:(e=at(n,(n|0)<0?-1:0,!0),r&&(Ua[n]=e),e)):(n|=0,(r=-128<=n&&n<128)&&(s=za[n],s)?s:(e=at(n,n<0?-1:0,!1),r&&(za[n]=e),e))}it.fromInt=Qe;function te(n,t){if(isNaN(n))return t?tn:ee;if(t){if(n<0)return tn;if(n>=Ga)return Ka}else{if(n<=-Va)return Ft;if(n+1>=Va)return Ha}return n<0?te(-n,t).neg():at(n%_n|0,n/_n|0,t)}it.fromNumber=te;function at(n,t,e){return new it(n,t,e)}it.fromBits=at;var qs=Math.pow;function wo(n,t,e){if(n.length===0)throw Error("empty string");if(n==="NaN"||n==="Infinity"||n==="+Infinity"||n==="-Infinity")return ee;if(typeof t=="number"?(e=t,t=!1):t=!!t,e=e||10,e<2||36<e)throw RangeError("radix");var s;if((s=n.indexOf("-"))>0)throw Error("interior hyphen");if(s===0)return wo(n.substring(1),t,e).neg();for(var r=te(qs(e,8)),o=ee,i=0;i<n.length;i+=8){var a=Math.min(8,n.length-i),l=parseInt(n.substring(i,i+a),e);if(a<8){var u=te(qs(e,a));o=o.mul(u).add(te(l))}else o=o.mul(r),o=o.add(te(l))}return o.unsigned=t,o}it.fromString=wo;function ue(n,t){return typeof n=="number"?te(n,t):typeof n=="string"?wo(n,t):at(n.low,n.high,typeof t=="boolean"?t:n.unsigned)}it.fromValue=ue;var Wa=65536,hp=1<<24,_n=Wa*Wa,Ga=_n*_n,Va=Ga/2,qa=Qe(hp),ee=Qe(0);it.ZERO=ee;var tn=Qe(0,!0);it.UZERO=tn;var Tn=Qe(1);it.ONE=Tn;var ja=Qe(1,!0);it.UONE=ja;var xo=Qe(-1);it.NEG_ONE=xo;var Ha=at(-1,2147483647,!1);it.MAX_VALUE=Ha;var Ka=at(-1,-1,!0);it.MAX_UNSIGNED_VALUE=Ka;var Ft=at(0,-2147483648,!1);it.MIN_VALUE=Ft;var R=it.prototype;R.toInt=function(){return this.unsigned?this.low>>>0:this.low},R.toNumber=function(){return this.unsigned?(this.high>>>0)*_n+(this.low>>>0):this.high*_n+(this.low>>>0)},R.toString=function(t){if(t=t||10,t<2||36<t)throw RangeError("radix");if(this.isZero())return"0";if(this.isNegative())if(this.eq(Ft)){var e=te(t),s=this.div(e),r=s.mul(e).sub(this);return s.toString(t)+r.toInt().toString(t)}else return"-"+this.neg().toString(t);for(var o=te(qs(t,6),this.unsigned),i=this,a="";;){var l=i.div(o),u=i.sub(l.mul(o)).toInt()>>>0,c=u.toString(t);if(i=l,i.isZero())return c+a;for(;c.length<6;)c="0"+c;a=""+c+a}},R.getHighBits=function(){return this.high},R.getHighBitsUnsigned=function(){return this.high>>>0},R.getLowBits=function(){return this.low},R.getLowBitsUnsigned=function(){return this.low>>>0},R.getNumBitsAbs=function(){if(this.isNegative())return this.eq(Ft)?64:this.neg().getNumBitsAbs();for(var t=this.high!=0?this.high:this.low,e=31;e>0&&!(t&1<<e);e--);return this.high!=0?e+33:e+1},R.isZero=function(){return this.high===0&&this.low===0},R.eqz=R.isZero,R.isNegative=function(){return!this.unsigned&&this.high<0},R.isPositive=function(){return this.unsigned||this.high>=0},R.isOdd=function(){return(this.low&1)===1},R.isEven=function(){return(this.low&1)===0},R.equals=function(t){return Bt(t)||(t=ue(t)),this.unsigned!==t.unsigned&&this.high>>>31===1&&t.high>>>31===1?!1:this.high===t.high&&this.low===t.low},R.eq=R.equals,R.notEquals=function(t){return!this.eq(t)},R.neq=R.notEquals,R.ne=R.notEquals,R.lessThan=function(t){return this.comp(t)<0},R.lt=R.lessThan,R.lessThanOrEqual=function(t){return this.comp(t)<=0},R.lte=R.lessThanOrEqual,R.le=R.lessThanOrEqual,R.greaterThan=function(t){return this.comp(t)>0},R.gt=R.greaterThan,R.greaterThanOrEqual=function(t){return this.comp(t)>=0},R.gte=R.greaterThanOrEqual,R.ge=R.greaterThanOrEqual,R.compare=function(t){if(Bt(t)||(t=ue(t)),this.eq(t))return 0;var e=this.isNegative(),s=t.isNegative();return e&&!s?-1:!e&&s?1:this.unsigned?t.high>>>0>this.high>>>0||t.high===this.high&&t.low>>>0>this.low>>>0?-1:1:this.sub(t).isNegative()?-1:1},R.comp=R.compare,R.negate=function(){return!this.unsigned&&this.eq(Ft)?Ft:this.not().add(Tn)},R.neg=R.negate,R.add=function(t){Bt(t)||(t=ue(t));var e=this.high>>>16,s=this.high&65535,r=this.low>>>16,o=this.low&65535,i=t.high>>>16,a=t.high&65535,l=t.low>>>16,u=t.low&65535,c=0,h=0,f=0,d=0;return d+=o+u,f+=d>>>16,d&=65535,f+=r+l,h+=f>>>16,f&=65535,h+=s+a,c+=h>>>16,h&=65535,c+=e+i,c&=65535,at(f<<16|d,c<<16|h,this.unsigned)},R.subtract=function(t){return Bt(t)||(t=ue(t)),this.add(t.neg())},R.sub=R.subtract,R.multiply=function(t){if(this.isZero())return ee;if(Bt(t)||(t=ue(t)),Qt){var e=Qt.mul(this.low,this.high,t.low,t.high);return at(e,Qt.get_high(),this.unsigned)}if(t.isZero())return ee;if(this.eq(Ft))return t.isOdd()?Ft:ee;if(t.eq(Ft))return this.isOdd()?Ft:ee;if(this.isNegative())return t.isNegative()?this.neg().mul(t.neg()):this.neg().mul(t).neg();if(t.isNegative())return this.mul(t.neg()).neg();if(this.lt(qa)&&t.lt(qa))return te(this.toNumber()*t.toNumber(),this.unsigned);var s=this.high>>>16,r=this.high&65535,o=this.low>>>16,i=this.low&65535,a=t.high>>>16,l=t.high&65535,u=t.low>>>16,c=t.low&65535,h=0,f=0,d=0,p=0;return p+=i*c,d+=p>>>16,p&=65535,d+=o*c,f+=d>>>16,d&=65535,d+=i*u,f+=d>>>16,d&=65535,f+=r*c,h+=f>>>16,f&=65535,f+=o*u,h+=f>>>16,f&=65535,f+=i*l,h+=f>>>16,f&=65535,h+=s*c+r*u+o*l+i*a,h&=65535,at(d<<16|p,h<<16|f,this.unsigned)},R.mul=R.multiply,R.divide=function(t){if(Bt(t)||(t=ue(t)),t.isZero())throw Error("division by zero");if(Qt){if(!this.unsigned&&this.high===-2147483648&&t.low===-1&&t.high===-1)return this;var e=(this.unsigned?Qt.div_u:Qt.div_s)(this.low,this.high,t.low,t.high);return at(e,Qt.get_high(),this.unsigned)}if(this.isZero())return this.unsigned?tn:ee;var s,r,o;if(this.unsigned){if(t.unsigned||(t=t.toUnsigned()),t.gt(this))return tn;if(t.gt(this.shru(1)))return ja;o=tn}else{if(this.eq(Ft)){if(t.eq(Tn)||t.eq(xo))return Ft;if(t.eq(Ft))return Tn;var i=this.shr(1);return s=i.div(t).shl(1),s.eq(ee)?t.isNegative()?Tn:xo:(r=this.sub(t.mul(s)),o=s.add(r.div(t)),o)}else if(t.eq(Ft))return this.unsigned?tn:ee;if(this.isNegative())return t.isNegative()?this.neg().div(t.neg()):this.neg().div(t).neg();if(t.isNegative())return this.div(t.neg()).neg();o=ee}for(r=this;r.gte(t);){s=Math.max(1,Math.floor(r.toNumber()/t.toNumber()));for(var a=Math.ceil(Math.log(s)/Math.LN2),l=a<=48?1:qs(2,a-48),u=te(s),c=u.mul(t);c.isNegative()||c.gt(r);)s-=l,u=te(s,this.unsigned),c=u.mul(t);u.isZero()&&(u=Tn),o=o.add(u),r=r.sub(c)}return o},R.div=R.divide,R.modulo=function(t){if(Bt(t)||(t=ue(t)),Qt){var e=(this.unsigned?Qt.rem_u:Qt.rem_s)(this.low,this.high,t.low,t.high);return at(e,Qt.get_high(),this.unsigned)}return this.sub(this.div(t).mul(t))},R.mod=R.modulo,R.rem=R.modulo,R.not=function(){return at(~this.low,~this.high,this.unsigned)},R.and=function(t){return Bt(t)||(t=ue(t)),at(this.low&t.low,this.high&t.high,this.unsigned)},R.or=function(t){return Bt(t)||(t=ue(t)),at(this.low|t.low,this.high|t.high,this.unsigned)},R.xor=function(t){return Bt(t)||(t=ue(t)),at(this.low^t.low,this.high^t.high,this.unsigned)},R.shiftLeft=function(t){return Bt(t)&&(t=t.toInt()),(t&=63)===0?this:t<32?at(this.low<<t,this.high<<t|this.low>>>32-t,this.unsigned):at(0,this.low<<t-32,this.unsigned)},R.shl=R.shiftLeft,R.shiftRight=function(t){return Bt(t)&&(t=t.toInt()),(t&=63)===0?this:t<32?at(this.low>>>t|this.high<<32-t,this.high>>t,this.unsigned):at(this.high>>t-32,this.high>=0?0:-1,this.unsigned)},R.shr=R.shiftRight,R.shiftRightUnsigned=function(t){if(Bt(t)&&(t=t.toInt()),t&=63,t===0)return this;var e=this.high;if(t<32){var s=this.low;return at(s>>>t|e<<32-t,e>>>t,this.unsigned)}else return t===32?at(e,0,this.unsigned):at(e>>>t-32,0,this.unsigned)},R.shru=R.shiftRightUnsigned,R.shr_u=R.shiftRightUnsigned,R.toSigned=function(){return this.unsigned?at(this.low,this.high,!1):this},R.toUnsigned=function(){return this.unsigned?this:at(this.low,this.high,!0)},R.toBytes=function(t){return t?this.toBytesLE():this.toBytesBE()},R.toBytesLE=function(){var t=this.high,e=this.low;return[e&255,e>>>8&255,e>>>16&255,e>>>24,t&255,t>>>8&255,t>>>16&255,t>>>24]},R.toBytesBE=function(){var t=this.high,e=this.low;return[t>>>24,t>>>16&255,t>>>8&255,t&255,e>>>24,e>>>16&255,e>>>8&255,e&255]},it.fromBytes=function(t,e,s){return s?it.fromBytesLE(t,e):it.fromBytesBE(t,e)},it.fromBytesLE=function(t,e){return new it(t[0]|t[1]<<8|t[2]<<16|t[3]<<24,t[4]|t[5]<<8|t[6]<<16|t[7]<<24,e)},it.fromBytesBE=function(t,e){return new it(t[4]<<24|t[5]<<16|t[6]<<8|t[7],t[0]<<24|t[1]<<16|t[2]<<8|t[3],e)};const Ya=up(Fa),fp=Ht({__proto__:null,default:Ya},[Fa]);/**
|
|
107
107
|
* @license
|
|
108
108
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
109
109
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -118,7 +118,7 @@
|
|
|
118
118
|
* See the License for the specific language governing permissions and
|
|
119
119
|
* limitations under the License.
|
|
120
120
|
* =============================================================================
|
|
121
|
-
*/const
|
|
121
|
+
*/const en=Ya||fp;function js(n){return en.fromString(n,!0,16)}const Xa=js("c3a5c85c97cb3127"),nn=js("b492b66fbe98f273"),kt=js("9ae16a3b2f90404f");function So(n){return n.xor(n.shru(47))}function Ja(n,t,e){const s=n.slice(t,t+e);return en.fromBytes(Array.from(s),!0,!0)}function rt(n,t){return Ja(n,t,8)}function Za(n,t){return Ja(n,t,4)}function yt(n,t){return t===0?n:n.shru(t).or(n.shl(64-t))}function Ue(n,t,e=js("9ddfea08eb382d69")){let s=n.xor(t).mul(e);s=s.xor(s.shru(47));let r=t.xor(s).mul(e);return r=r.xor(r.shru(47)),r=r.mul(e),r}function dp(n,t,e,s,r,o){r=r.add(n),o=yt(o.add(r).add(s),21);const i=r;return r=r.add(t),r=r.add(e),o=o.add(yt(r,44)),[r.add(s),o.add(i)]}function Hs(n,t,e,s){return dp(rt(n,t),rt(n,t+8),rt(n,t+16),rt(n,t+24),e,s)}function pp(n,t=n.length){if(t>=8){const e=kt.add(t*2),s=rt(n,0).add(kt),r=rt(n,t-8),o=yt(r,37).mul(e).add(s),i=yt(s,25).add(r).mul(e);return Ue(o,i,e)}if(t>=4){const e=kt.add(t*2),s=Za(n,0);return Ue(s.shl(3).add(t),Za(n,t-4),e)}if(t>0){const e=n[0],s=n[t>>1],r=n[t-1],o=e+(s<<8),i=t+(r<<2);return So(kt.mul(o).xor(Xa.mul(i))).mul(kt)}return kt}function mp(n,t=n.length){const e=kt.add(t*2),s=rt(n,0).mul(nn),r=rt(n,8),o=rt(n,t-8).mul(e),i=rt(n,t-16).mul(kt);return Ue(yt(s.add(r),43).add(yt(o,30)).add(i),s.add(yt(r.add(kt),18)).add(o),e)}function gp(n,t=n.length){const e=kt.add(t*2),s=rt(n,0).mul(kt),r=rt(n,8),o=rt(n,t-8).mul(e),i=rt(n,t-16).mul(kt),a=yt(s.add(r),43).add(yt(o,30)).add(i),l=Ue(a,s.add(yt(r.add(kt),18)).add(o),e),u=rt(n,16).mul(e),c=rt(n,24),h=a.add(rt(n,t-32)).mul(e),f=l.add(rt(n,t-24)).mul(e);return Ue(yt(u.add(c),43).add(yt(h,30)).add(f),u.add(yt(c.add(s),18)).add(h),e)}function bp(n,t=n.length){const e=en.fromNumber(81,!0);if(t<=32)return t<=16?pp(n,t):mp(n,t);if(t<=64)return gp(n,t);let s=e,r=e.mul(nn).add(113),o=So(r.mul(kt).add(113)).mul(kt),i=[en.UZERO,en.UZERO],a=[en.UZERO,en.UZERO];s=s.mul(kt).add(rt(n,0));let l=0;const u=(t-1>>6)*64,c=u+(t-1&63)-63;do s=yt(s.add(r).add(i[0]).add(rt(n,l+8)),37).mul(nn),r=yt(r.add(i[1]).add(rt(n,l+48)),42).mul(nn),s=s.xor(a[1]),r=r.add(i[0]).add(rt(n,l+40)),o=yt(o.add(a[0]),33).mul(nn),i=Hs(n,l,i[1].mul(nn),s.add(a[0])),a=Hs(n,l+32,o.add(a[1]),r.add(rt(n,l+16))),[o,s]=[s,o],l+=64;while(l!==u);const h=nn.add(o.and(255).shl(1));return l=c,a[0]=a[0].add(t-1&63),i[0]=i[0].add(a[0]),a[0]=a[0].add(i[0]),s=yt(s.add(r).add(i[0]).add(rt(n,l+8)),37).mul(h),r=yt(r.add(i[1]).add(rt(n,l+48)),42).mul(h),s=s.xor(a[1].mul(9)),r=r.add(i[0].mul(9).add(rt(n,l+40))),o=yt(o.add(a[0]),33).mul(h),i=Hs(n,l,i[1].mul(h),s.add(a[0])),a=Hs(n,l+32,o.add(a[1]),r.add(rt(n,l+16))),[o,s]=[s,o],Ue(Ue(i[0],a[0],h).add(So(r).mul(Xa)).add(o),Ue(i[1],a[1],h).add(s),h)}/**
|
|
122
122
|
* @license
|
|
123
123
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
124
124
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -133,7 +133,7 @@
|
|
|
133
133
|
* See the License for the specific language governing permissions and
|
|
134
134
|
* limitations under the License.
|
|
135
135
|
* =============================================================================
|
|
136
|
-
*/function yp(n,t){return t==="string"?
|
|
136
|
+
*/function yp(n,t){return t==="string"?sn(n):Ks([n],t)}function wp(n,t){return n instanceof Float32Array&&t==="float32"||n instanceof Int32Array&&t==="int32"||n instanceof Uint8Array&&t==="bool"}function Ks(n,t){if(t==="string")throw new Error("Cannot convert a string[] to a TypedArray");if(Array.isArray(n)&&(n=rn(n)),W().getBool("DEBUG")&&_f(n,t),wp(n,t))return n;if(t==null||t==="float32"||t==="complex64")return new Float32Array(n);if(t==="int32")return new Int32Array(n);if(t==="bool"){const e=new Uint8Array(n.length);for(let s=0;s<e.length;++s)Math.round(n[s])!==0&&(e[s]=1);return e}else throw new Error(`Unknown data type ${t}`)}function Nn(){return W().platform.now()}function sn(n,t="utf-8"){return t=t||"utf-8",W().platform.encode(n,t)}function Ys(n,t="utf-8"){return t=t||"utf-8",W().platform.decode(n,t)}function ne(n){return W().platform.isTypedArray!=null?W().platform.isTypedArray(n):Ba(n)}function rn(n,t=[],e=!1){if(t==null&&(t=[]),typeof n=="boolean"||typeof n=="number"||typeof n=="string"||fo(n)||n==null||ne(n)&&e)t.push(n);else if(Array.isArray(n)||ne(n))for(let s=0;s<n.length;++s)rn(n[s],t,e);else{let s=-1;for(const r of Object.keys(n))/^([1-9]+[0-9]*|0)$/.test(r)&&(s=Math.max(s,Number(r)));for(let r=0;r<=s;r++)rn(n[r],t,e)}return t}/**
|
|
137
137
|
* @license
|
|
138
138
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
139
139
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -148,7 +148,7 @@
|
|
|
148
148
|
* See the License for the specific language governing permissions and
|
|
149
149
|
* limitations under the License.
|
|
150
150
|
* =============================================================================
|
|
151
|
-
*/class xp{constructor(t,e){this.backendTimer=t,this.logger=e,e==null&&(this.logger=new $p)}profileKernel(t,e,s){let r;const o=()=>{r=s()};let i;const a=Nn();if(this.backendTimer.timerAvailable())i=this.backendTimer.time(o);else{o();for(const u of r)u.dataSync();i=Promise.resolve({kernelMs:Nn()-a})}if(W().getBool("CHECK_COMPUTATION_FOR_ERRORS"))for(let u=0;u<r.length;u++){const c=r[u];c.data().then(h=>{Sp(h,c.dtype,t)})}return{kernelName:t,outputs:r,inputs:e,timeMs:i.then(u=>u.kernelMs),extraInfo:i.then(u=>u.getExtraProfileInfo!=null?u.getExtraProfileInfo():"")}}logKernelProfile(t){const{kernelName:e,outputs:s,timeMs:r,inputs:o,extraInfo:i}=t;s.forEach(a=>{Promise.all([a.data(),r,i]).then(l=>{this.logger.logKernelProfile(e,a,l[0],l[1],o,l[2])})})}}function Sp(n,t,e){if(t!=="float32")return!1;for(let s=0;s<n.length;s++){const r=n[s];if(isNaN(r)||!isFinite(r))return console.warn(`Found ${r} in the result of '${e}'`),!0}return!1}class $p{logKernelProfile(t,e,s,r,o,i){const a=typeof r=="number"?
|
|
151
|
+
*/class xp{constructor(t,e){this.backendTimer=t,this.logger=e,e==null&&(this.logger=new $p)}profileKernel(t,e,s){let r;const o=()=>{r=s()};let i;const a=Nn();if(this.backendTimer.timerAvailable())i=this.backendTimer.time(o);else{o();for(const u of r)u.dataSync();i=Promise.resolve({kernelMs:Nn()-a})}if(W().getBool("CHECK_COMPUTATION_FOR_ERRORS"))for(let u=0;u<r.length;u++){const c=r[u];c.data().then(h=>{Sp(h,c.dtype,t)})}return{kernelName:t,outputs:r,inputs:e,timeMs:i.then(u=>u.kernelMs),extraInfo:i.then(u=>u.getExtraProfileInfo!=null?u.getExtraProfileInfo():"")}}logKernelProfile(t){const{kernelName:e,outputs:s,timeMs:r,inputs:o,extraInfo:i}=t;s.forEach(a=>{Promise.all([a.data(),r,i]).then(l=>{this.logger.logKernelProfile(e,a,l[0],l[1],o,l[2])})})}}function Sp(n,t,e){if(t!=="float32")return!1;for(let s=0;s<n.length;s++){const r=n[s];if(isNaN(r)||!isFinite(r))return console.warn(`Found ${r} in the result of '${e}'`),!0}return!1}class $p{logKernelProfile(t,e,s,r,o,i){const a=typeof r=="number"?Ws(`${r}ms`,9):r.error,l=Ws(t,25),u=e.rank,c=e.size,h=Ws(e.shape.toString(),14);let f="";for(const d in o){const p=o[d];if(p!=null){const g=p.shape||e.shape,m=g.length;f+=`${d}: ${m}D ${m>0?g:""} `}}console.log(`%c${l} %c${a} %c${u}D ${h} %c${c} %c${f} %c${i}`,"font-weight:bold","color:red","color:blue","color: orange","color: green","color: steelblue")}}/**
|
|
152
152
|
* @license
|
|
153
153
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
154
154
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -178,11 +178,11 @@
|
|
|
178
178
|
* See the License for the specific language governing permissions and
|
|
179
179
|
* limitations under the License.
|
|
180
180
|
* =============================================================================
|
|
181
|
-
*/const Qa=20,
|
|
181
|
+
*/const Qa=20,ls=3,$o=7;function Ap(n,t,e,s){const r=Kt(t),o=Ep(n,t,e,r),i=t.length,a=Xs(n,t,e,r,o),l=["Tensor"];return s&&(l.push(` dtype: ${e}`),l.push(` rank: ${i}`),l.push(` shape: [${t}]`),l.push(" values:")),l.push(a.map(u=>" "+u).join(`
|
|
182
182
|
`)),l.join(`
|
|
183
|
-
`)}function Ep(n,t,e,s){const r=z(t),o=s[s.length-1],i=new Array(o).fill(0),a=t.length,l=e==="complex64"?
|
|
183
|
+
`)}function Ep(n,t,e,s){const r=z(t),o=s[s.length-1],i=new Array(o).fill(0),a=t.length,l=e==="complex64"?cs(n):n;if(a>1)for(let u=0;u<r/o;u++){const c=u*o;for(let h=0;h<o;h++)i[h]=Math.max(i[h],us(l[c+h],0,e).length)}return i}function us(n,t,e){let s;return Array.isArray(n)?s=`${parseFloat(n[0].toFixed($o))} + ${parseFloat(n[1].toFixed($o))}j`:Gs(n)?s=`'${n}'`:e==="bool"?s=tl(n):s=parseFloat(n.toFixed($o)).toString(),Ws(s,t)}function tl(n){return n===0?"false":"true"}function Xs(n,t,e,s,r,o=!0){const i=e==="complex64"?2:1,a=t[0],l=t.length;if(l===0){if(e==="complex64"){const g=cs(n);return[us(g[0],0,e)]}return e==="bool"?[tl(n[0])]:[n[0].toString()]}if(l===1){if(a>Qa){const m=ls*i;let b=Array.from(n.slice(0,m)),y=Array.from(n.slice((a-ls)*i,a*i));return e==="complex64"&&(b=cs(b),y=cs(y)),["["+b.map((S,x)=>us(S,r[x],e)).join(", ")+", ..., "+y.map((S,x)=>us(S,r[a-ls+x],e)).join(", ")+"]"]}return["["+(e==="complex64"?cs(n):Array.from(n)).map((m,b)=>us(m,r[b],e)).join(", ")+"]"]}const u=t.slice(1),c=s.slice(1),h=s[0]*i,f=[];if(a>Qa){for(let g=0;g<ls;g++){const m=g*h,b=m+h;f.push(...Xs(n.slice(m,b),u,e,c,r,!1))}f.push("...");for(let g=a-ls;g<a;g++){const m=g*h,b=m+h;f.push(...Xs(n.slice(m,b),u,e,c,r,g===a-1))}}else for(let g=0;g<a;g++){const m=g*h,b=m+h;f.push(...Xs(n.slice(m,b),u,e,c,r,g===a-1))}const d=l===2?",":"";f[0]="["+(a>0?f[0]+d:"");for(let g=1;g<f.length-1;g++)f[g]=" "+f[g]+d;let p=`,
|
|
184
184
|
`;for(let g=2;g<l;g++)p+=`
|
|
185
|
-
`;return f[f.length-1]=" "+f[f.length-1]+"]"+(o?"":p),f}function
|
|
185
|
+
`;return f[f.length-1]=" "+f[f.length-1]+"]"+(o?"":p),f}function cs(n){const t=[];for(let e=0;e<n.length;e+=2)t.push([n[e],n[e+1]]);return t}/**
|
|
186
186
|
* @license
|
|
187
187
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
188
188
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -197,7 +197,7 @@
|
|
|
197
197
|
* See the License for the specific language governing permissions and
|
|
198
198
|
* limitations under the License.
|
|
199
199
|
* =============================================================================
|
|
200
|
-
*/class
|
|
200
|
+
*/class Js{constructor(t,e,s){if(this.dtype=e,this.shape=t.slice(),this.size=z(t),s!=null){const r=s.length;w(r===this.size,()=>`Length of values '${r}' does not match the size inferred by the shape '${this.size}'.`)}if(e==="complex64")throw new Error("complex64 dtype TensorBuffers are not supported. Please create a TensorBuffer for the real and imaginary parts separately and call tf.complex(real, imag).");this.values=s||bt(e,this.size),this.strides=Kt(t)}set(t,...e){e.length===0&&(e=[0]),w(e.length===this.rank,()=>`The number of provided coordinates (${e.length}) must match the rank (${this.rank})`);const s=this.locToIndex(e);this.values[s]=t}get(...t){t.length===0&&(t=[0]);let e=0;for(const r of t){if(r<0||r>=this.shape[e]){const o=`Requested out of range element at ${t}. Buffer shape=${this.shape}`;throw new Error(o)}e++}let s=t[t.length-1];for(let r=0;r<t.length-1;++r)s+=this.strides[r]*t[r];return this.values[s]}locToIndex(t){if(this.rank===0)return 0;if(this.rank===1)return t[0];let e=t[t.length-1];for(let s=0;s<t.length-1;++s)e+=this.strides[s]*t[s];return e}indexToLoc(t){if(this.rank===0)return[];if(this.rank===1)return[t];const e=new Array(this.shape.length);for(let s=0;s<e.length-1;++s)e[s]=Math.floor(t/this.strides[s]),t-=e[s]*this.strides[s];return e[e.length-1]=t,e}get rank(){return this.shape.length}toTensor(){return ce().makeTensor(this.values,this.shape,this.dtype)}}let ce=null,Dn=null;function Cp(n){ce=n}function kp(n){Dn=n}class At{constructor(t,e,s,r){this.kept=!1,this.isDisposedInternal=!1,this.shape=t.slice(),this.dtype=e||"float32",this.size=z(t),this.strides=Kt(t),this.dataId=s,this.id=r,this.rankType=this.rank<5?this.rank.toString():"higher"}get rank(){return this.shape.length}async buffer(){const t=await this.data();return Dn.buffer(this.shape,this.dtype,t)}bufferSync(){return Dn.buffer(this.shape,this.dtype,this.dataSync())}async array(){const t=await this.data();return xa(this.shape,t,this.dtype==="complex64")}arraySync(){return xa(this.shape,this.dataSync(),this.dtype==="complex64")}async data(){this.throwIfDisposed();const t=ce().read(this.dataId);if(this.dtype==="string"){const e=await t;try{return e.map(s=>Ys(s))}catch{throw new Error("Failed to decode the string bytes into utf-8. To get the original bytes, call tensor.bytes().")}}return t}dataToGPU(t){return this.throwIfDisposed(),ce().readToGPU(this.dataId,t)}dataSync(){this.throwIfDisposed();const t=ce().readSync(this.dataId);if(this.dtype==="string")try{return t.map(e=>Ys(e))}catch{throw new Error("Failed to decode the string bytes into utf-8. To get the original bytes, call tensor.bytes().")}return t}async bytes(){this.throwIfDisposed();const t=await ce().read(this.dataId);return this.dtype==="string"?t:new Uint8Array(t.buffer)}dispose(){this.isDisposed||(this.kerasMask&&this.kerasMask.dispose(),ce().disposeTensor(this),this.isDisposedInternal=!0)}get isDisposed(){return this.isDisposedInternal}throwIfDisposed(){if(this.isDisposed)throw new Error("Tensor is disposed.")}print(t=!1){return Dn.print(this,t)}clone(){return this.throwIfDisposed(),Dn.clone(this)}toString(t=!1){const e=this.dataSync();return Ap(e,this.shape,this.dtype,t)}cast(t){return this.throwIfDisposed(),Dn.cast(this,t)}variable(t=!0,e,s){return this.throwIfDisposed(),ce().makeVariable(this,t,e,s)}}Object.defineProperty(At,Symbol.hasInstance,{value:n=>!!n&&n.data!=null&&n.dataSync!=null&&n.throwIfDisposed!=null});function el(){return mo("Tensor",()=>At)}el();class Zs extends At{constructor(t,e,s,r){super(t.shape,t.dtype,t.dataId,r),this.trainable=e,this.name=s}assign(t){if(t.dtype!==this.dtype)throw new Error(`dtype of the new value (${t.dtype}) and previous value (${this.dtype}) must match`);if(!Zt(t.shape,this.shape))throw new Error(`shape of the new value (${t.shape}) and previous value (${this.shape}) must match`);ce().disposeTensor(this),this.dataId=t.dataId,ce().incRef(this,null)}dispose(){ce().disposeVariable(this),this.isDisposedInternal=!0}}Object.defineProperty(Zs,Symbol.hasInstance,{value:n=>n instanceof At&&n.assign!=null&&n.assign instanceof Function});/**
|
|
201
201
|
* @license
|
|
202
202
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
203
203
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -242,7 +242,7 @@
|
|
|
242
242
|
* See the License for the specific language governing permissions and
|
|
243
243
|
* limitations under the License.
|
|
244
244
|
* =============================================================================
|
|
245
|
-
*/function ko(n){return n.kernelName!=null}class al{constructor(){this.registeredVariables={},this.nextTapeNodeId=0,this.numBytes=0,this.numTensors=0,this.numStringTensors=0,this.numDataBuffers=0,this.gradientDepth=0,this.kernelDepth=0,this.scopeStack=[],this.numDataMovesStack=[],this.nextScopeId=0,this.tensorInfo=new WeakMap,this.profiling=!1,this.activeProfile={newBytes:0,newTensors:0,peakBytes:0,kernels:[],result:null,get kernelNames(){return Array.from(new Set(this.kernels.map(t=>t.name)))}}}dispose(){for(const t in this.registeredVariables)this.registeredVariables[t].dispose()}}class Rn{constructor(t){this.ENV=t,this.registry={},this.registryFactory={},this.pendingBackendInitId=0,this.state=new al}async ready(){if(this.pendingBackendInit!=null)return this.pendingBackendInit.then(()=>{});if(this.backendInstance!=null)return;const t=this.getSortedBackends();for(let e=0;e<t.length;e++){const s=t[e];if(await this.initializeBackend(s).success){await this.setBackend(s);return}}throw new Error("Could not initialize any backends, all backend initializations failed.")}get backend(){if(this.pendingBackendInit!=null)throw new Error(`Backend '${this.backendName}' has not yet been initialized. Make sure to await tf.ready() or await tf.setBackend() before calling other methods`);if(this.backendInstance==null){const{name:t,asyncInit:e}=this.initializeBackendsAndReturnBest();if(e)throw new Error(`The highest priority backend '${t}' has not yet been initialized. Make sure to await tf.ready() or await tf.setBackend() before calling other methods`);this.setBackend(t)}return this.backendInstance}backendNames(){return Object.keys(this.registryFactory)}findBackend(t){if(!(t in this.registry))if(t in this.registryFactory){const{asyncInit:e}=this.initializeBackend(t);if(e)return null}else return null;return this.registry[t]}findBackendFactory(t){return t in this.registryFactory?this.registryFactory[t].factory:null}registerBackend(t,e,s=1){return t in this.registryFactory?(kn(`${t} backend was already registered. Reusing existing backend factory.`),!1):(this.registryFactory[t]={factory:e,priority:s},!0)}async setBackend(t){if(this.registryFactory[t]==null)throw new Error(`Backend name '${t}' not found in registry`);if(this.backendName=t,this.registry[t]==null){this.backendInstance=null;const{success:e,asyncInit:s}=this.initializeBackend(t);if(!(s?await e:e))return!1}return this.backendInstance=this.registry[t],this.setupRegisteredKernels(),this.profiler=new xp(this.backendInstance),!0}setupRegisteredKernels(){Ma(this.backendName).forEach(e=>{e.setupFunc!=null&&e.setupFunc(this.backendInstance)})}disposeRegisteredKernels(t){Ma(t).forEach(s=>{s.disposeFunc!=null&&s.disposeFunc(this.registry[t])})}initializeBackend(t){const e=this.registryFactory[t];if(e==null)throw new Error(`Cannot initialize backend ${t}, no registration found.`);try{const s=e.factory();if(s&&!(s instanceof ba)&&typeof s.then=="function"){const r=++this.pendingBackendInitId,o=s.then(i=>r<this.pendingBackendInitId?!1:(this.registry[t]=i,this.pendingBackendInit=null,!0)).catch(i=>(r<this.pendingBackendInitId||(this.pendingBackendInit=null,kn(`Initialization of backend ${t} failed`),kn(i.stack||i.message)),!1));return this.pendingBackendInit=o,{success:o,asyncInit:!0}}else return this.registry[t]=s,{success:!0,asyncInit:!1}}catch(s){return kn(`Initialization of backend ${t} failed`),kn(s.stack||s.message),{success:!1,asyncInit:!1}}}removeBackend(t){if(!(t in this.registryFactory))throw new Error(`${t} backend not found in registry`);this.backendName===t&&this.pendingBackendInit!=null&&this.pendingBackendInitId++,t in this.registry&&(this.disposeRegisteredKernels(t),this.registry[t].dispose(),delete this.registry[t]),delete this.registryFactory[t],this.backendName===t&&(this.pendingBackendInit=null,this.backendName=null,this.backendInstance=null)}getSortedBackends(){if(Object.keys(this.registryFactory).length===0)throw new Error("No backend found in registry.");return Object.keys(this.registryFactory).sort((t,e)=>this.registryFactory[e].priority-this.registryFactory[t].priority)}initializeBackendsAndReturnBest(){const t=this.getSortedBackends();for(let e=0;e<t.length;e++){const s=t[e],{success:r,asyncInit:o}=this.initializeBackend(s);if(o||r)return{name:s,asyncInit:o}}throw new Error("Could not initialize any backends, all backend initializations failed.")}moveData(t,e){const s=this.state.tensorInfo.get(e),r=s.backend,o=this.readSync(e),i=r.refCount(e);r.disposeData(e,!0),s.backend=t,t.move(e,o,s.shape,s.dtype,i),this.shouldCheckForMemLeaks()&&this.state.numDataMovesStack[this.state.numDataMovesStack.length-1]++}tidy(t,e){let s=null;if(e==null){if(typeof t!="function")throw new Error("Please provide a function to tidy()");e=t}else{if(typeof t!="string"&&!(t instanceof String))throw new Error("When calling with two arguments, the first argument to tidy() must be a string");if(typeof e!="function")throw new Error("When calling with two arguments, the 2nd argument to tidy() must be a function");s=t}let r;return this.scopedRun(()=>this.startScope(s),()=>this.endScope(r),()=>(r=e(),r instanceof Promise&&console.error("Cannot return a Promise inside of tidy."),r))}scopedRun(t,e,s){t();try{const r=s();return e(),r}catch(r){throw e(),r}}nextTensorId(){return Rn.nextTensorId++}nextVariableId(){return Rn.nextVariableId++}clone(t){const e=A.runKernel(go,{x:t}),s={x:t},r=i=>({x:()=>{const a="float32",l={x:i},u={dtype:a};return A.runKernel(Ea,l,u)}}),o=[];return this.addTapeNode(this.state.activeScope.name,s,[e],r,o,{}),e}runKernel(t,e,s){if(this.backendName==null&&this.backend,!(Pa(t,this.backendName)!=null))throw new Error(`Kernel '${t}' not registered for backend '${this.backendName}'`);return this.runKernelFunc({kernelName:t,inputs:e,attrs:s})}shouldCheckForMemLeaks(){return this.ENV.getBool("IS_TEST")}checkKernelForMemLeak(t,e,s){const r=this.backend.numDataIds();let o=0;s.forEach(l=>{o+=l.dtype==="complex64"?3:1});const i=this.state.numDataMovesStack[this.state.numDataMovesStack.length-1],a=r-e-o-i;if(a>0)throw new Error(`Backend '${this.backendName}' has an internal memory leak (${a} data ids) after running '${t}'`)}runKernelFunc(t){let e,s=[];const r=this.isTapeOn(),o=this.state.numBytes,i=this.state.numTensors;this.shouldCheckForMemLeaks()&&this.state.numDataMovesStack.push(0);let a;this.backendName==null&&this.backend;let l;const u=ko(t)?t.kernelName:this.state.activeScope!=null?this.state.activeScope.name:"";if(ko(t)){const{kernelName:p,inputs:g,attrs:m}=t;this.backendName==null&&this.backend;const b=Pa(p,this.backendName);w(b!=null,()=>`Cannot find registered kernel '${p}' for backend '${this.backendName}'`),a=()=>{const y=this.backend.numDataIds();l=b.kernelFunc({inputs:g,attrs:m,backend:this.backend});const S=Array.isArray(l)?l:[l];this.shouldCheckForMemLeaks()&&this.checkKernelForMemLeak(p,y,S);const x=S.map($=>$.rank!=null?$:this.makeTensorFromTensorInfo($));if(r){const $=this.getTensorsForGradient(p,g,x);s=this.saveTensorsForBackwardMode($)}return x}}else{const{forwardFunc:p}=t,g=m=>{r&&(s=m.map(b=>this.keep(this.clone(b))))};a=()=>{const m=this.backend.numDataIds();l=this.tidy(()=>p(this.backend,g));const b=Array.isArray(l)?l:[l];return this.shouldCheckForMemLeaks()&&this.checkKernelForMemLeak(u,m,b),b}}const{inputs:c,attrs:h}=t,f=ko(t)?null:t.backwardsFunc;let d;return this.scopedRun(()=>this.state.kernelDepth++,()=>this.state.kernelDepth--,()=>{!this.ENV.getBool("DEBUG")&&!this.state.profiling?e=a():(d=this.profiler.profileKernel(u,c,()=>a()),this.ENV.getBool("DEBUG")&&this.profiler.logKernelProfile(d),e=d.outputs)}),r&&this.addTapeNode(u,c,e,f,s,h),this.state.profiling&&this.state.activeProfile.kernels.push({name:u,bytesAdded:this.state.numBytes-o,totalBytesSnapshot:this.state.numBytes,tensorsAdded:this.state.numTensors-i,totalTensorsSnapshot:this.state.numTensors,inputShapes:Object.keys(c).map(p=>c[p]!=null?c[p].shape:null),outputShapes:e.map(p=>p.shape),kernelTimeMs:d.timeMs,extraInfo:d.extraInfo}),Array.isArray(l)?e:e[0]}saveTensorsForBackwardMode(t){return t.map(s=>this.keep(this.clone(s)))}getTensorsForGradient(t,e,s){const r=La(t);if(r!=null){const o=r.inputsToSave||[],i=r.outputsToSave||[];let a;r.saveAllInputs?(w(Array.isArray(e),()=>"saveAllInputs is true, expected inputs to be an array."),a=Object.keys(e).map(u=>e[u])):a=o.map(u=>e[u]);const l=s.filter((u,c)=>i[c]);return a.concat(l)}return[]}makeTensor(t,e,s,r){if(t==null)throw new Error("Values passed to engine.makeTensor() are null");s=s||"float32",r=r||this.backend;let o=t;s==="string"&&Ws(t[0])&&(o=t.map(l=>nn(l)));const i=r.write(o,e,s),a=new At(e,s,i,this.nextTensorId());if(this.trackTensor(a,r),s==="string"){const l=this.state.tensorInfo.get(i),u=Nf(o);this.state.numBytes+=u-l.bytes,l.bytes=u}return a}makeTensorFromDataId(t,e,s,r){s=s||"float32";const o={dataId:t,shape:e,dtype:s};return this.makeTensorFromTensorInfo(o,r)}makeTensorFromTensorInfo(t,e){const{dataId:s,shape:r,dtype:o}=t,i=new At(r,o,s,this.nextTensorId());return this.trackTensor(i,e),i}makeVariable(t,e=!0,s,r){s=s||this.nextVariableId().toString(),r!=null&&r!==t.dtype&&(t=t.cast(r));const o=new Js(t,e,s,this.nextTensorId());if(this.state.registeredVariables[o.name]!=null)throw new Error(`Variable with name ${o.name} was already registered`);return this.state.registeredVariables[o.name]=o,this.incRef(o,this.backend),o}trackTensor(t,e){this.state.numTensors++,t.dtype==="string"&&this.state.numStringTensors++;let s=0;t.dtype!=="complex64"&&t.dtype!=="string"&&(s=t.size*io(t.dtype)),this.state.numBytes+=s,this.state.tensorInfo.has(t.dataId)||(this.state.numDataBuffers++,this.state.tensorInfo.set(t.dataId,{backend:e||this.backend,dtype:t.dtype,shape:t.shape,bytes:s})),t instanceof Js||this.track(t)}incRef(t,e){this.trackTensor(t,e),this.backend.incRef(t.dataId)}removeDataId(t,e){this.state.tensorInfo.has(t)&&this.state.tensorInfo.get(t).backend===e&&(this.state.tensorInfo.delete(t),this.state.numDataBuffers--)}disposeTensor(t){if(!this.state.tensorInfo.has(t.dataId))return;const e=this.state.tensorInfo.get(t.dataId);if(this.state.numTensors--,t.dtype==="string"&&(this.state.numStringTensors--,this.state.numBytes-=e.bytes),t.dtype!=="complex64"&&t.dtype!=="string"){const s=t.size*io(t.dtype);this.state.numBytes-=s}e.backend.disposeData(t.dataId)&&this.removeDataId(t.dataId,e.backend)}disposeVariables(){for(const t in this.state.registeredVariables){const e=this.state.registeredVariables[t];this.disposeVariable(e)}}disposeVariable(t){this.disposeTensor(t),this.state.registeredVariables[t.name]!=null&&delete this.state.registeredVariables[t.name]}memory(){const t=this.backend.memory();return t.numTensors=this.state.numTensors,t.numDataBuffers=this.state.numDataBuffers,t.numBytes=this.state.numBytes,this.state.numStringTensors>0&&(t.unreliable=!0,t.reasons==null&&(t.reasons=[]),t.reasons.push("Memory usage by string tensors is approximate (2 bytes per character)")),t}async profile(t){this.state.profiling=!0;const e=this.state.numBytes,s=this.state.numTensors;this.state.activeProfile.kernels=[],this.state.activeProfile.result=await t(),this.state.profiling=!1,this.state.activeProfile.peakBytes=Math.max(...this.state.activeProfile.kernels.map(r=>r.totalBytesSnapshot)),this.state.activeProfile.newBytes=this.state.numBytes-e,this.state.activeProfile.newTensors=this.state.numTensors-s;for(const r of this.state.activeProfile.kernels)r.kernelTimeMs=await r.kernelTimeMs,r.extraInfo=await r.extraInfo;return this.state.activeProfile}isTapeOn(){return this.state.gradientDepth>0&&this.state.kernelDepth===0}addTapeNode(t,e,s,r,o,i){const a={id:this.state.nextTapeNodeId++,kernelName:t,inputs:e,outputs:s,saved:o},l=La(t);l!=null&&(r=l.gradFunc),r!=null&&(a.gradient=u=>(u=u.map((c,h)=>{if(c==null){const f=s[h],d=ze(f.size,f.dtype);return this.makeTensor(d,f.shape,f.dtype)}return c}),r(u.length>1?u:u[0],o,i))),this.state.activeTape.push(a)}keep(t){return t.kept=!0,t}startTape(){this.state.gradientDepth===0&&(this.state.activeTape=[]),this.state.gradientDepth++}endTape(){this.state.gradientDepth--}startScope(t){const e={track:[],name:"unnamed scope",id:this.state.nextScopeId++};t&&(e.name=t),this.state.scopeStack.push(e),this.state.activeScope=e}endScope(t){const e=ol(t),s=new Set(e.map(o=>o.id));for(let o=0;o<this.state.activeScope.track.length;o++){const i=this.state.activeScope.track[o];!i.kept&&!s.has(i.id)&&i.dispose()}const r=this.state.scopeStack.pop();this.state.activeScope=this.state.scopeStack.length===0?null:this.state.scopeStack[this.state.scopeStack.length-1],e.forEach(o=>{!o.kept&&o.scopeId===r.id&&this.track(o)})}gradients(t,e,s,r=!1){if(w(e.length>0,()=>"gradients() received an empty list of xs."),s!=null&&s.dtype!=="float32")throw new Error(`dy must have 'float32' dtype, but has '${s.dtype}'`);const o=this.scopedRun(()=>this.startTape(),()=>this.endTape(),()=>this.tidy("forward",t));w(o instanceof At,()=>"The result y returned by f() must be a tensor.");const i=vp(this.state.activeTape,e,o);if(!r&&i.length===0&&e.length>0)throw new Error("Cannot compute gradient of y=f(x) with respect to x. Make sure that the f you passed encloses all operations that lead from x to y.");return this.tidy("backward",()=>{const a={};a[o.id]=s??Dp(o.shape),Ip(a,i,u=>this.tidy(u),Rp);const l=e.map(u=>a[u.id]);return this.state.gradientDepth===0&&(this.state.activeTape.forEach(u=>{for(const c of u.saved)c.dispose()}),this.state.activeTape=null),{value:o,grads:l}})}customGrad(t){return w(lo(t),()=>"The f passed in customGrad(f) must be a function."),(...e)=>{w(e.every(a=>a instanceof At),()=>"The args passed in customGrad(f)(x1, x2,...) must all be tensors");let s;const r={};e.forEach((a,l)=>{r[l]=a});const o=(a,l)=>(s=t(...e,l),w(s.value instanceof At,()=>"The function f passed in customGrad(f) must return an object where `obj.value` is a tensor"),w(lo(s.gradFunc),()=>"The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function."),s.value),i=(a,l)=>{const u=s.gradFunc(a,l),c=Array.isArray(u)?u:[u];w(c.length===e.length,()=>"The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function that returns the same number of tensors as inputs passed to f(...)."),w(c.every(f=>f instanceof At),()=>"The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function that returns a list of only tensors.");const h={};return c.forEach((f,d)=>{h[d]=()=>f}),h};return this.runKernelFunc({forwardFunc:o,backwardsFunc:i,inputs:r})}}readSync(t){return this.state.tensorInfo.get(t).backend.readSync(t)}read(t){return this.state.tensorInfo.get(t).backend.read(t)}readToGPU(t,e){return this.state.tensorInfo.get(t).backend.readToGPU(t,e)}async time(t){const e=Nn(),s=await this.backend.time(t);return s.wallMs=Nn()-e,s}track(t){return this.state.activeScope!=null&&(t.scopeId=this.state.activeScope.id,this.state.activeScope.track.push(t)),t}get registeredVariables(){return this.state.registeredVariables}reset(){this.pendingBackendInitId++,this.state.dispose(),this.ENV.reset(),this.state=new al;for(const t in this.registry)this.disposeRegisteredKernels(t),this.registry[t].dispose(),delete this.registry[t];this.backendName=null,this.backendInstance=null,this.pendingBackendInit=null}}Rn.nextTensorId=0,Rn.nextVariableId=0;function Dp(n){const t=Sa(z(n),"float32");return A.makeTensor(t,n,"float32")}function ll(){const n=Ia();if(n._tfengine==null){const t=new Rf(n);n._tfengine=new Rn(t)}return Of(n._tfengine.ENV),Cp(()=>n._tfengine),n._tfengine}const A=ll();function Rp(n,t){const e={a:n,b:t};return A.runKernel(Aa,e)}/**
|
|
245
|
+
*/function ko(n){return n.kernelName!=null}class al{constructor(){this.registeredVariables={},this.nextTapeNodeId=0,this.numBytes=0,this.numTensors=0,this.numStringTensors=0,this.numDataBuffers=0,this.gradientDepth=0,this.kernelDepth=0,this.scopeStack=[],this.numDataMovesStack=[],this.nextScopeId=0,this.tensorInfo=new WeakMap,this.profiling=!1,this.activeProfile={newBytes:0,newTensors:0,peakBytes:0,kernels:[],result:null,get kernelNames(){return Array.from(new Set(this.kernels.map(t=>t.name)))}}}dispose(){for(const t in this.registeredVariables)this.registeredVariables[t].dispose()}}class Rn{constructor(t){this.ENV=t,this.registry={},this.registryFactory={},this.pendingBackendInitId=0,this.state=new al}async ready(){if(this.pendingBackendInit!=null)return this.pendingBackendInit.then(()=>{});if(this.backendInstance!=null)return;const t=this.getSortedBackends();for(let e=0;e<t.length;e++){const s=t[e];if(await this.initializeBackend(s).success){await this.setBackend(s);return}}throw new Error("Could not initialize any backends, all backend initializations failed.")}get backend(){if(this.pendingBackendInit!=null)throw new Error(`Backend '${this.backendName}' has not yet been initialized. Make sure to await tf.ready() or await tf.setBackend() before calling other methods`);if(this.backendInstance==null){const{name:t,asyncInit:e}=this.initializeBackendsAndReturnBest();if(e)throw new Error(`The highest priority backend '${t}' has not yet been initialized. Make sure to await tf.ready() or await tf.setBackend() before calling other methods`);this.setBackend(t)}return this.backendInstance}backendNames(){return Object.keys(this.registryFactory)}findBackend(t){if(!(t in this.registry))if(t in this.registryFactory){const{asyncInit:e}=this.initializeBackend(t);if(e)return null}else return null;return this.registry[t]}findBackendFactory(t){return t in this.registryFactory?this.registryFactory[t].factory:null}registerBackend(t,e,s=1){return t in this.registryFactory?(kn(`${t} backend was already registered. Reusing existing backend factory.`),!1):(this.registryFactory[t]={factory:e,priority:s},!0)}async setBackend(t){if(this.registryFactory[t]==null)throw new Error(`Backend name '${t}' not found in registry`);if(this.backendName=t,this.registry[t]==null){this.backendInstance=null;const{success:e,asyncInit:s}=this.initializeBackend(t);if(!(s?await e:e))return!1}return this.backendInstance=this.registry[t],this.setupRegisteredKernels(),this.profiler=new xp(this.backendInstance),!0}setupRegisteredKernels(){Ma(this.backendName).forEach(e=>{e.setupFunc!=null&&e.setupFunc(this.backendInstance)})}disposeRegisteredKernels(t){Ma(t).forEach(s=>{s.disposeFunc!=null&&s.disposeFunc(this.registry[t])})}initializeBackend(t){const e=this.registryFactory[t];if(e==null)throw new Error(`Cannot initialize backend ${t}, no registration found.`);try{const s=e.factory();if(s&&!(s instanceof ba)&&typeof s.then=="function"){const r=++this.pendingBackendInitId,o=s.then(i=>r<this.pendingBackendInitId?!1:(this.registry[t]=i,this.pendingBackendInit=null,!0)).catch(i=>(r<this.pendingBackendInitId||(this.pendingBackendInit=null,kn(`Initialization of backend ${t} failed`),kn(i.stack||i.message)),!1));return this.pendingBackendInit=o,{success:o,asyncInit:!0}}else return this.registry[t]=s,{success:!0,asyncInit:!1}}catch(s){return kn(`Initialization of backend ${t} failed`),kn(s.stack||s.message),{success:!1,asyncInit:!1}}}removeBackend(t){if(!(t in this.registryFactory))throw new Error(`${t} backend not found in registry`);this.backendName===t&&this.pendingBackendInit!=null&&this.pendingBackendInitId++,t in this.registry&&(this.disposeRegisteredKernels(t),this.registry[t].dispose(),delete this.registry[t]),delete this.registryFactory[t],this.backendName===t&&(this.pendingBackendInit=null,this.backendName=null,this.backendInstance=null)}getSortedBackends(){if(Object.keys(this.registryFactory).length===0)throw new Error("No backend found in registry.");return Object.keys(this.registryFactory).sort((t,e)=>this.registryFactory[e].priority-this.registryFactory[t].priority)}initializeBackendsAndReturnBest(){const t=this.getSortedBackends();for(let e=0;e<t.length;e++){const s=t[e],{success:r,asyncInit:o}=this.initializeBackend(s);if(o||r)return{name:s,asyncInit:o}}throw new Error("Could not initialize any backends, all backend initializations failed.")}moveData(t,e){const s=this.state.tensorInfo.get(e),r=s.backend,o=this.readSync(e),i=r.refCount(e);r.disposeData(e,!0),s.backend=t,t.move(e,o,s.shape,s.dtype,i),this.shouldCheckForMemLeaks()&&this.state.numDataMovesStack[this.state.numDataMovesStack.length-1]++}tidy(t,e){let s=null;if(e==null){if(typeof t!="function")throw new Error("Please provide a function to tidy()");e=t}else{if(typeof t!="string"&&!(t instanceof String))throw new Error("When calling with two arguments, the first argument to tidy() must be a string");if(typeof e!="function")throw new Error("When calling with two arguments, the 2nd argument to tidy() must be a function");s=t}let r;return this.scopedRun(()=>this.startScope(s),()=>this.endScope(r),()=>(r=e(),r instanceof Promise&&console.error("Cannot return a Promise inside of tidy."),r))}scopedRun(t,e,s){t();try{const r=s();return e(),r}catch(r){throw e(),r}}nextTensorId(){return Rn.nextTensorId++}nextVariableId(){return Rn.nextVariableId++}clone(t){const e=A.runKernel(go,{x:t}),s={x:t},r=i=>({x:()=>{const a="float32",l={x:i},u={dtype:a};return A.runKernel(Ea,l,u)}}),o=[];return this.addTapeNode(this.state.activeScope.name,s,[e],r,o,{}),e}runKernel(t,e,s){if(this.backendName==null&&this.backend,!(Pa(t,this.backendName)!=null))throw new Error(`Kernel '${t}' not registered for backend '${this.backendName}'`);return this.runKernelFunc({kernelName:t,inputs:e,attrs:s})}shouldCheckForMemLeaks(){return this.ENV.getBool("IS_TEST")}checkKernelForMemLeak(t,e,s){const r=this.backend.numDataIds();let o=0;s.forEach(l=>{o+=l.dtype==="complex64"?3:1});const i=this.state.numDataMovesStack[this.state.numDataMovesStack.length-1],a=r-e-o-i;if(a>0)throw new Error(`Backend '${this.backendName}' has an internal memory leak (${a} data ids) after running '${t}'`)}runKernelFunc(t){let e,s=[];const r=this.isTapeOn(),o=this.state.numBytes,i=this.state.numTensors;this.shouldCheckForMemLeaks()&&this.state.numDataMovesStack.push(0);let a;this.backendName==null&&this.backend;let l;const u=ko(t)?t.kernelName:this.state.activeScope!=null?this.state.activeScope.name:"";if(ko(t)){const{kernelName:p,inputs:g,attrs:m}=t;this.backendName==null&&this.backend;const b=Pa(p,this.backendName);w(b!=null,()=>`Cannot find registered kernel '${p}' for backend '${this.backendName}'`),a=()=>{const y=this.backend.numDataIds();l=b.kernelFunc({inputs:g,attrs:m,backend:this.backend});const S=Array.isArray(l)?l:[l];this.shouldCheckForMemLeaks()&&this.checkKernelForMemLeak(p,y,S);const x=S.map($=>$.rank!=null?$:this.makeTensorFromTensorInfo($));if(r){const $=this.getTensorsForGradient(p,g,x);s=this.saveTensorsForBackwardMode($)}return x}}else{const{forwardFunc:p}=t,g=m=>{r&&(s=m.map(b=>this.keep(this.clone(b))))};a=()=>{const m=this.backend.numDataIds();l=this.tidy(()=>p(this.backend,g));const b=Array.isArray(l)?l:[l];return this.shouldCheckForMemLeaks()&&this.checkKernelForMemLeak(u,m,b),b}}const{inputs:c,attrs:h}=t,f=ko(t)?null:t.backwardsFunc;let d;return this.scopedRun(()=>this.state.kernelDepth++,()=>this.state.kernelDepth--,()=>{!this.ENV.getBool("DEBUG")&&!this.state.profiling?e=a():(d=this.profiler.profileKernel(u,c,()=>a()),this.ENV.getBool("DEBUG")&&this.profiler.logKernelProfile(d),e=d.outputs)}),r&&this.addTapeNode(u,c,e,f,s,h),this.state.profiling&&this.state.activeProfile.kernels.push({name:u,bytesAdded:this.state.numBytes-o,totalBytesSnapshot:this.state.numBytes,tensorsAdded:this.state.numTensors-i,totalTensorsSnapshot:this.state.numTensors,inputShapes:Object.keys(c).map(p=>c[p]!=null?c[p].shape:null),outputShapes:e.map(p=>p.shape),kernelTimeMs:d.timeMs,extraInfo:d.extraInfo}),Array.isArray(l)?e:e[0]}saveTensorsForBackwardMode(t){return t.map(s=>this.keep(this.clone(s)))}getTensorsForGradient(t,e,s){const r=La(t);if(r!=null){const o=r.inputsToSave||[],i=r.outputsToSave||[];let a;r.saveAllInputs?(w(Array.isArray(e),()=>"saveAllInputs is true, expected inputs to be an array."),a=Object.keys(e).map(u=>e[u])):a=o.map(u=>e[u]);const l=s.filter((u,c)=>i[c]);return a.concat(l)}return[]}makeTensor(t,e,s,r){if(t==null)throw new Error("Values passed to engine.makeTensor() are null");s=s||"float32",r=r||this.backend;let o=t;s==="string"&&Gs(t[0])&&(o=t.map(l=>sn(l)));const i=r.write(o,e,s),a=new At(e,s,i,this.nextTensorId());if(this.trackTensor(a,r),s==="string"){const l=this.state.tensorInfo.get(i),u=Nf(o);this.state.numBytes+=u-l.bytes,l.bytes=u}return a}makeTensorFromDataId(t,e,s,r){s=s||"float32";const o={dataId:t,shape:e,dtype:s};return this.makeTensorFromTensorInfo(o,r)}makeTensorFromTensorInfo(t,e){const{dataId:s,shape:r,dtype:o}=t,i=new At(r,o,s,this.nextTensorId());return this.trackTensor(i,e),i}makeVariable(t,e=!0,s,r){s=s||this.nextVariableId().toString(),r!=null&&r!==t.dtype&&(t=t.cast(r));const o=new Zs(t,e,s,this.nextTensorId());if(this.state.registeredVariables[o.name]!=null)throw new Error(`Variable with name ${o.name} was already registered`);return this.state.registeredVariables[o.name]=o,this.incRef(o,this.backend),o}trackTensor(t,e){this.state.numTensors++,t.dtype==="string"&&this.state.numStringTensors++;let s=0;t.dtype!=="complex64"&&t.dtype!=="string"&&(s=t.size*io(t.dtype)),this.state.numBytes+=s,this.state.tensorInfo.has(t.dataId)||(this.state.numDataBuffers++,this.state.tensorInfo.set(t.dataId,{backend:e||this.backend,dtype:t.dtype,shape:t.shape,bytes:s})),t instanceof Zs||this.track(t)}incRef(t,e){this.trackTensor(t,e),this.backend.incRef(t.dataId)}removeDataId(t,e){this.state.tensorInfo.has(t)&&this.state.tensorInfo.get(t).backend===e&&(this.state.tensorInfo.delete(t),this.state.numDataBuffers--)}disposeTensor(t){if(!this.state.tensorInfo.has(t.dataId))return;const e=this.state.tensorInfo.get(t.dataId);if(this.state.numTensors--,t.dtype==="string"&&(this.state.numStringTensors--,this.state.numBytes-=e.bytes),t.dtype!=="complex64"&&t.dtype!=="string"){const s=t.size*io(t.dtype);this.state.numBytes-=s}e.backend.disposeData(t.dataId)&&this.removeDataId(t.dataId,e.backend)}disposeVariables(){for(const t in this.state.registeredVariables){const e=this.state.registeredVariables[t];this.disposeVariable(e)}}disposeVariable(t){this.disposeTensor(t),this.state.registeredVariables[t.name]!=null&&delete this.state.registeredVariables[t.name]}memory(){const t=this.backend.memory();return t.numTensors=this.state.numTensors,t.numDataBuffers=this.state.numDataBuffers,t.numBytes=this.state.numBytes,this.state.numStringTensors>0&&(t.unreliable=!0,t.reasons==null&&(t.reasons=[]),t.reasons.push("Memory usage by string tensors is approximate (2 bytes per character)")),t}async profile(t){this.state.profiling=!0;const e=this.state.numBytes,s=this.state.numTensors;this.state.activeProfile.kernels=[],this.state.activeProfile.result=await t(),this.state.profiling=!1,this.state.activeProfile.peakBytes=Math.max(...this.state.activeProfile.kernels.map(r=>r.totalBytesSnapshot)),this.state.activeProfile.newBytes=this.state.numBytes-e,this.state.activeProfile.newTensors=this.state.numTensors-s;for(const r of this.state.activeProfile.kernels)r.kernelTimeMs=await r.kernelTimeMs,r.extraInfo=await r.extraInfo;return this.state.activeProfile}isTapeOn(){return this.state.gradientDepth>0&&this.state.kernelDepth===0}addTapeNode(t,e,s,r,o,i){const a={id:this.state.nextTapeNodeId++,kernelName:t,inputs:e,outputs:s,saved:o},l=La(t);l!=null&&(r=l.gradFunc),r!=null&&(a.gradient=u=>(u=u.map((c,h)=>{if(c==null){const f=s[h],d=ze(f.size,f.dtype);return this.makeTensor(d,f.shape,f.dtype)}return c}),r(u.length>1?u:u[0],o,i))),this.state.activeTape.push(a)}keep(t){return t.kept=!0,t}startTape(){this.state.gradientDepth===0&&(this.state.activeTape=[]),this.state.gradientDepth++}endTape(){this.state.gradientDepth--}startScope(t){const e={track:[],name:"unnamed scope",id:this.state.nextScopeId++};t&&(e.name=t),this.state.scopeStack.push(e),this.state.activeScope=e}endScope(t){const e=ol(t),s=new Set(e.map(o=>o.id));for(let o=0;o<this.state.activeScope.track.length;o++){const i=this.state.activeScope.track[o];!i.kept&&!s.has(i.id)&&i.dispose()}const r=this.state.scopeStack.pop();this.state.activeScope=this.state.scopeStack.length===0?null:this.state.scopeStack[this.state.scopeStack.length-1],e.forEach(o=>{!o.kept&&o.scopeId===r.id&&this.track(o)})}gradients(t,e,s,r=!1){if(w(e.length>0,()=>"gradients() received an empty list of xs."),s!=null&&s.dtype!=="float32")throw new Error(`dy must have 'float32' dtype, but has '${s.dtype}'`);const o=this.scopedRun(()=>this.startTape(),()=>this.endTape(),()=>this.tidy("forward",t));w(o instanceof At,()=>"The result y returned by f() must be a tensor.");const i=vp(this.state.activeTape,e,o);if(!r&&i.length===0&&e.length>0)throw new Error("Cannot compute gradient of y=f(x) with respect to x. Make sure that the f you passed encloses all operations that lead from x to y.");return this.tidy("backward",()=>{const a={};a[o.id]=s??Dp(o.shape),Ip(a,i,u=>this.tidy(u),Rp);const l=e.map(u=>a[u.id]);return this.state.gradientDepth===0&&(this.state.activeTape.forEach(u=>{for(const c of u.saved)c.dispose()}),this.state.activeTape=null),{value:o,grads:l}})}customGrad(t){return w(lo(t),()=>"The f passed in customGrad(f) must be a function."),(...e)=>{w(e.every(a=>a instanceof At),()=>"The args passed in customGrad(f)(x1, x2,...) must all be tensors");let s;const r={};e.forEach((a,l)=>{r[l]=a});const o=(a,l)=>(s=t(...e,l),w(s.value instanceof At,()=>"The function f passed in customGrad(f) must return an object where `obj.value` is a tensor"),w(lo(s.gradFunc),()=>"The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function."),s.value),i=(a,l)=>{const u=s.gradFunc(a,l),c=Array.isArray(u)?u:[u];w(c.length===e.length,()=>"The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function that returns the same number of tensors as inputs passed to f(...)."),w(c.every(f=>f instanceof At),()=>"The function f passed in customGrad(f) must return an object where `obj.gradFunc` is a function that returns a list of only tensors.");const h={};return c.forEach((f,d)=>{h[d]=()=>f}),h};return this.runKernelFunc({forwardFunc:o,backwardsFunc:i,inputs:r})}}readSync(t){return this.state.tensorInfo.get(t).backend.readSync(t)}read(t){return this.state.tensorInfo.get(t).backend.read(t)}readToGPU(t,e){return this.state.tensorInfo.get(t).backend.readToGPU(t,e)}async time(t){const e=Nn(),s=await this.backend.time(t);return s.wallMs=Nn()-e,s}track(t){return this.state.activeScope!=null&&(t.scopeId=this.state.activeScope.id,this.state.activeScope.track.push(t)),t}get registeredVariables(){return this.state.registeredVariables}reset(){this.pendingBackendInitId++,this.state.dispose(),this.ENV.reset(),this.state=new al;for(const t in this.registry)this.disposeRegisteredKernels(t),this.registry[t].dispose(),delete this.registry[t];this.backendName=null,this.backendInstance=null,this.pendingBackendInit=null}}Rn.nextTensorId=0,Rn.nextVariableId=0;function Dp(n){const t=Sa(z(n),"float32");return A.makeTensor(t,n,"float32")}function ll(){const n=Ia();if(n._tfengine==null){const t=new Rf(n);n._tfengine=new Rn(t)}return Of(n._tfengine.ENV),Cp(()=>n._tfengine),n._tfengine}const A=ll();function Rp(n,t){const e={a:n,b:t};return A.runKernel(Aa,e)}/**
|
|
246
246
|
* @license
|
|
247
247
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
248
248
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -257,7 +257,7 @@
|
|
|
257
257
|
* See the License for the specific language governing permissions and
|
|
258
258
|
* limitations under the License.
|
|
259
259
|
* =============================================================================
|
|
260
|
-
*/function
|
|
260
|
+
*/function Qs(n,t){let e=n;if(ne(n))return t==="string"?[]:[n.length];if(sl(n)){const r=n.channels||"RGBA";return[n.height,n.width*r.length]}else if(rl(n))return[n.buffer.size/(t==null?4:io(t))];if(!Array.isArray(n))return[];const s=[];for(;Array.isArray(e)||ne(e)&&t!=="string";)s.push(e.length),e=e[0];return Array.isArray(n)&&W().getBool("TENSORLIKE_CHECK_SHAPE_CONSISTENCY")&&ul(n,s,[]),s}function ul(n,t,e){if(e=e||[],!Array.isArray(n)&&!ne(n)){w(t.length===0,()=>`Element arr[${e.join("][")}] is a primitive, but should be an array/TypedArray of ${t[0]} elements`);return}w(t.length>0,()=>`Element arr[${e.join("][")}] should be a primitive, but is an array of ${n.length} elements`),w(n.length===t[0],()=>`Element arr[${e.join("][")}] should have ${t[0]} elements, but has ${n.length} elements`);const s=t.slice(1);for(let r=0;r<n.length;++r)ul(n[r],s,e.concat(r))}function cl(n,t,e,s){if(n!=="string_or_numeric"){if(n==null)throw new Error("Expected dtype cannot be null.");if(n!=="numeric"&&n!==t||n==="numeric"&&t==="string")throw new Error(`Argument '${e}' passed to '${s}' must be ${n} tensor, but got ${t} tensor`)}}function v(n,t,e,s="numeric"){if(n instanceof el())return cl(s,n.dtype,t,e),n;let r=as(n);if(r!=="string"&&["bool","int32","float32"].indexOf(s)>=0&&(r=s),cl(s,r,t,e),n==null||!ne(n)&&!Array.isArray(n)&&typeof n!="number"&&typeof n!="boolean"&&typeof n!="string"){const l=n==null?"null":n.constructor.name;throw new Error(`Argument '${t}' passed to '${e}' must be a Tensor or TensorLike, but got '${l}'`)}const o=Qs(n,r);!ne(n)&&!Array.isArray(n)&&(n=[n]);const a=r!=="string"?Ks(n,r):rn(n,[],!0);return A.makeTensor(a,o,r)}function hl(n,t,e,s="numeric"){if(!Array.isArray(n))throw new Error(`Argument ${t} passed to ${e} must be a \`Tensor[]\` or \`TensorLike[]\``);return n.map((o,i)=>v(o,`${t}[${i}]`,e,s))}/**
|
|
261
261
|
* @license
|
|
262
262
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
263
263
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -272,7 +272,7 @@
|
|
|
272
272
|
* See the License for the specific language governing permissions and
|
|
273
273
|
* limitations under the License.
|
|
274
274
|
* =============================================================================
|
|
275
|
-
*/function
|
|
275
|
+
*/function tr(n,t,e,s){if(s==null)s=as(n);else if(s==="complex64")throw new Error("Cannot construct a complex64 tensor directly. Please use tf.complex(real, imag).");if(rl(n)||sl(n)){if(s!=="float32"&&s!=="int32")throw new Error(`Creating tensor from GPU data only supports 'float32'|'int32' dtype, while the dtype is ${s}.`);return A.backend.createTensorFromGPUData(n,t||e,s)}if(!ne(n)&&!Array.isArray(n)&&typeof n!="number"&&typeof n!="boolean"&&typeof n!="string")throw new Error("values passed to tensor(values) must be a number/boolean/string or an array of numbers/booleans/strings, or a TypedArray");if(t!=null){Re(t);const r=z(t),o=z(e);w(r===o,()=>`Based on the provided shape, [${t}], the tensor should have ${r} values but has ${o}`);for(let i=0;i<e.length;++i){const a=e[i],l=i===e.length-1?a!==z(t.slice(i)):!0;w(e[i]===t[i]||!l,()=>`Error creating a new Tensor. Inferred shape (${e}) does not match the provided shape (${t}). `)}}return!ne(n)&&!Array.isArray(n)&&(n=[n]),t=t||e,n=s!=="string"?Ks(n,s):rn(n,[],!0),A.makeTensor(n,t,s)}/**
|
|
276
276
|
* @license
|
|
277
277
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
278
278
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -287,7 +287,7 @@
|
|
|
287
287
|
* See the License for the specific language governing permissions and
|
|
288
288
|
* limitations under the License.
|
|
289
289
|
* =============================================================================
|
|
290
|
-
*/function
|
|
290
|
+
*/function er(n,t,e){const s=Qs(n,e);return tr(n,t,s,e)}/**
|
|
291
291
|
* @license
|
|
292
292
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
293
293
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -302,7 +302,7 @@
|
|
|
302
302
|
* See the License for the specific language governing permissions and
|
|
303
303
|
* limitations under the License.
|
|
304
304
|
* =============================================================================
|
|
305
|
-
*/function Rt(n,t){ya(n);const e=
|
|
305
|
+
*/function Rt(n,t){ya(n);const e=Qs(n,t);if(e.length!==1)throw new Error("tensor1d() requires values to be a flat/TypedArray");return tr(n,null,e,t)}/**
|
|
306
306
|
* @license
|
|
307
307
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
308
308
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -377,7 +377,7 @@
|
|
|
377
377
|
* See the License for the specific language governing permissions and
|
|
378
378
|
* limitations under the License.
|
|
379
379
|
* =============================================================================
|
|
380
|
-
*/function Wp(n,t,e){const s=v(n,"x","slice4d");return w(s.rank===4,()=>`slice4d expects a rank-4 tensor, but got a rank-${s.rank} tensor`),Et(s,t,e)}const
|
|
380
|
+
*/function Wp(n,t,e){const s=v(n,"x","slice4d");return w(s.rank===4,()=>`slice4d expects a rank-4 tensor, but got a rank-${s.rank} tensor`),Et(s,t,e)}const hs=C({slice4d_:Wp});/**
|
|
381
381
|
* @license
|
|
382
382
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
383
383
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -392,7 +392,7 @@
|
|
|
392
392
|
* See the License for the specific language governing permissions and
|
|
393
393
|
* limitations under the License.
|
|
394
394
|
* =============================================================================
|
|
395
|
-
*/function Gp(n){const e={x:v(n,"x","clone","string_or_numeric")};return A.runKernel(go,e)}const
|
|
395
|
+
*/function Gp(n){const e={x:v(n,"x","clone","string_or_numeric")};return A.runKernel(go,e)}const on=C({clone_:Gp});/**
|
|
396
396
|
* @license
|
|
397
397
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
398
398
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -408,7 +408,7 @@
|
|
|
408
408
|
* limitations under the License.
|
|
409
409
|
* =============================================================================
|
|
410
410
|
*/function Vp(n,t=0){w(n.length>=1,()=>"Pass at least one tensor to concat");const e=hl(n,"tensors","concat","string_or_numeric");if(e[0].dtype==="complex64"&&e.forEach(o=>{if(o.dtype!=="complex64")throw new Error(`Cannot concatenate complex64 tensors with a tensor
|
|
411
|
-
with dtype ${o.dtype}. `)}),e.length===1)return
|
|
411
|
+
with dtype ${o.dtype}. `)}),e.length===1)return on(e[0]);const s=e,r={axis:t};return A.runKernel(Ca,s,r)}const an=C({concat_:Vp});function qp(n,t){return an(n,t)}const jp=C({concat4d_:qp});/**
|
|
412
412
|
* @license
|
|
413
413
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
414
414
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -513,7 +513,7 @@
|
|
|
513
513
|
* See the License for the specific language governing permissions and
|
|
514
514
|
* limitations under the License.
|
|
515
515
|
* =============================================================================
|
|
516
|
-
*/const Do="tensorflowjs",Ro=1,
|
|
516
|
+
*/const Do="tensorflowjs",Ro=1,ln="models_store",We="model_info_store";function bl(){if(!W().getBool("IS_BROWSER"))throw new Error("Failed to obtain IndexedDB factory because the current environmentis not a web browser.");const n=typeof window>"u"?self:window,t=n.indexedDB||n.mozIndexedDB||n.webkitIndexedDB||n.msIndexedDB||n.shimIndexedDB;if(t==null)throw new Error("The current browser does not appear to support IndexedDB.");return t}function Po(n){const t=n.result;t.createObjectStore(ln,{keyPath:"modelPath"}),t.createObjectStore(We,{keyPath:"modelPath"})}class un{constructor(t){if(this.indexedDB=bl(),t==null||!t)throw new Error("For IndexedDB, modelPath must not be null, undefined or empty.");this.modelPath=t}async save(t){if(t.modelTopology instanceof ArrayBuffer)throw new Error("BrowserLocalStorage.save() does not support saving model topology in binary formats yet.");return this.databaseAction(this.modelPath,t)}async load(){return this.databaseAction(this.modelPath)}databaseAction(t,e){return new Promise((s,r)=>{const o=this.indexedDB.open(Do,Ro);o.onupgradeneeded=()=>Po(o),o.onsuccess=()=>{const i=o.result;if(e==null){const a=i.transaction(ln,"readonly"),u=a.objectStore(ln).get(this.modelPath);u.onsuccess=()=>{if(u.result==null)return i.close(),r(new Error(`Cannot find model with path '${this.modelPath}' in IndexedDB.`));s(u.result.modelArtifacts)},u.onerror=c=>(i.close(),r(u.error)),a.oncomplete=()=>i.close()}else{e.weightData=Pn.join(e.weightData);const a=gl(e),l=i.transaction(We,"readwrite");let u=l.objectStore(We),c;try{c=u.put({modelPath:this.modelPath,modelArtifactsInfo:a})}catch(f){return r(f)}let h;c.onsuccess=()=>{h=i.transaction(ln,"readwrite");const f=h.objectStore(ln);let d;try{d=f.put({modelPath:this.modelPath,modelArtifacts:e,modelArtifactsInfo:a})}catch(p){return r(p)}d.onsuccess=()=>s({modelArtifactsInfo:a}),d.onerror=p=>{u=l.objectStore(We);const g=u.delete(this.modelPath);g.onsuccess=()=>(i.close(),r(d.error)),g.onerror=m=>(i.close(),r(d.error))}},c.onerror=f=>(i.close(),r(c.error)),l.oncomplete=()=>{h==null?i.close():h.oncomplete=()=>i.close()}}},o.onerror=i=>r(o.error)})}}un.URL_SCHEME="indexeddb://";const yl=n=>W().getBool("IS_BROWSER")&&!Array.isArray(n)&&n.startsWith(un.URL_SCHEME)?sm(n.slice(un.URL_SCHEME.length)):null;_t.registerSaveRouter(yl),_t.registerLoadRouter(yl);function sm(n){return new un(n)}function rm(n){return n.startsWith(un.URL_SCHEME)?n.slice(un.URL_SCHEME.length):n}class om{constructor(){this.indexedDB=bl()}async listModels(){return new Promise((t,e)=>{const s=this.indexedDB.open(Do,Ro);s.onupgradeneeded=()=>Po(s),s.onsuccess=()=>{const r=s.result,o=r.transaction(We,"readonly"),a=o.objectStore(We).getAll();a.onsuccess=()=>{const l={};for(const u of a.result)l[u.modelPath]=u.modelArtifactsInfo;t(l)},a.onerror=l=>(r.close(),e(a.error)),o.oncomplete=()=>r.close()},s.onerror=r=>e(s.error)})}async removeModel(t){return t=rm(t),new Promise((e,s)=>{const r=this.indexedDB.open(Do,Ro);r.onupgradeneeded=()=>Po(r),r.onsuccess=()=>{const o=r.result,i=o.transaction(We,"readwrite"),a=i.objectStore(We),l=a.get(t);let u;l.onsuccess=()=>{if(l.result==null)return o.close(),s(new Error(`Cannot find model with path '${t}' in IndexedDB.`));{const c=a.delete(t),h=()=>{u=o.transaction(ln,"readwrite");const d=u.objectStore(ln).delete(t);d.onsuccess=()=>e(l.result.modelArtifactsInfo),d.onerror=p=>s(l.error)};c.onsuccess=h,c.onerror=f=>(h(),o.close(),s(l.error))}},l.onerror=c=>(o.close(),s(l.error)),i.oncomplete=()=>{u==null?o.close():u.oncomplete=()=>o.close()}},r.onerror=o=>s(r.error)})}}/**
|
|
517
517
|
* @license
|
|
518
518
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
519
519
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -528,7 +528,7 @@
|
|
|
528
528
|
* See the License for the specific language governing permissions and
|
|
529
529
|
* limitations under the License.
|
|
530
530
|
* =============================================================================
|
|
531
|
-
*/const Pe="/",Mn="tensorflowjs_models",wl="info",im="model_topology",am="weight_specs",lm="weight_data",um="model_metadata";function xl(n){return{info:[Mn,n,wl].join(Pe),topology:[Mn,n,im].join(Pe),weightSpecs:[Mn,n,am].join(Pe),weightData:[Mn,n,lm].join(Pe),modelMetadata:[Mn,n,um].join(Pe)}}function Sl(n){for(const t of Object.values(n))window.localStorage.removeItem(t)}function cm(n){const t=n.split(Pe);if(t.length<3)throw new Error(`Invalid key format: ${n}`);return t.slice(1,t.length-1).join(Pe)}function hm(n){return n.startsWith(
|
|
531
|
+
*/const Pe="/",Mn="tensorflowjs_models",wl="info",im="model_topology",am="weight_specs",lm="weight_data",um="model_metadata";function xl(n){return{info:[Mn,n,wl].join(Pe),topology:[Mn,n,im].join(Pe),weightSpecs:[Mn,n,am].join(Pe),weightData:[Mn,n,lm].join(Pe),modelMetadata:[Mn,n,um].join(Pe)}}function Sl(n){for(const t of Object.values(n))window.localStorage.removeItem(t)}function cm(n){const t=n.split(Pe);if(t.length<3)throw new Error(`Invalid key format: ${n}`);return t.slice(1,t.length-1).join(Pe)}function hm(n){return n.startsWith(cn.URL_SCHEME)?n.slice(cn.URL_SCHEME.length):n}class cn{constructor(t){if(!W().getBool("IS_BROWSER")||typeof window>"u"||typeof window.localStorage>"u")throw new Error("The current environment does not support local storage.");if(this.LS=window.localStorage,t==null||!t)throw new Error("For local storage, modelPath must not be null, undefined or empty.");this.modelPath=t,this.keys=xl(this.modelPath)}async save(t){if(t.modelTopology instanceof ArrayBuffer)throw new Error("BrowserLocalStorage.save() does not support saving model topology in binary formats yet.");{const e=JSON.stringify(t.modelTopology),s=JSON.stringify(t.weightSpecs),r=gl(t),o=Pn.join(t.weightData);try{this.LS.setItem(this.keys.info,JSON.stringify(r)),this.LS.setItem(this.keys.topology,e),this.LS.setItem(this.keys.weightSpecs,s),this.LS.setItem(this.keys.weightData,Qp(o));const i={format:t.format,generatedBy:t.generatedBy,convertedBy:t.convertedBy,signature:t.signature!=null?t.signature:void 0,userDefinedMetadata:t.userDefinedMetadata!=null?t.userDefinedMetadata:void 0,modelInitializer:t.modelInitializer!=null?t.modelInitializer:void 0,initializerSignature:t.initializerSignature!=null?t.initializerSignature:void 0,trainingConfig:t.trainingConfig!=null?t.trainingConfig:void 0};return this.LS.setItem(this.keys.modelMetadata,JSON.stringify(i)),{modelArtifactsInfo:r}}catch{throw Sl(this.keys),new Error(`Failed to save model '${this.modelPath}' to local storage: size quota being exceeded is a possible cause of this failure: modelTopologyBytes=${r.modelTopologyBytes}, weightSpecsBytes=${r.weightSpecsBytes}, weightDataBytes=${r.weightDataBytes}.`)}}}async load(){const t=JSON.parse(this.LS.getItem(this.keys.info));if(t==null)throw new Error(`In local storage, there is no model with name '${this.modelPath}'`);if(t.modelTopologyType!=="JSON")throw new Error("BrowserLocalStorage does not support loading non-JSON model topology yet.");const e={},s=JSON.parse(this.LS.getItem(this.keys.topology));if(s==null)throw new Error(`In local storage, the topology of model '${this.modelPath}' is missing.`);e.modelTopology=s;const r=JSON.parse(this.LS.getItem(this.keys.weightSpecs));if(r==null)throw new Error(`In local storage, the weight specs of model '${this.modelPath}' are missing.`);e.weightSpecs=r;const o=this.LS.getItem(this.keys.modelMetadata);if(o!=null){const a=JSON.parse(o);e.format=a.format,e.generatedBy=a.generatedBy,e.convertedBy=a.convertedBy,a.signature!=null&&(e.signature=a.signature),a.userDefinedMetadata!=null&&(e.userDefinedMetadata=a.userDefinedMetadata),a.modelInitializer!=null&&(e.modelInitializer=a.modelInitializer),a.initializerSignature!=null&&(e.initializerSignature=a.initializerSignature),a.trainingConfig!=null&&(e.trainingConfig=a.trainingConfig)}const i=this.LS.getItem(this.keys.weightData);if(i==null)throw new Error(`In local storage, the binary weight values of model '${this.modelPath}' are missing.`);return e.weightData=tm(i),e}}cn.URL_SCHEME="localstorage://";const $l=n=>W().getBool("IS_BROWSER")&&!Array.isArray(n)&&n.startsWith(cn.URL_SCHEME)?fm(n.slice(cn.URL_SCHEME.length)):null;_t.registerSaveRouter($l),_t.registerLoadRouter($l);function fm(n){return new cn(n)}class dm{constructor(){w(W().getBool("IS_BROWSER"),()=>"Current environment is not a web browser"),w(typeof window>"u"||typeof window.localStorage<"u",()=>"Current browser does not appear to support localStorage"),this.LS=window.localStorage}async listModels(){const t={},e=Mn+Pe,s=Pe+wl;for(let r=0;r<this.LS.length;++r){const o=this.LS.key(r);if(o.startsWith(e)&&o.endsWith(s)){const i=cm(o);t[i]=JSON.parse(this.LS.getItem(o))}}return t}async removeModel(t){t=hm(t);const e=xl(t);if(this.LS.getItem(e.info)==null)throw new Error(`Cannot find model at path '${t}'`);const s=JSON.parse(this.LS.getItem(e.info));return Sl(e),s}}/**
|
|
532
532
|
* @license
|
|
533
533
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
534
534
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -558,7 +558,7 @@
|
|
|
558
558
|
* See the License for the specific language governing permissions and
|
|
559
559
|
* limitations under the License.
|
|
560
560
|
* =============================================================================
|
|
561
|
-
*/class pm{constructor(){this.messageName="setTimeoutCustom",this.functionRefs=[],this.handledMessageCount=0,this.hasEventListener=!1}fetch(t,e){return fetch(t,e)}now(){return performance.now()}encode(t,e){if(e!=="utf-8"&&e!=="utf8")throw new Error(`Browser's encoder only supports utf-8, but got ${e}`);return this.textEncoder==null&&(this.textEncoder=new TextEncoder),this.textEncoder.encode(t)}decode(t,e){return new TextDecoder(e).decode(t)}setTimeoutCustom(t,e){if(typeof window>"u"||!W().getBool("USE_SETTIMEOUTCUSTOM")){setTimeout(t,e);return}this.functionRefs.push(t),setTimeout(()=>{window.postMessage({name:this.messageName,index:this.functionRefs.length-1},"*")},e),this.hasEventListener||(this.hasEventListener=!0,window.addEventListener("message",s=>{if(s.source===window&&s.data.name===this.messageName){s.stopPropagation();const r=this.functionRefs[s.data.index];r(),this.handledMessageCount++,this.handledMessageCount===this.functionRefs.length&&(this.functionRefs=[],this.handledMessageCount=0)}},!0))}isTypedArray(t){return Ba(t)}}if(W().get("IS_BROWSER")){W().setPlatform("browser",new pm);try{we.registerManager(
|
|
561
|
+
*/class pm{constructor(){this.messageName="setTimeoutCustom",this.functionRefs=[],this.handledMessageCount=0,this.hasEventListener=!1}fetch(t,e){return fetch(t,e)}now(){return performance.now()}encode(t,e){if(e!=="utf-8"&&e!=="utf8")throw new Error(`Browser's encoder only supports utf-8, but got ${e}`);return this.textEncoder==null&&(this.textEncoder=new TextEncoder),this.textEncoder.encode(t)}decode(t,e){return new TextDecoder(e).decode(t)}setTimeoutCustom(t,e){if(typeof window>"u"||!W().getBool("USE_SETTIMEOUTCUSTOM")){setTimeout(t,e);return}this.functionRefs.push(t),setTimeout(()=>{window.postMessage({name:this.messageName,index:this.functionRefs.length-1},"*")},e),this.hasEventListener||(this.hasEventListener=!0,window.addEventListener("message",s=>{if(s.source===window&&s.data.name===this.messageName){s.stopPropagation();const r=this.functionRefs[s.data.index];r(),this.handledMessageCount++,this.handledMessageCount===this.functionRefs.length&&(this.functionRefs=[],this.handledMessageCount=0)}},!0))}isTypedArray(t){return Ba(t)}}if(W().get("IS_BROWSER")){W().setPlatform("browser",new pm);try{we.registerManager(cn.URL_SCHEME,new dm)}catch{}try{we.registerManager(un.URL_SCHEME,new om)}catch{}}/**
|
|
562
562
|
* @license
|
|
563
563
|
* Copyright 2019 Google LLC. All Rights Reserved.
|
|
564
564
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -588,7 +588,7 @@
|
|
|
588
588
|
* See the License for the specific language governing permissions and
|
|
589
589
|
* limitations under the License.
|
|
590
590
|
* =============================================================================
|
|
591
|
-
*/function xt(n,t="float32",e){return t=t||"float32",Re(n),new
|
|
591
|
+
*/function xt(n,t="float32",e){return t=t||"float32",Re(n),new Js(n,t,e)}/**
|
|
592
592
|
* @license
|
|
593
593
|
* Copyright 2020 Google Inc. All Rights Reserved.
|
|
594
594
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -633,7 +633,7 @@
|
|
|
633
633
|
* See the License for the specific language governing permissions and
|
|
634
634
|
* limitations under the License.
|
|
635
635
|
* =============================================================================
|
|
636
|
-
*/ll(),kp({buffer:xt,cast:ot,clone:
|
|
636
|
+
*/ll(),kp({buffer:xt,cast:ot,clone:on,print:ym});/**
|
|
637
637
|
* @license
|
|
638
638
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
639
639
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -738,7 +738,7 @@
|
|
|
738
738
|
* See the License for the specific language governing permissions and
|
|
739
739
|
* limitations under the License.
|
|
740
740
|
* =============================================================================
|
|
741
|
-
*/function Cm(n,t=0){const s={x:v(n,"x","argMax")},r={axis:t};return A.runKernel(Uf,s,r)}const
|
|
741
|
+
*/function Cm(n,t=0){const s={x:v(n,"x","argMax")},r={axis:t};return A.runKernel(Uf,s,r)}const nr=C({argMax_:Cm});/**
|
|
742
742
|
* @license
|
|
743
743
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
744
744
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -753,7 +753,7 @@
|
|
|
753
753
|
* See the License for the specific language governing permissions and
|
|
754
754
|
* limitations under the License.
|
|
755
755
|
* =============================================================================
|
|
756
|
-
*/function km(n,t,e,s,r,o,i="channelsLast"){const[a,l]=
|
|
756
|
+
*/function km(n,t,e,s,r,o,i="channelsLast"){const[a,l]=fs(t);let u;if(i==="channelsLast")u=[a,l,n[3],n[3]];else if(i==="channelsFirst")u=[a,l,n[1],n[1]];else throw new Error(`Unknown dataFormat ${i}`);return Mo(n,u,e,s,r,o,!1,i)}function Mo(n,t,e,s,r,o,i=!1,a="channelsLast"){let[l,u,c,h]=[-1,-1,-1,-1];if(a==="channelsLast")[l,u,c,h]=n;else if(a==="channelsFirst")[l,h,u,c]=n;else throw new Error(`Unknown dataFormat ${a}`);const[f,d,,p]=t,[g,m]=fs(e),[b,y]=fs(s),S=Oo(f,b),x=Oo(d,y),{padInfo:$,outHeight:E,outWidth:D}=Nm(r,u,c,g,m,S,x,o,a),_=i?p*h:p;let T;return a==="channelsFirst"?T=[l,_,E,D]:a==="channelsLast"&&(T=[l,E,D,_]),{batchSize:l,dataFormat:a,inHeight:u,inWidth:c,inChannels:h,outHeight:E,outWidth:D,outChannels:_,padInfo:$,strideHeight:g,strideWidth:m,filterHeight:f,filterWidth:d,effectiveFilterHeight:S,effectiveFilterWidth:x,dilationHeight:b,dilationWidth:y,inShape:n,outShape:T,filterShape:t}}function _m(n,t,e,s,r){s==null&&(s=Tm(n,t,e));const o=n[0],i=n[1],a=sr((o-t+2*s)/e+1,r),l=sr((i-t+2*s)/e+1,r);return[a,l]}function Tm(n,t,e,s=1){const r=Oo(t,s);return Math.floor((n[0]*(e-1)-e+r)/2)}function fs(n){return typeof n=="number"?[n,n,n]:n.length===2?[n[0],n[1],1]:n}function Oo(n,t){return t<=1?n:n+(n-1)*(t-1)}function Nm(n,t,e,s,r,o,i,a,l){let u,c,h;if(typeof n=="number"){u={top:n,bottom:n,left:n,right:n,type:n===0?"VALID":"NUMBER"};const d=_m([t,e],o,s,n,a);c=d[0],h=d[1]}else if(n==="same"){c=Math.ceil(t/s),h=Math.ceil(e/r);const f=Math.max(0,(c-1)*s+o-t),d=Math.max(0,(h-1)*r+i-e),p=Math.floor(f/2),g=f-p,m=Math.floor(d/2),b=d-m;u={top:p,bottom:g,left:m,right:b,type:"SAME"}}else if(n==="valid")u={top:0,bottom:0,left:0,right:0,type:"VALID"},c=Math.ceil((t-o+1)/s),h=Math.ceil((e-i+1)/r);else if(typeof n=="object"){const f=l==="channelsLast"?n[1][0]:n[2][0],d=l==="channelsLast"?n[1][1]:n[2][1],p=l==="channelsLast"?n[2][0]:n[3][0],g=l==="channelsLast"?n[2][1]:n[3][1];u={top:f,bottom:d,left:p,right:g,type:f===0&&d===0&&p===0&&g===0?"VALID":"EXPLICIT"},c=sr((t-o+f+d)/s+1,a),h=sr((e-i+p+g)/r+1,a)}else throw Error(`Unknown padding parameter: ${n}`);return{padInfo:u,outHeight:c,outWidth:h}}function sr(n,t){if(!t)return Math.trunc(n);switch(t){case"round":return Math.round(n);case"ceil":return Math.ceil(n);case"floor":return Math.floor(n);default:throw new Error(`Unknown roundingMode ${t}`)}}function Bo(n){const[t,e,s]=fs(n);return t===1&&e===1&&s===1}function On(n,t){return Bo(n)||Bo(t)}function Bn(n){return fs(n).every(t=>t>0)}function Dm(n){if(n==="NHWC")return"channelsLast";if(n==="NCHW")return"channelsFirst";throw new Error(`Unknown dataFormat ${n}`)}function xe(n,t,e){if(e!=null){if(typeof t=="string")throw Error(`Error in ${n}: pad must be an integer when using dimRoundingMode ${e} but got pad ${t}.`);if(typeof t=="number")w(oo(t),()=>`Error in ${n}: pad must be an integer when using dimRoundingMode ${e} but got pad ${t}.`);else if(typeof t=="object")t.forEach(s=>{s.forEach(r=>{w(oo(r),()=>`Error in ${n}: pad must be an integer when using dimRoundingMode ${e} but got pad ${r}.`)})});else throw Error(`Error in ${n}: Unknown padding parameter: ${t}`)}}/**
|
|
757
757
|
* @license
|
|
758
758
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
759
759
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -873,7 +873,7 @@
|
|
|
873
873
|
* See the License for the specific language governing permissions and
|
|
874
874
|
* limitations under the License.
|
|
875
875
|
* =============================================================================
|
|
876
|
-
*/function Gm(n,t){let e=v(n,"broadcastTo","x");const s=e.shape;if(Re(t),t.length<e.rank)throw new Error(`broadcastTo(): shape.length=${t.length} < input.rank=${e.rank}.`);if(t.length>e.rank){const u=e.shape.slice();for(;u.length<t.length;)u.unshift(1);e=L(e,u)}const r=e.shape,o=Array.from(t);for(let u=t.length-1;u>=0;u--)if(r[u]===t[u])o[u]=1;else if(e.shape[u]!==1)throw new Error(`broadcastTo(): [${s}] cannot be broadcast to [${t}].`);if(o.map((u,c)=>u>1?c:-1).filter(u=>u>=0).length===0)return
|
|
876
|
+
*/function Gm(n,t){let e=v(n,"broadcastTo","x");const s=e.shape;if(Re(t),t.length<e.rank)throw new Error(`broadcastTo(): shape.length=${t.length} < input.rank=${e.rank}.`);if(t.length>e.rank){const u=e.shape.slice();for(;u.length<t.length;)u.unshift(1);e=L(e,u)}const r=e.shape,o=Array.from(t);for(let u=t.length-1;u>=0;u--)if(r[u]===t[u])o[u]=1;else if(e.shape[u]!==1)throw new Error(`broadcastTo(): [${s}] cannot be broadcast to [${t}].`);if(o.map((u,c)=>u>1?c:-1).filter(u=>u>=0).length===0)return on(e);const a={x:e},l={reps:o};return A.runKernel(Ra,a,l)}const rr=C({broadcastTo_:Gm});/**
|
|
877
877
|
* @license
|
|
878
878
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
879
879
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -888,7 +888,7 @@
|
|
|
888
888
|
* See the License for the specific language governing permissions and
|
|
889
889
|
* limitations under the License.
|
|
890
890
|
* =============================================================================
|
|
891
|
-
*/function
|
|
891
|
+
*/function or(n,t,e){Re(n),e=e||as(t);const s={shape:n,value:t,dtype:e};return A.runKernel(ud,{},s)}/**
|
|
892
892
|
* @license
|
|
893
893
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
894
894
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -903,7 +903,7 @@
|
|
|
903
903
|
* See the License for the specific language governing permissions and
|
|
904
904
|
* limitations under the License.
|
|
905
905
|
* =============================================================================
|
|
906
|
-
*/function Vm(n,t,e){const s=v(n,"x","clipByValue");if(w(t<=e,()=>`Error in clip: min (${t}) must be less than or equal to max (${e}).`),t===e)return
|
|
906
|
+
*/function Vm(n,t,e){const s=v(n,"x","clipByValue");if(w(t<=e,()=>`Error in clip: min (${t}) must be less than or equal to max (${e}).`),t===e)return or(s.shape,t,s.dtype);const r={x:s},o={clipValueMin:t,clipValueMax:e};return A.runKernel(jf,r,o)}const he=C({clipByValue_:Vm});/**
|
|
907
907
|
* @license
|
|
908
908
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
909
909
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -993,7 +993,7 @@
|
|
|
993
993
|
* See the License for the specific language governing permissions and
|
|
994
994
|
* limitations under the License.
|
|
995
995
|
* =============================================================================
|
|
996
|
-
*/function
|
|
996
|
+
*/function ir(n,t){const e=n.length,s=[];for(let r=0;r<e;r++){const o=e-1-r,i=n[o]||1;(t[t.length-1-r]||1)>1&&i===1&&s.unshift(o)}return s}function og(n,t){const e=[];for(let s=0;s<t.length;s++){const r=n[n.length-s-1],o=t.length-s-1,i=t[o];(r==null||r===1&&i>1)&&e.unshift(o)}return e}function zt(n,t){const e=Math.max(n.length,t.length),s=new Array(e);for(let r=0;r<e;r++){let o=n[n.length-r-1];o==null&&(o=1);let i=t[t.length-r-1];if(i==null&&(i=1),o===1)s[e-r-1]=i;else if(i===1)s[e-r-1]=o;else if(o!==i){const a=`Operands could not be broadcast together with shapes ${n} and ${t}.`;throw Error(a)}else s[e-r-1]=o}return s}/**
|
|
997
997
|
* @license
|
|
998
998
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
999
999
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1008,7 +1008,7 @@
|
|
|
1008
1008
|
* See the License for the specific language governing permissions and
|
|
1009
1009
|
* limitations under the License.
|
|
1010
1010
|
* =============================================================================
|
|
1011
|
-
*/function ig(n,t){let e=v(n,"a","equal","string_or_numeric"),s=v(t,"b","equal","string_or_numeric");[e,s]=Dt(e,s),zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel(id,r)}const
|
|
1011
|
+
*/function ig(n,t){let e=v(n,"a","equal","string_or_numeric"),s=v(t,"b","equal","string_or_numeric");[e,s]=Dt(e,s),zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel(id,r)}const hn=C({equal_:ig});/**
|
|
1012
1012
|
* @license
|
|
1013
1013
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1014
1014
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1023,7 +1023,7 @@
|
|
|
1023
1023
|
* See the License for the specific language governing permissions and
|
|
1024
1024
|
* limitations under the License.
|
|
1025
1025
|
* =============================================================================
|
|
1026
|
-
*/function ag(n,t,e){const s=v(t,"a","where"),r=v(e,"b","where"),o=v(n,"condition","where","bool"),i=zt(zt(o.shape,s.shape),r.shape),a=
|
|
1026
|
+
*/function ag(n,t,e){const s=v(t,"a","where"),r=v(e,"b","where"),o=v(n,"condition","where","bool"),i=zt(zt(o.shape,s.shape),r.shape),a=rr(o,i),l=rr(s,i),u=rr(r,i),c={condition:a,t:l,e:u};return A.runKernel(jd,c)}const fn=C({where_:ag});/**
|
|
1027
1027
|
* @license
|
|
1028
1028
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1029
1029
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1053,7 +1053,7 @@
|
|
|
1053
1053
|
* See the License for the specific language governing permissions and
|
|
1054
1054
|
* limitations under the License.
|
|
1055
1055
|
* =============================================================================
|
|
1056
|
-
*/function ug(n,...t){const e=t.map((r,o)=>v(r,`tensors${o}`,"einsum")),s={equation:n};return A.runKernel(sd,e,s)}const
|
|
1056
|
+
*/function ug(n,...t){const e=t.map((r,o)=>v(r,`tensors${o}`,"einsum")),s={equation:n};return A.runKernel(sd,e,s)}const ds=C({einsum_:ug});/**
|
|
1057
1057
|
* @license
|
|
1058
1058
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1059
1059
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1143,7 +1143,7 @@
|
|
|
1143
1143
|
* See the License for the specific language governing permissions and
|
|
1144
1144
|
* limitations under the License.
|
|
1145
1145
|
* =============================================================================
|
|
1146
|
-
*/function wg(n,t){let e=v(n,"base","pow"),s=v(t,"exp","pow");[e,s]=Dt(e,s);const r={a:e,b:s};return A.runKernel(Od,r)}const
|
|
1146
|
+
*/function wg(n,t){let e=v(n,"base","pow"),s=v(t,"exp","pow");[e,s]=Dt(e,s);const r={a:e,b:s};return A.runKernel(Od,r)}const ar=C({pow_:wg});/**
|
|
1147
1147
|
* @license
|
|
1148
1148
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1149
1149
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1158,7 +1158,7 @@
|
|
|
1158
1158
|
* See the License for the specific language governing permissions and
|
|
1159
1159
|
* limitations under the License.
|
|
1160
1160
|
* =============================================================================
|
|
1161
|
-
*/function
|
|
1161
|
+
*/function Yt(n,t){if((ne(n)&&t!=="string"||Array.isArray(n))&&t!=="complex64")throw new Error("Error creating a new Scalar: value must be a primitive (number|boolean|string)");if(t==="string"&&ne(n)&&!(n instanceof Uint8Array))throw new Error("When making a scalar from encoded string, the value must be `Uint8Array`.");return tr(n,[],[],t)}/**
|
|
1162
1162
|
* @license
|
|
1163
1163
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1164
1164
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1218,7 +1218,7 @@
|
|
|
1218
1218
|
* See the License for the specific language governing permissions and
|
|
1219
1219
|
* limitations under the License.
|
|
1220
1220
|
* =============================================================================
|
|
1221
|
-
*/function vg(n,t="euclidean",e=null,s=!1){n=v(n,"x","norm");const r=_l(n,t,e);let o=r.shape;if(s){const i=
|
|
1221
|
+
*/function vg(n,t="euclidean",e=null,s=!1){n=v(n,"x","norm");const r=_l(n,t,e);let o=r.shape;if(s){const i=is(e,n.shape);o=Cl(r.shape,i)}return L(r,o)}function _l(n,t,e=null){if(n.rank===0)return Lt(n);if(n.rank!==1&&e===null)return _l(L(n,[-1]),t,e);if(n.rank===1||typeof e=="number"||Array.isArray(e)&&e.length===1){if(t===1)return et(Lt(n),e);if(t===1/0)return Ge(Lt(n),e);if(t===-1/0)return kl(Lt(n),e);if(t==="euclidean"||t===2)return fe(et(ar(Lt(n),Yt(2,"int32")),e));throw new Error(`Error in norm: invalid ord value: ${t}`)}if(Array.isArray(e)&&e.length===2){if(t===1)return Ge(et(Lt(n),e[0]),e[1]-1);if(t===1/0)return Ge(et(Lt(n),e[1]),e[0]);if(t===-1/0)return kl(et(Lt(n),e[1]),e[0]);if(t==="fro"||t==="euclidean")return fe(et(Ve(n),e));throw new Error(`Error in norm: invalid ord value: ${t}`)}throw new Error(`Error in norm: invalid axis: ${e}`)}const Tl=C({norm_:vg});/**
|
|
1222
1222
|
* @license
|
|
1223
1223
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1224
1224
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1263,7 +1263,7 @@
|
|
|
1263
1263
|
* See the License for the specific language governing permissions and
|
|
1264
1264
|
* limitations under the License.
|
|
1265
1265
|
* =============================================================================
|
|
1266
|
-
*/function Eg(n,t){const e=v(n,"x","tile","string_or_numeric");w(e.rank===t.length,()=>`Error in transpose: rank of input ${e.rank} must match length of reps ${t}.`);const s={x:e},r={reps:t};return A.runKernel(Ra,s,r)}const
|
|
1266
|
+
*/function Eg(n,t){const e=v(n,"x","tile","string_or_numeric");w(e.rank===t.length,()=>`Error in transpose: rank of input ${e.rank} must match length of reps ${t}.`);const s={x:e},r={reps:t};return A.runKernel(Ra,s,r)}const lr=C({tile_:Eg});/**
|
|
1267
1267
|
* @license
|
|
1268
1268
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1269
1269
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1278,7 +1278,7 @@
|
|
|
1278
1278
|
* See the License for the specific language governing permissions and
|
|
1279
1279
|
* limitations under the License.
|
|
1280
1280
|
* =============================================================================
|
|
1281
|
-
*/function Cg(n,t,e,s="float32"){t==null&&(t=n);const r=xt([n,t],s),o=n<=t?n:t;for(let a=0;a<o;++a)r.set(1,a,a);const i=L(r.toTensor(),[n,t]);if(e==null)return i;if(e.length===1)return
|
|
1281
|
+
*/function Cg(n,t,e,s="float32"){t==null&&(t=n);const r=xt([n,t],s),o=n<=t?n:t;for(let a=0;a<o;++a)r.set(1,a,a);const i=L(r.toTensor(),[n,t]);if(e==null)return i;if(e.length===1)return lr(ve(i,0),[e[0],1,1]);if(e.length===2)return lr(ve(ve(i,0),0),[e[0],e[1],1,1]);if(e.length===3)return lr(ve(ve(ve(i,0),0),0),[e[0],e[1],e[2],1,1]);throw new Error(`eye() currently supports only 1D and 2D batchShapes, but received ${e.length}D.`)}const Nl=C({eye_:Cg});/**
|
|
1282
1282
|
* @license
|
|
1283
1283
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1284
1284
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1323,7 +1323,7 @@
|
|
|
1323
1323
|
* See the License for the specific language governing permissions and
|
|
1324
1324
|
* limitations under the License.
|
|
1325
1325
|
* =============================================================================
|
|
1326
|
-
*/function Dg(n,t){let e=v(n,"a","greater","string_or_numeric"),s=v(t,"b","greater","string_or_numeric");[e,s]=Dt(e,s),zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel(pd,r)}const
|
|
1326
|
+
*/function Dg(n,t){let e=v(n,"a","greater","string_or_numeric"),s=v(t,"b","greater","string_or_numeric");[e,s]=Dt(e,s),zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel(pd,r)}const ps=C({greater_:Dg});/**
|
|
1327
1327
|
* @license
|
|
1328
1328
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1329
1329
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1413,7 +1413,7 @@
|
|
|
1413
1413
|
* See the License for the specific language governing permissions and
|
|
1414
1414
|
* limitations under the License.
|
|
1415
1415
|
* =============================================================================
|
|
1416
|
-
*/function Ug(n){const e={x:v(n,"x","log","float32")};return A.runKernel(xd,e)}const
|
|
1416
|
+
*/function Ug(n){const e={x:v(n,"x","log","float32")};return A.runKernel(xd,e)}const dn=C({log_:Ug});/**
|
|
1417
1417
|
* @license
|
|
1418
1418
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1419
1419
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1443,7 +1443,7 @@
|
|
|
1443
1443
|
* See the License for the specific language governing permissions and
|
|
1444
1444
|
* limitations under the License.
|
|
1445
1445
|
* =============================================================================
|
|
1446
|
-
*/function Vg(n,t){w(lo(n),()=>"The f passed in variableGrads(f) must be a function"),w(t==null||Array.isArray(t)&&t.every(u=>u instanceof
|
|
1446
|
+
*/function Vg(n,t){w(lo(n),()=>"The f passed in variableGrads(f) must be a function"),w(t==null||Array.isArray(t)&&t.every(u=>u instanceof Zs),()=>"The varList passed in variableGrads(f, varList) must be an array of variables");const e=t!=null;if(!e){t=[];for(const u in A.registeredVariables)t.push(A.registeredVariables[u])}const s=e?t.filter(u=>!u.trainable):null,r=t.length;t=t.filter(u=>u.trainable),w(t.length>0,()=>`variableGrads() expects at least one of the input variables to be trainable, but none of the ${r} variables is trainable.`);const o=!0,{value:i,grads:a}=A.gradients(n,t,null,o);w(a.some(u=>u!=null),()=>"Cannot find a connection between any variable and the result of the loss function y=f(x). Please make sure the operations that use variables are inside the function f passed to minimize()."),w(i.rank===0,()=>`The f passed in variableGrads(f) must return a scalar, but it returned a rank-${i.rank} tensor`);const l={};return t.forEach((u,c)=>{a[c]!=null&&(l[u.name]=a[c])}),s!=null&&s.forEach(u=>l[u.name]=null),{value:i,grads:l}}function Vo(n){return A.customGrad(n)}/**
|
|
1447
1447
|
* @license
|
|
1448
1448
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1449
1449
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1503,7 +1503,7 @@
|
|
|
1503
1503
|
* See the License for the specific language governing permissions and
|
|
1504
1504
|
* limitations under the License.
|
|
1505
1505
|
* =============================================================================
|
|
1506
|
-
*/function Kg(n,t=-1){const e=v(n,"logits","logSoftmax");if(t===-1&&(t=e.rank-1),t!==e.rank-1)throw Error(`Log Softmax along a non-last dimension is not yet supported. Logits was rank ${e.rank} and axis was ${t}`);return Vo((r,o)=>{const a=Ge(r,t,!0),l=J(r,a),u=J(ot(l,"float32"),
|
|
1506
|
+
*/function Kg(n,t=-1){const e=v(n,"logits","logSoftmax");if(t===-1&&(t=e.rank-1),t!==e.rank-1)throw Error(`Log Softmax along a non-last dimension is not yet supported. Logits was rank ${e.rank} and axis was ${t}`);return Vo((r,o)=>{const a=Ge(r,t,!0),l=J(r,a),u=J(ot(l,"float32"),dn(et(Go(l),t,!0)));return o([u]),{value:u,gradFunc:(h,f)=>{const[d]=f,p=!0,g=Go(d);return J(h,N(et(h,t,p),g))}}})(e)}const Yg=C({logSoftmax_:Kg});/**
|
|
1507
1507
|
* @license
|
|
1508
1508
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1509
1509
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1518,7 +1518,7 @@
|
|
|
1518
1518
|
* See the License for the specific language governing permissions and
|
|
1519
1519
|
* limitations under the License.
|
|
1520
1520
|
* =============================================================================
|
|
1521
|
-
*/function Xg(n,t){const e=v(n,"a","logicalAnd","bool"),s=v(t,"b","logicalAnd","bool");zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel($d,r)}const
|
|
1521
|
+
*/function Xg(n,t){const e=v(n,"a","logicalAnd","bool"),s=v(t,"b","logicalAnd","bool");zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel($d,r)}const ur=C({logicalAnd_:Xg});/**
|
|
1522
1522
|
* @license
|
|
1523
1523
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1524
1524
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1623,7 +1623,7 @@
|
|
|
1623
1623
|
* See the License for the specific language governing permissions and
|
|
1624
1624
|
* limitations under the License.
|
|
1625
1625
|
* =============================================================================
|
|
1626
|
-
*/function s0(n,t){let e=v(n,"a","minimum"),s=v(t,"b","minimum");[e,s]=Dt(e,s),e.dtype==="bool"&&(e=ot(e,"int32"),s=ot(s,"int32")),zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel(kd,r)}const
|
|
1626
|
+
*/function s0(n,t){let e=v(n,"a","minimum"),s=v(t,"b","minimum");[e,s]=Dt(e,s),e.dtype==="bool"&&(e=ot(e,"int32"),s=ot(s,"int32")),zt(e.shape,s.shape);const r={a:e,b:s};return A.runKernel(kd,r)}const cr=C({minimum_:s0});/**
|
|
1627
1627
|
* @license
|
|
1628
1628
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1629
1629
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1668,7 +1668,7 @@
|
|
|
1668
1668
|
* See the License for the specific language governing permissions and
|
|
1669
1669
|
* limitations under the License.
|
|
1670
1670
|
* =============================================================================
|
|
1671
|
-
*/function a0(n,t){const e=v(n,"x","prelu"),s=v(t,"alpha","prelu"),r={x:e,alpha:s};return A.runKernel(Bd,r)}const l0=C({prelu_:a0});var Ho={exports:{}};Ho.exports,function(n){(function(t,e,s){function r(l){var u=this,c=a();u.next=function(){var h=2091639*u.s0+u.c*23283064365386963e-26;return u.s0=u.s1,u.s1=u.s2,u.s2=h-(u.c=h|0)},u.c=1,u.s0=c(" "),u.s1=c(" "),u.s2=c(" "),u.s0-=c(l),u.s0<0&&(u.s0+=1),u.s1-=c(l),u.s1<0&&(u.s1+=1),u.s2-=c(l),u.s2<0&&(u.s2+=1),c=null}function o(l,u){return u.c=l.c,u.s0=l.s0,u.s1=l.s1,u.s2=l.s2,u}function i(l,u){var c=new r(l),h=u&&u.state,f=c.next;return f.int32=function(){return c.next()*4294967296|0},f.double=function(){return f()+(f()*2097152|0)*11102230246251565e-32},f.quick=f,h&&(typeof h=="object"&&o(h,c),f.state=function(){return o(c,{})}),f}function a(){var l=4022871197,u=function(c){c=String(c);for(var h=0;h<c.length;h++){l+=c.charCodeAt(h);var f=.02519603282416938*l;l=f>>>0,f-=l,f*=l,l=f>>>0,f-=l,l+=f*4294967296}return(l>>>0)*23283064365386963e-26};return u}e&&e.exports?e.exports=i:this.alea=i})(
|
|
1671
|
+
*/function a0(n,t){const e=v(n,"x","prelu"),s=v(t,"alpha","prelu"),r={x:e,alpha:s};return A.runKernel(Bd,r)}const l0=C({prelu_:a0});var Ho={exports:{}};Ho.exports,function(n){(function(t,e,s){function r(l){var u=this,c=a();u.next=function(){var h=2091639*u.s0+u.c*23283064365386963e-26;return u.s0=u.s1,u.s1=u.s2,u.s2=h-(u.c=h|0)},u.c=1,u.s0=c(" "),u.s1=c(" "),u.s2=c(" "),u.s0-=c(l),u.s0<0&&(u.s0+=1),u.s1-=c(l),u.s1<0&&(u.s1+=1),u.s2-=c(l),u.s2<0&&(u.s2+=1),c=null}function o(l,u){return u.c=l.c,u.s0=l.s0,u.s1=l.s1,u.s2=l.s2,u}function i(l,u){var c=new r(l),h=u&&u.state,f=c.next;return f.int32=function(){return c.next()*4294967296|0},f.double=function(){return f()+(f()*2097152|0)*11102230246251565e-32},f.quick=f,h&&(typeof h=="object"&&o(h,c),f.state=function(){return o(c,{})}),f}function a(){var l=4022871197,u=function(c){c=String(c);for(var h=0;h<c.length;h++){l+=c.charCodeAt(h);var f=.02519603282416938*l;l=f>>>0,f-=l,f*=l,l=f>>>0,f-=l,l+=f*4294967296}return(l>>>0)*23283064365386963e-26};return u}e&&e.exports?e.exports=i:this.alea=i})(Ze,n)}(Ho);var u0=Ho.exports,Ko={exports:{}};Ko.exports,function(n){(function(t,e,s){function r(a){var l=this,u="";l.x=0,l.y=0,l.z=0,l.w=0,l.next=function(){var h=l.x^l.x<<11;return l.x=l.y,l.y=l.z,l.z=l.w,l.w^=l.w>>>19^h^h>>>8},a===(a|0)?l.x=a:u+=a;for(var c=0;c<u.length+64;c++)l.x^=u.charCodeAt(c)|0,l.next()}function o(a,l){return l.x=a.x,l.y=a.y,l.z=a.z,l.w=a.w,l}function i(a,l){var u=new r(a),c=l&&l.state,h=function(){return(u.next()>>>0)/4294967296};return h.double=function(){do var f=u.next()>>>11,d=(u.next()>>>0)/4294967296,p=(f+d)/(1<<21);while(p===0);return p},h.int32=u.next,h.quick=h,c&&(typeof c=="object"&&o(c,u),h.state=function(){return o(u,{})}),h}e&&e.exports?e.exports=i:this.xor128=i})(Ze,n)}(Ko);var c0=Ko.exports,Yo={exports:{}};Yo.exports,function(n){(function(t,e,s){function r(a){var l=this,u="";l.next=function(){var h=l.x^l.x>>>2;return l.x=l.y,l.y=l.z,l.z=l.w,l.w=l.v,(l.d=l.d+362437|0)+(l.v=l.v^l.v<<4^(h^h<<1))|0},l.x=0,l.y=0,l.z=0,l.w=0,l.v=0,a===(a|0)?l.x=a:u+=a;for(var c=0;c<u.length+64;c++)l.x^=u.charCodeAt(c)|0,c==u.length&&(l.d=l.x<<10^l.x>>>4),l.next()}function o(a,l){return l.x=a.x,l.y=a.y,l.z=a.z,l.w=a.w,l.v=a.v,l.d=a.d,l}function i(a,l){var u=new r(a),c=l&&l.state,h=function(){return(u.next()>>>0)/4294967296};return h.double=function(){do var f=u.next()>>>11,d=(u.next()>>>0)/4294967296,p=(f+d)/(1<<21);while(p===0);return p},h.int32=u.next,h.quick=h,c&&(typeof c=="object"&&o(c,u),h.state=function(){return o(u,{})}),h}e&&e.exports?e.exports=i:this.xorwow=i})(Ze,n)}(Yo);var h0=Yo.exports,Xo={exports:{}};Xo.exports,function(n){(function(t,e,s){function r(a){var l=this;l.next=function(){var c=l.x,h=l.i,f,d;return f=c[h],f^=f>>>7,d=f^f<<24,f=c[h+1&7],d^=f^f>>>10,f=c[h+3&7],d^=f^f>>>3,f=c[h+4&7],d^=f^f<<7,f=c[h+7&7],f=f^f<<13,d^=f^f<<9,c[h]=d,l.i=h+1&7,d};function u(c,h){var f,d=[];if(h===(h|0))d[0]=h;else for(h=""+h,f=0;f<h.length;++f)d[f&7]=d[f&7]<<15^h.charCodeAt(f)+d[f+1&7]<<13;for(;d.length<8;)d.push(0);for(f=0;f<8&&d[f]===0;++f);for(f==8?d[7]=-1:d[f],c.x=d,c.i=0,f=256;f>0;--f)c.next()}u(l,a)}function o(a,l){return l.x=a.x.slice(),l.i=a.i,l}function i(a,l){a==null&&(a=+new Date);var u=new r(a),c=l&&l.state,h=function(){return(u.next()>>>0)/4294967296};return h.double=function(){do var f=u.next()>>>11,d=(u.next()>>>0)/4294967296,p=(f+d)/(1<<21);while(p===0);return p},h.int32=u.next,h.quick=h,c&&(c.x&&o(c,u),h.state=function(){return o(u,{})}),h}e&&e.exports?e.exports=i:this.xorshift7=i})(Ze,n)}(Xo);var f0=Xo.exports,Jo={exports:{}};Jo.exports,function(n){(function(t,e,s){function r(a){var l=this;l.next=function(){var c=l.w,h=l.X,f=l.i,d,p;return l.w=c=c+1640531527|0,p=h[f+34&127],d=h[f=f+1&127],p^=p<<13,d^=d<<17,p^=p>>>15,d^=d>>>12,p=h[f]=p^d,l.i=f,p+(c^c>>>16)|0};function u(c,h){var f,d,p,g,m,b=[],y=128;for(h===(h|0)?(d=h,h=null):(h=h+"\0",d=0,y=Math.max(y,h.length)),p=0,g=-32;g<y;++g)h&&(d^=h.charCodeAt((g+32)%h.length)),g===0&&(m=d),d^=d<<10,d^=d>>>15,d^=d<<4,d^=d>>>13,g>=0&&(m=m+1640531527|0,f=b[g&127]^=d+m,p=f==0?p+1:0);for(p>=128&&(b[(h&&h.length||0)&127]=-1),p=127,g=4*128;g>0;--g)d=b[p+34&127],f=b[p=p+1&127],d^=d<<13,f^=f<<17,d^=d>>>15,f^=f>>>12,b[p]=d^f;c.w=m,c.X=b,c.i=p}u(l,a)}function o(a,l){return l.i=a.i,l.w=a.w,l.X=a.X.slice(),l}function i(a,l){a==null&&(a=+new Date);var u=new r(a),c=l&&l.state,h=function(){return(u.next()>>>0)/4294967296};return h.double=function(){do var f=u.next()>>>11,d=(u.next()>>>0)/4294967296,p=(f+d)/(1<<21);while(p===0);return p},h.int32=u.next,h.quick=h,c&&(c.X&&o(c,u),h.state=function(){return o(u,{})}),h}e&&e.exports?e.exports=i:this.xor4096=i})(Ze,n)}(Jo);var d0=Jo.exports,Zo={exports:{}};Zo.exports,function(n){(function(t,e,s){function r(a){var l=this,u="";l.next=function(){var h=l.b,f=l.c,d=l.d,p=l.a;return h=h<<25^h>>>7^f,f=f-d|0,d=d<<24^d>>>8^p,p=p-h|0,l.b=h=h<<20^h>>>12^f,l.c=f=f-d|0,l.d=d<<16^f>>>16^p,l.a=p-h|0},l.a=0,l.b=0,l.c=-1640531527,l.d=1367130551,a===Math.floor(a)?(l.a=a/4294967296|0,l.b=a|0):u+=a;for(var c=0;c<u.length+20;c++)l.b^=u.charCodeAt(c)|0,l.next()}function o(a,l){return l.a=a.a,l.b=a.b,l.c=a.c,l.d=a.d,l}function i(a,l){var u=new r(a),c=l&&l.state,h=function(){return(u.next()>>>0)/4294967296};return h.double=function(){do var f=u.next()>>>11,d=(u.next()>>>0)/4294967296,p=(f+d)/(1<<21);while(p===0);return p},h.int32=u.next,h.quick=h,c&&(typeof c=="object"&&o(c,u),h.state=function(){return o(u,{})}),h}e&&e.exports?e.exports=i:this.tychei=i})(Ze,n)}(Zo);var p0=Zo.exports,Ll={exports:{}};const m0=cp(Object.freeze(Object.defineProperty({__proto__:null,default:{}},Symbol.toStringTag,{value:"Module"})));(function(n){(function(t,e,s){var r=256,o=6,i=52,a="random",l=s.pow(r,o),u=s.pow(2,i),c=u*2,h=r-1,f;function d(x,$,E){var D=[];$=$==!0?{entropy:!0}:$||{};var _=b(m($.entropy?[x,S(e)]:x??y(),3),D),T=new p(D),P=function(){for(var B=T.g(o),Y=l,j=0;B<u;)B=(B+j)*r,Y*=r,j=T.g(1);for(;B>=c;)B/=2,Y/=2,j>>>=1;return(B+j)/Y};return P.int32=function(){return T.g(4)|0},P.quick=function(){return T.g(4)/4294967296},P.double=P,b(S(T.S),e),($.pass||E||function(B,Y,j,F){return F&&(F.S&&g(F,T),B.state=function(){return g(T,{})}),j?(s[a]=B,Y):B})(P,_,"global"in $?$.global:this==s,$.state)}function p(x){var $,E=x.length,D=this,_=0,T=D.i=D.j=0,P=D.S=[];for(E||(x=[E++]);_<r;)P[_]=_++;for(_=0;_<r;_++)P[_]=P[T=h&T+x[_%E]+($=P[_])],P[T]=$;(D.g=function(B){for(var Y,j=0,F=D.i,G=D.j,q=D.S;B--;)Y=q[F=h&F+1],j=j*r+q[h&(q[F]=q[G=h&G+Y])+(q[G]=Y)];return D.i=F,D.j=G,j})(r)}function g(x,$){return $.i=x.i,$.j=x.j,$.S=x.S.slice(),$}function m(x,$){var E=[],D=typeof x,_;if($&&D=="object")for(_ in x)try{E.push(m(x[_],$-1))}catch{}return E.length?E:D=="string"?x:x+"\0"}function b(x,$){for(var E=x+"",D,_=0;_<E.length;)$[h&_]=h&(D^=$[h&_]*19)+E.charCodeAt(_++);return S($)}function y(){try{var x;return f&&(x=f.randomBytes)?x=x(r):(x=new Uint8Array(r),(t.crypto||t.msCrypto).getRandomValues(x)),S(x)}catch{var $=t.navigator,E=$&&$.plugins;return[+new Date,t,E,t.screen,S(e)]}}function S(x){return String.fromCharCode.apply(0,x)}if(b(s.random(),e),n.exports){n.exports=d;try{f=m0}catch{}}else s["seed"+a]=d})(typeof self<"u"?self:Ze,[],Math)})(Ll);var g0=Ll.exports,b0=u0,y0=c0,w0=h0,x0=f0,S0=d0,$0=p0,pn=g0;pn.alea=b0,pn.xor128=y0,pn.xorwow=w0,pn.xorshift7=x0,pn.xor4096=S0,pn.tychei=$0;var Ml=pn;/**
|
|
1672
1672
|
* @license
|
|
1673
1673
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1674
1674
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1728,7 +1728,7 @@
|
|
|
1728
1728
|
* See the License for the specific language governing permissions and
|
|
1729
1729
|
* limitations under the License.
|
|
1730
1730
|
* =============================================================================
|
|
1731
|
-
*/function
|
|
1731
|
+
*/function hr(n,t,e=1,s="float32"){if(e===0)throw new Error("Cannot have a step of zero");const r={start:n,stop:t,step:e,dtype:s};return A.runKernel(Fd,{},r)}/**
|
|
1732
1732
|
* @license
|
|
1733
1733
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1734
1734
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1758,7 +1758,7 @@
|
|
|
1758
1758
|
* See the License for the specific language governing permissions and
|
|
1759
1759
|
* limitations under the License.
|
|
1760
1760
|
* =============================================================================
|
|
1761
|
-
*/function _0(n){const e={x:v(n,"x","relu")};return A.runKernel(Ud,e)}const
|
|
1761
|
+
*/function _0(n){const e={x:v(n,"x","relu")};return A.runKernel(Ud,e)}const ms=C({relu_:_0});/**
|
|
1762
1762
|
* @license
|
|
1763
1763
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1764
1764
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1893,7 +1893,7 @@
|
|
|
1893
1893
|
* See the License for the specific language governing permissions and
|
|
1894
1894
|
* limitations under the License.
|
|
1895
1895
|
* =============================================================================
|
|
1896
|
-
*/function G0(n,t){const e=v(n,"x","squeeze","string_or_numeric");return L(e,kf(e.shape,t).newShape)}const
|
|
1896
|
+
*/function G0(n,t){const e=v(n,"x","squeeze","string_or_numeric");return L(e,kf(e.shape,t).newShape)}const fr=C({squeeze_:G0});/**
|
|
1897
1897
|
* @license
|
|
1898
1898
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1899
1899
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1908,7 +1908,7 @@
|
|
|
1908
1908
|
* See the License for the specific language governing permissions and
|
|
1909
1909
|
* limitations under the License.
|
|
1910
1910
|
* =============================================================================
|
|
1911
|
-
*/function V0(n,t=0){const e=hl(n,"tensors","stack","string_or_numeric");w(e.length>=1,()=>"Pass at least one tensor to tf.stack"),e.length>0&&w(t<=e[0].rank,()=>"Axis must be <= rank of the tensor");const s=e,r={axis:t};return A.runKernel(Md,s,r)}const
|
|
1911
|
+
*/function V0(n,t=0){const e=hl(n,"tensors","stack","string_or_numeric");w(e.length>=1,()=>"Pass at least one tensor to tf.stack"),e.length>0&&w(t<=e[0].rank,()=>"Axis must be <= rank of the tensor");const s=e,r={axis:t};return A.runKernel(Md,s,r)}const dr=C({stack_:V0});/**
|
|
1912
1912
|
* @license
|
|
1913
1913
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
1914
1914
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -1938,7 +1938,7 @@
|
|
|
1938
1938
|
* See the License for the specific language governing permissions and
|
|
1939
1939
|
* limitations under the License.
|
|
1940
1940
|
* =============================================================================
|
|
1941
|
-
*/function ei(n,t,e){if(ya(n),t!=null&&t.length!==2)throw new Error("tensor2d() requires shape to have two numbers");const s=
|
|
1941
|
+
*/function ei(n,t,e){if(ya(n),t!=null&&t.length!==2)throw new Error("tensor2d() requires shape to have two numbers");const s=Qs(n,e);if(s.length!==2&&s.length!==1)throw new Error("tensor2d() requires values to be number[][] or flat/TypedArray");if(s.length===1&&t==null)throw new Error("tensor2d() requires shape to be provided when `values` are a flat/TypedArray");return tr(n,t,s,e)}/**
|
|
1942
1942
|
* @license
|
|
1943
1943
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
1944
1944
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2028,7 +2028,7 @@
|
|
|
2028
2028
|
* See the License for the specific language governing permissions and
|
|
2029
2029
|
* limitations under the License.
|
|
2030
2030
|
* =============================================================================
|
|
2031
|
-
*/function Q0(n,t,e){if(e==null||e==="linear")return n;if(e==="relu")return N(n,j0(t));throw new Error(`Cannot compute gradient for fused activation ${e}.`)}function tb(n,t){let e=t;const s=og(n.shape,t.shape);return s.length>0&&(e=et(e,s)),L(e,n.shape)}function eb(n,t,e,s){if(t==="linear")return n;if(t==="relu")return
|
|
2031
|
+
*/function Q0(n,t,e){if(e==null||e==="linear")return n;if(e==="relu")return N(n,j0(t));throw new Error(`Cannot compute gradient for fused activation ${e}.`)}function tb(n,t){let e=t;const s=og(n.shape,t.shape);return s.length>0&&(e=et(e,s)),L(e,n.shape)}function eb(n,t,e,s){if(t==="linear")return n;if(t==="relu")return ms(n);if(t==="elu")return Al(n);if(t==="relu6")return N0(n);if(t==="prelu")return l0(n,e);if(t==="leakyrelu")return Bg(n,s);if(t==="sigmoid")return Fo(n);throw new Error(`Unknown fused activation ${t}.`)}const nb=(n,t)=>!(n>0)||t==="linear";/**
|
|
2032
2032
|
* @license
|
|
2033
2033
|
* Copyright 2019 Google LLC. All Rights Reserved.
|
|
2034
2034
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2088,7 +2088,7 @@
|
|
|
2088
2088
|
* See the License for the specific language governing permissions and
|
|
2089
2089
|
* limitations under the License.
|
|
2090
2090
|
* =============================================================================
|
|
2091
|
-
*/function ub(n){const t=v(n,"image","grayscaleToRGB"),e=t.rank-1,s=t.shape[e];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}.`);const r=new Array(t.rank);return r.fill(1,0,e),r[e]=3,
|
|
2091
|
+
*/function ub(n){const t=v(n,"image","grayscaleToRGB"),e=t.rank-1,s=t.shape[e];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}.`);const r=new Array(t.rank);return r.fill(1,0,e),r[e]=3,lr(t,r)}const cb=C({grayscaleToRGB_:ub});/**
|
|
2092
2092
|
* @license
|
|
2093
2093
|
* Copyright 2023 Google LLC.
|
|
2094
2094
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2103,7 +2103,7 @@
|
|
|
2103
2103
|
* See the License for the specific language governing permissions and
|
|
2104
2104
|
* limitations under the License.
|
|
2105
2105
|
* =============================================================================
|
|
2106
|
-
*/function hb(n){const t=v(n,"image","RGBToGrayscale"),e=t.rank-1,s=t.shape[e];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}.`);const r=t.dtype,o=ot(t,"float32"),i=Rt([.2989,.587,.114]);let a;switch(t.rank){case 2:a=
|
|
2106
|
+
*/function hb(n){const t=v(n,"image","RGBToGrayscale"),e=t.rank-1,s=t.shape[e];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}.`);const r=t.dtype,o=ot(t,"float32"),i=Rt([.2989,.587,.114]);let a;switch(t.rank){case 2:a=ds("ij,j->i",o,i);break;case 3:a=ds("ijk,k->ij",o,i);break;case 4:a=ds("ijkl,l->ijk",o,i);break;case 5:a=ds("ijklm,m->ijkl",o,i);break;case 6:a=ds("ijklmn,n->ijklm",o,i);break;default:throw new Error("Not a valid tensor rank.")}return a=ve(a,-1),ot(a,r)}const fb=C({rgbToGrayscale_:hb});/**
|
|
2107
2107
|
* @license
|
|
2108
2108
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
2109
2109
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2253,7 +2253,7 @@
|
|
|
2253
2253
|
* See the License for the specific language governing permissions and
|
|
2254
2254
|
* limitations under the License.
|
|
2255
2255
|
* =============================================================================
|
|
2256
|
-
*/async function Pb(n,t,e,s=.5,r=Number.NEGATIVE_INFINITY,o=!1){const i=v(n,"boxes","nonMaxSuppressionAsync"),a=v(t,"scores","nonMaxSuppressionAsync"),l=Wn(i,a,e,s,r,null),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);return i!==n&&i.dispose(),a!==t&&a.dispose(),{selectedIndices:Rt(p,"int32"),validOutputs:
|
|
2256
|
+
*/async function Pb(n,t,e,s=.5,r=Number.NEGATIVE_INFINITY,o=!1){const i=v(n,"boxes","nonMaxSuppressionAsync"),a=v(t,"scores","nonMaxSuppressionAsync"),l=Wn(i,a,e,s,r,null),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);return i!==n&&i.dispose(),a!==t&&a.dispose(),{selectedIndices:Rt(p,"int32"),validOutputs:Yt(g,"int32")}}const Lb=Pb;/**
|
|
2257
2257
|
* @license
|
|
2258
2258
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
2259
2259
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2298,7 +2298,7 @@
|
|
|
2298
2298
|
* See the License for the specific language governing permissions and
|
|
2299
2299
|
* limitations under the License.
|
|
2300
2300
|
* =============================================================================
|
|
2301
|
-
*/function zb(n,t="binary",e=!1,s=.5){const r=v(n,"image","threshold"),o=.2989,i=.587,a=.114,l=r.shape[0]*r.shape[1];let u=N(Rt([s]),255),c,h,f,d;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){[c,h,f]=Ul(r,[1,1,1],-1);const m=N(c,o),b=N(h,i),y=N(f,a);d=O(O(m,b),y)}else d=n;if(t==="otsu"){const m=Wm(ot(R0(d),"int32"),
|
|
2301
|
+
*/function zb(n,t="binary",e=!1,s=.5){const r=v(n,"image","threshold"),o=.2989,i=.587,a=.114,l=r.shape[0]*r.shape[1];let u=N(Rt([s]),255),c,h,f,d;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){[c,h,f]=Ul(r,[1,1,1],-1);const m=N(c,o),b=N(h,i),y=N(f,a);d=O(O(m,b),y)}else d=n;if(t==="otsu"){const m=Wm(ot(R0(d),"int32"),er([]),256);u=Ub(m,l)}const p=e?Rl(d,u):ps(d,u);return ot(N(p,255),"int32")}function Ub(n,t){let e=Rt([-1]),s=Rt([0]),r=Rt([0]),o,i,a,l,u,c;for(let h=0;h<n.size-1;h++){o=Et(n,0,h+1),i=Et(n,h+1),u=X(et(o),t),c=X(et(i),t);const f=et(N(o,hr(0,o.size)));a=X(f,et(o));const d=or(i.shape,o.size),p=O(hr(0,i.size),d),g=N(i,p);l=X(et(g),et(i));const m=J(a,l),b=J(a,l),y=N(u,c);r=N(N(y,m),b);const S=ps(r,s);s=fn(S,r,s),e=fn(S,Rt([h]),e)}return e}const Wb=C({threshold_:zb});/**
|
|
2302
2302
|
* @license
|
|
2303
2303
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
2304
2304
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2328,7 +2328,7 @@
|
|
|
2328
2328
|
* See the License for the specific language governing permissions and
|
|
2329
2329
|
* limitations under the License.
|
|
2330
2330
|
* =============================================================================
|
|
2331
|
-
*/function qb(n,t,e){const s=v(n,"a","bandPart");w(s.rank>=2,()=>`bandPart(): Rank must be at least 2, got ${s.rank}.`);const r=s.shape,[o,i]=s.shape.slice(-2);let a,l;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=v(t<0?o:t,"numLower","bandPart")):(w(t.dtype==="int32",()=>"bandPart(): numLower's dtype must be an int32."),a=
|
|
2331
|
+
*/function qb(n,t,e){const s=v(n,"a","bandPart");w(s.rank>=2,()=>`bandPart(): Rank must be at least 2, got ${s.rank}.`);const r=s.shape,[o,i]=s.shape.slice(-2);let a,l;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=v(t<0?o:t,"numLower","bandPart")):(w(t.dtype==="int32",()=>"bandPart(): numLower's dtype must be an int32."),a=fn(Dl(t,0),o,cr(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=v(e<0?i:e,"numUpper","bandPart")):(w(e.dtype==="int32",()=>"bandPart(): numUpper's dtype must be an int32."),l=fn(Dl(e,0),i,cr(e,i)));const u=L(hr(0,o,1,"int32"),[-1,1]),c=hr(0,i,1,"int32"),h=J(u,c),f=ur(Rl(h,a),Pg(h,Fn(l))),d=Un([o,i],s.dtype);return L(dr(Gl(L(s,[-1,o,i])).map(p=>fn(f,p,d))),r)}const jb=C({bandPart_:qb});/**
|
|
2332
2332
|
* @license
|
|
2333
2333
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
2334
2334
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2343,7 +2343,7 @@
|
|
|
2343
2343
|
* See the License for the specific language governing permissions and
|
|
2344
2344
|
* limitations under the License.
|
|
2345
2345
|
* =============================================================================
|
|
2346
|
-
*/function Hb(n){let t;if(Array.isArray(n)){t=!1,w(n!=null&&n.length>0,()=>"Gram-Schmidt process: input must not be null, undefined, or empty");const r=n[0].shape[0];for(let o=1;o<n.length;++o)w(n[o].shape[0]===r,()=>`Gram-Schmidt: Non-unique lengths found in the input vectors: (${n[o].shape[0]} vs. ${r})`)}else t=!0,n=Ul(n,n.shape[0],0).map(r=>
|
|
2346
|
+
*/function Hb(n){let t;if(Array.isArray(n)){t=!1,w(n!=null&&n.length>0,()=>"Gram-Schmidt process: input must not be null, undefined, or empty");const r=n[0].shape[0];for(let o=1;o<n.length;++o)w(n[o].shape[0]===r,()=>`Gram-Schmidt: Non-unique lengths found in the input vectors: (${n[o].shape[0]} vs. ${r})`)}else t=!0,n=Ul(n,n.shape[0],0).map(r=>fr(r,[0]));w(n.length<=n[0].shape[0],()=>`Gram-Schmidt: Number of vectors (${n.length}) exceeds number of dimensions (${n[0].shape[0]}).`);const e=[],s=n;for(let r=0;r<n.length;++r)e.push(A.tidy(()=>{let o=s[r];if(r>0)for(let i=0;i<r;++i){const a=N(et(N(e[i],o)),e[i]);o=J(o,a)}return X(o,Tl(o,"euclidean"))}));return t?dr(e,0):e}const Kb=C({gramSchmidt_:Hb});/**
|
|
2347
2347
|
* @license
|
|
2348
2348
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
2349
2349
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2358,7 +2358,7 @@
|
|
|
2358
2358
|
* See the License for the specific language governing permissions and
|
|
2359
2359
|
* limitations under the License.
|
|
2360
2360
|
* =============================================================================
|
|
2361
|
-
*/function Yb(n,t=!1){if(w(n.rank>=2,()=>`qr() requires input tensor to have a rank >= 2, but got rank ${n.rank}`),n.rank===2)return ql(n,t);{const e=n.shape.slice(0,n.shape.length-2).reduce((l,u)=>l*u),s=Gl(L(n,[e,n.shape[n.shape.length-2],n.shape[n.shape.length-1]]),0),r=[],o=[];s.forEach(l=>{const[u,c]=ql(l,t);r.push(u),o.push(c)});const i=L(
|
|
2361
|
+
*/function Yb(n,t=!1){if(w(n.rank>=2,()=>`qr() requires input tensor to have a rank >= 2, but got rank ${n.rank}`),n.rank===2)return ql(n,t);{const e=n.shape.slice(0,n.shape.length-2).reduce((l,u)=>l*u),s=Gl(L(n,[e,n.shape[n.shape.length-2],n.shape[n.shape.length-1]]),0),r=[],o=[];s.forEach(l=>{const[u,c]=ql(l,t);r.push(u),o.push(c)});const i=L(dr(r,0),n.shape),a=L(dr(o,0),n.shape);return[i,a]}}function ql(n,t=!1){return A.tidy(()=>{w(n.shape.length===2,()=>`qr2d() requires a 2D Tensor, but got a ${n.shape.length}D Tensor.`);const e=n.shape[0],s=n.shape[1];let r=Nl(e),o=on(n);const i=ei([[1]],[1,1]);let a=on(i);const l=e>=s?s:e;for(let u=0;u<l;++u){const c=o,h=a,f=r;[a,o,r]=A.tidy(()=>{const d=Et(o,[u,u],[e-u,1]),p=Tl(d),g=Et(o,[u,u],[1,1]),m=fn(ps(g,0),ei([[-1]]),ei([[1]])),b=J(g,N(m,p)),y=X(d,b);y.shape[0]===1?a=on(i):a=an([i,Et(y,[1,0],[y.shape[0]-1,y.shape[1]])],0);const S=Fn(X(Se(m,b),p)),x=Et(o,[u,0],[e-u,s]),$=N(S,a),E=pt(a);if(u===0)o=J(x,Se($,Se(E,x)));else{const T=J(x,Se($,Se(E,x)));o=an([Et(o,[0,0],[u,s]),T],0)}const D=pt($),_=Et(r,[0,u],[e,r.shape[1]-u]);if(u===0)r=J(_,Se(Se(_,a),D));else{const T=J(_,Se(Se(_,a),D));r=an([Et(r,[0,0],[e,u]),T],1)}return[a,o,r]}),ut([c,h,f])}return!t&&e>s&&(r=Et(r,[0,0],[e,s]),o=Et(o,[0,0],[s,s])),[r,o]})}const Xb=C({qr_:Yb});/**
|
|
2362
2362
|
* @license
|
|
2363
2363
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
2364
2364
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2373,7 +2373,7 @@
|
|
|
2373
2373
|
* See the License for the specific language governing permissions and
|
|
2374
2374
|
* limitations under the License.
|
|
2375
2375
|
* =============================================================================
|
|
2376
|
-
*/const
|
|
2376
|
+
*/const pr={flipLeftRight:lb,grayscaleToRGB:cb,resizeNearestNeighbor:Fb,resizeBilinear:Ob,rgbToGrayscale:fb,rotateWithOffset:pb,cropAndResize:ib,nonMaxSuppression:gb,nonMaxSuppressionAsync:Cb,nonMaxSuppressionWithScore:_b,nonMaxSuppressionWithScoreAsync:Nb,nonMaxSuppressionPadded:Rb,nonMaxSuppressionPaddedAsync:Lb,threshold:Wb,transform:Vb},Jb={bandPart:jb,gramSchmidt:Kb,qr:Xb};/**
|
|
2377
2377
|
* @license
|
|
2378
2378
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2379
2379
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2403,7 +2403,7 @@
|
|
|
2403
2403
|
* See the License for the specific language governing permissions and
|
|
2404
2404
|
* limitations under the License.
|
|
2405
2405
|
* =============================================================================
|
|
2406
|
-
*/class qe extends Gn{minimize(t,e=!1,s){const{value:r,grads:o}=this.computeGradients(t,s);if(s!=null){const i=s.map(a=>({name:a.name,tensor:o[a.name]}));this.applyGradients(i)}else this.applyGradients(o);return ut(o),e?r:(r.dispose(),null)}get iterations(){return this.iterations_==null&&(this.iterations_=0),this.iterations_}incrementIterations(){this.iterations_=this.iterations+1}computeGradients(t,e){return Vg(t,e)}dispose(){this.iterations_!=null&&ut(this.iterations_)}async saveIterations(){return this.iterations_==null&&(this.iterations_=0),{name:"iter",tensor:
|
|
2406
|
+
*/class qe extends Gn{minimize(t,e=!1,s){const{value:r,grads:o}=this.computeGradients(t,s);if(s!=null){const i=s.map(a=>({name:a.name,tensor:o[a.name]}));this.applyGradients(i)}else this.applyGradients(o);return ut(o),e?r:(r.dispose(),null)}get iterations(){return this.iterations_==null&&(this.iterations_=0),this.iterations_}incrementIterations(){this.iterations_=this.iterations+1}computeGradients(t,e){return Vg(t,e)}dispose(){this.iterations_!=null&&ut(this.iterations_)}async saveIterations(){return this.iterations_==null&&(this.iterations_=0),{name:"iter",tensor:Yt(this.iterations_,"int32")}}async getWeights(){throw new Error("getWeights() is not implemented for this optimizer yet.")}async setWeights(t){throw new Error(`setWeights() is not implemented for this optimizer class ${this.getClassName()}`)}async extractIterations(t){return this.iterations_=(await t[0].tensor.data())[0],t.slice(1)}}Object.defineProperty(qe,Symbol.hasInstance,{value:n=>n.minimize!=null&&n.computeGradients!=null&&n.applyGradients!=null});/**
|
|
2407
2407
|
* @license
|
|
2408
2408
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2409
2409
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2433,7 +2433,7 @@
|
|
|
2433
2433
|
* See the License for the specific language governing permissions and
|
|
2434
2434
|
* limitations under the License.
|
|
2435
2435
|
* =============================================================================
|
|
2436
|
-
*/class Hl extends qe{static get className(){return"Adagrad"}constructor(t,e=.1){super(),this.learningRate=t,this.initialAccumulatorValue=e,this.accumulatedGrads=[]}applyGradients(t){(Array.isArray(t)?t.map(s=>s.name):Object.keys(t)).forEach((s,r)=>{const o=A.registeredVariables[s];this.accumulatedGrads[r]==null&&(this.accumulatedGrads[r]={originalName:`${s}/accumulator`,variable:k(()=>
|
|
2436
|
+
*/class Hl extends qe{static get className(){return"Adagrad"}constructor(t,e=.1){super(),this.learningRate=t,this.initialAccumulatorValue=e,this.accumulatedGrads=[]}applyGradients(t){(Array.isArray(t)?t.map(s=>s.name):Object.keys(t)).forEach((s,r)=>{const o=A.registeredVariables[s];this.accumulatedGrads[r]==null&&(this.accumulatedGrads[r]={originalName:`${s}/accumulator`,variable:k(()=>or(o.shape,this.initialAccumulatorValue).variable(!1))});const i=Array.isArray(t)?t[r].tensor:t[s];if(i==null)return;const a=this.accumulatedGrads[r].variable;k(()=>{const l=O(a,Ve(i));a.assign(l);const u=O(N(X(i,fe(O(l,A.backend.epsilon()))),-this.learningRate),o);o.assign(u)})}),this.incrementIterations()}dispose(){this.accumulatedGrads!=null&&ut(this.accumulatedGrads.map(t=>t.variable))}async getWeights(){return[await this.saveIterations()].concat(this.accumulatedGrads.map(t=>({name:t.originalName,tensor:t.variable})))}async setWeights(t){t=await this.extractIterations(t);const e=!1;this.accumulatedGrads=t.map(s=>({originalName:s.name,variable:s.tensor.variable(e)}))}getConfig(){return{learningRate:this.learningRate,initialAccumulatorValue:this.initialAccumulatorValue}}static fromConfig(t,e){return new t(e.learningRate,e.initialAccumulatorValue)}}/**
|
|
2437
2437
|
* @license
|
|
2438
2438
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2439
2439
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2448,7 +2448,7 @@
|
|
|
2448
2448
|
* See the License for the specific language governing permissions and
|
|
2449
2449
|
* limitations under the License.
|
|
2450
2450
|
* =============================================================================
|
|
2451
|
-
*/class Kl extends qe{static get className(){return"Adam"}constructor(t,e,s,r=null){super(),this.learningRate=t,this.beta1=e,this.beta2=s,this.epsilon=r,this.accumulatedFirstMoment=[],this.accumulatedSecondMoment=[],k(()=>{this.accBeta1=
|
|
2451
|
+
*/class Kl extends qe{static get className(){return"Adam"}constructor(t,e,s,r=null){super(),this.learningRate=t,this.beta1=e,this.beta2=s,this.epsilon=r,this.accumulatedFirstMoment=[],this.accumulatedSecondMoment=[],k(()=>{this.accBeta1=Yt(e).variable(),this.accBeta2=Yt(s).variable()}),r==null&&(this.epsilon=A.backend.epsilon())}applyGradients(t){const e=Array.isArray(t)?t.map(s=>s.name):Object.keys(t);k(()=>{const s=J(1,this.accBeta1),r=J(1,this.accBeta2);e.forEach((o,i)=>{const a=A.registeredVariables[o],l=!1;this.accumulatedFirstMoment[i]==null&&(this.accumulatedFirstMoment[i]={originalName:`${o}/m`,variable:k(()=>$e(a).variable(l))}),this.accumulatedSecondMoment[i]==null&&(this.accumulatedSecondMoment[i]={originalName:`${o}/v`,variable:k(()=>$e(a).variable(l))});const u=Array.isArray(t)?t[i].tensor:t[o];if(u==null)return;const c=this.accumulatedFirstMoment[i].variable,h=this.accumulatedSecondMoment[i].variable,f=O(N(c,this.beta1),N(u,1-this.beta1)),d=O(N(h,this.beta2),N(Ve(u),1-this.beta2)),p=X(f,s),g=X(d,r);c.assign(f),h.assign(d);const m=O(N(X(p,O(fe(g),this.epsilon)),-this.learningRate),a);a.assign(m)}),this.accBeta1.assign(N(this.accBeta1,this.beta1)),this.accBeta2.assign(N(this.accBeta2,this.beta2))}),this.incrementIterations()}dispose(){this.accBeta1.dispose(),this.accBeta2.dispose(),this.accumulatedFirstMoment!=null&&ut(this.accumulatedFirstMoment.map(t=>t.variable)),this.accumulatedSecondMoment!=null&&ut(this.accumulatedSecondMoment.map(t=>t.variable))}async getWeights(){const t=[...this.accumulatedFirstMoment,...this.accumulatedSecondMoment];return[await this.saveIterations()].concat(t.map(e=>({name:e.originalName,tensor:e.variable})))}async setWeights(t){t=await this.extractIterations(t),k(()=>{this.accBeta1.assign(ar(this.beta1,this.iterations_+1)),this.accBeta2.assign(ar(this.beta2,this.iterations_+1))});const e=t.length/2,s=!1;this.accumulatedFirstMoment=t.slice(0,e).map(r=>({originalName:r.name,variable:r.tensor.variable(s)})),this.accumulatedSecondMoment=t.slice(e,e*2).map(r=>({originalName:r.name,variable:r.tensor.variable(s)}))}getConfig(){return{learningRate:this.learningRate,beta1:this.beta1,beta2:this.beta2,epsilon:this.epsilon}}static fromConfig(t,e){return new t(e.learningRate,e.beta1,e.beta2,e.epsilon)}}/**
|
|
2452
2452
|
* @license
|
|
2453
2453
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2454
2454
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2463,7 +2463,7 @@
|
|
|
2463
2463
|
* See the License for the specific language governing permissions and
|
|
2464
2464
|
* limitations under the License.
|
|
2465
2465
|
* =============================================================================
|
|
2466
|
-
*/class Yl extends qe{static get className(){return"Adamax"}constructor(t,e,s,r=null,o=0){super(),this.learningRate=t,this.beta1=e,this.beta2=s,this.epsilon=r,this.decay=o,this.accumulatedFirstMoment=[],this.accumulatedWeightedInfNorm=[],k(()=>{this.iteration=
|
|
2466
|
+
*/class Yl extends qe{static get className(){return"Adamax"}constructor(t,e,s,r=null,o=0){super(),this.learningRate=t,this.beta1=e,this.beta2=s,this.epsilon=r,this.decay=o,this.accumulatedFirstMoment=[],this.accumulatedWeightedInfNorm=[],k(()=>{this.iteration=Yt(0).variable(),this.accBeta1=Yt(e).variable()}),r==null&&(this.epsilon=A.backend.epsilon())}applyGradients(t){const e=Array.isArray(t)?t.map(s=>s.name):Object.keys(t);k(()=>{const s=J(1,this.accBeta1),r=X(-this.learningRate,O(N(this.iteration,this.decay),1));e.forEach((o,i)=>{const a=A.registeredVariables[o],l=!1;this.accumulatedFirstMoment[i]==null&&(this.accumulatedFirstMoment[i]={originalName:`${o}/m`,variable:$e(a).variable(l)}),this.accumulatedWeightedInfNorm[i]==null&&(this.accumulatedWeightedInfNorm[i]={originalName:`${o}/v`,variable:$e(a).variable(l)});const u=Array.isArray(t)?t[i].tensor:t[o];if(u==null)return;const c=this.accumulatedFirstMoment[i].variable,h=this.accumulatedWeightedInfNorm[i].variable,f=O(N(c,this.beta1),N(u,1-this.beta1)),d=N(h,this.beta2),p=Lt(u),g=zn(d,p);c.assign(f),h.assign(g);const m=O(N(X(r,s),X(f,O(g,this.epsilon))),a);a.assign(m)}),this.iteration.assign(O(this.iteration,1)),this.accBeta1.assign(N(this.accBeta1,this.beta1))}),this.incrementIterations()}dispose(){this.accBeta1.dispose(),this.iteration.dispose(),this.accumulatedFirstMoment!=null&&ut(this.accumulatedFirstMoment.map(t=>t.variable)),this.accumulatedWeightedInfNorm!=null&&ut(this.accumulatedWeightedInfNorm.map(t=>t.variable))}async getWeights(){throw new Error("getWeights() is not implemented for Adamax yet.")}async setWeights(t){throw new Error("setWeights() is not implemented for Adamax yet.")}getConfig(){return{learningRate:this.learningRate,beta1:this.beta1,beta2:this.beta2,epsilon:this.epsilon,decay:this.decay}}static fromConfig(t,e){return new t(e.learningRate,e.beta1,e.beta2,e.epsilon,e.decay)}}/**
|
|
2467
2467
|
* @license
|
|
2468
2468
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2469
2469
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2478,7 +2478,7 @@
|
|
|
2478
2478
|
* See the License for the specific language governing permissions and
|
|
2479
2479
|
* limitations under the License.
|
|
2480
2480
|
* =============================================================================
|
|
2481
|
-
*/class si extends qe{static get className(){return"SGD"}constructor(t){super(),this.learningRate=t,this.setLearningRate(t)}applyGradients(t){(Array.isArray(t)?t.map(s=>s.name):Object.keys(t)).forEach((s,r)=>{const o=Array.isArray(t)?t[r].tensor:t[s];if(o==null)return;const i=A.registeredVariables[s];k(()=>{const a=O(N(this.c,o),i);i.assign(a)})}),this.incrementIterations()}setLearningRate(t){this.learningRate=t,this.c!=null&&this.c.dispose(),this.c=Ln(
|
|
2481
|
+
*/class si extends qe{static get className(){return"SGD"}constructor(t){super(),this.learningRate=t,this.setLearningRate(t)}applyGradients(t){(Array.isArray(t)?t.map(s=>s.name):Object.keys(t)).forEach((s,r)=>{const o=Array.isArray(t)?t[r].tensor:t[s];if(o==null)return;const i=A.registeredVariables[s];k(()=>{const a=O(N(this.c,o),i);i.assign(a)})}),this.incrementIterations()}setLearningRate(t){this.learningRate=t,this.c!=null&&this.c.dispose(),this.c=Ln(Yt(-t))}dispose(){this.c.dispose()}async getWeights(){return[await this.saveIterations()]}async setWeights(t){if(t=await this.extractIterations(t),t.length!==0)throw new Error("SGD optimizer does not have settable weights.")}getConfig(){return{learningRate:this.learningRate}}static fromConfig(t,e){return new t(e.learningRate)}}/**
|
|
2482
2482
|
* @license
|
|
2483
2483
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2484
2484
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2493,7 +2493,7 @@
|
|
|
2493
2493
|
* See the License for the specific language governing permissions and
|
|
2494
2494
|
* limitations under the License.
|
|
2495
2495
|
* =============================================================================
|
|
2496
|
-
*/class Xl extends si{static get className(){return"Momentum"}constructor(t,e,s=!1){super(t),this.learningRate=t,this.momentum=e,this.useNesterov=s,this.accumulations=[],this.m=
|
|
2496
|
+
*/class Xl extends si{static get className(){return"Momentum"}constructor(t,e,s=!1){super(t),this.learningRate=t,this.momentum=e,this.useNesterov=s,this.accumulations=[],this.m=Yt(this.momentum)}applyGradients(t){(Array.isArray(t)?t.map(s=>s.name):Object.keys(t)).forEach((s,r)=>{const o=A.registeredVariables[s];this.accumulations[r]==null&&(this.accumulations[r]={originalName:`${s}/momentum`,variable:k(()=>$e(o).variable(!1))});const i=this.accumulations[r].variable,a=Array.isArray(t)?t[r].tensor:t[s];a!=null&&k(()=>{let l;const u=O(N(this.m,i),a);this.useNesterov?l=O(N(this.c,O(a,N(u,this.m))),o):l=O(N(this.c,u),o),i.assign(u),o.assign(l)})}),this.incrementIterations()}dispose(){this.m.dispose(),this.accumulations!=null&&ut(this.accumulations.map(t=>t.variable))}setMomentum(t){this.momentum=t}async getWeights(){return[await this.saveIterations()].concat(this.accumulations.map(t=>({name:t.originalName,tensor:t.variable})))}async setWeights(t){t=await this.extractIterations(t);const e=!1;this.accumulations=t.map(s=>({originalName:s.name,variable:s.tensor.variable(e)}))}getConfig(){return{learningRate:this.learningRate,momentum:this.momentum,useNesterov:this.useNesterov}}static fromConfig(t,e){return new t(e.learningRate,e.momentum,e.useNesterov)}}/**
|
|
2497
2497
|
* @license
|
|
2498
2498
|
* Copyright 2018 Google LLC. All Rights Reserved.
|
|
2499
2499
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2598,7 +2598,7 @@
|
|
|
2598
2598
|
* See the License for the specific language governing permissions and
|
|
2599
2599
|
* limitations under the License.
|
|
2600
2600
|
* =============================================================================
|
|
2601
|
-
*/function uy(n,t){const e=n[0].length;n.forEach((r,o)=>{w(r.length===e,()=>`Error in concat${e}D: rank of tensors[${o}] must be the same as the rank of the rest (${e})`)}),w(t>=0&&t<e,()=>`Error in concat${e}D: axis must be between 0 and ${e-1}.`);const s=n[0];n.forEach((r,o)=>{for(let i=0;i<e;i++)w(i===t||r[i]===s[i],()=>`Error in concat${e}D: Shape of tensors[${o}] (${r}) does not match the shape of the rest (${s}) along the non-concatenated axis ${o}.`)})}function
|
|
2601
|
+
*/function uy(n,t){const e=n[0].length;n.forEach((r,o)=>{w(r.length===e,()=>`Error in concat${e}D: rank of tensors[${o}] must be the same as the rank of the rest (${e})`)}),w(t>=0&&t<e,()=>`Error in concat${e}D: axis must be between 0 and ${e-1}.`);const s=n[0];n.forEach((r,o)=>{for(let i=0;i<e;i++)w(i===t||r[i]===s[i],()=>`Error in concat${e}D: Shape of tensors[${o}] (${r}) does not match the shape of the rest (${s}) along the non-concatenated axis ${o}.`)})}function gs(n,t){const e=n[0].slice();for(let s=1;s<n.length;s++)e[t]+=n[s][t];return e}/**
|
|
2602
2602
|
* @license
|
|
2603
2603
|
* Copyright 2022 Google LLC. All Rights Reserved.
|
|
2604
2604
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2720,7 +2720,7 @@
|
|
|
2720
2720
|
* See the License for the specific language governing permissions and
|
|
2721
2721
|
* limitations under the License.
|
|
2722
2722
|
* =============================================================================
|
|
2723
|
-
*/function tu(n){try{return n.map(t=>
|
|
2723
|
+
*/function tu(n){try{return n.map(t=>Ys(t))}catch(t){throw new Error(`Failed to decode encoded string bytes into utf-8, error: ${t}`)}}function Ry(n){return n.map(t=>sn(t))}/**
|
|
2724
2724
|
* @license
|
|
2725
2725
|
* Copyright 2017 Google LLC. All Rights Reserved.
|
|
2726
2726
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -2759,12 +2759,12 @@
|
|
|
2759
2759
|
* license that can be found in the LICENSE file or at
|
|
2760
2760
|
* https://opensource.org/licenses/MIT.
|
|
2761
2761
|
* =============================================================================
|
|
2762
|
-
*/function
|
|
2762
|
+
*/function mr(n,t){if(Array.isArray(n)){let e=[];for(let s=0;s<t;s++)e=e.concat(n);return e}else{const e=new Array(t);return e.fill(n),e}}function Ae(n,t){if(!n)throw new ri(t)}function eu(n,t){let e=0;for(const s of n)s===t&&e++;return e}function Ut(n){return n.length===1?n[0]:n}function nt(n){return Array.isArray(n)?n:[n]}function Le(n){const e=n.replace(/(.)([A-Z][a-z0-9]+)/g,"$1_$2").replace(/([a-z])([A-Z])/g,"$1_$2").toLowerCase();return e[0]!=="_"?e:"private"+e}function mn(n){return n.length<=1||n.indexOf("_")===-1?n:n.replace(/[_]+(\w|$)/g,(t,e)=>e.toUpperCase())}let re={};function oi(n){if(n==null)return null;const t={};return t.className=n.getClassName(),t.config=n.getConfig(),t}function ii(n){if(!(n==null||typeof n!="object"))if(Array.isArray(n))n.forEach(t=>ii(t));else{const t=Object.keys(n);for(const e of t){const s=n[e];s!=null&&typeof s=="object"&&(!Array.isArray(s)&&s.type==="ndarray"&&typeof s.value=="number"?n[e]=s.value:ii(s))}}}function bs(n,t={},e={},s="object",r=!1){if(typeof n=="string"){const o=n;let i;if(o in e)i=e[o];else if(o in re)i=re[o];else if(i=t[o],i==null)throw new I(`Unknown ${s}: ${n}. This may be due to one of the following reasons:
|
|
2763
2763
|
1. The ${s} is defined in Python, in which case it needs to be ported to TensorFlow.js or your JavaScript code.
|
|
2764
2764
|
2. The custom ${s} is defined in JavaScript, but is not registered properly with tf.serialization.registerClass().`);return i}else{const o=n;if(o.className==null||o.config==null)throw new I(`${s}: Improper config format: ${JSON.stringify(o)}.
|
|
2765
2765
|
'className' and 'config' must set.`);const i=o.className;let a,l;if(i in e?[a,l]=e[i]:i in re?[a,l]=re.className:i in t&&([a,l]=t[i]),a==null)throw new I(`Unknown ${s}: ${i}. This may be due to one of the following reasons:
|
|
2766
2766
|
1. The ${s} is defined in Python, in which case it needs to be ported to TensorFlow.js or your JavaScript code.
|
|
2767
|
-
2. The custom ${s} is defined in JavaScript, but is not registered properly with tf.serialization.registerClass().`);if(l!=null){const u={};for(const d of Object.keys(re))u[d]=re[d];for(const d of Object.keys(e))u[d]=e[d];const c=o.config;c.customObjects=u;const h=Object.assign({},re);for(const d of Object.keys(e))re[d]=e[d];ii(o.config);const f=l(a,o.config,e,r);return re=Object.assign({},h),f}else{const u=Object.assign({},re);for(const h of Object.keys(e))re[h]=e[h];const c=new a(o.config);return re=Object.assign({},u),c}}}function By(n,t){return n<t?-1:n>t?1:0}function
|
|
2767
|
+
2. The custom ${s} is defined in JavaScript, but is not registered properly with tf.serialization.registerClass().`);if(l!=null){const u={};for(const d of Object.keys(re))u[d]=re[d];for(const d of Object.keys(e))u[d]=e[d];const c=o.config;c.customObjects=u;const h=Object.assign({},re);for(const d of Object.keys(e))re[d]=e[d];ii(o.config);const f=l(a,o.config,e,r);return re=Object.assign({},h),f}else{const u=Object.assign({},re);for(const h of Object.keys(e))re[h]=e[h];const c=new a(o.config);return re=Object.assign({},u),c}}}function By(n,t){return n<t?-1:n>t?1:0}function gr(n,t){return-1*By(n,t)}function gn(n){if(n==null)return n;const t=[];for(const e of n)t.indexOf(e)===-1&&t.push(e);return t}function Fy(n){if(n==null)throw new I(`Invalid value in obj: ${JSON.stringify(n)}`);for(const t in n)if(n.hasOwnProperty(t))return!1;return!0}function qn(n,t,e){if(e!=null&&n.indexOf(e)<0)throw new I(`${e} is not a valid ${t}. Valid values are ${n} or null/undefined.`)}function ai(n,t,e=0,s=1/0){return Ae(e>=0),Ae(s>=e),Array.isArray(n)&&n.length>=e&&n.length<=s&&n.every(r=>typeof r===t)}function Me(n,t){Array.isArray(n)?(w(n.length>0,()=>`${t} is unexpectedly an empty array.`),n.forEach((e,s)=>Me(e,`element ${s+1} of ${t}`))):w(Number.isInteger(n)&&n>0,()=>`Expected ${t} to be a positive integer, but got ${nu(n)}.`)}function nu(n){return n===null?"null":Array.isArray(n)?"["+n.map(t=>nu(t)).join(",")+"]":typeof n=="string"?`"${n}"`:`${n}`}function zy(n,t,e){let s=e!=null?e():Nn(),r;return(...i)=>{const a=e!=null?e():Nn();return a-s<t||(s=a,r=n(...i)),r}}function Uy(n){return n==="relu"?"relu":n==="linear"?"linear":n==="elu"?"elu":null}/**
|
|
2768
2768
|
* @license
|
|
2769
2769
|
* Copyright 2018 Google LLC
|
|
2770
2770
|
*
|
|
@@ -2772,7 +2772,7 @@
|
|
|
2772
2772
|
* license that can be found in the LICENSE file or at
|
|
2773
2773
|
* https://opensource.org/licenses/MIT.
|
|
2774
2774
|
* =============================================================================
|
|
2775
|
-
*/const jn=new Map;function mt(n){qn(Py,"DataFormat",n)}function Wy(n){qn(Ly,"InterpolationFormat",n)}function oe(n){qn(My,"PaddingMode",n)}function su(n){qn(Oy,"PoolMode",n)}const
|
|
2775
|
+
*/const jn=new Map;function mt(n){qn(Py,"DataFormat",n)}function Wy(n){qn(Ly,"InterpolationFormat",n)}function oe(n){qn(My,"PaddingMode",n)}function su(n){qn(Oy,"PoolMode",n)}const ys=[],ru="/";function br(n,t){ys.push(n);try{const e=t();return ys.pop(),e}catch(e){throw ys.pop(),e}}function Gy(){return ys.length===0?"":ys.join(ru)+ru}function ou(n){if(!au(n))throw new Error("Not a valid tensor name: '"+n+"'");return Gy()+n}function iu(n){if(!au(n))throw new Error("Not a valid tensor name: '"+n+"'");jn.has(n)||jn.set(n,0);const t=jn.get(n);if(jn.set(n,jn.get(n)+1),t>0){const e=`${n}_${t}`;return jn.set(e,1),e}else return n}const Vy=new RegExp(/^[A-Za-z0-9][-A-Za-z0-9\._\/]*$/);function au(n){return!!n.match(Vy)}/**
|
|
2776
2776
|
* @license
|
|
2777
2777
|
* Copyright 2018 Google LLC
|
|
2778
2778
|
*
|
|
@@ -2780,7 +2780,7 @@
|
|
|
2780
2780
|
* license that can be found in the LICENSE file or at
|
|
2781
2781
|
* https://opensource.org/licenses/MIT.
|
|
2782
2782
|
* =============================================================================
|
|
2783
|
-
*/function qy(n){return n===parseInt(n.toString(),10)}function
|
|
2783
|
+
*/function qy(n){return n===parseInt(n.toString(),10)}function ws(n,t,e){t==null&&(t=0),e==null&&(e=n.length);let s=1;for(let r=t;r<e;++r)s*=n[r];return s}function lu(n){if(n.length===0)return Number.NaN;let t=Number.NEGATIVE_INFINITY;for(let e=0;e<n.length;e++){const s=n[e];s>t&&(t=s)}return t}function yr(n,t){if(t<n)throw new I(`end (${t}) < begin (${n}) is forbidden.`);const e=[];for(let s=n;s<t;++s)e.push(s);return e}/**
|
|
2784
2784
|
* @license
|
|
2785
2785
|
* Copyright 2018 Google LLC
|
|
2786
2786
|
*
|
|
@@ -2796,7 +2796,7 @@
|
|
|
2796
2796
|
* license that can be found in the LICENSE file or at
|
|
2797
2797
|
* https://opensource.org/licenses/MIT.
|
|
2798
2798
|
* =============================================================================
|
|
2799
|
-
*/function uu(n,t){return ot(n,t)}function ui(n,t=-1){const e=n.shape.slice();return t<0&&(t=e.length+t+1),e.splice(t,0,1),L(n,e)}function jy(n){const t=[
|
|
2799
|
+
*/function uu(n,t){return ot(n,t)}function ui(n,t=-1){const e=n.shape.slice();return t<0&&(t=e.length+t+1),e.splice(t,0,1),L(n,e)}function jy(n){const t=[ws(n.shape)];return L(n,t)}function bn(n,t,e){return k(()=>{switch(n.rank){case 1:return Qo(n,t,e);case 2:return Fl(n,[t,0],[e,n.shape[1]]);case 3:return ti(n,[t,0,0],[e,n.shape[1],n.shape[2]]);case 4:return hs(n,[t,0,0,0],[e,n.shape[1],n.shape[2],n.shape[3]]);case 5:return Et(n,[t,0,0,0,0],[e,n.shape[1],n.shape[2],n.shape[3],n.shape[4]]);case 6:return Et(n,[t,0,0,0,0,0],[e,n.shape[1],n.shape[2],n.shape[3],n.shape[4],n.shape[5]]);default:throw new I(`sliceAlongFirstAxis() received an unsupported tensor rank: ${n.rank}`)}})}function ci(n,t,e){return k(()=>{switch(n.rank){case 1:return Qo(n,t,e);case 2:return Fl(n,[0,t],[n.shape[0],e]);case 3:return ti(n,[0,0,t],[n.shape[0],n.shape[1],e]);case 4:return hs(n,[0,0,0,t],[n.shape[0],n.shape[1],n.shape[2],e]);default:throw new I(`sliceAlongLastAxis() received an unsupported tensor rank: ${n.rank}`)}})}function wr(n,t,e,s){return k(()=>{switch(n.rank){case 1:return Qo(n,t,e);case 2:switch(s){case 1:return bn(n,t,e);case 2:return ci(n,t,e);default:throw new I(`The axis is not within the rank of the tensor ${s}`)}case 3:switch(s){case 1:return bn(n,t,e);case 2:return ti(n,[0,t,0],[n.shape[0],e,n.shape[2]]);case 3:return ci(n,t,e);default:throw new I(`The axis is not within the rank of the tensor ${s}`)}case 4:switch(s){case 1:return bn(n,t,e);case 2:return hs(n,[0,t,0,0],[n.shape[0],e,n.shape[2],n.shape[3]]);case 3:return hs(n,[0,0,t,0],[n.shape[0],n.shape[1],e,n.shape[3]]);case 4:return ci(n,t,e);default:throw new I(`The axis is not within the rank of the tensor ${s}`)}default:throw new I(`sliceAlongLastAxis() received an unsupported tensor rank: ${n.rank}`)}})}function Hy(n,t=-1){let e;return t<0&&(e=n[0].rank,e!==0?t=e:t=0),t===n[0].rank&&(t=-1),an(n,t)}function cu(n,t=0,e=1,s,r){return A0(n,t,e,s,r)}function Ky(n,t,e){return k(()=>(Array.isArray(t)?t=Rt(t,"int32"):t=ot(t,"int32"),Ng(n,t,e)))}function xs(n){return N(n,n)}function Yy(n,t,e){const s=t.shape;if(t.rank!==1&&t.rank!==n)throw new I(`Unexpected bias dimensions: ${t.rank}; expected it to be 1 or ${n}`);if(n===5){if(e==="channelsFirst")return s.length===1?L(t,[1,s[0],1,1,1]):L(t,[1,s[3],s[0],s[1],s[2]]);if(e==="channelsLast")return s.length===1?L(t,[1,1,1,1,s[0]]):L(t,[1].concat(s))}else if(n===4){if(e==="channelsFirst")return s.length===1?L(t,[1,s[0],1,1]):L(t,[1,s[2],s[0],s[1]]);if(e==="channelsLast")return s.length===1?L(t,[1,1,1,s[0]]):L(t,[1].concat(s))}else if(n===3){if(e==="channelsFirst")return s.length===1?L(t,[1,s[0],1]):L(t,[1,s[1],s[0]]);if(e==="channelsLast")return s.length===1?L(t,[1,1,s[0]]):L(t,[1].concat(s))}else if(n<3)return t;throw new I(`Unsupported input rank by biasAdd: ${t.rank}`)}function Ss(n,t,e){return k(()=>(e==null&&(e=Hn()),mt(e),O(n,Yy(n.rank,t,e))))}function Xy(n,t=1){if(t!==1)throw new Z(`Support for alpha values other than 1 (${t}) is not implemented yet.`);return Al(n)}function Jy(n){return k(()=>X(n,O(Lt(n),1)))}function Zy(n){return k(()=>{const t=O(.5,N(.2,n));return he(t,0,1)})}/**
|
|
2800
2800
|
* @license
|
|
2801
2801
|
* Copyright 2018 Google LLC
|
|
2802
2802
|
*
|
|
@@ -2804,7 +2804,7 @@
|
|
|
2804
2804
|
* license that can be found in the LICENSE file or at
|
|
2805
2805
|
* https://opensource.org/licenses/MIT.
|
|
2806
2806
|
* =============================================================================
|
|
2807
|
-
*/class Ct extends Gn{getConfig(){return{}}}class hu extends Ct{apply(t,e=1){return Xy(t,e)}}hu.className="elu",M(hu);class fu extends Ct{apply(t){return L0(t)}}fu.className="selu",M(fu);class du extends Ct{apply(t){return
|
|
2807
|
+
*/class Ct extends Gn{getConfig(){return{}}}class hu extends Ct{apply(t,e=1){return Xy(t,e)}}hu.className="elu",M(hu);class fu extends Ct{apply(t){return L0(t)}}fu.className="selu",M(fu);class du extends Ct{apply(t){return ms(t)}}du.className="relu",M(du);class pu extends Ct{apply(t){return k(()=>cr(6,ms(t)))}}pu.className="relu6",M(pu);class mu extends Ct{apply(t){return t}}mu.className="linear",M(mu);class gu extends Ct{apply(t){return Fo(t)}}gu.className="sigmoid",M(gu);class bu extends Ct{apply(t){return Zy(t)}}bu.className="hardSigmoid",M(bu);class yu extends Ct{apply(t){return qo(t)}}yu.className="softplus",M(yu);class wu extends Ct{apply(t){return Jy(t)}}wu.className="softsign",M(wu);class xu extends Ct{apply(t){return zo(t)}}xu.className="tanh",M(xu);class Su extends Ct{apply(t,e=-1){return zl(t,e)}}Su.className="softmax",M(Su);class $u extends Ct{apply(t,e=-1){return Yg(t,e)}}$u.className="logSoftmax",M($u);class vu extends Ct{apply(t){return k(()=>k(()=>{const e=Math.sqrt(2),s=N(.5,O(1,fg(X(t,e))));return N(t,s)}))}}vu.className="gelu",M(vu);class Iu extends Ct{apply(t){return k(()=>N(.5,N(t,O(1,zo(N(fe(X(2,Math.PI)),O(t,N(.044715,ar(t,3)))))))))}}Iu.className="gelu_new",M(Iu);class Au extends Ct{apply(t){return k(()=>N(t,zo(qo(t))))}}Au.className="mish",M(Au);class Eu extends Ct{apply(t,e=1){return k(()=>N(Fo(N(t,e)),t))}}Eu.className="swish",M(Eu);function Qy(n){return n.getClassName()}function hi(n,t={}){return bs(n,se.getMap().classNameMap,t,"activation")}function tw(n){if(n==null){const t={};return t.className="linear",t.config={},hi(t)}if(typeof n=="string"){const t={};return t.className=n,t.config={},hi(t)}else return n instanceof Ct?n:hi(n)}/**
|
|
2808
2808
|
* @license
|
|
2809
2809
|
* Copyright 2018 Google LLC
|
|
2810
2810
|
*
|
|
@@ -2812,7 +2812,7 @@
|
|
|
2812
2812
|
* license that can be found in the LICENSE file or at
|
|
2813
2813
|
* https://opensource.org/licenses/MIT.
|
|
2814
2814
|
* =============================================================================
|
|
2815
|
-
*/function fi(n,t){return k(()=>fe(et(N(n,n),t,!0)))}class
|
|
2815
|
+
*/function fi(n,t){return k(()=>fe(et(N(n,n),t,!0)))}class $s extends Gn{getConfig(){return{}}}class Cu extends $s{constructor(t){super(),this.defaultMaxValue=2,this.defaultAxis=0,this.maxValue=t.maxValue!=null?t.maxValue:this.defaultMaxValue,this.axis=t.axis!=null?t.axis:this.defaultAxis}apply(t){return k(()=>{const e=fi(t,this.axis),s=he(e,0,this.maxValue);return N(t,X(s,O(gt(),e)))})}getConfig(){return{maxValue:this.maxValue,axis:this.axis}}}Cu.className="MaxNorm",M(Cu);class ku extends $s{constructor(t){super(),this.defaultAxis=0,this.axis=t.axis!=null?t.axis:this.defaultAxis}apply(t){return k(()=>X(t,O(gt(),fi(t,this.axis))))}getConfig(){return{axis:this.axis}}}ku.className="UnitNorm",M(ku);class _u extends $s{apply(t){return ms(t)}}_u.className="NonNeg",M(_u);class Tu extends $s{constructor(t){super(),this.defaultMinValue=0,this.defaultMaxValue=1,this.defaultRate=1,this.defaultAxis=0,this.minValue=t.minValue!=null?t.minValue:this.defaultMinValue,this.maxValue=t.maxValue!=null?t.maxValue:this.defaultMaxValue,this.rate=t.rate!=null?t.rate:this.defaultRate,this.axis=t.axis!=null?t.axis:this.defaultAxis}apply(t){return k(()=>{const e=fi(t,this.axis),s=O(N(this.rate,he(e,this.minValue,this.maxValue)),N(1-this.rate,e));return N(t,X(s,O(gt(),e)))})}getConfig(){return{minValue:this.minValue,maxValue:this.maxValue,rate:this.rate,axis:this.axis}}}Tu.className="MinMaxNorm",M(Tu);const Nu={maxNorm:"MaxNorm",minMaxNorm:"MinMaxNorm",nonNeg:"NonNeg",unitNorm:"UnitNorm"};function xr(n){return oi(n)}function Du(n,t={}){return bs(n,se.getMap().classNameMap,t,"constraint")}function Sr(n){if(n==null)return null;if(typeof n=="string"){const e={className:n in Nu?Nu[n]:n,config:{}};return Du(e)}else return n instanceof $s?n:Du(n)}/**
|
|
2816
2816
|
* @license
|
|
2817
2817
|
* Copyright 2018 Google LLC
|
|
2818
2818
|
*
|
|
@@ -2820,7 +2820,7 @@
|
|
|
2820
2820
|
* license that can be found in the LICENSE file or at
|
|
2821
2821
|
* https://opensource.org/licenses/MIT.
|
|
2822
2822
|
* =============================================================================
|
|
2823
|
-
*/let ew=0;function Ru(){return ew++}const
|
|
2823
|
+
*/let ew=0;function Ru(){return ew++}const $r={};function di(n=""){return n in $r||($r[n]=0),$r[n]+=1,n+$r[n].toString()}/**
|
|
2824
2824
|
* @license
|
|
2825
2825
|
* Copyright 2018 Google LLC
|
|
2826
2826
|
*
|
|
@@ -2836,7 +2836,7 @@
|
|
|
2836
2836
|
* license that can be found in the LICENSE file or at
|
|
2837
2837
|
* https://opensource.org/licenses/MIT.
|
|
2838
2838
|
* =============================================================================
|
|
2839
|
-
*/function rw(n){qn(nw,"FanMode",n)}function ow(n){qn(sw,"Distribution",n)}class Ee extends Gn{fromConfigUsesCustomObjects(){return!1}getConfig(){return{}}}class Pu extends Ee{apply(t,e){return Un(t,e)}}Pu.className="Zeros",M(Pu);class Lu extends Ee{apply(t,e){return jo(t,e)}}Lu.className="Ones",M(Lu);class Mu extends Ee{constructor(t){if(super(),typeof t!="object")throw new I(`Expected argument of type ConstantConfig but got ${t}`);if(t.value===void 0)throw new I(`config must have value set but got ${t}`);this.value=t.value}apply(t,e){return k(()=>N(
|
|
2839
|
+
*/function rw(n){qn(nw,"FanMode",n)}function ow(n){qn(sw,"Distribution",n)}class Ee extends Gn{fromConfigUsesCustomObjects(){return!1}getConfig(){return{}}}class Pu extends Ee{apply(t,e){return Un(t,e)}}Pu.className="Zeros",M(Pu);class Lu extends Ee{apply(t,e){return jo(t,e)}}Lu.className="Ones",M(Lu);class Mu extends Ee{constructor(t){if(super(),typeof t!="object")throw new I(`Expected argument of type ConstantConfig but got ${t}`);if(t.value===void 0)throw new I(`config must have value set but got ${t}`);this.value=t.value}apply(t,e){return k(()=>N(Yt(this.value),jo(t,e)))}getConfig(){return{value:this.value}}}Mu.className="Constant",M(Mu);class Ou extends Ee{constructor(t){super(),this.DEFAULT_MINVAL=-.05,this.DEFAULT_MAXVAL=.05,this.minval=t.minval||this.DEFAULT_MINVAL,this.maxval=t.maxval||this.DEFAULT_MAXVAL,this.seed=t.seed}apply(t,e){return Bl(t,this.minval,this.maxval,e,this.seed)}getConfig(){return{minval:this.minval,maxval:this.maxval,seed:this.seed}}}Ou.className="RandomUniform",M(Ou);class Bu extends Ee{constructor(t){super(),this.DEFAULT_MEAN=0,this.DEFAULT_STDDEV=.05,this.mean=t.mean||this.DEFAULT_MEAN,this.stddev=t.stddev||this.DEFAULT_STDDEV,this.seed=t.seed}apply(t,e){if(e=e||"float32",e!=="float32"&&e!=="int32")throw new Z(`randomNormal does not support dType ${e}.`);return cu(t,this.mean,this.stddev,e,this.seed)}getConfig(){return{mean:this.mean,stddev:this.stddev,seed:this.seed}}}Bu.className="RandomNormal",M(Bu);class Fu extends Ee{constructor(t){super(),this.DEFAULT_MEAN=0,this.DEFAULT_STDDEV=.05,this.mean=t.mean||this.DEFAULT_MEAN,this.stddev=t.stddev||this.DEFAULT_STDDEV,this.seed=t.seed}apply(t,e){if(e=e||"float32",e!=="float32"&&e!=="int32")throw new Z(`truncatedNormal does not support dType ${e}.`);return Wl(t,this.mean,this.stddev,e,this.seed)}getConfig(){return{mean:this.mean,stddev:this.stddev,seed:this.seed}}}Fu.className="TruncatedNormal",M(Fu);class zu extends Ee{constructor(t){super(),this.gain=t.gain!=null?t.gain:1}apply(t,e){return k(()=>{if(t.length!==2||t[0]!==t[1])throw new I("Identity matrix initializer can only be used for 2D square matrices.");return N(this.gain,Nl(t[0]))})}getConfig(){return{gain:this.gain}}}zu.className="Identity",M(zu);function iw(n,t="channelsLast"){let e,s;if(mt(t),n.length===2)e=n[0],s=n[1];else if([3,4,5].indexOf(n.length)!==-1){if(t==="channelsFirst"){const r=ws(n,2);e=n[1]*r,s=n[0]*r}else if(t==="channelsLast"){const r=ws(n,0,n.length-2);e=n[n.length-2]*r,s=n[n.length-1]*r}}else{const r=ws(n);e=Math.sqrt(r),s=Math.sqrt(r)}return[e,s]}class Wt extends Ee{constructor(t){if(super(),t.scale<0)throw new I(`scale must be a positive float. Got: ${t.scale}`);this.scale=t.scale==null?1:t.scale,this.mode=t.mode==null?"fanIn":t.mode,rw(this.mode),this.distribution=t.distribution==null?"normal":t.distribution,ow(this.distribution),this.seed=t.seed}apply(t,e){const s=iw(t),r=s[0],o=s[1];let i=this.scale;if(this.mode==="fanIn"?i/=Math.max(1,r):this.mode==="fanOut"?i/=Math.max(1,o):i/=Math.max(1,(r+o)/2),this.distribution==="normal"){const a=Math.sqrt(i);if(e=e||"float32",e!=="float32"&&e!=="int32")throw new Z(`${this.getClassName()} does not support dType ${e}.`);return Wl(t,0,a,e,this.seed)}else{const a=Math.sqrt(3*i);return Bl(t,-a,a,e,this.seed)}}getConfig(){return{scale:this.scale,mode:this.mode,distribution:this.distribution,seed:this.seed}}}Wt.className="VarianceScaling",M(Wt);class pi extends Wt{constructor(t){super({scale:1,mode:"fanAvg",distribution:"uniform",seed:t==null?null:t.seed})}getClassName(){return Wt.className}}pi.className="GlorotUniform",M(pi);class mi extends Wt{constructor(t){super({scale:1,mode:"fanAvg",distribution:"normal",seed:t==null?null:t.seed})}getClassName(){return Wt.className}}mi.className="GlorotNormal",M(mi);class gi extends Wt{constructor(t){super({scale:2,mode:"fanIn",distribution:"normal",seed:t==null?null:t.seed})}getClassName(){return Wt.className}}gi.className="HeNormal",M(gi);class bi extends Wt{constructor(t){super({scale:2,mode:"fanIn",distribution:"uniform",seed:t==null?null:t.seed})}getClassName(){return Wt.className}}bi.className="HeUniform",M(bi);class yi extends Wt{constructor(t){super({scale:1,mode:"fanIn",distribution:"normal",seed:t==null?null:t.seed})}getClassName(){return Wt.className}}yi.className="LeCunNormal",M(yi);class wi extends Wt{constructor(t){super({scale:1,mode:"fanIn",distribution:"uniform",seed:t==null?null:t.seed})}getClassName(){return Wt.className}}wi.className="LeCunUniform",M(wi);class Uu extends Ee{constructor(t){super(),this.DEFAULT_GAIN=1,this.ELEMENTS_WARN_SLOW=2e3,this.gain=t.gain==null?this.DEFAULT_GAIN:t.gain,this.seed=t.seed}apply(t,e){return k(()=>{if(t.length<2)throw new Z("Shape must be at least 2D.");if(e!=="int32"&&e!=="float32"&&e!==void 0)throw new TypeError(`Unsupported data type ${e}.`);e=e;const s=z(t.slice(0,-1)),r=t[t.length-1],o=s*r;o>this.ELEMENTS_WARN_SLOW&&console.warn(`Orthogonal initializer is being called on a matrix with more than ${this.ELEMENTS_WARN_SLOW} (${o}) elements: Slowness may result.`);const i=[Math.max(r,s),Math.min(r,s)],a=cu(i,0,1,e,this.seed),l=Jb.qr(a,!1);let u=l[0];const h=l[1].flatten().stridedSlice([0],[Math.min(r,s)*Math.min(r,s)],[Math.min(r,s)+1]);return u=N(u,h.sign()),s<r&&(u=u.transpose()),N(Yt(this.gain),u.reshape(t))})}getConfig(){return{gain:this.gain,seed:this.seed}}}Uu.className="Orthogonal",M(Uu);const Wu={constant:"Constant",glorotNormal:"GlorotNormal",glorotUniform:"GlorotUniform",heNormal:"HeNormal",heUniform:"HeUniform",identity:"Identity",leCunNormal:"LeCunNormal",leCunUniform:"LeCunUniform",ones:"Ones",orthogonal:"Orthogonal",randomNormal:"RandomNormal",randomUniform:"RandomUniform",truncatedNormal:"TruncatedNormal",varianceScaling:"VarianceScaling",zeros:"Zeros"};function Gu(n,t={}){return bs(n,se.getMap().classNameMap,t,"initializer")}function vr(n){return oi(n)}function vs(n){if(typeof n=="string"){const t=n in Wu?Wu[n]:n;if(t==="GlorotNormal")return new mi;if(t==="GlorotUniform")return new pi;if(t==="HeNormal")return new gi;if(t==="HeUniform")return new bi;if(t==="LeCunNormal")return new yi;if(t==="LeCunUniform")return new wi;{const e={};return e.className=t,e.config={},Gu(e)}}else return n instanceof Ee?n:Gu(n)}/**
|
|
2840
2840
|
* @license
|
|
2841
2841
|
* Copyright 2018 Google LLC
|
|
2842
2842
|
*
|
|
@@ -2844,7 +2844,7 @@
|
|
|
2844
2844
|
* license that can be found in the LICENSE file or at
|
|
2845
2845
|
* https://opensource.org/licenses/MIT.
|
|
2846
2846
|
* =============================================================================
|
|
2847
|
-
*/function
|
|
2847
|
+
*/function Ir(n){return n.length===0?[]:Array.isArray(n[0])?n:[n]}function Gt(n){let t;if(Array.isArray(n)){if(n.length!==1)throw new I(`Expected Tensor length to be 1; got ${n.length}`);t=n[0]}else t=n;return t}function de(n){if(Array.isArray(n)&&Array.isArray(n[0])){if(n.length===1)return n=n,n[0];throw new I(`Expected exactly 1 Shape; got ${n.length}`)}else return n}/**
|
|
2848
2848
|
* @license
|
|
2849
2849
|
* Copyright 2018 Google LLC
|
|
2850
2850
|
*
|
|
@@ -2852,7 +2852,7 @@
|
|
|
2852
2852
|
* license that can be found in the LICENSE file or at
|
|
2853
2853
|
* https://opensource.org/licenses/MIT.
|
|
2854
2854
|
* =============================================================================
|
|
2855
|
-
*/function
|
|
2855
|
+
*/function Ar(n){let t=0;for(const e of n)e.shape.length===0?t+=1:t+=e.shape.reduce((s,r)=>s*r);return t}/**
|
|
2856
2856
|
* @license
|
|
2857
2857
|
* Copyright 2018 Google LLC
|
|
2858
2858
|
*
|
|
@@ -2868,7 +2868,7 @@
|
|
|
2868
2868
|
* license that can be found in the LICENSE file or at
|
|
2869
2869
|
* https://opensource.org/licenses/MIT.
|
|
2870
2870
|
* =============================================================================
|
|
2871
|
-
*/class Ce{constructor(t){this.dtype=t.dtype,this.shape=t.shape,t.shape!=null?this.ndim=t.shape.length:this.ndim=t.ndim,this.maxNDim=t.maxNDim,this.minNDim=t.minNDim,this.axes=t.axes||{}}}class bn{constructor(t,e,s,r,o,i,a){this.dtype=t,this.shape=e,this.sourceLayer=s,this.inputs=r,this.callArgs=o,this.outputTensorIndex=a,this.id=Ru(),i!=null&&(this.originalName=ou(i),this.name=iu(this.originalName)),this.rank=e.length}}let uw=0;class xi{constructor(t,e){this.callArgs=e,this.id=uw++,this.outboundLayer=t.outboundLayer,this.inboundLayers=t.inboundLayers,this.nodeIndices=t.nodeIndices,this.tensorIndices=t.tensorIndices,this.inputTensors=t.inputTensors,this.outputTensors=t.outputTensors,this.inputMasks=t.inputMasks,this.outputMasks=t.outputMasks,this.inputShapes=t.inputShapes,this.outputShapes=t.outputShapes;for(const s of t.inboundLayers)s!=null&&s.outboundNodes.push(this);t.outboundLayer.inboundNodes.push(this)}getConfig(){const t=[];for(const e of this.inboundLayers)e!=null?t.push(e.name):t.push(null);return{outboundLayer:this.outboundLayer?this.outboundLayer.name:null,inboundLayers:t,nodeIndices:this.nodeIndices,tensorIndices:this.tensorIndices}}}let cw=0;class pe extends Gn{constructor(t={}){super(),this._callHook=null,this._addedWeightNames=[],this._stateful=!1,this.id=cw++,this.activityRegularizer=null,this.inputSpec=null,this.supportsMasking=!1,this._trainableWeights=[],this._nonTrainableWeights=[],this._losses=[],this._updates=[],this._built=!1,this.inboundNodes=[],this.outboundNodes=[];let e=t.name;if(!e){const s=this.getClassName();e=Le(s)+"_"+di(s)}if(this.name=e,this.trainable_=t.trainable==null?!0:t.trainable,t.inputShape!=null||t.batchInputShape!=null){let s;if(t.batchInputShape!=null)s=t.batchInputShape;else if(t.inputShape!=null){let o=null;t.batchSize!=null&&(o=t.batchSize),s=[o].concat(t.inputShape)}this.batchInputShape=s;let r=t.dtype;r==null&&(r=t.inputDType),r==null&&(r="float32"),this.dtype=r}t.weights!=null?this.initialWeights=t.weights:this.initialWeights=null,this._refCount=null,this.fastWeightInitDuringBuild=!1}static nodeKey(t,e){return t.name+"_ib-"+e.toString()}getNodeAtIndex(t,e){if(this.inboundNodes.length===0)throw new He(`The layer has never been called and thus has no defined ${e}.`);if(this.inboundNodes.length<=t)throw new I(`Asked to get ${e} at node ${t}, but the layer has only ${this.inboundNodes.length} inbound nodes.`);return this.inboundNodes[t]}getInputAt(t){return Ut(this.getNodeAtIndex(t,"input").inputTensors)}getOutputAt(t){return Ut(this.getNodeAtIndex(t,"output").outputTensors)}get input(){if(this.inboundNodes.length>1)throw new je(`Layer ${this.name} has multiple inbound nodes, hence the notion of "layer input" is ill-defined. Use \`getInputAt(nodeIndex)\` instead.`);if(this.inboundNodes.length===0)throw new je(`Layer ${this.name} is not connected, no input to return.`);return Ut(this.getNodeAtIndex(0,"input").inputTensors)}get output(){if(this.inboundNodes.length===0)throw new je(`Layer ${this.name} has no inbound nodes.`);if(this.inboundNodes.length>1)throw new je(`Layer ${this.name} has multiple inbound nodes, hence the notion of "layer output" is ill-defined. Use \`getOutputAt(nodeIndex)\` instead.`);return Ut(this.getNodeAtIndex(0,"output").outputTensors)}get losses(){return this._losses}calculateLosses(){return this.losses.map(t=>t())}get updates(){return this._updates}get built(){return this._built}set built(t){this._built=t}get trainable(){return this.trainable_}set trainable(t){this._trainableWeights.forEach(e=>e.trainable=t),this.trainable_=t}get trainableWeights(){return this.trainable_?this._trainableWeights.filter(t=>t.trainable):[]}set trainableWeights(t){this._trainableWeights=t}get nonTrainableWeights(){return this.trainable?this._trainableWeights.filter(t=>!t.trainable).concat(this._nonTrainableWeights):this._trainableWeights.concat(this._nonTrainableWeights)}set nonTrainableWeights(t){this._nonTrainableWeights=t}get weights(){return this.trainableWeights.concat(this.nonTrainableWeights)}get stateful(){return this._stateful}resetStates(){if(!this.stateful)throw new Error("Cannot call the resetStates() method of a non-stateful Layer object.")}assertInputCompatibility(t){const e=nt(t);if(this.inputSpec==null||this.inputSpec.length===0)return;const s=nt(this.inputSpec);if(e.length!==s.length)throw new I(`Layer ${this.name} expects ${s.length} inputs, but it received ${e.length} input tensors. Input received: ${t}`);for(let r=0;r<e.length;r++){const o=e[r],i=s[r];if(i==null)continue;const a=o.rank;if(i.ndim!=null&&a!==i.ndim)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected ndim=${i.ndim}, found ndim=${a}`);if(i.maxNDim!=null&&a>i.maxNDim)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected max_ndim=${i.maxNDim}, found ndim=${a}`);if(i.minNDim!=null&&a<i.minNDim)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected min_ndim=${i.minNDim}, found ndim=${a}.`);if(i.dtype!=null&&o.dtype!==i.dtype)throw new I(`Input ${r} is incompatible with layer ${this.name} : expected dtype=${i.dtype}, found dtype=${o.dtype}.`);if(i.axes){const l=o.shape;for(const u in i.axes){const c=Number(u),h=i.axes[u],f=c>=0?l[c]:l[l.length+c];if(h!=null&&[h,null].indexOf(f)===-1)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}.`)}}if(i.shape!=null)for(let l=0;l<i.shape.length;++l){const u=i.shape[l],c=o.shape[l];if(u!=null&&c!=null&&u!==c)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected shape=${i.shape}, found shape=${o.shape}.`)}}}call(t,e){return t}invokeCallHook(t,e){this._callHook!=null&&this._callHook(t,e)}setCallHook(t){this._callHook=t}clearCallHook(){this._callHook=null}apply(t,e){e=e||{},this.assertNotDisposed();const s=nt(t),r=dw(t),o=pw(t);if(r===o)throw new I("Arguments to apply() must be all SymbolicTensors or all Tensors");return gr(this.name,()=>{if(!this.built){this.assertInputCompatibility(t);const i=[];for(const a of nt(t))i.push(a.shape);this.build(Ut(i)),this.built=!0,this.initialWeights&&this.setWeights(this.initialWeights),this._refCount===null&&o&&(this._refCount=1)}if(this.assertInputCompatibility(t),o){let i=this.call(t,e);this.supportsMasking&&this.setMaskMetadata(t,i);const a=nt(i),l=[];for(let u of a)s.indexOf(u)!==-1&&(u=u.clone()),l.push(u);if(i=Ut(l),this.activityRegularizer!=null)throw new Z("Layer invocation in the presence of activity regularizer(s) is not supported yet.");return i}else{const i=hw(t),a=this.computeOutputShape(i);let l;const u=fw(t);if(this.warnOnIncompatibleInputShape(Array.isArray(t)?i[0]:i),a!=null&&a.length>0&&Array.isArray(a[0])?l=a.map((c,h)=>new bn(u,c,this,nt(t),e,this.name,h)):l=new bn(u,a,this,nt(t),e,this.name),this.addInboundNode(t,l,null,null,i,a,e),this._refCount++,this.activityRegularizer!=null)throw new Z("Layer invocation in the presence of activity regularizer(s) is not supported yet.");return l}})}warnOnIncompatibleInputShape(t){if(this.batchInputShape!=null)if(t.length!==this.batchInputShape.length)console.warn(`The rank of the input tensor provided (shape: ${JSON.stringify(t)}) does not match that of the batchInputShape (${JSON.stringify(this.batchInputShape)}) of the layer ${this.name}`);else{let e=!1;this.batchInputShape.forEach((s,r)=>{s!=null&&t[r]!=null&&t[r]!==s&&(e=!0)}),e&&console.warn(`The shape of the input tensor (${JSON.stringify(t)}) does not match the expectation of layer ${this.name}: ${JSON.stringify(this.batchInputShape)}`)}}get outputShape(){if(this.inboundNodes==null||this.inboundNodes.length===0)throw new je(`The layer ${this.name} has never been called and thus has no defined output shape.`);const t=[];for(const e of this.inboundNodes){const s=JSON.stringify(e.outputShapes);t.indexOf(s)===-1&&t.push(s)}if(t.length===1){const e=this.inboundNodes[0].outputShapes;return Array.isArray(e)&&Array.isArray(e[0])&&e.length===1?e[0]:e}else throw new je(`The layer ${this.name} has multiple inbound nodes with different output shapes. Hence the notion of "output shape" is ill-defined for the layer.`)}countParams(){if(!this.built)throw new He(`You tried to call countParams() on ${this.name}, but the layer is not built yet. Build it first by calling build(batchInputShape).`);return Ir(this.weights)}build(t){this.built=!0}getWeights(t=!1){return qu(t?this.trainableWeights:this.weights)}setWeights(t){k(()=>{const e=this.weights;if(e.length!==t.length)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}...`);if(e.length===0)return;const s=[],r=qu(e);for(let o=0;o<r.length;++o){const i=r[o],a=e[o],l=t[o];if(!Zt(i.shape,l.shape))throw new I(`Layer weight shape ${i.shape} not compatible with provided weight shape ${l.shape}`);s.push([a,l])}ju(s)})}addWeight(t,e,s,r,o,i,a,l){if(this._addedWeightNames.indexOf(t)!==-1)throw new I(`Duplicate weight name ${t} for layer ${this.name}`);this._addedWeightNames.push(t),s==null&&(s="float32"),this.fastWeightInitDuringBuild&&(r=l!=null?l():$s("zeros"));const u=r.apply(e,s),c=new aw(u,s,t,i,a);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}setFastWeightInitDuringBuild(t){this.fastWeightInitDuringBuild=t}addLoss(t){t==null||Array.isArray(t)&&t.length===0||(t=nt(t),this._losses!==void 0&&this._losses!==null&&this.losses.push(...t))}computeOutputShape(t){return t}computeMask(t,e){if(!this.supportsMasking){if(e!=null)if(Array.isArray(e))e.forEach(s=>{if(s!=null)throw new TypeError(`Layer ${this.name} does not support masking, but was passed an inputMask.`)});else throw new TypeError(`Layer ${this.name} does not support masking, but was passed an inputMask.`);return null}return e}setMaskMetadata(t,e,s){if(!this.supportsMasking)return;const r=this.computeMask(t,s),o=nt(e),i=nt(r);if(o.length!==i.length)throw new Error(`${this.name} outputs ${o.length} tensors but ${o.length} masks for those tensors`);for(let a=0;a<o.length;a++)o[a].kerasMask=i[a]}addInboundNode(t,e,s,r,o,i,a=null){const l=nt(t);e=nt(e),s=nt(s),r=nt(r),o=vr(o),i=vr(i);const u=[],c=[],h=[];for(const f of l)u.push(f.sourceLayer),c.push(f.nodeIndex),h.push(f.tensorIndex);new xi({outboundLayer:this,inboundLayers:u,nodeIndices:c,tensorIndices:h,inputTensors:l,outputTensors:e,inputMasks:s,outputMasks:r,inputShapes:o,outputShapes:i},a);for(let f=0;f<e.length;f++)e[f].sourceLayer=this,e[f].nodeIndex=this.inboundNodes.length-1,e[f].tensorIndex=f}getConfig(){const t={name:this.name,trainable:this.trainable};return this.batchInputShape!=null&&(t.batchInputShape=this.batchInputShape),this.dtype!=null&&(t.dtype=this.dtype),t}disposeWeights(){return this.weights.forEach(t=>t.dispose()),this.weights.length}assertNotDisposed(){if(this._refCount===0)throw new Error(`Layer '${this.name}' is already disposed.`)}dispose(){if(!this.built)throw new Error(`Cannot dispose Layer ${this.name} because it has not been built yet.`);if(this._refCount===null)throw new Error(`Cannot dispose Layer ${this.name} because it has not been used yet.`);this.assertNotDisposed();let t=0;return--this._refCount===0&&(t=this.disposeWeights()),{refCountAfterDispose:this._refCount,numDisposedVariables:t}}}function hw(n){n=nt(n);const t=[];for(const e of n)t.push(e.shape);return Ut(t)}function fw(n){return"float32"}function dw(n){let t=!0;for(const e of nt(n))if(!(e instanceof bn)){t=!1;break}return t}function pw(n){let t=!0;for(const e of nt(n))if(e instanceof bn){t=!1;break}return t}/**
|
|
2871
|
+
*/class Ce{constructor(t){this.dtype=t.dtype,this.shape=t.shape,t.shape!=null?this.ndim=t.shape.length:this.ndim=t.ndim,this.maxNDim=t.maxNDim,this.minNDim=t.minNDim,this.axes=t.axes||{}}}class yn{constructor(t,e,s,r,o,i,a){this.dtype=t,this.shape=e,this.sourceLayer=s,this.inputs=r,this.callArgs=o,this.outputTensorIndex=a,this.id=Ru(),i!=null&&(this.originalName=ou(i),this.name=iu(this.originalName)),this.rank=e.length}}let uw=0;class xi{constructor(t,e){this.callArgs=e,this.id=uw++,this.outboundLayer=t.outboundLayer,this.inboundLayers=t.inboundLayers,this.nodeIndices=t.nodeIndices,this.tensorIndices=t.tensorIndices,this.inputTensors=t.inputTensors,this.outputTensors=t.outputTensors,this.inputMasks=t.inputMasks,this.outputMasks=t.outputMasks,this.inputShapes=t.inputShapes,this.outputShapes=t.outputShapes;for(const s of t.inboundLayers)s!=null&&s.outboundNodes.push(this);t.outboundLayer.inboundNodes.push(this)}getConfig(){const t=[];for(const e of this.inboundLayers)e!=null?t.push(e.name):t.push(null);return{outboundLayer:this.outboundLayer?this.outboundLayer.name:null,inboundLayers:t,nodeIndices:this.nodeIndices,tensorIndices:this.tensorIndices}}}let cw=0;class pe extends Gn{constructor(t={}){super(),this._callHook=null,this._addedWeightNames=[],this._stateful=!1,this.id=cw++,this.activityRegularizer=null,this.inputSpec=null,this.supportsMasking=!1,this._trainableWeights=[],this._nonTrainableWeights=[],this._losses=[],this._updates=[],this._built=!1,this.inboundNodes=[],this.outboundNodes=[];let e=t.name;if(!e){const s=this.getClassName();e=Le(s)+"_"+di(s)}if(this.name=e,this.trainable_=t.trainable==null?!0:t.trainable,t.inputShape!=null||t.batchInputShape!=null){let s;if(t.batchInputShape!=null)s=t.batchInputShape;else if(t.inputShape!=null){let o=null;t.batchSize!=null&&(o=t.batchSize),s=[o].concat(t.inputShape)}this.batchInputShape=s;let r=t.dtype;r==null&&(r=t.inputDType),r==null&&(r="float32"),this.dtype=r}t.weights!=null?this.initialWeights=t.weights:this.initialWeights=null,this._refCount=null,this.fastWeightInitDuringBuild=!1}static nodeKey(t,e){return t.name+"_ib-"+e.toString()}getNodeAtIndex(t,e){if(this.inboundNodes.length===0)throw new He(`The layer has never been called and thus has no defined ${e}.`);if(this.inboundNodes.length<=t)throw new I(`Asked to get ${e} at node ${t}, but the layer has only ${this.inboundNodes.length} inbound nodes.`);return this.inboundNodes[t]}getInputAt(t){return Ut(this.getNodeAtIndex(t,"input").inputTensors)}getOutputAt(t){return Ut(this.getNodeAtIndex(t,"output").outputTensors)}get input(){if(this.inboundNodes.length>1)throw new je(`Layer ${this.name} has multiple inbound nodes, hence the notion of "layer input" is ill-defined. Use \`getInputAt(nodeIndex)\` instead.`);if(this.inboundNodes.length===0)throw new je(`Layer ${this.name} is not connected, no input to return.`);return Ut(this.getNodeAtIndex(0,"input").inputTensors)}get output(){if(this.inboundNodes.length===0)throw new je(`Layer ${this.name} has no inbound nodes.`);if(this.inboundNodes.length>1)throw new je(`Layer ${this.name} has multiple inbound nodes, hence the notion of "layer output" is ill-defined. Use \`getOutputAt(nodeIndex)\` instead.`);return Ut(this.getNodeAtIndex(0,"output").outputTensors)}get losses(){return this._losses}calculateLosses(){return this.losses.map(t=>t())}get updates(){return this._updates}get built(){return this._built}set built(t){this._built=t}get trainable(){return this.trainable_}set trainable(t){this._trainableWeights.forEach(e=>e.trainable=t),this.trainable_=t}get trainableWeights(){return this.trainable_?this._trainableWeights.filter(t=>t.trainable):[]}set trainableWeights(t){this._trainableWeights=t}get nonTrainableWeights(){return this.trainable?this._trainableWeights.filter(t=>!t.trainable).concat(this._nonTrainableWeights):this._trainableWeights.concat(this._nonTrainableWeights)}set nonTrainableWeights(t){this._nonTrainableWeights=t}get weights(){return this.trainableWeights.concat(this.nonTrainableWeights)}get stateful(){return this._stateful}resetStates(){if(!this.stateful)throw new Error("Cannot call the resetStates() method of a non-stateful Layer object.")}assertInputCompatibility(t){const e=nt(t);if(this.inputSpec==null||this.inputSpec.length===0)return;const s=nt(this.inputSpec);if(e.length!==s.length)throw new I(`Layer ${this.name} expects ${s.length} inputs, but it received ${e.length} input tensors. Input received: ${t}`);for(let r=0;r<e.length;r++){const o=e[r],i=s[r];if(i==null)continue;const a=o.rank;if(i.ndim!=null&&a!==i.ndim)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected ndim=${i.ndim}, found ndim=${a}`);if(i.maxNDim!=null&&a>i.maxNDim)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected max_ndim=${i.maxNDim}, found ndim=${a}`);if(i.minNDim!=null&&a<i.minNDim)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected min_ndim=${i.minNDim}, found ndim=${a}.`);if(i.dtype!=null&&o.dtype!==i.dtype)throw new I(`Input ${r} is incompatible with layer ${this.name} : expected dtype=${i.dtype}, found dtype=${o.dtype}.`);if(i.axes){const l=o.shape;for(const u in i.axes){const c=Number(u),h=i.axes[u],f=c>=0?l[c]:l[l.length+c];if(h!=null&&[h,null].indexOf(f)===-1)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}.`)}}if(i.shape!=null)for(let l=0;l<i.shape.length;++l){const u=i.shape[l],c=o.shape[l];if(u!=null&&c!=null&&u!==c)throw new I(`Input ${r} is incompatible with layer ${this.name}: expected shape=${i.shape}, found shape=${o.shape}.`)}}}call(t,e){return t}invokeCallHook(t,e){this._callHook!=null&&this._callHook(t,e)}setCallHook(t){this._callHook=t}clearCallHook(){this._callHook=null}apply(t,e){e=e||{},this.assertNotDisposed();const s=nt(t),r=dw(t),o=pw(t);if(r===o)throw new I("Arguments to apply() must be all SymbolicTensors or all Tensors");return br(this.name,()=>{if(!this.built){this.assertInputCompatibility(t);const i=[];for(const a of nt(t))i.push(a.shape);this.build(Ut(i)),this.built=!0,this.initialWeights&&this.setWeights(this.initialWeights),this._refCount===null&&o&&(this._refCount=1)}if(this.assertInputCompatibility(t),o){let i=this.call(t,e);this.supportsMasking&&this.setMaskMetadata(t,i);const a=nt(i),l=[];for(let u of a)s.indexOf(u)!==-1&&(u=u.clone()),l.push(u);if(i=Ut(l),this.activityRegularizer!=null)throw new Z("Layer invocation in the presence of activity regularizer(s) is not supported yet.");return i}else{const i=hw(t),a=this.computeOutputShape(i);let l;const u=fw(t);if(this.warnOnIncompatibleInputShape(Array.isArray(t)?i[0]:i),a!=null&&a.length>0&&Array.isArray(a[0])?l=a.map((c,h)=>new yn(u,c,this,nt(t),e,this.name,h)):l=new yn(u,a,this,nt(t),e,this.name),this.addInboundNode(t,l,null,null,i,a,e),this._refCount++,this.activityRegularizer!=null)throw new Z("Layer invocation in the presence of activity regularizer(s) is not supported yet.");return l}})}warnOnIncompatibleInputShape(t){if(this.batchInputShape!=null)if(t.length!==this.batchInputShape.length)console.warn(`The rank of the input tensor provided (shape: ${JSON.stringify(t)}) does not match that of the batchInputShape (${JSON.stringify(this.batchInputShape)}) of the layer ${this.name}`);else{let e=!1;this.batchInputShape.forEach((s,r)=>{s!=null&&t[r]!=null&&t[r]!==s&&(e=!0)}),e&&console.warn(`The shape of the input tensor (${JSON.stringify(t)}) does not match the expectation of layer ${this.name}: ${JSON.stringify(this.batchInputShape)}`)}}get outputShape(){if(this.inboundNodes==null||this.inboundNodes.length===0)throw new je(`The layer ${this.name} has never been called and thus has no defined output shape.`);const t=[];for(const e of this.inboundNodes){const s=JSON.stringify(e.outputShapes);t.indexOf(s)===-1&&t.push(s)}if(t.length===1){const e=this.inboundNodes[0].outputShapes;return Array.isArray(e)&&Array.isArray(e[0])&&e.length===1?e[0]:e}else throw new je(`The layer ${this.name} has multiple inbound nodes with different output shapes. Hence the notion of "output shape" is ill-defined for the layer.`)}countParams(){if(!this.built)throw new He(`You tried to call countParams() on ${this.name}, but the layer is not built yet. Build it first by calling build(batchInputShape).`);return Ar(this.weights)}build(t){this.built=!0}getWeights(t=!1){return qu(t?this.trainableWeights:this.weights)}setWeights(t){k(()=>{const e=this.weights;if(e.length!==t.length)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}...`);if(e.length===0)return;const s=[],r=qu(e);for(let o=0;o<r.length;++o){const i=r[o],a=e[o],l=t[o];if(!Zt(i.shape,l.shape))throw new I(`Layer weight shape ${i.shape} not compatible with provided weight shape ${l.shape}`);s.push([a,l])}ju(s)})}addWeight(t,e,s,r,o,i,a,l){if(this._addedWeightNames.indexOf(t)!==-1)throw new I(`Duplicate weight name ${t} for layer ${this.name}`);this._addedWeightNames.push(t),s==null&&(s="float32"),this.fastWeightInitDuringBuild&&(r=l!=null?l():vs("zeros"));const u=r.apply(e,s),c=new aw(u,s,t,i,a);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}setFastWeightInitDuringBuild(t){this.fastWeightInitDuringBuild=t}addLoss(t){t==null||Array.isArray(t)&&t.length===0||(t=nt(t),this._losses!==void 0&&this._losses!==null&&this.losses.push(...t))}computeOutputShape(t){return t}computeMask(t,e){if(!this.supportsMasking){if(e!=null)if(Array.isArray(e))e.forEach(s=>{if(s!=null)throw new TypeError(`Layer ${this.name} does not support masking, but was passed an inputMask.`)});else throw new TypeError(`Layer ${this.name} does not support masking, but was passed an inputMask.`);return null}return e}setMaskMetadata(t,e,s){if(!this.supportsMasking)return;const r=this.computeMask(t,s),o=nt(e),i=nt(r);if(o.length!==i.length)throw new Error(`${this.name} outputs ${o.length} tensors but ${o.length} masks for those tensors`);for(let a=0;a<o.length;a++)o[a].kerasMask=i[a]}addInboundNode(t,e,s,r,o,i,a=null){const l=nt(t);e=nt(e),s=nt(s),r=nt(r),o=Ir(o),i=Ir(i);const u=[],c=[],h=[];for(const f of l)u.push(f.sourceLayer),c.push(f.nodeIndex),h.push(f.tensorIndex);new xi({outboundLayer:this,inboundLayers:u,nodeIndices:c,tensorIndices:h,inputTensors:l,outputTensors:e,inputMasks:s,outputMasks:r,inputShapes:o,outputShapes:i},a);for(let f=0;f<e.length;f++)e[f].sourceLayer=this,e[f].nodeIndex=this.inboundNodes.length-1,e[f].tensorIndex=f}getConfig(){const t={name:this.name,trainable:this.trainable};return this.batchInputShape!=null&&(t.batchInputShape=this.batchInputShape),this.dtype!=null&&(t.dtype=this.dtype),t}disposeWeights(){return this.weights.forEach(t=>t.dispose()),this.weights.length}assertNotDisposed(){if(this._refCount===0)throw new Error(`Layer '${this.name}' is already disposed.`)}dispose(){if(!this.built)throw new Error(`Cannot dispose Layer ${this.name} because it has not been built yet.`);if(this._refCount===null)throw new Error(`Cannot dispose Layer ${this.name} because it has not been used yet.`);this.assertNotDisposed();let t=0;return--this._refCount===0&&(t=this.disposeWeights()),{refCountAfterDispose:this._refCount,numDisposedVariables:t}}}function hw(n){n=nt(n);const t=[];for(const e of n)t.push(e.shape);return Ut(t)}function fw(n){return"float32"}function dw(n){let t=!0;for(const e of nt(n))if(!(e instanceof yn)){t=!1;break}return t}function pw(n){let t=!0;for(const e of nt(n))if(e instanceof yn){t=!1;break}return t}/**
|
|
2872
2872
|
* @license
|
|
2873
2873
|
* Copyright 2018 Google LLC
|
|
2874
2874
|
*
|
|
@@ -2876,7 +2876,7 @@
|
|
|
2876
2876
|
* license that can be found in the LICENSE file or at
|
|
2877
2877
|
* https://opensource.org/licenses/MIT.
|
|
2878
2878
|
* =============================================================================
|
|
2879
|
-
*/function mw(n){if(n!=null&&typeof n!="object")throw new Error(`Argument to L1L2 regularizer's constructor is expected to be an object, but received: ${n}`)}class Hu extends Gn{}class Ku extends Hu{constructor(t){super(),mw(t),this.l1=t==null||t.l1==null?.01:t.l1,this.l2=t==null||t.l2==null?.01:t.l2,this.hasL1=this.l1!==0,this.hasL2=this.l2!==0}apply(t){return k(()=>{let e=Un([1]);return this.hasL1&&(e=O(e,et(N(this.l1,Lt(t))))),this.hasL2&&(e=O(e,et(N(this.l2,
|
|
2879
|
+
*/function mw(n){if(n!=null&&typeof n!="object")throw new Error(`Argument to L1L2 regularizer's constructor is expected to be an object, but received: ${n}`)}class Hu extends Gn{}class Ku extends Hu{constructor(t){super(),mw(t),this.l1=t==null||t.l1==null?.01:t.l1,this.l2=t==null||t.l2==null?.01:t.l2,this.hasL1=this.l1!==0,this.hasL2=this.l2!==0}apply(t){return k(()=>{let e=Un([1]);return this.hasL1&&(e=O(e,et(N(this.l1,Lt(t))))),this.hasL2&&(e=O(e,et(N(this.l2,xs(t))))),L(e,[])})}getConfig(){return{l1:this.l1,l2:this.l2}}static fromConfig(t,e){return new t({l1:e.l1,l2:e.l2})}}Ku.className="L1L2",M(Ku);const Yu={l1l2:"L1L2"};function Is(n){return oi(n)}function Xu(n,t={}){return bs(n,se.getMap().classNameMap,t,"regularizer")}function As(n){if(n==null)return null;if(typeof n=="string"){const e={className:n in Yu?Yu[n]:n,config:{}};return Xu(e)}else return n instanceof Hu?n:Xu(n)}/**
|
|
2880
2880
|
* @license
|
|
2881
2881
|
* Copyright 2018 Google LLC
|
|
2882
2882
|
*
|
|
@@ -2884,7 +2884,7 @@
|
|
|
2884
2884
|
* license that can be found in the LICENSE file or at
|
|
2885
2885
|
* https://opensource.org/licenses/MIT.
|
|
2886
2886
|
* =============================================================================
|
|
2887
|
-
*/function Si(n,t,e){if(typeof n=="number")return
|
|
2887
|
+
*/function Si(n,t,e){if(typeof n=="number")return mr(n,t);if(n.length!==t)throw new I(`The ${e} argument must be an integer or tuple of ${t} integers. Received: ${n.length} elements.`);for(let s=0;s<t;++s){const r=n[s];if(!qy(r))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}`)}return n}function wn(n,t,e,s,r=1){if(n==null)return n;const o=t+(t-1)*(r-1);let i;return e==="same"?i=n:i=n-o+1,Math.floor((i+s-1)/s)}function ke(n,t,e,s){if(n==null)return null;if(s==="valid")n=n*t+lu([e-t,0]);else if(s==="same")n=n*t;else throw new I(`Unsupport padding mode: ${s}.`);return n}/**
|
|
2888
2888
|
* @license
|
|
2889
2889
|
* Copyright 2018 Google LLC
|
|
2890
2890
|
*
|
|
@@ -2892,7 +2892,7 @@
|
|
|
2892
2892
|
* license that can be found in the LICENSE file or at
|
|
2893
2893
|
* https://opensource.org/licenses/MIT.
|
|
2894
2894
|
* =============================================================================
|
|
2895
|
-
*/function Ju(n,t){return k(()=>(mt(t),t==="channelsFirst"?pt(n,[0,2,3,1]):n))}function Zu(n,t){return k(()=>(mt(t),t==="channelsFirst"?pt(n,[0,2,3,4,1]):n))}function gw(n,t,e,s=1,r="valid",o,i=1){return k(()=>{if(o==null&&(o=Hn()),mt(o),n.shape.length!==3)throw new I(`The input of a conv1dWithBias operation should be 3, but is ${n.shape.length} instead.`);if(t.shape.length!==3)throw new I(`The kernel for a conv1dWithBias operation should be 3, but is ${t.shape.length} instead`);if(e!=null&&e.shape.length!==1)throw new I(`The bias for a conv1dWithBias operation should be 1, but is ${e.shape.length} instead`);if(o==="channelsFirst"&&(n=pt(n,[0,2,1])),r==="causal")throw new Z("The support for CAUSAL padding mode in conv1dWithBias is not implemented yet.");let a=Hm(n,t,s,r==="same"?"same":"valid","NWC",i);return e!=null&&(a=xs(a,e)),a})}function Qu(n,t,e,s=[1,1],r="valid",o,i,a=null){return k(()=>{if(o==null&&(o=Hn()),mt(o),n.rank!==3&&n.rank!==4)throw new I(`conv2dWithBiasActivation expects input to be of rank 3 or 4, but received ${n.rank}.`);if(t.rank!==3&&t.rank!==4)throw new I(`conv2dWithBiasActivation expects kernel to be of rank 3 or 4, but received ${n.rank}.`);let l=Ju(n,o);if(r==="causal")throw new Z("The support for CAUSAL padding mode in conv1dWithBias is not implemented yet.");return l=rb({x:l,filter:t,strides:s,pad:r==="same"?"same":"valid",dilations:i,dataFormat:"NHWC",bias:e,activation:a}),o==="channelsFirst"&&(l=pt(l,[0,3,1,2])),l})}function bw(n,t,e,s=[1,1,1],r="valid",o,i){return k(()=>{if(o==null&&(o=Hn()),mt(o),n.rank!==4&&n.rank!==5)throw new I(`conv3dWithBias expects input to be of rank 4 or 5, but received ${n.rank}.`);if(t.rank!==4&&t.rank!==5)throw new I(`conv3dWithBias expects kernel to be of rank 4 or 5, but received ${n.rank}.`);let a=Zu(n,o);if(r==="causal")throw new Z("The support for CAUSAL padding mode in conv3dWithBias is not implemented yet.");return a=Zm(a,t,s,r==="same"?"same":"valid","NDHWC",i),e!=null&&(a=xs(a,e)),o==="channelsFirst"&&(a=pt(a,[0,4,1,2,3])),a})}class $i extends pe{constructor(t,e){if(super(e),this.bias=null,this.DEFAULT_KERNEL_INITIALIZER="glorotNormal",this.DEFAULT_BIAS_INITIALIZER="zeros",$i.verifyArgs(e),this.rank=t,Me(this.rank,"rank"),this.rank!==1&&this.rank!==2&&this.rank!==3)throw new Z(`Convolution layer for rank other than 1, 2, or 3 (${this.rank}) is not implemented yet.`);if(this.kernelSize=Si(e.kernelSize,t,"kernelSize"),this.strides=Si(e.strides==null?1:e.strides,t,"strides"),this.padding=e.padding==null?"valid":e.padding,oe(this.padding),this.dataFormat=e.dataFormat==null?"channelsLast":e.dataFormat,mt(this.dataFormat),this.activation=tw(e.activation),this.useBias=e.useBias==null?!0:e.useBias,this.biasInitializer=$s(e.biasInitializer||this.DEFAULT_BIAS_INITIALIZER),this.biasConstraint=xr(e.biasConstraint),this.biasRegularizer=Is(e.biasRegularizer),this.activityRegularizer=Is(e.activityRegularizer),this.dilationRate=Si(e.dilationRate==null?1:e.dilationRate,t,"dilationRate"),this.rank===1&&Array.isArray(this.dilationRate)&&this.dilationRate.length!==1)throw new I(`dilationRate must be a number or an array of a single number for 1D convolution, but received ${JSON.stringify(this.dilationRate)}`);if(this.rank===2){if(typeof this.dilationRate=="number")this.dilationRate=[this.dilationRate,this.dilationRate];else if(this.dilationRate.length!==2)throw new I(`dilationRate must be a number or array of two numbers for 2D convolution, but received ${JSON.stringify(this.dilationRate)}`)}else if(this.rank===3){if(typeof this.dilationRate=="number")this.dilationRate=[this.dilationRate,this.dilationRate,this.dilationRate];else if(this.dilationRate.length!==3)throw new I(`dilationRate must be a number or array of three numbers for 3D convolution, but received ${JSON.stringify(this.dilationRate)}`)}}static verifyArgs(t){if(Ae("kernelSize"in t,"required key 'kernelSize' not in config"),typeof t.kernelSize!="number"&&!ai(t.kernelSize,"number",1,3))throw new I(`BaseConv expects config.kernelSize to be number or number[] with length 1, 2, or 3, but received ${JSON.stringify(t.kernelSize)}.`)}getConfig(){const t={kernelSize:this.kernelSize,strides:this.strides,padding:this.padding,dataFormat:this.dataFormat,dilationRate:this.dilationRate,activation:Qy(this.activation),useBias:this.useBias,biasInitializer:$r(this.biasInitializer),biasRegularizer:vs(this.biasRegularizer),activityRegularizer:vs(this.activityRegularizer),biasConstraint:wr(this.biasConstraint)},e=super.getConfig();return Object.assign(t,e),t}}class Kn extends $i{constructor(t,e){super(t,e),this.kernel=null,Kn.verifyArgs(e),this.filters=e.filters,Me(this.filters,"filters"),this.kernelInitializer=$s(e.kernelInitializer||this.DEFAULT_KERNEL_INITIALIZER),this.kernelConstraint=xr(e.kernelConstraint),this.kernelRegularizer=Is(e.kernelRegularizer)}build(t){t=de(t);const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null)throw new I(`The channel dimension of the input should be defined. Found ${t[e]}`);const s=t[e],r=this.kernelSize.concat([s,this.filters]);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}call(t,e){return k(()=>{t=Gt(t);let s;const r=this.bias==null?null:this.bias.read(),o=Uy(this.activation.getClassName());if(o!=null&&this.rank===2)s=Qu(t,this.kernel.read(),r,this.strides,this.padding,this.dataFormat,this.dilationRate,o);else{if(this.rank===1)s=gw(t,this.kernel.read(),r,this.strides[0],this.padding,this.dataFormat,this.dilationRate[0]);else if(this.rank===2)s=Qu(t,this.kernel.read(),r,this.strides,this.padding,this.dataFormat,this.dilationRate);else if(this.rank===3)s=bw(t,this.kernel.read(),r,this.strides,this.padding,this.dataFormat,this.dilationRate);else throw new Z("convolutions greater than 3D are not implemented yet.");this.activation!=null&&(s=this.activation.apply(s))}return s})}computeOutputShape(t){t=de(t);const e=[],s=this.dataFormat==="channelsLast"?t.slice(1,t.length-1):t.slice(2);for(let o=0;o<s.length;++o){const i=yn(s[o],this.kernelSize[o],this.padding,this.strides[o],typeof this.dilationRate=="number"?this.dilationRate:this.dilationRate[o]);e.push(i)}let r=[t[0]];return this.dataFormat==="channelsLast"?(r=r.concat(e),r.push(this.filters)):(r.push(this.filters),r=r.concat(e)),r}getConfig(){const t={filters:this.filters,kernelInitializer:$r(this.kernelInitializer),kernelRegularizer:vs(this.kernelRegularizer),kernelConstraint:wr(this.kernelConstraint)},e=super.getConfig();return Object.assign(t,e),t}static verifyArgs(t){if(!("filters"in t)||typeof t.filters!="number"||t.filters<1)throw new I(`Convolution layer expected config.filters to be a 'number' > 0 but got ${JSON.stringify(t.filters)}`)}}class Yn extends Kn{constructor(t){super(2,t),Yn.verifyArgs(t)}getConfig(){const t=super.getConfig();return delete t.rank,t}static verifyArgs(t){if(typeof t.kernelSize!="number"&&!ai(t.kernelSize,"number",1,2))throw new I(`Conv2D expects config.kernelSize to be number or number[] with length 1 or 2, but received ${JSON.stringify(t.kernelSize)}.`)}}Yn.className="Conv2D",M(Yn);class As extends Kn{constructor(t){super(3,t),As.verifyArgs(t)}getConfig(){const t=super.getConfig();return delete t.rank,t}static verifyArgs(t){if(typeof t.kernelSize!="number"&&!(Array.isArray(t.kernelSize)&&(t.kernelSize.length===1||t.kernelSize.length===3)))throw new I(`Conv3D expects config.kernelSize to be number or [number, number, number], but received ${JSON.stringify(t.kernelSize)}.`)}}As.className="Conv3D",M(As);class tc extends Yn{constructor(t){if(super(t),this.inputSpec=[new Ce({ndim:4})],this.padding!=="same"&&this.padding!=="valid")throw new I(`Conv2DTranspose currently supports only padding modes 'same' and 'valid', but received padding mode ${this.padding}`)}build(t){if(t=de(t),t.length!==4)throw new I("Input should have rank 4; Received input shape: "+JSON.stringify(t));const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null)throw new I("The channel dimension of the inputs should be defined. Found `None`.");const s=t[e],r=this.kernelSize.concat([this.filters,s]);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}call(t,e){return k(()=>{let s=Gt(t);if(s.shape.length!==4)throw new I(`Conv2DTranspose.call() expects input tensor to be rank-4, but received a tensor of rank-${s.shape.length}`);const r=s.shape,o=r[0];let i,a;this.dataFormat==="channelsFirst"?(i=2,a=3):(i=1,a=2);const l=r[i],u=r[a],c=this.kernelSize[0],h=this.kernelSize[1],f=this.strides[0],d=this.strides[1],p=ke(l,f,c,this.padding),g=ke(u,d,h,this.padding),m=[o,p,g,this.filters];this.dataFormat!=="channelsLast"&&(s=pt(s,[0,2,3,1]));let b=Xm(s,this.kernel.read(),m,this.strides,this.padding);return this.dataFormat!=="channelsLast"&&(b=pt(b,[0,3,1,2])),this.bias!=null&&(b=xs(b,this.bias.read(),this.dataFormat)),this.activation!=null&&(b=this.activation.apply(b)),b})}computeOutputShape(t){t=de(t);const e=t.slice();let s,r,o;this.dataFormat==="channelsFirst"?(s=1,r=2,o=3):(s=3,r=1,o=2);const i=this.kernelSize[0],a=this.kernelSize[1],l=this.strides[0],u=this.strides[1];return e[s]=this.filters,e[r]=ke(e[r],l,i,this.padding),e[o]=ke(e[o],u,a,this.padding),e}getConfig(){const t=super.getConfig();return delete t.dilationRate,t}}tc.className="Conv2DTranspose",M(tc);class ec extends As{constructor(t){if(super(t),this.inputSpec=[new Ce({ndim:5})],this.padding!=="same"&&this.padding!=="valid")throw new I(`Conv3DTranspose currently supports only padding modes 'same' and 'valid', but received padding mode ${this.padding}`)}build(t){if(t=de(t),t.length!==5)throw new I("Input should have rank 5; Received input shape: "+JSON.stringify(t));const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null)throw new I("The channel dimension of the inputs should be defined. Found `None`.");const s=t[e],r=this.kernelSize.concat([this.filters,s]);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}call(t,e){return k(()=>{let s=Gt(t);if(s.shape.length!==5)throw new I(`Conv3DTranspose.call() expects input tensor to be rank-4, but received a tensor of rank-${s.shape.length}`);const r=s.shape,o=r[0];let i,a,l;this.dataFormat==="channelsFirst"?(l=2,i=3,a=4):(l=1,i=2,a=3);const u=r[l],c=r[i],h=r[a],f=this.kernelSize[0],d=this.kernelSize[1],p=this.kernelSize[2],g=this.strides[0],m=this.strides[1],b=this.strides[2],y=ke(u,g,f,this.padding),S=ke(c,m,d,this.padding),x=ke(h,b,p,this.padding),$=[o,y,S,x,this.filters];this.dataFormat!=="channelsLast"&&(s=pt(s,[0,2,3,4,1]));let E=ng(s,this.kernel.read(),$,this.strides,this.padding);return this.dataFormat!=="channelsLast"&&(E=pt(E,[0,4,1,2,3])),this.bias!==null&&(E=xs(E,this.bias.read(),this.dataFormat)),this.activation!==null&&(E=this.activation.apply(E)),E})}computeOutputShape(t){t=de(t);const e=t.slice();let s,r,o,i;this.dataFormat==="channelsFirst"?(s=1,r=2,o=3,i=4):(s=4,r=1,o=2,i=3);const a=this.kernelSize[0],l=this.kernelSize[1],u=this.kernelSize[2],c=this.strides[0],h=this.strides[1],f=this.strides[2];return e[s]=this.filters,e[r]=ke(e[r],c,a,this.padding),e[o]=ke(e[o],h,l,this.padding),e[i]=ke(e[i],f,u,this.padding),e}getConfig(){const t=super.getConfig();return delete t.dilationRate,t}}ec.className="Conv3DTranspose",M(ec);class nc extends Kn{constructor(t,e){if(super(t,e),this.DEFAULT_DEPTHWISE_INITIALIZER="glorotUniform",this.DEFAULT_POINTWISE_INITIALIZER="glorotUniform",this.depthwiseKernel=null,this.pointwiseKernel=null,e.filters==null)throw new I("The `filters` configuration field is required by SeparableConv, but is unspecified.");if(e.kernelInitializer!=null||e.kernelRegularizer!=null||e.kernelConstraint!=null)throw new I("Fields kernelInitializer, kernelRegularizer and kernelConstraint are invalid for SeparableConv2D. Use depthwiseInitializer, depthwiseRegularizer, depthwiseConstraint, pointwiseInitializer, pointwiseRegularizer and pointwiseConstraint instead.");if(e.padding!=null&&e.padding!=="same"&&e.padding!=="valid")throw new I(`SeparableConv${this.rank}D supports only padding modes: 'same' and 'valid', but received ${JSON.stringify(e.padding)}`);this.depthMultiplier=e.depthMultiplier==null?1:e.depthMultiplier,this.depthwiseInitializer=$s(e.depthwiseInitializer||this.DEFAULT_DEPTHWISE_INITIALIZER),this.depthwiseRegularizer=Is(e.depthwiseRegularizer),this.depthwiseConstraint=xr(e.depthwiseConstraint),this.pointwiseInitializer=$s(e.depthwiseInitializer||this.DEFAULT_POINTWISE_INITIALIZER),this.pointwiseRegularizer=Is(e.pointwiseRegularizer),this.pointwiseConstraint=xr(e.pointwiseConstraint)}build(t){if(t=de(t),t.length<this.rank+2)throw new I(`Inputs to SeparableConv${this.rank}D should have rank ${this.rank+2}, but received input shape: ${JSON.stringify(t)}`);const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null||t[e]<0)throw new I(`The channel dimension of the inputs should be defined, but found ${JSON.stringify(t[e])}`);const s=t[e],r=this.kernelSize.concat([s,this.depthMultiplier]),o=[];for(let a=0;a<this.rank;++a)o.push(1);o.push(s*this.depthMultiplier,this.filters);const i=!0;this.depthwiseKernel=this.addWeight("depthwise_kernel",r,"float32",this.depthwiseInitializer,this.depthwiseRegularizer,i,this.depthwiseConstraint),this.pointwiseKernel=this.addWeight("pointwise_kernel",o,"float32",this.pointwiseInitializer,this.pointwiseRegularizer,i,this.pointwiseConstraint),this.useBias?this.bias=this.addWeight("bias",[this.filters],"float32",this.biasInitializer,this.biasRegularizer,i,this.biasConstraint):this.bias=null,this.inputSpec=[new Ce({ndim:this.rank+2,axes:{[e]:s}})],this.built=!0}call(t,e){return k(()=>{t=Gt(t);let s;if(this.rank===1)throw new Z("1D separable convolution is not implemented yet.");return this.rank===2&&(this.dataFormat==="channelsFirst"&&(t=pt(t,[0,2,3,1])),s=O0(t,this.depthwiseKernel.read(),this.pointwiseKernel.read(),this.strides,this.padding,this.dilationRate,"NHWC")),this.useBias&&(s=xs(s,this.bias.read(),this.dataFormat)),this.activation!=null&&(s=this.activation.apply(s)),this.dataFormat==="channelsFirst"&&(s=pt(s,[0,3,1,2])),s})}getConfig(){const t=super.getConfig();return delete t.rank,delete t.kernelInitializer,delete t.kernelRegularizer,delete t.kernelConstraint,t.depthwiseInitializer=$r(this.depthwiseInitializer),t.pointwiseInitializer=$r(this.pointwiseInitializer),t.depthwiseRegularizer=vs(this.depthwiseRegularizer),t.pointwiseRegularizer=vs(this.pointwiseRegularizer),t.depthwiseConstraint=wr(this.depthwiseConstraint),t.pointwiseConstraint=wr(this.pointwiseConstraint),t}}nc.className="SeparableConv";class sc extends nc{constructor(t){super(2,t)}}sc.className="SeparableConv2D",M(sc);class Ar extends Kn{constructor(t){super(1,t),Ar.verifyArgs(t),this.inputSpec=[{ndim:3}]}getConfig(){const t=super.getConfig();return delete t.rank,delete t.dataFormat,t}static verifyArgs(t){if(typeof t.kernelSize!="number"&&!ai(t.kernelSize,"number",1,1))throw new I(`Conv1D expects config.kernelSize to be number or number[] with length 1, but received ${JSON.stringify(t.kernelSize)}.`)}}Ar.className="Conv1D",M(Ar);class rc extends pe{constructor(t){super(t),typeof t.cropping=="number"?this.cropping=[[t.cropping,t.cropping],[t.cropping,t.cropping]]:typeof t.cropping[0]=="number"?this.cropping=[[t.cropping[0],t.cropping[0]],[t.cropping[1],t.cropping[1]]]:this.cropping=t.cropping,this.dataFormat=t.dataFormat===void 0?"channelsLast":t.dataFormat,this.inputSpec=[{ndim:4}]}computeOutputShape(t){return this.dataFormat==="channelsFirst"?[t[0],t[1],t[2]-this.cropping[0][0]-this.cropping[0][1],t[3]-this.cropping[1][0]-this.cropping[1][1]]:[t[0],t[1]-this.cropping[0][0]-this.cropping[0][1],t[2]-this.cropping[1][0]-this.cropping[1][1],t[3]]}call(t,e){return k(()=>{if(t=Gt(t),this.dataFormat==="channelsLast"){const s=yr(t,this.cropping[0][0],t.shape[1]-this.cropping[0][0]-this.cropping[0][1],2);return yr(s,this.cropping[1][0],t.shape[2]-this.cropping[1][1]-this.cropping[1][0],3)}else{const s=yr(t,this.cropping[0][0],t.shape[2]-this.cropping[0][0]-this.cropping[0][1],3);return yr(s,this.cropping[1][0],t.shape[3]-this.cropping[1][1]-this.cropping[1][0],4)}})}getConfig(){const t={cropping:this.cropping,dataFormat:this.dataFormat},e=super.getConfig();return Object.assign(t,e),t}}rc.className="Cropping2D",M(rc);class vi extends pe{constructor(t){super(t),this.DEFAULT_SIZE=[2,2],this.inputSpec=[{ndim:4}],this.size=t.size==null?this.DEFAULT_SIZE:t.size,this.dataFormat=t.dataFormat==null?"channelsLast":t.dataFormat,mt(this.dataFormat),this.interpolation=t.interpolation==null?"nearest":t.interpolation,Wy(this.interpolation)}computeOutputShape(t){if(this.dataFormat==="channelsFirst"){const e=t[2]==null?null:this.size[0]*t[2],s=t[3]==null?null:this.size[1]*t[3];return[t[0],t[1],e,s]}else{const e=t[1]==null?null:this.size[0]*t[1],s=t[2]==null?null:this.size[1]*t[2];return[t[0],e,s,t[3]]}}call(t,e){return k(()=>{let s=Gt(t);const r=s.shape;if(this.dataFormat==="channelsFirst"){s=pt(s,[0,2,3,1]);const o=this.size[0]*r[2],i=this.size[1]*r[3],a=this.interpolation==="nearest"?dr.resizeNearestNeighbor(s,[o,i]):dr.resizeBilinear(s,[o,i]);return pt(a,[0,3,1,2])}else{const o=this.size[0]*r[1],i=this.size[1]*r[2];return this.interpolation==="nearest"?dr.resizeNearestNeighbor(s,[o,i]):dr.resizeBilinear(s,[o,i])}})}getConfig(){const t={size:this.size,dataFormat:this.dataFormat,interpolation:this.interpolation},e=super.getConfig();return Object.assign(t,e),t}}vi.className="UpSampling2D",M(vi);/**
|
|
2895
|
+
*/function Ju(n,t){return k(()=>(mt(t),t==="channelsFirst"?pt(n,[0,2,3,1]):n))}function Zu(n,t){return k(()=>(mt(t),t==="channelsFirst"?pt(n,[0,2,3,4,1]):n))}function gw(n,t,e,s=1,r="valid",o,i=1){return k(()=>{if(o==null&&(o=Hn()),mt(o),n.shape.length!==3)throw new I(`The input of a conv1dWithBias operation should be 3, but is ${n.shape.length} instead.`);if(t.shape.length!==3)throw new I(`The kernel for a conv1dWithBias operation should be 3, but is ${t.shape.length} instead`);if(e!=null&&e.shape.length!==1)throw new I(`The bias for a conv1dWithBias operation should be 1, but is ${e.shape.length} instead`);if(o==="channelsFirst"&&(n=pt(n,[0,2,1])),r==="causal")throw new Z("The support for CAUSAL padding mode in conv1dWithBias is not implemented yet.");let a=Hm(n,t,s,r==="same"?"same":"valid","NWC",i);return e!=null&&(a=Ss(a,e)),a})}function Qu(n,t,e,s=[1,1],r="valid",o,i,a=null){return k(()=>{if(o==null&&(o=Hn()),mt(o),n.rank!==3&&n.rank!==4)throw new I(`conv2dWithBiasActivation expects input to be of rank 3 or 4, but received ${n.rank}.`);if(t.rank!==3&&t.rank!==4)throw new I(`conv2dWithBiasActivation expects kernel to be of rank 3 or 4, but received ${n.rank}.`);let l=Ju(n,o);if(r==="causal")throw new Z("The support for CAUSAL padding mode in conv1dWithBias is not implemented yet.");return l=rb({x:l,filter:t,strides:s,pad:r==="same"?"same":"valid",dilations:i,dataFormat:"NHWC",bias:e,activation:a}),o==="channelsFirst"&&(l=pt(l,[0,3,1,2])),l})}function bw(n,t,e,s=[1,1,1],r="valid",o,i){return k(()=>{if(o==null&&(o=Hn()),mt(o),n.rank!==4&&n.rank!==5)throw new I(`conv3dWithBias expects input to be of rank 4 or 5, but received ${n.rank}.`);if(t.rank!==4&&t.rank!==5)throw new I(`conv3dWithBias expects kernel to be of rank 4 or 5, but received ${n.rank}.`);let a=Zu(n,o);if(r==="causal")throw new Z("The support for CAUSAL padding mode in conv3dWithBias is not implemented yet.");return a=Zm(a,t,s,r==="same"?"same":"valid","NDHWC",i),e!=null&&(a=Ss(a,e)),o==="channelsFirst"&&(a=pt(a,[0,4,1,2,3])),a})}class $i extends pe{constructor(t,e){if(super(e),this.bias=null,this.DEFAULT_KERNEL_INITIALIZER="glorotNormal",this.DEFAULT_BIAS_INITIALIZER="zeros",$i.verifyArgs(e),this.rank=t,Me(this.rank,"rank"),this.rank!==1&&this.rank!==2&&this.rank!==3)throw new Z(`Convolution layer for rank other than 1, 2, or 3 (${this.rank}) is not implemented yet.`);if(this.kernelSize=Si(e.kernelSize,t,"kernelSize"),this.strides=Si(e.strides==null?1:e.strides,t,"strides"),this.padding=e.padding==null?"valid":e.padding,oe(this.padding),this.dataFormat=e.dataFormat==null?"channelsLast":e.dataFormat,mt(this.dataFormat),this.activation=tw(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=As(e.biasRegularizer),this.activityRegularizer=As(e.activityRegularizer),this.dilationRate=Si(e.dilationRate==null?1:e.dilationRate,t,"dilationRate"),this.rank===1&&Array.isArray(this.dilationRate)&&this.dilationRate.length!==1)throw new I(`dilationRate must be a number or an array of a single number for 1D convolution, but received ${JSON.stringify(this.dilationRate)}`);if(this.rank===2){if(typeof this.dilationRate=="number")this.dilationRate=[this.dilationRate,this.dilationRate];else if(this.dilationRate.length!==2)throw new I(`dilationRate must be a number or array of two numbers for 2D convolution, but received ${JSON.stringify(this.dilationRate)}`)}else if(this.rank===3){if(typeof this.dilationRate=="number")this.dilationRate=[this.dilationRate,this.dilationRate,this.dilationRate];else if(this.dilationRate.length!==3)throw new I(`dilationRate must be a number or array of three numbers for 3D convolution, but received ${JSON.stringify(this.dilationRate)}`)}}static verifyArgs(t){if(Ae("kernelSize"in t,"required key 'kernelSize' not in config"),typeof t.kernelSize!="number"&&!ai(t.kernelSize,"number",1,3))throw new I(`BaseConv expects config.kernelSize to be number or number[] with length 1, 2, or 3, but received ${JSON.stringify(t.kernelSize)}.`)}getConfig(){const t={kernelSize:this.kernelSize,strides:this.strides,padding:this.padding,dataFormat:this.dataFormat,dilationRate:this.dilationRate,activation:Qy(this.activation),useBias:this.useBias,biasInitializer:vr(this.biasInitializer),biasRegularizer:Is(this.biasRegularizer),activityRegularizer:Is(this.activityRegularizer),biasConstraint:xr(this.biasConstraint)},e=super.getConfig();return Object.assign(t,e),t}}class Kn extends $i{constructor(t,e){super(t,e),this.kernel=null,Kn.verifyArgs(e),this.filters=e.filters,Me(this.filters,"filters"),this.kernelInitializer=vs(e.kernelInitializer||this.DEFAULT_KERNEL_INITIALIZER),this.kernelConstraint=Sr(e.kernelConstraint),this.kernelRegularizer=As(e.kernelRegularizer)}build(t){t=de(t);const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null)throw new I(`The channel dimension of the input should be defined. Found ${t[e]}`);const s=t[e],r=this.kernelSize.concat([s,this.filters]);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}call(t,e){return k(()=>{t=Gt(t);let s;const r=this.bias==null?null:this.bias.read(),o=Uy(this.activation.getClassName());if(o!=null&&this.rank===2)s=Qu(t,this.kernel.read(),r,this.strides,this.padding,this.dataFormat,this.dilationRate,o);else{if(this.rank===1)s=gw(t,this.kernel.read(),r,this.strides[0],this.padding,this.dataFormat,this.dilationRate[0]);else if(this.rank===2)s=Qu(t,this.kernel.read(),r,this.strides,this.padding,this.dataFormat,this.dilationRate);else if(this.rank===3)s=bw(t,this.kernel.read(),r,this.strides,this.padding,this.dataFormat,this.dilationRate);else throw new Z("convolutions greater than 3D are not implemented yet.");this.activation!=null&&(s=this.activation.apply(s))}return s})}computeOutputShape(t){t=de(t);const e=[],s=this.dataFormat==="channelsLast"?t.slice(1,t.length-1):t.slice(2);for(let o=0;o<s.length;++o){const i=wn(s[o],this.kernelSize[o],this.padding,this.strides[o],typeof this.dilationRate=="number"?this.dilationRate:this.dilationRate[o]);e.push(i)}let r=[t[0]];return this.dataFormat==="channelsLast"?(r=r.concat(e),r.push(this.filters)):(r.push(this.filters),r=r.concat(e)),r}getConfig(){const t={filters:this.filters,kernelInitializer:vr(this.kernelInitializer),kernelRegularizer:Is(this.kernelRegularizer),kernelConstraint:xr(this.kernelConstraint)},e=super.getConfig();return Object.assign(t,e),t}static verifyArgs(t){if(!("filters"in t)||typeof t.filters!="number"||t.filters<1)throw new I(`Convolution layer expected config.filters to be a 'number' > 0 but got ${JSON.stringify(t.filters)}`)}}class Yn extends Kn{constructor(t){super(2,t),Yn.verifyArgs(t)}getConfig(){const t=super.getConfig();return delete t.rank,t}static verifyArgs(t){if(typeof t.kernelSize!="number"&&!ai(t.kernelSize,"number",1,2))throw new I(`Conv2D expects config.kernelSize to be number or number[] with length 1 or 2, but received ${JSON.stringify(t.kernelSize)}.`)}}Yn.className="Conv2D",M(Yn);class Es extends Kn{constructor(t){super(3,t),Es.verifyArgs(t)}getConfig(){const t=super.getConfig();return delete t.rank,t}static verifyArgs(t){if(typeof t.kernelSize!="number"&&!(Array.isArray(t.kernelSize)&&(t.kernelSize.length===1||t.kernelSize.length===3)))throw new I(`Conv3D expects config.kernelSize to be number or [number, number, number], but received ${JSON.stringify(t.kernelSize)}.`)}}Es.className="Conv3D",M(Es);class tc extends Yn{constructor(t){if(super(t),this.inputSpec=[new Ce({ndim:4})],this.padding!=="same"&&this.padding!=="valid")throw new I(`Conv2DTranspose currently supports only padding modes 'same' and 'valid', but received padding mode ${this.padding}`)}build(t){if(t=de(t),t.length!==4)throw new I("Input should have rank 4; Received input shape: "+JSON.stringify(t));const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null)throw new I("The channel dimension of the inputs should be defined. Found `None`.");const s=t[e],r=this.kernelSize.concat([this.filters,s]);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}call(t,e){return k(()=>{let s=Gt(t);if(s.shape.length!==4)throw new I(`Conv2DTranspose.call() expects input tensor to be rank-4, but received a tensor of rank-${s.shape.length}`);const r=s.shape,o=r[0];let i,a;this.dataFormat==="channelsFirst"?(i=2,a=3):(i=1,a=2);const l=r[i],u=r[a],c=this.kernelSize[0],h=this.kernelSize[1],f=this.strides[0],d=this.strides[1],p=ke(l,f,c,this.padding),g=ke(u,d,h,this.padding),m=[o,p,g,this.filters];this.dataFormat!=="channelsLast"&&(s=pt(s,[0,2,3,1]));let b=Xm(s,this.kernel.read(),m,this.strides,this.padding);return this.dataFormat!=="channelsLast"&&(b=pt(b,[0,3,1,2])),this.bias!=null&&(b=Ss(b,this.bias.read(),this.dataFormat)),this.activation!=null&&(b=this.activation.apply(b)),b})}computeOutputShape(t){t=de(t);const e=t.slice();let s,r,o;this.dataFormat==="channelsFirst"?(s=1,r=2,o=3):(s=3,r=1,o=2);const i=this.kernelSize[0],a=this.kernelSize[1],l=this.strides[0],u=this.strides[1];return e[s]=this.filters,e[r]=ke(e[r],l,i,this.padding),e[o]=ke(e[o],u,a,this.padding),e}getConfig(){const t=super.getConfig();return delete t.dilationRate,t}}tc.className="Conv2DTranspose",M(tc);class ec extends Es{constructor(t){if(super(t),this.inputSpec=[new Ce({ndim:5})],this.padding!=="same"&&this.padding!=="valid")throw new I(`Conv3DTranspose currently supports only padding modes 'same' and 'valid', but received padding mode ${this.padding}`)}build(t){if(t=de(t),t.length!==5)throw new I("Input should have rank 5; Received input shape: "+JSON.stringify(t));const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null)throw new I("The channel dimension of the inputs should be defined. Found `None`.");const s=t[e],r=this.kernelSize.concat([this.filters,s]);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}call(t,e){return k(()=>{let s=Gt(t);if(s.shape.length!==5)throw new I(`Conv3DTranspose.call() expects input tensor to be rank-4, but received a tensor of rank-${s.shape.length}`);const r=s.shape,o=r[0];let i,a,l;this.dataFormat==="channelsFirst"?(l=2,i=3,a=4):(l=1,i=2,a=3);const u=r[l],c=r[i],h=r[a],f=this.kernelSize[0],d=this.kernelSize[1],p=this.kernelSize[2],g=this.strides[0],m=this.strides[1],b=this.strides[2],y=ke(u,g,f,this.padding),S=ke(c,m,d,this.padding),x=ke(h,b,p,this.padding),$=[o,y,S,x,this.filters];this.dataFormat!=="channelsLast"&&(s=pt(s,[0,2,3,4,1]));let E=ng(s,this.kernel.read(),$,this.strides,this.padding);return this.dataFormat!=="channelsLast"&&(E=pt(E,[0,4,1,2,3])),this.bias!==null&&(E=Ss(E,this.bias.read(),this.dataFormat)),this.activation!==null&&(E=this.activation.apply(E)),E})}computeOutputShape(t){t=de(t);const e=t.slice();let s,r,o,i;this.dataFormat==="channelsFirst"?(s=1,r=2,o=3,i=4):(s=4,r=1,o=2,i=3);const a=this.kernelSize[0],l=this.kernelSize[1],u=this.kernelSize[2],c=this.strides[0],h=this.strides[1],f=this.strides[2];return e[s]=this.filters,e[r]=ke(e[r],c,a,this.padding),e[o]=ke(e[o],h,l,this.padding),e[i]=ke(e[i],f,u,this.padding),e}getConfig(){const t=super.getConfig();return delete t.dilationRate,t}}ec.className="Conv3DTranspose",M(ec);class nc extends Kn{constructor(t,e){if(super(t,e),this.DEFAULT_DEPTHWISE_INITIALIZER="glorotUniform",this.DEFAULT_POINTWISE_INITIALIZER="glorotUniform",this.depthwiseKernel=null,this.pointwiseKernel=null,e.filters==null)throw new I("The `filters` configuration field is required by SeparableConv, but is unspecified.");if(e.kernelInitializer!=null||e.kernelRegularizer!=null||e.kernelConstraint!=null)throw new I("Fields kernelInitializer, kernelRegularizer and kernelConstraint are invalid for SeparableConv2D. Use depthwiseInitializer, depthwiseRegularizer, depthwiseConstraint, pointwiseInitializer, pointwiseRegularizer and pointwiseConstraint instead.");if(e.padding!=null&&e.padding!=="same"&&e.padding!=="valid")throw new I(`SeparableConv${this.rank}D supports only padding modes: 'same' and 'valid', but received ${JSON.stringify(e.padding)}`);this.depthMultiplier=e.depthMultiplier==null?1:e.depthMultiplier,this.depthwiseInitializer=vs(e.depthwiseInitializer||this.DEFAULT_DEPTHWISE_INITIALIZER),this.depthwiseRegularizer=As(e.depthwiseRegularizer),this.depthwiseConstraint=Sr(e.depthwiseConstraint),this.pointwiseInitializer=vs(e.depthwiseInitializer||this.DEFAULT_POINTWISE_INITIALIZER),this.pointwiseRegularizer=As(e.pointwiseRegularizer),this.pointwiseConstraint=Sr(e.pointwiseConstraint)}build(t){if(t=de(t),t.length<this.rank+2)throw new I(`Inputs to SeparableConv${this.rank}D should have rank ${this.rank+2}, but received input shape: ${JSON.stringify(t)}`);const e=this.dataFormat==="channelsFirst"?1:t.length-1;if(t[e]==null||t[e]<0)throw new I(`The channel dimension of the inputs should be defined, but found ${JSON.stringify(t[e])}`);const s=t[e],r=this.kernelSize.concat([s,this.depthMultiplier]),o=[];for(let a=0;a<this.rank;++a)o.push(1);o.push(s*this.depthMultiplier,this.filters);const i=!0;this.depthwiseKernel=this.addWeight("depthwise_kernel",r,"float32",this.depthwiseInitializer,this.depthwiseRegularizer,i,this.depthwiseConstraint),this.pointwiseKernel=this.addWeight("pointwise_kernel",o,"float32",this.pointwiseInitializer,this.pointwiseRegularizer,i,this.pointwiseConstraint),this.useBias?this.bias=this.addWeight("bias",[this.filters],"float32",this.biasInitializer,this.biasRegularizer,i,this.biasConstraint):this.bias=null,this.inputSpec=[new Ce({ndim:this.rank+2,axes:{[e]:s}})],this.built=!0}call(t,e){return k(()=>{t=Gt(t);let s;if(this.rank===1)throw new Z("1D separable convolution is not implemented yet.");return this.rank===2&&(this.dataFormat==="channelsFirst"&&(t=pt(t,[0,2,3,1])),s=O0(t,this.depthwiseKernel.read(),this.pointwiseKernel.read(),this.strides,this.padding,this.dilationRate,"NHWC")),this.useBias&&(s=Ss(s,this.bias.read(),this.dataFormat)),this.activation!=null&&(s=this.activation.apply(s)),this.dataFormat==="channelsFirst"&&(s=pt(s,[0,3,1,2])),s})}getConfig(){const t=super.getConfig();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}}nc.className="SeparableConv";class sc extends nc{constructor(t){super(2,t)}}sc.className="SeparableConv2D",M(sc);class Er extends Kn{constructor(t){super(1,t),Er.verifyArgs(t),this.inputSpec=[{ndim:3}]}getConfig(){const t=super.getConfig();return delete t.rank,delete t.dataFormat,t}static verifyArgs(t){if(typeof t.kernelSize!="number"&&!ai(t.kernelSize,"number",1,1))throw new I(`Conv1D expects config.kernelSize to be number or number[] with length 1, but received ${JSON.stringify(t.kernelSize)}.`)}}Er.className="Conv1D",M(Er);class rc extends pe{constructor(t){super(t),typeof t.cropping=="number"?this.cropping=[[t.cropping,t.cropping],[t.cropping,t.cropping]]:typeof t.cropping[0]=="number"?this.cropping=[[t.cropping[0],t.cropping[0]],[t.cropping[1],t.cropping[1]]]:this.cropping=t.cropping,this.dataFormat=t.dataFormat===void 0?"channelsLast":t.dataFormat,this.inputSpec=[{ndim:4}]}computeOutputShape(t){return this.dataFormat==="channelsFirst"?[t[0],t[1],t[2]-this.cropping[0][0]-this.cropping[0][1],t[3]-this.cropping[1][0]-this.cropping[1][1]]:[t[0],t[1]-this.cropping[0][0]-this.cropping[0][1],t[2]-this.cropping[1][0]-this.cropping[1][1],t[3]]}call(t,e){return k(()=>{if(t=Gt(t),this.dataFormat==="channelsLast"){const s=wr(t,this.cropping[0][0],t.shape[1]-this.cropping[0][0]-this.cropping[0][1],2);return wr(s,this.cropping[1][0],t.shape[2]-this.cropping[1][1]-this.cropping[1][0],3)}else{const s=wr(t,this.cropping[0][0],t.shape[2]-this.cropping[0][0]-this.cropping[0][1],3);return wr(s,this.cropping[1][0],t.shape[3]-this.cropping[1][1]-this.cropping[1][0],4)}})}getConfig(){const t={cropping:this.cropping,dataFormat:this.dataFormat},e=super.getConfig();return Object.assign(t,e),t}}rc.className="Cropping2D",M(rc);class vi extends pe{constructor(t){super(t),this.DEFAULT_SIZE=[2,2],this.inputSpec=[{ndim:4}],this.size=t.size==null?this.DEFAULT_SIZE:t.size,this.dataFormat=t.dataFormat==null?"channelsLast":t.dataFormat,mt(this.dataFormat),this.interpolation=t.interpolation==null?"nearest":t.interpolation,Wy(this.interpolation)}computeOutputShape(t){if(this.dataFormat==="channelsFirst"){const e=t[2]==null?null:this.size[0]*t[2],s=t[3]==null?null:this.size[1]*t[3];return[t[0],t[1],e,s]}else{const e=t[1]==null?null:this.size[0]*t[1],s=t[2]==null?null:this.size[1]*t[2];return[t[0],e,s,t[3]]}}call(t,e){return k(()=>{let s=Gt(t);const r=s.shape;if(this.dataFormat==="channelsFirst"){s=pt(s,[0,2,3,1]);const o=this.size[0]*r[2],i=this.size[1]*r[3],a=this.interpolation==="nearest"?pr.resizeNearestNeighbor(s,[o,i]):pr.resizeBilinear(s,[o,i]);return pt(a,[0,3,1,2])}else{const o=this.size[0]*r[1],i=this.size[1]*r[2];return this.interpolation==="nearest"?pr.resizeNearestNeighbor(s,[o,i]):pr.resizeBilinear(s,[o,i])}})}getConfig(){const t={size:this.size,dataFormat:this.dataFormat,interpolation:this.interpolation},e=super.getConfig();return Object.assign(t,e),t}}vi.className="UpSampling2D",M(vi);/**
|
|
2896
2896
|
* @license
|
|
2897
2897
|
* Copyright 2018 Google LLC
|
|
2898
2898
|
*
|
|
@@ -2900,7 +2900,7 @@
|
|
|
2900
2900
|
* license that can be found in the LICENSE file or at
|
|
2901
2901
|
* https://opensource.org/licenses/MIT.
|
|
2902
2902
|
* =============================================================================
|
|
2903
|
-
*/function
|
|
2903
|
+
*/function Cr(n,t,e,s,r,o){return k(()=>{mt(r),su(o),oe(s),e==null&&(e=[1,1]),s==null&&(s="valid"),r==null&&(r=Hn()),o==null&&(o="max"),n=Ju(n,r);let i;const a=s==="same"?"same":"valid";return o==="max"?i=Zg(n,t,e,a):i=Lm(n,t,e,a),r==="channelsFirst"&&(i=pt(i,[0,3,1,2])),i})}function oc(n,t,e,s,r,o){return k(()=>{mt(r),su(o),oe(s),e==null&&(e=[1,1,1]),s==null&&(s="valid"),r==null&&(r=Hn()),o==null&&(o="max"),n=Zu(n,r);let i;const a=s==="same"?"same":"valid";return o==="max"?i=t0(n,t,e,a):i=Om(n,t,e,a),r==="channelsFirst"&&(i=pt(i,[0,4,1,2,3])),i})}class ic extends pe{constructor(t){if(t.poolSize==null&&(t.poolSize=2),super(t),typeof t.poolSize=="number")this.poolSize=[t.poolSize];else if(Array.isArray(t.poolSize)&&t.poolSize.length===1&&typeof t.poolSize[0]=="number")this.poolSize=t.poolSize;else 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)}`);if(Me(this.poolSize,"poolSize"),t.strides==null)this.strides=this.poolSize;else if(typeof t.strides=="number")this.strides=[t.strides];else if(Array.isArray(t.strides)&&t.strides.length===1&&typeof t.strides[0]=="number")this.strides=t.strides;else 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)}`);Me(this.strides,"strides"),this.padding=t.padding==null?"valid":t.padding,oe(this.padding),this.inputSpec=[new Ce({ndim:3})]}computeOutputShape(t){t=de(t);const e=wn(t[1],this.poolSize[0],this.padding,this.strides[0]);return[t[0],e,t[2]]}call(t,e){return k(()=>{this.invokeCallHook(t,e),t=ui(Gt(t),2);const s=this.poolingFunction(Gt(t),[this.poolSize[0],1],[this.strides[0],1],this.padding,"channelsLast");return fr(s,[2])})}getConfig(){const t={poolSize:this.poolSize,padding:this.padding,strides:this.strides},e=super.getConfig();return Object.assign(t,e),t}}class ac extends ic{constructor(t){super(t)}poolingFunction(t,e,s,r,o){return mt(o),oe(r),Cr(t,e,s,r,o,"max")}}ac.className="MaxPooling1D",M(ac);class lc extends ic{constructor(t){super(t)}poolingFunction(t,e,s,r,o){return mt(o),oe(r),Cr(t,e,s,r,o,"avg")}}lc.className="AveragePooling1D",M(lc);class uc extends pe{constructor(t){if(t.poolSize==null&&(t.poolSize=[2,2]),super(t),this.poolSize=Array.isArray(t.poolSize)?t.poolSize:[t.poolSize,t.poolSize],t.strides==null)this.strides=this.poolSize;else if(Array.isArray(t.strides)){if(t.strides.length!==2)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}.`);this.strides=t.strides}else this.strides=[t.strides,t.strides];Me(this.poolSize,"poolSize"),Me(this.strides,"strides"),this.padding=t.padding==null?"valid":t.padding,this.dataFormat=t.dataFormat==null?"channelsLast":t.dataFormat,mt(this.dataFormat),oe(this.padding),this.inputSpec=[new Ce({ndim:4})]}computeOutputShape(t){t=de(t);let e=this.dataFormat==="channelsFirst"?t[2]:t[1],s=this.dataFormat==="channelsFirst"?t[3]:t[2];return e=wn(e,this.poolSize[0],this.padding,this.strides[0]),s=wn(s,this.poolSize[1],this.padding,this.strides[1]),this.dataFormat==="channelsFirst"?[t[0],t[1],e,s]:[t[0],e,s,t[3]]}call(t,e){return k(()=>(this.invokeCallHook(t,e),this.poolingFunction(Gt(t),this.poolSize,this.strides,this.padding,this.dataFormat)))}getConfig(){const t={poolSize:this.poolSize,padding:this.padding,strides:this.strides,dataFormat:this.dataFormat},e=super.getConfig();return Object.assign(t,e),t}}class Ii extends uc{constructor(t){super(t)}poolingFunction(t,e,s,r,o){return mt(o),oe(r),Cr(t,e,s,r,o,"max")}}Ii.className="MaxPooling2D",M(Ii);class cc extends uc{constructor(t){super(t)}poolingFunction(t,e,s,r,o){return mt(o),oe(r),Cr(t,e,s,r,o,"avg")}}cc.className="AveragePooling2D",M(cc);class hc extends pe{constructor(t){if(t.poolSize==null&&(t.poolSize=[2,2,2]),super(t),this.poolSize=Array.isArray(t.poolSize)?t.poolSize:[t.poolSize,t.poolSize,t.poolSize],t.strides==null)this.strides=this.poolSize;else if(Array.isArray(t.strides)){if(t.strides.length!==3)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}.`);this.strides=t.strides}else this.strides=[t.strides,t.strides,t.strides];Me(this.poolSize,"poolSize"),Me(this.strides,"strides"),this.padding=t.padding==null?"valid":t.padding,this.dataFormat=t.dataFormat==null?"channelsLast":t.dataFormat,mt(this.dataFormat),oe(this.padding),this.inputSpec=[new Ce({ndim:5})]}computeOutputShape(t){t=de(t);let e=this.dataFormat==="channelsFirst"?t[2]:t[1],s=this.dataFormat==="channelsFirst"?t[3]:t[2],r=this.dataFormat==="channelsFirst"?t[4]:t[3];return e=wn(e,this.poolSize[0],this.padding,this.strides[0]),s=wn(s,this.poolSize[1],this.padding,this.strides[1]),r=wn(r,this.poolSize[2],this.padding,this.strides[2]),this.dataFormat==="channelsFirst"?[t[0],t[1],e,s,r]:[t[0],e,s,r,t[4]]}call(t,e){return k(()=>(this.invokeCallHook(t,e),this.poolingFunction(Gt(t),this.poolSize,this.strides,this.padding,this.dataFormat)))}getConfig(){const t={poolSize:this.poolSize,padding:this.padding,strides:this.strides,dataFormat:this.dataFormat},e=super.getConfig();return Object.assign(t,e),t}}class fc extends hc{constructor(t){super(t)}poolingFunction(t,e,s,r,o){return mt(o),oe(r),oc(t,e,s,r,o,"max")}}fc.className="MaxPooling3D",M(fc);class dc extends hc{constructor(t){super(t)}poolingFunction(t,e,s,r,o){return mt(o),oe(r),oc(t,e,s,r,o,"avg")}}dc.className="AveragePooling3D",M(dc);class pc extends pe{constructor(t){super(t),this.inputSpec=[new Ce({ndim:3})]}computeOutputShape(t){return[t[0],t[2]]}call(t,e){throw new Z}}class mc extends pc{constructor(t){super(t||{})}call(t,e){return k(()=>{const s=Gt(t);return St(s,1)})}}mc.className="GlobalAveragePooling1D",M(mc);class gc extends pc{constructor(t){super(t||{})}call(t,e){return k(()=>{const s=Gt(t);return Ge(s,1)})}}gc.className="GlobalMaxPooling1D",M(gc);class bc extends pe{constructor(t){super(t),this.dataFormat=t.dataFormat==null?"channelsLast":t.dataFormat,mt(this.dataFormat),this.inputSpec=[new Ce({ndim:4})]}computeOutputShape(t){return t=t,this.dataFormat==="channelsLast"?[t[0],t[3]]:[t[0],t[1]]}call(t,e){throw new Z}getConfig(){const t={dataFormat:this.dataFormat},e=super.getConfig();return Object.assign(t,e),t}}class yc extends bc{call(t,e){return k(()=>{const s=Gt(t);return this.dataFormat==="channelsLast"?St(s,[1,2]):St(s,[2,3])})}}yc.className="GlobalAveragePooling2D",M(yc);class wc extends bc{call(t,e){return k(()=>{const s=Gt(t);return this.dataFormat==="channelsLast"?Ge(s,[1,2]):Ge(s,[2,3])})}}wc.className="GlobalMaxPooling2D",M(wc);/**
|
|
2904
2904
|
* @license
|
|
2905
2905
|
* Copyright 2018 Google LLC
|
|
2906
2906
|
*
|
|
@@ -2908,7 +2908,7 @@
|
|
|
2908
2908
|
* license that can be found in the LICENSE file or at
|
|
2909
2909
|
* https://opensource.org/licenses/MIT.
|
|
2910
2910
|
* =============================================================================
|
|
2911
|
-
*/function
|
|
2911
|
+
*/function kr(n,t){return k(()=>{n.dtype!=="float32"&&(n=ot(n,"float32"));const e=et(xs(n),t,!0),s=or(e.shape,gt()),r=fe(zn(e,s));return X(n,r)})}function _r(n,t){return k(()=>St(xs(J(t,n)),-1))}function Ai(n,t){return k(()=>St(Lt(J(t,n)),-1))}function Ei(n,t){return k(()=>{const e=J(n,t),s=he(Lt(n),gt(),Number.MAX_VALUE),r=Lt(X(e,s));return N(100,St(r,-1))})}function yw(n,t){return k(()=>{const e=he(t,gt(),Number.MAX_VALUE),s=dn(O(1,e)),r=he(n,gt(),Number.MAX_VALUE),o=dn(O(1,r));return St(xs(J(s,o)),-1)})}function ww(n,t){return k(()=>{const e=zn(0,J(1,N(n,t)));return St(xs(e),-1)})}function xw(n,t){return k(()=>{const e=zn(0,J(1,N(n,t)));return St(e,-1)})}function Sw(n,t){return k(()=>{const e=et(N(n,t),-1),s=Ge(N(J(1,n),t),-1);return zn(0,O(1,J(s,e)))})}function $w(n,t){return k(()=>{const e=Math.log(2),s=J(t,n),r=J(O(s,qo(N(-2,s))),e);return St(r,-1)})}function Cs(n,t,e=!1){return k(()=>{if(e)t=zl(t);else{const s=et(t,t.shape.length-1,!0);t=X(t,s)}return t=he(t,gt(),1-gt()),Fn(et(N(ot(n,"float32"),dn(t)),t.shape.length-1))})}function Tr(n,t,e=!1){return k(()=>{const s=ot(_g(jy(n)),"int32");t=he(t,gt(),1-gt());const r=t.shape,o=L(o0(s,r[r.length-1]),r);return Cs(o,t,e)})}function vw(n,t){if(!Zt(n.shape,t.shape))throw new I(`logits and labels must have the same shape, but got shapes ${JSON.stringify(n.shape)} and ${JSON.stringify(t.shape)}`);return k(()=>{const e=ms(t),s=Fn(Lt(t));return O(J(e,N(t,n)),Gg(Go(s)))})}function Nr(n,t){return k(()=>{let e;return e=he(t,gt(),1-gt()),e=dn(X(e,J(1,e))),St(vw(n,e),-1)})}function Iw(n,t){return k(()=>{const e=he(n,gt(),1),s=he(t,gt(),1);return et(N(n,dn(X(e,s))),-1)})}function Aw(n,t){return k(()=>{const e=dn(O(gt(),t));return St(J(t,N(n,e)),-1)})}function xc(n,t){return k(()=>{const e=kr(n,-1),s=kr(t,-1),r=N(e,s);return Fn(et(r,-1))})}const Dr={meanSquaredError:_r,meanAbsoluteError:Ai,meanAbsolutePercentageError:Ei,meanSquaredLogarithmicError:yw,squaredHinge:ww,hinge:xw,categoricalHinge:Sw,logcosh:$w,categoricalCrossentropy:Cs,sparseCategoricalCrossentropy:Tr,binaryCrossentropy:Nr,kullbackLeiblerDivergence:Iw,poisson:Aw,cosineProximity:xc};function Ci(n){if(typeof n=="string"){if(n in Dr)return Dr[n];let t=`Unknown loss ${n}`;throw n.toLowerCase().includes("softmaxcrossentropy")&&(t=`Unknown loss ${n}. Use "categoricalCrossentropy" as the string name for tf.losses.softmaxCrossEntropy`),new I(t)}else return n}/**
|
|
2912
2912
|
* @license
|
|
2913
2913
|
* Copyright 2018 Google LLC
|
|
2914
2914
|
*
|
|
@@ -2916,7 +2916,7 @@
|
|
|
2916
2916
|
* license that can be found in the LICENSE file or at
|
|
2917
2917
|
* https://opensource.org/licenses/MIT.
|
|
2918
2918
|
* =============================================================================
|
|
2919
|
-
*/class
|
|
2919
|
+
*/class xn extends pe{constructor(t){super(t||{}),this.supportsMasking=!0}mergeFunction(t){throw new Z}computeElementwiseOpOutputShape(t,e){if(t==null||e==null)return null;if(t.length<e.length)return this.computeElementwiseOpOutputShape(e,t);if(e.length===0)return t;const s=t.slice(0,t.length-e.length);for(let r=0;r<e.length;++r){const o=t[t.length-e.length+r],i=e[r];if(o==null||i==null||o<0||i<0)s.push(null);else if(o===1)s.push(i);else if(i===1)s.push(o);else{if(o!==i)throw new I("Operands could not be broadcast together with shapes "+JSON.stringify(t)+" "+JSON.stringify(e));s.push(o)}}return s}build(t){if(Array.isArray(t)&&!Array.isArray(t[0])&&(t=[de(t)]),t=t,t.length<2)throw new I(`A merge layer should be called on an Array of at least 2 inputs. Got ${t.length} input(s).`);let e=[];for(const o of t)o!=null&&o[0]!==null&&e.push(o[0]);if(e=gn(e),e.length>1)throw new I(`Can not merge tensors with different batch sizes. Got tensors with shapes: ${JSON.stringify(t)}.`);let s=t[0]==null?null:t[0].slice(1);for(let o=1;o<t.length;++o){const i=t[o]==null?null:t[o].slice(1);s=this.computeElementwiseOpOutputShape(s,i)}const r=t.map(o=>o.length);t.indexOf(null)===-1&&gn(r).length===1?this.reshapeRequired=!1:this.reshapeRequired=!0}call(t,e){return k(()=>{if(t=t,this.reshapeRequired){const s=[],r=t.map(o=>o.rank);if(r.indexOf(null)===-1){const o=lu(r);for(let i of t){const a=i.rank;for(let l=0;l<o-a;++l)i=ui(i,1);s.push(i)}return this.mergeFunction(s)}else{let o=!1;for(const l of t){const u=l.rank;if(u==null){const c=l.shape,h=c[0],f=c.slice(1).concat([h]);let d=L(l,[h].concat(ws(c.slice(1))));d=pt(d,[1,0]),d=L(d,f),s.push(d),o=!0}else if(u>1){const c=yr(1,u).concat([0]);s.push(pt(l,c)),o=!0}else s.push(l)}let i=this.mergeFunction(s);const a=i.rank;if(o){if(a==null){const l=i.shape,u=l.length,c=l[u-1],h=[c].concat(l.slice(0,l.length-1));i=L(pt(L(i,[-1,c]),[1,0]),h)}else if(a>1){const l=[a-1].concat(yr(0,a-1));i=pt(i,l)}}return i}}else return this.mergeFunction(t)})}computeOutputShape(t){t=t;let e;t[0]==null?e=null:e=t[0].slice(1);for(let r=1;r<t.length;++r){const o=t[r]==null?null:t[r].slice(1);e=this.computeElementwiseOpOutputShape(e,o)}let s=[];for(const r of t)r!=null&&r[0]!==null&&s.push(r[0]);return s=gn(s),s.length===1?e=s.concat(e):e=[null].concat(e),e}computeMask(t,e){return k(()=>{if(e==null)return null;if(!Array.isArray(e))throw new I("`mask` should be an Array");if(!Array.isArray(t))throw new I("`inputs` should be an Array");if(e.length!==t.length)throw new I(`The Array 'inputs' and 'mask' are expected to have the same length, but have different lengths (${t.length} vs ${e.length})`);if(e.every(r=>r==null))return null;e=e.map(r=>r==null?r:ve(r,0));let s=e[0];for(let r=1;r<e.length-1;++r)s=ur(s,e[r]);return s})}}class Sc extends xn{constructor(t){super(t)}mergeFunction(t){return k(()=>{let e=t[0].clone();for(let s=1;s<t.length;++s)e=O(e,t[s]);return e})}}Sc.className="Add",M(Sc);class $c extends xn{constructor(t){super(t)}mergeFunction(t){return k(()=>{let e=t[0].clone();for(let s=1;s<t.length;++s)e=N(e,t[s]);return e})}}$c.className="Multiply",M($c);class vc extends xn{constructor(t){super(t)}mergeFunction(t){return k(()=>{let e=t[0].clone();for(let s=1;s<t.length;++s)e=O(e,t[s]);return N(1/t.length,e)})}}vc.className="Average",M(vc);class Ic extends xn{constructor(t){super(t)}mergeFunction(t){return k(()=>{let e=t[0];for(let s=1;s<t.length;++s)e=zn(e,t[s]);return e})}}Ic.className="Maximum",M(Ic);class Ac extends xn{constructor(t){super(t)}mergeFunction(t){return k(()=>{let e=t[0];for(let s=1;s<t.length;++s)e=cr(e,t[s]);return e})}}Ac.className="Minimum",M(Ac);class ki extends xn{constructor(t){super(t),this.DEFAULT_AXIS=-1,t==null&&(t={}),this.axis=t.axis==null?this.DEFAULT_AXIS:t.axis,this.supportsMasking=!0,this.reshapeRequired=!1}build(t){if(!(Array.isArray(t)&&Array.isArray(t[0]))||t.length===1)throw new I("A `Concatenate` layer should be called on a list of at least 2 inputs");t=t;let e=!0;for(const r of t)if(r!=null){e=!1;break}if(e)return;const s=[];for(let r=0;r<t.length;++r){const o=t[r].slice();o.splice(this.axis,1);let i=!1;for(const a of s)if(Zt(a,o)){i=!0;break}i||s.push(o)}if(s.length>1)throw new I("A `Concatenate` layer requires inputs with matching shapes except for the concat axis. Got input shapes: "+JSON.stringify(t))}mergeFunction(t){return k(()=>Hy(t,this.axis))}computeOutputShape(t){if(!(Array.isArray(t)&&Array.isArray(t[0])))throw new I("A `Concatenate` layer should be called on a list of inputs.");const e=t,s=e[0].slice(),r=this.axis<0?s.length+this.axis:this.axis;for(const o of e.slice(1)){if(s[r]==null||o[r]==null){s[r]=null;break}s[r]+=o[r]}return s}computeMask(t,e){if(e==null)return null;if(!Array.isArray(e))throw new I("`mask` should be an array for Concatenate");if(!Array.isArray(t))throw new I("`inputs` should be an array for Concatenate");if(e.length!==t.length)throw new I(`Mismatch in the length of mask (${e.length}) and the legnth of inputs (${t.length})`);return k(()=>{let s=!0;if(e.forEach(i=>{if(i!=null){s=!1;return}}),s)return null;const r=[];for(let i=0;i<t.length;++i)e[i]==null?r.push(ot(Pl(t[i]),"bool")):e[i].rank<t[i].rank?r.push(ve(e[i],-1)):r.push(e[i]);const o=an(r,this.axis);return Em(o,-1,!1)})}getConfig(){const t={axis:this.axis},e=super.getConfig();return Object.assign(t,e),t}}ki.className="Concatenate",M(ki);function ks(n,t){for(;n<0;)n+=t;return n}function Ew(n,t,e){if(n.shape.length>3||t.shape.length>3)throw new Z("batchDot is not implemented for tensors of 4D or higher rank yet");if(w(n.shape.length>=2,()=>`batchDot requires the rank of x to be >= 2, but got ${n.shape.length}`),w(n.shape.length>=2,()=>`batchDot requires the rank of y to be >= 2, but got ${t.shape.length}`),typeof e=="number"&&(e=[e,e]),n.dtype==="complex64"||t.dtype==="complex64")throw new Z("batchDot is not implemented for complex64-type Tensors yet.");const s=n.shape.length,r=t.shape.length;e==null&&(e=[s-1,r-2]);const o=e;return k(()=>{let i;if(s>r){i=s-r;const l=[];for(let u=0;u<i;++u)l.push(1);t=L(t,t.shape.concat(l))}else if(r>s){i=r-s;const l=[];for(let u=0;u<i;++u)l.push(1);n=L(n,n.shape.concat(l))}else i=0;let a;if(n.shape.length===2&&t.shape.length===2)o[0]===o[1]?a=et(N(n,t),o[0]):a=et(N(pt(n,[1,0]),t),o[1]);else{const l=o[0]!==n.shape.length-1,u=o[1]===t.shape.length-1;a=Se(n,t,l,u)}if(i>0){let l;s>r?l=s+r-3:l=s-1;const u=[];for(let c=l;c<l+i;++c)u.push(c);a=fr(a,u)}return a.shape.length===1&&(a=ve(a,1)),a})}class Ec extends xn{constructor(t){super(t),this.axes=t.axes,this.normalize=t.normalize==null?!1:t.normalize,this.supportsMasking=!0,this.reshapeRequired=!1}build(t){w(Array.isArray(t)&&t.length===2&&Array.isArray(t[0])&&Array.isArray(t[1]),()=>"A `Dot` layer should be called on a list of exactly 2 inputs.");const e=t[0],s=t[1];if(e.length>3||s.length>3)throw new Z("Dot layer does not support tensors of 4D or higher rank yet.");const r=this.interpretAxes(e,s);if(e[r[0]]!==s[r[1]])throw new I(`Dimension incompatibility: ${e[r[0]]} !== ${s[r[1]]}`)}mergeFunction(t){if(t.length!==2)throw new I(`A \`Dot\` layer must be called on exactly 2 inputs, but received ${t.length} input(s).`);let e=t[0],s=t[1],r;return Array.isArray(this.axes)?r=this.axes.map((o,i)=>ks(o,t[i].shape.length)):r=[ks(this.axes,e.shape.length),ks(this.axes,s.shape.length)],this.normalize&&(e=kr(e,r[0]),s=kr(s,r[1])),Ew(e,s,r)}interpretAxes(t,e){let s;return Array.isArray(this.axes)?s=this.axes:s=[ks(this.axes,t.length),ks(this.axes,e.length)],s}computeOutputShape(t){w(Array.isArray(t)&&t.length===2&&Array.isArray(t[0])&&Array.isArray(t[1]),()=>"A `Dot` layer should be called on a list of exactly 2 inputs.");const e=t[0].slice(),s=t[1].slice();if(e.length>3||s.length>3)throw new Z("Dot layer does not support tensors of 4D or higher rank yet.");const r=this.interpretAxes(e,s);e.splice(r[0],1),s.splice(r[1],1),s.splice(0,1);const o=e.concat(s);return o.length===1&&o.push(1),o}computeMask(t,e){return null}getConfig(){const t={axes:this.axes,normalize:this.normalize},e=super.getConfig();return Object.assign(t,e),t}}Ec.className="Dot",M(Ec);/**
|
|
2920
2920
|
* @license
|
|
2921
2921
|
* Copyright 2018 Google LLC
|
|
2922
2922
|
*
|
|
@@ -2924,7 +2924,7 @@
|
|
|
2924
2924
|
* license that can be found in the LICENSE file or at
|
|
2925
2925
|
* https://opensource.org/licenses/MIT.
|
|
2926
2926
|
* =============================================================================
|
|
2927
|
-
*/async function
|
|
2927
|
+
*/async function Sn(n){if(n==null)return;const t=[],e=[],s=[];for(const r in n){const o=n[r];if(typeof o!="number"){const i=o;t.push(i.data()),e.push(r),s.push(i)}}if(t.length>0){const r=await Promise.all(t);for(let o=0;o<r.length;++o)n[e[o]]=r[o][0];ut(s)}}function Cc(n){if(n!=null)for(const t in n){const e=n[t];typeof e!="number"&&e.dispose()}}/**
|
|
2928
2928
|
* @license
|
|
2929
2929
|
* Copyright 2018 Google LLC
|
|
2930
2930
|
*
|
|
@@ -2932,7 +2932,7 @@
|
|
|
2932
2932
|
* license that can be found in the LICENSE file or at
|
|
2933
2933
|
* https://opensource.org/licenses/MIT.
|
|
2934
2934
|
* =============================================================================
|
|
2935
|
-
*/var kc;(function(n){n[n.SILENT=0]="SILENT",n[n.VERBOSE=1]="VERBOSE"})(kc||(kc={}));const Cw=125;class
|
|
2935
|
+
*/var kc;(function(n){n[n.SILENT=0]="SILENT",n[n.VERBOSE=1]="VERBOSE"})(kc||(kc={}));const Cw=125;class _s{constructor(){this.validationData=null}setParams(t){this.params=t}async onEpochBegin(t,e){}async onEpochEnd(t,e){}async onBatchBegin(t,e){}async onBatchEnd(t,e){}async onTrainBegin(t){}async onTrainEnd(t){}setModel(t){}}class kw{constructor(t,e=10){t==null&&(t=[]),this.callbacks=t,this.queueLength=e}append(t){this.callbacks.push(t)}setParams(t){for(const e of this.callbacks)e.setParams(t)}setModel(t){for(const e of this.callbacks)e.setModel(t)}async onEpochBegin(t,e){e==null&&(e={});for(const s of this.callbacks)await s.onEpochBegin(t,e)}async onEpochEnd(t,e){e==null&&(e={});for(const s of this.callbacks)await s.onEpochEnd(t,e)}async onBatchBegin(t,e){e==null&&(e={});for(const s of this.callbacks)await s.onBatchBegin(t,e)}async onBatchEnd(t,e){e==null&&(e={});for(const s of this.callbacks)await s.onBatchEnd(t,e)}async onTrainBegin(t){t==null&&(t={});for(const e of this.callbacks)await e.onTrainBegin(t)}async onTrainEnd(t){t==null&&(t={});for(const e of this.callbacks)await e.onTrainEnd(t)}}class _w extends _s{constructor(){super()}async onEpochBegin(t){this.seen=0,this.totals={}}async onBatchEnd(t,e){e==null&&(e={});const s=e.size==null?0:e.size;this.seen+=s;for(const r in e){const o=e[r];if(typeof o=="number")this.totals.hasOwnProperty(r)||(this.totals[r]=0),this.totals[r]=this.totals[r]+o*s;else{let i;r in this.totals?i=this.totals[r]:this.totals[r]=0;const a=k(()=>O(this.totals[r],N(o,s)));this.totals[r]=a,i!=null&&i.dispose()}}}async onEpochEnd(t,e){if(e!=null)for(const s of this.params.metrics)this.totals[s]!=null&&(typeof this.totals[s]=="number"?e[s]=this.totals[s]/this.seen:k(()=>{const r=N(X(1,this.seen),this.totals[s]);e[s]=r,this.totals[s].dispose(),Ln(e[s])}))}}class Tw extends _s{async onTrainBegin(t){this.epoch=[],this.history={}}async onEpochEnd(t,e){e==null&&(e={}),this.epoch.push(t);for(const s in e)this.history[s]==null&&(this.history[s]=[]),this.history[s].push(e[s])}async syncData(){const t=[],e=[],s=[];for(const o in this.history){const i=this.history[o];for(let a=0;a<i.length;++a)if(typeof i[a]!="number"){const l=i[a];t.push(l.data()),e.push(o),s.push(a)}}const r=await Promise.all(t);for(let o=0;o<r.length;++o)this.history[e[o]][s[o]].dispose(),this.history[e[o]][s[o]]=r[o][0]}}class Nw extends _s{constructor(t,e){if(super(),this.currentEpoch=0,this.nowFunc=t.nowFunc,this.nextFrameFunc=t.nextFrameFunc||ly,this.yieldEvery=e||"auto",this.yieldEvery==="auto"&&(this.yieldEvery=Cw),this.yieldEvery==="never"&&t.onYield!=null)throw new Error("yieldEvery is `never` but you provided an `onYield` callback. Either change `yieldEvery` or remove the callback");ao(this.yieldEvery)&&(this.maybeWait=zy(this.maybeWait.bind(this),this.yieldEvery,this.nowFunc)),this.trainBegin=t.onTrainBegin,this.trainEnd=t.onTrainEnd,this.epochBegin=t.onEpochBegin,this.epochEnd=t.onEpochEnd,this.batchBegin=t.onBatchBegin,this.batchEnd=t.onBatchEnd,this.yield=t.onYield}async maybeWait(t,e,s){const r=[];this.yield!=null&&(await Sn(s),r.push(this.yield(t,e,s))),r.push(this.nextFrameFunc()),await Promise.all(r)}async onEpochBegin(t,e){this.currentEpoch=t,this.epochBegin!=null&&(await Sn(e),await this.epochBegin(t,e))}async onEpochEnd(t,e){const s=[];this.epochEnd!=null&&(await Sn(e),s.push(this.epochEnd(t,e))),this.yieldEvery==="epoch"&&s.push(this.nextFrameFunc()),await Promise.all(s)}async onBatchBegin(t,e){this.batchBegin!=null&&(await Sn(e),await this.batchBegin(t,e))}async onBatchEnd(t,e){const s=[];this.batchEnd!=null&&(await Sn(e),s.push(this.batchEnd(t,e))),this.yieldEvery==="batch"?s.push(this.nextFrameFunc()):ao(this.yieldEvery)&&s.push(this.maybeWait(this.currentEpoch,t,e)),await Promise.all(s)}async onTrainBegin(t){this.trainBegin!=null&&(await Sn(t),await this.trainBegin(t))}async onTrainEnd(t){this.trainEnd!=null&&(await Sn(t),await this.trainEnd(t))}}function _c(n,t){return n==null&&(n={}),n instanceof _s?[n]:Array.isArray(n)&&n[0]instanceof _s?n:nt(n).map(s=>new Nw(s,t))}class ie{constructor(){}static registerCallbackConstructor(t,e){w(t>=0&&Number.isInteger(t),()=>`Verbosity level is expected to be an integer >= 0, but got ${t}`),ie.checkForDuplicate(e),ie.constructors[t]==null&&(ie.constructors[t]=[]),ie.constructors[t].push(e)}static checkForDuplicate(t){for(const e in ie.constructors)ie.constructors[+e].forEach(r=>{if(r===t)throw new I("Duplicate callback constructor.")})}static clear(){ie.constructors={}}static createCallbacks(t){const e=[];for(const s in ie.constructors){const r=+s;t>=r&&e.push(...ie.constructors[r])}return e.map(s=>new s)}}ie.constructors={};function Tc(n,t,e,s,r,o,i,a,l){const u=new Tw,c=[new _w,...ie.createCallbacks(t)];n!=null&&c.push(...n),c.push(u);const h=new kw(c);return h.setParams({epochs:e,initialEpoch:s,samples:r,steps:o,batchSize:i,verbose:t,doValidation:a,metrics:l}),{callbackList:h,history:u}}/**
|
|
2936
2936
|
* @license
|
|
2937
2937
|
* Copyright 2018 Google LLC
|
|
2938
2938
|
*
|
|
@@ -2940,7 +2940,7 @@
|
|
|
2940
2940
|
* license that can be found in the LICENSE file or at
|
|
2941
2941
|
* https://opensource.org/licenses/MIT.
|
|
2942
2942
|
* =============================================================================
|
|
2943
|
-
*/function Nc(n,t={},e=!1){return
|
|
2943
|
+
*/function Nc(n,t={},e=!1){return bs(n,se.getMap().classNameMap,t,"layer",e)}/**
|
|
2944
2944
|
* @license
|
|
2945
2945
|
* Copyright 2018 Google LLC
|
|
2946
2946
|
*
|
|
@@ -2948,7 +2948,7 @@
|
|
|
2948
2948
|
* license that can be found in the LICENSE file or at
|
|
2949
2949
|
* https://opensource.org/licenses/MIT.
|
|
2950
2950
|
* =============================================================================
|
|
2951
|
-
*/function Dc(n,t){return k(()=>{const e=N(.5,Pl(t)),s=uu(
|
|
2951
|
+
*/function Dc(n,t){return k(()=>{const e=N(.5,Pl(t)),s=uu(ps(t,e),n.dtype);return St(hn(n,s),-1)})}function Rc(n,t){return k(()=>uu(hn(nr(n,-1),nr(t,-1)),"float32"))}function Dw(n,t){return k(()=>ot(et(ur(hn(n,1),hn(t,1))),"float32"))}function Rw(n,t){return k(()=>ot(et(ur(hn(n,0),hn(t,1))),"float32"))}function Pw(n,t){return k(()=>{const e=Dw(n,t),s=Rw(n,t),r=O(e,s);return ot(fn(ps(r,0),X(e,r),0),"float32")})}function Lw(n,t){return Nr(n,t)}function Mw(n,t){return n.rank===t.rank&&(n=fr(n,[n.rank-1])),t=nr(t,-1),t.dtype!==n.dtype&&(t=ot(t,n.dtype)),ot(hn(n,t),"float32")}const Ow=_r,Bw=_r,Fw=Ai,zw=Ai,Uw=Ei,Ww=Ei,Pc=Cs,Gw=xc,Lc=Tr,Rr={binaryAccuracy:Dc,categoricalAccuracy:Rc,precision:Pw,categoricalCrossentropy:Pc,sparseCategoricalCrossentropy:Lc,mse:Ow,MSE:Bw,mae:Fw,MAE:zw,mape:Uw,MAPE:Ww,cosine:Gw};function Vw(n){if(typeof n=="string"&&n in Rr)return Rr[n];if(typeof n!="string"&&n!=null)return n;throw new I(`Unknown metric ${n}`)}function Pr(n){if(Ae(n!==null,`Unknown LossOrMetricFn ${n}`),typeof n=="string")return n;{let t;for(const e of Object.keys(Dr))if(Dr[e]===n){t=e;break}if(t!==void 0)return t;for(const e of Object.keys(Rr))if(Rr[e]===n){t=e;break}return t!==void 0?t:n.name}}/**
|
|
2952
2952
|
* @license
|
|
2953
2953
|
* Copyright 2018 Google LLC
|
|
2954
2954
|
*
|
|
@@ -2972,7 +2972,7 @@
|
|
|
2972
2972
|
* license that can be found in the LICENSE file or at
|
|
2973
2973
|
* https://opensource.org/licenses/MIT.
|
|
2974
2974
|
* =============================================================================
|
|
2975
|
-
*/function jw(n,t,e,s=console.log){const r=Kw(n),o=["Layer (type)","Input Shape","Output shape","Param #"];r?(t=t||90,e=e||[.32,.61,.89,1]):(t=t||115,e=e||[.24,.48,.7,.8,1]),e[e.length-1]<=1&&(e=e.map(c=>Math.floor(t*c)));let i;if(!r){o.push("Receives inputs"),i=[];for(const c in n.nodesByDepth)i.push(...n.nodesByDepth[c])}s("_".repeat(t)),
|
|
2975
|
+
*/function jw(n,t,e,s=console.log){const r=Kw(n),o=["Layer (type)","Input Shape","Output shape","Param #"];r?(t=t||90,e=e||[.32,.61,.89,1]):(t=t||115,e=e||[.24,.48,.7,.8,1]),e[e.length-1]<=1&&(e=e.map(c=>Math.floor(t*c)));let i;if(!r){o.push("Receives inputs"),i=[];for(const c in n.nodesByDepth)i.push(...n.nodesByDepth[c])}s("_".repeat(t)),Lr(o,e,s),s("=".repeat(t));const a=n.layers;for(let c=0;c<a.length;++c)r?Yw(a[c],e,s):Xw(a[c],e,i,s),s((c===a.length-1?"=":"_").repeat(t));n.checkTrainableWeightsConsistency();const l=Hw(n),u=Ar(n.nonTrainableWeights);s(`Total params: ${l+u}`),s(`Trainable params: ${l}`),s(`Non-trainable params: ${u}`),s("_".repeat(t))}function Hw(n){let t;return n.collectedTrainableWeights!=null?t=Ar(n.collectedTrainableWeights):t=Ar(n.trainableWeights),t}function Kw(n){let t=!0;const e=[],s=[];for(const r in n.nodesByDepth)e.push(n.nodesByDepth[r]);for(const r of e){if(r.length>1||r.length===1&&r[0].inboundLayers.length>1){t=!1;break}s.push(...r)}if(t)for(const r of n.layers){let o=!1;for(const i of r.inboundNodes)if(s.indexOf(i)!==-1)if(o){t=!1;break}else o=!0;if(!t)break}return t}function Lr(n,t,e=console.log){let s="";for(let r=0;r<n.length;++r)r>0&&(s=s.slice(0,s.length-1)+" "),s+=n[r],s=s.slice(0,t[r]),s+=" ".repeat(t[r]-s.length);e(s)}function Yw(n,t,e){let s,r;try{r=n.inboundNodes.map(l=>JSON.stringify(l.inputShapes)).join(",")}catch{r="multiple"}try{s=JSON.stringify(n.outputShape)}catch{s="multiple"}const o=n.name,i=n.getClassName(),a=[`${o} (${i})`,r,s,n.countParams().toString()];Lr(a,t,e)}function Xw(n,t,e,s){let r,o;try{o=n.inboundNodes.map(h=>JSON.stringify(h.inputShapes)).join(",")}catch{o="multiple"}try{r=JSON.stringify(n.outputShape)}catch{r="multiple"}const i=[];for(const h of n.inboundNodes)if(!(e!=null&&e.length>0&&e.indexOf(h)===-1))for(let f=0;f<h.inboundLayers.length;++f){const d=h.inboundLayers[f].name,p=h.nodeIndices[f],g=h.tensorIndices[f];i.push(`${d}[${p}][${g}]`)}const a=n.name,l=n.getClassName(),u=i.length===0?"":i[0],c=[`${a} (${l})`,o,r,n.countParams().toString(),u];Lr(c,t,s);for(let h=1;h<i.length;++h)Lr(["","","","",i[h]],t,s)}/**
|
|
2976
2976
|
* @license
|
|
2977
2977
|
* Copyright 2018 Google LLC
|
|
2978
2978
|
*
|
|
@@ -2980,7 +2980,7 @@
|
|
|
2980
2980
|
* license that can be found in the LICENSE file or at
|
|
2981
2981
|
* https://opensource.org/licenses/MIT.
|
|
2982
2982
|
* =============================================================================
|
|
2983
|
-
*/function Bc(n,t,e){return(n==="inboundNodes"||n==="outputLayers"||n==="inputLayers")&&t===0&&typeof e=="string"}function Ti(n,t){if(n===null)return null;if(typeof n=="string")return
|
|
2983
|
+
*/function Bc(n,t,e){return(n==="inboundNodes"||n==="outputLayers"||n==="inputLayers")&&t===0&&typeof e=="string"}function Ti(n,t){if(n===null)return null;if(typeof n=="string")return mn(n);if(typeof n=="number"||typeof n=="boolean")return n;if(n instanceof Array){const e=[],s=n.length;for(let r=0;r<s;++r){const o=n[r];Bc(t,r,o)?e.push(o):e.push(Ti(o,t))}return e}else{const e={};for(const s of Object.keys(n)){const r=n[s];if(s==="name"&&typeof r=="string")e[s]=r;else{const o=mn(s);e[o]=Ti(r,o)}}return e}}function Ni(n,t){if(n==null)return null;if(typeof n=="string")return Le(n);if(typeof n=="number"||typeof n=="boolean")return n;if(n instanceof Array){const e=[],s=n.length;for(let r=0;r<s;++r){const o=n[r];Bc(t,r,o)?e.push(o):e.push(Ni(o,t))}return e}else{const e={};for(const s of Object.keys(n)){const r=n[s],o=Le(s);(s==="name"||s==="className")&&typeof r=="string"?e[o]=r:e[o]=Ni(r,s)}return e}}/** @license See the LICENSE file. */const Fc="4.20.0";/**
|
|
2984
2984
|
* @license
|
|
2985
2985
|
* Copyright 2022 Google LLC
|
|
2986
2986
|
*
|
|
@@ -2996,7 +2996,7 @@
|
|
|
2996
2996
|
* license that can be found in the LICENSE file or at
|
|
2997
2997
|
* https://opensource.org/licenses/MIT.
|
|
2998
2998
|
* =============================================================================
|
|
2999
|
-
*/class
|
|
2999
|
+
*/class Ts extends pe{constructor(t){if(super({dtype:t.dtype,name:t.name!=null?t.name:di("input").toString()}),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)throw new I("Only provide the inputShape OR batchInputShape argument to inputLayer, not both at the same time.");let e=t.batchInputShape;if(e==null){if(t.inputShape==null)throw new I("An InputLayer should be passed either a `batchInputShape` or an `inputShape`.");e=[t.batchSize].concat(t.inputShape)}else if(t.batchSize!=null)throw new I("Cannot specify batchSize if batchInputShape is specified when creating an InputLayer.");const s=t.dtype||"float32";this.batchInputShape=e,this.dtype=s,this.inputSpec=[{shape:e}];const r=new yn(this.dtype,this.batchInputShape,this,[],{},this.name);r.nodeIndex=0,r.tensorIndex=0,new xi({outboundLayer:this,inboundLayers:[],nodeIndices:[],tensorIndices:[],inputTensors:[r],outputTensors:[r],inputMasks:[null],outputMasks:[null],inputShapes:[e],outputShapes:[e]})}apply(t,e){throw new I(`Cannot pass any input to an InputLayer's apply() method. InputLayer name: ${this.name}`)}dispose(){return{refCountAfterDispose:this._refCount,numDisposedVariables:0}}getConfig(){return{batchInputShape:this.batchInputShape,dtype:this.dtype,sparse:this.sparse,name:this.name}}}Ts.className="InputLayer",M(Ts);function Jw(n){if(n.batchShape==null&&n.shape==null)throw new Error("Please provide to Input either a `shape` or a `batchShape` argument. Note that `shape` does not include the batch dimension.");if(n.batchShape!=null&&n.shape!=null)throw new I("Please provide either a `shape` or `batchShape` argument to Input, but not both.");let t=n.batchShape;n.shape!=null&&t==null&&(t=[null].concat(n.shape));let e=n.dtype;return e==null&&(e="float32"),new Ts({batchInputShape:t,name:n.name,dtype:e,sparse:n.sparse}).inboundNodes[0].outputTensors[0]}/**
|
|
3000
3000
|
* @license
|
|
3001
3001
|
* Copyright 2018 Google LLC
|
|
3002
3002
|
*
|
|
@@ -3004,7 +3004,7 @@
|
|
|
3004
3004
|
* license that can be found in the LICENSE file or at
|
|
3005
3005
|
* https://opensource.org/licenses/MIT.
|
|
3006
3006
|
* =============================================================================
|
|
3007
|
-
*/function Zw(n,t){if(n.dtype==null||n.dtype===t.dtype)return t;try{return ot(t,n.dtype)}catch{throw new I(`The dtype of the feed (${t.dtype}) can not be cast to the dtype of the key '${n.name}' (${n.dtype}).`)}}class Ke{constructor(t){if(this.id2Value={},this.id2Mask={},this.name2Id={},t instanceof Ke)for(const e in t.id2Value)this.id2Value[e]=t.id2Value[e],e in t.id2Mask&&(this.id2Mask[e]=t.id2Mask[e]);else{if(t==null)return;for(const e of t)this.add(e.key,e.value)}}add(t,e,s){if(this.id2Value[t.id]==null)this.id2Value[t.id]=Zw(t,e),this.name2Id[t.name]=t.id,s!=null&&(this.id2Mask[t.id]=s);else throw new I(`Duplicate key: name=${t.name}, id=${t.id}`);return this}addFeed(t){this.add(t.key,t.value)}hasKey(t){return this.id2Value[t.id]!=null}names(){return Object.keys(this.name2Id)}getValue(t){if(t instanceof
|
|
3007
|
+
*/function Zw(n,t){if(n.dtype==null||n.dtype===t.dtype)return t;try{return ot(t,n.dtype)}catch{throw new I(`The dtype of the feed (${t.dtype}) can not be cast to the dtype of the key '${n.name}' (${n.dtype}).`)}}class Ke{constructor(t){if(this.id2Value={},this.id2Mask={},this.name2Id={},t instanceof Ke)for(const e in t.id2Value)this.id2Value[e]=t.id2Value[e],e in t.id2Mask&&(this.id2Mask[e]=t.id2Mask[e]);else{if(t==null)return;for(const e of t)this.add(e.key,e.value)}}add(t,e,s){if(this.id2Value[t.id]==null)this.id2Value[t.id]=Zw(t,e),this.name2Id[t.name]=t.id,s!=null&&(this.id2Mask[t.id]=s);else throw new I(`Duplicate key: name=${t.name}, id=${t.id}`);return this}addFeed(t){this.add(t.key,t.value)}hasKey(t){return this.id2Value[t.id]!=null}names(){return Object.keys(this.name2Id)}getValue(t){if(t instanceof yn){if(this.id2Value[t.id]==null)throw new I(`Nonexistent key: ${t.name}`);return this.id2Value[t.id]}else{const e=this.name2Id[t];if(e==null)throw new I(`Feed dict has no SymbolicTensor name: ${t}`);return this.id2Value[e]}}getMask(t){if(t instanceof yn){if(this.id2Value[t.id]==null)throw new I(`Nonexistent key: ${t.name}`);return this.id2Mask[t.id]}else{const e=this.name2Id[t];if(e==null)throw new I(`Feed dict has no SymbolicTensor name: ${t}`);return this.id2Mask[e]}}disposeMasks(){this.id2Mask!=null&&ut(this.id2Mask)}}const Uc=new zc,Wc=new zc;function Ns(n,t,e,s){const r=e==null?!1:e.training,o=Array.isArray(n),i=o?n:[n],a=i.map(p=>p.name),l=[],u=t.names();for(const p of a)u.indexOf(p)!==-1?l.push(t.getValue(p)):l.push(null);const c=a.join(",")+"|"+t.names().sort().join(",");let h=Uc.get(c),f;if(h==null){const p=Qw(i,t);h=p.sorted,f=p.recipientCounts,Uc.put(c,h),Wc.put(c,f)}f={},r||Object.assign(f,Wc.get(c));const d=new Ke(t);for(let p=0;p<h.length;++p){const g=h[p],m=g.sourceLayer;if(m instanceof Ts)continue;const b=[],y=[],S=[];let x=!1;for(const T of g.inputs){const P=d.getValue(T),B=d.getMask(T);b.push(P),y.push(B),B!=null&&(x=!0),r||(f[T.name]--,f[T.name]===0&&!t.hasKey(T)&&a.indexOf(T.name)===-1&&!P.isDisposed&&T.sourceLayer.stateful!==!0&&S.push(P))}x&&(e=e||{},e.mask=y[0]);const $=nt(m.apply(b,e));let E=null;m.supportsMasking&&(E=m.computeMask(b,y));const D=e1(g),_=Array.isArray(D)?D:[D];for(let T=0;T<_.length;++T){d.hasKey(_[T])||d.add(_[T],$[T],Array.isArray(E)?E[0]:E);const P=a.indexOf(_[T].name);P!==-1&&(l[P]=$[T])}r||ut(S)}return d.disposeMasks(),o?l:l[0]}function Qw(n,t){w(n!=null&&n.length>0,()=>"Expected at least one fetch, got none");let e=[],s={};if(n.length===1){const r=Gc(n[0],t);e=r.sorted,s=r.recipientMap}else{const r=new Set;for(const o of n){const{sorted:i,recipientMap:a}=Gc(o,t);for(const l of i)r.has(l.name)||(e.push(l),r.add(l.name));for(const l in a)s[l]==null&&(s[l]=new Set),a[l].forEach(u=>s[l].add(u))}}return{sorted:e,recipientCounts:t1(s)}}function t1(n){const t={};for(const e in n)t[e]=n[e].size;return t}function Gc(n,t){const e=new Set,s=[],r={};for(const a of t.names())e.add(a);const o=[],i=[];for(o.push(n);o.length>0;){const a=o[o.length-1];if(e.has(a.name)){o.pop();continue}const l=i[i.length-1]===o.length-1;if(a.inputs.length===0||l)o.pop(),s.push(a),e.add(a.name),l&&i.pop();else{i.push(o.length-1);for(const u of a.inputs)r[u.name]==null&&(r[u.name]=new Set),r[u.name].add(a.name),!e.has(u.name)&&o.push(u)}}return{sorted:s,recipientMap:r}}function e1(n){let t;if(n.sourceLayer.inboundNodes.length===1)t=n.sourceLayer.output;else{let e=null;for(let s=0;s<n.sourceLayer.inboundNodes.length;++s)for(const r of n.sourceLayer.inboundNodes[s].outputTensors)if(r.id===n.id){e=s;break}t=n.sourceLayer.getOutputAt(e)}return t}/**
|
|
3008
3008
|
* @license
|
|
3009
3009
|
* Copyright 2018 Google LLC
|
|
3010
3010
|
*
|
|
@@ -3012,7 +3012,7 @@
|
|
|
3012
3012
|
* license that can be found in the LICENSE file or at
|
|
3013
3013
|
* https://opensource.org/licenses/MIT.
|
|
3014
3014
|
* =============================================================================
|
|
3015
|
-
*/const n1=n=>{const t=Object.keys(n);if(t.length===0)return!1;const e=t[0].split("/");return!isNaN(parseInt(e[e.length-1],10))};class me extends pe{constructor(t){if(super({}),this.containerNodes=new Set,this.name=t.name,this.name==null){const y=this.getClassName().toLowerCase();this.name=di(y)}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],mn(this.inputs).length!==this.inputs.length)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)}`);mn(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=[];for(const y of this.outputs){const S=y.sourceLayer,x=y.nodeIndex,$=y.tensorIndex;this.outputLayers.push(S),this.outputLayersNodeIndices.push(x),this.outputLayersTensorIndices.push($)}for(const y of this.inputs){const S=y.sourceLayer,x=y.nodeIndex,$=y.tensorIndex;Ae(x===0,"input layer has >1 nodes"),Ae($===0,"input layer has >1 tensors"),this.inputLayers.push(S),this.inputLayersNodeIndices.push(x),this.inputLayersTensorIndices.push($)}this.inputNames=[],this.outputNames=[],this.feedInputShapes=[],this.feedInputNames=[],this.feedOutputNames=[];for(let y=0;y<this.inputLayers.length;y++){const S=this.inputLayers[y];if(!(S instanceof _s))throw new TypeError(`Input layers to a LayersModel must be InputLayer objects. Received inputs: ${t.inputs}. Input ${y} (0-based) originates from layer type ${S.getClassName()}.`);this.inputNames.push(S.name),this.feedInputShapes.push(S.batchInputShape),this.feedInputNames.push(S.name)}for(const y of this.outputLayers)this.outputNames.push(y.name);this.internalInputShapes=this.inputs.map(y=>y.shape),this.internalOutputShapes=this.outputs.map(y=>y.shape);const e={},s={},r={},o={},i={},a=[],l=(y,S,x,$,E,D)=>{($==null||E==null||D==null)&&($=y.sourceLayer,E=y.nodeIndex,D=y.tensorIndex);const _=$.inboundNodes[E];if(x.indexOf(_)!==-1)throw new He(`The tensor ${y.name} at layer "${$.name}" is part of a cycle.`);if(S.indexOf(_)!==-1)return;this.containerNodes.add(me.nodeKey($,E)),$.id in i||(i[$.id]=Object.keys(i).length),x.indexOf(_)===-1&&x.push(_);const T=_.inboundLayers.length;for(let P=0;P<T;P++){const B=_.inputTensors[P],Y=_.inboundLayers[P],j=_.nodeIndices[P],F=_.tensorIndices[P];l(B,S,x,Y,j,F)}for(S.push(_);x.indexOf(_)>=0;)x.splice(x.indexOf(_),1);a.push(_)},u=[],c=[];for(const y of this.outputs)l(y,u,c);const h=a.slice().reverse();for(const y of h){s[y.id]=y,y.id in e||(e[y.id]=0);let S=e[y.id];const x=r[y.outboundLayer.id]==null?0:r[y.outboundLayer.id];S=Math.max(S,x),r[y.outboundLayer.id]=S,o[y.outboundLayer.id]=y.outboundLayer,e[y.id]=S;for(let $=0;$<y.inboundLayers.length;$++){const E=y.inboundLayers[$],D=y.nodeIndices[$],_=E.inboundNodes[D],T=e[_.id]==null?0:e[_.id];e[_.id]=Math.max(S+1,T),s[_.id]=_}}const f={};for(const y in e){const S=e[y];S in f||(f[S]=[]),f[S].push(s[y])}const d={};for(const y in r){const S=r[y];S in d||(d[S]=[]),d[S].push(o[y])}let p=Object.keys(d).map(y=>parseInt(y,10)).sort(mr);this.layers=[];for(const y of p){const S=d[y];S.sort((x,$)=>{const E=i[x.id],D=i[$.id];return E<D?-1:E>D?1:0});for(const x of S)x instanceof me&&this.internalContainerRefs.push(x),this.layers.push(x)}this.layersByDepth=d,p=Object.keys(f).map(y=>parseInt(y,10)).sort(mr);const g=this.inputs.slice(),m=[];for(const y of p)for(const S of f[y]){const x=S.outboundLayer;if(x!=null){for(const $ of S.inputTensors)if(g.indexOf($)===-1)throw new He(`Graph disconnected: cannot obtain value for tensor ${$} at layer "${x.name}". The following previous layers were accessed without issue: ${m}`);for(const $ of S.outputTensors)g.push($);m.push(x.name)}}this.nodesByDepth=f;const b=this.layers.map(y=>y.name);for(const y of b){const S=b.filter(x=>x===y).length;if(S!==1)throw new He(`The name "${y}" is used ${S} times in the model. All layer names should be unique. Layer names: `+JSON.stringify(b))}this.outboundNodes=[],this.inboundNodes=[],new xi({outboundLayer:this,inboundLayers:[],nodeIndices:[],tensorIndices:[],inputTensors:this.inputs,outputTensors:this.outputs,inputMasks:this.inputs.map(y=>null),outputMasks:this.outputs.map(y=>null),inputShapes:this.inputs.map(y=>y.shape),outputShapes:this.outputs.map(y=>y.shape)}),this.built=!0,this._refCount=1}assertNotDisposed(){if(this._refCount===0)throw new Error(`Container '${this.name}' is already disposed.`)}dispose(){this.assertNotDisposed();const t={refCountAfterDispose:null,numDisposedVariables:0};if(--this._refCount===0){for(const e of this.layers)t.numDisposedVariables+=e.dispose().numDisposedVariables;for(const e of this.internalContainerRefs)t.numDisposedVariables+=e.dispose().numDisposedVariables}return t.refCountAfterDispose=this._refCount,t}get trainable(){return this.trainable_}set trainable(t){this.layers.forEach(e=>{e._trainableWeights.forEach(s=>s.trainable=t)}),this.trainable_=t}get trainableWeights(){if(this._trainableWeights.length>0)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.");if(!this.trainable)return[];let t=[];for(const e of this.layers)t=t.concat(e.trainableWeights);return t}get nonTrainableWeights(){const t=[];for(const e of this.layers)t.push(...e.nonTrainableWeights);if(!this.trainable){const e=[];for(const s of this.layers)e.push(...s.trainableWeights);return e.concat(t)}return t}get weights(){return this.trainableWeights.concat(this.nonTrainableWeights)}loadWeights(t,e=!0){const s={};let r=0;const o=n1(t);o&&this.parseWeights(t);for(const a of this.layers)for(const[l,u]of a.weights.entries()){const c=o?`${u.name.split("/").slice(0,-1).join("/")+"/"}${l}`:u.originalName;if(s[c]!=null)throw new I(`Duplicate weight name: ${c}`);s[c]=u,r++}const i=[];for(const a in t){let l=a;if(s[a]==null){const u=a.split("/");l=u.slice(0,-2).concat([u[u.length-1]]).join("/")}if(s[l]!=null)i.push([s[l],t[a]]);else if(e)throw new I(`Provided weight data has no target variable: ${a}`);delete s[l]}if(e){const a=[];for(const l in s)a.push(l);if(a.length>0)throw new I(`${a.length} of ${r} weights are not set: ${a}`)}ju(i)}parseWeights(t){for(const e in Object.keys(t)){const s=e.split("/"),r=["vars","layer_checkpoint_dependencies"],o=s.map(i=>i.startsWith("_")?i.slice(1):i).filter(i=>!r.includes(i)).join("/");o!==e&&(t[o]=t[e],delete t[e])}}updatedConfig(){const t=this.getConfig(),e={};return e.className=this.getClassName(),e.config=t,e.kerasVersion=`tfjs-layers ${Fc}`,e.backend="TensorFlow.js",e}toJSON(t,e=!0){const s=Ni(this.updatedConfig());return e?JSON.stringify(s):s}call(t,e){return k(()=>{t=nt(t);const s=new Ke;for(let r=0;r<this.inputs.length;++r)s.add(this.inputs[r],t[r]);return Ts(this.outputs,s,e)})}computeMask(t,e){return k(()=>{t=nt(t);let s;return e==null?s=pr(null,t.length):s=nt(e),this.runInternalGraph(t,s)[1]})}computeOutputShape(t){const e=vr(t);if(e.length!==this.inputLayers.length)throw new I(`Invalid inputShape argument ${t}: model has ${this.inputLayers.length} tensor inputs.`);const s={};for(let a=0;a<e.length;a++){const l=this.inputLayers[a],u=e[a],c=l.name+"_0_0";s[c]=u}const r=Object.keys(this.nodesByDepth).map(a=>parseInt(a,10)).sort(mr);if(r.length>1)for(const a of r){const l=this.nodesByDepth[a];for(const u of l){const c=u.outboundLayer;if(this.inputLayers.map(g=>g.id).indexOf(c.id)!==-1)continue;const h=[];for(let g=0;g<u.inboundLayers.length;g++){const m=u.inboundLayers[g],b=u.nodeIndices[g],y=u.tensorIndices[g],S=`${m.name}_${b}_${y}`,x=s[S];h.push(x)}const f=c.computeOutputShape(Ut(h)),d=vr(f),p=c.inboundNodes.indexOf(u);for(let g=0;g<d.length;g++){const m=`${c.name}_${p}_${g}`;s[m]=d[g]}}}const o=[],i=[];for(let a=0;a<this.outputLayers.length;a++){const l=this.outputLayers[a],u=this.outputLayersNodeIndices[a],c=this.outputLayersTensorIndices[a],h=`${l.name}_${u}_${c}`;i.push(h)}for(let a=0;a<i.length;a++){const l=i[a];Ae(l in s),o.push(s[l])}return Ut(o)}runInternalGraph(t,e){e==null&&(e=pr(null,t.length));const s={};for(let l=0;l<this.inputs.length;++l){const u=this.inputs[l],c=t[l],h=e[l];s[u.id]=[c,h]}const r=Object.keys(this.nodesByDepth).map(l=>parseInt(l,10)).sort(mr);for(const l of r){const u=this.nodesByDepth[l];for(const c of u){const h=c.outboundLayer,f=c.inputTensors,d=c.outputTensors,p=new Array;for(const g of f)g.id in s&&p.push(s[g.id]);if(p.length===f.length){let g={},m,b,y,S;if(c.callArgs!=null&&(g=c.callArgs),p.length===1){const[x,$]=p[0];g.mask==null&&(g.mask=$),y=nt(h.call(x,g)),S=nt(h.computeMask(x,$)),m=[x],b=[$]}else m=p.map(x=>x[0]),b=p.map(x=>x[1]),g.mask==null&&(g.mask=b),y=nt(h.call(m,g)),S=nt(h.computeMask(m,b));if(h.activityRegularizer)throw new Z("LayersModel invocation with concrete Tensor value(s) in the presence of activity regularizer(s) is not supported yet.");for(let x=0;x<d.length;++x){const $=d[x],E=y[x],D=S[x];s[$.id]=[E,D]}}}}const o=[],i=[],a=[];for(const l of this.outputs){Ae(l.id in s,`Could not compute output ${l.name} : ${l.id}`);const[u,c]=s[l.id];a.push(u.shape),o.push(u),i.push(c)}return[o,i,a]}buildNodeConversionMap(t){const e={};let s;for(const r of this.layers){s=r instanceof me?1:0;for(let o=0;o<r.inboundNodes.length;o++){const i=me.nodeKey(r,o);this.containerNodes.has(i)&&(e[i]=s,s+=1)}}return e}getLayer(t,e){if(e!=null)return this.findLayer(e);if(t==null)throw new I("Provide either a layer name or layer index");if(typeof t=="number")return this.findLayer(t);for(const s of this.layers)if(s.name===t)return s;throw new I(`No such layer: ${t}`)}findLayer(t){if(this.layers.length<=t)throw new I(`Was asked to retrieve layer at index ${t}, but model only has ${this.layers.length} layer(s).`);return this.layers[t]}calculateLosses(){return k(()=>{const t=[];for(const e of this.layers)for(let s=0;s<e.inboundNodes.length;++s){const r=me.nodeKey(e,s);this.containerNodes.has(r)&&t.push(...e.calculateLosses())}return t})}getConfig(){const t={name:this.name},e=this.buildNodeConversionMap(this.layers),s=[];for(const i of this.layers){const a=i.getClassName(),l=i.getConfig(),u=[];for(let h=0;h<i.inboundNodes.length;h++){const f=i.inboundNodes[h],d=me.nodeKey(i,h);let p={};if(this.containerNodes.has(d)){if(f.callArgs)try{JSON.stringify(f.callArgs),p=f.callArgs}catch{console.warn(`Layer ${i.name} was passed non-serializable keyword arguments: ${f.callArgs}. They will not be included in the serialized model (and thus will be missing at deserialization time).`),p={}}if(f.inboundLayers.length>0){const g=[];for(let m=0;m<f.inboundLayers.length;m++){const b=f.inboundLayers[m],y=f.nodeIndices[m],S=f.tensorIndices[m],x=me.nodeKey(b,y);let $=e[x];$==null&&($=0),g.push([b.name,$,S,p])}u.push(g)}}}const c={};c.name=i.name,c.className=a,c.config=l,c.inboundNodes=u,s.push(c)}t.layers=s;const r=[];for(let i=0;i<this.inputLayers.length;i++){const a=this.inputLayers[i],l=this.inputLayersNodeIndices[i],u=me.nodeKey(a,l);if(!this.containerNodes.has(u))continue;let c=e[u];c==null&&(c=0);const h=this.inputLayersTensorIndices[i];r.push([a.name,c,h])}t.inputLayers=r;const o=[];for(let i=0;i<this.outputLayers.length;i++){const a=this.outputLayers[i],l=this.outputLayersNodeIndices[i],u=me.nodeKey(a,l);if(!this.containerNodes.has(u))continue;let c=e[u];c==null&&(c=0);const h=this.outputLayersTensorIndices[i];o.push([a.name,c,h])}return t.outputLayers=o,t}static fromConfig(t,e,s={},r=!1){const o={},i={};function a(m,b){m.name in i?i[m.name].push(b):i[m.name]=[b]}function l(m,b){const y=[];let S;for(const x of b){const $=x[0],E=x[1],D=x[2];if(S=x[3]==null?{}:x[3],!($ in o)){a(m,b);return}const _=o[$];if(_.inboundNodes.length<=E){a(m,b);return}const T=_.inboundNodes[E];y.push(T.outputTensors[D])}y.length>0&&m.apply(Ut(y),S)}function u(m){const b=m.name,y=Nc(m,e.customObjects!=null?e.customObjects:{});y.setFastWeightInitDuringBuild(r),o[b]=y,m.inboundNodes.forEach(x=>{if(!(x instanceof Array))throw new I(`Corrupted configuration, expected array for nodeData: ${x}`);a(y,x)})}const c=e.name,h=e.layers;for(const m of h)u(m);for(;!Fy(i);)for(const m of h){const b=o[m.name];if(b.name in i){const y=i[b.name];delete i[b.name];for(const S of y)l(b,S)}}const f=[],d=[],p=e.inputLayers;for(const m of p){const b=m[0],y=m[1],S=m[2];Ae(b in o);const $=o[b].inboundNodes[y].outputTensors;f.push($[S])}const g=e.outputLayers;for(const m of g){const b=m[0],y=m[1],S=m[2];Ae(b in o);const $=o[b].inboundNodes[y].outputTensors;d.push($[S])}return new t({inputs:f,outputs:d,name:c})}get stateful(){if(this._stateful)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.");for(const t of this.layers)if(t.stateful)return!0;return!1}resetStates(){k(()=>{this.layers.forEach(t=>{t.stateful&&t.resetStates()})})}}/**
|
|
3015
|
+
*/const n1=n=>{const t=Object.keys(n);if(t.length===0)return!1;const e=t[0].split("/");return!isNaN(parseInt(e[e.length-1],10))};class me extends pe{constructor(t){if(super({}),this.containerNodes=new Set,this.name=t.name,this.name==null){const y=this.getClassName().toLowerCase();this.name=di(y)}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],gn(this.inputs).length!==this.inputs.length)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)}`);gn(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=[];for(const y of this.outputs){const S=y.sourceLayer,x=y.nodeIndex,$=y.tensorIndex;this.outputLayers.push(S),this.outputLayersNodeIndices.push(x),this.outputLayersTensorIndices.push($)}for(const y of this.inputs){const S=y.sourceLayer,x=y.nodeIndex,$=y.tensorIndex;Ae(x===0,"input layer has >1 nodes"),Ae($===0,"input layer has >1 tensors"),this.inputLayers.push(S),this.inputLayersNodeIndices.push(x),this.inputLayersTensorIndices.push($)}this.inputNames=[],this.outputNames=[],this.feedInputShapes=[],this.feedInputNames=[],this.feedOutputNames=[];for(let y=0;y<this.inputLayers.length;y++){const S=this.inputLayers[y];if(!(S instanceof Ts))throw new TypeError(`Input layers to a LayersModel must be InputLayer objects. Received inputs: ${t.inputs}. Input ${y} (0-based) originates from layer type ${S.getClassName()}.`);this.inputNames.push(S.name),this.feedInputShapes.push(S.batchInputShape),this.feedInputNames.push(S.name)}for(const y of this.outputLayers)this.outputNames.push(y.name);this.internalInputShapes=this.inputs.map(y=>y.shape),this.internalOutputShapes=this.outputs.map(y=>y.shape);const e={},s={},r={},o={},i={},a=[],l=(y,S,x,$,E,D)=>{($==null||E==null||D==null)&&($=y.sourceLayer,E=y.nodeIndex,D=y.tensorIndex);const _=$.inboundNodes[E];if(x.indexOf(_)!==-1)throw new He(`The tensor ${y.name} at layer "${$.name}" is part of a cycle.`);if(S.indexOf(_)!==-1)return;this.containerNodes.add(me.nodeKey($,E)),$.id in i||(i[$.id]=Object.keys(i).length),x.indexOf(_)===-1&&x.push(_);const T=_.inboundLayers.length;for(let P=0;P<T;P++){const B=_.inputTensors[P],Y=_.inboundLayers[P],j=_.nodeIndices[P],F=_.tensorIndices[P];l(B,S,x,Y,j,F)}for(S.push(_);x.indexOf(_)>=0;)x.splice(x.indexOf(_),1);a.push(_)},u=[],c=[];for(const y of this.outputs)l(y,u,c);const h=a.slice().reverse();for(const y of h){s[y.id]=y,y.id in e||(e[y.id]=0);let S=e[y.id];const x=r[y.outboundLayer.id]==null?0:r[y.outboundLayer.id];S=Math.max(S,x),r[y.outboundLayer.id]=S,o[y.outboundLayer.id]=y.outboundLayer,e[y.id]=S;for(let $=0;$<y.inboundLayers.length;$++){const E=y.inboundLayers[$],D=y.nodeIndices[$],_=E.inboundNodes[D],T=e[_.id]==null?0:e[_.id];e[_.id]=Math.max(S+1,T),s[_.id]=_}}const f={};for(const y in e){const S=e[y];S in f||(f[S]=[]),f[S].push(s[y])}const d={};for(const y in r){const S=r[y];S in d||(d[S]=[]),d[S].push(o[y])}let p=Object.keys(d).map(y=>parseInt(y,10)).sort(gr);this.layers=[];for(const y of p){const S=d[y];S.sort((x,$)=>{const E=i[x.id],D=i[$.id];return E<D?-1:E>D?1:0});for(const x of S)x instanceof me&&this.internalContainerRefs.push(x),this.layers.push(x)}this.layersByDepth=d,p=Object.keys(f).map(y=>parseInt(y,10)).sort(gr);const g=this.inputs.slice(),m=[];for(const y of p)for(const S of f[y]){const x=S.outboundLayer;if(x!=null){for(const $ of S.inputTensors)if(g.indexOf($)===-1)throw new He(`Graph disconnected: cannot obtain value for tensor ${$} at layer "${x.name}". The following previous layers were accessed without issue: ${m}`);for(const $ of S.outputTensors)g.push($);m.push(x.name)}}this.nodesByDepth=f;const b=this.layers.map(y=>y.name);for(const y of b){const S=b.filter(x=>x===y).length;if(S!==1)throw new He(`The name "${y}" is used ${S} times in the model. All layer names should be unique. Layer names: `+JSON.stringify(b))}this.outboundNodes=[],this.inboundNodes=[],new xi({outboundLayer:this,inboundLayers:[],nodeIndices:[],tensorIndices:[],inputTensors:this.inputs,outputTensors:this.outputs,inputMasks:this.inputs.map(y=>null),outputMasks:this.outputs.map(y=>null),inputShapes:this.inputs.map(y=>y.shape),outputShapes:this.outputs.map(y=>y.shape)}),this.built=!0,this._refCount=1}assertNotDisposed(){if(this._refCount===0)throw new Error(`Container '${this.name}' is already disposed.`)}dispose(){this.assertNotDisposed();const t={refCountAfterDispose:null,numDisposedVariables:0};if(--this._refCount===0){for(const e of this.layers)t.numDisposedVariables+=e.dispose().numDisposedVariables;for(const e of this.internalContainerRefs)t.numDisposedVariables+=e.dispose().numDisposedVariables}return t.refCountAfterDispose=this._refCount,t}get trainable(){return this.trainable_}set trainable(t){this.layers.forEach(e=>{e._trainableWeights.forEach(s=>s.trainable=t)}),this.trainable_=t}get trainableWeights(){if(this._trainableWeights.length>0)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.");if(!this.trainable)return[];let t=[];for(const e of this.layers)t=t.concat(e.trainableWeights);return t}get nonTrainableWeights(){const t=[];for(const e of this.layers)t.push(...e.nonTrainableWeights);if(!this.trainable){const e=[];for(const s of this.layers)e.push(...s.trainableWeights);return e.concat(t)}return t}get weights(){return this.trainableWeights.concat(this.nonTrainableWeights)}loadWeights(t,e=!0){const s={};let r=0;const o=n1(t);o&&this.parseWeights(t);for(const a of this.layers)for(const[l,u]of a.weights.entries()){const c=o?`${u.name.split("/").slice(0,-1).join("/")+"/"}${l}`:u.originalName;if(s[c]!=null)throw new I(`Duplicate weight name: ${c}`);s[c]=u,r++}const i=[];for(const a in t){let l=a;if(s[a]==null){const u=a.split("/");l=u.slice(0,-2).concat([u[u.length-1]]).join("/")}if(s[l]!=null)i.push([s[l],t[a]]);else if(e)throw new I(`Provided weight data has no target variable: ${a}`);delete s[l]}if(e){const a=[];for(const l in s)a.push(l);if(a.length>0)throw new I(`${a.length} of ${r} weights are not set: ${a}`)}ju(i)}parseWeights(t){for(const e in Object.keys(t)){const s=e.split("/"),r=["vars","layer_checkpoint_dependencies"],o=s.map(i=>i.startsWith("_")?i.slice(1):i).filter(i=>!r.includes(i)).join("/");o!==e&&(t[o]=t[e],delete t[e])}}updatedConfig(){const t=this.getConfig(),e={};return e.className=this.getClassName(),e.config=t,e.kerasVersion=`tfjs-layers ${Fc}`,e.backend="TensorFlow.js",e}toJSON(t,e=!0){const s=Ni(this.updatedConfig());return e?JSON.stringify(s):s}call(t,e){return k(()=>{t=nt(t);const s=new Ke;for(let r=0;r<this.inputs.length;++r)s.add(this.inputs[r],t[r]);return Ns(this.outputs,s,e)})}computeMask(t,e){return k(()=>{t=nt(t);let s;return e==null?s=mr(null,t.length):s=nt(e),this.runInternalGraph(t,s)[1]})}computeOutputShape(t){const e=Ir(t);if(e.length!==this.inputLayers.length)throw new I(`Invalid inputShape argument ${t}: model has ${this.inputLayers.length} tensor inputs.`);const s={};for(let a=0;a<e.length;a++){const l=this.inputLayers[a],u=e[a],c=l.name+"_0_0";s[c]=u}const r=Object.keys(this.nodesByDepth).map(a=>parseInt(a,10)).sort(gr);if(r.length>1)for(const a of r){const l=this.nodesByDepth[a];for(const u of l){const c=u.outboundLayer;if(this.inputLayers.map(g=>g.id).indexOf(c.id)!==-1)continue;const h=[];for(let g=0;g<u.inboundLayers.length;g++){const m=u.inboundLayers[g],b=u.nodeIndices[g],y=u.tensorIndices[g],S=`${m.name}_${b}_${y}`,x=s[S];h.push(x)}const f=c.computeOutputShape(Ut(h)),d=Ir(f),p=c.inboundNodes.indexOf(u);for(let g=0;g<d.length;g++){const m=`${c.name}_${p}_${g}`;s[m]=d[g]}}}const o=[],i=[];for(let a=0;a<this.outputLayers.length;a++){const l=this.outputLayers[a],u=this.outputLayersNodeIndices[a],c=this.outputLayersTensorIndices[a],h=`${l.name}_${u}_${c}`;i.push(h)}for(let a=0;a<i.length;a++){const l=i[a];Ae(l in s),o.push(s[l])}return Ut(o)}runInternalGraph(t,e){e==null&&(e=mr(null,t.length));const s={};for(let l=0;l<this.inputs.length;++l){const u=this.inputs[l],c=t[l],h=e[l];s[u.id]=[c,h]}const r=Object.keys(this.nodesByDepth).map(l=>parseInt(l,10)).sort(gr);for(const l of r){const u=this.nodesByDepth[l];for(const c of u){const h=c.outboundLayer,f=c.inputTensors,d=c.outputTensors,p=new Array;for(const g of f)g.id in s&&p.push(s[g.id]);if(p.length===f.length){let g={},m,b,y,S;if(c.callArgs!=null&&(g=c.callArgs),p.length===1){const[x,$]=p[0];g.mask==null&&(g.mask=$),y=nt(h.call(x,g)),S=nt(h.computeMask(x,$)),m=[x],b=[$]}else m=p.map(x=>x[0]),b=p.map(x=>x[1]),g.mask==null&&(g.mask=b),y=nt(h.call(m,g)),S=nt(h.computeMask(m,b));if(h.activityRegularizer)throw new Z("LayersModel invocation with concrete Tensor value(s) in the presence of activity regularizer(s) is not supported yet.");for(let x=0;x<d.length;++x){const $=d[x],E=y[x],D=S[x];s[$.id]=[E,D]}}}}const o=[],i=[],a=[];for(const l of this.outputs){Ae(l.id in s,`Could not compute output ${l.name} : ${l.id}`);const[u,c]=s[l.id];a.push(u.shape),o.push(u),i.push(c)}return[o,i,a]}buildNodeConversionMap(t){const e={};let s;for(const r of this.layers){s=r instanceof me?1:0;for(let o=0;o<r.inboundNodes.length;o++){const i=me.nodeKey(r,o);this.containerNodes.has(i)&&(e[i]=s,s+=1)}}return e}getLayer(t,e){if(e!=null)return this.findLayer(e);if(t==null)throw new I("Provide either a layer name or layer index");if(typeof t=="number")return this.findLayer(t);for(const s of this.layers)if(s.name===t)return s;throw new I(`No such layer: ${t}`)}findLayer(t){if(this.layers.length<=t)throw new I(`Was asked to retrieve layer at index ${t}, but model only has ${this.layers.length} layer(s).`);return this.layers[t]}calculateLosses(){return k(()=>{const t=[];for(const e of this.layers)for(let s=0;s<e.inboundNodes.length;++s){const r=me.nodeKey(e,s);this.containerNodes.has(r)&&t.push(...e.calculateLosses())}return t})}getConfig(){const t={name:this.name},e=this.buildNodeConversionMap(this.layers),s=[];for(const i of this.layers){const a=i.getClassName(),l=i.getConfig(),u=[];for(let h=0;h<i.inboundNodes.length;h++){const f=i.inboundNodes[h],d=me.nodeKey(i,h);let p={};if(this.containerNodes.has(d)){if(f.callArgs)try{JSON.stringify(f.callArgs),p=f.callArgs}catch{console.warn(`Layer ${i.name} was passed non-serializable keyword arguments: ${f.callArgs}. They will not be included in the serialized model (and thus will be missing at deserialization time).`),p={}}if(f.inboundLayers.length>0){const g=[];for(let m=0;m<f.inboundLayers.length;m++){const b=f.inboundLayers[m],y=f.nodeIndices[m],S=f.tensorIndices[m],x=me.nodeKey(b,y);let $=e[x];$==null&&($=0),g.push([b.name,$,S,p])}u.push(g)}}}const c={};c.name=i.name,c.className=a,c.config=l,c.inboundNodes=u,s.push(c)}t.layers=s;const r=[];for(let i=0;i<this.inputLayers.length;i++){const a=this.inputLayers[i],l=this.inputLayersNodeIndices[i],u=me.nodeKey(a,l);if(!this.containerNodes.has(u))continue;let c=e[u];c==null&&(c=0);const h=this.inputLayersTensorIndices[i];r.push([a.name,c,h])}t.inputLayers=r;const o=[];for(let i=0;i<this.outputLayers.length;i++){const a=this.outputLayers[i],l=this.outputLayersNodeIndices[i],u=me.nodeKey(a,l);if(!this.containerNodes.has(u))continue;let c=e[u];c==null&&(c=0);const h=this.outputLayersTensorIndices[i];o.push([a.name,c,h])}return t.outputLayers=o,t}static fromConfig(t,e,s={},r=!1){const o={},i={};function a(m,b){m.name in i?i[m.name].push(b):i[m.name]=[b]}function l(m,b){const y=[];let S;for(const x of b){const $=x[0],E=x[1],D=x[2];if(S=x[3]==null?{}:x[3],!($ in o)){a(m,b);return}const _=o[$];if(_.inboundNodes.length<=E){a(m,b);return}const T=_.inboundNodes[E];y.push(T.outputTensors[D])}y.length>0&&m.apply(Ut(y),S)}function u(m){const b=m.name,y=Nc(m,e.customObjects!=null?e.customObjects:{});y.setFastWeightInitDuringBuild(r),o[b]=y,m.inboundNodes.forEach(x=>{if(!(x instanceof Array))throw new I(`Corrupted configuration, expected array for nodeData: ${x}`);a(y,x)})}const c=e.name,h=e.layers;for(const m of h)u(m);for(;!Fy(i);)for(const m of h){const b=o[m.name];if(b.name in i){const y=i[b.name];delete i[b.name];for(const S of y)l(b,S)}}const f=[],d=[],p=e.inputLayers;for(const m of p){const b=m[0],y=m[1],S=m[2];Ae(b in o);const $=o[b].inboundNodes[y].outputTensors;f.push($[S])}const g=e.outputLayers;for(const m of g){const b=m[0],y=m[1],S=m[2];Ae(b in o);const $=o[b].inboundNodes[y].outputTensors;d.push($[S])}return new t({inputs:f,outputs:d,name:c})}get stateful(){if(this._stateful)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.");for(const t of this.layers)if(t.stateful)return!0;return!1}resetStates(){k(()=>{this.layers.forEach(t=>{t.stateful&&t.resetStates()})})}}/**
|
|
3016
3016
|
* @license
|
|
3017
3017
|
* Copyright 2018 Google LLC
|
|
3018
3018
|
*
|
|
@@ -3020,7 +3020,7 @@
|
|
|
3020
3020
|
* license that can be found in the LICENSE file or at
|
|
3021
3021
|
* https://opensource.org/licenses/MIT.
|
|
3022
3022
|
* =============================================================================
|
|
3023
|
-
*/function s1(n,t,e){const s=t.length;if(n==null||Array.isArray(n)&&n.length===0)return t.map(r=>null);if(s===1)return Array.isArray(n)&&n.length===1?n:typeof n=="object"&&t[0]in n?[n[t[0]]]:[n];if(Array.isArray(n)){if(n.length!==s)throw new Error(`Provided ${e} is an array of ${n.length} element(s), but the model has ${s} outputs. Make sure a set of weights is provided for each model output.`);return n}else if(typeof n=="object"&&Object.keys(n).length>0&&typeof n[Object.keys(n)[0]]=="object"){const r=[];return t.forEach(o=>{o in n?r.push(n[o]):r.push(null)}),r}else throw new Error(`The model has multiple (${s}) outputs, so ${e} must be either an array with ${s} elements or an object with ${t} keys. Provided ${e} not understood: ${JSON.stringify(n)}`)}function Vc(n,t){return s1(n,t,"classWeight")}async function qc(n,t,e,s){if(e!=null){const r=k(()=>{if(n.shape.length===1)return
|
|
3023
|
+
*/function s1(n,t,e){const s=t.length;if(n==null||Array.isArray(n)&&n.length===0)return t.map(r=>null);if(s===1)return Array.isArray(n)&&n.length===1?n:typeof n=="object"&&t[0]in n?[n[t[0]]]:[n];if(Array.isArray(n)){if(n.length!==s)throw new Error(`Provided ${e} is an array of ${n.length} element(s), but the model has ${s} outputs. Make sure a set of weights is provided for each model output.`);return n}else if(typeof n=="object"&&Object.keys(n).length>0&&typeof n[Object.keys(n)[0]]=="object"){const r=[];return t.forEach(o=>{o in n?r.push(n[o]):r.push(null)}),r}else throw new Error(`The model has multiple (${s}) outputs, so ${e} must be either an array with ${s} elements or an object with ${t} keys. Provided ${e} not understood: ${JSON.stringify(n)}`)}function Vc(n,t){return s1(n,t,"classWeight")}async function qc(n,t,e,s){if(e!=null){const r=k(()=>{if(n.shape.length===1)return on(n);if(n.shape.length===2){if(n.shape[1]>1)return nr(n,1);if(n.shape[1]===1)return L(n,[n.shape[0]]);throw new Error(`Encountered unexpected last-dimension size (${n.shape[1]}) during handling of class weights. The size is expected to be >= 1.`)}else throw new Error(`Unexpected rank of target (y) tensor (${n.rank}) during handling of class weights. The rank is expected to be 1 or 2.`)}),o=Array.from(await r.data());ut(r);const i=[];return o.forEach(a=>{if(e[a]==null)throw new Error(`classWeight must contain all classes in the training data. The class ${a} exists in the data but not in classWeight`);i.push(e[a])}),Rt(i,"float32")}else return null}function r1(n,t){return N(n,t)}/**
|
|
3024
3024
|
* @license
|
|
3025
3025
|
* Copyright 2018 Google LLC
|
|
3026
3026
|
*
|
|
@@ -3028,7 +3028,7 @@
|
|
|
3028
3028
|
* license that can be found in the LICENSE file or at
|
|
3029
3029
|
* https://opensource.org/licenses/MIT.
|
|
3030
3030
|
* =============================================================================
|
|
3031
|
-
*/const o1=32;function jc(n,t){let e,s;const r=t;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}`);const o=Hc("input",n.inputNames,e),i=Hc("output",n.outputNames,s),a=o[0].shape[0];w(o.length===n.inputs.length,()=>`LayersModel has ${n.inputs.length} inputs, but the dataset provides ${o.length} inputs. (Expected input keys: ${JSON.stringify(n.inputNames)})`),w(i.length===n.outputs.length,()=>`LayersModel has ${n.outputs.length} outputs, but the dataset provides ${i.length} outputs. (Expected output keys: ${JSON.stringify(n.outputNames)})`);for(let l=0;l<o.length;l++)w(o[l].shape[0]===a,()=>`Batch size mismatch: input ${n.inputNames[l]} has ${o[l].shape[0]}; expected ${a} based on input ${n.inputNames[0]}.`);for(let l=0;l<i.length;l++)w(i[l].shape[0]===a,()=>`Batch size mismatch: output ${n.outputNames[l]} has ${i[l].shape[0]}; expected ${a} based on input ${n.inputNames[0]}.`);return{xs:o,ys:i}}function Hc(n,t,e){if(e instanceof At)return[e];if(Array.isArray(e))return w(e.length===t.length,()=>`Received an array of ${e.length} Tensors, but expected ${t.length} to match the ${n} keys ${t}.`),e;{const s=[];for(const r of t){if(e[r]==null)throw new I(`The feature data generated by the dataset lacks the required ${n} key '${r}'.`);s.push(e[r])}return s}}function i1(n){if(n.length===3)throw new Z("Validation with sample weights is not implemented yet.");return{xs:n[0],ys:n[1]}}async function a1(n,t,e){const s=e.batchesPerEpoch!=null;if(w(n.optimizer!=null,()=>"You must compile a model before training/testing. Use LayersModel.compile(modelCompileConfig)."),w(e!=null,()=>"For fitDataset(), the 2nd argument (config) is required, but it is not provided in this call."),w(e.epochs!=null&&e.epochs>0&&Number.isInteger(e.epochs),()=>`For fitDataset(), config.epochs is expected to be a positive integer, but got ${e.epochs}`),w(!s||e.batchesPerEpoch>0&&Number.isInteger(e.batchesPerEpoch),()=>`For fitDataset(), config.batchesPerEpoch is expected to be a positive integer if specified, but got ${e.batchesPerEpoch}`),w(e.validationSplit==null,()=>"`validationSplit` is not supported by `fitDataset()`. Use validationData instead."),n.isTraining)throw new Error("Cannot start training because another fit() call is ongoing.");n.isTraining=!0;try{const r=e.validationData!=null;let o,i;if(r)if(Kc(e.validationData))w(e.validationBatches==null||e.validationBatches>0&&Number.isInteger(e.validationBatches),()=>`For fitDataset() with dataset-based validation, config.validationBatches is expected not to be provided, or to be a positive integer, but got ${e.validationBatches}`);else{const m=i1(e.validationData);o=m.xs,i=m.ys}const a=n.makeTrainFunction(),l=n.getDedupedMetricsNames();let u;r?u=l.slice().concat(l.map(m=>"val_"+m)):u=l.slice();const c=_c(e.callbacks,e.yieldEvery),h=e.verbose==null?1:e.verbose,{callbackList:f,history:d}=Tc(c,h,e.epochs,null,null,l1(t,e),null,r,u);f.setModel(n),n.history=d,await f.onTrainBegin(),n.stopTraining_=!1;let p=e.initialEpoch==null?0:e.initialEpoch,g=await t.iterator();for(;p<e.epochs;){const m={};await f.onEpochBegin(p);let b=0,y=0;for(s||(g=await t.iterator());!s||b<e.batchesPerEpoch;){const S=await g.next();if(s&&S.done){console.warn(`You provided \`batchesPerEpoch\` as ${e.batchesPerEpoch}, but your dataset iterator ran out of data after ${b} batches; interrupting training. Make sure that your dataset can generate at least \`batchesPerEpoch * epochs\` batches (in this case, ${e.batchesPerEpoch*e.epochs} batches). You may need to use the repeat() function when building your dataset.`);break}if(S.value!=null){const{xs:x,ys:$}=jc(n,S.value),E={};E.batch=y,E.size=x[0].shape[0],await f.onBatchBegin(y,E);const D=[];if(e.classWeight!=null){const P=Vc(e.classWeight,n.outputNames);for(let B=0;B<P.length;++B)D.push(await qc($[B],null,P[B]))}const _=x.concat($).concat(D),T=a(_);ut(_);for(let P=0;P<l.length;++P){const B=l[P],Y=T[P];E[B]=Y,Ln(Y)}await f.onBatchEnd(y,E),Cc(E),y++,b++}if(s?b>=e.batchesPerEpoch:S.done){if(r){let x;Kc(e.validationData)?x=nt(await n.evaluateDataset(e.validationData,{batches:e.validationBatches})):x=nt(n.evaluate(o,i,{batchSize:e.validationBatchSize==null?o1:e.validationBatchSize,verbose:0}));for(let $=0;$<n.metricsNames.length;++$)m[`val_${n.metricsNames[$]}`]=x[$]}break}if(n.stopTraining_)break}if(await f.onEpochEnd(p,m),p++,n.stopTraining_)break}return await f.onTrainEnd(),await n.history.syncData(),n.history}finally{n.isTraining=!1}}function l1(n,t){let e=null;return t.batchesPerEpoch!=null?e=t.batchesPerEpoch:Number.isFinite(n.size)&&(e=n.size),e}function Kc(n){return typeof n.iterator=="function"}function u1(n){return typeof n.next=="function"}async function c1(n,t,e){e=e||{};const s=e.batches!=null,r=n.testFunction;let o=[];if(e.verbose>0)throw new Z("Verbose mode is not implemented yet.");w(!s||e.batches>0&&Number.isInteger(e.batches),()=>`Test loop expects \`batches\` to be a positive integer, but received ${JSON.stringify(e.batches)}`);const i=u1(t)?t:await t.iterator();let a=0,l=0;for(;!s||l<e.batches;){const u=await i.next();if(o=k(()=>{if(u.value){const{xs:c,ys:h}=jc(n,u.value),f=c.concat(h),d=k(()=>r(f));if(ut(f),l===0)for(let g=0;g<d.length;++g)o.push(
|
|
3031
|
+
*/const o1=32;function jc(n,t){let e,s;const r=t;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}`);const o=Hc("input",n.inputNames,e),i=Hc("output",n.outputNames,s),a=o[0].shape[0];w(o.length===n.inputs.length,()=>`LayersModel has ${n.inputs.length} inputs, but the dataset provides ${o.length} inputs. (Expected input keys: ${JSON.stringify(n.inputNames)})`),w(i.length===n.outputs.length,()=>`LayersModel has ${n.outputs.length} outputs, but the dataset provides ${i.length} outputs. (Expected output keys: ${JSON.stringify(n.outputNames)})`);for(let l=0;l<o.length;l++)w(o[l].shape[0]===a,()=>`Batch size mismatch: input ${n.inputNames[l]} has ${o[l].shape[0]}; expected ${a} based on input ${n.inputNames[0]}.`);for(let l=0;l<i.length;l++)w(i[l].shape[0]===a,()=>`Batch size mismatch: output ${n.outputNames[l]} has ${i[l].shape[0]}; expected ${a} based on input ${n.inputNames[0]}.`);return{xs:o,ys:i}}function Hc(n,t,e){if(e instanceof At)return[e];if(Array.isArray(e))return w(e.length===t.length,()=>`Received an array of ${e.length} Tensors, but expected ${t.length} to match the ${n} keys ${t}.`),e;{const s=[];for(const r of t){if(e[r]==null)throw new I(`The feature data generated by the dataset lacks the required ${n} key '${r}'.`);s.push(e[r])}return s}}function i1(n){if(n.length===3)throw new Z("Validation with sample weights is not implemented yet.");return{xs:n[0],ys:n[1]}}async function a1(n,t,e){const s=e.batchesPerEpoch!=null;if(w(n.optimizer!=null,()=>"You must compile a model before training/testing. Use LayersModel.compile(modelCompileConfig)."),w(e!=null,()=>"For fitDataset(), the 2nd argument (config) is required, but it is not provided in this call."),w(e.epochs!=null&&e.epochs>0&&Number.isInteger(e.epochs),()=>`For fitDataset(), config.epochs is expected to be a positive integer, but got ${e.epochs}`),w(!s||e.batchesPerEpoch>0&&Number.isInteger(e.batchesPerEpoch),()=>`For fitDataset(), config.batchesPerEpoch is expected to be a positive integer if specified, but got ${e.batchesPerEpoch}`),w(e.validationSplit==null,()=>"`validationSplit` is not supported by `fitDataset()`. Use validationData instead."),n.isTraining)throw new Error("Cannot start training because another fit() call is ongoing.");n.isTraining=!0;try{const r=e.validationData!=null;let o,i;if(r)if(Kc(e.validationData))w(e.validationBatches==null||e.validationBatches>0&&Number.isInteger(e.validationBatches),()=>`For fitDataset() with dataset-based validation, config.validationBatches is expected not to be provided, or to be a positive integer, but got ${e.validationBatches}`);else{const m=i1(e.validationData);o=m.xs,i=m.ys}const a=n.makeTrainFunction(),l=n.getDedupedMetricsNames();let u;r?u=l.slice().concat(l.map(m=>"val_"+m)):u=l.slice();const c=_c(e.callbacks,e.yieldEvery),h=e.verbose==null?1:e.verbose,{callbackList:f,history:d}=Tc(c,h,e.epochs,null,null,l1(t,e),null,r,u);f.setModel(n),n.history=d,await f.onTrainBegin(),n.stopTraining_=!1;let p=e.initialEpoch==null?0:e.initialEpoch,g=await t.iterator();for(;p<e.epochs;){const m={};await f.onEpochBegin(p);let b=0,y=0;for(s||(g=await t.iterator());!s||b<e.batchesPerEpoch;){const S=await g.next();if(s&&S.done){console.warn(`You provided \`batchesPerEpoch\` as ${e.batchesPerEpoch}, but your dataset iterator ran out of data after ${b} batches; interrupting training. Make sure that your dataset can generate at least \`batchesPerEpoch * epochs\` batches (in this case, ${e.batchesPerEpoch*e.epochs} batches). You may need to use the repeat() function when building your dataset.`);break}if(S.value!=null){const{xs:x,ys:$}=jc(n,S.value),E={};E.batch=y,E.size=x[0].shape[0],await f.onBatchBegin(y,E);const D=[];if(e.classWeight!=null){const P=Vc(e.classWeight,n.outputNames);for(let B=0;B<P.length;++B)D.push(await qc($[B],null,P[B]))}const _=x.concat($).concat(D),T=a(_);ut(_);for(let P=0;P<l.length;++P){const B=l[P],Y=T[P];E[B]=Y,Ln(Y)}await f.onBatchEnd(y,E),Cc(E),y++,b++}if(s?b>=e.batchesPerEpoch:S.done){if(r){let x;Kc(e.validationData)?x=nt(await n.evaluateDataset(e.validationData,{batches:e.validationBatches})):x=nt(n.evaluate(o,i,{batchSize:e.validationBatchSize==null?o1:e.validationBatchSize,verbose:0}));for(let $=0;$<n.metricsNames.length;++$)m[`val_${n.metricsNames[$]}`]=x[$]}break}if(n.stopTraining_)break}if(await f.onEpochEnd(p,m),p++,n.stopTraining_)break}return await f.onTrainEnd(),await n.history.syncData(),n.history}finally{n.isTraining=!1}}function l1(n,t){let e=null;return t.batchesPerEpoch!=null?e=t.batchesPerEpoch:Number.isFinite(n.size)&&(e=n.size),e}function Kc(n){return typeof n.iterator=="function"}function u1(n){return typeof n.next=="function"}async function c1(n,t,e){e=e||{};const s=e.batches!=null,r=n.testFunction;let o=[];if(e.verbose>0)throw new Z("Verbose mode is not implemented yet.");w(!s||e.batches>0&&Number.isInteger(e.batches),()=>`Test loop expects \`batches\` to be a positive integer, but received ${JSON.stringify(e.batches)}`);const i=u1(t)?t:await t.iterator();let a=0,l=0;for(;!s||l<e.batches;){const u=await i.next();if(o=k(()=>{if(u.value){const{xs:c,ys:h}=jc(n,u.value),f=c.concat(h),d=k(()=>r(f));if(ut(f),l===0)for(let g=0;g<d.length;++g)o.push(Yt(0));const p=f[0].shape[0];for(let g=0;g<d.length;++g){const m=d[g],b=o[g];o[g]=k(()=>O(o[g],N(p,m))),l>0&&ut(b)}ut(d),a+=p,++l}return o}),u.done){s&&console.warn(`Your dataset iterator ran out of data during evaluateDataset(). Interrupting evalution. Make sure that your dataset can generate at least \`batches\` batches (in this case, ${e.batches} batches). You may need to use the repeat() function when building your dataset.`);break}}for(let u=0;u<o.length;++u){const c=o[u];o[u]=X(o[u],a),ut(c)}return Ut(o)}/**
|
|
3032
3032
|
* @license
|
|
3033
3033
|
* Copyright 2018 Google LLC
|
|
3034
3034
|
*
|
|
@@ -3036,7 +3036,7 @@
|
|
|
3036
3036
|
* license that can be found in the LICENSE file or at
|
|
3037
3037
|
* https://opensource.org/licenses/MIT.
|
|
3038
3038
|
* =============================================================================
|
|
3039
|
-
*/function Di(n){w(n>0&&Number.isInteger(n),()=>`batchSize is required to be a positive integer, but got ${n}`)}function
|
|
3039
|
+
*/function Di(n){w(n>0&&Number.isInteger(n),()=>`batchSize is required to be a positive integer, but got ${n}`)}function Ds(n,t,e){return n==null?[null]:Array.isArray(n)?n.map(s=>bn(s,t,e-t)):bn(n,t,e-t)}function Ri(n,t){return k(()=>n==null?null:Array.isArray(n)?n.map(e=>Ri(e,t)):Ky(n,t.dtype==="int32"?t:ot(t,"int32")))}function Pi(n,t){const e=[];let s=0,r=null;for(;s<n;)r=s+t,r>=n&&(r=n),e.push([s,r]),s=r;return e}function Yc(n){const t=[];n instanceof At&&(n=[n]);for(let e=0;e<n.length;++e){const s=n[e];if(s.rank===1)t.push(ui(s,1));else{if(s.rank===0)throw new Error("Expected tensor to be at least 1D, but received a 0D tensor (scalar).");t.push(s)}}return t}function ge(n,t){if(n==null)return;const e=[];if(t instanceof At)e.push(t.id);else if(Array.isArray(t))t.forEach(r=>e.push(r.id));else if(t!=null)for(const r in t){const o=t[r];e.push(o.id)}const s=[];if(n instanceof At)e.indexOf(n.id)===-1&&s.push(n);else if(Array.isArray(n))n.forEach(r=>{e.indexOf(r.id)===-1&&s.push(r)});else if(n!=null)for(const r in n){const o=n[r];e.indexOf(o.id)===-1&&s.push(o)}s.forEach(r=>{r.isDisposed||r.dispose()})}/**
|
|
3040
3040
|
* @license
|
|
3041
3041
|
* Copyright 2018 Google LLC
|
|
3042
3042
|
*
|
|
@@ -3044,7 +3044,7 @@
|
|
|
3044
3044
|
* license that can be found in the LICENSE file or at
|
|
3045
3045
|
* https://opensource.org/licenses/MIT.
|
|
3046
3046
|
* =============================================================================
|
|
3047
|
-
*/function h1(n){return n instanceof At}function Li(n){return Array.isArray(n)}function Xc(n){return!h1(n)&&!Li(n)}function Jc(n,t,e,s=!0,r=""){if(t==null||t.length===0){if(n!=null){let i=!1;if(Li(n)&&n.length>0)i=!0;else if(Xc(n)){for(const a in n)if(n.hasOwnProperty(a)){i=!0;break}}else i=!0;if(i)throw new I(`Error when checking model ${r} expected no data, but got ${n}`)}return[]}if(n==null)return t.map(i=>null);let o;if(Xc(n)){n=n,o=[];for(const i of t){if(n[i]==null)throw new I(`No data provided for "${i}". Need data for each key in: ${t}`);o.push(n[i])}}else if(Li(n)){if(n=n,n.length!==t.length)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}`);o=n}else{if(n=n,t.length>1)throw new I(`The model ${r} expects ${t.length} Tensor(s), but only received one Tensor. Found: Tensor with shape ${n.shape}`);o=[n]}if(o=Yc(o),e!=null)for(let i=0;i<t.length;++i){if(e[i]==null)continue;const a=o[i];if(a.shape.length!==e[i].length)throw new I(`Error when checking ${r}: expected ${t[i]} to have ${e[i].length} dimension(s). but got array with shape ${a.shape}`);for(let l=0;l<e[i].length;++l){if(l===0&&!s)continue;const u=a.shape[l],c=e[i][l];if(c!=null&&c>=0&&u!==c)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}])`)}}return o}function f1(n,t,e){const s=mn(n.map(o=>o.shape[0]));s.sort();const r=mn(t.map(o=>o.shape[0]));if(r.sort(),s.length>1)throw new I(`All input Tensors (x) should have the same number of samples. Got array shapes: ${JSON.stringify(n.map(o=>o.shape))}`);if(r.length>1)throw new I(`All target Tensors (y) should have the same number of samples. Got array shapes: ${JSON.stringify(t.map(o=>o.shape))}`);if(s.length>0&&r.length>0&&!Zt(s,r))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).`)}function d1(n,t,e){const s=[kr,Tr,Es];for(let r=0;r<n.length;++r){const o=n[r],i=t[r],a=e[r];if(i!=null){if(i===Es&&o.shape[o.shape.length-1]===1)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].`);if(s.indexOf(i)!==-1){const l=o.shape.slice(1),u=a.slice(1);for(let c=0;c<l.length;++c){const h=l[c],f=u[c];if(f!=null&&h!==f)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.`)}}}}}function Zc(n,t,e,s=!0,r=""){let o;if(Array.isArray(n)){if(n.length!==t.length)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).`);o=n}else{if(t.length>1)throw new I(`The model expects ${t.length} ${r} Tensors, but only received one Tensor. Found: array with shape ${JSON.stringify(n.shape)}.`);o=[n]}if(e!=null)for(let i=0;i<t.length;++i){if(e[i]==null)continue;const a=o[i];if(a.shape.length!==e[i].length)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)}`);for(let l=0;l<e[i].length;++l){if(l===0&&!s)continue;const u=a.shape[l],c=e[i][l];if(c!=null&&c!==u)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)}.`)}}}function p1(n,t){if(n==null||Array.isArray(n)&&n.length===0)return t.map(s=>[]);let e;if(typeof n=="string"||typeof n=="function")e=[n];else if(Array.isArray(n)||typeof n=="object")e=n;else throw new TypeError(`Type of metrics argument not understood. Expected an string,function, Array, or Object, found: ${n}`);if(Array.isArray(e))return t.map(s=>e);{const s=[];for(const r of t){let o=e.hasOwnProperty(r)?e[r]:[];Array.isArray(o)||(o=[o]),s.push(o)}return s}}const m1="layers-model";class Lr extends me{constructor(t){super(t),this.isTraining=!1}summary(t,e,s=console.log){if(!this.built)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).");jw(this,t,e,s)}compile(t){if(t.loss==null&&(t.loss=[]),this.loss=t.loss,typeof t.optimizer=="string")this.optimizer_=qw(t.optimizer),this.isOptimizerOwned=!0;else{if(!(t.optimizer instanceof qe))throw new I("User-defined optimizer must be an instance of tf.Optimizer.");this.optimizer_=t.optimizer,this.isOptimizerOwned=!1}let e=[];if(!Array.isArray(t.loss)&&typeof t.loss!="string"&&typeof t.loss!="function"){t.loss=t.loss;for(const i in t.loss)if(this.outputNames.indexOf(i)===-1)throw new I(`Unknown entry in loss dictionary: "${i}". Only expected the following keys: ${this.outputNames}`);for(const i of this.outputNames)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(Ci(t.loss[i]))}else if(Array.isArray(t.loss)){if(t.loss.length!==this.outputs.length)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}.`);e=t.loss.map(a=>Ci(a))}else{const i=Ci(t.loss);this.outputs.forEach(a=>{e.push(i)})}this.lossFunctions=e,this.feedOutputNames=[],this.feedOutputShapes=[],this.feedLossFns=[];for(let i=0;i<this.outputs.length;++i){const a=this.internalOutputShapes[i],l=this.outputNames[i];this.feedOutputNames.push(l),this.feedOutputShapes.push(a),this.feedLossFns.push(this.lossFunctions[i])}const s=[];this.metrics=t.metrics,this.metricsNames=["loss"],this.metricsTensors=[],gr("loss",()=>{for(let i=0;i<this.outputs.length;++i){if(s.indexOf(i)!==-1)continue;const a=this.lossFunctions[i];this.outputs.length>1&&(this.metricsTensors.push([a,i]),this.metricsNames.push(this.outputNames[i]+"_loss"))}});const r=p1(t.metrics,this.outputNames),o=(i,a,l)=>{this.outputNames.length>1&&(a=this.outputNames[i]+"_"+a),this.metricsNames.push(a),this.metricsTensors.push([l,i])};gr("metric",()=>{for(let i=0;i<this.outputs.length;++i){if(s.indexOf(i)!==-1)continue;const a=r[i];(u=>{const c="";let h,f,d;for(const p of u){if(typeof p=="string"&&["accuracy","acc","crossentropy","ce"].indexOf(p)!==-1){const m=this.internalOutputShapes[i];m[m.length-1]===1||this.lossFunctions[i]===Tr?["accuracy","acc"].indexOf(p)!==-1?f=Dc:["crossentropy","ce"].indexOf(p)!==-1&&(f=Lw):this.lossFunctions[i]===_r?["accuracy","acc"].indexOf(p)!==-1?f=Mw:["crossentropy","ce"].indexOf(p)!==-1&&(f=Lc):["accuracy","acc"].indexOf(p)!==-1?f=Rc:["crossentropy","ce"].indexOf(p)!==-1&&(f=Pc);let b;["accuracy","acc"].indexOf(p)!==-1?b="acc":["crossentropy","ce"].indexOf(p)!==-1&&(b="ce"),d=f,h=c+b}else d=Vw(p),h=c+Rr(p);let g;gr(h,()=>{g=d}),o(i,h,g)}})(a)}}),this.collectedTrainableWeights=this.trainableWeights}checkTrainableWeightsConsistency(){this.collectedTrainableWeights!=null&&this.trainableWeights.length!==this.collectedTrainableWeights.length&&console.warn("Discrepancy between trainableweights and collected trainable weights. Did you set `model.trainable` without calling `model.compile()` afterwards?")}evaluate(t,e,s={}){const r=s.batchSize==null?32:s.batchSize;Di(r);const i=this.standardizeUserDataXY(t,e,!0,r);try{const a=i[0].concat(i[1]);this.makeTestFunction();const l=this.testFunction,u=this.testLoop(l,a,r,s.verbose,s.steps);return Ut(u)}finally{ge(i[0],t),ge(i[1],e)}}async evaluateDataset(t,e){return this.makeTestFunction(),c1(this,t,e)}checkNumSamples(t,e,s,r="steps"){let o;if(s!=null){if(o=null,e!=null)throw new I(`If ${r} is set, batchSize must be null or undefined.Got batchSize = ${e}`)}else if(t!=null)Array.isArray(t)?o=t[0].shape[0]:o=t.shape[0];else throw new I(`Either the input data should have a defined shape, or ${r} shoud be specified.`);return o}execute(t,e){if(Array.isArray(e)&&e.length===0)throw new I("`outputs` is an empty Array, which is not allowed.");const s=Array.isArray(e),r=s?e:[e],o=this.retrieveSymbolicTensors(r),i=new Ke;if(t instanceof At&&(t=[t]),Array.isArray(t)){if(t.length!==this.inputs.length)throw new I(`The number of inputs provided (${t.length}) does not match the number of inputs of this model (${this.inputs.length}).`);for(let l=0;l<this.inputs.length;++l)i.add(this.inputs[l],t[l])}else for(const l of this.inputs){const u=t[l.name];if(u==null)throw new I(`No value is provided for the model's input ${l.name}`);i.add(l,u)}const a=Ts(o,i);return s?a:a[0]}retrieveSymbolicTensors(t){const e=pr(null,t.length);let s=t.length;for(const r of this.layers){const o=Array.isArray(r.output)?r.output:[r.output],i=o.map(a=>a.name);for(let a=0;a<t.length;++a){const l=i.indexOf(t[a]);if(l!==-1&&(e[a]=o[l],s--),s===0)break}if(s===0)break}if(s>0){const r=[];throw e.forEach((o,i)=>{o==null&&r.push(t[i])}),new I(`Cannot find SymbolicTensors for output name(s): ${JSON.stringify(r)}`)}return e}predictLoop(t,e=32,s=!1){return k(()=>{const r=this.checkNumSamples(t);if(s)throw new Z("Verbose predictLoop() is not implemented yet.");const o=Pi(r,e),i=this.outputs.map(a=>[]);for(let a=0;a<o.length;++a)k(()=>{const u=o[a][0],c=o[a][1],h=Ns(t,u,c),f=[];if(Array.isArray(h))for(let p=0;p<h.length;++p)f.push({key:this.inputs[p],value:h[p]});else f.push({key:this.inputs[0],value:h});const d=new Ke(f);return Ts(this.outputs,d)}).forEach((u,c)=>i[c].push(u));return Ut(i.map(a=>on(a,0)))})}predict(t,e={}){const s=Yc(t);Zc(s,this.inputNames,this.feedInputShapes,!1);try{const r=e.batchSize==null?32:e.batchSize;return Di(r),this.predictLoop(s,r)}finally{ge(s,t)}}predictOnBatch(t){Zc(t,this.inputNames,this.feedInputShapes,!0);const e=(Array.isArray(t)?t[0]:t).shape[0];return this.predictLoop(t,e)}standardizeUserDataXY(t,e,s=!0,r){if(this.optimizer_==null)throw new He("You must compile a model before training/testing. Use LayersModel.compile(modelCompileArgs).");const o=[];for(let i=0;i<this.feedOutputShapes.length;++i){const a=this.feedOutputShapes[i];this.feedLossFns[i]===_r?o.push(a.slice(0,a.length-1).concat([1])):o.push(a)}if(t=Jc(t,this.feedInputNames,this.feedInputShapes,!1,"input"),e=Jc(e,this.feedOutputNames,o,!1,"target"),f1(t,e),d1(e,this.feedLossFns,this.feedOutputShapes),this.stateful&&r!=null&&r>0&&t[0].shape[0]%r!==0)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).`);return[t,e]}async standardizeUserData(t,e,s,r,o=!0,i){const[a,l]=this.standardizeUserDataXY(t,e,o,i);if(s!=null)throw new Error("sample weight is not supported yet.");let u=null;if(r!=null){const c=Vc(r,this.outputNames);u=[];for(let h=0;h<c.length;++h)u.push(await qc(l[h],null,c[h]))}return[a,l,u]}testLoop(t,e,s,r=0,o){return k(()=>{const i=this.checkNumSamples(e,s,o,"steps"),a=[];if(r>0)throw new Z("Verbose mode is not implemented yet.");if(o!=null)throw new Z("steps mode in testLoop() is not implemented yet");{const l=Pi(i,s),u=Rt(br(0,i));for(let c=0;c<l.length;++c){const h=l[c][0],f=l[c][1],d=gn(u,h,f-h),p=Ri(e,d),g=t(p);if(c===0)for(let m=0;m<g.length;++m)a.push(Ht(0));for(let m=0;m<g.length;++m){const b=g[m];a[m]=O(a[m],N(f-h,b))}}for(let c=0;c<a.length;++c)a[c]=X(a[c],i)}return a})}getDedupedMetricsNames(){const t=this.metricsNames,e=[];for(let s=0;s<t.length;++s){const r=t[s];let o=r;if(eu(t,r)>1){const i=eu(t.slice(0,s),r);o+=`_${i}`}e.push(o)}return e}makeTrainFunction(){return t=>{const e=[],s=t.slice(0,this.inputs.length),r=t.slice(this.inputs.length,this.inputs.length+this.outputs.length),o=t.slice(this.inputs.length+this.outputs.length,this.inputs.length+this.outputs.length*2),i=[],a=()=>{const h=[];for(let g=0;g<this.inputs.length;++g)h.push({key:this.inputs[g],value:s[g]});const f=new Ke(h),d=Ts(this.outputs,f,{training:!0});let p;for(let g=0;g<this.lossFunctions.length;++g){const m=this.lossFunctions[g];let b=m(r[g],d[g]);o[g]!=null&&(b=r1(b,o[g]));const y=St(b);e.push(y),g===0?p=b:p=O(p,b)}for(let g=0;g<this.metricsTensors.length;++g){let m;if(this.outputs.length>1&&g<this.outputs.length)m=e[g];else{const b=this.metricsTensors[g][0],y=this.metricsTensors[g][1];m=St(b(r[y],d[y]))}Ln(m),i.push(m)}return p=St(p),this.calculateLosses().forEach(g=>{p=O(p,g)}),p},l=this.collectedTrainableWeights.map(h=>h.read());return[this.optimizer_.minimize(a,!0,l)].concat(i)}}makeTestFunction(){this.testFunction=t=>k(()=>{const e=[];let s;const r=t.slice(0,this.inputs.length),o=t.slice(this.inputs.length,this.inputs.length+this.outputs.length),i=[];for(let u=0;u<this.inputs.length;++u)i.push({key:this.inputs[u],value:r[u]});const a=new Ke(i),l=Ts(this.outputs,a);for(let u=0;u<this.lossFunctions.length;++u){const c=this.lossFunctions[u],h=St(c(o[u],l[u]));u===0?s=h:s=O(s,h),e.push(s)}for(let u=0;u<this.metricsTensors.length;++u){const c=this.metricsTensors[u][0],h=this.metricsTensors[u][1],f=St(c(o[h],l[h]));e.push(f)}return e})}async fit(t,e,s={}){if(this.isTraining)throw new Error("Cannot start training because another fit() call is ongoing.");this.isTraining=!0;let r,o,i,a,l,u,c,h,f;try{const d=s.batchSize==null?32:s.batchSize;Di(d);const g=await this.standardizeUserData(t,e,s.sampleWeight,s.classWeight,!1,d);r=g[0],o=g[1],f=g[2];let m=!1,b;if(s.validationData!=null&&s.validationData.length>0){if(m=!0,s.validationData.length===2)l=s.validationData[0],u=s.validationData[1];else throw s.validationData.length===3?new Z("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.`);const P=await this.standardizeUserData(l,u,null,null,!0,d);c=P[0],h=P[1],b=c.concat(h)}else if(s.validationSplit!=null&&s.validationSplit>0&&s.validationSplit<1){m=!0;const T=Math.floor(r[0].shape[0]*(1-s.validationSplit)),P=r[0].shape[0];c=Ns(r,T,P),i=r,r=Ns(r,0,T),h=Ns(o,T,P),a=o,o=Ns(o,0,T),b=c.concat(h)}else s.validationSteps!=null&&(m=!0);const y=r.concat(o).concat(f);this.checkTrainableWeightsConsistency();const S=this.makeTrainFunction(),x=this.getDedupedMetricsNames();let $,E;m?(this.makeTestFunction(),$=this.testFunction,E=x.slice().concat(x.map(T=>"val_"+T))):($=null,b=[],E=x.slice());const D=_c(s.callbacks,s.yieldEvery);return await this.fitLoop(S,y,x,d,s.epochs,s.verbose,D,$,b,s.shuffle,E,s.initialEpoch,null,null)}finally{this.isTraining=!1,ge(r,t),ge(o,e),ge(i,t),ge(a,e),ge(c,l),ge(h,u),f!=null&&ut(f)}}async fitLoop(t,e,s,r,o,i,a,l,u,c,h,f,d,p){r==null&&(r=32),o==null&&(o=1),c==null&&(c=!0),f==null&&(f=0);let g=!1;if(l!=null&&u!=null&&(g=!0),p!=null&&(g=!0,d==null))throw new I("Can only use `validationSteps` when doing step-wise training, i.e., `stepsPerEpoch` must be set.");const m=this.checkNumSamples(e,r,d,"steps_per_epoch");let b;m!=null&&(b=br(0,m)),i==null&&(i=1);const{callbackList:y,history:S}=Tc(a,i,o,f,m,d,r,g,h);y.setModel(this),this.history=S,await y.onTrainBegin(),this.stopTraining_=!1;for(let x=f;x<o;++x){await y.onEpochBegin(x);const $={};if(d!=null)throw new Z("stepsPerEpoch mode is not implemented yet.");{if(c==="batch")throw new Z("batch shuffling is not implemneted yet");c&&If(b);const E=Rt(b),D=Pi(m,r);for(let _=0;_<D.length;++_){const T={};if(await y.onBatchBegin(_,T),k(()=>{const P=D[_][0],B=D[_][1],Y=gn(E,P,B-P);T.batch=_,T.size=B-P;const j=Ri(e,Y),F=t(j);for(let G=0;G<s.length;++G){const q=s[G],De=F[G];T[q]=De,Ln(De)}if(_===D.length-1&&g){const G=this.testLoop(l,u,r);for(let q=0;q<s.length;++q){const De=s[q],Xt=G[q];Ln(Xt),$["val_"+De]=Xt}}}),await y.onBatchEnd(_,T),Cc(T),this.stopTraining_)break}E.dispose()}if(await y.onEpochEnd(x,$),this.stopTraining_)break}return await y.onTrainEnd(),await this.history.syncData(),this.history}async fitDataset(t,e){return a1(this,t,e)}async trainOnBatch(t,e){const s=await this.standardizeUserData(t,e),r=s[0],o=s[1],a=this.makeTrainFunction()(r.concat(o)),l=[];for(const u of a){const c=await u.data();l.push(c[0])}return ut(a),ge(s[0],t),ge(s[1],e),Ut(l)}getNamedWeights(t){const e=[],s=t!=null&&t.trainableOnly,r=s?this.trainableWeights:this.weights,o=this.getWeights(s);for(let i=0;i<r.length;++i)s&&!r[i].trainable||e.push({name:r[i].originalName,tensor:o[i]});return e}set stopTraining(t){this.stopTraining_=t}get stopTraining(){return this.stopTraining_}get optimizer(){return this.optimizer_}set optimizer(t){this.optimizer_!==t&&(this.optimizer_=t,this.isOptimizerOwned=!1)}dispose(){const t=super.dispose();if(t.refCountAfterDispose===0&&this.optimizer!=null&&this.isOptimizerOwned){const e=fl().numTensors;this.optimizer_.dispose(),t.numDisposedVariables+=e-fl().numTensors}return t}getLossIdentifiers(){let t;if(typeof this.loss=="string")t=Le(this.loss);else if(Array.isArray(this.loss)){for(const e of this.loss)if(typeof e!="string")throw new Error("Serialization of non-string loss is not supported.");t=this.loss.map(e=>Le(e))}else{const e=Object.keys(this.loss);t={};const s=this.loss;for(const r of e)if(typeof s[r]=="string")t[r]=Le(s[r]);else throw new Error("Serialization of non-string loss is not supported.")}return t}getMetricIdentifiers(){if(typeof this.metrics=="string"||typeof this.metrics=="function")return[Le(Rr(this.metrics))];if(Array.isArray(this.metrics))return this.metrics.map(t=>Le(Rr(t)));{const t={};for(const e in this.metrics)t[e]=Le(Rr(this.metrics[e]));return t}}getTrainingConfig(){return{loss:this.getLossIdentifiers(),metrics:this.getMetricIdentifiers(),optimizer_config:{class_name:this.optimizer.getClassName(),config:this.optimizer.getConfig()}}}loadTrainingConfig(t){if(t.weighted_metrics!=null)throw new Error("Loading weight_metrics is not supported yet.");if(t.loss_weights!=null)throw new Error("Loading loss_weights is not supported yet.");if(t.sample_weight_mode!=null)throw new Error("Loading sample_weight_mode is not supported yet.");const e=Ti(t.optimizer_config),s=Nc(e);let r;if(typeof t.loss=="string")r=pn(t.loss);else if(Array.isArray(t.loss))r=t.loss.map(i=>pn(i));else if(t.loss!=null){r={};for(const i in t.loss)r[i]=pn(t.loss[i])}let o;if(Array.isArray(t.metrics))o=t.metrics.map(i=>pn(i));else if(t.metrics!=null){o={};for(const i in t.metrics)o[i]=pn(t.metrics[i])}this.compile({loss:r,metrics:o,optimizer:s})}async save(t,e){if(typeof t=="string"){const u=nm(t);if(u.length===0)throw new I(`Cannot find any save handlers for URL '${t}'`);if(u.length>1)throw new I(`Found more than one (${u.length}) save handlers for URL '${t}'`);t=u[0]}if(t.save==null)throw new I("LayersModel.save() cannot proceed because the IOHandler provided does not have the `save` attribute defined.");const s=await pl(this.getNamedWeights(e)),a={modelTopology:this.toJSON(null,!1),format:m1,generatedBy:`TensorFlow.js tfjs-layers v${Fc}`,convertedBy:null};if((e==null?!1:e.includeOptimizer)&&this.optimizer!=null){a.trainingConfig=this.getTrainingConfig();const u="optimizer",{data:c,specs:h}=await pl(await this.optimizer.getWeights(),u);s.specs.push(...h),s.data=em([s.data,c])}return this.userDefinedMetadata!=null&&(Oc(this.userDefinedMetadata,this.name,!0),a.userDefinedMetadata=this.userDefinedMetadata),a.weightData=s.data,a.weightSpecs=s.specs,t.save(a)}setUserDefinedMetadata(t){Oc(t,this.name),this.userDefinedMetadata=t}getUserDefinedMetadata(){return this.userDefinedMetadata}}Lr.className="Model",M(Lr);class Qc extends Lr{}Qc.className="Functional",M(Qc);const g1="This is not an object",b1="This is not a Float16Array object",th="This constructor is not a subclass of Float16Array",eh="The constructor property value is not an object",y1="Species constructor didn't return TypedArray object",w1="Derived constructor created TypedArray object which was too small length",Ds="Attempting to access detached ArrayBuffer",Mi="Cannot convert undefined or null to object",Oi="Cannot mix BigInt and other types, use explicit conversions",nh="@@iterator property is not callable",sh="Reduce of empty array with no initial value",x1="The comparison function must be either a function or undefined",Bi="Offset is out of bounds";function ct(n){return(t,...e)=>Vt(n,t,e)}function Xn(n,t){return ct(Jn(n,t).get)}const{apply:Vt,construct:Rs,defineProperty:rh,get:Fi,getOwnPropertyDescriptor:Jn,getPrototypeOf:Ps,has:zi,ownKeys:oh,set:ih,setPrototypeOf:ah}=Reflect,S1=Proxy,{EPSILON:$1,MAX_SAFE_INTEGER:lh,isFinite:uh,isNaN:Zn}=Number,{iterator:_e,species:v1,toStringTag:Ui,for:I1}=Symbol,Qn=Object,{create:Mr,defineProperty:Ls,freeze:A1,is:ch}=Qn,Wi=Qn.prototype,E1=Wi.__lookupGetter__?ct(Wi.__lookupGetter__):(n,t)=>{if(n==null)throw ht(Mi);let e=Qn(n);do{const s=Jn(e,t);if(s!==void 0)return Oe(s,"get")?s.get:void 0}while((e=Ps(e))!==null)},Oe=Qn.hasOwn||ct(Wi.hasOwnProperty),hh=Array,fh=hh.isArray,Or=hh.prototype,C1=ct(Or.join),k1=ct(Or.push),_1=ct(Or.toLocaleString),Gi=Or[_e],T1=ct(Gi),{abs:N1,trunc:dh}=Math,Br=ArrayBuffer,D1=Br.isView,ph=Br.prototype,R1=ct(ph.slice),P1=Xn(ph,"byteLength"),Vi=typeof SharedArrayBuffer<"u"?SharedArrayBuffer:null,L1=Vi&&Xn(Vi.prototype,"byteLength"),qi=Ps(Uint8Array),M1=qi.from,$t=qi.prototype,O1=$t[_e],B1=ct($t.keys),F1=ct($t.values),z1=ct($t.entries),U1=ct($t.set),mh=ct($t.reverse),W1=ct($t.fill),G1=ct($t.copyWithin),gh=ct($t.sort),Ms=ct($t.slice),V1=ct($t.subarray),vt=Xn($t,"buffer"),Sn=Xn($t,"byteOffset"),Q=Xn($t,"length"),bh=Xn($t,Ui),q1=Uint8Array,Kt=Uint16Array,yh=(...n)=>Vt(M1,Kt,n),ji=Uint32Array,j1=Float32Array,$n=Ps([][_e]()),Fr=ct($n.next),H1=ct(function*(){}().next),K1=Ps($n),ht=TypeError,Hi=RangeError,wh=WeakSet,xh=wh.prototype,Y1=ct(xh.add),X1=ct(xh.has),zr=WeakMap,Ki=zr.prototype,Ur=ct(Ki.get),J1=ct(Ki.has),Yi=ct(Ki.set),Sh=new zr,Z1=Mr(null,{next:{value:function(){const t=Ur(Sh,this);return Fr(t)}},[_e]:{value:function(){return this}}});function Wr(n){if(n[_e]===Gi&&$n.next===Fr)return n;const t=Mr(Z1);return Yi(Sh,t,T1(n)),t}const $h=new zr,vh=Mr(K1,{next:{value:function(){const t=Ur($h,this);return H1(t)},writable:!0,configurable:!0}});for(const n of oh($n))n!=="next"&&Ls(vh,n,Jn($n,n));function Ih(n){const t=Mr(vh);return Yi($h,t,n),t}function Gr(n){return n!==null&&typeof n=="object"||typeof n=="function"}function Ah(n){return n!==null&&typeof n=="object"}function Vr(n){return bh(n)!==void 0}function Xi(n){const t=bh(n);return t==="BigInt64Array"||t==="BigUint64Array"}function Q1(n){try{return fh(n)?!1:(P1(n),!0)}catch{return!1}}function Eh(n){if(Vi===null)return!1;try{return L1(n),!0}catch{return!1}}function tx(n){return Q1(n)||Eh(n)}function Ch(n){return fh(n)?n[_e]===Gi&&$n.next===Fr:!1}function ex(n){return Vr(n)?n[_e]===O1&&$n.next===Fr:!1}function qr(n){if(typeof n!="string")return!1;const t=+n;return n!==t+""||!uh(t)?!1:t===dh(t)}const jr=I1("__Float16Array__");function nx(n){if(!Ah(n))return!1;const t=Ps(n);if(!Ah(t))return!1;const e=t.constructor;if(e===void 0)return!1;if(!Gr(e))throw ht(eh);return zi(e,jr)}const Ji=1/$1;function sx(n){return n+Ji-Ji}const kh=6103515625e-14,rx=65504,_h=.0009765625,Th=_h*kh,ox=_h*Ji;function ix(n){const t=+n;if(!uh(t)||t===0)return t;const e=t>0?1:-1,s=N1(t);if(s<kh)return e*sx(s/Th)*Th;const r=(1+ox)*s,o=r-(r-s);return o>rx||Zn(o)?e*(1/0):e*o}const Nh=new Br(4),Dh=new j1(Nh),Rh=new ji(Nh),be=new Kt(512),ye=new q1(512);for(let n=0;n<256;++n){const t=n-127;t<-24?(be[n]=0,be[n|256]=32768,ye[n]=24,ye[n|256]=24):t<-14?(be[n]=1024>>-t-14,be[n|256]=1024>>-t-14|32768,ye[n]=-t-1,ye[n|256]=-t-1):t<=15?(be[n]=t+15<<10,be[n|256]=t+15<<10|32768,ye[n]=13,ye[n|256]=13):t<128?(be[n]=31744,be[n|256]=64512,ye[n]=24,ye[n|256]=24):(be[n]=31744,be[n|256]=64512,ye[n]=13,ye[n|256]=13)}function Te(n){Dh[0]=ix(n);const t=Rh[0],e=t>>23&511;return be[e]+((t&8388607)>>ye[e])}const Zi=new ji(2048);for(let n=1;n<1024;++n){let t=n<<13,e=0;for(;!(t&8388608);)t<<=1,e-=8388608;t&=-8388609,e+=947912704,Zi[n]=t|e}for(let n=1024;n<2048;++n)Zi[n]=939524096+(n-1024<<13);const ts=new ji(64);for(let n=1;n<31;++n)ts[n]=n<<23;ts[31]=1199570944,ts[32]=2147483648;for(let n=33;n<63;++n)ts[n]=2147483648+(n-32<<23);ts[63]=3347054592;const Ph=new Kt(64);for(let n=1;n<64;++n)n!==32&&(Ph[n]=1024);function st(n){const t=n>>10;return Rh[0]=Zi[Ph[t]+(n&1023)]+ts[t],Dh[0]}function Be(n){const t=+n;return Zn(t)||t===0?0:dh(t)}function Qi(n){const t=Be(n);return t<0?0:t<lh?t:lh}function Hr(n,t){if(!Gr(n))throw ht(g1);const e=n.constructor;if(e===void 0)return t;if(!Gr(e))throw ht(eh);const s=e[v1];return s??t}function Os(n){if(Eh(n))return!1;try{return R1(n,0,0),!1}catch{}return!0}function Lh(n,t){const e=Zn(n),s=Zn(t);if(e&&s)return 0;if(e)return 1;if(s||n<t)return-1;if(n>t)return 1;if(n===0&&t===0){const r=ch(n,0),o=ch(t,0);if(!r&&o)return-1;if(r&&!o)return 1}return 0}const ta=2,Kr=new zr;function es(n){return J1(Kr,n)||!D1(n)&&nx(n)}function tt(n){if(!es(n))throw ht(b1)}function Yr(n,t){const e=es(n),s=Vr(n);if(!e&&!s)throw ht(y1);if(typeof t=="number"){let r;if(e){const o=V(n);r=Q(o)}else r=Q(n);if(r<t)throw ht(w1)}if(Xi(n))throw ht(Oi)}function V(n){const t=Ur(Kr,n);if(t!==void 0){const r=vt(t);if(Os(r))throw ht(Ds);return t}const e=n.buffer;if(Os(e))throw ht(Ds);const s=Rs(ft,[e,n.byteOffset,n.length],n.constructor);return Ur(Kr,s)}function Mh(n){const t=Q(n),e=[];for(let s=0;s<t;++s)e[s]=st(n[s]);return e}const Oh=new wh;for(const n of oh($t)){if(n===Ui)continue;const t=Jn($t,n);Oe(t,"get")&&typeof t.get=="function"&&Y1(Oh,t.get)}const ax=A1({get(n,t,e){return qr(t)&&Oe(n,t)?st(Fi(n,t)):X1(Oh,E1(n,t))?Fi(n,t):Fi(n,t,e)},set(n,t,e,s){return qr(t)&&Oe(n,t)?ih(n,t,Te(e)):ih(n,t,e,s)},getOwnPropertyDescriptor(n,t){if(qr(t)&&Oe(n,t)){const e=Jn(n,t);return e.value=st(e.value),e}return Jn(n,t)},defineProperty(n,t,e){return qr(t)&&Oe(n,t)&&Oe(e,"value")&&(e.value=Te(e.value)),rh(n,t,e)}});class ft{constructor(t,e,s){let r;if(es(t))r=Rs(Kt,[V(t)],new.target);else if(Gr(t)&&!tx(t)){let i,a;if(Vr(t)){i=t,a=Q(t);const l=vt(t);if(Os(l))throw ht(Ds);if(Xi(t))throw ht(Oi);const u=new Br(a*ta);r=Rs(Kt,[u],new.target)}else{const l=t[_e];if(l!=null&&typeof l!="function")throw ht(nh);l!=null?Ch(t)?(i=t,a=t.length):(i=[...t],a=i.length):(i=t,a=Qi(i.length)),r=Rs(Kt,[a],new.target)}for(let l=0;l<a;++l)r[l]=Te(i[l])}else r=Rs(Kt,arguments,new.target);const o=new S1(r,ax);return Yi(Kr,o,r),o}static from(t,...e){const s=this;if(!zi(s,jr))throw ht(th);if(s===ft){if(es(t)&&e.length===0){const c=V(t),h=new Kt(vt(c),Sn(c),Q(c));return new ft(vt(Ms(h)))}if(e.length===0)return new ft(vt(yh(t,Te)));const l=e[0],u=e[1];return new ft(vt(yh(t,function(c,...h){return Te(Vt(l,this,[c,...Wr(h)]))},u)))}let r,o;const i=t[_e];if(i!=null&&typeof i!="function")throw ht(nh);if(i!=null)Ch(t)?(r=t,o=t.length):ex(t)?(r=t,o=Q(t)):(r=[...t],o=r.length);else{if(t==null)throw ht(Mi);r=Qn(t),o=Qi(r.length)}const a=new s(o);if(e.length===0)for(let l=0;l<o;++l)a[l]=r[l];else{const l=e[0],u=e[1];for(let c=0;c<o;++c)a[c]=Vt(l,u,[r[c],c])}return a}static of(...t){const e=this;if(!zi(e,jr))throw ht(th);const s=t.length;if(e===ft){const o=new ft(s),i=V(o);for(let a=0;a<s;++a)i[a]=Te(t[a]);return o}const r=new e(s);for(let o=0;o<s;++o)r[o]=t[o];return r}keys(){tt(this);const t=V(this);return B1(t)}values(){tt(this);const t=V(this);return Ih(function*(){for(const e of F1(t))yield st(e)}())}entries(){tt(this);const t=V(this);return Ih(function*(){for(const[e,s]of z1(t))yield[e,st(s)]}())}at(t){tt(this);const e=V(this),s=Q(e),r=Be(t),o=r>=0?r:s+r;if(!(o<0||o>=s))return st(e[o])}with(t,e){tt(this);const s=V(this),r=Q(s),o=Be(t),i=o>=0?o:r+o,a=+e;if(i<0||i>=r)throw Hi(Bi);const l=new Kt(vt(s),Sn(s),Q(s)),u=new ft(vt(Ms(l))),c=V(u);return c[i]=Te(a),u}map(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0],i=Hr(s,ft);if(i===ft){const l=new ft(r),u=V(l);for(let c=0;c<r;++c){const h=st(s[c]);u[c]=Te(Vt(t,o,[h,c,this]))}return l}const a=new i(r);Yr(a,r);for(let l=0;l<r;++l){const u=st(s[l]);a[l]=Vt(t,o,[u,l,this])}return a}filter(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0],i=[];for(let u=0;u<r;++u){const c=st(s[u]);Vt(t,o,[c,u,this])&&k1(i,c)}const a=Hr(s,ft),l=new a(i);return Yr(l),l}reduce(t,...e){tt(this);const s=V(this),r=Q(s);if(r===0&&e.length===0)throw ht(sh);let o,i;e.length===0?(o=st(s[0]),i=1):(o=e[0],i=0);for(let a=i;a<r;++a)o=t(o,st(s[a]),a,this);return o}reduceRight(t,...e){tt(this);const s=V(this),r=Q(s);if(r===0&&e.length===0)throw ht(sh);let o,i;e.length===0?(o=st(s[r-1]),i=r-2):(o=e[0],i=r-1);for(let a=i;a>=0;--a)o=t(o,st(s[a]),a,this);return o}forEach(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i)Vt(t,o,[st(s[i]),i,this])}find(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return a}}findIndex(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return i}return-1}findLast(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=r-1;i>=0;--i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return a}}findLastIndex(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=r-1;i>=0;--i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return i}return-1}every(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i)if(!Vt(t,o,[st(s[i]),i,this]))return!1;return!0}some(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i)if(Vt(t,o,[st(s[i]),i,this]))return!0;return!1}set(t,...e){tt(this);const s=V(this),r=Be(e[0]);if(r<0)throw Hi(Bi);if(t==null)throw ht(Mi);if(Xi(t))throw ht(Oi);if(es(t))return U1(V(this),V(t),r);if(Vr(t)){const l=vt(t);if(Os(l))throw ht(Ds)}const o=Q(s),i=Qn(t),a=Qi(i.length);if(r===1/0||a+r>o)throw Hi(Bi);for(let l=0;l<a;++l)s[l+r]=Te(i[l])}reverse(){tt(this);const t=V(this);return mh(t),this}toReversed(){tt(this);const t=V(this),e=new Kt(vt(t),Sn(t),Q(t)),s=new ft(vt(Ms(e))),r=V(s);return mh(r),s}fill(t,...e){tt(this);const s=V(this);return W1(s,Te(t),...Wr(e)),this}copyWithin(t,e,...s){tt(this);const r=V(this);return G1(r,t,e,...Wr(s)),this}sort(t){tt(this);const e=V(this),s=t!==void 0?t:Lh;return gh(e,(r,o)=>s(st(r),st(o))),this}toSorted(t){tt(this);const e=V(this);if(t!==void 0&&typeof t!="function")throw new ht(x1);const s=t!==void 0?t:Lh,r=new Kt(vt(e),Sn(e),Q(e)),o=new ft(vt(Ms(r))),i=V(o);return gh(i,(a,l)=>s(st(a),st(l))),o}slice(t,e){tt(this);const s=V(this),r=Hr(s,ft);if(r===ft){const p=new Kt(vt(s),Sn(s),Q(s));return new ft(vt(Ms(p,t,e)))}const o=Q(s),i=Be(t),a=e===void 0?o:Be(e);let l;i===-1/0?l=0:i<0?l=o+i>0?o+i:0:l=o<i?o:i;let u;a===-1/0?u=0:a<0?u=o+a>0?o+a:0:u=o<a?o:a;const c=u-l>0?u-l:0,h=new r(c);if(Yr(h,c),c===0)return h;const f=vt(s);if(Os(f))throw ht(Ds);let d=0;for(;l<u;)h[d]=st(s[l]),++l,++d;return h}subarray(t,e){tt(this);const s=V(this),r=Hr(s,ft),o=new Kt(vt(s),Sn(s),Q(s)),i=V1(o,t,e),a=new r(vt(i),Sn(i),Q(i));return Yr(a),a}indexOf(t,...e){tt(this);const s=V(this),r=Q(s);let o=Be(e[0]);if(o===1/0)return-1;o<0&&(o+=r,o<0&&(o=0));for(let i=o;i<r;++i)if(Oe(s,i)&&st(s[i])===t)return i;return-1}lastIndexOf(t,...e){tt(this);const s=V(this),r=Q(s);let o=e.length>=1?Be(e[0]):r-1;if(o===-1/0)return-1;o>=0?o=o<r-1?o:r-1:o+=r;for(let i=o;i>=0;--i)if(Oe(s,i)&&st(s[i])===t)return i;return-1}includes(t,...e){tt(this);const s=V(this),r=Q(s);let o=Be(e[0]);if(o===1/0)return!1;o<0&&(o+=r,o<0&&(o=0));const i=Zn(t);for(let a=o;a<r;++a){const l=st(s[a]);if(i&&Zn(l)||l===t)return!0}return!1}join(t){tt(this);const e=V(this),s=Mh(e);return C1(s,t)}toLocaleString(...t){tt(this);const e=V(this),s=Mh(e);return _1(s,...Wr(t))}get[Ui](){if(es(this))return"Float16Array"}}Ls(ft,"BYTES_PER_ELEMENT",{value:ta}),Ls(ft,jr,{}),ah(ft,qi);const Xr=ft.prototype;Ls(Xr,"BYTES_PER_ELEMENT",{value:ta}),Ls(Xr,_e,{value:Xr.values,writable:!0,configurable:!0}),ah(Xr,$t);function lx(n,t){return n.channels===t.channels}const Jr=8;class Bh{constructor(t,e,s){K(this,"_label");K(this,"_device");K(this,"_outputBuffers",{});K(this,"_pipeline");K(this,"_bindGroups",[]);K(this,"_needsUpdatePipeline",!0);K(this,"_inputs",[]);K(this,"_outputs",[]);K(this,"_uniforms",[]);K(this,"_uniformBuffers",{});K(this,"_width",10);K(this,"_height",10);K(this,"_execWidth");K(this,"_execHeight");K(this,"_csCode","");K(this,"_csMain");K(this,"_csDefine");this._label=t,this._device=e,this._csMain=s.csMain,this._csDefine=s.csDefine,this._inputs=s.inputs,this._outputs=s.outputs,this._uniforms=s.uniforms,s.uniforms.forEach(r=>{this._uniformBuffers[r.label]=e.createBuffer({label:this._label,size:r.data.byteLength,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST}),this._device.queue.writeBuffer(this._uniformBuffers[r.label],0,r.data)})}setSize(t,e){t=Math.ceil(t),e=Math.ceil(e);const s=t!==this._width||e!==this._height;this._width=t,this._height=e,s&&(this._resizeOutputBuffers(),this._needsUpdatePipeline=!0)}setExecuteSize(t,e){t=Math.ceil(t),e=Math.ceil(e),this._execWidth=t,this._execHeight=e}setOutputParams(t){this._updateOutputBuffers(t),this._needsUpdatePipeline=!0}setUniform(t,e){const s=this._uniformBuffers[t];this._device.queue.writeBuffer(s,0,e)}getOutputBuffer(t){return this._outputBuffers[t].buffer}dispose(){Object.keys(this._uniformBuffers).forEach(t=>{this._uniformBuffers[t].destroy()}),Object.keys(this._outputBuffers).forEach(t=>{this._outputBuffers[t].texture.destroy()})}_createBuffer(t){const e=this._width*this._height*4*4;return this._device.createBuffer({label:this._label,size:Math.max(e,80),usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST|GPUBufferUsage.COPY_SRC})}_resizeOutputBuffers(){const t=this._outputBuffers;for(const e in t){const{buffer:s,params:r}=t[e];s.destroy(),t[e].buffer=this._createBuffer(r)}}_updateOutputBuffers(t){var s,r;const e=this._outputBuffers;for(const o in t){const i=t[o];if(!lx(i,((s=e[o])==null?void 0:s.params)||{})){(r=e[o])==null||r.buffer.destroy();const a=this._createBuffer(i);e[o]={buffer:a,params:i}}}}_updatePipeline(t){if(!this._needsUpdatePipeline)return;this._needsUpdatePipeline=!1;const e=this._device,s=this._getFullCs(t);s!==this._csCode&&(this._csCode=s,this._pipeline=e.createComputePipeline({label:this._label,layout:"auto",compute:{module:e.createShaderModule({label:this._label,code:s}),entryPoint:"main"}}),this._updateBindGroups())}_getFullCs(t){const e=this._inputs,s=e.length>0;return`
|
|
3047
|
+
*/function h1(n){return n instanceof At}function Li(n){return Array.isArray(n)}function Xc(n){return!h1(n)&&!Li(n)}function Jc(n,t,e,s=!0,r=""){if(t==null||t.length===0){if(n!=null){let i=!1;if(Li(n)&&n.length>0)i=!0;else if(Xc(n)){for(const a in n)if(n.hasOwnProperty(a)){i=!0;break}}else i=!0;if(i)throw new I(`Error when checking model ${r} expected no data, but got ${n}`)}return[]}if(n==null)return t.map(i=>null);let o;if(Xc(n)){n=n,o=[];for(const i of t){if(n[i]==null)throw new I(`No data provided for "${i}". Need data for each key in: ${t}`);o.push(n[i])}}else if(Li(n)){if(n=n,n.length!==t.length)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}`);o=n}else{if(n=n,t.length>1)throw new I(`The model ${r} expects ${t.length} Tensor(s), but only received one Tensor. Found: Tensor with shape ${n.shape}`);o=[n]}if(o=Yc(o),e!=null)for(let i=0;i<t.length;++i){if(e[i]==null)continue;const a=o[i];if(a.shape.length!==e[i].length)throw new I(`Error when checking ${r}: expected ${t[i]} to have ${e[i].length} dimension(s). but got array with shape ${a.shape}`);for(let l=0;l<e[i].length;++l){if(l===0&&!s)continue;const u=a.shape[l],c=e[i][l];if(c!=null&&c>=0&&u!==c)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}])`)}}return o}function f1(n,t,e){const s=gn(n.map(o=>o.shape[0]));s.sort();const r=gn(t.map(o=>o.shape[0]));if(r.sort(),s.length>1)throw new I(`All input Tensors (x) should have the same number of samples. Got array shapes: ${JSON.stringify(n.map(o=>o.shape))}`);if(r.length>1)throw new I(`All target Tensors (y) should have the same number of samples. Got array shapes: ${JSON.stringify(t.map(o=>o.shape))}`);if(s.length>0&&r.length>0&&!Zt(s,r))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).`)}function d1(n,t,e){const s=[_r,Nr,Cs];for(let r=0;r<n.length;++r){const o=n[r],i=t[r],a=e[r];if(i!=null){if(i===Cs&&o.shape[o.shape.length-1]===1)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].`);if(s.indexOf(i)!==-1){const l=o.shape.slice(1),u=a.slice(1);for(let c=0;c<l.length;++c){const h=l[c],f=u[c];if(f!=null&&h!==f)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.`)}}}}}function Zc(n,t,e,s=!0,r=""){let o;if(Array.isArray(n)){if(n.length!==t.length)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).`);o=n}else{if(t.length>1)throw new I(`The model expects ${t.length} ${r} Tensors, but only received one Tensor. Found: array with shape ${JSON.stringify(n.shape)}.`);o=[n]}if(e!=null)for(let i=0;i<t.length;++i){if(e[i]==null)continue;const a=o[i];if(a.shape.length!==e[i].length)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)}`);for(let l=0;l<e[i].length;++l){if(l===0&&!s)continue;const u=a.shape[l],c=e[i][l];if(c!=null&&c!==u)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)}.`)}}}function p1(n,t){if(n==null||Array.isArray(n)&&n.length===0)return t.map(s=>[]);let e;if(typeof n=="string"||typeof n=="function")e=[n];else if(Array.isArray(n)||typeof n=="object")e=n;else throw new TypeError(`Type of metrics argument not understood. Expected an string,function, Array, or Object, found: ${n}`);if(Array.isArray(e))return t.map(s=>e);{const s=[];for(const r of t){let o=e.hasOwnProperty(r)?e[r]:[];Array.isArray(o)||(o=[o]),s.push(o)}return s}}const m1="layers-model";class Mr extends me{constructor(t){super(t),this.isTraining=!1}summary(t,e,s=console.log){if(!this.built)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).");jw(this,t,e,s)}compile(t){if(t.loss==null&&(t.loss=[]),this.loss=t.loss,typeof t.optimizer=="string")this.optimizer_=qw(t.optimizer),this.isOptimizerOwned=!0;else{if(!(t.optimizer instanceof qe))throw new I("User-defined optimizer must be an instance of tf.Optimizer.");this.optimizer_=t.optimizer,this.isOptimizerOwned=!1}let e=[];if(!Array.isArray(t.loss)&&typeof t.loss!="string"&&typeof t.loss!="function"){t.loss=t.loss;for(const i in t.loss)if(this.outputNames.indexOf(i)===-1)throw new I(`Unknown entry in loss dictionary: "${i}". Only expected the following keys: ${this.outputNames}`);for(const i of this.outputNames)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(Ci(t.loss[i]))}else if(Array.isArray(t.loss)){if(t.loss.length!==this.outputs.length)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}.`);e=t.loss.map(a=>Ci(a))}else{const i=Ci(t.loss);this.outputs.forEach(a=>{e.push(i)})}this.lossFunctions=e,this.feedOutputNames=[],this.feedOutputShapes=[],this.feedLossFns=[];for(let i=0;i<this.outputs.length;++i){const a=this.internalOutputShapes[i],l=this.outputNames[i];this.feedOutputNames.push(l),this.feedOutputShapes.push(a),this.feedLossFns.push(this.lossFunctions[i])}const s=[];this.metrics=t.metrics,this.metricsNames=["loss"],this.metricsTensors=[],br("loss",()=>{for(let i=0;i<this.outputs.length;++i){if(s.indexOf(i)!==-1)continue;const a=this.lossFunctions[i];this.outputs.length>1&&(this.metricsTensors.push([a,i]),this.metricsNames.push(this.outputNames[i]+"_loss"))}});const r=p1(t.metrics,this.outputNames),o=(i,a,l)=>{this.outputNames.length>1&&(a=this.outputNames[i]+"_"+a),this.metricsNames.push(a),this.metricsTensors.push([l,i])};br("metric",()=>{for(let i=0;i<this.outputs.length;++i){if(s.indexOf(i)!==-1)continue;const a=r[i];(u=>{const c="";let h,f,d;for(const p of u){if(typeof p=="string"&&["accuracy","acc","crossentropy","ce"].indexOf(p)!==-1){const m=this.internalOutputShapes[i];m[m.length-1]===1||this.lossFunctions[i]===Nr?["accuracy","acc"].indexOf(p)!==-1?f=Dc:["crossentropy","ce"].indexOf(p)!==-1&&(f=Lw):this.lossFunctions[i]===Tr?["accuracy","acc"].indexOf(p)!==-1?f=Mw:["crossentropy","ce"].indexOf(p)!==-1&&(f=Lc):["accuracy","acc"].indexOf(p)!==-1?f=Rc:["crossentropy","ce"].indexOf(p)!==-1&&(f=Pc);let b;["accuracy","acc"].indexOf(p)!==-1?b="acc":["crossentropy","ce"].indexOf(p)!==-1&&(b="ce"),d=f,h=c+b}else d=Vw(p),h=c+Pr(p);let g;br(h,()=>{g=d}),o(i,h,g)}})(a)}}),this.collectedTrainableWeights=this.trainableWeights}checkTrainableWeightsConsistency(){this.collectedTrainableWeights!=null&&this.trainableWeights.length!==this.collectedTrainableWeights.length&&console.warn("Discrepancy between trainableweights and collected trainable weights. Did you set `model.trainable` without calling `model.compile()` afterwards?")}evaluate(t,e,s={}){const r=s.batchSize==null?32:s.batchSize;Di(r);const i=this.standardizeUserDataXY(t,e,!0,r);try{const a=i[0].concat(i[1]);this.makeTestFunction();const l=this.testFunction,u=this.testLoop(l,a,r,s.verbose,s.steps);return Ut(u)}finally{ge(i[0],t),ge(i[1],e)}}async evaluateDataset(t,e){return this.makeTestFunction(),c1(this,t,e)}checkNumSamples(t,e,s,r="steps"){let o;if(s!=null){if(o=null,e!=null)throw new I(`If ${r} is set, batchSize must be null or undefined.Got batchSize = ${e}`)}else if(t!=null)Array.isArray(t)?o=t[0].shape[0]:o=t.shape[0];else throw new I(`Either the input data should have a defined shape, or ${r} shoud be specified.`);return o}execute(t,e){if(Array.isArray(e)&&e.length===0)throw new I("`outputs` is an empty Array, which is not allowed.");const s=Array.isArray(e),r=s?e:[e],o=this.retrieveSymbolicTensors(r),i=new Ke;if(t instanceof At&&(t=[t]),Array.isArray(t)){if(t.length!==this.inputs.length)throw new I(`The number of inputs provided (${t.length}) does not match the number of inputs of this model (${this.inputs.length}).`);for(let l=0;l<this.inputs.length;++l)i.add(this.inputs[l],t[l])}else for(const l of this.inputs){const u=t[l.name];if(u==null)throw new I(`No value is provided for the model's input ${l.name}`);i.add(l,u)}const a=Ns(o,i);return s?a:a[0]}retrieveSymbolicTensors(t){const e=mr(null,t.length);let s=t.length;for(const r of this.layers){const o=Array.isArray(r.output)?r.output:[r.output],i=o.map(a=>a.name);for(let a=0;a<t.length;++a){const l=i.indexOf(t[a]);if(l!==-1&&(e[a]=o[l],s--),s===0)break}if(s===0)break}if(s>0){const r=[];throw e.forEach((o,i)=>{o==null&&r.push(t[i])}),new I(`Cannot find SymbolicTensors for output name(s): ${JSON.stringify(r)}`)}return e}predictLoop(t,e=32,s=!1){return k(()=>{const r=this.checkNumSamples(t);if(s)throw new Z("Verbose predictLoop() is not implemented yet.");const o=Pi(r,e),i=this.outputs.map(a=>[]);for(let a=0;a<o.length;++a)k(()=>{const u=o[a][0],c=o[a][1],h=Ds(t,u,c),f=[];if(Array.isArray(h))for(let p=0;p<h.length;++p)f.push({key:this.inputs[p],value:h[p]});else f.push({key:this.inputs[0],value:h});const d=new Ke(f);return Ns(this.outputs,d)}).forEach((u,c)=>i[c].push(u));return Ut(i.map(a=>an(a,0)))})}predict(t,e={}){const s=Yc(t);Zc(s,this.inputNames,this.feedInputShapes,!1);try{const r=e.batchSize==null?32:e.batchSize;return Di(r),this.predictLoop(s,r)}finally{ge(s,t)}}predictOnBatch(t){Zc(t,this.inputNames,this.feedInputShapes,!0);const e=(Array.isArray(t)?t[0]:t).shape[0];return this.predictLoop(t,e)}standardizeUserDataXY(t,e,s=!0,r){if(this.optimizer_==null)throw new He("You must compile a model before training/testing. Use LayersModel.compile(modelCompileArgs).");const o=[];for(let i=0;i<this.feedOutputShapes.length;++i){const a=this.feedOutputShapes[i];this.feedLossFns[i]===Tr?o.push(a.slice(0,a.length-1).concat([1])):o.push(a)}if(t=Jc(t,this.feedInputNames,this.feedInputShapes,!1,"input"),e=Jc(e,this.feedOutputNames,o,!1,"target"),f1(t,e),d1(e,this.feedLossFns,this.feedOutputShapes),this.stateful&&r!=null&&r>0&&t[0].shape[0]%r!==0)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).`);return[t,e]}async standardizeUserData(t,e,s,r,o=!0,i){const[a,l]=this.standardizeUserDataXY(t,e,o,i);if(s!=null)throw new Error("sample weight is not supported yet.");let u=null;if(r!=null){const c=Vc(r,this.outputNames);u=[];for(let h=0;h<c.length;++h)u.push(await qc(l[h],null,c[h]))}return[a,l,u]}testLoop(t,e,s,r=0,o){return k(()=>{const i=this.checkNumSamples(e,s,o,"steps"),a=[];if(r>0)throw new Z("Verbose mode is not implemented yet.");if(o!=null)throw new Z("steps mode in testLoop() is not implemented yet");{const l=Pi(i,s),u=Rt(yr(0,i));for(let c=0;c<l.length;++c){const h=l[c][0],f=l[c][1],d=bn(u,h,f-h),p=Ri(e,d),g=t(p);if(c===0)for(let m=0;m<g.length;++m)a.push(Yt(0));for(let m=0;m<g.length;++m){const b=g[m];a[m]=O(a[m],N(f-h,b))}}for(let c=0;c<a.length;++c)a[c]=X(a[c],i)}return a})}getDedupedMetricsNames(){const t=this.metricsNames,e=[];for(let s=0;s<t.length;++s){const r=t[s];let o=r;if(eu(t,r)>1){const i=eu(t.slice(0,s),r);o+=`_${i}`}e.push(o)}return e}makeTrainFunction(){return t=>{const e=[],s=t.slice(0,this.inputs.length),r=t.slice(this.inputs.length,this.inputs.length+this.outputs.length),o=t.slice(this.inputs.length+this.outputs.length,this.inputs.length+this.outputs.length*2),i=[],a=()=>{const h=[];for(let g=0;g<this.inputs.length;++g)h.push({key:this.inputs[g],value:s[g]});const f=new Ke(h),d=Ns(this.outputs,f,{training:!0});let p;for(let g=0;g<this.lossFunctions.length;++g){const m=this.lossFunctions[g];let b=m(r[g],d[g]);o[g]!=null&&(b=r1(b,o[g]));const y=St(b);e.push(y),g===0?p=b:p=O(p,b)}for(let g=0;g<this.metricsTensors.length;++g){let m;if(this.outputs.length>1&&g<this.outputs.length)m=e[g];else{const b=this.metricsTensors[g][0],y=this.metricsTensors[g][1];m=St(b(r[y],d[y]))}Ln(m),i.push(m)}return p=St(p),this.calculateLosses().forEach(g=>{p=O(p,g)}),p},l=this.collectedTrainableWeights.map(h=>h.read());return[this.optimizer_.minimize(a,!0,l)].concat(i)}}makeTestFunction(){this.testFunction=t=>k(()=>{const e=[];let s;const r=t.slice(0,this.inputs.length),o=t.slice(this.inputs.length,this.inputs.length+this.outputs.length),i=[];for(let u=0;u<this.inputs.length;++u)i.push({key:this.inputs[u],value:r[u]});const a=new Ke(i),l=Ns(this.outputs,a);for(let u=0;u<this.lossFunctions.length;++u){const c=this.lossFunctions[u],h=St(c(o[u],l[u]));u===0?s=h:s=O(s,h),e.push(s)}for(let u=0;u<this.metricsTensors.length;++u){const c=this.metricsTensors[u][0],h=this.metricsTensors[u][1],f=St(c(o[h],l[h]));e.push(f)}return e})}async fit(t,e,s={}){if(this.isTraining)throw new Error("Cannot start training because another fit() call is ongoing.");this.isTraining=!0;let r,o,i,a,l,u,c,h,f;try{const d=s.batchSize==null?32:s.batchSize;Di(d);const g=await this.standardizeUserData(t,e,s.sampleWeight,s.classWeight,!1,d);r=g[0],o=g[1],f=g[2];let m=!1,b;if(s.validationData!=null&&s.validationData.length>0){if(m=!0,s.validationData.length===2)l=s.validationData[0],u=s.validationData[1];else throw s.validationData.length===3?new Z("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.`);const P=await this.standardizeUserData(l,u,null,null,!0,d);c=P[0],h=P[1],b=c.concat(h)}else if(s.validationSplit!=null&&s.validationSplit>0&&s.validationSplit<1){m=!0;const T=Math.floor(r[0].shape[0]*(1-s.validationSplit)),P=r[0].shape[0];c=Ds(r,T,P),i=r,r=Ds(r,0,T),h=Ds(o,T,P),a=o,o=Ds(o,0,T),b=c.concat(h)}else s.validationSteps!=null&&(m=!0);const y=r.concat(o).concat(f);this.checkTrainableWeightsConsistency();const S=this.makeTrainFunction(),x=this.getDedupedMetricsNames();let $,E;m?(this.makeTestFunction(),$=this.testFunction,E=x.slice().concat(x.map(T=>"val_"+T))):($=null,b=[],E=x.slice());const D=_c(s.callbacks,s.yieldEvery);return await this.fitLoop(S,y,x,d,s.epochs,s.verbose,D,$,b,s.shuffle,E,s.initialEpoch,null,null)}finally{this.isTraining=!1,ge(r,t),ge(o,e),ge(i,t),ge(a,e),ge(c,l),ge(h,u),f!=null&&ut(f)}}async fitLoop(t,e,s,r,o,i,a,l,u,c,h,f,d,p){r==null&&(r=32),o==null&&(o=1),c==null&&(c=!0),f==null&&(f=0);let g=!1;if(l!=null&&u!=null&&(g=!0),p!=null&&(g=!0,d==null))throw new I("Can only use `validationSteps` when doing step-wise training, i.e., `stepsPerEpoch` must be set.");const m=this.checkNumSamples(e,r,d,"steps_per_epoch");let b;m!=null&&(b=yr(0,m)),i==null&&(i=1);const{callbackList:y,history:S}=Tc(a,i,o,f,m,d,r,g,h);y.setModel(this),this.history=S,await y.onTrainBegin(),this.stopTraining_=!1;for(let x=f;x<o;++x){await y.onEpochBegin(x);const $={};if(d!=null)throw new Z("stepsPerEpoch mode is not implemented yet.");{if(c==="batch")throw new Z("batch shuffling is not implemneted yet");c&&If(b);const E=Rt(b),D=Pi(m,r);for(let _=0;_<D.length;++_){const T={};if(await y.onBatchBegin(_,T),k(()=>{const P=D[_][0],B=D[_][1],Y=bn(E,P,B-P);T.batch=_,T.size=B-P;const j=Ri(e,Y),F=t(j);for(let G=0;G<s.length;++G){const q=s[G],De=F[G];T[q]=De,Ln(De)}if(_===D.length-1&&g){const G=this.testLoop(l,u,r);for(let q=0;q<s.length;++q){const De=s[q],qt=G[q];Ln(qt),$["val_"+De]=qt}}}),await y.onBatchEnd(_,T),Cc(T),this.stopTraining_)break}E.dispose()}if(await y.onEpochEnd(x,$),this.stopTraining_)break}return await y.onTrainEnd(),await this.history.syncData(),this.history}async fitDataset(t,e){return a1(this,t,e)}async trainOnBatch(t,e){const s=await this.standardizeUserData(t,e),r=s[0],o=s[1],a=this.makeTrainFunction()(r.concat(o)),l=[];for(const u of a){const c=await u.data();l.push(c[0])}return ut(a),ge(s[0],t),ge(s[1],e),Ut(l)}getNamedWeights(t){const e=[],s=t!=null&&t.trainableOnly,r=s?this.trainableWeights:this.weights,o=this.getWeights(s);for(let i=0;i<r.length;++i)s&&!r[i].trainable||e.push({name:r[i].originalName,tensor:o[i]});return e}set stopTraining(t){this.stopTraining_=t}get stopTraining(){return this.stopTraining_}get optimizer(){return this.optimizer_}set optimizer(t){this.optimizer_!==t&&(this.optimizer_=t,this.isOptimizerOwned=!1)}dispose(){const t=super.dispose();if(t.refCountAfterDispose===0&&this.optimizer!=null&&this.isOptimizerOwned){const e=fl().numTensors;this.optimizer_.dispose(),t.numDisposedVariables+=e-fl().numTensors}return t}getLossIdentifiers(){let t;if(typeof this.loss=="string")t=Le(this.loss);else if(Array.isArray(this.loss)){for(const e of this.loss)if(typeof e!="string")throw new Error("Serialization of non-string loss is not supported.");t=this.loss.map(e=>Le(e))}else{const e=Object.keys(this.loss);t={};const s=this.loss;for(const r of e)if(typeof s[r]=="string")t[r]=Le(s[r]);else throw new Error("Serialization of non-string loss is not supported.")}return t}getMetricIdentifiers(){if(typeof this.metrics=="string"||typeof this.metrics=="function")return[Le(Pr(this.metrics))];if(Array.isArray(this.metrics))return this.metrics.map(t=>Le(Pr(t)));{const t={};for(const e in this.metrics)t[e]=Le(Pr(this.metrics[e]));return t}}getTrainingConfig(){return{loss:this.getLossIdentifiers(),metrics:this.getMetricIdentifiers(),optimizer_config:{class_name:this.optimizer.getClassName(),config:this.optimizer.getConfig()}}}loadTrainingConfig(t){if(t.weighted_metrics!=null)throw new Error("Loading weight_metrics is not supported yet.");if(t.loss_weights!=null)throw new Error("Loading loss_weights is not supported yet.");if(t.sample_weight_mode!=null)throw new Error("Loading sample_weight_mode is not supported yet.");const e=Ti(t.optimizer_config),s=Nc(e);let r;if(typeof t.loss=="string")r=mn(t.loss);else if(Array.isArray(t.loss))r=t.loss.map(i=>mn(i));else if(t.loss!=null){r={};for(const i in t.loss)r[i]=mn(t.loss[i])}let o;if(Array.isArray(t.metrics))o=t.metrics.map(i=>mn(i));else if(t.metrics!=null){o={};for(const i in t.metrics)o[i]=mn(t.metrics[i])}this.compile({loss:r,metrics:o,optimizer:s})}async save(t,e){if(typeof t=="string"){const u=nm(t);if(u.length===0)throw new I(`Cannot find any save handlers for URL '${t}'`);if(u.length>1)throw new I(`Found more than one (${u.length}) save handlers for URL '${t}'`);t=u[0]}if(t.save==null)throw new I("LayersModel.save() cannot proceed because the IOHandler provided does not have the `save` attribute defined.");const s=await pl(this.getNamedWeights(e)),a={modelTopology:this.toJSON(null,!1),format:m1,generatedBy:`TensorFlow.js tfjs-layers v${Fc}`,convertedBy:null};if((e==null?!1:e.includeOptimizer)&&this.optimizer!=null){a.trainingConfig=this.getTrainingConfig();const u="optimizer",{data:c,specs:h}=await pl(await this.optimizer.getWeights(),u);s.specs.push(...h),s.data=em([s.data,c])}return this.userDefinedMetadata!=null&&(Oc(this.userDefinedMetadata,this.name,!0),a.userDefinedMetadata=this.userDefinedMetadata),a.weightData=s.data,a.weightSpecs=s.specs,t.save(a)}setUserDefinedMetadata(t){Oc(t,this.name),this.userDefinedMetadata=t}getUserDefinedMetadata(){return this.userDefinedMetadata}}Mr.className="Model",M(Mr);class Qc extends Mr{}Qc.className="Functional",M(Qc);const g1="This is not an object",b1="This is not a Float16Array object",th="This constructor is not a subclass of Float16Array",eh="The constructor property value is not an object",y1="Species constructor didn't return TypedArray object",w1="Derived constructor created TypedArray object which was too small length",Rs="Attempting to access detached ArrayBuffer",Mi="Cannot convert undefined or null to object",Oi="Cannot mix BigInt and other types, use explicit conversions",nh="@@iterator property is not callable",sh="Reduce of empty array with no initial value",x1="The comparison function must be either a function or undefined",Bi="Offset is out of bounds";function ct(n){return(t,...e)=>Vt(n,t,e)}function Xn(n,t){return ct(Jn(n,t).get)}const{apply:Vt,construct:Ps,defineProperty:rh,get:Fi,getOwnPropertyDescriptor:Jn,getPrototypeOf:Ls,has:zi,ownKeys:oh,set:ih,setPrototypeOf:ah}=Reflect,S1=Proxy,{EPSILON:$1,MAX_SAFE_INTEGER:lh,isFinite:uh,isNaN:Zn}=Number,{iterator:_e,species:v1,toStringTag:Ui,for:I1}=Symbol,Qn=Object,{create:Or,defineProperty:Ms,freeze:A1,is:ch}=Qn,Wi=Qn.prototype,E1=Wi.__lookupGetter__?ct(Wi.__lookupGetter__):(n,t)=>{if(n==null)throw ht(Mi);let e=Qn(n);do{const s=Jn(e,t);if(s!==void 0)return Oe(s,"get")?s.get:void 0}while((e=Ls(e))!==null)},Oe=Qn.hasOwn||ct(Wi.hasOwnProperty),hh=Array,fh=hh.isArray,Br=hh.prototype,C1=ct(Br.join),k1=ct(Br.push),_1=ct(Br.toLocaleString),Gi=Br[_e],T1=ct(Gi),{abs:N1,trunc:dh}=Math,Fr=ArrayBuffer,D1=Fr.isView,ph=Fr.prototype,R1=ct(ph.slice),P1=Xn(ph,"byteLength"),Vi=typeof SharedArrayBuffer<"u"?SharedArrayBuffer:null,L1=Vi&&Xn(Vi.prototype,"byteLength"),qi=Ls(Uint8Array),M1=qi.from,$t=qi.prototype,O1=$t[_e],B1=ct($t.keys),F1=ct($t.values),z1=ct($t.entries),U1=ct($t.set),mh=ct($t.reverse),W1=ct($t.fill),G1=ct($t.copyWithin),gh=ct($t.sort),Os=ct($t.slice),V1=ct($t.subarray),vt=Xn($t,"buffer"),$n=Xn($t,"byteOffset"),Q=Xn($t,"length"),bh=Xn($t,Ui),q1=Uint8Array,Xt=Uint16Array,yh=(...n)=>Vt(M1,Xt,n),ji=Uint32Array,j1=Float32Array,vn=Ls([][_e]()),zr=ct(vn.next),H1=ct(function*(){}().next),K1=Ls(vn),ht=TypeError,Hi=RangeError,wh=WeakSet,xh=wh.prototype,Y1=ct(xh.add),X1=ct(xh.has),Ur=WeakMap,Ki=Ur.prototype,Wr=ct(Ki.get),J1=ct(Ki.has),Yi=ct(Ki.set),Sh=new Ur,Z1=Or(null,{next:{value:function(){const t=Wr(Sh,this);return zr(t)}},[_e]:{value:function(){return this}}});function Gr(n){if(n[_e]===Gi&&vn.next===zr)return n;const t=Or(Z1);return Yi(Sh,t,T1(n)),t}const $h=new Ur,vh=Or(K1,{next:{value:function(){const t=Wr($h,this);return H1(t)},writable:!0,configurable:!0}});for(const n of oh(vn))n!=="next"&&Ms(vh,n,Jn(vn,n));function Ih(n){const t=Or(vh);return Yi($h,t,n),t}function Vr(n){return n!==null&&typeof n=="object"||typeof n=="function"}function Ah(n){return n!==null&&typeof n=="object"}function qr(n){return bh(n)!==void 0}function Xi(n){const t=bh(n);return t==="BigInt64Array"||t==="BigUint64Array"}function Q1(n){try{return fh(n)?!1:(P1(n),!0)}catch{return!1}}function Eh(n){if(Vi===null)return!1;try{return L1(n),!0}catch{return!1}}function tx(n){return Q1(n)||Eh(n)}function Ch(n){return fh(n)?n[_e]===Gi&&vn.next===zr:!1}function ex(n){return qr(n)?n[_e]===O1&&vn.next===zr:!1}function jr(n){if(typeof n!="string")return!1;const t=+n;return n!==t+""||!uh(t)?!1:t===dh(t)}const Hr=I1("__Float16Array__");function nx(n){if(!Ah(n))return!1;const t=Ls(n);if(!Ah(t))return!1;const e=t.constructor;if(e===void 0)return!1;if(!Vr(e))throw ht(eh);return zi(e,Hr)}const Ji=1/$1;function sx(n){return n+Ji-Ji}const kh=6103515625e-14,rx=65504,_h=.0009765625,Th=_h*kh,ox=_h*Ji;function ix(n){const t=+n;if(!uh(t)||t===0)return t;const e=t>0?1:-1,s=N1(t);if(s<kh)return e*sx(s/Th)*Th;const r=(1+ox)*s,o=r-(r-s);return o>rx||Zn(o)?e*(1/0):e*o}const Nh=new Fr(4),Dh=new j1(Nh),Rh=new ji(Nh),be=new Xt(512),ye=new q1(512);for(let n=0;n<256;++n){const t=n-127;t<-24?(be[n]=0,be[n|256]=32768,ye[n]=24,ye[n|256]=24):t<-14?(be[n]=1024>>-t-14,be[n|256]=1024>>-t-14|32768,ye[n]=-t-1,ye[n|256]=-t-1):t<=15?(be[n]=t+15<<10,be[n|256]=t+15<<10|32768,ye[n]=13,ye[n|256]=13):t<128?(be[n]=31744,be[n|256]=64512,ye[n]=24,ye[n|256]=24):(be[n]=31744,be[n|256]=64512,ye[n]=13,ye[n|256]=13)}function Te(n){Dh[0]=ix(n);const t=Rh[0],e=t>>23&511;return be[e]+((t&8388607)>>ye[e])}const Zi=new ji(2048);for(let n=1;n<1024;++n){let t=n<<13,e=0;for(;!(t&8388608);)t<<=1,e-=8388608;t&=-8388609,e+=947912704,Zi[n]=t|e}for(let n=1024;n<2048;++n)Zi[n]=939524096+(n-1024<<13);const ts=new ji(64);for(let n=1;n<31;++n)ts[n]=n<<23;ts[31]=1199570944,ts[32]=2147483648;for(let n=33;n<63;++n)ts[n]=2147483648+(n-32<<23);ts[63]=3347054592;const Ph=new Xt(64);for(let n=1;n<64;++n)n!==32&&(Ph[n]=1024);function st(n){const t=n>>10;return Rh[0]=Zi[Ph[t]+(n&1023)]+ts[t],Dh[0]}function Be(n){const t=+n;return Zn(t)||t===0?0:dh(t)}function Qi(n){const t=Be(n);return t<0?0:t<lh?t:lh}function Kr(n,t){if(!Vr(n))throw ht(g1);const e=n.constructor;if(e===void 0)return t;if(!Vr(e))throw ht(eh);const s=e[v1];return s??t}function Bs(n){if(Eh(n))return!1;try{return R1(n,0,0),!1}catch{}return!0}function Lh(n,t){const e=Zn(n),s=Zn(t);if(e&&s)return 0;if(e)return 1;if(s||n<t)return-1;if(n>t)return 1;if(n===0&&t===0){const r=ch(n,0),o=ch(t,0);if(!r&&o)return-1;if(r&&!o)return 1}return 0}const ta=2,Yr=new Ur;function es(n){return J1(Yr,n)||!D1(n)&&nx(n)}function tt(n){if(!es(n))throw ht(b1)}function Xr(n,t){const e=es(n),s=qr(n);if(!e&&!s)throw ht(y1);if(typeof t=="number"){let r;if(e){const o=V(n);r=Q(o)}else r=Q(n);if(r<t)throw ht(w1)}if(Xi(n))throw ht(Oi)}function V(n){const t=Wr(Yr,n);if(t!==void 0){const r=vt(t);if(Bs(r))throw ht(Rs);return t}const e=n.buffer;if(Bs(e))throw ht(Rs);const s=Ps(ft,[e,n.byteOffset,n.length],n.constructor);return Wr(Yr,s)}function Mh(n){const t=Q(n),e=[];for(let s=0;s<t;++s)e[s]=st(n[s]);return e}const Oh=new wh;for(const n of oh($t)){if(n===Ui)continue;const t=Jn($t,n);Oe(t,"get")&&typeof t.get=="function"&&Y1(Oh,t.get)}const ax=A1({get(n,t,e){return jr(t)&&Oe(n,t)?st(Fi(n,t)):X1(Oh,E1(n,t))?Fi(n,t):Fi(n,t,e)},set(n,t,e,s){return jr(t)&&Oe(n,t)?ih(n,t,Te(e)):ih(n,t,e,s)},getOwnPropertyDescriptor(n,t){if(jr(t)&&Oe(n,t)){const e=Jn(n,t);return e.value=st(e.value),e}return Jn(n,t)},defineProperty(n,t,e){return jr(t)&&Oe(n,t)&&Oe(e,"value")&&(e.value=Te(e.value)),rh(n,t,e)}});class ft{constructor(t,e,s){let r;if(es(t))r=Ps(Xt,[V(t)],new.target);else if(Vr(t)&&!tx(t)){let i,a;if(qr(t)){i=t,a=Q(t);const l=vt(t);if(Bs(l))throw ht(Rs);if(Xi(t))throw ht(Oi);const u=new Fr(a*ta);r=Ps(Xt,[u],new.target)}else{const l=t[_e];if(l!=null&&typeof l!="function")throw ht(nh);l!=null?Ch(t)?(i=t,a=t.length):(i=[...t],a=i.length):(i=t,a=Qi(i.length)),r=Ps(Xt,[a],new.target)}for(let l=0;l<a;++l)r[l]=Te(i[l])}else r=Ps(Xt,arguments,new.target);const o=new S1(r,ax);return Yi(Yr,o,r),o}static from(t,...e){const s=this;if(!zi(s,Hr))throw ht(th);if(s===ft){if(es(t)&&e.length===0){const c=V(t),h=new Xt(vt(c),$n(c),Q(c));return new ft(vt(Os(h)))}if(e.length===0)return new ft(vt(yh(t,Te)));const l=e[0],u=e[1];return new ft(vt(yh(t,function(c,...h){return Te(Vt(l,this,[c,...Gr(h)]))},u)))}let r,o;const i=t[_e];if(i!=null&&typeof i!="function")throw ht(nh);if(i!=null)Ch(t)?(r=t,o=t.length):ex(t)?(r=t,o=Q(t)):(r=[...t],o=r.length);else{if(t==null)throw ht(Mi);r=Qn(t),o=Qi(r.length)}const a=new s(o);if(e.length===0)for(let l=0;l<o;++l)a[l]=r[l];else{const l=e[0],u=e[1];for(let c=0;c<o;++c)a[c]=Vt(l,u,[r[c],c])}return a}static of(...t){const e=this;if(!zi(e,Hr))throw ht(th);const s=t.length;if(e===ft){const o=new ft(s),i=V(o);for(let a=0;a<s;++a)i[a]=Te(t[a]);return o}const r=new e(s);for(let o=0;o<s;++o)r[o]=t[o];return r}keys(){tt(this);const t=V(this);return B1(t)}values(){tt(this);const t=V(this);return Ih(function*(){for(const e of F1(t))yield st(e)}())}entries(){tt(this);const t=V(this);return Ih(function*(){for(const[e,s]of z1(t))yield[e,st(s)]}())}at(t){tt(this);const e=V(this),s=Q(e),r=Be(t),o=r>=0?r:s+r;if(!(o<0||o>=s))return st(e[o])}with(t,e){tt(this);const s=V(this),r=Q(s),o=Be(t),i=o>=0?o:r+o,a=+e;if(i<0||i>=r)throw Hi(Bi);const l=new Xt(vt(s),$n(s),Q(s)),u=new ft(vt(Os(l))),c=V(u);return c[i]=Te(a),u}map(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0],i=Kr(s,ft);if(i===ft){const l=new ft(r),u=V(l);for(let c=0;c<r;++c){const h=st(s[c]);u[c]=Te(Vt(t,o,[h,c,this]))}return l}const a=new i(r);Xr(a,r);for(let l=0;l<r;++l){const u=st(s[l]);a[l]=Vt(t,o,[u,l,this])}return a}filter(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0],i=[];for(let u=0;u<r;++u){const c=st(s[u]);Vt(t,o,[c,u,this])&&k1(i,c)}const a=Kr(s,ft),l=new a(i);return Xr(l),l}reduce(t,...e){tt(this);const s=V(this),r=Q(s);if(r===0&&e.length===0)throw ht(sh);let o,i;e.length===0?(o=st(s[0]),i=1):(o=e[0],i=0);for(let a=i;a<r;++a)o=t(o,st(s[a]),a,this);return o}reduceRight(t,...e){tt(this);const s=V(this),r=Q(s);if(r===0&&e.length===0)throw ht(sh);let o,i;e.length===0?(o=st(s[r-1]),i=r-2):(o=e[0],i=r-1);for(let a=i;a>=0;--a)o=t(o,st(s[a]),a,this);return o}forEach(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i)Vt(t,o,[st(s[i]),i,this])}find(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return a}}findIndex(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return i}return-1}findLast(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=r-1;i>=0;--i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return a}}findLastIndex(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=r-1;i>=0;--i){const a=st(s[i]);if(Vt(t,o,[a,i,this]))return i}return-1}every(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i)if(!Vt(t,o,[st(s[i]),i,this]))return!1;return!0}some(t,...e){tt(this);const s=V(this),r=Q(s),o=e[0];for(let i=0;i<r;++i)if(Vt(t,o,[st(s[i]),i,this]))return!0;return!1}set(t,...e){tt(this);const s=V(this),r=Be(e[0]);if(r<0)throw Hi(Bi);if(t==null)throw ht(Mi);if(Xi(t))throw ht(Oi);if(es(t))return U1(V(this),V(t),r);if(qr(t)){const l=vt(t);if(Bs(l))throw ht(Rs)}const o=Q(s),i=Qn(t),a=Qi(i.length);if(r===1/0||a+r>o)throw Hi(Bi);for(let l=0;l<a;++l)s[l+r]=Te(i[l])}reverse(){tt(this);const t=V(this);return mh(t),this}toReversed(){tt(this);const t=V(this),e=new Xt(vt(t),$n(t),Q(t)),s=new ft(vt(Os(e))),r=V(s);return mh(r),s}fill(t,...e){tt(this);const s=V(this);return W1(s,Te(t),...Gr(e)),this}copyWithin(t,e,...s){tt(this);const r=V(this);return G1(r,t,e,...Gr(s)),this}sort(t){tt(this);const e=V(this),s=t!==void 0?t:Lh;return gh(e,(r,o)=>s(st(r),st(o))),this}toSorted(t){tt(this);const e=V(this);if(t!==void 0&&typeof t!="function")throw new ht(x1);const s=t!==void 0?t:Lh,r=new Xt(vt(e),$n(e),Q(e)),o=new ft(vt(Os(r))),i=V(o);return gh(i,(a,l)=>s(st(a),st(l))),o}slice(t,e){tt(this);const s=V(this),r=Kr(s,ft);if(r===ft){const p=new Xt(vt(s),$n(s),Q(s));return new ft(vt(Os(p,t,e)))}const o=Q(s),i=Be(t),a=e===void 0?o:Be(e);let l;i===-1/0?l=0:i<0?l=o+i>0?o+i:0:l=o<i?o:i;let u;a===-1/0?u=0:a<0?u=o+a>0?o+a:0:u=o<a?o:a;const c=u-l>0?u-l:0,h=new r(c);if(Xr(h,c),c===0)return h;const f=vt(s);if(Bs(f))throw ht(Rs);let d=0;for(;l<u;)h[d]=st(s[l]),++l,++d;return h}subarray(t,e){tt(this);const s=V(this),r=Kr(s,ft),o=new Xt(vt(s),$n(s),Q(s)),i=V1(o,t,e),a=new r(vt(i),$n(i),Q(i));return Xr(a),a}indexOf(t,...e){tt(this);const s=V(this),r=Q(s);let o=Be(e[0]);if(o===1/0)return-1;o<0&&(o+=r,o<0&&(o=0));for(let i=o;i<r;++i)if(Oe(s,i)&&st(s[i])===t)return i;return-1}lastIndexOf(t,...e){tt(this);const s=V(this),r=Q(s);let o=e.length>=1?Be(e[0]):r-1;if(o===-1/0)return-1;o>=0?o=o<r-1?o:r-1:o+=r;for(let i=o;i>=0;--i)if(Oe(s,i)&&st(s[i])===t)return i;return-1}includes(t,...e){tt(this);const s=V(this),r=Q(s);let o=Be(e[0]);if(o===1/0)return!1;o<0&&(o+=r,o<0&&(o=0));const i=Zn(t);for(let a=o;a<r;++a){const l=st(s[a]);if(i&&Zn(l)||l===t)return!0}return!1}join(t){tt(this);const e=V(this),s=Mh(e);return C1(s,t)}toLocaleString(...t){tt(this);const e=V(this),s=Mh(e);return _1(s,...Gr(t))}get[Ui](){if(es(this))return"Float16Array"}}Ms(ft,"BYTES_PER_ELEMENT",{value:ta}),Ms(ft,Hr,{}),ah(ft,qi);const Jr=ft.prototype;Ms(Jr,"BYTES_PER_ELEMENT",{value:ta}),Ms(Jr,_e,{value:Jr.values,writable:!0,configurable:!0}),ah(Jr,$t);function lx(n,t){return n.channels===t.channels}const Zr=8;class Bh{constructor(t,e,s){K(this,"_label");K(this,"_device");K(this,"_outputBuffers",{});K(this,"_pipeline");K(this,"_bindGroups",[]);K(this,"_needsUpdatePipeline",!0);K(this,"_inputs",[]);K(this,"_outputs",[]);K(this,"_uniforms",[]);K(this,"_uniformBuffers",{});K(this,"_width",10);K(this,"_height",10);K(this,"_execWidth");K(this,"_execHeight");K(this,"_csCode","");K(this,"_csMain");K(this,"_csDefine");this._label=t,this._device=e,this._csMain=s.csMain,this._csDefine=s.csDefine,this._inputs=s.inputs,this._outputs=s.outputs,this._uniforms=s.uniforms,s.uniforms.forEach(r=>{this._uniformBuffers[r.label]=e.createBuffer({label:this._label,size:r.data.byteLength,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST}),this._device.queue.writeBuffer(this._uniformBuffers[r.label],0,r.data)})}setSize(t,e){t=Math.ceil(t),e=Math.ceil(e);const s=t!==this._width||e!==this._height;this._width=t,this._height=e,s&&(this._resizeOutputBuffers(),this._needsUpdatePipeline=!0)}setExecuteSize(t,e){t=Math.ceil(t),e=Math.ceil(e),this._execWidth=t,this._execHeight=e}setOutputParams(t){this._updateOutputBuffers(t),this._needsUpdatePipeline=!0}setUniform(t,e){const s=this._uniformBuffers[t];this._device.queue.writeBuffer(s,0,e)}getOutputBuffer(t){return this._outputBuffers[t].buffer}dispose(){Object.keys(this._uniformBuffers).forEach(t=>{this._uniformBuffers[t].destroy()}),Object.keys(this._outputBuffers).forEach(t=>{this._outputBuffers[t].texture.destroy()})}_createBuffer(t){const e=this._width*this._height*4*4;return this._device.createBuffer({label:this._label,size:Math.max(e,80),usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST|GPUBufferUsage.COPY_SRC})}_resizeOutputBuffers(){const t=this._outputBuffers;for(const e in t){const{buffer:s,params:r}=t[e];s.destroy(),t[e].buffer=this._createBuffer(r)}}_updateOutputBuffers(t){var s,r;const e=this._outputBuffers;for(const o in t){const i=t[o];if(!lx(i,((s=e[o])==null?void 0:s.params)||{})){(r=e[o])==null||r.buffer.destroy();const a=this._createBuffer(i);e[o]={buffer:a,params:i}}}}_updatePipeline(t){if(!this._needsUpdatePipeline)return;this._needsUpdatePipeline=!1;const e=this._device,s=this._getFullCs(t);s!==this._csCode&&(this._csCode=s,this._pipeline=e.createComputePipeline({label:this._label,layout:"auto",compute:{module:e.createShaderModule({label:this._label,code:s}),entryPoint:"main"}}),this._updateBindGroups())}_getFullCs(t){const e=this._inputs,s=e.length>0;return`
|
|
3048
3048
|
${e.sort().map((o,i)=>`@group(0) @binding(${i}) var<storage, read> in_${o}: array<vec${t[o].channels}f>;`).join(`
|
|
3049
3049
|
`)}
|
|
3050
3050
|
${this._uniforms.map((o,i)=>`@group(${s?1:0}) @binding(${i}) var<uniform> ${o.label}: ${o.type};`).join(`
|
|
@@ -3053,11 +3053,11 @@ ${this._uniforms.map((o,i)=>`@group(${s?1:0}) @binding(${i}) var<uniform> ${o.la
|
|
|
3053
3053
|
${this._outputs.map((o,i)=>`@group(${s?2:1}) @binding(${i}) var<storage, read_write> out_${o}: array<vec${this._outputBuffers[o].params.channels}f>;`).join(`
|
|
3054
3054
|
`)}
|
|
3055
3055
|
${this._csDefine??""}
|
|
3056
|
-
@compute @workgroup_size(${
|
|
3056
|
+
@compute @workgroup_size(${Zr}, ${Zr}, 1)
|
|
3057
3057
|
fn main(@builtin(global_invocation_id) globalId: vec3u) {
|
|
3058
3058
|
${this._csMain}
|
|
3059
3059
|
}
|
|
3060
|
-
`}_updateBindGroups(){const t=[],e=this._device,s=this._inputs.length>0?1:0;this._uniforms.length>0&&(t[s]=e.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(s),entries:this._uniforms.map((r,o)=>({binding:o,resource:{buffer:this._uniformBuffers[r.label]}}))})),this._bindGroups=t}createPass(t,e){this._updatePipeline(e);const s=this._inputs.length>0;s&&(this._bindGroups[0]=this._device.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(0),entries:this._inputs.map((o,i)=>({binding:i,resource:{buffer:e[o].buffer}}))})),this._bindGroups[s?2:1]=this._device.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(s?2:1),entries:this._outputs.map((o,i)=>({binding:i,resource:{buffer:this._outputBuffers[o].buffer}}))});const r=t.beginComputePass();r.setPipeline(this._pipeline),this._bindGroups.forEach((o,i)=>{r.setBindGroup(i,o)}),r.dispatchWorkgroups(Math.ceil((this._execWidth??this._width)/
|
|
3060
|
+
`}_updateBindGroups(){const t=[],e=this._device,s=this._inputs.length>0?1:0;this._uniforms.length>0&&(t[s]=e.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(s),entries:this._uniforms.map((r,o)=>({binding:o,resource:{buffer:this._uniformBuffers[r.label]}}))})),this._bindGroups=t}createPass(t,e){this._updatePipeline(e);const s=this._inputs.length>0;s&&(this._bindGroups[0]=this._device.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(0),entries:this._inputs.map((o,i)=>({binding:i,resource:{buffer:e[o].buffer}}))})),this._bindGroups[s?2:1]=this._device.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(s?2:1),entries:this._outputs.map((o,i)=>({binding:i,resource:{buffer:this._outputBuffers[o].buffer}}))});const r=t.beginComputePass();r.setPipeline(this._pipeline),this._bindGroups.forEach((o,i)=>{r.setBindGroup(i,o)}),r.dispatchWorkgroups(Math.ceil((this._execWidth??this._width)/Zr),Math.ceil((this._execHeight??this._height)/Zr),1),r.end()}}const ea=1412.83765,na=1.64593172,sa=.431384981,ra=-.00294139609,oa=.192653254,ia=.00626026094,aa=.998620152,Fh=15794576e-13,zh=.0322087631,Uh=.00223151711,Wh=.370974749;function Gh(n){return n<=Fh?n=ea*n:n<=zh?n=na*Math.pow(n,sa)+ra:n=oa*Math.log(n+ia)+aa,n}function ux(n){return n<=Uh?n=n/ea:n<=Wh?n=Math.pow((n-ra)/na,1/sa):n=Math.exp((n-aa)/oa)-ia,n}const Vh=Gh(65504),qh=1/Vh,jh=Vh;class la{constructor(t,e,s,r){this.x=t,this.y=e,this.width=s,this.height=r}}function cx({data:n,channels:t}){let e=0;for(let i=0;i<n.length;i+=t){const a=n[i],l=n[i+1],u=n[i+2],c=.212671*a+.71516*l+.072169*u;e+=Math.log2(c+1e-4)}const s=n.length/t,r=e/s;return .18/Math.pow(2,r)}function hx({data:n,channels:t,inputScale:e}){const s=new Float32Array(n.length);s.set(n);for(let r=0;r<s.length;r+=t)for(let o=0;o<3;o++){let i=s[r+o]*e;s[r+o]=Gh(i)*qh}return s}function fx({data:n,channels:t,inputScale:e}){const s=new Float32Array(n.length);s.set(n);const r=1/e;for(let o=0;o<s.length;o+=t)for(let i=0;i<3;i++){let a=s[o+i]*jh;s[o+i]=ux(a)*r}return s}const Hh=`
|
|
3061
3061
|
const a = ${ea};
|
|
3062
3062
|
const b = ${na};
|
|
3063
3063
|
const c = ${sa};
|
|
@@ -3121,7 +3121,7 @@ out_color[outIdx] = vec4f(
|
|
|
3121
3121
|
// Pick the alpha
|
|
3122
3122
|
raw.a
|
|
3123
3123
|
);
|
|
3124
|
-
`}),this._inputPass.setOutputParams({color:{channels:3},albedo:{channels:3},normal:{channels:3}}),this._outputPass.setOutputParams({color:{channels:4}})}setImageSize(t,e){this._inputPass.setUniform("inputSize",new Float32Array([t,e])),this._outputPass.setUniform("imageSize",new Float32Array([t,e])),this._outputPass.setSize(t,e)}setInputTile(t){const e=this._inputPass,s=new Float32Array([t.width,t.height]);e.setUniform("inputOffset",new Float32Array([t.x,t.y])),e.setUniform("outputSize",s),e.setSize(s[0],s[1]),this._outputPass.setUniform("inputSize",s)}setOutputTile(t,e){const s=this._outputPass,r=new Float32Array([t.width,t.height]),o=t.x-e.x,i=t.y-e.y;s.setUniform("outputSize",r),s.setUniform("inputOffset",new Float32Array([o,i])),s.setUniform("outputOffset",new Float32Array([t.x,t.y])),s.setExecuteSize(r[0],r[1])}forward(t,e,s){const r=this._inputPass,o=this._device.createCommandEncoder();return r.createPass(o,{color:{buffer:t,channels:4},albedo:{buffer:e,channels:4},normal:{buffer:s,channels:4}}),this._device.queue.submit([o.finish()]),{color:r.getOutputBuffer("color"),albedo:r.getOutputBuffer("albedo"),normal:r.getOutputBuffer("normal")}}inverse(t,e){const r=this._device.createCommandEncoder(),o=this._outputPass;return o.createPass(r,{color:{buffer:t,channels:4},raw:{buffer:e,channels:4}}),this._device.queue.submit([r.finish()]),o.getOutputBuffer("color")}copyInputDataToOutput(t){const e=this._device.createCommandEncoder(),r=this._outputPass.getOutputBuffer("color");e.copyBufferToBuffer(t,0,r,0,r.size),this._device.queue.submit([e.finish()])}dispose(){this._outputPass.dispose(),this._inputPass.dispose()}}function Kh(n,t){const e=n.buffer;if(t==="Float32")return new Float32Array(n.buffer);const s=new ft(e),r=new Float32Array(s.length);for(let o=0;o<r.length;++o)r[o]=s[o];return r}function px(n,t){const[e,s,r,o]=t,i=new Float32Array(n.length);for(let a=0;a<e;++a)for(let l=0;l<s;++l)for(let u=0;u<r;++u)for(let c=0;c<o;++c){const h=a*s*r*o+l*r*o+u*o+c,f=u*o*s*e+c*s*e+l*e+a;i[f]=n[h]}return i}function ua(n,t){return Math.ceil(n/t)*t}function
|
|
3124
|
+
`}),this._inputPass.setOutputParams({color:{channels:3},albedo:{channels:3},normal:{channels:3}}),this._outputPass.setOutputParams({color:{channels:4}})}setImageSize(t,e){this._inputPass.setUniform("inputSize",new Float32Array([t,e])),this._outputPass.setUniform("imageSize",new Float32Array([t,e])),this._outputPass.setSize(t,e)}setInputTile(t){const e=this._inputPass,s=new Float32Array([t.width,t.height]);e.setUniform("inputOffset",new Float32Array([t.x,t.y])),e.setUniform("outputSize",s),e.setSize(s[0],s[1]),this._outputPass.setUniform("inputSize",s)}setOutputTile(t,e){const s=this._outputPass,r=new Float32Array([t.width,t.height]),o=t.x-e.x,i=t.y-e.y;s.setUniform("outputSize",r),s.setUniform("inputOffset",new Float32Array([o,i])),s.setUniform("outputOffset",new Float32Array([t.x,t.y])),s.setExecuteSize(r[0],r[1])}forward(t,e,s){const r=this._inputPass,o=this._device.createCommandEncoder();return r.createPass(o,{color:{buffer:t,channels:4},albedo:{buffer:e,channels:4},normal:{buffer:s,channels:4}}),this._device.queue.submit([o.finish()]),{color:r.getOutputBuffer("color"),albedo:r.getOutputBuffer("albedo"),normal:r.getOutputBuffer("normal")}}inverse(t,e){const r=this._device.createCommandEncoder(),o=this._outputPass;return o.createPass(r,{color:{buffer:t,channels:4},raw:{buffer:e,channels:4}}),this._device.queue.submit([r.finish()]),o.getOutputBuffer("color")}copyInputDataToOutput(t){const e=this._device.createCommandEncoder(),r=this._outputPass.getOutputBuffer("color");e.copyBufferToBuffer(t,0,r,0,r.size),this._device.queue.submit([e.finish()])}dispose(){this._outputPass.dispose(),this._inputPass.dispose()}}function Kh(n,t){const e=n.buffer;if(t==="Float32")return new Float32Array(n.buffer);const s=new ft(e),r=new Float32Array(s.length);for(let o=0;o<r.length;++o)r[o]=s[o];return r}function px(n,t){const[e,s,r,o]=t,i=new Float32Array(n.length);for(let a=0;a<e;++a)for(let l=0;l<s;++l)for(let u=0;u<r;++u)for(let c=0;c<o;++c){const h=a*s*r*o+l*r*o+u*o+c,f=u*o*s*e+c*s*e+l*e+a;i[f]=n[h]}return i}function ua(n,t){return Math.ceil(n/t)*t}function Qr(n){return n.data instanceof GPUBuffer}const mx=174,ca=16,to=ua(mx/2,ca);class Yh{constructor(t,e,s={}){K(this,"_tfModel");K(this,"_device");K(this,"_tileWidth",0);K(this,"_tileHeight",0);K(this,"_tileOverlapX",0);K(this,"_tileOverlapY",0);K(this,"_aux",!1);K(this,"_hdr",!1);K(this,"_dataProcessGPU");K(this,"_maxTileSize");this._tensors=t,this._backend=e,this._aux=s.aux||!1,this._hdr=s.hdr||!1,this._maxTileSize=s.maxTileSize??512,this._device=this._backend.device}_createConv(t,e,s){const r=this._tensors.get(t+".weight"),o=this._tensors.get(t+".bias"),i=r.desc.dims,a=er(px(Kh(r.data,r.desc.dataType),i),[i[2],i[3],i[1],i[0]],"float32"),l=Rt(Kh(o.data,o.desc.dataType),"float32");return new Yn({name:t,filters:r.desc.dims[0],kernelSize:r.desc.dims.slice(2,4),useBias:!0,activation:s,padding:"same",weights:[a,l],trainable:!1}).apply(e)}_createConcatConv(t,e,s){return this._createConv(t,new ki({trainable:!1,axis:3}).apply([e,s]),"relu")}_createPooling(t){return new Ii({name:t.name+"/pooling",poolSize:[2,2],strides:[2,2],padding:"same",trainable:!1}).apply(t)}_addUpsamplingLayer(t){return new vi({name:t.name+"/upsampling",size:[2,2],trainable:!1}).apply(t)}getDevice(){return this._device}buildModel(){const e=3+(this._aux?6:0),s=this._getTileSizeWithOverlap(),r=Jw({shape:[s.height,s.width,e],dtype:"float32"}),o=this._createConv("enc_conv0",r,"relu"),i=this._createPooling(this._createConv("enc_conv1",o,"relu")),a=this._createPooling(this._createConv("enc_conv2",i,"relu")),l=this._createPooling(this._createConv("enc_conv3",a,"relu")),u=this._createPooling(this._createConv("enc_conv4",l,"relu")),c=this._createConv("enc_conv5a",u,"relu"),h=this._addUpsamplingLayer(this._createConv("enc_conv5b",c,"relu")),f=this._createConcatConv("dec_conv4a",h,l),d=this._addUpsamplingLayer(this._createConv("dec_conv4b",f,"relu")),p=this._createConcatConv("dec_conv3a",d,a),g=this._addUpsamplingLayer(this._createConv("dec_conv3b",p,"relu")),m=this._createConcatConv("dec_conv2a",g,i),b=this._addUpsamplingLayer(this._createConv("dec_conv2b",m,"relu")),y=this._createConcatConv("dec_conv1a",b,r),S=this._createConv("dec_conv1b",y,"relu"),x=this._createConv("dec_conv0",S,"relu");this._tfModel=new Mr({inputs:[r],outputs:x})}_updateModel(t,e){const s=this._maxTileSize;let r=s,o=s,i=to,a=to;t<s+to*2&&(r=ua(t,ca),i=0),e<s+to*2&&(o=ua(e,ca),a=0),(r!==this._tileWidth||o!==this._tileHeight||i!==this._tileOverlapX||a!==this._tileOverlapY||!this._tfModel)&&(this._tileWidth=r,this._tileHeight=o,this._tileOverlapX=i,this._tileOverlapY=a,this._tfModel&&this._tfModel.dispose(),this.buildModel())}_getTileSizeWithOverlap(){return{width:this._tileWidth+2*this._tileOverlapX,height:this._tileHeight+2*this._tileOverlapY}}_processImageData(t,e,s,r){const o=t.data,i=o.length/4,a=this._aux?9:3,l=new Float32Array(i*a);if(e&&!s||s&&!e)throw new Error("Normal map and albedo map are both required");if(e&&s&&(e.width!==s.width||e.height!==s.height||t.width!==e.width||t.height!==e.height))throw new Error("Image size mismatch");const u=e==null?void 0:e.data,c=s==null?void 0:s.data;for(let h=0;h<o.length;h+=4){const f=h/4*a;for(let d=0;d<3;d++)r?l[f+d]=o[h+d]:l[f+d]=o[h+d]/255,u&&(l[f+d+3]=u[h+d]/255),c&&(l[f+d+6]=c[h+d]/255)}return l}_readTile(t,e,s,r){const o=new Float32Array(s.width*s.height*e);for(let i=0;i<s.height;i++)for(let a=0;a<s.width;a++){const l=((i+s.y)*r+(a+s.x))*e,u=(i*s.width+a)*e;for(let c=0;c<e;c++)o[u+c]=t[l+c]}return o}_writeTile(t,e,s,r,o,i){const{data:a,width:l}=t,u=s.x-e.x,c=s.y-e.y;for(let h=0;h<s.height;h++)for(let f=0;f<s.width;f++){const d=((h+c)*o+f+u)*3,p=((h+s.y)*l+(f+s.x))*4;for(let g=0;g<3;g++)i?a[p+g]=r[d+g]:a[p+g]=Math.min(Math.max(r[d+g]*255,0),255);t.data[p+3]=i?1:255}}_executeTile(t,e,s,r,o,i,a,l){const u=this._aux?9:3,c=this._tileOverlapX,h=this._tileOverlapY;let f=this._getTileSizeWithOverlap(),d={width:this._tileWidth,height:this._tileHeight},p=r>0?r*d.width-c:0,g=Math.min(p+f.width,i);p=Math.max(g-f.width,0);let m=o>0?o*d.height-h:0,b=Math.min(m+f.height,a);m=Math.max(b-f.height,0);const y=Math.min(f.width,i),S=Math.min(f.height,a),x=i<d.width||a<d.height,$=new la(p,m,y,S);let E,D=1;const _=this._device;let T=this._dataProcessGPU;if(t instanceof Float32Array){let F=this._readTile(t,u,$,i);l&&(D=cx({data:F,channels:9}),F=hx({data:F,channels:9,inputScale:D})),E=er(F,[1,S,y,u],"float32")}else{if(!l)throw new Error("Only hdr is supported for webgpu data.");T||(T=this._dataProcessGPU=new dx(_)),T.setImageSize(i,a),T.setInputTile($),r===0&&o===0&&T.copyInputDataToOutput(t.color);const{color:F,albedo:G,normal:q}=T.forward(t.color,t.albedo,t.normal),De=[1,S,y,4],qt=[F,G,q].map(jt=>{const Xe=er({buffer:jt,zeroCopy:!0},De),os=hs(Xe,[0,0,0,0],[1,S,y,3]);return Xe.dispose(),os});E=jp(qt,3),qt.forEach(jt=>jt.dispose())}if(x){const F=E;E=Mp(F,[[0,0],[0,f.height-a],[0,f.width-i],[0,0]],"reflect"),F.dispose()}const P=this._tfModel.predict(E);E.dispose();const B=Math.min(d.width,i),Y=Math.min(d.height,a),j=new la(r*B,o*Y,B,Y);if(j.width=Math.min(j.width,i-j.x),j.height=Math.min(j.height,a-j.y),t instanceof Float32Array){let F=P.dataSync();l&&(F=fx({data:F,channels:3,inputScale:D})),this._writeTile(s,$,j,F,f.width,l);for(let G=0;G<Y;G++)for(let q=0;q<B;q++){const De=(G*B+q)*4,qt=((G+j.y)*i+(q+j.x))*4;for(let jt=0;jt<4;jt++)e.data[De+jt]=s.data[qt+jt]}P.dispose()}else{T.setOutputTile(j,$);const F=zp(P,[[0,0],[0,0],[0,0],[0,1]]),G=T.inverse(F.dataToGPU().buffer,t.color);return P.dispose(),F.dispose(),G}}progressiveExecute({color:t,albedo:e,normal:s,done:r,progress:o}){if(this._aux&&(!e||!s))throw new Error("Normal map and albedo map are both required");if(!this._aux&&(e||s))throw new Error("Normal map and albedo map are not required");const i=t.width,a=t.height;this._updateModel(i,a);const l=this._hdr||!1;let u;Qr(t)||(u=this._processImageData(t,e,s,l));const c=this._tileWidth,h=this._tileHeight,f=Math.ceil(a/h),d=Math.ceil(i/c);function p(S,x){return l?{data:new Float32Array(S*x*4),width:S,height:x}:new ImageData(S,x)}const g=Qr(t)?void 0:p(i,a),m=Qr(t)?void 0:p(Math.min(c,i),Math.min(h,a));let b=!1;const y=(S,x)=>{if(b)return;let $;$=this._executeTile(Qr(t)?{color:t.data,albedo:e.data,normal:s.data}:u,m,g,S,x,i,a,l);const E=g||{data:$,width:i,height:a};o==null||o(E,m,new la(S*c,x*h,c,h),S+x*d,d*f),S+1<d||x+1<f?requestAnimationFrame(()=>{S+1<d?y(S+1,x):x+1<f&&y(0,x+1)}):r(E)};return y(0,0),()=>{b=!0}}dispose(){var t,e;(t=this._tfModel)==null||t.dispose(),(e=this._dataProcessGPU)==null||e.dispose()}}/**
|
|
3125
3125
|
* @license
|
|
3126
3126
|
* Copyright 2019 Google LLC. All Rights Reserved.
|
|
3127
3127
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3136,7 +3136,7 @@ out_color[outIdx] = vec4f(
|
|
|
3136
3136
|
* See the License for the specific language governing permissions and
|
|
3137
3137
|
* limitations under the License.
|
|
3138
3138
|
* =============================================================================
|
|
3139
|
-
*/const
|
|
3139
|
+
*/const Jt=W();Jt.registerFlag("WEBGPU_DEFERRED_SUBMIT_BATCH_SIZE",()=>15),Jt.registerFlag("WEBGPU_CPU_FORWARD",()=>!0),Jt.registerFlag("WEBGPU_MATMUL_PROGRAM_TYPE",()=>-1),Jt.registerFlag("WEBGPU_USE_NAIVE_CONV2D_TRANSPOSE",()=>!0),Jt.registerFlag("WEBGPU_USE_LOW_POWER_GPU",()=>!1),Jt.registerFlag("WEBGPU_CPU_HANDOFF_SIZE_THRESHOLD",()=>1e3),Jt.registerFlag("WEBGPU_USE_PROFILE_TOOL",()=>!1),Jt.registerFlag("WEBGPU_IMPORT_EXTERNAL_TEXTURE",()=>!0),Jt.registerFlag("WEBGPU_USE_NAIVE_CONV2D_DEBUG",()=>!1),Jt.registerFlag("WEBGPU_THRESHOLD_TO_INCREASE_WORKGROUPS_FOR_MATMUL",()=>-1),Jt.registerFlag("WEBGPU_CONV_SEPARATE_IM2COL_SHADER",()=>!1),Jt.registerFlag("WEBGPU_PRINT_SHADER",()=>""),Jt.registerFlag("WEBGPU_ENGINE_COMPILE_ONLY",()=>!1);/**
|
|
3140
3140
|
* @license
|
|
3141
3141
|
* Copyright 2022 Google LLC.
|
|
3142
3142
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3223,7 +3223,7 @@ out_color[outIdx] = vec4f(
|
|
|
3223
3223
|
* See the License for the specific language governing permissions and
|
|
3224
3224
|
* limitations under the License.
|
|
3225
3225
|
* =============================================================================
|
|
3226
|
-
*/var
|
|
3226
|
+
*/var eo;(function(n){n[n.FROM_PIXELS=0]="FROM_PIXELS",n[n.DRAW=1]="DRAW"})(eo||(eo={}));const Sx=(n,t,e,s,r)=>{const o={dtype:s.dtype,shape:s.shape},i=vx(e,o,t),a=n.createShaderModule({code:i,label:t.constructor.name});let l=W().get("WEBGPU_PRINT_SHADER");if(l!==""){l=l.toLowerCase();const u=l.split(",");(l==="all"||u.some(c=>t.shaderKey.toLowerCase().includes(c)))&&(console.group(t.shaderKey),console.debug(i),console.groupEnd())}return r?n.createComputePipelineAsync({compute:{module:a,entryPoint:"_start"},label:t.constructor.name,layout:"auto"}):n.createComputePipeline({compute:{module:a,entryPoint:"_start"},label:t.constructor.name,layout:"auto"})},H=(n,t="f32")=>{switch(n){case 1:return`${t}`;case 2:return`vec2<${t}>`;case 3:return`vec3<${t}>`;case 4:return`vec4<${t}>`;default:throw new Error(`${n}-component ${t} is not supported.`)}};function Tt(n){if(n<=1)return"i32";if(n===2)return"vec2<i32>";if(n===3)return"vec3<i32>";if(n===4)return"vec4<i32>";if(n===5)return"vec5";if(n===6)return"vec6";throw Error(`GPU for rank ${n} is not yet supported`)}function In(n){if(n===0)return"x";if(n===1)return"y";if(n===2)return"z";if(n===3)return"w";if(n===4)return"u";if(n===5)return"v";throw Error(`Index ${n} is not yet supported`)}function wt(...n){let t;switch(n.length){case 0:t=`
|
|
3227
3227
|
fn main()
|
|
3228
3228
|
`;break;case 1:t=`
|
|
3229
3229
|
fn main(${n[0]} : i32)
|
|
@@ -3258,7 +3258,7 @@ out_color[outIdx] = vec4f(
|
|
|
3258
3258
|
localIndex);
|
|
3259
3259
|
`}
|
|
3260
3260
|
}
|
|
3261
|
-
`),e.pixelsOpType!=null){const p=e.pixelsOpType===
|
|
3261
|
+
`),e.pixelsOpType!=null){const p=e.pixelsOpType===eo.FROM_PIXELS?`@group(0) @binding(0) var<storage, read_write> result: array<${ns(t.dtype,e.outputComponent)}>;`:`@group(0) @binding(1) var<storage, read> inBuf : array<${ns(n[0].dtype,e.outputComponent)}>;`,g=t.shape.length===3?"vec2<i32>":"i32";s.push(`
|
|
3262
3262
|
struct Uniform {
|
|
3263
3263
|
outShapeStrides : ${g},
|
|
3264
3264
|
size : i32,
|
|
@@ -3282,7 +3282,7 @@ out_color[outIdx] = vec4f(
|
|
|
3282
3282
|
`);const u=_x(t.shape,e.dispatchLayout),c=[tf,s.join(`
|
|
3283
3283
|
`)+Ax,ha(t.shape),u,Tx(t.shape.length)];e.atomic||c.push(Nx(t.shape,t.dtype,e.outputComponent)),e.variableNames.forEach((p,g)=>{c.push(`${ha(n[g].shape,p)}`)});const h=n.map((p,g)=>kx(p,t.shape,e.variableComponents?e.variableComponents[g]:e.outputComponent,e.dispatchLayout.x.length===t.shape.length)).join(`
|
|
3284
3284
|
`);c.push(h),c.push(e.getUserCode());const f=nf(e);return c.push(Qh(f,e)),c.join(`
|
|
3285
|
-
`)}function Ix(n,t,e){let s=n.shaderKey;if(n.pixelsOpType!=null)return s;const r=[],o=[];t.forEach(c=>{r.push(c.shape),o.push(c.dtype)}),r.push(e.shape),o.push(e.dtype);const i=t.map(c=>
|
|
3285
|
+
`)}function Ix(n,t,e){let s=n.shaderKey;if(n.pixelsOpType!=null)return s;const r=[],o=[];t.forEach(c=>{r.push(c.shape),o.push(c.dtype)}),r.push(e.shape),o.push(e.dtype);const i=t.map(c=>ir(c.shape,e.shape)),a=t.map(c=>Zt(c.shape,e.shape)).join("_"),l=i.map(c=>c.join("_")).join(";"),u=ef(n)?"flatDispatch":"";return s+="_"+(n.workgroupSize?n.workgroupSize.join(","):"")+r.map(c=>c.length).join(",")+o.join(",")+n.variableNames.join(",")+l+a+u,s}const tf=`
|
|
3286
3286
|
struct vec5 {x: i32, y: i32, z: i32, w: i32, u: i32};
|
|
3287
3287
|
struct vec6 {x: i32, y: i32, z: i32, w: i32, u: i32, v: i32};
|
|
3288
3288
|
|
|
@@ -3336,10 +3336,10 @@ out_color[outIdx] = vec4f(
|
|
|
3336
3336
|
fn isinf(val: f32) -> bool {
|
|
3337
3337
|
return abs(val) == uniforms.INFINITY;
|
|
3338
3338
|
}
|
|
3339
|
-
`;function ha(n,t=""){const e=n.length,s=t!==""?`get${t.charAt(0).toUpperCase()+t.slice(1)}CoordsFromIndex`:"getCoordsFromIndex",r=t!==""?`${t.charAt(0).toLowerCase()+t.slice(1)}ShapeStrides`:"outShapeStrides";if(e<=1)return`fn ${s}(index : i32) -> i32 { return index; }`;const o=
|
|
3339
|
+
`;function ha(n,t=""){const e=n.length,s=t!==""?`get${t.charAt(0).toUpperCase()+t.slice(1)}CoordsFromIndex`:"getCoordsFromIndex",r=t!==""?`${t.charAt(0).toLowerCase()+t.slice(1)}ShapeStrides`:"outShapeStrides";if(e<=1)return`fn ${s}(index : i32) -> i32 { return index; }`;const o=Kt(n),i=Tt(e),a=[];for(let u=0;u<e;u++)a.push(`d${u}`);if(o.length===1)return` fn ${s}(index : i32) -> vec2<i32> {
|
|
3340
3340
|
let d0 = index / uniforms.${r}; let d1 = index - d0 * uniforms.${r};
|
|
3341
3341
|
return vec2<i32>(d0, d1);
|
|
3342
|
-
}`;let l;return l="var index2 = index;"+o.map((u,c)=>{const h=`let ${a[c]} = index2 / uniforms.${r}.${
|
|
3342
|
+
}`;let l;return l="var index2 = index;"+o.map((u,c)=>{const h=`let ${a[c]} = index2 / uniforms.${r}.${In(c)}`,f=c===o.length-1?`let ${a[c+1]} = index2 - ${a[c]} * uniforms.${r}.${In(c)}`:`index2 = index2 - ${a[c]} * uniforms.${r}.${In(c)}`;return`${h}; ${f};`}).join(""),`
|
|
3343
3343
|
fn ${s}(index : i32) -> ${i} {
|
|
3344
3344
|
${l}
|
|
3345
3345
|
return ${i}(${a.join(",")});
|
|
@@ -3361,7 +3361,7 @@ out_color[outIdx] = vec4f(
|
|
|
3361
3361
|
fn ${i}Coords(coords : ${u}) -> ${H(e)} {
|
|
3362
3362
|
return ${H(e)}(${r}[${l>1?"getOutputIndexFromCoords(coords)":"coords"}${e===1?"":` / ${e}`}]);
|
|
3363
3363
|
}
|
|
3364
|
-
`;const c=
|
|
3364
|
+
`;const c=ir(n.shape,t),h=l-a;let f="";if(a===0)return`
|
|
3365
3365
|
fn ${i}Index(globalIndex : i32) -> ${H(e)}{
|
|
3366
3366
|
return get${o}();
|
|
3367
3367
|
}
|
|
@@ -3369,8 +3369,8 @@ out_color[outIdx] = vec4f(
|
|
|
3369
3369
|
fn ${i}Coords(coords : ${u}) -> ${H(e)}{
|
|
3370
3370
|
return get${o}();
|
|
3371
3371
|
}
|
|
3372
|
-
`;l<2&&c.length>=1?f="coords = 0;":f=c.map(m=>`coords.${
|
|
3373
|
-
`);let d="";if(l<2&&a>0)d="coords";else if(l>1){const m=Tt(a),b=n.shape.map((y,S)=>`coords.${
|
|
3372
|
+
`;l<2&&c.length>=1?f="coords = 0;":f=c.map(m=>`coords.${In(m+h)} = 0;`).join(`
|
|
3373
|
+
`);let d="";if(l<2&&a>0)d="coords";else if(l>1){const m=Tt(a),b=n.shape.map((y,S)=>`coords.${In(S+h)}`).join(", ");d=`${m}(${b})`}else d="coords";const p=`uniforms.${r.charAt(0).toLowerCase()+r.slice(1)}Shape`,g=`${a}D`;return`
|
|
3374
3374
|
fn ${i}Index(globalIndex : i32) -> ${H(e)} {
|
|
3375
3375
|
var coords = getCoordsFromIndex(globalIndex);
|
|
3376
3376
|
${f}
|
|
@@ -3453,7 +3453,7 @@ out_color[outIdx] = vec4f(
|
|
|
3453
3453
|
* See the License for the specific language governing permissions and
|
|
3454
3454
|
* limitations under the License.
|
|
3455
3455
|
* =============================================================================
|
|
3456
|
-
*/const
|
|
3456
|
+
*/const An=n=>{let t=1;for(let e=0;e<n.length;e++)t*=n[e];return t};function It(n,t,e=[1,1,1],s=[1,1,1]){const[r,o,i]=[Math.ceil(An(n.x.map(a=>t[a]))/(e[0]*s[0])),n.y?Math.ceil(An(n.y.map(a=>t[a]))/(e[1]*s[1])):1,n.z?Math.ceil(An(n.z.map(a=>t[a]))/(e[2]*s[2])):1];return[r,o,i]}function Rx(n,t,e,s=!1){const r=[8,8,1],o=[4,4,1];return s||(n<=8&&(o[1]=1),t<=16&&e<=16&&(r[0]=4)),{workgroupSize:r,elementsPerThread:o}}function Px(n,t,e=!1){if(e)return[8,8,1];const s=An(n.x.map(o=>t[o])),r=An(n.y.map(o=>t[o]));return s<=4?[4,16,1]:r<=4?[16,4,1]:[16,16,1]}function Lx(n,t,e=!1){if(e)return[4,4,1];const s=An(n.x.map(o=>t[o])),r=An(n.y.map(o=>t[o]));return s<=4?[1,2,1]:r<=4?[2,1,1]:[2,2,1]}function ae(n){return{x:n.map((t,e)=>e)}}function sf(n){if(n==="float32"||n==="int32"||n==="bool"||n==="string")return 4;if(n==="complex64")return 8;throw new Error(`Unknown dtype ${n}`)}function rf(){return!!(typeof globalThis<"u"&&globalThis.navigator&&globalThis.navigator.gpu)}var Ne;(function(n){n[n.MatMulReduceProgram=0]="MatMulReduceProgram",n[n.MatMulSplitKProgram=1]="MatMulSplitKProgram",n[n.MatMulSmallOutputSizeProgram=2]="MatMulSmallOutputSizeProgram",n[n.MatMulPackedProgram=3]="MatMulPackedProgram",n[n.MatMulMax=4]="MatMulMax"})(Ne||(Ne={}));/**
|
|
3457
3457
|
* @license
|
|
3458
3458
|
* Copyright 2019 Google LLC. All Rights Reserved.
|
|
3459
3459
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3468,7 +3468,7 @@ out_color[outIdx] = vec4f(
|
|
|
3468
3468
|
* See the License for the specific language governing permissions and
|
|
3469
3469
|
* limitations under the License.
|
|
3470
3470
|
* =============================================================================
|
|
3471
|
-
*/const Mx=W().getNumber("WEBGPU_CPU_HANDOFF_SIZE_THRESHOLD"),Ox=(n,t)=>{const e=n.limits.maxComputeWorkgroupsPerDimension,s=t.dispatchLayout,r=t.dispatch;if(r.every(i=>i<=e))return r;w(r[0]>e&&s.y===void 0&&s.z===void 0,()=>"Dispatch size exceeds WebGPU limits in Y or Z dimension.");let o=Math.ceil(Math.sqrt(r[0]));return o>e?(o=Math.ceil(Math.cbrt(r[0])),w(o<=e,()=>"Total dispatch size exceeds WebGPU maximum."),[o,o,o]):[o,o,1]};class Bs extends ba{nextDataId(){return Bs.nextDataId++}constructor(t,e){if(super(),this.commandQueueOwnedIds=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())throw new Error("WebGPU is not supported on this device");this.pipelineCache={},this.device=t,this.queue=t.queue,this.commandEncoder=null,this.computePassEncoder=null,this.adapterInfo=new gx(e),this.supportTimestampQuery=this.device.features.has("timestamp-query"),this.thresholdToIncreaseWorkgroups=this.adapterInfo.intelGPUGeneration>=12?16:8,this.bufferManager=new bx(this.device),this.textureManager=new yx(this.device),this.tensorMap=new vf(this,To()),W().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({device:t,format:"bgra8unorm"}),document.body.appendChild(this.dummyCanvas))}floatPrecision(){return 32}disposeData(t,e=!1){if(!this.tensorMap.has(t))return!0;const s=this.tensorMap.get(t);return e?s.refCount=0:s.refCount--,s.refCount>0?!1:(s.complexTensorInfos!=null&&(this.disposeData(s.complexTensorInfos.real.dataId),this.disposeData(s.complexTensorInfos.imag.dataId)),this.commandQueueOwnedIds.has(t)?(this.tensorDataPendingDisposal.push(t),!0):(this.releaseResource(t),this.tensorMap.delete(t),!0))}memory(){return{numBytesInGPU:this.bufferManager.numBytesUsed,numBytesAllocatedInGPU:this.bufferManager.numBytesAllocated,unreliable:!1}}releaseResource(t){const e=this.tensorMap.get(t);if(!(!e||!e.resource)){if(e.external){e.resource=null;return}e.resource instanceof GPUBuffer?this.bufferManager.releaseBuffer(e.resource):e.resource instanceof GPUTexture&&this.textureManager.releaseTexture(e.resource),e.resource=null}}refCount(t){return this.tensorMap.has(t)?this.tensorMap.get(t).refCount:0}incRef(t){const e=this.tensorMap.get(t);e.refCount++}decRef(t){if(this.tensorMap.has(t)){const e=this.tensorMap.get(t);e.refCount--}}write(t,e,s){if(s==="complex64"&&t!=null)throw new Error("Cannot write to a complex64 dtype. Please use tf.complex(real, imag).");const r={id:this.nextDataId()};return this.tensorMap.set(r,{dtype:s,shape:e,values:t,refCount:1}),r}move(t,e,s,r,o){if(r==="complex64")throw new Error("Cannot write to a complex64 dtype. Please use tf.complex(real, imag).");this.tensorMap.set(t,{dtype:r,shape:s,values:e,refCount:o})}submitQueue(){this.queue.submit([this.commandEncoder.finish()]),this.commandEncoder=null,this.dispatchCountInPass=0,this.commandQueueOwnedIds=new WeakSet,this.tensorDataPendingDisposal.forEach(t=>{this.releaseResource(t),this.tensorMap.delete(t)}),this.uniformPendingDisposal.forEach(t=>this.bufferManager.releaseBuffer(t)),this.stagingPendingDisposal.forEach(t=>this.bufferManager.releaseBuffer(t,!1)),this.tensorDataPendingDisposal=[],this.uniformPendingDisposal=[],this.stagingPendingDisposal=[]}ensureCommandEncoderReady(){this.commandEncoder||(this.commandEncoder=this.device.createCommandEncoder())}endComputePassEncoder(){this.computePassEncoder&&(this.computePassEncoder.end(),this.computePassEncoder=null)}async checkCompileCompletionAsync(){let t;try{t=await Promise.all(Object.values(this.pipelineCache))}catch(e){throw new Error(e.message)}Object.keys(this.pipelineCache).map((e,s)=>{this.pipelineCache[e]=t[s]})}async getBufferData(t){if(W().getBool("WEBGPU_ENGINE_COMPILE_ONLY"))return console.warn("The data may be invalid since WEBGPU_ENGINE_COMPILE_ONLY is true, this can only be called when WEBGPU_ENGINE_COMPILE_ONLY is false"),null;const e=t.size,s=this.bufferManager.acquireBuffer(e,GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ);this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(t,0,s,0,e),this.submitQueue(),await s.mapAsync(GPUMapMode.READ);const r=s.getMappedRange().slice(0);return s.unmap(),s!=null&&this.bufferManager.releaseBuffer(s),W().getBool("WEBGPU_USE_PROFILE_TOOL")&&(w(this.dummyContext!==void 0,()=>"Fail to get context for profiling tool"),this.dummyContext.getCurrentTexture()),r}convertAndCacheOnCPU(t,e){const s=this.tensorMap.get(t);return s.values=e,s.values}readSync(t){const e=this.tensorMap.get(t),{values:s,complexTensorInfos:r}=e;if(s!=null||e.dtype==="string")return s;if(e.dtype==="complex64"){const g=this.readSync(r.real.dataId),m=this.readSync(r.imag.dataId),b=uo(Zl(g,m).buffer,"float32");return this.convertAndCacheOnCPU(t,b),b}this.hasReadSyncWarned||(this.hasReadSyncWarned=!0,console.warn("The performance of synchronously reading data from GPU to CPU is poor on the webgpu backend, please use asynchronous APIs instead."));const o=["opaque","premultiplied"],i=e.resource,a=i.size;w(a%4===0,()=>"Because there is 4 bytes for one pixel, buffer size must be multiple of 4.");const l=a/4,u=new ArrayBuffer(a),c=256,h=256,f=o.map(g=>new OffscreenCanvas(c,h)),d=new OffscreenCanvas(c,h);this.endComputePassEncoder(),f.map((g,m)=>{const b=g.getContext("webgpu");return b.configure({device:this.device,format:"bgra8unorm",usage:GPUTextureUsage.COPY_DST,alphaMode:o[m]}),b.getCurrentTexture()}).map((g,m)=>{const b=c*4,y=(_,T,P)=>{this.ensureCommandEncoderReady(),this.commandEncoder.copyBufferToTexture({buffer:i,bytesPerRow:b,offset:P},{texture:g},{width:_,height:T}),this.submitQueue();const B=d.getContext("2d",{willReadFrequently:!0});B.clearRect(0,0,_,T),B.drawImage(f[m],0,0);const Y=B.getImageData(0,0,_,T).data,j=o[m],F=new Uint8ClampedArray(u,P,_*T*4);for(let G=0;G<F.length;G+=4)if(j==="premultiplied")F[G+3]=Y[G+3];else{const q=Y[G];F[G]=Y[G+2],F[G+1]=Y[G+1],F[G+2]=q}},S=Math.floor(l/(c*h));let x=c,$=h,E=0;for(let _=0;_<S;_++)y(x,$,E),E+=c*h*4;const D=l%(c*h);$=Math.floor(D/c),$>0&&(y(x,$,E),E+=$*(c*4)),x=D%c,x>0&&y(x,1,E)});const p=uo(u,e.dtype);return this.convertAndCacheOnCPU(t,p),p}async read(t){if(!this.tensorMap.has(t))throw new Error(`Tensor ${t} was not registered!`);const e=this.tensorMap.get(t),{values:s}=e;if(s!=null)return s;let r;if(e.dtype==="complex64"){const o=await Promise.all([this.read(e.complexTensorInfos.real.dataId),this.read(e.complexTensorInfos.imag.dataId)]),i=o[0],a=o[1];r=Zl(i,a)}else{const o=await this.getBufferData(e.resource);r=uo(o,e.dtype)}return this.convertAndCacheOnCPU(t,r),r}copyBuffer(t){const e=t.size,s=t.usage,r=this.bufferManager.acquireBuffer(e,s);return this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(t,0,r,0,e),this.submitQueue(),r}createTensorFromGPUData(t,e,s){let r=t.buffer;if(s==="complex64")throw new Error("Cannot write to a complex64 dtype. ");const o={id:this.nextDataId()};this.tensorMap.set(o,{dtype:s,shape:e,values:null,refCount:1,external:t.zeroCopy});const i=this.tensorMap.get(o),a=sf(i.dtype)*z(i.shape);if(t.buffer.size<a)throw new Error(`GPUBuffer size(${t.buffer.size}) is smaller than tensor size(${a})!`);if((t.buffer.usage&(GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC))!==(GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC))throw new Error("GPUBuffer.usage should include GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC!");return t.zeroCopy!==!0&&(r=this.copyBuffer(r)),i.resource=r,To().makeTensorFromDataId(o,e,s,this)}readToGPU(t){const e=this.tensorMap.get(t),{values:s,dtype:r,shape:o,resource:i}=e;if(r==="complex64")throw new Error("Does not support reading buffer for complex64 dtype.");if(i==null)throw s!=null?new Error("Data is not on GPU but on CPU."):new Error("There is no data on GPU or CPU.");const a=i,l=a.size,u=a.usage,c=this.bufferManager.acquireBuffer(l,u);this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(i,0,c,0,l),this.submitQueue();const h=this.makeTensorInfo(o,r),f=To().makeTensorFromTensorInfo(h),d=this.tensorMap.get(h.dataId);return d.resource=c,{tensorRef:f,buffer:c}}bufferSync(t){const e=this.readSync(t.dataId);if(t.dtype==="string")try{const s=e.map(r=>Ks(r));return xt(t.shape,t.dtype,s)}catch{throw new Error("Failed to decode encoded string bytes into utf-8")}return xt(t.shape,t.dtype,e)}async time(t){!this.supportTimestampQuery&&!this.hasTimestampQueryWarned&&(console.warn("This device doesn't support timestamp-query extension. Start Chrome browser with flag --enable-dawn-features=allow_unsafe_apis to try it again. Otherwise, zero will be shown for the kernel time when profiling mode is enabled."),this.hasTimestampQueryWarned=!0);const e=this.activeTimers,s=[];let r=!1;this.programTimersStack==null?(this.programTimersStack=s,r=!0):this.activeTimers.push(s),this.activeTimers=s,t();const o=sn(this.activeTimers.map(u=>u.query)).filter(u=>u!=null),i=sn(this.activeTimers.map(u=>u.name)).filter(u=>u!=null);this.activeTimers=e,r&&(this.programTimersStack=null);const a={uploadWaitMs:this.uploadWaitMs,downloadWaitMs:this.downloadWaitMs,kernelMs:null,wallMs:null},l=await Promise.all(o);return a.kernelMs=Af(l),a.getExtraProfileInfo=()=>l.map((u,c)=>({name:i[c],ms:u})).map(u=>`${u.name}: ${u.ms}`).join(", "),this.uploadWaitMs=0,this.downloadWaitMs=0,a}makeTensorInfo(t,e,s){return e==="string"&&s!=null&&s.length>0&&Ws(s[0])&&(s=s.map(o=>nn(o))),{dataId:this.write(s,t,e),shape:t,dtype:e}}tensorToBinding(t){if(!t)return null;const s=this.tensorMap.get(t.dataId).resource;return s instanceof GPUBuffer?{buffer:s}:s instanceof GPUTexture?s.createView():s}uploadToGPU(t){const e=this.tensorMap.get(t);if(e.resource!=null)return;const s=sf(e.dtype)*z(e.shape);let r;const o=GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST;if(e.values){if(r=this.bufferManager.acquireBuffer(s,o,!0),r.mapState==="unmapped"){const i=this.bufferManager.acquireBuffer(s,GPUBufferUsage.MAP_WRITE|GPUBufferUsage.COPY_SRC,!0,!1),a=i.getMappedRange();e.dtype==="int32"||e.dtype==="bool"?new Int32Array(a).set(e.values):new Float32Array(a).set(e.values),i.unmap(),this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(i,0,r,0,s),this.stagingPendingDisposal.push(i)}else{const i=r.getMappedRange();e.dtype==="int32"||e.dtype==="bool"?new Int32Array(i).set(e.values):new Float32Array(i).set(e.values),r.unmap()}e.values=null}else r=this.bufferManager.acquireBuffer(s,o);e.resource=r}makeUniforms(t){let e=0,s=0;const r=[];let o=1;t.forEach(l=>{l.data.length===0&&(l.data=[1]);let u;switch(l.data.length){case 1:u=4;break;case 2:u=8;break;case 3:u=16;break;case 4:u=16;break;case 5:u=16;break;case 6:u=16;break;default:w(!1,()=>`Unsupported ${l.data.length}D shape`)}(s===5||s===6)&&(u=16),u>o&&(o=u),e=Math.ceil(e/u)*u,s=l.data.length,r.push(e),e+=l.data.length*4}),e=Math.ceil(e/o)*o;const i=new ArrayBuffer(e);t.forEach((l,u)=>{const c=r[u];l.type==="int32"?new Int32Array(i,c,l.data.length).set(l.data):l.type==="uint32"?new Uint32Array(i,c,l.data.length).set(l.data):new Float32Array(i,c,l.data.length).set(l.data)});const a=this.bufferManager.acquireBuffer(e,GPUBufferUsage.COPY_DST|GPUBufferUsage.UNIFORM);return this.queue.writeBuffer(a,0,i,0,e),this.uniformPendingDisposal.push(a),{offset:0,size:e,buffer:a}}runWebGPUProgram(t,e,s,r,o){if(o||(o=this.makeTensorInfo(t.outputShape,s)),z(o.shape)===0)return this.tensorMap.get(o.dataId).values=Cn(o.dtype,0),o;this.uploadToGPU(o.dataId),t.dispatch=Ox(this.device,t);const i=e.map((l,u)=>{if(l.dtype==="complex64")throw new Error("GPGPUProgram does not support complex64 input. For complex64 dtypes, please separate the program into real and imaginary parts.");return this.uploadToGPU(l.dataId),{dtype:this.tensorMap.get(l.dataId).dtype,shape:l.shape,name:t.variableNames[u]}});t.shaderKey=Ix(t,i,o);const a=W().getBool("WEBGPU_ENGINE_COMPILE_ONLY");return t.shaderKey in this.pipelineCache||(this.pipelineCache[t.shaderKey]=Sx(this.device,t,i,o,a)),t.pipeline=this.pipelineCache[t.shaderKey],a||this.recordAndSubmit(t,o,e,r),o}recordAndSubmit(t,e,s,r){if(t.pipeline instanceof Promise)throw new Error("Please call checkCompileCompletionAsync to ensure parallel compilation is done!");let o=[],i=[];const a="int32";if(t.pixelsOpType==null){o.push({type:"float32",data:[NaN]},{type:"float32",data:[1/0]}),i=s.concat(e).map(d=>d.shape);const f="int32";i.map(d=>{o.push({type:f,data:d});const p=jt(d);o.push({type:f,data:p})})}else{const f=jt(e.shape);o.push({type:a,data:f})}if(t.size){const f=z(t.outputShape);o.push({type:a,data:[t.outputComponent?f/t.outputComponent:f]})}r&&(o=[...o,...r]);const l=[this.tensorToBinding(e),...s.map(f=>this.tensorToBinding(f)),this.makeUniforms(o)];s.forEach(f=>{this.commandQueueOwnedIds.add(f.dataId)}),this.commandQueueOwnedIds.add(e.dataId);const u=this.device.createBindGroup({layout:t.pipeline.getBindGroupLayout(0),entries:l.map((f,d)=>({binding:d,resource:f}))}),c=this.activeTimers!=null;this.ensureCommandEncoderReady();const h={};c&&this.supportTimestampQuery?(this.endComputePassEncoder(),this.querySet==null&&(this.querySet=this.device.createQuerySet({type:"timestamp",count:this.querySetCount})),h.timestampWrites={querySet:this.querySet,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1},this.computePassEncoder=this.commandEncoder.beginComputePass(h)):this.computePassEncoder||(this.computePassEncoder=this.commandEncoder.beginComputePass(h)),this.computePassEncoder.setPipeline(t.pipeline),this.computePassEncoder.setBindGroup(0,u),this.computePassEncoder.dispatchWorkgroups(t.dispatch[0],t.dispatch[1],t.dispatch[2]),this.dispatchCountInPass++,(c||W().get("WEBGPU_DEFERRED_SUBMIT_BATCH_SIZE")<=this.dispatchCountInPass||t.pixelsOpType===to.DRAW)&&(this.endComputePassEncoder(),c?this.activeTimers.push({name:t.constructor.name,query:this.getQueryTime()}):this.submitQueue())}async getQueryTime(){if(!this.supportTimestampQuery)return 0;this.queryResolveBuffer==null&&(this.queryResolveBuffer=this.bufferManager.acquireBuffer(this.querySetCount*8,GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST|GPUBufferUsage.QUERY_RESOLVE)),this.commandEncoder.resolveQuerySet(this.querySet,0,this.querySetCount,this.queryResolveBuffer,0);const t=this.bufferManager.acquireBuffer(this.querySetCount*8,GPUBufferUsage.MAP_READ|GPUBufferUsage.COPY_DST);this.commandEncoder.copyBufferToBuffer(this.queryResolveBuffer,0,t,0,this.querySetCount*8),this.submitQueue(),await t.mapAsync(GPUMapMode.READ);const e=new BigUint64Array(t.getMappedRange()),s=Number(e[1]-e[0])/1e6;return t.unmap(),this.bufferManager.releaseBuffer(t),s}shouldExecuteOnCPU(t,e=Mx){return W().getBool("WEBGPU_CPU_FORWARD")&&t.every(s=>this.tensorMap.get(s.dataId).resource==null&&z(s.shape)<e)}numDataIds(){return this.tensorMap.numDataIds()-this.tensorDataPendingDisposal.length}dispose(){this.disposed||(this.querySet!=null&&this.querySet.destroy(),this.bufferManager.dispose(),this.textureManager.dispose(),this.disposed=!0)}}Bs.nextDataId=0;/**
|
|
3471
|
+
*/const Mx=W().getNumber("WEBGPU_CPU_HANDOFF_SIZE_THRESHOLD"),Ox=(n,t)=>{const e=n.limits.maxComputeWorkgroupsPerDimension,s=t.dispatchLayout,r=t.dispatch;if(r.every(i=>i<=e))return r;w(r[0]>e&&s.y===void 0&&s.z===void 0,()=>"Dispatch size exceeds WebGPU limits in Y or Z dimension.");let o=Math.ceil(Math.sqrt(r[0]));return o>e?(o=Math.ceil(Math.cbrt(r[0])),w(o<=e,()=>"Total dispatch size exceeds WebGPU maximum."),[o,o,o]):[o,o,1]};class Fs extends ba{nextDataId(){return Fs.nextDataId++}constructor(t,e){if(super(),this.commandQueueOwnedIds=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())throw new Error("WebGPU is not supported on this device");this.pipelineCache={},this.device=t,this.queue=t.queue,this.commandEncoder=null,this.computePassEncoder=null,this.adapterInfo=new gx(e),this.supportTimestampQuery=this.device.features.has("timestamp-query"),this.thresholdToIncreaseWorkgroups=this.adapterInfo.intelGPUGeneration>=12?16:8,this.bufferManager=new bx(this.device),this.textureManager=new yx(this.device),this.tensorMap=new vf(this,To()),W().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({device:t,format:"bgra8unorm"}),document.body.appendChild(this.dummyCanvas))}floatPrecision(){return 32}disposeData(t,e=!1){if(!this.tensorMap.has(t))return!0;const s=this.tensorMap.get(t);return e?s.refCount=0:s.refCount--,s.refCount>0?!1:(s.complexTensorInfos!=null&&(this.disposeData(s.complexTensorInfos.real.dataId),this.disposeData(s.complexTensorInfos.imag.dataId)),this.commandQueueOwnedIds.has(t)?(this.tensorDataPendingDisposal.push(t),!0):(this.releaseResource(t),this.tensorMap.delete(t),!0))}memory(){return{numBytesInGPU:this.bufferManager.numBytesUsed,numBytesAllocatedInGPU:this.bufferManager.numBytesAllocated,unreliable:!1}}releaseResource(t){const e=this.tensorMap.get(t);if(!(!e||!e.resource)){if(e.external){e.resource=null;return}e.resource instanceof GPUBuffer?this.bufferManager.releaseBuffer(e.resource):e.resource instanceof GPUTexture&&this.textureManager.releaseTexture(e.resource),e.resource=null}}refCount(t){return this.tensorMap.has(t)?this.tensorMap.get(t).refCount:0}incRef(t){const e=this.tensorMap.get(t);e.refCount++}decRef(t){if(this.tensorMap.has(t)){const e=this.tensorMap.get(t);e.refCount--}}write(t,e,s){if(s==="complex64"&&t!=null)throw new Error("Cannot write to a complex64 dtype. Please use tf.complex(real, imag).");const r={id:this.nextDataId()};return this.tensorMap.set(r,{dtype:s,shape:e,values:t,refCount:1}),r}move(t,e,s,r,o){if(r==="complex64")throw new Error("Cannot write to a complex64 dtype. Please use tf.complex(real, imag).");this.tensorMap.set(t,{dtype:r,shape:s,values:e,refCount:o})}submitQueue(){this.queue.submit([this.commandEncoder.finish()]),this.commandEncoder=null,this.dispatchCountInPass=0,this.commandQueueOwnedIds=new WeakSet,this.tensorDataPendingDisposal.forEach(t=>{this.releaseResource(t),this.tensorMap.delete(t)}),this.uniformPendingDisposal.forEach(t=>this.bufferManager.releaseBuffer(t)),this.stagingPendingDisposal.forEach(t=>this.bufferManager.releaseBuffer(t,!1)),this.tensorDataPendingDisposal=[],this.uniformPendingDisposal=[],this.stagingPendingDisposal=[]}ensureCommandEncoderReady(){this.commandEncoder||(this.commandEncoder=this.device.createCommandEncoder())}endComputePassEncoder(){this.computePassEncoder&&(this.computePassEncoder.end(),this.computePassEncoder=null)}async checkCompileCompletionAsync(){let t;try{t=await Promise.all(Object.values(this.pipelineCache))}catch(e){throw new Error(e.message)}Object.keys(this.pipelineCache).map((e,s)=>{this.pipelineCache[e]=t[s]})}async getBufferData(t){if(W().getBool("WEBGPU_ENGINE_COMPILE_ONLY"))return console.warn("The data may be invalid since WEBGPU_ENGINE_COMPILE_ONLY is true, this can only be called when WEBGPU_ENGINE_COMPILE_ONLY is false"),null;const e=t.size,s=this.bufferManager.acquireBuffer(e,GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ);this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(t,0,s,0,e),this.submitQueue(),await s.mapAsync(GPUMapMode.READ);const r=s.getMappedRange().slice(0);return s.unmap(),s!=null&&this.bufferManager.releaseBuffer(s),W().getBool("WEBGPU_USE_PROFILE_TOOL")&&(w(this.dummyContext!==void 0,()=>"Fail to get context for profiling tool"),this.dummyContext.getCurrentTexture()),r}convertAndCacheOnCPU(t,e){const s=this.tensorMap.get(t);return s.values=e,s.values}readSync(t){const e=this.tensorMap.get(t),{values:s,complexTensorInfos:r}=e;if(s!=null||e.dtype==="string")return s;if(e.dtype==="complex64"){const g=this.readSync(r.real.dataId),m=this.readSync(r.imag.dataId),b=uo(Zl(g,m).buffer,"float32");return this.convertAndCacheOnCPU(t,b),b}this.hasReadSyncWarned||(this.hasReadSyncWarned=!0,console.warn("The performance of synchronously reading data from GPU to CPU is poor on the webgpu backend, please use asynchronous APIs instead."));const o=["opaque","premultiplied"],i=e.resource,a=i.size;w(a%4===0,()=>"Because there is 4 bytes for one pixel, buffer size must be multiple of 4.");const l=a/4,u=new ArrayBuffer(a),c=256,h=256,f=o.map(g=>new OffscreenCanvas(c,h)),d=new OffscreenCanvas(c,h);this.endComputePassEncoder(),f.map((g,m)=>{const b=g.getContext("webgpu");return b.configure({device:this.device,format:"bgra8unorm",usage:GPUTextureUsage.COPY_DST,alphaMode:o[m]}),b.getCurrentTexture()}).map((g,m)=>{const b=c*4,y=(_,T,P)=>{this.ensureCommandEncoderReady(),this.commandEncoder.copyBufferToTexture({buffer:i,bytesPerRow:b,offset:P},{texture:g},{width:_,height:T}),this.submitQueue();const B=d.getContext("2d",{willReadFrequently:!0});B.clearRect(0,0,_,T),B.drawImage(f[m],0,0);const Y=B.getImageData(0,0,_,T).data,j=o[m],F=new Uint8ClampedArray(u,P,_*T*4);for(let G=0;G<F.length;G+=4)if(j==="premultiplied")F[G+3]=Y[G+3];else{const q=Y[G];F[G]=Y[G+2],F[G+1]=Y[G+1],F[G+2]=q}},S=Math.floor(l/(c*h));let x=c,$=h,E=0;for(let _=0;_<S;_++)y(x,$,E),E+=c*h*4;const D=l%(c*h);$=Math.floor(D/c),$>0&&(y(x,$,E),E+=$*(c*4)),x=D%c,x>0&&y(x,1,E)});const p=uo(u,e.dtype);return this.convertAndCacheOnCPU(t,p),p}async read(t){if(!this.tensorMap.has(t))throw new Error(`Tensor ${t} was not registered!`);const e=this.tensorMap.get(t),{values:s}=e;if(s!=null)return s;let r;if(e.dtype==="complex64"){const o=await Promise.all([this.read(e.complexTensorInfos.real.dataId),this.read(e.complexTensorInfos.imag.dataId)]),i=o[0],a=o[1];r=Zl(i,a)}else{const o=await this.getBufferData(e.resource);r=uo(o,e.dtype)}return this.convertAndCacheOnCPU(t,r),r}copyBuffer(t){const e=t.size,s=t.usage,r=this.bufferManager.acquireBuffer(e,s);return this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(t,0,r,0,e),this.submitQueue(),r}createTensorFromGPUData(t,e,s){let r=t.buffer;if(s==="complex64")throw new Error("Cannot write to a complex64 dtype. ");const o={id:this.nextDataId()};this.tensorMap.set(o,{dtype:s,shape:e,values:null,refCount:1,external:t.zeroCopy});const i=this.tensorMap.get(o),a=sf(i.dtype)*z(i.shape);if(t.buffer.size<a)throw new Error(`GPUBuffer size(${t.buffer.size}) is smaller than tensor size(${a})!`);if((t.buffer.usage&(GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC))!==(GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC))throw new Error("GPUBuffer.usage should include GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC!");return t.zeroCopy!==!0&&(r=this.copyBuffer(r)),i.resource=r,To().makeTensorFromDataId(o,e,s,this)}readToGPU(t){const e=this.tensorMap.get(t),{values:s,dtype:r,shape:o,resource:i}=e;if(r==="complex64")throw new Error("Does not support reading buffer for complex64 dtype.");if(i==null)throw s!=null?new Error("Data is not on GPU but on CPU."):new Error("There is no data on GPU or CPU.");const a=i,l=a.size,u=a.usage,c=this.bufferManager.acquireBuffer(l,u);this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(i,0,c,0,l),this.submitQueue();const h=this.makeTensorInfo(o,r),f=To().makeTensorFromTensorInfo(h),d=this.tensorMap.get(h.dataId);return d.resource=c,{tensorRef:f,buffer:c}}bufferSync(t){const e=this.readSync(t.dataId);if(t.dtype==="string")try{const s=e.map(r=>Ys(r));return xt(t.shape,t.dtype,s)}catch{throw new Error("Failed to decode encoded string bytes into utf-8")}return xt(t.shape,t.dtype,e)}async time(t){!this.supportTimestampQuery&&!this.hasTimestampQueryWarned&&(console.warn("This device doesn't support timestamp-query extension. Start Chrome browser with flag --enable-dawn-features=allow_unsafe_apis to try it again. Otherwise, zero will be shown for the kernel time when profiling mode is enabled."),this.hasTimestampQueryWarned=!0);const e=this.activeTimers,s=[];let r=!1;this.programTimersStack==null?(this.programTimersStack=s,r=!0):this.activeTimers.push(s),this.activeTimers=s,t();const o=rn(this.activeTimers.map(u=>u.query)).filter(u=>u!=null),i=rn(this.activeTimers.map(u=>u.name)).filter(u=>u!=null);this.activeTimers=e,r&&(this.programTimersStack=null);const a={uploadWaitMs:this.uploadWaitMs,downloadWaitMs:this.downloadWaitMs,kernelMs:null,wallMs:null},l=await Promise.all(o);return a.kernelMs=Af(l),a.getExtraProfileInfo=()=>l.map((u,c)=>({name:i[c],ms:u})).map(u=>`${u.name}: ${u.ms}`).join(", "),this.uploadWaitMs=0,this.downloadWaitMs=0,a}makeTensorInfo(t,e,s){return e==="string"&&s!=null&&s.length>0&&Gs(s[0])&&(s=s.map(o=>sn(o))),{dataId:this.write(s,t,e),shape:t,dtype:e}}tensorToBinding(t){if(!t)return null;const s=this.tensorMap.get(t.dataId).resource;return s instanceof GPUBuffer?{buffer:s}:s instanceof GPUTexture?s.createView():s}uploadToGPU(t){const e=this.tensorMap.get(t);if(e.resource!=null)return;const s=sf(e.dtype)*z(e.shape);let r;const o=GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST;if(e.values){if(r=this.bufferManager.acquireBuffer(s,o,!0),r.mapState==="unmapped"){const i=this.bufferManager.acquireBuffer(s,GPUBufferUsage.MAP_WRITE|GPUBufferUsage.COPY_SRC,!0,!1),a=i.getMappedRange();e.dtype==="int32"||e.dtype==="bool"?new Int32Array(a).set(e.values):new Float32Array(a).set(e.values),i.unmap(),this.ensureCommandEncoderReady(),this.endComputePassEncoder(),this.commandEncoder.copyBufferToBuffer(i,0,r,0,s),this.stagingPendingDisposal.push(i)}else{const i=r.getMappedRange();e.dtype==="int32"||e.dtype==="bool"?new Int32Array(i).set(e.values):new Float32Array(i).set(e.values),r.unmap()}e.values=null}else r=this.bufferManager.acquireBuffer(s,o);e.resource=r}makeUniforms(t){let e=0,s=0;const r=[];let o=1;t.forEach(l=>{l.data.length===0&&(l.data=[1]);let u;switch(l.data.length){case 1:u=4;break;case 2:u=8;break;case 3:u=16;break;case 4:u=16;break;case 5:u=16;break;case 6:u=16;break;default:w(!1,()=>`Unsupported ${l.data.length}D shape`)}(s===5||s===6)&&(u=16),u>o&&(o=u),e=Math.ceil(e/u)*u,s=l.data.length,r.push(e),e+=l.data.length*4}),e=Math.ceil(e/o)*o;const i=new ArrayBuffer(e);t.forEach((l,u)=>{const c=r[u];l.type==="int32"?new Int32Array(i,c,l.data.length).set(l.data):l.type==="uint32"?new Uint32Array(i,c,l.data.length).set(l.data):new Float32Array(i,c,l.data.length).set(l.data)});const a=this.bufferManager.acquireBuffer(e,GPUBufferUsage.COPY_DST|GPUBufferUsage.UNIFORM);return this.queue.writeBuffer(a,0,i,0,e),this.uniformPendingDisposal.push(a),{offset:0,size:e,buffer:a}}runWebGPUProgram(t,e,s,r,o){if(o||(o=this.makeTensorInfo(t.outputShape,s)),z(o.shape)===0)return this.tensorMap.get(o.dataId).values=Cn(o.dtype,0),o;this.uploadToGPU(o.dataId),t.dispatch=Ox(this.device,t);const i=e.map((l,u)=>{if(l.dtype==="complex64")throw new Error("GPGPUProgram does not support complex64 input. For complex64 dtypes, please separate the program into real and imaginary parts.");return this.uploadToGPU(l.dataId),{dtype:this.tensorMap.get(l.dataId).dtype,shape:l.shape,name:t.variableNames[u]}});t.shaderKey=Ix(t,i,o);const a=W().getBool("WEBGPU_ENGINE_COMPILE_ONLY");return t.shaderKey in this.pipelineCache||(this.pipelineCache[t.shaderKey]=Sx(this.device,t,i,o,a)),t.pipeline=this.pipelineCache[t.shaderKey],a||this.recordAndSubmit(t,o,e,r),o}recordAndSubmit(t,e,s,r){if(t.pipeline instanceof Promise)throw new Error("Please call checkCompileCompletionAsync to ensure parallel compilation is done!");let o=[],i=[];const a="int32";if(t.pixelsOpType==null){o.push({type:"float32",data:[NaN]},{type:"float32",data:[1/0]}),i=s.concat(e).map(d=>d.shape);const f="int32";i.map(d=>{o.push({type:f,data:d});const p=Kt(d);o.push({type:f,data:p})})}else{const f=Kt(e.shape);o.push({type:a,data:f})}if(t.size){const f=z(t.outputShape);o.push({type:a,data:[t.outputComponent?f/t.outputComponent:f]})}r&&(o=[...o,...r]);const l=[this.tensorToBinding(e),...s.map(f=>this.tensorToBinding(f)),this.makeUniforms(o)];s.forEach(f=>{this.commandQueueOwnedIds.add(f.dataId)}),this.commandQueueOwnedIds.add(e.dataId);const u=this.device.createBindGroup({layout:t.pipeline.getBindGroupLayout(0),entries:l.map((f,d)=>({binding:d,resource:f}))}),c=this.activeTimers!=null;this.ensureCommandEncoderReady();const h={};c&&this.supportTimestampQuery?(this.endComputePassEncoder(),this.querySet==null&&(this.querySet=this.device.createQuerySet({type:"timestamp",count:this.querySetCount})),h.timestampWrites={querySet:this.querySet,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1},this.computePassEncoder=this.commandEncoder.beginComputePass(h)):this.computePassEncoder||(this.computePassEncoder=this.commandEncoder.beginComputePass(h)),this.computePassEncoder.setPipeline(t.pipeline),this.computePassEncoder.setBindGroup(0,u),this.computePassEncoder.dispatchWorkgroups(t.dispatch[0],t.dispatch[1],t.dispatch[2]),this.dispatchCountInPass++,(c||W().get("WEBGPU_DEFERRED_SUBMIT_BATCH_SIZE")<=this.dispatchCountInPass||t.pixelsOpType===eo.DRAW)&&(this.endComputePassEncoder(),c?this.activeTimers.push({name:t.constructor.name,query:this.getQueryTime()}):this.submitQueue())}async getQueryTime(){if(!this.supportTimestampQuery)return 0;this.queryResolveBuffer==null&&(this.queryResolveBuffer=this.bufferManager.acquireBuffer(this.querySetCount*8,GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST|GPUBufferUsage.QUERY_RESOLVE)),this.commandEncoder.resolveQuerySet(this.querySet,0,this.querySetCount,this.queryResolveBuffer,0);const t=this.bufferManager.acquireBuffer(this.querySetCount*8,GPUBufferUsage.MAP_READ|GPUBufferUsage.COPY_DST);this.commandEncoder.copyBufferToBuffer(this.queryResolveBuffer,0,t,0,this.querySetCount*8),this.submitQueue(),await t.mapAsync(GPUMapMode.READ);const e=new BigUint64Array(t.getMappedRange()),s=Number(e[1]-e[0])/1e6;return t.unmap(),this.bufferManager.releaseBuffer(t),s}shouldExecuteOnCPU(t,e=Mx){return W().getBool("WEBGPU_CPU_FORWARD")&&t.every(s=>this.tensorMap.get(s.dataId).resource==null&&z(s.shape)<e)}numDataIds(){return this.tensorMap.numDataIds()-this.tensorDataPendingDisposal.length}dispose(){this.disposed||(this.querySet!=null&&this.querySet.destroy(),this.bufferManager.dispose(),this.textureManager.dispose(),this.disposed=!0)}}Fs.nextDataId=0;/**
|
|
3472
3472
|
* @license
|
|
3473
3473
|
* Copyright 2022 Google Inc. All Rights Reserved.
|
|
3474
3474
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3483,7 +3483,7 @@ out_color[outIdx] = vec4f(
|
|
|
3483
3483
|
* See the License for the specific language governing permissions and
|
|
3484
3484
|
* limitations under the License.
|
|
3485
3485
|
* =============================================================================
|
|
3486
|
-
*/rf()&&Xp("webgpu",async()=>{const n={powerPreference:W().get("WEBGPU_USE_LOW_POWER_GPU")?"low-power":"high-performance"},t=await navigator.gpu.requestAdapter(n),e={},s=[];t.features.has("timestamp-query")&&s.push("timestamp-query"),t.features.has("bgra8unorm-storage")&&s.push(["bgra8unorm-storage"]),e.requiredFeatures=s;const r=t.limits;e.requiredLimits={maxComputeWorkgroupStorageSize:r.maxComputeWorkgroupStorageSize,maxComputeWorkgroupsPerDimension:r.maxComputeWorkgroupsPerDimension,maxStorageBufferBindingSize:r.maxStorageBufferBindingSize,maxBufferSize:r.maxBufferSize,maxComputeWorkgroupSizeX:r.maxComputeWorkgroupSizeX,maxComputeInvocationsPerWorkgroup:r.maxComputeInvocationsPerWorkgroup};const o=await t.requestDevice(e),i=await t.requestAdapterInfo();return new
|
|
3486
|
+
*/rf()&&Xp("webgpu",async()=>{const n={powerPreference:W().get("WEBGPU_USE_LOW_POWER_GPU")?"low-power":"high-performance"},t=await navigator.gpu.requestAdapter(n),e={},s=[];t.features.has("timestamp-query")&&s.push("timestamp-query"),t.features.has("bgra8unorm-storage")&&s.push(["bgra8unorm-storage"]),e.requiredFeatures=s;const r=t.limits;e.requiredLimits={maxComputeWorkgroupStorageSize:r.maxComputeWorkgroupStorageSize,maxComputeWorkgroupsPerDimension:r.maxComputeWorkgroupsPerDimension,maxStorageBufferBindingSize:r.maxStorageBufferBindingSize,maxBufferSize:r.maxBufferSize,maxComputeWorkgroupSizeX:r.maxComputeWorkgroupSizeX,maxComputeInvocationsPerWorkgroup:r.maxComputeInvocationsPerWorkgroup};const o=await t.requestDevice(e),i=await t.requestAdapterInfo();return new Fs(o,i)},3);/**
|
|
3487
3487
|
* @license
|
|
3488
3488
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
3489
3489
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3612,7 +3612,7 @@ out_color[outIdx] = vec4f(
|
|
|
3612
3612
|
* See the License for the specific language governing permissions and
|
|
3613
3613
|
* limitations under the License.
|
|
3614
3614
|
* =============================================================================
|
|
3615
|
-
*/function of(n){const{backend:t,attrs:e}=n,{shape:s,value:r}=e;let{dtype:o}=e;if(o=o||
|
|
3615
|
+
*/function of(n){const{backend:t,attrs:e}=n,{shape:s,value:r}=e;let{dtype:o}=e;if(o=o||as(r),o==="string"){const i=bt(o,z(s));return i.fill(r),t.makeTensorInfo(s,o,i)}else{const i=new Gx(s),a=[{type:"float32",data:[r]}];return t.runWebGPUProgram(i,[],o,a)}}/**
|
|
3616
3616
|
* @license
|
|
3617
3617
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
3618
3618
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3657,7 +3657,7 @@ out_color[outIdx] = vec4f(
|
|
|
3657
3657
|
* See the License for the specific language governing permissions and
|
|
3658
3658
|
* limitations under the License.
|
|
3659
3659
|
* =============================================================================
|
|
3660
|
-
*/function Mt(n){return(t,e,s,r,o)=>{const i=zt(t,e),a=i.length,l=
|
|
3660
|
+
*/function Mt(n){return(t,e,s,r,o)=>{const i=zt(t,e),a=i.length,l=Kt(i),u=z(i),c=Cn(o,u),h=t.length,f=e.length,d=Kt(t),p=Kt(e),g=ir(t,i),m=ir(e,i);if(g.length+m.length===0)for(let b=0;b<c.length;++b)c[b]=n(s[b%s.length],r[b%r.length]);else for(let b=0;b<c.length;++b){const y=ho(b,a,l),S=y.slice(-h);g.forEach(D=>S[D]=0);const x=co(S,h,d),$=y.slice(-f);m.forEach(D=>$[D]=0);const E=co($,f,p);c[b]=n(s[x],r[E])}return[c,i]}}/**
|
|
3661
3661
|
* @license
|
|
3662
3662
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
3663
3663
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -3672,7 +3672,7 @@ out_color[outIdx] = vec4f(
|
|
|
3672
3672
|
* See the License for the specific language governing permissions and
|
|
3673
3673
|
* limitations under the License.
|
|
3674
3674
|
* =============================================================================
|
|
3675
|
-
*/function jx(n,t,e,s){if(s==="int32"){const r=Int32Array.from(n);return[t,"int32",r]}if(s==="bool"){const r=
|
|
3675
|
+
*/function jx(n,t,e,s){if(s==="int32"){const r=Int32Array.from(n);return[t,"int32",r]}if(s==="bool"){const r=Ks([0],e),[o,i]=Mt((a,l)=>a!==l?1:0)(t,[],n,r,"bool");return[i,"bool",o]}throw new Error(`Error in Cast: failed to cast ${e} to ${s}`)}/**
|
|
3676
3676
|
* @license
|
|
3677
3677
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
3678
3678
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4062,7 +4062,7 @@ out_color[outIdx] = vec4f(
|
|
|
4062
4062
|
* See the License for the specific language governing permissions and
|
|
4063
4063
|
* limitations under the License.
|
|
4064
4064
|
* =============================================================================
|
|
4065
|
-
*/function bS(n,t,e,s,r){const o=t.length,i=z(t),a=
|
|
4065
|
+
*/function bS(n,t,e,s,r){const o=t.length,i=z(t),a=Kt(t),l=Kt(r),u=Cn(e,z(r));for(let c=0;c<i;++c){const h=ho(c,o,a),f=new Array(h.length);for(let p=0;p<f.length;p++)f[p]=h[s[p]];const d=co(f,o,l);u[d]=n[c]}return u}/**
|
|
4066
4066
|
* @license
|
|
4067
4067
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
4068
4068
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4092,7 +4092,7 @@ out_color[outIdx] = vec4f(
|
|
|
4092
4092
|
* See the License for the specific language governing permissions and
|
|
4093
4093
|
* limitations under the License.
|
|
4094
4094
|
* =============================================================================
|
|
4095
|
-
*/function wS(n,t,e){n.forEach((s,r)=>{if(s<0||s>=e){const o=ho(r,t.length,
|
|
4095
|
+
*/function wS(n,t,e){n.forEach((s,r)=>{if(s<0||s>=e){const o=ho(r,t.length,Kt(t)).join(",");throw new Error(`indices[${o}] = ${s} is not in [0, ${e})`)}})}function xS(n,t){for(let e=0;e<n.length;++e){const s=n[e],r=e===n.length-1?t:n[e+1].length;if(s.length===0)throw new Error("Ragged splits may not be empty");if(s[0]<0)throw new Error("Ragged splits must be non-negative");if(s[s.length-1]>r)throw new Error("Ragged splits must not point past values");for(let o=1;o<s.length;++o)if(s[o-1]>s[o])throw new Error("Ragged splits must be sorted in ascending order")}}function SS(n,t,e,s){const r=[];let o=0;const i=t.length-1+e.length,a=new Array(i).fill(null).map(()=>[0]);xS(e,s);let l=1;for(let u=0;u<t.length-1;++u){l*=t[u];const c=t[u+1];for(let h=1;h<l+1;++h)a[u].push(h*c)}for(let u=0;u<n.length;++u){let c=n[u],h=n[u]+1;for(let f=0;f<e.length;++f){const d=e[f],p=f+t.length-1;if(p>=0){const g=a[p],m=g[g.length-1]-d[c];for(let b=c;b<h;++b)a[p].push(d[b+1]+m)}c=d[c],h=d[h]}h!==c&&(r.push([c,h]),o+=h-c)}return{outSplits:a,valueSlices:r,numValues:o}}function $S(n){const t=[];for(let e=0;e<n.length;++e){const s=n[e].length,r=bt("int32",s);t.push(r),n[e].forEach((o,i)=>r[i]=o)}return t}function lf(n,t){const e=n.slice(0,t);for(;e.length<t;)e.push(1);for(let s=t;s<n.length;s++)e[t-1]*=n[s];return e}function vS(n,t,e,s,r,o){const i=lf(t,2)[1],a=lf(o,2)[1];let l=0;for(const u of e)for(let c=u[0];c<u[1];++c){for(let h=0;h<s;++h)r[l*a+h]=n[c*i+h];++l}}function IS(n,t,e,s,r){const o=t.slice();o[0]=r;const i=bt(e,z(o)),a=n.length,l=a===0?0:a/t[0];return vS(n,t,s,l,i,o),[i,o]}function AS(n,t,e,s,r,o,i,a){if(n.length===0)throw new Error("paramsNestedSplits must be non empty");if(t[0].length===0)throw new Error("Split tensors must not be scalars");const l=t[0][0]-1;if(wS(o,i,l),s.length===0)throw new Error("params.rank must be nonzero");const u=s[0],{outSplits:c,valueSlices:h,numValues:f}=SS(o,i,n,u),d=$S(c),p=IS(e,s,r,h,f);return[d,p[0],p[1]]}/**
|
|
4096
4096
|
* @license
|
|
4097
4097
|
* Copyright 2022 Google LLC.
|
|
4098
4098
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4122,7 +4122,7 @@ out_color[outIdx] = vec4f(
|
|
|
4122
4122
|
* See the License for the specific language governing permissions and
|
|
4123
4123
|
* limitations under the License.
|
|
4124
4124
|
* =============================================================================
|
|
4125
|
-
*/var le=Ie;class
|
|
4125
|
+
*/var le=Ie;class no{constructor(t,e,s,r,o,i,a,l,u,c){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=hy(c),this.raggedRank=fy(this.rowPartitionTypes)}getRowPartitionTypeByDimension(t){return this.rowPartitionTypes[0]===le.FIRST_DIM_SIZE?this.rowPartitionTypes[t+1]:this.rowPartitionTypes[t]}getRowPartitionTensor(t){return this.rowPartitionTypes[0]===le.FIRST_DIM_SIZE?this.rowPartitionValues[t+1]:this.rowPartitionValues[t]}getMaxWidth(t){const e=this.getRowPartitionTensor(t-1);switch(this.getRowPartitionTypeByDimension(t-1)){case le.VALUE_ROWIDS:return no.getMaxWidthValueRowID(e);case le.ROW_SPLITS:return no.getMaxWidthRowSplit(e);default:throw new Error(`Cannot handle partition type ${le[this.getRowPartitionTypeByDimension(t-1)]}`)}}static getMaxWidthRowSplit(t){const e=t.length;if(e===0||e===1)return 0;let s=0;for(let r=0;r<e-1;++r){const o=t[r+1]-t[r];o>s&&(s=o)}return s}static getMaxWidthValueRowID(t){const e=t.length;if(e===0)return 0;let s=0,r=t[0],o=0;for(let i=1;i<e;++i){const a=t[i];a!==r&&(r=a,o=Math.max(i-s,o),s=i)}return Math.max(e-s,o)}tensorShapeFromTensor(t,e,s=!0){if(e.length===0){if(t[0]===-1)return[];throw new Error("The only valid scalar shape tensor is the fully unknown shape specified as -1.")}return hf(t,s)}calculateOutputSize(t){const e=this.valuesShape,s=this.defaultValueShape;dy(s,e);const r=this.tensorShapeFromTensor(this.shape,this.shapeShape),i=cy(this.raggedRank,r,e);i[0]<0&&(i[0]=t);for(let a=1;a<=this.raggedRank;++a)i[a]<0&&(i[a]=this.getMaxWidth(a));return i}calculateFirstParentOutputIndex(t,e,s){const r=Math.min(t,s),o=[];let i=0;for(let a=0;a<r;++a,i+=e)o.push(i);for(let a=r;a<t;++a)o.push(-1);return w(o.length===t,()=>"Final length of result must be equal to firstDimension."),o}calculateOutputIndexRowSplit(t,e,s,r){const o=t.length,i=[];for(let a=0;a<o-1;++a){const l=t[a+1]-t[a];let u=Math.min(r,l),c=e[a];c===-1&&(u=0);for(let h=0;h<u;++h)i.push(c),c+=s;for(let h=0;h<l-u;++h)i.push(-1)}if(o>0&&i.length!==t[o-1])throw new Error("Invalid row split size.");return i}calculateOutputIndexValueRowID(t,e,s,r){const o=t.length,i=[];if(o===0)return[];let a=0,l=t[0];if(l>=e.length)throw new Error(`Got currentValueRowId=${l}, which is not less than ${e.length}`);let u=e[l];i.push(u);for(let c=1;c<o;++c){const h=t[c];if(h===l)u>=0&&(++a,a<r?u+=s:u=-1);else{if(a=0,l=h,h>=e.length)throw new Error(`Got nextValueRowId=${h} which is not less than ${e.length}`);u=e[h]}i.push(u)}if(i.length!==t.length)throw new Error("Invalid row ids.");return i}calculateOutputIndex(t,e,s,r){const o=this.getRowPartitionTensor(t),i=this.getRowPartitionTypeByDimension(t);switch(i){case le.VALUE_ROWIDS:return this.calculateOutputIndexValueRowID(o,e,s,r);case le.ROW_SPLITS:if(o.length-1>e.length)throw new Error(`Row partition size is greater than output size: ${o.length-1} > ${e.length}`);return this.calculateOutputIndexRowSplit(o,e,s,r);default:throw new Error(`Unsupported partition type: ${le[i]}`)}}getFirstDimensionSize(){const t=this.rowPartitionValues[0];if(this.rowPartitionTypes.length===0)throw new Error("No row_partition_types given.");const e=this.rowPartitionTypes[0];switch(e){case le.FIRST_DIM_SIZE:return t[0];case le.VALUE_ROWIDS:throw new Error("Cannot handle VALUE_ROWIDS in first dimension.");case le.ROW_SPLITS:return this.rowPartitionValuesShapes[0][0]-1;default:throw new Error(`Cannot handle type ${le[e]}`)}}compute(){if(this.rowPartitionValues[0].length<=0)throw new Error("Invalid first partition input. Tensor requires at least one element.");const e=this.getFirstDimensionSize(),s=this.calculateOutputSize(e),r=new Array(this.raggedRank+1);r[r.length-1]=1;for(let l=r.length-2;l>=0;--l)r[l]=r[l+1]*s[l+1];const o=hf(s,!1),i=bt(this.valuesDType,z(o));if(r[0]*s[0]>0){let l=this.calculateFirstParentOutputIndex(e,r[0],s[0]);for(let u=1;u<=this.raggedRank;++u)l=this.calculateOutputIndex(u-1,l,r[u],s[u]);this.setOutput(this.raggedRank,l,i,o)}return[o,i]}setOutput(t,e,s,r){if(s.length===0)return;const o=this.values,i=s;let a=r.slice();a=a.slice(t+1);const l=z(a),u=e.length;let c=this.defaultValue;if(c.length!==l&&c.length!==1){const p=this.defaultValueShape;k(()=>{const g=L(c,p);c=rr(g,a).dataSync()})}let h=0,f=0,d=0;for(let p=0;p<=u;++p){let g=p<u?e[p]:-1;if(g===d){++d;continue}if(f<d){const m=o.subarray(h*l),b=i.subarray(f*l),y=(d-f)*l;cf(b,m,y)}if(p>=u){const m=s.length;g=Math.floor(m/l)}if(g>d)if(this.defaultValue.length===1)i.subarray(d*l,g*l).fill(this.defaultValue[0]),d=g;else for(;g>d;){const m=i.slice(d*l);cf(m,c,l),++d}g<0?(h=p+1,f=d):(h=p,f=d,d=f+1)}}}function cf(n,t,e){for(let s=0;s<e;s++)n[s]=t[s]}function hf(n,t){const e=[];for(let s of n){if(s<0){if(!t)throw new Error(`Dimension ${s} must be >= 0`);if(s<-1)throw new Error(`Dimension ${s} must be >= -1`);s=-1}e.push(s)}return e}function CS(n,t,e,s,r,o,i,a,l,u){return new no(n,t,e,s,r,o,i,a,l,u).compute()}/**
|
|
4126
4126
|
* @license
|
|
4127
4127
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
4128
4128
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4167,7 +4167,7 @@ out_color[outIdx] = vec4f(
|
|
|
4167
4167
|
* See the License for the specific language governing permissions and
|
|
4168
4168
|
* limitations under the License.
|
|
4169
4169
|
* =============================================================================
|
|
4170
|
-
*/function TS(n,t,e,s,r,o,i,a,l,u){const c=[s/r,r],h=n.values,f=t.values;if(s===0)return xt(e,t.dtype);const d=l instanceof
|
|
4170
|
+
*/function TS(n,t,e,s,r,o,i,a,l,u){const c=[s/r,r],h=n.values,f=t.values;if(s===0)return xt(e,t.dtype);const d=l instanceof Js?l:xt(c,t.dtype);typeof l=="string"||typeof l=="number"?d.values.fill(l):typeof l=="boolean"&&d.values.fill(+l);for(let p=0;p<o;p++){const g=[];let m=0;for(let b=0;b<i;b++){const y=h[p*i+b];g.push(y),m+=y*a[b]}if(m<0||m>=s/r)throw new Error(`Invalid indices: ${g} does not index into ${e}`);for(let b=0;b<r;b++)u?d.values[m*r+b]+=f[p*r+b]:d.values[m*r+b]=t.rank===0?f[0]:f[p*r+b]}return d}/**
|
|
4171
4171
|
* @license
|
|
4172
4172
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
4173
4173
|
* Licensed under the Apache License, Version 2.0 (the License);
|
|
@@ -4197,7 +4197,7 @@ out_color[outIdx] = vec4f(
|
|
|
4197
4197
|
* See the License for the specific language governing permissions and
|
|
4198
4198
|
* limitations under the License.
|
|
4199
4199
|
* =============================================================================
|
|
4200
|
-
*/function DS(n,t,e,s,r){const o=sy(s,t,e),i=z(e),a=
|
|
4200
|
+
*/function DS(n,t,e,s,r){const o=sy(s,t,e),i=z(e),a=Kt(s);if(o){const h=ry(t,a);return r==="string"?n.slice(h,h+i):n.subarray(h,h+i)}const l=r==="string"?tu(n):n,u=xt(s,r,l),c=xt(e,r);for(let h=0;h<c.size;++h){const f=c.indexToLoc(h),d=f.map((p,g)=>p+t[g]);c.set(u.get(...d),...f)}return r==="string"?Ry(c.values):c.values}/**
|
|
4201
4201
|
* @license
|
|
4202
4202
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
4203
4203
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4317,7 +4317,7 @@ out_color[outIdx] = vec4f(
|
|
|
4317
4317
|
* See the License for the specific language governing permissions and
|
|
4318
4318
|
* limitations under the License.
|
|
4319
4319
|
* =============================================================================
|
|
4320
|
-
*/class zS{constructor(t,e,s,r,o,i){this.separator=
|
|
4320
|
+
*/class zS{constructor(t,e,s,r,o,i){this.separator=sn(t),this.nGramWidths=e,this.leftPad=sn(s),this.rightPad=sn(r),this.padWidth=o,this.preserveShort=i}getPadWidth(t){return Math.min(this.padWidth<0?t-1:this.padWidth,t-1)}getNumNGrams(t,e){const s=this.getPadWidth(e);return Math.max(0,t+2*s-e+1)}createNGrams(t,e,s,r,o,i){for(let a=0;a<o;++a){const l=this.getPadWidth(i),u=Math.max(0,l-a),c=Math.max(0,l-(o-(a+1))),h=i-(u+c),f=e+(u>0?0:a-l);let d=0;d+=u*this.leftPad.length;for(let y=0;y<h;++y)d+=t[f+y].length;d+=c*this.rightPad.length;const p=u+c+h-1;d+=p*this.separator.length,s[r+a]=new Uint8Array(d);const g=s[r+a];let m=0;const b=y=>y.forEach(S=>g[m++]=S);for(let y=0;y<u;++y)b(this.leftPad),b(this.separator);for(let y=0;y<h-1;++y)b(t[f+y]),b(this.separator);if(h>0){b(t[f+h-1]);for(let y=0;y<c;++y)b(this.separator),b(this.rightPad)}else{for(let y=0;y<c-1;++y)b(this.rightPad),b(this.separator);b(this.rightPad)}}}compute(t,e){const s=t.length,r=e.length;if(r>0){let l=e[0];if(l!==0)throw new Error(`First split value must be 0, got ${l}`);for(let u=1;u<r;++u){let c=e[u]>=l;if(c=c&&e[u]<=s,!c)throw new Error(`Invalid split value ${e[u]}, must be in [${l}, ${s}]`);l=e[u]}if(l!==s)throw new Error(`Last split value must be data size. Expected ${s}, got ${l}`)}const o=r-1,i=bt("int32",r);if(s===0||r===0){const l=new Array(s);for(let u=0;u<=o;++u)i[u]=0;return[l,i]}i[0]=0;for(let l=1;l<=o;++l){const u=e[l]-e[l-1];let c=0;this.nGramWidths.forEach(h=>{c+=this.getNumNGrams(u,h)}),this.preserveShort&&u>0&&c===0&&(c=1),i[l]=i[l-1]+c}const a=new Array(i[o]);for(let l=0;l<o;++l){const u=e[l];let c=i[l];if(this.nGramWidths.forEach(h=>{const f=e[l+1]-e[l],d=this.getNumNGrams(f,h);this.createNGrams(t,u,a,c,d,h),c+=d}),this.preserveShort&&c===i[l]){const h=e[l+1]-e[l];if(h===0)continue;const f=h+2*this.padWidth;this.createNGrams(t,u,a,c,1,f)}}return[a,i]}}function US(n,t,e,s,r,o,i,a){return new zS(e,s,r,o,i,a).compute(n,t)}/**
|
|
4321
4321
|
* @license
|
|
4322
4322
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
4323
4323
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4392,7 +4392,7 @@ out_color[outIdx] = vec4f(
|
|
|
4392
4392
|
* See the License for the specific language governing permissions and
|
|
4393
4393
|
* limitations under the License.
|
|
4394
4394
|
* =============================================================================
|
|
4395
|
-
*/const
|
|
4395
|
+
*/const zs=(n,t)=>{const e=t.value-n.value;return e===0?n.index-t.index:e};function ff(n,t,e=0,s=n.length-1){for(;s>e;){if(s-e>600){const a=s-e+1,l=t-e+1,u=Math.log(a),c=.5*Math.exp(2*u/3),h=.5*Math.sqrt(u*c*(a-c)/a)*Math.sign(l-a/2),f=Math.max(e,Math.floor(t-l*c/a+h)),d=Math.min(s,Math.floor(t+(a-l)*c/a+h));ff(n,t,f,d)}const r=n[t];let o=e,i=s;for(En(n,e,t),zs(n[s],r)>0&&En(n,e,s);o<i;){for(En(n,o,i),o++,i--;zs(n[o],r)<0;)o=o+1;for(;zs(n[i],r)>0;)i=i-1}zs(n[e],r)===0?En(n,e,i):(i=i+1,En(n,i,s)),i<=t&&(e=i+1),t<=i&&(s=i-1)}}function HS(n,t,e,s,r){const o=t[t.length-1],[i,a]=[n.length/o,o],l=Cn(e,i*s),u=Cn("int32",i*s);for(let h=0;h<i;h++){const f=h*a,d=n.subarray(f,f+a);let p=new Array(d.length);d.forEach((y,S)=>p[S]={value:y,index:S}),s<p.length&&(ff(p,s),p=p.slice(0,s)),r&&p.sort(zs);const g=h*s,m=l.subarray(g,g+s),b=u.subarray(g,g+s);for(let y=0;y<s;y++)m[y]=p[y].value,b[y]=p[y].index}const c=t.slice();return c[c.length-1]=s,[xt(c,e,l),xt(c,"int32",u)]}/**
|
|
4396
4396
|
* @license
|
|
4397
4397
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
4398
4398
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4407,7 +4407,7 @@ out_color[outIdx] = vec4f(
|
|
|
4407
4407
|
* See the License for the specific language governing permissions and
|
|
4408
4408
|
* limitations under the License.
|
|
4409
4409
|
* =============================================================================
|
|
4410
|
-
*/function KS(n,t,e,s){const r=
|
|
4410
|
+
*/function KS(n,t,e,s){const r=is(t,e)[0],o=[1,e[0],1];for(let p=0;p<r;p++)o[0]*=e[p];o[1]=e[r];for(let p=r+1;p<e.length;p++)o[2]*=e[p];const i=new Map,a=new Int32Array(e[r]),l=new Js(o,s,n),u=[],c=o[0]===1&&o[2]===1;for(let p=0;p<e[r];p++){let g;if(c)g=n[p].toString();else{const b=[];for(let y=0;y<o[0];y++)for(let S=0;S<o[2];S++)b.push(l.get(y,p,S));g=b.join(",")}const m=i.get(g);if(m!=null)a[p]=m;else{const b=i.size;i.set(g,b),a[p]=b,u.push(p)}}const h=o.slice();h[1]=i.size;const f=new Js(h,s);u.forEach((p,g)=>{for(let m=0;m<o[0];m++)for(let b=0;b<o[2];b++)f.set(l.get(m,p,b),m,g,b)});const d=e.slice();return d[r]=h[1],{outputValues:f.values,outputShape:d,indices:a}}/**
|
|
4411
4411
|
* @license
|
|
4412
4412
|
* Copyright 2020 Google LLC. All Rights Reserved.
|
|
4413
4413
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -4452,7 +4452,7 @@ out_color[outIdx] = vec4f(
|
|
|
4452
4452
|
* See the License for the specific language governing permissions and
|
|
4453
4453
|
* limitations under the License.
|
|
4454
4454
|
* =============================================================================
|
|
4455
|
-
*/class e2{constructor(t,e){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=It(this.dispatchLayout,this.outputShape,this.workgroupSize,[this.workPerThread,1,1]),this.start=t,this.uniforms=`start : ${Tt(t.length)}, `,this.shaderKey="slice"}getUserCode(){const t=Tt(this.rank),e=n2(this.rank);let s;return this.start.length===1?s=this.outputShape.map((o,i)=>"sourceLoc = uniforms.start + coords;"):s=this.outputShape.map((o,i)=>`sourceLoc.${fa[i]} = uniforms.start.${
|
|
4455
|
+
*/class e2{constructor(t,e){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=It(this.dispatchLayout,this.outputShape,this.workgroupSize,[this.workPerThread,1,1]),this.start=t,this.uniforms=`start : ${Tt(t.length)}, `,this.shaderKey="slice"}getUserCode(){const t=Tt(this.rank),e=n2(this.rank);let s;return this.start.length===1?s=this.outputShape.map((o,i)=>"sourceLoc = uniforms.start + coords;"):s=this.outputShape.map((o,i)=>`sourceLoc.${fa[i]} = uniforms.start.${In(i)} + coords.${fa[i]};`),`
|
|
4456
4456
|
${wt("index")} {
|
|
4457
4457
|
if (index < uniforms.size) {
|
|
4458
4458
|
var sourceLoc : ${t};
|
|
@@ -4738,7 +4738,7 @@ out_color[outIdx] = vec4f(
|
|
|
4738
4738
|
}`:a=`
|
|
4739
4739
|
fn activation(a : ${i}, coords : vec${s}<i32>) -> ${i} {
|
|
4740
4740
|
${r}
|
|
4741
|
-
}`,a}function
|
|
4741
|
+
}`,a}function so(n,t){return`
|
|
4742
4742
|
${n?"value = value + getBiasByOutputCoords(coords);":""}
|
|
4743
4743
|
${t?"value = activation(value, coords);":""}
|
|
4744
4744
|
`}/**
|
|
@@ -4783,7 +4783,7 @@ out_color[outIdx] = vec4f(
|
|
|
4783
4783
|
{
|
|
4784
4784
|
var value = valueIn;
|
|
4785
4785
|
let coords = vec3<i32>(batch, row, col);
|
|
4786
|
-
${
|
|
4786
|
+
${so(n,t)}
|
|
4787
4787
|
setOutputAtCoords(coords[0], coords[1], coords[2], value);
|
|
4788
4788
|
}
|
|
4789
4789
|
}
|
|
@@ -5125,7 +5125,7 @@ out_color[outIdx] = vec4f(
|
|
|
5125
5125
|
var value = valueIn;
|
|
5126
5126
|
let outWidth = ${n?"uniforms.outShape[2]":"uniforms.outShape[3]"};
|
|
5127
5127
|
${d}
|
|
5128
|
-
${
|
|
5128
|
+
${so(r,o)}
|
|
5129
5129
|
setOutputAtCoords(coords[0], coords[1], coords[2], coords[3], value);
|
|
5130
5130
|
}
|
|
5131
5131
|
}`}class R${constructor(t,e,s,r,o=!1,i=null,a=!1,l=!1){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=Px(this.dispatchLayout,this.outputShape,this.isVec4),this.elementsPerThread=Lx(this.dispatchLayout,this.outputShape,this.isVec4),this.dispatch=It(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}`}getUserCode(){const t=this.isVec4?pa(this.elementsPerThread,this.workgroupSize,!this.isChannelsLast,this.tileInner):ma(this.elementsPerThread,this.workgroupSize,!this.isChannelsLast,this.tileInner,!1,null,this.sequentialAccessByThreads),e=this.isVec4?[this.innerElementSize,4,4]:[1,1,1];return`
|
|
@@ -5168,7 +5168,7 @@ out_color[outIdx] = vec4f(
|
|
|
5168
5168
|
let coords = ${this.isChannelsLast?"vec4<i32>(batch, row, col, chan);":"vec4<i32>(batch, chan, row, col);"}
|
|
5169
5169
|
if (coordsInBounds4D(coords, uniforms.outShape)) {
|
|
5170
5170
|
var value = valueIn;
|
|
5171
|
-
${
|
|
5171
|
+
${so(this.addBias,this.activation)}
|
|
5172
5172
|
setOutputAtCoords(coords.x, coords.y, coords.z, coords.w, value);
|
|
5173
5173
|
}
|
|
5174
5174
|
}
|
|
@@ -5389,7 +5389,7 @@ out_color[outIdx] = vec4f(
|
|
|
5389
5389
|
if (index < uniforms.size) {
|
|
5390
5390
|
let coords = getCoordsFromIndex(index);
|
|
5391
5391
|
var value = getXByOutputIndex(index);
|
|
5392
|
-
${
|
|
5392
|
+
${so(this.addBias,this.activation)}
|
|
5393
5393
|
setOutputAtIndex(index, value);
|
|
5394
5394
|
}
|
|
5395
5395
|
}
|
|
@@ -5423,7 +5423,7 @@ out_color[outIdx] = vec4f(
|
|
|
5423
5423
|
* See the License for the specific language governing permissions and
|
|
5424
5424
|
* limitations under the License.
|
|
5425
5425
|
* =============================================================================
|
|
5426
|
-
*/function mf({a:n,b:t,transposeA:e,transposeB:s,backend:r,bias:o=null,preluActivationWeights:i=null,leakyreluAlpha:a=0,activation:l=null}){const u=n.shape.length,c=t.shape.length,h=e?n.shape[u-2]:n.shape[u-1],f=s?t.shape[c-1]:t.shape[c-2],d=e?n.shape[u-1]:n.shape[u-2],p=s?t.shape[c-2]:t.shape[c-1],g=n.shape.slice(0,-2),m=t.shape.slice(0,-2),b=z(g),y=z(m),x=zt(n.shape.slice(0,-2),t.shape.slice(0,-2)).concat([d,p]);w(h===f,()=>`Error in matMul: inner shapes (${h}) and (${f}) of Tensors with shapes ${n.shape} and ${t.shape} and transposeA=${e} and transposeB=${s} must match.`);const $=e?[b,h,d]:[b,d,h],E=s?[y,p,f]:[y,f,p],D=dt({inputs:{x:n},backend:r,attrs:{shape:$}}),_=dt({inputs:{x:t},backend:r,attrs:{shape:E}}),T=[D,_],P=Math.max(b,y),B=[D,_],Y=[{type:"int32",data:[d]},{type:"int32",data:[p]},{type:"int32",data:[h]}];let j,F;const G=[P,d,p];let q=W().get("WEBGPU_MATMUL_PROGRAM_TYPE");if(q<0){const
|
|
5426
|
+
*/function mf({a:n,b:t,transposeA:e,transposeB:s,backend:r,bias:o=null,preluActivationWeights:i=null,leakyreluAlpha:a=0,activation:l=null}){const u=n.shape.length,c=t.shape.length,h=e?n.shape[u-2]:n.shape[u-1],f=s?t.shape[c-1]:t.shape[c-2],d=e?n.shape[u-1]:n.shape[u-2],p=s?t.shape[c-2]:t.shape[c-1],g=n.shape.slice(0,-2),m=t.shape.slice(0,-2),b=z(g),y=z(m),x=zt(n.shape.slice(0,-2),t.shape.slice(0,-2)).concat([d,p]);w(h===f,()=>`Error in matMul: inner shapes (${h}) and (${f}) of Tensors with shapes ${n.shape} and ${t.shape} and transposeA=${e} and transposeB=${s} must match.`);const $=e?[b,h,d]:[b,d,h],E=s?[y,p,f]:[y,f,p],D=dt({inputs:{x:n},backend:r,attrs:{shape:$}}),_=dt({inputs:{x:t},backend:r,attrs:{shape:E}}),T=[D,_],P=Math.max(b,y),B=[D,_],Y=[{type:"int32",data:[d]},{type:"int32",data:[p]},{type:"int32",data:[h]}];let j,F;const G=[P,d,p];let q=W().get("WEBGPU_MATMUL_PROGRAM_TYPE");if(q<0){const qt=W().getNumber("WEBGPU_THRESHOLD_TO_INCREASE_WORKGROUPS_FOR_MATMUL"),jt=qt>0?qt:r.thresholdToIncreaseWorkgroups,Xe=P*Math.ceil(d/32)*Math.ceil(p/32);Xe<=jt||d<=8&&Xe<=jt*2?P*d*p<=128?q=Ne.MatMulReduceProgram:P===1&&f>=2e3?q=Ne.MatMulSplitKProgram:q=Ne.MatMulSmallOutputSizeProgram:q=Ne.MatMulPackedProgram}switch(q){case Ne.MatMulReduceProgram:j=new O$(G,e,s,o,l,i);break;case Ne.MatMulSplitKProgram:{if(F=of({backend:r,attrs:{shape:G,value:0,dtype:n.dtype}}),j=new z$(G,f,e,s),o||l){F=r.runWebGPUProgram(j,B,n.dtype,Y,F);const jt=new U$(F.shape,o,l,i);let Xe=null;const os=[F];o&&os.push(o),i&&os.push(i),l==="leakyrelu"&&(Xe=[{type:"float32",data:[a]}],jt.uniforms+=" alpha : f32,");const yf=r.runWebGPUProgram(jt,os,F.dtype,Xe);T.push(F);const yv=dt({inputs:{x:yf},backend:r,attrs:{shape:x}});T.push(yf);for(const wv of T)r.disposeData(wv.dataId);return yv}break}case Ne.MatMulSmallOutputSizeProgram:j=new F$($,E,G,e,s,o,l,i);break;case Ne.MatMulPackedProgram:const qt=r.adapterInfo.isIntel();j=new N$($,G,e,s,o,l,i,qt);break;default:throw new Error(`Unsupported MatMulProgramType ${q}.`)}o&&B.push(o),i&&B.push(i),l==="leakyrelu"&&(Y.push({type:"float32",data:[a]}),j.uniforms+=" alpha : f32,"),F=r.runWebGPUProgram(j,B,n.dtype,Y,F);const De=dt({inputs:{x:F},backend:r,attrs:{shape:x}});T.push(F);for(const qt of T)r.disposeData(qt.dataId);return De}/**
|
|
5427
5427
|
* @license
|
|
5428
5428
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
5429
5429
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -5438,7 +5438,7 @@ out_color[outIdx] = vec4f(
|
|
|
5438
5438
|
* See the License for the specific language governing permissions and
|
|
5439
5439
|
* limitations under the License.
|
|
5440
5440
|
* =============================================================================
|
|
5441
|
-
*/function
|
|
5441
|
+
*/function ro(n,t){const e=n.length;return e>=3?t?[...n.slice(0,-3),n[e-3]*n[e-2],n[e-1]]:[...n.slice(0,-3),n[e-3],n[e-2]*n[e-1]]:!t&&e===1&&n[0]>1?[n[0],1]:null}function W$({x:n,filter:t,convInfo:e,backend:s,bias:r=null,preluActivationWeights:o=null,leakyreluAlpha:i=0,activation:a=null}){const l=e.dataFormat==="channelsLast",u=!l,c=!1,h=l&&e.filterHeight===e.inHeight&&e.filterWidth===e.inWidth&&e.padInfo.type==="VALID",f=[];let d,p;if(h){const b=e.inHeight*e.inWidth*e.inChannels;d=dt({inputs:{x:n},backend:s,attrs:{shape:[1,e.batchSize,b]}}),p=dt({inputs:{x:t},backend:s,attrs:{shape:[1,b,e.outChannels]}})}else d=dt({inputs:{x:n},backend:s,attrs:{shape:l?[e.batchSize,e.inHeight*e.inWidth,e.inChannels]:[e.batchSize,e.inChannels,e.inHeight*e.inWidth]}}),p=dt({inputs:{x:t},backend:s,attrs:{shape:[1,e.inChannels,e.outChannels]}});if(f.push(d),f.push(p),o!=null){const b=ro(o.shape,l);b!=null&&(o=dt({inputs:{x:o},backend:s,attrs:{shape:b}}),f.push(o))}if(r!=null){const b=ro(r.shape,l);b!=null&&(r=dt({inputs:{x:r},backend:s,attrs:{shape:b}}),f.push(r))}const g=mf({a:l?d:p,b:l?p:d,transposeA:u,transposeB:c,backend:s,bias:r,activation:a,preluActivationWeights:o,leakyreluAlpha:i}),m=dt({inputs:{x:g},backend:s,attrs:{shape:e.outShape}});f.push(g);for(const b of f)s.disposeData(b.dataId);return m}function G$({x:n,filter:t,convInfo:e,backend:s,bias:r=null,preluActivationWeights:o=null,leakyreluAlpha:i=0,activation:a=null}){const{filterWidth:l,filterHeight:u,inChannels:c,strideWidth:h,strideHeight:f,padInfo:d,outWidth:p,outHeight:g,dilationWidth:m,dilationHeight:b,dataFormat:y}=e,S=y==="channelsLast",x=l*u*c,$=g*p,E=S?[e.batchSize,$,x]:[e.batchSize,x,$],D=new L$(E,S),_=[{type:"int32",data:[d.top,d.left]},{type:"int32",data:[f,h]},{type:"int32",data:[b,m]},{type:"int32",data:[p]},{type:"int32",data:[c*l]},{type:"int32",data:[c]}],T=s.runWebGPUProgram(D,[n],n.dtype,_),P=[];P.push(T);const B=dt({inputs:{x:t},backend:s,attrs:{shape:[1,x,-1]}});if(P.push(B),o!=null){const q=ro(o.shape,S);q!=null&&(o=dt({inputs:{x:o},backend:s,attrs:{shape:q}}),P.push(o))}if(r!=null){const q=ro(r.shape,S);q!=null&&(r=dt({inputs:{x:r},backend:s,attrs:{shape:q}}),P.push(r))}const F=mf({a:S?T:B,b:S?B:T,transposeA:!S,transposeB:!1,backend:s,bias:r,activation:a,preluActivationWeights:o,leakyreluAlpha:i}),G=dt({inputs:{x:F},backend:s,attrs:{shape:e.outShape}});P.push(F);for(const q of P)s.disposeData(q.dataId);return G}function V$({x:n,filter:t,convInfo:e,backend:s,bias:r=null,preluActivationWeights:o=null,leakyreluAlpha:i=0,activation:a=null}){const l=r!=null,u=o!=null,c=e.dataFormat==="channelsLast",h=c&&e.filterHeight===e.inHeight&&e.filterWidth===e.inWidth&&e.padInfo.type==="VALID",f=W().getBool("WEBGPU_USE_NAIVE_CONV2D_DEBUG");if(!f&&(h||e.filterHeight===1&&e.filterWidth===1&&e.dilationHeight===1&&e.dilationWidth===1&&e.strideHeight===1&&e.strideWidth===1&&(e.padInfo.type==="SAME"||e.padInfo.type==="VALID")))return W$({x:n,filter:t,convInfo:e,backend:s,bias:r,activation:a,preluActivationWeights:o,leakyreluAlpha:i});const d=W().getNumber("WEBGPU_THRESHOLD_TO_INCREASE_WORKGROUPS_FOR_MATMUL"),p=d>-1?d:s.thresholdToIncreaseWorkgroups,g=e.batchSize*Math.ceil(e.outHeight*e.outWidth/32)*Math.ceil(e.outChannels/32);if(W().getBool("WEBGPU_CONV_SEPARATE_IM2COL_SHADER")||g<=p)return G$({x:n,filter:t,convInfo:e,backend:s,bias:r,preluActivationWeights:o,leakyreluAlpha:i,activation:a});let m;const b=[e.padInfo.top,e.padInfo.left],y=[{type:"int32",data:[e.filterHeight,e.filterWidth]},{type:"int32",data:[...b]},{type:"int32",data:[e.strideHeight,e.strideWidth]},{type:"int32",data:[e.dilationHeight,e.dilationWidth]}];if(f)m=new P$(e,l,a,u);else{const E=c?e.outHeight*e.outWidth:e.outChannels,D=c?e.outChannels:e.outHeight*e.outWidth,_=e.filterHeight*e.filterWidth*e.inChannels;y.push({type:"int32",data:[E]},{type:"int32",data:[D]},{type:"int32",data:[_]});const T=s.adapterInfo.isIntel();m=new R$(e,E,D,_,l,a,u,T)}const S=[],x=[n,t];l&&(!c&&r.shape.length===1&&(r=dt({inputs:{x:r},backend:s,attrs:{shape:[r.shape[0],1,1]}}),S.push(r)),x.push(r)),u&&(!c&&o.shape.length===1&&(o=dt({inputs:{x:o},backend:s,attrs:{shape:[o.shape[0],1,1]}}),S.push(o)),x.push(o)),a==="leakyrelu"&&(y.push({type:"float32",data:[i]}),m.uniforms+=" alpha : f32,");const $=s.runWebGPUProgram(m,x,n.dtype,y);for(const E of S)s.disposeData(E.dataId);return $}/**
|
|
5442
5442
|
* @license
|
|
5443
5443
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
5444
5444
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -5599,7 +5599,7 @@ out_color[outIdx] = vec4f(
|
|
|
5599
5599
|
}
|
|
5600
5600
|
}
|
|
5601
5601
|
}
|
|
5602
|
-
`}}function J$(n){const t=n.length;if(t>6)throw Error(`Transpose for rank ${t} is not yet supported`);const e=new Array(t);for(let s=0;s<n.length;s++)e[n[s]]=`coords.${
|
|
5602
|
+
`}}function J$(n){const t=n.length;if(t>6)throw Error(`Transpose for rank ${t} is not yet supported`);const e=new Array(t);for(let s=0;s<n.length;s++)e[n[s]]=`coords.${In(s)}`;return e.join()}/**
|
|
5603
5603
|
* @license
|
|
5604
5604
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
5605
5605
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -5692,7 +5692,7 @@ out_color[outIdx] = vec4f(
|
|
|
5692
5692
|
* See the License for the specific language governing permissions and
|
|
5693
5693
|
* limitations under the License.
|
|
5694
5694
|
* =============================================================================
|
|
5695
|
-
*/const tv={mean:"float32",all:"bool",any:"bool"};function ev(n,t,e,s,r){const o=n.shape.length,i=[],a=
|
|
5695
|
+
*/const tv={mean:"float32",all:"bool",any:"bool"};function ev(n,t,e,s,r){const o=n.shape.length,i=[],a=is(t,n.shape);let l=a;const u=mg(l,o);let c=n;u!=null&&(c=Z$({inputs:{x:n},attrs:{perm:u},backend:r}),l=gg(l.length,o),i.push(c)),pg(s,l,o);const[h,f]=Wo(c.shape,l);let d=h;e&&(d=Cl(h,a));let p;if(r.shouldExecuteOnCPU([c])){const g=r.tensorMap.get(c.dataId).values;switch(s){case"max":const m=JS(g,z(f),d,n.dtype);p=r.makeTensorInfo(d,n.dtype,m);break;case"prod":const{outVals:b,outShape:y,outDtype:S}=ZS(c.shape,c.dtype,g,l);p=r.makeTensorInfo(y,S,b);break;default:throw new Error(`${s} CPU implementation is not yet supported.`)}}else{const g=z(f),b=z(c.shape)/g,y={windowSize:g,inSize:g,batchSize:b,outSize:1},S=tv[s]||Tp(n.dtype),x=[{type:"int32",data:[g]}],$=new Q$(y,s,r.device.limits.maxComputeWorkgroupSizeX),E=r.runWebGPUProgram($,[c],S,x);i.push(E),p=dt({inputs:{x:E},attrs:{shape:d},backend:r})}return i.forEach(g=>r.disposeData(g.dataId)),p}/**
|
|
5696
5696
|
* @license
|
|
5697
5697
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
5698
5698
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -5813,7 +5813,7 @@ out_color[outIdx] = vec4f(
|
|
|
5813
5813
|
* See the License for the specific language governing permissions and
|
|
5814
5814
|
* limitations under the License.
|
|
5815
5815
|
* =============================================================================
|
|
5816
|
-
*/class uv{constructor(t){this.uniforms="",this.workPerThread=1,this.workgroupSize=[64,1,1],this.size=!0,this.outputShape=
|
|
5816
|
+
*/class uv{constructor(t){this.uniforms="",this.workPerThread=1,this.workgroupSize=[64,1,1],this.size=!0,this.outputShape=gs(t,1),this.variableNames=t.map((e,s)=>`T${s}`),this.dispatchLayout=ae(this.outputShape),this.dispatch=It(this.dispatchLayout,this.outputShape,this.workgroupSize,[this.workPerThread,1,1]),this.offsetLength=t.length-1;for(let e=0;e<this.offsetLength;e++)this.uniforms+=`offset${e} : i32,`;this.shaderKey="concat"}getUserCode(){const t=[];if(this.offsetLength>0){t.push("if (yC < uniforms.offset0){ setOutputAtCoords(coords.x, coords.y, getT0(yR, yC)); }");for(let o=1;o<this.offsetLength;o++)t.push(`else if (yC < uniforms.offset${[o]}){ setOutputAtCoords(coords.x, coords.y, getT${o}(yR, yC - uniforms.offset${o-1})); }`);const s=this.offsetLength,r=this.offsetLength-1;t.push(`else { setOutputAtCoords(coords.x, coords.y, getT${s}(yR, yC - uniforms.offset${r})); }`)}else t.push("setOutputAtCoords(coords.x, coords.y, getT0(yR, yC));");return`
|
|
5817
5817
|
${wt("index")} {
|
|
5818
5818
|
for(var i = 0; i < ${this.workPerThread}; i = i + 1) {
|
|
5819
5819
|
let flatIndex = index * ${this.workPerThread} + i;
|
|
@@ -5887,7 +5887,7 @@ out_color[outIdx] = vec4f(
|
|
|
5887
5887
|
* See the License for the specific language governing permissions and
|
|
5888
5888
|
* limitations under the License.
|
|
5889
5889
|
* =============================================================================
|
|
5890
|
-
*/function
|
|
5890
|
+
*/function Us(n,t,e){const s=n[0].dtype;if(s==="complex64"){const p=n.map(S=>fv({inputs:{input:S},backend:e})),g=n.map(S=>hv({inputs:{input:S},backend:e})),m=Us(p,t,e),b=Us(g,t,e),y=cv({inputs:{real:m,imag:b},backend:e});return p.forEach(S=>e.disposeData(S.dataId)),g.forEach(S=>e.disposeData(S.dataId)),e.disposeData(m.dataId),e.disposeData(b.dataId),y}let r=e.shouldExecuteOnCPU(n);if(s==="string"&&(r=!0),r){const p=n.map($=>{const D=[-1,z($.shape.slice(t))];return dt({inputs:{x:$},backend:e,attrs:{shape:D}})}),g=p.map($=>({vals:e.readSync($.dataId),shape:$.shape})),m=gs(p.map($=>$.shape),1),b=p[0].shape[0]===1,y=XS(g,m,s,b),S=gs(n.map($=>$.shape),t),x=e.makeTensorInfo(S,s,y);return p.forEach($=>e.disposeData($.dataId)),x}const o=e.device.limits.maxStorageBuffersPerShaderStage-1;if(n.length>o){const p=[];for(let m=0;m<n.length;m+=o){const b=n.slice(m,m+o);p.push(Us(b,t,e))}const g=Us(p,t,e);for(const m of p)e.disposeData(m.dataId);return g}const{tensors2D:i,outShape:a}=dv(n,t,e),l=i.map(p=>p.shape),u=new uv(l),c=[],h=new Array(l.length-1);if(h.length>0){h[0]=l[0][1],c.push({type:"int32",data:[h[0]]});for(let p=1;p<h.length;p++)h[p]=h[p-1]+l[p][1],c.push({type:"int32",data:[h[p]]})}const f=e.runWebGPUProgram(u,i,i[0].dtype,c);i.forEach(p=>e.disposeData(p.dataId));const d=dt({inputs:{x:f},backend:e,attrs:{shape:a}});return e.disposeData(f.dataId),d}function dv(n,t,e){const s=gs(n.map(o=>o.shape),t);return{tensors2D:n.map(o=>dt({inputs:{x:o},backend:e,attrs:{shape:[z(o.shape.slice(0,t)),z(o.shape.slice(t))]}})),outShape:s}}/**
|
|
5891
5891
|
* @license
|
|
5892
5892
|
* Copyright 2021 Google LLC. All Rights Reserved.
|
|
5893
5893
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
|
@@ -5902,4 +5902,4 @@ out_color[outIdx] = vec4f(
|
|
|
5902
5902
|
* See the License for the specific language governing permissions and
|
|
5903
5903
|
* limitations under the License.
|
|
5904
5904
|
* =============================================================================
|
|
5905
|
-
*/function pv(n){const{inputs:t,backend:e,attrs:s}=n,{axis:r}=s,o=
|
|
5905
|
+
*/function pv(n){const{inputs:t,backend:e,attrs:s}=n,{axis:r}=s,o=is(r,t[0].shape)[0],i=t.map(u=>u.shape);uy(i,o);const a=gs(t.map(u=>u.shape),o);if(z(a)===0)return e.makeTensorInfo(a,t[0].dtype,[]);const l=t.filter(u=>z(u.shape)>0);return l.length===1?Ye({inputs:{x:l[0]},backend:e}):Us(l,o,e)}const mv=[Fx,Vx,r2,j$,ov,lv,{kernelName:Ca,backendName:"webgpu",kernelFunc:pv},zx];for(const n of mv)lp({...n,backendName:"webgpu-oidn"});async function gv(){try{const n={powerPreference:"high-performance"},t=await navigator.gpu.requestAdapter(n),e={},s=[];t.features.has("timestamp-query")&&s.push("timestamp-query"),t.features.has("bgra8unorm-storage")&&s.push(["bgra8unorm-storage"]),e.requiredFeatures=s;const r=t.limits;e.requiredLimits={maxComputeWorkgroupStorageSize:r.maxComputeWorkgroupStorageSize,maxComputeWorkgroupsPerDimension:r.maxComputeWorkgroupsPerDimension,maxStorageBufferBindingSize:r.maxStorageBufferBindingSize,maxBufferSize:r.maxBufferSize,maxComputeWorkgroupSizeX:r.maxComputeWorkgroupSizeX,maxComputeInvocationsPerWorkgroup:r.maxComputeInvocationsPerWorkgroup};const o=await t.requestDevice(e),i=await t.requestAdapterInfo();return gf(o,i)}catch{}}async function gf(n,t){const e=new Fs(n,t);return A.registerBackend("webgpu-oidn",()=>e),await A.setBackend("webgpu-oidn"),e}async function bf(n,t,e){const s=await(t?gf(t.device,t.adapterInfo):gv()),r=ga(n);return new Yh(r,s,e)}async function bv(n,t,e){return fetch(n).then(s=>s.arrayBuffer()).then(s=>bf(s,t,e))}Nt.UNet=Yh,Nt.initUNetFromBuffer=bf,Nt.initUNetFromURL=bv,Nt.parseTZA=ga,Object.defineProperty(Nt,Symbol.toStringTag,{value:"Module"})});
|