()
| 25 | |
| 26 | |
| 27 | def sanity_check(): |
| 28 | from bitsandbytes.optim import Adam |
| 29 | |
| 30 | p = torch.nn.Parameter(torch.rand(10, 10).cuda()) |
| 31 | a = torch.rand(10, 10).cuda() |
| 32 | p1 = p.data.sum().item() |
| 33 | adam = Adam([p]) |
| 34 | out = a * p |
| 35 | loss = out.sum() |
| 36 | loss.backward() |
| 37 | adam.step() |
| 38 | p2 = p.data.sum().item() |
| 39 | assert p1 != p2 |
| 40 | |
| 41 | |
| 42 | def get_package_version(name: str) -> str: |