This commit is contained in:
parent
8ce3afbe19
commit
f5e535d16d
19
.env
19
.env
|
|
@ -2,21 +2,26 @@
|
||||||
MODELSCOPE_TOKEN=ms-00353db2-f626-47b9-8c59-d535824fd7fb
|
MODELSCOPE_TOKEN=ms-00353db2-f626-47b9-8c59-d535824fd7fb
|
||||||
MODEL_CACHE_DIR=/app/model_cache
|
MODEL_CACHE_DIR=/app/model_cache
|
||||||
|
|
||||||
|
# ==================== 统一服务配置 ====================
|
||||||
|
SERVICE_PORT=5003
|
||||||
|
DEVICE=cpu
|
||||||
|
MAX_WORKERS=4
|
||||||
|
|
||||||
# ==================== 税务分类模型 ====================
|
# ==================== 税务分类模型 ====================
|
||||||
TAX_MODEL_ID=circles1/tax_rate
|
TAX_MODEL_ID=circles1/tax_rate
|
||||||
TAX_MODEL_DIR=/app/services/tax/model
|
TAX_MODEL_DIR=/app/services/tax/model
|
||||||
TAX_PORT=5004
|
|
||||||
TAX_SERVICE_NAME=tax_classifier
|
TAX_SERVICE_NAME=tax_classifier
|
||||||
|
TAX_TYPE=tax
|
||||||
|
|
||||||
# ==================== 订单地址模型 ====================
|
# ==================== 订单地址模型 ====================
|
||||||
ADDRESS_MODEL_ID=circles1/order_address
|
ADDRESS_MODEL_ID=circles1/order_address
|
||||||
ADDRESS_MODEL_DIR=/app/services/address/model
|
ADDRESS_MODEL_DIR=/app/services/address/model
|
||||||
ADDRESS_PORT=5003
|
|
||||||
ADDRESS_SERVICE_NAME=address_classifier
|
ADDRESS_SERVICE_NAME=address_classifier
|
||||||
|
ADDRESS_TYPE=address
|
||||||
|
|
||||||
|
# ==================== 品牌模型 ====================
|
||||||
|
BRAND_MODEL_ID=circles1/brand
|
||||||
|
BRAND_MODEL_DIR=/app/services/brand/model
|
||||||
|
BRAND_SERVICE_NAME=brand_classifier
|
||||||
|
BRAND_TYPE=brand
|
||||||
|
|
||||||
# ==================== 运行配置 ====================
|
|
||||||
DEVICE=cpu
|
|
||||||
MAX_WORKERS=4
|
|
||||||
BATCH_SIZE=4
|
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,21 @@
|
||||||
|
@echo off
|
||||||
|
chcp 65001 >nul
|
||||||
|
|
||||||
|
echo ==========================================
|
||||||
|
echo Deploying bert-all-services...
|
||||||
|
echo ==========================================
|
||||||
|
|
||||||
|
|
||||||
|
echo.
|
||||||
|
echo [2] Stopping and removing old container...
|
||||||
|
docker compose down 2>nul
|
||||||
|
|
||||||
|
echo.
|
||||||
|
echo [3] Building and starting services...
|
||||||
|
docker compose up -d --build
|
||||||
|
|
||||||
|
echo.
|
||||||
|
echo [4] Showing logs (Ctrl+C to exit)...
|
||||||
|
docker compose logs -f
|
||||||
|
|
||||||
|
pause
|
||||||
|
|
@ -6,7 +6,7 @@ services:
|
||||||
container_name: bert-all-services
|
container_name: bert-all-services
|
||||||
user: "0:0" # 显式使用 root
|
user: "0:0" # 显式使用 root
|
||||||
ports:
|
ports:
|
||||||
- "5002-5010:5002-5010"
|
- "5003:5003"
|
||||||
env_file:
|
env_file:
|
||||||
- .env
|
- .env
|
||||||
environment:
|
environment:
|
||||||
|
|
|
||||||
|
|
@ -1,15 +1,12 @@
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import shutil
|
import shutil
|
||||||
|
import subprocess
|
||||||
from modelscope.hub.snapshot_download import snapshot_download
|
from modelscope.hub.snapshot_download import snapshot_download
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
load_dotenv()
|
load_dotenv()
|
||||||
|
|
||||||
# 使用持久化的缓存目录(挂载到宿主机)
|
|
||||||
CACHE_DIR = os.getenv('MODEL_CACHE_DIR', '/app/model_cache')
|
|
||||||
os.makedirs(CACHE_DIR, exist_ok=True)
|
|
||||||
|
|
||||||
|
|
||||||
def discover_models():
|
def discover_models():
|
||||||
"""从环境变量中发现所有模型"""
|
"""从环境变量中发现所有模型"""
|
||||||
|
|
@ -19,100 +16,32 @@ def discover_models():
|
||||||
prefix = key[:-9]
|
prefix = key[:-9]
|
||||||
service_name = os.getenv(f'{prefix}_SERVICE_NAME', f'{prefix.lower()}_classifier')
|
service_name = os.getenv(f'{prefix}_SERVICE_NAME', f'{prefix.lower()}_classifier')
|
||||||
service_folder = service_name.replace('_classifier', '')
|
service_folder = service_name.replace('_classifier', '')
|
||||||
|
model_type = os.getenv(f'{prefix}_TYPE', 'general')
|
||||||
|
|
||||||
models[prefix.lower()] = {
|
models[prefix.lower()] = {
|
||||||
'prefix': prefix,
|
'prefix': prefix,
|
||||||
'model_id': value,
|
'model_id': value,
|
||||||
'port': os.getenv(f'{prefix}_PORT', '5002'),
|
'port': os.getenv(f'{prefix}_PORT', '5002'),
|
||||||
'service_name': service_name,
|
'service_name': service_name,
|
||||||
'model_dir': os.getenv(f'{prefix}_MODEL_DIR', f'/app/services/bert_{service_folder}/model'),
|
'model_dir': os.getenv(f'{prefix}_MODEL_DIR', f'/app/services/{service_folder}/model'),
|
||||||
'service_folder': service_folder,
|
'service_folder': service_folder,
|
||||||
'cache_dir': os.path.join(CACHE_DIR, prefix.lower())
|
'model_type': model_type,
|
||||||
|
'cache_dir': os.path.join('/app/model_cache', prefix.lower())
|
||||||
}
|
}
|
||||||
return models
|
return models
|
||||||
|
|
||||||
|
|
||||||
def check_model_complete(model_dir):
|
|
||||||
"""检查模型是否完整"""
|
|
||||||
if not os.path.exists(model_dir):
|
|
||||||
return False
|
|
||||||
|
|
||||||
try:
|
|
||||||
files = os.listdir(model_dir)
|
|
||||||
if not files:
|
|
||||||
return False
|
|
||||||
|
|
||||||
has_weights = any(f.endswith(('.safetensors', '.bin', '.pt', '.pth')) for f in files)
|
|
||||||
has_config = any(f == 'config.json' for f in files)
|
|
||||||
|
|
||||||
if has_weights and has_config:
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
except Exception:
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def find_model_files(cache_dir):
|
def find_model_files(cache_dir):
|
||||||
"""在下载的缓存目录中查找实际的模型文件"""
|
"""查找实际的模型文件位置"""
|
||||||
# 检查是否直接有模型文件
|
if os.path.exists(cache_dir):
|
||||||
direct_files = [f for f in os.listdir(cache_dir) if
|
for root, dirs, files in os.walk(cache_dir):
|
||||||
not f.startswith('.') and os.path.isfile(os.path.join(cache_dir, f))]
|
if any(f.endswith(('.safetensors', '.bin')) for f in files):
|
||||||
if any(f.endswith(('.safetensors', '.bin', '.pt', '.pth')) for f in direct_files):
|
|
||||||
return cache_dir
|
|
||||||
|
|
||||||
# 检查 models/ 子目录(ModelScope 的标准结构)
|
|
||||||
models_dir = os.path.join(cache_dir, 'models')
|
|
||||||
if os.path.exists(models_dir):
|
|
||||||
for root, dirs, files in os.walk(models_dir):
|
|
||||||
if 'snapshots' in root:
|
|
||||||
for f in files:
|
|
||||||
if f.endswith(('.safetensors', '.bin', '.pt', '.pth')):
|
|
||||||
return root
|
|
||||||
if any(f.endswith(('.safetensors', '.bin', '.pt', '.pth')) for f in files):
|
|
||||||
return root
|
return root
|
||||||
|
|
||||||
# 检查是否存在单层子目录
|
|
||||||
subdirs = [d for d in os.listdir(cache_dir) if os.path.isdir(os.path.join(cache_dir, d)) and not d.startswith('.')]
|
|
||||||
for subdir in subdirs:
|
|
||||||
subpath = os.path.join(cache_dir, subdir)
|
|
||||||
if any(f.endswith(('.safetensors', '.bin', '.pt', '.pth')) for f in os.listdir(subpath)):
|
|
||||||
return subpath
|
|
||||||
|
|
||||||
# 递归查找任何包含模型文件的目录
|
|
||||||
for root, dirs, files in os.walk(cache_dir):
|
|
||||||
if any(f.endswith(('.safetensors', '.bin', '.pt', '.pth')) for f in files):
|
|
||||||
return root
|
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def copy_model_from_cache(cache_dir, target_dir):
|
def download_model(name, config):
|
||||||
"""从缓存目录复制模型到目标目录"""
|
"""下载单个模型"""
|
||||||
src_dir = find_model_files(cache_dir)
|
|
||||||
|
|
||||||
if src_dir is None:
|
|
||||||
print(f" ✗ Could not find model files in cache: {cache_dir}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
print(f" ✓ Found cached model files in: {src_dir}")
|
|
||||||
|
|
||||||
os.makedirs(target_dir, exist_ok=True)
|
|
||||||
|
|
||||||
copied_count = 0
|
|
||||||
for file in os.listdir(src_dir):
|
|
||||||
src_file = os.path.join(src_dir, file)
|
|
||||||
if os.path.isfile(src_file):
|
|
||||||
dest_file = os.path.join(target_dir, file)
|
|
||||||
shutil.copy2(src_file, dest_file)
|
|
||||||
copied_count += 1
|
|
||||||
|
|
||||||
print(f" ✓ Copied {copied_count} files to {target_dir}")
|
|
||||||
return copied_count > 0
|
|
||||||
|
|
||||||
|
|
||||||
def download_model_with_fallback(name, config):
|
|
||||||
"""下载单个模型,带备选方案"""
|
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
print(f"Processing {name} model...")
|
print(f"Processing {name} model...")
|
||||||
print(f"Model ID: {config['model_id']}")
|
print(f"Model ID: {config['model_id']}")
|
||||||
|
|
@ -121,92 +50,71 @@ def download_model_with_fallback(name, config):
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
# 确保目录存在
|
||||||
os.makedirs(config['model_dir'], exist_ok=True)
|
os.makedirs(config['model_dir'], exist_ok=True)
|
||||||
os.makedirs(config['cache_dir'], exist_ok=True)
|
os.makedirs(config['cache_dir'], exist_ok=True)
|
||||||
|
|
||||||
# 1. 检查目标目录是否已有完整模型
|
# 检查是否已有模型文件
|
||||||
if check_model_complete(config['model_dir']):
|
if os.path.exists(config['model_dir']):
|
||||||
print(f"✅ Model already exists in target, skipping")
|
files = [f for f in os.listdir(config['model_dir']) if f.endswith(('.safetensors', '.bin'))]
|
||||||
return True
|
if files:
|
||||||
|
print(f"✓ Model already exists in {config['model_dir']} (found {len(files)} files)")
|
||||||
|
return True
|
||||||
|
|
||||||
# 2. 检查缓存目录是否有模型
|
# 登录 ModelScope
|
||||||
if check_model_complete(config['cache_dir']):
|
token = os.getenv('MODELSCOPE_TOKEN', '')
|
||||||
print(f"✅ Found complete model in local cache")
|
if token:
|
||||||
print(f" Copying from cache to target...")
|
result = subprocess.run(
|
||||||
|
['modelscope', 'login', '--token', token],
|
||||||
if copy_model_from_cache(config['cache_dir'], config['model_dir']):
|
capture_output=True,
|
||||||
if check_model_complete(config['model_dir']):
|
text=True
|
||||||
print(f"✅ Model copied successfully from cache")
|
|
||||||
return True
|
|
||||||
else:
|
|
||||||
print(f"⚠️ Model copied but appears incomplete, will re-download")
|
|
||||||
|
|
||||||
# 3. 缓存中没有,需要下载
|
|
||||||
print(f"📥 Downloading model from ModelScope...")
|
|
||||||
|
|
||||||
download_success = False
|
|
||||||
|
|
||||||
# 尝试方式1:直接下载
|
|
||||||
try:
|
|
||||||
snapshot_download(
|
|
||||||
model_id=config['model_id'],
|
|
||||||
cache_dir=config['cache_dir'],
|
|
||||||
revision="master",
|
|
||||||
ignore_file_pattern=[".git", ".gitattributes"]
|
|
||||||
)
|
)
|
||||||
print(f"✅ Download completed")
|
if result.returncode == 0:
|
||||||
download_success = True
|
print("✓ ModelScope login successful")
|
||||||
except Exception as e:
|
|
||||||
print(f"⚠️ Primary download failed: {e}")
|
|
||||||
|
|
||||||
# 尝试方式2:使用 HubApi(如果有 token)
|
|
||||||
token = os.getenv('MODELSCOPE_TOKEN', '')
|
|
||||||
if token:
|
|
||||||
print(" Trying alternative download with HubApi...")
|
|
||||||
try:
|
|
||||||
from modelscope.hub.api import HubApi
|
|
||||||
api = HubApi()
|
|
||||||
api.login(token=token)
|
|
||||||
|
|
||||||
snapshot_download(
|
|
||||||
model_id=config['model_id'],
|
|
||||||
cache_dir=config['cache_dir'],
|
|
||||||
revision="master",
|
|
||||||
ignore_file_pattern=[".git", ".gitattributes"]
|
|
||||||
)
|
|
||||||
print(f"✅ Alternative download completed")
|
|
||||||
download_success = True
|
|
||||||
except Exception as alt_error:
|
|
||||||
print(f"✗ Alternative download also failed: {alt_error}")
|
|
||||||
raise
|
|
||||||
else:
|
else:
|
||||||
print("✗ No MODELSCOPE_TOKEN set, cannot use alternative method")
|
print(f"⚠️ ModelScope login failed: {result.stderr}")
|
||||||
raise
|
|
||||||
|
|
||||||
if not download_success:
|
# 下载模型
|
||||||
|
print(f"Downloading model to {config['cache_dir']}...")
|
||||||
|
snapshot_download(
|
||||||
|
model_id=config['model_id'],
|
||||||
|
cache_dir=config['cache_dir'],
|
||||||
|
revision="master",
|
||||||
|
ignore_file_pattern=[".git", ".gitattributes"]
|
||||||
|
)
|
||||||
|
|
||||||
|
# 查找实际模型文件
|
||||||
|
src_dir = find_model_files(config['cache_dir'])
|
||||||
|
|
||||||
|
if src_dir is None:
|
||||||
|
print(f"✗ Could not find model files in downloaded content")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# 验证并复制
|
print(f"✓ Found model files in: {src_dir}")
|
||||||
if not check_model_complete(config['cache_dir']):
|
|
||||||
print(f"⚠️ Downloaded model appears incomplete in cache")
|
|
||||||
print(f" Attempting to copy anyway...")
|
|
||||||
|
|
||||||
print(f"Copying from cache to target...")
|
# 复制文件到目标目录
|
||||||
if copy_model_from_cache(config['cache_dir'], config['model_dir']):
|
print(f"Copying files to {config['model_dir']}...")
|
||||||
if check_model_complete(config['model_dir']):
|
copied_count = 0
|
||||||
print(f"✅ Model deployed successfully!")
|
for file in os.listdir(src_dir):
|
||||||
return True
|
src_file = os.path.join(src_dir, file)
|
||||||
else:
|
if os.path.isfile(src_file):
|
||||||
print(f"⚠️ Model deployed but appears incomplete")
|
dest_file = os.path.join(config['model_dir'], file)
|
||||||
return True
|
shutil.copy2(src_file, dest_file)
|
||||||
|
copied_count += 1
|
||||||
|
print(f" Copied: {file}")
|
||||||
|
|
||||||
|
if copied_count > 0:
|
||||||
|
print(f"✓ {name} model downloaded successfully! ({copied_count} files)")
|
||||||
|
return True
|
||||||
else:
|
else:
|
||||||
print(f"✗ Failed to copy model to target")
|
print(f"✗ No files copied")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
except PermissionError as e:
|
||||||
|
print(f"✗ Permission error: {e}")
|
||||||
|
return False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"✗ Failed to process {name} model: {e}")
|
print(f"✗ Failed to download {name} model: {e}")
|
||||||
import traceback
|
|
||||||
traceback.print_exc()
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -214,13 +122,12 @@ def main():
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
print("Model Download Script")
|
print("Model Download Script")
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
print(f"Cache directory: {CACHE_DIR}")
|
|
||||||
|
|
||||||
force_download = os.getenv('FORCE_MODEL_DOWNLOAD', 'false').lower() == 'true'
|
import pwd
|
||||||
if force_download:
|
current_user = pwd.getpwuid(os.getuid()).pw_name
|
||||||
print("⚠️ FORCE_MODEL_DOWNLOAD=true - will re-download all models")
|
print(f"Running as user: {current_user}")
|
||||||
|
print("=" * 50)
|
||||||
|
|
||||||
print("\nDiscovering models from .env...")
|
|
||||||
models = discover_models()
|
models = discover_models()
|
||||||
|
|
||||||
if not models:
|
if not models:
|
||||||
|
|
@ -228,39 +135,19 @@ def main():
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
|
|
||||||
print(f"\nFound {len(models)} model(s):")
|
print(f"\nFound {len(models)} model(s):")
|
||||||
all_exist = True
|
|
||||||
for name, config in models.items():
|
for name, config in models.items():
|
||||||
target_exists = check_model_complete(config['model_dir'])
|
|
||||||
cache_exists = check_model_complete(config['cache_dir'])
|
|
||||||
|
|
||||||
print(f" {name}: {config['model_id']}")
|
print(f" {name}: {config['model_id']}")
|
||||||
print(f" Target: {'✅' if target_exists else '❌'} {config['model_dir']}")
|
print(f" Target: {config['model_dir']}")
|
||||||
print(f" Cache: {'✅' if cache_exists else '❌'} {config['cache_dir']}")
|
print(f" Type: {config['model_type']}")
|
||||||
|
|
||||||
if not target_exists and not cache_exists:
|
|
||||||
all_exist = False
|
|
||||||
print("")
|
print("")
|
||||||
|
|
||||||
if all_exist and not force_download:
|
print("=" * 50)
|
||||||
print("✅ All models exist!")
|
|
||||||
print(" Skipping downloads.")
|
|
||||||
sys.exit(0)
|
|
||||||
|
|
||||||
if force_download:
|
|
||||||
print("⚠️ Force download enabled, clearing caches...")
|
|
||||||
for name, config in models.items():
|
|
||||||
if os.path.exists(config['cache_dir']):
|
|
||||||
print(f" Clearing cache for {name}")
|
|
||||||
shutil.rmtree(config['cache_dir'])
|
|
||||||
os.makedirs(config['cache_dir'], exist_ok=True)
|
|
||||||
|
|
||||||
print("\n" + "=" * 50)
|
|
||||||
print("Processing models...")
|
print("Processing models...")
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
|
|
||||||
success_count = 0
|
success_count = 0
|
||||||
for name, config in models.items():
|
for name, config in models.items():
|
||||||
if download_model_with_fallback(name, config):
|
if download_model(name, config):
|
||||||
success_count += 1
|
success_count += 1
|
||||||
print("")
|
print("")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -26,7 +26,7 @@ def discover_models():
|
||||||
|
|
||||||
|
|
||||||
def generate_supervisor_conf(models):
|
def generate_supervisor_conf(models):
|
||||||
"""生成supervisor配置文件"""
|
"""生成supervisor配置文件 - 统一使用一个服务"""
|
||||||
conf = []
|
conf = []
|
||||||
|
|
||||||
conf.append("""[supervisord]
|
conf.append("""[supervisord]
|
||||||
|
|
@ -37,28 +37,23 @@ childlogdir=/var/log/supervisor
|
||||||
|
|
||||||
""")
|
""")
|
||||||
|
|
||||||
programs = []
|
# 只启动一个统一服务
|
||||||
for name, config in models.items():
|
conf.append(f"""[program:unified_service]
|
||||||
program_name = config['service_name']
|
command=gunicorn --bind 0.0.0.0:{os.environ.get('SERVICE_PORT', 5003)} --workers %(ENV_MAX_WORKERS)s --threads 2 --timeout 120 services.unified.app:app
|
||||||
service_folder = config['service_folder']
|
|
||||||
programs.append(program_name)
|
|
||||||
|
|
||||||
conf.append(f"""[program:{program_name}]
|
|
||||||
command=gunicorn --bind 0.0.0.0:{config['port']} --workers %(ENV_MAX_WORKERS)s --threads 2 --timeout 120 services.{service_folder}.app:app
|
|
||||||
directory=/app
|
directory=/app
|
||||||
autostart=true
|
autostart=true
|
||||||
autorestart=true
|
autorestart=true
|
||||||
startretries=3
|
startretries=3
|
||||||
stdout_logfile=/app/logs/{program_name}.log
|
stdout_logfile=/app/logs/unified_service.log
|
||||||
stdout_logfile_maxbytes=50MB
|
stdout_logfile_maxbytes=50MB
|
||||||
stderr_logfile=/app/logs/{program_name}_error.log
|
stderr_logfile=/app/logs/unified_service_error.log
|
||||||
stderr_logfile_maxbytes=50MB
|
stderr_logfile_maxbytes=50MB
|
||||||
environment=SERVICE_NAME="{program_name}",SERVICE_PORT="{config['port']}",MODEL_ID="{config['model_id']}",MODEL_DIR="{config['model_dir']}"
|
environment=PYTHONUNBUFFERED="1"
|
||||||
|
|
||||||
""") # 注意:移除了 user=appuser
|
""")
|
||||||
|
|
||||||
conf.append(f"""[group:bert_services]
|
conf.append(f"""[group:bert_services]
|
||||||
programs={','.join(programs)}
|
programs=unified_service
|
||||||
""")
|
""")
|
||||||
|
|
||||||
return ''.join(conf)
|
return ''.join(conf)
|
||||||
|
|
@ -84,7 +79,7 @@ def main():
|
||||||
|
|
||||||
print(f"Found {len(models)} model(s):")
|
print(f"Found {len(models)} model(s):")
|
||||||
for name, config in models.items():
|
for name, config in models.items():
|
||||||
print(f" - {name}: {config['model_id']} (port {config['port']})")
|
print(f" - {name}: {config['model_id']}")
|
||||||
|
|
||||||
conf_content = generate_supervisor_conf(models)
|
conf_content = generate_supervisor_conf(models)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -48,6 +48,11 @@ logging.basicConfig(
|
||||||
)
|
)
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# 如果使用 GPU,限制并发线程数为 1 以避免并发 GPU 推理导致 OOM
|
||||||
|
if isinstance(DEVICE, str) and DEVICE.startswith('cuda') and MAX_WORKERS > 1:
|
||||||
|
logger.warning("CUDA in use — limiting MAX_WORKERS to 1 to avoid concurrent GPU inference")
|
||||||
|
MAX_WORKERS = 1
|
||||||
|
|
||||||
# 打印启动信息
|
# 打印启动信息
|
||||||
logger.info("=" * 50)
|
logger.info("=" * 50)
|
||||||
logger.info(f"Starting service: {SERVICE_NAME}")
|
logger.info(f"Starting service: {SERVICE_NAME}")
|
||||||
|
|
@ -187,17 +192,45 @@ class Predictor:
|
||||||
logger.error(f"Prediction error for text '{text[:50]}...': {e}")
|
logger.error(f"Prediction error for text '{text[:50]}...': {e}")
|
||||||
return {"address": "", "confidence": 0.0, "error": str(e)}
|
return {"address": "", "confidence": 0.0, "error": str(e)}
|
||||||
|
|
||||||
def batch_predict(self, texts: List[str]) -> List[Dict]:
|
def batch_predict(self, texts: List[str], batch_size: int = 32) -> List[Dict]:
|
||||||
"""批量预测(并发处理)"""
|
"""按 batch 进行单次前向,避免为每个文本并发调用模型导致内存/GPU竞争。"""
|
||||||
if not texts:
|
if not texts:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
logger.info(f"Processing {len(texts)} texts")
|
logger.info(f"Processing {len(texts)} texts with batch_size={batch_size}")
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
# 使用线程池并发处理
|
results = []
|
||||||
futures = [self.executor.submit(self.predict_single, text) for text in texts]
|
total_batches = (len(texts) + batch_size - 1) // batch_size
|
||||||
results = [future.result() for future in futures]
|
|
||||||
|
for batch_idx in range(total_batches):
|
||||||
|
batch_start = batch_idx * batch_size
|
||||||
|
batch_end = min(batch_start + batch_size, len(texts))
|
||||||
|
batch_texts = texts[batch_start:batch_end]
|
||||||
|
|
||||||
|
# Tokenize batch on CPU then move to device
|
||||||
|
inputs = self.tokenizer(
|
||||||
|
batch_texts,
|
||||||
|
return_tensors='pt',
|
||||||
|
truncation=True,
|
||||||
|
padding=True,
|
||||||
|
max_length=MAX_LENGTH
|
||||||
|
)
|
||||||
|
inputs = {k: v.to(DEVICE) for k, v in inputs.items()}
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = self.model(**inputs)
|
||||||
|
probs = torch.softmax(outputs.logits, dim=1).cpu()
|
||||||
|
top_prob, top_idx = torch.topk(probs, k=1)
|
||||||
|
|
||||||
|
for i in range(len(batch_texts)):
|
||||||
|
results.append({
|
||||||
|
"address": self.label_map[top_idx[i].item()],
|
||||||
|
"confidence": round(top_prob[i].item(), 4)
|
||||||
|
})
|
||||||
|
|
||||||
|
if batch_idx % 10 == 0 and batch_idx > 0:
|
||||||
|
logger.info(f"Processed {batch_end}/{len(texts)} texts")
|
||||||
|
|
||||||
elapsed_time = time.time() - start_time
|
elapsed_time = time.time() - start_time
|
||||||
logger.info(f"Batch prediction completed in {elapsed_time:.2f}s for {len(texts)} texts")
|
logger.info(f"Batch prediction completed in {elapsed_time:.2f}s for {len(texts)} texts")
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,382 @@
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from flask import Flask, request, jsonify
|
||||||
|
import torch
|
||||||
|
from transformers import BertTokenizer, BertForSequenceClassification
|
||||||
|
import joblib
|
||||||
|
import re
|
||||||
|
from functools import lru_cache
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
from typing import List, Dict
|
||||||
|
import threading
|
||||||
|
import logging
|
||||||
|
import atexit
|
||||||
|
from datetime import datetime
|
||||||
|
import time
|
||||||
|
|
||||||
|
# ==================== 从环境变量读取配置 ====================
|
||||||
|
SERVICE_NAME = os.environ.get('SERVICE_NAME', 'brand')
|
||||||
|
SERVICE_PORT = int(os.environ.get('SERVICE_PORT', 5003))
|
||||||
|
MODEL_ID = os.environ.get('MODEL_ID', 'circles1/brand_model')
|
||||||
|
MODEL_DIR = os.environ.get('MODEL_DIR', '/app/services/bert_brand/model')
|
||||||
|
|
||||||
|
# 模型文件路径
|
||||||
|
LABEL_ENCODER_PATH = os.path.join(MODEL_DIR, 'label_encoder.pkl')
|
||||||
|
# 如果上面路径不存在,尝试其他可能的文件名
|
||||||
|
if not os.path.exists(LABEL_ENCODER_PATH):
|
||||||
|
import glob
|
||||||
|
|
||||||
|
encoder_files = glob.glob(os.path.join(MODEL_DIR, 'label_encoder*.pkl'))
|
||||||
|
if encoder_files:
|
||||||
|
LABEL_ENCODER_PATH = encoder_files[0]
|
||||||
|
|
||||||
|
# ==================== 配置参数 ====================
|
||||||
|
MAX_LENGTH = 512
|
||||||
|
DEVICE = os.environ.get('DEVICE', 'cuda' if torch.cuda.is_available() else 'cpu')
|
||||||
|
# 如果设备设置为cuda但不可用,回退到cpu
|
||||||
|
if DEVICE == 'cuda' and not torch.cuda.is_available():
|
||||||
|
DEVICE = 'cpu'
|
||||||
|
|
||||||
|
MAX_WORKERS = int(os.environ.get('MAX_WORKERS', 4))
|
||||||
|
CACHE_SIZE = int(os.environ.get('CACHE_SIZE', 2000))
|
||||||
|
TOKEN_CACHE_SIZE = int(os.environ.get('TOKEN_CACHE_SIZE', 1000))
|
||||||
|
|
||||||
|
# ==================== 日志配置 ====================
|
||||||
|
logging.basicConfig(
|
||||||
|
level=logging.INFO,
|
||||||
|
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||||
|
)
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# 如果使用 GPU,限制并发线程数为 1 以避免并发 GPU 推理导致 OOM
|
||||||
|
if isinstance(DEVICE, str) and DEVICE.startswith('cuda') and MAX_WORKERS > 1:
|
||||||
|
logger.warning("CUDA in use — limiting MAX_WORKERS to 1 to avoid concurrent GPU inference")
|
||||||
|
MAX_WORKERS = 1
|
||||||
|
|
||||||
|
# 打印启动信息
|
||||||
|
logger.info("=" * 50)
|
||||||
|
logger.info(f"Starting service: {SERVICE_NAME}")
|
||||||
|
logger.info(f"Service port: {SERVICE_PORT}")
|
||||||
|
logger.info(f"Model ID: {MODEL_ID}")
|
||||||
|
logger.info(f"Model directory: {MODEL_DIR}")
|
||||||
|
logger.info(f"Label encoder: {LABEL_ENCODER_PATH}")
|
||||||
|
logger.info(f"Device: {DEVICE}")
|
||||||
|
logger.info(f"Max workers: {MAX_WORKERS}")
|
||||||
|
logger.info("=" * 50)
|
||||||
|
|
||||||
|
app = Flask(__name__)
|
||||||
|
app.config['JSON_AS_ASCII'] = False
|
||||||
|
|
||||||
|
# 全局变量锁
|
||||||
|
model_lock = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
class Predictor:
|
||||||
|
"""预测器类,用于管理模型和分词器的生命周期"""
|
||||||
|
_instance = None
|
||||||
|
|
||||||
|
def __new__(cls):
|
||||||
|
if cls._instance is None:
|
||||||
|
with model_lock:
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = super().__new__(cls)
|
||||||
|
cls._instance._initialized = False
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
if self._initialized:
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.info(f"Initializing Predictor on device: {DEVICE}")
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 检查模型目录是否存在
|
||||||
|
if not os.path.exists(MODEL_DIR):
|
||||||
|
raise FileNotFoundError(f"Model directory not found: {MODEL_DIR}")
|
||||||
|
|
||||||
|
# 检查label_encoder是否存在
|
||||||
|
if not os.path.exists(LABEL_ENCODER_PATH):
|
||||||
|
raise FileNotFoundError(f"Label encoder not found: {LABEL_ENCODER_PATH}")
|
||||||
|
|
||||||
|
# 加载模型和分词器
|
||||||
|
logger.info(f"Loading tokenizer and model from {MODEL_DIR}...")
|
||||||
|
self.tokenizer = BertTokenizer.from_pretrained(MODEL_DIR)
|
||||||
|
self.model = BertForSequenceClassification.from_pretrained(MODEL_DIR).to(DEVICE)
|
||||||
|
self.model.eval()
|
||||||
|
logger.info("Model loaded successfully")
|
||||||
|
|
||||||
|
# 加载标签映射器
|
||||||
|
logger.info(f"Loading label encoder from {LABEL_ENCODER_PATH}...")
|
||||||
|
self.label_encoder = joblib.load(LABEL_ENCODER_PATH)
|
||||||
|
self.label_map = {i: label for i, label in enumerate(self.label_encoder.classes_)}
|
||||||
|
logger.info(f"Loaded {len(self.label_map)} labels: {list(self.label_map.values())}")
|
||||||
|
|
||||||
|
# 创建线程池
|
||||||
|
self.executor = ThreadPoolExecutor(max_workers=MAX_WORKERS)
|
||||||
|
logger.info(f"ThreadPoolExecutor created with {MAX_WORKERS} workers")
|
||||||
|
|
||||||
|
# 注册关闭钩子
|
||||||
|
atexit.register(self._shutdown)
|
||||||
|
|
||||||
|
# 模型预热
|
||||||
|
self._warmup()
|
||||||
|
|
||||||
|
self._initialized = True
|
||||||
|
elapsed_time = time.time() - start_time
|
||||||
|
logger.info(f"Predictor initialized successfully in {elapsed_time:.2f}s")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to initialize Predictor: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
def _shutdown(self):
|
||||||
|
"""清理资源"""
|
||||||
|
logger.info("Shutting down executor...")
|
||||||
|
self.executor.shutdown(wait=True)
|
||||||
|
logger.info("Executor shutdown complete")
|
||||||
|
|
||||||
|
def _warmup(self):
|
||||||
|
"""模型预热,确保CUDA初始化"""
|
||||||
|
logger.info("Warming up model...")
|
||||||
|
try:
|
||||||
|
test_text = "测试地址"
|
||||||
|
self.predict_single(test_text)
|
||||||
|
logger.info("Model warmup completed")
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Model warmup failed: {e}")
|
||||||
|
|
||||||
|
@lru_cache(maxsize=CACHE_SIZE)
|
||||||
|
def clean_text(self, text: str) -> str:
|
||||||
|
"""文本清洗函数,带缓存"""
|
||||||
|
if not isinstance(text, str):
|
||||||
|
return ""
|
||||||
|
# 只保留中文字符
|
||||||
|
cleaned_text = re.sub(r'[^\u4e00-\u9fa5]', '', text)
|
||||||
|
return cleaned_text.strip()
|
||||||
|
|
||||||
|
def predict_single(self, text: str) -> Dict:
|
||||||
|
"""单个文本预测"""
|
||||||
|
if not text or not isinstance(text, str):
|
||||||
|
return {"address": "", "confidence": 0.0, "error": "Invalid input"}
|
||||||
|
|
||||||
|
try:
|
||||||
|
cleaned_text = self.clean_text(text)
|
||||||
|
if not cleaned_text:
|
||||||
|
logger.warning(f"Empty text after cleaning: {text[:50]}...")
|
||||||
|
return {"address": "", "confidence": 0.0}
|
||||||
|
|
||||||
|
inputs = self.tokenizer(
|
||||||
|
cleaned_text,
|
||||||
|
return_tensors="pt",
|
||||||
|
truncation=True,
|
||||||
|
padding=True,
|
||||||
|
max_length=MAX_LENGTH
|
||||||
|
).to(DEVICE)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = self.model(**inputs)
|
||||||
|
probs = torch.softmax(outputs.logits, dim=1).cpu()
|
||||||
|
top_prob, top_idx = torch.topk(probs, k=1)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"address": self.label_map[top_idx.item()],
|
||||||
|
"confidence": round(top_prob.item(), 4)
|
||||||
|
}
|
||||||
|
|
||||||
|
except torch.cuda.OutOfMemoryError as e:
|
||||||
|
logger.error(f"CUDA OOM error: {e}")
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
return {"address": "", "confidence": 0.0, "error": "GPU memory exhausted"}
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Prediction error for text '{text[:50]}...': {e}")
|
||||||
|
return {"address": "", "confidence": 0.0, "error": str(e)}
|
||||||
|
|
||||||
|
def batch_predict(self, texts: List[str], batch_size: int = 32) -> List[Dict]:
|
||||||
|
"""按 batch 进行单次前向,避免为每个文本并发调用模型导致内存/GPU竞争。"""
|
||||||
|
if not texts:
|
||||||
|
return []
|
||||||
|
|
||||||
|
logger.info(f"Processing {len(texts)} texts with batch_size={batch_size}")
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
results = []
|
||||||
|
total_batches = (len(texts) + batch_size - 1) // batch_size
|
||||||
|
|
||||||
|
for batch_idx in range(total_batches):
|
||||||
|
batch_start = batch_idx * batch_size
|
||||||
|
batch_end = min(batch_start + batch_size, len(texts))
|
||||||
|
batch_texts = texts[batch_start:batch_end]
|
||||||
|
|
||||||
|
# Tokenize batch on CPU then move to device
|
||||||
|
inputs = self.tokenizer(
|
||||||
|
batch_texts,
|
||||||
|
return_tensors='pt',
|
||||||
|
truncation=True,
|
||||||
|
padding=True,
|
||||||
|
max_length=MAX_LENGTH
|
||||||
|
)
|
||||||
|
inputs = {k: v.to(DEVICE) for k, v in inputs.items()}
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = self.model(**inputs)
|
||||||
|
probs = torch.softmax(outputs.logits, dim=1).cpu()
|
||||||
|
top_prob, top_idx = torch.topk(probs, k=1)
|
||||||
|
|
||||||
|
for i in range(len(batch_texts)):
|
||||||
|
results.append({
|
||||||
|
"address": self.label_map[top_idx[i].item()],
|
||||||
|
"confidence": round(top_prob[i].item(), 4)
|
||||||
|
})
|
||||||
|
|
||||||
|
if batch_idx % 10 == 0 and batch_idx > 0:
|
||||||
|
logger.info(f"Processed {batch_end}/{len(texts)} texts")
|
||||||
|
|
||||||
|
elapsed_time = time.time() - start_time
|
||||||
|
logger.info(f"Batch prediction completed in {elapsed_time:.2f}s for {len(texts)} texts")
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict:
|
||||||
|
"""获取预测器统计信息"""
|
||||||
|
return {
|
||||||
|
"service_name": SERVICE_NAME,
|
||||||
|
"model_id": MODEL_ID,
|
||||||
|
"model_dir": MODEL_DIR,
|
||||||
|
"device": DEVICE,
|
||||||
|
"num_labels": len(self.label_map),
|
||||||
|
"labels": list(self.label_map.values()),
|
||||||
|
"max_length": MAX_LENGTH,
|
||||||
|
"max_workers": MAX_WORKERS
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# 初始化预测器
|
||||||
|
try:
|
||||||
|
predictor = Predictor()
|
||||||
|
logger.info("Predictor instance created successfully")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to create predictor instance: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/predict', methods=['POST'])
|
||||||
|
def predict():
|
||||||
|
"""预测接口"""
|
||||||
|
try:
|
||||||
|
data = request.get_json()
|
||||||
|
if not data:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Invalid request, JSON body required"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
if 'address' not in data:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Invalid request, 'address' is required"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
addresses = data['address']
|
||||||
|
if not isinstance(addresses, list):
|
||||||
|
addresses = [addresses]
|
||||||
|
|
||||||
|
# 限制最大处理数量
|
||||||
|
max_items = data.get('max_items', 1000)
|
||||||
|
if len(addresses) > max_items:
|
||||||
|
logger.warning(f"Too many items: {len(addresses)} > {max_items}, truncating")
|
||||||
|
addresses = addresses[:max_items]
|
||||||
|
|
||||||
|
# 去重(可选)
|
||||||
|
if data.get('deduplicate', False):
|
||||||
|
addresses = list(dict.fromkeys(addresses))
|
||||||
|
logger.info(f"After deduplication: {len(addresses)} unique texts")
|
||||||
|
|
||||||
|
# 批量预测
|
||||||
|
results = predictor.batch_predict(addresses)
|
||||||
|
|
||||||
|
return jsonify({
|
||||||
|
"status": "success",
|
||||||
|
"predictions": results,
|
||||||
|
"metadata": {
|
||||||
|
"total": len(results),
|
||||||
|
"timestamp": datetime.now().isoformat(),
|
||||||
|
"service": SERVICE_NAME,
|
||||||
|
"device": DEVICE
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Prediction endpoint error: {e}", exc_info=True)
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": f"Internal server error: {str(e)}"
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/health', methods=['GET'])
|
||||||
|
def health_check():
|
||||||
|
"""健康检查接口"""
|
||||||
|
try:
|
||||||
|
stats = predictor.get_stats()
|
||||||
|
return jsonify({
|
||||||
|
"status": "healthy",
|
||||||
|
"timestamp": datetime.now().isoformat(),
|
||||||
|
"stats": stats
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
return jsonify({
|
||||||
|
"status": "unhealthy",
|
||||||
|
"error": str(e)
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/stats', methods=['GET'])
|
||||||
|
def get_stats():
|
||||||
|
"""获取预测器统计信息"""
|
||||||
|
try:
|
||||||
|
stats = predictor.get_stats()
|
||||||
|
return jsonify({
|
||||||
|
"status": "success",
|
||||||
|
"stats": stats
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": str(e)
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.errorhandler(404)
|
||||||
|
def not_found(error):
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Endpoint not found"
|
||||||
|
}), 404
|
||||||
|
|
||||||
|
|
||||||
|
@app.errorhandler(500)
|
||||||
|
def internal_error(error):
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Internal server error"
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
logger.info(f"Starting Flask application: {SERVICE_NAME}")
|
||||||
|
logger.info(f"Model path: {MODEL_DIR}")
|
||||||
|
logger.info(f"Device: {DEVICE}")
|
||||||
|
logger.info(f"Port: {SERVICE_PORT}")
|
||||||
|
|
||||||
|
app.run(
|
||||||
|
host='0.0.0.0',
|
||||||
|
port=SERVICE_PORT,
|
||||||
|
threaded=True,
|
||||||
|
debug=False
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
# 兼容 WSGI 标准(如 Gunicorn)
|
||||||
|
application = app
|
||||||
|
logger.info(f"Application loaded for WSGI server: {SERVICE_NAME}")
|
||||||
|
|
@ -52,6 +52,11 @@ logging.basicConfig(
|
||||||
)
|
)
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# 如果使用 GPU,限制并发线程数为 1 以避免并发 GPU 推理导致 OOM
|
||||||
|
if isinstance(DEVICE, str) and DEVICE.startswith('cuda') and MAX_WORKERS > 1:
|
||||||
|
logger.warning("CUDA in use — limiting MAX_WORKERS to 1 to avoid concurrent GPU inference")
|
||||||
|
MAX_WORKERS = 1
|
||||||
|
|
||||||
# 打印启动信息
|
# 打印启动信息
|
||||||
logger.info("=" * 50)
|
logger.info("=" * 50)
|
||||||
logger.info(f"Starting service: {SERVICE_NAME}")
|
logger.info(f"Starting service: {SERVICE_NAME}")
|
||||||
|
|
@ -165,9 +170,10 @@ class Predictor:
|
||||||
cleaned_text = re.sub(r'[^\u4e00-\u9fa5]', '', text)
|
cleaned_text = re.sub(r'[^\u4e00-\u9fa5]', '', text)
|
||||||
return cleaned_text.strip()
|
return cleaned_text.strip()
|
||||||
|
|
||||||
@lru_cache(maxsize=TOKEN_CACHE_SIZE)
|
|
||||||
def tokenize_text(self, text: str):
|
def tokenize_text(self, text: str):
|
||||||
"""缓存tokenization结果"""
|
"""Tokenize input text (no large-cache to avoid memory growth).
|
||||||
|
Returns tokenizer output on CPU.
|
||||||
|
"""
|
||||||
if not text:
|
if not text:
|
||||||
return None
|
return None
|
||||||
return self.tokenizer(
|
return self.tokenizer(
|
||||||
|
|
@ -234,7 +240,7 @@ class Predictor:
|
||||||
return {"type": "", "tax": "", "confidence": 0.0, "error": str(e)}
|
return {"type": "", "tax": "", "confidence": 0.0, "error": str(e)}
|
||||||
|
|
||||||
def batch_predict(self, texts: List[str], batch_size: int = BATCH_SIZE) -> List[Dict]:
|
def batch_predict(self, texts: List[str], batch_size: int = BATCH_SIZE) -> List[Dict]:
|
||||||
"""批量预测(并发处理)"""
|
"""批量预测(按 batch 进行单次前向,避免为每个文本并发调用模型导致内存/GPU竞争)。"""
|
||||||
if not texts:
|
if not texts:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
@ -249,15 +255,31 @@ class Predictor:
|
||||||
batch_end = min(batch_start + batch_size, len(texts))
|
batch_end = min(batch_start + batch_size, len(texts))
|
||||||
batch_texts = texts[batch_start:batch_end]
|
batch_texts = texts[batch_start:batch_end]
|
||||||
|
|
||||||
# 提交批次任务
|
# Tokenize whole batch on CPU, then move tensors to device
|
||||||
futures = [
|
inputs = self.tokenizer(
|
||||||
self.executor.submit(self.predict_single, text)
|
batch_texts,
|
||||||
for text in batch_texts
|
return_tensors='pt',
|
||||||
]
|
truncation=True,
|
||||||
|
padding=True,
|
||||||
|
max_length=MAX_LENGTH
|
||||||
|
)
|
||||||
|
inputs = {k: v.to(DEVICE) for k, v in inputs.items()}
|
||||||
|
|
||||||
# 收集结果
|
with torch.no_grad():
|
||||||
batch_results = [future.result() for future in futures]
|
outputs = self.model(**inputs)
|
||||||
results.extend(batch_results)
|
probs = torch.softmax(outputs.logits, dim=1).cpu()
|
||||||
|
top_prob, top_idx = torch.topk(probs, k=1)
|
||||||
|
|
||||||
|
for i in range(len(batch_texts)):
|
||||||
|
label_idx = top_idx[i].item()
|
||||||
|
tax_label = self.label_encoder.inverse_transform([label_idx])[0]
|
||||||
|
confidence = round(top_prob[i].item(), 4)
|
||||||
|
tax_arr = tax_label.split("_")
|
||||||
|
if len(tax_arr) >= 2:
|
||||||
|
res = {"type": tax_arr[0], "tax": tax_arr[1], "confidence": confidence}
|
||||||
|
else:
|
||||||
|
res = {"type": tax_label, "tax": "", "confidence": confidence}
|
||||||
|
results.append(res)
|
||||||
|
|
||||||
if batch_idx % 10 == 0 and batch_idx > 0:
|
if batch_idx % 10 == 0 and batch_idx > 0:
|
||||||
logger.info(f"Processed {batch_end}/{len(texts)} texts")
|
logger.info(f"Processed {batch_end}/{len(texts)} texts")
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,506 @@
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import warnings
|
||||||
|
from flask import Flask, request, jsonify
|
||||||
|
import torch
|
||||||
|
from transformers import BertTokenizer, BertForSequenceClassification
|
||||||
|
import joblib
|
||||||
|
import re
|
||||||
|
from functools import lru_cache
|
||||||
|
from typing import List, Dict
|
||||||
|
import threading
|
||||||
|
import logging
|
||||||
|
import atexit
|
||||||
|
from datetime import datetime
|
||||||
|
import time
|
||||||
|
|
||||||
|
# 忽略 scikit-learn 版本警告
|
||||||
|
warnings.filterwarnings("ignore", category=UserWarning, module="sklearn")
|
||||||
|
|
||||||
|
# ==================== 从环境变量读取配置 ====================
|
||||||
|
SERVICE_PORT = int(os.environ.get('SERVICE_PORT', 5003))
|
||||||
|
MAX_WORKERS = int(os.environ.get('MAX_WORKERS', 2))
|
||||||
|
CACHE_SIZE = 2000
|
||||||
|
TOKEN_CACHE_SIZE = 1000
|
||||||
|
|
||||||
|
# ==================== 设备配置 ====================
|
||||||
|
MAX_LENGTH = 512
|
||||||
|
DEVICE = os.environ.get('DEVICE', 'cuda' if torch.cuda.is_available() else 'cpu')
|
||||||
|
if DEVICE == 'cuda' and not torch.cuda.is_available():
|
||||||
|
DEVICE = 'cpu'
|
||||||
|
print(f"⚠️ CUDA not available, falling back to {DEVICE}")
|
||||||
|
|
||||||
|
# ==================== 日志配置 ====================
|
||||||
|
logging.basicConfig(
|
||||||
|
level=logging.INFO,
|
||||||
|
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||||
|
)
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# 如果使用 GPU,限制并发线程数为 1 以避免并发 GPU 推理导致 OOM
|
||||||
|
if isinstance(DEVICE, str) and DEVICE.startswith('cuda') and MAX_WORKERS > 1:
|
||||||
|
logger.warning("CUDA in use — limiting MAX_WORKERS to 1 to avoid concurrent GPU inference")
|
||||||
|
MAX_WORKERS = 1
|
||||||
|
|
||||||
|
app = Flask(__name__)
|
||||||
|
app.config['JSON_AS_ASCII'] = False
|
||||||
|
|
||||||
|
# 全局变量锁
|
||||||
|
model_lock = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
def discover_models():
|
||||||
|
"""
|
||||||
|
从环境变量中发现所有模型配置
|
||||||
|
约定:每个模型需要以下环境变量:
|
||||||
|
- {PREFIX}_MODEL_ID: 模型ID
|
||||||
|
- {PREFIX}_MODEL_DIR: 模型目录
|
||||||
|
- {PREFIX}_SERVICE_NAME: 服务名称(用于日志和标识)
|
||||||
|
- {PREFIX}_TYPE: 模型类型(tax/address/other),用于决定输出格式
|
||||||
|
"""
|
||||||
|
models = {}
|
||||||
|
for key, value in os.environ.items():
|
||||||
|
if key.endswith('_MODEL_ID'):
|
||||||
|
prefix = key[:-9] # 去掉 _MODEL_ID
|
||||||
|
service_name = os.getenv(f'{prefix}_SERVICE_NAME', f'{prefix.lower()}_classifier')
|
||||||
|
model_type = os.getenv(f'{prefix}_TYPE', 'general')
|
||||||
|
|
||||||
|
models[prefix.lower()] = {
|
||||||
|
'prefix': prefix,
|
||||||
|
'model_id': value,
|
||||||
|
'model_dir': os.getenv(f'{prefix}_MODEL_DIR', f'/app/services/{prefix.lower()}/model'),
|
||||||
|
'service_name': service_name,
|
||||||
|
'type': model_type, # tax, address, general
|
||||||
|
}
|
||||||
|
return models
|
||||||
|
|
||||||
|
|
||||||
|
class BasePredictor:
|
||||||
|
"""基础预测器类 - 每个模型实例"""
|
||||||
|
|
||||||
|
def __init__(self, config):
|
||||||
|
self.config = config
|
||||||
|
self.model_type = config['type']
|
||||||
|
self.model_dir = config['model_dir']
|
||||||
|
self.service_name = config['service_name']
|
||||||
|
self.model_id = config['model_id']
|
||||||
|
self._initialized = False
|
||||||
|
self._init_model()
|
||||||
|
|
||||||
|
def _find_label_encoder(self):
|
||||||
|
"""自动查找 label_encoder 文件"""
|
||||||
|
# 尝试常见的文件名
|
||||||
|
possible_names = [
|
||||||
|
'label_encoder.pkl',
|
||||||
|
'label_encoder_roberta_large.pkl',
|
||||||
|
'label_encoder_roberta.pkl',
|
||||||
|
'label_encoder_bert.pkl'
|
||||||
|
]
|
||||||
|
for name in possible_names:
|
||||||
|
path = os.path.join(self.model_dir, name)
|
||||||
|
if os.path.exists(path):
|
||||||
|
return path
|
||||||
|
|
||||||
|
# 尝试通配符查找
|
||||||
|
import glob
|
||||||
|
encoder_files = glob.glob(os.path.join(self.model_dir, 'label_encoder*.pkl'))
|
||||||
|
if encoder_files:
|
||||||
|
return encoder_files[0]
|
||||||
|
|
||||||
|
raise FileNotFoundError(f"No label_encoder found in {self.model_dir}")
|
||||||
|
|
||||||
|
def _init_model(self):
|
||||||
|
"""初始化模型"""
|
||||||
|
logger.info(f"Loading model '{self.service_name}' from {self.model_dir}...")
|
||||||
|
try:
|
||||||
|
if not os.path.exists(self.model_dir):
|
||||||
|
raise FileNotFoundError(f"Model directory not found: {self.model_dir}")
|
||||||
|
|
||||||
|
# 查找 label_encoder
|
||||||
|
label_encoder_path = self._find_label_encoder()
|
||||||
|
logger.info(f"Found label_encoder: {label_encoder_path}")
|
||||||
|
|
||||||
|
# 加载模型和分词器
|
||||||
|
self.tokenizer = BertTokenizer.from_pretrained(self.model_dir)
|
||||||
|
self.model = BertForSequenceClassification.from_pretrained(self.model_dir).to(DEVICE)
|
||||||
|
self.model.eval()
|
||||||
|
|
||||||
|
# 加载标签映射器
|
||||||
|
self.label_encoder = joblib.load(label_encoder_path)
|
||||||
|
self.label_map = {i: label for i, label in enumerate(self.label_encoder.classes_)}
|
||||||
|
self.num_labels = len(self.label_map)
|
||||||
|
|
||||||
|
self._initialized = True
|
||||||
|
logger.info(f"✅ Model '{self.service_name}' loaded, {self.num_labels} labels")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"❌ Failed to load model '{self.service_name}': {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
@lru_cache(maxsize=CACHE_SIZE)
|
||||||
|
def clean_text(self, text: str) -> str:
|
||||||
|
"""文本清洗函数,带缓存"""
|
||||||
|
if not isinstance(text, str):
|
||||||
|
return ""
|
||||||
|
cleaned_text = re.sub(r'[^\u4e00-\u9fa5]', '', text)
|
||||||
|
return cleaned_text.strip()
|
||||||
|
|
||||||
|
def tokenize_text(self, text: str):
|
||||||
|
"""Tokenize input text (no large-cache to avoid memory growth).
|
||||||
|
Returns tokenizer output on CPU.
|
||||||
|
"""
|
||||||
|
if not text:
|
||||||
|
return None
|
||||||
|
return self.tokenizer(
|
||||||
|
text,
|
||||||
|
return_tensors="pt",
|
||||||
|
truncation=True,
|
||||||
|
padding=True,
|
||||||
|
max_length=MAX_LENGTH
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_single(self, text: str) -> Dict:
|
||||||
|
"""单个文本预测"""
|
||||||
|
if not text or not isinstance(text, str):
|
||||||
|
return {"error": "Invalid input"}
|
||||||
|
|
||||||
|
try:
|
||||||
|
cleaned_text = self.clean_text(text)
|
||||||
|
if not cleaned_text:
|
||||||
|
return {"error": "Empty text after cleaning"}
|
||||||
|
|
||||||
|
inputs = self.tokenize_text(cleaned_text)
|
||||||
|
if inputs is None:
|
||||||
|
return {"error": "Tokenization failed"}
|
||||||
|
|
||||||
|
inputs = {k: v.to(DEVICE) for k, v in inputs.items()}
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = self.model(**inputs)
|
||||||
|
probs = torch.softmax(outputs.logits, dim=1).cpu()
|
||||||
|
top_prob, top_idx = torch.topk(probs, k=1)
|
||||||
|
|
||||||
|
label = self.label_map[top_idx.item()]
|
||||||
|
confidence = round(top_prob.item(), 4)
|
||||||
|
|
||||||
|
# 根据模型类型格式化输出
|
||||||
|
return self._format_output(label, confidence)
|
||||||
|
|
||||||
|
except torch.cuda.OutOfMemoryError as e:
|
||||||
|
logger.error(f"CUDA OOM in {self.service_name}: {e}")
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
return {"error": "GPU memory exhausted"}
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Prediction error in {self.service_name}: {e}")
|
||||||
|
return {"error": str(e)}
|
||||||
|
|
||||||
|
def _format_output(self, label: str, confidence: float) -> Dict:
|
||||||
|
"""根据模型类型格式化输出"""
|
||||||
|
result = {"confidence": confidence}
|
||||||
|
|
||||||
|
if self.model_type == 'tax':
|
||||||
|
# tax 模型:type_tax 格式
|
||||||
|
parts = label.split('_')
|
||||||
|
if len(parts) >= 2:
|
||||||
|
result['type'] = parts[0]
|
||||||
|
result['tax'] = '_'.join(parts[1:])
|
||||||
|
else:
|
||||||
|
result['type'] = label
|
||||||
|
result['tax'] = ''
|
||||||
|
elif self.model_type == 'address':
|
||||||
|
result['address'] = label
|
||||||
|
elif self.model_type == 'brand':
|
||||||
|
result['brand'] = label
|
||||||
|
else:
|
||||||
|
# 通用模型:直接输出 label
|
||||||
|
result['label'] = label
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict:
|
||||||
|
"""获取模型统计信息"""
|
||||||
|
return {
|
||||||
|
"service_name": self.service_name,
|
||||||
|
"model_id": self.model_id,
|
||||||
|
"model_type": self.model_type,
|
||||||
|
"model_dir": self.model_dir,
|
||||||
|
"device": DEVICE,
|
||||||
|
"num_labels": self.num_labels,
|
||||||
|
"labels": list(self.label_map.values())[:10] # 只显示前10个
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class ModelManager:
|
||||||
|
"""模型管理器 - 管理所有模型实例"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.predictors = {}
|
||||||
|
self._load_all_models()
|
||||||
|
|
||||||
|
def _load_all_models(self):
|
||||||
|
"""加载所有模型"""
|
||||||
|
logger.info("=" * 50)
|
||||||
|
logger.info("Initializing Model Manager...")
|
||||||
|
logger.info(f"Device: {DEVICE}")
|
||||||
|
logger.info("=" * 50)
|
||||||
|
|
||||||
|
models_config = discover_models()
|
||||||
|
|
||||||
|
if not models_config:
|
||||||
|
logger.warning("⚠️ No models configured in .env")
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.info(f"Found {len(models_config)} model(s):")
|
||||||
|
for name, config in models_config.items():
|
||||||
|
logger.info(f" - {name}: {config['model_id']} (type: {config['type']})")
|
||||||
|
|
||||||
|
logger.info("-" * 50)
|
||||||
|
|
||||||
|
for name, config in models_config.items():
|
||||||
|
try:
|
||||||
|
self.predictors[name] = BasePredictor(config)
|
||||||
|
logger.info(f"✅ {name} loaded successfully")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"❌ Failed to load {name}: {e}")
|
||||||
|
|
||||||
|
logger.info("=" * 50)
|
||||||
|
logger.info(f"Loaded {len(self.predictors)}/{len(models_config)} models")
|
||||||
|
logger.info(f"Available: {list(self.predictors.keys())}")
|
||||||
|
logger.info("=" * 50)
|
||||||
|
|
||||||
|
def get_predictor(self, name: str):
|
||||||
|
"""获取指定名称的预测器。
|
||||||
|
尝试多种匹配:精确 key,大小写无关,service_name 或 model_id 匹配。
|
||||||
|
返回匹配的 Predictor 或 None。
|
||||||
|
"""
|
||||||
|
if not name:
|
||||||
|
return None
|
||||||
|
# 直接按 key 查找(优先)
|
||||||
|
if name in self.predictors:
|
||||||
|
logger.debug(f"Predictor matched by key: {name}")
|
||||||
|
return self.predictors[name]
|
||||||
|
# 尝试大小写不敏感的 key
|
||||||
|
lower_name = name.lower()
|
||||||
|
for k in self.predictors.keys():
|
||||||
|
if k.lower() == lower_name:
|
||||||
|
logger.debug(f"Predictor matched by case-insensitive key: {k} for request '{name}'")
|
||||||
|
return self.predictors[k]
|
||||||
|
# 尝试按 predictor 的 service_name 或 model_id 匹配
|
||||||
|
for k, predictor in self.predictors.items():
|
||||||
|
try:
|
||||||
|
if getattr(predictor, 'service_name', '').lower() == lower_name:
|
||||||
|
logger.debug(f"Predictor matched by service_name: {k} -> {predictor.service_name}")
|
||||||
|
return predictor
|
||||||
|
if getattr(predictor, 'model_id', '').lower() == lower_name:
|
||||||
|
logger.debug(f"Predictor matched by model_id: {k} -> {predictor.model_id}")
|
||||||
|
return predictor
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
# 不再做模糊/子串匹配以避免错误映射;只做严格或基于 service_name/model_id 的匹配
|
||||||
|
return None
|
||||||
|
|
||||||
|
def list_models(self) -> List[str]:
|
||||||
|
"""列出所有可用模型"""
|
||||||
|
return list(self.predictors.keys())
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict:
|
||||||
|
"""获取所有模型的统计信息"""
|
||||||
|
return {
|
||||||
|
name: predictor.get_stats()
|
||||||
|
for name, predictor in self.predictors.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== 初始化模型管理器 ====================
|
||||||
|
model_manager = ModelManager()
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/predict', methods=['POST'])
|
||||||
|
def predict():
|
||||||
|
"""统一预测接口"""
|
||||||
|
try:
|
||||||
|
# 检查 Content-Type
|
||||||
|
if not request.is_json:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Content-Type must be application/json"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
# 解析请求
|
||||||
|
data = request.get_json(silent=True)
|
||||||
|
if not data:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Invalid JSON body"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
# 检查必填字段
|
||||||
|
if 'model' not in data:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": f"Missing required field: 'model'. Available: {model_manager.list_models()}"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
if 'text' not in data:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Missing required field: 'text'"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
model_name = data['model']
|
||||||
|
text = data['text']
|
||||||
|
|
||||||
|
# 验证模型是否存在
|
||||||
|
predictor = model_manager.get_predictor(model_name)
|
||||||
|
if predictor is None:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": f"Invalid model: {model_name}. Available: {model_manager.list_models()}"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
# 验证文本
|
||||||
|
if not isinstance(text, str):
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Invalid 'text' field, must be a string"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
if not text.strip():
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Empty text"
|
||||||
|
}), 400
|
||||||
|
|
||||||
|
# 记录并返回所使用的 predictor 信息(仅当请求中包含 debug=True 时会把详细信息返给客户端)
|
||||||
|
predictor_info = {
|
||||||
|
"key": None,
|
||||||
|
"service_name": getattr(predictor, 'service_name', None),
|
||||||
|
"model_id": getattr(predictor, 'model_id', None),
|
||||||
|
"model_dir": getattr(predictor, 'model_dir', None)
|
||||||
|
}
|
||||||
|
# 尝试找出 predictor 的注册 key
|
||||||
|
for k, p in model_manager.predictors.items():
|
||||||
|
if p is predictor:
|
||||||
|
predictor_info['key'] = k
|
||||||
|
break
|
||||||
|
|
||||||
|
# 进行预测
|
||||||
|
result = predictor.predict_single(text)
|
||||||
|
|
||||||
|
response = {
|
||||||
|
"status": "success",
|
||||||
|
"model": model_name,
|
||||||
|
"prediction": result,
|
||||||
|
"metadata": {
|
||||||
|
"timestamp": datetime.now().isoformat(),
|
||||||
|
"device": DEVICE
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
# 如果请求开启 debug,则在响应中包含 predictor 信息
|
||||||
|
if data.get('debug', False):
|
||||||
|
response['predictor'] = predictor_info
|
||||||
|
|
||||||
|
logger.info(f"Request for model='{model_name}' routed to predictor key='{predictor_info.get('key')}', service_name='{predictor_info.get('service_name')}', model_id='{predictor_info.get('model_id')}'")
|
||||||
|
|
||||||
|
return jsonify(response)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Prediction endpoint error: {e}", exc_info=True)
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": f"Internal server error: {str(e)}"
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/models', methods=['GET'])
|
||||||
|
def list_models():
|
||||||
|
"""列出所有可用的模型"""
|
||||||
|
return jsonify({
|
||||||
|
"status": "success",
|
||||||
|
"models": model_manager.list_models(),
|
||||||
|
"details": model_manager.get_stats()
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/debug_models', methods=['GET'])
|
||||||
|
def debug_models():
|
||||||
|
"""返回 model_manager 的映射信息,便于诊断模型与 key 的对应关系(仅内部/运维使用)。"""
|
||||||
|
try:
|
||||||
|
data = {}
|
||||||
|
for k, p in model_manager.predictors.items():
|
||||||
|
data[k] = {
|
||||||
|
'service_name': getattr(p, 'service_name', None),
|
||||||
|
'model_id': getattr(p, 'model_id', None),
|
||||||
|
'model_dir': getattr(p, 'model_dir', None),
|
||||||
|
}
|
||||||
|
return jsonify({'status': 'success', 'mappings': data})
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"debug_models error: {e}")
|
||||||
|
return jsonify({'status': 'error', 'error': str(e)}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/health', methods=['GET'])
|
||||||
|
def health_check():
|
||||||
|
"""健康检查接口"""
|
||||||
|
try:
|
||||||
|
return jsonify({
|
||||||
|
"status": "healthy",
|
||||||
|
"timestamp": datetime.now().isoformat(),
|
||||||
|
"device": DEVICE,
|
||||||
|
"models": model_manager.list_models(),
|
||||||
|
"loaded": len(model_manager.predictors)
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
return jsonify({
|
||||||
|
"status": "unhealthy",
|
||||||
|
"error": str(e)
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.route('/stats', methods=['GET'])
|
||||||
|
def get_stats():
|
||||||
|
"""获取统计信息"""
|
||||||
|
try:
|
||||||
|
return jsonify({
|
||||||
|
"status": "success",
|
||||||
|
"device": DEVICE,
|
||||||
|
"models": model_manager.get_stats()
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": str(e)
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
@app.errorhandler(404)
|
||||||
|
def not_found(error):
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Endpoint not found"
|
||||||
|
}), 404
|
||||||
|
|
||||||
|
|
||||||
|
@app.errorhandler(500)
|
||||||
|
def internal_error(error):
|
||||||
|
return jsonify({
|
||||||
|
"status": "error",
|
||||||
|
"error": "Internal server error"
|
||||||
|
}), 500
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
logger.info(f"Starting unified Flask application on port {SERVICE_PORT}")
|
||||||
|
logger.info(f"Loaded models: {model_manager.list_models()}")
|
||||||
|
logger.info(f"Device: {DEVICE}")
|
||||||
|
|
||||||
|
app.run(
|
||||||
|
host='0.0.0.0',
|
||||||
|
port=SERVICE_PORT,
|
||||||
|
threaded=True,
|
||||||
|
debug=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
application = app
|
||||||
|
logger.info(f"Application loaded for WSGI server with models: {model_manager.list_models()}")
|
||||||
|
|
@ -0,0 +1,301 @@
|
||||||
|
#!/usr/bin/env python3
|
||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
|
||||||
|
import pandas as pd
|
||||||
|
import sys
|
||||||
|
from difflib import SequenceMatcher
|
||||||
|
from fuzzywuzzy import fuzz
|
||||||
|
from tqdm import tqdm
|
||||||
|
import hashlib
|
||||||
|
|
||||||
|
|
||||||
|
def calculate_similarity(s1, s2, method='fuzzy'):
|
||||||
|
"""计算两个字符串的相似度"""
|
||||||
|
if not s1 or not s2:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
s1 = str(s1).strip()
|
||||||
|
s2 = str(s2).strip()
|
||||||
|
|
||||||
|
if s1 == s2:
|
||||||
|
return 1.0
|
||||||
|
|
||||||
|
if method == 'fuzzy':
|
||||||
|
return fuzz.ratio(s1, s2) / 100.0
|
||||||
|
elif method == 'partial':
|
||||||
|
return fuzz.partial_ratio(s1, s2) / 100.0
|
||||||
|
elif method == 'sequence':
|
||||||
|
return SequenceMatcher(None, s1, s2).ratio()
|
||||||
|
else:
|
||||||
|
return fuzz.ratio(s1, s2) / 100.0
|
||||||
|
|
||||||
|
|
||||||
|
def extract_keyword(text):
|
||||||
|
"""
|
||||||
|
提取关键词:只用前3个字符
|
||||||
|
"""
|
||||||
|
if not text:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
text = str(text).strip()
|
||||||
|
if not text:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
# 取前3个字符,如果长度小于3则取全部
|
||||||
|
return text[:3]
|
||||||
|
|
||||||
|
|
||||||
|
def get_text_hash(text):
|
||||||
|
"""获取文本的哈希值,用于精确去重"""
|
||||||
|
return hashlib.md5(str(text).encode('utf-8')).hexdigest()[:8]
|
||||||
|
|
||||||
|
|
||||||
|
class SafeDedup:
|
||||||
|
"""安全的去重器"""
|
||||||
|
|
||||||
|
def __init__(self, threshold=0.8, method='fuzzy', max_cache_size=50000):
|
||||||
|
self.threshold = threshold
|
||||||
|
self.method = method
|
||||||
|
self.max_cache_size = max_cache_size
|
||||||
|
self.similarity_cache = {}
|
||||||
|
self.cache_hits = 0
|
||||||
|
self.cache_misses = 0
|
||||||
|
|
||||||
|
def get_from_cache(self, key):
|
||||||
|
if key in self.similarity_cache:
|
||||||
|
self.cache_hits += 1
|
||||||
|
return self.similarity_cache[key]
|
||||||
|
self.cache_misses += 1
|
||||||
|
return None
|
||||||
|
|
||||||
|
def add_to_cache(self, key, value):
|
||||||
|
if len(self.similarity_cache) >= self.max_cache_size:
|
||||||
|
items = list(self.similarity_cache.items())
|
||||||
|
self.similarity_cache = dict(items[len(items) // 2:])
|
||||||
|
self.similarity_cache[key] = value
|
||||||
|
|
||||||
|
def is_similar(self, text1, text2):
|
||||||
|
"""判断两个文本是否相似"""
|
||||||
|
if not text1 or not text2:
|
||||||
|
return False
|
||||||
|
|
||||||
|
text1 = str(text1).strip()
|
||||||
|
text2 = str(text2).strip()
|
||||||
|
|
||||||
|
if text1 == text2:
|
||||||
|
return True
|
||||||
|
|
||||||
|
cache_key = tuple(sorted([text1, text2]))
|
||||||
|
cached_result = self.get_from_cache(cache_key)
|
||||||
|
if cached_result is not None:
|
||||||
|
return cached_result
|
||||||
|
|
||||||
|
# 快速预筛选:长度差异检查
|
||||||
|
len_diff = abs(len(text1) - len(text2))
|
||||||
|
max_len = max(len(text1), len(text2))
|
||||||
|
if max_len > 0 and len_diff / max_len > 0.6:
|
||||||
|
self.add_to_cache(cache_key, False)
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 计算相似度
|
||||||
|
similarity = calculate_similarity(text1, text2, self.method)
|
||||||
|
result = similarity >= self.threshold
|
||||||
|
|
||||||
|
self.add_to_cache(cache_key, result)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def dedup_by_similarity(input_file='input.xlsx', output_file='output.xlsx',
|
||||||
|
key_column='sentence', threshold=0.8, method='fuzzy'):
|
||||||
|
"""
|
||||||
|
基于相似度去重
|
||||||
|
使用前3个字符作为关键词分组
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
# 1. 读取数据
|
||||||
|
print(f"📂 读取文件: {input_file}")
|
||||||
|
df = pd.read_excel(input_file, engine='openpyxl')
|
||||||
|
|
||||||
|
print(f"📋 列名: {list(df.columns)}")
|
||||||
|
print(f"🔑 判重列: {key_column}")
|
||||||
|
print(f"📊 相似度阈值: {threshold}")
|
||||||
|
print(f"🔧 相似度方法: {method}")
|
||||||
|
|
||||||
|
if key_column not in df.columns:
|
||||||
|
print(f"❌ 错误: 找不到列 '{key_column}'")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
# 2. 预处理数据
|
||||||
|
df = df.dropna(subset=[key_column])
|
||||||
|
df[key_column] = df[key_column].astype(str).str.strip()
|
||||||
|
df = df[df[key_column] != '']
|
||||||
|
|
||||||
|
total_rows = len(df)
|
||||||
|
print(f"\n📊 有效数据: {total_rows} 行")
|
||||||
|
|
||||||
|
# 3. 精确去重
|
||||||
|
print(f"\n🔍 步骤1: 精确去重...")
|
||||||
|
df['text_hash'] = df[key_column].apply(get_text_hash)
|
||||||
|
df = df.drop_duplicates(subset=['text_hash'], keep='first')
|
||||||
|
df = df.drop(columns=['text_hash'])
|
||||||
|
|
||||||
|
after_exact_dedup = len(df)
|
||||||
|
print(f" 精确去重后: {after_exact_dedup} 行 (减少 {total_rows - after_exact_dedup} 行)")
|
||||||
|
|
||||||
|
# 4. 关键词分组:只用前3个字符
|
||||||
|
print(f"\n🔍 步骤2: 关键词分组 (前3个字符)...")
|
||||||
|
df['keyword'] = df[key_column].apply(extract_keyword)
|
||||||
|
|
||||||
|
grouped = df.groupby('keyword')
|
||||||
|
group_count = len(grouped)
|
||||||
|
print(f" 分为 {group_count} 个组")
|
||||||
|
print(f" 平均每组 {len(df) / group_count:.1f} 条")
|
||||||
|
|
||||||
|
# 显示分组情况
|
||||||
|
print(f"\n📊 分组统计 (前10组):")
|
||||||
|
for i, (keyword, group) in enumerate(grouped):
|
||||||
|
if i >= 10:
|
||||||
|
break
|
||||||
|
print(f" 关键词 '{keyword}': {len(group)} 条")
|
||||||
|
# 显示该组第一条数据作为示例
|
||||||
|
first_text = group[key_column].iloc[0]
|
||||||
|
print(f" 示例: {first_text[:40]}...")
|
||||||
|
|
||||||
|
# 5. 相似度去重
|
||||||
|
print(f"\n🔄 步骤3: 相似度去重...")
|
||||||
|
|
||||||
|
deduper = SafeDedup(threshold=threshold, method=method, max_cache_size=50000)
|
||||||
|
kept_indices = []
|
||||||
|
kept_texts = []
|
||||||
|
|
||||||
|
pbar = tqdm(total=len(df), desc="去重进度",
|
||||||
|
bar_format='{l_bar}{bar}| {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]')
|
||||||
|
|
||||||
|
skipped_count = 0
|
||||||
|
kept_count = 0
|
||||||
|
|
||||||
|
# 按组处理
|
||||||
|
for keyword, group in grouped:
|
||||||
|
group_texts = group[key_column].tolist()
|
||||||
|
group_indices = group.index.tolist()
|
||||||
|
|
||||||
|
# 组内只有一条,直接保留
|
||||||
|
if len(group_texts) == 1:
|
||||||
|
kept_indices.append(group_indices[0])
|
||||||
|
kept_texts.append(group_texts[0])
|
||||||
|
kept_count += 1
|
||||||
|
pbar.update(1)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 组内去重
|
||||||
|
group_kept = []
|
||||||
|
group_kept_texts = []
|
||||||
|
|
||||||
|
for i, text in enumerate(group_texts):
|
||||||
|
is_duplicate = False
|
||||||
|
|
||||||
|
for kept_text in group_kept_texts:
|
||||||
|
if deduper.is_similar(text, kept_text):
|
||||||
|
is_duplicate = True
|
||||||
|
skipped_count += 1
|
||||||
|
break
|
||||||
|
|
||||||
|
if not is_duplicate:
|
||||||
|
group_kept.append(group_indices[i])
|
||||||
|
group_kept_texts.append(text)
|
||||||
|
kept_count += 1
|
||||||
|
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
kept_indices.extend(group_kept)
|
||||||
|
kept_texts.extend(group_kept_texts)
|
||||||
|
|
||||||
|
pbar.set_postfix({
|
||||||
|
'保留': kept_count,
|
||||||
|
'跳过': skipped_count,
|
||||||
|
'保留率': f'{kept_count / (kept_count + skipped_count) * 100:.1f}%' if (
|
||||||
|
kept_count + skipped_count) > 0 else '0%'
|
||||||
|
})
|
||||||
|
|
||||||
|
pbar.close()
|
||||||
|
|
||||||
|
# 6. 导出结果
|
||||||
|
print(f"\n📝 正在导出结果...")
|
||||||
|
df_result = df.loc[kept_indices].copy()
|
||||||
|
df_result = df_result.drop(columns=['keyword'])
|
||||||
|
df_result.to_excel(output_file, index=False, engine='openpyxl')
|
||||||
|
|
||||||
|
# 7. 统计信息
|
||||||
|
print(f"\n{'=' * 60}")
|
||||||
|
print(f"✅ 相似度去重完成!")
|
||||||
|
print(f"{'=' * 60}")
|
||||||
|
print(f"📊 统计信息:")
|
||||||
|
print(f" 原始行数: {total_rows}")
|
||||||
|
print(f" 精确去重后: {after_exact_dedup}")
|
||||||
|
print(f" 相似度去重后: {len(df_result)}")
|
||||||
|
print(f" 总删除行数: {total_rows - len(df_result)}")
|
||||||
|
print(f" 去重率: {(total_rows - len(df_result)) / total_rows * 100:.1f}%")
|
||||||
|
print(f" 相似度阈值: {threshold}")
|
||||||
|
print(f" 关键词组数: {group_count}")
|
||||||
|
print(f" 缓存命中率: {deduper.cache_hits / (deduper.cache_hits + deduper.cache_misses) * 100:.1f}%"
|
||||||
|
if (deduper.cache_hits + deduper.cache_misses) > 0 else "N/A")
|
||||||
|
print(f"📁 输出文件: {output_file}")
|
||||||
|
print(f"{'=' * 60}")
|
||||||
|
|
||||||
|
# 显示保留的示例
|
||||||
|
if len(kept_texts) > 0:
|
||||||
|
print(f"\n📝 保留的示例 (前5个):")
|
||||||
|
for i in range(min(5, len(kept_texts))):
|
||||||
|
print(f" {i + 1}. {kept_texts[i][:50]}...")
|
||||||
|
|
||||||
|
except FileNotFoundError:
|
||||||
|
print(f"❌ 错误: 找不到文件 '{input_file}'")
|
||||||
|
sys.exit(1)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"❌ 处理失败: {e}")
|
||||||
|
import traceback
|
||||||
|
traceback.print_exc()
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description='基于前3个字符分组的相似度去重',
|
||||||
|
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||||
|
epilog="""
|
||||||
|
说明:
|
||||||
|
- 使用前3个字符作为关键词分组
|
||||||
|
- 只在同组内进行相似度比较
|
||||||
|
- 大幅提升处理速度
|
||||||
|
|
||||||
|
示例:
|
||||||
|
# 使用默认阈值0.8
|
||||||
|
python dedup_similarity.py -i input.xlsx
|
||||||
|
|
||||||
|
# 调整阈值
|
||||||
|
python dedup_similarity.py -i input.xlsx -t 0.75
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument('-i', '--input', default='input.xlsx',
|
||||||
|
help='输入文件路径 (默认: input.xlsx)')
|
||||||
|
parser.add_argument('-o', '--output', default='output.xlsx',
|
||||||
|
help='输出文件路径 (默认: output.xlsx)')
|
||||||
|
parser.add_argument('-c', '--column', default='sentence',
|
||||||
|
help='判重列名 (默认: sentence)')
|
||||||
|
parser.add_argument('-t', '--threshold', type=float, default=0.8,
|
||||||
|
help='相似度阈值 0-1 (默认: 0.8)')
|
||||||
|
parser.add_argument('-m', '--method', default='fuzzy',
|
||||||
|
choices=['fuzzy', 'partial', 'sequence'],
|
||||||
|
help='相似度计算方法 (默认: fuzzy)')
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
dedup_by_similarity(args.input, args.output, args.column,
|
||||||
|
args.threshold, args.method)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Binary file not shown.
|
|
@ -0,0 +1,6 @@
|
||||||
|
from modelscope.pipelines import pipeline
|
||||||
|
from modelscope.utils.constant import Tasks
|
||||||
|
|
||||||
|
ner_pipeline = pipeline(Tasks.named_entity_recognition, 'iic/nlp_raner_named-entity-recognition_chinese-base-ecom-50cls')
|
||||||
|
result = ner_pipeline('苏泊尔(SUPOR)空气炸锅钴蓝色KJ55D705')
|
||||||
|
print(result)
|
||||||
Binary file not shown.
|
|
@ -0,0 +1,327 @@
|
||||||
|
import pandas as pd
|
||||||
|
import sys
|
||||||
|
from modelscope.pipelines import pipeline
|
||||||
|
from modelscope.utils.constant import Tasks
|
||||||
|
import re
|
||||||
|
|
||||||
|
# 定义需要跳过的标签(这些标签的商品不进行推理)
|
||||||
|
SKIP_LABELS = ['当当图书']
|
||||||
|
|
||||||
|
# 全局 NER pipeline
|
||||||
|
ner_pipeline = None
|
||||||
|
|
||||||
|
# 品牌缓存 {品牌关键词: 完整品牌名}
|
||||||
|
brand_cache = {}
|
||||||
|
# 缓存命中统计
|
||||||
|
cache_hit_count = 0
|
||||||
|
cache_miss_count = 0
|
||||||
|
|
||||||
|
|
||||||
|
def load_model():
|
||||||
|
"""
|
||||||
|
加载 ModelScope NER 模型
|
||||||
|
"""
|
||||||
|
global ner_pipeline
|
||||||
|
print("🔧 加载 ModelScope NER 模型中...")
|
||||||
|
ner_pipeline = pipeline(
|
||||||
|
Tasks.named_entity_recognition,
|
||||||
|
'damo/nlp_raner_named-entity-recognition_chinese-base-ecom-50cls'
|
||||||
|
)
|
||||||
|
print("✅ 模型加载完成!")
|
||||||
|
return ner_pipeline
|
||||||
|
|
||||||
|
|
||||||
|
def merge_brand_names(brand_spans):
|
||||||
|
"""
|
||||||
|
合并多个品牌名称
|
||||||
|
规则:如果有中文和英文,拼接为 中文(英文)
|
||||||
|
如果只有一个,直接返回
|
||||||
|
"""
|
||||||
|
if not brand_spans:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# 分离中文和英文
|
||||||
|
chinese_brands = []
|
||||||
|
english_brands = []
|
||||||
|
other_brands = []
|
||||||
|
|
||||||
|
for brand in brand_spans:
|
||||||
|
# 判断是否包含中文字符
|
||||||
|
if re.search(r'[\u4e00-\u9fff]', brand):
|
||||||
|
chinese_brands.append(brand)
|
||||||
|
# 判断是否为英文(只包含英文字母、数字、空格、.、-)
|
||||||
|
elif re.match(r'^[a-zA-Z0-9\s\.\-]+$', brand):
|
||||||
|
english_brands.append(brand)
|
||||||
|
else:
|
||||||
|
other_brands.append(brand)
|
||||||
|
|
||||||
|
# 构建最终品牌名
|
||||||
|
result_parts = []
|
||||||
|
|
||||||
|
# 如果有中文品牌
|
||||||
|
if chinese_brands:
|
||||||
|
chinese_name = chinese_brands[0]
|
||||||
|
# 如果有多个中文品牌,用 / 连接
|
||||||
|
if len(chinese_brands) > 1:
|
||||||
|
chinese_name = '/'.join(chinese_brands)
|
||||||
|
result_parts.append(chinese_name)
|
||||||
|
|
||||||
|
# 如果有英文品牌
|
||||||
|
if english_brands:
|
||||||
|
english_name = english_brands[0]
|
||||||
|
if len(english_brands) > 1:
|
||||||
|
english_name = '/'.join(english_brands)
|
||||||
|
# 如果已经有中文,用括号包裹英文
|
||||||
|
if result_parts:
|
||||||
|
result_parts.append(f"({english_name})")
|
||||||
|
else:
|
||||||
|
result_parts.append(english_name)
|
||||||
|
|
||||||
|
# 如果有其他品牌(非中英文)
|
||||||
|
if other_brands:
|
||||||
|
for brand in other_brands:
|
||||||
|
if brand not in result_parts:
|
||||||
|
result_parts.append(brand)
|
||||||
|
|
||||||
|
return ''.join(result_parts)
|
||||||
|
|
||||||
|
|
||||||
|
def extract_brand_spans(ner_result):
|
||||||
|
"""
|
||||||
|
从 ModelScope NER 结果中提取所有品牌 span
|
||||||
|
返回: [品牌span列表]
|
||||||
|
"""
|
||||||
|
brands = []
|
||||||
|
|
||||||
|
if isinstance(ner_result, dict):
|
||||||
|
output_list = ner_result.get('output', [])
|
||||||
|
for item in output_list:
|
||||||
|
if item.get('type') == '品牌':
|
||||||
|
span = item.get('span', '')
|
||||||
|
if span:
|
||||||
|
brands.append(span)
|
||||||
|
|
||||||
|
return brands
|
||||||
|
|
||||||
|
|
||||||
|
def check_cache(text):
|
||||||
|
"""
|
||||||
|
检查缓存中是否有匹配的品牌
|
||||||
|
遍历缓存,如果缓存中的key包含在商品名称中,则返回对应的品牌
|
||||||
|
返回: (是否命中, 品牌名)
|
||||||
|
"""
|
||||||
|
global cache_hit_count, cache_miss_count
|
||||||
|
|
||||||
|
# 按key长度从长到短排序,优先匹配更长的关键词(更精确)
|
||||||
|
sorted_keys = sorted(brand_cache.keys(), key=len, reverse=True)
|
||||||
|
|
||||||
|
for key in sorted_keys:
|
||||||
|
if key in text:
|
||||||
|
cache_hit_count += 1
|
||||||
|
print(f" ✅ 缓存命中: '{key}' -> '{brand_cache[key]}'")
|
||||||
|
return True, brand_cache[key]
|
||||||
|
|
||||||
|
cache_miss_count += 1
|
||||||
|
return False, None
|
||||||
|
|
||||||
|
|
||||||
|
def update_cache(product_text, brand_spans, merged_brand):
|
||||||
|
"""
|
||||||
|
更新缓存
|
||||||
|
直接将模型返回的品牌 span 作为 key,merged_brand 作为 value
|
||||||
|
"""
|
||||||
|
global brand_cache
|
||||||
|
|
||||||
|
if not merged_brand or merged_brand == "其他":
|
||||||
|
return
|
||||||
|
|
||||||
|
for span in brand_spans:
|
||||||
|
if span and span not in brand_cache:
|
||||||
|
brand_cache[span] = merged_brand
|
||||||
|
print(f" 💾 缓存更新: '{span}' -> '{merged_brand}'")
|
||||||
|
|
||||||
|
|
||||||
|
def reco_single(text, use_cache=True):
|
||||||
|
"""
|
||||||
|
对单个商品名称进行推理,提取品牌信息
|
||||||
|
支持缓存机制
|
||||||
|
"""
|
||||||
|
global brand_cache
|
||||||
|
|
||||||
|
if ner_pipeline is None:
|
||||||
|
raise RuntimeError("模型未加载,请先调用 load_model()")
|
||||||
|
|
||||||
|
# 1. 先检查缓存
|
||||||
|
if use_cache and brand_cache:
|
||||||
|
hit, cached_brand = check_cache(text)
|
||||||
|
if hit:
|
||||||
|
return {'品牌': [[cached_brand, 1.0]]}
|
||||||
|
|
||||||
|
# 2. 缓存未命中,调用模型
|
||||||
|
try:
|
||||||
|
ner_result = ner_pipeline(text)
|
||||||
|
brand_spans = extract_brand_spans(ner_result)
|
||||||
|
|
||||||
|
if not brand_spans:
|
||||||
|
return {'品牌': []}
|
||||||
|
|
||||||
|
# 3. 合并品牌名称
|
||||||
|
merged_brand = merge_brand_names(brand_spans)
|
||||||
|
|
||||||
|
if merged_brand:
|
||||||
|
# 4. 更新缓存:用模型返回的 span 作为 key
|
||||||
|
update_cache(text, brand_spans, merged_brand)
|
||||||
|
return {'品牌': [[merged_brand, 1.0]]}
|
||||||
|
else:
|
||||||
|
return {'品牌': []}
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"⚠️ 推理失败: {text} - {e}")
|
||||||
|
return {'品牌': []}
|
||||||
|
|
||||||
|
|
||||||
|
def extract_brand_from_ner(ner_tags):
|
||||||
|
"""
|
||||||
|
从NER推理结果中提取品牌
|
||||||
|
返回: {商品名称: 品牌名称}
|
||||||
|
"""
|
||||||
|
result = {}
|
||||||
|
|
||||||
|
for product_name, tags in ner_tags.items():
|
||||||
|
brand_list = tags.get('品牌', [])
|
||||||
|
|
||||||
|
if not brand_list:
|
||||||
|
brand_name = "其他"
|
||||||
|
else:
|
||||||
|
# 取第一个品牌(已经合并过的)
|
||||||
|
brand_name = brand_list[0][0]
|
||||||
|
|
||||||
|
result[product_name] = brand_name
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def should_skip_inference(label_value):
|
||||||
|
"""
|
||||||
|
判断是否应该跳过推理
|
||||||
|
"""
|
||||||
|
if pd.isna(label_value):
|
||||||
|
return False
|
||||||
|
return str(label_value).strip() in SKIP_LABELS
|
||||||
|
|
||||||
|
|
||||||
|
def batch_process(input_file='input.xlsx', output_file='output.xlsx',
|
||||||
|
key_column='sentence', label_column='label', batch_size=32):
|
||||||
|
"""
|
||||||
|
批量处理大文件
|
||||||
|
"""
|
||||||
|
global cache_hit_count, cache_miss_count, brand_cache
|
||||||
|
|
||||||
|
try:
|
||||||
|
print(f"📂 读取文件: {input_file}")
|
||||||
|
df = pd.read_excel(input_file, engine='openpyxl')
|
||||||
|
|
||||||
|
if key_column not in df.columns:
|
||||||
|
print(f"❌ 错误: 找不到列 '{key_column}'")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
if label_column not in df.columns:
|
||||||
|
print(f"⚠️ 警告: 找不到列 '{label_column}',将创建新列")
|
||||||
|
df[label_column] = ""
|
||||||
|
|
||||||
|
# 数据清洗
|
||||||
|
df = df.dropna(subset=[key_column])
|
||||||
|
df[key_column] = df[key_column].astype(str).str.strip()
|
||||||
|
df[label_column] = df[label_column].fillna("").astype(str).str.strip()
|
||||||
|
|
||||||
|
# 分离需要推理和不需要推理的数据
|
||||||
|
print("🔍 检查标签,筛选需要推理的数据...")
|
||||||
|
skip_mask = df[label_column].apply(should_skip_inference)
|
||||||
|
need_inference_mask = ~skip_mask
|
||||||
|
|
||||||
|
df_need_inference = df[need_inference_mask].copy()
|
||||||
|
df_skip = df[skip_mask].copy()
|
||||||
|
|
||||||
|
print(f" 📊 总共 {len(df)} 行数据")
|
||||||
|
print(f" ✅ 需要推理: {len(df_need_inference)} 行")
|
||||||
|
print(f" ⏭️ 跳过推理: {len(df_skip)} 行")
|
||||||
|
|
||||||
|
# 处理需要推理的数据
|
||||||
|
if len(df_need_inference) > 0:
|
||||||
|
inference_data = df_need_inference[key_column].tolist()
|
||||||
|
|
||||||
|
print(f"🤖 开始逐条推理 (批次大小: {batch_size})...")
|
||||||
|
print(f"📝 当前缓存条目数: {len(brand_cache)}")
|
||||||
|
all_ner_tags = {}
|
||||||
|
|
||||||
|
total = len(inference_data)
|
||||||
|
for i, text in enumerate(inference_data):
|
||||||
|
if (i + 1) % batch_size == 0 or i == total - 1:
|
||||||
|
print(
|
||||||
|
f" 进度: {i + 1}/{total} | 缓存命中: {cache_hit_count} | 未命中: {cache_miss_count} | 缓存条目: {len(brand_cache)}")
|
||||||
|
|
||||||
|
ner_result = reco_single(text, use_cache=True)
|
||||||
|
all_ner_tags[text] = ner_result
|
||||||
|
|
||||||
|
# 提取品牌
|
||||||
|
print("📊 提取品牌信息...")
|
||||||
|
brand_dict = extract_brand_from_ner(all_ner_tags)
|
||||||
|
|
||||||
|
# 添加品牌列
|
||||||
|
df_need_inference['brand'] = df_need_inference[key_column].map(brand_dict)
|
||||||
|
df_need_inference['brand'] = df_need_inference['brand'].fillna("其他")
|
||||||
|
|
||||||
|
# 处理跳过的数据
|
||||||
|
if len(df_skip) > 0:
|
||||||
|
df_skip['brand'] = df_skip[label_column]
|
||||||
|
|
||||||
|
# 合并数据并恢复原始顺序
|
||||||
|
df_result = pd.concat([df_need_inference, df_skip], ignore_index=True)
|
||||||
|
df_result = df_result.sort_index().reset_index(drop=True)
|
||||||
|
|
||||||
|
# 保存结果
|
||||||
|
df_result.to_excel(output_file, index=False, engine='openpyxl')
|
||||||
|
print(f"✅ 处理完成!结果已保存到: {output_file}")
|
||||||
|
|
||||||
|
# 统计信息
|
||||||
|
total = len(df_result)
|
||||||
|
found = (df_result['brand'] != "其他").sum()
|
||||||
|
skipped_count = len(df_result[df_result[label_column].apply(should_skip_inference)])
|
||||||
|
|
||||||
|
print(f"\n📊 统计信息:")
|
||||||
|
print(f" 总行数: {total}")
|
||||||
|
print(f" 识别品牌: {found}")
|
||||||
|
print(f" 跳过推理: {skipped_count}")
|
||||||
|
print(f" 缓存命中: {cache_hit_count}")
|
||||||
|
print(f" 缓存未命中: {cache_miss_count}")
|
||||||
|
if (cache_hit_count + cache_miss_count) > 0:
|
||||||
|
print(f" 缓存命中率: {cache_hit_count / (cache_hit_count + cache_miss_count) * 100:.1f}%")
|
||||||
|
print(f" 缓存条目数: {len(brand_cache)}")
|
||||||
|
|
||||||
|
print(f"\n📋 结果预览 (前10行):")
|
||||||
|
print(df_result[[key_column, label_column, 'brand']].head(10))
|
||||||
|
|
||||||
|
# 打印缓存内容供参考
|
||||||
|
if brand_cache:
|
||||||
|
print(f"\n📝 缓存内容:")
|
||||||
|
for k, v in brand_cache.items():
|
||||||
|
print(f" '{k}' -> '{v}'")
|
||||||
|
|
||||||
|
return df_result
|
||||||
|
|
||||||
|
except FileNotFoundError:
|
||||||
|
print(f"❌ 错误: 找不到文件 '{input_file}'")
|
||||||
|
sys.exit(1)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"❌ 处理失败: {e}")
|
||||||
|
import traceback
|
||||||
|
traceback.print_exc()
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# 加载模型
|
||||||
|
load_model()
|
||||||
|
|
||||||
|
# 执行批量处理
|
||||||
|
batch_process('input.xlsx', 'output.xlsx', 'sentence', 'label', batch_size=32)
|
||||||
|
|
@ -0,0 +1 @@
|
||||||
|
transformers==4.30.0
|
||||||
Loading…
Reference in New Issue