mrs83 commited on
Commit
c4fc7ae
Β·
verified Β·
1 Parent(s): 7f85abe

Upload folder using huggingface_hub

Browse files
Files changed (3) hide show
  1. __init__.py +3 -0
  2. configuration_echo.py +1 -0
  3. modeling_echo.py +198 -46
__init__.py CHANGED
@@ -14,6 +14,7 @@ from transformers import (
14
 
15
  from .configuration_echo import EchoConfig
16
  from .modeling_echo import EchoForCausalLM, EchoForSequenceClassification, EchoModel
 
17
 
18
  # Register with HuggingFace so AutoClass routing works
19
  AutoConfig.register("echo", EchoConfig)
@@ -25,4 +26,6 @@ __all__ = [
25
  "EchoModel",
26
  "EchoForCausalLM",
27
  "EchoForSequenceClassification",
 
 
28
  ]
 
14
 
15
  from .configuration_echo import EchoConfig
16
  from .modeling_echo import EchoForCausalLM, EchoForSequenceClassification, EchoModel
17
+ from .pipelines import ChatTextClassificationPipeline, pipeline
18
 
19
  # Register with HuggingFace so AutoClass routing works
20
  AutoConfig.register("echo", EchoConfig)
 
26
  "EchoModel",
27
  "EchoForCausalLM",
28
  "EchoForSequenceClassification",
29
+ "ChatTextClassificationPipeline",
30
+ "pipeline",
31
  ]
configuration_echo.py CHANGED
@@ -27,6 +27,7 @@ class EchoConfig(PretrainedConfig):
27
  id2label: Optional[dict] = None,
28
  label2id: Optional[dict] = None,
29
  classifier_dropout: float = 0.0,
 
30
  **kwargs,
31
  ):
32
  # Synchronize hidden_size / embed_dim (HF synonym pair).
 
27
  id2label: Optional[dict] = None,
28
  label2id: Optional[dict] = None,
29
  classifier_dropout: float = 0.0,
30
+ classification_use_chat_template: bool = True,
31
  **kwargs,
32
  ):
33
  # Synchronize hidden_size / embed_dim (HF synonym pair).
modeling_echo.py CHANGED
@@ -924,15 +924,22 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
924
  # tie_word_embeddings=True and get_input/output_embeddings() are both defined.
925
  _tied_weights_keys = {"lm_head.weight": "model.embedding.weight"}
926
 
 
 
 
 
 
 
927
  @property
928
  def _keys_to_ignore_on_load_missing(self):
929
- # When mlp_bias=False (the default, and the setting for all v0.1.2 checkpoints),
930
- # bias tensors are not present in the checkpoint and should not trigger warnings.
931
- # When mlp_bias=True, these keys WILL exist in the checkpoint β€” do not silence them.
932
  if not getattr(self.config, "mlp_bias", False):
933
  return [r"model\.blocks\.\d+\.mlp_(up|down)\.bias"]
934
  return []
935
 
 
 
 
 
936
  @classmethod
937
  def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
938
  model = super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
@@ -998,6 +1005,7 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
998
  output_hidden_states: Optional[bool] = None,
999
  return_dict: Optional[bool] = None,
1000
  output_dsrn_telemetry: Optional[bool] = False,
 
1001
  **kwargs,
1002
  ) -> Union[Tuple, CausalLMOutputWithPast]:
1003
 
@@ -1063,27 +1071,54 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
1063
  self._latest_c_states = model_out[2]
1064
  self._latest_gate_stats = model_out[3]
1065
 
1066
- # Project using Causal LM head
1067
- logits = self.lm_head(hidden_states)
1068
-
1069
- # ── Surprise-gate temperature modulation ─────────────────────────
1070
- # When alpha > 0, the surprise gate Ξ»_t modulates the output logits:
1071
- # logits = logits / (1 + Ξ± Β· Ξ»_t)
1072
- # High surprise flattens the distribution; low surprise leaves it alone.
1073
- alpha = getattr(self.config, "surprise_temperature_alpha", 0.0)
1074
  if alpha > 0.0:
1075
  gate_stats = getattr(model_out, "all_gate_stats", None)
1076
  if gate_stats is not None and len(gate_stats) > 0:
