Skip to content

Commit 9154e45

Browse files
committed
fix ci failed test
1 parent 703edac commit 9154e45

13 files changed

Lines changed: 60 additions & 30 deletions

File tree

.github/workflows/fleet-release-test.yml

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1276,14 +1276,13 @@ jobs:
12761276
docker exec -i -e "python_version=${{ env.python_version }}" ${{ env.container_name }} /bin/bash <<-DOCKER_EOF
12771277
echo "Change python version to $python_version"
12781278
set -e
1279-
cp ${{ github.workspace }}/../../../Miniconda3-latest-Linux-x86_64.sh /tmp/miniconda.sh
1280-
bash /tmp/miniconda.sh -b -p /opt/conda
1281-
export PATH=/opt/conda/bin:$PATH
1282-
conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/main
1283-
conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/r
1284-
conda create -n py_$python_version python=$python_version -y
1279+
cp /home/models/Miniforge3-Linux-x86_64.sh /tmp/miniforge.sh
1280+
bash /tmp/miniforge.sh -b -p /opt/conda
1281+
/opt/conda/bin/conda create -n py_\$python_version python=\$python_version -y \
1282+
--override-channels \
1283+
-c conda-forge
12851284
. /opt/conda/etc/profile.d/conda.sh
1286-
conda activate py_$python_version
1285+
conda activate py_\$python_version
12871286
python --version
12881287
DOCKER_EOF
12891288
- name: Test

paddleformers/trainer/argparser.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
import json
2121
import os
2222
import sys
23+
import types
2324
from argparse import (
2425
ArgumentDefaultsHelpFormatter,
2526
ArgumentParser,
@@ -112,6 +113,15 @@ def _parse_dataclass_field(
112113
)
113114

114115
origin_type = getattr(field.type, "__origin__", field.type)
116+
# Python 3.10+ `str | None` syntax produces types.UnionType, normalize it to typing.Union
117+
if isinstance(field.type, types.UnionType):
118+
args = field.type.__args__
119+
field.type = (
120+
Optional[args[0]]
121+
if len(args) == 2 and type(None) in args
122+
else Union[args]
123+
)
124+
origin_type = Union
115125
if origin_type is Union:
116126
if (
117127
len(field.type.__args__) != 2

paddleformers/utils/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,7 @@
5252
"is_datasets_available",
5353
"is_transformers_available",
5454
"is_paddleformers_available",
55+
"is_paddlefleet_available",
5556
"dynamic_graph_pybind_context",
5657
"custom_import",
5758
"direct_paddleformers_import",

paddleformers/utils/import_utils.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -246,6 +246,16 @@ def is_paddle_available() -> bool:
246246
return is_package_available("paddle")
247247

248248

249+
def is_paddlefleet_available() -> bool:
250+
"""Check if paddleformers.fleet is available (PaddleFleet was merged into PaddleFormers)."""
251+
try:
252+
import paddleformers.fleet # noqa: F401
253+
254+
return True
255+
except ImportError:
256+
return False
257+
258+
249259
def is_tiktoken_available():
250260
return importlib.util.find_spec("tiktoken") is not None
251261

tests/fleet/multi_card_tests/pipeline_parallel/test_vpp_balanced_memory_with_shared_weight.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,9 @@
1919
0,
2020
os.path.dirname(
2121
os.path.dirname(
22-
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
22+
os.path.dirname(
23+
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
24+
)
2325
)
2426
),
2527
)

tests/fleet/single_card_tests/ai_edited_test/extensions/test_ai_block_mask_utils_2.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,9 @@
6565
_project_root = os.path.dirname(
6666
os.path.dirname(
6767
os.path.dirname(
68-
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
68+
os.path.dirname(
69+
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
70+
)
6971
)
7072
)
7173
)

tests/fleet/single_card_tests/ai_edited_test/extensions/test_ai_rr_attn_estimate_triton_op_3.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,9 @@
6565
_project_root = os.path.dirname(
6666
os.path.dirname(
6767
os.path.dirname(
68-
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
68+
os.path.dirname(
69+
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
70+
)
6971
)
7072
)
7173
)

tests/fleet/single_card_tests/ai_edited_test/fp8/test_ai_fp8_module_structure.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -35,8 +35,9 @@ def _load_fp8_module(mod_name, filename):
3535
"..",
3636
"..",
3737
"..",
38-
"src",
39-
"paddleformers.fleet",
38+
"..",
39+
"paddleformers",
40+
"fleet",
4041
"fp8",
4142
filename,
4243
)
@@ -74,8 +75,9 @@ def test_fp8_has_linear(self):
7475
"..",
7576
"..",
7677
"..",
77-
"src",
78-
"paddleformers.fleet",
78+
"..",
79+
"paddleformers",
80+
"fleet",
7981
"fp8",
8082
"linear.py",
8183
)

tests/fleet/single_card_tests/ai_edited_test/fp8/test_ai_fp8_utils.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,8 +39,9 @@ def _load_fp8_utils():
3939
"..",
4040
"..",
4141
"..",
42-
"src",
43-
"paddleformers.fleet",
42+
"..",
43+
"paddleformers",
44+
"fleet",
4445
"fp8",
4546
"utils.py",
4647
),

tests/fleet/single_card_tests/ai_edited_test/fp8/test_ai_quantization_3.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -41,8 +41,9 @@ def _load_quantization_module():
4141
"..",
4242
"..",
4343
"..",
44-
"src",
45-
"paddleformers.fleet",
44+
"..",
45+
"paddleformers",
46+
"fleet",
4647
"fp8",
4748
"quantization.py",
4849
),

0 commit comments

Comments
 (0)