Apply one-sided Jacobi on elements at indices pp, pq, qp, qq.
| 449 | |
| 450 | // Apply one-sided Jacobi on elements at indices pp, pq, qp, qq. |
| 451 | StatusOr<SVDResult> OneSidedJacobiUpdate(SVDResult svd_result, XlaOp p, XlaOp q, |
| 452 | XlaOp eps) { |
| 453 | XlaOp u = svd_result.u; |
| 454 | XlaOp v = svd_result.v; |
| 455 | XlaOp d = svd_result.d; |
| 456 | XlaBuilder* builder = d.builder(); |
| 457 | TF_ASSIGN_OR_RETURN(Shape d_shape, builder->GetShape(d)); |
| 458 | const int64 num_dims = d_shape.rank(); |
| 459 | const int64 num_batch_dims = num_dims - 2; |
| 460 | std::vector<int64> batch_dims(num_batch_dims); |
| 461 | for (int i = 0; i < num_batch_dims; ++i) { |
| 462 | batch_dims[i] = ShapeUtil::GetDimension(d_shape, i); |
| 463 | } |
| 464 | const int64 m = ShapeUtil::GetDimension(d_shape, -2); |
| 465 | const int64 n = ShapeUtil::GetDimension(d_shape, -1); |
| 466 | |
| 467 | TF_ASSIGN_OR_RETURN(OneSidedJacobiRotation onesided_jacobi, |
| 468 | GetOneSidedJacobiRotation(d, p, q, eps)); |
| 469 | |
| 470 | auto zero = ScalarLike(p, 0); |
| 471 | |
| 472 | // Zero out a_{pq} explicitly. |
| 473 | std::vector<int64> pq_dims(batch_dims.begin(), batch_dims.end()); |
| 474 | pq_dims.push_back(1); |
| 475 | pq_dims.push_back(1); |
| 476 | auto pq_zero = ScalarLike(d, 0.0); |
| 477 | auto pq_zeros = Broadcast(pq_zero, pq_dims); |
| 478 | |
| 479 | std::vector<int64> broadcast_dims(batch_dims.size()); |
| 480 | std::iota(broadcast_dims.begin(), broadcast_dims.end(), 0); |
| 481 | broadcast_dims.push_back(num_dims - 1); |
| 482 | |
| 483 | // Apply Jacobi Rotation on the left. |
| 484 | auto slice_p = DynamicSliceInMinorDims(d, {p, zero}, {1, n}); |
| 485 | auto slice_q = DynamicSliceInMinorDims(d, {q, zero}, {1, n}); |
| 486 | auto slice_p_new = |
| 487 | onesided_jacobi.rot_l.c * slice_p - onesided_jacobi.rot_l.s * slice_q; |
| 488 | auto slice_q_new = |
| 489 | onesided_jacobi.rot_l.s * slice_p + onesided_jacobi.rot_l.c * slice_q; |
| 490 | d = DynamicUpdateSliceInMinorDims(d, slice_p_new, {p, zero}); |
| 491 | d = DynamicUpdateSliceInMinorDims(d, slice_q_new, {q, zero}); |
| 492 | |
| 493 | // Apply Jacobi Rotation on the right. |
| 494 | slice_p = DynamicSliceInMinorDims(d, {zero, p}, {m, 1}); |
| 495 | slice_q = DynamicSliceInMinorDims(d, {zero, q}, {m, 1}); |
| 496 | slice_p_new = |
| 497 | onesided_jacobi.rot_r.c * slice_p - onesided_jacobi.rot_r.s * slice_q; |
| 498 | slice_q_new = |
| 499 | onesided_jacobi.rot_r.s * slice_p + onesided_jacobi.rot_r.c * slice_q; |
| 500 | d = DynamicUpdateSliceInMinorDims(d, slice_p_new, {zero, p}); |
| 501 | d = DynamicUpdateSliceInMinorDims(d, slice_q_new, {zero, q}); |
| 502 | |
| 503 | d = DynamicUpdateSliceInMinorDims(d, pq_zeros, {p, q}); |
| 504 | d = DynamicUpdateSliceInMinorDims(d, pq_zeros, {q, p}); |
| 505 | |
| 506 | // Apply left Jacobi Rotation on U. |
| 507 | slice_p = DynamicSliceInMinorDims(u, {zero, p}, {m, 1}); |
| 508 | slice_q = DynamicSliceInMinorDims(u, {zero, q}, {m, 1}); |
no test coverage detected