Jiqing commited on
Commit
e4cf79f
·
verified ·
1 Parent(s): aaa02f3

Resolve the execution device instead of hardcoding CUDA

Browse files

`infer()` hardcoded `.cuda()` and `torch.autocast("cuda", ...)`, so the model could only run on NVIDIA GPUs even though nothing in it is CUDA-specific. The device is now read back from the model own parameters, which keeps the existing CUDA behaviour identical while letting the same code run on CPU and on other accelerator backends.

`ATTENTION_CLASSES` only offered `eager` and `flash_attention_2`. FlashAttention is CUDA-only, which left `eager` as the single portable option, so the `sdpa` variants are registered as well and `_supports_sdpa` is set. MLA has no dedicated SDPA kernel here, so `mla_sdpa` maps to the math implementation.

Verified on an Intel Arc Pro B60 (`attn_implementation="sdpa"`) and on CPU (`attn_implementation="eager"`): both produce identical output.

Files changed (2) hide show
  1. modeling_deepseekocr2.py +15 -12
  2. modeling_deepseekv2.py +8 -2
modeling_deepseekocr2.py CHANGED
@@ -492,7 +492,7 @@ class DeepseekOCR2Model(DeepseekV2Model):
492
  images_in_this_batch = torch.cat(images_in_this_batch, dim=0)
493
  # exit()
494
 
495
- inputs_embeds[idx].masked_scatter_(images_seq_mask[idx].unsqueeze(-1).cuda(), images_in_this_batch)
496
 
497
  idx += 1
498
 
@@ -693,6 +693,9 @@ class DeepseekOCR2ForCausalLM(DeepseekV2ForCausalLM):
693
  def infer(self, tokenizer, prompt='', image_file='', output_path = '', base_size=1024, image_size=640, crop_mode=True, test_compress=False, save_results=False, eval_mode=False):
694
  self.disable_torch_init()
695
 
 
 
 
696
  os.makedirs(output_path, exist_ok=True)
697
  os.makedirs(f'{output_path}/images', exist_ok=True)
698
 
@@ -903,12 +906,12 @@ class DeepseekOCR2ForCausalLM(DeepseekV2ForCausalLM):
903
 
904
  if not eval_mode:
905
  streamer = NoEOSTextStreamer(tokenizer, skip_prompt=True, skip_special_tokens=False)
906
- with torch.autocast("cuda", dtype=torch.bfloat16):
907
  with torch.no_grad():
908
  output_ids = self.generate(
909
- input_ids.unsqueeze(0).cuda(),
910
- images=[(images_crop.cuda(), images_ori.cuda())],
911
- images_seq_mask = images_seq_mask.unsqueeze(0).cuda(),
912
  images_spatial_crop = images_spatial_crop,
913
  # do_sample=False,
914
  # num_beams = 1,
@@ -921,12 +924,12 @@ class DeepseekOCR2ForCausalLM(DeepseekV2ForCausalLM):
921
  )
922
 
923
  else:
924
- with torch.autocast("cuda", dtype=torch.bfloat16):
925
  with torch.no_grad():
