CoolFace
Apppublic

OpenGVLab/InternVL

sourceHugging Facemitupdated 2y agoView on Hugging Face
510likes
utils.py164 linesDownload Raw Back to root
1from ast import Dict2import logging3import logging.handlers4import os5import sys6import base647from PIL import Image8from io import BytesIO9import json10import requests11from constants import LOGDIR12import datetime13 14server_error_msg = (15    "**NETWORK ERROR DUE TO HIGH TRAFFIC. PLEASE REGENERATE OR REFRESH THIS PAGE.**"16)17moderation_msg = (18    "YOUR INPUT VIOLATES OUR CONTENT MODERATION GUIDELINES. PLEASE TRY AGAIN."19)20 21handler = None22 23 24def build_logger(logger_name, logger_filename):25    global handler26 27    formatter = logging.Formatter(28        fmt="%(asctime)s | %(levelname)s | %(name)s | %(message)s",29        datefmt="%Y-%m-%d %H:%M:%S",30    )31 32    # Set the format of root handlers33    if not logging.getLogger().handlers:34        logging.basicConfig(level=logging.INFO)35    logging.getLogger().handlers[0].setFormatter(formatter)36 37    # Redirect stdout and stderr to loggers38    stdout_logger = logging.getLogger("stdout")39    stdout_logger.setLevel(logging.INFO)40    sl = StreamToLogger(stdout_logger, logging.INFO)41    sys.stdout = sl42 43    stderr_logger = logging.getLogger("stderr")44    stderr_logger.setLevel(logging.ERROR)45    sl = StreamToLogger(stderr_logger, logging.ERROR)46    sys.stderr = sl47 48    # Get logger49    logger = logging.getLogger(logger_name)50    logger.setLevel(logging.INFO)51 52    # Add a file handler for all loggers53    if handler is None:54        os.makedirs(LOGDIR, exist_ok=True)55        filename = os.path.join(LOGDIR, logger_filename)56        handler = logging.handlers.TimedRotatingFileHandler(57            filename, when="D", utc=True58        )59        handler.setFormatter(formatter)60 61        for name, item in logging.root.manager.loggerDict.items():62            if isinstance(item, logging.Logger):63                item.addHandler(handler)64 65    return logger66 67 68class StreamToLogger(object):69    """70    Fake file-like stream object that redirects writes to a logger instance.71    """72 73    def __init__(self, logger, log_level=logging.INFO):74        self.terminal = sys.stdout75        self.logger = logger76        self.log_level = log_level77        self.linebuf = ""78 79    def __getattr__(self, attr):80        return getattr(self.terminal, attr)81 82    def write(self, buf):83        temp_linebuf = self.linebuf + buf84        self.linebuf = ""85        for line in temp_linebuf.splitlines(True):86            # From the io.TextIOWrapper docs:87            #   On output, if newline is None, any '\n' characters written88            #   are translated to the system default line separator.89            # By default sys.stdout.write() expects '\n' newlines and then90            # translates them so this is still cross platform.91            if line[-1] == "\n":92                self.logger.log(self.log_level, line.rstrip())93            else:94                self.linebuf += line95 96    def flush(self):97        if self.linebuf != "":98            self.logger.log(self.log_level, self.linebuf.rstrip())99        self.linebuf = ""100 101 102def disable_torch_init():103    """104    Disable the redundant torch default initialization to accelerate model creation.105    """106    import torch107 108    setattr(torch.nn.Linear, "reset_parameters", lambda self: None)109    setattr(torch.nn.LayerNorm, "reset_parameters", lambda self: None)110 111 112def violates_moderation(text):113    """114    Check whether the text violates OpenAI moderation API.115    """116    url = "https://api.openai.com/v1/moderations"117    headers = {118        "Content-Type": "application/json",119        "Authorization": "Bearer " + os.environ["OPENAI_API_KEY"],120    }121    text = text.replace("\n", "")122    data = "{" + '"input": ' + f'"{text}"' + "}"123    data = data.encode("utf-8")124    try:125        ret = requests.post(url, headers=headers, data=data, timeout=5)126        flagged = ret.json()["results"][0]["flagged"]127    except requests.exceptions.RequestException as e:128        flagged = False129    except KeyError as e:130        flagged = False131 132    return flagged133 134 135def pretty_print_semaphore(semaphore):136    if semaphore is None:137        return "None"138    return f"Semaphore(value={semaphore._value}, locked={semaphore.locked()})"139 140 141def load_image_from_base64(image):142    return Image.open(BytesIO(base64.b64decode(image)))143 144 145def get_log_filename():146    t = datetime.datetime.now()147    name = os.path.join(LOGDIR, f"{t.year}-{t.month:02d}-{t.day:02d}-conv.json")148    return name149 150 151def data_wrapper(data):152    if isinstance(data, bytes):153        return data154    elif isinstance(data, Image.Image):155        buffered = BytesIO()156        data.save(buffered, format="PNG")157        return buffered.getvalue()158    elif isinstance(data, str):159        return data.encode()160    elif isinstance(data, Dict):161        return json.dumps(data).encode()162    else:163        raise ValueError(f"Unsupported data type: {type(data)}")164