dmlc--dgl
5b9147c464
* add rtfd * rrr * update * change env * temp fix * update * fix * fix * add * conf * Move file_pattern from Makefile to conf.py * remove yml * fix * fix * fix * fix * remove yml * remove yml * add doc docker * add dgl install script * change name * change dockerfile * fix * name * add * fix * fix * fix * fix * fix docker * delete sphinx.py for doc-build backend * Add softmax to test backend * Add group apply function and tests * Delete unnecessary file * Update comments and test * Fix lint * remove unused bucketing code * group apply edge bucketing code * gen degree bucket schedule for group apply edge * schedule and graph code * fix compiling * fix * fix lint * naming * harder test case * fix comments * more comments * tweak function name
62 行
1.1 KiB
Python
62 行
1.1 KiB
Python
from __future__ import absolute_import
|
|
|
|
import torch as th
|
|
|
|
def cuda():
|
|
return th.device('cuda')
|
|
|
|
def array_equal(a, b):
|
|
return th.equal(a, b)
|
|
|
|
def allclose(a, b):
|
|
return th.allclose(a.float(), b.float(), rtol=1e-4, atol=1e-4)
|
|
|
|
def randn(shape):
|
|
return th.randn(*shape)
|
|
|
|
def attach_grad(x):
|
|
if x.grad is not None:
|
|
x.grad.zero_()
|
|
return x
|
|
else:
|
|
return x.requires_grad_()
|
|
|
|
def backward(x, head_gradient=None):
|
|
x.backward(head_gradient)
|
|
|
|
def grad(x):
|
|
return x.grad
|
|
|
|
def is_no_grad(x):
|
|
return x.grad is None or (x.grad == 0).all()
|
|
|
|
def full(shape, fill_value, dtype, ctx):
|
|
return th.full(shape, fill_value, dtype=dtype, device=ctx)
|
|
|
|
def narrow_row_set(x, start, stop, new):
|
|
x[start:stop] = new
|
|
|
|
def sparse_to_numpy(x):
|
|
return x.to_dense().numpy()
|
|
|
|
def clone(x):
|
|
return x.clone()
|
|
|
|
def reduce_sum(x):
|
|
return x.sum()
|
|
|
|
def softmax(x, dim):
|
|
return th.softmax(x, dim)
|
|
|
|
class record_grad(object):
|
|
def __init__(self):
|
|
pass
|
|
|
|
def __enter__(self):
|
|
pass
|
|
|
|
def __exit__(self, exc_type, exc_value, exc_traceback):
|
|
pass
|
|
|
|
no_grad = th.no_grad
|