#!/usr/bin/env python3
"""
vLLM 最小化推理示例 - 使用 Hugging Face transformers 库进行大模型推理
这个脚本演示了如何使用预训练的大语言模型进行文本生成。
"""

from transformers import AutoTokenizer, AutoModelForCausalLM
import torch

# ============================================================================
# 第一步：加载预训练的大模型和分词器
# ============================================================================
print("Loading model...")
# 模型标识符，指向 Hugging Face 模型库中的 Qwen 2.5 1.5B 指令微调版本
# Qwen/Qwen2.5-1.5B-Instruct 是一个 15 亿参数的轻量级指令跟随模型
model_name = "Qwen/Qwen2.5-1.5B-Instruct"

# AutoTokenizer: 自动加载与模型配套的分词器
# 分词器的作用是将文本转换为模型能理解的 token（令牌/词元）
# 具有相反的解码功能：将 token 转回文本
tokenizer = AutoTokenizer.from_pretrained(model_name)

# AutoModelForCausalLM: 加载因果语言模型（Causal Language Model）
# 因果语言模型：只能看到当前位置之前的 token，用于文本生成任务
#
# 数据类型选择与优化：
# - torch.float16 (FP16): GPU 推理的标准选择，节省 50% 显存，加速计算
#   但在某些 CPU 场景下可能数值不稳定
# - torch.bfloat16 (BF16): 更好的数值稳定性，特别是在 CPU 上，推荐优先使用
#   BF16 保留与 FP32 相同的指数范围，只牺牲尾数精度，更适合深度学习
# - torch.float32: 精度最高但显存占用最多（2 倍 FP16），一般不用于推理
# 设备选择策略：
# - 优先选择 CUDA（NVIDIA GPU）- 性能最好，生态最完善
# - 其次 CPU - 稳定可靠，但速度较慢
# - 避免 MPS（Apple GPU）- 某些操作支持不完整，容易出现 OOM 错误
device = torch.device("cpu")  # Mac 用户推荐使用 CPU，避免 MPS 兼容性问题
if torch.cuda.is_available():
    device = torch.device("cuda")

# 选择 bfloat16 作为默认数据类型，兼容 GPU 和 CPU，更稳定
dtype = torch.bfloat16

# 直接指定设备而不用 device_map="auto"，避免设备选择的不确定性
model = AutoModelForCausalLM.from_pretrained(model_name, dtype=dtype, device_map=device)

# ============================================================================
# 性能和稳定性优化
# ============================================================================
# 启用梯度检查点（Gradient Checkpointing）：
# 虽然推理时不需要梯度，但启用此优化对以下场景有益：
# 1. 为后续微调预留显存空间
# 2. 在显存受限的环境中运行较大的模型
# 原理：不缓存所有前向传播的中间激活值，而是需要时重新计算
# 权衡：用计算时间换取显存空间（通常减少 50% 显存占用，增加约 20% 推理时间）
model.gradient_checkpointing_enable()

# 启用 eval 模式：禁用 dropout、batch normalization 等训练特定的行为
# 优势：
# 1. 输出更稳定，同样输入每次得到完全相同的输出（复现性强）
# 2. 略微提升推理速度
# 3. 避免随机失活层的影响
model.eval()

# 设置最大序列长度限制（用于防止显存溢出 OOM）
# Qwen2.5-1.5B 的完整上下文窗口可达 32k tokens，但这会占用大量显存
# 这里限制到 2048 是实际应用中的合理折衷
# 如果需要更长的上下文，可以使用 sliding window attention 等优化技术
max_sequence_length = 2048

print(f"Model loaded on {device} with dtype {dtype}!")

# ============================================================================
# 第二步：准备输入数据 - 构建对话消息
# ============================================================================
# 定义聊天消息，遵循标准的对话格式
# role: 消息的角色（user=用户提问，assistant=模型回复）
messages = [
    {"role": "user", "content": "Hello, how are you?"},
]

