cadazar commited on
Commit
8a947f2
·
verified ·
1 Parent(s): 7cb7d4d

modeling_han2han: the subword tables follow the model's device on load; continuous batching withheld (serve falls back to sequential generation)

Browse files
Files changed (1) hide show
  1. modeling_han2han.py +24 -11
modeling_han2han.py CHANGED
@@ -1066,18 +1066,16 @@ class Han2HanPreTrainedModel(PreTrainedModel):
1066
  )
1067
 
1068
  if safetensors_path is not None and os.path.exists(safetensors_path):
 
 
1069
  with safe_open(safetensors_path, framework="pt") as f:
1070
- # Load encoder buffers
1071
- if "encoder.jbu" in f.keys() and not hasattr(base_model.encoder, 'jbu'):
1072
- base_model.encoder.register_buffer('jbu', f.get_tensor("encoder.jbu"))
1073
- if "encoder.cbu" in f.keys() and not hasattr(base_model.encoder, 'cbu'):
1074
- base_model.encoder.register_buffer('cbu', f.get_tensor("encoder.cbu"))
1075
-
1076
- # Load decoder buffers
1077
- if "decoder.jbu" in f.keys() and not hasattr(base_model.decoder, 'jbu'):
1078
- base_model.decoder.register_buffer('jbu', f.get_tensor("decoder.jbu"))
1079
- if "decoder.cbu" in f.keys() and not hasattr(base_model.decoder, 'cbu'):
1080
- base_model.decoder.register_buffer('cbu', f.get_tensor("decoder.cbu"))
1081
 
1082
  if loading_info is not None:
1083
  return model, loading_info
@@ -1582,6 +1580,21 @@ class Han2Han(Han2HanPreTrainedModel, GenerationMixin):
1582
  self.encoder.jbu = self.decoder.jbu = torch.tensor(jamo_buckets.copy())
1583
  self.register_buffer('jbu', self.encoder.jbu)
1584
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1585
  def tie_weights(self, missing_keys=None, *args, **kwargs):
1586
  """Override HF's generic tie_weights.
1587
 
 
1066
  )
1067
 
1068
  if safetensors_path is not None and os.path.exists(safetensors_path):
1069
+ # the base class has already placed the weights (device_map, .to()), so each
1070
+ # table goes to the device of the embedding it is indexed beside
1071
  with safe_open(safetensors_path, framework="pt") as f:
1072
+ for module, module_name in ((base_model.encoder, "encoder"), (base_model.decoder, "decoder")):
1073
+ for table_name in ("jbu", "cbu"):
1074
+ key = f"{module_name}.{table_name}"
1075
+ if key in f.keys() and not hasattr(module, table_name):
1076
+ module.register_buffer(
1077
+ table_name, f.get_tensor(key).to(module.wte.weight.device)
1078
+ )
 
 
 
 
1079
 
1080
  if loading_info is not None:
1081
  return model, loading_info
 
1580
  self.encoder.jbu = self.decoder.jbu = torch.tensor(jamo_buckets.copy())
1581
  self.register_buffer('jbu', self.encoder.jbu)
1582
 
1583
+ @property
1584
+ def init_continuous_batching(self):
1585
+ """Continuous batching is not available for this model.
1586
+
1587
+ `GenerationMixin` offers it to every model, but it drives a decoder-only forward
1588
+ pass over one packed sequence with a paged KV cache. Han2Han is an encoder-decoder
1589
+ with its own attention, so the attribute is withheld: `transformers serve
1590
+ --continuous-batching` then falls back to sequential generation with a warning
1591
+ instead of failing inside `forward()`.
1592
+ """
1593
+ raise AttributeError(
1594
+ f"{type(self).__name__} does not support continuous batching (encoder-decoder model); "
1595
+ "use generate()."
1596
+ )
1597
+
1598
  def tie_weights(self, missing_keys=None, *args, **kwargs):
1599
  """Override HF's generic tie_weights.
1600