浏览代码

Remove `image_weights` DDP code (#4579)

* Initial commit

* Update
modifyDataloader
Glenn Jocher GitHub 3 年前
父节点
当前提交
d7aa3f153d
找不到此签名对应的密钥 GPG 密钥 ID: 4AEE18F83AFDEB23
共有 1 个文件被更改,包括 6 次插入14 次删除
  1. +6
    -14
      train.py

+ 6
- 14
train.py 查看文件

@@ -265,21 +265,13 @@ def train(hyp, # path/to/hyp.yaml or hyp dictionary
for epoch in range(start_epoch, epochs): # epoch ------------------------------------------------------------------
model.train()

# Update image weights (optional)
# Update image weights (optional, single-GPU only)
if opt.image_weights:
# Generate indices
if RANK in [-1, 0]:
cw = model.class_weights.cpu().numpy() * (1 - maps) ** 2 / nc # class weights
iw = labels_to_image_weights(dataset.labels, nc=nc, class_weights=cw) # image weights
dataset.indices = random.choices(range(dataset.n), weights=iw, k=dataset.n) # rand weighted idx
# Broadcast if DDP
if RANK != -1:
indices = (torch.tensor(dataset.indices) if RANK == 0 else torch.zeros(dataset.n)).int()
dist.broadcast(indices, 0)
if RANK != 0:
dataset.indices = indices.cpu().numpy()

# Update mosaic border
cw = model.class_weights.cpu().numpy() * (1 - maps) ** 2 / nc # class weights
iw = labels_to_image_weights(dataset.labels, nc=nc, class_weights=cw) # image weights
dataset.indices = random.choices(range(dataset.n), weights=iw, k=dataset.n) # rand weighted idx

# Update mosaic border (optional)
# b = int(random.uniform(0.25 * imgsz, 0.75 * imgsz + gs) // gs * gs)
# dataset.mosaic_border = [b - imgsz, -b] # height, width borders


正在加载...
取消
保存