| 20 | |
| 21 | |
| 22 | class FakeDeepSpeedEngine: |
| 23 | def __init__(self, model, optimizer): |
| 24 | self.model = model |
| 25 | self.optimizer = optimizer |
| 26 | self.forward_calls = 0 |
| 27 | self.backward_calls = 0 |
| 28 | self.step_calls = 0 |
| 29 | self.zero_grad_calls = 0 |
| 30 | |
| 31 | def __call__(self, *args, **kwargs): |
| 32 | self.forward_calls += 1 |
| 33 | return self.model(*args, **kwargs) |
| 34 | |
| 35 | def zero_grad(self, *args, **kwargs): |
| 36 | self.zero_grad_calls += 1 |
| 37 | self.optimizer.zero_grad(*args, **kwargs) |
| 38 | |
| 39 | def backward(self, loss): |
| 40 | self.backward_calls += 1 |
| 41 | loss.backward() |
| 42 | |
| 43 | def step(self): |
| 44 | self.step_calls += 1 |
| 45 | self.optimizer.step() |
| 46 | |
| 47 | |
| 48 | class FakeDeepSpeedModule: |