Skip to content

Commit 177df2a

Browse files
authored
Merge pull request #42 from yuchen0cc/main
add yolo training integration and examples
2 parents 6efda81 + 43b28ec commit 177df2a

6 files changed

Lines changed: 1434 additions & 11 deletions

File tree

docs/torchconnector/examples.md

Lines changed: 219 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -316,6 +316,225 @@ DCP.load(
316316

317317
```
318318

319+
## YOLO
320+
321+
OSS connector for AI/ML provides integration with popular YOLO frameworks for training object detection models directly from OSS storage.
322+
323+
### Training with Ultralytics
324+
325+
OSS connector for AI/ML provides integration with [Ultralytics](https://docs.ultralytics.com/) framework for training YOLO models directly from OSS storage.
326+
327+
#### Dataset Configuration
328+
329+
A YAML (Yet Another Markup Language) file is used to define the dataset configuration. It contains information about the dataset's paths, classes, and other relevant information. Two formats are supported:
330+
331+
**Format 1: Directory-based**
332+
333+
In this format, `train` and `val` specify directory prefixes containing the images. The connector will list all objects under these prefixes.
334+
335+
```yaml
336+
path: oss://ossconnectorbucket/COCO_YOLO
337+
train: images/train2017
338+
val: images/val2017
339+
340+
nc: 80
341+
342+
names:
343+
0: person
344+
1: bicycle
345+
2: car
346+
# ... (class names continue)
347+
```
348+
349+
**Expected OSS object layout:**
350+
```
351+
<bucket>/<base_key>/
352+
images/
353+
train2017/ ← ~118k images
354+
val2017/ ← 5k images
355+
labels/
356+
train2017/ ← YOLO .txt annotations
357+
val2017/
358+
```
359+
360+
**Label path rule:** Images in `images/train2017/x.jpg` will automatically resolve to labels in `labels/train2017/x.txt`.
361+
362+
**Format 2: Manifest file-based**
363+
364+
In this format, `train` and `val` specify manifest files containing lists of image paths. This format is suitable for datasets with a large number of objects and repeated dataset loading, as it avoids the overhead of listing objects in OSS.
365+
366+
```yaml
367+
path: oss://ossconnectorbucket/COCO_YOLO
368+
train: images/train2017.txt
369+
val: images/val2017.txt
370+
371+
nc: 80
372+
373+
names:
374+
0: person
375+
1: bicycle
376+
2: car
377+
# ... (class names continue)
378+
```
379+
380+
**Manifest file format (`train2017.txt`):**
381+
```
382+
/bucket/COCO_YOLO/images/train2017/image001.jpg
383+
/bucket/COCO_YOLO/images/train2017/image002.jpg
384+
/bucket/COCO_YOLO/images/train2017/image003.jpg
385+
```
386+
387+
Each line in the manifest file is a `/bucket/key` path. Lines starting with `/` are used as-is; other lines are joined with the manifest's directory prefix.
388+
389+
#### Training Example
390+
391+
```py
392+
from ultralytics import YOLO
393+
from osstorchconnector import make_oss_trainer
394+
395+
396+
ENDPOINT = "http://oss-cn-beijing-internal.aliyuncs.com"
397+
REGION = "cn-beijing"
398+
CONFIG_PATH = "/etc/oss-connector/config.json"
399+
CRED_PATH = "/root/.alibabacloud/credentials"
400+
OSS_DATA = "coco_oss.yaml"
401+
402+
403+
custom_trainer = make_oss_trainer(
404+
endpoint=ENDPOINT,
405+
cred_path=CRED_PATH,
406+
config_path=CONFIG_PATH,
407+
)
408+
409+
# Load model
410+
model = YOLO("yolo26n.pt")
411+
412+
# Train the model
413+
results = model.train(
414+
trainer=custom_trainer,
415+
data=OSS_DATA,
416+
epochs=1, batch=16, imgsz=640, fraction=0.005)
417+
```
418+
419+
The `make_oss_trainer` function creates a custom trainer that enables Ultralytics to read training data directly from OSS. The trainer handles:
420+
- Resolving OSS URIs from the YAML configuration
421+
- Reading images and labels from OSS storage
422+
- Supporting both directory-based and manifest file-based dataset configurations
423+
424+
### Training with MMDetection
425+
426+
OSS connector for AI/ML provides integration with [MMDetection](https://github.com/open-mmlab/mmdetection) framework for training object detection models directly from OSS storage.
427+
428+
#### Dataset Configuration
429+
430+
MMDetection uses COCO-format JSON annotation files. The dataset configuration is specified programmatically when building the MMEngine config, rather than through a YAML file.
431+
432+
**Expected OSS object layout for COCO dataset:**
433+
```
434+
<bucket>/COCO/
435+
annotations/
436+
instances_train2017.json
437+
instances_val2017.json
438+
train2017/ ← training images
439+
val2017/ ← validation images
440+
```
441+
442+
#### Training Example
443+
444+
```py
445+
import os
446+
from mmengine.runner import Runner
447+
from mmengine.config import Config
448+
449+
# Import osstorchconnector — registers OSSDetDataset + OSSLoadImageFromFile
450+
from osstorchconnector import OSSDetDataset, OSSLoadImageFromFile, get_oss_ann_path
451+
452+
453+
ENDPOINT = "http://oss-cn-beijing-internal.aliyuncs.com"
454+
REGION = "cn-beijing"
455+
CONFIG_PATH = "/etc/oss-connector/config.json"
456+
CRED_PATH = "/root/.alibabacloud/credentials"
457+
458+
OSS_DATA_ROOT = 'oss://ossconnectorbucket/COCO'
459+
OSS_TRAIN_ANN = 'annotations/instances_train2017.json'
460+
OSS_VAL_ANN = 'annotations/instances_val2017.json'
461+
OSS_TRAIN_PREFIX = 'train2017/'
462+
OSS_VAL_PREFIX = 'val2017/'
463+
OSS_WORK_DIR = './work_dirs/mmdet_oss'
464+
465+
MAX_EPOCHS = 1
466+
BATCH_SIZE = 16
467+
NUM_WORKERS = 8
468+
469+
# Load MMDetection config
470+
cfg = Config.fromfile('rtmdet_tiny_8xb32-300e_coco.py')
471+
472+
# Patch image loader: LoadImageFromFile → OSSLoadImageFromFile
473+
def _patch_oss_pipeline(pipeline):
474+
out = []
475+
for t in pipeline:
476+
t = dict(t)
477+
if t.get("type") in ("LoadImageFromFile", "mmdet.LoadImageFromFile"):
478+
t = dict(type="OSSLoadImageFromFile")
479+
out.append(t)
480+
return out
481+
482+
cfg.train_dataloader.dataset.pipeline = _patch_oss_pipeline(
483+
cfg.train_dataloader.dataset.pipeline)
484+
cfg.val_dataloader.dataset.pipeline = _patch_oss_pipeline(
485+
cfg.val_dataloader.dataset.pipeline)
486+
487+
# Replace dataset type and OSS connection params
488+
cfg.merge_from_dict({
489+
"train_dataloader.dataset.type": "OSSDetDataset",
490+
"train_dataloader.dataset.data_root": OSS_DATA_ROOT,
491+
"train_dataloader.dataset.ann_file": OSS_TRAIN_ANN,
492+
"train_dataloader.dataset.data_prefix": dict(img=OSS_TRAIN_PREFIX),
493+
"train_dataloader.dataset.oss_endpoint": ENDPOINT,
494+
"train_dataloader.dataset.oss_cred_path": CRED_PATH,
495+
"train_dataloader.dataset.oss_config_path": CONFIG_PATH,
496+
"train_dataloader.dataset.oss_region": REGION,
497+
"val_dataloader.dataset.type": "OSSDetDataset",
498+
"val_dataloader.dataset.data_root": OSS_DATA_ROOT,
499+
"val_dataloader.dataset.ann_file": OSS_VAL_ANN,
500+
"val_dataloader.dataset.data_prefix": dict(img=OSS_VAL_PREFIX),
501+
"val_dataloader.dataset.oss_endpoint": ENDPOINT,
502+
"val_dataloader.dataset.oss_cred_path": CRED_PATH,
503+
"val_dataloader.dataset.oss_config_path": CONFIG_PATH,
504+
"val_dataloader.dataset.oss_region": REGION,
505+
})
506+
507+
# ann_cache_dir: downloaded annotation JSONs live inside work_dir/ann_cache/
508+
ann_cache_dir = os.path.join(OSS_WORK_DIR, "ann_cache")
509+
cfg.merge_from_dict({
510+
'train_dataloader.dataset.ann_cache_dir': ann_cache_dir,
511+
'val_dataloader.dataset.ann_cache_dir': ann_cache_dir,
512+
})
513+
514+
# CocoMetric: point at the local cache of the downloaded OSS annotation
515+
_, _, _val_local_ann = get_oss_ann_path(OSS_DATA_ROOT, OSS_VAL_ANN, ann_cache_dir)
516+
cfg.val_evaluator.ann_file = _val_local_ann
517+
cfg.test_evaluator.ann_file = _val_local_ann
518+
519+
# Training config overrides
520+
cfg.max_epochs = MAX_EPOCHS
521+
cfg.train_cfg.max_epochs = MAX_EPOCHS
522+
cfg.train_dataloader.batch_size = BATCH_SIZE
523+
cfg.train_dataloader.num_workers = NUM_WORKERS
524+
cfg.val_dataloader.batch_size = 1
525+
cfg.val_dataloader.num_workers = NUM_WORKERS
526+
cfg.work_dir = OSS_WORK_DIR
527+
528+
# Train
529+
runner = Runner.from_cfg(cfg)
530+
runner.train()
531+
```
532+
533+
The MMDetection integration provides:
534+
- `OSSDetDataset`: A custom dataset class that reads COCO-format annotations and images from OSS
535+
- `OSSLoadImageFromFile`: A data pipeline transform that loads images directly from OSS storage
536+
- `get_oss_ann_path`: A helper function to manage local caching of annotation files
537+
319538
## Safetensor
320539

321540
OSS connector for AI/ML supports saving/loading safetensors since v1.2.0rc6.
Lines changed: 69 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,23 +1,82 @@
1-
from .oss_iterable_dataset import OssIterableDataset
2-
from .oss_map_dataset import OssMapDataset
1+
# Core modules (no heavy dependencies - safe to import directly)
32
from .oss_checkpoint import OssCheckpoint
4-
from .oss_safetensor import OssSafetensor
5-
from .oss_dcp_filesystem import OssDCPFileSystem, OssStorageReader, OssStorageWriter
63
from ._oss_client import OssClient
74
from ._oss_connector import new_data_object
85
from ._oss_bucket_iterable import imagenet_manifest_parser
96
from ._oss_tar_iterable import generate_tar_archive
107

118
__all__ = [
12-
"OssIterableDataset",
13-
"OssMapDataset",
9+
# Core (no heavy deps)
1410
"OssCheckpoint",
15-
"OssSafetensor"
16-
"OssDCPFileSystem",
17-
"OssStorageReader",
18-
"OssStorageWriter",
1911
"OssClient",
2012
"new_data_object",
2113
"imagenet_manifest_parser",
2214
"generate_tar_archive",
15+
# Torch-dependent (lazy)
16+
"OssIterableDataset",
17+
"OssMapDataset",
18+
"OssSafetensor",
19+
# Torch distributed checkpoint-dependent (lazy)
20+
"OssDCPFileSystem",
21+
"OssStorageReader",
22+
"OssStorageWriter",
23+
# Ultralytics-dependent (lazy)
24+
"OSSYOLODataset",
25+
"make_oss_trainer",
26+
# MMDetection-dependent (lazy)
27+
"OSSDetDataset",
28+
"OSSLoadImageFromFile",
29+
"get_oss_ann_path",
2330
]
31+
32+
33+
def __getattr__(name: str):
34+
"""Lazy import for modules with heavy dependencies.
35+
36+
This avoids forcing users to install torch, safetensors, or ultralytics
37+
when they only need core functionality like OssCheckpoint.
38+
"""
39+
# Torch-dependent: OssIterableDataset, OssMapDataset
40+
if name == "OssIterableDataset":
41+
from .oss_iterable_dataset import OssIterableDataset
42+
return OssIterableDataset
43+
if name == "OssMapDataset":
44+
from .oss_map_dataset import OssMapDataset
45+
return OssMapDataset
46+
47+
# Torch + safetensors-dependent: OssSafetensor
48+
if name == "OssSafetensor":
49+
from .oss_safetensor import OssSafetensor
50+
return OssSafetensor
51+
52+
# Torch distributed checkpoint-dependent: OssDCPFileSystem, OssStorageReader, OssStorageWriter
53+
if name == "OssDCPFileSystem":
54+
from .oss_dcp_filesystem import OssDCPFileSystem
55+
return OssDCPFileSystem
56+
if name == "OssStorageReader":
57+
from .oss_dcp_filesystem import OssStorageReader
58+
return OssStorageReader
59+
if name == "OssStorageWriter":
60+
from .oss_dcp_filesystem import OssStorageWriter
61+
return OssStorageWriter
62+
63+
# Ultralytics-dependent: OSSYOLODataset, make_oss_trainer
64+
if name == "OSSYOLODataset":
65+
from .integrations.ultralytics import OSSYOLODataset
66+
return OSSYOLODataset
67+
if name == "make_oss_trainer":
68+
from .integrations.ultralytics import make_oss_trainer
69+
return make_oss_trainer
70+
71+
# MMDetection-dependent: OSSDetDataset, OSSLoadImageFromFile, get_oss_ann_path
72+
if name == "OSSDetDataset":
73+
from .integrations.mmdet import OSSDetDataset
74+
return OSSDetDataset
75+
if name == "OSSLoadImageFromFile":
76+
from .integrations.mmdet import OSSLoadImageFromFile
77+
return OSSLoadImageFromFile
78+
if name == "get_oss_ann_path":
79+
from .integrations.mmdet import get_oss_ann_path
80+
return get_oss_ann_path
81+
82+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
"""
2+
Framework-specific integrations for OSS data loading.
3+
4+
Each submodule provides drop-in replacements for popular ML frameworks:
5+
- ultralytics: OSSYOLODataset, make_oss_trainer for YOLO training
6+
- mmdet: OSSDetDataset, OSSLoadImageFromFile for MMDetection
7+
8+
All imports are lazy to avoid forcing dependency installation.
9+
"""
10+
11+
__all__ = [
12+
# Ultralytics-dependent
13+
"OSSYOLODataset",
14+
"make_oss_trainer",
15+
# MMDetection-dependent
16+
"get_oss_ann_path",
17+
"OSSDetDataset",
18+
"OSSLoadImageFromFile",
19+
]
20+
21+
22+
def __getattr__(name: str):
23+
"""Lazy import for framework-specific integrations."""
24+
# Ultralytics-dependent: OSSYOLODataset, make_oss_trainer
25+
if name == "OSSYOLODataset":
26+
from .ultralytics import OSSYOLODataset
27+
return OSSYOLODataset
28+
if name == "make_oss_trainer":
29+
from .ultralytics import make_oss_trainer
30+
return make_oss_trainer
31+
32+
# MMDetection-dependent: get_oss_ann_path, OSSDetDataset, OSSLoadImageFromFile
33+
if name == "get_oss_ann_path":
34+
from .mmdet import get_oss_ann_path
35+
return get_oss_ann_path
36+
if name == "OSSDetDataset":
37+
from .mmdet import OSSDetDataset
38+
return OSSDetDataset
39+
if name == "OSSLoadImageFromFile":
40+
from .mmdet import OSSLoadImageFromFile
41+
return OSSLoadImageFromFile
42+
43+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")

0 commit comments

Comments
 (0)