import argparse, os, json, time, random import numpy as np import torch from torch.utils.data import DataLoader from tqdm import tqdm import datasets as hf_datasets from datasets import load_dataset from transformers import DetrImageProcessor, DetrForObjectDetection from PIL import Image def parse_args(): p = argparse.ArgumentParser() p.add_argument('--data_dir', default='/data') p.add_argument('--model_id', default='facebook/detr-resnet-50') p.add_argument('--output_dir', default='/output') p.add_argument('--repo_id', default='harness-race/control-r1') p.add_argument('--epochs', type=int, default=4) p.add_argument('--batch_size', type=int, default=4) p.add_argument('--lr', type=float, default=1e-4) p.add_argument('--shortest', type=int, default=800) p.add_argument('--longest', type=int, default=1333) p.add_argument('--seed', type=int, default=42) p.add_argument('--max_train_size', type=int, default=2846) p.add_argument('--resume_from', default=None) p.add_argument('--push', type=int, default=1) return p.parse_args() CLASS_NAMES = ["Photograph","Illustration","Map","Comics/Cartoon","Editorial Cartoon","Headline","Advertisement"] def set_seed(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) try: torch.cuda.manual_seed_all(s) except: pass def main(): args = parse_args() set_seed(args.seed) device = 'cuda' if torch.cuda.is_available() else 'cpu' print(f"[train] device={device} cuda_name={torch.cuda.get_device_name(0) if device=='cuda' else 'n/a'}", flush=True) id2label = {i: nm for i, nm in enumerate(CLASS_NAMES)} label2id = {v: k for k, v in id2label.items()} processor = DetrImageProcessor.from_pretrained(args.model_id) print("[train] loading dataset", flush=True) ds = None if args.data_dir and os.path.isdir(args.data_dir) and os.path.isdir(os.path.join(args.data_dir, 'data')): try: import glob files = sorted(glob.glob(os.path.join(args.data_dir, 'data', '*.parquet'))) train_files = [f for f in files if 'train-' in os.path.basename(f)] val_files = [f for f in files if 'validation-' in os.path.basename(f)] files_dict = {} if train_files: files_dict['train'] = train_files if val_files: files_dict['validation'] = val_files print("[train] loading parquet:", files_dict, flush=True) ds = hf_datasets.load_dataset("parquet", data_files=files_dict) # verify image decoding works _ = ds['train'][0]['image'] print("[train] parquet image decoding OK", flush=True) except Exception as e: print(f"[train] parquet load failed ({e}); falling back to hub load", flush=True) ds = None if ds is None: ds = load_dataset("biglam/loc_beyond_words") train_ds = ds['train'] val_ds = ds['validation'] # limit train size for quick runs if args.max_train_size and args.max_train_size < len(train_ds): train_ds = train_ds.select(range(args.max_train_size)) print(f"[train] train={len(train_ds)} val={len(val_ds)}", flush=True) class DetrDataset(torch.utils.data.Dataset): def __init__(self, hf_ds, processor): self.ds = hf_ds self.processor = processor def __len__(self): return len(self.ds) def __getitem__(self, i): ex = self.ds[i] image = ex['image'].convert('RGB') objs = ex['objects'] alist = [{'image_id': int(o['image_id']), 'bbox': o['bbox'], 'category_id': o['category_id'], 'area': o['area'], 'iscrowd': int(o['iscrowd'])} for o in objs] anns = {'image_id': int(ex['image_id']), 'annotations': alist} enc = self.processor(images=image, annotations=anns, return_tensors='pt') return {'pixel_values': enc['pixel_values'][0], 'pixel_mask': enc['pixel_mask'][0], 'labels': enc['labels'][0], 'image_id': int(ex['image_id'])} def collate(batch): pvs = [b['pixel_values'] for b in batch] pms = [b['pixel_mask'] for b in batch] bs = len(pvs) max_h = max(pv.shape[1] for pv in pvs) max_w = max(pv.shape[2] for pv in pvs) out_pv = torch.zeros((bs, 3, max_h, max_w), dtype=pvs[0].dtype) out_pm = torch.zeros((bs, max_h, max_w), dtype=torch.int64) for i in range(bs): h, w = pvs[i].shape[1], pvs[i].shape[2] out_pv[i, :, :h, :w] = pvs[i] out_pm[i, :h, :w] = pms[i] labels = [b['labels'] for b in batch] return {'pixel_values': out_pv, 'pixel_mask': out_pm, 'labels': labels, 'img_id': [b['image_id'] for b in batch]} train_dl = DataLoader(DetrDataset(train_ds, processor), batch_size=args.batch_size, shuffle=True, collate_fn=collate, num_workers=4, pin_memory=True, drop_last=False) val_dl = DataLoader(DetrDataset(val_ds, processor), batch_size=args.batch_size, shuffle=False, collate_fn=collate, num_workers=4, pin_memory=True) model = DetrForObjectDetection.from_pretrained(args.model_id, ignore_mismatched_sizes=True, num_labels=len(CLASS_NAMES), id2label=id2label, label2id=label2id) model.to(device) no_decay = ['bias','LayerNorm.weight','layernorm.weight'] pnames = [n for n,_ in model.named_parameters()] opt_groups = [ {'params': [p for n,p in model.named_parameters() if not any(nd in n for nd in no_decay)], 'weight_decay': 1e-4}, {'params': [p for n,p in model.named_parameters() if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}, ] optimizer = torch.optim.AdamW(opt_groups, lr=args.lr) total_steps = args.epochs * len(train_dl) warmup = int(0.05*total_steps) def lr_lambda(step): if step < warmup: return step/max(1,warmup) progress = (step-warmup)/max(1,total_steps-warmup) return 0.5*(1+np.cos(np.pi*progress)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) use_amp = (device == 'cuda') if device == 'cuda' and torch.cuda.is_bf16_supported(): amp_dtype = torch.bfloat16 use_scaler = False elif device == 'cuda': amp_dtype = torch.float16 use_scaler = True else: amp_dtype = torch.float32 use_scaler = False scalar = torch.cuda.amp.GradScaler(enabled=use_scaler) print(f"[train] AMP dtype={amp_dtype} scaler={use_scaler}", flush=True) grad_accum = 2 start_epoch = 0 if args.resume_from: ck = torch.load(args.resume_from, map_location=device) model.load_state_dict(ck['model']) optimizer.load_state_dict(ck['optim']) start_epoch = ck['epoch'] print('[train] resumed epoch', start_epoch, flush=True) model.train() step = start_epoch*len(train_dl) print(f"[train] total_steps={total_steps} warmup={warmup} per_epoch={len(train_dl)}", flush=True) metric_dir = os.path.join(args.output_dir, 'metrics') os.makedirs(metric_dir, exist_ok=True) for epoch in range(start_epoch, args.epochs): ep_start = time.time() tot_loss = 0.0 tpbar = tqdm(train_dl, desc=f'epoch {epoch+1}/{args.epochs}') for bi, batch in enumerate(tpbar): pv = batch['pixel_values'].to(device) pm = batch['pixel_mask'].to(device) tgt = [{k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in t.items()} for t in batch['labels']] with torch.autocast(device_type='cuda', dtype=amp_dtype, enabled=use_amp): out = model(pixel_values=pv, pixel_mask=pm, labels=tgt) loss = out.loss / grad_accum if use_scaler: scalar.scale(loss).backward() else: loss.backward() if (bi+1) % grad_accum == 0: if use_scaler: scalar.step(optimizer); scalar.update() else: optimizer.step() optimizer.zero_grad() scheduler.step(); step += 1 loss_it = loss.item()*grad_accum if torch.isfinite(loss) else float('nan') tot_loss += loss_it tpbar.set_postfix(loss=round(loss_it,3)) torch.cuda.synchronize() if device=='cuda' else None ep_loss = tot_loss/len(train_dl) print(f"[train] epoch {epoch+1} loss={ep_loss:.4f} time={time.time()-ep_start:.1f}s (epoch elapsed {time.time()-ep_start:.1f}s)", flush=True) # save final ck_path = os.path.join(args.output_dir, 'checkpoint.pt') torch.save({'model': model.state_dict(), 'epoch': args.epochs}, ck_path) model.save_pretrained(os.path.join(args.output_dir, 'model')) processor.save_pretrained(os.path.join(args.output_dir, 'model')) print("[train] saved model", flush=True) # ---- evaluation ---- print("[eval] starting", flush=True) from object_eval import evaluate model.eval() metrics = evaluate(model, processor, val_ds, val_dl, device) print("[eval] metrics:", json.dumps(metrics, indent=2), flush=True) with open(os.path.join(args.output_dir, 'metrics.json'), 'w') as f: json.dump(metrics, f, indent=2) # ---- push to hub ---- token = os.environ.get('HF_TOKEN') if args.push and token: try: push_to_repo(os.path.join(args.output_dir, 'model'), args.repo_id, token, metrics, args) except Exception as e: print(f"[push] FAILED: {e}", flush=True) import traceback; traceback.print_exc() else: print("[push] skipped (no token or --push flag)", flush=True) def build_model_card(metrics, args): ma = metrics.get('mean_AP_0.50_0.95', 0.0) ap50 = metrics.get('AP_0.50', 0.0) ap75 = metrics.get('AP_0.75', 0.0) ar100 = metrics.get('AR_max_100', 0.0) per = metrics.get('per_class_AP', {}) per_line = "\n".join(f"- **{k}**: {v:.4f}" for k,v in per.items()) base = args.model_id import time ts = time.strftime("%Y-%m-%d") md = f"""--- language: - en license: apache-2.0 base_model: {base} tags: - object-detection - detr - vision - doc-layout - generated_from_trainer library_name: transformers datasets: - biglam/loc_beyond_words metrics: - name: mean_average_precision (COCO 0.50:0.95) type: mean_average_precision value: {ma:.4f} --- # Control-R1: Doc-layout object detection (fine-tuned DETR) This model is a fine-tuned **DETR (ResNet-50)** object-detection model for **document / newspaper page-layout analysis**. It detects 7 region types in scanned historical newspaper pages and was trained on the `biglam/loc_beyond_words` dataset. ## Base model & license - **Base model**: [`facebook/detr-resnet-50`](https://huggingface.co/facebook/detr-resnet-50) - **License**: Apache-2.0 (open license — free to share, modify and use commercially) - The fine-tuned weights in this repository inherit the **Apache-2.0** license. ## Classes (7) {", ".join(CLASS_NAMES)} ## Training data - **Dataset**: [biglam/loc_beyond_words](https://huggingface.co/datasets/biglam/loc_beyond_words) - **Train**: {args.max_train_size} images - **Validation**: 712 images - COCO-style annotations (xywh bounding boxes). ## Training details - **Model**: `{base}`, all parameters fine-tuned (backbone unfrozen) - **Epochs**: {args.epochs} - **Batch size**: {args.batch_size} (gradient accumulation = 2) - **Optimizer**: AdamW (lr = {args.lr}, weight decay on non-bias/LN params) - **LR schedule**: linear warmup + cosine decay - **Image preprocessing**: resize to shortest edge = {args.shortest}px, longest = {args.longest}px (aspect preserved), ImageNet normalization - **Mixed precision**: bf16/fp16 autocast (GPU) - **Hardware**: Hugging Face Jobs GPU ## Validation results (COCO eval on the 712-image validation split) | Metric | Value | |---|---| | **mAP (IoU 0.50:0.95)** | {ma:.4f} | | AP @ IoU 0.50 | {ap50:.4f} | | AP @ IoU 0.75 | {ap75:.4f} | | AR (max 100 dets) | {ar100:.4f} | Per-class AP (IoU 0.50:0.95): {per_line} ## Quick usage ```python from transformers import AutoModelForObjectDetection, AutoImageProcessor from PIL import Image model = AutoModelForObjectDetection.from_pretrained("harness-race/control-r1") processor = AutoImageProcessor.from_pretrained("harness-race/control-r1") image = Image.open("page.png").convert("RGB") inputs = processor(images=image, return_tensors="pt") outputs = model(**inputs) results = processor.post_process_object_detection(outputs, target_sizes=[(image.height, image.width)], threshold=0.4) for r in results[0]: print(model.config.id2label[int(r['labels'])], round(r['scores'].item(),3) if hasattr(r['scores'],'item') else r['scores'], [round(c,1) for c in r['boxes'].tolist()]) ``` ## Intended use & limitations Trained for research on historical newspaper layout analysis. Best on page layouts similar to the `loc_beyond_words` training distribution; large format/styled pages not seen in training may be missed. Evaluation was done on the dataset's official 712-image validation split; run at {ts} (UTC) on HF Jobs. --- *Control-R1 — a layout-model entry. Trained on Hugging Face Jobs (< \$5 budget).* """ return md def push_to_repo(model_dir, repo_id, token, metrics, args): from huggingface_hub import HfApi api = HfApi() api.create_repo(repo_id=repo_id, repo_type='model', exist_ok=True, token=token) md = build_model_card(metrics, args) with open(os.path.join(model_dir, 'README.md'), 'w') as f: f.write(md) api.upload_folder(folder_path=model_dir, repo_id=repo_id, repo_type='model', commit_message='Fine-tune DETR (ResNet-50) on biglam/loc_beyond_words (control-r1)', token=token) print(f"[push] pushed {repo_id}", flush=True) if __name__ == '__main__': main()