From 7959ed561c2ced8d99901381114e81b470956e63 Mon Sep 17 00:00:00 2001 From: ModelHub XC Date: Fri, 11 Sep 2026 05:36:12 +0800 Subject: [PATCH] =?UTF-8?q?=E5=88=9D=E5=A7=8B=E5=8C=96=E9=A1=B9=E7=9B=AE?= =?UTF-8?q?=EF=BC=8C=E7=94=B1ModelHub=20XC=E7=A4=BE=E5=8C=BA=E6=8F=90?= =?UTF-8?q?=E4=BE=9B=E6=A8=A1=E5=9E=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Model: PKU-DS-LAB/Fairy2i-W2 Source: Original Platform --- .gitattributes | 53 +++++ README.md | 144 ++++++++++++ config.json | 29 +++ configuration.json | 1 + generation_config.json | 10 + load_model.py | 143 ++++++++++++ model-00001-of-00003.safetensors | 3 + model-00002-of-00003.safetensors | 3 + model-00003-of-00003.safetensors | 3 + model.safetensors.index.json | 298 ++++++++++++++++++++++++ qat_modules.py | 388 +++++++++++++++++++++++++++++++ quantization.py | 308 ++++++++++++++++++++++++ special_tokens_map.json | 24 ++ tokenizer.json | 3 + tokenizer_config.json | 43 ++++ training_args.bin | 3 + 16 files changed, 1456 insertions(+) create mode 100644 .gitattributes create mode 100644 README.md create mode 100644 config.json create mode 100644 configuration.json create mode 100644 generation_config.json create mode 100644 load_model.py create mode 100644 model-00001-of-00003.safetensors create mode 100644 model-00002-of-00003.safetensors create mode 100644 model-00003-of-00003.safetensors create mode 100644 model.safetensors.index.json create mode 100644 qat_modules.py create mode 100644 quantization.py create mode 100644 special_tokens_map.json create mode 100644 tokenizer.json create mode 100644 tokenizer_config.json create mode 100644 training_args.bin diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..27f6db5 --- /dev/null +++ b/.gitattributes @@ -0,0 +1,53 @@ +*.7z filter=lfs diff=lfs merge=lfs -text +*.arrow filter=lfs diff=lfs merge=lfs -text + + +*.bz2 filter=lfs diff=lfs merge=lfs -text +*.ftz filter=lfs diff=lfs merge=lfs -text +*.gz filter=lfs diff=lfs merge=lfs -text +*.h5 filter=lfs diff=lfs merge=lfs -text +*.joblib filter=lfs diff=lfs merge=lfs -text +*.lfs.* filter=lfs diff=lfs merge=lfs -text +*.model filter=lfs diff=lfs merge=lfs -text +*.msgpack filter=lfs diff=lfs merge=lfs -text +*.onnx filter=lfs diff=lfs merge=lfs -text +*.ot filter=lfs diff=lfs merge=lfs -text +*.parquet filter=lfs diff=lfs merge=lfs -text +*.pb filter=lfs diff=lfs merge=lfs -text +*.pt filter=lfs diff=lfs merge=lfs -text +*.pth filter=lfs diff=lfs merge=lfs -text +*.rar filter=lfs diff=lfs merge=lfs -text +saved_model/**/* filter=lfs diff=lfs merge=lfs -text +*.tar.* filter=lfs diff=lfs merge=lfs -text +*.tflite filter=lfs diff=lfs merge=lfs -text +*.tgz filter=lfs diff=lfs merge=lfs -text +*.xz filter=lfs diff=lfs merge=lfs -text +*.zip filter=lfs diff=lfs merge=lfs -text +*.zstandard filter=lfs diff=lfs merge=lfs -text +*.tfevents* filter=lfs diff=lfs merge=lfs -text +*.db* filter=lfs diff=lfs merge=lfs -text +*.ark* filter=lfs diff=lfs merge=lfs -text +**/*ckpt*data* filter=lfs diff=lfs merge=lfs -text +**/*ckpt*.meta filter=lfs diff=lfs merge=lfs -text +**/*ckpt*.index filter=lfs diff=lfs merge=lfs -text + +*.ckpt filter=lfs diff=lfs merge=lfs -text +*.gguf* filter=lfs diff=lfs merge=lfs -text +*.ggml filter=lfs diff=lfs merge=lfs -text +*.llamafile* filter=lfs diff=lfs merge=lfs -text +*.pt2 filter=lfs diff=lfs merge=lfs -text +*.mlmodel filter=lfs diff=lfs merge=lfs -text +*.npy filter=lfs diff=lfs merge=lfs -text +*.npz filter=lfs diff=lfs merge=lfs -text +*.pickle filter=lfs diff=lfs merge=lfs -text +*.pkl filter=lfs diff=lfs merge=lfs -text +*.tar filter=lfs diff=lfs merge=lfs -text +*.wasm filter=lfs diff=lfs merge=lfs -text +*.zst filter=lfs diff=lfs merge=lfs -text +*tfevents* filter=lfs diff=lfs merge=lfs -text + +model-00002-of-00003.safetensors filter=lfs diff=lfs merge=lfs -text +training_args.bin filter=lfs diff=lfs merge=lfs -text +tokenizer.json filter=lfs diff=lfs merge=lfs -text +model-00001-of-00003.safetensors filter=lfs diff=lfs merge=lfs -text +model-00003-of-00003.safetensors filter=lfs diff=lfs merge=lfs -text \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..99dc09c --- /dev/null +++ b/README.md @@ -0,0 +1,144 @@ +--- +license: llama2 +base_model: meta-llama/Llama-2-7b-hf +tags: + - llama-2 + - quantization + - qat + - complex-valued + - 2-bit + - text-generation + - recursive + - safetensors +language: + - en +pipeline_tag: text-generation +--- + +# Fairy2i-W2 + +**🔗 Links** + +[![Paper](https://img.shields.io/badge/Paper-arXiv-red?logo=arxiv&logoColor=white)](https://arxiv.org/abs/2512.02901) [![GitHub](https://img.shields.io/badge/GitHub-181717?logo=github&logoColor=white)](https://github.com/PKULab1806/Fairy2i-W2) [![ModelScope](https://img.shields.io/badge/ModelScope-624AFF?logoColor=white)](https://modelscope.cn/models/PKULab1806/Fairy2i-W2) + +## Abstract + +Large language models (LLMs) have revolutionized artificial intelligence, yet their massive memory and computational demands necessitate aggressive quantization, increasingly pushing representations toward the theoretical limit of a single bit. While complex-valued LLMs, such as iFairy, offer a superior chance for low-bit representation compared to real-valued counterparts, they require training from scratch, preventing the utilization of the vast ecosystem of pre-trained real-valued foundation models. + +Here we present **Fairy2i**, a universal framework that transforms pre-trained real-valued layers into an equivalent widely-linear complex form, enabling extremely low-bit quantization while reusing existing checkpoints. By proving a lossless mathematical equivalence between real and widely-linear maps, we convert standard Transformers into the complex domain and employ a phase-aware quantization scheme with a highly efficient codebook of fourth roots of unity ({±1, ±i}). Furthermore, we introduce a recursive residual quantization mechanism that iteratively minimizes quantization error, allowing inference to proceed via efficient multiplication-free accumulation. + +We demonstrate that **Fairy2i-W2** restores the performance of LLaMA-2 7B at an effective 2-bit precision to levels nearly comparable with full-precision baselines, significantly outperforming state-of-the-art real-valued binary and ternary quantization methods. + +This work bridges the gap between the representational efficiency of complex-valued arithmetic and the practical utility of pre-trained models, paving a new way for efficient inference on commodity hardware. + +## Method + +Fairy2i-W2 consists of three key components: + +### Widely-Linear Transformation + +We transform pre-trained real-valued linear layers into an equivalent **widely-linear complex form** without altering the model's behavior. Each real linear layer R (a real matrix of size 2n×2m) is reparameterized into two complex matrices U and W (each of size n×m) such that y = Ux + Wx̅, where x̅ denotes the complex conjugate of x. This transformation is **lossless** and **unique**, preserving the original forward computation before quantization. + +### Phase-Aware Complex Quantization + +We quantize complex weights using a phase-based scheme with the codebook {±1, ±i} (fourth roots of unity). For each complex weight, we project it to the nearest codeword by angle and apply axis-wise scaling factors. During QAT training, we maintain full-precision master weights and use quantized copies in the forward pass with straight-through estimator (STE) gradients. + +### Recursive Residual Quantization + +To further reduce quantization error, we recursively quantize the residual error. Each complex weight is represented as a sum of low-bit terms: W_q ≈ Σ W^(t) (sum over t from 0 to T-1), where each term is quantized using the same phase-aware mechanism. For **Fairy2i-W2** (T=2), we use 2 recursive stages, achieving an effective **2 bits per real parameter**. + +## Evaluation + +### Main Results on LLaMA-2 7B + +| Method | Bits | C4 PPL↓ | ARC-e | ARC-c | HellaSwag | PIQA | Winogrande | Avg. | +|---------|------|---------|-------|-------|-----------|------|------------|------| +| LLaMA-2 (FP16) | 16 | 6.63 | 75.59 | 43.17 | 57.06 | 77.91 | 69.85 | 64.72 | +| **Fairy2i-W2** | **2** | **7.85** | **72.73** | **39.76** | **53.33** | **76.17** | **68.03** | **62.00** | +| AQLM | 2 | 8.54 | 63.68 | 32.76 | 49.55 | 74.76 | 65.67 | 57.28 | +| QuIP# | 2 | 11.01 | 55.56 | 28.84 | 42.94 | 71.38 | 62.43 | 52.23 | +| Real-Ternary (QAT) | 1.58 | 11.06 | 55.93 | 24.15 | 38.43 | 69.80 | 55.17 | 48.70 | +| **Fairy2i-W1** | **1** | **11.03** | **56.56** | **24.82** | **38.19** | **70.08** | **53.67** | **48.66** | +| Real-Binary (QAT) | 1 | 11.75 | 53.32 | 22.70 | 35.57 | 66.81 | 52.64 | 46.21 | +| GPTQ | 3 | 10.61 | 58.46 | 31.06 | 45.21 | 71.49 | 59.19 | 53.08 | + +**Key Results:** +- **Fairy2i-W2 (2-bit)** achieves a perplexity of 7.85, closing the gap to FP16 (6.63) while outperforming all 2-bit PTQ methods +- **Fairy2i-W2** achieves 62.00% average accuracy on zero-shot tasks, highly competitive with FP16 (64.72%) +- **Fairy2i-W1 (1-bit)** outperforms real-valued binary and ternary baselines at the same or lower bit budgets + +## Quick Start + +**Fairy2i-W2** is based on LLaMA-2 7B architecture, with only the linear layers replaced by complex-valued QAT layers. The model structure is otherwise identical to LLaMA-2. + +### Installation + +```bash +pip install torch transformers safetensors huggingface_hub +``` + +### Loading the Model + +Please refer to `load_model.py` for detailed implementation. Basic usage: + +```python +from load_model import load_model + +# Load Fairy2i-W2 model +model, tokenizer = load_model() + +# The model is ready to use! +prompt = "Hello, how are you?" +inputs = tokenizer(prompt, return_tensors="pt").to(model.device) + +with torch.no_grad(): + outputs = model.generate( + **inputs, + max_new_tokens=50, + do_sample=True, + temperature=0.7 + ) + +response = tokenizer.decode(outputs[0], skip_special_tokens=True) +print(response) +``` + +### Model Details + +- **Base Model**: LLaMA-2 7B +- **Quantization Method**: Complex-Phase V2 (2-step recursive residual quantization) +- **Effective Bit Width**: 2 bits per real parameter +- **Codebook**: {±1, ±i} (fourth roots of unity) +- **Training**: QAT (Quantization-Aware Training) on 30B tokens from RedPajama dataset + +### Files in Repository + +- `load_model.py`: Model loading script +- `qat_modules.py`: QAT linear layer implementations +- `quantization.py`: Quantization functions (PhaseQuant, BitNet, etc.) +- `config.json`: Model configuration (identical to LLaMA-2 7B) +- `model.safetensors.index.json`: Weight file index +- `model-0000X-of-00003.safetensors`: Sharded model weights +- Tokenizer files: `tokenizer.json`, `tokenizer_config.json`, etc. + +### Citation + +If you use Fairy2i-W2 in your research, please cite: + +```bibtex +@article{wang2025fairy2i, + title={Fairy2i: Training Complex LLMs from Real LLMs with All Parameters in {±1, ±i}}, + author={Wang, Feiyu and Tan, Xinyu and Huang, Bokai and Zhang, Yihao and Wang, Guoan and Cong, Peizhuang and Yang, Tong}, + journal={arXiv preprint}, + year={2025} +} +``` + +### License + +This model follows the same license as LLaMA-2. Please refer to the original LLaMA-2 license for details. + +### Contact + +For questions or issues, please contact: tanxinyu330@gmail.com + diff --git a/config.json b/config.json new file mode 100644 index 0000000..9761429 --- /dev/null +++ b/config.json @@ -0,0 +1,29 @@ +{ + "architectures": [ + "LlamaForCausalLM" + ], + "attention_bias": false, + "attention_dropout": 0.0, + "bos_token_id": 1, + "eos_token_id": 2, + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 4096, + "initializer_range": 0.02, + "intermediate_size": 11008, + "max_position_embeddings": 4096, + "mlp_bias": false, + "model_type": "llama", + "num_attention_heads": 32, + "num_hidden_layers": 32, + "num_key_value_heads": 32, + "pretraining_tp": 1, + "rms_norm_eps": 1e-05, + "rope_scaling": null, + "rope_theta": 10000.0, + "tie_word_embeddings": false, + "torch_dtype": "bfloat16", + "transformers_version": "4.52.4", + "use_cache": true, + "vocab_size": 32000 +} diff --git a/configuration.json b/configuration.json new file mode 100644 index 0000000..bbeeda1 --- /dev/null +++ b/configuration.json @@ -0,0 +1 @@ +{"framework": "pytorch", "task": "text-generation", "allow_remote": true} \ No newline at end of file diff --git a/generation_config.json b/generation_config.json new file mode 100644 index 0000000..7bc2dcc --- /dev/null +++ b/generation_config.json @@ -0,0 +1,10 @@ +{ + "bos_token_id": 1, + "do_sample": true, + "eos_token_id": 2, + "max_length": 4096, + "pad_token_id": 0, + "temperature": 0.6, + "top_p": 0.9, + "transformers_version": "4.52.4" +} diff --git a/load_model.py b/load_model.py new file mode 100644 index 0000000..b3df6bb --- /dev/null +++ b/load_model.py @@ -0,0 +1,143 @@ +""" +Model loading script + +Load Fairy2i-W2 model from Hugging Face repository. + +Usage: + from load_model import load_model + model, tokenizer = load_model() +""" + +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer +from safetensors.torch import load_file +from huggingface_hub import hf_hub_download +import os +import sys + +# Add current directory to path for importing qat_modules +current_dir = os.path.dirname(os.path.abspath(__file__)) +if current_dir not in sys.path: + sys.path.insert(0, current_dir) + +from qat_modules import replace_modules_for_qat + + +def load_model( + device="cuda" if torch.cuda.is_available() else "cpu", + torch_dtype=torch.bfloat16 +): + """ + Load Fairy2i-W2 model: standard architecture + custom weights + QAT linear layer replacement + + Load weights and tokenizer from Hugging Face repository. + + Args: + device: Device, default auto-select + torch_dtype: Data type, default torch.bfloat16 + + Returns: + model, tokenizer + """ + # Configuration parameters + base_model_id = "meta-llama/Llama-2-7b-hf" + weights_repo_id = "PKU-DS-LAB/Fairy2i-W2" + quant_method = "complex_phase_v2" + skip_lm_head = False + print("=" * 70) + print("Loading Fairy2i-W2 Model") + print("=" * 70) + + # Step 1: Load standard model architecture + print(f"\n📥 Step 1/4: Loading standard model architecture: {base_model_id}") + model = AutoModelForCausalLM.from_pretrained( + base_model_id, + torch_dtype=torch_dtype, + device_map=device, + trust_remote_code=False + ) + print("✅ Standard model architecture loaded") + + # Step 2: Load custom weights + print(f"\n💾 Step 2/4: Loading weights from Hugging Face repository: {weights_repo_id}") + + # Check for sharded weights + try: + index_path = hf_hub_download( + repo_id=weights_repo_id, + filename="model.safetensors.index.json", + local_dir=None + ) + + # Sharded weights + from safetensors import safe_open + import json + + with open(index_path, 'r') as f: + weight_map = json.load(f)["weight_map"] + + state_dict = {} + for weight_file in set(weight_map.values()): + file_path = hf_hub_download( + repo_id=weights_repo_id, + filename=weight_file, + local_dir=None + ) + with safe_open(file_path, framework="pt", device="cpu") as f: + for key in f.keys(): + state_dict[key] = f.get_tensor(key) + + model.load_state_dict(state_dict, strict=False) + print(f"✅ Weights loaded (sharded)") + except Exception: + # Single weight file + try: + weights_path = hf_hub_download( + repo_id=weights_repo_id, + filename="model.safetensors", + local_dir=None + ) + state_dict = load_file(weights_path) + model.load_state_dict(state_dict, strict=False) + print(f"✅ Weights loaded (single file)") + except Exception as e: + raise RuntimeError(f"Failed to load weights from Hugging Face: {e}") + + # Step 3: Apply QAT replacement + print(f"\n🔧 Step 3/4: Applying QAT replacement ({quant_method})...") + replace_modules_for_qat(model, quant_method, skip_lm_head=skip_lm_head) + print("✅ QAT replacement completed") + + # Step 4: Load tokenizer + print(f"\n📝 Step 4/4: Loading Tokenizer from Hugging Face repository: {weights_repo_id}") + tokenizer = AutoTokenizer.from_pretrained(weights_repo_id) + print("✅ Tokenizer loaded") + + print("\n" + "=" * 70) + print("✅ Model loading completed!") + print("=" * 70) + + return model, tokenizer + + +if __name__ == "__main__": + # Example: Load model + model, tokenizer = load_model() + + # Test generation + print("\n🧪 Testing generation...") + prompt = "Hello, how are you?" + inputs = tokenizer(prompt, return_tensors="pt").to(model.device) + + with torch.no_grad(): + outputs = model.generate( + **inputs, + max_new_tokens=50, + do_sample=True, + temperature=0.7 + ) + + response = tokenizer.decode(outputs[0], skip_special_tokens=True) + print(f"Prompt: {prompt}") + print(f"Response: {response}") + diff --git a/model-00001-of-00003.safetensors b/model-00001-of-00003.safetensors new file mode 100644 index 0000000..38b98dd --- /dev/null +++ b/model-00001-of-00003.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7faa5a86a5d9968018d5ec54582e4c8d953541fdda69dd3855cad02b0ba62f8c +size 4938985352 diff --git a/model-00002-of-00003.safetensors b/model-00002-of-00003.safetensors new file mode 100644 index 0000000..25c41f5 --- /dev/null +++ b/model-00002-of-00003.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c10a92f9aa154ead36f02411ed284ac5a4abcba038d43c3d3b4599b3c4fd132e +size 4947390880 diff --git a/model-00003-of-00003.safetensors b/model-00003-of-00003.safetensors new file mode 100644 index 0000000..41b705d --- /dev/null +++ b/model-00003-of-00003.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:33494b29047ea5f2696e949935b46b7e9372ec5de81d9718656ae8237766bf2d +size 3590488816 diff --git a/model.safetensors.index.json b/model.safetensors.index.json new file mode 100644 index 0000000..13674e5 --- /dev/null +++ b/model.safetensors.index.json @@ -0,0 +1,298 @@ +{ + "metadata": { + "total_size": 13476831232 + }, + "weight_map": { + "lm_head.weight": "model-00003-of-00003.safetensors", + "model.embed_tokens.weight": "model-00001-of-00003.safetensors", + "model.layers.0.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.0.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.0.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.0.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.0.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.0.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.0.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.0.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.0.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.1.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.1.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.1.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.10.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.10.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.10.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.11.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.11.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.11.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.11.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.11.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.11.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.11.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.11.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.11.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.12.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.12.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.12.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.12.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.12.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.12.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.12.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.12.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.12.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.13.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.13.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.13.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.14.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.14.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.14.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.15.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.15.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.15.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.16.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.16.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.16.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.17.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.17.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.17.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.18.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.18.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.18.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.19.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.19.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.19.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.2.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.2.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.2.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.2.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.2.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.2.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.2.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.2.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.2.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.20.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.20.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.20.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.20.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.20.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.20.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.20.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.20.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.20.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.21.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.21.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.21.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.input_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.22.mlp.down_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.post_attention_layernorm.weight": "model-00002-of-00003.safetensors", + "model.layers.22.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.22.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.23.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.23.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.23.mlp.gate_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.23.mlp.up_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.23.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.23.self_attn.k_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.23.self_attn.o_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.23.self_attn.q_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.23.self_attn.v_proj.weight": "model-00002-of-00003.safetensors", + "model.layers.24.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.24.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.24.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.24.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.24.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.24.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.24.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.24.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.24.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.25.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.25.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.25.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.26.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.26.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.26.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.27.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.27.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.27.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.28.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.28.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.28.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.29.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.29.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.29.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.3.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.3.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.3.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.3.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.3.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.3.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.3.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.3.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.3.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.30.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.30.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.30.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.30.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.30.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.30.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.30.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.30.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.30.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.input_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.31.mlp.down_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.mlp.gate_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.mlp.up_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.post_attention_layernorm.weight": "model-00003-of-00003.safetensors", + "model.layers.31.self_attn.k_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.self_attn.o_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.self_attn.q_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.31.self_attn.v_proj.weight": "model-00003-of-00003.safetensors", + "model.layers.4.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.4.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.4.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.4.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.4.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.4.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.4.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.4.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.4.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.5.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.5.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.5.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.6.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.6.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.6.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.7.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.7.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.7.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.8.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.8.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.8.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.input_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.9.mlp.down_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.mlp.gate_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.mlp.up_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.post_attention_layernorm.weight": "model-00001-of-00003.safetensors", + "model.layers.9.self_attn.k_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.self_attn.o_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.self_attn.q_proj.weight": "model-00001-of-00003.safetensors", + "model.layers.9.self_attn.v_proj.weight": "model-00001-of-00003.safetensors", + "model.norm.weight": "model-00003-of-00003.safetensors" + } +} diff --git a/qat_modules.py b/qat_modules.py new file mode 100644 index 0000000..dc31ea5 --- /dev/null +++ b/qat_modules.py @@ -0,0 +1,388 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +from quantization import BitNetQuantSTE, PhaseQuantSTE, PhaseQuantSTE_V2, PhaseQuantSTE_V3, PhaseQuantSTE_V4 +import math + +class QATLinearBitNet(nn.Linear): + """BitNet QAT linear layer""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def forward(self, x): + quantized_weight = BitNetQuantSTE.apply(self.weight) + return F.linear(x, quantized_weight, self.bias) + +class QATLinearComplexPhaseV1(nn.Linear): + """Complex-Phase V1 QAT linear layer""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if self.in_features % 2 != 0 or self.out_features % 2 != 0: + raise ValueError("Complex-Phase QAT requires even in/out features for Linear layers.") + + def forward(self, x): + A = self.weight + n, m = A.shape[0] // 2, A.shape[1] // 2 + A11, A12 = A[:n, :m], A[:n, m:] + A21, A22 = A[n:, :m], A[n:, m:] + + U_re = 0.5 * (A11 + A22) + U_im = 0.5 * (A21 - A12) + W_re = 0.5 * (A11 - A22) + W_im = 0.5 * (A12 + A21) + + U_re_q, U_im_q = PhaseQuantSTE.apply(U_re, U_im) + W_re_q, W_im_q = PhaseQuantSTE.apply(W_re, W_im) + + A11_q = W_re_q + U_re_q + A12_q = W_im_q - U_im_q + A21_q = W_im_q + U_im_q + A22_q = -W_re_q + U_re_q + + A_quant_top = torch.cat([A11_q, A12_q], dim=1) + A_quant_bottom = torch.cat([A21_q, A22_q], dim=1) + A_quant = torch.cat([A_quant_top, A_quant_bottom], dim=0) + + return F.linear(x, A_quant, self.bias) + +class QATLinearComplexPhaseV2(nn.Linear): + """Complex-Phase V2 QAT linear layer (1-step residual)""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if self.in_features % 2 != 0 or self.out_features % 2 != 0: + raise ValueError("Complex-Phase QAT requires even in/out features for Linear layers.") + + def forward(self, x): + A = self.weight + n, m = A.shape[0] // 2, A.shape[1] // 2 + A11, A12 = A[:n, :m], A[:n, m:] + A21, A22 = A[n:, :m], A[n:, m:] + + U_re = 0.5 * (A11 + A22) + U_im = 0.5 * (A21 - A12) + W_re = 0.5 * (A11 - A22) + W_im = 0.5 * (A12 + A21) + + U_re_q, U_im_q = PhaseQuantSTE_V2.apply(U_re, U_im) + W_re_q, W_im_q = PhaseQuantSTE_V2.apply(W_re, W_im) + + A11_q = W_re_q + U_re_q + A12_q = W_im_q - U_im_q + A21_q = W_im_q + U_im_q + A22_q = -W_re_q + U_re_q + + A_quant_top = torch.cat([A11_q, A12_q], dim=1) + A_quant_bottom = torch.cat([A21_q, A22_q], dim=1) + A_quant = torch.cat([A_quant_top, A_quant_bottom], dim=0) + + return F.linear(x, A_quant, self.bias) + +class QATLinearComplexPhaseV3(nn.Linear): + """Complex-Phase V3 QAT linear layer (2-step residual)""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if self.in_features % 2 != 0 or self.out_features % 2 != 0: + raise ValueError("Complex-Phase QAT requires even in/out features for Linear layers.") + + def forward(self, x): + A = self.weight + n, m = A.shape[0] // 2, A.shape[1] // 2 + A11, A12 = A[:n, :m], A[:n, m:] + A21, A22 = A[n:, :m], A[n:, m:] + + U_re = 0.5 * (A11 + A22) + U_im = 0.5 * (A21 - A12) + W_re = 0.5 * (A11 - A22) + W_im = 0.5 * (A12 + A21) + + U_re_q, U_im_q = PhaseQuantSTE_V3.apply(U_re, U_im) + W_re_q, W_im_q = PhaseQuantSTE_V3.apply(W_re, W_im) + + A11_q = W_re_q + U_re_q + A12_q = W_im_q - U_im_q + A21_q = W_im_q + U_im_q + A22_q = -W_re_q + U_re_q + + A_quant_top = torch.cat([A11_q, A12_q], dim=1) + A_quant_bottom = torch.cat([A21_q, A22_q], dim=1) + A_quant = torch.cat([A_quant_top, A_quant_bottom], dim=0) + + return F.linear(x, A_quant, self.bias) + +class QATLinearComplexPhaseV4(nn.Linear): + """Complex-Phase V4 QAT linear layer (3-step residual)""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if self.in_features % 2 != 0 or self.out_features % 2 != 0: + raise ValueError("Complex-Phase QAT requires even in/out features for Linear layers.") + + def forward(self, x): + A = self.weight + n, m = A.shape[0] // 2, A.shape[1] // 2 + A11, A12 = A[:n, :m], A[:n, m:] + A21, A22 = A[n:, :m], A[n:, m:] + + U_re = 0.5 * (A11 + A22) + U_im = 0.5 * (A21 - A12) + W_re = 0.5 * (A11 - A22) + W_im = 0.5 * (A12 + A21) + + U_re_q, U_im_q = PhaseQuantSTE_V4.apply(U_re, U_im) + W_re_q, W_im_q = PhaseQuantSTE_V4.apply(W_re, W_im) + + A11_q = W_re_q + U_re_q + A12_q = W_im_q - U_im_q + A21_q = W_im_q + U_im_q + A22_q = -W_re_q + U_re_q + + A_quant_top = torch.cat([A11_q, A12_q], dim=1) + A_quant_bottom = torch.cat([A21_q, A22_q], dim=1) + A_quant = torch.cat([A_quant_top, A_quant_bottom], dim=0) + + return F.linear(x, A_quant, self.bias) + +METHOD_MAP = { + 'bitnet': QATLinearBitNet, + 'complex_phase_v1': QATLinearComplexPhaseV1, + 'complex_phase_v2': QATLinearComplexPhaseV2, + 'complex_phase_v3': QATLinearComplexPhaseV3, + 'complex_phase_v4': QATLinearComplexPhaseV4, +} + +def replace_modules_for_qat(model: nn.Module, method: str, skip_lm_head: bool = False): + """Recursively replace nn.Linear layers in the model with QAT layers""" + if method not in METHOD_MAP: + raise ValueError(f"Unknown method: {method}. Available methods: {list(METHOD_MAP.keys())}") + + TargetQATClass = METHOD_MAP[method] + + for name, module in model.named_children(): + if len(list(module.children())) > 0: + replace_modules_for_qat(module, method, skip_lm_head) + + if isinstance(module, nn.Linear): + if skip_lm_head and name == 'lm_head': + print(f" -> Skipping lm_head layer (skip_lm_head=True)") + continue + + if 'complex_phase' in method: + if module.in_features % 2 != 0 or module.out_features % 2 != 0: + print(f" -> Skipping Complex-Phase replacement (non-even dimensions): {name} ({module.in_features}, {module.out_features})") + continue + + print(f" -> Replacing layer: {name} with {TargetQATClass.__name__}") + new_module = TargetQATClass( + module.in_features, + module.out_features, + bias=module.bias is not None, + dtype=module.weight.dtype, + device=module.weight.device + ) + new_module.weight.data.copy_(module.weight.data) + if module.bias is not None: + new_module.bias.data.copy_(module.bias.data) + + setattr(model, name, new_module) + +class InferenceOptimizedBitNet(nn.Linear): + """Inference-optimized BitNet linear layer, in-place weight replacement to save memory""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._is_quantized = False + + def _ensure_quantized(self): + """Ensure weights are quantized, executed only once""" + if not self._is_quantized: + with torch.no_grad(): + w = self.weight + scale = w.abs().mean() + alpha = w.mean() + centered_w = w - alpha + binarized_w = torch.where(centered_w > 0, 1.0, -1.0).to(w.dtype) + quantized_w = binarized_w * scale + self.weight.data = quantized_w + self._is_quantized = True + + def forward(self, x): + self._ensure_quantized() + return F.linear(x, self.weight, self.bias) + +class InferenceOptimizedComplexPhase(nn.Linear): + """Inference-optimized Complex Phase linear layer, supports V1-V4""" + def __init__(self, version="v1", *args, **kwargs): + super().__init__(*args, **kwargs) + if self.in_features % 2 != 0 or self.out_features % 2 != 0: + raise ValueError("Complex-Phase requires even in/out features.") + self._is_quantized = False + self._version = version.lower() + if self._version not in ["v1", "v2", "v3", "v4"]: + raise ValueError(f"Unsupported version: {version}. Must be one of ['v1', 'v2', 'v3', 'v4']") + + def _ensure_quantized(self): + """Ensure weights are quantized, executed only once""" + if not self._is_quantized: + with torch.no_grad(): + A = self.weight + n, m = A.shape[0] // 2, A.shape[1] // 2 + A11, A12 = A[:n, :m], A[:n, m:] + A21, A22 = A[n:, :m], A[n:, m:] + + U_re = 0.5 * (A11 + A22) + U_im = 0.5 * (A21 - A12) + W_re = 0.5 * (A11 - A22) + W_im = 0.5 * (A12 + A21) + + if self._version == "v1": + U_re_q, U_im_q = self._phase_quant_v1(U_re, U_im) + W_re_q, W_im_q = self._phase_quant_v1(W_re, W_im) + elif self._version == "v2": + U_re_q, U_im_q = self._phase_quant_v2(U_re, U_im) + W_re_q, W_im_q = self._phase_quant_v2(W_re, W_im) + elif self._version == "v3": + U_re_q, U_im_q = self._phase_quant_v3(U_re, U_im) + W_re_q, W_im_q = self._phase_quant_v3(W_re, W_im) + elif self._version == "v4": + U_re_q, U_im_q = self._phase_quant_v4(U_re, U_im) + W_re_q, W_im_q = self._phase_quant_v4(W_re, W_im) + + A11_q = W_re_q + U_re_q + A12_q = W_im_q - U_im_q + A21_q = W_im_q + U_im_q + A22_q = -W_re_q + U_re_q + + A_quant_top = torch.cat([A11_q, A12_q], dim=1) + A_quant_bottom = torch.cat([A21_q, A22_q], dim=1) + A_quant = torch.cat([A_quant_top, A_quant_bottom], dim=0) + + self.weight.data = A_quant + self._is_quantized = True + + def _phase_quant_v1(self, w_real, w_imag): + """V1: Basic PhaseQuant""" + phase = torch.angle(w_real + 1j * w_imag) + + real_pos = (phase >= -math.pi / 4) & (phase < math.pi / 4) + real_neg = (phase >= 3 * math.pi / 4) | (phase < -3 * math.pi / 4) + imag_pos = (phase >= math.pi / 4) & (phase < 3 * math.pi / 4) + imag_neg = (phase >= -3 * math.pi / 4) & (phase < -math.pi / 4) + + mask_real = real_pos | real_neg + mask_imag = imag_pos | imag_neg + + s_re = w_real[mask_real].abs().mean() if mask_real.any() else torch.tensor(0.0, device=w_real.device) + s_im = w_imag[mask_imag].abs().mean() if mask_imag.any() else torch.tensor(0.0, device=w_imag.device) + + s_re = torch.clamp(s_re, min=1e-6) + s_im = torch.clamp(s_im, min=1e-6) + + qw_real = torch.zeros_like(w_real) + qw_imag = torch.zeros_like(w_imag) + + qw_real[real_pos] = 1.0 + qw_real[real_neg] = -1.0 + qw_imag[imag_pos] = 1.0 + qw_imag[imag_neg] = -1.0 + + return qw_real * s_re, qw_imag * s_im + + def _phase_quant_v2(self, w_real, w_imag): + """V2: 1-step residual quantization""" + qw_real_o1, qw_imag_o1 = self._phase_quant_v1(w_real, w_imag) + error_real = w_real - qw_real_o1 + error_imag = w_imag - qw_imag_o1 + qw_real_o2, qw_imag_o2 = self._phase_quant_v1(error_real, error_imag) + qw_real = qw_real_o1 + qw_real_o2 + qw_imag = qw_imag_o1 + qw_imag_o2 + return qw_real, qw_imag + + def _phase_quant_v3(self, w_real, w_imag): + """V3: 2-step residual quantization""" + qw_real_o1, qw_imag_o1 = self._phase_quant_v1(w_real, w_imag) + error_real_1 = w_real - qw_real_o1 + error_imag_1 = w_imag - qw_imag_o1 + qw_real_o2, qw_imag_o2 = self._phase_quant_v1(error_real_1, error_imag_1) + error_real_2 = error_real_1 - qw_real_o2 + error_imag_2 = error_imag_1 - qw_imag_o2 + qw_real_o3, qw_imag_o3 = self._phase_quant_v1(error_real_2, error_imag_2) + qw_real = qw_real_o1 + qw_real_o2 + qw_real_o3 + qw_imag = qw_imag_o1 + qw_imag_o2 + qw_imag_o3 + return qw_real, qw_imag + + def _phase_quant_v4(self, w_real, w_imag): + """V4: 3-step residual quantization""" + qw_real_o1, qw_imag_o1 = self._phase_quant_v1(w_real, w_imag) + error_real_1 = w_real - qw_real_o1 + error_imag_1 = w_imag - qw_imag_o1 + qw_real_o2, qw_imag_o2 = self._phase_quant_v1(error_real_1, error_imag_1) + error_real_2 = error_real_1 - qw_real_o2 + error_imag_2 = error_imag_1 - qw_imag_o2 + qw_real_o3, qw_imag_o3 = self._phase_quant_v1(error_real_2, error_imag_2) + error_real_3 = error_real_2 - qw_real_o3 + error_imag_3 = error_imag_2 - qw_imag_o3 + qw_real_o4, qw_imag_o4 = self._phase_quant_v1(error_real_3, error_imag_3) + qw_real = qw_real_o1 + qw_real_o2 + qw_real_o3 + qw_real_o4 + qw_imag = qw_imag_o1 + qw_imag_o2 + qw_imag_o3 + qw_imag_o4 + return qw_real, qw_imag + + def forward(self, x): + self._ensure_quantized() + return F.linear(x, self.weight, self.bias) + +def convert_to_inference_mode(model): + """Convert QAT modules to inference-optimized version (permanently modifies model weights)""" + converted_count = 0 + + def _convert_module(module, name_path=""): + nonlocal converted_count + + for name, child in list(module.named_children()): + full_name = f"{name_path}.{name}" if name_path else name + + if isinstance(child, QATLinearBitNet): + new_module = InferenceOptimizedBitNet( + child.in_features, + child.out_features, + bias=child.bias is not None, + device=child.weight.device, + dtype=child.weight.dtype + ) + new_module.weight.data.copy_(child.weight.data) + if child.bias is not None: + new_module.bias.data.copy_(child.bias.data) + + setattr(module, name, new_module) + converted_count += 1 + print(f" -> Converting BitNet layer: {full_name}") + + elif isinstance(child, (QATLinearComplexPhaseV1, QATLinearComplexPhaseV2, + QATLinearComplexPhaseV3, QATLinearComplexPhaseV4)): + if isinstance(child, QATLinearComplexPhaseV1): + version = "v1" + elif isinstance(child, QATLinearComplexPhaseV2): + version = "v2" + elif isinstance(child, QATLinearComplexPhaseV3): + version = "v3" + elif isinstance(child, QATLinearComplexPhaseV4): + version = "v4" + + new_module = InferenceOptimizedComplexPhase( + version=version, + in_features=child.in_features, + out_features=child.out_features, + bias=child.bias is not None, + device=child.weight.device, + dtype=child.weight.dtype + ) + new_module.weight.data.copy_(child.weight.data) + if child.bias is not None: + new_module.bias.data.copy_(child.bias.data) + + setattr(module, name, new_module) + converted_count += 1 + print(f" -> Converting ComplexPhase{version.upper()} layer: {full_name}") + else: + _convert_module(child, full_name) + + _convert_module(model) + print(f"Converted {converted_count} QAT layers to inference-optimized version") + return model \ No newline at end of file diff --git a/quantization.py b/quantization.py new file mode 100644 index 0000000..f68f180 --- /dev/null +++ b/quantization.py @@ -0,0 +1,308 @@ +import torch +import torch.nn as nn +import math + +@torch.no_grad() +def quantize_complex_tensor(w_real: torch.Tensor, w_imag: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Apply PhaseQuant logic to complex weight tensors""" + phase = torch.angle(w_real + 1j * w_imag) + + real_pos = (phase >= -math.pi / 4) & (phase < math.pi / 4) + real_neg = (phase >= 3 * math.pi / 4) | (phase < -3 * math.pi / 4) + imag_pos = (phase >= math.pi / 4) & (phase < 3 * math.pi / 4) + imag_neg = (phase >= -3 * math.pi / 4) & (phase < -math.pi / 4) + + mask_real = real_pos | real_neg + mask_imag = imag_pos | imag_neg + + s_re = w_real[mask_real].abs().mean() if mask_real.any() else torch.tensor(0.0, device=w_real.device) + s_im = w_imag[mask_imag].abs().mean() if mask_imag.any() else torch.tensor(0.0, device=w_imag.device) + + s_re = torch.clamp(s_re, min=1e-6) + s_im = torch.clamp(s_im, min=1e-6) + if torch.isnan(s_re) or torch.isinf(s_re): s_re = torch.tensor(1e-6, device=w_real.device) + if torch.isnan(s_im) or torch.isinf(s_im): s_im = torch.tensor(1e-6, device=w_imag.device) + + qw_real = torch.zeros_like(w_real) + qw_imag = torch.zeros_like(w_imag) + + qw_real[real_pos] = 1.0 + qw_real[real_neg] = -1.0 + qw_imag[imag_pos] = 1.0 + qw_imag[imag_neg] = -1.0 + + qw_real_scaled = qw_real * s_re + qw_imag_scaled = qw_imag * s_im + return qw_real_scaled.to(w_real.dtype), qw_imag_scaled.to(w_imag.dtype) + +def apply_complex_inspired_quantization(model: nn.Module): + """Apply complex-inspired quantization to real-valued model""" + print("Applying complex-inspired quantization (PhaseQuant-based)...") + + @torch.no_grad() + def quantize_linear_layer(module: nn.Linear): + A = module.weight.data + if A.shape[0] % 2 != 0 or A.shape[1] % 2 != 0: + print(f" -> Skipping layer (non-even dimensions): {A.shape}") + return + + n, m = A.shape[0] // 2, A.shape[1] // 2 + A11, A12 = A[:n, :m], A[:n, m:] + A21, A22 = A[n:, :m], A[n:, m:] + + U_re = 0.5 * (A11 + A22) + U_im = 0.5 * (A21 - A12) + W_re = 0.5 * (A11 - A22) + W_im = 0.5 * (A12 + A21) + + U_re_q, U_im_q = quantize_complex_tensor(U_re, U_im) + W_re_q, W_im_q = quantize_complex_tensor(W_re, W_im) + + A11_q = W_re_q + U_re_q + A12_q = W_im_q - U_im_q + A21_q = W_im_q + U_im_q + A22_q = -W_re_q + U_re_q + + A_quant_top = torch.cat([A11_q, A12_q], dim=1) + A_quant_bottom = torch.cat([A21_q, A22_q], dim=1) + A_quant = torch.cat([A_quant_top, A_quant_bottom], dim=0) + + module.weight.data = A_quant.to(A.dtype) + + model.apply(lambda module: quantize_linear_layer(module) if isinstance(module, nn.Linear) else None) + print("Complex-inspired quantization completed.") + return model + +def apply_bitnet_quantization(model: nn.Module): + """Apply BitNet 1-bit quantization to real-valued model""" + print("Applying BitNet (true 1-bit, affine) quantization to real-valued model...") + + @torch.no_grad() + def quantize_linear_layer(module: nn.Linear): + scale = module.weight.data.abs().mean() + alpha = module.weight.data.mean() + centered_weights = module.weight.data - alpha + binarized_weights = torch.where(centered_weights > 0, 1.0, -1.0) + module.weight.data = binarized_weights.to(module.weight.data.dtype) * scale + + model.apply(lambda module: quantize_linear_layer(module) if isinstance(module, nn.Linear) else None) + print("BitNet quantization completed.") + return model + +def apply_bitnet_1_58bit_quantization_standard(model: nn.Module): + """Apply BitNet 1.58-bit quantization to real-valued model (quantize to {-1, 0, +1})""" + print("Applying BitNet 1.58-bit (absmean threshold) quantization to real-valued model...") + + @torch.no_grad() + def quantize_linear_layer(module: nn.Linear): + W = module.weight.data + gamma = W.abs().mean() + W_normalized = W / (gamma + 1e-5) + W_quantized = torch.clamp(torch.round(W_normalized), -1.0, 1.0) + module.weight.data = W_quantized.to(W.dtype) * gamma + + model.apply(lambda module: quantize_linear_layer(module) if isinstance(module, nn.Linear) else None) + print("BitNet 1.58-bit (absmean threshold) quantization completed.") + return model + +def apply_bitnet_1_58bit_quantization_variant(model: nn.Module, threshold: float = 0.5): + """Apply BitNet 1.58-bit quantization to real-valued model (quantize to {-1, 0, +1})""" + print("Applying BitNet 1.58-bit (ternary) quantization to real-valued model...") + + @torch.no_grad() + def quantize_linear_layer(module: nn.Linear): + gamma = module.weight.data.abs().mean() + normalized_weights = module.weight.data / (gamma + 1e-5) + adaptive_threshold = threshold + ternary_weights = torch.zeros_like(normalized_weights) + ternary_weights[normalized_weights > adaptive_threshold] = 1.0 + ternary_weights[normalized_weights < -adaptive_threshold] = -1.0 + module.weight.data = ternary_weights.to(module.weight.data.dtype) * gamma + + model.apply(lambda module: quantize_linear_layer(module) if isinstance(module, nn.Linear) else None) + print("BitNet 1.58-bit quantization completed.") + return model + +def minmax_1bit_quantize_dequantize(w: torch.Tensor) -> torch.Tensor: + """Apply 1-bit Min-Max quantization and dequantization to weight tensor""" + min_val = w.min() + max_val = w.max() + scale = (max_val - min_val) / 1.0 + zero_point = min_val + + if abs(scale) < 1e-9: + return w + + quantized_w = torch.round((w - zero_point) / scale) + dequantized_w = quantized_w * scale + zero_point + + return dequantized_w.to(w.dtype) + +def apply_minmax_1bit_quantization(model: nn.Module): + """Apply Min-Max 1-bit quantization to real-valued model""" + print("Applying Min-Max (1-bit) quantization to real-valued model...") + + @torch.no_grad() + def quantize_linear_layer(module: nn.Linear): + module.weight.data = minmax_1bit_quantize_dequantize(module.weight.data) + + model.apply(lambda module: quantize_linear_layer(module) if isinstance(module, nn.Linear) else None) + + print("Min-Max 1-bit quantization completed.") + return model + +def symmetric_minmax_1bit_quantize_dequantize(w: torch.Tensor) -> torch.Tensor: + """Apply symmetric 1-bit Min-Max quantization to weight tensor (quantize to {-1, 1})""" + max_abs = w.abs().max() + scale = max_abs + + if scale < 1e-9: + return w + + quantized_w = (w / scale).sign() + dequantized_w = quantized_w * scale + + return dequantized_w.to(w.dtype) + +def apply_symmetric_minmax_1bit_quantization(model: nn.Module): + """Apply symmetric Min-Max 1-bit quantization to real-valued model""" + print("Applying symmetric Min-Max (1-bit, to {-1, 1}) quantization to real-valued model...") + + @torch.no_grad() + def quantize_linear_layer(module: nn.Linear): + module.weight.data = symmetric_minmax_1bit_quantize_dequantize(module.weight.data) + + model.apply(lambda module: quantize_linear_layer(module) if isinstance(module, nn.Linear) else None) + + print("Symmetric Min-Max 1-bit quantization completed.") + return model + +class BitNetQuantSTE(torch.autograd.Function): + """BitNet STE: quantize in forward, pass gradients in backward""" + @staticmethod + def forward(ctx, w): + scale = w.abs().mean() + alpha = w.mean() + centered_w = w - alpha + binarized_w = torch.where(centered_w > 0, 1.0, -1.0).to(w.dtype) + quantized_w = binarized_w * scale + return quantized_w + + @staticmethod + def backward(ctx, grad_output): + return grad_output + + +class BitNet1_58QuantSTE(torch.autograd.Function): + """BitNet 1.58-bit STE: quantize to {-1, 0, +1}, pass gradients in backward""" + @staticmethod + def forward(ctx, w): + gamma = w.abs().mean() + w_normalized = w / (gamma + 1e-5) + w_quantized = torch.clamp(torch.round(w_normalized), -1.0, 1.0) + quantized_w = (w_quantized * gamma).to(w.dtype) + return quantized_w + + @staticmethod + def backward(ctx, grad_output): + return grad_output + + +class PhaseQuantSTE(torch.autograd.Function): + """Complex-Phase STE: quantize in forward, pass gradients in backward""" + @staticmethod + def forward(ctx, w_real, w_imag): + phase = torch.angle(w_real + 1j * w_imag) + + real_pos = (phase >= -math.pi / 4) & (phase < math.pi / 4) + real_neg = (phase >= 3 * math.pi / 4) | (phase < -3 * math.pi / 4) + imag_pos = (phase >= math.pi / 4) & (phase < 3 * math.pi / 4) + imag_neg = (phase >= -3 * math.pi / 4) & (phase < -math.pi / 4) + + mask_real = real_pos | real_neg + mask_imag = imag_pos | imag_neg + + s_re = w_real[mask_real].abs().mean() if mask_real.any() else torch.tensor(0.0, device=w_real.device) + s_im = w_imag[mask_imag].abs().mean() if mask_imag.any() else torch.tensor(0.0, device=w_imag.device) + + s_re = torch.clamp(s_re, min=1e-6) + s_im = torch.clamp(s_im, min=1e-6) + + qw_real = torch.zeros_like(w_real) + qw_imag = torch.zeros_like(w_imag) + + qw_real[real_pos] = 1.0 + qw_real[real_neg] = -1.0 + qw_imag[imag_pos] = 1.0 + qw_imag[imag_neg] = -1.0 + + qw_real_scaled = qw_real * s_re + qw_imag_scaled = qw_imag * s_im + + return qw_real_scaled.to(w_real.dtype), qw_imag_scaled.to(w_imag.dtype) + + @staticmethod + def backward(ctx, grad_w_real, grad_w_imag): + return grad_w_real, grad_w_imag + +class PhaseQuantSTE_V2(torch.autograd.Function): + """Two-step residual quantization""" + + @staticmethod + def forward(ctx, w_real: torch.Tensor, w_imag: torch.Tensor): + qw_real_o1, qw_imag_o1 = PhaseQuantSTE.apply(w_real, w_imag) + error_real = w_real - qw_real_o1 + error_imag = w_imag - qw_imag_o1 + qw_real_o2, qw_imag_o2 = PhaseQuantSTE.apply(error_real, error_imag) + qw_real = qw_real_o1 + qw_real_o2 + qw_imag = qw_imag_o1 + qw_imag_o2 + return qw_real, qw_imag + + @staticmethod + def backward(ctx, grad_real, grad_imag): + return grad_real, grad_imag + + +class PhaseQuantSTE_V3(torch.autograd.Function): + """Three-step residual quantization""" + + @staticmethod + def forward(ctx, w_real: torch.Tensor, w_imag: torch.Tensor): + qw_real_o1, qw_imag_o1 = PhaseQuantSTE.apply(w_real, w_imag) + error_real_1 = w_real - qw_real_o1 + error_imag_1 = w_imag - qw_imag_o1 + qw_real_o2, qw_imag_o2 = PhaseQuantSTE.apply(error_real_1, error_imag_1) + error_real_2 = error_real_1 - qw_real_o2 + error_imag_2 = error_imag_1 - qw_imag_o2 + qw_real_o3, qw_imag_o3 = PhaseQuantSTE.apply(error_real_2, error_imag_2) + qw_real = qw_real_o1 + qw_real_o2 + qw_real_o3 + qw_imag = qw_imag_o1 + qw_imag_o2 + qw_imag_o3 + return qw_real, qw_imag + + @staticmethod + def backward(ctx, grad_real, grad_imag): + return grad_real, grad_imag + + +class PhaseQuantSTE_V4(torch.autograd.Function): + """Four-step residual quantization""" + + @staticmethod + def forward(ctx, w_real: torch.Tensor, w_imag: torch.Tensor): + qw_real_o1, qw_imag_o1 = PhaseQuantSTE.apply(w_real, w_imag) + error_real_1 = w_real - qw_real_o1 + error_imag_1 = w_imag - qw_imag_o1 + qw_real_o2, qw_imag_o2 = PhaseQuantSTE.apply(error_real_1, error_imag_1) + error_real_2 = error_real_1 - qw_real_o2 + error_imag_2 = error_imag_1 - qw_imag_o2 + qw_real_o3, qw_imag_o3 = PhaseQuantSTE.apply(error_real_2, error_imag_2) + error_real_3 = error_real_2 - qw_real_o3 + error_imag_3 = error_imag_2 - qw_imag_o3 + qw_real_o4, qw_imag_o4 = PhaseQuantSTE.apply(error_real_3, error_imag_3) + qw_real = qw_real_o1 + qw_real_o2 + qw_real_o3 + qw_real_o4 + qw_imag = qw_imag_o1 + qw_imag_o2 + qw_imag_o3 + qw_imag_o4 + return qw_real, qw_imag + + @staticmethod + def backward(ctx, grad_real, grad_imag): + return grad_real, grad_imag diff --git a/special_tokens_map.json b/special_tokens_map.json new file mode 100644 index 0000000..72ecfee --- /dev/null +++ b/special_tokens_map.json @@ -0,0 +1,24 @@ +{ + "bos_token": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false + }, + "eos_token": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false + }, + "pad_token": "", + "unk_token": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false + } +} diff --git a/tokenizer.json b/tokenizer.json new file mode 100644 index 0000000..bdc9404 --- /dev/null +++ b/tokenizer.json @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c88bda6bdd84543eadebdf4bd2ec325ae43b71a3ff1fead6a766562c00b29bd8 +size 3619016 diff --git a/tokenizer_config.json b/tokenizer_config.json new file mode 100644 index 0000000..2952c7c --- /dev/null +++ b/tokenizer_config.json @@ -0,0 +1,43 @@ +{ + "add_bos_token": true, + "add_eos_token": false, + "add_prefix_space": null, + "added_tokens_decoder": { + "0": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "1": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "2": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + } + }, + "bos_token": "", + "clean_up_tokenization_spaces": false, + "eos_token": "", + "extra_special_tokens": {}, + "legacy": false, + "model_max_length": 1000000000000000019884624838656, + "pad_token": "", + "padding_side": "right", + "sp_model_kwargs": {}, + "tokenizer_class": "LlamaTokenizer", + "unk_token": "", + "use_default_system_prompt": false +} diff --git a/training_args.bin b/training_args.bin new file mode 100644 index 0000000..2152a12 --- /dev/null +++ b/training_args.bin @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e9aacb37f66b32727647179ccb831ff830505e9721317432e224e6a2abb2dae7 +size 6929