Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 15 additions & 8 deletions deepspeed/runtime/sequence_parallel/ulysses_sp.py
Original file line number Diff line number Diff line change
Expand Up @@ -444,19 +444,24 @@ def register_with_transformers(
mpu.initialize_sequence_parallel(sequence_parallel_size=sequence_parallel_size)

from transformers import PreTrainedModel
if hasattr(model_name_or_path, "config") or isinstance(model_name_or_path, PreTrainedModel):
model_was_loaded = hasattr(model_name_or_path, "config") or isinstance(model_name_or_path, PreTrainedModel)
if model_was_loaded:
# we already have the model (or a PEFT wrapper with config attribute)
hf_model_config = model_name_or_path.config
else:
# if we don't have the model yet at this stage
hf_model_config = AutoConfig.from_pretrained(model_name_or_path)

model_attn_implementation = getattr(hf_model_config, "_attn_implementation", None)
if model_attn_implementation is not None and model_attn_implementation != core_attn_implementation:
raise ValueError(
f"core_attn_implementation='{core_attn_implementation}' does not match "
f"model config attn_implementation='{model_attn_implementation}'. "
"Set both to the same value so sequence-parallel wrapper can intercept the active attention path.")
# Only a loaded model's config carries the attn implementation resolved at load time;
# a bare AutoConfig still holds the unresolved 'eager' default, so there is nothing
# meaningful to compare for a string model path.
if model_was_loaded:
model_attn_implementation = getattr(hf_model_config, "_attn_implementation", None)
if model_attn_implementation is not None and model_attn_implementation != core_attn_implementation:
raise ValueError(
f"core_attn_implementation='{core_attn_implementation}' does not match "
f"model config attn_implementation='{model_attn_implementation}'. "
"Set both to the same value so sequence-parallel wrapper can intercept the active attention path.")

# eager always materializes a 4D attention_mask (O(n²) memory) and cannot fall back
# to is_causal=True like sdpa — so it's incompatible with SP which discards masks.
Expand Down Expand Up @@ -666,7 +671,9 @@ def refill(self):
"Ensure your data collator includes position_ids in its output.")

# we have batches of variable seqlen so in order to do all_gather on batches - we need to know the exact length of each tensor on each rank
seqlen = torch.tensor(batch["input_ids"].shape[1], dtype=torch.int64, device=self.device)
# gloo validates gather shapes strictly, so send a 1-element tensor to match the
# receive list; a 0-dim scalar only passes on backends that move raw bytes.
seqlen = torch.full((1, ), batch["input_ids"].shape[1], dtype=torch.int64, device=self.device)
seqlens = [torch.zeros(1, dtype=torch.int64, device=self.device) for _ in range(self.sp_world_size)]
dist.all_gather(seqlens, seqlen, group=self.sp_group)
seqlens = [x[0].item() for x in seqlens]
Expand Down
16 changes: 11 additions & 5 deletions tests/unit/ulysses_alst/test_ulysses_sp_hf.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from unit.util import torch_assert_equal, torch_assert_close, torch_assert_dicts_of_tensors_equal
import deepspeed
import deepspeed.comm as dist
from deepspeed.accelerator import get_accelerator
import pytest
import torch

Expand Down Expand Up @@ -207,6 +208,10 @@ def test_ulysses_sp_hf_with_peft_model(self):
# Create a mock PEFT model object that has config but doesn't inherit from PreTrainedModel
from transformers import AutoConfig
hf_config = AutoConfig.from_pretrained(model_name_or_path)
# A real PEFT wrapper carries the base model's load-resolved attention implementation;
# a bare AutoConfig still holds the unresolved 'eager' default, which would trip the
# consistency check in register_with_transformers.
hf_config._attn_implementation = "sdpa"

class MockPEFTModel:
"""Mock PEFT model that simulates PeftModel behavior"""
Expand Down Expand Up @@ -255,16 +260,16 @@ def test_disable_in_eval(self):
micro_batch_size = 1

dtype = preferred_dtype()
rank = dist.get_rank()
device = get_accelerator().current_device_name()

# Full sequence input (not sharded) - this is what users would pass during eval
# when they want to bypass SP and process sequences independently per rank
input_ids = tensor([[1, 10, 10, 10, 2, 2]], device=f"cuda:{rank}")
position_ids = tensor([[0, 1, 2, 3, 4, 5]], device=f"cuda:{rank}")
input_ids = tensor([[1, 10, 10, 10, 2, 2]], device=device)
position_ids = tensor([[0, 1, 2, 3, 4, 5]], device=device)

# 1. Baseline: model without SP, processing full sequence
model_baseline = AutoModelForCausalLM.from_pretrained(model_name_or_path, torch_dtype=dtype)
model_baseline = model_baseline.to(f"cuda:{rank}")
model_baseline = model_baseline.to(device)
model_baseline.eval()

# Save original attention function for comparison
Expand Down Expand Up @@ -294,7 +299,7 @@ def test_disable_in_eval(self):
"register_with_transformers should have replaced the attention function"

model_sp = AutoModelForCausalLM.from_pretrained(model_name_or_path, torch_dtype=dtype)
model_sp = model_sp.to(f"cuda:{rank}")
model_sp = model_sp.to(device)
model_sp.eval()

with torch.no_grad():
Expand Down Expand Up @@ -390,6 +395,7 @@ def __init__(self, config):
)


@pytest.mark.skipif(get_accelerator().device_name() != 'cuda', reason="flex_attention requires CUDA")
@pytest.mark.parametrize("zero_stage", [2, 3])
class TestUlyssesSPHFFlexAttention(DistributedTest):
"""Separate class for flex_attention tests — requires non_daemonic_procs
Expand Down
Loading