Team Ai
Apppublic

OpenVINO/nncf-quantization

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
21likes
app.py287 linesDownload Raw Back to root
1import os2import shutil3import gradio as gr4from huggingface_hub import HfApi, whoami, ModelCard, model_info5from gradio_huggingfacehub_search import HuggingfaceHubSearch6from textwrap import dedent7from pathlib import Path8 9from tempfile import TemporaryDirectory10 11from huggingface_hub.file_download import repo_folder_name12from optimum.exporters import TasksManager13from optimum.intel import (14    OVModelForAudioClassification,15    OVModelForCausalLM,16    OVModelForFeatureExtraction,17    OVModelForImageClassification,18    OVModelForMaskedLM,19    OVModelForQuestionAnswering,20    OVModelForSeq2SeqLM,21    OVModelForSequenceClassification,22    OVModelForTokenClassification,23    OVStableDiffusionPipeline,24    OVStableDiffusionXLPipeline,25    OVLatentConsistencyModelPipeline,26    OVWeightQuantizationConfig,27)28from diffusers import ConfigMixin29 30_HEAD_TO_AUTOMODELS = {31    "feature-extraction": "OVModelForFeatureExtraction",32    "fill-mask": "OVModelForMaskedLM",33    "text-generation": "OVModelForCausalLM",34    "text-classification": "OVModelForSequenceClassification",35    "token-classification": "OVModelForTokenClassification",36    "question-answering": "OVModelForQuestionAnswering",37    "image-classification": "OVModelForImageClassification",38    "audio-classification": "OVModelForAudioClassification",39    "stable-diffusion": "OVStableDiffusionPipeline",40    "stable-diffusion-xl": "OVStableDiffusionXLPipeline",41    "latent-consistency": "OVLatentConsistencyModelPipeline",42}43 44def quantize_model(45    model_id: str,46    dtype: str,47    calibration_dataset: str,48    ratio: str,49    private_repo: bool,50    overwritte: bool,51    oauth_token: gr.OAuthToken,52):53    if oauth_token.token is None:54        return "You must be logged in to use this space"55 56    if not model_id:57        return f"### Invalid input ๐Ÿž Please specify a model name, got {model_id}"58 59    try:60        model_name = model_id.split("/")[-1]61        username = whoami(oauth_token.token)["name"]62        w_t = dtype.replace("-", "")63        suffix = f"{w_t}" if model_name.endswith("openvino") else f"openvino-{w_t}"64        new_repo_id = f"{username}/{model_name}-{suffix}"65        library_name = TasksManager.infer_library_from_model(model_id, token=oauth_token.token)66 67        if library_name == "diffusers":68            ConfigMixin.config_name = "model_index.json"69            class_name = ConfigMixin.load_config(model_id, token=oauth_token.token)["_class_name"].lower()70            if "xl" in class_name:71                task = "stable-diffusion-xl"72            elif "consistency" in class_name:73                task = "latent-consistency"74            else:75                task = "stable-diffusion"76        else:77            task = TasksManager.infer_task_from_model(model_id, token=oauth_token.token)78 79        if task == "text2text-generation":80            return "Export of Seq2Seq models is currently disabled."81 82        if task not in _HEAD_TO_AUTOMODELS:83            return f"The task '{task}' is not supported, only {_HEAD_TO_AUTOMODELS.keys()} tasks are supported"84 85        auto_model_class = _HEAD_TO_AUTOMODELS[task]86        if calibration_dataset == "None":87            calibration_dataset = None88 89        is_int8 = dtype == "8-bit"90        # if library_name == "diffusers":91        # quant_method = "hybrid"92        if not is_int8 and calibration_dataset is not None:93            quant_method = "awq"94        else:95            if calibration_dataset is not None:96                print("Default quantization was selected, calibration dataset won't be used")97            quant_method = "default"98 99        quantization_config = OVWeightQuantizationConfig(100            bits=8 if is_int8 else 4,101            quant_method=quant_method,102            dataset=None if quant_method=="default" else calibration_dataset,103            ratio=1.0 if is_int8 else ratio,104            num_samples=None if quant_method=="default" else 20,105        )106 107        api = HfApi(token=oauth_token.token)108        if api.repo_exists(new_repo_id) and not overwritte:109            return f"Model {new_repo_id} already exist, please tick the overwritte box to push on an existing repository"110 111        with TemporaryDirectory() as d:112            folder = os.path.join(d, repo_folder_name(repo_id=model_id, repo_type="models"))113            os.makedirs(folder)114 115            try:116                api.snapshot_download(repo_id=model_id, local_dir=folder, allow_patterns=["*.json"])117                ov_model = eval(auto_model_class).from_pretrained(118                    model_id,119                    cache_dir=folder,120                    token=oauth_token.token,121                    quantization_config=quantization_config122                )123                ov_model.save_pretrained(folder)124                new_repo_url = api.create_repo(repo_id=new_repo_id, exist_ok=True, private=private_repo)125                new_repo_id = new_repo_url.repo_id126                print("Repository created successfully!", new_repo_url)127 128                folder = Path(folder)129                for dir_name in (130                    "",131                    "vae_encoder",132                    "vae_decoder",133                    "text_encoder",134                    "text_encoder_2",135                    "unet",136                    "tokenizer",137                    "tokenizer_2",138                    "scheduler",139                    "feature_extractor",140                ):141                    if not (folder / dir_name).is_dir():142                        continue143                    for file_path in (folder / dir_name).iterdir():144                        if file_path.is_file():145                            try:146                                api.upload_file(147                                    path_or_fileobj=file_path,148                                    path_in_repo=os.path.join(dir_name, file_path.name),149                                    repo_id=new_repo_id,150                                )151                            except Exception as e:152                                return f"Error uploading file {file_path}: {e}"153 154                try:155                    card = ModelCard.load(model_id, token=oauth_token.token)156                except:157                    card = ModelCard("")158 159                if card.data.tags is None:160                    card.data.tags = []161                if "openvino" not in card.data.tags:162                    card.data.tags.append("openvino")163                card.data.tags.append("nncf")164                card.data.tags.append(dtype)165                card.data.base_model = model_id166 167                card.text = dedent(168                    f"""169                    This model is a quantized version of [`{model_id}`](https://huggingface.co/{model_id}) and is converted to the OpenVINO format. This model was obtained via the [nncf-quantization](https://huggingface.co/spaces/echarlaix/nncf-quantization) space with [optimum-intel](https://github.com/huggingface/optimum-intel).170 171                    First make sure you have `optimum-intel` installed:172 173                    ```bash174                    pip install optimum[openvino]175                    ```176 177                    To load your model you can do as follows:178 179                    ```python180                    from optimum.intel import {auto_model_class}181 182                    model_id = "{new_repo_id}"183                    model = {auto_model_class}.from_pretrained(model_id)184                    ```185                    """186                )187                card_path = os.path.join(folder, "README.md")188                card.save(card_path)189 190                api.upload_file(191                    path_or_fileobj=card_path,192                    path_in_repo="README.md",193                    repo_id=new_repo_id,194                )195                return f"This model was successfully quantized, find it under your repository {new_repo_url}"196            finally:197                shutil.rmtree(folder, ignore_errors=True)198    except Exception as e:199        return f"### Error: {e}"200 201DESCRIPTION = """202This Space uses [Optimum Intel](https://github.com/huggingface/optimum-intel) to automatically apply NNCF [Weight Only Quantization](https://huggingface.co/docs/optimum/main/en/intel/openvino/optimization) (WOQ) on your model and convert it to the [OpenVINO format](https://docs.openvino.ai/2024/documentation/openvino-ir-format.html) if not already.203 204After conversion, a repository will be pushed under your namespace with the resulting model.205 206The list of the supported architectures can be found in the [documentation](https://huggingface.co/docs/optimum/main/en/intel/openvino/models)207"""208 209model_id = HuggingfaceHubSearch(210    label="Hub Model ID",211    placeholder="Search for model id on the hub",212    search_type="model",213)214dtype = gr.Dropdown(215    ["8-bit", "4-bit"],216    value="8-bit",217    label="Weights precision",218    filterable=False,219    visible=True,220)221"""222quant_method = gr.Dropdown(223    ["default", "awq", "hybrid"],224    value="default",225    label="Quantization method",226    filterable=False,227    visible=True,228)229"""230calibration_dataset = gr.Dropdown(231    [232        "None",233        "wikitext2",234        "c4",235        "c4-new",236        "conceptual_captions",237        "laion/220k-GPT4Vision-captions-from-LIVIS",238        "laion/filtered-wit",239    ],240    value="None",241    label="Calibration dataset",242    filterable=False,243    visible=True,244)245ratio = gr.Slider(246    label="Ratio",247    info="Parameter used when applying 4-bit quantization to control the ratio between 4-bit and 8-bit quantization",248    minimum=0.0,249    maximum=1.0,250    step=0.1,251    value=1.0,252)253private_repo = gr.Checkbox(254    value=False,255    label="Private repository",256    info="Create a private repository instead of a public one",257)258overwritte = gr.Checkbox(259    value=False,260    label="Overwrite repository content",261    info="Enable pushing files on existing repositories, potentially overwriting existing files",262)263interface = gr.Interface(264    fn=quantize_model,265    inputs=[266        model_id,267        dtype,268        calibration_dataset,269        ratio,270        private_repo,271        overwritte,272    ],273    outputs=[274        gr.Markdown(label="output"),275    ],276    title="Quantize your model with NNCF",277    description=DESCRIPTION,278    api_name=False,279)280 281with gr.Blocks() as demo:282    gr.Markdown("You must be logged in to use this space")283    gr.LoginButton(min_width=250)284    interface.render()285 286demo.launch()287