Files changed (1) hide show
  1. modeling_unlimitedocr.py +21 -18
modeling_unlimitedocr.py CHANGED
@@ -579,7 +579,8 @@ class UnlimitedOCRModel(DeepseekV2Model):
579
  images_in_this_batch = torch.cat(images_in_this_batch, dim=0)
580
  # exit()
581
 
582
- inputs_embeds[idx].masked_scatter_(images_seq_mask[idx].unsqueeze(-1).cuda(), images_in_this_batch)
 
583
 
584
  idx += 1
585
 
@@ -786,6 +787,7 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
786
 
787
  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, max_length=32768, tps_interval=0, no_repeat_ngram_size=0, ngram_window=0, temperature=0.0):
788
  self.disable_torch_init()
 
789
 
790
  os.makedirs(output_path, exist_ok=True)
791
  os.makedirs(f'{output_path}/images', exist_ok=True)
@@ -1000,9 +1002,9 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1000
  self.config.sliding_window = None
1001
  # Build logits processors for ngram
1002
  gen_kwargs = dict(
1003
- input_ids=input_ids.unsqueeze(0).cuda(),
1004
- images=[(images_crop.cuda(), images_ori.cuda())],
1005
- images_seq_mask=images_seq_mask.unsqueeze(0).cuda(),
1006
  images_spatial_crop=images_spatial_crop,
1007
  do_sample=temperature > 0,
1008
  temperature=temperature if temperature > 0 else None,
@@ -1015,7 +1017,7 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1015
  gen_kwargs['logits_processor'] = [SlidingWindowNoRepeatNgramProcessor(no_repeat_ngram_size, ngram_window)]
1016
  elif no_repeat_ngram_size > 0:
1017
  gen_kwargs['no_repeat_ngram_size'] = no_repeat_ngram_size
1018
- with torch.autocast("cuda", dtype=torch.bfloat16):
1019
  with torch.no_grad():
1020
  output_ids = self.generate(**gen_kwargs)
1021
  self.config.sliding_window = _orig_sw
@@ -1025,9 +1027,9 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1025
  self.config._ring_window = _orig_sw
1026
  self.config.sliding_window = None
1027
  gen_kwargs = dict(
1028
- input_ids=input_ids.unsqueeze(0).cuda(),
1029
- images=[(images_crop.cuda(), images_ori.cuda())],
1030
- images_seq_mask=images_seq_mask.unsqueeze(0).cuda(),
1031
  images_spatial_crop=images_spatial_crop,
1032
  do_sample=temperature > 0,
1033
  temperature=temperature if temperature > 0 else None,
@@ -1039,14 +1041,14 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1039
  gen_kwargs['logits_processor'] = [SlidingWindowNoRepeatNgramProcessor(no_repeat_ngram_size, ngram_window)]
1040
  elif no_repeat_ngram_size > 0:
1041
  gen_kwargs['no_repeat_ngram_size'] = no_repeat_ngram_size
1042
- with torch.autocast("cuda", dtype=torch.bfloat16):
1043
  with torch.no_grad():
1044
  output_ids = self.generate(**gen_kwargs)
1045
  self.config.sliding_window = _orig_sw
1046
 
1047
 
1048
  if '<image>' in conversation[0]['content'] and eval_mode:
1049
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
1050
  stop_str = '<|end▁of▁sentence|>'
1051
  if outputs.endswith(stop_str):
1052
  outputs = outputs[:-len(stop_str)]
@@ -1056,7 +1058,7 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1056
  return outputs
1057
 
1058
  if '<image>' in conversation[0]['content'] and test_compress:
1059
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
1060
  pure_texts_outputs_token_length = len(text_encode(tokenizer, outputs, bos=False, eos=False))
1061
  print('='*50)
1062
  print('image size: ', (w, h))
@@ -1067,7 +1069,7 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1067
 
1068
 
1069
  if '<image>' in conversation[0]['content'] and save_results:
1070
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
1071
  stop_str = '<|end▁of▁sentence|>'
1072
 
1073
  print('='*15 + 'save results:' + '='*15)
@@ -1150,7 +1152,8 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1150
  save_results: whether to save output to file
1151
  """
1152
  self.disable_torch_init()
1153
-
 
1154
  if image_files is None or len(image_files) == 0:
1155
  assert False, 'image_files must be a non-empty list for multi-image inference!'
1156
 
@@ -1235,12 +1238,12 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1235
  _orig_sw = getattr(self.config, 'sliding_window_size', None) or getattr(self.config, 'sliding_window', None)
1236
  self.config._ring_window = _orig_sw # Save for ring buffer to read
1237
  self.config.sliding_window = None
1238
- with torch.autocast("cuda", dtype=torch.bfloat16):
1239
  with torch.no_grad():
1240
  gen_kwargs = dict(
1241
- input_ids=input_ids.unsqueeze(0).cuda(),
1242
- images=[(dummy_crop.cuda(), images_ori.cuda())],
1243
- images_seq_mask=images_seq_mask.unsqueeze(0).cuda(),
1244
  images_spatial_crop=images_spatial_crop,
1245
  do_sample=temperature > 0,
1246
  temperature=temperature if temperature > 0 else None,
@@ -1256,7 +1259,7 @@ class UnlimitedOCRForCausalLM(DeepseekV2ForCausalLM):
1256
  output_ids = self.generate(**gen_kwargs)
1257
  self.config.sliding_window = _orig_sw # Restore
1258
 
1259
- outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).cuda().shape[1]:])
1260
  stop_str = '<|end▁of▁sentence|>'
1261
  if outputs.endswith(stop_str):
1262
  outputs = outputs[:-len(stop_str)]
 
579
  images_in_this_batch = torch.cat(images_in_this_batch, dim=0)
580
  # exit()
581
 
582
+ scatter_mask = images_seq_mask[idx].unsqueeze(-1).to(inputs_embeds.device).expand_as(inputs_embeds[idx]).contiguous()
583
+ inputs_embeds[idx].masked_scatter_(scatter_mask, images_in_this_batch)
584
 
585
  idx += 1
586
 
 
787
 
788
  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, max_length=32768, tps_interval=0, no_repeat_ngram_size=0, ngram_window=0, temperature=0.0):
789
  self.disable_torch_init()
790
+ device = self.device
791
 
792
  os.makedirs(output_path, exist_ok=True)
793
  os.makedirs(f'{output_path}/images', exist_ok=True)
 
1002
  self.config.sliding_window = None
1003
  # Build logits processors for ngram
1004
  gen_kwargs = dict(
1005
+ input_ids=input_ids.unsqueeze(0).to(device),
1006
+ images=[(images_crop.to(device), images_ori.to(device))],
1007
+ images_seq_mask=images_seq_mask.unsqueeze(0).to(device),
1008
  images_spatial_crop=images_spatial_crop,
1009
  do_sample=temperature > 0,
1010
  temperature=temperature if temperature > 0 else None,
 
1017
  gen_kwargs['logits_processor'] = [SlidingWindowNoRepeatNgramProcessor(no_repeat_ngram_size, ngram_window)]
1018
  elif no_repeat_ngram_size > 0:
1019
  gen_kwargs['no_repeat_ngram_size'] = no_repeat_ngram_size
1020
+ with torch.autocast(device.type, dtype=torch.bfloat16):
1021
  with torch.no_grad():
1022
  output_ids = self.generate(**gen_kwargs)
1023
  self.config.sliding_window = _orig_sw
 
1027
  self.config._ring_window = _orig_sw
1028
  self.config.sliding_window = None
1029
  gen_kwargs = dict(
1030
+ input_ids=input_ids.unsqueeze(0).to(device),
1031
+ images=[(images_crop.to(device), images_ori.to(device))],
1032
+ images_seq_mask=images_seq_mask.unsqueeze(0).to(device),
1033
  images_spatial_crop=images_spatial_crop,
1034
  do_sample=temperature > 0,
1035
  temperature=temperature if temperature > 0 else None,
 
1041
  gen_kwargs['logits_processor'] = [SlidingWindowNoRepeatNgramProcessor(no_repeat_ngram_size, ngram_window)]
1042
  elif no_repeat_ngram_size > 0:
1043
  gen_kwargs['no_repeat_ngram_size'] = no_repeat_ngram_size
1044
+ with torch.autocast(device.type, dtype=torch.bfloat16):
1045
  with torch.no_grad():
1046
  output_ids = self.generate(**gen_kwargs)
1047
  self.config.sliding_window = _orig_sw
1048
 
1049
 
1050
  if '<image>' in conversation[0]['content'] and eval_mode:
1051
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
1052
  stop_str = '<|end▁of▁sentence|>'
1053
  if outputs.endswith(stop_str):
1054
  outputs = outputs[:-len(stop_str)]
 
1058
  return outputs
1059
 
1060
  if '<image>' in conversation[0]['content'] and test_compress:
1061
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
1062
  pure_texts_outputs_token_length = len(text_encode(tokenizer, outputs, bos=False, eos=False))
1063
  print('='*50)
1064
  print('image size: ', (w, h))
 
1069
 
1070
 
1071
  if '<image>' in conversation[0]['content'] and save_results:
1072
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
1073
  stop_str = '<|end▁of▁sentence|>'
1074
 
1075
  print('='*15 + 'save results:' + '='*15)
 
1152
  save_results: whether to save output to file
1153
  """
1154
  self.disable_torch_init()
1155
+ device = self.device
1156
+
1157
  if image_files is None or len(image_files) == 0:
1158
  assert False, 'image_files must be a non-empty list for multi-image inference!'
1159
 
 
1238
  _orig_sw = getattr(self.config, 'sliding_window_size', None) or getattr(self.config, 'sliding_window', None)
1239
  self.config._ring_window = _orig_sw # Save for ring buffer to read
1240
  self.config.sliding_window = None
1241
+ with torch.autocast(device.type, dtype=torch.bfloat16):
1242
  with torch.no_grad():
1243
  gen_kwargs = dict(
1244
+ input_ids=input_ids.unsqueeze(0).to(device),
1245
+ images=[(dummy_crop.to(device), images_ori.to(device))],
1246
+ images_seq_mask=images_seq_mask.unsqueeze(0).to(device),
1247
  images_spatial_crop=images_spatial_crop,
1248
  do_sample=temperature > 0,
1249
  temperature=temperature if temperature > 0 else None,
 
1259
  output_ids = self.generate(**gen_kwargs)
1260
  self.config.sliding_window = _orig_sw # Restore
1261
 
1262
+ outputs = tokenizer.decode(output_ids[0, input_ids.unsqueeze(0).to(device).shape[1]:])
1263
  stop_str = '<|end▁of▁sentence|>'
1264
  if outputs.endswith(stop_str):
1265
  outputs = outputs[:-len(stop_str)]