(args: {
backend: BackendWasm,
inputs: TensorScatterUpdateInputs,
attrs: TensorScatterUpdateAttrs
})
| 43 | } |
| 44 | |
| 45 | function 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 | |
| 81 | export const tensorScatterUpdateConfig: KernelConfig = { |
| 82 | kernelName: TensorScatterUpdate, |
nothing calls this directly
no test coverage detected
searching dependent graphs…