umutonuryasar commited on
Commit
07d6008
·
1 Parent(s): fcc3719

fix: 10 critical bugs + 5 warnings from pre-training audit

Browse files

- train.py: add --resume flag, fix grad accum flush, fix COCO annotations format
- evaluate.py: fix label2cat unhashable dict crash
- api/main.py: replace deprecated on_event with lifespan pattern
- train.py + evaluate.py: migrate to torch.amp (PyTorch 2.x)
- configs: A100 tuning — bf16, batch 16, grad_accum 4, effective batch 64
- notebook: fix threshold kwarg, box dict access, nested config keys, checkpoint glob, eval CLI args

api/main.py CHANGED
@@ -2,6 +2,7 @@ from __future__ import annotations
2
 
3
  import io
4
  import os
 
5
 
6
  from fastapi import FastAPI, File, HTTPException, Query, UploadFile
7
  from PIL import Image, UnidentifiedImageError
@@ -12,8 +13,6 @@ from api.schemas import Detection, BoundingBox, PredictResponse
12
  MODEL_ID = os.getenv("MODEL_ID", "PekingU/rtdetr_r50vd")
13
  CONFIDENCE_THRESHOLD = float(os.getenv("CONFIDENCE_THRESHOLD", "0.5"))
14
 
15
- app = FastAPI(title="detrflow", description="RT-DETR object detection API", version="0.1.0")
16
-
17
  _predictor: RTDetrPredictor | None = None
18
 
19
 
@@ -27,9 +26,13 @@ def get_predictor() -> RTDetrPredictor:
27
  return _predictor
28
 
29
 
30
- @app.on_event("startup")
31
- async def _warm_up() -> None:
32
  get_predictor()
 
 
 
 
33
 
34
 
35
  @app.post("/predict", response_model=PredictResponse, summary="Detect objects in an image")
 
2
 
3
  import io
4
  import os
5
+ from contextlib import asynccontextmanager
6
 
7
  from fastapi import FastAPI, File, HTTPException, Query, UploadFile
8
  from PIL import Image, UnidentifiedImageError
 
13
  MODEL_ID = os.getenv("MODEL_ID", "PekingU/rtdetr_r50vd")
14
  CONFIDENCE_THRESHOLD = float(os.getenv("CONFIDENCE_THRESHOLD", "0.5"))
15
 
 
 
16
  _predictor: RTDetrPredictor | None = None
17
 
18
 
 
26
  return _predictor
27
 
28
 
29
+ @asynccontextmanager
30
+ async def lifespan(app: FastAPI):
31
  get_predictor()
32
+ yield
33
+
34
+
35
+ app = FastAPI(title="detrflow", description="RT-DETR object detection API", version="0.1.0", lifespan=lifespan)
36
 
37
 
38
  @app.post("/predict", response_model=PredictResponse, summary="Detect objects in an image")
configs/rtdetr_r50_coco.yaml CHANGED
@@ -7,12 +7,12 @@ data:
7
  val_ann: data/coco/annotations/instances_val2017.json
8
  train_img: data/coco/train2017
9
  val_img: data/coco/val2017
10
- num_workers: 4
11
 
12
  training:
13
  epochs: 12
14
- batch_size: 8 # fits 8 GB VRAM with fp16 + gradient checkpointing
15
- grad_accum_steps: 2 # effective batch = 16
16
 
17
  optimizer:
18
  name: AdamW
@@ -25,7 +25,8 @@ training:
25
  warmup_epochs: 1
26
  min_lr: 1.0e-6
27
 
28
- fp16: true
 
29
  gradient_checkpointing: true
30
  clip_grad_norm: 0.1
31
 
 
7
  val_ann: data/coco/annotations/instances_val2017.json
8
  train_img: data/coco/train2017
9
  val_img: data/coco/val2017
10
+ num_workers: 8
11
 
12
  training:
13
  epochs: 12
14
+ batch_size: 16 # A100 40 GB with bf16 + gradient checkpointing
15
+ grad_accum_steps: 4 # effective batch = 64 (matches RT-DETR paper)
16
 
17
  optimizer:
18
  name: AdamW
 
25
  warmup_epochs: 1
26
  min_lr: 1.0e-6
