Spaces:
Running on Zero
Running on Zero
Download count.py from broyang/CrowdCounter: direct link, hf CLI and curl.
- Browser
- Download file 12.6 kB
-
https://huggingface.co/spaces/broyang/CrowdCounter/resolve/main/count.py
- Command line
-
hf download hf://spaces/broyang/CrowdCounter/count.py
-
curl -L -o count.py https://huggingface.co/spaces/broyang/CrowdCounter/resolve/main/count.py
12.6 kB
| import os | |
| os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| import cv2 | |
| import numpy as np | |
| import torch | |
| from PIL import Image, ImageDraw, ImageFont | |
| REPO_ROOT = Path(__file__).resolve().parent | |
| PET_ROOT = REPO_ROOT / "PET" | |
| sys.path.insert(0, str(PET_ROOT)) | |
| import util.misc as utils | |
| from models import build_model | |
| DEFAULT_WEIGHTS = REPO_ROOT / "weights" / "PET_Finetuned.safetensors" | |
| MAX_INFERENCE_SIDE = 2560 | |
| IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1) | |
| IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1) | |
| def normalize(rgb): | |
| tensor = torch.from_numpy(np.ascontiguousarray(rgb)).permute(2, 0, 1).float().div_(255.0) | |
| return (tensor - IMAGENET_MEAN) / IMAGENET_STD | |
| MODEL_ARGS = dict( | |
| backbone="vgg16_bn", | |
| position_embedding="sine", | |
| dec_layers=2, | |
| dim_feedforward=512, | |
| hidden_dim=256, | |
| dropout=0.0, | |
| nheads=8, | |
| set_cost_class=1, | |
| set_cost_point=0.05, | |
| ce_loss_coef=1.0, | |
| point_loss_coef=5.0, | |
| eos_coef=0.5, | |
| dataset_file="SHA", | |
| data_path="./data/ShanghaiTech/PartA", | |
| ) | |
| MAC_FONTS = [ | |
| "/System/Library/Fonts/Supplemental/Arial Bold.ttf", | |
| "/System/Library/Fonts/Helvetica.ttc", | |
| "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf", | |
| "/usr/share/fonts/truetype/liberation/LiberationSans-Bold.ttf", | |
| ] | |
| def select_device(requested): | |
| if requested and requested != "auto": | |
| return torch.device(requested) | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| return torch.device("cpu") | |
| def load_state_dict(path): | |
| path = Path(path) | |
| if not path.exists(): | |
| raise FileNotFoundError(f"Weights not found: {path}") | |
| if path.suffix == ".safetensors": | |
| from safetensors.torch import load_file | |
| return load_file(str(path), device="cpu") | |
| checkpoint = torch.load(str(path), map_location="cpu") | |
| if isinstance(checkpoint, dict) and isinstance(checkpoint.get("model"), dict): | |
| return checkpoint["model"] | |
| return checkpoint | |
| def build_pet(device): | |
| args = SimpleNamespace(device=str(device), **MODEL_ARGS) | |
| previous_cwd = os.getcwd() | |
| os.chdir(PET_ROOT) | |
| try: | |
| model, _ = build_model(args) | |
| finally: | |
| os.chdir(previous_cwd) | |
| model.to(device) | |
| model.eval() | |
| return model | |
| def resize_for_inference(rgb): | |
| h, w = rgb.shape[:2] | |
| longest = max(h, w) | |
| if longest <= MAX_INFERENCE_SIDE: | |
| return rgb, 1.0 | |
| scale = MAX_INFERENCE_SIDE / float(longest) | |
| resized = cv2.resize(rgb, (round(w * scale), round(h * scale)), interpolation=cv2.INTER_LINEAR) | |
| return resized, scale | |
| def predict_batch(model, crops_rgb, device, thrs): | |
| tensors = [normalize(crop) for crop in crops_rgb] | |
| samples = utils.nested_tensor_from_tensor_list(tensors).to(device) | |
| padded_h, padded_w = samples.tensors.shape[-2:] | |
| outputs = model(samples, test=True, thrs=thrs) | |
| points = outputs["pred_points"] | |
| logits = outputs["pred_logits"] | |
| if points.dim() == 3: | |
| points, logits = points[0], logits[0] | |
| points = points.detach().cpu().numpy() | |
| scores = torch.softmax(logits.detach().float().cpu(), -1)[:, 1].numpy() | |
| sparse_ends = np.cumsum([0] + outputs["sparse_keep_counts"].tolist()) | |
| dense_ends = np.cumsum([0] + outputs["dense_keep_counts"].tolist()) + sparse_ends[-1] | |
| results = [] | |
| for i, crop in enumerate(crops_rgb): | |
| idx = np.concatenate([ | |
| np.arange(sparse_ends[i], sparse_ends[i + 1]), | |
| np.arange(dense_ends[i], dense_ends[i + 1]), | |
| ]).astype(int) | |
| crop_h, crop_w = crop.shape[:2] | |
| ys = np.clip(points[idx, 0] * float(padded_h), 0, crop_h - 1) | |
| xs = np.clip(points[idx, 1] * float(padded_w), 0, crop_w - 1) | |
| results.append((np.stack([xs, ys], axis=1), scores[idx])) | |
| return results | |
| def tile_starts(extent, tile, stride, offset): | |
| if extent <= tile: | |
| return [0] | |
| last = extent - tile | |
| return sorted({0, last} | set(range(offset, last + 1, stride))) | |
| def grid_pass(model, up_rgb, device, tile, overlap, thrs, offset): | |
| up_h, up_w = up_rgb.shape[:2] | |
| stride = max(1, tile - overlap) | |
| margin = overlap // 2 | |
| ys_starts = tile_starts(up_h, tile, stride, offset) | |
| xs_starts = tile_starts(up_w, tile, stride, offset) | |
| tiles = [] | |
| for iy, y0 in enumerate(ys_starts): | |
| for ix, x0 in enumerate(xs_starts): | |
| x1, y1 = min(x0 + tile, up_w), min(y0 + tile, up_h) | |
| own_x0 = 0 if ix == 0 else x0 + margin | |
| own_y0 = 0 if iy == 0 else y0 + margin | |
| own_x1 = up_w if ix == len(xs_starts) - 1 else x1 - margin | |
| own_y1 = up_h if iy == len(ys_starts) - 1 else y1 - margin | |
| tiles.append((x0, y0, x1, y1, own_x0, own_y0, own_x1, own_y1)) | |
| collected_pts, collected_scores = [], [] | |
| for x0, y0, x1, y1, own_x0, own_y0, own_x1, own_y1 in tiles: | |
| local, local_scores = predict_batch(model, [up_rgb[y0:y1, x0:x1]], device, thrs)[0] | |
| if local.size == 0: | |
| continue | |
| gx, gy = local[:, 0] + x0, local[:, 1] + y0 | |
| keep = (gx >= own_x0) & (gx < own_x1) & (gy >= own_y0) & (gy < own_y1) | |
| collected_pts.append(np.stack([gx[keep], gy[keep]], axis=1)) | |
| collected_scores.append(local_scores[keep]) | |
| if not collected_pts: | |
| return np.zeros((0, 2), dtype=np.float32), np.zeros((0,), dtype=np.float32) | |
| return np.concatenate(collected_pts), np.concatenate(collected_scores) | |
| def dedup_points(points, scores, max_radius): | |
| if len(points) < 6: | |
| return points | |
| from scipy.spatial import cKDTree | |
| order = np.argsort(-scores) | |
| ranked = points[order] | |
| tree = cKDTree(ranked) | |
| spacing = tree.query(ranked, k=5)[0][:, 4] | |
| radii = np.clip(0.45 * spacing, 1.5, max_radius) | |
| keep = np.ones(len(ranked), dtype=bool) | |
| for i, neighbors in enumerate(tree.query_ball_point(ranked, radii)): | |
| if keep[i]: | |
| for j in neighbors: | |
| if j > i: | |
| keep[j] = False | |
| return ranked[keep] | |
| def predict_candidates_tiled(model, image_bgr, device, upscale, tile, overlap, thrs=0.1, grids=2, time_budget=None): | |
| import time | |
| started = time.monotonic() | |
| orig_h, orig_w = image_bgr.shape[:2] | |
| rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB) | |
| if upscale != 1.0: | |
| rgb = cv2.resize(rgb, (round(orig_w * upscale), round(orig_h * upscale)), interpolation=cv2.INTER_CUBIC) | |
| stride = max(1, tile - overlap) | |
| offsets = [0, stride // 2][:max(1, grids)] | |
| passes = [] | |
| for pass_index, offset in enumerate(offsets): | |
| if pass_index > 0 and time_budget is not None and time.monotonic() - started > time_budget * 0.45: | |
| break | |
| passes.append(grid_pass(model, rgb, device, tile, overlap, thrs, offset)) | |
| points = np.concatenate([p for p, _ in passes]) / upscale | |
| scores = np.concatenate([s for _, s in passes]) | |
| if points.size == 0: | |
| return np.zeros((0, 2), dtype=np.float32), scores | |
| points[:, 0] = np.clip(points[:, 0], 0, orig_w - 1) | |
| points[:, 1] = np.clip(points[:, 1], 0, orig_h - 1) | |
| return points, scores | |
| def filter_predictions(points, scores, threshold, dedup_radius=4.0): | |
| keep = scores > float(threshold) | |
| return dedup_points(points[keep], scores[keep], dedup_radius) | |
| def predict_points_tiled(model, image_bgr, device, upscale, tile, overlap, thrs=0.5, grids=2, dedup_radius=4.0, time_budget=None): | |
| points, scores = predict_candidates_tiled( | |
| model, image_bgr, device, upscale, tile, overlap, thrs, grids, time_budget, | |
| ) | |
| return dedup_points(points, scores, dedup_radius) | |
| def predict_points(model, image_bgr, device, thrs=0.5): | |
| rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB) | |
| resized, scale = resize_for_inference(rgb) | |
| points, _ = predict_batch(model, [resized], device, thrs)[0] | |
| if points.size == 0: | |
| return points | |
| points /= scale | |
| orig_h, orig_w = image_bgr.shape[:2] | |
| points[:, 0] = np.clip(points[:, 0], 0, orig_w - 1) | |
| points[:, 1] = np.clip(points[:, 1], 0, orig_h - 1) | |
| return points | |
| def load_font(size): | |
| for path in MAC_FONTS: | |
| if os.path.exists(path): | |
| try: | |
| return ImageFont.truetype(path, size) | |
| except OSError: | |
| continue | |
| return ImageFont.load_default() | |
| def render_overlay(image_bgr, points, count, dot_color=(90, 17, 255), logo=None): | |
| rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB) | |
| canvas = Image.fromarray(rgb).convert("RGB") | |
| draw = ImageDraw.Draw(canvas) | |
| radius = max(2, round(max(canvas.size) * 0.0022)) | |
| for x, y in points: | |
| draw.ellipse((x - radius, y - radius, x + radius, y + radius), fill=dot_color) | |
| width, height = canvas.size | |
| bar_height = max(64, round(height * 0.085)) | |
| footer = Image.new("RGBA", (width, height + bar_height), color=(15, 15, 18, 255)) | |
| footer.paste(canvas, (0, 0)) | |
| footer_draw = ImageDraw.Draw(footer) | |
| label = f"Estimated crowd count: {count:,}" | |
| stroke = max(1, round(bar_height * 0.03)) | |
| font = load_font(max(26, round(bar_height * 0.5))) | |
| bbox = footer_draw.textbbox((0, 0), label, font=font, stroke_width=stroke) | |
| text_w, text_h = bbox[2] - bbox[0], bbox[3] - bbox[1] | |
| footer_draw.text( | |
| ((width - text_w) / 2, height + (bar_height - text_h) / 2 - bbox[1]), | |
| label, fill=(255, 255, 255), font=font, | |
| stroke_width=stroke, stroke_fill=(15, 15, 18), | |
| ) | |
| if logo is not None: | |
| logo_height = round(bar_height * 0.55) | |
| logo_width = round(logo_height * logo.width / logo.height) | |
| bar_margin = round(bar_height * 0.35) | |
| if bar_margin + logo_width < (width - text_w) / 2: | |
| resized_logo = logo.resize((logo_width, logo_height), Image.LANCZOS) | |
| footer.alpha_composite( | |
| resized_logo, | |
| (bar_margin, height + (bar_height - logo_height) // 2), | |
| ) | |
| return footer.convert("RGB") | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Aerial crowd counter (PET) with count overlay") | |
| parser.add_argument("image", type=str, help="Path to input image") | |
| parser.add_argument("-o", "--output", type=str, default="", help="Output image path") | |
| parser.add_argument("--weights", type=str, default=str(DEFAULT_WEIGHTS)) | |
| parser.add_argument("--device", type=str, default="auto", help="auto | cpu | mps | cuda") | |
| parser.add_argument("--json", type=str, default="", help="Optional path to write count + points JSON") | |
| parser.add_argument("--plain", action="store_true", help="Disable tiling (single-pass, faster, less accurate on dense scenes)") | |
| parser.add_argument("--upscale", type=float, default=1.0, help="Upscale factor before tiling") | |
| parser.add_argument("--tile", type=int, default=768, help="Tile size in upscaled pixels") | |
| parser.add_argument("--overlap", type=int, default=192, help="Tile overlap in upscaled pixels") | |
| parser.add_argument("--thrs", type=float, default=0.35, help="Head-confidence threshold; lower recovers more dense-core heads (0.5 = PET default)") | |
| args = parser.parse_args() | |
| image_path = Path(args.image) | |
| image_bgr = cv2.imread(str(image_path)) | |
| if image_bgr is None: | |
| raise SystemExit(f"Failed to read image: {image_path}") | |
| device = select_device(args.device) | |
| print(f"device: {device}") | |
| model = build_pet(device) | |
| model.load_state_dict(load_state_dict(args.weights), strict=True) | |
| if args.plain: | |
| points = predict_points(model, image_bgr, device, args.thrs) | |
| else: | |
| points = predict_points_tiled(model, image_bgr, device, args.upscale, args.tile, args.overlap, args.thrs) | |
| count = int(points.shape[0]) | |
| print(f"estimated_count: {count}") | |
| output_path = Path(args.output) if args.output else REPO_ROOT / "outputs" / f"{image_path.stem}_counted.jpg" | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| logo_path = REPO_ROOT / "assets" / "logo_white.png" | |
| logo = Image.open(logo_path).convert("RGBA") if logo_path.exists() else None | |
| render_overlay(image_bgr, points, count, logo=logo).save(str(output_path)) | |
| print(f"overlay_saved: {output_path}") | |
| if args.json: | |
| json_path = Path(args.json) | |
| json_path.parent.mkdir(parents=True, exist_ok=True) | |
| json_path.write_text(json.dumps({"image": str(image_path), "count": count, "points_xy": points.tolist()}, indent=2)) | |
| print(f"json_saved: {json_path}") | |
| if __name__ == "__main__": | |
| main() | |