Skip to content

Commit b27df42

Browse files
authored
Support getitem with Index (#75)
Fix behaviour when passing Index to __getitem__ Fixes #74
1 parent 4751444 commit b27df42

2 files changed

Lines changed: 14 additions & 2 deletions

File tree

sparsity/sparse_frame.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -789,7 +789,7 @@ def drop_duplicate_idx(self, **kwargs):
789789
790790
Returns
791791
-------
792-
dropeed: SparseFrame
792+
dropped: SparseFrame
793793
"""
794794
mask = ~self.index.duplicated(**kwargs)
795795
return SparseFrame(self.data[mask], index=self.index.values[mask],
@@ -798,7 +798,10 @@ def drop_duplicate_idx(self, **kwargs):
798798
def __getitem__(self, item):
799799
if item is None:
800800
raise ValueError('Cannot label index with a null key.')
801-
if not isinstance(item, (tuple, list)):
801+
if not isinstance(item, (pd.Series, np.ndarray, pd.Index, list,
802+
tuple)):
803+
# TODO: tuple probably should be a separate case as in Pandas
804+
# where it is used with Multiindex
802805
item = [item]
803806
if len(item) > 0:
804807
return self.reindex_axis(item, axis=1)

sparsity/test/test_sparse_frame.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66

77
import numpy as np
88
import pandas as pd
9+
import pandas.testing as pdt
910
import pytest
1011
from moto import mock_s3
1112
from scipy import sparse
@@ -575,6 +576,7 @@ def test_npz_io_s3(complex_example):
575576
def test_getitem():
576577
id_ = np.identity(10)
577578
sf = SparseFrame(id_, columns=list('abcdefghij'))
579+
578580
assert sf['a'].data.todense()[0] == 1
579581
assert sf['j'].data.todense()[9] == 1
580582
assert np.all(sf[['a', 'b']].data.todense() == np.asmatrix(id_[:, [0, 1]]))
@@ -588,6 +590,13 @@ def test_getitem():
588590
with pytest.raises(ValueError):
589591
sf[None]
590592

593+
idx = pd.Index(list('abc'))
594+
pdt.assert_index_equal(idx, sf[idx].columns)
595+
pdt.assert_index_equal(idx, sf[idx.to_series()].columns)
596+
pdt.assert_index_equal(idx, sf[idx.tolist()].columns)
597+
pdt.assert_index_equal(idx, sf[tuple(idx)].columns)
598+
pdt.assert_index_equal(idx, sf[idx.values].columns)
599+
591600

592601
def test_vstack():
593602
frames = []

0 commit comments

Comments
 (0)