27
 
28
+ fp16: false
29
+ bf16: true # A100 native bfloat16 — wider range, no GradScaler needed
30
  gradient_checkpointing: true
31
  clip_grad_norm: 0.1
32
 
notebooks/detrflow_training.ipynb CHANGED
@@ -171,27 +171,7 @@
171
  "id": "show-config",
172
  "metadata": {},
173
  "outputs": [],
174
- "source": [
175
- "import yaml\n",
176
- "\n",
177
- "with open('configs/rtdetr_r50_coco.yaml') as f:\n",
178
- " cfg = yaml.safe_load(f)\n",
179
- "\n",
180
- "# Override output dir to point to Drive\n",
181
- "cfg['output_dir'] = CKPT_DIR\n",
182
- "\n",
183
- "# Colab T4 optimizations\n",
184
- "cfg['batch_size'] = 8 # T4 16 GB → batch 8 safe with fp16\n",
185
- "cfg['grad_accum_steps'] = 2 # effective batch = 16\n",
186
- "cfg['num_epochs'] = 12 # standard RT-DETR schedule\n",
187
- "cfg['fp16'] = True\n",
188
- "\n",
189
- "# Save overridden config\n",
190
- "with open('configs/rtdetr_r50_coco_colab.yaml', 'w') as f:\n",
191
- " yaml.dump(cfg, f)\n",
192
- "\n",
193
- "print(yaml.dump(cfg, default_flow_style=False))"
194
- ]
195
  },
196
  {
197
  "cell_type": "markdown",
@@ -209,32 +189,7 @@
209
  "id": "smoke-test",
210
  "metadata": {},
211
  "outputs": [],
212
- "source": [
213
- "import sys\n",
214
- "sys.path.insert(0, '/content/detrflow')\n",
215
- "\n",
216
- "from inference.predictor import RTDetrPredictor\n",
217
- "from inference.visualizer import draw_detections\n",
218
- "from PIL import Image\n",
219
- "import requests\n",
220
- "from IPython.display import display\n",
221
- "\n",
222
- "# Load sample image\n",
223
- "url = 'http://images.cocodataset.org/val2017/000000039769.jpg'\n",
224
- "img = Image.open(requests.get(url, stream=True).raw)\n",
225
- "\n",
226
- "# Run inference\n",
227
- "predictor = RTDetrPredictor(model_id='PekingU/rtdetr_r50vd', threshold=0.5)\n",
228
- "detections = predictor.predict(img)\n",
229
- "\n",
230
- "print(f'Detections: {len(detections)}')\n",
231
- "for d in detections[:5]:\n",
232
- " print(f\" {d['label']:20s} score={d['score']:.3f} box={[round(x) for x in d['box']]}\")\n",
233
- "\n",
234
- "# Visualize\n",
235
- "vis = draw_detections(img, detections)\n",
236
- "display(vis)"
237
- ]
238
  },
239
  {
240
  "cell_type": "markdown",
@@ -274,24 +229,7 @@
274
  "id": "resume-training",
275
  "metadata": {},
276
  "outputs": [],
277
- "source": [
278
- "import os\n",
279
- "\n",
280
- "# Find the latest checkpoint\n",
281
- "ckpts = sorted([\n",
282
- " f for f in os.listdir(CKPT_DIR)\n",
283
- " if f.startswith('checkpoint-epoch')\n",
284
- "])\n",
285
- "\n",
286
- "if ckpts:\n",
287
- " latest = f'{CKPT_DIR}/{ckpts[-1]}'\n",
288
- " print(f'Resuming from: {latest}')\n",
289
- " !python scripts/train.py \\\n",
290
- " --config configs/rtdetr_r50_coco_colab.yaml \\\n",
291
- " --resume {latest}\n",
292
- "else:\n",
293
- " print('No checkpoint found. Run cell 6 first.')"
294
- ]
295
  },
