CoolFace
Apppublic

rakesh9177/Quantization

sourceHugging Faceupdated 3y agoView on Hugging Face
2likes
Untitled.ipynb192 linesDownload Raw Back to root
1{2 "cells": [3  {4   "cell_type": "code",5   "execution_count": 1,6   "id": "25d9b213-f625-4327-98b3-d9e67db11687",7   "metadata": {},8   "outputs": [9    {10     "name": "stdout",11     "output_type": "stream",12     "text": [13      "saving model\n",14      "Model Saved\n"15     ]16    },17    {18     "data": {19      "text/plain": [20       "'\\ngenerator = pipeline(\\'text-generation\\', model=\"facebook/opt-1.3b\")\\n\\ndef generate_text_pip(prompt):\\n    generated_text = generator(prompt, max_length=1000, num_return_sequences=1)[0][\\'generated_text\\']\\n    return generated_text\\nprint(generator(\"I went to boston and\"))\\n\\ndef generate_text(prompt):\\n    inputs = tokenizer(prompt, return_tensors=\"pt\")\\n    output = model(**inputs)\\n    logits = output.logits\\n    predicted_ids = logits.argmax(-1)\\n    generated_text = tokenizer.decode(predicted_ids[0], skip_special_tokens=True)\\n    return generated_text\\n\\ndef generate_text_from_quantized(prompt):\\n    inputs = tokenizer(prompt, return_tensors=\"pt\")\\n    output = model_q(**inputs)\\n    logits = output.logits\\n    predicted_ids = logits.argmax(-1)\\n    generated_text = tokenizer.decode(predicted_ids[0], skip_special_tokens=True)\\n    return generated_text\\n\\n# Create a Gradio interface\\niface = gr.Interface(fn=generate_text_pip, inputs=\"text\", outputs=\"text\", live=True)\\n\\niface_2 = gr.Interface(fn=generate_text_from_quantized, inputs=\"text\", outputs=\"text\", live=True)\\n\\n\\napp = gr.TabbedInterface([iface, iface_2],[\"Normal\", \"Quantized\"])\\n\\n# Launch the Gradio app\\napp.launch()\\n'"21      ]22     },23     "execution_count": 1,24     "metadata": {},25     "output_type": "execute_result"26    }27   ],28   "source": [29    "import gradio as gr\n",30    "import tqdm\n",31    "import torch\n",32    "from torch import nn\n",33    "from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline\n",34    "from functools import partial\n",35    "import gc\n",36    "\n",37    "\n",38    "# core quantization method (simulated quantization)\n",39    "def pseudo_quantize_tensor(w, n_bit=4, q_group_size=-1):\n",40    "    org_w_shape = w.shape\n",41    "    if q_group_size > 0:\n",42    "        assert org_w_shape[-1] % q_group_size == 0\n",43    "        w = w.reshape(-1, q_group_size)\n",44    "\n",45    "    assert w.dim() == 2\n",46    "\n",47    "    # Calculate the maximum (\\alpha) and minimum values (\\beta) in the tensor.\n",48    "    max_val = w.amax(dim=1, keepdim=True)\n",49    "    assert max_val.dim() == 2 and max_val.size(0) == w.size(0) and max_val.size(1) == 1\n",50    "    min_val = w.amin(dim=1, keepdim=True)\n",51    "    assert min_val.dim() == 2 and min_val.size(0) == w.size(0) and min_val.size(1) == 1\n",52    "\n",53    "    # Calculate the scale factor and zero point.  (Formula 1 & 2)\n",54    "    max_int = 2 ** n_bit - 1\n",55    "    scales = (max_val - min_val).clamp(min=1e-5) / max_int\n",56    "    assert scales.shape == max_val.shape\n",57    "    zeros = (-torch.round(min_val / scales)).clamp_(0, max_int)\n",58    "    assert scales.shape == min_val.shape\n",59    "\n",60    "    assert torch.isnan(scales).sum() == 0\n",61    "    assert torch.isnan(w).sum() == 0\n",62    "\n",63    "    # Quantize W: Map values in the range [\\beta, \\alpha] to lie within [0, 2^b - 1] (Formula 3)\n",64    "    w = torch.clamp(torch.round(w / scales) + zeros, 0, max_int)\n",65    "    assert w.dim() == 2 and w.size(0) == scales.size(0) and w.size(1) == q_group_size\n",66    "\n",67    "    # Dequantize W (pseudo quantization, the inverse transformation of Formula 3)\n",68    "    w = (w - zeros) * scales\n",69    "    assert w.dim() == 2 and w.size(0) == scales.size(0) and w.size(1) == q_group_size\n",70    "\n",71    "    assert torch.isnan(w).sum() == 0\n",72    "\n",73    "    w = w.reshape(org_w_shape)\n",74    "    return w\n",75    "\n",76    "@torch.no_grad()\n",77    "def pseudo_quantize_model_weight(\n",78    "    model, w_bit, q_group_size,\n",79    "):\n",80    "    for n, m in model.named_modules():\n",81    "        if isinstance(m, nn.Linear):\n",82    "            m.weight.data = pseudo_quantize_tensor(m.weight.data, n_bit=w_bit, q_group_size=q_group_size)\n",83    "            \n",84    "            \n",85    "    \n",86    "            \n",87    "# Load the tokenizer and model\n",88    "model_path = \"facebook/opt-125m\"\n",89    "offload_folder = \"offload\"\n",90    "tokenizer = AutoTokenizer.from_pretrained(model_path, use_fast=False)\n",91    "model = AutoModelForCausalLM.from_pretrained(model_path, device_map=\"auto\", offload_folder=offload_folder)\n",92    "model_q = AutoModelForCausalLM.from_pretrained(model_path, device_map=\"auto\",offload_folder=offload_folder)\n",93    "pseudo_quantize_model_weight(model_q, w_bit=3, q_group_size=128)\n",94    "# Define a function for model inference\n",95    "\n",96    "\n",97    "quantized_model_path = \"facebook/opt-125m_3bit\"\n",98    "print(\"saving model\")\n",99    "model_q.save_pretrained(quantized_model_path)\n",100    "tokenizer.save_pretrained(quantized_model_path)\n",101    "print(\"Model Saved\")\n",102    "'''\n",103    "generator = pipeline('text-generation', model=\"facebook/opt-1.3b\")\n",104    "\n",105    "def generate_text_pip(prompt):\n",106    "    generated_text = generator(prompt, max_length=1000, num_return_sequences=1)[0]['generated_text']\n",107    "    return generated_text\n",108    "print(generator(\"I went to boston and\"))\n",109    "\n",110    "def generate_text(prompt):\n",111    "    inputs = tokenizer(prompt, return_tensors=\"pt\")\n",112    "    output = model(**inputs)\n",113    "    logits = output.logits\n",114    "    predicted_ids = logits.argmax(-1)\n",115    "    generated_text = tokenizer.decode(predicted_ids[0], skip_special_tokens=True)\n",116    "    return generated_text\n",117    "\n",118    "def generate_text_from_quantized(prompt):\n",119    "    inputs = tokenizer(prompt, return_tensors=\"pt\")\n",120    "    output = model_q(**inputs)\n",121    "    logits = output.logits\n",122    "    predicted_ids = logits.argmax(-1)\n",123    "    generated_text = tokenizer.decode(predicted_ids[0], skip_special_tokens=True)\n",124    "    return generated_text\n",125    "\n",126    "# Create a Gradio interface\n",127    "iface = gr.Interface(fn=generate_text_pip, inputs=\"text\", outputs=\"text\", live=True)\n",128    "\n",129    "iface_2 = gr.Interface(fn=generate_text_from_quantized, inputs=\"text\", outputs=\"text\", live=True)\n",130    "\n",131    "\n",132    "app = gr.TabbedInterface([iface, iface_2],[\"Normal\", \"Quantized\"])\n",133    "\n",134    "# Launch the Gradio app\n",135    "app.launch()\n",136    "'''"137   ]138  },139  {140   "cell_type": "code",141   "execution_count": 6,142   "id": "5b3171de-b63f-4f5e-8347-c4a5da79c397",143   "metadata": {},144   "outputs": [145    {146     "name": "stdout",147     "output_type": "stream",148     "text": [149      "quantized model [{'generated_text': 'I went to boston and was hoping for a good time for a good time for a good time'}]\n"150     ]151    }152   ],153   "source": [154    "model_q_path = \"facebook/opt-125m_3bit\"\n",155    "\n",156    "\n",157    "generator_q = pipeline('text-generation', model=model_q_path)\n",158    "print(\"quantized model\",generator_q(\"I went to boston and\"))"159   ]160  },161  {162   "cell_type": "code",163   "execution_count": null,164   "id": "8e969157-5261-4245-a6d0-d394e971b347",165   "metadata": {},166   "outputs": [],167   "source": []168  }169 ],170 "metadata": {171  "kernelspec": {172   "display_name": "Python 3 (ipykernel)",173   "language": "python",174   "name": "python3"175  },176  "language_info": {177   "codemirror_mode": {178    "name": "ipython",179    "version": 3180   },181   "file_extension": ".py",182   "mimetype": "text/x-python",183   "name": "python",184   "nbconvert_exporter": "python",185   "pygments_lexer": "ipython3",186   "version": "3.9.12"187  }188 },189 "nbformat": 4,190 "nbformat_minor": 5191}192