FALcon6 commited on
Commit
c3730a4
·
1 Parent(s): 756dc95

Fixed `past_key_values` evaluation in `MiniCPMModel.forward`

Browse files
Files changed (1) hide show
  1. modeling_minicpm.py +12 -9
modeling_minicpm.py CHANGED
@@ -1954,18 +1954,21 @@ class MiniCPMModel(MiniCPMPreTrainedModel):
1954
  past_key_values_length = 0
1955
 
1956
  if use_cache:
1957
- use_legacy_cache = not isinstance(past_key_values, Cache)
1958
- if use_legacy_cache:
1959
  raise ValueError(
1960
  'You must use the new past_key_values format, such as the Cache class, instead of the old tuple format.'
1961
  )
1962
-
 
 
 
 
 
 
 
1963
  # Calculate the usable length of past key values
1964
- past_key_values_length = past_key_values.get_seq_length() if isinstance(past_key_values, InfLLMv2Cache) else 0
1965
-
1966
- # Initialize InfLLMv2Cache if needed
1967
- if self.config.sparse_config is not None and torch.cuda.is_available() and past_key_values_length == 0:
1968
- past_key_values = InfLLMv2Cache(config = self.config, num_hidden_layers=self.config.num_hidden_layers)
1969
 
1970
  if position_ids is None:
1971
  device = input_ids.device if input_ids is not None else inputs_embeds.device
@@ -2047,7 +2050,7 @@ class MiniCPMModel(MiniCPMPreTrainedModel):
2047
 
2048
  next_cache = None
2049
  if use_cache:
2050
- next_cache = next_decoder_cache.to_legacy_cache() if use_legacy_cache else next_decoder_cache
2051
  if not return_dict:
2052
  return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
2053
  return BaseModelOutputWithPast(
 
1954
  past_key_values_length = 0
1955
 
1956
  if use_cache:
1957
+ # Reject old tuple-style cache, but allow None (first forward pass)
1958
+ if past_key_values is not None and not isinstance(past_key_values, Cache):
1959
  raise ValueError(
1960
  'You must use the new past_key_values format, such as the Cache class, instead of the old tuple format.'
1961
  )
1962
+
1963
+ # Initialize cache if None (first forward pass)
1964
+ if past_key_values is None:
1965
+ if self.config.sparse_config is not None and torch.cuda.is_available():
1966
+ past_key_values = InfLLMv2Cache(config=self.config, num_hidden_layers=self.config.num_hidden_layers)
1967
+ else:
1968
+ past_key_values = DynamicCache()
1969
+
1970
  # Calculate the usable length of past key values
1971
+ past_key_values_length = past_key_values.get_seq_length()
 
 
 
 
1972
 
1973
  if position_ids is None:
1974
  device = input_ids.device if input_ids is not None else inputs_embeds.device
 
2050
 
2051
  next_cache = None
2052
  if use_cache:
2053
+ next_cache = next_decoder_cache
2054
  if not return_dict:
2055
  return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
2056
  return BaseModelOutputWithPast(