From 72371b9f11958f03ebb2921f97839f16c04d6155 Mon Sep 17 00:00:00 2001 From: henryruhs Date: Wed, 5 Mar 2025 12:53:24 +0100 Subject: [PATCH] Add basic test for aad and unet --- .github/workflows/ci.yml | 19 +++++++++++++++++-- face_swapper/tests/test_networks.py | 19 +++++++++++++++++++ 2 files changed, 36 insertions(+), 2 deletions(-) create mode 100644 face_swapper/tests/test_networks.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8bbeef5..d22407d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -8,12 +8,27 @@ jobs: steps: - name: Checkout uses: actions/checkout@v4 - - name: Set up Python 3.10 + - name: Set up Python 3.12 uses: actions/setup-python@v5 with: - python-version: '3.10' + python-version: '3.12' - run: pip install flake8 - run: pip install flake8-import-order - run: pip install mypy - run: flake8 embedding_converter face_swapper - run: mypy embedding_converter face_swapper + test: + strategy: + matrix: + os: [ macos-latest, ubuntu-latest, windows-latest ] + runs-on: ${{ matrix.os }} + steps: + - name: Checkout + uses: actions/checkout@v4 + - name: Set up Python 3.12 + uses: actions/setup-python@v5 + with: + python-version: '3.12' + - run: pip install torch torchvision + - run: pip install pytest + - run: pytest diff --git a/face_swapper/tests/test_networks.py b/face_swapper/tests/test_networks.py new file mode 100644 index 0000000..e1d4381 --- /dev/null +++ b/face_swapper/tests/test_networks.py @@ -0,0 +1,19 @@ +import torch +import pytest + +from face_swapper.src.networks.aad import AAD +from face_swapper.src.networks.unet import UNet + + +@pytest.mark.parametrize('output_size', [ 256 ]) +def test_aad_with_unet(output_size : int) -> None: + generator = AAD(512, 4096, output_size, 2).eval() + encoder = UNet(output_size).eval() + + source_tensor = torch.randn(1, 512) + target_tensor = torch.randn(1, 3, output_size, output_size) + + target_attributes = encoder(target_tensor) + output_tensor = generator(source_tensor, target_attributes) + + assert output_tensor.shape == (1, 3, output_size, output_size)