(self, sequential_model, simple_config, batch_input, activation_checkpoints, use_compile)
| 66 | @pytest.mark.parametrize("activation_checkpoints", [False, True]) |
| 67 | @pytest.mark.parametrize("use_compile", [False, True]) |
| 68 | def test(self, sequential_model, simple_config, batch_input, activation_checkpoints, use_compile): |
| 69 | base_model = copy.deepcopy(sequential_model) |
| 70 | base_input = batch_input.clone().detach() |
| 71 | base_output = base_model(base_input) |
| 72 | base_output = base_output |
| 73 | base_params = sum(p.numel() for p in base_model.parameters()) |
| 74 | |
| 75 | pipe_model = copy.deepcopy(sequential_model) |
| 76 | pipe_model = PipelineModule(layers=pipe_model, num_stages=2) |
| 77 | if (use_compile): |
| 78 | pipe_model.compile() |
| 79 | # Ensure all parameters are accounted for. |
| 80 | my_params = sum(p.numel() for p in pipe_model.parameters()) |
| 81 | total_pipe_params = torch.LongTensor([my_params]).to(get_accelerator().device_name()) |
| 82 | dist.all_reduce(total_pipe_params) |
| 83 | total_pipe_params = total_pipe_params.item() |
| 84 | assert total_pipe_params == base_params |
| 85 | |
| 86 | pipe_model, _, _, _ = deepspeed.initialize(config=simple_config, |
| 87 | model=pipe_model, |
| 88 | model_parameters=[p for p in pipe_model.parameters()]) |
| 89 | |
| 90 | if activation_checkpoints: |
| 91 | deepspeed.checkpointing.configure(None, |
| 92 | deepspeed_config=pipe_model.config, |
| 93 | partition_activations=True, |
| 94 | contiguous_checkpointing=True, |
| 95 | num_checkpoints=9) |
| 96 | |
| 97 | if pipe_model.is_first_stage or pipe_model.is_last_stage: |
| 98 | pipe_input = base_input.clone().detach().to(get_accelerator().device_name()) |
| 99 | # label 0 is meaningless |
| 100 | dataset = [(pipe_input, 0)] |
| 101 | loader = RepeatingLoader(dataset) |
| 102 | data_iter = iter(loader) |
| 103 | else: |
| 104 | data_iter = None |
| 105 | |
| 106 | pipe_output = pipe_model.eval_batch(data_iter=data_iter) |
| 107 | |
| 108 | base_output = base_output.to('cpu') |
| 109 | pipe_output = pipe_output.to('cpu') |
| 110 | |
| 111 | assert torch.allclose(base_output, pipe_output, atol=1e-4) |
nothing calls this directly
no test coverage detected