From 22300586147669a4e6a19d24fb0fcb379b16259c Mon Sep 17 00:00:00 2001 From: Vitaliy Chiley Date: Fri, 7 Jul 2023 00:02:27 +0000 Subject: [PATCH 1/3] adding sp to te ln_mlp --- llmfoundry/models/layers/ffn.py | 42 ++++++++++++++++++++++- llmfoundry/models/mpt/modeling_mpt.py | 2 ++ llmfoundry/models/utils/param_init_fns.py | 12 +++++++ 3 files changed, 55 insertions(+), 1 deletion(-) diff --git a/llmfoundry/models/layers/ffn.py b/llmfoundry/models/layers/ffn.py index df7a2ccffd..6521ec0e6d 100644 --- a/llmfoundry/models/layers/ffn.py +++ b/llmfoundry/models/layers/ffn.py @@ -3,10 +3,14 @@ """GPT Blocks used for the GPT Model.""" +import warnings from typing import Optional import torch import torch.nn as nn +from torch import distributed + +from composer.utils import dist from llmfoundry.models.layers.attention import ATTN_CLASS_REGISTRY from llmfoundry.models.layers.fc import FC_CLASS_REGISTRY @@ -75,10 +79,46 @@ def build_ffn( device=device, ) elif ffn_type == 'te_ln_mlp': - return te.LayerNormMLP( + parallel_mode = kwargs.get('set_parallel_mode', False) + if parallel_mode: + if not kwargs.get('sequence_parallel', False): + warnings.warn( + 'Unexpected usage: te.LayerNormMLP args are `set_parallel_mode: true` and `sequence_parallel: false`.' + ) + tp_group = kwargs.get('tp_group', None) + tp_size = kwargs.get('tp_size', 1) + if tp_group is None and tp_size == 1: + warnings.warn(f'tp (sp) not configured correctly and therefore will be disabled.') + # kwargs.pop('set_parallel_mode', None) + # kwargs.pop('sequence_parallel', None) + # kwargs.pop('tp_group', None) + # kwargs.pop('tp_size', None) + + if tp_group is None: # and tp_size != 1: + world_size = dist.get_world_size() + if world_size % tp_size != 0: + raise RuntimeError(f'{world_size} must be divisible by {tp_size=}.') + start = dist.get_global_rank() // tp_size * tp_size + ranks = tuple(range(start, start + tp_size)) + ranks_per_subgroup_list = list(set(dist.all_gather_object(ranks))) + current_group, _subgroups = distributed.distributed_c10d.new_subgroups_by_enumeration(ranks_per_subgroup_list) + tp_group = current_group + kwargs['tp_group'] = tp_group + + # if tp_group is not None and tp_size == 1: + # # TODO init tp_group + # tp_size = tp_group.size() + # kwargs['tp_size'] = tp_size + + mlp = te.LayerNormMLP( hidden_size=d_model, ffn_hidden_size=d_model * expansion_ratio, **kwargs, ) + if parallel_mode: + mlp._fsdp_process_group = f"mod{kwargs.get('tp_size')}" + + return mlp + raise ValueError(f'{ffn_type=} not recognized.') diff --git a/llmfoundry/models/mpt/modeling_mpt.py b/llmfoundry/models/mpt/modeling_mpt.py index d67b01c9fe..c1345ef0d2 100644 --- a/llmfoundry/models/mpt/modeling_mpt.py +++ b/llmfoundry/models/mpt/modeling_mpt.py @@ -476,6 +476,8 @@ def param_init_fn(self, module): # FSDP Wrap function def fsdp_wrap_fn(self, module): + if hasattr(module, '_fsdp_process_group'): + return {'process_group': module._fsdp_process_group} return isinstance(module, MPTBlock) # Activation Checkpointing diff --git a/llmfoundry/models/utils/param_init_fns.py b/llmfoundry/models/utils/param_init_fns.py index 3df6387b1d..f3b29072d2 100644 --- a/llmfoundry/models/utils/param_init_fns.py +++ b/llmfoundry/models/utils/param_init_fns.py @@ -209,6 +209,18 @@ def generic_param_init_fn_( if module.fc2_bias is not None: torch.nn.init.zeros_(module.fc2_bias) + if module.tp_size > 1: + if 'kaiming_' in init_fn_.func.__name__: + with torch.no_grad(): + if init_fn_.keywords.get('mode', 'fan_in') == 'fan_in': + module.fc1_weight.div_(math.sqrt(module.tp_size)) + else: + module.fc2_weight.div_(math.sqrt(module.tp_size)) + else: + warnings.warn( + f'te.LayerNormMLP layer is using tp; init_fn ({init_fn_.func.__name__}) not being adjusted for TP split.' + ) + with torch.no_grad(): module.fc2_weight.div_(div_is_residual) From 5ea969a4d8cbb11794fd9e6b8e68d06a6586dde2 Mon Sep 17 00:00:00 2001 From: Vitaliy Chiley Date: Fri, 7 Jul 2023 00:09:44 +0000 Subject: [PATCH 2/3] updt --- llmfoundry/models/layers/ffn.py | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/llmfoundry/models/layers/ffn.py b/llmfoundry/models/layers/ffn.py index 6521ec0e6d..73ee71213a 100644 --- a/llmfoundry/models/layers/ffn.py +++ b/llmfoundry/models/layers/ffn.py @@ -89,12 +89,12 @@ def build_ffn( tp_size = kwargs.get('tp_size', 1) if tp_group is None and tp_size == 1: warnings.warn(f'tp (sp) not configured correctly and therefore will be disabled.') - # kwargs.pop('set_parallel_mode', None) - # kwargs.pop('sequence_parallel', None) - # kwargs.pop('tp_group', None) - # kwargs.pop('tp_size', None) + kwargs.pop('set_parallel_mode', None) + kwargs.pop('sequence_parallel', None) + kwargs.pop('tp_group', None) + kwargs.pop('tp_size', None) - if tp_group is None: # and tp_size != 1: + if tp_group is None and tp_size != 1: world_size = dist.get_world_size() if world_size % tp_size != 0: raise RuntimeError(f'{world_size} must be divisible by {tp_size=}.') @@ -105,10 +105,9 @@ def build_ffn( tp_group = current_group kwargs['tp_group'] = tp_group - # if tp_group is not None and tp_size == 1: - # # TODO init tp_group - # tp_size = tp_group.size() - # kwargs['tp_size'] = tp_size + if tp_group is not None and tp_size == 1: + tp_size = tp_group.size() + kwargs['tp_size'] = tp_size mlp = te.LayerNormMLP( hidden_size=d_model, @@ -117,7 +116,7 @@ def build_ffn( ) if parallel_mode: - mlp._fsdp_process_group = f"mod{kwargs.get('tp_size')}" + mlp._fsdp_process_group = f"mod{tp_size}" return mlp From f9beaea20c2b880872f1bdef516205b12578795a Mon Sep 17 00:00:00 2001 From: Vitaliy Chiley Date: Sat, 8 Jul 2023 03:52:58 +0000 Subject: [PATCH 3/3] fix n_active_params with tp --- llmfoundry/models/layers/ffn.py | 45 +++++++++++++++++++++------ llmfoundry/models/mpt/modeling_mpt.py | 17 +++++++++- scripts/train/train.py | 5 ++- 3 files changed, 56 insertions(+), 11 deletions(-) diff --git a/llmfoundry/models/layers/ffn.py b/llmfoundry/models/layers/ffn.py index 69e0c399d2..da3e1fb8b5 100644 --- a/llmfoundry/models/layers/ffn.py +++ b/llmfoundry/models/layers/ffn.py @@ -8,9 +8,8 @@ import torch import torch.nn as nn -from torch import distributed - from composer.utils import dist +from torch import distributed from llmfoundry.models.layers.attention import ATTN_CLASS_REGISTRY from llmfoundry.models.layers.fc import FC_CLASS_REGISTRY @@ -57,6 +56,29 @@ def forward(self, x): } if te is not None: + + def te_ln_mlp_n_params(parent_cls): + n_params = 0 + for m_n, m in parent_cls.named_modules(): + for p_n, p in m.named_parameters(): + if '.' not in p_n: + # local params + if isinstance(m, te.LayerNormMLP): + if p_n in [ + 'layer_norm_weight', 'layer_norm_bias', + 'fc2_bias' + ]: + n_params += p.numel() + elif p_n in ['fc1_weight', 'fc1_bias', 'fc2_weight']: + n_params += (p.numel() * m.tp_size) + else: + RuntimeError(f'te_ln_mlp_n_params fn has error.') + else: + n_params += p.numel() + return n_params + + te.LayerNormMLP.parent_n_active_params = staticmethod(te_ln_mlp_n_params) + te.LayerNormMLP._has_norm = True FFN_CLASS_REGISTRY['te_ln_mlp'] = te.LayerNormMLP @@ -89,23 +111,28 @@ def build_ffn( tp_group = kwargs.get('tp_group', None) tp_size = kwargs.get('tp_size', 1) if tp_group is None and tp_size == 1: - warnings.warn(f'tp (sp) not configured correctly and therefore will be disabled.') + warnings.warn( + f'tp (sp) not configured correctly and therefore will be disabled.' + ) kwargs.pop('set_parallel_mode', None) kwargs.pop('sequence_parallel', None) kwargs.pop('tp_group', None) kwargs.pop('tp_size', None) - + if tp_group is None and tp_size != 1: world_size = dist.get_world_size() if world_size % tp_size != 0: - raise RuntimeError(f'{world_size} must be divisible by {tp_size=}.') + raise RuntimeError( + f'{world_size} must be divisible by {tp_size=}.') start = dist.get_global_rank() // tp_size * tp_size ranks = tuple(range(start, start + tp_size)) - ranks_per_subgroup_list = list(set(dist.all_gather_object(ranks))) - current_group, _subgroups = distributed.distributed_c10d.new_subgroups_by_enumeration(ranks_per_subgroup_list) + ranks_per_subgroup_list = list( + set(dist.all_gather_object(ranks))) + current_group, _subgroups = distributed.distributed_c10d.new_subgroups_by_enumeration( + ranks_per_subgroup_list) tp_group = current_group kwargs['tp_group'] = tp_group - + if tp_group is not None and tp_size == 1: tp_size = tp_group.size() kwargs['tp_size'] = tp_size @@ -117,7 +144,7 @@ def build_ffn( ) if parallel_mode: - mlp._fsdp_process_group = f"mod{tp_size}" + mlp._fsdp_process_group = f'mod{tp_size}' return mlp diff --git a/llmfoundry/models/mpt/modeling_mpt.py b/llmfoundry/models/mpt/modeling_mpt.py index acb06e4a68..2b671d1308 100644 --- a/llmfoundry/models/mpt/modeling_mpt.py +++ b/llmfoundry/models/mpt/modeling_mpt.py @@ -708,7 +708,7 @@ def __init__( allow_embedding_resizing=True, ) - self.n_active_params = sum(p.numel() for p in self.parameters()) + self._n_active_params = None loss_fn_config = om_model_config.get('loss_fn', 'fused_crossentropy') if loss_fn_config == 'fused_crossentropy': @@ -753,6 +753,21 @@ def loss(self, outputs, batch): return self.loss_fn(outputs.logits.view(-1, outputs.logits.size(-1)), targets.view(-1)) + @property + def n_active_params(self): + if self._n_active_params is not None: + return self._n_active_params + + ffn_type = self.model.config.ffn_config['ffn_type'] + if ffn_type == 'te_ln_mlp': + _cls = FFN_CLASS_REGISTRY[ffn_type] + self._n_active_params = _cls.parent_n_active_params(self) + else: + # default behavior + self._n_active_params = sum(p.numel() for p in self.parameters()) + + return self._n_active_params + def flops_per_batch(self, batch): # Note: this computation does not take into account padding, and assumes # that the dataset has been constructed without padding. Additionally, we diff --git a/scripts/train/train.py b/scripts/train/train.py index 188c0cdd32..b2b243b275 100644 --- a/scripts/train/train.py +++ b/scripts/train/train.py @@ -239,7 +239,10 @@ def main(cfg): print_trainable_parameters(model) # should not be 100% else: # standard model model = build_composer_model(cfg.model, tokenizer) - cfg.n_params = sum(p.numel() for p in model.parameters()) + if hasattr(model, 'n_active_params'): + cfg.n_params = model.n_active_params + else: + cfg.n_params = sum(p.numel() for p in model.parameters()) print(f'{cfg.n_params=:.2e}') # Dataloaders