haolun/llm-profiler
0
1import gradio as gr2import io3import logging4 5from llm_profiler import *6import sys7from contextlib import redirect_stdout8 9# 模型列表10model_names = [11 "opt-1.3b",12 "opt-6.7b",13 "opt-13b",14 "opt-66b",15 "opt-175b",16 "gpt2",17 "gpt2-medium",18 "gpt2-large",19 "gpt2-xl",20 "bloom-560m",21 "bloom-7b",22 "bloom-175b",23 "llama-7b",24 "llama-13b",25 "llama-30b",26 "llama-65b",27 "llama2-13b",28 "llama2-70b",29 "internlm-20b",30 "baichuan2-13b",31]32# GPU 列表33gpu_names = [34 "t4-pcie-15gb",35 "v100-pcie-32gb",36 "v100-sxm-32gb",37 "br104p",38 "a100-pcie-40gb",39 "a100-sxm-40gb",40 "a100-pcie-80gb",41 "a100-sxm-80gb",42 "910b-64gb",43 "h100-sxm-80gb",44 "h100-pcie-80gb",45 "a30-pcie-24gb",46 "a30-sxm-24gb",47 "a40-pcie-48gb",48]49 50 51# 创建一个日志处理器,将日志消息写入 StringIO 对象52class StringHandler(logging.Handler):53 def __init__(self):54 super().__init__()55 self.stream = io.StringIO()56 self.setFormatter(logging.Formatter("%(message)s"))57 58 def emit(self, record):59 self.stream.write(self.format(record) + "\n")60 61 def get_value(self):62 return self.stream.getvalue()63 64 65# 创建一个日志记录器并添加 StringHandler66logger = logging.getLogger(__name__)67logger.setLevel(logging.INFO)68string_handler = StringHandler()69logger.addHandler(string_handler)70 71 72def gradio_interface(73 model_name="llama2-70b",74 gpu_name: str = "t4-pcie-15gb",75 bytes_per_param: int = BYTES_FP16,76 batch_size_per_gpu: int = 2,77 seq_len: int = 300,78 generate_len: int = 40,79 ds_zero: int = 0,80 dp_size: int = 1,81 tp_size: int = 4,82 pp_size: int = 1,83 sp_size: int = 1,84 use_kv_cache: bool = True,85 layernorm_dtype_bytes: int = BYTES_FP16,86 kv_cache_dtype_bytes: int = BYTES_FP16,87 flops_efficiency: float = FLOPS_EFFICIENCY,88 hbm_memory_efficiency: float = HBM_MEMORY_EFFICIENCY,89 intra_node_memory_efficiency: float = INTRA_NODE_MEMORY_EFFICIENCY,90 inter_node_memory_efficiency: float = INTER_NODE_MEMORY_EFFICIENCY,91 mode: str = "inference",92 print_flag: bool = True,93) -> list:94 # 清空 StringIO 对象95 string_handler.stream.seek(0)96 string_handler.stream.truncate()97 98 # 重定向 sys.stdout 到 StringHandler99 original_stdout = sys.stdout100 sys.stdout = string_handler.stream101 102 # 调用你的推理函数103 results = llm_profile_infer(104 model_name,105 gpu_name,106 bytes_per_param,107 batch_size_per_gpu,108 seq_len,109 generate_len,110 ds_zero,111 dp_size,112 tp_size,113 pp_size,114 sp_size,115 use_kv_cache,116 layernorm_dtype_bytes,117 kv_cache_dtype_bytes,118 flops_efficiency,119 hbm_memory_efficiency,120 intra_node_memory_efficiency,121 inter_node_memory_efficiency,122 mode,123 print_flag,124 )125 126 # 恢复 sys.stdout127 sys.stdout = original_stdout128 129 # 获取日志消息130 log_output = string_handler.get_value()131 132 # 返回推理结果和日志输出133 return results, log_output134 135 136# 创建 Gradio 界面137iface = gr.Interface(138 fn=gradio_interface,139 inputs=[140 gr.Dropdown(choices=model_names, label="Model Name", value="llama2-70b"),141 gr.Dropdown(choices=gpu_names, label="GPU Name", value="a100-sxm-80gb"),142 gr.Number(label="Bytes per Param", value=BYTES_FP16),143 gr.Number(label="Batch Size per GPU", value=2),144 gr.Number(label="Sequence Length", value=300),145 gr.Number(label="Generate Length", value=40),146 gr.Number(label="DS Zero", value=0),147 gr.Number(label="DP Size", value=1),148 gr.Number(label="TP Size", value=4),149 gr.Number(label="PP Size", value=1),150 gr.Number(label="SP Size", value=1),151 gr.Checkbox(label="Use KV Cache", value=True),152 gr.Number(label="Layernorm dtype Bytes", value=BYTES_FP16),153 gr.Number(label="KV Cache dtype Bytes", value=BYTES_FP16),154 gr.Number(label="FLOPS Efficiency", value=FLOPS_EFFICIENCY),155 gr.Number(label="HBM Memory Efficiency", value=HBM_MEMORY_EFFICIENCY),156 gr.Number(157 label="Intra Node Memory Efficiency", value=INTRA_NODE_MEMORY_EFFICIENCY158 ),159 gr.Number(160 label="Inter Node Memory Efficiency", value=INTER_NODE_MEMORY_EFFICIENCY161 ),162 gr.Radio(choices=["inference", "other_mode"], label="Mode", value="inference"),163 gr.Checkbox(label="Print Flag", value=True),164 ],165 outputs=[166 gr.Textbox(label="Inference Results"), # 推理结果输出,带标签167 gr.Textbox(label="Detailed Analysis"), # 日志输出,带标签168 ],169 title="LLM Profiler",170 description="Input parameters to profile your LLM.",171)172 173# 启动 Gradio 界面174iface.launch(auth=("xtrt-llm", "xtrt-llm"), share=False)175# iface.launch()176 