(self, s, u, v)
| 32 | class SvdOpTest(xla_test.XLATestCase, parameterized.TestCase): |
| 33 | |
| 34 | def _compute_usvt(self, s, u, v): |
| 35 | m = u.shape[-1] |
| 36 | n = v.shape[-1] |
| 37 | if m <= n: |
| 38 | v = v[..., :m] |
| 39 | else: |
| 40 | u = u[..., :n] |
| 41 | |
| 42 | return np.matmul(u * s[..., None, :], np.swapaxes(v, -1, -2)) |
| 43 | |
| 44 | def _testSvdCorrectness(self, dtype, shape): |
| 45 | np.random.seed(1) |
no test coverage detected