(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]])
| 354 | |
| 355 | @register_lower_rule(mops.SetSubtensor) |
| 356 | def setsubtensor_lower(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]]): |
| 357 | assert len(ctx.vars_out) == 1 |
| 358 | opr, dst, src = ctx.op, args[0], args[1] # dst[indices] = src |
| 359 | |
| 360 | raw_slices = _parse_subtensor_items(opr.items, dst.shape, ctx.vars_in[2:], args[2:]) |
| 361 | ( |
| 362 | update_window_dims, |
| 363 | inserted_window_dims, |
| 364 | scattered_dims_to_operand_dims, |
| 365 | indices, |
| 366 | slice_shape, |
| 367 | ) = _get_scatter_configs_and_indices(dst.shape, src.shape, raw_slices) |
| 368 | |
| 369 | if len(slice_shape) == 0 or np.prod(slice_shape) == 0: |
| 370 | return [dst] |
| 371 | |
| 372 | src = src.broadcast_to(slice_shape) |
| 373 | |
| 374 | out = xla_scatter( |
| 375 | dst, |
| 376 | indices, |
| 377 | src, |
| 378 | update_window_dims=update_window_dims, |
| 379 | inserted_window_dims=inserted_window_dims, |
| 380 | scattered_dims_to_operand_dims=scattered_dims_to_operand_dims, |
| 381 | reduce_mode=None, |
| 382 | ) |
| 383 | |
| 384 | return out |
| 385 | |
| 386 | |
| 387 | def _check_tensor_indexing_arg(src, index, axis): |
nothing calls this directly
no test coverage detected