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)# pip install -U transformers accelerate # 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 +28 -0
- chat_template.jinja +5 -1
- configuration_echo.py +9 -0
- modeling_echo.py +76 -8
- utils.py +51 -0
__init__.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
echo_dsrn/__init__.py
|
| 3 |
+
────────────────────────────────────────────────────────────────────────────
|
| 4 |
+
Package init: registers echo_dsrn classes with HuggingFace AutoClass so
|
| 5 |
+
that AutoConfig.from_pretrained() and AutoModelForCausalLM.from_pretrained()
|
| 6 |
+
work transparently without trust_remote_code=True.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
from transformers import (
|
| 10 |
+
AutoConfig,
|
| 11 |
+
AutoModelForCausalLM,
|
| 12 |
+
AutoModelForSequenceClassification,
|
| 13 |
+
)
|
| 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)
|
| 20 |
+
AutoModelForCausalLM.register(EchoConfig, EchoForCausalLM)
|
| 21 |
+
AutoModelForSequenceClassification.register(EchoConfig, EchoForSequenceClassification)
|
| 22 |
+
|
| 23 |
+
__all__ = [
|
| 24 |
+
"EchoConfig",
|
| 25 |
+
"EchoModel",
|
| 26 |
+
"EchoForCausalLM",
|
| 27 |
+
"EchoForSequenceClassification",
|
| 28 |
+
]
|
chat_template.jinja
CHANGED
|
@@ -28,7 +28,11 @@
|
|
| 28 |
{%- if tool_call['function'] is defined %}
|
| 29 |
{%- set tool_call = tool_call['function'] %}
|
| 30 |
{%- endif %}
|
| 31 |
-
{
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
{%- endfor %}
|
| 33 |
{%- endif %}
|
| 34 |
{{- '<|end|>\n' }}
|
|
|
|
| 28 |
{%- if tool_call['function'] is defined %}
|
| 29 |
{%- set tool_call = tool_call['function'] %}
|
| 30 |
{%- endif %}
|
| 31 |
+
{%- if tool_call['arguments'] is string %}
|
| 32 |
+
{{- '\n<tool_call>\n{"name": "' + tool_call['name'] + '", "arguments": ' + tool_call['arguments'] + '}\n</tool_call>' }}
|
| 33 |
+
{%- else %}
|
| 34 |
+
{{- '\n<tool_call>\n{"name": "' + tool_call['name'] + '", "arguments": ' + tool_call['arguments'] | tojson + '}\n</tool_call>' }}
|
| 35 |
+
{%- endif %}
|
| 36 |
{%- endfor %}
|
| 37 |
{%- endif %}
|
| 38 |
{{- '<|end|>\n' }}
|
configuration_echo.py
CHANGED
|
@@ -17,6 +17,11 @@ class EchoConfig(PretrainedConfig):
|
|
| 17 |
use_hybrid_attention=True,
|
| 18 |
use_rmsnorm=True,
|
| 19 |
mlp_bias: bool = False,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
# --- Classification fields (optional, ignored by CausalLM) ---
|
| 21 |
num_labels: int = 2,
|
| 22 |
id2label: Optional[dict] = None,
|
|
@@ -52,6 +57,10 @@ class EchoConfig(PretrainedConfig):
|
|
| 52 |
self.use_hybrid_attention = use_hybrid_attention
|
| 53 |
self.use_rmsnorm = use_rmsnorm
|
| 54 |
self.mlp_bias = mlp_bias
|
|
|
|
|
|
|
|
|
|
|
|
|
| 55 |
self.classifier_dropout = classifier_dropout
|
| 56 |
|
| 57 |
# Standard HF aliases
|
|
|
|
| 17 |
use_hybrid_attention=True,
|
| 18 |
use_rmsnorm=True,
|
| 19 |
mlp_bias: bool = False,
|
| 20 |
+
pooling_mode: str = "c_T",
|
| 21 |
+
attention_masking: str = "causal",
|
| 22 |
+
# --- DSpark speculative decoding integration ---
|
| 23 |
+
output_surprise_gate_logits: bool = False,
|
| 24 |
+
surprise_temperature_alpha: float = 0.0,
|
| 25 |
# --- Classification fields (optional, ignored by CausalLM) ---
|
| 26 |
num_labels: int = 2,
|
| 27 |
id2label: Optional[dict] = None,
|
|
|
|
| 57 |
self.use_hybrid_attention = use_hybrid_attention
|
| 58 |
self.use_rmsnorm = use_rmsnorm
|
| 59 |
self.mlp_bias = mlp_bias
|
| 60 |
+
self.pooling_mode = pooling_mode
|
| 61 |
+
self.attention_masking = attention_masking
|
| 62 |
+
self.output_surprise_gate_logits = output_surprise_gate_logits
|
| 63 |
+
self.surprise_temperature_alpha = surprise_temperature_alpha
|
| 64 |
self.classifier_dropout = classifier_dropout
|
| 65 |
|
| 66 |
# Standard HF aliases
|
modeling_echo.py
CHANGED
|
@@ -13,8 +13,9 @@ from transformers.modeling_outputs import (
|
|
| 13 |
from .configuration_echo import EchoConfig
|
| 14 |
|
| 15 |
if TYPE_CHECKING:
|
| 16 |
-
# Force HF trust_remote_code AST parser to bundle triton_scan.py
|
| 17 |
-
|
|
|
|
| 18 |
|
| 19 |
try:
|
| 20 |
# pyrefly: ignore [missing-import]
|
|
@@ -280,7 +281,7 @@ def dsrn_parallel_kernel_legacy(
|
|
| 280 |
# Enabled on Legacy to fix Disconnected Slow State bug while keeping LayerNorm
|
| 281 |
x_out = x_out + model_block.linear_read(c_all)
|
| 282 |
|
| 283 |
-
return x_out, h_new, c_new, gate_stats, h_all, c_all
|
| 284 |
|
| 285 |
|
| 286 |
def dsrn_parallel_kernel_hybrid(
|
|
@@ -359,7 +360,7 @@ def dsrn_parallel_kernel_hybrid(
|
|
| 359 |
if model_block.use_hybrid_attention:
|
| 360 |
x_out = x_out + model_block.linear_read(c_all)
|
| 361 |
|
| 362 |
-
return x_out, h_new, c_new, gate_stats, h_all, c_all
|
| 363 |
|
| 364 |
|
| 365 |
def dsrn_parallel_kernel(
|
|
@@ -632,8 +633,10 @@ class DSRNBlock(nn.Module):
|
|
| 632 |
# Placeholder for Triton
|
| 633 |
pass
|
| 634 |
|
|
|
|
|
|
|
| 635 |
# Use Parallel Kernel
|
| 636 |
-
x_out, h_new, c_new, gate_stats, h_all, c_all = dsrn_parallel_kernel(
|
| 637 |
self, x, h_prev, c_prev
|
| 638 |
)
|
| 639 |
|
|
@@ -662,7 +665,11 @@ class DSRNBlock(nn.Module):
|
|
| 662 |
h_new_full = (h_new, c_new)
|
| 663 |
|
| 664 |
if kwargs.get("output_all_states", False):
|
|
|
|
|
|
|
| 665 |
return x_out, h_new_full, gate_stats, h_all, c_all
|
|
|
|
|
|
|
| 666 |
return x_out, h_new_full, gate_stats
|
| 667 |
|
| 668 |
|
|
@@ -804,6 +811,10 @@ class EchoModel(EchoPreTrainedModel):
|
|
| 804 |
all_h_all = [] if output_all_states else None
|
| 805 |
all_c_all = [] if output_all_states else None
|
| 806 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 807 |
# Layer-Major Execution
|
| 808 |
for i, block in enumerate(self.blocks):
|
| 809 |
|
|
@@ -834,10 +845,17 @@ class EchoModel(EchoPreTrainedModel):
|
|
| 834 |
state_i,
|
| 835 |
use_reentrant=False,
|
| 836 |
output_all_states=output_all_states,
|
|
|
|
| 837 |
**kwargs,
|
| 838 |
)
|
| 839 |
else:
|
| 840 |
-
out = block(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 841 |
|
| 842 |
x = out[0]
|
| 843 |
next_states.append(out[1])
|
|
@@ -846,6 +864,11 @@ class EchoModel(EchoPreTrainedModel):
|
|
| 846 |
all_gate_stats.append(out[2])
|
| 847 |
all_c_states.append(out[1][1])
|
| 848 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 849 |
if output_all_states:
|
| 850 |
all_h_all.append(out[3])
|
| 851 |
all_c_all.append(out[4])
|
|
@@ -864,6 +887,10 @@ class EchoModel(EchoPreTrainedModel):
|
|
| 864 |
if output_all_states:
|
| 865 |
return x, next_states, all_c_states, all_gate_stats, all_h_all, all_c_all
|
| 866 |
return x, next_states, all_c_states, all_gate_stats
|
|
|
|
|
|
|
|
|
|
|
|
|
| 867 |
if output_all_states:
|
| 868 |
return x, next_states, all_h_all, all_c_all
|
| 869 |
return x, next_states
|
|
@@ -878,6 +905,8 @@ class EchoModel(EchoPreTrainedModel):
|
|
| 878 |
if output_dsrn_telemetry:
|
| 879 |
output_obj.all_c_states = all_c_states
|
| 880 |
output_obj.all_gate_stats = all_gate_stats
|
|
|
|
|
|
|
| 881 |
if output_all_states:
|
| 882 |
output_obj.all_h_all = all_h_all
|
| 883 |
output_obj.all_c_all = all_c_all
|
|
@@ -990,6 +1019,12 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
|
|
| 990 |
else getattr(self.config, "use_return_dict", True)
|
| 991 |
)
|
| 992 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 993 |
'''
|
| 994 |
If kwargs is getting overloaded with extra args HF generate passes,
|
| 995 |
we safely extract kwargs here.
|
|
@@ -1013,9 +1048,12 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
|
|
| 1013 |
if hasattr(model_out, "last_hidden_state"):
|
| 1014 |
hidden_states = model_out.last_hidden_state
|
| 1015 |
new_states = model_out.past_key_values
|
|
|
|
|
|
|
| 1016 |
else:
|
| 1017 |
hidden_states = model_out[0]
|
| 1018 |
new_states = model_out[1]
|
|
|
|
| 1019 |
|
| 1020 |
# Extract telemetry if model returned raw tuple (or via custom properties)
|
| 1021 |
if hasattr(model_out, "all_c_states"):
|
|
@@ -1028,6 +1066,17 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
|
|
| 1028 |
# Project using Causal LM head
|
| 1029 |
logits = self.lm_head(hidden_states)
|
| 1030 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1031 |
loss = None
|
| 1032 |
if labels is not None:
|
| 1033 |
# Shift so that tokens < n predict n
|
|
@@ -1038,15 +1087,20 @@ class EchoForCausalLM(EchoPreTrainedModel, GenerationMixin):
|
|
| 1038 |
|
| 1039 |
if not return_dict:
|
| 1040 |
output = (logits, new_states)
|
|
|
|
|
|
|
| 1041 |
return ((loss,) + output) if loss is not None else output
|
| 1042 |
|
| 1043 |
-
|
| 1044 |
loss=loss,
|
| 1045 |
logits=logits,
|
| 1046 |
past_key_values=new_states if use_cache else None,
|
| 1047 |
hidden_states=(hidden_states,) if output_hidden_states else None,
|
| 1048 |
attentions=None,
|
| 1049 |
)
|
|
|
|
|
|
|
|
|
|
| 1050 |
|
| 1051 |
def prepare_inputs_for_generation(
|
| 1052 |
self, input_ids, past_key_values=None, attention_mask=None, **kwargs
|
|
@@ -1220,6 +1274,15 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
|
|
| 1220 |
hidden_states = model_out[0] # (B, T, D)
|
| 1221 |
new_states = model_out[1]
|
| 1222 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1223 |
# --- Pooling: last non-padding token ---
|
| 1224 |
if attention_mask is not None:
|
| 1225 |
# Find the index of the last 1 in each row of attention_mask
|
|
@@ -1267,15 +1330,20 @@ class EchoForSequenceClassification(EchoPreTrainedModel):
|
|
| 1267 |
|
| 1268 |
if not return_dict:
|
| 1269 |
output = (logits, new_states)
|
|
|
|
|
|
|
| 1270 |
return ((loss,) + output) if loss is not None else output
|
| 1271 |
|
| 1272 |
-
|
| 1273 |
loss=loss,
|
| 1274 |
logits=logits,
|
| 1275 |
past_key_values=new_states if use_cache else None,
|
| 1276 |
hidden_states=None,
|
| 1277 |
attentions=None,
|
| 1278 |
)
|
|
|
|
|
|
|
|
|
|
| 1279 |
|
| 1280 |
# ------------------------------------------------------------------
|
| 1281 |
# Convenience inference API
|
|
|
|
| 13 |
from .configuration_echo import EchoConfig
|
| 14 |
|
| 15 |
if TYPE_CHECKING:
|
| 16 |
+
# Force HF trust_remote_code AST parser to bundle triton_scan.py and utils.py
|
| 17 |
+
from .triton_scan import triton_dsrn_parallel_scan
|
| 18 |
+
from .utils import rms_norm_fn
|
| 19 |
|
| 20 |
try:
|
| 21 |
# pyrefly: ignore [missing-import]
|
|
|
|
| 281 |
# Enabled on Legacy to fix Disconnected Slow State bug while keeping LayerNorm
|
| 282 |
x_out = x_out + model_block.linear_read(c_all)
|
| 283 |
|
| 284 |
+
return x_out, h_new, c_new, gate_stats, h_all, c_all, gate_logits
|
| 285 |
|
| 286 |
|
| 287 |
def dsrn_parallel_kernel_hybrid(
|
|
|
|
| 360 |
if model_block.use_hybrid_attention:
|
| 361 |
x_out = x_out + model_block.linear_read(c_all)
|
| 362 |
|
| 363 |
+
return x_out, h_new, c_new, gate_stats, h_all, c_all, gate_logits
|
| 364 |
|
| 365 |
|
| 366 |
def dsrn_parallel_kernel(
|
|
|
|
| 633 |
# Placeholder for Triton
|
| 634 |
pass
|
| 635 |
|
| 636 |
+
output_gate_logits = kwargs.get("output_gate_logits", False)
|
| 637 |
+
|
| 638 |
# Use Parallel Kernel
|
| 639 |
+
x_out, h_new, c_new, gate_stats, h_all, c_all, gate_logits = dsrn_parallel_kernel(
|
| 640 |
self, x, h_prev, c_prev
|
| 641 |
)
|
| 642 |
|
|
|
|
| 665 |
h_new_full = (h_new, c_new)
|
| 666 |
|
| 667 |
if kwargs.get("output_all_states", False):
|
| 668 |
+
if output_gate_logits:
|
| 669 |
+
return x_out, h_new_full, gate_stats, h_all, c_all, gate_logits
|
| 670 |
return x_out, h_new_full, gate_stats, h_all, c_all
|
| 671 |
+
if output_gate_logits:
|
| 672 |
+
return x_out, h_new_full, gate_stats, gate_logits
|
| 673 |
return x_out, h_new_full, gate_stats
|
| 674 |
|
| 675 |
|
|
|
|
| 811 |
all_h_all = [] if output_all_states else None
|
| 812 |
all_c_all = [] if output_all_states else None
|
| 813 |
|
| 814 |
+
# Gate logits for DSpark speculative decoding integration
|
| 815 |
+
_output_gl = getattr(self.config, "output_surprise_gate_logits", False)
|
| 816 |
+
all_gate_logits = [] if _output_gl else None
|
| 817 |
+
|
| 818 |
# Layer-Major Execution
|
| 819 |
for i, block in enumerate(self.blocks):
|
| 820 |
|
|
|
|
| 845 |
state_i,
|
| 846 |
use_reentrant=False,
|
| 847 |
output_all_states=output_all_states,
|
| 848 |
+
output_gate_logits=_output_gl,
|
| 849 |
**kwargs,
|
| 850 |
)
|
| 851 |
else:
|
| 852 |
+
out = block(
|
| 853 |
+
x,
|
| 854 |
+
state_i,
|
| 855 |
+
output_all_states=output_all_states,
|
| 856 |
+
output_gate_logits=_output_gl,
|
| 857 |
+
**kwargs,
|
| 858 |
+
)
|
| 859 |
|
| 860 |
x = out[0]
|
| 861 |
next_states.append(out[1])
|
|
|
|
| 864 |
all_gate_stats.append(out[2])
|
| 865 |
all_c_states.append(out[1][1])
|
| 866 |
|
| 867 |
+
if _output_gl:
|
| 868 |
+
# gate_logits is at index 5 when output_all_states=False,
|
| 869 |
+
# and at index 5 when output_all_states=True (before h_all, c_all)
|
| 870 |
+
all_gate_logits.append(out[3] if not output_all_states else out[5])
|
| 871 |
+
|
| 872 |
if output_all_states:
|
| 873 |
all_h_all.append(out[3])
|
| 874 |
all_c_all.append(out[4])
|
|
|
|
| 887 |
if output_all_states:
|
| 888 |
return x, next_states, all_c_states, all_gate_stats, all_h_all, all_c_all
|
| 889 |
return x, next_states, all_c_states, all_gate_stats
|
| 890 |
+
if _output_gl:
|
| 891 |
+
if output_all_states:
|
| 892 |
+
return x, next_states, all_gate_logits, all_h_all, all_c_all
|
| 893 |
+
return x, next_states, all_gate_logits
|
| 894 |
if output_all_states:
|
| 895 |
return x, next_states, all_h_all, all_c_all
|
| 896 |
return x, next_states
|
|
|
|
| 905 |
if output_dsrn_telemetry:
|
| 906 |
output_obj.all_c_states = all_c_states
|
| 907 |
output_obj.all_gate_stats = all_gate_stats
|
| 908 |
+
if _output_gl:
|
| 909 |
+
output_obj.all_gate_logits = all_gate_logits
|
| 910 |
if output_all_states:
|
| 911 |
output_obj.all_h_all = all_h_all
|
| 912 |
output_obj.all_c_all = all_c_all
|
|
|
|
| 1019 |
else getattr(self.config, "use_return_dict", True)
|
| 1020 |
)
|
| 1021 |
|
| 1022 |
+
_output_gl = getattr(self.config, "output_surprise_gate_logits", False)
|
| 1023 |
+
alpha = getattr(self.config, "surprise_temperature_alpha", 0.0)
|
| 1024 |
+
# Force telemetry when alpha > 0 (needed for gate_stats)
|
| 1025 |
+
if alpha > 0.0:
|
| 1026 |
+
output_dsrn_telemetry = True
|
| 1027 |
+
|
| 1028 |
'''
|
| 1029 |
If kwargs is getting overloaded with extra args HF generate passes,
|
| 1030 |
we safely extract kwargs here.
|
|
|
|
| 1048 |
if hasattr(model_out, "last_hidden_state"):
|
| 1049 |
hidden_states = model_out.last_hidden_state
|
| 1050 |
new_states = model_out.past_key_values
|
| 1051 |
+
# Extract gate_logits when available (DSpark speculative decoding)
|
| 1052 |
+
gate_logits = getattr(model_out, "all_gate_logits", None)
|
| 1053 |
else:
|
| 1054 |
hidden_states = model_out[0]
|
| 1055 |
new_states = model_out[1]
|
| 1056 |
+
gate_logits = model_out[2] if len(model_out) > 2 and _output_gl else None
|
| 1057 |
|
| 1058 |
# Extract telemetry if model returned raw tuple (or via custom properties)
|
| 1059 |
if hasattr(model_out, "all_c_states"):
|
|
|
|
| 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
|
|
|
|
| 1087 |
|
| 1088 |
if not return_dict:
|
| 1089 |
output = (logits, new_states)
|
| 1090 |
+
if gate_logits is not None:
|
| 1091 |
+
output = output + (gate_logits,)
|
| 1092 |
return ((loss,) + output) if loss is not None else output
|
| 1093 |
|
| 1094 |
+
out = CausalLMOutputWithPast(
|
| 1095 |
loss=loss,
|
| 1096 |
logits=logits,
|
| 1097 |
past_key_values=new_states if use_cache else None,
|
| 1098 |
hidden_states=(hidden_states,) if output_hidden_states else None,
|
| 1099 |
attentions=None,
|
| 1100 |
)
|
| 1101 |
+
if gate_logits is not None:
|
| 1102 |
+
out.all_gate_logits = gate_logits
|
| 1103 |
+
return out
|
| 1104 |
|
| 1105 |
def prepare_inputs_for_generation(
|
| 1106 |
self, input_ids, past_key_values=None, attention_mask=None, **kwargs
|
|
|
|
| 1274 |
hidden_states = model_out[0] # (B, T, D)
|
| 1275 |
new_states = model_out[1]
|
| 1276 |
|
| 1277 |
+
# Extract gate_logits when available (DSpark speculative decoding)
|
| 1278 |
+
_output_gl = getattr(self.config, "output_surprise_gate_logits", False)
|
| 1279 |
+
if hasattr(model_out, "all_gate_logits"):
|
| 1280 |
+
gate_logits = model_out.all_gate_logits
|
| 1281 |
+
elif isinstance(model_out, tuple) and len(model_out) > 2 and _output_gl:
|
| 1282 |
+
gate_logits = model_out[2]
|
| 1283 |
+
else:
|
| 1284 |
+
gate_logits = None
|
| 1285 |
+
|
| 1286 |
# --- Pooling: last non-padding token ---
|
| 1287 |
if attention_mask is not None:
|
| 1288 |
# Find the index of the last 1 in each row of attention_mask
|
|
|
|
| 1330 |
|
| 1331 |
if not return_dict:
|
| 1332 |
output = (logits, new_states)
|
| 1333 |
+
if gate_logits is not None:
|
| 1334 |
+
output = output + (gate_logits,)
|
| 1335 |
return ((loss,) + output) if loss is not None else output
|
| 1336 |
|
| 1337 |
+
out = SequenceClassifierOutputWithPast(
|
| 1338 |
loss=loss,
|
| 1339 |
logits=logits,
|
| 1340 |
past_key_values=new_states if use_cache else None,
|
| 1341 |
hidden_states=None,
|
| 1342 |
attentions=None,
|
| 1343 |
)
|
| 1344 |
+
if gate_logits is not None:
|
| 1345 |
+
out.all_gate_logits = gate_logits
|
| 1346 |
+
return out
|
| 1347 |
|
| 1348 |
# ------------------------------------------------------------------
|
| 1349 |
# Convenience inference API
|
utils.py
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
def visualize_masked_samples(dataset, tokenizer, num_samples=2):
|
| 2 |
+
"""
|
| 3 |
+
Prints tokenized samples and color-codes tokens based on whether they are masked for loss.
|
| 4 |
+
"""
|
| 5 |
+
# ANSI escape codes
|
| 6 |
+
RESET = '\033[0m'
|
| 7 |
+
MASKED = '\033[90m\033[9m' # Dark gray + strikethrough
|
| 8 |
+
LOSS = '\033[92m\033[1m' # Bright green + bold
|
| 9 |
+
|
| 10 |
+
print("\n" + "=" * 80)
|
| 11 |
+
print(f"🔍 PREVIEWING TOKENIZED DATASET MASKS ({min(num_samples, len(dataset))} samples)")
|
| 12 |
+
print(
|
| 13 |
+
f"Legend: {MASKED}Gray Strikethrough{RESET} = Masked (prompt, ignored), {LOSS}Bright Green{RESET} = Loss Calculated (completion)"
|
| 14 |
+
)
|
| 15 |
+
print("=" * 80 + "\n")
|
| 16 |
+
|
| 17 |
+
for i in range(min(num_samples, len(dataset))):
|
| 18 |
+
row = dataset[i]
|
| 19 |
+
ids = row['input_ids']
|
| 20 |
+
labels = row['labels']
|
| 21 |
+
|
| 22 |
+
out = []
|
| 23 |
+
span_ids = []
|
| 24 |
+
# Fallback if somehow empty
|
| 25 |
+
if not labels:
|
| 26 |
+
continue
|
| 27 |
+
|
| 28 |
+
span_masked = labels[0] == -100
|
| 29 |
+
|
| 30 |
+
for tid, lbl in zip(ids, labels):
|
| 31 |
+
is_masked = lbl == -100
|
| 32 |
+
if is_masked == span_masked:
|
| 33 |
+
span_ids.append(tid)
|
| 34 |
+
else:
|
| 35 |
+
text = tokenizer.decode(span_ids, skip_special_tokens=False)
|
| 36 |
+
color = MASKED if span_masked else LOSS
|
| 37 |
+
# Ensure newlines don't break the ANSI formatting
|
| 38 |
+
text = text.replace('\n', f'{RESET}\n{color}')
|
| 39 |
+
out.append(f"{color}{text}{RESET}")
|
| 40 |
+
span_ids = [tid]
|
| 41 |
+
span_masked = is_masked
|
| 42 |
+
|
| 43 |
+
if span_ids:
|
| 44 |
+
text = tokenizer.decode(span_ids, skip_special_tokens=False)
|
| 45 |
+
color = MASKED if span_masked else LOSS
|
| 46 |
+
text = text.replace('\n', f'{RESET}\n{color}')
|
| 47 |
+
out.append(f"{color}{text}{RESET}")
|
| 48 |
+
|
| 49 |
+
print(f"--- Sample {i} ({len(ids)} tokens, {labels.count(-100)} masked) ---")
|
| 50 |
+
print("".join(out))
|
| 51 |
+
print("\n" + "-" * 80 + "\n")
|