mrs83 commited on
Commit
1b1f42b
·
verified ·
1 Parent(s): c4f72b2

Upload folder using huggingface_hub

Browse files
Files changed (5) hide show
  1. __init__.py +28 -0
  2. chat_template.jinja +5 -1
  3. configuration_echo.py +9 -0
  4. modeling_echo.py +76 -8
  5. 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
- {{- '\n<tool_call>\n{"name": "' + tool_call['name'] + '", "arguments": ' + tool_call['arguments'] | tojson + '}\n</tool_call>' }}
 
 
 
 
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
- pass
 
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(x, state_i, output_all_states=output_all_states, **kwargs)
 
 
 
 
 
 
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
- return CausalLMOutputWithPast(
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
- return SequenceClassifierOutputWithPast(
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")