File size: 2,603 Bytes
a3a407d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
"""Data augmentation pipeline using albumentations.
This script augments images and correspondingly transforms YOLO bboxes.
"""
import os
from glob import glob
import cv2
import albumentations as A
from pathlib import Path
import argparse

AUG = A.Compose([
    A.HorizontalFlip(p=0.5),
    A.RandomRotate90(p=0.5),
    A.RandomBrightnessContrast(p=0.5),
    A.GaussianBlur(p=0.3),
    A.GaussNoise(p=0.3),
], bbox_params=A.BboxParams(format='yolo', label_fields=['category_ids']))


def load_label_txt(txt_path):
    bboxes = []
    labels = []
    if not os.path.exists(txt_path):
        return bboxes, labels
    with open(txt_path,'r') as f:
        for line in f:
            vals = line.strip().split()
            if not vals:
                continue
            cls = int(vals[0])
            bbox = list(map(float, vals[1:5]))
            bboxes.append(bbox)
            labels.append(cls)
    return bboxes, labels


def save_label_txt(txt_path, bboxes, labels):
    os.makedirs(os.path.dirname(txt_path), exist_ok=True)
    with open(txt_path,'w') as f:
        for cls,b in zip(labels,bboxes):
            f.write(f"{cls} {b[0]:.6f} {b[1]:.6f} {b[2]:.6f} {b[3]:.6f}\n")


def augment_image(img_path, label_path, out_img_path, out_lbl_path, n=3):
    img = cv2.imread(img_path)
    h,w = img.shape[:2]
    bboxes, labels = load_label_txt(label_path)
    for i in range(n):
        try:
            augmented = AUG(image=img, bboxes=bboxes, category_ids=labels)
        except Exception:
            continue
        aug_img = augmented['image']
        aug_bboxes = augmented['bboxes']
        save_img_p = out_img_path.replace('{i}',str(i))
        save_lbl_p = out_lbl_path.replace('{i}',str(i))
        cv2.imwrite(save_img_p, aug_img)
        save_label_txt(save_lbl_p, aug_bboxes, augmented['category_ids'])


if __name__ == '__main__':
    parser = argparse.ArgumentParser()
    parser.add_argument('--src', default='dataset/images/train')
    parser.add_argument('--labels', default='dataset/labels/train')
    parser.add_argument('--out', default='dataset_aug')
    parser.add_argument('--n', type=int, default=3)
    args = parser.parse_args()

    img_files = glob(os.path.join(args.src,'*.jpg')) + glob(os.path.join(args.src,'*.png'))
    for img_path in img_files:
        stem = Path(img_path).stem
        lbl_path = os.path.join(args.labels, stem + '.txt')
        out_img = os.path.join(args.out, 'images', stem + '_aug_{i}.jpg')
        out_lbl = os.path.join(args.out, 'labels', stem + '_aug_{i}.txt')
        augment_image(img_path, lbl_path, out_img, out_lbl, n=args.n)