Restrict Transformers to data-parallel replicas
Browse files- configuration_limite.py +0 -12
- 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)
|