Text Generation
Transformers
Safetensors
PyTorch
nemotron_h
nvidia
nemotron-3
latent-moe
mtp
conversational
custom_code
Eval Results

generation using transformers hit device mismatch issue

#23
by shengliangx - opened
NVIDIA org

Test environment:

transformers                  5.4.0
accelerate                    1.13.0

Nvidia B200 8 GPUs

test code:

import torch
from transformers import AutoTokenizer, AutoModelForCausalLM

model_path = "nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16"
tokenizer = AutoTokenizer.from_pretrained(model_path)

model = AutoModelForCausalLM.from_pretrained(
    model_path ,
    torch_dtype=torch.bfloat16,
    device_map="auto"
)

messages = [
    {"role": "user", "content": "Write a haiku about GPUs"},
]

tokenized_chat = tokenizer.apply_chat_template(
    messages,
    tokenize=True,
    add_generation_prompt=True,
    return_tensors="pt"
).to(model.device)

if not isinstance(tokenized_chat, torch.Tensor):
    input_ids = tokenized_chat["input_ids"]
else:
    input_ids = tokenized_chat

with torch.backends.cuda.sdp_kernel(
    enable_flash=True,
    enable_math=False,
    enable_cudnn=False
):
    outputs = model.generate(
        input_ids,
        max_new_tokens=50,
        temperature=1.0,
        top_p=0.95,
        eos_token_id=tokenizer.eos_token_id
    )

print(tokenizer.decode(outputs[0]))

error:

  Traceback (most recent call last):
    File "/workspace/test.py", line 34, in <module>
      outputs = model.generate(
                ^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/utils/_contextlib.py", line 124, in decorate_context
      return func(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/generation/utils.py", line 2521, in generate
      result = decoding_method(
               ^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/generation/utils.py", line 2728, in _sample
      outputs = model_forward(**model_inputs, return_dict=True)
                ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
      return self._call_impl(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
      return forward_call(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/accelerate/hooks.py", line 192, in new_forward
      output = module._old_forward(*args, **kwargs)
               ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/utils/generic.py", line 857, in wrapper
      output = func(self, *args, **kwargs)
               ^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/models/nemotron_h/modeling_nemotron_h.py", line 1292, in forward
      outputs = self.model(
                ^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
      return self._call_impl(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
      return forward_call(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/utils/generic.py", line 931, in wrapper
      output = func(self, *args, **kwargs)
               ^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/utils/output_capturing.py", line 248, in wrapper
      outputs = func(self, *args, **kwargs)
                ^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/models/nemotron_h/modeling_nemotron_h.py", line 1210, in forward
      hidden_states = mixer_block(
                      ^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/modeling_layers.py", line 93, in __call__
      return super().__call__(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
      return self._call_impl(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
      return forward_call(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/accelerate/hooks.py", line 192, in new_forward
      output = module._old_forward(*args, **kwargs)
               ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/models/nemotron_h/modeling_nemotron_h.py", line 1049, in forward
      hidden_states = self.mixer(hidden_states, cache_params=past_key_values, attention_mask=attention_mask)
                      ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
      return self._call_impl(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
      return forward_call(*args, **kwargs)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/accelerate/hooks.py", line 192, in new_forward
      output = module._old_forward(*args, **kwargs)
               ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/models/nemotron_h/modeling_nemotron_h.py", line 688, in forward
      return self.torch_forward(hidden_states, cache_params, attention_mask)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
    File "/workspace/lib/python3.12/site-packages/transformers/models/nemotron_h/modeling_nemotron_h.py", line 561, in torch_forward
      cache_params.ssm_states[self.layer_idx] * dA + dBx
      ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~^~~~
  RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cuda:1!

The issue seems to be that at generation, the NemotronHHybridDynamicCache's ssm caches are instantiated at one device but with accelerate's big model inference, the parameters may get assigned to different devices. The tensors in NemotronHHybridDynamicCache won't get automatically send to the execution device because it is not a simple tensor or container of tensor.

Sign up or log in to comment