oidn-web 0.4.0 → 0.5.0
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/CHANGELOG.md +45 -0
- package/README.md +83 -19
- package/dist/oidn.js +3239 -2642
- package/dist/oidn.umd.cjs +512 -296
- package/lib/UNet.d.ts +54 -20
- package/lib/UNet.js +194 -118
- package/lib/UNet.js.map +1 -1
- package/lib/backend.d.ts +1 -8
- package/lib/backend.js +1 -9
- package/lib/backend.js.map +1 -1
- package/lib/finalRgbShader.d.ts +13 -0
- package/lib/finalRgbShader.js +160 -0
- package/lib/finalRgbShader.js.map +1 -0
- package/lib/graphOptimizer.js +1 -2
- package/lib/graphOptimizer.js.map +1 -1
- package/lib/hdrTransfer.d.ts +14 -0
- package/lib/hdrTransfer.js +61 -0
- package/lib/hdrTransfer.js.map +1 -0
- package/lib/main.d.ts +16 -7
- package/lib/main.js +4 -5
- package/lib/main.js.map +1 -1
- package/lib/nativeUNet.d.ts +39 -3
- package/lib/nativeUNet.js +449 -120
- package/lib/nativeUNet.js.map +1 -1
- package/lib/process.d.ts +5 -11
- package/lib/process.js +35 -49
- package/lib/process.js.map +1 -1
- package/lib/tileScheduler.d.ts +32 -4
- package/lib/tileScheduler.js +133 -20
- package/lib/tileScheduler.js.map +1 -1
- package/package.json +9 -2
- package/src/UNet.ts +287 -158
- package/src/backend.ts +1 -14
- package/src/finalRgbShader.ts +186 -0
- package/src/graphOptimizer.ts +1 -2
- package/src/hdrTransfer.ts +88 -0
- package/src/main.ts +28 -13
- package/src/nativeUNet.ts +515 -116
- package/src/process.ts +43 -70
- package/src/tileScheduler.ts +216 -24
- package/benchmarks/compare.mjs +0 -651
- package/benchmarks/leak.mjs +0 -255
- package/benchmarks/results/before-spatial.json +0 -391
- package/benchmarks/results/before-spatial.md +0 -47
- package/benchmarks/results/int8-scan.json +0 -2007
- package/benchmarks/results/int8-scan.md +0 -160
- package/benchmarks/results/int8-w8a8-scan.json +0 -2007
- package/benchmarks/results/int8-w8a8-scan.md +0 -160
- package/benchmarks/results/int8-weight-channel.json +0 -1413
- package/benchmarks/results/int8-weight-channel.md +0 -118
- package/benchmarks/results/int8-weight-only.json +0 -1437
- package/benchmarks/results/int8-weight-only.md +0 -118
- package/benchmarks/results/kernel-webnn-final.json +0 -1115
- package/benchmarks/results/kernel-webnn-final.md +0 -104
- package/benchmarks/results/latest-optimized.json +0 -375
- package/benchmarks/results/latest-optimized.md +0 -47
- package/benchmarks/results/latest.json +0 -391
- package/benchmarks/results/latest.md +0 -47
- package/benchmarks/results/profile-baseline.json +0 -331
- package/benchmarks/results/profile-baseline.md +0 -12
- package/benchmarks/results/profile-conv2x.json +0 -331
- package/benchmarks/results/profile-conv2x.md +0 -12
- package/benchmarks/results/profile-fast-init.json +0 -385
- package/benchmarks/results/profile-fast-init.md +0 -47
- package/benchmarks/results/profile-fp16-fma.json +0 -369
- package/benchmarks/results/profile-fp16-fma.md +0 -47
- package/benchmarks/results/profile-fp16-tiled-decoder.json +0 -369
- package/benchmarks/results/profile-fp16-tiled-decoder.md +0 -47
- package/benchmarks/results/profile-fp16-tiled-encoder.json +0 -369
- package/benchmarks/results/profile-fp16-tiled-encoder.md +0 -47
- package/benchmarks/results/profile-fp16-unfused-pool.json +0 -385
- package/benchmarks/results/profile-fp16-unfused-pool.md +0 -47
- package/benchmarks/results/profile-input-major.json +0 -347
- package/benchmarks/results/profile-input-major.md +0 -12
- package/benchmarks/results/profile-k16.json +0 -347
- package/benchmarks/results/profile-k16.md +0 -12
- package/benchmarks/results/profile-k4.json +0 -347
- package/benchmarks/results/profile-k4.md +0 -12
- package/benchmarks/results/profile-pool-reuse.json +0 -331
- package/benchmarks/results/profile-pool-reuse.md +0 -12
- package/benchmarks/results/profile-precompiled.json +0 -385
- package/benchmarks/results/profile-precompiled.md +0 -47
- package/benchmarks/results/profile-static-channels.json +0 -385
- package/benchmarks/results/profile-static-channels.md +0 -47
- package/benchmarks/results/profile-static-io.json +0 -385
- package/benchmarks/results/profile-static-io.md +0 -47
- package/benchmarks/results/profile-tiled-conv.json +0 -331
- package/benchmarks/results/profile-tiled-conv.md +0 -12
- package/benchmarks/results/profile-tiled-decoder.json +0 -347
- package/benchmarks/results/profile-tiled-decoder.md +0 -12
- package/benchmarks/results/profile-tiled-matmul.json +0 -331
- package/benchmarks/results/profile-tiled-matmul.md +0 -12
- package/benchmarks/results/profile-unfused-decoder.json +0 -379
- package/benchmarks/results/profile-unfused-decoder.md +0 -12
- package/benchmarks/results/profile-unfused-pool.json +0 -347
- package/benchmarks/results/profile-unfused-pool.md +0 -12
- package/benchmarks/results/spatial-auto.json +0 -575
- package/benchmarks/results/spatial-auto.md +0 -61
- package/benchmarks/results/subgroup-smoke.json +0 -1094
- package/benchmarks/results/subgroup-smoke.md +0 -104
- package/benchmarks/results/webnn-smoke.json +0 -739
- package/benchmarks/results/webnn-smoke.md +0 -76
- package/scripts/inspect-model.mjs +0 -64
- package/tests/modelSpec.test.mjs +0 -128
- package/tests/resourceLifecycle.test.mjs +0 -383
- package/tests/tileScheduler.test.mjs +0 -90
package/dist/oidn.umd.cjs
CHANGED
|
@@ -1,49 +1,57 @@
|
|
|
1
|
-
(function(
|
|
2
|
-
${
|
|
1
|
+
(function(A,K){typeof exports=="object"&&typeof module<"u"?K(exports):typeof define=="function"&&define.amd?define(["exports"],K):(A=typeof globalThis<"u"?globalThis:A||self,K(A.oidn={}))})(this,function(A){"use strict";var No=Object.defineProperty;var Uo=(A,K,ve)=>K in A?No(A,K,{enumerable:!0,configurable:!0,writable:!0,value:ve}):A[K]=ve;var m=(A,K,ve)=>Uo(A,typeof K!="symbol"?K+"":K,ve);class K{constructor(){m(this,"dims",[]);m(this,"paddedDims",[]);m(this,"layout","x");m(this,"dataType","Float32")}getByteSize(){let e=1;for(const t of this.paddedDims)e*=t;return this.dataType==="Float32"?e*=4:this.dataType==="Float16"&&(e*=2),e}}class ve{constructor(e,t){this.desc=e,this.data=t}}class Sn{constructor(e){m(this,"offset",0);this._view=e}read(e){const t=this._view,i=this.offset;switch(this.offset+=e,e){case 1:return t.getUint8(i);case 2:return t.getUint16(i,!0);case 4:return t.getUint32(i,!0);case 8:return Number(t.getBigUint64(i,!0));default:throw new Error("unsupported read size")}}}function Dt(n){const e=new Uint8Array(n),t=new Sn(new DataView(n));if(t.read(2)!==16855)throw new Error("invalid or corrupted weights blob");const r=t.read(1);if(t.read(1),r!==2)throw new Error("unsupported weights blob version");const o=t.read(8);t.offset=o;const a=t.read(4),u=new Map;for(let s=0;s<a;++s){const l=new K,p=t.read(2),c=new TextDecoder().decode(e.subarray(t.offset,t.offset+p));t.offset+=p;const d=t.read(1);for(let v=0;v<d;++v)l.dims.push(t.read(4));l.paddedDims=[...l.dims],new TextDecoder().decode(e.subarray(t.offset,t.offset+d))==="oihw"&&(l.layout="oihw"),t.offset+=d;const f=String.fromCharCode(t.read(1));if(f==="f")l.dataType="Float32";else if(f==="h")l.dataType="Float16";else throw new Error("invalid tensor data type");const g=t.read(8),y=e.slice(g,g+l.getByteSize());u.set(c,new ve(l,y))}return u}function En(n,e){return n.channels===e.channels}const We=8;class De{constructor(e,t,i){m(this,"autoUpdateOutputBuffer",!0);m(this,"_label");m(this,"_device");m(this,"_outputBuffers",{});m(this,"_pipeline");m(this,"_bindGroups",[]);m(this,"_needsUpdatePipeline",!0);m(this,"_needsResizeBuffer",!0);m(this,"_inputs",[]);m(this,"_outputs",[]);m(this,"_uniforms",[]);m(this,"_uniformBuffers",{});m(this,"_width",10);m(this,"_height",10);m(this,"_execWidth");m(this,"_execHeight");m(this,"_csCode","");m(this,"_csMain");m(this,"_csDefine");m(this,"_groupOffsets",{inputs:0,uniforms:1,outputs:2});this._label=e,this._device=t,this._csMain=i.csMain,this._csDefine=i.csDefine,this._inputs=i.inputs,this._outputs=i.outputs,this._uniforms=i.uniforms,this.autoUpdateOutputBuffer=i.autoUpdateOutputBuffer??!0,i.uniforms.forEach(r=>{this._uniformBuffers[r.label]=t.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)})}setCSCode({csDefine:e,csMain:t}){this._csDefine=e,this._csMain=t,this._needsUpdatePipeline=!0}setSize(e,t){e=Math.ceil(e),t=Math.ceil(t);const i=e!==this._width||t!==this._height;this._width=e,this._height=t,i&&(this._needsResizeBuffer=!0,this._needsUpdatePipeline=!0)}setExecuteSize(e,t){e=Math.ceil(e),t=Math.ceil(t),this._execWidth=e,this._execHeight=t}setOutputParams(e){this.autoUpdateOutputBuffer&&this._updateOutputBuffers(e),this._needsUpdatePipeline=!0}setOutputBuffers(e){this._outputBuffers=Object.keys(e).reduce((t,i)=>(t[i]={buffer:e[i],params:{channels:4}},t),{})}setUniform(e,t){const i=this._uniformBuffers[e];this._device.queue.writeBuffer(i,0,t)}getOutput(e){return this._needsResizeBuffer&&this.autoUpdateOutputBuffer&&(this._resizeOutputBuffers(),this._needsResizeBuffer=!1),this._outputBuffers[e].buffer}dispose(e=!0){Object.keys(this._uniformBuffers).forEach(t=>{this._uniformBuffers[t].destroy()}),e&&Object.keys(this._outputBuffers).forEach(t=>{this._outputBuffers[t].buffer.destroy()})}_createBuffer(e){const t=this._width*this._height*4*4;return this._device.createBuffer({label:this._label,size:Math.max(t,80),usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST|GPUBufferUsage.COPY_SRC})}_resizeOutputBuffers(){const e=this._outputBuffers;for(const t in e){const{buffer:i,params:r}=e[t];i.destroy(),e[t].buffer=this._createBuffer(r)}}_updateOutputBuffers(e){var i,r;const t=this._outputBuffers;for(const o in e){const a=e[o];if(!En(a,((i=t[o])==null?void 0:i.params)||{})){(r=t[o])==null||r.buffer.destroy();const u=this._createBuffer(a);t[o]={buffer:u,params:a}}}}_updatePipeline(e,t){if(!this._needsUpdatePipeline)return;this._needsUpdatePipeline=!1;const i=this._device,r=this._getFullCs(e,t);r!==this._csCode&&(this._csCode=r,this._pipeline=i.createComputePipeline({label:this._label,layout:"auto",compute:{module:i.createShaderModule({label:this._label,code:r}),entryPoint:"main"}}),this._updateBindGroups())}_getFullCs(e,t){const i=this._inputs,r=this._uniforms;let o=0;const a=this._groupOffsets={inputs:0,uniforms:0,outputs:0};return i.length>0&&o++,r.length>0&&(a.uniforms=o,o++),a.outputs=o,`
|
|
2
|
+
${i.sort().map((s,l)=>{const p=`@group(${a.inputs}) @binding(${l}) `,c=`in_${s}`;return t[s]==="texture"?`${p} var ${c}: texture_2d<f32>;`:`${p} var<storage, read> ${c}: array<vec${e[s].channels}f>;`}).join(`
|
|
3
3
|
`)}
|
|
4
|
-
${this._uniforms.map((
|
|
4
|
+
${this._uniforms.map((s,l)=>`@group(${a.uniforms}) @binding(${l}) var<uniform> ${s.label}: ${s.type};`).join(`
|
|
5
5
|
`)}
|
|
6
6
|
|
|
7
|
-
${this._outputs.map((
|
|
7
|
+
${this._outputs.map((s,l)=>`@group(${a.outputs}) @binding(${l}) var<storage, read_write> out_${s}: array<vec${this._outputBuffers[s].params.channels}f>;`).join(`
|
|
8
8
|
`)}
|
|
9
9
|
${this._csDefine??""}
|
|
10
10
|
@compute @workgroup_size(${We}, ${We}, 1)
|
|
11
11
|
fn main(@builtin(global_invocation_id) globalId: vec3u) {
|
|
12
12
|
${this._csMain}
|
|
13
13
|
}
|
|
14
|
-
`}_updateBindGroups(){const e=[],t=this._device,
|
|
15
|
-
const a = ${
|
|
16
|
-
const b = ${
|
|
17
|
-
const c = ${
|
|
18
|
-
const d = ${
|
|
19
|
-
const e = ${
|
|
20
|
-
const f = ${
|
|
21
|
-
const g = ${
|
|
22
|
-
const y0 =${
|
|
23
|
-
const y1 =${
|
|
24
|
-
const x0 =${
|
|
25
|
-
const x1 =${
|
|
26
|
-
|
|
27
|
-
const normScale = ${
|
|
28
|
-
const rcpNormScale = ${
|
|
29
|
-
`;class
|
|
14
|
+
`}_updateBindGroups(){const e=[],t=this._device,i=this._groupOffsets;this._uniforms.length>0&&(e[i.uniforms]=t.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(i.uniforms),entries:this._uniforms.map((r,o)=>({binding:o,resource:{buffer:this._uniformBuffers[r.label]}}))})),this._bindGroups=e}createPass(e,t){this._needsResizeBuffer&&this.autoUpdateOutputBuffer&&(this._resizeOutputBuffers(),this._needsResizeBuffer=!1);const i=this._inputs.reduce((a,u)=>(a[u]=t[u].buffer?"buffer":"texture",a),{});this._updatePipeline(t,i);const r=this._groupOffsets;this._inputs.length>0&&(this._bindGroups[r.inputs]=this._device.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(r.inputs),entries:this._inputs.map((a,u)=>({binding:u,resource:t[a].buffer?{buffer:t[a].buffer}:t[a].texture.createView()}))})),this._bindGroups[r.outputs]=this._device.createBindGroup({label:this._label,layout:this._pipeline.getBindGroupLayout(r.outputs),entries:this._outputs.map((a,u)=>({binding:u,resource:{buffer:this._outputBuffers[a].buffer}}))});const o=e.beginComputePass();o.setPipeline(this._pipeline),this._bindGroups.forEach((a,u)=>{o.setBindGroup(u,a)}),o.dispatchWorkgroups(Math.ceil((this._execWidth??this._width)/We),Math.ceil((this._execHeight??this._height)/We),1),o.end()}}const Rt=1412.83765,Gt=1.64593172,Yt=.431384981,Xt=-.00294139609,Ft=.192653254,Vt=.00626026094,Ht=.998620152,An=15794576e-13,Mn=.0322087631,On=.00223151711,Nn=.370974749,jt=65504,Kt=Zt(jt),Un=1/Kt,zn=Kt,qt=Math.log(jt+1),Ln=1/qt;function Zt(n){return n<=An?Rt*n:n<=Mn?Gt*Math.pow(n,Yt)+Xt:Ft*Math.log(n+Vt)+Ht}function Wn(n){return n<=On?n/Rt:n<=Nn?Math.pow((n-Xt)/Gt,1/Yt):Math.exp((n-Ht)/Ft)-Vt}function Dn(n,e){return e==="log"?Math.log(n+1)*Ln:Zt(n)*Un}function Rn(n,e){return e==="log"?Math.exp(n*qt)-1:Wn(n*zn)}function Gn({data:n,channels:e,inputScale:t,transfer:i="pu"}){const r=new Float32Array(n);for(let o=0;o<r.length;o+=e)for(let a=0;a<3;a++)r[o+a]=Dn(r[o+a]*t,i);return r}function Yn({data:n,channels:e,inputScale:t,transfer:i="pu"}){const r=new Float32Array(n),o=1/t;for(let a=0;a<r.length;a+=e)for(let u=0;u<3;u++)r[a+u]=Rn(r[a+u],i)*o;return r}const Jt=1412.83765,Qt=1.64593172,ei=.431384981,ti=-.00294139609,ii=.192653254,ni=.00626026094,ri=.998620152,oi=15794576e-13,ai=.0322087631,Xn=.00223151711,Fn=.370974749;function Vn(n){return n<=oi?n=Jt*n:n<=ai?n=Qt*Math.pow(n,ei)+ti:n=ii*Math.log(n+ni)+ri,n}const si=65504,ui=Vn(si),Hn=1/ui,jn=ui,ct=Math.log(si+1),Kn=1/ct;class lt{constructor(e,t,i,r){this.x=e,this.y=t,this.width=i,this.height=r}}function qn({data:n,channels:e}){let t=0;for(let a=0;a<n.length;a+=e){const u=n[a],s=n[a+1],l=n[a+2],p=.212671*u+.71516*s+.072169*l;t+=Math.log2(p+1e-4)}const i=n.length/e,r=t/i;return .18/Math.pow(2,r)}const ci=`
|
|
15
|
+
const a = ${Jt};
|
|
16
|
+
const b = ${Qt};
|
|
17
|
+
const c = ${ei};
|
|
18
|
+
const d = ${ti};
|
|
19
|
+
const e = ${ii};
|
|
20
|
+
const f = ${ni};
|
|
21
|
+
const g = ${ri};
|
|
22
|
+
const y0 =${oi};
|
|
23
|
+
const y1 =${ai};
|
|
24
|
+
const x0 =${Xn};
|
|
25
|
+
const x1 =${Fn};
|
|
26
|
+
|
|
27
|
+
const normScale = ${Hn};
|
|
28
|
+
const rcpNormScale = ${jn};
|
|
29
|
+
`;class Zn{constructor(e,t,i="pu"){m(this,"_inputPassAux");m(this,"_inputPassColor");m(this,"_outputPass");m(this,"_copyPass");m(this,"_isInputTexture");this._device=e,this._isHDR=t,this._hdrTransfer=i;const r=[{label:"inputScale",type:"f32",data:new Float32Array([1])},{label:"inputSize",type:"vec2i",data:new Int32Array(2)},{label:"outputSize",type:"vec2i",data:new Int32Array(2)},{label:"inputOffset",type:"vec2i",data:new Int32Array(2)}];this._inputPassAux=new De("inputPassAux",this._device,{inputs:["color","albedo","normal"],outputs:["color","albedo","normal"],uniforms:r,csDefine:"",csMain:""}),this._inputPassColor=new De("inputPassColor",this._device,{inputs:["color"],outputs:["color"],uniforms:r,csDefine:"",csMain:""}),this._outputPass=new De("outputPass",this._device,{inputs:["color","raw"],outputs:["color"],uniforms:[{label:"inputScale",type:"f32",data:new Float32Array([1])},{label:"inputSize",type:"vec2i",data:new Int32Array(2)},{label:"outputSize",type:"vec2i",data:new Int32Array(2)},{label:"imageSize",type:"vec2i",data:new Int32Array(2)},{label:"inputOffset",type:"vec2i",data:new Int32Array(2)},{label:"outputOffset",type:"vec2i",data:new Int32Array(2)}],csDefine:"",csMain:""}),this._copyPass=new De("copyPass",this._device,{inputs:["color"],outputs:["color"],autoUpdateOutputBuffer:!1,uniforms:[{label:"size",type:"vec2i",data:new Int32Array(2)}],csMain:`
|
|
30
30
|
let outIdx = i32(globalId.x + globalId.y * u32(size.x));
|
|
31
31
|
out_color[outIdx] = textureLoad(in_color, globalId.xy, 0);
|
|
32
|
-
`}),this._inputPassAux.setOutputParams({color:{channels:3},albedo:{channels:3},normal:{channels:3}}),this._inputPassColor.setOutputParams({color:{channels:3}}),this._outputPass.setOutputParams({color:{channels:4}})}_updatePasses(e,t=!1){if(this._isInputTexture!=null&&this._isInputTexture===e)return;this._isInputTexture=e;const
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
42
|
-
|
|
32
|
+
`}),this._inputPassAux.setOutputParams({color:{channels:3},albedo:{channels:3},normal:{channels:3}}),this._inputPassColor.setOutputParams({color:{channels:3}}),this._outputPass.setOutputParams({color:{channels:4}})}_updatePasses(e,t=!1){if(this._isInputTexture!=null&&this._isInputTexture===e)return;this._isInputTexture=e;const i=this._isHDR,r=this._hdrTransfer==="log"?"fn HDRForward(y: f32) -> f32 { return log(y + 1.0) * logNormScale; }":`fn HDRForward(y: f32) -> f32 {
|
|
33
|
+
if (y <= y0) { return a * y * normScale; }
|
|
34
|
+
else if (y <= y1) { return (b * pow(y, c) + d) * normScale; }
|
|
35
|
+
else { return (e * log(y + f) + g) * normScale; }
|
|
36
|
+
}`,o=this._hdrTransfer==="log"?"fn HDRInverse(x: f32) -> f32 { return exp(x * logXMax) - 1.0; }":`fn HDRInverse(x: f32) -> f32 {
|
|
37
|
+
let y = x * rcpNormScale;
|
|
38
|
+
if (y <= x0) { return y / a; }
|
|
39
|
+
else if (y <= x1) { return pow((y - d) / b, 1 / c); }
|
|
40
|
+
else { return exp((y - g) / e) - f; }
|
|
41
|
+
}`,a=`
|
|
42
|
+
${ci}
|
|
43
|
+
const logXMax = ${ct};
|
|
44
|
+
const logNormScale = ${Kn};
|
|
45
|
+
${r}`;function u(l){return e?`textureLoad(in_${l}, vec2u(inputPosition), 0)`:`in_${l}[inIdx]`}const s=`
|
|
43
46
|
let x = i32(globalId.x);
|
|
44
47
|
let y = i32(globalId.y);
|
|
45
|
-
let
|
|
46
|
-
|
|
48
|
+
let inputPosition = clamp(
|
|
49
|
+
vec2i(x, y) + inputOffset,
|
|
50
|
+
vec2i(0),
|
|
51
|
+
inputSize - vec2i(1)
|
|
52
|
+
);
|
|
53
|
+
let inIdx = inputPosition.y * inputSize.x + inputPosition.x;
|
|
54
|
+
let col = ${u("color")};
|
|
47
55
|
|
|
48
56
|
let outIdx = y * outputSize.x + x;
|
|
49
57
|
|
|
@@ -51,31 +59,24 @@ if (${t}) {
|
|
|
51
59
|
// Denoise the inversed alpha. Or the anti aliased edge will be too dark after denoised
|
|
52
60
|
out_color[outIdx] = vec3f(1.0 - col.a);
|
|
53
61
|
}
|
|
54
|
-
else if (${
|
|
55
|
-
out_color[outIdx] = vec3f(
|
|
62
|
+
else if (${i}) {
|
|
63
|
+
out_color[outIdx] = vec3f(HDRForward(col.r * inputScale), HDRForward(col.g * inputScale), HDRForward(col.b * inputScale));
|
|
56
64
|
}
|
|
57
65
|
else {
|
|
58
66
|
out_color[outIdx] = col.rgb;
|
|
59
67
|
}
|
|
60
|
-
`;this._inputPassAux.setCSCode({csDefine:
|
|
68
|
+
`;this._inputPassAux.setCSCode({csDefine:a,csMain:`
|
|
61
69
|
${s}
|
|
62
|
-
let alb = ${
|
|
63
|
-
let nor = ${
|
|
70
|
+
let alb = ${u("albedo")};
|
|
71
|
+
let nor = ${u("normal")};
|
|
64
72
|
out_normal[outIdx] = nor.rgb;
|
|
65
73
|
out_albedo[outIdx] = alb.rgb;
|
|
66
|
-
`}),this._inputPassColor.setCSCode({csDefine:
|
|
74
|
+
`}),this._inputPassColor.setCSCode({csDefine:a,csMain:`
|
|
67
75
|
${s}
|
|
68
76
|
`}),this._outputPass.setCSCode({csDefine:`
|
|
69
|
-
${
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
return y / a;
|
|
73
|
-
} else if (y <= x1) {
|
|
74
|
-
return pow((y - d) / b, 1 / c);
|
|
75
|
-
} else {
|
|
76
|
-
return exp((y - g) / e) - f;
|
|
77
|
-
}
|
|
78
|
-
}
|
|
77
|
+
${ci}
|
|
78
|
+
const logXMax = ${ct};
|
|
79
|
+
${o}
|
|
79
80
|
`,csMain:`
|
|
80
81
|
let x = i32(globalId.x);
|
|
81
82
|
let y = i32(globalId.y);
|
|
@@ -90,9 +91,9 @@ let raw = ${e?"textureLoad(in_raw, globalId.xy + vec2u(outputOffset), 0)":"in_ra
|
|
|
90
91
|
if (${t}) {
|
|
91
92
|
out_color[outIdx] = vec4f(raw.rgb, 1.0 - col.r);
|
|
92
93
|
}
|
|
93
|
-
else if (${
|
|
94
|
+
else if (${i}) {
|
|
94
95
|
out_color[outIdx] = vec4f(
|
|
95
|
-
vec3f(
|
|
96
|
+
vec3f(HDRInverse(col.r), HDRInverse(col.g), HDRInverse(col.b)) / inputScale,
|
|
96
97
|
// Pick the alpha
|
|
97
98
|
raw.a
|
|
98
99
|
);
|
|
@@ -100,16 +101,126 @@ else if (${n}) {
|
|
|
100
101
|
else {
|
|
101
102
|
out_color[outIdx] = vec4f(col.rgb, raw.a);
|
|
102
103
|
}
|
|
103
|
-
`})}setImageSize(e,t){this._inputPassAux.setUniform("inputSize",new Int32Array([e,t])),this._inputPassColor.setUniform("inputSize",new Int32Array([e,t])),this._outputPass.setUniform("imageSize",new Int32Array([e,t])),this._outputPass.setSize(e,t),this._copyPass.setSize(e,t),this._copyPass.setUniform("size",new Int32Array([e,t]))}setInputTile(e){const t=new Int32Array([e.width,e.height]);[this._inputPassAux,this._inputPassColor].forEach(n=>{n.setUniform("inputOffset",new Int32Array([e.x,e.y])),n.setUniform("outputSize",t),n.setSize(t[0],t[1])}),this._outputPass.setUniform("inputSize",t)}setOutputTile(e,t){const n=this._outputPass,o=new Int32Array([e.width,e.height]),r=e.x-t.x,s=e.y-t.y;n.setUniform("outputSize",o),n.setUniform("inputOffset",new Int32Array([r,s])),n.setUniform("outputOffset",new Int32Array([e.x,e.y])),n.setExecuteSize(o[0],o[1])}forward(e,t,n,o){const r=e instanceof GPUTexture;this._updatePasses(r,o);const s=this._inputPassAux,a=this._inputPassColor,u=this._device.createCommandEncoder();function l(p){return p instanceof GPUTexture?{texture:p,channels:4}:{buffer:p,channels:4}}return t&&n?s.createPass(u,{color:l(e),albedo:l(t),normal:l(n)}):a.createPass(u,{color:l(e)}),this._device.queue.submit([u.finish()]),t&&n?{color:s.getOutput("color"),albedo:s.getOutput("albedo"),normal:s.getOutput("normal")}:{color:a.getOutput("color")}}inverse(e,t){const o=this._device.createCommandEncoder(),r=this._outputPass;return r.createPass(o,{color:{buffer:e,channels:4},raw:t instanceof GPUBuffer?{buffer:t,channels:4}:{texture:t,channels:4}}),this._device.queue.submit([o.finish()]),r.getOutput("color")}copyInputDataToOutput(e){const t=this._device.createCommandEncoder(),o=this._outputPass.getOutput("color"),r=this._copyPass;e instanceof GPUTexture?(r.setOutputBuffers({color:o}),r.createPass(t,{color:{texture:e,channels:4}})):t.copyBufferToBuffer(e,0,o,0,o.size),this._device.queue.submit([t.finish()])}dispose(){this._outputPass.dispose(),this._inputPassAux.dispose(),this._inputPassColor.dispose(),this._copyPass.dispose(!1)}}const hi=256,gi=384,_i=16,yi=128,ne=16;function Ce(i,e){return Math.ceil(i/e)*e}function mi(i,e){return Math.floor(i/e)*e}function ft(i,e,t){return Math.min(Math.max(i,e),t)}function vi(i){const e=[...i].sort((n,o)=>n-o),t=Math.floor(e.length/2);return e.length%2?e[t]:(e[t-1]+e[t])/2}function qt(i,e){return i<=e?Math.min(Ce(i,ne),e):e}class wi{constructor(e,t=!0){g(this,"enabled");g(this,"maxTileSize");g(this,"minTileSize");g(this,"targetTileTimeMs");g(this,"_tileSize");g(this,"_adjustmentStep");const n=typeof t=="object"?t:{};this.enabled=t!==!1,this.maxTileSize=Math.max(ne,mi(e,ne)),this.minTileSize=ft(Ce(n.minTileSize??hi,ne),ne,this.maxTileSize),this.targetTileTimeMs=Math.max(1,n.targetTileTimeMs??_i),this._adjustmentStep=Math.max(ne,Ce(n.adjustmentStep??yi,ne)),this._tileSize=this.enabled?ft(Ce(n.initialTileSize??gi,ne),this.minTileSize,this.maxTileSize):this.maxTileSize}get tileSize(){return this._tileSize}observe(e){if(!this.enabled||e.length===0)return!1;const t=e.filter(r=>Number.isFinite(r)&&r>=0);if(t.length===0)return!1;const n=vi(t);let o=this._tileSize;return n>this.targetTileTimeMs*1.25?o-=this._adjustmentStep:n<this.targetTileTimeMs*.65&&(o+=this._adjustmentStep),o=ft(Ce(o,ne),this.minTileSize,this.maxTileSize),o===this._tileSize?!1:(this._tileSize=o,!0)}}async function xi(i){try{await i.onSubmittedWorkDone()}catch{}}function w(i,e){return{op:"conv2d",id:i,input:e,weight:`${i}.weight`,bias:`${i}.bias`,activation:"relu",padding:"same"}}function fe(i,e){return{op:"maxPool2d",id:i,input:e,size:2,stride:2,padding:"same"}}function de(i,e){return{op:"upsample2d",id:i,input:e,scale:2,mode:"nearest"}}function he(i,e,t){return{op:"concat",id:i,inputs:[e,t],axis:"channels"}}const jt={schemaVersion:1,id:"oidn-unet-small-v1",family:"oidn-unet-small",input:"input",output:"dec_conv0",receptiveField:174,nodes:[w("enc_conv0","input"),w("enc_conv1","enc_conv0"),fe("pool1","enc_conv1"),w("enc_conv2","pool1"),fe("pool2","enc_conv2"),w("enc_conv3","pool2"),fe("pool3","enc_conv3"),w("enc_conv4","pool3"),fe("pool4","enc_conv4"),w("enc_conv5a","pool4"),w("enc_conv5b","enc_conv5a"),de("up4","enc_conv5b"),he("concat4","up4","pool3"),w("dec_conv4a","concat4"),w("dec_conv4b","dec_conv4a"),de("up3","dec_conv4b"),he("concat3","up3","pool2"),w("dec_conv3a","concat3"),w("dec_conv3b","dec_conv3a"),de("up2","dec_conv3b"),he("concat2","up2","pool1"),w("dec_conv2a","concat2"),w("dec_conv2b","dec_conv2a"),de("up1","dec_conv2b"),he("concat1","up1","input"),w("dec_conv1a","concat1"),w("dec_conv1b","dec_conv1a"),w("dec_conv0","dec_conv1b")]},Zt={schemaVersion:1,id:"oidn-unet-large-v1",family:"oidn-unet-large",input:"input",output:"dec_conv1c",receptiveField:202,nodes:[w("enc_conv1a","input"),w("enc_conv1b","enc_conv1a"),fe("pool1","enc_conv1b"),w("enc_conv2a","pool1"),w("enc_conv2b","enc_conv2a"),fe("pool2","enc_conv2b"),w("enc_conv3a","pool2"),w("enc_conv3b","enc_conv3a"),fe("pool3","enc_conv3b"),w("enc_conv4a","pool3"),w("enc_conv4b","enc_conv4a"),fe("pool4","enc_conv4b"),w("enc_conv5a","pool4"),w("enc_conv5b","enc_conv5a"),de("up4","enc_conv5b"),he("concat4","up4","pool3"),w("dec_conv4a","concat4"),w("dec_conv4b","dec_conv4a"),de("up3","dec_conv4b"),he("concat3","up3","pool2"),w("dec_conv3a","concat3"),w("dec_conv3b","dec_conv3a"),de("up2","dec_conv3b"),he("concat2","up2","pool1"),w("dec_conv2a","concat2"),w("dec_conv2b","dec_conv2a"),de("up1","dec_conv2b"),he("concat1","up1","input"),w("dec_conv1a","concat1"),w("dec_conv1b","dec_conv1a"),w("dec_conv1c","dec_conv1b")]},bi=[jt,Zt];function Jt(i){const e=new Set;for(const t of i.nodes)t.op==="conv2d"&&(e.add(t.weight),e.add(t.bias));return e}function Qt(i){return i.desc.getByteSize()}function en(i){return[...i].sort().join(", ")}function dt(i,e=bi){const t=e.filter(n=>{const o=Jt(n);return[...o].some(r=>!i.has(r))?!1:n.allowAdditionalTensors===!0||[...i.keys()].every(r=>o.has(r))});if(t.length===1)return t[0];throw t.length>1?new Error(`Ambiguous OIDN model topology: ${t.map(n=>n.id).join(", ")}`):new Error(`Unsupported OIDN model topology. TZA tensors: ${en(i.keys())}`)}function tn(i,e,t){const n=i.get(e);if(!n)throw new Error(`Model ${t} is missing tensor ${e}`);if(n.data.byteLength!==Qt(n))throw new Error(`Tensor ${e} has ${n.data.byteLength} bytes, expected ${Qt(n)}`);return n}function nn(i,e=dt(i)){if(e.schemaVersion!==1)throw new Error(`Unsupported model descriptor schema ${e.schemaVersion}`);const t=Jt(e);if(!e.allowAdditionalTensors){const c=[...i.keys()].filter(f=>!t.has(f));if(c.length>0)throw new Error(`Model ${e.id} has unexpected tensors: ${en(c)}`)}const n=new Map,o=new Map,r=new Map,s=new Set([e.input]);let a,u;const l=(c,f)=>{const d=n.get(c);if(d===void 0)throw new Error(`Model ${e.id} node ${f} reads unknown or forward value ${c}`);return d};for(const c of e.nodes){if(s.has(c.id))throw new Error(`Model ${e.id} produces duplicate value ${c.id}`);if(c.op==="conv2d"){const f=tn(i,c.weight,e.id),d=tn(i,c.bias,e.id),h=f.desc.dims;if(f.desc.layout!=="oihw"||h.length!==4)throw new Error(`Tensor ${c.weight} must use OIHW layout`);if(h[2]!==3||h[3]!==3)throw new Error(`Tensor ${c.weight} must use a 3x3 kernel`);if(d.desc.layout!=="x"||d.desc.dims.length!==1)throw new Error(`Tensor ${c.bias} must be a one-dimensional bias`);if(d.desc.dims[0]!==h[0])throw new Error(`Tensor ${c.bias} has ${d.desc.dims[0]} channels, expected ${h[0]}`);if(f.desc.dataType!==d.desc.dataType)throw new Error(`Weight and bias dtype differ for ${c.id}`);if(u&&u!==f.desc.dataType)throw new Error(`Mixed tensor dtypes are not supported by model ${e.id}`);u=f.desc.dataType,c.input===e.input&&a===void 0&&(a=h[1],n.set(e.input,a));const _=l(c.input,c.id);if(_!==h[1])throw new Error(`Tensor ${c.weight} expects ${h[1]} input channels, but ${c.input} provides ${_}`);n.set(c.id,h[0]),o.set(c.id,{weight:f,bias:d,inputChannels:h[1],outputChannels:h[0],kernelHeight:h[2],kernelWidth:h[3]}),r.set(c.id,{inputChannels:h[1],outputChannels:h[0]})}else if(c.op==="concat"){if(c.inputs.length<2)throw new Error(`Concat ${c.id} requires at least two inputs`);const f=c.inputs.reduce((d,h)=>d+l(h,c.id),0);n.set(c.id,f)}else n.set(c.id,l(c.input,c.id));s.add(c.id)}if(a===void 0||u===void 0)throw new Error(`Model ${e.id} has no convolution reading its input`);const p=n.get(e.output);if(p===void 0)throw new Error(`Model ${e.id} output ${e.output} is not produced`);if(p!==3)throw new Error(`Model ${e.id} must produce 3 channels, got ${p}`);return{spec:e,inputChannels:a,outputChannels:p,tensorDataType:u,channelsByValue:n,convChannels:r,convTensors:o}}const $i="This is not an object",ki="This is not a Float16Array object",rn="This constructor is not a subclass of Float16Array",on="The constructor property value is not an object",Bi="Species constructor didn't return TypedArray object",Pi="Derived constructor created TypedArray object which was too small length",Se="Attempting to access detached ArrayBuffer",ht="Cannot convert undefined or null to object",gt="Cannot mix BigInt and other types, use explicit conversions",sn="@@iterator property is not callable",an="Reduce of empty array with no initial value",Ii="The comparison function must be either a function or undefined",_t="Offset is out of bounds";function T(i){return(e,...t)=>V(i,e,t)}function be(i,e){return T($e(i,e).get)}const{apply:V,construct:Ee,defineProperty:un,get:yt,getOwnPropertyDescriptor:$e,getPrototypeOf:Ae,has:mt,ownKeys:cn,set:ln,setPrototypeOf:pn}=Reflect,Ti=Proxy,{EPSILON:Ci,MAX_SAFE_INTEGER:fn,isFinite:dn,isNaN:ke}=Number,{iterator:oe,species:Si,toStringTag:vt,for:Ei}=Symbol,Be=Object,{create:Re,defineProperty:Oe,freeze:Ai,is:hn}=Be,wt=Be.prototype,Oi=wt.__lookupGetter__?T(wt.__lookupGetter__):(i,e)=>{if(i==null)throw S(ht);let t=Be(i);do{const n=$e(t,e);if(n!==void 0)return ce(n,"get")?n.get:void 0}while((t=Ae(t))!==null)},ce=Be.hasOwn||T(wt.hasOwnProperty),gn=Array,_n=gn.isArray,Ge=gn.prototype,Ui=T(Ge.join),Ni=T(Ge.push),Mi=T(Ge.toLocaleString),xt=Ge[oe],zi=T(xt),{abs:Di,trunc:yn}=Math,Ye=ArrayBuffer,Wi=Ye.isView,mn=Ye.prototype,Li=T(mn.slice),Ri=be(mn,"byteLength"),bt=typeof SharedArrayBuffer<"u"?SharedArrayBuffer:null,Gi=bt&&be(bt.prototype,"byteLength"),$t=Ae(Uint8Array),Yi=$t.from,D=$t.prototype,Fi=D[oe],Xi=T(D.keys),Hi=T(D.values),Vi=T(D.entries),Ki=T(D.set),vn=T(D.reverse),qi=T(D.fill),ji=T(D.copyWithin),wn=T(D.sort),Ue=T(D.slice),Zi=T(D.subarray),W=be(D,"buffer"),me=be(D,"byteOffset"),$=be(D,"length"),xn=be(D,vt),Ji=Uint8Array,J=Uint16Array,bn=(...i)=>V(Yi,J,i),kt=Uint32Array,Qi=Float32Array,ve=Ae([][oe]()),Fe=T(ve.next),er=T(function*(){}().next),tr=Ae(ve),S=TypeError,Bt=RangeError,$n=WeakSet,kn=$n.prototype,nr=T(kn.add),ir=T(kn.has),Xe=WeakMap,Pt=Xe.prototype,He=T(Pt.get),rr=T(Pt.has),It=T(Pt.set),Bn=new Xe,or=Re(null,{next:{value:function(){const e=He(Bn,this);return Fe(e)}},[oe]:{value:function(){return this}}});function Ve(i){if(i[oe]===xt&&ve.next===Fe)return i;const e=Re(or);return It(Bn,e,zi(i)),e}const Pn=new Xe,In=Re(tr,{next:{value:function(){const e=He(Pn,this);return er(e)},writable:!0,configurable:!0}});for(const i of cn(ve))i!=="next"&&Oe(In,i,$e(ve,i));function Tn(i){const e=Re(In);return It(Pn,e,i),e}function Ke(i){return i!==null&&typeof i=="object"||typeof i=="function"}function Cn(i){return i!==null&&typeof i=="object"}function qe(i){return xn(i)!==void 0}function Tt(i){const e=xn(i);return e==="BigInt64Array"||e==="BigUint64Array"}function sr(i){try{return _n(i)?!1:(Ri(i),!0)}catch{return!1}}function Sn(i){if(bt===null)return!1;try{return Gi(i),!0}catch{return!1}}function ar(i){return sr(i)||Sn(i)}function En(i){return _n(i)?i[oe]===xt&&ve.next===Fe:!1}function ur(i){return qe(i)?i[oe]===Fi&&ve.next===Fe:!1}function je(i){if(typeof i!="string")return!1;const e=+i;return i!==e+""||!dn(e)?!1:e===yn(e)}const Ze=Ei("__Float16Array__");function cr(i){if(!Cn(i))return!1;const e=Ae(i);if(!Cn(e))return!1;const t=e.constructor;if(t===void 0)return!1;if(!Ke(t))throw S(on);return mt(t,Ze)}const Ct=1/Ci;function lr(i){return i+Ct-Ct}const An=6103515625e-14,pr=65504,On=.0009765625,Un=On*An,fr=On*Ct;function dr(i){const e=+i;if(!dn(e)||e===0)return e;const t=e>0?1:-1,n=Di(e);if(n<An)return t*lr(n/Un)*Un;const o=(1+fr)*n,r=o-(o-n);return r>pr||ke(r)?t*(1/0):t*r}const Nn=new Ye(4),Mn=new Qi(Nn),zn=new kt(Nn),ie=new J(512),re=new Ji(512);for(let i=0;i<256;++i){const e=i-127;e<-24?(ie[i]=0,ie[i|256]=32768,re[i]=24,re[i|256]=24):e<-14?(ie[i]=1024>>-e-14,ie[i|256]=1024>>-e-14|32768,re[i]=-e-1,re[i|256]=-e-1):e<=15?(ie[i]=e+15<<10,ie[i|256]=e+15<<10|32768,re[i]=13,re[i|256]=13):e<128?(ie[i]=31744,ie[i|256]=64512,re[i]=24,re[i|256]=24):(ie[i]=31744,ie[i|256]=64512,re[i]=13,re[i|256]=13)}function se(i){Mn[0]=dr(i);const e=zn[0],t=e>>23&511;return ie[t]+((e&8388607)>>re[t])}const St=new kt(2048);for(let i=1;i<1024;++i){let e=i<<13,t=0;for(;!(e&8388608);)e<<=1,t-=8388608;e&=-8388609,t+=947912704,St[i]=e|t}for(let i=1024;i<2048;++i)St[i]=939524096+(i-1024<<13);const Pe=new kt(64);for(let i=1;i<31;++i)Pe[i]=i<<23;Pe[31]=1199570944,Pe[32]=2147483648;for(let i=33;i<63;++i)Pe[i]=2147483648+(i-32<<23);Pe[63]=3347054592;const Dn=new J(64);for(let i=1;i<64;++i)i!==32&&(Dn[i]=1024);function B(i){const e=i>>10;return zn[0]=St[Dn[e]+(i&1023)]+Pe[e],Mn[0]}function le(i){const e=+i;return ke(e)||e===0?0:yn(e)}function Et(i){const e=le(i);return e<0?0:e<fn?e:fn}function Je(i,e){if(!Ke(i))throw S($i);const t=i.constructor;if(t===void 0)return e;if(!Ke(t))throw S(on);const n=t[Si];return n??e}function Ne(i){if(Sn(i))return!1;try{return Li(i,0,0),!1}catch{}return!0}function Wn(i,e){const t=ke(i),n=ke(e);if(t&&n)return 0;if(t)return 1;if(n||i<e)return-1;if(i>e)return 1;if(i===0&&e===0){const o=hn(i,0),r=hn(e,0);if(!o&&r)return-1;if(o&&!r)return 1}return 0}const At=2,Qe=new Xe;function Ie(i){return rr(Qe,i)||!Wi(i)&&cr(i)}function k(i){if(!Ie(i))throw S(ki)}function et(i,e){const t=Ie(i),n=qe(i);if(!t&&!n)throw S(Bi);if(typeof e=="number"){let o;if(t){const r=v(i);o=$(r)}else o=$(i);if(o<e)throw S(Pi)}if(Tt(i))throw S(gt)}function v(i){const e=He(Qe,i);if(e!==void 0){const o=W(e);if(Ne(o))throw S(Se);return e}const t=i.buffer;if(Ne(t))throw S(Se);const n=Ee(P,[t,i.byteOffset,i.length],i.constructor);return He(Qe,n)}function Ln(i){const e=$(i),t=[];for(let n=0;n<e;++n)t[n]=B(i[n]);return t}const Rn=new $n;for(const i of cn(D)){if(i===vt)continue;const e=$e(D,i);ce(e,"get")&&typeof e.get=="function"&&nr(Rn,e.get)}const hr=Ai({get(i,e,t){return je(e)&&ce(i,e)?B(yt(i,e)):ir(Rn,Oi(i,e))?yt(i,e):yt(i,e,t)},set(i,e,t,n){return je(e)&&ce(i,e)?ln(i,e,se(t)):ln(i,e,t,n)},getOwnPropertyDescriptor(i,e){if(je(e)&&ce(i,e)){const t=$e(i,e);return t.value=B(t.value),t}return $e(i,e)},defineProperty(i,e,t){return je(e)&&ce(i,e)&&ce(t,"value")&&(t.value=se(t.value)),un(i,e,t)}});class P{constructor(e,t,n){let o;if(Ie(e))o=Ee(J,[v(e)],new.target);else if(Ke(e)&&!ar(e)){let s,a;if(qe(e)){s=e,a=$(e);const u=W(e);if(Ne(u))throw S(Se);if(Tt(e))throw S(gt);const l=new Ye(a*At);o=Ee(J,[l],new.target)}else{const u=e[oe];if(u!=null&&typeof u!="function")throw S(sn);u!=null?En(e)?(s=e,a=e.length):(s=[...e],a=s.length):(s=e,a=Et(s.length)),o=Ee(J,[a],new.target)}for(let u=0;u<a;++u)o[u]=se(s[u])}else o=Ee(J,arguments,new.target);const r=new Ti(o,hr);return It(Qe,r,o),r}static from(e,...t){const n=this;if(!mt(n,Ze))throw S(rn);if(n===P){if(Ie(e)&&t.length===0){const p=v(e),c=new J(W(p),me(p),$(p));return new P(W(Ue(c)))}if(t.length===0)return new P(W(bn(e,se)));const u=t[0],l=t[1];return new P(W(bn(e,function(p,...c){return se(V(u,this,[p,...Ve(c)]))},l)))}let o,r;const s=e[oe];if(s!=null&&typeof s!="function")throw S(sn);if(s!=null)En(e)?(o=e,r=e.length):ur(e)?(o=e,r=$(e)):(o=[...e],r=o.length);else{if(e==null)throw S(ht);o=Be(e),r=Et(o.length)}const a=new n(r);if(t.length===0)for(let u=0;u<r;++u)a[u]=o[u];else{const u=t[0],l=t[1];for(let p=0;p<r;++p)a[p]=V(u,l,[o[p],p])}return a}static of(...e){const t=this;if(!mt(t,Ze))throw S(rn);const n=e.length;if(t===P){const r=new P(n),s=v(r);for(let a=0;a<n;++a)s[a]=se(e[a]);return r}const o=new t(n);for(let r=0;r<n;++r)o[r]=e[r];return o}keys(){k(this);const e=v(this);return Xi(e)}values(){k(this);const e=v(this);return Tn(function*(){for(const t of Hi(e))yield B(t)}())}entries(){k(this);const e=v(this);return Tn(function*(){for(const[t,n]of Vi(e))yield[t,B(n)]}())}at(e){k(this);const t=v(this),n=$(t),o=le(e),r=o>=0?o:n+o;if(!(r<0||r>=n))return B(t[r])}with(e,t){k(this);const n=v(this),o=$(n),r=le(e),s=r>=0?r:o+r,a=+t;if(s<0||s>=o)throw Bt(_t);const u=new J(W(n),me(n),$(n)),l=new P(W(Ue(u))),p=v(l);return p[s]=se(a),l}map(e,...t){k(this);const n=v(this),o=$(n),r=t[0],s=Je(n,P);if(s===P){const u=new P(o),l=v(u);for(let p=0;p<o;++p){const c=B(n[p]);l[p]=se(V(e,r,[c,p,this]))}return u}const a=new s(o);et(a,o);for(let u=0;u<o;++u){const l=B(n[u]);a[u]=V(e,r,[l,u,this])}return a}filter(e,...t){k(this);const n=v(this),o=$(n),r=t[0],s=[];for(let l=0;l<o;++l){const p=B(n[l]);V(e,r,[p,l,this])&&Ni(s,p)}const a=Je(n,P),u=new a(s);return et(u),u}reduce(e,...t){k(this);const n=v(this),o=$(n);if(o===0&&t.length===0)throw S(an);let r,s;t.length===0?(r=B(n[0]),s=1):(r=t[0],s=0);for(let a=s;a<o;++a)r=e(r,B(n[a]),a,this);return r}reduceRight(e,...t){k(this);const n=v(this),o=$(n);if(o===0&&t.length===0)throw S(an);let r,s;t.length===0?(r=B(n[o-1]),s=o-2):(r=t[0],s=o-1);for(let a=s;a>=0;--a)r=e(r,B(n[a]),a,this);return r}forEach(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=0;s<o;++s)V(e,r,[B(n[s]),s,this])}find(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=0;s<o;++s){const a=B(n[s]);if(V(e,r,[a,s,this]))return a}}findIndex(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=0;s<o;++s){const a=B(n[s]);if(V(e,r,[a,s,this]))return s}return-1}findLast(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=o-1;s>=0;--s){const a=B(n[s]);if(V(e,r,[a,s,this]))return a}}findLastIndex(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=o-1;s>=0;--s){const a=B(n[s]);if(V(e,r,[a,s,this]))return s}return-1}every(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=0;s<o;++s)if(!V(e,r,[B(n[s]),s,this]))return!1;return!0}some(e,...t){k(this);const n=v(this),o=$(n),r=t[0];for(let s=0;s<o;++s)if(V(e,r,[B(n[s]),s,this]))return!0;return!1}set(e,...t){k(this);const n=v(this),o=le(t[0]);if(o<0)throw Bt(_t);if(e==null)throw S(ht);if(Tt(e))throw S(gt);if(Ie(e))return Ki(v(this),v(e),o);if(qe(e)){const u=W(e);if(Ne(u))throw S(Se)}const r=$(n),s=Be(e),a=Et(s.length);if(o===1/0||a+o>r)throw Bt(_t);for(let u=0;u<a;++u)n[u+o]=se(s[u])}reverse(){k(this);const e=v(this);return vn(e),this}toReversed(){k(this);const e=v(this),t=new J(W(e),me(e),$(e)),n=new P(W(Ue(t))),o=v(n);return vn(o),n}fill(e,...t){k(this);const n=v(this);return qi(n,se(e),...Ve(t)),this}copyWithin(e,t,...n){k(this);const o=v(this);return ji(o,e,t,...Ve(n)),this}sort(e){k(this);const t=v(this),n=e!==void 0?e:Wn;return wn(t,(o,r)=>n(B(o),B(r))),this}toSorted(e){k(this);const t=v(this);if(e!==void 0&&typeof e!="function")throw new S(Ii);const n=e!==void 0?e:Wn,o=new J(W(t),me(t),$(t)),r=new P(W(Ue(o))),s=v(r);return wn(s,(a,u)=>n(B(a),B(u))),r}slice(e,t){k(this);const n=v(this),o=Je(n,P);if(o===P){const h=new J(W(n),me(n),$(n));return new P(W(Ue(h,e,t)))}const r=$(n),s=le(e),a=t===void 0?r:le(t);let u;s===-1/0?u=0:s<0?u=r+s>0?r+s:0:u=r<s?r:s;let l;a===-1/0?l=0:a<0?l=r+a>0?r+a:0:l=r<a?r:a;const p=l-u>0?l-u:0,c=new o(p);if(et(c,p),p===0)return c;const f=W(n);if(Ne(f))throw S(Se);let d=0;for(;u<l;)c[d]=B(n[u]),++u,++d;return c}subarray(e,t){k(this);const n=v(this),o=Je(n,P),r=new J(W(n),me(n),$(n)),s=Zi(r,e,t),a=new o(W(s),me(s),$(s));return et(a),a}indexOf(e,...t){k(this);const n=v(this),o=$(n);let r=le(t[0]);if(r===1/0)return-1;r<0&&(r+=o,r<0&&(r=0));for(let s=r;s<o;++s)if(ce(n,s)&&B(n[s])===e)return s;return-1}lastIndexOf(e,...t){k(this);const n=v(this),o=$(n);let r=t.length>=1?le(t[0]):o-1;if(r===-1/0)return-1;r>=0?r=r<o-1?r:o-1:r+=o;for(let s=r;s>=0;--s)if(ce(n,s)&&B(n[s])===e)return s;return-1}includes(e,...t){k(this);const n=v(this),o=$(n);let r=le(t[0]);if(r===1/0)return!1;r<0&&(r+=o,r<0&&(r=0));const s=ke(e);for(let a=r;a<o;++a){const u=B(n[a]);if(s&&ke(u)||u===e)return!0}return!1}join(e){k(this);const t=v(this),n=Ln(t);return Ui(n,e)}toLocaleString(...e){k(this);const t=v(this),n=Ln(t);return Mi(n,...Ve(e))}get[vt](){if(Ie(this))return"Float16Array"}}Oe(P,"BYTES_PER_ELEMENT",{value:At}),Oe(P,Ze,{}),pn(P,$t);const tt=P.prototype;Oe(tt,"BYTES_PER_ELEMENT",{value:At}),Oe(tt,oe,{value:tt.values,writable:!0,configurable:!0}),pn(tt,D);function gr(i){return i.op==="concat"?i.inputs:[i.input]}function _r(i){const e=new Map;for(const t of i.nodes)for(const n of gr(t)){const o=e.get(n)??[];o.push(t),e.set(n,o)}return e}function Ot(i,e,t){const n=i.get(e);if(!((n==null?void 0:n.length)!==1||n[0].op!==t))return n[0]}function nt(i,e={}){const t=i.spec,n=_r(t),o=new Map(t.nodes.map(p=>[p.id,p])),r=new Set,s=new Map;let a=0,u=0;if(e.fuseConvPool!==!1)for(const p of t.nodes){if(p.op!=="conv2d"||p.activation!=="relu")continue;const c=Ot(n,p.id,"maxPool2d");!c||c.size!==2||c.stride!==2||(r.add(p.id),s.set(c.id,{op:"fusedConvReluMaxPool2d",id:c.id,input:p.input,conv:p,pool:c}),a++)}if(e.fuseUpsampleConcatConv!==!1)for(const p of t.nodes){if(p.op!=="conv2d")continue;const c=o.get(p.input);if((c==null?void 0:c.op)!=="concat"||c.inputs.length!==2||Ot(n,c.id,"conv2d")!==p)continue;const f=i.channelsByValue.get(c.inputs[0]);if(f===void 0||f%4!==0)continue;const d=c.inputs.map(_=>{const y=o.get(_);return(y==null?void 0:y.op)==="upsample2d"&&y.scale===2&&y.mode==="nearest"&&Ot(n,y.id,"concat")===c?{value:y.input,upsample:y}:{value:_}});if(d.filter(_=>_.upsample).length===1){r.add(c.id);for(const _ of c.inputs){const y=o.get(_);(y==null?void 0:y.op)==="upsample2d"&&r.add(y.id)}s.set(p.id,{op:"fusedUpsampleConcatConv2d",id:p.id,inputs:d,conv:p}),u++}}const l=[];for(const p of t.nodes){const c=s.get(p.id);c?l.push(c):r.has(p.id)||l.push(p)}return{spec:t,nodes:l,fusions:{convPool:a,upsampleConcatConv:u}}}function yr(i){return i.op==="concat"?i.inputs:i.op==="fusedUpsampleConcatConv2d"?i.inputs.map(e=>e.value):[i.input]}function Gn(i,e){return i.width===e.width&&i.height===e.height}function Yn(i,e,t,n){if(!Number.isInteger(e)||e<=0||!Number.isInteger(t)||t<=0)throw new Error(`Invalid model input size ${e}x${t}`);const o=nt(i,n),r={width:e,height:t,channels:i.inputChannels},s=new Map([[i.spec.input,r]]),a=[],u=(c,f)=>{const d=s.get(c);if(!d)throw new Error(`Planned node ${f} reads missing value ${c}`);return d};for(const c of o.nodes){let f;if(c.op==="conv2d"){const d=u(c.input,c.id);f={width:d.width,height:d.height,channels:i.convChannels.get(c.id).outputChannels}}else if(c.op==="maxPool2d"){const d=u(c.input,c.id);f={width:Math.ceil(d.width/2),height:Math.ceil(d.height/2),channels:d.channels}}else if(c.op==="upsample2d"){const d=u(c.input,c.id);f={width:d.width*2,height:d.height*2,channels:d.channels}}else if(c.op==="concat"){const d=c.inputs.map(h=>u(h,c.id));if(d.some(h=>!Gn(h,d[0])))throw new Error(`Concat ${c.id} has mismatched spatial shapes`);f={width:d[0].width,height:d[0].height,channels:d.reduce((h,_)=>h+_.channels,0)}}else if(c.op==="fusedConvReluMaxPool2d"){const d=u(c.input,c.id);f={width:Math.ceil(d.width/2),height:Math.ceil(d.height/2),channels:i.convChannels.get(c.conv.id).outputChannels}}else{const d=c.inputs.map(h=>{const _=u(h.value,c.id);return h.upsample?{..._,width:_.width*2,height:_.height*2}:_});if(d.some(h=>!Gn(h,d[0])))throw new Error(`Fused decoder ${c.id} has mismatched spatial shapes`);f={width:d[0].width,height:d[0].height,channels:i.convChannels.get(c.conv.id).outputChannels}}s.set(c.id,f),a.push(f)}const l=new Map;o.nodes.forEach((c,f)=>{for(const d of yr(c))l.set(d,f)}),l.set(i.spec.output,o.nodes.length);const p=o.nodes.map((c,f)=>({node:c,outputShape:a[f],lastUse:l.get(c.id)??f}));return{...o,inputShape:r,valueShapes:s,plannedNodes:p}}class Fn{constructor(){g(this,"_stats",new Map)}track(e,t){let n=this._stats.get(e);return n||(n={created:0,destroyed:0,live:0,peakLive:0,resources:new Set},this._stats.set(e,n)),n.resources.has(t)||(n.resources.add(t),n.created++,n.live++,n.peakLive=Math.max(n.peakLive,n.live)),t}release(e,t,n){if(!t)return!1;const o=this._stats.get(e);if(!(o!=null&&o.resources.delete(t)))return!1;try{n()}finally{o.destroyed++,o.live--}return!0}snapshot(e=0){let t=0,n=0,o=0,r=0;const s={};for(const[a,u]of this._stats){const l={created:u.created,destroyed:u.destroyed,live:u.live,peakLive:u.peakLive};s[a]=l,t+=l.live,n+=l.created,o+=l.destroyed,r+=l.peakLive}return{live:t,created:n,destroyed:o,peakLive:r,pending:e,byKind:s}}}const Xn=new WeakMap;function mr(i){let e=Xn.get(i);return e||(e={ready:new Map,pending:new Map},Xn.set(i,e)),e}const Y=8,ee=8,F=4,we=ee*F,Q=ee,K=8,E=8,te=E+2;function Ut(i,e){return Math.ceil(i/e)*e}function U(i){return Math.ceil(i/4)}function vr(i,e){return i.width*i.height*U(i.channels)*4*e}let Me;function Hn(i){const e=i&32768?-1:1,t=i>>>10&31,n=i&1023;return t===0?e*n*2**-24:t===31?n===0?e*(1/0):NaN:e*(1+n/1024)*2**(t-15)}function wr(){if(!Me){Me=new Float32Array(65536);for(let i=0;i<Me.length;i++)Me[i]=Hn(i)}return Me}function Vn(i){if(i.desc.dataType==="Float32")return new Float32Array(i.data.buffer,i.data.byteOffset,i.data.byteLength/4);const e=new Uint16Array(i.data.buffer,i.data.byteOffset,i.data.byteLength/2),t=new Float32Array(e.length);if(e.length<4096)for(let n=0;n<e.length;n++)t[n]=Hn(e[n]);else{const n=wr();for(let o=0;o<e.length;o++)t[o]=n[e[o]]}return t}function Nt(i,e,t,n){const o=Ut(t.byteLength,4),r=i.createBuffer({label:e,size:o,usage:n,mappedAtCreation:!0});return new Uint8Array(r.getMappedRange()).set(new Uint8Array(t.buffer,t.byteOffset,t.byteLength)),r.unmap(),r}function Kn(i,e,t){const n=new Uint32Array(Ut(t.length,4));return n.set(t),Nt(i,e,n,GPUBufferUsage.UNIFORM)}function xr(i,e,t,n){const o=U(t.inputChannels),r=U(t.outputChannels),s=r*t.kernelHeight*t.kernelWidth*o*4*4,a=n==="fp16"&&t.weight.desc.dataType==="Float16",u=a?new Uint16Array(s):n==="fp16"?new P(s):new Float32Array(s),l=a?new Uint16Array(t.weight.data.buffer,t.weight.data.byteOffset,t.weight.data.byteLength/2):Vn(t.weight);for(let f=0;f<r;f++)for(let d=0;d<t.kernelHeight;d++)for(let h=0;h<t.kernelWidth;h++)for(let _=0;_<o;_++)for(let y=0;y<4;y++){const x=f*4+y;for(let m=0;m<4;m++){const b=_*4+m,L=(((f*t.kernelHeight+d)*t.kernelWidth+h)*o+_)*16+m*4+y;if(x<t.outputChannels&&b<t.inputChannels){const q=((x*t.inputChannels+b)*t.kernelHeight+d)*t.kernelWidth+h;u[L]=l[q]}}}const p=new Float32Array(r*4);p.set(Vn(t.bias));const c=Nt(i,`oidn/${e}/weights/${n}`,u,GPUBufferUsage.STORAGE);try{return{weights:c,bias:Nt(i,`oidn/${e}/bias`,p,GPUBufferUsage.STORAGE)}}catch(f){throw c.destroy(),f}}function X(i){return i==="fp16"?"vec4<f16>":"vec4<f32>"}function ae(i){return i==="fp16"?`enable f16;
|
|
104
|
-
`:""}function
|
|
105
|
-
let inputValue = vec4<
|
|
104
|
+
`})}setImageSize(e,t){this._inputPassAux.setUniform("inputSize",new Int32Array([e,t])),this._inputPassColor.setUniform("inputSize",new Int32Array([e,t])),this._outputPass.setUniform("imageSize",new Int32Array([e,t])),this._outputPass.setSize(e,t),this._copyPass.setSize(e,t),this._copyPass.setUniform("size",new Int32Array([e,t]))}setInputTile(e){const t=new Int32Array([e.width,e.height]);[this._inputPassAux,this._inputPassColor].forEach(i=>{i.setUniform("inputOffset",new Int32Array([e.x,e.y])),i.setUniform("outputSize",t),i.setSize(t[0],t[1])}),this._outputPass.setUniform("inputSize",t)}setOutputTile(e,t){const i=this._outputPass,r=new Int32Array([e.width,e.height]),o=e.x-t.x,a=e.y-t.y;i.setUniform("outputSize",r),i.setUniform("inputOffset",new Int32Array([o,a])),i.setUniform("outputOffset",new Int32Array([e.x,e.y])),i.setExecuteSize(r[0],r[1])}forward(e,t,i,r){const o=e instanceof GPUTexture;this._updatePasses(o,r);const a=this._inputPassAux,u=this._inputPassColor,s=this._device.createCommandEncoder();function l(p){return p instanceof GPUTexture?{texture:p,channels:4}:{buffer:p,channels:4}}return t&&i?a.createPass(s,{color:l(e),albedo:l(t),normal:l(i)}):u.createPass(s,{color:l(e)}),this._device.queue.submit([s.finish()]),t&&i?{color:a.getOutput("color"),albedo:a.getOutput("albedo"),normal:a.getOutput("normal")}:{color:u.getOutput("color")}}inverse(e,t){const r=this._device.createCommandEncoder(),o=this._outputPass;return o.createPass(r,{color:{buffer:e,channels:4},raw:t instanceof GPUBuffer?{buffer:t,channels:4}:{texture:t,channels:4}}),this._device.queue.submit([r.finish()]),o.getOutput("color")}copyInputDataToOutput(e){const t=this._device.createCommandEncoder(),r=this._outputPass.getOutput("color"),o=this._copyPass;e instanceof GPUTexture?(o.setOutputBuffers({color:r}),o.createPass(t,{color:{texture:e,channels:4}})):t.copyBufferToBuffer(e,0,r,0,r.size),this._device.queue.submit([t.finish()])}dispose(){this._outputPass.dispose(),this._inputPassAux.dispose(),this._inputPassColor.dispose(),this._copyPass.dispose(!1)}}const Jn=256,Qn=432,er=16,tr=16,D=16;function xe(n,e){return Math.ceil(n/e)*e}function li(n,e){return Math.floor(n/e)*e}function pi(n,e,t){const i=Math.round(t*n/e),r=Math.round((t+1)*n/e);return{start:i,end:r}}function Re(n,e,t,i=0){const r=e-n,o=Math.max(i,xe(r,D));let a=n-Math.floor((o-r)/2);return a=Ge(a,0,Math.max(0,t-o)),{start:a,end:a+o}}function pt(n,e,t,i){if(!Number.isInteger(n)||n<=0||!Number.isInteger(e)||e<=0)throw new Error("Tile grid dimensions must be positive integers");if(!Number.isFinite(t)||t<=0)throw new Error("Maximum tile size must be positive");if(!Number.isFinite(i)||i<0)throw new Error("Tile overlap must be non-negative");const r=Math.max(D,li(t,D)),o=xe(i,D),a=Math.max(1,Math.ceil(n/r)),u=Math.max(1,Math.ceil(e/r)),s=[];let l=0,p=0;for(let f=0;f<u;f++){const g=pi(e,u,f);for(let y=0;y<a;y++){const v=pi(n,a,y),_={start:y===0?v.start:Math.max(0,v.start-o),end:y===a-1?v.end:Math.min(n,v.end+o)},w={start:f===0?g.start:Math.max(0,g.start-o),end:f===u-1?g.end:Math.min(e,g.end+o)},b=Re(_.start,_.end,n),x=Re(w.start,w.end,e),B={x:v.start,y:g.start,width:v.end-v.start,height:g.end-g.start},k={x:b.start,y:x.start,width:b.end-b.start,height:x.end-x.start};l=Math.max(l,B.width),p=Math.max(p,B.height),s.push({column:y,row:f,input:k,output:B})}}const c=()=>new Set(s.map(({input:f})=>`${f.width}x${f.height}`));if(c().size>2){const f=Math.max(...s.map(({input:_})=>_.width)),g=Math.max(...s.map(({input:_})=>_.height)),y=s.reduce((_,{input:w})=>_+f*w.height,0),v=s.reduce((_,{input:w})=>_+w.width*g,0);if(y<=v)for(const _ of s){const w=Re(_.input.x,_.input.x+_.input.width,n,f);_.input.x=w.start,_.input.width=w.end-w.start}else for(const _ of s){const w=Re(_.input.y,_.input.y+_.input.height,e,g);_.input.y=w.start,_.input.height=w.end-w.start}}const d=c(),h=s.reduce((f,{input:g})=>f+g.width*g.height,0);return{columns:a,rows:u,overlap:o,maxOutputWidth:l,maxOutputHeight:p,inputPixelCount:h,inputShapeCount:d.size,tiles:s}}function Ge(n,e,t){return Math.min(Math.max(n,e),t)}function ir(n,e){const t=[...n].sort((i,r)=>i-r);return t[Math.ceil(t.length*e)-1]}class nr{constructor(e,t=!0){m(this,"enabled");m(this,"maxTileSize");m(this,"minTileSize");m(this,"targetTileTimeMs");m(this,"_tileSize");m(this,"_adjustmentStep");m(this,"_smoothedTileTimeMs");m(this,"_completeExecutionsSinceChange",0);const i=typeof t=="object"?t:{};this.enabled=t!==!1,this.maxTileSize=Math.max(D,li(e,D)),this.minTileSize=Ge(xe(i.minTileSize??Jn,D),D,this.maxTileSize),this.targetTileTimeMs=Math.max(1,i.targetTileTimeMs??er),this._adjustmentStep=Math.max(D,xe(i.adjustmentStep??tr,D)),this._tileSize=this.enabled?Ge(xe(i.initialTileSize??Qn,D),this.minTileSize,this.maxTileSize):this.maxTileSize}get tileSize(){return this._tileSize}observe(e){if(!this.enabled||e.length===0)return!1;const t=e.filter(a=>Number.isFinite(a)&&a>=0);if(t.length===0)return!1;const i=t.length>=3?t.slice(1):t,r=ir(i,.75);if(this._smoothedTileTimeMs=this._smoothedTileTimeMs===void 0?r:this._smoothedTileTimeMs*.65+r*.35,this._completeExecutionsSinceChange++,this._completeExecutionsSinceChange<2)return!1;let o=this._tileSize;return this._smoothedTileTimeMs>this.targetTileTimeMs*1.25?o-=this._adjustmentStep:this._smoothedTileTimeMs<this.targetTileTimeMs*.65&&(o+=this._adjustmentStep),o=Ge(xe(o,D),this.minTileSize,this.maxTileSize),o===this._tileSize?!1:(this._tileSize=o,this._smoothedTileTimeMs=void 0,this._completeExecutionsSinceChange=0,!0)}}function I(n,e){return{op:"conv2d",id:n,input:e,weight:`${n}.weight`,bias:`${n}.bias`,activation:"relu",padding:"same"}}function le(n,e){return{op:"maxPool2d",id:n,input:e,size:2,stride:2,padding:"same"}}function pe(n,e){return{op:"upsample2d",id:n,input:e,scale:2,mode:"nearest"}}function de(n,e,t){return{op:"concat",id:n,inputs:[e,t],axis:"channels"}}const di={schemaVersion:1,id:"oidn-unet-small-v1",family:"oidn-unet-small",input:"input",output:"dec_conv0",receptiveField:174,nodes:[I("enc_conv0","input"),I("enc_conv1","enc_conv0"),le("pool1","enc_conv1"),I("enc_conv2","pool1"),le("pool2","enc_conv2"),I("enc_conv3","pool2"),le("pool3","enc_conv3"),I("enc_conv4","pool3"),le("pool4","enc_conv4"),I("enc_conv5a","pool4"),I("enc_conv5b","enc_conv5a"),pe("up4","enc_conv5b"),de("concat4","up4","pool3"),I("dec_conv4a","concat4"),I("dec_conv4b","dec_conv4a"),pe("up3","dec_conv4b"),de("concat3","up3","pool2"),I("dec_conv3a","concat3"),I("dec_conv3b","dec_conv3a"),pe("up2","dec_conv3b"),de("concat2","up2","pool1"),I("dec_conv2a","concat2"),I("dec_conv2b","dec_conv2a"),pe("up1","dec_conv2b"),de("concat1","up1","input"),I("dec_conv1a","concat1"),I("dec_conv1b","dec_conv1a"),I("dec_conv0","dec_conv1b")]},hi={schemaVersion:1,id:"oidn-unet-large-v1",family:"oidn-unet-large",input:"input",output:"dec_conv1c",receptiveField:202,nodes:[I("enc_conv1a","input"),I("enc_conv1b","enc_conv1a"),le("pool1","enc_conv1b"),I("enc_conv2a","pool1"),I("enc_conv2b","enc_conv2a"),le("pool2","enc_conv2b"),I("enc_conv3a","pool2"),I("enc_conv3b","enc_conv3a"),le("pool3","enc_conv3b"),I("enc_conv4a","pool3"),I("enc_conv4b","enc_conv4a"),le("pool4","enc_conv4b"),I("enc_conv5a","pool4"),I("enc_conv5b","enc_conv5a"),pe("up4","enc_conv5b"),de("concat4","up4","pool3"),I("dec_conv4a","concat4"),I("dec_conv4b","dec_conv4a"),pe("up3","dec_conv4b"),de("concat3","up3","pool2"),I("dec_conv3a","concat3"),I("dec_conv3b","dec_conv3a"),pe("up2","dec_conv3b"),de("concat2","up2","pool1"),I("dec_conv2a","concat2"),I("dec_conv2b","dec_conv2a"),pe("up1","dec_conv2b"),de("concat1","up1","input"),I("dec_conv1a","concat1"),I("dec_conv1b","dec_conv1a"),I("dec_conv1c","dec_conv1b")]},rr=[di,hi];function fi(n){const e=new Set;for(const t of n.nodes)t.op==="conv2d"&&(e.add(t.weight),e.add(t.bias));return e}function gi(n){return n.desc.getByteSize()}function mi(n){return[...n].sort().join(", ")}function dt(n,e=rr){const t=e.filter(i=>{const r=fi(i);return[...r].some(o=>!n.has(o))?!1:i.allowAdditionalTensors===!0||[...n.keys()].every(o=>r.has(o))});if(t.length===1)return t[0];throw t.length>1?new Error(`Ambiguous OIDN model topology: ${t.map(i=>i.id).join(", ")}`):new Error(`Unsupported OIDN model topology. TZA tensors: ${mi(n.keys())}`)}function _i(n,e,t){const i=n.get(e);if(!i)throw new Error(`Model ${t} is missing tensor ${e}`);if(i.data.byteLength!==gi(i))throw new Error(`Tensor ${e} has ${i.data.byteLength} bytes, expected ${gi(i)}`);return i}function yi(n,e=dt(n)){if(e.schemaVersion!==1)throw new Error(`Unsupported model descriptor schema ${e.schemaVersion}`);const t=fi(e);if(!e.allowAdditionalTensors){const c=[...n.keys()].filter(d=>!t.has(d));if(c.length>0)throw new Error(`Model ${e.id} has unexpected tensors: ${mi(c)}`)}const i=new Map,r=new Map,o=new Map,a=new Set([e.input]);let u,s;const l=(c,d)=>{const h=i.get(c);if(h===void 0)throw new Error(`Model ${e.id} node ${d} reads unknown or forward value ${c}`);return h};for(const c of e.nodes){if(a.has(c.id))throw new Error(`Model ${e.id} produces duplicate value ${c.id}`);if(c.op==="conv2d"){const d=_i(n,c.weight,e.id),h=_i(n,c.bias,e.id),f=d.desc.dims;if(d.desc.layout!=="oihw"||f.length!==4)throw new Error(`Tensor ${c.weight} must use OIHW layout`);if(f[2]!==3||f[3]!==3)throw new Error(`Tensor ${c.weight} must use a 3x3 kernel`);if(h.desc.layout!=="x"||h.desc.dims.length!==1)throw new Error(`Tensor ${c.bias} must be a one-dimensional bias`);if(h.desc.dims[0]!==f[0])throw new Error(`Tensor ${c.bias} has ${h.desc.dims[0]} channels, expected ${f[0]}`);if(d.desc.dataType!==h.desc.dataType)throw new Error(`Weight and bias dtype differ for ${c.id}`);if(s&&s!==d.desc.dataType)throw new Error(`Mixed tensor dtypes are not supported by model ${e.id}`);s=d.desc.dataType,c.input===e.input&&u===void 0&&(u=f[1],i.set(e.input,u));const g=l(c.input,c.id);if(g!==f[1])throw new Error(`Tensor ${c.weight} expects ${f[1]} input channels, but ${c.input} provides ${g}`);i.set(c.id,f[0]),r.set(c.id,{weight:d,bias:h,inputChannels:f[1],outputChannels:f[0],kernelHeight:f[2],kernelWidth:f[3]}),o.set(c.id,{inputChannels:f[1],outputChannels:f[0]})}else if(c.op==="concat"){if(c.inputs.length<2)throw new Error(`Concat ${c.id} requires at least two inputs`);const d=c.inputs.reduce((h,f)=>h+l(f,c.id),0);i.set(c.id,d)}else i.set(c.id,l(c.input,c.id));a.add(c.id)}if(u===void 0||s===void 0)throw new Error(`Model ${e.id} has no convolution reading its input`);const p=i.get(e.output);if(p===void 0)throw new Error(`Model ${e.id} output ${e.output} is not produced`);if(p!==3)throw new Error(`Model ${e.id} must produce 3 channels, got ${p}`);return{spec:e,inputChannels:u,outputChannels:p,tensorDataType:s,channelsByValue:i,convChannels:o,convTensors:r}}const or="This is not an object",ar="This is not a Float16Array object",wi="This constructor is not a subclass of Float16Array",vi="The constructor property value is not an object",sr="Species constructor didn't return TypedArray object",ur="Derived constructor created TypedArray object which was too small length",Ce="Attempting to access detached ArrayBuffer",ht="Cannot convert undefined or null to object",ft="Cannot mix BigInt and other types, use explicit conversions",xi="@@iterator property is not callable",bi="Reduce of empty array with no initial value",cr="The comparison function must be either a function or undefined",gt="Offset is out of bounds";function M(n){return(e,...t)=>j(n,e,t)}function be(n,e){return M($e(n,e).get)}const{apply:j,construct:Se,defineProperty:$i,get:mt,getOwnPropertyDescriptor:$e,getPrototypeOf:Ee,has:_t,ownKeys:ki,set:Ii,setPrototypeOf:Bi}=Reflect,lr=Proxy,{EPSILON:pr,MAX_SAFE_INTEGER:Ti,isFinite:Pi,isNaN:ke}=Number,{iterator:te,species:dr,toStringTag:yt,for:hr}=Symbol,Ie=Object,{create:Ye,defineProperty:Ae,freeze:fr,is:Ci}=Ie,wt=Ie.prototype,gr=wt.__lookupGetter__?M(wt.__lookupGetter__):(n,e)=>{if(n==null)throw O(ht);let t=Ie(n);do{const i=$e(t,e);if(i!==void 0)return ue(i,"get")?i.get:void 0}while((t=Ee(t))!==null)},ue=Ie.hasOwn||M(wt.hasOwnProperty),Si=Array,Ei=Si.isArray,Xe=Si.prototype,mr=M(Xe.join),_r=M(Xe.push),yr=M(Xe.toLocaleString),vt=Xe[te],wr=M(vt),{abs:vr,trunc:Ai}=Math,Fe=ArrayBuffer,xr=Fe.isView,Mi=Fe.prototype,br=M(Mi.slice),$r=be(Mi,"byteLength"),xt=typeof SharedArrayBuffer<"u"?SharedArrayBuffer:null,kr=xt&&be(xt.prototype,"byteLength"),bt=Ee(Uint8Array),Ir=bt.from,G=bt.prototype,Br=G[te],Tr=M(G.keys),Pr=M(G.values),Cr=M(G.entries),Sr=M(G.set),Oi=M(G.reverse),Er=M(G.fill),Ar=M(G.copyWithin),Ni=M(G.sort),Me=M(G.slice),Mr=M(G.subarray),Y=be(G,"buffer"),ge=be(G,"byteOffset"),T=be(G,"length"),Ui=be(G,yt),Or=Uint8Array,q=Uint16Array,zi=(...n)=>j(Ir,q,n),$t=Uint32Array,Nr=Float32Array,me=Ee([][te]()),Ve=M(me.next),Ur=M(function*(){}().next),zr=Ee(me),O=TypeError,kt=RangeError,Li=WeakSet,Wi=Li.prototype,Lr=M(Wi.add),Wr=M(Wi.has),He=WeakMap,It=He.prototype,je=M(It.get),Dr=M(It.has),Bt=M(It.set),Di=new He,Rr=Ye(null,{next:{value:function(){const e=je(Di,this);return Ve(e)}},[te]:{value:function(){return this}}});function Ke(n){if(n[te]===vt&&me.next===Ve)return n;const e=Ye(Rr);return Bt(Di,e,wr(n)),e}const Ri=new He,Gi=Ye(zr,{next:{value:function(){const e=je(Ri,this);return Ur(e)},writable:!0,configurable:!0}});for(const n of ki(me))n!=="next"&&Ae(Gi,n,$e(me,n));function Yi(n){const e=Ye(Gi);return Bt(Ri,e,n),e}function qe(n){return n!==null&&typeof n=="object"||typeof n=="function"}function Xi(n){return n!==null&&typeof n=="object"}function Ze(n){return Ui(n)!==void 0}function Tt(n){const e=Ui(n);return e==="BigInt64Array"||e==="BigUint64Array"}function Gr(n){try{return Ei(n)?!1:($r(n),!0)}catch{return!1}}function Fi(n){if(xt===null)return!1;try{return kr(n),!0}catch{return!1}}function Yr(n){return Gr(n)||Fi(n)}function Vi(n){return Ei(n)?n[te]===vt&&me.next===Ve:!1}function Xr(n){return Ze(n)?n[te]===Br&&me.next===Ve:!1}function Je(n){if(typeof n!="string")return!1;const e=+n;return n!==e+""||!Pi(e)?!1:e===Ai(e)}const Qe=hr("__Float16Array__");function Fr(n){if(!Xi(n))return!1;const e=Ee(n);if(!Xi(e))return!1;const t=e.constructor;if(t===void 0)return!1;if(!qe(t))throw O(vi);return _t(t,Qe)}const Pt=1/pr;function Vr(n){return n+Pt-Pt}const Hi=6103515625e-14,Hr=65504,ji=.0009765625,Ki=ji*Hi,jr=ji*Pt;function Kr(n){const e=+n;if(!Pi(e)||e===0)return e;const t=e>0?1:-1,i=vr(e);if(i<Hi)return t*Vr(i/Ki)*Ki;const r=(1+jr)*i,o=r-(r-i);return o>Hr||ke(o)?t*(1/0):t*o}const qi=new Fe(4),Zi=new Nr(qi),Ji=new $t(qi),Q=new q(512),ee=new Or(512);for(let n=0;n<256;++n){const e=n-127;e<-24?(Q[n]=0,Q[n|256]=32768,ee[n]=24,ee[n|256]=24):e<-14?(Q[n]=1024>>-e-14,Q[n|256]=1024>>-e-14|32768,ee[n]=-e-1,ee[n|256]=-e-1):e<=15?(Q[n]=e+15<<10,Q[n|256]=e+15<<10|32768,ee[n]=13,ee[n|256]=13):e<128?(Q[n]=31744,Q[n|256]=64512,ee[n]=24,ee[n|256]=24):(Q[n]=31744,Q[n|256]=64512,ee[n]=13,ee[n|256]=13)}function ie(n){Zi[0]=Kr(n);const e=Ji[0],t=e>>23&511;return Q[t]+((e&8388607)>>ee[t])}const Ct=new $t(2048);for(let n=1;n<1024;++n){let e=n<<13,t=0;for(;!(e&8388608);)e<<=1,t-=8388608;e&=-8388609,t+=947912704,Ct[n]=e|t}for(let n=1024;n<2048;++n)Ct[n]=939524096+(n-1024<<13);const Be=new $t(64);for(let n=1;n<31;++n)Be[n]=n<<23;Be[31]=1199570944,Be[32]=2147483648;for(let n=33;n<63;++n)Be[n]=2147483648+(n-32<<23);Be[63]=3347054592;const Qi=new q(64);for(let n=1;n<64;++n)n!==32&&(Qi[n]=1024);function S(n){const e=n>>10;return Ji[0]=Ct[Qi[e]+(n&1023)]+Be[e],Zi[0]}function ce(n){const e=+n;return ke(e)||e===0?0:Ai(e)}function St(n){const e=ce(n);return e<0?0:e<Ti?e:Ti}function et(n,e){if(!qe(n))throw O(or);const t=n.constructor;if(t===void 0)return e;if(!qe(t))throw O(vi);const i=t[dr];return i??e}function Oe(n){if(Fi(n))return!1;try{return br(n,0,0),!1}catch{}return!0}function en(n,e){const t=ke(n),i=ke(e);if(t&&i)return 0;if(t)return 1;if(i||n<e)return-1;if(n>e)return 1;if(n===0&&e===0){const r=Ci(n,0),o=Ci(e,0);if(!r&&o)return-1;if(r&&!o)return 1}return 0}const Et=2,tt=new He;function Te(n){return Dr(tt,n)||!xr(n)&&Fr(n)}function P(n){if(!Te(n))throw O(ar)}function it(n,e){const t=Te(n),i=Ze(n);if(!t&&!i)throw O(sr);if(typeof e=="number"){let r;if(t){const o=$(n);r=T(o)}else r=T(n);if(r<e)throw O(ur)}if(Tt(n))throw O(ft)}function $(n){const e=je(tt,n);if(e!==void 0){const r=Y(e);if(Oe(r))throw O(Ce);return e}const t=n.buffer;if(Oe(t))throw O(Ce);const i=Se(E,[t,n.byteOffset,n.length],n.constructor);return je(tt,i)}function tn(n){const e=T(n),t=[];for(let i=0;i<e;++i)t[i]=S(n[i]);return t}const nn=new Li;for(const n of ki(G)){if(n===yt)continue;const e=$e(G,n);ue(e,"get")&&typeof e.get=="function"&&Lr(nn,e.get)}const qr=fr({get(n,e,t){return Je(e)&&ue(n,e)?S(mt(n,e)):Wr(nn,gr(n,e))?mt(n,e):mt(n,e,t)},set(n,e,t,i){return Je(e)&&ue(n,e)?Ii(n,e,ie(t)):Ii(n,e,t,i)},getOwnPropertyDescriptor(n,e){if(Je(e)&&ue(n,e)){const t=$e(n,e);return t.value=S(t.value),t}return $e(n,e)},defineProperty(n,e,t){return Je(e)&&ue(n,e)&&ue(t,"value")&&(t.value=ie(t.value)),$i(n,e,t)}});class E{constructor(e,t,i){let r;if(Te(e))r=Se(q,[$(e)],new.target);else if(qe(e)&&!Yr(e)){let a,u;if(Ze(e)){a=e,u=T(e);const s=Y(e);if(Oe(s))throw O(Ce);if(Tt(e))throw O(ft);const l=new Fe(u*Et);r=Se(q,[l],new.target)}else{const s=e[te];if(s!=null&&typeof s!="function")throw O(xi);s!=null?Vi(e)?(a=e,u=e.length):(a=[...e],u=a.length):(a=e,u=St(a.length)),r=Se(q,[u],new.target)}for(let s=0;s<u;++s)r[s]=ie(a[s])}else r=Se(q,arguments,new.target);const o=new lr(r,qr);return Bt(tt,o,r),o}static from(e,...t){const i=this;if(!_t(i,Qe))throw O(wi);if(i===E){if(Te(e)&&t.length===0){const p=$(e),c=new q(Y(p),ge(p),T(p));return new E(Y(Me(c)))}if(t.length===0)return new E(Y(zi(e,ie)));const s=t[0],l=t[1];return new E(Y(zi(e,function(p,...c){return ie(j(s,this,[p,...Ke(c)]))},l)))}let r,o;const a=e[te];if(a!=null&&typeof a!="function")throw O(xi);if(a!=null)Vi(e)?(r=e,o=e.length):Xr(e)?(r=e,o=T(e)):(r=[...e],o=r.length);else{if(e==null)throw O(ht);r=Ie(e),o=St(r.length)}const u=new i(o);if(t.length===0)for(let s=0;s<o;++s)u[s]=r[s];else{const s=t[0],l=t[1];for(let p=0;p<o;++p)u[p]=j(s,l,[r[p],p])}return u}static of(...e){const t=this;if(!_t(t,Qe))throw O(wi);const i=e.length;if(t===E){const o=new E(i),a=$(o);for(let u=0;u<i;++u)a[u]=ie(e[u]);return o}const r=new t(i);for(let o=0;o<i;++o)r[o]=e[o];return r}keys(){P(this);const e=$(this);return Tr(e)}values(){P(this);const e=$(this);return Yi(function*(){for(const t of Pr(e))yield S(t)}())}entries(){P(this);const e=$(this);return Yi(function*(){for(const[t,i]of Cr(e))yield[t,S(i)]}())}at(e){P(this);const t=$(this),i=T(t),r=ce(e),o=r>=0?r:i+r;if(!(o<0||o>=i))return S(t[o])}with(e,t){P(this);const i=$(this),r=T(i),o=ce(e),a=o>=0?o:r+o,u=+t;if(a<0||a>=r)throw kt(gt);const s=new q(Y(i),ge(i),T(i)),l=new E(Y(Me(s))),p=$(l);return p[a]=ie(u),l}map(e,...t){P(this);const i=$(this),r=T(i),o=t[0],a=et(i,E);if(a===E){const s=new E(r),l=$(s);for(let p=0;p<r;++p){const c=S(i[p]);l[p]=ie(j(e,o,[c,p,this]))}return s}const u=new a(r);it(u,r);for(let s=0;s<r;++s){const l=S(i[s]);u[s]=j(e,o,[l,s,this])}return u}filter(e,...t){P(this);const i=$(this),r=T(i),o=t[0],a=[];for(let l=0;l<r;++l){const p=S(i[l]);j(e,o,[p,l,this])&&_r(a,p)}const u=et(i,E),s=new u(a);return it(s),s}reduce(e,...t){P(this);const i=$(this),r=T(i);if(r===0&&t.length===0)throw O(bi);let o,a;t.length===0?(o=S(i[0]),a=1):(o=t[0],a=0);for(let u=a;u<r;++u)o=e(o,S(i[u]),u,this);return o}reduceRight(e,...t){P(this);const i=$(this),r=T(i);if(r===0&&t.length===0)throw O(bi);let o,a;t.length===0?(o=S(i[r-1]),a=r-2):(o=t[0],a=r-1);for(let u=a;u>=0;--u)o=e(o,S(i[u]),u,this);return o}forEach(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=0;a<r;++a)j(e,o,[S(i[a]),a,this])}find(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=0;a<r;++a){const u=S(i[a]);if(j(e,o,[u,a,this]))return u}}findIndex(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=0;a<r;++a){const u=S(i[a]);if(j(e,o,[u,a,this]))return a}return-1}findLast(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=r-1;a>=0;--a){const u=S(i[a]);if(j(e,o,[u,a,this]))return u}}findLastIndex(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=r-1;a>=0;--a){const u=S(i[a]);if(j(e,o,[u,a,this]))return a}return-1}every(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=0;a<r;++a)if(!j(e,o,[S(i[a]),a,this]))return!1;return!0}some(e,...t){P(this);const i=$(this),r=T(i),o=t[0];for(let a=0;a<r;++a)if(j(e,o,[S(i[a]),a,this]))return!0;return!1}set(e,...t){P(this);const i=$(this),r=ce(t[0]);if(r<0)throw kt(gt);if(e==null)throw O(ht);if(Tt(e))throw O(ft);if(Te(e))return Sr($(this),$(e),r);if(Ze(e)){const s=Y(e);if(Oe(s))throw O(Ce)}const o=T(i),a=Ie(e),u=St(a.length);if(r===1/0||u+r>o)throw kt(gt);for(let s=0;s<u;++s)i[s+r]=ie(a[s])}reverse(){P(this);const e=$(this);return Oi(e),this}toReversed(){P(this);const e=$(this),t=new q(Y(e),ge(e),T(e)),i=new E(Y(Me(t))),r=$(i);return Oi(r),i}fill(e,...t){P(this);const i=$(this);return Er(i,ie(e),...Ke(t)),this}copyWithin(e,t,...i){P(this);const r=$(this);return Ar(r,e,t,...Ke(i)),this}sort(e){P(this);const t=$(this),i=e!==void 0?e:en;return Ni(t,(r,o)=>i(S(r),S(o))),this}toSorted(e){P(this);const t=$(this);if(e!==void 0&&typeof e!="function")throw new O(cr);const i=e!==void 0?e:en,r=new q(Y(t),ge(t),T(t)),o=new E(Y(Me(r))),a=$(o);return Ni(a,(u,s)=>i(S(u),S(s))),o}slice(e,t){P(this);const i=$(this),r=et(i,E);if(r===E){const f=new q(Y(i),ge(i),T(i));return new E(Y(Me(f,e,t)))}const o=T(i),a=ce(e),u=t===void 0?o:ce(t);let s;a===-1/0?s=0:a<0?s=o+a>0?o+a:0:s=o<a?o:a;let l;u===-1/0?l=0:u<0?l=o+u>0?o+u:0:l=o<u?o:u;const p=l-s>0?l-s:0,c=new r(p);if(it(c,p),p===0)return c;const d=Y(i);if(Oe(d))throw O(Ce);let h=0;for(;s<l;)c[h]=S(i[s]),++s,++h;return c}subarray(e,t){P(this);const i=$(this),r=et(i,E),o=new q(Y(i),ge(i),T(i)),a=Mr(o,e,t),u=new r(Y(a),ge(a),T(a));return it(u),u}indexOf(e,...t){P(this);const i=$(this),r=T(i);let o=ce(t[0]);if(o===1/0)return-1;o<0&&(o+=r,o<0&&(o=0));for(let a=o;a<r;++a)if(ue(i,a)&&S(i[a])===e)return a;return-1}lastIndexOf(e,...t){P(this);const i=$(this),r=T(i);let o=t.length>=1?ce(t[0]):r-1;if(o===-1/0)return-1;o>=0?o=o<r-1?o:r-1:o+=r;for(let a=o;a>=0;--a)if(ue(i,a)&&S(i[a])===e)return a;return-1}includes(e,...t){P(this);const i=$(this),r=T(i);let o=ce(t[0]);if(o===1/0)return!1;o<0&&(o+=r,o<0&&(o=0));const a=ke(e);for(let u=o;u<r;++u){const s=S(i[u]);if(a&&ke(s)||s===e)return!0}return!1}join(e){P(this);const t=$(this),i=tn(t);return mr(i,e)}toLocaleString(...e){P(this);const t=$(this),i=tn(t);return yr(i,...Ke(e))}get[yt](){if(Te(this))return"Float16Array"}}Ae(E,"BYTES_PER_ELEMENT",{value:Et}),Ae(E,Qe,{}),Bi(E,bt);const nt=E.prototype;Ae(nt,"BYTES_PER_ELEMENT",{value:Et}),Ae(nt,te,{value:nt.values,writable:!0,configurable:!0}),Bi(nt,G);const ne=8,_e=ne+2,rn=3*3;function on(n,e,t=!1){const i=n==="fp16"?2:4,r=_e*_e*e,o=t?rn*e*4:0;return(r+o)*4*i}function Zr(n){return n==="fp16"?"vec4<f16>":"vec4<f32>"}function Jr(n){return n==="fp16"?`enable f16;
|
|
105
|
+
`:""}function Qr(n,e,t){const i=r=>`${t}[weightBase + ${r}u]`;return n==="fp16"?`
|
|
106
|
+
let inputValue = vec4<f16>(${e});
|
|
107
|
+
var partial = vec4<f16>(0.0h);
|
|
108
|
+
partial = fma(${i(0)}, vec4<f16>(inputValue.x), partial);
|
|
109
|
+
partial = fma(${i(1)}, vec4<f16>(inputValue.y), partial);
|
|
110
|
+
partial = fma(${i(2)}, vec4<f16>(inputValue.z), partial);
|
|
111
|
+
partial = fma(${i(3)}, vec4<f16>(inputValue.w), partial);
|
|
112
|
+
acc += vec4<f32>(partial);
|
|
113
|
+
`:`
|
|
114
|
+
let inputValue = vec4<f32>(${e});
|
|
115
|
+
acc = fma(vec4<f32>(${i(0)}), vec4<f32>(inputValue.x), acc);
|
|
116
|
+
acc = fma(vec4<f32>(${i(1)}), vec4<f32>(inputValue.y), acc);
|
|
117
|
+
acc = fma(vec4<f32>(${i(2)}), vec4<f32>(inputValue.z), acc);
|
|
118
|
+
acc = fma(vec4<f32>(${i(3)}), vec4<f32>(inputValue.w), acc);
|
|
119
|
+
`}function eo(n,e,t,i=!1){if(!Number.isInteger(t)||t<1)throw new Error(`Final RGB shader requires positive input blocks, got ${t}`);const r=Zr(n),o="vec4<f32>",a=_e*_e*t,u=rn*t*4,s=e==="relu"?"max(acc, vec4<f32>(0.0))":"acc",l=Qr(n,"inputTile[patchBase + inputBlock]",i?"weightTile":"weights");return`${Jr(n)}
|
|
120
|
+
struct Params {
|
|
121
|
+
inputWidth: u32,
|
|
122
|
+
inputHeight: u32,
|
|
123
|
+
outputWidth: u32,
|
|
124
|
+
outputHeight: u32,
|
|
125
|
+
inputBlocks: u32,
|
|
126
|
+
outputBlocks: u32,
|
|
127
|
+
}
|
|
128
|
+
|
|
129
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${r}>;
|
|
130
|
+
@group(0) @binding(1) var<storage, read> weights: array<${r}>;
|
|
131
|
+
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
132
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${o}>;
|
|
133
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
134
|
+
|
|
135
|
+
var<workgroup> inputTile: array<${r}, ${a}>;
|
|
136
|
+
${i?`var<workgroup> weightTile: array<${r}, ${u}>;`:""}
|
|
137
|
+
|
|
138
|
+
@compute @workgroup_size(${ne}, ${ne}, 1)
|
|
139
|
+
fn main(
|
|
140
|
+
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
141
|
+
@builtin(global_invocation_id) gid: vec3<u32>,
|
|
142
|
+
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
143
|
+
) {
|
|
144
|
+
let localLinear = localId.y * ${ne}u + localId.x;
|
|
145
|
+
for (
|
|
146
|
+
var loadIndex = localLinear;
|
|
147
|
+
loadIndex < ${a}u;
|
|
148
|
+
loadIndex += ${ne*ne}u
|
|
149
|
+
) {
|
|
150
|
+
let tilePixel = loadIndex / ${t}u;
|
|
151
|
+
let inputBlock = loadIndex % ${t}u;
|
|
152
|
+
let tileX = tilePixel % ${_e}u;
|
|
153
|
+
let tileY = tilePixel / ${_e}u;
|
|
154
|
+
let inputX = i32(workgroupId.x * ${ne}u + tileX) - 1;
|
|
155
|
+
let inputY = i32(workgroupId.y * ${ne}u + tileY) - 1;
|
|
156
|
+
var value = ${r}(0.0);
|
|
157
|
+
if (
|
|
158
|
+
inputX >= 0 && inputX < i32(params.inputWidth) &&
|
|
159
|
+
inputY >= 0 && inputY < i32(params.inputHeight)
|
|
160
|
+
) {
|
|
161
|
+
let inputIndex =
|
|
162
|
+
(u32(inputY) * params.inputWidth + u32(inputX)) *
|
|
163
|
+
${t}u + inputBlock;
|
|
164
|
+
value = inputData[inputIndex];
|
|
165
|
+
}
|
|
166
|
+
inputTile[loadIndex] = value;
|
|
167
|
+
}
|
|
168
|
+
|
|
169
|
+
${i?` for (
|
|
170
|
+
var loadIndex = localLinear;
|
|
171
|
+
loadIndex < ${u}u;
|
|
172
|
+
loadIndex += ${ne*ne}u
|
|
173
|
+
) {
|
|
174
|
+
weightTile[loadIndex] = weights[loadIndex];
|
|
175
|
+
}
|
|
176
|
+
|
|
177
|
+
`:""} // Out-of-range invocations must reach this barrier before returning.
|
|
178
|
+
workgroupBarrier();
|
|
179
|
+
|
|
180
|
+
let outputInBounds =
|
|
181
|
+
gid.x < params.outputWidth &&
|
|
182
|
+
gid.y < params.outputHeight &&
|
|
183
|
+
gid.z < params.outputBlocks;
|
|
184
|
+
if (!outputInBounds) {
|
|
185
|
+
return;
|
|
186
|
+
}
|
|
187
|
+
|
|
188
|
+
var acc = bias[gid.z];
|
|
189
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
190
|
+
let inputY = i32(gid.y) + i32(ky) - 1;
|
|
191
|
+
if (inputY < 0 || inputY >= i32(params.inputHeight)) {
|
|
192
|
+
continue;
|
|
193
|
+
}
|
|
194
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
195
|
+
let inputX = i32(gid.x) + i32(kx) - 1;
|
|
196
|
+
if (inputX < 0 || inputX >= i32(params.inputWidth)) {
|
|
197
|
+
continue;
|
|
198
|
+
}
|
|
199
|
+
let patchBase =
|
|
200
|
+
((localId.y + ky) * ${_e}u + localId.x + kx) *
|
|
201
|
+
${t}u;
|
|
202
|
+
for (var inputBlock = 0u; inputBlock < ${t}u; inputBlock++) {
|
|
203
|
+
let weightBase =
|
|
204
|
+
((ky * 3u + kx) * ${t}u + inputBlock) * 4u;
|
|
205
|
+
${l}
|
|
206
|
+
}
|
|
207
|
+
}
|
|
208
|
+
}
|
|
209
|
+
|
|
210
|
+
let outputIndex =
|
|
211
|
+
(gid.y * params.outputWidth + gid.x) * 1u + gid.z;
|
|
212
|
+
outputData[outputIndex] = ${s};
|
|
213
|
+
}
|
|
214
|
+
`}function to(n){return n.op==="concat"?n.inputs:[n.input]}function io(n){const e=new Map;for(const t of n.nodes)for(const i of to(t)){const r=e.get(i)??[];r.push(t),e.set(i,r)}return e}function At(n,e,t){const i=n.get(e);if(!((i==null?void 0:i.length)!==1||i[0].op!==t))return i[0]}function rt(n,e={}){const t=n.spec,i=io(t),r=new Map(t.nodes.map(p=>[p.id,p])),o=new Set,a=new Map;let u=0,s=0;if(e.fuseConvPool!==!1)for(const p of t.nodes){if(p.op!=="conv2d"||p.activation!=="relu")continue;const c=At(i,p.id,"maxPool2d");!c||c.size!==2||c.stride!==2||(o.add(p.id),a.set(c.id,{op:"fusedConvReluMaxPool2d",id:c.id,input:p.input,conv:p,pool:c}),u++)}if(e.fuseUpsampleConcatConv!==!1)for(const p of t.nodes){if(p.op!=="conv2d")continue;const c=r.get(p.input);if((c==null?void 0:c.op)!=="concat"||c.inputs.length!==2||At(i,c.id,"conv2d")!==p)continue;const d=n.channelsByValue.get(c.inputs[0]);if(d===void 0||d%4!==0)continue;const h=c.inputs.map(g=>{const y=r.get(g);return(y==null?void 0:y.op)==="upsample2d"&&y.scale===2&&y.mode==="nearest"&&At(i,y.id,"concat")===c?{value:y.input,upsample:y}:{value:g}});if(h.filter(g=>g.upsample).length===1){o.add(c.id);for(const g of c.inputs){const y=r.get(g);(y==null?void 0:y.op)==="upsample2d"&&o.add(y.id)}a.set(p.id,{op:"fusedUpsampleConcatConv2d",id:p.id,inputs:h,conv:p}),s++}}const l=[];for(const p of t.nodes){const c=a.get(p.id);c?l.push(c):o.has(p.id)||l.push(p)}return{spec:t,nodes:l,fusions:{convPool:u,upsampleConcatConv:s}}}function no(n){return n.op==="concat"?n.inputs:n.op==="fusedUpsampleConcatConv2d"?n.inputs.map(e=>e.value):[n.input]}function an(n,e){return n.width===e.width&&n.height===e.height}function sn(n,e,t,i){if(!Number.isInteger(e)||e<=0||!Number.isInteger(t)||t<=0)throw new Error(`Invalid model input size ${e}x${t}`);const r=rt(n,i),o={width:e,height:t,channels:n.inputChannels},a=new Map([[n.spec.input,o]]),u=[],s=(c,d)=>{const h=a.get(c);if(!h)throw new Error(`Planned node ${d} reads missing value ${c}`);return h};for(const c of r.nodes){let d;if(c.op==="conv2d"){const h=s(c.input,c.id);d={width:h.width,height:h.height,channels:n.convChannels.get(c.id).outputChannels}}else if(c.op==="maxPool2d"){const h=s(c.input,c.id);d={width:Math.ceil(h.width/2),height:Math.ceil(h.height/2),channels:h.channels}}else if(c.op==="upsample2d"){const h=s(c.input,c.id);d={width:h.width*2,height:h.height*2,channels:h.channels}}else if(c.op==="concat"){const h=c.inputs.map(f=>s(f,c.id));if(h.some(f=>!an(f,h[0])))throw new Error(`Concat ${c.id} has mismatched spatial shapes`);d={width:h[0].width,height:h[0].height,channels:h.reduce((f,g)=>f+g.channels,0)}}else if(c.op==="fusedConvReluMaxPool2d"){const h=s(c.input,c.id);d={width:Math.ceil(h.width/2),height:Math.ceil(h.height/2),channels:n.convChannels.get(c.conv.id).outputChannels}}else{const h=c.inputs.map(f=>{const g=s(f.value,c.id);return f.upsample?{...g,width:g.width*2,height:g.height*2}:g});if(h.some(f=>!an(f,h[0])))throw new Error(`Fused decoder ${c.id} has mismatched spatial shapes`);d={width:h[0].width,height:h[0].height,channels:n.convChannels.get(c.conv.id).outputChannels}}a.set(c.id,d),u.push(d)}const l=new Map;r.nodes.forEach((c,d)=>{for(const h of no(c))l.set(h,d)}),l.set(n.spec.output,r.nodes.length);const p=r.nodes.map((c,d)=>({node:c,outputShape:u[d],lastUse:l.get(c.id)??d}));return{...r,inputShape:o,valueShapes:a,plannedNodes:p}}class un{constructor(){m(this,"_stats",new Map)}track(e,t){let i=this._stats.get(e);return i||(i={created:0,destroyed:0,live:0,peakLive:0,resources:new Set},this._stats.set(e,i)),i.resources.has(t)||(i.resources.add(t),i.created++,i.live++,i.peakLive=Math.max(i.peakLive,i.live)),t}release(e,t,i){if(!t)return!1;const r=this._stats.get(e);if(!(r!=null&&r.resources.delete(t)))return!1;try{i()}finally{r.destroyed++,r.live--}return!0}snapshot(e=0){let t=0,i=0,r=0,o=0;const a={};for(const[u,s]of this._stats){const l={created:s.created,destroyed:s.destroyed,live:s.live,peakLive:s.peakLive};a[u]=l,t+=l.live,i+=l.created,r+=l.destroyed,o+=l.peakLive}return{live:t,created:i,destroyed:r,peakLive:o,pending:e,byKind:a}}}const cn=new WeakMap;function ro(n){let e=cn.get(n);return e||(e={ready:new Map,pending:new Map},cn.set(n,e)),e}const R=8,ln=8,U=8,J=U+2;function re(n){const[e,t]=n.workgroupSize;return{workgroupX:e,workgroupY:t,rowsPerThread:n.rowsPerThread,tileM:t*n.rowsPerThread,tileNBlocks:e,tileKBlocks:ln}}function ye(n){const e=re(n),t=n.sharedLayout==="padded"||n.sharedLayout==="padded-input",i=n.sharedLayout==="padded"||n.sharedLayout==="padded-weights";return{input:e.tileM*e.tileKBlocks+(t?e.workgroupY:0),weights:e.tileKBlocks*e.tileNBlocks*(i?5:4)}}function ot(n,e){return e.sharedLayout==="padded"||e.sharedLayout==="padded-input"?`(${n}) + (${n}) / ${e.rowsPerThread*ln}u`:n}function pn(n,e){return e.sharedLayout==="padded"||e.sharedLayout==="padded-weights"?`(${n}) + (${n}) / 4u`:n}function Mt(n,e){return Math.ceil(n/e)*e}function N(n){return Math.ceil(n/4)}function oo(n,e){return n.width*n.height*N(n.channels)*4*e}let Ne;function dn(n){const e=n&32768?-1:1,t=n>>>10&31,i=n&1023;return t===0?e*i*2**-24:t===31?i===0?e*(1/0):NaN:e*(1+i/1024)*2**(t-15)}function ao(){if(!Ne){Ne=new Float32Array(65536);for(let n=0;n<Ne.length;n++)Ne[n]=dn(n)}return Ne}function hn(n){if(n.desc.dataType==="Float32")return new Float32Array(n.data.buffer,n.data.byteOffset,n.data.byteLength/4);const e=new Uint16Array(n.data.buffer,n.data.byteOffset,n.data.byteLength/2),t=new Float32Array(e.length);if(e.length<4096)for(let i=0;i<e.length;i++)t[i]=dn(e[i]);else{const i=ao();for(let r=0;r<e.length;r++)t[r]=i[e[r]]}return t}function Ot(n,e,t,i){const r=Mt(t.byteLength,4),o=n.createBuffer({label:e,size:r,usage:i,mappedAtCreation:!0});return new Uint8Array(o.getMappedRange()).set(new Uint8Array(t.buffer,t.byteOffset,t.byteLength)),o.unmap(),o}function fn(n,e,t){const i=new Uint32Array(Mt(t.length,4));return i.set(t),Ot(n,e,i,GPUBufferUsage.UNIFORM)}function so(n,e,t,i,r){const o=N(t.inputChannels),a=N(t.outputChannels),u=a*t.kernelHeight*t.kernelWidth*o*4*4,s=i==="fp16"&&t.weight.desc.dataType==="Float16",l=s?new Uint16Array(u):i==="fp16"?new E(u):new Float32Array(u),p=s?new Uint16Array(t.weight.data.buffer,t.weight.data.byteOffset,t.weight.data.byteLength/2):hn(t.weight);for(let h=0;h<a;h++)for(let f=0;f<t.kernelHeight;f++)for(let g=0;g<t.kernelWidth;g++)for(let y=0;y<o;y++)for(let v=0;v<4;v++){const _=h*4+v;for(let w=0;w<4;w++){const b=y*4+w,B=(r==="k-major"?(((f*t.kernelWidth+g)*o+y)*a+h)*16:(((h*t.kernelHeight+f)*t.kernelWidth+g)*o+y)*16)+w*4+v;if(_<t.outputChannels&&b<t.inputChannels){const k=((_*t.inputChannels+b)*t.kernelHeight+f)*t.kernelWidth+g;l[B]=p[k]}}}const c=new Float32Array(a*4);c.set(hn(t.bias));const d=Ot(n,`oidn/${e}/weights/${i}/${r}`,l,GPUBufferUsage.STORAGE);try{return{weights:d,weightLayout:r,bias:Ot(n,`oidn/${e}/bias`,c,GPUBufferUsage.STORAGE)}}catch(h){throw d.destroy(),h}}function W(n){return n==="fp16"?"vec4<f16>":"vec4<f32>"}function oe(n){return n==="fp16"?`enable f16;
|
|
215
|
+
`:""}function ae(n,e){return e==="fp16"?`vec4<f16>(${n})`:n}function he(n,e){return e==="relu"?`max(${n}, vec4<f32>(0.0))`:n}function Ue(n,e,t){return t==="fp32"?`
|
|
216
|
+
let inputValue = vec4<f32>(${n});
|
|
106
217
|
let weightBase = ${e};
|
|
107
218
|
acc = fma(vec4<f32>(weights[weightBase]), vec4<f32>(inputValue.x), acc);
|
|
108
219
|
acc = fma(vec4<f32>(weights[weightBase + 1u]), vec4<f32>(inputValue.y), acc);
|
|
109
220
|
acc = fma(vec4<f32>(weights[weightBase + 2u]), vec4<f32>(inputValue.z), acc);
|
|
110
221
|
acc = fma(vec4<f32>(weights[weightBase + 3u]), vec4<f32>(inputValue.w), acc);
|
|
111
222
|
`:`
|
|
112
|
-
let inputValue = vec4<f16>(${
|
|
223
|
+
let inputValue = vec4<f16>(${n});
|
|
113
224
|
let weightBase = ${e};
|
|
114
225
|
var partial = vec4<f16>(0.0h);
|
|
115
226
|
partial = fma(weights[weightBase], vec4<f16>(inputValue.x), partial);
|
|
@@ -117,31 +228,45 @@ partial = fma(weights[weightBase + 1u], vec4<f16>(inputValue.y), partial);
|
|
|
117
228
|
partial = fma(weights[weightBase + 2u], vec4<f16>(inputValue.z), partial);
|
|
118
229
|
partial = fma(weights[weightBase + 3u], vec4<f16>(inputValue.w), partial);
|
|
119
230
|
acc += vec4<f32>(partial);
|
|
120
|
-
`}function
|
|
121
|
-
partial = fma(subgroupBroadcast(weights[weightBase], 0u), ${
|
|
122
|
-
partial = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), ${
|
|
123
|
-
partial = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), ${
|
|
124
|
-
partial = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), ${
|
|
231
|
+
`}function uo(n,e,t){const i=t==="fp16"?"vec4<f16>":"vec4<f32>",r=t==="fp16"?`var partial = vec4<f16>(0.0h);
|
|
232
|
+
partial = fma(subgroupBroadcast(weights[weightBase], 0u), ${i}(inputValue.x), partial);
|
|
233
|
+
partial = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), ${i}(inputValue.y), partial);
|
|
234
|
+
partial = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), ${i}(inputValue.z), partial);
|
|
235
|
+
partial = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), ${i}(inputValue.w), partial);
|
|
125
236
|
acc += vec4<f32>(partial);`:`acc = fma(subgroupBroadcast(weights[weightBase], 0u), vec4<f32>(inputValue.x), acc);
|
|
126
237
|
acc = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), vec4<f32>(inputValue.y), acc);
|
|
127
238
|
acc = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), vec4<f32>(inputValue.z), acc);
|
|
128
239
|
acc = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), vec4<f32>(inputValue.w), acc);`;return`
|
|
129
|
-
let inputValue = ${
|
|
240
|
+
let inputValue = ${i}(${n});
|
|
130
241
|
let weightBase = ${e};
|
|
131
|
-
${
|
|
132
|
-
`}function
|
|
133
|
-
var
|
|
134
|
-
|
|
242
|
+
${r}
|
|
243
|
+
`}function at(n,e,t){return n==="fp16"&&(e.loadMode==="packed-all"||t&&e.loadMode==="packed-weights")}function ze(n,e){return e?`bitcast<vec4<f16>>(${n})`:n}function gn(n,e){const{rowsPerThread:t,tileNBlocks:i,tileKBlocks:r}=re(e),o=ot(`tileSpatial * ${r}u + tileK`,e),a=`(tileK * ${i}u + localId.x) * ${e.sharedLayout==="padded"||e.sharedLayout==="padded-weights"?5:4}u`;if(e.accumulationOrder==="row-major"){const u=W(n),s=n==="fp16"?"partial":"acc[row]";return`
|
|
244
|
+
for (var row = 0u; row < ${t}u; row++) {
|
|
245
|
+
${n==="fp16"?"var partial = vec4<f16>(0.0h);":""}
|
|
246
|
+
let tileSpatial = localId.y * ${t}u + row;
|
|
247
|
+
for (var tileK = 0u; tileK < ${r}u; tileK++) {
|
|
248
|
+
let weightBase = ${a};
|
|
249
|
+
let inputValue = inputTile[${o}];
|
|
250
|
+
${s} = fma(weightTile[weightBase], ${u}(inputValue.x), ${s});
|
|
251
|
+
${s} = fma(weightTile[weightBase + 1u], ${u}(inputValue.y), ${s});
|
|
252
|
+
${s} = fma(weightTile[weightBase + 2u], ${u}(inputValue.z), ${s});
|
|
253
|
+
${s} = fma(weightTile[weightBase + 3u], ${u}(inputValue.w), ${s});
|
|
254
|
+
}
|
|
255
|
+
${n==="fp16"?"acc[row] += vec4<f32>(partial);":""}
|
|
256
|
+
}
|
|
257
|
+
`}return n==="fp16"?`
|
|
258
|
+
var partial: array<vec4<f16>, ${t}>;
|
|
259
|
+
for (var row = 0u; row < ${t}u; row++) {
|
|
135
260
|
partial[row] = vec4<f16>(0.0h);
|
|
136
261
|
}
|
|
137
|
-
for (var tileK = 0u; tileK < ${
|
|
262
|
+
for (var tileK = 0u; tileK < ${r}u; tileK++) {
|
|
138
263
|
let weightBase =
|
|
139
|
-
|
|
140
|
-
for (var row = 0u; row < ${
|
|
264
|
+
${a};
|
|
265
|
+
for (var row = 0u; row < ${t}u; row++) {
|
|
141
266
|
let tileSpatial =
|
|
142
|
-
localId.y * ${
|
|
267
|
+
localId.y * ${t}u + row;
|
|
143
268
|
let inputValue =
|
|
144
|
-
inputTile[
|
|
269
|
+
inputTile[${o}];
|
|
145
270
|
partial[row] = fma(
|
|
146
271
|
weightTile[weightBase],
|
|
147
272
|
vec4<f16>(inputValue.x),
|
|
@@ -164,18 +289,18 @@ ${o}
|
|
|
164
289
|
);
|
|
165
290
|
}
|
|
166
291
|
}
|
|
167
|
-
for (var row = 0u; row < ${
|
|
292
|
+
for (var row = 0u; row < ${t}u; row++) {
|
|
168
293
|
acc[row] += vec4<f32>(partial[row]);
|
|
169
294
|
}
|
|
170
295
|
`:`
|
|
171
|
-
for (var tileK = 0u; tileK < ${
|
|
296
|
+
for (var tileK = 0u; tileK < ${r}u; tileK++) {
|
|
172
297
|
let weightBase =
|
|
173
|
-
|
|
174
|
-
for (var row = 0u; row < ${
|
|
298
|
+
${a};
|
|
299
|
+
for (var row = 0u; row < ${t}u; row++) {
|
|
175
300
|
let tileSpatial =
|
|
176
|
-
localId.y * ${
|
|
301
|
+
localId.y * ${t}u + row;
|
|
177
302
|
let inputValue =
|
|
178
|
-
inputTile[
|
|
303
|
+
inputTile[${o}];
|
|
179
304
|
acc[row] = fma(
|
|
180
305
|
weightTile[weightBase],
|
|
181
306
|
vec4<f32>(inputValue.x),
|
|
@@ -198,7 +323,7 @@ ${o}
|
|
|
198
323
|
);
|
|
199
324
|
}
|
|
200
325
|
}
|
|
201
|
-
`}function
|
|
326
|
+
`}function co(n,e,t,i,r){const o=W(n),a=W(n),u=W(e),s=ae(he("acc",t),e);return`${oe(n)}
|
|
202
327
|
struct Params {
|
|
203
328
|
inputWidth: u32,
|
|
204
329
|
inputHeight: u32,
|
|
@@ -208,15 +333,15 @@ struct Params {
|
|
|
208
333
|
outputBlocks: u32,
|
|
209
334
|
}
|
|
210
335
|
|
|
211
|
-
@group(0) @binding(0) var<storage, read> inputData: array<${
|
|
212
|
-
@group(0) @binding(1) var<storage, read> weights: array<${
|
|
336
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${o}>;
|
|
337
|
+
@group(0) @binding(1) var<storage, read> weights: array<${a}>;
|
|
213
338
|
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
214
|
-
@group(0) @binding(3) var<storage, read_write> outputData: array<${
|
|
339
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${u}>;
|
|
215
340
|
@group(0) @binding(4) var<uniform> params: Params;
|
|
216
341
|
|
|
217
|
-
@compute @workgroup_size(${
|
|
342
|
+
@compute @workgroup_size(${R}, ${R}, 1)
|
|
218
343
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
219
|
-
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${
|
|
344
|
+
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${r}u) {
|
|
220
345
|
return;
|
|
221
346
|
}
|
|
222
347
|
var acc = bias[gid.z];
|
|
@@ -226,16 +351,16 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
226
351
|
for (var kx = 0u; kx < 3u; kx++) {
|
|
227
352
|
let inputX = i32(gid.x) + i32(kx) - 1;
|
|
228
353
|
if (inputX < 0 || inputX >= i32(params.inputWidth)) { continue; }
|
|
229
|
-
let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${
|
|
230
|
-
for (var inputBlock = 0u; inputBlock < ${
|
|
231
|
-
${
|
|
354
|
+
let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${i}u;
|
|
355
|
+
for (var inputBlock = 0u; inputBlock < ${i}u; inputBlock++) {
|
|
356
|
+
${Ue("inputData[pixelBase + inputBlock]",`((((gid.z * 3u + ky) * 3u + kx) * ${i}u + inputBlock) * 4u)`,n)}
|
|
232
357
|
}
|
|
233
358
|
}
|
|
234
359
|
}
|
|
235
|
-
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${
|
|
236
|
-
outputData[outputIndex] = ${
|
|
360
|
+
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${r}u + gid.z;
|
|
361
|
+
outputData[outputIndex] = ${s};
|
|
237
362
|
}
|
|
238
|
-
`}function
|
|
363
|
+
`}function lo(n,e,t,i,r){const o=W(n),a=W(n),u=W(e),s=ae(he("acc",t),e);return`${oe(n)}
|
|
239
364
|
enable subgroups;
|
|
240
365
|
struct Params {
|
|
241
366
|
inputWidth: u32,
|
|
@@ -245,13 +370,13 @@ struct Params {
|
|
|
245
370
|
inputBlocks: u32,
|
|
246
371
|
outputBlocks: u32,
|
|
247
372
|
}
|
|
248
|
-
@group(0) @binding(0) var<storage, read> inputData: array<${
|
|
249
|
-
@group(0) @binding(1) var<storage, read> weights: array<${
|
|
373
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${o}>;
|
|
374
|
+
@group(0) @binding(1) var<storage, read> weights: array<${a}>;
|
|
250
375
|
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
251
|
-
@group(0) @binding(3) var<storage, read_write> outputData: array<${
|
|
376
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${u}>;
|
|
252
377
|
@group(0) @binding(4) var<uniform> params: Params;
|
|
253
378
|
|
|
254
|
-
@compute @workgroup_size(${
|
|
379
|
+
@compute @workgroup_size(${R}, ${R}, 1)
|
|
255
380
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
256
381
|
let outputInBounds =
|
|
257
382
|
gid.x < params.outputWidth && gid.y < params.outputHeight;
|
|
@@ -266,19 +391,19 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
266
391
|
outputInBounds && inputX >= 0 && inputY >= 0 &&
|
|
267
392
|
inputX < i32(params.inputWidth) && inputY < i32(params.inputHeight);
|
|
268
393
|
let pixelBase =
|
|
269
|
-
(clampedY * params.inputWidth + clampedX) * ${
|
|
270
|
-
for (var inputBlock = 0u; inputBlock < ${
|
|
271
|
-
${
|
|
394
|
+
(clampedY * params.inputWidth + clampedX) * ${i}u;
|
|
395
|
+
for (var inputBlock = 0u; inputBlock < ${i}u; inputBlock++) {
|
|
396
|
+
${uo(`select(${o}(0.0), inputData[pixelBase + inputBlock], inputInBounds)`,`((((gid.z * 3u + ky) * 3u + kx) * ${i}u + inputBlock) * 4u)`,n)}
|
|
272
397
|
}
|
|
273
398
|
}
|
|
274
399
|
}
|
|
275
400
|
if (outputInBounds) {
|
|
276
401
|
let outputIndex =
|
|
277
|
-
(gid.y * params.outputWidth + gid.x) * ${
|
|
278
|
-
outputData[outputIndex] = ${
|
|
402
|
+
(gid.y * params.outputWidth + gid.x) * ${r}u + gid.z;
|
|
403
|
+
outputData[outputIndex] = ${s};
|
|
279
404
|
}
|
|
280
405
|
}
|
|
281
|
-
`}function
|
|
406
|
+
`}function po(n,e,t,i,r){const o=W(n),a=W(e),u=ae(he("acc",t),e),s=J*J*i,l=U*U;return`${oe(n)}
|
|
282
407
|
struct Params {
|
|
283
408
|
inputWidth: u32,
|
|
284
409
|
inputHeight: u32,
|
|
@@ -287,42 +412,42 @@ struct Params {
|
|
|
287
412
|
inputBlocks: u32,
|
|
288
413
|
outputBlocks: u32,
|
|
289
414
|
}
|
|
290
|
-
@group(0) @binding(0) var<storage, read> inputData: array<${
|
|
291
|
-
@group(0) @binding(1) var<storage, read> weights: array<${
|
|
415
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${o}>;
|
|
416
|
+
@group(0) @binding(1) var<storage, read> weights: array<${o}>;
|
|
292
417
|
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
293
|
-
@group(0) @binding(3) var<storage, read_write> outputData: array<${
|
|
418
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${a}>;
|
|
294
419
|
@group(0) @binding(4) var<uniform> params: Params;
|
|
295
420
|
|
|
296
|
-
var<workgroup> inputPatch: array<${
|
|
421
|
+
var<workgroup> inputPatch: array<${o}, ${s}>;
|
|
297
422
|
|
|
298
|
-
@compute @workgroup_size(${
|
|
423
|
+
@compute @workgroup_size(${U}, ${U}, 1)
|
|
299
424
|
fn main(
|
|
300
425
|
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
301
426
|
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
302
427
|
) {
|
|
303
428
|
let localLinear =
|
|
304
|
-
localId.y * ${
|
|
429
|
+
localId.y * ${U}u + localId.x;
|
|
305
430
|
for (
|
|
306
431
|
var loadIndex = localLinear;
|
|
307
|
-
loadIndex < ${
|
|
432
|
+
loadIndex < ${s}u;
|
|
308
433
|
loadIndex += ${l}u
|
|
309
434
|
) {
|
|
310
|
-
let patchPixel = loadIndex / ${
|
|
311
|
-
let inputBlock = loadIndex % ${
|
|
312
|
-
let patchX = patchPixel % ${
|
|
313
|
-
let patchY = patchPixel / ${
|
|
435
|
+
let patchPixel = loadIndex / ${i}u;
|
|
436
|
+
let inputBlock = loadIndex % ${i}u;
|
|
437
|
+
let patchX = patchPixel % ${J}u;
|
|
438
|
+
let patchY = patchPixel / ${J}u;
|
|
314
439
|
let inputX =
|
|
315
|
-
i32(workgroupId.x * ${
|
|
440
|
+
i32(workgroupId.x * ${U}u + patchX) - 1;
|
|
316
441
|
let inputY =
|
|
317
|
-
i32(workgroupId.y * ${
|
|
318
|
-
var value = ${
|
|
442
|
+
i32(workgroupId.y * ${U}u + patchY) - 1;
|
|
443
|
+
var value = ${o}(0.0);
|
|
319
444
|
if (
|
|
320
445
|
inputX >= 0 && inputX < i32(params.inputWidth) &&
|
|
321
446
|
inputY >= 0 && inputY < i32(params.inputHeight)
|
|
322
447
|
) {
|
|
323
448
|
let inputIndex =
|
|
324
449
|
(u32(inputY) * params.inputWidth + u32(inputX)) *
|
|
325
|
-
${
|
|
450
|
+
${i}u + inputBlock;
|
|
326
451
|
value = inputData[inputIndex];
|
|
327
452
|
}
|
|
328
453
|
inputPatch[loadIndex] = value;
|
|
@@ -330,13 +455,13 @@ fn main(
|
|
|
330
455
|
workgroupBarrier();
|
|
331
456
|
|
|
332
457
|
let outputX =
|
|
333
|
-
workgroupId.x * ${
|
|
458
|
+
workgroupId.x * ${U}u + localId.x;
|
|
334
459
|
let outputY =
|
|
335
|
-
workgroupId.y * ${
|
|
460
|
+
workgroupId.y * ${U}u + localId.y;
|
|
336
461
|
let outputBlock = workgroupId.z;
|
|
337
462
|
if (
|
|
338
463
|
outputX >= params.outputWidth || outputY >= params.outputHeight ||
|
|
339
|
-
outputBlock >= ${
|
|
464
|
+
outputBlock >= ${r}u
|
|
340
465
|
) {
|
|
341
466
|
return;
|
|
342
467
|
}
|
|
@@ -345,18 +470,62 @@ fn main(
|
|
|
345
470
|
for (var ky = 0u; ky < 3u; ky++) {
|
|
346
471
|
for (var kx = 0u; kx < 3u; kx++) {
|
|
347
472
|
let patchBase =
|
|
348
|
-
((localId.y + ky) * ${
|
|
349
|
-
${
|
|
350
|
-
for (var inputBlock = 0u; inputBlock < ${
|
|
351
|
-
${
|
|
473
|
+
((localId.y + ky) * ${J}u + localId.x + kx) *
|
|
474
|
+
${i}u;
|
|
475
|
+
for (var inputBlock = 0u; inputBlock < ${i}u; inputBlock++) {
|
|
476
|
+
${Ue("inputPatch[patchBase + inputBlock]",`((((outputBlock * 3u + ky) * 3u + kx) * ${i}u + inputBlock) * 4u)`,n)}
|
|
352
477
|
}
|
|
353
478
|
}
|
|
354
479
|
}
|
|
355
480
|
let outputIndex =
|
|
356
|
-
(outputY * params.outputWidth + outputX) * ${
|
|
357
|
-
outputData[outputIndex] = ${
|
|
481
|
+
(outputY * params.outputWidth + outputX) * ${r}u + outputBlock;
|
|
482
|
+
outputData[outputIndex] = ${u};
|
|
358
483
|
}
|
|
359
|
-
`}function
|
|
484
|
+
`}function mn(n,e,t,i){const{workgroupX:r,workgroupY:o,tileM:a,tileKBlocks:u}=re(t),s=e?"params.outputWidth":"params.inputWidth",l=e?"params.outputHeight":"params.inputHeight",p=t.addressMode==="base-offset",c=i?i.blocks.map((h,f)=>`cachedBase${f}[load] = (${f===i.upsampled?"y / 2u":"y"} * params.source${f}Width + ${f===i.upsampled?"x / 2u":"x"}) * ${h}u;`).join(`
|
|
485
|
+
`):`cachedBase0[load] = (y * params.inputWidth + x) * ${n}u;`,d=a*u/(r*o);return`
|
|
486
|
+
${p?`var cachedBase0: array<u32, ${d}>;
|
|
487
|
+
${i?`var cachedBase1: array<u32, ${d}>;`:""}`:`var cachedX: array<u32, ${d}>;
|
|
488
|
+
var cachedY: array<u32, ${d}>;`}
|
|
489
|
+
var cachedMask: array<u32, ${d}>;
|
|
490
|
+
for (var load = 0u; load < ${d}u; load++) {
|
|
491
|
+
let loadIndex = localLinear + load * ${r*o}u;
|
|
492
|
+
let spatial = workgroupId.x * ${a}u + loadIndex / ${u}u;
|
|
493
|
+
let x = spatial % params.outputWidth;
|
|
494
|
+
let y = spatial / params.outputWidth;
|
|
495
|
+
${p?c:`cachedX[load] = x;
|
|
496
|
+
cachedY[load] = y;`}
|
|
497
|
+
let columns =
|
|
498
|
+
select(0u, 0x049u, x > 0u && x - 1u < ${s}) |
|
|
499
|
+
select(0u, 0x092u, x < ${s}) |
|
|
500
|
+
select(0u, 0x124u, x + 1u < ${s});
|
|
501
|
+
let rows =
|
|
502
|
+
select(0u, 0x007u, y > 0u && y - 1u < ${l}) |
|
|
503
|
+
select(0u, 0x038u, y < ${l}) |
|
|
504
|
+
select(0u, 0x1c0u, y + 1u < ${l});
|
|
505
|
+
cachedMask[load] = select(0u, columns & rows, spatial < spatialCount)${p&&i?" | ((x & 1u) << 9u) | ((y & 1u) << 10u)":""};
|
|
506
|
+
}
|
|
507
|
+
var channelBlock = (localLinear % ${u}u) % ${n}u;
|
|
508
|
+
var filterPosition = (localLinear % ${u}u) / ${n}u;
|
|
509
|
+
`}function _n(n,e){const{tileKBlocks:t}=re(e);return`
|
|
510
|
+
channelBlock += ${t%n}u;
|
|
511
|
+
let carry = channelBlock >= ${n}u;
|
|
512
|
+
channelBlock -= select(0u, ${n}u, carry);
|
|
513
|
+
filterPosition += ${Math.floor(t/n)}u + select(0u, 1u, carry);
|
|
514
|
+
`}function st(n,e,t){const{workgroupX:i,workgroupY:r,tileM:o,tileKBlocks:a}=re(t);return`
|
|
515
|
+
let inputBlock = channelBlock;
|
|
516
|
+
let kernelY = filterPosition / 3u;
|
|
517
|
+
let kernelX = filterPosition % 3u;
|
|
518
|
+
for (var load = 0u; load < ${o*a/(i*r)}u; load++) {
|
|
519
|
+
let loadIndex = localLinear + load * ${i*r}u;
|
|
520
|
+
var value = ${n}(0.0);
|
|
521
|
+
if (filterPosition < 9u && (cachedMask[load] & (1u << filterPosition)) != 0u) {
|
|
522
|
+
${t.addressMode==="base-offset"?"":`let inputX = i32(cachedX[load]) + i32(kernelX) - 1;
|
|
523
|
+
let inputY = i32(cachedY[load]) + i32(kernelY) - 1;`}
|
|
524
|
+
${e}
|
|
525
|
+
}
|
|
526
|
+
inputTile[${ot("loadIndex",t)}] = value;
|
|
527
|
+
}
|
|
528
|
+
`}function ho(n,e,t,i,r,o){const{workgroupX:a,workgroupY:u,rowsPerThread:s,tileM:l,tileNBlocks:p,tileKBlocks:c}=re(o),{addressMode:d,weightLayout:h}=o,f=W(n),g=W(e),y=at(n,o,!1),v=at(n,o,!0),_=y?"vec2<u32>":f,w=v?"vec2<u32>":f,b=ae(he("acc",t),e),x=a*u,B=l*c,k=c*p*4;return`${oe(n)}
|
|
360
529
|
struct Params {
|
|
361
530
|
inputWidth: u32,
|
|
362
531
|
inputHeight: u32,
|
|
@@ -365,53 +534,60 @@ struct Params {
|
|
|
365
534
|
inputBlocks: u32,
|
|
366
535
|
outputBlocks: u32,
|
|
367
536
|
}
|
|
368
|
-
@group(0) @binding(0) var<storage, read> inputData: array<${
|
|
369
|
-
@group(0) @binding(1) var<storage, read> weights: array<${
|
|
537
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${_}>;
|
|
538
|
+
@group(0) @binding(1) var<storage, read> weights: array<${w}>;
|
|
370
539
|
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
371
|
-
@group(0) @binding(3) var<storage, read_write> outputData: array<${
|
|
540
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${g}>;
|
|
372
541
|
@group(0) @binding(4) var<uniform> params: Params;
|
|
373
542
|
|
|
374
|
-
var<workgroup> inputTile: array<${
|
|
375
|
-
var<workgroup> weightTile: array<${
|
|
543
|
+
var<workgroup> inputTile: array<${f}, ${ye(o).input}>;
|
|
544
|
+
var<workgroup> weightTile: array<${f}, ${ye(o).weights}>;
|
|
376
545
|
|
|
377
|
-
@compute @workgroup_size(${
|
|
546
|
+
@compute @workgroup_size(${a}, ${u}, 1)
|
|
378
547
|
fn main(
|
|
379
548
|
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
380
549
|
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
381
550
|
) {
|
|
382
551
|
let spatialBase =
|
|
383
|
-
workgroupId.x * ${
|
|
384
|
-
localId.y * ${
|
|
552
|
+
workgroupId.x * ${l}u +
|
|
553
|
+
localId.y * ${s}u;
|
|
385
554
|
let outputBlock =
|
|
386
|
-
workgroupId.y * ${
|
|
555
|
+
workgroupId.y * ${p}u + localId.x;
|
|
387
556
|
let spatialCount = params.outputWidth * params.outputHeight;
|
|
388
|
-
var acc: array<vec4<f32>, ${
|
|
389
|
-
if (outputBlock < ${
|
|
390
|
-
for (var row = 0u; row < ${
|
|
557
|
+
var acc: array<vec4<f32>, ${s}>;
|
|
558
|
+
if (outputBlock < ${r}u) {
|
|
559
|
+
for (var row = 0u; row < ${s}u; row++) {
|
|
391
560
|
acc[row] = bias[outputBlock];
|
|
392
561
|
}
|
|
393
562
|
}
|
|
394
563
|
|
|
395
564
|
let localLinear =
|
|
396
|
-
localId.y * ${
|
|
397
|
-
|
|
398
|
-
|
|
565
|
+
localId.y * ${a}u + localId.x;
|
|
566
|
+
${d!=="analytic"?mn(i,!1,o):""}
|
|
567
|
+
let totalK = ${i*9}u;
|
|
568
|
+
for (var kBase = 0u; kBase < totalK; kBase += ${c}u) {
|
|
569
|
+
${d!=="analytic"?st(f,`
|
|
570
|
+
${d==="base-offset"?`let offset = (i32(kernelY) - 1) * i32(params.inputWidth) + i32(kernelX) - 1;
|
|
571
|
+
let inputIndex = cachedBase0[load] + u32(offset * ${i}) + inputBlock;`:`let inputIndex = (u32(inputY) * params.inputWidth + u32(inputX)) *
|
|
572
|
+
${i}u + inputBlock;`}
|
|
573
|
+
value = ${ze("inputData[inputIndex]",y)};
|
|
574
|
+
`,o):`
|
|
399
575
|
for (
|
|
400
576
|
var loadIndex = localLinear;
|
|
401
|
-
loadIndex < ${
|
|
402
|
-
loadIndex += ${
|
|
577
|
+
loadIndex < ${B}u;
|
|
578
|
+
loadIndex += ${x}u
|
|
403
579
|
) {
|
|
404
|
-
let tileSpatial = loadIndex / ${
|
|
405
|
-
let tileK = loadIndex % ${
|
|
580
|
+
let tileSpatial = loadIndex / ${c}u;
|
|
581
|
+
let tileK = loadIndex % ${c}u;
|
|
406
582
|
let inputSpatialIndex =
|
|
407
|
-
workgroupId.x * ${
|
|
583
|
+
workgroupId.x * ${l}u + tileSpatial;
|
|
408
584
|
let kIndex = kBase + tileK;
|
|
409
|
-
var value = ${
|
|
585
|
+
var value = ${f}(0.0);
|
|
410
586
|
if (inputSpatialIndex < spatialCount && kIndex < totalK) {
|
|
411
587
|
let outputY = inputSpatialIndex / params.outputWidth;
|
|
412
588
|
let outputX = inputSpatialIndex % params.outputWidth;
|
|
413
|
-
let inputBlock = kIndex % ${
|
|
414
|
-
let kernelIndex = kIndex / ${
|
|
589
|
+
let inputBlock = kIndex % ${i}u;
|
|
590
|
+
let kernelIndex = kIndex / ${i}u;
|
|
415
591
|
let kernelY = kernelIndex / 3u;
|
|
416
592
|
let kernelX = kernelIndex % 3u;
|
|
417
593
|
let inputY = i32(outputY) + i32(kernelY) - 1;
|
|
@@ -422,55 +598,63 @@ fn main(
|
|
|
422
598
|
) {
|
|
423
599
|
let inputIndex =
|
|
424
600
|
(u32(inputY) * params.inputWidth + u32(inputX)) *
|
|
425
|
-
${
|
|
426
|
-
value = inputData[inputIndex];
|
|
601
|
+
${i}u + inputBlock;
|
|
602
|
+
value = ${ze("inputData[inputIndex]",y)};
|
|
427
603
|
}
|
|
428
604
|
}
|
|
429
|
-
inputTile[loadIndex] = value;
|
|
605
|
+
inputTile[${ot("loadIndex",o)}] = value;
|
|
430
606
|
}
|
|
607
|
+
`}
|
|
431
608
|
|
|
432
609
|
for (
|
|
433
610
|
var loadIndex = localLinear;
|
|
434
|
-
loadIndex < ${
|
|
435
|
-
loadIndex += ${
|
|
611
|
+
loadIndex < ${k}u;
|
|
612
|
+
loadIndex += ${x}u
|
|
436
613
|
) {
|
|
437
|
-
let tileK = loadIndex / ${
|
|
438
|
-
let outputRemainder = loadIndex % ${
|
|
614
|
+
let tileK = loadIndex / ${p*4}u;
|
|
615
|
+
let outputRemainder = loadIndex % ${p*4}u;
|
|
439
616
|
let tileOutputBlock = outputRemainder / 4u;
|
|
440
617
|
let outputLane = outputRemainder % 4u;
|
|
441
618
|
let loadedOutputBlock =
|
|
442
|
-
workgroupId.y * ${
|
|
619
|
+
workgroupId.y * ${p}u + tileOutputBlock;
|
|
443
620
|
let kIndex = kBase + tileK;
|
|
444
|
-
var value = ${
|
|
445
|
-
if (loadedOutputBlock < ${
|
|
446
|
-
|
|
447
|
-
let
|
|
621
|
+
var value = ${f}(0.0);
|
|
622
|
+
if (loadedOutputBlock < ${r}u && kIndex < totalK) {
|
|
623
|
+
${h==="k-major"?`
|
|
624
|
+
let weightIndex = (kIndex * ${r}u + loadedOutputBlock) * 4u + outputLane;
|
|
625
|
+
`:d!=="analytic"?`
|
|
626
|
+
let weightIndex = (loadedOutputBlock * ${i*9}u + kIndex) * 4u + outputLane;
|
|
627
|
+
`:`
|
|
628
|
+
let inputBlock = kIndex % ${i}u;
|
|
629
|
+
let kernelIndex = kIndex / ${i}u;
|
|
448
630
|
let kernelY = kernelIndex / 3u;
|
|
449
631
|
let kernelX = kernelIndex % 3u;
|
|
450
632
|
let weightIndex =
|
|
451
633
|
((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
|
|
452
|
-
${
|
|
453
|
-
|
|
634
|
+
${i}u + inputBlock) * 4u + outputLane);
|
|
635
|
+
`}
|
|
636
|
+
value = ${ze("weights[weightIndex]",v)};
|
|
454
637
|
}
|
|
455
|
-
weightTile[loadIndex] = value;
|
|
638
|
+
weightTile[${pn("loadIndex",o)}] = value;
|
|
456
639
|
}
|
|
457
640
|
|
|
458
641
|
workgroupBarrier();
|
|
459
|
-
${
|
|
642
|
+
${gn(n,o)}
|
|
460
643
|
workgroupBarrier();
|
|
644
|
+
${d!=="analytic"?_n(i,o):""}
|
|
461
645
|
}
|
|
462
646
|
|
|
463
|
-
if (outputBlock < ${
|
|
464
|
-
for (var row = 0u; row < ${
|
|
647
|
+
if (outputBlock < ${r}u) {
|
|
648
|
+
for (var row = 0u; row < ${s}u; row++) {
|
|
465
649
|
let spatialIndex = spatialBase + row;
|
|
466
650
|
if (spatialIndex < spatialCount) {
|
|
467
|
-
let outputIndex = spatialIndex * ${
|
|
468
|
-
outputData[outputIndex] = ${
|
|
651
|
+
let outputIndex = spatialIndex * ${r}u + outputBlock;
|
|
652
|
+
outputData[outputIndex] = ${b.replaceAll("acc","acc[row]")};
|
|
469
653
|
}
|
|
470
654
|
}
|
|
471
655
|
}
|
|
472
656
|
}
|
|
473
|
-
`}function
|
|
657
|
+
`}function fo(n,e,t,i){const r=W(n),o=he("acc",e);return`${oe(n)}
|
|
474
658
|
struct Params {
|
|
475
659
|
inputWidth: u32,
|
|
476
660
|
inputHeight: u32,
|
|
@@ -479,15 +663,15 @@ struct Params {
|
|
|
479
663
|
inputBlocks: u32,
|
|
480
664
|
outputBlocks: u32,
|
|
481
665
|
}
|
|
482
|
-
@group(0) @binding(0) var<storage, read> inputData: array<${
|
|
483
|
-
@group(0) @binding(1) var<storage, read> weights: array<${
|
|
666
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${r}>;
|
|
667
|
+
@group(0) @binding(1) var<storage, read> weights: array<${r}>;
|
|
484
668
|
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
485
|
-
@group(0) @binding(3) var<storage, read_write> outputData: array<${
|
|
669
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${r}>;
|
|
486
670
|
@group(0) @binding(4) var<uniform> params: Params;
|
|
487
671
|
|
|
488
|
-
@compute @workgroup_size(${
|
|
672
|
+
@compute @workgroup_size(${R}, ${R}, 1)
|
|
489
673
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
490
|
-
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${
|
|
674
|
+
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${i}u) {
|
|
491
675
|
return;
|
|
492
676
|
}
|
|
493
677
|
var pooled = vec4<f32>(-3.402823466e+38);
|
|
@@ -506,17 +690,17 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
506
690
|
if (inputX < 0 || inputX >= i32(params.inputWidth)) { continue; }
|
|
507
691
|
let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${t}u;
|
|
508
692
|
for (var inputBlock = 0u; inputBlock < ${t}u; inputBlock++) {
|
|
509
|
-
${
|
|
693
|
+
${Ue("inputData[pixelBase + inputBlock]",`((((gid.z * 3u + ky) * 3u + kx) * ${t}u + inputBlock) * 4u)`,n)}
|
|
510
694
|
}
|
|
511
695
|
}
|
|
512
696
|
}
|
|
513
|
-
pooled = max(pooled, ${
|
|
697
|
+
pooled = max(pooled, ${o});
|
|
514
698
|
}
|
|
515
699
|
}
|
|
516
|
-
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${
|
|
517
|
-
outputData[outputIndex] = ${
|
|
700
|
+
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${i}u + gid.z;
|
|
701
|
+
outputData[outputIndex] = ${ae("pooled",n)};
|
|
518
702
|
}
|
|
519
|
-
`}function
|
|
703
|
+
`}function go(n,e,t){const i=W(n);return`${oe(n)}
|
|
520
704
|
struct Params {
|
|
521
705
|
inputWidth: u32,
|
|
522
706
|
inputHeight: u32,
|
|
@@ -524,16 +708,23 @@ struct Params {
|
|
|
524
708
|
outputHeight: u32,
|
|
525
709
|
outputBlocks: u32,
|
|
526
710
|
}
|
|
527
|
-
@group(0) @binding(0) var<storage, read> inputData: array<${
|
|
528
|
-
@group(0) @binding(1) var<storage, read_write> outputData: array<${
|
|
711
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${i}>;
|
|
712
|
+
@group(0) @binding(1) var<storage, read_write> outputData: array<${i}>;
|
|
529
713
|
@group(0) @binding(2) var<uniform> params: Params;
|
|
530
714
|
|
|
531
|
-
@compute @workgroup_size(${
|
|
532
|
-
fn main(
|
|
715
|
+
@compute @workgroup_size(${R}, ${R}, 1)
|
|
716
|
+
fn main(
|
|
717
|
+
@builtin(global_invocation_id) gid: vec3<u32>${t?`,
|
|
718
|
+
@builtin(num_workgroups) groupCount: vec3<u32>`:""}
|
|
719
|
+
) {
|
|
720
|
+
${t?`let linearX = gid.x + gid.z * groupCount.x * ${R}u;
|
|
721
|
+
let outputBlock = linearX % ${e}u;
|
|
722
|
+
let outputX = linearX / ${e}u;`:`let outputBlock = gid.z;
|
|
723
|
+
let outputX = gid.x;`}
|
|
533
724
|
if (
|
|
534
|
-
|
|
725
|
+
outputX >= params.outputWidth ||
|
|
535
726
|
gid.y >= params.outputHeight ||
|
|
536
|
-
|
|
727
|
+
outputBlock >= ${e}u
|
|
537
728
|
) {
|
|
538
729
|
return;
|
|
539
730
|
}
|
|
@@ -542,27 +733,30 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
542
733
|
let inputY = gid.y * 2u + py;
|
|
543
734
|
if (inputY >= params.inputHeight) { continue; }
|
|
544
735
|
for (var px = 0u; px < 2u; px++) {
|
|
545
|
-
let inputX =
|
|
736
|
+
let inputX = outputX * 2u + px;
|
|
546
737
|
if (inputX >= params.inputWidth) { continue; }
|
|
547
738
|
let inputIndex =
|
|
548
|
-
(inputY * params.inputWidth + inputX) * ${e}u +
|
|
739
|
+
(inputY * params.inputWidth + inputX) * ${e}u + outputBlock;
|
|
549
740
|
pooled = max(pooled, vec4<f32>(inputData[inputIndex]));
|
|
550
741
|
}
|
|
551
742
|
}
|
|
552
743
|
let outputIndex =
|
|
553
|
-
(gid.y * params.outputWidth +
|
|
554
|
-
outputData[outputIndex] = ${
|
|
744
|
+
(gid.y * params.outputWidth + outputX) * ${e}u + outputBlock;
|
|
745
|
+
outputData[outputIndex] = ${ae("pooled",n)};
|
|
555
746
|
}
|
|
556
|
-
`}function
|
|
747
|
+
`}function mo(n,e,t,i,r,o,a){const{workgroupX:u,workgroupY:s,rowsPerThread:l,tileM:p,tileNBlocks:c,tileKBlocks:d}=re(a),{addressMode:h,weightLayout:f}=a,g=W(n),y=W(e),v=at(n,a,!1),_=at(n,a,!0),w=v?"vec2<u32>":g,b=_?"vec2<u32>":g,x=i[0]+i[1],B=ae(he("acc[row]",t),e),k=u*s,C=p*d,L=d*c*4,F=(z,Pe)=>`
|
|
557
748
|
{
|
|
558
|
-
let sourceBlock = ${
|
|
559
|
-
let
|
|
560
|
-
let
|
|
749
|
+
let sourceBlock = ${Pe};
|
|
750
|
+
${h==="base-offset"?`let dx = ${z===r?"(i32((cachedMask[load] >> 9u) & 1u) + i32(kernelX) - 1) >> 1u":"i32(kernelX) - 1"};
|
|
751
|
+
let dy = ${z===r?"(i32((cachedMask[load] >> 10u) & 1u) + i32(kernelY) - 1) >> 1u":"i32(kernelY) - 1"};
|
|
752
|
+
let offset = (dy * i32(params.source${z}Width) + dx) * ${i[z]};
|
|
753
|
+
let sourceIndex = cachedBase${z}[load] + u32(offset) + sourceBlock;`:`let sourceX = ${z===r?"u32(inputX) / 2u":"u32(inputX)"};
|
|
754
|
+
let sourceY = ${z===r?"u32(inputY) / 2u":"u32(inputY)"};
|
|
561
755
|
let sourceIndex =
|
|
562
|
-
(sourceY * params.source${
|
|
563
|
-
${
|
|
564
|
-
value = input${
|
|
565
|
-
}`;return`${
|
|
756
|
+
(sourceY * params.source${z}Width + sourceX) *
|
|
757
|
+
${i[z]}u + sourceBlock;`}
|
|
758
|
+
value = ${ze(`input${z}[sourceIndex]`,v)};
|
|
759
|
+
}`;return`${oe(n)}
|
|
566
760
|
struct Params {
|
|
567
761
|
outputWidth: u32,
|
|
568
762
|
outputHeight: u32,
|
|
@@ -573,54 +767,68 @@ struct Params {
|
|
|
573
767
|
source1Width: u32,
|
|
574
768
|
source1Height: u32,
|
|
575
769
|
}
|
|
576
|
-
@group(0) @binding(0) var<storage, read> input0: array<${
|
|
577
|
-
@group(0) @binding(1) var<storage, read> input1: array<${
|
|
578
|
-
@group(0) @binding(2) var<storage, read> weights: array<${
|
|
770
|
+
@group(0) @binding(0) var<storage, read> input0: array<${w}>;
|
|
771
|
+
@group(0) @binding(1) var<storage, read> input1: array<${w}>;
|
|
772
|
+
@group(0) @binding(2) var<storage, read> weights: array<${b}>;
|
|
579
773
|
@group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
|
|
580
|
-
@group(0) @binding(4) var<storage, read_write> outputData: array<${
|
|
774
|
+
@group(0) @binding(4) var<storage, read_write> outputData: array<${y}>;
|
|
581
775
|
@group(0) @binding(5) var<uniform> params: Params;
|
|
582
776
|
|
|
583
|
-
var<workgroup> inputTile: array<${
|
|
584
|
-
var<workgroup> weightTile: array<${
|
|
777
|
+
var<workgroup> inputTile: array<${g}, ${ye(a).input}>;
|
|
778
|
+
var<workgroup> weightTile: array<${g}, ${ye(a).weights}>;
|
|
585
779
|
|
|
586
|
-
@compute @workgroup_size(${
|
|
780
|
+
@compute @workgroup_size(${u}, ${s}, 1)
|
|
587
781
|
fn main(
|
|
588
782
|
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
589
783
|
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
590
784
|
) {
|
|
591
785
|
let spatialBase =
|
|
592
|
-
workgroupId.x * ${
|
|
593
|
-
localId.y * ${
|
|
786
|
+
workgroupId.x * ${p}u +
|
|
787
|
+
localId.y * ${l}u;
|
|
594
788
|
let outputBlock =
|
|
595
|
-
workgroupId.y * ${
|
|
789
|
+
workgroupId.y * ${c}u + localId.x;
|
|
596
790
|
let spatialCount = params.outputWidth * params.outputHeight;
|
|
597
|
-
var acc: array<vec4<f32>, ${
|
|
791
|
+
var acc: array<vec4<f32>, ${l}>;
|
|
598
792
|
if (outputBlock < ${o}u) {
|
|
599
|
-
for (var row = 0u; row < ${
|
|
793
|
+
for (var row = 0u; row < ${l}u; row++) {
|
|
600
794
|
acc[row] = bias[outputBlock];
|
|
601
795
|
}
|
|
602
796
|
}
|
|
603
797
|
|
|
604
798
|
let localLinear =
|
|
605
|
-
localId.y * ${
|
|
606
|
-
|
|
607
|
-
|
|
799
|
+
localId.y * ${u}u + localId.x;
|
|
800
|
+
${h!=="analytic"?mn(x,!0,a,{blocks:i,upsampled:r}):""}
|
|
801
|
+
let totalK = ${x*9}u;
|
|
802
|
+
for (var kBase = 0u; kBase < totalK; kBase += ${d}u) {
|
|
803
|
+
${h!=="analytic"?a.decoderLoad==="source-first"?`
|
|
804
|
+
if (channelBlock < ${i[0]}u) {
|
|
805
|
+
${st(g,F(0,"inputBlock"),a)}
|
|
806
|
+
} else {
|
|
807
|
+
${st(g,F(1,`inputBlock - ${i[0]}u`),a)}
|
|
808
|
+
}
|
|
809
|
+
`:st(g,`
|
|
810
|
+
if (inputBlock < ${i[0]}u) {
|
|
811
|
+
${F(0,"inputBlock")}
|
|
812
|
+
} else {
|
|
813
|
+
${F(1,`inputBlock - ${i[0]}u`)}
|
|
814
|
+
}
|
|
815
|
+
`,a):`
|
|
608
816
|
for (
|
|
609
817
|
var loadIndex = localLinear;
|
|
610
|
-
loadIndex < ${
|
|
611
|
-
loadIndex += ${
|
|
818
|
+
loadIndex < ${C}u;
|
|
819
|
+
loadIndex += ${k}u
|
|
612
820
|
) {
|
|
613
|
-
let tileSpatial = loadIndex / ${
|
|
614
|
-
let tileK = loadIndex % ${
|
|
821
|
+
let tileSpatial = loadIndex / ${d}u;
|
|
822
|
+
let tileK = loadIndex % ${d}u;
|
|
615
823
|
let outputSpatialIndex =
|
|
616
|
-
workgroupId.x * ${
|
|
824
|
+
workgroupId.x * ${p}u + tileSpatial;
|
|
617
825
|
let kIndex = kBase + tileK;
|
|
618
|
-
var value = ${
|
|
826
|
+
var value = ${g}(0.0);
|
|
619
827
|
if (outputSpatialIndex < spatialCount && kIndex < totalK) {
|
|
620
828
|
let outputY = outputSpatialIndex / params.outputWidth;
|
|
621
829
|
let outputX = outputSpatialIndex % params.outputWidth;
|
|
622
|
-
let inputBlock = kIndex % ${
|
|
623
|
-
let kernelIndex = kIndex / ${
|
|
830
|
+
let inputBlock = kIndex % ${x}u;
|
|
831
|
+
let kernelIndex = kIndex / ${x}u;
|
|
624
832
|
let kernelY = kernelIndex / 3u;
|
|
625
833
|
let kernelX = kernelIndex % 3u;
|
|
626
834
|
let inputY = i32(outputY) + i32(kernelY) - 1;
|
|
@@ -629,68 +837,76 @@ fn main(
|
|
|
629
837
|
inputY >= 0 && inputY < i32(params.outputHeight) &&
|
|
630
838
|
inputX >= 0 && inputX < i32(params.outputWidth)
|
|
631
839
|
) {
|
|
632
|
-
if (inputBlock < ${
|
|
633
|
-
${
|
|
840
|
+
if (inputBlock < ${i[0]}u) {
|
|
841
|
+
${F(0,"inputBlock")}
|
|
634
842
|
} else {
|
|
635
|
-
${
|
|
843
|
+
${F(1,`inputBlock - ${i[0]}u`)}
|
|
636
844
|
}
|
|
637
845
|
}
|
|
638
846
|
}
|
|
639
|
-
inputTile[loadIndex] = value;
|
|
847
|
+
inputTile[${ot("loadIndex",a)}] = value;
|
|
640
848
|
}
|
|
849
|
+
`}
|
|
641
850
|
|
|
642
851
|
for (
|
|
643
852
|
var loadIndex = localLinear;
|
|
644
|
-
loadIndex < ${
|
|
645
|
-
loadIndex += ${
|
|
853
|
+
loadIndex < ${L}u;
|
|
854
|
+
loadIndex += ${k}u
|
|
646
855
|
) {
|
|
647
|
-
let tileK = loadIndex / ${
|
|
648
|
-
let outputRemainder = loadIndex % ${
|
|
856
|
+
let tileK = loadIndex / ${c*4}u;
|
|
857
|
+
let outputRemainder = loadIndex % ${c*4}u;
|
|
649
858
|
let tileOutputBlock = outputRemainder / 4u;
|
|
650
859
|
let outputLane = outputRemainder % 4u;
|
|
651
860
|
let loadedOutputBlock =
|
|
652
|
-
workgroupId.y * ${
|
|
861
|
+
workgroupId.y * ${c}u + tileOutputBlock;
|
|
653
862
|
let kIndex = kBase + tileK;
|
|
654
|
-
var value = ${
|
|
863
|
+
var value = ${g}(0.0);
|
|
655
864
|
if (loadedOutputBlock < ${o}u && kIndex < totalK) {
|
|
656
|
-
|
|
657
|
-
let
|
|
865
|
+
${f==="k-major"?`
|
|
866
|
+
let weightIndex = (kIndex * ${o}u + loadedOutputBlock) * 4u + outputLane;
|
|
867
|
+
`:h!=="analytic"?`
|
|
868
|
+
let weightIndex = (loadedOutputBlock * ${x*9}u + kIndex) * 4u + outputLane;
|
|
869
|
+
`:`
|
|
870
|
+
let inputBlock = kIndex % ${x}u;
|
|
871
|
+
let kernelIndex = kIndex / ${x}u;
|
|
658
872
|
let kernelY = kernelIndex / 3u;
|
|
659
873
|
let kernelX = kernelIndex % 3u;
|
|
660
874
|
let weightIndex =
|
|
661
875
|
((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
|
|
662
|
-
${
|
|
663
|
-
|
|
876
|
+
${x}u + inputBlock) * 4u + outputLane);
|
|
877
|
+
`}
|
|
878
|
+
value = ${ze("weights[weightIndex]",_)};
|
|
664
879
|
}
|
|
665
|
-
weightTile[loadIndex] = value;
|
|
880
|
+
weightTile[${pn("loadIndex",a)}] = value;
|
|
666
881
|
}
|
|
667
882
|
|
|
668
883
|
workgroupBarrier();
|
|
669
|
-
${
|
|
884
|
+
${gn(n,a)}
|
|
670
885
|
workgroupBarrier();
|
|
886
|
+
${h!=="analytic"?_n(x,a):""}
|
|
671
887
|
}
|
|
672
888
|
|
|
673
889
|
if (outputBlock < ${o}u) {
|
|
674
|
-
for (var row = 0u; row < ${
|
|
890
|
+
for (var row = 0u; row < ${l}u; row++) {
|
|
675
891
|
let spatialIndex = spatialBase + row;
|
|
676
892
|
if (spatialIndex < spatialCount) {
|
|
677
893
|
let outputIndex = spatialIndex * ${o}u + outputBlock;
|
|
678
|
-
outputData[outputIndex] = ${
|
|
894
|
+
outputData[outputIndex] = ${B};
|
|
679
895
|
}
|
|
680
896
|
}
|
|
681
897
|
}
|
|
682
898
|
}
|
|
683
|
-
`}function
|
|
899
|
+
`}function _o(n,e,t,i,r,o){const a=W(n),u=W(e),s=i[0]+i[1],l=(c,d)=>{const h=c===r;return`
|
|
684
900
|
{
|
|
685
|
-
let sourceX = ${
|
|
686
|
-
let sourceY = ${
|
|
687
|
-
let sourcePixelBase = (sourceY * params.source${
|
|
688
|
-
for (var sourceBlock = 0u; sourceBlock < ${
|
|
689
|
-
let inputBlock = ${
|
|
690
|
-
${
|
|
901
|
+
let sourceX = ${h?"u32(inputX) / 2u":"u32(inputX)"};
|
|
902
|
+
let sourceY = ${h?"u32(inputY) / 2u":"u32(inputY)"};
|
|
903
|
+
let sourcePixelBase = (sourceY * params.source${c}Width + sourceX) * ${i[c]}u;
|
|
904
|
+
for (var sourceBlock = 0u; sourceBlock < ${i[c]}u; sourceBlock++) {
|
|
905
|
+
let inputBlock = ${d}u + sourceBlock;
|
|
906
|
+
${Ue(`input${c}[sourcePixelBase + sourceBlock]`,`((((gid.z * 3u + ky) * 3u + kx) * ${s}u + inputBlock) * 4u)`,n)}
|
|
691
907
|
}
|
|
692
908
|
}
|
|
693
|
-
`},
|
|
909
|
+
`},p=ae(he("acc",t),e);return`${oe(n)}
|
|
694
910
|
struct Params {
|
|
695
911
|
outputWidth: u32,
|
|
696
912
|
outputHeight: u32,
|
|
@@ -701,14 +917,14 @@ struct Params {
|
|
|
701
917
|
source1Width: u32,
|
|
702
918
|
source1Height: u32,
|
|
703
919
|
}
|
|
704
|
-
@group(0) @binding(0) var<storage, read> input0: array<${
|
|
705
|
-
@group(0) @binding(1) var<storage, read> input1: array<${
|
|
706
|
-
@group(0) @binding(2) var<storage, read> weights: array<${
|
|
920
|
+
@group(0) @binding(0) var<storage, read> input0: array<${a}>;
|
|
921
|
+
@group(0) @binding(1) var<storage, read> input1: array<${a}>;
|
|
922
|
+
@group(0) @binding(2) var<storage, read> weights: array<${a}>;
|
|
707
923
|
@group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
|
|
708
|
-
@group(0) @binding(4) var<storage, read_write> outputData: array<${
|
|
924
|
+
@group(0) @binding(4) var<storage, read_write> outputData: array<${u}>;
|
|
709
925
|
@group(0) @binding(5) var<uniform> params: Params;
|
|
710
926
|
|
|
711
|
-
@compute @workgroup_size(${
|
|
927
|
+
@compute @workgroup_size(${R}, ${R}, 1)
|
|
712
928
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
713
929
|
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${o}u) {
|
|
714
930
|
return;
|
|
@@ -720,19 +936,19 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
720
936
|
for (var kx = 0u; kx < 3u; kx++) {
|
|
721
937
|
let inputX = i32(gid.x) + i32(kx) - 1;
|
|
722
938
|
if (inputX < 0 || inputX >= i32(params.outputWidth)) { continue; }
|
|
723
|
-
${
|
|
724
|
-
${
|
|
939
|
+
${l(0,0)}
|
|
940
|
+
${l(1,i[0])}
|
|
725
941
|
}
|
|
726
942
|
}
|
|
727
943
|
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${o}u + gid.z;
|
|
728
|
-
outputData[outputIndex] = ${
|
|
944
|
+
outputData[outputIndex] = ${p};
|
|
729
945
|
}
|
|
730
|
-
`}function
|
|
946
|
+
`}function yo(n,e,t,i,r,o){const a=W(n),u=W(e),s=i[0]+i[1],l=J*J*s,p=U*U,c=ae(he("acc",t),e),d=(h,f)=>`
|
|
731
947
|
let sourceBlock = ${f};
|
|
732
948
|
let sourceIndex =
|
|
733
|
-
(${
|
|
734
|
-
${
|
|
735
|
-
value = input${
|
|
949
|
+
(${h===r?"u32(inputY) / 2u":"u32(inputY)"} * params.source${h}Width + ${h===r?"u32(inputX) / 2u":"u32(inputX)"}) *
|
|
950
|
+
${i[h]}u + sourceBlock;
|
|
951
|
+
value = input${h}[sourceIndex];`;return`${oe(n)}
|
|
736
952
|
struct Params {
|
|
737
953
|
outputWidth: u32,
|
|
738
954
|
outputHeight: u32,
|
|
@@ -743,44 +959,44 @@ struct Params {
|
|
|
743
959
|
source1Width: u32,
|
|
744
960
|
source1Height: u32,
|
|
745
961
|
}
|
|
746
|
-
@group(0) @binding(0) var<storage, read> input0: array<${
|
|
747
|
-
@group(0) @binding(1) var<storage, read> input1: array<${
|
|
748
|
-
@group(0) @binding(2) var<storage, read> weights: array<${
|
|
962
|
+
@group(0) @binding(0) var<storage, read> input0: array<${a}>;
|
|
963
|
+
@group(0) @binding(1) var<storage, read> input1: array<${a}>;
|
|
964
|
+
@group(0) @binding(2) var<storage, read> weights: array<${a}>;
|
|
749
965
|
@group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
|
|
750
|
-
@group(0) @binding(4) var<storage, read_write> outputData: array<${
|
|
966
|
+
@group(0) @binding(4) var<storage, read_write> outputData: array<${u}>;
|
|
751
967
|
@group(0) @binding(5) var<uniform> params: Params;
|
|
752
968
|
|
|
753
|
-
var<workgroup> inputPatch: array<${
|
|
969
|
+
var<workgroup> inputPatch: array<${a}, ${l}>;
|
|
754
970
|
|
|
755
|
-
@compute @workgroup_size(${
|
|
971
|
+
@compute @workgroup_size(${U}, ${U}, 1)
|
|
756
972
|
fn main(
|
|
757
973
|
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
758
974
|
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
759
975
|
) {
|
|
760
976
|
let localLinear =
|
|
761
|
-
localId.y * ${
|
|
977
|
+
localId.y * ${U}u + localId.x;
|
|
762
978
|
for (
|
|
763
979
|
var loadIndex = localLinear;
|
|
764
|
-
loadIndex < ${
|
|
765
|
-
loadIndex += ${
|
|
980
|
+
loadIndex < ${l}u;
|
|
981
|
+
loadIndex += ${p}u
|
|
766
982
|
) {
|
|
767
983
|
let patchPixel = loadIndex / ${s}u;
|
|
768
984
|
let inputBlock = loadIndex % ${s}u;
|
|
769
|
-
let patchX = patchPixel % ${
|
|
770
|
-
let patchY = patchPixel / ${
|
|
985
|
+
let patchX = patchPixel % ${J}u;
|
|
986
|
+
let patchY = patchPixel / ${J}u;
|
|
771
987
|
let inputX =
|
|
772
|
-
i32(workgroupId.x * ${
|
|
988
|
+
i32(workgroupId.x * ${U}u + patchX) - 1;
|
|
773
989
|
let inputY =
|
|
774
|
-
i32(workgroupId.y * ${
|
|
775
|
-
var value = ${
|
|
990
|
+
i32(workgroupId.y * ${U}u + patchY) - 1;
|
|
991
|
+
var value = ${a}(0.0);
|
|
776
992
|
if (
|
|
777
993
|
inputX >= 0 && inputX < i32(params.outputWidth) &&
|
|
778
994
|
inputY >= 0 && inputY < i32(params.outputHeight)
|
|
779
995
|
) {
|
|
780
|
-
if (inputBlock < ${
|
|
781
|
-
${
|
|
996
|
+
if (inputBlock < ${i[0]}u) {
|
|
997
|
+
${d(0,"inputBlock")}
|
|
782
998
|
} else {
|
|
783
|
-
${
|
|
999
|
+
${d(1,`inputBlock - ${i[0]}u`)}
|
|
784
1000
|
}
|
|
785
1001
|
}
|
|
786
1002
|
inputPatch[loadIndex] = value;
|
|
@@ -788,9 +1004,9 @@ fn main(
|
|
|
788
1004
|
workgroupBarrier();
|
|
789
1005
|
|
|
790
1006
|
let outputX =
|
|
791
|
-
workgroupId.x * ${
|
|
1007
|
+
workgroupId.x * ${U}u + localId.x;
|
|
792
1008
|
let outputY =
|
|
793
|
-
workgroupId.y * ${
|
|
1009
|
+
workgroupId.y * ${U}u + localId.y;
|
|
794
1010
|
let outputBlock = workgroupId.z;
|
|
795
1011
|
if (
|
|
796
1012
|
outputX >= params.outputWidth || outputY >= params.outputHeight ||
|
|
@@ -803,36 +1019,36 @@ fn main(
|
|
|
803
1019
|
for (var ky = 0u; ky < 3u; ky++) {
|
|
804
1020
|
for (var kx = 0u; kx < 3u; kx++) {
|
|
805
1021
|
let patchBase =
|
|
806
|
-
((localId.y + ky) * ${
|
|
1022
|
+
((localId.y + ky) * ${J}u + localId.x + kx) *
|
|
807
1023
|
${s}u;
|
|
808
1024
|
for (var inputBlock = 0u; inputBlock < ${s}u; inputBlock++) {
|
|
809
|
-
${
|
|
1025
|
+
${Ue("inputPatch[patchBase + inputBlock]",`((((outputBlock * 3u + ky) * 3u + kx) * ${s}u + inputBlock) * 4u)`,n)}
|
|
810
1026
|
}
|
|
811
1027
|
}
|
|
812
1028
|
}
|
|
813
1029
|
let outputIndex =
|
|
814
1030
|
(outputY * params.outputWidth + outputX) * ${o}u + outputBlock;
|
|
815
|
-
outputData[outputIndex] = ${
|
|
1031
|
+
outputData[outputIndex] = ${c};
|
|
816
1032
|
}
|
|
817
|
-
`}function
|
|
818
|
-
`),
|
|
819
|
-
`),
|
|
1033
|
+
`}function yn(n,e){const t=W(n),i=Array.from({length:e},(u,s)=>`@group(0) @binding(${s}) var<storage, read> input${s}: array<vec4<f32>>;`).join(`
|
|
1034
|
+
`),r=Array.from({length:e},(u,s)=>{const l=s*3;return`if (channel < ${l+3}u) { return input${s}[pixel][channel - ${l}u]; }`}).join(`
|
|
1035
|
+
`),o=e,a=e+1;return`${oe(n)}
|
|
820
1036
|
struct Params {
|
|
821
1037
|
width: u32,
|
|
822
1038
|
height: u32,
|
|
823
1039
|
outputBlocks: u32,
|
|
824
1040
|
inputChannels: u32,
|
|
825
1041
|
}
|
|
826
|
-
${
|
|
827
|
-
@group(0) @binding(${
|
|
828
|
-
@group(0) @binding(${
|
|
1042
|
+
${i}
|
|
1043
|
+
@group(0) @binding(${o}) var<storage, read_write> outputData: array<${t}>;
|
|
1044
|
+
@group(0) @binding(${a}) var<uniform> params: Params;
|
|
829
1045
|
|
|
830
1046
|
fn readChannel(pixel: u32, channel: u32) -> f32 {
|
|
831
|
-
${
|
|
1047
|
+
${r}
|
|
832
1048
|
return 0.0;
|
|
833
1049
|
}
|
|
834
1050
|
|
|
835
|
-
@compute @workgroup_size(${
|
|
1051
|
+
@compute @workgroup_size(${R}, ${R}, 1)
|
|
836
1052
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
837
1053
|
if (gid.x >= params.width || gid.y >= params.height || gid.z >= params.outputBlocks) {
|
|
838
1054
|
return;
|
|
@@ -845,24 +1061,24 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
845
1061
|
readChannel(pixel, firstChannel + 2u),
|
|
846
1062
|
readChannel(pixel, firstChannel + 3u)
|
|
847
1063
|
);
|
|
848
|
-
outputData[pixel * params.outputBlocks + gid.z] = ${
|
|
1064
|
+
outputData[pixel * params.outputBlocks + gid.z] = ${ae("value",n)};
|
|
849
1065
|
}
|
|
850
|
-
`}function Ar(i){return i.op==="concat"?i.inputs:i.op==="fusedUpsampleConcatConv2d"?i.inputs.map(e=>e.value):[i.input]}function Zn(i,e="auto"){const t=i.features.has("shader-f16");if(e==="fp16"&&!t)throw new Error("OIDN FP16 was requested but the GPUDevice does not have shader-f16 enabled");return e==="auto"?t?"fp16":"fp32":e}class Jn{constructor(e,t,n={}){g(this,"precision");g(this,"kernelSetting");g(this,"maxSpatialInputBlocks");g(this,"subgroupsAvailable");g(this,"_model");g(this,"_packedConvs",new Map);g(this,"_pipelineCache");g(this,"_pipelinePromises");g(this,"_executionCache",new Map);g(this,"_retiredExecutions",new Set);g(this,"_clock",0);g(this,"_shapeCacheSize");g(this,"_profileNextExecution",!1);g(this,"_lastExecutionProfile");g(this,"_profileOperations",0);g(this,"_resources",new Fn);g(this,"_disposed",!1);this._device=e;const o=mr(e);if(this._pipelineCache=o.ready,this._pipelinePromises=o.pending,this.precision=Zn(e,n.precision??"auto"),this.kernelSetting=n.kernel??"auto",this.subgroupsAvailable=e.features.has("subgroups"),this.maxSpatialInputBlocks=this.precision==="fp16"&&e.limits.maxComputeInvocationsPerWorkgroup>=E*E&&e.limits.maxComputeWorkgroupSizeX>=E&&e.limits.maxComputeWorkgroupSizeY>=E?Math.floor(e.limits.maxComputeWorkgroupStorageSize/(te*te*4*2)):0,this._shapeCacheSize=Math.max(1,n.shapeCacheSize??2),t.inputChannels%3!==0||t.inputChannels<3||t.inputChannels>9)throw new Error(`Native OIDN expects 3, 6, or 9 input channels, got ${t.inputChannels}`);this._model={spec:t.spec,inputChannels:t.inputChannels,outputChannels:t.outputChannels,channelsByValue:new Map(t.channelsByValue),convChannels:new Map(t.convChannels)};for(const r of nt(this._model,{fuseConvPool:!1}).nodes)if(r.op!=="conv2d"&&r.op!=="maxPool2d"&&r.op!=="fusedConvReluMaxPool2d"&&r.op!=="fusedUpsampleConcatConv2d")throw new Error(`Native OIDN descriptor ${t.spec.id} leaves unsupported ${r.op} node ${r.id} after graph optimization`);try{for(const[r,s]of t.convTensors){const a=xr(e,r,s,this.precision);this._resources.track("gpu-buffer",a.weights),this._resources.track("gpu-buffer",a.bias),this._packedConvs.set(r,a)}}catch(r){for(const s of this._packedConvs.values())this._releaseBuffer(s.weights),this._releaseBuffer(s.bias);throw this._packedConvs.clear(),r}}_pipeline(e,t){let n=this._pipelineCache.get(e);return n||(n=this._device.createComputePipeline({label:`oidn/${e}`,layout:"auto",compute:{module:this._device.createShaderModule({label:`oidn/${e}`,code:t}),entryPoint:"main"}}),this._pipelineCache.set(e,n)),n}_pipelineAsync(e,t){const n=this._pipelineCache.get(e);if(n)return Promise.resolve(n);const o=this._pipelinePromises.get(e);if(o)return o;const r=this._device.createComputePipelineAsync({label:`oidn/${e}`,layout:"auto",compute:{module:this._device.createShaderModule({label:`oidn/${e}`,code:t}),entryPoint:"main"}}).then(s=>(this._pipelineCache.set(e,s),this._pipelinePromises.delete(e),s),s=>{throw this._pipelinePromises.delete(e),s});return this._pipelinePromises.set(e,r),r}_nodePipelineSpec(e,t){if(e.op==="conv2d"){const n=t?"fp32":this.precision,o=U(this._model.convChannels.get(e.id).inputChannels),r=U(this._model.convChannels.get(e.id).outputChannels),s=this._selectConvKernel(o,t);return{key:`conv-${s}/${this.precision}/${n}/${e.activation}/in${o}/out${r}`,kernel:s,code:s==="implicit-gemm"?Pr(this.precision,n,e.activation,o,r):s==="spatial"?Br(this.precision,n,e.activation,o,r):s==="subgroup"?kr(this.precision,n,e.activation,o,r):$r(this.precision,n,e.activation,o,r)}}if(e.op==="maxPool2d"){const n=U(this._model.channelsByValue.get(e.id));return{key:`max-pool/${this.precision}/out${n}`,kernel:"direct",code:Tr(this.precision,n)}}if(e.op==="fusedConvReluMaxPool2d"){const n=U(this._model.convChannels.get(e.conv.id).inputChannels),o=U(this._model.convChannels.get(e.conv.id).outputChannels);return{key:`conv-pool/${this.precision}/${e.conv.activation}/in${n}/out${o}`,kernel:"direct",code:Ir(this.precision,e.conv.activation,n,o)}}if(e.op==="fusedUpsampleConcatConv2d"){if(e.inputs.length!==2)throw new Error(`Native fused decoder ${e.id} requires two inputs`);const n=e.inputs.map(p=>U(this._model.channelsByValue.get(p.value))),o=e.inputs.findIndex(p=>p.upsample);if(o!==0&&o!==1)throw new Error(`Native fused decoder ${e.id} has no upsample input`);const r=n[0]+n[1],s=this._selectConvKernel(r,!1),a=s==="subgroup"?"direct":s,u=U(this._model.convChannels.get(e.conv.id).outputChannels);return{key:`decoder-${a}/${this.precision}/${e.conv.activation}/${n.join("+")}/out${u}/up${o}`,kernel:a,code:a==="implicit-gemm"?Cr(this.precision,e.conv.activation,n,o,u):a==="spatial"?Er(this.precision,e.conv.activation,n,o,u):Sr(this.precision,e.conv.activation,n,o,u)}}throw new Error(`Native OIDN does not implement unfused ${e.op} node ${e.id}`)}_selectConvKernel(e,t){const n=this.precision==="fp16"&&e<=this.maxSpatialInputBlocks;return this.kernelSetting==="direct"?"direct":this.kernelSetting==="spatial"?n?"spatial":"direct":this.kernelSetting==="implicit-gemm"?t?"direct":"implicit-gemm":this.kernelSetting==="subgroup"?this.subgroupsAvailable?"subgroup":"direct":this.precision==="fp32"&&!t?"implicit-gemm":"direct"}_nodePipeline(e,t){const{key:n,code:o}=this._nodePipelineSpec(e,t);return this._pipeline(n,o)}async prepare(){if(this._disposed)throw new Error("Native OIDN executor is disposed");const e=nt(this._model,{fuseConvPool:!1}),t=this._model.inputChannels/3,n=[{key:`pack/${this.precision}/${t}`,code:jn(this.precision,t)},...e.nodes.map(o=>this._nodePipelineSpec(o,o.id===e.spec.output))];await Promise.all(n.map(({key:o,code:r})=>this._pipelineAsync(o,r)))}_createExecution(e,t){const n=Yn(this._model,e,t,{fuseConvPool:!1}),o=new Map,r=[],s=new Map;n.nodes.forEach((l,p)=>{for(const c of Ar(l))s.set(c,p)}),s.set(n.spec.output,n.nodes.length);const a=[],u=l=>(a.push(l),this._resources.track("gpu-buffer",l));try{const l=(m,b,N,L)=>{for(const M of r)M.activeValue&&(s.get(M.activeValue)??-1)<L&&(M.activeValue=void 0);const q=vr(b,N);let R=r.filter(M=>!M.activeValue&&M.capacity>=q).sort((M,H)=>M.capacity-H.capacity)[0];R||(R={buffer:u(this._device.createBuffer({label:`oidn/activation/${e}x${t}/${r.length}`,size:Ut(q,4),usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST})),capacity:q},r.push(R)),R.activeValue=m,o.set(m,R.buffer)};l(n.spec.input,n.inputShape,this.precision==="fp16"?2:4,-1),n.plannedNodes.forEach(({node:m,outputShape:b},N)=>{const L=m.id===n.spec.output;l(m.id,b,L||this.precision==="fp32"?4:2,N)});const p=this._model.inputChannels/3,c=`pack/${this.precision}/${p}`,f=this._pipeline(c,jn(this.precision,p)),d=[],h=[],_=[],y=[];n.plannedNodes.forEach(({node:m,outputShape:b},N)=>{const L=m.id===n.spec.output,q=this._nodePipelineSpec(m,L),R=this._pipeline(q.key,q.code);d.push(R),h.push(q.kernel??"direct");const M=o.get(m.id);let H,I,A;if(m.op==="conv2d"){const z=n.valueShapes.get(m.input);A=m.id,I=[z.width,z.height,b.width,b.height,U(z.channels),U(b.channels)];const O=this._packedConvs.get(A);H=[{binding:0,resource:{buffer:o.get(m.input)}},{binding:1,resource:{buffer:O.weights}},{binding:2,resource:{buffer:O.bias}},{binding:3,resource:{buffer:M}}]}else if(m.op==="maxPool2d"){const z=n.valueShapes.get(m.input);I=[z.width,z.height,b.width,b.height,U(b.channels)],H=[{binding:0,resource:{buffer:o.get(m.input)}},{binding:1,resource:{buffer:M}}]}else if(m.op==="fusedConvReluMaxPool2d"){const z=n.valueShapes.get(m.input);A=m.conv.id,I=[z.width,z.height,b.width,b.height,U(z.channels),U(b.channels)];const O=this._packedConvs.get(A);H=[{binding:0,resource:{buffer:o.get(m.input)}},{binding:1,resource:{buffer:O.weights}},{binding:2,resource:{buffer:O.bias}},{binding:3,resource:{buffer:M}}]}else if(m.op==="fusedUpsampleConcatConv2d"){A=m.conv.id;const z=n.valueShapes.get(m.inputs[0].value),O=n.valueShapes.get(m.inputs[1].value);I=[b.width,b.height,U(b.channels),U(this._model.convChannels.get(A).inputChannels),z.width,z.height,O.width,O.height];const G=this._packedConvs.get(A);H=[{binding:0,resource:{buffer:o.get(m.inputs[0].value)}},{binding:1,resource:{buffer:o.get(m.inputs[1].value)}},{binding:2,resource:{buffer:G.weights}},{binding:3,resource:{buffer:G.bias}},{binding:4,resource:{buffer:M}}]}else throw new Error(`Unexpected native node ${m.op}`);const pe=u(Kn(this._device,`oidn/${m.id}/params/${e}x${t}`,I));y.push(pe),H.push({binding:H.length,resource:{buffer:pe}}),_.push(this._device.createBindGroup({label:`oidn/${m.id}/bindings`,layout:R.getBindGroupLayout(0),entries:H}))});const x=u(Kn(this._device,`oidn/input/params/${e}x${t}`,[e,t,U(this._model.inputChannels),this._model.inputChannels]));return y.push(x),{plan:n,valueBuffers:o,slots:r,nodeBindings:_,nodePipelines:d,nodeKernels:h,inputPipeline:f,inputUniform:x,ownedBuffers:y,lastUsed:++this._clock}}catch(l){for(const p of a)this._releaseBuffer(p);throw l}}_execution(e,t){if(this._disposed)throw new Error("Native OIDN executor is disposed");const n=`${e}x${t}`;let o=this._executionCache.get(n);if(!o&&(o=this._createExecution(e,t),this._executionCache.set(n,o),this._executionCache.size>this._shapeCacheSize)){const r=[...this._executionCache.entries()].filter(([s])=>s!==n).sort((s,a)=>s[1].lastUsed-a[1].lastUsed)[0];r&&(this._executionCache.delete(r[0]),this._retiredExecutions.add(r[1]),this._device.queue.onSubmittedWorkDone().catch(()=>{}).then(()=>{this._retiredExecutions.delete(r[1]),this._destroyExecution(r[1])}))}return o.lastUsed=++this._clock,o}profileNextExecution(){return this._device.features.has("timestamp-query")?(this._profileNextExecution=!0,!0):!1}getLastExecutionProfile(){return this._lastExecutionProfile}execute(e,t,n){const o=this._model.inputChannels/3;if(e.length!==o)throw new Error(`Native OIDN expected ${o} input buffers, got ${e.length}`);const r=this._execution(t,n),s=["input-pack",...r.plan.nodes.map(h=>h.id)],a=this._profileNextExecution&&this._device.features.has("timestamp-query");this._profileNextExecution=!1;const u=s.length*2,l=a?this._resources.track("gpu-query-set",this._device.createQuerySet({type:"timestamp",count:u})):void 0,p=u*8;let c,f;try{c=a?this._resources.track("gpu-buffer",this._device.createBuffer({label:`oidn/profile/resolve/${t}x${n}`,size:p,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC})):void 0,f=a?this._resources.track("gpu-buffer",this._device.createBuffer({label:`oidn/profile/readback/${t}x${n}`,size:p,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ})):void 0}catch(h){throw this._releaseBuffer(f),this._releaseBuffer(c),this._releaseQuerySet(l),h}const d=(h,_)=>({label:h,...l?{timestampWrites:{querySet:l,beginningOfPassWriteIndex:_*2,endOfPassWriteIndex:_*2+1}}:{}});try{const h=this._device.createCommandEncoder({label:`oidn/native/${t}x${n}`}),_=e.map((x,m)=>({binding:m,resource:{buffer:x}}));_.push({binding:o,resource:{buffer:r.valueBuffers.get(r.plan.spec.input)}}),_.push({binding:o+1,resource:{buffer:r.inputUniform}});const y=this._device.createBindGroup({label:"oidn/input/bindings",layout:r.inputPipeline.getBindGroupLayout(0),entries:_});{const x=h.beginComputePass(d("oidn/input-pack",0));x.setPipeline(r.inputPipeline),x.setBindGroup(0,y),x.dispatchWorkgroups(Math.ceil(t/Y),Math.ceil(n/Y),U(this._model.inputChannels)),x.end()}r.plan.plannedNodes.forEach(({node:x,outputShape:m},b)=>{const N=h.beginComputePass(d(`oidn/${r.plan.nodes[b].id}`,b+1));N.setPipeline(r.nodePipelines[b]),N.setBindGroup(0,r.nodeBindings[b]),r.nodeKernels[b]==="implicit-gemm"?N.dispatchWorkgroups(Math.ceil(m.width*m.height/we),Math.ceil(U(m.channels)/Q),1):N.dispatchWorkgroups(Math.ceil(m.width/Y),Math.ceil(m.height/Y),U(m.channels)),N.end()}),l&&(h.resolveQuerySet(l,0,u,c,0),h.copyBufferToBuffer(c,0,f,0,p)),this._device.queue.submit([h.finish()])}catch(h){throw this._releaseBuffer(f),this._releaseBuffer(c),this._releaseQuerySet(l),h}return l&&(this._profileOperations++,this._lastExecutionProfile=(async()=>{try{await f.mapAsync(GPUMapMode.READ);const h=new BigUint64Array(f.getMappedRange()),_=s.map((y,x)=>({id:y,durationMs:Number(h[x*2+1]-h[x*2])/1e6}));return{totalMs:_.reduce((y,x)=>y+x.durationMs,0),layers:_}}finally{f.mapState==="mapped"&&f.unmap(),this._releaseQuerySet(l),this._releaseBuffer(c),this._releaseBuffer(f),this._profileOperations--}})()),r.valueBuffers.get(r.plan.spec.output)}async executeCPU(e,t,n){const o=t*n*this._model.inputChannels;if(e.length!==o)throw new Error(`Native OIDN CPU input has ${e.length} values, expected ${o}`);const r=this._execution(t,n),s=this._model.inputChannels/3,a=t*n;r.cpuInputBuffers||(r.cpuInputBuffers=Array.from({length:s},(f,d)=>{const h=this._device.createBuffer({label:`oidn/cpu-input/${t}x${n}/${d}`,size:a*16,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST});return this._resources.track("gpu-buffer",h),r.ownedBuffers.push(h),h}),r.cpuReadbackBuffer=this._device.createBuffer({label:`oidn/cpu-readback/${t}x${n}`,size:a*16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ}),this._resources.track("gpu-buffer",r.cpuReadbackBuffer),r.ownedBuffers.push(r.cpuReadbackBuffer));for(let f=0;f<s;f++){const d=new Float32Array(a*4);for(let h=0;h<a;h++){const _=h*this._model.inputChannels+f*3,y=h*4;d[y]=e[_],d[y+1]=e[_+1],d[y+2]=e[_+2]}this._device.queue.writeBuffer(r.cpuInputBuffers[f],0,d)}const u=this.execute(r.cpuInputBuffers,t,n),l=this._device.createCommandEncoder({label:`oidn/cpu-readback/${t}x${n}`});l.copyBufferToBuffer(u,0,r.cpuReadbackBuffer,0,a*16),this._device.queue.submit([l.finish()]),await r.cpuReadbackBuffer.mapAsync(GPUMapMode.READ);const p=new Float32Array(r.cpuReadbackBuffer.getMappedRange()),c=new Float32Array(a*3);for(let f=0;f<a;f++)c[f*3]=p[f*4],c[f*3+1]=p[f*4+1],c[f*3+2]=p[f*4+2];return r.cpuReadbackBuffer.unmap(),c}_releaseBuffer(e){this._resources.release("gpu-buffer",e,()=>e.destroy())}_releaseQuerySet(e){this._resources.release("gpu-query-set",e,()=>e.destroy())}_destroyExecution(e){e.slots.forEach(t=>this._releaseBuffer(t.buffer)),e.ownedBuffers.forEach(t=>this._releaseBuffer(t))}getResourceInfo(){return this._resources.snapshot(this._retiredExecutions.size+this._profileOperations)}dispose(){if(!this._disposed){this._disposed=!0;for(const e of this._packedConvs.values())this._releaseBuffer(e.weights),this._releaseBuffer(e.bias);this._packedConvs.clear();for(const e of this._executionCache.values())this._destroyExecution(e);this._executionCache.clear();for(const e of this._retiredExecutions)this._destroyExecution(e);this._retiredExecutions.clear()}}}const Qn=new WeakMap;function Or(i,e){let t=Qn.get(i);t||(t=new Map,Qn.set(i,t));let n=t.get(e);if(!n){const o=i.createShaderModule({label:`oidn/webnn/input-pack/${e}`,code:Mr(e)}),r=i.createShaderModule({label:"oidn/webnn/output-unpack",code:zr()});n={input:i.createComputePipeline({label:`oidn/webnn/input-pack/${e}`,layout:"auto",compute:{module:o,entryPoint:"main"}}),output:i.createComputePipeline({label:"oidn/webnn/output-unpack",layout:"auto",compute:{module:r,entryPoint:"main"}})},t.set(e,n)}return n}const _e=8;function ei(i,e){return Math.ceil(i/e)*e}function Ur(i,e,t,n){const o=i.createBuffer({label:e,size:ei(t.byteLength,4),usage:n,mappedAtCreation:!0});return new Uint8Array(o.getMappedRange()).set(new Uint8Array(t.buffer,t.byteOffset,t.byteLength)),o.unmap(),o}function ti(i,e,t){const n=new Uint32Array(ei(t.length,4));return n.set(t),Ur(i,e,n,GPUBufferUsage.UNIFORM)}function Mt(i,e,t,n){var o,r,s,a;return!!((a=(s=(r=(o=i==null?void 0:i[e])==null?void 0:o[t])==null?void 0:r.dataTypes)==null?void 0:s.includes)!=null&&a.call(s,n))}function Nr(i,e){if(e==="fp16"&&i.desc.dataType==="Float16")return new Uint8Array(i.data.buffer,i.data.byteOffset,i.data.byteLength);if(e==="fp32"&&i.desc.dataType==="Float32")return new Uint8Array(i.data.buffer,i.data.byteOffset,i.data.byteLength);const t=i.desc.dataType==="Float32"?new Float32Array(i.data.buffer,i.data.byteOffset,i.data.byteLength/4):new P(i.data.buffer,i.data.byteOffset,i.data.byteLength/2),n=e==="fp16"?new P(t):new Float32Array(t);return new Uint8Array(n.buffer,n.byteOffset,n.byteLength)}function Mr(i){const e=Array.from({length:i},(n,o)=>`@group(0) @binding(${o}) var<storage, read> input${o}: array<vec4<f32>>;`).join(`
|
|
851
|
-
`),t=Array.from({length:
|
|
852
|
-
return input${
|
|
1066
|
+
`}function wo(n){return n.op==="concat"?n.inputs:n.op==="fusedUpsampleConcatConv2d"?n.inputs.map(e=>e.value):[n.input]}function wn(n,e="auto"){const t=n.features.has("shader-f16");if(e==="fp16"&&!t)throw new Error("OIDN FP16 was requested but the GPUDevice does not have shader-f16 enabled");return e==="auto"?t?"fp16":"fp32":e}class vn{constructor(e,t,i={}){m(this,"precision");m(this,"kernelSetting");m(this,"gemm");m(this,"maxSpatialInputBlocks");m(this,"subgroupsAvailable");m(this,"_model");m(this,"_gemmByOutputBlocks",new Map);m(this,"_packedConvs",new Map);m(this,"_pipelineCache");m(this,"_pipelinePromises");m(this,"_executionCache",new Map);m(this,"_retiredExecutions",new Set);m(this,"_clock",0);m(this,"_shapeCacheSize");m(this,"_profileNextExecution",!1);m(this,"_lastExecutionProfile");m(this,"_profileOperations",0);m(this,"_resources",new un);m(this,"_disposed",!1);var s,l,p,c,d,h,f,g,y,v,_,w;this._device=e;const r=ro(e);this._pipelineCache=r.ready,this._pipelinePromises=r.pending,this.precision=wn(e,i.precision??"auto"),this.kernelSetting=i.kernel??"auto";const o=((s=i.gemm)==null?void 0:s.workgroupSize)??[8,8];if(this.gemm=Object.freeze({tilePolicy:((l=i.gemm)==null?void 0:l.tilePolicy)??((p=i.gemm)!=null&&p.workgroupSize?"fixed":"output-aligned"),decoderLoad:((c=i.gemm)==null?void 0:c.decoderLoad)??"per-load",poolLayout:((d=i.gemm)==null?void 0:d.poolLayout)??"channels",sharedLayout:((h=i.gemm)==null?void 0:h.sharedLayout)??(this.precision==="fp16"?"padded-input":"padded"),accumulationOrder:((f=i.gemm)==null?void 0:f.accumulationOrder)??"k-major",finalLayer:((g=i.gemm)==null?void 0:g.finalLayer)??"shared-auto",loadMode:((y=i.gemm)==null?void 0:y.loadMode)??"native",addressMode:((v=i.gemm)==null?void 0:v.addressMode)??"incremental",weightLayout:((_=i.gemm)==null?void 0:_.weightLayout)??"k-major",rowsPerThread:((w=i.gemm)==null?void 0:w.rowsPerThread)??8,workgroupSize:Object.freeze([o[0],o[1]])}),!["fixed","output-aligned"].includes(this.gemm.tilePolicy))throw new Error(`Unsupported GEMM tile policy: ${this.gemm.tilePolicy}`);if(!["per-load","source-first"].includes(this.gemm.decoderLoad))throw new Error(`Unsupported GEMM decoder load: ${this.gemm.decoderLoad}`);if(!["spatial","channels"].includes(this.gemm.poolLayout))throw new Error(`Unsupported GEMM pool layout: ${this.gemm.poolLayout}`);if(!["linear","padded","padded-input","padded-weights"].includes(this.gemm.sharedLayout))throw new Error(`Unsupported GEMM shared layout: ${this.gemm.sharedLayout}`);if(!["k-major","row-major"].includes(this.gemm.accumulationOrder))throw new Error(`Unsupported GEMM accumulation order: ${this.gemm.accumulationOrder}`);if(!["direct","shared-input","shared-input-weights","shared-auto"].includes(this.gemm.finalLayer))throw new Error(`Unsupported GEMM final layer: ${this.gemm.finalLayer}`);if(!["native","packed-weights","packed-all"].includes(this.gemm.loadMode))throw new Error(`Unsupported GEMM load mode: ${this.gemm.loadMode}`);if(!["analytic","incremental","base-offset"].includes(this.gemm.addressMode))throw new Error(`Unsupported GEMM address mode: ${this.gemm.addressMode}`);if(!["output-major","k-major"].includes(this.gemm.weightLayout))throw new Error(`Unsupported GEMM weight layout: ${this.gemm.weightLayout}`);const a=re(this.gemm),u=(ye(this.gemm).input+ye(this.gemm).weights)*4*(this.precision==="fp16"?2:4);if(![2,4,8].includes(a.rowsPerThread)||![4,8,16].includes(a.workgroupX)||![4,8].includes(a.workgroupY)||o.length!==2||a.workgroupX>e.limits.maxComputeWorkgroupSizeX||a.workgroupY>e.limits.maxComputeWorkgroupSizeY||a.workgroupX*a.workgroupY>e.limits.maxComputeInvocationsPerWorkgroup||u>e.limits.maxComputeWorkgroupStorageSize)throw new Error("Unsupported GEMM tile configuration for this GPUDevice");if(this.subgroupsAvailable=e.features.has("subgroups"),this.maxSpatialInputBlocks=this.precision==="fp16"&&e.limits.maxComputeInvocationsPerWorkgroup>=U*U&&e.limits.maxComputeWorkgroupSizeX>=U&&e.limits.maxComputeWorkgroupSizeY>=U?Math.floor(e.limits.maxComputeWorkgroupStorageSize/(J*J*4*2)):0,this._shapeCacheSize=Math.max(1,i.shapeCacheSize??2),t.inputChannels%3!==0||t.inputChannels<3||t.inputChannels>9)throw new Error(`Native OIDN expects 3, 6, or 9 input channels, got ${t.inputChannels}`);this._model={spec:t.spec,inputChannels:t.inputChannels,outputChannels:t.outputChannels,channelsByValue:new Map(t.channelsByValue),convChannels:new Map(t.convChannels)};for(const b of rt(this._model,{fuseConvPool:!1}).nodes)if(b.op!=="conv2d"&&b.op!=="maxPool2d"&&b.op!=="fusedConvReluMaxPool2d"&&b.op!=="fusedUpsampleConcatConv2d")throw new Error(`Native OIDN descriptor ${t.spec.id} leaves unsupported ${b.op} node ${b.id} after graph optimization`);try{for(const[b,x]of t.convTensors){const k=this._selectConvKernel(N(x.inputChannels),b===t.spec.output)==="implicit-gemm"?this.gemm.weightLayout:"output-major",C=so(e,b,x,this.precision,k);this._resources.track("gpu-buffer",C.weights),this._resources.track("gpu-buffer",C.bias),this._packedConvs.set(b,C)}}catch(b){for(const x of this._packedConvs.values())this._releaseBuffer(x.weights),this._releaseBuffer(x.bias);throw this._packedConvs.clear(),b}}_pipeline(e,t){let i=this._pipelineCache.get(e);return i||(i=this._device.createComputePipeline({label:`oidn/${e}`,layout:"auto",compute:{module:this._device.createShaderModule({label:`oidn/${e}`,code:t}),entryPoint:"main"}}),this._pipelineCache.set(e,i)),i}_pipelineAsync(e,t){const i=this._pipelineCache.get(e);if(i)return Promise.resolve(i);const r=this._pipelinePromises.get(e);if(r)return r;const o=this._device.createComputePipelineAsync({label:`oidn/${e}`,layout:"auto",compute:{module:this._device.createShaderModule({label:`oidn/${e}`,code:t}),entryPoint:"main"}}).then(a=>(this._pipelineCache.set(e,a),this._pipelinePromises.delete(e),a),a=>{throw this._pipelinePromises.delete(e),a});return this._pipelinePromises.set(e,o),o}_nodePipelineSpec(e,t){if(e.op==="conv2d"){const i=t?"fp32":this.precision,r=N(this._model.convChannels.get(e.id).inputChannels),o=N(this._model.convChannels.get(e.id).outputChannels),a=this._gemmForOutput(o),u=this._selectConvKernel(r,t),s=a.finalLayer==="shared-input-weights"||a.finalLayer==="shared-auto"&&on(this.precision,r,!0)<=this._device.limits.maxComputeWorkgroupStorageSize;return t&&this._model.convChannels.get(e.id).outputChannels===3&&(this.kernelSetting==="auto"||this.kernelSetting==="implicit-gemm")&&a.finalLayer!=="direct"&&this._device.limits.maxComputeWorkgroupSizeX>=8&&this._device.limits.maxComputeWorkgroupSizeY>=8&&this._device.limits.maxComputeInvocationsPerWorkgroup>=64&&on(this.precision,r,s)<=this._device.limits.maxComputeWorkgroupStorageSize?{key:`conv-final-rgb/${this.precision}/${e.activation}/in${r}/weights-${s?"shared":"storage"}`,kernel:"direct",code:eo(this.precision,e.activation,r,s)}:{key:`conv-${u}/${this.precision}/${i}/${e.activation}/in${r}/out${o}`+(u==="implicit-gemm"?`/address-${a.addressMode}/weights-${a.weightLayout}-v1/tile-${a.workgroupSize.join("x")}-r${a.rowsPerThread}/loads-${a.loadMode}/shared-${a.sharedLayout}/acc-${a.accumulationOrder}`:""),kernel:u,code:u==="implicit-gemm"?ho(this.precision,i,e.activation,r,o,a):u==="spatial"?po(this.precision,i,e.activation,r,o):u==="subgroup"?lo(this.precision,i,e.activation,r,o):co(this.precision,i,e.activation,r,o)}}if(e.op==="maxPool2d"){const i=N(this._model.channelsByValue.get(e.id)),r=this._coalescedPool();return{key:`max-pool/${this.precision}/out${i}/${r?"channels":"spatial"}`,kernel:"direct",code:go(this.precision,i,r)}}if(e.op==="fusedConvReluMaxPool2d"){const i=N(this._model.convChannels.get(e.conv.id).inputChannels),r=N(this._model.convChannels.get(e.conv.id).outputChannels);return{key:`conv-pool/${this.precision}/${e.conv.activation}/in${i}/out${r}`,kernel:"direct",code:fo(this.precision,e.conv.activation,i,r)}}if(e.op==="fusedUpsampleConcatConv2d"){if(e.inputs.length!==2)throw new Error(`Native fused decoder ${e.id} requires two inputs`);const i=e.inputs.map(d=>N(this._model.channelsByValue.get(d.value))),r=e.inputs.findIndex(d=>d.upsample);if(r!==0&&r!==1)throw new Error(`Native fused decoder ${e.id} has no upsample input`);const o=i[0]+i[1],a=t?"fp32":this.precision,u=this._selectConvKernel(o,t),s=u==="subgroup"?"direct":u,l=N(this._model.convChannels.get(e.conv.id).outputChannels),p=this._gemmForOutput(l);return{key:`decoder-${s}/${this.precision}/${a}/${e.conv.activation}/${i.join("+")}/out${l}/up${r}`+(s==="implicit-gemm"?`/address-${p.addressMode}/weights-${p.weightLayout}-v1/tile-${p.workgroupSize.join("x")}-r${p.rowsPerThread}/loads-${p.loadMode}/shared-${p.sharedLayout}/acc-${p.accumulationOrder}/decoder-${p.decoderLoad}`:""),kernel:s,code:s==="implicit-gemm"?mo(this.precision,a,e.conv.activation,i,r,l,p):s==="spatial"?yo(this.precision,a,e.conv.activation,i,r,l):_o(this.precision,a,e.conv.activation,i,r,l)}}throw new Error(`Native OIDN does not implement unfused ${e.op} node ${e.id}`)}_gemmForOutput(e){if(this.gemm.tilePolicy!=="output-aligned"||this.gemm.workgroupSize[0]!==8||e<=0||e%16!==0)return this.gemm;const t=this._gemmByOutputBlocks.get(e);if(t)return t;const i=Object.freeze({...this.gemm,workgroupSize:Object.freeze([16,this.gemm.workgroupSize[1]])}),r=re(i),o=ye(i),a=this._device.limits,s=r.workgroupX<=a.maxComputeWorkgroupSizeX&&r.workgroupY<=a.maxComputeWorkgroupSizeY&&r.workgroupX*r.workgroupY<=a.maxComputeInvocationsPerWorkgroup&&(o.input+o.weights)*4*(this.precision==="fp16"?2:4)<=a.maxComputeWorkgroupStorageSize?i:this.gemm;return this._gemmByOutputBlocks.set(e,s),s}_coalescedPool(){return this.gemm.poolLayout==="channels"&&(this.kernelSetting==="auto"||this.kernelSetting==="implicit-gemm")}_selectConvKernel(e,t){const i=this.precision==="fp16"&&e<=this.maxSpatialInputBlocks;return this.kernelSetting==="direct"?"direct":this.kernelSetting==="spatial"?i?"spatial":"direct":this.kernelSetting==="implicit-gemm"?t?"direct":"implicit-gemm":this.kernelSetting==="subgroup"?this.subgroupsAvailable?"subgroup":"direct":t?"direct":"implicit-gemm"}_nodePipeline(e,t){const{key:i,code:r}=this._nodePipelineSpec(e,t);return this._pipeline(i,r)}async prepare(){if(this._disposed)throw new Error("Native OIDN executor is disposed");const e=rt(this._model,{fuseConvPool:!1}),t=this._model.inputChannels/3,i=[{key:`pack/${this.precision}/${t}`,code:yn(this.precision,t)},...e.nodes.map(r=>this._nodePipelineSpec(r,r.id===e.spec.output))];await Promise.all(i.map(({key:r,code:o})=>this._pipelineAsync(r,o)))}_createExecution(e,t){const i=sn(this._model,e,t,{fuseConvPool:!1}),r=new Map,o=[],a=new Map;i.nodes.forEach((l,p)=>{for(const c of wo(l))a.set(c,p)}),a.set(i.spec.output,i.nodes.length);const u=[],s=l=>(u.push(l),this._resources.track("gpu-buffer",l));try{const l=(_,w,b,x)=>{for(const C of o)C.activeValue&&(a.get(C.activeValue)??-1)<x&&(C.activeValue=void 0);const B=oo(w,b);let k=o.filter(C=>!C.activeValue&&C.capacity>=B).sort((C,L)=>C.capacity-L.capacity)[0];k||(k={buffer:s(this._device.createBuffer({label:`oidn/activation/${e}x${t}/${o.length}`,size:Mt(B,4),usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST})),capacity:B},o.push(k)),k.activeValue=_,r.set(_,k.buffer)};l(i.spec.input,i.inputShape,this.precision==="fp16"?2:4,-1),i.plannedNodes.forEach(({node:_,outputShape:w},b)=>{const x=_.id===i.spec.output;l(_.id,w,x||this.precision==="fp32"?4:2,b)});const p=this._model.inputChannels/3,c=`pack/${this.precision}/${p}`,d=this._pipeline(c,yn(this.precision,p)),h=[],f=[],g=[],y=[];i.plannedNodes.forEach(({node:_,outputShape:w},b)=>{const x=_.id===i.spec.output,B=this._nodePipelineSpec(_,x),k=this._pipeline(B.key,B.code);h.push(k),f.push(B.kernel??"direct");const C=r.get(_.id);let L,F,z;if(_.op==="conv2d"){const X=i.valueShapes.get(_.input);z=_.id,F=[X.width,X.height,w.width,w.height,N(X.channels),N(w.channels)];const Z=this._packedConvs.get(z);L=[{binding:0,resource:{buffer:r.get(_.input)}},{binding:1,resource:{buffer:Z.weights}},{binding:2,resource:{buffer:Z.bias}},{binding:3,resource:{buffer:C}}]}else if(_.op==="maxPool2d"){const X=i.valueShapes.get(_.input);F=[X.width,X.height,w.width,w.height,N(w.channels)],L=[{binding:0,resource:{buffer:r.get(_.input)}},{binding:1,resource:{buffer:C}}]}else if(_.op==="fusedConvReluMaxPool2d"){const X=i.valueShapes.get(_.input);z=_.conv.id,F=[X.width,X.height,w.width,w.height,N(X.channels),N(w.channels)];const Z=this._packedConvs.get(z);L=[{binding:0,resource:{buffer:r.get(_.input)}},{binding:1,resource:{buffer:Z.weights}},{binding:2,resource:{buffer:Z.bias}},{binding:3,resource:{buffer:C}}]}else if(_.op==="fusedUpsampleConcatConv2d"){z=_.conv.id;const X=i.valueShapes.get(_.inputs[0].value),Z=i.valueShapes.get(_.inputs[1].value);F=[w.width,w.height,N(w.channels),N(this._model.convChannels.get(z).inputChannels),X.width,X.height,Z.width,Z.height];const Le=this._packedConvs.get(z);L=[{binding:0,resource:{buffer:r.get(_.inputs[0].value)}},{binding:1,resource:{buffer:r.get(_.inputs[1].value)}},{binding:2,resource:{buffer:Le.weights}},{binding:3,resource:{buffer:Le.bias}},{binding:4,resource:{buffer:C}}]}else throw new Error(`Unexpected native node ${_.op}`);const Pe=s(fn(this._device,`oidn/${_.id}/params/${e}x${t}`,F));y.push(Pe),L.push({binding:L.length,resource:{buffer:Pe}}),g.push(this._device.createBindGroup({label:`oidn/${_.id}/bindings`,layout:k.getBindGroupLayout(0),entries:L}))});const v=s(fn(this._device,`oidn/input/params/${e}x${t}`,[e,t,N(this._model.inputChannels),this._model.inputChannels]));return y.push(v),{plan:i,valueBuffers:r,slots:o,nodeBindings:g,nodePipelines:h,nodeKernels:f,inputPipeline:d,inputUniform:v,ownedBuffers:y,lastUsed:++this._clock}}catch(l){for(const p of u)this._releaseBuffer(p);throw l}}_execution(e,t){if(this._disposed)throw new Error("Native OIDN executor is disposed");const i=`${e}x${t}`;let r=this._executionCache.get(i);if(!r&&(r=this._createExecution(e,t),this._executionCache.set(i,r),this._executionCache.size>this._shapeCacheSize)){const o=[...this._executionCache.entries()].filter(([a])=>a!==i).sort((a,u)=>a[1].lastUsed-u[1].lastUsed)[0];o&&(this._executionCache.delete(o[0]),this._retiredExecutions.add(o[1]),this._device.queue.onSubmittedWorkDone().catch(()=>{}).then(()=>{this._retiredExecutions.delete(o[1]),this._destroyExecution(o[1])}))}return r.lastUsed=++this._clock,r}prewarm(e){for(const t of e)this._execution(t.width,t.height)}profileNextExecution(){return this._device.features.has("timestamp-query")?(this._profileNextExecution=!0,!0):!1}getLastExecutionProfile(){return this._lastExecutionProfile}execute(e,t,i){const r=this._model.inputChannels/3;if(e.length!==r)throw new Error(`Native OIDN expected ${r} input buffers, got ${e.length}`);const o=this._execution(t,i),a=["input-pack",...o.plan.nodes.map(f=>f.id)],u=this._profileNextExecution&&this._device.features.has("timestamp-query");this._profileNextExecution=!1;const s=a.length*2,l=u?this._resources.track("gpu-query-set",this._device.createQuerySet({type:"timestamp",count:s})):void 0,p=s*8;let c,d;try{c=u?this._resources.track("gpu-buffer",this._device.createBuffer({label:`oidn/profile/resolve/${t}x${i}`,size:p,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC})):void 0,d=u?this._resources.track("gpu-buffer",this._device.createBuffer({label:`oidn/profile/readback/${t}x${i}`,size:p,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ})):void 0}catch(f){throw this._releaseBuffer(d),this._releaseBuffer(c),this._releaseQuerySet(l),f}const h=(f,g)=>({label:f,...l?{timestampWrites:{querySet:l,beginningOfPassWriteIndex:g*2,endOfPassWriteIndex:g*2+1}}:{}});try{const f=this._device.createCommandEncoder({label:`oidn/native/${t}x${i}`}),g=e.map((v,_)=>({binding:_,resource:{buffer:v}}));g.push({binding:r,resource:{buffer:o.valueBuffers.get(o.plan.spec.input)}}),g.push({binding:r+1,resource:{buffer:o.inputUniform}});const y=this._device.createBindGroup({label:"oidn/input/bindings",layout:o.inputPipeline.getBindGroupLayout(0),entries:g});{const v=f.beginComputePass(h("oidn/input-pack",0));v.setPipeline(o.inputPipeline),v.setBindGroup(0,y),v.dispatchWorkgroups(Math.ceil(t/R),Math.ceil(i/R),N(this._model.inputChannels)),v.end()}o.plan.plannedNodes.forEach(({node:v,outputShape:_},w)=>{const b=f.beginComputePass(h(`oidn/${o.plan.nodes[w].id}`,w+1));if(b.setPipeline(o.nodePipelines[w]),b.setBindGroup(0,o.nodeBindings[w]),v.op==="maxPool2d"&&this._coalescedPool()){const x=Math.ceil(_.width*N(_.channels)/R),B=this._device.limits.maxComputeWorkgroupsPerDimension;b.dispatchWorkgroups(Math.min(x,B),Math.ceil(_.height/R),Math.ceil(x/B))}else if(o.nodeKernels[w]==="implicit-gemm"){const x=re(this._gemmForOutput(N(_.channels)));b.dispatchWorkgroups(Math.ceil(_.width*_.height/x.tileM),Math.ceil(N(_.channels)/x.tileNBlocks),1)}else b.dispatchWorkgroups(Math.ceil(_.width/R),Math.ceil(_.height/R),N(_.channels));b.end()}),l&&(f.resolveQuerySet(l,0,s,c,0),f.copyBufferToBuffer(c,0,d,0,p)),this._device.queue.submit([f.finish()])}catch(f){throw this._releaseBuffer(d),this._releaseBuffer(c),this._releaseQuerySet(l),f}return l&&(this._profileOperations++,this._lastExecutionProfile=(async()=>{try{await d.mapAsync(GPUMapMode.READ);const f=new BigUint64Array(d.getMappedRange()),g=a.map((y,v)=>({id:y,durationMs:Number(f[v*2+1]-f[v*2])/1e6}));return{totalMs:g.reduce((y,v)=>y+v.durationMs,0),layers:g}}finally{d.mapState==="mapped"&&d.unmap(),this._releaseQuerySet(l),this._releaseBuffer(c),this._releaseBuffer(d),this._profileOperations--}})()),o.valueBuffers.get(o.plan.spec.output)}async executeCPU(e,t,i){const r=t*i*this._model.inputChannels;if(e.length!==r)throw new Error(`Native OIDN CPU input has ${e.length} values, expected ${r}`);const o=this._execution(t,i),a=this._model.inputChannels/3,u=t*i;o.cpuInputBuffers||(o.cpuInputBuffers=Array.from({length:a},(d,h)=>{const f=this._device.createBuffer({label:`oidn/cpu-input/${t}x${i}/${h}`,size:u*16,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST});return this._resources.track("gpu-buffer",f),o.ownedBuffers.push(f),f}),o.cpuReadbackBuffer=this._device.createBuffer({label:`oidn/cpu-readback/${t}x${i}`,size:u*16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ}),this._resources.track("gpu-buffer",o.cpuReadbackBuffer),o.ownedBuffers.push(o.cpuReadbackBuffer));for(let d=0;d<a;d++){const h=new Float32Array(u*4);for(let f=0;f<u;f++){const g=f*this._model.inputChannels+d*3,y=f*4;h[y]=e[g],h[y+1]=e[g+1],h[y+2]=e[g+2]}this._device.queue.writeBuffer(o.cpuInputBuffers[d],0,h)}const s=this.execute(o.cpuInputBuffers,t,i),l=this._device.createCommandEncoder({label:`oidn/cpu-readback/${t}x${i}`});l.copyBufferToBuffer(s,0,o.cpuReadbackBuffer,0,u*16),this._device.queue.submit([l.finish()]),await o.cpuReadbackBuffer.mapAsync(GPUMapMode.READ);const p=new Float32Array(o.cpuReadbackBuffer.getMappedRange()),c=new Float32Array(u*3);for(let d=0;d<u;d++)c[d*3]=p[d*4],c[d*3+1]=p[d*4+1],c[d*3+2]=p[d*4+2];return o.cpuReadbackBuffer.unmap(),c}_releaseBuffer(e){this._resources.release("gpu-buffer",e,()=>e.destroy())}_releaseQuerySet(e){this._resources.release("gpu-query-set",e,()=>e.destroy())}_destroyExecution(e){e.slots.forEach(t=>this._releaseBuffer(t.buffer)),e.ownedBuffers.forEach(t=>this._releaseBuffer(t))}getResourceInfo(){return this._resources.snapshot(this._retiredExecutions.size+this._profileOperations)}dispose(){if(!this._disposed){this._disposed=!0;for(const e of this._packedConvs.values())this._releaseBuffer(e.weights),this._releaseBuffer(e.bias);this._packedConvs.clear();for(const e of this._executionCache.values())this._destroyExecution(e);this._executionCache.clear();for(const e of this._retiredExecutions)this._destroyExecution(e);this._retiredExecutions.clear()}}}const xn=new WeakMap;function vo(n,e){let t=xn.get(n);t||(t=new Map,xn.set(n,t));let i=t.get(e);if(!i){const r=n.createShaderModule({label:`oidn/webnn/input-pack/${e}`,code:$o(e)}),o=n.createShaderModule({label:"oidn/webnn/output-unpack",code:ko()});i={input:n.createComputePipeline({label:`oidn/webnn/input-pack/${e}`,layout:"auto",compute:{module:r,entryPoint:"main"}}),output:n.createComputePipeline({label:"oidn/webnn/output-unpack",layout:"auto",compute:{module:o,entryPoint:"main"}})},t.set(e,i)}return i}const fe=8;function bn(n,e){return Math.ceil(n/e)*e}function xo(n,e,t,i){const r=n.createBuffer({label:e,size:bn(t.byteLength,4),usage:i,mappedAtCreation:!0});return new Uint8Array(r.getMappedRange()).set(new Uint8Array(t.buffer,t.byteOffset,t.byteLength)),r.unmap(),r}function $n(n,e,t){const i=new Uint32Array(bn(t.length,4));return i.set(t),xo(n,e,i,GPUBufferUsage.UNIFORM)}function Nt(n,e,t,i){var r,o,a,u;return!!((u=(a=(o=(r=n==null?void 0:n[e])==null?void 0:r[t])==null?void 0:o.dataTypes)==null?void 0:a.includes)!=null&&u.call(a,i))}function bo(n,e){if(e==="fp16"&&n.desc.dataType==="Float16")return new Uint8Array(n.data.buffer,n.data.byteOffset,n.data.byteLength);if(e==="fp32"&&n.desc.dataType==="Float32")return new Uint8Array(n.data.buffer,n.data.byteOffset,n.data.byteLength);const t=n.desc.dataType==="Float32"?new Float32Array(n.data.buffer,n.data.byteOffset,n.data.byteLength/4):new E(n.data.buffer,n.data.byteOffset,n.data.byteLength/2),i=e==="fp16"?new E(t):new Float32Array(t);return new Uint8Array(i.buffer,i.byteOffset,i.byteLength)}function $o(n){const e=Array.from({length:n},(i,r)=>`@group(0) @binding(${r}) var<storage, read> input${r}: array<vec4<f32>>;`).join(`
|
|
1067
|
+
`),t=Array.from({length:n},(i,r)=>{const o=r*3;return`if (channel < ${o+3}u) {
|
|
1068
|
+
return input${r}[pixel][channel - ${o}u];
|
|
853
1069
|
}`}).join(`
|
|
854
1070
|
`);return`enable f16;
|
|
855
1071
|
struct Params { width: u32, height: u32, channels: u32, padding: u32 }
|
|
856
1072
|
${e}
|
|
857
|
-
@group(0) @binding(${
|
|
858
|
-
@group(0) @binding(${
|
|
1073
|
+
@group(0) @binding(${n}) var<storage, read_write> outputData: array<f16>;
|
|
1074
|
+
@group(0) @binding(${n+1}) var<uniform> params: Params;
|
|
859
1075
|
|
|
860
1076
|
fn readChannel(pixel: u32, channel: u32) -> f32 {
|
|
861
1077
|
${t}
|
|
862
1078
|
return 0.0;
|
|
863
1079
|
}
|
|
864
1080
|
|
|
865
|
-
@compute @workgroup_size(${
|
|
1081
|
+
@compute @workgroup_size(${fe}, ${fe}, 1)
|
|
866
1082
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
867
1083
|
if (gid.x >= params.width || gid.y >= params.height || gid.z >= params.channels) {
|
|
868
1084
|
return;
|
|
@@ -871,13 +1087,13 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
871
1087
|
let outputIndex = (gid.z * params.height + gid.y) * params.width + gid.x;
|
|
872
1088
|
outputData[outputIndex] = f16(readChannel(pixel, gid.z));
|
|
873
1089
|
}
|
|
874
|
-
`}function
|
|
1090
|
+
`}function ko(){return`enable f16;
|
|
875
1091
|
struct Params { width: u32, height: u32, padding0: u32, padding1: u32 }
|
|
876
1092
|
@group(0) @binding(0) var<storage, read> inputData: array<f16>;
|
|
877
1093
|
@group(0) @binding(1) var<storage, read_write> outputData: array<vec4<f32>>;
|
|
878
1094
|
@group(0) @binding(2) var<uniform> params: Params;
|
|
879
1095
|
|
|
880
|
-
@compute @workgroup_size(${
|
|
1096
|
+
@compute @workgroup_size(${fe}, ${fe}, 1)
|
|
881
1097
|
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
882
1098
|
if (gid.x >= params.width || gid.y >= params.height) { return; }
|
|
883
1099
|
let pixel = gid.y * params.width + gid.x;
|
|
@@ -889,4 +1105,4 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
889
1105
|
0.0
|
|
890
1106
|
);
|
|
891
1107
|
}
|
|
892
|
-
`}function Dr(i,e){if(e==="fp32")throw new Error("OIDN WebNN GPU interop currently requires FP16 exportable tensors");if(!i.features.has("shader-f16"))throw new Error("OIDN WebNN requires shader-f16 on the shared GPUDevice");return"fp16"}function Wr(i,e,t){return t==="relu"?i.relu(e):e}class ni{constructor(e,t,n={}){g(this,"precision");g(this,"support");g(this,"_context");g(this,"_builderConstructor");g(this,"_shapeCache",new Map);g(this,"_shapePromises",new Map);g(this,"_retiredExecutions",new Set);g(this,"_pendingCreationCount",0);g(this,"_shapeCacheSize");g(this,"_clock",0);g(this,"_inputPipeline");g(this,"_outputPipeline");g(this,"_resources",new Fn);g(this,"_disposed",!1);this._device=e,this._model=t,this.precision=Dr(e,n.precision??"auto"),this._shapeCacheSize=Math.max(1,n.shapeCacheSize??2),this.support={available:!1,fp16Conv:!1,gpuInterop:!1};const o=Or(e,t.inputChannels/3);this._inputPipeline=o.input,this._outputPipeline=o.output}async prepare(){var n,o,r;if(this._disposed)throw new Error("OIDN WebNN executor is disposed");const e=(n=globalThis.navigator)==null?void 0:n.ml,t=globalThis.MLGraphBuilder;if(!(e!=null&&e.createContext)||typeof t!="function")throw this.support.reason="WebNN is not exposed by this browser",new Error(this.support.reason);this._builderConstructor=t;try{try{this._context=await e.createContext({deviceType:"gpu",powerPreference:"high-performance"})}catch{this._context=await e.createContext({deviceType:"gpu"})}if(this._resources.track("ml-context",this._context),this._disposed)throw new Error("OIDN WebNN executor is disposed");if(typeof this._context.createExportableTensor!="function"||typeof this._context.exportToGPU!="function")throw this.support.reason="WebNN WebGPU tensor interop is unavailable",new Error(this.support.reason);const s=((r=(o=this._context).opSupportLimits)==null?void 0:r.call(o))??{};if(this.support.fp16Conv=Mt(s,"conv2d","input","float16")&&Mt(s,"conv2d","filter","float16")&&Mt(s,"conv2d","output","float16"),!this.support.fp16Conv)throw this.support.reason="WebNN does not support FP16 conv2d",new Error(this.support.reason);let a,u;try{a=this._resources.track("ml-tensor",await this._context.createExportableTensor({dataType:"float16",shape:[4]},this._device)),u=this._resources.track("gpu-buffer",await this._context.exportToGPU(a)),this.support.gpuInterop=!0}catch(l){throw this.support.reason=`WebNN FP16 WebGPU interop failed: ${String(l)}`,new Error(this.support.reason)}finally{this._releaseBuffer(u),this._releaseTensor(a)}this.support.available=!0}catch(s){throw this._releaseContext(),s}}_constant(e,t){return e.constant({dataType:"float16",shape:[...t.desc.dims]},Nr(t,this.precision))}async _createExecution(e,t){if(this._disposed)throw new Error("OIDN WebNN executor is disposed");const n=new this._builderConstructor(this._context),o=new Map,r=new Map;o.set(this._model.spec.input,n.input("input",{dataType:"float16",shape:[1,this._model.inputChannels,t,e]})),r.set(this._model.spec.input,[this._model.inputChannels,t,e]);for(const f of this._model.spec.nodes){let d,h;if(f.op==="conv2d"){const _=r.get(f.input),y=this._model.convTensors.get(f.id),x=n.conv2d(o.get(f.input),this._constant(n,y.weight),{bias:this._constant(n,y.bias),padding:[1,1,1,1],inputLayout:"nchw",filterLayout:"oihw"});d=Wr(n,x,f.activation),h=[y.outputChannels,_[1],_[2]]}else if(f.op==="maxPool2d"){const _=r.get(f.input);d=n.maxPool2d(o.get(f.input),{windowDimensions:[2,2],strides:[2,2],padding:[0,_[1]%2,0,_[2]%2],layout:"nchw"}),h=[_[0],Math.ceil(_[1]/2),Math.ceil(_[2]/2)]}else if(f.op==="upsample2d"){const _=r.get(f.input);d=n.resample2d(o.get(f.input),{mode:"nearest-neighbor",axes:[2,3],scales:[2,2]}),h=[_[0],_[1]*2,_[2]*2]}else{const _=f.inputs.map(y=>r.get(y));if(_.some(y=>y[1]!==_[0][1]||y[2]!==_[0][2]))throw new Error(`WebNN concat ${f.id} has mismatched spatial shapes`);d=n.concat(f.inputs.map(y=>o.get(y)),1),h=[_.reduce((y,x)=>y+x[0],0),_[0][1],_[0][2]]}o.set(f.id,d),r.set(f.id,h)}let s,a,u,l,p,c;try{return s=this._resources.track("ml-graph",await n.build({output:o.get(this._model.spec.output)})),a=this._resources.track("ml-tensor",await this._context.createExportableTensor({dataType:"float16",shape:[1,this._model.inputChannels,t,e],writable:!0},this._device)),u=this._resources.track("ml-tensor",await this._context.createExportableTensor({dataType:"float16",shape:[1,this._model.outputChannels,t,e],readable:!0},this._device)),l=this._resources.track("gpu-buffer",this._device.createBuffer({label:`oidn/webnn/output/${e}x${t}`,size:e*t*4*4,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST})),p=this._resources.track("gpu-buffer",ti(this._device,`oidn/webnn/input/${e}x${t}`,[e,t,this._model.inputChannels])),c=this._resources.track("gpu-buffer",ti(this._device,`oidn/webnn/output/${e}x${t}`,[e,t])),{graph:s,inputTensor:a,outputTensor:u,outputBuffer:l,inputUniform:p,outputUniform:c,width:e,height:t,lastUsed:++this._clock}}catch(f){throw this._releaseBuffer(c),this._releaseBuffer(p),this._releaseBuffer(l),this._releaseTensor(u),this._releaseTensor(a),this._releaseGraph(s),f}}async _execution(e,t){const n=`${e}x${t}`;let o=this._shapeCache.get(n);if(!o){let r=this._shapePromises.get(n);r||(r=(async()=>{this._pendingCreationCount++;try{return await this._createExecution(e,t)}finally{this._pendingCreationCount--}})(),this._shapePromises.set(n,r));try{if(o=await r,this._disposed)throw this._destroyExecution(o),new Error("OIDN WebNN executor is disposed");this._shapeCache.set(n,o)}finally{this._shapePromises.get(n)===r&&this._shapePromises.delete(n)}if(this._shapeCache.size>this._shapeCacheSize){const s=[...this._shapeCache.entries()].filter(([a])=>a!==n).sort((a,u)=>a[1].lastUsed-u[1].lastUsed)[0];s&&(this._shapeCache.delete(s[0]),this._retireExecution(s[1]))}}return o.lastUsed=++this._clock,o}async prewarm(e){for(const t of e)await this._execution(t.width,t.height)}async execute(e,t,n){const o=this._model.inputChannels/3;if(e.length!==o)throw new Error(`OIDN WebNN expected ${o} input buffers, got ${e.length}`);const r=await this._execution(t,n),s=this._resources.track("gpu-buffer",await this._context.exportToGPU(r.inputTensor));try{const u=e.map((f,d)=>({binding:d,resource:{buffer:f}}));u.push({binding:o,resource:{buffer:s}}),u.push({binding:o+1,resource:{buffer:r.inputUniform}});const l=this._device.createBindGroup({label:"oidn/webnn/input-bindings",layout:this._inputPipeline.getBindGroupLayout(0),entries:u}),p=this._device.createCommandEncoder({label:"oidn/webnn/input-pack"}),c=p.beginComputePass();c.setPipeline(this._inputPipeline),c.setBindGroup(0,l),c.dispatchWorkgroups(Math.ceil(t/_e),Math.ceil(n/_e),this._model.inputChannels),c.end(),this._device.queue.submit([p.finish()])}finally{this._releaseBuffer(s)}this._context.dispatch(r.graph,{input:r.inputTensor},{output:r.outputTensor});const a=this._resources.track("gpu-buffer",await this._context.exportToGPU(r.outputTensor));try{const u=this._device.createBindGroup({label:"oidn/webnn/output-bindings",layout:this._outputPipeline.getBindGroupLayout(0),entries:[{binding:0,resource:{buffer:a}},{binding:1,resource:{buffer:r.outputBuffer}},{binding:2,resource:{buffer:r.outputUniform}}]}),l=this._device.createCommandEncoder({label:"oidn/webnn/output-unpack"}),p=l.beginComputePass();p.setPipeline(this._outputPipeline),p.setBindGroup(0,u),p.dispatchWorkgroups(Math.ceil(t/_e),Math.ceil(n/_e)),p.end(),this._device.queue.submit([l.finish()])}finally{this._releaseBuffer(a)}return r.outputBuffer}async executeCPU(e,t,n){const o=await this._execution(t,n),r=t*n,s=new P(r*this._model.inputChannels);for(let l=0;l<r;l++)for(let p=0;p<this._model.inputChannels;p++)s[p*r+l]=e[l*this._model.inputChannels+p];this._context.writeTensor(o.inputTensor,s),this._context.dispatch(o.graph,{input:o.inputTensor},{output:o.outputTensor});const a=new P(await this._context.readTensor(o.outputTensor)),u=new Float32Array(r*this._model.outputChannels);for(let l=0;l<r;l++)for(let p=0;p<this._model.outputChannels;p++)u[l*this._model.outputChannels+p]=a[p*r+l];return u}_destroyExecution(e){this._releaseGraph(e.graph),this._releaseTensor(e.inputTensor),this._releaseTensor(e.outputTensor),this._releaseBuffer(e.outputBuffer),this._releaseBuffer(e.inputUniform),this._releaseBuffer(e.outputUniform)}_retireExecution(e){this._retiredExecutions.add(e),this._device.queue.onSubmittedWorkDone().catch(()=>{}).then(()=>{this._retiredExecutions.delete(e),this._destroyExecution(e)})}_releaseBuffer(e){this._resources.release("gpu-buffer",e,()=>e.destroy())}_releaseTensor(e){this._resources.release("ml-tensor",e,()=>e.destroy())}_releaseGraph(e){this._resources.release("ml-graph",e,()=>{var t;return(t=e.destroy)==null?void 0:t.call(e)})}_releaseContext(){this._resources.release("ml-context",this._context,()=>{var e,t;return(t=(e=this._context).destroy)==null?void 0:t.call(e)})}getResourceInfo(){return this._resources.snapshot(this._pendingCreationCount+this._retiredExecutions.size)}dispose(){if(!this._disposed){this._disposed=!0;for(const e of this._shapeCache.values())this._destroyExecution(e);this._shapeCache.clear();for(const e of this._retiredExecutions)this._destroyExecution(e);this._retiredExecutions.clear(),this._shapePromises.clear(),this._releaseContext()}}}function ii(i,e){return Math.ceil(i/e)*e}function it(i){return i.data instanceof GPUBuffer||i.data instanceof GPUTexture}class ri{constructor(e,t,n={}){g(this,"_device");g(this,"_tileWidth",0);g(this,"_tileHeight",0);g(this,"_tileOverlapX",0);g(this,"_tileOverlapY",0);g(this,"_aux");g(this,"_hdr");g(this,"_dataProcessGPU");g(this,"_nativeExecutor");g(this,"_webNNExecutor");g(this,"_modelSpec");g(this,"_inputChannels");g(this,"_engine");g(this,"_dynamicTileController");g(this,"_lastExecution");this._aux=n.aux||!1,this._hdr=n.hdr||!1,this._engine=n.engine??"auto";const o=n.modelSpec??dt(e),r=nn(e,o);this._modelSpec=r.spec,this._inputChannels=r.inputChannels;const s=this._aux?9:3;if(r.inputChannels!==s)throw new Error(`OIDN model expects ${r.inputChannels} input channels, but aux=${this._aux} provides ${s}`);this._dynamicTileController=new wi(n.maxTileSize??512,n.dynamicTile),this._device=t.device,this._engine==="webnn"?this._webNNExecutor=new ni(this._device,r,{precision:n.precision}):this._nativeExecutor=new Jn(this._device,r,{precision:n.precision,kernel:n.kernel})}getDevice(){return this._device}async prepare(){if(this._webNNExecutor){await this._webNNExecutor.prepare();const e=ii(this._modelSpec.receptiveField/2,ne),t=[this._dynamicTileController.tileSize,this._dynamicTileController.minTileSize];await this._webNNExecutor.prewarm([...new Set(t)].map(n=>({width:n+2*e,height:n+2*e})));return}await this._nativeExecutor.prepare()}getRuntimeInfo(){var e;return{configuredEngine:this._engine,gpuEngine:this._webNNExecutor?"webnn":"wgsl",precision:(this._webNNExecutor??this._nativeExecutor).precision,kernel:this._nativeExecutor?{configured:this._nativeExecutor.kernelSetting,maxSpatialInputBlocks:this._nativeExecutor.maxSpatialInputBlocks,subgroupsAvailable:this._nativeExecutor.subgroupsAvailable}:void 0,webnn:(e=this._webNNExecutor)==null?void 0:e.support,resources:(this._webNNExecutor??this._nativeExecutor).getResourceInfo(),model:this._modelSpec.id,modelFamily:this._modelSpec.family,inputChannels:this._inputChannels,dynamicTile:{enabled:this._dynamicTileController.enabled,currentTileSize:this._dynamicTileController.tileSize,minTileSize:this._dynamicTileController.minTileSize,maxTileSize:this._dynamicTileController.maxTileSize,targetTileTimeMs:this._dynamicTileController.targetTileTimeMs},lastExecution:this._lastExecution}}profileNextExecution(){var e;return((e=this._nativeExecutor)==null?void 0:e.profileNextExecution())??!1}getLastExecutionProfile(){var e;return(e=this._nativeExecutor)==null?void 0:e.getLastExecutionProfile()}_updateModel(e,t){const n=this._dynamicTileController.tileSize;let o=qt(e,n),r=qt(t,n);const s=ii(this._modelSpec.receptiveField/2,ne);let a=s,u=s;e<=n&&(a=0),t<=n&&(u=0);const l=Math.max(o,r),p=Math.max(a,u);o=l,r=l,a=p,u=p,(o!==this._tileWidth||r!==this._tileHeight||a!==this._tileOverlapX||u!==this._tileOverlapY)&&(this._tileWidth=o,this._tileHeight=r,this._tileOverlapX=a,this._tileOverlapY=u)}_getTileSizeWithOverlap(){return{width:this._tileWidth+2*this._tileOverlapX,height:this._tileHeight+2*this._tileOverlapY}}_processImageData(e,t,n,o){const r=e.data,s=r.length/4,a=this._aux?9:3,u=new Float32Array(s*a);if(t&&!n||n&&!t)throw new Error("Normal map and albedo map are both required");if(t&&n&&(t.width!==n.width||t.height!==n.height||e.width!==t.width||e.height!==t.height))throw new Error("Image size mismatch");const l=t==null?void 0:t.data,p=n==null?void 0:n.data;for(let c=0;c<r.length;c+=4){const f=c/4*a;for(let d=0;d<3;d++)o?u[f+d]=r[c+d]:u[f+d]=r[c+d]/255,l&&(u[f+d+3]=l[c+d]/255),p&&(u[f+d+6]=p[c+d]/255)}return u}_readTile(e,t,n,o){const r=new Float32Array(n.width*n.height*t);for(let s=0;s<n.height;s++)for(let a=0;a<n.width;a++){const u=((s+n.y)*o+(a+n.x))*t,l=(s*n.width+a)*t;for(let p=0;p<t;p++)r[l+p]=e[u+p]}return r}_writeTile(e,t,n,o,r,s){const{data:a,width:u}=e,l=n.x-t.x,p=n.y-t.y;for(let c=0;c<n.height;c++)for(let f=0;f<n.width;f++){const d=((c+p)*r+f+l)*3,h=((c+n.y)*u+(f+n.x))*4;for(let _=0;_<3;_++)s?a[h+_]=o[d+_]:a[h+_]=Math.min(Math.max(o[d+_]*255,0),255);e.data[h+3]=s?1:255}}async _executeTile(e,t,n,o,r,s,a,u,l){const p=this._aux?9:3,c=this._tileOverlapX,f=this._tileOverlapY;let d=this._getTileSizeWithOverlap(),h={width:this._tileWidth,height:this._tileHeight},_=o>0?o*h.width-c:0,y=Math.min(_+d.width,s);_=Math.max(y-d.width,0);let x=r>0?r*h.height-f:0,m=Math.min(x+d.height,a);x=Math.max(m-d.height,0);const b=d.width,N=d.height,L=new pt(_,x,b,N);let q,R,M=1;const H=this._device;let I=this._dataProcessGPU;if(e instanceof Float32Array){let G=this._readTile(e,p,L,s);u&&(M=li({data:G,channels:p}),G=pi({data:G,channels:p,inputScale:M})),R=await(this._webNNExecutor??this._nativeExecutor).executeCPU(G,b,N)}else{I||(I=this._dataProcessGPU=new di(H,u)),I.setImageSize(s,a),I.setInputTile(L),o===0&&r===0&&I.copyInputDataToOutput(e.color);const{color:G,albedo:ye,normal:j}=I.forward(e.color,this._aux?e.albedo:void 0,this._aux?e.normal:void 0,l);q=await(this._webNNExecutor??this._nativeExecutor).execute(this._aux?[G,ye,j]:[G],b,N)}let A;const pe=Math.min(h.width,s),z=Math.min(h.height,a),O=new pt(o*pe,r*z,pe,z);if(O.width=Math.min(O.width,s-O.x),O.height=Math.min(O.height,a-O.y),e instanceof Float32Array){u&&(R=fi({data:R,channels:3,inputScale:M})),this._writeTile(n,L,O,R,d.width,u);for(let G=0;G<z;G++)for(let ye=0;ye<pe;ye++){const j=(G*pe+ye)*4,De=((G+O.y)*s+(ye+O.x))*4;for(let Te=0;Te<4;Te++)t.data[j+Te]=n.data[De+Te]}}else I.setOutputTile(O,L),A=I.inverse(q,e.color);return A}tileExecute({color:e,albedo:t,normal:n,done:o,progress:r,denoiseAlpha:s}){if(this._aux&&(!t||!n))throw new Error("Normal map and albedo map are both required");if(!this._aux&&(t||n))throw new Error("Normal map and albedo map are not required");const a=e.width,u=e.height,l=this._dynamicTileController.tileSize,p=a>l||u>l;this._updateModel(a,u);const c=this._hdr||!1;let f;it(e)||(f=this._processImageData(e,t,n,c));const d=this._tileWidth,h=this._tileHeight,_=Math.ceil(u/h),y=Math.ceil(a/d);function x(I,A){return c?{data:new Float32Array(I*A*4),width:I,height:A}:new ImageData(I,A)}const m=it(e)?void 0:x(a,u),b=it(e)?void 0:x(Math.min(d,a),Math.min(h,u));let N=!1;const L=()=>typeof performance>"u"?Date.now():performance.now(),q=L(),R=[],M=I=>{typeof requestAnimationFrame>"u"?setTimeout(I,0):requestAnimationFrame(I)},H=async(I,A)=>{if(N)return;const pe=L(),z=await this._executeTile(it(e)?{color:e.data,albedo:t==null?void 0:t.data,normal:n==null?void 0:n.data}:f,b,m,I,A,a,u,c,s);if(N)return;const O=m||{data:z,width:a,height:u};r==null||r(O,b,new pt(I*d,A*h,d,h),I+A*y,y*_);const G=I+1<y||A+1<_,ye=()=>{if(R.push(L()-pe),!N)if(G)M(()=>{N||(I+1<y?H(I+1,A):A+1<_&&H(0,A+1))});else{const j=[...R].sort((zt,Dt)=>zt-Dt),De=Math.floor(j.length/2),Te=j.length%2?j[De]:(j[De-1]+j[De])/2;this._lastExecution={width:a,height:u,tileWidth:d,tileHeight:h,tileCount:y*_,durationMs:L()-q,tileTimeMs:{min:j[0],median:Te,mean:j.reduce((zt,Dt)=>zt+Dt,0)/j.length,max:j[j.length-1]}},p&&this._dynamicTileController.observe(R),o(O)}};xi(this._device.queue).then(ye)};return H(0,0),()=>{N=!0}}dispose(){var e,t,n;(e=this._dataProcessGPU)==null||e.dispose(),(t=this._nativeExecutor)==null||t.dispose(),(n=this._webNNExecutor)==null||n.dispose()}}async function Lr(){var a;if(!navigator.gpu)throw new Error("WebGPU is not available");const i={powerPreference:"high-performance"},e=await navigator.gpu.requestAdapter(i);if(!e)throw new Error("No WebGPU adapter is available");const t={},n=[];e.features.has("timestamp-query")&&n.push("timestamp-query"),e.features.has("bgra8unorm-storage")&&n.push("bgra8unorm-storage"),e.features.has("shader-f16")&&n.push("shader-f16"),t.requiredFeatures=n;const o=e.limits;t.requiredLimits={maxComputeWorkgroupStorageSize:o.maxComputeWorkgroupStorageSize,maxComputeWorkgroupsPerDimension:o.maxComputeWorkgroupsPerDimension,maxStorageBufferBindingSize:o.maxStorageBufferBindingSize,maxBufferSize:o.maxBufferSize,maxComputeWorkgroupSizeX:o.maxComputeWorkgroupSizeX,maxComputeInvocationsPerWorkgroup:o.maxComputeInvocationsPerWorkgroup};const r=await e.requestDevice(t),s=e.info??await((a=e.requestAdapterInfo)==null?void 0:a.call(e));return oi(r,s)}async function oi(i,e){return{device:i,adapterInfo:e}}async function si(i,e,t){const n=await(e?oi(e.device,e.adapterInfo):Lr()),o=Wt(i),r=new ri(o,n,t);return await r.prepare(),r}async function Rr(i,e,t){return fetch(i).then(n=>n.arrayBuffer()).then(n=>si(n,e,t))}C.NativeUNetExecutor=Jn,C.OIDN_UNET_LARGE_SPEC=Zt,C.OIDN_UNET_SMALL_SPEC=jt,C.UNet=ri,C.WebNNUNetExecutor=ni,C.detectUNetModelSpec=dt,C.initUNetFromBuffer=si,C.initUNetFromURL=Rr,C.optimizeModelGraph=nt,C.parseTZA=Wt,C.planModelExecution=Yn,C.resolveNativeUNetPrecision=Zn,C.validateUNetModel=nn,Object.defineProperty(C,Symbol.toStringTag,{value:"Module"})});
|
|
1108
|
+
`}function Io(n,e){if(e==="fp32")throw new Error("OIDN WebNN GPU interop currently requires FP16 exportable tensors");if(!n.features.has("shader-f16"))throw new Error("OIDN WebNN requires shader-f16 on the shared GPUDevice");return"fp16"}function Bo(n,e,t){return t==="relu"?n.relu(e):e}class kn{constructor(e,t,i={}){m(this,"precision");m(this,"support");m(this,"_context");m(this,"_builderConstructor");m(this,"_shapeCache",new Map);m(this,"_shapePromises",new Map);m(this,"_retiredExecutions",new Set);m(this,"_pendingCreationCount",0);m(this,"_shapeCacheSize");m(this,"_clock",0);m(this,"_inputPipeline");m(this,"_outputPipeline");m(this,"_resources",new un);m(this,"_disposed",!1);this._device=e,this._model=t,this.precision=Io(e,i.precision??"auto"),this._shapeCacheSize=Math.max(1,i.shapeCacheSize??2),this.support={available:!1,fp16Conv:!1,gpuInterop:!1};const r=vo(e,t.inputChannels/3);this._inputPipeline=r.input,this._outputPipeline=r.output}async prepare(){var i,r,o;if(this._disposed)throw new Error("OIDN WebNN executor is disposed");const e=(i=globalThis.navigator)==null?void 0:i.ml,t=globalThis.MLGraphBuilder;if(!(e!=null&&e.createContext)||typeof t!="function")throw this.support.reason="WebNN is not exposed by this browser",new Error(this.support.reason);this._builderConstructor=t;try{try{this._context=await e.createContext({deviceType:"gpu",powerPreference:"high-performance"})}catch{this._context=await e.createContext({deviceType:"gpu"})}if(this._resources.track("ml-context",this._context),this._disposed)throw new Error("OIDN WebNN executor is disposed");if(typeof this._context.createExportableTensor!="function"||typeof this._context.exportToGPU!="function")throw this.support.reason="WebNN WebGPU tensor interop is unavailable",new Error(this.support.reason);const a=((o=(r=this._context).opSupportLimits)==null?void 0:o.call(r))??{};if(this.support.fp16Conv=Nt(a,"conv2d","input","float16")&&Nt(a,"conv2d","filter","float16")&&Nt(a,"conv2d","output","float16"),!this.support.fp16Conv)throw this.support.reason="WebNN does not support FP16 conv2d",new Error(this.support.reason);let u,s;try{u=this._resources.track("ml-tensor",await this._context.createExportableTensor({dataType:"float16",shape:[4]},this._device)),s=this._resources.track("gpu-buffer",await this._context.exportToGPU(u)),this.support.gpuInterop=!0}catch(l){throw this.support.reason=`WebNN FP16 WebGPU interop failed: ${String(l)}`,new Error(this.support.reason)}finally{this._releaseBuffer(s),this._releaseTensor(u)}this.support.available=!0}catch(a){throw this._releaseContext(),a}}_constant(e,t){return e.constant({dataType:"float16",shape:[...t.desc.dims]},bo(t,this.precision))}async _createExecution(e,t){if(this._disposed)throw new Error("OIDN WebNN executor is disposed");const i=new this._builderConstructor(this._context),r=new Map,o=new Map;r.set(this._model.spec.input,i.input("input",{dataType:"float16",shape:[1,this._model.inputChannels,t,e]})),o.set(this._model.spec.input,[this._model.inputChannels,t,e]);for(const d of this._model.spec.nodes){let h,f;if(d.op==="conv2d"){const g=o.get(d.input),y=this._model.convTensors.get(d.id),v=i.conv2d(r.get(d.input),this._constant(i,y.weight),{bias:this._constant(i,y.bias),padding:[1,1,1,1],inputLayout:"nchw",filterLayout:"oihw"});h=Bo(i,v,d.activation),f=[y.outputChannels,g[1],g[2]]}else if(d.op==="maxPool2d"){const g=o.get(d.input);h=i.maxPool2d(r.get(d.input),{windowDimensions:[2,2],strides:[2,2],padding:[0,g[1]%2,0,g[2]%2],layout:"nchw"}),f=[g[0],Math.ceil(g[1]/2),Math.ceil(g[2]/2)]}else if(d.op==="upsample2d"){const g=o.get(d.input);h=i.resample2d(r.get(d.input),{mode:"nearest-neighbor",axes:[2,3],scales:[2,2]}),f=[g[0],g[1]*2,g[2]*2]}else{const g=d.inputs.map(y=>o.get(y));if(g.some(y=>y[1]!==g[0][1]||y[2]!==g[0][2]))throw new Error(`WebNN concat ${d.id} has mismatched spatial shapes`);h=i.concat(d.inputs.map(y=>r.get(y)),1),f=[g.reduce((y,v)=>y+v[0],0),g[0][1],g[0][2]]}r.set(d.id,h),o.set(d.id,f)}let a,u,s,l,p,c;try{return a=this._resources.track("ml-graph",await i.build({output:r.get(this._model.spec.output)})),u=this._resources.track("ml-tensor",await this._context.createExportableTensor({dataType:"float16",shape:[1,this._model.inputChannels,t,e],writable:!0},this._device)),s=this._resources.track("ml-tensor",await this._context.createExportableTensor({dataType:"float16",shape:[1,this._model.outputChannels,t,e],readable:!0},this._device)),l=this._resources.track("gpu-buffer",this._device.createBuffer({label:`oidn/webnn/output/${e}x${t}`,size:e*t*4*4,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST})),p=this._resources.track("gpu-buffer",$n(this._device,`oidn/webnn/input/${e}x${t}`,[e,t,this._model.inputChannels])),c=this._resources.track("gpu-buffer",$n(this._device,`oidn/webnn/output/${e}x${t}`,[e,t])),{graph:a,inputTensor:u,outputTensor:s,outputBuffer:l,inputUniform:p,outputUniform:c,width:e,height:t,lastUsed:++this._clock}}catch(d){throw this._releaseBuffer(c),this._releaseBuffer(p),this._releaseBuffer(l),this._releaseTensor(s),this._releaseTensor(u),this._releaseGraph(a),d}}async _execution(e,t){const i=`${e}x${t}`;let r=this._shapeCache.get(i);if(!r){let o=this._shapePromises.get(i);o||(o=(async()=>{this._pendingCreationCount++;try{return await this._createExecution(e,t)}finally{this._pendingCreationCount--}})(),this._shapePromises.set(i,o));try{if(r=await o,this._disposed)throw this._destroyExecution(r),new Error("OIDN WebNN executor is disposed");this._shapeCache.set(i,r)}finally{this._shapePromises.get(i)===o&&this._shapePromises.delete(i)}if(this._shapeCache.size>this._shapeCacheSize){const a=[...this._shapeCache.entries()].filter(([u])=>u!==i).sort((u,s)=>u[1].lastUsed-s[1].lastUsed)[0];a&&(this._shapeCache.delete(a[0]),this._retireExecution(a[1]))}}return r.lastUsed=++this._clock,r}async prewarm(e){for(const t of e)await this._execution(t.width,t.height)}async execute(e,t,i){const r=this._model.inputChannels/3;if(e.length!==r)throw new Error(`OIDN WebNN expected ${r} input buffers, got ${e.length}`);const o=await this._execution(t,i),a=this._resources.track("gpu-buffer",await this._context.exportToGPU(o.inputTensor));try{const s=e.map((d,h)=>({binding:h,resource:{buffer:d}}));s.push({binding:r,resource:{buffer:a}}),s.push({binding:r+1,resource:{buffer:o.inputUniform}});const l=this._device.createBindGroup({label:"oidn/webnn/input-bindings",layout:this._inputPipeline.getBindGroupLayout(0),entries:s}),p=this._device.createCommandEncoder({label:"oidn/webnn/input-pack"}),c=p.beginComputePass();c.setPipeline(this._inputPipeline),c.setBindGroup(0,l),c.dispatchWorkgroups(Math.ceil(t/fe),Math.ceil(i/fe),this._model.inputChannels),c.end(),this._device.queue.submit([p.finish()])}finally{this._releaseBuffer(a)}this._context.dispatch(o.graph,{input:o.inputTensor},{output:o.outputTensor});const u=this._resources.track("gpu-buffer",await this._context.exportToGPU(o.outputTensor));try{const s=this._device.createBindGroup({label:"oidn/webnn/output-bindings",layout:this._outputPipeline.getBindGroupLayout(0),entries:[{binding:0,resource:{buffer:u}},{binding:1,resource:{buffer:o.outputBuffer}},{binding:2,resource:{buffer:o.outputUniform}}]}),l=this._device.createCommandEncoder({label:"oidn/webnn/output-unpack"}),p=l.beginComputePass();p.setPipeline(this._outputPipeline),p.setBindGroup(0,s),p.dispatchWorkgroups(Math.ceil(t/fe),Math.ceil(i/fe)),p.end(),this._device.queue.submit([l.finish()])}finally{this._releaseBuffer(u)}return o.outputBuffer}async executeCPU(e,t,i){const r=await this._execution(t,i),o=t*i,a=new E(o*this._model.inputChannels);for(let l=0;l<o;l++)for(let p=0;p<this._model.inputChannels;p++)a[p*o+l]=e[l*this._model.inputChannels+p];this._context.writeTensor(r.inputTensor,a),this._context.dispatch(r.graph,{input:r.inputTensor},{output:r.outputTensor});const u=new E(await this._context.readTensor(r.outputTensor)),s=new Float32Array(o*this._model.outputChannels);for(let l=0;l<o;l++)for(let p=0;p<this._model.outputChannels;p++)s[l*this._model.outputChannels+p]=u[p*o+l];return s}_destroyExecution(e){this._releaseGraph(e.graph),this._releaseTensor(e.inputTensor),this._releaseTensor(e.outputTensor),this._releaseBuffer(e.outputBuffer),this._releaseBuffer(e.inputUniform),this._releaseBuffer(e.outputUniform)}_retireExecution(e){this._retiredExecutions.add(e),this._device.queue.onSubmittedWorkDone().catch(()=>{}).then(()=>{this._retiredExecutions.delete(e),this._destroyExecution(e)})}_releaseBuffer(e){this._resources.release("gpu-buffer",e,()=>e.destroy())}_releaseTensor(e){this._resources.release("ml-tensor",e,()=>e.destroy())}_releaseGraph(e){this._resources.release("ml-graph",e,()=>{var t;return(t=e.destroy)==null?void 0:t.call(e)})}_releaseContext(){this._resources.release("ml-context",this._context,()=>{var e,t;return(t=(e=this._context).destroy)==null?void 0:t.call(e)})}getResourceInfo(){return this._resources.snapshot(this._pendingCreationCount+this._retiredExecutions.size)}dispose(){if(!this._disposed){this._disposed=!0;for(const e of this._shapeCache.values())this._destroyExecution(e);this._shapeCache.clear();for(const e of this._retiredExecutions)this._destroyExecution(e);this._retiredExecutions.clear(),this._shapePromises.clear(),this._releaseContext()}}}const To=100;function we(n,e){return Math.ceil(n/e)*e}function ut(n){return n.data instanceof GPUBuffer||n.data instanceof GPUTexture}class In{constructor(e,t,i={}){m(this,"_device");m(this,"_aux");m(this,"_hdr");m(this,"_hdrTransfer");m(this,"_dataProcessGPU");m(this,"_nativeExecutor");m(this,"_webNNExecutor");m(this,"_modelSpec");m(this,"_inputChannels");m(this,"_engine");m(this,"_dynamicTileController");m(this,"_lastExecution");m(this,"_activeExecutionFailures",new Set);m(this,"_deviceLostObserved",!1);m(this,"_deviceLostSettled",!1);m(this,"_deviceLostReason");this._aux=i.aux||!1,this._hdr=i.hdr||!1,this._hdrTransfer=i.hdrTransfer??"pu",this._engine=i.engine??"auto";const r=i.modelSpec??dt(e),o=yi(e,r);this._modelSpec=o.spec,this._inputChannels=o.inputChannels;const a=this._aux?9:3;if(o.inputChannels!==a)throw new Error(`OIDN model expects ${o.inputChannels} input channels, but aux=${this._aux} provides ${a}`);this._dynamicTileController=new nr(i.maxTileSize??512,i.dynamicTile),this._device=t,this._observeDeviceLoss(),this._engine==="webnn"?this._webNNExecutor=new kn(this._device,o,{precision:i.precision}):this._nativeExecutor=new vn(this._device,o,{precision:i.precision,kernel:i.kernel,gemm:i.gemm})}getDevice(){return this._device}async prepare(){if(this._webNNExecutor){await this._webNNExecutor.prepare();const e=we(this._modelSpec.receptiveField/2,D),t=[this._dynamicTileController.tileSize,this._dynamicTileController.minTileSize];await this._webNNExecutor.prewarm([...new Set(t)].map(i=>({width:i+2*e,height:i+2*e})));return}await this._nativeExecutor.prepare()}async prepareForImage(e,t,i={}){var s,l;const r=we(this._modelSpec.receptiveField/2,D),o=i.tileOverlap===void 0?r:we(Math.max(0,i.tileOverlap),D),a=pt(e,t,i.wholeImage?we(Math.max(e,t),D):this._dynamicTileController.tileSize,o),u=[...new Map(a.tiles.map(({input:p})=>[`${p.width}x${p.height}`,{width:p.width,height:p.height}])).values()];await((s=this._webNNExecutor)==null?void 0:s.prewarm(u)),(l=this._nativeExecutor)==null||l.prewarm(u)}getRuntimeInfo(){var e;return{configuredEngine:this._engine,gpuEngine:this._webNNExecutor?"webnn":"wgsl",precision:(this._webNNExecutor??this._nativeExecutor).precision,kernel:this._nativeExecutor?{configured:this._nativeExecutor.kernelSetting,gemm:this._nativeExecutor.gemm,maxSpatialInputBlocks:this._nativeExecutor.maxSpatialInputBlocks,subgroupsAvailable:this._nativeExecutor.subgroupsAvailable}:void 0,webnn:(e=this._webNNExecutor)==null?void 0:e.support,resources:(this._webNNExecutor??this._nativeExecutor).getResourceInfo(),model:this._modelSpec.id,modelFamily:this._modelSpec.family,inputChannels:this._inputChannels,hdrTransfer:this._hdrTransfer,dynamicTile:{enabled:this._dynamicTileController.enabled,currentTileSize:this._dynamicTileController.tileSize,minTileSize:this._dynamicTileController.minTileSize,maxTileSize:this._dynamicTileController.maxTileSize,targetTileTimeMs:this._dynamicTileController.targetTileTimeMs},lastExecution:this._lastExecution,activeExecutionCount:this._activeExecutionFailures.size}}_observeDeviceLoss(){if(this._activeExecutionFailures??(this._activeExecutionFailures=new Set),this._deviceLostObserved)return;this._deviceLostObserved=!0;const e=this._device.lost;if(!e)return;const t=i=>{if(this._deviceLostSettled)return;this._deviceLostSettled=!0,this._deviceLostReason=i;const r=[...this._activeExecutionFailures];this._activeExecutionFailures.clear();for(const o of r)o(i)};e.then(i=>t(new Error(`WebGPU device lost: ${i.message}`)),t)}_registerExecutionFailure(e){return this._observeDeviceLoss(),this._activeExecutionFailures.add(e),this._deviceLostSettled&&(this._activeExecutionFailures.delete(e),queueMicrotask(()=>e(this._deviceLostReason))),()=>this._activeExecutionFailures.delete(e)}profileNextExecution(){var e;return((e=this._nativeExecutor)==null?void 0:e.profileNextExecution())??!1}getLastExecutionProfile(){var e;return(e=this._nativeExecutor)==null?void 0:e.getLastExecutionProfile()}_processImageData(e,t,i,r){const o=e.data,a=o.length/4,u=this._aux?9:3,s=new Float32Array(a*u);if(t&&!i||i&&!t)throw new Error("Normal map and albedo map are both required");if(t&&i&&(t.width!==i.width||t.height!==i.height||e.width!==t.width||e.height!==t.height))throw new Error("Image size mismatch");const l=t==null?void 0:t.data,p=i==null?void 0:i.data;for(let c=0;c<o.length;c+=4){const d=c/4*u;for(let h=0;h<3;h++)r?s[d+h]=o[c+h]:s[d+h]=o[c+h]/255,l&&(s[d+h+3]=l[c+h]/255),p&&(s[d+h+6]=p[c+h]/255)}return s}_readTile(e,t,i,r){const o=new Float32Array(i.width*i.height*t),a=e.length/(r*t);for(let u=0;u<i.height;u++)for(let s=0;s<i.width;s++){const l=Math.min(r-1,s+i.x),c=(Math.min(a-1,u+i.y)*r+l)*t,d=(u*i.width+s)*t;for(let h=0;h<t;h++)o[d+h]=e[c+h]}return o}_writeTile(e,t,i,r,o,a){const{data:u,width:s}=e,l=i.x-t.x,p=i.y-t.y;for(let c=0;c<i.height;c++)for(let d=0;d<i.width;d++){const h=((c+p)*o+d+l)*3,f=((c+i.y)*s+(d+i.x))*4;for(let g=0;g<3;g++)a?u[f+g]=r[h+g]:u[f+g]=Math.min(Math.max(r[h+g]*255,0),255);e.data[f+3]=a?1:255}}async _executeTile(e,t,i,r,o,a,u,s,l){const p=this._aux?9:3,c=new lt(r.input.x,r.input.y,r.input.width,r.input.height),d=new lt(r.output.x,r.output.y,r.output.width,r.output.height),h=c.width,f=c.height;let g,y,v=1;const _=this._device;let w=this._dataProcessGPU;if(e instanceof Float32Array){let x=this._readTile(e,p,c,a);s&&(v=qn({data:x,channels:p}),x=Gn({data:x,channels:p,inputScale:v,transfer:this._hdrTransfer})),y=await(this._webNNExecutor??this._nativeExecutor).executeCPU(x,h,f)}else{w||(w=this._dataProcessGPU=new Zn(_,s,this._hdrTransfer)),w.setImageSize(a,u),w.setInputTile(c),o&&w.copyInputDataToOutput(e.color);const{color:x,albedo:B,normal:k}=w.forward(e.color,this._aux?e.albedo:void 0,this._aux?e.normal:void 0,l);g=await(this._webNNExecutor??this._nativeExecutor).execute(this._aux?[x,B,k]:[x],h,f)}let b;if(e instanceof Float32Array){s&&(y=Yn({data:y,channels:3,inputScale:v,transfer:this._hdrTransfer})),this._writeTile(i,c,d,y,c.width,s);for(let x=0;x<d.height;x++)for(let B=0;B<d.width;B++){const k=(x*d.width+B)*4,C=((x+d.y)*a+(B+d.x))*4;for(let L=0;L<4;L++)t.data[k+L]=i.data[C+L]}}else w.setOutputTile(d,c),b=w.inverse(g,e.color);return b}tileExecute({color:e,albedo:t,normal:i,done:r,progress:o,denoiseAlpha:a,tileOverlap:u,wholeImage:s,scheduling:l="event-loop",error:p}){if(this._aux&&(!t||!i))throw new Error("Normal map and albedo map are both required");if(!this._aux&&(t||i))throw new Error("Normal map and albedo map are not required");const c=e.width,d=e.height,h=this._dynamicTileController.tileSize,f=s?we(Math.max(c,d),D):h,g=we(this._modelSpec.receptiveField/2,D),y=u===void 0?g:we(Math.max(0,u),D),v=pt(c,d,f,y),_=v.tiles.length>1,w=this._hdr||!1;let b;ut(e)||(b=this._processImageData(e,t,i,w));function x(V,H){return w?{data:new Float32Array(V*H*4),width:V,height:H}:new ImageData(V,H)}const B=ut(e)?void 0:x(c,d);let k="active",C,L,F=()=>!1;const z=()=>typeof performance>"u"?Date.now():performance.now(),Pe=z(),X=[],Z=()=>{C!==void 0&&(clearTimeout(C),C=void 0),L!==void 0&&typeof cancelAnimationFrame<"u"&&(cancelAnimationFrame(L),L=void 0)},Le=V=>{console.error("OIDN error callback failed",V)},Ut=V=>{if(k==="active")if(k="settled",Z(),F(),p)try{Promise.resolve(p(V)).catch(Le)}catch(H){Le(H)}else console.error("OIDN execution failed",V)},So=V=>{if(l==="event-loop"||typeof requestAnimationFrame>"u")C=setTimeout(()=>{C=void 0,V()},0);else{const H=()=>{Z(),V()};L=requestAnimationFrame(H),C=setTimeout(H,To)}},Tn=async V=>{if(k!=="active")return;const H=v.tiles[V],Pn=ut(e)?void 0:x(H.output.width,H.output.height),Eo=z(),Ao=await this._executeTile(ut(e)?{color:e.data,albedo:t==null?void 0:t.data,normal:i==null?void 0:i.data}:b,Pn,B,H,V===0,c,d,w,a);if(k!=="active")return;const Cn=B||{data:Ao,width:c,height:d};if(o&&await o(Cn,Pn,new lt(H.output.x,H.output.y,H.output.width,H.output.height),V,v.tiles.length),k!=="active")return;const Mo=V+1<v.tiles.length;if(await this._device.queue.onSubmittedWorkDone(),k!=="active")return;await(async()=>{if(X.push(z()-Eo),k==="active")if(Mo)So(()=>{k==="active"&&Tn(V+1).catch(Ut)});else{const se=[...X].sort((Lt,Wt)=>Lt-Wt),zt=Math.floor(se.length/2),Oo=se.length%2?se[zt]:(se[zt-1]+se[zt])/2;this._lastExecution={width:c,height:d,tileCount:v.tiles.length,tileColumns:v.columns,tileRows:v.rows,tileOverlap:v.overlap,inputPixelCount:v.inputPixelCount,inputShapeCount:v.inputShapeCount,durationMs:z()-Pe,tileTimeMs:{min:se[0],median:Oo,mean:se.reduce((Lt,Wt)=>Lt+Wt,0)/se.length,max:se[se.length-1]}},_&&this._dynamicTileController.observe(X),F(),await r(Cn),k==="active"&&(k="settled")}})()};return F=this._registerExecutionFailure(Ut),Tn(0).catch(Ut),()=>{k==="active"&&(k="aborted",Z(),F())}}dispose(){var e,t,i;(e=this._dataProcessGPU)==null||e.dispose(),(t=this._nativeExecutor)==null||t.dispose(),(i=this._webNNExecutor)==null||i.dispose()}}async function Po(){if(!navigator.gpu)throw new Error("WebGPU is not available");const n={powerPreference:"high-performance"},e=await navigator.gpu.requestAdapter(n);if(!e)throw new Error("No WebGPU adapter is available");const t={},i=[];e.features.has("timestamp-query")&&i.push("timestamp-query"),e.features.has("bgra8unorm-storage")&&i.push("bgra8unorm-storage"),e.features.has("shader-f16")&&i.push("shader-f16"),t.requiredFeatures=i;const r=e.limits;return t.requiredLimits={maxComputeWorkgroupStorageSize:r.maxComputeWorkgroupStorageSize,maxComputeWorkgroupsPerDimension:r.maxComputeWorkgroupsPerDimension,maxStorageBufferBindingSize:r.maxStorageBufferBindingSize,maxBufferSize:r.maxBufferSize,maxComputeWorkgroupSizeX:r.maxComputeWorkgroupSizeX,maxComputeInvocationsPerWorkgroup:r.maxComputeInvocationsPerWorkgroup},e.requestDevice(t)}async function Bn(n,e,t){const i=(e==null?void 0:e.device)??await Po(),r=Dt(n),o=new In(r,i,t);return await o.prepare(),o}async function Co(n,e,t){return fetch(n).then(i=>i.arrayBuffer()).then(i=>Bn(i,e,t))}A.NativeUNetExecutor=vn,A.OIDN_UNET_LARGE_SPEC=hi,A.OIDN_UNET_SMALL_SPEC=di,A.UNet=In,A.WebNNUNetExecutor=kn,A.detectUNetModelSpec=dt,A.initUNetFromBuffer=Bn,A.initUNetFromURL=Co,A.optimizeModelGraph=rt,A.parseTZA=Dt,A.planModelExecution=sn,A.planTileGrid=pt,A.resolveNativeUNetPrecision=wn,A.validateUNetModel=yi,Object.defineProperty(A,Symbol.toStringTag,{value:"Module"})});
|