diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..52373fe24473b1aa44333d318f578ae6bf04b49b 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +tokenizer.json filter=lfs diff=lfs merge=lfs -text diff --git a/README.md b/README.md new file mode 100644 index 0000000000000000000000000000000000000000..ce2a642901ce1e5e94d0aeb148e0a040dcc44ec0 --- /dev/null +++ b/README.md @@ -0,0 +1,45 @@ +--- +library_name: transformers +pipeline_tag: text-generation +tags: +- tinycenn +- cenn +- language-modeling +- text-generation +- research +--- + +# Qwen3.5-0.8B-PDelta3-CLVR-Local32 + +Research artifact from **TinyCeNN-LM**. Architecture: `TinyCeNN-LM experiment`. + +## Architecture + +- Architecture/run type: `TinyCeNN-LM experiment` +- Base model: `not recorded` +- Dataset: `Not recorded` +- Source code: https://github.com/vtavakkoli/TinyCeNN-LM + +## Latest saved results + +No structured training report was found in this upload. + +The Hugging Face repository keeps timestamped run artifacts under `runs/`. This preserves training reports, configs and run metadata independently of the temporary Colab filesystem. + +## Saved experiment files + +- `config.json` +- `generation_config.json` +- `tokenizer_config.json` + +## Reproducibility + +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. + +## Limitations + +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. + +## Citation + +If you use this experimental checkpoint, cite the TinyCeNN-LM repository and the upstream base model. diff --git a/chat_template.jinja b/chat_template.jinja new file mode 100644 index 0000000000000000000000000000000000000000..0ef09f214eaa6d9bca297988afc1454b5827b2c7 --- /dev/null +++ b/chat_template.jinja @@ -0,0 +1,154 @@ +{%- set image_count = namespace(value=0) %} +{%- set video_count = namespace(value=0) %} +{%- macro render_content(content, do_vision_count, is_system_content=false) %} + {%- if content is string %} + {{- content }} + {%- elif content is iterable and content is not mapping %} + {%- for item in content %} + {%- if 'image' in item or 'image_url' in item or item.type == 'image' %} + {%- if is_system_content %} + {{- raise_exception('System message cannot contain images.') }} + {%- endif %} + {%- if do_vision_count %} + {%- set image_count.value = image_count.value + 1 %} + {%- endif %} + {%- if add_vision_id %} + {{- 'Picture ' ~ image_count.value ~ ': ' }} + {%- endif %} + {{- '<|vision_start|><|image_pad|><|vision_end|>' }} + {%- elif 'video' in item or item.type == 'video' %} + {%- if is_system_content %} + {{- raise_exception('System message cannot contain videos.') }} + {%- endif %} + {%- if do_vision_count %} + {%- set video_count.value = video_count.value + 1 %} + {%- endif %} + {%- if add_vision_id %} + {{- 'Video ' ~ video_count.value ~ ': ' }} + {%- endif %} + {{- '<|vision_start|><|video_pad|><|vision_end|>' }} + {%- elif 'text' in item %} + {{- item.text }} + {%- else %} + {{- raise_exception('Unexpected item type in content.') }} + {%- endif %} + {%- endfor %} + {%- elif content is none or content is undefined %} + {{- '' }} + {%- else %} + {{- raise_exception('Unexpected content type.') }} + {%- endif %} +{%- endmacro %} +{%- if not messages %} + {{- raise_exception('No messages provided.') }} +{%- endif %} +{%- if tools and tools is iterable and tools is not mapping %} + {{- '<|im_start|>system\n' }} + {{- "# Tools\n\nYou have access to the following functions:\n\n" }} + {%- for tool in tools %} + {{- "\n" }} + {{- tool | tojson }} + {%- endfor %} + {{- "\n" }} + {{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n\n\n\nvalue_1\n\n\nThis is the value for the second parameter\nthat can span\nmultiple lines\n\n\n\n\n\nReminder:\n- Function calls MUST follow the specified format: an inner block must be nested within 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' }} + {%- if messages[0].role == 'system' %} + {%- set content = render_content(messages[0].content, false, true)|trim %} + {%- if content %} + {{- '\n\n' + content }} + {%- endif %} + {%- endif %} + {{- '<|im_end|>\n' }} +{%- else %} + {%- if messages[0].role == 'system' %} + {%- set content = render_content(messages[0].content, false, true)|trim %} + {{- '<|im_start|>system\n' + content + '<|im_end|>\n' }} + {%- endif %} +{%- endif %} +{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %} +{%- for message in messages[::-1] %} + {%- set index = (messages|length - 1) - loop.index0 %} + {%- if ns.multi_step_tool and message.role == "user" %} + {%- set content = render_content(message.content, false)|trim %} + {%- if not(content.startswith('') and content.endswith('')) %} + {%- set ns.multi_step_tool = false %} + {%- set ns.last_query_index = index %} + {%- endif %} + {%- endif %} +{%- endfor %} +{%- if ns.multi_step_tool %} + {{- raise_exception('No user query found in messages.') }} +{%- endif %} +{%- for message in messages %} + {%- set content = render_content(message.content, true)|trim %} + {%- if message.role == "system" %} + {%- if not loop.first %} + {{- raise_exception('System message must be at the beginning.') }} + {%- endif %} + {%- elif message.role == "user" %} + {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }} + {%- elif message.role == "assistant" %} + {%- set reasoning_content = '' %} + {%- if message.reasoning_content is string %} + {%- set reasoning_content = message.reasoning_content %} + {%- else %} + {%- if '' in content %} + {%- set reasoning_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') %} + {%- set content = content.split('')[-1].lstrip('\n') %} + {%- endif %} + {%- endif %} + {%- set reasoning_content = reasoning_content|trim %} + {%- if loop.index0 > ns.last_query_index %} + {{- '<|im_start|>' + message.role + '\n\n' + reasoning_content + '\n\n\n' + content }} + {%- else %} + {{- '<|im_start|>' + message.role + '\n' + content }} + {%- endif %} + {%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %} + {%- for tool_call in message.tool_calls %} + {%- if tool_call.function is defined %} + {%- set tool_call = tool_call.function %} + {%- endif %} + {%- if loop.first %} + {%- if content|trim %} + {{- '\n\n\n\n' }} + {%- else %} + {{- '\n\n' }} + {%- endif %} + {%- else %} + {{- '\n\n\n' }} + {%- endif %} + {%- if tool_call.arguments is defined %} + {%- for args_name, args_value in tool_call.arguments|items %} + {{- '\n' }} + {%- 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 %} + {{- args_value }} + {{- '\n\n' }} + {%- endfor %} + {%- endif %} + {{- '\n' }} + {%- endfor %} + {%- endif %} + {{- '<|im_end|>\n' }} + {%- elif message.role == "tool" %} + {%- if loop.previtem and loop.previtem.role != "tool" %} + {{- '<|im_start|>user' }} + {%- endif %} + {{- '\n\n' }} + {{- content }} + {{- '\n' }} + {%- if not loop.last and loop.nextitem.role != "tool" %} + {{- '<|im_end|>\n' }} + {%- elif loop.last %} + {{- '<|im_end|>\n' }} + {%- endif %} + {%- else %} + {{- raise_exception('Unexpected message role.') }} + {%- endif %} +{%- endfor %} +{%- if add_generation_prompt %} + {{- '<|im_start|>assistant\n' }} + {%- if enable_thinking is defined and enable_thinking is true %} + {{- '\n' }} + {%- else %} + {{- '\n\n\n\n' }} + {%- endif %} +{%- endif %} \ No newline at end of file diff --git a/config.json b/config.json new file mode 100644 index 0000000000000000000000000000000000000000..dd30a41c12217d7a5aaff97a20891c33b4471d78 --- /dev/null +++ b/config.json @@ -0,0 +1,75 @@ +{ + "architectures": [ + "Qwen3_5ForCausalLM" + ], + "attention_bias": false, + "attention_dropout": 0.0, + "attn_output_gate": true, + "bos_token_id": null, + "dtype": "bfloat16", + "eos_token_id": 248044, + "full_attention_interval": 4, + "head_dim": 256, + "hidden_act": "silu", + "hidden_size": 1024, + "initializer_range": 0.02, + "intermediate_size": 3584, + "layer_types": [ + "linear_attention", + "linear_attention", + "linear_attention", + "full_attention", + "linear_attention", + "linear_attention", + "linear_attention", + "full_attention", + "linear_attention", + "linear_attention", + "linear_attention", + "full_attention", + "linear_attention", + "linear_attention", + "linear_attention", + "full_attention", + "linear_attention", + "linear_attention", + "linear_attention", + "full_attention", + "linear_attention", + "linear_attention", + "linear_attention", + "full_attention" + ], + "linear_conv_kernel_dim": 4, + "linear_key_head_dim": 128, + "linear_num_key_heads": 16, + "linear_num_value_heads": 16, + "linear_value_head_dim": 128, + "mamba_ssm_dtype": "float32", + "max_position_embeddings": 262144, + "mlp_only_layers": [], + "model_type": "qwen3_5_text", + "mtp_num_hidden_layers": 1, + "mtp_use_dedicated_embeddings": false, + "num_attention_heads": 8, + "num_hidden_layers": 24, + "num_key_value_heads": 2, + "pad_token_id": null, + "partial_rotary_factor": 0.25, + "rms_norm_eps": 1e-06, + "rope_parameters": { + "mrope_interleaved": true, + "mrope_section": [ + 11, + 11, + 10 + ], + "partial_rotary_factor": 0.25, + "rope_theta": 10000000, + "rope_type": "default" + }, + "tie_word_embeddings": true, + "transformers_version": "5.17.0", + "use_cache": false, + "vocab_size": 248320 +} diff --git a/generation_config.json b/generation_config.json new file mode 100644 index 0000000000000000000000000000000000000000..31ec2a53cc9cde720b0a8fbc68b9c4e73ac5a18f --- /dev/null +++ b/generation_config.json @@ -0,0 +1,6 @@ +{ + "_from_model_config": true, + "eos_token_id": 248044, + "transformers_version": "5.17.0", + "use_cache": true +} diff --git a/load_model.py b/load_model.py new file mode 100644 index 0000000000000000000000000000000000000000..8335005afc7bffa51cace825d8e284c8a6ac2f3d --- /dev/null +++ b/load_model.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +import json +import sys +from pathlib import Path +import torch +from safetensors.torch import load_file +from transformers import AutoTokenizer, Qwen3_5ForCausalLM + +def load_model(path=None, device="cpu", dtype=None): + root = Path(path or Path(__file__).resolve().parent) + for p in (root / "src", root, root / "scripts"): + if str(p) not in sys.path: + sys.path.insert(0, str(p)) + from train_qwen35_pdelta3_clvr_sequential import QwenPDelta3CLVRConfig, replace_full_attention_layers + meta = json.loads((root / "tinycenn_qwen35.json").read_text()) + accepted = [int(x) for x in meta["accepted_full_attention_layers"]] + cfg = QwenPDelta3CLVRConfig.from_dict(meta["replacement_config"]) + if dtype is None: + dtype = torch.bfloat16 if device.startswith("cuda") and torch.cuda.is_bf16_supported() else (torch.float16 if device.startswith("cuda") else torch.float32) + model = Qwen3_5ForCausalLM.from_pretrained(root, dtype=dtype, local_files_only=True, attn_implementation="eager") + replace_full_attention_layers(model, cfg, accepted) + single = root / "model.safetensors" + if single.exists(): + model.load_state_dict(load_file(str(single), device="cpu"), strict=False) + else: + index = json.loads((root / "model.safetensors.index.json").read_text()) + for shard in sorted(set(index["weight_map"].values())): + model.load_state_dict(load_file(str(root / shard), device="cpu"), strict=False) + model.config.use_cache = False + model.to(device).eval() + tokenizer = AutoTokenizer.from_pretrained(root, local_files_only=True, use_fast=True) + return model, tokenizer diff --git a/model.safetensors b/model.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c65d9eb8544c8a8bac32b5ecc6420d20e2a62e42 --- /dev/null +++ b/model.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fd2298e1a9cadb70609cd7fbe6a91cf99776c61406520e47556c643ac818c229 +size 1513872992 diff --git a/qwen35_progress.json b/qwen35_progress.json new file mode 100644 index 0000000000000000000000000000000000000000..ece1f7a79664fa3c835ae5b2309ebec4c493a8c0 --- /dev/null +++ b/qwen35_progress.json @@ -0,0 +1,57 @@ +{ + "format_version": 1, + "architecture": "qwen3.5-pdelta3-gdn2-clvr-localw", + "accepted_full_attention_layers": [ + 3, + 7, + 11 + ], + "config": { + "feature_dim": 96, + "local_window": 32, + "chunk_size": 32, + "conv_kernel": 4, + "state_dtype": "fp16", + "variant": "conv4_gdn2_clvr_f96", + "local_gate_init": 0.72, + "warm_start_previous_core": true + }, + "reports": [ + { + "layer": 3, + "nmse": 0.062080949544906616, + "cosine": 0.9744633436203003, + "probe_nll": 2.8579028447469077, + "incremental_delta_nll": -0.003246148427327178, + "cumulative_delta_nll": -0.003246148427327178, + "local_gate_mean": 0.7206810712814331, + "step": 75, + "round": 1, + "accepted": true + }, + { + "layer": 7, + "nmse": 0.03889689967036247, + "cosine": 0.9704566597938538, + "probe_nll": 2.8711770375569663, + "incremental_delta_nll": 0.013274192810058594, + "cumulative_delta_nll": 0.010028044382731416, + "local_gate_mean": 0.7213350534439087, + "step": 150, + "round": 1, + "accepted": true + }, + { + "layer": 11, + "nmse": 0.10417895764112473, + "cosine": 0.9411033391952515, + "probe_nll": 2.8818757136662803, + "incremental_delta_nll": 0.010698676109313965, + "cumulative_delta_nll": 0.02072672049204538, + "local_gate_mean": 0.7214324474334717, + "step": 75, + "round": 1, + "accepted": true + } + ] +} \ No newline at end of file diff --git a/qwen35_run_status.json b/qwen35_run_status.json new file mode 100644 index 0000000000000000000000000000000000000000..1ea38ddaf9aad875fde2128bb826aca2aaa52e09 --- /dev/null +++ b/qwen35_run_status.json @@ -0,0 +1,38 @@ +{ + "status": "target_full_attention_prefix_accepted", + "architecture": "Qwen3.5-PDelta3-GDN2-CLVR+LocalW", + "base_model": "Qwen/Qwen3.5-0.8B", + "native_full_attention_layers": [ + 3, + 7, + 11, + 15, + 19, + 23 + ], + "target_full_attention_layers": [ + 3, + 7, + 11 + ], + "accepted_full_attention_layers": [ + 3, + 7, + 11 + ], + "teacher_probe_nll": 2.861148993174235, + "final_probe_nll": 2.8818757136662803, + "final_delta_nll": 0.02072672049204538, + "config": { + "feature_dim": 96, + "local_window": 32, + "chunk_size": 32, + "conv_kernel": 4, + "state_dtype": "fp16", + "variant": "conv4_gdn2_clvr_f96", + "local_gate_init": 0.72, + "warm_start_previous_core": true + }, + "elapsed_minutes": 2.422297354216666, + "peak_vram_gib": 4.466190814971924 +} \ No newline at end of file diff --git a/qwen35_verification.json b/qwen35_verification.json new file mode 100644 index 0000000000000000000000000000000000000000..678289986869fcbedd5cad2788db7ad429cfd5bd --- /dev/null +++ b/qwen35_verification.json @@ -0,0 +1,80 @@ +{ + "verified": true, + "base_model": "Qwen/Qwen3.5-0.8B", + "accepted_full_attention_layers": [ + 3, + 7, + 11 + ], + "config": { + "feature_dim": 96, + "local_window": 32, + "chunk_size": 32, + "conv_kernel": 4, + "state_dtype": "fp16", + "variant": "conv4_gdn2_clvr_f96", + "local_gate_init": 0.72, + "warm_start_previous_core": true + }, + "probe_context": 128, + "probe_blocks": 6, + "baseline_probe_nll": 2.861148993174235, + "candidate_probe_nll": 2.8818757136662803, + "delta_nll": 0.02072672049204538, + "release_cumulative_delta_nll_limit": 0.05, + "quality_pass": true, + "all_saved_layer_gates_pass": true, + "layer_checks": [ + { + "layer": 3, + "pass": true, + "nmse": 0.062080949544906616, + "cosine": 0.9744633436203003, + "incremental_delta_nll": -0.003246148427327178, + "cumulative_delta_nll": -0.003246148427327178 + }, + { + "layer": 7, + "pass": true, + "nmse": 0.03889689967036247, + "cosine": 0.9704566597938538, + "incremental_delta_nll": 0.013274192810058594, + "cumulative_delta_nll": 0.010028044382731416 + }, + { + "layer": 11, + "pass": true, + "nmse": 0.10417895764112473, + "cosine": 0.9411033391952515, + "incremental_delta_nll": 0.010698676109313965, + "cumulative_delta_nll": 0.02072672049204538 + } + ], + "prompt_examples": [ + { + "prompt": "The future of small language models is", + "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.", + "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" + }, + { + "prompt": "Artificial intelligence can help scientists by", + "baseline": "Artificial intelligence can help scientists by providing a way to predict the future.\n\n\n\n\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", + "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\n\n\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" + }, + { + "prompt": "A good software architecture should", + "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.", + "pdelta3_clvr": "A good software architecture should be able to handle a large number of concurrent users.\nA. True\nB. False\n\n\n\n\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" + }, + { + "prompt": "The capital of Austria is", + "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", + "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" + }, + { + "prompt": "Once upon a time, a small robot", + "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", + "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" + } + ] +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..03a8f6e763db01512de3d5f8d5000d505c77ea62 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +torch +transformers==5.17.0 +datasets>=3,<5 +safetensors diff --git a/run_manifest.json b/run_manifest.json new file mode 100644 index 0000000000000000000000000000000000000000..898a1963e7db197caf436db80180c08cd02c9b6d --- /dev/null +++ b/run_manifest.json @@ -0,0 +1,16 @@ +{ + "run_id": "qwen35_pdelta3_clvr_hf_e-20260915T234639Z", + "created_utc": "2026-09-15T23:46:39.791454+00:00", + "repo_id": "vtava/Qwen3.5-0.8B-PDelta3-CLVR-Local32", + "notebook": null, + "python": "3.13.15", + "platform": "Linux-6.6.122+-x86_64-with-glibc2.39", + "reports": [ + "config.json", + "generation_config.json", + "tokenizer_config.json" + ], + "torch": "2.11.0+cu128", + "cuda_available": true, + "gpu": "Tesla T4" +} \ No newline at end of file diff --git a/scripts/train_qwen35_pdelta3_clvr_sequential.py b/scripts/train_qwen35_pdelta3_clvr_sequential.py new file mode 100644 index 0000000000000000000000000000000000000000..cc1dd5de4a0b0f790e888df2cc534a64d98a14a9 --- /dev/null +++ b/scripts/train_qwen35_pdelta3_clvr_sequential.py @@ -0,0 +1,491 @@ +#!/usr/bin/env python3 +"""Sequential PDelta3/GDN2-CLVR replacement of Qwen3.5 full-attention layers. + +Qwen3.5 already uses a 3:1 hybrid stack: most layers are Gated DeltaNet +linear-attention layers and only every fourth layer is full attention. This +experiment leaves native linear-attention layers untouched and replaces only +full-attention layers, one at a time, with PDelta3/GDN2 + bounded Local-W + +cross-layer value routing. +""" +from __future__ import annotations + +import argparse, copy, json, math, os, random, sys, time, weakref +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Any + +os.environ.setdefault("HF_HUB_DISABLE_IMPLICIT_TOKEN", "1") +os.environ.setdefault("HF_HUB_DISABLE_TELEMETRY", "1") + +import torch +import torch.nn.functional as F +from datasets import load_dataset +from torch import Tensor, nn +from transformers import AutoTokenizer, Qwen3_5ForCausalLM +from transformers.models.qwen3_5.modeling_qwen3_5 import apply_rotary_pos_emb + +REPO_ROOT = Path(__file__).resolve().parents[1] +SRC_ROOT = REPO_ROOT / "src" +for p in (str(SRC_ROOT), str(REPO_ROOT), str(REPO_ROOT / "scripts")): + if p not in sys.path: + sys.path.insert(0, p) + +import train_smollm2_memory_fusion_sequential as seq +from tinycenn_lm.pdelta3_frontier import FrontierPDelta3Layer + + +@dataclass(frozen=True) +class QwenPDelta3CLVRConfig: + feature_dim: int = 96 + local_window: int = 32 + chunk_size: int = 32 + conv_kernel: int = 4 + state_dtype: str = "fp16" + variant: str = "conv4_gdn2_clvr_f96" + local_gate_init: float = 0.72 + warm_start_previous_core: bool = True + + def validate(self, cfg): + if self.feature_dim < 16 or self.local_window < 1 or self.conv_kernel < 1: + raise ValueError("invalid PDelta3-CLVR dimensions") + if not 1 <= self.chunk_size <= 32: + raise ValueError("chunk_size must be in [1,32]") + if self.state_dtype not in {"fp16", "fp32"}: + raise ValueError("state_dtype must be fp16 or fp32") + if int(cfg.num_attention_heads) % int(cfg.num_key_value_heads): + raise ValueError("attention heads must be divisible by KV heads") + + def to_dict(self): + return asdict(self) + + @classmethod + def from_dict(cls, value): + return cls(**dict(value)) + + +def _text_config(model): + return getattr(model.config, "text_config", model.config) + + +def full_attention_layers(model): + kinds = list(getattr(_text_config(model), "layer_types", [])) + if not kinds: + raise RuntimeError("Qwen3.5 config does not expose layer_types") + return [i for i, kind in enumerate(kinds) if kind == "full_attention"] + + +def _repeat_kv(x, groups): + return x.repeat_interleave(groups, dim=1) + + +class QwenPDelta3CLVRAttention(nn.Module): + """Qwen3.5 Gated Attention replacement: Local-W + recurrent PDelta3/GDN2.""" + def __init__(self, original, cfg, config, layer_idx, previous_attention=None): + super().__init__() + config.validate(cfg) + self.config = cfg + self.layer_idx = int(layer_idx) + self.hidden_size = int(cfg.hidden_size) + self.num_heads = int(cfg.num_attention_heads) + self.num_kv_heads = int(cfg.num_key_value_heads) + self.head_dim = int(getattr(cfg, "head_dim", self.hidden_size // self.num_heads)) + self.groups = self.num_heads // self.num_kv_heads + self.scaling = self.head_dim ** -0.5 + self.attention_dropout = float(getattr(cfg, "attention_dropout", 0.0)) + self.is_causal = True + self.local_window = int(config.local_window) + object.__setattr__(self, "_previous_attention_ref", weakref.ref(previous_attention) if previous_attention is not None else None) + + # Qwen3.5 q_proj is doubled: query + post-attention output gate. + self.q_proj = copy.deepcopy(original.q_proj) + self.k_proj = copy.deepcopy(original.k_proj) + self.v_proj = copy.deepcopy(original.v_proj) + self.o_proj = copy.deepcopy(original.o_proj) + self.q_norm = copy.deepcopy(original.q_norm) + self.k_norm = copy.deepcopy(original.k_norm) + + self.core = FrontierPDelta3Layer( + self.num_heads, self.num_kv_heads, self.head_dim, + feature_dim=config.feature_dim, variant=config.variant, + chunk_size=config.chunk_size, conv_kernel=config.conv_kernel, + state_dtype=config.state_dtype, + ) + init = min(max(float(config.local_gate_init), 1e-4), 1 - 1e-4) + logit = math.log(init / (1 - init)) + self.local_gate_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim)) + self.local_gate_b = nn.Parameter(torch.full((self.num_heads,), logit)) + self.last_value = None + + def _previous_attention(self): + ref = object.__getattribute__(self, "_previous_attention_ref") + return None if ref is None else ref() + + def _local_attention(self, q, k, v, attention_mask): + kh, vh = _repeat_kv(k, self.groups), _repeat_kv(v, self.groups) + scores = torch.matmul(q.float(), kh.float().transpose(-2, -1)) * self.scaling + t = q.shape[-2] + qi = torch.arange(t, device=q.device)[:, None] + kj = torch.arange(t, device=q.device)[None, :] + allowed = (kj <= qi) & (kj >= qi - self.local_window + 1) + bias = torch.zeros((t, t), device=q.device, dtype=scores.dtype) + bias.masked_fill_(~allowed, torch.finfo(scores.dtype).min) + scores = scores + bias[None, None] + if attention_mask is not None: + if attention_mask.ndim == 4: + scores = scores + attention_mask[..., :t, :t].float() + elif attention_mask.ndim == 2: + key_mask = 1.0 - attention_mask[:, None, None, :t].float() + scores = scores + key_mask * torch.finfo(scores.dtype).min + probs = torch.softmax(scores, dim=-1, dtype=torch.float32).to(q.dtype) + if self.training and self.attention_dropout: + probs = F.dropout(probs, p=self.attention_dropout) + return torch.matmul(probs, vh.to(probs.dtype)) + + def forward(self, hidden_states, position_embeddings, attention_mask=None, past_key_values=None, **kwargs): + if past_key_values is not None: + raise ValueError("research replacement requires use_cache=False") + shape = hidden_states.shape[:-1] + qg = self.q_proj(hidden_states).view(*shape, self.num_heads, self.head_dim * 2) + q, out_gate = torch.chunk(qg, 2, dim=-1) + out_gate = out_gate.reshape(*shape, self.num_heads * self.head_dim) + q = self.q_norm(q).transpose(1, 2) + k = self.k_norm(self.k_proj(hidden_states).view(*shape, self.num_kv_heads, self.head_dim)).transpose(1, 2) + v = self.v_proj(hidden_states).view(*shape, self.num_kv_heads, self.head_dim).transpose(1, 2) + cos, sin = position_embeddings + q, k = apply_rotary_pos_emb(q, k, cos, sin) + + self.last_value = v.detach() + previous = self._previous_attention() + routed_v = None + if previous is not None and previous.last_value is not None and previous.last_value.shape == v.shape: + routed_v = previous.last_value.to(device=v.device, dtype=v.dtype) + if routed_v is None: + routed_v = v + + global_out = self.core(q, k, v, routed_v=routed_v) + local_out = self._local_attention(q, k, v, attention_mask) + 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) + mixed = gate[..., None] * local_out + (1 - gate[..., None]) * global_out + out = mixed.transpose(1, 2).contiguous().reshape(*shape, self.num_heads * self.head_dim) + out = out * torch.sigmoid(out_gate) + return self.o_proj(out.to(hidden_states.dtype)), None + + @torch.no_grad() + def local_gate_mean(self): + return float(torch.sigmoid(self.local_gate_b.float()).mean()) + + +def replace_full_attention_layers(model, config, indices): + cfg = _text_config(model) + allowed = set(full_attention_layers(model)) + previous = None + made = [] + for idx in sorted(indices): + if idx not in allowed: + raise ValueError(f"layer {idx} is not full attention") + layer = model.model.layers[idx] + if isinstance(layer.self_attn, QwenPDelta3CLVRAttention): + wrapper = layer.self_attn + object.__setattr__(wrapper, "_previous_attention_ref", weakref.ref(previous) if previous is not None else None) + else: + wrapper = QwenPDelta3CLVRAttention(layer.self_attn, cfg, config, idx, previous) + layer.self_attn = wrapper + previous = wrapper + made.append(wrapper) + return made + + +def _prefix(idx): + return f"model.layers.{idx}.self_attn." + + +def selected_state(model, indices): + prefixes = tuple(_prefix(i) for i in indices) + return {k: v.detach().cpu() for k, v in model.state_dict().items() if prefixes and k.startswith(prefixes)} + + +def atomic_torch(payload, path): + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + torch.save(payload, tmp); tmp.replace(path) + + +def atomic_json(payload, path): + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + tmp.write_text(json.dumps(payload, indent=2), encoding="utf-8"); tmp.replace(path) + + +def _load_selected(model, state, layers): + inc = model.load_state_dict(state, strict=False) + prefixes = tuple(_prefix(i) for i in layers) + missing = [k for k in inc.missing_keys if prefixes and k.startswith(prefixes)] + if missing: + raise RuntimeError(f"checkpoint missing keys: {missing[:8]}") + + +def save_progress(out, model, config, accepted, reports): + payload = { + "format_version": 1, + "architecture": "qwen3.5-pdelta3-gdn2-clvr-localw", + "accepted_full_attention_layers": list(accepted), + "config": config.to_dict(), "reports": list(reports), + "attention_state": selected_state(model, accepted), + } + atomic_torch(payload, out / "qwen35_progress.pt") + atomic_json({k:v for k,v in payload.items() if k != "attention_state"}, out / "qwen35_progress.json") + + +def save_in_progress(out, model, config, accepted, current, rounds, pre_nll, reports, best): + layers = accepted + [current] + payload = { + "format_version": 1, "status": "current_layer_needs_more_training", + "accepted_full_attention_layers": list(accepted), + "current_full_attention_layer": int(current), "rounds_completed": int(rounds), + "pre_probe_nll": float(pre_nll), "config": config.to_dict(), + "reports": list(reports), "best_report": dict(best), + "attention_state": selected_state(model, layers), + } + atomic_torch(payload, out / "qwen35_in_progress.pt") + atomic_json({k:v for k,v in payload.items() if k != "attention_state"}, out / "qwen35_in_progress.json") + + +def load_progress(out, model, config, resume): + path = out / "qwen35_progress.pt" + if not resume or not path.exists(): return [], [] + p = torch.load(path, map_location="cpu", weights_only=False) + accepted = [int(x) for x in p.get("accepted_full_attention_layers", [])] + if QwenPDelta3CLVRConfig.from_dict(p["config"]) != config: + raise RuntimeError("resume architecture differs from saved config") + if accepted: + replace_full_attention_layers(model, config, accepted) + _load_selected(model, p["attention_state"], accepted) + print(f"RESUME accepted full-attention layers={accepted}", flush=True) + return accepted, list(p.get("reports", [])) + + +def load_in_progress(out, model, config, accepted, targets, resume): + path = out / "qwen35_in_progress.pt" + if not resume or not path.exists(): return None + p = torch.load(path, map_location="cpu", weights_only=False) + if [int(x) for x in p.get("accepted_full_attention_layers", [])] != accepted: return None + expected = targets[len(accepted)] + if int(p.get("current_full_attention_layer", -1)) != expected: return None + layers = accepted + [expected] + replace_full_attention_layers(model, config, layers) + _load_selected(model, p["attention_state"], layers) + print(f"RESUME current layer={expected}, rounds={p.get('rounds_completed',0)}", flush=True) + return p + + +def remove_in_progress(out): + for name in ("qwen35_in_progress.pt", "qwen35_in_progress.json"): + p = out / name + if p.exists(): p.unlink() + + +def warm_start(model, current, accepted): + if not accepted: return + prev = model.model.layers[accepted[-1]].self_attn + cur = model.model.layers[current].self_attn + if isinstance(prev, QwenPDelta3CLVRAttention) and isinstance(cur, QwenPDelta3CLVRAttention): + cur.core.load_state_dict(prev.core.state_dict(), strict=True) + cur.local_gate_w.data.copy_(prev.local_gate_w.data) + cur.local_gate_b.data.copy_(prev.local_gate_b.data) + print(f" warm-started layer {current} from full-attention layer {accepted[-1]}", flush=True) + + +def trainable_groups(model, idx, lr, qkv_scale, train_qkv): + for p in model.parameters(): p.requires_grad = False + m = model.model.layers[idx].self_attn + main = [] + for p in m.core.parameters(): p.requires_grad = True; main.append(p) + for p in (m.local_gate_w, m.local_gate_b): p.requires_grad = True; main.append(p) + slow = [] + if train_qkv: + for child in (m.q_proj, m.k_proj, m.v_proj, m.q_norm, m.k_norm): + for p in child.parameters(): p.requires_grad = True; slow.append(p) + groups = [{"params": main, "lr": lr}] + if slow: groups.append({"params": slow, "lr": lr * qkv_scale}) + return groups, main + slow + + +def make_optimizer(groups, device): + try: return torch.optim.AdamW(groups, weight_decay=0.01, fused=device.type == "cuda") + except Exception: return torch.optim.AdamW(groups, weight_decay=0.01) + + +def capture_attention_input(model, ids, idx, amp, with_output): + cap: dict[str, Any] = {} + module = model.model.layers[idx].self_attn + def hook(mod, args, kwargs): + h = args[0] if args else kwargs.get("hidden_states") + if h is None: raise RuntimeError("hidden_states not found") + cap["hidden"] = h + for key in ("position_embeddings","position_ids","attention_mask","cache_position"): + if key in kwargs and kwargs[key] is not None: + v = kwargs[key] + if torch.is_tensor(v): v = v.detach() + elif isinstance(v, tuple): v = tuple(x.detach() if torch.is_tensor(x) else x for x in v) + cap[key] = v + handle = module.register_forward_pre_hook(hook, with_kwargs=True) + try: + if with_output: + with amp(): out = model(input_ids=ids, labels=ids, use_cache=False, return_dict=True) + else: + with torch.no_grad(), amp(): out = model(input_ids=ids, use_cache=False, return_dict=True) + finally: handle.remove() + if "hidden" not in cap: raise RuntimeError(f"failed to capture layer {idx}") + return cap, out + + +def attention_kwargs(cap): + return {k:cap[k] for k in ("position_embeddings","position_ids","attention_mask","cache_position") if k in cap} + + +def call_attention(module, hidden, kwargs): + out = module(hidden, past_key_values=None, **kwargs) + return out[0] if isinstance(out, (tuple,list)) else out + + +@torch.no_grad() +def function_metrics(teacher, student, ids, idx, amp): + tc, _ = capture_attention_input(teacher, ids, idx, amp, False) + sc, _ = capture_attention_input(student, ids, idx, amp, False) + hidden = sc["hidden"].detach(); kwargs = attention_kwargs(tc) + target = call_attention(teacher.model.layers[idx].self_attn, hidden, kwargs) + pred = call_attention(student.model.layers[idx].self_attn, hidden, kwargs) + nmse, cosine = seq.alignment_metrics(pred, target) + return float(nmse), float(cosine) + + +def distill_kl(student_logits, teacher_logits, temperature): + s, t = student_logits.float()/temperature, teacher_logits.float()/temperature + 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]) + + +def passes(nmse, cosine, inc, total, args): + nll_ok = inc <= args.accept_incremental_delta_nll and total <= args.accept_cumulative_delta_nll + return nll_ok if not args.strict_acceptance else (nmse <= args.accept_nmse and cosine >= args.accept_cosine and nll_ok) + + +def score(r): + return (float(r["cumulative_delta_nll"]), float(r["incremental_delta_nll"]), float(r["nmse"])) + + +def train_round(teacher, student, idx, batch_iter, probe_blocks, teacher_nll, pre_nll, args, device, amp, round_idx): + rescue = round_idx > 1 + lr = args.layer_lr * (args.rescue_lr_scale if rescue else 1.0) + fw = args.rescue_functional_weight if rescue else args.functional_weight + kw = args.rescue_kl_weight if rescue else args.kl_weight + cw = args.rescue_ce_weight if rescue else args.ce_weight + groups, trainable = trainable_groups(student, idx, lr, args.qkv_lr_scale, args.train_qkv) + opt = make_optimizer(groups, device) + scaler = torch.amp.GradScaler("cuda", enabled=(device.type == "cuda" and seq.choose_dtype(device) == torch.float16)) + student.train(); module = student.model.layers[idx].self_attn + best = best_state = None + for step in range(1, args.max_layer_steps + 1): + ids = next(batch_iter).to(device, non_blocking=True) + tc, to = capture_attention_input(teacher, ids, idx, amp, True) + sc, so = capture_attention_input(student, ids, idx, amp, True) + hidden = sc["hidden"].detach(); kwargs = attention_kwargs(tc) + with torch.no_grad(), amp(): target = call_attention(teacher.model.layers[idx].self_attn, hidden, kwargs) + with amp(): + pred = call_attention(module, hidden, kwargs) + functional = seq.alignment_loss(pred, target, args.cosine_weight) + kl = distill_kl(so.logits, to.logits, args.temperature) + local_penalty = torch.sigmoid(module.local_gate_b.float()).mean() + loss = fw*functional + kw*kl + cw*so.loss + args.local_gate_penalty*local_penalty + if not torch.isfinite(loss): raise RuntimeError(f"non-finite loss layer {idx} step {step}") + opt.zero_grad(set_to_none=True); scaler.scale(loss).backward(); scaler.unscale_(opt) + grad = float(torch.nn.utils.clip_grad_norm_(trainable, 1.0)); scaler.step(opt); scaler.update() + if step == 1 or step % args.log_every == 0: + bn, bc = seq.alignment_metrics(pred.detach(), target.detach()) + 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) + if step >= args.min_layer_steps and (step % args.check_every == 0 or step == args.max_layer_steps): + nmse, cosine = function_metrics(teacher, student, ids, idx, amp) + nll = seq.probe_nll(student, probe_blocks, device, amp); student.train() + inc, cum = nll - pre_nll, nll - teacher_nll + ok = passes(nmse, cosine, inc, cum, args) + 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} + if best is None or score(cand) < score(best): best, best_state = cand, selected_state(student, [idx]) + 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) + if ok: best, best_state = cand, selected_state(student, [idx]); break + del to, so, pred, target, loss + if best is None or best_state is None: raise RuntimeError("acceptance was never evaluated") + _load_selected(student, best_state, [idx]) + return best + + +def parse_args(): + p = argparse.ArgumentParser() + p.add_argument("--base-model", default="Qwen/Qwen3.5-0.8B"); p.add_argument("--output-dir", required=True) + 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) + 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) + 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) + 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) + 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) + 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) + 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) + return p.parse_args() + + +def main(): + args = parse_args(); random.seed(args.seed); torch.manual_seed(args.seed) + if torch.cuda.is_available(): torch.cuda.manual_seed_all(args.seed) + out = Path(args.output_dir); out.mkdir(parents=True, exist_ok=True) + max_rounds = max(1, int(os.environ.get("SEQUENTIAL_MAX_ROUNDS_PER_RUN", "2"))) + device = torch.device("cuda" if torch.cuda.is_available() else "cpu"); dtype = seq.choose_dtype(device); amp = seq.amp_factory(device, dtype) + if device.type == "cuda": torch.cuda.reset_peak_memory_stats(); torch.backends.cuda.matmul.allow_tf32 = True + print(f"device={device} dtype={dtype} base={args.base_model} feature_dim={args.feature_dim} local_window={args.local_window}", flush=True) + tok = AutoTokenizer.from_pretrained(args.base_model, use_fast=True, token=False) + if tok.pad_token_id is None: tok.pad_token = tok.eos_token + load_kwargs = dict(dtype=dtype, token=False, attn_implementation="eager") + teacher = Qwen3_5ForCausalLM.from_pretrained(args.base_model, **load_kwargs).to(device); teacher.eval(); teacher.config.use_cache=False; teacher.requires_grad_(False) + student = Qwen3_5ForCausalLM.from_pretrained(args.base_model, **load_kwargs).to(device); student.config.use_cache=False + all_full = full_attention_layers(student) + if not 1 <= args.target_full_layers <= len(all_full): raise ValueError(f"target-full-layers must be in [1,{len(all_full)}]") + targets = all_full[:args.target_full_layers] + print(f"Qwen3.5 full-attention layers: {all_full}", flush=True); print(f"target prefix: {targets}", flush=True) + 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) + accepted, reports = load_progress(out, student, config, args.resume) + if accepted != targets[:len(accepted)]: raise RuntimeError(f"accepted={accepted} is not target prefix={targets}") + inprog = load_in_progress(out, student, config, accepted, targets, args.resume) if len(accepted)= args.max_runtime_minutes*0.75: + 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 + if not accepted_now: + 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."} + atomic_json(status,out/"qwen35_run_status.json"); print("\nNOT A CRASH:",json.dumps(status,indent=2),flush=True); return 0 + final_nll=seq.probe_nll(student,probes,device,amp) + 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} + 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 + + +if __name__ == "__main__": + try: raise SystemExit(main()) + except KeyboardInterrupt: + print("Interrupted. Persistent best checkpoints remain resumable.", file=sys.stderr, flush=True); raise diff --git a/scripts/train_smollm2_memory_fusion_sequential.py b/scripts/train_smollm2_memory_fusion_sequential.py new file mode 100644 index 0000000000000000000000000000000000000000..32f8e7ad9878e20e822db998b32deb57f6916bdb --- /dev/null +++ b/scripts/train_smollm2_memory_fusion_sequential.py @@ -0,0 +1,986 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import argparse +import json +import random +import time +from contextlib import nullcontext +from pathlib import Path +from typing import Any + +import torch +import torch.nn.functional as F +from datasets import load_dataset +from transformers import AutoModelForCausalLM, AutoTokenizer + +from tinycenn_lm.smollm2_memory_fusion import ( + DEFAULT_SMOLLM2, + MemoryFusionLlamaAttention, + SmolMemoryFusionConfig, + parameter_summary, + replace_attention_layers, + structural_summary, +) + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser( + description=( + "Sequential teacher-guided SmolLM2 Memory Fusion conversion with " + "per-layer acceptance gates and staged whole-model distillation." + ) + ) + p.add_argument("--base-model", default=DEFAULT_SMOLLM2) + p.add_argument("--output-dir", default="checkpoints/smollm2-memory-fusion-sequential-r64") + 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("--context-length", type=int, default=128) + p.add_argument("--batch-size", type=int, default=1) + p.add_argument("--feature-dim", type=int, default=32) + p.add_argument("--memory-rank", type=int, choices=(32, 48, 64), default=64) + p.add_argument("--seed", type=int, default=73) + p.add_argument("--shuffle-buffer", type=int, default=2048) + + p.add_argument("--min-layer-steps", type=int, default=50) + p.add_argument("--max-layer-steps", type=int, default=300) + p.add_argument("--check-every", type=int, default=25) + p.add_argument("--layer-lr", type=float, default=2e-4) + p.add_argument("--teacher-alpha-start", type=float, default=0.90) + p.add_argument("--teacher-alpha-end", type=float, default=0.00) + p.add_argument("--layer-kl-weight", type=float, default=0.10) + p.add_argument("--layer-ce-weight", type=float, default=0.05) + p.add_argument("--cosine-weight", type=float, default=0.25) + p.add_argument("--accept-nmse", type=float, default=0.20) + p.add_argument("--accept-cosine", type=float, default=0.90) + 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("--probe-blocks", type=int, default=4) + p.add_argument("--probe-context", type=int, default=128) + p.add_argument("--strict-acceptance", action=argparse.BooleanOptionalAction, default=True) + + p.add_argument("--core-o-tokens", type=int, default=50_000) + p.add_argument("--core-o-lr", type=float, default=3e-5) + p.add_argument("--norm-tokens", type=int, default=50_000) + p.add_argument("--norm-lr", type=float, default=8e-6) + p.add_argument("--full-tokens", type=int, default=100_000) + p.add_argument("--full-lr", type=float, default=3e-6) + p.add_argument("--grad-accum", type=int, default=4) + p.add_argument("--temperature", type=float, default=2.0) + p.add_argument("--ce-weight", type=float, default=0.25) + p.add_argument("--kl-weight", type=float, default=1.0) + p.add_argument("--hidden-weight", type=float, default=0.5) + + p.add_argument("--resume", action=argparse.BooleanOptionalAction, default=True) + p.add_argument("--backup-every-updates", type=int, default=50) + p.add_argument("--max-runtime-minutes", type=float, default=240.0) + p.add_argument("--log-every", type=int, default=10) + return p.parse_args() + + +def set_seed(seed: int) -> None: + random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def choose_dtype(device: torch.device) -> torch.dtype: + if device.type != "cuda": + return torch.float32 + return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 + + +def amp_factory(device: torch.device, dtype: torch.dtype): + if device.type == "cuda": + return lambda: torch.autocast("cuda", dtype=dtype) + return nullcontext + + +def token_blocks(dataset, tokenizer, text_field: str, context_length: int): + eos = tokenizer.eos_token_id + if eos is None: + raise ValueError("tokenizer must define eos_token_id") + buffer: list[int] = [] + for row in dataset: + text = str(row.get(text_field, "")).strip() + if not text: + continue + ids = tokenizer(text, add_special_tokens=False, verbose=False)["input_ids"] + if not ids: + continue + buffer.extend(ids) + buffer.append(eos) + while len(buffer) >= context_length: + yield torch.tensor(buffer[:context_length], dtype=torch.long) + del buffer[:context_length] + + +def batches(blocks, batch_size: int): + pending = [] + for block in blocks: + pending.append(block) + if len(pending) == batch_size: + yield torch.stack(pending) + pending.clear() + + +def make_optimizer(params, lr: float, device: torch.device): + params = list(params) + if not params: + raise ValueError("optimizer received no trainable parameters") + try: + return torch.optim.AdamW(params, lr=lr, weight_decay=0.01, fused=(device.type == "cuda")) + except Exception: + return torch.optim.AdamW(params, lr=lr, weight_decay=0.01) + + +def distillation_kl(student_logits, teacher_logits, temperature: float) -> torch.Tensor: + s = student_logits.float() / temperature + t = teacher_logits.float() / temperature + per_token = F.kl_div( + F.log_softmax(s, dim=-1), + F.softmax(t, dim=-1), + reduction="none", + ).sum(dim=-1) + return per_token.mean() * (temperature ** 2) + + +def representation_loss(student_hidden, teacher_hidden) -> torch.Tensor: + max_idx = min(len(student_hidden), len(teacher_hidden)) - 1 + candidates = (6, 12, 18, 24, 30) + indices = [i for i in candidates if i <= max_idx] + if not indices: + indices = [max_idx] + terms = [] + for idx in indices: + s = student_hidden[idx].float() + t = teacher_hidden[idx].float() + cosine = 1.0 - F.cosine_similarity(s, t, dim=-1).mean() + nmse = (s - t).square().mean() / t.square().mean().clamp_min(1e-5) + terms.append(cosine + 0.25 * nmse) + return torch.stack(terms).mean() + + +def alignment_metrics(pred: torch.Tensor, target: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + pred32 = pred.float() + target32 = target.float() + mse = (pred32 - target32).square().mean() + nmse = mse / target32.square().mean().clamp_min(1e-5) + cosine = F.cosine_similarity(pred32, target32, dim=-1).mean() + return nmse, cosine + + +def alignment_loss(pred: torch.Tensor, target: torch.Tensor, cosine_weight: float) -> torch.Tensor: + nmse, cosine = alignment_metrics(pred, target) + return nmse + cosine_weight * (1.0 - cosine) + + +def _detach_tree(value): + if torch.is_tensor(value): + return value.detach() + if isinstance(value, tuple): + return tuple(_detach_tree(v) for v in value) + if isinstance(value, list): + return [_detach_tree(v) for v in value] + return value + + +def capture_attention_input(model, ids: torch.Tensor, layer_idx: int, amp, *, with_output: bool): + capture: dict[str, Any] = {} + module = model.model.layers[layer_idx].self_attn + + def pre_hook(mod, args, kwargs): + hidden = args[0] if args else kwargs.get("hidden_states") + if hidden is None: + raise RuntimeError("attention hidden_states not found") + capture["hidden"] = hidden + for key in ( + "position_embeddings", + "position_ids", + "attention_mask", + "cache_position", + ): + if key in kwargs and kwargs[key] is not None: + capture[key] = _detach_tree(kwargs[key]) + + handle = module.register_forward_pre_hook(pre_hook, with_kwargs=True) + try: + if with_output: + with amp(): + out = model( + input_ids=ids, + labels=ids, + use_cache=False, + output_hidden_states=True, + return_dict=True, + ) + else: + with torch.no_grad(), amp(): + out = model( + input_ids=ids, + use_cache=False, + output_hidden_states=False, + return_dict=True, + ) + finally: + handle.remove() + if "hidden" not in capture: + raise RuntimeError(f"failed to capture layer {layer_idx} attention input") + return capture, out + + +def attention_kwargs_from_capture(capture: dict[str, Any]) -> dict[str, Any]: + kwargs: dict[str, Any] = {"use_cache": False} + for key in ( + "position_embeddings", + "position_ids", + "attention_mask", + "cache_position", + ): + if key in capture: + kwargs[key] = capture[key] + return kwargs + + +def call_attention(module, hidden: torch.Tensor, kwargs: dict[str, Any]) -> torch.Tensor: + out = module(hidden, **kwargs) + if isinstance(out, (tuple, list)): + return out[0] + return out + + +def alpha_for_step(step: int, max_steps: int, start: float, end: float) -> float: + if max_steps <= 1: + return float(end) + progress = min(max((step - 1) / (max_steps - 1), 0.0), 1.0) + return float(start + progress * (end - start)) + + +def freeze_current_layer_only(student, layer_idx: int) -> list[torch.nn.Parameter]: + for p in student.parameters(): + p.requires_grad = False + module = student.model.layers[layer_idx].self_attn + if not isinstance(module, MemoryFusionLlamaAttention): + raise TypeError(f"layer {layer_idx} is not MemoryFusionLlamaAttention") + trainable = [] + for p in module.core.parameters(): + p.requires_grad = True + trainable.append(p) + for p in module.o_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def select_integrated_stage_params(student, stage: str) -> list[torch.nn.Parameter]: + for p in student.parameters(): + p.requires_grad = False + + if stage in {"core_o", "core_o_norm"}: + for module in student.modules(): + if isinstance(module, MemoryFusionLlamaAttention): + for p in module.core.parameters(): + p.requires_grad = True + for p in module.o_proj.parameters(): + p.requires_grad = True + + if stage == "core_o_norm": + for name, p in student.named_parameters(): + if "norm" in name.lower(): + p.requires_grad = True + elif stage == "full": + for p in student.parameters(): + p.requires_grad = True + else: + raise ValueError(f"unknown integrated stage {stage!r}") + + return [p for p in student.parameters() if p.requires_grad] + + +def _layer_prefix(idx: int) -> str: + return f"model.layers.{idx}.self_attn." + + +def progressive_attention_state(student, accepted_layers: list[int]) -> dict[str, torch.Tensor]: + prefixes = tuple(_layer_prefix(i) for i in accepted_layers) + return { + k: v.detach().cpu() + for k, v in student.state_dict().items() + if prefixes and k.startswith(prefixes) + } + + +def save_progress( + output_dir: Path, + student, + config: SmolMemoryFusionConfig, + accepted_layers: list[int], + layer_reports: list[dict], + *, + stage: str, +) -> None: + output_dir.mkdir(parents=True, exist_ok=True) + payload = { + "format_version": 1, + "stage": stage, + "accepted_layers": accepted_layers, + "config": config.to_dict(), + "layer_reports": layer_reports, + "attention_state": progressive_attention_state(student, accepted_layers), + } + tmp = output_dir / "sequential_progress.tmp" + final = output_dir / "sequential_progress.pt" + torch.save(payload, tmp) + tmp.replace(final) + (output_dir / "sequential_progress.json").write_text( + json.dumps( + { + "format_version": 1, + "stage": stage, + "accepted_layers": accepted_layers, + "config": config.to_dict(), + "layer_reports": layer_reports, + }, + indent=2, + ), + encoding="utf-8", + ) + + +def load_progress_if_available( + output_dir: Path, + student, + config: SmolMemoryFusionConfig, + *, + resume: bool, +) -> tuple[list[int], list[dict]]: + path = output_dir / "sequential_progress.pt" + if not resume or not path.exists(): + return [], [] + payload = torch.load(path, map_location="cpu", weights_only=False) + accepted = [int(x) for x in payload.get("accepted_layers", [])] + if payload.get("config", {}).get("memory_rank") != config.memory_rank: + raise RuntimeError("resume checkpoint memory rank differs from requested rank") + if accepted: + replace_attention_layers(student, config, accepted) + incompatible = student.load_state_dict(payload["attention_state"], strict=False) + expected = { + k + for k in student.state_dict() + if any(k.startswith(_layer_prefix(i)) for i in accepted) + } + missing = [k for k in incompatible.missing_keys if k in expected] + if missing: + raise RuntimeError(f"resume checkpoint missing accepted-layer keys: {missing[:8]}") + print(f"RESUME: accepted layers={accepted}") + return accepted, list(payload.get("layer_reports", [])) + + +def save_full_state( + output_dir: Path, + student, + config: SmolMemoryFusionConfig, + metadata: dict, + *, + filename: str = "smollm2_memory_fusion_sequential_full.pt", +) -> None: + output_dir.mkdir(parents=True, exist_ok=True) + state_path = output_dir / filename + tmp = output_dir / (filename + ".tmp") + torch.save({k: v.detach().cpu() for k, v in student.state_dict().items()}, tmp) + tmp.replace(state_path) + meta = { + "format_version": 1, + "architecture": "smollm2-memory-fusion-sequential-full", + "base_model": metadata["base_model"], + "memory_fusion": config.to_dict(), + "training": metadata, + "state_file": filename, + } + (output_dir / "smollm2_memory_fusion_sequential_config.json").write_text( + json.dumps(meta, indent=2), + encoding="utf-8", + ) + + +def build_probe_blocks(tokenizer, context: int, count: int) -> list[torch.Tensor]: + try: + wiki = load_dataset( + "Salesforce/wikitext", + "wikitext-2-raw-v1", + split="validation", + ) + except Exception: + texts = [ + "The history of science is shaped by observation, measurement, and careful comparison.", + "A language model predicts the next token from the sequence that came before it.", + "Vienna is the capital of Austria and has a long history of music, science, and public administration.", + "Neural networks can approximate complex functions when their parameters are trained on representative data.", + "The experiment should separate training data from the probe used to decide whether a replacement is acceptable.", + "Reliable evaluation requires the same inputs, tokenizer, context length, and scoring rule for every model.", + "A recurrent memory can carry information through a sequence without constructing a full dense attention matrix.", + "When a model component is replaced, small approximation errors can accumulate across many layers.", + ] + stream = "\n".join(texts) + else: + stream = "\n".join(str(x) for x in wiki["text"] if str(x).strip()) + + ids = tokenizer(stream, add_special_tokens=False)["input_ids"] + need = context + blocks = [] + for start in range(0, len(ids) - need + 1, need): + blocks.append(torch.tensor(ids[start : start + need], dtype=torch.long)) + if len(blocks) >= count: + break + if not blocks: + raise RuntimeError("could not construct probe blocks") + return blocks + + +@torch.no_grad() +def probe_nll(model, probe_blocks: list[torch.Tensor], device: torch.device, amp) -> float: + model.eval() + losses = [] + for block in probe_blocks: + x = block.unsqueeze(0).to(device) + with amp(): + out = model(input_ids=x, labels=x, use_cache=False, return_dict=True) + losses.append(float(out.loss.detach().float())) + return sum(losses) / len(losses) + + +@torch.no_grad() +def real_hidden_function_metrics( + teacher, + student, + ids: torch.Tensor, + layer_idx: int, + amp, +) -> tuple[float, float]: + t_capture, _ = capture_attention_input(teacher, ids, layer_idx, amp, with_output=False) + s_capture, _ = capture_attention_input(student, ids, layer_idx, amp, with_output=False) + real_hidden = s_capture["hidden"].detach() + kwargs = attention_kwargs_from_capture(t_capture) + target = call_attention(teacher.model.layers[layer_idx].self_attn, real_hidden, kwargs) + pred = call_attention(student.model.layers[layer_idx].self_attn, real_hidden, kwargs) + nmse, cosine = alignment_metrics(pred, target) + return float(nmse), float(cosine) + + +def acceptance_passes( + *, + nmse: float, + cosine: float, + incremental_delta_nll: float, + cumulative_delta_nll: float, + args: argparse.Namespace, +) -> bool: + return ( + nmse <= args.accept_nmse + and cosine >= args.accept_cosine + and incremental_delta_nll <= args.accept_incremental_delta_nll + and cumulative_delta_nll <= args.accept_cumulative_delta_nll + ) + + +def train_one_replacement( + *, + teacher, + student, + layer_idx: int, + batch_iter, + probe_blocks, + teacher_probe_nll: float, + pre_replacement_probe_nll: float, + args, + device, + amp, +) -> dict: + trainable = freeze_current_layer_only(student, layer_idx) + student.train() + optimizer = make_optimizer(trainable, args.layer_lr, device) + scaler = torch.amp.GradScaler( + "cuda", enabled=(device.type == "cuda" and choose_dtype(device) == torch.float16) + ) + best: dict[str, float | int | bool] | None = None + accepted = False + + for step in range(1, args.max_layer_steps + 1): + ids = next(batch_iter).to(device, non_blocking=True) + + t_capture, t_out = capture_attention_input( + teacher, ids, layer_idx, amp, with_output=True + ) + s_capture, s_out = capture_attention_input( + student, ids, layer_idx, amp, with_output=True + ) + + alpha = alpha_for_step( + step, + args.max_layer_steps, + args.teacher_alpha_start, + args.teacher_alpha_end, + ) + mixed_hidden = ( + alpha * t_capture["hidden"].detach() + + (1.0 - alpha) * s_capture["hidden"].detach() + ) + kwargs = attention_kwargs_from_capture(t_capture) + + with torch.no_grad(), amp(): + target = call_attention( + teacher.model.layers[layer_idx].self_attn, + mixed_hidden, + kwargs, + ) + + with amp(): + pred = call_attention( + student.model.layers[layer_idx].self_attn, + mixed_hidden, + kwargs, + ) + functional = alignment_loss(pred, target, args.cosine_weight) + kl = distillation_kl(s_out.logits, t_out.logits, args.temperature) + total = functional + args.layer_kl_weight * kl + args.layer_ce_weight * s_out.loss + + if not torch.isfinite(total): + raise RuntimeError(f"non-finite loss at layer {layer_idx}, step {step}") + + optimizer.zero_grad(set_to_none=True) + scaler.scale(total).backward() + scaler.unscale_(optimizer) + grad_norm = float(torch.nn.utils.clip_grad_norm_(trainable, 1.0)) + scaler.step(optimizer) + scaler.update() + + if step == 1 or step % args.log_every == 0: + nmse_batch, cosine_batch = alignment_metrics(pred.detach(), target.detach()) + print( + f"layer={layer_idx:02d} step={step:03d}/{args.max_layer_steps} " + f"alpha={alpha:.3f} functional={float(functional.detach()):.4f} " + f"nmse={float(nmse_batch):.4f} cos={float(cosine_batch):.4f} " + f"kl={float(kl.detach()):.4f} ce={float(s_out.loss.detach()):.4f} " + f"grad={grad_norm:.3f}" + ) + + should_check = ( + step >= args.min_layer_steps + and (step % args.check_every == 0 or step == args.max_layer_steps) + ) + if should_check: + nmse, cosine = real_hidden_function_metrics( + teacher, student, ids, layer_idx, amp + ) + current_probe_nll = probe_nll(student, probe_blocks, device, amp) + student.train() + incremental = current_probe_nll - pre_replacement_probe_nll + cumulative = current_probe_nll - teacher_probe_nll + passed = acceptance_passes( + nmse=nmse, + cosine=cosine, + incremental_delta_nll=incremental, + cumulative_delta_nll=cumulative, + args=args, + ) + candidate = { + "step": step, + "nmse": nmse, + "cosine": cosine, + "probe_nll": current_probe_nll, + "incremental_delta_nll": incremental, + "cumulative_delta_nll": cumulative, + "passed": passed, + } + if best is None or ( + candidate["nmse"] < best["nmse"] + and candidate["cumulative_delta_nll"] <= best["cumulative_delta_nll"] + 0.01 + ): + best = candidate + print( + f" ACCEPTANCE CHECK layer={layer_idx:02d}: " + f"NMSE={nmse:.4f} (≤{args.accept_nmse:.4f}) " + f"cos={cosine:.4f} (≥{args.accept_cosine:.4f}) " + f"ΔNLL_inc={incremental:+.5f} (≤{args.accept_incremental_delta_nll:+.5f}) " + f"ΔNLL_total={cumulative:+.5f} (≤{args.accept_cumulative_delta_nll:+.5f}) " + f"=> {'PASS' if passed else 'continue'}" + ) + if passed: + accepted = True + best = candidate + break + + del t_out, s_out, pred, target, total + + if best is None: + raise RuntimeError("acceptance was never evaluated") + return { + "layer": layer_idx, + "accepted": accepted, + "steps": int(best["step"]), + "nmse": float(best["nmse"]), + "cosine": float(best["cosine"]), + "probe_nll": float(best["probe_nll"]), + "incremental_delta_nll": float(best["incremental_delta_nll"]), + "cumulative_delta_nll": float(best["cumulative_delta_nll"]), + } + + +def run_integrated_stage( + *, + teacher, + student, + batch_iter, + stage_name: str, + token_budget: int, + lr: float, + args, + device, + amp, + output_dir: Path, + config: SmolMemoryFusionConfig, + report: dict, +) -> dict: + if token_budget <= 0: + return {"stage": stage_name, "tokens": 0, "updates": 0, "skipped": True} + + trainable = select_integrated_stage_params(student, stage_name) + optimizer = make_optimizer(trainable, lr, device) + scaler = torch.amp.GradScaler( + "cuda", enabled=(device.type == "cuda" and choose_dtype(device) == torch.float16) + ) + student.train() + optimizer.zero_grad(set_to_none=True) + + seen_tokens = 0 + micro = 0 + updates = 0 + last = {} + start = time.perf_counter() + + while seen_tokens < token_budget: + ids = next(batch_iter).to(device, non_blocking=True) + with torch.no_grad(), amp(): + t_out = teacher( + input_ids=ids, + use_cache=False, + output_hidden_states=True, + return_dict=True, + ) + with amp(): + s_out = student( + input_ids=ids, + labels=ids, + use_cache=False, + output_hidden_states=True, + return_dict=True, + ) + kl = distillation_kl(s_out.logits, t_out.logits, args.temperature) + hidden = representation_loss(s_out.hidden_states, t_out.hidden_states) + loss = ( + args.ce_weight * s_out.loss + + args.kl_weight * kl + + args.hidden_weight * hidden + ) + scaled = loss / args.grad_accum + + if not torch.isfinite(scaled): + raise RuntimeError(f"non-finite loss during integrated stage {stage_name}") + + scaler.scale(scaled).backward() + micro += 1 + seen_tokens += ids.numel() + last = { + "ce": float(s_out.loss.detach().float()), + "kl": float(kl.detach().float()), + "hidden": float(hidden.detach().float()), + "total": float(loss.detach().float()), + } + del t_out, s_out + + if micro % args.grad_accum: + continue + + scaler.unscale_(optimizer) + grad_norm = float(torch.nn.utils.clip_grad_norm_(trainable, 1.0)) + scaler.step(optimizer) + scaler.update() + optimizer.zero_grad(set_to_none=True) + updates += 1 + + if updates == 1 or updates % args.log_every == 0: + print( + f"stage={stage_name} update={updates} tokens={seen_tokens:,}/{token_budget:,} " + f"ce={last['ce']:.4f} kl={last['kl']:.4f} hidden={last['hidden']:.4f} " + f"grad={grad_norm:.3f}" + ) + + if args.backup_every_updates > 0 and updates % args.backup_every_updates == 0: + backup_meta = dict(report) + backup_meta["integrated_stage"] = stage_name + backup_meta["integrated_stage_tokens"] = seen_tokens + save_full_state( + output_dir, + student, + config, + backup_meta, + filename="live_sequential_full_state.pt", + ) + print(" persistent full-state backup saved") + + elapsed = time.perf_counter() - start + return { + "stage": stage_name, + "tokens": seen_tokens, + "updates": updates, + "lr": lr, + "elapsed_minutes": elapsed / 60.0, + "last": last, + "trainable_parameters": sum(p.numel() for p in trainable), + } + + +def main() -> None: + args = parse_args() + set_seed(args.seed) + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + dtype = choose_dtype(device) + amp = amp_factory(device, dtype) + if device.type == "cuda": + torch.cuda.reset_peak_memory_stats() + torch.backends.cuda.matmul.allow_tf32 = True + + print(f"device={device} dtype={dtype} memory_rank={args.memory_rank}") + print(f"output_dir={output_dir}") + + tokenizer = AutoTokenizer.from_pretrained(args.base_model, use_fast=True) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + + teacher = AutoModelForCausalLM.from_pretrained(args.base_model, dtype=dtype).to(device) + teacher.eval() + teacher.config.use_cache = False + teacher.requires_grad_(False) + + student = AutoModelForCausalLM.from_pretrained(args.base_model, dtype=dtype).to(device) + student.config.use_cache = False + + config = SmolMemoryFusionConfig( + feature_dim=args.feature_dim, + memory_rank=args.memory_rank, + train_output_projection=True, + ) + + accepted_layers, layer_reports = load_progress_if_available( + output_dir, student, config, resume=args.resume + ) + + raw = load_dataset( + args.dataset, + name=args.dataset_config, + split=args.split, + streaming=True, + ).shuffle(seed=args.seed, buffer_size=args.shuffle_buffer) + batch_iter = iter( + batches( + token_blocks(raw, tokenizer, args.text_field, args.context_length), + args.batch_size, + ) + ) + + probe_blocks = build_probe_blocks( + tokenizer, + context=args.probe_context, + count=args.probe_blocks, + ) + teacher_probe_nll = probe_nll(teacher, probe_blocks, device, amp) + print(f"teacher probe NLL={teacher_probe_nll:.6f}") + + start_time = time.perf_counter() + num_layers = int(student.config.num_hidden_layers) + + for layer_idx in range(num_layers): + if layer_idx in accepted_layers: + continue + if accepted_layers != list(range(layer_idx)): + raise RuntimeError( + f"accepted layers must be a contiguous prefix before layer {layer_idx}: " + f"{accepted_layers}" + ) + if (time.perf_counter() - start_time) / 60.0 >= args.max_runtime_minutes * 0.70: + raise RuntimeError( + "runtime budget reached during sequential replacement; progress was " + "saved and the same command can resume from the last accepted layer" + ) + + print("\n" + "=" * 110) + print(f"SEQUENTIAL REPLACEMENT: layer {layer_idx}/{num_layers - 1}") + print("=" * 110) + + pre_probe_nll = probe_nll(student, probe_blocks, device, amp) + print( + f"before replacement: probe NLL={pre_probe_nll:.6f}, " + f"Δ vs teacher={pre_probe_nll - teacher_probe_nll:+.6f}" + ) + + replace_attention_layers(student, config, [layer_idx]) + + layer_report = train_one_replacement( + teacher=teacher, + student=student, + layer_idx=layer_idx, + batch_iter=batch_iter, + probe_blocks=probe_blocks, + teacher_probe_nll=teacher_probe_nll, + pre_replacement_probe_nll=pre_probe_nll, + args=args, + device=device, + amp=amp, + ) + layer_reports.append(layer_report) + + if not layer_report["accepted"] and args.strict_acceptance: + save_progress( + output_dir, + student, + config, + accepted_layers, + layer_reports, + stage=f"layer_{layer_idx}_rejected", + ) + (output_dir / "sequential_training_report.json").write_text( + json.dumps( + { + "status": "stopped_on_rejected_layer", + "base_model": args.base_model, + "memory_rank": args.memory_rank, + "accepted_layers": accepted_layers, + "layer_reports": layer_reports, + "thresholds": { + "nmse": args.accept_nmse, + "cosine": args.accept_cosine, + "incremental_delta_nll": args.accept_incremental_delta_nll, + "cumulative_delta_nll": args.accept_cumulative_delta_nll, + }, + }, + indent=2, + ), + encoding="utf-8", + ) + raise RuntimeError( + f"layer {layer_idx} did not satisfy acceptance criteria; " + "the next Transformer layer was NOT replaced" + ) + + if not layer_report["accepted"]: + print("WARNING: relaxed mode accepts a layer that did not pass all thresholds") + + accepted_layers.append(layer_idx) + save_progress( + output_dir, + student, + config, + accepted_layers, + layer_reports, + stage=f"accepted_layer_{layer_idx}", + ) + print(f"✅ accepted layer {layer_idx}; persistent progress saved") + + if accepted_layers != list(range(num_layers)): + raise RuntimeError("not all layers were accepted") + + summary = structural_summary(student) + if summary["memory_fusion_layers"] != num_layers or summary["transformer_attention_layers"] != 0: + raise RuntimeError(f"unexpected final structure: {summary}") + + report = { + "status": "all_layers_accepted", + "architecture": "smollm2-memory-fusion-sequential-full", + "base_model": args.base_model, + "memory_rank": args.memory_rank, + "feature_dim": args.feature_dim, + "context_length": args.context_length, + "teacher_probe_nll": teacher_probe_nll, + "accepted_layers": accepted_layers, + "layer_reports": layer_reports, + "thresholds": { + "nmse": args.accept_nmse, + "cosine": args.accept_cosine, + "incremental_delta_nll": args.accept_incremental_delta_nll, + "cumulative_delta_nll": args.accept_cumulative_delta_nll, + }, + "integrated_stages": [], + } + + print("\n" + "=" * 110) + print("ALL 30 REPLACEMENTS ACCEPTED — STARTING INTEGRATED TRAINING") + print("=" * 110) + + integrated_specs = [ + ("core_o", args.core_o_tokens, args.core_o_lr), + ("core_o_norm", args.norm_tokens, args.norm_lr), + ("full", args.full_tokens, args.full_lr), + ] + for stage_name, token_budget, lr in integrated_specs: + if (time.perf_counter() - start_time) / 60.0 >= args.max_runtime_minutes: + print("runtime budget reached before remaining integrated stages") + report["status"] = "runtime_budget_after_acceptance" + break + print(f"\n--- integrated stage: {stage_name} ---") + stage_report = run_integrated_stage( + teacher=teacher, + student=student, + batch_iter=batch_iter, + stage_name=stage_name, + token_budget=token_budget, + lr=lr, + args=args, + device=device, + amp=amp, + output_dir=output_dir, + config=config, + report=report, + ) + report["integrated_stages"].append(stage_report) + save_full_state( + output_dir, + student, + config, + report, + filename=f"stage_{stage_name}_full_state.pt", + ) + print(f"✅ completed {stage_name}; full persistent checkpoint saved") + + student.eval() + final_probe_nll = probe_nll(student, probe_blocks, device, amp) + report["final_probe_nll"] = final_probe_nll + report["final_probe_delta_nll"] = final_probe_nll - teacher_probe_nll + report["parameters"] = parameter_summary(student) + report["structure"] = structural_summary(student) + report["elapsed_minutes"] = (time.perf_counter() - start_time) / 60.0 + report["peak_vram_gib"] = ( + torch.cuda.max_memory_allocated() / (1024 ** 3) + if device.type == "cuda" + else 0.0 + ) + + save_full_state(output_dir, student, config, report) + (output_dir / "sequential_training_report.json").write_text( + json.dumps(report, indent=2), + encoding="utf-8", + ) + tokenizer.save_pretrained(output_dir) + + print("\nFINAL REPORT") + print(json.dumps(report, indent=2)) + print("saved:", output_dir) + + +if __name__ == "__main__": + main() diff --git a/src/tinycenn_lm/__init__.py b/src/tinycenn_lm/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..220525a4b49bb64df733f2fa7b228f7f23d17673 --- /dev/null +++ b/src/tinycenn_lm/__init__.py @@ -0,0 +1,144 @@ +from .live_console import configure_live_console + +# Configure the current process before importing model modules. Every train_*.py +# script imports tinycenn_lm, so notebook-launched trainers inherit immediate, +# line-buffered stdout/stderr even when the notebook uses subprocess.run(...). +configure_live_console() + +from .cenn import CeNNConfig, FastCeNNCore +from .modeling import ( + DEFAULT_BASE_MODEL, + HybridDecoderLayer, + build_from_adapter, + freeze_for_adapter_training, + inject_cenn, + load_adapter, + save_adapter, + trainable_parameter_summary, +) +from .student import ( + CeNNReplacementLayer, + build_cenn_student, + freeze_student_interfaces, + load_cenn_student_weights, + replace_transformer_with_cenn, + save_cenn_student, + student_parameter_summary, +) +from .moe import ( + FastMoECeNNCore, + MoECeNNConfig, + MoECeNNReplacementLayer, + build_moe_cenn_student, + freeze_moe_student_interfaces, + load_moe_cenn_student_weights, + moe_router_stats, + replace_transformer_with_moe_cenn, + save_moe_cenn_student, + warmstart_moe_from_plain_cenn, +) +from .sharded_moe import ( + FastShardedMoECeNNCore, + ShardedMoECeNNConfig, + ShardedMoECeNNReplacementLayer, + build_sharded_moe_student, + freeze_sharded_moe_interfaces, + load_sharded_moe_student_weights, + replace_transformer_with_sharded_moe_cenn, + save_sharded_moe_student, + sharded_router_stats, + warmstart_sharded_moe_from_plain_cenn, +) +from .story_v2 import ( + CausalStoryMemory, + LowRankLMHeadAdapter, + StoryV2Config, + StoryV2ReplacementLayer, + build_story_v2_from_story_v1, + build_story_v2_student, + freeze_story_v2_interfaces, + load_story_v2_weights, + save_story_v2_student, + story_v2_parameter_summary, + story_v2_router_stats, + upgrade_sharded_model_to_story_v2, +) +from .smollm2_amcenn import ( + DEFAULT_SMOLLM2, + AMCeNNAttention, + PositiveSoftmaxFeatures, + ShardedTop2LlamaMLP, + SmolAMCeNNConfig, + amcenn_parameter_summary, + amcenn_router_stats, + build_smollm2_amcenn, + freeze_smollm2_for_amcenn_training, + load_smollm2_amcenn_weights, + replace_smollm2_core, + save_smollm2_amcenn, +) +from .smollm2_amcenn_v2 import ( + AMCeNNAttentionV2, + AdaptivePositiveSoftmaxFeatures, + SmolAMCeNNV2Config, + build_smollm2_amcenn_v2, + convert_all_ffns_to_sharded_top2, + freeze_for_global_training, + freeze_for_group_calibration, + load_smollm2_amcenn_v2_weights, + replace_all_smollm2_attention, + replace_attention_layers, + save_smollm2_amcenn_v2, + v2_parameter_summary, +) +from .hf_persistence import ( + build_model_card, + collect_reports, + install_colab_hf_upload_enhancer, + persist_hf_run, + redact_secrets, + utc_run_id, +) +from .colab_live_backup import ( + install_colab_training_backup, + is_tinycenn_training_command, + output_dir_from_command, +) +from .direct_colab_backup import install_direct_training_backup + +# In Colab, make Hugging Face backup mandatory. The parent notebook wrapper handles +# normal subprocess-launched trainers. A trainer-side fallback covers notebooks that +# launch train_*.py before the notebook kernel imports tinycenn_lm. +install_colab_training_backup() +install_direct_training_backup() +install_colab_hf_upload_enhancer() + +__all__ = [ + "configure_live_console", + "CeNNConfig", "FastCeNNCore", "DEFAULT_BASE_MODEL", "HybridDecoderLayer", + "build_from_adapter", "freeze_for_adapter_training", "inject_cenn", "load_adapter", + "save_adapter", "trainable_parameter_summary", "CeNNReplacementLayer", "build_cenn_student", + "freeze_student_interfaces", "load_cenn_student_weights", "replace_transformer_with_cenn", + "save_cenn_student", "student_parameter_summary", "MoECeNNConfig", "FastMoECeNNCore", + "MoECeNNReplacementLayer", "build_moe_cenn_student", "freeze_moe_student_interfaces", + "load_moe_cenn_student_weights", "moe_router_stats", "replace_transformer_with_moe_cenn", + "save_moe_cenn_student", "warmstart_moe_from_plain_cenn", "ShardedMoECeNNConfig", + "FastShardedMoECeNNCore", "ShardedMoECeNNReplacementLayer", "build_sharded_moe_student", + "freeze_sharded_moe_interfaces", "load_sharded_moe_student_weights", + "replace_transformer_with_sharded_moe_cenn", "save_sharded_moe_student", "sharded_router_stats", + "warmstart_sharded_moe_from_plain_cenn", "StoryV2Config", "CausalStoryMemory", + "StoryV2ReplacementLayer", "LowRankLMHeadAdapter", "upgrade_sharded_model_to_story_v2", + "freeze_story_v2_interfaces", "story_v2_router_stats", "story_v2_parameter_summary", + "save_story_v2_student", "load_story_v2_weights", "build_story_v2_student", + "build_story_v2_from_story_v1", "DEFAULT_SMOLLM2", "SmolAMCeNNConfig", + "PositiveSoftmaxFeatures", "AMCeNNAttention", "ShardedTop2LlamaMLP", "replace_smollm2_core", + "freeze_smollm2_for_amcenn_training", "amcenn_router_stats", "amcenn_parameter_summary", + "save_smollm2_amcenn", "load_smollm2_amcenn_weights", "build_smollm2_amcenn", + "SmolAMCeNNV2Config", "AdaptivePositiveSoftmaxFeatures", "AMCeNNAttentionV2", + "convert_all_ffns_to_sharded_top2", "replace_attention_layers", "replace_all_smollm2_attention", + "freeze_for_group_calibration", "freeze_for_global_training", "v2_parameter_summary", + "save_smollm2_amcenn_v2", "load_smollm2_amcenn_v2_weights", "build_smollm2_amcenn_v2", + "build_model_card", "collect_reports", "persist_hf_run", "install_colab_hf_upload_enhancer", + "redact_secrets", "utc_run_id", "install_colab_training_backup", "install_direct_training_backup", + "is_tinycenn_training_command", "output_dir_from_command", +] diff --git a/src/tinycenn_lm/__pycache__/__init__.cpython-313.pyc b/src/tinycenn_lm/__pycache__/__init__.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..ab4e0b977fd51a673e0e1125eec218b0c301b331 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/__init__.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/cellular_attention.cpython-313.pyc b/src/tinycenn_lm/__pycache__/cellular_attention.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..04920b7811284f4a3fdf2e3885804fe7a0c4321e Binary files /dev/null and b/src/tinycenn_lm/__pycache__/cellular_attention.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/cenn.cpython-313.pyc b/src/tinycenn_lm/__pycache__/cenn.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a094505f24f6886bcc6aea3511997930519ef2cf Binary files /dev/null and b/src/tinycenn_lm/__pycache__/cenn.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/colab_live_backup.cpython-313.pyc b/src/tinycenn_lm/__pycache__/colab_live_backup.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..35269dfdf6595b3f6e5effac7a82ee9c4e27df0b Binary files /dev/null and b/src/tinycenn_lm/__pycache__/colab_live_backup.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/direct_colab_backup.cpython-313.pyc b/src/tinycenn_lm/__pycache__/direct_colab_backup.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..f4344c81cce3f6673f2cc07f1f79dc905141fedf Binary files /dev/null and b/src/tinycenn_lm/__pycache__/direct_colab_backup.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/hf_persistence.cpython-313.pyc b/src/tinycenn_lm/__pycache__/hf_persistence.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..4f4bb7378dfe37336410c22d1bc6f5966b75f05a Binary files /dev/null and b/src/tinycenn_lm/__pycache__/hf_persistence.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/live_console.cpython-313.pyc b/src/tinycenn_lm/__pycache__/live_console.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..f64596bfe7e3380492d97169ea75d65852105ca0 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/live_console.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/memory_attention.cpython-313.pyc b/src/tinycenn_lm/__pycache__/memory_attention.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..41227902f75a7db489f79f9891ab56305c208193 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/memory_attention.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/modeling.cpython-313.pyc b/src/tinycenn_lm/__pycache__/modeling.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..47af8ceedff231d29fa2c9fb53e030b08d09c70b Binary files /dev/null and b/src/tinycenn_lm/__pycache__/modeling.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/moe.cpython-313.pyc b/src/tinycenn_lm/__pycache__/moe.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..455ac87e442e3c7414aa1ee4f2927929e36e34ec Binary files /dev/null and b/src/tinycenn_lm/__pycache__/moe.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/pdelta2_er.cpython-313.pyc b/src/tinycenn_lm/__pycache__/pdelta2_er.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..50d57d3394dbab0a92a14f1568450e19732f02bb Binary files /dev/null and b/src/tinycenn_lm/__pycache__/pdelta2_er.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/pdelta2_features.cpython-313.pyc b/src/tinycenn_lm/__pycache__/pdelta2_features.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..00f6eddf39753813f9c4e7e68267bf5c0d1d198d Binary files /dev/null and b/src/tinycenn_lm/__pycache__/pdelta2_features.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/pdelta3_frontier.cpython-313.pyc b/src/tinycenn_lm/__pycache__/pdelta3_frontier.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..b7369b9023662d8c24c4ab4e991753cae9242b1f Binary files /dev/null and b/src/tinycenn_lm/__pycache__/pdelta3_frontier.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/research_layers.cpython-313.pyc b/src/tinycenn_lm/__pycache__/research_layers.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c59fd1c234bcfef8eb099d7c2cb30da6d997f8b0 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/research_layers.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/sharded_moe.cpython-313.pyc b/src/tinycenn_lm/__pycache__/sharded_moe.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..f29d34776956adcbb24ca3b41d518ec4178dc2ae Binary files /dev/null and b/src/tinycenn_lm/__pycache__/sharded_moe.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/smollm2_amcenn.cpython-313.pyc b/src/tinycenn_lm/__pycache__/smollm2_amcenn.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..86cb90736ade75b1fe11a89da5f79071cd314090 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/smollm2_amcenn.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/smollm2_amcenn_v2.cpython-313.pyc b/src/tinycenn_lm/__pycache__/smollm2_amcenn_v2.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..1148707d5ea5a67a0d1ba7ecac52a59ae10bd227 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/smollm2_amcenn_v2.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/smollm2_memory_fusion.cpython-313.pyc b/src/tinycenn_lm/__pycache__/smollm2_memory_fusion.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..ebcf619d10d59c4ac56d49b9c4fcc3c4dd03c470 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/smollm2_memory_fusion.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/story_v2.cpython-313.pyc b/src/tinycenn_lm/__pycache__/story_v2.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..d8ddddb589ed00804a9ac6fffc661f35f7bde28e Binary files /dev/null and b/src/tinycenn_lm/__pycache__/story_v2.cpython-313.pyc differ diff --git a/src/tinycenn_lm/__pycache__/student.cpython-313.pyc b/src/tinycenn_lm/__pycache__/student.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..fa51b6928ce9e8a42f51c90317cc8545a50348b2 Binary files /dev/null and b/src/tinycenn_lm/__pycache__/student.cpython-313.pyc differ diff --git a/src/tinycenn_lm/cellular_attention.py b/src/tinycenn_lm/cellular_attention.py new file mode 100644 index 0000000000000000000000000000000000000000..917bc45d0981c2657f80e124b84c14ed796857c4 --- /dev/null +++ b/src/tinycenn_lm/cellular_attention.py @@ -0,0 +1,572 @@ +"""Causal 1-D Cellular Attention for TinyCeNN-LM research experiments. + +The layer keeps attention sparse and local at each cellular step, but changes +the neighborhood across steps. Power-of-two dilations create an exponentially +growing receptive field without constructing a T-by-T attention matrix. + +Advanced variants add token/head-specific routing, old-school max/mean pooling, +gated RMS normalization, and a strictly-causal two-level encoder-decoder path. +The U-AMP path uses only current/past states during downsampling and nearest-left +upsampling, so no future token can enter an earlier output. + +This is a research reference implementation. It favors clarity and auditable +causality over fused-kernel speed. +""" +from __future__ import annotations + +import math +from typing import Iterable + +import torch +from torch import Tensor, nn +import torch.nn.functional as F + + +VARIANTS = ( + "cellular_local3", + "cellular_dilated3", + "cellular_dilated5", + "cellular_multiscale5", + "cellular_shifted8", + "cellular_adaptive_multiscale5", + "cellular_multiscale5_maxpool", + "cellular_adaptive_maxpool5", + "cellular_adaptive_maxpool5_rms", + "cellular_adaptive_mixedpool5_rms", + "cellular_uamp5", + "cellular_uamp5_channelgate", + "cellular_uamp5_varlatent", +) + +_ADAPTIVE_VARIANTS = { + "cellular_adaptive_multiscale5", + "cellular_adaptive_maxpool5", + "cellular_adaptive_maxpool5_rms", + "cellular_adaptive_mixedpool5_rms", + "cellular_uamp5", + "cellular_uamp5_channelgate", + "cellular_uamp5_varlatent", +} +_MAXPOOL_VARIANTS = { + "cellular_multiscale5_maxpool", + "cellular_adaptive_maxpool5", + "cellular_adaptive_maxpool5_rms", +} +_MIXEDPOOL_VARIANTS = { + "cellular_adaptive_mixedpool5_rms", + "cellular_uamp5", + "cellular_uamp5_channelgate", + "cellular_uamp5_varlatent", +} +_RMS_VARIANTS = { + "cellular_adaptive_maxpool5_rms", + "cellular_adaptive_mixedpool5_rms", + "cellular_uamp5", + "cellular_uamp5_channelgate", + "cellular_uamp5_varlatent", +} +_UNET_VARIANTS = { + "cellular_uamp5", + "cellular_uamp5_channelgate", + "cellular_uamp5_varlatent", +} +_CHANNEL_GATE_VARIANTS = {"cellular_uamp5_channelgate"} +_VARLATENT_VARIANTS = {"cellular_uamp5_varlatent"} +_MULTISCALE_VARIANTS = { + "cellular_multiscale5", + "cellular_adaptive_multiscale5", + "cellular_multiscale5_maxpool", + "cellular_adaptive_maxpool5", + "cellular_adaptive_maxpool5_rms", + "cellular_adaptive_mixedpool5_rms", + "cellular_uamp5", + "cellular_uamp5_channelgate", + "cellular_uamp5_varlatent", +} + + +class CellularAttentionLayer(nn.Module): + """Sparse causal attention with optional U-AMP multi-resolution refinement.""" + + def __init__( + self, + num_heads: int, + num_kv_heads: int, + head_dim: int, + feature_dim: int = 64, + variant: str = "cellular_dilated3", + dilations: Iterable[int] = (1, 2, 4, 8, 16, 32, 64, 128), + shifted_window: int = 8, + ): + super().__init__() + if variant not in VARIANTS: + raise ValueError(f"unknown variant {variant!r}; choose from {VARIANTS}") + if min(num_heads, num_kv_heads, head_dim, feature_dim, shifted_window) < 1: + raise ValueError("all dimensions must be positive") + if num_heads % num_kv_heads: + raise ValueError("num_heads must be divisible by num_kv_heads") + dilations = tuple(int(d) for d in dilations) + if not dilations or min(dilations) < 1: + raise ValueError("dilations must contain positive integers") + + self.num_heads = int(num_heads) + self.num_kv_heads = int(num_kv_heads) + self.head_dim = int(head_dim) + self.feature_dim = int(feature_dim) + self.groups = self.num_heads // self.num_kv_heads + self.variant = variant + self.dilations = dilations + self.shifted_window = int(shifted_window) + self.latent_dim = max(8, self.head_dim // 2) + self._last_aux_loss: Tensor | None = None + + q_base = torch.zeros(self.num_heads, self.feature_dim, self.head_dim) + k_base = torch.zeros(self.num_kv_heads, self.feature_dim, self.head_dim) + for h in range(self.num_heads): + nn.init.orthogonal_(q_base[h]) + for h in range(self.num_kv_heads): + nn.init.orthogonal_(k_base[h]) + self.wq = nn.Parameter(q_base) + self.wk = nn.Parameter(k_base) + + self.state_q = nn.Parameter(torch.zeros( + self.num_heads, self.feature_dim, self.head_dim + )) + self.state_q_gate = nn.Parameter(torch.full( + (len(self.dilations), self.num_heads), -2.0 + )) + + max_neighbors = max( + self._max_neighbors_for_variant(variant), self.shifted_window + ) + self.relative_bias = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads, max_neighbors + )) + self.log_temperature = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads + )) + self.step_gate = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads + )) + + if self.uses_adaptive_routing(): + self.route_key = nn.Parameter(torch.empty( + len(self.dilations), self.num_heads, max_neighbors, self.feature_dim + )) + nn.init.normal_(self.route_key, mean=0.0, std=0.02) + self.route_prior = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads, max_neighbors + )) + self.route_strength = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads + )) + else: + self.register_parameter("route_key", None) + self.register_parameter("route_prior", None) + self.register_parameter("route_strength", None) + + if self.uses_maxpool_branch(): + self.pool_mix_logit = nn.Parameter(torch.full( + (len(self.dilations), self.num_heads), -1.5 + )) + self.log_pool_gain = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads + )) + else: + self.register_parameter("pool_mix_logit", None) + self.register_parameter("log_pool_gain", None) + + if self.uses_mixedpool_branch(): + # [attention, max, mean], initialized to preserve the attention path. + initial = torch.tensor([2.0, -1.0, -1.0]).view(1, 1, 3) + self.mixed_pool_logits = nn.Parameter( + initial.expand(len(self.dilations), self.num_heads, 3).clone() + ) + self.log_mixed_pool_gain = nn.Parameter(torch.zeros( + len(self.dilations), self.num_heads, 2 + )) + else: + self.register_parameter("mixed_pool_logits", None) + self.register_parameter("log_mixed_pool_gain", None) + + if self.uses_rms_refinement(): + self.pre_rms_weight = nn.Parameter(torch.ones(self.num_heads, self.head_dim)) + self.post_rms_weight = nn.Parameter(torch.ones(self.num_heads, self.head_dim)) + self.pre_rms_gate = nn.Parameter(torch.full((self.num_heads,), -2.0)) + self.post_rms_gate = nn.Parameter(torch.full((self.num_heads,), -2.0)) + else: + self.register_parameter("pre_rms_weight", None) + self.register_parameter("post_rms_weight", None) + self.register_parameter("pre_rms_gate", None) + self.register_parameter("post_rms_gate", None) + + if self.uses_unet_refinement(): + # Causal two-level U-Net. Downsampling endpoints are 0,2,4,...; + # nearest-left upsampling never uses a future endpoint. + self.unet_pool_logits = nn.Parameter(torch.zeros(2, self.num_heads, 2)) + self.unet_encoder = nn.Parameter(torch.empty( + self.num_heads, self.latent_dim, self.head_dim + )) + self.unet_decoder = nn.Parameter(torch.empty( + self.num_heads, self.head_dim, self.latent_dim + )) + self.unet_swiglu_gate = nn.Parameter(torch.empty( + self.num_heads, self.latent_dim, self.latent_dim + )) + self.unet_swiglu_value = nn.Parameter(torch.empty( + self.num_heads, self.latent_dim, self.latent_dim + )) + for tensor in ( + self.unet_encoder, self.unet_decoder, + self.unet_swiglu_gate, self.unet_swiglu_value, + ): + for h in range(self.num_heads): + nn.init.orthogonal_(tensor[h]) + self.unet_skip1_gate = nn.Parameter(torch.full((self.num_heads,), -1.5)) + self.unet_skip0_gate = nn.Parameter(torch.full((self.num_heads,), -1.5)) + self.unet_output_gate = nn.Parameter(torch.full((self.num_heads,), -2.0)) + else: + self.register_parameter("unet_pool_logits", None) + self.register_parameter("unet_encoder", None) + self.register_parameter("unet_decoder", None) + self.register_parameter("unet_swiglu_gate", None) + self.register_parameter("unet_swiglu_value", None) + self.register_parameter("unet_skip1_gate", None) + self.register_parameter("unet_skip0_gate", None) + self.register_parameter("unet_output_gate", None) + + if self.uses_channel_gate(): + rank = max(4, self.head_dim // 4) + self.channel_down = nn.Parameter(torch.empty( + self.num_heads, rank, self.head_dim + )) + self.channel_up = nn.Parameter(torch.empty( + self.num_heads, self.head_dim, rank + )) + for h in range(self.num_heads): + nn.init.orthogonal_(self.channel_down[h]) + nn.init.orthogonal_(self.channel_up[h]) + self.channel_strength = nn.Parameter(torch.full((self.num_heads,), -2.0)) + else: + self.register_parameter("channel_down", None) + self.register_parameter("channel_up", None) + self.register_parameter("channel_strength", None) + + if self.uses_variational_latent(): + self.var_logvar = nn.Parameter(torch.empty( + self.num_heads, self.latent_dim, self.head_dim + )) + for h in range(self.num_heads): + nn.init.normal_(self.var_logvar[h], mean=0.0, std=0.01) + self.var_logvar_bias = nn.Parameter(torch.full( + (self.num_heads, self.latent_dim), -3.0 + )) + self.var_kl_weight = 1e-4 + else: + self.register_parameter("var_logvar", None) + self.register_parameter("var_logvar_bias", None) + self.var_kl_weight = 0.0 + + eye = torch.eye(self.head_dim).expand(self.num_heads, -1, -1).clone() + self.out_proj = nn.Parameter(eye) + self.final_gate = nn.Parameter(torch.full((self.num_heads,), 2.0)) + self.log_gain = nn.Parameter(torch.zeros(self.num_heads)) + + def uses_adaptive_routing(self) -> bool: + return self.variant in _ADAPTIVE_VARIANTS + + def uses_maxpool_branch(self) -> bool: + return self.variant in _MAXPOOL_VARIANTS + + def uses_mixedpool_branch(self) -> bool: + return self.variant in _MIXEDPOOL_VARIANTS + + def uses_rms_refinement(self) -> bool: + return self.variant in _RMS_VARIANTS + + def uses_unet_refinement(self) -> bool: + return self.variant in _UNET_VARIANTS + + def uses_channel_gate(self) -> bool: + return self.variant in _CHANNEL_GATE_VARIANTS + + def uses_variational_latent(self) -> bool: + return self.variant in _VARLATENT_VARIANTS + + @staticmethod + def _max_neighbors_for_variant(variant: str) -> int: + if variant in ("cellular_local3", "cellular_dilated3"): + return 3 + if variant in ("cellular_dilated5", *_MULTISCALE_VARIANTS): + return 5 + if variant == "cellular_shifted8": + return 8 + raise ValueError(variant) + + @property + def config(self) -> dict: + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "variant": self.variant, + "dilations": list(self.dilations), + "shifted_window": self.shifted_window, + } + + def _offsets(self, step: int) -> tuple[int, ...]: + d = self.dilations[step] + if self.variant == "cellular_local3": + return (0, 1, 2) + if self.variant == "cellular_dilated3": + return (0, d, 2 * d) + if self.variant == "cellular_dilated5": + return (0, d, 2 * d, 3 * d, 4 * d) + if self.variant in _MULTISCALE_VARIANTS: + return tuple(dict.fromkeys((0, 1, d, 2 * d, 4 * d))) + if self.variant == "cellular_shifted8": + return tuple(range(self.shifted_window)) + raise ValueError(self.variant) + + def _valid_mask(self, length: int, step: int, device) -> tuple[Tensor, Tensor]: + offsets = torch.tensor(self._offsets(step), device=device, dtype=torch.long) + pos = torch.arange(length, device=device, dtype=torch.long) + index = pos[:, None] - offsets[None, :] + valid = index >= 0 + + if self.variant == "cellular_shifted8": + shift = 0 if step % 2 == 0 else self.shifted_window // 2 + block_start = ((pos + shift) // self.shifted_window) * self.shifted_window - shift + valid = valid & (index >= block_start[:, None]) + + return index.clamp_min(0), valid + + def _features(self, q: Tensor, k: Tensor) -> tuple[Tensor, Tensor]: + q = F.normalize(q, dim=-1) + k = F.normalize(k, dim=-1) + qf = F.normalize(torch.einsum("bhtd,hfd->bhtf", q, self.wq), dim=-1) + kf = F.normalize(torch.einsum("bhtd,hfd->bhtf", k, self.wk), dim=-1) + return qf, kf.repeat_interleave(self.groups, dim=1) + + @staticmethod + def _rms(x: Tensor, weight: Tensor) -> Tensor: + normed = x * torch.rsqrt(x.square().mean(dim=-1, keepdim=True) + 1e-6) + return normed * weight[None, :, None, :] + + def _adaptive_route_log_prior( + self, query: Tensor, valid: Tensor, step: int, width: int, + ) -> Tensor: + assert self.route_key is not None + assert self.route_prior is not None + assert self.route_strength is not None + prototypes = self.route_key[step, :, :width, :] + route_logits = torch.einsum("bhtf,hwf->bhtw", query, prototypes) + route_logits = route_logits / math.sqrt(self.feature_dim) + route_logits = route_logits + self.route_prior[step, :, :width][None, :, None, :] + route_logits = route_logits.masked_fill( + ~valid[None, None, :, :], float("-inf") + ) + route_log_prob = route_logits.log_softmax(dim=-1) + route_log_prob = route_log_prob.masked_fill( + ~valid[None, None, :, :], 0.0 + ) + strength = F.softplus(self.route_strength[step])[None, :, None, None] + return route_log_prob * strength + + @staticmethod + def _pool_messages(values: Tensor, valid: Tensor) -> tuple[Tensor, Tensor]: + mask = valid[None, None, :, :, None] + maximum = values.masked_fill(~mask, float("-inf")).amax(dim=-2) + count = mask.sum(dim=-2).clamp_min(1).to(values.dtype) + mean = values.masked_fill(~mask, 0.0).sum(dim=-2) / count + return maximum, mean + + def _cellular_step( + self, qf: Tensor, kf: Tensor, state: Tensor, step: int, + ) -> Tensor: + length = state.shape[2] + index, valid = self._valid_mask(length, step, state.device) + keys = kf[:, :, index, :] + values = state[:, :, index, :] + + dynamic_q = torch.einsum("bhtd,hfd->bhtf", state, self.state_q) + mix = self.state_q_gate[step].sigmoid()[None, :, None, None] + query = F.normalize(qf + mix * dynamic_q, dim=-1) + + scores = torch.einsum("bhtf,bhtwf->bhtw", query, keys) + scores = scores / math.sqrt(self.feature_dim) + scores = scores * self.log_temperature[step].clamp(-3, 3).exp()[None, :, None, None] + width = index.shape[1] + scores = scores + self.relative_bias[step, :, :width][None, :, None, :] + scores = scores.masked_fill(~valid[None, None, :, :], float("-inf")) + + if self.uses_adaptive_routing(): + scores = scores + self._adaptive_route_log_prior(query, valid, step, width) + + weights = scores.softmax(dim=-1) + attention_message = torch.einsum("bhtw,bhtwd->bhtd", weights, values) + message = attention_message + + if self.uses_mixedpool_branch(): + assert self.mixed_pool_logits is not None + assert self.log_mixed_pool_gain is not None + maximum, mean = self._pool_messages(values, valid) + gains = self.log_mixed_pool_gain[step].clamp(-3, 3).exp() + maximum = maximum * gains[:, 0][None, :, None, None] + mean = mean * gains[:, 1][None, :, None, None] + mixture = self.mixed_pool_logits[step].softmax(dim=-1) + message = ( + mixture[:, 0][None, :, None, None] * attention_message + + mixture[:, 1][None, :, None, None] * maximum + + mixture[:, 2][None, :, None, None] * mean + ) + elif self.uses_maxpool_branch(): + assert self.pool_mix_logit is not None + assert self.log_pool_gain is not None + maximum, _ = self._pool_messages(values, valid) + gain = self.log_pool_gain[step].clamp(-3, 3).exp()[None, :, None, None] + pooled = maximum * gain + pool_mix = self.pool_mix_logit[step].sigmoid()[None, :, None, None] + message = message + pool_mix * (pooled - message) + + gate = self.step_gate[step].sigmoid()[None, :, None, None] + return state + gate * (message - state) + + def _causal_stride2_pool(self, x: Tensor, level: int) -> Tensor: + assert self.unet_pool_logits is not None + length = x.shape[2] + endpoints = torch.arange(0, length, 2, device=x.device) + previous = (endpoints - 1).clamp_min(0) + pair = torch.stack((x[:, :, previous, :], x[:, :, endpoints, :]), dim=-2) + maximum = pair.amax(dim=-2) + mean = pair.mean(dim=-2) + mix = self.unet_pool_logits[level].softmax(dim=-1) + return ( + mix[:, 0][None, :, None, None] * maximum + + mix[:, 1][None, :, None, None] * mean + ) + + @staticmethod + def _causal_upsample(x: Tensor, target_length: int) -> Tensor: + # Reduced element j represents an endpoint <= 2*j. floor(t/2) is + # therefore always current/past relative to target token t. + index = torch.arange(target_length, device=x.device) // 2 + return x[:, :, index.clamp_max(x.shape[2] - 1), :] + + def _unet_refine(self, state: Tensor) -> Tensor: + assert self.unet_encoder is not None + assert self.unet_decoder is not None + assert self.unet_swiglu_gate is not None + assert self.unet_swiglu_value is not None + assert self.unet_skip1_gate is not None + assert self.unet_skip0_gate is not None + assert self.unet_output_gate is not None + + e0 = state + e1 = self._causal_stride2_pool(e0, 0) + e2 = self._causal_stride2_pool(e1, 1) + mu = torch.einsum("bhtd,hld->bhtl", e2, self.unet_encoder) + + if self.uses_variational_latent(): + assert self.var_logvar is not None + assert self.var_logvar_bias is not None + logvar = torch.einsum("bhtd,hld->bhtl", e2, self.var_logvar) + logvar = (logvar + self.var_logvar_bias[None, :, None, :]).clamp(-8, 4) + # Deterministic mean path keeps evaluation reproducible; KL still + # regularizes a variational latent family during fitting. + self._last_aux_loss = self.var_kl_weight * 0.5 * ( + mu.square() + logvar.exp() - 1.0 - logvar + ).mean() + else: + self._last_aux_loss = None + + gate_part = torch.einsum("bhtl,hlm->bhtm", mu, self.unet_swiglu_gate) + value_part = torch.einsum("bhtl,hlm->bhtm", mu, self.unet_swiglu_value) + latent = F.silu(gate_part) * value_part + decoded = torch.einsum("bhtl,hdl->bhtd", latent, self.unet_decoder) + + up1 = self._causal_upsample(decoded, e1.shape[2]) + g1 = self.unet_skip1_gate.sigmoid()[None, :, None, None] + d1 = e1 + g1 * (up1 - e1) + up0 = self._causal_upsample(d1, e0.shape[2]) + g0 = self.unet_skip0_gate.sigmoid()[None, :, None, None] + d0 = e0 + g0 * (up0 - e0) + gout = self.unet_output_gate.sigmoid()[None, :, None, None] + return state + gout * (d0 - state) + + def _channel_gate(self, state: Tensor) -> Tensor: + assert self.channel_down is not None + assert self.channel_up is not None + assert self.channel_strength is not None + time = torch.arange(1, state.shape[2] + 1, device=state.device, dtype=state.dtype) + prefix_mean = state.cumsum(dim=2) / time[None, None, :, None] + hidden = F.silu(torch.einsum("bhtd,hrd->bhtr", prefix_mean, self.channel_down)) + logits = torch.einsum("bhtr,hdr->bhtd", hidden, self.channel_up) + strength = self.channel_strength.sigmoid()[None, :, None, None] + return state * (1.0 + strength * torch.tanh(logits)) + + def auxiliary_loss(self) -> Tensor: + if self._last_aux_loss is None: + return self.wq.sum() * 0.0 + return self._last_aux_loss + + def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor: + if q.ndim != 4 or k.ndim != 4 or v.ndim != 4: + raise ValueError("expected Q/K/V as [batch, heads, time, dim]") + if k.shape != v.shape: + raise ValueError("K and V must have the same shape") + if q.shape[0] != k.shape[0] or q.shape[2:] != k.shape[2:]: + raise ValueError("Q/K/V batch, time and head_dim must match") + if q.shape[1] != self.num_heads or k.shape[1] != self.num_kv_heads: + raise ValueError("Q/K head counts do not match layer configuration") + if q.shape[-1] != self.head_dim or q.shape[2] < 1: + raise ValueError("head_dim mismatch or empty sequence") + + self._last_aux_loss = None + q, k, v = (x.to(self.wq.dtype) for x in (q, k, v)) + qf, kf = self._features(q, k) + base = v.repeat_interleave(self.groups, dim=1) + state = base + + if self.uses_rms_refinement(): + assert self.pre_rms_weight is not None and self.pre_rms_gate is not None + normed = self._rms(state, self.pre_rms_weight) + gate = self.pre_rms_gate.sigmoid()[None, :, None, None] + state = state + gate * (normed - state) + + for step in range(len(self.dilations)): + state = self._cellular_step(qf, kf, state, step) + + if self.uses_unet_refinement(): + state = self._unet_refine(state) + if self.uses_channel_gate(): + state = self._channel_gate(state) + + if self.uses_rms_refinement(): + assert self.post_rms_weight is not None and self.post_rms_gate is not None + normed = self._rms(state, self.post_rms_weight) + gate = self.post_rms_gate.sigmoid()[None, :, None, None] + state = state + gate * (normed - state) + + projected = torch.einsum("bhtd,hde->bhte", state, self.out_proj) + final_gate = self.final_gate.sigmoid()[None, :, None, None] + output = base + final_gate * (projected - base) + return output * self.log_gain.clamp(-4, 4).exp()[None, :, None, None] + + def receptive_field_tokens(self) -> int: + reach = 0 + for step in range(len(self.dilations)): + reach += max(self._offsets(step)) + return reach + 1 + + def max_score_pairs(self, context: int) -> int: + total = 0 + device = self.wq.device + for step in range(len(self.dilations)): + _, valid = self._valid_mask(context, step, device) + total += int(valid.sum().item()) + return total + + def max_neighbors_per_step(self) -> int: + return max(len(self._offsets(step)) for step in range(len(self.dilations))) diff --git a/src/tinycenn_lm/cenn.py b/src/tinycenn_lm/cenn.py new file mode 100644 index 0000000000000000000000000000000000000000..c59b36f18ea70b095e21674f677ac3d1da442603 --- /dev/null +++ b/src/tinycenn_lm/cenn.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +from dataclasses import asdict, dataclass + +import torch +from torch import Tensor, nn +import torch.nn.functional as F + + +@dataclass(frozen=True) +class CeNNConfig: + """Configuration for the causal 1-D Cellular Neural Network adapter.""" + + hidden_size: int = 192 + kernel_size: int = 3 + expansion: int = 4 + steps: int = 4 + dilations: tuple[int, ...] = (1, 2, 4, 8) + rms_norm_eps: float = 1e-5 + dropout: float = 0.0 + + def validate(self) -> None: + if self.hidden_size <= 0: + raise ValueError("hidden_size must be positive") + if self.kernel_size < 2: + raise ValueError("kernel_size must be >= 2") + if self.expansion <= 0: + raise ValueError("expansion must be positive") + if self.steps <= 0: + raise ValueError("steps must be positive") + if not self.dilations or any(d <= 0 for d in self.dilations): + raise ValueError("dilations must contain positive integers") + if not 0.0 <= self.dropout < 1.0: + raise ValueError("dropout must be in [0, 1)") + + def to_dict(self) -> dict: + data = asdict(self) + data["dilations"] = list(self.dilations) + return data + + @classmethod + def from_dict(cls, data: dict) -> "CeNNConfig": + data = dict(data) + if "dilations" in data: + data["dilations"] = tuple(data["dilations"]) + return cls(**data) + + +class StableRMSNorm(nn.Module): + """Small RMSNorm with fp32 variance accumulation for training stability.""" + + def __init__(self, hidden_size: int, eps: float = 1e-5) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.eps = eps + + def forward(self, x: Tensor) -> Tensor: + dtype = x.dtype + variance = x.float().pow(2).mean(dim=-1, keepdim=True) + x = x * torch.rsqrt(variance + self.eps).to(dtype=dtype) + return x * self.weight.to(dtype=dtype) + + +class CausalDepthwiseNeighborhood(nn.Module): + """A causal, channel-wise cellular neighborhood operator. + + The same kernel is reused at every recurrent step. Only the dilation changes, + which grows the receptive field while keeping the parameter count fixed. + """ + + def __init__(self, hidden_size: int, kernel_size: int) -> None: + super().__init__() + self.hidden_size = hidden_size + self.kernel_size = kernel_size + self.weight = nn.Parameter(torch.empty(hidden_size, 1, kernel_size)) + nn.init.kaiming_uniform_(self.weight, a=5**0.5) + + def forward(self, x: Tensor, dilation: int) -> Tensor: + if x.ndim != 3: + raise ValueError(f"expected [batch, seq, hidden], got {tuple(x.shape)}") + if x.shape[-1] != self.hidden_size: + raise ValueError( + f"expected hidden size {self.hidden_size}, got {x.shape[-1]}" + ) + left_pad = dilation * (self.kernel_size - 1) + y = x.transpose(1, 2) + y = F.pad(y, (left_pad, 0)) + y = F.conv1d( + y, + self.weight, + bias=None, + stride=1, + padding=0, + dilation=dilation, + groups=self.hidden_size, + ) + return y.transpose(1, 2) + + +class SharedCeNNCell(nn.Module): + """Shared recurrent CeNN cell using causal local mixing + gated SwiGLU update.""" + + def __init__(self, config: CeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + inner = config.hidden_size * config.expansion + + self.norm = StableRMSNorm(config.hidden_size, config.rms_norm_eps) + self.neighborhood = CausalDepthwiseNeighborhood( + config.hidden_size, config.kernel_size + ) + self.in_proj = nn.Linear(config.hidden_size, inner * 2, bias=False) + self.gate_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=True) + self.out_proj = nn.Linear(inner, config.hidden_size, bias=False) + self.dropout = nn.Dropout(config.dropout) + + nn.init.zeros_(self.out_proj.weight) + nn.init.constant_(self.gate_proj.bias, -1.0) + + def forward(self, state: Tensor, dilation: int, step_scale: float) -> Tensor: + x = self.norm(state) + local = self.neighborhood(x, dilation=dilation) + a, b = self.in_proj(local).chunk(2, dim=-1) + update = F.silu(a) * b + update = self.out_proj(update) + update = self.dropout(update) + gate = torch.sigmoid(self.gate_proj(local)) + return state + (step_scale * gate * update) + + +class FastCeNNCore(nn.Module): + """Iterates one shared CeNN cell multiple times with a dilation schedule.""" + + def __init__(self, config: CeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + self.cell = SharedCeNNCell(config) + + def forward(self, hidden_states: Tensor) -> Tensor: + initial = hidden_states + state = hidden_states + step_scale = self.config.steps ** -0.5 + for step in range(self.config.steps): + dilation = self.config.dilations[step % len(self.config.dilations)] + state = self.cell(state, dilation=dilation, step_scale=step_scale) + return state - initial + + @property + def receptive_field(self) -> int: + radius = sum( + self.config.dilations[i % len(self.config.dilations)] + for i in range(self.config.steps) + ) + return 1 + (self.config.kernel_size - 1) * radius diff --git a/src/tinycenn_lm/colab_live_backup.py b/src/tinycenn_lm/colab_live_backup.py new file mode 100644 index 0000000000000000000000000000000000000000..cf80ec5e66a3282b682246c0f498e19b2c0d1ca6 --- /dev/null +++ b/src/tinycenn_lm/colab_live_backup.py @@ -0,0 +1,459 @@ +from __future__ import annotations + +import json +import os +import shlex +import subprocess +import sys +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Callable + +from .hf_persistence import redact_secrets, utc_run_id + +_SMALL_SUFFIXES = {".json", ".jsonl", ".csv", ".txt", ".md", ".log", ".yaml", ".yml"} + + +def is_tinycenn_training_command(cmd: Any) -> bool: + if not isinstance(cmd, (list, tuple)): + return False + parts = [str(x) for x in cmd] + for part in parts: + name = Path(part).name.lower() + if name.startswith("train_") and name.endswith(".py"): + return True + return False + + +def training_script_name(cmd: list[str] | tuple[str, ...]) -> str: + for part in cmd: + name = Path(str(part)).name + if name.lower().startswith("train_") and name.lower().endswith(".py"): + return name[:-3] + return "training" + + +def output_dir_from_command(cmd: list[str] | tuple[str, ...], cwd: str | Path | None = None) -> Path | None: + parts = [str(x) for x in cmd] + for i, part in enumerate(parts): + if part == "--output-dir" and i + 1 < len(parts): + path = Path(parts[i + 1]) + if not path.is_absolute() and cwd is not None: + path = Path(cwd) / path + return path + if part.startswith("--output-dir="): + path = Path(part.split("=", 1)[1]) + if not path.is_absolute() and cwd is not None: + path = Path(cwd) / path + return path + return None + + +def _small_files(folder: Path) -> list[Path]: + if not folder.exists(): + return [] + return [ + p for p in folder.rglob("*") + if p.is_file() + and p.suffix.lower() in _SMALL_SUFFIXES + and p.stat().st_size <= 10 * 1024 * 1024 + and ".hf_run_archive" not in p.parts + and ".hf_live_redacted" not in p.parts + ] + + +def _retry_required( + action: Callable[[], Any], + *, + label: str, + attempts: int = 3, + delay_seconds: float = 5.0, +) -> Any: + """Run one mandatory Hugging Face operation with bounded retries.""" + last_error: Exception | None = None + for attempt in range(1, max(attempts, 1) + 1): + try: + return action() + except Exception as exc: # network/API failures must fail closed after retries + last_error = exc + print( + f"[TinyCeNN][BACKUP] {label} failed " + f"(attempt {attempt}/{attempts}): {exc}", + flush=True, + ) + if attempt < attempts: + time.sleep(delay_seconds) + raise RuntimeError( + f"Mandatory Hugging Face backup failed during {label} after {attempts} attempts" + ) from last_error + + +def _upload_live_required( + api, + *, + repo_id: str, + run_id: str, + output_dir: Path | None, + log_file: Path, + include_checkpoint: bool = True, +) -> None: + """Persist the live log and reports, optionally mirroring checkpoint files.""" + if log_file.exists(): + clean_log = log_file.with_name("train_redacted.log") + clean_log.write_text( + redact_secrets(log_file.read_text(encoding="utf-8", errors="replace")), + encoding="utf-8", + ) + _retry_required( + lambda: api.upload_file( + repo_id=repo_id, + repo_type="model", + path_or_fileobj=str(clean_log), + path_in_repo=f"runs/{run_id}/train.log", + commit_message=f"Live backup {run_id}", + ), + label="live log upload", + ) + + if output_dir is None or not output_dir.exists(): + return + + # Upload small human-readable artifacts separately after redaction. + for path in _small_files(output_dir): + rel = path.relative_to(output_dir) + upload_path = path + try: + clean_dir = output_dir / ".hf_live_redacted" + dst = clean_dir / rel + dst.parent.mkdir(parents=True, exist_ok=True) + dst.write_text( + redact_secrets(path.read_text(encoding="utf-8", errors="replace")), + encoding="utf-8", + ) + upload_path = dst + except Exception: + upload_path = path + + _retry_required( + lambda p=upload_path, r=rel: api.upload_file( + repo_id=repo_id, + repo_type="model", + path_or_fileobj=str(p), + path_in_repo=f"runs/{run_id}/artifacts/{r.as_posix()}", + commit_message=f"Live results {run_id}", + ), + label=f"artifact upload: {rel.as_posix()}", + ) + + if not include_checkpoint: + return + + # Mandatory periodic checkpoint mirror. This is disabled when the output + # directory already existed before the run, because otherwise stale weights from + # an earlier execution can trigger a large upload unrelated to the current run. + checkpoint_files = [ + p for p in output_dir.rglob("*") + if p.is_file() + and ".hf_run_archive" not in p.parts + and ".hf_live_redacted" not in p.parts + ] + if checkpoint_files: + _retry_required( + lambda: api.upload_folder( + repo_id=repo_id, + repo_type="model", + folder_path=str(output_dir), + path_in_repo=f"runs/{run_id}/checkpoint", + ignore_patterns=[".hf_run_archive/**", ".hf_live_redacted/**"], + commit_message=f"Live checkpoint backup {run_id}", + ), + label="live checkpoint mirror", + ) + + +def _backup_metadata(cmd, output_dir: Path | None, run_id: str) -> dict[str, Any]: + return { + "run_id": run_id, + "created_utc": datetime.now(timezone.utc).isoformat(), + "command": [redact_secrets(str(x)) for x in cmd], + "output_dir": str(output_dir) if output_dir else None, + "status": "running", + "backup_policy": "mandatory-fail-closed", + } + + +def _format_elapsed(seconds: float) -> str: + seconds = max(0, int(seconds)) + hours, remainder = divmod(seconds, 3600) + minutes, secs = divmod(remainder, 60) + if hours: + return f"{hours:d}h {minutes:02d}m {secs:02d}s" + return f"{minutes:d}m {secs:02d}s" + + +def install_colab_training_backup(*, interval_seconds: int = 180) -> bool: + """Stream TinyCeNN Colab training and require a private Hugging Face backup. + + Every ``subprocess.run([... train_*.py ...])`` call in Colab is converted to a + line-streaming ``Popen`` execution. Before the child process starts, a valid + Hugging Face login and a writable private ``TinyCeNN-LM-Colab-Backups`` repo are + required. Live logs and reports are mirrored periodically, along with checkpoints + when the output directory is new for the current run. If a mandatory sync fails + after retries, training is terminated. A successful run must finish by uploading + its complete checkpoint. A failed trainer still uploads its log/status promptly, + but does not block on a large or stale checkpoint directory. + """ + if not (os.environ.get("COLAB_RELEASE_TAG") or os.environ.get("COLAB_GPU") or Path("/content").exists()): + return False + if getattr(subprocess.run, "_tinycenn_live_backup", False): + return True + + original_run = subprocess.run + original_popen = subprocess.Popen + + def run_with_backup(cmd, *args, **kwargs): + if not is_tinycenn_training_command(cmd): + return original_run(cmd, *args, **kwargs) + + unsupported = {"input", "capture_output", "stdout", "stderr", "timeout"} & set(kwargs) + if unsupported or args: + return original_run(cmd, *args, **kwargs) + + cwd = kwargs.pop("cwd", None) + requested_env = kwargs.pop("env", None) + check = bool(kwargs.pop("check", False)) + if kwargs: + return original_run(cmd, cwd=cwd, env=requested_env, check=check, **kwargs) + + child_env = os.environ.copy() + if requested_env is not None: + child_env.update({str(k): str(v) for k, v in requested_env.items()}) + child_env["PYTHONUNBUFFERED"] = "1" + + script = training_script_name(cmd) + run_id = utc_run_id(script[:32]) + output_dir = output_dir_from_command(cmd, cwd=cwd) + output_preexisting = bool(output_dir is not None and output_dir.exists()) + work_root = Path(cwd).resolve() if cwd is not None else Path.cwd().resolve() + live_root = work_root / ".colab_live_backup" / run_id + live_root.mkdir(parents=True, exist_ok=True) + log_file = live_root / "train.log" + meta_file = live_root / "run_status.json" + metadata = _backup_metadata(cmd, output_dir, run_id) + metadata["output_dir_preexisting"] = output_preexisting + meta_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + + # Mandatory preflight. Do not spend GPU time unless the remote safety target + # is authenticated, private, writable, and has accepted the run metadata. + try: + from huggingface_hub import HfApi, get_token + except Exception as exc: + raise RuntimeError( + "Mandatory Hugging Face backup requires huggingface_hub. " + "Install it before training." + ) from exc + + token = get_token() or os.environ.get("HF_TOKEN") + if not token: + raise RuntimeError( + "Mandatory Hugging Face backup is enabled. Log in first or provide " + "HF_TOKEN (in Colab: add HF_TOKEN to Secrets and run the login cell)." + ) + + api = HfApi(token=token) + identity = _retry_required(api.whoami, label="Hugging Face authentication") + user = identity["name"] + backup_repo = f"{user}/TinyCeNN-LM-Colab-Backups" + _retry_required( + lambda: api.create_repo( + backup_repo, + repo_type="model", + private=True, + exist_ok=True, + ), + label="private backup repository preflight", + ) + _retry_required( + lambda: api.upload_file( + repo_id=backup_repo, + repo_type="model", + path_or_fileobj=str(meta_file), + path_in_repo=f"runs/{run_id}/run_status.json", + commit_message=f"Start mandatory backup {run_id}", + ), + label="initial backup metadata upload", + ) + + command_text = " ".join(shlex.quote(str(x)) for x in cmd) + print("\n" + "=" * 88, flush=True) + print(f"[TinyCeNN][START] {script}", flush=True) + print(f"[TinyCeNN][COMMAND] {command_text}", flush=True) + if output_dir is not None: + print(f"[TinyCeNN][OUTPUT] {output_dir}", flush=True) + if output_preexisting: + print( + "[TinyCeNN][BACKUP] output directory already exists; periodic backup " + "will save logs/reports only and avoid re-uploading stale checkpoint files.", + flush=True, + ) + print(f"[TinyCeNN][LOCAL LOG] {log_file}", flush=True) + print( + f"[TinyCeNN][BACKUP REQUIRED] " + f"https://huggingface.co/{backup_repo}/tree/main/runs/{run_id}", + flush=True, + ) + print( + f"[TinyCeNN][BACKUP POLICY] mandatory sync every {interval_seconds}s; " + "training aborts if backup cannot be persisted after retries", + flush=True, + ) + print("[TinyCeNN][LIVE] streaming training output...", flush=True) + print("-" * 88, flush=True) + + started = time.monotonic() + proc = original_popen( + cmd, + cwd=cwd, + env=child_env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + bufsize=1, + ) + last_sync = started + backup_failure: Exception | None = None + + with log_file.open("a", encoding="utf-8") as log: + assert proc.stdout is not None + for line in proc.stdout: + sys.stdout.write(line) + sys.stdout.flush() + log.write(line) + log.flush() + now = time.monotonic() + if now - last_sync >= interval_seconds: + try: + _upload_live_required( + api, + repo_id=backup_repo, + run_id=run_id, + output_dir=output_dir, + log_file=log_file, + include_checkpoint=not output_preexisting, + ) + print( + f"[TinyCeNN][BACKUP OK] live state persisted at " + f"{_format_elapsed(now - started)}", + flush=True, + ) + last_sync = now + except Exception as exc: + backup_failure = exc + print( + f"[TinyCeNN][BACKUP FATAL] {exc}. Terminating training to protect the run.", + flush=True, + ) + proc.terminate() + try: + proc.wait(timeout=30) + except subprocess.TimeoutExpired: + proc.kill() + proc.wait() + break + + returncode = proc.wait() + elapsed = time.monotonic() - started + + if backup_failure is not None: + metadata["status"] = "failed-backup" + metadata["returncode"] = returncode + metadata["elapsed_seconds"] = elapsed + metadata["finished_utc"] = datetime.now(timezone.utc).isoformat() + metadata["backup_error"] = redact_secrets(str(backup_failure)) + meta_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + # Best effort only for the failure marker; the triggering sync already + # proved that the remote is unavailable. + try: + api.upload_file( + repo_id=backup_repo, + repo_type="model", + path_or_fileobj=str(meta_file), + path_in_repo=f"runs/{run_id}/run_status.json", + commit_message=f"Backup failure {run_id}", + ) + except Exception: + pass + raise RuntimeError( + "Training was terminated because mandatory Hugging Face backup could not be maintained." + ) from backup_failure + + metadata["status"] = "completed" if returncode == 0 else "failed-training" + metadata["returncode"] = returncode + metadata["elapsed_seconds"] = elapsed + metadata["finished_utc"] = datetime.now(timezone.utc).isoformat() + meta_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + + # Always preserve the run's log, small reports and final status. On failure, + # do not upload the checkpoint directory: it may be a stale fixed-name folder + # from an earlier run and was the source of long post-crash Colab hangs. + _upload_live_required( + api, + repo_id=backup_repo, + run_id=run_id, + output_dir=output_dir, + log_file=log_file, + include_checkpoint=False, + ) + _retry_required( + lambda: api.upload_file( + repo_id=backup_repo, + repo_type="model", + path_or_fileobj=str(meta_file), + path_in_repo=f"runs/{run_id}/run_status.json", + commit_message=f"Finish mandatory backup {run_id}", + ), + label="final status upload", + ) + + if returncode == 0 and output_dir is not None and output_dir.exists(): + print("[TinyCeNN][BACKUP] uploading successful final checkpoint...", flush=True) + _retry_required( + lambda: api.upload_folder( + repo_id=backup_repo, + repo_type="model", + folder_path=str(output_dir), + path_in_repo=f"runs/{run_id}/checkpoint", + ignore_patterns=[".hf_run_archive/**", ".hf_live_redacted/**"], + commit_message=f"Final checkpoint backup {run_id}", + ), + label="final checkpoint upload", + ) + elif returncode != 0: + print( + "[TinyCeNN][BACKUP] trainer failed; log/status backed up, large checkpoint upload skipped.", + flush=True, + ) + + print("-" * 88, flush=True) + print( + f"[TinyCeNN][BACKUP COMPLETE] private Hugging Face backup committed for {run_id}", + flush=True, + ) + if returncode == 0: + print(f"[TinyCeNN][DONE] {script} completed in {_format_elapsed(elapsed)}", flush=True) + else: + print( + f"[TinyCeNN][FAILED] {script} exited with code {returncode} after {_format_elapsed(elapsed)}", + flush=True, + ) + print("=" * 88 + "\n", flush=True) + + completed = subprocess.CompletedProcess(cmd, returncode) + if check and returncode: + raise subprocess.CalledProcessError(returncode, cmd) + return completed + + run_with_backup._tinycenn_live_backup = True + subprocess.run = run_with_backup + return True \ No newline at end of file diff --git a/src/tinycenn_lm/direct_colab_backup.py b/src/tinycenn_lm/direct_colab_backup.py new file mode 100644 index 0000000000000000000000000000000000000000..7476cff2e48bdcb143f3c90f6a4f78ac31b38f6a --- /dev/null +++ b/src/tinycenn_lm/direct_colab_backup.py @@ -0,0 +1,276 @@ +from __future__ import annotations + +import atexit +import json +import os +import sys +import threading +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from .colab_live_backup import _retry_required, _upload_live_required, output_dir_from_command +from .hf_persistence import redact_secrets, utc_run_id + +_INSTALLED = False + + +def _is_colab() -> bool: + return bool(os.environ.get("COLAB_RELEASE_TAG") or os.environ.get("COLAB_GPU") or Path("/content").exists()) + + +def _training_script_name() -> str | None: + name = Path(sys.argv[0]).name + if name.lower().startswith("train_") and name.lower().endswith(".py"): + return name[:-3] + return None + + +def _candidate_parent_roots(root: Path) -> list[Path]: + """Return likely repository roots used by notebook-side backup wrappers.""" + candidates = [root.resolve()] + + # A trainer is normally /content/TinyCeNN-LM/scripts/train_*.py. Deriving the + # repository from argv makes detection independent of the notebook's cwd. + try: + script_path = Path(sys.argv[0]).resolve() + if script_path.parent.name == "scripts": + candidates.append(script_path.parent.parent) + except Exception: + pass + + # Older notebook wrappers used this fixed Colab repository path. + candidates.append(Path("/content/TinyCeNN-LM")) + + configured = os.environ.get("TINYCENN_PARENT_BACKUP_ROOT") + if configured: + candidates.append(Path(configured)) + + unique: list[Path] = [] + seen: set[str] = set() + for candidate in candidates: + key = str(candidate.resolve()) + if key not in seen: + seen.add(key) + unique.append(candidate.resolve()) + return unique + + +def _parent_backup_is_active(script: str, root: Path) -> bool: + """Detect a recent parent-wrapper run so the child does not duplicate uploads.""" + if os.environ.get("TINYCENN_PARENT_BACKUP_ACTIVE") == "1": + return True + + now = time.time() + for candidate in _candidate_parent_roots(root): + backup_root = candidate / ".colab_live_backup" + if not backup_root.exists(): + continue + for meta in backup_root.glob("*/run_status.json"): + try: + if now - meta.stat().st_mtime > 10 * 60: + continue + data = json.loads(meta.read_text(encoding="utf-8")) + if data.get("status") != "running": + continue + command = " ".join(str(x) for x in data.get("command", [])) + if script in command: + return True + except Exception: + continue + return False + + +class _TeeStream: + def __init__(self, primary, log_file): + self._primary = primary + self._log = log_file + + def write(self, text): + result = self._primary.write(text) + self._primary.flush() + self._log.write(text) + self._log.flush() + return result + + def flush(self): + self._primary.flush() + self._log.flush() + + def __getattr__(self, name): + return getattr(self._primary, name) + + +def install_direct_training_backup(*, interval_seconds: int = 180) -> bool: + """Mandatory trainer-side HF backup when no parent Colab wrapper is active. + + This is a safety fallback for notebooks that launch ``train_*.py`` before the + notebook kernel imports ``tinycenn_lm``. Normal notebooks are backed up by the + parent wrapper; this function detects that active run and stays out of the way. + """ + global _INSTALLED + if _INSTALLED or not _is_colab(): + return _INSTALLED + + script = _training_script_name() + if script is None: + return False + + root = Path.cwd().resolve() + if _parent_backup_is_active(script, root): + print("[TinyCeNN][BACKUP] parent mandatory backup detected; child fallback not needed.", flush=True) + _INSTALLED = True + return True + + try: + from huggingface_hub import HfApi, get_token + except Exception as exc: + raise RuntimeError( + "Mandatory Hugging Face backup requires huggingface_hub. Install it before training." + ) from exc + + token = get_token() or os.environ.get("HF_TOKEN") + if not token: + raise RuntimeError( + "Mandatory Hugging Face backup is enabled. Add HF_TOKEN to Colab Secrets, " + "run the Hugging Face login cell, then start training." + ) + + api = HfApi(token=token) + identity = _retry_required(api.whoami, label="trainer-side Hugging Face authentication") + user = identity["name"] + repo_id = f"{user}/TinyCeNN-LM-Colab-Backups" + _retry_required( + lambda: api.create_repo(repo_id, repo_type="model", private=True, exist_ok=True), + label="trainer-side private backup repository preflight", + ) + + run_id = utc_run_id(f"{script}-direct"[:32]) + output_dir = output_dir_from_command(sys.argv, cwd=root) + output_preexisting = bool(output_dir is not None and output_dir.exists()) + live_root = root / ".colab_live_backup" / run_id + live_root.mkdir(parents=True, exist_ok=True) + log_path = live_root / "train.log" + meta_path = live_root / "run_status.json" + metadata: dict[str, Any] = { + "run_id": run_id, + "created_utc": datetime.now(timezone.utc).isoformat(), + "command": [redact_secrets(str(x)) for x in sys.argv], + "output_dir": str(output_dir) if output_dir else None, + "output_dir_preexisting": output_preexisting, + "status": "running", + "backup_policy": "mandatory-fail-closed-direct-trainer-fallback", + } + meta_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + _retry_required( + lambda: api.upload_file( + repo_id=repo_id, + repo_type="model", + path_or_fileobj=str(meta_path), + path_in_repo=f"runs/{run_id}/run_status.json", + commit_message=f"Start mandatory direct backup {run_id}", + ), + label="trainer-side initial metadata upload", + ) + + log_handle = log_path.open("a", encoding="utf-8", buffering=1) + sys.stdout = _TeeStream(sys.stdout, log_handle) + sys.stderr = _TeeStream(sys.stderr, log_handle) + + print( + f"[TinyCeNN][BACKUP REQUIRED][DIRECT] " + f"https://huggingface.co/{repo_id}/tree/main/runs/{run_id}", + flush=True, + ) + print( + f"[TinyCeNN][BACKUP POLICY][DIRECT] mandatory sync every {interval_seconds}s; " + "process exits if backup cannot be persisted after retries", + flush=True, + ) + if output_preexisting: + print( + "[TinyCeNN][BACKUP][DIRECT] output directory already existed; periodic " + "checkpoint mirroring is disabled to avoid uploading stale files.", + flush=True, + ) + + stop_event = threading.Event() + started = time.monotonic() + + def heartbeat() -> None: + while not stop_event.wait(interval_seconds): + try: + _upload_live_required( + api, + repo_id=repo_id, + run_id=run_id, + output_dir=output_dir, + log_file=log_path, + include_checkpoint=not output_preexisting, + ) + print("[TinyCeNN][BACKUP OK][DIRECT] live state persisted.", flush=True) + except Exception as exc: + print( + f"[TinyCeNN][BACKUP FATAL][DIRECT] {exc}. " + "Stopping training because remote safety cannot be guaranteed.", + flush=True, + ) + os._exit(74) + + worker = threading.Thread(target=heartbeat, name="tinycenn-hf-backup", daemon=True) + worker.start() + + def finalize() -> None: + stop_event.set() + metadata["status"] = "process-exit" + metadata["finished_utc"] = datetime.now(timezone.utc).isoformat() + metadata["elapsed_seconds"] = time.monotonic() - started + meta_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + try: + # The explicit final upload below is the only full-folder upload here. + # This avoids sending the same checkpoint twice at process exit. + _upload_live_required( + api, + repo_id=repo_id, + run_id=run_id, + output_dir=output_dir, + log_file=log_path, + include_checkpoint=False, + ) + _retry_required( + lambda: api.upload_file( + repo_id=repo_id, + repo_type="model", + path_or_fileobj=str(meta_path), + path_in_repo=f"runs/{run_id}/run_status.json", + commit_message=f"Finish mandatory direct backup {run_id}", + ), + label="trainer-side final status upload", + ) + if output_dir is not None and output_dir.exists(): + _retry_required( + lambda: api.upload_folder( + repo_id=repo_id, + repo_type="model", + folder_path=str(output_dir), + path_in_repo=f"runs/{run_id}/checkpoint", + ignore_patterns=[".hf_run_archive/**", ".hf_live_redacted/**"], + commit_message=f"Final direct checkpoint backup {run_id}", + ), + label="trainer-side final checkpoint upload", + ) + print("[TinyCeNN][BACKUP COMPLETE][DIRECT] final backup committed.", flush=True) + except Exception as exc: + print( + f"[TinyCeNN][BACKUP FATAL][DIRECT] final backup failed: {exc}", + flush=True, + ) + try: + log_handle.flush() + finally: + os._exit(75) + + atexit.register(finalize) + _INSTALLED = True + return True diff --git a/src/tinycenn_lm/distill_utils.py b/src/tinycenn_lm/distill_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..c12f12b37a0ba8b4866d3b95dc0e52a56ceec497 --- /dev/null +++ b/src/tinycenn_lm/distill_utils.py @@ -0,0 +1,232 @@ +from __future__ import annotations + +import hashlib +import math +import random +import struct +from collections.abc import Iterable, Iterator +from contextlib import nullcontext + +import torch +import torch.nn.functional as F + + +def holdout_bucket(text: str, buckets: int = 1000) -> int: + digest = hashlib.blake2b(text.encode("utf-8", errors="ignore"), digest_size=8).digest() + return int.from_bytes(digest, "little") % buckets + + +def partition_rows(dataset: Iterable[dict], text_field: str, *, validation: bool) -> Iterator[dict]: + """Deterministic document-level split: 99% train / 1% validation.""" + for example in dataset: + text = example.get(text_field) + if not isinstance(text, str) or not text.strip(): + continue + is_validation = holdout_bucket(text) >= 990 + if is_validation == validation: + yield example + + +def buffered_shuffle(rows: Iterable[dict], *, buffer_size: int, seed: int) -> Iterator[dict]: + """Deterministic bounded-memory shuffle for streaming datasets.""" + if buffer_size <= 1: + yield from rows + return + + rng = random.Random(seed) + buffer: list[dict] = [] + for row in rows: + if len(buffer) < buffer_size: + buffer.append(row) + continue + index = rng.randrange(len(buffer)) + yield buffer[index] + buffer[index] = row + + rng.shuffle(buffer) + yield from buffer + + +def token_blocks( + rows: Iterable[dict], tokenizer, text_field: str, block_size: int, *, skip_tokens: int = 0 +) -> Iterator[torch.Tensor]: + """Pack documents, optionally advancing a deterministic stream before packing. + + Skipping is before tensor allocation and works across document boundaries and + changes in batch/context length. Reconstructing a cursor still needs reading + and tokenizing the prefix; it does not run the teacher or student on it. + """ + if block_size < 2 or skip_tokens < 0: + raise ValueError("block_size must be >= 2 and skip_tokens must be nonnegative") + buffer: list[int] = [] + offset = 0 + eos = tokenizer.eos_token_id + for example in rows: + ids = tokenizer(example[text_field], add_special_tokens=False)["input_ids"] + if eos is not None: + ids.append(eos) + if skip_tokens: + skipped = min(skip_tokens, len(ids)) + skip_tokens -= skipped + ids = ids[skipped:] + buffer.extend(ids) + while len(buffer) - offset >= block_size: + yield torch.tensor(buffer[offset : offset + block_size], dtype=torch.long) + offset += block_size + if offset > 1_000_000: + buffer = buffer[offset:] + offset = 0 + + +def batch_blocks(blocks: Iterator[torch.Tensor], batch_size: int) -> Iterator[torch.Tensor]: + batch: list[torch.Tensor] = [] + for block in blocks: + batch.append(block) + if len(batch) == batch_size: + yield torch.stack(batch) + batch.clear() + + +def collect_eval_batches(rows, tokenizer, text_field: str, block_size: int, batch_size: int, count: int): + batches = batch_blocks(token_blocks(rows, tokenizer, text_field, block_size), batch_size) + out: list[torch.Tensor] = [] + for _ in range(count): + try: + out.append(next(batches)) + except StopIteration: + break + if not out: + raise RuntimeError("could not build held-out evaluation batches") + return out + + +def evaluation_fingerprint(batches: list[torch.Tensor]) -> str: + """Stable SHA256 fingerprint of the exact held-out token batches. + + Tokens are encoded explicitly as little-endian signed int64 values, avoiding + NumPy and platform-dependent tensor byte representations. + """ + digest = hashlib.sha256() + for batch in batches: + tensor = batch.detach().to(device="cpu", dtype=torch.int64).contiguous().view(-1) + for token_id in tensor.tolist(): + digest.update(struct.pack(" torch.Tensor: + if temperature <= 0 or chunk_rows < 1: + raise ValueError("temperature and chunk_rows must be positive") + s = student_logits.reshape(-1, student_logits.shape[-1]) + t = teacher_logits.reshape(-1, teacher_logits.shape[-1]) + total = s.new_zeros((), dtype=torch.float32) + rows = s.shape[0] + for start in range(0, rows, chunk_rows): + end = min(start + chunk_rows, rows) + s_chunk = s[start:end].float() / temperature + t_chunk = t[start:end].float() / temperature + total = total + F.kl_div( + F.log_softmax(s_chunk, dim=-1), + F.softmax(t_chunk, dim=-1), + reduction="sum", + ) + return total * (temperature * temperature) / max(rows, 1) + + +def hidden_cosine_loss(student_hidden: torch.Tensor, teacher_hidden: torch.Tensor) -> torch.Tensor: + return (1.0 - F.cosine_similarity(student_hidden.float(), teacher_hidden.float(), dim=-1)).mean() + + +def combined_loss( + student_out, + teacher_out, + *, + temperature: float, + kl_chunk_rows: int, + ce_weight: float, + kl_weight: float, + hidden_weight: float, +) -> tuple[torch.Tensor, dict[str, float]]: + ce = student_out.loss.float() + kl = chunked_kl(student_out.logits, teacher_out.logits, temperature, kl_chunk_rows) + hidden = hidden_cosine_loss(student_out.hidden_states[-1], teacher_out.hidden_states[-1]) + total = ce_weight * ce + kl_weight * kl + hidden_weight * hidden + return total, { + "ce": float(ce.detach()), + "kl": float(kl.detach()), + "hidden": float(hidden.detach()), + "total": float(total.detach()), + } + + +@torch.inference_mode() +def evaluate_distillation( + teacher, + student, + batches: list[torch.Tensor], + *, + device: torch.device, + dtype: torch.dtype, + temperature: float, + kl_chunk_rows: int, + ce_weight: float, + kl_weight: float, + hidden_weight: float, +) -> dict[str, float | int]: + teacher.eval() + student.eval() + sums = { + "student_ce": 0.0, + "teacher_ce": 0.0, + "kl": 0.0, + "hidden": 0.0, + "total": 0.0, + } + n_batches = 0 + eval_tokens = 0 + amp = (lambda: torch.autocast("cuda", dtype=dtype)) if device.type == "cuda" else nullcontext + for cpu_ids in batches: + ids = cpu_ids.to(device, non_blocking=True) + with amp(): + teacher_out = teacher( + input_ids=ids, + labels=ids, + output_hidden_states=True, + use_cache=False, + ) + student_out = student( + input_ids=ids, + labels=ids, + output_hidden_states=True, + use_cache=False, + ) + total, parts = combined_loss( + student_out, + teacher_out, + temperature=temperature, + kl_chunk_rows=kl_chunk_rows, + ce_weight=ce_weight, + kl_weight=kl_weight, + hidden_weight=hidden_weight, + ) + sums["student_ce"] += float(student_out.loss.detach().float()) + sums["teacher_ce"] += float(teacher_out.loss.detach().float()) + sums["kl"] += parts["kl"] + sums["hidden"] += parts["hidden"] + sums["total"] += float(total.detach().float()) + n_batches += 1 + eval_tokens += ids.numel() + + for key in sums: + sums[key] /= max(n_batches, 1) + result: dict[str, float | int] = dict(sums) + result["student_ppl"] = math.exp(min(sums["student_ce"], 30.0)) + result["teacher_ppl"] = math.exp(min(sums["teacher_ce"], 30.0)) + result["eval_batches"] = n_batches + result["eval_tokens"] = eval_tokens + return result diff --git a/src/tinycenn_lm/gemma3_integrated_memory.py b/src/tinycenn_lm/gemma3_integrated_memory.py new file mode 100644 index 0000000000000000000000000000000000000000..e2e9c4ceda99908bc1ba95c78215df910f69e735 --- /dev/null +++ b/src/tinycenn_lm/gemma3_integrated_memory.py @@ -0,0 +1,319 @@ +"""Gemma 3 / FunctionGemma integrated TinyCeNN memory. + +This module mirrors the SmolLM2 V3 experiment but respects Gemma 3's +Q/K RMS normalization and hybrid sliding/full-attention cache. It is intended +for replacing *full-attention* Gemma 3 text layers; the original sliding-window +layers remain untouched. +""" +import copy +from contextlib import contextmanager + +import torch +import torch.nn.functional as F +from torch import nn +from transformers.cache_utils import DynamicCache +from transformers.models.gemma3.modeling_gemma3 import apply_rotary_pos_emb, repeat_kv + +from .optimized_memory import OptimizedMemory + + +FORMAT = "functiongemma-cenn-integrated-v1" + + +def native_dtype(device): + """Use BF16 only when the GPU has native BF16 arithmetic.""" + if torch.device(device).type != "cuda": + return torch.float32 + return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16 + + +class Gemma3IntegratedCache(DynamicCache): + """Gemma 3 hybrid cache plus bounded TinyCeNN states at replaced layers.""" + + def __init__(self, config, memory_layers=()): + super().__init__(config=config) + self.memory_layers = frozenset(memory_layers) + self.memory_states = {} + + def get_seq_length(self, layer_idx=0): + if layer_idx in self.memory_layers: + state = self.memory_states.get(layer_idx) + return state.position if state is not None else 0 + return super().get_seq_length(layer_idx) + + def get_mask_sizes(self, cache_position, layer_idx): + if layer_idx in self.memory_layers: + query_length = cache_position.shape[0] if isinstance(cache_position, torch.Tensor) else int(cache_position) + return self.get_seq_length(layer_idx) + query_length, 0 + return super().get_mask_sizes(cache_position, layer_idx) + + @property + def nbytes(self): + tensors = [ + tensor + for layer in self.layers + for tensor in (getattr(layer, "keys", None), getattr(layer, "values", None)) + if isinstance(tensor, torch.Tensor) + ] + return sum(x.numel() * x.element_size() for x in tensors) + sum( + state.nbytes for state in self.memory_states.values() + ) + + def reorder_cache(self, *args, **kwargs): + raise NotImplementedError("Use batch-one greedy decoding; beam-search cache reordering is unsupported") + + def crop(self, *args, **kwargs): + raise NotImplementedError("Compressed TinyCeNN history cannot be cropped; start a fresh cache") + + def batch_repeat_interleave(self, *args, **kwargs): + raise NotImplementedError("Batch expansion is unsupported for TinyCeNN memory states") + + def batch_select_indices(self, *args, **kwargs): + raise NotImplementedError("Batch selection is unsupported for TinyCeNN memory states") + + +class Gemma3IntegratedAttention(nn.Module): + """Full-attention Gemma 3 layer replaced by a TinyCeNN memory core. + + Gemma3DecoderLayer inspects attributes on ``self_attn`` before calling it + (most importantly ``is_sliding`` to select local vs global RoPE). Mirror + the lightweight structural attributes of Gemma3Attention so replacing the + module preserves the Transformers 4.57.x decoder contract. + """ + + def __init__(self, original, core, layer_idx): + super().__init__() + self.original = original + self.core = core + self.layer_idx = layer_idx + + # Structural Gemma3Attention API used by Gemma3DecoderLayer and by + # attention tooling. These are plain metadata values, not duplicate + # module registrations; Q/K/V/O and RMSNorm modules remain under + # ``self.original`` only. + self.is_sliding = bool(original.is_sliding) + self.config = original.config + self.head_dim = original.head_dim + self.num_key_value_groups = original.num_key_value_groups + self.scaling = original.scaling + self.attention_dropout = original.attention_dropout + self.is_causal = original.is_causal + self.attn_logit_softcapping = original.attn_logit_softcapping + self.sliding_window = original.sliding_window + + if self.is_sliding: + raise ValueError("Gemma3IntegratedAttention only supports full-attention layers") + + self.register_buffer("fused_weight", None, persistent=False) + + def fuse(self, enabled=True): + if not enabled: + self.fused_weight = None + return + with torch.no_grad(): + weight = self.original.o_proj.weight.float().reshape( + -1, self.core.num_heads, self.core.head_dim + ) + self.fused_weight = torch.einsum( + "ohd,hkd->ohk", weight, self.core.readout.float() + ).reshape_as(self.original.o_proj.weight).to(self.original.o_proj.weight.dtype) + + def forward( + self, + hidden_states, + position_embeddings=None, + attention_mask=None, + past_key_values=None, + past_key_value=None, + cache_position=None, + **kwargs, + ): + if position_embeddings is None: + raise ValueError("Gemma 3 position_embeddings are required") + cache = past_key_values if past_key_values is not None else past_key_value + if cache is not None and not isinstance(cache, Gemma3IntegratedCache): + raise TypeError("Use Gemma3IntegratedCache with a FunctionGemma TinyCeNN model") + if cache is not None and torch.is_grad_enabled(): + raise RuntimeError("Train with use_cache=False") + if self.fused_weight is not None and torch.is_grad_enabled(): + raise RuntimeError("Unfuse the readout before training") + + if self.is_sliding: + raise RuntimeError("TinyCeNN FunctionGemma V1 only supports replacing full-attention layers") + + if attention_mask is not None: + if attention_mask.ndim != 4 or bool((attention_mask[..., -1, :] < 0).any()): + raise ValueError("Only unpadded causal batches are supported") + + b, t, _ = hidden_states.shape + h = self.core.num_heads + hk = self.core.num_kv_heads + d = self.core.head_dim + + q = self.original.q_proj(hidden_states).view(b, t, h, d).transpose(1, 2) + k = self.original.k_proj(hidden_states).view(b, t, hk, d).transpose(1, 2) + v = self.original.v_proj(hidden_states).view(b, t, hk, d).transpose(1, 2) + + q = self.original.q_norm(q) + k = self.original.k_norm(k) + cos, sin = position_embeddings + q, k = apply_rotary_pos_emb(q, k, cos, sin) + + if self.core.variant == "transformer_readout": + if cache is not None: + cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + k, v = cache.update(k, v, self.layer_idx, cache_kwargs) + output = F.scaled_dot_product_attention( + q, + repeat_kv(k, h // hk), + repeat_kv(v, h // hk), + attn_mask=attention_mask, + is_causal=attention_mask is None and t > 1, + scale=float(self.scaling), + ) + if self.fused_weight is None: + output = self.core.calibrate(output.float()) + elif cache is not None: + output, state = self.core( + q, + k, + v, + state=cache.memory_states.get(self.layer_idx), + return_state=True, + apply_readout=self.fused_weight is None, + ) + cache.memory_states[self.layer_idx] = state + else: + output = self.core(q, k, v, apply_readout=self.fused_weight is None) + + flat = output.transpose(1, 2).reshape(b, t, h * d).to(hidden_states.dtype) + if self.fused_weight is None: + return self.original.o_proj(flat), None + return F.linear(flat, self.fused_weight, self.original.o_proj.bias), None + + +def wrappers(model): + return [ + layer.self_attn + for layer in model.model.layers + if isinstance(layer.self_attn, Gemma3IntegratedAttention) + ] + + +def full_attention_layers(model): + return [ + i for i, layer in enumerate(model.model.layers) + if not getattr(layer.self_attn, "is_sliding", False) + ] + + +def build_student(teacher, layers, variant="cenn_partition", features=64, block_size=32, sinks=4): + model = copy.deepcopy(teacher).eval().requires_grad_(False) + config = model.config.get_text_config(decoder=True) + if config.model_type != "gemma3_text": + raise ValueError(f"Expected gemma3_text, got {config.model_type}") + + for index in layers: + if not 0 <= index < len(model.model.layers): + raise ValueError(f"Invalid layer {index}") + original = model.model.layers[index].self_attn + if getattr(original, "is_sliding", False): + raise ValueError( + f"Layer {index} is sliding_attention. FunctionGemma V1 intentionally replaces only full_attention layers." + ) + # FunctionGemma uses query_pre_attn_scalar == head_dim (256), so its + # native attention scale matches the OptimizedMemory softmax scale. + # Keep this guard explicit for future Gemma3 checkpoints where that + # assumption may not hold. + expected_scale = original.head_dim ** -0.5 + if abs(float(original.scaling) - float(expected_scale)) > 1e-8: + raise ValueError( + f"Layer {index} uses attention scale {original.scaling}, but TinyCeNN currently expects {expected_scale}." + ) + core = OptimizedMemory( + config.num_attention_heads, + config.num_key_value_heads, + original.head_dim, + features, + variant, + block_size, + sinks, + ).to(original.q_proj.weight.device) + model.model.layers[index].self_attn = Gemma3IntegratedAttention(original, core, index) + return model + + +@contextmanager +def inference_mode(model, compute_dtype="float32"): + adapters = wrappers(model) + previous = [a.core.compute_dtype for a in adapters] + try: + for adapter in adapters: + adapter.core.compute_dtype = compute_dtype + adapter.fuse() + with torch.no_grad(): + yield model + finally: + for adapter, dtype in zip(adapters, previous): + adapter.fuse(False) + adapter.core.compute_dtype = dtype + + +def new_cache(model): + memory_layers = [ + a.layer_idx for a in wrappers(model) + if a.core.variant != "transformer_readout" + ] + return Gemma3IntegratedCache(model.config, memory_layers) + + +def adapter_payload(model, metadata=None): + return { + "format": FORMAT, + "metadata": metadata or {}, + "adapters": { + str(a.layer_idx): { + "config": a.core.config, + "state_dict": { + k: v.detach().cpu().clone() for k, v in a.core.state_dict().items() + }, + } + for a in wrappers(model) + }, + } + + +def restore_student(teacher, payload): + if payload["format"] != FORMAT: + raise ValueError(f"Not a {FORMAT} checkpoint") + model = copy.deepcopy(teacher).eval().requires_grad_(False) + for key, value in payload["adapters"].items(): + index = int(key) + original = model.model.layers[index].self_attn + if getattr(original, "is_sliding", False): + raise ValueError(f"Checkpoint attempts to replace sliding-attention layer {index}") + core = OptimizedMemory(**value["config"]).to(original.q_proj.weight.device) + core.load_state_dict(value["state_dict"]) + model.model.layers[index].self_attn = Gemma3IntegratedAttention(original, core, index) + return model + + +@torch.no_grad() +def greedy_generate(model, ids, tokens=32, stop_token_ids=()): + """Batch-one cached greedy generation for the hybrid FunctionGemma model.""" + if ids.shape[0] != 1 or ids.shape[1] < 1 or tokens < 1: + raise ValueError("Use batch size one, a nonempty prompt, and positive tokens") + stop_token_ids = set(int(x) for x in stop_token_ids) + cache = new_cache(model) + output = model(input_ids=ids, past_key_values=cache, use_cache=True).logits[:, -1] + continuation = [] + token = output.argmax(-1, keepdim=True) + continuation.append(token) + for _ in range(tokens - 1): + if int(continuation[-1].item()) in stop_token_ids: + break + output = model( + input_ids=continuation[-1], past_key_values=cache, use_cache=True + ).logits[:, -1] + continuation.append(output.argmax(-1, keepdim=True)) + return torch.cat(continuation, dim=1), cache diff --git a/src/tinycenn_lm/gemma3_memory_fusion.py b/src/tinycenn_lm/gemma3_memory_fusion.py new file mode 100644 index 0000000000000000000000000000000000000000..1d1c08f8988abde0a0ba3da6eff7baaf1b4d4cb0 --- /dev/null +++ b/src/tinycenn_lm/gemma3_memory_fusion.py @@ -0,0 +1,255 @@ +from __future__ import annotations + +import copy +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable + +import torch +from torch import Tensor, nn +from transformers.models.gemma3.modeling_gemma3 import apply_rotary_pos_emb + +from .memory_attention import MemoryAugmentedCellularLayer + +DEFAULT_FUNCTIONGEMMA = "vtava/functiongemma-270m-it-simple-tool-calling" +FORMAT = "functiongemma-memory-fusion-sequential-v1" + + +@dataclass(frozen=True) +class Gemma3MemoryFusionConfig: + feature_dim: int = 32 + memory_rank: int = 64 + dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64, 128) + shifted_window: int = 8 + train_output_projection: bool = True + + def validate(self, model_config) -> None: + config = model_config.get_text_config(decoder=True) if hasattr(model_config, "get_text_config") else model_config + if getattr(config, "model_type", None) != "gemma3_text": + raise ValueError(f"expected gemma3_text, got {getattr(config, 'model_type', None)!r}") + if self.feature_dim < 4: + raise ValueError("feature_dim must be >= 4") + if self.memory_rank < 4: + raise ValueError("memory_rank must be >= 4") + if not self.dilations or min(self.dilations) < 1: + raise ValueError("dilations must be positive") + if int(config.num_attention_heads) % int(config.num_key_value_heads): + raise ValueError("num_attention_heads must be divisible by num_key_value_heads") + + def to_dict(self) -> dict: + value = asdict(self) + value["dilations"] = list(self.dilations) + return value + + @classmethod + def from_dict(cls, data: dict) -> "Gemma3MemoryFusionConfig": + value = dict(data) + value["dilations"] = tuple(value.get("dilations", (1, 2, 4, 8, 16, 32, 64, 128))) + return cls(**value) + + +class MemoryFusionGemma3Attention(nn.Module): + """Gemma3 full-attention replacement using the TinyCeNN Memory Fusion core. + + The pretrained Q/K/V/O projections and Gemma3 Q/K RMS normalizers are copied + exactly. Only original *full-attention* layers are supported in V1; the model's + sliding-window layers remain untouched. Training and prompt checks use full + prefixes with ``use_cache=False`` so the experiment is intentionally simple + and auditable before adding a hybrid recurrent cache. + """ + + def __init__(self, original_attn: nn.Module, model_config, config: Gemma3MemoryFusionConfig, layer_idx: int): + super().__init__() + config.validate(model_config) + text_config = model_config.get_text_config(decoder=True) if hasattr(model_config, "get_text_config") else model_config + if bool(getattr(original_attn, "is_sliding", False)): + raise ValueError("Memory Fusion V1 replaces only Gemma3 full-attention layers") + + self.layer_idx = int(layer_idx) + self.config = getattr(original_attn, "config", text_config) + self.is_sliding = False + self.hidden_size = int(text_config.hidden_size) + self.num_heads = int(text_config.num_attention_heads) + self.num_key_value_heads = int(text_config.num_key_value_heads) + self.head_dim = int(getattr(original_attn, "head_dim", text_config.head_dim)) + self.attention_width = self.num_heads * self.head_dim + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.scaling = float(getattr(original_attn, "scaling", self.head_dim ** -0.5)) + self.attention_dropout = float(getattr(original_attn, "attention_dropout", 0.0)) + self.is_causal = bool(getattr(original_attn, "is_causal", True)) + self.attn_logit_softcapping = getattr(original_attn, "attn_logit_softcapping", None) + self.sliding_window = getattr(original_attn, "sliding_window", None) + + self.q_proj = copy.deepcopy(original_attn.q_proj) + self.k_proj = copy.deepcopy(original_attn.k_proj) + self.v_proj = copy.deepcopy(original_attn.v_proj) + self.o_proj = copy.deepcopy(original_attn.o_proj) + self.q_norm = copy.deepcopy(original_attn.q_norm) + self.k_norm = copy.deepcopy(original_attn.k_norm) + + self.core = MemoryAugmentedCellularLayer( + num_heads=self.num_heads, + num_kv_heads=self.num_key_value_heads, + head_dim=self.head_dim, + feature_dim=config.feature_dim, + variant="cellular_memory_fusion", + dilations=config.dilations, + shifted_window=config.shifted_window, + memory_rank=config.memory_rank, + ) + self.last_core_output: Tensor | None = None + + def forward( + self, + hidden_states: Tensor, + position_embeddings=None, + attention_mask=None, + position_ids=None, + past_key_values=None, + past_key_value=None, + use_cache: bool = False, + cache_position=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache or past_key_values is not None or past_key_value is not None: + raise RuntimeError("FunctionGemma Memory Fusion V1 currently requires use_cache=False") + if position_embeddings is None: + raise ValueError("Gemma3 position_embeddings are required") + + bsz, seq_len, _ = hidden_states.shape + if attention_mask is not None: + if attention_mask.ndim != 4 or attention_mask.shape[-1] != seq_len: + raise ValueError("only unpadded full causal blocks are supported") + if bool((attention_mask[..., -1, :] < -1e4).any()): + raise ValueError("padded batches are not supported") + + q = self.q_proj(hidden_states).view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + k = self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + q = self.q_norm(q) + k = self.k_norm(k) + cos, sin = position_embeddings + q, k = apply_rotary_pos_emb(q, k, cos, sin) + + core_out = self.core(q.float(), k.float(), v.float()) + self.last_core_output = core_out + # Gemma3 can use num_heads * head_dim != hidden_size (FunctionGemma does). + # The original o_proj maps the attention width back to hidden_size. + flat = core_out.transpose(1, 2).reshape(bsz, seq_len, self.attention_width) + return self.o_proj(flat.to(hidden_states.dtype)), None + + +def full_attention_layers(model: nn.Module) -> list[int]: + return [ + i for i, layer in enumerate(model.model.layers) + if not bool(getattr(layer.self_attn, "is_sliding", False)) + ] + + +def replace_attention_layers(model: nn.Module, config: Gemma3MemoryFusionConfig, layer_indices: Iterable[int]) -> nn.Module: + config.validate(model.config) + for raw_idx in layer_indices: + idx = int(raw_idx) + layer = model.model.layers[idx] + if isinstance(layer.self_attn, MemoryFusionGemma3Attention): + continue + old = layer.self_attn + if bool(getattr(old, "is_sliding", False)): + raise ValueError(f"layer {idx} is sliding_attention; V1 replaces only full_attention") + device = old.q_proj.weight.device + projection_dtype = old.q_proj.weight.dtype + new = MemoryFusionGemma3Attention(old, model.config, config, idx) + for module in (new.q_proj, new.k_proj, new.v_proj, new.o_proj, new.q_norm, new.k_norm): + module.to(device=device, dtype=projection_dtype) + new.core.to(device=device, dtype=torch.float32) + layer.self_attn = new + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def freeze_current_layer_only(model: nn.Module, layer_idx: int, *, train_output_projection: bool = True) -> list[nn.Parameter]: + for p in model.parameters(): + p.requires_grad = False + module = model.model.layers[int(layer_idx)].self_attn + if not isinstance(module, MemoryFusionGemma3Attention): + raise TypeError(f"layer {layer_idx} is not MemoryFusionGemma3Attention") + trainable: list[nn.Parameter] = [] + for p in module.core.parameters(): + p.requires_grad = True + trainable.append(p) + if train_output_projection: + for p in module.o_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def freeze_all_memory_fusion(model: nn.Module, *, train_output_projection: bool = True) -> list[nn.Parameter]: + for p in model.parameters(): + p.requires_grad = False + trainable: list[nn.Parameter] = [] + for layer in model.model.layers: + module = layer.self_attn + if not isinstance(module, MemoryFusionGemma3Attention): + continue + for p in module.core.parameters(): + p.requires_grad = True + trainable.append(p) + if train_output_projection: + for p in module.o_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def structural_summary(model: nn.Module) -> dict[str, object]: + fusion = [ + i for i, layer in enumerate(model.model.layers) + if isinstance(layer.self_attn, MemoryFusionGemma3Attention) + ] + full = [ + i for i, layer in enumerate(model.model.layers) + if not isinstance(layer.self_attn, MemoryFusionGemma3Attention) + and not bool(getattr(layer.self_attn, "is_sliding", False)) + ] + sliding = [ + i for i, layer in enumerate(model.model.layers) + if not isinstance(layer.self_attn, MemoryFusionGemma3Attention) + and bool(getattr(layer.self_attn, "is_sliding", False)) + ] + return { + "memory_fusion_layers": fusion, + "remaining_full_attention_layers": full, + "sliding_attention_layers": sliding, + } + + +def selected_attention_state(model: nn.Module, layers: Iterable[int]) -> dict[str, Tensor]: + prefixes = tuple(f"model.layers.{int(i)}.self_attn." for i in layers) + return { + key: value.detach().cpu() + for key, value in model.state_dict().items() + if prefixes and key.startswith(prefixes) + } + + +def save_adapter(model: nn.Module, output_dir: str | Path, *, config: Gemma3MemoryFusionConfig, base_model: str, accepted_layers: list[int], metadata: dict | None = None) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(selected_attention_state(model, accepted_layers), output_dir / "functiongemma_memory_fusion.pt") + payload = { + "format": FORMAT, + "base_model": base_model, + "accepted_layers": list(accepted_layers), + "memory_fusion": config.to_dict(), + "metadata": metadata or {}, + } + (output_dir / "functiongemma_memory_fusion_config.json").write_text(json.dumps(payload, indent=2), encoding="utf-8") + return output_dir diff --git a/src/tinycenn_lm/hf_persistence.py b/src/tinycenn_lm/hf_persistence.py new file mode 100644 index 0000000000000000000000000000000000000000..cb4f2f45be47fb6b41e17698029fb8882db699e5 --- /dev/null +++ b/src/tinycenn_lm/hf_persistence.py @@ -0,0 +1,411 @@ +from __future__ import annotations + +import hashlib +import json +import os +import platform +import re +import shutil +import sys +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +_TOKEN_RE = re.compile(r"hf_[A-Za-z0-9]{20,}") +_SMALL_ARTIFACT_SUFFIXES = {".json", ".jsonl", ".csv", ".txt", ".md", ".log", ".yaml", ".yml"} + + +def utc_run_id(prefix: str = "run") -> str: + stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") + return f"{prefix}-{stamp}" + + +def redact_secrets(text: str) -> str: + return _TOKEN_RE.sub("hf_REDACTED", text) + + +def _json_files(folder: Path) -> list[Path]: + return sorted( + p for p in folder.rglob("*.json") + if p.is_file() and p.stat().st_size <= 5 * 1024 * 1024 + ) + + +def _load_json(path: Path) -> Any | None: + try: + return json.loads(path.read_text(encoding="utf-8")) + except Exception: + return None + + +def collect_reports(folder: str | Path) -> dict[str, Any]: + folder = Path(folder) + reports: dict[str, Any] = {} + for path in _json_files(folder): + name = path.name.lower() + if any(key in name for key in ("report", "result", "metric", "summary", "config")): + value = _load_json(path) + if value is not None: + reports[str(path.relative_to(folder))] = value + return reports + + +def _first_value(reports: dict[str, Any], keys: tuple[str, ...]) -> Any | None: + def walk(value: Any) -> Any | None: + if isinstance(value, dict): + for key in keys: + if key in value and value[key] not in (None, ""): + return value[key] + for child in value.values(): + found = walk(child) + if found is not None: + return found + elif isinstance(value, list): + for child in value: + found = walk(child) + if found is not None: + return found + return None + + return walk(reports) + + +def _flatten_scalars(value: Any, prefix: str = "") -> dict[str, Any]: + out: dict[str, Any] = {} + if isinstance(value, dict): + for key, child in value.items(): + name = f"{prefix}.{key}" if prefix else str(key) + out.update(_flatten_scalars(child, name)) + elif isinstance(value, (str, int, float, bool)) or value is None: + out[prefix] = value + return out + + +def _interesting_metrics(reports: dict[str, Any]) -> list[tuple[str, Any]]: + wanted = ( + "status", "stop_reason", "seen_tokens", "updates", "context_length", + "feature_dim", "num_shards", "top_k", "trainable", "trainable_percent", + "last_training_ce", "last_distillation_kl", "best.student_ce", "best.teacher_ce", + "teacher_gap_recovery_fraction", "mean_route_mix", "mean_router_entropy", + "elapsed_minutes", "peak_vram_gib", "evaluation_performed", + ) + flat: dict[str, Any] = {} + for report in reports.values(): + if isinstance(report, dict): + flat.update(_flatten_scalars(report)) + rows: list[tuple[str, Any]] = [] + seen: set[str] = set() + for want in wanted: + for key, value in flat.items(): + if key == want or key.endswith("." + want): + label = want.split(".")[-1] + if label not in seen: + rows.append((label, value)) + seen.add(label) + break + return rows + + +def _format_value(value: Any) -> str: + if isinstance(value, float): + if abs(value) >= 1000: + return f"{value:,.2f}" + if abs(value) < 0.001 and value != 0: + return f"{value:.3e}" + return f"{value:.6g}" + if isinstance(value, int): + return f"{value:,}" + return str(value) + + +def build_model_card( + folder: str | Path, + *, + title: str | None = None, + architecture: str | None = None, + base_model: str | None = None, + source_repo: str = "https://github.com/vtavakkoli/TinyCeNN-LM", + extra_notes: str | None = None, +) -> str: + folder = Path(folder) + reports = collect_reports(folder) + inferred_arch = architecture or _first_value(reports, ("architecture", "task")) or "TinyCeNN-LM experiment" + inferred_base = base_model or _first_value(reports, ("base_model", "teacher_model")) + dataset = _first_value(reports, ("dataset",)) + dataset_config = _first_value(reports, ("dataset_config",)) + display_title = title or folder.name.replace("-", " ") + + tags = ["tinycenn", "cenn", "language-modeling", "text-generation", "research"] + arch_lower = str(inferred_arch).lower() + if "distill" in arch_lower: + tags.append("knowledge-distillation") + if "moe" in arch_lower or _first_value(reports, ("num_shards", "num_experts")): + tags.append("mixture-of-experts") + if "amcenn" in arch_lower or "attention-free" in arch_lower: + tags.extend(["attention-free", "linear-attention", "recurrent-memory"]) + if "story" in arch_lower: + tags.append("story-generation") + + yaml_lines = ["---", "library_name: transformers", "pipeline_tag: text-generation"] + if inferred_base: + yaml_lines.append(f"base_model: {inferred_base}") + if dataset: + yaml_lines.append("datasets:") + yaml_lines.append(f"- {dataset}") + yaml_lines.append("tags:") + for tag in dict.fromkeys(tags): + yaml_lines.append(f"- {tag}") + yaml_lines.append("---") + + metrics = _interesting_metrics(reports) + if metrics: + lines = ["| Metric | Value |", "|---|---:|"] + lines += [f"| `{name}` | {_format_value(value)} |" for name, value in metrics] + metric_table = "\n".join(lines) + else: + metric_table = "No structured training report was found in this upload." + + report_files = sorted(reports) + files_text = "\n".join(f"- `{name}`" for name in report_files) or "- No JSON report files detected." + dataset_text = str(dataset) if dataset else "Not recorded" + if dataset_config: + dataset_text += f" (`{dataset_config}`)" + + limitations = ( + "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." + ) + notes = f"\n## Notes\n\n{extra_notes.strip()}\n" if extra_notes else "" + return redact_secrets("\n".join(yaml_lines) + f"\n\n# {display_title}\n\n" + f"Research artifact from **TinyCeNN-LM**. Architecture: `{inferred_arch}`.\n\n" + "## Architecture\n\n" + f"- Architecture/run type: `{inferred_arch}`\n" + f"- Base model: `{inferred_base or 'not recorded'}`\n" + f"- Dataset: `{dataset_text}`\n" + f"- Source code: {source_repo}\n\n" + "## Latest saved results\n\n" + f"{metric_table}\n\n" + "The Hugging Face repository keeps timestamped run artifacts under `runs/`. This preserves training " + "reports, configs and run metadata independently of the temporary Colab filesystem.\n\n" + "## Saved experiment files\n\n" + f"{files_text}\n\n" + "## Reproducibility\n\n" + "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.\n\n" + "## Limitations\n\n" + f"{limitations}\n" + f"{notes}\n" + "## Citation\n\n" + "If you use this experimental checkpoint, cite the TinyCeNN-LM repository and the upstream base model.\n" + ) + + +def write_run_manifest( + folder: str | Path, + *, + run_id: str, + notebook: str | None = None, + repo_id: str | None = None, +) -> Path: + folder = Path(folder) + reports = collect_reports(folder) + manifest = { + "run_id": run_id, + "created_utc": datetime.now(timezone.utc).isoformat(), + "repo_id": repo_id, + "notebook": notebook, + "python": sys.version.split()[0], + "platform": platform.platform(), + "reports": sorted(reports), + } + try: + import torch + manifest["torch"] = torch.__version__ + manifest["cuda_available"] = bool(torch.cuda.is_available()) + if torch.cuda.is_available(): + manifest["gpu"] = torch.cuda.get_device_name(0) + except Exception: + pass + path = folder / "run_manifest.json" + path.write_text(json.dumps(manifest, indent=2), encoding="utf-8") + return path + + +def _copy_small_artifacts(folder: Path, archive_dir: Path) -> list[str]: + copied: list[str] = [] + for path in folder.rglob("*"): + if not path.is_file(): + continue + try: + if path.is_relative_to(archive_dir.parent): + continue + except AttributeError: + pass + if path.suffix.lower() not in _SMALL_ARTIFACT_SUFFIXES: + continue + if path.stat().st_size > 10 * 1024 * 1024: + continue + rel = path.relative_to(folder) + dst = archive_dir / rel + dst.parent.mkdir(parents=True, exist_ok=True) + try: + text = path.read_text(encoding="utf-8") + dst.write_text(redact_secrets(text), encoding="utf-8") + except Exception: + shutil.copy2(path, dst) + copied.append(str(rel)) + return copied + + +def _prepare_run_archive(folder: Path, *, repo_id: str, run_id: str, notebook: str | None = None) -> Path: + archive_dir = folder / ".hf_run_archive" / run_id + if archive_dir.exists(): + shutil.rmtree(archive_dir) + archive_dir.mkdir(parents=True, exist_ok=True) + copied = _copy_small_artifacts(folder, archive_dir) + latest = { + "run_id": run_id, + "repo_id": repo_id, + "created_utc": datetime.now(timezone.utc).isoformat(), + "notebook": notebook, + "artifacts": copied, + } + (archive_dir / "latest_run.json").write_text(json.dumps(latest, indent=2), encoding="utf-8") + return archive_dir + + +def persist_hf_run( + *, + api, + repo_id: str, + folder_path: str | Path, + token: str | None = None, + title: str | None = None, + architecture: str | None = None, + base_model: str | None = None, + notebook: str | None = None, + run_id: str | None = None, + commit_message: str | None = None, + upload_model_files: bool = True, +) -> dict[str, Any]: + """Persist a completed Colab run and a timestamped result archive to Hugging Face.""" + from huggingface_hub import HfApi + + folder = Path(folder_path) + if not folder.exists(): + raise FileNotFoundError(folder) + run_id = run_id or utc_run_id(folder.name[:24] or "run") + client = api if api is not None else HfApi(token=token) + client.create_repo(repo_id=repo_id, repo_type="model", exist_ok=True) + + card = build_model_card(folder, title=title, architecture=architecture, base_model=base_model) + (folder / "README.md").write_text(card, encoding="utf-8") + write_run_manifest(folder, run_id=run_id, notebook=notebook, repo_id=repo_id) + + if upload_model_files: + client.upload_folder( + repo_id=repo_id, repo_type="model", folder_path=str(folder), + commit_message=commit_message or f"Publish TinyCeNN run {run_id}", + ignore_patterns=[".hf_run_archive/**"], + ) + + archive_dir = _prepare_run_archive(folder, repo_id=repo_id, run_id=run_id, notebook=notebook) + latest_file = archive_dir / "latest_run.json" + client.upload_folder( + repo_id=repo_id, repo_type="model", folder_path=str(archive_dir), + path_in_repo=f"runs/{run_id}", commit_message=f"Archive TinyCeNN results {run_id}", + ) + client.upload_file( + repo_id=repo_id, repo_type="model", path_or_fileobj=str(latest_file), + path_in_repo="runs/latest_run.json", + commit_message=f"Update latest TinyCeNN run pointer to {run_id}", + ) + return json.loads(latest_file.read_text(encoding="utf-8")) + + +def _looks_like_tinycenn_folder(folder: Path) -> bool: + if "tinycenn" in str(folder).lower() or "smollm2-amcenn" in str(folder).lower(): + return True + names = {p.name.lower() for p in folder.iterdir()} if folder.exists() and folder.is_dir() else set() + return any("cenn" in name or "story_v2" in name or "sharded" in name for name in names) + + +def install_colab_hf_upload_enhancer() -> bool: + """Enhance existing notebook HfApi.upload_folder calls without duplicating notebook code. + + In Colab, TinyCeNN notebooks already import tinycenn_lm before publishing. This wrapper regenerates a + report-backed README and archives small result files under runs// whenever a TinyCeNN checkpoint + folder is uploaded. Outside Colab it is a no-op. + """ + if not (os.environ.get("COLAB_RELEASE_TAG") or os.environ.get("COLAB_GPU") or Path("/content").exists()): + return False + try: + from huggingface_hub import HfApi + except Exception: + return False + if getattr(HfApi.upload_folder, "_tinycenn_enhanced", False): + return True + + original_upload_folder = HfApi.upload_folder + original_upload_file = HfApi.upload_file + + def enhanced_upload_folder(self, *args, **kwargs): + folder_value = kwargs.get("folder_path") + repo_id = kwargs.get("repo_id") + repo_type = kwargs.get("repo_type", "model") + if folder_value is None and len(args) >= 1: + folder_value = args[0] + if repo_id is None and len(args) >= 2: + repo_id = args[1] + folder = Path(folder_value) if folder_value else None + should_enhance = ( + repo_type == "model" and folder is not None and folder.exists() and folder.is_dir() + and repo_id and _looks_like_tinycenn_folder(folder) + and ".hf_run_archive" not in str(folder) + ) + if not should_enhance: + return original_upload_folder(self, *args, **kwargs) + + run_id = utc_run_id(folder.name[:24] or "run") + reports = collect_reports(folder) + architecture = _first_value(reports, ("architecture", "task")) + base_model = _first_value(reports, ("base_model", "teacher_model")) + (folder / "README.md").write_text( + build_model_card(folder, title=str(repo_id).split("/")[-1], architecture=architecture, base_model=base_model), + encoding="utf-8", + ) + write_run_manifest(folder, run_id=run_id, repo_id=str(repo_id)) + ignore = list(kwargs.get("ignore_patterns") or []) + if ".hf_run_archive/**" not in ignore: + ignore.append(".hf_run_archive/**") + kwargs["ignore_patterns"] = ignore + result = original_upload_folder(self, *args, **kwargs) + + try: + archive_dir = _prepare_run_archive(folder, repo_id=str(repo_id), run_id=run_id) + original_upload_folder( + self, repo_id=repo_id, repo_type="model", folder_path=str(archive_dir), + path_in_repo=f"runs/{run_id}", commit_message=f"Archive TinyCeNN results {run_id}", + ) + original_upload_file( + self, repo_id=repo_id, repo_type="model", + path_or_fileobj=str(archive_dir / "latest_run.json"), + path_in_repo="runs/latest_run.json", + commit_message=f"Update latest TinyCeNN run pointer to {run_id}", + ) + except Exception as exc: + print(f"TinyCeNN HF result archive warning: {exc}") + return result + + enhanced_upload_folder._tinycenn_enhanced = True + HfApi.upload_folder = enhanced_upload_folder + return True + + +def fingerprint_file(path: str | Path) -> str: + h = hashlib.sha256() + with Path(path).open("rb") as f: + for chunk in iter(lambda: f.read(1024 * 1024), b""): + h.update(chunk) + return h.hexdigest() diff --git a/src/tinycenn_lm/integrated_memory.py b/src/tinycenn_lm/integrated_memory.py new file mode 100644 index 0000000000000000000000000000000000000000..32a9549847a4143bf3cb5b27fcd81e62a0272277 --- /dev/null +++ b/src/tinycenn_lm/integrated_memory.py @@ -0,0 +1,186 @@ +"""Jointly trainable V2 mixers with full-model causal caching for Llama/SmolLM2. + +Unpadded, batch-one greedy decoding is the supported inference protocol. This +adapter deliberately does not implement beam search, cache cropping or offload. +""" +import copy +from contextlib import contextmanager + +import torch +import torch.nn.functional as F +from torch import nn +from transformers.cache_utils import DynamicCache +from transformers.models.llama.modeling_llama import apply_rotary_pos_emb, repeat_kv + +from .optimized_memory import OptimizedMemory + + +def native_dtype(device): + """Do not mistake emulated BF16 allocation support for native T4 arithmetic.""" + if torch.device(device).type != "cuda": + return torch.float32 + return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16 + + +class IntegratedCache(DynamicCache): + def __init__(self, memory_layers=()): + super().__init__() + self.memory_layers = frozenset(memory_layers) + self.memory_states = {} + + def get_seq_length(self, layer_idx=0): + if layer_idx in self.memory_layers: + state = self.memory_states.get(layer_idx) + return state.position if state is not None else 0 + return super().get_seq_length(layer_idx) + + def get_mask_sizes(self, cache_position, layer_idx): + if layer_idx in self.memory_layers: + # Transformers 4.x passes positions; 5.x passes the query length. + query_length = cache_position.shape[0] if isinstance(cache_position, torch.Tensor) else cache_position + return self.get_seq_length(layer_idx) + query_length, 0 + return super().get_mask_sizes(cache_position, layer_idx) + + @property + def nbytes(self): + tensors = [x for layer in self.layers for x in (layer.keys, layer.values) + if isinstance(x, torch.Tensor)] + return sum(x.numel() * x.element_size() for x in tensors) + sum( + state.nbytes for state in self.memory_states.values()) + + def reorder_cache(self, *args, **kwargs): + raise NotImplementedError("Use the supplied batch-one greedy decoder; beam search is unsupported") + + def crop(self, *args, **kwargs): + raise NotImplementedError("Compressed history cannot be cropped; start a fresh cache") + + +class IntegratedAttention(nn.Module): + def __init__(self, original, core, layer_idx): + super().__init__() + self.original, self.core, self.layer_idx = original, core, layer_idx + self.register_buffer("fused_weight", None, persistent=False) + + def fuse(self, enabled=True): + if not enabled: + self.fused_weight = None + return + with torch.no_grad(): + weight = self.original.o_proj.weight.float().reshape( + -1, self.core.num_heads, self.core.head_dim) + self.fused_weight = torch.einsum("ohd,hkd->ohk", weight, self.core.readout.float()).reshape_as( + self.original.o_proj.weight).to(self.original.o_proj.weight.dtype) + + def forward(self, hidden_states, position_embeddings=None, attention_mask=None, + past_key_values=None, past_key_value=None, **kwargs): + if position_embeddings is None: + raise ValueError("Llama rotary position embeddings are required") + cache = past_key_values if past_key_values is not None else past_key_value + if cache is not None and not isinstance(cache, IntegratedCache): + raise TypeError("Use IntegratedCache for this model") + if cache is not None and torch.is_grad_enabled(): + raise RuntimeError("Train with use_cache=False") + if self.fused_weight is not None and torch.is_grad_enabled(): + raise RuntimeError("Unfuse the readout before training") + if attention_mask is not None: + if attention_mask.ndim != 4 or bool((attention_mask[..., -1, :] < 0).any()): + raise ValueError("Only unpadded causal blocks are supported") + b, t, _ = hidden_states.shape + h, hk, d = self.core.num_heads, self.core.num_kv_heads, self.core.head_dim + q = self.original.q_proj(hidden_states).view(b, t, h, d).transpose(1, 2) + k = self.original.k_proj(hidden_states).view(b, t, hk, d).transpose(1, 2) + v = self.original.v_proj(hidden_states).view(b, t, hk, d).transpose(1, 2) + q, k = apply_rotary_pos_emb(q, k, *position_embeddings) + if self.core.variant == "transformer_readout": + if cache is not None: + k, v = cache.update(k, v, self.layer_idx) + output = F.scaled_dot_product_attention(q, repeat_kv(k, h // hk), repeat_kv(v, h // hk), + attn_mask=attention_mask, is_causal=attention_mask is None and t > 1) + if self.fused_weight is None: + output = self.core.calibrate(output.float()) + elif cache is not None: + output, state = self.core(q, k, v, state=cache.memory_states.get(self.layer_idx), + return_state=True, apply_readout=self.fused_weight is None) + cache.memory_states[self.layer_idx] = state + else: + output = self.core(q, k, v, apply_readout=self.fused_weight is None) + flat = output.transpose(1, 2).reshape(b, t, h * d).to(hidden_states.dtype) + if self.fused_weight is None: + return self.original.o_proj(flat), None + return F.linear(flat, self.fused_weight, self.original.o_proj.bias), None + + +def wrappers(model): + return [layer.self_attn for layer in model.model.layers + if isinstance(layer.self_attn, IntegratedAttention)] + + +def build_student(teacher, layers, variant="cenn_partition", features=64, block_size=32, sinks=4): + model = copy.deepcopy(teacher).eval().requires_grad_(False) + config = model.config + if config.model_type != "llama": + raise ValueError("Only Llama-family models are supported") + for index in layers: + if not 0 <= index < len(model.model.layers): + raise ValueError(f"Invalid layer {index}") + original = model.model.layers[index].self_attn + core = OptimizedMemory(config.num_attention_heads, config.num_key_value_heads, + config.hidden_size // config.num_attention_heads, features, + variant, block_size, sinks).to(original.q_proj.weight.device) + model.model.layers[index].self_attn = IntegratedAttention(original, core, index) + return model + + +@contextmanager +def inference_mode(model, compute_dtype="float32"): + adapters = wrappers(model) + previous = [a.core.compute_dtype for a in adapters] + try: + for a in adapters: + a.core.compute_dtype = compute_dtype + a.fuse() + with torch.no_grad(): + yield model + finally: + for a, dtype in zip(adapters, previous): + a.fuse(False) + a.core.compute_dtype = dtype + + +def new_cache(model): + return IntegratedCache(a.layer_idx for a in wrappers(model) + if a.core.variant != "transformer_readout") + + +def adapter_payload(model, metadata=None): + return {"format": "smollm2-integrated-memory-v3", "metadata": metadata or {}, "adapters": { + str(a.layer_idx): {"config": a.core.config, + "state_dict": {k: v.detach().cpu().clone() for k, v in a.core.state_dict().items()}} + for a in wrappers(model)}} + + +def restore_student(teacher, payload): + if payload["format"] != "smollm2-integrated-memory-v3": + raise ValueError("Not an integrated memory checkpoint") + model = copy.deepcopy(teacher).eval().requires_grad_(False) + for key, value in payload["adapters"].items(): + index = int(key) + original = model.model.layers[index].self_attn + core = OptimizedMemory(**value["config"]).to(original.q_proj.weight.device) + core.load_state_dict(value["state_dict"]) + model.model.layers[index].self_attn = IntegratedAttention(original, core, index) + return model + + +@torch.no_grad() +def greedy_generate(model, ids, tokens=32): + """Deterministic fixed-length generation; no early EOS for comparable timing.""" + if ids.shape[0] != 1 or ids.shape[1] < 1 or tokens < 1: + raise ValueError("Use batch size one, a nonempty prompt, and positive tokens") + cache = new_cache(model) + output = model(input_ids=ids, past_key_values=cache, use_cache=True).logits[:, -1] + continuation = [output.argmax(-1, keepdim=True)] + for _ in range(tokens - 1): + output = model(input_ids=continuation[-1], past_key_values=cache, use_cache=True).logits[:, -1] + continuation.append(output.argmax(-1, keepdim=True)) + return torch.cat(continuation, dim=1), cache diff --git a/src/tinycenn_lm/live_console.py b/src/tinycenn_lm/live_console.py new file mode 100644 index 0000000000000000000000000000000000000000..9d55e86982bba50f00ae8d9ead2e747e52d56145 --- /dev/null +++ b/src/tinycenn_lm/live_console.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +import atexit +import os +import sys +import time +from pathlib import Path + +_STATUS_INSTALLED = False +_PROCESS_FAILED = False +_PROCESS_STARTED = 0.0 +_PROCESS_NAME = "" + + +def _format_elapsed(seconds: float) -> str: + seconds = max(0, int(seconds)) + hours, remainder = divmod(seconds, 3600) + minutes, secs = divmod(remainder, 60) + if hours: + return f"{hours:d}h {minutes:02d}m {secs:02d}s" + return f"{minutes:d}m {secs:02d}s" + + +def _install_training_process_status() -> None: + global _STATUS_INSTALLED, _PROCESS_STARTED, _PROCESS_NAME + if _STATUS_INSTALLED: + return + + script = Path(sys.argv[0]).name + if not (script.lower().startswith("train_") and script.lower().endswith(".py")): + return + + _STATUS_INSTALLED = True + _PROCESS_STARTED = time.monotonic() + _PROCESS_NAME = script[:-3] + print(f"[TinyCeNN][PROCESS START] {_PROCESS_NAME}", flush=True) + + original_excepthook = sys.excepthook + + def status_excepthook(exc_type, exc_value, traceback): + global _PROCESS_FAILED + _PROCESS_FAILED = True + elapsed = _format_elapsed(time.monotonic() - _PROCESS_STARTED) + print( + f"[TinyCeNN][PROCESS FAILED] {_PROCESS_NAME} after {elapsed}: " + f"{exc_type.__name__}: {exc_value}", + flush=True, + ) + original_excepthook(exc_type, exc_value, traceback) + + sys.excepthook = status_excepthook + + def final_status() -> None: + elapsed = _format_elapsed(time.monotonic() - _PROCESS_STARTED) + if _PROCESS_FAILED: + print(f"[TinyCeNN][PROCESS END] {_PROCESS_NAME} failed after {elapsed}", flush=True) + else: + print(f"[TinyCeNN][PROCESS DONE] {_PROCESS_NAME} completed in {elapsed}", flush=True) + + atexit.register(final_status) + + +def configure_live_console() -> None: + """Prefer immediate stdout/stderr visibility for notebook-launched trainers. + + Existing Colab notebooks launch ``train_*.py`` with ``subprocess.run``. Setting + PYTHONUNBUFFERED before child interpreters start and making the current process + line-buffered keeps progress messages visible as they are produced instead of + appearing in a large block at the end of a run. + + Trainer processes also emit their own PROCESS START/DONE/FAILED markers. This + covers notebooks that have not imported ``tinycenn_lm`` in the parent kernel + before launching the trainer. + """ + os.environ.setdefault("PYTHONUNBUFFERED", "1") + + for stream in (sys.stdout, sys.stderr): + reconfigure = getattr(stream, "reconfigure", None) + if reconfigure is None: + continue + try: + reconfigure(line_buffering=True, write_through=True) + except (TypeError, ValueError, OSError): + # Some notebook stream wrappers do not expose every TextIO option. + try: + reconfigure(line_buffering=True) + except (TypeError, ValueError, OSError): + pass + + _install_training_process_status() diff --git a/src/tinycenn_lm/memory_attention.py b/src/tinycenn_lm/memory_attention.py new file mode 100644 index 0000000000000000000000000000000000000000..41a1a0ee4f1c756b1644e5a78f49b59f3a0c8b2e --- /dev/null +++ b/src/tinycenn_lm/memory_attention.py @@ -0,0 +1,473 @@ +"""Research-only global-memory augmentations for TinyCeNN-LM. + +The module keeps the proven adaptive+MaxPool Cellular Attention path as a local/ +multiscale branch and adds optional global causal memories inspired by recent +efficient sequence models: + +* Hedgehog: learned positive feature maps for softmax-mimicking linear attention. +* Kimi Delta Attention (KDA): fine-grained per-channel forgetting plus delta updates. +* Gated DeltaNet-2: KDA-like decay with decoupled erase and write gates. +* xLSTM/mLSTM: normalized matrix memory with gated covariance-style updates. +* Differential attention: subtracts a second learned linear-attention map. +* Memory fusion: token-wise mixture of sparse Cellular, Hedgehog and GDN2 paths. + +These are deliberately small, auditable reference implementations for controlled +ablation inside this repository. They are inspired by the papers, not drop-in +copies of the authors' optimized kernels. +""" +from __future__ import annotations + +import math +from typing import Iterable + +import torch +from torch import Tensor, nn +import torch.nn.functional as F + +from tinycenn_lm.cellular_attention import CellularAttentionLayer + + +VARIANTS = ( + "cellular_adaptive_maxpool5", + "cellular_hedgehog_global", + "cellular_kda_global", + "cellular_gdn2_global", + "cellular_xlstm_global", + "cellular_diff_hedgehog", + "cellular_memory_fusion", +) + +_HEDGEHOG_VARIANTS = { + "cellular_hedgehog_global", + "cellular_diff_hedgehog", + "cellular_memory_fusion", +} +_DIFF_VARIANTS = {"cellular_diff_hedgehog"} +_KDA_VARIANTS = {"cellular_kda_global"} +_GDN2_VARIANTS = {"cellular_gdn2_global", "cellular_memory_fusion"} +_XLSTM_VARIANTS = {"cellular_xlstm_global"} +_GLOBAL_VARIANTS = set(VARIANTS) - {"cellular_adaptive_maxpool5"} + + +class MemoryAugmentedCellularLayer(nn.Module): + """Adaptive+MaxPool Cellular Attention plus an optional global causal memory.""" + + def __init__( + self, + num_heads: int, + num_kv_heads: int, + head_dim: int, + feature_dim: int = 32, + variant: str = "cellular_adaptive_maxpool5", + dilations: Iterable[int] = (1, 2, 4, 8, 16, 32, 64, 128), + shifted_window: int = 8, + memory_rank: int = 16, + ): + super().__init__() + if variant not in VARIANTS: + raise ValueError(f"unknown variant {variant!r}; choose from {VARIANTS}") + if memory_rank < 4: + raise ValueError("memory_rank must be >= 4") + self.num_heads = int(num_heads) + self.num_kv_heads = int(num_kv_heads) + self.head_dim = int(head_dim) + self.feature_dim = int(feature_dim) + self.groups = self.num_heads // self.num_kv_heads + self.variant = variant + self.dilations = tuple(int(x) for x in dilations) + self.shifted_window = int(shifted_window) + self.memory_rank = int(memory_rank) + + # Local branch: preserve the strongest result already measured in the repo. + self.local = CellularAttentionLayer( + num_heads=self.num_heads, + num_kv_heads=self.num_kv_heads, + head_dim=self.head_dim, + feature_dim=self.feature_dim, + variant="cellular_adaptive_maxpool5", + dilations=self.dilations, + shifted_window=self.shifted_window, + ) + + # Alternate branches always begin as small perturbations of the local winner. + if self.has_global_memory() and self.variant != "cellular_memory_fusion": + self.branch_mix_logit = nn.Parameter(torch.full((self.num_heads,), -2.0)) + self.branch_log_gain = nn.Parameter(torch.zeros(self.num_heads)) + else: + self.register_parameter("branch_mix_logit", None) + self.register_parameter("branch_log_gain", None) + + # Shared low-rank Q/K maps for recurrent matrix memories. + if self.uses_kda() or self.uses_gdn2() or self.uses_xlstm(): + self.mem_q = nn.Parameter(torch.empty( + self.num_heads, self.memory_rank, self.head_dim + )) + self.mem_k = nn.Parameter(torch.empty( + self.num_heads, self.memory_rank, self.head_dim + )) + for h in range(self.num_heads): + nn.init.orthogonal_(self.mem_q[h]) + nn.init.orthogonal_(self.mem_k[h]) + else: + self.register_parameter("mem_q", None) + self.register_parameter("mem_k", None) + + # KDA/GDN2-style fine-grained forgetting and token-wise write rate. + if self.uses_kda() or self.uses_gdn2(): + self.decay_w = nn.Parameter(torch.zeros( + self.num_heads, self.memory_rank, self.head_dim + )) + self.decay_bias = nn.Parameter(torch.full( + (self.num_heads, self.memory_rank), 4.0 + )) + self.beta_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim)) + self.beta_bias = nn.Parameter(torch.full((self.num_heads,), -1.5)) + else: + self.register_parameter("decay_w", None) + self.register_parameter("decay_bias", None) + self.register_parameter("beta_w", None) + self.register_parameter("beta_bias", None) + + # GDN2-inspired decoupled key-side erase and value-side write controls. + if self.uses_gdn2(): + self.erase_w = nn.Parameter(torch.zeros( + self.num_heads, self.memory_rank, self.head_dim + )) + self.erase_bias = nn.Parameter(torch.full( + (self.num_heads, self.memory_rank), -0.5 + )) + self.write_scale = nn.Parameter(torch.zeros( + self.num_heads, self.head_dim + )) + self.write_bias = nn.Parameter(torch.full( + (self.num_heads, self.head_dim), -0.5 + )) + else: + self.register_parameter("erase_w", None) + self.register_parameter("erase_bias", None) + self.register_parameter("write_scale", None) + self.register_parameter("write_bias", None) + + # mLSTM-inspired matrix-memory gates. + if self.uses_xlstm(): + self.x_forget_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim)) + self.x_forget_bias = nn.Parameter(torch.full((self.num_heads,), 3.0)) + self.x_input_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim)) + self.x_input_bias = nn.Parameter(torch.full((self.num_heads,), -1.0)) + else: + self.register_parameter("x_forget_w", None) + self.register_parameter("x_forget_bias", None) + self.register_parameter("x_input_w", None) + self.register_parameter("x_input_bias", None) + + # Hedgehog-inspired trainable positive feature maps. The softmax over the + # learned feature axis enforces positivity and can become low-entropy/spiky. + if self.uses_hedgehog(): + self.hedge_q = nn.Parameter(torch.empty( + self.num_heads, self.memory_rank, self.head_dim + )) + self.hedge_k = nn.Parameter(torch.empty( + self.num_heads, self.memory_rank, self.head_dim + )) + self.hedge_q_bias = nn.Parameter(torch.zeros( + self.num_heads, self.memory_rank + )) + self.hedge_k_bias = nn.Parameter(torch.zeros( + self.num_heads, self.memory_rank + )) + self.hedge_log_sharpness = nn.Parameter(torch.zeros(self.num_heads)) + for h in range(self.num_heads): + nn.init.orthogonal_(self.hedge_q[h]) + nn.init.orthogonal_(self.hedge_k[h]) + else: + self.register_parameter("hedge_q", None) + self.register_parameter("hedge_k", None) + self.register_parameter("hedge_q_bias", None) + self.register_parameter("hedge_k_bias", None) + self.register_parameter("hedge_log_sharpness", None) + + if self.uses_differential(): + self.hedge2_q = nn.Parameter(torch.empty( + self.num_heads, self.memory_rank, self.head_dim + )) + self.hedge2_k = nn.Parameter(torch.empty( + self.num_heads, self.memory_rank, self.head_dim + )) + self.hedge2_q_bias = nn.Parameter(torch.zeros( + self.num_heads, self.memory_rank + )) + self.hedge2_k_bias = nn.Parameter(torch.zeros( + self.num_heads, self.memory_rank + )) + self.hedge2_log_sharpness = nn.Parameter(torch.zeros(self.num_heads)) + self.diff_lambda_logit = nn.Parameter(torch.full((self.num_heads,), -1.0)) + for h in range(self.num_heads): + nn.init.orthogonal_(self.hedge2_q[h]) + nn.init.orthogonal_(self.hedge2_k[h]) + else: + self.register_parameter("hedge2_q", None) + self.register_parameter("hedge2_k", None) + self.register_parameter("hedge2_q_bias", None) + self.register_parameter("hedge2_k_bias", None) + self.register_parameter("hedge2_log_sharpness", None) + self.register_parameter("diff_lambda_logit", None) + + # The fusion candidate makes the branch choice input-dependent. + if self.variant == "cellular_memory_fusion": + self.fusion_gate_w = nn.Parameter(torch.zeros( + self.num_heads, 3, self.head_dim + )) + prior = torch.tensor([2.0, -1.0, -1.0]) + self.fusion_gate_bias = nn.Parameter( + prior[None, :].expand(self.num_heads, -1).clone() + ) + else: + self.register_parameter("fusion_gate_w", None) + self.register_parameter("fusion_gate_bias", None) + + @property + def config(self) -> dict: + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "variant": self.variant, + "dilations": list(self.dilations), + "shifted_window": self.shifted_window, + "memory_rank": self.memory_rank, + } + + def has_global_memory(self) -> bool: + return self.variant in _GLOBAL_VARIANTS + + def uses_hedgehog(self) -> bool: + return self.variant in _HEDGEHOG_VARIANTS + + def uses_differential(self) -> bool: + return self.variant in _DIFF_VARIANTS + + def uses_kda(self) -> bool: + return self.variant in _KDA_VARIANTS + + def uses_gdn2(self) -> bool: + return self.variant in _GDN2_VARIANTS + + def uses_xlstm(self) -> bool: + return self.variant in _XLSTM_VARIANTS + + def _repeat_kv(self, x: Tensor) -> Tensor: + return x.repeat_interleave(self.groups, dim=1) + + @staticmethod + def _project(x: Tensor, weight: Tensor) -> Tensor: + return torch.einsum("bhtd,hrd->bhtr", x, weight) + + def _memory_qk(self, q: Tensor, k: Tensor) -> tuple[Tensor, Tensor, Tensor]: + assert self.mem_q is not None and self.mem_k is not None + kh = self._repeat_kv(k) + qm = F.normalize(self._project(q, self.mem_q), dim=-1) + km = F.normalize(self._project(kh, self.mem_k), dim=-1) + return qm, km, kh + + def _hedgehog_features( + self, + x: Tensor, + weight: Tensor, + bias: Tensor, + log_sharpness: Tensor, + ) -> Tensor: + logits = self._project(x, weight) + bias[None, :, None, :] + sharpness = log_sharpness.clamp(-1.4, 2.1).exp()[None, :, None, None] + # sqrt(rank) keeps q.k magnitudes from vanishing as rank grows. + return logits.mul(sharpness).softmax(dim=-1) * math.sqrt(self.memory_rank) + + def _hedgehog_linear( + self, + q: Tensor, + k: Tensor, + v: Tensor, + *, + second: bool = False, + ) -> Tensor: + kh, vh = self._repeat_kv(k), self._repeat_kv(v) + if second: + assert self.hedge2_q is not None and self.hedge2_k is not None + assert self.hedge2_q_bias is not None and self.hedge2_k_bias is not None + assert self.hedge2_log_sharpness is not None + qf = self._hedgehog_features( + q, self.hedge2_q, self.hedge2_q_bias, self.hedge2_log_sharpness + ) + kf = self._hedgehog_features( + kh, self.hedge2_k, self.hedge2_k_bias, self.hedge2_log_sharpness + ) + else: + assert self.hedge_q is not None and self.hedge_k is not None + assert self.hedge_q_bias is not None and self.hedge_k_bias is not None + assert self.hedge_log_sharpness is not None + qf = self._hedgehog_features( + q, self.hedge_q, self.hedge_q_bias, self.hedge_log_sharpness + ) + kf = self._hedgehog_features( + kh, self.hedge_k, self.hedge_k_bias, self.hedge_log_sharpness + ) + + kv = torch.einsum("bhtr,bhtd->bhtrd", kf, vh).cumsum(dim=2) + kz = kf.cumsum(dim=2) + numerator = torch.einsum("bhtr,bhtrd->bhtd", qf, kv) + denominator = torch.einsum("bhtr,bhtr->bht", qf, kz) + return numerator / denominator.clamp_min(1e-6)[..., None] + + def _delta_memory(self, q: Tensor, k: Tensor, v: Tensor, *, gdn2: bool) -> Tensor: + qm, km, kh = self._memory_qk(q, k) + vh = self._repeat_kv(v) + assert self.decay_w is not None and self.decay_bias is not None + assert self.beta_w is not None and self.beta_bias is not None + + decay = torch.sigmoid( + torch.einsum("bhtd,hrd->bhtr", kh, self.decay_w) + + self.decay_bias[None, :, None, :] + ) + beta = torch.sigmoid( + torch.einsum("bhtd,hd->bht", kh, self.beta_w) + + self.beta_bias[None, :, None] + ) + + if gdn2: + assert self.erase_w is not None and self.erase_bias is not None + assert self.write_scale is not None and self.write_bias is not None + erase = torch.sigmoid( + torch.einsum("bhtd,hrd->bhtr", kh, self.erase_w) + + self.erase_bias[None, :, None, :] + ) + write = torch.sigmoid( + vh * self.write_scale[None, :, None, :] + + self.write_bias[None, :, None, :] + ) + else: + erase = None + write = None + + b, h, t, _ = qm.shape + state = torch.zeros( + b, h, self.memory_rank, self.head_dim, + device=q.device, dtype=q.dtype, + ) + outputs = [] + for i in range(t): + state = state * decay[:, :, i, :, None] + pred = torch.einsum("bhr,bhrd->bhd", km[:, :, i], state) + error = vh[:, :, i] - pred + key_write = km[:, :, i] + if gdn2: + assert erase is not None and write is not None + key_write = key_write * erase[:, :, i] + error = error * write[:, :, i] + update = torch.einsum("bhr,bhd->bhrd", key_write, error) + state = state + beta[:, :, i, None, None] * update + outputs.append(torch.einsum("bhr,bhrd->bhd", qm[:, :, i], state)) + return torch.stack(outputs, dim=2) + + def _xlstm_memory(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor: + qm, km, kh = self._memory_qk(q, k) + vh = self._repeat_kv(v) + assert self.x_forget_w is not None and self.x_forget_bias is not None + assert self.x_input_w is not None and self.x_input_bias is not None + + forget = torch.sigmoid( + torch.einsum("bhtd,hd->bht", kh, self.x_forget_w) + + self.x_forget_bias[None, :, None] + ) + inp = torch.sigmoid( + torch.einsum("bhtd,hd->bht", kh, self.x_input_w) + + self.x_input_bias[None, :, None] + ) + b, h, t, _ = qm.shape + memory = torch.zeros( + b, h, self.memory_rank, self.head_dim, + device=q.device, dtype=q.dtype, + ) + normalizer = torch.zeros( + b, h, self.memory_rank, device=q.device, dtype=q.dtype + ) + outputs = [] + for i in range(t): + f = forget[:, :, i, None, None] + ii = inp[:, :, i, None, None] + outer = torch.einsum("bhr,bhd->bhrd", km[:, :, i], vh[:, :, i]) + memory = f * memory + ii * outer + normalizer = ( + forget[:, :, i, None] * normalizer + + inp[:, :, i, None] * km[:, :, i] + ) + numerator = torch.einsum("bhr,bhrd->bhd", qm[:, :, i], memory) + denominator = torch.einsum( + "bhr,bhr->bh", qm[:, :, i], normalizer + ).abs().clamp_min(1.0) + outputs.append(numerator / denominator[..., None]) + return torch.stack(outputs, dim=2) + + def _merge(self, local: Tensor, branch: Tensor) -> Tensor: + assert self.branch_mix_logit is not None and self.branch_log_gain is not None + gate = self.branch_mix_logit.sigmoid()[None, :, None, None] + gain = self.branch_log_gain.clamp(-2, 2).exp()[None, :, None, None] + branch = branch * gain + return local + gate * (branch - local) + + def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor: + if q.ndim != 4 or k.ndim != 4 or v.ndim != 4: + raise ValueError("expected Q/K/V as [batch, heads, time, dim]") + local = self.local(q, k, v) + if self.variant == "cellular_adaptive_maxpool5": + return local + + q = q.to(local.dtype) + k = k.to(local.dtype) + v = v.to(local.dtype) + + if self.variant == "cellular_hedgehog_global": + return self._merge(local, self._hedgehog_linear(q, k, v)) + + if self.variant == "cellular_diff_hedgehog": + assert self.diff_lambda_logit is not None + first = self._hedgehog_linear(q, k, v) + second = self._hedgehog_linear(q, k, v, second=True) + lam = 0.5 * self.diff_lambda_logit.sigmoid()[None, :, None, None] + return self._merge(local, first - lam * second) + + if self.variant == "cellular_kda_global": + return self._merge(local, self._delta_memory(q, k, v, gdn2=False)) + + if self.variant == "cellular_gdn2_global": + return self._merge(local, self._delta_memory(q, k, v, gdn2=True)) + + if self.variant == "cellular_xlstm_global": + return self._merge(local, self._xlstm_memory(q, k, v)) + + if self.variant == "cellular_memory_fusion": + assert self.fusion_gate_w is not None and self.fusion_gate_bias is not None + hedge = self._hedgehog_linear(q, k, v) + gdn2 = self._delta_memory(q, k, v, gdn2=True) + logits = ( + torch.einsum("bhtd,hcd->bhtc", q, self.fusion_gate_w) + + self.fusion_gate_bias[None, :, None, :] + ) + weights = logits.softmax(dim=-1) + return ( + weights[..., 0, None] * local + + weights[..., 1, None] * hedge + + weights[..., 2, None] * gdn2 + ) + + raise ValueError(self.variant) + + def max_score_pairs(self, context: int) -> int: + # This counts only the sparse Cellular softmax branch. Global memories + # are O(T * memory_rank * head_dim), not pairwise T^2 score matrices. + return self.local.max_score_pairs(context) + + def receptive_field_tokens(self) -> int: + return self.local.receptive_field_tokens() + + def max_neighbors_per_step(self) -> int: + return self.local.max_neighbors_per_step() diff --git a/src/tinycenn_lm/modeling.py b/src/tinycenn_lm/modeling.py new file mode 100644 index 0000000000000000000000000000000000000000..e910a6f5286fe77dbceb55f9ca51c9058bdb6140 --- /dev/null +++ b/src/tinycenn_lm/modeling.py @@ -0,0 +1,251 @@ +from __future__ import annotations + +import json +from pathlib import Path +from typing import Sequence + +import torch +from torch import Tensor, nn + +from .cenn import CeNNConfig, FastCeNNCore + + +DEFAULT_BASE_MODEL = "arnir0/Tiny-LLM" + + +def _floating_reference_parameter(module: nn.Module) -> nn.Parameter | None: + """Return a representative floating-point parameter for device/dtype alignment.""" + return next((p for p in module.parameters() if p.is_floating_point()), None) + + +class HybridDecoderLayer(nn.Module): + """Wrap a pretrained decoder layer with a zero-init recurrent CeNN residual. + + The base layer is kept intact, including its attention behavior. TinyCeNN-LM + v0.1 intentionally disables Transformer-only KV caching because the CeNN branch + also needs its own recurrent neighborhood state for exact incremental decoding. + + The newly created CeNN branch inherits the pretrained decoder layer's floating + point dtype/device. This is essential for BF16/FP16 inference: a BF16 hidden + state cannot be convolved with an FP32 CeNN kernel without an explicit cast. + """ + + def __init__(self, base_layer: nn.Module, config: CeNNConfig) -> None: + super().__init__() + self.base_layer = base_layer + self.cenn = FastCeNNCore(config) + + reference = _floating_reference_parameter(base_layer) + if reference is None: + self.residual_scale = nn.Parameter(torch.ones(())) + else: + self.cenn.to(device=reference.device, dtype=reference.dtype) + self.residual_scale = nn.Parameter( + torch.ones((), device=reference.device, dtype=reference.dtype) + ) + + def forward(self, *args, **kwargs): + if kwargs.get("use_cache", False): + raise RuntimeError( + "TinyCeNN-LM v0.1 requires use_cache=False. The Transformer KV cache " + "does not contain the per-step CeNN neighborhood state needed for exact " + "incremental generation. Full-prefix generation is correct; a dedicated " + "streaming CeNN cache is planned for a later version." + ) + outputs = self.base_layer(*args, **kwargs) + + if torch.is_tensor(outputs): + return outputs + self.residual_scale * self.cenn(outputs) + + if isinstance(outputs, tuple): + hidden = outputs[0] + hidden = hidden + self.residual_scale * self.cenn(hidden) + return (hidden, *outputs[1:]) + + if isinstance(outputs, list): + hidden = outputs[0] + hidden = hidden + self.residual_scale * self.cenn(hidden) + return [hidden, *outputs[1:]] + + raise TypeError( + "Unsupported decoder-layer output type: " + f"{type(outputs)!r}. Expected Tensor, tuple, or list." + ) + + +def _get_decoder_layers(model: nn.Module) -> nn.ModuleList: + candidates = ( + ("model", "layers"), + ("model", "model", "layers"), + ) + for path in candidates: + obj = model + try: + for name in path: + obj = getattr(obj, name) + except AttributeError: + continue + if isinstance(obj, nn.ModuleList): + return obj + raise ValueError( + "Could not locate decoder layers. TinyCeNN-LM currently targets " + "Llama-family causal language models such as arnir0/Tiny-LLM." + ) + + +def inject_cenn( + model: nn.Module, + config: CeNNConfig | None = None, + layer_indices: Sequence[int] = (0,), +) -> nn.Module: + layers = _get_decoder_layers(model) + hidden_size = int(getattr(model.config, "hidden_size")) + if config is None: + config = CeNNConfig(hidden_size=hidden_size) + elif config.hidden_size != hidden_size: + raise ValueError( + f"CeNN hidden_size={config.hidden_size} does not match " + f"model hidden_size={hidden_size}" + ) + + # A Transformer KV cache alone is insufficient for the recurrent CeNN state. + # Disable it globally to prevent a silent train/inference mismatch. + if hasattr(model, "config"): + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + + for index in layer_indices: + if index < 0 or index >= len(layers): + raise IndexError(f"layer index {index} out of range [0, {len(layers)})") + if isinstance(layers[index], HybridDecoderLayer): + raise ValueError(f"layer {index} already has a CeNN adapter") + layers[index] = HybridDecoderLayer(layers[index], config) + return model + + +def freeze_for_adapter_training( + model: nn.Module, + train_lm_head: bool = False, + train_embeddings: bool = False, +) -> None: + for parameter in model.parameters(): + parameter.requires_grad = False + + for module in model.modules(): + if isinstance(module, HybridDecoderLayer): + for parameter in module.cenn.parameters(): + parameter.requires_grad = True + module.residual_scale.requires_grad = True + + if train_lm_head and hasattr(model, "lm_head"): + for parameter in model.lm_head.parameters(): + parameter.requires_grad = True + if train_embeddings: + embeddings = model.get_input_embeddings() + for parameter in embeddings.parameters(): + parameter.requires_grad = True + + +def trainable_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + return { + "total": total, + "trainable": trainable, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def _adapter_state_dict(model: nn.Module) -> dict[str, Tensor]: + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".cenn." in name or name.endswith(".residual_scale"): + state[name] = tensor.detach().cpu() + if not state: + raise ValueError("no CeNN adapter found in model") + return state + + +def save_adapter( + model: nn.Module, + output_dir: str | Path, + *, + base_model: str = DEFAULT_BASE_MODEL, + layer_indices: Sequence[int] = (0,), + config: CeNNConfig, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + weights_path = output_dir / "cenn_adapter.pt" + metadata_path = output_dir / "cenn_config.json" + + torch.save(_adapter_state_dict(model), weights_path) + metadata = { + "format_version": 1, + "base_model": base_model, + "layer_indices": list(layer_indices), + "cenn": config.to_dict(), + } + metadata_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + return output_dir + + +def load_adapter( + model: nn.Module, + adapter_dir: str | Path, + *, + map_location: str | torch.device = "cpu", + strict: bool = True, +) -> nn.Module: + adapter_dir = Path(adapter_dir) + state = torch.load( + adapter_dir / "cenn_adapter.pt", + map_location=map_location, + weights_only=True, + ) + incompatible = model.load_state_dict(state, strict=False) + unexpected = [k for k in incompatible.unexpected_keys if ".cenn." not in k] + if strict and unexpected: + raise RuntimeError(f"unexpected adapter keys: {unexpected}") + missing_adapter = [ + k + for k in _adapter_state_dict(model) + if k in incompatible.missing_keys + ] + if strict and missing_adapter: + raise RuntimeError(f"missing adapter keys: {missing_adapter}") + return model + + +def build_from_adapter( + adapter_dir: str | Path, + *, + device: str | torch.device | None = None, + dtype: torch.dtype | None = None, + attn_implementation: str = "sdpa", +): + from transformers import AutoModelForCausalLM + + adapter_dir = Path(adapter_dir) + metadata = json.loads((adapter_dir / "cenn_config.json").read_text()) + base_model = metadata["base_model"] + config = CeNNConfig.from_dict(metadata["cenn"]) + layer_indices = tuple(metadata["layer_indices"]) + + kwargs = {"attn_implementation": attn_implementation} + if dtype is not None: + # Modern Transformers uses `dtype`; `torch_dtype` is deprecated. + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(base_model, **kwargs) + inject_cenn(model, config=config, layer_indices=layer_indices) + load_adapter(model, adapter_dir) + + move_kwargs: dict[str, object] = {} + if device is not None: + move_kwargs["device"] = device + if dtype is not None: + move_kwargs["dtype"] = dtype + if move_kwargs: + model.to(**move_kwargs) + return model diff --git a/src/tinycenn_lm/moe.py b/src/tinycenn_lm/moe.py new file mode 100644 index 0000000000000000000000000000000000000000..e548096b9b71ea41512ec43dd4b909965c96cd1a --- /dev/null +++ b/src/tinycenn_lm/moe.py @@ -0,0 +1,334 @@ +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Sequence + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from .cenn import CeNNConfig, CausalDepthwiseNeighborhood, StableRMSNorm +from .modeling import DEFAULT_BASE_MODEL, _get_decoder_layers + + +@dataclass(frozen=True) +class MoECeNNConfig: + hidden_size: int = 192 + kernel_size: int = 3 + expansion: int = 4 + steps: int = 7 + dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64) + rms_norm_eps: float = 1e-5 + dropout: float = 0.0 + num_experts: int = 8 + top_k: int = 2 + router_noise_std: float = 1e-3 + + def validate(self) -> None: + CeNNConfig( + hidden_size=self.hidden_size, + kernel_size=self.kernel_size, + expansion=self.expansion, + steps=self.steps, + dilations=self.dilations, + rms_norm_eps=self.rms_norm_eps, + dropout=self.dropout, + ).validate() + if self.num_experts < 2: + raise ValueError("num_experts must be >= 2") + if not 1 <= self.top_k <= self.num_experts: + raise ValueError("top_k must be in [1, num_experts]") + if self.router_noise_std < 0: + raise ValueError("router_noise_std must be >= 0") + + def to_dict(self) -> dict: + data = asdict(self) + data["dilations"] = list(self.dilations) + return data + + @classmethod + def from_dict(cls, data: dict) -> "MoECeNNConfig": + data = dict(data) + if "dilations" in data: + data["dilations"] = tuple(data["dilations"]) + return cls(**data) + + +class SwiGLUExpert(nn.Module): + def __init__(self, hidden_size: int, expansion: int, dropout: float = 0.0) -> None: + super().__init__() + inner = hidden_size * expansion + self.in_proj = nn.Linear(hidden_size, inner * 2, bias=False) + self.out_proj = nn.Linear(inner, hidden_size, bias=False) + self.dropout = nn.Dropout(dropout) + nn.init.zeros_(self.out_proj.weight) + + def forward(self, x: Tensor) -> Tensor: + a, b = self.in_proj(x).chunk(2, dim=-1) + return self.dropout(self.out_proj(F.silu(a) * b)) + + +class Top2Router(nn.Module): + def __init__(self, hidden_size: int, num_experts: int, top_k: int, noise_std: float) -> None: + super().__init__() + self.num_experts = num_experts + self.top_k = top_k + self.noise_std = noise_std + self.proj = nn.Linear(hidden_size, num_experts, bias=False) + nn.init.normal_(self.proj.weight, mean=0.0, std=noise_std) + + def forward(self, x: Tensor) -> tuple[Tensor, Tensor, dict[str, Tensor]]: + logits = self.proj(x).float() + probs = F.softmax(logits, dim=-1) + top_values, top_indices = torch.topk(probs, k=self.top_k, dim=-1) + top_weights = top_values / top_values.sum(dim=-1, keepdim=True).clamp_min(1e-9) + + assignment = F.one_hot(top_indices, num_classes=self.num_experts).float().sum(dim=-2) + assignment = assignment / float(self.top_k) + expert_fraction = assignment.mean(dim=(0, 1)) + probability_fraction = probs.mean(dim=(0, 1)) + load_balance = self.num_experts * torch.sum(expert_fraction * probability_fraction) + z_loss = torch.logsumexp(logits, dim=-1).pow(2).mean() + entropy = -(probs * probs.clamp_min(1e-9).log()).sum(dim=-1).mean() + return top_indices, top_weights.to(dtype=x.dtype), { + "load_balance": load_balance, + "z_loss": z_loss, + "entropy": entropy, + "expert_fraction": expert_fraction, + "probability_fraction": probability_fraction, + } + + +class MoESharedCeNNCell(nn.Module): + def __init__(self, config: MoECeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + self.norm = StableRMSNorm(config.hidden_size, config.rms_norm_eps) + self.neighborhood = CausalDepthwiseNeighborhood(config.hidden_size, config.kernel_size) + self.gate_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=True) + nn.init.constant_(self.gate_proj.bias, -1.0) + self.router = Top2Router( + config.hidden_size, config.num_experts, config.top_k, config.router_noise_std + ) + self.experts = nn.ModuleList( + SwiGLUExpert(config.hidden_size, config.expansion, config.dropout) + for _ in range(config.num_experts) + ) + + def forward(self, state: Tensor, dilation: int, step_scale: float) -> tuple[Tensor, dict[str, Tensor]]: + x = self.norm(state) + local = self.neighborhood(x, dilation=dilation) + top_idx, top_weight, stats = self.router(local) + + flat = local.reshape(-1, local.shape[-1]) + idx_flat = top_idx.reshape(-1, self.config.top_k) + weight_flat = top_weight.reshape(-1, self.config.top_k) + update = torch.zeros_like(flat) + for expert_id, expert in enumerate(self.experts): + selected = idx_flat.eq(expert_id) + positions = selected.nonzero(as_tuple=False) + if positions.numel() == 0: + continue + rows = positions[:, 0] + slots = positions[:, 1] + expert_out = expert(flat.index_select(0, rows)) + weighted = expert_out * weight_flat[rows, slots].unsqueeze(-1) + update = update.index_add(0, rows, weighted) + update = update.view_as(local) + gate = torch.sigmoid(self.gate_proj(local)) + return state + step_scale * gate * update, stats + + +class FastMoECeNNCore(nn.Module): + def __init__(self, config: MoECeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + self.cell = MoESharedCeNNCell(config) + self.last_router_stats: dict[str, Tensor] = {} + + def forward(self, hidden_states: Tensor) -> Tensor: + initial = hidden_states + state = hidden_states + step_scale = self.config.steps ** -0.5 + accum: dict[str, Tensor] = {} + expert_fraction = None + probability_fraction = None + for step in range(self.config.steps): + dilation = self.config.dilations[step % len(self.config.dilations)] + state, stats = self.cell(state, dilation=dilation, step_scale=step_scale) + for key in ("load_balance", "z_loss", "entropy"): + accum[key] = accum.get(key, stats[key].new_zeros(())) + stats[key] + expert_fraction = stats["expert_fraction"] if expert_fraction is None else expert_fraction + stats["expert_fraction"] + probability_fraction = stats["probability_fraction"] if probability_fraction is None else probability_fraction + stats["probability_fraction"] + self.last_router_stats = { + "load_balance": accum["load_balance"] / self.config.steps, + "z_loss": accum["z_loss"] / self.config.steps, + "entropy": accum["entropy"] / self.config.steps, + "expert_fraction": expert_fraction / self.config.steps, + "probability_fraction": probability_fraction / self.config.steps, + } + return state - initial + + @property + def receptive_field(self) -> int: + radius = sum(self.config.dilations[i % len(self.config.dilations)] for i in range(self.config.steps)) + return 1 + (self.config.kernel_size - 1) * radius + + +class MoECeNNReplacementLayer(nn.Module): + def __init__(self, config: MoECeNNConfig, *, device=None, dtype=None) -> None: + super().__init__() + self.config = config + self.cenn = FastMoECeNNCore(config) + if device is not None or dtype is not None: + kwargs = {} + if device is not None: + kwargs["device"] = device + if dtype is not None: + kwargs["dtype"] = dtype + self.cenn.to(**kwargs) + + def forward(self, hidden_states: Tensor, *args, **kwargs) -> Tensor: + if kwargs.get("use_cache", False): + raise RuntimeError("MoE-CeNN student requires use_cache=False") + if kwargs.get("output_attentions", False): + raise RuntimeError("MoE-CeNN student has no attention matrices") + return hidden_states + self.cenn(hidden_states) + + +def replace_transformer_with_moe_cenn(model: nn.Module, config: MoECeNNConfig, layer_indices: Sequence[int] = (0,)) -> nn.Module: + layers = _get_decoder_layers(model) + if config.hidden_size != int(model.config.hidden_size): + raise ValueError("MoE-CeNN hidden size does not match base model") + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + for index in layer_indices: + old = layers[index] + reference = next((p for p in old.parameters() if p.is_floating_point()), None) + layers[index] = MoECeNNReplacementLayer( + config, + device=reference.device if reference is not None else None, + dtype=reference.dtype if reference is not None else None, + ) + return model + + +def freeze_moe_student_interfaces(model: nn.Module) -> None: + for parameter in model.parameters(): + parameter.requires_grad = False + for module in model.modules(): + if isinstance(module, MoECeNNReplacementLayer): + for parameter in module.parameters(): + parameter.requires_grad = True + + +def warmstart_moe_from_plain_cenn(model: nn.Module, plain_student_dir: str | Path) -> None: + """Initialize shared dynamics and clone the trained dense FFN into every expert. + + With identical expert weights, Top-2 weighted routing initially reproduces the + dense CeNN FFN output (weights sum to one), while tiny router noise allows the + experts to specialize during training. + """ + state = torch.load(Path(plain_student_dir) / "cenn_student.pt", map_location="cpu", weights_only=True) + target = model.state_dict() + copied = 0 + for target_name in list(target): + if ".cenn.cell.experts." in target_name: + suffix = target_name.split(".experts.", 1)[1].split(".", 1)[1] + prefix = target_name.split(".cenn.cell.experts.", 1)[0] + ".cenn.cell." + source_name = prefix + suffix + elif any(part in target_name for part in (".cenn.cell.norm.", ".cenn.cell.neighborhood.", ".cenn.cell.gate_proj.")): + source_name = target_name + elif ".cenn." not in target_name and target_name in state: + # v2 dense checkpoints may include adapted language interfaces. + source_name = target_name + else: + continue + if source_name in state and state[source_name].shape == target[target_name].shape: + target[target_name].copy_(state[source_name].to(dtype=target[target_name].dtype)) + copied += 1 + if copied == 0: + raise RuntimeError("could not map plain CeNN weights into MoE-CeNN model") + model.load_state_dict(target, strict=False) + model._cenn_interface_keys = tuple(name for name in state if ".cenn." not in name) + + +def moe_router_stats(model: nn.Module) -> dict[str, Tensor]: + layer = next((m for m in model.modules() if isinstance(m, MoECeNNReplacementLayer)), None) + if layer is None or not layer.cenn.last_router_stats: + raise RuntimeError("router statistics unavailable; run a forward pass first") + return layer.cenn.last_router_stats + + +def _moe_state_dict(model: nn.Module) -> dict[str, Tensor]: + interfaces = set(getattr(model, "_cenn_interface_keys", ())) + state = {name: tensor.detach().cpu() for name, tensor in model.state_dict().items() + if ".cenn." in name or name in interfaces} + if not state: + raise ValueError("no MoE-CeNN weights found") + return state + + +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: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + state = _moe_state_dict(model) + torch.save(state, output_dir / "moe_cenn_student.pt") + metadata = { + "format_version": 2, + "architecture": "moe-cenn-top2-replacement", + "base_model": base_model, + "layer_indices": list(layer_indices), + "moe_cenn": config.to_dict(), + "state_keys": sorted(state), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "moe_student_config.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8") + return output_dir + + +def load_moe_cenn_student_weights(model: nn.Module, student_dir: str | Path, *, map_location="cpu", strict: bool = True) -> nn.Module: + state = torch.load(Path(student_dir) / "moe_cenn_student.pt", map_location=map_location, weights_only=True) + metadata = json.loads((Path(student_dir) / "moe_student_config.json").read_text()) + core_keys = {name for name in model.state_dict() if ".cenn." in name} + expected = set(metadata.get("state_keys", core_keys)) | core_keys + missing = sorted(expected - state.keys()) + unexpected = sorted(state.keys() - expected | state.keys() - model.state_dict().keys()) + if strict and missing: + raise RuntimeError(f"missing MoE-CeNN keys: {missing}") + if strict and unexpected: + raise RuntimeError(f"unexpected MoE-CeNN keys: {unexpected}") + model.load_state_dict(state, strict=False) + model._cenn_interface_keys = tuple(name for name in state if ".cenn." not in name) + return model + + +def build_moe_cenn_student(student_dir: str | Path, *, device=None, dtype=None, attn_implementation: str = "sdpa"): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + metadata = json.loads((student_dir / "moe_student_config.json").read_text()) + if metadata.get("architecture") != "moe-cenn-top2-replacement": + raise ValueError("checkpoint is not an MoE-CeNN Top-2 student") + kwargs = {"attn_implementation": attn_implementation} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs) + config = MoECeNNConfig.from_dict(metadata["moe_cenn"]) + replace_transformer_with_moe_cenn(model, config, tuple(metadata["layer_indices"])) + load_moe_cenn_student_weights(model, student_dir) + move_kwargs = {} + if device is not None: + move_kwargs["device"] = device + if dtype is not None: + move_kwargs["dtype"] = dtype + if move_kwargs: + model.to(**move_kwargs) + model.config.use_cache = False + return model diff --git a/src/tinycenn_lm/optimized_memory.py b/src/tinycenn_lm/optimized_memory.py new file mode 100644 index 0000000000000000000000000000000000000000..17871c10d44887883c26ff54536fb3b75237f50f --- /dev/null +++ b/src/tinycenn_lm/optimized_memory.py @@ -0,0 +1,308 @@ +"""Parallel normalized memory with disjoint exact sink/local attention. + +Independent research implementation. See OPTIMIZED_MEMORY.md for derivation. +The block partition is terraced: current and previous block are exact, older +non-sink blocks are compressed. No T-by-T matrix or triangular solve is used +by the new memory candidates. +""" +from dataclasses import dataclass +import math + +import torch +import torch.nn.functional as F +from torch import nn + +VARIANTS = ("cenn_linear", "cenn_partition", "sink_window", "transformer_readout") + + +@dataclass +class MemoryState: + numerator: torch.Tensor + denominator: torch.Tensor + keys: torch.Tensor + values: torch.Tensor + sinks_k: torch.Tensor + sinks_v: torch.Tensor + position: int + + @property + def nbytes(self): + return sum(t.numel() * t.element_size() for t in ( + self.numerator, self.denominator, self.keys, self.values, + self.sinks_k, self.sinks_v)) + + +class OptimizedMemory(nn.Module): + def __init__(self, num_heads, num_kv_heads, head_dim, feature_dim=64, + variant="cenn_partition", block_size=32, sink_tokens=4, + compute_dtype="float32"): + super().__init__() + if variant not in VARIANTS: + raise ValueError(f"Unknown variant {variant}") + if min(num_heads, num_kv_heads, head_dim, feature_dim, block_size) < 1: + raise ValueError("Dimensions must be positive") + if num_heads % num_kv_heads or feature_dim % 2 or sink_tokens < 0: + raise ValueError("Require valid GQA, even feature_dim, nonnegative sink count") + if compute_dtype not in ("float32", "float16", "bfloat16"): + raise ValueError("Unknown compute dtype") + self.num_heads, self.num_kv_heads = num_heads, num_kv_heads + self.head_dim, self.feature_dim = head_dim, feature_dim + self.groups = num_heads // num_kv_heads + self.variant, self.block_size, self.sink_tokens = variant, block_size, sink_tokens + self.compute_dtype = compute_dtype + self.has_memory = variant in ("cenn_linear", "cenn_partition") + if self.has_memory: + weight = torch.randn(num_kv_heads, feature_dim // 2, head_dim) / math.sqrt(head_dim) + self.wk = nn.Parameter(weight.clone()) + self.wq = nn.Parameter(weight.repeat_interleave(self.groups, dim=0).clone()) + if variant == "cenn_partition": + self.log_mass = nn.Parameter(torch.full((num_heads,), math.log(feature_dim))) + self.mass_w = nn.Parameter(torch.zeros(num_heads, head_dim)) + readout = torch.eye(head_dim).repeat(num_heads, 1, 1) + if variant == "sink_window": + self.register_buffer("readout", readout) + else: + self.readout = nn.Parameter(readout) + + @property + def config(self): + return dict(num_heads=self.num_heads, num_kv_heads=self.num_kv_heads, + head_dim=self.head_dim, feature_dim=self.feature_dim, + variant=self.variant, block_size=self.block_size, + sink_tokens=self.sink_tokens, compute_dtype=self.compute_dtype) + + def mm(self, a, b): + dtype = getattr(torch, self.compute_dtype) if a.is_cuda else torch.float32 + return torch.matmul(a.to(dtype), b.to(dtype)).float() + + def attend(self, q, k, v, mask=None, causal=False): + dtype = getattr(torch, self.compute_dtype) if q.is_cuda else torch.float32 + shape = q.shape + if q.ndim == 5: + b, h, n, c, d = shape + length = k.shape[-2] + q, k, v = (x.reshape(b * h * n, 1, x.shape[-2], d) for x in (q, k, v)) + if mask is not None: + mask = mask.expand(b, h, n, c, length).reshape(b * h * n, 1, c, length) + out = F.scaled_dot_product_attention( + q.to(dtype), k.to(dtype), v.to(dtype), attn_mask=mask, is_causal=causal + ).float() + return out.reshape(shape) + + def features(self, x, query=False): + weight = self.wq if query else self.wk + # Retain input norms: do not impose the previous experiment's L2 normalization. + logits = self.mm(x.float() / self.head_dim ** 0.25, + weight.view(1, weight.shape[0], *([1] * (x.ndim - 4)), + weight.shape[1], weight.shape[2]).transpose(-1, -2)) + return torch.cat((logits, -logits), dim=-1).softmax(dim=-1) + + def calibrate(self, output): + return self.mm(output, self.readout[None]) + + def empty_state(self, k): + b = k.shape[0] + empty = k.new_empty(b, self.num_kv_heads, 0, self.head_dim) + n = k.new_zeros(b, self.num_kv_heads, self.feature_dim, self.head_dim + ) if self.has_memory else k.new_empty(0) + z = k.new_zeros(b, self.num_kv_heads, self.feature_dim + ) if self.has_memory else k.new_empty(0) + return MemoryState(n, z, empty.clone(), empty.clone(), empty.clone(), + empty.clone(), 0) + + def combine(self, q, local_k, local_v, valid, numerator=None, denominator=None): + """Shared log-stabilized denominator for exact and compressed contributions.""" + if self.variant == "sink_window": + return self.attend(q, local_k.repeat_interleave(self.groups, 1), + local_v.repeat_interleave(self.groups, 1), mask=valid) + scores = self.mm(q / math.sqrt(self.head_dim), + local_k.repeat_interleave(self.groups, dim=1).transpose(-1, -2)) + scores = scores.masked_fill(~valid, float("-inf")) + maximum = scores.amax(dim=-1, keepdim=True) + if self.variant == "cenn_partition": + qp = self.features(q, query=True) + num = self.mm(qp, numerator.repeat_interleave(self.groups, dim=1)) + den = self.mm(qp, denominator.repeat_interleave(self.groups, dim=1).unsqueeze(-1)) + raw_q = F.normalize(q.float(), dim=-1) + # Works for [B,H,T,D] and [B,H,N,C,D]. + mass_shape = [1, self.num_heads] + [1] * (q.ndim - 3) + mass_w_shape = [1, self.num_heads] + [1] * (q.ndim - 3) + [self.head_dim] + mass = self.log_mass.view(mass_shape) + ( + raw_q * self.mass_w.view(mass_w_shape)).sum(-1) + log_global = mass.clamp(-12, 12).unsqueeze(-1) + den.clamp_min(1e-30).log() + log_global = torch.where(den > 0, log_global, float("-inf")) + maximum = torch.maximum(maximum, log_global) + global_weight = (log_global - maximum).exp() + global_value = num / den.clamp_min(1e-20) + weights = (scores - maximum).exp() + local_num = self.mm(weights, local_v.repeat_interleave(self.groups, dim=1)) + local_den = weights.sum(dim=-1, keepdim=True) + if self.variant == "cenn_partition": + return (local_num + global_weight * global_value) / ( + local_den + global_weight).clamp_min(1e-20) + return local_num / local_den.clamp_min(1e-20) + + @staticmethod + def prefix(blocks, delay): + cumulative = blocks.cumsum(dim=2) + zeros = torch.zeros_like(blocks[:, :, :1]).expand( + *blocks.shape[:2], delay, *blocks.shape[3:]) + return torch.cat((zeros, cumulative), dim=2)[:, :, :blocks.shape[2]] + + def prefill(self, q, k, v, need_state=True): + b, _, t, d = q.shape + state = self.empty_state(k) if need_state else None + if self.variant == "transformer_readout": + output = self.attend(q, k.repeat_interleave(self.groups, 1), + v.repeat_interleave(self.groups, 1), causal=True) + if need_state: + state.keys, state.values, state.position = k.clone(), v.clone(), t + return output, state + c = self.block_size + n = (t + c - 1) // c + pad = n * c - t + kb = F.pad(k, (0, 0, 0, pad)).reshape(b, self.num_kv_heads, n, c, d) + vb = F.pad(v, (0, 0, 0, pad)).reshape(b, self.num_kv_heads, n, c, d) + if self.has_memory: + phi_k = self.features(k) + if self.variant == "cenn_partition": + phi_k = phi_k * (torch.arange(t, device=q.device) >= self.sink_tokens)[None, None, :, None] + pk = F.pad(phi_k, (0, 0, 0, pad)).reshape( + b, self.num_kv_heads, n, c, self.feature_dim) + writes = self.mm(pk.transpose(-1, -2), vb) + masses = pk.sum(dim=-2) + delay = 1 if self.variant == "cenn_linear" else 2 + past_n, past_z = self.prefix(writes, delay), self.prefix(masses, delay) + if self.variant == "cenn_linear": + pq = F.pad(self.features(q, query=True), (0, 0, 0, pad)).reshape( + b, self.num_heads, n, c, self.feature_dim) + within = self.mm(pq, pk.repeat_interleave(self.groups, 1).transpose(-1, -2)).tril() + numerator = self.mm(pq, past_n.repeat_interleave(self.groups, 1)) + self.mm( + within, vb.repeat_interleave(self.groups, 1)) + denominator = self.mm( + pq, past_z.repeat_interleave(self.groups, 1).unsqueeze(-1) + ) + within.sum(-1, keepdim=True) + output = (numerator / denominator.clamp_min(1e-20)).reshape(b, self.num_heads, n * c, d)[:, :, :t] + if need_state: + state.numerator, state.denominator = writes.sum(2), masses.sum(2) + else: + qb = F.pad(q, (0, 0, 0, pad)).reshape(b, self.num_heads, n, c, d) + previous_k = torch.cat((torch.zeros_like(kb[:, :, :1]), kb[:, :, :-1]), dim=2) + previous_v = torch.cat((torch.zeros_like(vb[:, :, :1]), vb[:, :, :-1]), dim=2) + s = min(t, self.sink_tokens) + sinks_k = k[:, :, :s].unsqueeze(2).expand(b, self.num_kv_heads, n, s, d) + sinks_v = v[:, :, :s].unsqueeze(2).expand(b, self.num_kv_heads, n, s, d) + local_k = torch.cat((sinks_k, previous_k, kb), dim=3) + local_v = torch.cat((sinks_v, previous_v, vb), dim=3) + block = torch.arange(n, device=q.device)[:, None] + offset = torch.arange(c, device=q.device)[None] + positions = block * c + offset + local_positions = torch.cat(( + torch.arange(s, device=q.device)[None].expand(n, s), + positions - c, positions), dim=1) + query_positions = positions[:, :, None] + valid = (local_positions[:, None, :] <= query_positions) & ( + local_positions[:, None, :] >= 0) & (local_positions[:, None, :] < t) + # Every sink occurs only in the sink columns, never twice in exact local attention. + valid[:, :, s:] &= local_positions[:, None, s:] >= self.sink_tokens + output = self.combine(qb, local_k, local_v, valid[None, None], + past_n if self.has_memory else None, + past_z if self.has_memory else None) + output = output.reshape(b, self.num_heads, n * c, d)[:, :, :t] + if need_state: + if self.has_memory: + state.numerator, state.denominator = past_n[:, :, -1].clone(), past_z[:, :, -1].clone() + keep = min(t, c + (t - 1) % c + 1) + state.keys, state.values = k[:, :, -keep:].clone(), v[:, :, -keep:].clone() + state.sinks_k, state.sinks_v = k[:, :, :s].clone(), v[:, :, :s].clone() + if need_state: + state.position = t + return output, state + + def step(self, q, k, v, state): + """One token; the cache and compressed state are bounded in context length.""" + t, c = state.position, self.block_size + num, den = state.numerator, state.denominator + keys, values = state.keys, state.values + sinks_k, sinks_v = state.sinks_k, state.sinks_v + if self.variant == "transformer_readout": + keys, values = torch.cat((keys, k), 2), torch.cat((values, v), 2) + output = self.attend(q, keys.repeat_interleave(self.groups, 1), + values.repeat_interleave(self.groups, 1), causal=False) + elif self.variant == "cenn_linear": + pk = self.features(k) + num = num + self.mm(pk.transpose(-1, -2), v) + den = den + pk[:, :, 0] + qp = self.features(q, query=True) + output = self.mm(qp, num.repeat_interleave(self.groups, 1)) / self.mm( + qp, den.repeat_interleave(self.groups, 1).unsqueeze(-1)).clamp_min(1e-20) + else: + if t % c == 0 and keys.shape[2] > c: + retired = keys.shape[2] - c + if self.has_memory: + pk = self.features(keys[:, :, :retired]) + positions = torch.arange(t - keys.shape[2], t - c, device=q.device) + pk = pk * (positions >= self.sink_tokens)[None, None, :, None] + num = num + self.mm(pk.transpose(-1, -2), values[:, :, :retired]) + den = den + pk.sum(2) + keys, values = keys[:, :, -c:].clone(), values[:, :, -c:].clone() + keys, values = torch.cat((keys, k), 2), torch.cat((values, v), 2) + if t < self.sink_tokens: + sinks_k, sinks_v = torch.cat((sinks_k, k), 2), torch.cat((sinks_v, v), 2) + lk, lv = torch.cat((sinks_k, keys), 2), torch.cat((sinks_v, values), 2) + positions = torch.arange(t + 1 - keys.shape[2], t + 1, device=q.device) + valid = torch.cat((torch.ones(sinks_k.shape[2], dtype=torch.bool, device=q.device), + positions >= self.sink_tokens))[None, None, None] + output = self.combine(q, lk, lv, valid, num, den) + return output, MemoryState(num, den, keys, values, sinks_k, sinks_v, t + 1) + + def forward(self, q, k, v, state=None, return_state=False, apply_readout=True): + if (q.ndim != 4 or k.shape != v.shape or q.shape[0] != k.shape[0] + or q.shape[2:] != k.shape[2:] or q.shape[1] != self.num_heads + or k.shape[1] != self.num_kv_heads or q.shape[-1] != self.head_dim + or q.shape[2] < 1): + raise ValueError("Incompatible Q/K/V") + q, k, v = q.float(), k.float(), v.float() + if state is None: + output, state = self.prefill(q, k, v, need_state=return_state) + else: + outputs = [] + for i in range(q.shape[2]): + value, state = self.step(q[:, :, i:i+1], k[:, :, i:i+1], v[:, :, i:i+1], state) + outputs.append(value) + output = torch.cat(outputs, 2) + if apply_readout and self.variant != "sink_window": + output = self.calibrate(output) + return (output, state) if return_state else output + + +@torch.no_grad() +def ridge_calibrate(core, samples, device, relative_ridge=0.01): + """Identity-prior ridge solution, fitted only on training attention outputs. + + Solve (Y^T Y + lambda I) R = Y^T T + lambda I. Diagonal symmetric + preconditioning preserves the exact solution while improving conditioning. + CPU float64 is used for the small D-by-D systems. + """ + h, d = core.num_heads, core.head_dim + gram = torch.zeros(h, d, d, dtype=torch.float64) + cross = torch.zeros_like(gram) + before, count = 0.0, 0 + for q, k, v, target in samples: + y = core(q.to(device), k.to(device), v.to(device), apply_readout=False).cpu().double() + target = target.double() + y = y.permute(1, 0, 2, 3).reshape(h, -1, d) + target = target.permute(1, 0, 2, 3).reshape(h, -1, d) + gram += y.transpose(-1, -2) @ y + cross += y.transpose(-1, -2) @ target + before += (y - target).square().sum().item() + count += y.numel() + identity = torch.eye(d, dtype=torch.float64).expand(h, d, d) + ridge = relative_ridge * gram.diagonal(dim1=-2, dim2=-1).mean(-1).clamp_min(1e-8) + a, rhs = gram + ridge[:, None, None] * identity, cross + ridge[:, None, None] * identity + scale = a.diagonal(dim1=-2, dim2=-1).rsqrt() + conditioned = scale[:, :, None] * a * scale[:, None, :] + solution = scale[:, :, None] * torch.linalg.solve(conditioned, scale[:, :, None] * rhs) + core.readout.copy_(solution.to(core.readout)) + residual = (a @ solution - rhs).norm() / rhs.norm().clamp_min(1e-20) + return {"ridge_relative_residual": float(residual), "uncalibrated_train_mse": before / count} diff --git a/src/tinycenn_lm/pdelta2_er.py b/src/tinycenn_lm/pdelta2_er.py new file mode 100644 index 0000000000000000000000000000000000000000..3859cb5bcfd68c9404cf00debfd8fd4959710d7c --- /dev/null +++ b/src/tinycenn_lm/pdelta2_er.py @@ -0,0 +1,267 @@ +"""Error-residual PDelta2 layer for closing the remaining Transformer NLL gap. + +The layer keeps the strongest previous TinyCeNN direction (PDelta2 F96 + causal +Conv4) and tests three focused ingredients: + +1. learnable per-feature retention initialized to a broad half-life spectrum; +2. a small secondary recurrent state trained on the teacher-minus-base residual; +3. teacher-error-directed weighting that emphasizes the tokens on which the + replacement diverges most from exact Transformer attention. + +The residual memory is deliberately small and sees the original values while the +main memory sees Conv4 values. Its output gain starts at exactly zero, so adding +it is function preserving before training. Persistent memory matrices are stored +as FP16 between streaming calls while curvature remains FP32. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from tinycenn_lm.pdelta2_features import PDelta2Core, PDeltaState + + +@dataclass +class ErrorResidualState: + base: PDeltaState + residual: PDeltaState | None = None + conv_tail: Tensor | None = None + + +def initialize_retention_spectrum(core: PDelta2Core, minimum: float = 8.0, + maximum: float = 2048.0) -> None: + """Initialize feature channels with logarithmically spaced decay half-lives.""" + if minimum <= 0 or maximum <= minimum: + raise ValueError("retention half-life range must satisfy 0 < minimum < maximum") + with torch.no_grad(): + half_life = torch.exp(torch.linspace( + math.log(minimum), math.log(maximum), core.feature_dim, + device=core.forget_b.device, dtype=core.forget_b.dtype, + )) + log_decay = math.log(0.5) / half_life + probability = (-log_decay / 0.25).clamp(1e-5, 1.0 - 1e-5) + raw = torch.logit(probability) + core.forget_b.copy_(raw[None].expand(core.num_kv_heads, -1)) + core.forget_w.zero_() + + +def retention_half_lives(core: PDelta2Core) -> Tensor: + """Return the bias-only half-life represented by each KV-head/feature channel.""" + with torch.no_grad(): + log_decay = -0.25 * core.forget_b.sigmoid() + return math.log(0.5) / log_decay.clamp_max(-1e-7) + + +def teacher_error_weights(prediction: Tensor, target: Tensor, hard_fraction: float = 0.25, + hard_boost: float = 3.0) -> Tensor: + """Return normalized token weights that upweight the current hardest tokens. + + Error is averaged over heads and value dimensions, leaving [batch,time]. + The top ``hard_fraction`` tokens receive ``1 + hard_boost`` weight. + The returned weights have mean one so the loss scale stays comparable. + """ + if not 0.0 < hard_fraction < 1.0: + raise ValueError("hard_fraction must be in (0,1)") + if hard_boost < 0: + raise ValueError("hard_boost must be non-negative") + error = (prediction.detach() - target.detach()).square().mean(dim=(-1, 1)) + threshold = torch.quantile(error, 1.0 - hard_fraction, dim=-1, keepdim=True) + weights = 1.0 + hard_boost * (error >= threshold).to(error.dtype) + return weights / weights.mean(dim=-1, keepdim=True).clamp_min(1e-8) + + +def causal_value_conv_stream(v: Tensor, weight: Tensor | None, tail: Tensor | None = None): + """Causal depthwise value convolution with a correct streaming prefix tail.""" + if weight is None: + return v.float(), None + if weight.ndim != 3 or weight.shape[1] != 1: + raise ValueError("weight must be [channels,1,kernel]") + b, h, t, d = v.shape + kernel = weight.shape[-1] + channels = h * d + if weight.shape[0] != channels: + raise ValueError("convolution channel count does not match values") + if tail is None: + tail = v.new_zeros(b, h, kernel - 1, d) + if tail.shape != (b, h, kernel - 1, d): + raise ValueError("streaming convolution tail has an incompatible shape") + history = torch.cat((tail.to(v.dtype), v), dim=2) + x = history.transpose(1, 2).reshape(b, history.shape[2], channels).transpose(1, 2) + y = F.conv1d(x.float(), weight.float(), groups=channels) + y = y.transpose(1, 2).reshape(b, t, h, d).transpose(1, 2) + new_tail = history[:, :, -(kernel - 1):].clone() if kernel > 1 else None + return y, new_tail + + +class ErrorResidualPDelta2Layer(nn.Module): + """Conv4 PDelta2 with retention spectrum and a compact residual-error memory.""" + + def __init__(self, num_heads: int, num_kv_heads: int, head_dim: int, + feature_dim: int = 96, residual_dim: int = 0, chunk_size: int = 32, + conv_kernel: int = 4, retention_spectrum: bool = False, + retention_min: float = 8.0, retention_max: float = 2048.0, + state_dtype: str = "fp16"): + super().__init__() + if num_heads % num_kv_heads: + raise ValueError("num_heads must be divisible by num_kv_heads") + if state_dtype not in {"fp16", "fp32"}: + raise ValueError("state_dtype must be fp16 or fp32") + if conv_kernel < 1 or residual_dim < 0: + raise ValueError("conv_kernel must be positive and residual_dim non-negative") + self.num_heads = int(num_heads) + self.num_kv_heads = int(num_kv_heads) + self.head_dim = int(head_dim) + self.feature_dim = int(feature_dim) + self.residual_dim = int(residual_dim) + self.chunk_size = int(chunk_size) + self.conv_kernel = int(conv_kernel) + self.retention_spectrum = bool(retention_spectrum) + self.retention_min = float(retention_min) + self.retention_max = float(retention_max) + self.state_dtype = state_dtype + self.groups = self.num_heads // self.num_kv_heads + + self.base = PDelta2Core( + self.num_heads, self.num_kv_heads, self.head_dim, + feature_dim=self.feature_dim, chunk_size=self.chunk_size, + ) + if self.retention_spectrum: + initialize_retention_spectrum(self.base, self.retention_min, self.retention_max) + + if self.conv_kernel > 1: + channels = self.num_kv_heads * self.head_dim + kernel = torch.zeros(channels, 1, self.conv_kernel) + kernel[:, 0, -1] = 1.0 + self.conv_weight = nn.Parameter(kernel) + else: + self.register_parameter("conv_weight", None) + + if self.residual_dim: + self.residual = PDelta2Core( + self.num_heads, self.num_kv_heads, self.head_dim, + feature_dim=self.residual_dim, chunk_size=self.chunk_size, + ) + if self.retention_spectrum: + initialize_retention_spectrum( + self.residual, + max(4.0, self.retention_min / 2.0), + self.retention_max * 2.0, + ) + # Zero makes the whole residual branch exactly inactive at initialization. + self.residual_gain = nn.Parameter(torch.zeros(self.num_heads)) + else: + self.residual = None + self.register_parameter("residual_gain", None) + + @property + def config(self): + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "residual_dim": self.residual_dim, + "chunk_size": self.chunk_size, + "conv_kernel": self.conv_kernel, + "retention_spectrum": self.retention_spectrum, + "retention_min": self.retention_min, + "retention_max": self.retention_max, + "state_dtype": self.state_dtype, + } + + @property + def storage_dtype(self): + return torch.float16 if self.state_dtype == "fp16" else torch.float32 + + def _working_state(self, state: PDeltaState | None, core: PDelta2Core): + if state is None: + return None + dtype = core.wq.dtype + return PDeltaState(state.memory.to(dtype), state.curvature.to(dtype)) + + def _stored_state(self, state: PDeltaState): + return PDeltaState( + state.memory.to(self.storage_dtype), + state.curvature.float(), + ) + + def _run(self, q: Tensor, k: Tensor, v: Tensor, state: ErrorResidualState | None): + tail = None if state is None else state.conv_tail + conv_v, new_tail = causal_value_conv_stream(v.float(), self.conv_weight, tail) + base_state = None if state is None else state.base + base_out, new_base = self.base( + q, k, conv_v, + state=self._working_state(base_state, self.base), + return_state=True, + ) + + residual_raw = None + new_residual = None + output = base_out + if self.residual is not None: + residual_state = None if state is None else state.residual + residual_raw, residual_state_out = self.residual( + q, k, v.float(), + state=self._working_state(residual_state, self.residual), + return_state=True, + ) + gain = self.residual_gain.clamp(-1.5, 1.5)[None, :, None, None] + output = base_out + gain * residual_raw + new_residual = self._stored_state(residual_state_out) + + new_state = ErrorResidualState( + base=self._stored_state(new_base), + residual=new_residual, + conv_tail=(None if new_tail is None else new_tail.to(self.storage_dtype)), + ) + components = { + "base": base_out, + "residual_raw": residual_raw, + "output": output, + } + return output, new_state, components + + def components(self, q: Tensor, k: Tensor, v: Tensor, + state: ErrorResidualState | None = None): + return self._run(q, k, v, state) + + def forward(self, q: Tensor, k: Tensor, v: Tensor, + state: ErrorResidualState | None = None, + return_state: bool = False, implementation: str = "chunk"): + if implementation != "chunk": + raise ValueError("ErrorResidualPDelta2Layer supports the chunk implementation") + output, new_state, _ = self._run(q, k, v, state) + return (output, new_state) if return_state else output + + def recurrent_state_bytes(self, batch_size: int = 1, context: int | None = None): + del context + memory_bytes = 2 if self.state_dtype == "fp16" else 4 + base_memory = self.num_kv_heads * self.feature_dim * self.head_dim * memory_bytes + base_curvature = self.num_kv_heads * self.feature_dim * 4 + total = base_memory + base_curvature + if self.residual_dim: + total += self.num_kv_heads * self.residual_dim * self.head_dim * memory_bytes + total += self.num_kv_heads * self.residual_dim * 4 + if self.conv_kernel > 1: + total += (self.conv_kernel - 1) * self.num_kv_heads * self.head_dim * memory_bytes + return batch_size * total + + def retention_statistics(self): + values = retention_half_lives(self.base).float().reshape(-1) + result = { + "base_half_life_min": float(values.min()), + "base_half_life_median": float(values.median()), + "base_half_life_max": float(values.max()), + } + if self.residual is not None: + rv = retention_half_lives(self.residual).float().reshape(-1) + result.update( + residual_half_life_min=float(rv.min()), + residual_half_life_median=float(rv.median()), + residual_half_life_max=float(rv.max()), + ) + return result diff --git a/src/tinycenn_lm/pdelta2_er2.py b/src/tinycenn_lm/pdelta2_er2.py new file mode 100644 index 0000000000000000000000000000000000000000..4ebdcbeac2a99eeaa481528e7ed36e8172d6e0b6 --- /dev/null +++ b/src/tinycenn_lm/pdelta2_er2.py @@ -0,0 +1,391 @@ +"""PDelta2-ER2: selective compressed error memory for long-context replacement. + +The layer keeps the strongest Conv4 + PDelta2 F96 core and focuses the secondary +Residual16 state on a low-rank teacher-error subspace. A learned query gate +predicts which head/token positions need the correction, so the residual branch +is not trained to imitate every attention output equally. + +This is a research reference implementation. Persistent recurrent matrices may +be stored in FP16 between streaming calls while curvature remains FP32. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from tinycenn_lm.pdelta2_er import ( + causal_value_conv_stream, + initialize_retention_spectrum, + retention_half_lives, +) +from tinycenn_lm.pdelta2_features import PDelta2Core, PDeltaState + + +@dataclass +class ER2State: + base: PDeltaState + residual: PDeltaState | None = None + conv_tail: Tensor | None = None + + +def hard_error_mask(base: Tensor, target: Tensor, fraction: float = 0.25) -> tuple[Tensor, Tensor]: + """Return per-head hard mask and squared error, both [B,H,T].""" + if not 0.0 < fraction < 1.0: + raise ValueError("fraction must be in (0,1)") + error = (target.detach() - base.detach()).square().mean(dim=-1) + threshold = torch.quantile(error, 1.0 - fraction, dim=-1, keepdim=True) + return error >= threshold, error + + +def normalized_hard_weights(mask: Tensor, boost: float = 3.0) -> Tensor: + """Weights with mean one, preserving the overall loss scale.""" + if boost < 0: + raise ValueError("boost must be non-negative") + weights = 1.0 + boost * mask.to(torch.float32) + return weights / weights.mean(dim=(-1, -2), keepdim=True).clamp_min(1e-8) + + +class SelectiveCompressedPDelta2Layer(nn.Module): + """Conv4 PDelta2 with optional retention and selective low-rank Residual16. + + ``residual_mode="compressed"`` constrains the correction to a learned + per-query-head rank-R subspace. ``residual_mode="raw"`` is included as a + same-budget control. + """ + + def __init__( + self, + num_heads: int, + num_kv_heads: int, + head_dim: int, + feature_dim: int = 96, + residual_dim: int = 16, + code_rank: int = 8, + chunk_size: int = 32, + conv_kernel: int = 4, + retention_spectrum: bool = True, + retention_min: float = 8.0, + retention_max: float = 4096.0, + residual_mode: str = "compressed", + state_dtype: str = "fp16", + ): + super().__init__() + if num_heads % num_kv_heads: + raise ValueError("num_heads must be divisible by num_kv_heads") + if residual_mode not in {"none", "raw", "compressed"}: + raise ValueError("residual_mode must be none, raw, or compressed") + if state_dtype not in {"fp16", "fp32"}: + raise ValueError("state_dtype must be fp16 or fp32") + if conv_kernel < 1 or residual_dim < 0: + raise ValueError("conv_kernel must be positive and residual_dim non-negative") + if residual_mode == "compressed" and not 1 <= code_rank <= head_dim: + raise ValueError("compressed mode needs 1 <= code_rank <= head_dim") + if residual_mode == "none": + residual_dim = 0 + code_rank = 0 + + self.num_heads = int(num_heads) + self.num_kv_heads = int(num_kv_heads) + self.head_dim = int(head_dim) + self.feature_dim = int(feature_dim) + self.residual_dim = int(residual_dim) + self.code_rank = int(code_rank) + self.chunk_size = int(chunk_size) + self.conv_kernel = int(conv_kernel) + self.retention_spectrum = bool(retention_spectrum) + self.retention_min = float(retention_min) + self.retention_max = float(retention_max) + self.residual_mode = residual_mode + self.state_dtype = state_dtype + self.groups = self.num_heads // self.num_kv_heads + + self.base = PDelta2Core( + self.num_heads, + self.num_kv_heads, + self.head_dim, + feature_dim=self.feature_dim, + chunk_size=self.chunk_size, + ) + if self.retention_spectrum: + initialize_retention_spectrum(self.base, self.retention_min, self.retention_max) + + if self.conv_kernel > 1: + channels = self.num_kv_heads * self.head_dim + kernel = torch.zeros(channels, 1, self.conv_kernel) + kernel[:, 0, -1] = 1.0 + self.conv_weight = nn.Parameter(kernel) + else: + self.register_parameter("conv_weight", None) + + if self.residual_dim: + self.residual = PDelta2Core( + self.num_heads, + self.num_kv_heads, + self.head_dim, + feature_dim=self.residual_dim, + chunk_size=self.chunk_size, + ) + if self.retention_spectrum: + initialize_retention_spectrum( + self.residual, + max(4.0, self.retention_min / 2.0), + self.retention_max * 2.0, + ) + self.residual_gain = nn.Parameter(torch.zeros(self.num_heads)) + else: + self.residual = None + self.register_parameter("residual_gain", None) + + if self.residual_mode == "compressed": + basis = torch.empty(self.num_heads, self.code_rank, self.head_dim) + for head in range(self.num_heads): + nn.init.orthogonal_(basis[head]) + self.residual_basis = nn.Parameter(basis) + self.error_gate_w = nn.Parameter(torch.zeros(self.num_heads, self.head_dim)) + self.error_gate_b = nn.Parameter(torch.full((self.num_heads,), math.log(0.2 / 0.8))) + else: + self.register_parameter("residual_basis", None) + self.register_parameter("error_gate_w", None) + self.register_parameter("error_gate_b", None) + + @property + def config(self): + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "residual_dim": self.residual_dim, + "code_rank": self.code_rank, + "chunk_size": self.chunk_size, + "conv_kernel": self.conv_kernel, + "retention_spectrum": self.retention_spectrum, + "retention_min": self.retention_min, + "retention_max": self.retention_max, + "residual_mode": self.residual_mode, + "state_dtype": self.state_dtype, + } + + @property + def storage_dtype(self): + return torch.float16 if self.state_dtype == "fp16" else torch.float32 + + def normalized_basis(self) -> Tensor | None: + if self.residual_basis is None: + return None + return F.normalize(self.residual_basis.float(), dim=-1) + + def project_to_code(self, x: Tensor) -> Tensor: + basis = self.normalized_basis() + if basis is None: + raise RuntimeError("code projection requires compressed residual mode") + return torch.einsum("bhtd,hrd->bhtr", x.float(), basis) + + def reconstruct_code(self, code: Tensor) -> Tensor: + basis = self.normalized_basis() + if basis is None: + raise RuntimeError("code reconstruction requires compressed residual mode") + return torch.einsum("bhtr,hrd->bhtd", code.float(), basis) + + def project_vector(self, x: Tensor) -> Tensor: + return self.reconstruct_code(self.project_to_code(x)) + + def orthogonality_penalty(self) -> Tensor: + basis = self.normalized_basis() + if basis is None: + return torch.zeros((), device=self.base.wq.device) + gram = torch.einsum("hrd,hsd->hrs", basis, basis) + eye = torch.eye(self.code_rank, device=gram.device, dtype=gram.dtype)[None] + return (gram - eye).square().mean() + + def predicted_error_gate(self, q: Tensor) -> Tensor: + if self.residual_mode != "compressed": + return q.new_ones(q.shape[0], q.shape[1], q.shape[2]) + qn = F.normalize(q.float(), dim=-1) + return ( + torch.einsum("bhtd,hd->bht", qn, self.error_gate_w.float()) + + self.error_gate_b.float()[None, :, None] + ).sigmoid() + + def _working_state(self, state: PDeltaState | None, core: PDelta2Core): + if state is None: + return None + dtype = core.wq.dtype + return PDeltaState(state.memory.to(dtype), state.curvature.to(dtype)) + + def _stored_state(self, state: PDeltaState): + return PDeltaState(state.memory.to(self.storage_dtype), state.curvature.float()) + + def _run(self, q: Tensor, k: Tensor, v: Tensor, state: ER2State | None): + tail = None if state is None else state.conv_tail + conv_v, new_tail = causal_value_conv_stream(v.float(), self.conv_weight, tail) + + base_state = None if state is None else state.base + base_out, new_base = self.base( + q, k, conv_v, + state=self._working_state(base_state, self.base), + return_state=True, + ) + + residual_raw = None + residual_projected = None + gate = None + new_residual = None + output = base_out + + if self.residual is not None: + residual_state = None if state is None else state.residual + residual_raw, residual_state_out = self.residual( + q, k, v.float(), + state=self._working_state(residual_state, self.residual), + return_state=True, + ) + new_residual = self._stored_state(residual_state_out) + gain = self.residual_gain.clamp(-1.5, 1.5)[None, :, None, None] + + if self.residual_mode == "compressed": + residual_projected = self.project_vector(residual_raw) + gate = self.predicted_error_gate(q).unsqueeze(-1) + correction = gain * gate * residual_projected + elif self.residual_mode == "raw": + residual_projected = residual_raw + gate = residual_raw.new_ones(residual_raw.shape[:-1] + (1,)) + correction = gain * residual_raw + else: + correction = 0.0 + output = base_out + correction + + new_state = ER2State( + base=self._stored_state(new_base), + residual=new_residual, + conv_tail=None if new_tail is None else new_tail.to(self.storage_dtype), + ) + return output, new_state, { + "base": base_out, + "residual_raw": residual_raw, + "residual_projected": residual_projected, + "gate": gate, + "output": output, + } + + def components(self, q: Tensor, k: Tensor, v: Tensor, state: ER2State | None = None): + return self._run(q, k, v, state) + + def forward( + self, + q: Tensor, + k: Tensor, + v: Tensor, + state: ER2State | None = None, + return_state: bool = False, + implementation: str = "chunk", + ): + if implementation != "chunk": + raise ValueError("SelectiveCompressedPDelta2Layer supports chunk implementation") + output, new_state, _ = self._run(q, k, v, state) + return (output, new_state) if return_state else output + + def auxiliary_losses( + self, + q: Tensor, + k: Tensor, + v: Tensor, + teacher: Tensor, + hard_fraction: float = 0.25, + hard_boost: float = 3.0, + ): + """Teacher-error-directed losses used only during training.""" + output, _, parts = self._run(q, k, v, None) + base = parts["base"] + mask, error = hard_error_mask(base, teacher, hard_fraction) + weights = normalized_hard_weights(mask, hard_boost).to(output.device) + weights4 = weights.unsqueeze(-1) + + teacher_den = (teacher.detach().square() * weights4).mean().clamp_min(1e-8) + hard_teacher_nmse = ((output - teacher.detach()).square() * weights4).mean() / teacher_den + result = { + "hard_teacher_nmse": hard_teacher_nmse, + "hard_fraction_observed": mask.float().mean(), + "teacher_error_mean": error.mean(), + "hard_error_mean": error[mask].mean(), + "easy_error_mean": error[~mask].mean(), + } + + if self.residual is None: + zero = hard_teacher_nmse.new_zeros(()) + result.update( + residual_code_nmse=zero, + residual_reconstruction_nmse=zero, + gate_bce=zero, + orthogonality=zero, + ) + return result + + residual_target = teacher.detach() - base.detach() + hard4 = mask.unsqueeze(-1).to(residual_target.dtype) + + if self.residual_mode == "compressed": + target_code = self.project_to_code(residual_target) + predicted_code = self.project_to_code(parts["residual_raw"]) + hard_code = mask.unsqueeze(-1).to(target_code.dtype) + code_den = (target_code.square() * hard_code).sum().clamp_min(1e-8) + code_nmse = ((predicted_code - target_code).square() * hard_code).sum() / code_den + + target_projection = self.reconstruct_code(target_code) + recon_den = (residual_target.square() * hard4).sum().clamp_min(1e-8) + recon_nmse = ((target_projection - residual_target).square() * hard4).sum() / recon_den + + gate = parts["gate"].squeeze(-1).clamp(1e-5, 1 - 1e-5) + gate_bce = F.binary_cross_entropy(gate, mask.to(gate.dtype)) + orth = self.orthogonality_penalty() + else: + pred = parts["residual_raw"] + den = (residual_target.square() * hard4).sum().clamp_min(1e-8) + code_nmse = ((pred - residual_target).square() * hard4).sum() / den + recon_nmse = code_nmse.new_zeros(()) + gate_bce = code_nmse.new_zeros(()) + orth = code_nmse.new_zeros(()) + + result.update( + residual_code_nmse=code_nmse, + residual_reconstruction_nmse=recon_nmse, + gate_bce=gate_bce, + orthogonality=orth, + ) + return result + + def recurrent_state_bytes(self, batch_size: int = 1, context: int | None = None): + del context + memory_bytes = 2 if self.state_dtype == "fp16" else 4 + total = self.num_kv_heads * self.feature_dim * self.head_dim * memory_bytes + total += self.num_kv_heads * self.feature_dim * 4 + if self.residual_dim: + total += self.num_kv_heads * self.residual_dim * self.head_dim * memory_bytes + total += self.num_kv_heads * self.residual_dim * 4 + if self.conv_kernel > 1: + total += (self.conv_kernel - 1) * self.num_kv_heads * self.head_dim * memory_bytes + return batch_size * total + + def retention_statistics(self): + values = retention_half_lives(self.base).float().reshape(-1) + result = { + "base_half_life_min": float(values.min()), + "base_half_life_median": float(values.median()), + "base_half_life_max": float(values.max()), + } + if self.residual is not None: + rv = retention_half_lives(self.residual).float().reshape(-1) + result.update( + residual_half_life_min=float(rv.min()), + residual_half_life_median=float(rv.median()), + residual_half_life_max=float(rv.max()), + ) + if self.residual_gain is not None: + result["residual_gain_abs_mean"] = float(self.residual_gain.detach().abs().mean()) + if self.residual_mode == "compressed": + result["predicted_gate_mean"] = float(self.error_gate_b.detach().sigmoid().mean()) + return result diff --git a/src/tinycenn_lm/pdelta2_features.py b/src/tinycenn_lm/pdelta2_features.py new file mode 100644 index 0000000000000000000000000000000000000000..c11a4136dbee2e29ba2d3b7f2470943e6a835f46 --- /dev/null +++ b/src/tinycenn_lm/pdelta2_features.py @@ -0,0 +1,294 @@ +"""Feature-lab components for faster, stronger P-Delta2 attention replacements. + +This module keeps the recurrent state bounded while testing three ingredients: +1) a chunk-vectorized curvature preconditioner (same recurrence as serial P-Delta2), +2) sparse dilated exact retrieval over logarithmic offsets, and +3) dual-timescale recurrent memories with query-dependent mixing. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Iterable + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from tinycenn_lm.research_layers import delta_recurrence, local_window_attention + + +@dataclass +class PDeltaState: + memory: Tensor + curvature: Tensor + + +@dataclass +class FeatureState: + fast: PDeltaState + slow: PDeltaState | None = None + keys: Tensor | None = None + values: Tensor | None = None + + +def precondition_reference(kp: Tensor, curvature: Tensor, alpha: Tensor, beta: Tensor, + log_x: Tensor, center: Tensor): + """Tokenwise oracle for the diagonal curvature preconditioner.""" + writes = [] + for t in range(kp.shape[2]): + kt = kp[:, :, t] + r = (curvature + 1e-4).log() - center[None] + s = r / (1.0 + r.abs()) + scale = torch.exp(-log_x[None] * s) + numerator = scale * kt + denominator = 1.0 + (kt * numerator).sum(-1, keepdim=True) + writes.append(numerator / denominator.clamp_min(1e-4)) + curvature = alpha[None] * curvature + beta[None] * kt.square() + return torch.stack(writes, dim=2), curvature + + +def precondition_chunked(kp: Tensor, curvature: Tensor, alpha: Tensor, beta: Tensor, + log_x: Tensor, center: Tensor, chunk_size: int = 32): + """Vectorize curvature states inside bounded chunks; recurrent only across chunks.""" + writes = [] + for start in range(0, kp.shape[2], chunk_size): + kc = kp[:, :, start:start + chunk_size] + length = kc.shape[2] + squared = kc.square() + t = torch.arange(length, device=kp.device) + j = torch.arange(length, device=kp.device) + lag = t[:, None] - 1 - j[None, :] + valid = lag >= 0 + weights = alpha[:, :, None, None].pow(lag.clamp_min(0)[None, None]) + weights = weights * valid[None, None] + contribution = torch.einsum("bhjf,hftj->bhtf", squared, weights) + contribution = contribution * beta[None, :, None, :] + powers = alpha[:, :, None].pow(t[None, None]) + before = curvature[:, :, None, :] * powers.permute(0, 2, 1)[None] + contribution + + r = (before + 1e-4).log() - center[None, :, None, :] + s = r / (1.0 + r.abs()) + scale = torch.exp(-log_x[None, :, None, :] * s) + numerator = scale * kc + denominator = 1.0 + (kc * numerator).sum(-1, keepdim=True) + writes.append(numerator / denominator.clamp_min(1e-4)) + + curvature = alpha[None] * before[:, :, -1] + beta[None] * squared[:, :, -1] + return torch.cat(writes, dim=2), curvature + + +class PDelta2Core(nn.Module): + """P-Delta2 recurrence with a chunk-vectorized diagonal preconditioner.""" + + def __init__(self, num_heads: int, num_kv_heads: int, head_dim: int, + feature_dim: int = 96, chunk_size: int = 32, + forget_bias: float | None = None): + super().__init__() + if num_heads % num_kv_heads: + raise ValueError("num_heads must be divisible by num_kv_heads") + if min(num_heads, num_kv_heads, head_dim, feature_dim) < 1: + raise ValueError("dimensions must be positive") + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.feature_dim = feature_dim + self.groups = num_heads // num_kv_heads + self.chunk_size = chunk_size + + base = torch.zeros(num_kv_heads, feature_dim, head_dim) + for head in range(num_kv_heads): + if feature_dim == head_dim: + base[head] = torch.eye(head_dim) + else: + nn.init.orthogonal_(base[head]) + self.wk = nn.Parameter(base.clone()) + self.wq = nn.Parameter(base.repeat_interleave(self.groups, dim=0).clone()) + self.forget_w = nn.Parameter(torch.zeros(num_kv_heads, feature_dim, head_dim)) + default_bias = math.log(0.04 / 0.96) if forget_bias is None else forget_bias + self.forget_b = nn.Parameter(torch.full((num_kv_heads, feature_dim), default_bias)) + self.erase_w = nn.Parameter(torch.zeros(num_kv_heads, feature_dim, head_dim)) + self.erase_b = nn.Parameter(torch.full((num_kv_heads, feature_dim), -1.0)) + self.write_w = nn.Parameter(torch.zeros(num_kv_heads, head_dim, head_dim)) + self.write_b = nn.Parameter(torch.full((num_kv_heads, head_dim), -1.0)) + self.pre_log_decay = nn.Parameter(torch.full((num_kv_heads, feature_dim), math.log(0.995))) + self.pre_gain_logit = nn.Parameter(torch.full((num_kv_heads, feature_dim), math.log(0.12 / 0.88))) + self.pre_range_raw = nn.Parameter(torch.zeros(num_kv_heads, 1)) + self.pre_center = nn.Parameter(torch.zeros(num_kv_heads, 1)) + self.log_gain = nn.Parameter(torch.zeros(num_heads)) + + @staticmethod + def _project(x, weight, bias): + return torch.einsum("bhtd,hfd->bhtf", x, weight) + bias[None, :, None] + + def features(self, q, k, v): + qn, kn = F.normalize(q, dim=-1), F.normalize(k, dim=-1) + qp = F.normalize(torch.einsum("bhtd,hfd->bhtf", qn, self.wq), dim=-1) + kp = F.normalize(torch.einsum("bhtd,hfd->bhtf", kn, self.wk), dim=-1) + log_decay = -0.25 * self._project(kn, self.forget_w, self.forget_b).sigmoid() + erase = kp * self._project(kn, self.erase_w, self.erase_b).sigmoid() + write_gate = self._project(F.normalize(v, dim=-1), self.write_w, self.write_b).sigmoid() + return qp, kp, v * write_gate, erase, log_decay + + def precondition_parameters(self): + alpha = self.pre_log_decay.clamp(math.log(0.98), math.log(0.9999)).exp() + beta = self.pre_gain_logit.sigmoid() + log_x = math.log(2.0) + self.pre_range_raw.sigmoid() * (math.log(8.0) - math.log(2.0)) + return alpha, beta, log_x, self.pre_center + + def precondition_keys(self, kp, curvature): + return precondition_chunked(kp, curvature, *self.precondition_parameters(), self.chunk_size) + + def precondition_keys_reference(self, kp, curvature): + return precondition_reference(kp, curvature, *self.precondition_parameters()) + + def forward(self, q, k, v, state: PDeltaState | None = None, return_state: bool = False): + q, k, v = (x.to(self.wq.dtype) for x in (q, k, v)) + if state is None: + state = PDeltaState( + memory=q.new_zeros(q.shape[0], self.num_kv_heads, self.feature_dim, self.head_dim), + curvature=q.new_ones(q.shape[0], self.num_kv_heads, self.feature_dim), + ) + qp, kp, z, erase, log_decay = self.features(q, k, v) + kpre, curvature = self.precondition_keys(kp, state.curvature) + output, memory = delta_recurrence( + qp, kpre, z, erase, log_decay, state.memory, self.groups, self.chunk_size + ) + output = output * self.log_gain.clamp(-4, 4).exp()[None, :, None, None] + new_state = PDeltaState(memory, curvature) + return (output, new_state) if return_state else output + + def recurrent_state_bytes(self, batch_size: int = 1): + elements = self.num_kv_heads * self.feature_dim * self.head_dim + elements += self.num_kv_heads * self.feature_dim + return batch_size * elements * self.wq.element_size() + + +def dilated_sparse_attention(q: Tensor, k: Tensor, v: Tensor, offsets: Iterable[int], groups: int): + """Exact causal softmax over a fixed set of logarithmic past offsets.""" + offsets = tuple(sorted(set(int(x) for x in offsets))) + if not offsets or offsets[0] != 0 or min(offsets) < 0: + raise ValueError("offsets must be non-negative and include 0") + k = k.repeat_interleave(groups, dim=1) + v = v.repeat_interleave(groups, dim=1) + total, length = k.shape[2], q.shape[2] + prefix = total - length + positions = prefix + torch.arange(length, device=q.device) + off = torch.tensor(offsets, device=q.device) + index = positions[:, None] - off[None, :] + valid = index >= 0 + index = index.clamp_min(0) + selected_k = k[:, :, index, :] + selected_v = v[:, :, index, :] + scores = (q.unsqueeze(-2) * selected_k).sum(-1) / math.sqrt(q.shape[-1]) + scores = scores.masked_fill(~valid[None, None], float("-inf")) + weights = scores.softmax(-1) + return (weights.unsqueeze(-1) * selected_v).sum(-2) + + +class FeaturePDelta2Layer(nn.Module): + """Composable P-Delta2 experiment with dense/dilated retrieval and optional dual time scales.""" + + def __init__(self, num_heads: int, num_kv_heads: int, head_dim: int, + feature_dim: int = 96, retrieval: str = "dense", window: int = 32, + offsets: Iterable[int] = (0, 1, 2, 4, 8, 16, 32), + dual_timescale: bool = False, chunk_size: int = 32): + super().__init__() + if retrieval not in {"none", "dense", "dilated"}: + raise ValueError("retrieval must be none, dense, or dilated") + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.feature_dim = feature_dim + self.groups = num_heads // num_kv_heads + self.retrieval = retrieval + self.window = window + self.offsets = tuple(int(x) for x in offsets) + self.dual_timescale = bool(dual_timescale) + self.chunk_size = chunk_size + + if dual_timescale: + fast_dim = feature_dim // 2 + slow_dim = feature_dim - fast_dim + self.fast = PDelta2Core(num_heads, num_kv_heads, head_dim, fast_dim, chunk_size, -1.5) + self.slow = PDelta2Core(num_heads, num_kv_heads, head_dim, slow_dim, chunk_size, -5.0) + self.timescale_w = nn.Parameter(torch.zeros(num_heads, head_dim)) + self.timescale_b = nn.Parameter(torch.zeros(num_heads)) + else: + self.fast = PDelta2Core(num_heads, num_kv_heads, head_dim, feature_dim, chunk_size) + self.slow = None + if retrieval != "none": + self.retrieval_w = nn.Parameter(torch.zeros(num_heads, head_dim)) + self.retrieval_b = nn.Parameter(torch.full((num_heads,), -0.5)) + + @property + def config(self): + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "retrieval": self.retrieval, + "window": self.window, + "offsets": list(self.offsets), + "dual_timescale": self.dual_timescale, + "chunk_size": self.chunk_size, + } + + def _retrieval_keep(self): + if self.retrieval == "dense": + return max(0, self.window - 1) + if self.retrieval == "dilated": + return max(self.offsets) + return 0 + + def forward(self, q, k, v, state: FeatureState | None = None, return_state: bool = False, + implementation: str = "chunk"): + if implementation != "chunk": + raise ValueError("feature lab supports the chunk implementation") + fast_state = None if state is None else state.fast + fast_out, new_fast = self.fast(q, k, v, fast_state, return_state=True) + recurrent = fast_out + new_slow = None + if self.slow is not None: + slow_state = None if state is None else state.slow + slow_out, new_slow = self.slow(q, k, v, slow_state, return_state=True) + gate = ( + torch.einsum("bhtd,hd->bht", F.normalize(q.float(), dim=-1), self.timescale_w) + + self.timescale_b[None, :, None] + ).sigmoid().unsqueeze(-1) + recurrent = gate * fast_out + (1.0 - gate) * slow_out + + new_state = FeatureState(new_fast, new_slow) + if self.retrieval != "none": + keys = k.float() if state is None or state.keys is None else torch.cat((state.keys, k.float()), dim=2) + values = v.float() if state is None or state.values is None else torch.cat((state.values, v.float()), dim=2) + if self.retrieval == "dense": + retrieved = local_window_attention(q.float(), keys, values, self.window, self.groups) + else: + retrieved = dilated_sparse_attention(q.float(), keys, values, self.offsets, self.groups) + mix = ( + torch.einsum("bhtd,hd->bht", F.normalize(q.float(), dim=-1), self.retrieval_w) + + self.retrieval_b[None, :, None] + ).sigmoid().unsqueeze(-1) + recurrent = mix * retrieved + (1.0 - mix) * recurrent + keep = self._retrieval_keep() + new_state.keys = keys[:, :, -keep:].clone() if keep else None + new_state.values = values[:, :, -keep:].clone() if keep else None + + return (recurrent, new_state) if return_state else recurrent + + def recurrent_state_bytes(self, batch_size: int = 1): + total = self.fast.recurrent_state_bytes(batch_size) + if self.slow is not None: + total += self.slow.recurrent_state_bytes(batch_size) + keep = self._retrieval_keep() + total += batch_size * 2 * keep * self.num_kv_heads * self.head_dim * self.fast.wq.element_size() + return total + + def retrieval_pairs_per_token(self): + if self.retrieval == "none": + return 0 + if self.retrieval == "dense": + return self.window + return len(self.offsets) diff --git a/src/tinycenn_lm/pdelta2_flash.py b/src/tinycenn_lm/pdelta2_flash.py new file mode 100644 index 0000000000000000000000000000000000000000..f9594295eb7d28e7d5ebcbe33f183ccb6dc77b70 --- /dev/null +++ b/src/tinycenn_lm/pdelta2_flash.py @@ -0,0 +1,325 @@ +"""PDelta2-Flash research ingredients for a stronger/faster single-layer replacement. + +The module keeps the proven P-Delta2 recurrent core and tests lightweight ideas +suggested by recent hybrid efficient-attention models: +- function-preserving per-head output gating; +- a short causal value convolution (kernel 4 by default); +- compact content-indexed block summaries for long-range recall; +- FP16 persistent recurrent-memory storage with FP32 curvature. + +The indexed path stores one K/V summary per completed block, not every token. +It is therefore a compact growing memory, while the P-Delta2 recurrent state +remains bounded. This is an independent TinyCeNN experiment. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from tinycenn_lm.pdelta2_features import PDelta2Core, PDeltaState + + +@dataclass +class FlashState: + recurrent: PDeltaState + conv_tail: Tensor | None = None + block_keys: Tensor | None = None + block_values: Tensor | None = None + pending_keys: Tensor | None = None + pending_values: Tensor | None = None + + +def causal_depthwise_value_conv(v: Tensor, weight: Tensor) -> Tensor: + """Causal depthwise 1-D convolution over KV-head/value channels.""" + if weight.ndim != 3 or weight.shape[1] != 1: + raise ValueError("weight must be [channels, 1, kernel]") + b, h, t, d = v.shape + channels = h * d + if weight.shape[0] != channels: + raise ValueError("weight channel count does not match values") + x = v.transpose(1, 2).reshape(b, t, channels).transpose(1, 2) + kernel = weight.shape[-1] + x = F.pad(x, (kernel - 1, 0)) + y = F.conv1d(x, weight, groups=channels) + return y.transpose(1, 2).reshape(b, t, h, d).transpose(1, 2) + + +def indexed_block_attention(q: Tensor, k: Tensor, v: Tensor, groups: int, + block_size: int = 16, topk: int = 4): + """Attend to compact summaries of completed causal blocks. + + A query at token ``t`` may only see blocks whose final token is < ``t``. + Each block contributes one mean-key and one mean-value summary, reducing + long-range storage by roughly ``block_size`` versus tokenwise KV storage. + """ + if block_size < 2 or topk < 1: + raise ValueError("block_size must be >=2 and topk positive") + b, h, t, d = q.shape + k = k.repeat_interleave(groups, dim=1) + v = v.repeat_interleave(groups, dim=1) + blocks = t // block_size + if blocks == 0: + return q.new_zeros(q.shape), torch.zeros((b, h, t, 1), dtype=torch.bool, device=q.device) + + usable = blocks * block_size + bk = k[:, :, :usable].reshape(b, h, blocks, block_size, d).mean(dim=3) + bv = v[:, :, :usable].reshape(b, h, blocks, block_size, d).mean(dim=3) + bk = F.normalize(bk, dim=-1) + qn = F.normalize(q, dim=-1) + scores = torch.einsum("bhtd,bhnd->bhtn", qn, bk) / math.sqrt(d) + + positions = torch.arange(t, device=q.device) + block_ends = torch.arange(blocks, device=q.device) * block_size + (block_size - 1) + valid = block_ends[None, :] < positions[:, None] + scores = scores.masked_fill(~valid[None, None], float("-inf")) + + ksel = min(topk, blocks) + top_scores, index = torch.topk(scores, k=ksel, dim=-1) + top_valid = torch.isfinite(top_scores) + safe = top_scores.masked_fill(~top_valid, -1e4) + weights = safe.softmax(dim=-1) * top_valid + weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8) + + bank = bv[:, :, None].expand(-1, -1, t, -1, -1) + selected = torch.gather(bank, 3, index.unsqueeze(-1).expand(-1, -1, -1, -1, d)) + output = (weights.unsqueeze(-1) * selected).sum(dim=3) + available = top_valid.any(dim=-1, keepdim=True) + return output, available + + +def indexed_summary_attention(q: Tensor, block_keys: Tensor | None, + block_values: Tensor | None, groups: int, topk: int): + """One-step indexed lookup against already-completed block summaries.""" + if block_keys is None or block_keys.shape[2] == 0: + shape = (q.shape[0], q.shape[1], q.shape[2], q.shape[3]) + return q.new_zeros(shape), torch.zeros( + (q.shape[0], q.shape[1], q.shape[2], 1), dtype=torch.bool, device=q.device + ) + k = block_keys.float().repeat_interleave(groups, dim=1) + v = block_values.float().repeat_interleave(groups, dim=1) + scores = torch.einsum("bhtd,bhnd->bhtn", F.normalize(q.float(), dim=-1), + F.normalize(k, dim=-1)) / math.sqrt(q.shape[-1]) + ksel = min(topk, scores.shape[-1]) + top_scores, index = torch.topk(scores, k=ksel, dim=-1) + bank = v[:, :, None].expand(-1, -1, q.shape[2], -1, -1) + selected = torch.gather( + bank, 3, index.unsqueeze(-1).expand(-1, -1, -1, -1, q.shape[-1]) + ) + weights = top_scores.softmax(dim=-1) + return (weights.unsqueeze(-1) * selected).sum(dim=3), torch.ones( + (q.shape[0], q.shape[1], q.shape[2], 1), dtype=torch.bool, device=q.device + ) + + +class FlashPDelta2Layer(nn.Module): + """P-Delta2 plus cheap gating, short causal convolution and indexed recall.""" + + def __init__(self, num_heads: int, num_kv_heads: int, head_dim: int, + feature_dim: int = 96, chunk_size: int = 32, + output_gate: bool = False, conv_kernel: int = 1, + indexed_retrieval: bool = False, block_size: int = 16, + index_topk: int = 4, state_dtype: str = "fp16"): + super().__init__() + if state_dtype not in {"fp16", "fp32"}: + raise ValueError("state_dtype must be fp16 or fp32") + if conv_kernel < 1: + raise ValueError("conv_kernel must be positive") + if num_heads % num_kv_heads: + raise ValueError("num_heads must be divisible by num_kv_heads") + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.feature_dim = feature_dim + self.groups = num_heads // num_kv_heads + self.chunk_size = chunk_size + self.output_gate = bool(output_gate) + self.conv_kernel = int(conv_kernel) + self.indexed_retrieval = bool(indexed_retrieval) + self.block_size = int(block_size) + self.index_topk = int(index_topk) + self.state_dtype = state_dtype + + self.core = PDelta2Core( + num_heads, num_kv_heads, head_dim, feature_dim=feature_dim, chunk_size=chunk_size + ) + if self.conv_kernel > 1: + channels = num_kv_heads * head_dim + kernel = torch.zeros(channels, 1, self.conv_kernel) + kernel[:, 0, -1] = 1.0 + self.conv_weight = nn.Parameter(kernel) + else: + self.register_parameter("conv_weight", None) + + if self.output_gate: + self.output_gate_w = nn.Parameter(torch.zeros(num_heads, head_dim)) + self.output_gate_b = nn.Parameter(torch.zeros(num_heads)) + else: + self.register_parameter("output_gate_w", None) + self.register_parameter("output_gate_b", None) + + if self.indexed_retrieval: + self.index_mix_w = nn.Parameter(torch.zeros(num_heads, head_dim)) + self.index_mix_b = nn.Parameter(torch.full((num_heads,), -2.0)) + else: + self.register_parameter("index_mix_w", None) + self.register_parameter("index_mix_b", None) + + @property + def config(self): + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "chunk_size": self.chunk_size, + "output_gate": self.output_gate, + "conv_kernel": self.conv_kernel, + "indexed_retrieval": self.indexed_retrieval, + "block_size": self.block_size, + "index_topk": self.index_topk, + "state_dtype": self.state_dtype, + } + + def _convolve_values(self, v: Tensor, tail: Tensor | None = None): + v = v.float() + if self.conv_weight is None: + return v + if tail is None: + return causal_depthwise_value_conv(v, self.conv_weight) + joined = torch.cat((tail.float(), v), dim=2) + return causal_depthwise_value_conv(joined, self.conv_weight)[:, :, -v.shape[2]:] + + def _pack_recurrent_state(self, state: PDeltaState): + memory = state.memory + if self.state_dtype == "fp16": + memory = memory.to(torch.float16) + else: + memory = memory.float() + return PDeltaState(memory=memory, curvature=state.curvature.float()) + + def _unpack_recurrent_state(self, state: PDeltaState | None): + if state is None: + return None + return PDeltaState(memory=state.memory.float(), curvature=state.curvature.float()) + + def _build_index_state(self, k: Tensor, v: Tensor): + if not self.indexed_retrieval: + return None, None, None, None + complete = (k.shape[2] // self.block_size) * self.block_size + # A completed block is only available to queries after its last token. + bk = bv = None + if complete: + blocks = complete // self.block_size + bk = k[:, :, :complete].reshape( + k.shape[0], k.shape[1], blocks, self.block_size, k.shape[-1] + ).mean(dim=3).to(torch.float16) + bv = v[:, :, :complete].reshape( + v.shape[0], v.shape[1], blocks, self.block_size, v.shape[-1] + ).mean(dim=3).to(torch.float16) + pk = k[:, :, complete:].float() + pv = v[:, :, complete:].float() + return bk, bv, pk, pv + + def _advance_index_state(self, state: FlashState, k: Tensor, v: Tensor): + pk = k.float() if state.pending_keys is None else torch.cat((state.pending_keys.float(), k.float()), dim=2) + pv = v.float() if state.pending_values is None else torch.cat((state.pending_values.float(), v.float()), dim=2) + bk, bv = state.block_keys, state.block_values + while pk.shape[2] >= self.block_size: + new_k = pk[:, :, :self.block_size].mean(dim=2, keepdim=True).to(torch.float16) + new_v = pv[:, :, :self.block_size].mean(dim=2, keepdim=True).to(torch.float16) + bk = new_k if bk is None else torch.cat((bk, new_k), dim=2) + bv = new_v if bv is None else torch.cat((bv, new_v), dim=2) + pk, pv = pk[:, :, self.block_size:], pv[:, :, self.block_size:] + return bk, bv, pk, pv + + def forward(self, q: Tensor, k: Tensor, v: Tensor, state: FlashState | None = None, + return_state: bool = False, implementation: str = "chunk"): + if implementation != "chunk": + raise ValueError("PDelta2-Flash uses the chunk implementation") + if state is not None and self.indexed_retrieval and q.shape[2] != 1: + raise ValueError("stateful indexed retrieval currently supports one decode token at a time") + + tail = None if state is None else state.conv_tail + conv_v = self._convolve_values(v, tail) + recurrent_state = None if state is None else self._unpack_recurrent_state(state.recurrent) + recurrent, new_recurrent = self.core( + q, k, conv_v, state=recurrent_state, return_state=True + ) + + output = recurrent + if self.output_gate: + raw = ( + torch.einsum("bhtd,hd->bht", F.normalize(q.float(), dim=-1), self.output_gate_w) + + self.output_gate_b[None, :, None] + ) + # Exactly 1.0 at initialization; bounded to [0.75, 1.25]. + gain = 1.0 + 0.25 * torch.tanh(raw) + output = output * gain.unsqueeze(-1) + + if self.indexed_retrieval: + if state is None: + indexed, available = indexed_block_attention( + q.float(), k.float(), conv_v, self.groups, self.block_size, self.index_topk + ) + else: + indexed, available = indexed_summary_attention( + q.float(), state.block_keys, state.block_values, self.groups, self.index_topk + ) + mix = ( + torch.einsum("bhtd,hd->bht", F.normalize(q.float(), dim=-1), self.index_mix_w) + + self.index_mix_b[None, :, None] + ).sigmoid().unsqueeze(-1) + mix = mix * available.to(mix.dtype) + output = (1.0 - mix) * output + mix * indexed + + if not return_state: + return output + + keep = self.conv_kernel - 1 + if keep: + raw = v.float() if tail is None else torch.cat((tail.float(), v.float()), dim=2) + new_tail = raw[:, :, -keep:].clone() + else: + new_tail = None + + if self.indexed_retrieval: + if state is None: + bk, bv, pk, pv = self._build_index_state(k.float(), conv_v) + else: + bk, bv, pk, pv = self._advance_index_state(state, k.float(), conv_v) + else: + bk = bv = pk = pv = None + + new_state = FlashState( + recurrent=self._pack_recurrent_state(new_recurrent), + conv_tail=new_tail, + block_keys=bk, + block_values=bv, + pending_keys=pk, + pending_values=pv, + ) + return output, new_state + + def recurrent_state_bytes(self, batch_size: int = 1, context: int | None = None): + memory_elements = self.num_kv_heads * self.feature_dim * self.head_dim + curvature_elements = self.num_kv_heads * self.feature_dim + memory_bytes = 2 if self.state_dtype == "fp16" else 4 + total = batch_size * (memory_elements * memory_bytes + curvature_elements * 4) + if self.conv_kernel > 1: + total += batch_size * (self.conv_kernel - 1) * self.num_kv_heads * self.head_dim * 4 + if self.indexed_retrieval and context is not None: + blocks = context // self.block_size + total += batch_size * 2 * blocks * self.num_kv_heads * self.head_dim * 2 + pending = context % self.block_size + total += batch_size * 2 * pending * self.num_kv_heads * self.head_dim * 4 + return total + + def index_pairs_per_token(self, context: int): + if not self.indexed_retrieval: + return 0 + blocks = context // self.block_size + return min(self.index_topk, blocks) diff --git a/src/tinycenn_lm/pdelta3_frontier.py b/src/tinycenn_lm/pdelta3_frontier.py new file mode 100644 index 0000000000000000000000000000000000000000..6a04e76927824284cff33c72c334d155481b9369 --- /dev/null +++ b/src/tinycenn_lm/pdelta3_frontier.py @@ -0,0 +1,390 @@ +"""Frontier-inspired PDelta3 research layers. + +Independent TinyCeNN adaptations for a controlled frozen-attention replacement lab: + +* ``conv4_pdelta_f96`` keeps the proven Conv4 + PDelta2 recurrence; +* ``conv4_channel_decay_f96`` replaces the homogeneous forget path with a + KDA/GDN2-style content-dependent channel decay; +* ``conv4_gdn2_f96`` uses the Gated DeltaNet-2 update structure with independent + key-channel erase and value-channel write gates; +* ``conv4_gdn2_clvr_f96`` additionally routes the previous layer's aligned value + representation into the current write target (a direct one-hop CLVR test). + +The implementation is written from the published recurrence rather than copied +from any external kernel. It is a PyTorch research reference, not a claim of +kernel-level reproduction or speed parity with fused frontier implementations. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from tinycenn_lm.pdelta2_er import causal_value_conv_stream +from tinycenn_lm.pdelta2_features import PDelta2Core, PDeltaState +from tinycenn_lm.research_layers import delta_recurrence + + +VARIANTS = ( + "conv4_pdelta_f96", + "conv4_channel_decay_f96", + "conv4_gdn2_f96", + "conv4_gdn2_clvr_f96", +) + + +@dataclass +class FrontierState: + memory: Tensor + curvature: Tensor | None = None + q_tail: Tensor | None = None + k_tail: Tensor | None = None + v_tail: Tensor | None = None + + +def _inverse_softplus(x: Tensor) -> Tensor: + return x + torch.log(-torch.expm1(-x)) + + +def _orthogonal_maps(heads: int, feature_dim: int, head_dim: int) -> Tensor: + base = torch.empty(heads, feature_dim, head_dim) + for head in range(heads): + if feature_dim == head_dim: + base[head] = torch.eye(head_dim) + else: + nn.init.orthogonal_(base[head]) + return base + + +class FrontierPDelta3Layer(nn.Module): + """Conv4 recurrent attention replacement with four controlled variants. + + GDN2-style variants implement, in feature space, + + S_t = (I - k_t (b_t * k_t)^T) D_t S_{t-1} + + k_t (w_t * v_t)^T + + where ``D_t`` is channel-wise decay, ``b_t`` is key-channel erase, and + ``w_t`` is value-channel write. The CLVR variant adds a learned aligned + value from the immediately preceding Transformer layer to the write target; + it does not add another temporal memory matrix. + """ + + def __init__( + self, + num_heads: int, + num_kv_heads: int, + head_dim: int, + feature_dim: int = 96, + variant: str = "conv4_pdelta_f96", + chunk_size: int = 32, + conv_kernel: int = 4, + state_dtype: str = "fp16", + ): + super().__init__() + if variant not in VARIANTS: + raise ValueError(f"variant must be one of {VARIANTS}") + if num_heads % num_kv_heads: + raise ValueError("num_heads must be divisible by num_kv_heads") + if state_dtype not in {"fp16", "fp32"}: + raise ValueError("state_dtype must be fp16 or fp32") + if min(num_heads, num_kv_heads, head_dim, feature_dim, conv_kernel) < 1: + raise ValueError("dimensions must be positive") + if not 1 <= chunk_size <= 32: + raise ValueError("chunk_size must be in [1,32]") + + self.num_heads = int(num_heads) + self.num_kv_heads = int(num_kv_heads) + self.head_dim = int(head_dim) + self.feature_dim = int(feature_dim) + self.variant = variant + self.chunk_size = int(chunk_size) + self.conv_kernel = int(conv_kernel) + self.state_dtype = state_dtype + self.groups = self.num_heads // self.num_kv_heads + + self.is_pdelta = variant == "conv4_pdelta_f96" + self.is_channel_decay = variant == "conv4_channel_decay_f96" + self.is_gdn2 = variant in {"conv4_gdn2_f96", "conv4_gdn2_clvr_f96"} + self.use_clvr = variant == "conv4_gdn2_clvr_f96" + + if self.is_pdelta or self.is_channel_decay: + self.pdelta = PDelta2Core( + self.num_heads, + self.num_kv_heads, + self.head_dim, + feature_dim=self.feature_dim, + chunk_size=self.chunk_size, + ) + else: + self.pdelta = None + + if self.is_gdn2: + kv_map = _orthogonal_maps(self.num_kv_heads, self.feature_dim, self.head_dim) + self.wk = nn.Parameter(kv_map) + self.wq = nn.Parameter(kv_map.repeat_interleave(self.groups, dim=0).clone()) + self.erase_w = nn.Parameter( + torch.zeros(self.num_kv_heads, self.feature_dim, self.head_dim) + ) + self.erase_b = nn.Parameter(torch.full((self.num_kv_heads, self.feature_dim), -1.0)) + self.write_w = nn.Parameter( + torch.zeros(self.num_kv_heads, self.head_dim, self.head_dim) + ) + self.write_b = nn.Parameter(torch.full((self.num_kv_heads, self.head_dim), -1.0)) + self.log_gain = nn.Parameter(torch.zeros(self.num_heads)) + else: + self.register_parameter("wk", None) + self.register_parameter("wq", None) + self.register_parameter("erase_w", None) + self.register_parameter("erase_b", None) + self.register_parameter("write_w", None) + self.register_parameter("write_b", None) + self.register_parameter("log_gain", None) + + if not self.is_pdelta: + self.decay_w = nn.Parameter( + torch.zeros(self.num_kv_heads, self.feature_dim, self.head_dim) + ) + dt = torch.exp(torch.linspace( + math.log(0.001), math.log(0.05), self.feature_dim + )).clamp_min(1e-4) + self.dt_bias = nn.Parameter( + _inverse_softplus(dt)[None].expand(self.num_kv_heads, -1).clone() + ) + self.A_log = nn.Parameter(torch.zeros(self.num_kv_heads, self.feature_dim)) + else: + self.register_parameter("decay_w", None) + self.register_parameter("dt_bias", None) + self.register_parameter("A_log", None) + + # The proven control uses Conv4 on V. GDN2-style candidates use the + # frontier-model Q/K/V short-convolution pattern. + if self.is_gdn2: + self.q_conv_weight = nn.Parameter(self._make_conv_weight(self.num_heads)) + self.k_conv_weight = nn.Parameter(self._make_conv_weight(self.num_kv_heads)) + else: + self.register_parameter("q_conv_weight", None) + self.register_parameter("k_conv_weight", None) + self.v_conv_weight = nn.Parameter(self._make_conv_weight(self.num_kv_heads)) + + if self.use_clvr: + route = torch.eye(self.head_dim)[None].repeat(self.num_kv_heads, 1, 1) + self.route_proj = nn.Parameter(route) + self.route_gate_w = nn.Parameter( + torch.zeros(self.num_kv_heads, self.head_dim, self.head_dim) + ) + self.route_gate_b = nn.Parameter( + torch.full((self.num_kv_heads, self.head_dim), -4.0) + ) + else: + self.register_parameter("route_proj", None) + self.register_parameter("route_gate_w", None) + self.register_parameter("route_gate_b", None) + + def _make_conv_weight(self, heads: int) -> Tensor: + channels = heads * self.head_dim + kernel = torch.zeros(channels, 1, self.conv_kernel) + kernel[:, 0, -1] = 1.0 + return kernel + + @property + def config(self): + return { + "num_heads": self.num_heads, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "feature_dim": self.feature_dim, + "variant": self.variant, + "chunk_size": self.chunk_size, + "conv_kernel": self.conv_kernel, + "state_dtype": self.state_dtype, + } + + @property + def storage_dtype(self): + return torch.float16 if self.state_dtype == "fp16" else torch.float32 + + def _conv(self, x: Tensor, weight: Tensor | None, tail: Tensor | None): + if weight is None: + return x.float(), None + return causal_value_conv_stream(x.float(), weight, tail) + + @staticmethod + def _project(x: Tensor, weight: Tensor, bias: Tensor | None = None): + y = torch.einsum("bhtd,hfd->bhtf", x, weight) + return y if bias is None else y + bias[None, :, None] + + def _channel_decay(self, k: Tensor) -> Tensor: + if self.decay_w is None: + raise RuntimeError("channel decay is not configured for this variant") + kn = F.normalize(k.float(), dim=-1) + raw_dt = self._project(kn, self.decay_w, self.dt_bias) + dt = F.softplus(raw_dt.float()) + rate = self.A_log.float().clamp(-4, 2).exp()[None, :, None, :] * dt + return -rate.clamp(1e-5, 0.25) + + def _working_state(self, state: FrontierState | None): + if state is None: + return None + dtype = next(self.parameters()).dtype + return FrontierState( + memory=state.memory.to(dtype), + curvature=None if state.curvature is None else state.curvature.to(dtype), + q_tail=state.q_tail, + k_tail=state.k_tail, + v_tail=state.v_tail, + ) + + def _store_state(self, memory: Tensor, curvature: Tensor | None, + q_tail: Tensor | None, k_tail: Tensor | None, v_tail: Tensor | None): + return FrontierState( + memory=memory.to(self.storage_dtype), + curvature=None if curvature is None else curvature.float(), + q_tail=None if q_tail is None else q_tail.to(self.storage_dtype), + k_tail=None if k_tail is None else k_tail.to(self.storage_dtype), + v_tail=None if v_tail is None else v_tail.to(self.storage_dtype), + ) + + def _initial_memory(self, q: Tensor): + return q.new_zeros( + q.shape[0], self.num_kv_heads, self.feature_dim, self.head_dim + ) + + def _run_pdelta(self, q: Tensor, k: Tensor, v: Tensor, state: FrontierState | None, + channel_decay: bool): + v_conv, v_tail = self._conv( + v, self.v_conv_weight, None if state is None else state.v_tail + ) + pstate = None + if state is not None: + if state.curvature is None: + raise ValueError("PDelta variants require a curvature state") + pstate = PDeltaState( + state.memory.to(self.pdelta.wq.dtype), + state.curvature.to(self.pdelta.wq.dtype), + ) + if not channel_decay: + output, new = self.pdelta(q, k, v_conv, state=pstate, return_state=True) + else: + qf, kf, z, erase, _ = self.pdelta.features( + q.to(self.pdelta.wq.dtype), + k.to(self.pdelta.wq.dtype), + v_conv.to(self.pdelta.wq.dtype), + ) + curvature = ( + qf.new_ones(qf.shape[0], self.num_kv_heads, self.feature_dim) + if pstate is None else pstate.curvature + ) + memory = self._initial_memory(qf) if pstate is None else pstate.memory + kpre, curvature = self.pdelta.precondition_keys(kf, curvature) + log_decay = self._channel_decay(k) + output, memory = delta_recurrence( + qf, kpre, z, erase, log_decay, memory, + self.groups, self.chunk_size, + ) + output = output * self.pdelta.log_gain.clamp(-4, 4).exp()[None, :, None, None] + new = PDeltaState(memory, curvature) + stored = self._store_state(new.memory, new.curvature, None, None, v_tail) + return output, stored + + def _run_gdn2(self, q: Tensor, k: Tensor, v: Tensor, routed_v: Tensor | None, + state: FrontierState | None): + q_conv, q_tail = self._conv( + q, self.q_conv_weight, None if state is None else state.q_tail + ) + k_conv, k_tail = self._conv( + k, self.k_conv_weight, None if state is None else state.k_tail + ) + v_conv, v_tail = self._conv( + v, self.v_conv_weight, None if state is None else state.v_tail + ) + + qn, kn = F.normalize(q_conv, dim=-1), F.normalize(k_conv, dim=-1) + qf = F.normalize(self._project(qn, self.wq), dim=-1) + kf = F.normalize(self._project(kn, self.wk), dim=-1) + log_decay = self._channel_decay(k_conv) + + erase_gate = self._project(kn, self.erase_w, self.erase_b).sigmoid() + write_gate = self._project( + F.normalize(v_conv, dim=-1), self.write_w, self.write_b + ).sigmoid() + write_value = v_conv + + if self.use_clvr: + if routed_v is None: + routed_v = torch.zeros_like(v_conv) + if routed_v.shape != v_conv.shape: + raise ValueError("routed_v must match current V shape") + aligned = torch.einsum( + "bhtd,hde->bhte", routed_v.float(), self.route_proj.float() + ) + route_gate = self._project(kn, self.route_gate_w, self.route_gate_b).sigmoid() + write_value = write_value + route_gate * aligned + + z = write_gate * write_value + erase = erase_gate * kf + memory = self._initial_memory(qf) if state is None else state.memory.to(qf.dtype) + output, memory = delta_recurrence( + qf, kf, z, erase, log_decay, memory, + self.groups, self.chunk_size, + ) + output = output * self.log_gain.clamp(-4, 4).exp()[None, :, None, None] + stored = self._store_state(memory, None, q_tail, k_tail, v_tail) + return output, stored + + def forward(self, q: Tensor, k: Tensor, v: Tensor, + state: FrontierState | None = None, return_state: bool = False, + implementation: str = "chunk", routed_v: Tensor | None = None): + if implementation != "chunk": + raise ValueError("FrontierPDelta3Layer supports the bounded chunk implementation") + if q.ndim != 4 or k.ndim != 4 or v.shape != k.shape: + raise ValueError("expected [B,H,T,D] Q/K/V") + if self.use_clvr and routed_v is not None and routed_v.shape != v.shape: + raise ValueError("CLVR routed value shape must match V") + + state = self._working_state(state) + if self.is_pdelta: + output, new_state = self._run_pdelta(q, k, v, state, channel_decay=False) + elif self.is_channel_decay: + output, new_state = self._run_pdelta(q, k, v, state, channel_decay=True) + else: + output, new_state = self._run_gdn2(q, k, v, routed_v, state) + return (output, new_state) if return_state else output + + def recurrent_state_bytes(self, batch_size: int = 1, context: int | None = None): + del context + memory_bytes = 2 if self.state_dtype == "fp16" else 4 + total = self.num_kv_heads * self.feature_dim * self.head_dim * memory_bytes + if self.is_pdelta or self.is_channel_decay: + total += self.num_kv_heads * self.feature_dim * 4 + conv_heads = self.num_kv_heads + else: + conv_heads = self.num_heads + 2 * self.num_kv_heads + if self.conv_kernel > 1: + total += (self.conv_kernel - 1) * conv_heads * self.head_dim * memory_bytes + return batch_size * total + + def decay_statistics(self): + if self.is_pdelta: + return {} + with torch.no_grad(): + dt = F.softplus(self.dt_bias.float()) + rate = self.A_log.float().clamp(-4, 2).exp() * dt + half_life = math.log(2.0) / rate.clamp_min(1e-8) + return { + "decay_half_life_min": float(half_life.min()), + "decay_half_life_median": float(half_life.median()), + "decay_half_life_max": float(half_life.max()), + } + + def gate_statistics(self): + result = {} + if self.is_gdn2: + result["erase_bias_mean"] = float(self.erase_b.detach().sigmoid().mean()) + result["write_bias_mean"] = float(self.write_b.detach().sigmoid().mean()) + if self.use_clvr: + result["route_bias_mean"] = float(self.route_gate_b.detach().sigmoid().mean()) + return result diff --git a/src/tinycenn_lm/qwen35_integrated_memory.py b/src/tinycenn_lm/qwen35_integrated_memory.py new file mode 100644 index 0000000000000000000000000000000000000000..2a43be53c49e7c066f374a0effbf017404b114b5 --- /dev/null +++ b/src/tinycenn_lm/qwen35_integrated_memory.py @@ -0,0 +1,264 @@ +"""Qwen3.5 text-backbone TinyCeNN integrated memory. + +Replaces selected full-attention layers in Qwen3.5 while preserving its native +Gated DeltaNet linear-attention layers and Qwen3.5's post-attention output gate. +""" +import copy +from contextlib import contextmanager +import torch +import torch.nn.functional as F +from torch import nn +from transformers.cache_utils import DynamicCache +from transformers.models.qwen3_5.modeling_qwen3_5 import apply_rotary_pos_emb, repeat_kv +from .optimized_memory import OptimizedMemory + +FORMAT = "qwen35-cenn-integrated-v1" + + +def native_dtype(device): + if torch.device(device).type != "cuda": + return torch.float32 + return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16 + + +def text_config(model): + cfg = model.config + return cfg.get_text_config(decoder=True) if hasattr(cfg, "get_text_config") else cfg + + +class Qwen35IntegratedCache(DynamicCache): + def __init__(self, config, memory_layers=()): + super().__init__(config=config) + self.memory_layers = frozenset(int(x) for x in memory_layers) + self.memory_states = {} + + def _memory_seq_length(self, layer_idx=None): + if layer_idx is not None: + state = self.memory_states.get(int(layer_idx)) + return int(state.position) if state is not None else 0 + return max((int(state.position) for state in self.memory_states.values()), default=0) + + def get_seq_length(self, layer_idx=0): + """Return the real sequence length even when the first attention layer is CeNN. + + Qwen3.5 layer 0 is linear attention. Hugging Face's DynamicCache therefore + redirects the default get_seq_length() call to the first full-attention + cache layer. When that first full-attention layer (layer 3 in Qwen3.5-0.8B) + is replaced by TinyCeNN, its DynamicLayer is intentionally never updated; + the position lives in memory_states instead. Without this override, cached + decoding repeatedly reports length 0 and reuses incorrect RoPE positions. + """ + if layer_idx in self.memory_layers: + return self._memory_seq_length(layer_idx) + if layer_idx == 0: + memory_length = self._memory_seq_length() + try: + native_length = int(super().get_seq_length(layer_idx)) + except (ValueError, StopIteration): + native_length = 0 + return max(memory_length, native_length) + return super().get_seq_length(layer_idx) + + def get_mask_sizes(self, query_length, layer_idx): + """Mirror get_seq_length() for causal-mask construction at CeNN layers.""" + if layer_idx in self.memory_layers: + return self._memory_seq_length(layer_idx) + int(query_length), 0 + if layer_idx == 0 and self.memory_states: + return self.get_seq_length(0) + int(query_length), 0 + return super().get_mask_sizes(query_length, layer_idx) + + @staticmethod + def _bytes(value): + if isinstance(value, torch.Tensor): + return value.numel() * value.element_size() + if isinstance(value, (list, tuple)): + return sum(Qwen35IntegratedCache._bytes(x) for x in value) + if isinstance(value, dict): + return sum(Qwen35IntegratedCache._bytes(x) for x in value.values()) + return 0 + + @property + def nbytes(self): + total = 0 + for layer in self.layers: + for name in ("keys", "values", "conv_states", "recurrent_states"): + total += self._bytes(getattr(layer, name, None)) + total += sum(state.nbytes for state in self.memory_states.values()) + return total + + def reorder_cache(self, beam_idx): + if self.memory_states: + raise NotImplementedError("Beam/batch reordering is unsupported for TinyCeNN states") + return super().reorder_cache(beam_idx) + + def crop(self, *args, **kwargs): + if self.memory_states: + raise NotImplementedError("Compressed TinyCeNN history cannot be cropped") + return super().crop(*args, **kwargs) + + +class Qwen35IntegratedAttention(nn.Module): + """Drop-in replacement for a Qwen3.5 full-attention module.""" + def __init__(self, original, core, layer_idx): + super().__init__() + self.original = original + self.core = core + self.layer_idx = int(layer_idx) + self.config = original.config + self.head_dim = original.head_dim + self.num_key_value_groups = original.num_key_value_groups + self.scaling = original.scaling + self.attention_dropout = original.attention_dropout + self.is_causal = original.is_causal + + def forward(self, hidden_states, position_embeddings=None, attention_mask=None, + past_key_values=None, past_key_value=None, **kwargs): + if position_embeddings is None: + raise ValueError("Qwen3.5 position_embeddings are required") + cache = past_key_values if past_key_values is not None else past_key_value + if cache is not None and not isinstance(cache, Qwen35IntegratedCache): + raise TypeError("Use Qwen35IntegratedCache with a Qwen3.5 TinyCeNN model") + if cache is not None and torch.is_grad_enabled(): + raise RuntimeError("Train with use_cache=False") + + input_shape = hidden_states.shape[:-1] + hidden_shape = (*input_shape, -1, self.head_dim) + t = hidden_states.shape[1] + + query_states, gate = torch.chunk( + self.original.q_proj(hidden_states).view(*input_shape, -1, self.head_dim * 2), 2, dim=-1 + ) + gate = gate.reshape(*input_shape, -1) + query_states = self.original.q_norm(query_states.view(hidden_shape)).transpose(1, 2) + key_states = self.original.k_norm(self.original.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2) + value_states = self.original.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) + cos, sin = position_embeddings + query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) + + if attention_mask is not None: + if attention_mask.ndim != 4 or bool((attention_mask[..., -1, :] < 0).any()): + raise ValueError("Only unpadded causal batches are supported") + + if self.core.variant == "transformer_readout": + if cache is not None: + key_states, value_states = cache.update(key_states, value_states, self.layer_idx) + output = F.scaled_dot_product_attention( + query_states, + repeat_kv(key_states, self.num_key_value_groups), + repeat_kv(value_states, self.num_key_value_groups), + attn_mask=attention_mask, + is_causal=attention_mask is None and t > 1, + scale=float(self.scaling), + ) + output = self.core.calibrate(output.float()) + elif cache is not None: + output, state = self.core( + query_states, key_states, value_states, + state=cache.memory_states.get(self.layer_idx), + return_state=True, apply_readout=True, + ) + cache.memory_states[self.layer_idx] = state + else: + output = self.core(query_states, key_states, value_states, apply_readout=True) + + flat = output.transpose(1, 2).reshape(*input_shape, -1).to(hidden_states.dtype) + # Important: Qwen3.5 gates after attention. Do not fold the TinyCeNN + # readout into o_proj because that would not commute with this gate. + flat = flat * torch.sigmoid(gate) + return self.original.o_proj(flat), None + + +def wrappers(model): + return [layer.self_attn for layer in model.model.layers + if getattr(layer, "block_type", None) == "full_attention" + and isinstance(getattr(layer, "self_attn", None), Qwen35IntegratedAttention)] + + +def full_attention_layers(model): + return [i for i, layer in enumerate(model.model.layers) + if getattr(layer, "block_type", None) == "full_attention"] + + +def build_student(teacher, layers, variant="cenn_partition", features=64, block_size=32, sinks=4): + model = copy.deepcopy(teacher).eval().requires_grad_(False) + cfg = text_config(model) + if cfg.model_type != "qwen3_5_text": + raise ValueError(f"Expected qwen3_5_text, got {cfg.model_type}") + for index in layers: + if not 0 <= index < len(model.model.layers): + raise ValueError(f"Invalid layer {index}") + layer = model.model.layers[index] + if getattr(layer, "block_type", None) != "full_attention": + raise ValueError(f"Layer {index} is not full_attention") + original = layer.self_attn + core = OptimizedMemory( + cfg.num_attention_heads, cfg.num_key_value_heads, original.head_dim, + features, variant, block_size, sinks, + ).to(original.q_proj.weight.device) + layer.self_attn = Qwen35IntegratedAttention(original, core, index) + return model + + +@contextmanager +def inference_mode(model, compute_dtype="float32"): + adapters = wrappers(model) + previous = [a.core.compute_dtype for a in adapters] + try: + for adapter in adapters: + adapter.core.compute_dtype = compute_dtype + with torch.no_grad(): + yield model + finally: + for adapter, dtype in zip(adapters, previous): + adapter.core.compute_dtype = dtype + + +def new_cache(model): + cfg = text_config(model) + memory_layers = [a.layer_idx for a in wrappers(model) if a.core.variant != "transformer_readout"] + return Qwen35IntegratedCache(cfg, memory_layers) + + +def adapter_payload(model, metadata=None): + return { + "format": FORMAT, + "metadata": metadata or {}, + "adapters": { + str(a.layer_idx): { + "config": a.core.config, + "state_dict": {k: v.detach().cpu().clone() for k, v in a.core.state_dict().items()}, + } for a in wrappers(model) + }, + } + + +def restore_student(teacher, payload): + if payload["format"] != FORMAT: + raise ValueError(f"Not a {FORMAT} checkpoint") + model = copy.deepcopy(teacher).eval().requires_grad_(False) + for key, value in payload["adapters"].items(): + index = int(key) + layer = model.model.layers[index] + if getattr(layer, "block_type", None) != "full_attention": + raise ValueError(f"Checkpoint targets non-full-attention layer {index}") + original = layer.self_attn + core = OptimizedMemory(**value["config"]).to(original.q_proj.weight.device) + core.load_state_dict(value["state_dict"]) + layer.self_attn = Qwen35IntegratedAttention(original, core, index) + return model + + +@torch.no_grad() +def greedy_generate(model, ids, tokens=32, stop_token_ids=()): + if ids.shape[0] != 1 or ids.shape[1] < 1 or tokens < 1: + raise ValueError("Use batch size one, a nonempty prompt, and positive tokens") + stop_token_ids = set(int(x) for x in stop_token_ids) + cache = new_cache(model) + logits = model(input_ids=ids, past_key_values=cache, use_cache=True).logits[:, -1] + continuation = [logits.argmax(-1, keepdim=True)] + for _ in range(tokens - 1): + if int(continuation[-1].item()) in stop_token_ids: + break + logits = model(input_ids=continuation[-1], past_key_values=cache, use_cache=True).logits[:, -1] + continuation.append(logits.argmax(-1, keepdim=True)) + return torch.cat(continuation, dim=1), cache diff --git a/src/tinycenn_lm/qwen3_5_memory_fusion.py b/src/tinycenn_lm/qwen3_5_memory_fusion.py new file mode 100644 index 0000000000000000000000000000000000000000..673446500ecfb17603d7508de335c709c311f1c6 --- /dev/null +++ b/src/tinycenn_lm/qwen3_5_memory_fusion.py @@ -0,0 +1,273 @@ +from __future__ import annotations + +import copy +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable + +import torch +from torch import Tensor, nn +from transformers.models.qwen3_5.modeling_qwen3_5 import apply_rotary_pos_emb + +from .memory_attention import MemoryAugmentedCellularLayer + +DEFAULT_QWEN35 = "Qwen/Qwen3.5-0.8B" +FORMAT = "qwen3.5-memory-fusion-sequential-v1" + + +@dataclass(frozen=True) +class Qwen35MemoryFusionConfig: + feature_dim: int = 32 + memory_rank: int = 64 + dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64, 128) + shifted_window: int = 8 + train_output_projection: bool = True + + def validate(self, model_config) -> None: + config = model_config.get_text_config(decoder=True) if hasattr(model_config, "get_text_config") else model_config + if getattr(config, "model_type", None) != "qwen3_5_text": + raise ValueError(f"expected qwen3_5_text, got {getattr(config, 'model_type', None)!r}") + if self.feature_dim < 4: + raise ValueError("feature_dim must be >= 4") + if self.memory_rank < 4: + raise ValueError("memory_rank must be >= 4") + if not self.dilations or min(self.dilations) < 1: + raise ValueError("dilations must be positive") + if int(config.num_attention_heads) % int(config.num_key_value_heads): + raise ValueError("num_attention_heads must be divisible by num_key_value_heads") + + def to_dict(self) -> dict: + value = asdict(self) + value["dilations"] = list(self.dilations) + return value + + @classmethod + def from_dict(cls, data: dict) -> "Qwen35MemoryFusionConfig": + value = dict(data) + value["dilations"] = tuple(value.get("dilations", (1, 2, 4, 8, 16, 32, 64, 128))) + return cls(**value) + + +class MemoryFusionQwen35Attention(nn.Module): + """Replace a Qwen3.5 *full-attention* anchor with TinyCeNN Memory Fusion. + + Qwen3.5 already contains native Gated DeltaNet linear-attention layers. This + adapter deliberately leaves those layers untouched and only targets the + original full-attention anchors. The pretrained Q/K/V/O projections, Q/K + RMS normalizers, partial RoPE path, and Qwen attention-output gate are kept. + + V1 is a research/full-prefix implementation and requires ``use_cache=False``. + """ + + def __init__(self, original_attn: nn.Module, model_config, config: Qwen35MemoryFusionConfig, layer_idx: int): + super().__init__() + config.validate(model_config) + text_config = model_config.get_text_config(decoder=True) if hasattr(model_config, "get_text_config") else model_config + self.layer_idx = int(layer_idx) + self.config = getattr(original_attn, "config", text_config) + self.hidden_size = int(text_config.hidden_size) + self.num_heads = int(text_config.num_attention_heads) + self.num_key_value_heads = int(text_config.num_key_value_heads) + self.head_dim = int(getattr(original_attn, "head_dim", text_config.head_dim)) + self.attention_width = self.num_heads * self.head_dim + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.scaling = float(getattr(original_attn, "scaling", self.head_dim ** -0.5)) + self.attention_dropout = float(getattr(original_attn, "attention_dropout", 0.0)) + self.is_causal = True + self.layer_type = "full_attention" + + # Qwen3.5 q_proj emits both query and output-gate channels. + self.q_proj = copy.deepcopy(original_attn.q_proj) + self.k_proj = copy.deepcopy(original_attn.k_proj) + self.v_proj = copy.deepcopy(original_attn.v_proj) + self.o_proj = copy.deepcopy(original_attn.o_proj) + self.q_norm = copy.deepcopy(original_attn.q_norm) + self.k_norm = copy.deepcopy(original_attn.k_norm) + + self.core = MemoryAugmentedCellularLayer( + num_heads=self.num_heads, + num_kv_heads=self.num_key_value_heads, + head_dim=self.head_dim, + feature_dim=config.feature_dim, + variant="cellular_memory_fusion", + dilations=config.dilations, + shifted_window=config.shifted_window, + memory_rank=config.memory_rank, + ) + self.last_core_output: Tensor | None = None + + def forward( + self, + hidden_states: Tensor, + position_embeddings=None, + attention_mask=None, + position_ids=None, + past_key_values=None, + past_key_value=None, + use_cache: bool = False, + cache_position=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache or past_key_values is not None or past_key_value is not None: + raise RuntimeError("Qwen3.5 Memory Fusion V1 currently requires use_cache=False") + if position_embeddings is None: + raise ValueError("Qwen3.5 position_embeddings are required") + + bsz, seq_len, _ = hidden_states.shape + q_and_gate = self.q_proj(hidden_states).view( + bsz, seq_len, self.num_heads, self.head_dim * 2 + ) + query, gate = torch.chunk(q_and_gate, 2, dim=-1) + gate = gate.reshape(bsz, seq_len, self.attention_width) + + q = self.q_norm(query).transpose(1, 2) + k = self.k_norm( + self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ) + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + + cos, sin = position_embeddings + q, k = apply_rotary_pos_emb(q, k, cos, sin) + + core_out = self.core(q.float(), k.float(), v.float()) + self.last_core_output = core_out + flat = core_out.transpose(1, 2).reshape(bsz, seq_len, self.attention_width) + flat = flat.to(hidden_states.dtype) * torch.sigmoid(gate).to(hidden_states.dtype) + return self.o_proj(flat), None + + +def text_model(model: nn.Module) -> nn.Module: + """Return Qwen3.5's text decoder for text-only or multimodal wrappers.""" + base = getattr(model, "model", model) + return getattr(base, "language_model", base) + + +def full_attention_layers(model: nn.Module) -> list[int]: + backbone = text_model(model) + return [ + i for i, kind in enumerate(backbone.config.layer_types) + if kind == "full_attention" + ] + + +def replace_attention_layers(model: nn.Module, config: Qwen35MemoryFusionConfig, layer_indices: Iterable[int]) -> nn.Module: + config.validate(model.config) + backbone = text_model(model) + available = set(full_attention_layers(model)) + for raw_idx in layer_indices: + idx = int(raw_idx) + if idx not in available: + raise ValueError( + f"layer {idx} is not a Qwen3.5 full-attention anchor; available={sorted(available)}" + ) + layer = backbone.layers[idx] + if isinstance(layer.self_attn, MemoryFusionQwen35Attention): + continue + old = layer.self_attn + device = old.q_proj.weight.device + projection_dtype = old.q_proj.weight.dtype + new = MemoryFusionQwen35Attention(old, model.config, config, idx) + for module in (new.q_proj, new.k_proj, new.v_proj, new.o_proj, new.q_norm, new.k_norm): + module.to(device=device, dtype=projection_dtype) + # Keep the research memory core numerically stable in FP32. + new.core.to(device=device, dtype=torch.float32) + layer.self_attn = new + model.config.use_cache = False + backbone.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def freeze_current_layer_only(model: nn.Module, layer_idx: int, *, train_output_projection: bool = True) -> list[nn.Parameter]: + for p in model.parameters(): + p.requires_grad = False + module = text_model(model).layers[int(layer_idx)].self_attn + if not isinstance(module, MemoryFusionQwen35Attention): + raise TypeError(f"layer {layer_idx} is not MemoryFusionQwen35Attention") + trainable: list[nn.Parameter] = [] + for p in module.core.parameters(): + p.requires_grad = True + trainable.append(p) + if train_output_projection: + for p in module.o_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def structural_summary(model: nn.Module) -> dict[str, object]: + backbone = text_model(model) + fusion = [ + i for i, layer in enumerate(backbone.layers) + if hasattr(layer, "self_attn") and isinstance(layer.self_attn, MemoryFusionQwen35Attention) + ] + remaining_full = [ + i for i, kind in enumerate(backbone.config.layer_types) + if kind == "full_attention" and i not in fusion + ] + linear = [ + i for i, kind in enumerate(backbone.config.layer_types) + if kind == "linear_attention" + ] + return { + "memory_fusion_layers": fusion, + "remaining_full_attention_layers": remaining_full, + "native_linear_attention_layers": linear, + } + + +def selected_attention_state(model: nn.Module, layers: Iterable[int]) -> dict[str, Tensor]: + backbone = text_model(model) + result: dict[str, Tensor] = {} + for idx in (int(i) for i in layers): + module = backbone.layers[idx].self_attn + for key, value in module.state_dict().items(): + result[f"layers.{idx}.self_attn.{key}"] = value.detach().cpu() + return result + + +def load_selected_attention_state(model: nn.Module, state: dict[str, Tensor], layers: Iterable[int]) -> None: + backbone = text_model(model) + for idx in (int(i) for i in layers): + prefix = f"layers.{idx}.self_attn." + local = {k[len(prefix):]: v for k, v in state.items() if k.startswith(prefix)} + incompatible = backbone.layers[idx].self_attn.load_state_dict(local, strict=False) + if incompatible.missing_keys or incompatible.unexpected_keys: + raise RuntimeError( + f"layer {idx} checkpoint mismatch: missing={incompatible.missing_keys[:6]} " + f"unexpected={incompatible.unexpected_keys[:6]}" + ) + + +def save_adapter( + model: nn.Module, + output_dir: str | Path, + *, + config: Qwen35MemoryFusionConfig, + base_model: str, + accepted_layers: list[int], + metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save( + selected_attention_state(model, accepted_layers), + output_dir / "qwen35_memory_fusion.pt", + ) + payload = { + "format": FORMAT, + "base_model": base_model, + "accepted_layers": list(accepted_layers), + "memory_fusion": config.to_dict(), + "metadata": metadata or {}, + } + (output_dir / "qwen35_memory_fusion_config.json").write_text( + json.dumps(payload, indent=2), encoding="utf-8" + ) + return output_dir diff --git a/src/tinycenn_lm/research_layers.py b/src/tinycenn_lm/research_layers.py new file mode 100644 index 0000000000000000000000000000000000000000..d5f493ab87d380c26e665f725c7881d847a0d597 --- /dev/null +++ b/src/tinycenn_lm/research_layers.py @@ -0,0 +1,222 @@ +"""Research adaptations of KDA and Gated DeltaNet-2 for frozen GQA projections. + +Equations and deviations are documented in RESEARCH_LAYERS.md. This is an +independent PyTorch reference implementation, not the authors' fused kernels. +All state updates use float32 by default; no T-by-T matrix is constructed. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +from torch import Tensor, nn +import torch.nn.functional as F + +VARIANTS = ("cenn_kda", "cenn_delta2", "cenn_delta2_window") + + +@dataclass +class DeltaState: + memory: Tensor + keys: Tensor | None = None + values: Tensor | None = None + + @property + def nbytes(self) -> int: + return sum(x.numel() * x.element_size() + for x in (self.memory, self.keys, self.values) if x is not None) + + +def delta_recurrence(q, k, z, erase, log_decay, memory, groups, chunk_size=32): + """Asymmetric delta update, with a triangular solve inside each bounded chunk. + + Sbar_t = diag(exp(log_decay_t)) S_{t-1} + S_t = Sbar_t + k_t (z_t - erase_t^T Sbar_t)^T. + q: [B,Hq,T,F]; other features: [B,Hkv,T,F]; z: [B,Hkv,T,Dv]. + The chunk solve is algebraically identical to the token recurrence. + """ + if not 1 <= chunk_size <= 32: + raise ValueError("chunk_size must be in [1,32] for the bounded-decay reference") + outputs = [] + for start in range(0, q.shape[2], chunk_size): + stop = min(start + chunk_size, q.shape[2]) + kc, ec, zc = k[:, :, start:stop], erase[:, :, start:stop], z[:, :, start:stop] + decay = log_decay[:, :, start:stop].cumsum(dim=2).exp() + write = kc / decay + read = ec * decay + n = stop - start + interaction = torch.matmul(read, write.transpose(-1, -2)) + lower = interaction.tril(diagonal=-1) + system = lower + torch.eye(n, device=q.device, dtype=q.dtype) + residual = torch.linalg.solve_triangular( + system, zc - torch.matmul(read, memory), upper=False, unitriangular=True + ) + qr = q[:, :, start:stop] * decay.repeat_interleave(groups, dim=1) + wr = write.repeat_interleave(groups, dim=1) + rr = residual.repeat_interleave(groups, dim=1) + old = torch.matmul(qr, memory.repeat_interleave(groups, dim=1)) + within = torch.matmul(qr, wr.transpose(-1, -2)).tril() + outputs.append(old + torch.matmul(within, rr)) + memory = decay[:, :, -1, :, None] * ( + memory + torch.matmul(write.transpose(-1, -2), residual) + ) + return torch.cat(outputs, dim=2), memory + + +def delta_token_reference(q, k, z, erase, log_decay, memory, groups): + """Independent, slow tokenwise oracle used for numerical validation.""" + outputs = [] + for t in range(q.shape[2]): + memory = log_decay[:, :, t].exp().unsqueeze(-1) * memory + old = torch.einsum("bhf,bhfv->bhv", erase[:, :, t], memory) + memory = memory + k[:, :, t, :, None] * (z[:, :, t] - old).unsqueeze(-2) + outputs.append(torch.einsum( + "bhf,bhfv->bhv", q[:, :, t], memory.repeat_interleave(groups, dim=1) + )) + return torch.stack(outputs, dim=2), memory + + +def local_window_attention(q, k, v, window, groups): + """Exact causal softmax over at most window keys, including the current token. + + k/v may include up to window-1 cached prefix tokens. Unfold creates views; + scores occupy [B,Hq,T,window], never [B,Hq,T,T]. + """ + k = k.repeat_interleave(groups, dim=1) + v = v.repeat_interleave(groups, dim=1) + total, t = k.shape[2], q.shape[2] + kw = F.pad(k, (0, 0, window - 1, 0)).unfold(2, window, 1) + vw = F.pad(v, (0, 0, window - 1, 0)).unfold(2, window, 1) + kw = kw[:, :, -t:].transpose(-1, -2) + vw = vw[:, :, -t:].transpose(-1, -2) + scores = torch.einsum("bhtd,bhtwd->bhtw", q, kw) / math.sqrt(q.shape[-1]) + positions = torch.arange(total - t, total, device=q.device) + offsets = torch.arange(window, device=q.device) - window + 1 + valid = positions[:, None] + offsets[None, :] >= 0 + weights = scores.masked_fill(~valid[None, None], float("-inf")).softmax(-1) + return torch.einsum("bhtw,bhtwd->bhtd", weights, vw) + + +class ResearchCeNNLayer(nn.Module): + """A GQA-compatible memory array with learned local state updates. + + Inputs are frozen, post-RoPE Q/K and frozen V. K/V heads share memory exactly + as indicated by num_kv_heads. Query feature maps remain head-specific. + """ + + def __init__(self, num_heads, num_kv_heads, head_dim, feature_dim=64, + variant="cenn_delta2", window=32, chunk_size=32): + super().__init__() + if variant not in VARIANTS: + raise ValueError(f"unknown variant: {variant}") + if min(num_heads, num_kv_heads, head_dim, feature_dim, window) < 1: + raise ValueError("dimensions and window must be positive") + if num_heads % num_kv_heads: + raise ValueError("query heads must be divisible by KV heads") + if not 1 <= chunk_size <= 32: + raise ValueError("chunk_size must be in [1,32]") + self.num_heads, self.num_kv_heads = num_heads, num_kv_heads + self.head_dim, self.feature_dim = head_dim, feature_dim + self.groups = num_heads // num_kv_heads + self.variant, self.window, self.chunk_size = variant, window, chunk_size + base = torch.zeros(num_kv_heads, feature_dim, head_dim) + for h in range(num_kv_heads): + if feature_dim == head_dim: + base[h] = torch.eye(head_dim) + else: + nn.init.orthogonal_(base[h]) + self.wk = nn.Parameter(base.clone()) + self.wq = nn.Parameter(base.repeat_interleave(self.groups, dim=0).clone()) + self.forget_w = nn.Parameter(torch.zeros(num_kv_heads, feature_dim, head_dim)) + # log(alpha) in [-0.25,0], initially -0.01. Bounds keep chunk rescaling safe. + self.forget_b = nn.Parameter(torch.full((num_kv_heads, feature_dim), + math.log(0.04 / 0.96))) + erase_dim = 1 if variant == "cenn_kda" else feature_dim + self.erase_w = nn.Parameter(torch.zeros(num_kv_heads, erase_dim, head_dim)) + self.erase_b = nn.Parameter(torch.full((num_kv_heads, erase_dim), -1.0)) + if variant != "cenn_kda": + self.write_w = nn.Parameter(torch.zeros(num_kv_heads, head_dim, head_dim)) + self.write_b = nn.Parameter(torch.full((num_kv_heads, head_dim), -1.0)) + self.log_gain = nn.Parameter(torch.zeros(num_heads)) + if variant == "cenn_delta2_window": + self.mix_w = nn.Parameter(torch.zeros(num_heads, head_dim)) + self.mix_b = nn.Parameter(torch.zeros(num_heads)) + + @property + def config(self): + return dict(num_heads=self.num_heads, num_kv_heads=self.num_kv_heads, + head_dim=self.head_dim, feature_dim=self.feature_dim, + variant=self.variant, window=self.window, chunk_size=self.chunk_size) + + @staticmethod + def _project(x, weight, bias): + return torch.einsum("bhtd,hfd->bhtf", x, weight) + bias[None, :, None] + + def features(self, q, k, v): + qn, kn = F.normalize(q, dim=-1), F.normalize(k, dim=-1) + qp = F.normalize(torch.einsum("bhtd,hfd->bhtf", qn, self.wq), dim=-1) + kp = F.normalize(torch.einsum("bhtd,hfd->bhtf", kn, self.wk), dim=-1) + log_decay = -0.25 * self._project(kn, self.forget_w, self.forget_b).sigmoid() + erase_gate = self._project(kn, self.erase_w, self.erase_b).sigmoid() + if self.variant == "cenn_kda": + write_gate = erase_gate + else: + write_gate = self._project( + F.normalize(v, dim=-1), self.write_w, self.write_b + ).sigmoid() + return qp, kp, v * write_gate, kp * erase_gate, log_decay + + def forward(self, q, k, v, state=None, return_state=False, implementation="chunk"): + if q.ndim != 4 or k.ndim != 4 or v.shape != k.shape: + raise ValueError("expected [batch,heads,time,dim] Q/K/V with equal K/V shapes") + if (q.shape[0] != k.shape[0] or q.shape[2:] != k.shape[2:] + or q.shape[1] != self.num_heads or k.shape[1] != self.num_kv_heads + or q.shape[-1] != self.head_dim or q.shape[2] < 1): + raise ValueError("Q/K/V shapes do not match the layer configuration") + q, k, v = (x.to(self.wq.dtype) for x in (q, k, v)) + b = q.shape[0] + if state is None: + state = DeltaState(q.new_zeros(b, self.num_kv_heads, + self.feature_dim, self.head_dim)) + qp, kp, z, erase, log_decay = self.features(q, k, v) + if implementation == "chunk": + output, memory = delta_recurrence( + qp, kp, z, erase, log_decay, state.memory, self.groups, self.chunk_size + ) + elif implementation == "token": + output, memory = delta_token_reference( + qp, kp, z, erase, log_decay, state.memory, self.groups + ) + else: + raise ValueError("implementation must be chunk or token") + output = output * self.log_gain.clamp(-4, 4).exp()[None, :, None, None] + new_state = DeltaState(memory) + if self.variant == "cenn_delta2_window": + keys = k if state.keys is None else torch.cat((state.keys, k), dim=2) + values = v if state.values is None else torch.cat((state.values, v), dim=2) + local = local_window_attention(q, keys, values, self.window, self.groups) + mix = (torch.einsum("bhtd,hd->bht", F.normalize(q, dim=-1), self.mix_w) + + self.mix_b[None, :, None]).sigmoid().unsqueeze(-1) + output = mix * local + (1.0 - mix) * output + keep = self.window - 1 + # Clone even contiguous slices: a view can retain the entire prefix storage. + new_state.keys = keys[:, :, -keep:].clone() if keep else keys.new_empty( + keys.shape[0], keys.shape[1], 0, keys.shape[-1]) + new_state.values = values[:, :, -keep:].clone() if keep else values.new_empty( + values.shape[0], values.shape[1], 0, values.shape[-1]) + return (output, new_state) if return_state else output + + def recurrent_state_bytes(self, batch_size=1): + # Matrix and cache both use this module's floating-point dtype. + elements = self.num_kv_heads * self.feature_dim * self.head_dim + if self.variant == "cenn_delta2_window": + elements += 2 * (self.window - 1) * self.num_kv_heads * self.head_dim + return batch_size * elements * self.wq.element_size() + + +def softmax_reference(q, k, v, groups): + return F.scaled_dot_product_attention( + q, k.repeat_interleave(groups, dim=1), v.repeat_interleave(groups, dim=1), + dropout_p=0.0, is_causal=True + ) diff --git a/src/tinycenn_lm/sharded_moe.py b/src/tinycenn_lm/sharded_moe.py new file mode 100644 index 0000000000000000000000000000000000000000..598a4eb0a4d1ae0798e578748f639c7d36c8d4e5 --- /dev/null +++ b/src/tinycenn_lm/sharded_moe.py @@ -0,0 +1,400 @@ +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Sequence + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from .cenn import CeNNConfig, CausalDepthwiseNeighborhood, StableRMSNorm +from .modeling import DEFAULT_BASE_MODEL, _get_decoder_layers + + +@dataclass(frozen=True) +class ShardedMoECeNNConfig: + """Parameter-neutral routed FFN built by partitioning one dense CeNN FFN. + + The original dense SwiGLU hidden dimension is split across ``num_shards``. + Therefore all shard parameters together equal one dense FFN (plus a tiny + router and one scalar route-mix parameter), rather than ``num_shards`` full + copies of the FFN. + """ + + hidden_size: int = 192 + kernel_size: int = 3 + expansion: int = 4 + steps: int = 7 + dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64) + rms_norm_eps: float = 1e-5 + dropout: float = 0.0 + num_shards: int = 8 + top_k: int = 2 + router_noise_std: float = 1e-3 + + @property + def dense_inner(self) -> int: + return self.hidden_size * self.expansion + + @property + def shard_inner(self) -> int: + return self.dense_inner // self.num_shards + + def validate(self) -> None: + CeNNConfig( + hidden_size=self.hidden_size, + kernel_size=self.kernel_size, + expansion=self.expansion, + steps=self.steps, + dilations=self.dilations, + rms_norm_eps=self.rms_norm_eps, + dropout=self.dropout, + ).validate() + if self.num_shards < 2: + raise ValueError("num_shards must be >= 2") + if self.dense_inner % self.num_shards: + raise ValueError( + f"dense FFN inner size {self.dense_inner} must divide evenly by " + f"num_shards={self.num_shards}" + ) + if not 1 <= self.top_k <= self.num_shards: + raise ValueError("top_k must be in [1, num_shards]") + if self.router_noise_std < 0: + raise ValueError("router_noise_std must be >= 0") + + def to_dict(self) -> dict: + data = asdict(self) + data["dilations"] = list(self.dilations) + return data + + @classmethod + def from_dict(cls, data: dict) -> "ShardedMoECeNNConfig": + data = dict(data) + if "dilations" in data: + data["dilations"] = tuple(data["dilations"]) + return cls(**data) + + +class TopKShardRouter(nn.Module): + def __init__(self, hidden_size: int, num_shards: int, top_k: int, noise_std: float) -> None: + super().__init__() + self.num_shards = num_shards + self.top_k = top_k + self.proj = nn.Linear(hidden_size, num_shards, bias=False) + nn.init.normal_(self.proj.weight, mean=0.0, std=noise_std) + + def forward(self, x: Tensor) -> tuple[Tensor, Tensor, dict[str, Tensor]]: + logits = self.proj(x).float() + probs = F.softmax(logits, dim=-1) + top_values, top_indices = torch.topk(probs, k=self.top_k, dim=-1) + top_weights = top_values / top_values.sum(dim=-1, keepdim=True).clamp_min(1e-9) + + assignment = F.one_hot(top_indices, num_classes=self.num_shards).float().sum(dim=-2) + assignment = assignment / float(self.top_k) + shard_fraction = assignment.mean(dim=(0, 1)) + probability_fraction = probs.mean(dim=(0, 1)) + load_balance = self.num_shards * torch.sum(shard_fraction * probability_fraction) + z_loss = torch.logsumexp(logits, dim=-1).pow(2).mean() + entropy = -(probs * probs.clamp_min(1e-9).log()).sum(dim=-1).mean() + return top_indices, top_weights.to(dtype=x.dtype), { + "load_balance": load_balance, + "z_loss": z_loss, + "entropy": entropy, + "shard_fraction": shard_fraction, + "probability_fraction": probability_fraction, + } + + +class ShardedSwiGLU(nn.Module): + """One dense SwiGLU split into parameter-neutral channel shards. + + All shards are evaluated to reconstruct the complete dense FFN. Top-k routing + produces a sparse, scaled estimate of that same FFN. A zero-initialized scalar + learns how much of the sparse routed specialization to mix into the complete + FFN. At initialization this module is exactly the original dense FFN. + """ + + def __init__(self, config: ShardedMoECeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + e = config.num_shards + h = config.hidden_size + s = config.shard_inner + + # Vectorized expert-shard weights. Across all 8 shards these contain + # exactly the same number of parameters as one dense SwiGLU FFN. + self.in_weight = nn.Parameter(torch.empty(e, 2 * s, h)) + self.out_weight = nn.Parameter(torch.empty(e, h, s)) + nn.init.kaiming_uniform_(self.in_weight, a=5**0.5) + nn.init.zeros_(self.out_weight) + + self.router = TopKShardRouter(h, e, config.top_k, config.router_noise_std) + self.route_mix = nn.Parameter(torch.zeros(())) + self.dropout = nn.Dropout(config.dropout) + self.last_router_stats: dict[str, Tensor] = {} + + def forward(self, x: Tensor) -> Tensor: + # [B,S,H] -> [B,S,E,2I] + projected = torch.einsum("bsh,eih->bsei", x, self.in_weight) + a, b = projected.chunk(2, dim=-1) + hidden = F.silu(a) * b + # [B,S,E,I] x [E,H,I] -> [B,S,E,H] + shard_outputs = torch.einsum("bsei,ehi->bseh", hidden, self.out_weight) + shard_outputs = self.dropout(shard_outputs) + + # Complete dense-FFN reconstruction from all disjoint shards. + dense_full = shard_outputs.sum(dim=-2) + + top_idx, top_weight, stats = self.router(x) + gather_index = top_idx.unsqueeze(-1).expand(*top_idx.shape, x.shape[-1]) + selected = torch.gather(shard_outputs, dim=-2, index=gather_index) + routed = (selected * top_weight.unsqueeze(-1)).sum(dim=-2) + + # Scale the Top-k estimate to the full shard count. route_mix starts at + # exactly zero, so warm-start output exactly equals the trained dense FFN. + sparse_scaled = routed * (self.config.num_shards / float(self.config.top_k)) + route_delta = sparse_scaled - dense_full + mixed = dense_full + self.route_mix * route_delta + + self.last_router_stats = { + **stats, + "route_mix": self.route_mix, + } + return mixed + + +class ShardedMoESharedCeNNCell(nn.Module): + def __init__(self, config: ShardedMoECeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + self.norm = StableRMSNorm(config.hidden_size, config.rms_norm_eps) + self.neighborhood = CausalDepthwiseNeighborhood(config.hidden_size, config.kernel_size) + self.ffn = ShardedSwiGLU(config) + self.gate_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=True) + nn.init.constant_(self.gate_proj.bias, -1.0) + + def forward(self, state: Tensor, dilation: int, step_scale: float) -> tuple[Tensor, dict[str, Tensor]]: + x = self.norm(state) + local = self.neighborhood(x, dilation=dilation) + update = self.ffn(local) + gate = torch.sigmoid(self.gate_proj(local)) + return state + step_scale * gate * update, self.ffn.last_router_stats + + +class FastShardedMoECeNNCore(nn.Module): + def __init__(self, config: ShardedMoECeNNConfig) -> None: + super().__init__() + config.validate() + self.config = config + self.cell = ShardedMoESharedCeNNCell(config) + self.last_router_stats: dict[str, Tensor] = {} + + def forward(self, hidden_states: Tensor) -> Tensor: + initial = hidden_states + state = hidden_states + step_scale = self.config.steps ** -0.5 + scalar_accum: dict[str, Tensor] = {} + shard_fraction = None + probability_fraction = None + for step in range(self.config.steps): + dilation = self.config.dilations[step % len(self.config.dilations)] + state, stats = self.cell(state, dilation=dilation, step_scale=step_scale) + for key in ("load_balance", "z_loss", "entropy"): + scalar_accum[key] = scalar_accum.get(key, stats[key].new_zeros(())) + stats[key] + shard_fraction = stats["shard_fraction"] if shard_fraction is None else shard_fraction + stats["shard_fraction"] + probability_fraction = stats["probability_fraction"] if probability_fraction is None else probability_fraction + stats["probability_fraction"] + + self.last_router_stats = { + "load_balance": scalar_accum["load_balance"] / self.config.steps, + "z_loss": scalar_accum["z_loss"] / self.config.steps, + "entropy": scalar_accum["entropy"] / self.config.steps, + "shard_fraction": shard_fraction / self.config.steps, + "probability_fraction": probability_fraction / self.config.steps, + "route_mix": self.cell.ffn.route_mix, + } + return state - initial + + @property + def receptive_field(self) -> int: + radius = sum(self.config.dilations[i % len(self.config.dilations)] for i in range(self.config.steps)) + return 1 + (self.config.kernel_size - 1) * radius + + +class ShardedMoECeNNReplacementLayer(nn.Module): + def __init__(self, config: ShardedMoECeNNConfig, *, device=None, dtype=None) -> None: + super().__init__() + self.config = config + self.cenn = FastShardedMoECeNNCore(config) + if device is not None or dtype is not None: + kwargs = {} + if device is not None: + kwargs["device"] = device + if dtype is not None: + kwargs["dtype"] = dtype + self.cenn.to(**kwargs) + + def forward(self, hidden_states: Tensor, *args, **kwargs) -> Tensor: + if kwargs.get("use_cache", False): + raise RuntimeError("Sharded MoE-CeNN student requires use_cache=False") + if kwargs.get("output_attentions", False): + raise RuntimeError("Sharded MoE-CeNN student has no attention matrices") + return hidden_states + self.cenn(hidden_states) + + +def replace_transformer_with_sharded_moe_cenn( + model: nn.Module, + config: ShardedMoECeNNConfig, + layer_indices: Sequence[int] = (0,), +) -> nn.Module: + layers = _get_decoder_layers(model) + if config.hidden_size != int(model.config.hidden_size): + raise ValueError("Sharded MoE-CeNN hidden size does not match base model") + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + for index in layer_indices: + old = layers[index] + reference = next((p for p in old.parameters() if p.is_floating_point()), None) + layers[index] = ShardedMoECeNNReplacementLayer( + config, + device=reference.device if reference is not None else None, + dtype=reference.dtype if reference is not None else None, + ) + return model + + +def freeze_sharded_moe_interfaces(model: nn.Module) -> None: + for parameter in model.parameters(): + parameter.requires_grad = False + for module in model.modules(): + if isinstance(module, ShardedMoECeNNReplacementLayer): + for parameter in module.parameters(): + parameter.requires_grad = True + + +def warmstart_sharded_moe_from_plain_cenn(model: nn.Module, plain_student_dir: str | Path) -> None: + """Slice the trained dense CeNN FFN exactly across 8 routed shards.""" + state = torch.load(Path(plain_student_dir) / "cenn_student.pt", map_location="cpu", weights_only=True) + layer = next((m for m in model.modules() if isinstance(m, ShardedMoECeNNReplacementLayer)), None) + if layer is None: + raise RuntimeError("Sharded MoE-CeNN replacement layer not found") + + # Locate source prefix independently of the wrapper path. + source_prefix = next( + (name.rsplit("in_proj.weight", 1)[0] for name in state if name.endswith(".cenn.cell.in_proj.weight")), + None, + ) + if source_prefix is None: + raise RuntimeError("plain CeNN checkpoint does not contain dense FFN weights") + + source_in = state[source_prefix + "in_proj.weight"] + source_out = state[source_prefix + "out_proj.weight"] + source_norm = state[source_prefix + "norm.weight"] + source_neighborhood = state[source_prefix + "neighborhood.weight"] + source_gate_w = state[source_prefix + "gate_proj.weight"] + source_gate_b = state[source_prefix + "gate_proj.bias"] + + cfg = layer.config + inner = cfg.dense_inner + shard = cfg.shard_inner + with torch.no_grad(): + layer.cenn.cell.norm.weight.copy_(source_norm.to(layer.cenn.cell.norm.weight)) + layer.cenn.cell.neighborhood.weight.copy_(source_neighborhood.to(layer.cenn.cell.neighborhood.weight)) + layer.cenn.cell.gate_proj.weight.copy_(source_gate_w.to(layer.cenn.cell.gate_proj.weight)) + layer.cenn.cell.gate_proj.bias.copy_(source_gate_b.to(layer.cenn.cell.gate_proj.bias)) + for expert_id in range(cfg.num_shards): + lo = expert_id * shard + hi = lo + shard + layer.cenn.cell.ffn.in_weight[expert_id, :shard].copy_( + source_in[lo:hi].to(layer.cenn.cell.ffn.in_weight) + ) + layer.cenn.cell.ffn.in_weight[expert_id, shard:].copy_( + source_in[inner + lo : inner + hi].to(layer.cenn.cell.ffn.in_weight) + ) + layer.cenn.cell.ffn.out_weight[expert_id].copy_( + source_out[:, lo:hi].to(layer.cenn.cell.ffn.out_weight) + ) + layer.cenn.cell.ffn.route_mix.zero_() + + +def sharded_router_stats(model: nn.Module) -> dict[str, Tensor]: + layer = next((m for m in model.modules() if isinstance(m, ShardedMoECeNNReplacementLayer)), None) + if layer is None or not layer.cenn.last_router_stats: + raise RuntimeError("router statistics unavailable; run a forward pass first") + return layer.cenn.last_router_stats + + +def _state_dict(model: nn.Module) -> dict[str, Tensor]: + state = {name: tensor.detach().cpu() for name, tensor in model.state_dict().items() if ".cenn." in name} + if not state: + raise ValueError("no Sharded MoE-CeNN weights found") + return state + + +def save_sharded_moe_student( + model: nn.Module, + output_dir: str | Path, + *, + config: ShardedMoECeNNConfig, + base_model: str = DEFAULT_BASE_MODEL, + layer_indices: Sequence[int] = (0,), + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_state_dict(model), output_dir / "sharded_moe_cenn_student.pt") + metadata = { + "format_version": 1, + "architecture": "sharded-moe-cenn-top2-replacement", + "base_model": base_model, + "layer_indices": list(layer_indices), + "sharded_moe_cenn": config.to_dict(), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "sharded_moe_student_config.json").write_text( + json.dumps(metadata, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_sharded_moe_student_weights(model: nn.Module, student_dir: str | Path, *, map_location="cpu", strict: bool = True) -> nn.Module: + state = torch.load(Path(student_dir) / "sharded_moe_cenn_student.pt", map_location=map_location, weights_only=True) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_state_dict(model)) + missing = [key for key in incompatible.missing_keys if key in expected] + unexpected = [key for key in incompatible.unexpected_keys if key not in expected] + if strict and missing: + raise RuntimeError(f"missing Sharded MoE-CeNN keys: {missing}") + if strict and unexpected: + raise RuntimeError(f"unexpected Sharded MoE-CeNN keys: {unexpected}") + return model + + +def build_sharded_moe_student(student_dir: str | Path, *, device=None, dtype=None, attn_implementation: str = "sdpa"): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + metadata = json.loads((student_dir / "sharded_moe_student_config.json").read_text()) + if metadata.get("architecture") != "sharded-moe-cenn-top2-replacement": + raise ValueError("checkpoint is not a Sharded MoE-CeNN Top-2 student") + kwargs = {"attn_implementation": attn_implementation} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs) + config = ShardedMoECeNNConfig.from_dict(metadata["sharded_moe_cenn"]) + replace_transformer_with_sharded_moe_cenn(model, config, tuple(metadata["layer_indices"])) + load_sharded_moe_student_weights(model, student_dir) + move_kwargs = {} + if device is not None: + move_kwargs["device"] = device + if dtype is not None: + move_kwargs["dtype"] = dtype + if move_kwargs: + model.to(**move_kwargs) + model.config.use_cache = False + return model diff --git a/src/tinycenn_lm/smollm2_amcenn.py b/src/tinycenn_lm/smollm2_amcenn.py new file mode 100644 index 0000000000000000000000000000000000000000..b4437cb7fe122790b28b3ca8489d74466e8de3b2 --- /dev/null +++ b/src/tinycenn_lm/smollm2_amcenn.py @@ -0,0 +1,356 @@ +from __future__ import annotations + +import json +import math +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +DEFAULT_SMOLLM2 = "HuggingFaceTB/SmolLM2-135M" + + +@dataclass(frozen=True) +class SmolAMCeNNConfig: + feature_dim: int = 32 + num_shards: int = 8 + top_k: int = 2 + feature_seed: int = 1234 + eps: float = 1e-6 + + def validate(self, model_config) -> None: + if self.feature_dim < 4: + raise ValueError("feature_dim must be >= 4") + if self.num_shards < 2: + raise ValueError("num_shards must be >= 2") + if not 1 <= self.top_k <= self.num_shards: + raise ValueError("top_k must be in [1, num_shards]") + if int(model_config.intermediate_size) % self.num_shards: + raise ValueError("intermediate_size must divide evenly by num_shards") + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> "SmolAMCeNNConfig": + return cls(**dict(data)) + + +class PositiveSoftmaxFeatures(nn.Module): + """Finite positive random features for exp(q^T k / sqrt(d)). + + With x=q/d^(1/4), y=k/d^(1/4) and omega~N(0,I), + E[phi(x)^T phi(y)] = exp(q^T k / sqrt(d)). + This is the finite-state approximation of the exact infinite-dimensional + recurrent softmax-kernel state discussed in the TinyCeNN derivation. + """ + + def __init__(self, head_dim: int, feature_dim: int, seed: int) -> None: + super().__init__() + g = torch.Generator(device="cpu").manual_seed(seed) + projection = torch.randn(feature_dim, head_dim, generator=g) + self.register_buffer("projection", projection, persistent=True) + self.head_dim = head_dim + self.feature_dim = feature_dim + self.scale = head_dim ** -0.25 + self.log_norm = 0.5 * math.log(float(feature_dim)) + + def forward(self, x: Tensor) -> Tensor: + work = x.float() * self.scale + projected = torch.einsum("...d,fd->...f", work, self.projection.float()) + norm = 0.5 * work.square().sum(dim=-1, keepdim=True) + log_phi = (projected - norm - self.log_norm).clamp(min=-20.0, max=20.0) + return torch.exp(log_phi) + + +class AMCeNNAttention(nn.Module): + """Causal recurrent associative-memory replacement for Llama self-attention. + + S_t = S_{t-1} + phi(k_t) v_t^T + z_t = z_{t-1} + phi(k_t) + a_t = phi(q_t)^T S_t / (phi(q_t)^T z_t + eps) + + Training uses a vectorized causal prefix scan (cumsum), so no T x T attention + matrix is constructed. Generation currently uses use_cache=False and rebuilds + the prefix state; a persistent recurrent cache can be added separately. + """ + + def __init__(self, original_attn: nn.Module, model_config, config: SmolAMCeNNConfig, layer_idx: int) -> None: + super().__init__() + self.hidden_size = int(model_config.hidden_size) + self.num_heads = int(model_config.num_attention_heads) + self.num_key_value_heads = int(model_config.num_key_value_heads) + self.head_dim = self.hidden_size // self.num_heads + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.eps = float(config.eps) + self.layer_idx = layer_idx + + self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False) + self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) + with torch.no_grad(): + self.q_proj.weight.copy_(original_attn.q_proj.weight) + self.k_proj.weight.copy_(original_attn.k_proj.weight) + self.v_proj.weight.copy_(original_attn.v_proj.weight) + self.o_proj.weight.copy_(original_attn.o_proj.weight) + + self.features = PositiveSoftmaxFeatures( + self.head_dim, config.feature_dim, seed=config.feature_seed + layer_idx + ) + self.last_state_norm = torch.tensor(0.0) + + def _apply_rope(self, q: Tensor, k: Tensor, position_embeddings) -> tuple[Tensor, Tensor]: + if position_embeddings is None: + return q, k + try: + from transformers.models.llama.modeling_llama import apply_rotary_pos_emb + + cos, sin = position_embeddings + return apply_rotary_pos_emb(q, k, cos, sin) + except Exception: + return q, k + + def forward( + self, + hidden_states: Tensor, + attention_mask=None, + position_ids=None, + past_key_values=None, + use_cache: bool = False, + cache_position=None, + position_embeddings=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache: + raise RuntimeError("AM-CeNN currently requires use_cache=False") + bsz, seq_len, _ = hidden_states.shape + + q = self.q_proj(hidden_states).view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + k = self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + q, k = self._apply_rope(q, k, position_embeddings) + + phi_q = self.features(q).transpose(1, 2) + phi_k = self.features(k).transpose(1, 2) + values = v.transpose(1, 2).float() + + kv_write = torch.einsum("btkf,btkd->btkfd", phi_k, values) + state_s = kv_write.cumsum(dim=1) + state_z = phi_k.cumsum(dim=1) + + state_s = state_s.repeat_interleave(self.num_key_value_groups, dim=2) + state_z = state_z.repeat_interleave(self.num_key_value_groups, dim=2) + numerator = torch.einsum("bthf,bthfd->bthd", phi_q, state_s) + denominator = torch.einsum("bthf,bthf->bth", phi_q, state_z).unsqueeze(-1) + out = numerator / denominator.clamp_min(self.eps) + out = out.to(dtype=hidden_states.dtype).reshape(bsz, seq_len, self.hidden_size) + out = self.o_proj(out) + self.last_state_norm = state_s[:, -1].detach().float().norm().cpu() + return out, None + + +class ShardedTop2LlamaMLP(nn.Module): + """Exact parameter-neutral 8-way partition of a pretrained Llama SwiGLU FFN. + + All disjoint shards sum to the original dense FFN. Top-2 routing supplies a + learned correction controlled by route_mix. route_mix=0 is exactly the + pretrained dense FFN, so the FFN conversion itself is function preserving. + """ + + def __init__(self, original_mlp: nn.Module, model_config, config: SmolAMCeNNConfig) -> None: + super().__init__() + h = int(model_config.hidden_size) + inner = int(model_config.intermediate_size) + e = config.num_shards + s = inner // e + self.hidden_size = h + self.inner_size = inner + self.num_shards = e + self.shard_inner = s + self.top_k = config.top_k + + self.gate_weight = nn.Parameter(torch.empty(e, s, h)) + self.up_weight = nn.Parameter(torch.empty(e, s, h)) + self.down_weight = nn.Parameter(torch.empty(e, h, s)) + self.router = nn.Linear(h, e, bias=False) + nn.init.normal_(self.router.weight, mean=0.0, std=1e-3) + self.route_mix = nn.Parameter(torch.zeros(())) + + with torch.no_grad(): + for i in range(e): + lo, hi = i * s, (i + 1) * s + self.gate_weight[i].copy_(original_mlp.gate_proj.weight[lo:hi]) + self.up_weight[i].copy_(original_mlp.up_proj.weight[lo:hi]) + self.down_weight[i].copy_(original_mlp.down_proj.weight[:, lo:hi]) + self.last_router_stats: dict[str, Tensor] = {} + + def forward(self, x: Tensor) -> Tensor: + gate = torch.einsum("bth,esh->btes", x, self.gate_weight) + up = torch.einsum("bth,esh->btes", x, self.up_weight) + hidden = F.silu(gate) * up + shard_out = torch.einsum("btes,ehs->bteh", hidden, self.down_weight) + dense_full = shard_out.sum(dim=2) + + logits = self.router(x).float() + probs = F.softmax(logits, dim=-1) + top_values, top_idx = torch.topk(probs, k=self.top_k, dim=-1) + top_weight = top_values / top_values.sum(dim=-1, keepdim=True).clamp_min(1e-9) + gather_idx = top_idx.unsqueeze(-1).expand(*top_idx.shape, self.hidden_size) + selected = torch.gather(shard_out, dim=2, index=gather_idx) + routed = (selected * top_weight.to(selected.dtype).unsqueeze(-1)).sum(dim=2) + sparse_scaled = routed * (self.num_shards / float(self.top_k)) + out = dense_full + self.route_mix * (sparse_scaled - dense_full) + + assignment = F.one_hot(top_idx, num_classes=self.num_shards).float().sum(dim=-2) + assignment = assignment / float(self.top_k) + shard_fraction = assignment.mean(dim=(0, 1)) + probability_fraction = probs.mean(dim=(0, 1)) + self.last_router_stats = { + "load_balance": self.num_shards * torch.sum(shard_fraction * probability_fraction), + "z_loss": torch.logsumexp(logits, dim=-1).pow(2).mean(), + "entropy": -(probs * probs.clamp_min(1e-9).log()).sum(dim=-1).mean(), + "shard_fraction": shard_fraction, + "probability_fraction": probability_fraction, + "route_mix": self.route_mix, + } + return out + + +def replace_smollm2_core(model: nn.Module, config: SmolAMCeNNConfig) -> nn.Module: + config.validate(model.config) + layers = model.model.layers + for idx, layer in enumerate(layers): + if not isinstance(layer.self_attn, AMCeNNAttention): + old_attn = layer.self_attn + new_attn = AMCeNNAttention(old_attn, model.config, config, idx) + new_attn.to(device=old_attn.q_proj.weight.device, dtype=old_attn.q_proj.weight.dtype) + layer.self_attn = new_attn + if not isinstance(layer.mlp, ShardedTop2LlamaMLP): + old_mlp = layer.mlp + new_mlp = ShardedTop2LlamaMLP(old_mlp, model.config, config) + new_mlp.to(device=old_mlp.gate_proj.weight.device, dtype=old_mlp.gate_proj.weight.dtype) + layer.mlp = new_mlp + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def freeze_smollm2_for_amcenn_training(model: nn.Module, *, train_ffn_shards: bool = False) -> None: + for p in model.parameters(): + p.requires_grad = False + for module in model.modules(): + if isinstance(module, AMCeNNAttention): + for p in module.parameters(): + p.requires_grad = True + elif isinstance(module, ShardedTop2LlamaMLP): + module.router.weight.requires_grad = True + module.route_mix.requires_grad = True + if train_ffn_shards: + module.gate_weight.requires_grad = True + module.up_weight.requires_grad = True + module.down_weight.requires_grad = True + + +def amcenn_router_stats(model: nn.Module) -> dict[str, Tensor]: + stats = [m.last_router_stats for m in model.modules() if isinstance(m, ShardedTop2LlamaMLP) and m.last_router_stats] + if not stats: + raise RuntimeError("router statistics unavailable; run a forward pass first") + keys = ("load_balance", "z_loss", "entropy", "route_mix") + out = {key: torch.stack([s[key].float() for s in stats]).mean() for key in keys} + out["shard_fraction"] = torch.stack([s["shard_fraction"].float() for s in stats]).mean(dim=0) + out["probability_fraction"] = torch.stack([s["probability_fraction"].float() for s in stats]).mean(dim=0) + return out + + +def amcenn_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + attention = sum(p.numel() for m in model.modules() if isinstance(m, AMCeNNAttention) for p in m.parameters()) + routers = sum( + m.router.weight.numel() + m.route_mix.numel() + for m in model.modules() + if isinstance(m, ShardedTop2LlamaMLP) + ) + return { + "total": total, + "trainable": trainable, + "attention_replacement": attention, + "router_and_mix": routers, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def _replacement_state(model: nn.Module) -> dict[str, Tensor]: + state = {} + for name, tensor in model.state_dict().items(): + if ".self_attn." in name or ".mlp." in name: + state[name] = tensor.detach().cpu() + if not state: + raise RuntimeError("no AM-CeNN replacement state found") + return state + + +def save_smollm2_amcenn( + model: nn.Module, + output_dir: str | Path, + *, + config: SmolAMCeNNConfig, + base_model: str = DEFAULT_SMOLLM2, + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_replacement_state(model), output_dir / "smollm2_amcenn_top2.pt") + metadata = { + "format_version": 1, + "architecture": "smollm2-amcenn-top2", + "base_model": base_model, + "amcenn": config.to_dict(), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "smollm2_amcenn_config.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8") + return output_dir + + +def load_smollm2_amcenn_weights(model: nn.Module, student_dir: str | Path) -> nn.Module: + state = torch.load(Path(student_dir) / "smollm2_amcenn_top2.pt", map_location="cpu", weights_only=True) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_replacement_state(model)) + missing = [k for k in incompatible.missing_keys if k in expected] + unexpected = [k for k in incompatible.unexpected_keys if k not in expected] + if missing: + raise RuntimeError(f"missing AM-CeNN keys: {missing[:8]}") + if unexpected: + raise RuntimeError(f"unexpected AM-CeNN keys: {unexpected[:8]}") + return model + + +def build_smollm2_amcenn(student_dir: str | Path, *, device=None, dtype=None): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + meta = json.loads((student_dir / "smollm2_amcenn_config.json").read_text()) + if meta.get("architecture") != "smollm2-amcenn-top2": + raise ValueError("checkpoint is not SmolLM2 AM-CeNN Top-2") + kwargs = {} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(meta["base_model"], **kwargs) + config = SmolAMCeNNConfig.from_dict(meta["amcenn"]) + replace_smollm2_core(model, config) + load_smollm2_amcenn_weights(model, student_dir) + if device is not None: + model.to(device) + if dtype is not None: + model.to(dtype=dtype) + model.config.use_cache = False + return model diff --git a/src/tinycenn_lm/smollm2_amcenn_v2.py b/src/tinycenn_lm/smollm2_amcenn_v2.py new file mode 100644 index 0000000000000000000000000000000000000000..aedb09d439d67f53ed729f5bf616950a6ccc30f5 --- /dev/null +++ b/src/tinycenn_lm/smollm2_amcenn_v2.py @@ -0,0 +1,359 @@ +from __future__ import annotations + +import json +import math +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable + +import torch +from torch import Tensor, nn + +from .smollm2_amcenn import DEFAULT_SMOLLM2, ShardedTop2LlamaMLP, SmolAMCeNNConfig + + +@dataclass(frozen=True) +class SmolAMCeNNV2Config: + feature_dim: int = 128 + num_shards: int = 8 + top_k: int = 2 + feature_seed: int = 2026 + eps: float = 1e-6 + antithetic_features: bool = True + learnable_feature_correction: bool = True + + def validate(self, model_config) -> None: + if self.feature_dim < 8: + raise ValueError("feature_dim must be >= 8") + if self.antithetic_features and self.feature_dim % 2: + raise ValueError("antithetic feature_dim must be even") + if self.num_shards < 2: + raise ValueError("num_shards must be >= 2") + if not 1 <= self.top_k <= self.num_shards: + raise ValueError("top_k must be in [1, num_shards]") + if int(model_config.intermediate_size) % self.num_shards: + raise ValueError("intermediate_size must divide evenly by num_shards") + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> "SmolAMCeNNV2Config": + return cls(**dict(data)) + + +class AdaptivePositiveSoftmaxFeatures(nn.Module): + """Positive softmax-kernel features with a zero-init learnable correction. + + The base projection is sampled from N(0,I). In v2 we optionally use + antithetic pairs (omega, -omega) to reduce estimator variance while keeping + Gaussian marginals. delta_projection starts at zero, so training begins from + the mathematical random-feature estimator and can adapt the finite feature + basis to each pretrained layer's actual Q/K distribution. + """ + + def __init__( + self, + head_dim: int, + feature_dim: int, + seed: int, + *, + antithetic: bool = True, + learnable_correction: bool = True, + ) -> None: + super().__init__() + g = torch.Generator(device="cpu").manual_seed(seed) + if antithetic: + half = torch.randn(feature_dim // 2, head_dim, generator=g) + projection = torch.cat([half, -half], dim=0) + else: + projection = torch.randn(feature_dim, head_dim, generator=g) + self.register_buffer("base_projection", projection, persistent=True) + if learnable_correction: + self.delta_projection = nn.Parameter(torch.zeros_like(projection)) + else: + self.register_parameter("delta_projection", None) + self.head_dim = int(head_dim) + self.feature_dim = int(feature_dim) + self.scale = head_dim ** -0.25 + self.log_norm = 0.5 * math.log(float(feature_dim)) + + @property + def projection(self) -> Tensor: + if self.delta_projection is None: + return self.base_projection + return self.base_projection + self.delta_projection + + def forward(self, x: Tensor) -> Tensor: + work = x.float() * self.scale + projected = torch.einsum("...d,fd->...f", work, self.projection.float()) + norm = 0.5 * work.square().sum(dim=-1, keepdim=True) + log_phi = (projected - norm - self.log_norm).clamp(min=-20.0, max=20.0) + return torch.exp(log_phi) + + +class AMCeNNAttentionV2(nn.Module): + """Causal finite-state approximation of softmax attention. + + S_t = S_{t-1} + phi(k_t) v_t^T + z_t = z_{t-1} + phi(k_t) + a_t = phi(q_t)^T S_t / (phi(q_t)^T z_t + eps) + + Unlike v1, the feature map is larger by default (m=128), variance reduced + with antithetic features, and the finite basis can adapt through a zero-init + delta projection. No T x T attention matrix is constructed. + """ + + def __init__(self, original_attn: nn.Module, model_config, config: SmolAMCeNNV2Config, layer_idx: int) -> None: + super().__init__() + self.hidden_size = int(model_config.hidden_size) + self.num_heads = int(model_config.num_attention_heads) + self.num_key_value_heads = int(model_config.num_key_value_heads) + self.head_dim = self.hidden_size // self.num_heads + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.eps = float(config.eps) + self.layer_idx = int(layer_idx) + + self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False) + self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) + with torch.no_grad(): + self.q_proj.weight.copy_(original_attn.q_proj.weight) + self.k_proj.weight.copy_(original_attn.k_proj.weight) + self.v_proj.weight.copy_(original_attn.v_proj.weight) + self.o_proj.weight.copy_(original_attn.o_proj.weight) + + self.features = AdaptivePositiveSoftmaxFeatures( + self.head_dim, + config.feature_dim, + seed=config.feature_seed + self.layer_idx, + antithetic=config.antithetic_features, + learnable_correction=config.learnable_feature_correction, + ) + self.last_state_norm = torch.tensor(0.0) + + def _apply_rope(self, q: Tensor, k: Tensor, position_embeddings) -> tuple[Tensor, Tensor]: + if position_embeddings is None: + return q, k + try: + from transformers.models.llama.modeling_llama import apply_rotary_pos_emb + + cos, sin = position_embeddings + return apply_rotary_pos_emb(q, k, cos, sin) + except Exception: + return q, k + + def forward( + self, + hidden_states: Tensor, + attention_mask=None, + position_ids=None, + past_key_values=None, + use_cache: bool = False, + cache_position=None, + position_embeddings=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache: + raise RuntimeError("AM-CeNN v2 currently requires use_cache=False") + bsz, seq_len, _ = hidden_states.shape + + q = self.q_proj(hidden_states).view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + k = self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + q, k = self._apply_rope(q, k, position_embeddings) + + phi_q = self.features(q).transpose(1, 2) + phi_k = self.features(k).transpose(1, 2) + values = v.transpose(1, 2).float() + + kv_write = torch.einsum("btkf,btkd->btkfd", phi_k, values) + state_s = kv_write.cumsum(dim=1) + state_z = phi_k.cumsum(dim=1) + + state_s = state_s.repeat_interleave(self.num_key_value_groups, dim=2) + state_z = state_z.repeat_interleave(self.num_key_value_groups, dim=2) + numerator = torch.einsum("bthf,bthfd->bthd", phi_q, state_s) + denominator = torch.einsum("bthf,bthf->bth", phi_q, state_z).unsqueeze(-1) + out = numerator / denominator.clamp_min(self.eps) + out = out.to(dtype=hidden_states.dtype).reshape(bsz, seq_len, self.hidden_size) + out = self.o_proj(out) + self.last_state_norm = state_s[:, -1].detach().float().norm().cpu() + return out, None + + +def _moe_config(config: SmolAMCeNNV2Config) -> SmolAMCeNNConfig: + return SmolAMCeNNConfig( + feature_dim=max(8, config.feature_dim), + num_shards=config.num_shards, + top_k=config.top_k, + feature_seed=config.feature_seed, + eps=config.eps, + ) + + +def convert_all_ffns_to_sharded_top2(model: nn.Module, config: SmolAMCeNNV2Config) -> nn.Module: + config.validate(model.config) + moe_cfg = _moe_config(config) + for layer in model.model.layers: + if not isinstance(layer.mlp, ShardedTop2LlamaMLP): + old = layer.mlp + new = ShardedTop2LlamaMLP(old, model.config, moe_cfg) + new.to(device=old.gate_proj.weight.device, dtype=old.gate_proj.weight.dtype) + layer.mlp = new + return model + + +def replace_attention_layers( + model: nn.Module, + config: SmolAMCeNNV2Config, + layer_indices: Iterable[int], +) -> nn.Module: + config.validate(model.config) + for idx in layer_indices: + layer = model.model.layers[int(idx)] + if isinstance(layer.self_attn, AMCeNNAttentionV2): + continue + old = layer.self_attn + new = AMCeNNAttentionV2(old, model.config, config, int(idx)) + new.to(device=old.q_proj.weight.device, dtype=old.q_proj.weight.dtype) + layer.self_attn = new + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def replace_all_smollm2_attention(model: nn.Module, config: SmolAMCeNNV2Config) -> nn.Module: + return replace_attention_layers(model, config, range(int(model.config.num_hidden_layers))) + + +def freeze_for_group_calibration(model: nn.Module, layer_indices: Iterable[int]) -> list[nn.Parameter]: + selected = {int(i) for i in layer_indices} + for p in model.parameters(): + p.requires_grad = False + trainable: list[nn.Parameter] = [] + for module in model.modules(): + if isinstance(module, AMCeNNAttentionV2) and module.layer_idx in selected: + for p in module.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def freeze_for_global_training(model: nn.Module, *, train_router: bool = True) -> list[nn.Parameter]: + for p in model.parameters(): + p.requires_grad = False + trainable: list[nn.Parameter] = [] + for module in model.modules(): + if isinstance(module, AMCeNNAttentionV2): + for p in module.parameters(): + p.requires_grad = True + trainable.append(p) + elif isinstance(module, ShardedTop2LlamaMLP) and train_router: + module.router.weight.requires_grad = True + module.route_mix.requires_grad = True + trainable.extend([module.router.weight, module.route_mix]) + return trainable + + +def v2_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + am = sum(p.numel() for m in model.modules() if isinstance(m, AMCeNNAttentionV2) for p in m.parameters()) + feature_delta = sum( + m.features.delta_projection.numel() + for m in model.modules() + if isinstance(m, AMCeNNAttentionV2) and m.features.delta_projection is not None + ) + routers = sum( + m.router.weight.numel() + m.route_mix.numel() + for m in model.modules() + if isinstance(m, ShardedTop2LlamaMLP) + ) + return { + "total": total, + "trainable": trainable, + "amcenn_attention": am, + "feature_delta": feature_delta, + "router_and_mix": routers, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def _v2_state(model: nn.Module) -> dict[str, Tensor]: + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".self_attn." in name or ".mlp." in name: + state[name] = tensor.detach().cpu() + if not state: + raise RuntimeError("no SmolLM2 AM-CeNN v2 state found") + return state + + +def save_smollm2_amcenn_v2( + model: nn.Module, + output_dir: str | Path, + *, + config: SmolAMCeNNV2Config, + base_model: str = DEFAULT_SMOLLM2, + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_v2_state(model), output_dir / "smollm2_amcenn_v2.pt") + meta = { + "format_version": 2, + "architecture": "smollm2-amcenn-top2-v2", + "base_model": base_model, + "amcenn_v2": config.to_dict(), + } + if extra_metadata: + meta["training"] = extra_metadata + (output_dir / "smollm2_amcenn_v2_config.json").write_text( + json.dumps(meta, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_smollm2_amcenn_v2_weights(model: nn.Module, student_dir: str | Path) -> nn.Module: + state = torch.load(Path(student_dir) / "smollm2_amcenn_v2.pt", map_location="cpu", weights_only=True) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_v2_state(model)) + missing = [k for k in incompatible.missing_keys if k in expected] + unexpected = [k for k in incompatible.unexpected_keys if k not in expected] + if missing: + raise RuntimeError(f"missing AM-CeNN v2 keys: {missing[:8]}") + if unexpected: + raise RuntimeError(f"unexpected AM-CeNN v2 keys: {unexpected[:8]}") + return model + + +def build_smollm2_amcenn_v2(student_dir: str | Path, *, device=None, dtype=None): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + meta = json.loads((student_dir / "smollm2_amcenn_v2_config.json").read_text()) + if meta.get("architecture") != "smollm2-amcenn-top2-v2": + raise ValueError("checkpoint is not SmolLM2 AM-CeNN Top-2 v2") + kwargs = {} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(meta["base_model"], **kwargs) + config = SmolAMCeNNV2Config.from_dict(meta["amcenn_v2"]) + convert_all_ffns_to_sharded_top2(model, config) + replace_all_smollm2_attention(model, config) + load_smollm2_amcenn_v2_weights(model, student_dir) + if device is not None: + model.to(device) + if dtype is not None: + model.to(dtype=dtype) + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model diff --git a/src/tinycenn_lm/smollm2_amcenn_v3.py b/src/tinycenn_lm/smollm2_amcenn_v3.py new file mode 100644 index 0000000000000000000000000000000000000000..a2997722091e2340ea62eafab0e4cc6fd107f7ac --- /dev/null +++ b/src/tinycenn_lm/smollm2_amcenn_v3.py @@ -0,0 +1,406 @@ +from __future__ import annotations + +import json +import math +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable + +import torch +from torch import Tensor, nn + +from .smollm2_amcenn import DEFAULT_SMOLLM2 +from .smollm2_amcenn_v2 import AdaptivePositiveSoftmaxFeatures + + +@dataclass(frozen=True) +class SmolAMCeNNV3Config: + """Hybrid exact-local + recurrent-global attention configuration.""" + + feature_dim: int = 256 + local_window: int = 32 + feature_seed: int = 3030 + eps: float = 1e-6 + antithetic_features: bool = True + learnable_feature_correction: bool = True + global_gate_init: float = 0.05 + + def validate(self, model_config) -> None: + if self.feature_dim < 16: + raise ValueError("feature_dim must be >= 16") + if self.antithetic_features and self.feature_dim % 2: + raise ValueError("antithetic feature_dim must be even") + if self.local_window < 1: + raise ValueError("local_window must be >= 1") + if not 0.0 < self.global_gate_init < 1.0: + raise ValueError("global_gate_init must be in (0, 1)") + hidden = int(model_config.hidden_size) + heads = int(model_config.num_attention_heads) + if hidden % heads: + raise ValueError("hidden_size must be divisible by num_attention_heads") + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> "SmolAMCeNNV3Config": + return cls(**dict(data)) + + +def _logit(probability: float) -> float: + p = min(max(float(probability), 1e-6), 1.0 - 1e-6) + return math.log(p / (1.0 - p)) + + +class HybridLocalAMCeNNAttention(nn.Module): + """Exact local causal attention plus gated recurrent AM-CeNN long memory. + + Recent tokens use exact softmax attention inside ``local_window``. Older tokens + are summarized by the positive-feature recurrent state used by AM-CeNN. A + per-head sigmoid gate, initialized close to zero, mixes the recurrent branch + into the exact-local branch. Q/K/V/O weights are copied from the pretrained + attention module. + + During calibration the copied projections stay frozen; only the finite-feature + correction and global gate learn. Global distillation can later update Q/K/V/O + with a much smaller learning rate. + """ + + def __init__( + self, + original_attn: nn.Module, + model_config, + config: SmolAMCeNNV3Config, + layer_idx: int, + ) -> None: + super().__init__() + config.validate(model_config) + self.hidden_size = int(model_config.hidden_size) + self.num_heads = int(model_config.num_attention_heads) + self.num_key_value_heads = int(model_config.num_key_value_heads) + self.head_dim = self.hidden_size // self.num_heads + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.local_window = int(config.local_window) + self.eps = float(config.eps) + self.layer_idx = int(layer_idx) + + self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False) + self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) + with torch.no_grad(): + self.q_proj.weight.copy_(original_attn.q_proj.weight) + self.k_proj.weight.copy_(original_attn.k_proj.weight) + self.v_proj.weight.copy_(original_attn.v_proj.weight) + self.o_proj.weight.copy_(original_attn.o_proj.weight) + + self.features = AdaptivePositiveSoftmaxFeatures( + self.head_dim, + config.feature_dim, + seed=config.feature_seed + self.layer_idx, + antithetic=config.antithetic_features, + learnable_correction=config.learnable_feature_correction, + ) + self.global_gate_logit = nn.Parameter( + torch.full((self.num_heads,), _logit(config.global_gate_init), dtype=torch.float32) + ) + self.last_global_gate = torch.tensor(float(config.global_gate_init)) + self.last_global_state_norm = torch.tensor(0.0) + + def _apply_rope(self, q: Tensor, k: Tensor, position_embeddings) -> tuple[Tensor, Tensor]: + if position_embeddings is None: + return q, k + try: + from transformers.models.llama.modeling_llama import apply_rotary_pos_emb + + cos, sin = position_embeddings + return apply_rotary_pos_emb(q, k, cos, sin) + except Exception: + return q, k + + def _local_exact(self, q: Tensor, k: Tensor, v: Tensor, attention_mask) -> Tensor: + k_heads = k.repeat_interleave(self.num_key_value_groups, dim=1) + v_heads = v.repeat_interleave(self.num_key_value_groups, dim=1) + scores = torch.einsum("bhtd,bhsd->bhts", q.float(), k_heads.float()) + scores = scores / math.sqrt(float(self.head_dim)) + + seq_len = q.shape[-2] + positions = torch.arange(seq_len, device=q.device) + qpos = positions[:, None] + kpos = positions[None, :] + local = (kpos <= qpos) & (kpos >= (qpos - self.local_window + 1)) + scores = scores.masked_fill(~local.view(1, 1, seq_len, seq_len), float("-inf")) + + if torch.is_tensor(attention_mask): + mask = attention_mask + try: + if mask.ndim == 4 and mask.shape[-2:] == (seq_len, seq_len): + scores = scores + mask.to(device=scores.device, dtype=scores.dtype) + elif mask.ndim == 2 and mask.shape[-1] == seq_len: + valid = mask.to(device=scores.device).bool().view(mask.shape[0], 1, 1, seq_len) + scores = scores.masked_fill(~valid, float("-inf")) + except Exception: + pass + + probs = torch.softmax(scores, dim=-1, dtype=torch.float32) + return torch.einsum("bhts,bhsd->bhtd", probs, v_heads.float()) + + def _global_old_memory(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]: + phi_q = self.features(q).transpose(1, 2) + phi_k = self.features(k).transpose(1, 2) + values = v.transpose(1, 2).float() + + writes = torch.einsum("btkf,btkd->btkfd", phi_k, values) + full_s = writes.cumsum(dim=1) + full_z = phi_k.cumsum(dim=1) + + old_s = torch.zeros_like(full_s) + old_z = torch.zeros_like(full_z) + if q.shape[-2] > self.local_window: + old_s[:, self.local_window:] = full_s[:, :-self.local_window] + old_z[:, self.local_window:] = full_z[:, :-self.local_window] + + old_s_h = old_s.repeat_interleave(self.num_key_value_groups, dim=2) + old_z_h = old_z.repeat_interleave(self.num_key_value_groups, dim=2) + numerator = torch.einsum("bthf,bthfd->bthd", phi_q, old_s_h) + denominator = torch.einsum("bthf,bthf->bth", phi_q, old_z_h).unsqueeze(-1) + global_out = numerator / denominator.clamp_min(self.eps) + valid = denominator > self.eps + global_out = torch.where(valid, global_out, torch.zeros_like(global_out)) + return global_out.transpose(1, 2), valid.transpose(1, 2) + + def forward( + self, + hidden_states: Tensor, + attention_mask=None, + position_ids=None, + past_key_values=None, + use_cache: bool = False, + cache_position=None, + position_embeddings=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache: + raise RuntimeError("AM-CeNN v3 currently requires use_cache=False") + + bsz, seq_len, _ = hidden_states.shape + q = self.q_proj(hidden_states).view( + bsz, seq_len, self.num_heads, self.head_dim + ).transpose(1, 2) + k = self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + q, k = self._apply_rope(q, k, position_embeddings) + + local_out = self._local_exact(q, k, v, attention_mask) + global_out, global_valid = self._global_old_memory(q, k, v) + + gate = torch.sigmoid(self.global_gate_logit.float()).view(1, self.num_heads, 1, 1) + effective_gate = gate * global_valid.to(dtype=gate.dtype) + mixed = local_out.float() + effective_gate * (global_out.float() - local_out.float()) + mixed = mixed.transpose(1, 2).contiguous().view(bsz, seq_len, self.hidden_size) + out = self.o_proj(mixed.to(dtype=hidden_states.dtype)) + + self.last_global_gate = gate.detach().mean().cpu() + self.last_global_state_norm = global_out[:, :, -1].detach().float().norm().cpu() + return out, None + + +def replace_attention_layers_v3( + model: nn.Module, + config: SmolAMCeNNV3Config, + layer_indices: Iterable[int], +) -> nn.Module: + config.validate(model.config) + for idx in layer_indices: + layer = model.model.layers[int(idx)] + if isinstance(layer.self_attn, HybridLocalAMCeNNAttention): + continue + old = layer.self_attn + new = HybridLocalAMCeNNAttention(old, model.config, config, int(idx)) + new.to(device=old.q_proj.weight.device, dtype=old.q_proj.weight.dtype) + new.global_gate_logit.data = new.global_gate_logit.data.float() + layer.self_attn = new + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def replace_all_attention_v3(model: nn.Module, config: SmolAMCeNNV3Config) -> nn.Module: + return replace_attention_layers_v3(model, config, range(int(model.config.num_hidden_layers))) + + +def freeze_for_v3_calibration( + model: nn.Module, + layer_indices: Iterable[int], +) -> list[nn.Parameter]: + selected = {int(i) for i in layer_indices} + for p in model.parameters(): + p.requires_grad = False + + trainable: list[nn.Parameter] = [] + for module in model.modules(): + if isinstance(module, HybridLocalAMCeNNAttention) and module.layer_idx in selected: + if module.features.delta_projection is not None: + module.features.delta_projection.requires_grad = True + trainable.append(module.features.delta_projection) + module.global_gate_logit.requires_grad = True + trainable.append(module.global_gate_logit) + return trainable + + +def v3_global_parameter_groups( + model: nn.Module, + *, + main_lr: float, + qkvo_lr: float, + weight_decay: float = 0.01, +) -> tuple[list[dict], list[nn.Parameter]]: + for p in model.parameters(): + p.requires_grad = False + + main: list[nn.Parameter] = [] + qkvo: list[nn.Parameter] = [] + for module in model.modules(): + if not isinstance(module, HybridLocalAMCeNNAttention): + continue + if module.features.delta_projection is not None: + module.features.delta_projection.requires_grad = True + main.append(module.features.delta_projection) + module.global_gate_logit.requires_grad = True + main.append(module.global_gate_logit) + for projection in (module.q_proj, module.k_proj, module.v_proj, module.o_proj): + projection.weight.requires_grad = True + qkvo.append(projection.weight) + + groups = [ + {"params": main, "lr": float(main_lr), "weight_decay": float(weight_decay)}, + {"params": qkvo, "lr": float(qkvo_lr), "weight_decay": float(weight_decay)}, + ] + return groups, [*main, *qkvo] + + +def v3_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + attention = sum( + p.numel() + for module in model.modules() + if isinstance(module, HybridLocalAMCeNNAttention) + for p in module.parameters() + ) + feature_delta = sum( + module.features.delta_projection.numel() + for module in model.modules() + if isinstance(module, HybridLocalAMCeNNAttention) + and module.features.delta_projection is not None + ) + gates = sum( + module.global_gate_logit.numel() + for module in model.modules() + if isinstance(module, HybridLocalAMCeNNAttention) + ) + return { + "total": total, + "trainable": trainable, + "hybrid_attention": attention, + "feature_delta": feature_delta, + "global_gates": gates, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def v3_attention_stats(model: nn.Module) -> dict[str, float]: + modules = [m for m in model.modules() if isinstance(m, HybridLocalAMCeNNAttention)] + if not modules: + raise RuntimeError("AM-CeNN v3 attention modules unavailable") + gates = [float(torch.sigmoid(m.global_gate_logit.detach().float()).mean().cpu()) for m in modules] + return { + "mean_global_gate": sum(gates) / len(gates), + "min_global_gate": min(gates), + "max_global_gate": max(gates), + } + + +def _v3_state(model: nn.Module) -> dict[str, Tensor]: + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".self_attn." in name: + state[name] = tensor.detach().cpu() + if not state: + raise RuntimeError("no AM-CeNN v3 state found") + return state + + +def save_smollm2_amcenn_v3( + model: nn.Module, + output_dir: str | Path, + *, + config: SmolAMCeNNV3Config, + base_model: str = DEFAULT_SMOLLM2, + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_v3_state(model), output_dir / "smollm2_amcenn_v3.pt") + metadata = { + "format_version": 3, + "architecture": "smollm2-amcenn-hybrid-v3", + "base_model": base_model, + "amcenn_v3": config.to_dict(), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "smollm2_amcenn_v3_config.json").write_text( + json.dumps(metadata, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_smollm2_amcenn_v3_weights(model: nn.Module, student_dir: str | Path) -> nn.Module: + state = torch.load( + Path(student_dir) / "smollm2_amcenn_v3.pt", + map_location="cpu", + weights_only=True, + ) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_v3_state(model)) + missing = [key for key in incompatible.missing_keys if key in expected] + unexpected = [key for key in incompatible.unexpected_keys if key not in expected] + if missing: + raise RuntimeError(f"missing AM-CeNN v3 keys: {missing[:8]}") + if unexpected: + raise RuntimeError(f"unexpected AM-CeNN v3 keys: {unexpected[:8]}") + return model + + +def build_smollm2_amcenn_v3(student_dir: str | Path, *, device=None, dtype=None): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + metadata = json.loads((student_dir / "smollm2_amcenn_v3_config.json").read_text()) + if metadata.get("architecture") != "smollm2-amcenn-hybrid-v3": + raise ValueError("checkpoint is not SmolLM2 AM-CeNN hybrid v3") + + kwargs = {} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs) + config = SmolAMCeNNV3Config.from_dict(metadata["amcenn_v3"]) + replace_all_attention_v3(model, config) + load_smollm2_amcenn_v3_weights(model, student_dir) + if device is not None: + model.to(device) + if dtype is not None: + model.to(dtype=dtype) + for module in model.modules(): + if isinstance(module, HybridLocalAMCeNNAttention): + module.global_gate_logit.data = module.global_gate_logit.data.float() + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model diff --git a/src/tinycenn_lm/smollm2_amcenn_v4.py b/src/tinycenn_lm/smollm2_amcenn_v4.py new file mode 100644 index 0000000000000000000000000000000000000000..75ec410a94cc2260932e208ec3bf1a6e251501c2 --- /dev/null +++ b/src/tinycenn_lm/smollm2_amcenn_v4.py @@ -0,0 +1,519 @@ +from __future__ import annotations + +import json +import math +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from .smollm2_amcenn import DEFAULT_SMOLLM2 +from .smollm2_amcenn_v2 import AdaptivePositiveSoftmaxFeatures + + +@dataclass(frozen=True) +class LayerProfileV4: + tier: str + local_window: int + feature_dim: int + gate_init: float + gate_cap: float + + def to_dict(self) -> dict: + return asdict(self) + + +@dataclass(frozen=True) +class SmolAMCeNNV4Config: + """Layer-adaptive exact-local/anchor + recurrent-global attention. + + The default profile comes directly from the v3 calibration pattern: + easy layers keep a smaller exact window, difficult layers get a larger exact + window and a larger AM-CeNN feature map, and the two most difficult layers get + a 96-token exact window plus 512 recurrent features. + """ + + easy_feature_dim: int = 256 + medium_feature_dim: int = 320 + hard_feature_dim: int = 384 + critical_feature_dim: int = 512 + easy_window: int = 32 + medium_window: int = 48 + hard_window: int = 64 + critical_window: int = 96 + anchor_tokens: int = 8 + feature_seed: int = 4040 + eps: float = 1e-6 + antithetic_features: bool = True + learnable_feature_correction: bool = True + gate_init: float = 0.03 + easy_gate_cap: float = 0.25 + medium_gate_cap: float = 0.20 + hard_gate_cap: float = 0.15 + critical_gate_cap: float = 0.10 + easy_layers: tuple[int, ...] = (0, 1, 2, 10, 11, 24, 29) + hard_layers: tuple[int, ...] = (6, 7, 8, 14, 17, 23, 26) + critical_layers: tuple[int, ...] = (18, 20) + + def validate(self, model_config) -> None: + dims = ( + self.easy_feature_dim, + self.medium_feature_dim, + self.hard_feature_dim, + self.critical_feature_dim, + ) + if min(dims) < 16: + raise ValueError("all v4 feature dimensions must be >= 16") + if self.antithetic_features and any(dim % 2 for dim in dims): + raise ValueError("antithetic v4 feature dimensions must be even") + windows = ( + self.easy_window, + self.medium_window, + self.hard_window, + self.critical_window, + ) + if min(windows) < 1: + raise ValueError("all v4 local windows must be >= 1") + if self.anchor_tokens < 0: + raise ValueError("anchor_tokens must be >= 0") + caps = ( + self.easy_gate_cap, + self.medium_gate_cap, + self.hard_gate_cap, + self.critical_gate_cap, + ) + if not 0.0 < self.gate_init < 1.0: + raise ValueError("gate_init must be in (0, 1)") + if any(not 0.0 < cap < 1.0 for cap in caps): + raise ValueError("all v4 gate caps must be in (0, 1)") + if any(self.gate_init >= cap for cap in caps): + raise ValueError("gate_init must be lower than every v4 gate cap") + hidden = int(model_config.hidden_size) + heads = int(model_config.num_attention_heads) + if hidden % heads: + raise ValueError("hidden_size must be divisible by num_attention_heads") + + def profile_for_layer(self, layer_idx: int) -> LayerProfileV4: + idx = int(layer_idx) + if idx in self.critical_layers: + return LayerProfileV4( + "critical", self.critical_window, self.critical_feature_dim, + self.gate_init, self.critical_gate_cap, + ) + if idx in self.hard_layers: + return LayerProfileV4( + "hard", self.hard_window, self.hard_feature_dim, + self.gate_init, self.hard_gate_cap, + ) + if idx in self.easy_layers: + return LayerProfileV4( + "easy", self.easy_window, self.easy_feature_dim, + self.gate_init, self.easy_gate_cap, + ) + return LayerProfileV4( + "medium", self.medium_window, self.medium_feature_dim, + self.gate_init, self.medium_gate_cap, + ) + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> "SmolAMCeNNV4Config": + values = dict(data) + for key in ("easy_layers", "hard_layers", "critical_layers"): + if key in values: + values[key] = tuple(int(v) for v in values[key]) + return cls(**values) + + +def _logit(probability: float) -> float: + p = min(max(float(probability), 1e-6), 1.0 - 1e-6) + return math.log(p / (1.0 - p)) + + +class AdaptiveHybridAMCeNNAttentionV4(nn.Module): + """Exact anchors + exact local window + token-dependent recurrent memory. + + For query position t, exact softmax covers: + * the first ``anchor_tokens`` tokens (causally masked), and + * the most recent ``local_window`` tokens. + + AM-CeNN summarizes only the older middle tokens not already covered exactly. + A token-dependent, per-head gate chooses how much of that recurrent memory to + use. The gate is intrinsically capped per layer tier, so difficult layers can + never over-rely on a poor recurrent approximation. + """ + + def __init__( + self, + original_attn: nn.Module, + model_config, + config: SmolAMCeNNV4Config, + layer_idx: int, + ) -> None: + super().__init__() + config.validate(model_config) + self.hidden_size = int(model_config.hidden_size) + self.num_heads = int(model_config.num_attention_heads) + self.num_key_value_heads = int(model_config.num_key_value_heads) + self.head_dim = self.hidden_size // self.num_heads + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.layer_idx = int(layer_idx) + self.anchor_tokens = int(config.anchor_tokens) + self.eps = float(config.eps) + self.profile = config.profile_for_layer(self.layer_idx) + self.local_window = int(self.profile.local_window) + self.feature_dim = int(self.profile.feature_dim) + self.gate_cap = float(self.profile.gate_cap) + self.tier = self.profile.tier + + self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False) + self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) + self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) + with torch.no_grad(): + self.q_proj.weight.copy_(original_attn.q_proj.weight) + self.k_proj.weight.copy_(original_attn.k_proj.weight) + self.v_proj.weight.copy_(original_attn.v_proj.weight) + self.o_proj.weight.copy_(original_attn.o_proj.weight) + + self.features = AdaptivePositiveSoftmaxFeatures( + self.head_dim, + self.feature_dim, + seed=config.feature_seed + self.layer_idx, + antithetic=config.antithetic_features, + learnable_correction=config.learnable_feature_correction, + ) + + # gate = gate_cap * sigmoid(W h + b). W starts at zero, therefore v4 + # begins as a safe constant-gate model and can become token-adaptive. + self.gate_proj = nn.Linear(self.hidden_size, self.num_heads, bias=True) + nn.init.zeros_(self.gate_proj.weight) + initial_fraction = float(self.profile.gate_init) / self.gate_cap + nn.init.constant_(self.gate_proj.bias, _logit(initial_fraction)) + + self.last_mean_gate = torch.tensor(float(self.profile.gate_init)) + self.last_max_gate = torch.tensor(float(self.profile.gate_init)) + self.last_global_state_norm = torch.tensor(0.0) + + def _apply_rope(self, q: Tensor, k: Tensor, position_embeddings) -> tuple[Tensor, Tensor]: + if position_embeddings is None: + return q, k + try: + from transformers.models.llama.modeling_llama import apply_rotary_pos_emb + + cos, sin = position_embeddings + return apply_rotary_pos_emb(q, k, cos, sin) + except Exception: + return q, k + + def _exact_attention(self, q: Tensor, k: Tensor, v: Tensor, attention_mask) -> Tensor: + k_heads = k.repeat_interleave(self.num_key_value_groups, dim=1) + v_heads = v.repeat_interleave(self.num_key_value_groups, dim=1) + scores = torch.einsum("bhtd,bhsd->bhts", q.float(), k_heads.float()) + scores = scores / math.sqrt(float(self.head_dim)) + + seq_len = q.shape[-2] + positions = torch.arange(seq_len, device=q.device) + qpos = positions[:, None] + kpos = positions[None, :] + causal = kpos <= qpos + local = kpos >= (qpos - self.local_window + 1) + anchors = kpos < min(self.anchor_tokens, seq_len) + allowed = causal & (local | anchors) + scores = scores.masked_fill(~allowed.view(1, 1, seq_len, seq_len), float("-inf")) + + if torch.is_tensor(attention_mask): + mask = attention_mask + try: + if mask.ndim == 4 and mask.shape[-2:] == (seq_len, seq_len): + scores = scores + mask.to(device=scores.device, dtype=scores.dtype) + elif mask.ndim == 2 and mask.shape[-1] == seq_len: + valid = mask.to(device=scores.device).bool().view(mask.shape[0], 1, 1, seq_len) + scores = scores.masked_fill(~valid, float("-inf")) + except Exception: + pass + + probs = torch.softmax(scores, dim=-1, dtype=torch.float32) + return torch.einsum("bhts,bhsd->bhtd", probs, v_heads.float()) + + def _global_middle_memory(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]: + phi_q = self.features(q).transpose(1, 2) # B,T,H,F or B,T,KV,F after repeat below + phi_k = self.features(k).transpose(1, 2) # B,T,KV,F + values = v.transpose(1, 2).float() # B,T,KV,D + + writes = torch.einsum("btkf,btkd->btkfd", phi_k, values) + if self.anchor_tokens > 0: + writes = writes.clone() + phi_k = phi_k.clone() + cutoff = min(self.anchor_tokens, writes.shape[1]) + writes[:, :cutoff] = 0 + phi_k[:, :cutoff] = 0 + + prefix_s = writes.cumsum(dim=1) + prefix_z = phi_k.cumsum(dim=1) + old_s = torch.zeros_like(prefix_s) + old_z = torch.zeros_like(prefix_z) + if q.shape[-2] > self.local_window: + old_s[:, self.local_window:] = prefix_s[:, :-self.local_window] + old_z[:, self.local_window:] = prefix_z[:, :-self.local_window] + + old_s_h = old_s.repeat_interleave(self.num_key_value_groups, dim=2) + old_z_h = old_z.repeat_interleave(self.num_key_value_groups, dim=2) + numerator = torch.einsum("bthf,bthfd->bthd", phi_q, old_s_h) + denominator = torch.einsum("bthf,bthf->bth", phi_q, old_z_h).unsqueeze(-1) + global_out = numerator / denominator.clamp_min(self.eps) + valid = denominator > self.eps + global_out = torch.where(valid, global_out, torch.zeros_like(global_out)) + return global_out.transpose(1, 2), valid.transpose(1, 2) + + def _token_gate(self, hidden_states: Tensor, global_valid: Tensor) -> Tensor: + logits = F.linear( + hidden_states.float(), + self.gate_proj.weight.float(), + self.gate_proj.bias.float(), + ) + gate = self.gate_cap * torch.sigmoid(logits) + gate = gate.transpose(1, 2).unsqueeze(-1) # B,H,T,1 + return gate * global_valid.to(dtype=gate.dtype) + + def forward( + self, + hidden_states: Tensor, + attention_mask=None, + position_ids=None, + past_key_values=None, + use_cache: bool = False, + cache_position=None, + position_embeddings=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache: + raise RuntimeError("AM-CeNN v4 currently requires use_cache=False") + + bsz, seq_len, _ = hidden_states.shape + q = self.q_proj(hidden_states).view( + bsz, seq_len, self.num_heads, self.head_dim + ).transpose(1, 2) + k = self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + q, k = self._apply_rope(q, k, position_embeddings) + + exact_out = self._exact_attention(q, k, v, attention_mask) + global_out, global_valid = self._global_middle_memory(q, k, v) + gate = self._token_gate(hidden_states, global_valid) + + mixed = exact_out.float() + gate * (global_out.float() - exact_out.float()) + mixed = mixed.transpose(1, 2).contiguous().view(bsz, seq_len, self.hidden_size) + out = self.o_proj(mixed.to(dtype=hidden_states.dtype)) + + valid_gates = gate.detach().float()[global_valid.expand_as(gate)] + if valid_gates.numel(): + self.last_mean_gate = valid_gates.mean().cpu() + self.last_max_gate = valid_gates.max().cpu() + else: + self.last_mean_gate = torch.tensor(0.0) + self.last_max_gate = torch.tensor(0.0) + self.last_global_state_norm = global_out[:, :, -1].detach().float().norm().cpu() + return out, None + + +def replace_attention_layers_v4( + model: nn.Module, + config: SmolAMCeNNV4Config, + layer_indices: Iterable[int], +) -> nn.Module: + config.validate(model.config) + for idx in layer_indices: + layer = model.model.layers[int(idx)] + if isinstance(layer.self_attn, AdaptiveHybridAMCeNNAttentionV4): + continue + old = layer.self_attn + new = AdaptiveHybridAMCeNNAttentionV4(old, model.config, config, int(idx)) + new.to(device=old.q_proj.weight.device, dtype=old.q_proj.weight.dtype) + layer.self_attn = new + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def replace_all_attention_v4(model: nn.Module, config: SmolAMCeNNV4Config) -> nn.Module: + return replace_attention_layers_v4(model, config, range(int(model.config.num_hidden_layers))) + + +def freeze_for_v4_calibration(model: nn.Module, layer_indices: Iterable[int]) -> list[nn.Parameter]: + selected = {int(i) for i in layer_indices} + for p in model.parameters(): + p.requires_grad = False + trainable: list[nn.Parameter] = [] + for module in model.modules(): + if not isinstance(module, AdaptiveHybridAMCeNNAttentionV4) or module.layer_idx not in selected: + continue + if module.features.delta_projection is not None: + module.features.delta_projection.requires_grad = True + trainable.append(module.features.delta_projection) + for p in module.gate_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def v4_global_parameter_groups( + model: nn.Module, + *, + memory_lr: float, + qkvo_lr: float, + weight_decay: float = 0.01, +) -> tuple[list[dict], list[nn.Parameter]]: + for p in model.parameters(): + p.requires_grad = False + memory: list[nn.Parameter] = [] + qkvo: list[nn.Parameter] = [] + for module in model.modules(): + if not isinstance(module, AdaptiveHybridAMCeNNAttentionV4): + continue + if module.features.delta_projection is not None: + module.features.delta_projection.requires_grad = True + memory.append(module.features.delta_projection) + for p in module.gate_proj.parameters(): + p.requires_grad = True + memory.append(p) + for projection in (module.q_proj, module.k_proj, module.v_proj, module.o_proj): + projection.weight.requires_grad = True + qkvo.append(projection.weight) + groups = [ + {"params": memory, "lr": float(memory_lr), "weight_decay": float(weight_decay)}, + {"params": qkvo, "lr": float(qkvo_lr), "weight_decay": float(weight_decay)}, + ] + return groups, [*memory, *qkvo] + + +def v4_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + modules = [m for m in model.modules() if isinstance(m, AdaptiveHybridAMCeNNAttentionV4)] + feature_delta = sum( + m.features.delta_projection.numel() + for m in modules + if m.features.delta_projection is not None + ) + gate_params = sum(p.numel() for m in modules for p in m.gate_proj.parameters()) + return { + "total": total, + "trainable": trainable, + "hybrid_attention": sum(p.numel() for m in modules for p in m.parameters()), + "feature_delta": feature_delta, + "token_gate": gate_params, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def v4_attention_stats(model: nn.Module) -> dict: + modules = [m for m in model.modules() if isinstance(m, AdaptiveHybridAMCeNNAttentionV4)] + if not modules: + raise RuntimeError("AM-CeNN v4 attention modules unavailable") + means = [float(m.last_mean_gate) for m in modules] + maxima = [float(m.last_max_gate) for m in modules] + by_tier: dict[str, list[float]] = {} + layer_profiles = [] + for m, mean_gate in zip(modules, means): + by_tier.setdefault(m.tier, []).append(mean_gate) + layer_profiles.append({ + "layer": m.layer_idx, + "tier": m.tier, + "local_window": m.local_window, + "feature_dim": m.feature_dim, + "gate_cap": m.gate_cap, + "mean_gate": mean_gate, + }) + return { + "mean_global_gate": sum(means) / len(means), + "max_observed_gate": max(maxima), + "mean_gate_by_tier": {k: sum(v) / len(v) for k, v in by_tier.items()}, + "layer_profiles": layer_profiles, + } + + +def _v4_state(model: nn.Module) -> dict[str, Tensor]: + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".self_attn." in name: + state[name] = tensor.detach().cpu() + if not state: + raise RuntimeError("no AM-CeNN v4 state found") + return state + + +def save_smollm2_amcenn_v4( + model: nn.Module, + output_dir: str | Path, + *, + config: SmolAMCeNNV4Config, + base_model: str = DEFAULT_SMOLLM2, + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_v4_state(model), output_dir / "smollm2_amcenn_v4.pt") + metadata = { + "format_version": 4, + "architecture": "smollm2-amcenn-adaptive-v4", + "base_model": base_model, + "amcenn_v4": config.to_dict(), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "smollm2_amcenn_v4_config.json").write_text( + json.dumps(metadata, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_smollm2_amcenn_v4_weights(model: nn.Module, student_dir: str | Path) -> nn.Module: + state = torch.load( + Path(student_dir) / "smollm2_amcenn_v4.pt", + map_location="cpu", + weights_only=True, + ) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_v4_state(model)) + missing = [k for k in incompatible.missing_keys if k in expected] + unexpected = [k for k in incompatible.unexpected_keys if k not in expected] + if missing: + raise RuntimeError(f"missing AM-CeNN v4 keys: {missing[:8]}") + if unexpected: + raise RuntimeError(f"unexpected AM-CeNN v4 keys: {unexpected[:8]}") + return model + + +def build_smollm2_amcenn_v4(student_dir: str | Path, *, device=None, dtype=None): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + metadata = json.loads((student_dir / "smollm2_amcenn_v4_config.json").read_text()) + if metadata.get("architecture") != "smollm2-amcenn-adaptive-v4": + raise ValueError("checkpoint is not SmolLM2 AM-CeNN adaptive v4") + kwargs = {} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs) + config = SmolAMCeNNV4Config.from_dict(metadata["amcenn_v4"]) + replace_all_attention_v4(model, config) + load_smollm2_amcenn_v4_weights(model, student_dir) + if device is not None: + model.to(device) + if dtype is not None: + model.to(dtype=dtype) + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model diff --git a/src/tinycenn_lm/smollm2_memory_fusion.py b/src/tinycenn_lm/smollm2_memory_fusion.py new file mode 100644 index 0000000000000000000000000000000000000000..04095075f70feaa8c11fd4cf0d3e3124611e04db --- /dev/null +++ b/src/tinycenn_lm/smollm2_memory_fusion.py @@ -0,0 +1,317 @@ +from __future__ import annotations + +import copy +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable + +import torch +from torch import Tensor, nn + +from .memory_attention import MemoryAugmentedCellularLayer + +DEFAULT_SMOLLM2 = "HuggingFaceTB/SmolLM2-135M" + + +@dataclass(frozen=True) +class SmolMemoryFusionConfig: + feature_dim: int = 32 + memory_rank: int = 48 + dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64, 128) + shifted_window: int = 8 + train_output_projection: bool = True + + def validate(self, model_config) -> None: + if self.feature_dim < 4: + raise ValueError("feature_dim must be >= 4") + if self.memory_rank < 4: + raise ValueError("memory_rank must be >= 4") + if not self.dilations or min(self.dilations) < 1: + raise ValueError("dilations must be positive") + if int(model_config.num_attention_heads) % int(model_config.num_key_value_heads): + raise ValueError("num_attention_heads must be divisible by num_key_value_heads") + + def to_dict(self) -> dict: + value = asdict(self) + value["dilations"] = list(self.dilations) + return value + + @classmethod + def from_dict(cls, data: dict) -> "SmolMemoryFusionConfig": + value = dict(data) + value["dilations"] = tuple(value.get("dilations", (1, 2, 4, 8, 16, 32, 64, 128))) + return cls(**value) + + +class MemoryFusionLlamaAttention(nn.Module): + """Drop-in no-cache Llama attention using TinyCeNN Memory Fusion. + + The pretrained Q/K/V/O projections are copied exactly. The new trainable core + combines adaptive-MaxPool Cellular attention, a Hedgehog-style global linear + path, and GDN2-style editable memory. The core stays in float32 for numerical + stability while the surrounding SmolLM2 projections can remain bf16/fp16. + """ + + def __init__( + self, + original_attn: nn.Module, + model_config, + config: SmolMemoryFusionConfig, + layer_idx: int, + ) -> None: + super().__init__() + config.validate(model_config) + self.layer_idx = int(layer_idx) + self.hidden_size = int(model_config.hidden_size) + self.num_heads = int(model_config.num_attention_heads) + self.num_key_value_heads = int(model_config.num_key_value_heads) + self.head_dim = int(getattr(model_config, "head_dim", self.hidden_size // self.num_heads)) + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + + self.q_proj = copy.deepcopy(original_attn.q_proj) + self.k_proj = copy.deepcopy(original_attn.k_proj) + self.v_proj = copy.deepcopy(original_attn.v_proj) + self.o_proj = copy.deepcopy(original_attn.o_proj) + + self.core = MemoryAugmentedCellularLayer( + num_heads=self.num_heads, + num_kv_heads=self.num_key_value_heads, + head_dim=self.head_dim, + feature_dim=config.feature_dim, + variant="cellular_memory_fusion", + dilations=config.dilations, + shifted_window=config.shifted_window, + memory_rank=config.memory_rank, + ) + self.last_core_output: Tensor | None = None + + def _apply_rope(self, q: Tensor, k: Tensor, position_embeddings) -> tuple[Tensor, Tensor]: + if position_embeddings is None: + raise ValueError("position_embeddings are required for Memory Fusion attention") + from transformers.models.llama.modeling_llama import apply_rotary_pos_emb + + cos, sin = position_embeddings + return apply_rotary_pos_emb(q, k, cos, sin) + + def forward( + self, + hidden_states: Tensor, + attention_mask=None, + position_ids=None, + past_key_values=None, + past_key_value=None, + use_cache: bool = False, + cache_position=None, + position_embeddings=None, + **kwargs, + ) -> tuple[Tensor, None]: + if use_cache or past_key_values is not None or past_key_value is not None: + raise RuntimeError("SmolLM2 Memory Fusion currently requires use_cache=False") + bsz, seq_len, _ = hidden_states.shape + + if attention_mask is not None and attention_mask.ndim == 4: + if attention_mask.shape[-1] != seq_len: + raise ValueError("attention mask length mismatch") + if bool((attention_mask[..., -1, :] < -1e4).any()): + raise ValueError("padded batches are not supported by Memory Fusion attention") + + q = self.q_proj(hidden_states).view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + k = self.k_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + v = self.v_proj(hidden_states).view( + bsz, seq_len, self.num_key_value_heads, self.head_dim + ).transpose(1, 2) + q, k = self._apply_rope(q, k, position_embeddings) + + core_out = self.core(q.float(), k.float(), v.float()) + self.last_core_output = core_out + flat = core_out.transpose(1, 2).reshape(bsz, seq_len, self.hidden_size) + out = self.o_proj(flat.to(dtype=hidden_states.dtype)) + return out, None + + +def replace_attention_layers( + model: nn.Module, + config: SmolMemoryFusionConfig, + layer_indices: Iterable[int], +) -> nn.Module: + config.validate(model.config) + for raw_idx in layer_indices: + idx = int(raw_idx) + layer = model.model.layers[idx] + if isinstance(layer.self_attn, MemoryFusionLlamaAttention): + continue + old = layer.self_attn + device = old.q_proj.weight.device + projection_dtype = old.q_proj.weight.dtype + new = MemoryFusionLlamaAttention(old, model.config, config, idx) + for projection in (new.q_proj, new.k_proj, new.v_proj, new.o_proj): + projection.to(device=device, dtype=projection_dtype) + new.core.to(device=device, dtype=torch.float32) + layer.self_attn = new + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def replace_all_attention(model: nn.Module, config: SmolMemoryFusionConfig) -> nn.Module: + return replace_attention_layers(model, config, range(int(model.config.num_hidden_layers))) + + +def freeze_for_group_calibration( + model: nn.Module, + layer_indices: Iterable[int], + *, + train_output_projection: bool = True, +) -> list[nn.Parameter]: + selected = {int(x) for x in layer_indices} + for p in model.parameters(): + p.requires_grad = False + trainable: list[nn.Parameter] = [] + for module in model.modules(): + if isinstance(module, MemoryFusionLlamaAttention) and module.layer_idx in selected: + for p in module.core.parameters(): + p.requires_grad = True + trainable.append(p) + if train_output_projection: + for p in module.o_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def freeze_for_global_training( + model: nn.Module, + *, + train_output_projection: bool = True, +) -> list[nn.Parameter]: + for p in model.parameters(): + p.requires_grad = False + trainable: list[nn.Parameter] = [] + for module in model.modules(): + if isinstance(module, MemoryFusionLlamaAttention): + for p in module.core.parameters(): + p.requires_grad = True + trainable.append(p) + if train_output_projection: + for p in module.o_proj.parameters(): + p.requires_grad = True + trainable.append(p) + return trainable + + +def parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + cores = sum( + p.numel() + for m in model.modules() + if isinstance(m, MemoryFusionLlamaAttention) + for p in m.core.parameters() + ) + output_proj = sum( + p.numel() + for m in model.modules() + if isinstance(m, MemoryFusionLlamaAttention) + for p in m.o_proj.parameters() + ) + return { + "total": total, + "trainable": trainable, + "memory_fusion_cores": cores, + "output_projections": output_proj, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def structural_summary(model: nn.Module) -> dict[str, int]: + fusion = sum(isinstance(m, MemoryFusionLlamaAttention) for m in model.modules()) + transformer = sum( + m.__class__.__name__ == "LlamaAttention" + for m in model.modules() + if not isinstance(m, MemoryFusionLlamaAttention) + ) + return {"memory_fusion_layers": fusion, "transformer_attention_layers": transformer} + + +def structural_assertions(model: nn.Module) -> None: + expected = int(model.config.num_hidden_layers) + summary = structural_summary(model) + if summary["memory_fusion_layers"] != expected: + raise RuntimeError(f"expected {expected} Memory Fusion layers, found {summary['memory_fusion_layers']}") + if summary["transformer_attention_layers"]: + raise RuntimeError("Transformer self-attention remains in the final student") + + +def _attention_state(model: nn.Module) -> dict[str, Tensor]: + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".self_attn." in name: + state[name] = tensor.detach().cpu() + if not state: + raise RuntimeError("no Memory Fusion attention state found") + return state + + +def save_smollm2_memory_fusion( + model: nn.Module, + output_dir: str | Path, + *, + config: SmolMemoryFusionConfig, + base_model: str = DEFAULT_SMOLLM2, + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_attention_state(model), output_dir / "smollm2_memory_fusion.pt") + meta = { + "format_version": 1, + "architecture": "smollm2-memory-fusion", + "base_model": base_model, + "memory_fusion": config.to_dict(), + } + if extra_metadata: + meta["training"] = extra_metadata + (output_dir / "smollm2_memory_fusion_config.json").write_text( + json.dumps(meta, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_smollm2_memory_fusion_weights(model: nn.Module, student_dir: str | Path) -> nn.Module: + path = Path(student_dir) / "smollm2_memory_fusion.pt" + state = torch.load(path, map_location="cpu", weights_only=True) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_attention_state(model)) + missing = [k for k in incompatible.missing_keys if k in expected] + unexpected = [k for k in incompatible.unexpected_keys if k not in expected] + if missing: + raise RuntimeError(f"missing Memory Fusion keys: {missing[:8]}") + if unexpected: + raise RuntimeError(f"unexpected Memory Fusion keys: {unexpected[:8]}") + return model + + +def build_smollm2_memory_fusion(student_dir: str | Path, *, device=None, dtype=None): + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + meta = json.loads((student_dir / "smollm2_memory_fusion_config.json").read_text()) + if meta.get("architecture") != "smollm2-memory-fusion": + raise ValueError("checkpoint is not SmolLM2 Memory Fusion") + kwargs = {} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(meta["base_model"], **kwargs) + if device is not None: + model.to(device) + config = SmolMemoryFusionConfig.from_dict(meta["memory_fusion"]) + replace_all_attention(model, config) + load_smollm2_memory_fusion_weights(model, student_dir) + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model diff --git a/src/tinycenn_lm/story.py b/src/tinycenn_lm/story.py new file mode 100644 index 0000000000000000000000000000000000000000..62188bf7a0d82cd8562ed1f0149672a26129132f --- /dev/null +++ b/src/tinycenn_lm/story.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from collections import Counter + +import torch +from torch import Tensor + + +def repetition_unlikelihood_loss( + logits: Tensor, + labels: Tensor, + *, + window: int = 32, + ignore_index: int = -100, +) -> Tensor: + """Penalize probability assigned to recently seen tokens. + + For every next-token prediction, tokens that appeared in the recent history + are treated as negative candidates unless the token is the true next target. + This is a lightweight unlikelihood objective aimed specifically at the short + repetition loops seen in TinyCeNN generation. + """ + if logits.ndim != 3 or labels.ndim != 2: + raise ValueError("expected logits [B,T,V] and labels [B,T]") + if logits.shape[:2] != labels.shape: + raise ValueError("logits and labels sequence dimensions must match") + if window <= 0: + return logits.new_zeros(()) + + pred = logits[:, :-1, :].float() + targets = labels[:, 1:] + history = labels[:, :-1] + log_z = torch.logsumexp(pred, dim=-1) + + total = pred.new_zeros(()) + count = pred.new_zeros(()) + max_back = min(window, history.shape[1]) + + for back in range(max_back): + if back == 0: + negatives = history + valid = torch.ones_like(history, dtype=torch.bool) + else: + negatives = torch.roll(history, shifts=back, dims=1) + valid = torch.ones_like(history, dtype=torch.bool) + valid[:, :back] = False + + valid &= targets.ne(ignore_index) + valid &= negatives.ne(ignore_index) + valid &= negatives.ne(targets) + + safe_negatives = negatives.clamp_min(0) + neg_logits = pred.gather(-1, safe_negatives.unsqueeze(-1)).squeeze(-1) + p_negative = torch.exp(neg_logits - log_z).clamp(max=1.0 - 1e-6) + penalties = -torch.log1p(-p_negative) + + total = total + penalties.masked_select(valid).sum() + count = count + valid.sum().to(dtype=total.dtype) + + return total / count.clamp_min(1.0) + + +def repeated_ngram_fraction(text: str, n: int = 3) -> float: + """Fraction of generated n-gram occurrences beyond their first occurrence.""" + words = text.split() + if len(words) < n or n <= 0: + return 0.0 + grams = [tuple(words[i : i + n]) for i in range(len(words) - n + 1)] + counts = Counter(grams) + repeats = sum(max(0, c - 1) for c in counts.values()) + return repeats / max(len(grams), 1) + + +def story_generation_kwargs(tokenizer, *, max_new_tokens: int = 120) -> dict: + """Decoding defaults chosen to suppress loops without making text deterministic.""" + return { + "max_new_tokens": max_new_tokens, + "min_new_tokens": min(40, max_new_tokens), + "do_sample": True, + "temperature": 0.78, + "top_p": 0.90, + "top_k": 40, + "repetition_penalty": 1.18, + "no_repeat_ngram_size": 4, + "renormalize_logits": True, + "use_cache": False, + "eos_token_id": tokenizer.eos_token_id, + "pad_token_id": tokenizer.pad_token_id or tokenizer.eos_token_id, + } diff --git a/src/tinycenn_lm/story_v2.py b/src/tinycenn_lm/story_v2.py new file mode 100644 index 0000000000000000000000000000000000000000..61e38dcd8c86d66cadd9e600be070ad55f7fdb75 --- /dev/null +++ b/src/tinycenn_lm/story_v2.py @@ -0,0 +1,271 @@ +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Sequence + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from .modeling import DEFAULT_BASE_MODEL, _get_decoder_layers +from .sharded_moe import ( + ShardedMoECeNNConfig, + ShardedMoECeNNReplacementLayer, + build_sharded_moe_student, +) + + +@dataclass(frozen=True) +class StoryV2Config: + memory_rank: int = 32 + head_rank: int = 4 + + def validate(self, hidden_size: int, vocab_size: int) -> None: + if not 1 <= self.memory_rank <= hidden_size: + raise ValueError("memory_rank must be in [1, hidden_size]") + if not 1 <= self.head_rank <= min(hidden_size, vocab_size): + raise ValueError("head_rank must be positive and <= hidden/vocab size") + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> "StoryV2Config": + return cls(**dict(data)) + + +class CausalStoryMemory(nn.Module): + """Cheap global prefix memory with no attention and no token-by-token Python loop.""" + + def __init__(self, hidden_size: int, rank: int) -> None: + super().__init__() + self.down = nn.Linear(hidden_size, rank, bias=False) + self.up = nn.Linear(rank, hidden_size, bias=False) + nn.init.normal_(self.down.weight, mean=0.0, std=0.02) + nn.init.zeros_(self.up.weight) + + def forward(self, hidden_states: Tensor) -> Tensor: + dtype = hidden_states.dtype + work = hidden_states.float() + prefix_sum = work.cumsum(dim=1) + denom = torch.arange( + 1, + hidden_states.shape[1] + 1, + device=hidden_states.device, + dtype=work.dtype, + ).view(1, -1, 1) + prefix_mean = (prefix_sum / denom).to(dtype=dtype) + return self.up(F.silu(self.down(prefix_mean))) + + +class StoryV2ReplacementLayer(nn.Module): + def __init__(self, cenn: nn.Module, hidden_size: int, story_config: StoryV2Config) -> None: + super().__init__() + self.cenn = cenn + self.story_memory = CausalStoryMemory(hidden_size, story_config.memory_rank) + + def forward(self, hidden_states: Tensor, *args, **kwargs) -> Tensor: + if kwargs.get("use_cache", False): + raise RuntimeError("TinyCeNN Story-v2 requires use_cache=False") + if kwargs.get("output_attentions", False): + raise RuntimeError("TinyCeNN Story-v2 has no attention matrices") + memory = self.story_memory(hidden_states) + enriched = hidden_states + memory + return hidden_states + self.cenn(enriched) + + +class LowRankLMHeadAdapter(nn.Module): + """Frozen language-model head plus a tiny trainable low-rank story adapter.""" + + def __init__(self, base_head: nn.Module, hidden_size: int, vocab_size: int, rank: int) -> None: + super().__init__() + self.base_head = base_head + for parameter in self.base_head.parameters(): + parameter.requires_grad = False + self.down = nn.Linear(hidden_size, rank, bias=False) + self.up = nn.Linear(rank, vocab_size, bias=False) + nn.init.normal_(self.down.weight, mean=0.0, std=0.02) + nn.init.zeros_(self.up.weight) + + @property + def weight(self): + return getattr(self.base_head, "weight", None) + + def forward(self, hidden_states: Tensor) -> Tensor: + return self.base_head(hidden_states) + self.up(F.silu(self.down(hidden_states))) + + +def upgrade_sharded_model_to_story_v2(model: nn.Module, story_config: StoryV2Config) -> nn.Module: + hidden_size = int(model.config.hidden_size) + vocab_size = int(model.config.vocab_size) + story_config.validate(hidden_size, vocab_size) + + layers = _get_decoder_layers(model) + replaced = 0 + for index, layer in enumerate(list(layers)): + if isinstance(layer, ShardedMoECeNNReplacementLayer): + layers[index] = StoryV2ReplacementLayer(layer.cenn, hidden_size, story_config) + replaced += 1 + if replaced == 0: + raise RuntimeError("no Sharded MoE-CeNN layer found to upgrade") + + base_head = model.get_output_embeddings() + if base_head is None: + raise RuntimeError("base model has no output embedding/head") + if not isinstance(base_head, LowRankLMHeadAdapter): + model.set_output_embeddings( + LowRankLMHeadAdapter(base_head, hidden_size, vocab_size, story_config.head_rank) + ) + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + return model + + +def freeze_story_v2_interfaces(model: nn.Module) -> None: + for parameter in model.parameters(): + parameter.requires_grad = False + for module in model.modules(): + if isinstance(module, StoryV2ReplacementLayer): + for parameter in module.cenn.parameters(): + parameter.requires_grad = True + for parameter in module.story_memory.parameters(): + parameter.requires_grad = True + elif isinstance(module, LowRankLMHeadAdapter): + for parameter in module.down.parameters(): + parameter.requires_grad = True + for parameter in module.up.parameters(): + parameter.requires_grad = True + + +def story_v2_router_stats(model: nn.Module) -> dict[str, Tensor]: + layer = next((m for m in model.modules() if isinstance(m, StoryV2ReplacementLayer)), None) + if layer is None or not getattr(layer.cenn, "last_router_stats", None): + raise RuntimeError("Story-v2 router statistics unavailable; run a forward pass first") + return layer.cenn.last_router_stats + + +def story_v2_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + memory = sum( + p.numel() for m in model.modules() if isinstance(m, CausalStoryMemory) for p in m.parameters() + ) + head_adapter = sum( + p.numel() + for m in model.modules() + if isinstance(m, LowRankLMHeadAdapter) + for sub in (m.down, m.up) + for p in sub.parameters() + ) + return { + "total": total, + "trainable": trainable, + "memory": memory, + "head_adapter": head_adapter, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def _story_v2_state_dict(model: nn.Module) -> dict[str, Tensor]: + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".cenn." in name or ".story_memory." in name: + state[name] = tensor.detach().cpu() + elif "lm_head.down." in name or "lm_head.up." in name: + state[name] = tensor.detach().cpu() + if not state: + raise RuntimeError("no Story-v2 trainable state found") + return state + + +def save_story_v2_student( + model: nn.Module, + output_dir: str | Path, + *, + story_config: StoryV2Config, + sharded_config: ShardedMoECeNNConfig, + base_model: str = DEFAULT_BASE_MODEL, + layer_indices: Sequence[int] = (0,), + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + torch.save(_story_v2_state_dict(model), output_dir / "story_v2_student.pt") + metadata = { + "format_version": 1, + "architecture": "tinycenn-story-v2-memory-head", + "base_model": base_model, + "layer_indices": list(layer_indices), + "sharded_moe_cenn": sharded_config.to_dict(), + "story_v2": story_config.to_dict(), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "story_v2_config.json").write_text( + json.dumps(metadata, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_story_v2_weights(model: nn.Module, student_dir: str | Path) -> nn.Module: + state = torch.load(Path(student_dir) / "story_v2_student.pt", map_location="cpu", weights_only=True) + incompatible = model.load_state_dict(state, strict=False) + expected = set(_story_v2_state_dict(model)) + missing = [key for key in incompatible.missing_keys if key in expected] + if missing: + raise RuntimeError(f"missing Story-v2 keys: {missing}") + if incompatible.unexpected_keys: + raise RuntimeError(f"unexpected Story-v2 keys: {incompatible.unexpected_keys}") + return model + + +def build_story_v2_student( + student_dir: str | Path, + *, + device=None, + dtype=None, + attn_implementation: str = "sdpa", +): + from transformers import AutoModelForCausalLM + from .sharded_moe import replace_transformer_with_sharded_moe_cenn + + student_dir = Path(student_dir) + metadata = json.loads((student_dir / "story_v2_config.json").read_text()) + if metadata.get("architecture") != "tinycenn-story-v2-memory-head": + raise ValueError("checkpoint is not a TinyCeNN Story-v2 model") + + kwargs = {"attn_implementation": attn_implementation} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs) + sharded_config = ShardedMoECeNNConfig.from_dict(metadata["sharded_moe_cenn"]) + replace_transformer_with_sharded_moe_cenn(model, sharded_config, tuple(metadata["layer_indices"])) + story_config = StoryV2Config.from_dict(metadata["story_v2"]) + upgrade_sharded_model_to_story_v2(model, story_config) + load_story_v2_weights(model, student_dir) + + move_kwargs = {} + if device is not None: + move_kwargs["device"] = device + if dtype is not None: + move_kwargs["dtype"] = dtype + if move_kwargs: + model.to(**move_kwargs) + model.config.use_cache = False + return model + + +def build_story_v2_from_story_v1( + story_v1_dir: str | Path, + *, + story_config: StoryV2Config, + device=None, + dtype=None, +): + model = build_sharded_moe_student(story_v1_dir, device=device, dtype=dtype) + upgrade_sharded_model_to_story_v2(model, story_config) + freeze_story_v2_interfaces(model) + return model diff --git a/src/tinycenn_lm/student.py b/src/tinycenn_lm/student.py new file mode 100644 index 0000000000000000000000000000000000000000..06bdcb9c536c578abb903115938d5b61bfa2a6a0 --- /dev/null +++ b/src/tinycenn_lm/student.py @@ -0,0 +1,223 @@ +from __future__ import annotations + +import json +from pathlib import Path +from typing import Sequence + +import torch +from torch import Tensor, nn + +from .cenn import CeNNConfig, FastCeNNCore +from .modeling import DEFAULT_BASE_MODEL, _get_decoder_layers + + +class CeNNReplacementLayer(nn.Module): + """Transformer-free decoder layer built only from a recurrent causal CeNN core. + + The surrounding pretrained language interface is intentionally kept: token + embeddings, final RMSNorm, tokenizer and LM head. The original Transformer + decoder layer itself is removed and no attention/MLP parameters remain in + this layer. + """ + + def __init__( + self, + config: CeNNConfig, + *, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + self.config = config + self.cenn = FastCeNNCore(config) + if device is not None or dtype is not None: + move_kwargs: dict[str, object] = {} + if device is not None: + move_kwargs["device"] = device + if dtype is not None: + move_kwargs["dtype"] = dtype + self.cenn.to(**move_kwargs) + + def forward(self, hidden_states: Tensor, *args, **kwargs) -> Tensor: + if kwargs.get("use_cache", False): + raise RuntimeError( + "CeNN student requires use_cache=False until a recurrent CeNN state cache " + "is implemented." + ) + if kwargs.get("output_attentions", False): + raise RuntimeError("CeNN-only student has no attention matrices to return.") + return hidden_states + self.cenn(hidden_states) + + +def _reference_parameter(module: nn.Module) -> nn.Parameter | None: + return next((p for p in module.parameters() if p.is_floating_point()), None) + + +def replace_transformer_with_cenn( + model: nn.Module, + config: CeNNConfig, + layer_indices: Sequence[int] = (0,), +) -> nn.Module: + """Remove selected Transformer decoder layers and replace them with CeNN layers.""" + layers = _get_decoder_layers(model) + hidden_size = int(getattr(model.config, "hidden_size")) + if config.hidden_size != hidden_size: + raise ValueError( + f"CeNN hidden_size={config.hidden_size} does not match model hidden_size={hidden_size}" + ) + + if hasattr(model, "config"): + model.config.use_cache = False + if hasattr(model, "generation_config"): + model.generation_config.use_cache = False + + for index in layer_indices: + if index < 0 or index >= len(layers): + raise IndexError(f"layer index {index} out of range [0, {len(layers)})") + if isinstance(layers[index], CeNNReplacementLayer): + raise ValueError(f"layer {index} is already CeNN-only") + reference = _reference_parameter(layers[index]) + device = reference.device if reference is not None else None + dtype = reference.dtype if reference is not None else None + layers[index] = CeNNReplacementLayer(config, device=device, dtype=dtype) + return model + + +def freeze_student_interfaces(model: nn.Module, train_interfaces: str = "none") -> None: + """Train the core and optionally adapt pretrained interfaces after distillation. + + ``norm`` trains the final normalization; ``all`` also trains the embeddings + and output head. The default retains the original core-only experiment. + """ + if train_interfaces not in {"none", "norm", "all"}: + raise ValueError("train_interfaces must be none, norm, or all") + for parameter in model.parameters(): + parameter.requires_grad = False + for module in model.modules(): + if isinstance(module, CeNNReplacementLayer): + for parameter in module.parameters(): + parameter.requires_grad = True + if train_interfaces != "none": + norm = getattr(getattr(model, "model", None), "norm", None) + if not isinstance(norm, nn.Module): + raise ValueError("student has no supported final model.norm") + norm.requires_grad_(True) + if train_interfaces == "all": + for interface in (model.get_input_embeddings(), model.get_output_embeddings()): + if interface is None: + raise ValueError("student must expose input and output embeddings") + interface.requires_grad_(True) + + +def student_parameter_summary(model: nn.Module) -> dict[str, int | float]: + total = sum(p.numel() for p in model.parameters()) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + return { + "total": total, + "trainable": trainable, + "trainable_percent": 100.0 * trainable / max(total, 1), + } + + +def _student_state_dict(model: nn.Module) -> dict[str, Tensor]: + # Keep interfaces loaded from a v2 checkpoint even when they are frozen in a + # later stage. Otherwise re-saving silently reverts them to the base model. + extra_keys = set(getattr(model, "_cenn_interface_keys", ())) + extra_keys.update(name for name, p in model.named_parameters() if p.requires_grad) + state: dict[str, Tensor] = {} + for name, tensor in model.state_dict().items(): + if ".cenn." in name or name in extra_keys: + state[name] = tensor.detach().cpu() + if not any(".cenn." in name for name in state): + raise ValueError("no CeNN replacement weights found") + return state + + +def save_cenn_student( + model: nn.Module, + output_dir: str | Path, + *, + config: CeNNConfig, + base_model: str = DEFAULT_BASE_MODEL, + layer_indices: Sequence[int] = (0,), + extra_metadata: dict | None = None, +) -> Path: + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + state = _student_state_dict(model) + torch.save(state, output_dir / "cenn_student.pt") + metadata = { + "format_version": 2, + "architecture": "cenn-only-replacement", + "base_model": base_model, + "layer_indices": list(layer_indices), + "cenn": config.to_dict(), + "state_keys": sorted(state), + } + if extra_metadata: + metadata["training"] = extra_metadata + (output_dir / "student_config.json").write_text( + json.dumps(metadata, indent=2), encoding="utf-8" + ) + return output_dir + + +def load_cenn_student_weights( + model: nn.Module, + student_dir: str | Path, + *, + map_location: str | torch.device = "cpu", + strict: bool = True, +) -> nn.Module: + student_dir = Path(student_dir) + state = torch.load( + student_dir / "cenn_student.pt", + map_location=map_location, + weights_only=True, + ) + metadata = json.loads((student_dir / "student_config.json").read_text()) + core_keys = {key for key in model.state_dict() if ".cenn." in key} + expected = set(metadata.get("state_keys", core_keys)) | core_keys + missing = sorted(expected - state.keys()) + unexpected = sorted(state.keys() - expected | state.keys() - model.state_dict().keys()) + if strict and missing: + raise RuntimeError(f"missing CeNN student keys: {missing}") + if strict and unexpected: + raise RuntimeError(f"unexpected CeNN student keys: {unexpected}") + model.load_state_dict(state, strict=False) + model._cenn_interface_keys = tuple(key for key in state if ".cenn." not in key) + return model + + +def build_cenn_student( + student_dir: str | Path, + *, + device: str | torch.device | None = None, + dtype: torch.dtype | None = None, + attn_implementation: str = "sdpa", +): + """Rebuild a Transformer-free CeNN student from the pretrained interfaces.""" + from transformers import AutoModelForCausalLM + + student_dir = Path(student_dir) + metadata = json.loads((student_dir / "student_config.json").read_text()) + if metadata.get("architecture") != "cenn-only-replacement": + raise ValueError("checkpoint is not a CeNN-only replacement student") + + kwargs: dict[str, object] = {"attn_implementation": attn_implementation} + if dtype is not None: + kwargs["dtype"] = dtype + model = AutoModelForCausalLM.from_pretrained(metadata["base_model"], **kwargs) + config = CeNNConfig.from_dict(metadata["cenn"]) + replace_transformer_with_cenn(model, config, tuple(metadata["layer_indices"])) + load_cenn_student_weights(model, student_dir) + + move_kwargs: dict[str, object] = {} + if device is not None: + move_kwargs["device"] = device + if dtype is not None: + move_kwargs["dtype"] = dtype + if move_kwargs: + model.to(**move_kwargs) + model.config.use_cache = False + return model diff --git a/src/tinycenn_lm/training.py b/src/tinycenn_lm/training.py new file mode 100644 index 0000000000000000000000000000000000000000..dc92d61fc443f5b839f6dfe46eb5d1dbd8ed3aba --- /dev/null +++ b/src/tinycenn_lm/training.py @@ -0,0 +1,137 @@ +"""Training policies shared by continuation and its offline regression tests.""" +from __future__ import annotations + +import hashlib +import json +import math + +import torch + + +def training_stream_signature(args, tokenizer) -> dict: + backend = getattr(tokenizer, "backend_tokenizer", None) + tokenizer_state = { + "vocab": tokenizer.get_vocab(), "eos": tokenizer.eos_token_id, + "backend": backend.to_str() if backend is not None else None, + } + digest = hashlib.sha256(json.dumps(tokenizer_state, sort_keys=True).encode()).hexdigest() + return { + "dataset": args.dataset, "dataset_config": args.dataset_config, + "dataset_revision": args.dataset_revision, "split": args.split, + "text_field": args.text_field, "shuffle_buffer": args.shuffle_buffer, + "seed": args.seed, "tokenizer_sha256": digest, + "partition": "text-hash-99-1-v1", "packing": "eos-v1", + } + + +def optimizer_groups(model, learning_rate: float, interface_lr_scale: float, weight_decay: float): + """Use a smaller LR for copied interfaces and no decay for norms/biases.""" + if learning_rate <= 0 or not 0 < interface_lr_scale <= 1 or weight_decay < 0: + raise ValueError("invalid learning rate, interface LR scale, or weight decay") + groups: dict[tuple[bool, bool], dict] = {} + for name, parameter in model.named_parameters(): + if not parameter.requires_grad: + continue + if parameter.dtype != torch.float32: + raise ValueError(f"trainable parameter {name} must be float32 before creating AdamW") + interface = ".cenn." not in name + decay = parameter.ndim >= 2 + key = (interface, decay) + if key not in groups: + lr = learning_rate * (interface_lr_scale if interface else 1.0) + groups[key] = { + "params": [], "lr": lr, "initial_lr": lr, + "weight_decay": weight_decay if decay else 0.0, + "name": ("interface" if interface else "cenn") + ("_decay" if decay else "_no_decay"), + } + groups[key]["params"].append(parameter) + if not groups: + raise ValueError("no trainable parameters") + return list(groups.values()) + + +def lr_multiplier( + step: int, total: int, warmup: int, *, schedule: str = "cosine", + decay_ratio: float = 0.2, min_lr_ratio: float = 0.1, +) -> float: + """Cosine or warmup/stable/decay, with the same explicit nonzero LR floor.""" + if total < 1 or warmup < 0 or warmup >= total: + raise ValueError("require total >= 1 and 0 <= warmup < total") + if schedule not in {"cosine", "wsd"} or not 0 < decay_ratio <= 1: + raise ValueError("invalid LR schedule or decay_ratio") + if not 0 <= min_lr_ratio <= 1: + raise ValueError("min_lr_ratio must be in [0, 1]") + if warmup and step <= warmup: + return max(step, 1) / warmup + decay_start = warmup + if schedule == "wsd": + decay_start = max(warmup, total - max(1, round(total * decay_ratio))) + progress = min(max((step - decay_start) / max(total - decay_start, 1), 0.0), 1.0) + return min_lr_ratio + (1 - min_lr_ratio) * 0.5 * (1 + math.cos(math.pi * progress)) + + +def annealed_weight(initial: float, final: float | None, progress: float) -> float: + if initial < 0 or (final is not None and final < 0): + raise ValueError("loss weights must be nonnegative") + end = initial if final is None else final + return initial + (end - initial) * min(max(progress, 0.0), 1.0) + + +def checkpoint_training_tokens(metadata: dict, report: dict | None) -> int: + """Read the saved weights' token count, never a best run's final token count.""" + training = metadata.get("training", {}) + for key in ("cumulative_tokens", "tokens"): + if key in training: + return int(training[key]) + report = report or {} + checkpoint = report.get("checkpoint", {}) + if "cumulative_training_tokens" in checkpoint: + return int(checkpoint["cumulative_training_tokens"]) + return int(report.get("cumulative_training_tokens", report.get("seen_tokens", 0))) + + +def stream_resume_offset( + metadata: dict, report: dict | None, signature: dict, *, explicit_offset: int | None = None, +) -> tuple[int, str]: + """Recover an exact new-format cursor, or label a legacy estimate explicitly. + + A v2 run restarted its stream at zero. Its run-local token count is therefore + the prefix to skip, not the cumulative token count across earlier runs. + """ + if explicit_offset is not None: + if explicit_offset < 0: + raise ValueError("train-skip-tokens must be nonnegative") + return explicit_offset, "explicit" + if not metadata: + return 0, "new_stream" + stream = metadata.get("training", {}).get("data_stream") + if stream: + if stream["signature"] != signature: + raise ValueError( + "resume data stream differs (dataset/revision/tokenizer/shuffle). " + "Use matching settings, or --train-skip-tokens for an intentional new stream." + ) + return int(stream["next_token_offset"]), "exact_token_offset" + report = report or {} + training = metadata.get("training", {}) + # Legacy metadata identifies the actual best snapshot; a copied report may + # describe a later final model. Favor the snapshot's run-local count. + offset = training.get("tokens_this_run", training.get("tokens")) + if offset is None: + offset = report.get("seen_tokens_this_run", report.get("seen_tokens", 0)) + return int(offset), "legacy_estimate" + + +def plateau_summary(history: list[dict], patience: int, min_delta: float) -> dict: + if patience < 1 or min_delta < 0: + raise ValueError("plateau patience must be positive and min_delta nonnegative") + if len(history) <= patience: + return {"detected": False, "evaluations_without_progress": 0} + prior_best = min(float(row["student_ce"]) for row in history[:-patience]) + recent_best = min(float(row["student_ce"]) for row in history[-patience:]) + improvement = prior_best - recent_best + return { + "detected": improvement < min_delta, + "evaluations_without_progress": patience if improvement < min_delta else 0, + "recent_best_ce_improvement": improvement, + } diff --git a/tinycenn_qwen35.json b/tinycenn_qwen35.json new file mode 100644 index 0000000000000000000000000000000000000000..56444ac0b3530aae6b2fe937b212f05c67e3d096 --- /dev/null +++ b/tinycenn_qwen35.json @@ -0,0 +1,22 @@ +{ + "format_version": 1, + "architecture": "qwen3.5-pdelta3-gdn2-clvr-localw", + "base_model": "Qwen/Qwen3.5-0.8B", + "accepted_full_attention_layers": [ + 3, + 7, + 11 + ], + "replacement_config": { + "feature_dim": 96, + "local_window": 32, + "chunk_size": 32, + "conv_kernel": 4, + "state_dtype": "fp16", + "variant": "conv4_gdn2_clvr_f96", + "local_gate_init": 0.72, + "warm_start_previous_core": true + }, + "verification_file": "qwen35_verification.json", + "loader": "load_model.py" +} \ No newline at end of file diff --git a/tokenizer.json b/tokenizer.json new file mode 100644 index 0000000000000000000000000000000000000000..5520bfd2dd834ce386c1312c410fa71af56db5ad --- /dev/null +++ b/tokenizer.json @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:06b9509352d2af50381ab2247e083b80d32d5c0aba91c272ca9ff729b6a0e523 +size 19989325 diff --git a/tokenizer_config.json b/tokenizer_config.json new file mode 100644 index 0000000000000000000000000000000000000000..ed1f99f3d50eecc9850623be644bd1131686cbe1 --- /dev/null +++ b/tokenizer_config.json @@ -0,0 +1,32 @@ +{ + "add_prefix_space": false, + "audio_bos_token": "<|audio_start|>", + "audio_eos_token": "<|audio_end|>", + "audio_token": "<|audio_pad|>", + "backend": "tokenizers", + "bos_token": null, + "clean_up_tokenization_spaces": false, + "eos_token": "<|im_end|>", + "errors": "replace", + "image_token": "<|image_pad|>", + "is_local": false, + "local_files_only": false, + "model_max_length": 262144, + "model_specific_special_tokens": { + "audio_bos_token": "<|audio_start|>", + "audio_eos_token": "<|audio_end|>", + "audio_token": "<|audio_pad|>", + "image_token": "<|image_pad|>", + "video_token": "<|video_pad|>", + "vision_bos_token": "<|vision_start|>", + "vision_eos_token": "<|vision_end|>" + }, + "pad_token": "<|endoftext|>", + "pretokenize_regex": "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + "split_special_tokens": false, + "tokenizer_class": "Qwen2Tokenizer", + "unk_token": null, + "video_token": "<|video_pad|>", + "vision_bos_token": "<|vision_start|>", + "vision_eos_token": "<|vision_end|>" +}