Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 2 additions & 10 deletions src/accelerate/commands/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
import argparse
import os
import platform
import subprocess
from shutil import which

import numpy as np
import psutil
Expand Down Expand Up @@ -82,15 +82,7 @@ def env_command(args):
if args.config_file is not None or os.path.isfile(default_config_file):
accelerate_config = load_config_from_file(args.config_file).to_dict()

# if we can run which, get it
command = None
bash_location = "Not found"
if os.name == "nt":
command = ["where", "accelerate"]
elif os.name == "posix":
command = ["which", "accelerate"]
if command is not None:
bash_location = subprocess.check_output(command, text=True, stderr=subprocess.STDOUT).strip()
bash_location = which("accelerate") or "Not found"
info = {
"`Accelerate` version": version,
"Platform": platform.platform(),
Expand Down
11 changes: 11 additions & 0 deletions tests/test_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import torch
from huggingface_hub.utils import GatedRepoError

import accelerate.commands.env as accelerate_env_cmd
import accelerate.commands.test as accelerate_test_cmd
from accelerate.commands.config.config_args import BaseConfig, ClusterConfig, SageMakerConfig, load_config_from_file
from accelerate.commands.estimate import estimate_command, estimate_command_parser, gather_data
Expand All @@ -43,6 +44,16 @@
from accelerate.utils.launch import prepare_simple_launcher_cmd_env


class EnvCommandTester(unittest.TestCase):
@patch("accelerate.commands.env.which", return_value=None)
def test_executable_not_found(self, _):
args = accelerate_env_cmd.env_command_parser().parse_args([])

info = accelerate_env_cmd.env_command(args)

self.assertEqual(info["`accelerate` bash location"], "Not found")


class AccelerateLauncherTester(unittest.TestCase):
"""
Test case for verifying the `accelerate launch` CLI operates correctly.
Expand Down
Loading