初始化项目,由ModelHub XC社区提供模型
Model: PKU-DS-LAB/Fairy2i-W2 Source: Original Platform
This commit is contained in:
53
.gitattributes
vendored
Normal file
53
.gitattributes
vendored
Normal file
@@ -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
|
||||
144
README.md
Normal file
144
README.md
Normal file
@@ -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**
|
||||
|
||||
[](https://arxiv.org/abs/2512.02901) [](https://github.com/PKULab1806/Fairy2i-W2) [](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
|
||||
|
||||
29
config.json
Normal file
29
config.json
Normal file
@@ -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
|
||||
}
|
||||
1
configuration.json
Normal file
1
configuration.json
Normal file
@@ -0,0 +1 @@
|
||||
{"framework": "pytorch", "task": "text-generation", "allow_remote": true}
|
||||
10
generation_config.json
Normal file
10
generation_config.json
Normal file
@@ -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"
|
||||
}
|
||||
143
load_model.py
Normal file
143
load_model.py
Normal file
@@ -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}")
|
||||
|
||||
3
model-00001-of-00003.safetensors
Normal file
3
model-00001-of-00003.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7faa5a86a5d9968018d5ec54582e4c8d953541fdda69dd3855cad02b0ba62f8c
|
||||
size 4938985352
|
||||
3
model-00002-of-00003.safetensors
Normal file
3
model-00002-of-00003.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:c10a92f9aa154ead36f02411ed284ac5a4abcba038d43c3d3b4599b3c4fd132e
|
||||
size 4947390880
|
||||
3
model-00003-of-00003.safetensors
Normal file
3
model-00003-of-00003.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:33494b29047ea5f2696e949935b46b7e9372ec5de81d9718656ae8237766bf2d
|
||||
size 3590488816
|
||||
298
model.safetensors.index.json
Normal file
298
model.safetensors.index.json
Normal file
@@ -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"
|
||||
}
|
||||
}
|
||||
388
qat_modules.py
Normal file
388
qat_modules.py
Normal file
@@ -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
|
||||
308
quantization.py
Normal file
308
quantization.py
Normal file
@@ -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
|
||||
24
special_tokens_map.json
Normal file
24
special_tokens_map.json
Normal file
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"bos_token": {
|
||||
"content": "<s>",
|
||||
"lstrip": false,
|
||||
"normalized": false,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
},
|
||||
"eos_token": {
|
||||
"content": "</s>",
|
||||
"lstrip": false,
|
||||
"normalized": false,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
},
|
||||
"pad_token": "</s>",
|
||||
"unk_token": {
|
||||
"content": "<unk>",
|
||||
"lstrip": false,
|
||||
"normalized": false,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
}
|
||||
}
|
||||
3
tokenizer.json
Normal file
3
tokenizer.json
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:c88bda6bdd84543eadebdf4bd2ec325ae43b71a3ff1fead6a766562c00b29bd8
|
||||
size 3619016
|
||||
43
tokenizer_config.json
Normal file
43
tokenizer_config.json
Normal file
@@ -0,0 +1,43 @@
|
||||
{
|
||||
"add_bos_token": true,
|
||||
"add_eos_token": false,
|
||||
"add_prefix_space": null,
|
||||
"added_tokens_decoder": {
|
||||
"0": {
|
||||
"content": "<unk>",
|
||||
"lstrip": false,
|
||||
"normalized": false,
|
||||
"rstrip": false,
|
||||
"single_word": false,
|
||||
"special": true
|
||||
},
|
||||
"1": {
|
||||
"content": "<s>",
|
||||
"lstrip": false,
|
||||
"normalized": false,
|
||||
"rstrip": false,
|
||||
"single_word": false,
|
||||
"special": true
|
||||
},
|
||||
"2": {
|
||||
"content": "</s>",
|
||||
"lstrip": false,
|
||||
"normalized": false,
|
||||
"rstrip": false,
|
||||
"single_word": false,
|
||||
"special": true
|
||||
}
|
||||
},
|
||||
"bos_token": "<s>",
|
||||
"clean_up_tokenization_spaces": false,
|
||||
"eos_token": "</s>",
|
||||
"extra_special_tokens": {},
|
||||
"legacy": false,
|
||||
"model_max_length": 1000000000000000019884624838656,
|
||||
"pad_token": "</s>",
|
||||
"padding_side": "right",
|
||||
"sp_model_kwargs": {},
|
||||
"tokenizer_class": "LlamaTokenizer",
|
||||
"unk_token": "<unk>",
|
||||
"use_default_system_prompt": false
|
||||
}
|
||||
3
training_args.bin
Normal file
3
training_args.bin
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:e9aacb37f66b32727647179ccb831ff830505e9721317432e224e6a2abb2dae7
|
||||
size 6929
|
||||
Reference in New Issue
Block a user