Selectable-Face-Swap / logging_utils.py
someone-in-the-world's picture
Claude Sonnet 4.6
Add pre-commit black formatter and apply initial formatting
db79f31
Raw History Blame Contribute Delete
8.41 kB
import os
import threading
import uuid
from io import BytesIO
from datetime import datetime, timezone, timedelta
from huggingface_hub import hf_hub_download
HF_TOKEN = os.environ.get("HF_TOKEN")
DATASET_REPO = os.environ.get("LOG_DATASET_REPO")
MAX_LOG_DAYS = int(os.environ.get("LOG_MAX_DAYS", "7"))
def _img_to_jpeg(img, quality=85):
if img is None:
return None
buf = BytesIO()
img.convert("RGB").save(buf, format="JPEG", quality=quality)
return buf.getvalue()
def _build_table(
target_pil,
source_pils,
output_pil,
selected_indices,
duration_seconds,
success,
error_message,
now,
):
import json as _json
import pyarrow as pa
img_struct = pa.struct([("bytes", pa.binary()), ("path", pa.string())])
hf_meta = _json.dumps(
{
"info": {
"features": {
"timestamp": {"dtype": "float64", "_type": "Value"},
"target_image": {"_type": "Image"},
"source_images": {
"feature": {"_type": "Image"},
"_type": "Sequence",
},
"output_image": {"_type": "Image"},
"selected_indices": {
"feature": {"dtype": "int32", "_type": "Value"},
"_type": "Sequence",
},
"n_faces_swapped": {"dtype": "int32", "_type": "Value"},
"duration_seconds": {"dtype": "float32", "_type": "Value"},
"success": {"dtype": "bool", "_type": "Value"},
"error_message": {"dtype": "string", "_type": "Value"},
}
}
}
).encode()
schema = pa.schema(
[
("timestamp", pa.float64()),
("target_image", img_struct),
("source_images", pa.list_(img_struct)),
("output_image", img_struct),
("selected_indices", pa.list_(pa.int32())),
("n_faces_swapped", pa.int32()),
("duration_seconds", pa.float32()),
("success", pa.bool_()),
("error_message", pa.string()),
],
metadata={b"huggingface": hf_meta},
)
def _img(b):
return {"bytes": b, "path": None}
target_jpeg = _img_to_jpeg(target_pil)
source_jpegs = [_img_to_jpeg(img) for img in source_pils]
output_jpeg = _img_to_jpeg(output_pil)
return pa.table(
{
"timestamp": pa.array([now.timestamp()], type=pa.float64()),
"target_image": pa.array(
[_img(target_jpeg) if target_jpeg else None], type=img_struct
),
"source_images": pa.array(
[[_img(b) for b in source_jpegs]], type=pa.list_(img_struct)
),
"output_image": pa.array(
[_img(output_jpeg) if output_jpeg else None], type=img_struct
),
"selected_indices": pa.array(
[[int(i) for i in selected_indices]], type=pa.list_(pa.int32())
),
"n_faces_swapped": pa.array([len(selected_indices)], type=pa.int32()),
"duration_seconds": pa.array([float(duration_seconds)], type=pa.float32()),
"success": pa.array([bool(success)], type=pa.bool_()),
"error_message": pa.array([str(error_message)], type=pa.string()),
},
schema=schema,
)
def _upload_parquet(api, repo_id, table, path_in_repo):
import tempfile
import pyarrow.parquet as pq
tmp_path = None
try:
with tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as tmp:
tmp_path = tmp.name
pq.write_table(table, tmp_path)
print(f"[log] uploading {path_in_repo} ({os.path.getsize(tmp_path)//1024}KB)")
api.upload_file(
path_or_fileobj=tmp_path,
path_in_repo=path_in_repo,
repo_id=repo_id,
repo_type="dataset",
)
print(f"[log] upload done — {repo_id}/{path_in_repo}")
finally:
if tmp_path:
try:
os.unlink(tmp_path)
except Exception as e:
print(f"[log] failed to delete temp file {tmp_path}: {e}")
def _make_path(now, uid):
return f"data/{now.strftime('%Y-%m-%d-%H%M%S')}-{uid}.parquet"
def _file_date(path):
return os.path.basename(path)[:10]
def _maybe_squash_history(api, repo_id, now):
marker = "metadata/last_squash.txt"
today = now.strftime("%Y-%m-%d")
try:
try:
local = hf_hub_download(
repo_id=repo_id, filename=marker, repo_type="dataset", token=api.token
)
if open(local).read().strip() == today:
return
except Exception as e:
print(f"[log] squash marker not found ({e}), proceeding with squash")
api.super_squash_history(repo_id=repo_id, repo_type="dataset")
print(f"[log] squashed history for {repo_id}")
api.upload_file(
path_or_fileobj=today.encode(),
path_in_repo=marker,
repo_id=repo_id,
repo_type="dataset",
)
print(f"[log] updated squash marker: {today}")
except Exception as e:
print(f"[log] squash warning: {e}")
def _prune_old_files(api, repo_id, keep_days, now):
if keep_days <= 0:
return
cutoff = (now - timedelta(days=keep_days)).strftime("%Y-%m-%d")
try:
to_delete = [
f.path
for f in api.list_repo_tree(
repo_id, repo_type="dataset", path_in_repo="data"
)
if f.path.endswith(".parquet") and _file_date(f.path) < cutoff
]
for path in to_delete:
api.delete_file(path_in_repo=path, repo_id=repo_id, repo_type="dataset")
print(f"[log] pruned: {path}")
if to_delete:
print(f"[log] pruned {len(to_delete)} old file(s)")
except Exception as e:
print(f"[log] prune warning: {e}")
def log_inference(
target_pil,
source_pils,
output_pil,
selected_indices,
duration_seconds,
success,
error_message="",
):
import time as _time
_t0 = _time.perf_counter()
if not HF_TOKEN or not DATASET_REPO:
print(
f"[log] skipped — HF_TOKEN={'set' if HF_TOKEN else 'missing'}, DATASET_REPO={'set' if DATASET_REPO else 'missing'}"
)
return
try:
from huggingface_hub import HfApi
now = datetime.now(timezone.utc)
_t1 = _time.perf_counter()
table = _build_table(
target_pil,
source_pils,
output_pil,
selected_indices,
duration_seconds,
success,
error_message,
now,
)
print(f"[log] build_table: {_time.perf_counter() - _t1:.3f}s")
uid = uuid.uuid4().hex[:8]
path_in_repo = _make_path(now, uid)
_t2 = _time.perf_counter()
api = HfApi(token=HF_TOKEN)
api.create_repo(
repo_id=DATASET_REPO, repo_type="dataset", private=True, exist_ok=True
)
print(f"[log] create_repo: {_time.perf_counter() - _t2:.3f}s")
_t3 = _time.perf_counter()
_upload_parquet(api, DATASET_REPO, table, path_in_repo)
print(f"[log] upload_parquet: {_time.perf_counter() - _t3:.3f}s")
_t4 = _time.perf_counter()
_prune_old_files(api, DATASET_REPO, MAX_LOG_DAYS, now)
print(f"[log] prune_old_files: {_time.perf_counter() - _t4:.3f}s")
_t5 = _time.perf_counter()
_maybe_squash_history(api, DATASET_REPO, now)
print(f"[log] squash_history: {_time.perf_counter() - _t5:.3f}s")
except Exception as log_err:
import traceback as _tb
print(f"[log] WARNING: {log_err}\n{_tb.format_exc()}")
finally:
print(f"[log] log_inference total: {_time.perf_counter() - _t0:.3f}s")
def spawn_log(
target_pil, source_pils, result_pil, selected_indices, duration, success, error=""
):
threading.Thread(
target=log_inference,
args=(
target_pil,
source_pils,
result_pil,
selected_indices,
duration,
success,
error,
),
daemon=True,
).start()