XaviXva/Video-LLaVA
0
1import datetime2import logging3import logging.handlers4import os5import sys6 7import requests8 9from llava.constants import LOGDIR10 11server_error_msg = "**NETWORK ERROR DUE TO HIGH TRAFFIC. PLEASE REGENERATE OR REFRESH THIS PAGE.**"12moderation_msg = "YOUR INPUT VIOLATES OUR CONTENT MODERATION GUIDELINES. PLEASE TRY AGAIN."13 14handler = None15 16 17def build_logger(logger_name, logger_filename):18 global handler19 20 formatter = logging.Formatter(21 fmt="%(asctime)s | %(levelname)s | %(name)s | %(message)s",22 datefmt="%Y-%m-%d %H:%M:%S",23 )24 25 # Set the format of root handlers26 if not logging.getLogger().handlers:27 logging.basicConfig(level=logging.INFO)28 logging.getLogger().handlers[0].setFormatter(formatter)29 30 # Redirect stdout and stderr to loggers31 stdout_logger = logging.getLogger("stdout")32 stdout_logger.setLevel(logging.INFO)33 sl = StreamToLogger(stdout_logger, logging.INFO)34 sys.stdout = sl35 36 stderr_logger = logging.getLogger("stderr")37 stderr_logger.setLevel(logging.ERROR)38 sl = StreamToLogger(stderr_logger, logging.ERROR)39 sys.stderr = sl40 41 # Get logger42 logger = logging.getLogger(logger_name)43 logger.setLevel(logging.INFO)44 45 # Add a file handler for all loggers46 if handler is None:47 os.makedirs(LOGDIR, exist_ok=True)48 filename = os.path.join(LOGDIR, logger_filename)49 handler = logging.handlers.TimedRotatingFileHandler(50 filename, when='D', utc=True)51 handler.setFormatter(formatter)52 53 for name, item in logging.root.manager.loggerDict.items():54 if isinstance(item, logging.Logger):55 item.addHandler(handler)56 57 return logger58 59 60class StreamToLogger(object):61 """62 Fake file-like stream object that redirects writes to a logger instance.63 """64 def __init__(self, logger, log_level=logging.INFO):65 self.terminal = sys.stdout66 self.logger = logger67 self.log_level = log_level68 self.linebuf = ''69 70 def __getattr__(self, attr):71 return getattr(self.terminal, attr)72 73 def write(self, buf):74 temp_linebuf = self.linebuf + buf75 self.linebuf = ''76 for line in temp_linebuf.splitlines(True):77 # From the io.TextIOWrapper docs:78 # On output, if newline is None, any '\n' characters written79 # are translated to the system default line separator.80 # By default sys.stdout.write() expects '\n' newlines and then81 # translates them so this is still cross platform.82 if line[-1] == '\n':83 self.logger.log(self.log_level, line.rstrip())84 else:85 self.linebuf += line86 87 def flush(self):88 if self.linebuf != '':89 self.logger.log(self.log_level, self.linebuf.rstrip())90 self.linebuf = ''91 92 93def disable_torch_init():94 """95 Disable the redundant torch default initialization to accelerate model creation.96 """97 import torch98 setattr(torch.nn.Linear, "reset_parameters", lambda self: None)99 setattr(torch.nn.LayerNorm, "reset_parameters", lambda self: None)100 101 102def violates_moderation(text):103 """104 Check whether the text violates OpenAI moderation API.105 """106 url = "https://api.openai.com/v1/moderations"107 headers = {"Content-Type": "application/json",108 "Authorization": "Bearer " + os.environ["OPENAI_API_KEY"]}109 text = text.replace("\n", "")110 data = "{" + '"input": ' + f'"{text}"' + "}"111 data = data.encode("utf-8")112 try:113 ret = requests.post(url, headers=headers, data=data, timeout=5)114 flagged = ret.json()["results"][0]["flagged"]115 except requests.exceptions.RequestException as e:116 flagged = False117 except KeyError as e:118 flagged = False119 120 return flagged121 122 123def pretty_print_semaphore(semaphore):124 if semaphore is None:125 return "None"126 return f"Semaphore(value={semaphore._value}, locked={semaphore.locked()})"127 