| 209 | } |
| 210 | |
| 211 | XlaOp TorchScatterDense(XlaOp input, XlaOp index, XlaOp src, int64 dim, |
| 212 | const std::function<XlaOp(XlaOp, XlaOp)>& combiner) { |
| 213 | XlaBuilder* builder = input.builder(); |
| 214 | return builder->ReportErrorOrReturn([&]() -> StatusOr<XlaOp> { |
| 215 | TF_ASSIGN_OR_RETURN(Shape index_shape, builder->GetShape(index)); |
| 216 | TF_ASSIGN_OR_RETURN(Shape input_shape, builder->GetShape(input)); |
| 217 | std::vector<int64> index_broadcast_dims; |
| 218 | std::vector<int64> sizes; |
| 219 | for (int64 i = 0; i < index_shape.rank(); ++i) { |
| 220 | if (i < dim) { |
| 221 | index_broadcast_dims.push_back(i); |
| 222 | } else { |
| 223 | if (i == dim) { |
| 224 | sizes.push_back(input_shape.dimensions(i)); |
| 225 | } |
| 226 | index_broadcast_dims.push_back(i + 1); |
| 227 | } |
| 228 | sizes.push_back(index_shape.dimensions(i)); |
| 229 | } |
| 230 | auto mask = |
| 231 | Eq(BroadcastInDim(index, sizes, index_broadcast_dims), |
| 232 | Iota(builder, |
| 233 | ShapeUtil::MakeShape(index_shape.element_type(), sizes), dim)); |
| 234 | auto masked_src = |
| 235 | Select(mask, BroadcastInDim(src, sizes, index_broadcast_dims), |
| 236 | Zeros(builder, |
| 237 | ShapeUtil::MakeShape(input_shape.element_type(), sizes))); |
| 238 | |
| 239 | return combiner( |
| 240 | input, |
| 241 | Reduce(masked_src, Zero(builder, input_shape.element_type()), |
| 242 | CreateScalarComputation("reducer", input_shape.element_type(), |
| 243 | builder, combiner), |
| 244 | {dim + 1})); |
| 245 | }); |
| 246 | } |
| 247 | |
| 248 | XlaOp TorchIndexSelect(XlaOp input, XlaOp index, int64 dim, int64 batch_dims) { |
| 249 | XlaBuilder* builder = input.builder(); |