k050506koch commited on
Commit
8192c59
·
verified ·
1 Parent(s): 3428622

Added model forward compatibility with transformers>=5.5.3

Browse files
Files changed (1) hide show
  1. modeling_gpt3dev.py +70 -0
modeling_gpt3dev.py CHANGED
@@ -408,3 +408,73 @@ class GPT3DevLMHeadModel(GPT2LMHeadModel):
408
  AutoConfig.register("gpt3dev", GPT3DevConfig)
409
  AutoModel.register(GPT3DevConfig, GPT3DevModel)
410
  AutoModelForCausalLM.register(GPT3DevConfig, GPT3DevLMHeadModel)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
408
  AutoConfig.register("gpt3dev", GPT3DevConfig)
409
  AutoModel.register(GPT3DevConfig, GPT3DevModel)
410
  AutoModelForCausalLM.register(GPT3DevConfig, GPT3DevLMHeadModel)
411
+
412
+ # ---- Transformers 5.x compatibility patch ----
413
+ _ORIG_GPT3DEV_BLOCK_FORWARD = GPT3DevBlock.forward
414
+ _ORIG_GPT3DEV_SPARSE_FORWARD = GPT3DevSparseAttention.forward
415
+
416
+ def _patched_gpt3dev_block_forward(
417
+ self,
418
+ hidden_states,
419
+ past_key_values=None,
420
+ attention_mask=None,
421
+ encoder_hidden_states=None,
422
+ encoder_attention_mask=None,
423
+ use_cache=False,
424
+ **kwargs,
425
+ ):
426
+ cache_position = kwargs.pop("cache_position", None)
427
+ output_attentions = kwargs.pop("output_attentions", False)
428
+ head_mask = kwargs.pop("head_mask", None)
429
+ past_key_value = kwargs.pop("past_key_value", None)
430
+ if past_key_values is None:
431
+ past_key_values = past_key_value
432
+
433
+ return _ORIG_GPT3DEV_BLOCK_FORWARD(
434
+ self,
435
+ hidden_states,
436
+ past_key_value=past_key_values,
437
+ cache_position=cache_position,
438
+ attention_mask=attention_mask,
439
+ head_mask=head_mask,
440
+ encoder_hidden_states=encoder_hidden_states,
441
+ encoder_attention_mask=encoder_attention_mask,
442
+ use_cache=use_cache,
443
+ output_attentions=output_attentions,
444
+ **kwargs,
445
+ )
446
+
447
+
448
+ def _patched_gpt3dev_sparse_forward(
449
+ self,
450
+ hidden_states,
451
+ past_key_values=None,
452
+ attention_mask=None,
453
+ encoder_hidden_states=None,
454
+ encoder_attention_mask=None,
455
+ output_attentions=False,
456
+ **kwargs,
457
+ ):
458
+ cache_position = kwargs.pop("cache_position", None)
459
+ head_mask = kwargs.pop("head_mask", None)
460
+ past_key_value = kwargs.pop("past_key_value", None)
461
+ if past_key_values is None:
462
+ past_key_values = past_key_value
463
+
464
+ return _ORIG_GPT3DEV_SPARSE_FORWARD(
465
+ self,
466
+ hidden_states,
467
+ past_key_value=past_key_values,
468
+ cache_position=cache_position,
469
+ attention_mask=attention_mask,
470
+ head_mask=head_mask,
471
+ encoder_hidden_states=encoder_hidden_states,
472
+ encoder_attention_mask=encoder_attention_mask,
473
+ output_attentions=output_attentions,
474
+ **kwargs,
475
+ )
476
+
477
+
478
+ GPT3DevBlock.forward = _patched_gpt3dev_block_forward
479
+ GPT3DevSparseAttention.forward = _patched_gpt3dev_sparse_forward
480
+ # ---- End compatibility patch ----