(a, b, rtol=1e-3, atol=1e-3, max_error_count=0)
| 17 | |
| 18 | |
| 19 | def assert_most_approx_close(a, b, rtol=1e-3, atol=1e-3, max_error_count=0): |
| 20 | idx = torch.isclose(a, b, rtol=rtol, atol=atol) |
| 21 | error_count = (idx == 0).sum().item() |
| 22 | if error_count > max_error_count: |
| 23 | print(f"Too many values not close: assert {error_count} < {max_error_count}") |
| 24 | torch.testing.assert_close(a, b, rtol=rtol, atol=atol) |
| 25 | |
| 26 | |
| 27 | str2optimizers = {} |
no outgoing calls
no test coverage detected