quantumiracle-git/OpenBiDexHand
5
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)