dmlc--dgl
9b4d60799a
* WIP. remove graph arg in NodeBatch and EdgeBatch * refactor: use graph adapter for scheduler * WIP: recv * draft impl * stuck at bipartite * bipartite->unitgraph; support dsttype == srctype * pass test_query * pass test_query * pass test_view * test apply * pass udf message passing tests * pass quan's test using builtins * WIP: wildcard slicing * new construct methods * broken * good * add stack cross reducer * fix bug; fix mx * fix bug in csrmm2 when the CSR is not square * lint * removed FlattenedHeteroGraph class * WIP * prop nodes, prop edges, filter nodes/edges * add DGLGraph tests to heterograph. Fix several bugs * finish nx<->hetero graph conversion * create bipartite from nx * more spec on hetero/homo conversion * silly fixes * check node and edge types * repr * to api * adj APIs * inc * fix some lints and bugs * fix some lints * hetero/homo conversion * fix flatten test * more spec in hetero_from_homo and test * flatten using concat names * WIP: creators * rewrite hetero_from_homo in a more efficient way * remove useless variables * fix lint * subgraphs and typed subgraphs * lint & removed heterosubgraph class * lint x2 * disable heterograph mutation test * docstring update * add edge id for nx graph test * fix mx unittests * fix bug * try fix * fix unittest when cross_reducer is stack * fix ci * fix nx bipartite bug; docstring * fix scipy creation bug * lint * fix bug when converting heterograph from homograph * fix bug in hetero_from_homo about ntype order * trailing white * docstring fixes for add_foo and data views * docstring for relation slice * to_hetero and to_homo with feature support * lint * lint * DGLGraph compatibility * incidence matrix & docstring fixes * example string fixes * feature in hetero_from_relations * deduplication of edge types in to_hetero * fix lint * fix
127 行
2.6 KiB
Python
127 行
2.6 KiB
Python
"""Temporary adapter to unify DGLGraph and HeteroGraph for scheduler.
|
|
NOTE(minjie): remove once all scheduler codes are migrated to heterograph
|
|
"""
|
|
from __future__ import absolute_import
|
|
|
|
from abc import ABC, abstractmethod
|
|
|
|
class GraphAdapter(ABC):
|
|
"""Temporary adapter class to unify DGLGraph and DGLHeteroGraph for schedulers."""
|
|
@property
|
|
@abstractmethod
|
|
def gidx(self):
|
|
"""Get graph index object."""
|
|
|
|
@abstractmethod
|
|
def num_src(self):
|
|
"""Number of source nodes."""
|
|
|
|
@abstractmethod
|
|
def num_dst(self):
|
|
"""Number of destination nodes."""
|
|
|
|
@abstractmethod
|
|
def num_edges(self):
|
|
"""Number of edges."""
|
|
|
|
@property
|
|
@abstractmethod
|
|
def srcframe(self):
|
|
"""Frame to store source node features."""
|
|
|
|
@property
|
|
@abstractmethod
|
|
def dstframe(self):
|
|
"""Frame to store source node features."""
|
|
|
|
@property
|
|
@abstractmethod
|
|
def edgeframe(self):
|
|
"""Frame to store edge features."""
|
|
|
|
@property
|
|
@abstractmethod
|
|
def msgframe(self):
|
|
"""Frame to store messages."""
|
|
|
|
|
|
@property
|
|
@abstractmethod
|
|
def msgindicator(self):
|
|
"""Message indicator tensor."""
|
|
|
|
@msgindicator.setter
|
|
@abstractmethod
|
|
def msgindicator(self, val):
|
|
"""Set new message indicator tensor."""
|
|
|
|
@abstractmethod
|
|
def in_edges(self, nodes):
|
|
"""Get in edges
|
|
|
|
Parameters
|
|
----------
|
|
nodes : utils.Index
|
|
Nodes
|
|
|
|
Returns
|
|
-------
|
|
tuple of utils.Index
|
|
(src, dst, eid)
|
|
"""
|
|
|
|
@abstractmethod
|
|
def out_edges(self, nodes):
|
|
"""Get out edges
|
|
|
|
Parameters
|
|
----------
|
|
nodes : utils.Index
|
|
Nodes
|
|
|
|
Returns
|
|
-------
|
|
tuple of utils.Index
|
|
(src, dst, eid)
|
|
"""
|
|
|
|
@abstractmethod
|
|
def edges(self, form):
|
|
"""Get all edges
|
|
|
|
Parameters
|
|
----------
|
|
form : str
|
|
"eid", "uv", etc.
|
|
|
|
Returns
|
|
-------
|
|
tuple of utils.Index
|
|
(src, dst, eid)
|
|
"""
|
|
|
|
@abstractmethod
|
|
def get_immutable_gidx(self, ctx):
|
|
"""Get immutable graph index for kernel computation.
|
|
|
|
Parameters
|
|
----------
|
|
ctx : DGLContext
|
|
The context of the returned graph.
|
|
|
|
Returns
|
|
-------
|
|
GraphIndex
|
|
|
|
"""
|
|
|
|
@abstractmethod
|
|
def bits_needed(self):
|
|
"""Return the number of integer bits needed to represent the graph
|
|
|
|
Returns
|
|
-------
|
|
int
|
|
The number of bits needed
|
|
"""
|