926
  output_ids = self.generate(
927
- input_ids.unsqueeze(0).cuda(),
928
- images=[(images_crop.cuda(), images_ori.cuda())],
929
- images_seq_mask = images_seq_mask.unsqueeze(0).cuda(),
930
  images_spatial_crop = images_spatial_crop,
931
  # do_sample=False,
932
  # num_beams = 1,
@@ -939,7 +942,7 @@ class DeepseekOCR2ForCausalLM(DeepseekV2ForCausalLM):
939
 
940
 
941
  if '<image>' in conversation[0]['content'] and eval_mode:
942
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
943
  stop_str = '<|end▁of▁sentence|>'
944
  if outputs.endswith(stop_str):
945
  outputs = outputs[:-len(stop_str)]
@@ -949,7 +952,7 @@ class DeepseekOCR2ForCausalLM(DeepseekV2ForCausalLM):
949
  return outputs
950
 
951
  if '<image>' in conversation[0]['content'] and test_compress:
952
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
953
  pure_texts_outputs_token_length = len(text_encode(tokenizer, outputs, bos=False, eos=False))
954
  print('='*50)
955
  print('image size: ', (w, h))
@@ -960,7 +963,7 @@ class DeepseekOCR2ForCausalLM(DeepseekV2ForCausalLM):
960
 
961
 
962
  if '<image>' in conversation[0]['content'] and save_results:
963
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
964
  stop_str = '<|end▁of▁sentence|>'
965
 
966
  print('='*15 + 'save results:' + '='*15)
 
492
  images_in_this_batch = torch.cat(images_in_this_batch, dim=0)
493
  # exit()
494
 
495
+ inputs_embeds[idx].masked_scatter_(images_seq_mask[idx].unsqueeze(-1).to(inputs_embeds.device), images_in_this_batch)
496
 
497
  idx += 1
498
 
 
693
  def infer(self, tokenizer, prompt='', image_file='', output_path = '', base_size=1024, image_size=640, crop_mode=True, test_compress=False, save_results=False, eval_mode=False):
694
  self.disable_torch_init()
695
 
696
+ device = next(self.parameters()).device
697
+ device_type = device.type
698
+
699
  os.makedirs(output_path, exist_ok=True)
700
  os.makedirs(f'{output_path}/images', exist_ok=True)
701
 
 
906
 
907
  if not eval_mode:
908
  streamer = NoEOSTextStreamer(tokenizer, skip_prompt=True, skip_special_tokens=False)
909
+ with torch.autocast(device_type, dtype=torch.bfloat16):
910
  with torch.no_grad():
911
  output_ids = self.generate(
912
+ input_ids.unsqueeze(0).to(device),
913
+ images=[(images_crop.to(device), images_ori.to(device))],
914
+ images_seq_mask = images_seq_mask.unsqueeze(0).to(device),
915
  images_spatial_crop = images_spatial_crop,
916
  # do_sample=False,
917
  # num_beams = 1,
 
924
  )
925
 
926
  else:
927
+ with torch.autocast(device_type, dtype=torch.bfloat16):
928
  with torch.no_grad():
929
  output_ids = self.generate(
930
+ input_ids.unsqueeze(0).to(device),
931
+ images=[(images_crop.to(device), images_ori.to(device))],
932
+ images_seq_mask = images_seq_mask.unsqueeze(0).to(device),
933
  images_spatial_crop = images_spatial_crop,
934
  # do_sample=False,
935
  # num_beams = 1,
 
942
 
943
 
944
  if '<image>' in conversation[0]['content'] and eval_mode:
945
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
946
  stop_str = '<|end▁of▁sentence|>'
947
  if outputs.endswith(stop_str):
948
  outputs = outputs[:-len(stop_str)]
 
952
  return outputs
953
 
954
  if '<image>' in conversation[0]['content'] and test_compress:
955
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
956
  pure_texts_outputs_token_length = len(text_encode(tokenizer, outputs, bos=False, eos=False))
957
  print('='*50)
958
  print('image size: ', (w, h))
 
963
 
964
 
965
  if '<image>' in conversation[0]['content'] and save_results:
966
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
967
  stop_str = '<|end▁of▁sentence|>'
968
 
969
  print('='*15 + 'save results:' + '='*15)
modeling_deepseekv2.py CHANGED
@@ -36,7 +36,8 @@ from transformers.cache_utils import Cache, DynamicCache
36
  from transformers.modeling_attn_mask_utils import _prepare_4d_causal_attention_mask
37
  from transformers.models.llama.modeling_llama import (
38
  LlamaAttention,
39
- LlamaFlashAttention2
 
40
  )
41
  from transformers.modeling_outputs import (
42
  BaseModelOutputWithPast,
@@ -1230,12 +1231,16 @@ class DeepseekV2FlashAttention2(DeepseekV2Attention):
1230
  ATTENTION_CLASSES = {
1231
  "eager": DeepseekV2Attention,
1232
  "flash_attention_2": DeepseekV2FlashAttention2,
 
1233
 
1234
  "mla_eager": DeepseekV2Attention,
1235
  "mla_flash_attention_2": DeepseekV2FlashAttention2,
 
 
1236
 
1237
  "mha_eager": LlamaAttention,
1238
- "mha_flash_attention_2": LlamaFlashAttention2
 
1239
  }
1240
 
1241
 
@@ -1361,6 +1366,7 @@ class DeepseekV2PreTrainedModel(PreTrainedModel):
1361
  _no_split_modules = ["DeepseekV2DecoderLayer"]
1362
  _skip_keys_device_placement = "past_key_values"
1363
  _supports_flash_attn_2 = True
 
1364
  _supports_cache_class = True
1365
 
1366
  def _init_weights(self, module):
 
36
  from transformers.modeling_attn_mask_utils import _prepare_4d_causal_attention_mask
37
  from transformers.models.llama.modeling_llama import (
38
  LlamaAttention,
39
+ LlamaFlashAttention2,
40
+ LlamaSdpaAttention,
41
  )
42
  from transformers.modeling_outputs import (
43
  BaseModelOutputWithPast,
 
1231
  ATTENTION_CLASSES = {
1232
  "eager": DeepseekV2Attention,
1233
  "flash_attention_2": DeepseekV2FlashAttention2,
1234
+ "sdpa": DeepseekV2Attention,
1235
 
1236
  "mla_eager": DeepseekV2Attention,
1237
  "mla_flash_attention_2": DeepseekV2FlashAttention2,
1238
+ # MLA has no dedicated SDPA path, fall back to the math implementation
1239
+ "mla_sdpa": DeepseekV2Attention,
1240
 
1241
  "mha_eager": LlamaAttention,
1242
+ "mha_flash_attention_2": LlamaFlashAttention2,
1243
+ "mha_sdpa": LlamaSdpaAttention,
1244
  }
1245
 
1246
 
 
1366
  _no_split_modules = ["DeepseekV2DecoderLayer"]
1367
  _skip_keys_device_placement = "past_key_values"
1368
  _supports_flash_attn_2 = True
1369
+ _supports_sdpa = True
1370
  _supports_cache_class = True
1371
 
1372
  def _init_weights(self, module):