| 169 | } |
| 170 | |
| 171 | XlaOp SetMatrixDiagonal(XlaOp matrix, XlaOp diag, int k) { |
| 172 | XlaBuilder* builder = matrix.builder(); |
| 173 | return builder->ReportErrorOrReturn([&]() -> StatusOr<XlaOp> { |
| 174 | TF_ASSIGN_OR_RETURN(Shape shape, builder->GetShape(matrix)); |
| 175 | TF_ASSIGN_OR_RETURN(Shape diag_shape, builder->GetShape(diag)); |
| 176 | auto n_dims = static_cast<int32>(shape.rank()); |
| 177 | TF_RET_CHECK(n_dims >= 2); |
| 178 | const int64 m = shape.dimensions(n_dims - 2); |
| 179 | const int64 n = shape.dimensions(n_dims - 1); |
| 180 | const int64 d = diag_shape.dimensions(n_dims - 2); |
| 181 | std::vector<int64> broadcast_dims(n_dims - 1); |
| 182 | absl::c_iota(broadcast_dims, 0); |
| 183 | int64 pad_high = m - d; |
| 184 | if (k < 0) { |
| 185 | ++(broadcast_dims.back()); |
| 186 | pad_high = n - d; |
| 187 | } |
| 188 | |
| 189 | if (pad_high != 0) { |
| 190 | PaddingConfig padding_config; |
| 191 | for (xla::int64 i = 0; i < diag_shape.rank() - 1; ++i) { |
| 192 | auto* dims = padding_config.add_dimensions(); |
| 193 | dims->set_edge_padding_low(0); |
| 194 | dims->set_interior_padding(0); |
| 195 | dims->set_edge_padding_high(0); |
| 196 | } |
| 197 | auto* dims = padding_config.add_dimensions(); |
| 198 | dims->set_edge_padding_low(0); |
| 199 | dims->set_interior_padding(0); |
| 200 | dims->set_edge_padding_high(pad_high); |
| 201 | diag = Pad(diag, ScalarLike(diag, 0), padding_config); |
| 202 | } |
| 203 | |
| 204 | return Select(GetDiagonalMask(matrix, k), |
| 205 | BroadcastInDim(diag, shape.dimensions(), broadcast_dims), |
| 206 | matrix); |
| 207 | }); |
| 208 | } |
| 209 | |
| 210 | XlaOp TriangleMask(XlaOp x, int diagonal) { |
| 211 | XlaBuilder* builder = x.builder(); |