Compare commits
281
Commits
0bbf53d495
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
70492acafe | ||
|
|
06913064c3 | ||
|
|
7190b6c438 | ||
|
|
7ce7a368c7 | ||
|
|
a2b818bbcc | ||
|
|
08062ad87d | ||
|
|
7176dfdf90 | ||
|
|
6399b88ade | ||
|
|
c5bcd09bba | ||
|
|
7d8af1c582 | ||
|
|
a379270add | ||
|
|
0b54ccaadb | ||
|
|
ac8ca830e1 | ||
|
|
49d2c48d31 | ||
|
|
280dbb579e | ||
|
|
ab7e346939 | ||
|
|
ad82ad4591 | ||
|
|
400706e730 | ||
|
|
096ffd56e9 | ||
|
|
790204af6a | ||
|
|
14a4125688 | ||
|
|
62519c371e | ||
|
|
9c8089d091 | ||
|
|
77a3219f39 | ||
|
|
66370a33c2 | ||
|
|
5870c3b0a3 | ||
|
|
e9340f17c4 | ||
|
|
282fee0806 | ||
|
|
33d30941a6 | ||
|
|
cb8d17cbc3 | ||
|
|
df2c2b7e5f | ||
|
|
6756ae4871 | ||
|
|
fa30aae3a8 | ||
|
|
15f3944d27 | ||
|
|
e6b73fe4fc | ||
|
|
1f72a2db8d | ||
|
|
a610bf4d89 | ||
|
|
9af3d9bc54 | ||
|
|
bbbeb403a8 | ||
|
|
9be7ec9e2e | ||
|
|
a6c717a401 | ||
|
|
4db59d2801 | ||
|
|
3a0c1e9067 | ||
|
|
5a8350d0ad | ||
|
|
62b6224de7 | ||
|
|
68bfff87a7 | ||
|
|
ac21c7a111 | ||
|
|
a18b4d304c | ||
|
|
a2f892e35c | ||
|
|
4bcfbbed0a | ||
|
|
31be0d7aac | ||
|
|
d4c9738db4 | ||
|
|
813daa0c90 | ||
|
|
29bf41c7da | ||
|
|
d9bd5be46d | ||
|
|
c2a2c34967 | ||
|
|
9e66981d3e | ||
|
|
6840ac72f1 | ||
|
|
ee304f1a75 | ||
|
|
4827a3f7dd | ||
|
|
d92051f5c4 | ||
|
|
10ce0d9124 | ||
|
|
88a20c4b0b | ||
|
|
090a1cd8bb | ||
|
|
ac6a305da8 | ||
|
|
cbedc395bc | ||
|
|
0c762b55b2 | ||
|
|
99b5820ae6 | ||
|
|
8214bd2669 | ||
|
|
aff0ae8fc6 | ||
|
|
16ad2b80ee | ||
|
|
88698dd735 | ||
|
|
f28bff7e47 | ||
|
|
9ee27b65a8 | ||
|
|
de4f984110 | ||
|
|
3b677e5e97 | ||
|
|
226fb1a206 | ||
|
|
7a1e17c29f | ||
|
|
f52720685e | ||
|
|
a3da7b048d | ||
|
|
f61eb59afb | ||
|
|
a430366bbe | ||
|
|
06d9753c46 | ||
|
|
4e6499cee3 | ||
|
|
7da3482366 | ||
|
|
7b3374d451 | ||
|
|
b17c1745ab | ||
|
|
db4ce6b425 | ||
|
|
aa6a097e4f | ||
|
|
150a213ae2 | ||
|
|
d860b71b1a | ||
|
|
1a3ee96932 | ||
|
|
8e366f26bc | ||
|
|
8a602264e6 | ||
|
|
99d5c88c8d | ||
|
|
c9da7378a1 | ||
|
|
b2773ecf09 | ||
|
|
084a13426a | ||
|
|
499996329f | ||
|
|
ce08ee4ccd | ||
|
|
cbbb044177 | ||
|
|
c630498168 | ||
|
|
f9d23c5bd4 | ||
|
|
3fa79faa9e | ||
|
|
556e26f21d | ||
|
|
01ee4664ff | ||
|
|
169e0531d9 | ||
|
|
3c46275015 | ||
|
|
0b86acd1a9 | ||
|
|
cddfced177 | ||
|
|
53d179663b | ||
|
|
6be837025b | ||
|
|
2854f74db4 | ||
|
|
462f7c79b4 | ||
|
|
0c9537975c | ||
|
|
10bfac73bd | ||
|
|
bdf3e73ba9 | ||
|
|
42400d6f32 | ||
|
|
0676011f98 | ||
|
|
e189b7c1f7 | ||
|
|
dc535167fa | ||
|
|
b722ddcb83 | ||
|
|
b755a81ef9 | ||
|
|
ae18041d6e | ||
|
|
d8e3d077e3 | ||
|
|
b54e2ed541 | ||
|
|
30f37b6ca5 | ||
|
|
860150905d | ||
|
|
66c3f38dba | ||
|
|
1a6cf69346 | ||
|
|
6f06485776 | ||
|
|
c7fd78dcf8 | ||
|
|
7a8eb93c48 | ||
|
|
4ba68d056b | ||
|
|
ba7a7299c2 | ||
|
|
6dba06704c | ||
|
|
20831c9bc3 | ||
|
|
e2338ad710 | ||
|
|
d1ee43e135 | ||
|
|
909f8f220c | ||
|
|
7fa88166ee | ||
|
|
321defb7ff | ||
|
|
00f1c07cfa | ||
|
|
312f525a46 | ||
|
|
78363988f7 | ||
|
|
b109ddb2b8 | ||
|
|
5b446a9f2f | ||
|
|
b462056e24 | ||
|
|
bdc0535ea4 | ||
|
|
5a1de4f2b2 | ||
|
|
4e499da168 | ||
|
|
8bfdec067b | ||
|
|
9d45f8a32b | ||
|
|
082ccf8382 | ||
|
|
8f10c16f3b | ||
|
|
5a674603ca | ||
|
|
ba0fc050fe | ||
|
|
7fc4d20262 | ||
|
|
d6036a95a3 | ||
|
|
14ccaace95 | ||
|
|
6d327bee22 | ||
|
|
270dc8e4d1 | ||
|
|
e49a0c37a6 | ||
|
|
a0726857a0 | ||
|
|
bf44947024 | ||
|
|
4f22b9755c | ||
|
|
ae3ac714d4 | ||
|
|
ad069269ac | ||
|
|
a54e7bbbb4 | ||
|
|
1f62054f36 | ||
|
|
59048a6862 | ||
|
|
a4764b8fc5 | ||
|
|
bf4cc9c483 | ||
|
|
f35cdefc21 | ||
|
|
493da70ba4 | ||
|
|
b415137e25 | ||
|
|
cead3d24c8 | ||
|
|
0bfbb1920c | ||
|
|
164102557e | ||
|
|
cf923a228d | ||
|
|
9f4a09a18c | ||
|
|
849b59e575 | ||
|
|
4a7e33b527 | ||
|
|
e5fad6d765 | ||
|
|
be6833ac6a | ||
|
|
6541a47668 | ||
|
|
96ccdb5c90 | ||
|
|
6e9c0a2626 | ||
|
|
396bf29f32 | ||
|
|
05774364fc | ||
|
|
74a9538feb | ||
|
|
9423c3378a | ||
|
|
e6fdbfe0b0 | ||
|
|
6c354b23bd | ||
|
|
3ad6c161b1 | ||
|
|
a1ff6a6f33 | ||
|
|
cf4f67d7a5 | ||
|
|
7c968db8ac | ||
|
|
e656cd4e12 | ||
|
|
8b1893b30a | ||
|
|
96d868a40c | ||
|
|
e698d058ff | ||
|
|
8aafbdae6c | ||
|
|
ec4162e34d | ||
|
|
99b7067cfe | ||
|
|
e62435cf22 | ||
|
|
3560a61277 | ||
|
|
7f0cbea937 | ||
|
|
31156eea54 | ||
|
|
ef11c33675 | ||
|
|
89aecfdef7 | ||
|
|
68e04c7653 | ||
|
|
a12f8c46e8 | ||
|
|
ff88dc1719 | ||
|
|
6502ea043a | ||
|
|
9cd4f70194 | ||
|
|
2cf55efcb3 | ||
|
|
d55adf87ac | ||
|
|
a621525d9e | ||
|
|
ba1ba027f0 | ||
|
|
24e1c44f5e | ||
|
|
94801efa6a | ||
|
|
4f7e8150e9 | ||
|
|
e9b5250e95 | ||
|
|
c1f38e3d77 | ||
|
|
9a107c7bb7 | ||
|
|
74d7307219 | ||
|
|
bdd0dfff10 | ||
|
|
e913bb28cc | ||
|
|
f9b845abd9 | ||
|
|
d46ad01e1d | ||
|
|
8d887b40b3 | ||
|
|
2ee5ab2b9c | ||
|
|
962d028c07 | ||
|
|
eaca23f27a | ||
|
|
0bd4008e72 | ||
|
|
123cab2500 | ||
|
|
d09cb9dd20 | ||
|
|
e1075b3551 | ||
|
|
d537cdf11d | ||
|
|
b8061466ce | ||
|
|
119282fee9 | ||
|
|
2e060f0776 | ||
|
|
9d0648f3c1 | ||
|
|
b093154193 | ||
|
|
f3aab68d5f | ||
|
|
0ebf7e6ce1 | ||
|
|
7229773371 | ||
|
|
ef985f6687 | ||
|
|
599769a8f6 | ||
|
|
70dca8f598 | ||
|
|
28d86e8762 | ||
|
|
fb278b641c | ||
|
|
3eb5f5e826 | ||
|
|
4a4ee95e31 | ||
|
|
4f8f32cf93 | ||
|
|
4e7f6cdee6 | ||
|
|
6a36215d82 | ||
|
|
8620d43e74 | ||
|
|
d26eeaddd1 | ||
|
|
b9d883484f | ||
|
|
833724912e | ||
|
|
717bb352e9 | ||
|
|
cb8e907ecb | ||
|
|
084273f5bc | ||
|
|
b98234267f | ||
|
|
11aed4497c | ||
|
|
93af7a69be | ||
|
|
385eb79f0e | ||
|
|
dcd01849ec | ||
|
|
cf03275da0 | ||
|
|
cecdf546b8 | ||
|
|
bd9c8136d1 | ||
|
|
675835a56b | ||
|
|
68a542be04 | ||
|
|
2ab4cf3b71 | ||
|
|
0fbdc3865c | ||
|
|
d516578d4b | ||
|
|
1579b15b4e | ||
|
|
9fbed5a9c4 | ||
|
|
6f0513ea2c |
@@ -1,25 +0,0 @@
|
|||||||
name: Code Quality Pipeline
|
|
||||||
run-name: ${{ gitea.actor }} is running the Code Quality Pipeline
|
|
||||||
on: push
|
|
||||||
jobs:
|
|
||||||
test:
|
|
||||||
name: Test
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
- name: Checkout Code
|
|
||||||
uses: actions/checkout@v3
|
|
||||||
- name: Setup Environment
|
|
||||||
uses: https://github.com/actions/setup-python@v3
|
|
||||||
with:
|
|
||||||
python-verison: "3.12"
|
|
||||||
architecture: "x64"
|
|
||||||
- name: Install Packages
|
|
||||||
run: |
|
|
||||||
pip install poetry
|
|
||||||
poetry install
|
|
||||||
- name: PEP8 Check
|
|
||||||
run: |
|
|
||||||
poetry run flake8 . --benchmark
|
|
||||||
- name: Type Check
|
|
||||||
run: |
|
|
||||||
poetry run mypy . --disable-error-code=import-untyped
|
|
||||||
@@ -1,9 +1,8 @@
|
|||||||
name: CI Pipeline
|
name: CI Pipeline
|
||||||
run-name: ${{ gitea.actor }} is running the CI Pipeline
|
run-name: ${{ gitea.actor }} is running the CI Pipeline
|
||||||
on:
|
on:
|
||||||
pull_request:
|
release:
|
||||||
branches:
|
types: [published]
|
||||||
- main
|
|
||||||
jobs:
|
jobs:
|
||||||
test:
|
test:
|
||||||
name: Test
|
name: Test
|
||||||
@@ -17,6 +16,9 @@ jobs:
|
|||||||
python-verison: "3.12"
|
python-verison: "3.12"
|
||||||
architecture: "x64"
|
architecture: "x64"
|
||||||
- name: Install Packages
|
- name: Install Packages
|
||||||
|
env:
|
||||||
|
PIP_INDEX_URL: http://192.168.1.2:5001/index/
|
||||||
|
PIP_TRUSTED_HOST: 192.168.1.2
|
||||||
run: |
|
run: |
|
||||||
pip install poetry
|
pip install poetry
|
||||||
poetry install
|
poetry install
|
||||||
@@ -25,7 +27,7 @@ jobs:
|
|||||||
poetry run flake8 . --benchmark
|
poetry run flake8 . --benchmark
|
||||||
- name: Type Check
|
- name: Type Check
|
||||||
run: |
|
run: |
|
||||||
poetry run mypy . --disable-error-code=import-untyped
|
poetry run mypy .
|
||||||
publish:
|
publish:
|
||||||
name: Build and Publish
|
name: Build and Publish
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
@@ -37,9 +39,9 @@ jobs:
|
|||||||
-
|
-
|
||||||
dockerfile: ./Dockerfile.web_ui
|
dockerfile: ./Dockerfile.web_ui
|
||||||
image: ${{ vars.docker_repo_url }}/${{ gitea.repository }}/web_ui
|
image: ${{ vars.docker_repo_url }}/${{ gitea.repository }}/web_ui
|
||||||
# -
|
-
|
||||||
# dockerfile: ./model/Dockerfile
|
dockerfile: ./Dockerfile.model
|
||||||
# image: ${{ vars.docker_repo_url }}/${{ gitea.repository }}/model
|
image: ${{ vars.docker_repo_url }}/${{ gitea.repository }}/model
|
||||||
steps:
|
steps:
|
||||||
-
|
-
|
||||||
name: Checkout Code
|
name: Checkout Code
|
||||||
|
|||||||
@@ -0,0 +1,50 @@
|
|||||||
|
name: Code Quality Pipeline
|
||||||
|
on:
|
||||||
|
pull_request:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
jobs:
|
||||||
|
job:
|
||||||
|
name: Check Code
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- name: Checkout Code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
- name: Setup Environment
|
||||||
|
uses: https://github.com/actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.12"
|
||||||
|
architecture: "x64"
|
||||||
|
- name: Setup poetry
|
||||||
|
env:
|
||||||
|
POETRY_VERSION: 2.1.1
|
||||||
|
POETRY_HOME: /opt/poetry
|
||||||
|
POETRY_NO_INTERACTION: 1
|
||||||
|
POETRY_NO_CACHE: 1
|
||||||
|
run: |
|
||||||
|
curl -sSL https://install.python-poetry.org | python3 -
|
||||||
|
export PATH=$POETRY_HOME/bin:$PATH
|
||||||
|
poetry --version
|
||||||
|
- name: Install Dependencies
|
||||||
|
env:
|
||||||
|
PIP_INDEX_URL: ${{ vars.PIP_INDEX_URL }}
|
||||||
|
PIP_TRUSTED_HOST: ${{ vars.PIP_TRUSTED_HOST }}
|
||||||
|
run: |
|
||||||
|
/opt/poetry/bin/poetry install
|
||||||
|
- name: PEP8 Check
|
||||||
|
run: |
|
||||||
|
/opt/poetry/bin/poetry run flake8 . --benchmark
|
||||||
|
- name: Type Check
|
||||||
|
run: |
|
||||||
|
/opt/poetry/bin/poetry run mypy .
|
||||||
|
- name: Pytest & Calculate Coverage
|
||||||
|
env:
|
||||||
|
MINIO_ENDPOINT: ${{ vars.MINIO_ENDPOINT }}
|
||||||
|
MINIO_ACCESS_KEY: ${{ vars.MINIO_ACCESS_KEY }}
|
||||||
|
MINIO_SECRET_KEY: ${{ secrets.MINIO_SECRET_KEY }}
|
||||||
|
MONGO_ENDPOINT: ${{ vars.MONGO_ENDPOINT }}
|
||||||
|
run: |
|
||||||
|
/opt/poetry/bin/poetry run coverage run -m pytest .
|
||||||
|
- name: Coverage Report
|
||||||
|
run: |
|
||||||
|
/opt/poetry/bin/poetry run coverage report -m
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
name: release-tag
|
||||||
|
on:
|
||||||
|
workflow_dispatch:
|
||||||
|
inputs:
|
||||||
|
version:
|
||||||
|
description: 'The version number to deploy'
|
||||||
|
required: true
|
||||||
|
jobs:
|
||||||
|
release-image:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
-
|
||||||
|
run: echo "manual trigger executed"
|
||||||
|
-
|
||||||
|
run: echo "$MESSAGE"
|
||||||
|
env:
|
||||||
|
MESSAGE; ${{ github.event.inputs.version}}
|
||||||
@@ -177,3 +177,6 @@ ipython_config.py
|
|||||||
|
|
||||||
# Remove previous ipynb_checkpoints
|
# Remove previous ipynb_checkpoints
|
||||||
# git rm -r .ipynb_checkpoints/
|
# git rm -r .ipynb_checkpoints/
|
||||||
|
|
||||||
|
# Logging data
|
||||||
|
runs/*
|
||||||
|
|||||||
+33
-15
@@ -1,44 +1,62 @@
|
|||||||
repos:
|
repos:
|
||||||
- repo: https://github.com/pre-commit/pre-commit-hooks
|
- repo: https://github.com/pre-commit/pre-commit-hooks
|
||||||
rev: v4.5.0
|
rev: v5.0.0
|
||||||
hooks:
|
hooks:
|
||||||
- id: trailing-whitespace
|
- id: trailing-whitespace
|
||||||
- id: end-of-file-fixer
|
- id: end-of-file-fixer
|
||||||
- id: check-yaml
|
- id: check-yaml
|
||||||
|
- id: check-added-large-files
|
||||||
- id: debug-statements
|
- id: debug-statements
|
||||||
- id: double-quote-string-fixer
|
- id: double-quote-string-fixer
|
||||||
- id: name-tests-test
|
- id: name-tests-test
|
||||||
- id: requirements-txt-fixer
|
|
||||||
- repo: https://github.com/asottile/setup-cfg-fmt
|
- repo: https://github.com/asottile/setup-cfg-fmt
|
||||||
rev: v2.5.0
|
rev: v2.7.0
|
||||||
hooks:
|
hooks:
|
||||||
- id: setup-cfg-fmt
|
- id: setup-cfg-fmt
|
||||||
- repo: https://github.com/asottile/reorder-python-imports
|
- repo: https://github.com/pre-commit/mirrors-isort
|
||||||
rev: v3.12.0
|
rev: v5.10.1
|
||||||
hooks:
|
hooks:
|
||||||
- id: reorder-python-imports
|
- id: isort
|
||||||
exclude: ^(pre_commit/resources/|testing/resources/python3_hooks_repo/)
|
language_version: python3.10
|
||||||
args: [--py39-plus, --add-import, 'from __future__ import annotations']
|
args: [ --tc ]
|
||||||
- repo: https://github.com/asottile/add-trailing-comma
|
- repo: https://github.com/asottile/add-trailing-comma
|
||||||
rev: v3.1.0
|
rev: v3.1.0
|
||||||
hooks:
|
hooks:
|
||||||
- id: add-trailing-comma
|
- id: add-trailing-comma
|
||||||
- repo: https://github.com/asottile/pyupgrade
|
- repo: https://github.com/asottile/pyupgrade
|
||||||
rev: v3.15.1
|
rev: v3.19.1
|
||||||
hooks:
|
hooks:
|
||||||
- id: pyupgrade
|
- id: pyupgrade
|
||||||
args: [--py39-plus]
|
args: [--py39-plus]
|
||||||
- repo: https://github.com/hhatto/autopep8
|
|
||||||
rev: v2.0.4
|
|
||||||
hooks:
|
|
||||||
- id: autopep8
|
|
||||||
- repo: https://github.com/PyCQA/flake8
|
- repo: https://github.com/PyCQA/flake8
|
||||||
rev: 7.0.0
|
rev: 7.0.0
|
||||||
hooks:
|
hooks:
|
||||||
- id: flake8
|
- id: flake8
|
||||||
|
entry: pflake8
|
||||||
|
additional_dependencies:
|
||||||
|
- "pyproject-flake8"
|
||||||
- repo: https://github.com/pre-commit/mirrors-mypy
|
- repo: https://github.com/pre-commit/mirrors-mypy
|
||||||
rev: v1.8.0
|
rev: v1.15.0
|
||||||
hooks:
|
hooks:
|
||||||
- id: mypy
|
- id: mypy
|
||||||
additional_dependencies: [types-all]
|
|
||||||
exclude: ^testing/resources/
|
exclude: ^testing/resources/
|
||||||
|
- repo: https://github.com/psf/black
|
||||||
|
rev: 25.1.0
|
||||||
|
hooks:
|
||||||
|
- id: black
|
||||||
|
language_version: python3.12
|
||||||
|
args: [ --skip-string-normalization ]
|
||||||
|
- repo: https://github.com/myint/docformatter
|
||||||
|
rev: v1.5.0
|
||||||
|
hooks:
|
||||||
|
- id: docformatter
|
||||||
|
name: docformatter
|
||||||
|
description: 'Formats docstrings to follow PEP 257.'
|
||||||
|
entry: docformatter
|
||||||
|
language: python
|
||||||
|
types: [ python ]
|
||||||
|
- repo: https://github.com/jendrikseipp/vulture
|
||||||
|
rev: 'v2.14'
|
||||||
|
hooks:
|
||||||
|
- id: vulture
|
||||||
|
entry: vulture . --min-confidence 90 --exclude */.venv/*.py,*/tests/*.py
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
# build stage
|
||||||
|
FROM python:3.12-slim-bookworm AS builder
|
||||||
|
|
||||||
|
# set environment variables
|
||||||
|
ENV PYTHONUNBUFFERED=1 \
|
||||||
|
PYTHONDONTWRITEBYTECODE=1 \
|
||||||
|
PIP_NO_CACHE_DIR=off \
|
||||||
|
PIP_DISABLE_PIP_VERSION_CHECK=ON \
|
||||||
|
PIP_DEFAULT_TIMEOUT=100 \
|
||||||
|
DEBIAN_FRONTEND=noninteractive \
|
||||||
|
POETRY_HOME=/etc/poetry \
|
||||||
|
POETRY_VERSION=1.7.1 \
|
||||||
|
POETRY_VIRTUALENVS_IN_PROJECT=1 \
|
||||||
|
POETRY_VIRTUALENVS_CREATE=1 \
|
||||||
|
POETRY_NO_INTERACTION=1 \
|
||||||
|
POETRY_CACHE_DIR=/tmp/poetry_cache \
|
||||||
|
APP_HOME=/home/app
|
||||||
|
|
||||||
|
# update system
|
||||||
|
RUN apt-get update \
|
||||||
|
&& apt-get install -y --no-install-recommends \
|
||||||
|
build-essential \
|
||||||
|
curl \
|
||||||
|
&& apt-get clean
|
||||||
|
|
||||||
|
# install poetry
|
||||||
|
RUN curl -sSL https://install.python-poetry.org | python3 -
|
||||||
|
ENV PATH="${POETRY_HOME}/bin:$PATH"
|
||||||
|
|
||||||
|
# install runtime dependencies
|
||||||
|
WORKDIR ${APP_HOME}
|
||||||
|
COPY ./poetry.lock ./pyproject.toml ./
|
||||||
|
RUN --mount=type=cache,target=${POETRY_CACHE_DIR} poetry install \
|
||||||
|
--with shared,model \
|
||||||
|
--no-root
|
||||||
|
|
||||||
|
# final stage
|
||||||
|
FROM python:3.12-slim-bookworm
|
||||||
|
|
||||||
|
# copy virtualenv made by poetry
|
||||||
|
ENV APP_HOME=/home/app \
|
||||||
|
VIRTUAL_ENV=/home/app/.venv
|
||||||
|
COPY --from=builder ${VIRTUAL_ENV} ${VIRTUAL_ENV}
|
||||||
|
ENV PATH="${VIRTUAL_ENV}/bin:${PATH}"
|
||||||
|
|
||||||
|
# create home directory and app user
|
||||||
|
RUN mkdir -p $APP_HOME
|
||||||
|
|
||||||
|
# add code while changing ownership
|
||||||
|
WORKDIR $APP_HOME
|
||||||
|
COPY ./shared ./shared
|
||||||
|
COPY ./model/src ./
|
||||||
|
|
||||||
|
ENTRYPOINT [ "python", "main.py" ]
|
||||||
+3
-2
@@ -1,5 +1,5 @@
|
|||||||
# build stage
|
# build stage
|
||||||
FROM python:3.12-slim-bookworm as BUILDER
|
FROM python:3.12-slim-bookworm AS builder
|
||||||
|
|
||||||
# set environment variables
|
# set environment variables
|
||||||
ENV PYTHONUNBUFFERED=1 \
|
ENV PYTHONUNBUFFERED=1 \
|
||||||
@@ -31,7 +31,8 @@ ENV PATH="${POETRY_HOME}/bin:$PATH"
|
|||||||
WORKDIR ${APP_HOME}
|
WORKDIR ${APP_HOME}
|
||||||
COPY ./poetry.lock ./pyproject.toml ./
|
COPY ./poetry.lock ./pyproject.toml ./
|
||||||
RUN --mount=type=cache,target=${POETRY_CACHE_DIR} poetry install \
|
RUN --mount=type=cache,target=${POETRY_CACHE_DIR} poetry install \
|
||||||
--with shared,web_ui
|
--with shared,web_ui \
|
||||||
|
--no-root
|
||||||
|
|
||||||
# final stage
|
# final stage
|
||||||
FROM python:3.12-slim-bookworm
|
FROM python:3.12-slim-bookworm
|
||||||
|
|||||||
@@ -5,8 +5,8 @@ from pathlib import Path
|
|||||||
|
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
from shared.database import connect
|
from shared.docstore import connect_mongodb
|
||||||
from shared.database.classes import VisualCommunication
|
from shared.docstore.classes import VisualCommunication
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
# prepare env vars
|
# prepare env vars
|
||||||
@@ -15,7 +15,7 @@ if __name__ == '__main__':
|
|||||||
load_dotenv(env_path)
|
load_dotenv(env_path)
|
||||||
os.environ['MONGO_HOST'] = 'localhost'
|
os.environ['MONGO_HOST'] = 'localhost'
|
||||||
# connect to database
|
# connect to database
|
||||||
collection, db, client = connect()
|
collection, db, client = connect_mongodb()
|
||||||
print(client.server_info())
|
print(client.server_info())
|
||||||
# download images
|
# download images
|
||||||
data = None
|
data = None
|
||||||
|
Before Width: | Height: | Size: 26 KiB After Width: | Height: | Size: 26 KiB |
|
Before Width: | Height: | Size: 29 KiB After Width: | Height: | Size: 29 KiB |
|
Before Width: | Height: | Size: 37 KiB After Width: | Height: | Size: 37 KiB |
@@ -0,0 +1,9 @@
|
|||||||
|
from shared.docstore.classes import ModelData
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
# instantiate data object
|
||||||
|
vis_com_list = [ModelData.from_random() for i in range(3)]
|
||||||
|
# generate random predictions
|
||||||
|
[vis_com.from_random() for vis_com in vis_com_list]
|
||||||
|
for vis_com in vis_com_list:
|
||||||
|
print(vis_com)
|
||||||
@@ -5,8 +5,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
from shared.database import connect
|
from shared.docstore import connect_mongodb, count_documents
|
||||||
from shared.database import count_documents
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
# prepare env vars
|
# prepare env vars
|
||||||
@@ -15,7 +14,7 @@ if __name__ == '__main__':
|
|||||||
load_dotenv(env_path)
|
load_dotenv(env_path)
|
||||||
os.environ['MONGO_HOST'] = 'localhost'
|
os.environ['MONGO_HOST'] = 'localhost'
|
||||||
# connect to database
|
# connect to database
|
||||||
collection, db, client = connect()
|
collection, db, client = connect_mongodb()
|
||||||
# get visual communication
|
# get visual communication
|
||||||
num_docs = count_documents(
|
num_docs = count_documents(
|
||||||
collection=collection,
|
collection=collection,
|
||||||
@@ -5,8 +5,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
from shared.database import connect
|
from shared.docstore import connect_mongodb, count_documents
|
||||||
from shared.database import count_documents
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
# prepare env vars
|
# prepare env vars
|
||||||
@@ -15,7 +14,7 @@ if __name__ == '__main__':
|
|||||||
load_dotenv(env_path)
|
load_dotenv(env_path)
|
||||||
os.environ['MONGO_HOST'] = 'localhost'
|
os.environ['MONGO_HOST'] = 'localhost'
|
||||||
# connect to database
|
# connect to database
|
||||||
collection, db, client = connect()
|
collection, db, client = connect_mongodb()
|
||||||
# get visual communication
|
# get visual communication
|
||||||
num_docs = count_documents(
|
num_docs = count_documents(
|
||||||
collection=collection,
|
collection=collection,
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
services:
|
||||||
|
vcda_model:
|
||||||
|
image: vcda_model:test
|
||||||
|
container_name: vcda_model_alone
|
||||||
|
build:
|
||||||
|
context: ../
|
||||||
|
dockerfile: ./Dockerfile.model
|
||||||
|
env_file:
|
||||||
|
- ../server.env
|
||||||
|
environment:
|
||||||
|
- ENV=TEST
|
||||||
|
deploy:
|
||||||
|
resources:
|
||||||
|
limits:
|
||||||
|
cpus: '2.000'
|
||||||
|
memory: 8G
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
{
|
||||||
|
"datasets": {
|
||||||
|
"train": {
|
||||||
|
"filelist": "datasets/train.csv"
|
||||||
|
},
|
||||||
|
"val": {
|
||||||
|
"filelist": "datasets/val.csv"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"backbone": {
|
||||||
|
"class": "model.src.models.VisualCommunicationModel",
|
||||||
|
"model_name": "pretrained"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,146 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import random
|
|
||||||
from io import BytesIO
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.utils.data import Dataset
|
|
||||||
from torchvision.io import read_image
|
|
||||||
from torchvision.transforms import ColorJitter
|
|
||||||
from torchvision.transforms import InterpolationMode
|
|
||||||
from torchvision.transforms import Normalize
|
|
||||||
from torchvision.transforms.functional import hflip
|
|
||||||
from torchvision.transforms.functional import pad
|
|
||||||
from torchvision.transforms.functional import resize
|
|
||||||
from torchvision.transforms.functional import rotate
|
|
||||||
|
|
||||||
from shared.database import connect
|
|
||||||
|
|
||||||
# resnet18 original normalization values
|
|
||||||
RESNET_NORMALIZE_MEAN = [0.485, 0.456, 0.406]
|
|
||||||
RESNET_NORMALIZE_STD = [0.229, 0.224, 0.225]
|
|
||||||
|
|
||||||
|
|
||||||
class VCDADataset(Dataset):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
data_name_list: list[str],
|
|
||||||
do_augment: bool = False,
|
|
||||||
random_annotations: bool = False,
|
|
||||||
normalize_mean: list[float] = RESNET_NORMALIZE_MEAN,
|
|
||||||
normalize_std: list[float] = RESNET_NORMALIZE_STD,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.data_name_list = data_name_list
|
|
||||||
self.do_augment = do_augment
|
|
||||||
self.random_annotations = random_annotations
|
|
||||||
self.normalize_mean = normalize_mean
|
|
||||||
self.normalize_std = normalize_std
|
|
||||||
# prepare augmentation functions
|
|
||||||
self.normalize = Normalize(
|
|
||||||
mean=normalize_mean,
|
|
||||||
std=normalize_std,
|
|
||||||
)
|
|
||||||
self.color_jitter = ColorJitter(
|
|
||||||
brightness=1e-1,
|
|
||||||
contrast=8e-2,
|
|
||||||
saturation=8e-2,
|
|
||||||
)
|
|
||||||
# connect to database
|
|
||||||
collection, _, _ = connect()
|
|
||||||
self.collection = collection
|
|
||||||
|
|
||||||
def __len__(self):
|
|
||||||
return len(self.data_name_list)
|
|
||||||
|
|
||||||
def __getitem__(self, idx):
|
|
||||||
# get image from database
|
|
||||||
name = self.data_name_list[idx]
|
|
||||||
query = {
|
|
||||||
'name': name,
|
|
||||||
}
|
|
||||||
projection = {
|
|
||||||
'_id': False,
|
|
||||||
'image': True,
|
|
||||||
}
|
|
||||||
img_bytes = self.collection.find_one(
|
|
||||||
filter=query,
|
|
||||||
projection=projection,
|
|
||||||
)
|
|
||||||
img = self.load_image(img_bytes)
|
|
||||||
if self.do_augment:
|
|
||||||
img = self.augment(img)
|
|
||||||
return img
|
|
||||||
|
|
||||||
def load_image(
|
|
||||||
self,
|
|
||||||
data: bytes,
|
|
||||||
) -> Tensor:
|
|
||||||
"""Load images tensor from bytes."""
|
|
||||||
assert isinstance(data, bytes)
|
|
||||||
img = read_image(BytesIO(data))
|
|
||||||
img /= 255 # normalize 8-bit image
|
|
||||||
img = self.square_pad(img)
|
|
||||||
img = resize(
|
|
||||||
img=img,
|
|
||||||
size=(512, 512),
|
|
||||||
interpolation=InterpolationMode.BICUBIC,
|
|
||||||
)
|
|
||||||
return img
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def square_pad(
|
|
||||||
img: Tensor,
|
|
||||||
) -> Tensor:
|
|
||||||
"""
|
|
||||||
Pads image to a square with side length
|
|
||||||
equal to the largest side of the input image.
|
|
||||||
"""
|
|
||||||
assert isinstance(img, Tensor)
|
|
||||||
# B, nc, w, h = img.shape
|
|
||||||
h = img.shape[-2]
|
|
||||||
w = img.shape[-1]
|
|
||||||
if h == w:
|
|
||||||
return img
|
|
||||||
max_wh = max([h, w])
|
|
||||||
hp = int((max_wh - w) / 2)
|
|
||||||
vp = int((max_wh - h) / 2)
|
|
||||||
padding = (hp, vp, hp, vp)
|
|
||||||
return pad(img, padding, 0, 'constant')
|
|
||||||
|
|
||||||
def augment(
|
|
||||||
self,
|
|
||||||
img: Tensor,
|
|
||||||
) -> Tensor:
|
|
||||||
"""
|
|
||||||
Augment image with random horizontal flips,
|
|
||||||
rotations and color jitter.
|
|
||||||
"""
|
|
||||||
assert isinstance(img, Tensor)
|
|
||||||
# left-right flip
|
|
||||||
if random.random() >= 0.5:
|
|
||||||
img = hflip(img)
|
|
||||||
# rotation
|
|
||||||
rnd = random.random()
|
|
||||||
if rnd < 0.25:
|
|
||||||
img = rotate(img, angle=90)
|
|
||||||
if rnd < 0.5:
|
|
||||||
img = rotate(img, angle=180)
|
|
||||||
if rnd < 0.75:
|
|
||||||
img = rotate(img, angle=270)
|
|
||||||
# color jitter
|
|
||||||
img = self.color_jitter(img)
|
|
||||||
return img
|
|
||||||
|
|
||||||
def reverse_normalise(
|
|
||||||
self,
|
|
||||||
img: Tensor,
|
|
||||||
) -> Tensor:
|
|
||||||
"""
|
|
||||||
Reverse normalization to get an image
|
|
||||||
that can be interpreted by humans.
|
|
||||||
"""
|
|
||||||
assert isinstance(img, Tensor)
|
|
||||||
img *= Tensor(self.normalize_std).reshape((3, 1, 1))
|
|
||||||
img += Tensor(self.normalize_mean).reshape((3, 1, 1))
|
|
||||||
return img
|
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
"""Main script to be run by service."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from traceback import print_exc
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from models import VisualCommunicationModel
|
||||||
|
from tqdm import tqdm
|
||||||
|
from utils import DEVICE, VCDADataset, load_model
|
||||||
|
|
||||||
|
from shared.utils import setup_logging
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
# setup logging
|
||||||
|
setup_logging()
|
||||||
|
# instantiate model
|
||||||
|
model: VisualCommunicationModel = load_model()
|
||||||
|
model.eval()
|
||||||
|
# setup dataset
|
||||||
|
dataset = VCDADataset(
|
||||||
|
data_name_list=[
|
||||||
|
'02dbaf48d713e4e6d3a6b98fd2dc866e',
|
||||||
|
],
|
||||||
|
do_augment=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
# make prediction
|
||||||
|
with torch.no_grad():
|
||||||
|
for i in tqdm(range(len(dataset))):
|
||||||
|
try:
|
||||||
|
# get image
|
||||||
|
image = dataset[i]
|
||||||
|
image = torch.unsqueeze(image, 0) # add artificial batch dimension
|
||||||
|
image = image.to(DEVICE)
|
||||||
|
# make prediction
|
||||||
|
pred: dict = model(image)
|
||||||
|
except Exception:
|
||||||
|
print_exc()
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
print(json.dumps(pred, indent=4))
|
||||||
|
logging.debug('finished')
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
2bbe2be17a663f4299657378ac5e993e
|
||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,26 +1,24 @@
|
|||||||
from __future__ import annotations
|
from torch import nn
|
||||||
|
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
|
|
||||||
class FullyConnectedModel(nn.Module):
|
class FullyConnectedModel(nn.Module):
|
||||||
|
"""Fully connected layers model template for intrepreting feature space-
|
||||||
|
output from ResNet18 head."""
|
||||||
|
|
||||||
def __init__(self, num_out_features: int):
|
def __init__(self, num_out_features: int):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
# define layers
|
# define layers
|
||||||
self.fc1 = nn.Linear(in_features=16*16*512, out_features=512)
|
self.fc1 = nn.Linear(in_features=512, out_features=128)
|
||||||
self.af1 = nn.ReLU()
|
self.af1 = nn.ReLU()
|
||||||
self.fc2 = nn.Linear(in_features=512, out_features=128)
|
self.fc2 = nn.Linear(in_features=128, out_features=32)
|
||||||
self.af2 = nn.ReLU()
|
self.af2 = nn.ReLU()
|
||||||
self.fc3 = nn.Linear(in_features=128, out_features=32)
|
self.fc3 = nn.Linear(in_features=32, out_features=num_out_features)
|
||||||
self.af3 = nn.ReLU()
|
|
||||||
self.fc4 = nn.Linear(in_features=32, out_features=num_out_features)
|
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
|
"""Pass input through model."""
|
||||||
x = self.fc1(x)
|
x = self.fc1(x)
|
||||||
x = self.af1(x)
|
x = self.af1(x)
|
||||||
x = self.fc2(x)
|
x = self.fc2(x)
|
||||||
x = self.af2(x)
|
x = self.af2(x)
|
||||||
x = self.fc3(x)
|
x = self.fc3(x)
|
||||||
x = self.af3(x)
|
|
||||||
x = self.fc4(x)
|
|
||||||
return x
|
return x
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,11 +5,15 @@ import torchvision
|
|||||||
|
|
||||||
|
|
||||||
class ResNet18Head(nn.Module):
|
class ResNet18Head(nn.Module):
|
||||||
def __init__(self):
|
def __init__(self, download_resnet_weights: bool = False):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
# copy out parts from ResNet18 with weights
|
# copy out parts from ResNet18 with weights
|
||||||
|
if download_resnet_weights:
|
||||||
|
weights = torchvision.models.ResNet18_Weights.IMAGENET1K_V1
|
||||||
|
else:
|
||||||
|
weights = None
|
||||||
resnet18 = torchvision.models.resnet18(
|
resnet18 = torchvision.models.resnet18(
|
||||||
weights=torchvision.models.ResNet18_Weights.IMAGENET1K_V1,
|
weights=weights,
|
||||||
)
|
)
|
||||||
# save relevant layers
|
# save relevant layers
|
||||||
self.conv1 = resnet18.conv1
|
self.conv1 = resnet18.conv1
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
|
"""Definition of VisualCommunicationModel class."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import torch.nn as nn
|
from torch import nn
|
||||||
|
|
||||||
from .angle import AngleTail
|
from .angle import AngleTail
|
||||||
from .contact import ContactTail
|
from .contact import ContactTail
|
||||||
@@ -14,24 +16,15 @@ from .point_of_view import PointOfViewTail
|
|||||||
from .resnet18_head import ResNet18Head
|
from .resnet18_head import ResNet18Head
|
||||||
from .salience import SalienceTail
|
from .salience import SalienceTail
|
||||||
from .visual_syntax import VisualSyntaxTail
|
from .visual_syntax import VisualSyntaxTail
|
||||||
from shared.dto import AngleData
|
|
||||||
from shared.dto import ContactData
|
|
||||||
from shared.dto import DistanceData
|
|
||||||
from shared.dto import FramingData
|
|
||||||
from shared.dto import InformationValueData
|
|
||||||
from shared.dto import ModalityColorData
|
|
||||||
from shared.dto import ModalityDepthData
|
|
||||||
from shared.dto import ModalityLightingData
|
|
||||||
from shared.dto import PointOfViewData
|
|
||||||
from shared.dto import SalienceData
|
|
||||||
from shared.dto import VisualSyntaxData
|
|
||||||
|
|
||||||
|
|
||||||
class VisualCommunicationModel(nn.Module):
|
class VisualCommunicationModel(nn.Module):
|
||||||
def __init__(self):
|
"""Visual communication model."""
|
||||||
|
|
||||||
|
def __init__(self, download_resnet_weights: bool = False):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
# store other models
|
# store other models
|
||||||
self.resnet_head = ResNet18Head()
|
self.resnet_head = ResNet18Head(download_resnet_weights)
|
||||||
self.visual_syntax_tail = VisualSyntaxTail()
|
self.visual_syntax_tail = VisualSyntaxTail()
|
||||||
self.contact_tail = ContactTail()
|
self.contact_tail = ContactTail()
|
||||||
self.angle_tail = AngleTail()
|
self.angle_tail = AngleTail()
|
||||||
@@ -44,75 +37,22 @@ class VisualCommunicationModel(nn.Module):
|
|||||||
self.framing_tail = FramingTail()
|
self.framing_tail = FramingTail()
|
||||||
self.salience_tail = SalienceTail()
|
self.salience_tail = SalienceTail()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x) -> dict:
|
||||||
|
"""Calculate model output on data."""
|
||||||
# generate visual representation
|
# generate visual representation
|
||||||
vis_rep = self.resnet_head(x)
|
features = self.resnet_head(x)
|
||||||
# prepare result map
|
# make predictions
|
||||||
results = {}
|
prediction_dict = {
|
||||||
# predict visual syntax
|
'visual_syntax': self.visual_syntax_tail(features),
|
||||||
results['visual_syntax'] = VisualSyntaxData.from_list(
|
'contact': self.contact_tail(features),
|
||||||
self.visual_syntax_tail(
|
'angle': self.angle_tail(features),
|
||||||
vis_rep,
|
'point_of_view': self.point_of_view_tail(features),
|
||||||
).cpu(),
|
'distance': self.distance_tail(features),
|
||||||
)
|
'modality_lighting': self.modality_lighting_tail(features),
|
||||||
# predict contact
|
'modality_color': self.modality_color_tail(features),
|
||||||
results['contact'] = ContactData.from_list(
|
'modality_depth': self.modality_depth_tail(features),
|
||||||
self.contact_tail(
|
'information_value': self.information_value_tail(features),
|
||||||
vis_rep,
|
'framing': self.framing_tail(features),
|
||||||
).cpu(),
|
'salience': self.salience_tail(features),
|
||||||
)
|
}
|
||||||
# predict angle
|
return prediction_dict
|
||||||
results['angle'] = AngleData.from_list(
|
|
||||||
self.angle_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict point of view
|
|
||||||
results['point_of_view'] = PointOfViewData.from_list(
|
|
||||||
self.point_of_view_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict distance
|
|
||||||
results['distance'] = DistanceData.from_list(
|
|
||||||
self.distance_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict modality lighting
|
|
||||||
results['modality_lighting'] = ModalityLightingData.from_list(
|
|
||||||
self.modality_lighting_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict modality color
|
|
||||||
results['modality_color'] = ModalityColorData.from_list(
|
|
||||||
self.modality_color_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict modality depth
|
|
||||||
results['modality_depth'] = ModalityDepthData.from_list(
|
|
||||||
self.modality_depth_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict information value
|
|
||||||
results['information_value'] = InformationValueData.from_list(
|
|
||||||
self.information_value_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict framing
|
|
||||||
results['framing'] = FramingData.from_list(
|
|
||||||
self.framing_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
# predict salience
|
|
||||||
results['salience'] = SalienceData.from_list(
|
|
||||||
self.salience_tail(
|
|
||||||
vis_rep,
|
|
||||||
).cpu(),
|
|
||||||
)
|
|
||||||
return results
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .fully_connected import FullyConnectedModel
|
from .fully_connected import FullyConnectedModel
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .get_class import get_class
|
||||||
|
from .load_model import DEVICE, load_model
|
||||||
|
from .vcda_dataset import VCDADataset
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
from importlib import import_module
|
||||||
|
|
||||||
|
from torch.nn import Module
|
||||||
|
|
||||||
|
|
||||||
|
def get_class(path: str) -> Module:
|
||||||
|
parts = path.split('.')
|
||||||
|
module_path = '.'.join(parts[:-1])
|
||||||
|
class_name = parts[-1]
|
||||||
|
module = import_module(module_path)
|
||||||
|
return getattr(module, class_name)
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
"""Definition of get_model_name function."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def get_model_name(
|
||||||
|
path: Path,
|
||||||
|
) -> str:
|
||||||
|
"""Get model name from model_name file."""
|
||||||
|
assert isinstance(path, Path)
|
||||||
|
assert path.exists(), f'{path} does not exist'
|
||||||
|
# read file contents
|
||||||
|
with open(file=path, encoding='utf-8') as fh:
|
||||||
|
model_name = fh.read()
|
||||||
|
# remove newline character
|
||||||
|
model_name = model_name.strip()
|
||||||
|
logging.debug('finished')
|
||||||
|
return model_name
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
"""Definition of load_model function."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from model.src.models import VisualCommunicationModel
|
||||||
|
from shared.repositories import ModelRepository
|
||||||
|
|
||||||
|
from .get_model_name import get_model_name
|
||||||
|
|
||||||
|
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||||
|
|
||||||
|
|
||||||
|
def load_model() -> VisualCommunicationModel:
|
||||||
|
"""Instantiate model with weights loaded from latest model saved in
|
||||||
|
MinIO."""
|
||||||
|
# instantiate model
|
||||||
|
model = VisualCommunicationModel()
|
||||||
|
# get model object name
|
||||||
|
model_name_path = Path('model_name.txt')
|
||||||
|
model_object_name = get_model_name(path=model_name_path)
|
||||||
|
logging.info('using model: %s', model_object_name)
|
||||||
|
# load model data
|
||||||
|
with ModelRepository() as repo:
|
||||||
|
model_data = repo.get_data(model_object_name)
|
||||||
|
if model_data is None:
|
||||||
|
raise FileNotFoundError(f'model {model_object_name} not found')
|
||||||
|
model_checkpoint = torch.load(model_data.buffer)
|
||||||
|
model.load_state_dict(model_checkpoint)
|
||||||
|
# clean memory
|
||||||
|
model_checkpoint.clear()
|
||||||
|
# move model to selected device
|
||||||
|
model = model.to(DEVICE)
|
||||||
|
logging.debug('finished')
|
||||||
|
return model
|
||||||
@@ -6,7 +6,7 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from dataloader import VCDADataset # noqa: F401
|
from dataloader import VCDADataset # noqa: F401
|
||||||
|
|
||||||
from .models import VisualCommunicationModel
|
from ..models import VisualCommunicationModel
|
||||||
|
|
||||||
PRE_WARMUP_LR = 1e-10
|
PRE_WARMUP_LR = 1e-10
|
||||||
POST_WARMUP_LR = 1e-5
|
POST_WARMUP_LR = 1e-5
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
"""Definition of Visual Critial Discourse Analysis dataloader."""
|
||||||
|
|
||||||
|
import random
|
||||||
|
|
||||||
|
from PIL import Image
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
from torchvision.transforms import ColorJitter, InterpolationMode, Normalize
|
||||||
|
from torchvision.transforms.functional import (
|
||||||
|
hflip,
|
||||||
|
pad,
|
||||||
|
resize,
|
||||||
|
rotate,
|
||||||
|
to_pil_image,
|
||||||
|
to_tensor,
|
||||||
|
)
|
||||||
|
|
||||||
|
from shared.repositories import ImageRepository
|
||||||
|
|
||||||
|
# resnet18 original normalization values
|
||||||
|
RESNET_NORMALIZE_MEAN = [0.485, 0.456, 0.406]
|
||||||
|
RESNET_NORMALIZE_STD = [0.229, 0.224, 0.225]
|
||||||
|
|
||||||
|
|
||||||
|
class VCDADataset(Dataset):
|
||||||
|
"""VCDA dataset class."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
data_name_list: list[str],
|
||||||
|
do_augment: bool = False,
|
||||||
|
random_annotations: bool = False,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.data_name_list = data_name_list
|
||||||
|
self.do_augment = do_augment
|
||||||
|
self.random_annotations = random_annotations
|
||||||
|
self.normalize_mean = RESNET_NORMALIZE_MEAN
|
||||||
|
self.normalize_std = RESNET_NORMALIZE_STD
|
||||||
|
# prepare augmentation functions
|
||||||
|
self.normalize = Normalize(
|
||||||
|
mean=self.normalize_mean,
|
||||||
|
std=self.normalize_std,
|
||||||
|
)
|
||||||
|
self.color_jitter = ColorJitter(
|
||||||
|
brightness=1e-1,
|
||||||
|
contrast=8e-2,
|
||||||
|
saturation=8e-2,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.data_name_list)
|
||||||
|
|
||||||
|
def __getitem__(self, idx):
|
||||||
|
# get image from database
|
||||||
|
image_name = self.data_name_list[idx]
|
||||||
|
with ImageRepository() as repo:
|
||||||
|
image_data = repo.get_data(image_name)
|
||||||
|
assert image_data is not None
|
||||||
|
tensor = self.image_to_tensor(image_data.image)
|
||||||
|
if self.do_augment:
|
||||||
|
tensor = self.augment(tensor)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
def image_to_tensor(
|
||||||
|
self,
|
||||||
|
image: Image.Image,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Load images tensor from bytes."""
|
||||||
|
tensor = to_tensor(image)
|
||||||
|
tensor /= 255 # normalize 8-bit image
|
||||||
|
tensor = self.square_pad(tensor=tensor)
|
||||||
|
tensor = resize(
|
||||||
|
img=tensor,
|
||||||
|
size=(512, 512),
|
||||||
|
interpolation=InterpolationMode.BICUBIC,
|
||||||
|
)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def square_pad(
|
||||||
|
tensor: Tensor,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Pads image to a square with side length equal to the largest side of
|
||||||
|
the input image."""
|
||||||
|
assert isinstance(tensor, Tensor)
|
||||||
|
# B, nc, w, h = img.shape
|
||||||
|
h = tensor.shape[-2]
|
||||||
|
w = tensor.shape[-1]
|
||||||
|
if h == w:
|
||||||
|
return tensor
|
||||||
|
max_wh = max([h, w])
|
||||||
|
hp = int((max_wh - w) / 2)
|
||||||
|
vp = int((max_wh - h) / 2)
|
||||||
|
padding = (hp, vp, hp, vp)
|
||||||
|
tensor = pad(tensor, padding, 0, 'constant')
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
def augment(
|
||||||
|
self,
|
||||||
|
tensor: Tensor,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Augment image with random horizontal flips, rotations and color
|
||||||
|
jitter."""
|
||||||
|
assert isinstance(tensor, Tensor)
|
||||||
|
# left-right flip
|
||||||
|
if random.random() >= 0.5:
|
||||||
|
tensor = hflip(tensor)
|
||||||
|
# rotation
|
||||||
|
rnd = random.random()
|
||||||
|
if rnd < 0.25:
|
||||||
|
tensor = rotate(tensor, angle=90)
|
||||||
|
if rnd < 0.5:
|
||||||
|
tensor = rotate(tensor, angle=180)
|
||||||
|
if rnd < 0.75:
|
||||||
|
tensor = rotate(tensor, angle=270)
|
||||||
|
# color jitter
|
||||||
|
tensor = self.color_jitter(tensor)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
def reverse_normalise(
|
||||||
|
self,
|
||||||
|
tensor: Tensor,
|
||||||
|
) -> Image.Image:
|
||||||
|
"""Reverse normalization to get an image that can be interpreted by
|
||||||
|
humans."""
|
||||||
|
assert isinstance(tensor, Tensor)
|
||||||
|
tensor *= Tensor(self.normalize_std).reshape((3, 1, 1))
|
||||||
|
tensor += Tensor(self.normalize_mean).reshape((3, 1, 1))
|
||||||
|
image = to_pil_image(
|
||||||
|
pic=tensor,
|
||||||
|
mode='RGB',
|
||||||
|
)
|
||||||
|
return image
|
||||||
+163
@@ -0,0 +1,163 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
from argparse import ArgumentParser
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from ignite.engine import Events, create_supervised_evaluator, create_supervised_trainer
|
||||||
|
from ignite.handlers import global_step_from_engine
|
||||||
|
from ignite.handlers.checkpoint import Checkpoint
|
||||||
|
from ignite.handlers.param_scheduler import create_lr_scheduler_with_warmup
|
||||||
|
from ignite.handlers.tensorboard_logger import TensorboardLogger
|
||||||
|
from ignite.handlers.tqdm_logger import ProgressBar
|
||||||
|
from ignite.metrics import Average, Loss, RunningAverage
|
||||||
|
from torch import nn
|
||||||
|
from torch.optim.lr_scheduler import ExponentialLR
|
||||||
|
from torch.utils.data import DataLoader
|
||||||
|
|
||||||
|
from model.src.utils import VCDADataset, get_class
|
||||||
|
|
||||||
|
|
||||||
|
def parse_arguments():
|
||||||
|
parser = ArgumentParser()
|
||||||
|
parser.add_argument('-c', '--config', type=Path, required=True)
|
||||||
|
parser.add_argument('-d', '--device', type=str, default='cpu')
|
||||||
|
parser.add_argument('-b', '--batch-size', type=int, default=32)
|
||||||
|
parser.add_argument('--epochs', type=int, default=50)
|
||||||
|
parser.add_argument('--epoch-length', type=int, default=128)
|
||||||
|
parser.add_argument('--checkpoint', type=Path)
|
||||||
|
parser.add_argument('--lr', type=float, default=1e-5)
|
||||||
|
parser.add_argument('--lr-gamma', type=float, default=0.95)
|
||||||
|
parser.add_argument('--lr-warmup-start', type=float, default=1e-8)
|
||||||
|
parser.add_argument('--lr-warmup-duration', type=int, default=5)
|
||||||
|
parser.add_argument('--train-seed', type=int, default=1312)
|
||||||
|
parser.add_argument('--val-seed', type=int, default=1313)
|
||||||
|
parser.add_argument('--loader-workers', type=int, default=8)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
args = parse_arguments()
|
||||||
|
|
||||||
|
load_dotenv('server.env')
|
||||||
|
|
||||||
|
# read model config
|
||||||
|
with open(args.config) as fh:
|
||||||
|
config = json.load(fh)
|
||||||
|
model_name = args.config.stem
|
||||||
|
run_name = datetime.now(timezone.utc).strftime('%Y-%m-%d-%H%M') + '_' + model_name
|
||||||
|
device = torch.device(args.device)
|
||||||
|
|
||||||
|
# create model
|
||||||
|
model_class = get_class(config['backbone']['class'])
|
||||||
|
model = model_class().to(device)
|
||||||
|
optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
|
||||||
|
loss_fn = nn.CrossEntropyLoss()
|
||||||
|
|
||||||
|
# create datasets and loaders
|
||||||
|
with open('model/src/dataset/train.csv', encoding='utf-8') as fh:
|
||||||
|
train_data_name_list = fh.read().split('\n')
|
||||||
|
train_dataset = VCDADataset(
|
||||||
|
data_name_list=train_data_name_list,
|
||||||
|
)
|
||||||
|
train_loader = DataLoader(dataset=train_dataset, num_workers=args.loader_workers)
|
||||||
|
|
||||||
|
with open('model/src/dataset/val.csv', encoding='utf-8') as fh:
|
||||||
|
val_data_name_list = fh.read().split('\n')
|
||||||
|
val_dataset = VCDADataset(data_name_list=val_data_name_list)
|
||||||
|
val_loader = DataLoader(dataset=val_dataset, num_workers=args.loader_workers)
|
||||||
|
|
||||||
|
# create trainer and evaluator
|
||||||
|
trainer = create_supervised_trainer(
|
||||||
|
model=model,
|
||||||
|
optimizer=optimizer,
|
||||||
|
loss_fn=loss_fn,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
Average(output_transform=lambda x: x).attach(trainer, name='loss')
|
||||||
|
RunningAverage(output_transform=lambda x: x, alpha=0.5).attach(
|
||||||
|
trainer,
|
||||||
|
'running_avg_loss',
|
||||||
|
)
|
||||||
|
ProgressBar(desc='Train', ncols=80).attach(trainer, ['running_avg_loss'])
|
||||||
|
|
||||||
|
val_metrics = {
|
||||||
|
'loss': Loss(loss_fn=loss_fn, device=device),
|
||||||
|
}
|
||||||
|
evaluator = create_supervised_evaluator(
|
||||||
|
model=model,
|
||||||
|
metrics=val_metrics, # type: ignore
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
ProgressBar(desc='Val', ncols=80).attach(evaluator)
|
||||||
|
|
||||||
|
# create lr scheduler
|
||||||
|
lr_scheduler = ExponentialLR(
|
||||||
|
optimizer=optimizer,
|
||||||
|
gamma=args.lr_gamma,
|
||||||
|
)
|
||||||
|
lr_handler = create_lr_scheduler_with_warmup(
|
||||||
|
lr_scheduler=lr_scheduler,
|
||||||
|
warmup_start_value=args.lr_warmup_start,
|
||||||
|
warmup_duration=args.lr_warmup_duration,
|
||||||
|
warmup_end_value=args.lr,
|
||||||
|
)
|
||||||
|
trainer.add_event_handler(Events.EPOCH_STARTED, lr_handler)
|
||||||
|
|
||||||
|
|
||||||
|
# log metrics to TensorBoard
|
||||||
|
@trainer.on(Events.EPOCH_COMPLETED)
|
||||||
|
def evaluate():
|
||||||
|
evaluator.run(val_loader, epoch_length=args.val_epoch_length)
|
||||||
|
|
||||||
|
|
||||||
|
tb_logger = TensorboardLogger(log_dir=f'runs/logs/{run_name}')
|
||||||
|
tb_logger.attach_opt_params_handler(
|
||||||
|
engine=trainer,
|
||||||
|
event_name=Events.EPOCH_COMPLETED,
|
||||||
|
optimizer=optimizer,
|
||||||
|
)
|
||||||
|
for tag, engine in [('train', trainer), ('val', evaluator)]:
|
||||||
|
tb_logger.attach_output_handler(
|
||||||
|
engine,
|
||||||
|
event_name=Events.EPOCH_COMPLETED,
|
||||||
|
tag=tag,
|
||||||
|
metric_names='all',
|
||||||
|
global_step_transform=global_step_from_engine(trainer),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Set up checkpoint saving
|
||||||
|
to_save = {
|
||||||
|
'model': model,
|
||||||
|
'optimizer': optimizer,
|
||||||
|
'trainer': trainer,
|
||||||
|
'lr_scheduler': lr_scheduler,
|
||||||
|
'lr_handler': lr_handler,
|
||||||
|
}
|
||||||
|
checkpoint_handler = Checkpoint(
|
||||||
|
to_save,
|
||||||
|
f'runs/checkpoints/{run_name}',
|
||||||
|
n_saved=3,
|
||||||
|
filename_prefix='best',
|
||||||
|
score_function=lambda engine: -engine.state.metrics['loss'],
|
||||||
|
score_name='neg_val_loss',
|
||||||
|
global_step_transform=global_step_from_engine(trainer),
|
||||||
|
)
|
||||||
|
evaluator.add_event_handler(Events.COMPLETED, checkpoint_handler)
|
||||||
|
|
||||||
|
if args.checkpoint:
|
||||||
|
Checkpoint.load_objects(to_load=to_save, checkpoint=str(args.checkpoint))
|
||||||
|
|
||||||
|
# save model config
|
||||||
|
os.makedirs('runs/configs/', exist_ok=True)
|
||||||
|
with open(f'runs/configs/{run_name}.json', 'w', encoding='utf-8') as fh:
|
||||||
|
json.dump(config, fh)
|
||||||
|
|
||||||
|
# start training
|
||||||
|
trainer.run(
|
||||||
|
train_loader,
|
||||||
|
max_epochs=args.epochs,
|
||||||
|
epoch_length=args.epoch_length,
|
||||||
|
)
|
||||||
|
tb_logger.close()
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Script to move all images from MongoDB to MinIO."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from bson import ObjectId
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from pymongo.collection import Collection
|
||||||
|
|
||||||
|
from shared.datastore import Datastore
|
||||||
|
from shared.docstore import connect_mongodb
|
||||||
|
from shared.docstore.src.classes import VisualCommunication
|
||||||
|
from shared.utils import check_env, setup_logging
|
||||||
|
from web_ui.src.main import NECESSARY_ENV_VAR_LIST
|
||||||
|
|
||||||
|
|
||||||
|
def list_mongo_document_ids(
|
||||||
|
collection: Collection,
|
||||||
|
) -> list[ObjectId]:
|
||||||
|
"""Get list of all VisualCommunication documents in MongoDB."""
|
||||||
|
# prepare query
|
||||||
|
query = collection.find(
|
||||||
|
filter={},
|
||||||
|
projection={'_id': True},
|
||||||
|
)
|
||||||
|
# execute query
|
||||||
|
res_list = list(query)
|
||||||
|
logging.debug('got %s document(s)', len(res_list))
|
||||||
|
# convert result
|
||||||
|
id_list = [elem['_id'] for elem in res_list]
|
||||||
|
return id_list
|
||||||
|
|
||||||
|
|
||||||
|
def get_visual_communication(
|
||||||
|
collection: Collection,
|
||||||
|
doc_id: ObjectId,
|
||||||
|
) -> VisualCommunication:
|
||||||
|
"""Get image from MongoDB."""
|
||||||
|
# prepare query
|
||||||
|
query = collection.find(
|
||||||
|
filter={'_id': doc_id},
|
||||||
|
projection={'_id': False},
|
||||||
|
)
|
||||||
|
# execute query
|
||||||
|
res_list = list(query)
|
||||||
|
logging.debug('got %s document(s)', len(res_list))
|
||||||
|
# convert result
|
||||||
|
vis_com = VisualCommunication.model_validate(res_list[0])
|
||||||
|
return vis_com
|
||||||
|
|
||||||
|
|
||||||
|
def update_visual_communication(
|
||||||
|
collection: Collection,
|
||||||
|
vis_com: VisualCommunication,
|
||||||
|
) -> None:
|
||||||
|
"""Update VisualCommunication document in MongoDB."""
|
||||||
|
try:
|
||||||
|
collection.update_one(
|
||||||
|
filter={
|
||||||
|
'name': vis_com.name,
|
||||||
|
},
|
||||||
|
update={
|
||||||
|
'$set': vis_com.model_dump(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logging.error('failed saving %s', vis_com.name)
|
||||||
|
logging.debug(exc)
|
||||||
|
raise exc
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
# load in env file
|
||||||
|
env_path = Path(__file__).parent.parent / 'server.env'
|
||||||
|
assert env_path.exists()
|
||||||
|
load_dotenv(env_path)
|
||||||
|
# ensure env vars set
|
||||||
|
check_env(NECESSARY_ENV_VAR_LIST)
|
||||||
|
# setup logging
|
||||||
|
setup_logging()
|
||||||
|
# connect to minIO
|
||||||
|
datastore = Datastore()
|
||||||
|
datastore.connect()
|
||||||
|
assert datastore._client is not None
|
||||||
|
# connect to MongoDB
|
||||||
|
collection, db, client = connect_mongodb()
|
||||||
|
# list documents in mongoDB
|
||||||
|
id_list = list_mongo_document_ids(collection)
|
||||||
|
for doc_id in id_list:
|
||||||
|
try:
|
||||||
|
# get image from MongoDB
|
||||||
|
vis_com = get_visual_communication(collection, doc_id)
|
||||||
|
except Exception as exc:
|
||||||
|
logging.debug(exc)
|
||||||
|
logging.error('failed getting image from document: %s', doc_id)
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
# get image
|
||||||
|
image = vis_com.get_image(
|
||||||
|
minio_client=datastore._client,
|
||||||
|
)
|
||||||
|
# put buffer in minio
|
||||||
|
object_name = datastore.put_image(
|
||||||
|
image=image,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logging.debug(exc)
|
||||||
|
logging.error('failed saving image to minio')
|
||||||
|
continue
|
||||||
|
# update visual communication in mongodb
|
||||||
|
try:
|
||||||
|
vis_com.image = None # type: ignore
|
||||||
|
vis_com.object_name = object_name
|
||||||
|
update_visual_communication(collection, vis_com)
|
||||||
|
except Exception as exc:
|
||||||
|
logging.debug(exc)
|
||||||
|
logging.error(
|
||||||
|
'failed updating visual communication %s',
|
||||||
|
vis_com.name,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
logging.debug('updated visual communication %s', vis_com.name)
|
||||||
Generated
+2541
-771
File diff suppressed because it is too large
Load Diff
+77
-1
@@ -16,6 +16,15 @@ python = "^3.12"
|
|||||||
flake8 = "^7.0.0"
|
flake8 = "^7.0.0"
|
||||||
mypy = "^1.8.0"
|
mypy = "^1.8.0"
|
||||||
types-pillow = "^10.2.0.20240213"
|
types-pillow = "^10.2.0.20240213"
|
||||||
|
types-requests = "^2.32.0.20240602"
|
||||||
|
types-retry = "^0.9.9.4"
|
||||||
|
flake8-pyproject = "^1.2.3"
|
||||||
|
pandas-stubs = "^2.2.2.240603"
|
||||||
|
types-tqdm = "^4.66.0.20240417"
|
||||||
|
pytest = "^8.3.3"
|
||||||
|
testcontainers = "^4.8.2"
|
||||||
|
coverage = "^7.6.4"
|
||||||
|
pytest-cov = "^6.0.0"
|
||||||
|
|
||||||
|
|
||||||
[tool.poetry.group.dev.dependencies]
|
[tool.poetry.group.dev.dependencies]
|
||||||
@@ -23,11 +32,17 @@ pandas = "^2.2.1"
|
|||||||
selenium = "^4.18.1"
|
selenium = "^4.18.1"
|
||||||
webdriver-manager = "^4.0.1"
|
webdriver-manager = "^4.0.1"
|
||||||
retry = "^0.9.2"
|
retry = "^0.9.2"
|
||||||
|
pre-commit = "^4.1.0"
|
||||||
|
|
||||||
|
|
||||||
[tool.poetry.group.model.dependencies]
|
[tool.poetry.group.model.dependencies]
|
||||||
torch = "^2.2.1"
|
torch = "^2.0.0"
|
||||||
torchvision = "^0.17.1"
|
torchvision = "^0.17.1"
|
||||||
|
torchinfo = "^1.8.0"
|
||||||
|
minio = "^7.2.7"
|
||||||
|
tqdm = "^4.66.4"
|
||||||
|
pytorch-ignite = "^0.5.1"
|
||||||
|
tensorboard = "^2.17.1"
|
||||||
|
|
||||||
|
|
||||||
[tool.poetry.group.shared.dependencies]
|
[tool.poetry.group.shared.dependencies]
|
||||||
@@ -43,6 +58,67 @@ dash-bootstrap-components = "^1.5.0"
|
|||||||
dash-mantine-components = "^0.12.1"
|
dash-mantine-components = "^0.12.1"
|
||||||
dash-auth = "^2.2.0"
|
dash-auth = "^2.2.0"
|
||||||
|
|
||||||
|
|
||||||
|
[[tool.poetry.source]]
|
||||||
|
name = "threadripper"
|
||||||
|
url = "http://192.168.1.2:5001/index/"
|
||||||
|
priority = "primary"
|
||||||
|
|
||||||
[build-system]
|
[build-system]
|
||||||
requires = ["poetry-core"]
|
requires = ["poetry-core"]
|
||||||
build-backend = "poetry.core.masonry.api"
|
build-backend = "poetry.core.masonry.api"
|
||||||
|
|
||||||
|
[tool.isort]
|
||||||
|
profile = "black"
|
||||||
|
|
||||||
|
[tool.flake8]
|
||||||
|
per-file-ignores = "__init__.py:F401"
|
||||||
|
max-line-length = 88
|
||||||
|
extend-ignore = "E203"
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
exclude = "image_download"
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "dash.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "dash_auth.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "dash_mantine_components.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "dash_bootstrap_components.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "torchvision.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "image_download.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "dataloader.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "utils.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "shared.datastore.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "models.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "shared.docstore.*"
|
||||||
|
ignore_missing_imports = true
|
||||||
|
|||||||
@@ -1,13 +0,0 @@
|
|||||||
"""Database module content."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .classes.dataset import Dataset
|
|
||||||
from .classes.exceptions import NoDocumentFoundException
|
|
||||||
from .classes.visual_communication import VisualCommunication
|
|
||||||
from .utils.connect import connect
|
|
||||||
from .utils.count_documents import count_documents
|
|
||||||
from .utils.get_visual_communication import get_visual_communication
|
|
||||||
from .utils.list_names import list_names
|
|
||||||
from .utils.upsert_annotation import upsert_annotation
|
|
||||||
from .utils.upsert_prediction import upsert_prediction
|
|
||||||
from .utils.upsert_visual_communication import upsert_visual_communication
|
|
||||||
@@ -1,6 +0,0 @@
|
|||||||
"""Database classes module content."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .dataset import Dataset
|
|
||||||
from .exceptions import NoDocumentFoundException
|
|
||||||
from .visual_communication import VisualCommunication
|
|
||||||
@@ -1,71 +0,0 @@
|
|||||||
"""Definition of database Dataset class."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
import random
|
|
||||||
from datetime import datetime
|
|
||||||
from datetime import UTC
|
|
||||||
|
|
||||||
from pydantic import BaseModel
|
|
||||||
from pydantic import Field
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
|
|
||||||
class Dataset(BaseModel):
|
|
||||||
"""Database Dataset model."""
|
|
||||||
create_time: datetime = Field(default_factory=lambda: datetime.now(UTC))
|
|
||||||
train_names: list[str]
|
|
||||||
test_names: list[str]
|
|
||||||
validation_names: list[str]
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def fraction_map(cls) -> dict[str, float]:
|
|
||||||
"""Dict with train, test and validation fractions."""
|
|
||||||
# define map
|
|
||||||
split_map = {
|
|
||||||
'train': 0.7,
|
|
||||||
'test': 0.2,
|
|
||||||
'validation': 0.1,
|
|
||||||
}
|
|
||||||
# sanity check
|
|
||||||
assert sum(split_map.values()) == 1.0
|
|
||||||
return split_map
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def new_from_name_list(
|
|
||||||
cls,
|
|
||||||
name_list: list[str],
|
|
||||||
) -> Dataset:
|
|
||||||
"""Generate new dataset from list of filenames."""
|
|
||||||
# calculate split fractions
|
|
||||||
fraction_map = cls.fraction_map()
|
|
||||||
num_total = len(name_list)
|
|
||||||
num_validation = round(num_total * fraction_map['validation'])
|
|
||||||
num_test = round(num_total * fraction_map['test'])
|
|
||||||
# split data
|
|
||||||
validation_name_list = random.choices(name_list, k=num_validation)
|
|
||||||
name_list = [
|
|
||||||
name for name in name_list if name not in validation_name_list
|
|
||||||
]
|
|
||||||
test_name_list = random.choices(name_list, k=num_test)
|
|
||||||
train_name_list = [
|
|
||||||
name for name in name_list if name not in test_name_list
|
|
||||||
]
|
|
||||||
# instantiate object
|
|
||||||
dataset = Dataset(
|
|
||||||
train_names=train_name_list,
|
|
||||||
test_names=test_name_list,
|
|
||||||
validation_names=validation_name_list,
|
|
||||||
)
|
|
||||||
logging.debug('finished')
|
|
||||||
return dataset
|
|
||||||
|
|
||||||
def save(
|
|
||||||
self,
|
|
||||||
collection: Collection,
|
|
||||||
) -> None:
|
|
||||||
"""Save dataset to database."""
|
|
||||||
res = collection.insert_one(
|
|
||||||
document=self.model_dump(),
|
|
||||||
)
|
|
||||||
logging.debug('inserted document: %s', res)
|
|
||||||
@@ -1,6 +0,0 @@
|
|||||||
"""Definition of database exception."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
|
|
||||||
class NoDocumentFoundException(Exception):
|
|
||||||
"""Database exception for when no documents are found."""
|
|
||||||
@@ -1,85 +0,0 @@
|
|||||||
"""Definition of VisualCommunication model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
from base64 import b64decode
|
|
||||||
from base64 import b64encode
|
|
||||||
from io import BytesIO
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from PIL import Image
|
|
||||||
from pydantic import BaseModel
|
|
||||||
from pydantic import field_serializer
|
|
||||||
from pydantic import field_validator
|
|
||||||
|
|
||||||
from shared.dto import ModelData
|
|
||||||
|
|
||||||
|
|
||||||
class VisualCommunication(BaseModel):
|
|
||||||
"""Visual communication model."""
|
|
||||||
|
|
||||||
name: str
|
|
||||||
image: Image.Image
|
|
||||||
annotation: ModelData | None = None
|
|
||||||
prediction: ModelData | None = None
|
|
||||||
|
|
||||||
class Config:
|
|
||||||
"""BaseModel configuration."""
|
|
||||||
arbitrary_types_allowed = True
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def classname(cls) -> str:
|
|
||||||
"""Return classname."""
|
|
||||||
return cls.__name__
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_file(cls, path: Path) -> VisualCommunication:
|
|
||||||
"""Instantiate from file."""
|
|
||||||
name = path.stem
|
|
||||||
image = Image.open(path)
|
|
||||||
image.load()
|
|
||||||
return VisualCommunication(name=name, image=image)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def decode_image(cls, content: str) -> Image.Image:
|
|
||||||
"""Extract image from webencoded content."""
|
|
||||||
_, content_data = content.split(',')
|
|
||||||
return Image.open(BytesIO(b64decode(content_data)))
|
|
||||||
|
|
||||||
@field_serializer('image')
|
|
||||||
@classmethod
|
|
||||||
def serialize_image(cls, image: Image.Image) -> bytes: # type: ignore
|
|
||||||
"""Convert image to bytes for storage in database."""
|
|
||||||
buffer = BytesIO()
|
|
||||||
image.save(buffer, format='JPEG')
|
|
||||||
return buffer.getvalue()
|
|
||||||
|
|
||||||
@field_validator('image', mode='before')
|
|
||||||
@classmethod
|
|
||||||
def convert_to_image(
|
|
||||||
cls,
|
|
||||||
image: Image.Image | BytesIO | bytes,
|
|
||||||
) -> Image.Image:
|
|
||||||
"""Convert bytes input from database into image."""
|
|
||||||
if isinstance(image, bytes):
|
|
||||||
image = BytesIO(image)
|
|
||||||
if isinstance(image, BytesIO):
|
|
||||||
image = Image.open(image)
|
|
||||||
return image
|
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
|
||||||
return f"{self.classname()}(name='{self.name}')"
|
|
||||||
|
|
||||||
def webencoded_image(self) -> str:
|
|
||||||
"""Convert image to be displayed on webpage."""
|
|
||||||
# convert images to bytes string
|
|
||||||
buffer = BytesIO()
|
|
||||||
self.image.save(buffer, format='png')
|
|
||||||
img_enc = b64encode(buffer.getvalue()).decode('utf-8')
|
|
||||||
return f"data:image/png;base64, {img_enc}"
|
|
||||||
|
|
||||||
def generate_random_prediction(self, force: bool = False) -> None:
|
|
||||||
"""Generate random prediction values."""
|
|
||||||
if not force and self.prediction is not None:
|
|
||||||
logging.warning('set force=True to overwrite existing values.')
|
|
||||||
self.prediction = ModelData.from_random()
|
|
||||||
@@ -1,4 +0,0 @@
|
|||||||
"""Database utils module content."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .connect import connect
|
|
||||||
@@ -1,33 +0,0 @@
|
|||||||
"""
|
|
||||||
Definition of function to connect to database
|
|
||||||
using environment variables.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
import os
|
|
||||||
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
from pymongo import MongoClient
|
|
||||||
|
|
||||||
|
|
||||||
def connect():
|
|
||||||
"""Connect to MongoDB using env vars."""
|
|
||||||
# load env vars
|
|
||||||
load_dotenv()
|
|
||||||
necessary_env_vars = [
|
|
||||||
'MONGO_HOST',
|
|
||||||
'MONGO_DB',
|
|
||||||
'MONGO_COLLECTION',
|
|
||||||
]
|
|
||||||
for env_var in necessary_env_vars:
|
|
||||||
assert env_var in os.environ, f"{env_var} not found"
|
|
||||||
# connect to database
|
|
||||||
client = MongoClient(os.getenv('MONGO_HOST'))
|
|
||||||
db = client[os.getenv('MONGO_DB')]
|
|
||||||
# extract collection
|
|
||||||
collection = db[os.getenv('MONGO_COLLECTION')]
|
|
||||||
# set unique index on "name"
|
|
||||||
collection.create_index('name', unique=True)
|
|
||||||
logging.debug('finished')
|
|
||||||
return collection, db, client
|
|
||||||
@@ -1,21 +0,0 @@
|
|||||||
"""Definition of function to count documents in database."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
|
|
||||||
def count_documents(
|
|
||||||
collection: Collection,
|
|
||||||
only_with_annotation: bool = False,
|
|
||||||
) -> int:
|
|
||||||
"""
|
|
||||||
Get the total number of documents
|
|
||||||
in database that matches the filters.
|
|
||||||
"""
|
|
||||||
assert isinstance(collection, Collection)
|
|
||||||
assert isinstance(only_with_annotation, bool)
|
|
||||||
# build query
|
|
||||||
query = {}
|
|
||||||
if only_with_annotation:
|
|
||||||
query['annotation'] = {'$ne': None}
|
|
||||||
return collection.count_documents(filter=query)
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from typing import Literal
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
|
|
||||||
def get_dataset(
|
|
||||||
collection: Collection,
|
|
||||||
type: Literal['train', 'test', 'validation'],
|
|
||||||
) -> list[str]:
|
|
||||||
"""Get list of data names for the corresponding type."""
|
|
||||||
return []
|
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
"""Definition of function to get visual communication from database."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
from shared.database import NoDocumentFoundException
|
|
||||||
from shared.database import VisualCommunication
|
|
||||||
|
|
||||||
|
|
||||||
def get_visual_communication(
|
|
||||||
collection: Collection,
|
|
||||||
with_annotation: bool = False,
|
|
||||||
) -> VisualCommunication:
|
|
||||||
"""Get a random visual communication from the database."""
|
|
||||||
query = {}
|
|
||||||
if with_annotation:
|
|
||||||
query['annotation'] = {'$ne': None}
|
|
||||||
else:
|
|
||||||
query['annotation'] = {'$eq': None}
|
|
||||||
data = collection.aggregate(
|
|
||||||
pipeline=[
|
|
||||||
{
|
|
||||||
'$match': query, # find using filters
|
|
||||||
},
|
|
||||||
{
|
|
||||||
'$sample': {
|
|
||||||
'size': 1, # get one random
|
|
||||||
},
|
|
||||||
},
|
|
||||||
],
|
|
||||||
)
|
|
||||||
data_list = list(data) # read data from cursor object
|
|
||||||
if len(data_list) == 0:
|
|
||||||
raise NoDocumentFoundException()
|
|
||||||
vis_com = VisualCommunication.model_validate(data_list[0])
|
|
||||||
logging.debug('finished')
|
|
||||||
return vis_com
|
|
||||||
@@ -1,31 +0,0 @@
|
|||||||
"""
|
|
||||||
Definition of function to list names
|
|
||||||
of all visual communication documents in database.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
|
|
||||||
def list_names(
|
|
||||||
collection: Collection,
|
|
||||||
only_with_annotation: bool = True,
|
|
||||||
) -> list[str]:
|
|
||||||
"""List the names of entries that match the filters."""
|
|
||||||
assert isinstance(collection, Collection)
|
|
||||||
assert isinstance(only_with_annotation, bool)
|
|
||||||
# build query
|
|
||||||
query = {}
|
|
||||||
if only_with_annotation:
|
|
||||||
query['annotation'] = {'$ne': None}
|
|
||||||
# execute query
|
|
||||||
res_list = collection.find(
|
|
||||||
filter=query,
|
|
||||||
projection={
|
|
||||||
'_id': False,
|
|
||||||
'name': True,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
# extract information
|
|
||||||
name_list = [elem['name'] for elem in res_list]
|
|
||||||
return name_list
|
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
from shared.database import Dataset
|
|
||||||
|
|
||||||
|
|
||||||
def save_dataset(
|
|
||||||
collection: Collection,
|
|
||||||
dataset: Dataset,
|
|
||||||
) -> None:
|
|
||||||
"""Save dataset to database."""
|
|
||||||
res = collection.insert_one(
|
|
||||||
document=dataset.model_dump(),
|
|
||||||
)
|
|
||||||
logging.debug('inserted document: %s', res)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
load_dotenv('local.env')
|
|
||||||
from shared.database import list_names, connect
|
|
||||||
# connect to database
|
|
||||||
collection, db, client = connect()
|
|
||||||
print(client.server_info())
|
|
||||||
|
|
||||||
name_list = list_names(collection=collection, only_with_annotation=True)
|
|
||||||
ds = Dataset.new_from_name_list(name_list=name_list)
|
|
||||||
|
|
||||||
print(ds)
|
|
||||||
# save_dataset(
|
|
||||||
# collection=collection,
|
|
||||||
# dataset=ds
|
|
||||||
# )
|
|
||||||
@@ -1,30 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
from shared.dto import ModelData
|
|
||||||
|
|
||||||
|
|
||||||
def upsert_annotation(
|
|
||||||
collection: Collection,
|
|
||||||
vis_com_name: str,
|
|
||||||
annotations: ModelData,
|
|
||||||
) -> None:
|
|
||||||
"""Upserts annotation data in the database."""
|
|
||||||
query = {
|
|
||||||
'name': vis_com_name,
|
|
||||||
}
|
|
||||||
update = {
|
|
||||||
'$set': {
|
|
||||||
'annotation': annotations.model_dump(),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
res = collection.update_one(
|
|
||||||
filter=query,
|
|
||||||
update=update,
|
|
||||||
upsert=True,
|
|
||||||
)
|
|
||||||
logging.info('upserted document: %s', res)
|
|
||||||
logging.info('finished')
|
|
||||||
@@ -1,30 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
from shared.dto import ModelData
|
|
||||||
|
|
||||||
|
|
||||||
def upsert_prediction(
|
|
||||||
collection: Collection,
|
|
||||||
vis_com_name: str,
|
|
||||||
predictions: ModelData,
|
|
||||||
) -> None:
|
|
||||||
"""Upsert prediction data in the database."""
|
|
||||||
query = {
|
|
||||||
'name': vis_com_name,
|
|
||||||
}
|
|
||||||
update = {
|
|
||||||
'$set': {
|
|
||||||
'prediction': predictions.model_dump(),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
res = collection.update_one(
|
|
||||||
filter=query,
|
|
||||||
update=update,
|
|
||||||
upsert=True,
|
|
||||||
)
|
|
||||||
logging.debug('upserted document: %s', res)
|
|
||||||
logging.info('finished')
|
|
||||||
@@ -1,23 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from pymongo.collection import Collection
|
|
||||||
|
|
||||||
from shared.database import VisualCommunication
|
|
||||||
|
|
||||||
|
|
||||||
def upsert_visual_communication(
|
|
||||||
collection: Collection,
|
|
||||||
visual_communication_list: list[VisualCommunication],
|
|
||||||
) -> bool:
|
|
||||||
"""
|
|
||||||
Upsert VisualCommunication object in the database.
|
|
||||||
Returns bool stating success.
|
|
||||||
"""
|
|
||||||
response = collection.insert_many(
|
|
||||||
[
|
|
||||||
vis_com.model_dump()
|
|
||||||
for vis_com
|
|
||||||
in visual_communication_list
|
|
||||||
],
|
|
||||||
)
|
|
||||||
return response.acknowledged
|
|
||||||
@@ -1,15 +0,0 @@
|
|||||||
"""Data transfer objects module content."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .angle import AngleData
|
|
||||||
from .contact import ContactData
|
|
||||||
from .distance import DistanceData
|
|
||||||
from .framing import FramingData
|
|
||||||
from .information_value import InformationValueData
|
|
||||||
from .modality_color import ModalityColorData
|
|
||||||
from .modality_depth import ModalityDepthData
|
|
||||||
from .modality_lighting import ModalityLightingData
|
|
||||||
from .model_data import ModelData
|
|
||||||
from .point_of_view import PointOfViewData
|
|
||||||
from .salience import SalienceData
|
|
||||||
from .visual_syntax import VisualSyntaxData
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
"""Definition of Angle data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class AngleData(DataModel):
|
|
||||||
"""Angle data model."""
|
|
||||||
high: float
|
|
||||||
eye_level: float
|
|
||||||
low: float
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
"""Definition of ContactData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class ContactData(DataModel):
|
|
||||||
"""ContactData data model."""
|
|
||||||
|
|
||||||
offer: float
|
|
||||||
demand: float
|
|
||||||
@@ -1,67 +0,0 @@
|
|||||||
"""Definition of DataModel base class."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import random
|
|
||||||
|
|
||||||
from pydantic import BaseModel
|
|
||||||
from pydantic import ValidationError
|
|
||||||
|
|
||||||
|
|
||||||
class DataModel(BaseModel):
|
|
||||||
"""DataModel base class."""
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def classname(cls) -> str:
|
|
||||||
"""Return classname."""
|
|
||||||
return cls.__name__
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def list_fields(cls) -> list[str]:
|
|
||||||
"""List options that are stored as attributes."""
|
|
||||||
return list(cls.model_fields.keys())
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_random(cls):
|
|
||||||
"""Instantiate with random numbers."""
|
|
||||||
kwargs = {field: random.random() for field in cls.list_fields()}
|
|
||||||
return cls(**kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_choice(cls, option: str):
|
|
||||||
"""Instantiate from choice."""
|
|
||||||
if option is None:
|
|
||||||
raise ValidationError()
|
|
||||||
assert isinstance(option, str), 'option is not a string'
|
|
||||||
allowed_options_list = cls.list_fields()
|
|
||||||
assert option in allowed_options_list, \
|
|
||||||
f"{option} is not among allowed fields {allowed_options_list}"
|
|
||||||
kwargs = {field: 0 for field in cls.list_fields()}
|
|
||||||
kwargs[option] = 1
|
|
||||||
return cls(**kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_list(cls, data_list: list[float]):
|
|
||||||
"""Instantiate from list of values."""
|
|
||||||
kwargs = {key: val for key, val in zip(cls.list_fields(), data_list)}
|
|
||||||
return cls(**kwargs)
|
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
|
||||||
model_dict = self.model_dump()
|
|
||||||
model_repr_str = f"{self.classname()}("
|
|
||||||
model_repr_str += ', '.join([
|
|
||||||
f"{field}={value:.3f}"
|
|
||||||
for field, value
|
|
||||||
in model_dict.items()
|
|
||||||
])
|
|
||||||
model_repr_str += ')'
|
|
||||||
return model_repr_str
|
|
||||||
|
|
||||||
def highest_score_field(self) -> str:
|
|
||||||
"""Return name of field with highest score."""
|
|
||||||
model_dict = self.model_dump()
|
|
||||||
return max(model_dict, key=lambda k: model_dict[k])
|
|
||||||
|
|
||||||
def highest_score_value(self) -> float:
|
|
||||||
"""Return value of field with highest score."""
|
|
||||||
model_dict = self.model_dump()
|
|
||||||
return max(model_dict.values())
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
"""Definition of DistanceData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class DistanceData(DataModel):
|
|
||||||
"""DistanceData data model."""
|
|
||||||
|
|
||||||
long: float
|
|
||||||
medium: float
|
|
||||||
close: float
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
"""Definition of FramingData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class FramingData(DataModel):
|
|
||||||
"""FramingData data model."""
|
|
||||||
|
|
||||||
frame_lines: float
|
|
||||||
empty_space: float
|
|
||||||
colour_contrast: float
|
|
||||||
form_contrast: float
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
"""Definition of InformationValueData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class InformationValueData(DataModel):
|
|
||||||
"""InformationValueData data model."""
|
|
||||||
|
|
||||||
given_new: float
|
|
||||||
ideal_real: float
|
|
||||||
central_marginal: float
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
"""Definition of ModalityColorData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class ModalityColorData(DataModel):
|
|
||||||
"""ModalityColorData data model."""
|
|
||||||
|
|
||||||
high: float
|
|
||||||
medium: float
|
|
||||||
low: float
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
"""Definition of ModalityDepthData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class ModalityDepthData(DataModel):
|
|
||||||
"""ModalityDepthData data model."""
|
|
||||||
|
|
||||||
high: float
|
|
||||||
medium: float
|
|
||||||
low: float
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
"""Definition of ModalityLightingData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class ModalityLightingData(DataModel):
|
|
||||||
"""ModalityLightingData data model."""
|
|
||||||
|
|
||||||
high: float
|
|
||||||
medium: float
|
|
||||||
low: float
|
|
||||||
@@ -1,82 +0,0 @@
|
|||||||
"""Definition of ModelData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .angle import AngleData
|
|
||||||
from .contact import ContactData
|
|
||||||
from .data_model import DataModel
|
|
||||||
from .distance import DistanceData
|
|
||||||
from .framing import FramingData
|
|
||||||
from .information_value import InformationValueData
|
|
||||||
from .modality_color import ModalityColorData
|
|
||||||
from .modality_depth import ModalityDepthData
|
|
||||||
from .modality_lighting import ModalityLightingData
|
|
||||||
from .point_of_view import PointOfViewData
|
|
||||||
from .salience import SalienceData
|
|
||||||
from .visual_syntax import VisualSyntaxData
|
|
||||||
|
|
||||||
|
|
||||||
class ModelData(DataModel):
|
|
||||||
"""ModelData model for data IO with combined ML model."""
|
|
||||||
visual_syntax: VisualSyntaxData
|
|
||||||
contact: ContactData
|
|
||||||
angle: AngleData
|
|
||||||
point_of_view: PointOfViewData
|
|
||||||
distance: DistanceData
|
|
||||||
modality_lighting: ModalityLightingData
|
|
||||||
modality_color: ModalityColorData
|
|
||||||
modality_depth: ModalityDepthData
|
|
||||||
information_value: InformationValueData
|
|
||||||
framing: FramingData
|
|
||||||
salience: SalienceData
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_random(cls) -> ModelData:
|
|
||||||
"""Instantiate with random numbers."""
|
|
||||||
kwargs = {
|
|
||||||
field: field_info.annotation.from_random() # type: ignore
|
|
||||||
for field, field_info
|
|
||||||
in cls.model_fields.items()
|
|
||||||
}
|
|
||||||
return cls(**kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_annotations(
|
|
||||||
cls,
|
|
||||||
visual_syntax: str,
|
|
||||||
contact: str,
|
|
||||||
angle: str,
|
|
||||||
point_of_view: str,
|
|
||||||
distance: str,
|
|
||||||
modality_lighting: str,
|
|
||||||
modality_color: str,
|
|
||||||
modality_depth: str,
|
|
||||||
information_value: str,
|
|
||||||
framing: str,
|
|
||||||
salience: str,
|
|
||||||
) -> ModelData:
|
|
||||||
"""Instantiate from annotation."""
|
|
||||||
kwargs = {
|
|
||||||
'visual_syntax': VisualSyntaxData
|
|
||||||
.from_choice(visual_syntax),
|
|
||||||
'contact': ContactData
|
|
||||||
.from_choice(contact),
|
|
||||||
'angle': AngleData
|
|
||||||
.from_choice(angle),
|
|
||||||
'point_of_view': PointOfViewData
|
|
||||||
.from_choice(point_of_view),
|
|
||||||
'distance': DistanceData
|
|
||||||
.from_choice(distance),
|
|
||||||
'modality_lighting': ModalityLightingData
|
|
||||||
.from_choice(modality_lighting),
|
|
||||||
'modality_color': ModalityColorData
|
|
||||||
.from_choice(modality_color),
|
|
||||||
'modality_depth': ModalityDepthData
|
|
||||||
.from_choice(modality_depth),
|
|
||||||
'information_value': InformationValueData
|
|
||||||
.from_choice(information_value),
|
|
||||||
'framing': FramingData
|
|
||||||
.from_choice(framing),
|
|
||||||
'salience': SalienceData
|
|
||||||
.from_choice(salience),
|
|
||||||
}
|
|
||||||
return cls(**kwargs)
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
"""Definition of PointOfViewData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class PointOfViewData(DataModel):
|
|
||||||
"""PointOfViewData data model."""
|
|
||||||
|
|
||||||
frontal: float
|
|
||||||
oblique: float
|
|
||||||
@@ -1,14 +0,0 @@
|
|||||||
"""Definition of SalienceData data model."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .data_model import DataModel
|
|
||||||
|
|
||||||
|
|
||||||
class SalienceData(DataModel):
|
|
||||||
"""SalienceData data model."""
|
|
||||||
|
|
||||||
size: float
|
|
||||||
colour: float
|
|
||||||
tone: float
|
|
||||||
form: float
|
|
||||||
positioning: float
|
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
from .src import ImageRepository, ModelRepository, VisualCommunicationRepository
|
||||||
|
from .src.dto import (
|
||||||
|
HexadecimalString,
|
||||||
|
ImageData,
|
||||||
|
ModelData,
|
||||||
|
VisualCommunicationData,
|
||||||
|
VisualCommunicationValues,
|
||||||
|
)
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .image_repository import ImageRepository
|
||||||
|
from .model_repository import ModelRepository
|
||||||
|
from .visual_communication_repository import VisualCommunicationRepository
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
from .hexadecimal_string import HexadecimalString
|
||||||
|
from .image_data import ImageData
|
||||||
|
from .model_data import ModelData
|
||||||
|
from .visual_communication_data import (
|
||||||
|
VisualCommunicationData,
|
||||||
|
VisualCommunicationValues,
|
||||||
|
)
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
"""Definition of BytesIO Pydantic Annotation."""
|
||||||
|
|
||||||
|
from io import BytesIO
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from pydantic.json_schema import JsonSchemaValue
|
||||||
|
from pydantic_core import core_schema
|
||||||
|
|
||||||
|
|
||||||
|
class BytesIOPydanticAnnotation:
|
||||||
|
"""Pydantic annotation that defines input validation, as well as general
|
||||||
|
and json serialization."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_input(cls, v: Any, handler) -> BytesIO:
|
||||||
|
"""Pydantic-related function to validate input on instantiation."""
|
||||||
|
if isinstance(v, BytesIO):
|
||||||
|
return v
|
||||||
|
s = handler(v)
|
||||||
|
return BytesIO(s)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def __get_pydantic_core_schema__(
|
||||||
|
cls,
|
||||||
|
source_type,
|
||||||
|
_handler,
|
||||||
|
) -> core_schema.CoreSchema:
|
||||||
|
assert source_type is BytesIO
|
||||||
|
return core_schema.no_info_wrap_validator_function(
|
||||||
|
function=cls.validate_input,
|
||||||
|
schema=core_schema.str_schema(),
|
||||||
|
serialization=core_schema.to_string_ser_schema(),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def __get_pydantic_json_schema__(cls, _core_schema, handler) -> JsonSchemaValue:
|
||||||
|
return handler(core_schema.str_schema())
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
"""Definition of Checksum DTO."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
|
||||||
|
|
||||||
|
class HexadecimalString(str):
|
||||||
|
"""Hexadecimal-string class."""
|
||||||
|
|
||||||
|
def __new__(cls, string):
|
||||||
|
# ensure proper input format
|
||||||
|
pattern = r'[0-9-a-fA-F]{32}'
|
||||||
|
match = re.match(pattern, string)
|
||||||
|
if match is None:
|
||||||
|
raise ValueError(f'format does not match a hexadecimal-string: {string}')
|
||||||
|
return super().__new__(cls, string)
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
class_name = self.__class__.__name__
|
||||||
|
return f"{class_name}('{self}')"
|
||||||
|
|
||||||
|
def __reduce__(self):
|
||||||
|
return self.__class__, (self,)
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
"""Definition of HexadecimalString Pydantic Annotation."""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from pydantic.json_schema import JsonSchemaValue
|
||||||
|
from pydantic_core import core_schema
|
||||||
|
|
||||||
|
from .hexadecimal_string import HexadecimalString
|
||||||
|
|
||||||
|
|
||||||
|
class HexadecimalStringPydanticAnnotation:
|
||||||
|
"""Pydantic annotation that defines input validation, as well as general
|
||||||
|
and json serialization."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_input(cls, v: Any, handler) -> HexadecimalString:
|
||||||
|
"""Pydantic-related function to validate input on instantiation."""
|
||||||
|
if isinstance(v, HexadecimalString):
|
||||||
|
return v
|
||||||
|
s = handler(v)
|
||||||
|
return HexadecimalString(s)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def __get_pydantic_core_schema__(
|
||||||
|
cls,
|
||||||
|
source_type,
|
||||||
|
_handler,
|
||||||
|
) -> core_schema.CoreSchema:
|
||||||
|
assert source_type is HexadecimalString
|
||||||
|
return core_schema.no_info_wrap_validator_function(
|
||||||
|
function=cls.validate_input,
|
||||||
|
schema=core_schema.str_schema(),
|
||||||
|
serialization=core_schema.to_string_ser_schema(),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def __get_pydantic_json_schema__(cls, _core_schema, handler) -> JsonSchemaValue:
|
||||||
|
return handler(core_schema.str_schema())
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
"""Definition of VisualData DTO."""
|
||||||
|
|
||||||
|
from typing import Annotated
|
||||||
|
|
||||||
|
from PIL import Image
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
|
from .image_pydantic_annotation import ImagePydanticAnnotation
|
||||||
|
from .type_checking_base_model import TypeCheckingBaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class ImageData(TypeCheckingBaseModel):
|
||||||
|
"""Visual data class."""
|
||||||
|
|
||||||
|
image: Annotated[Image.Image, ImagePydanticAnnotation]
|
||||||
|
name: str = Field(min_length=1)
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
"""Definition of BytesIO Pydantic Annotation."""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from PIL import Image
|
||||||
|
from pydantic.json_schema import JsonSchemaValue
|
||||||
|
from pydantic_core import core_schema
|
||||||
|
|
||||||
|
|
||||||
|
class ImagePydanticAnnotation:
|
||||||
|
"""Pydantic annotation that defines input validation, as well as general
|
||||||
|
and json serialization."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_input(cls, v: Any, handler) -> Image.Image:
|
||||||
|
"""Pydantic-related function to validate input on instantiation."""
|
||||||
|
if isinstance(v, Image.Image):
|
||||||
|
return v
|
||||||
|
s = handler(v)
|
||||||
|
return Image.open(s)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def __get_pydantic_core_schema__(
|
||||||
|
cls,
|
||||||
|
source_type,
|
||||||
|
_handler,
|
||||||
|
) -> core_schema.CoreSchema:
|
||||||
|
assert source_type is Image.Image
|
||||||
|
return core_schema.no_info_wrap_validator_function(
|
||||||
|
function=cls.validate_input,
|
||||||
|
schema=core_schema.str_schema(),
|
||||||
|
serialization=core_schema.to_string_ser_schema(),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def __get_pydantic_json_schema__(cls, _core_schema, handler) -> JsonSchemaValue:
|
||||||
|
return handler(core_schema.str_schema())
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
"""Definition of ModelData DTO."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from hashlib import md5
|
||||||
|
from io import BytesIO
|
||||||
|
from typing import Annotated
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
|
from .bytes_io_pydantic_annotation import BytesIOPydanticAnnotation
|
||||||
|
from .hexadecimal_string import HexadecimalString
|
||||||
|
from .hexadecimal_string_pydantic_annotation import HexadecimalStringPydanticAnnotation
|
||||||
|
from .type_checking_base_model import TypeCheckingBaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class ModelData(TypeCheckingBaseModel):
|
||||||
|
"""Model Data DTO."""
|
||||||
|
|
||||||
|
buffer: Annotated[BytesIO, BytesIOPydanticAnnotation]
|
||||||
|
buffer_checksum: Annotated[HexadecimalString, HexadecimalStringPydanticAnnotation]
|
||||||
|
class_name: str = Field(
|
||||||
|
min_length=1,
|
||||||
|
description='name of model class to generate data.',
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def calculate_checksum(buffer: BytesIO) -> HexadecimalString:
|
||||||
|
"""Calculate buffer checksum."""
|
||||||
|
checksum = md5(buffer.getbuffer()).hexdigest()
|
||||||
|
return HexadecimalString(checksum)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def model_to_buffer(model: torch.nn.Module) -> BytesIO:
|
||||||
|
"""Save model to buffer."""
|
||||||
|
assert isinstance(model, torch.nn.Module)
|
||||||
|
buffer = BytesIO()
|
||||||
|
torch.save(model.state_dict(), buffer)
|
||||||
|
return buffer
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_model(cls, model: torch.nn.Module) -> ModelData:
|
||||||
|
"""Instantiate from torch module."""
|
||||||
|
assert isinstance(model, torch.nn.Module)
|
||||||
|
# get model name
|
||||||
|
class_name = type(model).__name__
|
||||||
|
# save data to buffer
|
||||||
|
buffer = cls.model_to_buffer(model)
|
||||||
|
buffer = BytesIO()
|
||||||
|
torch.save(model.state_dict(), buffer)
|
||||||
|
# calculate checksum
|
||||||
|
buffer_checksum = cls.calculate_checksum(buffer)
|
||||||
|
# instantiate from buffer
|
||||||
|
data = cls(
|
||||||
|
buffer=buffer,
|
||||||
|
buffer_checksum=buffer_checksum,
|
||||||
|
class_name=class_name,
|
||||||
|
)
|
||||||
|
return data
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of TypeCheckingBaseModel class."""
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict
|
||||||
|
|
||||||
|
|
||||||
|
class TypeCheckingBaseModel(BaseModel):
|
||||||
|
"""BaseModel with added type checking on input types."""
|
||||||
|
|
||||||
|
model_config = ConfigDict(
|
||||||
|
frozen=True, # ensure data immutability
|
||||||
|
)
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
from .visual_communication_data import VisualCommunicationData
|
||||||
|
from .visual_communication_values import VisualCommunicationValues
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of AngleValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class AngleValues(ValuesModel):
|
||||||
|
"""Angle values DTO."""
|
||||||
|
|
||||||
|
high: float
|
||||||
|
eye_level: float
|
||||||
|
low: float
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
"""Definition of ContactValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class ContactValues(ValuesModel):
|
||||||
|
"""Contact values DTO."""
|
||||||
|
|
||||||
|
offer: float
|
||||||
|
demand: float
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of DistanceValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class DistanceValues(ValuesModel):
|
||||||
|
"""Distance values DTO."""
|
||||||
|
|
||||||
|
long: float
|
||||||
|
medium: float
|
||||||
|
close: float
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
"""Definition of FramingValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class FramingValues(ValuesModel):
|
||||||
|
"""Framing values DTO."""
|
||||||
|
|
||||||
|
frame_lines: float
|
||||||
|
empty_space: float
|
||||||
|
colour_contrast: float
|
||||||
|
form_contrast: float
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of InformationValueValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class InformationValueValues(ValuesModel):
|
||||||
|
"""Information value values DTO."""
|
||||||
|
|
||||||
|
given_new: float
|
||||||
|
ideal_real: float
|
||||||
|
central_marginal: float
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of ModalityColorValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class ModalityColorValues(ValuesModel):
|
||||||
|
"""Modality color values DTO."""
|
||||||
|
|
||||||
|
high: float
|
||||||
|
medium: float
|
||||||
|
low: float
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of ModalityDepthValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class ModalityDepthValues(ValuesModel):
|
||||||
|
"""Modality depth values DTO."""
|
||||||
|
|
||||||
|
high: float
|
||||||
|
medium: float
|
||||||
|
low: float
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
"""Definition of ModalityLightingValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class ModalityLightingValues(ValuesModel):
|
||||||
|
"""Modality lighting values DTO."""
|
||||||
|
|
||||||
|
high: float
|
||||||
|
medium: float
|
||||||
|
low: float
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
"""Definition of PointOfViewValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class PointOfViewValues(ValuesModel):
|
||||||
|
"""Point-of-view values DTO."""
|
||||||
|
|
||||||
|
frontal: float
|
||||||
|
oblique: float
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
"""Definition of SalienceValues DTO."""
|
||||||
|
|
||||||
|
from .values_model import ValuesModel
|
||||||
|
|
||||||
|
|
||||||
|
class SalienceValues(ValuesModel):
|
||||||
|
"""Salience values DTO."""
|
||||||
|
|
||||||
|
size: float
|
||||||
|
colour: float
|
||||||
|
tone: float
|
||||||
|
form: float
|
||||||
|
positioning: float
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
"""Definition of ValuesModel base class."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import random
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class ValuesModel(BaseModel):
|
||||||
|
"""ValuesModel base class."""
|
||||||
|
|
||||||
|
model_config = ConfigDict(
|
||||||
|
validate_assignment=True, # argument type checking
|
||||||
|
frozen=True, # ensure data immutability
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_fields(cls) -> list[str]:
|
||||||
|
"""List options that are stored as attributes."""
|
||||||
|
return list(cls.model_fields.keys())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_random(cls):
|
||||||
|
"""Instantiate with random numbers."""
|
||||||
|
kwargs = {field: random.random() for field in cls.list_fields()}
|
||||||
|
return cls(**kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_choice(cls, option: str) -> ValuesModel:
|
||||||
|
"""Instantiate from choice."""
|
||||||
|
assert isinstance(option, str)
|
||||||
|
assert len(option) > 0
|
||||||
|
allowed_options_list = cls.list_fields()
|
||||||
|
if option not in allowed_options_list:
|
||||||
|
raise ValueError(f'option {option} must be in {allowed_options_list}')
|
||||||
|
# generate field values
|
||||||
|
kwargs = {field: 0 for field in allowed_options_list}
|
||||||
|
# set chosen value to max probability
|
||||||
|
kwargs[option] = 1
|
||||||
|
return cls(**kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_tensor(cls, tensor: Tensor):
|
||||||
|
"""Instantiate from list of values."""
|
||||||
|
assert tensor.size(dim=0) == 1, f'tensor batch larger than 1: {tensor}'
|
||||||
|
data_list = [float(t.item()) for t in tensor[0]]
|
||||||
|
kwargs = dict(zip(cls.list_fields(), data_list))
|
||||||
|
return cls(**kwargs)
|
||||||
|
|
||||||
|
def highest_score_field(self) -> str:
|
||||||
|
"""Return name of field with highest score."""
|
||||||
|
model_dict = self.model_dump()
|
||||||
|
return max(model_dict, key=lambda k: model_dict[k])
|
||||||
|
|
||||||
|
def highest_score_value(self) -> float:
|
||||||
|
"""Return value of field with highest score."""
|
||||||
|
model_dict = self.model_dump()
|
||||||
|
return max(model_dict.values())
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
"""Definition of VisualCommunicationData DTO."""
|
||||||
|
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
|
from ..type_checking_base_model import TypeCheckingBaseModel
|
||||||
|
from .visual_communication_values import VisualCommunicationValues
|
||||||
|
|
||||||
|
|
||||||
|
class VisualCommunicationData(TypeCheckingBaseModel):
|
||||||
|
"""Visual communication data class."""
|
||||||
|
|
||||||
|
name: str = Field(min_length=1)
|
||||||
|
annotation: VisualCommunicationValues | None = None
|
||||||
|
prediction: VisualCommunicationValues | None = None
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user