@@ -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