vtava commited on
Commit
a38f163
·
verified ·
1 Parent(s): 3ea5a9b

Upload verified PDelta3-CLVR checkpoint for layers [3, 7, 11]

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. README.md +45 -0
  3. chat_template.jinja +154 -0
  4. config.json +75 -0
  5. generation_config.json +6 -0
  6. load_model.py +33 -0
  7. model.safetensors +3 -0
  8. qwen35_progress.json +57 -0
  9. qwen35_run_status.json +38 -0
  10. qwen35_verification.json +80 -0
  11. requirements.txt +4 -0
  12. run_manifest.json +16 -0
  13. scripts/train_qwen35_pdelta3_clvr_sequential.py +491 -0
  14. scripts/train_smollm2_memory_fusion_sequential.py +986 -0
  15. src/tinycenn_lm/__init__.py +144 -0
  16. src/tinycenn_lm/__pycache__/__init__.cpython-313.pyc +0 -0
  17. src/tinycenn_lm/__pycache__/cellular_attention.cpython-313.pyc +0 -0
  18. src/tinycenn_lm/__pycache__/cenn.cpython-313.pyc +0 -0
  19. src/tinycenn_lm/__pycache__/colab_live_backup.cpython-313.pyc +0 -0
  20. src/tinycenn_lm/__pycache__/direct_colab_backup.cpython-313.pyc +0 -0
  21. src/tinycenn_lm/__pycache__/hf_persistence.cpython-313.pyc +0 -0
  22. src/tinycenn_lm/__pycache__/live_console.cpython-313.pyc +0 -0
  23. src/tinycenn_lm/__pycache__/memory_attention.cpython-313.pyc +0 -0
  24. src/tinycenn_lm/__pycache__/modeling.cpython-313.pyc +0 -0
  25. src/tinycenn_lm/__pycache__/moe.cpython-313.pyc +0 -0
  26. src/tinycenn_lm/__pycache__/pdelta2_er.cpython-313.pyc +0 -0
  27. src/tinycenn_lm/__pycache__/pdelta2_features.cpython-313.pyc +0 -0
  28. src/tinycenn_lm/__pycache__/pdelta3_frontier.cpython-313.pyc +0 -0
  29. src/tinycenn_lm/__pycache__/research_layers.cpython-313.pyc +0 -0
  30. src/tinycenn_lm/__pycache__/sharded_moe.cpython-313.pyc +0 -0
  31. src/tinycenn_lm/__pycache__/smollm2_amcenn.cpython-313.pyc +0 -0
  32. src/tinycenn_lm/__pycache__/smollm2_amcenn_v2.cpython-313.pyc +0 -0
  33. src/tinycenn_lm/__pycache__/smollm2_memory_fusion.cpython-313.pyc +0 -0
  34. src/tinycenn_lm/__pycache__/story_v2.cpython-313.pyc +0 -0
  35. src/tinycenn_lm/__pycache__/student.cpython-313.pyc +0 -0
  36. src/tinycenn_lm/cellular_attention.py +572 -0
  37. src/tinycenn_lm/cenn.py +156 -0
  38. src/tinycenn_lm/colab_live_backup.py +459 -0
  39. src/tinycenn_lm/direct_colab_backup.py +276 -0
  40. src/tinycenn_lm/distill_utils.py +232 -0
  41. src/tinycenn_lm/gemma3_integrated_memory.py +319 -0
  42. src/tinycenn_lm/gemma3_memory_fusion.py +255 -0
  43. src/tinycenn_lm/hf_persistence.py +411 -0
  44. src/tinycenn_lm/integrated_memory.py +186 -0
  45. src/tinycenn_lm/live_console.py +90 -0
  46. src/tinycenn_lm/memory_attention.py +473 -0
  47. src/tinycenn_lm/modeling.py +251 -0
  48. src/tinycenn_lm/moe.py +334 -0
  49. src/tinycenn_lm/optimized_memory.py +308 -0
  50. src/tinycenn_lm/pdelta2_er.py +267 -0
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ tokenizer.json filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ library_name: transformers
3
+ pipeline_tag: text-generation
4
+ tags:
5
+ - tinycenn
6
+ - cenn
7
+ - language-modeling
8
+ - text-generation
9
+ - research
10
+ ---
11
+
12
+ # Qwen3.5-0.8B-PDelta3-CLVR-Local32
13
+
14
+ Research artifact from **TinyCeNN-LM**. Architecture: `TinyCeNN-LM experiment`.
15
+
16
+ ## Architecture
17
+
18
+ - Architecture/run type: `TinyCeNN-LM experiment`
19
+ - Base model: `not recorded`
20
+ - Dataset: `Not recorded`
21
+ - Source code: https://github.com/vtavakkoli/TinyCeNN-LM
22
+
23
+ ## Latest saved results
24
+
25
+ No structured training report was found in this upload.
26
+
27
+ The Hugging Face repository keeps timestamped run artifacts under `runs/`. This preserves training reports, configs and run metadata independently of the temporary Colab filesystem.
28
+
29
+ ## Saved experiment files
30
+
31
+ - `config.json`
32
+ - `generation_config.json`
33
+ - `tokenizer_config.json`
34
+
35
+ ## Reproducibility
36
+
37
+ Run the matching notebook from the TinyCeNN-LM repository. Colab notebooks use a Hugging Face write token from the `HF_TOKEN` Colab Secret; tokens should never be pasted into notebook source.
38
+
39
+ ## Limitations
40
+
41
+ This is a research checkpoint. Metrics saved here are the metrics produced by the corresponding training notebook/script; unless explicitly marked as held-out evaluation, they should not be treated as publication-grade benchmark results. Generation quality can differ substantially from the base model.
42
+
43
+ ## Citation
44
+
45
+ If you use this experimental checkpoint, cite the TinyCeNN-LM repository and the upstream base model.
chat_template.jinja ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {%- set image_count = namespace(value=0) %}
2
+ {%- set video_count = namespace(value=0) %}
3
+ {%- macro render_content(content, do_vision_count, is_system_content=false) %}
4
+ {%- if content is string %}
5
+ {{- content }}
6
+ {%- elif content is iterable and content is not mapping %}
7
+ {%- for item in content %}
8
+ {%- if 'image' in item or 'image_url' in item or item.type == 'image' %}
9
+ {%- if is_system_content %}
10
+ {{- raise_exception('System message cannot contain images.') }}
11
+ {%- endif %}
12
+ {%- if do_vision_count %}
13
+ {%- set image_count.value = image_count.value + 1 %}
14
+ {%- endif %}
15
+ {%- if add_vision_id %}
16
+ {{- 'Picture ' ~ image_count.value ~ ': ' }}
17
+ {%- endif %}
18
+ {{- '<|vision_start|><|image_pad|><|vision_end|>' }}
19
+ {%- elif 'video' in item or item.type == 'video' %}
20
+ {%- if is_system_content %}
21
+ {{- raise_exception('System message cannot contain videos.') }}
22
+ {%- endif %}
23
+ {%- if do_vision_count %}
24
+ {%- set video_count.value = video_count.value + 1 %}
25
+ {%- endif %}
26
+ {%- if add_vision_id %}
27
+ {{- 'Video ' ~ video_count.value ~ ': ' }}
28
+ {%- endif %}
29
+ {{- '<|vision_start|><|video_pad|><|vision_end|>' }}
30
+ {%- elif 'text' in item %}
31
+ {{- item.text }}
32
+ {%- else %}
33
+ {{- raise_exception('Unexpected item type in content.') }}
34
+ {%- endif %}
35
+ {%- endfor %}
36
+ {%- elif content is none or content is undefined %}
37
+ {{- '' }}
38
+ {%- else %}
39
+ {{- raise_exception('Unexpected content type.') }}
40
+ {%- endif %}
41
+ {%- endmacro %}
42
+ {%- if not messages %}
43
+ {{- raise_exception('No messages provided.') }}
44
+ {%- endif %}
45
+ {%- if tools and tools is iterable and tools is not mapping %}
46
+ {{- '<|im_start|>system\n' }}
47
+ {{- "# Tools\n\nYou have access to the following functions:\n\n<tools>" }}
48
+ {%- for tool in tools %}
49
+ {{- "\n" }}
50
+ {{- tool | tojson }}
51
+ {%- endfor %}
52
+ {{- "\n</tools>" }}
53
+ {{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>' }}
54
+ {%- if messages[0].role == 'system' %}
55
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
56
+ {%- if content %}
57
+ {{- '\n\n' + content }}
58
+ {%- endif %}
59
+ {%- endif %}
60
+ {{- '<|im_end|>\n' }}
61
+ {%- else %}
62
+ {%- if messages[0].role == 'system' %}
63
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
64
+ {{- '<|im_start|>system\n' + content + '<|im_end|>\n' }}
65
+ {%- endif %}
66
+ {%- endif %}
67
+ {%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
68
+ {%- for message in messages[::-1] %}
69
+ {%- set index = (messages|length - 1) - loop.index0 %}
70
+ {%- if ns.multi_step_tool and message.role == "user" %}
71
+ {%- set content = render_content(message.content, false)|trim %}
72
+ {%- if not(content.startswith('<tool_response>') and content.endswith('</tool_response>')) %}
73
+ {%- set ns.multi_step_tool = false %}
74
+ {%- set ns.last_query_index = index %}
75
+ {%- endif %}
76
+ {%- endif %}
77
+ {%- endfor %}
78
+ {%- if ns.multi_step_tool %}
79
+ {{- raise_exception('No user query found in messages.') }}
80
+ {%- endif %}
81
+ {%- for message in messages %}
82
+ {%- set content = render_content(message.content, true)|trim %}
83
+ {%- if message.role == "system" %}
84
+ {%- if not loop.first %}
85
+ {{- raise_exception('System message must be at the beginning.') }}
86
+ {%- endif %}
87
+ {%- elif message.role == "user" %}
88
+ {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
89
+ {%- elif message.role == "assistant" %}
90
+ {%- set reasoning_content = '' %}
91
+ {%- if message.reasoning_content is string %}
92
+ {%- set reasoning_content = message.reasoning_content %}
93
+ {%- else %}
94
+ {%- if '</think>' in content %}
95
+ {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
96
+ {%- set content = content.split('</think>')[-1].lstrip('\n') %}
97
+ {%- endif %}
98
+ {%- endif %}
99
+ {%- set reasoning_content = reasoning_content|trim %}
100
+ {%- if loop.index0 > ns.last_query_index %}
101
+ {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content + '\n</think>\n\n' + content }}
102
+ {%- else %}
103
+ {{- '<|im_start|>' + message.role + '\n' + content }}
104
+ {%- endif %}
105
+ {%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
106
+ {%- for tool_call in message.tool_calls %}
107
+ {%- if tool_call.function is defined %}
108
+ {%- set tool_call = tool_call.function %}
109
+ {%- endif %}
110
+ {%- if loop.first %}
111
+ {%- if content|trim %}
112
+ {{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
113
+ {%- else %}
114
+ {{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
115
+ {%- endif %}
116
+ {%- else %}
117
+ {{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
118
+ {%- endif %}
119
+ {%- if tool_call.arguments is defined %}
120
+ {%- for args_name, args_value in tool_call.arguments|items %}
121
+ {{- '<parameter=' + args_name + '>\n' }}
122
+ {%- set args_value = args_value | tojson | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %}
123
+ {{- args_value }}
124
+ {{- '\n</parameter>\n' }}
125
+ {%- endfor %}
126
+ {%- endif %}
127
+ {{- '</function>\n</tool_call>' }}
128
+ {%- endfor %}
129
+ {%- endif %}
130
+ {{- '<|im_end|>\n' }}
131
+ {%- elif message.role == "tool" %}
132
+ {%- if loop.previtem and loop.previtem.role != "tool" %}
133
+ {{- '<|im_start|>user' }}
134
+ {%- endif %}
135
+ {{- '\n<tool_response>\n' }}
136
+ {{- content }}
137
+ {{- '\n</tool_response>' }}
138
+ {%- if not loop.last and loop.nextitem.role != "tool" %}
139
+ {{- '<|im_end|>\n' }}
140
+ {%- elif loop.last %}
141
+ {{- '<|im_end|>\n' }}
142
+ {%- endif %}
143
+ {%- else %}
144
+ {{- raise_exception('Unexpected message role.') }}
145
+ {%- endif %}
146
+ {%- endfor %}
147
+ {%- if add_generation_prompt %}
148
+ {{- '<|im_start|>assistant\n' }}
149
+ {%- if enable_thinking is defined and enable_thinking is true %}
150
+ {{- '<think>\n' }}
151
+ {%- else %}
152
+ {{- '<think>\n\n</think>\n\n' }}
153
+ {%- endif %}
154
+ {%- endif %}
config.json ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "Qwen3_5ForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "attn_output_gate": true,
8
+ "bos_token_id": null,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 248044,
11
+ "full_attention_interval": 4,
12
+ "head_dim": 256,
13
+ "hidden_act": "silu",
14
+ "hidden_size": 1024,
15
+ "initializer_range": 0.02,
16
+ "intermediate_size": 3584,
17
+ "layer_types": [
18
+ "linear_attention",
19
+ "linear_attention",
20
+ "linear_attention",
21
+ "full_attention",
22
+ "linear_attention",
23
+ "linear_attention",
24
+ "linear_attention",
25
+ "full_attention",
26
+ "linear_attention",
27
+ "linear_attention",
28
+ "linear_attention",
29
+ "full_attention",
30
+ "linear_attention",
31
+ "linear_attention",
32
+ "linear_attention",
33
+ "full_attention",
34
+ "linear_attention",
35
+ "linear_attention",
36
+ "linear_attention",
37
+ "full_attention",
38
+ "linear_attention",
39
+ "linear_attention",
40
+ "linear_attention",
41
+ "full_attention"
42
+ ],
43
+ "linear_conv_kernel_dim": 4,
44
+ "linear_key_head_dim": 128,
45
+ "linear_num_key_heads": 16,
46
+ "linear_num_value_heads": 16,
47
+ "linear_value_head_dim": 128,
48
+ "mamba_ssm_dtype": "float32",
49
+ "max_position_embeddings": 262144,
50
+ "mlp_only_layers": [],
51
+ "model_type": "qwen3_5_text",
52
+ "mtp_num_hidden_layers": 1,
53
+ "mtp_use_dedicated_embeddings": false,
54
+ "num_attention_heads": 8,
55
+ "num_hidden_layers": 24,
56
+ "num_key_value_heads": 2,
57
+ "pad_token_id": null,
58
+ "partial_rotary_factor": 0.25,
59
+ "rms_norm_eps": 1e-06,
60
+ "rope_parameters": {
61
+ "mrope_interleaved": true,
62
+ "mrope_section": [
63
+ 11,
64
+ 11,
65
+ 10
66
+ ],
67
+ "partial_rotary_factor": 0.25,
68
+ "rope_theta": 10000000,
69
+ "rope_type": "default"
70
+ },
71
+ "tie_word_embeddings": true,
72
+ "transformers_version": "5.17.0",
73
+ "use_cache": false,
74
+ "vocab_size": 248320
75
+ }
generation_config.json ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "eos_token_id": 248044,
4
+ "transformers_version": "5.17.0",
5
+ "use_cache": true
6
+ }
load_model.py ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import sys
5
+ from pathlib import Path
6
+ import torch
7
+ from safetensors.torch import load_file
8
+ from transformers import AutoTokenizer, Qwen3_5ForCausalLM
9
+
10
+ def load_model(path=None, device="cpu", dtype=None):
11
+ root = Path(path or Path(__file__).resolve().parent)
12
+ for p in (root / "src", root, root / "scripts"):
13
+ if str(p) not in sys.path:
14
+ sys.path.insert(0, str(p))
15
+ from train_qwen35_pdelta3_clvr_sequential import QwenPDelta3CLVRConfig, replace_full_attention_layers
16
+ meta = json.loads((root / "tinycenn_qwen35.json").read_text())
17
+ accepted = [int(x) for x in meta["accepted_full_attention_layers"]]
18
+ cfg = QwenPDelta3CLVRConfig.from_dict(meta["replacement_config"])
19
+ if dtype is None:
20
+ dtype = torch.bfloat16 if device.startswith("cuda") and torch.cuda.is_bf16_supported() else (torch.float16 if device.startswith("cuda") else torch.float32)
21
+ model = Qwen3_5ForCausalLM.from_pretrained(root, dtype=dtype, local_files_only=True, attn_implementation="eager")
22
+ replace_full_attention_layers(model, cfg, accepted)
23
+ single = root / "model.safetensors"
24
+ if single.exists():
25
+ model.load_state_dict(load_file(str(single), device="cpu"), strict=False)
26
+ else:
27
+ index = json.loads((root / "model.safetensors.index.json").read_text())
28
+ for shard in sorted(set(index["weight_map"].values())):
29
+ model.load_state_dict(load_file(str(root / shard), device="cpu"), strict=False)
30
+ model.config.use_cache = False
31
+ model.to(device).eval()
32
+ tokenizer = AutoTokenizer.from_pretrained(root, local_files_only=True, use_fast=True)
33
+ return model, tokenizer
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fd2298e1a9cadb70609cd7fbe6a91cf99776c61406520e47556c643ac818c229
3
+ size 1513872992
qwen35_progress.json ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "format_version": 1,
3
+ "architecture": "qwen3.5-pdelta3-gdn2-clvr-localw",
4
+ "accepted_full_attention_layers": [
5
+ 3,
6
+ 7,
7
+ 11
8
+ ],
9
+ "config": {
10
+ "feature_dim": 96,
11
+ "local_window": 32,
12
+ "chunk_size": 32,
13
+ "conv_kernel": 4,
14
+ "state_dtype": "fp16",
15
+ "variant": "conv4_gdn2_clvr_f96",
16
+ "local_gate_init": 0.72,
17
+ "warm_start_previous_core": true
18
+ },
19
+ "reports": [
20
+ {
21
+ "layer": 3,
22
+ "nmse": 0.062080949544906616,
23
+ "cosine": 0.9744633436203003,
24
+ "probe_nll": 2.8579028447469077,
25
+ "incremental_delta_nll": -0.003246148427327178,
26
+ "cumulative_delta_nll": -0.003246148427327178,
27
+ "local_gate_mean": 0.7206810712814331,
28
+ "step": 75,
29
+ "round": 1,
30
+ "accepted": true
31
+ },
32
+ {
33
+ "layer": 7,
34
+ "nmse": 0.03889689967036247,
35
+ "cosine": 0.9704566597938538,
36
+ "probe_nll": 2.8711770375569663,
37
+ "incremental_delta_nll": 0.013274192810058594,
38
+ "cumulative_delta_nll": 0.010028044382731416,
39
+ "local_gate_mean": 0.7213350534439087,
40
+ "step": 150,
41
+ "round": 1,
42
+ "accepted": true
43
+ },
44
+ {
45
+ "layer": 11,
46
+ "nmse": 0.10417895764112473,
47
+ "cosine": 0.9411033391952515,
48
+ "probe_nll": 2.8818757136662803,
49
+ "incremental_delta_nll": 0.010698676109313965,
50
+ "cumulative_delta_nll": 0.02072672049204538,
51
+ "local_gate_mean": 0.7214324474334717,
52
+ "step": 75,
53
+ "round": 1,
54
+ "accepted": true
55
+ }
56
+ ]
57
+ }
qwen35_run_status.json ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "status": "target_full_attention_prefix_accepted",
3
+ "architecture": "Qwen3.5-PDelta3-GDN2-CLVR+LocalW",
4
+ "base_model": "Qwen/Qwen3.5-0.8B",
5
+ "native_full_attention_layers": [
6
+ 3,
7
+ 7,
8
+ 11,
9
+ 15,
10
+ 19,
11
+ 23
12
+ ],
13
+ "target_full_attention_layers": [
14
+ 3,
15
+ 7,
16
+ 11
17
+ ],
18
+ "accepted_full_attention_layers": [
19
+ 3,
20
+ 7,
21
+ 11
22
+ ],
23
+ "teacher_probe_nll": 2.861148993174235,
24
+ "final_probe_nll": 2.8818757136662803,
25
+ "final_delta_nll": 0.02072672049204538,
26
+ "config": {
27
+ "feature_dim": 96,
28
+ "local_window": 32,
29
+ "chunk_size": 32,
30
+ "conv_kernel": 4,
31
+ "state_dtype": "fp16",
32
+ "variant": "conv4_gdn2_clvr_f96",
33
+ "local_gate_init": 0.72,
34
+ "warm_start_previous_core": true
35
+ },
36
+ "elapsed_minutes": 2.422297354216666,
37
+ "peak_vram_gib": 4.466190814971924
38
+ }
qwen35_verification.json ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "verified": true,
3
+ "base_model": "Qwen/Qwen3.5-0.8B",
4
+ "accepted_full_attention_layers": [
5
+ 3,
6
+ 7,
7
+ 11
8
+ ],
9
+ "config": {
10
+ "feature_dim": 96,
11
+ "local_window": 32,
12
+ "chunk_size": 32,
13
+ "conv_kernel": 4,
14
+ "state_dtype": "fp16",
15
+ "variant": "conv4_gdn2_clvr_f96",
16
+ "local_gate_init": 0.72,
17
+ "warm_start_previous_core": true
18
+ },
19
+ "probe_context": 128,
20
+ "probe_blocks": 6,
21
+ "baseline_probe_nll": 2.861148993174235,
22
+ "candidate_probe_nll": 2.8818757136662803,
23
+ "delta_nll": 0.02072672049204538,
24
+ "release_cumulative_delta_nll_limit": 0.05,
25
+ "quality_pass": true,
26
+ "all_saved_layer_gates_pass": true,
27
+ "layer_checks": [
28
+ {
29
+ "layer": 3,
30
+ "pass": true,
31
+ "nmse": 0.062080949544906616,
32
+ "cosine": 0.9744633436203003,
33
+ "incremental_delta_nll": -0.003246148427327178,
34
+ "cumulative_delta_nll": -0.003246148427327178
35
+ },
36
+ {
37
+ "layer": 7,
38
+ "pass": true,
39
+ "nmse": 0.03889689967036247,
40
+ "cosine": 0.9704566597938538,
41
+ "incremental_delta_nll": 0.013274192810058594,
42
+ "cumulative_delta_nll": 0.010028044382731416
43
+ },
44
+ {
45
+ "layer": 11,
46
+ "pass": true,
47
+ "nmse": 0.10417895764112473,
48
+ "cosine": 0.9411033391952515,
49
+ "incremental_delta_nll": 0.010698676109313965,
50
+ "cumulative_delta_nll": 0.02072672049204538
51
+ }
52
+ ],
53
+ "prompt_examples": [
54
+ {
55
+ "prompt": "The future of small language models is",
56
+ "baseline": "The future of small language models is not just about the technology, but also about the human side of the conversation.\n\nIn the last few years, small language models (SLMs) have become a game changer in the world of AI. They are fast, cheap, and capable of handling complex tasks. However, they are also not without their limitations.",
57
+ "pdelta3_clvr": "The future of small language models is not just about the technology, but also about the human side of the conversation.\n\nIn the last few years, small language models (SLMs) have been gaining traction in the tech industry. They are becoming increasingly popular for their ability to generate text, code, and other tasks. However, they are also facing challenges"
58
+ },
59
+ {
60
+ "prompt": "Artificial intelligence can help scientists by",
61
+ "baseline": "Artificial intelligence can help scientists by providing a way to predict the future.\n\n<think>\n\n</think>\n\nYes, that is a correct statement. Artificial intelligence (AI) offers scientists a powerful tool for **predictive modeling**, which allows them to analyze vast amounts of historical data to identify patterns and trends that would otherwise be difficult to spot.\n\nHere are a",
62
+ "pdelta3_clvr": "Artificial intelligence can help scientists by providing a way to predict the future of a system. This is a common example of what?\n\n<think>\n\n</think>\n\nThis is a common example of **artificial intelligence (AI)**.\n\nSpecifically, this scenario describes **machine learning** or **predictive modeling**. In this context, the AI system analyzes historical data"
63
+ },
64
+ {
65
+ "prompt": "A good software architecture should",
66
+ "baseline": "A good software architecture should be able to handle the following:\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.\nA.",
67
+ "pdelta3_clvr": "A good software architecture should be able to handle a large number of concurrent users.\nA. True\nB. False\n\n<think>\n\n</think>\n\nThe statement is **True**.\n\nIn modern software development, especially for web applications, mobile devices, and distributed systems, the ability to handle a large number of concurrent users is a fundamental requirement. This"
68
+ },
69
+ {
70
+ "prompt": "The capital of Austria is",
71
+ "baseline": "The capital of Austria is Vienna.\nThe following is a list of the most recent changes to the following:\nThe following is a list of the most recent changes to the following:\nThe following is a list of the most recent changes to the following:\nThe following is a list of the most recent changes to the following:\nThe",
72
+ "pdelta3_clvr": "The capital of Austria is Vienna.\nThe following are the most common questions and answers about the topic of \"The Great Wall of China\".\nThe following are the most common questions and answers about the topic of \"The Great Wall of China\".\nThe following are the most common questions and answers about the topic of \"The Great Wall of China"
73
+ },
74
+ {
75
+ "prompt": "Once upon a time, a small robot",
76
+ "baseline": "Once upon a time, a small robot named \"Blinky\" was living in a small room. One day, Blinky decided to play a game with his friends.\nBlinky's friends were:\n- A robot named \"Blinky\"\n- A robot named \"Blinky\"\n- A robot named \"Blinky\"\n- A robot",
77
+ "pdelta3_clvr": "Once upon a time, a small robot named \"Blinky\" was exploring the world. One day, he found a mysterious box that contained a special tool. This tool was called a \"safety net\" and it had a special feature.\n\nThe safety net had a special shape. It was like a triangle with a base of 100 units"
78
+ }
79
+ ]
80
+ }
requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ torch
2
+ transformers==5.17.0
3
+ datasets>=3,<5
4
+ safetensors
run_manifest.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "run_id": "qwen35_pdelta3_clvr_hf_e-20260915T234639Z",
3
+ "created_utc": "2026-09-15T23:46:39.791454+00:00",
4
+ "repo_id": "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32",
5
+ "notebook": null,
6
+ "python": "3.13.15",
7
+ "platform": "Linux-6.6.122+-x86_64-with-glibc2.39",
8
+ "reports": [
9
+ "config.json",
10
+ "generation_config.json",
11
+ "tokenizer_config.json"
12
+ ],
13
+ "torch": "2.11.0+cu128",
14
+ "cuda_available": true,
15
+ "gpu": "Tesla T4"
16
+ }
scripts/train_qwen35_pdelta3_clvr_sequential.py ADDED
@@ -0,0 +1,491 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Sequential PDelta3/GDN2-CLVR replacement of Qwen3.5 full-attention layers.
3
+
4
+ Qwen3.5 already uses a 3:1 hybrid stack: most layers are Gated DeltaNet
5
+ linear-attention layers and only every fourth layer is full attention. This
6
+ experiment leaves native linear-attention layers untouched and replaces only
7
+ full-attention layers, one at a time, with PDelta3/GDN2 + bounded Local-W +
8
+ cross-layer value routing.
9
+ """
10
+ from __future__ import annotations
11
+
12
+ import argparse, copy, json, math, os, random, sys, time, weakref
13
+ from dataclasses import asdict, dataclass
14
+ from pathlib import Path
15
+ from typing import Any
16
+
17
+ os.environ.setdefault("HF_HUB_DISABLE_IMPLICIT_TOKEN", "1")
18
+ os.environ.setdefault("HF_HUB_DISABLE_TELEMETRY", "1")
19
+
20
+ import torch
21
+ import torch.nn.functional as F
22
+ from datasets import load_dataset
23
+ from torch import Tensor, nn
24
+ from transformers import AutoTokenizer, Qwen3_5ForCausalLM
25
+ from transformers.models.qwen3_5.modeling_qwen3_5 import apply_rotary_pos_emb
26
+
27
+ REPO_ROOT = Path(__file__).resolve().parents[1]
28
+ SRC_ROOT = REPO_ROOT / "src"
29
+ for p in (str(SRC_ROOT), str(REPO_ROOT), str(REPO_ROOT / "scripts")):
30
+ if p not in sys.path:
31
+ sys.path.insert(0, p)
32
+
33
+ import train_smollm2_memory_fusion_sequential as seq
34
+ from tinycenn_lm.pdelta3_frontier import FrontierPDelta3Layer
35
+
36
+
37
+ @dataclass(frozen=True)
38
+ class QwenPDelta3CLVRConfig:
39
+ feature_dim: int = 96
40
+ local_window: int = 32
41
+ chunk_size: int = 32
42
+ conv_kernel: int = 4
43
+ state_dtype: str = "fp16"
44
+ variant: str = "conv4_gdn2_clvr_f96"
45
+ local_gate_init: float = 0.72
46
+ warm_start_previous_core: bool = True
47
+
48
+ def validate(self, cfg):
49
+ if self.feature_dim < 16 or self.local_window < 1 or self.conv_kernel < 1:
50
+ raise ValueError("invalid PDelta3-CLVR dimensions")
51
+ if not 1 <= self.chunk_size <= 32:
52
+ raise ValueError("chunk_size must be in [1,32]")
53
+ if self.state_dtype not in {"fp16", "fp32"}:
54
+ raise ValueError("state_dtype must be fp16 or fp32")
55
+ if int(cfg.num_attention_heads) % int(cfg.num_key_value_heads):
56
+ raise ValueError("attention heads must be divisible by KV heads")
57
+
58
+ def to_dict(self):
59
+ return asdict(self)
60
+
61
+ @classmethod
62
+ def from_dict(cls, value):
63
+ return cls(**dict(value))
64
+
65
+
66
+ def _text_config(model):
67
+ return getattr(model.config, "text_config", model.config)
68
+
69
+
70
+ def full_attention_layers(model):
71
+ kinds = list(getattr(_text_config(model), "layer_types", []))
72
+ if not kinds:
73
+ raise RuntimeError("Qwen3.5 config does not expose layer_types")
74
+ return [i for i, kind in enumerate(kinds) if kind == "full_attention"]
75
+
76
+
77
+ def _repeat_kv(x, groups):
78
+ return x.repeat_interleave(groups, dim=1)
79
+
80
+
81
+ class QwenPDelta3CLVRAttention(nn.Module):
82
+ """Qwen3.5 Gated Attention replacement: Local-W + recurrent PDelta3/GDN2."""
83
+ def __init__(self, original, cfg, config, layer_idx, previous_attention=None):
84
+ super().__init__()
85
+ config.validate(cfg)
86
+ self.config = cfg
87
+ self.layer_idx = int(layer_idx)
88
+ self.hidden_size = int(cfg.hidden_size)
89
+ self.num_heads = int(cfg.num_attention_heads)
90
+ self.num_kv_heads = int(cfg.num_key_value_heads)
91
+ self.head_dim = int(getattr(cfg, "head_dim", self.hidden_size // self.num_heads))
92
+ self.groups = self.num_heads // self.num_kv_heads
93
+ self.scaling = self.head_dim ** -0.5
94
+ self.attention_dropout = float(getattr(cfg, "attention_dropout", 0.0))
95
+ self.is_causal = True
96
+ self.local_window = int(config.local_window)
97
+ object.__setattr__(self, "_previous_attention_ref", weakref.ref(previous_attention) if previous_attention is not None else None)
98
+
99
+ # Qwen3.5 q_proj is doubled: query + post-attention output gate.
100
+ self.q_proj = copy.deepcopy(original.q_proj)
101
+ self.k_proj = copy.deepcopy(original.k_proj)
102
+ self.v_proj = copy.deepcopy(original.v_proj)
103
+ self.o_proj = copy.deepcopy(original.o_proj)
104
+ self.q_norm = copy.deepcopy(original.q_norm)
105
+ self.k_norm = copy.deepcopy(original.k_norm)
106
+
107
+ self.core = FrontierPDelta3Layer(
108
+ self.num_heads, self.num_kv_heads, self.head_dim,
109
+ feature_dim=config.feature_dim, variant=config.variant,
110
+ chunk_size=config.chunk_size, conv_kernel=config.conv_kernel,
111
+ state_dtype=config.state_dtype,
112
+ )
113
+ init = min(max(float(config.local_gate_init), 1e-4), 1 - 1e-4)
114
+ logit = math.log(init / (1 - init))
115
+ self.local_gate_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim))
116
+ self.local_gate_b = nn.Parameter(torch.full((self.num_heads,), logit))
117
+ self.last_value = None
118
+
119
+ def _previous_attention(self):
120
+ ref = object.__getattribute__(self, "_previous_attention_ref")
121
+ return None if ref is None else ref()
122
+
123
+ def _local_attention(self, q, k, v, attention_mask):
124
+ kh, vh = _repeat_kv(k, self.groups), _repeat_kv(v, self.groups)
125
+ scores = torch.matmul(q.float(), kh.float().transpose(-2, -1)) * self.scaling
126
+ t = q.shape[-2]
127
+ qi = torch.arange(t, device=q.device)[:, None]
128
+ kj = torch.arange(t, device=q.device)[None, :]
129
+ allowed = (kj <= qi) & (kj >= qi - self.local_window + 1)
130
+ bias = torch.zeros((t, t), device=q.device, dtype=scores.dtype)
131
+ bias.masked_fill_(~allowed, torch.finfo(scores.dtype).min)
132
+ scores = scores + bias[None, None]
133
+ if attention_mask is not None:
134
+ if attention_mask.ndim == 4:
135
+ scores = scores + attention_mask[..., :t, :t].float()
136
+ elif attention_mask.ndim == 2:
137
+ key_mask = 1.0 - attention_mask[:, None, None, :t].float()
138
+ scores = scores + key_mask * torch.finfo(scores.dtype).min
139
+ probs = torch.softmax(scores, dim=-1, dtype=torch.float32).to(q.dtype)
140
+ if self.training and self.attention_dropout:
141
+ probs = F.dropout(probs, p=self.attention_dropout)
142
+ return torch.matmul(probs, vh.to(probs.dtype))
143
+
144
+ def forward(self, hidden_states, position_embeddings, attention_mask=None, past_key_values=None, **kwargs):
145
+ if past_key_values is not None:
146
+ raise ValueError("research replacement requires use_cache=False")
147
+ shape = hidden_states.shape[:-1]
148
+ qg = self.q_proj(hidden_states).view(*shape, self.num_heads, self.head_dim * 2)
149
+ q, out_gate = torch.chunk(qg, 2, dim=-1)
150
+ out_gate = out_gate.reshape(*shape, self.num_heads * self.head_dim)
151
+ q = self.q_norm(q).transpose(1, 2)
152
+ k = self.k_norm(self.k_proj(hidden_states).view(*shape, self.num_kv_heads, self.head_dim)).transpose(1, 2)
153
+ v = self.v_proj(hidden_states).view(*shape, self.num_kv_heads, self.head_dim).transpose(1, 2)
154
+ cos, sin = position_embeddings
155
+ q, k = apply_rotary_pos_emb(q, k, cos, sin)
156
+
157
+ self.last_value = v.detach()
158
+ previous = self._previous_attention()
159
+ routed_v = None
160
+ if previous is not None and previous.last_value is not None and previous.last_value.shape == v.shape:
161
+ routed_v = previous.last_value.to(device=v.device, dtype=v.dtype)
162
+ if routed_v is None:
163
+ routed_v = v
164
+
165
+ global_out = self.core(q, k, v, routed_v=routed_v)
166
+ local_out = self._local_attention(q, k, v, attention_mask)
167
+ gate = torch.sigmoid(torch.einsum("bhtd,hd->bht", q.float(), self.local_gate_w.float()) + self.local_gate_b.float()[None, :, None]).to(q.dtype)
168
+ mixed = gate[..., None] * local_out + (1 - gate[..., None]) * global_out
169
+ out = mixed.transpose(1, 2).contiguous().reshape(*shape, self.num_heads * self.head_dim)
170
+ out = out * torch.sigmoid(out_gate)
171
+ return self.o_proj(out.to(hidden_states.dtype)), None
172
+
173
+ @torch.no_grad()
174
+ def local_gate_mean(self):
175
+ return float(torch.sigmoid(self.local_gate_b.float()).mean())
176
+
177
+
178
+ def replace_full_attention_layers(model, config, indices):
179
+ cfg = _text_config(model)
180
+ allowed = set(full_attention_layers(model))
181
+ previous = None
182
+ made = []
183
+ for idx in sorted(indices):
184
+ if idx not in allowed:
185
+ raise ValueError(f"layer {idx} is not full attention")
186
+ layer = model.model.layers[idx]
187
+ if isinstance(layer.self_attn, QwenPDelta3CLVRAttention):
188
+ wrapper = layer.self_attn
189
+ object.__setattr__(wrapper, "_previous_attention_ref", weakref.ref(previous) if previous is not None else None)
190
+ else:
191
+ wrapper = QwenPDelta3CLVRAttention(layer.self_attn, cfg, config, idx, previous)
192
+ layer.self_attn = wrapper
193
+ previous = wrapper
194
+ made.append(wrapper)
195
+ return made
196
+
197
+
198
+ def _prefix(idx):
199
+ return f"model.layers.{idx}.self_attn."
200
+
201
+
202
+ def selected_state(model, indices):
203
+ prefixes = tuple(_prefix(i) for i in indices)
204
+ return {k: v.detach().cpu() for k, v in model.state_dict().items() if prefixes and k.startswith(prefixes)}
205
+
206
+
207
+ def atomic_torch(payload, path):
208
+ path.parent.mkdir(parents=True, exist_ok=True)
209
+ tmp = path.with_suffix(path.suffix + ".tmp")
210
+ torch.save(payload, tmp); tmp.replace(path)
211
+
212
+
213
+ def atomic_json(payload, path):
214
+ path.parent.mkdir(parents=True, exist_ok=True)
215
+ tmp = path.with_suffix(path.suffix + ".tmp")
216
+ tmp.write_text(json.dumps(payload, indent=2), encoding="utf-8"); tmp.replace(path)
217
+
218
+
219
+ def _load_selected(model, state, layers):
220
+ inc = model.load_state_dict(state, strict=False)
221
+ prefixes = tuple(_prefix(i) for i in layers)
222
+ missing = [k for k in inc.missing_keys if prefixes and k.startswith(prefixes)]
223
+ if missing:
224
+ raise RuntimeError(f"checkpoint missing keys: {missing[:8]}")
225
+
226
+
227
+ def save_progress(out, model, config, accepted, reports):
228
+ payload = {
229
+ "format_version": 1,
230
+ "architecture": "qwen3.5-pdelta3-gdn2-clvr-localw",
231
+ "accepted_full_attention_layers": list(accepted),
232
+ "config": config.to_dict(), "reports": list(reports),
233
+ "attention_state": selected_state(model, accepted),
234
+ }
235
+ atomic_torch(payload, out / "qwen35_progress.pt")
236
+ atomic_json({k:v for k,v in payload.items() if k != "attention_state"}, out / "qwen35_progress.json")
237
+
238
+
239
+ def save_in_progress(out, model, config, accepted, current, rounds, pre_nll, reports, best):
240
+ layers = accepted + [current]
241
+ payload = {
242
+ "format_version": 1, "status": "current_layer_needs_more_training",
243
+ "accepted_full_attention_layers": list(accepted),
244
+ "current_full_attention_layer": int(current), "rounds_completed": int(rounds),
245
+ "pre_probe_nll": float(pre_nll), "config": config.to_dict(),
246
+ "reports": list(reports), "best_report": dict(best),
247
+ "attention_state": selected_state(model, layers),
248
+ }
249
+ atomic_torch(payload, out / "qwen35_in_progress.pt")
250
+ atomic_json({k:v for k,v in payload.items() if k != "attention_state"}, out / "qwen35_in_progress.json")
251
+
252
+
253
+ def load_progress(out, model, config, resume):
254
+ path = out / "qwen35_progress.pt"
255
+ if not resume or not path.exists(): return [], []
256
+ p = torch.load(path, map_location="cpu", weights_only=False)
257
+ accepted = [int(x) for x in p.get("accepted_full_attention_layers", [])]
258
+ if QwenPDelta3CLVRConfig.from_dict(p["config"]) != config:
259
+ raise RuntimeError("resume architecture differs from saved config")
260
+ if accepted:
261
+ replace_full_attention_layers(model, config, accepted)
262
+ _load_selected(model, p["attention_state"], accepted)
263
+ print(f"RESUME accepted full-attention layers={accepted}", flush=True)
264
+ return accepted, list(p.get("reports", []))
265
+
266
+
267
+ def load_in_progress(out, model, config, accepted, targets, resume):
268
+ path = out / "qwen35_in_progress.pt"
269
+ if not resume or not path.exists(): return None
270
+ p = torch.load(path, map_location="cpu", weights_only=False)
271
+ if [int(x) for x in p.get("accepted_full_attention_layers", [])] != accepted: return None
272
+ expected = targets[len(accepted)]
273
+ if int(p.get("current_full_attention_layer", -1)) != expected: return None
274
+ layers = accepted + [expected]
275
+ replace_full_attention_layers(model, config, layers)
276
+ _load_selected(model, p["attention_state"], layers)
277
+ print(f"RESUME current layer={expected}, rounds={p.get('rounds_completed',0)}", flush=True)
278
+ return p
279
+
280
+
281
+ def remove_in_progress(out):
282
+ for name in ("qwen35_in_progress.pt", "qwen35_in_progress.json"):
283
+ p = out / name
284
+ if p.exists(): p.unlink()
285
+
286
+
287
+ def warm_start(model, current, accepted):
288
+ if not accepted: return
289
+ prev = model.model.layers[accepted[-1]].self_attn
290
+ cur = model.model.layers[current].self_attn
291
+ if isinstance(prev, QwenPDelta3CLVRAttention) and isinstance(cur, QwenPDelta3CLVRAttention):
292
+ cur.core.load_state_dict(prev.core.state_dict(), strict=True)
293
+ cur.local_gate_w.data.copy_(prev.local_gate_w.data)
294
+ cur.local_gate_b.data.copy_(prev.local_gate_b.data)
295
+ print(f" warm-started layer {current} from full-attention layer {accepted[-1]}", flush=True)
296
+
297
+
298
+ def trainable_groups(model, idx, lr, qkv_scale, train_qkv):
299
+ for p in model.parameters(): p.requires_grad = False
300
+ m = model.model.layers[idx].self_attn
301
+ main = []
302
+ for p in m.core.parameters(): p.requires_grad = True; main.append(p)
303
+ for p in (m.local_gate_w, m.local_gate_b): p.requires_grad = True; main.append(p)
304
+ slow = []
305
+ if train_qkv:
306
+ for child in (m.q_proj, m.k_proj, m.v_proj, m.q_norm, m.k_norm):
307
+ for p in child.parameters(): p.requires_grad = True; slow.append(p)
308
+ groups = [{"params": main, "lr": lr}]
309
+ if slow: groups.append({"params": slow, "lr": lr * qkv_scale})
310
+ return groups, main + slow
311
+
312
+
313
+ def make_optimizer(groups, device):
314
+ try: return torch.optim.AdamW(groups, weight_decay=0.01, fused=device.type == "cuda")
315
+ except Exception: return torch.optim.AdamW(groups, weight_decay=0.01)
316
+
317
+
318
+ def capture_attention_input(model, ids, idx, amp, with_output):
319
+ cap: dict[str, Any] = {}
320
+ module = model.model.layers[idx].self_attn
321
+ def hook(mod, args, kwargs):
322
+ h = args[0] if args else kwargs.get("hidden_states")
323
+ if h is None: raise RuntimeError("hidden_states not found")
324
+ cap["hidden"] = h
325
+ for key in ("position_embeddings","position_ids","attention_mask","cache_position"):
326
+ if key in kwargs and kwargs[key] is not None:
327
+ v = kwargs[key]
328
+ if torch.is_tensor(v): v = v.detach()
329
+ elif isinstance(v, tuple): v = tuple(x.detach() if torch.is_tensor(x) else x for x in v)
330
+ cap[key] = v
331
+ handle = module.register_forward_pre_hook(hook, with_kwargs=True)
332
+ try:
333
+ if with_output:
334
+ with amp(): out = model(input_ids=ids, labels=ids, use_cache=False, return_dict=True)
335
+ else:
336
+ with torch.no_grad(), amp(): out = model(input_ids=ids, use_cache=False, return_dict=True)
337
+ finally: handle.remove()
338
+ if "hidden" not in cap: raise RuntimeError(f"failed to capture layer {idx}")
339
+ return cap, out
340
+
341
+
342
+ def attention_kwargs(cap):
343
+ return {k:cap[k] for k in ("position_embeddings","position_ids","attention_mask","cache_position") if k in cap}
344
+
345
+
346
+ def call_attention(module, hidden, kwargs):
347
+ out = module(hidden, past_key_values=None, **kwargs)
348
+ return out[0] if isinstance(out, (tuple,list)) else out
349
+
350
+
351
+ @torch.no_grad()
352
+ def function_metrics(teacher, student, ids, idx, amp):
353
+ tc, _ = capture_attention_input(teacher, ids, idx, amp, False)
354
+ sc, _ = capture_attention_input(student, ids, idx, amp, False)
355
+ hidden = sc["hidden"].detach(); kwargs = attention_kwargs(tc)
356
+ target = call_attention(teacher.model.layers[idx].self_attn, hidden, kwargs)
357
+ pred = call_attention(student.model.layers[idx].self_attn, hidden, kwargs)
358
+ nmse, cosine = seq.alignment_metrics(pred, target)
359
+ return float(nmse), float(cosine)
360
+
361
+
362
+ def distill_kl(student_logits, teacher_logits, temperature):
363
+ s, t = student_logits.float()/temperature, teacher_logits.float()/temperature
364
+ return F.kl_div(F.log_softmax(s, dim=-1), F.softmax(t, dim=-1), reduction="batchmean") * temperature**2 / max(1, student_logits.shape[1])
365
+
366
+
367
+ def passes(nmse, cosine, inc, total, args):
368
+ nll_ok = inc <= args.accept_incremental_delta_nll and total <= args.accept_cumulative_delta_nll
369
+ return nll_ok if not args.strict_acceptance else (nmse <= args.accept_nmse and cosine >= args.accept_cosine and nll_ok)
370
+
371
+
372
+ def score(r):
373
+ return (float(r["cumulative_delta_nll"]), float(r["incremental_delta_nll"]), float(r["nmse"]))
374
+
375
+
376
+ def train_round(teacher, student, idx, batch_iter, probe_blocks, teacher_nll, pre_nll, args, device, amp, round_idx):
377
+ rescue = round_idx > 1
378
+ lr = args.layer_lr * (args.rescue_lr_scale if rescue else 1.0)
379
+ fw = args.rescue_functional_weight if rescue else args.functional_weight
380
+ kw = args.rescue_kl_weight if rescue else args.kl_weight
381
+ cw = args.rescue_ce_weight if rescue else args.ce_weight
382
+ groups, trainable = trainable_groups(student, idx, lr, args.qkv_lr_scale, args.train_qkv)
383
+ opt = make_optimizer(groups, device)
384
+ scaler = torch.amp.GradScaler("cuda", enabled=(device.type == "cuda" and seq.choose_dtype(device) == torch.float16))
385
+ student.train(); module = student.model.layers[idx].self_attn
386
+ best = best_state = None
387
+ for step in range(1, args.max_layer_steps + 1):
388
+ ids = next(batch_iter).to(device, non_blocking=True)
389
+ tc, to = capture_attention_input(teacher, ids, idx, amp, True)
390
+ sc, so = capture_attention_input(student, ids, idx, amp, True)
391
+ hidden = sc["hidden"].detach(); kwargs = attention_kwargs(tc)
392
+ with torch.no_grad(), amp(): target = call_attention(teacher.model.layers[idx].self_attn, hidden, kwargs)
393
+ with amp():
394
+ pred = call_attention(module, hidden, kwargs)
395
+ functional = seq.alignment_loss(pred, target, args.cosine_weight)
396
+ kl = distill_kl(so.logits, to.logits, args.temperature)
397
+ local_penalty = torch.sigmoid(module.local_gate_b.float()).mean()
398
+ loss = fw*functional + kw*kl + cw*so.loss + args.local_gate_penalty*local_penalty
399
+ if not torch.isfinite(loss): raise RuntimeError(f"non-finite loss layer {idx} step {step}")
400
+ opt.zero_grad(set_to_none=True); scaler.scale(loss).backward(); scaler.unscale_(opt)
401
+ grad = float(torch.nn.utils.clip_grad_norm_(trainable, 1.0)); scaler.step(opt); scaler.update()
402
+ if step == 1 or step % args.log_every == 0:
403
+ bn, bc = seq.alignment_metrics(pred.detach(), target.detach())
404
+ print(f"layer={idx:02d} round={round_idx:02d} step={step:03d}/{args.max_layer_steps} loss={float(loss.detach()):.4f} func={float(functional.detach()):.4f} nmse={float(bn):.4f} cos={float(bc):.4f} kl={float(kl.detach()):.4f} ce={float(so.loss.detach()):.4f} local={module.local_gate_mean():.3f} grad={grad:.3f}", flush=True)
405
+ if step >= args.min_layer_steps and (step % args.check_every == 0 or step == args.max_layer_steps):
406
+ nmse, cosine = function_metrics(teacher, student, ids, idx, amp)
407
+ nll = seq.probe_nll(student, probe_blocks, device, amp); student.train()
408
+ inc, cum = nll - pre_nll, nll - teacher_nll
409
+ ok = passes(nmse, cosine, inc, cum, args)
410
+ cand = {"layer":idx,"nmse":nmse,"cosine":cosine,"probe_nll":nll,"incremental_delta_nll":inc,"cumulative_delta_nll":cum,"local_gate_mean":module.local_gate_mean(),"step":step,"round":round_idx,"accepted":ok}
411
+ if best is None or score(cand) < score(best): best, best_state = cand, selected_state(student, [idx])
412
+ print(f" CHECK layer={idx:02d} round={round_idx:02d} NMSE={nmse:.4f} (≤{args.accept_nmse:.4f}) cos={cosine:.4f} (≥{args.accept_cosine:.4f}) ΔNLL_inc={inc:+.5f} (≤{args.accept_incremental_delta_nll:+.5f}) ΔNLL_total={cum:+.5f} (≤{args.accept_cumulative_delta_nll:+.5f}) local={module.local_gate_mean():.3f} => {'PASS' if ok else 'continue'}", flush=True)
413
+ if ok: best, best_state = cand, selected_state(student, [idx]); break
414
+ del to, so, pred, target, loss
415
+ if best is None or best_state is None: raise RuntimeError("acceptance was never evaluated")
416
+ _load_selected(student, best_state, [idx])
417
+ return best
418
+
419
+
420
+ def parse_args():
421
+ p = argparse.ArgumentParser()
422
+ p.add_argument("--base-model", default="Qwen/Qwen3.5-0.8B"); p.add_argument("--output-dir", required=True)
423
+ p.add_argument("--dataset", default="HuggingFaceFW/fineweb-edu"); p.add_argument("--dataset-config", default="sample-10BT"); p.add_argument("--split", default="train"); p.add_argument("--text-field", default="text"); p.add_argument("--shuffle-buffer", type=int, default=2048); p.add_argument("--batch-size", type=int, default=1)
424
+ p.add_argument("--feature-dim", type=int, default=96); p.add_argument("--local-window", type=int, default=32); p.add_argument("--chunk-size", type=int, default=32); p.add_argument("--conv-kernel", type=int, default=4); p.add_argument("--state-dtype", choices=("fp16","fp32"), default="fp16"); p.add_argument("--local-gate-init", type=float, default=0.72); p.add_argument("--warm-start-previous-core", action=argparse.BooleanOptionalAction, default=True); p.add_argument("--target-full-layers", type=int, default=3)
425
+ p.add_argument("--context-length", type=int, default=128); p.add_argument("--probe-context", type=int, default=128); p.add_argument("--probe-blocks", type=int, default=6); p.add_argument("--seed", type=int, default=2026)
426
+ p.add_argument("--min-layer-steps", type=int, default=60); p.add_argument("--max-layer-steps", type=int, default=250); p.add_argument("--check-every", type=int, default=25); p.add_argument("--layer-lr", type=float, default=2e-4); p.add_argument("--qkv-lr-scale", type=float, default=0.10); p.add_argument("--train-qkv", action=argparse.BooleanOptionalAction, default=True); p.add_argument("--temperature", type=float, default=1.5)
427
+ p.add_argument("--functional-weight", type=float, default=0.30); p.add_argument("--kl-weight", type=float, default=1.0); p.add_argument("--ce-weight", type=float, default=0.08); p.add_argument("--cosine-weight", type=float, default=0.20); p.add_argument("--local-gate-penalty", type=float, default=0.001)
428
+ p.add_argument("--rescue-lr-scale", type=float, default=0.50); p.add_argument("--rescue-functional-weight", type=float, default=0.15); p.add_argument("--rescue-kl-weight", type=float, default=1.50); p.add_argument("--rescue-ce-weight", type=float, default=0.12)
429
+ p.add_argument("--accept-nmse", type=float, default=0.15); p.add_argument("--accept-cosine", type=float, default=0.94); p.add_argument("--accept-incremental-delta-nll", type=float, default=0.015); p.add_argument("--accept-cumulative-delta-nll", type=float, default=0.05); p.add_argument("--strict-acceptance", action=argparse.BooleanOptionalAction, default=True); p.add_argument("--resume", action=argparse.BooleanOptionalAction, default=True); p.add_argument("--max-runtime-minutes", type=float, default=240.0); p.add_argument("--log-every", type=int, default=10)
430
+ return p.parse_args()
431
+
432
+
433
+ def main():
434
+ args = parse_args(); random.seed(args.seed); torch.manual_seed(args.seed)
435
+ if torch.cuda.is_available(): torch.cuda.manual_seed_all(args.seed)
436
+ out = Path(args.output_dir); out.mkdir(parents=True, exist_ok=True)
437
+ max_rounds = max(1, int(os.environ.get("SEQUENTIAL_MAX_ROUNDS_PER_RUN", "2")))
438
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu"); dtype = seq.choose_dtype(device); amp = seq.amp_factory(device, dtype)
439
+ if device.type == "cuda": torch.cuda.reset_peak_memory_stats(); torch.backends.cuda.matmul.allow_tf32 = True
440
+ print(f"device={device} dtype={dtype} base={args.base_model} feature_dim={args.feature_dim} local_window={args.local_window}", flush=True)
441
+ tok = AutoTokenizer.from_pretrained(args.base_model, use_fast=True, token=False)
442
+ if tok.pad_token_id is None: tok.pad_token = tok.eos_token
443
+ load_kwargs = dict(dtype=dtype, token=False, attn_implementation="eager")
444
+ teacher = Qwen3_5ForCausalLM.from_pretrained(args.base_model, **load_kwargs).to(device); teacher.eval(); teacher.config.use_cache=False; teacher.requires_grad_(False)
445
+ student = Qwen3_5ForCausalLM.from_pretrained(args.base_model, **load_kwargs).to(device); student.config.use_cache=False
446
+ all_full = full_attention_layers(student)
447
+ if not 1 <= args.target_full_layers <= len(all_full): raise ValueError(f"target-full-layers must be in [1,{len(all_full)}]")
448
+ targets = all_full[:args.target_full_layers]
449
+ print(f"Qwen3.5 full-attention layers: {all_full}", flush=True); print(f"target prefix: {targets}", flush=True)
450
+ config = QwenPDelta3CLVRConfig(args.feature_dim,args.local_window,args.chunk_size,args.conv_kernel,args.state_dtype,"conv4_gdn2_clvr_f96",args.local_gate_init,args.warm_start_previous_core)
451
+ accepted, reports = load_progress(out, student, config, args.resume)
452
+ if accepted != targets[:len(accepted)]: raise RuntimeError(f"accepted={accepted} is not target prefix={targets}")
453
+ inprog = load_in_progress(out, student, config, accepted, targets, args.resume) if len(accepted)<len(targets) else None
454
+ raw = load_dataset(args.dataset, name=args.dataset_config, split=args.split, streaming=True).shuffle(seed=args.seed, buffer_size=args.shuffle_buffer)
455
+ batch_iter = iter(seq.batches(seq.token_blocks(raw, tok, args.text_field, args.context_length), args.batch_size))
456
+ probes = seq.build_probe_blocks(tok, context=args.probe_context, count=args.probe_blocks)
457
+ teacher_nll = seq.probe_nll(teacher, probes, device, amp); print(f"teacher probe NLL={teacher_nll:.6f}", flush=True)
458
+ start=time.perf_counter()
459
+ while len(accepted)<len(targets):
460
+ slot=len(accepted); idx=targets[slot]
461
+ print("\n"+"="*112, flush=True); print(f"QWEN3.5 REPLACEMENT slot {slot+1}/{len(targets)} — actual layer {idx}", flush=True); print("="*112, flush=True)
462
+ if inprog is not None and int(inprog["current_full_attention_layer"])==idx:
463
+ pre_nll=float(inprog["pre_probe_nll"]); rounds=int(inprog.get("rounds_completed",0)); best=dict(inprog.get("best_report",{})); print(f"continuing layer {idx} from saved best checkpoint", flush=True)
464
+ else:
465
+ pre_nll=seq.probe_nll(student, probes, device, amp); print(f"before replacement NLL={pre_nll:.6f} Δteacher={pre_nll-teacher_nll:+.6f}", flush=True)
466
+ replace_full_attention_layers(student, config, accepted+[idx]);
467
+ if config.warm_start_previous_core: warm_start(student, idx, accepted)
468
+ rounds=0; best={}
469
+ accepted_now=False
470
+ for local_round in range(1,max_rounds+1):
471
+ round_idx=rounds+local_round; print(f"\n--- full-attention layer {idx} round {round_idx} ---", flush=True)
472
+ report=train_round(teacher,student,idx,batch_iter,probes,teacher_nll,pre_nll,args,device,amp,round_idx); reports.append(report)
473
+ if not best or score(report)<score(best): best=report
474
+ if report["accepted"]:
475
+ accepted.append(idx); save_progress(out,student,config,accepted,reports); remove_in_progress(out); inprog=None; accepted_now=True
476
+ print(f"✅ ACCEPTED Qwen3.5 full-attention layer {idx}; prefix={accepted}", flush=True); break
477
+ save_in_progress(out,student,config,accepted,idx,round_idx,pre_nll,reports,best); print(f"Layer {idx} not accepted; best checkpoint saved.", flush=True)
478
+ if (time.perf_counter()-start)/60 >= args.max_runtime_minutes*0.75:
479
+ atomic_json({"status":"paused_runtime_budget","accepted_full_attention_layers":accepted,"current_full_attention_layer":idx,"best_report":best}, out/"qwen35_run_status.json"); return 0
480
+ if not accepted_now:
481
+ status={"status":"current_layer_needs_more_training","architecture":"Qwen3.5-PDelta3-GDN2-CLVR+LocalW","base_model":args.base_model,"native_full_attention_layers":all_full,"target_full_attention_layers":targets,"accepted_full_attention_layers":accepted,"current_full_attention_layer":idx,"rounds_completed":rounds+max_rounds,"best_report":best,"message":"Rerun with RESUME=True to continue from the best saved checkpoint."}
482
+ atomic_json(status,out/"qwen35_run_status.json"); print("\nNOT A CRASH:",json.dumps(status,indent=2),flush=True); return 0
483
+ final_nll=seq.probe_nll(student,probes,device,amp)
484
+ status={"status":"target_full_attention_prefix_accepted","architecture":"Qwen3.5-PDelta3-GDN2-CLVR+LocalW","base_model":args.base_model,"native_full_attention_layers":all_full,"target_full_attention_layers":targets,"accepted_full_attention_layers":accepted,"teacher_probe_nll":teacher_nll,"final_probe_nll":final_nll,"final_delta_nll":final_nll-teacher_nll,"config":config.to_dict(),"elapsed_minutes":(time.perf_counter()-start)/60,"peak_vram_gib":torch.cuda.max_memory_allocated()/(1024**3) if device.type=="cuda" else 0.0}
485
+ atomic_json(status,out/"qwen35_run_status.json"); tok.save_pretrained(out/"tokenizer"); print("\nFINAL STATUS",json.dumps(status,indent=2),flush=True); return 0
486
+
487
+
488
+ if __name__ == "__main__":
489
+ try: raise SystemExit(main())
490
+ except KeyboardInterrupt:
491
+ print("Interrupted. Persistent best checkpoints remain resumable.", file=sys.stderr, flush=True); raise
scripts/train_smollm2_memory_fusion_sequential.py ADDED
@@ -0,0 +1,986 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ from __future__ import annotations
3
+
4
+ import argparse
5
+ import json
6
+ import random
7
+ import time
8
+ from contextlib import nullcontext
9
+ from pathlib import Path
10
+ from typing import Any
11
+
12
+ import torch
13
+ import torch.nn.functional as F
14
+ from datasets import load_dataset
15
+ from transformers import AutoModelForCausalLM, AutoTokenizer
16
+
17
+ from tinycenn_lm.smollm2_memory_fusion import (
18
+ DEFAULT_SMOLLM2,
19
+ MemoryFusionLlamaAttention,
20
+ SmolMemoryFusionConfig,
21
+ parameter_summary,
22
+ replace_attention_layers,
23
+ structural_summary,
24
+ )
25
+
26
+
27
+ def parse_args() -> argparse.Namespace:
28
+ p = argparse.ArgumentParser(
29
+ description=(
30
+ "Sequential teacher-guided SmolLM2 Memory Fusion conversion with "
31
+ "per-layer acceptance gates and staged whole-model distillation."
32
+ )
33
+ )
34
+ p.add_argument("--base-model", default=DEFAULT_SMOLLM2)
35
+ p.add_argument("--output-dir", default="checkpoints/smollm2-memory-fusion-sequential-r64")
36
+ p.add_argument("--dataset", default="HuggingFaceFW/fineweb-edu")
37
+ p.add_argument("--dataset-config", default="sample-10BT")
38
+ p.add_argument("--split", default="train")
39
+ p.add_argument("--text-field", default="text")
40
+ p.add_argument("--context-length", type=int, default=128)
41
+ p.add_argument("--batch-size", type=int, default=1)
42
+ p.add_argument("--feature-dim", type=int, default=32)
43
+ p.add_argument("--memory-rank", type=int, choices=(32, 48, 64), default=64)
44
+ p.add_argument("--seed", type=int, default=73)
45
+ p.add_argument("--shuffle-buffer", type=int, default=2048)
46
+
47
+ p.add_argument("--min-layer-steps", type=int, default=50)
48
+ p.add_argument("--max-layer-steps", type=int, default=300)
49
+ p.add_argument("--check-every", type=int, default=25)
50
+ p.add_argument("--layer-lr", type=float, default=2e-4)
51
+ p.add_argument("--teacher-alpha-start", type=float, default=0.90)
52
+ p.add_argument("--teacher-alpha-end", type=float, default=0.00)
53
+ p.add_argument("--layer-kl-weight", type=float, default=0.10)
54
+ p.add_argument("--layer-ce-weight", type=float, default=0.05)
55
+ p.add_argument("--cosine-weight", type=float, default=0.25)
56
+ p.add_argument("--accept-nmse", type=float, default=0.20)
57
+ p.add_argument("--accept-cosine", type=float, default=0.90)
58
+ p.add_argument("--accept-incremental-delta-nll", type=float, default=0.015)
59
+ p.add_argument("--accept-cumulative-delta-nll", type=float, default=0.05)
60
+ p.add_argument("--probe-blocks", type=int, default=4)
61
+ p.add_argument("--probe-context", type=int, default=128)
62
+ p.add_argument("--strict-acceptance", action=argparse.BooleanOptionalAction, default=True)
63
+
64
+ p.add_argument("--core-o-tokens", type=int, default=50_000)
65
+ p.add_argument("--core-o-lr", type=float, default=3e-5)
66
+ p.add_argument("--norm-tokens", type=int, default=50_000)
67
+ p.add_argument("--norm-lr", type=float, default=8e-6)
68
+ p.add_argument("--full-tokens", type=int, default=100_000)
69
+ p.add_argument("--full-lr", type=float, default=3e-6)
70
+ p.add_argument("--grad-accum", type=int, default=4)
71
+ p.add_argument("--temperature", type=float, default=2.0)
72
+ p.add_argument("--ce-weight", type=float, default=0.25)
73
+ p.add_argument("--kl-weight", type=float, default=1.0)
74
+ p.add_argument("--hidden-weight", type=float, default=0.5)
75
+
76
+ p.add_argument("--resume", action=argparse.BooleanOptionalAction, default=True)
77
+ p.add_argument("--backup-every-updates", type=int, default=50)
78
+ p.add_argument("--max-runtime-minutes", type=float, default=240.0)
79
+ p.add_argument("--log-every", type=int, default=10)
80
+ return p.parse_args()
81
+
82
+
83
+ def set_seed(seed: int) -> None:
84
+ random.seed(seed)
85
+ torch.manual_seed(seed)
86
+ if torch.cuda.is_available():
87
+ torch.cuda.manual_seed_all(seed)
88
+
89
+
90
+ def choose_dtype(device: torch.device) -> torch.dtype:
91
+ if device.type != "cuda":
92
+ return torch.float32
93
+ return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
94
+
95
+
96
+ def amp_factory(device: torch.device, dtype: torch.dtype):
97
+ if device.type == "cuda":
98
+ return lambda: torch.autocast("cuda", dtype=dtype)
99
+ return nullcontext
100
+
101
+
102
+ def token_blocks(dataset, tokenizer, text_field: str, context_length: int):
103
+ eos = tokenizer.eos_token_id
104
+ if eos is None:
105
+ raise ValueError("tokenizer must define eos_token_id")
106
+ buffer: list[int] = []
107
+ for row in dataset:
108
+ text = str(row.get(text_field, "")).strip()
109
+ if not text:
110
+ continue
111
+ ids = tokenizer(text, add_special_tokens=False, verbose=False)["input_ids"]
112
+ if not ids:
113
+ continue
114
+ buffer.extend(ids)
115
+ buffer.append(eos)
116
+ while len(buffer) >= context_length:
117
+ yield torch.tensor(buffer[:context_length], dtype=torch.long)
118
+ del buffer[:context_length]
119
+
120
+
121
+ def batches(blocks, batch_size: int):
122
+ pending = []
123
+ for block in blocks:
124
+ pending.append(block)
125
+ if len(pending) == batch_size:
126
+ yield torch.stack(pending)
127
+ pending.clear()
128
+
129
+
130
+ def make_optimizer(params, lr: float, device: torch.device):
131
+ params = list(params)
132
+ if not params:
133
+ raise ValueError("optimizer received no trainable parameters")
134
+ try:
135
+ return torch.optim.AdamW(params, lr=lr, weight_decay=0.01, fused=(device.type == "cuda"))
136
+ except Exception:
137
+ return torch.optim.AdamW(params, lr=lr, weight_decay=0.01)
138
+
139
+
140
+ def distillation_kl(student_logits, teacher_logits, temperature: float) -> torch.Tensor:
141
+ s = student_logits.float() / temperature
142
+ t = teacher_logits.float() / temperature
143
+ per_token = F.kl_div(
144
+ F.log_softmax(s, dim=-1),
145
+ F.softmax(t, dim=-1),
146
+ reduction="none",
147
+ ).sum(dim=-1)
148
+ return per_token.mean() * (temperature ** 2)
149
+
150
+
151
+ def representation_loss(student_hidden, teacher_hidden) -> torch.Tensor:
152
+ max_idx = min(len(student_hidden), len(teacher_hidden)) - 1
153
+ candidates = (6, 12, 18, 24, 30)
154
+ indices = [i for i in candidates if i <= max_idx]
155
+ if not indices:
156
+ indices = [max_idx]
157
+ terms = []
158
+ for idx in indices:
159
+ s = student_hidden[idx].float()
160
+ t = teacher_hidden[idx].float()
161
+ cosine = 1.0 - F.cosine_similarity(s, t, dim=-1).mean()
162
+ nmse = (s - t).square().mean() / t.square().mean().clamp_min(1e-5)
163
+ terms.append(cosine + 0.25 * nmse)
164
+ return torch.stack(terms).mean()
165
+
166
+
167
+ def alignment_metrics(pred: torch.Tensor, target: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
168
+ pred32 = pred.float()
169
+ target32 = target.float()
170
+ mse = (pred32 - target32).square().mean()
171
+ nmse = mse / target32.square().mean().clamp_min(1e-5)
172
+ cosine = F.cosine_similarity(pred32, target32, dim=-1).mean()
173
+ return nmse, cosine
174
+
175
+
176
+ def alignment_loss(pred: torch.Tensor, target: torch.Tensor, cosine_weight: float) -> torch.Tensor:
177
+ nmse, cosine = alignment_metrics(pred, target)
178
+ return nmse + cosine_weight * (1.0 - cosine)
179
+
180
+
181
+ def _detach_tree(value):
182
+ if torch.is_tensor(value):
183
+ return value.detach()
184
+ if isinstance(value, tuple):
185
+ return tuple(_detach_tree(v) for v in value)
186
+ if isinstance(value, list):
187
+ return [_detach_tree(v) for v in value]
188
+ return value
189
+
190
+
191
+ def capture_attention_input(model, ids: torch.Tensor, layer_idx: int, amp, *, with_output: bool):
192
+ capture: dict[str, Any] = {}
193
+ module = model.model.layers[layer_idx].self_attn
194
+
195
+ def pre_hook(mod, args, kwargs):
196
+ hidden = args[0] if args else kwargs.get("hidden_states")
197
+ if hidden is None:
198
+ raise RuntimeError("attention hidden_states not found")
199
+ capture["hidden"] = hidden
200
+ for key in (
201
+ "position_embeddings",
202
+ "position_ids",
203
+ "attention_mask",
204
+ "cache_position",
205
+ ):
206
+ if key in kwargs and kwargs[key] is not None:
207
+ capture[key] = _detach_tree(kwargs[key])
208
+
209
+ handle = module.register_forward_pre_hook(pre_hook, with_kwargs=True)
210
+ try:
211
+ if with_output:
212
+ with amp():
213
+ out = model(
214
+ input_ids=ids,
215
+ labels=ids,
216
+ use_cache=False,
217
+ output_hidden_states=True,
218
+ return_dict=True,
219
+ )
220
+ else:
221
+ with torch.no_grad(), amp():
222
+ out = model(
223
+ input_ids=ids,
224
+ use_cache=False,
225
+ output_hidden_states=False,
226
+ return_dict=True,
227
+ )
228
+ finally:
229
+ handle.remove()
230
+ if "hidden" not in capture:
231
+ raise RuntimeError(f"failed to capture layer {layer_idx} attention input")
232
+ return capture, out
233
+
234
+
235
+ def attention_kwargs_from_capture(capture: dict[str, Any]) -> dict[str, Any]:
236
+ kwargs: dict[str, Any] = {"use_cache": False}
237
+ for key in (
238
+ "position_embeddings",
239
+ "position_ids",
240
+ "attention_mask",
241
+ "cache_position",
242
+ ):
243
+ if key in capture:
244
+ kwargs[key] = capture[key]
245
+ return kwargs
246
+
247
+
248
+ def call_attention(module, hidden: torch.Tensor, kwargs: dict[str, Any]) -> torch.Tensor:
249
+ out = module(hidden, **kwargs)
250
+ if isinstance(out, (tuple, list)):
251
+ return out[0]
252
+ return out
253
+
254
+
255
+ def alpha_for_step(step: int, max_steps: int, start: float, end: float) -> float:
256
+ if max_steps <= 1:
257
+ return float(end)
258
+ progress = min(max((step - 1) / (max_steps - 1), 0.0), 1.0)
259
+ return float(start + progress * (end - start))
260
+
261
+
262
+ def freeze_current_layer_only(student, layer_idx: int) -> list[torch.nn.Parameter]:
263
+ for p in student.parameters():
264
+ p.requires_grad = False
265
+ module = student.model.layers[layer_idx].self_attn
266
+ if not isinstance(module, MemoryFusionLlamaAttention):
267
+ raise TypeError(f"layer {layer_idx} is not MemoryFusionLlamaAttention")
268
+ trainable = []
269
+ for p in module.core.parameters():
270
+ p.requires_grad = True
271
+ trainable.append(p)
272
+ for p in module.o_proj.parameters():
273
+ p.requires_grad = True
274
+ trainable.append(p)
275
+ return trainable
276
+
277
+
278
+ def select_integrated_stage_params(student, stage: str) -> list[torch.nn.Parameter]:
279
+ for p in student.parameters():
280
+ p.requires_grad = False
281
+
282
+ if stage in {"core_o", "core_o_norm"}:
283
+ for module in student.modules():
284
+ if isinstance(module, MemoryFusionLlamaAttention):
285
+ for p in module.core.parameters():
286
+ p.requires_grad = True
287
+ for p in module.o_proj.parameters():
288
+ p.requires_grad = True
289
+
290
+ if stage == "core_o_norm":
291
+ for name, p in student.named_parameters():
292
+ if "norm" in name.lower():
293
+ p.requires_grad = True
294
+ elif stage == "full":
295
+ for p in student.parameters():
296
+ p.requires_grad = True
297
+ else:
298
+ raise ValueError(f"unknown integrated stage {stage!r}")
299
+
300
+ return [p for p in student.parameters() if p.requires_grad]
301
+
302
+
303
+ def _layer_prefix(idx: int) -> str:
304
+ return f"model.layers.{idx}.self_attn."
305
+
306
+
307
+ def progressive_attention_state(student, accepted_layers: list[int]) -> dict[str, torch.Tensor]:
308
+ prefixes = tuple(_layer_prefix(i) for i in accepted_layers)
309
+ return {
310
+ k: v.detach().cpu()
311
+ for k, v in student.state_dict().items()
312
+ if prefixes and k.startswith(prefixes)
313
+ }
314
+
315
+
316
+ def save_progress(
317
+ output_dir: Path,
318
+ student,
319
+ config: SmolMemoryFusionConfig,
320
+ accepted_layers: list[int],
321
+ layer_reports: list[dict],
322
+ *,
323
+ stage: str,
324
+ ) -> None:
325
+ output_dir.mkdir(parents=True, exist_ok=True)
326
+ payload = {
327
+ "format_version": 1,
328
+ "stage": stage,
329
+ "accepted_layers": accepted_layers,
330
+ "config": config.to_dict(),
331
+ "layer_reports": layer_reports,
332
+ "attention_state": progressive_attention_state(student, accepted_layers),
333
+ }
334
+ tmp = output_dir / "sequential_progress.tmp"
335
+ final = output_dir / "sequential_progress.pt"
336
+ torch.save(payload, tmp)
337
+ tmp.replace(final)
338
+ (output_dir / "sequential_progress.json").write_text(
339
+ json.dumps(
340
+ {
341
+ "format_version": 1,
342
+ "stage": stage,
343
+ "accepted_layers": accepted_layers,
344
+ "config": config.to_dict(),
345
+ "layer_reports": layer_reports,
346
+ },
347
+ indent=2,
348
+ ),
349
+ encoding="utf-8",
350
+ )
351
+
352
+
353
+ def load_progress_if_available(
354
+ output_dir: Path,
355
+ student,
356
+ config: SmolMemoryFusionConfig,
357
+ *,
358
+ resume: bool,
359
+ ) -> tuple[list[int], list[dict]]:
360
+ path = output_dir / "sequential_progress.pt"
361
+ if not resume or not path.exists():
362
+ return [], []
363
+ payload = torch.load(path, map_location="cpu", weights_only=False)
364
+ accepted = [int(x) for x in payload.get("accepted_layers", [])]
365
+ if payload.get("config", {}).get("memory_rank") != config.memory_rank:
366
+ raise RuntimeError("resume checkpoint memory rank differs from requested rank")
367
+ if accepted:
368
+ replace_attention_layers(student, config, accepted)
369
+ incompatible = student.load_state_dict(payload["attention_state"], strict=False)
370
+ expected = {
371
+ k
372
+ for k in student.state_dict()
373
+ if any(k.startswith(_layer_prefix(i)) for i in accepted)
374
+ }
375
+ missing = [k for k in incompatible.missing_keys if k in expected]
376
+ if missing:
377
+ raise RuntimeError(f"resume checkpoint missing accepted-layer keys: {missing[:8]}")
378
+ print(f"RESUME: accepted layers={accepted}")
379
+ return accepted, list(payload.get("layer_reports", []))
380
+
381
+
382
+ def save_full_state(
383
+ output_dir: Path,
384
+ student,
385
+ config: SmolMemoryFusionConfig,
386
+ metadata: dict,
387
+ *,
388
+ filename: str = "smollm2_memory_fusion_sequential_full.pt",
389
+ ) -> None:
390
+ output_dir.mkdir(parents=True, exist_ok=True)
391
+ state_path = output_dir / filename
392
+ tmp = output_dir / (filename + ".tmp")
393
+ torch.save({k: v.detach().cpu() for k, v in student.state_dict().items()}, tmp)
394
+ tmp.replace(state_path)
395
+ meta = {
396
+ "format_version": 1,
397
+ "architecture": "smollm2-memory-fusion-sequential-full",
398
+ "base_model": metadata["base_model"],
399
+ "memory_fusion": config.to_dict(),
400
+ "training": metadata,
401
+ "state_file": filename,
402
+ }
403
+ (output_dir / "smollm2_memory_fusion_sequential_config.json").write_text(
404
+ json.dumps(meta, indent=2),
405
+ encoding="utf-8",
406
+ )
407
+
408
+
409
+ def build_probe_blocks(tokenizer, context: int, count: int) -> list[torch.Tensor]:
410
+ try:
411
+ wiki = load_dataset(
412
+ "Salesforce/wikitext",
413
+ "wikitext-2-raw-v1",
414
+ split="validation",
415
+ )
416
+ except Exception:
417
+ texts = [
418
+ "The history of science is shaped by observation, measurement, and careful comparison.",
419
+ "A language model predicts the next token from the sequence that came before it.",
420
+ "Vienna is the capital of Austria and has a long history of music, science, and public administration.",
421
+ "Neural networks can approximate complex functions when their parameters are trained on representative data.",
422
+ "The experiment should separate training data from the probe used to decide whether a replacement is acceptable.",
423
+ "Reliable evaluation requires the same inputs, tokenizer, context length, and scoring rule for every model.",
424
+ "A recurrent memory can carry information through a sequence without constructing a full dense attention matrix.",
425
+ "When a model component is replaced, small approximation errors can accumulate across many layers.",
426
+ ]
427
+ stream = "\n".join(texts)
428
+ else:
429
+ stream = "\n".join(str(x) for x in wiki["text"] if str(x).strip())
430
+
431
+ ids = tokenizer(stream, add_special_tokens=False)["input_ids"]
432
+ need = context
433
+ blocks = []
434
+ for start in range(0, len(ids) - need + 1, need):
435
+ blocks.append(torch.tensor(ids[start : start + need], dtype=torch.long))
436
+ if len(blocks) >= count:
437
+ break
438
+ if not blocks:
439
+ raise RuntimeError("could not construct probe blocks")
440
+ return blocks
441
+
442
+
443
+ @torch.no_grad()
444
+ def probe_nll(model, probe_blocks: list[torch.Tensor], device: torch.device, amp) -> float:
445
+ model.eval()
446
+ losses = []
447
+ for block in probe_blocks:
448
+ x = block.unsqueeze(0).to(device)
449
+ with amp():
450
+ out = model(input_ids=x, labels=x, use_cache=False, return_dict=True)
451
+ losses.append(float(out.loss.detach().float()))
452
+ return sum(losses) / len(losses)
453
+
454
+
455
+ @torch.no_grad()
456
+ def real_hidden_function_metrics(
457
+ teacher,
458
+ student,
459
+ ids: torch.Tensor,
460
+ layer_idx: int,
461
+ amp,
462
+ ) -> tuple[float, float]:
463
+ t_capture, _ = capture_attention_input(teacher, ids, layer_idx, amp, with_output=False)
464
+ s_capture, _ = capture_attention_input(student, ids, layer_idx, amp, with_output=False)
465
+ real_hidden = s_capture["hidden"].detach()
466
+ kwargs = attention_kwargs_from_capture(t_capture)
467
+ target = call_attention(teacher.model.layers[layer_idx].self_attn, real_hidden, kwargs)
468
+ pred = call_attention(student.model.layers[layer_idx].self_attn, real_hidden, kwargs)
469
+ nmse, cosine = alignment_metrics(pred, target)
470
+ return float(nmse), float(cosine)
471
+
472
+
473
+ def acceptance_passes(
474
+ *,
475
+ nmse: float,
476
+ cosine: float,
477
+ incremental_delta_nll: float,
478
+ cumulative_delta_nll: float,
479
+ args: argparse.Namespace,
480
+ ) -> bool:
481
+ return (
482
+ nmse <= args.accept_nmse
483
+ and cosine >= args.accept_cosine
484
+ and incremental_delta_nll <= args.accept_incremental_delta_nll
485
+ and cumulative_delta_nll <= args.accept_cumulative_delta_nll
486
+ )
487
+
488
+
489
+ def train_one_replacement(
490
+ *,
491
+ teacher,
492
+ student,
493
+ layer_idx: int,
494
+ batch_iter,
495
+ probe_blocks,
496
+ teacher_probe_nll: float,
497
+ pre_replacement_probe_nll: float,
498
+ args,
499
+ device,
500
+ amp,
501
+ ) -> dict:
502
+ trainable = freeze_current_layer_only(student, layer_idx)
503
+ student.train()
504
+ optimizer = make_optimizer(trainable, args.layer_lr, device)
505
+ scaler = torch.amp.GradScaler(
506
+ "cuda", enabled=(device.type == "cuda" and choose_dtype(device) == torch.float16)
507
+ )
508
+ best: dict[str, float | int | bool] | None = None
509
+ accepted = False
510
+
511
+ for step in range(1, args.max_layer_steps + 1):
512
+ ids = next(batch_iter).to(device, non_blocking=True)
513
+
514
+ t_capture, t_out = capture_attention_input(
515
+ teacher, ids, layer_idx, amp, with_output=True
516
+ )
517
+ s_capture, s_out = capture_attention_input(
518
+ student, ids, layer_idx, amp, with_output=True
519
+ )
520
+
521
+ alpha = alpha_for_step(
522
+ step,
523
+ args.max_layer_steps,
524
+ args.teacher_alpha_start,
525
+ args.teacher_alpha_end,
526
+ )
527
+ mixed_hidden = (
528
+ alpha * t_capture["hidden"].detach()
529
+ + (1.0 - alpha) * s_capture["hidden"].detach()
530
+ )
531
+ kwargs = attention_kwargs_from_capture(t_capture)
532
+
533
+ with torch.no_grad(), amp():
534
+ target = call_attention(
535
+ teacher.model.layers[layer_idx].self_attn,
536
+ mixed_hidden,
537
+ kwargs,
538
+ )
539
+
540
+ with amp():
541
+ pred = call_attention(
542
+ student.model.layers[layer_idx].self_attn,
543
+ mixed_hidden,
544
+ kwargs,
545
+ )
546
+ functional = alignment_loss(pred, target, args.cosine_weight)
547
+ kl = distillation_kl(s_out.logits, t_out.logits, args.temperature)
548
+ total = functional + args.layer_kl_weight * kl + args.layer_ce_weight * s_out.loss
549
+
550
+ if not torch.isfinite(total):
551
+ raise RuntimeError(f"non-finite loss at layer {layer_idx}, step {step}")
552
+
553
+ optimizer.zero_grad(set_to_none=True)
554
+ scaler.scale(total).backward()
555
+ scaler.unscale_(optimizer)
556
+ grad_norm = float(torch.nn.utils.clip_grad_norm_(trainable, 1.0))
557
+ scaler.step(optimizer)
558
+ scaler.update()
559
+
560
+ if step == 1 or step % args.log_every == 0:
561
+ nmse_batch, cosine_batch = alignment_metrics(pred.detach(), target.detach())
562
+ print(
563
+ f"layer={layer_idx:02d} step={step:03d}/{args.max_layer_steps} "
564
+ f"alpha={alpha:.3f} functional={float(functional.detach()):.4f} "
565
+ f"nmse={float(nmse_batch):.4f} cos={float(cosine_batch):.4f} "
566
+ f"kl={float(kl.detach()):.4f} ce={float(s_out.loss.detach()):.4f} "
567
+ f"grad={grad_norm:.3f}"
568
+ )
569
+
570
+ should_check = (
571
+ step >= args.min_layer_steps
572
+ and (step % args.check_every == 0 or step == args.max_layer_steps)
573
+ )
574
+ if should_check:
575
+ nmse, cosine = real_hidden_function_metrics(
576
+ teacher, student, ids, layer_idx, amp
577
+ )
578
+ current_probe_nll = probe_nll(student, probe_blocks, device, amp)
579
+ student.train()
580
+ incremental = current_probe_nll - pre_replacement_probe_nll
581
+ cumulative = current_probe_nll - teacher_probe_nll
582
+ passed = acceptance_passes(
583
+ nmse=nmse,
584
+ cosine=cosine,
585
+ incremental_delta_nll=incremental,
586
+ cumulative_delta_nll=cumulative,
587
+ args=args,
588
+ )
589
+ candidate = {
590
+ "step": step,
591
+ "nmse": nmse,
592
+ "cosine": cosine,
593
+ "probe_nll": current_probe_nll,
594
+ "incremental_delta_nll": incremental,
595
+ "cumulative_delta_nll": cumulative,
596
+ "passed": passed,
597
+ }
598
+ if best is None or (
599
+ candidate["nmse"] < best["nmse"]
600
+ and candidate["cumulative_delta_nll"] <= best["cumulative_delta_nll"] + 0.01
601
+ ):
602
+ best = candidate
603
+ print(
604
+ f" ACCEPTANCE CHECK layer={layer_idx:02d}: "
605
+ f"NMSE={nmse:.4f} (≤{args.accept_nmse:.4f}) "
606
+ f"cos={cosine:.4f} (≥{args.accept_cosine:.4f}) "
607
+ f"ΔNLL_inc={incremental:+.5f} (≤{args.accept_incremental_delta_nll:+.5f}) "
608
+ f"ΔNLL_total={cumulative:+.5f} (≤{args.accept_cumulative_delta_nll:+.5f}) "
609
+ f"=> {'PASS' if passed else 'continue'}"
610
+ )
611
+ if passed:
612
+ accepted = True
613
+ best = candidate
614
+ break
615
+
616
+ del t_out, s_out, pred, target, total
617
+
618
+ if best is None:
619
+ raise RuntimeError("acceptance was never evaluated")
620
+ return {
621
+ "layer": layer_idx,
622
+ "accepted": accepted,
623
+ "steps": int(best["step"]),
624
+ "nmse": float(best["nmse"]),
625
+ "cosine": float(best["cosine"]),
626
+ "probe_nll": float(best["probe_nll"]),
627
+ "incremental_delta_nll": float(best["incremental_delta_nll"]),
628
+ "cumulative_delta_nll": float(best["cumulative_delta_nll"]),
629
+ }
630
+
631
+
632
+ def run_integrated_stage(
633
+ *,
634
+ teacher,
635
+ student,
636
+ batch_iter,
637
+ stage_name: str,
638
+ token_budget: int,
639
+ lr: float,
640
+ args,
641
+ device,
642
+ amp,
643
+ output_dir: Path,
644
+ config: SmolMemoryFusionConfig,
645
+ report: dict,
646
+ ) -> dict:
647
+ if token_budget <= 0:
648
+ return {"stage": stage_name, "tokens": 0, "updates": 0, "skipped": True}
649
+
650
+ trainable = select_integrated_stage_params(student, stage_name)
651
+ optimizer = make_optimizer(trainable, lr, device)
652
+ scaler = torch.amp.GradScaler(
653
+ "cuda", enabled=(device.type == "cuda" and choose_dtype(device) == torch.float16)
654
+ )
655
+ student.train()
656
+ optimizer.zero_grad(set_to_none=True)
657
+
658
+ seen_tokens = 0
659
+ micro = 0
660
+ updates = 0
661
+ last = {}
662
+ start = time.perf_counter()
663
+
664
+ while seen_tokens < token_budget:
665
+ ids = next(batch_iter).to(device, non_blocking=True)
666
+ with torch.no_grad(), amp():
667
+ t_out = teacher(
668
+ input_ids=ids,
669
+ use_cache=False,
670
+ output_hidden_states=True,
671
+ return_dict=True,
672
+ )
673
+ with amp():
674
+ s_out = student(
675
+ input_ids=ids,
676
+ labels=ids,
677
+ use_cache=False,
678
+ output_hidden_states=True,
679
+ return_dict=True,
680
+ )
681
+ kl = distillation_kl(s_out.logits, t_out.logits, args.temperature)
682
+ hidden = representation_loss(s_out.hidden_states, t_out.hidden_states)
683
+ loss = (
684
+ args.ce_weight * s_out.loss
685
+ + args.kl_weight * kl
686
+ + args.hidden_weight * hidden
687
+ )
688
+ scaled = loss / args.grad_accum
689
+
690
+ if not torch.isfinite(scaled):
691
+ raise RuntimeError(f"non-finite loss during integrated stage {stage_name}")
692
+
693
+ scaler.scale(scaled).backward()
694
+ micro += 1
695
+ seen_tokens += ids.numel()
696
+ last = {
697
+ "ce": float(s_out.loss.detach().float()),
698
+ "kl": float(kl.detach().float()),
699
+ "hidden": float(hidden.detach().float()),
700
+ "total": float(loss.detach().float()),
701
+ }
702
+ del t_out, s_out
703
+
704
+ if micro % args.grad_accum:
705
+ continue
706
+
707
+ scaler.unscale_(optimizer)
708
+ grad_norm = float(torch.nn.utils.clip_grad_norm_(trainable, 1.0))
709
+ scaler.step(optimizer)
710
+ scaler.update()
711
+ optimizer.zero_grad(set_to_none=True)
712
+ updates += 1
713
+
714
+ if updates == 1 or updates % args.log_every == 0:
715
+ print(
716
+ f"stage={stage_name} update={updates} tokens={seen_tokens:,}/{token_budget:,} "
717
+ f"ce={last['ce']:.4f} kl={last['kl']:.4f} hidden={last['hidden']:.4f} "
718
+ f"grad={grad_norm:.3f}"
719
+ )
720
+
721
+ if args.backup_every_updates > 0 and updates % args.backup_every_updates == 0:
722
+ backup_meta = dict(report)
723
+ backup_meta["integrated_stage"] = stage_name
724
+ backup_meta["integrated_stage_tokens"] = seen_tokens
725
+ save_full_state(
726
+ output_dir,
727
+ student,
728
+ config,
729
+ backup_meta,
730
+ filename="live_sequential_full_state.pt",
731
+ )
732
+ print(" persistent full-state backup saved")
733
+
734
+ elapsed = time.perf_counter() - start
735
+ return {
736
+ "stage": stage_name,
737
+ "tokens": seen_tokens,
738
+ "updates": updates,
739
+ "lr": lr,
740
+ "elapsed_minutes": elapsed / 60.0,
741
+ "last": last,
742
+ "trainable_parameters": sum(p.numel() for p in trainable),
743
+ }
744
+
745
+
746
+ def main() -> None:
747
+ args = parse_args()
748
+ set_seed(args.seed)
749
+ output_dir = Path(args.output_dir)
750
+ output_dir.mkdir(parents=True, exist_ok=True)
751
+
752
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
753
+ dtype = choose_dtype(device)
754
+ amp = amp_factory(device, dtype)
755
+ if device.type == "cuda":
756
+ torch.cuda.reset_peak_memory_stats()
757
+ torch.backends.cuda.matmul.allow_tf32 = True
758
+
759
+ print(f"device={device} dtype={dtype} memory_rank={args.memory_rank}")
760
+ print(f"output_dir={output_dir}")
761
+
762
+ tokenizer = AutoTokenizer.from_pretrained(args.base_model, use_fast=True)
763
+ if tokenizer.pad_token_id is None:
764
+ tokenizer.pad_token = tokenizer.eos_token
765
+
766
+ teacher = AutoModelForCausalLM.from_pretrained(args.base_model, dtype=dtype).to(device)
767
+ teacher.eval()
768
+ teacher.config.use_cache = False
769
+ teacher.requires_grad_(False)
770
+
771
+ student = AutoModelForCausalLM.from_pretrained(args.base_model, dtype=dtype).to(device)
772
+ student.config.use_cache = False
773
+
774
+ config = SmolMemoryFusionConfig(
775
+ feature_dim=args.feature_dim,
776
+ memory_rank=args.memory_rank,
777
+ train_output_projection=True,
778
+ )
779
+
780
+ accepted_layers, layer_reports = load_progress_if_available(
781
+ output_dir, student, config, resume=args.resume
782
+ )
783
+
784
+ raw = load_dataset(
785
+ args.dataset,
786
+ name=args.dataset_config,
787
+ split=args.split,
788
+ streaming=True,
789
+ ).shuffle(seed=args.seed, buffer_size=args.shuffle_buffer)
790
+ batch_iter = iter(
791
+ batches(
792
+ token_blocks(raw, tokenizer, args.text_field, args.context_length),
793
+ args.batch_size,
794
+ )
795
+ )
796
+
797
+ probe_blocks = build_probe_blocks(
798
+ tokenizer,
799
+ context=args.probe_context,
800
+ count=args.probe_blocks,
801
+ )
802
+ teacher_probe_nll = probe_nll(teacher, probe_blocks, device, amp)
803
+ print(f"teacher probe NLL={teacher_probe_nll:.6f}")
804
+
805
+ start_time = time.perf_counter()
806
+ num_layers = int(student.config.num_hidden_layers)
807
+
808
+ for layer_idx in range(num_layers):
809
+ if layer_idx in accepted_layers:
810
+ continue
811
+ if accepted_layers != list(range(layer_idx)):
812
+ raise RuntimeError(
813
+ f"accepted layers must be a contiguous prefix before layer {layer_idx}: "
814
+ f"{accepted_layers}"
815
+ )
816
+ if (time.perf_counter() - start_time) / 60.0 >= args.max_runtime_minutes * 0.70:
817
+ raise RuntimeError(
818
+ "runtime budget reached during sequential replacement; progress was "
819
+ "saved and the same command can resume from the last accepted layer"
820
+ )
821
+
822
+ print("\n" + "=" * 110)
823
+ print(f"SEQUENTIAL REPLACEMENT: layer {layer_idx}/{num_layers - 1}")
824
+ print("=" * 110)
825
+
826
+ pre_probe_nll = probe_nll(student, probe_blocks, device, amp)
827
+ print(
828
+ f"before replacement: probe NLL={pre_probe_nll:.6f}, "
829
+ f"Δ vs teacher={pre_probe_nll - teacher_probe_nll:+.6f}"
830
+ )
831
+
832
+ replace_attention_layers(student, config, [layer_idx])
833
+
834
+ layer_report = train_one_replacement(
835
+ teacher=teacher,
836
+ student=student,
837
+ layer_idx=layer_idx,
838
+ batch_iter=batch_iter,
839
+ probe_blocks=probe_blocks,
840
+ teacher_probe_nll=teacher_probe_nll,
841
+ pre_replacement_probe_nll=pre_probe_nll,
842
+ args=args,
843
+ device=device,
844
+ amp=amp,
845
+ )
846
+ layer_reports.append(layer_report)
847
+
848
+ if not layer_report["accepted"] and args.strict_acceptance:
849
+ save_progress(
850
+ output_dir,
851
+ student,
852
+ config,
853
+ accepted_layers,
854
+ layer_reports,
855
+ stage=f"layer_{layer_idx}_rejected",
856
+ )
857
+ (output_dir / "sequential_training_report.json").write_text(
858
+ json.dumps(
859
+ {
860
+ "status": "stopped_on_rejected_layer",
861
+ "base_model": args.base_model,
862
+ "memory_rank": args.memory_rank,
863
+ "accepted_layers": accepted_layers,
864
+ "layer_reports": layer_reports,
865
+ "thresholds": {
866
+ "nmse": args.accept_nmse,
867
+ "cosine": args.accept_cosine,
868
+ "incremental_delta_nll": args.accept_incremental_delta_nll,
869
+ "cumulative_delta_nll": args.accept_cumulative_delta_nll,
870
+ },
871
+ },
872
+ indent=2,
873
+ ),
874
+ encoding="utf-8",
875
+ )
876
+ raise RuntimeError(
877
+ f"layer {layer_idx} did not satisfy acceptance criteria; "
878
+ "the next Transformer layer was NOT replaced"
879
+ )
880
+
881
+ if not layer_report["accepted"]:
882
+ print("WARNING: relaxed mode accepts a layer that did not pass all thresholds")
883
+
884
+ accepted_layers.append(layer_idx)
885
+ save_progress(
886
+ output_dir,
887
+ student,
888
+ config,
889
+ accepted_layers,
890
+ layer_reports,
891
+ stage=f"accepted_layer_{layer_idx}",
892
+ )
893
+ print(f"✅ accepted layer {layer_idx}; persistent progress saved")
894
+
895
+ if accepted_layers != list(range(num_layers)):
896
+ raise RuntimeError("not all layers were accepted")
897
+
898
+ summary = structural_summary(student)
899
+ if summary["memory_fusion_layers"] != num_layers or summary["transformer_attention_layers"] != 0:
900
+ raise RuntimeError(f"unexpected final structure: {summary}")
901
+
902
+ report = {
903
+ "status": "all_layers_accepted",
904
+ "architecture": "smollm2-memory-fusion-sequential-full",
905
+ "base_model": args.base_model,
906
+ "memory_rank": args.memory_rank,
907
+ "feature_dim": args.feature_dim,
908
+ "context_length": args.context_length,
909
+ "teacher_probe_nll": teacher_probe_nll,
910
+ "accepted_layers": accepted_layers,
911
+ "layer_reports": layer_reports,
912
+ "thresholds": {
913
+ "nmse": args.accept_nmse,
914
+ "cosine": args.accept_cosine,
915
+ "incremental_delta_nll": args.accept_incremental_delta_nll,
916
+ "cumulative_delta_nll": args.accept_cumulative_delta_nll,
917
+ },
918
+ "integrated_stages": [],
919
+ }
920
+
921
+ print("\n" + "=" * 110)
922
+ print("ALL 30 REPLACEMENTS ACCEPTED — STARTING INTEGRATED TRAINING")
923
+ print("=" * 110)
924
+
925
+ integrated_specs = [
926
+ ("core_o", args.core_o_tokens, args.core_o_lr),
927
+ ("core_o_norm", args.norm_tokens, args.norm_lr),
928
+ ("full", args.full_tokens, args.full_lr),
929
+ ]
930
+ for stage_name, token_budget, lr in integrated_specs:
931
+ if (time.perf_counter() - start_time) / 60.0 >= args.max_runtime_minutes:
932
+ print("runtime budget reached before remaining integrated stages")
933
+ report["status"] = "runtime_budget_after_acceptance"
934
+ break
935
+ print(f"\n--- integrated stage: {stage_name} ---")
936
+ stage_report = run_integrated_stage(
937
+ teacher=teacher,
938
+ student=student,
939
+ batch_iter=batch_iter,
940
+ stage_name=stage_name,
941
+ token_budget=token_budget,
942
+ lr=lr,
943
+ args=args,
944
+ device=device,
945
+ amp=amp,
946
+ output_dir=output_dir,
947
+ config=config,
948
+ report=report,
949
+ )
950
+ report["integrated_stages"].append(stage_report)
951
+ save_full_state(
952
+ output_dir,
953
+ student,
954
+ config,
955
+ report,
956
+ filename=f"stage_{stage_name}_full_state.pt",
957
+ )
958
+ print(f"✅ completed {stage_name}; full persistent checkpoint saved")
959
+
960
+ student.eval()
961
+ final_probe_nll = probe_nll(student, probe_blocks, device, amp)
962
+ report["final_probe_nll"] = final_probe_nll
963
+ report["final_probe_delta_nll"] = final_probe_nll - teacher_probe_nll
964
+ report["parameters"] = parameter_summary(student)
965
+ report["structure"] = structural_summary(student)
966
+ report["elapsed_minutes"] = (time.perf_counter() - start_time) / 60.0
967
+ report["peak_vram_gib"] = (
968
+ torch.cuda.max_memory_allocated() / (1024 ** 3)
969
+ if device.type == "cuda"
970
+ else 0.0
971
+ )
972
+
973
+ save_full_state(output_dir, student, config, report)
974
+ (output_dir / "sequential_training_report.json").write_text(
975
+ json.dumps(report, indent=2),
976
+ encoding="utf-8",
977
+ )
978
+ tokenizer.save_pretrained(output_dir)
979
+
980
+ print("\nFINAL REPORT")
981
+ print(json.dumps(report, indent=2))
982
+ print("saved:", output_dir)
983
+
984
+
985
+ if __name__ == "__main__":
986
+ main()
src/tinycenn_lm/__init__.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .live_console import configure_live_console
2
+
3
+ # Configure the current process before importing model modules. Every train_*.py
4
+ # script imports tinycenn_lm, so notebook-launched trainers inherit immediate,
5
+ # line-buffered stdout/stderr even when the notebook uses subprocess.run(...).
6
+ configure_live_console()
7
+
8
+ from .cenn import CeNNConfig, FastCeNNCore
9
+ from .modeling import (
10
+ DEFAULT_BASE_MODEL,
11
+ HybridDecoderLayer,
12
+ build_from_adapter,
13
+ freeze_for_adapter_training,
14
+ inject_cenn,
15
+ load_adapter,
16
+ save_adapter,
17
+ trainable_parameter_summary,
18
+ )
19
+ from .student import (
20
+ CeNNReplacementLayer,
21
+ build_cenn_student,
22
+ freeze_student_interfaces,
23
+ load_cenn_student_weights,
24
+ replace_transformer_with_cenn,
25
+ save_cenn_student,
26
+ student_parameter_summary,
27
+ )
28
+ from .moe import (
29
+ FastMoECeNNCore,
30
+ MoECeNNConfig,
31
+ MoECeNNReplacementLayer,
32
+ build_moe_cenn_student,
33
+ freeze_moe_student_interfaces,
34
+ load_moe_cenn_student_weights,
35
+ moe_router_stats,
36
+ replace_transformer_with_moe_cenn,
37
+ save_moe_cenn_student,
38
+ warmstart_moe_from_plain_cenn,
39
+ )
40
+ from .sharded_moe import (
41
+ FastShardedMoECeNNCore,
42
+ ShardedMoECeNNConfig,
43
+ ShardedMoECeNNReplacementLayer,
44
+ build_sharded_moe_student,
45
+ freeze_sharded_moe_interfaces,
46
+ load_sharded_moe_student_weights,
47
+ replace_transformer_with_sharded_moe_cenn,
48
+ save_sharded_moe_student,
49
+ sharded_router_stats,
50
+ warmstart_sharded_moe_from_plain_cenn,
51
+ )
52
+ from .story_v2 import (
53
+ CausalStoryMemory,
54
+ LowRankLMHeadAdapter,
55
+ StoryV2Config,
56
+ StoryV2ReplacementLayer,
57
+ build_story_v2_from_story_v1,
58
+ build_story_v2_student,
59
+ freeze_story_v2_interfaces,
60
+ load_story_v2_weights,
61
+ save_story_v2_student,
62
+ story_v2_parameter_summary,
63
+ story_v2_router_stats,
64
+ upgrade_sharded_model_to_story_v2,
65
+ )
66
+ from .smollm2_amcenn import (
67
+ DEFAULT_SMOLLM2,
68
+ AMCeNNAttention,
69
+ PositiveSoftmaxFeatures,
70
+ ShardedTop2LlamaMLP,
71
+ SmolAMCeNNConfig,
72
+ amcenn_parameter_summary,
73
+ amcenn_router_stats,
74
+ build_smollm2_amcenn,
75
+ freeze_smollm2_for_amcenn_training,
76
+ load_smollm2_amcenn_weights,
77
+ replace_smollm2_core,
78
+ save_smollm2_amcenn,
79
+ )
80
+ from .smollm2_amcenn_v2 import (
81
+ AMCeNNAttentionV2,
82
+ AdaptivePositiveSoftmaxFeatures,
83
+ SmolAMCeNNV2Config,
84
+ build_smollm2_amcenn_v2,
85
+ convert_all_ffns_to_sharded_top2,
86
+ freeze_for_global_training,
87
+ freeze_for_group_calibration,
88
+ load_smollm2_amcenn_v2_weights,
89
+ replace_all_smollm2_attention,
90
+ replace_attention_layers,
91
+ save_smollm2_amcenn_v2,
92
+ v2_parameter_summary,
93
+ )
94
+ from .hf_persistence import (
95
+ build_model_card,
96
+ collect_reports,
97
+ install_colab_hf_upload_enhancer,
98
+ persist_hf_run,
99
+ redact_secrets,
100
+ utc_run_id,
101
+ )
102
+ from .colab_live_backup import (
103
+ install_colab_training_backup,
104
+ is_tinycenn_training_command,
105
+ output_dir_from_command,
106
+ )
107
+ from .direct_colab_backup import install_direct_training_backup
108
+
109
+ # In Colab, make Hugging Face backup mandatory. The parent notebook wrapper handles
110
+ # normal subprocess-launched trainers. A trainer-side fallback covers notebooks that
111
+ # launch train_*.py before the notebook kernel imports tinycenn_lm.
112
+ install_colab_training_backup()
113
+ install_direct_training_backup()
114
+ install_colab_hf_upload_enhancer()
115
+
116
+ __all__ = [
117
+ "configure_live_console",
118
+ "CeNNConfig", "FastCeNNCore", "DEFAULT_BASE_MODEL", "HybridDecoderLayer",
119
+ "build_from_adapter", "freeze_for_adapter_training", "inject_cenn", "load_adapter",
120
+ "save_adapter", "trainable_parameter_summary", "CeNNReplacementLayer", "build_cenn_student",
121
+ "freeze_student_interfaces", "load_cenn_student_weights", "replace_transformer_with_cenn",
122
+ "save_cenn_student", "student_parameter_summary", "MoECeNNConfig", "FastMoECeNNCore",
123
+ "MoECeNNReplacementLayer", "build_moe_cenn_student", "freeze_moe_student_interfaces",
124
+ "load_moe_cenn_student_weights", "moe_router_stats", "replace_transformer_with_moe_cenn",
125
+ "save_moe_cenn_student", "warmstart_moe_from_plain_cenn", "ShardedMoECeNNConfig",
126
+ "FastShardedMoECeNNCore", "ShardedMoECeNNReplacementLayer", "build_sharded_moe_student",
127
+ "freeze_sharded_moe_interfaces", "load_sharded_moe_student_weights",
128
+ "replace_transformer_with_sharded_moe_cenn", "save_sharded_moe_student", "sharded_router_stats",
129
+ "warmstart_sharded_moe_from_plain_cenn", "StoryV2Config", "CausalStoryMemory",
130
+ "StoryV2ReplacementLayer", "LowRankLMHeadAdapter", "upgrade_sharded_model_to_story_v2",
131
+ "freeze_story_v2_interfaces", "story_v2_router_stats", "story_v2_parameter_summary",
132
+ "save_story_v2_student", "load_story_v2_weights", "build_story_v2_student",
133
+ "build_story_v2_from_story_v1", "DEFAULT_SMOLLM2", "SmolAMCeNNConfig",
134
+ "PositiveSoftmaxFeatures", "AMCeNNAttention", "ShardedTop2LlamaMLP", "replace_smollm2_core",
135
+ "freeze_smollm2_for_amcenn_training", "amcenn_router_stats", "amcenn_parameter_summary",
136
+ "save_smollm2_amcenn", "load_smollm2_amcenn_weights", "build_smollm2_amcenn",
137
+ "SmolAMCeNNV2Config", "AdaptivePositiveSoftmaxFeatures", "AMCeNNAttentionV2",
138
+ "convert_all_ffns_to_sharded_top2", "replace_attention_layers", "replace_all_smollm2_attention",
139
+ "freeze_for_group_calibration", "freeze_for_global_training", "v2_parameter_summary",
140
+ "save_smollm2_amcenn_v2", "load_smollm2_amcenn_v2_weights", "build_smollm2_amcenn_v2",
141
+ "build_model_card", "collect_reports", "persist_hf_run", "install_colab_hf_upload_enhancer",
142
+ "redact_secrets", "utc_run_id", "install_colab_training_backup", "install_direct_training_backup",
143
+ "is_tinycenn_training_command", "output_dir_from_command",
144
+ ]
src/tinycenn_lm/__pycache__/__init__.cpython-313.pyc ADDED
Binary file (3.92 kB). View file
 
src/tinycenn_lm/__pycache__/cellular_attention.cpython-313.pyc ADDED
Binary file (37.2 kB). View file
 
src/tinycenn_lm/__pycache__/cenn.cpython-313.pyc ADDED
Binary file (11 kB). View file
 
src/tinycenn_lm/__pycache__/colab_live_backup.cpython-313.pyc ADDED
Binary file (20.9 kB). View file
 
src/tinycenn_lm/__pycache__/direct_colab_backup.cpython-313.pyc ADDED
Binary file (14.3 kB). View file
 
src/tinycenn_lm/__pycache__/hf_persistence.cpython-313.pyc ADDED
Binary file (22.5 kB). View file
 
src/tinycenn_lm/__pycache__/live_console.cpython-313.pyc ADDED
Binary file (4.36 kB). View file
 
src/tinycenn_lm/__pycache__/memory_attention.cpython-313.pyc ADDED
Binary file (28.5 kB). View file
 
src/tinycenn_lm/__pycache__/modeling.cpython-313.pyc ADDED
Binary file (12.4 kB). View file
 
src/tinycenn_lm/__pycache__/moe.cpython-313.pyc ADDED
Binary file (24.1 kB). View file
 
src/tinycenn_lm/__pycache__/pdelta2_er.cpython-313.pyc ADDED
Binary file (17 kB). View file
 
src/tinycenn_lm/__pycache__/pdelta2_features.cpython-313.pyc ADDED
Binary file (23.1 kB). View file
 
src/tinycenn_lm/__pycache__/pdelta3_frontier.cpython-313.pyc ADDED
Binary file (24.8 kB). View file
 
src/tinycenn_lm/__pycache__/research_layers.cpython-313.pyc ADDED
Binary file (17.9 kB). View file
 
src/tinycenn_lm/__pycache__/sharded_moe.cpython-313.pyc ADDED
Binary file (27.6 kB). View file
 
src/tinycenn_lm/__pycache__/smollm2_amcenn.cpython-313.pyc ADDED
Binary file (26 kB). View file
 
src/tinycenn_lm/__pycache__/smollm2_amcenn_v2.cpython-313.pyc ADDED
Binary file (23.5 kB). View file
 
src/tinycenn_lm/__pycache__/smollm2_memory_fusion.cpython-313.pyc ADDED
Binary file (19.1 kB). View file
 
src/tinycenn_lm/__pycache__/story_v2.cpython-313.pyc ADDED
Binary file (16.9 kB). View file
 
src/tinycenn_lm/__pycache__/student.cpython-313.pyc ADDED
Binary file (12.7 kB). View file
 
src/tinycenn_lm/cellular_attention.py ADDED
@@ -0,0 +1,572 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Causal 1-D Cellular Attention for TinyCeNN-LM research experiments.
2
+
3
+ The layer keeps attention sparse and local at each cellular step, but changes
4
+ the neighborhood across steps. Power-of-two dilations create an exponentially
5
+ growing receptive field without constructing a T-by-T attention matrix.
6
+
7
+ Advanced variants add token/head-specific routing, old-school max/mean pooling,
8
+ gated RMS normalization, and a strictly-causal two-level encoder-decoder path.
9
+ The U-AMP path uses only current/past states during downsampling and nearest-left
10
+ upsampling, so no future token can enter an earlier output.
11
+
12
+ This is a research reference implementation. It favors clarity and auditable
13
+ causality over fused-kernel speed.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ from typing import Iterable
19
+
20
+ import torch
21
+ from torch import Tensor, nn
22
+ import torch.nn.functional as F
23
+
24
+
25
+ VARIANTS = (
26
+ "cellular_local3",
27
+ "cellular_dilated3",
28
+ "cellular_dilated5",
29
+ "cellular_multiscale5",
30
+ "cellular_shifted8",
31
+ "cellular_adaptive_multiscale5",
32
+ "cellular_multiscale5_maxpool",
33
+ "cellular_adaptive_maxpool5",
34
+ "cellular_adaptive_maxpool5_rms",
35
+ "cellular_adaptive_mixedpool5_rms",
36
+ "cellular_uamp5",
37
+ "cellular_uamp5_channelgate",
38
+ "cellular_uamp5_varlatent",
39
+ )
40
+
41
+ _ADAPTIVE_VARIANTS = {
42
+ "cellular_adaptive_multiscale5",
43
+ "cellular_adaptive_maxpool5",
44
+ "cellular_adaptive_maxpool5_rms",
45
+ "cellular_adaptive_mixedpool5_rms",
46
+ "cellular_uamp5",
47
+ "cellular_uamp5_channelgate",
48
+ "cellular_uamp5_varlatent",
49
+ }
50
+ _MAXPOOL_VARIANTS = {
51
+ "cellular_multiscale5_maxpool",
52
+ "cellular_adaptive_maxpool5",
53
+ "cellular_adaptive_maxpool5_rms",
54
+ }
55
+ _MIXEDPOOL_VARIANTS = {
56
+ "cellular_adaptive_mixedpool5_rms",
57
+ "cellular_uamp5",
58
+ "cellular_uamp5_channelgate",
59
+ "cellular_uamp5_varlatent",
60
+ }
61
+ _RMS_VARIANTS = {
62
+ "cellular_adaptive_maxpool5_rms",
63
+ "cellular_adaptive_mixedpool5_rms",
64
+ "cellular_uamp5",
65
+ "cellular_uamp5_channelgate",
66
+ "cellular_uamp5_varlatent",
67
+ }
68
+ _UNET_VARIANTS = {
69
+ "cellular_uamp5",
70
+ "cellular_uamp5_channelgate",
71
+ "cellular_uamp5_varlatent",
72
+ }
73
+ _CHANNEL_GATE_VARIANTS = {"cellular_uamp5_channelgate"}
74
+ _VARLATENT_VARIANTS = {"cellular_uamp5_varlatent"}
75
+ _MULTISCALE_VARIANTS = {
76
+ "cellular_multiscale5",
77
+ "cellular_adaptive_multiscale5",
78
+ "cellular_multiscale5_maxpool",
79
+ "cellular_adaptive_maxpool5",
80
+ "cellular_adaptive_maxpool5_rms",
81
+ "cellular_adaptive_mixedpool5_rms",
82
+ "cellular_uamp5",
83
+ "cellular_uamp5_channelgate",
84
+ "cellular_uamp5_varlatent",
85
+ }
86
+
87
+
88
+ class CellularAttentionLayer(nn.Module):
89
+ """Sparse causal attention with optional U-AMP multi-resolution refinement."""
90
+
91
+ def __init__(
92
+ self,
93
+ num_heads: int,
94
+ num_kv_heads: int,
95
+ head_dim: int,
96
+ feature_dim: int = 64,
97
+ variant: str = "cellular_dilated3",
98
+ dilations: Iterable[int] = (1, 2, 4, 8, 16, 32, 64, 128),
99
+ shifted_window: int = 8,
100
+ ):
101
+ super().__init__()
102
+ if variant not in VARIANTS:
103
+ raise ValueError(f"unknown variant {variant!r}; choose from {VARIANTS}")
104
+ if min(num_heads, num_kv_heads, head_dim, feature_dim, shifted_window) < 1:
105
+ raise ValueError("all dimensions must be positive")
106
+ if num_heads % num_kv_heads:
107
+ raise ValueError("num_heads must be divisible by num_kv_heads")
108
+ dilations = tuple(int(d) for d in dilations)
109
+ if not dilations or min(dilations) < 1:
110
+ raise ValueError("dilations must contain positive integers")
111
+
112
+ self.num_heads = int(num_heads)
113
+ self.num_kv_heads = int(num_kv_heads)
114
+ self.head_dim = int(head_dim)
115
+ self.feature_dim = int(feature_dim)
116
+ self.groups = self.num_heads // self.num_kv_heads
117
+ self.variant = variant
118
+ self.dilations = dilations
119
+ self.shifted_window = int(shifted_window)
120
+ self.latent_dim = max(8, self.head_dim // 2)
121
+ self._last_aux_loss: Tensor | None = None
122
+
123
+ q_base = torch.zeros(self.num_heads, self.feature_dim, self.head_dim)
124
+ k_base = torch.zeros(self.num_kv_heads, self.feature_dim, self.head_dim)
125
+ for h in range(self.num_heads):
126
+ nn.init.orthogonal_(q_base[h])
127
+ for h in range(self.num_kv_heads):
128
+ nn.init.orthogonal_(k_base[h])
129
+ self.wq = nn.Parameter(q_base)
130
+ self.wk = nn.Parameter(k_base)
131
+
132
+ self.state_q = nn.Parameter(torch.zeros(
133
+ self.num_heads, self.feature_dim, self.head_dim
134
+ ))
135
+ self.state_q_gate = nn.Parameter(torch.full(
136
+ (len(self.dilations), self.num_heads), -2.0
137
+ ))
138
+
139
+ max_neighbors = max(
140
+ self._max_neighbors_for_variant(variant), self.shifted_window
141
+ )
142
+ self.relative_bias = nn.Parameter(torch.zeros(
143
+ len(self.dilations), self.num_heads, max_neighbors
144
+ ))
145
+ self.log_temperature = nn.Parameter(torch.zeros(
146
+ len(self.dilations), self.num_heads
147
+ ))
148
+ self.step_gate = nn.Parameter(torch.zeros(
149
+ len(self.dilations), self.num_heads
150
+ ))
151
+
152
+ if self.uses_adaptive_routing():
153
+ self.route_key = nn.Parameter(torch.empty(
154
+ len(self.dilations), self.num_heads, max_neighbors, self.feature_dim
155
+ ))
156
+ nn.init.normal_(self.route_key, mean=0.0, std=0.02)
157
+ self.route_prior = nn.Parameter(torch.zeros(
158
+ len(self.dilations), self.num_heads, max_neighbors
159
+ ))
160
+ self.route_strength = nn.Parameter(torch.zeros(
161
+ len(self.dilations), self.num_heads
162
+ ))
163
+ else:
164
+ self.register_parameter("route_key", None)
165
+ self.register_parameter("route_prior", None)
166
+ self.register_parameter("route_strength", None)
167
+
168
+ if self.uses_maxpool_branch():
169
+ self.pool_mix_logit = nn.Parameter(torch.full(
170
+ (len(self.dilations), self.num_heads), -1.5
171
+ ))
172
+ self.log_pool_gain = nn.Parameter(torch.zeros(
173
+ len(self.dilations), self.num_heads
174
+ ))
175
+ else:
176
+ self.register_parameter("pool_mix_logit", None)
177
+ self.register_parameter("log_pool_gain", None)
178
+
179
+ if self.uses_mixedpool_branch():
180
+ # [attention, max, mean], initialized to preserve the attention path.
181
+ initial = torch.tensor([2.0, -1.0, -1.0]).view(1, 1, 3)
182
+ self.mixed_pool_logits = nn.Parameter(
183
+ initial.expand(len(self.dilations), self.num_heads, 3).clone()
184
+ )
185
+ self.log_mixed_pool_gain = nn.Parameter(torch.zeros(
186
+ len(self.dilations), self.num_heads, 2
187
+ ))
188
+ else:
189
+ self.register_parameter("mixed_pool_logits", None)
190
+ self.register_parameter("log_mixed_pool_gain", None)
191
+
192
+ if self.uses_rms_refinement():
193
+ self.pre_rms_weight = nn.Parameter(torch.ones(self.num_heads, self.head_dim))
194
+ self.post_rms_weight = nn.Parameter(torch.ones(self.num_heads, self.head_dim))
195
+ self.pre_rms_gate = nn.Parameter(torch.full((self.num_heads,), -2.0))
196
+ self.post_rms_gate = nn.Parameter(torch.full((self.num_heads,), -2.0))
197
+ else:
198
+ self.register_parameter("pre_rms_weight", None)
199
+ self.register_parameter("post_rms_weight", None)
200
+ self.register_parameter("pre_rms_gate", None)
201
+ self.register_parameter("post_rms_gate", None)
202
+
203
+ if self.uses_unet_refinement():
204
+ # Causal two-level U-Net. Downsampling endpoints are 0,2,4,...;
205
+ # nearest-left upsampling never uses a future endpoint.
206
+ self.unet_pool_logits = nn.Parameter(torch.zeros(2, self.num_heads, 2))
207
+ self.unet_encoder = nn.Parameter(torch.empty(
208
+ self.num_heads, self.latent_dim, self.head_dim
209
+ ))
210
+ self.unet_decoder = nn.Parameter(torch.empty(
211
+ self.num_heads, self.head_dim, self.latent_dim
212
+ ))
213
+ self.unet_swiglu_gate = nn.Parameter(torch.empty(
214
+ self.num_heads, self.latent_dim, self.latent_dim
215
+ ))
216
+ self.unet_swiglu_value = nn.Parameter(torch.empty(
217
+ self.num_heads, self.latent_dim, self.latent_dim
218
+ ))
219
+ for tensor in (
220
+ self.unet_encoder, self.unet_decoder,
221
+ self.unet_swiglu_gate, self.unet_swiglu_value,
222
+ ):
223
+ for h in range(self.num_heads):
224
+ nn.init.orthogonal_(tensor[h])
225
+ self.unet_skip1_gate = nn.Parameter(torch.full((self.num_heads,), -1.5))
226
+ self.unet_skip0_gate = nn.Parameter(torch.full((self.num_heads,), -1.5))
227
+ self.unet_output_gate = nn.Parameter(torch.full((self.num_heads,), -2.0))
228
+ else:
229
+ self.register_parameter("unet_pool_logits", None)
230
+ self.register_parameter("unet_encoder", None)
231
+ self.register_parameter("unet_decoder", None)
232
+ self.register_parameter("unet_swiglu_gate", None)
233
+ self.register_parameter("unet_swiglu_value", None)
234
+ self.register_parameter("unet_skip1_gate", None)
235
+ self.register_parameter("unet_skip0_gate", None)
236
+ self.register_parameter("unet_output_gate", None)
237
+
238
+ if self.uses_channel_gate():
239
+ rank = max(4, self.head_dim // 4)
240
+ self.channel_down = nn.Parameter(torch.empty(
241
+ self.num_heads, rank, self.head_dim
242
+ ))
243
+ self.channel_up = nn.Parameter(torch.empty(
244
+ self.num_heads, self.head_dim, rank
245
+ ))
246
+ for h in range(self.num_heads):
247
+ nn.init.orthogonal_(self.channel_down[h])
248
+ nn.init.orthogonal_(self.channel_up[h])
249
+ self.channel_strength = nn.Parameter(torch.full((self.num_heads,), -2.0))
250
+ else:
251
+ self.register_parameter("channel_down", None)
252
+ self.register_parameter("channel_up", None)
253
+ self.register_parameter("channel_strength", None)
254
+
255
+ if self.uses_variational_latent():
256
+ self.var_logvar = nn.Parameter(torch.empty(
257
+ self.num_heads, self.latent_dim, self.head_dim
258
+ ))
259
+ for h in range(self.num_heads):
260
+ nn.init.normal_(self.var_logvar[h], mean=0.0, std=0.01)
261
+ self.var_logvar_bias = nn.Parameter(torch.full(
262
+ (self.num_heads, self.latent_dim), -3.0
263
+ ))
264
+ self.var_kl_weight = 1e-4
265
+ else:
266
+ self.register_parameter("var_logvar", None)
267
+ self.register_parameter("var_logvar_bias", None)
268
+ self.var_kl_weight = 0.0
269
+
270
+ eye = torch.eye(self.head_dim).expand(self.num_heads, -1, -1).clone()
271
+ self.out_proj = nn.Parameter(eye)
272
+ self.final_gate = nn.Parameter(torch.full((self.num_heads,), 2.0))
273
+ self.log_gain = nn.Parameter(torch.zeros(self.num_heads))
274
+
275
+ def uses_adaptive_routing(self) -> bool:
276
+ return self.variant in _ADAPTIVE_VARIANTS
277
+
278
+ def uses_maxpool_branch(self) -> bool:
279
+ return self.variant in _MAXPOOL_VARIANTS
280
+
281
+ def uses_mixedpool_branch(self) -> bool:
282
+ return self.variant in _MIXEDPOOL_VARIANTS
283
+
284
+ def uses_rms_refinement(self) -> bool:
285
+ return self.variant in _RMS_VARIANTS
286
+
287
+ def uses_unet_refinement(self) -> bool:
288
+ return self.variant in _UNET_VARIANTS
289
+
290
+ def uses_channel_gate(self) -> bool:
291
+ return self.variant in _CHANNEL_GATE_VARIANTS
292
+
293
+ def uses_variational_latent(self) -> bool:
294
+ return self.variant in _VARLATENT_VARIANTS
295
+
296
+ @staticmethod
297
+ def _max_neighbors_for_variant(variant: str) -> int:
298
+ if variant in ("cellular_local3", "cellular_dilated3"):
299
+ return 3
300
+ if variant in ("cellular_dilated5", *_MULTISCALE_VARIANTS):
301
+ return 5
302
+ if variant == "cellular_shifted8":
303
+ return 8
304
+ raise ValueError(variant)
305
+
306
+ @property
307
+ def config(self) -> dict:
308
+ return {
309
+ "num_heads": self.num_heads,
310
+ "num_kv_heads": self.num_kv_heads,
311
+ "head_dim": self.head_dim,
312
+ "feature_dim": self.feature_dim,
313
+ "variant": self.variant,
314
+ "dilations": list(self.dilations),
315
+ "shifted_window": self.shifted_window,
316
+ }
317
+
318
+ def _offsets(self, step: int) -> tuple[int, ...]:
319
+ d = self.dilations[step]
320
+ if self.variant == "cellular_local3":
321
+ return (0, 1, 2)
322
+ if self.variant == "cellular_dilated3":
323
+ return (0, d, 2 * d)
324
+ if self.variant == "cellular_dilated5":
325
+ return (0, d, 2 * d, 3 * d, 4 * d)
326
+ if self.variant in _MULTISCALE_VARIANTS:
327
+ return tuple(dict.fromkeys((0, 1, d, 2 * d, 4 * d)))
328
+ if self.variant == "cellular_shifted8":
329
+ return tuple(range(self.shifted_window))
330
+ raise ValueError(self.variant)
331
+
332
+ def _valid_mask(self, length: int, step: int, device) -> tuple[Tensor, Tensor]:
333
+ offsets = torch.tensor(self._offsets(step), device=device, dtype=torch.long)
334
+ pos = torch.arange(length, device=device, dtype=torch.long)
335
+ index = pos[:, None] - offsets[None, :]
336
+ valid = index >= 0
337
+
338
+ if self.variant == "cellular_shifted8":
339
+ shift = 0 if step % 2 == 0 else self.shifted_window // 2
340
+ block_start = ((pos + shift) // self.shifted_window) * self.shifted_window - shift
341
+ valid = valid & (index >= block_start[:, None])
342
+
343
+ return index.clamp_min(0), valid
344
+
345
+ def _features(self, q: Tensor, k: Tensor) -> tuple[Tensor, Tensor]:
346
+ q = F.normalize(q, dim=-1)
347
+ k = F.normalize(k, dim=-1)
348
+ qf = F.normalize(torch.einsum("bhtd,hfd->bhtf", q, self.wq), dim=-1)
349
+ kf = F.normalize(torch.einsum("bhtd,hfd->bhtf", k, self.wk), dim=-1)
350
+ return qf, kf.repeat_interleave(self.groups, dim=1)
351
+
352
+ @staticmethod
353
+ def _rms(x: Tensor, weight: Tensor) -> Tensor:
354
+ normed = x * torch.rsqrt(x.square().mean(dim=-1, keepdim=True) + 1e-6)
355
+ return normed * weight[None, :, None, :]
356
+
357
+ def _adaptive_route_log_prior(
358
+ self, query: Tensor, valid: Tensor, step: int, width: int,
359
+ ) -> Tensor:
360
+ assert self.route_key is not None
361
+ assert self.route_prior is not None
362
+ assert self.route_strength is not None
363
+ prototypes = self.route_key[step, :, :width, :]
364
+ route_logits = torch.einsum("bhtf,hwf->bhtw", query, prototypes)
365
+ route_logits = route_logits / math.sqrt(self.feature_dim)
366
+ route_logits = route_logits + self.route_prior[step, :, :width][None, :, None, :]
367
+ route_logits = route_logits.masked_fill(
368
+ ~valid[None, None, :, :], float("-inf")
369
+ )
370
+ route_log_prob = route_logits.log_softmax(dim=-1)
371
+ route_log_prob = route_log_prob.masked_fill(
372
+ ~valid[None, None, :, :], 0.0
373
+ )
374
+ strength = F.softplus(self.route_strength[step])[None, :, None, None]
375
+ return route_log_prob * strength
376
+
377
+ @staticmethod
378
+ def _pool_messages(values: Tensor, valid: Tensor) -> tuple[Tensor, Tensor]:
379
+ mask = valid[None, None, :, :, None]
380
+ maximum = values.masked_fill(~mask, float("-inf")).amax(dim=-2)
381
+ count = mask.sum(dim=-2).clamp_min(1).to(values.dtype)
382
+ mean = values.masked_fill(~mask, 0.0).sum(dim=-2) / count
383
+ return maximum, mean
384
+
385
+ def _cellular_step(
386
+ self, qf: Tensor, kf: Tensor, state: Tensor, step: int,
387
+ ) -> Tensor:
388
+ length = state.shape[2]
389
+ index, valid = self._valid_mask(length, step, state.device)
390
+ keys = kf[:, :, index, :]
391
+ values = state[:, :, index, :]
392
+
393
+ dynamic_q = torch.einsum("bhtd,hfd->bhtf", state, self.state_q)
394
+ mix = self.state_q_gate[step].sigmoid()[None, :, None, None]
395
+ query = F.normalize(qf + mix * dynamic_q, dim=-1)
396
+
397
+ scores = torch.einsum("bhtf,bhtwf->bhtw", query, keys)
398
+ scores = scores / math.sqrt(self.feature_dim)
399
+ scores = scores * self.log_temperature[step].clamp(-3, 3).exp()[None, :, None, None]
400
+ width = index.shape[1]
401
+ scores = scores + self.relative_bias[step, :, :width][None, :, None, :]
402
+ scores = scores.masked_fill(~valid[None, None, :, :], float("-inf"))
403
+
404
+ if self.uses_adaptive_routing():
405
+ scores = scores + self._adaptive_route_log_prior(query, valid, step, width)
406
+
407
+ weights = scores.softmax(dim=-1)
408
+ attention_message = torch.einsum("bhtw,bhtwd->bhtd", weights, values)
409
+ message = attention_message
410
+
411
+ if self.uses_mixedpool_branch():
412
+ assert self.mixed_pool_logits is not None
413
+ assert self.log_mixed_pool_gain is not None
414
+ maximum, mean = self._pool_messages(values, valid)
415
+ gains = self.log_mixed_pool_gain[step].clamp(-3, 3).exp()
416
+ maximum = maximum * gains[:, 0][None, :, None, None]
417
+ mean = mean * gains[:, 1][None, :, None, None]
418
+ mixture = self.mixed_pool_logits[step].softmax(dim=-1)
419
+ message = (
420
+ mixture[:, 0][None, :, None, None] * attention_message
421
+ + mixture[:, 1][None, :, None, None] * maximum
422
+ + mixture[:, 2][None, :, None, None] * mean
423
+ )
424
+ elif self.uses_maxpool_branch():
425
+ assert self.pool_mix_logit is not None
426
+ assert self.log_pool_gain is not None
427
+ maximum, _ = self._pool_messages(values, valid)
428
+ gain = self.log_pool_gain[step].clamp(-3, 3).exp()[None, :, None, None]
429
+ pooled = maximum * gain
430
+ pool_mix = self.pool_mix_logit[step].sigmoid()[None, :, None, None]
431
+ message = message + pool_mix * (pooled - message)
432
+
433
+ gate = self.step_gate[step].sigmoid()[None, :, None, None]
434
+ return state + gate * (message - state)
435
+
436
+ def _causal_stride2_pool(self, x: Tensor, level: int) -> Tensor:
437
+ assert self.unet_pool_logits is not None
438
+ length = x.shape[2]
439
+ endpoints = torch.arange(0, length, 2, device=x.device)
440
+ previous = (endpoints - 1).clamp_min(0)
441
+ pair = torch.stack((x[:, :, previous, :], x[:, :, endpoints, :]), dim=-2)
442
+ maximum = pair.amax(dim=-2)
443
+ mean = pair.mean(dim=-2)
444
+ mix = self.unet_pool_logits[level].softmax(dim=-1)
445
+ return (
446
+ mix[:, 0][None, :, None, None] * maximum
447
+ + mix[:, 1][None, :, None, None] * mean
448
+ )
449
+
450
+ @staticmethod
451
+ def _causal_upsample(x: Tensor, target_length: int) -> Tensor:
452
+ # Reduced element j represents an endpoint <= 2*j. floor(t/2) is
453
+ # therefore always current/past relative to target token t.
454
+ index = torch.arange(target_length, device=x.device) // 2
455
+ return x[:, :, index.clamp_max(x.shape[2] - 1), :]
456
+
457
+ def _unet_refine(self, state: Tensor) -> Tensor:
458
+ assert self.unet_encoder is not None
459
+ assert self.unet_decoder is not None
460
+ assert self.unet_swiglu_gate is not None
461
+ assert self.unet_swiglu_value is not None
462
+ assert self.unet_skip1_gate is not None
463
+ assert self.unet_skip0_gate is not None
464
+ assert self.unet_output_gate is not None
465
+
466
+ e0 = state
467
+ e1 = self._causal_stride2_pool(e0, 0)
468
+ e2 = self._causal_stride2_pool(e1, 1)
469
+ mu = torch.einsum("bhtd,hld->bhtl", e2, self.unet_encoder)
470
+
471
+ if self.uses_variational_latent():
472
+ assert self.var_logvar is not None
473
+ assert self.var_logvar_bias is not None
474
+ logvar = torch.einsum("bhtd,hld->bhtl", e2, self.var_logvar)
475
+ logvar = (logvar + self.var_logvar_bias[None, :, None, :]).clamp(-8, 4)
476
+ # Deterministic mean path keeps evaluation reproducible; KL still
477
+ # regularizes a variational latent family during fitting.
478
+ self._last_aux_loss = self.var_kl_weight * 0.5 * (
479
+ mu.square() + logvar.exp() - 1.0 - logvar
480
+ ).mean()
481
+ else:
482
+ self._last_aux_loss = None
483
+
484
+ gate_part = torch.einsum("bhtl,hlm->bhtm", mu, self.unet_swiglu_gate)
485
+ value_part = torch.einsum("bhtl,hlm->bhtm", mu, self.unet_swiglu_value)
486
+ latent = F.silu(gate_part) * value_part
487
+ decoded = torch.einsum("bhtl,hdl->bhtd", latent, self.unet_decoder)
488
+
489
+ up1 = self._causal_upsample(decoded, e1.shape[2])
490
+ g1 = self.unet_skip1_gate.sigmoid()[None, :, None, None]
491
+ d1 = e1 + g1 * (up1 - e1)
492
+ up0 = self._causal_upsample(d1, e0.shape[2])
493
+ g0 = self.unet_skip0_gate.sigmoid()[None, :, None, None]
494
+ d0 = e0 + g0 * (up0 - e0)
495
+ gout = self.unet_output_gate.sigmoid()[None, :, None, None]
496
+ return state + gout * (d0 - state)
497
+
498
+ def _channel_gate(self, state: Tensor) -> Tensor:
499
+ assert self.channel_down is not None
500
+ assert self.channel_up is not None
501
+ assert self.channel_strength is not None
502
+ time = torch.arange(1, state.shape[2] + 1, device=state.device, dtype=state.dtype)
503
+ prefix_mean = state.cumsum(dim=2) / time[None, None, :, None]
504
+ hidden = F.silu(torch.einsum("bhtd,hrd->bhtr", prefix_mean, self.channel_down))
505
+ logits = torch.einsum("bhtr,hdr->bhtd", hidden, self.channel_up)
506
+ strength = self.channel_strength.sigmoid()[None, :, None, None]
507
+ return state * (1.0 + strength * torch.tanh(logits))
508
+
509
+ def auxiliary_loss(self) -> Tensor:
510
+ if self._last_aux_loss is None:
511
+ return self.wq.sum() * 0.0
512
+ return self._last_aux_loss
513
+
514
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
515
+ if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
516
+ raise ValueError("expected Q/K/V as [batch, heads, time, dim]")
517
+ if k.shape != v.shape:
518
+ raise ValueError("K and V must have the same shape")
519
+ if q.shape[0] != k.shape[0] or q.shape[2:] != k.shape[2:]:
520
+ raise ValueError("Q/K/V batch, time and head_dim must match")
521
+ if q.shape[1] != self.num_heads or k.shape[1] != self.num_kv_heads:
522
+ raise ValueError("Q/K head counts do not match layer configuration")
523
+ if q.shape[-1] != self.head_dim or q.shape[2] < 1:
524
+ raise ValueError("head_dim mismatch or empty sequence")
525
+
526
+ self._last_aux_loss = None
527
+ q, k, v = (x.to(self.wq.dtype) for x in (q, k, v))
528
+ qf, kf = self._features(q, k)
529
+ base = v.repeat_interleave(self.groups, dim=1)
530
+ state = base
531
+
532
+ if self.uses_rms_refinement():
533
+ assert self.pre_rms_weight is not None and self.pre_rms_gate is not None
534
+ normed = self._rms(state, self.pre_rms_weight)
535
+ gate = self.pre_rms_gate.sigmoid()[None, :, None, None]
536
+ state = state + gate * (normed - state)
537
+
538
+ for step in range(len(self.dilations)):
539
+ state = self._cellular_step(qf, kf, state, step)
540
+
541
+ if self.uses_unet_refinement():
542
+ state = self._unet_refine(state)
543
+ if self.uses_channel_gate():
544
+ state = self._channel_gate(state)
545
+
546
+ if self.uses_rms_refinement():
547
+ assert self.post_rms_weight is not None and self.post_rms_gate is not None
548
+ normed = self._rms(state, self.post_rms_weight)
549
+ gate = self.post_rms_gate.sigmoid()[None, :, None, None]
550
+ state = state + gate * (normed - state)
551
+
552
+ projected = torch.einsum("bhtd,hde->bhte", state, self.out_proj)
553
+ final_gate = self.final_gate.sigmoid()[None, :, None, None]
554
+ output = base + final_gate * (projected - base)
555
+ return output * self.log_gain.clamp(-4, 4).exp()[None, :, None, None]
556
+
557
+ def receptive_field_tokens(self) -> int:
558
+ reach = 0
559
+ for step in range(len(self.dilations)):
560
+ reach += max(self._offsets(step))
561
+ return reach + 1
562
+
563
+ def max_score_pairs(self, context: int) -> int:
564
+ total = 0
565
+ device = self.wq.device
566
+ for step in range(len(self.dilations)):
567
+ _, valid = self._valid_mask(context, step, device)
568
+ total += int(valid.sum().item())
569
+ return total
570
+
571
+ def max_neighbors_per_step(self) -> int:
572
+ return max(len(self._offsets(step)) for step in range(len(self.dilations)))
src/tinycenn_lm/cenn.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import asdict, dataclass
4
+
5
+ import torch
6
+ from torch import Tensor, nn
7
+ import torch.nn.functional as F
8
+
9
+
10
+ @dataclass(frozen=True)
11
+ class CeNNConfig:
12
+ """Configuration for the causal 1-D Cellular Neural Network adapter."""
13
+
14
+ hidden_size: int = 192
15
+ kernel_size: int = 3
16
+ expansion: int = 4
17
+ steps: int = 4
18
+ dilations: tuple[int, ...] = (1, 2, 4, 8)
19
+ rms_norm_eps: float = 1e-5
20
+ dropout: float = 0.0
21
+
22
+ def validate(self) -> None:
23
+ if self.hidden_size <= 0:
24
+ raise ValueError("hidden_size must be positive")
25
+ if self.kernel_size < 2:
26
+ raise ValueError("kernel_size must be >= 2")
27
+ if self.expansion <= 0:
28
+ raise ValueError("expansion must be positive")
29
+ if self.steps <= 0:
30
+ raise ValueError("steps must be positive")
31
+ if not self.dilations or any(d <= 0 for d in self.dilations):
32
+ raise ValueError("dilations must contain positive integers")
33
+ if not 0.0 <= self.dropout < 1.0:
34
+ raise ValueError("dropout must be in [0, 1)")
35
+
36
+ def to_dict(self) -> dict:
37
+ data = asdict(self)
38
+ data["dilations"] = list(self.dilations)
39
+ return data
40
+
41
+ @classmethod
42
+ def from_dict(cls, data: dict) -> "CeNNConfig":
43
+ data = dict(data)
44
+ if "dilations" in data:
45
+ data["dilations"] = tuple(data["dilations"])
46
+ return cls(**data)
47
+
48
+
49
+ class StableRMSNorm(nn.Module):
50
+ """Small RMSNorm with fp32 variance accumulation for training stability."""
51
+
52
+ def __init__(self, hidden_size: int, eps: float = 1e-5) -> None:
53
+ super().__init__()
54
+ self.weight = nn.Parameter(torch.ones(hidden_size))
55
+ self.eps = eps
56
+
57
+ def forward(self, x: Tensor) -> Tensor:
58
+ dtype = x.dtype
59
+ variance = x.float().pow(2).mean(dim=-1, keepdim=True)
60
+ x = x * torch.rsqrt(variance + self.eps).to(dtype=dtype)
61
+ return x * self.weight.to(dtype=dtype)
62
+
63
+
64
+ class CausalDepthwiseNeighborhood(nn.Module):
65
+ """A causal, channel-wise cellular neighborhood operator.
66
+
67
+ The same kernel is reused at every recurrent step. Only the dilation changes,
68
+ which grows the receptive field while keeping the parameter count fixed.
69
+ """
70
+
71
+ def __init__(self, hidden_size: int, kernel_size: int) -> None:
72
+ super().__init__()
73
+ self.hidden_size = hidden_size
74
+ self.kernel_size = kernel_size
75
+ self.weight = nn.Parameter(torch.empty(hidden_size, 1, kernel_size))
76
+ nn.init.kaiming_uniform_(self.weight, a=5**0.5)
77
+
78
+ def forward(self, x: Tensor, dilation: int) -> Tensor:
79
+ if x.ndim != 3:
80
+ raise ValueError(f"expected [batch, seq, hidden], got {tuple(x.shape)}")
81
+ if x.shape[-1] != self.hidden_size:
82
+ raise ValueError(
83
+ f"expected hidden size {self.hidden_size}, got {x.shape[-1]}"
84
+ )
85
+ left_pad = dilation * (self.kernel_size - 1)
86
+ y = x.transpose(1, 2)
87
+ y = F.pad(y, (left_pad, 0))
88
+ y = F.conv1d(
89
+ y,
90
+ self.weight,
91
+ bias=None,
92
+ stride=1,
93
+ padding=0,
94
+ dilation=dilation,
95
+ groups=self.hidden_size,
96
+ )
97
+ return y.transpose(1, 2)
98
+
99
+
100
+ class SharedCeNNCell(nn.Module):
101
+ """Shared recurrent CeNN cell using causal local mixing + gated SwiGLU update."""
102
+
103
+ def __init__(self, config: CeNNConfig) -> None:
104
+ super().__init__()
105
+ config.validate()
106
+ self.config = config
107
+ inner = config.hidden_size * config.expansion
108
+
109
+ self.norm = StableRMSNorm(config.hidden_size, config.rms_norm_eps)
110
+ self.neighborhood = CausalDepthwiseNeighborhood(
111
+ config.hidden_size, config.kernel_size
112
+ )
113
+ self.in_proj = nn.Linear(config.hidden_size, inner * 2, bias=False)
114
+ self.gate_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=True)
115
+ self.out_proj = nn.Linear(inner, config.hidden_size, bias=False)
116
+ self.dropout = nn.Dropout(config.dropout)
117
+
118
+ nn.init.zeros_(self.out_proj.weight)
119
+ nn.init.constant_(self.gate_proj.bias, -1.0)
120
+
121
+ def forward(self, state: Tensor, dilation: int, step_scale: float) -> Tensor:
122
+ x = self.norm(state)
123
+ local = self.neighborhood(x, dilation=dilation)
124
+ a, b = self.in_proj(local).chunk(2, dim=-1)
125
+ update = F.silu(a) * b
126
+ update = self.out_proj(update)
127
+ update = self.dropout(update)
128
+ gate = torch.sigmoid(self.gate_proj(local))
129
+ return state + (step_scale * gate * update)
130
+
131
+
132
+ class FastCeNNCore(nn.Module):
133
+ """Iterates one shared CeNN cell multiple times with a dilation schedule."""
134
+
135
+ def __init__(self, config: CeNNConfig) -> None:
136
+ super().__init__()
137
+ config.validate()
138
+ self.config = config
139
+ self.cell = SharedCeNNCell(config)
140
+
141
+ def forward(self, hidden_states: Tensor) -> Tensor:
142
+ initial = hidden_states
143
+ state = hidden_states
144
+ step_scale = self.config.steps ** -0.5
145
+ for step in range(self.config.steps):
146
+ dilation = self.config.dilations[step % len(self.config.dilations)]
147
+ state = self.cell(state, dilation=dilation, step_scale=step_scale)
148
+ return state - initial
149
+
150
+ @property
151
+ def receptive_field(self) -> int:
152
+ radius = sum(
153
+ self.config.dilations[i % len(self.config.dilations)]
154
+ for i in range(self.config.steps)
155
+ )
156
+ return 1 + (self.config.kernel_size - 1) * radius
src/tinycenn_lm/colab_live_backup.py ADDED
@@ -0,0 +1,459 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import os
5
+ import shlex
6
+ import subprocess
7
+ import sys
8
+ import time
9
+ from datetime import datetime, timezone
10
+ from pathlib import Path
11
+ from typing import Any, Callable
12
+
13
+ from .hf_persistence import redact_secrets, utc_run_id
14
+
15
+ _SMALL_SUFFIXES = {".json", ".jsonl", ".csv", ".txt", ".md", ".log", ".yaml", ".yml"}
16
+
17
+
18
+ def is_tinycenn_training_command(cmd: Any) -> bool:
19
+ if not isinstance(cmd, (list, tuple)):
20
+ return False
21
+ parts = [str(x) for x in cmd]
22
+ for part in parts:
23
+ name = Path(part).name.lower()
24
+ if name.startswith("train_") and name.endswith(".py"):
25
+ return True
26
+ return False
27
+
28
+
29
+ def training_script_name(cmd: list[str] | tuple[str, ...]) -> str:
30
+ for part in cmd:
31
+ name = Path(str(part)).name
32
+ if name.lower().startswith("train_") and name.lower().endswith(".py"):
33
+ return name[:-3]
34
+ return "training"
35
+
36
+
37
+ def output_dir_from_command(cmd: list[str] | tuple[str, ...], cwd: str | Path | None = None) -> Path | None:
38
+ parts = [str(x) for x in cmd]
39
+ for i, part in enumerate(parts):
40
+ if part == "--output-dir" and i + 1 < len(parts):
41
+ path = Path(parts[i + 1])
42
+ if not path.is_absolute() and cwd is not None:
43
+ path = Path(cwd) / path
44
+ return path
45
+ if part.startswith("--output-dir="):
46
+ path = Path(part.split("=", 1)[1])
47
+ if not path.is_absolute() and cwd is not None:
48
+ path = Path(cwd) / path
49
+ return path
50
+ return None
51
+
52
+
53
+ def _small_files(folder: Path) -> list[Path]:
54
+ if not folder.exists():
55
+ return []
56
+ return [
57
+ p for p in folder.rglob("*")
58
+ if p.is_file()
59
+ and p.suffix.lower() in _SMALL_SUFFIXES
60
+ and p.stat().st_size <= 10 * 1024 * 1024
61
+ and ".hf_run_archive" not in p.parts
62
+ and ".hf_live_redacted" not in p.parts
63
+ ]
64
+
65
+
66
+ def _retry_required(
67
+ action: Callable[[], Any],
68
+ *,
69
+ label: str,
70
+ attempts: int = 3,
71
+ delay_seconds: float = 5.0,
72
+ ) -> Any:
73
+ """Run one mandatory Hugging Face operation with bounded retries."""
74
+ last_error: Exception | None = None
75
+ for attempt in range(1, max(attempts, 1) + 1):
76
+ try:
77
+ return action()
78
+ except Exception as exc: # network/API failures must fail closed after retries
79
+ last_error = exc
80
+ print(
81
+ f"[TinyCeNN][BACKUP] {label} failed "
82
+ f"(attempt {attempt}/{attempts}): {exc}",
83
+ flush=True,
84
+ )
85
+ if attempt < attempts:
86
+ time.sleep(delay_seconds)
87
+ raise RuntimeError(
88
+ f"Mandatory Hugging Face backup failed during {label} after {attempts} attempts"
89
+ ) from last_error
90
+
91
+
92
+ def _upload_live_required(
93
+ api,
94
+ *,
95
+ repo_id: str,
96
+ run_id: str,
97
+ output_dir: Path | None,
98
+ log_file: Path,
99
+ include_checkpoint: bool = True,
100
+ ) -> None:
101
+ """Persist the live log and reports, optionally mirroring checkpoint files."""
102
+ if log_file.exists():
103
+ clean_log = log_file.with_name("train_redacted.log")
104
+ clean_log.write_text(
105
+ redact_secrets(log_file.read_text(encoding="utf-8", errors="replace")),
106
+ encoding="utf-8",
107
+ )
108
+ _retry_required(
109
+ lambda: api.upload_file(
110
+ repo_id=repo_id,
111
+ repo_type="model",
112
+ path_or_fileobj=str(clean_log),
113
+ path_in_repo=f"runs/{run_id}/train.log",
114
+ commit_message=f"Live backup {run_id}",
115
+ ),
116
+ label="live log upload",
117
+ )
118
+
119
+ if output_dir is None or not output_dir.exists():
120
+ return
121
+
122
+ # Upload small human-readable artifacts separately after redaction.
123
+ for path in _small_files(output_dir):
124
+ rel = path.relative_to(output_dir)
125
+ upload_path = path
126
+ try:
127
+ clean_dir = output_dir / ".hf_live_redacted"
128
+ dst = clean_dir / rel
129
+ dst.parent.mkdir(parents=True, exist_ok=True)
130
+ dst.write_text(
131
+ redact_secrets(path.read_text(encoding="utf-8", errors="replace")),
132
+ encoding="utf-8",
133
+ )
134
+ upload_path = dst
135
+ except Exception:
136
+ upload_path = path
137
+
138
+ _retry_required(
139
+ lambda p=upload_path, r=rel: api.upload_file(
140
+ repo_id=repo_id,
141
+ repo_type="model",
142
+ path_or_fileobj=str(p),
143
+ path_in_repo=f"runs/{run_id}/artifacts/{r.as_posix()}",
144
+ commit_message=f"Live results {run_id}",
145
+ ),
146
+ label=f"artifact upload: {rel.as_posix()}",
147
+ )
148
+
149
+ if not include_checkpoint:
150
+ return
151
+
152
+ # Mandatory periodic checkpoint mirror. This is disabled when the output
153
+ # directory already existed before the run, because otherwise stale weights from
154
+ # an earlier execution can trigger a large upload unrelated to the current run.
155
+ checkpoint_files = [
156
+ p for p in output_dir.rglob("*")
157
+ if p.is_file()
158
+ and ".hf_run_archive" not in p.parts
159
+ and ".hf_live_redacted" not in p.parts
160
+ ]
161
+ if checkpoint_files:
162
+ _retry_required(
163
+ lambda: api.upload_folder(
164
+ repo_id=repo_id,
165
+ repo_type="model",
166
+ folder_path=str(output_dir),
167
+ path_in_repo=f"runs/{run_id}/checkpoint",
168
+ ignore_patterns=[".hf_run_archive/**", ".hf_live_redacted/**"],
169
+ commit_message=f"Live checkpoint backup {run_id}",
170
+ ),
171
+ label="live checkpoint mirror",
172
+ )
173
+
174
+
175
+ def _backup_metadata(cmd, output_dir: Path | None, run_id: str) -> dict[str, Any]:
176
+ return {
177
+ "run_id": run_id,
178
+ "created_utc": datetime.now(timezone.utc).isoformat(),
179
+ "command": [redact_secrets(str(x)) for x in cmd],
180
+ "output_dir": str(output_dir) if output_dir else None,
181
+ "status": "running",
182
+ "backup_policy": "mandatory-fail-closed",
183
+ }
184
+
185
+
186
+ def _format_elapsed(seconds: float) -> str:
187
+ seconds = max(0, int(seconds))
188
+ hours, remainder = divmod(seconds, 3600)
189
+ minutes, secs = divmod(remainder, 60)
190
+ if hours:
191
+ return f"{hours:d}h {minutes:02d}m {secs:02d}s"
192
+ return f"{minutes:d}m {secs:02d}s"
193
+
194
+
195
+ def install_colab_training_backup(*, interval_seconds: int = 180) -> bool:
196
+ """Stream TinyCeNN Colab training and require a private Hugging Face backup.
197
+
198
+ Every ``subprocess.run([... train_*.py ...])`` call in Colab is converted to a
199
+ line-streaming ``Popen`` execution. Before the child process starts, a valid
200
+ Hugging Face login and a writable private ``TinyCeNN-LM-Colab-Backups`` repo are
201
+ required. Live logs and reports are mirrored periodically, along with checkpoints
202
+ when the output directory is new for the current run. If a mandatory sync fails
203
+ after retries, training is terminated. A successful run must finish by uploading
204
+ its complete checkpoint. A failed trainer still uploads its log/status promptly,
205
+ but does not block on a large or stale checkpoint directory.
206
+ """
207
+ if not (os.environ.get("COLAB_RELEASE_TAG") or os.environ.get("COLAB_GPU") or Path("/content").exists()):
208
+ return False
209
+ if getattr(subprocess.run, "_tinycenn_live_backup", False):
210
+ return True
211
+
212
+ original_run = subprocess.run
213
+ original_popen = subprocess.Popen
214
+
215
+ def run_with_backup(cmd, *args, **kwargs):
216
+ if not is_tinycenn_training_command(cmd):
217
+ return original_run(cmd, *args, **kwargs)
218
+
219
+ unsupported = {"input", "capture_output", "stdout", "stderr", "timeout"} & set(kwargs)
220
+ if unsupported or args:
221
+ return original_run(cmd, *args, **kwargs)
222
+
223
+ cwd = kwargs.pop("cwd", None)
224
+ requested_env = kwargs.pop("env", None)
225
+ check = bool(kwargs.pop("check", False))
226
+ if kwargs:
227
+ return original_run(cmd, cwd=cwd, env=requested_env, check=check, **kwargs)
228
+
229
+ child_env = os.environ.copy()
230
+ if requested_env is not None:
231
+ child_env.update({str(k): str(v) for k, v in requested_env.items()})
232
+ child_env["PYTHONUNBUFFERED"] = "1"
233
+
234
+ script = training_script_name(cmd)
235
+ run_id = utc_run_id(script[:32])
236
+ output_dir = output_dir_from_command(cmd, cwd=cwd)
237
+ output_preexisting = bool(output_dir is not None and output_dir.exists())
238
+ work_root = Path(cwd).resolve() if cwd is not None else Path.cwd().resolve()
239
+ live_root = work_root / ".colab_live_backup" / run_id
240
+ live_root.mkdir(parents=True, exist_ok=True)
241
+ log_file = live_root / "train.log"
242
+ meta_file = live_root / "run_status.json"
243
+ metadata = _backup_metadata(cmd, output_dir, run_id)
244
+ metadata["output_dir_preexisting"] = output_preexisting
245
+ meta_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
246
+
247
+ # Mandatory preflight. Do not spend GPU time unless the remote safety target
248
+ # is authenticated, private, writable, and has accepted the run metadata.
249
+ try:
250
+ from huggingface_hub import HfApi, get_token
251
+ except Exception as exc:
252
+ raise RuntimeError(
253
+ "Mandatory Hugging Face backup requires huggingface_hub. "
254
+ "Install it before training."
255
+ ) from exc
256
+
257
+ token = get_token() or os.environ.get("HF_TOKEN")
258
+ if not token:
259
+ raise RuntimeError(
260
+ "Mandatory Hugging Face backup is enabled. Log in first or provide "
261
+ "HF_TOKEN (in Colab: add HF_TOKEN to Secrets and run the login cell)."
262
+ )
263
+
264
+ api = HfApi(token=token)
265
+ identity = _retry_required(api.whoami, label="Hugging Face authentication")
266
+ user = identity["name"]
267
+ backup_repo = f"{user}/TinyCeNN-LM-Colab-Backups"
268
+ _retry_required(
269
+ lambda: api.create_repo(
270
+ backup_repo,
271
+ repo_type="model",
272
+ private=True,
273
+ exist_ok=True,
274
+ ),
275
+ label="private backup repository preflight",
276
+ )
277
+ _retry_required(
278
+ lambda: api.upload_file(
279
+ repo_id=backup_repo,
280
+ repo_type="model",
281
+ path_or_fileobj=str(meta_file),
282
+ path_in_repo=f"runs/{run_id}/run_status.json",
283
+ commit_message=f"Start mandatory backup {run_id}",
284
+ ),
285
+ label="initial backup metadata upload",
286
+ )
287
+
288
+ command_text = " ".join(shlex.quote(str(x)) for x in cmd)
289
+ print("\n" + "=" * 88, flush=True)
290
+ print(f"[TinyCeNN][START] {script}", flush=True)
291
+ print(f"[TinyCeNN][COMMAND] {command_text}", flush=True)
292
+ if output_dir is not None:
293
+ print(f"[TinyCeNN][OUTPUT] {output_dir}", flush=True)
294
+ if output_preexisting:
295
+ print(
296
+ "[TinyCeNN][BACKUP] output directory already exists; periodic backup "
297
+ "will save logs/reports only and avoid re-uploading stale checkpoint files.",
298
+ flush=True,
299
+ )
300
+ print(f"[TinyCeNN][LOCAL LOG] {log_file}", flush=True)
301
+ print(
302
+ f"[TinyCeNN][BACKUP REQUIRED] "
303
+ f"https://huggingface.co/{backup_repo}/tree/main/runs/{run_id}",
304
+ flush=True,
305
+ )
306
+ print(
307
+ f"[TinyCeNN][BACKUP POLICY] mandatory sync every {interval_seconds}s; "
308
+ "training aborts if backup cannot be persisted after retries",
309
+ flush=True,
310
+ )
311
+ print("[TinyCeNN][LIVE] streaming training output...", flush=True)
312
+ print("-" * 88, flush=True)
313
+
314
+ started = time.monotonic()
315
+ proc = original_popen(
316
+ cmd,
317
+ cwd=cwd,
318
+ env=child_env,
319
+ stdout=subprocess.PIPE,
320
+ stderr=subprocess.STDOUT,
321
+ text=True,
322
+ bufsize=1,
323
+ )
324
+ last_sync = started
325
+ backup_failure: Exception | None = None
326
+
327
+ with log_file.open("a", encoding="utf-8") as log:
328
+ assert proc.stdout is not None
329
+ for line in proc.stdout:
330
+ sys.stdout.write(line)
331
+ sys.stdout.flush()
332
+ log.write(line)
333
+ log.flush()
334
+ now = time.monotonic()
335
+ if now - last_sync >= interval_seconds:
336
+ try:
337
+ _upload_live_required(
338
+ api,
339
+ repo_id=backup_repo,
340
+ run_id=run_id,
341
+ output_dir=output_dir,
342
+ log_file=log_file,
343
+ include_checkpoint=not output_preexisting,
344
+ )
345
+ print(
346
+ f"[TinyCeNN][BACKUP OK] live state persisted at "
347
+ f"{_format_elapsed(now - started)}",
348
+ flush=True,
349
+ )
350
+ last_sync = now
351
+ except Exception as exc:
352
+ backup_failure = exc
353
+ print(
354
+ f"[TinyCeNN][BACKUP FATAL] {exc}. Terminating training to protect the run.",
355
+ flush=True,
356
+ )
357
+ proc.terminate()
358
+ try:
359
+ proc.wait(timeout=30)
360
+ except subprocess.TimeoutExpired:
361
+ proc.kill()
362
+ proc.wait()
363
+ break
364
+
365
+ returncode = proc.wait()
366
+ elapsed = time.monotonic() - started
367
+
368
+ if backup_failure is not None:
369
+ metadata["status"] = "failed-backup"
370
+ metadata["returncode"] = returncode
371
+ metadata["elapsed_seconds"] = elapsed
372
+ metadata["finished_utc"] = datetime.now(timezone.utc).isoformat()
373
+ metadata["backup_error"] = redact_secrets(str(backup_failure))
374
+ meta_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
375
+ # Best effort only for the failure marker; the triggering sync already
376
+ # proved that the remote is unavailable.
377
+ try:
378
+ api.upload_file(
379
+ repo_id=backup_repo,
380
+ repo_type="model",
381
+ path_or_fileobj=str(meta_file),
382
+ path_in_repo=f"runs/{run_id}/run_status.json",
383
+ commit_message=f"Backup failure {run_id}",
384
+ )
385
+ except Exception:
386
+ pass
387
+ raise RuntimeError(
388
+ "Training was terminated because mandatory Hugging Face backup could not be maintained."
389
+ ) from backup_failure
390
+
391
+ metadata["status"] = "completed" if returncode == 0 else "failed-training"
392
+ metadata["returncode"] = returncode
393
+ metadata["elapsed_seconds"] = elapsed
394
+ metadata["finished_utc"] = datetime.now(timezone.utc).isoformat()
395
+ meta_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
396
+
397
+ # Always preserve the run's log, small reports and final status. On failure,
398
+ # do not upload the checkpoint directory: it may be a stale fixed-name folder
399
+ # from an earlier run and was the source of long post-crash Colab hangs.
400
+ _upload_live_required(
401
+ api,
402
+ repo_id=backup_repo,
403
+ run_id=run_id,
404
+ output_dir=output_dir,
405
+ log_file=log_file,
406
+ include_checkpoint=False,
407
+ )
408
+ _retry_required(
409
+ lambda: api.upload_file(
410
+ repo_id=backup_repo,
411
+ repo_type="model",
412
+ path_or_fileobj=str(meta_file),
413
+ path_in_repo=f"runs/{run_id}/run_status.json",
414
+ commit_message=f"Finish mandatory backup {run_id}",
415
+ ),
416
+ label="final status upload",
417
+ )
418
+
419
+ if returncode == 0 and output_dir is not None and output_dir.exists():
420
+ print("[TinyCeNN][BACKUP] uploading successful final checkpoint...", flush=True)
421
+ _retry_required(
422
+ lambda: api.upload_folder(
423
+ repo_id=backup_repo,
424
+ repo_type="model",
425
+ folder_path=str(output_dir),
426
+ path_in_repo=f"runs/{run_id}/checkpoint",
427
+ ignore_patterns=[".hf_run_archive/**", ".hf_live_redacted/**"],
428
+ commit_message=f"Final checkpoint backup {run_id}",
429
+ ),
430
+ label="final checkpoint upload",
431
+ )
432
+ elif returncode != 0:
433
+ print(
434
+ "[TinyCeNN][BACKUP] trainer failed; log/status backed up, large checkpoint upload skipped.",
435
+ flush=True,
436
+ )
437
+
438
+ print("-" * 88, flush=True)
439
+ print(
440
+ f"[TinyCeNN][BACKUP COMPLETE] private Hugging Face backup committed for {run_id}",
441
+ flush=True,
442
+ )
443
+ if returncode == 0:
444
+ print(f"[TinyCeNN][DONE] {script} completed in {_format_elapsed(elapsed)}", flush=True)
445
+ else:
446
+ print(
447
+ f"[TinyCeNN][FAILED] {script} exited with code {returncode} after {_format_elapsed(elapsed)}",
448
+ flush=True,
449
+ )
450
+ print("=" * 88 + "\n", flush=True)
451
+
452
+ completed = subprocess.CompletedProcess(cmd, returncode)
453
+ if check and returncode:
454
+ raise subprocess.CalledProcessError(returncode, cmd)
455
+ return completed
456
+
457
+ run_with_backup._tinycenn_live_backup = True
458
+ subprocess.run = run_with_backup
459
+ return True
src/tinycenn_lm/direct_colab_backup.py ADDED
@@ -0,0 +1,276 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import atexit
4
+ import json
5
+ import os
6
+ import sys
7
+ import threading
8
+ import time
9
+ from datetime import datetime, timezone
10
+ from pathlib import Path
11
+ from typing import Any
12
+
13
+ from .colab_live_backup import _retry_required, _upload_live_required, output_dir_from_command
14
+ from .hf_persistence import redact_secrets, utc_run_id
15
+
16
+ _INSTALLED = False
17
+
18
+
19
+ def _is_colab() -> bool:
20
+ return bool(os.environ.get("COLAB_RELEASE_TAG") or os.environ.get("COLAB_GPU") or Path("/content").exists())
21
+
22
+
23
+ def _training_script_name() -> str | None:
24
+ name = Path(sys.argv[0]).name
25
+ if name.lower().startswith("train_") and name.lower().endswith(".py"):
26
+ return name[:-3]
27
+ return None
28
+
29
+
30
+ def _candidate_parent_roots(root: Path) -> list[Path]:
31
+ """Return likely repository roots used by notebook-side backup wrappers."""
32
+ candidates = [root.resolve()]
33
+
34
+ # A trainer is normally /content/TinyCeNN-LM/scripts/train_*.py. Deriving the
35
+ # repository from argv makes detection independent of the notebook's cwd.
36
+ try:
37
+ script_path = Path(sys.argv[0]).resolve()
38
+ if script_path.parent.name == "scripts":
39
+ candidates.append(script_path.parent.parent)
40
+ except Exception:
41
+ pass
42
+
43
+ # Older notebook wrappers used this fixed Colab repository path.
44
+ candidates.append(Path("/content/TinyCeNN-LM"))
45
+
46
+ configured = os.environ.get("TINYCENN_PARENT_BACKUP_ROOT")
47
+ if configured:
48
+ candidates.append(Path(configured))
49
+
50
+ unique: list[Path] = []
51
+ seen: set[str] = set()
52
+ for candidate in candidates:
53
+ key = str(candidate.resolve())
54
+ if key not in seen:
55
+ seen.add(key)
56
+ unique.append(candidate.resolve())
57
+ return unique
58
+
59
+
60
+ def _parent_backup_is_active(script: str, root: Path) -> bool:
61
+ """Detect a recent parent-wrapper run so the child does not duplicate uploads."""
62
+ if os.environ.get("TINYCENN_PARENT_BACKUP_ACTIVE") == "1":
63
+ return True
64
+
65
+ now = time.time()
66
+ for candidate in _candidate_parent_roots(root):
67
+ backup_root = candidate / ".colab_live_backup"
68
+ if not backup_root.exists():
69
+ continue
70
+ for meta in backup_root.glob("*/run_status.json"):
71
+ try:
72
+ if now - meta.stat().st_mtime > 10 * 60:
73
+ continue
74
+ data = json.loads(meta.read_text(encoding="utf-8"))
75
+ if data.get("status") != "running":
76
+ continue
77
+ command = " ".join(str(x) for x in data.get("command", []))
78
+ if script in command:
79
+ return True
80
+ except Exception:
81
+ continue
82
+ return False
83
+
84
+
85
+ class _TeeStream:
86
+ def __init__(self, primary, log_file):
87
+ self._primary = primary
88
+ self._log = log_file
89
+
90
+ def write(self, text):
91
+ result = self._primary.write(text)
92
+ self._primary.flush()
93
+ self._log.write(text)
94
+ self._log.flush()
95
+ return result
96
+
97
+ def flush(self):
98
+ self._primary.flush()
99
+ self._log.flush()
100
+
101
+ def __getattr__(self, name):
102
+ return getattr(self._primary, name)
103
+
104
+
105
+ def install_direct_training_backup(*, interval_seconds: int = 180) -> bool:
106
+ """Mandatory trainer-side HF backup when no parent Colab wrapper is active.
107
+
108
+ This is a safety fallback for notebooks that launch ``train_*.py`` before the
109
+ notebook kernel imports ``tinycenn_lm``. Normal notebooks are backed up by the
110
+ parent wrapper; this function detects that active run and stays out of the way.
111
+ """
112
+ global _INSTALLED
113
+ if _INSTALLED or not _is_colab():
114
+ return _INSTALLED
115
+
116
+ script = _training_script_name()
117
+ if script is None:
118
+ return False
119
+
120
+ root = Path.cwd().resolve()
121
+ if _parent_backup_is_active(script, root):
122
+ print("[TinyCeNN][BACKUP] parent mandatory backup detected; child fallback not needed.", flush=True)
123
+ _INSTALLED = True
124
+ return True
125
+
126
+ try:
127
+ from huggingface_hub import HfApi, get_token
128
+ except Exception as exc:
129
+ raise RuntimeError(
130
+ "Mandatory Hugging Face backup requires huggingface_hub. Install it before training."
131
+ ) from exc
132
+
133
+ token = get_token() or os.environ.get("HF_TOKEN")
134
+ if not token:
135
+ raise RuntimeError(
136
+ "Mandatory Hugging Face backup is enabled. Add HF_TOKEN to Colab Secrets, "
137
+ "run the Hugging Face login cell, then start training."
138
+ )
139
+
140
+ api = HfApi(token=token)
141
+ identity = _retry_required(api.whoami, label="trainer-side Hugging Face authentication")
142
+ user = identity["name"]
143
+ repo_id = f"{user}/TinyCeNN-LM-Colab-Backups"
144
+ _retry_required(
145
+ lambda: api.create_repo(repo_id, repo_type="model", private=True, exist_ok=True),
146
+ label="trainer-side private backup repository preflight",
147
+ )
148
+
149
+ run_id = utc_run_id(f"{script}-direct"[:32])
150
+ output_dir = output_dir_from_command(sys.argv, cwd=root)
151
+ output_preexisting = bool(output_dir is not None and output_dir.exists())
152
+ live_root = root / ".colab_live_backup" / run_id
153
+ live_root.mkdir(parents=True, exist_ok=True)
154
+ log_path = live_root / "train.log"
155
+ meta_path = live_root / "run_status.json"
156
+ metadata: dict[str, Any] = {
157
+ "run_id": run_id,
158
+ "created_utc": datetime.now(timezone.utc).isoformat(),
159
+ "command": [redact_secrets(str(x)) for x in sys.argv],
160
+ "output_dir": str(output_dir) if output_dir else None,
161
+ "output_dir_preexisting": output_preexisting,
162
+ "status": "running",
163
+ "backup_policy": "mandatory-fail-closed-direct-trainer-fallback",
164
+ }
165
+ meta_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
166
+ _retry_required(
167
+ lambda: api.upload_file(
168
+ repo_id=repo_id,
169
+ repo_type="model",
170
+ path_or_fileobj=str(meta_path),
171
+ path_in_repo=f"runs/{run_id}/run_status.json",
172
+ commit_message=f"Start mandatory direct backup {run_id}",
173
+ ),
174
+ label="trainer-side initial metadata upload",
175
+ )
176
+
177
+ log_handle = log_path.open("a", encoding="utf-8", buffering=1)
178
+ sys.stdout = _TeeStream(sys.stdout, log_handle)
179
+ sys.stderr = _TeeStream(sys.stderr, log_handle)
180
+
181
+ print(
182
+ f"[TinyCeNN][BACKUP REQUIRED][DIRECT] "
183
+ f"https://huggingface.co/{repo_id}/tree/main/runs/{run_id}",
184
+ flush=True,
185
+ )
186
+ print(
187
+ f"[TinyCeNN][BACKUP POLICY][DIRECT] mandatory sync every {interval_seconds}s; "
188
+ "process exits if backup cannot be persisted after retries",
189
+ flush=True,
190
+ )
191
+ if output_preexisting:
192
+ print(
193
+ "[TinyCeNN][BACKUP][DIRECT] output directory already existed; periodic "
194
+ "checkpoint mirroring is disabled to avoid uploading stale files.",
195
+ flush=True,
196
+ )
197
+
198
+ stop_event = threading.Event()
199
+ started = time.monotonic()
200
+
201
+ def heartbeat() -> None:
202
+ while not stop_event.wait(interval_seconds):
203
+ try:
204
+ _upload_live_required(
205
+ api,
206
+ repo_id=repo_id,
207
+ run_id=run_id,
208
+ output_dir=output_dir,
209
+ log_file=log_path,
210
+ include_checkpoint=not output_preexisting,
211
+ )
212
+ print("[TinyCeNN][BACKUP OK][DIRECT] live state persisted.", flush=True)
213
+ except Exception as exc:
214
+ print(
215
+ f"[TinyCeNN][BACKUP FATAL][DIRECT] {exc}. "
216
+ "Stopping training because remote safety cannot be guaranteed.",
217
+ flush=True,
218
+ )
219
+ os._exit(74)
220
+
221
+ worker = threading.Thread(target=heartbeat, name="tinycenn-hf-backup", daemon=True)
222
+ worker.start()
223
+
224
+ def finalize() -> None:
225
+ stop_event.set()
226
+ metadata["status"] = "process-exit"
227
+ metadata["finished_utc"] = datetime.now(timezone.utc).isoformat()
228
+ metadata["elapsed_seconds"] = time.monotonic() - started
229
+ meta_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
230
+ try:
231
+ # The explicit final upload below is the only full-folder upload here.
232
+ # This avoids sending the same checkpoint twice at process exit.
233
+ _upload_live_required(
234
+ api,
235
+ repo_id=repo_id,
236
+ run_id=run_id,
237
+ output_dir=output_dir,
238
+ log_file=log_path,
239
+ include_checkpoint=False,
240
+ )
241
+ _retry_required(
242
+ lambda: api.upload_file(
243
+ repo_id=repo_id,
244
+ repo_type="model",
245
+ path_or_fileobj=str(meta_path),
246
+ path_in_repo=f"runs/{run_id}/run_status.json",
247
+ commit_message=f"Finish mandatory direct backup {run_id}",
248
+ ),
249
+ label="trainer-side final status upload",
250
+ )
251
+ if output_dir is not None and output_dir.exists():
252
+ _retry_required(
253
+ lambda: api.upload_folder(
254
+ repo_id=repo_id,
255
+ repo_type="model",
256
+ folder_path=str(output_dir),
257
+ path_in_repo=f"runs/{run_id}/checkpoint",
258
+ ignore_patterns=[".hf_run_archive/**", ".hf_live_redacted/**"],
259
+ commit_message=f"Final direct checkpoint backup {run_id}",
260
+ ),
261
+ label="trainer-side final checkpoint upload",
262
+ )
263
+ print("[TinyCeNN][BACKUP COMPLETE][DIRECT] final backup committed.", flush=True)
264
+ except Exception as exc:
265
+ print(
266
+ f"[TinyCeNN][BACKUP FATAL][DIRECT] final backup failed: {exc}",
267
+ flush=True,
268
+ )
269
+ try:
270
+ log_handle.flush()
271
+ finally:
272
+ os._exit(75)
273
+
274
+ atexit.register(finalize)
275
+ _INSTALLED = True
276
+ return True
src/tinycenn_lm/distill_utils.py ADDED
@@ -0,0 +1,232 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import hashlib
4
+ import math
5
+ import random
6
+ import struct
7
+ from collections.abc import Iterable, Iterator
8
+ from contextlib import nullcontext
9
+
10
+ import torch
11
+ import torch.nn.functional as F
12
+
13
+
14
+ def holdout_bucket(text: str, buckets: int = 1000) -> int:
15
+ digest = hashlib.blake2b(text.encode("utf-8", errors="ignore"), digest_size=8).digest()
16
+ return int.from_bytes(digest, "little") % buckets
17
+
18
+
19
+ def partition_rows(dataset: Iterable[dict], text_field: str, *, validation: bool) -> Iterator[dict]:
20
+ """Deterministic document-level split: 99% train / 1% validation."""
21
+ for example in dataset:
22
+ text = example.get(text_field)
23
+ if not isinstance(text, str) or not text.strip():
24
+ continue
25
+ is_validation = holdout_bucket(text) >= 990
26
+ if is_validation == validation:
27
+ yield example
28
+
29
+
30
+ def buffered_shuffle(rows: Iterable[dict], *, buffer_size: int, seed: int) -> Iterator[dict]:
31
+ """Deterministic bounded-memory shuffle for streaming datasets."""
32
+ if buffer_size <= 1:
33
+ yield from rows
34
+ return
35
+
36
+ rng = random.Random(seed)
37
+ buffer: list[dict] = []
38
+ for row in rows:
39
+ if len(buffer) < buffer_size:
40
+ buffer.append(row)
41
+ continue
42
+ index = rng.randrange(len(buffer))
43
+ yield buffer[index]
44
+ buffer[index] = row
45
+
46
+ rng.shuffle(buffer)
47
+ yield from buffer
48
+
49
+
50
+ def token_blocks(
51
+ rows: Iterable[dict], tokenizer, text_field: str, block_size: int, *, skip_tokens: int = 0
52
+ ) -> Iterator[torch.Tensor]:
53
+ """Pack documents, optionally advancing a deterministic stream before packing.
54
+
55
+ Skipping is before tensor allocation and works across document boundaries and
56
+ changes in batch/context length. Reconstructing a cursor still needs reading
57
+ and tokenizing the prefix; it does not run the teacher or student on it.
58
+ """
59
+ if block_size < 2 or skip_tokens < 0:
60
+ raise ValueError("block_size must be >= 2 and skip_tokens must be nonnegative")
61
+ buffer: list[int] = []
62
+ offset = 0
63
+ eos = tokenizer.eos_token_id
64
+ for example in rows:
65
+ ids = tokenizer(example[text_field], add_special_tokens=False)["input_ids"]
66
+ if eos is not None:
67
+ ids.append(eos)
68
+ if skip_tokens:
69
+ skipped = min(skip_tokens, len(ids))
70
+ skip_tokens -= skipped
71
+ ids = ids[skipped:]
72
+ buffer.extend(ids)
73
+ while len(buffer) - offset >= block_size:
74
+ yield torch.tensor(buffer[offset : offset + block_size], dtype=torch.long)
75
+ offset += block_size
76
+ if offset > 1_000_000:
77
+ buffer = buffer[offset:]
78
+ offset = 0
79
+
80
+
81
+ def batch_blocks(blocks: Iterator[torch.Tensor], batch_size: int) -> Iterator[torch.Tensor]:
82
+ batch: list[torch.Tensor] = []
83
+ for block in blocks:
84
+ batch.append(block)
85
+ if len(batch) == batch_size:
86
+ yield torch.stack(batch)
87
+ batch.clear()
88
+
89
+
90
+ def collect_eval_batches(rows, tokenizer, text_field: str, block_size: int, batch_size: int, count: int):
91
+ batches = batch_blocks(token_blocks(rows, tokenizer, text_field, block_size), batch_size)
92
+ out: list[torch.Tensor] = []
93
+ for _ in range(count):
94
+ try:
95
+ out.append(next(batches))
96
+ except StopIteration:
97
+ break
98
+ if not out:
99
+ raise RuntimeError("could not build held-out evaluation batches")
100
+ return out
101
+
102
+
103
+ def evaluation_fingerprint(batches: list[torch.Tensor]) -> str:
104
+ """Stable SHA256 fingerprint of the exact held-out token batches.
105
+
106
+ Tokens are encoded explicitly as little-endian signed int64 values, avoiding
107
+ NumPy and platform-dependent tensor byte representations.
108
+ """
109
+ digest = hashlib.sha256()
110
+ for batch in batches:
111
+ tensor = batch.detach().to(device="cpu", dtype=torch.int64).contiguous().view(-1)
112
+ for token_id in tensor.tolist():
113
+ digest.update(struct.pack("<q", int(token_id)))
114
+ return digest.hexdigest()
115
+
116
+
117
+ def chunked_kl(
118
+ student_logits: torch.Tensor,
119
+ teacher_logits: torch.Tensor,
120
+ temperature: float,
121
+ chunk_rows: int,
122
+ ) -> torch.Tensor:
123
+ if temperature <= 0 or chunk_rows < 1:
124
+ raise ValueError("temperature and chunk_rows must be positive")
125
+ s = student_logits.reshape(-1, student_logits.shape[-1])
126
+ t = teacher_logits.reshape(-1, teacher_logits.shape[-1])
127
+ total = s.new_zeros((), dtype=torch.float32)
128
+ rows = s.shape[0]
129
+ for start in range(0, rows, chunk_rows):
130
+ end = min(start + chunk_rows, rows)
131
+ s_chunk = s[start:end].float() / temperature
132
+ t_chunk = t[start:end].float() / temperature
133
+ total = total + F.kl_div(
134
+ F.log_softmax(s_chunk, dim=-1),
135
+ F.softmax(t_chunk, dim=-1),
136
+ reduction="sum",
137
+ )
138
+ return total * (temperature * temperature) / max(rows, 1)
139
+
140
+
141
+ def hidden_cosine_loss(student_hidden: torch.Tensor, teacher_hidden: torch.Tensor) -> torch.Tensor:
142
+ return (1.0 - F.cosine_similarity(student_hidden.float(), teacher_hidden.float(), dim=-1)).mean()
143
+
144
+
145
+ def combined_loss(
146
+ student_out,
147
+ teacher_out,
148
+ *,
149
+ temperature: float,
150
+ kl_chunk_rows: int,
151
+ ce_weight: float,
152
+ kl_weight: float,
153
+ hidden_weight: float,
154
+ ) -> tuple[torch.Tensor, dict[str, float]]:
155
+ ce = student_out.loss.float()
156
+ kl = chunked_kl(student_out.logits, teacher_out.logits, temperature, kl_chunk_rows)
157
+ hidden = hidden_cosine_loss(student_out.hidden_states[-1], teacher_out.hidden_states[-1])
158
+ total = ce_weight * ce + kl_weight * kl + hidden_weight * hidden
159
+ return total, {
160
+ "ce": float(ce.detach()),
161
+ "kl": float(kl.detach()),
162
+ "hidden": float(hidden.detach()),
163
+ "total": float(total.detach()),
164
+ }
165
+
166
+
167
+ @torch.inference_mode()
168
+ def evaluate_distillation(
169
+ teacher,
170
+ student,
171
+ batches: list[torch.Tensor],
172
+ *,
173
+ device: torch.device,
174
+ dtype: torch.dtype,
175
+ temperature: float,
176
+ kl_chunk_rows: int,
177
+ ce_weight: float,
178
+ kl_weight: float,
179
+ hidden_weight: float,
180
+ ) -> dict[str, float | int]:
181
+ teacher.eval()
182
+ student.eval()
183
+ sums = {
184
+ "student_ce": 0.0,
185
+ "teacher_ce": 0.0,
186
+ "kl": 0.0,
187
+ "hidden": 0.0,
188
+ "total": 0.0,
189
+ }
190
+ n_batches = 0
191
+ eval_tokens = 0
192
+ amp = (lambda: torch.autocast("cuda", dtype=dtype)) if device.type == "cuda" else nullcontext
193
+ for cpu_ids in batches:
194
+ ids = cpu_ids.to(device, non_blocking=True)
195
+ with amp():
196
+ teacher_out = teacher(
197
+ input_ids=ids,
198
+ labels=ids,
199
+ output_hidden_states=True,
200
+ use_cache=False,
201
+ )
202
+ student_out = student(
203
+ input_ids=ids,
204
+ labels=ids,
205
+ output_hidden_states=True,
206
+ use_cache=False,
207
+ )
208
+ total, parts = combined_loss(
209
+ student_out,
210
+ teacher_out,
211
+ temperature=temperature,
212
+ kl_chunk_rows=kl_chunk_rows,
213
+ ce_weight=ce_weight,
214
+ kl_weight=kl_weight,
215
+ hidden_weight=hidden_weight,
216
+ )
217
+ sums["student_ce"] += float(student_out.loss.detach().float())
218
+ sums["teacher_ce"] += float(teacher_out.loss.detach().float())
219
+ sums["kl"] += parts["kl"]
220
+ sums["hidden"] += parts["hidden"]
221
+ sums["total"] += float(total.detach().float())
222
+ n_batches += 1
223
+ eval_tokens += ids.numel()
224
+
225
+ for key in sums:
226
+ sums[key] /= max(n_batches, 1)
227
+ result: dict[str, float | int] = dict(sums)
228
+ result["student_ppl"] = math.exp(min(sums["student_ce"], 30.0))
229
+ result["teacher_ppl"] = math.exp(min(sums["teacher_ce"], 30.0))
230
+ result["eval_batches"] = n_batches
231
+ result["eval_tokens"] = eval_tokens
232
+ return result
src/tinycenn_lm/gemma3_integrated_memory.py ADDED
@@ -0,0 +1,319 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Gemma 3 / FunctionGemma integrated TinyCeNN memory.
2
+
3
+ This module mirrors the SmolLM2 V3 experiment but respects Gemma 3's
4
+ Q/K RMS normalization and hybrid sliding/full-attention cache. It is intended
5
+ for replacing *full-attention* Gemma 3 text layers; the original sliding-window
6
+ layers remain untouched.
7
+ """
8
+ import copy
9
+ from contextlib import contextmanager
10
+
11
+ import torch
12
+ import torch.nn.functional as F
13
+ from torch import nn
14
+ from transformers.cache_utils import DynamicCache
15
+ from transformers.models.gemma3.modeling_gemma3 import apply_rotary_pos_emb, repeat_kv
16
+
17
+ from .optimized_memory import OptimizedMemory
18
+
19
+
20
+ FORMAT = "functiongemma-cenn-integrated-v1"
21
+
22
+
23
+ def native_dtype(device):
24
+ """Use BF16 only when the GPU has native BF16 arithmetic."""
25
+ if torch.device(device).type != "cuda":
26
+ return torch.float32
27
+ return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16
28
+
29
+
30
+ class Gemma3IntegratedCache(DynamicCache):
31
+ """Gemma 3 hybrid cache plus bounded TinyCeNN states at replaced layers."""
32
+
33
+ def __init__(self, config, memory_layers=()):
34
+ super().__init__(config=config)
35
+ self.memory_layers = frozenset(memory_layers)
36
+ self.memory_states = {}
37
+
38
+ def get_seq_length(self, layer_idx=0):
39
+ if layer_idx in self.memory_layers:
40
+ state = self.memory_states.get(layer_idx)
41
+ return state.position if state is not None else 0
42
+ return super().get_seq_length(layer_idx)
43
+
44
+ def get_mask_sizes(self, cache_position, layer_idx):
45
+ if layer_idx in self.memory_layers:
46
+ query_length = cache_position.shape[0] if isinstance(cache_position, torch.Tensor) else int(cache_position)
47
+ return self.get_seq_length(layer_idx) + query_length, 0
48
+ return super().get_mask_sizes(cache_position, layer_idx)
49
+
50
+ @property
51
+ def nbytes(self):
52
+ tensors = [
53
+ tensor
54
+ for layer in self.layers
55
+ for tensor in (getattr(layer, "keys", None), getattr(layer, "values", None))
56
+ if isinstance(tensor, torch.Tensor)
57
+ ]
58
+ return sum(x.numel() * x.element_size() for x in tensors) + sum(
59
+ state.nbytes for state in self.memory_states.values()
60
+ )
61
+
62
+ def reorder_cache(self, *args, **kwargs):
63
+ raise NotImplementedError("Use batch-one greedy decoding; beam-search cache reordering is unsupported")
64
+
65
+ def crop(self, *args, **kwargs):
66
+ raise NotImplementedError("Compressed TinyCeNN history cannot be cropped; start a fresh cache")
67
+
68
+ def batch_repeat_interleave(self, *args, **kwargs):
69
+ raise NotImplementedError("Batch expansion is unsupported for TinyCeNN memory states")
70
+
71
+ def batch_select_indices(self, *args, **kwargs):
72
+ raise NotImplementedError("Batch selection is unsupported for TinyCeNN memory states")
73
+
74
+
75
+ class Gemma3IntegratedAttention(nn.Module):
76
+ """Full-attention Gemma 3 layer replaced by a TinyCeNN memory core.
77
+
78
+ Gemma3DecoderLayer inspects attributes on ``self_attn`` before calling it
79
+ (most importantly ``is_sliding`` to select local vs global RoPE). Mirror
80
+ the lightweight structural attributes of Gemma3Attention so replacing the
81
+ module preserves the Transformers 4.57.x decoder contract.
82
+ """
83
+
84
+ def __init__(self, original, core, layer_idx):
85
+ super().__init__()
86
+ self.original = original
87
+ self.core = core
88
+ self.layer_idx = layer_idx
89
+
90
+ # Structural Gemma3Attention API used by Gemma3DecoderLayer and by
91
+ # attention tooling. These are plain metadata values, not duplicate
92
+ # module registrations; Q/K/V/O and RMSNorm modules remain under
93
+ # ``self.original`` only.
94
+ self.is_sliding = bool(original.is_sliding)
95
+ self.config = original.config
96
+ self.head_dim = original.head_dim
97
+ self.num_key_value_groups = original.num_key_value_groups
98
+ self.scaling = original.scaling
99
+ self.attention_dropout = original.attention_dropout
100
+ self.is_causal = original.is_causal
101
+ self.attn_logit_softcapping = original.attn_logit_softcapping
102
+ self.sliding_window = original.sliding_window
103
+
104
+ if self.is_sliding:
105
+ raise ValueError("Gemma3IntegratedAttention only supports full-attention layers")
106
+
107
+ self.register_buffer("fused_weight", None, persistent=False)
108
+
109
+ def fuse(self, enabled=True):
110
+ if not enabled:
111
+ self.fused_weight = None
112
+ return
113
+ with torch.no_grad():
114
+ weight = self.original.o_proj.weight.float().reshape(
115
+ -1, self.core.num_heads, self.core.head_dim
116
+ )
117
+ self.fused_weight = torch.einsum(
118
+ "ohd,hkd->ohk", weight, self.core.readout.float()
119
+ ).reshape_as(self.original.o_proj.weight).to(self.original.o_proj.weight.dtype)
120
+
121
+ def forward(
122
+ self,
123
+ hidden_states,
124
+ position_embeddings=None,
125
+ attention_mask=None,
126
+ past_key_values=None,
127
+ past_key_value=None,
128
+ cache_position=None,
129
+ **kwargs,
130
+ ):
131
+ if position_embeddings is None:
132
+ raise ValueError("Gemma 3 position_embeddings are required")
133
+ cache = past_key_values if past_key_values is not None else past_key_value
134
+ if cache is not None and not isinstance(cache, Gemma3IntegratedCache):
135
+ raise TypeError("Use Gemma3IntegratedCache with a FunctionGemma TinyCeNN model")
136
+ if cache is not None and torch.is_grad_enabled():
137
+ raise RuntimeError("Train with use_cache=False")
138
+ if self.fused_weight is not None and torch.is_grad_enabled():
139
+ raise RuntimeError("Unfuse the readout before training")
140
+
141
+ if self.is_sliding:
142
+ raise RuntimeError("TinyCeNN FunctionGemma V1 only supports replacing full-attention layers")
143
+
144
+ if attention_mask is not None:
145
+ if attention_mask.ndim != 4 or bool((attention_mask[..., -1, :] < 0).any()):
146
+ raise ValueError("Only unpadded causal batches are supported")
147
+
148
+ b, t, _ = hidden_states.shape
149
+ h = self.core.num_heads
150
+ hk = self.core.num_kv_heads
151
+ d = self.core.head_dim
152
+
153
+ q = self.original.q_proj(hidden_states).view(b, t, h, d).transpose(1, 2)
154
+ k = self.original.k_proj(hidden_states).view(b, t, hk, d).transpose(1, 2)
155
+ v = self.original.v_proj(hidden_states).view(b, t, hk, d).transpose(1, 2)
156
+
157
+ q = self.original.q_norm(q)
158
+ k = self.original.k_norm(k)
159
+ cos, sin = position_embeddings
160
+ q, k = apply_rotary_pos_emb(q, k, cos, sin)
161
+
162
+ if self.core.variant == "transformer_readout":
163
+ if cache is not None:
164
+ cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
165
+ k, v = cache.update(k, v, self.layer_idx, cache_kwargs)
166
+ output = F.scaled_dot_product_attention(
167
+ q,
168
+ repeat_kv(k, h // hk),
169
+ repeat_kv(v, h // hk),
170
+ attn_mask=attention_mask,
171
+ is_causal=attention_mask is None and t > 1,
172
+ scale=float(self.scaling),
173
+ )
174
+ if self.fused_weight is None:
175
+ output = self.core.calibrate(output.float())
176
+ elif cache is not None:
177
+ output, state = self.core(
178
+ q,
179
+ k,
180
+ v,
181
+ state=cache.memory_states.get(self.layer_idx),
182
+ return_state=True,
183
+ apply_readout=self.fused_weight is None,
184
+ )
185
+ cache.memory_states[self.layer_idx] = state
186
+ else:
187
+ output = self.core(q, k, v, apply_readout=self.fused_weight is None)
188
+
189
+ flat = output.transpose(1, 2).reshape(b, t, h * d).to(hidden_states.dtype)
190
+ if self.fused_weight is None:
191
+ return self.original.o_proj(flat), None
192
+ return F.linear(flat, self.fused_weight, self.original.o_proj.bias), None
193
+
194
+
195
+ def wrappers(model):
196
+ return [
197
+ layer.self_attn
198
+ for layer in model.model.layers
199
+ if isinstance(layer.self_attn, Gemma3IntegratedAttention)
200
+ ]
201
+
202
+
203
+ def full_attention_layers(model):
204
+ return [
205
+ i for i, layer in enumerate(model.model.layers)
206
+ if not getattr(layer.self_attn, "is_sliding", False)
207
+ ]
208
+
209
+
210
+ def build_student(teacher, layers, variant="cenn_partition", features=64, block_size=32, sinks=4):
211
+ model = copy.deepcopy(teacher).eval().requires_grad_(False)
212
+ config = model.config.get_text_config(decoder=True)
213
+ if config.model_type != "gemma3_text":
214
+ raise ValueError(f"Expected gemma3_text, got {config.model_type}")
215
+
216
+ for index in layers:
217
+ if not 0 <= index < len(model.model.layers):
218
+ raise ValueError(f"Invalid layer {index}")
219
+ original = model.model.layers[index].self_attn
220
+ if getattr(original, "is_sliding", False):
221
+ raise ValueError(
222
+ f"Layer {index} is sliding_attention. FunctionGemma V1 intentionally replaces only full_attention layers."
223
+ )
224
+ # FunctionGemma uses query_pre_attn_scalar == head_dim (256), so its
225
+ # native attention scale matches the OptimizedMemory softmax scale.
226
+ # Keep this guard explicit for future Gemma3 checkpoints where that
227
+ # assumption may not hold.
228
+ expected_scale = original.head_dim ** -0.5
229
+ if abs(float(original.scaling) - float(expected_scale)) > 1e-8:
230
+ raise ValueError(
231
+ f"Layer {index} uses attention scale {original.scaling}, but TinyCeNN currently expects {expected_scale}."
232
+ )
233
+ core = OptimizedMemory(
234
+ config.num_attention_heads,
235
+ config.num_key_value_heads,
236
+ original.head_dim,
237
+ features,
238
+ variant,
239
+ block_size,
240
+ sinks,
241
+ ).to(original.q_proj.weight.device)
242
+ model.model.layers[index].self_attn = Gemma3IntegratedAttention(original, core, index)
243
+ return model
244
+
245
+
246
+ @contextmanager
247
+ def inference_mode(model, compute_dtype="float32"):
248
+ adapters = wrappers(model)
249
+ previous = [a.core.compute_dtype for a in adapters]
250
+ try:
251
+ for adapter in adapters:
252
+ adapter.core.compute_dtype = compute_dtype
253
+ adapter.fuse()
254
+ with torch.no_grad():
255
+ yield model
256
+ finally:
257
+ for adapter, dtype in zip(adapters, previous):
258
+ adapter.fuse(False)
259
+ adapter.core.compute_dtype = dtype
260
+
261
+
262
+ def new_cache(model):
263
+ memory_layers = [
264
+ a.layer_idx for a in wrappers(model)
265
+ if a.core.variant != "transformer_readout"
266
+ ]
267
+ return Gemma3IntegratedCache(model.config, memory_layers)
268
+
269
+
270
+ def adapter_payload(model, metadata=None):
271
+ return {
272
+ "format": FORMAT,
273
+ "metadata": metadata or {},
274
+ "adapters": {
275
+ str(a.layer_idx): {
276
+ "config": a.core.config,
277
+ "state_dict": {
278
+ k: v.detach().cpu().clone() for k, v in a.core.state_dict().items()
279
+ },
280
+ }
281
+ for a in wrappers(model)
282
+ },
283
+ }
284
+
285
+
286
+ def restore_student(teacher, payload):
287
+ if payload["format"] != FORMAT:
288
+ raise ValueError(f"Not a {FORMAT} checkpoint")
289
+ model = copy.deepcopy(teacher).eval().requires_grad_(False)
290
+ for key, value in payload["adapters"].items():
291
+ index = int(key)
292
+ original = model.model.layers[index].self_attn
293
+ if getattr(original, "is_sliding", False):
294
+ raise ValueError(f"Checkpoint attempts to replace sliding-attention layer {index}")
295
+ core = OptimizedMemory(**value["config"]).to(original.q_proj.weight.device)
296
+ core.load_state_dict(value["state_dict"])
297
+ model.model.layers[index].self_attn = Gemma3IntegratedAttention(original, core, index)
298
+ return model
299
+
300
+
301
+ @torch.no_grad()
302
+ def greedy_generate(model, ids, tokens=32, stop_token_ids=()):
303
+ """Batch-one cached greedy generation for the hybrid FunctionGemma model."""
304
+ if ids.shape[0] != 1 or ids.shape[1] < 1 or tokens < 1:
305
+ raise ValueError("Use batch size one, a nonempty prompt, and positive tokens")
306
+ stop_token_ids = set(int(x) for x in stop_token_ids)
307
+ cache = new_cache(model)
308
+ output = model(input_ids=ids, past_key_values=cache, use_cache=True).logits[:, -1]
309
+ continuation = []
310
+ token = output.argmax(-1, keepdim=True)
311
+ continuation.append(token)
312
+ for _ in range(tokens - 1):
313
+ if int(continuation[-1].item()) in stop_token_ids:
314
+ break
315
+ output = model(
316
+ input_ids=continuation[-1], past_key_values=cache, use_cache=True
317
+ ).logits[:, -1]
318
+ continuation.append(output.argmax(-1, keepdim=True))
319
+ return torch.cat(continuation, dim=1), cache
src/tinycenn_lm/gemma3_memory_fusion.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import copy
4
+ import json
5
+ from dataclasses import asdict, dataclass
6
+ from pathlib import Path
7
+ from typing import Iterable
8
+
9
+ import torch
10
+ from torch import Tensor, nn
11
+ from transformers.models.gemma3.modeling_gemma3 import apply_rotary_pos_emb
12
+
13
+ from .memory_attention import MemoryAugmentedCellularLayer
14
+
15
+ DEFAULT_FUNCTIONGEMMA = "vtava/functiongemma-270m-it-simple-tool-calling"
16
+ FORMAT = "functiongemma-memory-fusion-sequential-v1"
17
+
18
+
19
+ @dataclass(frozen=True)
20
+ class Gemma3MemoryFusionConfig:
21
+ feature_dim: int = 32
22
+ memory_rank: int = 64
23
+ dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64, 128)
24
+ shifted_window: int = 8
25
+ train_output_projection: bool = True
26
+
27
+ def validate(self, model_config) -> None:
28
+ config = model_config.get_text_config(decoder=True) if hasattr(model_config, "get_text_config") else model_config
29
+ if getattr(config, "model_type", None) != "gemma3_text":
30
+ raise ValueError(f"expected gemma3_text, got {getattr(config, 'model_type', None)!r}")
31
+ if self.feature_dim < 4:
32
+ raise ValueError("feature_dim must be >= 4")
33
+ if self.memory_rank < 4:
34
+ raise ValueError("memory_rank must be >= 4")
35
+ if not self.dilations or min(self.dilations) < 1:
36
+ raise ValueError("dilations must be positive")
37
+ if int(config.num_attention_heads) % int(config.num_key_value_heads):
38
+ raise ValueError("num_attention_heads must be divisible by num_key_value_heads")
39
+
40
+ def to_dict(self) -> dict:
41
+ value = asdict(self)
42
+ value["dilations"] = list(self.dilations)
43
+ return value
44
+
45
+ @classmethod
46
+ def from_dict(cls, data: dict) -> "Gemma3MemoryFusionConfig":
47
+ value = dict(data)
48
+ value["dilations"] = tuple(value.get("dilations", (1, 2, 4, 8, 16, 32, 64, 128)))
49
+ return cls(**value)
50
+
51
+
52
+ class MemoryFusionGemma3Attention(nn.Module):
53
+ """Gemma3 full-attention replacement using the TinyCeNN Memory Fusion core.
54
+
55
+ The pretrained Q/K/V/O projections and Gemma3 Q/K RMS normalizers are copied
56
+ exactly. Only original *full-attention* layers are supported in V1; the model's
57
+ sliding-window layers remain untouched. Training and prompt checks use full
58
+ prefixes with ``use_cache=False`` so the experiment is intentionally simple
59
+ and auditable before adding a hybrid recurrent cache.
60
+ """
61
+
62
+ def __init__(self, original_attn: nn.Module, model_config, config: Gemma3MemoryFusionConfig, layer_idx: int):
63
+ super().__init__()
64
+ config.validate(model_config)
65
+ text_config = model_config.get_text_config(decoder=True) if hasattr(model_config, "get_text_config") else model_config
66
+ if bool(getattr(original_attn, "is_sliding", False)):
67
+ raise ValueError("Memory Fusion V1 replaces only Gemma3 full-attention layers")
68
+
69
+ self.layer_idx = int(layer_idx)
70
+ self.config = getattr(original_attn, "config", text_config)
71
+ self.is_sliding = False
72
+ self.hidden_size = int(text_config.hidden_size)
73
+ self.num_heads = int(text_config.num_attention_heads)
74
+ self.num_key_value_heads = int(text_config.num_key_value_heads)
75
+ self.head_dim = int(getattr(original_attn, "head_dim", text_config.head_dim))
76
+ self.attention_width = self.num_heads * self.head_dim
77
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
78
+ self.scaling = float(getattr(original_attn, "scaling", self.head_dim ** -0.5))
79
+ self.attention_dropout = float(getattr(original_attn, "attention_dropout", 0.0))
80
+ self.is_causal = bool(getattr(original_attn, "is_causal", True))
81
+ self.attn_logit_softcapping = getattr(original_attn, "attn_logit_softcapping", None)
82
+ self.sliding_window = getattr(original_attn, "sliding_window", None)
83
+
84
+ self.q_proj = copy.deepcopy(original_attn.q_proj)
85
+ self.k_proj = copy.deepcopy(original_attn.k_proj)
86
+ self.v_proj = copy.deepcopy(original_attn.v_proj)
87
+ self.o_proj = copy.deepcopy(original_attn.o_proj)
88
+ self.q_norm = copy.deepcopy(original_attn.q_norm)
89
+ self.k_norm = copy.deepcopy(original_attn.k_norm)
90
+
91
+ self.core = MemoryAugmentedCellularLayer(
92
+ num_heads=self.num_heads,
93
+ num_kv_heads=self.num_key_value_heads,
94
+ head_dim=self.head_dim,
95
+ feature_dim=config.feature_dim,
96
+ variant="cellular_memory_fusion",
97
+ dilations=config.dilations,
98
+ shifted_window=config.shifted_window,
99
+ memory_rank=config.memory_rank,
100
+ )
101
+ self.last_core_output: Tensor | None = None
102
+
103
+ def forward(
104
+ self,
105
+ hidden_states: Tensor,
106
+ position_embeddings=None,
107
+ attention_mask=None,
108
+ position_ids=None,
109
+ past_key_values=None,
110
+ past_key_value=None,
111
+ use_cache: bool = False,
112
+ cache_position=None,
113
+ **kwargs,
114
+ ) -> tuple[Tensor, None]:
115
+ if use_cache or past_key_values is not None or past_key_value is not None:
116
+ raise RuntimeError("FunctionGemma Memory Fusion V1 currently requires use_cache=False")
117
+ if position_embeddings is None:
118
+ raise ValueError("Gemma3 position_embeddings are required")
119
+
120
+ bsz, seq_len, _ = hidden_states.shape
121
+ if attention_mask is not None:
122
+ if attention_mask.ndim != 4 or attention_mask.shape[-1] != seq_len:
123
+ raise ValueError("only unpadded full causal blocks are supported")
124
+ if bool((attention_mask[..., -1, :] < -1e4).any()):
125
+ raise ValueError("padded batches are not supported")
126
+
127
+ q = self.q_proj(hidden_states).view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
128
+ k = self.k_proj(hidden_states).view(
129
+ bsz, seq_len, self.num_key_value_heads, self.head_dim
130
+ ).transpose(1, 2)
131
+ v = self.v_proj(hidden_states).view(
132
+ bsz, seq_len, self.num_key_value_heads, self.head_dim
133
+ ).transpose(1, 2)
134
+ q = self.q_norm(q)
135
+ k = self.k_norm(k)
136
+ cos, sin = position_embeddings
137
+ q, k = apply_rotary_pos_emb(q, k, cos, sin)
138
+
139
+ core_out = self.core(q.float(), k.float(), v.float())
140
+ self.last_core_output = core_out
141
+ # Gemma3 can use num_heads * head_dim != hidden_size (FunctionGemma does).
142
+ # The original o_proj maps the attention width back to hidden_size.
143
+ flat = core_out.transpose(1, 2).reshape(bsz, seq_len, self.attention_width)
144
+ return self.o_proj(flat.to(hidden_states.dtype)), None
145
+
146
+
147
+ def full_attention_layers(model: nn.Module) -> list[int]:
148
+ return [
149
+ i for i, layer in enumerate(model.model.layers)
150
+ if not bool(getattr(layer.self_attn, "is_sliding", False))
151
+ ]
152
+
153
+
154
+ def replace_attention_layers(model: nn.Module, config: Gemma3MemoryFusionConfig, layer_indices: Iterable[int]) -> nn.Module:
155
+ config.validate(model.config)
156
+ for raw_idx in layer_indices:
157
+ idx = int(raw_idx)
158
+ layer = model.model.layers[idx]
159
+ if isinstance(layer.self_attn, MemoryFusionGemma3Attention):
160
+ continue
161
+ old = layer.self_attn
162
+ if bool(getattr(old, "is_sliding", False)):
163
+ raise ValueError(f"layer {idx} is sliding_attention; V1 replaces only full_attention")
164
+ device = old.q_proj.weight.device
165
+ projection_dtype = old.q_proj.weight.dtype
166
+ new = MemoryFusionGemma3Attention(old, model.config, config, idx)
167
+ for module in (new.q_proj, new.k_proj, new.v_proj, new.o_proj, new.q_norm, new.k_norm):
168
+ module.to(device=device, dtype=projection_dtype)
169
+ new.core.to(device=device, dtype=torch.float32)
170
+ layer.self_attn = new
171
+ model.config.use_cache = False
172
+ if hasattr(model, "generation_config"):
173
+ model.generation_config.use_cache = False
174
+ return model
175
+
176
+
177
+ def freeze_current_layer_only(model: nn.Module, layer_idx: int, *, train_output_projection: bool = True) -> list[nn.Parameter]:
178
+ for p in model.parameters():
179
+ p.requires_grad = False
180
+ module = model.model.layers[int(layer_idx)].self_attn
181
+ if not isinstance(module, MemoryFusionGemma3Attention):
182
+ raise TypeError(f"layer {layer_idx} is not MemoryFusionGemma3Attention")
183
+ trainable: list[nn.Parameter] = []
184
+ for p in module.core.parameters():
185
+ p.requires_grad = True
186
+ trainable.append(p)
187
+ if train_output_projection:
188
+ for p in module.o_proj.parameters():
189
+ p.requires_grad = True
190
+ trainable.append(p)
191
+ return trainable
192
+
193
+
194
+ def freeze_all_memory_fusion(model: nn.Module, *, train_output_projection: bool = True) -> list[nn.Parameter]:
195
+ for p in model.parameters():
196
+ p.requires_grad = False
197
+ trainable: list[nn.Parameter] = []
198
+ for layer in model.model.layers:
199
+ module = layer.self_attn
200
+ if not isinstance(module, MemoryFusionGemma3Attention):
201
+ continue
202
+ for p in module.core.parameters():
203
+ p.requires_grad = True
204
+ trainable.append(p)
205
+ if train_output_projection:
206
+ for p in module.o_proj.parameters():
207
+ p.requires_grad = True
208
+ trainable.append(p)
209
+ return trainable
210
+
211
+
212
+ def structural_summary(model: nn.Module) -> dict[str, object]:
213
+ fusion = [
214
+ i for i, layer in enumerate(model.model.layers)
215
+ if isinstance(layer.self_attn, MemoryFusionGemma3Attention)
216
+ ]
217
+ full = [
218
+ i for i, layer in enumerate(model.model.layers)
219
+ if not isinstance(layer.self_attn, MemoryFusionGemma3Attention)
220
+ and not bool(getattr(layer.self_attn, "is_sliding", False))
221
+ ]
222
+ sliding = [
223
+ i for i, layer in enumerate(model.model.layers)
224
+ if not isinstance(layer.self_attn, MemoryFusionGemma3Attention)
225
+ and bool(getattr(layer.self_attn, "is_sliding", False))
226
+ ]
227
+ return {
228
+ "memory_fusion_layers": fusion,
229
+ "remaining_full_attention_layers": full,
230
+ "sliding_attention_layers": sliding,
231
+ }
232
+
233
+
234
+ def selected_attention_state(model: nn.Module, layers: Iterable[int]) -> dict[str, Tensor]:
235
+ prefixes = tuple(f"model.layers.{int(i)}.self_attn." for i in layers)
236
+ return {
237
+ key: value.detach().cpu()
238
+ for key, value in model.state_dict().items()
239
+ if prefixes and key.startswith(prefixes)
240
+ }
241
+
242
+
243
+ def save_adapter(model: nn.Module, output_dir: str | Path, *, config: Gemma3MemoryFusionConfig, base_model: str, accepted_layers: list[int], metadata: dict | None = None) -> Path:
244
+ output_dir = Path(output_dir)
245
+ output_dir.mkdir(parents=True, exist_ok=True)
246
+ torch.save(selected_attention_state(model, accepted_layers), output_dir / "functiongemma_memory_fusion.pt")
247
+ payload = {
248
+ "format": FORMAT,
249
+ "base_model": base_model,
250
+ "accepted_layers": list(accepted_layers),
251
+ "memory_fusion": config.to_dict(),
252
+ "metadata": metadata or {},
253
+ }
254
+ (output_dir / "functiongemma_memory_fusion_config.json").write_text(json.dumps(payload, indent=2), encoding="utf-8")
255
+ return output_dir
src/tinycenn_lm/hf_persistence.py ADDED
@@ -0,0 +1,411 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import hashlib
4
+ import json
5
+ import os
6
+ import platform
7
+ import re
8
+ import shutil
9
+ import sys
10
+ from datetime import datetime, timezone
11
+ from pathlib import Path
12
+ from typing import Any
13
+
14
+ _TOKEN_RE = re.compile(r"hf_[A-Za-z0-9]{20,}")
15
+ _SMALL_ARTIFACT_SUFFIXES = {".json", ".jsonl", ".csv", ".txt", ".md", ".log", ".yaml", ".yml"}
16
+
17
+
18
+ def utc_run_id(prefix: str = "run") -> str:
19
+ stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
20
+ return f"{prefix}-{stamp}"
21
+
22
+
23
+ def redact_secrets(text: str) -> str:
24
+ return _TOKEN_RE.sub("hf_REDACTED", text)
25
+
26
+
27
+ def _json_files(folder: Path) -> list[Path]:
28
+ return sorted(
29
+ p for p in folder.rglob("*.json")
30
+ if p.is_file() and p.stat().st_size <= 5 * 1024 * 1024
31
+ )
32
+
33
+
34
+ def _load_json(path: Path) -> Any | None:
35
+ try:
36
+ return json.loads(path.read_text(encoding="utf-8"))
37
+ except Exception:
38
+ return None
39
+
40
+
41
+ def collect_reports(folder: str | Path) -> dict[str, Any]:
42
+ folder = Path(folder)
43
+ reports: dict[str, Any] = {}
44
+ for path in _json_files(folder):
45
+ name = path.name.lower()
46
+ if any(key in name for key in ("report", "result", "metric", "summary", "config")):
47
+ value = _load_json(path)
48
+ if value is not None:
49
+ reports[str(path.relative_to(folder))] = value
50
+ return reports
51
+
52
+
53
+ def _first_value(reports: dict[str, Any], keys: tuple[str, ...]) -> Any | None:
54
+ def walk(value: Any) -> Any | None:
55
+ if isinstance(value, dict):
56
+ for key in keys:
57
+ if key in value and value[key] not in (None, ""):
58
+ return value[key]
59
+ for child in value.values():
60
+ found = walk(child)
61
+ if found is not None:
62
+ return found
63
+ elif isinstance(value, list):
64
+ for child in value:
65
+ found = walk(child)
66
+ if found is not None:
67
+ return found
68
+ return None
69
+
70
+ return walk(reports)
71
+
72
+
73
+ def _flatten_scalars(value: Any, prefix: str = "") -> dict[str, Any]:
74
+ out: dict[str, Any] = {}
75
+ if isinstance(value, dict):
76
+ for key, child in value.items():
77
+ name = f"{prefix}.{key}" if prefix else str(key)
78
+ out.update(_flatten_scalars(child, name))
79
+ elif isinstance(value, (str, int, float, bool)) or value is None:
80
+ out[prefix] = value
81
+ return out
82
+
83
+
84
+ def _interesting_metrics(reports: dict[str, Any]) -> list[tuple[str, Any]]:
85
+ wanted = (
86
+ "status", "stop_reason", "seen_tokens", "updates", "context_length",
87
+ "feature_dim", "num_shards", "top_k", "trainable", "trainable_percent",
88
+ "last_training_ce", "last_distillation_kl", "best.student_ce", "best.teacher_ce",
89
+ "teacher_gap_recovery_fraction", "mean_route_mix", "mean_router_entropy",
90
+ "elapsed_minutes", "peak_vram_gib", "evaluation_performed",
91
+ )
92
+ flat: dict[str, Any] = {}
93
+ for report in reports.values():
94
+ if isinstance(report, dict):
95
+ flat.update(_flatten_scalars(report))
96
+ rows: list[tuple[str, Any]] = []
97
+ seen: set[str] = set()
98
+ for want in wanted:
99
+ for key, value in flat.items():
100
+ if key == want or key.endswith("." + want):
101
+ label = want.split(".")[-1]
102
+ if label not in seen:
103
+ rows.append((label, value))
104
+ seen.add(label)
105
+ break
106
+ return rows
107
+
108
+
109
+ def _format_value(value: Any) -> str:
110
+ if isinstance(value, float):
111
+ if abs(value) >= 1000:
112
+ return f"{value:,.2f}"
113
+ if abs(value) < 0.001 and value != 0:
114
+ return f"{value:.3e}"
115
+ return f"{value:.6g}"
116
+ if isinstance(value, int):
117
+ return f"{value:,}"
118
+ return str(value)
119
+
120
+
121
+ def build_model_card(
122
+ folder: str | Path,
123
+ *,
124
+ title: str | None = None,
125
+ architecture: str | None = None,
126
+ base_model: str | None = None,
127
+ source_repo: str = "https://github.com/vtavakkoli/TinyCeNN-LM",
128
+ extra_notes: str | None = None,
129
+ ) -> str:
130
+ folder = Path(folder)
131
+ reports = collect_reports(folder)
132
+ inferred_arch = architecture or _first_value(reports, ("architecture", "task")) or "TinyCeNN-LM experiment"
133
+ inferred_base = base_model or _first_value(reports, ("base_model", "teacher_model"))
134
+ dataset = _first_value(reports, ("dataset",))
135
+ dataset_config = _first_value(reports, ("dataset_config",))
136
+ display_title = title or folder.name.replace("-", " ")
137
+
138
+ tags = ["tinycenn", "cenn", "language-modeling", "text-generation", "research"]
139
+ arch_lower = str(inferred_arch).lower()
140
+ if "distill" in arch_lower:
141
+ tags.append("knowledge-distillation")
142
+ if "moe" in arch_lower or _first_value(reports, ("num_shards", "num_experts")):
143
+ tags.append("mixture-of-experts")
144
+ if "amcenn" in arch_lower or "attention-free" in arch_lower:
145
+ tags.extend(["attention-free", "linear-attention", "recurrent-memory"])
146
+ if "story" in arch_lower:
147
+ tags.append("story-generation")
148
+
149
+ yaml_lines = ["---", "library_name: transformers", "pipeline_tag: text-generation"]
150
+ if inferred_base:
151
+ yaml_lines.append(f"base_model: {inferred_base}")
152
+ if dataset:
153
+ yaml_lines.append("datasets:")
154
+ yaml_lines.append(f"- {dataset}")
155
+ yaml_lines.append("tags:")
156
+ for tag in dict.fromkeys(tags):
157
+ yaml_lines.append(f"- {tag}")
158
+ yaml_lines.append("---")
159
+
160
+ metrics = _interesting_metrics(reports)
161
+ if metrics:
162
+ lines = ["| Metric | Value |", "|---|---:|"]
163
+ lines += [f"| `{name}` | {_format_value(value)} |" for name, value in metrics]
164
+ metric_table = "\n".join(lines)
165
+ else:
166
+ metric_table = "No structured training report was found in this upload."
167
+
168
+ report_files = sorted(reports)
169
+ files_text = "\n".join(f"- `{name}`" for name in report_files) or "- No JSON report files detected."
170
+ dataset_text = str(dataset) if dataset else "Not recorded"
171
+ if dataset_config:
172
+ dataset_text += f" (`{dataset_config}`)"
173
+
174
+ limitations = (
175
+ "This is a research checkpoint. Metrics saved here are the metrics produced by the corresponding "
176
+ "training notebook/script; unless explicitly marked as held-out evaluation, they should not be treated "
177
+ "as publication-grade benchmark results. Generation quality can differ substantially from the base model."
178
+ )
179
+ notes = f"\n## Notes\n\n{extra_notes.strip()}\n" if extra_notes else ""
180
+ return redact_secrets("\n".join(yaml_lines) + f"\n\n# {display_title}\n\n"
181
+ f"Research artifact from **TinyCeNN-LM**. Architecture: `{inferred_arch}`.\n\n"
182
+ "## Architecture\n\n"
183
+ f"- Architecture/run type: `{inferred_arch}`\n"
184
+ f"- Base model: `{inferred_base or 'not recorded'}`\n"
185
+ f"- Dataset: `{dataset_text}`\n"
186
+ f"- Source code: {source_repo}\n\n"
187
+ "## Latest saved results\n\n"
188
+ f"{metric_table}\n\n"
189
+ "The Hugging Face repository keeps timestamped run artifacts under `runs/`. This preserves training "
190
+ "reports, configs and run metadata independently of the temporary Colab filesystem.\n\n"
191
+ "## Saved experiment files\n\n"
192
+ f"{files_text}\n\n"
193
+ "## Reproducibility\n\n"
194
+ "Run the matching notebook from the TinyCeNN-LM repository. Colab notebooks use a Hugging Face write "
195
+ "token from the `HF_TOKEN` Colab Secret; tokens should never be pasted into notebook source.\n\n"
196
+ "## Limitations\n\n"
197
+ f"{limitations}\n"
198
+ f"{notes}\n"
199
+ "## Citation\n\n"
200
+ "If you use this experimental checkpoint, cite the TinyCeNN-LM repository and the upstream base model.\n"
201
+ )
202
+
203
+
204
+ def write_run_manifest(
205
+ folder: str | Path,
206
+ *,
207
+ run_id: str,
208
+ notebook: str | None = None,
209
+ repo_id: str | None = None,
210
+ ) -> Path:
211
+ folder = Path(folder)
212
+ reports = collect_reports(folder)
213
+ manifest = {
214
+ "run_id": run_id,
215
+ "created_utc": datetime.now(timezone.utc).isoformat(),
216
+ "repo_id": repo_id,
217
+ "notebook": notebook,
218
+ "python": sys.version.split()[0],
219
+ "platform": platform.platform(),
220
+ "reports": sorted(reports),
221
+ }
222
+ try:
223
+ import torch
224
+ manifest["torch"] = torch.__version__
225
+ manifest["cuda_available"] = bool(torch.cuda.is_available())
226
+ if torch.cuda.is_available():
227
+ manifest["gpu"] = torch.cuda.get_device_name(0)
228
+ except Exception:
229
+ pass
230
+ path = folder / "run_manifest.json"
231
+ path.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
232
+ return path
233
+
234
+
235
+ def _copy_small_artifacts(folder: Path, archive_dir: Path) -> list[str]:
236
+ copied: list[str] = []
237
+ for path in folder.rglob("*"):
238
+ if not path.is_file():
239
+ continue
240
+ try:
241
+ if path.is_relative_to(archive_dir.parent):
242
+ continue
243
+ except AttributeError:
244
+ pass
245
+ if path.suffix.lower() not in _SMALL_ARTIFACT_SUFFIXES:
246
+ continue
247
+ if path.stat().st_size > 10 * 1024 * 1024:
248
+ continue
249
+ rel = path.relative_to(folder)
250
+ dst = archive_dir / rel
251
+ dst.parent.mkdir(parents=True, exist_ok=True)
252
+ try:
253
+ text = path.read_text(encoding="utf-8")
254
+ dst.write_text(redact_secrets(text), encoding="utf-8")
255
+ except Exception:
256
+ shutil.copy2(path, dst)
257
+ copied.append(str(rel))
258
+ return copied
259
+
260
+
261
+ def _prepare_run_archive(folder: Path, *, repo_id: str, run_id: str, notebook: str | None = None) -> Path:
262
+ archive_dir = folder / ".hf_run_archive" / run_id
263
+ if archive_dir.exists():
264
+ shutil.rmtree(archive_dir)
265
+ archive_dir.mkdir(parents=True, exist_ok=True)
266
+ copied = _copy_small_artifacts(folder, archive_dir)
267
+ latest = {
268
+ "run_id": run_id,
269
+ "repo_id": repo_id,
270
+ "created_utc": datetime.now(timezone.utc).isoformat(),
271
+ "notebook": notebook,
272
+ "artifacts": copied,
273
+ }
274
+ (archive_dir / "latest_run.json").write_text(json.dumps(latest, indent=2), encoding="utf-8")
275
+ return archive_dir
276
+
277
+
278
+ def persist_hf_run(
279
+ *,
280
+ api,
281
+ repo_id: str,
282
+ folder_path: str | Path,
283
+ token: str | None = None,
284
+ title: str | None = None,
285
+ architecture: str | None = None,
286
+ base_model: str | None = None,
287
+ notebook: str | None = None,
288
+ run_id: str | None = None,
289
+ commit_message: str | None = None,
290
+ upload_model_files: bool = True,
291
+ ) -> dict[str, Any]:
292
+ """Persist a completed Colab run and a timestamped result archive to Hugging Face."""
293
+ from huggingface_hub import HfApi
294
+
295
+ folder = Path(folder_path)
296
+ if not folder.exists():
297
+ raise FileNotFoundError(folder)
298
+ run_id = run_id or utc_run_id(folder.name[:24] or "run")
299
+ client = api if api is not None else HfApi(token=token)
300
+ client.create_repo(repo_id=repo_id, repo_type="model", exist_ok=True)
301
+
302
+ card = build_model_card(folder, title=title, architecture=architecture, base_model=base_model)
303
+ (folder / "README.md").write_text(card, encoding="utf-8")
304
+ write_run_manifest(folder, run_id=run_id, notebook=notebook, repo_id=repo_id)
305
+
306
+ if upload_model_files:
307
+ client.upload_folder(
308
+ repo_id=repo_id, repo_type="model", folder_path=str(folder),
309
+ commit_message=commit_message or f"Publish TinyCeNN run {run_id}",
310
+ ignore_patterns=[".hf_run_archive/**"],
311
+ )
312
+
313
+ archive_dir = _prepare_run_archive(folder, repo_id=repo_id, run_id=run_id, notebook=notebook)
314
+ latest_file = archive_dir / "latest_run.json"
315
+ client.upload_folder(
316
+ repo_id=repo_id, repo_type="model", folder_path=str(archive_dir),
317
+ path_in_repo=f"runs/{run_id}", commit_message=f"Archive TinyCeNN results {run_id}",
318
+ )
319
+ client.upload_file(
320
+ repo_id=repo_id, repo_type="model", path_or_fileobj=str(latest_file),
321
+ path_in_repo="runs/latest_run.json",
322
+ commit_message=f"Update latest TinyCeNN run pointer to {run_id}",
323
+ )
324
+ return json.loads(latest_file.read_text(encoding="utf-8"))
325
+
326
+
327
+ def _looks_like_tinycenn_folder(folder: Path) -> bool:
328
+ if "tinycenn" in str(folder).lower() or "smollm2-amcenn" in str(folder).lower():
329
+ return True
330
+ names = {p.name.lower() for p in folder.iterdir()} if folder.exists() and folder.is_dir() else set()
331
+ return any("cenn" in name or "story_v2" in name or "sharded" in name for name in names)
332
+
333
+
334
+ def install_colab_hf_upload_enhancer() -> bool:
335
+ """Enhance existing notebook HfApi.upload_folder calls without duplicating notebook code.
336
+
337
+ In Colab, TinyCeNN notebooks already import tinycenn_lm before publishing. This wrapper regenerates a
338
+ report-backed README and archives small result files under runs/<timestamp>/ whenever a TinyCeNN checkpoint
339
+ folder is uploaded. Outside Colab it is a no-op.
340
+ """
341
+ if not (os.environ.get("COLAB_RELEASE_TAG") or os.environ.get("COLAB_GPU") or Path("/content").exists()):
342
+ return False
343
+ try:
344
+ from huggingface_hub import HfApi
345
+ except Exception:
346
+ return False
347
+ if getattr(HfApi.upload_folder, "_tinycenn_enhanced", False):
348
+ return True
349
+
350
+ original_upload_folder = HfApi.upload_folder
351
+ original_upload_file = HfApi.upload_file
352
+
353
+ def enhanced_upload_folder(self, *args, **kwargs):
354
+ folder_value = kwargs.get("folder_path")
355
+ repo_id = kwargs.get("repo_id")
356
+ repo_type = kwargs.get("repo_type", "model")
357
+ if folder_value is None and len(args) >= 1:
358
+ folder_value = args[0]
359
+ if repo_id is None and len(args) >= 2:
360
+ repo_id = args[1]
361
+ folder = Path(folder_value) if folder_value else None
362
+ should_enhance = (
363
+ repo_type == "model" and folder is not None and folder.exists() and folder.is_dir()
364
+ and repo_id and _looks_like_tinycenn_folder(folder)
365
+ and ".hf_run_archive" not in str(folder)
366
+ )
367
+ if not should_enhance:
368
+ return original_upload_folder(self, *args, **kwargs)
369
+
370
+ run_id = utc_run_id(folder.name[:24] or "run")
371
+ reports = collect_reports(folder)
372
+ architecture = _first_value(reports, ("architecture", "task"))
373
+ base_model = _first_value(reports, ("base_model", "teacher_model"))
374
+ (folder / "README.md").write_text(
375
+ build_model_card(folder, title=str(repo_id).split("/")[-1], architecture=architecture, base_model=base_model),
376
+ encoding="utf-8",
377
+ )
378
+ write_run_manifest(folder, run_id=run_id, repo_id=str(repo_id))
379
+ ignore = list(kwargs.get("ignore_patterns") or [])
380
+ if ".hf_run_archive/**" not in ignore:
381
+ ignore.append(".hf_run_archive/**")
382
+ kwargs["ignore_patterns"] = ignore
383
+ result = original_upload_folder(self, *args, **kwargs)
384
+
385
+ try:
386
+ archive_dir = _prepare_run_archive(folder, repo_id=str(repo_id), run_id=run_id)
387
+ original_upload_folder(
388
+ self, repo_id=repo_id, repo_type="model", folder_path=str(archive_dir),
389
+ path_in_repo=f"runs/{run_id}", commit_message=f"Archive TinyCeNN results {run_id}",
390
+ )
391
+ original_upload_file(
392
+ self, repo_id=repo_id, repo_type="model",
393
+ path_or_fileobj=str(archive_dir / "latest_run.json"),
394
+ path_in_repo="runs/latest_run.json",
395
+ commit_message=f"Update latest TinyCeNN run pointer to {run_id}",
396
+ )
397
+ except Exception as exc:
398
+ print(f"TinyCeNN HF result archive warning: {exc}")
399
+ return result
400
+
401
+ enhanced_upload_folder._tinycenn_enhanced = True
402
+ HfApi.upload_folder = enhanced_upload_folder
403
+ return True
404
+
405
+
406
+ def fingerprint_file(path: str | Path) -> str:
407
+ h = hashlib.sha256()
408
+ with Path(path).open("rb") as f:
409
+ for chunk in iter(lambda: f.read(1024 * 1024), b""):
410
+ h.update(chunk)
411
+ return h.hexdigest()
src/tinycenn_lm/integrated_memory.py ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Jointly trainable V2 mixers with full-model causal caching for Llama/SmolLM2.
2
+
3
+ Unpadded, batch-one greedy decoding is the supported inference protocol. This
4
+ adapter deliberately does not implement beam search, cache cropping or offload.
5
+ """
6
+ import copy
7
+ from contextlib import contextmanager
8
+
9
+ import torch
10
+ import torch.nn.functional as F
11
+ from torch import nn
12
+ from transformers.cache_utils import DynamicCache
13
+ from transformers.models.llama.modeling_llama import apply_rotary_pos_emb, repeat_kv
14
+
15
+ from .optimized_memory import OptimizedMemory
16
+
17
+
18
+ def native_dtype(device):
19
+ """Do not mistake emulated BF16 allocation support for native T4 arithmetic."""
20
+ if torch.device(device).type != "cuda":
21
+ return torch.float32
22
+ return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16
23
+
24
+
25
+ class IntegratedCache(DynamicCache):
26
+ def __init__(self, memory_layers=()):
27
+ super().__init__()
28
+ self.memory_layers = frozenset(memory_layers)
29
+ self.memory_states = {}
30
+
31
+ def get_seq_length(self, layer_idx=0):
32
+ if layer_idx in self.memory_layers:
33
+ state = self.memory_states.get(layer_idx)
34
+ return state.position if state is not None else 0
35
+ return super().get_seq_length(layer_idx)
36
+
37
+ def get_mask_sizes(self, cache_position, layer_idx):
38
+ if layer_idx in self.memory_layers:
39
+ # Transformers 4.x passes positions; 5.x passes the query length.
40
+ query_length = cache_position.shape[0] if isinstance(cache_position, torch.Tensor) else cache_position
41
+ return self.get_seq_length(layer_idx) + query_length, 0
42
+ return super().get_mask_sizes(cache_position, layer_idx)
43
+
44
+ @property
45
+ def nbytes(self):
46
+ tensors = [x for layer in self.layers for x in (layer.keys, layer.values)
47
+ if isinstance(x, torch.Tensor)]
48
+ return sum(x.numel() * x.element_size() for x in tensors) + sum(
49
+ state.nbytes for state in self.memory_states.values())
50
+
51
+ def reorder_cache(self, *args, **kwargs):
52
+ raise NotImplementedError("Use the supplied batch-one greedy decoder; beam search is unsupported")
53
+
54
+ def crop(self, *args, **kwargs):
55
+ raise NotImplementedError("Compressed history cannot be cropped; start a fresh cache")
56
+
57
+
58
+ class IntegratedAttention(nn.Module):
59
+ def __init__(self, original, core, layer_idx):
60
+ super().__init__()
61
+ self.original, self.core, self.layer_idx = original, core, layer_idx
62
+ self.register_buffer("fused_weight", None, persistent=False)
63
+
64
+ def fuse(self, enabled=True):
65
+ if not enabled:
66
+ self.fused_weight = None
67
+ return
68
+ with torch.no_grad():
69
+ weight = self.original.o_proj.weight.float().reshape(
70
+ -1, self.core.num_heads, self.core.head_dim)
71
+ self.fused_weight = torch.einsum("ohd,hkd->ohk", weight, self.core.readout.float()).reshape_as(
72
+ self.original.o_proj.weight).to(self.original.o_proj.weight.dtype)
73
+
74
+ def forward(self, hidden_states, position_embeddings=None, attention_mask=None,
75
+ past_key_values=None, past_key_value=None, **kwargs):
76
+ if position_embeddings is None:
77
+ raise ValueError("Llama rotary position embeddings are required")
78
+ cache = past_key_values if past_key_values is not None else past_key_value
79
+ if cache is not None and not isinstance(cache, IntegratedCache):
80
+ raise TypeError("Use IntegratedCache for this model")
81
+ if cache is not None and torch.is_grad_enabled():
82
+ raise RuntimeError("Train with use_cache=False")
83
+ if self.fused_weight is not None and torch.is_grad_enabled():
84
+ raise RuntimeError("Unfuse the readout before training")
85
+ if attention_mask is not None:
86
+ if attention_mask.ndim != 4 or bool((attention_mask[..., -1, :] < 0).any()):
87
+ raise ValueError("Only unpadded causal blocks are supported")
88
+ b, t, _ = hidden_states.shape
89
+ h, hk, d = self.core.num_heads, self.core.num_kv_heads, self.core.head_dim
90
+ q = self.original.q_proj(hidden_states).view(b, t, h, d).transpose(1, 2)
91
+ k = self.original.k_proj(hidden_states).view(b, t, hk, d).transpose(1, 2)
92
+ v = self.original.v_proj(hidden_states).view(b, t, hk, d).transpose(1, 2)
93
+ q, k = apply_rotary_pos_emb(q, k, *position_embeddings)
94
+ if self.core.variant == "transformer_readout":
95
+ if cache is not None:
96
+ k, v = cache.update(k, v, self.layer_idx)
97
+ output = F.scaled_dot_product_attention(q, repeat_kv(k, h // hk), repeat_kv(v, h // hk),
98
+ attn_mask=attention_mask, is_causal=attention_mask is None and t > 1)
99
+ if self.fused_weight is None:
100
+ output = self.core.calibrate(output.float())
101
+ elif cache is not None:
102
+ output, state = self.core(q, k, v, state=cache.memory_states.get(self.layer_idx),
103
+ return_state=True, apply_readout=self.fused_weight is None)
104
+ cache.memory_states[self.layer_idx] = state
105
+ else:
106
+ output = self.core(q, k, v, apply_readout=self.fused_weight is None)
107
+ flat = output.transpose(1, 2).reshape(b, t, h * d).to(hidden_states.dtype)
108
+ if self.fused_weight is None:
109
+ return self.original.o_proj(flat), None
110
+ return F.linear(flat, self.fused_weight, self.original.o_proj.bias), None
111
+
112
+
113
+ def wrappers(model):
114
+ return [layer.self_attn for layer in model.model.layers
115
+ if isinstance(layer.self_attn, IntegratedAttention)]
116
+
117
+
118
+ def build_student(teacher, layers, variant="cenn_partition", features=64, block_size=32, sinks=4):
119
+ model = copy.deepcopy(teacher).eval().requires_grad_(False)
120
+ config = model.config
121
+ if config.model_type != "llama":
122
+ raise ValueError("Only Llama-family models are supported")
123
+ for index in layers:
124
+ if not 0 <= index < len(model.model.layers):
125
+ raise ValueError(f"Invalid layer {index}")
126
+ original = model.model.layers[index].self_attn
127
+ core = OptimizedMemory(config.num_attention_heads, config.num_key_value_heads,
128
+ config.hidden_size // config.num_attention_heads, features,
129
+ variant, block_size, sinks).to(original.q_proj.weight.device)
130
+ model.model.layers[index].self_attn = IntegratedAttention(original, core, index)
131
+ return model
132
+
133
+
134
+ @contextmanager
135
+ def inference_mode(model, compute_dtype="float32"):
136
+ adapters = wrappers(model)
137
+ previous = [a.core.compute_dtype for a in adapters]
138
+ try:
139
+ for a in adapters:
140
+ a.core.compute_dtype = compute_dtype
141
+ a.fuse()
142
+ with torch.no_grad():
143
+ yield model
144
+ finally:
145
+ for a, dtype in zip(adapters, previous):
146
+ a.fuse(False)
147
+ a.core.compute_dtype = dtype
148
+
149
+
150
+ def new_cache(model):
151
+ return IntegratedCache(a.layer_idx for a in wrappers(model)
152
+ if a.core.variant != "transformer_readout")
153
+
154
+
155
+ def adapter_payload(model, metadata=None):
156
+ return {"format": "smollm2-integrated-memory-v3", "metadata": metadata or {}, "adapters": {
157
+ str(a.layer_idx): {"config": a.core.config,
158
+ "state_dict": {k: v.detach().cpu().clone() for k, v in a.core.state_dict().items()}}
159
+ for a in wrappers(model)}}
160
+
161
+
162
+ def restore_student(teacher, payload):
163
+ if payload["format"] != "smollm2-integrated-memory-v3":
164
+ raise ValueError("Not an integrated memory checkpoint")
165
+ model = copy.deepcopy(teacher).eval().requires_grad_(False)
166
+ for key, value in payload["adapters"].items():
167
+ index = int(key)
168
+ original = model.model.layers[index].self_attn
169
+ core = OptimizedMemory(**value["config"]).to(original.q_proj.weight.device)
170
+ core.load_state_dict(value["state_dict"])
171
+ model.model.layers[index].self_attn = IntegratedAttention(original, core, index)
172
+ return model
173
+
174
+
175
+ @torch.no_grad()
176
+ def greedy_generate(model, ids, tokens=32):
177
+ """Deterministic fixed-length generation; no early EOS for comparable timing."""
178
+ if ids.shape[0] != 1 or ids.shape[1] < 1 or tokens < 1:
179
+ raise ValueError("Use batch size one, a nonempty prompt, and positive tokens")
180
+ cache = new_cache(model)
181
+ output = model(input_ids=ids, past_key_values=cache, use_cache=True).logits[:, -1]
182
+ continuation = [output.argmax(-1, keepdim=True)]
183
+ for _ in range(tokens - 1):
184
+ output = model(input_ids=continuation[-1], past_key_values=cache, use_cache=True).logits[:, -1]
185
+ continuation.append(output.argmax(-1, keepdim=True))
186
+ return torch.cat(continuation, dim=1), cache
src/tinycenn_lm/live_console.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import atexit
4
+ import os
5
+ import sys
6
+ import time
7
+ from pathlib import Path
8
+
9
+ _STATUS_INSTALLED = False
10
+ _PROCESS_FAILED = False
11
+ _PROCESS_STARTED = 0.0
12
+ _PROCESS_NAME = ""
13
+
14
+
15
+ def _format_elapsed(seconds: float) -> str:
16
+ seconds = max(0, int(seconds))
17
+ hours, remainder = divmod(seconds, 3600)
18
+ minutes, secs = divmod(remainder, 60)
19
+ if hours:
20
+ return f"{hours:d}h {minutes:02d}m {secs:02d}s"
21
+ return f"{minutes:d}m {secs:02d}s"
22
+
23
+
24
+ def _install_training_process_status() -> None:
25
+ global _STATUS_INSTALLED, _PROCESS_STARTED, _PROCESS_NAME
26
+ if _STATUS_INSTALLED:
27
+ return
28
+
29
+ script = Path(sys.argv[0]).name
30
+ if not (script.lower().startswith("train_") and script.lower().endswith(".py")):
31
+ return
32
+
33
+ _STATUS_INSTALLED = True
34
+ _PROCESS_STARTED = time.monotonic()
35
+ _PROCESS_NAME = script[:-3]
36
+ print(f"[TinyCeNN][PROCESS START] {_PROCESS_NAME}", flush=True)
37
+
38
+ original_excepthook = sys.excepthook
39
+
40
+ def status_excepthook(exc_type, exc_value, traceback):
41
+ global _PROCESS_FAILED
42
+ _PROCESS_FAILED = True
43
+ elapsed = _format_elapsed(time.monotonic() - _PROCESS_STARTED)
44
+ print(
45
+ f"[TinyCeNN][PROCESS FAILED] {_PROCESS_NAME} after {elapsed}: "
46
+ f"{exc_type.__name__}: {exc_value}",
47
+ flush=True,
48
+ )
49
+ original_excepthook(exc_type, exc_value, traceback)
50
+
51
+ sys.excepthook = status_excepthook
52
+
53
+ def final_status() -> None:
54
+ elapsed = _format_elapsed(time.monotonic() - _PROCESS_STARTED)
55
+ if _PROCESS_FAILED:
56
+ print(f"[TinyCeNN][PROCESS END] {_PROCESS_NAME} failed after {elapsed}", flush=True)
57
+ else:
58
+ print(f"[TinyCeNN][PROCESS DONE] {_PROCESS_NAME} completed in {elapsed}", flush=True)
59
+
60
+ atexit.register(final_status)
61
+
62
+
63
+ def configure_live_console() -> None:
64
+ """Prefer immediate stdout/stderr visibility for notebook-launched trainers.
65
+
66
+ Existing Colab notebooks launch ``train_*.py`` with ``subprocess.run``. Setting
67
+ PYTHONUNBUFFERED before child interpreters start and making the current process
68
+ line-buffered keeps progress messages visible as they are produced instead of
69
+ appearing in a large block at the end of a run.
70
+
71
+ Trainer processes also emit their own PROCESS START/DONE/FAILED markers. This
72
+ covers notebooks that have not imported ``tinycenn_lm`` in the parent kernel
73
+ before launching the trainer.
74
+ """
75
+ os.environ.setdefault("PYTHONUNBUFFERED", "1")
76
+
77
+ for stream in (sys.stdout, sys.stderr):
78
+ reconfigure = getattr(stream, "reconfigure", None)
79
+ if reconfigure is None:
80
+ continue
81
+ try:
82
+ reconfigure(line_buffering=True, write_through=True)
83
+ except (TypeError, ValueError, OSError):
84
+ # Some notebook stream wrappers do not expose every TextIO option.
85
+ try:
86
+ reconfigure(line_buffering=True)
87
+ except (TypeError, ValueError, OSError):
88
+ pass
89
+
90
+ _install_training_process_status()
src/tinycenn_lm/memory_attention.py ADDED
@@ -0,0 +1,473 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Research-only global-memory augmentations for TinyCeNN-LM.
2
+
3
+ The module keeps the proven adaptive+MaxPool Cellular Attention path as a local/
4
+ multiscale branch and adds optional global causal memories inspired by recent
5
+ efficient sequence models:
6
+
7
+ * Hedgehog: learned positive feature maps for softmax-mimicking linear attention.
8
+ * Kimi Delta Attention (KDA): fine-grained per-channel forgetting plus delta updates.
9
+ * Gated DeltaNet-2: KDA-like decay with decoupled erase and write gates.
10
+ * xLSTM/mLSTM: normalized matrix memory with gated covariance-style updates.
11
+ * Differential attention: subtracts a second learned linear-attention map.
12
+ * Memory fusion: token-wise mixture of sparse Cellular, Hedgehog and GDN2 paths.
13
+
14
+ These are deliberately small, auditable reference implementations for controlled
15
+ ablation inside this repository. They are inspired by the papers, not drop-in
16
+ copies of the authors' optimized kernels.
17
+ """
18
+ from __future__ import annotations
19
+
20
+ import math
21
+ from typing import Iterable
22
+
23
+ import torch
24
+ from torch import Tensor, nn
25
+ import torch.nn.functional as F
26
+
27
+ from tinycenn_lm.cellular_attention import CellularAttentionLayer
28
+
29
+
30
+ VARIANTS = (
31
+ "cellular_adaptive_maxpool5",
32
+ "cellular_hedgehog_global",
33
+ "cellular_kda_global",
34
+ "cellular_gdn2_global",
35
+ "cellular_xlstm_global",
36
+ "cellular_diff_hedgehog",
37
+ "cellular_memory_fusion",
38
+ )
39
+
40
+ _HEDGEHOG_VARIANTS = {
41
+ "cellular_hedgehog_global",
42
+ "cellular_diff_hedgehog",
43
+ "cellular_memory_fusion",
44
+ }
45
+ _DIFF_VARIANTS = {"cellular_diff_hedgehog"}
46
+ _KDA_VARIANTS = {"cellular_kda_global"}
47
+ _GDN2_VARIANTS = {"cellular_gdn2_global", "cellular_memory_fusion"}
48
+ _XLSTM_VARIANTS = {"cellular_xlstm_global"}
49
+ _GLOBAL_VARIANTS = set(VARIANTS) - {"cellular_adaptive_maxpool5"}
50
+
51
+
52
+ class MemoryAugmentedCellularLayer(nn.Module):
53
+ """Adaptive+MaxPool Cellular Attention plus an optional global causal memory."""
54
+
55
+ def __init__(
56
+ self,
57
+ num_heads: int,
58
+ num_kv_heads: int,
59
+ head_dim: int,
60
+ feature_dim: int = 32,
61
+ variant: str = "cellular_adaptive_maxpool5",
62
+ dilations: Iterable[int] = (1, 2, 4, 8, 16, 32, 64, 128),
63
+ shifted_window: int = 8,
64
+ memory_rank: int = 16,
65
+ ):
66
+ super().__init__()
67
+ if variant not in VARIANTS:
68
+ raise ValueError(f"unknown variant {variant!r}; choose from {VARIANTS}")
69
+ if memory_rank < 4:
70
+ raise ValueError("memory_rank must be >= 4")
71
+ self.num_heads = int(num_heads)
72
+ self.num_kv_heads = int(num_kv_heads)
73
+ self.head_dim = int(head_dim)
74
+ self.feature_dim = int(feature_dim)
75
+ self.groups = self.num_heads // self.num_kv_heads
76
+ self.variant = variant
77
+ self.dilations = tuple(int(x) for x in dilations)
78
+ self.shifted_window = int(shifted_window)
79
+ self.memory_rank = int(memory_rank)
80
+
81
+ # Local branch: preserve the strongest result already measured in the repo.
82
+ self.local = CellularAttentionLayer(
83
+ num_heads=self.num_heads,
84
+ num_kv_heads=self.num_kv_heads,
85
+ head_dim=self.head_dim,
86
+ feature_dim=self.feature_dim,
87
+ variant="cellular_adaptive_maxpool5",
88
+ dilations=self.dilations,
89
+ shifted_window=self.shifted_window,
90
+ )
91
+
92
+ # Alternate branches always begin as small perturbations of the local winner.
93
+ if self.has_global_memory() and self.variant != "cellular_memory_fusion":
94
+ self.branch_mix_logit = nn.Parameter(torch.full((self.num_heads,), -2.0))
95
+ self.branch_log_gain = nn.Parameter(torch.zeros(self.num_heads))
96
+ else:
97
+ self.register_parameter("branch_mix_logit", None)
98
+ self.register_parameter("branch_log_gain", None)
99
+
100
+ # Shared low-rank Q/K maps for recurrent matrix memories.
101
+ if self.uses_kda() or self.uses_gdn2() or self.uses_xlstm():
102
+ self.mem_q = nn.Parameter(torch.empty(
103
+ self.num_heads, self.memory_rank, self.head_dim
104
+ ))
105
+ self.mem_k = nn.Parameter(torch.empty(
106
+ self.num_heads, self.memory_rank, self.head_dim
107
+ ))
108
+ for h in range(self.num_heads):
109
+ nn.init.orthogonal_(self.mem_q[h])
110
+ nn.init.orthogonal_(self.mem_k[h])
111
+ else:
112
+ self.register_parameter("mem_q", None)
113
+ self.register_parameter("mem_k", None)
114
+
115
+ # KDA/GDN2-style fine-grained forgetting and token-wise write rate.
116
+ if self.uses_kda() or self.uses_gdn2():
117
+ self.decay_w = nn.Parameter(torch.zeros(
118
+ self.num_heads, self.memory_rank, self.head_dim
119
+ ))
120
+ self.decay_bias = nn.Parameter(torch.full(
121
+ (self.num_heads, self.memory_rank), 4.0
122
+ ))
123
+ self.beta_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim))
124
+ self.beta_bias = nn.Parameter(torch.full((self.num_heads,), -1.5))
125
+ else:
126
+ self.register_parameter("decay_w", None)
127
+ self.register_parameter("decay_bias", None)
128
+ self.register_parameter("beta_w", None)
129
+ self.register_parameter("beta_bias", None)
130
+
131
+ # GDN2-inspired decoupled key-side erase and value-side write controls.
132
+ if self.uses_gdn2():
133
+ self.erase_w = nn.Parameter(torch.zeros(
134
+ self.num_heads, self.memory_rank, self.head_dim
135
+ ))
136
+ self.erase_bias = nn.Parameter(torch.full(
137
+ (self.num_heads, self.memory_rank), -0.5
138
+ ))
139
+ self.write_scale = nn.Parameter(torch.zeros(
140
+ self.num_heads, self.head_dim
141
+ ))
142
+ self.write_bias = nn.Parameter(torch.full(
143
+ (self.num_heads, self.head_dim), -0.5
144
+ ))
145
+ else:
146
+ self.register_parameter("erase_w", None)
147
+ self.register_parameter("erase_bias", None)
148
+ self.register_parameter("write_scale", None)
149
+ self.register_parameter("write_bias", None)
150
+
151
+ # mLSTM-inspired matrix-memory gates.
152
+ if self.uses_xlstm():
153
+ self.x_forget_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim))
154
+ self.x_forget_bias = nn.Parameter(torch.full((self.num_heads,), 3.0))
155
+ self.x_input_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim))
156
+ self.x_input_bias = nn.Parameter(torch.full((self.num_heads,), -1.0))
157
+ else:
158
+ self.register_parameter("x_forget_w", None)
159
+ self.register_parameter("x_forget_bias", None)
160
+ self.register_parameter("x_input_w", None)
161
+ self.register_parameter("x_input_bias", None)
162
+
163
+ # Hedgehog-inspired trainable positive feature maps. The softmax over the
164
+ # learned feature axis enforces positivity and can become low-entropy/spiky.
165
+ if self.uses_hedgehog():
166
+ self.hedge_q = nn.Parameter(torch.empty(
167
+ self.num_heads, self.memory_rank, self.head_dim
168
+ ))
169
+ self.hedge_k = nn.Parameter(torch.empty(
170
+ self.num_heads, self.memory_rank, self.head_dim
171
+ ))
172
+ self.hedge_q_bias = nn.Parameter(torch.zeros(
173
+ self.num_heads, self.memory_rank
174
+ ))
175
+ self.hedge_k_bias = nn.Parameter(torch.zeros(
176
+ self.num_heads, self.memory_rank
177
+ ))
178
+ self.hedge_log_sharpness = nn.Parameter(torch.zeros(self.num_heads))
179
+ for h in range(self.num_heads):
180
+ nn.init.orthogonal_(self.hedge_q[h])
181
+ nn.init.orthogonal_(self.hedge_k[h])
182
+ else:
183
+ self.register_parameter("hedge_q", None)
184
+ self.register_parameter("hedge_k", None)
185
+ self.register_parameter("hedge_q_bias", None)
186
+ self.register_parameter("hedge_k_bias", None)
187
+ self.register_parameter("hedge_log_sharpness", None)
188
+
189
+ if self.uses_differential():
190
+ self.hedge2_q = nn.Parameter(torch.empty(
191
+ self.num_heads, self.memory_rank, self.head_dim
192
+ ))
193
+ self.hedge2_k = nn.Parameter(torch.empty(
194
+ self.num_heads, self.memory_rank, self.head_dim
195
+ ))
196
+ self.hedge2_q_bias = nn.Parameter(torch.zeros(
197
+ self.num_heads, self.memory_rank
198
+ ))
199
+ self.hedge2_k_bias = nn.Parameter(torch.zeros(
200
+ self.num_heads, self.memory_rank
201
+ ))
202
+ self.hedge2_log_sharpness = nn.Parameter(torch.zeros(self.num_heads))
203
+ self.diff_lambda_logit = nn.Parameter(torch.full((self.num_heads,), -1.0))
204
+ for h in range(self.num_heads):
205
+ nn.init.orthogonal_(self.hedge2_q[h])
206
+ nn.init.orthogonal_(self.hedge2_k[h])
207
+ else:
208
+ self.register_parameter("hedge2_q", None)
209
+ self.register_parameter("hedge2_k", None)
210
+ self.register_parameter("hedge2_q_bias", None)
211
+ self.register_parameter("hedge2_k_bias", None)
212
+ self.register_parameter("hedge2_log_sharpness", None)
213
+ self.register_parameter("diff_lambda_logit", None)
214
+
215
+ # The fusion candidate makes the branch choice input-dependent.
216
+ if self.variant == "cellular_memory_fusion":
217
+ self.fusion_gate_w = nn.Parameter(torch.zeros(
218
+ self.num_heads, 3, self.head_dim
219
+ ))
220
+ prior = torch.tensor([2.0, -1.0, -1.0])
221
+ self.fusion_gate_bias = nn.Parameter(
222
+ prior[None, :].expand(self.num_heads, -1).clone()
223
+ )
224
+ else:
225
+ self.register_parameter("fusion_gate_w", None)
226
+ self.register_parameter("fusion_gate_bias", None)
227
+
228
+ @property
229
+ def config(self) -> dict:
230
+ return {
231
+ "num_heads": self.num_heads,
232
+ "num_kv_heads": self.num_kv_heads,
233
+ "head_dim": self.head_dim,
234
+ "feature_dim": self.feature_dim,
235
+ "variant": self.variant,
236
+ "dilations": list(self.dilations),
237
+ "shifted_window": self.shifted_window,
238
+ "memory_rank": self.memory_rank,
239
+ }
240
+
241
+ def has_global_memory(self) -> bool:
242
+ return self.variant in _GLOBAL_VARIANTS
243
+
244
+ def uses_hedgehog(self) -> bool:
245
+ return self.variant in _HEDGEHOG_VARIANTS
246
+
247
+ def uses_differential(self) -> bool:
248
+ return self.variant in _DIFF_VARIANTS
249
+
250
+ def uses_kda(self) -> bool:
251
+ return self.variant in _KDA_VARIANTS
252
+
253
+ def uses_gdn2(self) -> bool:
254
+ return self.variant in _GDN2_VARIANTS
255
+
256
+ def uses_xlstm(self) -> bool:
257
+ return self.variant in _XLSTM_VARIANTS
258
+
259
+ def _repeat_kv(self, x: Tensor) -> Tensor:
260
+ return x.repeat_interleave(self.groups, dim=1)
261
+
262
+ @staticmethod
263
+ def _project(x: Tensor, weight: Tensor) -> Tensor:
264
+ return torch.einsum("bhtd,hrd->bhtr", x, weight)
265
+
266
+ def _memory_qk(self, q: Tensor, k: Tensor) -> tuple[Tensor, Tensor, Tensor]:
267
+ assert self.mem_q is not None and self.mem_k is not None
268
+ kh = self._repeat_kv(k)
269
+ qm = F.normalize(self._project(q, self.mem_q), dim=-1)
270
+ km = F.normalize(self._project(kh, self.mem_k), dim=-1)
271
+ return qm, km, kh
272
+
273
+ def _hedgehog_features(
274
+ self,
275
+ x: Tensor,
276
+ weight: Tensor,
277
+ bias: Tensor,
278
+ log_sharpness: Tensor,
279
+ ) -> Tensor:
280
+ logits = self._project(x, weight) + bias[None, :, None, :]
281
+ sharpness = log_sharpness.clamp(-1.4, 2.1).exp()[None, :, None, None]
282
+ # sqrt(rank) keeps q.k magnitudes from vanishing as rank grows.
283
+ return logits.mul(sharpness).softmax(dim=-1) * math.sqrt(self.memory_rank)
284
+
285
+ def _hedgehog_linear(
286
+ self,
287
+ q: Tensor,
288
+ k: Tensor,
289
+ v: Tensor,
290
+ *,
291
+ second: bool = False,
292
+ ) -> Tensor:
293
+ kh, vh = self._repeat_kv(k), self._repeat_kv(v)
294
+ if second:
295
+ assert self.hedge2_q is not None and self.hedge2_k is not None
296
+ assert self.hedge2_q_bias is not None and self.hedge2_k_bias is not None
297
+ assert self.hedge2_log_sharpness is not None
298
+ qf = self._hedgehog_features(
299
+ q, self.hedge2_q, self.hedge2_q_bias, self.hedge2_log_sharpness
300
+ )
301
+ kf = self._hedgehog_features(
302
+ kh, self.hedge2_k, self.hedge2_k_bias, self.hedge2_log_sharpness
303
+ )
304
+ else:
305
+ assert self.hedge_q is not None and self.hedge_k is not None
306
+ assert self.hedge_q_bias is not None and self.hedge_k_bias is not None
307
+ assert self.hedge_log_sharpness is not None
308
+ qf = self._hedgehog_features(
309
+ q, self.hedge_q, self.hedge_q_bias, self.hedge_log_sharpness
310
+ )
311
+ kf = self._hedgehog_features(
312
+ kh, self.hedge_k, self.hedge_k_bias, self.hedge_log_sharpness
313
+ )
314
+
315
+ kv = torch.einsum("bhtr,bhtd->bhtrd", kf, vh).cumsum(dim=2)
316
+ kz = kf.cumsum(dim=2)
317
+ numerator = torch.einsum("bhtr,bhtrd->bhtd", qf, kv)
318
+ denominator = torch.einsum("bhtr,bhtr->bht", qf, kz)
319
+ return numerator / denominator.clamp_min(1e-6)[..., None]
320
+
321
+ def _delta_memory(self, q: Tensor, k: Tensor, v: Tensor, *, gdn2: bool) -> Tensor:
322
+ qm, km, kh = self._memory_qk(q, k)
323
+ vh = self._repeat_kv(v)
324
+ assert self.decay_w is not None and self.decay_bias is not None
325
+ assert self.beta_w is not None and self.beta_bias is not None
326
+
327
+ decay = torch.sigmoid(
328
+ torch.einsum("bhtd,hrd->bhtr", kh, self.decay_w)
329
+ + self.decay_bias[None, :, None, :]
330
+ )
331
+ beta = torch.sigmoid(
332
+ torch.einsum("bhtd,hd->bht", kh, self.beta_w)
333
+ + self.beta_bias[None, :, None]
334
+ )
335
+
336
+ if gdn2:
337
+ assert self.erase_w is not None and self.erase_bias is not None
338
+ assert self.write_scale is not None and self.write_bias is not None
339
+ erase = torch.sigmoid(
340
+ torch.einsum("bhtd,hrd->bhtr", kh, self.erase_w)
341
+ + self.erase_bias[None, :, None, :]
342
+ )
343
+ write = torch.sigmoid(
344
+ vh * self.write_scale[None, :, None, :]
345
+ + self.write_bias[None, :, None, :]
346
+ )
347
+ else:
348
+ erase = None
349
+ write = None
350
+
351
+ b, h, t, _ = qm.shape
352
+ state = torch.zeros(
353
+ b, h, self.memory_rank, self.head_dim,
354
+ device=q.device, dtype=q.dtype,
355
+ )
356
+ outputs = []
357
+ for i in range(t):
358
+ state = state * decay[:, :, i, :, None]
359
+ pred = torch.einsum("bhr,bhrd->bhd", km[:, :, i], state)
360
+ error = vh[:, :, i] - pred
361
+ key_write = km[:, :, i]
362
+ if gdn2:
363
+ assert erase is not None and write is not None
364
+ key_write = key_write * erase[:, :, i]
365
+ error = error * write[:, :, i]
366
+ update = torch.einsum("bhr,bhd->bhrd", key_write, error)
367
+ state = state + beta[:, :, i, None, None] * update
368
+ outputs.append(torch.einsum("bhr,bhrd->bhd", qm[:, :, i], state))
369
+ return torch.stack(outputs, dim=2)
370
+
371
+ def _xlstm_memory(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
372
+ qm, km, kh = self._memory_qk(q, k)
373
+ vh = self._repeat_kv(v)
374
+ assert self.x_forget_w is not None and self.x_forget_bias is not None
375
+ assert self.x_input_w is not None and self.x_input_bias is not None
376
+
377
+ forget = torch.sigmoid(
378
+ torch.einsum("bhtd,hd->bht", kh, self.x_forget_w)
379
+ + self.x_forget_bias[None, :, None]
380
+ )
381
+ inp = torch.sigmoid(
382
+ torch.einsum("bhtd,hd->bht", kh, self.x_input_w)
383
+ + self.x_input_bias[None, :, None]
384
+ )
385
+ b, h, t, _ = qm.shape
386
+ memory = torch.zeros(
387
+ b, h, self.memory_rank, self.head_dim,
388
+ device=q.device, dtype=q.dtype,
389
+ )
390
+ normalizer = torch.zeros(
391
+ b, h, self.memory_rank, device=q.device, dtype=q.dtype
392
+ )
393
+ outputs = []
394
+ for i in range(t):
395
+ f = forget[:, :, i, None, None]
396
+ ii = inp[:, :, i, None, None]
397
+ outer = torch.einsum("bhr,bhd->bhrd", km[:, :, i], vh[:, :, i])
398
+ memory = f * memory + ii * outer
399
+ normalizer = (
400
+ forget[:, :, i, None] * normalizer
401
+ + inp[:, :, i, None] * km[:, :, i]
402
+ )
403
+ numerator = torch.einsum("bhr,bhrd->bhd", qm[:, :, i], memory)
404
+ denominator = torch.einsum(
405
+ "bhr,bhr->bh", qm[:, :, i], normalizer
406
+ ).abs().clamp_min(1.0)
407
+ outputs.append(numerator / denominator[..., None])
408
+ return torch.stack(outputs, dim=2)
409
+
410
+ def _merge(self, local: Tensor, branch: Tensor) -> Tensor:
411
+ assert self.branch_mix_logit is not None and self.branch_log_gain is not None
412
+ gate = self.branch_mix_logit.sigmoid()[None, :, None, None]
413
+ gain = self.branch_log_gain.clamp(-2, 2).exp()[None, :, None, None]
414
+ branch = branch * gain
415
+ return local + gate * (branch - local)
416
+
417
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
418
+ if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
419
+ raise ValueError("expected Q/K/V as [batch, heads, time, dim]")
420
+ local = self.local(q, k, v)
421
+ if self.variant == "cellular_adaptive_maxpool5":
422
+ return local
423
+
424
+ q = q.to(local.dtype)
425
+ k = k.to(local.dtype)
426
+ v = v.to(local.dtype)
427
+
428
+ if self.variant == "cellular_hedgehog_global":
429
+ return self._merge(local, self._hedgehog_linear(q, k, v))
430
+
431
+ if self.variant == "cellular_diff_hedgehog":
432
+ assert self.diff_lambda_logit is not None
433
+ first = self._hedgehog_linear(q, k, v)
434
+ second = self._hedgehog_linear(q, k, v, second=True)
435
+ lam = 0.5 * self.diff_lambda_logit.sigmoid()[None, :, None, None]
436
+ return self._merge(local, first - lam * second)
437
+
438
+ if self.variant == "cellular_kda_global":
439
+ return self._merge(local, self._delta_memory(q, k, v, gdn2=False))
440
+
441
+ if self.variant == "cellular_gdn2_global":
442
+ return self._merge(local, self._delta_memory(q, k, v, gdn2=True))
443
+
444
+ if self.variant == "cellular_xlstm_global":
445
+ return self._merge(local, self._xlstm_memory(q, k, v))
446
+
447
+ if self.variant == "cellular_memory_fusion":
448
+ assert self.fusion_gate_w is not None and self.fusion_gate_bias is not None
449
+ hedge = self._hedgehog_linear(q, k, v)
450
+ gdn2 = self._delta_memory(q, k, v, gdn2=True)
451
+ logits = (
452
+ torch.einsum("bhtd,hcd->bhtc", q, self.fusion_gate_w)
453
+ + self.fusion_gate_bias[None, :, None, :]
454
+ )
455
+ weights = logits.softmax(dim=-1)
456
+ return (
457
+ weights[..., 0, None] * local
458
+ + weights[..., 1, None] * hedge
459
+ + weights[..., 2, None] * gdn2
460
+ )
461
+
462
+ raise ValueError(self.variant)
463
+
464
+ def max_score_pairs(self, context: int) -> int:
465
+ # This counts only the sparse Cellular softmax branch. Global memories
466
+ # are O(T * memory_rank * head_dim), not pairwise T^2 score matrices.
467
+ return self.local.max_score_pairs(context)
468
+
469
+ def receptive_field_tokens(self) -> int:
470
+ return self.local.receptive_field_tokens()
471
+
472
+ def max_neighbors_per_step(self) -> int:
473
+ return self.local.max_neighbors_per_step()
src/tinycenn_lm/modeling.py ADDED
@@ -0,0 +1,251 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from pathlib import Path
5
+ from typing import Sequence
6
+
7
+ import torch
8
+ from torch import Tensor, nn
9
+
10
+ from .cenn import CeNNConfig, FastCeNNCore
11
+
12
+
13
+ DEFAULT_BASE_MODEL = "arnir0/Tiny-LLM"
14
+
15
+
16
+ def _floating_reference_parameter(module: nn.Module) -> nn.Parameter | None:
17
+ """Return a representative floating-point parameter for device/dtype alignment."""
18
+ return next((p for p in module.parameters() if p.is_floating_point()), None)
19
+
20
+
21
+ class HybridDecoderLayer(nn.Module):
22
+ """Wrap a pretrained decoder layer with a zero-init recurrent CeNN residual.
23
+
24
+ The base layer is kept intact, including its attention behavior. TinyCeNN-LM
25
+ v0.1 intentionally disables Transformer-only KV caching because the CeNN branch
26
+ also needs its own recurrent neighborhood state for exact incremental decoding.
27
+
28
+ The newly created CeNN branch inherits the pretrained decoder layer's floating
29
+ point dtype/device. This is essential for BF16/FP16 inference: a BF16 hidden
30
+ state cannot be convolved with an FP32 CeNN kernel without an explicit cast.
31
+ """
32
+
33
+ def __init__(self, base_layer: nn.Module, config: CeNNConfig) -> None:
34
+ super().__init__()
35
+ self.base_layer = base_layer
36
+ self.cenn = FastCeNNCore(config)
37
+
38
+ reference = _floating_reference_parameter(base_layer)
39
+ if reference is None:
40
+ self.residual_scale = nn.Parameter(torch.ones(()))
41
+ else:
42
+ self.cenn.to(device=reference.device, dtype=reference.dtype)
43
+ self.residual_scale = nn.Parameter(
44
+ torch.ones((), device=reference.device, dtype=reference.dtype)
45
+ )
46
+
47
+ def forward(self, *args, **kwargs):
48
+ if kwargs.get("use_cache", False):
49
+ raise RuntimeError(
50
+ "TinyCeNN-LM v0.1 requires use_cache=False. The Transformer KV cache "
51
+ "does not contain the per-step CeNN neighborhood state needed for exact "
52
+ "incremental generation. Full-prefix generation is correct; a dedicated "
53
+ "streaming CeNN cache is planned for a later version."
54
+ )
55
+ outputs = self.base_layer(*args, **kwargs)
56
+
57
+ if torch.is_tensor(outputs):
58
+ return outputs + self.residual_scale * self.cenn(outputs)
59
+
60
+ if isinstance(outputs, tuple):
61
+ hidden = outputs[0]
62
+ hidden = hidden + self.residual_scale * self.cenn(hidden)
63
+ return (hidden, *outputs[1:])
64
+
65
+ if isinstance(outputs, list):
66
+ hidden = outputs[0]
67
+ hidden = hidden + self.residual_scale * self.cenn(hidden)
68
+ return [hidden, *outputs[1:]]
69
+
70
+ raise TypeError(
71
+ "Unsupported decoder-layer output type: "
72
+ f"{type(outputs)!r}. Expected Tensor, tuple, or list."
73
+ )
74
+
75
+
76
+ def _get_decoder_layers(model: nn.Module) -> nn.ModuleList:
77
+ candidates = (
78
+ ("model", "layers"),
79
+ ("model", "model", "layers"),
80
+ )
81
+ for path in candidates:
82
+ obj = model
83
+ try:
84
+ for name in path:
85
+ obj = getattr(obj, name)
86
+ except AttributeError:
87
+ continue
88
+ if isinstance(obj, nn.ModuleList):
89
+ return obj
90
+ raise ValueError(
91
+ "Could not locate decoder layers. TinyCeNN-LM currently targets "
92
+ "Llama-family causal language models such as arnir0/Tiny-LLM."
93
+ )
94
+
95
+
96
+ def inject_cenn(
97
+ model: nn.Module,
98
+ config: CeNNConfig | None = None,
99
+ layer_indices: Sequence[int] = (0,),
100
+ ) -> nn.Module:
101
+ layers = _get_decoder_layers(model)
102
+ hidden_size = int(getattr(model.config, "hidden_size"))
103
+ if config is None:
104
+ config = CeNNConfig(hidden_size=hidden_size)
105
+ elif config.hidden_size != hidden_size:
106
+ raise ValueError(
107
+ f"CeNN hidden_size={config.hidden_size} does not match "
108
+ f"model hidden_size={hidden_size}"
109
+ )
110
+
111
+ # A Transformer KV cache alone is insufficient for the recurrent CeNN state.
112
+ # Disable it globally to prevent a silent train/inference mismatch.
113
+ if hasattr(model, "config"):
114
+ model.config.use_cache = False
115
+ if hasattr(model, "generation_config"):
116
+ model.generation_config.use_cache = False
117
+
118
+ for index in layer_indices:
119
+ if index < 0 or index >= len(layers):
120
+ raise IndexError(f"layer index {index} out of range [0, {len(layers)})")
121
+ if isinstance(layers[index], HybridDecoderLayer):
122
+ raise ValueError(f"layer {index} already has a CeNN adapter")
123
+ layers[index] = HybridDecoderLayer(layers[index], config)
124
+ return model
125
+
126
+
127
+ def freeze_for_adapter_training(
128
+ model: nn.Module,
129
+ train_lm_head: bool = False,
130
+ train_embeddings: bool = False,
131
+ ) -> None:
132
+ for parameter in model.parameters():
133
+ parameter.requires_grad = False
134
+
135
+ for module in model.modules():
136
+ if isinstance(module, HybridDecoderLayer):
137
+ for parameter in module.cenn.parameters():
138
+ parameter.requires_grad = True
139
+ module.residual_scale.requires_grad = True
140
+
141
+ if train_lm_head and hasattr(model, "lm_head"):
142
+ for parameter in model.lm_head.parameters():
143
+ parameter.requires_grad = True
144
+ if train_embeddings:
145
+ embeddings = model.get_input_embeddings()
146
+ for parameter in embeddings.parameters():
147
+ parameter.requires_grad = True
148
+
149
+
150
+ def trainable_parameter_summary(model: nn.Module) -> dict[str, int | float]:
151
+ total = sum(p.numel() for p in model.parameters())
152
+ trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
153
+ return {
154
+ "total": total,
155
+ "trainable": trainable,
156
+ "trainable_percent": 100.0 * trainable / max(total, 1),
157
+ }
158
+
159
+
160
+ def _adapter_state_dict(model: nn.Module) -> dict[str, Tensor]:
161
+ state: dict[str, Tensor] = {}
162
+ for name, tensor in model.state_dict().items():
163
+ if ".cenn." in name or name.endswith(".residual_scale"):
164
+ state[name] = tensor.detach().cpu()
165
+ if not state:
166
+ raise ValueError("no CeNN adapter found in model")
167
+ return state
168
+
169
+
170
+ def save_adapter(
171
+ model: nn.Module,
172
+ output_dir: str | Path,
173
+ *,
174
+ base_model: str = DEFAULT_BASE_MODEL,
175
+ layer_indices: Sequence[int] = (0,),
176
+ config: CeNNConfig,
177
+ ) -> Path:
178
+ output_dir = Path(output_dir)
179
+ output_dir.mkdir(parents=True, exist_ok=True)
180
+ weights_path = output_dir / "cenn_adapter.pt"
181
+ metadata_path = output_dir / "cenn_config.json"
182
+
183
+ torch.save(_adapter_state_dict(model), weights_path)
184
+ metadata = {
185
+ "format_version": 1,
186
+ "base_model": base_model,
187
+ "layer_indices": list(layer_indices),
188
+ "cenn": config.to_dict(),
189
+ }
190
+ metadata_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8")
191
+ return output_dir
192
+
193
+
194
+ def load_adapter(
195
+ model: nn.Module,
196
+ adapter_dir: str | Path,
197
+ *,
198
+ map_location: str | torch.device = "cpu",
199
+ strict: bool = True,
200
+ ) -> nn.Module:
201
+ adapter_dir = Path(adapter_dir)
202
+ state = torch.load(
203
+ adapter_dir / "cenn_adapter.pt",
204
+ map_location=map_location,
205
+ weights_only=True,
206
+ )
207
+ incompatible = model.load_state_dict(state, strict=False)
208
+ unexpected = [k for k in incompatible.unexpected_keys if ".cenn." not in k]
209
+ if strict and unexpected:
210
+ raise RuntimeError(f"unexpected adapter keys: {unexpected}")
211
+ missing_adapter = [
212
+ k
213
+ for k in _adapter_state_dict(model)
214
+ if k in incompatible.missing_keys
215
+ ]
216
+ if strict and missing_adapter:
217
+ raise RuntimeError(f"missing adapter keys: {missing_adapter}")
218
+ return model
219
+
220
+
221
+ def build_from_adapter(
222
+ adapter_dir: str | Path,
223
+ *,
224
+ device: str | torch.device | None = None,
225
+ dtype: torch.dtype | None = None,
226
+ attn_implementation: str = "sdpa",
227
+ ):
228
+ from transformers import AutoModelForCausalLM
229
+
230
+ adapter_dir = Path(adapter_dir)
231
+ metadata = json.loads((adapter_dir / "cenn_config.json").read_text())
232
+ base_model = metadata["base_model"]
233
+ config = CeNNConfig.from_dict(metadata["cenn"])
234
+ layer_indices = tuple(metadata["layer_indices"])
235
+
236
+ kwargs = {"attn_implementation": attn_implementation}
237
+ if dtype is not None:
238
+ # Modern Transformers uses `dtype`; `torch_dtype` is deprecated.
239
+ kwargs["dtype"] = dtype
240
+ model = AutoModelForCausalLM.from_pretrained(base_model, **kwargs)
241
+ inject_cenn(model, config=config, layer_indices=layer_indices)
242
+ load_adapter(model, adapter_dir)
243
+
244
+ move_kwargs: dict[str, object] = {}
245
+ if device is not None:
246
+ move_kwargs["device"] = device
247
+ if dtype is not None:
248
+ move_kwargs["dtype"] = dtype
249
+ if move_kwargs:
250
+ model.to(**move_kwargs)
251
+ return model
src/tinycenn_lm/moe.py ADDED
@@ -0,0 +1,334 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from dataclasses import asdict, dataclass
5
+ from pathlib import Path
6
+ from typing import Sequence
7
+
8
+ import torch
9
+ import torch.nn.functional as F
10
+ from torch import Tensor, nn
11
+
12
+ from .cenn import CeNNConfig, CausalDepthwiseNeighborhood, StableRMSNorm
13
+ from .modeling import DEFAULT_BASE_MODEL, _get_decoder_layers
14
+
15
+
16
+ @dataclass(frozen=True)
17
+ class MoECeNNConfig:
18
+ hidden_size: int = 192
19
+ kernel_size: int = 3
20
+ expansion: int = 4
21
+ steps: int = 7
22
+ dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64)
23
+ rms_norm_eps: float = 1e-5
24
+ dropout: float = 0.0
25
+ num_experts: int = 8
26
+ top_k: int = 2
27
+ router_noise_std: float = 1e-3
28
+
29
+ def validate(self) -> None:
30
+ CeNNConfig(
31
+ hidden_size=self.hidden_size,
32
+ kernel_size=self.kernel_size,
33
+ expansion=self.expansion,
34
+ steps=self.steps,
35
+ dilations=self.dilations,
36
+ rms_norm_eps=self.rms_norm_eps,
37
+ dropout=self.dropout,
38
+ ).validate()
39
+ if self.num_experts < 2:
40
+ raise ValueError("num_experts must be >= 2")
41
+ if not 1 <= self.top_k <= self.num_experts:
42
+ raise ValueError("top_k must be in [1, num_experts]")
43
+ if self.router_noise_std < 0:
44
+ raise ValueError("router_noise_std must be >= 0")
45
+
46
+ def to_dict(self) -> dict:
47
+ data = asdict(self)
48
+ data["dilations"] = list(self.dilations)
49
+ return data
50
+
51
+ @classmethod
52
+ def from_dict(cls, data: dict) -> "MoECeNNConfig":
53
+ data = dict(data)
54
+ if "dilations" in data:
55
+ data["dilations"] = tuple(data["dilations"])
56
+ return cls(**data)
57
+
58
+
59
+ class SwiGLUExpert(nn.Module):
60
+ def __init__(self, hidden_size: int, expansion: int, dropout: float = 0.0) -> None:
61
+ super().__init__()
62
+ inner = hidden_size * expansion
63
+ self.in_proj = nn.Linear(hidden_size, inner * 2, bias=False)
64
+ self.out_proj = nn.Linear(inner, hidden_size, bias=False)
65
+ self.dropout = nn.Dropout(dropout)
66
+ nn.init.zeros_(self.out_proj.weight)
67
+
68
+ def forward(self, x: Tensor) -> Tensor:
69
+ a, b = self.in_proj(x).chunk(2, dim=-1)
70
+ return self.dropout(self.out_proj(F.silu(a) * b))
71
+
72
+
73
+ class Top2Router(nn.Module):
74
+ def __init__(self, hidden_size: int, num_experts: int, top_k: int, noise_std: float) -> None:
75
+ super().__init__()
76
+ self.num_experts = num_experts
77
+ self.top_k = top_k
78
+ self.noise_std = noise_std
79
+ self.proj = nn.Linear(hidden_size, num_experts, bias=False)
80
+ nn.init.normal_(self.proj.weight, mean=0.0, std=noise_std)
81
+
82
+ def forward(self, x: Tensor) -> tuple[Tensor, Tensor, dict[str, Tensor]]:
83
+ logits = self.proj(x).float()
84
+ probs = F.softmax(logits, dim=-1)
85
+ top_values, top_indices = torch.topk(probs, k=self.top_k, dim=-1)
86
+ top_weights = top_values / top_values.sum(dim=-1, keepdim=True).clamp_min(1e-9)
87
+
88
+ assignment = F.one_hot(top_indices, num_classes=self.num_experts).float().sum(dim=-2)
89
+ assignment = assignment / float(self.top_k)
90
+ expert_fraction = assignment.mean(dim=(0, 1))
91
+ probability_fraction = probs.mean(dim=(0, 1))
92
+ load_balance = self.num_experts * torch.sum(expert_fraction * probability_fraction)
93
+ z_loss = torch.logsumexp(logits, dim=-1).pow(2).mean()
94
+ entropy = -(probs * probs.clamp_min(1e-9).log()).sum(dim=-1).mean()
95
+ return top_indices, top_weights.to(dtype=x.dtype), {
96
+ "load_balance": load_balance,
97
+ "z_loss": z_loss,
98
+ "entropy": entropy,
99
+ "expert_fraction": expert_fraction,
100
+ "probability_fraction": probability_fraction,
101
+ }
102
+
103
+
104
+ class MoESharedCeNNCell(nn.Module):
105
+ def __init__(self, config: MoECeNNConfig) -> None:
106
+ super().__init__()
107
+ config.validate()
108
+ self.config = config
109
+ self.norm = StableRMSNorm(config.hidden_size, config.rms_norm_eps)
110
+ self.neighborhood = CausalDepthwiseNeighborhood(config.hidden_size, config.kernel_size)
111
+ self.gate_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=True)
112
+ nn.init.constant_(self.gate_proj.bias, -1.0)
113
+ self.router = Top2Router(
114
+ config.hidden_size, config.num_experts, config.top_k, config.router_noise_std
115
+ )
116
+ self.experts = nn.ModuleList(
117
+ SwiGLUExpert(config.hidden_size, config.expansion, config.dropout)
118
+ for _ in range(config.num_experts)
119
+ )
120
+
121
+ def forward(self, state: Tensor, dilation: int, step_scale: float) -> tuple[Tensor, dict[str, Tensor]]:
122
+ x = self.norm(state)
123
+ local = self.neighborhood(x, dilation=dilation)
124
+ top_idx, top_weight, stats = self.router(local)
125
+
126
+ flat = local.reshape(-1, local.shape[-1])
127
+ idx_flat = top_idx.reshape(-1, self.config.top_k)
128
+ weight_flat = top_weight.reshape(-1, self.config.top_k)
129
+ update = torch.zeros_like(flat)
130
+ for expert_id, expert in enumerate(self.experts):
131
+ selected = idx_flat.eq(expert_id)
132
+ positions = selected.nonzero(as_tuple=False)
133
+ if positions.numel() == 0:
134
+ continue
135
+ rows = positions[:, 0]
136
+ slots = positions[:, 1]
137
+ expert_out = expert(flat.index_select(0, rows))
138
+ weighted = expert_out * weight_flat[rows, slots].unsqueeze(-1)
139
+ update = update.index_add(0, rows, weighted)
140
+ update = update.view_as(local)
141
+ gate = torch.sigmoid(self.gate_proj(local))
142
+ return state + step_scale * gate * update, stats
143
+
144
+
145
+ class FastMoECeNNCore(nn.Module):
146
+ def __init__(self, config: MoECeNNConfig) -> None:
147
+ super().__init__()
148
+ config.validate()
149
+ self.config = config
150
+ self.cell = MoESharedCeNNCell(config)
151
+ self.last_router_stats: dict[str, Tensor] = {}
152
+
153
+ def forward(self, hidden_states: Tensor) -> Tensor:
154
+ initial = hidden_states
155
+ state = hidden_states
156
+ step_scale = self.config.steps ** -0.5
157
+ accum: dict[str, Tensor] = {}
158
+ expert_fraction = None
159
+ probability_fraction = None
160
+ for step in range(self.config.steps):
161
+ dilation = self.config.dilations[step % len(self.config.dilations)]
162
+ state, stats = self.cell(state, dilation=dilation, step_scale=step_scale)
163
+ for key in ("load_balance", "z_loss", "entropy"):
164
+ accum[key] = accum.get(key, stats[key].new_zeros(())) + stats[key]
165
+ expert_fraction = stats["expert_fraction"] if expert_fraction is None else expert_fraction + stats["expert_fraction"]
166
+ probability_fraction = stats["probability_fraction"] if probability_fraction is None else probability_fraction + stats["probability_fraction"]
167
+ self.last_router_stats = {
168
+ "load_balance": accum["load_balance"] / self.config.steps,
169
+ "z_loss": accum["z_loss"] / self.config.steps,
170
+ "entropy": accum["entropy"] / self.config.steps,
171
+ "expert_fraction": expert_fraction / self.config.steps,
172
+ "probability_fraction": probability_fraction / self.config.steps,
173
+ }
174
+ return state - initial
175
+
176
+ @property
177
+ def receptive_field(self) -> int:
178
+ radius = sum(self.config.dilations[i % len(self.config.dilations)] for i in range(self.config.steps))
179
+ return 1 + (self.config.kernel_size - 1) * radius
180
+
181
+
182
+ class MoECeNNReplacementLayer(nn.Module):
183
+ def __init__(self, config: MoECeNNConfig, *, device=None, dtype=None) -> None:
184
+ super().__init__()
185
+ self.config = config
186
+ self.cenn = FastMoECeNNCore(config)
187
+ if device is not None or dtype is not None:
188
+ kwargs = {}
189
+ if device is not None:
190
+ kwargs["device"] = device
191
+ if dtype is not None:
192
+ kwargs["dtype"] = dtype
193
+ self.cenn.to(**kwargs)
194
+
195
+ def forward(self, hidden_states: Tensor, *args, **kwargs) -> Tensor:
196
+ if kwargs.get("use_cache", False):
197
+ raise RuntimeError("MoE-CeNN student requires use_cache=False")
198
+ if kwargs.get("output_attentions", False):
199
+ raise RuntimeError("MoE-CeNN student has no attention matrices")
200
+ return hidden_states + self.cenn(hidden_states)
201
+
202
+
203
+ def replace_transformer_with_moe_cenn(model: nn.Module, config: MoECeNNConfig, layer_indices: Sequence[int] = (0,)) -> nn.Module:
204
+ layers = _get_decoder_layers(model)
205
+ if config.hidden_size != int(model.config.hidden_size):
206
+ raise ValueError("MoE-CeNN hidden size does not match base model")
207
+ model.config.use_cache = False
208
+ if hasattr(model, "generation_config"):
209
+ model.generation_config.use_cache = False
210
+ for index in layer_indices:
211
+ old = layers[index]
212
+ reference = next((p for p in old.parameters() if p.is_floating_point()), None)
213
+ layers[index] = MoECeNNReplacementLayer(
214
+ config,
215
+ device=reference.device if reference is not None else None,
216
+ dtype=reference.dtype if reference is not None else None,
217
+ )
218
+ return model
219
+
220
+
221
+ def freeze_moe_student_interfaces(model: nn.Module) -> None:
222
+ for parameter in model.parameters():
223
+ parameter.requires_grad = False
224
+ for module in model.modules():
225
+ if isinstance(module, MoECeNNReplacementLayer):
226
+ for parameter in module.parameters():
227
+ parameter.requires_grad = True
228
+
229
+
230
+ def warmstart_moe_from_plain_cenn(model: nn.Module, plain_student_dir: str | Path) -> None:
231
+ """Initialize shared dynamics and clone the trained dense FFN into every expert.
232
+
233
+ With identical expert weights, Top-2 weighted routing initially reproduces the
234
+ dense CeNN FFN output (weights sum to one), while tiny router noise allows the
235
+ experts to specialize during training.
236
+ """
237
+ state = torch.load(Path(plain_student_dir) / "cenn_student.pt", map_location="cpu", weights_only=True)
238
+ target = model.state_dict()
239
+ copied = 0
240
+ for target_name in list(target):
241
+ if ".cenn.cell.experts." in target_name:
242
+ suffix = target_name.split(".experts.", 1)[1].split(".", 1)[1]
243
+ prefix = target_name.split(".cenn.cell.experts.", 1)[0] + ".cenn.cell."
244
+ source_name = prefix + suffix
245
+ elif any(part in target_name for part in (".cenn.cell.norm.", ".cenn.cell.neighborhood.", ".cenn.cell.gate_proj.")):
246
+ source_name = target_name
247
+ elif ".cenn." not in target_name and target_name in state:
248
+ # v2 dense checkpoints may include adapted language interfaces.
249
+ source_name = target_name
250
+ else:
251
+ continue
252
+ if source_name in state and state[source_name].shape == target[target_name].shape:
253
+ target[target_name].copy_(state[source_name].to(dtype=target[target_name].dtype))
254
+ copied += 1
255
+ if copied == 0:
256
+ raise RuntimeError("could not map plain CeNN weights into MoE-CeNN model")
257
+ model.load_state_dict(target, strict=False)
258
+ model._cenn_interface_keys = tuple(name for name in state if ".cenn." not in name)
259
+
260
+
261
+ def moe_router_stats(model: nn.Module) -> dict[str, Tensor]:
262
+ layer = next((m for m in model.modules() if isinstance(m, MoECeNNReplacementLayer)), None)
263
+ if layer is None or not layer.cenn.last_router_stats:
264
+ raise RuntimeError("router statistics unavailable; run a forward pass first")
265
+ return layer.cenn.last_router_stats
266
+
267
+
268
+ def _moe_state_dict(model: nn.Module) -> dict[str, Tensor]:
269
+ interfaces = set(getattr(model, "_cenn_interface_keys", ()))
270
+ state = {name: tensor.detach().cpu() for name, tensor in model.state_dict().items()
271
+ if ".cenn." in name or name in interfaces}
272
+ if not state:
273
+ raise ValueError("no MoE-CeNN weights found")
274
+ return state
275
+
276
+
277
+ def save_moe_cenn_student(model: nn.Module, output_dir: str | Path, *, config: MoECeNNConfig, base_model: str = DEFAULT_BASE_MODEL, layer_indices: Sequence[int] = (0,), extra_metadata: dict | None = None) -> Path:
278
+ output_dir = Path(output_dir)
279
+ output_dir.mkdir(parents=True, exist_ok=True)
280
+ state = _moe_state_dict(model)
281
+ torch.save(state, output_dir / "moe_cenn_student.pt")
282
+ metadata = {
283
+ "format_version": 2,
284
+ "architecture": "moe-cenn-top2-replacement",
285
+ "base_model": base_model,
286
+ "layer_indices": list(layer_indices),
287
+ "moe_cenn": config.to_dict(),
288
+ "state_keys": sorted(state),
289
+ }
290
+ if extra_metadata:
291
+ metadata["training"] = extra_metadata
292
+ (output_dir / "moe_student_config.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8")
293
+ return output_dir
294
+
295
+
296
+ def load_moe_cenn_student_weights(model: nn.Module, student_dir: str | Path, *, map_location="cpu", strict: bool = True) -> nn.Module:
297
+ state = torch.load(Path(student_dir) / "moe_cenn_student.pt", map_location=map_location, weights_only=True)
298
+ metadata = json.loads((Path(student_dir) / "moe_student_config.json").read_text())
299
+ core_keys = {name for name in model.state_dict() if ".cenn." in name}
300
+ expected = set(metadata.get("state_keys", core_keys)) | core_keys
301
+ missing = sorted(expected - state.keys())
302
+ unexpected = sorted(state.keys() - expected | state.keys() - model.state_dict().keys())
303
+ if strict and missing:
304
+ raise RuntimeError(f"missing MoE-CeNN keys: {missing}")
305
+ if strict and unexpected:
306
+ raise RuntimeError(f"unexpected MoE-CeNN keys: {unexpected}")
307
+ model.load_state_dict(state, strict=False)
308
+ model._cenn_interface_keys = tuple(name for name in state if ".cenn." not in name)
309
+ return model
310
+
311
+
312
+ def build_moe_cenn_student(student_dir: str | Path, *, device=None, dtype=None, attn_implementation: str = "sdpa"):
313
+ from transformers import AutoModelForCausalLM
314
+
315
+ student_dir = Path(student_dir)
316
+ metadata = json.loads((student_dir / "moe_student_config.json").read_text())
317
+ if metadata.get("architecture") != "moe-cenn-top2-replacement":
318
+ raise ValueError("checkpoint is not an MoE-CeNN Top-2 student")
319
+ kwargs = {"attn_implementation": attn_implementation}
320
+ if dtype is not None:
321
+ kwargs["dtype"] = dtype
322
+ model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs)
323
+ config = MoECeNNConfig.from_dict(metadata["moe_cenn"])
324
+ replace_transformer_with_moe_cenn(model, config, tuple(metadata["layer_indices"]))
325
+ load_moe_cenn_student_weights(model, student_dir)
326
+ move_kwargs = {}
327
+ if device is not None:
328
+ move_kwargs["device"] = device
329
+ if dtype is not None:
330
+ move_kwargs["dtype"] = dtype
331
+ if move_kwargs:
332
+ model.to(**move_kwargs)
333
+ model.config.use_cache = False
334
+ return model
src/tinycenn_lm/optimized_memory.py ADDED
@@ -0,0 +1,308 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Parallel normalized memory with disjoint exact sink/local attention.
2
+
3
+ Independent research implementation. See OPTIMIZED_MEMORY.md for derivation.
4
+ The block partition is terraced: current and previous block are exact, older
5
+ non-sink blocks are compressed. No T-by-T matrix or triangular solve is used
6
+ by the new memory candidates.
7
+ """
8
+ from dataclasses import dataclass
9
+ import math
10
+
11
+ import torch
12
+ import torch.nn.functional as F
13
+ from torch import nn
14
+
15
+ VARIANTS = ("cenn_linear", "cenn_partition", "sink_window", "transformer_readout")
16
+
17
+
18
+ @dataclass
19
+ class MemoryState:
20
+ numerator: torch.Tensor
21
+ denominator: torch.Tensor
22
+ keys: torch.Tensor
23
+ values: torch.Tensor
24
+ sinks_k: torch.Tensor
25
+ sinks_v: torch.Tensor
26
+ position: int
27
+
28
+ @property
29
+ def nbytes(self):
30
+ return sum(t.numel() * t.element_size() for t in (
31
+ self.numerator, self.denominator, self.keys, self.values,
32
+ self.sinks_k, self.sinks_v))
33
+
34
+
35
+ class OptimizedMemory(nn.Module):
36
+ def __init__(self, num_heads, num_kv_heads, head_dim, feature_dim=64,
37
+ variant="cenn_partition", block_size=32, sink_tokens=4,
38
+ compute_dtype="float32"):
39
+ super().__init__()
40
+ if variant not in VARIANTS:
41
+ raise ValueError(f"Unknown variant {variant}")
42
+ if min(num_heads, num_kv_heads, head_dim, feature_dim, block_size) < 1:
43
+ raise ValueError("Dimensions must be positive")
44
+ if num_heads % num_kv_heads or feature_dim % 2 or sink_tokens < 0:
45
+ raise ValueError("Require valid GQA, even feature_dim, nonnegative sink count")
46
+ if compute_dtype not in ("float32", "float16", "bfloat16"):
47
+ raise ValueError("Unknown compute dtype")
48
+ self.num_heads, self.num_kv_heads = num_heads, num_kv_heads
49
+ self.head_dim, self.feature_dim = head_dim, feature_dim
50
+ self.groups = num_heads // num_kv_heads
51
+ self.variant, self.block_size, self.sink_tokens = variant, block_size, sink_tokens
52
+ self.compute_dtype = compute_dtype
53
+ self.has_memory = variant in ("cenn_linear", "cenn_partition")
54
+ if self.has_memory:
55
+ weight = torch.randn(num_kv_heads, feature_dim // 2, head_dim) / math.sqrt(head_dim)
56
+ self.wk = nn.Parameter(weight.clone())
57
+ self.wq = nn.Parameter(weight.repeat_interleave(self.groups, dim=0).clone())
58
+ if variant == "cenn_partition":
59
+ self.log_mass = nn.Parameter(torch.full((num_heads,), math.log(feature_dim)))
60
+ self.mass_w = nn.Parameter(torch.zeros(num_heads, head_dim))
61
+ readout = torch.eye(head_dim).repeat(num_heads, 1, 1)
62
+ if variant == "sink_window":
63
+ self.register_buffer("readout", readout)
64
+ else:
65
+ self.readout = nn.Parameter(readout)
66
+
67
+ @property
68
+ def config(self):
69
+ return dict(num_heads=self.num_heads, num_kv_heads=self.num_kv_heads,
70
+ head_dim=self.head_dim, feature_dim=self.feature_dim,
71
+ variant=self.variant, block_size=self.block_size,
72
+ sink_tokens=self.sink_tokens, compute_dtype=self.compute_dtype)
73
+
74
+ def mm(self, a, b):
75
+ dtype = getattr(torch, self.compute_dtype) if a.is_cuda else torch.float32
76
+ return torch.matmul(a.to(dtype), b.to(dtype)).float()
77
+
78
+ def attend(self, q, k, v, mask=None, causal=False):
79
+ dtype = getattr(torch, self.compute_dtype) if q.is_cuda else torch.float32
80
+ shape = q.shape
81
+ if q.ndim == 5:
82
+ b, h, n, c, d = shape
83
+ length = k.shape[-2]
84
+ q, k, v = (x.reshape(b * h * n, 1, x.shape[-2], d) for x in (q, k, v))
85
+ if mask is not None:
86
+ mask = mask.expand(b, h, n, c, length).reshape(b * h * n, 1, c, length)
87
+ out = F.scaled_dot_product_attention(
88
+ q.to(dtype), k.to(dtype), v.to(dtype), attn_mask=mask, is_causal=causal
89
+ ).float()
90
+ return out.reshape(shape)
91
+
92
+ def features(self, x, query=False):
93
+ weight = self.wq if query else self.wk
94
+ # Retain input norms: do not impose the previous experiment's L2 normalization.
95
+ logits = self.mm(x.float() / self.head_dim ** 0.25,
96
+ weight.view(1, weight.shape[0], *([1] * (x.ndim - 4)),
97
+ weight.shape[1], weight.shape[2]).transpose(-1, -2))
98
+ return torch.cat((logits, -logits), dim=-1).softmax(dim=-1)
99
+
100
+ def calibrate(self, output):
101
+ return self.mm(output, self.readout[None])
102
+
103
+ def empty_state(self, k):
104
+ b = k.shape[0]
105
+ empty = k.new_empty(b, self.num_kv_heads, 0, self.head_dim)
106
+ n = k.new_zeros(b, self.num_kv_heads, self.feature_dim, self.head_dim
107
+ ) if self.has_memory else k.new_empty(0)
108
+ z = k.new_zeros(b, self.num_kv_heads, self.feature_dim
109
+ ) if self.has_memory else k.new_empty(0)
110
+ return MemoryState(n, z, empty.clone(), empty.clone(), empty.clone(),
111
+ empty.clone(), 0)
112
+
113
+ def combine(self, q, local_k, local_v, valid, numerator=None, denominator=None):
114
+ """Shared log-stabilized denominator for exact and compressed contributions."""
115
+ if self.variant == "sink_window":
116
+ return self.attend(q, local_k.repeat_interleave(self.groups, 1),
117
+ local_v.repeat_interleave(self.groups, 1), mask=valid)
118
+ scores = self.mm(q / math.sqrt(self.head_dim),
119
+ local_k.repeat_interleave(self.groups, dim=1).transpose(-1, -2))
120
+ scores = scores.masked_fill(~valid, float("-inf"))
121
+ maximum = scores.amax(dim=-1, keepdim=True)
122
+ if self.variant == "cenn_partition":
123
+ qp = self.features(q, query=True)
124
+ num = self.mm(qp, numerator.repeat_interleave(self.groups, dim=1))
125
+ den = self.mm(qp, denominator.repeat_interleave(self.groups, dim=1).unsqueeze(-1))
126
+ raw_q = F.normalize(q.float(), dim=-1)
127
+ # Works for [B,H,T,D] and [B,H,N,C,D].
128
+ mass_shape = [1, self.num_heads] + [1] * (q.ndim - 3)
129
+ mass_w_shape = [1, self.num_heads] + [1] * (q.ndim - 3) + [self.head_dim]
130
+ mass = self.log_mass.view(mass_shape) + (
131
+ raw_q * self.mass_w.view(mass_w_shape)).sum(-1)
132
+ log_global = mass.clamp(-12, 12).unsqueeze(-1) + den.clamp_min(1e-30).log()
133
+ log_global = torch.where(den > 0, log_global, float("-inf"))
134
+ maximum = torch.maximum(maximum, log_global)
135
+ global_weight = (log_global - maximum).exp()
136
+ global_value = num / den.clamp_min(1e-20)
137
+ weights = (scores - maximum).exp()
138
+ local_num = self.mm(weights, local_v.repeat_interleave(self.groups, dim=1))
139
+ local_den = weights.sum(dim=-1, keepdim=True)
140
+ if self.variant == "cenn_partition":
141
+ return (local_num + global_weight * global_value) / (
142
+ local_den + global_weight).clamp_min(1e-20)
143
+ return local_num / local_den.clamp_min(1e-20)
144
+
145
+ @staticmethod
146
+ def prefix(blocks, delay):
147
+ cumulative = blocks.cumsum(dim=2)
148
+ zeros = torch.zeros_like(blocks[:, :, :1]).expand(
149
+ *blocks.shape[:2], delay, *blocks.shape[3:])
150
+ return torch.cat((zeros, cumulative), dim=2)[:, :, :blocks.shape[2]]
151
+
152
+ def prefill(self, q, k, v, need_state=True):
153
+ b, _, t, d = q.shape
154
+ state = self.empty_state(k) if need_state else None
155
+ if self.variant == "transformer_readout":
156
+ output = self.attend(q, k.repeat_interleave(self.groups, 1),
157
+ v.repeat_interleave(self.groups, 1), causal=True)
158
+ if need_state:
159
+ state.keys, state.values, state.position = k.clone(), v.clone(), t
160
+ return output, state
161
+ c = self.block_size
162
+ n = (t + c - 1) // c
163
+ pad = n * c - t
164
+ kb = F.pad(k, (0, 0, 0, pad)).reshape(b, self.num_kv_heads, n, c, d)
165
+ vb = F.pad(v, (0, 0, 0, pad)).reshape(b, self.num_kv_heads, n, c, d)
166
+ if self.has_memory:
167
+ phi_k = self.features(k)
168
+ if self.variant == "cenn_partition":
169
+ phi_k = phi_k * (torch.arange(t, device=q.device) >= self.sink_tokens)[None, None, :, None]
170
+ pk = F.pad(phi_k, (0, 0, 0, pad)).reshape(
171
+ b, self.num_kv_heads, n, c, self.feature_dim)
172
+ writes = self.mm(pk.transpose(-1, -2), vb)
173
+ masses = pk.sum(dim=-2)
174
+ delay = 1 if self.variant == "cenn_linear" else 2
175
+ past_n, past_z = self.prefix(writes, delay), self.prefix(masses, delay)
176
+ if self.variant == "cenn_linear":
177
+ pq = F.pad(self.features(q, query=True), (0, 0, 0, pad)).reshape(
178
+ b, self.num_heads, n, c, self.feature_dim)
179
+ within = self.mm(pq, pk.repeat_interleave(self.groups, 1).transpose(-1, -2)).tril()
180
+ numerator = self.mm(pq, past_n.repeat_interleave(self.groups, 1)) + self.mm(
181
+ within, vb.repeat_interleave(self.groups, 1))
182
+ denominator = self.mm(
183
+ pq, past_z.repeat_interleave(self.groups, 1).unsqueeze(-1)
184
+ ) + within.sum(-1, keepdim=True)
185
+ output = (numerator / denominator.clamp_min(1e-20)).reshape(b, self.num_heads, n * c, d)[:, :, :t]
186
+ if need_state:
187
+ state.numerator, state.denominator = writes.sum(2), masses.sum(2)
188
+ else:
189
+ qb = F.pad(q, (0, 0, 0, pad)).reshape(b, self.num_heads, n, c, d)
190
+ previous_k = torch.cat((torch.zeros_like(kb[:, :, :1]), kb[:, :, :-1]), dim=2)
191
+ previous_v = torch.cat((torch.zeros_like(vb[:, :, :1]), vb[:, :, :-1]), dim=2)
192
+ s = min(t, self.sink_tokens)
193
+ sinks_k = k[:, :, :s].unsqueeze(2).expand(b, self.num_kv_heads, n, s, d)
194
+ sinks_v = v[:, :, :s].unsqueeze(2).expand(b, self.num_kv_heads, n, s, d)
195
+ local_k = torch.cat((sinks_k, previous_k, kb), dim=3)
196
+ local_v = torch.cat((sinks_v, previous_v, vb), dim=3)
197
+ block = torch.arange(n, device=q.device)[:, None]
198
+ offset = torch.arange(c, device=q.device)[None]
199
+ positions = block * c + offset
200
+ local_positions = torch.cat((
201
+ torch.arange(s, device=q.device)[None].expand(n, s),
202
+ positions - c, positions), dim=1)
203
+ query_positions = positions[:, :, None]
204
+ valid = (local_positions[:, None, :] <= query_positions) & (
205
+ local_positions[:, None, :] >= 0) & (local_positions[:, None, :] < t)
206
+ # Every sink occurs only in the sink columns, never twice in exact local attention.
207
+ valid[:, :, s:] &= local_positions[:, None, s:] >= self.sink_tokens
208
+ output = self.combine(qb, local_k, local_v, valid[None, None],
209
+ past_n if self.has_memory else None,
210
+ past_z if self.has_memory else None)
211
+ output = output.reshape(b, self.num_heads, n * c, d)[:, :, :t]
212
+ if need_state:
213
+ if self.has_memory:
214
+ state.numerator, state.denominator = past_n[:, :, -1].clone(), past_z[:, :, -1].clone()
215
+ keep = min(t, c + (t - 1) % c + 1)
216
+ state.keys, state.values = k[:, :, -keep:].clone(), v[:, :, -keep:].clone()
217
+ state.sinks_k, state.sinks_v = k[:, :, :s].clone(), v[:, :, :s].clone()
218
+ if need_state:
219
+ state.position = t
220
+ return output, state
221
+
222
+ def step(self, q, k, v, state):
223
+ """One token; the cache and compressed state are bounded in context length."""
224
+ t, c = state.position, self.block_size
225
+ num, den = state.numerator, state.denominator
226
+ keys, values = state.keys, state.values
227
+ sinks_k, sinks_v = state.sinks_k, state.sinks_v
228
+ if self.variant == "transformer_readout":
229
+ keys, values = torch.cat((keys, k), 2), torch.cat((values, v), 2)
230
+ output = self.attend(q, keys.repeat_interleave(self.groups, 1),
231
+ values.repeat_interleave(self.groups, 1), causal=False)
232
+ elif self.variant == "cenn_linear":
233
+ pk = self.features(k)
234
+ num = num + self.mm(pk.transpose(-1, -2), v)
235
+ den = den + pk[:, :, 0]
236
+ qp = self.features(q, query=True)
237
+ output = self.mm(qp, num.repeat_interleave(self.groups, 1)) / self.mm(
238
+ qp, den.repeat_interleave(self.groups, 1).unsqueeze(-1)).clamp_min(1e-20)
239
+ else:
240
+ if t % c == 0 and keys.shape[2] > c:
241
+ retired = keys.shape[2] - c
242
+ if self.has_memory:
243
+ pk = self.features(keys[:, :, :retired])
244
+ positions = torch.arange(t - keys.shape[2], t - c, device=q.device)
245
+ pk = pk * (positions >= self.sink_tokens)[None, None, :, None]
246
+ num = num + self.mm(pk.transpose(-1, -2), values[:, :, :retired])
247
+ den = den + pk.sum(2)
248
+ keys, values = keys[:, :, -c:].clone(), values[:, :, -c:].clone()
249
+ keys, values = torch.cat((keys, k), 2), torch.cat((values, v), 2)
250
+ if t < self.sink_tokens:
251
+ sinks_k, sinks_v = torch.cat((sinks_k, k), 2), torch.cat((sinks_v, v), 2)
252
+ lk, lv = torch.cat((sinks_k, keys), 2), torch.cat((sinks_v, values), 2)
253
+ positions = torch.arange(t + 1 - keys.shape[2], t + 1, device=q.device)
254
+ valid = torch.cat((torch.ones(sinks_k.shape[2], dtype=torch.bool, device=q.device),
255
+ positions >= self.sink_tokens))[None, None, None]
256
+ output = self.combine(q, lk, lv, valid, num, den)
257
+ return output, MemoryState(num, den, keys, values, sinks_k, sinks_v, t + 1)
258
+
259
+ def forward(self, q, k, v, state=None, return_state=False, apply_readout=True):
260
+ if (q.ndim != 4 or k.shape != v.shape or q.shape[0] != k.shape[0]
261
+ or q.shape[2:] != k.shape[2:] or q.shape[1] != self.num_heads
262
+ or k.shape[1] != self.num_kv_heads or q.shape[-1] != self.head_dim
263
+ or q.shape[2] < 1):
264
+ raise ValueError("Incompatible Q/K/V")
265
+ q, k, v = q.float(), k.float(), v.float()
266
+ if state is None:
267
+ output, state = self.prefill(q, k, v, need_state=return_state)
268
+ else:
269
+ outputs = []
270
+ for i in range(q.shape[2]):
271
+ value, state = self.step(q[:, :, i:i+1], k[:, :, i:i+1], v[:, :, i:i+1], state)
272
+ outputs.append(value)
273
+ output = torch.cat(outputs, 2)
274
+ if apply_readout and self.variant != "sink_window":
275
+ output = self.calibrate(output)
276
+ return (output, state) if return_state else output
277
+
278
+
279
+ @torch.no_grad()
280
+ def ridge_calibrate(core, samples, device, relative_ridge=0.01):
281
+ """Identity-prior ridge solution, fitted only on training attention outputs.
282
+
283
+ Solve (Y^T Y + lambda I) R = Y^T T + lambda I. Diagonal symmetric
284
+ preconditioning preserves the exact solution while improving conditioning.
285
+ CPU float64 is used for the small D-by-D systems.
286
+ """
287
+ h, d = core.num_heads, core.head_dim
288
+ gram = torch.zeros(h, d, d, dtype=torch.float64)
289
+ cross = torch.zeros_like(gram)
290
+ before, count = 0.0, 0
291
+ for q, k, v, target in samples:
292
+ y = core(q.to(device), k.to(device), v.to(device), apply_readout=False).cpu().double()
293
+ target = target.double()
294
+ y = y.permute(1, 0, 2, 3).reshape(h, -1, d)
295
+ target = target.permute(1, 0, 2, 3).reshape(h, -1, d)
296
+ gram += y.transpose(-1, -2) @ y
297
+ cross += y.transpose(-1, -2) @ target
298
+ before += (y - target).square().sum().item()
299
+ count += y.numel()
300
+ identity = torch.eye(d, dtype=torch.float64).expand(h, d, d)
301
+ ridge = relative_ridge * gram.diagonal(dim1=-2, dim2=-1).mean(-1).clamp_min(1e-8)
302
+ a, rhs = gram + ridge[:, None, None] * identity, cross + ridge[:, None, None] * identity
303
+ scale = a.diagonal(dim1=-2, dim2=-1).rsqrt()
304
+ conditioned = scale[:, :, None] * a * scale[:, None, :]
305
+ solution = scale[:, :, None] * torch.linalg.solve(conditioned, scale[:, :, None] * rhs)
306
+ core.readout.copy_(solution.to(core.readout))
307
+ residual = (a @ solution - rhs).norm() / rhs.norm().clamp_min(1e-20)
308
+ return {"ridge_relative_residual": float(residual), "uncalibrated_train_mse": before / count}
src/tinycenn_lm/pdelta2_er.py ADDED
@@ -0,0 +1,267 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Error-residual PDelta2 layer for closing the remaining Transformer NLL gap.
2
+
3
+ The layer keeps the strongest previous TinyCeNN direction (PDelta2 F96 + causal
4
+ Conv4) and tests three focused ingredients:
5
+
6
+ 1. learnable per-feature retention initialized to a broad half-life spectrum;
7
+ 2. a small secondary recurrent state trained on the teacher-minus-base residual;
8
+ 3. teacher-error-directed weighting that emphasizes the tokens on which the
9
+ replacement diverges most from exact Transformer attention.
10
+
11
+ The residual memory is deliberately small and sees the original values while the
12
+ main memory sees Conv4 values. Its output gain starts at exactly zero, so adding
13
+ it is function preserving before training. Persistent memory matrices are stored
14
+ as FP16 between streaming calls while curvature remains FP32.
15
+ """
16
+ from __future__ import annotations
17
+
18
+ import math
19
+ from dataclasses import dataclass
20
+
21
+ import torch
22
+ import torch.nn.functional as F
23
+ from torch import Tensor, nn
24
+
25
+ from tinycenn_lm.pdelta2_features import PDelta2Core, PDeltaState
26
+
27
+
28
+ @dataclass
29
+ class ErrorResidualState:
30
+ base: PDeltaState
31
+ residual: PDeltaState | None = None
32
+ conv_tail: Tensor | None = None
33
+
34
+
35
+ def initialize_retention_spectrum(core: PDelta2Core, minimum: float = 8.0,
36
+ maximum: float = 2048.0) -> None:
37
+ """Initialize feature channels with logarithmically spaced decay half-lives."""
38
+ if minimum <= 0 or maximum <= minimum:
39
+ raise ValueError("retention half-life range must satisfy 0 < minimum < maximum")
40
+ with torch.no_grad():
41
+ half_life = torch.exp(torch.linspace(
42
+ math.log(minimum), math.log(maximum), core.feature_dim,
43
+ device=core.forget_b.device, dtype=core.forget_b.dtype,
44
+ ))
45
+ log_decay = math.log(0.5) / half_life
46
+ probability = (-log_decay / 0.25).clamp(1e-5, 1.0 - 1e-5)
47
+ raw = torch.logit(probability)
48
+ core.forget_b.copy_(raw[None].expand(core.num_kv_heads, -1))
49
+ core.forget_w.zero_()
50
+
51
+
52
+ def retention_half_lives(core: PDelta2Core) -> Tensor:
53
+ """Return the bias-only half-life represented by each KV-head/feature channel."""
54
+ with torch.no_grad():
55
+ log_decay = -0.25 * core.forget_b.sigmoid()
56
+ return math.log(0.5) / log_decay.clamp_max(-1e-7)
57
+
58
+
59
+ def teacher_error_weights(prediction: Tensor, target: Tensor, hard_fraction: float = 0.25,
60
+ hard_boost: float = 3.0) -> Tensor:
61
+ """Return normalized token weights that upweight the current hardest tokens.
62
+
63
+ Error is averaged over heads and value dimensions, leaving [batch,time].
64
+ The top ``hard_fraction`` tokens receive ``1 + hard_boost`` weight.
65
+ The returned weights have mean one so the loss scale stays comparable.
66
+ """
67
+ if not 0.0 < hard_fraction < 1.0:
68
+ raise ValueError("hard_fraction must be in (0,1)")
69
+ if hard_boost < 0:
70
+ raise ValueError("hard_boost must be non-negative")
71
+ error = (prediction.detach() - target.detach()).square().mean(dim=(-1, 1))
72
+ threshold = torch.quantile(error, 1.0 - hard_fraction, dim=-1, keepdim=True)
73
+ weights = 1.0 + hard_boost * (error >= threshold).to(error.dtype)
74
+ return weights / weights.mean(dim=-1, keepdim=True).clamp_min(1e-8)
75
+
76
+
77
+ def causal_value_conv_stream(v: Tensor, weight: Tensor | None, tail: Tensor | None = None):
78
+ """Causal depthwise value convolution with a correct streaming prefix tail."""
79
+ if weight is None:
80
+ return v.float(), None
81
+ if weight.ndim != 3 or weight.shape[1] != 1:
82
+ raise ValueError("weight must be [channels,1,kernel]")
83
+ b, h, t, d = v.shape
84
+ kernel = weight.shape[-1]
85
+ channels = h * d
86
+ if weight.shape[0] != channels:
87
+ raise ValueError("convolution channel count does not match values")
88
+ if tail is None:
89
+ tail = v.new_zeros(b, h, kernel - 1, d)
90
+ if tail.shape != (b, h, kernel - 1, d):
91
+ raise ValueError("streaming convolution tail has an incompatible shape")
92
+ history = torch.cat((tail.to(v.dtype), v), dim=2)
93
+ x = history.transpose(1, 2).reshape(b, history.shape[2], channels).transpose(1, 2)
94
+ y = F.conv1d(x.float(), weight.float(), groups=channels)
95
+ y = y.transpose(1, 2).reshape(b, t, h, d).transpose(1, 2)
96
+ new_tail = history[:, :, -(kernel - 1):].clone() if kernel > 1 else None
97
+ return y, new_tail
98
+
99
+
100
+ class ErrorResidualPDelta2Layer(nn.Module):
101
+ """Conv4 PDelta2 with retention spectrum and a compact residual-error memory."""
102
+
103
+ def __init__(self, num_heads: int, num_kv_heads: int, head_dim: int,
104
+ feature_dim: int = 96, residual_dim: int = 0, chunk_size: int = 32,
105
+ conv_kernel: int = 4, retention_spectrum: bool = False,
106
+ retention_min: float = 8.0, retention_max: float = 2048.0,
107
+ state_dtype: str = "fp16"):
108
+ super().__init__()
109
+ if num_heads % num_kv_heads:
110
+ raise ValueError("num_heads must be divisible by num_kv_heads")
111
+ if state_dtype not in {"fp16", "fp32"}:
112
+ raise ValueError("state_dtype must be fp16 or fp32")
113
+ if conv_kernel < 1 or residual_dim < 0:
114
+ raise ValueError("conv_kernel must be positive and residual_dim non-negative")
115
+ self.num_heads = int(num_heads)
116
+ self.num_kv_heads = int(num_kv_heads)
117
+ self.head_dim = int(head_dim)
118
+ self.feature_dim = int(feature_dim)
119
+ self.residual_dim = int(residual_dim)
120
+ self.chunk_size = int(chunk_size)
121
+ self.conv_kernel = int(conv_kernel)
122
+ self.retention_spectrum = bool(retention_spectrum)
123
+ self.retention_min = float(retention_min)
124
+ self.retention_max = float(retention_max)
125
+ self.state_dtype = state_dtype
126
+ self.groups = self.num_heads // self.num_kv_heads
127
+
128
+ self.base = PDelta2Core(
129
+ self.num_heads, self.num_kv_heads, self.head_dim,
130
+ feature_dim=self.feature_dim, chunk_size=self.chunk_size,
131
+ )
132
+ if self.retention_spectrum:
133
+ initialize_retention_spectrum(self.base, self.retention_min, self.retention_max)
134
+
135
+ if self.conv_kernel > 1:
136
+ channels = self.num_kv_heads * self.head_dim
137
+ kernel = torch.zeros(channels, 1, self.conv_kernel)
138
+ kernel[:, 0, -1] = 1.0
139
+ self.conv_weight = nn.Parameter(kernel)
140
+ else:
141
+ self.register_parameter("conv_weight", None)
142
+
143
+ if self.residual_dim:
144
+ self.residual = PDelta2Core(
145
+ self.num_heads, self.num_kv_heads, self.head_dim,
146
+ feature_dim=self.residual_dim, chunk_size=self.chunk_size,
147
+ )
148
+ if self.retention_spectrum:
149
+ initialize_retention_spectrum(
150
+ self.residual,
151
+ max(4.0, self.retention_min / 2.0),
152
+ self.retention_max * 2.0,
153
+ )
154
+ # Zero makes the whole residual branch exactly inactive at initialization.
155
+ self.residual_gain = nn.Parameter(torch.zeros(self.num_heads))
156
+ else:
157
+ self.residual = None
158
+ self.register_parameter("residual_gain", None)
159
+
160
+ @property
161
+ def config(self):
162
+ return {
163
+ "num_heads": self.num_heads,
164
+ "num_kv_heads": self.num_kv_heads,
165
+ "head_dim": self.head_dim,
166
+ "feature_dim": self.feature_dim,
167
+ "residual_dim": self.residual_dim,
168
+ "chunk_size": self.chunk_size,
169
+ "conv_kernel": self.conv_kernel,
170
+ "retention_spectrum": self.retention_spectrum,
171
+ "retention_min": self.retention_min,
172
+ "retention_max": self.retention_max,
173
+ "state_dtype": self.state_dtype,
174
+ }
175
+
176
+ @property
177
+ def storage_dtype(self):
178
+ return torch.float16 if self.state_dtype == "fp16" else torch.float32
179
+
180
+ def _working_state(self, state: PDeltaState | None, core: PDelta2Core):
181
+ if state is None:
182
+ return None
183
+ dtype = core.wq.dtype
184
+ return PDeltaState(state.memory.to(dtype), state.curvature.to(dtype))
185
+
186
+ def _stored_state(self, state: PDeltaState):
187
+ return PDeltaState(
188
+ state.memory.to(self.storage_dtype),
189
+ state.curvature.float(),
190
+ )
191
+
192
+ def _run(self, q: Tensor, k: Tensor, v: Tensor, state: ErrorResidualState | None):
193
+ tail = None if state is None else state.conv_tail
194
+ conv_v, new_tail = causal_value_conv_stream(v.float(), self.conv_weight, tail)
195
+ base_state = None if state is None else state.base
196
+ base_out, new_base = self.base(
197
+ q, k, conv_v,
198
+ state=self._working_state(base_state, self.base),
199
+ return_state=True,
200
+ )
201
+
202
+ residual_raw = None
203
+ new_residual = None
204
+ output = base_out
205
+ if self.residual is not None:
206
+ residual_state = None if state is None else state.residual
207
+ residual_raw, residual_state_out = self.residual(
208
+ q, k, v.float(),
209
+ state=self._working_state(residual_state, self.residual),
210
+ return_state=True,
211
+ )
212
+ gain = self.residual_gain.clamp(-1.5, 1.5)[None, :, None, None]
213
+ output = base_out + gain * residual_raw
214
+ new_residual = self._stored_state(residual_state_out)
215
+
216
+ new_state = ErrorResidualState(
217
+ base=self._stored_state(new_base),
218
+ residual=new_residual,
219
+ conv_tail=(None if new_tail is None else new_tail.to(self.storage_dtype)),
220
+ )
221
+ components = {
222
+ "base": base_out,
223
+ "residual_raw": residual_raw,
224
+ "output": output,
225
+ }
226
+ return output, new_state, components
227
+
228
+ def components(self, q: Tensor, k: Tensor, v: Tensor,
229
+ state: ErrorResidualState | None = None):
230
+ return self._run(q, k, v, state)
231
+
232
+ def forward(self, q: Tensor, k: Tensor, v: Tensor,
233
+ state: ErrorResidualState | None = None,
234
+ return_state: bool = False, implementation: str = "chunk"):
235
+ if implementation != "chunk":
236
+ raise ValueError("ErrorResidualPDelta2Layer supports the chunk implementation")
237
+ output, new_state, _ = self._run(q, k, v, state)
238
+ return (output, new_state) if return_state else output
239
+
240
+ def recurrent_state_bytes(self, batch_size: int = 1, context: int | None = None):
241
+ del context
242
+ memory_bytes = 2 if self.state_dtype == "fp16" else 4
243
+ base_memory = self.num_kv_heads * self.feature_dim * self.head_dim * memory_bytes
244
+ base_curvature = self.num_kv_heads * self.feature_dim * 4
245
+ total = base_memory + base_curvature
246
+ if self.residual_dim:
247
+ total += self.num_kv_heads * self.residual_dim * self.head_dim * memory_bytes
248
+ total += self.num_kv_heads * self.residual_dim * 4
249
+ if self.conv_kernel > 1:
250
+ total += (self.conv_kernel - 1) * self.num_kv_heads * self.head_dim * memory_bytes
251
+ return batch_size * total
252
+
253
+ def retention_statistics(self):
254
+ values = retention_half_lives(self.base).float().reshape(-1)
255
+ result = {
256
+ "base_half_life_min": float(values.min()),
257
+ "base_half_life_median": float(values.median()),
258
+ "base_half_life_max": float(values.max()),
259
+ }
260
+ if self.residual is not None:
261
+ rv = retention_half_lives(self.residual).float().reshape(-1)
262
+ result.update(
263
+ residual_half_life_min=float(rv.min()),
264
+ residual_half_life_median=float(rv.median()),
265
+ residual_half_life_max=float(rv.max()),
266
+ )
267
+ return result