Files
change_hair_3090/meidaojia/worker.py
T
colomi 0eb61f3e60 初始化换发型项目:3个微服务代码 + 部署脚本
包含:
- hair_service_sd: 换发型/换发色算法服务 (端口 8801)
- photo_service: LoRA 训练调度服务 (端口 32678)
- stable-diffusion-webui: SD WebUI 推理服务 (端口 57860)
- kohya_ss_home: 训练环境代码
- meidaojia: 监控测试脚本
- setup.sh: 一键部署脚本 (conda环境恢复 + 配置生成 + 完整性检查)
- start_all_services.sh: 启动3个服务
- configure.ini.template: 路径模板化 (BASE_DIR自动推导)
- conda_envs/py310.yml: py310 环境定义

大文件 (weights/, models/, data/, conda_envs/*.tar.gz 等) 通过 .gitignore 排除,
由网盘单独上传。
2026-07-11 18:11:49 +08:00

428 lines
15 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import time
import datetime
import uuid
import redis
import logging
import requests
import threading
import json
from datetime import datetime, timedelta
from config import REDIS_HOST
from config import REDIS_PORT
from config import REDIS_DB
from config import QUEUE_NAME
from config import GPU_SERVER_LIST
from config import SERVER_LOG_EVENT_LEN
from config import SERVER_LOG_
from config import GPU_SERVER_TIME_OUT
from config import KEY_QUEUE_LOCK_NAME, acquire_lock, SERVER_LIST_LOCK_NAME, SERVER_LIST_LOG_LOCK_NAME
import copy
import logging
import csv
# 创建 logger
logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO) # 设置 logger 的级别
# 获取当前日期和时间
now = datetime.now()
# 格式化为字符串(例如:2023-10-25 14:30:45
date_time_str = now.strftime("%Y-%m-%d_%H:%M:%S")
# 创建文件 handler
file_handler = logging.FileHandler(f'/var/log/meidaojia/worker_{date_time_str}.log')
file_handler.setLevel(logging.INFO)
file_formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
file_handler.setFormatter(file_formatter)
# 创建控制台 handler
console_handler = logging.StreamHandler()
console_handler.setLevel(logging.INFO)
console_formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
console_handler.setFormatter(console_formatter)
# 添加 handlers 到 logger
logger.addHandler(file_handler)
logger.addHandler(console_handler)
# 创建 Redis 连接池
redis_pool = redis.ConnectionPool(
host=REDIS_HOST,
port=REDIS_PORT,
db=REDIS_DB,
max_connections=2000 # 根据实际情况调整
)
def get_redis_conn():
"""获取 Redis 连接"""
return redis.Redis(connection_pool=redis_pool)
redis_conn = get_redis_conn()
def get_time(timestamp):
# 转换为UTC时间
utc_time = datetime.utcfromtimestamp(timestamp)
# 转换为北京时间 (UTC+8)
beijing_time = utc_time + timedelta(hours=8) # 正确
# 格式化输出
formatted_time = beijing_time.strftime("%Y-%m-%d %H:%M:%S")
return formatted_time
def server_log_event(name, action, in_data, success=True):
log_data = copy.deepcopy(in_data)
server_log_name = f"{SERVER_LOG_}{name}"
lock = acquire_lock(redis_conn, f"server_log_lock_{name}")
if lock:
try:
logger.info(f"server_log_event name:{name} action:{action} success:{success}")
server_log_str = redis_conn.get(server_log_name).decode("utf-8")
server_log = json.loads(server_log_str)
event = {}
t = time.time()
event['time'] = t
event['time_str'] = get_time(t)
event['action'] = action
server_log['last_call'] = "success"
if 'state' in log_data:
if log_data['state'] != 0:
server_log['last_call'] = "failure"
if not success:
server_log['last_call'] = "failure"
if 'data' in log_data:
imgstr = log_data['data']
imglen = len(imgstr)
if imglen > 256:
log_data['data'] = imgstr[:64]
if 'img' in log_data:
imgstr = log_data['img']
imglen = len(imgstr)
if imglen > 256:
log_data['img'] = imgstr[:64]
if 'result' in log_data:
imgstr = log_data['result']
imglen = len(imgstr)
if imglen > 256:
log_data['result'] = imgstr[:64]
if 'user_img_path' in log_data:
imgstr = log_data['user_img_path']
imglen = len(imgstr)
if imglen > 256:
log_data['user_img_path'] = imgstr[:32]
event['data'] = log_data
if(len(server_log['events']) > SERVER_LOG_EVENT_LEN):
server_log['events'].pop(0)
server_log['events'].append(event)
server_log['last_event'] = event
server_log_str = json.dumps(server_log)
redis_conn.set(f"{SERVER_LOG_}{name}", server_log_str)
finally:
lock.release()
else:
logger.error(f"get lock none {server_log_name}")
def registerGpuServer(name, url, can_use):
try:
server_list_str = redis_conn.get(GPU_SERVER_LIST).decode("utf-8")
server_list = json.loads(server_list_str)
if name not in server_list:
server_list.append(name)
server = {}
server['name'] = name
server['url'] = url
server['can_use'] = can_use
server_str = json.dumps(server)
logger.info(f"registerGpuServer, {server_str}")
redis_conn.set(name, server_str)
server_log = {}
server_log['name'] = name
event = {}
t = time.time()
event['time'] = t
event['time_str'] = get_time(t)
event['action'] = "register"
event['data'] = f"url {url}"
server_log['last_event'] = event
server_log['events'] = []
server_log['last_call'] = "success"
if(len(server_log['events']) > SERVER_LOG_EVENT_LEN):
server_log['events'].pop(0)
server_log['events'].append(event)
server_log_str = json.dumps(server_log)
redis_conn.set(f"{SERVER_LOG_}{name}", server_log_str)
redis_conn.set(GPU_SERVER_LIST, json.dumps(server_list))
except Exception as e:
logger.error(f"registerGpuServer{e.message}")
def get_remote_gpu_server():
start_time = datetime.now()
while (datetime.now() - start_time).seconds < GPU_SERVER_TIME_OUT:
lock = acquire_lock(redis_conn, SERVER_LIST_LOCK_NAME)
try:
server_list_str = redis_conn.get(GPU_SERVER_LIST).decode("utf-8")
server_list = json.loads(server_list_str)
for name in server_list:
server_str = redis_conn.get(name).decode("utf-8")
server = json.loads(server_str)
if server['can_use']:
return server
finally:
lock.release()
time.sleep(0.2)
return None
def call_remote_gpu_server(task_data_str, server=None):
task_data = json.loads(task_data_str)
key = task_data['key']
if server == None:
server = get_remote_gpu_server()
result_str = ''
result_status_code = '200'
if server:
headers = {
'Content-Type': 'application/json'
}
if 'hairColor' in task_data['api']:
data = {
"img": task_data['request']['img'],
"rgb": task_data['request']["rgb"],
"ratio": task_data['request']['ratio'],
"userId": task_data['request']['userId'],
"output_format":task_data['request']['output_format']
}
else:
data = {
"hair_id": task_data['request']['hair_id'],
"task_id": task_data['request']["task_id"],
"user_img_path": task_data['request']['user_img_path'],
"is_hr": task_data['request']['is_hr'],
"output_format":task_data['request']['output_format']
}
name = server['name']
# set server busy
server['can_use'] = False
server_str = json.dumps(server)
redis_conn.set(name, server_str)
server_log_event(server['name'], "Request_GpuServer", data, True)
url = server['url'] + task_data['api']
try:
logger.info(f"call get_remote_gpu_server {url}")
start_ms = int(time.time() * 1000) # 转换为毫秒级整数
base64 = False
if data['output_format'] == 'base64':
base64 = True
# if 'hairColor' in task_data['api']:
# if data['img'] == 'base64':
# imgstr = redis_conn.get(f"img_{key}").decode("utf-8")
# data['img'] = imgstr
# else:
# if data['user_img_path'] == 'base64':
# imgstr = redis_conn.get(f"img_{key}").decode("utf-8")
# data['user_img_path'] = imgstr
response = requests.post(url, headers=headers, json=data, timeout=30)
logger.info(f"call finshed {response.status_code} {response.headers} {response.text[:128]}")
if response.status_code == 200:
result = response.json()
end_ms = int(time.time() * 1000)
if 'hairColor' in task_data['api']:
if base64:
result['result'] = f"data:image/jpeg;base64,{result['result']}"
else:
if base64:
result['data'] = f"data:image/jpeg;base64,{result['data']}"
result['process_time_ms'] = (end_ms - start_ms)
server_log_event(server['name'], "GpuServer_Response", result, True)
result_str = json.dumps(result)
result_status_code = '200'
if 'state' in result:
if result['state'] != 0:
result_status_code = '500'
else:
try:
result = response.json()
except json.JSONDecodeError:
result = {
"state":-1,
"data":"",
'msg': f'status_code error{response.status_code} call_remote_gpu_server{url} {response.text}'
}
time.sleep(1)
result_str = json.dumps(result)
result_status_code = str(response.status_code)
server_log_event(server['name'], f"error GpuServer_Response ", result, False)
except Exception as e:
logger.info(f"End call Exception {url} {type(e).__name__}, 错误信息: {e}")
time.sleep(1)
result = {
"state":-1,
"data":"",
'msg': f'Exception call call_remote_gpu_server{url}'
}
result_str = json.dumps(result)
result_status_code = '500'
server_log_event(server['name'], f"error GpuServer_Response ", result, False)
logger.info(f"End call {url}")
# release server
server['can_use'] = True
server_str = json.dumps(server)
redis_conn.set(name, server_str)
else:
result_str = json.dumps({
"state":-1,
"data":"",
'msg': f'Timeout after all gpu server is busy'
})
result_status_code = '500'
result_key = f"result_{key}"
redis_conn.set(result_key, result_str, ex=60)
result_status_key = f"result_status_code_{key}"
redis_conn.set(result_status_key, result_status_code, ex=60)
def check_server_work_call(server):
task_data = {}
request = {}
request['hair_id'] = "1907651680352395265"
request['task_id'] = "1907651680352395265"
request['user_img_path'] = "https://cdn.meidaojia.com/ZoeFiles/user2_1_%E5%89%AF%E6%9C%AC.JPG"
request['output_format'] = "url"
request['is_hr'] = "false"
task_data['request'] = request
key = str(uuid.uuid4())
task_data['key'] = key
task_data['api'] = "/api/swapHair/v1"
task_data_str = json.dumps(task_data)
threading.Thread(target=call_remote_gpu_server, args=(task_data_str, server)).start()
def check_server_work():
server_list_str = redis_conn.get(GPU_SERVER_LIST).decode("utf-8")
server_list = json.loads(server_list_str)
for name in server_list:
server_str = redis_conn.get(name).decode("utf-8")
server = json.loads(server_str)
if not server['can_use']:
continue
server_log_name = f"{SERVER_LOG_}{name}"
server_log_str = redis_conn.get(server_log_name).decode("utf-8")
server_log = json.loads(server_log_str)
if time.time() - server_log['last_event']['time'] > (5*60):
logger.info("check_server_work and call name")
check_server_work_call(server)
def get_task_queue_key():
redis_conn = get_redis_conn()
lock = acquire_lock(redis_conn, KEY_QUEUE_LOCK_NAME)
if not lock:
logger.error("get_task_queue_key get lock error")
return
try:
key_queue_str = redis_conn.get(QUEUE_NAME).decode("utf-8")
key_queue = json.loads(key_queue_str)
queue_len = len(key_queue)
if(queue_len > 0):
logger.info(f"get_task_queue_key queue len:{queue_len}")
key = key_queue.pop(0)
key_queue_str = json.dumps(key_queue)
logger.info(f"get_task_queue_key queue len:{len(key_queue)}")
redis_conn.set(QUEUE_NAME, key_queue_str)
return key
else:
return None
finally:
lock.release()
def main_worker():
while True:
try:
key = get_task_queue_key()
if key:
task_data_str = redis_conn.get(key)
if task_data_str:
threading.Thread(target=call_remote_gpu_server, args=(task_data_str,)).start()
else:
logging.error(f"main_worker get task_data_str error with key{key}")
continue
except Exception as e:
logger.info(f"main_worker {e}")
# check_server_work()
time.sleep(0.1)
serverindex = 0
def regServer(instance):
global serverindex
registerGpuServer(f"s{serverindex}_1_{instance}", f"https://{instance}-http-8801.northwest1.gpugeek.com:8443", True)
# registerGpuServer(f"s{serverindex}_2_{instance}", f"https://{instance}-http-8801.northwest1.gpugeek.com:8443", True)
serverindex += 1
def get_server_names(csv_file='servers.csv'):
"""
从servers.csv文件中提取所有主机名称
参数:
csv_file (str): CSV文件路径,默认为'servers.csv'
返回:
list: 包含所有主机名称的列表
"""
server_names = []
try:
with open(csv_file, mode='r', encoding='utf-8') as file:
csv_reader = csv.reader(file)
for row in csv_reader:
if row: # 确保不是空行
# 取第一列并去除前后空格
server_name = row[0].strip()
if server_name: # 确保名称不为空
server_names.append(server_name)
except FileNotFoundError:
print(f"错误:文件 {csv_file} 未找到")
except Exception as e:
print(f"读取文件时出错: {e}")
return server_names
if __name__ == '__main__':
import os
print(f"[Worker] Starting with PID: {os.getpid()}")
redis_conn.set(GPU_SERVER_LIST, "[]")
# servers = get_server_names()
# for s in servers:
# regServer(s)
regServer("693570664882181")
regServer("693565557084165")
main_worker()