From 8b34bf487c3e489f1778416cc9ae69e118a099d9 Mon Sep 17 00:00:00 2001 From: Rasic2 <1051987201@qq.com> Date: Wed, 4 Feb 2026 22:25:29 +0800 Subject: [PATCH] =?UTF-8?q?feat=EF=BC=9A=E6=96=B0=E5=A2=9Ellm=5Fcompletion?= =?UTF-8?q?.py=E8=84=9A=E6=9C=AC=E4=BB=A5=E6=94=AF=E6=8C=81Azure=E5=92=8CG?= =?UTF-8?q?PUGeek=20API=E4=B8=B2=E6=B5=81=E7=BB=9F=E8=AE=A1=E5=8A=9F?= =?UTF-8?q?=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/llm_completion.py | 92 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 92 insertions(+) create mode 100644 scripts/llm_completion.py diff --git a/scripts/llm_completion.py b/scripts/llm_completion.py new file mode 100644 index 00000000..85bcf49f --- /dev/null +++ b/scripts/llm_completion.py @@ -0,0 +1,92 @@ +import os +import time + +from dotenv import find_dotenv, load_dotenv +from openai import AzureOpenAI, OpenAI + +load_dotenv(find_dotenv()) + +USER_MESSAGE = 'hi' + + +def run_stream_stats(client, model, label, *, max_tokens=16384, temperature=None): + """流式请求并统计 TTFT、总耗时、token 数。""" + kwargs = { + 'messages': [{'role': 'user', 'content': USER_MESSAGE}], + 'max_completion_tokens': max_tokens, + 'model': model, + 'stream': True, + 'stream_options': {'include_usage': True}, + } + if temperature is not None: + kwargs['temperature'] = temperature + + t0 = time.perf_counter() + ttft = None + content_parts = [] + usage = None + + stream = client.chat.completions.create(**kwargs) + + for chunk in stream: + if ttft is None and chunk.choices and chunk.choices[0].delta.content: + ttft = time.perf_counter() - t0 + if chunk.choices and chunk.choices[0].delta.content: + part = chunk.choices[0].delta.content + content_parts.append(part) + print(part, end='', flush=True) + if getattr(chunk, 'usage', None) is not None: + usage = chunk.usage + + elapsed = time.perf_counter() - t0 + print() + print(f"\n[{label}]") + print( + f" 首 token 耗时 (TTFT): {ttft:.3f}s" + if ttft is not None + else ' 首 token 耗时: N/A' + ) + print(f" 总耗时: {elapsed:.3f}s") + if usage and getattr(usage, 'total_tokens', None): + total_tokens = usage.total_tokens + if total_tokens > 0: + print( + f" 总 token 数: {total_tokens}, 约 {total_tokens / elapsed:.1f} tokens/s" + ) + return ttft, elapsed, usage + + +# ---------- 1. Azure(.env:AZURE_API_KEY 等)---------- +endpoint = os.environ.get('AZURE_API_BASE', 'https://physical4.openai.azure.com/') +api_key = os.environ.get('AZURE_API_KEY') +if not api_key: + raise ValueError('请在 .env 中设置 AZURE_API_KEY') +deployment = os.environ.get('GPT_CHAT_MODEL', 'gpt-5-chat') +api_version = os.environ.get('AZURE_API_VERSION', '2025-01-01-preview') + +client_azure = AzureOpenAI( + api_version=api_version, + azure_endpoint=endpoint, + api_key=api_key, +) + +print('=== Azure (GPT_CHAT_MODEL) ===') +run_stream_stats(client_azure, deployment, 'Azure', max_tokens=16384) + +# ---------- 2. GPUGeek(OpenAI/Azure-GPT-5.1,.env:GPUGEEK_API_KEY)---------- +gpugeek_key = os.environ.get('GPUGEEK_API_KEY') +gpugeek_base = os.environ.get('GPUGEEK_BASE_URL', 'https://api.gpugeek.com/v1') +gpugeek_model = os.environ.get('GPUGEEK_MODEL', 'OpenAI/Azure-GPT-5.1') + +if gpugeek_key: + client_gpugeek = OpenAI(api_key=gpugeek_key, base_url=gpugeek_base) + print('\n=== GPUGeek (OpenAI/Azure-GPT-5.1) ===') + run_stream_stats( + client_gpugeek, + gpugeek_model, + 'GPUGeek', + max_tokens=2048, + temperature=0.7, + ) +else: + print('\n(跳过 GPUGeek:未设置 GPUGEEK_API_KEY)')