Text Classification
Transformers
Safetensors
English
echo
research-intent
echo-dsrn
openaire-2026-hackathon
vllm
custom_code
Instructions to use ethicalabs/Echo-DSRN-v0.1.3-Research-Intent-CLF with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ethicalabs/Echo-DSRN-v0.1.3-Research-Intent-CLF with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="ethicalabs/Echo-DSRN-v0.1.3-Research-Intent-CLF", trust_remote_code=True)# Load model directly from transformers import AutoModelForSequenceClassification model = AutoModelForSequenceClassification.from_pretrained("ethicalabs/Echo-DSRN-v0.1.3-Research-Intent-CLF", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Upload folder using huggingface_hub
Browse files- __init__.py +3 -0
- configuration_echo.py +1 -0
- 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 |
-
#
|
| 1067 |
-
|
| 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 |
-
|
| 1078 |
-
logits = logits / (1.0 + alpha * gate_mean.unsqueeze(-1))
|
| 1079 |
|
| 1080 |
-
|
| 1081 |
-
if labels is not None:
|
| 1082 |
-
|
| 1083 |
-
|
| 1084 |
shift_labels = labels[..., 1:].contiguous()
|
| 1085 |
-
|
| 1086 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 1292 |
-
if
|
| 1293 |
-
#
|
| 1294 |
-
|
| 1295 |
-
|
| 1296 |
-
|
| 1297 |
-
|
| 1298 |
-
|
| 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 |
-
|
| 1307 |
-
|
| 1308 |
-
|
| 1309 |
-
|
| 1310 |
-
|
| 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
|