SAM2 微调代码详解-----dataloader
官方提供的微调代码是在MOSE数据集上进行微调,如下所示:

由于我暂未下载MOSE数据集,所以采用了与MOSE数据集结构相似的DAVIS数据集。
查看train.py文件,更改配置,比如yaml文件的位置、GPU数目等,如下所示:
if __name__ == "__main__":
initialize_config_module("sam2", version_base="1.2")
parser = ArgumentParser()
parser.add_argument(
"-c",
"--config",
# required=True,
default="configs/sam2.1_training/sam2.1_hiera_b+_MOSE_finetune.yaml",
type=str,
help="path to config file (e.g. configs/sam2.1_training/sam2.1_hiera_b+_MOSE_finetune.yaml)",
)
parser.add_argument(
"--use-cluster",
type=int,
default=0,
help="whether to launch on a cluster, 0: run locally, 1: run on a cluster",
)
parser.add_argument("--partition", type=str, default=None, help="SLURM partition")
parser.add_argument("--account", type=str, default=None, help="SLURM account")
parser.add_argument("--qos", type=str, default=None, help="SLURM qos")
parser.add_argument(
"--num-gpus", type=int, default=3, help="number of GPUS per node"
)
parser.add_argument("--num-nodes", type=int, default=None, help="Number of nodes")
args = parser.parse_args()
args.use_cluster = bool(args.use_cluster) if args.use_cluster is not None else None
register_omegaconf_resolvers()
main(args)
通过查看main()函数,我们可以发现,模型使用Hydra 框架进行配置,Hydra 主要提供以下几项核心功能:
- 配置文件的层次结构:支持配置文件之间的继承和组合,使得配置更加灵活和可扩展。
- 动态配置选择:允许在命令行上动态选择或覆盖配置中的某些字段。
- 多种配置源支持:可以使用 YAML 文件、Python 脚本、命令行参数等多种方式来定义和加载配置。
- 配置组(Config Groups):允许你为配置文件提供一组预定义的选项,从而轻松选择不同的配置组合。
有关Hydra的知识可参考:https://zhuanlan.zhihu.com/p/662221581
https://zhuanlan.zhihu.com/p/662221581
代码首先导入配置初始化trainer,然后再其基础上进行模型训练:
def single_proc_run(local_rank, main_port, cfg, world_size):
"""Single GPU process"""
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = str(main_port)
os.environ["RANK"] = str(local_rank)
os.environ["LOCAL_RANK"] = str(local_rank)
os.environ["WORLD_SIZE"] = str(world_size)
try:
register_omegaconf_resolvers()
except Exception as e:
logging.info(e)
trainer = instantiate(cfg.trainer, _recursive_=False)
trainer.run()
有关trainer的配置查看sam2.1_hiera_b+_MOSE_finetune.yaml如下:
trainer:
_target_: training.trainer.Trainer
mode: train_only
max_epochs: ${times:${scratch.num_epochs},${scratch.phases_per_epoch}}
accelerator: cuda
seed_value: 123
model:
_target_: training.model.sam2.SAM2Train
image_encoder:
_target_: sam2.modeling.backbones.image_encoder.ImageEncoder
scalp: 1
trunk:
_target_: sam2.modeling.backbones.hieradet.Hiera
embed_dim: 112
num_heads: 2
drop_path_rate: 0.1
neck:
_target_: sam2.modeling.backbones.image_encoder.FpnNeck
position_encoding:
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
num_pos_feats: 256
normalize: true
scale: null
temperature: 10000
d_model: 256
backbone_channel_list: [896, 448, 224, 112]
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
fpn_interp_model: nearest
memory_attention:
_target_: sam2.modeling.memory_attention.MemoryAttention
d_model: 256
pos_enc_at_input: true
layer:
_target_: sam2.modeling.memory_attention.MemoryAttentionLayer
activation: relu
dim_feedforward: 2048
dropout: 0.1
pos_enc_at_attn: false
self_attention:
_target_: sam2.modeling.sam.transformer.RoPEAttention
rope_theta: 10000.0
feat_sizes: [32, 32]
embedding_dim: 256
num_heads: 1
downsample_rate: 1
dropout: 0.1
d_model: 256
pos_enc_at_cross_attn_keys: true
pos_enc_at_cross_attn_queries: false
cross_attention:
_target_: sam2.modeling.sam.transformer.RoPEAttention
rope_theta: 10000.0
feat_sizes: [32, 32]
rope_k_repeat: True
embedding_dim: 256
num_heads: 1
downsample_rate: 1
dropout: 0.1
kv_in_dim: 64
num_layers: 4
memory_encoder:
_target_: sam2.modeling.memory_encoder.MemoryEncoder
out_dim: 64
position_encoding:
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
num_pos_feats: 64
normalize: true
scale: null
temperature: 10000
mask_downsampler:
_target_: sam2.modeling.memory_encoder.MaskDownSampler
kernel_size: 3
stride: 2
padding: 1
fuser:
_target_: sam2.modeling.memory_encoder.Fuser
layer:
_target_: sam2.modeling.memory_encoder.CXBlock
dim: 256
kernel_size: 7
padding: 3
layer_scale_init_value: 1e-6
use_dwconv: True # depth-wise convs
num_layers: 2
num_maskmem: 7
image_size: ${scratch.resolution}
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
sigmoid_scale_for_mem_enc: 20.0
sigmoid_bias_for_mem_enc: -10.0
use_mask_input_as_output_without_sam: true
# Memory
directly_add_no_mem_embed: true
no_obj_embed_spatial: true
# use high-resolution feature map in the SAM mask decoder
use_high_res_features_in_sam: true
# output 3 masks on the first click on initial conditioning frames
multimask_output_in_sam: true
# SAM heads
iou_prediction_use_sigmoid: True
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
use_obj_ptrs_in_encoder: true
add_tpos_enc_to_obj_ptrs: true
proj_tpos_enc_in_obj_ptrs: true
use_signed_tpos_enc_to_obj_ptrs: true
only_obj_ptrs_in_the_past_for_eval: true
# object occlusion prediction
pred_obj_scores: true
pred_obj_scores_mlp: true
fixed_no_obj_ptr: true
# multimask tracking settings
multimask_output_for_tracking: true
use_multimask_token_for_obj_ptr: true
multimask_min_pt_num: 0
multimask_max_pt_num: 1
use_mlp_for_obj_ptr_proj: true
# Compilation flag
# compile_image_encoder: False
####### Training specific params #######
# box/point input and corrections
prob_to_use_pt_input_for_train: 0.5
prob_to_use_pt_input_for_eval: 0.0
prob_to_use_box_input_for_train: 0.5 # 0.5*0.5 = 0.25 prob to use box instead of points
prob_to_use_box_input_for_eval: 0.0
prob_to_sample_from_gt_for_train: 0.1 # with a small prob, sampling correction points from GT mask instead of prediction errors
num_frames_to_correct_for_train: 2 # iteratively sample on random 1~2 frames (always include the first frame)
num_frames_to_correct_for_eval: 1 # only iteratively sample on first frame
rand_frames_to_correct_for_train: True # random #init-cond-frame ~ 2
add_all_frames_to_correct_as_cond: True # when a frame receives a correction click, it becomes a conditioning frame (even if it's not initially a conditioning frame)
# maximum 2 initial conditioning frames
num_init_cond_frames_for_train: 2
rand_init_cond_frames_for_train: True # random 1~2
num_correction_pt_per_frame: 7
use_act_ckpt_iterative_pt_sampling: false
num_init_cond_frames_for_eval: 1 # only mask on the first frame
forward_backbone_per_frame_for_eval: True
data:
train:
_target_: training.dataset.sam2_datasets.TorchTrainMixedDataset
phases_per_epoch: ${scratch.phases_per_epoch}
batch_sizes:
- ${scratch.train_batch_size}
datasets:
- _target_: training.dataset.utils.RepeatFactorWrapper
dataset:
_target_: training.dataset.utils.ConcatDataset
datasets:
- _target_: training.dataset.vos_dataset.VOSDataset
transforms: ${vos.train_transforms}
training: true
video_dataset:
_target_: training.dataset.vos_raw_dataset.PNGRawDataset
img_folder: ${dataset.img_folder}
gt_folder: ${dataset.gt_folder}
file_list_txt: ${dataset.file_list_txt}
sampler:
_target_: training.dataset.vos_sampler.RandomUniformSampler
num_frames: ${scratch.num_frames}
max_num_objects: ${scratch.max_num_objects}
multiplier: ${dataset.multiplier}
shuffle: True
num_workers: ${scratch.num_train_workers}
pin_memory: True
drop_last: True
collate_fn:
_target_: training.utils.data_utils.collate_fn
_partial_: true
dict_key: all
optim:
amp:
enabled: True
amp_dtype: bfloat16
optimizer:
_target_: torch.optim.AdamW
gradient_clip:
_target_: training.optimizer.GradientClipper
max_norm: 0.1
norm_type: 2
param_group_modifiers:
- _target_: training.optimizer.layer_decay_param_modifier
_partial_: True
layer_decay_value: 0.9
apply_to: 'image_encoder.trunk'
overrides:
- pattern: '*pos_embed*'
value: 1.0
options:
lr:
- scheduler:
_target_: fvcore.common.param_scheduler.CosineParamScheduler
start_value: ${scratch.base_lr}
end_value: ${divide:${scratch.base_lr},10}
- scheduler:
_target_: fvcore.common.param_scheduler.CosineParamScheduler
start_value: ${scratch.vision_lr}
end_value: ${divide:${scratch.vision_lr},10}
param_names:
- 'image_encoder.*'
weight_decay:
- scheduler:
_target_: fvcore.common.param_scheduler.ConstantParamScheduler
value: 0.1
- scheduler:
_target_: fvcore.common.param_scheduler.ConstantParamScheduler
value: 0.0
param_names:
- '*bias*'
module_cls_names: ['torch.nn.LayerNorm']
loss:
all:
_target_: training.loss_fns.MultiStepMultiMasksAndIous
weight_dict:
loss_mask: 20
loss_dice: 1
loss_iou: 1
loss_class: 1
supervise_all_iou: true
iou_use_l1_loss: true
pred_obj_scores: true
focal_gamma_obj_score: 0.0
focal_alpha_obj_score: -1.0
distributed:
backend: nccl
find_unused_parameters: True
logging:
tensorboard_writer:
_target_: training.utils.logger.make_tensorboard_logger
log_dir: ${launcher.experiment_log_dir}/tensorboard
flush_secs: 120
should_log: True
log_dir: ${launcher.experiment_log_dir}/logs
log_freq: 10
# initialize from a SAM 2 checkpoint
checkpoint:
save_dir: ${launcher.experiment_log_dir}/checkpoints
save_freq: 0 # 0 only last checkpoint is saved.
model_weight_initializer:
_partial_: True
_target_: training.utils.checkpoint_utils.load_state_dict_into_model
strict: True
ignore_unexpected_keys: null
ignore_missing_keys: null
state_dict:
_target_: training.utils.checkpoint_utils.load_checkpoint_and_apply_kernels
checkpoint_path: /data/seg/sam2-main/checkpoints/sam2.1_hiera_base_plus.pt # PATH to SAM 2.1 checkpoint
ckpt_state_dict_keys: ['model']
我们可以看到,里面有关于model、data、loss等配置,因此需要修改配置中涉及的一些文件地址:
- 修改dataset的image和annotation的位置:
dataset:
# PATHS to Dataset
img_folder: /data/seg/DAVIS/2017/trainval/JPEGImages/480p # PATH to MOSE JPEGImages folder
gt_folder: /data/seg/DAVIS/2017/trainval/Annotations/480p # PATH to MOSE Annotations folder
file_list_txt: /data/seg/sam2-main/training/assets/DAVIS_2017_train.txt # Optional PATH to filelist containing a subset of videos to be used for training
multiplier: 2
- 修改checkpoint_path的位置:
# initialize from a SAM 2 checkpoint checkpoint: save_dir: ${launcher.experiment_log_dir}/checkpoints save_freq: 0 # 0 only last checkpoint is saved. model_weight_initializer: _partial_: True _target_: training.utils.checkpoint_utils.load_state_dict_into_model strict: True ignore_unexpected_keys: null ignore_missing_keys: null state_dict: _target_: training.utils.checkpoint_utils.load_checkpoint_and_apply_kernels checkpoint_path: /data/seg/sam2-main/checkpoints/sam2.1_hiera_base_plus.pt # PATH to SAM 2.1 checkpoint ckpt_state_dict_keys: ['model']
现在,进入trainer.py文件,其中的使用配置文件中的配置信息配置后的class Trainer对应着train文件中实例化的trainer,__init__()函数中传入的参数即对应着yaml配置文件中的信息:
class Trainer:
"""
Trainer supporting the DDP training strategies.
"""
EPSILON = 1e-8
def __init__(
self,
*, # the order of these args can change at any time, so they are keyword-only
data: Dict[str, Any],
model: Dict[str, Any],
logging: Dict[str, Any],
checkpoint: Dict[str, Any],
max_epochs: int,
mode: str = "train",
accelerator: str = "cuda",
seed_value: int = 123,
val_epoch_freq: int = 1,
distributed: Dict[str, bool] = None,
cuda: Dict[str, bool] = None,
env_variables: Optional[Dict[str, Any]] = None,
optim: Optional[Dict[str, Any]] = None,
optim_overrides: Optional[List[Dict[str, Any]]] = None,
meters: Optional[Dict[str, Any]] = None,
loss: Optional[Dict[str, Any]] = None,
):
我们首先来看data,这涉及数据是如何导入的,如果我们要使用自己的数据集进行微调,我们需要进行哪些数据预处理。查看data配置:
data:
train:
_target_: training.dataset.sam2_datasets.TorchTrainMixedDataset
phases_per_epoch: ${scratch.phases_per_epoch}
batch_sizes:
- ${scratch.train_batch_size}
datasets:
- _target_: training.dataset.utils.RepeatFactorWrapper
dataset:
_target_: training.dataset.utils.ConcatDataset
datasets:
- _target_: training.dataset.vos_dataset.VOSDataset
transforms: ${vos.train_transforms}
training: true
video_dataset:
_target_: training.dataset.vos_raw_dataset.PNGRawDataset
img_folder: ${dataset.img_folder}
gt_folder: ${dataset.gt_folder}
file_list_txt: ${dataset.file_list_txt}
sampler:
_target_: training.dataset.vos_sampler.RandomUniformSampler
num_frames: ${scratch.num_frames}
max_num_objects: ${scratch.max_num_objects}
multiplier: ${dataset.multiplier}
shuffle: True
num_workers: ${scratch.num_train_workers}
pin_memory: True
drop_last: True
collate_fn:
_target_: training.utils.data_utils.collate_fn
_partial_: true
dict_key: all
因此,我们可以发现,在训练阶段,data由TorchTrainMixedDataset实例化:
class TorchTrainMixedDataset:
def __init__(
self,
datasets: List[Dataset],
batch_sizes: List[int],
num_workers: int,
shuffle: bool,
pin_memory: bool,
drop_last: bool,
collate_fn: Optional[Callable] = None,
worker_init_fn: Optional[Callable] = None,
phases_per_epoch: int = 1,
dataset_prob: Optional[List[float]] = None,
) -> None:
而TorchTrainMixedDataset中的datasets由RepeatFactorWrapper实例化:
class RepeatFactorWrapper(Dataset):
"""
Thin wrapper around a dataset to implement repeat factor sampling.
The underlying dataset must have a repeat_factors member to indicate the per-image factor.
Set it to uniformly ones to disable repeat factor sampling
"""
def __init__(self, dataset, seed: int = 0):
self.dataset = dataset
self.epoch_ids = None
self._seed = seed
# Split into whole number (_int_part) and fractional (_frac_part) parts.
self._int_part = torch.trunc(dataset.repeat_factors)
self._frac_part = dataset.repeat_factors - self._int_part
但RepeatFactorWrapper类中的dataset又调用了类ConcatDataset:
class ConcatDataset(TorchConcatDataset):
def __init__(self, datasets: Iterable[Dataset]) -> None:
super(ConcatDataset, self).__init__(datasets)
self.repeat_factors = torch.cat([d.repeat_factors for d in datasets])
类ConcatDataset中的datasets又调用了类VOSDataset:
class VOSDataset(VisionDataset):
def __init__(
self,
transforms,
training: bool,
video_dataset: VOSRawDataset,
sampler: VOSSampler,
multiplier: int,
always_target=True,
target_segments_available=True,
):
self._transforms = transforms
self.training = training
self.video_dataset = video_dataset
self.sampler = sampler
self.repeat_factors = torch.ones(len(self.video_dataset), dtype=torch.float32)
self.repeat_factors *= multiplier
print(f"Raw dataset length = {len(self.video_dataset)}")
self.curr_epoch = 0 # Used in case data loader behavior changes across epochs
self.always_target = always_target
self.target_segments_available = target_segments_available
类VOSDataset中的video_dataset又调用了类PNGRawDataset:
class PNGRawDataset(VOSRawDataset):
def __init__(
self,
img_folder,
gt_folder,
file_list_txt=None,
excluded_videos_list_txt=None,
sample_rate=1,
is_palette=True,
single_object_mode=False,
truncate_video=-1,
frames_sampling_mult=False,
):
这是一个层层嵌套的过程,从而也加大了我们的阅读难度,我们从最底层的类开始逐步解析。
1. PNGRawDataset
- __init__方法
该方法首先读取了file_list_txt文件中待处理的视频的名字,通过视频的名字,以及提供的img和gt的路径,我们可以定位到每一个待处理的视频的位置。
class PNGRawDataset(VOSRawDataset):
def __init__(
self,
img_folder,
gt_folder,
file_list_txt=None,
excluded_videos_list_txt=None,
sample_rate=1,
is_palette=True,
single_object_mode=False,
truncate_video=-1,
frames_sampling_mult=False,
):
self.img_folder = img_folder
self.gt_folder = gt_folder
self.sample_rate = sample_rate
self.is_palette = is_palette
self.single_object_mode = single_object_mode # False
self.truncate_video = truncate_video
# Read the subset defined in file_list_txt
if file_list_txt is not None:
with g_pathmgr.open(file_list_txt, "r") as f:
subset = [os.path.splitext(line.strip())[0] for line in f]
else:
subset = os.listdir(self.img_folder)
# Read and process excluded files if provided
if excluded_videos_list_txt is not None:
with g_pathmgr.open(excluded_videos_list_txt, "r") as f:
excluded_files = [os.path.splitext(line.strip())[0] for line in f]
else:
excluded_files = []
# Check if it's not in excluded_files
self.video_names = sorted(
[video_name for video_name in subset if video_name not in excluded_files]
)
if self.single_object_mode:
# single object mode
self.video_names = sorted(
[
os.path.join(video_name, obj)
for video_name in self.video_names
for obj in os.listdir(os.path.join(self.gt_folder, video_name))
]
)
if frames_sampling_mult:
video_names_mult = []
for video_name in self.video_names:
num_frames = len(os.listdir(os.path.join(self.img_folder, video_name)))
video_names_mult.extend([video_name] * num_frames)
self.video_names = video_names_mult
- get_video方法:
该方法主要根据idx,即索引去检索得到当前的video_name之后,获取得到当前视频对应的img path以及gt path,分别包装为对应的segment_loader以及video,其中主要调用了类PalettisedPNGSegmentLoader以及VOSVideo处理每一帧信息。
def get_video(self, idx):
"""
Given a VOSVideo object, return the mask tensors.
"""
video_name = self.video_names[idx]
if self.single_object_mode:
video_frame_root = os.path.join(
self.img_folder, os.path.dirname(video_name)
)
else:
video_frame_root = os.path.join(self.img_folder, video_name)
video_mask_root = os.path.join(self.gt_folder, video_name)
if self.is_palette:
segment_loader = PalettisedPNGSegmentLoader(video_mask_root)
else:
segment_loader = MultiplePNGSegmentLoader(
video_mask_root, self.single_object_mode
)
all_frames = sorted(glob.glob(os.path.join(video_frame_root, "*.jpg")))
if self.truncate_video > 0:
all_frames = all_frames[: self.truncate_video]
frames = []
for _, fpath in enumerate(all_frames[:: self.sample_rate]):
fid = int(os.path.basename(fpath).split(".")[0])
frames.append(VOSFrame(fid, image_path=fpath))
video = VOSVideo(video_name, idx, frames)
return video, segment_loader
- __len__方法返回当前使用数据集的样本数目,即有多少个视频参与训练:
def __len__(self):
return len(self.video_names)
2. VOSDataset
- __init__方法
class VOSDataset(VisionDataset):
def __init__(
self,
transforms,
training: bool,
video_dataset: VOSRawDataset,
sampler: VOSSampler,
multiplier: int,
always_target=True,
target_segments_available=True,
):
self._transforms = transforms
self.training = training
self.video_dataset = video_dataset
self.sampler = sampler
self.repeat_factors = torch.ones(len(self.video_dataset), dtype=torch.float32)
self.repeat_factors *= multiplier
print(f"Raw dataset length = {len(self.video_dataset)}")
self.curr_epoch = 0 # Used in case data loader behavior changes across epochs
self.always_target = always_target
self.target_segments_available = target_segments_available
在该方法中使用RandomUniformSampler对self.sampler进行初始化,RandomUniformSampler主要是涉及从每个视频中采样一定的帧数作为训练,因为不同的视频长度是不一样的。随机采样固定帧数之后,以一定的概率确定是否将视频帧进行反转,这可以模拟从视频的后期到前期的反向推理。并且,载入第一帧的mask信息,统计第一帧中有多少个object,确保首帧一定有object存在。最后返回采样后的frames以及object id信息。
与此同时,__init__方法还对self._transforms进行初始化,主要涉及的transform方法在配置文件中给出,通过观察transforms.py文件中的方法可知,它们处理的数据类型均为VideoDatapoint,因此在后续的处理过程中我们需要根据我们的video_dataset得到相应的VideoDatapoint。
vos:
train_transforms:
- _target_: training.dataset.transforms.ComposeAPI
transforms:
- _target_: training.dataset.transforms.RandomHorizontalFlip
consistent_transform: True
- _target_: training.dataset.transforms.RandomAffine
degrees: 25
shear: 20
image_interpolation: bilinear
consistent_transform: True
- _target_: training.dataset.transforms.RandomResizeAPI
sizes: ${scratch.resolution}
square: true
consistent_transform: True
- _target_: training.dataset.transforms.ColorJitter
consistent_transform: True
brightness: 0.1
contrast: 0.03
saturation: 0.03
hue: null
- _target_: training.dataset.transforms.RandomGrayscale
p: 0.05
consistent_transform: True
- _target_: training.dataset.transforms.ColorJitter
consistent_transform: False
brightness: 0.1
contrast: 0.05
saturation: 0.05
hue: null
- _target_: training.dataset.transforms.ToTensorAPI
- _target_: training.dataset.transforms.NormalizeAPI
mean: [0.485, 0.456, 0.406]
std: [0.229, 0.224, 0.225]
- construct方法
该方法传入video、segment_loader以及采样后的视频帧信息,每个object在对应的帧中都有相应的mask,如果某一个object在当前帧中未出现,则对应的mask为全0的矩阵。将每一帧的RGB三维数据存储到Frame类下的data域中,而object mask信息以list的形式存储到Frame类下的object域中。最后将所有帧的信息封装到类VideoDatapoint中返回。
@dataclass
class Frame:
data: Union[torch.Tensor, PILImage.Image]
objects: List[Object]
def construct(self, video, sampled_frms_and_objs, segment_loader):
"""
Constructs a VideoDatapoint sample to pass to transforms
"""
sampled_frames = sampled_frms_and_objs.frames
sampled_object_ids = sampled_frms_and_objs.object_ids
images = []
rgb_images = load_images(sampled_frames)
# Iterate over the sampled frames and store their rgb data and object data (bbox, segment)
for frame_idx, frame in enumerate(sampled_frames):
w, h = rgb_images[frame_idx].size
images.append(
Frame(
data=rgb_images[frame_idx],
objects=[],
)
)
# We load the gt segments associated with the current frame
if isinstance(segment_loader, JSONSegmentLoader):
segments = segment_loader.load(
frame.frame_idx, obj_ids=sampled_object_ids
)
else:
segments = segment_loader.load(frame.frame_idx)
for obj_id in sampled_object_ids:
# Extract the segment
if obj_id in segments:
assert (
segments[obj_id] is not None
), "None targets are not supported"
# segment is uint8 and remains uint8 throughout the transforms
segment = segments[obj_id].to(torch.uint8)
else:
# There is no target, we either use a zero mask target or drop this object
if not self.always_target:
continue
segment = torch.zeros(h, w, dtype=torch.uint8)
images[frame_idx].objects.append(
Object(
object_id=obj_id,
frame_index=frame.frame_idx,
segment=segment,
)
)
return VideoDatapoint(
frames=images,
video_id=video.video_id,
size=(h, w),
)
- _get_datapoint方法:
该方法基于传入的索引idx,调用PNGRawDataset的get_video()方法获取得到当前索引对应的视频的video和segment_loader信息,获取完成之后,调用RandomUniformSampler的sample()方法对视频帧进行采样,将采样的结果传入self.construct()方法中最后得到datapoint数据,再使用调用transforms去对当前视频中的视频帧进行数据增强、Resize等操作。
def _get_datapoint(self, idx):
for retry in range(MAX_RETRIES):
try:
if isinstance(idx, torch.Tensor):
idx = idx.item()
# sample a video
video, segment_loader = self.video_dataset.get_video(idx)
# sample frames and object indices to be used in a datapoint
sampled_frms_and_objs = self.sampler.sample(
video, segment_loader, epoch=self.curr_epoch
)
break # Succesfully loaded video
except Exception as e:
if self.training:
logging.warning(
f"Loading failed (id={idx}); Retry {retry} with exception: {e}"
)
idx = random.randrange(0, len(self.video_dataset))
else:
# Shouldn't fail to load a val video
raise e
datapoint = self.construct(video, sampled_frms_and_objs, segment_loader)
for transform in self._transforms:
datapoint = transform(datapoint, epoch=self.curr_epoch)
return datapoint
- __getitem__方法 :
该方法调用_get_datapoint()方法,根据索引 idx 获取具体的数据样本。
def __getitem__(self, idx):
return self._get_datapoint(idx)
3. ConcatDataset
- __init__方法 :
调用父类的构造方法,将可迭代对象 datasets 传递给父类即PyTorch 自带的ConcatDataset,它会将多个数据集组合为一个数据集,允许通过单个索引访问不同子数据集中的样本。
class ConcatDataset(TorchConcatDataset):
def __init__(self, datasets: Iterable[Dataset]) -> None:
super(ConcatDataset, self).__init__(datasets)
self.repeat_factors = torch.cat([d.repeat_factors for d in datasets])
- set_epoch方法:
该方法为所有子数据集设置当前的 epoch,以便在采样或数据加载过程中依据 epoch 实现可重复的行为(例如随机种子的初始化)。
def set_epoch(self, epoch: int):
for dataset in self.datasets:
if hasattr(dataset, "epoch"):
dataset.epoch = epoch
if hasattr(dataset, "set_epoch"):
dataset.set_epoch(epoch)
4. RepeatFactorWrapper
- __init__方法:
传入ConcatDataset,它必须具有 repeat_factors 属性,该属性指定了每个样本的重复因子。_seed用于伪随机数生成器的种子,默认为0 。_int_part将 repeat_factors 的整数部分提取出来,表示每个样本的整数重复次数。_frac_part获取 repeat_factors 的小数部分,表示每个样本的小数部分,后续用于决定是否额外增加一个重复。
# Adapted from Detectron2
class RepeatFactorWrapper(Dataset):
"""
Thin wrapper around a dataset to implement repeat factor sampling.
The underlying dataset must have a repeat_factors member to indicate the per-image factor.
Set it to uniformly ones to disable repeat factor sampling
"""
def __init__(self, dataset, seed: int = 0):
self.dataset = dataset
self.epoch_ids = None
self._seed = seed
# Split into whole number (_int_part) and fractional (_frac_part) parts.
self._int_part = torch.trunc(dataset.repeat_factors)
self._frac_part = dataset.repeat_factors - self._int_part
- get_epoch_indices方法:
该方法根据每个样本的重复因子生成一组训练数据索引,索引中包含重复的样本,以使得样本的重复次数符合其对应的重复因子。
def _get_epoch_indices(self, generator):
"""
Create a list of dataset indices (with repeats) to use for one epoch.
Args:
generator (torch.Generator): pseudo random number generator used for
stochastic rounding.
Returns:
torch.Tensor: list of dataset indices to use in one epoch. Each index
is repeated based on its calculated repeat factor.
"""
# Since repeat factors are fractional, we use stochastic rounding so
# that the target repeat factor is achieved in expectation over the
# course of training
rands = torch.rand(len(self._frac_part), generator=generator)
rep_factors = self._int_part + (rands < self._frac_part).float()
# Construct a list of indices in which we repeat images as specified
indices = []
for dataset_index, rep_factor in enumerate(rep_factors):
indices.extend([dataset_index] * int(rep_factor.item()))
return torch.tensor(indices, dtype=torch.int64)
- set_epoch方法:
调用 self._get_epoch_indices(g) 来根据当前的种子和重复因子计算出样本的索引顺序,并保存在 self.epoch_ids 中。
def set_epoch(self, epoch: int):
g = torch.Generator()
g.manual_seed(self._seed + epoch)
self.epoch_ids = self._get_epoch_indices(g)
if hasattr(self.dataset, "set_epoch"):
self.dataset.set_epoch(epoch)
- __getitem__方法:
根据 epoch_ids 提供的索引从原始数据集中获取数据样本。
def __getitem__(self, idx):
if self.epoch_ids is None:
raise RuntimeError(
"Repeat ids haven't been computed. Did you forget to call set_epoch?"
)
return self.dataset[self.epoch_ids[idx]]
5. TorchTrainMixedDataset
- __init__方法:
datasets表示需要混合的多个 Dataset 实例列表。这里仅传入RepeatFactorWrapper。batch_sizes对应每个数据集的批量大小。num_workers表示每个数据加载器的工作线程数。shuffle表示是否对数据进行洗牌。pin_memory代表是否在加载数据时使用固定内存,这通常可以提高数据加载性能。drop_last表示是否丢弃数据集的最后一批数据,通常用于确保每批数据的大小一致。
class TorchTrainMixedDataset:
def __init__(
self,
datasets: List[Dataset],
batch_sizes: List[int],
num_workers: int,
shuffle: bool,
pin_memory: bool,
drop_last: bool,
collate_fn: Optional[Callable] = None,
worker_init_fn: Optional[Callable] = None,
phases_per_epoch: int = 1,
dataset_prob: Optional[List[float]] = None,
) -> None:
"""
Args:
datasets (List[Dataset]): List of Datasets to be mixed.
batch_sizes (List[int]): Batch sizes for each dataset in the list.
num_workers (int): Number of workers per dataloader.
shuffle (bool): Whether or not to shuffle data.
pin_memory (bool): If True, use pinned memory when loading tensors from disk.
drop_last (bool): Whether or not to drop the last batch of data.
collate_fn (Callable): Function to merge a list of samples into a mini-batch.
worker_init_fn (Callable): Function to init each dataloader worker.
phases_per_epoch (int): Number of phases per epoch.
dataset_prob (List[float]): Probability of choosing the dataloader to sample from. Should sum to 1.0
"""
self.datasets = datasets
self.batch_sizes = batch_sizes
self.num_workers = num_workers
self.shuffle = shuffle
self.pin_memory = pin_memory
self.drop_last = drop_last
self.collate_fn = collate_fn
self.worker_init_fn = worker_init_fn
assert len(self.datasets) > 0
for dataset in self.datasets:
assert not isinstance(dataset, IterableDataset), "Not supported"
# `RepeatFactorWrapper` requires calling set_epoch first to get its length
self._set_dataset_epoch(dataset, 0)
self.phases_per_epoch = phases_per_epoch
self.chunks = [None] * len(datasets)
if dataset_prob is None:
# If not provided, assign each dataset a probability proportional to its length.
dataset_lens = [
(math.floor(len(d) / bs) if drop_last else math.ceil(len(d) / bs))
for d, bs in zip(datasets, batch_sizes)
]
total_len = sum(dataset_lens)
dataset_prob = torch.tensor([d_len / total_len for d_len in dataset_lens])
else:
assert len(dataset_prob) == len(datasets)
dataset_prob = torch.tensor(dataset_prob)
logging.info(f"Dataset mixing probabilities: {dataset_prob.tolist()}")
assert dataset_prob.sum().item() == 1.0, "Probabilities should sum to 1.0"
self.dataset_prob = dataset_prob
- get_loader函数:
该方法根据当前的 epoch 创建一个数据加载器(DataLoader)列表:如果 phases_per_epoch > 1,则每个 epoch 被拆分成多个阶段(local_phase),在每个阶段中,数据集会被重新分配并进行处理。通过使用 torch.chunk() 将数据集划分为多个子集,并在每个阶段循环访问这些子集。每个数据集的 Sampler 被封装在 BatchSampler 中,使用 DataLoader 生成对应的加载器。 最后,返回一个 MixedDataLoader 对象,该对象会同时管理多个数据加载器,并根据指定的概率从多个数据集中采样数据。
def get_loader(self, epoch) -> Iterable:
dataloaders = []
for d_idx, (dataset, batch_size) in enumerate(
zip(self.datasets, self.batch_sizes)
):
if self.phases_per_epoch > 1:
# Major epoch that looops over entire dataset
# len(main_epoch) == phases_per_epoch * len(epoch)
main_epoch = epoch // self.phases_per_epoch
# Phase with in the main epoch
local_phase = epoch % self.phases_per_epoch
# Start of new data-epoch or job is resumed after preemtion.
if local_phase == 0 or self.chunks[d_idx] is None:
# set seed for dataset epoch
# If using RepeatFactorWrapper, this step currectly re-samples indices before chunking.
self._set_dataset_epoch(dataset, main_epoch)
# Separate random generator for subset sampling
g = torch.Generator()
g.manual_seed(main_epoch)
self.chunks[d_idx] = torch.chunk(
torch.randperm(len(dataset), generator=g),
self.phases_per_epoch,
)
dataset = Subset(dataset, self.chunks[d_idx][local_phase])
else:
self._set_dataset_epoch(dataset, epoch)
sampler = DistributedSampler(dataset, shuffle=self.shuffle)
sampler.set_epoch(epoch)
batch_sampler = BatchSampler(sampler, batch_size, drop_last=self.drop_last)
dataloaders.append(
DataLoader(
dataset,
num_workers=self.num_workers,
pin_memory=self.pin_memory,
batch_sampler=batch_sampler,
collate_fn=self.collate_fn,
worker_init_fn=self.worker_init_fn,
)
)
return MixedDataLoader(dataloaders, self.dataset_prob)
魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。
更多推荐


所有评论(0)