FireRed-Experimental-Turbo / logging_utils.py
someone-in-the-world's picture
Sync develop into main (#52)
8c202fb unverified
Raw
History Blame Contribute Delete
18.3 kB
import os
import uuid
import threading
import time as _time
from io import BytesIO
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any
import torch
from huggingface_hub import hf_hub_download, CommitOperationAdd, CommitOperationDelete
from PIL.Image import Image as PILImage
from mode import Mode
if TYPE_CHECKING:
from huggingface_hub import HfApi
def print_cuda_visible_devices() -> None:
print("CUDA_VISIBLE_DEVICES=", os.environ.get("CUDA_VISIBLE_DEVICES"), flush=True)
def print_torch_version() -> None:
print("torch.__version__ =", torch.__version__, flush=True)
def print_using_device(device: torch.device) -> None:
print("Using device:", device, flush=True)
def print_cuda_device_count() -> None:
print(f"CUDA device_count={torch.cuda.device_count()}, is_available={torch.cuda.is_available()}", flush=True)
def print_env_gpu(p: torch.cuda._CudaDeviceProperties) -> None:
print(f"[env] GPU: {p.name}, VRAM={p.total_memory/1024**3:.1f}GB, cap={p.major}.{p.minor}", flush=True)
def print_env_cuda_version() -> None:
print(f"[env] CUDA (torch build): {torch.version.cuda}", flush=True)
def print_env_cudnn_version() -> None:
print(f"[env] cuDNN: {torch.backends.cudnn.version()}", flush=True) # type: ignore[no-untyped-call]
def print_env_package_version(pkg: str, version: str) -> None:
print(f"[env] {pkg}=={version}", flush=True)
def print_env_package_version_unavailable(pkg: str, error: Exception) -> None:
print(f"[env] {pkg}==? ({error})", flush=True)
def print_env_ram(total_gb: float, avail_gb: float) -> None:
print(f"[env] RAM: {total_gb:.0f}GB total, {avail_gb:.0f}GB available", flush=True)
def print_env_ram_unavailable(error: Exception) -> None:
print(f"[env] RAM: unavailable ({error})", flush=True)
def print_tf32_enabled() -> None:
print("[startup] TF32 enabled", flush=True)
def print_heartbeat(label: str, elapsed: float) -> None:
print(f"[startup] {label} still loading... ({elapsed:.0f}s)", flush=True)
def print_unpatched_fp8_param(name: str, pname: str, module_type: str) -> None:
print(
f"[startup] WARNING: unpatched fp8 parameter {name}.{pname} "
f"({module_type}) — will likely error at inference",
flush=True,
)
def print_loading_transformer() -> None:
print("[startup] loading transformer from_pretrained (prithivMLmods/Qwen-Image-Edit-Rapid-AIO-V23)...", flush=True)
def print_transformer_loaded(elapsed: float) -> None:
print(f"[startup] transformer loaded in {elapsed:.1f}s", flush=True)
def print_transformer_patched(n_patched: int) -> None:
print(f"[startup] patched {n_patched} fp8-resident nn.Linear/RMSNorm modules for just-in-time upcast", flush=True)
def print_lm_head_dropped(tag: str) -> None:
print(f"[{tag}] dropped text_encoder.lm_head (unused — pipeline only reads hidden_states)", flush=True)
def print_transformer_memory_footprint(gb: float) -> None:
print(f"[startup] transformer memory footprint: {gb:.2f}GB", flush=True)
def print_transformer_memory_footprint_unavailable(error: Exception) -> None:
print(f"[startup] transformer memory footprint: unavailable ({error})", flush=True)
def print_loading_pipeline() -> None:
print("[startup] loading pipeline from_pretrained (FireRedTeam/FireRed-Image-Edit-1.1)...", flush=True)
def print_vae_tiling(height: int, width: int, use_tiling: bool) -> None:
print(f"[startup] VAE tiling: threshold={height}x{width}px use_tiling={use_tiling}", flush=True)
def print_pipeline_loaded(elapsed: float) -> None:
print(f"[startup] pipeline loaded in {elapsed:.1f}s", flush=True)
def print_setting_attn_processor() -> None:
print("[startup] setting cuDNN SDPA attention processor...", flush=True)
def print_attn_processor_set() -> None:
print("[startup] cuDNN SDPA attention processor set.", flush=True)
def print_timing_divider() -> None:
print("[timing] ─────────────────────────────────────")
def print_timing_lines(lines: list[str]) -> None:
print("\n".join(lines))
def print_infer_exception(e: Exception) -> None:
print(f"[infer] EXCEPTION type={type(e).__module__}.{type(e).__qualname__} repr={e!r}")
def print_infer_start_header() -> None:
print("[infer] ===== START =====")
def print_infer_params(steps: int, guidance_scale: float, seed: int, gpu_duration: int, mode: Mode) -> None:
print(f"[infer] steps={steps}, guidance={guidance_scale}, seed={seed}, gpu_duration={gpu_duration}s, mode={mode.value}")
def print_infer_prompt(prompt: str) -> None:
print(f"[infer] prompt={repr(prompt[:120])}")
def print_infer_gpu_properties(p: torch.cuda._CudaDeviceProperties) -> None:
print(f"[infer] GPU: {p.name}, total={p.total_memory/1024**3:.1f}GB, cap={p.major}.{p.minor}")
def print_loading_int8_text_encoder(repo: str) -> None:
print(f"[startup] loading int8 text_encoder from_pretrained ({repo})...", flush=True)
def print_int8_text_encoder_loaded(elapsed: float) -> None:
print(f"[startup] int8 text_encoder loaded in {elapsed:.1f}s", flush=True)
def print_first_call_into_module(name: str, mem_str: str, elapsed: float) -> None:
print(f"[infer] first call into {name}{mem_str} | t={elapsed:.1f}s")
def print_step_done(step_idx: int, steps: int, delta_ms: float, tag: str, elapsed: float) -> None:
print(f"[infer] step {step_idx+1}/{steps} done — {delta_ms:.0f}ms{tag} | t={elapsed:.1f}s")
def print_text_encoder_offload_skipped_int8() -> None:
print("[infer] skipping text_encoder offload (int8, ~8.75GB footprint doesn't need it; .to() device-move works but isn't wired up here)")
def print_pre_vae_decode(mem_str: str, elapsed: float) -> None:
print(f"[infer] pre-VAE-decode — {mem_str} | t={elapsed:.1f}s")
def print_infer_error(e: Exception, elapsed: float) -> None:
print(f"[infer] ERROR: {type(e).__name__}: {e} | t={elapsed:.1f}s")
def print_infer_traceback() -> None:
import traceback
print(traceback.format_exc())
def print_cuda_sync_after_error(cuda_err: Exception) -> None:
print(f"[infer] CUDA synchronize after error: {cuda_err}")
def print_gpu_mem_status(mem_str: str, elapsed: float) -> None:
print(f"[infer] {mem_str} — t={elapsed:.1f}s")
def print_images_predecoded(n: int, width: int, height: int, seed: int) -> None:
print(f"[infer] {n} image(s) pre-decoded, output={width}x{height}, seed={seed}")
def print_vae_tiling_activation(will_tile: bool, height: int, width: int) -> None:
print(f"[infer] VAE tiling will {'activate' if will_tile else 'NOT activate'} "
f"(threshold={height}x{width}px)")
def print_calling_pipe(elapsed: float) -> None:
print(f"[infer] calling pipe... t={elapsed:.1f}s")
def print_vae_decode_done(mem_str: str, elapsed: float) -> None:
print(f"[infer] VAE decode + postprocess done — {mem_str} | t={elapsed:.1f}s")
def print_infer_end(elapsed: float) -> None:
print(f"[infer] ===== END t={elapsed:.1f}s =====")
def print_building_example_thumbnails() -> None:
print("Building example thumbnails...")
def print_built_example_cards(n: int) -> None:
print(f"Built {n} example cards.")
def print_built_suggestion_chips(n: int) -> None:
print(f"Built {n} suggestion chips.")
def print_thumbnail_error(path: str, e: Exception) -> None:
print(f"Thumbnail error for {path}: {e}")
def print_encode_error(path: str, e: Exception) -> None:
print(f"Encode error for {path}: {e}")
def print_decode_error(e: Exception) -> None:
print(f"Error decoding image: {e}")
def print_cudnn_sdpa_fallback(e: RuntimeError) -> None:
print(f"[attn] cuDNN SDPA backend unavailable ({e}), falling back to default", flush=True)
def print_pipe_phase_timing(label: str, delta_ms: float, elapsed_s: float) -> None:
print(f"[pipe] {label}{delta_ms:.0f}ms | t={elapsed_s:.1f}s", flush=True)
def _img_to_jpeg(img: PILImage | None, quality: int = 85) -> bytes | None:
if img is None:
return None
buf = BytesIO()
img.convert("RGB").save(buf, format="JPEG", quality=quality)
return buf.getvalue()
def _build_table(pil_inputs: list[PILImage], output_pil: PILImage | None, prompt: str, seed: int,
steps: int, guidance_scale: float, input_width: int, input_height: int,
duration_seconds: float, success: bool, error_message: str, now: datetime) -> Any:
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"},
"prompt": {"dtype": "string", "_type": "Value"},
"seed": {"dtype": "int32", "_type": "Value"},
"steps": {"dtype": "int32", "_type": "Value"},
"guidance_scale": {"dtype": "float32", "_type": "Value"},
"input_images": {"feature": {"_type": "Image"}, "_type": "Sequence"},
"output_image": {"_type": "Image"},
"duration_seconds": {"dtype": "float32", "_type": "Value"},
"input_width": {"dtype": "int32", "_type": "Value"},
"input_height": {"dtype": "int32", "_type": "Value"},
"success": {"dtype": "bool", "_type": "Value"},
"error_message": {"dtype": "string", "_type": "Value"},
}}}).encode()
schema = pa.schema([
("timestamp", pa.float64()),
("prompt", pa.string()),
("seed", pa.int32()),
("steps", pa.int32()),
("guidance_scale", pa.float32()),
("input_images", pa.list_(img_struct)),
("output_image", img_struct),
("duration_seconds", pa.float32()),
("input_width", pa.int32()),
("input_height", pa.int32()),
("success", pa.bool_()),
("error_message", pa.string()),
], metadata={b"huggingface": hf_meta})
def _img(b: bytes | None) -> dict[str, Any]:
return {"bytes": b, "path": None}
input_jpegs = [_img_to_jpeg(img) for img in pil_inputs]
output_jpeg = _img_to_jpeg(output_pil)
return pa.table({
"timestamp": pa.array([now.timestamp()], type=pa.float64()),
"prompt": pa.array([prompt], type=pa.string()),
"seed": pa.array([int(seed)], type=pa.int32()),
"steps": pa.array([int(steps)], type=pa.int32()),
"guidance_scale": pa.array([float(guidance_scale)], type=pa.float32()),
"input_images": pa.array([[_img(b) for b in input_jpegs]], type=pa.list_(img_struct)),
"output_image": pa.array([_img(output_jpeg) if output_jpeg else None], type=img_struct),
"duration_seconds": pa.array([float(duration_seconds)], type=pa.float32()),
"input_width": pa.array([int(input_width)], type=pa.int32()),
"input_height": pa.array([int(input_height)], type=pa.int32()),
"success": pa.array([bool(success)], type=pa.bool_()),
"error_message": pa.array([str(error_message)], type=pa.string()),
}, schema=schema)
def _write_parquet(table: Any) -> str:
import tempfile
import pyarrow.parquet as pq
with tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as tmp:
path = tmp.name
pq.write_table(table, path)
return path
def _make_path(now: datetime, uid: str) -> str:
return f"data/{now.strftime('%Y-%m-%d-%H%M%S')}-{uid}.parquet"
def print_log_list_existing_files_failed(e: Exception) -> None:
print(f"[log] could not list existing files (empty repo?): {e}")
def _list_existing_files(api: "HfApi", repo_id: str) -> list[str]:
try:
entries = list(api.list_repo_tree(repo_id, repo_type="dataset", path_in_repo="data"))
except Exception as e:
print_log_list_existing_files_failed(e)
return []
return sorted(f.path for f in entries if f.path.endswith(".parquet"))
def _build_add_ops(batch: list[tuple[str, str]]) -> list[CommitOperationAdd]:
return [CommitOperationAdd(path_in_repo=p, path_or_fileobj=local)
for p, local in batch]
def _build_delete_ops(existing_files: list[str], n_new: int, max_files: int) -> list[CommitOperationDelete]:
total_after = len(existing_files) + n_new
if max_files <= 0 or total_after <= max_files:
return []
n_delete = total_after - max_files
return [CommitOperationDelete(path_in_repo=p) for p in existing_files[:n_delete]]
def _delete_temp_files(batch: list[tuple[str, str]]) -> None:
for _, local in batch:
try:
os.unlink(local)
except Exception:
pass
def print_log_squash_marker_not_found(e: Exception) -> None:
print(f"[log] squash marker not found ({e}), proceeding with squash")
def print_log_squashed_history(repo_id: str) -> None:
print(f"[log] squashed history for {repo_id}")
def print_log_squash_warning(e: Exception) -> None:
print(f"[log] squash warning: {e}")
def _squash_if_needed(api: "HfApi", repo_id: str) -> None:
marker = "metadata/last_squash.txt"
today = datetime.now(timezone.utc).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_log_squash_marker_not_found(e)
api.super_squash_history(repo_id=repo_id, repo_type="dataset")
api.upload_file(path_or_fileobj=today.encode(), path_in_repo=marker,
repo_id=repo_id, repo_type="dataset")
print_log_squashed_history(repo_id)
except Exception as e:
print_log_squash_warning(e)
def print_log_skipped(has_token: bool, has_repo: bool) -> None:
print(f"[log] skipped — token={'set' if has_token else 'missing'}, repo={'set' if has_repo else 'missing'}")
def print_log_queued(path_in_repo: str, pending: int) -> None:
print(f"[log] queued {path_in_repo} (pending={pending})")
def print_log_inference_warning(e: Exception, tb: str) -> None:
print(f"[log] WARNING: {e}\n{tb}")
def print_log_inference_total(elapsed: float) -> None:
print(f"[log] log_inference total: {elapsed:.3f}s")
def print_log_batch_upload_warning(e: Exception) -> None:
print(f"[log] batch upload warning: {e}")
def print_log_committed(n_files: int, n_pruned: int) -> None:
print(f"[log] committed {n_files} file(s), pruned {n_pruned}")
class LogUploader:
def __init__(self, token: str | None, repo_id: str | None, max_files: int = 5000, batch_interval: int = 60) -> None:
self._token = token
self._repo_id = repo_id
self._max_files = max_files
self._batch_interval = batch_interval
self._pending: list[tuple[str, str]] = []
self._lock = threading.Lock()
if token and repo_id:
threading.Thread(target=self._loop, daemon=True, name="log-uploader").start()
def log_inference(self, pil_inputs: list[PILImage], output_pil: PILImage | None, prompt: str, seed: int,
steps: int, guidance_scale: float, input_width: int, input_height: int,
duration_seconds: float, success: bool, error_message: str = "") -> None:
if not self._token or not self._repo_id:
print_log_skipped(bool(self._token), bool(self._repo_id))
return
t0 = _time.perf_counter()
try:
now = datetime.now(timezone.utc)
table = _build_table(pil_inputs, output_pil, prompt, seed, steps, guidance_scale,
input_width, input_height, duration_seconds, success, error_message, now)
local_path = _write_parquet(table)
path_in_repo = _make_path(now, uuid.uuid4().hex[:8])
self._enqueue(path_in_repo, local_path)
print_log_queued(path_in_repo, len(self._pending))
except Exception as e:
import traceback as _tb
print_log_inference_warning(e, _tb.format_exc())
print_log_inference_total(_time.perf_counter() - t0)
def _enqueue(self, path_in_repo: str, local_path: str) -> None:
with self._lock:
self._pending.append((path_in_repo, local_path))
def _drain(self) -> list[tuple[str, str]]:
with self._lock:
batch = self._pending[:]
self._pending.clear()
return batch
def _requeue(self, batch: list[tuple[str, str]]) -> None:
with self._lock:
self._pending[:0] = batch
def _loop(self) -> None:
while True:
_time.sleep(self._batch_interval)
self._flush()
def _flush(self) -> None:
batch = self._drain()
if not batch:
return
try:
self._commit_batch(batch)
_delete_temp_files(batch)
except Exception as e:
print_log_batch_upload_warning(e)
self._requeue(batch)
def _commit_batch(self, batch: list[tuple[str, str]]) -> None:
from huggingface_hub import HfApi
assert self._repo_id is not None
api = HfApi(token=self._token)
api.create_repo(repo_id=self._repo_id, repo_type="dataset", private=True, exist_ok=True)
existing = _list_existing_files(api, self._repo_id)
add_ops = _build_add_ops(batch)
del_ops = _build_delete_ops(existing, len(batch), self._max_files)
api.create_commit(
repo_id=self._repo_id, repo_type="dataset",
operations=[*add_ops, *del_ops],
commit_message=f"[log] batch {len(batch)}" + (f", prune {len(del_ops)}" if del_ops else ""),
)
print_log_committed(len(batch), len(del_ops))
_squash_if_needed(api, self._repo_id)