MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / test

Method test

tests/unit/pipe/test_pipe_module.py:68–111  ·  view source on GitHub ↗
(self, sequential_model, simple_config, batch_input, activation_checkpoints, use_compile)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 13

compileMethod · 0.95
PipelineModuleClass · 0.90
get_acceleratorFunction · 0.90
RepeatingLoaderClass · 0.85
numelMethod · 0.80
initializeMethod · 0.80
configureMethod · 0.80
eval_batchMethod · 0.80
parametersMethod · 0.45
toMethod · 0.45
LongTensorMethod · 0.45
device_nameMethod · 0.45

Tested by

no test coverage detected