allow user to save an empty label dataframe (#1012)

This commit is contained in:
Bruce Martin
2019-11-04 14:14:49 -08:00
committed by GitHub
parent a16c2e2747
commit 823d819eee
+14 -10
View File
@@ -350,34 +350,37 @@ class ScanpyEngine(CXGDriver):
) )
@requires_data @requires_data
def _validate_label_data(self): def _validate_label_data(self, labels=None):
""" """
labels is None if disabled, empty if enabled by no data labels is None if disabled, empty if enabled by no data
""" """
if self.labels is None or self.labels.empty: if labels is None:
labels = self.labels
if labels is None or labels.empty:
return return
# all lables must have a name, which must be unique and not used in obs column names # all lables must have a name, which must be unique and not used in obs column names
if not self.labels.columns.is_unique: if not labels.columns.is_unique:
raise KeyError(f"All column names specified in {self.config['label_file']} must be unique.") raise KeyError(f"All column names specified in {self.config['label_file']} must be unique.")
# the label index must be unique, and must have same values the anndata obs index # the label index must be unique, and must have same values the anndata obs index
if not self.labels.index.is_unique: if not labels.index.is_unique:
raise KeyError(f"All row index values specified in the label file " raise KeyError(f"All row index values specified in the label file "
f"`{self.config['label_file']}` must be unique.") f"`{self.config['label_file']}` must be unique.")
if not self.labels.index.equals(self.original_obs_index): if not labels.index.equals(self.original_obs_index):
raise KeyError("Label file row index does not match H5AD file index. " raise KeyError("Label file row index does not match H5AD file index. "
"Please ensure that column zero (0) in the label file contain the same " "Please ensure that column zero (0) in the label file contain the same "
"index values as the H5AD file.") "index values as the H5AD file.")
duplicate_columns = list(set(self.labels.columns) & set(self.data.obs.columns)) duplicate_columns = list(set(labels.columns) & set(self.data.obs.columns))
if len(duplicate_columns) > 0: if len(duplicate_columns) > 0:
raise KeyError(f"Labels file may not contain column names which overlap " raise KeyError(f"Labels file may not contain column names which overlap "
f"with h5ad obs columns {duplicate_columns}") f"with h5ad obs columns {duplicate_columns}")
# labels must have same count as obs annotations # labels must have same count as obs annotations
if self.labels.shape[0] != self.data.obs.shape[0]: if labels.shape[0] != self.data.obs.shape[0]:
raise ValueError("Labels file must have same number of rows as h5ad file.") raise ValueError("Labels file must have same number of rows as h5ad file.")
@staticmethod @staticmethod
@@ -467,8 +470,9 @@ class ScanpyEngine(CXGDriver):
raise ValueError("Only OBS dimension access is supported") raise ValueError("Only OBS dimension access is supported")
new_label_df = decode_matrix_fbs(fbs) new_label_df = decode_matrix_fbs(fbs)
new_label_df.index = self.original_obs_index if not new_label_df.empty:
self._validate_label_data() # paranoia new_label_df.index = self.original_obs_index
self._validate_label_data(labels=new_label_df) # paranoia
# if any of the new column labels overlap with our existing labels, raise error # if any of the new column labels overlap with our existing labels, raise error
duplicate_columns = list(set(new_label_df.columns) & set(self.data.obs.columns)) duplicate_columns = list(set(new_label_df.columns) & set(self.data.obs.columns))
@@ -477,7 +481,7 @@ class ScanpyEngine(CXGDriver):
f"with h5ad obs columns {duplicate_columns}") f"with h5ad obs columns {duplicate_columns}")
# update our internal state and save it. Multi-threading often enabled, # update our internal state and save it. Multi-threading often enabled,
# so treat this as a critical section critical section. # so treat this as a critical section.
with self.label_lock: with self.label_lock:
self.labels = new_label_df self.labels = new_label_df
write_labels(fname, self.labels) write_labels(fname, self.labels)