bnitokyo/InternVL
0
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 