Update configuration_neollm.py
Browse files- configuration_neollm.py +19 -3
configuration_neollm.py
CHANGED
|
@@ -207,6 +207,12 @@ class NeoLLMConfig(PretrainedConfig):
|
|
| 207 |
stack_memory_cache_size (:obj:`int`, *optional*, defaults to ``2048``):
|
| 208 |
Cache length kept for parity with the released StackTrans source.
|
| 209 |
The default training path keeps this cache disabled.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 210 |
fan_ratio (:obj:`float`, *optional*, defaults to 0.125):
|
| 211 |
Ratio controlling the periodic-dimension size in FANformer attention.
|
| 212 |
The transformed representation has dimension
|
|
@@ -707,11 +713,12 @@ class NeoLLMConfig(PretrainedConfig):
|
|
| 707 |
use_attn_res=False,
|
| 708 |
attn_res_num_blocks=4,
|
| 709 |
# ── StackMemory / STACKTRANS (Zhang et al., NeurIPS 2025) ───────
|
| 710 |
-
|
| 711 |
-
stack_d_model=24, # H * d_s = 4 *
|
| 712 |
num_mem_heads=4, # H = 4
|
| 713 |
-
stack_slots=16, # S =
|
| 714 |
stack_memory_cache_size=2048,
|
|
|
|
| 715 |
# ── ResFormer cross-layer FAN residual (He et al., 2023) ─────────
|
| 716 |
use_fan_residual=False,
|
| 717 |
fan_ratio=0.125,
|
|
@@ -1190,6 +1197,14 @@ class NeoLLMConfig(PretrainedConfig):
|
|
| 1190 |
f"`stack_d_model` must be divisible by `num_mem_heads`, "
|
| 1191 |
f"got stack_d_model={stack_d_model}, num_mem_heads={num_mem_heads}."
|
| 1192 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1193 |
|
| 1194 |
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
|
| 1195 |
|
|
@@ -1277,6 +1292,7 @@ class NeoLLMConfig(PretrainedConfig):
|
|
| 1277 |
self.num_mem_heads = num_mem_heads
|
| 1278 |
self.stack_slots = stack_slots
|
| 1279 |
self.stack_memory_cache_size = stack_memory_cache_size
|
|
|
|
| 1280 |
rope_config_validation(self)
|
| 1281 |
|
| 1282 |
# ── FANformer periodicity ─────────────────────────────────────────
|
|
|
|
| 207 |
stack_memory_cache_size (:obj:`int`, *optional*, defaults to ``2048``):
|
| 208 |
Cache length kept for parity with the released StackTrans source.
|
| 209 |
The default training path keeps this cache disabled.
|
| 210 |
+
stack_entropy_loss_weight (:obj:`float`, *optional*, defaults to ``0.001``):
|
| 211 |
+
Weight ``lambda_St`` for the StackTrans action-entropy regularizer.
|
| 212 |
+
When StackMemory is active, training adds
|
| 213 |
+
``lambda_St * mean(H(push, pop, noop))`` across valid tokens,
|
| 214 |
+
memory heads, and decoder layers. Set to ``0.0`` to monitor the
|
| 215 |
+
entropy without applying it to the objective.
|
| 216 |
fan_ratio (:obj:`float`, *optional*, defaults to 0.125):
|
| 217 |
Ratio controlling the periodic-dimension size in FANformer attention.
|
| 218 |
The transformed representation has dimension
|
|
|
|
| 713 |
use_attn_res=False,
|
| 714 |
attn_res_num_blocks=4,
|
| 715 |
# ── StackMemory / STACKTRANS (Zhang et al., NeurIPS 2025) ───────
|
| 716 |
+
use_stack_memory=False,
|
| 717 |
+
stack_d_model=24, # H * d_s = 4 * 6
|
| 718 |
num_mem_heads=4, # H = 4
|
| 719 |
+
stack_slots=16, # S = 16
|
| 720 |
stack_memory_cache_size=2048,
|
| 721 |
+
stack_entropy_loss_weight=1e-3,
|
| 722 |
# ── ResFormer cross-layer FAN residual (He et al., 2023) ─────────
|
| 723 |
use_fan_residual=False,
|
| 724 |
fan_ratio=0.125,
|
|
|
|
| 1197 |
f"`stack_d_model` must be divisible by `num_mem_heads`, "
|
| 1198 |
f"got stack_d_model={stack_d_model}, num_mem_heads={num_mem_heads}."
|
| 1199 |
)
|
| 1200 |
+
if (
|
| 1201 |
+
not math.isfinite(float(stack_entropy_loss_weight))
|
| 1202 |
+
or float(stack_entropy_loss_weight) < 0.0
|
| 1203 |
+
):
|
| 1204 |
+
raise ValueError(
|
| 1205 |
+
"`stack_entropy_loss_weight` must be finite and >= 0, "
|
| 1206 |
+
f"got {stack_entropy_loss_weight}."
|
| 1207 |
+
)
|
| 1208 |
|
| 1209 |
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
|
| 1210 |
|
|
|
|
| 1292 |
self.num_mem_heads = num_mem_heads
|
| 1293 |
self.stack_slots = stack_slots
|
| 1294 |
self.stack_memory_cache_size = stack_memory_cache_size
|
| 1295 |
+
self.stack_entropy_loss_weight = float(stack_entropy_loss_weight)
|
| 1296 |
rope_config_validation(self)
|
| 1297 |
|
| 1298 |
# ── FANformer periodicity ─────────────────────────────────────────
|