File size: 4,278 Bytes
9e14838
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
# -*- coding: utf-8 -*-
import os
import random

import numpy as np


def _extract_data_based_dist(
    data_type: str, image_paths: list, labels: list, dist: list, **params
):
    def sampling_frames(image_paths: list, labels: list, dist: list, **params):
        """
        Extracting image paths and labels based on given distribution
        """
        total_f = sum(labels)
        total_r = len(labels) - total_f

        assert (
            total_f > 0
        ), "Number of fake images must be greater than 0 for distribution sampling!"
        assert (
            total_r > 0
        ), "Number of real images must be greater than 0 for distribution sampling!"

        r_dist, f_dist = dist[0], dist[1]

        idxes = sorted(range(0, len(labels)), key=lambda k: labels[k])
        r_idxes = idxes[:total_r]
        f_idxes = idxes[total_r:]

        print(f"Original Number of Fake images --- {len(f_idxes)}")
        print(f"Original Number of Real images --- {len(r_idxes)}")

        if int((total_f / f_dist) * r_dist) > total_r:
            total_f = int((total_r / r_dist) * f_dist)
            f_idxes = random.sample(f_idxes, total_f)
        else:
            total_r = int((total_f / f_dist) * r_dist)
            r_idxes = random.sample(r_idxes, total_r)

        print(
            f"Number of Fake images --- {len(f_idxes)} given Fake distribution --- {f_dist}"
        )
        print(
            f"Number of Real images --- {len(r_idxes)} given Real distribution --- {r_dist}"
        )

        new_idxes = r_idxes + f_idxes

        image_paths = [image_paths[i] for i in new_idxes]
        labels = np.array(labels)[new_idxes]

        for k, v in params.items():
            if v is not None and len(v):
                params[k] = [v[i] for i in new_idxes]

        return image_paths, labels, params

    def sampling_videos(image_paths: list, labels: list, dist: list, **params):
        """
        Extracting image paths and labels based on given distribution for video data
        """
        f_vid_ids = []
        r_vid_ids = []
        for ip in image_paths:
            faketype = ip.split("/")[8]
            vid_id = os.path.dirname(ip)
            if (
                faketype == "real_videos"
                or "real" in faketype
                or "original" in faketype
            ):
                r_vid_ids.append(vid_id)
            else:
                f_vid_ids.append(vid_id)

        f_vid_ids = list(set(f_vid_ids))
        r_vid_ids = list(set(r_vid_ids))
        total_f = len(f_vid_ids)
        total_r = len(r_vid_ids)

        assert (
            total_f > 0
        ), "Number of fake videos must be greater than 0 for distribution sampling!"
        assert (
            total_r > 0
        ), "Number of real videos must be greater than 0 for distribution sampling!"
        print(f"Original Number of Fake videos --- {total_f}")
        print(f"Original Number of Real videos --- {total_r}")

        r_dist, f_dist = dist[0], dist[1]

        if int((total_f / f_dist) * r_dist) > total_r:
            total_f = int((total_r / r_dist) * f_dist)
            f_vid_ids = random.sample(f_vid_ids, total_f)
        else:
            total_r = int((total_f / f_dist) * r_dist)
            r_vid_ids = random.sample(r_vid_ids, total_r)

        print(
            f"Number of Fake videos --- {len(f_vid_ids)} given Fake distribution --- {f_dist}"
        )
        print(
            f"Number of Real videos --- {len(r_vid_ids)} given Real distribution --- {r_dist}"
        )

        vid_ids = r_vid_ids + f_vid_ids
        new_idxes = []

        for i in range(len(labels)):
            ip = image_paths[i]
            vid_id = "/".join([ip.split("/")[-3], ip.split("/")[-2]])
            if vid_id in vid_ids:
                new_idxes.append(i)

        image_paths = [image_paths[i] for i in new_idxes]
        labels = np.array(labels)[new_idxes]

        for k, v in params.items():
            if v is not None and len(v):
                params[k] = [v[i] for i in new_idxes]

        return image_paths, labels, params

    if data_type == "image":
        return sampling_frames(image_paths, labels, dist, **params)
    else:
        return sampling_videos(image_paths, labels, dist, **params)