MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / OneSidedJacobiUpdate

Function OneSidedJacobiUpdate

tensorflow/compiler/xla/client/lib/svd.cc:451–563  ·  view source on GitHub ↗

Apply one-sided Jacobi on elements at indices pp, pq, qp, qq.

Source from the content-addressed store, hash-verified

449
450// Apply one-sided Jacobi on elements at indices pp, pq, qp, qq.
451StatusOr<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});

Callers 1

WhileLoopFnFunction · 0.85

Calls 15

ScalarLikeFunction · 0.85
BroadcastFunction · 0.85
DynamicSliceInMinorDimsFunction · 0.85
SquareFunction · 0.70
MulFunction · 0.50
RsqrtFunction · 0.50
ReduceFunction · 0.50
builderMethod · 0.45
rankMethod · 0.45
beginMethod · 0.45

Tested by

no test coverage detected