MisterOss commited on
Commit
eae5ed4
·
verified ·
1 Parent(s): b1f3d57

Restrict Transformers to data-parallel replicas

Browse files
Files changed (2) hide show
  1. configuration_limite.py +0 -12
  2. modeling_limite.py +40 -0
configuration_limite.py CHANGED
@@ -129,18 +129,6 @@ class LimiteConfig(PretrainedConfig):
129
  model_type = MODEL_TYPE
130
  keys_to_ignore_at_inference = ["past_key_values"]
131
 
132
- base_model_tp_plan = {
133
- "layers.*.self_attn.qkv_proj": "colwise",
134
- "layers.*.self_attn.o_proj": "rowwise",
135
- "layers.*.mlp.gate_up_proj": "packed_colwise",
136
- "layers.*.mlp.down_proj": "rowwise",
137
- }
138
- base_model_pp_plan = {
139
- "embed_tokens": (["input_ids"], ["inputs_embeds"]),
140
- "layers": (["hidden_states"], ["hidden_states"]),
141
- "norm": (["hidden_states"], ["hidden_states"]),
142
- }
143
-
144
  @classmethod
145
  def from_dict(cls, config_dict: dict[str, Any], **kwargs: Any) -> "LimiteConfig":
146
  missing = [
 
129
  model_type = MODEL_TYPE
130
  keys_to_ignore_at_inference = ["past_key_values"]
131
 
 
 
 
 
 
 
 
 
 
 
 
 
132
  @classmethod
133
  def from_dict(cls, config_dict: dict[str, Any], **kwargs: Any) -> "LimiteConfig":
134
  missing = [
modeling_limite.py CHANGED
@@ -671,6 +671,46 @@ class LimitePreTrainedModel(PreTrainedModel):
671
  _supports_attention_backend = True
672
  _can_compile_fullgraph = True
673
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
674
  @torch.no_grad()
675
  def _init_weights(self, module: nn.Module) -> None:
676
  super()._init_weights(module)
 
671
  _supports_attention_backend = True
672
  _can_compile_fullgraph = True
673
 
674
+ @classmethod
675
+ def from_pretrained(
676
+ cls,
677
+ pretrained_model_name_or_path: str | None,
678
+ *model_args: Any,
679
+ **kwargs: Any,
680
+ ) -> LimitePreTrainedModel:
681
+ requested_parallelism = [
682
+ name
683
+ for name in ("tp_plan", "tp_size", "distributed_config")
684
+ if kwargs.get(name) is not None
685
+ ]
686
+ device_map = kwargs.get("device_map")
687
+ if isinstance(device_map, str) and device_map in {
688
+ "auto",
689
+ "balanced",
690
+ "balanced_low_0",
691
+ "sequential",
692
+ }:
693
+ requested_parallelism.append(f"device_map={device_map!r}")
694
+ elif isinstance(device_map, dict):
695
+ placements = {str(device) for device in device_map.values()}
696
+ if len(placements) > 1:
697
+ requested_parallelism.append("multi-device device_map")
698
+
699
+ if requested_parallelism:
700
+ requested = ", ".join(requested_parallelism)
701
+ raise NotImplementedError(
702
+ "Limite supports one complete model replica per process; "
703
+ "tensor parallelism, pipeline parallelism, and multi-device "
704
+ f"model sharding are not supported (requested: {requested}). "
705
+ "Use process-level data parallelism with one explicit device "
706
+ "per replica."
707
+ )
708
+ return super().from_pretrained(
709
+ pretrained_model_name_or_path,
710
+ *model_args,
711
+ **kwargs,
712
+ )
713
+
714
  @torch.no_grad()
715
  def _init_weights(self, module: nn.Module) -> None:
716
  super()._init_weights(module)