MultiModal / infer.py
szxllm's picture
Update infer.py
2e9a238 verified
Raw
History Blame Contribute Delete
12.6 kB
import os
import argparse
from pathlib import Path
import json
from typing import Optional
import torch
from PIL import Image
from transformers import AutoTokenizer
import gradio as gr
from model import MultiModalDenseTransformer
from continual_learning import UnifiedMultiModalPreprocessor
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
from torchvision import transforms
image_transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
class ModelInference:
def __init__(
self,
checkpoint_path: str,
tokenizer_name: str,
config_path: Optional[str] = None,
device: str = 'cuda' if torch.cuda.is_available() else 'cpu'
):
self.device = torch.device(device)
self.tokenizer = AutoTokenizer.from_pretrained(
tokenizer_name,
use_fast=True,
trust_remote_code=True
)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
self.tokenizer.pad_token_id = self.tokenizer.eos_token_id
if config_path and Path(config_path).exists():
with open(config_path, 'r') as f:
self.config = json.load(f)
else:
self.config = {
'model_dim': 1536,
'vocab_size': len(self.tokenizer),
'n_layers': 12,
'n_heads': 12,
'n_kv_heads': 4,
'head_dim': None,
'max_seq_len': 512,
'dropout': 0.0,
'use_moe': False,
'use_adapter': False,
'use_lora': False,
'rope_scaling_type': "yarn",
'use_multimodal_fusion': False,
'use_contrastive': False
}
self.model = MultiModalDenseTransformer(**self.config)
self.preprocessor = UnifiedMultiModalPreprocessor(model_dim=self.config['model_dim'])
print(f"Loading checkpoint from {checkpoint_path}...")
checkpoint = torch.load(checkpoint_path, map_location=self.device)
state_dict = None
if 'actor_state_dict' in checkpoint:
print("Detected GRPO checkpoint format (actor_state_dict)")
state_dict = checkpoint['actor_state_dict']
elif 'model_state_dict' in checkpoint:
print("Detected Standard/SFT checkpoint format (model_state_dict)")
state_dict = checkpoint['model_state_dict']
else:
print("Detected raw state dict format")
state_dict = checkpoint
new_state_dict = {}
for k, v in state_dict.items():
if k.startswith('module.'):
new_state_dict[k[7:]] = v
else:
new_state_dict[k] = v
missing, unexpected = self.model.load_state_dict(new_state_dict, strict=False)
if missing:
print(f"Warning: Missing keys: {len(missing)}")
if len(missing) <= 10:
print(f"Missing keys: {missing}")
if unexpected:
print(f"Warning: Unexpected keys: {len(unexpected)}")
if len(unexpected) <= 10:
print(f"Unexpected keys: {unexpected}")
self.model.to(self.device)
self.preprocessor.to(self.device)
self.model.eval()
def _build_position_ids(self, attention_mask: torch.Tensor) -> torch.Tensor:
batch_size, seq_len = attention_mask.shape
position_ids = torch.zeros((batch_size, seq_len), dtype=torch.long, device=self.device)
for i in range(batch_size):
non_pad_positions = (attention_mask[i] == 1).nonzero(as_tuple=True)[0]
if len(non_pad_positions) > 0:
start_pos = non_pad_positions[0].item()
valid_len = len(non_pad_positions)
# 从 0 开始编号有效 token 的位置
position_ids[i, start_pos:start_pos + valid_len] = torch.arange(
valid_len,
device=self.device
)
return position_ids
@torch.no_grad()
def generate_text(
self,
prompt: str,
max_new_tokens: int = 128,
temperature: float = 0.7,
top_k: int = 40,
top_p: float = 0.9,
repetition_penalty: float = 1.1,
image: Optional[Image.Image] = None
) -> str:
formatted_prompt = f"user: {prompt}\nassistant:\n<think>\n"
inputs = self.tokenizer(
formatted_prompt,
return_tensors="pt",
padding=False
)
input_ids = inputs['input_ids'].to(self.device)
attention_mask = inputs['attention_mask'].to(self.device)
segments = []
segments.append({
'type': 'text',
'data': input_ids,
'modality_id': 0
})
has_image = False
if image is not None:
try:
if image.mode != 'RGB':
image = image.convert('RGB')
image_tensor = image_transform(image).unsqueeze(0).to(self.device)
segments.append({
'type': 'image',
'data': image_tensor,
'modality_id': 1
})
has_image = True
print("Image added to input")
except Exception as e:
print(f"Warning: Image processing error: {e}")
position_ids = self._build_position_ids(attention_mask)
input_data = {
'segments': segments,
}
input_data['attention_mask'] = attention_mask
if not has_image:
input_data['position_ids'] = position_ids
try:
generated_ids = self.model.generate(
input_data,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_k=top_k,
top_p=top_p,
repetition_penalty=repetition_penalty,
do_sample=True,
eos_token_id=self.tokenizer.eos_token_id,
pad_token_id=self.tokenizer.pad_token_id
)
output_text = self.tokenizer.decode(generated_ids[0], skip_special_tokens=True)
return output_text.strip()
except Exception as e:
import traceback
traceback.print_exc()
return f"Error during generation: {str(e)}"
def build_ui(model_instance):
with gr.Blocks(title="MultiModal Dense Transformer - Gradio", css="""
.gradio-container { max-width: 900px; margin: auto; }
""") as demo:
gr.Markdown("## 在线推理(文本)")
with gr.Row():
with gr.Column(scale=3):
txt = gr.Textbox(
label="Prompt (Instruction)",
placeholder="请输入指令或问题...",
lines=5
)
img = gr.Image(type="pil", label="(可选) 上传图片(支持多模态)")
btn = gr.Button("生成 (Generate)", variant="primary")
with gr.Column(scale=2):
max_tokens = gr.Slider(
label="Max New Tokens",
minimum=16,
maximum=1024,
step=1,
value=128
)
temperature = gr.Slider(
label="Temperature",
minimum=0.1,
maximum=1.5,
step=0.01,
value=0.7
)
top_k = gr.Slider(
label="Top-k",
minimum=0,
maximum=200,
step=1,
value=40
)
top_p = gr.Slider(
label="Top-p",
minimum=0.0,
maximum=1.0,
step=0.01,
value=0.9
)
rep_pen = gr.Slider(
label="Repetition Penalty",
minimum=0.5,
maximum=2.0,
step=0.01,
value=1.1
)
status = gr.Textbox(
label="Status",
value="Ready",
interactive=False
)
output = gr.Textbox(label="Output", lines=12, interactive=False)
gr.Examples(
examples=[
["请解释什么是深度学习", None],
["计算 123 + 456 等于多少?", None],
["写一首关于春天的诗", None],
],
inputs=[txt, img],
)
def gr_generate(prompt, image, max_tokens_v, temp_v, topk_v, topp_v, rep_v):
if not prompt or str(prompt).strip() == "":
return "", " 请输入 Prompt"
try:
status_msg = " Generating..."
# 调用模型生成
out = model_instance.generate_text(
prompt=prompt,
max_new_tokens=int(max_tokens_v),
temperature=float(temp_v),
top_k=int(topk_v),
top_p=float(topp_v),
repetition_penalty=float(rep_v),
image=image
)
return out, " Done"
except Exception as e:
return f"Error: {str(e)}", " Error"
btn.click(
fn=gr_generate,
inputs=[txt, img, max_tokens, temperature, top_k, top_p, rep_pen],
outputs=[output, status]
)
return demo
def main():
parser = argparse.ArgumentParser(
description="Gradio inference interface for MultiModal Dense Transformer"
)
parser.add_argument(
"--checkpoint",
type=str,
default="/root/checkpoints/dcpo_posttrain_round3/step_15600.pt",
help="Path to model checkpoint"
)
parser.add_argument(
"--tokenizer",
type=str,
default="Qwen/Qwen2.5-7B-Instruct",
help="Tokenizer name or path"
)
parser.add_argument(
"--config",
type=str,
default=None,
help="Path to model config JSON (optional)"
)
parser.add_argument(
"--port",
type=int,
default=5001,
help="Port to run Gradio server"
)
parser.add_argument(
"--share",
type=lambda x: x.lower() in ("true","1","yes"),
default=True,
help="Create public link (True/False)"
)
args = parser.parse_args()
if not Path(args.checkpoint).exists():
print(f" Checkpoint not found: {args.checkpoint}")
possible_dirs = [
Path("/root/checkpoints/posttrain/grpo"),
Path("/root/checkpoints/dcpo_training"),
Path("/root/checkpoints/r1_zero_reproduction"),
]
for checkpoint_dir in possible_dirs:
if checkpoint_dir.exists():
grpo_files = sorted(
[p for p in checkpoint_dir.glob("grpo_iter_*.pt")],
key=lambda p: int(p.stem.split('_')[-1]) if p.stem.split('_')[-1].isdigit() else 0
)
step_files = sorted(
[p for p in checkpoint_dir.glob("step_*.pt")],
key=lambda p: int(p.stem.split('_')[-1]) if p.stem.split('_')[-1].isdigit() else 0
)
candidates = grpo_files + step_files
if candidates:
args.checkpoint = str(candidates[-1])
print(f" Using latest checkpoint: {args.checkpoint}")
break
if not Path(args.checkpoint).exists():
raise FileNotFoundError(f"找不到可用的检查点文件")
global model_instance
model_instance = ModelInference(
args.checkpoint,
args.tokenizer,
args.config
)
demo = build_ui(model_instance)
demo.launch(
server_port=args.port,
share=args.share,
server_name="0.0.0.0" # 允许外部访问
)
if __name__ == "__main__":
main()