dmlc--dgl
cbee427839
* add working scripts * add frcnn training script * remove redundent files * refactor validation computation, will optimize sgdet and training * validation finally finished * f-rcnn training * test reldn * rm file * update reldn training * data preprocess to h5 * temp * use coco json * fix conflict * new obj dataset for detection * update training * before cleanup * remove abundant files * add arg parse to train * cleanup code file * update * fix * add readme * add ipynb as demo * add demo pic * update readme * add demo script * improve paths * improve readme * add docstrings * fix args description * update readme * add models from s3 * update README Co-authored-by: Minjie Wang <minjie.wang@nyu.edu>
19 行
628 B
Python
19 行
628 B
Python
"""DataLoader utils."""
|
|
import dgl
|
|
from mxnet import nd
|
|
from gluoncv.data.batchify import Pad
|
|
|
|
def dgl_mp_batchify_fn(data):
|
|
if isinstance(data[0], tuple):
|
|
data = zip(*data)
|
|
return [dgl_mp_batchify_fn(i) for i in data]
|
|
|
|
for dt in data:
|
|
if dt is not None:
|
|
if isinstance(dt, dgl.DGLGraph):
|
|
return [d for d in data if isinstance(d, dgl.DGLGraph)]
|
|
elif isinstance(dt, nd.NDArray):
|
|
pad = Pad(axis=(1, 2), num_shards=1, ret_length=False)
|
|
data_list = [dt for dt in data if dt is not None]
|
|
return pad(data_list)
|