(self, zero_stage)
| 36 | world_size = 1 |
| 37 | |
| 38 | def test(self, zero_stage): |
| 39 | if zero_stage == 3 and not required_torch_version(min_version=1.8): |
| 40 | pytest.skip("zero-3 param offload requires at least torch 1.8") |
| 41 | |
| 42 | ds_config = { |
| 43 | 'train_batch_size': self.world_size, |
| 44 | 'zero_optimization': { |
| 45 | "stage": zero_stage, |
| 46 | "offload_param": { |
| 47 | "device": "cpu" |
| 48 | } |
| 49 | } |
| 50 | } |
| 51 | if get_accelerator().is_bf16_supported(): |
| 52 | ds_config["bf16"] = {"enabled": True} |
| 53 | elif get_accelerator().is_fp16_supported(): |
| 54 | ds_config["fp16"] = {"enabled": True} |
| 55 | # 20B test |
| 56 | #hidden_dim = 16 * 1024 |
| 57 | hidden_dim = 4 |
| 58 | |
| 59 | with deepspeed.zero.Init(enabled=zero_stage == 3, config_dict_or_path=ds_config): |
| 60 | model = SimpleModel(hidden_dim, nlayers=78) |
| 61 | see_memory_usage('pre-init', force=True) |
| 62 | model, _, _, _ = deepspeed.initialize(model=model, config=ds_config) |
| 63 | see_memory_usage('post-init', force=True) |
| 64 | data_loader = random_dataloader(model=model, total_samples=50, hidden_dim=hidden_dim, device=model.device) |
| 65 | for batch in data_loader: |
| 66 | model(batch[0], batch[1]) |
| 67 | see_memory_usage('post-fwds', force=True) |
| 68 | |
| 69 | |
| 70 | @pytest.mark.parametrize('optimizer_type', [None, Optimizer, Callable]) |
nothing calls this directly
no test coverage detected