296
  {
297
  "cell_type": "markdown",
@@ -307,25 +245,7 @@
307
  "id": "run-eval",
308
  "metadata": {},
309
  "outputs": [],
310
- "source": [
311
- "import os\n",
312
- "\n",
313
- "# Pick the latest checkpoint (last epoch)\n",
314
- "ckpts = sorted([\n",
315
- " f for f in os.listdir(CKPT_DIR)\n",
316
- " if f.startswith('checkpoint-epoch')\n",
317
- "])\n",
318
- "\n",
319
- "if not ckpts:\n",
320
- " print('No checkpoint found. Run training first.')\n",
321
- "else:\n",
322
- " best = f'{CKPT_DIR}/{ckpts[-1]}'\n",
323
- " print(f'Evaluating: {best}')\n",
324
- " !python scripts/evaluate.py \\\n",
325
- " --model_path {best} \\\n",
326
- " --coco_dir data/coco \\\n",
327
- " --split val"
328
- ]
329
  },
330
  {
331
  "cell_type": "markdown",
@@ -385,4 +305,4 @@
385
  },
386
  "nbformat": 4,
387
  "nbformat_minor": 5
388
- }
 
171
  "id": "show-config",
172
  "metadata": {},
173
  "outputs": [],
174
+ "source": "import yaml\n\nwith open('configs/rtdetr_r50_coco.yaml') as f:\n cfg = yaml.safe_load(f)\n\n# Override output dir to point to Drive\ncfg['training']['save_dir'] = CKPT_DIR\n\n# Colab T4 optimizations (nested keys matching train.py's cfg structure)\ncfg['training']['batch_size'] = 8 # T4 16 GB → batch 8 safe with fp16\ncfg['training']['grad_accum_steps'] = 2 # effective batch = 16\ncfg['training']['epochs'] = 12 # standard RT-DETR schedule\ncfg['training']['fp16'] = True\n\n# Save overridden config\nwith open('configs/rtdetr_r50_coco_colab.yaml', 'w') as f:\n yaml.dump(cfg, f)\n\nprint(yaml.dump(cfg, default_flow_style=False))"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
175
  },
176
  {
177
  "cell_type": "markdown",
 
189
  "id": "smoke-test",
190
  "metadata": {},
191
  "outputs": [],
192
+ "source": "import sys\nsys.path.insert(0, '/content/detrflow')\n\nfrom inference.predictor import RTDetrPredictor\nfrom inference.visualizer import draw_detections\nfrom PIL import Image\nimport requests\nfrom IPython.display import display\n\n# Load sample image\nurl = 'http://images.cocodataset.org/val2017/000000039769.jpg'\nimg = Image.open(requests.get(url, stream=True).raw)\n\n# Run inference\npredictor = RTDetrPredictor(model_id='PekingU/rtdetr_r50vd', confidence_threshold=0.5)\ndetections = predictor.predict(img)\n\nprint(f'Detections: {len(detections)}')\nfor d in detections[:5]:\n b = d['box']\n print(f\" {d['label']:20s} score={d['score']:.3f} box=[{b['x1']:.0f},{b['y1']:.0f},{b['x2']:.0f},{b['y2']:.0f}]\")\n\n# Visualize\nvis = draw_detections(img, detections)\ndisplay(vis)"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
193
  },
194
  {
195
  "cell_type": "markdown",
 
229
  "id": "resume-training",
230
  "metadata": {},
231
  "outputs": [],
232
+ "source": "import os\n\n# Find the latest checkpoint (train.py saves as epoch_001, epoch_002, ...)\nckpts = sorted([\n f for f in os.listdir(CKPT_DIR)\n if f.startswith('epoch_')\n])\n\nif ckpts:\n latest = f'{CKPT_DIR}/{ckpts[-1]}'\n print(f'Resuming from: {latest}')\n !python scripts/train.py \\\n --config configs/rtdetr_r50_coco_colab.yaml \\\n --resume {latest}\nelse:\n print('No checkpoint found. Run cell 6 first.')"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
233
  },
234
  {
235
  "cell_type": "markdown",
 
245
  "id": "run-eval",
246
  "metadata": {},
247
  "outputs": [],
248
+ "source": "import os\n\n# Pick the latest checkpoint (train.py saves as epoch_001, epoch_002, ...)\nckpts = sorted([\n f for f in os.listdir(CKPT_DIR)\n if f.startswith('epoch_')\n])\n\nif not ckpts:\n print('No checkpoint found. Run training first.')\nelse:\n best = f'{CKPT_DIR}/{ckpts[-1]}'\n print(f'Evaluating: {best}')\n !python scripts/evaluate.py \\\n --config configs/rtdetr_r50_coco_colab.yaml \\\n --checkpoint {best}"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
249
  },
