control-r1-trainer / train.py
davanstrien's picture
davanstrien HF Staff
Upload train.py with huggingface_hub
1d5e038 verified
Raw
History Blame Contribute Delete
14 kB
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()