CrowdCounter / count.py
Bobby
Cache Crowd candidates per session and shorten small-image GPU reservations
cc83161
Raw History Blame Contribute Delete
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
@torch.no_grad()
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]
@torch.no_grad()
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)
@torch.no_grad()
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()