# apply_chat_template: 将对话消息转换为模型期望的格式
# 模型被微调为处理特定的提示模板，此方法自动应用这个模板
text = tokenizer.apply_chat_template(
    messages,
    # tokenize=False: 返回格式化的字符串而不是 token（我们稍后手动 tokenize）
    tokenize=False,
    # add_generation_prompt=True: 在末尾添加特殊 token，指示模型开始生成
    # 这确保了模型会继续文本而不是停止或反复输入
    add_generation_prompt=True
)

# tokenizer(...): 现在将格式化的文本转换为 token ID 张量
# return_tensors="pt": 返回 PyTorch 张量而不是列表
# model_inputs 包含 input_ids（token ID 序列）和 attention_mask（注意力掩码）
model_inputs = tokenizer([text], return_tensors="pt")

# 可选：如果输入序列过长，可以截断以节省显存
# 但会丢失上文信息，需要权衡
if model_inputs.input_ids.shape[1] > max_sequence_length:
    model_inputs.input_ids = model_inputs.input_ids[:, -max_sequence_length:]
    if "attention_mask" in model_inputs:
        model_inputs.attention_mask = model_inputs.attention_mask[:, -max_sequence_length:]

# 重要：将输入张量移动到模型所在的设备
# 如果设备不匹配（如模型在 GPU，输入在 CPU），会导致运行时错误
# model_inputs 是一个字典，包含 input_ids、attention_mask 等张量
model_inputs = {k: v.to(device) for k, v in model_inputs.items()}

# ============================================================================
# 第三步：执行模型推理 - 生成文本响应
# ============================================================================
# torch.no_grad(): 禁用梯度计算以节省显存和加速推理
# 梯度只在训练时需要，推理过程不需要
with torch.no_grad():
    # model.generate: 自回归文本生成方法
    # 逐个 token 地生成文本，每次预测下一个最可能的 token
    generated_ids = model.generate(
        # input_ids: 输入的 token ID 张量
        model_inputs["input_ids"],
        # attention_mask: 明确传入注意力掩码以避免 pad==eos 时无法推断掩码的问题
        attention_mask=model_inputs["attention_mask"],
        # max_new_tokens: 最多生成多少个新 token（不包括输入的 token）
        # 限制输出长度，防止生成过长的文本和显存溢出
        # 对于实时应用，建议控制在 256-512 之间以保证响应速度
        max_new_tokens=256,
        
        # temperature: 控制生成的随机性和多样性（范围 0.0-2.0）
        # 0.0: 完全确定性（总是选概率最高的 token），可能导致重复和单调
        # 0.7: 平衡点（推荐），既保留多样性又保持相对连贯
        # 1.0: 中等随机性，适合创意写作
        # > 1.0: 高度随机，可能导致不连贯或质量下降
        # 对话场景推荐 0.6-0.8，创意写作推荐 0.8-1.2
        temperature=0.7,
        
        # top_p (核采样，Nucleus Sampling): 动态选择生成的 token
        # 只从累计概率达到 p 的最可能的 token 中采样
        # 0.95: 保留概率最高的 token，直到累计概率 ≥ 95%
        # 优势：比 top_k 更自适应，避免低概率 token 造成的胡言乱语
        # 比 temperature 更高效，能更好地控制生成质量
        # 推荐范围：0.8-0.95（平衡质量和多样性）
        top_p=0.95,
        
        # do_sample: 采样策略选择
        # True: 根据概率分布随机采样下一个 token（推荐用于对话和创意任务）
        #   - 多次运行同一输入会得到不同的输出
        #   - 结合 temperature 和 top_p 使用效果最好
        # False: 贪心搜索，总是选择概率最高的 token（推荐用于翻译等确定性任务）
        #   - 输出完全确定，适合需要复现性的应用
        #   - 可能导致生成重复或单调的文本
        do_sample=True
    )

# ============================================================================
# 第四步：解码和输出结果
# ============================================================================
# tokenizer.decode: 将 token ID 序列转换回可读的文本
# skip_special_tokens=True: 不显示特殊 token（如 <pad>, <eos> 等）
response = tokenizer.decode(generated_ids[0], skip_special_tokens=True)

print("\nResponse:")
print(response)