mirror of
https://github.com/hpcaitech/ColossalAI.git
synced 2025-09-06 19:40:28 +00:00
[test] fixed gemini plugin test (#3411)
* [test] fixed gemini plugin test * polish code * polish code
This commit is contained in:
@@ -21,9 +21,6 @@ def check_gemini_plugin(early_stop: bool = True):
|
||||
Args:
|
||||
early_stop (bool, optional): Whether to stop when getting the first error. Defaults to True.
|
||||
"""
|
||||
plugin = GeminiPlugin(placement_policy='cuda', strict_ddp_mode=True, max_norm=1.0, initial_scale=2**5)
|
||||
booster = Booster(plugin=plugin)
|
||||
|
||||
passed_models = []
|
||||
failed_info = {} # (model_name, error) pair
|
||||
|
||||
@@ -34,46 +31,23 @@ def check_gemini_plugin(early_stop: bool = True):
|
||||
continue
|
||||
# These models are not compatible with gemini
|
||||
if name in [
|
||||
'diffusers_clip_vision_model',
|
||||
'timm_resnet',
|
||||
'timm_beit',
|
||||
'timm_beitv2',
|
||||
'timm_eca_nfnet',
|
||||
'timm_efficientformer',
|
||||
'timm_hrnet_w18_small',
|
||||
'timm_nf_ecaresnet101',
|
||||
'timm_nf_regnet_b0',
|
||||
'timm_skresnet18',
|
||||
'timm_wide_resnet50_2',
|
||||
'timm_convit',
|
||||
'timm_dm_nfnet',
|
||||
'timm_swin_transformer',
|
||||
'torchaudio_conformer',
|
||||
'torchaudio_deepspeech',
|
||||
'torchaudio_wavernn',
|
||||
'torchaudio_tacotron',
|
||||
'deepfm_interactionarch',
|
||||
'deepfm_simpledeepfmnn',
|
||||
'dlrm',
|
||||
'dlrm_interactionarch',
|
||||
'torchvision_googlenet',
|
||||
'torchvision_inception_v3',
|
||||
'torchvision_mobilenet_v3_small',
|
||||
'torchvision_resnet18',
|
||||
'torchvision_resnext50_32x4d',
|
||||
'torchvision_wide_resnet50_2',
|
||||
'torchvision_vit_b_16',
|
||||
'torchvision_convnext_base',
|
||||
'torchvision_swin_s',
|
||||
'transformers_albert',
|
||||
'transformers_albert_for_pretraining',
|
||||
'transformers_bert',
|
||||
'transformers_bert_for_pretraining',
|
||||
'transformers_gpt_double_heads',
|
||||
'torchaudio_hubert_base',
|
||||
'diffusers_clip_vision_model', 'timm_resnet', 'timm_beit', 'timm_beitv2', 'timm_eca_nfnet',
|
||||
'timm_efficientformer', 'timm_hrnet_w18_small', 'timm_nf_ecaresnet101', 'timm_nf_regnet_b0',
|
||||
'timm_skresnet18', 'timm_wide_resnet50_2', 'timm_convit', 'timm_dm_nfnet', 'timm_swin_transformer',
|
||||
'torchaudio_conformer', 'torchaudio_deepspeech', 'torchaudio_wavernn', 'torchaudio_tacotron',
|
||||
'deepfm_interactionarch', 'deepfm_simpledeepfmnn', 'dlrm', 'dlrm_interactionarch',
|
||||
'torchvision_googlenet', 'torchvision_inception_v3', 'torchvision_mobilenet_v3_small',
|
||||
'torchvision_resnet18', 'torchvision_resnext50_32x4d', 'torchvision_wide_resnet50_2',
|
||||
'torchvision_vit_b_16', 'torchvision_convnext_base', 'torchvision_swin_s', 'transformers_albert',
|
||||
'transformers_albert_for_pretraining', 'transformers_bert', 'transformers_bert_for_pretraining',
|
||||
'transformers_gpt_double_heads', 'torchaudio_hubert_base', 'torchaudio_wav2vec2_base',
|
||||
'transformers_t5_for_conditional_generation', 'transformers_t5', 'transformers_t5_encoder_model'
|
||||
]:
|
||||
continue
|
||||
|
||||
try:
|
||||
plugin = GeminiPlugin(placement_policy='cuda', strict_ddp_mode=True, max_norm=1.0, initial_scale=2**5)
|
||||
booster = Booster(plugin=plugin)
|
||||
model = model_fn()
|
||||
optimizer = HybridAdam(model.parameters(), lr=1e-3)
|
||||
criterion = lambda x: x.mean()
|
||||
@@ -97,10 +71,15 @@ def check_gemini_plugin(early_stop: bool = True):
|
||||
booster.backward(loss, optimizer)
|
||||
optimizer.step()
|
||||
passed_models.append(name)
|
||||
|
||||
del booster, plugin, model, optimizer, criterion, data, output, loss
|
||||
except Exception as e:
|
||||
failed_info[name] = e
|
||||
if early_stop:
|
||||
raise e
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if dist.get_rank() == 0:
|
||||
print(f'Passed models({len(passed_models)}): {passed_models}\n\n')
|
||||
print(f'Failed models({len(failed_info)}): {list(failed_info.keys())}\n\n')
|
||||
@@ -138,7 +117,6 @@ def run_dist(rank, world_size, port, early_stop: bool = True):
|
||||
check_gemini_plugin(early_stop=early_stop)
|
||||
|
||||
|
||||
@pytest.mark.skip(reason='Skip gemini plugin test due to OOM')
|
||||
@rerun_if_address_is_in_use()
|
||||
def test_gemini_plugin(early_stop: bool = True):
|
||||
world_size = 2
|
||||
|
Reference in New Issue
Block a user