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)
|