diff --git a/requirements-training.txt b/requirements-training.txt new file mode 100644 index 0000000000..71243a1d52 --- /dev/null +++ b/requirements-training.txt @@ -0,0 +1,8 @@ +cerberus +flatbuffers +h5py +numpy >= 1.16.6 +onnx +packaging +protobuf +sympy \ No newline at end of file diff --git a/setup.py b/setup.py index 1320ed5793..d774fe27c9 100644 --- a/setup.py +++ b/setup.py @@ -232,11 +232,14 @@ packages = [ 'onnxruntime.transformers.longformer', ] +requirements_file = "requirements.txt" + if '--enable_training' in sys.argv: packages.extend(['onnxruntime.training', 'onnxruntime.training.amp', 'onnxruntime.training.optim']) sys.argv.remove('--enable_training') + requirements_file = "requirements-training.txt" package_data = {} data_files = [] @@ -310,12 +313,12 @@ if bdist_wheel is not None : cmd_classes['bdist_wheel'] = bdist_wheel cmd_classes['build_ext'] = build_ext -requirements_path = path.join(getcwd(), "requirements.txt") +requirements_path = path.join(getcwd(), requirements_file) if not path.exists(requirements_path): this = path.dirname(__file__) - requirements_path = path.join(this, "requirements.txt") + requirements_path = path.join(this, requirements_file) if not path.exists(requirements_path): - raise FileNotFoundError("Unable to find 'requirements.txt'") + raise FileNotFoundError("Unable to find " + requirements_file) with open(requirements_path) as f: install_requires = f.read().splitlines()