Team Ai
Apppublic

quantumiracle-git/OpenBiDexHand

sourceHugging Faceupdated 4y agoView on Hugging Face
5likes
hfserver.py551 linesDownload Raw Back to root
1from __future__ import annotations2 3import csv4import datetime5import io6import json7import os8import uuid9from abc import ABC, abstractmethod10from typing import TYPE_CHECKING, Any, List, Optional11 12import gradio as gr13from gradio import encryptor, utils14from gradio.documentation import document, set_documentation_group15 16if TYPE_CHECKING:17    from gradio.components import IOComponent18 19set_documentation_group("flagging")20 21 22def _get_dataset_features_info(is_new, components):23    """24    Takes in a list of components and returns a dataset features info25    Parameters:26    is_new: boolean, whether the dataset is new or not27    components: list of components28    Returns:29    infos: a dictionary of the dataset features30    file_preview_types: dictionary mapping of gradio components to appropriate string.31    header: list of header strings32    """33    infos = {"flagged": {"features": {}}}34    # File previews for certain input and output types35    file_preview_types = {gr.Audio: "Audio", gr.Image: "Image"}36    headers = []37 38    # Generate the headers and dataset_infos39    if is_new:40 41        for component in components:42            headers.append(component.label)43            infos["flagged"]["features"][component.label] = {44                "dtype": "string",45                "_type": "Value",46            }47            if isinstance(component, tuple(file_preview_types)):48                headers.append(component.label + " file")49                for _component, _type in file_preview_types.items():50                    if isinstance(component, _component):51                        infos["flagged"]["features"][component.label + " file"] = {52                            "_type": _type53                        }54                        break55 56        headers.append("flag")57        infos["flagged"]["features"]["flag"] = {58            "dtype": "string",59            "_type": "Value",60        }61 62    return infos, file_preview_types, headers63 64 65class FlaggingCallback(ABC):66    """67    An abstract class for defining the methods that any FlaggingCallback should have.68    """69 70    @abstractmethod71    def setup(self, components: List[IOComponent], flagging_dir: str):72        """73        This method should be overridden and ensure that everything is set up correctly for flag().74        This method gets called once at the beginning of the Interface.launch() method.75        Parameters:76        components: Set of components that will provide flagged data.77        flagging_dir: A string, typically containing the path to the directory where the flagging file should be storied (provided as an argument to Interface.__init__()).78        """79        pass80 81    @abstractmethod82    def flag(83        self,84        flag_data: List[Any],85        flag_option: Optional[str] = None,86        flag_index: Optional[int] = None,87        username: Optional[str] = None,88    ) -> int:89        """90        This method should be overridden by the FlaggingCallback subclass and may contain optional additional arguments.91        This gets called every time the <flag> button is pressed.92        Parameters:93        interface: The Interface object that is being used to launch the flagging interface.94        flag_data: The data to be flagged.95        flag_option (optional): In the case that flagging_options are provided, the flag option that is being used.96        flag_index (optional): The index of the sample that is being flagged.97        username (optional): The username of the user that is flagging the data, if logged in.98        Returns:99        (int) The total number of samples that have been flagged.100        """101        pass102 103 104@document()105class SimpleCSVLogger(FlaggingCallback):106    """107    A simplified implementation of the FlaggingCallback abstract class108    provided for illustrative purposes.  Each flagged sample (both the input and output data)109    is logged to a CSV file on the machine running the gradio app.110    Example:111        import gradio as gr112        def image_classifier(inp):113            return {'cat': 0.3, 'dog': 0.7}114        demo = gr.Interface(fn=image_classifier, inputs="image", outputs="label",115                            flagging_callback=SimpleCSVLogger())116    """117 118    def __init__(self):119        pass120 121    def setup(self, components: List[IOComponent], flagging_dir: str):122        self.components = components123        self.flagging_dir = flagging_dir124        os.makedirs(flagging_dir, exist_ok=True)125 126    def flag(127        self,128        flag_data: List[Any],129        flag_option: Optional[str] = None,130        flag_index: Optional[int] = None,131        username: Optional[str] = None,132    ) -> int:133        flagging_dir = self.flagging_dir134        log_filepath = os.path.join(flagging_dir, "log.csv")135 136        csv_data = []137        for component, sample in zip(self.components, flag_data):138            save_dir = os.path.join(139                flagging_dir, utils.strip_invalid_filename_characters(component.label)140            )141            csv_data.append(142                component.deserialize(143                    sample,144                    save_dir,145                    None,146                )147            )148 149        with open(log_filepath, "a", newline="") as csvfile:150            writer = csv.writer(csvfile)151            writer.writerow(utils.sanitize_list_for_csv(csv_data))152 153        with open(log_filepath, "r") as csvfile:154            line_count = len([None for row in csv.reader(csvfile)]) - 1155        return line_count156 157 158@document()159class CSVLogger(FlaggingCallback):160    """161    The default implementation of the FlaggingCallback abstract class. Each flagged162    sample (both the input and output data) is logged to a CSV file with headers on the machine running the gradio app.163    Example:164        import gradio as gr165        def image_classifier(inp):166            return {'cat': 0.3, 'dog': 0.7}167        demo = gr.Interface(fn=image_classifier, inputs="image", outputs="label",168                            flagging_callback=CSVLogger())169    Guides: using_flagging170    """171 172    def __init__(self):173        pass174 175    def setup(176        self,177        components: List[IOComponent],178        flagging_dir: str,179        encryption_key: Optional[str] = None,180    ):181        self.components = components182        self.flagging_dir = flagging_dir183        self.encryption_key = encryption_key184        os.makedirs(flagging_dir, exist_ok=True)185 186    def flag(187        self,188        flag_data: List[Any],189        flag_option: Optional[str] = None,190        flag_index: Optional[int] = None,191        username: Optional[str] = None,192    ) -> int:193        flagging_dir = self.flagging_dir194        log_filepath = os.path.join(flagging_dir, "log.csv")195        is_new = not os.path.exists(log_filepath)196 197        if flag_index is None:198            csv_data = []199            for idx, (component, sample) in enumerate(zip(self.components, flag_data)):200                save_dir = os.path.join(201                    flagging_dir,202                    utils.strip_invalid_filename_characters(203                        component.label or f"component {idx}"204                    ),205                )206                if utils.is_update(sample):207                    csv_data.append(str(sample))208                else:209                    csv_data.append(210                        component.deserialize(211                            sample,212                            save_dir=save_dir,213                            encryption_key=self.encryption_key,214                        )215                        if sample is not None216                        else ""217                    )218            csv_data.append(flag_option if flag_option is not None else "")219            csv_data.append(username if username is not None else "")220            csv_data.append(str(datetime.datetime.now()))221            if is_new:222                headers = [223                    component.label or f"component {idx}"224                    for idx, component in enumerate(self.components)225                ] + [226                    "flag",227                    "username",228                    "timestamp",229                ]230 231        def replace_flag_at_index(file_content):232            file_content = io.StringIO(file_content)233            content = list(csv.reader(file_content))234            header = content[0]235            flag_col_index = header.index("flag")236            content[flag_index][flag_col_index] = flag_option237            output = io.StringIO()238            writer = csv.writer(output)239            writer.writerows(utils.sanitize_list_for_csv(content))240            return output.getvalue()241 242        if self.encryption_key:243            output = io.StringIO()244            if not is_new:245                with open(log_filepath, "rb", encoding="utf-8") as csvfile:246                    encrypted_csv = csvfile.read()247                    decrypted_csv = encryptor.decrypt(248                        self.encryption_key, encrypted_csv249                    )250                    file_content = decrypted_csv.decode()251                    if flag_index is not None:252                        file_content = replace_flag_at_index(file_content)253                    output.write(file_content)254            writer = csv.writer(output)255            if flag_index is None:256                if is_new:257                    writer.writerow(utils.sanitize_list_for_csv(headers))258                writer.writerow(utils.sanitize_list_for_csv(csv_data))259            with open(log_filepath, "wb", encoding="utf-8") as csvfile:260                csvfile.write(261                    encryptor.encrypt(self.encryption_key, output.getvalue().encode())262                )263        else:264            if flag_index is None:265                with open(log_filepath, "a", newline="", encoding="utf-8") as csvfile:266                    writer = csv.writer(csvfile)267                    if is_new:268                        writer.writerow(utils.sanitize_list_for_csv(headers))269                    writer.writerow(utils.sanitize_list_for_csv(csv_data))270            else:271                with open(log_filepath, encoding="utf-8") as csvfile:272                    file_content = csvfile.read()273                    file_content = replace_flag_at_index(file_content)274                with open(275                    log_filepath, "w", newline="", encoding="utf-8"276                ) as csvfile:  # newline parameter needed for Windows277                    csvfile.write(utils.sanitize_list_for_csv(file_content))278        with open(log_filepath, "r", encoding="utf-8") as csvfile:279            line_count = len([None for row in csv.reader(csvfile)]) - 1280        return line_count281 282 283@document()284class HuggingFaceDatasetSaver(FlaggingCallback):285    """286    A callback that saves each flagged sample (both the input and output data)287    to a HuggingFace dataset.288    Example:289        import gradio as gr290        hf_writer = gr.HuggingFaceDatasetSaver(HF_API_TOKEN, "image-classification-mistakes")291        def image_classifier(inp):292            return {'cat': 0.3, 'dog': 0.7}293        demo = gr.Interface(fn=image_classifier, inputs="image", outputs="label",294                            allow_flagging="manual", flagging_callback=hf_writer)295    Guides: using_flagging296    """297 298    def __init__(299        self,300        hf_token: str,301        dataset_name: str,302        organization: Optional[str] = None,303        private: bool = False,304    ):305        """306        Parameters:307            hf_token: The HuggingFace token to use to create (and write the flagged sample to) the HuggingFace dataset.308            dataset_name: The name of the dataset to save the data to, e.g. "image-classifier-1"309            organization: The organization to save the dataset under. The hf_token must provide write access to this organization. If not provided, saved under the name of the user corresponding to the hf_token.310            private: Whether the dataset should be private (defaults to False).311        """312        self.hf_token = hf_token313        self.dataset_name = dataset_name314        self.organization_name = organization315        self.dataset_private = private316 317    def setup(self, components: List[IOComponent], flagging_dir: str):318        """319        Params:320        flagging_dir (str): local directory where the dataset is cloned,321        updated, and pushed from.322        """323        try:324            import huggingface_hub325        except (ImportError, ModuleNotFoundError):326            raise ImportError(327                "Package `huggingface_hub` not found is needed "328                "for HuggingFaceDatasetSaver. Try 'pip install huggingface_hub'."329            )330        path_to_dataset_repo = huggingface_hub.create_repo(331            # name=self.dataset_name,332            repo_id=self.dataset_name,333            token=self.hf_token,334            private=self.dataset_private,335            repo_type="dataset",336            exist_ok=True,337        )338        self.path_to_dataset_repo = path_to_dataset_repo  # e.g. "https://huggingface.co/datasets/abidlabs/test-audio-10"339        self.components = components340        self.flagging_dir = flagging_dir341        self.dataset_dir = os.path.join(flagging_dir, self.dataset_name)342        self.repo = huggingface_hub.Repository(343            local_dir=self.dataset_dir,344            clone_from=path_to_dataset_repo,345            use_auth_token=self.hf_token,346        )347        self.repo.git_pull(lfs=True)348 349        # Should filename be user-specified?350        self.log_file = os.path.join(self.dataset_dir, "data.csv")351        self.infos_file = os.path.join(self.dataset_dir, "dataset_infos.json")352 353    def flag(354        self,355        flag_data: List[Any],356        flag_option: Optional[str] = None,357        flag_index: Optional[int] = None,358        username: Optional[str] = None,359    ) -> int:360        self.repo.git_pull(lfs=True)361 362        is_new = not os.path.exists(self.log_file)363 364        with open(self.log_file, "a", newline="", encoding="utf-8") as csvfile:365            writer = csv.writer(csvfile)366 367            # File previews for certain input and output types368            infos, file_preview_types, headers = _get_dataset_features_info(369                is_new, self.components370            )371 372            # Generate the headers and dataset_infos373            if is_new:374                writer.writerow(utils.sanitize_list_for_csv(headers))375 376            # Generate the row corresponding to the flagged sample377            csv_data = []378            for component, sample in zip(self.components, flag_data):379                save_dir = os.path.join(380                    self.dataset_dir,381                    utils.strip_invalid_filename_characters(component.label),382                )383                # filepath = component.deserialize(sample, save_dir, None)384                if sample is not None and str(component)!='image':385                    filepath = component.deserialize(sample, save_dir, None)386                else:387                    filepath = component.deserialize(sample, None, None)  # not saving image388                csv_data.append(filepath)389                if isinstance(component, tuple(file_preview_types)):390                    csv_data.append(391                        "{}/resolve/main/{}".format(self.path_to_dataset_repo, filepath)392                    )393            csv_data.append(flag_option if flag_option is not None else "")394            writer.writerow(utils.sanitize_list_for_csv(csv_data))395 396        if is_new:397            json.dump(infos, open(self.infos_file, "w"))398 399        with open(self.log_file, "r", encoding="utf-8") as csvfile:400            line_count = len([None for row in csv.reader(csvfile)]) - 1401 402        self.repo.push_to_hub(commit_message="Flagged sample #{}".format(line_count))403 404        return line_count405 406 407class HuggingFaceDatasetJSONSaver(FlaggingCallback):408    """409    A FlaggingCallback that saves flagged data to a Hugging Face dataset in JSONL format.410    Each data sample is saved in a different JSONL file,411    allowing multiple users to use flagging simultaneously.412    Saving to a single CSV would cause errors as only one user can edit at the same time.413    """414 415    def __init__(416        self,417        hf_foken: str,418        dataset_name: str,419        organization: Optional[str] = None,420        private: bool = False,421        verbose: bool = True,422    ):423        """424        Params:425        hf_token (str): The token to use to access the huggingface API.426        dataset_name (str): The name of the dataset to save the data to, e.g.427            "image-classifier-1"428        organization (str): The name of the organization to which to attach429            the datasets. If None, the dataset attaches to the user only.430        private (bool): If the dataset does not already exist, whether it431            should be created as a private dataset or public. Private datasets432            may require paid huggingface.co accounts433        verbose (bool): Whether to print out the status of the dataset434            creation.435        """436        self.hf_foken = hf_foken437        self.dataset_name = dataset_name438        self.organization_name = organization439        self.dataset_private = private440        self.verbose = verbose441 442    def setup(self, components: List[IOComponent], flagging_dir: str):443        """444        Params:445        components List[Component]: list of components for flagging446        flagging_dir (str): local directory where the dataset is cloned,447        updated, and pushed from.448        """449        try:450            import huggingface_hub451        except (ImportError, ModuleNotFoundError):452            raise ImportError(453                "Package `huggingface_hub` not found is needed "454                "for HuggingFaceDatasetJSONSaver. Try 'pip install huggingface_hub'."455            )456        path_to_dataset_repo = huggingface_hub.create_repo(457            # name=self.dataset_name,  https://github.com/huggingface/huggingface_hub/blob/main/src/huggingface_hub/hf_api.py458            repo_id=self.dataset_name,459            token=self.hf_foken,460            private=self.dataset_private,461            repo_type="dataset",462            exist_ok=True,463        )464        self.path_to_dataset_repo = path_to_dataset_repo  # e.g. "https://huggingface.co/datasets/abidlabs/test-audio-10"465        self.components = components466        self.flagging_dir = flagging_dir467        self.dataset_dir = os.path.join(flagging_dir, self.dataset_name)468        self.repo = huggingface_hub.Repository(469            local_dir=self.dataset_dir,470            clone_from=path_to_dataset_repo,471            use_auth_token=self.hf_foken,472        )473        self.repo.git_pull(lfs=True)474 475        self.infos_file = os.path.join(self.dataset_dir, "dataset_infos.json")476 477    def flag(478        self,479        flag_data: List[Any],480        flag_option: Optional[str] = None,481        flag_index: Optional[int] = None,482        username: Optional[str] = None,483    ) -> int:484        self.repo.git_pull(lfs=True)485 486        # Generate unique folder for the flagged sample487        unique_name = self.get_unique_name()  # unique name for folder488        folder_name = os.path.join(489            self.dataset_dir, unique_name490        )  # unique folder for specific example491        os.makedirs(folder_name)492 493        # Now uses the existence of `dataset_infos.json` to determine if new494        is_new = not os.path.exists(self.infos_file)495 496        # File previews for certain input and output types497        infos, file_preview_types, _ = _get_dataset_features_info(498            is_new, self.components499        )500 501        # Generate the row and header corresponding to the flagged sample502        csv_data = []503        headers = []504 505        for component, sample in zip(self.components, flag_data):506            headers.append(component.label)507 508            try:509                filepath = component.save_flagged(510                    folder_name, component.label, sample, None511                )512            except Exception:513                # Could not parse 'sample' (mostly) because it was None and `component.save_flagged`514                # does not handle None cases.515                # for example: Label (line 3109 of components.py raises an error if data is None)516                filepath = None517 518            if isinstance(component, tuple(file_preview_types)):519                headers.append(str(component.label) + " file")520 521                csv_data.append(522                    "{}/resolve/main/{}/{}".format(523                        self.path_to_dataset_repo, unique_name, filepath524                    )525                    if filepath is not None526                    else None527                )528 529            csv_data.append(filepath)530        headers.append("flag")531        csv_data.append(flag_option if flag_option is not None else "")532 533        # Creates metadata dict from row data and dumps it534        metadata_dict = {535            header: _csv_data for header, _csv_data in zip(headers, csv_data)536        }537        self.dump_json(metadata_dict, os.path.join(folder_name, "metadata.jsonl"))538 539        if is_new:540            json.dump(infos, open(self.infos_file, "w"))541 542        self.repo.push_to_hub(commit_message="Flagged sample {}".format(unique_name))543        return unique_name544 545    def get_unique_name(self):546        id = uuid.uuid4()547        return str(id)548 549    def dump_json(self, thing: dict, file_path: str) -> None:550        with open(file_path, "w+", encoding="utf8") as f:551            json.dump(thing, f)