1077
- gate_mean = torch.stack(gate_stats).mean(dim=0) # (B, T)
1078
- logits = logits / (1.0 + alpha * gate_mean.unsqueeze(-1))
1079
 
1080
- loss = None
1081
- if labels is not None:
1082
- # Shift so that tokens < n predict n
1083
- shift_logits = logits[..., :-1, :].contiguous()
1084
  shift_labels = labels[..., 1:].contiguous()
1085
- loss_fct = nn.CrossEntropyLoss()
1086
- loss = loss_fct(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1087
 
1088
  if not return_dict:
1089
  output = (logits, new_states)
@@ -1203,12 +1238,27 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
1203
 
1204
  def __init__(self, config: EchoConfig):
1205
  super().__init__(config)
1206
- self.num_labels = getattr(config, "num_labels", 2)
 
 
 
 
1207
  self.model = EchoModel(config)
1208
 
1209
  classifier_dropout = getattr(config, "classifier_dropout", 0.0)
1210
  self.dropout = nn.Dropout(classifier_dropout) if classifier_dropout > 0.0 else nn.Identity()
1211
- self.classifier = EchoClassifier(config.embed_dim, self.num_labels, bias=True)
 
 
 
 
 
 
 
 
 
 
 
1212
 
1213
  self.post_init()
1214
 
@@ -1269,6 +1319,11 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
1269
  if alpha > 0.0:
1270
  kwargs["output_dsrn_telemetry"] = True
1271
 
 
 
 
 
 
1272
  model_out = self.model(
1273
  input_ids=input_ids,
1274
  past_key_values=past_key_values,
@@ -1288,32 +1343,41 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
1288
  else:
1289
  gate_logits = None
1290
 
1291
- # --- Pooling: last non-padding token ---
1292
- if attention_mask is not None:
1293
- # Find the index of the last 1 in each row of attention_mask
1294
- seq_lengths = attention_mask.sum(dim=1) - 1 # (B,)
1295
- seq_lengths = seq_lengths.clamp(min=0)
1296
- else:
1297
- # No mask: use the true last token
1298
- if input_ids is not None:
1299
- seq_lengths = torch.full(
1300
- (hidden_states.size(0),),
1301
- hidden_states.size(1) - 1,
1302
- dtype=torch.long,
1303
- device=hidden_states.device,
1304
  )
1305
  else:
1306
- seq_lengths = torch.full(
1307
- (hidden_states.size(0),),
1308
- hidden_states.size(1) - 1,
1309
- dtype=torch.long,
1310
- device=hidden_states.device,
1311
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1312
 
1313
- # Gather last-token hidden states: (B, D)
1314
- pooled = hidden_states[
1315
- torch.arange(hidden_states.size(0), device=hidden_states.device), seq_lengths
1316
- ]
1317
  pooled = self.dropout(pooled)
1318
  logits = self.classifier(pooled) # (B, num_labels)
1319
 
@@ -1402,11 +1466,12 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
1402
 
1403
  self.eval()
1404
 
1405
- # Format text if baked-in templates exist
1406
  sys_prompt = getattr(self.config, "system_prompt", None)
1407
  usr_template = getattr(self.config, "user_template", None)
 
1408
 
1409
- if sys_prompt and usr_template:
1410
  messages = [{"role": "system", "content": sys_prompt}]
1411
  messages.append({"role": "user", "content": usr_template.format(text=text)})
1412
  # Format using the tokenizer's chat template
@@ -1575,3 +1640,90 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
1575
  config.dtype = str(src_dtype).replace("torch.", "")
1576
 
1577
  return clf_model
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
924
  # tie_word_embeddings=True and get_input/output_embeddings() are both defined.
925
  _tied_weights_keys = {"lm_head.weight": "model.embedding.weight"}
926
 
927
+ # _keys_to_ignore_on_load_missing: some transformers releases (β‰₯5.7)
928
+ # declare this as a read-only @property, which conflicts with
929
+ # register_buffer / __setattr__. We provide a trivial setter
930
+ # to avoid AttributeError during model.__init__.
931
+ _keys = []
932
+
933
  @property
934
  def _keys_to_ignore_on_load_missing(self):
 
 
 
935
  if not getattr(self.config, "mlp_bias", False):
936
  return [r"model\.blocks\.\d+\.mlp_(up|down)\.bias"]
937
  return []
938
 
939
+ @_keys_to_ignore_on_load_missing.setter
940
+ def _keys_to_ignore_on_load_missing(self, value):
941
+ self._keys = value
942
+
943
  @classmethod
944
  def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
945
  model = super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
 
1005
  output_hidden_states: Optional[bool] = None,
1006
  return_dict: Optional[bool] = None,
1007
  output_dsrn_telemetry: Optional[bool] = False,
1008
+ skip_logits: Optional[bool] = False,
1009
  **kwargs,
1010
  ) -> Union[Tuple, CausalLMOutputWithPast]:
1011
 
 
1071
  self._latest_c_states = model_out[2]
1072
  self._latest_gate_stats = model_out[3]
1073
 
1074
+ # ── Surprise-gate temperature scale (computed once for both paths) ──
1075
+ _surprise_scale = None # (B, T) or None
 
 
 
 
 
 
1076
  if alpha > 0.0:
1077
  gate_stats = getattr(model_out, "all_gate_stats", None)
1078
  if gate_stats is not None and len(gate_stats) > 0:
1079
+ _surprise_scale = 1.0 + alpha * torch.stack(gate_stats).mean(dim=0) # (B, T)
 
1080
 
1081
+ # ── skip_logits: avoid materialising the full [B, T, V] logits ──
1082
+ if skip_logits and labels is not None:
1083
+ logits = None
1084
+ shift_hidden = hidden_states[..., :-1, :]
1085
  shift_labels = labels[..., 1:].contiguous()
1086
+ _CHUNK = 512
1087
+ _flat_hidden = shift_hidden.reshape(-1, shift_hidden.size(-1))
1088
+ _flat_labels = shift_labels.reshape(-1)
1089
+ # Shift surprise scale to align with shifted logits
1090
+ _flat_scale = None
1091
+ if _surprise_scale is not None:
1092
+ _flat_scale = _surprise_scale[..., :-1].reshape(-1) # (B*(T-1),)
1093
+ _loss_fct = nn.CrossEntropyLoss(ignore_index=-100, reduction="sum")
1094
+ _total_loss = torch.zeros((), dtype=torch.float32, device=_flat_hidden.device)
1095
+ _total_tokens = _flat_hidden.new_zeros((), dtype=torch.long)
1096
+ for _i in range(0, _flat_hidden.size(0), _CHUNK):
1097
+ _ch = _flat_hidden[_i : _i + _CHUNK]
1098
+ _cl = self.lm_head(_ch).float()
1099
+ if _flat_scale is not None:
1100
+ _cl = _cl / _flat_scale[_i : _i + _CHUNK].unsqueeze(-1)
1101
+ _ll = _flat_labels[_i : _i + _CHUNK]
1102
+ _total_loss = _total_loss + _loss_fct(_cl, _ll)
1103
+ _total_tokens = _total_tokens + (_ll != -100).sum()
1104
+ loss = _total_loss / _total_tokens.clamp(min=1)
1105
+ else:
1106
+ # Project using Causal LM head
1107
+ logits = self.lm_head(hidden_states)
1108
+
1109
+ # ── Surprise-gate temperature modulation ─────────────────────────
1110
+ if _surprise_scale is not None:
1111
+ logits = logits / _surprise_scale.unsqueeze(-1)
1112
+
1113
+ loss = None
1114
+ if labels is not None:
1115
+ # Shift so that tokens < n predict n
1116
+ shift_logits = logits[..., :-1, :].contiguous()
1117
+ shift_labels = labels[..., 1:].contiguous()
1118
+ loss_fct = nn.CrossEntropyLoss()
1119
+ loss = loss_fct(
1120
+ shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1)
1121
+ )
1122
 
1123
  if not return_dict:
1124
  output = (logits, new_states)
 
1238
 
1239
  def __init__(self, config: EchoConfig):
1240
  super().__init__(config)
1241
+ # PretrainedConfig.to_dict() strips num_labels β€” infer from id2label
1242
+ if config.id2label is not None and len(config.id2label) > 0:
1243
+ self.num_labels = len(config.id2label)
1244
+ else:
1245
+ self.num_labels = getattr(config, "num_labels", 2)
1246
  self.model = EchoModel(config)
1247
 
1248
  classifier_dropout = getattr(config, "classifier_dropout", 0.0)
1249
  self.dropout = nn.Dropout(classifier_dropout) if classifier_dropout > 0.0 else nn.Identity()
1250
+
1251
+ # Classifier head dimension depends on pooling mode
1252
+ pooling_mode = getattr(config, "pooling_mode", None)
1253
+ if pooling_mode == "mean_c_all":
1254
+ head_dim = config.hidden_size * config.num_heads # c_all: 2048
1255
+ else:
1256
+ head_dim = config.embed_dim # hidden state: 512
1257
+
1258
+ self.classifier = EchoClassifier(head_dim, self.num_labels, bias=True)
1259
+
1260
+ # Persist pooling mode from config (survives save/load roundtrip)
1261
+ self._pooling_mode = getattr(config, "pooling_mode", None)
1262
 
1263
  self.post_init()
1264
 
 
1319
  if alpha > 0.0:
1320
  kwargs["output_dsrn_telemetry"] = True
1321
 
1322
+ # mean_c_all pooling needs the recurrent slow state per token
1323
+ pooling_mode = getattr(self, "_pooling_mode", None)
1324
+ if pooling_mode in ("mean_c_all", "hybrid"):
1325
+ kwargs["output_all_states"] = True
1326
+
1327
  model_out = self.model(
1328
  input_ids=input_ids,
1329
  past_key_values=past_key_values,
 
1343
  else:
1344
  gate_logits = None
1345
 
1346
+ # --- Pooling ---
1347
+ if pooling_mode == "mean_c_all":
1348
+ # Mean pool the recurrent slow states c_all from the last layer
1349
+ c_all = model_out.all_c_all[-1] # (B, T, hidden_size * num_heads)
1350
+ if attention_mask is not None:
1351
+ mask_expanded = attention_mask.unsqueeze(-1).expand(c_all.size()).float()
1352
+ pooled = (c_all * mask_expanded).sum(dim=1) / mask_expanded.sum(dim=1).clamp(
1353
+ min=1e-9
 
 
 
 
 
1354
  )
1355
  else:
1356
+ pooled = c_all.mean(dim=1)
1357
+ else:
1358
+ # Default: last non-padding token
1359
+ if attention_mask is not None:
1360
+ seq_lengths = attention_mask.sum(dim=1) - 1 # (B,)
1361
+ seq_lengths = seq_lengths.clamp(min=0)
1362
+ else:
1363
+ if input_ids is not None:
1364
+ seq_lengths = torch.full(
1365
+ (hidden_states.size(0),),
1366
+ hidden_states.size(1) - 1,
1367
+ dtype=torch.long,
1368
+ device=hidden_states.device,
1369
+ )
1370
+ else:
1371
+ seq_lengths = torch.full(
1372
+ (hidden_states.size(0),),
1373
+ hidden_states.size(1) - 1,
1374
+ dtype=torch.long,
1375
+ device=hidden_states.device,
1376
+ )
1377
+ pooled = hidden_states[
1378
+ torch.arange(hidden_states.size(0), device=hidden_states.device), seq_lengths
1379
+ ] # (B, D)
1380
 
 
 
 
 
1381
  pooled = self.dropout(pooled)
1382
  logits = self.classifier(pooled) # (B, num_labels)
1383
 
 
1466
 
1467
  self.eval()
1468
 
1469
+ # Format text if baked-in templates exist AND not disabled
1470
  sys_prompt = getattr(self.config, "system_prompt", None)
1471
  usr_template = getattr(self.config, "user_template", None)
1472
+ use_chat = getattr(self.config, "classification_use_chat_template", True)
1473
 
1474
+ if use_chat and sys_prompt and usr_template:
1475
  messages = [{"role": "system", "content": sys_prompt}]
1476
  messages.append({"role": "user", "content": usr_template.format(text=text)})
1477
  # Format using the tokenizer's chat template
 
1640
  config.dtype = str(src_dtype).replace("torch.", "")
1641
 
1642
  return clf_model
1643
+
1644
+ # ──────────────────────────────────────────────────────────────────
1645
+ # from_embedding β€” convert an embedding model to a classifier
1646
+ # ──────────────────────────────────────────────────────────────────
1647
+ @classmethod
1648
+ def from_embedding(
1649
+ cls,
1650
+ embed_model,
1651
+ num_labels: int = 60,
1652
+ id2label: Optional[dict] = None,
1653
+ label2id: Optional[dict] = None,
1654
+ classifier_dropout: float = 0.0,
1655
+ ) -> "EchoForSequenceClassification":
1656
+ """
1657
+ Construct an :class:`EchoForSequenceClassification` from an
1658
+ :class:`~echo_embedding.modeling_embedding.EchoModelForSentenceEmbedding`
1659
+ instance (or HF path).
1660
+
1661
+ The backbone weights are copied; the pooling mode (``mean_c_all``)
1662
+ is inherited from the embedding model. The classifier head is
1663
+ **randomly initialised** β€” this factory is dataset-agnostic.
1664
+ Fine-tune on your target dataset afterward.
1665
+
1666
+ Parameters
1667
+ ----------
1668
+ embed_model:
1669
+ An ``EchoModelForSentenceEmbedding`` instance or a HuggingFace
1670
+ model path / hub ID.
1671
+ num_labels:
1672
+ Number of output classes.
1673
+ id2label:
1674
+ Optional mapping ``{int -> str}`` for label names.
1675
+ label2id:
1676
+ Optional reverse mapping ``{str -> int}``.
1677
+ classifier_dropout:
1678
+ Dropout probability before the classification head.
1679
+
1680
+ Returns
1681
+ -------
1682
+ EchoForSequenceClassification
1683
+ """
1684
+ # ── 1. Resolve the embedding model ─────────────────────────
1685
+ if isinstance(embed_model, str):
1686
+ from echo_embedding.modeling_embedding import EchoModelForSentenceEmbedding
1687
+
1688
+ embed_model = EchoModelForSentenceEmbedding.from_pretrained(
1689
+ embed_model, trust_remote_code=True
1690
+ )
1691
+
1692
+ if id2label is None:
1693
+ id2label = {i: str(i) for i in range(num_labels)}
1694
+ if label2id is None:
1695
+ label2id = {v: k for k, v in id2label.items()}
1696
+
1697
+ # ── 2. Clone config and inject classification fields ──────
1698
+ config = embed_model.config
1699
+ config.num_labels = num_labels
1700
+ config.id2label = id2label
1701
+ config.label2id = label2id
1702
+ config.classifier_dropout = classifier_dropout
1703
+ config.pooling_mode = getattr(config, "pooling_mode", "c_T")
1704
+ config.classification_use_chat_template = False # embed→CLF: no chat template
1705
+
1706
+ # Update auto_map so Hub pipeline + AutoModel load correctly
1707
+ config.auto_map = {
1708
+ "AutoConfig": "configuration_echo.EchoConfig",
1709
+ "AutoModel": "modeling_echo.EchoForSequenceClassification",
1710
+ "AutoModelForSequenceClassification": "modeling_echo.EchoForSequenceClassification",
1711
+ "DSRNScan": "triton_scan.DSRNScanTriton",
1712
+ }
1713
+
1714
+ # ── 3. Build classifier (random init) ─────────────────────
1715
+ clf_model = cls(config)
1716
+
1717
+ # ── 4. Copy backbone weights ──────────────────────────────
1718
+ clf_model.model.load_state_dict(embed_model.model.state_dict(), strict=True)
1719
+
1720
+ # ── 5. Cast to source dtype ───────────────────────────────
1721
+ src_dtype = embed_model.dtype
1722
+ if src_dtype != torch.float32:
1723
+ clf_model = clf_model.to(src_dtype)
1724
+ config.dtype = str(src_dtype).replace("torch.", "")
1725
+
1726
+ # ── 6. Store pooling mode for forward pass ─────────────────
1727
+ clf_model._pooling_mode = getattr(config, "pooling_mode", "c_T")
1728
+
1729
+ return clf_model