MCPcopy Create free account
hub / github.com/tensorflow/tfjs / tensorScatterUpdate

Function tensorScatterUpdate

tfjs-backend-wasm/src/kernels/TensorScatterUpdate.ts:45–79  ·  view source on GitHub ↗
(args: {
  backend: BackendWasm,
  inputs: TensorScatterUpdateInputs,
  attrs: TensorScatterUpdateAttrs
})

Source from the content-addressed store, hash-verified

43}
44
45function tensorScatterUpdate(args: {
46 backend: BackendWasm,
47 inputs: TensorScatterUpdateInputs,
48 attrs: TensorScatterUpdateAttrs
49}): TensorInfo {
50 const {backend, inputs, attrs} = args;
51 const {tensor, indices, updates} = inputs;
52 const {} = attrs;
53
54 const out = backend.makeOutput(tensor.shape, tensor.dtype);
55 if (util.sizeFromShape(tensor.shape) === 0) {
56 return out;
57 }
58
59 const {sliceRank, numUpdates, sliceSize, strides, outputSize} =
60 scatter_util.calculateShapes(updates, indices, tensor.shape);
61
62 const indicesData = backend.dataIdMap.get(indices.dataId);
63 const indicesId = indicesData.id;
64
65 const updatesData = backend.dataIdMap.get(updates.dataId);
66 const updatesId = updatesData.id;
67
68 const tensorData = backend.dataIdMap.get(tensor.dataId);
69 const tensorId = tensorData.id;
70
71 const stridesBytes = new Uint8Array(new Int32Array(strides).buffer);
72
73 const outId = backend.dataIdMap.get(out.dataId).id;
74 wasmTensorScatterUpdate(
75 indicesId, updatesId, CppDType[updates.dtype], sliceRank, numUpdates,
76 sliceSize, stridesBytes, outputSize, outId, tensorId);
77
78 return out;
79}
80
81export const tensorScatterUpdateConfig: KernelConfig = {
82 kernelName: TensorScatterUpdate,

Callers

nothing calls this directly

Calls 2

makeOutputMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…