diff --git a/python/cudf/cudf/io/csv.py b/python/cudf/cudf/io/csv.py index eccedc7d9d02..e5b99b68fac1 100644 --- a/python/cudf/cudf/io/csv.py +++ b/python/cudf/cudf/io/csv.py @@ -5,7 +5,7 @@ import csv import itertools import os -from collections.abc import Collection, Mapping +from collections.abc import Collection, Mapping, Sequence from io import BytesIO, StringIO, TextIOBase from typing import TYPE_CHECKING, cast @@ -177,9 +177,6 @@ def read_csv( if byte_range is None: byte_range = (0, 0) - # We need this later when setting index cols - orig_header = header - if names is not None: # explicitly mentioned name, so don't check header if header is None or header == "infer": @@ -358,7 +355,7 @@ def read_csv( if ( isinstance(index_col_name, str) and names is None - and orig_header == "infer" + and header != -1 ): if index_col_name.startswith("Unnamed:"): # TODO: Try to upstream it to libcudf @@ -366,6 +363,20 @@ def read_csv( df.index.name = None elif names is None: df.index.name = index_col + elif isinstance(index_col, Sequence) and all( + isinstance(col, int) for col in index_col + ): + index_col_labels = list(df._data.get_labels_by_index(index_col)) + df = df.set_index(index_col_labels) + if names is None: + df.index.names = [ + (None if label.startswith("Unnamed:") else label) + if isinstance(label, str) and header != -1 + else position + for label, position in zip( + index_col_labels, index_col, strict=True + ) + ] else: df = df.set_index(index_col) diff --git a/python/cudf/cudf/tests/input_output/test_csv.py b/python/cudf/cudf/tests/input_output/test_csv.py index 3be019c1058c..d2fffc125278 100644 --- a/python/cudf/cudf/tests/input_output/test_csv.py +++ b/python/cudf/cudf/tests/input_output/test_csv.py @@ -1263,6 +1263,28 @@ def test_csv_reader_index_col(): pd_df = pd.read_csv(StringIO(buffer), header=None, index_col=False) assert_eq(cu_df.index, pd_df.index) + # using a single column index wrapped in a list + cu_df = read_csv(StringIO(buffer), header=None, index_col=[0]) + pd_df = pd.read_csv(StringIO(buffer), header=None, index_col=[0]) + assert_eq(cu_df.index, pd_df.index) + + # using multiple column indices + cu_df = read_csv(StringIO(buffer), header=None, index_col=[0, 1]) + pd_df = pd.read_csv(StringIO(buffer), header=None, index_col=[0, 1]) + assert_eq(cu_df.index, pd_df.index) + + +def test_csv_reader_index_col_position_with_explicit_header(): + buffer = "a,b,c\n3,4,5\n6,7,8" + + cu_df = read_csv(StringIO(buffer), header=0, index_col=[0]) + pd_df = pd.read_csv(StringIO(buffer), header=0, index_col=[0]) + assert_eq(cu_df.index, pd_df.index) + + cu_df = read_csv(StringIO(buffer), header=0, index_col=0) + pd_df = pd.read_csv(StringIO(buffer), header=0, index_col=0) + assert_eq(cu_df.index, pd_df.index) + @pytest.mark.parametrize("index_name", [None, "custom name", 124]) @pytest.mark.parametrize("index_col", [None, 0, "a"])