Text Generation
Transformers
Safetensors
qwen3_5_text
tinycenn
cenn
language-modeling
research
conversational
Instructions to use vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32") messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoTokenizer, AutoModelForCausalLM tokenizer = AutoTokenizer.from_pretrained("vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32") model = AutoModelForCausalLM.from_pretrained("vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32", device_map="auto") messages = [ {"role": "user", "content": "Who are you?"}, ] inputs = tokenizer.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ).to(model.device) outputs = model.generate(**inputs, max_new_tokens=40) print(tokenizer.decode(outputs[0][inputs["input_ids"].shape[-1]:])) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32
- SGLang
How to use vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32 with Docker Model Runner:
docker model run hf.co/vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32
Upload verified PDelta3-CLVR checkpoint for layers [3, 7, 11]
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +1 -0
- README.md +45 -0
- chat_template.jinja +154 -0
- config.json +75 -0
- generation_config.json +6 -0
- load_model.py +33 -0
- model.safetensors +3 -0
- qwen35_progress.json +57 -0
- qwen35_run_status.json +38 -0
- qwen35_verification.json +80 -0
- requirements.txt +4 -0
- run_manifest.json +16 -0
- scripts/train_qwen35_pdelta3_clvr_sequential.py +491 -0
- scripts/train_smollm2_memory_fusion_sequential.py +986 -0
- src/tinycenn_lm/__init__.py +144 -0
- src/tinycenn_lm/__pycache__/__init__.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/cellular_attention.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/cenn.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/colab_live_backup.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/direct_colab_backup.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/hf_persistence.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/live_console.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/memory_attention.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/modeling.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/moe.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/pdelta2_er.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/pdelta2_features.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/pdelta3_frontier.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/research_layers.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/sharded_moe.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/smollm2_amcenn.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/smollm2_amcenn_v2.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/smollm2_memory_fusion.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/story_v2.cpython-313.pyc +0 -0
- src/tinycenn_lm/__pycache__/student.cpython-313.pyc +0 -0
- src/tinycenn_lm/cellular_attention.py +572 -0
- src/tinycenn_lm/cenn.py +156 -0
- src/tinycenn_lm/colab_live_backup.py +459 -0
- src/tinycenn_lm/direct_colab_backup.py +276 -0
- src/tinycenn_lm/distill_utils.py +232 -0
- src/tinycenn_lm/gemma3_integrated_memory.py +319 -0
- src/tinycenn_lm/gemma3_memory_fusion.py +255 -0
- src/tinycenn_lm/hf_persistence.py +411 -0
- src/tinycenn_lm/integrated_memory.py +186 -0
- src/tinycenn_lm/live_console.py +90 -0
- src/tinycenn_lm/memory_attention.py +473 -0
- src/tinycenn_lm/modeling.py +251 -0
- src/tinycenn_lm/moe.py +334 -0
- src/tinycenn_lm/optimized_memory.py +308 -0
- 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
|