final demo #2
Browse files- .gitattributes +3 -0
- Uni_Sign/__pycache__/datasets.cpython-311.pyc +0 -0
- Uni_Sign/__pycache__/deformable_attention_2d.cpython-311.pyc +0 -0
- Uni_Sign/__pycache__/models.cpython-311.pyc +0 -0
- Uni_Sign/__pycache__/normalization.cpython-311.pyc +0 -0
- Uni_Sign/datasets.py +1178 -0
- Uni_Sign/deformable_attention_2d.py +311 -0
- Uni_Sign/models.py +415 -0
- Uni_Sign/normalization.py +117 -0
- Uni_Sign/stgcn_layers/.DS_Store +0 -0
- Uni_Sign/stgcn_layers/__init__.py +2 -0
- Uni_Sign/stgcn_layers/__pycache__/__init__.cpython-311.pyc +0 -0
- Uni_Sign/stgcn_layers/__pycache__/__init__.cpython-39.pyc +0 -0
- Uni_Sign/stgcn_layers/__pycache__/gcn_utils.cpython-311.pyc +0 -0
- Uni_Sign/stgcn_layers/__pycache__/gcn_utils.cpython-39.pyc +0 -0
- Uni_Sign/stgcn_layers/__pycache__/stgcn_block.cpython-311.pyc +0 -0
- Uni_Sign/stgcn_layers/__pycache__/stgcn_block.cpython-39.pyc +0 -0
- Uni_Sign/stgcn_layers/gcn_utils.py +350 -0
- Uni_Sign/stgcn_layers/stgcn_block.py +128 -0
- Uni_Sign/unisign_model/config.json +28 -0
- Uni_Sign/unisign_model/generation_config.json +7 -0
- Uni_Sign/unisign_model/special_tokens_map.json +1 -0
- Uni_Sign/unisign_model/spiece.model +3 -0
- Uni_Sign/unisign_model/tokenizer_config.json +1 -0
- checkpoints/pose/face_landmarker.task +3 -0
- checkpoints/pose/hand_landmarker.task +3 -0
- checkpoints/pose/pose_landmarker_full.task +3 -0
- checkpoints/pose/yolov8n-pose.pt +3 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,6 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
checkpoints/pose/face_landmarker.task filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
checkpoints/pose/hand_landmarker.task filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
checkpoints/pose/pose_landmarker_full.task filter=lfs diff=lfs merge=lfs -text
|
Uni_Sign/__pycache__/datasets.cpython-311.pyc
ADDED
|
Binary file (54.1 kB). View file
|
|
|
Uni_Sign/__pycache__/deformable_attention_2d.cpython-311.pyc
ADDED
|
Binary file (18.2 kB). View file
|
|
|
Uni_Sign/__pycache__/models.cpython-311.pyc
ADDED
|
Binary file (21.3 kB). View file
|
|
|
Uni_Sign/__pycache__/normalization.cpython-311.pyc
ADDED
|
Binary file (7 kB). View file
|
|
|
Uni_Sign/datasets.py
ADDED
|
@@ -0,0 +1,1178 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.utils.data.dataset as Dataset
|
| 3 |
+
from torch.nn.utils.rnn import pad_sequence
|
| 4 |
+
from PIL import Image
|
| 5 |
+
import os
|
| 6 |
+
import random
|
| 7 |
+
import numpy as np
|
| 8 |
+
import copy
|
| 9 |
+
import pickle
|
| 10 |
+
from decord import VideoReader, cpu
|
| 11 |
+
import json
|
| 12 |
+
import pathlib
|
| 13 |
+
import re
|
| 14 |
+
from torchvision import transforms
|
| 15 |
+
# from config import rgb_dirs, pose_dirs
|
| 16 |
+
from Uni_Sign.normalization import (local_keypoint_normalization, global_keypoint_normalization)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def all_same(keypoints):
|
| 20 |
+
return np.sum(keypoints == keypoints[0, 0]) == keypoints.size
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def sign_space_normalization(raw_keypoints, missing_values=None, layout='default'):
|
| 24 |
+
local_landmarks = {}
|
| 25 |
+
global_landmarks = {}
|
| 26 |
+
kp_normalization = ('global-body', 'local-right', 'local-left', 'local-face_all')
|
| 27 |
+
part_order = [i.removeprefix('local-').removeprefix('global-') for i in kp_normalization]
|
| 28 |
+
part_order = {k: v for v, k in enumerate(part_order)}
|
| 29 |
+
|
| 30 |
+
for idx, landmarks in enumerate(kp_normalization):
|
| 31 |
+
prefix, landmarks = landmarks.split("-")
|
| 32 |
+
if prefix == "local":
|
| 33 |
+
local_landmarks[idx] = landmarks
|
| 34 |
+
elif prefix == "global":
|
| 35 |
+
global_landmarks[idx] = landmarks
|
| 36 |
+
|
| 37 |
+
# local normalization
|
| 38 |
+
for idx, landmarks in local_landmarks.items():
|
| 39 |
+
normalized_keypoints = local_keypoint_normalization(raw_keypoints, landmarks, padding=0.2)
|
| 40 |
+
local_landmarks[idx] = normalized_keypoints
|
| 41 |
+
|
| 42 |
+
# global normalization
|
| 43 |
+
additional_landmarks = list(global_landmarks.values())
|
| 44 |
+
if "body" in additional_landmarks:
|
| 45 |
+
additional_landmarks.remove("body")
|
| 46 |
+
|
| 47 |
+
if layout == 'default':
|
| 48 |
+
l_shoulder_idx, r_shoulder_idx = 11, 12
|
| 49 |
+
else:
|
| 50 |
+
l_shoulder_idx, r_shoulder_idx = 3, 4
|
| 51 |
+
keypoints, additional_keypoints = global_keypoint_normalization(
|
| 52 |
+
raw_keypoints,
|
| 53 |
+
"body",
|
| 54 |
+
additional_landmarks,
|
| 55 |
+
l_shoulder_idx=l_shoulder_idx,
|
| 56 |
+
r_shoulder_idx=r_shoulder_idx,
|
| 57 |
+
)
|
| 58 |
+
|
| 59 |
+
for k, landmark in global_landmarks.items():
|
| 60 |
+
if landmark == "body":
|
| 61 |
+
global_landmarks[k] = keypoints
|
| 62 |
+
else:
|
| 63 |
+
global_landmarks[k] = additional_keypoints[landmark]
|
| 64 |
+
|
| 65 |
+
all_landmarks = {**local_landmarks, **global_landmarks}
|
| 66 |
+
all_landmarks_per_part = {k: all_landmarks[v] for k, v in part_order.items()}
|
| 67 |
+
|
| 68 |
+
if missing_values is not None:
|
| 69 |
+
for part, data in all_landmarks_per_part.items():
|
| 70 |
+
for fidx in range(len(data)):
|
| 71 |
+
if not all_same(data[fidx]):
|
| 72 |
+
continue
|
| 73 |
+
all_landmarks_per_part[part][fidx] = np.zeros_like(data[fidx]) + missing_values
|
| 74 |
+
|
| 75 |
+
return all_landmarks_per_part
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def load_part_kp_YTASL(skeletons, confs, normalization, layout):
|
| 79 |
+
thr = 0.3
|
| 80 |
+
# kps_with_scores = {}
|
| 81 |
+
kps_all_parts = {}
|
| 82 |
+
confs_all_parts = {}
|
| 83 |
+
scale = None
|
| 84 |
+
|
| 85 |
+
for part in ['body', 'left', 'right', 'face_all']:
|
| 86 |
+
kps = []
|
| 87 |
+
confidences = []
|
| 88 |
+
for i, (skeleton, conf) in enumerate(zip(skeletons, confs)):
|
| 89 |
+
|
| 90 |
+
if part == 'body':
|
| 91 |
+
if layout == 'default':
|
| 92 |
+
hand_kp2d = np.stack(skeleton['pose_landmarks'][:25])
|
| 93 |
+
confidence = np.stack(conf['pose_landmarks'][:25])
|
| 94 |
+
elif layout == 'pruned':
|
| 95 |
+
pose_landmarks = [0, 7, 8, 11, 12, 13, 14, 15, 16]
|
| 96 |
+
hand_kp2d = np.stack([skeleton['pose_landmarks'][i] for i in pose_landmarks])
|
| 97 |
+
confidence = np.stack([conf['pose_landmarks'][i] for i in pose_landmarks])
|
| 98 |
+
elif layout == 'isharah':
|
| 99 |
+
pose_landmarks = [0, 7, 8, 11, 12, 13, 14, 15, 16]
|
| 100 |
+
hand_kp2d = np.stack([skeleton['pose_landmarks'][i] for i in pose_landmarks])
|
| 101 |
+
confidence = np.stack([conf['pose_landmarks'][i] for i in pose_landmarks])
|
| 102 |
+
|
| 103 |
+
elif part == 'left':
|
| 104 |
+
if layout in ['default', 'pruned', 'isharah']:
|
| 105 |
+
hand_kp2d = np.stack(skeleton['left_hand_landmarks'])
|
| 106 |
+
confidence = np.stack(conf['left_hand_landmarks'])
|
| 107 |
+
|
| 108 |
+
elif part == 'right':
|
| 109 |
+
if layout in ['default', 'pruned', 'isharah']:
|
| 110 |
+
hand_kp2d = np.stack(skeleton['right_hand_landmarks'])
|
| 111 |
+
confidence = np.stack(conf['right_hand_landmarks'])
|
| 112 |
+
|
| 113 |
+
elif part == 'face_all':
|
| 114 |
+
if layout == 'default':
|
| 115 |
+
face_landmarks = [
|
| 116 |
+
0, 4, 13, 14, 17, 33, 39, 46, 52, 55, 61, 64, 81,
|
| 117 |
+
93, 133, 151, 152, 159, 172, 178, 181, 263, 269, 276,
|
| 118 |
+
282, 285, 291, 294, 311, 323, 362, 386, 397, 402, 405, 468, 473
|
| 119 |
+
]
|
| 120 |
+
hand_kp2d = np.stack([skeleton['face_landmarks'][i] for i in face_landmarks])
|
| 121 |
+
confidence = np.stack([conf['face_landmarks'][i] for i in face_landmarks])
|
| 122 |
+
elif layout == 'pruned':
|
| 123 |
+
face_landmarks = [4, 13, 14, 61, 81, 93, 152, 159, 172, 178, 291, 311, 323, 386, 397, 402, 472, 477]
|
| 124 |
+
hand_kp2d = np.stack([skeleton['face_landmarks'][i] for i in face_landmarks])
|
| 125 |
+
confidence = np.stack([conf['face_landmarks'][i] for i in face_landmarks])
|
| 126 |
+
elif layout == 'isharah':
|
| 127 |
+
face_landmarks = [0, 17, 37, 39, 40, 61, 84, 91, 146, 181, 185, 267, 269, 270, 291, 314, 321, 375, 405]
|
| 128 |
+
hand_kp2d = np.stack([skeleton['face_landmarks'][i] for i in face_landmarks])
|
| 129 |
+
confidence = np.stack([conf['face_landmarks'][i] for i in face_landmarks])
|
| 130 |
+
|
| 131 |
+
else:
|
| 132 |
+
raise NotImplementedError
|
| 133 |
+
kps.append(hand_kp2d)
|
| 134 |
+
confidences.append(confidence)
|
| 135 |
+
|
| 136 |
+
kps = np.stack(kps, axis=0)
|
| 137 |
+
confidences = np.stack(confidences, axis=0)
|
| 138 |
+
|
| 139 |
+
kps_all_parts[part] = kps
|
| 140 |
+
confs_all_parts[part] = confidences[..., None]
|
| 141 |
+
|
| 142 |
+
if normalization == 'signspace':
|
| 143 |
+
normalized_kps = sign_space_normalization(kps_all_parts.copy(), layout=layout)
|
| 144 |
+
else:
|
| 145 |
+
normalized_kps = kps_all_parts
|
| 146 |
+
|
| 147 |
+
kps_with_scores = {}
|
| 148 |
+
for part in normalized_kps.keys():
|
| 149 |
+
kps_with_scores[part] = np.concatenate([normalized_kps[part], confs_all_parts[part]], axis=-1)
|
| 150 |
+
|
| 151 |
+
kps_with_scores = {k: torch.as_tensor(v, dtype=torch.float32) for k, v in kps_with_scores.items()}
|
| 152 |
+
return kps_with_scores
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def load_part_kp_Isharah(skeletons, confs, normalization, layout):
|
| 156 |
+
# kps_with_scores = {}
|
| 157 |
+
kps_all_parts = {}
|
| 158 |
+
confs_all_parts = {}
|
| 159 |
+
|
| 160 |
+
for part in ['body', 'left', 'right', 'face_all']:
|
| 161 |
+
kps = []
|
| 162 |
+
confidences = []
|
| 163 |
+
for i, (skeleton, conf) in enumerate(zip(skeletons, confs)):
|
| 164 |
+
|
| 165 |
+
if part == 'body':
|
| 166 |
+
pose_landmarks = [0, 7, 8, 11, 12, 13, 14, 15, 16]
|
| 167 |
+
hand_kp2d = np.stack([skeleton['pose_landmarks'][i] for i in pose_landmarks])
|
| 168 |
+
confidence = np.stack([conf['pose_landmarks'][i] for i in pose_landmarks])
|
| 169 |
+
|
| 170 |
+
elif part == 'left':
|
| 171 |
+
hand_kp2d = np.stack(skeleton['left_hand_landmarks'])
|
| 172 |
+
confidence = np.stack(conf['left_hand_landmarks'])
|
| 173 |
+
|
| 174 |
+
elif part == 'right':
|
| 175 |
+
hand_kp2d = np.stack(skeleton['right_hand_landmarks'])
|
| 176 |
+
confidence = np.stack(conf['right_hand_landmarks'])
|
| 177 |
+
|
| 178 |
+
elif part == 'face_all':
|
| 179 |
+
hand_kp2d = np.stack(skeleton['face_landmarks'])
|
| 180 |
+
confidence = np.stack(conf['face_landmarks'])
|
| 181 |
+
|
| 182 |
+
else:
|
| 183 |
+
raise NotImplementedError
|
| 184 |
+
kps.append(hand_kp2d)
|
| 185 |
+
confidences.append(confidence)
|
| 186 |
+
|
| 187 |
+
kps = np.stack(kps, axis=0)
|
| 188 |
+
confidences = np.stack(confidences, axis=0)
|
| 189 |
+
|
| 190 |
+
kps_all_parts[part] = kps
|
| 191 |
+
confs_all_parts[part] = confidences[..., None]
|
| 192 |
+
|
| 193 |
+
if normalization == 'signspace':
|
| 194 |
+
normalized_kps = sign_space_normalization(kps_all_parts.copy(), layout=layout)
|
| 195 |
+
else:
|
| 196 |
+
normalized_kps = kps_all_parts
|
| 197 |
+
|
| 198 |
+
kps_with_scores = {}
|
| 199 |
+
for part in normalized_kps.keys():
|
| 200 |
+
kps_with_scores[part] = np.concatenate([normalized_kps[part], confs_all_parts[part]], axis=-1)
|
| 201 |
+
|
| 202 |
+
kps_with_scores = {k: torch.as_tensor(v, dtype=torch.float32) for k, v in kps_with_scores.items()}
|
| 203 |
+
return kps_with_scores
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
YTASL_GROUP_SIZES = {
|
| 207 |
+
'pose_landmarks': 33,
|
| 208 |
+
'right_hand_landmarks': 21,
|
| 209 |
+
'left_hand_landmarks': 21,
|
| 210 |
+
'face_landmarks': 478,
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
YTASL_GROUP_ERROR_LABELS = {
|
| 214 |
+
'pose_landmarks': 'a pose group',
|
| 215 |
+
'right_hand_landmarks': 'a Rhand group',
|
| 216 |
+
'left_hand_landmarks': 'a Lhand group',
|
| 217 |
+
'face_landmarks': 'a face group',
|
| 218 |
+
}
|
| 219 |
+
|
| 220 |
+
ISHARAH_GROUP_SIZES = {
|
| 221 |
+
'pose_landmarks': 25,
|
| 222 |
+
'right_hand_landmarks': 21,
|
| 223 |
+
'left_hand_landmarks': 21,
|
| 224 |
+
'face_landmarks': 19,
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
def _fill_missing_landmarks(
|
| 229 |
+
skeleton,
|
| 230 |
+
conf,
|
| 231 |
+
group_name,
|
| 232 |
+
expected_size,
|
| 233 |
+
clip_name,
|
| 234 |
+
frame_idx,
|
| 235 |
+
error_group_label=None,
|
| 236 |
+
include_size_details=True,
|
| 237 |
+
strict_key_access=False,
|
| 238 |
+
):
|
| 239 |
+
points = skeleton[group_name] if strict_key_access else skeleton.get(group_name, [])
|
| 240 |
+
if len(points) == 0:
|
| 241 |
+
conf[group_name] = [0] * expected_size
|
| 242 |
+
skeleton[group_name] = [[0.0, 0.0]] * expected_size
|
| 243 |
+
elif len(points) != expected_size:
|
| 244 |
+
group_label = error_group_label or f"group '{group_name}'"
|
| 245 |
+
if include_size_details:
|
| 246 |
+
raise NotImplementedError(
|
| 247 |
+
f"Unexpected number of keypoints in {group_label}: {clip_name}, frame {frame_idx}, "
|
| 248 |
+
f"expected {expected_size}, got {len(points)}"
|
| 249 |
+
)
|
| 250 |
+
raise NotImplementedError(f"Unexpected number of keypoints in {group_label}: {clip_name}, {frame_idx}")
|
| 251 |
+
else:
|
| 252 |
+
conf[group_name] = [1] * expected_size
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
# load sub-pose
|
| 256 |
+
def load_part_kp(skeletons, confs, force_ok=False):
|
| 257 |
+
thr = 0.3
|
| 258 |
+
kps_with_scores = {}
|
| 259 |
+
scale = None
|
| 260 |
+
|
| 261 |
+
for part in ['body', 'left', 'right', 'face_all']:
|
| 262 |
+
kps = []
|
| 263 |
+
confidences = []
|
| 264 |
+
|
| 265 |
+
for skeleton, conf in zip(skeletons, confs):
|
| 266 |
+
if skeleton.ndim == 4: # if (1,133,2) - wrapped in list
|
| 267 |
+
skeleton = skeleton[0]
|
| 268 |
+
conf = conf[0]
|
| 269 |
+
|
| 270 |
+
if part == 'body': # [0, 3, 4, 5, 6, 7, 8, 9, 10]
|
| 271 |
+
hand_kp2d = skeleton[[0] + [i for i in range(3, 11)], :]
|
| 272 |
+
confidence = conf[[0] + [i for i in range(3, 11)]]
|
| 273 |
+
elif part == 'left': # [91, 92, 93, 94, 95, 96, 97, 98, 99, 100, 101, 102, 103, 104, 105, 106, 107, 108, 109, 110, 111]
|
| 274 |
+
hand_kp2d = skeleton[91:112, :]
|
| 275 |
+
hand_kp2d = hand_kp2d - hand_kp2d[0, :]
|
| 276 |
+
confidence = conf[91:112]
|
| 277 |
+
elif part == 'right': # [112, 113, 114, 115, 116, 117, 118, 119, 120, 121, 122, 123, 124, 125, 126, 127, 128, 129, 130, 131, 132]
|
| 278 |
+
hand_kp2d = skeleton[112:133, :]
|
| 279 |
+
hand_kp2d = hand_kp2d - hand_kp2d[0, :]
|
| 280 |
+
confidence = conf[112:133]
|
| 281 |
+
elif part == 'face_all': # [23, 25, 27, 29, 31, 33, 35, 37, 39, 83, 84, 85, 86, 87, 88, 89, 90, 53]
|
| 282 |
+
hand_kp2d = skeleton[[i for i in list(range(23, 23 + 17))[::2]] + [i for i in range(83, 83 + 8)] + [53], :]
|
| 283 |
+
hand_kp2d = hand_kp2d - hand_kp2d[-1, :]
|
| 284 |
+
confidence = conf[[i for i in list(range(23, 23 + 17))[::2]] + [i for i in range(83, 83 + 8)] + [53]]
|
| 285 |
+
|
| 286 |
+
else:
|
| 287 |
+
raise NotImplementedError
|
| 288 |
+
|
| 289 |
+
kps.append(hand_kp2d)
|
| 290 |
+
confidences.append(confidence)
|
| 291 |
+
|
| 292 |
+
kps = np.stack(kps, axis=0)
|
| 293 |
+
confidences = np.stack(confidences, axis=0)
|
| 294 |
+
|
| 295 |
+
if part == 'body':
|
| 296 |
+
if force_ok:
|
| 297 |
+
result, scale, _ = crop_scale(np.concatenate([kps, confidences[..., None]], axis=-1), thr)
|
| 298 |
+
|
| 299 |
+
else:
|
| 300 |
+
result, scale, _ = crop_scale(np.concatenate([kps, confidences[..., None]], axis=-1), thr)
|
| 301 |
+
else:
|
| 302 |
+
assert not scale is None
|
| 303 |
+
result = np.concatenate([kps, confidences[..., None]], axis=-1)
|
| 304 |
+
if scale == 0:
|
| 305 |
+
result = np.zeros(result.shape)
|
| 306 |
+
else:
|
| 307 |
+
result[..., :2] = (result[..., :2]) / scale
|
| 308 |
+
result = np.clip(result, -1, 1)
|
| 309 |
+
# mask useless kp
|
| 310 |
+
result[result[..., 2] <= thr] = 0
|
| 311 |
+
|
| 312 |
+
kps_with_scores[part] = torch.tensor(result)
|
| 313 |
+
|
| 314 |
+
return kps_with_scores
|
| 315 |
+
|
| 316 |
+
|
| 317 |
+
# input: T, N, 3
|
| 318 |
+
# input is un-normed joints
|
| 319 |
+
def crop_scale(motion, thr):
|
| 320 |
+
'''
|
| 321 |
+
Motion: [(M), T, 17, 3].
|
| 322 |
+
Normalize to [-1, 1]
|
| 323 |
+
'''
|
| 324 |
+
result = copy.deepcopy(motion)
|
| 325 |
+
valid_coords = motion[motion[..., 2] > thr][:, :2]
|
| 326 |
+
if len(valid_coords) < 4:
|
| 327 |
+
return np.zeros(motion.shape), 0, None
|
| 328 |
+
xmin = min(valid_coords[:, 0])
|
| 329 |
+
xmax = max(valid_coords[:, 0])
|
| 330 |
+
ymin = min(valid_coords[:, 1])
|
| 331 |
+
ymax = max(valid_coords[:, 1])
|
| 332 |
+
# ratio = np.random.uniform(low=scale_range[0], high=scale_range[1], size=1)[0]
|
| 333 |
+
ratio = 1
|
| 334 |
+
scale = max(xmax - xmin, ymax - ymin) * ratio
|
| 335 |
+
if scale == 0:
|
| 336 |
+
return np.zeros(motion.shape), 0, None
|
| 337 |
+
xs = (xmin + xmax - scale) / 2
|
| 338 |
+
ys = (ymin + ymax - scale) / 2
|
| 339 |
+
result[..., :2] = (motion[..., :2] - [xs, ys]) / scale
|
| 340 |
+
result[..., :2] = (result[..., :2] - 0.5) * 2
|
| 341 |
+
result = np.clip(result, -1, 1)
|
| 342 |
+
# mask useless kp
|
| 343 |
+
result[result[..., 2] <= thr] = 0
|
| 344 |
+
return result, scale, [xs, ys]
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
# bbox of hands
|
| 348 |
+
def bbox_4hands(left_keypoints, right_keypoints, hw):
|
| 349 |
+
# keypoints --> T,21,2
|
| 350 |
+
# keypoints --> T,21,2
|
| 351 |
+
|
| 352 |
+
def compute_bbox(keypoints):
|
| 353 |
+
min_x = np.min(keypoints[..., 0], axis=1)
|
| 354 |
+
min_y = np.min(keypoints[..., 1], axis=1)
|
| 355 |
+
max_x = np.max(keypoints[..., 0], axis=1)
|
| 356 |
+
max_y = np.max(keypoints[..., 1], axis=1)
|
| 357 |
+
|
| 358 |
+
return (max_x + min_x) / 2, (max_y + min_y) / 2, (max_x - min_x), (max_y - min_y)
|
| 359 |
+
|
| 360 |
+
H, W = hw
|
| 361 |
+
|
| 362 |
+
if left_keypoints is None:
|
| 363 |
+
left_keypoints = np.zeros([1, 21, 2])
|
| 364 |
+
|
| 365 |
+
if right_keypoints is None:
|
| 366 |
+
right_keypoints = np.zeros([1, 21, 2])
|
| 367 |
+
# [T, 21, 2]
|
| 368 |
+
left_mean_x, left_mean_y, left_diff_x, left_diff_y = compute_bbox(left_keypoints)
|
| 369 |
+
left_mean_x = W * left_mean_x
|
| 370 |
+
left_mean_y = H * left_mean_y
|
| 371 |
+
|
| 372 |
+
left_diff_x = W * left_diff_x
|
| 373 |
+
left_diff_y = H * left_diff_y
|
| 374 |
+
|
| 375 |
+
left_diff_x = max(left_diff_x)
|
| 376 |
+
left_diff_y = max(left_diff_y)
|
| 377 |
+
left_box_hw = max(left_diff_x, left_diff_y)
|
| 378 |
+
|
| 379 |
+
right_mean_x, right_mean_y, right_diff_x, right_diff_y = compute_bbox(right_keypoints)
|
| 380 |
+
right_mean_x = W * right_mean_x
|
| 381 |
+
right_mean_y = H * right_mean_y
|
| 382 |
+
|
| 383 |
+
right_diff_x = W * right_diff_x
|
| 384 |
+
right_diff_y = H * right_diff_y
|
| 385 |
+
|
| 386 |
+
right_diff_x = max(right_diff_x)
|
| 387 |
+
right_diff_y = max(right_diff_y)
|
| 388 |
+
right_box_hw = max(right_diff_x, right_diff_y)
|
| 389 |
+
|
| 390 |
+
box_hw = int(max(left_box_hw, right_box_hw) * 1.2 / 2) * 2
|
| 391 |
+
box_hw = max(box_hw, 0)
|
| 392 |
+
|
| 393 |
+
left_new_box = np.stack([left_mean_x - box_hw / 2, left_mean_y - box_hw / 2, left_mean_x + box_hw / 2,
|
| 394 |
+
left_mean_y + box_hw / 2]).astype(np.int16)
|
| 395 |
+
right_new_box = np.stack([right_mean_x - box_hw / 2, right_mean_y - box_hw / 2, right_mean_x + box_hw / 2,
|
| 396 |
+
right_mean_y + box_hw / 2]).astype(np.int16)
|
| 397 |
+
|
| 398 |
+
return left_new_box.transpose(1, 0), right_new_box.transpose(1, 0), box_hw
|
| 399 |
+
|
| 400 |
+
|
| 401 |
+
def load_support_rgb_dict(tmp, skeletons, confs, full_path, data_transform):
|
| 402 |
+
support_rgb_dict = {}
|
| 403 |
+
|
| 404 |
+
confs = np.array(confs)
|
| 405 |
+
skeletons = np.array(skeletons)
|
| 406 |
+
|
| 407 |
+
# sample index of low scores
|
| 408 |
+
left_confs_filter = confs[:, 0, 91:112].mean(-1)
|
| 409 |
+
left_confs_filter_indices = np.where(left_confs_filter > 0.3)[0]
|
| 410 |
+
|
| 411 |
+
if len(left_confs_filter_indices) == 0:
|
| 412 |
+
left_sampled_indices = None
|
| 413 |
+
left_skeletons = None
|
| 414 |
+
else:
|
| 415 |
+
|
| 416 |
+
left_confs = confs[left_confs_filter_indices]
|
| 417 |
+
left_confs = left_confs[:, 0, [95, 99, 103, 107, 111]].min(-1)
|
| 418 |
+
|
| 419 |
+
left_weights = np.max(left_confs) - left_confs + 1e-5
|
| 420 |
+
left_probabilities = left_weights / np.sum(left_weights)
|
| 421 |
+
|
| 422 |
+
left_sample_size = int(np.ceil(0.1 * len(left_confs_filter_indices)))
|
| 423 |
+
|
| 424 |
+
left_sampled_indices = np.random.choice(left_confs_filter_indices.tolist(),
|
| 425 |
+
size=left_sample_size,
|
| 426 |
+
replace=False,
|
| 427 |
+
p=left_probabilities)
|
| 428 |
+
# left_sampled_indices: values: 0-255(0,max_len)
|
| 429 |
+
# tmp: values: 0-(end-start)
|
| 430 |
+
left_sampled_indices = np.sort(left_sampled_indices)
|
| 431 |
+
|
| 432 |
+
left_skeletons = skeletons[left_sampled_indices, 0, 91:112]
|
| 433 |
+
|
| 434 |
+
right_confs_filter = confs[:, 0, 112:].mean(-1)
|
| 435 |
+
right_confs_filter_indices = np.where(right_confs_filter > 0.3)[0]
|
| 436 |
+
if len(right_confs_filter_indices) == 0:
|
| 437 |
+
right_sampled_indices = None
|
| 438 |
+
right_skeletons = None
|
| 439 |
+
|
| 440 |
+
else:
|
| 441 |
+
right_confs = confs[right_confs_filter_indices]
|
| 442 |
+
right_confs = right_confs[:, 0, [95 + 21, 99 + 21, 103 + 21, 107 + 21, 111 + 21]].min(-1)
|
| 443 |
+
|
| 444 |
+
right_weights = np.max(right_confs) - right_confs + 1e-5
|
| 445 |
+
right_probabilities = right_weights / np.sum(right_weights)
|
| 446 |
+
|
| 447 |
+
right_sample_size = int(np.ceil(0.1 * len(right_confs_filter_indices)))
|
| 448 |
+
|
| 449 |
+
right_sampled_indices = np.random.choice(right_confs_filter_indices.tolist(),
|
| 450 |
+
size=right_sample_size,
|
| 451 |
+
replace=False,
|
| 452 |
+
p=right_probabilities)
|
| 453 |
+
right_sampled_indices = np.sort(right_sampled_indices)
|
| 454 |
+
|
| 455 |
+
right_skeletons = skeletons[right_sampled_indices, 0, 112:133]
|
| 456 |
+
|
| 457 |
+
image_size = 112
|
| 458 |
+
all_indices = []
|
| 459 |
+
if not left_sampled_indices is None:
|
| 460 |
+
all_indices.append(left_sampled_indices)
|
| 461 |
+
if not right_sampled_indices is None:
|
| 462 |
+
all_indices.append(right_sampled_indices)
|
| 463 |
+
if len(all_indices) == 0:
|
| 464 |
+
support_rgb_dict['left_sampled_indices'] = torch.tensor([-1])
|
| 465 |
+
support_rgb_dict['left_hands'] = torch.zeros(1, 3, image_size, image_size)
|
| 466 |
+
support_rgb_dict['left_skeletons_norm'] = torch.zeros(1, 21, 2)
|
| 467 |
+
|
| 468 |
+
support_rgb_dict['right_sampled_indices'] = torch.tensor([-1])
|
| 469 |
+
support_rgb_dict['right_hands'] = torch.zeros(1, 3, image_size, image_size)
|
| 470 |
+
support_rgb_dict['right_skeletons_norm'] = torch.zeros(1, 21, 2)
|
| 471 |
+
|
| 472 |
+
return support_rgb_dict
|
| 473 |
+
|
| 474 |
+
sampled_indices = np.concatenate(all_indices)
|
| 475 |
+
sampled_indices = np.unique(sampled_indices)
|
| 476 |
+
sampled_indices_real = tmp[sampled_indices]
|
| 477 |
+
|
| 478 |
+
# load image sample
|
| 479 |
+
imgs = load_video_support_rgb(full_path, sampled_indices_real)
|
| 480 |
+
|
| 481 |
+
# get hand bbox
|
| 482 |
+
left_new_box, right_new_box, box_hw = bbox_4hands(left_skeletons,
|
| 483 |
+
right_skeletons,
|
| 484 |
+
imgs[0].shape[:2])
|
| 485 |
+
|
| 486 |
+
# crop left and right hand
|
| 487 |
+
image_size = 112
|
| 488 |
+
if box_hw == 0:
|
| 489 |
+
support_rgb_dict['left_sampled_indices'] = torch.tensor([-1])
|
| 490 |
+
support_rgb_dict['left_hands'] = torch.zeros(1, 3, image_size, image_size)
|
| 491 |
+
support_rgb_dict['left_skeletons_norm'] = torch.zeros(1, 21, 2)
|
| 492 |
+
|
| 493 |
+
support_rgb_dict['right_sampled_indices'] = torch.tensor([-1])
|
| 494 |
+
support_rgb_dict['right_hands'] = torch.zeros(1, 3, image_size, image_size)
|
| 495 |
+
support_rgb_dict['right_skeletons_norm'] = torch.zeros(1, 21, 2)
|
| 496 |
+
|
| 497 |
+
return support_rgb_dict
|
| 498 |
+
|
| 499 |
+
factor = image_size / box_hw
|
| 500 |
+
|
| 501 |
+
if left_sampled_indices is None:
|
| 502 |
+
left_hands = torch.zeros(1, 3, image_size, image_size)
|
| 503 |
+
left_skeletons_norm = torch.zeros(1, 21, 2)
|
| 504 |
+
|
| 505 |
+
else:
|
| 506 |
+
left_hands = torch.zeros(len(left_sampled_indices), 3, image_size, image_size)
|
| 507 |
+
|
| 508 |
+
left_skeletons_norm = left_skeletons * imgs[0].shape[:2][::-1] - left_new_box[:, None, [0, 1]]
|
| 509 |
+
left_skeletons_norm = left_skeletons_norm / box_hw
|
| 510 |
+
left_skeletons_norm = left_skeletons_norm.clip(0, 1)
|
| 511 |
+
|
| 512 |
+
if right_sampled_indices is None:
|
| 513 |
+
right_hands = torch.zeros(1, 3, image_size, image_size)
|
| 514 |
+
right_skeletons_norm = torch.zeros(1, 21, 2)
|
| 515 |
+
|
| 516 |
+
else:
|
| 517 |
+
right_hands = torch.zeros(len(right_sampled_indices), 3, image_size, image_size)
|
| 518 |
+
|
| 519 |
+
right_skeletons_norm = right_skeletons * imgs[0].shape[:2][::-1] - right_new_box[:, None, [0, 1]]
|
| 520 |
+
right_skeletons_norm = right_skeletons_norm / box_hw
|
| 521 |
+
right_skeletons_norm = right_skeletons_norm.clip(0, 1)
|
| 522 |
+
left_idx = 0
|
| 523 |
+
right_idx = 0
|
| 524 |
+
|
| 525 |
+
for idx, img in enumerate(imgs):
|
| 526 |
+
mapping_idx = sampled_indices[idx]
|
| 527 |
+
if not left_sampled_indices is None and left_idx < len(left_sampled_indices) and mapping_idx == \
|
| 528 |
+
left_sampled_indices[left_idx]:
|
| 529 |
+
box = left_new_box[left_idx]
|
| 530 |
+
|
| 531 |
+
img_draw = np.uint8(copy.deepcopy(img))[box[1]:box[3], box[0]:box[2], :]
|
| 532 |
+
img_draw = np.pad(img_draw,
|
| 533 |
+
((0, max(0, box_hw - img_draw.shape[0])), (0, max(0, box_hw - img_draw.shape[1])),
|
| 534 |
+
(0, 0)), mode='constant', constant_values=0)
|
| 535 |
+
|
| 536 |
+
f_img = Image.fromarray(img_draw).convert('RGB').resize((image_size, image_size))
|
| 537 |
+
f_img = data_transform(f_img).unsqueeze(0)
|
| 538 |
+
left_hands[left_idx] = f_img
|
| 539 |
+
left_idx += 1
|
| 540 |
+
|
| 541 |
+
if not right_sampled_indices is None and right_idx < len(right_sampled_indices) and mapping_idx == \
|
| 542 |
+
right_sampled_indices[right_idx]:
|
| 543 |
+
box = right_new_box[right_idx]
|
| 544 |
+
|
| 545 |
+
img_draw = np.uint8(copy.deepcopy(img))[box[1]:box[3], box[0]:box[2], :]
|
| 546 |
+
img_draw = np.pad(img_draw,
|
| 547 |
+
((0, max(0, box_hw - img_draw.shape[0])), (0, max(0, box_hw - img_draw.shape[1])),
|
| 548 |
+
(0, 0)), mode='constant', constant_values=0)
|
| 549 |
+
|
| 550 |
+
f_img = Image.fromarray(img_draw).convert('RGB').resize((image_size, image_size))
|
| 551 |
+
f_img = data_transform(f_img).unsqueeze(0)
|
| 552 |
+
right_hands[right_idx] = f_img
|
| 553 |
+
right_idx += 1
|
| 554 |
+
|
| 555 |
+
if left_sampled_indices is None:
|
| 556 |
+
left_sampled_indices = np.array([-1])
|
| 557 |
+
|
| 558 |
+
if right_sampled_indices is None:
|
| 559 |
+
right_sampled_indices = np.array([-1])
|
| 560 |
+
|
| 561 |
+
# get index, images and keypoints priors
|
| 562 |
+
support_rgb_dict['left_sampled_indices'] = torch.tensor(left_sampled_indices)
|
| 563 |
+
support_rgb_dict['left_hands'] = left_hands
|
| 564 |
+
support_rgb_dict['left_skeletons_norm'] = torch.tensor(left_skeletons_norm)
|
| 565 |
+
|
| 566 |
+
support_rgb_dict['right_sampled_indices'] = torch.tensor(right_sampled_indices)
|
| 567 |
+
support_rgb_dict['right_hands'] = right_hands
|
| 568 |
+
support_rgb_dict['right_skeletons_norm'] = torch.tensor(right_skeletons_norm)
|
| 569 |
+
|
| 570 |
+
return support_rgb_dict
|
| 571 |
+
|
| 572 |
+
|
| 573 |
+
# use split rgb video for save time
|
| 574 |
+
def load_video_support_rgb(path, tmp):
|
| 575 |
+
vr = VideoReader(path, num_threads=1, ctx=cpu(0))
|
| 576 |
+
|
| 577 |
+
vr.seek(0)
|
| 578 |
+
buffer = vr.get_batch(tmp).asnumpy()
|
| 579 |
+
batch_image = buffer
|
| 580 |
+
del vr
|
| 581 |
+
|
| 582 |
+
return batch_image
|
| 583 |
+
|
| 584 |
+
|
| 585 |
+
def load_json(path):
|
| 586 |
+
with open(path, "r", encoding="utf-8") as f:
|
| 587 |
+
return json.load(f)
|
| 588 |
+
|
| 589 |
+
def is_valid_metric_label(text):
|
| 590 |
+
if text is None:
|
| 591 |
+
return False
|
| 592 |
+
text = " ".join(str(text).split()).strip()
|
| 593 |
+
if not text:
|
| 594 |
+
return False
|
| 595 |
+
# Require at least one alnum/letter token after punctuation/symbol stripping.
|
| 596 |
+
return re.search(r"\w", text, flags=re.UNICODE) is not None
|
| 597 |
+
|
| 598 |
+
def select_frame_indices(duration, max_length, phase):
|
| 599 |
+
if duration <= max_length:
|
| 600 |
+
return list(range(duration))
|
| 601 |
+
if phase == 'train':
|
| 602 |
+
return sorted(random.sample(range(duration), k=max_length))
|
| 603 |
+
# Deterministic, near-uniform coverage for dev/test.
|
| 604 |
+
return ((np.arange(max_length) * duration) // max_length).tolist()
|
| 605 |
+
|
| 606 |
+
|
| 607 |
+
# build base dataset
|
| 608 |
+
class Base_Dataset(Dataset.Dataset):
|
| 609 |
+
def collate_fn(self, batch):
|
| 610 |
+
tgt_batch, src_length_batch, name_batch, pose_tmp, gloss_batch = [], [], [], [], []
|
| 611 |
+
|
| 612 |
+
for name_sample, pose_sample, text, gloss, _ in batch:
|
| 613 |
+
name_batch.append(name_sample)
|
| 614 |
+
pose_tmp.append(pose_sample)
|
| 615 |
+
tgt_batch.append(text)
|
| 616 |
+
gloss_batch.append(gloss)
|
| 617 |
+
|
| 618 |
+
src_input = {}
|
| 619 |
+
|
| 620 |
+
keys = pose_tmp[0].keys()
|
| 621 |
+
for key in keys:
|
| 622 |
+
max_len = max([len(vid[key]) for vid in pose_tmp])
|
| 623 |
+
video_length = torch.LongTensor([len(vid[key]) for vid in pose_tmp])
|
| 624 |
+
|
| 625 |
+
padded_video = [torch.cat(
|
| 626 |
+
(
|
| 627 |
+
vid[key],
|
| 628 |
+
vid[key][-1][None].expand(max_len - len(vid[key]), -1, -1),
|
| 629 |
+
)
|
| 630 |
+
, dim=0)
|
| 631 |
+
for vid in pose_tmp]
|
| 632 |
+
|
| 633 |
+
img_batch = torch.stack(padded_video, 0)
|
| 634 |
+
|
| 635 |
+
src_input[key] = img_batch
|
| 636 |
+
if 'attention_mask' not in src_input.keys():
|
| 637 |
+
src_length_batch = video_length
|
| 638 |
+
|
| 639 |
+
mask_gen = []
|
| 640 |
+
for i in src_length_batch:
|
| 641 |
+
tmp = torch.ones([i]) + 7
|
| 642 |
+
mask_gen.append(tmp)
|
| 643 |
+
mask_gen = pad_sequence(mask_gen, padding_value=0, batch_first=True)
|
| 644 |
+
img_padding_mask = (mask_gen != 0).long()
|
| 645 |
+
src_input['attention_mask'] = img_padding_mask
|
| 646 |
+
|
| 647 |
+
src_input['name_batch'] = name_batch
|
| 648 |
+
src_input['src_length_batch'] = src_length_batch
|
| 649 |
+
|
| 650 |
+
if self.rgb_support:
|
| 651 |
+
support_rgb_dicts = {key: [] for key in batch[0][-1].keys()}
|
| 652 |
+
for _, _, _, _, support_rgb_dict in batch:
|
| 653 |
+
for key in support_rgb_dict.keys():
|
| 654 |
+
support_rgb_dicts[key].append(support_rgb_dict[key])
|
| 655 |
+
|
| 656 |
+
for part in ['left', 'right']:
|
| 657 |
+
index_key = f'{part}_sampled_indices'
|
| 658 |
+
skeletons_key = f'{part}_skeletons_norm'
|
| 659 |
+
rgb_key = f'{part}_hands'
|
| 660 |
+
len_key = f'{part}_rgb_len'
|
| 661 |
+
|
| 662 |
+
index_batch = torch.cat(support_rgb_dicts[index_key], 0)
|
| 663 |
+
skeletons_batch = torch.cat(support_rgb_dicts[skeletons_key], 0)
|
| 664 |
+
img_batch = torch.cat(support_rgb_dicts[rgb_key], 0)
|
| 665 |
+
|
| 666 |
+
src_input[index_key] = index_batch
|
| 667 |
+
src_input[skeletons_key] = skeletons_batch
|
| 668 |
+
src_input[rgb_key] = img_batch
|
| 669 |
+
src_input[len_key] = [len(index) for index in support_rgb_dicts[index_key]]
|
| 670 |
+
|
| 671 |
+
tgt_input = {}
|
| 672 |
+
tgt_input['gt_sentence'] = tgt_batch
|
| 673 |
+
tgt_input['gt_gloss'] = gloss_batch
|
| 674 |
+
|
| 675 |
+
return src_input, tgt_input
|
| 676 |
+
|
| 677 |
+
import gzip
|
| 678 |
+
def load_dataset_file(filename):
|
| 679 |
+
with gzip.open(filename, "rb") as f:
|
| 680 |
+
loaded_object = pickle.load(f)
|
| 681 |
+
return loaded_object
|
| 682 |
+
|
| 683 |
+
class S2T_Dataset(Base_Dataset):
|
| 684 |
+
def __init__(self, path, args, phase):
|
| 685 |
+
super(S2T_Dataset, self).__init__()
|
| 686 |
+
self.args = args
|
| 687 |
+
self.rgb_support = self.args.rgb_support
|
| 688 |
+
self.max_length = args.max_length
|
| 689 |
+
self.raw_data = load_dataset_file(path)
|
| 690 |
+
self.phase = phase
|
| 691 |
+
|
| 692 |
+
if self.args.dataset == "CSL_Daily":
|
| 693 |
+
self.pose_dir = pose_dirs[args.dataset]
|
| 694 |
+
self.rgb_dir = rgb_dirs[args.dataset]
|
| 695 |
+
|
| 696 |
+
elif "WLASL" in self.args.dataset:
|
| 697 |
+
self.pose_dir = os.path.join(pose_dirs[args.dataset], phase)
|
| 698 |
+
self.rgb_dir = os.path.join(rgb_dirs[args.dataset], phase)
|
| 699 |
+
|
| 700 |
+
elif self.args.dataset == "Isharah":
|
| 701 |
+
self.pose_dir = pose_dirs[args.dataset] # pose only
|
| 702 |
+
self.rgb_dir = rgb_dirs[args.dataset] # ""
|
| 703 |
+
|
| 704 |
+
else:
|
| 705 |
+
raise NotImplementedError(f"dataset {self.args.dataset} not supported")
|
| 706 |
+
|
| 707 |
+
self.list = list(self.raw_data.keys())
|
| 708 |
+
|
| 709 |
+
self.data_transform = transforms.Compose([
|
| 710 |
+
transforms.ToTensor(),
|
| 711 |
+
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
| 712 |
+
])
|
| 713 |
+
|
| 714 |
+
def __len__(self):
|
| 715 |
+
return len(self.list)
|
| 716 |
+
|
| 717 |
+
def __getitem__(self, index):
|
| 718 |
+
key = self.list[index]
|
| 719 |
+
sample = self.raw_data[key]
|
| 720 |
+
|
| 721 |
+
text = sample['text']
|
| 722 |
+
if "gloss" in sample.keys():
|
| 723 |
+
gloss = " ".join(sample['gloss'])
|
| 724 |
+
else:
|
| 725 |
+
gloss = ''
|
| 726 |
+
|
| 727 |
+
name_sample = sample['name']
|
| 728 |
+
pose_sample, support_rgb_dict = self.load_pose(sample['video_path'])
|
| 729 |
+
|
| 730 |
+
return name_sample, pose_sample, text, gloss, support_rgb_dict
|
| 731 |
+
|
| 732 |
+
def load_pose(self, path):
|
| 733 |
+
pose = pickle.load(open(os.path.join(self.pose_dir, path.replace(".mp4", '.pkl')), 'rb'))
|
| 734 |
+
|
| 735 |
+
if 'start' in pose.keys():
|
| 736 |
+
assert pose['start'] < pose['end']
|
| 737 |
+
duration = pose['end'] - pose['start']
|
| 738 |
+
start = pose['start']
|
| 739 |
+
else:
|
| 740 |
+
duration = len(pose['scores'])
|
| 741 |
+
start = 0
|
| 742 |
+
|
| 743 |
+
tmp = select_frame_indices(duration, self.max_length, self.phase)
|
| 744 |
+
|
| 745 |
+
tmp = np.array(tmp) + start
|
| 746 |
+
|
| 747 |
+
skeletons = pose['keypoints']
|
| 748 |
+
confs = pose['scores']
|
| 749 |
+
skeletons_tmp = []
|
| 750 |
+
confs_tmp = []
|
| 751 |
+
for index in tmp:
|
| 752 |
+
skeletons_tmp.append(skeletons[index])
|
| 753 |
+
confs_tmp.append(confs[index])
|
| 754 |
+
|
| 755 |
+
skeletons = skeletons_tmp
|
| 756 |
+
confs = confs_tmp
|
| 757 |
+
|
| 758 |
+
kps_with_scores = load_part_kp(skeletons, confs, force_ok=True)
|
| 759 |
+
|
| 760 |
+
support_rgb_dict = {}
|
| 761 |
+
if self.rgb_support:
|
| 762 |
+
full_path = os.path.join(self.rgb_dir, path)
|
| 763 |
+
support_rgb_dict = load_support_rgb_dict(tmp, skeletons, confs, full_path, self.data_transform)
|
| 764 |
+
|
| 765 |
+
return kps_with_scores, support_rgb_dict
|
| 766 |
+
|
| 767 |
+
def __str__(self):
|
| 768 |
+
return f'#total {len(self)}'
|
| 769 |
+
|
| 770 |
+
|
| 771 |
+
class S2T_Dataset_YTASL(Base_Dataset):
|
| 772 |
+
def __init__(self, path, args, phase):
|
| 773 |
+
super(S2T_Dataset_YTASL, self).__init__()
|
| 774 |
+
self.args = args
|
| 775 |
+
self.max_length = args.max_length
|
| 776 |
+
self.phase = phase
|
| 777 |
+
self.annotation = load_json(path)
|
| 778 |
+
self.rgb_support = self.args.rgb_support
|
| 779 |
+
self.normalization = self.args.normalization
|
| 780 |
+
self.layout = args.layout
|
| 781 |
+
|
| 782 |
+
self.pose_dir = pose_dirs[args.dataset]
|
| 783 |
+
self.rgb_dir = rgb_dirs[args.dataset]
|
| 784 |
+
|
| 785 |
+
self.list_data = [] # [(video_id, clip_id), ...]
|
| 786 |
+
self.clip_order_to_int = {}
|
| 787 |
+
self.clip_order_from_int = {}
|
| 788 |
+
|
| 789 |
+
for video_id in self.annotation.keys():
|
| 790 |
+
co = self.annotation[video_id]['clip_order']
|
| 791 |
+
self.clip_order_from_int[video_id] = dict(zip(range(len(co)), co))
|
| 792 |
+
self.clip_order_to_int[video_id] = dict(zip(co, range(len(co))))
|
| 793 |
+
|
| 794 |
+
for video_id, clip_dict in self.annotation.items():
|
| 795 |
+
for clip_name in clip_dict['clip_order']:
|
| 796 |
+
self.list_data.append((video_id, self.clip_order_to_int[video_id][clip_name]))
|
| 797 |
+
|
| 798 |
+
available_clip_names = {
|
| 799 |
+
pathlib.Path(clip).stem # remove suffix
|
| 800 |
+
for clip in os.listdir(self.pose_dir)
|
| 801 |
+
if clip.endswith(".json")
|
| 802 |
+
}
|
| 803 |
+
video_clips = set()
|
| 804 |
+
for video_id, clip_dict in self.annotation.items():
|
| 805 |
+
for clip_name in clip_dict['clip_order']:
|
| 806 |
+
if clip_name in available_clip_names:
|
| 807 |
+
video_clips.add((video_id, self.clip_order_to_int[video_id][clip_name]))
|
| 808 |
+
|
| 809 |
+
self.remove_missing_annotation(video_clips) # Remove data in annotations that are missing in h5 file
|
| 810 |
+
|
| 811 |
+
def remove_missing_annotation(self, h5_video_clip):
|
| 812 |
+
annotations_to_delete = set(self.list_data) - h5_video_clip
|
| 813 |
+
for a in annotations_to_delete:
|
| 814 |
+
self.list_data.remove(a)
|
| 815 |
+
|
| 816 |
+
def __getitem__(self, index):
|
| 817 |
+
video_id, clip_id = self.list_data[index]
|
| 818 |
+
clip_name = self.clip_order_from_int[video_id][clip_id]
|
| 819 |
+
|
| 820 |
+
# Get translation
|
| 821 |
+
clip_dict = self.annotation[video_id][clip_name]
|
| 822 |
+
text = clip_dict['translation']
|
| 823 |
+
|
| 824 |
+
# Get the pose features
|
| 825 |
+
pose_sample = self.load_pose(clip_name)
|
| 826 |
+
|
| 827 |
+
# TODO: rgb support
|
| 828 |
+
video_path = ""
|
| 829 |
+
support_rgb_dict = {}
|
| 830 |
+
|
| 831 |
+
# sample = {"name": clip_name,
|
| 832 |
+
# "video_path": video_path,
|
| 833 |
+
# "pose_features": pose_features,
|
| 834 |
+
# "text": translation}
|
| 835 |
+
# Crop long sequences to desired max length. Random sample
|
| 836 |
+
|
| 837 |
+
# skeletons = pose['keypoints']
|
| 838 |
+
# confs = pose['scores']
|
| 839 |
+
# skeletons_tmp = []
|
| 840 |
+
# confs_tmp = []
|
| 841 |
+
# for index in tmp:
|
| 842 |
+
# skeletons_tmp.append(skeletons[index])
|
| 843 |
+
# confs_tmp.append(confs[index])
|
| 844 |
+
#
|
| 845 |
+
# skeletons = skeletons_tmp
|
| 846 |
+
# confs = confs_tmp
|
| 847 |
+
|
| 848 |
+
# confs = [np.ones(int(pose_features.shape[1]/2)) for _ in range(pose_features.shape[0])]
|
| 849 |
+
# confs = [np.ones(pose_features.shape[0])] * pose_features.shape[1]
|
| 850 |
+
# skeletons = [] # List of ndarrays (133,2) - full keypoints
|
| 851 |
+
# kps_with_scores = load_part_kp(skeletons, confs, force_ok=True)
|
| 852 |
+
|
| 853 |
+
name_sample = clip_name
|
| 854 |
+
gloss = ''
|
| 855 |
+
|
| 856 |
+
return name_sample, pose_sample, text, gloss, support_rgb_dict
|
| 857 |
+
|
| 858 |
+
def load_pose(self, clip_name):
|
| 859 |
+
path = os.path.join(self.pose_dir, f"{clip_name}.json")
|
| 860 |
+
pose_data = load_json(path)
|
| 861 |
+
pose = pose_data['cropped_keypoints']
|
| 862 |
+
|
| 863 |
+
duration = len(pose)
|
| 864 |
+
start = 0
|
| 865 |
+
|
| 866 |
+
tmp = select_frame_indices(duration, self.max_length, self.phase)
|
| 867 |
+
tmp = np.array(tmp) + start
|
| 868 |
+
skeletons = [pose[i] for i in tmp]
|
| 869 |
+
|
| 870 |
+
confs = []
|
| 871 |
+
for i, skeleton in enumerate(skeletons):
|
| 872 |
+
conf = {}
|
| 873 |
+
for group_name, expected_size in YTASL_GROUP_SIZES.items():
|
| 874 |
+
_fill_missing_landmarks(
|
| 875 |
+
skeleton=skeleton,
|
| 876 |
+
conf=conf,
|
| 877 |
+
group_name=group_name,
|
| 878 |
+
expected_size=expected_size,
|
| 879 |
+
clip_name=clip_name,
|
| 880 |
+
frame_idx=i,
|
| 881 |
+
error_group_label=YTASL_GROUP_ERROR_LABELS[group_name],
|
| 882 |
+
include_size_details=False,
|
| 883 |
+
strict_key_access=True,
|
| 884 |
+
)
|
| 885 |
+
|
| 886 |
+
confs.append(conf)
|
| 887 |
+
|
| 888 |
+
kps_with_scores = load_part_kp_YTASL(skeletons, confs, self.normalization, self.layout)
|
| 889 |
+
return kps_with_scores
|
| 890 |
+
|
| 891 |
+
def __len__(self):
|
| 892 |
+
return len(self.list_data)
|
| 893 |
+
|
| 894 |
+
def __str__(self):
|
| 895 |
+
return f'#total {len(self)}'
|
| 896 |
+
|
| 897 |
+
|
| 898 |
+
class S2T_Dataset_Isharah(S2T_Dataset_YTASL):
|
| 899 |
+
def __init__(self, path, args, phase):
|
| 900 |
+
super(S2T_Dataset_Isharah, self).__init__(path=path, args=args, phase=phase)
|
| 901 |
+
|
| 902 |
+
def load_pose(self, clip_name):
|
| 903 |
+
path = os.path.join(self.pose_dir, f"{clip_name}.json")
|
| 904 |
+
pose_data = load_json(path)
|
| 905 |
+
pose = pose_data['cropped_keypoints']
|
| 906 |
+
|
| 907 |
+
duration = len(pose)
|
| 908 |
+
tmp = select_frame_indices(duration, self.max_length, self.phase)
|
| 909 |
+
tmp = np.array(tmp)
|
| 910 |
+
skeletons = [pose[i] for i in tmp]
|
| 911 |
+
|
| 912 |
+
confs = []
|
| 913 |
+
for i, skeleton in enumerate(skeletons):
|
| 914 |
+
conf = {}
|
| 915 |
+
for group_name, expected_size in ISHARAH_GROUP_SIZES.items():
|
| 916 |
+
_fill_missing_landmarks(
|
| 917 |
+
skeleton=skeleton,
|
| 918 |
+
conf=conf,
|
| 919 |
+
group_name=group_name,
|
| 920 |
+
expected_size=expected_size,
|
| 921 |
+
clip_name=clip_name,
|
| 922 |
+
frame_idx=i,
|
| 923 |
+
)
|
| 924 |
+
confs.append(conf)
|
| 925 |
+
|
| 926 |
+
kps_with_scores = load_part_kp_Isharah(skeletons, confs, self.normalization, self.layout)
|
| 927 |
+
return kps_with_scores
|
| 928 |
+
|
| 929 |
+
|
| 930 |
+
# class S2T_Dataset_YTASL_h5(Base_Dataset):
|
| 931 |
+
# def __init__(self, path, args, phase):
|
| 932 |
+
# super(S2T_Dataset_YTASL_h5, self).__init__()
|
| 933 |
+
# self.args = args
|
| 934 |
+
# self.max_length = args.max_length
|
| 935 |
+
# self.phase = phase
|
| 936 |
+
# self.annotation = load_json(path)
|
| 937 |
+
# self.rgb_support = self.args.rgb_support
|
| 938 |
+
#
|
| 939 |
+
# # Load poses
|
| 940 |
+
# self.list_data = [] # [(video_id, clip_id), ...]
|
| 941 |
+
# self.h5_data = {}
|
| 942 |
+
# self.h5shard = defaultdict(lambda: defaultdict(dict))
|
| 943 |
+
# self.clip_order_to_int = {}
|
| 944 |
+
# self.clip_order_from_int = {}
|
| 945 |
+
#
|
| 946 |
+
# for video_id in self.annotation.keys():
|
| 947 |
+
# co = self.annotation[video_id]['clip_order']
|
| 948 |
+
# self.clip_order_from_int[video_id] = dict(zip(range(len(co)), co))
|
| 949 |
+
# self.clip_order_to_int[video_id] = dict(zip(co, range(len(co))))
|
| 950 |
+
#
|
| 951 |
+
# for video_id, clip_dict in self.annotation.items():
|
| 952 |
+
# for clip_name in clip_dict:
|
| 953 |
+
# if clip_name != "clip_order":
|
| 954 |
+
# self.list_data.append((video_id, self.clip_order_to_int[video_id][clip_name]))
|
| 955 |
+
#
|
| 956 |
+
# self.vf_path = os.path.join(pose_dirs["YTASL"], "YouTubeASL.keypoints.{}.json".format(phase))
|
| 957 |
+
#
|
| 958 |
+
# h5_video_clip = self.read_multih5_json(self.vf_path)
|
| 959 |
+
# self.remove_missing_annotation(h5_video_clip) # Remove data in annotations that are missing in h5 file
|
| 960 |
+
#
|
| 961 |
+
# def read_multih5_json(self, json_filename):
|
| 962 |
+
# """Helper function for reading json specifications of multiple H5 files for visual features"""
|
| 963 |
+
# h5_video_clip = set()
|
| 964 |
+
# with open(json_filename, 'r') as F:
|
| 965 |
+
# self.h5shard = json.load(F)
|
| 966 |
+
# self.h5_data = {}
|
| 967 |
+
# print(f"Pose {self.phase} data are loaded from: ")
|
| 968 |
+
# for k in set(self.h5shard.values()):
|
| 969 |
+
# h5file = json_filename.replace('metadata_', '').replace('.json', ".%s.h5" % k)
|
| 970 |
+
# print("--" + h5file) # ,k,json_filename,data_dir)
|
| 971 |
+
# self.h5_data[k] = h5py.File(h5file, 'r')
|
| 972 |
+
#
|
| 973 |
+
# for vi in self.h5_data[k].keys():
|
| 974 |
+
# for ci in self.h5_data[k][vi].keys():
|
| 975 |
+
# if vi in self.clip_order_to_int:
|
| 976 |
+
# if ci in self.clip_order_to_int[vi]:
|
| 977 |
+
# clip_id = self.clip_order_to_int[vi][ci]
|
| 978 |
+
# h5_video_clip.add((vi, clip_id))
|
| 979 |
+
# return h5_video_clip
|
| 980 |
+
#
|
| 981 |
+
# def remove_missing_annotation(self, h5_video_clip):
|
| 982 |
+
# annotations_to_delete = set(self.list_data) - h5_video_clip
|
| 983 |
+
# for a in annotations_to_delete:
|
| 984 |
+
# self.list_data.remove(a)
|
| 985 |
+
#
|
| 986 |
+
# def __getitem__(self, index):
|
| 987 |
+
# video_id, clip_id = self.list_data[index]
|
| 988 |
+
# clip_name = self.clip_order_from_int[video_id][clip_id]
|
| 989 |
+
#
|
| 990 |
+
# # Get the pose features
|
| 991 |
+
# shard = self.h5shard[video_id]
|
| 992 |
+
# pose_features = torch.tensor(np.array(self.h5_data[shard][video_id][clip_name]))
|
| 993 |
+
#
|
| 994 |
+
# # TODO: rgb support
|
| 995 |
+
# video_path = ""
|
| 996 |
+
#
|
| 997 |
+
# # Get translation
|
| 998 |
+
# clip_dict = self.annotation[video_id][clip_name]
|
| 999 |
+
# translation = clip_dict['translation']
|
| 1000 |
+
#
|
| 1001 |
+
# # sample = {"name": clip_name,
|
| 1002 |
+
# # "video_path": video_path,
|
| 1003 |
+
# # "pose_features": pose_features,
|
| 1004 |
+
# # "text": translation}
|
| 1005 |
+
# #
|
| 1006 |
+
# # # Crop long sequences to desired max length. Random sample
|
| 1007 |
+
# # duration = len(pose_features) # TODO: works?
|
| 1008 |
+
# # if duration > self.max_length:
|
| 1009 |
+
# # tmp = sorted(random.sample(range(duration), k=self.max_length))
|
| 1010 |
+
# # else:
|
| 1011 |
+
# # tmp = list(range(duration))
|
| 1012 |
+
# #
|
| 1013 |
+
# # tmp = np.array(tmp)
|
| 1014 |
+
#
|
| 1015 |
+
# # skeletons = pose['keypoints']
|
| 1016 |
+
# # confs = pose['scores']
|
| 1017 |
+
# # skeletons_tmp = []
|
| 1018 |
+
# # confs_tmp = []
|
| 1019 |
+
# # for index in tmp:
|
| 1020 |
+
# # skeletons_tmp.append(skeletons[index])
|
| 1021 |
+
# # confs_tmp.append(confs[index])
|
| 1022 |
+
# #
|
| 1023 |
+
# # skeletons = skeletons_tmp
|
| 1024 |
+
# # confs = confs_tmp
|
| 1025 |
+
#
|
| 1026 |
+
# # confs = [np.ones(int(pose_features.shape[1]/2)) for _ in range(pose_features.shape[0])]
|
| 1027 |
+
# # confs = [np.ones(pose_features.shape[0])] * pose_features.shape[1]
|
| 1028 |
+
# # skeletons = [] # List of ndarrays (133,2) - full keypoints
|
| 1029 |
+
# # kps_with_scores = load_part_kp(skeletons, confs, force_ok=True)
|
| 1030 |
+
#
|
| 1031 |
+
# # decoded = self.tokenizer(
|
| 1032 |
+
# # translation,
|
| 1033 |
+
# # max_length=self.max_token_length,
|
| 1034 |
+
# # padding="max_length",
|
| 1035 |
+
# # truncation=True,
|
| 1036 |
+
# # return_tensors="pt",
|
| 1037 |
+
# # )
|
| 1038 |
+
# # labels = decoded.input_ids
|
| 1039 |
+
#
|
| 1040 |
+
# # Skip frames for the keypoints
|
| 1041 |
+
# # if self.skip_frames:
|
| 1042 |
+
# # if type(self.skip_frames) == bool:
|
| 1043 |
+
# # for input_type in INPUT_TYPES:
|
| 1044 |
+
# # if visual_features[input_type] is not None:
|
| 1045 |
+
# # visual_features[input_type] = visual_features[input_type][::2]
|
| 1046 |
+
# # elif type(self.skip_frames) == int:
|
| 1047 |
+
# # for input_type in INPUT_TYPES:
|
| 1048 |
+
# # if visual_features[input_type] is not None:
|
| 1049 |
+
# # visual_features[input_type] = visual_features[input_type][::self.skip_frames]
|
| 1050 |
+
# #
|
| 1051 |
+
# # # Trim the keypoints to the max sequence length
|
| 1052 |
+
# # if self.max_sequence_length:
|
| 1053 |
+
# # for input_type in INPUT_TYPES:
|
| 1054 |
+
# # if visual_features[input_type] is not None:
|
| 1055 |
+
# # visual_features[input_type] = visual_features[input_type][: self.max_sequence_length]
|
| 1056 |
+
# # seq_len = len(visual_features[input_type])
|
| 1057 |
+
# #
|
| 1058 |
+
# # assert seq_len, "No modality provided or clip has no length!"
|
| 1059 |
+
# # attention_mask = torch.ones(seq_len)
|
| 1060 |
+
# #
|
| 1061 |
+
# # return {
|
| 1062 |
+
# # "sign_inputs": {'pose': visual_features['pose'],
|
| 1063 |
+
# # 'mae': visual_features['mae'],
|
| 1064 |
+
# # 'dino': visual_features['dino'],
|
| 1065 |
+
# # 'sign2vec': visual_features['sign2vec']},
|
| 1066 |
+
# # "sentence": translation,
|
| 1067 |
+
# # "labels": labels,
|
| 1068 |
+
# # "attention_mask": attention_mask,
|
| 1069 |
+
# # }
|
| 1070 |
+
# return kps
|
| 1071 |
+
#
|
| 1072 |
+
# def __len__(self):
|
| 1073 |
+
# return len(self.list_data)
|
| 1074 |
+
#
|
| 1075 |
+
# def __str__(self):
|
| 1076 |
+
# return f'#total {len(self)}'
|
| 1077 |
+
|
| 1078 |
+
|
| 1079 |
+
class S2T_Dataset_news(Base_Dataset):
|
| 1080 |
+
def __init__(self, path, args, phase):
|
| 1081 |
+
super(S2T_Dataset_news, self).__init__()
|
| 1082 |
+
self.args = args
|
| 1083 |
+
self.rgb_support = self.args.rgb_support
|
| 1084 |
+
self.phase = phase
|
| 1085 |
+
self.max_length = args.max_length
|
| 1086 |
+
|
| 1087 |
+
path = pathlib.Path(path)
|
| 1088 |
+
|
| 1089 |
+
with path.open(encoding='utf-8') as f:
|
| 1090 |
+
self.annotation = json.load(f)
|
| 1091 |
+
|
| 1092 |
+
if self.args.dataset == "CSL_News":
|
| 1093 |
+
self.pose_dir = pose_dirs[args.dataset]
|
| 1094 |
+
self.rgb_dir = rgb_dirs[args.dataset]
|
| 1095 |
+
|
| 1096 |
+
else:
|
| 1097 |
+
raise NotImplementedError
|
| 1098 |
+
sum_sample = len(self.annotation)
|
| 1099 |
+
self.data_transform = transforms.Compose([
|
| 1100 |
+
transforms.ToTensor(),
|
| 1101 |
+
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
| 1102 |
+
])
|
| 1103 |
+
|
| 1104 |
+
if phase == 'train':
|
| 1105 |
+
self.start_idx = int(sum_sample * 0.0)
|
| 1106 |
+
self.end_idx = int(sum_sample * 0.99)
|
| 1107 |
+
else:
|
| 1108 |
+
self.start_idx = int(sum_sample * 0.99)
|
| 1109 |
+
self.end_idx = int(sum_sample)
|
| 1110 |
+
|
| 1111 |
+
def __len__(self):
|
| 1112 |
+
return self.end_idx - self.start_idx
|
| 1113 |
+
|
| 1114 |
+
def __getitem__(self, index):
|
| 1115 |
+
num_retries = 10
|
| 1116 |
+
|
| 1117 |
+
# skip some invalid video sample
|
| 1118 |
+
for _ in range(num_retries):
|
| 1119 |
+
sample = self.annotation[self.start_idx:self.end_idx][index]
|
| 1120 |
+
|
| 1121 |
+
text = sample['text']
|
| 1122 |
+
name_sample = sample['video']
|
| 1123 |
+
|
| 1124 |
+
try:
|
| 1125 |
+
pose_sample, support_rgb_dict = self.load_pose(sample['pose'], sample['video'])
|
| 1126 |
+
|
| 1127 |
+
except:
|
| 1128 |
+
import traceback
|
| 1129 |
+
|
| 1130 |
+
traceback.print_exc()
|
| 1131 |
+
print(f"Failed to load examples with video: {name_sample}. "
|
| 1132 |
+
f"Will randomly sample an example as a replacement.")
|
| 1133 |
+
index = random.randint(0, len(self) - 1)
|
| 1134 |
+
continue
|
| 1135 |
+
|
| 1136 |
+
break
|
| 1137 |
+
|
| 1138 |
+
else:
|
| 1139 |
+
raise RuntimeError(f"Failed to fetch video after {num_retries} retries.")
|
| 1140 |
+
|
| 1141 |
+
return name_sample, pose_sample, text, _, support_rgb_dict
|
| 1142 |
+
|
| 1143 |
+
def load_pose(self, pose_name, rgb_name):
|
| 1144 |
+
pose = pickle.load(open(os.path.join(self.pose_dir, pose_name), 'rb'))
|
| 1145 |
+
full_path = os.path.join(self.rgb_dir, rgb_name)
|
| 1146 |
+
|
| 1147 |
+
duration = len(pose['scores'])
|
| 1148 |
+
|
| 1149 |
+
tmp = select_frame_indices(duration, self.max_length, self.phase)
|
| 1150 |
+
|
| 1151 |
+
tmp = np.array(tmp)
|
| 1152 |
+
|
| 1153 |
+
# dict_keys(['keypoints', 'scores'])
|
| 1154 |
+
# keypoints (1, 133, 2)
|
| 1155 |
+
# scores (1, 133)
|
| 1156 |
+
|
| 1157 |
+
skeletons = pose['keypoints']
|
| 1158 |
+
confs = pose['scores']
|
| 1159 |
+
skeletons_tmp = []
|
| 1160 |
+
confs_tmp = []
|
| 1161 |
+
|
| 1162 |
+
for index in tmp:
|
| 1163 |
+
skeletons_tmp.append(skeletons[index])
|
| 1164 |
+
confs_tmp.append(confs[index])
|
| 1165 |
+
|
| 1166 |
+
skeletons = skeletons_tmp
|
| 1167 |
+
confs = confs_tmp
|
| 1168 |
+
|
| 1169 |
+
kps_with_scores = load_part_kp(skeletons, confs)
|
| 1170 |
+
|
| 1171 |
+
support_rgb_dict = {}
|
| 1172 |
+
if self.rgb_support:
|
| 1173 |
+
support_rgb_dict = load_support_rgb_dict(tmp, skeletons, confs, full_path, self.data_transform)
|
| 1174 |
+
|
| 1175 |
+
return kps_with_scores, support_rgb_dict
|
| 1176 |
+
|
| 1177 |
+
def __str__(self):
|
| 1178 |
+
return f'#total {len(self)}'
|
Uni_Sign/deformable_attention_2d.py
ADDED
|
@@ -0,0 +1,311 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Clone from https://github.com/lucidrains/deformable-attention
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
from torch import nn, einsum
|
| 5 |
+
|
| 6 |
+
from einops import rearrange, repeat
|
| 7 |
+
from einops.layers.torch import Rearrange
|
| 8 |
+
# helper functions
|
| 9 |
+
from typing import Any, Optional, Tuple, Type
|
| 10 |
+
import numpy as np
|
| 11 |
+
|
| 12 |
+
def exists(val):
|
| 13 |
+
return val is not None
|
| 14 |
+
|
| 15 |
+
def default(val, d):
|
| 16 |
+
return val if exists(val) else d
|
| 17 |
+
|
| 18 |
+
def divisible_by(numer, denom):
|
| 19 |
+
return (numer % denom) == 0
|
| 20 |
+
|
| 21 |
+
# tensor helpers
|
| 22 |
+
|
| 23 |
+
def create_grid_like(t, dim = 0):
|
| 24 |
+
h, w, device = *t.shape[-2:], t.device
|
| 25 |
+
|
| 26 |
+
grid = torch.stack(torch.meshgrid(
|
| 27 |
+
torch.arange(w, device = device),
|
| 28 |
+
torch.arange(h, device = device),
|
| 29 |
+
indexing = 'xy'), dim = dim)
|
| 30 |
+
|
| 31 |
+
grid.requires_grad = False
|
| 32 |
+
grid = grid.type_as(t)
|
| 33 |
+
return grid
|
| 34 |
+
|
| 35 |
+
def normalize_grid(grid, dim = 1, out_dim = -1):
|
| 36 |
+
# normalizes a grid to range from -1 to 1
|
| 37 |
+
h, w = grid.shape[-2:]
|
| 38 |
+
grid_h, grid_w = grid.unbind(dim = dim)
|
| 39 |
+
|
| 40 |
+
grid_h = 2.0 * grid_h / max(h - 1, 1) - 1.0
|
| 41 |
+
grid_w = 2.0 * grid_w / max(w - 1, 1) - 1.0
|
| 42 |
+
|
| 43 |
+
return torch.stack((grid_h, grid_w), dim = out_dim)
|
| 44 |
+
|
| 45 |
+
def reshape_grid_1d(grid, dim = 1, out_dim = -1):
|
| 46 |
+
# normalizes a grid to range from -1 to 1
|
| 47 |
+
n = grid.shape[-1]
|
| 48 |
+
grid_h, grid_w = grid.unbind(dim = dim)
|
| 49 |
+
|
| 50 |
+
return torch.stack((grid_h, grid_w), dim = out_dim)
|
| 51 |
+
|
| 52 |
+
class Scale(nn.Module):
|
| 53 |
+
def __init__(self, scale):
|
| 54 |
+
super().__init__()
|
| 55 |
+
self.scale = scale
|
| 56 |
+
|
| 57 |
+
def forward(self, x):
|
| 58 |
+
return x * self.scale
|
| 59 |
+
|
| 60 |
+
# continuous positional bias from SwinV2
|
| 61 |
+
|
| 62 |
+
class CPB(nn.Module):
|
| 63 |
+
""" https://arxiv.org/abs/2111.09883v1 """
|
| 64 |
+
|
| 65 |
+
def __init__(self, dim, *, heads, offset_groups, depth):
|
| 66 |
+
super().__init__()
|
| 67 |
+
self.heads = heads
|
| 68 |
+
self.offset_groups = offset_groups
|
| 69 |
+
|
| 70 |
+
self.mlp = nn.ModuleList([])
|
| 71 |
+
|
| 72 |
+
self.mlp.append(nn.Sequential(
|
| 73 |
+
nn.Linear(2, dim),
|
| 74 |
+
nn.ReLU()
|
| 75 |
+
))
|
| 76 |
+
|
| 77 |
+
for _ in range(depth - 1):
|
| 78 |
+
self.mlp.append(nn.Sequential(
|
| 79 |
+
nn.Linear(dim, dim),
|
| 80 |
+
nn.ReLU()
|
| 81 |
+
))
|
| 82 |
+
|
| 83 |
+
self.mlp.append(nn.Linear(dim, heads // offset_groups))
|
| 84 |
+
|
| 85 |
+
def forward(self, grid_q, grid_kv):
|
| 86 |
+
device, dtype = grid_q.device, grid_kv.dtype
|
| 87 |
+
|
| 88 |
+
if grid_q.ndim == 3:
|
| 89 |
+
raise AttributeError
|
| 90 |
+
grid_q = rearrange(grid_q, 'h w c -> 1 (h w) c')
|
| 91 |
+
elif grid_q.ndim == 4:
|
| 92 |
+
grid_q = rearrange(grid_q, 'b h w c -> b (h w) c')
|
| 93 |
+
else:
|
| 94 |
+
raise AttributeError
|
| 95 |
+
grid_kv = rearrange(grid_kv, 'b h w c -> b (h w) c')
|
| 96 |
+
|
| 97 |
+
pos = rearrange(grid_q, 'b i c -> b i 1 c') - rearrange(grid_kv, 'b j c -> b 1 j c')
|
| 98 |
+
bias = torch.sign(pos) * torch.log(pos.abs() + 1) # log of distance is sign(rel_pos) * log(abs(rel_pos) + 1)
|
| 99 |
+
|
| 100 |
+
for layer in self.mlp:
|
| 101 |
+
bias = layer(bias)
|
| 102 |
+
|
| 103 |
+
bias = rearrange(bias, '(b g) i j o -> b (g o) i j', g = self.offset_groups)
|
| 104 |
+
|
| 105 |
+
return bias
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
class PositionEmbeddingRandom(nn.Module):
|
| 109 |
+
"""
|
| 110 |
+
Positional encoding using random spatial frequencies.
|
| 111 |
+
"""
|
| 112 |
+
|
| 113 |
+
def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
|
| 114 |
+
super().__init__()
|
| 115 |
+
if scale is None or scale <= 0.0:
|
| 116 |
+
scale = 1.0
|
| 117 |
+
self.register_buffer(
|
| 118 |
+
"positional_encoding_gaussian_matrix",
|
| 119 |
+
scale * torch.randn((2, num_pos_feats)),
|
| 120 |
+
)
|
| 121 |
+
|
| 122 |
+
def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
|
| 123 |
+
"""Positionally encode points that are normalized to [0,1]."""
|
| 124 |
+
# assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
|
| 125 |
+
coords = 2 * coords - 1
|
| 126 |
+
coords = coords @ self.positional_encoding_gaussian_matrix
|
| 127 |
+
coords = 2 * np.pi * coords
|
| 128 |
+
# outputs d_1 x ... x d_n x C shape
|
| 129 |
+
return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
|
| 130 |
+
|
| 131 |
+
def forward(self, size: Tuple[int, int]) -> torch.Tensor:
|
| 132 |
+
"""Generate positional encoding for a grid of the specified size."""
|
| 133 |
+
h, w = size
|
| 134 |
+
device: Any = self.positional_encoding_gaussian_matrix.device
|
| 135 |
+
grid = torch.ones((h, w), device=device, dtype=torch.float32)
|
| 136 |
+
y_embed = grid.cumsum(dim=0) - 0.5
|
| 137 |
+
x_embed = grid.cumsum(dim=1) - 0.5
|
| 138 |
+
y_embed = y_embed / h
|
| 139 |
+
x_embed = x_embed / w
|
| 140 |
+
|
| 141 |
+
pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
|
| 142 |
+
return pe.permute(2, 0, 1) # C x H x W
|
| 143 |
+
|
| 144 |
+
def forward_with_coords(
|
| 145 |
+
self, coords_input: torch.Tensor, image_size: Tuple[int, int]
|
| 146 |
+
) -> torch.Tensor:
|
| 147 |
+
"""Positionally encode points that are not normalized to [0,1]."""
|
| 148 |
+
coords = coords_input.clone()
|
| 149 |
+
coords[:, :, 0] = coords[:, :, 0] / image_size[1]
|
| 150 |
+
coords[:, :, 1] = coords[:, :, 1] / image_size[0]
|
| 151 |
+
return self._pe_encoding(coords.to(torch.float)) # B x N x C
|
| 152 |
+
|
| 153 |
+
def get_sinusoid_encoding_table(n_position, d_hid):
|
| 154 |
+
''' Sinusoid position encoding table '''
|
| 155 |
+
# TODO: make it with torch instead of numpy
|
| 156 |
+
def get_position_angle_vec(position):
|
| 157 |
+
return [position / np.power(10000, 2 * (hid_j // 2) / d_hid) for hid_j in range(d_hid)]
|
| 158 |
+
|
| 159 |
+
sinusoid_table = np.array([get_position_angle_vec(pos_i) for pos_i in range(n_position)])
|
| 160 |
+
sinusoid_table[:, 0::2] = np.sin(sinusoid_table[:, 0::2]) # dim 2i
|
| 161 |
+
sinusoid_table[:, 1::2] = np.cos(sinusoid_table[:, 1::2]) # dim 2i+1
|
| 162 |
+
|
| 163 |
+
return torch.tensor(sinusoid_table,dtype=torch.float, requires_grad=False).unsqueeze(0)
|
| 164 |
+
|
| 165 |
+
# main class
|
| 166 |
+
|
| 167 |
+
class DeformableAttention2D(nn.Module):
|
| 168 |
+
def __init__(
|
| 169 |
+
self,
|
| 170 |
+
*,
|
| 171 |
+
dim,
|
| 172 |
+
dim_head = 64,
|
| 173 |
+
heads = 8,
|
| 174 |
+
dropout = 0.,
|
| 175 |
+
downsample_factor = 1,
|
| 176 |
+
offset_scale = None,
|
| 177 |
+
offset_groups = None,
|
| 178 |
+
offset_kernel_size = 1,
|
| 179 |
+
group_queries = True,
|
| 180 |
+
group_key_values = True
|
| 181 |
+
):
|
| 182 |
+
super().__init__()
|
| 183 |
+
offset_scale = default(offset_scale, downsample_factor)
|
| 184 |
+
assert offset_kernel_size >= downsample_factor, 'offset kernel size must be greater than or equal to the downsample factor'
|
| 185 |
+
assert divisible_by(offset_kernel_size - downsample_factor, 2)
|
| 186 |
+
|
| 187 |
+
offset_groups = default(offset_groups, heads)
|
| 188 |
+
assert divisible_by(heads, offset_groups)
|
| 189 |
+
|
| 190 |
+
inner_dim = dim_head * heads
|
| 191 |
+
self.scale = dim_head ** -0.5
|
| 192 |
+
self.heads = heads
|
| 193 |
+
self.offset_groups = offset_groups
|
| 194 |
+
|
| 195 |
+
offset_dims = inner_dim // offset_groups
|
| 196 |
+
|
| 197 |
+
self.downsample_factor = downsample_factor
|
| 198 |
+
|
| 199 |
+
self.to_offsets = nn.Sequential(
|
| 200 |
+
nn.Conv1d(offset_dims, offset_dims, kernel_size=1, groups = offset_dims, stride = 1, padding = 0),
|
| 201 |
+
nn.GELU(),
|
| 202 |
+
nn.Conv1d(offset_dims, 2, 1, bias = False),
|
| 203 |
+
# Rearrange('b 1 n -> b n'),
|
| 204 |
+
nn.Tanh(),
|
| 205 |
+
Scale(4/6)
|
| 206 |
+
)
|
| 207 |
+
|
| 208 |
+
self.rel_pos_bias = CPB(dim // 4, offset_groups = offset_groups, heads = heads, depth = 2)
|
| 209 |
+
|
| 210 |
+
self.dropout = nn.Dropout(dropout)
|
| 211 |
+
self.to_q = nn.Conv1d(dim, inner_dim, 1, groups = offset_groups if group_queries else 1, bias = False)
|
| 212 |
+
self.to_k = nn.Conv2d(dim, inner_dim, 1, groups = offset_groups if group_key_values else 1, bias = False)
|
| 213 |
+
self.to_v = nn.Conv2d(dim, inner_dim, 1, groups = offset_groups if group_key_values else 1, bias = False)
|
| 214 |
+
self.to_out = nn.Conv1d(inner_dim, dim, 1)
|
| 215 |
+
|
| 216 |
+
self.cross_attn = nn.MultiheadAttention(embed_dim=dim, num_heads=heads, batch_first=True)
|
| 217 |
+
# for pose
|
| 218 |
+
self.pe_layer = PositionEmbeddingRandom(dim//2)
|
| 219 |
+
# for rgb
|
| 220 |
+
self.pos_embed = get_sinusoid_encoding_table(4*4, dim)
|
| 221 |
+
|
| 222 |
+
def forward(self, pose_feat, rgb_feat, pose_init, return_vgrid = False):
|
| 223 |
+
"""
|
| 224 |
+
b - batch
|
| 225 |
+
h - heads
|
| 226 |
+
x - height
|
| 227 |
+
y - width
|
| 228 |
+
d - dimension
|
| 229 |
+
g - offset groups
|
| 230 |
+
"""
|
| 231 |
+
|
| 232 |
+
heads, b, h, w, downsample_factor, device = self.heads, rgb_feat.shape[0], *rgb_feat.shape[-2:], self.downsample_factor, rgb_feat.device
|
| 233 |
+
|
| 234 |
+
# queries
|
| 235 |
+
# pose_feat: bt, c, n
|
| 236 |
+
pose_feat_cross = rearrange(pose_feat, 'b d n -> b n d')
|
| 237 |
+
rgb_feat_cross = rearrange(rgb_feat, 'b d h w -> b (h w) d')
|
| 238 |
+
pose_init_cross = rearrange(pose_init, 'b d n -> b n d')
|
| 239 |
+
|
| 240 |
+
point_embedding = self.pe_layer._pe_encoding(pose_init_cross.detach())
|
| 241 |
+
|
| 242 |
+
kv = rgb_feat_cross + self.pos_embed.expand(b, -1, -1).type_as(rgb_feat_cross).to(rgb_feat_cross.device).clone().detach()
|
| 243 |
+
|
| 244 |
+
pose_feat_cross, _ = self.cross_attn(pose_feat_cross + point_embedding,
|
| 245 |
+
kv, kv)
|
| 246 |
+
|
| 247 |
+
pose_feat_cross = rearrange(pose_feat_cross, 'b n d -> b d n')
|
| 248 |
+
|
| 249 |
+
q = self.to_q(pose_feat + pose_feat_cross)
|
| 250 |
+
|
| 251 |
+
# calculate offsets - offset MLP shared across all groups
|
| 252 |
+
|
| 253 |
+
group_1d = lambda t: rearrange(t, 'b (g d) n -> (b g) d n', g = self.offset_groups)
|
| 254 |
+
group_2d = lambda t: rearrange(t, 'b (g d) ... -> (b g) d ...', g = self.offset_groups)
|
| 255 |
+
|
| 256 |
+
grouped_queries = group_1d(q)
|
| 257 |
+
|
| 258 |
+
offsets = self.to_offsets(grouped_queries)
|
| 259 |
+
|
| 260 |
+
# calculate grid + offsets
|
| 261 |
+
|
| 262 |
+
# pose_init --> [0, 1]
|
| 263 |
+
grid = pose_init[:,None].repeat(1, self.offset_groups, 1, 1) * 2 - 1
|
| 264 |
+
# grid --> [-1, 1]
|
| 265 |
+
grid = grid.reshape(-1, *pose_init.shape[-2:])
|
| 266 |
+
|
| 267 |
+
vgrid = grid + offsets
|
| 268 |
+
vgrid_scaled = reshape_grid_1d(vgrid)[:,None]
|
| 269 |
+
# vgrid_scaled = normalize_grid(vgrid)
|
| 270 |
+
|
| 271 |
+
kv_feats = F.grid_sample(
|
| 272 |
+
group_2d(rgb_feat),
|
| 273 |
+
vgrid_scaled,
|
| 274 |
+
mode = 'bilinear', padding_mode = 'zeros', align_corners = False)
|
| 275 |
+
|
| 276 |
+
kv_feats = rearrange(kv_feats, '(b g) d ... -> b (g d) ...', b = b)
|
| 277 |
+
|
| 278 |
+
# derive key / values
|
| 279 |
+
k, v = self.to_k(kv_feats), self.to_v(kv_feats)
|
| 280 |
+
|
| 281 |
+
# scale queries
|
| 282 |
+
q = q * self.scale
|
| 283 |
+
|
| 284 |
+
# split out heads
|
| 285 |
+
q, k, v = map(lambda t: rearrange(t, 'b (h d) ... -> b h (...) d', h = heads), (q, k, v))
|
| 286 |
+
|
| 287 |
+
# query / key similarity
|
| 288 |
+
sim = einsum('b h i d, b h j d -> b h i j', q, k)
|
| 289 |
+
|
| 290 |
+
# relative positional bias
|
| 291 |
+
grid = reshape_grid_1d(grid)[:,None]
|
| 292 |
+
|
| 293 |
+
rel_pos_bias = self.rel_pos_bias(grid, vgrid_scaled)
|
| 294 |
+
sim = sim + rel_pos_bias
|
| 295 |
+
|
| 296 |
+
# numerical stability
|
| 297 |
+
sim = sim - sim.amax(dim = -1, keepdim = True).detach()
|
| 298 |
+
|
| 299 |
+
# attention
|
| 300 |
+
attn = sim.softmax(dim = -1)
|
| 301 |
+
attn = self.dropout(attn)
|
| 302 |
+
|
| 303 |
+
# aggregate and combine heads
|
| 304 |
+
out = einsum('b h i j, b h j d -> b h i d', attn, v)
|
| 305 |
+
out = rearrange(out, 'b h n d -> b (h d) n')
|
| 306 |
+
out = self.to_out(out)
|
| 307 |
+
|
| 308 |
+
if return_vgrid:
|
| 309 |
+
return out, vgrid
|
| 310 |
+
|
| 311 |
+
return out
|
Uni_Sign/models.py
ADDED
|
@@ -0,0 +1,415 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from torch import Tensor
|
| 2 |
+
import torch
|
| 3 |
+
from torch import nn
|
| 4 |
+
import torch.utils.checkpoint
|
| 5 |
+
import contextlib
|
| 6 |
+
import torchvision
|
| 7 |
+
from einops import rearrange
|
| 8 |
+
|
| 9 |
+
import math
|
| 10 |
+
from Uni_Sign.stgcn_layers import Graph, get_stgcn_chain
|
| 11 |
+
from Uni_Sign.deformable_attention_2d import DeformableAttention2D
|
| 12 |
+
from transformers import MT5ForConditionalGeneration, T5Tokenizer, MT5Config
|
| 13 |
+
import warnings
|
| 14 |
+
|
| 15 |
+
mt5_path = r"./Uni_Sign/unisign_model"
|
| 16 |
+
|
| 17 |
+
def _no_grad_trunc_normal_(tensor, mean, std, a, b):
|
| 18 |
+
# Cut & paste from PyTorch official master until it's in a few official releases - RW
|
| 19 |
+
# Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf
|
| 20 |
+
def norm_cdf(x):
|
| 21 |
+
# Computes standard normal cumulative distribution function
|
| 22 |
+
return (1. + math.erf(x / math.sqrt(2.))) / 2.
|
| 23 |
+
|
| 24 |
+
if (mean < a - 2 * std) or (mean > b + 2 * std):
|
| 25 |
+
warnings.warn("mean is more than 2 std from [a, b] in nn.init.trunc_normal_. "
|
| 26 |
+
"The distribution of values may be incorrect.",
|
| 27 |
+
stacklevel=2)
|
| 28 |
+
|
| 29 |
+
with torch.no_grad():
|
| 30 |
+
# Values are generated by using a truncated uniform distribution and
|
| 31 |
+
# then using the inverse CDF for the normal distribution.
|
| 32 |
+
# Get upper and lower cdf values
|
| 33 |
+
l = norm_cdf((a - mean) / std)
|
| 34 |
+
u = norm_cdf((b - mean) / std)
|
| 35 |
+
|
| 36 |
+
# Uniformly fill tensor with values from [l, u], then translate to
|
| 37 |
+
# [2l-1, 2u-1].
|
| 38 |
+
tensor.uniform_(2 * l - 1, 2 * u - 1)
|
| 39 |
+
|
| 40 |
+
# Use inverse cdf transform for normal distribution to get truncated
|
| 41 |
+
# standard normal
|
| 42 |
+
tensor.erfinv_()
|
| 43 |
+
|
| 44 |
+
# Transform to proper mean, std
|
| 45 |
+
tensor.mul_(std * math.sqrt(2.))
|
| 46 |
+
tensor.add_(mean)
|
| 47 |
+
|
| 48 |
+
# Clamp to ensure it's in the proper range
|
| 49 |
+
tensor.clamp_(min=a, max=b)
|
| 50 |
+
return tensor
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.):
|
| 54 |
+
# type: (Tensor, float, float, float, float) -> Tensor
|
| 55 |
+
r"""Fills the input Tensor with values drawn from a truncated
|
| 56 |
+
normal distribution. The values are effectively drawn from the
|
| 57 |
+
normal distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)`
|
| 58 |
+
with values outside :math:`[a, b]` redrawn until they are within
|
| 59 |
+
the bounds. The method used for generating the random values works
|
| 60 |
+
best when :math:`a \leq \text{mean} \leq b`.
|
| 61 |
+
Args:
|
| 62 |
+
tensor: an n-dimensional `torch.Tensor`
|
| 63 |
+
mean: the mean of the normal distribution
|
| 64 |
+
std: the standard deviation of the normal distribution
|
| 65 |
+
a: the minimum cutoff value
|
| 66 |
+
b: the maximum cutoff value
|
| 67 |
+
Examples:
|
| 68 |
+
>>> w = torch.empty(3, 5)
|
| 69 |
+
>>> nn.init.trunc_normal_(w)
|
| 70 |
+
"""
|
| 71 |
+
return _no_grad_trunc_normal_(tensor, mean, std, a, b)
|
| 72 |
+
|
| 73 |
+
class Uni_Sign(nn.Module):
|
| 74 |
+
def __init__(self, args):
|
| 75 |
+
super(Uni_Sign, self).__init__()
|
| 76 |
+
self.args = args
|
| 77 |
+
|
| 78 |
+
self.modes = ['body', 'left', 'right', 'face_all']
|
| 79 |
+
|
| 80 |
+
self.graph, A = {}, []
|
| 81 |
+
# project (x,y,score) to hidden dim
|
| 82 |
+
hidden_dim = args.hidden_dim
|
| 83 |
+
self.proj_linear = nn.ModuleDict()
|
| 84 |
+
for mode in self.modes:
|
| 85 |
+
graph_layout = f'{args.layout}_ytasl_{mode}' if self.args.dataset in ["YTASL", "Isharah"] else f'{args.layout}_{mode}'
|
| 86 |
+
self.graph[mode] = Graph(layout=graph_layout, strategy='distance', max_hop=1)
|
| 87 |
+
A.append(torch.tensor(self.graph[mode].A, dtype=torch.float32, requires_grad=False))
|
| 88 |
+
self.proj_linear[mode] = nn.Linear(3, 64)
|
| 89 |
+
|
| 90 |
+
self.gcn_modules = nn.ModuleDict()
|
| 91 |
+
self.fusion_gcn_modules = nn.ModuleDict()
|
| 92 |
+
spatial_kernel_size = A[0].size(0)
|
| 93 |
+
for index, mode in enumerate(self.modes):
|
| 94 |
+
self.gcn_modules[mode], final_dim = get_stgcn_chain(64, 'spatial', (1, spatial_kernel_size), A[index].clone(), adaptive=not self.args.no_adaptive_gcn)
|
| 95 |
+
self.fusion_gcn_modules[mode], _ = get_stgcn_chain(final_dim, 'temporal', (5, spatial_kernel_size), A[index].clone(), adaptive=not self.args.no_adaptive_gcn)
|
| 96 |
+
|
| 97 |
+
self.gcn_modules['left'] = self.gcn_modules['right']
|
| 98 |
+
self.fusion_gcn_modules['left'] = self.fusion_gcn_modules['right']
|
| 99 |
+
self.proj_linear['left'] = self.proj_linear['right']
|
| 100 |
+
|
| 101 |
+
self.part_para = nn.Parameter(torch.zeros(hidden_dim*len(self.modes)))
|
| 102 |
+
self.pose_proj = nn.Linear(256*4, 768)
|
| 103 |
+
|
| 104 |
+
self.apply(self._init_weights)
|
| 105 |
+
|
| 106 |
+
if self.args.dataset == "Isharah":
|
| 107 |
+
self.lang = 'Arabic'
|
| 108 |
+
elif "CSL" in self.args.dataset:
|
| 109 |
+
self.lang = 'Chinese'
|
| 110 |
+
else:
|
| 111 |
+
self.lang = 'English'
|
| 112 |
+
|
| 113 |
+
if self.args.rgb_support:
|
| 114 |
+
self.rgb_support_backbone = torch.nn.Sequential(*list(torchvision.models.efficientnet_b0(pretrained=True).children())[:-2])
|
| 115 |
+
self.rgb_proj = nn.Conv2d(1280, hidden_dim, kernel_size=1)
|
| 116 |
+
|
| 117 |
+
self.fusion_pose_rgb_linear = nn.Linear(hidden_dim, hidden_dim)
|
| 118 |
+
|
| 119 |
+
# PGF
|
| 120 |
+
self.fusion_pose_rgb_DA = DeformableAttention2D(
|
| 121 |
+
dim = hidden_dim, # feature dimensions
|
| 122 |
+
dim_head = 32, # dimension per head
|
| 123 |
+
heads = 8, # attention heads
|
| 124 |
+
dropout = 0., # dropout
|
| 125 |
+
downsample_factor = 1, # downsample factor (r in paper)
|
| 126 |
+
offset_scale = None, # scale of offset, maximum offset
|
| 127 |
+
offset_groups = None, # number of offset groups, should be multiple of heads
|
| 128 |
+
offset_kernel_size = 1, # offset kernel size
|
| 129 |
+
)
|
| 130 |
+
|
| 131 |
+
self.fusion_gate = nn.Sequential(nn.Conv1d(hidden_dim*2, hidden_dim, 1),
|
| 132 |
+
nn.GELU(),
|
| 133 |
+
nn.Conv1d(hidden_dim, 1, 1),
|
| 134 |
+
nn.Tanh(),
|
| 135 |
+
nn.ReLU(),
|
| 136 |
+
)
|
| 137 |
+
|
| 138 |
+
for layer in self.fusion_gate:
|
| 139 |
+
try:
|
| 140 |
+
if isinDataLoaderance(layer, nn.Conv1d):
|
| 141 |
+
nn.init.constant_(layer.weight, 0)
|
| 142 |
+
nn.init.constant_(layer.bias, 0)
|
| 143 |
+
except:
|
| 144 |
+
print("NOT IMPLEMENTED...")
|
| 145 |
+
|
| 146 |
+
# Načte pouze strukturu architektury z config.json
|
| 147 |
+
mt5_config = MT5Config.from_pretrained(mt5_path)
|
| 148 |
+
# Vytvoří model s prázdnými vahami, které hned v dalším kroku přepíšeme
|
| 149 |
+
self.mt5_model = MT5ForConditionalGeneration(mt5_config)
|
| 150 |
+
|
| 151 |
+
self.mt5_tokenizer = T5Tokenizer.from_pretrained(mt5_path, legacy=False)
|
| 152 |
+
|
| 153 |
+
self.n_registers = args.n_registers
|
| 154 |
+
self.register_position = args.register_position
|
| 155 |
+
self.d_model = self.mt5_model.config.d_model # should be 768
|
| 156 |
+
|
| 157 |
+
if self.n_registers > 0:
|
| 158 |
+
self.register_tokens = nn.Parameter(torch.zeros(self.n_registers, self.d_model))
|
| 159 |
+
# init like other embeddings
|
| 160 |
+
trunc_normal_(self.register_tokens, std=0.02)
|
| 161 |
+
else:
|
| 162 |
+
self.register_tokens = None
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def _init_weights(self, m):
|
| 166 |
+
if isinstance(m, nn.Linear):
|
| 167 |
+
trunc_normal_(m.weight, std=.02)
|
| 168 |
+
if isinstance(m, nn.Linear) and m.bias is not None:
|
| 169 |
+
nn.init.constant_(m.bias, 0)
|
| 170 |
+
elif isinstance(m, nn.LayerNorm):
|
| 171 |
+
nn.init.constant_(m.bias, 0)
|
| 172 |
+
nn.init.constant_(m.weight, 1.0)
|
| 173 |
+
|
| 174 |
+
def maybe_autocast(self, dtype=torch.float32):
|
| 175 |
+
# if on cpu, don't use autocast
|
| 176 |
+
# if on gpu, use autocast with dtype if provided, otherwise use torch.float16
|
| 177 |
+
# enable_autocast = self.device != torch.device("cpu")
|
| 178 |
+
enable_autocast = True
|
| 179 |
+
|
| 180 |
+
if enable_autocast:
|
| 181 |
+
return torch.cuda.amp.autocast(dtype=dtype)
|
| 182 |
+
else:
|
| 183 |
+
return contextlib.nullcontext()
|
| 184 |
+
|
| 185 |
+
def gather_feat_pose_rgb(self, gcn_feat, rgb_feat, indices, rgb_len, pose_init):
|
| 186 |
+
b, c, T, n = gcn_feat.shape
|
| 187 |
+
assert rgb_feat.shape[0] == indices.shape[0]
|
| 188 |
+
rgb_feat = self.rgb_proj(rgb_feat)
|
| 189 |
+
|
| 190 |
+
assert len(rgb_len) == b
|
| 191 |
+
start = 0
|
| 192 |
+
for batch in range(b):
|
| 193 |
+
index = indices[start:start + rgb_len[batch]].to(torch.long)
|
| 194 |
+
# ignore some invalid rgb clip
|
| 195 |
+
if rgb_len[batch] == 1 and -1 in index:
|
| 196 |
+
start = start + rgb_len[batch]
|
| 197 |
+
continue
|
| 198 |
+
|
| 199 |
+
# index selection
|
| 200 |
+
gcn_feat_selected = gcn_feat[batch, :, index]
|
| 201 |
+
rgb_feat_selected = rgb_feat[start:start + rgb_len[batch]]
|
| 202 |
+
pose_init_selected = pose_init[start:start + rgb_len[batch]]
|
| 203 |
+
|
| 204 |
+
gcn_feat_selected = rearrange(gcn_feat_selected, 'c t n -> t c n')
|
| 205 |
+
pose_init_selected = rearrange(pose_init_selected, 't n c -> t c n')
|
| 206 |
+
|
| 207 |
+
# PGF forward
|
| 208 |
+
with self.maybe_autocast():
|
| 209 |
+
fused_transposed = self.fusion_pose_rgb_DA(pose_feat=gcn_feat_selected,
|
| 210 |
+
rgb_feat=rgb_feat_selected,
|
| 211 |
+
pose_init=pose_init_selected, )
|
| 212 |
+
|
| 213 |
+
fused_transposed = fused_transposed.to(gcn_feat.dtype)
|
| 214 |
+
gate_feature = torch.concat([fused_transposed, gcn_feat_selected,], dim=-2)
|
| 215 |
+
gate_score = self.fusion_gate(gate_feature)
|
| 216 |
+
fused_transposed_post = (gate_score) * fused_transposed + (1 - gate_score) * gcn_feat_selected
|
| 217 |
+
|
| 218 |
+
gcn_feat = gcn_feat.clone()
|
| 219 |
+
fused_transposed_post = rearrange(fused_transposed_post, 't c n -> c t n')
|
| 220 |
+
|
| 221 |
+
# replace gcn feature
|
| 222 |
+
gcn_feat[batch, :, index] = fused_transposed_post
|
| 223 |
+
start = start + rgb_len[batch]
|
| 224 |
+
|
| 225 |
+
assert start == rgb_feat.shape[0]
|
| 226 |
+
return gcn_feat
|
| 227 |
+
|
| 228 |
+
def forward(self, src_input, tgt_input):
|
| 229 |
+
# RGB branch forward
|
| 230 |
+
if self.args.rgb_support:
|
| 231 |
+
rgb_support_dict = {}
|
| 232 |
+
for index_key, rgb_key in zip(['left_sampled_indices', 'right_sampled_indices'], ['left_hands', 'right_hands']):
|
| 233 |
+
rgb_feat = self.rgb_support_backbone(src_input[rgb_key])
|
| 234 |
+
|
| 235 |
+
rgb_support_dict[index_key] = src_input[index_key]
|
| 236 |
+
rgb_support_dict[rgb_key] = rgb_feat
|
| 237 |
+
|
| 238 |
+
# Pose branch forward
|
| 239 |
+
features = []
|
| 240 |
+
|
| 241 |
+
body_feat = None
|
| 242 |
+
for part in self.modes:
|
| 243 |
+
# project position to hidden dim
|
| 244 |
+
proj_feat = self.proj_linear[part](src_input[part]).permute(0,3,1,2) #B,C,T,V
|
| 245 |
+
# spatial gcn forward
|
| 246 |
+
gcn_feat = self.gcn_modules[part](proj_feat)
|
| 247 |
+
if part == 'body':
|
| 248 |
+
body_feat = gcn_feat
|
| 249 |
+
|
| 250 |
+
else:
|
| 251 |
+
assert not body_feat is None
|
| 252 |
+
if part == 'left':
|
| 253 |
+
# Pose RGB fusion
|
| 254 |
+
if self.args.rgb_support:
|
| 255 |
+
gcn_feat = self.gather_feat_pose_rgb(gcn_feat,
|
| 256 |
+
rgb_support_dict[f'{part}_hands'],
|
| 257 |
+
rgb_support_dict[f'{part}_sampled_indices'],
|
| 258 |
+
src_input[f'{part}_rgb_len'],
|
| 259 |
+
src_input[f'{part}_skeletons_norm'],
|
| 260 |
+
)
|
| 261 |
+
|
| 262 |
+
gcn_feat = gcn_feat + body_feat[..., -2][...,None].detach()
|
| 263 |
+
|
| 264 |
+
elif part == 'right':
|
| 265 |
+
# Pose RGB fusion
|
| 266 |
+
if self.args.rgb_support:
|
| 267 |
+
gcn_feat = self.gather_feat_pose_rgb(gcn_feat,
|
| 268 |
+
rgb_support_dict[f'{part}_hands'],
|
| 269 |
+
rgb_support_dict[f'{part}_sampled_indices'],
|
| 270 |
+
src_input[f'{part}_rgb_len'],
|
| 271 |
+
src_input[f'{part}_skeletons_norm'],
|
| 272 |
+
)
|
| 273 |
+
|
| 274 |
+
gcn_feat = gcn_feat + body_feat[..., -1][...,None].detach()
|
| 275 |
+
|
| 276 |
+
elif part == 'face_all':
|
| 277 |
+
gcn_feat = gcn_feat + body_feat[..., 0][...,None].detach()
|
| 278 |
+
|
| 279 |
+
else:
|
| 280 |
+
raise NotImplementedError
|
| 281 |
+
|
| 282 |
+
# temporal gcn forward
|
| 283 |
+
gcn_feat = self.fusion_gcn_modules[part](gcn_feat) #B,C,T,V
|
| 284 |
+
pool_feat = gcn_feat.mean(-1).transpose(1,2) #B,T,C
|
| 285 |
+
features.append(pool_feat)
|
| 286 |
+
|
| 287 |
+
# concat sub-pose feature across token dimension
|
| 288 |
+
inputs_embeds = torch.cat(features, dim=-1) + self.part_para
|
| 289 |
+
inputs_embeds = self.pose_proj(inputs_embeds)
|
| 290 |
+
|
| 291 |
+
prefix_token = self.mt5_tokenizer(
|
| 292 |
+
[f"Translate sign language video to {self.lang}: "] * len(tgt_input["gt_sentence"]),
|
| 293 |
+
padding="longest",
|
| 294 |
+
truncation=True,
|
| 295 |
+
return_tensors="pt",
|
| 296 |
+
).to(inputs_embeds.device)
|
| 297 |
+
|
| 298 |
+
prefix_embeds = self.mt5_model.encoder.embed_tokens(prefix_token['input_ids'])
|
| 299 |
+
|
| 300 |
+
if self.n_registers > 0:
|
| 301 |
+
B = inputs_embeds.size(0)
|
| 302 |
+
|
| 303 |
+
# expand registers for batch
|
| 304 |
+
register_embeds = self.register_tokens.unsqueeze(0).expand(B, -1, -1)
|
| 305 |
+
# shape: (B, 4, 768)
|
| 306 |
+
|
| 307 |
+
register_mask = torch.ones((B, self.n_registers), device=inputs_embeds.device,dtype=prefix_token['attention_mask'].dtype)
|
| 308 |
+
|
| 309 |
+
if self.register_position == 'before_all':
|
| 310 |
+
# prepend order: [registers | prefix | pose_tokens]
|
| 311 |
+
inputs_embeds = torch.cat([register_embeds, prefix_embeds, inputs_embeds], dim=1)
|
| 312 |
+
attention_mask = torch.cat([register_mask, prefix_token['attention_mask'], src_input['attention_mask']], dim=1)
|
| 313 |
+
|
| 314 |
+
elif self.register_position == 'after_prefix' or self.register_position == 'before_pose':
|
| 315 |
+
# prepend order: [prefix | registers | pose_tokens]
|
| 316 |
+
inputs_embeds = torch.cat([prefix_embeds, register_embeds, inputs_embeds], dim=1)
|
| 317 |
+
attention_mask = torch.cat([prefix_token['attention_mask'], register_mask, src_input['attention_mask']], dim=1)
|
| 318 |
+
|
| 319 |
+
elif self.register_position == 'after_valid_pose':
|
| 320 |
+
# prepend order: [prefix | valid_pose_tokens | registers | padded_pose_tokens]
|
| 321 |
+
inputs_list = []
|
| 322 |
+
mask_list = []
|
| 323 |
+
|
| 324 |
+
for b in range(B):
|
| 325 |
+
valid_len = int(src_input['attention_mask'][b].sum().item())
|
| 326 |
+
|
| 327 |
+
pose_valid = inputs_embeds[b, :valid_len]
|
| 328 |
+
pose_pad = inputs_embeds[b, valid_len:]
|
| 329 |
+
|
| 330 |
+
emb = torch.cat(
|
| 331 |
+
[prefix_embeds[b],
|
| 332 |
+
pose_valid,
|
| 333 |
+
register_embeds[b],
|
| 334 |
+
pose_pad],
|
| 335 |
+
dim=0
|
| 336 |
+
)
|
| 337 |
+
|
| 338 |
+
m = torch.cat(
|
| 339 |
+
[prefix_token['attention_mask'][b],
|
| 340 |
+
src_input['attention_mask'][b, :valid_len],
|
| 341 |
+
register_mask[b],
|
| 342 |
+
src_input['attention_mask'][b, valid_len:]],
|
| 343 |
+
dim=0
|
| 344 |
+
)
|
| 345 |
+
|
| 346 |
+
inputs_list.append(emb)
|
| 347 |
+
mask_list.append(m)
|
| 348 |
+
|
| 349 |
+
inputs_embeds = torch.stack(inputs_list, dim=0)
|
| 350 |
+
attention_mask = torch.stack(mask_list, dim=0)
|
| 351 |
+
|
| 352 |
+
elif self.register_position == 'after_all':
|
| 353 |
+
# prepend order: [prefix | pose_tokens | registers]
|
| 354 |
+
inputs_embeds = torch.cat([prefix_embeds, inputs_embeds, register_embeds], dim=1)
|
| 355 |
+
attention_mask = torch.cat([prefix_token['attention_mask'], src_input['attention_mask'], register_mask], dim=1)
|
| 356 |
+
|
| 357 |
+
else:
|
| 358 |
+
# prepend order: [prefix | pose_tokens]
|
| 359 |
+
inputs_embeds = torch.cat([prefix_embeds, inputs_embeds], dim=1)
|
| 360 |
+
attention_mask = torch.cat([prefix_token['attention_mask'], src_input['attention_mask']], dim=1)
|
| 361 |
+
|
| 362 |
+
tgt_input_tokenizer = self.mt5_tokenizer(tgt_input['gt_sentence'],
|
| 363 |
+
return_tensors="pt",
|
| 364 |
+
padding=True,
|
| 365 |
+
truncation=True,
|
| 366 |
+
max_length=50)
|
| 367 |
+
|
| 368 |
+
labels = tgt_input_tokenizer['input_ids']
|
| 369 |
+
labels[labels == self.mt5_tokenizer.pad_token_id] = -100
|
| 370 |
+
|
| 371 |
+
out = self.mt5_model(inputs_embeds = inputs_embeds,
|
| 372 |
+
attention_mask = attention_mask,
|
| 373 |
+
labels = labels.to(inputs_embeds.device),
|
| 374 |
+
return_dict = True,
|
| 375 |
+
)
|
| 376 |
+
|
| 377 |
+
label = labels.reshape(-1)
|
| 378 |
+
out_logits = out['logits']
|
| 379 |
+
logits = out_logits.reshape(-1,out_logits.shape[-1])
|
| 380 |
+
loss_fct = torch.nn.CrossEntropyLoss(label_smoothing=self.args.label_smoothing, ignore_index=-100)
|
| 381 |
+
loss = loss_fct(logits, label.to(out_logits.device, non_blocking=True))
|
| 382 |
+
|
| 383 |
+
stack_out = {
|
| 384 |
+
# use for inference
|
| 385 |
+
'inputs_embeds':inputs_embeds,
|
| 386 |
+
'attention_mask':attention_mask,
|
| 387 |
+
'loss':loss,
|
| 388 |
+
}
|
| 389 |
+
|
| 390 |
+
return stack_out
|
| 391 |
+
|
| 392 |
+
@torch.no_grad()
|
| 393 |
+
def generate(self,pre_compute_item,max_new_tokens,num_beams):
|
| 394 |
+
inputs_embeds = pre_compute_item['inputs_embeds']
|
| 395 |
+
attention_mask = pre_compute_item['attention_mask']
|
| 396 |
+
|
| 397 |
+
out = self.mt5_model.generate(inputs_embeds = inputs_embeds,
|
| 398 |
+
attention_mask = attention_mask,
|
| 399 |
+
max_new_tokens=max_new_tokens,
|
| 400 |
+
num_beams = num_beams,
|
| 401 |
+
)
|
| 402 |
+
|
| 403 |
+
return out
|
| 404 |
+
|
| 405 |
+
def get_requires_grad_dict(model):
|
| 406 |
+
param_requires_grad = {name: True for name, param in model.named_parameters()}
|
| 407 |
+
param_requires_grad_right = {}
|
| 408 |
+
for key in param_requires_grad.keys():
|
| 409 |
+
if 'left' in key:
|
| 410 |
+
param_requires_grad_right[key.replace("left", 'right')] = param_requires_grad[key]
|
| 411 |
+
param_requires_grad = {**param_requires_grad,
|
| 412 |
+
**param_requires_grad_right}
|
| 413 |
+
params_to_update = {k: v for k, v in model.state_dict().items() if param_requires_grad.get(k, True)}
|
| 414 |
+
|
| 415 |
+
return params_to_update
|
Uni_Sign/normalization.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
from typing import Tuple
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
def get_keypoints(joints, landmarks_name):
|
| 6 |
+
frames_keypoints = joints[landmarks_name]
|
| 7 |
+
frames_keypoints = frames_keypoints[:, :, :2]
|
| 8 |
+
|
| 9 |
+
return frames_keypoints
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def output_keypoints(joints, valid_frames, frames_keypoints):
|
| 13 |
+
frames_names = np.array(list(joints.keys()))[valid_frames]
|
| 14 |
+
frames_keypoints = frames_keypoints[valid_frames]
|
| 15 |
+
|
| 16 |
+
return dict(zip(frames_names, frames_keypoints))
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def safe_divide(a, b):
|
| 20 |
+
return np.divide(a, b, out=np.zeros_like(a), where=b != 0)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def local_keypoint_normalization(joints: dict, landmarks: str, select_idx: list = [], padding: float = 0.1) -> dict:
|
| 24 |
+
frames_keypoints = get_keypoints(joints, landmarks)
|
| 25 |
+
|
| 26 |
+
if select_idx:
|
| 27 |
+
frames_keypoints = frames_keypoints[:, select_idx, :]
|
| 28 |
+
|
| 29 |
+
# move to origin
|
| 30 |
+
xmin = np.min(frames_keypoints[:, :, 0], axis=1)
|
| 31 |
+
ymin = np.min(frames_keypoints[:, :, 1], axis=1)
|
| 32 |
+
|
| 33 |
+
frames_keypoints[:, :, 0] -= xmin[:, np.newaxis]
|
| 34 |
+
frames_keypoints[:, :, 1] -= ymin[:, np.newaxis]
|
| 35 |
+
|
| 36 |
+
# pad to square
|
| 37 |
+
xmax = np.max(frames_keypoints[:, :, 0], axis=1)
|
| 38 |
+
ymax = np.max(frames_keypoints[:, :, 1], axis=1)
|
| 39 |
+
|
| 40 |
+
dif_full = np.abs(xmax - ymax)
|
| 41 |
+
dif = np.floor(dif_full / 2)
|
| 42 |
+
|
| 43 |
+
for i in range(len(dif)):
|
| 44 |
+
if xmax[i] > ymax[i]:
|
| 45 |
+
ymax[i] += dif_full[i]
|
| 46 |
+
frames_keypoints[i, :, 1] += dif[i]
|
| 47 |
+
else:
|
| 48 |
+
xmax[i] += dif_full[i]
|
| 49 |
+
frames_keypoints[i, :, 0] += dif[i]
|
| 50 |
+
|
| 51 |
+
# add padding to all sides
|
| 52 |
+
side_size = np.max([xmax, ymax], axis=0)
|
| 53 |
+
padding = side_size * padding
|
| 54 |
+
|
| 55 |
+
frames_keypoints += padding[:, np.newaxis, np.newaxis]
|
| 56 |
+
xmax += padding * 2
|
| 57 |
+
ymax += padding * 2
|
| 58 |
+
|
| 59 |
+
# normalize to [-1, 1]
|
| 60 |
+
frames_keypoints = safe_divide(frames_keypoints, xmax[:, np.newaxis, np.newaxis])
|
| 61 |
+
# frames_keypoints /= xmax[:, np.newaxis, np.newaxis]
|
| 62 |
+
frames_keypoints = frames_keypoints * 2 - 1
|
| 63 |
+
|
| 64 |
+
return frames_keypoints
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def global_keypoint_normalization(
|
| 68 |
+
joints: dict,
|
| 69 |
+
landmarks: str,
|
| 70 |
+
add_landmarks_names: list,
|
| 71 |
+
face_select_idx: list = [],
|
| 72 |
+
sign_area_size: tuple = (1.5, 1.5),
|
| 73 |
+
l_shoulder_idx: int = 11,
|
| 74 |
+
r_shoulder_idx: int = 12) -> Tuple[dict, dict]:
|
| 75 |
+
frames_keypoints = get_keypoints(joints, landmarks)
|
| 76 |
+
|
| 77 |
+
# get distance between right and left shoulder
|
| 78 |
+
l_shoulder_points = frames_keypoints[:, l_shoulder_idx, :]
|
| 79 |
+
r_shoulder_points = frames_keypoints[:, r_shoulder_idx, :]
|
| 80 |
+
distance = np.sqrt((l_shoulder_points[:, 0] - r_shoulder_points[:, 0]) ** 2 + (
|
| 81 |
+
l_shoulder_points[:, 1] - r_shoulder_points[:, 1]) ** 2)
|
| 82 |
+
|
| 83 |
+
# get center point between shoulders
|
| 84 |
+
center_x = np.abs(l_shoulder_points[:, 0] - r_shoulder_points[:, 0]) / 2 + np.min(
|
| 85 |
+
[l_shoulder_points[:, 0], r_shoulder_points[:, 0]], 0)
|
| 86 |
+
center_y = np.abs(l_shoulder_points[:, 1] - r_shoulder_points[:, 1]) / 2 + np.min(
|
| 87 |
+
[l_shoulder_points[:, 1], r_shoulder_points[:, 1]], 0)
|
| 88 |
+
sign_area_size = np.array(sign_area_size) * distance[:, np.newaxis]
|
| 89 |
+
|
| 90 |
+
# normalize
|
| 91 |
+
frames_keypoints[:, :, 0] -= center_x[:, np.newaxis]
|
| 92 |
+
frames_keypoints[:, :, 1] -= center_y[:, np.newaxis]
|
| 93 |
+
|
| 94 |
+
# frames_keypoints[:, :, 0] /= sign_area_size[:, 1, np.newaxis]
|
| 95 |
+
# frames_keypoints[:, :, 1] /= sign_area_size[:, 0, np.newaxis]
|
| 96 |
+
frames_keypoints[:, :, 0] = safe_divide(frames_keypoints[:, :, 0], sign_area_size[:, 0, np.newaxis])
|
| 97 |
+
frames_keypoints[:, :, 1] = safe_divide(frames_keypoints[:, :, 1], sign_area_size[:, 1, np.newaxis])
|
| 98 |
+
|
| 99 |
+
# normalize additional landmarks
|
| 100 |
+
add_landmarks = {}
|
| 101 |
+
for add_landmarks_name in add_landmarks_names:
|
| 102 |
+
add_frames_keypoints = get_keypoints(joints, add_landmarks_name)
|
| 103 |
+
|
| 104 |
+
if face_select_idx and add_landmarks_name == "face_landmarks":
|
| 105 |
+
add_frames_keypoints = add_frames_keypoints[:, face_select_idx, :]
|
| 106 |
+
|
| 107 |
+
add_frames_keypoints[:, :, 0] -= center_x[:, np.newaxis]
|
| 108 |
+
add_frames_keypoints[:, :, 1] -= center_y[:, np.newaxis]
|
| 109 |
+
|
| 110 |
+
# add_frames_keypoints[:, :, 0] /= sign_area_size[:, 1, np.newaxis]
|
| 111 |
+
# add_frames_keypoints[:, :, 1] /= sign_area_size[:, 0, np.newaxis]
|
| 112 |
+
add_frames_keypoints[:, :, 0] = safe_divide(add_frames_keypoints[:, :, 0], sign_area_size[:, 0, np.newaxis])
|
| 113 |
+
add_frames_keypoints[:, :, 1] = safe_divide(add_frames_keypoints[:, :, 1], sign_area_size[:, 1, np.newaxis])
|
| 114 |
+
|
| 115 |
+
add_landmarks[add_landmarks_name] = add_frames_keypoints
|
| 116 |
+
|
| 117 |
+
return frames_keypoints, add_landmarks
|
Uni_Sign/stgcn_layers/.DS_Store
ADDED
|
Binary file (6.15 kB). View file
|
|
|
Uni_Sign/stgcn_layers/__init__.py
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .gcn_utils import Graph
|
| 2 |
+
from .stgcn_block import get_stgcn_chain
|
Uni_Sign/stgcn_layers/__pycache__/__init__.cpython-311.pyc
ADDED
|
Binary file (294 Bytes). View file
|
|
|
Uni_Sign/stgcn_layers/__pycache__/__init__.cpython-39.pyc
ADDED
|
Binary file (246 Bytes). View file
|
|
|
Uni_Sign/stgcn_layers/__pycache__/gcn_utils.cpython-311.pyc
ADDED
|
Binary file (13.3 kB). View file
|
|
|
Uni_Sign/stgcn_layers/__pycache__/gcn_utils.cpython-39.pyc
ADDED
|
Binary file (7.52 kB). View file
|
|
|
Uni_Sign/stgcn_layers/__pycache__/stgcn_block.cpython-311.pyc
ADDED
|
Binary file (6.54 kB). View file
|
|
|
Uni_Sign/stgcn_layers/__pycache__/stgcn_block.cpython-39.pyc
ADDED
|
Binary file (3.44 kB). View file
|
|
|
Uni_Sign/stgcn_layers/gcn_utils.py
ADDED
|
@@ -0,0 +1,350 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
import pdb
|
| 5 |
+
import math
|
| 6 |
+
import copy
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
class Graph:
|
| 10 |
+
"""The Graph to model the skeletons extracted by the openpose
|
| 11 |
+
|
| 12 |
+
Args:
|
| 13 |
+
strategy (string): must be one of the follow candidates
|
| 14 |
+
- uniform: Uniform Labeling
|
| 15 |
+
- distance: Distance Partitioning
|
| 16 |
+
- spatial: Spatial Configuration
|
| 17 |
+
For more information, please refer to the section 'Partition Strategies'
|
| 18 |
+
in our paper (https://arxiv.org/abs/1801.07455).
|
| 19 |
+
|
| 20 |
+
layout (string): must be one of the follow candidates
|
| 21 |
+
- openpose: Is consists of 18 joints. For more information, please
|
| 22 |
+
refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose#output
|
| 23 |
+
- ntu-rgb+d: Is consists of 25 joints. For more information, please
|
| 24 |
+
refer to https://github.com/shahroudy/NTURGB-D
|
| 25 |
+
|
| 26 |
+
max_hop (int): the maximal distance between two connected nodes
|
| 27 |
+
dilation (int): controls the spacing between the kernel points
|
| 28 |
+
|
| 29 |
+
"""
|
| 30 |
+
|
| 31 |
+
def __init__(self, layout='custom', strategy='uniform', max_hop=1, dilation=1):
|
| 32 |
+
self.max_hop = max_hop
|
| 33 |
+
self.dilation = dilation
|
| 34 |
+
|
| 35 |
+
self.get_edge(layout)
|
| 36 |
+
self.hop_dis = get_hop_distance(self.num_node, self.edge, max_hop=max_hop)
|
| 37 |
+
self.get_adjacency(strategy)
|
| 38 |
+
|
| 39 |
+
def __str__(self):
|
| 40 |
+
return self.A
|
| 41 |
+
|
| 42 |
+
def get_edge(self, layout):
|
| 43 |
+
# 'body', 'left', 'right', 'mouth', 'face'
|
| 44 |
+
# if layout == 'custom_hand21':
|
| 45 |
+
|
| 46 |
+
if layout == 'default_left' or layout == 'default_right':
|
| 47 |
+
self.num_node = 21
|
| 48 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 49 |
+
neighbor_1base = [
|
| 50 |
+
[0, 1],
|
| 51 |
+
[1, 2],
|
| 52 |
+
[2, 3],
|
| 53 |
+
[3, 4],
|
| 54 |
+
[0, 5],
|
| 55 |
+
[5, 6],
|
| 56 |
+
[6, 7],
|
| 57 |
+
[7, 8],
|
| 58 |
+
[0, 9],
|
| 59 |
+
[9, 10],
|
| 60 |
+
[10, 11],
|
| 61 |
+
[11, 12],
|
| 62 |
+
[0, 13],
|
| 63 |
+
[13, 14],
|
| 64 |
+
[14, 15],
|
| 65 |
+
[15, 16],
|
| 66 |
+
[0, 17],
|
| 67 |
+
[17, 18],
|
| 68 |
+
[18, 19],
|
| 69 |
+
[19, 20],
|
| 70 |
+
]
|
| 71 |
+
neighbor_link = neighbor_1base
|
| 72 |
+
self.edge = self_link + neighbor_link
|
| 73 |
+
self.center = 0
|
| 74 |
+
|
| 75 |
+
elif layout == 'default_body':
|
| 76 |
+
self.num_node = 9
|
| 77 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 78 |
+
neighbor_1base = [
|
| 79 |
+
[0, 1],
|
| 80 |
+
[0, 2],
|
| 81 |
+
[0, 3],
|
| 82 |
+
[0, 4],
|
| 83 |
+
[3, 5],
|
| 84 |
+
[5, 7],
|
| 85 |
+
[4, 6],
|
| 86 |
+
[6, 8],
|
| 87 |
+
]
|
| 88 |
+
neighbor_link = neighbor_1base
|
| 89 |
+
self.edge = self_link + neighbor_link
|
| 90 |
+
self.center = 0
|
| 91 |
+
|
| 92 |
+
elif layout == 'default_face_all':
|
| 93 |
+
self.num_node = 9 + 8 + 1
|
| 94 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 95 |
+
neighbor_1base = [[i, i + 1] for i in range(9 - 1)] + \
|
| 96 |
+
[[i, i + 1] for i in range(9, 9 + 8 - 1)] + \
|
| 97 |
+
[[9 + 8 - 1, 9]] + \
|
| 98 |
+
[[17, i] for i in range(17)]
|
| 99 |
+
neighbor_link = neighbor_1base
|
| 100 |
+
self.edge = self_link + neighbor_link
|
| 101 |
+
self.center = self.num_node - 1
|
| 102 |
+
|
| 103 |
+
elif layout in ['default_ytasl_left', 'default_ytasl_right', 'pruned_ytasl_left', 'pruned_ytasl_right', 'isharah_ytasl_left', 'isharah_ytasl_right']:
|
| 104 |
+
self.num_node = 21
|
| 105 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 106 |
+
neighbor_1base = [
|
| 107 |
+
[3, 4],
|
| 108 |
+
[0, 5],
|
| 109 |
+
[17, 18],
|
| 110 |
+
[0, 17],
|
| 111 |
+
[13, 14],
|
| 112 |
+
[13, 17],
|
| 113 |
+
[18, 19],
|
| 114 |
+
[5, 6],
|
| 115 |
+
[5, 9],
|
| 116 |
+
[14, 15],
|
| 117 |
+
[0, 1],
|
| 118 |
+
[9, 10],
|
| 119 |
+
[1, 2],
|
| 120 |
+
[9, 13],
|
| 121 |
+
[10, 11],
|
| 122 |
+
[19, 20],
|
| 123 |
+
[6, 7],
|
| 124 |
+
[15, 16],
|
| 125 |
+
[2, 3],
|
| 126 |
+
[11, 12],
|
| 127 |
+
[7, 8]
|
| 128 |
+
]
|
| 129 |
+
neighbor_link = neighbor_1base
|
| 130 |
+
self.edge = self_link + neighbor_link
|
| 131 |
+
self.center = 0
|
| 132 |
+
|
| 133 |
+
elif layout == 'default_ytasl_body':
|
| 134 |
+
self.num_node = 25
|
| 135 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 136 |
+
neighbor_1base = [
|
| 137 |
+
[15, 21],
|
| 138 |
+
[16, 20],
|
| 139 |
+
[18, 20],
|
| 140 |
+
[3, 7],
|
| 141 |
+
[14, 16],
|
| 142 |
+
[11, 23],
|
| 143 |
+
[6, 8],
|
| 144 |
+
[15, 17],
|
| 145 |
+
[16, 22],
|
| 146 |
+
[4, 5],
|
| 147 |
+
[5, 6],
|
| 148 |
+
[12, 24],
|
| 149 |
+
[23, 24],
|
| 150 |
+
[0, 1],
|
| 151 |
+
[9, 10],
|
| 152 |
+
[1, 2],
|
| 153 |
+
[0, 4],
|
| 154 |
+
[11, 13],
|
| 155 |
+
[15, 19],
|
| 156 |
+
[16, 18],
|
| 157 |
+
[12, 14],
|
| 158 |
+
[17, 19],
|
| 159 |
+
[2, 3],
|
| 160 |
+
[11, 12],
|
| 161 |
+
[13, 15]
|
| 162 |
+
]
|
| 163 |
+
neighbor_link = neighbor_1base
|
| 164 |
+
self.edge = self_link + neighbor_link
|
| 165 |
+
self.center = 0
|
| 166 |
+
|
| 167 |
+
elif layout == 'default_ytasl_face_all':
|
| 168 |
+
self.num_node = 37
|
| 169 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 170 |
+
neighbor_1base = [
|
| 171 |
+
[16, 18],
|
| 172 |
+
[18, 13],
|
| 173 |
+
[13, 7],
|
| 174 |
+
[7, 8],
|
| 175 |
+
[8, 15],
|
| 176 |
+
[15, 24],
|
| 177 |
+
[24, 23],
|
| 178 |
+
[23, 29],
|
| 179 |
+
[29, 32],
|
| 180 |
+
[32, 16],
|
| 181 |
+
[5, 17],
|
| 182 |
+
[17, 14],
|
| 183 |
+
[30, 31],
|
| 184 |
+
[31, 21],
|
| 185 |
+
[11, 1],
|
| 186 |
+
[1, 27],
|
| 187 |
+
[10, 6],
|
| 188 |
+
[6, 0],
|
| 189 |
+
[0, 22],
|
| 190 |
+
[22, 26],
|
| 191 |
+
[26, 34],
|
| 192 |
+
[34, 4],
|
| 193 |
+
[4, 20],
|
| 194 |
+
[20, 10],
|
| 195 |
+
[10, 12],
|
| 196 |
+
[12, 2],
|
| 197 |
+
[2, 28],
|
| 198 |
+
[28, 26],
|
| 199 |
+
[10, 19],
|
| 200 |
+
[19, 3],
|
| 201 |
+
[3, 33],
|
| 202 |
+
[33, 26]
|
| 203 |
+
]
|
| 204 |
+
neighbor_link = neighbor_1base
|
| 205 |
+
self.edge = self_link + neighbor_link
|
| 206 |
+
self.center = self.num_node - 1
|
| 207 |
+
|
| 208 |
+
elif layout in ['pruned_ytasl_body', 'isharah_ytasl_body']:
|
| 209 |
+
self.num_node = 9
|
| 210 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 211 |
+
neighbor_1base = [
|
| 212 |
+
[0, 1],
|
| 213 |
+
[0, 2],
|
| 214 |
+
[0, 3],
|
| 215 |
+
[0, 4],
|
| 216 |
+
[3, 5],
|
| 217 |
+
[4, 6],
|
| 218 |
+
[5, 7],
|
| 219 |
+
[6, 8],
|
| 220 |
+
]
|
| 221 |
+
neighbor_link = neighbor_1base
|
| 222 |
+
self.edge = self_link + neighbor_link
|
| 223 |
+
self.center = 0
|
| 224 |
+
|
| 225 |
+
elif layout == 'pruned_ytasl_face_all':
|
| 226 |
+
self.num_node = 18
|
| 227 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 228 |
+
neighbor_1base = [
|
| 229 |
+
[5, 8],
|
| 230 |
+
[8, 6],
|
| 231 |
+
[6, 14],
|
| 232 |
+
[14, 12],
|
| 233 |
+
[3, 4],
|
| 234 |
+
[4, 1],
|
| 235 |
+
[1, 11],
|
| 236 |
+
[11, 10],
|
| 237 |
+
[3, 9],
|
| 238 |
+
[9, 2],
|
| 239 |
+
[2, 15],
|
| 240 |
+
[15, 10],
|
| 241 |
+
[7, 16],
|
| 242 |
+
[13, 17],
|
| 243 |
+
]
|
| 244 |
+
neighbor_link = neighbor_1base
|
| 245 |
+
self.edge = self_link + neighbor_link
|
| 246 |
+
self.center = self.num_node - 1
|
| 247 |
+
|
| 248 |
+
elif layout == 'isharah_ytasl_face_all':
|
| 249 |
+
self.num_node = 19
|
| 250 |
+
self_link = [(i, i) for i in range(self.num_node)]
|
| 251 |
+
neighbor_1base = [
|
| 252 |
+
[0, 2],
|
| 253 |
+
[2, 3],
|
| 254 |
+
[3, 4],
|
| 255 |
+
[4, 10],
|
| 256 |
+
[10, 5],
|
| 257 |
+
[5, 8],
|
| 258 |
+
[8, 7],
|
| 259 |
+
[7, 9],
|
| 260 |
+
[9, 6],
|
| 261 |
+
[6, 1],
|
| 262 |
+
[1, 15],
|
| 263 |
+
[15, 18],
|
| 264 |
+
[18, 16],
|
| 265 |
+
[16, 17],
|
| 266 |
+
[17, 14],
|
| 267 |
+
[14, 13],
|
| 268 |
+
[13, 12],
|
| 269 |
+
[12, 11],
|
| 270 |
+
[11, 0],
|
| 271 |
+
]
|
| 272 |
+
neighbor_link = neighbor_1base
|
| 273 |
+
self.edge = self_link + neighbor_link
|
| 274 |
+
self.center = self.num_node - 1
|
| 275 |
+
|
| 276 |
+
else:
|
| 277 |
+
raise NotImplementedError(f"Layout not implemented for: {layout}")
|
| 278 |
+
|
| 279 |
+
def get_adjacency(self, strategy):
|
| 280 |
+
valid_hop = range(0, self.max_hop + 1, self.dilation)
|
| 281 |
+
adjacency = np.zeros((self.num_node, self.num_node))
|
| 282 |
+
for hop in valid_hop:
|
| 283 |
+
adjacency[self.hop_dis == hop] = 1
|
| 284 |
+
normalize_adjacency = normalize_digraph(adjacency)
|
| 285 |
+
|
| 286 |
+
if strategy == 'uniform':
|
| 287 |
+
A = np.zeros((1, self.num_node, self.num_node))
|
| 288 |
+
A[0] = normalize_adjacency
|
| 289 |
+
self.A = A
|
| 290 |
+
elif strategy == 'distance':
|
| 291 |
+
A = np.zeros((len(valid_hop), self.num_node, self.num_node))
|
| 292 |
+
for i, hop in enumerate(valid_hop):
|
| 293 |
+
A[i][self.hop_dis == hop] = normalize_adjacency[self.hop_dis == hop]
|
| 294 |
+
self.A = A
|
| 295 |
+
elif strategy == 'spatial':
|
| 296 |
+
A = []
|
| 297 |
+
for hop in valid_hop:
|
| 298 |
+
a_root = np.zeros((self.num_node, self.num_node))
|
| 299 |
+
a_close = np.zeros((self.num_node, self.num_node))
|
| 300 |
+
a_further = np.zeros((self.num_node, self.num_node))
|
| 301 |
+
for i in range(self.num_node):
|
| 302 |
+
for j in range(self.num_node):
|
| 303 |
+
if self.hop_dis[j, i] == hop:
|
| 304 |
+
if (
|
| 305 |
+
self.hop_dis[j, self.center]
|
| 306 |
+
== self.hop_dis[i, self.center]
|
| 307 |
+
):
|
| 308 |
+
a_root[j, i] = normalize_adjacency[j, i]
|
| 309 |
+
elif (
|
| 310 |
+
self.hop_dis[j, self.center]
|
| 311 |
+
> self.hop_dis[i, self.center]
|
| 312 |
+
):
|
| 313 |
+
a_close[j, i] = normalize_adjacency[j, i]
|
| 314 |
+
else:
|
| 315 |
+
a_further[j, i] = normalize_adjacency[j, i]
|
| 316 |
+
if hop == 0:
|
| 317 |
+
A.append(a_root)
|
| 318 |
+
else:
|
| 319 |
+
A.append(a_root + a_close)
|
| 320 |
+
A.append(a_further)
|
| 321 |
+
A = np.stack(A)
|
| 322 |
+
self.A = A
|
| 323 |
+
else:
|
| 324 |
+
raise ValueError("Do Not Exist This Strategy")
|
| 325 |
+
|
| 326 |
+
|
| 327 |
+
def get_hop_distance(num_node, edge, max_hop=1):
|
| 328 |
+
A = np.zeros((num_node, num_node))
|
| 329 |
+
for i, j in edge:
|
| 330 |
+
A[j, i] = 1
|
| 331 |
+
A[i, j] = 1
|
| 332 |
+
|
| 333 |
+
# compute hop steps
|
| 334 |
+
hop_dis = np.zeros((num_node, num_node)) + np.inf
|
| 335 |
+
transfer_mat = [np.linalg.matrix_power(A, d) for d in range(max_hop + 1)]
|
| 336 |
+
arrive_mat = np.stack(transfer_mat) > 0
|
| 337 |
+
for d in range(max_hop, -1, -1):
|
| 338 |
+
hop_dis[arrive_mat[d]] = d
|
| 339 |
+
return hop_dis
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
def normalize_digraph(A):
|
| 343 |
+
Dl = np.sum(A, 0)
|
| 344 |
+
num_node = A.shape[0]
|
| 345 |
+
Dn = np.zeros((num_node, num_node))
|
| 346 |
+
for i in range(num_node):
|
| 347 |
+
if Dl[i] > 0:
|
| 348 |
+
Dn[i, i] = Dl[i] ** (-1)
|
| 349 |
+
AD = np.dot(A, Dn)
|
| 350 |
+
return AD
|
Uni_Sign/stgcn_layers/stgcn_block.py
ADDED
|
@@ -0,0 +1,128 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
import pdb
|
| 5 |
+
import math
|
| 6 |
+
import copy
|
| 7 |
+
|
| 8 |
+
class GCN_unit(nn.Module):
|
| 9 |
+
def __init__(
|
| 10 |
+
self,
|
| 11 |
+
in_channels,
|
| 12 |
+
out_channels,
|
| 13 |
+
kernel_size,
|
| 14 |
+
A,
|
| 15 |
+
adaptive=True,
|
| 16 |
+
t_kernel_size=1,
|
| 17 |
+
t_stride=1,
|
| 18 |
+
t_padding=0,
|
| 19 |
+
t_dilation=1,
|
| 20 |
+
bias=True,
|
| 21 |
+
):
|
| 22 |
+
super().__init__()
|
| 23 |
+
self.kernel_size = kernel_size
|
| 24 |
+
assert A.size(0) == self.kernel_size
|
| 25 |
+
self.conv = nn.Conv2d(
|
| 26 |
+
in_channels,
|
| 27 |
+
out_channels * kernel_size,
|
| 28 |
+
kernel_size=(t_kernel_size, 1),
|
| 29 |
+
padding=(t_padding, 0),
|
| 30 |
+
stride=(t_stride, 1),
|
| 31 |
+
dilation=(t_dilation, 1),
|
| 32 |
+
bias=bias,
|
| 33 |
+
)
|
| 34 |
+
self.adaptive = adaptive
|
| 35 |
+
# print(self.adaptive)
|
| 36 |
+
if self.adaptive:
|
| 37 |
+
self.A = nn.Parameter(A.clone())
|
| 38 |
+
else:
|
| 39 |
+
self.register_buffer('A', A)
|
| 40 |
+
self.bn = nn.BatchNorm2d(out_channels)
|
| 41 |
+
self.relu = nn.ReLU(inplace=True)
|
| 42 |
+
|
| 43 |
+
def forward(self, x, len_x):
|
| 44 |
+
x = self.conv(x)
|
| 45 |
+
|
| 46 |
+
n, kc, t, v = x.size()
|
| 47 |
+
x = x.view(n, self.kernel_size, kc // self.kernel_size, t, v)
|
| 48 |
+
x = torch.einsum('nkctv,kvw->nctw', (x, self.A)).contiguous()
|
| 49 |
+
y = self.bn(x)
|
| 50 |
+
y = self.relu(y)
|
| 51 |
+
return y
|
| 52 |
+
|
| 53 |
+
class STGCN_block(nn.Module):
|
| 54 |
+
def __init__(
|
| 55 |
+
self,
|
| 56 |
+
in_channels,
|
| 57 |
+
out_channels,
|
| 58 |
+
kernel_size,
|
| 59 |
+
A,
|
| 60 |
+
adaptive=True,
|
| 61 |
+
stride=1,
|
| 62 |
+
dropout=0,
|
| 63 |
+
residual=True,
|
| 64 |
+
):
|
| 65 |
+
super().__init__()
|
| 66 |
+
|
| 67 |
+
assert len(kernel_size) == 2
|
| 68 |
+
assert kernel_size[0] % 2 == 1
|
| 69 |
+
padding = ((kernel_size[0] - 1) // 2, 0)
|
| 70 |
+
self.gcn = GCN_unit(
|
| 71 |
+
in_channels,
|
| 72 |
+
out_channels,
|
| 73 |
+
kernel_size[1],
|
| 74 |
+
A,
|
| 75 |
+
adaptive=adaptive,
|
| 76 |
+
)
|
| 77 |
+
if kernel_size[0] > 1:
|
| 78 |
+
self.tcn = nn.Sequential(
|
| 79 |
+
nn.Conv2d(
|
| 80 |
+
out_channels,
|
| 81 |
+
out_channels,
|
| 82 |
+
(kernel_size[0], 1),
|
| 83 |
+
(stride, 1),
|
| 84 |
+
padding,
|
| 85 |
+
),
|
| 86 |
+
nn.BatchNorm2d(out_channels),
|
| 87 |
+
nn.Dropout(dropout, inplace=True),
|
| 88 |
+
)
|
| 89 |
+
else:
|
| 90 |
+
self.tcn = nn.Identity()
|
| 91 |
+
|
| 92 |
+
if not residual:
|
| 93 |
+
self.residual = lambda x: 0
|
| 94 |
+
|
| 95 |
+
elif (in_channels == out_channels) and (stride == 1):
|
| 96 |
+
self.residual = lambda x: x
|
| 97 |
+
|
| 98 |
+
else:
|
| 99 |
+
self.residual = nn.Sequential(
|
| 100 |
+
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=(stride, 1)),
|
| 101 |
+
nn.BatchNorm2d(out_channels),
|
| 102 |
+
)
|
| 103 |
+
|
| 104 |
+
self.relu = nn.ReLU(inplace=True)
|
| 105 |
+
|
| 106 |
+
def forward(self, x, len_x=None):
|
| 107 |
+
res = self.residual(x)
|
| 108 |
+
x = self.gcn(x, len_x)
|
| 109 |
+
x = self.tcn(x) + res
|
| 110 |
+
return self.relu(x)
|
| 111 |
+
|
| 112 |
+
class STGCNChain(nn.Sequential):
|
| 113 |
+
def __init__(self, in_dim, block_args, kernel_size, A, adaptive):
|
| 114 |
+
super(STGCNChain, self).__init__()
|
| 115 |
+
last_dim = in_dim
|
| 116 |
+
for i, [channel, depth] in enumerate(block_args):
|
| 117 |
+
for j in range(depth):
|
| 118 |
+
self.add_module(f'layer{i}_{j}', STGCN_block(last_dim, channel, kernel_size, A.clone(), adaptive))
|
| 119 |
+
last_dim = channel
|
| 120 |
+
|
| 121 |
+
def get_stgcn_chain(in_dim, level, kernel_size, A, adaptive):
|
| 122 |
+
if level == 'spatial':
|
| 123 |
+
block_args = [[64,1], [128,1], [256,1]]
|
| 124 |
+
elif level == 'temporal':
|
| 125 |
+
block_args = [[256,3]]
|
| 126 |
+
else:
|
| 127 |
+
raise NotImplementedError
|
| 128 |
+
return STGCNChain(in_dim, block_args, kernel_size, A, adaptive), block_args[-1][0]
|
Uni_Sign/unisign_model/config.json
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_name_or_path": "/home/patrick/hugging_face/t5/mt5-base",
|
| 3 |
+
"architectures": [
|
| 4 |
+
"MT5ForConditionalGeneration"
|
| 5 |
+
],
|
| 6 |
+
"d_ff": 2048,
|
| 7 |
+
"d_kv": 64,
|
| 8 |
+
"d_model": 768,
|
| 9 |
+
"decoder_start_token_id": 0,
|
| 10 |
+
"dropout_rate": 0.1,
|
| 11 |
+
"eos_token_id": 1,
|
| 12 |
+
"feed_forward_proj": "gated-gelu",
|
| 13 |
+
"initializer_factor": 1.0,
|
| 14 |
+
"is_encoder_decoder": true,
|
| 15 |
+
"layer_norm_epsilon": 1e-06,
|
| 16 |
+
"model_type": "mt5",
|
| 17 |
+
"num_decoder_layers": 12,
|
| 18 |
+
"num_heads": 12,
|
| 19 |
+
"num_layers": 12,
|
| 20 |
+
"output_past": true,
|
| 21 |
+
"pad_token_id": 0,
|
| 22 |
+
"relative_attention_num_buckets": 32,
|
| 23 |
+
"tie_word_embeddings": false,
|
| 24 |
+
"tokenizer_class": "T5Tokenizer",
|
| 25 |
+
"transformers_version": "4.10.0.dev0",
|
| 26 |
+
"use_cache": true,
|
| 27 |
+
"vocab_size": 250112
|
| 28 |
+
}
|
Uni_Sign/unisign_model/generation_config.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_from_model_config": true,
|
| 3 |
+
"decoder_start_token_id": 0,
|
| 4 |
+
"eos_token_id": 1,
|
| 5 |
+
"pad_token_id": 0,
|
| 6 |
+
"transformers_version": "4.27.0.dev0"
|
| 7 |
+
}
|
Uni_Sign/unisign_model/special_tokens_map.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"eos_token": "</s>", "unk_token": "<unk>", "pad_token": "<pad>"}
|
Uni_Sign/unisign_model/spiece.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:ef78f86560d809067d12bac6c09f19a462cb3af3f54d2b8acbba26e1433125d6
|
| 3 |
+
size 4309802
|
Uni_Sign/unisign_model/tokenizer_config.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"eos_token": "</s>", "unk_token": "<unk>", "pad_token": "<pad>", "extra_ids": 0, "additional_special_tokens": null, "special_tokens_map_file": "/home/patrick/.cache/torch/transformers/685ac0ca8568ec593a48b61b0a3c272beee9bc194a3c7241d15dcadb5f875e53.f76030f3ec1b96a8199b2593390c610e76ca8028ef3d24680000619ffb646276", "tokenizer_file": null, "name_or_path": "google/mt5-small"}
|
checkpoints/pose/face_landmarker.task
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:64184e229b263107bc2b804c6625db1341ff2bb731874b0bcc2fe6544e0bc9ff
|
| 3 |
+
size 3758596
|
checkpoints/pose/hand_landmarker.task
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fbc2a30080c3c557093b5ddfc334698132eb341044ccee322ccf8bcf3607cde1
|
| 3 |
+
size 7819105
|
checkpoints/pose/pose_landmarker_full.task
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:4eaa5eb7a98365221087693fcc286334cf0858e2eb6e15b506aa4a7ecdcec4ad
|
| 3 |
+
size 9398198
|
checkpoints/pose/yolov8n-pose.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c6fa93dd1ee4a2c18c900a45c1d864a1c6f7aba75d84f91648a30b7fb641d212
|
| 3 |
+
size 6832633
|