feat: annotate command improvements (#2568)

* Replace --input-h5ad-file with a positional argument, for consistency with other CLI commands
* Replace --update-h5ad-file with --overwrite, for consistency with `prepare` command.
* Fix/clarify various help descriptions
* Fix final output message when input file is overwritten
* Fix annotate top-level help description
This commit is contained in:
Andrew Tolopko
2022-09-15 11:55:58 -04:00
committed by GitHub
parent 450261f109
commit ddb601c103
2 changed files with 89 additions and 59 deletions
+43 -32
View File
@@ -4,7 +4,7 @@ import os.path
import shlex import shlex
import shutil import shutil
import subprocess import subprocess
import sys from os.path import isfile
from subprocess import STDOUT, PIPE from subprocess import STDOUT, PIPE
from tempfile import NamedTemporaryFile from tempfile import NamedTemporaryFile
@@ -27,23 +27,42 @@ def annotate_args(func):
@sort_options @sort_options
@click.command( @click.command(
short_help="Annotate H5AD file columns. Run `cellxgene annotation --help` for more information.", options_metavar="<options>"
options_metavar="<options>",
) )
@click.option( @click.argument(
"-i", "input_h5ad_file",
"--input-h5ad-file", type=click.Path(exists=True, dir_okay=False, readable=True),
nargs=1,
metavar="<path to H5AD input file>",
required=True, required=True,
type=str,
help="The input H5AD file containing the missing annotations.",
) )
@click.option( @click.option(
"-m", "-m",
"--model-url", "--model-url",
# Making this a required "option", rather than an "argument", since we support automatic model selection in the
# future, in which case the user would not need to specify this option at all and we can make it optional at
# that time.
required=True, required=True,
help="The URL of the model used to prediction annotated labels. May be a local filesystem directory " help="The URL of the model used to prediction annotated labels. May be a local filesystem directory "
"or S3 path (s3://)", "or S3 path (s3://)",
) )
@click.option(
"-o",
"--output-h5ad-file",
default="",
help="The output H5AD file that will contain the generated annotation values. If this option is not provided, "
"the input file will be overwritten to include the new annotations; in this case you must specify "
"--overwrite.",
metavar="<filename>",
)
@click.option(
"--overwrite",
default=False,
is_flag=True,
help="Allow overwriting of the specified H5AD output file, if it exists. For safety, you must specify this "
"flag if the specified output file already exists or if the --output-h5ad-file option is not provided.",
show_default=True,
)
@click.option( @click.option(
"-l", "-l",
"--counts-layer", "--counts-layer",
@@ -53,8 +72,8 @@ def annotate_args(func):
@click.option( @click.option(
"-g", "-g",
"--gene-column-name", "--gene-column-name",
help="The name of the `var` column that contains gene identifiers. The values in this column will be used to match " help="The name of the `var` column that contains gene names. The values in this column will be used to match "
"genes between the query and reference datasets. If not specified, the gene identifiers are expected to exist " "genes between the query and reference datasets. If not specified, the gene names are expected to exist "
"in `var.index`.", "in `var.index`.",
) )
# TODO: Useful if we want to support discoverability of models # TODO: Useful if we want to support discoverability of models
@@ -91,19 +110,6 @@ def annotate_args(func):
"will store the predicted annotation values and confidence scores. This can be used to allow multiple " "will store the predicted annotation values and confidence scores. This can be used to allow multiple "
"annotation predictions to be run on a single AnnData object.", "annotation predictions to be run on a single AnnData object.",
) )
@click.option(
"-u",
"--update-h5ad-file",
is_flag=True,
help="Flag indicating whether to update the input h5ad file with annotation values. This option is mutually "
"exclusive with --output-h5ad-file.",
)
@click.option(
"-o",
"--output-h5ad-file",
help="The output H5AD file that will contain the generated annotation values. This option is mutually "
"exclusive with --update-h5ad-file.",
)
@click.option("--use-model-cache/--no-use-model-cache", default=True) @click.option("--use-model-cache/--no-use-model-cache", default=True)
@click.option( @click.option(
"--use-gpu/--no-use-gpu", "--use-gpu/--no-use-gpu",
@@ -141,6 +147,9 @@ def annotate_args(func):
) )
@click.help_option("--help", "-h", help="Show this message and exit.") @click.help_option("--help", "-h", help="Show this message and exit.")
def annotate(**cli_args): def annotate(**cli_args):
"""
Add predicted annotations to an H5AD file. Run `cellxgene annotate --help` for more information.
"""
_validate_options(cli_args) _validate_options(cli_args)
print(f"Reading query dataset {cli_args['input_h5ad_file']}...") print(f"Reading query dataset {cli_args['input_h5ad_file']}...")
@@ -149,7 +158,11 @@ def annotate(**cli_args):
filter(None, [cli_args.get("annotation_prefix"), cli_args.get("annotation_type"), cli_args.get("run_name")]) filter(None, [cli_args.get("annotation_prefix"), cli_args.get("annotation_type"), cli_args.get("run_name")])
) )
output_h5ad_file = cli_args["input_h5ad_file"] if cli_args["update_h5ad_file"] else cli_args["output_h5ad_file"] output_h5ad_file = (
cli_args["input_h5ad_file"]
if cli_args["overwrite"] and not cli_args["output_h5ad_file"]
else cli_args["output_h5ad_file"]
)
model_url = cli_args.get("model_url") model_url = cli_args.get("model_url")
local_model_path = _retrieve_model(cli_args.get("model_cache_dir"), model_url, cli_args.get("use_model_cache")) local_model_path = _retrieve_model(cli_args.get("model_cache_dir"), model_url, cli_args.get("use_model_cache"))
@@ -196,7 +209,7 @@ def annotate(**cli_args):
p.wait() p.wait()
if p.returncode == 0: if p.returncode == 0:
print(f"Wrote annotations to {cli_args.get('output_h5ad_file')}") print(f"Wrote annotations to {output_h5ad_file}")
else: else:
print("Annotation failed!") print("Annotation failed!")
else: else:
@@ -218,13 +231,11 @@ def _retrieve_model(model_cache_dir, model_url, use_cache=True):
def _validate_options(cli_args): def _validate_options(cli_args):
# TODO(atolopko): Use cloup library for this logic output = cli_args["output_h5ad_file"]
if cli_args["update_h5ad_file"] and cli_args["output_h5ad_file"]: overwrite = cli_args["overwrite"]
click.echo("--update_h5ad_file and --output_h5ad_file are mutually exclusive")
sys.exit(1) if isfile(output) and not overwrite:
if not (cli_args["update_h5ad_file"] or cli_args["output_h5ad_file"]): raise click.UsageError(f"Cannot overwrite existing file {output}, try using the flag --overwrite")
click.echo("--update_h5ad_file or --output_h5ad_file must be specified")
sys.exit(1)
if __name__ == "__main__": if __name__ == "__main__":
+46 -27
View File
@@ -1,7 +1,7 @@
import os import os
import shutil import shutil
import unittest import unittest
from tempfile import mkstemp, TemporaryDirectory from tempfile import mkstemp, TemporaryDirectory, NamedTemporaryFile
import mlflow import mlflow
from click.testing import CliRunner from click.testing import CliRunner
@@ -12,10 +12,8 @@ from test.unit.cli.fixtures.mlflow_model_fixture import FakeModel
def write_model(model) -> str: def write_model(model) -> str:
with TemporaryDirectory() as mlflow_model_dir: with TemporaryDirectory() as mlflow_model_dir:
fixtures_path = os.path.join(os.path.dirname(__file__), 'fixtures') fixtures_path = os.path.join(os.path.dirname(__file__), "fixtures")
mlflow.pyfunc.save_model(mlflow_model_dir, mlflow.pyfunc.save_model(mlflow_model_dir, loader_module="fixtures", code_path=[fixtures_path])
loader_module='fixtures',
code_path=[fixtures_path])
return shutil.make_archive(mkstemp()[1], "zip", mlflow_model_dir) return shutil.make_archive(mkstemp()[1], "zip", mlflow_model_dir)
@@ -46,7 +44,6 @@ class TestCliAnnotate(unittest.TestCase):
result = CliRunner().invoke( result = CliRunner().invoke(
annotate, annotate,
[ [
"--input-h5ad-file",
query_dataset_file_path, query_dataset_file_path,
"--model-url", "--model-url",
model_file_path, model_file_path,
@@ -54,7 +51,8 @@ class TestCliAnnotate(unittest.TestCase):
f"{query_dataset_file_path}.output", f"{query_dataset_file_path}.output",
# avoid having mflow create conda env or virtualenv when in test env; # avoid having mflow create conda env or virtualenv when in test env;
# this avoids making pip remote requests and is also faster # this avoids making pip remote requests and is also faster
"--mlflow-env-manager", "local" "--mlflow-env-manager",
"local",
], ],
) )
@@ -74,31 +72,52 @@ class TestCliAnnotate(unittest.TestCase):
result.stdout, result.stdout,
"inputs passed correctly", "inputs passed correctly",
) )
def test__annotate__verifies_mutually_exclusive_options(self):
required_options = ["--input-h5ad-file", "some.h5ad", "--model-url", "some_url"]
result = CliRunner().invoke(
annotate,
required_options + [],
)
self.assertNotEqual(0, result.exit_code, "aborts with non-success code")
self.assertIn( self.assertIn(
"--update_h5ad_file or --output_h5ad_file must be specified", f"Wrote annotations to {query_dataset_file_path}.output",
result.stdout, result.stdout,
"error message displayed", "success message is correct",
) )
result = CliRunner().invoke( def test__annotate__requires_overwrite_option_when_output_file_exists(self):
annotate, required_options + ["--output-h5ad-file", "some_arg", "--update-h5ad-file"]
)
self.assertNotEqual(0, result.exit_code, "aborts with non-success code") with NamedTemporaryFile() as input_h5ad, NamedTemporaryFile() as existing_file:
self.assertIn( required_options = [input_h5ad.name, "--output-h5ad-file", existing_file.name, "--model-url", "some_url"]
"--update_h5ad_file and --output_h5ad_file are mutually exclusive", result = CliRunner().invoke(
result.stdout, annotate,
"error message displayed", required_options + [],
) )
self.assertNotEqual(0, result.exit_code, "aborts with non-success code")
self.assertIn(
"try using the flag --overwrite",
result.stdout,
"error message displayed",
)
def test__annotate__overwrite_option_allows_overwrite_of_existing_output_file(self):
model_file_path = write_model(FakeModel())
with NamedTemporaryFile() as existing_file:
required_options = [
existing_file.name,
"--output-h5ad-file",
existing_file.name,
"--overwrite",
"--model-url",
model_file_path,
]
result = CliRunner().invoke(
annotate,
required_options + [],
)
print(result.stdout)
self.assertNotEqual(1, result.exit_code, "aborts with non-success code")
self.assertIn(
f"Wrote annotations to {existing_file.name}",
result.stdout,
"success message is correct on output file overwrite",
)
# TODO: # TODO: