Skip to content

Commit 88bc842

Browse files
committed
bug fix, edge IDs are passed as edge weights
1 parent 3f4bd9b commit 88bc842

1 file changed

Lines changed: 9 additions & 12 deletions

File tree

python/cugraph/cugraph/structure/replicate_edgelist.py

Lines changed: 9 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -48,9 +48,11 @@ def _call_plc_replicate_edgelist(
4848
resource_handle=ResourceHandle(Comms.get_handle(sID).getHandle()),
4949
src_array=edgelist_df[col_names[0]],
5050
dst_array=edgelist_df[col_names[1]],
51-
weight_array=edgelist_df[col_names[2]] if len(col_names) > 2 else None,
52-
edge_id_array=edgelist_df[col_names[3]] if len(col_names) > 3 else None,
53-
edge_type_id_array=edgelist_df[col_names[4]] if len(col_names) > 4 else None,
51+
weight_array=edgelist_df[col_names[2]] if col_names[2] is not None else None,
52+
edge_id_array=edgelist_df[col_names[3]] if col_names[3] is not None else None,
53+
edge_type_id_array=edgelist_df[col_names[4]]
54+
if col_names[4] is not None
55+
else None,
5456
)
5557
return _convert_to_cudf(cp_arrays, col_names)
5658

@@ -201,16 +203,11 @@ def replicate_edgelist(
201203
edgelist_ddf = dask_cudf.from_cudf(
202204
edgelist_ddf, npartitions=len(Comms.get_workers())
203205
)
204-
col_names = [source, destination]
206+
col_names = [source, destination, weight, edge_id, edge_type]
205207

206-
if weight is not None:
207-
col_names.append(weight)
208-
if edge_id is not None:
209-
col_names.append(edge_id)
210-
if edge_type is not None:
211-
col_names.append(edge_type)
212-
213-
if not (set(col_names).issubset(set(edgelist_ddf.columns))):
208+
if not {name for name in col_names if name is not None}.issubset(
209+
set(edgelist_ddf.columns)
210+
):
214211
raise ValueError(
215212
"Invalid column names were provided: valid columns names are "
216213
f"{edgelist_ddf.columns}"

0 commit comments

Comments
 (0)