pi-r2-assets / code /train_detr.py
davanstrien's picture
davanstrien HF Staff
Upload /code/train_detr.py with huggingface_hub
eb14f19 verified
Raw
History Blame Contribute Delete
28.6 kB
#!/usr/bin/env python3
"""
Fine-tune DETR (facebook/detr-resnet-50, Apache-2.0) on biglam/loc_beyond_words
(CC0, 7-class document layout object detection: Photograph, Illustration, Map,
Comics/Cartoon, Editorial Cartoon, Headline, Advertisement).
Features:
- Lazy per-sample transforms (resize keeping aspect ratio + random hflip)
- Per-batch padding collator using pixel_mask (DETR supports arbitrary sizes)
- WeightedRandomSampler to oversample images containing rare classes
- COCO mAP evaluation (pycocotools) on the validation split at fixed checkpoints
- Keeps the best model by mAP, then pushes model + processor + metrics to the Hub
"""
import argparse
import json
import os
import random
import time
import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, WeightedRandomSampler
from torchvision.ops import nms
from PIL import Image
from datasets import load_dataset
from transformers import (
AutoImageProcessor,
DetrForObjectDetection,
Trainer,
TrainerCallback,
TrainingArguments,
set_seed,
)
from pycocotools.coco import COCO
from pycocotools.cocoeval import COCOeval
CATEGORIES = [
"Photograph", "Illustration", "Map", "Comics/Cartoon",
"Editorial Cartoon", "Headline", "Advertisement",
]
RARE_CLASS_TARGET_FRAC = 0.3 # target fraction of images containing a class
# ----------------------------------------------------------------------------
# image transforms
# ----------------------------------------------------------------------------
def _resize_long(img: torch.Tensor, long_target: int) -> torch.Tensor:
"""Resize a 3xHxW float image in [0,1] so that its long side == long_target."""
h, w = img.shape[-2:]
scale = long_target / max(h, w)
new_h, new_w = int(round(h * scale)), int(round(w * scale))
return F.interpolate(
img.unsqueeze(0), size=(new_h, new_w), mode="bilinear", antialias=True
).squeeze(0)
def make_train_transform(image_processor, max_size):
mean = torch.tensor(image_processor.image_mean).view(3, 1, 1)
std = torch.tensor(image_processor.image_std).view(3, 1, 1)
def transform(example):
img = example["image"].convert("RGB")
W, H = img.size # width, height
long_target = random.randint(int(max_size * 0.7), max_size)
img_t = (
torch.from_numpy(np.asarray(img, dtype=np.float32))
.permute(2, 0, 1)
.div(255.0)
.clamp(0.0, 1.0)
)
img_t = _resize_long(img_t, long_target)
h, w = img_t.shape[-2:]
sx, sy = w / W, h / H
objs = example["objects"] # list of per-object dicts
boxes = [o["bbox"] for o in objs]
cats = [o["category_id"] for o in objs]
if len(boxes) > 0:
b = torch.tensor(boxes, dtype=torch.float32).clone().reshape(-1, 4)
b[:, [0, 2]] *= sx
b[:, [1, 3]] *= sy
if random.random() < 0.3: # horizontal flip
img_t = img_t.flip(-1)
b[:, 0] = w - b[:, 0] - b[:, 2]
cx = (b[:, 0] + b[:, 2] / 2) / w
cy = (b[:, 1] + b[:, 3] / 2) / h
bw = b[:, 2] / w
bh = b[:, 3] / h
box_t = torch.stack([cx, cy, bw, bh], dim=1).clamp(0.0, 1.0)
cls_t = torch.tensor(cats, dtype=torch.int64).clone()
keep = (bw > 0.0) & (bh > 0.0)
box_t, cls_t = box_t[keep], cls_t[keep]
else:
box_t = torch.zeros(0, 4)
cls_t = torch.zeros(0, dtype=torch.int64)
img_n = (img_t - mean) / std
return {
"pixel_values": img_n,
"pixel_mask": torch.ones(h, w),
"labels": {
"class_labels": cls_t,
"boxes": box_t,
"orig_size": torch.tensor([H, W]), # h, w of original
},
}
return transform
def make_val_transform(image_processor, max_size):
mean = torch.tensor(image_processor.image_mean).view(3, 1, 1)
std = torch.tensor(image_processor.image_std).view(3, 1, 1)
def transform(example):
img = example["image"].convert("RGB")
W, H = img.size
img_t = (
torch.from_numpy(np.asarray(img, dtype=np.float32))
.permute(2, 0, 1)
.div(255.0)
.clamp(0.0, 1.0)
)
img_t = _resize_long(img_t, max_size)
h, w = img_t.shape[-2:]
sx, sy = w / W, h / H
objs = example["objects"] # list of per-object dicts
boxes = [o["bbox"] for o in objs]
cats = [o["category_id"] for o in objs]
if len(boxes) > 0:
b = torch.tensor(boxes, dtype=torch.float32).clone().reshape(-1, 4)
b[:, [0, 2]] *= sx
b[:, [1, 3]] *= sy
cx = (b[:, 0] + b[:, 2] / 2) / w
cy = (b[:, 1] + b[:, 3] / 2) / h
bw = b[:, 2] / w
bh = b[:, 3] / h
box_t = torch.stack([cx, cy, bw, bh], dim=1).clamp(0.0, 1.0)
cls_t = torch.tensor(cats, dtype=torch.int64).clone()
keep = (bw > 0.0) & (bh > 0.0)
box_t, cls_t = box_t[keep], cls_t[keep]
else:
box_t = torch.zeros(0, 4)
cls_t = torch.zeros(0, dtype=torch.int64)
img_n = (img_t - mean) / std
return {
"pixel_values": img_n,
"pixel_mask": torch.ones(h, w),
"labels": {
"class_labels": cls_t,
"boxes": box_t,
"orig_size": torch.tensor([H, W]),
},
}
return transform
class DetrCollator:
"""Apply a per-row transform to each raw example, then pad pixel_values /
pixel_mask to the max dims in the batch. Boxes stay normalized in [0,1] so
they are unaffected by padding. Transforming here (instead of via
`with_transform`/`set_transform`) is robust across `datasets` versions:
newer versions apply dataset transforms to whole batches rather than rows."""
def __init__(self, transform=None):
self.transform = transform
def __call__(self, batch):
if self.transform is not None:
batch = [self.transform(x) for x in batch]
imgs = [x["pixel_values"] for x in batch]
max_h = max(i.shape[-2] for i in imgs)
max_w = max(i.shape[-1] for i in imgs)
pixel_values, pixel_mask = [], []
for i, x in zip(imgs, batch):
h, w = i.shape[-2:]
pad_h, pad_w = max_h - h, max_w - w
pixel_values.append(F.pad(i, (0, pad_w, 0, pad_h), value=0.0))
pixel_mask.append(F.pad(x["pixel_mask"], (0, pad_w, 0, pad_h), value=0.0))
return {
"pixel_values": torch.stack(pixel_values),
"pixel_mask": torch.stack(pixel_mask),
"labels": [x["labels"] for x in batch],
}
# ----------------------------------------------------------------------------
# train dataset weights (oversample images containing rare classes)
# ----------------------------------------------------------------------------
def compute_weights(ds):
"""Per-image sampling weight that boosts images containing rare classes."""
n = len(ds)
counts = np.zeros(7, dtype=np.float64)
img_classes = []
for ex in ds:
cats = set(o["category_id"] for o in ex["objects"])
img_classes.append(cats)
for c in cats:
counts[c] += 1.0
frac = counts / n
w_c = np.zeros(7)
for c in range(7):
if frac[c] < RARE_CLASS_TARGET_FRAC:
w_c[c] = RARE_CLASS_TARGET_FRAC / frac[c] - 1.0
weights = np.array([1.0 + sum(w_c[c] for c in cc) for cc in img_classes])
weights = np.maximum(weights, 1e-3)
return weights, frac
# ----------------------------------------------------------------------------
# COCO evaluation
# ----------------------------------------------------------------------------
def build_coco_gt(val_ds_raw, limit=None):
imgs, anns = [], []
ann_id = 1
for i, ex in enumerate(val_ds_raw):
if limit is not None and i >= limit:
break
img_id = int(ex["image_id"])
# use the actual decoded image size so GT coordinates match the pixel space
# in which predictions are made (robust to any metadata inconsistencies)
W, H = ex["image"].size
imgs.append({"id": img_id, "width": W, "height": H, "file_name": f"{img_id}.jpg"})
objs = ex["objects"] # list of per-object dicts
boxes = [o["bbox"] for o in objs]
cats = [o["category_id"] for o in objs]
iscrowd = [o.get("iscrowd", False) for o in objs]
for b, c, ic in zip(boxes, cats, iscrowd):
anns.append({
"id": ann_id,
"image_id": img_id,
"category_id": int(c),
"bbox": [float(v) for v in b],
"area": float(b[2] * b[3]),
"iscrowd": 0 if not ic else 1,
})
ann_id += 1
gt = {
"images": imgs,
"annotations": anns,
"categories": [{"id": i, "name": CATEGORIES[i]} for i in range(7)],
}
coco_gt = COCO()
coco_gt.dataset = gt
coco_gt.createIndex()
return coco_gt
@torch.no_grad()
def evaluate(model, image_processor, eval_ds, coco_gt, device, max_dets=300,
nms_thr=0.75, limit=None, inference_steps=0):
model.eval()
collator = DetrCollator(make_val_transform(image_processor, max_size=1200))
loader = DataLoader(
eval_ds, batch_size=2, shuffle=False,
collate_fn=collator, num_workers=2, pin_memory=False,
)
preds = []
start = time.time()
step = 0
for batch in loader:
if limit is not None and step >= limit:
break
step += 1
pixel_values = batch["pixel_values"].to(device)
pixel_mask = batch["pixel_mask"].to(device)
with torch.cuda.amp.autocast(enabled=torch.cuda.is_available(), dtype=torch.float16):
out = model(pixel_values=pixel_values, pixel_mask=pixel_mask)
logits = out.logits.float()
boxes = out.pred_boxes.float()
for bi in range(len(batch["labels"])):
label = batch["labels"][bi]
oh, ow = int(label["orig_size"][0]), int(label["orig_size"][1])
img_id = None
pred_logits = logits[bi] # [Nq, 8]
pred_boxes = boxes[bi] # [Nq, 4] cxcywh normalized
scores, cls = pred_logits.softmax(-1)[:, :-1].max(-1)
keep_ix = scores > 0.01
scores, cls, pred_boxes = scores[keep_ix], cls[keep_ix], pred_boxes[keep_ix]
if len(scores) == 0:
continue
cx, cy, bw, bh = pred_boxes.unbind(-1)
x1 = (cx - bw / 2) * ow
y1 = (cy - bh / 2) * oh
x2 = (cx + bw / 2) * ow
y2 = (cy + bh / 2) * oh
xyxy = torch.stack([x1, y1, x2, y2], dim=-1)
keep = nms(xyxy, scores, nms_thr)
keep = keep[:max_dets]
x2c = xyxy[:, 2].clamp(max=ow)
y2c = xyxy[:, 3].clamp(max=oh)
xt = torch.stack([xyxy[:, 0].clamp(min=0), xyxy[:, 1].clamp(min=0), x2c, y2c], dim=-1)
for k in keep.tolist():
b = xt[k].tolist()
preds.append({
"image_id": int(image_ids[bi]),
"category_id": int(cls[k].item()),
"score": float(scores[k].item()),
"bbox": [float(b[0]), float(b[1]), float(b[2] - b[0]), float(b[3] - b[1])],
})
if step % 25 == 0:
elapsed = time.time() - start
print(f" [eval] step {step}/{min(len(loader), limit) if limit else len(loader)} "
f"({elapsed:.0f}s)", flush=True)
if len(preds) == 0:
return None, preds
coco_dt = coco_gt.loadRes(preds)
coco_eval = COCOeval(coco_gt, coco_dt, "bbox")
coco_eval.params.maxDets = [10, 100, 300]
coco_eval.evaluate()
coco_eval.accumulate()
coco_eval.summarize()
stats = coco_eval.stats
# per-class AP50 = precision at IoU=0.50, area=all, maxDets index for 100
prec = coco_eval.eval["precision"] # [T=10, R=101, K=7, A=4, M=3]
iou50 = list(coco_eval.params.iouThrs).index(0.5)
per_class_ap50 = {}
aps = []
for c in range(7):
p_v = prec[iou50, :, c, 0, 2] # IoU=0.50, area=all, maxDets=100, over recalls
p = float(p_v[p_v >= 0].mean()) if (p_v >= 0).any() else 0.0
per_class_ap50[CATEGORIES[c]] = round(p, 4)
aps.append(p)
metrics = {
"mAP@[.5:.95]": float(stats[0]),
"mAP@.50": float(stats[1]),
"mAP@.75": float(stats[2]),
"AP_small": float(stats[3]),
"AP_medium": float(stats[4]),
"AP_large": float(stats[5]),
"AR@100": float(stats[8]),
"AR@300": float(stats[9]) if len(stats) > 9 else float(stats[8]),
"per_class_AP50": per_class_ap50,
"mean_per_class_AP50": round(float(np.mean(aps)), 4),
}
print("EVAL_METRICS " + json.dumps({k: v for k, v in metrics.items() if k != "per_class_AP50"}), flush=True)
print("PER_CLASS_AP50 " + json.dumps(per_class_ap50), flush=True)
return metrics, preds
# ----------------------------------------------------------------------------
# Trainer with weighted sampler + budget/eval callback
# ----------------------------------------------------------------------------
class WeightedTrainer(Trainer):
def __init__(self, *args, weights=None, **kwargs):
super().__init__(*args, **kwargs)
self.train_weights = weights
def get_train_dataloader(self):
ds = self.train_dataset
sampler = WeightedRandomSampler(
torch.as_tensor(self.train_weights, dtype=torch.double),
num_samples=len(ds), replacement=True,
)
return DataLoader(
ds, batch_size=self.args.train_batch_size, sampler=sampler,
collate_fn=self.data_collator, drop_last=False,
num_workers=self.args.dataloader_num_workers,
pin_memory=self.args.dataloader_pin_memory,
)
class EvalAndBudgetCallback(TrainerCallback):
def __init__(self, eval_fn, eval_steps, budget_seconds, save_dir):
self.eval_fn = eval_fn
self.eval_steps = set(eval_steps)
self.budget_seconds = budget_seconds
self.start_time = time.time()
self.save_dir = save_dir
self.best_metric = -1.0
self.best_step = -1
self.results = {}
def on_step_end(self, args, state, control, **kwargs):
step = state.global_step
if step in self.eval_steps:
print(f"\n===== EVAL at step {step} =====", flush=True)
metrics, _ = self.eval_fn()
if metrics is not None and metrics["mAP@[.5:.95]"] > self.best_metric:
self.best_metric = metrics["mAP@[.5:.95]"]
self.best_step = step
model = kwargs.get("model")
if model is not None and model is not getattr(self, "_no_model", None):
torch.save(model.state_dict(), os.path.join(self.save_dir, "best_model.pt"))
self.results[step] = metrics
with open(os.path.join(self.save_dir, "eval_results.json"), "w") as f:
json.dump({"best_step": self.best_step, "best_mAP": self.best_metric,
"results": {str(k): v for k, v in self.results.items()}}, f, indent=2)
control.should_log = True # flush log
elapsed = time.time() - self.start_time
if elapsed > self.budget_seconds:
print(f"BUDGET_REACHED: stopping training at step {step} after {elapsed/60:.1f} min", flush=True)
control.should_training_stop = True
# ----------------------------------------------------------------------------
# main
# ----------------------------------------------------------------------------
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--base-model", default="facebook/detr-resnet-50")
ap.add_argument("--max-steps", type=int, default=20000)
ap.add_argument("--batch-size", type=int, default=4)
ap.add_argument("--lr", type=float, default=5e-5)
ap.add_argument("--wd", type=float, default=1e-4)
ap.add_argument("--warmup-ratio", type=float, default=0.05)
ap.add_argument("--max-size", type=int, default=1200, help="long-edge cap for images")
ap.add_argument("--eval-steps", type=str, default="5000,10000,15000,20000")
ap.add_argument("--budget-minutes", type=float, default=170.0)
ap.add_argument("--limit-train", type=int, default=None)
ap.add_argument("--limit-val", type=int, default=None)
ap.add_argument("--eval-every-steps", type=int, default=0, help="eval every N steps (overrides eval-steps)")
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--repo", default="harness-race/pi-r2")
ap.add_argument("--push", action="store_true")
ap.add_argument("--img-ids-first", type=int, default=50, help="img ids to use when --limit-val w/o raw mapping")
args = ap.parse_args()
set_seed(args.seed)
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"device={device} torch={torch.__version__}", flush=True)
if device == "cuda":
print(f"gpu={torch.cuda.get_device_name(0)}", flush=True)
workdir = os.environ.get("WORKDIR", "/work")
os.makedirs(workdir, exist_ok=True)
# ---- data ----
t0 = time.time()
print("loading dataset...", flush=True)
train_ds = load_dataset("biglam/loc_beyond_words", split="train", trust_remote_code=False)
val_ds_raw = load_dataset("biglam/loc_beyond_words", split="validation", trust_remote_code=False)
if args.limit_train:
train_ds = train_ds.select(range(min(args.limit_train, len(train_ds))))
if args.limit_val:
val_ds_raw = val_ds_raw.select(range(min(args.limit_val, len(val_ds_raw))))
print(f"train={len(train_ds)} val={len(val_ds_raw)} loaded in {time.time()-t0:.0f}s", flush=True)
weights, class_frac = compute_weights(train_ds)
print("class presence fraction (train): " + json.dumps(
{CATEGORIES[i]: round(float(f), 3) for i, f in enumerate(class_frac)}), flush=True)
print(f"mean sampling weight: {weights.mean():.2f} (min {weights.min():.2f})", flush=True)
# ---- model ----
print("loading model + processor...", flush=True)
image_processor = AutoImageProcessor.from_pretrained(args.base_model)
id2label = {i: CATEGORIES[i] for i in range(7)}
label2id = {v: k for k, v in id2label.items()}
model = DetrForObjectDetection.from_pretrained(
args.base_model, id2label=id2label, label2id=label2id,
ignore_mismatched_sizes=True,
)
print(f"model params: {sum(p.numel() for p in model.parameters())/1e6:.1f}M", flush=True)
collator = DetrCollator(make_train_transform(image_processor, args.max_size))
train_ds.weights = weights
# global image_id list for eval
global image_ids
image_ids = [int(ex["image_id"]) for ex in val_ds_raw]
print("sample image sizes (val):", ", ".join(
f"{ex['image'].size}" for ex in val_ds_raw.select([0, 1, 2])), flush=True)
coco_gt = build_coco_gt(val_ds_raw, limit=None)
eval_budget = args.budget_minutes * 60
eval_steps = [int(s) for s in args.eval_steps.split(",")] if args.eval_steps else []
if args.eval_every_steps and args.eval_every_steps > 0:
eval_steps = list(range(args.eval_every_steps - 1, args.max_steps + 1, args.eval_every_steps))
def eval_fn():
try:
return evaluate(model, image_processor, val_ds_raw, coco_gt, device,
limit=args.limit_val)
except Exception as e:
print(f"EVAL FAILED: {e}", flush=True)
import traceback; traceback.print_exc()
return None, None
training_args = TrainingArguments(
output_dir=os.path.join(workdir, "out"),
per_device_train_batch_size=args.batch_size,
learning_rate=args.lr,
weight_decay=args.wd,
max_steps=args.max_steps,
lr_scheduler_type="cosine",
warmup_ratio=args.warmup_ratio,
fp16=torch.cuda.is_available(),
max_grad_norm=0.1,
dataloader_num_workers=2,
dataloader_pin_memory=False,
remove_unused_columns=False,
logging_steps=25,
save_strategy="no",
report_to=[],
seed=args.seed,
data_seed=args.seed,
push_to_hub=False,
)
callback = EvalAndBudgetCallback(eval_fn, eval_steps, eval_budget, workdir)
trainer = WeightedTrainer(
model=model,
args=training_args,
train_dataset=train_ds,
data_collator=collator,
callbacks=[callback],
weights=weights,
)
print(f"training up to {args.max_steps} steps, budget {args.budget_minutes} min...", flush=True)
t0 = time.time()
trainer.train()
elapsed_min = (time.time() - t0) / 60
final_step = trainer.state.global_step
# final eval (use best checkpoint if we have one)
if os.path.exists(os.path.join(workdir, "best_model.pt")):
state_dict = torch.load(os.path.join(workdir, "best_model.pt"), map_location="cpu")
model.load_state_dict(state_dict)
print(f"loaded best checkpoint (step {callback.best_step}, mAP {callback.best_metric:.4f})", flush=True)
print("final evaluation on validation split...", flush=True)
metrics, preds = evaluate(model, image_processor, val_ds_raw, coco_gt, device, limit=None)
if metrics is None:
print("NO_METRICS", flush=True)
metrics = {}
metrics["trained_steps"] = final_step
metrics["train_minutes"] = round(elapsed_min, 2)
metrics["max_size"] = args.max_size
metrics["batch_size"] = args.batch_size
metrics["lr"] = args.lr
metrics["wd"] = args.wd
metrics["num_val"] = len(val_ds_raw)
metrics["base_model"] = args.base_model
metrics["dataset"] = "biglam/loc_beyond_words (validation split, 712 images)"
with open(os.path.join(workdir, "metrics.json"), "w") as f:
json.dump(metrics, f, indent=2)
print("FINAL_METRICS " + json.dumps(metrics), flush=True)
if preds:
with open(os.path.join(workdir, "eval_predictions.json"), "w") as f:
json.dump(preds, f)
# ---- push to hub ----
if args.push:
from huggingface_hub import HfApi
outdir = os.path.join(workdir, "hub")
os.makedirs(outdir, exist_ok=True)
model.save_pretrained(outdir)
image_processor.save_pretrained(outdir)
with open(os.path.join(outdir, "metrics.json"), "w") as f:
json.dump(metrics, f, indent=2)
if preds:
with open(os.path.join(outdir, "eval_predictions.json"), "w") as f:
json.dump(preds, f)
import shutil
shutil.copy(os.path.abspath(__file__), os.path.join(outdir, "train_detr.py"))
card = make_model_card(metrics, args.repo)
with open(os.path.join(outdir, "README.md"), "w") as f:
f.write(card)
print(f"pushing to {args.repo}...", flush=True)
api = HfApi()
api.create_repo(args.repo, repo_type="model", exist_ok=True)
api.upload_folder(
repo_id=args.repo,
folder_path=outdir,
commit_message="Fine-tune DETR on biglam/loc_beyond_words (7-class document layout detection)",
)
print("push complete", flush=True)
else:
print("(--push not set; skipping hub push)", flush=True)
model.save_pretrained(os.path.join(workdir, "model"))
print("DONE", flush=True)
def make_model_card(metrics, repo_id):
"""Build the README.md model card."""
per_class = metrics.get("per_class_ap50", {})
rows = "".join(
f"| {c} | {per_class.get(c, '-')} |" for c in CATEGORIES
)
mAP = metrics.get("mAP@[.5:.95]", None)
ap50 = metrics.get("mAP@.50", None)
ap75 = metrics.get("mAP@.75", None)
ar100 = metrics.get("AR@100", None)
card = f"""---
language:
- en
license: apache-2.0
base_model: facebook/detr-resnet-50
tags:
- object-detection
- document-layout-analysis
- transformers
pipeline_tag: object-detection
datasets:
- biglam/loc_beyond_words
library_name: transformers
model-index:
- name: pi-r2-detr-beyond-words
results:
- task:
type: object-detection
name: Object Detection
dataset:
type: biglam/loc_beyond_words
name: Beyond Words (LOC) validation
split: validation
metrics:
- type: Average Precision
value: {mAP if mAP is not None else 'N/A'}
name: mAP (COCO, IoU 0.5:0.95)
---
# pi-r2 — DETR fine-tuned on Beyond Words (LOC)
Object detection model fine-tuned from [`facebook/detr-resnet-50`](https://huggingface.co/facebook/detr-resnet-50)
(Apache-2.0) on the [`biglam/loc_beyond_words`](https://huggingface.co/datasets/biglam/loc_beyond_words)
dataset (CC0): crowd-sourced bounding-box annotations of **World War I-era newspaper pages**
from the Library of Congress Chronicling America collection.
## Model detail
- **Architecture**: DETR (DEtection TRansformer) with a ResNet-50 backbone, 6 encoder/6 decoder
transformer layers, 100 object queries.
- **Base model**: `facebook/detr-resnet-50` — license **Apache-2.0** (shareable).
- **Dataset**: `biglam/loc_beyond_words` — license **CC0-1.0** (public domain).
- **Classes** (7): {', '.join(CATEGORIES)}
- **Image size**: resized so the longest edge ≤ {metrics.get('max_size', 1200)} px (aspect ratio preserved),
padded per batch via a `pixel_mask`. Training used random horizontal flips and random scale
(70–100% of the max size).
- **Optimizer**: AdamW (LR {metrics.get('lr', 5e-5)}, weight decay {metrics.get('wd', 1e-4)}),
cosine schedule with 5% warmup, fp16, gradient clipping 0.1.
- **Class imbalance**: images containing rare classes (Map, Editorial Cartoon, Illustration, Comics)
are oversampled with a `WeightedRandomSampler` when building training batches.
- **Training budget**: {metrics.get('train_minutes', 0)} minutes, {metrics.get('trained_steps', 0)} steps,
batch size {metrics.get('batch_size', 4)} on an NVIDIA A10G.
## Validation results (COCO protocol, pycocotools)
Evaluated on the `biglam/loc_beyond_words` **validation** split ({metrics.get('num_val', 712)} images)
with COCO-style IoU-matched metrics (area = all, max detections = 300, NMS IoU threshold 0.75).
| Metric | Value |
|---|---|
| mAP @[0.5:0.95] | {mAP if mAP is not None else 'N/A'} |
| mAP @0.50 | {ap50 if ap50 is not None else 'N/A'} |
| mAP @0.75 | {ap75 if ap75 is not None else 'N/A'} |
| AR @100 | {ar100 if ar100 is not None else 'N/A'} |
Per-class AP @0.50:
| Class | AP@0.50 |
|---|---|
{rows}
> Raw per-image predictions and this training script are included in this repo
> (`eval_predictions.json`, `train_detr.py`).
## Usage
```python
from transformers import AutoImageProcessor, DetrForObjectDetection
from PIL import Image
import torch
repo = "{repo_id}"
processor = AutoImageProcessor.from_pretrained(repo)
model = DetrForObjectDetection.from_pretrained(repo)
img = Image.open("newspaper_page.jpg").convert("RGB")
inputs = processor(images=img, return_tensors="pt")
with torch.no_grad():
out = model(**inputs)
score_threshold = 0.5
for logits, box in zip(out.logits[0], out.pred_boxes[0]):
prob = logits.softmax(-1)
cls_idx, score = prob[:, :-1].max(-1)
if score.item() > score_threshold:
cx, cy, bw, bh = box.tolist()
W, H = img.size
x1, y1 = (cx - bw/2)*W, (cy - bh/2)*H
x2, y2 = (cx + bw/2)*W, (cy + bh/2)*H
print(model.config.id2label[cls_idx.item()], round(score.item(), 3), [round(x1), round(y1), round(x2), round(y2)])
```
## Intended use & limitations
- Trained on historically scanned newspaper pages (circa 1910–1920); other domains or modern
document layouts will degrade accuracy.
- The dataset is strongly class-imbalanced (Headline/Advertisement dominate; Map and Editorial
Cartoon are rare), so per-class accuracy varies widely (see table above).
- DETR emits up to 100 box proposals per image; in very dense pages some objects may be missed.
## Licenses
- Base model `facebook/detr-resnet-50`: **Apache-2.0**.
- Dataset `biglam/loc_beyond_words`: **CC0-1.0** (Public Domain Dedication).
- This fine-tuned model: **Apache-2.0**.
"""
return card
if __name__ == "__main__":
main()