diff --git a/src/twinkle/model/megatron/megatron.py b/src/twinkle/model/megatron/megatron.py index 89d0a171d..cae3870d7 100644 --- a/src/twinkle/model/megatron/megatron.py +++ b/src/twinkle/model/megatron/megatron.py @@ -125,7 +125,7 @@ def __init__( self._model_path = HubOperation.download_model(model_id) self.tokenizer_id = kwargs.get('tokenizer_id', self.model_id) self._default_tokenizer = None - self.use_distributed_optimizer = kwargs.get('use_distributed_optimizer', True) + self.use_distributed_optimizer = kwargs.pop('use_distributed_optimizer', True) self.variable_seq_lengths = kwargs.get('variable_seq_lengths', True) torch_util.set_device() self._try_init_process_group() @@ -773,7 +773,7 @@ def _create_megatron_optimizer(self, **kwargs): - weight_decay: Weight decay (default: 0.0) - use_distributed_optimizer: Shard optimizer states (default: True) - clip_grad: Gradient clipping threshold (default: 1.0) - - bf16: Use bf16 training (default: True) + - bf16 / fp16: precision flags (default: derived from self.mixed_precision) - adam_beta1, adam_beta2, adam_eps: Adam parameters Returns: @@ -794,7 +794,8 @@ def _create_megatron_optimizer(self, **kwargs): adam_beta2=kwargs.pop('adam_beta2', 0.999), adam_eps=kwargs.pop('adam_eps', 1e-8), clip_grad=kwargs.pop('clip_grad', 1.0), - bf16=kwargs.pop('bf16', True), + bf16=kwargs.pop('bf16', self.mixed_precision == 'bf16'), + fp16=kwargs.pop('fp16', self.mixed_precision == 'fp16'), use_distributed_optimizer=self.use_distributed_optimizer, overlap_param_gather=kwargs.pop('overlap_param_gather', False), log_num_zeros_in_grad=kwargs.pop('log_num_zeros_in_grad', False), diff --git a/src/twinkle/processor/base.py b/src/twinkle/processor/base.py index ae37dbdd1..0a819e1ca 100644 --- a/src/twinkle/processor/base.py +++ b/src/twinkle/processor/base.py @@ -70,8 +70,8 @@ def __init__(self, self.framework = framework self.process_pipeline = [ self.prepare_inputs, - self.align_routed_experts, self.pad_cp, + self.align_routed_experts, self.collate_fn, self.to_transformers_dict, self.add_extra_padding_free_args, @@ -827,10 +827,8 @@ def align_routed_experts(self, inputs: Union[List[InputFeature], InputFeature], def align_to(_input): routed_experts = _input.get('routed_experts', None) - input_seq_len = _input.get('length', None) - if input_seq_len is None: - input_ids = _input.get('input_ids', None) - input_seq_len = input_ids.shape[1] if input_ids is not None else 0 + input_ids = _input.get('input_ids', None) + input_seq_len = input_ids.shape[-1] if input_ids is not None else 0 if routed_experts is not None: # The number of experts in the output can be 1 less than (prompt_length + response_token_count) # This gap of 1 is expected