USER
给我一个推理和动作控制的单线程代码
import os
import torch
import torch.nn as nn
from torchvision import transforms
from torchvision.models import resnet50, ResNet50_Weights
from PIL import Image
import queue
import threading
import mss
from pynput.keyboard import Key, Controller as KeyboardController
import pydirectinput
# ---- 模型定义 ----
class MultiTaskResNet(nn.Module):
def __init__(self, base_model):
super(MultiTaskResNet, self).__init__()
self.shared_features = nn.Sequential(*list(base_model.children())[:-1]) # 特征提取器
self.action_head = nn.Sequential(
nn.Linear(base_model.fc.in_features, 128),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(128, 8), # 动作头(预测8类动作)
)
self.offset_head = nn.Sequential(
nn.Linear(base_model.fc.in_features, 128),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(128, 2) # 偏移头(预测 dx 和 dy)
)
def forward(self, x):
features = self.shared_features(x).flatten(1) # 特征提取后展平
actions = self.action_head(features) # 动作预测
offsets = self.offset_head(features) # 偏移预测
return actions, offsets
# ---- 加载模型 ----
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 加载 ResNet50 基础模型
base_model = resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)
model = MultiTaskResNet(base_model).to(device).eval()
# 加载训练好的权重(确保路径正确)
model.load_state_dict(torch.load("model_final.pth", map_location=device)) # 替换为你的模型路径
model = model.float()
# ---- 图像预处理 ----
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]) # 标准化
])
# ---- 分段反归一化工具函数 ----
def denormalize(value, core_min, core_max, global_min, global_max):
if value < 0.2:
return value * (core_min - global_min) / 0.2 + global_min
elif value > 0.8:
return (value - 0.8) * (global_max - core_max) / 0.2 + core_max
else:
return (value - 0.2) * (core_max - core_min) / 0.6 + core_min
# 鼠标 dx/dy 偏移值范围
core_dx_min, core_dx_max = -8, 8
core_dy_min, core_dy_max = -3, 3
dx_min, dx_max = -312, 230
dy_min, dy_max = -64, 140
# ---- 推理函数 ----
def predict(image):
"""
推理函数,用于对输入图像进行动作和鼠标偏移预测。
输入:
- image: 一个 PIL 图像
输出:
- 动作数组和鼠标坐标偏移 [W, S, A, D, R, Space, Shift, MouseLeft, dx, dy]
"""
# 图像预处理
image_tensor = transform(image).unsqueeze(0).to(device)
# 禁用梯度计算,提高推理效率
with torch.no_grad():
# 模型推理
actions_logits, offsets_raw = model(image_tensor)
# 对动作信号应用 Sigmoid 激活,二值化输出
actions = (torch.sigmoid(actions_logits) > 0.5).squeeze(0).cpu().numpy().astype(int)
# 对鼠标 dx 和 dy 应用反归一化
dx = denormalize(offsets_raw[0, 0].item(), core_dx_min, core_dx_max, dx_min, dx_max)
dy = denormalize(offsets_raw[0, 1].item(), core_dy_min, core_dy_max, dy_min, dy_max)
# 打印调试信息
print(f"Predicted Actions: {actions}, dx: {dx:.2f}, dy: {dy:.2f}")
return list(actions) + [dx, dy]
# ---- 游戏控制逻辑 ----
keyboard = KeyboardController()
key_mappings = {
"W": "w",
"S": "s",
"A": "a",
"D": "d",
"R": "r",
"Space": Key.space,
"Shift": Key.shift,
}
def move_mouse_relative(dx, dy):
"""
使用 pydirectinput 相对控制鼠标移动。
"""
pydirectinput.moveRel(int(dx * 10), int(dy * 10)) # 控制鼠标移动幅度
def control_game(actions):
"""
根据推理结果控制游戏。
输入:
- actions: [W, S, A, D, R, Space, Shift, MouseLeft, dx, dy]
"""
keys = ["W", "S", "A", "D", "R", "Space", "Shift", "MouseLeft"]
# 控制按键
for i, key in enumerate(keys[:-2]): # 遍历动作按键
if actions[i] == 1:
if key in key_mappings:
keyboard.press(key_mappings[key])
else:
if key in key_mappings:
keyboard.release(key_mappings[key])
# 鼠标操作
dx, dy = actions[8], actions[9]
move_mouse_relative(dx, dy)
# 鼠标左键点击
if actions[7] == 1:
pydirectinput.mouseDown()
else:
pydirectinput.mouseUp()
# ---- 多线程逻辑 ----
def capture_screen(monitor, frame_queue):
"""
截屏线程:捕获屏幕图像并传入队列
"""
with mss.mss() as sct:
while True:
screenshot = sct.grab(monitor)
image = Image.frombytes("RGB", screenshot.size, screenshot.rgb)
frame_queue.put(image) # 放入帧队列
def run_inference(frame_queue, action_queue):
"""
推理线程:从图像队列中取出帧并进行预测,将动作放入动作队列
"""
while True:
if not frame_queue.empty():
image = frame_queue.get()
actions = predict(image)
action_queue.put(actions)
def execute_controls(action_queue):
"""
控制线程:从动作队列中取出动作并实时执行控制
"""
while True:
if not action_queue.empty():
actions = action_queue.get()
control_game(actions)
def listen_for_exit():
"""
按下 F4 键时设置退出事件
"""
def on_press(key):
if key == keyboard.Key.f4: # 如果按下 F4
print("F4 pressed. Exiting.")
exit_event.set() # 触发退出事件
return False # 停止监听
with keyboard.Listener(on_press=on_press) as listener:
listener.join()
# ---- 主程序入口 ----
if __name__ == "__main__":
# 屏幕区域
monitor = {"top": 0, "left": 0, "width": 1360, "height": 768}
# 创建队列
frame_queue = queue.Queue(maxsize=10) # 图像缓存队列
action_queue = queue.Queue(maxsize=10) # 动作缓存队列
# 启动线程
capture_thread = threading.Thread(target=capture_screen, args=(monitor, frame_queue))
inference_thread = threading.Thread(target=run_inference, args=(frame_queue, action_queue))
control_thread = threading.Thread(target=execute_controls, args=(action_queue,))
exit_thread = threading.Thread(target=listen_for_exit, daemon=True)
capture_thread.start()
inference_thread.start()
control_thread.start()
exit_thread.start()
try:
capture_thread.join()
inference_thread.join()
control_thread.join()
except KeyboardInterrupt:
print("程序已停止")ASSISTANT
你给出的代码相当全面,其中包含了神经网络推理、图像截取、动作执行、线程管理等功能。以下是你实现的单线程和实时推理动作控制代码的主要架构解析,以及一些小问题确认和改进建议。
---
### **代码功能总结**
1. **神经网络模型定义 (`MultiTaskResNet`):**
- 通过 ResNet50 提取特征。
- 用两个任务头:`action_head` 处理动作分类,`offset_head` 处理鼠标坐标偏移 (dx/dy)。
- 使用模型预测玩家动作 (按键) 和鼠标移动。
2. **图像预处理和推理:**
- 使用 `torchvision.transforms` 对输入图像进行标准化和尺寸调整。
- 对每帧截屏图像进行模型推理,得到包括按键和鼠标偏移量的指令。
3. **动作映射与控制执行:**
- 按照模型输出,控制按键 (如 "W", "A", "D", 等) 按下或释放。
- 使用 `pydirectinput` 实现鼠标移动和点击。
4. **多线程实现:**
- 线程 1:实时捕获屏幕图像。
- 线程 2:对图像进行推理,获取动作和鼠标偏移。
- 线程 3:执行动作到游戏。
- 辅助功能:监听 `F4` 退出程序。
---
### **改进建议与注意事项**
#### 1. **模型加载时的设备问题**
```python
model = model.float()
```
可以移到 `model.to(device).eval()` 前,确保 `.float()` 操作能正确切换至适配的设备。或者直接合成一行:
```python
model = model.to(device).float().eval()
```
#### 2. **`dx`, `dy` 解算映射边界问题**
你的 `denormalize` 函数实现存在多段区间归一化的逻辑,这是合理的,但是提供的范围 `(dx_min, dx_max), (core_dx_min, core_dx_max)` 应该根据你的鼠标控制特性调整,小心因为鼠标移动出问题导致偏移不可用:
- 增加范围边界的验证:
```python
# 在 denormalize 结果上加一个范围限制
dx = max(dx_min, min(dx_max, dx))
dy = max(dy_min, min(dy_max, dy))
```
- 同时用于调试,建议打印实际偏移值输入输出:
```python
print(f"Offset RAW: dx_raw={offsets_raw[0,0]:.2f} dy_raw={offsets_raw[0,1]:.2f}")
print(f"Offset Mapped: dx={dx:.2f}, dy={dy:.2f}")
```
#### 3. **`frame_queue` 和 `action_queue` 的调整**
目前的 `Queue` 设置 `maxsize=10`,可能会因为生产-消费不均而引起阻塞、延迟。如果你的模型推理扩展性需要保证高吞吐量,合适的方向是:减小队列缓存。
例子:
```python
frame_queue = queue.Queue(maxsize=3) # 限制缓存队列大小
action_queue = queue.Queue(maxsize=3)
```
#### 4. **线程退出的友好处理**
代码实现了捕获 `KeyboardInterrupt` 的退出机制,但未能保证各个线程安全终止,建议引入线程共享 `event` 控制全局运行状态。例如:
```python
exit_event = threading.Event()
# 修改线程运行的主循环条件
while not exit_event.is_set():
...
```
对于 `capture_screen`、`run_inference` 和 `execute_controls` 等运行时的 `True` 循环改为:
```python
if exit_event.is_set(): # 判断退出
break
```
#### 5. **线程 Listener 报错问题**
原有的 `listen_for_exit()` 代码使用了 `keyboard.Listener`,但脚本没有正确导入 `pynput.keyboard.Listener`,需要补充:
```python
from pynput.keyboard import Listener
```
---
### 更新后的代码主函数入口(改进版)
综合线程管理、退出处理和调试优化后的主代码框架如下:
```python
if __name__ == "__main__":
# 屏幕区域设置
monitor = {"top": 0, "left": 0, "width": 1360, "height": 768}
# 创建队列
frame_queue = queue.Queue(maxsize=5) # 图像缓存队列
action_queue = queue.Queue(maxsize=5) # 动作缓存队列
# 创建全局退出事件
exit_event = threading.Event()
# 定义线程
capture_thread = threading.Thread(target=capture_screen, args=(monitor, frame_queue))
inference_thread = threading.Thread(target=run_inference, args=(frame_queue, action_queue))
control_thread = threading.Thread(target=execute_controls, args=(action_queue,))
def listen_for_exit():
"""
捕获 F4 按键退出程序
"""
def on_press(key):
if key == Key.f4: # 如果按下 F4
print("F4 pressed. Exiting.")
exit_event.set() # 触发退出事件
return False # 停止监听
with Listener(on_press=on_press) as listener:
listener.join()
# 启动主线程
try:
# 启动子线程
capture_thread.start()
threading.Thread(target=listen_for_exit, daemon=True).start()
inference_thread.start()
control_thread.start()
# 等待线程运行
while not exit_event.is_set():
pass # 主线程等待退出
except KeyboardInterrupt:
print("捕获到 Ctrl+C,正在安全退出...")
exit_event.set()
# 等待线程退出完成
capture_thread.join()
inference_thread.join()
control_thread.join()
print("程序已安全退出。")
```
---
**小结**
- 上述代码优化了退出机制、调试建议和队列缓存限制。单线程配合 GPU 推理已经足够应对大多数实时推理任务。
- 若需进一步优化,可考虑 `multiprocessing` 替代 `threading`,提高多核推理能力,或基于更高吞吐需求缩小 `transform` 和 `sigmoid` 等非必要模块的运行负担。