250
  {
251
  "cell_type": "markdown",
 
305
  },
306
  "nbformat": 4,
307
  "nbformat_minor": 5
308
+ }
scripts/evaluate.py CHANGED
@@ -13,7 +13,7 @@ import tempfile
13
 
14
  import torch
15
  import yaml
16
- from torch.cuda.amp import autocast
17
  from torch.utils.data import DataLoader
18
  from transformers import RTDetrForObjectDetection, RTDetrImageProcessor
19
 
@@ -62,12 +62,15 @@ def build_val_loader(img_dir: str, ann_file: str, processor, batch_size: int, nu
62
  @torch.inference_mode()
63
  def evaluate(cfg: dict, checkpoint: str | None = None) -> None:
64
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
65
- use_fp16 = cfg["training"]["fp16"] and device.type == "cuda"
 
 
 
66
 
67
  model_id = checkpoint or cfg["model"]["id"]
68
  processor = RTDetrImageProcessor.from_pretrained(model_id)
69
  model = RTDetrForObjectDetection.from_pretrained(
70
- model_id, torch_dtype=torch.float16 if use_fp16 else torch.float32
71
  ).to(device)
72
  model.eval()
73
 
@@ -81,14 +84,14 @@ def evaluate(cfg: dict, checkpoint: str | None = None) -> None:
81
 
82
  coco_gt = COCO(cfg["data"]["val_ann"])
83
  # Build HF label → COCO category_id mapping
84
- label2cat = {v: k for k, v in coco_gt.cats.items()}
85
 
86
  results = []
87
  for batch in loader:
88
  pixel_values = batch["pixel_values"].to(device)
89
  orig_sizes = batch["orig_sizes"].to(device)
90
 
91
- with autocast(enabled=use_fp16):
92
  outputs = model(pixel_values=pixel_values)
93
 
94
  preds = processor.post_process_object_detection(
 
13
 
14
  import torch
15
  import yaml
16
+ from torch.amp import autocast
17
  from torch.utils.data import DataLoader
18
  from transformers import RTDetrForObjectDetection, RTDetrImageProcessor
19
 
 
62
  @torch.inference_mode()
63
  def evaluate(cfg: dict, checkpoint: str | None = None) -> None:
64
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
65
+ use_fp16 = cfg["training"].get("fp16", False) and device.type == "cuda"
66
+ use_bf16 = cfg["training"].get("bf16", False) and device.type == "cuda"
67
+ use_amp = use_fp16 or use_bf16
68
+ amp_dtype = torch.bfloat16 if use_bf16 else torch.float16
69
 
70
  model_id = checkpoint or cfg["model"]["id"]
71
  processor = RTDetrImageProcessor.from_pretrained(model_id)
72
  model = RTDetrForObjectDetection.from_pretrained(
73
+ model_id, torch_dtype=amp_dtype if use_amp else torch.float32
74
  ).to(device)
75
  model.eval()
76
 
 
84
 
85
  coco_gt = COCO(cfg["data"]["val_ann"])
86
  # Build HF label → COCO category_id mapping
87
+ label2cat = {cat["name"]: cat["id"] for cat in coco_gt.cats.values()}
88
 
89
  results = []
90
  for batch in loader:
91
  pixel_values = batch["pixel_values"].to(device)
92
  orig_sizes = batch["orig_sizes"].to(device)
93
 
94
+ with autocast("cuda", enabled=use_amp, dtype=amp_dtype):
95
  outputs = model(pixel_values=pixel_values)
96
 
97
  preds = processor.post_process_object_detection(
scripts/train.py CHANGED
@@ -13,7 +13,7 @@ import random
13
  import numpy as np
14
  import torch
15
  import yaml
16
- from torch.cuda.amp import GradScaler, autocast
17
  from torch.optim import AdamW
18
  from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR
19
  from torch.utils.data import DataLoader
@@ -36,14 +36,14 @@ def build_coco_dataset(img_dir: str, ann_file: str, processor: RTDetrImageProces
36
  class _Wrapped(torch.utils.data.Dataset):
37
  def __getitem__(self, idx):
38
  img, targets = base[idx]
39
- annotations = [
40
- {
41
- "bbox": t["bbox"], # [x, y, w, h]
42
- "category_id": t["category_id"],
43
- "image_id": t["image_id"],
44
- }
45
- for t in targets
46
- ]
47
  encoding = processor(
48
  images=img,
49
  annotations=annotations,
@@ -67,22 +67,35 @@ def collate_fn(batch):
67
  # Training
68
  # ---------------------------------------------------------------------------
69
 
70
- def train(cfg: dict) -> None:
71
  seed = cfg["training"]["seed"]
72
  random.seed(seed)
73
  np.random.seed(seed)
74
  torch.manual_seed(seed)
75
 
76
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
77
- use_fp16: bool = cfg["training"]["fp16"] and device.type == "cuda"
 
 
 
78
 
79
- processor = RTDetrImageProcessor.from_pretrained(cfg["model"]["id"])
 
80
  model = RTDetrForObjectDetection.from_pretrained(
81
- cfg["model"]["id"],
82
  num_labels=cfg["model"]["num_labels"],
83
- ignore_mismatched_sizes=True,
84
  ).to(device)
85
 
 
 
 
 
 
 
 
 
 
86
  if cfg["training"]["gradient_checkpointing"]:
87
  model.gradient_checkpointing_enable()
88
 
@@ -112,7 +125,7 @@ def train(cfg: dict) -> None:
112
  cosine = CosineAnnealingLR(optimizer, T_max=epochs - warmup_epochs, eta_min=min_lr)
113
  scheduler = SequentialLR(optimizer, schedulers=[warmup, cosine], milestones=[warmup_epochs])
114
 
115
- scaler = GradScaler(enabled=use_fp16)
116
 
117
  train_ds = build_coco_dataset(
118
  cfg["data"]["train_img"], cfg["data"]["train_ann"], processor
@@ -131,7 +144,7 @@ def train(cfg: dict) -> None:
131
  grad_accum = cfg["training"]["grad_accum_steps"]
132
  clip_norm = cfg["training"]["clip_grad_norm"]
133
 
134
- for epoch in range(1, epochs + 1):
135
  model.train()
136
  running_loss = 0.0
137
  optimizer.zero_grad()
@@ -140,7 +153,7 @@ def train(cfg: dict) -> None:
140
  pixel_values = batch["pixel_values"].to(device)
141
  labels = [{k: v.to(device) for k, v in lbl.items()} for lbl in batch["labels"]]
142
 
143
- with autocast(enabled=use_fp16):
144
  outputs = model(pixel_values=pixel_values, labels=labels)
145
  loss = outputs.loss / grad_accum
146
 
@@ -157,6 +170,14 @@ def train(cfg: dict) -> None:
157
  if step % 50 == 0:
158
  print(f"[epoch {epoch}/{epochs} step {step}] loss={running_loss/step:.4f}")
159
 
 
 
 
 
 
 
 
 
160
  scheduler.step()
161
 
162
  if epoch % cfg["training"]["save_every_n_epochs"] == 0:
@@ -173,6 +194,7 @@ def train(cfg: dict) -> None:
173
  def parse_args() -> argparse.Namespace:
174
  parser = argparse.ArgumentParser()
175
  parser.add_argument("--config", default="configs/rtdetr_r50_coco.yaml")
 
176
  return parser.parse_args()
177
 
178
 
@@ -180,4 +202,4 @@ if __name__ == "__main__":
180
  args = parse_args()
181
  with open(args.config) as f:
182
  cfg = yaml.safe_load(f)
183
- train(cfg)
 
13
  import numpy as np
14
  import torch
15
  import yaml
16
+ from torch.amp import GradScaler, autocast
17
  from torch.optim import AdamW
18
  from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR
19
  from torch.utils.data import DataLoader
 
36
  class _Wrapped(torch.utils.data.Dataset):
37
  def __getitem__(self, idx):
38
  img, targets = base[idx]
39
+ image_id = targets[0]["image_id"] if targets else 0
40
+ annotations = {
41
+ "image_id": image_id,
42
+ "annotations": [
43
+ {"bbox": t["bbox"], "category_id": t["category_id"]}
44
+ for t in targets
45
+ ],
46
+ }
47
  encoding = processor(
48
  images=img,
49
  annotations=annotations,
 
67
  # Training
68
  # ---------------------------------------------------------------------------
69
 
70
+ def train(cfg: dict, resume: str | None = None) -> None:
71
  seed = cfg["training"]["seed"]
72
  random.seed(seed)
73
  np.random.seed(seed)
74
  torch.manual_seed(seed)
75
 
76
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
77
+ use_fp16: bool = cfg["training"].get("fp16", False) and device.type == "cuda"
78
+ use_bf16: bool = cfg["training"].get("bf16", False) and device.type == "cuda"
79
+ use_amp: bool = use_fp16 or use_bf16
80
+ amp_dtype = torch.bfloat16 if use_bf16 else torch.float16
81
 
82
+ model_src = resume if resume else cfg["model"]["id"]
83
+ processor = RTDetrImageProcessor.from_pretrained(model_src)
84
  model = RTDetrForObjectDetection.from_pretrained(
85
+ model_src,
86
  num_labels=cfg["model"]["num_labels"],
87
+ ignore_mismatched_sizes=(resume is None),
88
  ).to(device)
89
 
90
+ # Infer which epoch to start from when resuming
91
+ start_epoch = 1
92
+ if resume:
93
+ import re as _re
94
+ m = _re.search(r"epoch_(\d+)", os.path.basename(resume.rstrip("/\\")))
95
+ if m:
96
+ start_epoch = int(m.group(1)) + 1
97
+ print(f"Resuming from {resume}, starting at epoch {start_epoch}")
98
+
99
  if cfg["training"]["gradient_checkpointing"]:
100
  model.gradient_checkpointing_enable()
101
 
 
125
  cosine = CosineAnnealingLR(optimizer, T_max=epochs - warmup_epochs, eta_min=min_lr)
126
  scheduler = SequentialLR(optimizer, schedulers=[warmup, cosine], milestones=[warmup_epochs])
127
 
128
+ scaler = GradScaler("cuda", enabled=use_fp16) # GradScaler only needed for fp16, not bf16
129
 
130
  train_ds = build_coco_dataset(
131
  cfg["data"]["train_img"], cfg["data"]["train_ann"], processor
 
144
  grad_accum = cfg["training"]["grad_accum_steps"]
145
  clip_norm = cfg["training"]["clip_grad_norm"]
146
 
147
+ for epoch in range(start_epoch, epochs + 1):
148
  model.train()
149
  running_loss = 0.0
150
  optimizer.zero_grad()
 
153
  pixel_values = batch["pixel_values"].to(device)
154
  labels = [{k: v.to(device) for k, v in lbl.items()} for lbl in batch["labels"]]
155
 
156
+ with autocast("cuda", enabled=use_amp, dtype=amp_dtype):
157
  outputs = model(pixel_values=pixel_values, labels=labels)
158
  loss = outputs.loss / grad_accum
159
 
 
170
  if step % 50 == 0:
171
  print(f"[epoch {epoch}/{epochs} step {step}] loss={running_loss/step:.4f}")
172
 
173
+ # Flush any remaining accumulated gradients at end of epoch
174
+ if step % grad_accum != 0:
175
+ scaler.unscale_(optimizer)
176
+ torch.nn.utils.clip_grad_norm_(model.parameters(), clip_norm)
177
+ scaler.step(optimizer)
178
+ scaler.update()
179
+ optimizer.zero_grad()
180
+
181
  scheduler.step()
182
 
183
  if epoch % cfg["training"]["save_every_n_epochs"] == 0:
 
194
  def parse_args() -> argparse.Namespace:
195
  parser = argparse.ArgumentParser()
196
  parser.add_argument("--config", default="configs/rtdetr_r50_coco.yaml")
197
+ parser.add_argument("--resume", default=None, help="Path to checkpoint dir to resume training from")
198
  return parser.parse_args()
199
 
200
 
 
202
  args = parse_args()
203
  with open(args.config) as f:
204
  cfg = yaml.safe_load(f)
205
+ train(cfg, resume=args.resume)