"use strict"; /** * @license * Copyright 2018 Google LLC. All Rights Reserved. * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. * You may obtain a copy of the License at * * http://www.apache.org/licenses/LICENSE-2.0 * * Unless required by applicable law or agreed to in writing, software * distributed under the License is distributed on an "AS IS" BASIS, * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. * See the License for the specific language governing permissions and * limitations under the License. * ============================================================================= */ Object.defineProperty(exports, "__esModule", { value: true }); var broadcast_util = require("../../ops/broadcast_util"); var BatchNormPackedProgram = /** @class */ (function () { function BatchNormPackedProgram(xShape, meanShape, varianceShape, offsetShape, scaleShape, varianceEpsilon) { this.packedInputs = true; this.packedOutput = true; this.variableNames = ['x', 'mean', 'variance']; broadcast_util.assertAndGetBroadcastShape(xShape, meanShape); broadcast_util.assertAndGetBroadcastShape(xShape, varianceShape); var offsetSnippet = 'vec4(0.0)'; if (offsetShape != null) { broadcast_util.assertAndGetBroadcastShape(xShape, offsetShape); this.variableNames.push('offset'); offsetSnippet = 'getOffsetAtOutCoords()'; } var scaleSnippet = 'vec4(1.0)'; if (scaleShape != null) { broadcast_util.assertAndGetBroadcastShape(xShape, scaleShape); this.variableNames.push('scale'); scaleSnippet = 'getScaleAtOutCoords()'; } this.outputShape = xShape; this.userCode = "\n void main() {\n vec4 offset = " + offsetSnippet + ";\n vec4 scale = " + scaleSnippet + ";\n\n vec4 x = getXAtOutCoords();\n vec4 mean = getMeanAtOutCoords();\n vec4 variance = getVarianceAtOutCoords();\n\n vec4 inv = scale * inversesqrt(variance + vec4(" + varianceEpsilon + "));\n\n setOutput((x - mean) * inv + offset);\n }\n "; } return BatchNormPackedProgram; }()); exports.BatchNormPackedProgram = BatchNormPackedProgram; //# sourceMappingURL=batchnorm_packed_gpu.js.map