Upload folder using huggingface_hub
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +42 -0
- =0.46.1 +25 -0
- Downloads/.ipynb_checkpoints/Convolutional Neural Networks Transfer Learning-checkpoint.ipynb +990 -0
- Downloads/.ipynb_checkpoints/EEEM071_CourseWork_ipynb-checkpoint.ipynb +379 -0
- Downloads/.ipynb_checkpoints/Human Action Recognition Tutorial-checkpoint.ipynb +820 -0
- Downloads/.ipynb_checkpoints/PyTorch Tutorial-checkpoint.ipynb +2030 -0
- Downloads/.ipynb_checkpoints/Python Tutorial(1)-checkpoint.ipynb +3313 -0
- Downloads/.ipynb_checkpoints/Python Tutorial-checkpoint.ipynb +3313 -0
- Downloads/.~lock.deformation_experiments(3).pptx# +1 -0
- Downloads/HjxNnnnu.html +0 -0
- Downloads/deformation_experiments(1).pptx +3 -0
- Downloads/deformation_experiments(2).pptx +3 -0
- Downloads/deformation_experiments(3).pptx +3 -0
- Downloads/deformation_experiments.pptx +3 -0
- Downloads/gap_zoomed_comparison.png +0 -0
- Downloads/geometric_solver(1).py +1147 -0
- Downloads/handover_summary(1).md +220 -0
- Downloads/handover_summary.md +220 -0
- Downloads/lbs_seam_constrained(1).py +222 -0
- Downloads/pipe_3.py +1465 -0
- Downloads/sketch_pipeline(1).pptx +0 -0
- Downloads/sketch_pipeline.pptx +0 -0
- LICENSE +21 -0
- P_1/band6/attempts/attempt_1/band6.png +3 -0
- P_1/band6/attempts/attempt_1/kf0.png +0 -0
- P_1/band6/attempts/attempt_1/kf1.png +0 -0
- P_1/band6/attempts/attempt_1/kf2.png +0 -0
- P_1/band6/attempts/attempt_1/kf3.png +0 -0
- P_1/band6/attempts/attempt_1/kf4.png +0 -0
- P_1/band6/attempts/attempt_2/band6.png +3 -0
- P_1/band6/attempts/attempt_2/kf0.png +0 -0
- P_1/band6/attempts/attempt_2/kf1.png +0 -0
- P_1/band6/attempts/attempt_2/kf2.png +0 -0
- P_1/band6/attempts/attempt_2/kf3.png +0 -0
- P_1/band6/attempts/attempt_2/kf4.png +0 -0
- P_1/band6/band6.png +3 -0
- P_1/band6/kf0.png +0 -0
- P_1/band6/kf1.png +0 -0
- P_1/band6/kf2.png +0 -0
- P_1/band6/kf3.png +0 -0
- P_1/band6/kf4.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_1/band6.png +3 -0
- P_1/band6__20260826_144524/attempts/attempt_1/kf0.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_1/kf1.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_1/kf2.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_1/kf3.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_1/kf4.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_2/band6.png +3 -0
- P_1/band6__20260826_144524/attempts/attempt_2/kf0.png +0 -0
- P_1/band6__20260826_144524/attempts/attempt_2/kf1.png +0 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,45 @@ 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 |
+
Downloads/deformation_experiments(1).pptx filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
Downloads/deformation_experiments(2).pptx filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
Downloads/deformation_experiments(3).pptx filter=lfs diff=lfs merge=lfs -text
|
| 39 |
+
Downloads/deformation_experiments.pptx filter=lfs diff=lfs merge=lfs -text
|
| 40 |
+
P_1/band6/attempts/attempt_1/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 41 |
+
P_1/band6/attempts/attempt_2/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 42 |
+
P_1/band6/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 43 |
+
P_1/band6__20260826_144524/attempts/attempt_1/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 44 |
+
P_1/band6__20260826_144524/attempts/attempt_2/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 45 |
+
P_1/band6__20260826_144524/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 46 |
+
P_1/band6__20260826_150549/attempts/attempt_1/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 47 |
+
P_1/band6__20260826_150549/attempts/attempt_2/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 48 |
+
P_1/band6__20260826_150549/attempts/attempt_3/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 49 |
+
P_1/band6__20260826_150549/attempts/attempt_4/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 50 |
+
P_1/band6__20260826_150549/attempts/attempt_5/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 51 |
+
P_1/band6__20260826_150549/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 52 |
+
P_1/band6__20260826_152911/attempts/attempt_1/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 53 |
+
P_1/band6__20260826_152911/attempts/attempt_2/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 54 |
+
P_1/band6__20260826_152911/attempts/attempt_3/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 55 |
+
P_1/band6__20260826_152911/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 56 |
+
P_1/band6__20260826_153547/attempts/attempt_1/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 57 |
+
P_1/band6__20260826_153547/attempts/attempt_2/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 58 |
+
P_1/band6__20260826_153547/attempts/attempt_3/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 59 |
+
P_1/band6__20260826_153547/band6.png filter=lfs diff=lfs merge=lfs -text
|
| 60 |
+
P_1/basketball5/attempts/attempt_1/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 61 |
+
P_1/basketball5/attempts/attempt_2/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 62 |
+
P_1/basketball5/attempts/attempt_3/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 63 |
+
P_1/basketball5/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 64 |
+
P_1/basketball5__20260826_140014/attempts/attempt_1/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 65 |
+
P_1/basketball5__20260826_140014/attempts/attempt_2/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 66 |
+
P_1/basketball5__20260826_140014/attempts/attempt_3/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 67 |
+
P_1/basketball5__20260826_140014/attempts/attempt_4/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 68 |
+
P_1/basketball5__20260826_140014/attempts/attempt_5/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 69 |
+
P_1/basketball5__20260826_140014/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 70 |
+
P_1/basketball5__20260826_142702/attempts/attempt_1/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 71 |
+
P_1/basketball5__20260826_142702/attempts/attempt_2/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 72 |
+
P_1/basketball5__20260826_142702/basketball5.png filter=lfs diff=lfs merge=lfs -text
|
| 73 |
+
P_1/cat1__20260826_163052/attempts/attempt_1/cat1.png filter=lfs diff=lfs merge=lfs -text
|
| 74 |
+
P_1/cat1__20260826_163052/attempts/attempt_2/cat1.png filter=lfs diff=lfs merge=lfs -text
|
| 75 |
+
P_1/cat1__20260826_163052/attempts/attempt_3/cat1.png filter=lfs diff=lfs merge=lfs -text
|
| 76 |
+
P_1/cat1__20260826_163052/attempts/attempt_4/cat1.png filter=lfs diff=lfs merge=lfs -text
|
| 77 |
+
P_1/cat1__20260826_163052/attempts/attempt_5/cat1.png filter=lfs diff=lfs merge=lfs -text
|
=0.46.1
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Requirement already satisfied: bitsandbytes in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (0.50.1)
|
| 2 |
+
Requirement already satisfied: torch<3,>=2.4 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from bitsandbytes) (2.5.1+cu121)
|
| 3 |
+
Requirement already satisfied: numpy>=1.17 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from bitsandbytes) (2.2.6)
|
| 4 |
+
Requirement already satisfied: packaging>=20.9 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from bitsandbytes) (26.0)
|
| 5 |
+
Requirement already satisfied: filelock in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (3.29.0)
|
| 6 |
+
Requirement already satisfied: typing-extensions>=4.8.0 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (4.15.0)
|
| 7 |
+
Requirement already satisfied: networkx in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (3.4.2)
|
| 8 |
+
Requirement already satisfied: jinja2 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (3.1.6)
|
| 9 |
+
Requirement already satisfied: fsspec in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (2026.4.0)
|
| 10 |
+
Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.1.105 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (12.1.105)
|
| 11 |
+
Requirement already satisfied: nvidia-cuda-runtime-cu12==12.1.105 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (12.1.105)
|
| 12 |
+
Requirement already satisfied: nvidia-cuda-cupti-cu12==12.1.105 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (12.1.105)
|
| 13 |
+
Requirement already satisfied: nvidia-cudnn-cu12==9.1.0.70 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (9.1.0.70)
|
| 14 |
+
Requirement already satisfied: nvidia-cublas-cu12==12.1.3.1 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (12.1.3.1)
|
| 15 |
+
Requirement already satisfied: nvidia-cufft-cu12==11.0.2.54 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (11.0.2.54)
|
| 16 |
+
Requirement already satisfied: nvidia-curand-cu12==10.3.2.106 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (10.3.2.106)
|
| 17 |
+
Requirement already satisfied: nvidia-cusolver-cu12==11.4.5.107 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (11.4.5.107)
|
| 18 |
+
Requirement already satisfied: nvidia-cusparse-cu12==12.1.0.106 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (12.1.0.106)
|
| 19 |
+
Requirement already satisfied: nvidia-nccl-cu12==2.21.5 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (2.21.5)
|
| 20 |
+
Requirement already satisfied: nvidia-nvtx-cu12==12.1.105 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (12.1.105)
|
| 21 |
+
Requirement already satisfied: triton==3.1.0 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (3.1.0)
|
| 22 |
+
Requirement already satisfied: sympy==1.13.1 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from torch<3,>=2.4->bitsandbytes) (1.13.1)
|
| 23 |
+
Requirement already satisfied: nvidia-nvjitlink-cu12 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from nvidia-cusolver-cu12==11.4.5.107->torch<3,>=2.4->bitsandbytes) (12.9.86)
|
| 24 |
+
Requirement already satisfied: mpmath<1.4,>=1.1.0 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from sympy==1.13.1->torch<3,>=2.4->bitsandbytes) (1.3.0)
|
| 25 |
+
Requirement already satisfied: MarkupSafe>=2.0 in /scratch/rk01499/anaconda3/envs/sketch/lib/python3.10/site-packages (from jinja2->torch<3,>=2.4->bitsandbytes) (3.0.3)
|
Downloads/.ipynb_checkpoints/Convolutional Neural Networks Transfer Learning-checkpoint.ipynb
ADDED
|
@@ -0,0 +1,990 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {
|
| 6 |
+
"colab_type": "text",
|
| 7 |
+
"id": "view-in-github"
|
| 8 |
+
},
|
| 9 |
+
"source": [
|
| 10 |
+
"<a href=\"https://colab.research.google.com/github/AnjanDutta/EEEM068/blob/main/Notebooks/Convolutional_Neural_Networks_Transfer_Learning.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
|
| 11 |
+
]
|
| 12 |
+
},
|
| 13 |
+
{
|
| 14 |
+
"cell_type": "markdown",
|
| 15 |
+
"metadata": {
|
| 16 |
+
"id": "D_5Tkl9SzhjN",
|
| 17 |
+
"pycharm": {}
|
| 18 |
+
},
|
| 19 |
+
"source": [
|
| 20 |
+
"<H1 style=\"text-align: center\">EEEM068 - Applied Machine Learning</H1>\n",
|
| 21 |
+
"<H1 style=\"text-align: center\">Workshop 04</H1>\n",
|
| 22 |
+
"<H1 style=\"text-align: center\">Convolutional Neural Networks and Transfer Learning Tutorial</H1>"
|
| 23 |
+
]
|
| 24 |
+
},
|
| 25 |
+
{
|
| 26 |
+
"cell_type": "markdown",
|
| 27 |
+
"metadata": {
|
| 28 |
+
"id": "e8f95r3cyd1S",
|
| 29 |
+
"pycharm": {}
|
| 30 |
+
},
|
| 31 |
+
"source": [
|
| 32 |
+
"## Introduction\n",
|
| 33 |
+
"In this tutorial, we will implement a convolutional neural network (CNN) model for classifying natural images. Specifically, we will use the STL-10 dataset for training and testing our model. In this workshop, we will use [PyTorch](https://pytorch.org/) deep learning framework to complete our task."
|
| 34 |
+
]
|
| 35 |
+
},
|
| 36 |
+
{
|
| 37 |
+
"cell_type": "markdown",
|
| 38 |
+
"metadata": {
|
| 39 |
+
"id": "LT-4QzaoAxQX",
|
| 40 |
+
"pycharm": {}
|
| 41 |
+
},
|
| 42 |
+
"source": [
|
| 43 |
+
"## STL-10 Dataset\n",
|
| 44 |
+
"\n",
|
| 45 |
+
"The [STL-10 dataset](https://cs.stanford.edu/~acoates/stl10/) is an image recognition dataset for developing supervised and unsupervised deep learning algorithms. It contains 10 classes: airplane, bird, car, cat, deer, dog, horse, monkey, ship, truck, containing 500 training and 800 test images per class. Each image is of size $96 \\times 96$ pixels. More detials on the STL-10 dataset can be found in here: https://cs.stanford.edu/~acoates/stl10\n",
|
| 46 |
+
"\n",
|
| 47 |
+
"<img src=\"https://cs.stanford.edu/~acoates/stl10/images.png\" width=\"400\" height=\"400\">\n",
|
| 48 |
+
"\n",
|
| 49 |
+
"Similar to MNIST, since this dataset is already implemented at the torchvision [dataset collections](https://pytorch.org/vision/stable/index.html), we don't have to implement the data generator for this dataset and will utilise the one available from torchvision. However, an example of custon data generator can be found within the transfer learning section of this tutorial."
|
| 50 |
+
]
|
| 51 |
+
},
|
| 52 |
+
{
|
| 53 |
+
"cell_type": "markdown",
|
| 54 |
+
"metadata": {
|
| 55 |
+
"id": "YDSiTaVKSXJW"
|
| 56 |
+
},
|
| 57 |
+
"source": [
|
| 58 |
+
"### Dataset and DataLoader"
|
| 59 |
+
]
|
| 60 |
+
},
|
| 61 |
+
{
|
| 62 |
+
"cell_type": "markdown",
|
| 63 |
+
"metadata": {
|
| 64 |
+
"id": "7Bi1u6k5P9wB"
|
| 65 |
+
},
|
| 66 |
+
"source": [
|
| 67 |
+
"In the following cell, we will be defining datasets and data loaders necessary for our training. Details on datasets and dataloaders can be found in the [documentation](https://pytorch.org/vision/stable/datasets.html)."
|
| 68 |
+
]
|
| 69 |
+
},
|
| 70 |
+
{
|
| 71 |
+
"cell_type": "code",
|
| 72 |
+
"execution_count": null,
|
| 73 |
+
"metadata": {
|
| 74 |
+
"id": "C7KkXL5OA8WX",
|
| 75 |
+
"pycharm": {
|
| 76 |
+
"is_executing": true
|
| 77 |
+
}
|
| 78 |
+
},
|
| 79 |
+
"outputs": [],
|
| 80 |
+
"source": [
|
| 81 |
+
"import torch\n",
|
| 82 |
+
"import torchvision\n",
|
| 83 |
+
"\n",
|
| 84 |
+
"# Before defining datasets, lets define how images should be transformed. This is \n",
|
| 85 |
+
"# because the transformations should go with the definitions of datasets. In this\n",
|
| 86 |
+
"# tutorial we will using simple transformations, such as (1) image to tensor, (2)\n",
|
| 87 |
+
"# normalization.\n",
|
| 88 |
+
"image_transform = torchvision.transforms.Compose([\n",
|
| 89 |
+
" torchvision.transforms.ToTensor(),\n",
|
| 90 |
+
" torchvision.transforms.Normalize(\n",
|
| 91 |
+
" (0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])\n",
|
| 92 |
+
"\n",
|
| 93 |
+
"# Once we have the transformations defined, lets define the train and test sets\n",
|
| 94 |
+
"train_dataset = torchvision.datasets.STL10('dataset/', \n",
|
| 95 |
+
" split='train', \n",
|
| 96 |
+
" download=True,\n",
|
| 97 |
+
" transform=image_transform)\n",
|
| 98 |
+
"test_dataset = torchvision.datasets.STL10('dataset/', \n",
|
| 99 |
+
" split='test', \n",
|
| 100 |
+
" download=True,\n",
|
| 101 |
+
" transform=image_transform)\n",
|
| 102 |
+
"\n",
|
| 103 |
+
"# Now, lets define batch size, batch size is how much data you feed for training\n",
|
| 104 |
+
"# in one iteration\n",
|
| 105 |
+
"batch_size_train = 256 # We use smaller batch size here for training\n",
|
| 106 |
+
"batch_size_test = 1024 # We use bigger batch size for testing\n",
|
| 107 |
+
"\n",
|
| 108 |
+
"# Once we have the datasets defined, lets define the data loaders as follows\n",
|
| 109 |
+
"train_loader = torch.utils.data.DataLoader(train_dataset,\n",
|
| 110 |
+
" batch_size=batch_size_train, \n",
|
| 111 |
+
" shuffle=True)\n",
|
| 112 |
+
"test_loader = torch.utils.data.DataLoader(test_dataset,\n",
|
| 113 |
+
" batch_size=batch_size_test, \n",
|
| 114 |
+
" shuffle=True)"
|
| 115 |
+
]
|
| 116 |
+
},
|
| 117 |
+
{
|
| 118 |
+
"cell_type": "markdown",
|
| 119 |
+
"metadata": {
|
| 120 |
+
"id": "WdYlqJ7xDvJa",
|
| 121 |
+
"pycharm": {}
|
| 122 |
+
},
|
| 123 |
+
"source": [
|
| 124 |
+
"### Example Image\n",
|
| 125 |
+
"Lets have a look on how images from STL-10 dataset looks like. Since the images are already normalized, their resolutions might have slightly changed. To visualize the original images, we should have ideally apply a reverse transformation which is avoided to keep this tutorial simple and brief."
|
| 126 |
+
]
|
| 127 |
+
},
|
| 128 |
+
{
|
| 129 |
+
"cell_type": "code",
|
| 130 |
+
"execution_count": null,
|
| 131 |
+
"metadata": {
|
| 132 |
+
"id": "U4FrTxdxDwnE",
|
| 133 |
+
"pycharm": {}
|
| 134 |
+
},
|
| 135 |
+
"outputs": [],
|
| 136 |
+
"source": [
|
| 137 |
+
"# import plot library\n",
|
| 138 |
+
"import matplotlib.pyplot as plt\n",
|
| 139 |
+
"# iterate the dataloader\n",
|
| 140 |
+
"_, (example_datas, labels) = next(enumerate(train_loader))\n",
|
| 141 |
+
"# get the first data\n",
|
| 142 |
+
"sample = example_datas[0]\n",
|
| 143 |
+
"# show the data\n",
|
| 144 |
+
"plt.imshow(sample.permute(1, 2, 0))\n",
|
| 145 |
+
"print(\"Label: \" + str(labels[0]))"
|
| 146 |
+
]
|
| 147 |
+
},
|
| 148 |
+
{
|
| 149 |
+
"cell_type": "markdown",
|
| 150 |
+
"metadata": {
|
| 151 |
+
"id": "MUWGhlsMCYZD",
|
| 152 |
+
"pycharm": {}
|
| 153 |
+
},
|
| 154 |
+
"source": [
|
| 155 |
+
"## Model\n",
|
| 156 |
+
"Now, we have to define trainable layers with parameters and put them inside a model. Have a look on the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Module.html#module) of `nn.Module` and read more about different layers and functionalities of PyTorch there. Here we are going to implement various versions of AlexNet model and use it for classification. In this model, we are going to use the following functions or modules:\n",
|
| 157 |
+
"\n",
|
| 158 |
+
"* `nn.Conv2d()`: It is a PyTorch module that applies a 2D convolution over an input signal composed of several input planes. More details are available on the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Conv2d.html).\n",
|
| 159 |
+
"\n",
|
| 160 |
+
"* `nn.MaxPool2d()`: It is also a module that applies a 2D max pooling over an input signal composed of several input planes. Please have a look on this [documentation](https://pytorch.org/docs/stable/generated/torch.nn.MaxPool2d.html) for more details.\n",
|
| 161 |
+
"\n",
|
| 162 |
+
"* `nn.AdaptiveAvgPool2d()`: It is a module that applies a 2D adaptive average pooling over an input signal composed of several input planes. Given an output size, this function automatically select the stride and kernel size to adapt the need of target size. More details on this can be found in the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.AdaptiveAvgPool2d.html).\n",
|
| 163 |
+
"\n",
|
| 164 |
+
"* `nn.Sequential()`: It is a sequential container. Modules will be added to it in the order they are passed in the constructor. Please check the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Sequential.html#torch.nn.Sequential) for more details.\n",
|
| 165 |
+
"\n",
|
| 166 |
+
"* `nn.Linear()`: It is a module that applies a linear transformation to the incoming data. More details can be found in its [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Linear.html#linear).\n",
|
| 167 |
+
"\n",
|
| 168 |
+
"* `nn.ReLU()`: It is also a module that applies element-wise the rectified linear unit function. Its [documentation](https://pytorch.org/docs/stable/generated/torch.nn.ReLU.html#relu) can explain more.\n",
|
| 169 |
+
"\n",
|
| 170 |
+
"* `nn.Dropout()`: This module randomly zeroes some of the elements of the input tensor with probability `p`. Check the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Dropout.html#dropout) for more details."
|
| 171 |
+
]
|
| 172 |
+
},
|
| 173 |
+
{
|
| 174 |
+
"cell_type": "markdown",
|
| 175 |
+
"metadata": {
|
| 176 |
+
"id": "yjGEol-IX31p"
|
| 177 |
+
},
|
| 178 |
+
"source": [
|
| 179 |
+
"One can define a model in several ways. Below, we show some of them."
|
| 180 |
+
]
|
| 181 |
+
},
|
| 182 |
+
{
|
| 183 |
+
"cell_type": "code",
|
| 184 |
+
"execution_count": null,
|
| 185 |
+
"metadata": {
|
| 186 |
+
"id": "LMtPp5OeCakG",
|
| 187 |
+
"pycharm": {}
|
| 188 |
+
},
|
| 189 |
+
"outputs": [],
|
| 190 |
+
"source": [
|
| 191 |
+
"## We first import the pytorch nn module and optimizer\n",
|
| 192 |
+
"import torch.nn as nn\n",
|
| 193 |
+
"import torch.nn.functional as F\n",
|
| 194 |
+
"import torch.optim as optim\n",
|
| 195 |
+
"## Below you can see one way of defining the model class, where each individual \n",
|
| 196 |
+
"## layer is defined as an instance variable.\n",
|
| 197 |
+
"class AlexNet1(nn.Module):\n",
|
| 198 |
+
" def __init__(self, num_classes):\n",
|
| 199 |
+
" super(AlexNet1, self).__init__()\n",
|
| 200 |
+
" # input channel 3, output channel 64\n",
|
| 201 |
+
" self.conv1 = nn.Conv2d(3, 64, kernel_size=11, stride=4, padding=2)\n",
|
| 202 |
+
" # relu non-linearity\n",
|
| 203 |
+
" self.relu1 = nn.ReLU()\n",
|
| 204 |
+
" # max pooling\n",
|
| 205 |
+
" self.max_pool2d1 = nn.MaxPool2d(kernel_size=3, stride=2)\n",
|
| 206 |
+
" # input channel 64, output channel 192\n",
|
| 207 |
+
" self.conv2 = nn.Conv2d(64, 192, kernel_size=5, stride=1, padding=2)\n",
|
| 208 |
+
" self.relu2 = nn.ReLU()\n",
|
| 209 |
+
" self.max_pool2d2 = nn.MaxPool2d(kernel_size=3, stride=2)\n",
|
| 210 |
+
" # input channel 192, output channel 384\n",
|
| 211 |
+
" self.conv3 = nn.Conv2d(192, 384, kernel_size=3, stride=1, padding=1)\n",
|
| 212 |
+
" self.relu3 = nn.ReLU()\n",
|
| 213 |
+
" # input channel 384, output channel 256\n",
|
| 214 |
+
" self.conv4 = nn.Conv2d(384, 256, kernel_size=3, stride=1, padding=1)\n",
|
| 215 |
+
" self.relu4 = nn.ReLU()\n",
|
| 216 |
+
" # input channel 256, output channel 256\n",
|
| 217 |
+
" self.conv5 = nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1)\n",
|
| 218 |
+
" self.relu5 = nn.ReLU()\n",
|
| 219 |
+
" self.max_pool2d5 = nn.MaxPool2d(kernel_size=3, stride=2)\n",
|
| 220 |
+
" # adaptive pooling\n",
|
| 221 |
+
" self.adapt_pool = nn.AdaptiveAvgPool2d(output_size=(6, 6))\n",
|
| 222 |
+
" #dropout layer\n",
|
| 223 |
+
" self.dropout1 = nn.Dropout()\n",
|
| 224 |
+
" # linear layer\n",
|
| 225 |
+
" self.linear1 = nn.Linear(in_features=9216, out_features=4096, bias=True)\n",
|
| 226 |
+
" self.relu6 = nn.ReLU()\n",
|
| 227 |
+
" self.dropout2 = nn.Dropout()\n",
|
| 228 |
+
" self.linear2 = nn.Linear(in_features=4096, out_features=4096, bias=True)\n",
|
| 229 |
+
" self.relu7 = nn.ReLU()\n",
|
| 230 |
+
" self.linear3 = nn.Linear(in_features=4096, out_features=num_classes, bias=True)\n",
|
| 231 |
+
"\n",
|
| 232 |
+
" def forward(self, x):\n",
|
| 233 |
+
" x = self.conv1(x)\n",
|
| 234 |
+
" x = self.relu1(x)\n",
|
| 235 |
+
" x = self.max_pool2d1(x)\n",
|
| 236 |
+
" x = self.conv2(x)\n",
|
| 237 |
+
" x = self.relu2(x)\n",
|
| 238 |
+
" x = self.max_pool2d2(x)\n",
|
| 239 |
+
" x = self.conv3(x)\n",
|
| 240 |
+
" x = self.relu3(x)\n",
|
| 241 |
+
" x = self.conv4(x)\n",
|
| 242 |
+
" x = self.relu4(x)\n",
|
| 243 |
+
" x = self.conv5(x)\n",
|
| 244 |
+
" x = self.relu5(x)\n",
|
| 245 |
+
" x = self.max_pool2d5(x)\n",
|
| 246 |
+
" x = self.adapt_pool(x)\n",
|
| 247 |
+
" # Note how we are flattening the feature map, B x C x H x W -> B x C*H*W\n",
|
| 248 |
+
" x = x.reshape(x.shape[0], -1)\n",
|
| 249 |
+
" x = self.dropout1(x)\n",
|
| 250 |
+
" x = self.linear1(x)\n",
|
| 251 |
+
" x = self.relu6(x)\n",
|
| 252 |
+
" x = self.dropout2(x)\n",
|
| 253 |
+
" x = self.linear2(x)\n",
|
| 254 |
+
" x = self.relu7(x)\n",
|
| 255 |
+
" x = self.linear3(x)\n",
|
| 256 |
+
" return x\n",
|
| 257 |
+
"\n",
|
| 258 |
+
"## Below you can see another way of defining the model class, where some layers \n",
|
| 259 |
+
"## together are defined as an instance variable.\n",
|
| 260 |
+
"class AlexNet2(nn.Module):\n",
|
| 261 |
+
" def __init__(self, num_classes):\n",
|
| 262 |
+
" super(AlexNet2, self).__init__()\n",
|
| 263 |
+
" self.features = nn.Sequential(\n",
|
| 264 |
+
" nn.Conv2d(3, 64, kernel_size=11, stride=4, padding=2),\n",
|
| 265 |
+
" nn.ReLU(),\n",
|
| 266 |
+
" nn.MaxPool2d(kernel_size=3, stride=2),\n",
|
| 267 |
+
" nn.Conv2d(64, 192, kernel_size=5, stride=1, padding=2),\n",
|
| 268 |
+
" nn.ReLU(),\n",
|
| 269 |
+
" nn.MaxPool2d(kernel_size=3, stride=2),\n",
|
| 270 |
+
" nn.Conv2d(192, 384, kernel_size=3, stride=1, padding=1),\n",
|
| 271 |
+
" nn.ReLU(),\n",
|
| 272 |
+
" nn.Conv2d(384, 256, kernel_size=3, stride=1, padding=1),\n",
|
| 273 |
+
" nn.ReLU(),\n",
|
| 274 |
+
" nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1),\n",
|
| 275 |
+
" nn.ReLU(),\n",
|
| 276 |
+
" nn.MaxPool2d(kernel_size=3, stride=2),\n",
|
| 277 |
+
" nn.AdaptiveAvgPool2d(output_size=(6, 6))\n",
|
| 278 |
+
" )\n",
|
| 279 |
+
" self.classifier = nn.Sequential(\n",
|
| 280 |
+
" nn.Dropout(),\n",
|
| 281 |
+
" nn.Linear(in_features=9216, out_features=4096, bias=True),\n",
|
| 282 |
+
" nn.ReLU(),\n",
|
| 283 |
+
" nn.Dropout(),\n",
|
| 284 |
+
" nn.Linear(in_features=4096, out_features=4096, bias=True),\n",
|
| 285 |
+
" nn.ReLU(),\n",
|
| 286 |
+
" nn.Linear(in_features=4096, out_features=num_classes, bias=True)\n",
|
| 287 |
+
" )\n",
|
| 288 |
+
"\n",
|
| 289 |
+
" def forward(self, x):\n",
|
| 290 |
+
" x = self.features(x)\n",
|
| 291 |
+
" # Note how we are flattening the feature map, B x C x H x W -> B x C*H*W\n",
|
| 292 |
+
" x = x.reshape(x.shape[0], -1)\n",
|
| 293 |
+
" x = self.classifier(x)\n",
|
| 294 |
+
" return x\n",
|
| 295 |
+
"\n",
|
| 296 |
+
"## Below we define the model as defined in the torchvision package and haven't \n",
|
| 297 |
+
"## initialised with pretrained weights (see pretrained=False flag)\n",
|
| 298 |
+
"class AlexNet3(nn.Module):\n",
|
| 299 |
+
" def __init__(self, num_classes):\n",
|
| 300 |
+
" super(AlexNet3, self).__init__()\n",
|
| 301 |
+
" from torchvision import models\n",
|
| 302 |
+
" alexnet = models.alexnet(weights=None)\n",
|
| 303 |
+
" self.features = alexnet.features\n",
|
| 304 |
+
" self.avgpool = alexnet.avgpool\n",
|
| 305 |
+
" self.classifier = alexnet.classifier\n",
|
| 306 |
+
" # Please note how to change the last layer of the classifier for a new dataset\n",
|
| 307 |
+
" # ImageNet-1K has 1000 classes, but STL-10 has 10 classes\n",
|
| 308 |
+
" self.classifier[6] = nn.Linear(in_features=4096, out_features=num_classes, bias=True)\n",
|
| 309 |
+
"\n",
|
| 310 |
+
" def forward(self, x):\n",
|
| 311 |
+
" x = self.features(x)\n",
|
| 312 |
+
" x = self.avgpool(x)\n",
|
| 313 |
+
" # Note how we are flattening the feature map, B x C x H x W -> B x C*H*W\n",
|
| 314 |
+
" x = x.reshape(x.shape[0], -1)\n",
|
| 315 |
+
" x = self.classifier(x)\n",
|
| 316 |
+
" return x\n",
|
| 317 |
+
"\n",
|
| 318 |
+
"## Below we define the model as defined in the torchvision package and initialised \n",
|
| 319 |
+
"## with pretrained weights (see pretrained=True flag)\n",
|
| 320 |
+
"class AlexNet4(nn.Module):\n",
|
| 321 |
+
" def __init__(self, num_classes):\n",
|
| 322 |
+
" super(AlexNet4, self).__init__()\n",
|
| 323 |
+
" from torchvision import models\n",
|
| 324 |
+
" alexnet = models.alexnet(weights='IMAGENET1K_V1')\n",
|
| 325 |
+
" self.features = alexnet.features\n",
|
| 326 |
+
" self.avgpool = alexnet.avgpool\n",
|
| 327 |
+
" self.classifier = alexnet.classifier\n",
|
| 328 |
+
" # Please note how to change the last layer of the classifier for a new dataset\n",
|
| 329 |
+
" # ImageNet-1K has 1000 classes, but STL-10 has 10 classes\n",
|
| 330 |
+
" self.classifier[6] = nn.Linear(in_features=4096, out_features=num_classes, bias=True)\n",
|
| 331 |
+
"\n",
|
| 332 |
+
" def forward(self, x):\n",
|
| 333 |
+
" x = self.features(x)\n",
|
| 334 |
+
" x = self.avgpool(x)\n",
|
| 335 |
+
" # Note how we are flattening the feature map, B x C x H x W -> B x C*H*W\n",
|
| 336 |
+
" x = x.reshape(x.shape[0], -1)\n",
|
| 337 |
+
" x = self.classifier(x)\n",
|
| 338 |
+
" return x"
|
| 339 |
+
]
|
| 340 |
+
},
|
| 341 |
+
{
|
| 342 |
+
"cell_type": "markdown",
|
| 343 |
+
"metadata": {
|
| 344 |
+
"id": "nayicPkJCkWy",
|
| 345 |
+
"pycharm": {}
|
| 346 |
+
},
|
| 347 |
+
"source": [
|
| 348 |
+
"## Initialization\n",
|
| 349 |
+
"Once we have the model defined, lets instantiate it and set other hyperparameters."
|
| 350 |
+
]
|
| 351 |
+
},
|
| 352 |
+
{
|
| 353 |
+
"cell_type": "markdown",
|
| 354 |
+
"metadata": {
|
| 355 |
+
"id": "Y6YBtqhvZGwG"
|
| 356 |
+
},
|
| 357 |
+
"source": [
|
| 358 |
+
"#### Model\n",
|
| 359 |
+
"We will initialize the model, transfer to the desired device and set the parameters to receive gradients."
|
| 360 |
+
]
|
| 361 |
+
},
|
| 362 |
+
{
|
| 363 |
+
"cell_type": "code",
|
| 364 |
+
"execution_count": null,
|
| 365 |
+
"metadata": {
|
| 366 |
+
"id": "TjyEGZSdCk_i",
|
| 367 |
+
"pycharm": {}
|
| 368 |
+
},
|
| 369 |
+
"outputs": [],
|
| 370 |
+
"source": [
|
| 371 |
+
"# define the model, we could use any of the models AlexNet1, AlexNet2, AlexNet3, AlexNet4 \n",
|
| 372 |
+
"model = AlexNet2(10) # since STL-10 dataset has 10 classes, we set num_classes = 10\n",
|
| 373 |
+
"# device: cuda (gpu) or cpu\n",
|
| 374 |
+
"device = \"cuda\"\n",
|
| 375 |
+
"# map to device\n",
|
| 376 |
+
"model = model.to(device) # `model.cuda()` will also do the same job\n",
|
| 377 |
+
"# make the parameters trainable\n",
|
| 378 |
+
"for param in model.parameters():\n",
|
| 379 |
+
" param.requires_grad = True"
|
| 380 |
+
]
|
| 381 |
+
},
|
| 382 |
+
{
|
| 383 |
+
"cell_type": "markdown",
|
| 384 |
+
"metadata": {
|
| 385 |
+
"id": "CETvGvW8Y5-U"
|
| 386 |
+
},
|
| 387 |
+
"source": [
|
| 388 |
+
"#### Optimizer\n",
|
| 389 |
+
"For updating the parameters, PyTorch provides the package torch.optim that has most popular optimizers implemented. In this tutorial, we will be using the `torch.optim.Adam` optimizer.\n"
|
| 390 |
+
]
|
| 391 |
+
},
|
| 392 |
+
{
|
| 393 |
+
"cell_type": "code",
|
| 394 |
+
"execution_count": null,
|
| 395 |
+
"metadata": {
|
| 396 |
+
"id": "jEBBRoh-Y-bU"
|
| 397 |
+
},
|
| 398 |
+
"outputs": [],
|
| 399 |
+
"source": [
|
| 400 |
+
"import torch.optim as optim\n",
|
| 401 |
+
"## some hyperparameters related to optimizer\n",
|
| 402 |
+
"learning_rate = 0.0001\n",
|
| 403 |
+
"weight_decay = 0.0005\n",
|
| 404 |
+
"# define optimizer\n",
|
| 405 |
+
"optimizer = optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)"
|
| 406 |
+
]
|
| 407 |
+
},
|
| 408 |
+
{
|
| 409 |
+
"cell_type": "markdown",
|
| 410 |
+
"metadata": {
|
| 411 |
+
"id": "-RkqO0VcZuRP",
|
| 412 |
+
"pycharm": {}
|
| 413 |
+
},
|
| 414 |
+
"source": [
|
| 415 |
+
"## Average Meter\n",
|
| 416 |
+
"It is a simple class for keeping training statistics, such as losses and accuracies etc. The `.val` field usually holds the statistics for the current batch, whereas the `.avg` field hold statistics for the current epoch."
|
| 417 |
+
]
|
| 418 |
+
},
|
| 419 |
+
{
|
| 420 |
+
"cell_type": "code",
|
| 421 |
+
"execution_count": null,
|
| 422 |
+
"metadata": {
|
| 423 |
+
"id": "JeLH7fbOHDhH",
|
| 424 |
+
"pycharm": {}
|
| 425 |
+
},
|
| 426 |
+
"outputs": [],
|
| 427 |
+
"source": [
|
| 428 |
+
"class AverageMeter(object):\n",
|
| 429 |
+
" \"\"\"Computes and stores the average and current value\"\"\"\n",
|
| 430 |
+
" def __init__(self):\n",
|
| 431 |
+
" self.reset()\n",
|
| 432 |
+
"\n",
|
| 433 |
+
" def reset(self):\n",
|
| 434 |
+
" self.val = 0\n",
|
| 435 |
+
" self.avg = 0\n",
|
| 436 |
+
" self.sum = 0\n",
|
| 437 |
+
" self.count = 0\n",
|
| 438 |
+
"\n",
|
| 439 |
+
" def update(self, val, n=1):\n",
|
| 440 |
+
" self.val = val\n",
|
| 441 |
+
" self.sum += val * n\n",
|
| 442 |
+
" self.count += n\n",
|
| 443 |
+
" self.avg = self.sum / self.count"
|
| 444 |
+
]
|
| 445 |
+
},
|
| 446 |
+
{
|
| 447 |
+
"cell_type": "markdown",
|
| 448 |
+
"metadata": {
|
| 449 |
+
"id": "CSk4VfO_C4tL",
|
| 450 |
+
"pycharm": {}
|
| 451 |
+
},
|
| 452 |
+
"source": [
|
| 453 |
+
"## Train and Test Functions"
|
| 454 |
+
]
|
| 455 |
+
},
|
| 456 |
+
{
|
| 457 |
+
"cell_type": "code",
|
| 458 |
+
"execution_count": null,
|
| 459 |
+
"metadata": {
|
| 460 |
+
"id": "NJeHMF_BC7bg",
|
| 461 |
+
"pycharm": {}
|
| 462 |
+
},
|
| 463 |
+
"outputs": [],
|
| 464 |
+
"source": [
|
| 465 |
+
"from tqdm.notebook import tqdm\n",
|
| 466 |
+
"##define train function\n",
|
| 467 |
+
"def train(model, device, train_loader, optimizer):\n",
|
| 468 |
+
" # meter\n",
|
| 469 |
+
" loss = AverageMeter()\n",
|
| 470 |
+
" # switch to train mode\n",
|
| 471 |
+
" model.train()\n",
|
| 472 |
+
" tk0 = tqdm(train_loader, total=int(len(train_loader)))\n",
|
| 473 |
+
" for batch_idx, (data, target) in enumerate(tk0):\n",
|
| 474 |
+
" # after fetching the data transfer the model to the \n",
|
| 475 |
+
" # required device, in this example the device is gpu\n",
|
| 476 |
+
" # transfer to gpu can also be done by \n",
|
| 477 |
+
" # data, target = data.cuda(), target.cuda()\n",
|
| 478 |
+
" data, target = data.to(device), target.to(device) \n",
|
| 479 |
+
" # compute the forward pass\n",
|
| 480 |
+
" # it can also be achieved by model.forward(data)\n",
|
| 481 |
+
" output = model(data) \n",
|
| 482 |
+
" # compute the loss function\n",
|
| 483 |
+
" loss_this = F.cross_entropy(output, target)\n",
|
| 484 |
+
" # initialize the optimizer\n",
|
| 485 |
+
" optimizer.zero_grad()\n",
|
| 486 |
+
" # compute the backward pass\n",
|
| 487 |
+
" loss_this.backward()\n",
|
| 488 |
+
" # update the parameters\n",
|
| 489 |
+
" optimizer.step()\n",
|
| 490 |
+
" # update the loss meter \n",
|
| 491 |
+
" loss.update(loss_this.item(), target.shape[0])\n",
|
| 492 |
+
" print('Train: Average loss: {:.4f}\\n'.format(loss.avg))\n",
|
| 493 |
+
" return loss.avg\n",
|
| 494 |
+
" \n",
|
| 495 |
+
"##define test function\n",
|
| 496 |
+
"def test(model, device, test_loader):\n",
|
| 497 |
+
" # meters\n",
|
| 498 |
+
" loss = AverageMeter()\n",
|
| 499 |
+
" acc = AverageMeter()\n",
|
| 500 |
+
" correct = 0\n",
|
| 501 |
+
" # switch to test mode\n",
|
| 502 |
+
" model.eval()\n",
|
| 503 |
+
" for data, target in test_loader:\n",
|
| 504 |
+
" # after fetching the data transfer the model to the \n",
|
| 505 |
+
" # required device, in this example the device is gpu\n",
|
| 506 |
+
" # transfer to gpu can also be done by \n",
|
| 507 |
+
" # data, target = data.cuda(), target.cuda()\n",
|
| 508 |
+
" data, target = data.to(device), target.to(device) # data, target = data.cuda(), target.cuda()\n",
|
| 509 |
+
" # since we dont need to backpropagate loss in testing,\n",
|
| 510 |
+
" # we dont keep the gradient\n",
|
| 511 |
+
" with torch.no_grad():\n",
|
| 512 |
+
" # compute the forward pass\n",
|
| 513 |
+
" # it can also be achieved by model.forward(data)\n",
|
| 514 |
+
" output = model(data)\n",
|
| 515 |
+
" # compute the loss function just for checking\n",
|
| 516 |
+
" loss_this = F.cross_entropy(output, target) # sum up batch loss\n",
|
| 517 |
+
" # get the index of the max log-probability\n",
|
| 518 |
+
" pred = output.argmax(dim=1, keepdim=True) \n",
|
| 519 |
+
" # check which of the predictions are correct\n",
|
| 520 |
+
" correct_this = pred.eq(target.view_as(pred)).sum().item()\n",
|
| 521 |
+
" # accumulate the correct ones\n",
|
| 522 |
+
" correct += correct_this\n",
|
| 523 |
+
" # compute accuracy\n",
|
| 524 |
+
" acc_this = correct_this/target.shape[0]*100.0\n",
|
| 525 |
+
" # update the loss and accuracy meter \n",
|
| 526 |
+
" acc.update(acc_this, target.shape[0])\n",
|
| 527 |
+
" loss.update(loss_this.item(), target.shape[0])\n",
|
| 528 |
+
" print('Test: Average loss: {:.4f}, Accuracy: {}/{} ({:.2f}%)\\n'.format(\n",
|
| 529 |
+
" loss.avg, correct, len(test_loader.dataset), acc.avg))"
|
| 530 |
+
]
|
| 531 |
+
},
|
| 532 |
+
{
|
| 533 |
+
"cell_type": "markdown",
|
| 534 |
+
"metadata": {
|
| 535 |
+
"id": "vzxivyrXDB7a",
|
| 536 |
+
"pycharm": {}
|
| 537 |
+
},
|
| 538 |
+
"source": [
|
| 539 |
+
"## Training Loop\n",
|
| 540 |
+
"Training loop containing alternating train and test phase. Below we are iterating the loops 5 times, you can iterate more times."
|
| 541 |
+
]
|
| 542 |
+
},
|
| 543 |
+
{
|
| 544 |
+
"cell_type": "code",
|
| 545 |
+
"execution_count": null,
|
| 546 |
+
"metadata": {
|
| 547 |
+
"id": "tna_R8TSDD4D",
|
| 548 |
+
"pycharm": {
|
| 549 |
+
"is_executing": true
|
| 550 |
+
}
|
| 551 |
+
},
|
| 552 |
+
"outputs": [],
|
| 553 |
+
"source": [
|
| 554 |
+
"# import tensorboard logger from PyTorch\n",
|
| 555 |
+
"from torch.utils.tensorboard import SummaryWriter\n",
|
| 556 |
+
"# create TensorBoard logger\n",
|
| 557 |
+
"writer = SummaryWriter('runs/stl10_experiment_1')\n",
|
| 558 |
+
"# number of epochs we decide to train\n",
|
| 559 |
+
"num_epoch = 10\n",
|
| 560 |
+
"for epoch in range(1, num_epoch + 1):\n",
|
| 561 |
+
" epoch_loss = train(model, device, train_loader, optimizer)\n",
|
| 562 |
+
" writer.add_scalar('training_loss', epoch_loss, global_step = epoch)\n",
|
| 563 |
+
"test(model, device, test_loader)"
|
| 564 |
+
]
|
| 565 |
+
},
|
| 566 |
+
{
|
| 567 |
+
"cell_type": "markdown",
|
| 568 |
+
"metadata": {
|
| 569 |
+
"id": "grSrfsp5BhTC"
|
| 570 |
+
},
|
| 571 |
+
"source": [
|
| 572 |
+
"### Training loss curve"
|
| 573 |
+
]
|
| 574 |
+
},
|
| 575 |
+
{
|
| 576 |
+
"cell_type": "markdown",
|
| 577 |
+
"metadata": {
|
| 578 |
+
"id": "5FhhYHF_BUBv"
|
| 579 |
+
},
|
| 580 |
+
"source": [
|
| 581 |
+
"The TensorBoard file in the folder runs/stl10_experiment_1 now contains a training loss curve over number of epochs. To start the TensorBoard visualizer, simply run the following statements."
|
| 582 |
+
]
|
| 583 |
+
},
|
| 584 |
+
{
|
| 585 |
+
"cell_type": "code",
|
| 586 |
+
"execution_count": null,
|
| 587 |
+
"metadata": {
|
| 588 |
+
"id": "oR9tsiHJ6LlS"
|
| 589 |
+
},
|
| 590 |
+
"outputs": [],
|
| 591 |
+
"source": [
|
| 592 |
+
"# Load tensorboard extension for Jupyter Notebook, only need to start TB in the notebook\n",
|
| 593 |
+
"%reload_ext tensorboard\n",
|
| 594 |
+
"%tensorboard --logdir runs/stl10_experiment_1"
|
| 595 |
+
]
|
| 596 |
+
},
|
| 597 |
+
{
|
| 598 |
+
"cell_type": "markdown",
|
| 599 |
+
"metadata": {
|
| 600 |
+
"id": "1B19bDzBDIV4",
|
| 601 |
+
"pycharm": {}
|
| 602 |
+
},
|
| 603 |
+
"source": [
|
| 604 |
+
"### Summary\n",
|
| 605 |
+
"Show the summary of the model. It shows the number of parameters in layerwise as well as the total number of parameters. It also shows the memories required for training the model."
|
| 606 |
+
]
|
| 607 |
+
},
|
| 608 |
+
{
|
| 609 |
+
"cell_type": "code",
|
| 610 |
+
"execution_count": null,
|
| 611 |
+
"metadata": {
|
| 612 |
+
"id": "9eztTJ4VDKCA",
|
| 613 |
+
"pycharm": {}
|
| 614 |
+
},
|
| 615 |
+
"outputs": [],
|
| 616 |
+
"source": [
|
| 617 |
+
"from torchsummary import summary\n",
|
| 618 |
+
"summary(model, (3, 96, 96))"
|
| 619 |
+
]
|
| 620 |
+
},
|
| 621 |
+
{
|
| 622 |
+
"cell_type": "markdown",
|
| 623 |
+
"metadata": {
|
| 624 |
+
"id": "Skq4kzyt7PpN",
|
| 625 |
+
"pycharm": {}
|
| 626 |
+
},
|
| 627 |
+
"source": [
|
| 628 |
+
"## Transfer Learning\n",
|
| 629 |
+
"Transfer learning is a machine learning paradigm where a model developed for a task is reused as the starting point for a model on a second task. In this workshop, you will learn how to classifiy sketches using pretrained model trained on [ImageNet-1K](http://image-net.org/). In this part of the workshop, we will use the pretrained weights of AlexNet (note the `pretrained=True` flag in `AlexNet4` model above) available from PyTorch to classify sketch images from the [TU-Berlin dataset](http://cybertron.cg.tu-berlin.de/eitz/projects/classifysketch/). This is a very good example where knowledge or weights learned from natural images could be used for solving a classification task on completely different domains, such as sketch."
|
| 630 |
+
]
|
| 631 |
+
},
|
| 632 |
+
{
|
| 633 |
+
"cell_type": "markdown",
|
| 634 |
+
"metadata": {
|
| 635 |
+
"id": "YxLFKtvlV9jN",
|
| 636 |
+
"pycharm": {}
|
| 637 |
+
},
|
| 638 |
+
"source": [
|
| 639 |
+
"## TU-Berlin Dataset\n",
|
| 640 |
+
"TU-Berlin dataset contains over 20,000 human drawn sketches evenly distributed over 250 object categories. Some of the sketches from the dataset can be seen below.\n",
|
| 641 |
+
"\n",
|
| 642 |
+
"\n",
|
| 643 |
+
"\n",
|
| 644 |
+
"More details on the dataset can be found here: http://cybertron.cg.tu-berlin.de/eitz/projects/classifysketch/. Lets download the dataset and prepare it for usage."
|
| 645 |
+
]
|
| 646 |
+
},
|
| 647 |
+
{
|
| 648 |
+
"cell_type": "code",
|
| 649 |
+
"execution_count": null,
|
| 650 |
+
"metadata": {
|
| 651 |
+
"id": "bxXqCPXvVI9J",
|
| 652 |
+
"pycharm": {
|
| 653 |
+
"is_executing": true
|
| 654 |
+
}
|
| 655 |
+
},
|
| 656 |
+
"outputs": [],
|
| 657 |
+
"source": [
|
| 658 |
+
"import os\n",
|
| 659 |
+
"if not os.path.exists('sketches_png.zip'):\n",
|
| 660 |
+
" !wget http://cybertron.cg.tu-berlin.de/eitz/projects/classifysketch/sketches_png.zip\n",
|
| 661 |
+
" !unzip -q sketches_png.zip\n",
|
| 662 |
+
" !rm sketches_png.zip\n",
|
| 663 |
+
" !mv png tu_berlin"
|
| 664 |
+
]
|
| 665 |
+
},
|
| 666 |
+
{
|
| 667 |
+
"cell_type": "markdown",
|
| 668 |
+
"metadata": {
|
| 669 |
+
"id": "kopBHyQu5K8t",
|
| 670 |
+
"pycharm": {}
|
| 671 |
+
},
|
| 672 |
+
"source": [
|
| 673 |
+
"### Split into Train and Test Set"
|
| 674 |
+
]
|
| 675 |
+
},
|
| 676 |
+
{
|
| 677 |
+
"cell_type": "code",
|
| 678 |
+
"execution_count": null,
|
| 679 |
+
"metadata": {
|
| 680 |
+
"id": "bgkRoHjIub9M",
|
| 681 |
+
"pycharm": {}
|
| 682 |
+
},
|
| 683 |
+
"outputs": [],
|
| 684 |
+
"source": [
|
| 685 |
+
"import numpy as np\n",
|
| 686 |
+
"from sklearn.model_selection import train_test_split\n",
|
| 687 |
+
"with open('tu_berlin/filelist.txt', 'r') as fp:\n",
|
| 688 |
+
" files = fp.read().splitlines()\n",
|
| 689 |
+
"classes_str = [file.split('/')[0] for file in files]\n",
|
| 690 |
+
"classes_str, classes = np.unique(classes_str, return_inverse=True)\n",
|
| 691 |
+
"train_files, test_files, train_classes, test_classes = train_test_split(files, classes, train_size=0.3, test_size=0.1, stratify=classes)"
|
| 692 |
+
]
|
| 693 |
+
},
|
| 694 |
+
{
|
| 695 |
+
"cell_type": "markdown",
|
| 696 |
+
"metadata": {
|
| 697 |
+
"id": "6_zDUF0R4xgR",
|
| 698 |
+
"pycharm": {}
|
| 699 |
+
},
|
| 700 |
+
"source": [
|
| 701 |
+
"### Custom Dataset\n",
|
| 702 |
+
"Since TU-Berlin is not implemented as a data generator within the torchvision package, we have to implement a custom data generator for this. One need to inherit the [`data.Dataset` class](https://pytorch.org/docs/stable/data.html) of PyTorch for designing a data generator for a dataset. The custom class should override the following methods:\n",
|
| 703 |
+
"\n",
|
| 704 |
+
"* `__len__` so that `len(dataset)` returns the size of the dataset.\n",
|
| 705 |
+
"* `__getitem__` to support the indexing such that `dataset[i]` can be used to get *i*th sample.\n",
|
| 706 |
+
"\n",
|
| 707 |
+
"Now lets create a dataset class for our TU-Berlin dataset. We will set the location of the sketches inside the `__init__` function, but leave the loading image sketch images for the `__getitem__` function. This way is memory efficient because all the images are not stored in the memory at once but read as required."
|
| 708 |
+
]
|
| 709 |
+
},
|
| 710 |
+
{
|
| 711 |
+
"cell_type": "code",
|
| 712 |
+
"execution_count": null,
|
| 713 |
+
"metadata": {
|
| 714 |
+
"id": "fCcARypv3hpL",
|
| 715 |
+
"pycharm": {}
|
| 716 |
+
},
|
| 717 |
+
"outputs": [],
|
| 718 |
+
"source": [
|
| 719 |
+
"from PIL import Image\n",
|
| 720 |
+
"import torch.utils.data as data\n",
|
| 721 |
+
"class TUBerlin(data.Dataset):\n",
|
| 722 |
+
" def __init__(self, root, files, classes, transforms=None): \n",
|
| 723 |
+
" # location of the dataset\n",
|
| 724 |
+
" self.root = root\n",
|
| 725 |
+
" # list of files\n",
|
| 726 |
+
" self.files = files\n",
|
| 727 |
+
" # list of classes\n",
|
| 728 |
+
" self.classes = classes\n",
|
| 729 |
+
" # transforms\n",
|
| 730 |
+
" self.transforms = transforms\n",
|
| 731 |
+
"\n",
|
| 732 |
+
" def __getitem__(self, item):\n",
|
| 733 |
+
" # read the image\n",
|
| 734 |
+
" image = Image.open(os.path.join(self.root, self.files[item])).convert(mode=\"RGB\")\n",
|
| 735 |
+
" # class for that image\n",
|
| 736 |
+
" class_ = self.classes[item]\n",
|
| 737 |
+
" # apply transformation\n",
|
| 738 |
+
" if self.transforms:\n",
|
| 739 |
+
" image = self.transforms(image)\n",
|
| 740 |
+
" # return the image and class\n",
|
| 741 |
+
" return image, class_\n",
|
| 742 |
+
"\n",
|
| 743 |
+
" def __len__(self):\n",
|
| 744 |
+
" # return the total number of images\n",
|
| 745 |
+
" return len(self.files)"
|
| 746 |
+
]
|
| 747 |
+
},
|
| 748 |
+
{
|
| 749 |
+
"cell_type": "markdown",
|
| 750 |
+
"metadata": {
|
| 751 |
+
"id": "NEFyh0BDo-l4",
|
| 752 |
+
"pycharm": {}
|
| 753 |
+
},
|
| 754 |
+
"source": [
|
| 755 |
+
"### Dataset and DataLoader\n",
|
| 756 |
+
"In the following cell, we are defining the datasets and data loaders. The usage of different functions are alike to the example mentioned above."
|
| 757 |
+
]
|
| 758 |
+
},
|
| 759 |
+
{
|
| 760 |
+
"cell_type": "code",
|
| 761 |
+
"execution_count": null,
|
| 762 |
+
"metadata": {
|
| 763 |
+
"id": "OTP0pmYG5oM9",
|
| 764 |
+
"pycharm": {}
|
| 765 |
+
},
|
| 766 |
+
"outputs": [],
|
| 767 |
+
"source": [
|
| 768 |
+
"import torch\n",
|
| 769 |
+
"import torchvision\n",
|
| 770 |
+
"# Define batch size, batch size is how much data you feed for training in one iteration\n",
|
| 771 |
+
"batch_size_train = 256 # We use a small batch size here for training\n",
|
| 772 |
+
"batch_size_test = 1024 # We use bigger batch size for testing\n",
|
| 773 |
+
"\n",
|
| 774 |
+
"# define how image transformed\n",
|
| 775 |
+
"image_transform = torchvision.transforms.Compose([\n",
|
| 776 |
+
" torchvision.transforms.Resize((224, 224)),\n",
|
| 777 |
+
" torchvision.transforms.ToTensor(),\n",
|
| 778 |
+
" torchvision.transforms.Normalize(\n",
|
| 779 |
+
" (0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])\n",
|
| 780 |
+
"# image datasets\n",
|
| 781 |
+
"train_dataset = TUBerlin('tu_berlin/', train_files, train_classes, \n",
|
| 782 |
+
" transforms=image_transform)\n",
|
| 783 |
+
"test_dataset = TUBerlin('tu_berlin/', test_files, test_classes, \n",
|
| 784 |
+
" transforms=image_transform)\n",
|
| 785 |
+
"# data loaders\n",
|
| 786 |
+
"train_loader = torch.utils.data.DataLoader(train_dataset,\n",
|
| 787 |
+
" batch_size=batch_size_train, \n",
|
| 788 |
+
" shuffle=True, num_workers=2)\n",
|
| 789 |
+
"test_loader = torch.utils.data.DataLoader(test_dataset,\n",
|
| 790 |
+
" batch_size=batch_size_test, \n",
|
| 791 |
+
" shuffle=True, num_workers=2)"
|
| 792 |
+
]
|
| 793 |
+
},
|
| 794 |
+
{
|
| 795 |
+
"cell_type": "markdown",
|
| 796 |
+
"metadata": {
|
| 797 |
+
"id": "oFhh272O7jkj",
|
| 798 |
+
"pycharm": {}
|
| 799 |
+
},
|
| 800 |
+
"source": [
|
| 801 |
+
"### Example Image"
|
| 802 |
+
]
|
| 803 |
+
},
|
| 804 |
+
{
|
| 805 |
+
"cell_type": "code",
|
| 806 |
+
"execution_count": null,
|
| 807 |
+
"metadata": {
|
| 808 |
+
"id": "4BSAiPFv7jkj",
|
| 809 |
+
"pycharm": {
|
| 810 |
+
"is_executing": true
|
| 811 |
+
}
|
| 812 |
+
},
|
| 813 |
+
"outputs": [],
|
| 814 |
+
"source": [
|
| 815 |
+
"# import library\n",
|
| 816 |
+
"import matplotlib.pyplot as plt\n",
|
| 817 |
+
"# We can check the dataloader\n",
|
| 818 |
+
"_, (example_datas, labels) = next(enumerate(train_loader))\n",
|
| 819 |
+
"sample = example_datas[0]\n",
|
| 820 |
+
"# show the data\n",
|
| 821 |
+
"plt.imshow(sample.permute(1, 2, 0));\n",
|
| 822 |
+
"print(\"Label: \" + str(classes_str[labels[0]]))"
|
| 823 |
+
]
|
| 824 |
+
},
|
| 825 |
+
{
|
| 826 |
+
"cell_type": "markdown",
|
| 827 |
+
"metadata": {
|
| 828 |
+
"id": "__Wx5zUY7L92",
|
| 829 |
+
"pycharm": {}
|
| 830 |
+
},
|
| 831 |
+
"source": [
|
| 832 |
+
"## Initialization\n",
|
| 833 |
+
"\n",
|
| 834 |
+
"Please read the comments and understand the purpose of different lines of code."
|
| 835 |
+
]
|
| 836 |
+
},
|
| 837 |
+
{
|
| 838 |
+
"cell_type": "markdown",
|
| 839 |
+
"metadata": {
|
| 840 |
+
"id": "fRqz56BYjgoU"
|
| 841 |
+
},
|
| 842 |
+
"source": [
|
| 843 |
+
"### Model\n",
|
| 844 |
+
"\n",
|
| 845 |
+
"Please check below how to make some of the layers not trainable and other trainable."
|
| 846 |
+
]
|
| 847 |
+
},
|
| 848 |
+
{
|
| 849 |
+
"cell_type": "code",
|
| 850 |
+
"execution_count": null,
|
| 851 |
+
"metadata": {
|
| 852 |
+
"id": "mEawkRwgjvDq"
|
| 853 |
+
},
|
| 854 |
+
"outputs": [],
|
| 855 |
+
"source": [
|
| 856 |
+
"# define the model which contains pretrained weights from ImageNet\n",
|
| 857 |
+
"model = AlexNet4(250) # note the pretrained=True flag in the AlexNet4 model\n",
|
| 858 |
+
"# device: cuda (gpu) or cpu\n",
|
| 859 |
+
"device = \"cuda\"\n",
|
| 860 |
+
"# map to device\n",
|
| 861 |
+
"model = model.to(device)\n",
|
| 862 |
+
"################################################################################\n",
|
| 863 |
+
"################################# IMPORTANT ####################################\n",
|
| 864 |
+
"################################################################################\n",
|
| 865 |
+
"# one can choose which parameters of the model to train or finetune\n",
|
| 866 |
+
"# Setting 1: make all the parameters of the model trainable\n",
|
| 867 |
+
"for param in model.parameters():\n",
|
| 868 |
+
" param.requires_grad = True\n",
|
| 869 |
+
"\n",
|
| 870 |
+
"# Setting 2: make only the last layer of the classifier handle trainable\n",
|
| 871 |
+
"for param in model.parameters():\n",
|
| 872 |
+
" param.requires_grad = False\n",
|
| 873 |
+
"for param in model.classifier[6].parameters():\n",
|
| 874 |
+
" param.requires_grad = True\n",
|
| 875 |
+
"\n",
|
| 876 |
+
"# Setting 3: make all the parameters of the conv layer (features handle) \n",
|
| 877 |
+
"# not trainable and others (classifier handle) trainable\n",
|
| 878 |
+
"for param in model.features.parameters():\n",
|
| 879 |
+
" param.requires_grad = False\n",
|
| 880 |
+
"for param in model.classifier.parameters():\n",
|
| 881 |
+
" param.requires_grad = True\n",
|
| 882 |
+
"\n",
|
| 883 |
+
"parameters = filter(lambda p: p.requires_grad, model.parameters())"
|
| 884 |
+
]
|
| 885 |
+
},
|
| 886 |
+
{
|
| 887 |
+
"cell_type": "markdown",
|
| 888 |
+
"metadata": {
|
| 889 |
+
"id": "DSg5gjryjkXV"
|
| 890 |
+
},
|
| 891 |
+
"source": [
|
| 892 |
+
"### Optimizer"
|
| 893 |
+
]
|
| 894 |
+
},
|
| 895 |
+
{
|
| 896 |
+
"cell_type": "code",
|
| 897 |
+
"execution_count": null,
|
| 898 |
+
"metadata": {
|
| 899 |
+
"id": "RxItpRN-7L93",
|
| 900 |
+
"pycharm": {}
|
| 901 |
+
},
|
| 902 |
+
"outputs": [],
|
| 903 |
+
"source": [
|
| 904 |
+
"## create model and optimizer\n",
|
| 905 |
+
"learning_rate = 0.0001\n",
|
| 906 |
+
"weight_decay = 0.0005\n",
|
| 907 |
+
"# define optimizer\n",
|
| 908 |
+
"optimizer = optim.Adam(parameters, lr=learning_rate, weight_decay=weight_decay)"
|
| 909 |
+
]
|
| 910 |
+
},
|
| 911 |
+
{
|
| 912 |
+
"cell_type": "markdown",
|
| 913 |
+
"metadata": {
|
| 914 |
+
"id": "f_ET0FAw7Va_",
|
| 915 |
+
"pycharm": {}
|
| 916 |
+
},
|
| 917 |
+
"source": [
|
| 918 |
+
"## Training Loop\n",
|
| 919 |
+
"Training loop for several epochs. Perform testing after training the model for some epochs."
|
| 920 |
+
]
|
| 921 |
+
},
|
| 922 |
+
{
|
| 923 |
+
"cell_type": "code",
|
| 924 |
+
"execution_count": null,
|
| 925 |
+
"metadata": {
|
| 926 |
+
"id": "4OyndchC7VbA",
|
| 927 |
+
"pycharm": {
|
| 928 |
+
"is_executing": true
|
| 929 |
+
}
|
| 930 |
+
},
|
| 931 |
+
"outputs": [],
|
| 932 |
+
"source": [
|
| 933 |
+
"num_epoch = 5\n",
|
| 934 |
+
"for epoch in range(1, num_epoch + 1):\n",
|
| 935 |
+
" train(model, device, train_loader, optimizer)\n",
|
| 936 |
+
"test(model, device, test_loader)"
|
| 937 |
+
]
|
| 938 |
+
},
|
| 939 |
+
{
|
| 940 |
+
"cell_type": "markdown",
|
| 941 |
+
"metadata": {
|
| 942 |
+
"id": "_hQ1PoEEvSQF"
|
| 943 |
+
},
|
| 944 |
+
"source": [
|
| 945 |
+
"## Summary\n",
|
| 946 |
+
"Show the summary of the model. It shows the number of parameters in layerwise as well as the total number of parameters. It also shows the memories required for training the model."
|
| 947 |
+
]
|
| 948 |
+
},
|
| 949 |
+
{
|
| 950 |
+
"cell_type": "code",
|
| 951 |
+
"execution_count": null,
|
| 952 |
+
"metadata": {
|
| 953 |
+
"id": "dxNHMSBPXNLp",
|
| 954 |
+
"pycharm": {}
|
| 955 |
+
},
|
| 956 |
+
"outputs": [],
|
| 957 |
+
"source": [
|
| 958 |
+
"from torchsummary import summary\n",
|
| 959 |
+
"summary(model, (3, 224, 224))"
|
| 960 |
+
]
|
| 961 |
+
}
|
| 962 |
+
],
|
| 963 |
+
"metadata": {
|
| 964 |
+
"accelerator": "GPU",
|
| 965 |
+
"colab": {
|
| 966 |
+
"include_colab_link": true,
|
| 967 |
+
"name": "ECMM426/ECMM441 - Convolutional Neural Networks and Transfer Learning.ipynb",
|
| 968 |
+
"provenance": []
|
| 969 |
+
},
|
| 970 |
+
"kernelspec": {
|
| 971 |
+
"display_name": "Python 3 (ipykernel)",
|
| 972 |
+
"language": "python",
|
| 973 |
+
"name": "python3"
|
| 974 |
+
},
|
| 975 |
+
"language_info": {
|
| 976 |
+
"codemirror_mode": {
|
| 977 |
+
"name": "ipython",
|
| 978 |
+
"version": 3
|
| 979 |
+
},
|
| 980 |
+
"file_extension": ".py",
|
| 981 |
+
"mimetype": "text/x-python",
|
| 982 |
+
"name": "python",
|
| 983 |
+
"nbconvert_exporter": "python",
|
| 984 |
+
"pygments_lexer": "ipython3",
|
| 985 |
+
"version": "3.12.3"
|
| 986 |
+
}
|
| 987 |
+
},
|
| 988 |
+
"nbformat": 4,
|
| 989 |
+
"nbformat_minor": 1
|
| 990 |
+
}
|
Downloads/.ipynb_checkpoints/EEEM071_CourseWork_ipynb-checkpoint.ipynb
ADDED
|
@@ -0,0 +1,379 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {
|
| 6 |
+
"id": "9X3IbRehFNjA"
|
| 7 |
+
},
|
| 8 |
+
"source": [
|
| 9 |
+
"# Step 1: GPU Selection\n",
|
| 10 |
+
"1. Find \"Edit\" tab above, select \"Hardware accelerator\" and choose \"GPU\"\n",
|
| 11 |
+
"2. Run bellow command to check what GPU you got"
|
| 12 |
+
]
|
| 13 |
+
},
|
| 14 |
+
{
|
| 15 |
+
"cell_type": "code",
|
| 16 |
+
"execution_count": null,
|
| 17 |
+
"metadata": {
|
| 18 |
+
"colab": {
|
| 19 |
+
"base_uri": "https://localhost:8080/"
|
| 20 |
+
},
|
| 21 |
+
"executionInfo": {
|
| 22 |
+
"elapsed": 120,
|
| 23 |
+
"status": "ok",
|
| 24 |
+
"timestamp": 1742400948017,
|
| 25 |
+
"user": {
|
| 26 |
+
"displayName": "Wenqing Wang",
|
| 27 |
+
"userId": "10666645302626123442"
|
| 28 |
+
},
|
| 29 |
+
"user_tz": 0
|
| 30 |
+
},
|
| 31 |
+
"id": "xHEpV2i0zQrW",
|
| 32 |
+
"outputId": "582b5d58-d87d-479e-f1ab-feed3516fe6c"
|
| 33 |
+
},
|
| 34 |
+
"outputs": [
|
| 35 |
+
{
|
| 36 |
+
"name": "stdout",
|
| 37 |
+
"output_type": "stream",
|
| 38 |
+
"text": [
|
| 39 |
+
"Wed Mar 19 16:15:47 2025 \n",
|
| 40 |
+
"+-----------------------------------------------------------------------------------------+\n",
|
| 41 |
+
"| NVIDIA-SMI 550.54.15 Driver Version: 550.54.15 CUDA Version: 12.4 |\n",
|
| 42 |
+
"|-----------------------------------------+------------------------+----------------------+\n",
|
| 43 |
+
"| GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC |\n",
|
| 44 |
+
"| Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. |\n",
|
| 45 |
+
"| | | MIG M. |\n",
|
| 46 |
+
"|=========================================+========================+======================|\n",
|
| 47 |
+
"| 0 Tesla T4 Off | 00000000:00:04.0 Off | 0 |\n",
|
| 48 |
+
"| N/A 36C P8 9W / 70W | 0MiB / 15360MiB | 0% Default |\n",
|
| 49 |
+
"| | | N/A |\n",
|
| 50 |
+
"+-----------------------------------------+------------------------+----------------------+\n",
|
| 51 |
+
" \n",
|
| 52 |
+
"+-----------------------------------------------------------------------------------------+\n",
|
| 53 |
+
"| Processes: |\n",
|
| 54 |
+
"| GPU GI CI PID Type Process name GPU Memory |\n",
|
| 55 |
+
"| ID ID Usage |\n",
|
| 56 |
+
"|=========================================================================================|\n",
|
| 57 |
+
"| No running processes found |\n",
|
| 58 |
+
"+-----------------------------------------------------------------------------------------+\n"
|
| 59 |
+
]
|
| 60 |
+
}
|
| 61 |
+
],
|
| 62 |
+
"source": [
|
| 63 |
+
"!nvidia-smi"
|
| 64 |
+
]
|
| 65 |
+
},
|
| 66 |
+
{
|
| 67 |
+
"cell_type": "markdown",
|
| 68 |
+
"metadata": {
|
| 69 |
+
"id": "oQzoka29GGzQ"
|
| 70 |
+
},
|
| 71 |
+
"source": [
|
| 72 |
+
"# Step 2: Code Preparation\n",
|
| 73 |
+
"\n",
|
| 74 |
+
"We need to maintain our codebase with git history, so a file system (Google Drive) is needed\n",
|
| 75 |
+
"1. Select the left file icon and mount your Google Drive\n",
|
| 76 |
+
"2. Move path to Google Drive\n",
|
| 77 |
+
"3. Git clone code base\n"
|
| 78 |
+
]
|
| 79 |
+
},
|
| 80 |
+
{
|
| 81 |
+
"cell_type": "code",
|
| 82 |
+
"execution_count": null,
|
| 83 |
+
"metadata": {
|
| 84 |
+
"colab": {
|
| 85 |
+
"base_uri": "https://localhost:8080/"
|
| 86 |
+
},
|
| 87 |
+
"executionInfo": {
|
| 88 |
+
"elapsed": 516,
|
| 89 |
+
"status": "ok",
|
| 90 |
+
"timestamp": 1742400952955,
|
| 91 |
+
"user": {
|
| 92 |
+
"displayName": "Wenqing Wang",
|
| 93 |
+
"userId": "10666645302626123442"
|
| 94 |
+
},
|
| 95 |
+
"user_tz": 0
|
| 96 |
+
},
|
| 97 |
+
"id": "4H0KIZFSGC4Z",
|
| 98 |
+
"outputId": "918bc8c9-00d1-4677-a68c-38c5977c3581"
|
| 99 |
+
},
|
| 100 |
+
"outputs": [
|
| 101 |
+
{
|
| 102 |
+
"name": "stdout",
|
| 103 |
+
"output_type": "stream",
|
| 104 |
+
"text": [
|
| 105 |
+
"[Errno 2] No such file or directory: '/content/MyDrive/'\n",
|
| 106 |
+
"/content/EEEM071-Coursework-2025\n",
|
| 107 |
+
"Cloning into 'EEEM071-Coursework-2025'...\n",
|
| 108 |
+
"remote: Enumerating objects: 55, done.\u001b[K\n",
|
| 109 |
+
"remote: Counting objects: 100% (15/15), done.\u001b[K\n",
|
| 110 |
+
"remote: Compressing objects: 100% (13/13), done.\u001b[K\n",
|
| 111 |
+
"remote: Total 55 (delta 5), reused 1 (delta 1), pack-reused 40 (from 1)\u001b[K\n",
|
| 112 |
+
"Receiving objects: 100% (55/55), 33.64 KiB | 4.20 MiB/s, done.\n",
|
| 113 |
+
"Resolving deltas: 100% (7/7), done.\n"
|
| 114 |
+
]
|
| 115 |
+
}
|
| 116 |
+
],
|
| 117 |
+
"source": [
|
| 118 |
+
"# %cd /content/drive/MyDrive/\n",
|
| 119 |
+
"%cd /content/MyDrive/\n",
|
| 120 |
+
"\n",
|
| 121 |
+
"!git clone https://github.com/Surrey-EEEM071-CVDL/EEEM071-Coursework-2025.git"
|
| 122 |
+
]
|
| 123 |
+
},
|
| 124 |
+
{
|
| 125 |
+
"cell_type": "code",
|
| 126 |
+
"execution_count": null,
|
| 127 |
+
"metadata": {
|
| 128 |
+
"colab": {
|
| 129 |
+
"base_uri": "https://localhost:8080/"
|
| 130 |
+
},
|
| 131 |
+
"executionInfo": {
|
| 132 |
+
"elapsed": 1549,
|
| 133 |
+
"status": "ok",
|
| 134 |
+
"timestamp": 1742400957932,
|
| 135 |
+
"user": {
|
| 136 |
+
"displayName": "Wenqing Wang",
|
| 137 |
+
"userId": "10666645302626123442"
|
| 138 |
+
},
|
| 139 |
+
"user_tz": 0
|
| 140 |
+
},
|
| 141 |
+
"id": "-LJSH0dbIeW5",
|
| 142 |
+
"outputId": "3fb3b507-65a8-45af-d33c-6ac920c4cd6c"
|
| 143 |
+
},
|
| 144 |
+
"outputs": [
|
| 145 |
+
{
|
| 146 |
+
"name": "stdout",
|
| 147 |
+
"output_type": "stream",
|
| 148 |
+
"text": [
|
| 149 |
+
"Drive already mounted at /content/drive; to attempt to forcibly remount, call drive.mount(\"/content/drive\", force_remount=True).\n"
|
| 150 |
+
]
|
| 151 |
+
}
|
| 152 |
+
],
|
| 153 |
+
"source": [
|
| 154 |
+
"from google.colab import drive\n",
|
| 155 |
+
"drive.mount('/content/drive')"
|
| 156 |
+
]
|
| 157 |
+
},
|
| 158 |
+
{
|
| 159 |
+
"cell_type": "markdown",
|
| 160 |
+
"metadata": {
|
| 161 |
+
"id": "SYUXrD24IV8x"
|
| 162 |
+
},
|
| 163 |
+
"source": [
|
| 164 |
+
"# Step 3: Data Preparation\n",
|
| 165 |
+
"\n",
|
| 166 |
+
"Because reading images from Google Drive is very slow, we download datasets to Colab temporary file\n",
|
| 167 |
+
"1. Install gdown\n",
|
| 168 |
+
"2. Download data\n",
|
| 169 |
+
"3. Unzip data with password"
|
| 170 |
+
]
|
| 171 |
+
},
|
| 172 |
+
{
|
| 173 |
+
"cell_type": "code",
|
| 174 |
+
"execution_count": null,
|
| 175 |
+
"metadata": {
|
| 176 |
+
"colab": {
|
| 177 |
+
"base_uri": "https://localhost:8080/"
|
| 178 |
+
},
|
| 179 |
+
"executionInfo": {
|
| 180 |
+
"elapsed": 3320,
|
| 181 |
+
"status": "ok",
|
| 182 |
+
"timestamp": 1742400323557,
|
| 183 |
+
"user": {
|
| 184 |
+
"displayName": "Wenqing Wang",
|
| 185 |
+
"userId": "10666645302626123442"
|
| 186 |
+
},
|
| 187 |
+
"user_tz": 0
|
| 188 |
+
},
|
| 189 |
+
"id": "_X3Y8Adk1xjd",
|
| 190 |
+
"outputId": "6c01c8c0-874e-4a1e-ddfd-eb45fb6133d1"
|
| 191 |
+
},
|
| 192 |
+
"outputs": [
|
| 193 |
+
{
|
| 194 |
+
"name": "stdout",
|
| 195 |
+
"output_type": "stream",
|
| 196 |
+
"text": [
|
| 197 |
+
"/content\n",
|
| 198 |
+
"Requirement already satisfied: gdown in /usr/local/lib/python3.11/dist-packages (5.2.0)\n",
|
| 199 |
+
"Requirement already satisfied: beautifulsoup4 in /usr/local/lib/python3.11/dist-packages (from gdown) (4.13.3)\n",
|
| 200 |
+
"Requirement already satisfied: filelock in /usr/local/lib/python3.11/dist-packages (from gdown) (3.17.0)\n",
|
| 201 |
+
"Requirement already satisfied: requests[socks] in /usr/local/lib/python3.11/dist-packages (from gdown) (2.32.3)\n",
|
| 202 |
+
"Requirement already satisfied: tqdm in /usr/local/lib/python3.11/dist-packages (from gdown) (4.67.1)\n",
|
| 203 |
+
"Requirement already satisfied: soupsieve>1.2 in /usr/local/lib/python3.11/dist-packages (from beautifulsoup4->gdown) (2.6)\n",
|
| 204 |
+
"Requirement already satisfied: typing-extensions>=4.0.0 in /usr/local/lib/python3.11/dist-packages (from beautifulsoup4->gdown) (4.12.2)\n",
|
| 205 |
+
"Requirement already satisfied: charset-normalizer<4,>=2 in /usr/local/lib/python3.11/dist-packages (from requests[socks]->gdown) (3.4.1)\n",
|
| 206 |
+
"Requirement already satisfied: idna<4,>=2.5 in /usr/local/lib/python3.11/dist-packages (from requests[socks]->gdown) (3.10)\n",
|
| 207 |
+
"Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.11/dist-packages (from requests[socks]->gdown) (2.3.0)\n",
|
| 208 |
+
"Requirement already satisfied: certifi>=2017.4.17 in /usr/local/lib/python3.11/dist-packages (from requests[socks]->gdown) (2025.1.31)\n",
|
| 209 |
+
"Requirement already satisfied: PySocks!=1.5.7,>=1.5.6 in /usr/local/lib/python3.11/dist-packages (from requests[socks]->gdown) (1.7.1)\n"
|
| 210 |
+
]
|
| 211 |
+
}
|
| 212 |
+
],
|
| 213 |
+
"source": [
|
| 214 |
+
"%cd /content\n",
|
| 215 |
+
"!pip install -U --no-cache-dir gdown --pre\n",
|
| 216 |
+
"\n",
|
| 217 |
+
"# please download datasets from assignment doc link and upload, then unzip it."
|
| 218 |
+
]
|
| 219 |
+
},
|
| 220 |
+
{
|
| 221 |
+
"cell_type": "markdown",
|
| 222 |
+
"metadata": {
|
| 223 |
+
"id": "u1z0Kb-LMfh-"
|
| 224 |
+
},
|
| 225 |
+
"source": [
|
| 226 |
+
"# Step 4: Training"
|
| 227 |
+
]
|
| 228 |
+
},
|
| 229 |
+
{
|
| 230 |
+
"cell_type": "code",
|
| 231 |
+
"execution_count": null,
|
| 232 |
+
"metadata": {
|
| 233 |
+
"colab": {
|
| 234 |
+
"base_uri": "https://localhost:8080/"
|
| 235 |
+
},
|
| 236 |
+
"executionInfo": {
|
| 237 |
+
"elapsed": 111,
|
| 238 |
+
"status": "ok",
|
| 239 |
+
"timestamp": 1742401034406,
|
| 240 |
+
"user": {
|
| 241 |
+
"displayName": "Wenqing Wang",
|
| 242 |
+
"userId": "10666645302626123442"
|
| 243 |
+
},
|
| 244 |
+
"user_tz": 0
|
| 245 |
+
},
|
| 246 |
+
"id": "GDKjqjjsMcao",
|
| 247 |
+
"outputId": "ae0b64ff-0afc-4822-b7a3-953d2a936877"
|
| 248 |
+
},
|
| 249 |
+
"outputs": [
|
| 250 |
+
{
|
| 251 |
+
"name": "stdout",
|
| 252 |
+
"output_type": "stream",
|
| 253 |
+
"text": [
|
| 254 |
+
"/content/VeRi.zip: Zip archive data, at least v2.0 to extract, compression method=store\n"
|
| 255 |
+
]
|
| 256 |
+
}
|
| 257 |
+
],
|
| 258 |
+
"source": [
|
| 259 |
+
"!file /content/VeRi.zip"
|
| 260 |
+
]
|
| 261 |
+
},
|
| 262 |
+
{
|
| 263 |
+
"cell_type": "code",
|
| 264 |
+
"execution_count": null,
|
| 265 |
+
"metadata": {
|
| 266 |
+
"colab": {
|
| 267 |
+
"base_uri": "https://localhost:8080/"
|
| 268 |
+
},
|
| 269 |
+
"executionInfo": {
|
| 270 |
+
"elapsed": 5,
|
| 271 |
+
"status": "ok",
|
| 272 |
+
"timestamp": 1742400487347,
|
| 273 |
+
"user": {
|
| 274 |
+
"displayName": "Wenqing Wang",
|
| 275 |
+
"userId": "10666645302626123442"
|
| 276 |
+
},
|
| 277 |
+
"user_tz": 0
|
| 278 |
+
},
|
| 279 |
+
"id": "xL87Cl0W6Wll",
|
| 280 |
+
"outputId": "7723f374-114e-42aa-c90a-6620e2e29517"
|
| 281 |
+
},
|
| 282 |
+
"outputs": [
|
| 283 |
+
{
|
| 284 |
+
"name": "stdout",
|
| 285 |
+
"output_type": "stream",
|
| 286 |
+
"text": [
|
| 287 |
+
"/content/EEEM071-Coursework-2025\n"
|
| 288 |
+
]
|
| 289 |
+
}
|
| 290 |
+
],
|
| 291 |
+
"source": [
|
| 292 |
+
"%cd /content/EEEM071-Coursework-2025/"
|
| 293 |
+
]
|
| 294 |
+
},
|
| 295 |
+
{
|
| 296 |
+
"cell_type": "code",
|
| 297 |
+
"execution_count": 20,
|
| 298 |
+
"metadata": {
|
| 299 |
+
"colab": {
|
| 300 |
+
"base_uri": "https://localhost:8080/",
|
| 301 |
+
"height": 108
|
| 302 |
+
},
|
| 303 |
+
"executionInfo": {
|
| 304 |
+
"elapsed": 20,
|
| 305 |
+
"status": "error",
|
| 306 |
+
"timestamp": 1742403141571,
|
| 307 |
+
"user": {
|
| 308 |
+
"displayName": "Wenqing Wang",
|
| 309 |
+
"userId": "10666645302626123442"
|
| 310 |
+
},
|
| 311 |
+
"user_tz": 0
|
| 312 |
+
},
|
| 313 |
+
"id": "w2BqyZO3Mqjz",
|
| 314 |
+
"outputId": "ad17ae56-2501-467a-d9ae-e6bdd0fb0328"
|
| 315 |
+
},
|
| 316 |
+
"outputs": [
|
| 317 |
+
{
|
| 318 |
+
"ename": "SyntaxError",
|
| 319 |
+
"evalue": "invalid syntax (<ipython-input-20-64eca099a521>, line 1)",
|
| 320 |
+
"output_type": "error",
|
| 321 |
+
"traceback": [
|
| 322 |
+
"\u001b[0;36m File \u001b[0;32m\"<ipython-input-20-64eca099a521>\"\u001b[0;36m, line \u001b[0;32m1\u001b[0m\n\u001b[0;31m STUDENT_ID='@#123s£' STUDENT_NAME=\"Jane_/*%~#¬Doe\" python main.py \\\u001b[0m\n\u001b[0m ^\u001b[0m\n\u001b[0;31mSyntaxError\u001b[0m\u001b[0;31m:\u001b[0m invalid syntax\n"
|
| 323 |
+
]
|
| 324 |
+
}
|
| 325 |
+
],
|
| 326 |
+
"source": [
|
| 327 |
+
"!STUDENT_ID='@#123s£' STUDENT_NAME=\"Jane_/*%~#¬Doe\" python main.py \\\n",
|
| 328 |
+
"-s veri \\\n",
|
| 329 |
+
"-t veri \\\n",
|
| 330 |
+
"-a mobilenet_v3_small \\\n",
|
| 331 |
+
"--root /content/drive/MyDrive/VeRi \\\n",
|
| 332 |
+
"--height 224 \\\n",
|
| 333 |
+
"--width 224 \\\n",
|
| 334 |
+
"--optim amsgrad \\\n",
|
| 335 |
+
"--lr 0.0003 \\\n",
|
| 336 |
+
"--max-epoch 10 \\\n",
|
| 337 |
+
"--stepsize 20 40 \\\n",
|
| 338 |
+
"--train-batch-size 64 \\\n",
|
| 339 |
+
"--test-batch-size 100 \\\n",
|
| 340 |
+
"--save-dir logs/mobilenet_v3_small-veri"
|
| 341 |
+
]
|
| 342 |
+
},
|
| 343 |
+
{
|
| 344 |
+
"cell_type": "code",
|
| 345 |
+
"execution_count": null,
|
| 346 |
+
"metadata": {
|
| 347 |
+
"id": "jA43L0cTNpiJ"
|
| 348 |
+
},
|
| 349 |
+
"outputs": [],
|
| 350 |
+
"source": []
|
| 351 |
+
}
|
| 352 |
+
],
|
| 353 |
+
"metadata": {
|
| 354 |
+
"accelerator": "GPU",
|
| 355 |
+
"colab": {
|
| 356 |
+
"provenance": []
|
| 357 |
+
},
|
| 358 |
+
"gpuClass": "standard",
|
| 359 |
+
"kernelspec": {
|
| 360 |
+
"display_name": "Python 3 (ipykernel)",
|
| 361 |
+
"language": "python",
|
| 362 |
+
"name": "python3"
|
| 363 |
+
},
|
| 364 |
+
"language_info": {
|
| 365 |
+
"codemirror_mode": {
|
| 366 |
+
"name": "ipython",
|
| 367 |
+
"version": 3
|
| 368 |
+
},
|
| 369 |
+
"file_extension": ".py",
|
| 370 |
+
"mimetype": "text/x-python",
|
| 371 |
+
"name": "python",
|
| 372 |
+
"nbconvert_exporter": "python",
|
| 373 |
+
"pygments_lexer": "ipython3",
|
| 374 |
+
"version": "3.12.3"
|
| 375 |
+
}
|
| 376 |
+
},
|
| 377 |
+
"nbformat": 4,
|
| 378 |
+
"nbformat_minor": 1
|
| 379 |
+
}
|
Downloads/.ipynb_checkpoints/Human Action Recognition Tutorial-checkpoint.ipynb
ADDED
|
@@ -0,0 +1,820 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {
|
| 6 |
+
"colab_type": "text",
|
| 7 |
+
"id": "view-in-github"
|
| 8 |
+
},
|
| 9 |
+
"source": [
|
| 10 |
+
"<a href=\"https://colab.research.google.com/github/AnjanDutta/EEEM068/blob/main/Notebooks/Human_Action_Recognition.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
|
| 11 |
+
]
|
| 12 |
+
},
|
| 13 |
+
{
|
| 14 |
+
"cell_type": "markdown",
|
| 15 |
+
"metadata": {
|
| 16 |
+
"id": "RakR_TVgNmE6"
|
| 17 |
+
},
|
| 18 |
+
"source": [
|
| 19 |
+
"<H1 style=\"text-align: center\">EEEM068 - Applied Machine Learning</H1>\n",
|
| 20 |
+
"<H1 style=\"text-align: center\">Workshop 05</H1>\n",
|
| 21 |
+
"<H1 style=\"text-align: center\">Human Action Recognition Tutorial</H1>\n",
|
| 22 |
+
"\n"
|
| 23 |
+
]
|
| 24 |
+
},
|
| 25 |
+
{
|
| 26 |
+
"cell_type": "markdown",
|
| 27 |
+
"metadata": {
|
| 28 |
+
"id": "vNvDHbU2zY1u"
|
| 29 |
+
},
|
| 30 |
+
"source": [
|
| 31 |
+
"##Introduction"
|
| 32 |
+
]
|
| 33 |
+
},
|
| 34 |
+
{
|
| 35 |
+
"cell_type": "markdown",
|
| 36 |
+
"metadata": {
|
| 37 |
+
"id": "1xVXuN1-zXI9"
|
| 38 |
+
},
|
| 39 |
+
"source": [
|
| 40 |
+
"In this tutorial, we will explore 2D and 3D convolutional neural network models in recognizing human actions occurring in the videos from the [KTH dataset](https://www.csc.kth.se/cvap/actions/)."
|
| 41 |
+
]
|
| 42 |
+
},
|
| 43 |
+
{
|
| 44 |
+
"cell_type": "markdown",
|
| 45 |
+
"metadata": {
|
| 46 |
+
"id": "RBqUFnmzx21o"
|
| 47 |
+
},
|
| 48 |
+
"source": [
|
| 49 |
+
"## KTH Dataset\n",
|
| 50 |
+
"The KTH dataset consists of videos of humans performing 6 types of action: *boxing*, *clapping*, *waving*, *jogging*, *running*, and *walking*. There are 25 subjects performing these actions in 4 scenarios: outdoor, outdoor with scale variation, outdoor with different clothes, and indoor. The total number of videos is therefore $25 \\times 4 \\times 6 = 600$. The videos' frame rate are 25fps and their resolution is $160 \\times 120$. More information about the dataset can be looked up at the website. More details on this dataset can be found at https://www.csc.kth.se/cvap/actions.\n",
|
| 51 |
+
"\n",
|
| 52 |
+
"We have preprocessed (i.e., extracted the frames and resized those frames etc) those videos and pickled them in the following link (https://empslocal.ex.ac.uk/people/staff/ad735/ECMM426/KTH_pickle.zip). In this experiment, we will be using the pickled version of the KTH dataset. However, the raw KTH videos can be found in this link (https://empslocal.ex.ac.uk/people/staff/ad735/ECMM426/KTH.zip).\n",
|
| 53 |
+
"\n",
|
| 54 |
+
"<img src=\"http://www.csc.kth.se/cvap/actions/actions.gif\" alt=\"action recognition\" width=\"500\"/>"
|
| 55 |
+
]
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"cell_type": "markdown",
|
| 59 |
+
"metadata": {
|
| 60 |
+
"id": "sVrBKULQUi5r"
|
| 61 |
+
},
|
| 62 |
+
"source": [
|
| 63 |
+
"Lets download the pickled dataset from the above link."
|
| 64 |
+
]
|
| 65 |
+
},
|
| 66 |
+
{
|
| 67 |
+
"cell_type": "code",
|
| 68 |
+
"execution_count": null,
|
| 69 |
+
"metadata": {
|
| 70 |
+
"id": "JC1zNcUjwDmQ"
|
| 71 |
+
},
|
| 72 |
+
"outputs": [],
|
| 73 |
+
"source": [
|
| 74 |
+
"# Download the dataset\n",
|
| 75 |
+
"import os\n",
|
| 76 |
+
"if not os.path.exists('KTH_pickle.zip'):\n",
|
| 77 |
+
" !wget --no-check-certificate https://empslocal.ex.ac.uk/people/staff/ad735/ECMM426/KTH_pickle.zip\n",
|
| 78 |
+
" !unzip -q KTH_pickle.zip\n",
|
| 79 |
+
"\n",
|
| 80 |
+
"# Dictionary of categories\n",
|
| 81 |
+
"CATEGORY_INDEX = {\n",
|
| 82 |
+
" \"boxing\": 0,\n",
|
| 83 |
+
" \"handclapping\": 1,\n",
|
| 84 |
+
" \"handwaving\": 2,\n",
|
| 85 |
+
" \"jogging\": 3,\n",
|
| 86 |
+
" \"running\": 4,\n",
|
| 87 |
+
" \"walking\": 5\n",
|
| 88 |
+
"}"
|
| 89 |
+
]
|
| 90 |
+
},
|
| 91 |
+
{
|
| 92 |
+
"cell_type": "markdown",
|
| 93 |
+
"metadata": {
|
| 94 |
+
"id": "9_3bTBMf26bP"
|
| 95 |
+
},
|
| 96 |
+
"source": [
|
| 97 |
+
"### Example Video"
|
| 98 |
+
]
|
| 99 |
+
},
|
| 100 |
+
{
|
| 101 |
+
"cell_type": "markdown",
|
| 102 |
+
"metadata": {
|
| 103 |
+
"id": "8W-FdHna6gvQ"
|
| 104 |
+
},
|
| 105 |
+
"source": [
|
| 106 |
+
"Lets plot 20 equidistant frames of the very first video in the train split."
|
| 107 |
+
]
|
| 108 |
+
},
|
| 109 |
+
{
|
| 110 |
+
"cell_type": "code",
|
| 111 |
+
"execution_count": null,
|
| 112 |
+
"metadata": {
|
| 113 |
+
"id": "irg1Pter2_33"
|
| 114 |
+
},
|
| 115 |
+
"outputs": [],
|
| 116 |
+
"source": [
|
| 117 |
+
"import pickle\n",
|
| 118 |
+
"import numpy as np\n",
|
| 119 |
+
"import matplotlib.pyplot as plt\n",
|
| 120 |
+
"videos = pickle.load(open('KTH_pickle/train.pickle', 'rb'))\n",
|
| 121 |
+
"# first video\n",
|
| 122 |
+
"video_0 = videos[0]['frames']\n",
|
| 123 |
+
"# number of frames\n",
|
| 124 |
+
"n = 20\n",
|
| 125 |
+
"# figure size\n",
|
| 126 |
+
"fig = plt.figure(figsize=(25, 25))\n",
|
| 127 |
+
"# index of equidistant frames\n",
|
| 128 |
+
"nth_frames = np.linspace(0, len(video_0) - 1, n).astype(int)\n",
|
| 129 |
+
"for i in range(n):\n",
|
| 130 |
+
" frame = video_0[i]\n",
|
| 131 |
+
" fig.add_subplot(1, n, i + 1)\n",
|
| 132 |
+
" plt.imshow(frame, cmap='gray')\n",
|
| 133 |
+
" plt.axis('off')\n",
|
| 134 |
+
"plt.show()\n",
|
| 135 |
+
"\n",
|
| 136 |
+
"print('Action class: ' + videos[0]['category'])"
|
| 137 |
+
]
|
| 138 |
+
},
|
| 139 |
+
{
|
| 140 |
+
"cell_type": "markdown",
|
| 141 |
+
"metadata": {
|
| 142 |
+
"id": "_b_RgpUI-LEG"
|
| 143 |
+
},
|
| 144 |
+
"source": [
|
| 145 |
+
"Now lets plot 20 equidistant frames of the 91st video in the train split"
|
| 146 |
+
]
|
| 147 |
+
},
|
| 148 |
+
{
|
| 149 |
+
"cell_type": "code",
|
| 150 |
+
"execution_count": null,
|
| 151 |
+
"metadata": {
|
| 152 |
+
"id": "O22r78_C-Xpv"
|
| 153 |
+
},
|
| 154 |
+
"outputs": [],
|
| 155 |
+
"source": [
|
| 156 |
+
"# tenth video\n",
|
| 157 |
+
"video_90 = videos[90]['frames']\n",
|
| 158 |
+
"# number of frames\n",
|
| 159 |
+
"n = 20\n",
|
| 160 |
+
"# figure size\n",
|
| 161 |
+
"fig = plt.figure(figsize=(25, 25))\n",
|
| 162 |
+
"# index of equidistant frames\n",
|
| 163 |
+
"nth_frames = np.linspace(0, len(video_90) - 1, n).astype(int)\n",
|
| 164 |
+
"for i in range(n):\n",
|
| 165 |
+
" frame = video_90[i]\n",
|
| 166 |
+
" fig.add_subplot(1, n, i + 1)\n",
|
| 167 |
+
" plt.imshow(frame, cmap='gray')\n",
|
| 168 |
+
" plt.axis('off')\n",
|
| 169 |
+
"plt.show()\n",
|
| 170 |
+
"\n",
|
| 171 |
+
"print('Action class: ' + videos[90]['category'])"
|
| 172 |
+
]
|
| 173 |
+
},
|
| 174 |
+
{
|
| 175 |
+
"cell_type": "markdown",
|
| 176 |
+
"metadata": {
|
| 177 |
+
"id": "oOyC0Tolsu1r"
|
| 178 |
+
},
|
| 179 |
+
"source": [
|
| 180 |
+
"### Dataset and DataLoader\n",
|
| 181 |
+
"Since KTH dataset is not available within the torchvision's collection of datasets, we have to create our own dataset class to train an action recognition model. As discussed in the lecture, we will build two different models for that. The first one is based on 2D CNN, which will consider each frame as an image and will classify each frame into one of the action classes. Therefore, this model will work as an image classification model rather than a video classification model. For the second model, we will use 3D CNN which will consider a sequence or block of frames and intend to classify the entire sequence of frames into one of the action classes. In contrast with the first model, the second model will consider temporal information and will act as a true action recognition model.\n",
|
| 182 |
+
"\n",
|
| 183 |
+
"Therefore, in order to feed appropriate data, we will design two different types of dataset: (1) **SingleFrameDataset:** the first will return single frame which we will consider as an individual image, (2) **BlockFrameDataset:** block or sequence of frames where temporal information will be considered. Below, we have the two datasets for in PyTorch format."
|
| 184 |
+
]
|
| 185 |
+
},
|
| 186 |
+
{
|
| 187 |
+
"cell_type": "markdown",
|
| 188 |
+
"metadata": {
|
| 189 |
+
"id": "_w7xu7m-D6lZ"
|
| 190 |
+
},
|
| 191 |
+
"source": [
|
| 192 |
+
"#### Single frame dataset"
|
| 193 |
+
]
|
| 194 |
+
},
|
| 195 |
+
{
|
| 196 |
+
"cell_type": "markdown",
|
| 197 |
+
"metadata": {
|
| 198 |
+
"id": "p30M6x3LD3L0"
|
| 199 |
+
},
|
| 200 |
+
"source": [
|
| 201 |
+
"This dataset is used for training the single frame model."
|
| 202 |
+
]
|
| 203 |
+
},
|
| 204 |
+
{
|
| 205 |
+
"cell_type": "code",
|
| 206 |
+
"execution_count": null,
|
| 207 |
+
"metadata": {
|
| 208 |
+
"id": "kJCZkUw0sxKv"
|
| 209 |
+
},
|
| 210 |
+
"outputs": [],
|
| 211 |
+
"source": [
|
| 212 |
+
"import os\n",
|
| 213 |
+
"import torch\n",
|
| 214 |
+
"import pickle\n",
|
| 215 |
+
"import numpy as np\n",
|
| 216 |
+
"from torch.utils.data import Dataset\n",
|
| 217 |
+
"\n",
|
| 218 |
+
"class SingleFrameDataset(Dataset):\n",
|
| 219 |
+
" def __init__(self, directory, dataset=\"train\"):\n",
|
| 220 |
+
" self.instances, self.labels = self.read_dataset(directory, dataset)\n",
|
| 221 |
+
" # convert them into tensor\n",
|
| 222 |
+
" self.instances = torch.from_numpy(self.instances)\n",
|
| 223 |
+
" self.labels = torch.from_numpy(self.labels)\n",
|
| 224 |
+
" # normalize\n",
|
| 225 |
+
" self.zero_center()\n",
|
| 226 |
+
"\n",
|
| 227 |
+
" def __len__(self):\n",
|
| 228 |
+
" return self.instances.shape[0]\n",
|
| 229 |
+
"\n",
|
| 230 |
+
" def __getitem__(self, idx):\n",
|
| 231 |
+
" return self.instances[idx], self.labels[idx]\n",
|
| 232 |
+
"\n",
|
| 233 |
+
" def zero_center(self):\n",
|
| 234 |
+
" self.instances -= float(self.mean)\n",
|
| 235 |
+
"\n",
|
| 236 |
+
" def read_dataset(self, directory, dataset=\"train\"):\n",
|
| 237 |
+
" # set paths according to split\n",
|
| 238 |
+
" if dataset == \"train\":\n",
|
| 239 |
+
" filepath = os.path.join(directory, \"train.pickle\")\n",
|
| 240 |
+
" elif dataset == \"val\":\n",
|
| 241 |
+
" filepath = os.path.join(directory, \"val.pickle\")\n",
|
| 242 |
+
" else:\n",
|
| 243 |
+
" filepath = os.path.join(directory, \"test.pickle\")\n",
|
| 244 |
+
" #read the pickle file\n",
|
| 245 |
+
" videos = pickle.load(open(filepath, \"rb\"))\n",
|
| 246 |
+
" # accumulate the instances and label\n",
|
| 247 |
+
" instances = []\n",
|
| 248 |
+
" labels = []\n",
|
| 249 |
+
" for video in videos:\n",
|
| 250 |
+
" for frame in video[\"frames\"]:\n",
|
| 251 |
+
" instances.append(frame.reshape((1, 60, 80)))\n",
|
| 252 |
+
" labels.append(CATEGORY_INDEX[video[\"category\"]])\n",
|
| 253 |
+
" # numpy array\n",
|
| 254 |
+
" instances = np.array(instances, dtype=np.float32)\n",
|
| 255 |
+
" labels = np.array(labels, dtype=np.uint8)\n",
|
| 256 |
+
" self.mean = np.mean(instances)\n",
|
| 257 |
+
" return instances, labels"
|
| 258 |
+
]
|
| 259 |
+
},
|
| 260 |
+
{
|
| 261 |
+
"cell_type": "markdown",
|
| 262 |
+
"metadata": {
|
| 263 |
+
"id": "qKKppTFJEDMK"
|
| 264 |
+
},
|
| 265 |
+
"source": [
|
| 266 |
+
"#### Block Frame Dataset"
|
| 267 |
+
]
|
| 268 |
+
},
|
| 269 |
+
{
|
| 270 |
+
"cell_type": "markdown",
|
| 271 |
+
"metadata": {
|
| 272 |
+
"id": "oWUpk0tytWjH"
|
| 273 |
+
},
|
| 274 |
+
"source": [
|
| 275 |
+
"This dataset is used for training the block frame model."
|
| 276 |
+
]
|
| 277 |
+
},
|
| 278 |
+
{
|
| 279 |
+
"cell_type": "code",
|
| 280 |
+
"execution_count": null,
|
| 281 |
+
"metadata": {
|
| 282 |
+
"id": "J0uYc1o_tcFm"
|
| 283 |
+
},
|
| 284 |
+
"outputs": [],
|
| 285 |
+
"source": [
|
| 286 |
+
"import os\n",
|
| 287 |
+
"import torch\n",
|
| 288 |
+
"import pickle\n",
|
| 289 |
+
"import numpy as np\n",
|
| 290 |
+
"from torch.utils.data import Dataset\n",
|
| 291 |
+
"\n",
|
| 292 |
+
"class BlockFrameDataset(Dataset):\n",
|
| 293 |
+
" def __init__(self, directory, dataset=\"train\"):\n",
|
| 294 |
+
" self.instances, self.labels = self.read_dataset(directory, dataset)\n",
|
| 295 |
+
" # convert them into tensor\n",
|
| 296 |
+
" self.instances = torch.from_numpy(self.instances)\n",
|
| 297 |
+
" self.labels = torch.from_numpy(self.labels)\n",
|
| 298 |
+
" # normalize\n",
|
| 299 |
+
" self.zero_center()\n",
|
| 300 |
+
"\n",
|
| 301 |
+
" def __len__(self):\n",
|
| 302 |
+
" return self.instances.shape[0]\n",
|
| 303 |
+
"\n",
|
| 304 |
+
" def __getitem__(self, idx):\n",
|
| 305 |
+
" return self.instances[idx], self.labels[idx]\n",
|
| 306 |
+
"\n",
|
| 307 |
+
" def zero_center(self):\n",
|
| 308 |
+
" self.instances -= float(self.mean)\n",
|
| 309 |
+
"\n",
|
| 310 |
+
" def read_dataset(self, directory, dataset=\"train\", mean=None):\n",
|
| 311 |
+
" # set paths according to split\n",
|
| 312 |
+
" if dataset == \"train\":\n",
|
| 313 |
+
" filepath = os.path.join(directory, \"train.pickle\")\n",
|
| 314 |
+
" elif dataset == \"val\":\n",
|
| 315 |
+
" filepath = os.path.join(directory, \"val.pickle\")\n",
|
| 316 |
+
" else:\n",
|
| 317 |
+
" filepath = os.path.join(directory, \"test.pickle\")\n",
|
| 318 |
+
" # read the pickle file\n",
|
| 319 |
+
" videos = pickle.load(open(filepath, \"rb\"))\n",
|
| 320 |
+
" # accumulate the instances and label\n",
|
| 321 |
+
" instances = []\n",
|
| 322 |
+
" labels = []\n",
|
| 323 |
+
" current_block = []\n",
|
| 324 |
+
" for video in videos:\n",
|
| 325 |
+
" for i, frame in enumerate(video[\"frames\"]):\n",
|
| 326 |
+
" current_block.append(frame)\n",
|
| 327 |
+
" # 15 consecutive frames\n",
|
| 328 |
+
" if len(current_block) % 15 == 0:\n",
|
| 329 |
+
" current_block = np.array(current_block)\n",
|
| 330 |
+
" instances.append(current_block.reshape((1, 15, 60, 80)))\n",
|
| 331 |
+
" current_block = []\n",
|
| 332 |
+
" labels.append(CATEGORY_INDEX[video[\"category\"]])\n",
|
| 333 |
+
" # numpy array\n",
|
| 334 |
+
" instances = np.array(instances, dtype=np.float32)\n",
|
| 335 |
+
" labels = np.array(labels, dtype=np.uint8)\n",
|
| 336 |
+
" self.mean = np.mean(instances)\n",
|
| 337 |
+
" return instances, labels"
|
| 338 |
+
]
|
| 339 |
+
},
|
| 340 |
+
{
|
| 341 |
+
"cell_type": "markdown",
|
| 342 |
+
"metadata": {
|
| 343 |
+
"id": "YE5seYpPxY_f"
|
| 344 |
+
},
|
| 345 |
+
"source": [
|
| 346 |
+
"## Models\n",
|
| 347 |
+
"\n",
|
| 348 |
+
"As mentioned above, we will create two different models. (1) **SingleFrameModel:** the first one is based on 2D CNN, which will consider each single frame as an image and will classify each frame into one of the action classes. In other words, this model will work as an image classification model rather than a video classification model, because it will not consider any temporal information. (2) **BlockFrameModel:** the second model we will use 3D CNN which will consider a sequence or block of frames and intend to classify the entire sequence of frames into one of the action classes. In contrast with the first model, the second model will consider temporal information and will act as a true action recognition model. Therefore, it is expected that the BlockFrameModel works better than the SingleFrameModel for the action classification task."
|
| 349 |
+
]
|
| 350 |
+
},
|
| 351 |
+
{
|
| 352 |
+
"cell_type": "markdown",
|
| 353 |
+
"metadata": {
|
| 354 |
+
"id": "eKlLD8NFFYoA"
|
| 355 |
+
},
|
| 356 |
+
"source": [
|
| 357 |
+
"### Single Frame Model"
|
| 358 |
+
]
|
| 359 |
+
},
|
| 360 |
+
{
|
| 361 |
+
"cell_type": "markdown",
|
| 362 |
+
"metadata": {
|
| 363 |
+
"id": "DtaTfXruFaZp"
|
| 364 |
+
},
|
| 365 |
+
"source": [
|
| 366 |
+
"Below we implement the single frame model. If this model is tested after training it for 20 epochs, the accuracy of this model on the test set should be around 55%. This accuracy could be increased if you train it longer. In this model, we are going to use the following functions or modules:\n",
|
| 367 |
+
"\n",
|
| 368 |
+
"* `nn.Sequential()`: It is a sequential container. Modules will be added to it in the order they are passed in the constructor. Please check the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Sequential.html#torch.nn.Sequential) for more details.\n",
|
| 369 |
+
"\n",
|
| 370 |
+
"* `nn.Conv2d()`: It is a PyTorch module that applies a 2D convolution over an input signal composed of several input planes. More details are available on the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Conv2d.html).\n",
|
| 371 |
+
"\n",
|
| 372 |
+
"* `nn.BatchNorm2d()`: This module applies batch normalization over a 4D input as described in the [Batch Normalization paper](https://arxiv.org/abs/1502.03167). More details can be found in the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.BatchNorm2d.html).\n",
|
| 373 |
+
"\n",
|
| 374 |
+
"* `nn.MaxPool2d()`: It is also a module that applies a 2D max pooling over an input signal composed of several input planes. Please have a look on this [documentation](https://pytorch.org/docs/stable/generated/torch.nn.MaxPool2d.html) for more details.\n",
|
| 375 |
+
"\n",
|
| 376 |
+
"* `nn.Linear()`: It is a module that applies a linear transformation to the incoming data. More details can be found in its [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Linear.html#linear).\n",
|
| 377 |
+
"\n",
|
| 378 |
+
"* `nn.ReLU()`: It is also a module that applies element-wise the rectified linear unit function. Its [documentation](https://pytorch.org/docs/stable/generated/torch.nn.ReLU.html#relu) can explain more.\n",
|
| 379 |
+
"\n",
|
| 380 |
+
"* `nn.Dropout()`: This module randomly zeroes some of the elements of the input tensor with probability `p`. Check the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Dropout.html#dropout) for more details."
|
| 381 |
+
]
|
| 382 |
+
},
|
| 383 |
+
{
|
| 384 |
+
"cell_type": "code",
|
| 385 |
+
"execution_count": null,
|
| 386 |
+
"metadata": {
|
| 387 |
+
"id": "w_mka7mhwppH"
|
| 388 |
+
},
|
| 389 |
+
"outputs": [],
|
| 390 |
+
"source": [
|
| 391 |
+
"import torch.nn as nn\n",
|
| 392 |
+
"\n",
|
| 393 |
+
"class SingleFrameModel(nn.Module):\n",
|
| 394 |
+
" def __init__(self, n_classes):\n",
|
| 395 |
+
" super(SingleFrameModel, self).__init__()\n",
|
| 396 |
+
"\n",
|
| 397 |
+
" self.conv = nn.Sequential(\n",
|
| 398 |
+
" nn.Conv2d(1, 16, kernel_size=5),\n",
|
| 399 |
+
" nn.BatchNorm2d(16),\n",
|
| 400 |
+
" nn.ReLU(),\n",
|
| 401 |
+
" nn.MaxPool2d(kernel_size=2),\n",
|
| 402 |
+
" nn.Dropout(0.5),\n",
|
| 403 |
+
" nn.Conv2d(16, 32, kernel_size=3),\n",
|
| 404 |
+
" nn.BatchNorm2d(32),\n",
|
| 405 |
+
" nn.ReLU(),\n",
|
| 406 |
+
" nn.MaxPool2d(kernel_size=2),\n",
|
| 407 |
+
" nn.Dropout(0.5),\n",
|
| 408 |
+
" nn.Conv2d(32, 64, kernel_size=3),\n",
|
| 409 |
+
" nn.BatchNorm2d(64),\n",
|
| 410 |
+
" nn.ReLU(),\n",
|
| 411 |
+
" nn.MaxPool2d(kernel_size=2),\n",
|
| 412 |
+
" nn.Dropout(0.5))\n",
|
| 413 |
+
"\n",
|
| 414 |
+
" self.fc = nn.Sequential(\n",
|
| 415 |
+
" nn.Linear(2560, 128),\n",
|
| 416 |
+
" nn.ReLU(),\n",
|
| 417 |
+
" nn.Dropout(0.5),\n",
|
| 418 |
+
" nn.Linear(128, n_classes))\n",
|
| 419 |
+
"\n",
|
| 420 |
+
" def forward(self, x):\n",
|
| 421 |
+
" out = self.conv(x)\n",
|
| 422 |
+
" out = out.view(out.size(0), -1)\n",
|
| 423 |
+
" out = self.fc(out)\n",
|
| 424 |
+
"\n",
|
| 425 |
+
" return out"
|
| 426 |
+
]
|
| 427 |
+
},
|
| 428 |
+
{
|
| 429 |
+
"cell_type": "markdown",
|
| 430 |
+
"metadata": {
|
| 431 |
+
"id": "5MjQIJjuxglx"
|
| 432 |
+
},
|
| 433 |
+
"source": [
|
| 434 |
+
"### Block Frame Model\n",
|
| 435 |
+
"Below we implement the 2nd model that considers sequence of frames. For each video, we devide it into blocks of 15 contiguous frames. The model is then trained on these blocks instead of individual frame. In the convolutional layers, we use 3D convolutional filters (i.e. 3D CNN) to train the model to learn to detect temporal features.\n",
|
| 436 |
+
"\n",
|
| 437 |
+
"To classify a video, we also divide it into blocks of 15 contiguous frames. We then run the model on each block to get the block's vector of class probabilities. If this model is tested after training the model for 20 epochs the obtained accuracy on the test set should be around 67%. This means that the model is able to detect capture temporal information appeared in consecutive frames. However, this accuracy could be increased further by training it longer. In this model, we are going to use the following functions or modules:\n",
|
| 438 |
+
"\n",
|
| 439 |
+
"* `nn.Sequential()`: It is a sequential container. Modules will be added to it in the order they are passed in the constructor. Please check the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Sequential.html#torch.nn.Sequential) for more details.\n",
|
| 440 |
+
"\n",
|
| 441 |
+
"* `nn.Conv3d()`: It is a PyTorch module that applies a 3D convolution over an input signal composed of several input planes. More details are available on the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Conv3d.html).\n",
|
| 442 |
+
"\n",
|
| 443 |
+
"* `nn.BatchNorm3d()`: This module applies batch normalization over a 5D input as described in the [Batch Normalization paper](https://arxiv.org/abs/1502.03167). More details can be found in the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.BatchNorm3d.html).\n",
|
| 444 |
+
"\n",
|
| 445 |
+
"* `nn.MaxPool3d()`: It is also a module that applies a 3D max pooling over an input signal composed of several input planes. Please have a look on this [documentation](https://pytorch.org/docs/stable/generated/torch.nn.MaxPool3d.html) for more details.\n",
|
| 446 |
+
"\n",
|
| 447 |
+
"* `nn.Linear()`: It is a module that applies a linear transformation to the incoming data. More details can be found in its [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Linear.html#linear).\n",
|
| 448 |
+
"\n",
|
| 449 |
+
"* `nn.ReLU()`: It is also a module that applies element-wise the rectified linear unit function. Its [documentation](https://pytorch.org/docs/stable/generated/torch.nn.ReLU.html#relu) can explain more.\n",
|
| 450 |
+
"\n",
|
| 451 |
+
"* `nn.Dropout()`: This module randomly zeroes some of the elements of the input tensor with probability `p`. Check the [documentation](https://pytorch.org/docs/stable/generated/torch.nn.Dropout.html#dropout) for more details."
|
| 452 |
+
]
|
| 453 |
+
},
|
| 454 |
+
{
|
| 455 |
+
"cell_type": "code",
|
| 456 |
+
"execution_count": null,
|
| 457 |
+
"metadata": {
|
| 458 |
+
"id": "2_FP6ZyGwLwC"
|
| 459 |
+
},
|
| 460 |
+
"outputs": [],
|
| 461 |
+
"source": [
|
| 462 |
+
"import torch.nn as nn\n",
|
| 463 |
+
"\n",
|
| 464 |
+
"class BlockFrameModel(nn.Module):\n",
|
| 465 |
+
" def __init__(self, n_classes):\n",
|
| 466 |
+
" super(BlockFrameModel, self).__init__()\n",
|
| 467 |
+
"\n",
|
| 468 |
+
" self.conv = nn.Sequential(\n",
|
| 469 |
+
" nn.Conv3d(1, 16, kernel_size=(4, 5, 5)),\n",
|
| 470 |
+
" nn.BatchNorm3d(16),\n",
|
| 471 |
+
" nn.ReLU(),\n",
|
| 472 |
+
" nn.MaxPool3d(kernel_size=(1, 2, 2)),\n",
|
| 473 |
+
" nn.Dropout(0.5),\n",
|
| 474 |
+
" nn.Conv3d(16, 32, kernel_size=(4, 3, 3)),\n",
|
| 475 |
+
" nn.BatchNorm3d(32),\n",
|
| 476 |
+
" nn.ReLU(),\n",
|
| 477 |
+
" nn.MaxPool3d(kernel_size=(2, 2, 2)),\n",
|
| 478 |
+
" nn.Dropout(0.5),\n",
|
| 479 |
+
" nn.Conv3d(32, 64, kernel_size=(3, 3, 3)),\n",
|
| 480 |
+
" nn.BatchNorm3d(64),\n",
|
| 481 |
+
" nn.ReLU(),\n",
|
| 482 |
+
" nn.MaxPool3d(kernel_size=(2, 2, 2)),\n",
|
| 483 |
+
" nn.Dropout(0.5))\n",
|
| 484 |
+
"\n",
|
| 485 |
+
" self.fc = nn.Sequential(\n",
|
| 486 |
+
" nn.Linear(2560, 128),\n",
|
| 487 |
+
" nn.ReLU(),\n",
|
| 488 |
+
" nn.Dropout(0.5),\n",
|
| 489 |
+
" nn.Linear(128, n_classes))\n",
|
| 490 |
+
"\n",
|
| 491 |
+
" def forward(self, x):\n",
|
| 492 |
+
" out = self.conv(x)\n",
|
| 493 |
+
" out = out.view(out.size(0), -1)\n",
|
| 494 |
+
" out = self.fc(out)\n",
|
| 495 |
+
" return out"
|
| 496 |
+
]
|
| 497 |
+
},
|
| 498 |
+
{
|
| 499 |
+
"cell_type": "markdown",
|
| 500 |
+
"metadata": {
|
| 501 |
+
"id": "oap6-7RKG0s7"
|
| 502 |
+
},
|
| 503 |
+
"source": [
|
| 504 |
+
"### Average Meter\n",
|
| 505 |
+
"It is a simple class for keeping training statistics, such as losses and accuracies etc. The `.val` field usually holds the statistics for the current batch, whereas the `.avg` field hold statistics for the current epoch."
|
| 506 |
+
]
|
| 507 |
+
},
|
| 508 |
+
{
|
| 509 |
+
"cell_type": "code",
|
| 510 |
+
"execution_count": null,
|
| 511 |
+
"metadata": {
|
| 512 |
+
"id": "JeLH7fbOHDhH"
|
| 513 |
+
},
|
| 514 |
+
"outputs": [],
|
| 515 |
+
"source": [
|
| 516 |
+
"class AverageMeter(object):\n",
|
| 517 |
+
" \"\"\"Computes and stores the average and current value\"\"\"\n",
|
| 518 |
+
" def __init__(self):\n",
|
| 519 |
+
" self.reset()\n",
|
| 520 |
+
"\n",
|
| 521 |
+
" def reset(self):\n",
|
| 522 |
+
" self.val = 0\n",
|
| 523 |
+
" self.avg = 0\n",
|
| 524 |
+
" self.sum = 0\n",
|
| 525 |
+
" self.count = 0\n",
|
| 526 |
+
"\n",
|
| 527 |
+
" def update(self, val, n=1):\n",
|
| 528 |
+
" self.val = val\n",
|
| 529 |
+
" self.sum += val * n\n",
|
| 530 |
+
" self.count += n\n",
|
| 531 |
+
" self.avg = self.sum / self.count"
|
| 532 |
+
]
|
| 533 |
+
},
|
| 534 |
+
{
|
| 535 |
+
"cell_type": "markdown",
|
| 536 |
+
"metadata": {
|
| 537 |
+
"id": "F-zN4n7LJrAY"
|
| 538 |
+
},
|
| 539 |
+
"source": [
|
| 540 |
+
"### Train and Test Functions\n",
|
| 541 |
+
"Dataset/model independent train and test functions."
|
| 542 |
+
]
|
| 543 |
+
},
|
| 544 |
+
{
|
| 545 |
+
"cell_type": "code",
|
| 546 |
+
"execution_count": null,
|
| 547 |
+
"metadata": {
|
| 548 |
+
"id": "SrWHdjf3waDZ"
|
| 549 |
+
},
|
| 550 |
+
"outputs": [],
|
| 551 |
+
"source": [
|
| 552 |
+
"from tqdm.notebook import tqdm\n",
|
| 553 |
+
"import torch.nn.functional as F\n",
|
| 554 |
+
"##define train function\n",
|
| 555 |
+
"def train(model, data_loader, optimizer, device):\n",
|
| 556 |
+
" # meter\n",
|
| 557 |
+
" loss_meter = AverageMeter()\n",
|
| 558 |
+
" # switch to train mode\n",
|
| 559 |
+
" model.train()\n",
|
| 560 |
+
" tk = tqdm(data_loader, total=int(len(data_loader)), desc='Training', unit='frames', leave=False)\n",
|
| 561 |
+
" for batch_idx, data in enumerate(tk):\n",
|
| 562 |
+
" # fetch the data\n",
|
| 563 |
+
" frame, label = data[0], data[1]\n",
|
| 564 |
+
" # after fetching the data, transfer the model to the\n",
|
| 565 |
+
" # required device, in this example the device is gpu\n",
|
| 566 |
+
" # transfer to gpu can also be done by\n",
|
| 567 |
+
" frame, label = frame.to(device), label.to(device)\n",
|
| 568 |
+
" # compute the forward pass\n",
|
| 569 |
+
" output = model(frame)\n",
|
| 570 |
+
" # compute the loss function\n",
|
| 571 |
+
" loss_this = F.cross_entropy(output, label)\n",
|
| 572 |
+
" # initialize the optimizer\n",
|
| 573 |
+
" optimizer.zero_grad()\n",
|
| 574 |
+
" # compute the backward pass\n",
|
| 575 |
+
" loss_this.backward()\n",
|
| 576 |
+
" # update the parameters\n",
|
| 577 |
+
" optimizer.step()\n",
|
| 578 |
+
" # update the loss meter\n",
|
| 579 |
+
" loss_meter.update(loss_this.item(), label.shape[0])\n",
|
| 580 |
+
" tk.set_postfix({\"loss\": loss_meter.avg})\n",
|
| 581 |
+
" print('Train: Average loss: {:.4f}\\n'.format(loss_meter.avg))\n",
|
| 582 |
+
"\n",
|
| 583 |
+
"##define test function\n",
|
| 584 |
+
"def test(model, data_loader, device):\n",
|
| 585 |
+
" # meters\n",
|
| 586 |
+
" loss_meter = AverageMeter()\n",
|
| 587 |
+
" acc_meter = AverageMeter()\n",
|
| 588 |
+
" # switch to test mode\n",
|
| 589 |
+
" correct = 0\n",
|
| 590 |
+
" model.eval()\n",
|
| 591 |
+
" tk = tqdm(data_loader, total=int(len(data_loader)), desc='Test', unit='frames', leave=False)\n",
|
| 592 |
+
" for batch_idx, data in enumerate(tk):\n",
|
| 593 |
+
" # fetch the data\n",
|
| 594 |
+
" frame, label = data[0], data[1]\n",
|
| 595 |
+
" # after fetching the data transfer the model to the\n",
|
| 596 |
+
" # required device, in this example the device is gpu\n",
|
| 597 |
+
" # transfer to gpu can also be done by\n",
|
| 598 |
+
" frame, label = frame.to(device), label.to(device)\n",
|
| 599 |
+
" # since we dont need to backpropagate loss in testing,\n",
|
| 600 |
+
" # we dont keep the gradient\n",
|
| 601 |
+
" with torch.no_grad():\n",
|
| 602 |
+
" output = model(frame)\n",
|
| 603 |
+
" # compute the loss function just for checking\n",
|
| 604 |
+
" loss_this = F.cross_entropy(output, label)\n",
|
| 605 |
+
" # get the index of the max log-probability\n",
|
| 606 |
+
" pred = output.argmax(dim=1, keepdim=True)\n",
|
| 607 |
+
" # check which of the predictions are correct\n",
|
| 608 |
+
" correct_this = pred.eq(label.view_as(pred)).sum().item()\n",
|
| 609 |
+
" # accumulate the correct ones\n",
|
| 610 |
+
" correct += correct_this\n",
|
| 611 |
+
" # compute accuracy\n",
|
| 612 |
+
" acc_this = correct_this / label.shape[0] * 100.0\n",
|
| 613 |
+
" # update the loss and accuracy meter\n",
|
| 614 |
+
" acc_meter.update(acc_this, label.shape[0])\n",
|
| 615 |
+
" loss_meter.update(loss_this.item(), label.shape[0])\n",
|
| 616 |
+
" print('Test: Average loss: {:.4f}, Accuracy: {}/{} ({:.2f}%)\\n'.format(\n",
|
| 617 |
+
" loss_meter.avg, correct, len(data_loader.dataset), acc_meter.avg))"
|
| 618 |
+
]
|
| 619 |
+
},
|
| 620 |
+
{
|
| 621 |
+
"cell_type": "markdown",
|
| 622 |
+
"metadata": {
|
| 623 |
+
"id": "6s8RdNhKSzK1"
|
| 624 |
+
},
|
| 625 |
+
"source": [
|
| 626 |
+
"## Train and Test"
|
| 627 |
+
]
|
| 628 |
+
},
|
| 629 |
+
{
|
| 630 |
+
"cell_type": "markdown",
|
| 631 |
+
"metadata": {
|
| 632 |
+
"id": "mYQgGE9gS2PC"
|
| 633 |
+
},
|
| 634 |
+
"source": [
|
| 635 |
+
"### Single Frame Model"
|
| 636 |
+
]
|
| 637 |
+
},
|
| 638 |
+
{
|
| 639 |
+
"cell_type": "markdown",
|
| 640 |
+
"metadata": {
|
| 641 |
+
"id": "sMWS0PreVZWU"
|
| 642 |
+
},
|
| 643 |
+
"source": [
|
| 644 |
+
"#### Parameters and Model Instantiation\n",
|
| 645 |
+
"Select the correct dataset and model."
|
| 646 |
+
]
|
| 647 |
+
},
|
| 648 |
+
{
|
| 649 |
+
"cell_type": "code",
|
| 650 |
+
"execution_count": null,
|
| 651 |
+
"metadata": {
|
| 652 |
+
"id": "4mAUBPogU6aY"
|
| 653 |
+
},
|
| 654 |
+
"outputs": [],
|
| 655 |
+
"source": [
|
| 656 |
+
"# 1. Create Dataset\n",
|
| 657 |
+
"from pathlib import Path\n",
|
| 658 |
+
"dir_pickle = Path('KTH_pickle/')\n",
|
| 659 |
+
"\n",
|
| 660 |
+
"train_set = SingleFrameDataset(dir_pickle, \"train\")\n",
|
| 661 |
+
"test_set = SingleFrameDataset(dir_pickle, \"test\")\n",
|
| 662 |
+
"\n",
|
| 663 |
+
"# 2. Create Dataloader\n",
|
| 664 |
+
"from torch.utils.data import DataLoader\n",
|
| 665 |
+
"batch_size = 64\n",
|
| 666 |
+
"loader_args = dict(batch_size=batch_size, num_workers=1, pin_memory=True)\n",
|
| 667 |
+
"train_loader = DataLoader(train_set, shuffle=True, **loader_args)\n",
|
| 668 |
+
"test_loader = DataLoader(test_set, shuffle=False, drop_last=True, **loader_args)\n",
|
| 669 |
+
"\n",
|
| 670 |
+
"# 3. Create Model\n",
|
| 671 |
+
"device = \"cuda\"\n",
|
| 672 |
+
"model = SingleFrameModel(n_classes=6)\n",
|
| 673 |
+
"model = model.to(device)\n",
|
| 674 |
+
"\n",
|
| 675 |
+
"# 4. Set up the optimizer, the loss, the learning rate scheduler and the loss scaling for AMP\n",
|
| 676 |
+
"from torch import optim\n",
|
| 677 |
+
"learning_rate = 0.001\n",
|
| 678 |
+
"optimizer = optim.Adam(model.parameters(), lr=learning_rate, weight_decay=1e-8)"
|
| 679 |
+
]
|
| 680 |
+
},
|
| 681 |
+
{
|
| 682 |
+
"cell_type": "markdown",
|
| 683 |
+
"metadata": {
|
| 684 |
+
"id": "-938nzOl1p4W"
|
| 685 |
+
},
|
| 686 |
+
"source": [
|
| 687 |
+
"#### Training Loop\n",
|
| 688 |
+
"Training loop containing 20 training epochs. Test accuracies should be around 55%."
|
| 689 |
+
]
|
| 690 |
+
},
|
| 691 |
+
{
|
| 692 |
+
"cell_type": "code",
|
| 693 |
+
"execution_count": null,
|
| 694 |
+
"metadata": {
|
| 695 |
+
"id": "ljL-MscZ1lou"
|
| 696 |
+
},
|
| 697 |
+
"outputs": [],
|
| 698 |
+
"source": [
|
| 699 |
+
"num_epoch = 20\n",
|
| 700 |
+
"for epoch in range(num_epoch):\n",
|
| 701 |
+
" train(model, train_loader, optimizer, device)\n",
|
| 702 |
+
"test(model, test_loader, device)"
|
| 703 |
+
]
|
| 704 |
+
},
|
| 705 |
+
{
|
| 706 |
+
"cell_type": "markdown",
|
| 707 |
+
"metadata": {
|
| 708 |
+
"id": "Ggun7cQCTAQ6"
|
| 709 |
+
},
|
| 710 |
+
"source": [
|
| 711 |
+
"### Block Frame Model"
|
| 712 |
+
]
|
| 713 |
+
},
|
| 714 |
+
{
|
| 715 |
+
"cell_type": "markdown",
|
| 716 |
+
"metadata": {
|
| 717 |
+
"id": "XdS9XXjuTJuk"
|
| 718 |
+
},
|
| 719 |
+
"source": [
|
| 720 |
+
"#### Parameters and Model Instantiation\n",
|
| 721 |
+
"Select the correct dataset and model."
|
| 722 |
+
]
|
| 723 |
+
},
|
| 724 |
+
{
|
| 725 |
+
"cell_type": "code",
|
| 726 |
+
"execution_count": null,
|
| 727 |
+
"metadata": {
|
| 728 |
+
"id": "bdAgghnaM4-M"
|
| 729 |
+
},
|
| 730 |
+
"outputs": [],
|
| 731 |
+
"source": [
|
| 732 |
+
"# 1. Create Dataset\n",
|
| 733 |
+
"from pathlib import Path\n",
|
| 734 |
+
"dir_pickle = Path('KTH_pickle/')\n",
|
| 735 |
+
"\n",
|
| 736 |
+
"train_set = BlockFrameDataset(dir_pickle, \"train\")\n",
|
| 737 |
+
"test_set = BlockFrameDataset(dir_pickle, \"test\")\n",
|
| 738 |
+
"\n",
|
| 739 |
+
"# 2. Create Dataloader\n",
|
| 740 |
+
"from torch.utils.data import DataLoader\n",
|
| 741 |
+
"batch_size = 64\n",
|
| 742 |
+
"loader_args = dict(batch_size=batch_size, num_workers=1, pin_memory=True)\n",
|
| 743 |
+
"train_loader = DataLoader(train_set, shuffle=True, **loader_args)\n",
|
| 744 |
+
"test_loader = DataLoader(test_set, shuffle=False, drop_last=True, **loader_args)\n",
|
| 745 |
+
"\n",
|
| 746 |
+
"# 3. Create Model\n",
|
| 747 |
+
"device = \"cuda\"\n",
|
| 748 |
+
"model = BlockFrameModel(n_classes=6)\n",
|
| 749 |
+
"model = model.to(device)\n",
|
| 750 |
+
"\n",
|
| 751 |
+
"# 4. Set up the optimizer, the loss, the learning rate scheduler and the loss scaling for AMP\n",
|
| 752 |
+
"from torch import optim\n",
|
| 753 |
+
"learning_rate = 0.001\n",
|
| 754 |
+
"optimizer = optim.Adam(model.parameters(), lr=learning_rate, weight_decay=1e-8)"
|
| 755 |
+
]
|
| 756 |
+
},
|
| 757 |
+
{
|
| 758 |
+
"cell_type": "markdown",
|
| 759 |
+
"metadata": {
|
| 760 |
+
"id": "TsPrD3fzYhG5"
|
| 761 |
+
},
|
| 762 |
+
"source": [
|
| 763 |
+
"#### Training Loop\n",
|
| 764 |
+
"Training loop containing 20 training epochs. Test accuracies should be around 67%."
|
| 765 |
+
]
|
| 766 |
+
},
|
| 767 |
+
{
|
| 768 |
+
"cell_type": "code",
|
| 769 |
+
"execution_count": null,
|
| 770 |
+
"metadata": {
|
| 771 |
+
"id": "qyf1Y_fZNDKi"
|
| 772 |
+
},
|
| 773 |
+
"outputs": [],
|
| 774 |
+
"source": [
|
| 775 |
+
"num_epoch = 20\n",
|
| 776 |
+
"for epoch in range(num_epoch):\n",
|
| 777 |
+
" train(model, train_loader, optimizer, device)\n",
|
| 778 |
+
"test(model, test_loader, device)"
|
| 779 |
+
]
|
| 780 |
+
},
|
| 781 |
+
{
|
| 782 |
+
"cell_type": "markdown",
|
| 783 |
+
"metadata": {
|
| 784 |
+
"id": "tTuOF4Fh2vpb"
|
| 785 |
+
},
|
| 786 |
+
"source": [
|
| 787 |
+
"### Conclusion\n",
|
| 788 |
+
"If all goes well, the `BlockFrameModel()` should achieve superior performance than the `SingleFrameModel()` because the latter does not consider temporal information in a video which is crucial for video recognition."
|
| 789 |
+
]
|
| 790 |
+
}
|
| 791 |
+
],
|
| 792 |
+
"metadata": {
|
| 793 |
+
"accelerator": "GPU",
|
| 794 |
+
"colab": {
|
| 795 |
+
"include_colab_link": true,
|
| 796 |
+
"name": "ECMM426/ECMM441 - Human Action Recognition.ipynb",
|
| 797 |
+
"private_outputs": true,
|
| 798 |
+
"provenance": []
|
| 799 |
+
},
|
| 800 |
+
"kernelspec": {
|
| 801 |
+
"display_name": "Python 3 (ipykernel)",
|
| 802 |
+
"language": "python",
|
| 803 |
+
"name": "python3"
|
| 804 |
+
},
|
| 805 |
+
"language_info": {
|
| 806 |
+
"codemirror_mode": {
|
| 807 |
+
"name": "ipython",
|
| 808 |
+
"version": 3
|
| 809 |
+
},
|
| 810 |
+
"file_extension": ".py",
|
| 811 |
+
"mimetype": "text/x-python",
|
| 812 |
+
"name": "python",
|
| 813 |
+
"nbconvert_exporter": "python",
|
| 814 |
+
"pygments_lexer": "ipython3",
|
| 815 |
+
"version": "3.12.3"
|
| 816 |
+
}
|
| 817 |
+
},
|
| 818 |
+
"nbformat": 4,
|
| 819 |
+
"nbformat_minor": 1
|
| 820 |
+
}
|
Downloads/.ipynb_checkpoints/PyTorch Tutorial-checkpoint.ipynb
ADDED
|
@@ -0,0 +1,2030 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {
|
| 6 |
+
"colab_type": "text",
|
| 7 |
+
"id": "view-in-github"
|
| 8 |
+
},
|
| 9 |
+
"source": [
|
| 10 |
+
"<a href=\"https://colab.research.google.com/github/AnjanDutta/EEEM068/blob/main/Notebooks/PyTorch_Tutorial.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
|
| 11 |
+
]
|
| 12 |
+
},
|
| 13 |
+
{
|
| 14 |
+
"cell_type": "markdown",
|
| 15 |
+
"metadata": {
|
| 16 |
+
"id": "hpDRECpayKDt"
|
| 17 |
+
},
|
| 18 |
+
"source": [
|
| 19 |
+
"<H1 style=\"text-align: center\">EEEM068 - Applied Machine Learning</H1>\n",
|
| 20 |
+
"<H1 style=\"text-align: center\">Workshop 02</H1>\n",
|
| 21 |
+
"<H1 style=\"text-align: center\">PyTorch Tutorial</H1>"
|
| 22 |
+
]
|
| 23 |
+
},
|
| 24 |
+
{
|
| 25 |
+
"cell_type": "markdown",
|
| 26 |
+
"metadata": {
|
| 27 |
+
"id": "3aZW7eWZWg_j"
|
| 28 |
+
},
|
| 29 |
+
"source": [
|
| 30 |
+
"## Introduction"
|
| 31 |
+
]
|
| 32 |
+
},
|
| 33 |
+
{
|
| 34 |
+
"cell_type": "markdown",
|
| 35 |
+
"metadata": {
|
| 36 |
+
"id": "CQ9aAM9AyKDx"
|
| 37 |
+
},
|
| 38 |
+
"source": [
|
| 39 |
+
"PyTorch is an open source machine/deep learning framework that allows you to write your own neural networks and optimize them efficiently. However, PyTorch is not the only framework of its kind. Alternatives to PyTorch include [TensorFlow](https://www.tensorflow.org/), [JAX](https://github.com/google/jax#quickstart-colab-in-the-cloud) and [Caffe](http://caffe.berkeleyvision.org/) etc. We choose PyTorch because it is well established, popular and has a huge developer community (supported by Meta/Facebook), is very flexible and especially used in research."
|
| 40 |
+
]
|
| 41 |
+
},
|
| 42 |
+
{
|
| 43 |
+
"cell_type": "markdown",
|
| 44 |
+
"metadata": {
|
| 45 |
+
"id": "WDqXKGwVw5VQ"
|
| 46 |
+
},
|
| 47 |
+
"source": [
|
| 48 |
+
"## PyTorch Versions"
|
| 49 |
+
]
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"cell_type": "markdown",
|
| 53 |
+
"metadata": {
|
| 54 |
+
"id": "YNagAs00YbRO"
|
| 55 |
+
},
|
| 56 |
+
"source": [
|
| 57 |
+
"As of February 2026, the current stable version of PyTorch on Colab is 2.9.0, eventually with some CUDA version (Cuda 12.8). The PyTorch version can be checked with the following line of code, which first import PyTorch and then print the `__version__` variable. Note that the package is called `torch`, based on its original framework [Torch](http://torch.ch/)."
|
| 58 |
+
]
|
| 59 |
+
},
|
| 60 |
+
{
|
| 61 |
+
"cell_type": "code",
|
| 62 |
+
"execution_count": null,
|
| 63 |
+
"metadata": {
|
| 64 |
+
"id": "zDcKMfz1ZeQB"
|
| 65 |
+
},
|
| 66 |
+
"outputs": [],
|
| 67 |
+
"source": [
|
| 68 |
+
"import torch\n",
|
| 69 |
+
"print(\"Using torch\", torch.__version__)"
|
| 70 |
+
]
|
| 71 |
+
},
|
| 72 |
+
{
|
| 73 |
+
"cell_type": "markdown",
|
| 74 |
+
"metadata": {
|
| 75 |
+
"id": "iFPzDt_2yKD1"
|
| 76 |
+
},
|
| 77 |
+
"source": [
|
| 78 |
+
"## Basics of PyTorch"
|
| 79 |
+
]
|
| 80 |
+
},
|
| 81 |
+
{
|
| 82 |
+
"cell_type": "markdown",
|
| 83 |
+
"metadata": {
|
| 84 |
+
"id": "hMyWsVfWyKD2"
|
| 85 |
+
},
|
| 86 |
+
"source": [
|
| 87 |
+
"As in every machine learning framework, PyTorch provides functions that are stochastic like generating random numbers. However, a very good practice is to setup your code to be reproducible with the exact same random numbers. This is why we set a seed below."
|
| 88 |
+
]
|
| 89 |
+
},
|
| 90 |
+
{
|
| 91 |
+
"cell_type": "code",
|
| 92 |
+
"execution_count": null,
|
| 93 |
+
"metadata": {
|
| 94 |
+
"id": "jXJtposLyKD2"
|
| 95 |
+
},
|
| 96 |
+
"outputs": [],
|
| 97 |
+
"source": [
|
| 98 |
+
"torch.manual_seed(42) # Setting the seed"
|
| 99 |
+
]
|
| 100 |
+
},
|
| 101 |
+
{
|
| 102 |
+
"cell_type": "markdown",
|
| 103 |
+
"metadata": {
|
| 104 |
+
"id": "kUqFQ7q1bRao"
|
| 105 |
+
},
|
| 106 |
+
"source": [
|
| 107 |
+
"### Tensors"
|
| 108 |
+
]
|
| 109 |
+
},
|
| 110 |
+
{
|
| 111 |
+
"cell_type": "markdown",
|
| 112 |
+
"metadata": {
|
| 113 |
+
"id": "Me76FJ5cyKD2"
|
| 114 |
+
},
|
| 115 |
+
"source": [
|
| 116 |
+
"Tensors are the PyTorch equivalent to Numpy arrays, with the addition to also have support for GPU acceleration. The name \"tensor\" is a generalization of concepts you already know. For instance, a vector is a 1-D tensor, and a matrix a 2-D tensor. When working with neural networks, we will use tensors of various shapes and number of dimensions. Most common functions you know from Numpy can be used on tensors as well. Actually, since Numpy arrays are so similar to tensors, we can convert most tensors to Numpy arrays and back."
|
| 117 |
+
]
|
| 118 |
+
},
|
| 119 |
+
{
|
| 120 |
+
"cell_type": "markdown",
|
| 121 |
+
"metadata": {
|
| 122 |
+
"id": "STQ4KUSwbXRi"
|
| 123 |
+
},
|
| 124 |
+
"source": [
|
| 125 |
+
"#### Initialization"
|
| 126 |
+
]
|
| 127 |
+
},
|
| 128 |
+
{
|
| 129 |
+
"cell_type": "markdown",
|
| 130 |
+
"metadata": {
|
| 131 |
+
"id": "-ldy1KWFbaBi"
|
| 132 |
+
},
|
| 133 |
+
"source": [
|
| 134 |
+
"Let's first start by looking at different ways of creating a tensor. There are many possible options, the simplest one is to call `torch.Tensor` passing the desired shape as input argument:"
|
| 135 |
+
]
|
| 136 |
+
},
|
| 137 |
+
{
|
| 138 |
+
"cell_type": "code",
|
| 139 |
+
"execution_count": null,
|
| 140 |
+
"metadata": {
|
| 141 |
+
"id": "Ws4H24EYyKD3"
|
| 142 |
+
},
|
| 143 |
+
"outputs": [],
|
| 144 |
+
"source": [
|
| 145 |
+
"x = torch.Tensor(2, 3, 4) # Creates a tensor of shape [2, 3, 4]\n",
|
| 146 |
+
"print(x)"
|
| 147 |
+
]
|
| 148 |
+
},
|
| 149 |
+
{
|
| 150 |
+
"cell_type": "markdown",
|
| 151 |
+
"metadata": {
|
| 152 |
+
"id": "uVpDx-T0yKD3"
|
| 153 |
+
},
|
| 154 |
+
"source": [
|
| 155 |
+
"The function `torch.Tensor` allocates memory for the desired tensor, but reuses any values that have already been in the memory. To directly assign values to the tensor during initialization, there are many alternatives including:\n",
|
| 156 |
+
"\n",
|
| 157 |
+
"* `torch.zeros`: Creates a tensor filled with zeros\n",
|
| 158 |
+
"* `torch.ones`: Creates a tensor filled with ones\n",
|
| 159 |
+
"* `torch.rand`: Creates a tensor with random values uniformly sampled between 0 and 1\n",
|
| 160 |
+
"* `torch.randn`: Creates a tensor with random values sampled from a normal distribution with mean 0 and variance 1\n",
|
| 161 |
+
"* `torch.arange`: Creates a tensor containing the values $N,N+1,N+2,...,M$\n",
|
| 162 |
+
"* `torch.Tensor` (input list): Creates a tensor from the list elements you provide"
|
| 163 |
+
]
|
| 164 |
+
},
|
| 165 |
+
{
|
| 166 |
+
"cell_type": "code",
|
| 167 |
+
"execution_count": null,
|
| 168 |
+
"metadata": {
|
| 169 |
+
"id": "yQtpXfWeyKD4"
|
| 170 |
+
},
|
| 171 |
+
"outputs": [],
|
| 172 |
+
"source": [
|
| 173 |
+
"# Create a tensor from a (nested) list\n",
|
| 174 |
+
"x = torch.Tensor([[1, 2], [3, 4]])\n",
|
| 175 |
+
"print(x)"
|
| 176 |
+
]
|
| 177 |
+
},
|
| 178 |
+
{
|
| 179 |
+
"cell_type": "code",
|
| 180 |
+
"execution_count": null,
|
| 181 |
+
"metadata": {
|
| 182 |
+
"id": "8awQ54BPyKD4"
|
| 183 |
+
},
|
| 184 |
+
"outputs": [],
|
| 185 |
+
"source": [
|
| 186 |
+
"# Create a tensor with random values between 0 and 1 with the shape [2, 3, 4]\n",
|
| 187 |
+
"x = torch.rand(2, 3, 4)\n",
|
| 188 |
+
"print(x)"
|
| 189 |
+
]
|
| 190 |
+
},
|
| 191 |
+
{
|
| 192 |
+
"cell_type": "markdown",
|
| 193 |
+
"metadata": {
|
| 194 |
+
"id": "NJCmVHByyKD4"
|
| 195 |
+
},
|
| 196 |
+
"source": [
|
| 197 |
+
"You can obtain the shape of a tensor in the same way as in Numpy (`x.shape`), or using the `.size` method:"
|
| 198 |
+
]
|
| 199 |
+
},
|
| 200 |
+
{
|
| 201 |
+
"cell_type": "code",
|
| 202 |
+
"execution_count": null,
|
| 203 |
+
"metadata": {
|
| 204 |
+
"id": "kbshkJFlyKD5"
|
| 205 |
+
},
|
| 206 |
+
"outputs": [],
|
| 207 |
+
"source": [
|
| 208 |
+
"print(\"Shape:\", x.shape)\n",
|
| 209 |
+
"\n",
|
| 210 |
+
"print(\"Size:\", x.size())\n",
|
| 211 |
+
"\n",
|
| 212 |
+
"dim1, dim2, dim3 = x.size()\n",
|
| 213 |
+
"print(\"Size:\", dim1, dim2, dim3)"
|
| 214 |
+
]
|
| 215 |
+
},
|
| 216 |
+
{
|
| 217 |
+
"cell_type": "markdown",
|
| 218 |
+
"metadata": {
|
| 219 |
+
"id": "wqU-hvDMfzXE"
|
| 220 |
+
},
|
| 221 |
+
"source": [
|
| 222 |
+
"#### Tensor to Numpy, and Numpy to Tensor"
|
| 223 |
+
]
|
| 224 |
+
},
|
| 225 |
+
{
|
| 226 |
+
"cell_type": "markdown",
|
| 227 |
+
"metadata": {
|
| 228 |
+
"id": "mlgyRypPyKD5"
|
| 229 |
+
},
|
| 230 |
+
"source": [
|
| 231 |
+
"Tensors can be converted to Numpy arrays, and Numpy arrays back to tensors. To transform a Numpy array into a tensor, we can use the function `torch.from_numpy`:"
|
| 232 |
+
]
|
| 233 |
+
},
|
| 234 |
+
{
|
| 235 |
+
"cell_type": "code",
|
| 236 |
+
"execution_count": null,
|
| 237 |
+
"metadata": {
|
| 238 |
+
"id": "5So3mR1fyKD5"
|
| 239 |
+
},
|
| 240 |
+
"outputs": [],
|
| 241 |
+
"source": [
|
| 242 |
+
"import numpy as np\n",
|
| 243 |
+
"np_arr = np.array([[1, 2], [3, 4]])\n",
|
| 244 |
+
"tensor = torch.from_numpy(np_arr)\n",
|
| 245 |
+
"\n",
|
| 246 |
+
"print(\"Numpy array:\", np_arr)\n",
|
| 247 |
+
"print(\"PyTorch tensor:\", tensor)"
|
| 248 |
+
]
|
| 249 |
+
},
|
| 250 |
+
{
|
| 251 |
+
"cell_type": "markdown",
|
| 252 |
+
"metadata": {
|
| 253 |
+
"id": "9ZhbSdmWyKD5"
|
| 254 |
+
},
|
| 255 |
+
"source": [
|
| 256 |
+
"To transform a PyTorch tensor back to a Numpy array, we can use the function `.numpy()` on tensors:"
|
| 257 |
+
]
|
| 258 |
+
},
|
| 259 |
+
{
|
| 260 |
+
"cell_type": "code",
|
| 261 |
+
"execution_count": null,
|
| 262 |
+
"metadata": {
|
| 263 |
+
"id": "lHpBvt7ryKD5"
|
| 264 |
+
},
|
| 265 |
+
"outputs": [],
|
| 266 |
+
"source": [
|
| 267 |
+
"tensor = torch.arange(4)\n",
|
| 268 |
+
"np_arr = tensor.numpy()\n",
|
| 269 |
+
"\n",
|
| 270 |
+
"print(\"PyTorch tensor:\", tensor)\n",
|
| 271 |
+
"print(\"Numpy array:\", np_arr)"
|
| 272 |
+
]
|
| 273 |
+
},
|
| 274 |
+
{
|
| 275 |
+
"cell_type": "markdown",
|
| 276 |
+
"metadata": {
|
| 277 |
+
"id": "6sj4hQR1yKD6"
|
| 278 |
+
},
|
| 279 |
+
"source": [
|
| 280 |
+
"The conversion of tensors to Numpy require the tensor to be on the CPU, and not the GPU (more on GPU support in a later section). In case you have a tensor on GPU, you need to call `.cpu()` on the tensor beforehand. Hence, you get a line like `np_arr = tensor.cpu().numpy()`."
|
| 281 |
+
]
|
| 282 |
+
},
|
| 283 |
+
{
|
| 284 |
+
"cell_type": "markdown",
|
| 285 |
+
"metadata": {
|
| 286 |
+
"id": "Q_Z4B31kfuwG"
|
| 287 |
+
},
|
| 288 |
+
"source": [
|
| 289 |
+
"#### Operations"
|
| 290 |
+
]
|
| 291 |
+
},
|
| 292 |
+
{
|
| 293 |
+
"cell_type": "markdown",
|
| 294 |
+
"metadata": {
|
| 295 |
+
"id": "CyPh-BetyKD6"
|
| 296 |
+
},
|
| 297 |
+
"source": [
|
| 298 |
+
"Most operations that exist in Numpy, also exist in PyTorch. A full list of operations can be found in the [PyTorch documentation](https://pytorch.org/docs/stable/tensors.html#), but we will review the most important ones here.\n",
|
| 299 |
+
"\n",
|
| 300 |
+
"The simplest operation is to add two tensors:"
|
| 301 |
+
]
|
| 302 |
+
},
|
| 303 |
+
{
|
| 304 |
+
"cell_type": "code",
|
| 305 |
+
"execution_count": null,
|
| 306 |
+
"metadata": {
|
| 307 |
+
"id": "ymI_kL7WyKD6"
|
| 308 |
+
},
|
| 309 |
+
"outputs": [],
|
| 310 |
+
"source": [
|
| 311 |
+
"x1 = torch.rand(2, 3)\n",
|
| 312 |
+
"x2 = torch.rand(2, 3)\n",
|
| 313 |
+
"y = x1 + x2\n",
|
| 314 |
+
"\n",
|
| 315 |
+
"print(\"X1\", x1)\n",
|
| 316 |
+
"print(\"X2\", x2)\n",
|
| 317 |
+
"print(\"Y\", y)"
|
| 318 |
+
]
|
| 319 |
+
},
|
| 320 |
+
{
|
| 321 |
+
"cell_type": "markdown",
|
| 322 |
+
"metadata": {
|
| 323 |
+
"id": "FzxUmB52yKD6"
|
| 324 |
+
},
|
| 325 |
+
"source": [
|
| 326 |
+
"Calling `x1 + x2` creates a new tensor containing the sum of the two inputs. However, we can also use in-place operations that are applied directly on the memory of a tensor. We therefore change the values of `x2` without the chance to re-accessing the values of `x2` before the operation. An example is shown below:"
|
| 327 |
+
]
|
| 328 |
+
},
|
| 329 |
+
{
|
| 330 |
+
"cell_type": "code",
|
| 331 |
+
"execution_count": null,
|
| 332 |
+
"metadata": {
|
| 333 |
+
"id": "lDVRRhYUyKD7"
|
| 334 |
+
},
|
| 335 |
+
"outputs": [],
|
| 336 |
+
"source": [
|
| 337 |
+
"x1 = torch.rand(2, 3)\n",
|
| 338 |
+
"x2 = torch.rand(2, 3)\n",
|
| 339 |
+
"print(\"X1 (before)\", x1)\n",
|
| 340 |
+
"print(\"X2 (before)\", x2)\n",
|
| 341 |
+
"\n",
|
| 342 |
+
"x2.add_(x1)\n",
|
| 343 |
+
"print(\"X1 (after)\", x1)\n",
|
| 344 |
+
"print(\"X2 (after)\", x2)"
|
| 345 |
+
]
|
| 346 |
+
},
|
| 347 |
+
{
|
| 348 |
+
"cell_type": "markdown",
|
| 349 |
+
"metadata": {
|
| 350 |
+
"id": "7_zEJjhUyKD7"
|
| 351 |
+
},
|
| 352 |
+
"source": [
|
| 353 |
+
"In-place operations are usually marked with a underscore postfix (e.g. \"add_\" instead of \"add\")."
|
| 354 |
+
]
|
| 355 |
+
},
|
| 356 |
+
{
|
| 357 |
+
"cell_type": "markdown",
|
| 358 |
+
"metadata": {
|
| 359 |
+
"id": "fkls4uTMjrlx"
|
| 360 |
+
},
|
| 361 |
+
"source": [
|
| 362 |
+
"#### `torch.arange()`"
|
| 363 |
+
]
|
| 364 |
+
},
|
| 365 |
+
{
|
| 366 |
+
"cell_type": "markdown",
|
| 367 |
+
"metadata": {
|
| 368 |
+
"id": "pLbFkfC6hZNt"
|
| 369 |
+
},
|
| 370 |
+
"source": [
|
| 371 |
+
"`torch.arange()` returns a 1-D tensor of size $\\lceil{\\frac{\\text{end} - \\text{start}}{\\text{step}}}\\rceil$ with values from the interval `[start, end)` taken with common difference `step` beginning from `start`. It mostly works as the `range()` function in Numpy. For more details, please see the [documentation](https://pytorch.org/docs/stable/generated/torch.arange.html#torch-arange)."
|
| 372 |
+
]
|
| 373 |
+
},
|
| 374 |
+
{
|
| 375 |
+
"cell_type": "code",
|
| 376 |
+
"execution_count": null,
|
| 377 |
+
"metadata": {
|
| 378 |
+
"id": "semrM0JbyKD7"
|
| 379 |
+
},
|
| 380 |
+
"outputs": [],
|
| 381 |
+
"source": [
|
| 382 |
+
"x = torch.arange(6)\n",
|
| 383 |
+
"print(\"X\", x)"
|
| 384 |
+
]
|
| 385 |
+
},
|
| 386 |
+
{
|
| 387 |
+
"cell_type": "markdown",
|
| 388 |
+
"metadata": {
|
| 389 |
+
"id": "YszyduPpkUBZ"
|
| 390 |
+
},
|
| 391 |
+
"source": [
|
| 392 |
+
"#### `.view()`"
|
| 393 |
+
]
|
| 394 |
+
},
|
| 395 |
+
{
|
| 396 |
+
"cell_type": "markdown",
|
| 397 |
+
"metadata": {
|
| 398 |
+
"id": "p5pKriuXg8VO"
|
| 399 |
+
},
|
| 400 |
+
"source": [
|
| 401 |
+
"A tensor of size (2, 3) can be re-organized to any other shape with the same number of elements (e.g. a tensor of size (6), or (3,2), ...). In PyTorch, this operation is called `view`:"
|
| 402 |
+
]
|
| 403 |
+
},
|
| 404 |
+
{
|
| 405 |
+
"cell_type": "code",
|
| 406 |
+
"execution_count": null,
|
| 407 |
+
"metadata": {
|
| 408 |
+
"id": "K-4HwV-1yKD7"
|
| 409 |
+
},
|
| 410 |
+
"outputs": [],
|
| 411 |
+
"source": [
|
| 412 |
+
"x = x.view(2, 3)\n",
|
| 413 |
+
"print(\"X\", x)"
|
| 414 |
+
]
|
| 415 |
+
},
|
| 416 |
+
{
|
| 417 |
+
"cell_type": "code",
|
| 418 |
+
"execution_count": null,
|
| 419 |
+
"metadata": {
|
| 420 |
+
"id": "DtPoljp1yKD7"
|
| 421 |
+
},
|
| 422 |
+
"outputs": [],
|
| 423 |
+
"source": [
|
| 424 |
+
"x = x.permute(1, 0) # Swapping dimension 0 and 1\n",
|
| 425 |
+
"print(\"X\", x)"
|
| 426 |
+
]
|
| 427 |
+
},
|
| 428 |
+
{
|
| 429 |
+
"cell_type": "markdown",
|
| 430 |
+
"metadata": {
|
| 431 |
+
"id": "ALaEhyWdkDdJ"
|
| 432 |
+
},
|
| 433 |
+
"source": [
|
| 434 |
+
"#### Matrix multiplication"
|
| 435 |
+
]
|
| 436 |
+
},
|
| 437 |
+
{
|
| 438 |
+
"cell_type": "markdown",
|
| 439 |
+
"metadata": {
|
| 440 |
+
"id": "LNyWsbTzyKD7"
|
| 441 |
+
},
|
| 442 |
+
"source": [
|
| 443 |
+
"Other commonly used operations include matrix multiplications, which are essential for neural networks. Quite often, we have an input vector $\\mathbf{x}$, which is transformed using a learned weight matrix $\\mathbf{W}$. There are multiple ways and functions to perform matrix multiplication, some of which are listed below:\n",
|
| 444 |
+
"\n",
|
| 445 |
+
"* `torch.matmul`: Performs the matrix product over two tensors, where the specific behavior depends on the dimensions. If both inputs are matrices (2-dimensional tensors), it performs the standard matrix product. For higher dimensional inputs, the function supports broadcasting (for details see the [documentation](https://pytorch.org/docs/stable/generated/torch.matmul.html?highlight=matmul#torch.matmul)). Similar to Numpy, it can also be written as `a @ b`.\n",
|
| 446 |
+
"* `torch.mm`: Performs the matrix product over two matrices, but doesn't support broadcasting (see [documentation](https://pytorch.org/docs/stable/generated/torch.mm.html?highlight=torch%20mm#torch.mm)).\n",
|
| 447 |
+
"* `torch.bmm`: Performs the matrix product with a support batch dimension. If the first tensor $T$ is of shape ($b\\times n\\times m$), and the second tensor $R$ ($b\\times m\\times p$), the output $O$ is of shape ($b\\times n\\times p$), and has been calculated by performing $b$ matrix multiplications of the submatrices of $T$ and $R$: $O_i = T_i @ R_i$.\n",
|
| 448 |
+
"* `torch.einsum`: Performs matrix multiplications and more (i.e. sums of products) using the Einstein summation convention. Explanation of the Einstein sum can be found in assignment 1.\n",
|
| 449 |
+
"\n",
|
| 450 |
+
"Usually, we use `torch.matmul` or `torch.bmm`. We can try a matrix multiplication with `torch.matmul` below."
|
| 451 |
+
]
|
| 452 |
+
},
|
| 453 |
+
{
|
| 454 |
+
"cell_type": "code",
|
| 455 |
+
"execution_count": null,
|
| 456 |
+
"metadata": {
|
| 457 |
+
"id": "IPfjNkQ9yKD8"
|
| 458 |
+
},
|
| 459 |
+
"outputs": [],
|
| 460 |
+
"source": [
|
| 461 |
+
"x = torch.arange(6)\n",
|
| 462 |
+
"x = x.view(2, 3)\n",
|
| 463 |
+
"print(\"X\", x)"
|
| 464 |
+
]
|
| 465 |
+
},
|
| 466 |
+
{
|
| 467 |
+
"cell_type": "code",
|
| 468 |
+
"execution_count": null,
|
| 469 |
+
"metadata": {
|
| 470 |
+
"id": "89BudWvryKD8"
|
| 471 |
+
},
|
| 472 |
+
"outputs": [],
|
| 473 |
+
"source": [
|
| 474 |
+
"W = torch.arange(9).view(3, 3) # We can also stack multiple operations in a single line\n",
|
| 475 |
+
"print(\"W\", W)"
|
| 476 |
+
]
|
| 477 |
+
},
|
| 478 |
+
{
|
| 479 |
+
"cell_type": "code",
|
| 480 |
+
"execution_count": null,
|
| 481 |
+
"metadata": {
|
| 482 |
+
"id": "VrwwqU5dyKD8"
|
| 483 |
+
},
|
| 484 |
+
"outputs": [],
|
| 485 |
+
"source": [
|
| 486 |
+
"h = torch.matmul(x, W) # Verify the result by calculating it by hand too!\n",
|
| 487 |
+
"print(\"h\", h)"
|
| 488 |
+
]
|
| 489 |
+
},
|
| 490 |
+
{
|
| 491 |
+
"cell_type": "markdown",
|
| 492 |
+
"metadata": {
|
| 493 |
+
"id": "NlL4qEXAkcVJ"
|
| 494 |
+
},
|
| 495 |
+
"source": [
|
| 496 |
+
"#### Indexing"
|
| 497 |
+
]
|
| 498 |
+
},
|
| 499 |
+
{
|
| 500 |
+
"cell_type": "markdown",
|
| 501 |
+
"metadata": {
|
| 502 |
+
"id": "uQ_URivAyKD8"
|
| 503 |
+
},
|
| 504 |
+
"source": [
|
| 505 |
+
"We often have the situation where we need to select a part of a tensor. Indexing in PyTorch works just like in Numpy, so let's try it:"
|
| 506 |
+
]
|
| 507 |
+
},
|
| 508 |
+
{
|
| 509 |
+
"cell_type": "code",
|
| 510 |
+
"execution_count": null,
|
| 511 |
+
"metadata": {
|
| 512 |
+
"id": "7SSPTOJ6yKD8"
|
| 513 |
+
},
|
| 514 |
+
"outputs": [],
|
| 515 |
+
"source": [
|
| 516 |
+
"x = torch.arange(12).view(3, 4)\n",
|
| 517 |
+
"print(\"X\", x)"
|
| 518 |
+
]
|
| 519 |
+
},
|
| 520 |
+
{
|
| 521 |
+
"cell_type": "code",
|
| 522 |
+
"execution_count": null,
|
| 523 |
+
"metadata": {
|
| 524 |
+
"id": "Ne9br4uSyKD8"
|
| 525 |
+
},
|
| 526 |
+
"outputs": [],
|
| 527 |
+
"source": [
|
| 528 |
+
"print(x[:, 1]) # Second column"
|
| 529 |
+
]
|
| 530 |
+
},
|
| 531 |
+
{
|
| 532 |
+
"cell_type": "code",
|
| 533 |
+
"execution_count": null,
|
| 534 |
+
"metadata": {
|
| 535 |
+
"id": "snXYszDAyKD9"
|
| 536 |
+
},
|
| 537 |
+
"outputs": [],
|
| 538 |
+
"source": [
|
| 539 |
+
"print(x[0]) # First row"
|
| 540 |
+
]
|
| 541 |
+
},
|
| 542 |
+
{
|
| 543 |
+
"cell_type": "code",
|
| 544 |
+
"execution_count": null,
|
| 545 |
+
"metadata": {
|
| 546 |
+
"id": "WmDokJuVyKD9"
|
| 547 |
+
},
|
| 548 |
+
"outputs": [],
|
| 549 |
+
"source": [
|
| 550 |
+
"print(x[:2, -1]) # First two rows, last column"
|
| 551 |
+
]
|
| 552 |
+
},
|
| 553 |
+
{
|
| 554 |
+
"cell_type": "code",
|
| 555 |
+
"execution_count": null,
|
| 556 |
+
"metadata": {
|
| 557 |
+
"id": "bXIJhtnbyKD9"
|
| 558 |
+
},
|
| 559 |
+
"outputs": [],
|
| 560 |
+
"source": [
|
| 561 |
+
"print(x[1:3, :]) # Middle two rows"
|
| 562 |
+
]
|
| 563 |
+
},
|
| 564 |
+
{
|
| 565 |
+
"cell_type": "markdown",
|
| 566 |
+
"metadata": {
|
| 567 |
+
"id": "XUpotMGSyKD9"
|
| 568 |
+
},
|
| 569 |
+
"source": [
|
| 570 |
+
"### Gradients, Computational graph and Backpropagation"
|
| 571 |
+
]
|
| 572 |
+
},
|
| 573 |
+
{
|
| 574 |
+
"cell_type": "markdown",
|
| 575 |
+
"metadata": {
|
| 576 |
+
"id": "yYttGusdKXn6"
|
| 577 |
+
},
|
| 578 |
+
"source": [
|
| 579 |
+
"#### Gradients"
|
| 580 |
+
]
|
| 581 |
+
},
|
| 582 |
+
{
|
| 583 |
+
"cell_type": "markdown",
|
| 584 |
+
"metadata": {
|
| 585 |
+
"id": "bvCU_mMuKaoU"
|
| 586 |
+
},
|
| 587 |
+
"source": [
|
| 588 |
+
"We use a deep learning framework for implementing neural networks which are effectively a combination of several differentiable functions parameterized with weights that we aim to learn or adjust during the training procedure. While training, the gradients of those parameters or weights are computed to update them via the delta rule. One of the main reasons for using deep learning framework, such as PyTorch, TensorFlow is that we can automatically obatin the **derivatives** or **gradients** of those weights if we have a valid differentiable functions."
|
| 589 |
+
]
|
| 590 |
+
},
|
| 591 |
+
{
|
| 592 |
+
"cell_type": "markdown",
|
| 593 |
+
"metadata": {
|
| 594 |
+
"id": "8EKJIwkfKef3"
|
| 595 |
+
},
|
| 596 |
+
"source": [
|
| 597 |
+
"#### Computational graph"
|
| 598 |
+
]
|
| 599 |
+
},
|
| 600 |
+
{
|
| 601 |
+
"cell_type": "markdown",
|
| 602 |
+
"metadata": {
|
| 603 |
+
"id": "uo3quHoWKh3a"
|
| 604 |
+
},
|
| 605 |
+
"source": [
|
| 606 |
+
"We define our function by manipulating the input, usually by a series of matrix multiplications with weight matrices ($\\mathbf{W}$) and additions with bias vectors ($b$). As we manipulate our input, we are automatically creating a **computational graph** showing how to arrive at the output from the input. In PyTorch, we just define the manipulations and it keeps track of that graph by design.\n"
|
| 607 |
+
]
|
| 608 |
+
},
|
| 609 |
+
{
|
| 610 |
+
"cell_type": "markdown",
|
| 611 |
+
"metadata": {
|
| 612 |
+
"id": "iEpXNRNqEeTC"
|
| 613 |
+
},
|
| 614 |
+
"source": [
|
| 615 |
+
"#### Backpropagation"
|
| 616 |
+
]
|
| 617 |
+
},
|
| 618 |
+
{
|
| 619 |
+
"cell_type": "markdown",
|
| 620 |
+
"metadata": {
|
| 621 |
+
"id": "wwdEDnmnEg29"
|
| 622 |
+
},
|
| 623 |
+
"source": [
|
| 624 |
+
"Given an input $\\mathbf{x}$, we obtain the output $y$ by manipulating that input via series of matrix multiplications with weight matrices ($\\mathbf{W}$) and additions with bias vectors ($b$). As we manipulate our input, we automatically create a computational graph showing how to arrive at the output from the input. We then define an error measure or **loss function** that tells us how wrong our network is. In other words, how good or bad it is in predicting output $y$ from input $\\mathbf{x}$. Based on this error measure, we can use the gradients to update or **backpropagate** the weights $\\mathbf{W}$ that were responsible for the output, so that the next time we present input $\\mathbf{x}$ to our network, the output will be closer to what we want."
|
| 625 |
+
]
|
| 626 |
+
},
|
| 627 |
+
{
|
| 628 |
+
"cell_type": "markdown",
|
| 629 |
+
"metadata": {
|
| 630 |
+
"id": "FLepWzW1L6tC"
|
| 631 |
+
},
|
| 632 |
+
"source": [
|
| 633 |
+
"#### `.requires_grad()`"
|
| 634 |
+
]
|
| 635 |
+
},
|
| 636 |
+
{
|
| 637 |
+
"cell_type": "markdown",
|
| 638 |
+
"metadata": {
|
| 639 |
+
"id": "tk32gBuMEwTx"
|
| 640 |
+
},
|
| 641 |
+
"source": [
|
| 642 |
+
"In PyTorch, whether a particular tensor containing a set of parameters requires gradient or not is determined by the associated flag `requires_grad`. By default, when we create a tensor, it does not require gradients."
|
| 643 |
+
]
|
| 644 |
+
},
|
| 645 |
+
{
|
| 646 |
+
"cell_type": "code",
|
| 647 |
+
"execution_count": null,
|
| 648 |
+
"metadata": {
|
| 649 |
+
"id": "ejXjJ9ECyKD9"
|
| 650 |
+
},
|
| 651 |
+
"outputs": [],
|
| 652 |
+
"source": [
|
| 653 |
+
"x = torch.ones((3,))\n",
|
| 654 |
+
"print(x.requires_grad)"
|
| 655 |
+
]
|
| 656 |
+
},
|
| 657 |
+
{
|
| 658 |
+
"cell_type": "markdown",
|
| 659 |
+
"metadata": {
|
| 660 |
+
"id": "bc3i5kVRyKD9"
|
| 661 |
+
},
|
| 662 |
+
"source": [
|
| 663 |
+
"We can change this for an existing tensor using the function `requires_grad_()` (underscore indicating that this is a in-place operation). Alternatively, when creating a tensor, you can pass the argument `requires_grad=True` to most initializers we have seen above."
|
| 664 |
+
]
|
| 665 |
+
},
|
| 666 |
+
{
|
| 667 |
+
"cell_type": "code",
|
| 668 |
+
"execution_count": null,
|
| 669 |
+
"metadata": {
|
| 670 |
+
"id": "7weDjnh5yKD9"
|
| 671 |
+
},
|
| 672 |
+
"outputs": [],
|
| 673 |
+
"source": [
|
| 674 |
+
"x.requires_grad_(True)\n",
|
| 675 |
+
"print(x.requires_grad)"
|
| 676 |
+
]
|
| 677 |
+
},
|
| 678 |
+
{
|
| 679 |
+
"cell_type": "markdown",
|
| 680 |
+
"metadata": {
|
| 681 |
+
"id": "RBQlwCyiDZ05"
|
| 682 |
+
},
|
| 683 |
+
"source": [
|
| 684 |
+
"#### Example"
|
| 685 |
+
]
|
| 686 |
+
},
|
| 687 |
+
{
|
| 688 |
+
"cell_type": "markdown",
|
| 689 |
+
"metadata": {
|
| 690 |
+
"id": "axzFVIkgyKD-"
|
| 691 |
+
},
|
| 692 |
+
"source": [
|
| 693 |
+
"In order to get familiar with the concept of a computation graph, we will create one for the following function:\n",
|
| 694 |
+
"\n",
|
| 695 |
+
"$$y = \\frac{1}{|x|}\\sum_i \\left[(x_i + 5)^3 + 7\\right]$$\n",
|
| 696 |
+
"\n",
|
| 697 |
+
"You could imagine that $x$ are our parameters, and we want to optimize (either maximize or minimize) the output $y$. For this, we want to obtain the gradients $\\partial y / \\partial \\mathbf{x}$. For our example, we'll use $\\mathbf{x}=[0,1,2,3,4]$ as our input."
|
| 698 |
+
]
|
| 699 |
+
},
|
| 700 |
+
{
|
| 701 |
+
"cell_type": "code",
|
| 702 |
+
"execution_count": null,
|
| 703 |
+
"metadata": {
|
| 704 |
+
"id": "T7Ts_CHgyKD-"
|
| 705 |
+
},
|
| 706 |
+
"outputs": [],
|
| 707 |
+
"source": [
|
| 708 |
+
"x = torch.arange(5, dtype=torch.float32, requires_grad=True) # Only float tensors can have gradients\n",
|
| 709 |
+
"print(\"X\", x)"
|
| 710 |
+
]
|
| 711 |
+
},
|
| 712 |
+
{
|
| 713 |
+
"cell_type": "markdown",
|
| 714 |
+
"metadata": {
|
| 715 |
+
"id": "lOoDHPT-yKD-"
|
| 716 |
+
},
|
| 717 |
+
"source": [
|
| 718 |
+
"Now let's build the computation graph step by step. You can combine multiple operations in a single line, but we will separate them here to get a better understanding of how each operation is added to the computation graph."
|
| 719 |
+
]
|
| 720 |
+
},
|
| 721 |
+
{
|
| 722 |
+
"cell_type": "code",
|
| 723 |
+
"execution_count": null,
|
| 724 |
+
"metadata": {
|
| 725 |
+
"id": "85m78jwHyKD-"
|
| 726 |
+
},
|
| 727 |
+
"outputs": [],
|
| 728 |
+
"source": [
|
| 729 |
+
"a = x + 5\n",
|
| 730 |
+
"b = a ** 3\n",
|
| 731 |
+
"c = b + 7\n",
|
| 732 |
+
"y = c.mean()\n",
|
| 733 |
+
"print(\"Y\", y)"
|
| 734 |
+
]
|
| 735 |
+
},
|
| 736 |
+
{
|
| 737 |
+
"cell_type": "markdown",
|
| 738 |
+
"metadata": {
|
| 739 |
+
"id": "U9Xjh2ZoyKD-"
|
| 740 |
+
},
|
| 741 |
+
"source": [
|
| 742 |
+
"Using the statements above, we have created a computation graph that looks similar to the figure below:\n",
|
| 743 |
+
"\n",
|
| 744 |
+
"<center style=\"width: 100%\"><img src=\"https://github.com/AnjanDutta/SharedFigures/blob/main/comp_graph.png?raw=true\" width=\"500px\"></center>\n",
|
| 745 |
+
"\n",
|
| 746 |
+
"We calculate $a$ based on the inputs $x$ and the constant $5$, $b$ is $a$ cubed, and so on. The visualization is an abstraction of the dependencies between inputs and outputs of the operations we have applied.\n",
|
| 747 |
+
"Each node of the computation graph has automatically defined a function for calculating the gradients with respect to its inputs, `grad_fn`. You can see this when we printed the output tensor $y$. This is why the computation graph is usually visualized in the reverse direction (arrows point from the result to the inputs). We can perform backpropagation on the computation graph by calling the function `backward()` on the last output, which effectively calculates the gradients for each tensor that has the property `requires_grad=True`:"
|
| 748 |
+
]
|
| 749 |
+
},
|
| 750 |
+
{
|
| 751 |
+
"cell_type": "code",
|
| 752 |
+
"execution_count": null,
|
| 753 |
+
"metadata": {
|
| 754 |
+
"id": "7_A1RRZUyKD-"
|
| 755 |
+
},
|
| 756 |
+
"outputs": [],
|
| 757 |
+
"source": [
|
| 758 |
+
"y.backward()"
|
| 759 |
+
]
|
| 760 |
+
},
|
| 761 |
+
{
|
| 762 |
+
"cell_type": "markdown",
|
| 763 |
+
"metadata": {
|
| 764 |
+
"id": "Lp9j77yKyKD_"
|
| 765 |
+
},
|
| 766 |
+
"source": [
|
| 767 |
+
"`x.grad` will now contain the gradient $\\partial y/ \\partial \\mathcal{x}$, and this gradient indicates how a change in $\\mathbf{x}$ will affect output $y$ given the current input $\\mathbf{x}=[0,1,2,3,4]$:"
|
| 768 |
+
]
|
| 769 |
+
},
|
| 770 |
+
{
|
| 771 |
+
"cell_type": "code",
|
| 772 |
+
"execution_count": null,
|
| 773 |
+
"metadata": {
|
| 774 |
+
"id": "jSnCN6REyKD_"
|
| 775 |
+
},
|
| 776 |
+
"outputs": [],
|
| 777 |
+
"source": [
|
| 778 |
+
"print(x.grad)"
|
| 779 |
+
]
|
| 780 |
+
},
|
| 781 |
+
{
|
| 782 |
+
"cell_type": "markdown",
|
| 783 |
+
"metadata": {
|
| 784 |
+
"id": "j5XNxMfryKD_"
|
| 785 |
+
},
|
| 786 |
+
"source": [
|
| 787 |
+
"We can also verify these gradients by hand. We will calculate the gradients using the chain rule, in the same way as PyTorch did it:\n",
|
| 788 |
+
"\n",
|
| 789 |
+
"$$\\frac{\\partial y}{\\partial x_i} = \\frac{\\partial y}{\\partial c_i}\\frac{\\partial c_i}{\\partial b_i}\\frac{\\partial b_i}{\\partial a_i}\\frac{\\partial a_i}{\\partial x_i}$$\n",
|
| 790 |
+
"\n",
|
| 791 |
+
"Note that we have simplified this equation to index notation, and by using the fact that all operation besides the mean do not combine the elements in the tensor. The partial derivatives are:\n",
|
| 792 |
+
"\n",
|
| 793 |
+
"$$\n",
|
| 794 |
+
"\\frac{\\partial a_i}{\\partial x_i} = 1,\\hspace{1cm}\n",
|
| 795 |
+
"\\frac{\\partial b_i}{\\partial a_i} = 3\\cdot a_i^2\\hspace{1cm}\n",
|
| 796 |
+
"\\frac{\\partial c_i}{\\partial b_i} = 1\\hspace{1cm}\n",
|
| 797 |
+
"\\frac{\\partial y}{\\partial c_i} = \\frac{1}{5}\n",
|
| 798 |
+
"$$\n",
|
| 799 |
+
"\n",
|
| 800 |
+
"Hence, with the input being $\\mathbf{x}=[0,1,2,3,4]$, our gradients are $\\frac{\\partial y}{\\partial \\mathbf{x}}=[15, \\frac{108}{5}, \\frac{147}{5}, \\frac{192}{5},\\frac{243}{5}]$. The previous code cell should have printed the same result."
|
| 801 |
+
]
|
| 802 |
+
},
|
| 803 |
+
{
|
| 804 |
+
"cell_type": "markdown",
|
| 805 |
+
"metadata": {
|
| 806 |
+
"id": "NbgfwdQlwudF"
|
| 807 |
+
},
|
| 808 |
+
"source": [
|
| 809 |
+
"### GPU support"
|
| 810 |
+
]
|
| 811 |
+
},
|
| 812 |
+
{
|
| 813 |
+
"cell_type": "markdown",
|
| 814 |
+
"metadata": {
|
| 815 |
+
"id": "LCd3NuIYyKD_"
|
| 816 |
+
},
|
| 817 |
+
"source": [
|
| 818 |
+
"A crucial feature of PyTorch is the support of Graphics Processing Unit (GPU). A GPU can perform many thousands of small operations in parallel, making it very well suitable for performing large matrix operations in neural networks. When comparing GPUs to CPUs, we can list the following main differences (credit: [Kevin Krewell, 2009](https://blogs.nvidia.com/blog/2009/12/16/whats-the-difference-between-a-cpu-and-a-gpu/))\n",
|
| 819 |
+
"\n",
|
| 820 |
+
"CPUs and GPUs have both different advantages and disadvantages, which is why many computers contain both components and use them for different tasks. In case you are not familiar with GPUs, you can read up more details in this [NVIDIA blog post](https://blogs.nvidia.com/blog/2009/12/16/whats-the-difference-between-a-cpu-and-a-gpu/) or [here](https://www.intel.com/content/www/us/en/products/docs/processors/what-is-a-gpu.html).\n",
|
| 821 |
+
"\n",
|
| 822 |
+
"GPUs can accelerate the training of your network up to a factor of $100$ which is essential for large neural networks. PyTorch implements a lot of functionality for supporting GPUs (mostly those of NVIDIA due to the libraries [CUDA](https://developer.nvidia.com/cuda-zone) and [cuDNN](https://developer.nvidia.com/cudnn)). First, let's check whether you have a GPU available:"
|
| 823 |
+
]
|
| 824 |
+
},
|
| 825 |
+
{
|
| 826 |
+
"cell_type": "code",
|
| 827 |
+
"execution_count": null,
|
| 828 |
+
"metadata": {
|
| 829 |
+
"id": "L7IUQhSZyKD_"
|
| 830 |
+
},
|
| 831 |
+
"outputs": [],
|
| 832 |
+
"source": [
|
| 833 |
+
"gpu_avail = torch.cuda.is_available()\n",
|
| 834 |
+
"print(f\"Is the GPU available? {gpu_avail}\")"
|
| 835 |
+
]
|
| 836 |
+
},
|
| 837 |
+
{
|
| 838 |
+
"cell_type": "markdown",
|
| 839 |
+
"metadata": {
|
| 840 |
+
"id": "YT-miN_jyKD_"
|
| 841 |
+
},
|
| 842 |
+
"source": [
|
| 843 |
+
"If you have a GPU on your computer but the command above returns False, make sure you have the correct CUDA-version installed. On Colab, please change it if necessary (CUDA 12.4 is currently common on Colab). On Google Colab, make sure that you have selected a GPU in your runtime setup (in the menu, check under `Runtime -> Change runtime type`).\n",
|
| 844 |
+
"\n",
|
| 845 |
+
"By default, all tensors you create are stored on the CPU. We can push a tensor to the GPU by using the function `.to(...)`, or `.cuda()`. However, it is often a good practice to define a `device` object in your code which points to the GPU if you have one, and otherwise to the CPU. Then, you can write your code with respect to this device object, and it allows you to run the same code on both a CPU-only system, and one with a GPU. Let's try it below. We can specify the device as follows:"
|
| 846 |
+
]
|
| 847 |
+
},
|
| 848 |
+
{
|
| 849 |
+
"cell_type": "code",
|
| 850 |
+
"execution_count": null,
|
| 851 |
+
"metadata": {
|
| 852 |
+
"id": "iYI9tssMyKD_"
|
| 853 |
+
},
|
| 854 |
+
"outputs": [],
|
| 855 |
+
"source": [
|
| 856 |
+
"device = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")\n",
|
| 857 |
+
"print(\"Device\", device)"
|
| 858 |
+
]
|
| 859 |
+
},
|
| 860 |
+
{
|
| 861 |
+
"cell_type": "markdown",
|
| 862 |
+
"metadata": {
|
| 863 |
+
"id": "l7gtUjj3yKEA"
|
| 864 |
+
},
|
| 865 |
+
"source": [
|
| 866 |
+
"Now let's create a tensor and push it to the device:"
|
| 867 |
+
]
|
| 868 |
+
},
|
| 869 |
+
{
|
| 870 |
+
"cell_type": "code",
|
| 871 |
+
"execution_count": null,
|
| 872 |
+
"metadata": {
|
| 873 |
+
"id": "AEKx_99SyKEA"
|
| 874 |
+
},
|
| 875 |
+
"outputs": [],
|
| 876 |
+
"source": [
|
| 877 |
+
"x = torch.zeros(2, 3)\n",
|
| 878 |
+
"x = x.to(device)\n",
|
| 879 |
+
"print(\"X\", x)"
|
| 880 |
+
]
|
| 881 |
+
},
|
| 882 |
+
{
|
| 883 |
+
"cell_type": "markdown",
|
| 884 |
+
"metadata": {
|
| 885 |
+
"id": "8UrSdpFzyKEA"
|
| 886 |
+
},
|
| 887 |
+
"source": [
|
| 888 |
+
"In case you have a GPU, you should now see the attribute `device='cuda:0'` being printed next to your tensor. The zero next to cuda indicates that this is the zero-th GPU device on your computer. PyTorch also supports multi-GPU systems, but this you will only need once you have very big networks to train (if interested, see the [PyTorch documentation](https://pytorch.org/docs/stable/distributed.html#distributed-basics)). We can also compare the runtime of a large matrix multiplication on the CPU with a operation on the GPU:"
|
| 889 |
+
]
|
| 890 |
+
},
|
| 891 |
+
{
|
| 892 |
+
"cell_type": "code",
|
| 893 |
+
"execution_count": null,
|
| 894 |
+
"metadata": {
|
| 895 |
+
"id": "PMEsWrjMyKEA"
|
| 896 |
+
},
|
| 897 |
+
"outputs": [],
|
| 898 |
+
"source": [
|
| 899 |
+
"import time\n",
|
| 900 |
+
"x = torch.randn(5000, 5000)\n",
|
| 901 |
+
"\n",
|
| 902 |
+
"## CPU version\n",
|
| 903 |
+
"start_time = time.time()\n",
|
| 904 |
+
"_ = torch.matmul(x, x)\n",
|
| 905 |
+
"end_time = time.time()\n",
|
| 906 |
+
"print(f\"CPU time: {(end_time - start_time):6.5f}s\")\n",
|
| 907 |
+
"\n",
|
| 908 |
+
"## GPU version\n",
|
| 909 |
+
"x = x.to(device)\n",
|
| 910 |
+
"_ = torch.matmul(x, x) # First operation to 'burn in' GPU\n",
|
| 911 |
+
"# CUDA is asynchronous, so we need to use different timing functions\n",
|
| 912 |
+
"start = torch.cuda.Event(enable_timing=True)\n",
|
| 913 |
+
"end = torch.cuda.Event(enable_timing=True)\n",
|
| 914 |
+
"start.record()\n",
|
| 915 |
+
"_ = torch.matmul(x, x)\n",
|
| 916 |
+
"end.record()\n",
|
| 917 |
+
"torch.cuda.synchronize() # Waits for everything to finish running on the GPU\n",
|
| 918 |
+
"print(f\"GPU time: {0.001 * start.elapsed_time(end):6.5f}s\") # Milliseconds to seconds"
|
| 919 |
+
]
|
| 920 |
+
},
|
| 921 |
+
{
|
| 922 |
+
"cell_type": "markdown",
|
| 923 |
+
"metadata": {
|
| 924 |
+
"id": "MFoSS6lCyKEA"
|
| 925 |
+
},
|
| 926 |
+
"source": [
|
| 927 |
+
"Depending on the size of the operation and the CPU/GPU in your system, the speedup of this operation can be >50x. As `matmul` operations are very common in neural networks, we can already see the great benefit of training a NN on a GPU. The time estimate can be relatively noisy here because we haven't run it for multiple times. Feel free to extend this, but it also takes longer to run.\n",
|
| 928 |
+
"\n",
|
| 929 |
+
"When generating random numbers, the seed between CPU and GPU is not synchronized. Hence, we need to set the seed on the GPU separately to ensure a reproducible code. Note that due to different GPU architectures, running the same code on different GPUs does not guarantee the same random numbers. Still, we don't want that our code gives us a different output every time we run it on the exact same hardware. Hence, we also set the seed on the GPU:"
|
| 930 |
+
]
|
| 931 |
+
},
|
| 932 |
+
{
|
| 933 |
+
"cell_type": "code",
|
| 934 |
+
"execution_count": null,
|
| 935 |
+
"metadata": {
|
| 936 |
+
"id": "Rw89vknyyKEA"
|
| 937 |
+
},
|
| 938 |
+
"outputs": [],
|
| 939 |
+
"source": [
|
| 940 |
+
"# GPU operations have a separate seed we also want to set\n",
|
| 941 |
+
"if torch.cuda.is_available():\n",
|
| 942 |
+
" torch.cuda.manual_seed(42)\n",
|
| 943 |
+
" torch.cuda.manual_seed_all(42)\n",
|
| 944 |
+
"\n",
|
| 945 |
+
"# Additionally, some operations on a GPU are implemented stochastic for efficiency\n",
|
| 946 |
+
"# We want to ensure that all operations are deterministic on GPU (if used) for reproducibility\n",
|
| 947 |
+
"torch.backends.cudnn.deterministic = True\n",
|
| 948 |
+
"torch.backends.cudnn.benchmark = False"
|
| 949 |
+
]
|
| 950 |
+
},
|
| 951 |
+
{
|
| 952 |
+
"cell_type": "markdown",
|
| 953 |
+
"metadata": {
|
| 954 |
+
"id": "HCjp1yYIyKEA"
|
| 955 |
+
},
|
| 956 |
+
"source": [
|
| 957 |
+
"## Example: Gaussian or continuous XNOR\n",
|
| 958 |
+
"\n",
|
| 959 |
+
"If we want to build a neural network in PyTorch, we could specify all our parameters (weight matrices, bias vectors) using `Tensors` (with `requires_grad=True`), ask PyTorch to calculate the gradients and then adjust the parameters. But things can quickly get cumbersome if we have a lot of parameters. In PyTorch, there is a package called `torch.nn` that makes building neural networks more convenient.\n",
|
| 960 |
+
"\n",
|
| 961 |
+
"We will introduce the libraries and all additional parts you might need to train a neural network in PyTorch, using a simple example classifier on a simple yet well known example: XNOR. Given two binary inputs $x_1$ and $x_2$, the label to predict is $1$ if $x_1$ equal to $x_2$ and $0$ if $x_1$ not equal to $x_2$. The example became famous by the fact that a single neuron, i.e. a linear classifier, cannot learn this simple function. Hence, we will learn how to build a small neural network that can learn this function. To make it a little bit more interesting, we move the XNOR into continuous space and introduce some Gaussian noise on the binary inputs. Our desired separation of an XNOR dataset could look as follows:\n",
|
| 962 |
+
"\n",
|
| 963 |
+
"<center style=\"width: 100%\"><img src=\"https://github.com/AnjanDutta/SharedFigures/blob/main/xnor_plot.png?raw=true\" width=\"500px\"></center>"
|
| 964 |
+
]
|
| 965 |
+
},
|
| 966 |
+
{
|
| 967 |
+
"cell_type": "markdown",
|
| 968 |
+
"metadata": {
|
| 969 |
+
"id": "gXP2REOmwnQM"
|
| 970 |
+
},
|
| 971 |
+
"source": [
|
| 972 |
+
"### The model"
|
| 973 |
+
]
|
| 974 |
+
},
|
| 975 |
+
{
|
| 976 |
+
"cell_type": "markdown",
|
| 977 |
+
"metadata": {
|
| 978 |
+
"id": "lvokYLAvyKEA"
|
| 979 |
+
},
|
| 980 |
+
"source": [
|
| 981 |
+
"The package `torch.nn` defines a series of useful classes like linear networks layers, activation functions, loss functions etc. A full list can be found [here](https://pytorch.org/docs/stable/nn.html). In case you need a certain network layer, check the documentation of the package first before writing the layer yourself as the package likely contains the code for it already. We import it below:"
|
| 982 |
+
]
|
| 983 |
+
},
|
| 984 |
+
{
|
| 985 |
+
"cell_type": "code",
|
| 986 |
+
"execution_count": null,
|
| 987 |
+
"metadata": {
|
| 988 |
+
"id": "G7rc32GwyKEA"
|
| 989 |
+
},
|
| 990 |
+
"outputs": [],
|
| 991 |
+
"source": [
|
| 992 |
+
"import torch.nn as nn"
|
| 993 |
+
]
|
| 994 |
+
},
|
| 995 |
+
{
|
| 996 |
+
"cell_type": "markdown",
|
| 997 |
+
"metadata": {
|
| 998 |
+
"id": "eRv10SNSyKEB"
|
| 999 |
+
},
|
| 1000 |
+
"source": [
|
| 1001 |
+
"Additionally to `torch.nn`, there is also `torch.nn.functional`. It contains functions that are used in network layers. This is in contrast to `torch.nn` which defines them as `nn.Modules` (more on it below), and `torch.nn` actually uses a lot of functionalities from `torch.nn.functional`. Hence, the functional package is useful in many situations, and so we import it as well here."
|
| 1002 |
+
]
|
| 1003 |
+
},
|
| 1004 |
+
{
|
| 1005 |
+
"cell_type": "code",
|
| 1006 |
+
"execution_count": null,
|
| 1007 |
+
"metadata": {
|
| 1008 |
+
"id": "BKV0ia61yKEB"
|
| 1009 |
+
},
|
| 1010 |
+
"outputs": [],
|
| 1011 |
+
"source": [
|
| 1012 |
+
"import torch.nn.functional as F"
|
| 1013 |
+
]
|
| 1014 |
+
},
|
| 1015 |
+
{
|
| 1016 |
+
"cell_type": "markdown",
|
| 1017 |
+
"metadata": {
|
| 1018 |
+
"id": "67AkppacyKEB"
|
| 1019 |
+
},
|
| 1020 |
+
"source": [
|
| 1021 |
+
"#### nn.Module\n",
|
| 1022 |
+
"\n",
|
| 1023 |
+
"In PyTorch, a neural network is built up out of modules. Modules can contain other modules, and a neural network is considered to be a module itself as well. The basic template of a module is as follows:"
|
| 1024 |
+
]
|
| 1025 |
+
},
|
| 1026 |
+
{
|
| 1027 |
+
"cell_type": "code",
|
| 1028 |
+
"execution_count": null,
|
| 1029 |
+
"metadata": {
|
| 1030 |
+
"id": "IMKZ1UiIyKEB"
|
| 1031 |
+
},
|
| 1032 |
+
"outputs": [],
|
| 1033 |
+
"source": [
|
| 1034 |
+
"class MyModule(nn.Module):\n",
|
| 1035 |
+
"\n",
|
| 1036 |
+
" def __init__(self):\n",
|
| 1037 |
+
" super().__init__()\n",
|
| 1038 |
+
" # Some init for my module\n",
|
| 1039 |
+
"\n",
|
| 1040 |
+
" def forward(self, x):\n",
|
| 1041 |
+
" # Function for performing the calculation of the module.\n",
|
| 1042 |
+
" pass"
|
| 1043 |
+
]
|
| 1044 |
+
},
|
| 1045 |
+
{
|
| 1046 |
+
"cell_type": "markdown",
|
| 1047 |
+
"metadata": {
|
| 1048 |
+
"id": "U_P3vY2dyKEB"
|
| 1049 |
+
},
|
| 1050 |
+
"source": [
|
| 1051 |
+
"The `forward()` function is where the computation of the module is taken place, and is executed when you call the module (`nn = MyModule(); nn(x)`). In the `__init__()` function, we usually create the parameters of the module, using `nn.Parameter`, or defining other modules that are used in the forward function. The backward calculation is done automatically, but could be overwritten as well if wanted.\n",
|
| 1052 |
+
"\n",
|
| 1053 |
+
"#### Simple classifier\n",
|
| 1054 |
+
"We can now make use of the pre-defined modules in the `torch.nn` package, and define our own small neural network. We will use a minimal network with a input layer, one hidden layer with tanh as activation function, and a output layer. In other words, our networks should look something like this:\n",
|
| 1055 |
+
"\n",
|
| 1056 |
+
"<center width=\"100%\"><img src=\"https://raw.githubusercontent.com/AnjanDutta/SharedFigures/5beba0ea58e54079bcff6663eb4d8e7a7969d142/small_neural_network.svg\" width=\"400px\"></center>\n",
|
| 1057 |
+
"\n",
|
| 1058 |
+
"The input neurons are shown in blue, which represent the coordinates $x_1$ and $x_2$ of a data point. The hidden neurons including a tanh activation are shown in white, and the output neuron in red.\n",
|
| 1059 |
+
"In PyTorch, we can define this as follows:"
|
| 1060 |
+
]
|
| 1061 |
+
},
|
| 1062 |
+
{
|
| 1063 |
+
"cell_type": "code",
|
| 1064 |
+
"execution_count": null,
|
| 1065 |
+
"metadata": {
|
| 1066 |
+
"id": "UX_CcgvTyKEB"
|
| 1067 |
+
},
|
| 1068 |
+
"outputs": [],
|
| 1069 |
+
"source": [
|
| 1070 |
+
"class SimpleClassifier(nn.Module):\n",
|
| 1071 |
+
"\n",
|
| 1072 |
+
" def __init__(self, num_inputs, num_hidden, num_outputs):\n",
|
| 1073 |
+
" super().__init__()\n",
|
| 1074 |
+
" # Initialize the modules we need to build the network\n",
|
| 1075 |
+
" self.linear1 = nn.Linear(num_inputs, num_hidden)\n",
|
| 1076 |
+
" self.act_fn = nn.Tanh()\n",
|
| 1077 |
+
" self.linear2 = nn.Linear(num_hidden, num_outputs)\n",
|
| 1078 |
+
"\n",
|
| 1079 |
+
" def forward(self, x):\n",
|
| 1080 |
+
" # Perform the calculation of the model to determine the prediction\n",
|
| 1081 |
+
" x = self.linear1(x)\n",
|
| 1082 |
+
" x = self.act_fn(x)\n",
|
| 1083 |
+
" x = self.linear2(x)\n",
|
| 1084 |
+
" return x"
|
| 1085 |
+
]
|
| 1086 |
+
},
|
| 1087 |
+
{
|
| 1088 |
+
"cell_type": "markdown",
|
| 1089 |
+
"metadata": {
|
| 1090 |
+
"id": "2KXcb56byKEB"
|
| 1091 |
+
},
|
| 1092 |
+
"source": [
|
| 1093 |
+
"For the examples in this notebook, we will use a tiny neural network with two input neurons and four hidden neurons. As we perform binary classification, we will use a single output neuron. Note that we do not apply a sigmoid on the output yet. This is because other functions, especially the loss, are more efficient and precise to calculate on the original outputs instead of the sigmoid output. We will discuss the detailed reason later."
|
| 1094 |
+
]
|
| 1095 |
+
},
|
| 1096 |
+
{
|
| 1097 |
+
"cell_type": "code",
|
| 1098 |
+
"execution_count": null,
|
| 1099 |
+
"metadata": {
|
| 1100 |
+
"id": "2nqMOh43yKEB"
|
| 1101 |
+
},
|
| 1102 |
+
"outputs": [],
|
| 1103 |
+
"source": [
|
| 1104 |
+
"model = SimpleClassifier(num_inputs=2, num_hidden=4, num_outputs=1)\n",
|
| 1105 |
+
"# Printing a module shows all its submodules\n",
|
| 1106 |
+
"print(model)"
|
| 1107 |
+
]
|
| 1108 |
+
},
|
| 1109 |
+
{
|
| 1110 |
+
"cell_type": "markdown",
|
| 1111 |
+
"metadata": {
|
| 1112 |
+
"id": "AuzsBqEqyKEB"
|
| 1113 |
+
},
|
| 1114 |
+
"source": [
|
| 1115 |
+
"Printing the model lists all submodules it contains. The parameters of a module can be obtained by using its `parameters()` functions, or `named_parameters()` to get a name to each parameter object. For our small neural network, we have the following parameters:"
|
| 1116 |
+
]
|
| 1117 |
+
},
|
| 1118 |
+
{
|
| 1119 |
+
"cell_type": "code",
|
| 1120 |
+
"execution_count": null,
|
| 1121 |
+
"metadata": {
|
| 1122 |
+
"id": "ZidXJhX4yKEC"
|
| 1123 |
+
},
|
| 1124 |
+
"outputs": [],
|
| 1125 |
+
"source": [
|
| 1126 |
+
"for name, param in model.named_parameters():\n",
|
| 1127 |
+
" print(f\"Parameter {name}, shape {param.shape}\")"
|
| 1128 |
+
]
|
| 1129 |
+
},
|
| 1130 |
+
{
|
| 1131 |
+
"cell_type": "markdown",
|
| 1132 |
+
"metadata": {
|
| 1133 |
+
"id": "W6LKZ0osyKEC"
|
| 1134 |
+
},
|
| 1135 |
+
"source": [
|
| 1136 |
+
"Each linear layer has a weight matrix of the shape `[output, input]`, and a bias of the shape `[output]`. The tanh activation function does not have any parameters. Note that parameters are only registered for `nn.Module` objects that are direct object attributes, i.e. `self.a = ...`. If you define a list of modules, the parameters of those are not registered for the outer module and can cause some issues when you try to optimize your module. There are alternatives, like `nn.ModuleList`, `nn.ModuleDict` and `nn.Sequential`, that allow you to have different data structures of modules. We will use them in a few later tutorials and explain them there."
|
| 1137 |
+
]
|
| 1138 |
+
},
|
| 1139 |
+
{
|
| 1140 |
+
"cell_type": "markdown",
|
| 1141 |
+
"metadata": {
|
| 1142 |
+
"id": "FUjzPv20yKEC"
|
| 1143 |
+
},
|
| 1144 |
+
"source": [
|
| 1145 |
+
"### The data\n",
|
| 1146 |
+
"\n",
|
| 1147 |
+
"PyTorch also provides a few functionalities to load the training and test data efficiently, summarized in the package `torch.utils.data`."
|
| 1148 |
+
]
|
| 1149 |
+
},
|
| 1150 |
+
{
|
| 1151 |
+
"cell_type": "code",
|
| 1152 |
+
"execution_count": null,
|
| 1153 |
+
"metadata": {
|
| 1154 |
+
"id": "Z-csKPaGyKEC"
|
| 1155 |
+
},
|
| 1156 |
+
"outputs": [],
|
| 1157 |
+
"source": [
|
| 1158 |
+
"import torch.utils.data as data"
|
| 1159 |
+
]
|
| 1160 |
+
},
|
| 1161 |
+
{
|
| 1162 |
+
"cell_type": "markdown",
|
| 1163 |
+
"metadata": {
|
| 1164 |
+
"id": "9TR6t5aVyKEC"
|
| 1165 |
+
},
|
| 1166 |
+
"source": [
|
| 1167 |
+
"The data package defines two classes which are the standard interface for handling data in PyTorch: `data.Dataset`, and `data.DataLoader`. The dataset class provides an uniform interface to access the training/test data, while the data loader makes sure to efficiently load and stack the data points from the dataset into batches during training."
|
| 1168 |
+
]
|
| 1169 |
+
},
|
| 1170 |
+
{
|
| 1171 |
+
"cell_type": "markdown",
|
| 1172 |
+
"metadata": {
|
| 1173 |
+
"id": "QGrZ9jr4wi5F"
|
| 1174 |
+
},
|
| 1175 |
+
"source": [
|
| 1176 |
+
"#### The dataset class"
|
| 1177 |
+
]
|
| 1178 |
+
},
|
| 1179 |
+
{
|
| 1180 |
+
"cell_type": "markdown",
|
| 1181 |
+
"metadata": {
|
| 1182 |
+
"id": "2yOzJkPuyKEC"
|
| 1183 |
+
},
|
| 1184 |
+
"source": [
|
| 1185 |
+
"The dataset class summarizes the basic functionality of a dataset in a natural way. To define a dataset in PyTorch, we simply specify two functions: `__getitem__`, and `__len__`. The get-item function has to return the $i$-th data point in the dataset, while the len function returns the size of the dataset. For the XNOR dataset, we can define the dataset class as follows:"
|
| 1186 |
+
]
|
| 1187 |
+
},
|
| 1188 |
+
{
|
| 1189 |
+
"cell_type": "code",
|
| 1190 |
+
"execution_count": null,
|
| 1191 |
+
"metadata": {
|
| 1192 |
+
"id": "_VOhsu4EyKEC"
|
| 1193 |
+
},
|
| 1194 |
+
"outputs": [],
|
| 1195 |
+
"source": [
|
| 1196 |
+
"class XNORDataset(data.Dataset):\n",
|
| 1197 |
+
"\n",
|
| 1198 |
+
" def __init__(self, size, std=0.1):\n",
|
| 1199 |
+
" \"\"\"\n",
|
| 1200 |
+
" Inputs:\n",
|
| 1201 |
+
" size - Number of data points we want to generate\n",
|
| 1202 |
+
" std - Standard deviation of the noise (see generate_continuous_xnor function)\n",
|
| 1203 |
+
" \"\"\"\n",
|
| 1204 |
+
" super().__init__()\n",
|
| 1205 |
+
" self.size = size\n",
|
| 1206 |
+
" self.std = std\n",
|
| 1207 |
+
" self.generate_continuous_xnor()\n",
|
| 1208 |
+
"\n",
|
| 1209 |
+
" def generate_continuous_xnor(self):\n",
|
| 1210 |
+
" # Each data point in the XNOR dataset has two variables, x and y, that can be either 0 or 1.\n",
|
| 1211 |
+
" # The label is their XNOR combination, i.e. 1 if only x equal to y and 0 if x not equal to y.\n",
|
| 1212 |
+
" data = torch.randint(low=0, high=2, size=(self.size, 2), dtype=torch.float32)\n",
|
| 1213 |
+
" label = (data[:, 0] == data[:, 1]).to(torch.long)\n",
|
| 1214 |
+
" # To make it slightly more challenging, we add a bit of gaussian noise to the data points.\n",
|
| 1215 |
+
" data += self.std * torch.randn(data.shape)\n",
|
| 1216 |
+
"\n",
|
| 1217 |
+
" self.data = data\n",
|
| 1218 |
+
" self.label = label\n",
|
| 1219 |
+
"\n",
|
| 1220 |
+
" def __len__(self):\n",
|
| 1221 |
+
" # Number of data point we have. Alternatively self.data.shape[0], or self.label.shape[0]\n",
|
| 1222 |
+
" return self.size\n",
|
| 1223 |
+
"\n",
|
| 1224 |
+
" def __getitem__(self, idx):\n",
|
| 1225 |
+
" # Return the idx-th data point of the dataset\n",
|
| 1226 |
+
" # If we have multiple things to return (data point and label), we can return them as tuple\n",
|
| 1227 |
+
" data_point = self.data[idx]\n",
|
| 1228 |
+
" data_label = self.label[idx]\n",
|
| 1229 |
+
" return data_point, data_label"
|
| 1230 |
+
]
|
| 1231 |
+
},
|
| 1232 |
+
{
|
| 1233 |
+
"cell_type": "markdown",
|
| 1234 |
+
"metadata": {
|
| 1235 |
+
"id": "8vuUacxOyKEC"
|
| 1236 |
+
},
|
| 1237 |
+
"source": [
|
| 1238 |
+
"Let's try to create such a dataset and inspect it:"
|
| 1239 |
+
]
|
| 1240 |
+
},
|
| 1241 |
+
{
|
| 1242 |
+
"cell_type": "code",
|
| 1243 |
+
"execution_count": null,
|
| 1244 |
+
"metadata": {
|
| 1245 |
+
"id": "LyHNFoEvyKEC"
|
| 1246 |
+
},
|
| 1247 |
+
"outputs": [],
|
| 1248 |
+
"source": [
|
| 1249 |
+
"dataset = XNORDataset(size=200)\n",
|
| 1250 |
+
"print(\"Size of dataset:\", len(dataset))\n",
|
| 1251 |
+
"print(\"Data point 0:\", dataset[0])"
|
| 1252 |
+
]
|
| 1253 |
+
},
|
| 1254 |
+
{
|
| 1255 |
+
"cell_type": "markdown",
|
| 1256 |
+
"metadata": {
|
| 1257 |
+
"id": "2V6Cfz7AyKED"
|
| 1258 |
+
},
|
| 1259 |
+
"source": [
|
| 1260 |
+
"To better relate to the dataset, we visualize the samples below."
|
| 1261 |
+
]
|
| 1262 |
+
},
|
| 1263 |
+
{
|
| 1264 |
+
"cell_type": "code",
|
| 1265 |
+
"execution_count": null,
|
| 1266 |
+
"metadata": {
|
| 1267 |
+
"id": "0wZalS9dyKED"
|
| 1268 |
+
},
|
| 1269 |
+
"outputs": [],
|
| 1270 |
+
"source": [
|
| 1271 |
+
"import matplotlib.pyplot as plt\n",
|
| 1272 |
+
"%matplotlib inline\n",
|
| 1273 |
+
"def visualize_samples(data, label):\n",
|
| 1274 |
+
" if isinstance(data, torch.Tensor):\n",
|
| 1275 |
+
" data = data.cpu().numpy()\n",
|
| 1276 |
+
" if isinstance(label, torch.Tensor):\n",
|
| 1277 |
+
" label = label.cpu().numpy()\n",
|
| 1278 |
+
" data_0 = data[label == 0]\n",
|
| 1279 |
+
" data_1 = data[label == 1]\n",
|
| 1280 |
+
"\n",
|
| 1281 |
+
" plt.figure(figsize=(4,4))\n",
|
| 1282 |
+
" plt.scatter(data_0[:,0], data_0[:,1], edgecolor=\"#333\", label=\"Class 0\")\n",
|
| 1283 |
+
" plt.scatter(data_1[:,0], data_1[:,1], edgecolor=\"#333\", label=\"Class 1\")\n",
|
| 1284 |
+
" plt.title(\"Dataset samples\")\n",
|
| 1285 |
+
" plt.ylabel(r\"$x_2$\")\n",
|
| 1286 |
+
" plt.xlabel(r\"$x_1$\")\n",
|
| 1287 |
+
" plt.legend()"
|
| 1288 |
+
]
|
| 1289 |
+
},
|
| 1290 |
+
{
|
| 1291 |
+
"cell_type": "code",
|
| 1292 |
+
"execution_count": null,
|
| 1293 |
+
"metadata": {
|
| 1294 |
+
"id": "BwIPc7RPyKED"
|
| 1295 |
+
},
|
| 1296 |
+
"outputs": [],
|
| 1297 |
+
"source": [
|
| 1298 |
+
"visualize_samples(dataset.data, dataset.label)\n",
|
| 1299 |
+
"plt.show()"
|
| 1300 |
+
]
|
| 1301 |
+
},
|
| 1302 |
+
{
|
| 1303 |
+
"cell_type": "markdown",
|
| 1304 |
+
"metadata": {
|
| 1305 |
+
"id": "vXxG7wkVyKED"
|
| 1306 |
+
},
|
| 1307 |
+
"source": [
|
| 1308 |
+
"#### The data loader class\n",
|
| 1309 |
+
"\n",
|
| 1310 |
+
"The class `torch.utils.data.DataLoader` represents a Python iterable over a dataset with support for automatic batching, multi-process data loading and many more features. The data loader communicates with the dataset using the function `__getitem__()`, and stacks its outputs as tensors over the first dimension to form a batch.\n",
|
| 1311 |
+
"In contrast to the dataset class, we usually don't have to define our own data loader class, but can just create an object of it with the dataset as input. Additionally, we can configure our data loader with the following input arguments (only a selection, see full list [here](https://pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader)):\n",
|
| 1312 |
+
"\n",
|
| 1313 |
+
"* `batch_size`: Number of samples to stack per batch\n",
|
| 1314 |
+
"* `shuffle`: If True, the data is returned in a random order. This is important during training for introducing stochasticity.\n",
|
| 1315 |
+
"* `num_workers`: Number of subprocesses to use for data loading. The default, 0, means that the data will be loaded in the main process which can slow down training for datasets where loading a data point takes a considerable amount of time (e.g. large images). More workers are recommended for those, but can cause issues on Windows computers. For tiny datasets as ours, 0 workers are usually faster.\n",
|
| 1316 |
+
"* `pin_memory`: If True, the data loader will copy Tensors into CUDA pinned memory before returning them. This can save some time for large data points on GPUs. Usually a good practice to use for a training set, but not necessarily for validation and test to save memory on the GPU.\n",
|
| 1317 |
+
"* `drop_last`: If True, the last batch is dropped in case it is smaller than the specified batch size. This occurs when the dataset size is not a multiple of the batch size. Only potentially helpful during training to keep a consistent batch size.\n",
|
| 1318 |
+
"\n",
|
| 1319 |
+
"Let's create a simple data loader below:"
|
| 1320 |
+
]
|
| 1321 |
+
},
|
| 1322 |
+
{
|
| 1323 |
+
"cell_type": "code",
|
| 1324 |
+
"execution_count": null,
|
| 1325 |
+
"metadata": {
|
| 1326 |
+
"id": "Zu_v2xxByKED"
|
| 1327 |
+
},
|
| 1328 |
+
"outputs": [],
|
| 1329 |
+
"source": [
|
| 1330 |
+
"data_loader = data.DataLoader(dataset, batch_size=8, shuffle=True)"
|
| 1331 |
+
]
|
| 1332 |
+
},
|
| 1333 |
+
{
|
| 1334 |
+
"cell_type": "code",
|
| 1335 |
+
"execution_count": null,
|
| 1336 |
+
"metadata": {
|
| 1337 |
+
"id": "mJoQcFgmyKED"
|
| 1338 |
+
},
|
| 1339 |
+
"outputs": [],
|
| 1340 |
+
"source": [
|
| 1341 |
+
"# next(iter(...)) catches the first batch of the data loader\n",
|
| 1342 |
+
"# If shuffle is True, this will return a different batch every time we run this cell\n",
|
| 1343 |
+
"# For iterating over the whole dataset, we can simple use \"for batch in data_loader: ...\"\n",
|
| 1344 |
+
"data_inputs, data_labels = next(iter(data_loader))\n",
|
| 1345 |
+
"\n",
|
| 1346 |
+
"# The shape of the outputs are [batch_size, d_1,...,d_N] where d_1,...,d_N are the\n",
|
| 1347 |
+
"# dimensions of the data point returned from the dataset class\n",
|
| 1348 |
+
"print(\"Data inputs\", data_inputs.shape, \"\\n\", data_inputs)\n",
|
| 1349 |
+
"print(\"Data labels\", data_labels.shape, \"\\n\", data_labels)"
|
| 1350 |
+
]
|
| 1351 |
+
},
|
| 1352 |
+
{
|
| 1353 |
+
"cell_type": "markdown",
|
| 1354 |
+
"metadata": {
|
| 1355 |
+
"id": "1zcfNLYtv4ic"
|
| 1356 |
+
},
|
| 1357 |
+
"source": [
|
| 1358 |
+
"### Optimization"
|
| 1359 |
+
]
|
| 1360 |
+
},
|
| 1361 |
+
{
|
| 1362 |
+
"cell_type": "markdown",
|
| 1363 |
+
"metadata": {
|
| 1364 |
+
"id": "2ShudEaKyKED"
|
| 1365 |
+
},
|
| 1366 |
+
"source": [
|
| 1367 |
+
"After defining the model and the dataset, it is time to prepare the optimization of the model. During training, we will perform the following steps:\n",
|
| 1368 |
+
"\n",
|
| 1369 |
+
"1. Get a batch from the data loader\n",
|
| 1370 |
+
"2. Obtain the predictions from the model for the batch\n",
|
| 1371 |
+
"3. Calculate the loss based on the difference between predictions and labels\n",
|
| 1372 |
+
"4. Backpropagation: calculate the gradients for every parameter with respect to the loss\n",
|
| 1373 |
+
"5. Update the parameters of the model in the direction of the gradients\n",
|
| 1374 |
+
"\n",
|
| 1375 |
+
"We have seen how we can do step 1, 2 and 4 in PyTorch. Now, we will look at step 3 and 5."
|
| 1376 |
+
]
|
| 1377 |
+
},
|
| 1378 |
+
{
|
| 1379 |
+
"cell_type": "markdown",
|
| 1380 |
+
"metadata": {
|
| 1381 |
+
"id": "mMcn4KdEv1JN"
|
| 1382 |
+
},
|
| 1383 |
+
"source": [
|
| 1384 |
+
"#### Loss modules"
|
| 1385 |
+
]
|
| 1386 |
+
},
|
| 1387 |
+
{
|
| 1388 |
+
"cell_type": "markdown",
|
| 1389 |
+
"metadata": {
|
| 1390 |
+
"id": "aR-L_Wa7yKEE"
|
| 1391 |
+
},
|
| 1392 |
+
"source": [
|
| 1393 |
+
"We can calculate the loss for a batch by simply performing a few tensor operations as those are automatically added to the computation graph. For instance, for binary classification, we can use Binary Cross Entropy (BCE) which is defined as follows:\n",
|
| 1394 |
+
"\n",
|
| 1395 |
+
"$$\\mathcal{L}_{BCE} = -\\sum_i \\left[ y_i \\log x_i + (1 - y_i) \\log (1 - x_i) \\right]$$\n",
|
| 1396 |
+
"\n",
|
| 1397 |
+
"where $y$ are our labels, and $x$ our predictions, both in the range of $[0,1]$. However, PyTorch already provides a list of predefined loss functions which we can use (see [here](https://pytorch.org/docs/stable/nn.html#loss-functions) for a full list). For instance, for BCE, PyTorch has two modules: `nn.BCELoss()`, `nn.BCEWithLogitsLoss()`. While `nn.BCELoss` expects the inputs $x$ to be in the range $[0,1]$, i.e. the output of a sigmoid, `nn.BCEWithLogitsLoss` combines a sigmoid layer and the BCE loss in a single class. This version is numerically more stable than using a plain Sigmoid followed by a BCE loss because of the logarithms applied in the loss function. Hence, it is adviced to use loss functions applied on \"logits\" where possible (remember to not apply a sigmoid on the output of the model in this case!). For our model defined above, we therefore use the module `nn.BCEWithLogitsLoss`."
|
| 1398 |
+
]
|
| 1399 |
+
},
|
| 1400 |
+
{
|
| 1401 |
+
"cell_type": "code",
|
| 1402 |
+
"execution_count": null,
|
| 1403 |
+
"metadata": {
|
| 1404 |
+
"id": "fw4yFEH6yKEE"
|
| 1405 |
+
},
|
| 1406 |
+
"outputs": [],
|
| 1407 |
+
"source": [
|
| 1408 |
+
"loss_module = nn.BCEWithLogitsLoss()"
|
| 1409 |
+
]
|
| 1410 |
+
},
|
| 1411 |
+
{
|
| 1412 |
+
"cell_type": "markdown",
|
| 1413 |
+
"metadata": {
|
| 1414 |
+
"id": "6GlWgf_5vyFo"
|
| 1415 |
+
},
|
| 1416 |
+
"source": [
|
| 1417 |
+
"#### Stochastic Gradient Descent"
|
| 1418 |
+
]
|
| 1419 |
+
},
|
| 1420 |
+
{
|
| 1421 |
+
"cell_type": "markdown",
|
| 1422 |
+
"metadata": {
|
| 1423 |
+
"id": "Kc7inCxbyKEE"
|
| 1424 |
+
},
|
| 1425 |
+
"source": [
|
| 1426 |
+
"For updating the parameters, PyTorch provides the package `torch.optim` that has most popular optimizers implemented. We will discuss the specific optimizers and their differences later in the course, but will for now use the simplest of them: `torch.optim.SGD`. Stochastic Gradient Descent updates parameters by multiplying the gradients with a small constant, called learning rate, and subtracting those from the parameters (hence minimizing the loss). Therefore, we slowly move towards the direction of minimizing the loss. A good default value of the learning rate for a small network as ours is 0.1."
|
| 1427 |
+
]
|
| 1428 |
+
},
|
| 1429 |
+
{
|
| 1430 |
+
"cell_type": "code",
|
| 1431 |
+
"execution_count": null,
|
| 1432 |
+
"metadata": {
|
| 1433 |
+
"id": "538d8YaIyKEE"
|
| 1434 |
+
},
|
| 1435 |
+
"outputs": [],
|
| 1436 |
+
"source": [
|
| 1437 |
+
"# Input to the optimizer are the parameters of the model: model.parameters()\n",
|
| 1438 |
+
"optimizer = torch.optim.SGD(model.parameters(), lr=0.1)"
|
| 1439 |
+
]
|
| 1440 |
+
},
|
| 1441 |
+
{
|
| 1442 |
+
"cell_type": "markdown",
|
| 1443 |
+
"metadata": {
|
| 1444 |
+
"id": "A3OqLcuAyKEE"
|
| 1445 |
+
},
|
| 1446 |
+
"source": [
|
| 1447 |
+
"The optimizer provides two useful functions: `optimizer.step()`, and `optimizer.zero_grad()`. The step function updates the parameters based on the gradients as explained above. The function `optimizer.zero_grad()` sets the gradients of all parameters to zero. While this function seems less relevant at first, it is a crucial pre-step before performing backpropagation. If we call the `backward` function on the loss while the parameter gradients are non-zero from the previous batch, the new gradients would actually be added to the previous ones instead of overwriting them. This is done because a parameter might occur multiple times in a computation graph, and we need to sum the gradients in this case instead of replacing them. Hence, remember to call `optimizer.zero_grad()` before calculating the gradients of a batch."
|
| 1448 |
+
]
|
| 1449 |
+
},
|
| 1450 |
+
{
|
| 1451 |
+
"cell_type": "markdown",
|
| 1452 |
+
"metadata": {
|
| 1453 |
+
"id": "FpAceRBPyKEE"
|
| 1454 |
+
},
|
| 1455 |
+
"source": [
|
| 1456 |
+
"### Training\n",
|
| 1457 |
+
"\n",
|
| 1458 |
+
"Finally, we are ready to train our model. As a first step, we create a slightly larger dataset and specify a data loader with a larger batch size."
|
| 1459 |
+
]
|
| 1460 |
+
},
|
| 1461 |
+
{
|
| 1462 |
+
"cell_type": "code",
|
| 1463 |
+
"execution_count": null,
|
| 1464 |
+
"metadata": {
|
| 1465 |
+
"id": "fZvApVhdyKEE"
|
| 1466 |
+
},
|
| 1467 |
+
"outputs": [],
|
| 1468 |
+
"source": [
|
| 1469 |
+
"train_dataset = XNORDataset(size=2500)\n",
|
| 1470 |
+
"train_data_loader = data.DataLoader(train_dataset, batch_size=128, shuffle=True)"
|
| 1471 |
+
]
|
| 1472 |
+
},
|
| 1473 |
+
{
|
| 1474 |
+
"cell_type": "markdown",
|
| 1475 |
+
"metadata": {
|
| 1476 |
+
"id": "-wgVr8C5yKEE"
|
| 1477 |
+
},
|
| 1478 |
+
"source": [
|
| 1479 |
+
"Now, we can write a small training function. Remember our five steps: load a batch, obtain the predictions, calculate the loss, backpropagate, and update. Additionally, we have to push all data and model parameters to the device of our choice (GPU if available). For the tiny neural network we have, communicating the data to the GPU actually takes much more time than we could save from running the operation on GPU. For large networks, the communication time is significantly smaller than the actual runtime making a GPU crucial in these cases. Still, to practice, we will push the data to GPU here."
|
| 1480 |
+
]
|
| 1481 |
+
},
|
| 1482 |
+
{
|
| 1483 |
+
"cell_type": "code",
|
| 1484 |
+
"execution_count": null,
|
| 1485 |
+
"metadata": {
|
| 1486 |
+
"id": "xQ_By4XfyKEE"
|
| 1487 |
+
},
|
| 1488 |
+
"outputs": [],
|
| 1489 |
+
"source": [
|
| 1490 |
+
"# Push model to device. Has to be only done once\n",
|
| 1491 |
+
"model.to(device)"
|
| 1492 |
+
]
|
| 1493 |
+
},
|
| 1494 |
+
{
|
| 1495 |
+
"cell_type": "markdown",
|
| 1496 |
+
"metadata": {
|
| 1497 |
+
"id": "EuusR5sTyKEE"
|
| 1498 |
+
},
|
| 1499 |
+
"source": [
|
| 1500 |
+
"In addition, we set our model to training mode. This is done by calling `model.train()`. There exist certain modules that need to perform a different forward step during training than during testing (e.g. BatchNorm and Dropout), and we can switch between them using `model.train()` and `model.eval()`."
|
| 1501 |
+
]
|
| 1502 |
+
},
|
| 1503 |
+
{
|
| 1504 |
+
"cell_type": "code",
|
| 1505 |
+
"execution_count": null,
|
| 1506 |
+
"metadata": {
|
| 1507 |
+
"id": "4u3-tu1fyKEE"
|
| 1508 |
+
},
|
| 1509 |
+
"outputs": [],
|
| 1510 |
+
"source": [
|
| 1511 |
+
"from tqdm.notebook import tqdm\n",
|
| 1512 |
+
"def train_model(model, optimizer, data_loader, loss_module, num_epochs=100):\n",
|
| 1513 |
+
" # Set model to train mode\n",
|
| 1514 |
+
" model.train()\n",
|
| 1515 |
+
"\n",
|
| 1516 |
+
" # Training loop\n",
|
| 1517 |
+
" for epoch in tqdm(range(num_epochs)):\n",
|
| 1518 |
+
" for data_inputs, data_labels in data_loader:\n",
|
| 1519 |
+
"\n",
|
| 1520 |
+
" ## Step 1: Move input data to device (only strictly necessary if we use GPU)\n",
|
| 1521 |
+
" data_inputs = data_inputs.to(device)\n",
|
| 1522 |
+
" data_labels = data_labels.to(device)\n",
|
| 1523 |
+
"\n",
|
| 1524 |
+
" ## Step 2: Run the model on the input data\n",
|
| 1525 |
+
" preds = model(data_inputs)\n",
|
| 1526 |
+
" preds = preds.squeeze(dim=1) # Output is [Batch size, 1], but we want [Batch size]\n",
|
| 1527 |
+
"\n",
|
| 1528 |
+
" ## Step 3: Calculate the loss\n",
|
| 1529 |
+
" loss = loss_module(preds, data_labels.float())\n",
|
| 1530 |
+
"\n",
|
| 1531 |
+
" ## Step 4: Perform backpropagation\n",
|
| 1532 |
+
" # Before calculating the gradients, we need to ensure that they are all zero.\n",
|
| 1533 |
+
" # The gradients would not be overwritten, but actually added to the existing ones.\n",
|
| 1534 |
+
" optimizer.zero_grad()\n",
|
| 1535 |
+
" # Perform backpropagation\n",
|
| 1536 |
+
" loss.backward()\n",
|
| 1537 |
+
"\n",
|
| 1538 |
+
" ## Step 5: Update the parameters\n",
|
| 1539 |
+
" optimizer.step()"
|
| 1540 |
+
]
|
| 1541 |
+
},
|
| 1542 |
+
{
|
| 1543 |
+
"cell_type": "code",
|
| 1544 |
+
"execution_count": null,
|
| 1545 |
+
"metadata": {
|
| 1546 |
+
"id": "miyLOI08yKEF"
|
| 1547 |
+
},
|
| 1548 |
+
"outputs": [],
|
| 1549 |
+
"source": [
|
| 1550 |
+
"train_model(model, optimizer, train_data_loader, loss_module)"
|
| 1551 |
+
]
|
| 1552 |
+
},
|
| 1553 |
+
{
|
| 1554 |
+
"cell_type": "markdown",
|
| 1555 |
+
"metadata": {
|
| 1556 |
+
"id": "eJdcLsO3wEG6"
|
| 1557 |
+
},
|
| 1558 |
+
"source": [
|
| 1559 |
+
"#### Saving a model"
|
| 1560 |
+
]
|
| 1561 |
+
},
|
| 1562 |
+
{
|
| 1563 |
+
"cell_type": "markdown",
|
| 1564 |
+
"metadata": {
|
| 1565 |
+
"id": "Hq__j4kOyKEF"
|
| 1566 |
+
},
|
| 1567 |
+
"source": [
|
| 1568 |
+
"After finish training a model, we save the model to disk so that we can load the same weights at a later time. For this, we extract the so-called `state_dict` from the model which contains all learnable parameters. For our simple model, the state dict contains the following entries:"
|
| 1569 |
+
]
|
| 1570 |
+
},
|
| 1571 |
+
{
|
| 1572 |
+
"cell_type": "code",
|
| 1573 |
+
"execution_count": null,
|
| 1574 |
+
"metadata": {
|
| 1575 |
+
"id": "Me3HsoD-yKEF"
|
| 1576 |
+
},
|
| 1577 |
+
"outputs": [],
|
| 1578 |
+
"source": [
|
| 1579 |
+
"state_dict = model.state_dict()\n",
|
| 1580 |
+
"print(state_dict)"
|
| 1581 |
+
]
|
| 1582 |
+
},
|
| 1583 |
+
{
|
| 1584 |
+
"cell_type": "markdown",
|
| 1585 |
+
"metadata": {
|
| 1586 |
+
"id": "0SqbNIWKyKEF"
|
| 1587 |
+
},
|
| 1588 |
+
"source": [
|
| 1589 |
+
"To save the state dictionary, we can use `torch.save`:"
|
| 1590 |
+
]
|
| 1591 |
+
},
|
| 1592 |
+
{
|
| 1593 |
+
"cell_type": "code",
|
| 1594 |
+
"execution_count": null,
|
| 1595 |
+
"metadata": {
|
| 1596 |
+
"id": "fooo7SBzyKEF"
|
| 1597 |
+
},
|
| 1598 |
+
"outputs": [],
|
| 1599 |
+
"source": [
|
| 1600 |
+
"# torch.save(object, filename). For the filename, any extension can be used\n",
|
| 1601 |
+
"torch.save(state_dict, \"our_model.tar\")"
|
| 1602 |
+
]
|
| 1603 |
+
},
|
| 1604 |
+
{
|
| 1605 |
+
"cell_type": "markdown",
|
| 1606 |
+
"metadata": {
|
| 1607 |
+
"id": "TJe90OhZWIK-"
|
| 1608 |
+
},
|
| 1609 |
+
"source": [
|
| 1610 |
+
"#### Loading a model"
|
| 1611 |
+
]
|
| 1612 |
+
},
|
| 1613 |
+
{
|
| 1614 |
+
"cell_type": "markdown",
|
| 1615 |
+
"metadata": {
|
| 1616 |
+
"id": "QGuwzmRPyKEF"
|
| 1617 |
+
},
|
| 1618 |
+
"source": [
|
| 1619 |
+
"To load a model from a state dict, we use the function `torch.load` to load the state dict from the disk, and the module function `load_state_dict` to overwrite our parameters with the new values:"
|
| 1620 |
+
]
|
| 1621 |
+
},
|
| 1622 |
+
{
|
| 1623 |
+
"cell_type": "code",
|
| 1624 |
+
"execution_count": null,
|
| 1625 |
+
"metadata": {
|
| 1626 |
+
"id": "dGwKGN9_yKEF"
|
| 1627 |
+
},
|
| 1628 |
+
"outputs": [],
|
| 1629 |
+
"source": [
|
| 1630 |
+
"# Load state dict from the disk (make sure it is the same name as above)\n",
|
| 1631 |
+
"state_dict = torch.load(\"our_model.tar\")\n",
|
| 1632 |
+
"\n",
|
| 1633 |
+
"# Create a new model and load the state\n",
|
| 1634 |
+
"new_model = SimpleClassifier(num_inputs=2, num_hidden=4, num_outputs=1)\n",
|
| 1635 |
+
"new_model.load_state_dict(state_dict)\n",
|
| 1636 |
+
"\n",
|
| 1637 |
+
"# Verify that the parameters are the same\n",
|
| 1638 |
+
"print(\"Original model\\n\", model.state_dict())\n",
|
| 1639 |
+
"print(\"\\nLoaded model\\n\", new_model.state_dict())"
|
| 1640 |
+
]
|
| 1641 |
+
},
|
| 1642 |
+
{
|
| 1643 |
+
"cell_type": "markdown",
|
| 1644 |
+
"metadata": {
|
| 1645 |
+
"id": "nGgtDnIiyKEG"
|
| 1646 |
+
},
|
| 1647 |
+
"source": [
|
| 1648 |
+
"A detailed tutorial on saving and loading models in PyTorch can be found [here](https://pytorch.org/tutorials/beginner/saving_loading_models.html)."
|
| 1649 |
+
]
|
| 1650 |
+
},
|
| 1651 |
+
{
|
| 1652 |
+
"cell_type": "markdown",
|
| 1653 |
+
"metadata": {
|
| 1654 |
+
"id": "ssTU_brHwIIW"
|
| 1655 |
+
},
|
| 1656 |
+
"source": [
|
| 1657 |
+
"### Evaluation"
|
| 1658 |
+
]
|
| 1659 |
+
},
|
| 1660 |
+
{
|
| 1661 |
+
"cell_type": "markdown",
|
| 1662 |
+
"metadata": {
|
| 1663 |
+
"id": "hoUN-UEqyKEG"
|
| 1664 |
+
},
|
| 1665 |
+
"source": [
|
| 1666 |
+
"Once we have trained a model, it is time to evaluate it on a held-out test set. As our dataset consist of randomly generated data points, we need to first create a test set with a corresponding data loader."
|
| 1667 |
+
]
|
| 1668 |
+
},
|
| 1669 |
+
{
|
| 1670 |
+
"cell_type": "code",
|
| 1671 |
+
"execution_count": null,
|
| 1672 |
+
"metadata": {
|
| 1673 |
+
"id": "SPrpqQx2yKEG"
|
| 1674 |
+
},
|
| 1675 |
+
"outputs": [],
|
| 1676 |
+
"source": [
|
| 1677 |
+
"test_dataset = XNORDataset(size=500)\n",
|
| 1678 |
+
"# drop_last -> Don't drop the last batch although it is smaller than 128\n",
|
| 1679 |
+
"test_data_loader = data.DataLoader(test_dataset, batch_size=128, shuffle=False, drop_last=False)"
|
| 1680 |
+
]
|
| 1681 |
+
},
|
| 1682 |
+
{
|
| 1683 |
+
"cell_type": "markdown",
|
| 1684 |
+
"metadata": {
|
| 1685 |
+
"id": "q2kpq9yKyKEG"
|
| 1686 |
+
},
|
| 1687 |
+
"source": [
|
| 1688 |
+
"As metric, we will use accuracy which is calculated as follows:\n",
|
| 1689 |
+
"\n",
|
| 1690 |
+
"$$acc = \\frac{\\#\\text{correct predictions}}{\\#\\text{all predictions}} = \\frac{TP+TN}{TP+TN+FP+FN}$$\n",
|
| 1691 |
+
"\n",
|
| 1692 |
+
"where TP are the true positives, TN true negatives, FP false positives, and FN the fale negatives.\n",
|
| 1693 |
+
"\n",
|
| 1694 |
+
"When evaluating the model, we don't need to keep track of the computation graph as we don't intend to calculate the gradients. This reduces the required memory and speed up the model. In PyTorch, we can deactivate the computation graph using `with torch.no_grad(): ...`. Remember to additionally set the model to eval mode."
|
| 1695 |
+
]
|
| 1696 |
+
},
|
| 1697 |
+
{
|
| 1698 |
+
"cell_type": "code",
|
| 1699 |
+
"execution_count": null,
|
| 1700 |
+
"metadata": {
|
| 1701 |
+
"id": "4tK9mh5NyKEG"
|
| 1702 |
+
},
|
| 1703 |
+
"outputs": [],
|
| 1704 |
+
"source": [
|
| 1705 |
+
"def eval_model(model, data_loader):\n",
|
| 1706 |
+
" model.eval() # Set model to eval mode\n",
|
| 1707 |
+
" true_preds, num_preds = 0., 0.\n",
|
| 1708 |
+
"\n",
|
| 1709 |
+
" with torch.no_grad(): # Deactivate gradients for the following code\n",
|
| 1710 |
+
" for data_inputs, data_labels in data_loader:\n",
|
| 1711 |
+
"\n",
|
| 1712 |
+
" # Determine prediction of model on dev set\n",
|
| 1713 |
+
" data_inputs, data_labels = data_inputs.to(device), data_labels.to(device)\n",
|
| 1714 |
+
" preds = model(data_inputs)\n",
|
| 1715 |
+
" preds = preds.squeeze(dim=1)\n",
|
| 1716 |
+
" preds = torch.sigmoid(preds) # Sigmoid to map predictions between 0 and 1\n",
|
| 1717 |
+
" pred_labels = (preds >= 0.5).long() # Binarize predictions to 0 and 1\n",
|
| 1718 |
+
"\n",
|
| 1719 |
+
" # Keep records of predictions for the accuracy metric (true_preds=TP+TN, num_preds=TP+TN+FP+FN)\n",
|
| 1720 |
+
" true_preds += (pred_labels == data_labels).sum()\n",
|
| 1721 |
+
" num_preds += data_labels.shape[0]\n",
|
| 1722 |
+
"\n",
|
| 1723 |
+
" acc = true_preds / num_preds\n",
|
| 1724 |
+
" print(f\"Accuracy of the model: {100.0*acc:4.2f}%\")"
|
| 1725 |
+
]
|
| 1726 |
+
},
|
| 1727 |
+
{
|
| 1728 |
+
"cell_type": "code",
|
| 1729 |
+
"execution_count": null,
|
| 1730 |
+
"metadata": {
|
| 1731 |
+
"id": "ZB_AtTqQyKEG"
|
| 1732 |
+
},
|
| 1733 |
+
"outputs": [],
|
| 1734 |
+
"source": [
|
| 1735 |
+
"eval_model(model, test_data_loader)"
|
| 1736 |
+
]
|
| 1737 |
+
},
|
| 1738 |
+
{
|
| 1739 |
+
"cell_type": "markdown",
|
| 1740 |
+
"metadata": {
|
| 1741 |
+
"id": "XBy98EtxyKEG"
|
| 1742 |
+
},
|
| 1743 |
+
"source": [
|
| 1744 |
+
"If we trained our model correctly, we should see a score close to 100% accuracy. However, this is only possible because of our simple task, and unfortunately, we usually don't get such high scores on test sets of more complex tasks."
|
| 1745 |
+
]
|
| 1746 |
+
},
|
| 1747 |
+
{
|
| 1748 |
+
"cell_type": "markdown",
|
| 1749 |
+
"metadata": {
|
| 1750 |
+
"id": "Z34M30BSwL44"
|
| 1751 |
+
},
|
| 1752 |
+
"source": [
|
| 1753 |
+
"#### Classification boundaries"
|
| 1754 |
+
]
|
| 1755 |
+
},
|
| 1756 |
+
{
|
| 1757 |
+
"cell_type": "markdown",
|
| 1758 |
+
"metadata": {
|
| 1759 |
+
"id": "CMLkSndpyKEH"
|
| 1760 |
+
},
|
| 1761 |
+
"source": [
|
| 1762 |
+
"To visualize what our model has learned, we can perform a prediction for every data point in a range of $[-0.5, 1.5]$, and visualize the predicted class as in the sample figure at the beginning of this section. This shows where the model has created decision boundaries, and which points would be classified as $0$, and which as $1$. We therefore get a background image out of blue (class 0) and orange (class 1). The spots where the model is uncertain we will see a blurry overlap. The specific code is less relevant compared to the output figure which should hopefully show us a clear separation of classes:"
|
| 1763 |
+
]
|
| 1764 |
+
},
|
| 1765 |
+
{
|
| 1766 |
+
"cell_type": "code",
|
| 1767 |
+
"execution_count": null,
|
| 1768 |
+
"metadata": {
|
| 1769 |
+
"id": "_T5bkqYnyKEH"
|
| 1770 |
+
},
|
| 1771 |
+
"outputs": [],
|
| 1772 |
+
"source": [
|
| 1773 |
+
"from matplotlib.colors import to_rgba\n",
|
| 1774 |
+
"@torch.no_grad() # Decorator, same effect as \"with torch.no_grad(): ...\" over the whole function.\n",
|
| 1775 |
+
"def visualize_classification(model, data, label):\n",
|
| 1776 |
+
" if isinstance(data, torch.Tensor):\n",
|
| 1777 |
+
" data = data.cpu().numpy()\n",
|
| 1778 |
+
" if isinstance(label, torch.Tensor):\n",
|
| 1779 |
+
" label = label.cpu().numpy()\n",
|
| 1780 |
+
" data_0 = data[label == 0]\n",
|
| 1781 |
+
" data_1 = data[label == 1]\n",
|
| 1782 |
+
"\n",
|
| 1783 |
+
" fig = plt.figure(figsize=(4,4), dpi=500)\n",
|
| 1784 |
+
" plt.scatter(data_0[:,0], data_0[:,1], edgecolor=\"#333\", label=\"Class 0\")\n",
|
| 1785 |
+
" plt.scatter(data_1[:,0], data_1[:,1], edgecolor=\"#333\", label=\"Class 1\")\n",
|
| 1786 |
+
" plt.title(\"Dataset samples\")\n",
|
| 1787 |
+
" plt.ylabel(r\"$x_2$\")\n",
|
| 1788 |
+
" plt.xlabel(r\"$x_1$\")\n",
|
| 1789 |
+
" plt.legend()\n",
|
| 1790 |
+
"\n",
|
| 1791 |
+
" # Let's make use of a lot of operations we have learned above\n",
|
| 1792 |
+
" model.to(device)\n",
|
| 1793 |
+
" c0 = torch.Tensor(to_rgba(\"C0\")).to(device)\n",
|
| 1794 |
+
" c1 = torch.Tensor(to_rgba(\"C1\")).to(device)\n",
|
| 1795 |
+
" x1 = torch.arange(-0.5, 1.5, step=0.01, device=device)\n",
|
| 1796 |
+
" x2 = torch.arange(-0.5, 1.5, step=0.01, device=device)\n",
|
| 1797 |
+
" xx1, xx2 = torch.meshgrid(x1, x2, indexing='ij') # Meshgrid function as in Numpy\n",
|
| 1798 |
+
" model_inputs = torch.stack([xx1, xx2], dim=-1)\n",
|
| 1799 |
+
" preds = model(model_inputs)\n",
|
| 1800 |
+
" preds = torch.sigmoid(preds)\n",
|
| 1801 |
+
" output_image = (1 - preds) * c0[None,None] + preds * c1[None,None] # Specifying \"None\" in a dimension creates a new one\n",
|
| 1802 |
+
" output_image = output_image.cpu().numpy() # Convert to Numpy array. This only works for tensors on CPU, hence first push to CPU\n",
|
| 1803 |
+
" plt.imshow(output_image, origin='lower', extent=(-0.5, 1.5, -0.5, 1.5))\n",
|
| 1804 |
+
" plt.grid(False)\n",
|
| 1805 |
+
" return fig\n",
|
| 1806 |
+
"\n",
|
| 1807 |
+
"_ = visualize_classification(model, dataset.data, dataset.label)\n",
|
| 1808 |
+
"plt.show()"
|
| 1809 |
+
]
|
| 1810 |
+
},
|
| 1811 |
+
{
|
| 1812 |
+
"cell_type": "markdown",
|
| 1813 |
+
"metadata": {
|
| 1814 |
+
"id": "Xd0V5xURyKEH"
|
| 1815 |
+
},
|
| 1816 |
+
"source": [
|
| 1817 |
+
"The decision boundaries might not look exactly as in the figure in the preamble of this section which can be caused by running it on CPU or a different GPU architecture. Nevertheless, the result on the accuracy metric should be the approximately the same."
|
| 1818 |
+
]
|
| 1819 |
+
},
|
| 1820 |
+
{
|
| 1821 |
+
"cell_type": "markdown",
|
| 1822 |
+
"metadata": {
|
| 1823 |
+
"id": "fD0bZOtWwO3N"
|
| 1824 |
+
},
|
| 1825 |
+
"source": [
|
| 1826 |
+
"## Additional features"
|
| 1827 |
+
]
|
| 1828 |
+
},
|
| 1829 |
+
{
|
| 1830 |
+
"cell_type": "markdown",
|
| 1831 |
+
"metadata": {
|
| 1832 |
+
"id": "SWS3SbxEyKEH"
|
| 1833 |
+
},
|
| 1834 |
+
"source": [
|
| 1835 |
+
"Finally, you are all set to start with your own PyTorch project! In summary, we have looked at how we can build neural networks in PyTorch, and train and test them on data. However, there is still much more to PyTorch we haven't discussed yet. In the coming series of Jupyter notebooks, we will discover more and more functionalities of PyTorch, so that you also get familiar to PyTorch concepts beyond the basics. If you are already interested in learning more of PyTorch, we recommend the official [tutorial website](https://pytorch.org/tutorials/) that contains many tutorials on various topics. Especially logging with Tensorboard ([official tutorial here](https://pytorch.org/tutorials/intermediate/tensorboard_tutorial.html)) is a very good practice. Nonetheless, let's check it shortly out how we could use TensorBoard in our small example."
|
| 1836 |
+
]
|
| 1837 |
+
},
|
| 1838 |
+
{
|
| 1839 |
+
"cell_type": "markdown",
|
| 1840 |
+
"metadata": {
|
| 1841 |
+
"id": "LJDPWpDCwW_b"
|
| 1842 |
+
},
|
| 1843 |
+
"source": [
|
| 1844 |
+
"### TensorBoard logging"
|
| 1845 |
+
]
|
| 1846 |
+
},
|
| 1847 |
+
{
|
| 1848 |
+
"cell_type": "markdown",
|
| 1849 |
+
"metadata": {
|
| 1850 |
+
"id": "T62gqr2ByKEH"
|
| 1851 |
+
},
|
| 1852 |
+
"source": [
|
| 1853 |
+
"TensorBoard is a logging and visualization tool that is a popular choice for training deep learning models. Although initially published for TensorFlow, TensorBoard is also integrated in PyTorch allowing us to easily use it. First, let's import it below."
|
| 1854 |
+
]
|
| 1855 |
+
},
|
| 1856 |
+
{
|
| 1857 |
+
"cell_type": "code",
|
| 1858 |
+
"execution_count": null,
|
| 1859 |
+
"metadata": {
|
| 1860 |
+
"id": "0gh6VV8dyKEH"
|
| 1861 |
+
},
|
| 1862 |
+
"outputs": [],
|
| 1863 |
+
"source": [
|
| 1864 |
+
"# Import tensorboard logger from PyTorch\n",
|
| 1865 |
+
"from torch.utils.tensorboard import SummaryWriter\n",
|
| 1866 |
+
"\n",
|
| 1867 |
+
"# Load tensorboard extension for Jupyter Notebook, only need to start TB in the notebook\n",
|
| 1868 |
+
"%load_ext tensorboard"
|
| 1869 |
+
]
|
| 1870 |
+
},
|
| 1871 |
+
{
|
| 1872 |
+
"cell_type": "markdown",
|
| 1873 |
+
"metadata": {
|
| 1874 |
+
"id": "Ilra_iscyKEH"
|
| 1875 |
+
},
|
| 1876 |
+
"source": [
|
| 1877 |
+
"The last line is required if you want to run TensorBoard directly in the Jupyter Notebook. Otherwise, you can start TensorBoard from the terminal.\n",
|
| 1878 |
+
"\n",
|
| 1879 |
+
"PyTorch's TensorBoard API is simple to use. We start the logging process by creating a new object, `writer = SummaryWriter(...)`, where we specify the directory in which the logging file should be saved. With this object, we can log different aspects of our model by calling functions of the style `writer.add_...`. For example, we can visualize the computation graph with the function `writer.add_graph`, or add a scalar value like the loss with `writer.add_scalar`. Let's adapt our initial training function with adding a TensorBoard logger below."
|
| 1880 |
+
]
|
| 1881 |
+
},
|
| 1882 |
+
{
|
| 1883 |
+
"cell_type": "code",
|
| 1884 |
+
"execution_count": null,
|
| 1885 |
+
"metadata": {
|
| 1886 |
+
"id": "BpiwqUmxyKEH"
|
| 1887 |
+
},
|
| 1888 |
+
"outputs": [],
|
| 1889 |
+
"source": [
|
| 1890 |
+
"def train_model_with_logger(model, optimizer, data_loader, loss_module, val_dataset, num_epochs=100, logging_dir='runs/our_experiment'):\n",
|
| 1891 |
+
" # Create TensorBoard logger\n",
|
| 1892 |
+
" writer = SummaryWriter(logging_dir)\n",
|
| 1893 |
+
" model_plotted = False\n",
|
| 1894 |
+
"\n",
|
| 1895 |
+
" # Set model to train mode\n",
|
| 1896 |
+
" model.train()\n",
|
| 1897 |
+
"\n",
|
| 1898 |
+
" # Training loop\n",
|
| 1899 |
+
" for epoch in tqdm(range(num_epochs)):\n",
|
| 1900 |
+
" epoch_loss = 0.0\n",
|
| 1901 |
+
" for data_inputs, data_labels in data_loader:\n",
|
| 1902 |
+
"\n",
|
| 1903 |
+
" ## Step 1: Move input data to device (only strictly necessary if we use GPU)\n",
|
| 1904 |
+
" data_inputs = data_inputs.to(device)\n",
|
| 1905 |
+
" data_labels = data_labels.to(device)\n",
|
| 1906 |
+
"\n",
|
| 1907 |
+
" # For the very first batch, we visualize the computation graph in TensorBoard\n",
|
| 1908 |
+
" if not model_plotted:\n",
|
| 1909 |
+
" writer.add_graph(model, data_inputs)\n",
|
| 1910 |
+
" model_plotted = True\n",
|
| 1911 |
+
"\n",
|
| 1912 |
+
" ## Step 2: Run the model on the input data\n",
|
| 1913 |
+
" preds = model(data_inputs)\n",
|
| 1914 |
+
" preds = preds.squeeze(dim=1) # Output is [Batch size, 1], but we want [Batch size]\n",
|
| 1915 |
+
"\n",
|
| 1916 |
+
" ## Step 3: Calculate the loss\n",
|
| 1917 |
+
" loss = loss_module(preds, data_labels.float())\n",
|
| 1918 |
+
"\n",
|
| 1919 |
+
" ## Step 4: Perform backpropagation\n",
|
| 1920 |
+
" # Before calculating the gradients, we need to ensure that they are all zero.\n",
|
| 1921 |
+
" # The gradients would not be overwritten, but actually added to the existing ones.\n",
|
| 1922 |
+
" optimizer.zero_grad()\n",
|
| 1923 |
+
" # Perform backpropagation\n",
|
| 1924 |
+
" loss.backward()\n",
|
| 1925 |
+
"\n",
|
| 1926 |
+
" ## Step 5: Update the parameters\n",
|
| 1927 |
+
" optimizer.step()\n",
|
| 1928 |
+
"\n",
|
| 1929 |
+
" ## Step 6: Take the running average of the loss\n",
|
| 1930 |
+
" epoch_loss += loss.item()\n",
|
| 1931 |
+
"\n",
|
| 1932 |
+
" # Add average loss to TensorBoard\n",
|
| 1933 |
+
" epoch_loss /= len(data_loader)\n",
|
| 1934 |
+
" writer.add_scalar('training_loss',\n",
|
| 1935 |
+
" epoch_loss,\n",
|
| 1936 |
+
" global_step = epoch + 1)\n",
|
| 1937 |
+
"\n",
|
| 1938 |
+
" # Visualize prediction and add figure to TensorBoard\n",
|
| 1939 |
+
" # Since matplotlib figures can be slow in rendering, we only do it every 10th epoch\n",
|
| 1940 |
+
" if (epoch + 1) % 10 == 0:\n",
|
| 1941 |
+
" fig = visualize_classification(model, val_dataset.data, val_dataset.label)\n",
|
| 1942 |
+
" writer.add_figure('predictions',\n",
|
| 1943 |
+
" fig,\n",
|
| 1944 |
+
" global_step = epoch + 1)\n",
|
| 1945 |
+
"\n",
|
| 1946 |
+
" writer.close()"
|
| 1947 |
+
]
|
| 1948 |
+
},
|
| 1949 |
+
{
|
| 1950 |
+
"cell_type": "markdown",
|
| 1951 |
+
"metadata": {
|
| 1952 |
+
"id": "2ionZe79yKEI"
|
| 1953 |
+
},
|
| 1954 |
+
"source": [
|
| 1955 |
+
"Let's use this method to train a model as before, with a new model and optimizer."
|
| 1956 |
+
]
|
| 1957 |
+
},
|
| 1958 |
+
{
|
| 1959 |
+
"cell_type": "code",
|
| 1960 |
+
"execution_count": null,
|
| 1961 |
+
"metadata": {
|
| 1962 |
+
"id": "w2qOWXWFyKEI"
|
| 1963 |
+
},
|
| 1964 |
+
"outputs": [],
|
| 1965 |
+
"source": [
|
| 1966 |
+
"model = SimpleClassifier(num_inputs=2, num_hidden=4, num_outputs=1).to(device)\n",
|
| 1967 |
+
"optimizer = torch.optim.SGD(model.parameters(), lr=0.1)\n",
|
| 1968 |
+
"train_model_with_logger(model, optimizer, train_data_loader, loss_module, val_dataset=dataset)"
|
| 1969 |
+
]
|
| 1970 |
+
},
|
| 1971 |
+
{
|
| 1972 |
+
"cell_type": "markdown",
|
| 1973 |
+
"metadata": {
|
| 1974 |
+
"id": "Nntp29iSyKEI"
|
| 1975 |
+
},
|
| 1976 |
+
"source": [
|
| 1977 |
+
"The TensorBoard file in the folder `runs/our_experiment` now contains a loss curve, the computation graph of our network, and a visualization of the learned predictions over number of epochs. To start the TensorBoard visualizer, simply run the following statement:"
|
| 1978 |
+
]
|
| 1979 |
+
},
|
| 1980 |
+
{
|
| 1981 |
+
"cell_type": "code",
|
| 1982 |
+
"execution_count": null,
|
| 1983 |
+
"metadata": {
|
| 1984 |
+
"id": "VTmYZzCCyKEI"
|
| 1985 |
+
},
|
| 1986 |
+
"outputs": [],
|
| 1987 |
+
"source": [
|
| 1988 |
+
"%tensorboard --logdir runs/our_experiment"
|
| 1989 |
+
]
|
| 1990 |
+
},
|
| 1991 |
+
{
|
| 1992 |
+
"cell_type": "markdown",
|
| 1993 |
+
"metadata": {
|
| 1994 |
+
"id": "QOlapcqzyKEI"
|
| 1995 |
+
},
|
| 1996 |
+
"source": [
|
| 1997 |
+
"<center><img src=\"https://github.com/AnjanDutta/SharedFigures/blob/main/tensorboard_screenshot.png?raw=true\" width=\"600px\"></center>\n",
|
| 1998 |
+
"\n",
|
| 1999 |
+
"TensorBoard visualizations can help to identify possible issues with your model, and identify situations such as overfitting. You can also track the training progress while a model is training, since the logger automatically writes everything added to it to the logging file. Feel free to explore the TensorBoard functionalities."
|
| 2000 |
+
]
|
| 2001 |
+
}
|
| 2002 |
+
],
|
| 2003 |
+
"metadata": {
|
| 2004 |
+
"accelerator": "GPU",
|
| 2005 |
+
"colab": {
|
| 2006 |
+
"include_colab_link": true,
|
| 2007 |
+
"provenance": []
|
| 2008 |
+
},
|
| 2009 |
+
"gpuClass": "standard",
|
| 2010 |
+
"kernelspec": {
|
| 2011 |
+
"display_name": "Python 3 (ipykernel)",
|
| 2012 |
+
"language": "python",
|
| 2013 |
+
"name": "python3"
|
| 2014 |
+
},
|
| 2015 |
+
"language_info": {
|
| 2016 |
+
"codemirror_mode": {
|
| 2017 |
+
"name": "ipython",
|
| 2018 |
+
"version": 3
|
| 2019 |
+
},
|
| 2020 |
+
"file_extension": ".py",
|
| 2021 |
+
"mimetype": "text/x-python",
|
| 2022 |
+
"name": "python",
|
| 2023 |
+
"nbconvert_exporter": "python",
|
| 2024 |
+
"pygments_lexer": "ipython3",
|
| 2025 |
+
"version": "3.12.3"
|
| 2026 |
+
}
|
| 2027 |
+
},
|
| 2028 |
+
"nbformat": 4,
|
| 2029 |
+
"nbformat_minor": 1
|
| 2030 |
+
}
|
Downloads/.ipynb_checkpoints/Python Tutorial(1)-checkpoint.ipynb
ADDED
|
@@ -0,0 +1,3313 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {
|
| 6 |
+
"colab_type": "text",
|
| 7 |
+
"id": "view-in-github"
|
| 8 |
+
},
|
| 9 |
+
"source": [
|
| 10 |
+
"<a href=\"https://colab.research.google.com/github/AnjanDutta/EEEM068/blob/main/Notebooks/Python_Tutorial.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
|
| 11 |
+
]
|
| 12 |
+
},
|
| 13 |
+
{
|
| 14 |
+
"cell_type": "markdown",
|
| 15 |
+
"metadata": {
|
| 16 |
+
"id": "dzNng6vCL9eP"
|
| 17 |
+
},
|
| 18 |
+
"source": [
|
| 19 |
+
"<H1 style=\"text-align: center\">EEEM068 - Applied Machine Learning</H1>\n",
|
| 20 |
+
"<H1 style=\"text-align: center\">Workshop 01</H1>\n",
|
| 21 |
+
"<H1 style=\"text-align: center\">Python Tutorial</H1>\n"
|
| 22 |
+
]
|
| 23 |
+
},
|
| 24 |
+
{
|
| 25 |
+
"cell_type": "markdown",
|
| 26 |
+
"metadata": {
|
| 27 |
+
"id": "qVrTo-LhL9eS"
|
| 28 |
+
},
|
| 29 |
+
"source": [
|
| 30 |
+
"##Introduction"
|
| 31 |
+
]
|
| 32 |
+
},
|
| 33 |
+
{
|
| 34 |
+
"cell_type": "markdown",
|
| 35 |
+
"metadata": {
|
| 36 |
+
"id": "9t1gKp9PL9eV"
|
| 37 |
+
},
|
| 38 |
+
"source": [
|
| 39 |
+
"Python is a great general purpose programming language on its own, but with the help of a few popular libraries, such as numpy, scipy, matplotlib it becomes a powerful environment for scientific computing.\n",
|
| 40 |
+
"\n",
|
| 41 |
+
"We expect that many of you to have some experience with Python and numpy. Nevertheless, for the rest of you, this tutorial will serve as a quick crash course both on the Python programming language and on the use of Python for scientific computing.\n",
|
| 42 |
+
"\n",
|
| 43 |
+
"Some of you may have previous knowledge in Matlab, in which case we also recommend the [NumPy for Matlab users](https://numpy.org/doc/stable/user/numpy-for-matlab-users.html) page."
|
| 44 |
+
]
|
| 45 |
+
},
|
| 46 |
+
{
|
| 47 |
+
"cell_type": "markdown",
|
| 48 |
+
"metadata": {
|
| 49 |
+
"id": "U1PvreR9L9eW"
|
| 50 |
+
},
|
| 51 |
+
"source": [
|
| 52 |
+
"In this tutorial, we will cover:\n",
|
| 53 |
+
"\n",
|
| 54 |
+
"* Basic Python: Basic data types (Containers, Lists, Dictionaries, Sets, Tuples), Functions, Classes\n",
|
| 55 |
+
"* Numpy: Arrays, Array indexing, Datatypes, Array math, Broadcasting\n",
|
| 56 |
+
"* Matplotlib: Plotting, Subplots, Images\n",
|
| 57 |
+
"* Scikit-learn: Toy dataset, Classifier, Confusion matrix, Regressor\n",
|
| 58 |
+
"* OpenCV: Image, Image representation, Colour and datatype conversion, Image processing\n",
|
| 59 |
+
"* SciPy: I/O of MATLAB files, Distance functions"
|
| 60 |
+
]
|
| 61 |
+
},
|
| 62 |
+
{
|
| 63 |
+
"cell_type": "markdown",
|
| 64 |
+
"metadata": {
|
| 65 |
+
"id": "-O99OrwPtGii"
|
| 66 |
+
},
|
| 67 |
+
"source": [
|
| 68 |
+
"## Python Versions"
|
| 69 |
+
]
|
| 70 |
+
},
|
| 71 |
+
{
|
| 72 |
+
"cell_type": "markdown",
|
| 73 |
+
"metadata": {
|
| 74 |
+
"id": "nxvEkGXPM3Xh"
|
| 75 |
+
},
|
| 76 |
+
"source": [
|
| 77 |
+
"Please note that as of February 2026, Colab is using Python 3.12.12. Therefore, we will be using Python 3.12 for this iteration of the course. More details on Python 3.12 can be found in the [documentation](https://docs.python.org/3.12/tutorial/index.html). You can check your Python version at the command line by running `python --version`."
|
| 78 |
+
]
|
| 79 |
+
},
|
| 80 |
+
{
|
| 81 |
+
"cell_type": "code",
|
| 82 |
+
"execution_count": null,
|
| 83 |
+
"metadata": {
|
| 84 |
+
"id": "1L4Am0QATgOc"
|
| 85 |
+
},
|
| 86 |
+
"outputs": [],
|
| 87 |
+
"source": [
|
| 88 |
+
"!python --version"
|
| 89 |
+
]
|
| 90 |
+
},
|
| 91 |
+
{
|
| 92 |
+
"cell_type": "markdown",
|
| 93 |
+
"metadata": {
|
| 94 |
+
"id": "JAFKYgrpL9eY"
|
| 95 |
+
},
|
| 96 |
+
"source": [
|
| 97 |
+
"##Basics of Python"
|
| 98 |
+
]
|
| 99 |
+
},
|
| 100 |
+
{
|
| 101 |
+
"cell_type": "markdown",
|
| 102 |
+
"metadata": {
|
| 103 |
+
"id": "RbFS6tdgL9ea"
|
| 104 |
+
},
|
| 105 |
+
"source": [
|
| 106 |
+
"Python is an easy to learn, high-level, dynamically typed multiparadigm programming language. Python code is often said to be almost like pseudocode, since it allows you to express very powerful ideas in very few lines of code while being very readable. As an example, here is an implementation of the classic quicksort algorithm in Python:"
|
| 107 |
+
]
|
| 108 |
+
},
|
| 109 |
+
{
|
| 110 |
+
"cell_type": "code",
|
| 111 |
+
"execution_count": null,
|
| 112 |
+
"metadata": {
|
| 113 |
+
"id": "cYb0pjh1L9eb"
|
| 114 |
+
},
|
| 115 |
+
"outputs": [],
|
| 116 |
+
"source": [
|
| 117 |
+
"def quicksort(arr):\n",
|
| 118 |
+
" if len(arr) <= 1:\n",
|
| 119 |
+
" return arr\n",
|
| 120 |
+
" pivot = arr[len(arr) // 2]\n",
|
| 121 |
+
" left = [x for x in arr if x < pivot]\n",
|
| 122 |
+
" middle = [x for x in arr if x == pivot]\n",
|
| 123 |
+
" right = [x for x in arr if x > pivot]\n",
|
| 124 |
+
" return quicksort(left) + middle + quicksort(right)\n",
|
| 125 |
+
"\n",
|
| 126 |
+
"print(quicksort([3,6,8,10,1,2,1]))"
|
| 127 |
+
]
|
| 128 |
+
},
|
| 129 |
+
{
|
| 130 |
+
"cell_type": "markdown",
|
| 131 |
+
"metadata": {
|
| 132 |
+
"id": "NwS_hu4xL9eo"
|
| 133 |
+
},
|
| 134 |
+
"source": [
|
| 135 |
+
"###Basic data types"
|
| 136 |
+
]
|
| 137 |
+
},
|
| 138 |
+
{
|
| 139 |
+
"cell_type": "markdown",
|
| 140 |
+
"metadata": {
|
| 141 |
+
"id": "DL5sMSZ9L9eq"
|
| 142 |
+
},
|
| 143 |
+
"source": [
|
| 144 |
+
"####Numbers"
|
| 145 |
+
]
|
| 146 |
+
},
|
| 147 |
+
{
|
| 148 |
+
"cell_type": "markdown",
|
| 149 |
+
"metadata": {
|
| 150 |
+
"id": "MGS0XEWoL9er"
|
| 151 |
+
},
|
| 152 |
+
"source": [
|
| 153 |
+
"Integers and floats work as you would expect from other languages:"
|
| 154 |
+
]
|
| 155 |
+
},
|
| 156 |
+
{
|
| 157 |
+
"cell_type": "code",
|
| 158 |
+
"execution_count": null,
|
| 159 |
+
"metadata": {
|
| 160 |
+
"id": "KheDr_zDL9es"
|
| 161 |
+
},
|
| 162 |
+
"outputs": [],
|
| 163 |
+
"source": [
|
| 164 |
+
"x = 3\n",
|
| 165 |
+
"print(x, type(x))"
|
| 166 |
+
]
|
| 167 |
+
},
|
| 168 |
+
{
|
| 169 |
+
"cell_type": "code",
|
| 170 |
+
"execution_count": null,
|
| 171 |
+
"metadata": {
|
| 172 |
+
"id": "sk_8DFcuL9ey"
|
| 173 |
+
},
|
| 174 |
+
"outputs": [],
|
| 175 |
+
"source": [
|
| 176 |
+
"print(x + 1) # Addition\n",
|
| 177 |
+
"print(x - 1) # Subtraction\n",
|
| 178 |
+
"print(x * 2) # Multiplication\n",
|
| 179 |
+
"print(x ** 2) # Exponentiation"
|
| 180 |
+
]
|
| 181 |
+
},
|
| 182 |
+
{
|
| 183 |
+
"cell_type": "code",
|
| 184 |
+
"execution_count": null,
|
| 185 |
+
"metadata": {
|
| 186 |
+
"id": "U4Jl8K0tL9e4"
|
| 187 |
+
},
|
| 188 |
+
"outputs": [],
|
| 189 |
+
"source": [
|
| 190 |
+
"x += 1\n",
|
| 191 |
+
"print(x)\n",
|
| 192 |
+
"x *= 2\n",
|
| 193 |
+
"print(x)"
|
| 194 |
+
]
|
| 195 |
+
},
|
| 196 |
+
{
|
| 197 |
+
"cell_type": "code",
|
| 198 |
+
"execution_count": null,
|
| 199 |
+
"metadata": {
|
| 200 |
+
"id": "w-nZ0Sg_L9e9"
|
| 201 |
+
},
|
| 202 |
+
"outputs": [],
|
| 203 |
+
"source": [
|
| 204 |
+
"y = 2.5\n",
|
| 205 |
+
"print(type(y))\n",
|
| 206 |
+
"print(y, y + 1, y * 2, y ** 2)"
|
| 207 |
+
]
|
| 208 |
+
},
|
| 209 |
+
{
|
| 210 |
+
"cell_type": "markdown",
|
| 211 |
+
"metadata": {
|
| 212 |
+
"id": "r2A9ApyaL9fB"
|
| 213 |
+
},
|
| 214 |
+
"source": [
|
| 215 |
+
"Note that unlike many languages (such as C and C++) Python does not have unary increment (x++) or decrement (x--) operators.\n",
|
| 216 |
+
"\n",
|
| 217 |
+
"Python also has built-in types for long integers and complex numbers; you can find all of the details in the [documentation](https://docs.python.org/3.8/library/stdtypes.html#numeric-types-int-float-long-complex)."
|
| 218 |
+
]
|
| 219 |
+
},
|
| 220 |
+
{
|
| 221 |
+
"cell_type": "markdown",
|
| 222 |
+
"metadata": {
|
| 223 |
+
"id": "EqRS7qhBL9fC"
|
| 224 |
+
},
|
| 225 |
+
"source": [
|
| 226 |
+
"####Booleans"
|
| 227 |
+
]
|
| 228 |
+
},
|
| 229 |
+
{
|
| 230 |
+
"cell_type": "markdown",
|
| 231 |
+
"metadata": {
|
| 232 |
+
"id": "Nv_LIVOJL9fD"
|
| 233 |
+
},
|
| 234 |
+
"source": [
|
| 235 |
+
"Python implements all of the usual operators for Boolean logic, but uses English words rather than symbols (`&&`, `||`, etc.):"
|
| 236 |
+
]
|
| 237 |
+
},
|
| 238 |
+
{
|
| 239 |
+
"cell_type": "code",
|
| 240 |
+
"execution_count": null,
|
| 241 |
+
"metadata": {
|
| 242 |
+
"id": "RvoImwgGL9fE"
|
| 243 |
+
},
|
| 244 |
+
"outputs": [],
|
| 245 |
+
"source": [
|
| 246 |
+
"t, f = True, False\n",
|
| 247 |
+
"print(type(t))"
|
| 248 |
+
]
|
| 249 |
+
},
|
| 250 |
+
{
|
| 251 |
+
"cell_type": "markdown",
|
| 252 |
+
"metadata": {
|
| 253 |
+
"id": "YQgmQfOgL9fI"
|
| 254 |
+
},
|
| 255 |
+
"source": [
|
| 256 |
+
"Now we let's look at the operations:"
|
| 257 |
+
]
|
| 258 |
+
},
|
| 259 |
+
{
|
| 260 |
+
"cell_type": "code",
|
| 261 |
+
"execution_count": null,
|
| 262 |
+
"metadata": {
|
| 263 |
+
"id": "6zYm7WzCL9fK"
|
| 264 |
+
},
|
| 265 |
+
"outputs": [],
|
| 266 |
+
"source": [
|
| 267 |
+
"print(t and f) # Logical AND;\n",
|
| 268 |
+
"print(t or f) # Logical OR;\n",
|
| 269 |
+
"print(not t) # Logical NOT;\n",
|
| 270 |
+
"print(t != f) # Logical XOR;"
|
| 271 |
+
]
|
| 272 |
+
},
|
| 273 |
+
{
|
| 274 |
+
"cell_type": "markdown",
|
| 275 |
+
"metadata": {
|
| 276 |
+
"id": "UQnQWFEyL9fP"
|
| 277 |
+
},
|
| 278 |
+
"source": [
|
| 279 |
+
"####Strings"
|
| 280 |
+
]
|
| 281 |
+
},
|
| 282 |
+
{
|
| 283 |
+
"cell_type": "code",
|
| 284 |
+
"execution_count": null,
|
| 285 |
+
"metadata": {
|
| 286 |
+
"id": "AijEDtPFL9fP"
|
| 287 |
+
},
|
| 288 |
+
"outputs": [],
|
| 289 |
+
"source": [
|
| 290 |
+
"hello = 'hello' # String literals can use single quotes\n",
|
| 291 |
+
"world = \"world\" # or double quotes; it does not matter\n",
|
| 292 |
+
"print(hello, len(hello))"
|
| 293 |
+
]
|
| 294 |
+
},
|
| 295 |
+
{
|
| 296 |
+
"cell_type": "code",
|
| 297 |
+
"execution_count": null,
|
| 298 |
+
"metadata": {
|
| 299 |
+
"id": "saDeaA7hL9fT"
|
| 300 |
+
},
|
| 301 |
+
"outputs": [],
|
| 302 |
+
"source": [
|
| 303 |
+
"hw = hello + ' ' + world # String concatenation\n",
|
| 304 |
+
"print(hw)"
|
| 305 |
+
]
|
| 306 |
+
},
|
| 307 |
+
{
|
| 308 |
+
"cell_type": "code",
|
| 309 |
+
"execution_count": null,
|
| 310 |
+
"metadata": {
|
| 311 |
+
"id": "Nji1_UjYL9fY"
|
| 312 |
+
},
|
| 313 |
+
"outputs": [],
|
| 314 |
+
"source": [
|
| 315 |
+
"hw12 = '{} {} {}'.format(hello, world, 12) # string formatting\n",
|
| 316 |
+
"print(hw12)"
|
| 317 |
+
]
|
| 318 |
+
},
|
| 319 |
+
{
|
| 320 |
+
"cell_type": "markdown",
|
| 321 |
+
"metadata": {
|
| 322 |
+
"id": "bUpl35bIL9fc"
|
| 323 |
+
},
|
| 324 |
+
"source": [
|
| 325 |
+
"String objects have a bunch of useful methods; for example:"
|
| 326 |
+
]
|
| 327 |
+
},
|
| 328 |
+
{
|
| 329 |
+
"cell_type": "code",
|
| 330 |
+
"execution_count": null,
|
| 331 |
+
"metadata": {
|
| 332 |
+
"id": "VOxGatlsL9fd"
|
| 333 |
+
},
|
| 334 |
+
"outputs": [],
|
| 335 |
+
"source": [
|
| 336 |
+
"s = \"hello\"\n",
|
| 337 |
+
"print(s.capitalize()) # Capitalize a string\n",
|
| 338 |
+
"print(s.upper()) # Convert a string to uppercase; prints \"HELLO\"\n",
|
| 339 |
+
"print(s.rjust(7)) # Right-justify a string, padding with spaces\n",
|
| 340 |
+
"print(s.center(7)) # Center a string, padding with spaces\n",
|
| 341 |
+
"print(s.replace('l', '(ell)')) # Replace all instances of one substring with another\n",
|
| 342 |
+
"print(' world '.strip()) # Strip leading and trailing whitespace"
|
| 343 |
+
]
|
| 344 |
+
},
|
| 345 |
+
{
|
| 346 |
+
"cell_type": "markdown",
|
| 347 |
+
"metadata": {
|
| 348 |
+
"id": "06cayXLtL9fi"
|
| 349 |
+
},
|
| 350 |
+
"source": [
|
| 351 |
+
"You can find a list of all string methods in the [documentation](https://docs.python.org/3.7/library/stdtypes.html#string-methods)."
|
| 352 |
+
]
|
| 353 |
+
},
|
| 354 |
+
{
|
| 355 |
+
"cell_type": "markdown",
|
| 356 |
+
"metadata": {
|
| 357 |
+
"id": "p-6hClFjL9fk"
|
| 358 |
+
},
|
| 359 |
+
"source": [
|
| 360 |
+
"###Containers"
|
| 361 |
+
]
|
| 362 |
+
},
|
| 363 |
+
{
|
| 364 |
+
"cell_type": "markdown",
|
| 365 |
+
"metadata": {
|
| 366 |
+
"id": "FD9H18eQL9fk"
|
| 367 |
+
},
|
| 368 |
+
"source": [
|
| 369 |
+
"Python includes several built-in container types: lists, dictionaries, sets, and tuples."
|
| 370 |
+
]
|
| 371 |
+
},
|
| 372 |
+
{
|
| 373 |
+
"cell_type": "markdown",
|
| 374 |
+
"metadata": {
|
| 375 |
+
"id": "UsIWOe0LL9fn"
|
| 376 |
+
},
|
| 377 |
+
"source": [
|
| 378 |
+
"####Lists"
|
| 379 |
+
]
|
| 380 |
+
},
|
| 381 |
+
{
|
| 382 |
+
"cell_type": "markdown",
|
| 383 |
+
"metadata": {
|
| 384 |
+
"id": "wzxX7rgWL9fn"
|
| 385 |
+
},
|
| 386 |
+
"source": [
|
| 387 |
+
"A list is the Python equivalent of an array, but is resizeable and can contain elements of different types:"
|
| 388 |
+
]
|
| 389 |
+
},
|
| 390 |
+
{
|
| 391 |
+
"cell_type": "code",
|
| 392 |
+
"execution_count": null,
|
| 393 |
+
"metadata": {
|
| 394 |
+
"id": "hk3A8pPcL9fp"
|
| 395 |
+
},
|
| 396 |
+
"outputs": [],
|
| 397 |
+
"source": [
|
| 398 |
+
"xs = [3, 1, 2] # Create a list\n",
|
| 399 |
+
"print(xs, xs[2])\n",
|
| 400 |
+
"print(xs[-1]) # Negative indices count from the end of the list; prints \"2\""
|
| 401 |
+
]
|
| 402 |
+
},
|
| 403 |
+
{
|
| 404 |
+
"cell_type": "code",
|
| 405 |
+
"execution_count": null,
|
| 406 |
+
"metadata": {
|
| 407 |
+
"id": "YCjCy_0_L9ft"
|
| 408 |
+
},
|
| 409 |
+
"outputs": [],
|
| 410 |
+
"source": [
|
| 411 |
+
"xs[2] = 'foo' # Lists can be heterogeneous, i.e. it can contain elements of different types\n",
|
| 412 |
+
"print(xs)"
|
| 413 |
+
]
|
| 414 |
+
},
|
| 415 |
+
{
|
| 416 |
+
"cell_type": "code",
|
| 417 |
+
"execution_count": null,
|
| 418 |
+
"metadata": {
|
| 419 |
+
"id": "vJ0x5cF-L9fx"
|
| 420 |
+
},
|
| 421 |
+
"outputs": [],
|
| 422 |
+
"source": [
|
| 423 |
+
"xs.append('bar') # Add a new element to the end of the list\n",
|
| 424 |
+
"print(xs)"
|
| 425 |
+
]
|
| 426 |
+
},
|
| 427 |
+
{
|
| 428 |
+
"cell_type": "code",
|
| 429 |
+
"execution_count": null,
|
| 430 |
+
"metadata": {
|
| 431 |
+
"id": "cxVCNRTNL9f1"
|
| 432 |
+
},
|
| 433 |
+
"outputs": [],
|
| 434 |
+
"source": [
|
| 435 |
+
"x = xs.pop() # Remove and return the last element of the list\n",
|
| 436 |
+
"print(x, xs)"
|
| 437 |
+
]
|
| 438 |
+
},
|
| 439 |
+
{
|
| 440 |
+
"cell_type": "markdown",
|
| 441 |
+
"metadata": {
|
| 442 |
+
"id": "ilyoyO34L9f4"
|
| 443 |
+
},
|
| 444 |
+
"source": [
|
| 445 |
+
"As usual, you can find all the gory details about lists in the [documentation](https://docs.python.org/3.7/tutorial/datastructures.html#more-on-lists)."
|
| 446 |
+
]
|
| 447 |
+
},
|
| 448 |
+
{
|
| 449 |
+
"cell_type": "markdown",
|
| 450 |
+
"metadata": {
|
| 451 |
+
"id": "ovahhxd_L9f5"
|
| 452 |
+
},
|
| 453 |
+
"source": [
|
| 454 |
+
"####Slicing"
|
| 455 |
+
]
|
| 456 |
+
},
|
| 457 |
+
{
|
| 458 |
+
"cell_type": "markdown",
|
| 459 |
+
"metadata": {
|
| 460 |
+
"id": "YeSYKhv9L9f6"
|
| 461 |
+
},
|
| 462 |
+
"source": [
|
| 463 |
+
"In addition to accessing list elements one at a time, Python provides concise syntax to access sublists; this is known as slicing:"
|
| 464 |
+
]
|
| 465 |
+
},
|
| 466 |
+
{
|
| 467 |
+
"cell_type": "code",
|
| 468 |
+
"execution_count": null,
|
| 469 |
+
"metadata": {
|
| 470 |
+
"id": "ninq666bL9f6"
|
| 471 |
+
},
|
| 472 |
+
"outputs": [],
|
| 473 |
+
"source": [
|
| 474 |
+
"nums = list(range(5)) # range is a built-in function that creates a list of integers\n",
|
| 475 |
+
"print(nums) # Prints \"[0, 1, 2, 3, 4]\"\n",
|
| 476 |
+
"print(nums[2:4]) # Get a slice from index 2 to 4 (exclusive); prints \"[2, 3]\"\n",
|
| 477 |
+
"print(nums[2:]) # Get a slice from index 2 to the end; prints \"[2, 3, 4]\"\n",
|
| 478 |
+
"print(nums[:2]) # Get a slice from the start to index 2 (exclusive); prints \"[0, 1]\"\n",
|
| 479 |
+
"print(nums[:]) # Get a slice of the whole list; prints [\"0, 1, 2, 3, 4]\"\n",
|
| 480 |
+
"print(nums[:-1]) # Slice indices can be negative; prints [\"0, 1, 2, 3]\"\n",
|
| 481 |
+
"nums[2:4] = [8, 9] # Assign a new sublist to a slice\n",
|
| 482 |
+
"print(nums) # Prints \"[0, 1, 8, 9, 4]\""
|
| 483 |
+
]
|
| 484 |
+
},
|
| 485 |
+
{
|
| 486 |
+
"cell_type": "markdown",
|
| 487 |
+
"metadata": {
|
| 488 |
+
"id": "arrLCcMyL9gK"
|
| 489 |
+
},
|
| 490 |
+
"source": [
|
| 491 |
+
"####List comprehensions:"
|
| 492 |
+
]
|
| 493 |
+
},
|
| 494 |
+
{
|
| 495 |
+
"cell_type": "markdown",
|
| 496 |
+
"metadata": {
|
| 497 |
+
"id": "5Qn2jU_pL9gL"
|
| 498 |
+
},
|
| 499 |
+
"source": [
|
| 500 |
+
"When programming, frequently we want to transform one type of data into another. As a simple example, consider the following code that computes square numbers:"
|
| 501 |
+
]
|
| 502 |
+
},
|
| 503 |
+
{
|
| 504 |
+
"cell_type": "code",
|
| 505 |
+
"execution_count": null,
|
| 506 |
+
"metadata": {
|
| 507 |
+
"id": "IVNEwoMXL9gL"
|
| 508 |
+
},
|
| 509 |
+
"outputs": [],
|
| 510 |
+
"source": [
|
| 511 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 512 |
+
"squares = []\n",
|
| 513 |
+
"for x in nums:\n",
|
| 514 |
+
" squares.append(x ** 2)\n",
|
| 515 |
+
"print(squares)"
|
| 516 |
+
]
|
| 517 |
+
},
|
| 518 |
+
{
|
| 519 |
+
"cell_type": "markdown",
|
| 520 |
+
"metadata": {
|
| 521 |
+
"id": "7DmKVUFaL9gQ"
|
| 522 |
+
},
|
| 523 |
+
"source": [
|
| 524 |
+
"You can make this code simpler using a list comprehension:"
|
| 525 |
+
]
|
| 526 |
+
},
|
| 527 |
+
{
|
| 528 |
+
"cell_type": "code",
|
| 529 |
+
"execution_count": null,
|
| 530 |
+
"metadata": {
|
| 531 |
+
"id": "kZxsUfV6L9gR"
|
| 532 |
+
},
|
| 533 |
+
"outputs": [],
|
| 534 |
+
"source": [
|
| 535 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 536 |
+
"squares = [x ** 2 for x in nums]\n",
|
| 537 |
+
"print(squares)"
|
| 538 |
+
]
|
| 539 |
+
},
|
| 540 |
+
{
|
| 541 |
+
"cell_type": "markdown",
|
| 542 |
+
"metadata": {
|
| 543 |
+
"id": "-D8ARK7tL9gV"
|
| 544 |
+
},
|
| 545 |
+
"source": [
|
| 546 |
+
"List comprehensions can also contain conditions:"
|
| 547 |
+
]
|
| 548 |
+
},
|
| 549 |
+
{
|
| 550 |
+
"cell_type": "code",
|
| 551 |
+
"execution_count": null,
|
| 552 |
+
"metadata": {
|
| 553 |
+
"id": "yUtgOyyYL9gV"
|
| 554 |
+
},
|
| 555 |
+
"outputs": [],
|
| 556 |
+
"source": [
|
| 557 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 558 |
+
"even_squares = [x ** 2 for x in nums if x % 2 == 0]\n",
|
| 559 |
+
"print(even_squares)"
|
| 560 |
+
]
|
| 561 |
+
},
|
| 562 |
+
{
|
| 563 |
+
"cell_type": "markdown",
|
| 564 |
+
"metadata": {
|
| 565 |
+
"id": "H8xsUEFpL9gZ"
|
| 566 |
+
},
|
| 567 |
+
"source": [
|
| 568 |
+
"####Dictionaries"
|
| 569 |
+
]
|
| 570 |
+
},
|
| 571 |
+
{
|
| 572 |
+
"cell_type": "markdown",
|
| 573 |
+
"metadata": {
|
| 574 |
+
"id": "kkjAGMAJL9ga"
|
| 575 |
+
},
|
| 576 |
+
"source": [
|
| 577 |
+
"A dictionary stores (key, value) pairs, similar to a `Map` in Java or an object in Javascript. You can use it like this:"
|
| 578 |
+
]
|
| 579 |
+
},
|
| 580 |
+
{
|
| 581 |
+
"cell_type": "code",
|
| 582 |
+
"execution_count": null,
|
| 583 |
+
"metadata": {
|
| 584 |
+
"id": "XBYI1MrYL9gb"
|
| 585 |
+
},
|
| 586 |
+
"outputs": [],
|
| 587 |
+
"source": [
|
| 588 |
+
"d = {'cat': 'cute', 'dog': 'furry'} # Create a new dictionary with some data\n",
|
| 589 |
+
"print(d['cat']) # Get an entry from a dictionary; prints \"cute\"\n",
|
| 590 |
+
"print('cat' in d) # Check if a dictionary has a given key; prints \"True\""
|
| 591 |
+
]
|
| 592 |
+
},
|
| 593 |
+
{
|
| 594 |
+
"cell_type": "code",
|
| 595 |
+
"execution_count": null,
|
| 596 |
+
"metadata": {
|
| 597 |
+
"id": "pS7e-G-HL9gf"
|
| 598 |
+
},
|
| 599 |
+
"outputs": [],
|
| 600 |
+
"source": [
|
| 601 |
+
"d['fish'] = 'wet' # Set an entry in a dictionary\n",
|
| 602 |
+
"print(d['fish']) # Prints \"wet\""
|
| 603 |
+
]
|
| 604 |
+
},
|
| 605 |
+
{
|
| 606 |
+
"cell_type": "code",
|
| 607 |
+
"execution_count": null,
|
| 608 |
+
"metadata": {
|
| 609 |
+
"id": "tFY065ItL9gi"
|
| 610 |
+
},
|
| 611 |
+
"outputs": [],
|
| 612 |
+
"source": [
|
| 613 |
+
"print(d['monkey']) # KeyError: 'monkey' not a key of d"
|
| 614 |
+
]
|
| 615 |
+
},
|
| 616 |
+
{
|
| 617 |
+
"cell_type": "code",
|
| 618 |
+
"execution_count": null,
|
| 619 |
+
"metadata": {
|
| 620 |
+
"id": "8TjbEWqML9gl"
|
| 621 |
+
},
|
| 622 |
+
"outputs": [],
|
| 623 |
+
"source": [
|
| 624 |
+
"print(d.get('monkey', 'N/A')) # Get an element with a default; prints \"N/A\"\n",
|
| 625 |
+
"print(d.get('fish', 'N/A')) # Get an element with a default; prints \"wet\""
|
| 626 |
+
]
|
| 627 |
+
},
|
| 628 |
+
{
|
| 629 |
+
"cell_type": "code",
|
| 630 |
+
"execution_count": null,
|
| 631 |
+
"metadata": {
|
| 632 |
+
"id": "0EItdNBJL9go"
|
| 633 |
+
},
|
| 634 |
+
"outputs": [],
|
| 635 |
+
"source": [
|
| 636 |
+
"del d['fish'] # Remove an element from a dictionary\n",
|
| 637 |
+
"print(d.get('fish', 'N/A')) # \"fish\" is no longer a key; prints \"N/A\""
|
| 638 |
+
]
|
| 639 |
+
},
|
| 640 |
+
{
|
| 641 |
+
"cell_type": "markdown",
|
| 642 |
+
"metadata": {
|
| 643 |
+
"id": "wqm4dRZNL9gr"
|
| 644 |
+
},
|
| 645 |
+
"source": [
|
| 646 |
+
"You can find all you need to know about dictionaries in the [documentation](https://docs.python.org/2/library/stdtypes.html#dict)."
|
| 647 |
+
]
|
| 648 |
+
},
|
| 649 |
+
{
|
| 650 |
+
"cell_type": "markdown",
|
| 651 |
+
"metadata": {
|
| 652 |
+
"id": "IxwEqHlGL9gr"
|
| 653 |
+
},
|
| 654 |
+
"source": [
|
| 655 |
+
"It is easy to iterate over the keys in a dictionary:"
|
| 656 |
+
]
|
| 657 |
+
},
|
| 658 |
+
{
|
| 659 |
+
"cell_type": "code",
|
| 660 |
+
"execution_count": null,
|
| 661 |
+
"metadata": {
|
| 662 |
+
"id": "rYfz7ZKNL9gs"
|
| 663 |
+
},
|
| 664 |
+
"outputs": [],
|
| 665 |
+
"source": [
|
| 666 |
+
"d = {'person': 2, 'cat': 4, 'spider': 8}\n",
|
| 667 |
+
"for animal, legs in d.items():\n",
|
| 668 |
+
" print('A {} has {} legs'.format(animal, legs))"
|
| 669 |
+
]
|
| 670 |
+
},
|
| 671 |
+
{
|
| 672 |
+
"cell_type": "markdown",
|
| 673 |
+
"metadata": {
|
| 674 |
+
"id": "17sxiOpzL9gz"
|
| 675 |
+
},
|
| 676 |
+
"source": [
|
| 677 |
+
"Dictionary comprehensions: These are similar to list comprehensions, but allow you to easily construct dictionaries. For example:"
|
| 678 |
+
]
|
| 679 |
+
},
|
| 680 |
+
{
|
| 681 |
+
"cell_type": "code",
|
| 682 |
+
"execution_count": null,
|
| 683 |
+
"metadata": {
|
| 684 |
+
"id": "8PB07imLL9gz"
|
| 685 |
+
},
|
| 686 |
+
"outputs": [],
|
| 687 |
+
"source": [
|
| 688 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 689 |
+
"even_num_to_square = {x: x ** 2 for x in nums if x % 2 == 0}\n",
|
| 690 |
+
"print(even_num_to_square)"
|
| 691 |
+
]
|
| 692 |
+
},
|
| 693 |
+
{
|
| 694 |
+
"cell_type": "markdown",
|
| 695 |
+
"metadata": {
|
| 696 |
+
"id": "V9MHfUdvL9g2"
|
| 697 |
+
},
|
| 698 |
+
"source": [
|
| 699 |
+
"####Sets"
|
| 700 |
+
]
|
| 701 |
+
},
|
| 702 |
+
{
|
| 703 |
+
"cell_type": "markdown",
|
| 704 |
+
"metadata": {
|
| 705 |
+
"id": "Rpm4UtNpL9g2"
|
| 706 |
+
},
|
| 707 |
+
"source": [
|
| 708 |
+
"A set is an unordered collection of distinct elements. As a simple example, consider the following:"
|
| 709 |
+
]
|
| 710 |
+
},
|
| 711 |
+
{
|
| 712 |
+
"cell_type": "code",
|
| 713 |
+
"execution_count": null,
|
| 714 |
+
"metadata": {
|
| 715 |
+
"id": "MmyaniLsL9g2"
|
| 716 |
+
},
|
| 717 |
+
"outputs": [],
|
| 718 |
+
"source": [
|
| 719 |
+
"animals = {'cat', 'dog'}\n",
|
| 720 |
+
"print('cat' in animals) # Check if an element is in a set; prints \"True\"\n",
|
| 721 |
+
"print('fish' in animals) # prints \"False\"\n"
|
| 722 |
+
]
|
| 723 |
+
},
|
| 724 |
+
{
|
| 725 |
+
"cell_type": "code",
|
| 726 |
+
"execution_count": null,
|
| 727 |
+
"metadata": {
|
| 728 |
+
"id": "ElJEyK86L9g6"
|
| 729 |
+
},
|
| 730 |
+
"outputs": [],
|
| 731 |
+
"source": [
|
| 732 |
+
"animals.add('fish') # Add an element to a set\n",
|
| 733 |
+
"print('fish' in animals)\n",
|
| 734 |
+
"print(len(animals)) # Number of elements in a set;"
|
| 735 |
+
]
|
| 736 |
+
},
|
| 737 |
+
{
|
| 738 |
+
"cell_type": "code",
|
| 739 |
+
"execution_count": null,
|
| 740 |
+
"metadata": {
|
| 741 |
+
"id": "5uGmrxdPL9g9"
|
| 742 |
+
},
|
| 743 |
+
"outputs": [],
|
| 744 |
+
"source": [
|
| 745 |
+
"animals.add('cat') # Adding an element that is already in the set does nothing\n",
|
| 746 |
+
"print(len(animals))\n",
|
| 747 |
+
"animals.remove('cat') # Remove an element from a set\n",
|
| 748 |
+
"print(len(animals))"
|
| 749 |
+
]
|
| 750 |
+
},
|
| 751 |
+
{
|
| 752 |
+
"cell_type": "markdown",
|
| 753 |
+
"metadata": {
|
| 754 |
+
"id": "zk2DbvLKL9g_"
|
| 755 |
+
},
|
| 756 |
+
"source": [
|
| 757 |
+
"_Loops_: Iterating over a set has the same syntax as iterating over a list; however since sets are unordered, you cannot make assumptions about the order in which you visit the elements of the set:"
|
| 758 |
+
]
|
| 759 |
+
},
|
| 760 |
+
{
|
| 761 |
+
"cell_type": "code",
|
| 762 |
+
"execution_count": null,
|
| 763 |
+
"metadata": {
|
| 764 |
+
"id": "K47KYNGyL9hA"
|
| 765 |
+
},
|
| 766 |
+
"outputs": [],
|
| 767 |
+
"source": [
|
| 768 |
+
"animals = {'cat', 'dog', 'fish'}\n",
|
| 769 |
+
"for idx, animal in enumerate(animals):\n",
|
| 770 |
+
" print('#{}: {}'.format(idx + 1, animal))"
|
| 771 |
+
]
|
| 772 |
+
},
|
| 773 |
+
{
|
| 774 |
+
"cell_type": "markdown",
|
| 775 |
+
"metadata": {
|
| 776 |
+
"id": "puq4S8buL9hC"
|
| 777 |
+
},
|
| 778 |
+
"source": [
|
| 779 |
+
"Set comprehensions: Like lists and dictionaries, we can easily construct sets using set comprehensions:"
|
| 780 |
+
]
|
| 781 |
+
},
|
| 782 |
+
{
|
| 783 |
+
"cell_type": "code",
|
| 784 |
+
"execution_count": null,
|
| 785 |
+
"metadata": {
|
| 786 |
+
"id": "iw7k90k3L9hC"
|
| 787 |
+
},
|
| 788 |
+
"outputs": [],
|
| 789 |
+
"source": [
|
| 790 |
+
"from math import sqrt\n",
|
| 791 |
+
"print({int(sqrt(x)) for x in range(30)})"
|
| 792 |
+
]
|
| 793 |
+
},
|
| 794 |
+
{
|
| 795 |
+
"cell_type": "markdown",
|
| 796 |
+
"metadata": {
|
| 797 |
+
"id": "qPsHSKB1L9hF"
|
| 798 |
+
},
|
| 799 |
+
"source": [
|
| 800 |
+
"####Tuples"
|
| 801 |
+
]
|
| 802 |
+
},
|
| 803 |
+
{
|
| 804 |
+
"cell_type": "markdown",
|
| 805 |
+
"metadata": {
|
| 806 |
+
"id": "kucc0LKVL9hG"
|
| 807 |
+
},
|
| 808 |
+
"source": [
|
| 809 |
+
"A tuple is an (immutable) ordered list of values. A tuple is in many ways similar to a list; one of the most important differences is that tuples can be used as keys in dictionaries and as elements of sets, while lists cannot. Here is a trivial example:"
|
| 810 |
+
]
|
| 811 |
+
},
|
| 812 |
+
{
|
| 813 |
+
"cell_type": "code",
|
| 814 |
+
"execution_count": null,
|
| 815 |
+
"metadata": {
|
| 816 |
+
"id": "9wHUyTKxL9hH"
|
| 817 |
+
},
|
| 818 |
+
"outputs": [],
|
| 819 |
+
"source": [
|
| 820 |
+
"d = {(x, x + 1): x for x in range(10)} # Create a dictionary with tuple keys\n",
|
| 821 |
+
"t = (5, 6) # Create a tuple\n",
|
| 822 |
+
"print(type(t))\n",
|
| 823 |
+
"print(d[t])\n",
|
| 824 |
+
"print(d[(1, 2)])"
|
| 825 |
+
]
|
| 826 |
+
},
|
| 827 |
+
{
|
| 828 |
+
"cell_type": "markdown",
|
| 829 |
+
"metadata": {
|
| 830 |
+
"id": "iFON3Tm0CfIg"
|
| 831 |
+
},
|
| 832 |
+
"source": [
|
| 833 |
+
"##Loops"
|
| 834 |
+
]
|
| 835 |
+
},
|
| 836 |
+
{
|
| 837 |
+
"cell_type": "markdown",
|
| 838 |
+
"metadata": {
|
| 839 |
+
"id": "7aXUmwT69XqK"
|
| 840 |
+
},
|
| 841 |
+
"source": [
|
| 842 |
+
"###`for` loop"
|
| 843 |
+
]
|
| 844 |
+
},
|
| 845 |
+
{
|
| 846 |
+
"cell_type": "markdown",
|
| 847 |
+
"metadata": {
|
| 848 |
+
"id": "_DYz1j6QL9f_"
|
| 849 |
+
},
|
| 850 |
+
"source": [
|
| 851 |
+
"You can loop over the elements of a list like this:"
|
| 852 |
+
]
|
| 853 |
+
},
|
| 854 |
+
{
|
| 855 |
+
"cell_type": "code",
|
| 856 |
+
"execution_count": null,
|
| 857 |
+
"metadata": {
|
| 858 |
+
"id": "4cCOysfWL9gA"
|
| 859 |
+
},
|
| 860 |
+
"outputs": [],
|
| 861 |
+
"source": [
|
| 862 |
+
"animals = ['cat', 'dog', 'monkey']\n",
|
| 863 |
+
"for animal in animals:\n",
|
| 864 |
+
" print(animal)"
|
| 865 |
+
]
|
| 866 |
+
},
|
| 867 |
+
{
|
| 868 |
+
"cell_type": "markdown",
|
| 869 |
+
"metadata": {
|
| 870 |
+
"id": "KxIaQs7pL9gE"
|
| 871 |
+
},
|
| 872 |
+
"source": [
|
| 873 |
+
"If you want access to the index of each element within the body of a loop, use the built-in `enumerate` function:"
|
| 874 |
+
]
|
| 875 |
+
},
|
| 876 |
+
{
|
| 877 |
+
"cell_type": "code",
|
| 878 |
+
"execution_count": null,
|
| 879 |
+
"metadata": {
|
| 880 |
+
"id": "JjGnDluWL9gF"
|
| 881 |
+
},
|
| 882 |
+
"outputs": [],
|
| 883 |
+
"source": [
|
| 884 |
+
"animals = ['cat', 'dog', 'monkey']\n",
|
| 885 |
+
"for idx, animal in enumerate(animals):\n",
|
| 886 |
+
" print('#{}: {}'.format(idx + 1, animal))"
|
| 887 |
+
]
|
| 888 |
+
},
|
| 889 |
+
{
|
| 890 |
+
"cell_type": "markdown",
|
| 891 |
+
"metadata": {
|
| 892 |
+
"id": "Tlf5gPRy9jfV"
|
| 893 |
+
},
|
| 894 |
+
"source": [
|
| 895 |
+
"###`range()` function"
|
| 896 |
+
]
|
| 897 |
+
},
|
| 898 |
+
{
|
| 899 |
+
"cell_type": "markdown",
|
| 900 |
+
"metadata": {
|
| 901 |
+
"id": "pzgK4H6j-PyN"
|
| 902 |
+
},
|
| 903 |
+
"source": [
|
| 904 |
+
"If you need to iterate over a sequence of numbers, the built-in function `range()` comes in handy. It generates arithmetic progressions:"
|
| 905 |
+
]
|
| 906 |
+
},
|
| 907 |
+
{
|
| 908 |
+
"cell_type": "code",
|
| 909 |
+
"execution_count": null,
|
| 910 |
+
"metadata": {
|
| 911 |
+
"id": "eNlqr8sn-TX5"
|
| 912 |
+
},
|
| 913 |
+
"outputs": [],
|
| 914 |
+
"source": [
|
| 915 |
+
"for i in range(5):\n",
|
| 916 |
+
" print(i)"
|
| 917 |
+
]
|
| 918 |
+
},
|
| 919 |
+
{
|
| 920 |
+
"cell_type": "markdown",
|
| 921 |
+
"metadata": {
|
| 922 |
+
"id": "vfigt0cE-fpp"
|
| 923 |
+
},
|
| 924 |
+
"source": [
|
| 925 |
+
"The given end point is never part of the generated sequence; `range(10)` generates 10 values, the legal indices for items of a sequence of length 10. It is possible to let the range start at another number, or to specify a different increment (even negative; sometimes this is called the ‘step’):"
|
| 926 |
+
]
|
| 927 |
+
},
|
| 928 |
+
{
|
| 929 |
+
"cell_type": "code",
|
| 930 |
+
"execution_count": null,
|
| 931 |
+
"metadata": {
|
| 932 |
+
"id": "apwUH4Ar-o4P"
|
| 933 |
+
},
|
| 934 |
+
"outputs": [],
|
| 935 |
+
"source": [
|
| 936 |
+
"print(list(range(5, 10)))\n",
|
| 937 |
+
"print(list(range(0, 10, 3)))\n",
|
| 938 |
+
"print(list(range(-10, -100, -30)))"
|
| 939 |
+
]
|
| 940 |
+
},
|
| 941 |
+
{
|
| 942 |
+
"cell_type": "markdown",
|
| 943 |
+
"metadata": {
|
| 944 |
+
"id": "mezbBTJqCmZy"
|
| 945 |
+
},
|
| 946 |
+
"source": [
|
| 947 |
+
"###`while` loop"
|
| 948 |
+
]
|
| 949 |
+
},
|
| 950 |
+
{
|
| 951 |
+
"cell_type": "markdown",
|
| 952 |
+
"metadata": {
|
| 953 |
+
"id": "A5qN9PZTCuvS"
|
| 954 |
+
},
|
| 955 |
+
"source": [
|
| 956 |
+
"With the `while` loop we can execute a set of statements as long as a condition is true."
|
| 957 |
+
]
|
| 958 |
+
},
|
| 959 |
+
{
|
| 960 |
+
"cell_type": "code",
|
| 961 |
+
"execution_count": null,
|
| 962 |
+
"metadata": {
|
| 963 |
+
"id": "T0NbKi1hCyCE"
|
| 964 |
+
},
|
| 965 |
+
"outputs": [],
|
| 966 |
+
"source": [
|
| 967 |
+
"i = 1\n",
|
| 968 |
+
"while i < 6:\n",
|
| 969 |
+
" print(i)\n",
|
| 970 |
+
" i += 1"
|
| 971 |
+
]
|
| 972 |
+
},
|
| 973 |
+
{
|
| 974 |
+
"cell_type": "markdown",
|
| 975 |
+
"metadata": {
|
| 976 |
+
"id": "uBXI2gMx9Dno"
|
| 977 |
+
},
|
| 978 |
+
"source": [
|
| 979 |
+
"## Control Flow Tools"
|
| 980 |
+
]
|
| 981 |
+
},
|
| 982 |
+
{
|
| 983 |
+
"cell_type": "markdown",
|
| 984 |
+
"metadata": {
|
| 985 |
+
"id": "q6aeyPu39PPC"
|
| 986 |
+
},
|
| 987 |
+
"source": [
|
| 988 |
+
"###`if` statement"
|
| 989 |
+
]
|
| 990 |
+
},
|
| 991 |
+
{
|
| 992 |
+
"cell_type": "markdown",
|
| 993 |
+
"metadata": {
|
| 994 |
+
"id": "eHUyBp-E_V-x"
|
| 995 |
+
},
|
| 996 |
+
"source": [
|
| 997 |
+
"Perhaps the most well-known statement type is the if statement. There can be zero or more `elif` parts, and the `else` part is optional. The keyword `elif` is short for `else if`, and is useful to avoid excessive indentation. An `if` … `elif` … `elif` … sequence is a substitute for the `switch` or `case` statements found in other languages. For example:"
|
| 998 |
+
]
|
| 999 |
+
},
|
| 1000 |
+
{
|
| 1001 |
+
"cell_type": "code",
|
| 1002 |
+
"execution_count": null,
|
| 1003 |
+
"metadata": {
|
| 1004 |
+
"id": "nD8ITrA__Z5D"
|
| 1005 |
+
},
|
| 1006 |
+
"outputs": [],
|
| 1007 |
+
"source": [
|
| 1008 |
+
"x = int(input(\"Please enter an integer: \"))\n",
|
| 1009 |
+
"if x < 0:\n",
|
| 1010 |
+
" x = 0\n",
|
| 1011 |
+
" print('Negative changed to zero')\n",
|
| 1012 |
+
"elif x == 0:\n",
|
| 1013 |
+
" print('Zero')\n",
|
| 1014 |
+
"elif x == 1:\n",
|
| 1015 |
+
" print('Single')\n",
|
| 1016 |
+
"else:\n",
|
| 1017 |
+
" print('More')"
|
| 1018 |
+
]
|
| 1019 |
+
},
|
| 1020 |
+
{
|
| 1021 |
+
"cell_type": "markdown",
|
| 1022 |
+
"metadata": {
|
| 1023 |
+
"id": "Y0CBMSRJ9sby"
|
| 1024 |
+
},
|
| 1025 |
+
"source": [
|
| 1026 |
+
"###`break` and `continue` statements"
|
| 1027 |
+
]
|
| 1028 |
+
},
|
| 1029 |
+
{
|
| 1030 |
+
"cell_type": "markdown",
|
| 1031 |
+
"metadata": {
|
| 1032 |
+
"id": "XrfS77Kzi91S"
|
| 1033 |
+
},
|
| 1034 |
+
"source": [
|
| 1035 |
+
"The `break` statement, like in C, breaks out of the innermost enclosing `for` or `while` loop.\n",
|
| 1036 |
+
"\n",
|
| 1037 |
+
"Loop statements may have an else clause; it is executed when the loop terminates through exhaustion of the iterable (with `for`) or when the condition becomes false (with `while`), but not when the loop is terminated by a `break` statement. This is exemplified by the following loop, which searches for prime numbers:"
|
| 1038 |
+
]
|
| 1039 |
+
},
|
| 1040 |
+
{
|
| 1041 |
+
"cell_type": "code",
|
| 1042 |
+
"execution_count": null,
|
| 1043 |
+
"metadata": {
|
| 1044 |
+
"id": "S2XoBEftjaXX"
|
| 1045 |
+
},
|
| 1046 |
+
"outputs": [],
|
| 1047 |
+
"source": [
|
| 1048 |
+
"for n in range(2, 10):\n",
|
| 1049 |
+
" for x in range(2, n):\n",
|
| 1050 |
+
" if n % x == 0:\n",
|
| 1051 |
+
" print(n, 'equals', x, '*', n//x)\n",
|
| 1052 |
+
" break\n",
|
| 1053 |
+
" else:\n",
|
| 1054 |
+
" # loop fell through without finding a factor\n",
|
| 1055 |
+
" print(n, 'is a prime number')"
|
| 1056 |
+
]
|
| 1057 |
+
},
|
| 1058 |
+
{
|
| 1059 |
+
"cell_type": "markdown",
|
| 1060 |
+
"metadata": {
|
| 1061 |
+
"id": "b5pf2Vdbkl0_"
|
| 1062 |
+
},
|
| 1063 |
+
"source": [
|
| 1064 |
+
"The `continue` statement, also borrowed from C, continues with the next iteration of the loop:"
|
| 1065 |
+
]
|
| 1066 |
+
},
|
| 1067 |
+
{
|
| 1068 |
+
"cell_type": "code",
|
| 1069 |
+
"execution_count": null,
|
| 1070 |
+
"metadata": {
|
| 1071 |
+
"id": "swr6-rEwksE2"
|
| 1072 |
+
},
|
| 1073 |
+
"outputs": [],
|
| 1074 |
+
"source": [
|
| 1075 |
+
"for num in range(2, 10):\n",
|
| 1076 |
+
" if num % 2 == 0:\n",
|
| 1077 |
+
" print(\"Found an even number\", num)\n",
|
| 1078 |
+
" continue\n",
|
| 1079 |
+
" print(\"Found an odd number\", num)"
|
| 1080 |
+
]
|
| 1081 |
+
},
|
| 1082 |
+
{
|
| 1083 |
+
"cell_type": "markdown",
|
| 1084 |
+
"metadata": {
|
| 1085 |
+
"id": "JVKNnslx95_d"
|
| 1086 |
+
},
|
| 1087 |
+
"source": [
|
| 1088 |
+
"###`pass` statement"
|
| 1089 |
+
]
|
| 1090 |
+
},
|
| 1091 |
+
{
|
| 1092 |
+
"cell_type": "markdown",
|
| 1093 |
+
"metadata": {
|
| 1094 |
+
"id": "-dZk28jllX6D"
|
| 1095 |
+
},
|
| 1096 |
+
"source": [
|
| 1097 |
+
"The `pass` statement does nothing. It can be used when a statement is required syntactically but the program requires no action. For example:"
|
| 1098 |
+
]
|
| 1099 |
+
},
|
| 1100 |
+
{
|
| 1101 |
+
"cell_type": "code",
|
| 1102 |
+
"execution_count": null,
|
| 1103 |
+
"metadata": {
|
| 1104 |
+
"id": "DnJfJPkTldUz"
|
| 1105 |
+
},
|
| 1106 |
+
"outputs": [],
|
| 1107 |
+
"source": [
|
| 1108 |
+
"while True:\n",
|
| 1109 |
+
" pass # Busy-wait for keyboard interrupt. Please press the stop button to stop execution."
|
| 1110 |
+
]
|
| 1111 |
+
},
|
| 1112 |
+
{
|
| 1113 |
+
"cell_type": "markdown",
|
| 1114 |
+
"metadata": {
|
| 1115 |
+
"id": "-tWb6by_lkn3"
|
| 1116 |
+
},
|
| 1117 |
+
"source": [
|
| 1118 |
+
"This is commonly used for creating minimal classes:"
|
| 1119 |
+
]
|
| 1120 |
+
},
|
| 1121 |
+
{
|
| 1122 |
+
"cell_type": "code",
|
| 1123 |
+
"execution_count": null,
|
| 1124 |
+
"metadata": {
|
| 1125 |
+
"id": "57_9LZkplsSz"
|
| 1126 |
+
},
|
| 1127 |
+
"outputs": [],
|
| 1128 |
+
"source": [
|
| 1129 |
+
"class MyEmptyClass:\n",
|
| 1130 |
+
" pass"
|
| 1131 |
+
]
|
| 1132 |
+
},
|
| 1133 |
+
{
|
| 1134 |
+
"cell_type": "markdown",
|
| 1135 |
+
"metadata": {
|
| 1136 |
+
"id": "45jdssyFlxqo"
|
| 1137 |
+
},
|
| 1138 |
+
"source": [
|
| 1139 |
+
"Another place `pass` can be used is as a place-holder for a function or conditional body when you are working on new code, allowing you to keep thinking at a more abstract level. The `pass` is silently ignored:"
|
| 1140 |
+
]
|
| 1141 |
+
},
|
| 1142 |
+
{
|
| 1143 |
+
"cell_type": "code",
|
| 1144 |
+
"execution_count": null,
|
| 1145 |
+
"metadata": {
|
| 1146 |
+
"id": "0r9Dikptl4-1"
|
| 1147 |
+
},
|
| 1148 |
+
"outputs": [],
|
| 1149 |
+
"source": [
|
| 1150 |
+
"def initlog(*args):\n",
|
| 1151 |
+
" pass # Remember to implement this!"
|
| 1152 |
+
]
|
| 1153 |
+
},
|
| 1154 |
+
{
|
| 1155 |
+
"cell_type": "markdown",
|
| 1156 |
+
"metadata": {
|
| 1157 |
+
"id": "AXA4jrEOL9hM"
|
| 1158 |
+
},
|
| 1159 |
+
"source": [
|
| 1160 |
+
"###Functions"
|
| 1161 |
+
]
|
| 1162 |
+
},
|
| 1163 |
+
{
|
| 1164 |
+
"cell_type": "markdown",
|
| 1165 |
+
"metadata": {
|
| 1166 |
+
"id": "WaRms-QfL9hN"
|
| 1167 |
+
},
|
| 1168 |
+
"source": [
|
| 1169 |
+
"Python functions are defined using the `def` keyword. For example:"
|
| 1170 |
+
]
|
| 1171 |
+
},
|
| 1172 |
+
{
|
| 1173 |
+
"cell_type": "code",
|
| 1174 |
+
"execution_count": null,
|
| 1175 |
+
"metadata": {
|
| 1176 |
+
"id": "kiMDUr58L9hN"
|
| 1177 |
+
},
|
| 1178 |
+
"outputs": [],
|
| 1179 |
+
"source": [
|
| 1180 |
+
"def sign(x):\n",
|
| 1181 |
+
" if x > 0:\n",
|
| 1182 |
+
" return 'positive'\n",
|
| 1183 |
+
" elif x < 0:\n",
|
| 1184 |
+
" return 'negative'\n",
|
| 1185 |
+
" else:\n",
|
| 1186 |
+
" return 'zero'\n",
|
| 1187 |
+
"\n",
|
| 1188 |
+
"for x in [-1, 0, 1]:\n",
|
| 1189 |
+
" print(sign(x))"
|
| 1190 |
+
]
|
| 1191 |
+
},
|
| 1192 |
+
{
|
| 1193 |
+
"cell_type": "markdown",
|
| 1194 |
+
"metadata": {
|
| 1195 |
+
"id": "U-QJFt8TL9hR"
|
| 1196 |
+
},
|
| 1197 |
+
"source": [
|
| 1198 |
+
"We will often define functions to take optional keyword arguments, like this:"
|
| 1199 |
+
]
|
| 1200 |
+
},
|
| 1201 |
+
{
|
| 1202 |
+
"cell_type": "code",
|
| 1203 |
+
"execution_count": null,
|
| 1204 |
+
"metadata": {
|
| 1205 |
+
"id": "PfsZ3DazL9hR"
|
| 1206 |
+
},
|
| 1207 |
+
"outputs": [],
|
| 1208 |
+
"source": [
|
| 1209 |
+
"def hello(name, loud=False):\n",
|
| 1210 |
+
" if loud:\n",
|
| 1211 |
+
" print('HELLO, {}'.format(name.upper()))\n",
|
| 1212 |
+
" else:\n",
|
| 1213 |
+
" print('Hello, {}!'.format(name))\n",
|
| 1214 |
+
"\n",
|
| 1215 |
+
"hello('Bob')\n",
|
| 1216 |
+
"hello('Fred', loud=True)"
|
| 1217 |
+
]
|
| 1218 |
+
},
|
| 1219 |
+
{
|
| 1220 |
+
"cell_type": "markdown",
|
| 1221 |
+
"metadata": {
|
| 1222 |
+
"id": "ObA9PRtQL9hT"
|
| 1223 |
+
},
|
| 1224 |
+
"source": [
|
| 1225 |
+
"###Classes"
|
| 1226 |
+
]
|
| 1227 |
+
},
|
| 1228 |
+
{
|
| 1229 |
+
"cell_type": "markdown",
|
| 1230 |
+
"metadata": {
|
| 1231 |
+
"id": "hAzL_lTkL9hU"
|
| 1232 |
+
},
|
| 1233 |
+
"source": [
|
| 1234 |
+
"In object-oriented programming, a class is a template definition of the methods and variables in a particular kind of object. Thus, an object is a specific instance of a class; it contains real values instead of variables. For more details, on class in object oriented programming, please have a look on this [link](https://www.w3schools.com/java/java_oop.asp). The syntax for defining classes in Python is straightforward and can be done as follows."
|
| 1235 |
+
]
|
| 1236 |
+
},
|
| 1237 |
+
{
|
| 1238 |
+
"cell_type": "code",
|
| 1239 |
+
"execution_count": null,
|
| 1240 |
+
"metadata": {
|
| 1241 |
+
"id": "RWdbaGigL9hU"
|
| 1242 |
+
},
|
| 1243 |
+
"outputs": [],
|
| 1244 |
+
"source": [
|
| 1245 |
+
"class Greeter:\n",
|
| 1246 |
+
"\n",
|
| 1247 |
+
" # Constructor\n",
|
| 1248 |
+
" def __init__(self, name):\n",
|
| 1249 |
+
" self.name = name # Create an instance variable\n",
|
| 1250 |
+
"\n",
|
| 1251 |
+
" # Instance method\n",
|
| 1252 |
+
" def greet(self, loud=False):\n",
|
| 1253 |
+
" if loud:\n",
|
| 1254 |
+
" print('HELLO, {}'.format(self.name.upper()))\n",
|
| 1255 |
+
" else:\n",
|
| 1256 |
+
" print('Hello, {}!'.format(self.name))\n",
|
| 1257 |
+
"\n",
|
| 1258 |
+
"g = Greeter('Fred') # Construct an instance of the Greeter class\n",
|
| 1259 |
+
"g.greet() # Call an instance method; prints \"Hello, Fred\"\n",
|
| 1260 |
+
"g.greet(loud=True) # Call an instance method; prints \"HELLO, FRED!\""
|
| 1261 |
+
]
|
| 1262 |
+
},
|
| 1263 |
+
{
|
| 1264 |
+
"cell_type": "markdown",
|
| 1265 |
+
"metadata": {
|
| 1266 |
+
"id": "3cfrOV4dL9hW"
|
| 1267 |
+
},
|
| 1268 |
+
"source": [
|
| 1269 |
+
"##Numpy"
|
| 1270 |
+
]
|
| 1271 |
+
},
|
| 1272 |
+
{
|
| 1273 |
+
"cell_type": "markdown",
|
| 1274 |
+
"metadata": {
|
| 1275 |
+
"id": "fY12nHhyL9hX"
|
| 1276 |
+
},
|
| 1277 |
+
"source": [
|
| 1278 |
+
"Numpy is the core library for scientific computing in Python. It provides a high-performance multidimensional array object, and tools for working with these arrays. If you are already familiar with MATLAB, you might find this [tutorial](http://wiki.scipy.org/NumPy_for_Matlab_Users) useful to get started with Numpy. To use Numpy, we first need to import the `numpy` package."
|
| 1279 |
+
]
|
| 1280 |
+
},
|
| 1281 |
+
{
|
| 1282 |
+
"cell_type": "markdown",
|
| 1283 |
+
"metadata": {
|
| 1284 |
+
"id": "2_lpLqwZpd-4"
|
| 1285 |
+
},
|
| 1286 |
+
"source": [
|
| 1287 |
+
"### Importing a package"
|
| 1288 |
+
]
|
| 1289 |
+
},
|
| 1290 |
+
{
|
| 1291 |
+
"cell_type": "markdown",
|
| 1292 |
+
"metadata": {
|
| 1293 |
+
"id": "hMmlsjljBbVE"
|
| 1294 |
+
},
|
| 1295 |
+
"source": [
|
| 1296 |
+
"In Python, a package or a module can be imported in many different ways, some of which are shown below. For more details, please have a look on this [documentation](https://docs.python.org/3/tutorial/modules.html#more-on-modules).\n",
|
| 1297 |
+
"\n",
|
| 1298 |
+
"\n",
|
| 1299 |
+
"```\n",
|
| 1300 |
+
"import numpy # import numpy, one can use it as numpy\n",
|
| 1301 |
+
"import numpy as np # import numpy and call it np\n",
|
| 1302 |
+
"from numpy import * # import all the modules from numpy\n",
|
| 1303 |
+
"from numpy import sum # import the \"sum\" function from numpy\n",
|
| 1304 |
+
"```\n",
|
| 1305 |
+
"\n"
|
| 1306 |
+
]
|
| 1307 |
+
},
|
| 1308 |
+
{
|
| 1309 |
+
"cell_type": "code",
|
| 1310 |
+
"execution_count": null,
|
| 1311 |
+
"metadata": {
|
| 1312 |
+
"id": "58QdX8BLL9hZ"
|
| 1313 |
+
},
|
| 1314 |
+
"outputs": [],
|
| 1315 |
+
"source": [
|
| 1316 |
+
"import numpy as np # import numpy and call it np. So the sum function of numpy can be called as np.sum()"
|
| 1317 |
+
]
|
| 1318 |
+
},
|
| 1319 |
+
{
|
| 1320 |
+
"cell_type": "markdown",
|
| 1321 |
+
"metadata": {
|
| 1322 |
+
"id": "DDx6v1EdL9hb"
|
| 1323 |
+
},
|
| 1324 |
+
"source": [
|
| 1325 |
+
"###Arrays"
|
| 1326 |
+
]
|
| 1327 |
+
},
|
| 1328 |
+
{
|
| 1329 |
+
"cell_type": "markdown",
|
| 1330 |
+
"metadata": {
|
| 1331 |
+
"id": "f-Zv3f7LL9hc"
|
| 1332 |
+
},
|
| 1333 |
+
"source": [
|
| 1334 |
+
"A numpy array is a grid of values, all of the same type, and is indexed by a tuple of nonnegative integers. The number of dimensions is the rank of the array; the shape of an array is a tuple of integers giving the size of the array along each dimension."
|
| 1335 |
+
]
|
| 1336 |
+
},
|
| 1337 |
+
{
|
| 1338 |
+
"cell_type": "markdown",
|
| 1339 |
+
"metadata": {
|
| 1340 |
+
"id": "_eMTRnZRL9hc"
|
| 1341 |
+
},
|
| 1342 |
+
"source": [
|
| 1343 |
+
"We can initialize numpy arrays from nested Python lists, and access elements using square brackets:"
|
| 1344 |
+
]
|
| 1345 |
+
},
|
| 1346 |
+
{
|
| 1347 |
+
"cell_type": "code",
|
| 1348 |
+
"execution_count": null,
|
| 1349 |
+
"metadata": {
|
| 1350 |
+
"id": "-l3JrGxCL9hc"
|
| 1351 |
+
},
|
| 1352 |
+
"outputs": [],
|
| 1353 |
+
"source": [
|
| 1354 |
+
"a = np.array([1, 2, 3]) # Create a rank 1 array\n",
|
| 1355 |
+
"print(type(a), a.shape, a[0], a[1], a[2])\n",
|
| 1356 |
+
"a[0] = 5 # Change an element of the array\n",
|
| 1357 |
+
"print(a)"
|
| 1358 |
+
]
|
| 1359 |
+
},
|
| 1360 |
+
{
|
| 1361 |
+
"cell_type": "code",
|
| 1362 |
+
"execution_count": null,
|
| 1363 |
+
"metadata": {
|
| 1364 |
+
"id": "ma6mk-kdL9hh"
|
| 1365 |
+
},
|
| 1366 |
+
"outputs": [],
|
| 1367 |
+
"source": [
|
| 1368 |
+
"b = np.array([[1,2,3],[4,5,6]]) # Create a rank 2 array\n",
|
| 1369 |
+
"print(b)"
|
| 1370 |
+
]
|
| 1371 |
+
},
|
| 1372 |
+
{
|
| 1373 |
+
"cell_type": "code",
|
| 1374 |
+
"execution_count": null,
|
| 1375 |
+
"metadata": {
|
| 1376 |
+
"id": "ymfSHAwtL9hj"
|
| 1377 |
+
},
|
| 1378 |
+
"outputs": [],
|
| 1379 |
+
"source": [
|
| 1380 |
+
"print(b.shape)\n",
|
| 1381 |
+
"print(b[0, 0], b[0, 1], b[1, 0])"
|
| 1382 |
+
]
|
| 1383 |
+
},
|
| 1384 |
+
{
|
| 1385 |
+
"cell_type": "markdown",
|
| 1386 |
+
"metadata": {
|
| 1387 |
+
"id": "F2qwdyvuL9hn"
|
| 1388 |
+
},
|
| 1389 |
+
"source": [
|
| 1390 |
+
"Numpy also provides many functions to create arrays:"
|
| 1391 |
+
]
|
| 1392 |
+
},
|
| 1393 |
+
{
|
| 1394 |
+
"cell_type": "code",
|
| 1395 |
+
"execution_count": null,
|
| 1396 |
+
"metadata": {
|
| 1397 |
+
"id": "mVTN_EBqL9hn"
|
| 1398 |
+
},
|
| 1399 |
+
"outputs": [],
|
| 1400 |
+
"source": [
|
| 1401 |
+
"a = np.zeros((2,2)) # Create an array of all zeros\n",
|
| 1402 |
+
"print(a)"
|
| 1403 |
+
]
|
| 1404 |
+
},
|
| 1405 |
+
{
|
| 1406 |
+
"cell_type": "code",
|
| 1407 |
+
"execution_count": null,
|
| 1408 |
+
"metadata": {
|
| 1409 |
+
"id": "skiKlNmlL9h5"
|
| 1410 |
+
},
|
| 1411 |
+
"outputs": [],
|
| 1412 |
+
"source": [
|
| 1413 |
+
"b = np.ones((1,2)) # Create an array of all ones\n",
|
| 1414 |
+
"print(b)"
|
| 1415 |
+
]
|
| 1416 |
+
},
|
| 1417 |
+
{
|
| 1418 |
+
"cell_type": "code",
|
| 1419 |
+
"execution_count": null,
|
| 1420 |
+
"metadata": {
|
| 1421 |
+
"id": "HtFsr03bL9h7"
|
| 1422 |
+
},
|
| 1423 |
+
"outputs": [],
|
| 1424 |
+
"source": [
|
| 1425 |
+
"c = np.full((2,2), 7) # Create a constant array\n",
|
| 1426 |
+
"print(c)"
|
| 1427 |
+
]
|
| 1428 |
+
},
|
| 1429 |
+
{
|
| 1430 |
+
"cell_type": "code",
|
| 1431 |
+
"execution_count": null,
|
| 1432 |
+
"metadata": {
|
| 1433 |
+
"id": "-QcALHvkL9h9"
|
| 1434 |
+
},
|
| 1435 |
+
"outputs": [],
|
| 1436 |
+
"source": [
|
| 1437 |
+
"d = np.eye(2) # Create a 2x2 identity matrix\n",
|
| 1438 |
+
"print(d)"
|
| 1439 |
+
]
|
| 1440 |
+
},
|
| 1441 |
+
{
|
| 1442 |
+
"cell_type": "code",
|
| 1443 |
+
"execution_count": null,
|
| 1444 |
+
"metadata": {
|
| 1445 |
+
"id": "RCpaYg9qL9iA"
|
| 1446 |
+
},
|
| 1447 |
+
"outputs": [],
|
| 1448 |
+
"source": [
|
| 1449 |
+
"e = np.random.random((2,2)) # Create an array filled with random values\n",
|
| 1450 |
+
"print(e)"
|
| 1451 |
+
]
|
| 1452 |
+
},
|
| 1453 |
+
{
|
| 1454 |
+
"cell_type": "markdown",
|
| 1455 |
+
"metadata": {
|
| 1456 |
+
"id": "jI5qcSDfL9iC"
|
| 1457 |
+
},
|
| 1458 |
+
"source": [
|
| 1459 |
+
"###Array indexing"
|
| 1460 |
+
]
|
| 1461 |
+
},
|
| 1462 |
+
{
|
| 1463 |
+
"cell_type": "markdown",
|
| 1464 |
+
"metadata": {
|
| 1465 |
+
"id": "M-E4MUeVL9iC"
|
| 1466 |
+
},
|
| 1467 |
+
"source": [
|
| 1468 |
+
"Numpy offers several ways to index into arrays.\n",
|
| 1469 |
+
"\n",
|
| 1470 |
+
"Slicing: Similar to Python lists, numpy arrays can be sliced. Since arrays may be multidimensional, you must specify a slice for each dimension of the array:"
|
| 1471 |
+
]
|
| 1472 |
+
},
|
| 1473 |
+
{
|
| 1474 |
+
"cell_type": "code",
|
| 1475 |
+
"execution_count": null,
|
| 1476 |
+
"metadata": {
|
| 1477 |
+
"id": "wLWA0udwL9iD"
|
| 1478 |
+
},
|
| 1479 |
+
"outputs": [],
|
| 1480 |
+
"source": [
|
| 1481 |
+
"import numpy as np\n",
|
| 1482 |
+
"\n",
|
| 1483 |
+
"# Create the following rank 2 array with shape (3, 4)\n",
|
| 1484 |
+
"# [[ 1 2 3 4]\n",
|
| 1485 |
+
"# [ 5 6 7 8]\n",
|
| 1486 |
+
"# [ 9 10 11 12]]\n",
|
| 1487 |
+
"a = np.array([[1,2,3,4], [5,6,7,8], [9,10,11,12]])\n",
|
| 1488 |
+
"\n",
|
| 1489 |
+
"# Use slicing to pull out the subarray consisting of the first 2 rows\n",
|
| 1490 |
+
"# and columns 1 and 2; b is the following array of shape (2, 2):\n",
|
| 1491 |
+
"# [[2 3]\n",
|
| 1492 |
+
"# [6 7]]\n",
|
| 1493 |
+
"b = a[:2, 1:3]\n",
|
| 1494 |
+
"print(b)"
|
| 1495 |
+
]
|
| 1496 |
+
},
|
| 1497 |
+
{
|
| 1498 |
+
"cell_type": "markdown",
|
| 1499 |
+
"metadata": {
|
| 1500 |
+
"id": "KahhtZKYL9iF"
|
| 1501 |
+
},
|
| 1502 |
+
"source": [
|
| 1503 |
+
"A slice of an array is a view into the same data, so modifying it will modify the original array."
|
| 1504 |
+
]
|
| 1505 |
+
},
|
| 1506 |
+
{
|
| 1507 |
+
"cell_type": "code",
|
| 1508 |
+
"execution_count": null,
|
| 1509 |
+
"metadata": {
|
| 1510 |
+
"id": "1kmtaFHuL9iG"
|
| 1511 |
+
},
|
| 1512 |
+
"outputs": [],
|
| 1513 |
+
"source": [
|
| 1514 |
+
"print(a[0, 1])\n",
|
| 1515 |
+
"b[0, 0] = 77 # b[0, 0] is the same piece of data as a[0, 1]\n",
|
| 1516 |
+
"print(a[0, 1])"
|
| 1517 |
+
]
|
| 1518 |
+
},
|
| 1519 |
+
{
|
| 1520 |
+
"cell_type": "markdown",
|
| 1521 |
+
"metadata": {
|
| 1522 |
+
"id": "_Zcf3zi-L9iI"
|
| 1523 |
+
},
|
| 1524 |
+
"source": [
|
| 1525 |
+
"You can also mix integer indexing with slice indexing. However, doing so will yield an array of lower rank than the original array. Note that this is quite different from the way that MATLAB handles array slicing:"
|
| 1526 |
+
]
|
| 1527 |
+
},
|
| 1528 |
+
{
|
| 1529 |
+
"cell_type": "code",
|
| 1530 |
+
"execution_count": null,
|
| 1531 |
+
"metadata": {
|
| 1532 |
+
"id": "G6lfbPuxL9iJ"
|
| 1533 |
+
},
|
| 1534 |
+
"outputs": [],
|
| 1535 |
+
"source": [
|
| 1536 |
+
"# Create the following rank 2 array with shape (3, 4)\n",
|
| 1537 |
+
"a = np.array([[1,2,3,4], [5,6,7,8], [9,10,11,12]])\n",
|
| 1538 |
+
"print(a)"
|
| 1539 |
+
]
|
| 1540 |
+
},
|
| 1541 |
+
{
|
| 1542 |
+
"cell_type": "markdown",
|
| 1543 |
+
"metadata": {
|
| 1544 |
+
"id": "NCye3NXhL9iL"
|
| 1545 |
+
},
|
| 1546 |
+
"source": [
|
| 1547 |
+
"Two ways of accessing the data in the middle row of the array.\n",
|
| 1548 |
+
"Mixing integer indexing with slices yields an array of lower rank,\n",
|
| 1549 |
+
"while using only slices yields an array of the same rank as the\n",
|
| 1550 |
+
"original array:"
|
| 1551 |
+
]
|
| 1552 |
+
},
|
| 1553 |
+
{
|
| 1554 |
+
"cell_type": "code",
|
| 1555 |
+
"execution_count": null,
|
| 1556 |
+
"metadata": {
|
| 1557 |
+
"id": "EOiEMsmNL9iL"
|
| 1558 |
+
},
|
| 1559 |
+
"outputs": [],
|
| 1560 |
+
"source": [
|
| 1561 |
+
"row_r1 = a[1, :] # Rank 1 view of the second row of a\n",
|
| 1562 |
+
"row_r2 = a[1:2, :] # Rank 2 view of the second row of a\n",
|
| 1563 |
+
"row_r3 = a[[1], :] # Rank 2 view of the second row of a\n",
|
| 1564 |
+
"print(row_r1, row_r1.shape)\n",
|
| 1565 |
+
"print(row_r2, row_r2.shape)\n",
|
| 1566 |
+
"print(row_r3, row_r3.shape)"
|
| 1567 |
+
]
|
| 1568 |
+
},
|
| 1569 |
+
{
|
| 1570 |
+
"cell_type": "code",
|
| 1571 |
+
"execution_count": null,
|
| 1572 |
+
"metadata": {
|
| 1573 |
+
"id": "JXu73pfDL9iN"
|
| 1574 |
+
},
|
| 1575 |
+
"outputs": [],
|
| 1576 |
+
"source": [
|
| 1577 |
+
"# We can make the same distinction when accessing columns of an array:\n",
|
| 1578 |
+
"col_r1 = a[:, 1]\n",
|
| 1579 |
+
"col_r2 = a[:, 1:2]\n",
|
| 1580 |
+
"print(col_r1, col_r1.shape)\n",
|
| 1581 |
+
"print()\n",
|
| 1582 |
+
"print(col_r2, col_r2.shape)"
|
| 1583 |
+
]
|
| 1584 |
+
},
|
| 1585 |
+
{
|
| 1586 |
+
"cell_type": "markdown",
|
| 1587 |
+
"metadata": {
|
| 1588 |
+
"id": "VP3916bOL9iP"
|
| 1589 |
+
},
|
| 1590 |
+
"source": [
|
| 1591 |
+
"Integer array indexing: When you index into numpy arrays using slicing, the resulting array view will always be a subarray of the original array. In contrast, integer array indexing allows you to construct arbitrary arrays using the data from another array. Here is an example:"
|
| 1592 |
+
]
|
| 1593 |
+
},
|
| 1594 |
+
{
|
| 1595 |
+
"cell_type": "code",
|
| 1596 |
+
"execution_count": null,
|
| 1597 |
+
"metadata": {
|
| 1598 |
+
"id": "TBnWonIDL9iP"
|
| 1599 |
+
},
|
| 1600 |
+
"outputs": [],
|
| 1601 |
+
"source": [
|
| 1602 |
+
"a = np.array([[1,2], [3, 4], [5, 6]])\n",
|
| 1603 |
+
"\n",
|
| 1604 |
+
"# An example of integer array indexing.\n",
|
| 1605 |
+
"# The returned array will have shape (3,) and\n",
|
| 1606 |
+
"print(a[[0, 1, 2], [0, 1, 0]])\n",
|
| 1607 |
+
"\n",
|
| 1608 |
+
"# The above example of integer array indexing is equivalent to this:\n",
|
| 1609 |
+
"print(np.array([a[0, 0], a[1, 1], a[2, 0]]))"
|
| 1610 |
+
]
|
| 1611 |
+
},
|
| 1612 |
+
{
|
| 1613 |
+
"cell_type": "code",
|
| 1614 |
+
"execution_count": null,
|
| 1615 |
+
"metadata": {
|
| 1616 |
+
"id": "n7vuati-L9iR"
|
| 1617 |
+
},
|
| 1618 |
+
"outputs": [],
|
| 1619 |
+
"source": [
|
| 1620 |
+
"# When using integer array indexing, you can reuse the same\n",
|
| 1621 |
+
"# element from the source array:\n",
|
| 1622 |
+
"print(a[[0, 0], [1, 1]])\n",
|
| 1623 |
+
"\n",
|
| 1624 |
+
"# Equivalent to the previous integer array indexing example\n",
|
| 1625 |
+
"print(np.array([a[0, 1], a[0, 1]]))"
|
| 1626 |
+
]
|
| 1627 |
+
},
|
| 1628 |
+
{
|
| 1629 |
+
"cell_type": "markdown",
|
| 1630 |
+
"metadata": {
|
| 1631 |
+
"id": "kaipSLafL9iU"
|
| 1632 |
+
},
|
| 1633 |
+
"source": [
|
| 1634 |
+
"One useful trick with integer array indexing is selecting or mutating one element from each row of a matrix:"
|
| 1635 |
+
]
|
| 1636 |
+
},
|
| 1637 |
+
{
|
| 1638 |
+
"cell_type": "code",
|
| 1639 |
+
"execution_count": null,
|
| 1640 |
+
"metadata": {
|
| 1641 |
+
"id": "ehqsV7TXL9iU"
|
| 1642 |
+
},
|
| 1643 |
+
"outputs": [],
|
| 1644 |
+
"source": [
|
| 1645 |
+
"# Create a new array from which we will select elements\n",
|
| 1646 |
+
"a = np.array([[1,2,3], [4,5,6], [7,8,9], [10, 11, 12]])\n",
|
| 1647 |
+
"print(a)"
|
| 1648 |
+
]
|
| 1649 |
+
},
|
| 1650 |
+
{
|
| 1651 |
+
"cell_type": "code",
|
| 1652 |
+
"execution_count": null,
|
| 1653 |
+
"metadata": {
|
| 1654 |
+
"id": "pAPOoqy5L9iV"
|
| 1655 |
+
},
|
| 1656 |
+
"outputs": [],
|
| 1657 |
+
"source": [
|
| 1658 |
+
"# Create an array of indices\n",
|
| 1659 |
+
"b = np.array([0, 2, 0, 1])\n",
|
| 1660 |
+
"\n",
|
| 1661 |
+
"# Select one element from each row of a using the indices in b\n",
|
| 1662 |
+
"print(a[np.arange(4), b]) # Prints \"[ 1 6 7 11]\""
|
| 1663 |
+
]
|
| 1664 |
+
},
|
| 1665 |
+
{
|
| 1666 |
+
"cell_type": "code",
|
| 1667 |
+
"execution_count": null,
|
| 1668 |
+
"metadata": {
|
| 1669 |
+
"id": "6v1PdI1DL9ib"
|
| 1670 |
+
},
|
| 1671 |
+
"outputs": [],
|
| 1672 |
+
"source": [
|
| 1673 |
+
"# Mutate one element from each row of a using the indices in b\n",
|
| 1674 |
+
"a[np.arange(4), b] += 10\n",
|
| 1675 |
+
"print(a)"
|
| 1676 |
+
]
|
| 1677 |
+
},
|
| 1678 |
+
{
|
| 1679 |
+
"cell_type": "markdown",
|
| 1680 |
+
"metadata": {
|
| 1681 |
+
"id": "kaE8dBGgL9id"
|
| 1682 |
+
},
|
| 1683 |
+
"source": [
|
| 1684 |
+
"Boolean array indexing: Boolean array indexing lets you pick out arbitrary elements of an array. Frequently this type of indexing is used to select the elements of an array that satisfy some condition. Here is an example:"
|
| 1685 |
+
]
|
| 1686 |
+
},
|
| 1687 |
+
{
|
| 1688 |
+
"cell_type": "code",
|
| 1689 |
+
"execution_count": null,
|
| 1690 |
+
"metadata": {
|
| 1691 |
+
"id": "32PusjtKL9id"
|
| 1692 |
+
},
|
| 1693 |
+
"outputs": [],
|
| 1694 |
+
"source": [
|
| 1695 |
+
"import numpy as np\n",
|
| 1696 |
+
"\n",
|
| 1697 |
+
"a = np.array([[1,2], [3, 4], [5, 6]])\n",
|
| 1698 |
+
"\n",
|
| 1699 |
+
"bool_idx = (a > 2) # Find the elements of a that are bigger than 2;\n",
|
| 1700 |
+
" # this returns a numpy array of Booleans of the same\n",
|
| 1701 |
+
" # shape as a, where each slot of bool_idx tells\n",
|
| 1702 |
+
" # whether that element of a is > 2.\n",
|
| 1703 |
+
"\n",
|
| 1704 |
+
"print(bool_idx)"
|
| 1705 |
+
]
|
| 1706 |
+
},
|
| 1707 |
+
{
|
| 1708 |
+
"cell_type": "code",
|
| 1709 |
+
"execution_count": null,
|
| 1710 |
+
"metadata": {
|
| 1711 |
+
"id": "cb2IRMXaL9if"
|
| 1712 |
+
},
|
| 1713 |
+
"outputs": [],
|
| 1714 |
+
"source": [
|
| 1715 |
+
"# We use boolean array indexing to construct a rank 1 array\n",
|
| 1716 |
+
"# consisting of the elements of a corresponding to the True values\n",
|
| 1717 |
+
"# of bool_idx\n",
|
| 1718 |
+
"print(a[bool_idx])\n",
|
| 1719 |
+
"\n",
|
| 1720 |
+
"# We can do all of the above in a single concise statement:\n",
|
| 1721 |
+
"print(a[a > 2])"
|
| 1722 |
+
]
|
| 1723 |
+
},
|
| 1724 |
+
{
|
| 1725 |
+
"cell_type": "markdown",
|
| 1726 |
+
"metadata": {
|
| 1727 |
+
"id": "CdofMonAL9ih"
|
| 1728 |
+
},
|
| 1729 |
+
"source": [
|
| 1730 |
+
"For brevity we have left out a lot of details about numpy array indexing; if you want to know more you should read the documentation."
|
| 1731 |
+
]
|
| 1732 |
+
},
|
| 1733 |
+
{
|
| 1734 |
+
"cell_type": "markdown",
|
| 1735 |
+
"metadata": {
|
| 1736 |
+
"id": "jTctwqdQL9ih"
|
| 1737 |
+
},
|
| 1738 |
+
"source": [
|
| 1739 |
+
"###Datatypes"
|
| 1740 |
+
]
|
| 1741 |
+
},
|
| 1742 |
+
{
|
| 1743 |
+
"cell_type": "markdown",
|
| 1744 |
+
"metadata": {
|
| 1745 |
+
"id": "kSZQ1WkIL9ih"
|
| 1746 |
+
},
|
| 1747 |
+
"source": [
|
| 1748 |
+
"Every numpy array is a grid of elements of the same type. Numpy provides a large set of numeric datatypes that you can use to construct arrays. Numpy tries to guess a datatype when you create an array, but functions that construct arrays usually also include an optional argument to explicitly specify the datatype. Here is an example:"
|
| 1749 |
+
]
|
| 1750 |
+
},
|
| 1751 |
+
{
|
| 1752 |
+
"cell_type": "code",
|
| 1753 |
+
"execution_count": null,
|
| 1754 |
+
"metadata": {
|
| 1755 |
+
"id": "4za4O0m5L9ih"
|
| 1756 |
+
},
|
| 1757 |
+
"outputs": [],
|
| 1758 |
+
"source": [
|
| 1759 |
+
"x = np.array([1, 2]) # Let numpy choose the datatype\n",
|
| 1760 |
+
"y = np.array([1.0, 2.0]) # Let numpy choose the datatype\n",
|
| 1761 |
+
"z = np.array([1, 2], dtype=np.int64) # Force a particular datatype\n",
|
| 1762 |
+
"\n",
|
| 1763 |
+
"print(x.dtype, y.dtype, z.dtype)"
|
| 1764 |
+
]
|
| 1765 |
+
},
|
| 1766 |
+
{
|
| 1767 |
+
"cell_type": "markdown",
|
| 1768 |
+
"metadata": {
|
| 1769 |
+
"id": "RLVIsZQpL9ik"
|
| 1770 |
+
},
|
| 1771 |
+
"source": [
|
| 1772 |
+
"You can read all about numpy datatypes in the [documentation](http://docs.scipy.org/doc/numpy/reference/arrays.dtypes.html)."
|
| 1773 |
+
]
|
| 1774 |
+
},
|
| 1775 |
+
{
|
| 1776 |
+
"cell_type": "markdown",
|
| 1777 |
+
"metadata": {
|
| 1778 |
+
"id": "TuB-fdhIL9ik"
|
| 1779 |
+
},
|
| 1780 |
+
"source": [
|
| 1781 |
+
"###Array math"
|
| 1782 |
+
]
|
| 1783 |
+
},
|
| 1784 |
+
{
|
| 1785 |
+
"cell_type": "markdown",
|
| 1786 |
+
"metadata": {
|
| 1787 |
+
"id": "18e8V8elL9ik"
|
| 1788 |
+
},
|
| 1789 |
+
"source": [
|
| 1790 |
+
"Basic mathematical functions operate elementwise on arrays, and are available both as operator overloads and as functions in the numpy module:"
|
| 1791 |
+
]
|
| 1792 |
+
},
|
| 1793 |
+
{
|
| 1794 |
+
"cell_type": "code",
|
| 1795 |
+
"execution_count": null,
|
| 1796 |
+
"metadata": {
|
| 1797 |
+
"id": "gHKvBrSKL9il"
|
| 1798 |
+
},
|
| 1799 |
+
"outputs": [],
|
| 1800 |
+
"source": [
|
| 1801 |
+
"x = np.array([[1,2],[3,4]], dtype=np.float64)\n",
|
| 1802 |
+
"y = np.array([[5,6],[7,8]], dtype=np.float64)\n",
|
| 1803 |
+
"\n",
|
| 1804 |
+
"# Elementwise sum; both produce the array\n",
|
| 1805 |
+
"print(x + y)\n",
|
| 1806 |
+
"print(np.add(x, y))"
|
| 1807 |
+
]
|
| 1808 |
+
},
|
| 1809 |
+
{
|
| 1810 |
+
"cell_type": "code",
|
| 1811 |
+
"execution_count": null,
|
| 1812 |
+
"metadata": {
|
| 1813 |
+
"id": "1fZtIAMxL9in"
|
| 1814 |
+
},
|
| 1815 |
+
"outputs": [],
|
| 1816 |
+
"source": [
|
| 1817 |
+
"# Elementwise difference; both produce the array\n",
|
| 1818 |
+
"print(x - y)\n",
|
| 1819 |
+
"print(np.subtract(x, y))"
|
| 1820 |
+
]
|
| 1821 |
+
},
|
| 1822 |
+
{
|
| 1823 |
+
"cell_type": "code",
|
| 1824 |
+
"execution_count": null,
|
| 1825 |
+
"metadata": {
|
| 1826 |
+
"id": "nil4AScML9io"
|
| 1827 |
+
},
|
| 1828 |
+
"outputs": [],
|
| 1829 |
+
"source": [
|
| 1830 |
+
"# Elementwise product; both produce the array\n",
|
| 1831 |
+
"print(x * y)\n",
|
| 1832 |
+
"print(np.multiply(x, y))"
|
| 1833 |
+
]
|
| 1834 |
+
},
|
| 1835 |
+
{
|
| 1836 |
+
"cell_type": "code",
|
| 1837 |
+
"execution_count": null,
|
| 1838 |
+
"metadata": {
|
| 1839 |
+
"id": "0JoA4lH6L9ip"
|
| 1840 |
+
},
|
| 1841 |
+
"outputs": [],
|
| 1842 |
+
"source": [
|
| 1843 |
+
"# Elementwise division; both produce the array\n",
|
| 1844 |
+
"# [[ 0.2 0.33333333]\n",
|
| 1845 |
+
"# [ 0.42857143 0.5 ]]\n",
|
| 1846 |
+
"print(x / y)\n",
|
| 1847 |
+
"print(np.divide(x, y))"
|
| 1848 |
+
]
|
| 1849 |
+
},
|
| 1850 |
+
{
|
| 1851 |
+
"cell_type": "code",
|
| 1852 |
+
"execution_count": null,
|
| 1853 |
+
"metadata": {
|
| 1854 |
+
"id": "g0iZuA6bL9ir"
|
| 1855 |
+
},
|
| 1856 |
+
"outputs": [],
|
| 1857 |
+
"source": [
|
| 1858 |
+
"# Elementwise square root; produces the array\n",
|
| 1859 |
+
"# [[ 1. 1.41421356]\n",
|
| 1860 |
+
"# [ 1.73205081 2. ]]\n",
|
| 1861 |
+
"print(np.sqrt(x))"
|
| 1862 |
+
]
|
| 1863 |
+
},
|
| 1864 |
+
{
|
| 1865 |
+
"cell_type": "markdown",
|
| 1866 |
+
"metadata": {
|
| 1867 |
+
"id": "a5d_uujuL9it"
|
| 1868 |
+
},
|
| 1869 |
+
"source": [
|
| 1870 |
+
"Note that unlike MATLAB, `*` is elementwise multiplication, not matrix multiplication. We instead use the dot function to compute inner products of vectors, to multiply a vector by a matrix, and to multiply matrices. dot is available both as a function in the numpy module and as an instance method of array objects:"
|
| 1871 |
+
]
|
| 1872 |
+
},
|
| 1873 |
+
{
|
| 1874 |
+
"cell_type": "code",
|
| 1875 |
+
"execution_count": null,
|
| 1876 |
+
"metadata": {
|
| 1877 |
+
"id": "I3FnmoSeL9iu"
|
| 1878 |
+
},
|
| 1879 |
+
"outputs": [],
|
| 1880 |
+
"source": [
|
| 1881 |
+
"x = np.array([[1,2],[3,4]])\n",
|
| 1882 |
+
"y = np.array([[5,6],[7,8]])\n",
|
| 1883 |
+
"\n",
|
| 1884 |
+
"v = np.array([9, 10])\n",
|
| 1885 |
+
"w = np.array([11, 12])\n",
|
| 1886 |
+
"\n",
|
| 1887 |
+
"# Inner product of vectors; both produce 219\n",
|
| 1888 |
+
"print(v.dot(w))\n",
|
| 1889 |
+
"print(np.dot(v, w))"
|
| 1890 |
+
]
|
| 1891 |
+
},
|
| 1892 |
+
{
|
| 1893 |
+
"cell_type": "markdown",
|
| 1894 |
+
"metadata": {
|
| 1895 |
+
"id": "vmxPbrHASVeA"
|
| 1896 |
+
},
|
| 1897 |
+
"source": [
|
| 1898 |
+
"You can also use the `@` operator which is equivalent to numpy's `dot` operator."
|
| 1899 |
+
]
|
| 1900 |
+
},
|
| 1901 |
+
{
|
| 1902 |
+
"cell_type": "code",
|
| 1903 |
+
"execution_count": null,
|
| 1904 |
+
"metadata": {
|
| 1905 |
+
"id": "vyrWA-mXSdtt"
|
| 1906 |
+
},
|
| 1907 |
+
"outputs": [],
|
| 1908 |
+
"source": [
|
| 1909 |
+
"print(v @ w)"
|
| 1910 |
+
]
|
| 1911 |
+
},
|
| 1912 |
+
{
|
| 1913 |
+
"cell_type": "code",
|
| 1914 |
+
"execution_count": null,
|
| 1915 |
+
"metadata": {
|
| 1916 |
+
"id": "zvUODeTxL9iw"
|
| 1917 |
+
},
|
| 1918 |
+
"outputs": [],
|
| 1919 |
+
"source": [
|
| 1920 |
+
"# Matrix / vector product; both produce the rank 1 array [29 67]\n",
|
| 1921 |
+
"print(x.dot(v))\n",
|
| 1922 |
+
"print(np.dot(x, v))\n",
|
| 1923 |
+
"print(x @ v)"
|
| 1924 |
+
]
|
| 1925 |
+
},
|
| 1926 |
+
{
|
| 1927 |
+
"cell_type": "code",
|
| 1928 |
+
"execution_count": null,
|
| 1929 |
+
"metadata": {
|
| 1930 |
+
"id": "3V_3NzNEL9iy"
|
| 1931 |
+
},
|
| 1932 |
+
"outputs": [],
|
| 1933 |
+
"source": [
|
| 1934 |
+
"# Matrix / matrix product; both produce the rank 2 array\n",
|
| 1935 |
+
"# [[19 22]\n",
|
| 1936 |
+
"# [43 50]]\n",
|
| 1937 |
+
"print(x.dot(y))\n",
|
| 1938 |
+
"print(np.dot(x, y))\n",
|
| 1939 |
+
"print(x @ y)"
|
| 1940 |
+
]
|
| 1941 |
+
},
|
| 1942 |
+
{
|
| 1943 |
+
"cell_type": "markdown",
|
| 1944 |
+
"metadata": {
|
| 1945 |
+
"id": "FbE-1If_L9i0"
|
| 1946 |
+
},
|
| 1947 |
+
"source": [
|
| 1948 |
+
"Numpy provides many useful functions for performing computations on arrays; one of the most useful is `sum`:"
|
| 1949 |
+
]
|
| 1950 |
+
},
|
| 1951 |
+
{
|
| 1952 |
+
"cell_type": "code",
|
| 1953 |
+
"execution_count": null,
|
| 1954 |
+
"metadata": {
|
| 1955 |
+
"id": "DZUdZvPrL9i0"
|
| 1956 |
+
},
|
| 1957 |
+
"outputs": [],
|
| 1958 |
+
"source": [
|
| 1959 |
+
"x = np.array([[1,2],[3,4]])\n",
|
| 1960 |
+
"\n",
|
| 1961 |
+
"print(np.sum(x)) # Compute sum of all elements; prints \"10\"\n",
|
| 1962 |
+
"print(np.sum(x, axis=0)) # Compute sum of each column; prints \"[4 6]\"\n",
|
| 1963 |
+
"print(np.sum(x, axis=1)) # Compute sum of each row; prints \"[3 7]\""
|
| 1964 |
+
]
|
| 1965 |
+
},
|
| 1966 |
+
{
|
| 1967 |
+
"cell_type": "markdown",
|
| 1968 |
+
"metadata": {
|
| 1969 |
+
"id": "ahdVW4iUL9i3"
|
| 1970 |
+
},
|
| 1971 |
+
"source": [
|
| 1972 |
+
"You can find the full list of mathematical functions provided by numpy in the [documentation](http://docs.scipy.org/doc/numpy/reference/routines.math.html).\n",
|
| 1973 |
+
"\n",
|
| 1974 |
+
"Apart from computing mathematical functions using arrays, we frequently need to reshape or otherwise manipulate data in arrays. The simplest example of this type of operation is transposing a matrix; to transpose a matrix, simply use the T attribute of an array object:"
|
| 1975 |
+
]
|
| 1976 |
+
},
|
| 1977 |
+
{
|
| 1978 |
+
"cell_type": "code",
|
| 1979 |
+
"execution_count": null,
|
| 1980 |
+
"metadata": {
|
| 1981 |
+
"id": "63Yl1f3oL9i3"
|
| 1982 |
+
},
|
| 1983 |
+
"outputs": [],
|
| 1984 |
+
"source": [
|
| 1985 |
+
"print(x)\n",
|
| 1986 |
+
"print(\"transpose\\n\", x.T)"
|
| 1987 |
+
]
|
| 1988 |
+
},
|
| 1989 |
+
{
|
| 1990 |
+
"cell_type": "code",
|
| 1991 |
+
"execution_count": null,
|
| 1992 |
+
"metadata": {
|
| 1993 |
+
"id": "mkk03eNIL9i4"
|
| 1994 |
+
},
|
| 1995 |
+
"outputs": [],
|
| 1996 |
+
"source": [
|
| 1997 |
+
"v = np.array([[1,2,3]])\n",
|
| 1998 |
+
"print(v )\n",
|
| 1999 |
+
"print(\"transpose\\n\", v.T)"
|
| 2000 |
+
]
|
| 2001 |
+
},
|
| 2002 |
+
{
|
| 2003 |
+
"cell_type": "markdown",
|
| 2004 |
+
"metadata": {
|
| 2005 |
+
"id": "REfLrUTcL9i7"
|
| 2006 |
+
},
|
| 2007 |
+
"source": [
|
| 2008 |
+
"###Broadcasting"
|
| 2009 |
+
]
|
| 2010 |
+
},
|
| 2011 |
+
{
|
| 2012 |
+
"cell_type": "markdown",
|
| 2013 |
+
"metadata": {
|
| 2014 |
+
"id": "EygGAMWqL9i7"
|
| 2015 |
+
},
|
| 2016 |
+
"source": [
|
| 2017 |
+
"Broadcasting is a powerful mechanism that allows numpy to work with arrays of different shapes when performing arithmetic operations. Frequently we have a smaller array and a larger array, and we want to use the smaller array multiple times to perform some operation on the larger array.\n",
|
| 2018 |
+
"\n",
|
| 2019 |
+
"For example, suppose that we want to add a constant vector to each row of a matrix. We could do it like this:"
|
| 2020 |
+
]
|
| 2021 |
+
},
|
| 2022 |
+
{
|
| 2023 |
+
"cell_type": "code",
|
| 2024 |
+
"execution_count": null,
|
| 2025 |
+
"metadata": {
|
| 2026 |
+
"id": "WEEvkV1ZL9i7"
|
| 2027 |
+
},
|
| 2028 |
+
"outputs": [],
|
| 2029 |
+
"source": [
|
| 2030 |
+
"# We will add the vector v to each row of the matrix x,\n",
|
| 2031 |
+
"# storing the result in the matrix y\n",
|
| 2032 |
+
"x = np.array([[1,2,3], [4,5,6], [7,8,9], [10, 11, 12]])\n",
|
| 2033 |
+
"v = np.array([1, 0, 1])\n",
|
| 2034 |
+
"y = np.empty_like(x) # Create an empty matrix with the same shape as x\n",
|
| 2035 |
+
"\n",
|
| 2036 |
+
"# Add the vector v to each row of the matrix x with an explicit loop\n",
|
| 2037 |
+
"for i in range(4):\n",
|
| 2038 |
+
" y[i, :] = x[i, :] + v\n",
|
| 2039 |
+
"\n",
|
| 2040 |
+
"print(y)"
|
| 2041 |
+
]
|
| 2042 |
+
},
|
| 2043 |
+
{
|
| 2044 |
+
"cell_type": "markdown",
|
| 2045 |
+
"metadata": {
|
| 2046 |
+
"id": "2OlXXupEL9i-"
|
| 2047 |
+
},
|
| 2048 |
+
"source": [
|
| 2049 |
+
"This works; however when the matrix `x` is very large, computing an explicit loop in Python could be slow. Note that adding the vector v to each row of the matrix `x` is equivalent to forming a matrix `vv` by stacking multiple copies of `v` vertically, then performing elementwise summation of `x` and `vv`. We could implement this approach like this:"
|
| 2050 |
+
]
|
| 2051 |
+
},
|
| 2052 |
+
{
|
| 2053 |
+
"cell_type": "code",
|
| 2054 |
+
"execution_count": null,
|
| 2055 |
+
"metadata": {
|
| 2056 |
+
"id": "vS7UwAQQL9i-"
|
| 2057 |
+
},
|
| 2058 |
+
"outputs": [],
|
| 2059 |
+
"source": [
|
| 2060 |
+
"vv = np.tile(v, (4, 1)) # Stack 4 copies of v on top of each other\n",
|
| 2061 |
+
"print(vv) # Prints \"[[1 0 1]\n",
|
| 2062 |
+
" # [1 0 1]\n",
|
| 2063 |
+
" # [1 0 1]\n",
|
| 2064 |
+
" # [1 0 1]]\""
|
| 2065 |
+
]
|
| 2066 |
+
},
|
| 2067 |
+
{
|
| 2068 |
+
"cell_type": "code",
|
| 2069 |
+
"execution_count": null,
|
| 2070 |
+
"metadata": {
|
| 2071 |
+
"id": "N0hJphSIL9jA"
|
| 2072 |
+
},
|
| 2073 |
+
"outputs": [],
|
| 2074 |
+
"source": [
|
| 2075 |
+
"y = x + vv # Add x and vv elementwise\n",
|
| 2076 |
+
"print(y)"
|
| 2077 |
+
]
|
| 2078 |
+
},
|
| 2079 |
+
{
|
| 2080 |
+
"cell_type": "markdown",
|
| 2081 |
+
"metadata": {
|
| 2082 |
+
"id": "zHos6RJnL9jB"
|
| 2083 |
+
},
|
| 2084 |
+
"source": [
|
| 2085 |
+
"Numpy broadcasting allows us to perform this computation without actually creating multiple copies of v. Consider this version, using broadcasting:"
|
| 2086 |
+
]
|
| 2087 |
+
},
|
| 2088 |
+
{
|
| 2089 |
+
"cell_type": "code",
|
| 2090 |
+
"execution_count": null,
|
| 2091 |
+
"metadata": {
|
| 2092 |
+
"id": "vnYFb-gYL9jC"
|
| 2093 |
+
},
|
| 2094 |
+
"outputs": [],
|
| 2095 |
+
"source": [
|
| 2096 |
+
"import numpy as np\n",
|
| 2097 |
+
"\n",
|
| 2098 |
+
"# We will add the vector v to each row of the matrix x,\n",
|
| 2099 |
+
"# storing the result in the matrix y\n",
|
| 2100 |
+
"x = np.array([[1,2,3], [4,5,6], [7,8,9], [10, 11, 12]])\n",
|
| 2101 |
+
"v = np.array([1, 0, 1])\n",
|
| 2102 |
+
"y = x + v # Add v to each row of x using broadcasting\n",
|
| 2103 |
+
"print(y)"
|
| 2104 |
+
]
|
| 2105 |
+
},
|
| 2106 |
+
{
|
| 2107 |
+
"cell_type": "markdown",
|
| 2108 |
+
"metadata": {
|
| 2109 |
+
"id": "08YyIURKL9jH"
|
| 2110 |
+
},
|
| 2111 |
+
"source": [
|
| 2112 |
+
"The line `y = x + v` works even though `x` has shape `(4, 3)` and `v` has shape `(3,)` due to broadcasting; this line works as if v actually had shape `(4, 3)`, where each row was a copy of `v`, and the sum was performed elementwise.\n",
|
| 2113 |
+
"\n",
|
| 2114 |
+
"Broadcasting two arrays together follows these rules:\n",
|
| 2115 |
+
"\n",
|
| 2116 |
+
"1. If the arrays do not have the same rank, prepend the shape of the lower rank array with 1s until both shapes have the same length.\n",
|
| 2117 |
+
"2. The two arrays are said to be compatible in a dimension if they have the same size in the dimension, or if one of the arrays has size 1 in that dimension.\n",
|
| 2118 |
+
"3. The arrays can be broadcast together if they are compatible in all dimensions.\n",
|
| 2119 |
+
"4. After broadcasting, each array behaves as if it had shape equal to the elementwise maximum of shapes of the two input arrays.\n",
|
| 2120 |
+
"5. In any dimension where one array had size 1 and the other array had size greater than 1, the first array behaves as if it were copied along that dimension\n",
|
| 2121 |
+
"\n",
|
| 2122 |
+
"If this explanation does not make sense, try reading the explanation from the [documentation](http://docs.scipy.org/doc/numpy/user/basics.broadcasting.html) or this [explanation](http://wiki.scipy.org/EricsBroadcastingDoc).\n",
|
| 2123 |
+
"\n",
|
| 2124 |
+
"Functions that support broadcasting are known as universal functions. You can find the list of all universal functions in the [documentation](http://docs.scipy.org/doc/numpy/reference/ufuncs.html#available-ufuncs).\n",
|
| 2125 |
+
"\n",
|
| 2126 |
+
"Here are some applications of broadcasting:"
|
| 2127 |
+
]
|
| 2128 |
+
},
|
| 2129 |
+
{
|
| 2130 |
+
"cell_type": "code",
|
| 2131 |
+
"execution_count": null,
|
| 2132 |
+
"metadata": {
|
| 2133 |
+
"id": "EmQnwoM9L9jH"
|
| 2134 |
+
},
|
| 2135 |
+
"outputs": [],
|
| 2136 |
+
"source": [
|
| 2137 |
+
"# Compute outer product of vectors\n",
|
| 2138 |
+
"v = np.array([1,2,3]) # v has shape (3,)\n",
|
| 2139 |
+
"w = np.array([4,5]) # w has shape (2,)\n",
|
| 2140 |
+
"# To compute an outer product, we first reshape v to be a column\n",
|
| 2141 |
+
"# vector of shape (3, 1); we can then broadcast it against w to yield\n",
|
| 2142 |
+
"# an output of shape (3, 2), which is the outer product of v and w:\n",
|
| 2143 |
+
"\n",
|
| 2144 |
+
"print(np.reshape(v, (3, 1)) * w)"
|
| 2145 |
+
]
|
| 2146 |
+
},
|
| 2147 |
+
{
|
| 2148 |
+
"cell_type": "code",
|
| 2149 |
+
"execution_count": null,
|
| 2150 |
+
"metadata": {
|
| 2151 |
+
"id": "PgotmpcnL9jK"
|
| 2152 |
+
},
|
| 2153 |
+
"outputs": [],
|
| 2154 |
+
"source": [
|
| 2155 |
+
"# Add a vector to each row of a matrix\n",
|
| 2156 |
+
"x = np.array([[1,2,3], [4,5,6]])\n",
|
| 2157 |
+
"# x has shape (2, 3) and v has shape (3,) so they broadcast to (2, 3),\n",
|
| 2158 |
+
"# giving the following matrix:\n",
|
| 2159 |
+
"\n",
|
| 2160 |
+
"print(x + v)"
|
| 2161 |
+
]
|
| 2162 |
+
},
|
| 2163 |
+
{
|
| 2164 |
+
"cell_type": "code",
|
| 2165 |
+
"execution_count": null,
|
| 2166 |
+
"metadata": {
|
| 2167 |
+
"id": "T5hKS1QaL9jK"
|
| 2168 |
+
},
|
| 2169 |
+
"outputs": [],
|
| 2170 |
+
"source": [
|
| 2171 |
+
"# Add a vector to each column of a matrix\n",
|
| 2172 |
+
"# x has shape (2, 3) and w has shape (2,).\n",
|
| 2173 |
+
"# If we transpose x then it has shape (3, 2) and can be broadcast\n",
|
| 2174 |
+
"# against w to yield a result of shape (3, 2); transposing this result\n",
|
| 2175 |
+
"# yields the final result of shape (2, 3) which is the matrix x with\n",
|
| 2176 |
+
"# the vector w added to each column. Gives the following matrix:\n",
|
| 2177 |
+
"\n",
|
| 2178 |
+
"print((x.T + w).T)"
|
| 2179 |
+
]
|
| 2180 |
+
},
|
| 2181 |
+
{
|
| 2182 |
+
"cell_type": "code",
|
| 2183 |
+
"execution_count": null,
|
| 2184 |
+
"metadata": {
|
| 2185 |
+
"id": "JDUrZUl6L9jN"
|
| 2186 |
+
},
|
| 2187 |
+
"outputs": [],
|
| 2188 |
+
"source": [
|
| 2189 |
+
"# Another solution is to reshape w to be a row vector of shape (2, 1);\n",
|
| 2190 |
+
"# we can then broadcast it directly against x to produce the same\n",
|
| 2191 |
+
"# output.\n",
|
| 2192 |
+
"print(x + np.reshape(w, (2, 1)))"
|
| 2193 |
+
]
|
| 2194 |
+
},
|
| 2195 |
+
{
|
| 2196 |
+
"cell_type": "code",
|
| 2197 |
+
"execution_count": null,
|
| 2198 |
+
"metadata": {
|
| 2199 |
+
"id": "VzrEo4KGL9jP"
|
| 2200 |
+
},
|
| 2201 |
+
"outputs": [],
|
| 2202 |
+
"source": [
|
| 2203 |
+
"# Multiply a matrix by a constant:\n",
|
| 2204 |
+
"# x has shape (2, 3). Numpy treats scalars as arrays of shape ();\n",
|
| 2205 |
+
"# these can be broadcast together to shape (2, 3), producing the\n",
|
| 2206 |
+
"# following array:\n",
|
| 2207 |
+
"print(x * 2)"
|
| 2208 |
+
]
|
| 2209 |
+
},
|
| 2210 |
+
{
|
| 2211 |
+
"cell_type": "markdown",
|
| 2212 |
+
"metadata": {
|
| 2213 |
+
"id": "89e2FXxFL9jQ"
|
| 2214 |
+
},
|
| 2215 |
+
"source": [
|
| 2216 |
+
"Broadcasting typically makes your code more concise and faster, so you should strive to use it where possible."
|
| 2217 |
+
]
|
| 2218 |
+
},
|
| 2219 |
+
{
|
| 2220 |
+
"cell_type": "markdown",
|
| 2221 |
+
"metadata": {
|
| 2222 |
+
"id": "yi90439hpLR0"
|
| 2223 |
+
},
|
| 2224 |
+
"source": [
|
| 2225 |
+
"### Numpy documentation"
|
| 2226 |
+
]
|
| 2227 |
+
},
|
| 2228 |
+
{
|
| 2229 |
+
"cell_type": "markdown",
|
| 2230 |
+
"metadata": {
|
| 2231 |
+
"id": "iF3ZtwVNL9jQ"
|
| 2232 |
+
},
|
| 2233 |
+
"source": [
|
| 2234 |
+
"This brief overview has touched on many of the important things that you need to know about numpy, but is far from complete. Check out the [numpy reference](http://docs.scipy.org/doc/numpy/reference/) to find out much more about numpy."
|
| 2235 |
+
]
|
| 2236 |
+
},
|
| 2237 |
+
{
|
| 2238 |
+
"cell_type": "markdown",
|
| 2239 |
+
"metadata": {
|
| 2240 |
+
"id": "tEINf4bEL9jR"
|
| 2241 |
+
},
|
| 2242 |
+
"source": [
|
| 2243 |
+
"##Matplotlib"
|
| 2244 |
+
]
|
| 2245 |
+
},
|
| 2246 |
+
{
|
| 2247 |
+
"cell_type": "markdown",
|
| 2248 |
+
"metadata": {
|
| 2249 |
+
"id": "0hgVWLaXL9jR"
|
| 2250 |
+
},
|
| 2251 |
+
"source": [
|
| 2252 |
+
"Matplotlib is a plotting library. In this section give a brief introduction to the `matplotlib.pyplot` module, which provides a plotting system similar to that of MATLAB."
|
| 2253 |
+
]
|
| 2254 |
+
},
|
| 2255 |
+
{
|
| 2256 |
+
"cell_type": "code",
|
| 2257 |
+
"execution_count": null,
|
| 2258 |
+
"metadata": {
|
| 2259 |
+
"id": "cmh_7c6KL9jR"
|
| 2260 |
+
},
|
| 2261 |
+
"outputs": [],
|
| 2262 |
+
"source": [
|
| 2263 |
+
"import matplotlib.pyplot as plt"
|
| 2264 |
+
]
|
| 2265 |
+
},
|
| 2266 |
+
{
|
| 2267 |
+
"cell_type": "markdown",
|
| 2268 |
+
"metadata": {
|
| 2269 |
+
"id": "jOsaA5hGL9jS"
|
| 2270 |
+
},
|
| 2271 |
+
"source": [
|
| 2272 |
+
"By running this special iPython command, we will be displaying plots inline:"
|
| 2273 |
+
]
|
| 2274 |
+
},
|
| 2275 |
+
{
|
| 2276 |
+
"cell_type": "code",
|
| 2277 |
+
"execution_count": null,
|
| 2278 |
+
"metadata": {
|
| 2279 |
+
"id": "ijpsmwGnL9jT"
|
| 2280 |
+
},
|
| 2281 |
+
"outputs": [],
|
| 2282 |
+
"source": [
|
| 2283 |
+
"%matplotlib inline"
|
| 2284 |
+
]
|
| 2285 |
+
},
|
| 2286 |
+
{
|
| 2287 |
+
"cell_type": "markdown",
|
| 2288 |
+
"metadata": {
|
| 2289 |
+
"id": "U5Z_oMoLL9jV"
|
| 2290 |
+
},
|
| 2291 |
+
"source": [
|
| 2292 |
+
"###Plotting"
|
| 2293 |
+
]
|
| 2294 |
+
},
|
| 2295 |
+
{
|
| 2296 |
+
"cell_type": "markdown",
|
| 2297 |
+
"metadata": {
|
| 2298 |
+
"id": "6QyFJ7dhL9jV"
|
| 2299 |
+
},
|
| 2300 |
+
"source": [
|
| 2301 |
+
"The most important function in `matplotlib` is plot, which allows you to plot 2D data. Here is a simple example:"
|
| 2302 |
+
]
|
| 2303 |
+
},
|
| 2304 |
+
{
|
| 2305 |
+
"cell_type": "code",
|
| 2306 |
+
"execution_count": null,
|
| 2307 |
+
"metadata": {
|
| 2308 |
+
"id": "pua52BGeL9jW"
|
| 2309 |
+
},
|
| 2310 |
+
"outputs": [],
|
| 2311 |
+
"source": [
|
| 2312 |
+
"# Compute the x and y coordinates for points on a sine curve\n",
|
| 2313 |
+
"x = np.arange(0, 3 * np.pi, 0.1)\n",
|
| 2314 |
+
"y = np.sin(x)\n",
|
| 2315 |
+
"\n",
|
| 2316 |
+
"# Plot the points using matplotlib\n",
|
| 2317 |
+
"plt.plot(x, y)"
|
| 2318 |
+
]
|
| 2319 |
+
},
|
| 2320 |
+
{
|
| 2321 |
+
"cell_type": "markdown",
|
| 2322 |
+
"metadata": {
|
| 2323 |
+
"id": "9W2VAcLiL9jX"
|
| 2324 |
+
},
|
| 2325 |
+
"source": [
|
| 2326 |
+
"With just a little bit of extra work we can easily plot multiple lines at once, and add a title, legend, and axis labels:"
|
| 2327 |
+
]
|
| 2328 |
+
},
|
| 2329 |
+
{
|
| 2330 |
+
"cell_type": "code",
|
| 2331 |
+
"execution_count": null,
|
| 2332 |
+
"metadata": {
|
| 2333 |
+
"id": "TfCQHJ5AL9jY"
|
| 2334 |
+
},
|
| 2335 |
+
"outputs": [],
|
| 2336 |
+
"source": [
|
| 2337 |
+
"y_sin = np.sin(x)\n",
|
| 2338 |
+
"y_cos = np.cos(x)\n",
|
| 2339 |
+
"\n",
|
| 2340 |
+
"# Plot the points using matplotlib\n",
|
| 2341 |
+
"plt.plot(x, y_sin)\n",
|
| 2342 |
+
"plt.plot(x, y_cos)\n",
|
| 2343 |
+
"plt.xlabel('x axis label')\n",
|
| 2344 |
+
"plt.ylabel('y axis label')\n",
|
| 2345 |
+
"plt.title('Sine and Cosine')\n",
|
| 2346 |
+
"plt.legend(['Sine', 'Cosine'])"
|
| 2347 |
+
]
|
| 2348 |
+
},
|
| 2349 |
+
{
|
| 2350 |
+
"cell_type": "markdown",
|
| 2351 |
+
"metadata": {
|
| 2352 |
+
"id": "R5IeAY03L9ja"
|
| 2353 |
+
},
|
| 2354 |
+
"source": [
|
| 2355 |
+
"###Subplots"
|
| 2356 |
+
]
|
| 2357 |
+
},
|
| 2358 |
+
{
|
| 2359 |
+
"cell_type": "markdown",
|
| 2360 |
+
"metadata": {
|
| 2361 |
+
"id": "CfUzwJg0L9ja"
|
| 2362 |
+
},
|
| 2363 |
+
"source": [
|
| 2364 |
+
"You can plot different things in the same figure using the subplot function. Here is an example:"
|
| 2365 |
+
]
|
| 2366 |
+
},
|
| 2367 |
+
{
|
| 2368 |
+
"cell_type": "code",
|
| 2369 |
+
"execution_count": null,
|
| 2370 |
+
"metadata": {
|
| 2371 |
+
"id": "dM23yGH9L9ja"
|
| 2372 |
+
},
|
| 2373 |
+
"outputs": [],
|
| 2374 |
+
"source": [
|
| 2375 |
+
"# Compute the x and y coordinates for points on sine and cosine curves\n",
|
| 2376 |
+
"x = np.arange(0, 3 * np.pi, 0.1)\n",
|
| 2377 |
+
"y_sin = np.sin(x)\n",
|
| 2378 |
+
"y_cos = np.cos(x)\n",
|
| 2379 |
+
"\n",
|
| 2380 |
+
"# Set up a subplot grid that has height 2 and width 1,\n",
|
| 2381 |
+
"# and set the first such subplot as active.\n",
|
| 2382 |
+
"plt.subplot(2, 1, 1)\n",
|
| 2383 |
+
"\n",
|
| 2384 |
+
"# Make the first plot\n",
|
| 2385 |
+
"plt.plot(x, y_sin)\n",
|
| 2386 |
+
"plt.title('Sine')\n",
|
| 2387 |
+
"\n",
|
| 2388 |
+
"# Set the second subplot as active, and make the second plot.\n",
|
| 2389 |
+
"plt.subplot(2, 1, 2)\n",
|
| 2390 |
+
"plt.plot(x, y_cos)\n",
|
| 2391 |
+
"plt.title('Cosine')\n",
|
| 2392 |
+
"\n",
|
| 2393 |
+
"# Show the figure.\n",
|
| 2394 |
+
"plt.show()"
|
| 2395 |
+
]
|
| 2396 |
+
},
|
| 2397 |
+
{
|
| 2398 |
+
"cell_type": "markdown",
|
| 2399 |
+
"metadata": {
|
| 2400 |
+
"id": "gLtsST5SL9jc"
|
| 2401 |
+
},
|
| 2402 |
+
"source": [
|
| 2403 |
+
"You can read much more about the `subplot` function in the [documentation](http://matplotlib.org/api/pyplot_api.html#matplotlib.pyplot.subplot)."
|
| 2404 |
+
]
|
| 2405 |
+
},
|
| 2406 |
+
{
|
| 2407 |
+
"cell_type": "markdown",
|
| 2408 |
+
"metadata": {
|
| 2409 |
+
"id": "7Zqndtogsq8J"
|
| 2410 |
+
},
|
| 2411 |
+
"source": [
|
| 2412 |
+
"### Download images"
|
| 2413 |
+
]
|
| 2414 |
+
},
|
| 2415 |
+
{
|
| 2416 |
+
"cell_type": "markdown",
|
| 2417 |
+
"metadata": {
|
| 2418 |
+
"id": "4FFozuUm7OE5"
|
| 2419 |
+
},
|
| 2420 |
+
"source": [
|
| 2421 |
+
"Lets download some images."
|
| 2422 |
+
]
|
| 2423 |
+
},
|
| 2424 |
+
{
|
| 2425 |
+
"cell_type": "code",
|
| 2426 |
+
"execution_count": null,
|
| 2427 |
+
"metadata": {
|
| 2428 |
+
"id": "cOEqkMl47VTP"
|
| 2429 |
+
},
|
| 2430 |
+
"outputs": [],
|
| 2431 |
+
"source": [
|
| 2432 |
+
"import os\n",
|
| 2433 |
+
"if not os.path.exists('images.zip'):\n",
|
| 2434 |
+
" !wget --no-check-certificate https://empslocal.ex.ac.uk/people/staff/ad735/ECMM426/images.zip\n",
|
| 2435 |
+
" !unzip -q images.zip"
|
| 2436 |
+
]
|
| 2437 |
+
},
|
| 2438 |
+
{
|
| 2439 |
+
"cell_type": "markdown",
|
| 2440 |
+
"metadata": {
|
| 2441 |
+
"id": "cjKpKTNYsxcF"
|
| 2442 |
+
},
|
| 2443 |
+
"source": [
|
| 2444 |
+
"You can use the `imread` and the`imshow` function to respectively read and show images. Here is an example:"
|
| 2445 |
+
]
|
| 2446 |
+
},
|
| 2447 |
+
{
|
| 2448 |
+
"cell_type": "code",
|
| 2449 |
+
"execution_count": null,
|
| 2450 |
+
"metadata": {
|
| 2451 |
+
"id": "LEep7wnTs5Si"
|
| 2452 |
+
},
|
| 2453 |
+
"outputs": [],
|
| 2454 |
+
"source": [
|
| 2455 |
+
"import numpy as np\n",
|
| 2456 |
+
"import matplotlib.pyplot as plt\n",
|
| 2457 |
+
"\n",
|
| 2458 |
+
"img = plt.imread('images/lena.png')\n",
|
| 2459 |
+
"img_tinted = img * [1, 0.85, 0.8]\n",
|
| 2460 |
+
"\n",
|
| 2461 |
+
"# Show the original image\n",
|
| 2462 |
+
"plt.subplot(1, 2, 1)\n",
|
| 2463 |
+
"plt.imshow(img)\n",
|
| 2464 |
+
"\n",
|
| 2465 |
+
"# Show the tinted image\n",
|
| 2466 |
+
"plt.subplot(1, 2, 2)\n",
|
| 2467 |
+
"plt.imshow(img_tinted)\n",
|
| 2468 |
+
"plt.show()"
|
| 2469 |
+
]
|
| 2470 |
+
},
|
| 2471 |
+
{
|
| 2472 |
+
"cell_type": "markdown",
|
| 2473 |
+
"metadata": {
|
| 2474 |
+
"id": "wOan27So8lpI"
|
| 2475 |
+
},
|
| 2476 |
+
"source": [
|
| 2477 |
+
"## Scikit-learn"
|
| 2478 |
+
]
|
| 2479 |
+
},
|
| 2480 |
+
{
|
| 2481 |
+
"cell_type": "markdown",
|
| 2482 |
+
"metadata": {
|
| 2483 |
+
"id": "uSki-HDorHgE"
|
| 2484 |
+
},
|
| 2485 |
+
"source": [
|
| 2486 |
+
"[Scikit-learn](https://scikit-learn.org/stable/) is an open source machine learning library that supports supervised and unsupervised learning. It also provides various tools for model fitting, data preprocessing, model selection, model evaluation, and many other utilities."
|
| 2487 |
+
]
|
| 2488 |
+
},
|
| 2489 |
+
{
|
| 2490 |
+
"cell_type": "markdown",
|
| 2491 |
+
"metadata": {
|
| 2492 |
+
"id": "BOwTaL8x2QJ1"
|
| 2493 |
+
},
|
| 2494 |
+
"source": [
|
| 2495 |
+
"### Moon dataset\n",
|
| 2496 |
+
"Below we will consider a toy dataset, such as moon dataset and consider some classifiers from the Scikit-learn library to classify them."
|
| 2497 |
+
]
|
| 2498 |
+
},
|
| 2499 |
+
{
|
| 2500 |
+
"cell_type": "code",
|
| 2501 |
+
"execution_count": null,
|
| 2502 |
+
"metadata": {
|
| 2503 |
+
"id": "-wudgYDp1xO_"
|
| 2504 |
+
},
|
| 2505 |
+
"outputs": [],
|
| 2506 |
+
"source": [
|
| 2507 |
+
"# Create the moon dataset and plot\n",
|
| 2508 |
+
"from sklearn.datasets import make_moons\n",
|
| 2509 |
+
"\n",
|
| 2510 |
+
"X, y = make_moons(n_samples=500, noise=0.30, random_state=42)\n",
|
| 2511 |
+
"\n",
|
| 2512 |
+
"id0 = y == 0\n",
|
| 2513 |
+
"id1 = y == 1\n",
|
| 2514 |
+
"plt.plot(X[id0, 0], X[id0, 1], 'bo', label='0')\n",
|
| 2515 |
+
"plt.plot(X[id1, 0], X[id1, 1], 'ro', label='1')\n",
|
| 2516 |
+
"plt.legend(loc=2)"
|
| 2517 |
+
]
|
| 2518 |
+
},
|
| 2519 |
+
{
|
| 2520 |
+
"cell_type": "markdown",
|
| 2521 |
+
"metadata": {
|
| 2522 |
+
"id": "XbIu90Rg2TvB"
|
| 2523 |
+
},
|
| 2524 |
+
"source": [
|
| 2525 |
+
"### Dataset split"
|
| 2526 |
+
]
|
| 2527 |
+
},
|
| 2528 |
+
{
|
| 2529 |
+
"cell_type": "code",
|
| 2530 |
+
"execution_count": null,
|
| 2531 |
+
"metadata": {
|
| 2532 |
+
"id": "QVDMPmL62EIr"
|
| 2533 |
+
},
|
| 2534 |
+
"outputs": [],
|
| 2535 |
+
"source": [
|
| 2536 |
+
"# Split into train and test sets\n",
|
| 2537 |
+
"from sklearn.model_selection import train_test_split\n",
|
| 2538 |
+
"X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)"
|
| 2539 |
+
]
|
| 2540 |
+
},
|
| 2541 |
+
{
|
| 2542 |
+
"cell_type": "markdown",
|
| 2543 |
+
"metadata": {
|
| 2544 |
+
"id": "uZNqoqgW31rp"
|
| 2545 |
+
},
|
| 2546 |
+
"source": [
|
| 2547 |
+
"### Random forest classifier\n",
|
| 2548 |
+
"Lets now train a random forest classifier from the scikit-learn library on the above training set and test it on the test set, and the compute the classification accuracy. Please check the [documentation](https://scikit-learn.org/stable/modules/generated/sklearn.ensemble.RandomForestClassifier.html#sklearn-ensemble-randomforestclassifier) of `RandomForestClassifier` for more details on its parameters."
|
| 2549 |
+
]
|
| 2550 |
+
},
|
| 2551 |
+
{
|
| 2552 |
+
"cell_type": "code",
|
| 2553 |
+
"execution_count": null,
|
| 2554 |
+
"metadata": {
|
| 2555 |
+
"id": "RC8sl3QE5BjD"
|
| 2556 |
+
},
|
| 2557 |
+
"outputs": [],
|
| 2558 |
+
"source": [
|
| 2559 |
+
"from sklearn.ensemble import RandomForestClassifier\n",
|
| 2560 |
+
"# Define the classifier\n",
|
| 2561 |
+
"rnd_clf = RandomForestClassifier(n_estimators=500, max_leaf_nodes=16, n_jobs=-1, random_state=42)\n",
|
| 2562 |
+
"# Training\n",
|
| 2563 |
+
"rnd_clf.fit(X_train, y_train)\n",
|
| 2564 |
+
"# Test\n",
|
| 2565 |
+
"y_pred_rf = rnd_clf.predict(X_test)\n",
|
| 2566 |
+
"# Classification accuracy\n",
|
| 2567 |
+
"from sklearn.metrics import accuracy_score\n",
|
| 2568 |
+
"print(accuracy_score(y_test, y_pred_rf))"
|
| 2569 |
+
]
|
| 2570 |
+
},
|
| 2571 |
+
{
|
| 2572 |
+
"cell_type": "markdown",
|
| 2573 |
+
"metadata": {
|
| 2574 |
+
"id": "T8JLXEV97i8c"
|
| 2575 |
+
},
|
| 2576 |
+
"source": [
|
| 2577 |
+
"### Non-linear Support Vector Machine\n",
|
| 2578 |
+
"Now lets do the same training and testing with a non-linear [Support Vector Machine (SVM)](https://scikit-learn.org/stable/modules/generated/sklearn.svm.SVC.html) classifier."
|
| 2579 |
+
]
|
| 2580 |
+
},
|
| 2581 |
+
{
|
| 2582 |
+
"cell_type": "code",
|
| 2583 |
+
"execution_count": null,
|
| 2584 |
+
"metadata": {
|
| 2585 |
+
"id": "av46bVVs8w08"
|
| 2586 |
+
},
|
| 2587 |
+
"outputs": [],
|
| 2588 |
+
"source": [
|
| 2589 |
+
"from sklearn.svm import SVC\n",
|
| 2590 |
+
"# Define the non-linear classifier with radial basis function (rbf) kernel\n",
|
| 2591 |
+
"nlin_svm_clf_1 = SVC(kernel=\"rbf\")\n",
|
| 2592 |
+
"# Training\n",
|
| 2593 |
+
"nlin_svm_clf_1.fit(X_train, y_train)\n",
|
| 2594 |
+
"# Test\n",
|
| 2595 |
+
"y_pred = nlin_svm_clf_1.predict(X_test)\n",
|
| 2596 |
+
"# Classification accuracy\n",
|
| 2597 |
+
"from sklearn.metrics import accuracy_score\n",
|
| 2598 |
+
"print(accuracy_score(y_pred, y_test))"
|
| 2599 |
+
]
|
| 2600 |
+
},
|
| 2601 |
+
{
|
| 2602 |
+
"cell_type": "markdown",
|
| 2603 |
+
"metadata": {
|
| 2604 |
+
"id": "74nKsfA99pkb"
|
| 2605 |
+
},
|
| 2606 |
+
"source": [
|
| 2607 |
+
"### Confusion matrix\n",
|
| 2608 |
+
"A confusion matrix is a table that is used to define the performance of a classification algorithm. A confusion matrix visualizes and summarizes the performance of a classification algorithm. More details on how to compute confusion matrix can be found in the [documentation](https://scikit-learn.org/stable/modules/generated/sklearn.metrics.confusion_matrix.html)."
|
| 2609 |
+
]
|
| 2610 |
+
},
|
| 2611 |
+
{
|
| 2612 |
+
"cell_type": "code",
|
| 2613 |
+
"execution_count": null,
|
| 2614 |
+
"metadata": {
|
| 2615 |
+
"id": "aB73WGra-aEF"
|
| 2616 |
+
},
|
| 2617 |
+
"outputs": [],
|
| 2618 |
+
"source": [
|
| 2619 |
+
"from sklearn.metrics import confusion_matrix\n",
|
| 2620 |
+
"confusion_matrix(y_test, y_pred)"
|
| 2621 |
+
]
|
| 2622 |
+
},
|
| 2623 |
+
{
|
| 2624 |
+
"cell_type": "markdown",
|
| 2625 |
+
"metadata": {
|
| 2626 |
+
"id": "Wp2DsEemAOub"
|
| 2627 |
+
},
|
| 2628 |
+
"source": [
|
| 2629 |
+
"### Regression\n",
|
| 2630 |
+
"\n",
|
| 2631 |
+
"Now, lets consider the following function and train an [MLP regressor](https://scikit-learn.org/stable/modules/generated/sklearn.neural_network.MLPRegressor.html) to learn it.\n",
|
| 2632 |
+
"\n",
|
| 2633 |
+
"\\begin{equation}\n",
|
| 2634 |
+
"y = f(x; \\mathbf{w}) = 5x^2 + 3\n",
|
| 2635 |
+
"\\end{equation}"
|
| 2636 |
+
]
|
| 2637 |
+
},
|
| 2638 |
+
{
|
| 2639 |
+
"cell_type": "code",
|
| 2640 |
+
"execution_count": null,
|
| 2641 |
+
"metadata": {
|
| 2642 |
+
"id": "C5Dzr0MLA3h8"
|
| 2643 |
+
},
|
| 2644 |
+
"outputs": [],
|
| 2645 |
+
"source": [
|
| 2646 |
+
"# Create the data that follow uniform distribution\n",
|
| 2647 |
+
"import numpy as np\n",
|
| 2648 |
+
"import matplotlib.pyplot as plt\n",
|
| 2649 |
+
"X = np.random.uniform(-100, 100, 1000)\n",
|
| 2650 |
+
"y = 5*(X*X) + 3\n",
|
| 2651 |
+
"plt.scatter(X, y, s=10);"
|
| 2652 |
+
]
|
| 2653 |
+
},
|
| 2654 |
+
{
|
| 2655 |
+
"cell_type": "code",
|
| 2656 |
+
"execution_count": null,
|
| 2657 |
+
"metadata": {
|
| 2658 |
+
"id": "Lr2qn7zxCKcK"
|
| 2659 |
+
},
|
| 2660 |
+
"outputs": [],
|
| 2661 |
+
"source": [
|
| 2662 |
+
"# Split the dataset into train and test sets\n",
|
| 2663 |
+
"from sklearn.model_selection import train_test_split\n",
|
| 2664 |
+
"X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)"
|
| 2665 |
+
]
|
| 2666 |
+
},
|
| 2667 |
+
{
|
| 2668 |
+
"cell_type": "code",
|
| 2669 |
+
"execution_count": null,
|
| 2670 |
+
"metadata": {
|
| 2671 |
+
"id": "UXpmCrg0CbeA"
|
| 2672 |
+
},
|
| 2673 |
+
"outputs": [],
|
| 2674 |
+
"source": [
|
| 2675 |
+
"from sklearn.neural_network import MLPRegressor\n",
|
| 2676 |
+
"# Define an MLPRegressor\n",
|
| 2677 |
+
"regr = MLPRegressor(hidden_layer_sizes=(10,), solver='lbfgs', activation='relu', max_iter=10000)\n",
|
| 2678 |
+
"# Fit on the training data\n",
|
| 2679 |
+
"regr = regr.fit(X_train.reshape(-1, 1), y_train)\n",
|
| 2680 |
+
"# Predict using the multi-layer perceptron model\n",
|
| 2681 |
+
"y_pred = regr.predict(X_test.reshape(-1, 1))\n",
|
| 2682 |
+
"# Return the coefficient of determination of the prediction. The best score can be 1.\n",
|
| 2683 |
+
"regr.score(X_test.reshape(-1, 1), y_test)"
|
| 2684 |
+
]
|
| 2685 |
+
},
|
| 2686 |
+
{
|
| 2687 |
+
"cell_type": "markdown",
|
| 2688 |
+
"metadata": {
|
| 2689 |
+
"id": "vxAwaKWhq6z5"
|
| 2690 |
+
},
|
| 2691 |
+
"source": [
|
| 2692 |
+
"## OpenCV"
|
| 2693 |
+
]
|
| 2694 |
+
},
|
| 2695 |
+
{
|
| 2696 |
+
"cell_type": "markdown",
|
| 2697 |
+
"metadata": {
|
| 2698 |
+
"id": "7S04s4KFz3EO"
|
| 2699 |
+
},
|
| 2700 |
+
"source": [
|
| 2701 |
+
"OpenCV is a library providing implementation of multitude of algorithms related to image processing, computer vision and machine learning. In this section, we will learn different image processing functions from the OpenCV library. For more details on OpenCV, please see the [OpenCV website](https://opencv.org/)."
|
| 2702 |
+
]
|
| 2703 |
+
},
|
| 2704 |
+
{
|
| 2705 |
+
"cell_type": "markdown",
|
| 2706 |
+
"metadata": {
|
| 2707 |
+
"id": "oTcZ703E2Zxh"
|
| 2708 |
+
},
|
| 2709 |
+
"source": [
|
| 2710 |
+
"### Data structures\n",
|
| 2711 |
+
"\n",
|
| 2712 |
+
"Colour images usually have three channels: red, green and blue and these channels are usually arranged in a certain order. Depending on this arrangement the image is termed in a certain way. For example, if the channels in an image are ordered in red (R), green (G) and blue (B), the image is called as RGB image. In OpenCV an image can be read by `cv2.imread()` function."
|
| 2713 |
+
]
|
| 2714 |
+
},
|
| 2715 |
+
{
|
| 2716 |
+
"cell_type": "code",
|
| 2717 |
+
"execution_count": null,
|
| 2718 |
+
"metadata": {
|
| 2719 |
+
"id": "W_6NRQ762_fP"
|
| 2720 |
+
},
|
| 2721 |
+
"outputs": [],
|
| 2722 |
+
"source": [
|
| 2723 |
+
"# read an image\n",
|
| 2724 |
+
"import cv2\n",
|
| 2725 |
+
"img = cv2.imread('images/lena.png')\n",
|
| 2726 |
+
"\n",
|
| 2727 |
+
"# show image format (basically a 3-d array of pixel colour info, in BGR format)\n",
|
| 2728 |
+
"print('Image shape: {}'.format(img.shape))\n",
|
| 2729 |
+
"print('Image: {}'.format(img))"
|
| 2730 |
+
]
|
| 2731 |
+
},
|
| 2732 |
+
{
|
| 2733 |
+
"cell_type": "markdown",
|
| 2734 |
+
"metadata": {
|
| 2735 |
+
"id": "aZPEVuGh_p8V",
|
| 2736 |
+
"pycharm": {}
|
| 2737 |
+
},
|
| 2738 |
+
"source": [
|
| 2739 |
+
"### Colour conversions\n",
|
| 2740 |
+
"By default, OpenCV loads images in BGR format. This is why the famous image of Lena looks a bit weird. **Note:** we will use imshow function from Matplotlib to display the image."
|
| 2741 |
+
]
|
| 2742 |
+
},
|
| 2743 |
+
{
|
| 2744 |
+
"cell_type": "code",
|
| 2745 |
+
"execution_count": null,
|
| 2746 |
+
"metadata": {
|
| 2747 |
+
"id": "uPBA122WEzXM",
|
| 2748 |
+
"pycharm": {}
|
| 2749 |
+
},
|
| 2750 |
+
"outputs": [],
|
| 2751 |
+
"source": [
|
| 2752 |
+
"# show image with matplotlib\n",
|
| 2753 |
+
"import matplotlib.pyplot as plt\n",
|
| 2754 |
+
"plt.imshow(img)"
|
| 2755 |
+
]
|
| 2756 |
+
},
|
| 2757 |
+
{
|
| 2758 |
+
"cell_type": "markdown",
|
| 2759 |
+
"metadata": {
|
| 2760 |
+
"id": "T89Km6qqQs3t"
|
| 2761 |
+
},
|
| 2762 |
+
"source": [
|
| 2763 |
+
"In OpenCV, a BGR image can be converted to an RGB image by the `cv2.cvtColor()` function as follows"
|
| 2764 |
+
]
|
| 2765 |
+
},
|
| 2766 |
+
{
|
| 2767 |
+
"cell_type": "code",
|
| 2768 |
+
"execution_count": null,
|
| 2769 |
+
"metadata": {
|
| 2770 |
+
"id": "_kIGhwKc_p8V",
|
| 2771 |
+
"pycharm": {},
|
| 2772 |
+
"scrolled": true
|
| 2773 |
+
},
|
| 2774 |
+
"outputs": [],
|
| 2775 |
+
"source": [
|
| 2776 |
+
"# convert image to RGB colour space\n",
|
| 2777 |
+
"img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n",
|
| 2778 |
+
"\n",
|
| 2779 |
+
"# show image with matplotlib\n",
|
| 2780 |
+
"plt.imshow(img)"
|
| 2781 |
+
]
|
| 2782 |
+
},
|
| 2783 |
+
{
|
| 2784 |
+
"cell_type": "markdown",
|
| 2785 |
+
"metadata": {
|
| 2786 |
+
"id": "8WYalGWXQs3u"
|
| 2787 |
+
},
|
| 2788 |
+
"source": [
|
| 2789 |
+
"In a similar way, a BGR image can also be converted to grayscale image which has only a single channel. Converting an RGB image into a grayscale image involves summing up the individual (RGB) components with the weights (0.299, 0.587, 0.114). The OpenCV function `cv2.cvtColor()` can also be used to convert an RGB image into a grayscale image."
|
| 2790 |
+
]
|
| 2791 |
+
},
|
| 2792 |
+
{
|
| 2793 |
+
"cell_type": "code",
|
| 2794 |
+
"execution_count": null,
|
| 2795 |
+
"metadata": {
|
| 2796 |
+
"id": "vpS6RcOV_p8Y",
|
| 2797 |
+
"pycharm": {}
|
| 2798 |
+
},
|
| 2799 |
+
"outputs": [],
|
| 2800 |
+
"source": [
|
| 2801 |
+
"# convert image to grayscale\n",
|
| 2802 |
+
"gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n",
|
| 2803 |
+
"\n",
|
| 2804 |
+
"print('Image shape: {}'.format(gray_img.shape))\n",
|
| 2805 |
+
"# grayscale image represented as a 2-d array\n",
|
| 2806 |
+
"print(gray_img)"
|
| 2807 |
+
]
|
| 2808 |
+
},
|
| 2809 |
+
{
|
| 2810 |
+
"cell_type": "markdown",
|
| 2811 |
+
"metadata": {
|
| 2812 |
+
"id": "E1yoDp2hCihi",
|
| 2813 |
+
"pycharm": {}
|
| 2814 |
+
},
|
| 2815 |
+
"source": [
|
| 2816 |
+
"Gray images have single channel"
|
| 2817 |
+
]
|
| 2818 |
+
},
|
| 2819 |
+
{
|
| 2820 |
+
"cell_type": "code",
|
| 2821 |
+
"execution_count": null,
|
| 2822 |
+
"metadata": {
|
| 2823 |
+
"id": "z2-K1pOh_p8a",
|
| 2824 |
+
"pycharm": {}
|
| 2825 |
+
},
|
| 2826 |
+
"outputs": [],
|
| 2827 |
+
"source": [
|
| 2828 |
+
"# plot the gray image, note the cmap parameter\n",
|
| 2829 |
+
"plt.imshow(gray_img, cmap='gray')"
|
| 2830 |
+
]
|
| 2831 |
+
},
|
| 2832 |
+
{
|
| 2833 |
+
"cell_type": "markdown",
|
| 2834 |
+
"metadata": {
|
| 2835 |
+
"id": "B_EOqF8qQs3v"
|
| 2836 |
+
},
|
| 2837 |
+
"source": [
|
| 2838 |
+
"Colour to grayscale is a lossy conversion. However, in OpenCV a grayscale image can approximately be converted to a colour image using the `cv2.applyColorMap()` function according to the colour maps described at this [link](https://docs.opencv.org/4.x/d3/d50/group__imgproc__colormap.html)."
|
| 2839 |
+
]
|
| 2840 |
+
},
|
| 2841 |
+
{
|
| 2842 |
+
"cell_type": "code",
|
| 2843 |
+
"execution_count": null,
|
| 2844 |
+
"metadata": {
|
| 2845 |
+
"id": "eOs27PQFSYrn",
|
| 2846 |
+
"pycharm": {}
|
| 2847 |
+
},
|
| 2848 |
+
"outputs": [],
|
| 2849 |
+
"source": [
|
| 2850 |
+
"gray_img_col = cv2.applyColorMap(gray_img, cv2.COLORMAP_JET)\n",
|
| 2851 |
+
"plt.imshow(gray_img_col)"
|
| 2852 |
+
]
|
| 2853 |
+
},
|
| 2854 |
+
{
|
| 2855 |
+
"cell_type": "markdown",
|
| 2856 |
+
"metadata": {
|
| 2857 |
+
"id": "cE_h8sgbQs3w"
|
| 2858 |
+
},
|
| 2859 |
+
"source": [
|
| 2860 |
+
"### Conversion from `uint8` to `float64` (`double`) and Normalization"
|
| 2861 |
+
]
|
| 2862 |
+
},
|
| 2863 |
+
{
|
| 2864 |
+
"cell_type": "code",
|
| 2865 |
+
"execution_count": null,
|
| 2866 |
+
"metadata": {
|
| 2867 |
+
"id": "-6aepcnKQs3w",
|
| 2868 |
+
"pycharm": {
|
| 2869 |
+
"name": "#%%\n"
|
| 2870 |
+
}
|
| 2871 |
+
},
|
| 2872 |
+
"outputs": [],
|
| 2873 |
+
"source": [
|
| 2874 |
+
"img_dble = cv2.normalize(img.astype('float64'), None, 0.0, 1.0, cv2.NORM_MINMAX)\n",
|
| 2875 |
+
"print(img_dble)"
|
| 2876 |
+
]
|
| 2877 |
+
},
|
| 2878 |
+
{
|
| 2879 |
+
"cell_type": "code",
|
| 2880 |
+
"execution_count": null,
|
| 2881 |
+
"metadata": {
|
| 2882 |
+
"id": "tOh1DQ5s5Iyg"
|
| 2883 |
+
},
|
| 2884 |
+
"outputs": [],
|
| 2885 |
+
"source": [
|
| 2886 |
+
"plt.imshow(img_dble)"
|
| 2887 |
+
]
|
| 2888 |
+
},
|
| 2889 |
+
{
|
| 2890 |
+
"cell_type": "markdown",
|
| 2891 |
+
"metadata": {
|
| 2892 |
+
"id": "PxNoVeWkwqDF"
|
| 2893 |
+
},
|
| 2894 |
+
"source": [
|
| 2895 |
+
"### Image processing\n",
|
| 2896 |
+
"Below we will review some brief image processing tasks, such as image filtering, binarization, edge detection etc with OpenCV."
|
| 2897 |
+
]
|
| 2898 |
+
},
|
| 2899 |
+
{
|
| 2900 |
+
"cell_type": "markdown",
|
| 2901 |
+
"metadata": {
|
| 2902 |
+
"id": "iIt6E6uC5aMS"
|
| 2903 |
+
},
|
| 2904 |
+
"source": [
|
| 2905 |
+
"#### Box Filtering"
|
| 2906 |
+
]
|
| 2907 |
+
},
|
| 2908 |
+
{
|
| 2909 |
+
"cell_type": "markdown",
|
| 2910 |
+
"metadata": {
|
| 2911 |
+
"id": "dAvdTtHeKQJJ",
|
| 2912 |
+
"pycharm": {}
|
| 2913 |
+
},
|
| 2914 |
+
"source": [
|
| 2915 |
+
"In this filtering, each pixel value in an image is replaced by the weighted average of the neighborhood (defined by the filter mask) intensity values. The most commonly used filter is the Box filter which has equal weights. A 3×3 normalized box filter is shown below\n",
|
| 2916 |
+
"\n",
|
| 2917 |
+
"\n",
|
| 2918 |
+
"\n",
|
| 2919 |
+
"It is a good practice to normalize the filter, this is why the above filter is divided by 9. This is to make sure that the image does not get brighter or darker. You can also use an unnormalized box filter.\n",
|
| 2920 |
+
"\n",
|
| 2921 |
+
"OpenCV provides two inbuilt functions for averaging namely:\n",
|
| 2922 |
+
"\n",
|
| 2923 |
+
"* `cv2.blur()` that blurs an image using only the normalized box filter and\n",
|
| 2924 |
+
"* `cv2.boxFilter()` which is more general, having the option of using either normalized or unnormalized box filter. Just pass an argument normalize=False to the function"
|
| 2925 |
+
]
|
| 2926 |
+
},
|
| 2927 |
+
{
|
| 2928 |
+
"cell_type": "code",
|
| 2929 |
+
"execution_count": null,
|
| 2930 |
+
"metadata": {
|
| 2931 |
+
"id": "HVb5aZSNLaGK",
|
| 2932 |
+
"pycharm": {}
|
| 2933 |
+
},
|
| 2934 |
+
"outputs": [],
|
| 2935 |
+
"source": [
|
| 2936 |
+
"img = cv2.cvtColor(cv2.imread('images/books.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 2937 |
+
"plt.imshow(img)"
|
| 2938 |
+
]
|
| 2939 |
+
},
|
| 2940 |
+
{
|
| 2941 |
+
"cell_type": "code",
|
| 2942 |
+
"execution_count": null,
|
| 2943 |
+
"metadata": {
|
| 2944 |
+
"id": "gS2OY6czd2oX",
|
| 2945 |
+
"pycharm": {}
|
| 2946 |
+
},
|
| 2947 |
+
"outputs": [],
|
| 2948 |
+
"source": [
|
| 2949 |
+
"gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n",
|
| 2950 |
+
"blur_img = cv2.blur(gray_img, (10, 10))\n",
|
| 2951 |
+
"plt.subplot(1, 2, 1); plt.imshow(gray_img, cmap='gray')\n",
|
| 2952 |
+
"plt.subplot(1, 2, 2); plt.imshow(blur_img, cmap='gray')"
|
| 2953 |
+
]
|
| 2954 |
+
},
|
| 2955 |
+
{
|
| 2956 |
+
"cell_type": "markdown",
|
| 2957 |
+
"metadata": {
|
| 2958 |
+
"id": "tOqyh7ZD5nW0"
|
| 2959 |
+
},
|
| 2960 |
+
"source": [
|
| 2961 |
+
"#### Gaussian Filtering"
|
| 2962 |
+
]
|
| 2963 |
+
},
|
| 2964 |
+
{
|
| 2965 |
+
"cell_type": "markdown",
|
| 2966 |
+
"metadata": {
|
| 2967 |
+
"id": "EhoKd7Is_p8y",
|
| 2968 |
+
"pycharm": {}
|
| 2969 |
+
},
|
| 2970 |
+
"source": [
|
| 2971 |
+
"In Gaussian filtering, instead of a box filter, a Gaussian kernel is used. In OpenCV, it is done with the function, `cv2.GaussianBlur()`. We should specify the width and height of the kernel which should be positive and odd. We also should specify the standard deviation in the X and Y directions, sigmaX and sigmaY respectively. If only sigmaX is specified, sigmaY is taken as the same as sigmaX. If both are given as zeros, they are calculated from the kernel size. Gaussian blurring is highly effective in removing Gaussian noise from an image."
|
| 2972 |
+
]
|
| 2973 |
+
},
|
| 2974 |
+
{
|
| 2975 |
+
"cell_type": "code",
|
| 2976 |
+
"execution_count": null,
|
| 2977 |
+
"metadata": {
|
| 2978 |
+
"id": "R2qNZNzB_p8z",
|
| 2979 |
+
"pycharm": {}
|
| 2980 |
+
},
|
| 2981 |
+
"outputs": [],
|
| 2982 |
+
"source": [
|
| 2983 |
+
"img = cv2.cvtColor(cv2.imread('images/oy.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 2984 |
+
"plt.imshow(img)"
|
| 2985 |
+
]
|
| 2986 |
+
},
|
| 2987 |
+
{
|
| 2988 |
+
"cell_type": "code",
|
| 2989 |
+
"execution_count": null,
|
| 2990 |
+
"metadata": {
|
| 2991 |
+
"id": "FcEcRBhT_p81",
|
| 2992 |
+
"pycharm": {}
|
| 2993 |
+
},
|
| 2994 |
+
"outputs": [],
|
| 2995 |
+
"source": [
|
| 2996 |
+
"# preproccess with blurring, with 5x5 kernel (note kernel size should be odd)\n",
|
| 2997 |
+
"img_blur_small = cv2.GaussianBlur(img, (5, 5), 0)\n",
|
| 2998 |
+
"plt.imshow(img_blur_small)"
|
| 2999 |
+
]
|
| 3000 |
+
},
|
| 3001 |
+
{
|
| 3002 |
+
"cell_type": "code",
|
| 3003 |
+
"execution_count": null,
|
| 3004 |
+
"metadata": {
|
| 3005 |
+
"id": "GcpJwBNU_p83",
|
| 3006 |
+
"pycharm": {}
|
| 3007 |
+
},
|
| 3008 |
+
"outputs": [],
|
| 3009 |
+
"source": [
|
| 3010 |
+
"img_blur_small = cv2.GaussianBlur(img, (5, 5), 25)\n",
|
| 3011 |
+
"plt.imshow(img_blur_small)"
|
| 3012 |
+
]
|
| 3013 |
+
},
|
| 3014 |
+
{
|
| 3015 |
+
"cell_type": "code",
|
| 3016 |
+
"execution_count": null,
|
| 3017 |
+
"metadata": {
|
| 3018 |
+
"id": "GW_zbFBx_p85",
|
| 3019 |
+
"pycharm": {}
|
| 3020 |
+
},
|
| 3021 |
+
"outputs": [],
|
| 3022 |
+
"source": [
|
| 3023 |
+
"img_blur_large = cv2.GaussianBlur(img, (15,15), 0)\n",
|
| 3024 |
+
"plt.imshow(img_blur_large)"
|
| 3025 |
+
]
|
| 3026 |
+
},
|
| 3027 |
+
{
|
| 3028 |
+
"cell_type": "markdown",
|
| 3029 |
+
"metadata": {
|
| 3030 |
+
"id": "knYfkspQ5y0F"
|
| 3031 |
+
},
|
| 3032 |
+
"source": [
|
| 3033 |
+
"#### Median Filtering"
|
| 3034 |
+
]
|
| 3035 |
+
},
|
| 3036 |
+
{
|
| 3037 |
+
"cell_type": "markdown",
|
| 3038 |
+
"metadata": {
|
| 3039 |
+
"id": "rFRJL4v1O7kt",
|
| 3040 |
+
"pycharm": {}
|
| 3041 |
+
},
|
| 3042 |
+
"source": [
|
| 3043 |
+
"This is a non-linear filtering technique. As clear from the name, this takes a median of all the pixels under the kernel area and replaces the central element with this median value. This is quite effective in reducing a certain type of noise (like salt-and-pepper noise) with considerably less edge blurring as compared to other linear filters of the same size. First create the function for creating noisy images with \"salt and pepper\" noise."
|
| 3044 |
+
]
|
| 3045 |
+
},
|
| 3046 |
+
{
|
| 3047 |
+
"cell_type": "code",
|
| 3048 |
+
"execution_count": null,
|
| 3049 |
+
"metadata": {
|
| 3050 |
+
"id": "dqMvcYl3u_S8",
|
| 3051 |
+
"pycharm": {}
|
| 3052 |
+
},
|
| 3053 |
+
"outputs": [],
|
| 3054 |
+
"source": [
|
| 3055 |
+
"def add_sp_noise(image, amount=0.1):\n",
|
| 3056 |
+
" row, col, ch = image.shape\n",
|
| 3057 |
+
" s_vs_p = 0.5\n",
|
| 3058 |
+
" out = np.copy(image)\n",
|
| 3059 |
+
" # Salt mode\n",
|
| 3060 |
+
" num_salt = np.ceil(amount * image.size * s_vs_p)\n",
|
| 3061 |
+
" coords = [np.random.randint(0, i - 1, int(num_salt))\n",
|
| 3062 |
+
" for i in image.shape]\n",
|
| 3063 |
+
" out[coords[0], coords[1], coords[2]] = 1\n",
|
| 3064 |
+
"\n",
|
| 3065 |
+
" # Pepper mode\n",
|
| 3066 |
+
" num_pepper = np.ceil(amount* image.size * (1. - s_vs_p))\n",
|
| 3067 |
+
" coords = [np.random.randint(0, i - 1, int(num_pepper))\n",
|
| 3068 |
+
" for i in image.shape]\n",
|
| 3069 |
+
" out[coords[0], coords[1], coords[2]] = 0\n",
|
| 3070 |
+
" return out"
|
| 3071 |
+
]
|
| 3072 |
+
},
|
| 3073 |
+
{
|
| 3074 |
+
"cell_type": "markdown",
|
| 3075 |
+
"metadata": {
|
| 3076 |
+
"id": "NBNSgJgfDgqv",
|
| 3077 |
+
"pycharm": {}
|
| 3078 |
+
},
|
| 3079 |
+
"source": [
|
| 3080 |
+
"Load an image and apply \"salt and pepper\" noise and then try to smooth it with Gaussian and Median filter"
|
| 3081 |
+
]
|
| 3082 |
+
},
|
| 3083 |
+
{
|
| 3084 |
+
"cell_type": "code",
|
| 3085 |
+
"execution_count": null,
|
| 3086 |
+
"metadata": {
|
| 3087 |
+
"id": "ea2GEwfFQLR1",
|
| 3088 |
+
"pycharm": {
|
| 3089 |
+
"is_executing": true
|
| 3090 |
+
}
|
| 3091 |
+
},
|
| 3092 |
+
"outputs": [],
|
| 3093 |
+
"source": [
|
| 3094 |
+
"img = cv2.cvtColor(cv2.imread('images/coins.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 3095 |
+
"noisy_img = add_sp_noise(img, amount=0.1)\n",
|
| 3096 |
+
"img_gaus = cv2.GaussianBlur(noisy_img, (5, 5), 3)\n",
|
| 3097 |
+
"img_med = cv2.medianBlur(noisy_img, 5)\n",
|
| 3098 |
+
"plt.subplot(1, 4, 1); plt.imshow(img); plt.title('Original')\n",
|
| 3099 |
+
"plt.subplot(1, 4, 2); plt.imshow(noisy_img); plt.title('Salt & Pepper Noise')\n",
|
| 3100 |
+
"plt.subplot(1, 4, 3); plt.imshow(img_gaus); plt.title('Gaussian Filtered')\n",
|
| 3101 |
+
"plt.subplot(1, 4, 4); plt.imshow(img_med); plt.title('Median Filtered')"
|
| 3102 |
+
]
|
| 3103 |
+
},
|
| 3104 |
+
{
|
| 3105 |
+
"cell_type": "markdown",
|
| 3106 |
+
"metadata": {
|
| 3107 |
+
"id": "jBSFKfNp58ka"
|
| 3108 |
+
},
|
| 3109 |
+
"source": [
|
| 3110 |
+
"#### Edge Detection"
|
| 3111 |
+
]
|
| 3112 |
+
},
|
| 3113 |
+
{
|
| 3114 |
+
"cell_type": "markdown",
|
| 3115 |
+
"metadata": {
|
| 3116 |
+
"id": "XxeFuSii_p9N",
|
| 3117 |
+
"pycharm": {}
|
| 3118 |
+
},
|
| 3119 |
+
"source": [
|
| 3120 |
+
"Edge detection is an image processing technique for finding the boundaries of objects within images. It works by detecting discontinuities in brightness, colour, surface etc. Edge detection is used for image segmentation and data extraction in areas such as image processing, computer vision, and machine vision. OpenCV provides the `cv2.Canny()` function to compute edges in an image."
|
| 3121 |
+
]
|
| 3122 |
+
},
|
| 3123 |
+
{
|
| 3124 |
+
"cell_type": "code",
|
| 3125 |
+
"execution_count": null,
|
| 3126 |
+
"metadata": {
|
| 3127 |
+
"id": "-utaqZp5SDP7",
|
| 3128 |
+
"pycharm": {}
|
| 3129 |
+
},
|
| 3130 |
+
"outputs": [],
|
| 3131 |
+
"source": [
|
| 3132 |
+
"cups = cv2.cvtColor(cv2.imread('images/cups.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 3133 |
+
"plt.imshow(cups)"
|
| 3134 |
+
]
|
| 3135 |
+
},
|
| 3136 |
+
{
|
| 3137 |
+
"cell_type": "code",
|
| 3138 |
+
"execution_count": null,
|
| 3139 |
+
"metadata": {
|
| 3140 |
+
"id": "7Ko1a2jmSM-M",
|
| 3141 |
+
"pycharm": {}
|
| 3142 |
+
},
|
| 3143 |
+
"outputs": [],
|
| 3144 |
+
"source": [
|
| 3145 |
+
"# preprocess by blurring and grayscale\n",
|
| 3146 |
+
"cups_preprocessed = cv2.cvtColor(cv2.GaussianBlur(cups, (7,7), 0), cv2.COLOR_RGB2GRAY)"
|
| 3147 |
+
]
|
| 3148 |
+
},
|
| 3149 |
+
{
|
| 3150 |
+
"cell_type": "code",
|
| 3151 |
+
"execution_count": null,
|
| 3152 |
+
"metadata": {
|
| 3153 |
+
"id": "a8-A44piSmwd",
|
| 3154 |
+
"pycharm": {}
|
| 3155 |
+
},
|
| 3156 |
+
"outputs": [],
|
| 3157 |
+
"source": [
|
| 3158 |
+
"# find binary image with thresholding\n",
|
| 3159 |
+
"low_thresh = 120\n",
|
| 3160 |
+
"high_thresh = 200\n",
|
| 3161 |
+
"_, cups_thresh = cv2.threshold(cups_preprocessed, low_thresh, 255, cv2.THRESH_BINARY)\n",
|
| 3162 |
+
"plt.imshow(cv2.cvtColor(cups_thresh, cv2.COLOR_GRAY2RGB))\n",
|
| 3163 |
+
"\n",
|
| 3164 |
+
"_, cups_thresh_hi = cv2.threshold(cups_preprocessed, high_thresh, 255, cv2.THRESH_BINARY)"
|
| 3165 |
+
]
|
| 3166 |
+
},
|
| 3167 |
+
{
|
| 3168 |
+
"cell_type": "code",
|
| 3169 |
+
"execution_count": null,
|
| 3170 |
+
"metadata": {
|
| 3171 |
+
"id": "lVNkDIgDRuci",
|
| 3172 |
+
"pycharm": {}
|
| 3173 |
+
},
|
| 3174 |
+
"outputs": [],
|
| 3175 |
+
"source": [
|
| 3176 |
+
"# find binary image with edges\n",
|
| 3177 |
+
"cups_edges = cv2.Canny(cups_preprocessed, threshold1=90, threshold2=110)\n",
|
| 3178 |
+
"plt.imshow(cv2.cvtColor(cups_edges, cv2.COLOR_GRAY2RGB))"
|
| 3179 |
+
]
|
| 3180 |
+
},
|
| 3181 |
+
{
|
| 3182 |
+
"cell_type": "markdown",
|
| 3183 |
+
"metadata": {
|
| 3184 |
+
"id": "-XOcVQ4hqP4R"
|
| 3185 |
+
},
|
| 3186 |
+
"source": [
|
| 3187 |
+
"## SciPy"
|
| 3188 |
+
]
|
| 3189 |
+
},
|
| 3190 |
+
{
|
| 3191 |
+
"cell_type": "markdown",
|
| 3192 |
+
"metadata": {
|
| 3193 |
+
"id": "tAWDvNu5qn4b"
|
| 3194 |
+
},
|
| 3195 |
+
"source": [
|
| 3196 |
+
"Numpy provides a high-performance multidimensional array and basic tools to compute with and manipulate these arrays. [SciPy](http://docs.scipy.org/doc/scipy/reference/) builds on this, and provides a large number of functions that operate on numpy arrays and are useful for different types of scientific and engineering applications. The best way to get familiar with SciPy is to [browse the documentation](https://docs.scipy.org/doc/scipy/reference/index.html). SciPy provides important functionalities for reading and writing MATLAB files, which show below.\n",
|
| 3197 |
+
"\n",
|
| 3198 |
+
"\n"
|
| 3199 |
+
]
|
| 3200 |
+
},
|
| 3201 |
+
{
|
| 3202 |
+
"cell_type": "markdown",
|
| 3203 |
+
"metadata": {
|
| 3204 |
+
"id": "ajs-UbqSrWk0"
|
| 3205 |
+
},
|
| 3206 |
+
"source": [
|
| 3207 |
+
"###MATLAB files"
|
| 3208 |
+
]
|
| 3209 |
+
},
|
| 3210 |
+
{
|
| 3211 |
+
"cell_type": "markdown",
|
| 3212 |
+
"metadata": {
|
| 3213 |
+
"id": "HoT2zazhrZ5m"
|
| 3214 |
+
},
|
| 3215 |
+
"source": [
|
| 3216 |
+
"The functions `scipy.io.loadmat` and `scipy.io.savemat` allow you to respectively read and write MATLAB files. You can read about them [in the documentation](http://docs.scipy.org/doc/scipy/reference/io.html)."
|
| 3217 |
+
]
|
| 3218 |
+
},
|
| 3219 |
+
{
|
| 3220 |
+
"cell_type": "markdown",
|
| 3221 |
+
"metadata": {
|
| 3222 |
+
"id": "iMtaY6Bzr7-w"
|
| 3223 |
+
},
|
| 3224 |
+
"source": [
|
| 3225 |
+
"###Distance between points"
|
| 3226 |
+
]
|
| 3227 |
+
},
|
| 3228 |
+
{
|
| 3229 |
+
"cell_type": "markdown",
|
| 3230 |
+
"metadata": {
|
| 3231 |
+
"id": "1tq9Mtwkr_s0"
|
| 3232 |
+
},
|
| 3233 |
+
"source": [
|
| 3234 |
+
"SciPy defines some useful functions for computing distances between sets of points.\n",
|
| 3235 |
+
"\n",
|
| 3236 |
+
"The function `scipy.spatial.distance.pdist` computes the distance between all pairs of points in a given set:"
|
| 3237 |
+
]
|
| 3238 |
+
},
|
| 3239 |
+
{
|
| 3240 |
+
"cell_type": "code",
|
| 3241 |
+
"execution_count": null,
|
| 3242 |
+
"metadata": {
|
| 3243 |
+
"id": "EwlHRO0jsJBI"
|
| 3244 |
+
},
|
| 3245 |
+
"outputs": [],
|
| 3246 |
+
"source": [
|
| 3247 |
+
"import numpy as np\n",
|
| 3248 |
+
"from scipy.spatial.distance import pdist, squareform\n",
|
| 3249 |
+
"\n",
|
| 3250 |
+
"# Create the following array where each row is a point in 2D space:\n",
|
| 3251 |
+
"# [[0 1]\n",
|
| 3252 |
+
"# [1 0]\n",
|
| 3253 |
+
"# [2 0]]\n",
|
| 3254 |
+
"x = np.array([[0, 1], [1, 0], [2, 0]])\n",
|
| 3255 |
+
"print(x)\n",
|
| 3256 |
+
"\n",
|
| 3257 |
+
"# Compute the Euclidean distance between all rows of x.\n",
|
| 3258 |
+
"# d[i, j] is the Euclidean distance between x[i, :] and x[j, :],\n",
|
| 3259 |
+
"# and d is the following array:\n",
|
| 3260 |
+
"# [[ 0. 1.41421356 2.23606798]\n",
|
| 3261 |
+
"# [ 1.41421356 0. 1. ]\n",
|
| 3262 |
+
"# [ 2.23606798 1. 0. ]]\n",
|
| 3263 |
+
"d = squareform(pdist(x, 'euclidean'))\n",
|
| 3264 |
+
"print(d)"
|
| 3265 |
+
]
|
| 3266 |
+
},
|
| 3267 |
+
{
|
| 3268 |
+
"cell_type": "markdown",
|
| 3269 |
+
"metadata": {
|
| 3270 |
+
"id": "dzYk_QSfsSXO"
|
| 3271 |
+
},
|
| 3272 |
+
"source": [
|
| 3273 |
+
"A similar function (`scipy.spatial.distance.cdist`) computes the distance between all pairs across two sets of points; you can read about it [in the documentation](https://docs.scipy.org/doc/scipy/reference/generated/scipy.spatial.distance.cdist.html)."
|
| 3274 |
+
]
|
| 3275 |
+
},
|
| 3276 |
+
{
|
| 3277 |
+
"cell_type": "markdown",
|
| 3278 |
+
"metadata": {
|
| 3279 |
+
"id": "d3XD7jVkU9Z3"
|
| 3280 |
+
},
|
| 3281 |
+
"source": [
|
| 3282 |
+
"#### Acknowledgement\n",
|
| 3283 |
+
"This tutorial was originally written by [Justin Johnson](https://web.eecs.umich.edu/~justincj/) for CS231n at the Stanford University. This version has been adapted and modified by [Anjan Dutta](https://www.surrey.ac.uk/people/anjan-dutta) for the Spring 2023 edition of [EEEM068](https://catalogue.surrey.ac.uk/2022-3/module/EEEM068) module at the University of Surrey."
|
| 3284 |
+
]
|
| 3285 |
+
}
|
| 3286 |
+
],
|
| 3287 |
+
"metadata": {
|
| 3288 |
+
"colab": {
|
| 3289 |
+
"include_colab_link": true,
|
| 3290 |
+
"name": "colab-tutorial.ipynb",
|
| 3291 |
+
"provenance": []
|
| 3292 |
+
},
|
| 3293 |
+
"kernelspec": {
|
| 3294 |
+
"display_name": "Python 3 (ipykernel)",
|
| 3295 |
+
"language": "python",
|
| 3296 |
+
"name": "python3"
|
| 3297 |
+
},
|
| 3298 |
+
"language_info": {
|
| 3299 |
+
"codemirror_mode": {
|
| 3300 |
+
"name": "ipython",
|
| 3301 |
+
"version": 3
|
| 3302 |
+
},
|
| 3303 |
+
"file_extension": ".py",
|
| 3304 |
+
"mimetype": "text/x-python",
|
| 3305 |
+
"name": "python",
|
| 3306 |
+
"nbconvert_exporter": "python",
|
| 3307 |
+
"pygments_lexer": "ipython3",
|
| 3308 |
+
"version": "3.12.3"
|
| 3309 |
+
}
|
| 3310 |
+
},
|
| 3311 |
+
"nbformat": 4,
|
| 3312 |
+
"nbformat_minor": 1
|
| 3313 |
+
}
|
Downloads/.ipynb_checkpoints/Python Tutorial-checkpoint.ipynb
ADDED
|
@@ -0,0 +1,3313 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {
|
| 6 |
+
"colab_type": "text",
|
| 7 |
+
"id": "view-in-github"
|
| 8 |
+
},
|
| 9 |
+
"source": [
|
| 10 |
+
"<a href=\"https://colab.research.google.com/github/AnjanDutta/EEEM068/blob/main/Notebooks/Python_Tutorial.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
|
| 11 |
+
]
|
| 12 |
+
},
|
| 13 |
+
{
|
| 14 |
+
"cell_type": "markdown",
|
| 15 |
+
"metadata": {
|
| 16 |
+
"id": "dzNng6vCL9eP"
|
| 17 |
+
},
|
| 18 |
+
"source": [
|
| 19 |
+
"<H1 style=\"text-align: center\">EEEM068 - Applied Machine Learning</H1>\n",
|
| 20 |
+
"<H1 style=\"text-align: center\">Workshop 01</H1>\n",
|
| 21 |
+
"<H1 style=\"text-align: center\">Python Tutorial</H1>\n"
|
| 22 |
+
]
|
| 23 |
+
},
|
| 24 |
+
{
|
| 25 |
+
"cell_type": "markdown",
|
| 26 |
+
"metadata": {
|
| 27 |
+
"id": "qVrTo-LhL9eS"
|
| 28 |
+
},
|
| 29 |
+
"source": [
|
| 30 |
+
"##Introduction"
|
| 31 |
+
]
|
| 32 |
+
},
|
| 33 |
+
{
|
| 34 |
+
"cell_type": "markdown",
|
| 35 |
+
"metadata": {
|
| 36 |
+
"id": "9t1gKp9PL9eV"
|
| 37 |
+
},
|
| 38 |
+
"source": [
|
| 39 |
+
"Python is a great general purpose programming language on its own, but with the help of a few popular libraries, such as numpy, scipy, matplotlib it becomes a powerful environment for scientific computing.\n",
|
| 40 |
+
"\n",
|
| 41 |
+
"We expect that many of you to have some experience with Python and numpy. Nevertheless, for the rest of you, this tutorial will serve as a quick crash course both on the Python programming language and on the use of Python for scientific computing.\n",
|
| 42 |
+
"\n",
|
| 43 |
+
"Some of you may have previous knowledge in Matlab, in which case we also recommend the [NumPy for Matlab users](https://numpy.org/doc/stable/user/numpy-for-matlab-users.html) page."
|
| 44 |
+
]
|
| 45 |
+
},
|
| 46 |
+
{
|
| 47 |
+
"cell_type": "markdown",
|
| 48 |
+
"metadata": {
|
| 49 |
+
"id": "U1PvreR9L9eW"
|
| 50 |
+
},
|
| 51 |
+
"source": [
|
| 52 |
+
"In this tutorial, we will cover:\n",
|
| 53 |
+
"\n",
|
| 54 |
+
"* Basic Python: Basic data types (Containers, Lists, Dictionaries, Sets, Tuples), Functions, Classes\n",
|
| 55 |
+
"* Numpy: Arrays, Array indexing, Datatypes, Array math, Broadcasting\n",
|
| 56 |
+
"* Matplotlib: Plotting, Subplots, Images\n",
|
| 57 |
+
"* Scikit-learn: Toy dataset, Classifier, Confusion matrix, Regressor\n",
|
| 58 |
+
"* OpenCV: Image, Image representation, Colour and datatype conversion, Image processing\n",
|
| 59 |
+
"* SciPy: I/O of MATLAB files, Distance functions"
|
| 60 |
+
]
|
| 61 |
+
},
|
| 62 |
+
{
|
| 63 |
+
"cell_type": "markdown",
|
| 64 |
+
"metadata": {
|
| 65 |
+
"id": "-O99OrwPtGii"
|
| 66 |
+
},
|
| 67 |
+
"source": [
|
| 68 |
+
"## Python Versions"
|
| 69 |
+
]
|
| 70 |
+
},
|
| 71 |
+
{
|
| 72 |
+
"cell_type": "markdown",
|
| 73 |
+
"metadata": {
|
| 74 |
+
"id": "nxvEkGXPM3Xh"
|
| 75 |
+
},
|
| 76 |
+
"source": [
|
| 77 |
+
"Please note that as of February 2026, Colab is using Python 3.12.12. Therefore, we will be using Python 3.12 for this iteration of the course. More details on Python 3.12 can be found in the [documentation](https://docs.python.org/3.12/tutorial/index.html). You can check your Python version at the command line by running `python --version`."
|
| 78 |
+
]
|
| 79 |
+
},
|
| 80 |
+
{
|
| 81 |
+
"cell_type": "code",
|
| 82 |
+
"execution_count": null,
|
| 83 |
+
"metadata": {
|
| 84 |
+
"id": "1L4Am0QATgOc"
|
| 85 |
+
},
|
| 86 |
+
"outputs": [],
|
| 87 |
+
"source": [
|
| 88 |
+
"!python --version"
|
| 89 |
+
]
|
| 90 |
+
},
|
| 91 |
+
{
|
| 92 |
+
"cell_type": "markdown",
|
| 93 |
+
"metadata": {
|
| 94 |
+
"id": "JAFKYgrpL9eY"
|
| 95 |
+
},
|
| 96 |
+
"source": [
|
| 97 |
+
"##Basics of Python"
|
| 98 |
+
]
|
| 99 |
+
},
|
| 100 |
+
{
|
| 101 |
+
"cell_type": "markdown",
|
| 102 |
+
"metadata": {
|
| 103 |
+
"id": "RbFS6tdgL9ea"
|
| 104 |
+
},
|
| 105 |
+
"source": [
|
| 106 |
+
"Python is an easy to learn, high-level, dynamically typed multiparadigm programming language. Python code is often said to be almost like pseudocode, since it allows you to express very powerful ideas in very few lines of code while being very readable. As an example, here is an implementation of the classic quicksort algorithm in Python:"
|
| 107 |
+
]
|
| 108 |
+
},
|
| 109 |
+
{
|
| 110 |
+
"cell_type": "code",
|
| 111 |
+
"execution_count": null,
|
| 112 |
+
"metadata": {
|
| 113 |
+
"id": "cYb0pjh1L9eb"
|
| 114 |
+
},
|
| 115 |
+
"outputs": [],
|
| 116 |
+
"source": [
|
| 117 |
+
"def quicksort(arr):\n",
|
| 118 |
+
" if len(arr) <= 1:\n",
|
| 119 |
+
" return arr\n",
|
| 120 |
+
" pivot = arr[len(arr) // 2]\n",
|
| 121 |
+
" left = [x for x in arr if x < pivot]\n",
|
| 122 |
+
" middle = [x for x in arr if x == pivot]\n",
|
| 123 |
+
" right = [x for x in arr if x > pivot]\n",
|
| 124 |
+
" return quicksort(left) + middle + quicksort(right)\n",
|
| 125 |
+
"\n",
|
| 126 |
+
"print(quicksort([3,6,8,10,1,2,1]))"
|
| 127 |
+
]
|
| 128 |
+
},
|
| 129 |
+
{
|
| 130 |
+
"cell_type": "markdown",
|
| 131 |
+
"metadata": {
|
| 132 |
+
"id": "NwS_hu4xL9eo"
|
| 133 |
+
},
|
| 134 |
+
"source": [
|
| 135 |
+
"###Basic data types"
|
| 136 |
+
]
|
| 137 |
+
},
|
| 138 |
+
{
|
| 139 |
+
"cell_type": "markdown",
|
| 140 |
+
"metadata": {
|
| 141 |
+
"id": "DL5sMSZ9L9eq"
|
| 142 |
+
},
|
| 143 |
+
"source": [
|
| 144 |
+
"####Numbers"
|
| 145 |
+
]
|
| 146 |
+
},
|
| 147 |
+
{
|
| 148 |
+
"cell_type": "markdown",
|
| 149 |
+
"metadata": {
|
| 150 |
+
"id": "MGS0XEWoL9er"
|
| 151 |
+
},
|
| 152 |
+
"source": [
|
| 153 |
+
"Integers and floats work as you would expect from other languages:"
|
| 154 |
+
]
|
| 155 |
+
},
|
| 156 |
+
{
|
| 157 |
+
"cell_type": "code",
|
| 158 |
+
"execution_count": null,
|
| 159 |
+
"metadata": {
|
| 160 |
+
"id": "KheDr_zDL9es"
|
| 161 |
+
},
|
| 162 |
+
"outputs": [],
|
| 163 |
+
"source": [
|
| 164 |
+
"x = 3\n",
|
| 165 |
+
"print(x, type(x))"
|
| 166 |
+
]
|
| 167 |
+
},
|
| 168 |
+
{
|
| 169 |
+
"cell_type": "code",
|
| 170 |
+
"execution_count": null,
|
| 171 |
+
"metadata": {
|
| 172 |
+
"id": "sk_8DFcuL9ey"
|
| 173 |
+
},
|
| 174 |
+
"outputs": [],
|
| 175 |
+
"source": [
|
| 176 |
+
"print(x + 1) # Addition\n",
|
| 177 |
+
"print(x - 1) # Subtraction\n",
|
| 178 |
+
"print(x * 2) # Multiplication\n",
|
| 179 |
+
"print(x ** 2) # Exponentiation"
|
| 180 |
+
]
|
| 181 |
+
},
|
| 182 |
+
{
|
| 183 |
+
"cell_type": "code",
|
| 184 |
+
"execution_count": null,
|
| 185 |
+
"metadata": {
|
| 186 |
+
"id": "U4Jl8K0tL9e4"
|
| 187 |
+
},
|
| 188 |
+
"outputs": [],
|
| 189 |
+
"source": [
|
| 190 |
+
"x += 1\n",
|
| 191 |
+
"print(x)\n",
|
| 192 |
+
"x *= 2\n",
|
| 193 |
+
"print(x)"
|
| 194 |
+
]
|
| 195 |
+
},
|
| 196 |
+
{
|
| 197 |
+
"cell_type": "code",
|
| 198 |
+
"execution_count": null,
|
| 199 |
+
"metadata": {
|
| 200 |
+
"id": "w-nZ0Sg_L9e9"
|
| 201 |
+
},
|
| 202 |
+
"outputs": [],
|
| 203 |
+
"source": [
|
| 204 |
+
"y = 2.5\n",
|
| 205 |
+
"print(type(y))\n",
|
| 206 |
+
"print(y, y + 1, y * 2, y ** 2)"
|
| 207 |
+
]
|
| 208 |
+
},
|
| 209 |
+
{
|
| 210 |
+
"cell_type": "markdown",
|
| 211 |
+
"metadata": {
|
| 212 |
+
"id": "r2A9ApyaL9fB"
|
| 213 |
+
},
|
| 214 |
+
"source": [
|
| 215 |
+
"Note that unlike many languages (such as C and C++) Python does not have unary increment (x++) or decrement (x--) operators.\n",
|
| 216 |
+
"\n",
|
| 217 |
+
"Python also has built-in types for long integers and complex numbers; you can find all of the details in the [documentation](https://docs.python.org/3.8/library/stdtypes.html#numeric-types-int-float-long-complex)."
|
| 218 |
+
]
|
| 219 |
+
},
|
| 220 |
+
{
|
| 221 |
+
"cell_type": "markdown",
|
| 222 |
+
"metadata": {
|
| 223 |
+
"id": "EqRS7qhBL9fC"
|
| 224 |
+
},
|
| 225 |
+
"source": [
|
| 226 |
+
"####Booleans"
|
| 227 |
+
]
|
| 228 |
+
},
|
| 229 |
+
{
|
| 230 |
+
"cell_type": "markdown",
|
| 231 |
+
"metadata": {
|
| 232 |
+
"id": "Nv_LIVOJL9fD"
|
| 233 |
+
},
|
| 234 |
+
"source": [
|
| 235 |
+
"Python implements all of the usual operators for Boolean logic, but uses English words rather than symbols (`&&`, `||`, etc.):"
|
| 236 |
+
]
|
| 237 |
+
},
|
| 238 |
+
{
|
| 239 |
+
"cell_type": "code",
|
| 240 |
+
"execution_count": null,
|
| 241 |
+
"metadata": {
|
| 242 |
+
"id": "RvoImwgGL9fE"
|
| 243 |
+
},
|
| 244 |
+
"outputs": [],
|
| 245 |
+
"source": [
|
| 246 |
+
"t, f = True, False\n",
|
| 247 |
+
"print(type(t))"
|
| 248 |
+
]
|
| 249 |
+
},
|
| 250 |
+
{
|
| 251 |
+
"cell_type": "markdown",
|
| 252 |
+
"metadata": {
|
| 253 |
+
"id": "YQgmQfOgL9fI"
|
| 254 |
+
},
|
| 255 |
+
"source": [
|
| 256 |
+
"Now we let's look at the operations:"
|
| 257 |
+
]
|
| 258 |
+
},
|
| 259 |
+
{
|
| 260 |
+
"cell_type": "code",
|
| 261 |
+
"execution_count": null,
|
| 262 |
+
"metadata": {
|
| 263 |
+
"id": "6zYm7WzCL9fK"
|
| 264 |
+
},
|
| 265 |
+
"outputs": [],
|
| 266 |
+
"source": [
|
| 267 |
+
"print(t and f) # Logical AND;\n",
|
| 268 |
+
"print(t or f) # Logical OR;\n",
|
| 269 |
+
"print(not t) # Logical NOT;\n",
|
| 270 |
+
"print(t != f) # Logical XOR;"
|
| 271 |
+
]
|
| 272 |
+
},
|
| 273 |
+
{
|
| 274 |
+
"cell_type": "markdown",
|
| 275 |
+
"metadata": {
|
| 276 |
+
"id": "UQnQWFEyL9fP"
|
| 277 |
+
},
|
| 278 |
+
"source": [
|
| 279 |
+
"####Strings"
|
| 280 |
+
]
|
| 281 |
+
},
|
| 282 |
+
{
|
| 283 |
+
"cell_type": "code",
|
| 284 |
+
"execution_count": null,
|
| 285 |
+
"metadata": {
|
| 286 |
+
"id": "AijEDtPFL9fP"
|
| 287 |
+
},
|
| 288 |
+
"outputs": [],
|
| 289 |
+
"source": [
|
| 290 |
+
"hello = 'hello' # String literals can use single quotes\n",
|
| 291 |
+
"world = \"world\" # or double quotes; it does not matter\n",
|
| 292 |
+
"print(hello, len(hello))"
|
| 293 |
+
]
|
| 294 |
+
},
|
| 295 |
+
{
|
| 296 |
+
"cell_type": "code",
|
| 297 |
+
"execution_count": null,
|
| 298 |
+
"metadata": {
|
| 299 |
+
"id": "saDeaA7hL9fT"
|
| 300 |
+
},
|
| 301 |
+
"outputs": [],
|
| 302 |
+
"source": [
|
| 303 |
+
"hw = hello + ' ' + world # String concatenation\n",
|
| 304 |
+
"print(hw)"
|
| 305 |
+
]
|
| 306 |
+
},
|
| 307 |
+
{
|
| 308 |
+
"cell_type": "code",
|
| 309 |
+
"execution_count": null,
|
| 310 |
+
"metadata": {
|
| 311 |
+
"id": "Nji1_UjYL9fY"
|
| 312 |
+
},
|
| 313 |
+
"outputs": [],
|
| 314 |
+
"source": [
|
| 315 |
+
"hw12 = '{} {} {}'.format(hello, world, 12) # string formatting\n",
|
| 316 |
+
"print(hw12)"
|
| 317 |
+
]
|
| 318 |
+
},
|
| 319 |
+
{
|
| 320 |
+
"cell_type": "markdown",
|
| 321 |
+
"metadata": {
|
| 322 |
+
"id": "bUpl35bIL9fc"
|
| 323 |
+
},
|
| 324 |
+
"source": [
|
| 325 |
+
"String objects have a bunch of useful methods; for example:"
|
| 326 |
+
]
|
| 327 |
+
},
|
| 328 |
+
{
|
| 329 |
+
"cell_type": "code",
|
| 330 |
+
"execution_count": null,
|
| 331 |
+
"metadata": {
|
| 332 |
+
"id": "VOxGatlsL9fd"
|
| 333 |
+
},
|
| 334 |
+
"outputs": [],
|
| 335 |
+
"source": [
|
| 336 |
+
"s = \"hello\"\n",
|
| 337 |
+
"print(s.capitalize()) # Capitalize a string\n",
|
| 338 |
+
"print(s.upper()) # Convert a string to uppercase; prints \"HELLO\"\n",
|
| 339 |
+
"print(s.rjust(7)) # Right-justify a string, padding with spaces\n",
|
| 340 |
+
"print(s.center(7)) # Center a string, padding with spaces\n",
|
| 341 |
+
"print(s.replace('l', '(ell)')) # Replace all instances of one substring with another\n",
|
| 342 |
+
"print(' world '.strip()) # Strip leading and trailing whitespace"
|
| 343 |
+
]
|
| 344 |
+
},
|
| 345 |
+
{
|
| 346 |
+
"cell_type": "markdown",
|
| 347 |
+
"metadata": {
|
| 348 |
+
"id": "06cayXLtL9fi"
|
| 349 |
+
},
|
| 350 |
+
"source": [
|
| 351 |
+
"You can find a list of all string methods in the [documentation](https://docs.python.org/3.7/library/stdtypes.html#string-methods)."
|
| 352 |
+
]
|
| 353 |
+
},
|
| 354 |
+
{
|
| 355 |
+
"cell_type": "markdown",
|
| 356 |
+
"metadata": {
|
| 357 |
+
"id": "p-6hClFjL9fk"
|
| 358 |
+
},
|
| 359 |
+
"source": [
|
| 360 |
+
"###Containers"
|
| 361 |
+
]
|
| 362 |
+
},
|
| 363 |
+
{
|
| 364 |
+
"cell_type": "markdown",
|
| 365 |
+
"metadata": {
|
| 366 |
+
"id": "FD9H18eQL9fk"
|
| 367 |
+
},
|
| 368 |
+
"source": [
|
| 369 |
+
"Python includes several built-in container types: lists, dictionaries, sets, and tuples."
|
| 370 |
+
]
|
| 371 |
+
},
|
| 372 |
+
{
|
| 373 |
+
"cell_type": "markdown",
|
| 374 |
+
"metadata": {
|
| 375 |
+
"id": "UsIWOe0LL9fn"
|
| 376 |
+
},
|
| 377 |
+
"source": [
|
| 378 |
+
"####Lists"
|
| 379 |
+
]
|
| 380 |
+
},
|
| 381 |
+
{
|
| 382 |
+
"cell_type": "markdown",
|
| 383 |
+
"metadata": {
|
| 384 |
+
"id": "wzxX7rgWL9fn"
|
| 385 |
+
},
|
| 386 |
+
"source": [
|
| 387 |
+
"A list is the Python equivalent of an array, but is resizeable and can contain elements of different types:"
|
| 388 |
+
]
|
| 389 |
+
},
|
| 390 |
+
{
|
| 391 |
+
"cell_type": "code",
|
| 392 |
+
"execution_count": null,
|
| 393 |
+
"metadata": {
|
| 394 |
+
"id": "hk3A8pPcL9fp"
|
| 395 |
+
},
|
| 396 |
+
"outputs": [],
|
| 397 |
+
"source": [
|
| 398 |
+
"xs = [3, 1, 2] # Create a list\n",
|
| 399 |
+
"print(xs, xs[2])\n",
|
| 400 |
+
"print(xs[-1]) # Negative indices count from the end of the list; prints \"2\""
|
| 401 |
+
]
|
| 402 |
+
},
|
| 403 |
+
{
|
| 404 |
+
"cell_type": "code",
|
| 405 |
+
"execution_count": null,
|
| 406 |
+
"metadata": {
|
| 407 |
+
"id": "YCjCy_0_L9ft"
|
| 408 |
+
},
|
| 409 |
+
"outputs": [],
|
| 410 |
+
"source": [
|
| 411 |
+
"xs[2] = 'foo' # Lists can be heterogeneous, i.e. it can contain elements of different types\n",
|
| 412 |
+
"print(xs)"
|
| 413 |
+
]
|
| 414 |
+
},
|
| 415 |
+
{
|
| 416 |
+
"cell_type": "code",
|
| 417 |
+
"execution_count": null,
|
| 418 |
+
"metadata": {
|
| 419 |
+
"id": "vJ0x5cF-L9fx"
|
| 420 |
+
},
|
| 421 |
+
"outputs": [],
|
| 422 |
+
"source": [
|
| 423 |
+
"xs.append('bar') # Add a new element to the end of the list\n",
|
| 424 |
+
"print(xs)"
|
| 425 |
+
]
|
| 426 |
+
},
|
| 427 |
+
{
|
| 428 |
+
"cell_type": "code",
|
| 429 |
+
"execution_count": null,
|
| 430 |
+
"metadata": {
|
| 431 |
+
"id": "cxVCNRTNL9f1"
|
| 432 |
+
},
|
| 433 |
+
"outputs": [],
|
| 434 |
+
"source": [
|
| 435 |
+
"x = xs.pop() # Remove and return the last element of the list\n",
|
| 436 |
+
"print(x, xs)"
|
| 437 |
+
]
|
| 438 |
+
},
|
| 439 |
+
{
|
| 440 |
+
"cell_type": "markdown",
|
| 441 |
+
"metadata": {
|
| 442 |
+
"id": "ilyoyO34L9f4"
|
| 443 |
+
},
|
| 444 |
+
"source": [
|
| 445 |
+
"As usual, you can find all the gory details about lists in the [documentation](https://docs.python.org/3.7/tutorial/datastructures.html#more-on-lists)."
|
| 446 |
+
]
|
| 447 |
+
},
|
| 448 |
+
{
|
| 449 |
+
"cell_type": "markdown",
|
| 450 |
+
"metadata": {
|
| 451 |
+
"id": "ovahhxd_L9f5"
|
| 452 |
+
},
|
| 453 |
+
"source": [
|
| 454 |
+
"####Slicing"
|
| 455 |
+
]
|
| 456 |
+
},
|
| 457 |
+
{
|
| 458 |
+
"cell_type": "markdown",
|
| 459 |
+
"metadata": {
|
| 460 |
+
"id": "YeSYKhv9L9f6"
|
| 461 |
+
},
|
| 462 |
+
"source": [
|
| 463 |
+
"In addition to accessing list elements one at a time, Python provides concise syntax to access sublists; this is known as slicing:"
|
| 464 |
+
]
|
| 465 |
+
},
|
| 466 |
+
{
|
| 467 |
+
"cell_type": "code",
|
| 468 |
+
"execution_count": null,
|
| 469 |
+
"metadata": {
|
| 470 |
+
"id": "ninq666bL9f6"
|
| 471 |
+
},
|
| 472 |
+
"outputs": [],
|
| 473 |
+
"source": [
|
| 474 |
+
"nums = list(range(5)) # range is a built-in function that creates a list of integers\n",
|
| 475 |
+
"print(nums) # Prints \"[0, 1, 2, 3, 4]\"\n",
|
| 476 |
+
"print(nums[2:4]) # Get a slice from index 2 to 4 (exclusive); prints \"[2, 3]\"\n",
|
| 477 |
+
"print(nums[2:]) # Get a slice from index 2 to the end; prints \"[2, 3, 4]\"\n",
|
| 478 |
+
"print(nums[:2]) # Get a slice from the start to index 2 (exclusive); prints \"[0, 1]\"\n",
|
| 479 |
+
"print(nums[:]) # Get a slice of the whole list; prints [\"0, 1, 2, 3, 4]\"\n",
|
| 480 |
+
"print(nums[:-1]) # Slice indices can be negative; prints [\"0, 1, 2, 3]\"\n",
|
| 481 |
+
"nums[2:4] = [8, 9] # Assign a new sublist to a slice\n",
|
| 482 |
+
"print(nums) # Prints \"[0, 1, 8, 9, 4]\""
|
| 483 |
+
]
|
| 484 |
+
},
|
| 485 |
+
{
|
| 486 |
+
"cell_type": "markdown",
|
| 487 |
+
"metadata": {
|
| 488 |
+
"id": "arrLCcMyL9gK"
|
| 489 |
+
},
|
| 490 |
+
"source": [
|
| 491 |
+
"####List comprehensions:"
|
| 492 |
+
]
|
| 493 |
+
},
|
| 494 |
+
{
|
| 495 |
+
"cell_type": "markdown",
|
| 496 |
+
"metadata": {
|
| 497 |
+
"id": "5Qn2jU_pL9gL"
|
| 498 |
+
},
|
| 499 |
+
"source": [
|
| 500 |
+
"When programming, frequently we want to transform one type of data into another. As a simple example, consider the following code that computes square numbers:"
|
| 501 |
+
]
|
| 502 |
+
},
|
| 503 |
+
{
|
| 504 |
+
"cell_type": "code",
|
| 505 |
+
"execution_count": null,
|
| 506 |
+
"metadata": {
|
| 507 |
+
"id": "IVNEwoMXL9gL"
|
| 508 |
+
},
|
| 509 |
+
"outputs": [],
|
| 510 |
+
"source": [
|
| 511 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 512 |
+
"squares = []\n",
|
| 513 |
+
"for x in nums:\n",
|
| 514 |
+
" squares.append(x ** 2)\n",
|
| 515 |
+
"print(squares)"
|
| 516 |
+
]
|
| 517 |
+
},
|
| 518 |
+
{
|
| 519 |
+
"cell_type": "markdown",
|
| 520 |
+
"metadata": {
|
| 521 |
+
"id": "7DmKVUFaL9gQ"
|
| 522 |
+
},
|
| 523 |
+
"source": [
|
| 524 |
+
"You can make this code simpler using a list comprehension:"
|
| 525 |
+
]
|
| 526 |
+
},
|
| 527 |
+
{
|
| 528 |
+
"cell_type": "code",
|
| 529 |
+
"execution_count": null,
|
| 530 |
+
"metadata": {
|
| 531 |
+
"id": "kZxsUfV6L9gR"
|
| 532 |
+
},
|
| 533 |
+
"outputs": [],
|
| 534 |
+
"source": [
|
| 535 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 536 |
+
"squares = [x ** 2 for x in nums]\n",
|
| 537 |
+
"print(squares)"
|
| 538 |
+
]
|
| 539 |
+
},
|
| 540 |
+
{
|
| 541 |
+
"cell_type": "markdown",
|
| 542 |
+
"metadata": {
|
| 543 |
+
"id": "-D8ARK7tL9gV"
|
| 544 |
+
},
|
| 545 |
+
"source": [
|
| 546 |
+
"List comprehensions can also contain conditions:"
|
| 547 |
+
]
|
| 548 |
+
},
|
| 549 |
+
{
|
| 550 |
+
"cell_type": "code",
|
| 551 |
+
"execution_count": null,
|
| 552 |
+
"metadata": {
|
| 553 |
+
"id": "yUtgOyyYL9gV"
|
| 554 |
+
},
|
| 555 |
+
"outputs": [],
|
| 556 |
+
"source": [
|
| 557 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 558 |
+
"even_squares = [x ** 2 for x in nums if x % 2 == 0]\n",
|
| 559 |
+
"print(even_squares)"
|
| 560 |
+
]
|
| 561 |
+
},
|
| 562 |
+
{
|
| 563 |
+
"cell_type": "markdown",
|
| 564 |
+
"metadata": {
|
| 565 |
+
"id": "H8xsUEFpL9gZ"
|
| 566 |
+
},
|
| 567 |
+
"source": [
|
| 568 |
+
"####Dictionaries"
|
| 569 |
+
]
|
| 570 |
+
},
|
| 571 |
+
{
|
| 572 |
+
"cell_type": "markdown",
|
| 573 |
+
"metadata": {
|
| 574 |
+
"id": "kkjAGMAJL9ga"
|
| 575 |
+
},
|
| 576 |
+
"source": [
|
| 577 |
+
"A dictionary stores (key, value) pairs, similar to a `Map` in Java or an object in Javascript. You can use it like this:"
|
| 578 |
+
]
|
| 579 |
+
},
|
| 580 |
+
{
|
| 581 |
+
"cell_type": "code",
|
| 582 |
+
"execution_count": null,
|
| 583 |
+
"metadata": {
|
| 584 |
+
"id": "XBYI1MrYL9gb"
|
| 585 |
+
},
|
| 586 |
+
"outputs": [],
|
| 587 |
+
"source": [
|
| 588 |
+
"d = {'cat': 'cute', 'dog': 'furry'} # Create a new dictionary with some data\n",
|
| 589 |
+
"print(d['cat']) # Get an entry from a dictionary; prints \"cute\"\n",
|
| 590 |
+
"print('cat' in d) # Check if a dictionary has a given key; prints \"True\""
|
| 591 |
+
]
|
| 592 |
+
},
|
| 593 |
+
{
|
| 594 |
+
"cell_type": "code",
|
| 595 |
+
"execution_count": null,
|
| 596 |
+
"metadata": {
|
| 597 |
+
"id": "pS7e-G-HL9gf"
|
| 598 |
+
},
|
| 599 |
+
"outputs": [],
|
| 600 |
+
"source": [
|
| 601 |
+
"d['fish'] = 'wet' # Set an entry in a dictionary\n",
|
| 602 |
+
"print(d['fish']) # Prints \"wet\""
|
| 603 |
+
]
|
| 604 |
+
},
|
| 605 |
+
{
|
| 606 |
+
"cell_type": "code",
|
| 607 |
+
"execution_count": null,
|
| 608 |
+
"metadata": {
|
| 609 |
+
"id": "tFY065ItL9gi"
|
| 610 |
+
},
|
| 611 |
+
"outputs": [],
|
| 612 |
+
"source": [
|
| 613 |
+
"print(d['monkey']) # KeyError: 'monkey' not a key of d"
|
| 614 |
+
]
|
| 615 |
+
},
|
| 616 |
+
{
|
| 617 |
+
"cell_type": "code",
|
| 618 |
+
"execution_count": null,
|
| 619 |
+
"metadata": {
|
| 620 |
+
"id": "8TjbEWqML9gl"
|
| 621 |
+
},
|
| 622 |
+
"outputs": [],
|
| 623 |
+
"source": [
|
| 624 |
+
"print(d.get('monkey', 'N/A')) # Get an element with a default; prints \"N/A\"\n",
|
| 625 |
+
"print(d.get('fish', 'N/A')) # Get an element with a default; prints \"wet\""
|
| 626 |
+
]
|
| 627 |
+
},
|
| 628 |
+
{
|
| 629 |
+
"cell_type": "code",
|
| 630 |
+
"execution_count": null,
|
| 631 |
+
"metadata": {
|
| 632 |
+
"id": "0EItdNBJL9go"
|
| 633 |
+
},
|
| 634 |
+
"outputs": [],
|
| 635 |
+
"source": [
|
| 636 |
+
"del d['fish'] # Remove an element from a dictionary\n",
|
| 637 |
+
"print(d.get('fish', 'N/A')) # \"fish\" is no longer a key; prints \"N/A\""
|
| 638 |
+
]
|
| 639 |
+
},
|
| 640 |
+
{
|
| 641 |
+
"cell_type": "markdown",
|
| 642 |
+
"metadata": {
|
| 643 |
+
"id": "wqm4dRZNL9gr"
|
| 644 |
+
},
|
| 645 |
+
"source": [
|
| 646 |
+
"You can find all you need to know about dictionaries in the [documentation](https://docs.python.org/2/library/stdtypes.html#dict)."
|
| 647 |
+
]
|
| 648 |
+
},
|
| 649 |
+
{
|
| 650 |
+
"cell_type": "markdown",
|
| 651 |
+
"metadata": {
|
| 652 |
+
"id": "IxwEqHlGL9gr"
|
| 653 |
+
},
|
| 654 |
+
"source": [
|
| 655 |
+
"It is easy to iterate over the keys in a dictionary:"
|
| 656 |
+
]
|
| 657 |
+
},
|
| 658 |
+
{
|
| 659 |
+
"cell_type": "code",
|
| 660 |
+
"execution_count": null,
|
| 661 |
+
"metadata": {
|
| 662 |
+
"id": "rYfz7ZKNL9gs"
|
| 663 |
+
},
|
| 664 |
+
"outputs": [],
|
| 665 |
+
"source": [
|
| 666 |
+
"d = {'person': 2, 'cat': 4, 'spider': 8}\n",
|
| 667 |
+
"for animal, legs in d.items():\n",
|
| 668 |
+
" print('A {} has {} legs'.format(animal, legs))"
|
| 669 |
+
]
|
| 670 |
+
},
|
| 671 |
+
{
|
| 672 |
+
"cell_type": "markdown",
|
| 673 |
+
"metadata": {
|
| 674 |
+
"id": "17sxiOpzL9gz"
|
| 675 |
+
},
|
| 676 |
+
"source": [
|
| 677 |
+
"Dictionary comprehensions: These are similar to list comprehensions, but allow you to easily construct dictionaries. For example:"
|
| 678 |
+
]
|
| 679 |
+
},
|
| 680 |
+
{
|
| 681 |
+
"cell_type": "code",
|
| 682 |
+
"execution_count": null,
|
| 683 |
+
"metadata": {
|
| 684 |
+
"id": "8PB07imLL9gz"
|
| 685 |
+
},
|
| 686 |
+
"outputs": [],
|
| 687 |
+
"source": [
|
| 688 |
+
"nums = [0, 1, 2, 3, 4]\n",
|
| 689 |
+
"even_num_to_square = {x: x ** 2 for x in nums if x % 2 == 0}\n",
|
| 690 |
+
"print(even_num_to_square)"
|
| 691 |
+
]
|
| 692 |
+
},
|
| 693 |
+
{
|
| 694 |
+
"cell_type": "markdown",
|
| 695 |
+
"metadata": {
|
| 696 |
+
"id": "V9MHfUdvL9g2"
|
| 697 |
+
},
|
| 698 |
+
"source": [
|
| 699 |
+
"####Sets"
|
| 700 |
+
]
|
| 701 |
+
},
|
| 702 |
+
{
|
| 703 |
+
"cell_type": "markdown",
|
| 704 |
+
"metadata": {
|
| 705 |
+
"id": "Rpm4UtNpL9g2"
|
| 706 |
+
},
|
| 707 |
+
"source": [
|
| 708 |
+
"A set is an unordered collection of distinct elements. As a simple example, consider the following:"
|
| 709 |
+
]
|
| 710 |
+
},
|
| 711 |
+
{
|
| 712 |
+
"cell_type": "code",
|
| 713 |
+
"execution_count": null,
|
| 714 |
+
"metadata": {
|
| 715 |
+
"id": "MmyaniLsL9g2"
|
| 716 |
+
},
|
| 717 |
+
"outputs": [],
|
| 718 |
+
"source": [
|
| 719 |
+
"animals = {'cat', 'dog'}\n",
|
| 720 |
+
"print('cat' in animals) # Check if an element is in a set; prints \"True\"\n",
|
| 721 |
+
"print('fish' in animals) # prints \"False\"\n"
|
| 722 |
+
]
|
| 723 |
+
},
|
| 724 |
+
{
|
| 725 |
+
"cell_type": "code",
|
| 726 |
+
"execution_count": null,
|
| 727 |
+
"metadata": {
|
| 728 |
+
"id": "ElJEyK86L9g6"
|
| 729 |
+
},
|
| 730 |
+
"outputs": [],
|
| 731 |
+
"source": [
|
| 732 |
+
"animals.add('fish') # Add an element to a set\n",
|
| 733 |
+
"print('fish' in animals)\n",
|
| 734 |
+
"print(len(animals)) # Number of elements in a set;"
|
| 735 |
+
]
|
| 736 |
+
},
|
| 737 |
+
{
|
| 738 |
+
"cell_type": "code",
|
| 739 |
+
"execution_count": null,
|
| 740 |
+
"metadata": {
|
| 741 |
+
"id": "5uGmrxdPL9g9"
|
| 742 |
+
},
|
| 743 |
+
"outputs": [],
|
| 744 |
+
"source": [
|
| 745 |
+
"animals.add('cat') # Adding an element that is already in the set does nothing\n",
|
| 746 |
+
"print(len(animals))\n",
|
| 747 |
+
"animals.remove('cat') # Remove an element from a set\n",
|
| 748 |
+
"print(len(animals))"
|
| 749 |
+
]
|
| 750 |
+
},
|
| 751 |
+
{
|
| 752 |
+
"cell_type": "markdown",
|
| 753 |
+
"metadata": {
|
| 754 |
+
"id": "zk2DbvLKL9g_"
|
| 755 |
+
},
|
| 756 |
+
"source": [
|
| 757 |
+
"_Loops_: Iterating over a set has the same syntax as iterating over a list; however since sets are unordered, you cannot make assumptions about the order in which you visit the elements of the set:"
|
| 758 |
+
]
|
| 759 |
+
},
|
| 760 |
+
{
|
| 761 |
+
"cell_type": "code",
|
| 762 |
+
"execution_count": null,
|
| 763 |
+
"metadata": {
|
| 764 |
+
"id": "K47KYNGyL9hA"
|
| 765 |
+
},
|
| 766 |
+
"outputs": [],
|
| 767 |
+
"source": [
|
| 768 |
+
"animals = {'cat', 'dog', 'fish'}\n",
|
| 769 |
+
"for idx, animal in enumerate(animals):\n",
|
| 770 |
+
" print('#{}: {}'.format(idx + 1, animal))"
|
| 771 |
+
]
|
| 772 |
+
},
|
| 773 |
+
{
|
| 774 |
+
"cell_type": "markdown",
|
| 775 |
+
"metadata": {
|
| 776 |
+
"id": "puq4S8buL9hC"
|
| 777 |
+
},
|
| 778 |
+
"source": [
|
| 779 |
+
"Set comprehensions: Like lists and dictionaries, we can easily construct sets using set comprehensions:"
|
| 780 |
+
]
|
| 781 |
+
},
|
| 782 |
+
{
|
| 783 |
+
"cell_type": "code",
|
| 784 |
+
"execution_count": null,
|
| 785 |
+
"metadata": {
|
| 786 |
+
"id": "iw7k90k3L9hC"
|
| 787 |
+
},
|
| 788 |
+
"outputs": [],
|
| 789 |
+
"source": [
|
| 790 |
+
"from math import sqrt\n",
|
| 791 |
+
"print({int(sqrt(x)) for x in range(30)})"
|
| 792 |
+
]
|
| 793 |
+
},
|
| 794 |
+
{
|
| 795 |
+
"cell_type": "markdown",
|
| 796 |
+
"metadata": {
|
| 797 |
+
"id": "qPsHSKB1L9hF"
|
| 798 |
+
},
|
| 799 |
+
"source": [
|
| 800 |
+
"####Tuples"
|
| 801 |
+
]
|
| 802 |
+
},
|
| 803 |
+
{
|
| 804 |
+
"cell_type": "markdown",
|
| 805 |
+
"metadata": {
|
| 806 |
+
"id": "kucc0LKVL9hG"
|
| 807 |
+
},
|
| 808 |
+
"source": [
|
| 809 |
+
"A tuple is an (immutable) ordered list of values. A tuple is in many ways similar to a list; one of the most important differences is that tuples can be used as keys in dictionaries and as elements of sets, while lists cannot. Here is a trivial example:"
|
| 810 |
+
]
|
| 811 |
+
},
|
| 812 |
+
{
|
| 813 |
+
"cell_type": "code",
|
| 814 |
+
"execution_count": null,
|
| 815 |
+
"metadata": {
|
| 816 |
+
"id": "9wHUyTKxL9hH"
|
| 817 |
+
},
|
| 818 |
+
"outputs": [],
|
| 819 |
+
"source": [
|
| 820 |
+
"d = {(x, x + 1): x for x in range(10)} # Create a dictionary with tuple keys\n",
|
| 821 |
+
"t = (5, 6) # Create a tuple\n",
|
| 822 |
+
"print(type(t))\n",
|
| 823 |
+
"print(d[t])\n",
|
| 824 |
+
"print(d[(1, 2)])"
|
| 825 |
+
]
|
| 826 |
+
},
|
| 827 |
+
{
|
| 828 |
+
"cell_type": "markdown",
|
| 829 |
+
"metadata": {
|
| 830 |
+
"id": "iFON3Tm0CfIg"
|
| 831 |
+
},
|
| 832 |
+
"source": [
|
| 833 |
+
"##Loops"
|
| 834 |
+
]
|
| 835 |
+
},
|
| 836 |
+
{
|
| 837 |
+
"cell_type": "markdown",
|
| 838 |
+
"metadata": {
|
| 839 |
+
"id": "7aXUmwT69XqK"
|
| 840 |
+
},
|
| 841 |
+
"source": [
|
| 842 |
+
"###`for` loop"
|
| 843 |
+
]
|
| 844 |
+
},
|
| 845 |
+
{
|
| 846 |
+
"cell_type": "markdown",
|
| 847 |
+
"metadata": {
|
| 848 |
+
"id": "_DYz1j6QL9f_"
|
| 849 |
+
},
|
| 850 |
+
"source": [
|
| 851 |
+
"You can loop over the elements of a list like this:"
|
| 852 |
+
]
|
| 853 |
+
},
|
| 854 |
+
{
|
| 855 |
+
"cell_type": "code",
|
| 856 |
+
"execution_count": null,
|
| 857 |
+
"metadata": {
|
| 858 |
+
"id": "4cCOysfWL9gA"
|
| 859 |
+
},
|
| 860 |
+
"outputs": [],
|
| 861 |
+
"source": [
|
| 862 |
+
"animals = ['cat', 'dog', 'monkey']\n",
|
| 863 |
+
"for animal in animals:\n",
|
| 864 |
+
" print(animal)"
|
| 865 |
+
]
|
| 866 |
+
},
|
| 867 |
+
{
|
| 868 |
+
"cell_type": "markdown",
|
| 869 |
+
"metadata": {
|
| 870 |
+
"id": "KxIaQs7pL9gE"
|
| 871 |
+
},
|
| 872 |
+
"source": [
|
| 873 |
+
"If you want access to the index of each element within the body of a loop, use the built-in `enumerate` function:"
|
| 874 |
+
]
|
| 875 |
+
},
|
| 876 |
+
{
|
| 877 |
+
"cell_type": "code",
|
| 878 |
+
"execution_count": null,
|
| 879 |
+
"metadata": {
|
| 880 |
+
"id": "JjGnDluWL9gF"
|
| 881 |
+
},
|
| 882 |
+
"outputs": [],
|
| 883 |
+
"source": [
|
| 884 |
+
"animals = ['cat', 'dog', 'monkey']\n",
|
| 885 |
+
"for idx, animal in enumerate(animals):\n",
|
| 886 |
+
" print('#{}: {}'.format(idx + 1, animal))"
|
| 887 |
+
]
|
| 888 |
+
},
|
| 889 |
+
{
|
| 890 |
+
"cell_type": "markdown",
|
| 891 |
+
"metadata": {
|
| 892 |
+
"id": "Tlf5gPRy9jfV"
|
| 893 |
+
},
|
| 894 |
+
"source": [
|
| 895 |
+
"###`range()` function"
|
| 896 |
+
]
|
| 897 |
+
},
|
| 898 |
+
{
|
| 899 |
+
"cell_type": "markdown",
|
| 900 |
+
"metadata": {
|
| 901 |
+
"id": "pzgK4H6j-PyN"
|
| 902 |
+
},
|
| 903 |
+
"source": [
|
| 904 |
+
"If you need to iterate over a sequence of numbers, the built-in function `range()` comes in handy. It generates arithmetic progressions:"
|
| 905 |
+
]
|
| 906 |
+
},
|
| 907 |
+
{
|
| 908 |
+
"cell_type": "code",
|
| 909 |
+
"execution_count": null,
|
| 910 |
+
"metadata": {
|
| 911 |
+
"id": "eNlqr8sn-TX5"
|
| 912 |
+
},
|
| 913 |
+
"outputs": [],
|
| 914 |
+
"source": [
|
| 915 |
+
"for i in range(5):\n",
|
| 916 |
+
" print(i)"
|
| 917 |
+
]
|
| 918 |
+
},
|
| 919 |
+
{
|
| 920 |
+
"cell_type": "markdown",
|
| 921 |
+
"metadata": {
|
| 922 |
+
"id": "vfigt0cE-fpp"
|
| 923 |
+
},
|
| 924 |
+
"source": [
|
| 925 |
+
"The given end point is never part of the generated sequence; `range(10)` generates 10 values, the legal indices for items of a sequence of length 10. It is possible to let the range start at another number, or to specify a different increment (even negative; sometimes this is called the ‘step’):"
|
| 926 |
+
]
|
| 927 |
+
},
|
| 928 |
+
{
|
| 929 |
+
"cell_type": "code",
|
| 930 |
+
"execution_count": null,
|
| 931 |
+
"metadata": {
|
| 932 |
+
"id": "apwUH4Ar-o4P"
|
| 933 |
+
},
|
| 934 |
+
"outputs": [],
|
| 935 |
+
"source": [
|
| 936 |
+
"print(list(range(5, 10)))\n",
|
| 937 |
+
"print(list(range(0, 10, 3)))\n",
|
| 938 |
+
"print(list(range(-10, -100, -30)))"
|
| 939 |
+
]
|
| 940 |
+
},
|
| 941 |
+
{
|
| 942 |
+
"cell_type": "markdown",
|
| 943 |
+
"metadata": {
|
| 944 |
+
"id": "mezbBTJqCmZy"
|
| 945 |
+
},
|
| 946 |
+
"source": [
|
| 947 |
+
"###`while` loop"
|
| 948 |
+
]
|
| 949 |
+
},
|
| 950 |
+
{
|
| 951 |
+
"cell_type": "markdown",
|
| 952 |
+
"metadata": {
|
| 953 |
+
"id": "A5qN9PZTCuvS"
|
| 954 |
+
},
|
| 955 |
+
"source": [
|
| 956 |
+
"With the `while` loop we can execute a set of statements as long as a condition is true."
|
| 957 |
+
]
|
| 958 |
+
},
|
| 959 |
+
{
|
| 960 |
+
"cell_type": "code",
|
| 961 |
+
"execution_count": null,
|
| 962 |
+
"metadata": {
|
| 963 |
+
"id": "T0NbKi1hCyCE"
|
| 964 |
+
},
|
| 965 |
+
"outputs": [],
|
| 966 |
+
"source": [
|
| 967 |
+
"i = 1\n",
|
| 968 |
+
"while i < 6:\n",
|
| 969 |
+
" print(i)\n",
|
| 970 |
+
" i += 1"
|
| 971 |
+
]
|
| 972 |
+
},
|
| 973 |
+
{
|
| 974 |
+
"cell_type": "markdown",
|
| 975 |
+
"metadata": {
|
| 976 |
+
"id": "uBXI2gMx9Dno"
|
| 977 |
+
},
|
| 978 |
+
"source": [
|
| 979 |
+
"## Control Flow Tools"
|
| 980 |
+
]
|
| 981 |
+
},
|
| 982 |
+
{
|
| 983 |
+
"cell_type": "markdown",
|
| 984 |
+
"metadata": {
|
| 985 |
+
"id": "q6aeyPu39PPC"
|
| 986 |
+
},
|
| 987 |
+
"source": [
|
| 988 |
+
"###`if` statement"
|
| 989 |
+
]
|
| 990 |
+
},
|
| 991 |
+
{
|
| 992 |
+
"cell_type": "markdown",
|
| 993 |
+
"metadata": {
|
| 994 |
+
"id": "eHUyBp-E_V-x"
|
| 995 |
+
},
|
| 996 |
+
"source": [
|
| 997 |
+
"Perhaps the most well-known statement type is the if statement. There can be zero or more `elif` parts, and the `else` part is optional. The keyword `elif` is short for `else if`, and is useful to avoid excessive indentation. An `if` … `elif` … `elif` … sequence is a substitute for the `switch` or `case` statements found in other languages. For example:"
|
| 998 |
+
]
|
| 999 |
+
},
|
| 1000 |
+
{
|
| 1001 |
+
"cell_type": "code",
|
| 1002 |
+
"execution_count": null,
|
| 1003 |
+
"metadata": {
|
| 1004 |
+
"id": "nD8ITrA__Z5D"
|
| 1005 |
+
},
|
| 1006 |
+
"outputs": [],
|
| 1007 |
+
"source": [
|
| 1008 |
+
"x = int(input(\"Please enter an integer: \"))\n",
|
| 1009 |
+
"if x < 0:\n",
|
| 1010 |
+
" x = 0\n",
|
| 1011 |
+
" print('Negative changed to zero')\n",
|
| 1012 |
+
"elif x == 0:\n",
|
| 1013 |
+
" print('Zero')\n",
|
| 1014 |
+
"elif x == 1:\n",
|
| 1015 |
+
" print('Single')\n",
|
| 1016 |
+
"else:\n",
|
| 1017 |
+
" print('More')"
|
| 1018 |
+
]
|
| 1019 |
+
},
|
| 1020 |
+
{
|
| 1021 |
+
"cell_type": "markdown",
|
| 1022 |
+
"metadata": {
|
| 1023 |
+
"id": "Y0CBMSRJ9sby"
|
| 1024 |
+
},
|
| 1025 |
+
"source": [
|
| 1026 |
+
"###`break` and `continue` statements"
|
| 1027 |
+
]
|
| 1028 |
+
},
|
| 1029 |
+
{
|
| 1030 |
+
"cell_type": "markdown",
|
| 1031 |
+
"metadata": {
|
| 1032 |
+
"id": "XrfS77Kzi91S"
|
| 1033 |
+
},
|
| 1034 |
+
"source": [
|
| 1035 |
+
"The `break` statement, like in C, breaks out of the innermost enclosing `for` or `while` loop.\n",
|
| 1036 |
+
"\n",
|
| 1037 |
+
"Loop statements may have an else clause; it is executed when the loop terminates through exhaustion of the iterable (with `for`) or when the condition becomes false (with `while`), but not when the loop is terminated by a `break` statement. This is exemplified by the following loop, which searches for prime numbers:"
|
| 1038 |
+
]
|
| 1039 |
+
},
|
| 1040 |
+
{
|
| 1041 |
+
"cell_type": "code",
|
| 1042 |
+
"execution_count": null,
|
| 1043 |
+
"metadata": {
|
| 1044 |
+
"id": "S2XoBEftjaXX"
|
| 1045 |
+
},
|
| 1046 |
+
"outputs": [],
|
| 1047 |
+
"source": [
|
| 1048 |
+
"for n in range(2, 10):\n",
|
| 1049 |
+
" for x in range(2, n):\n",
|
| 1050 |
+
" if n % x == 0:\n",
|
| 1051 |
+
" print(n, 'equals', x, '*', n//x)\n",
|
| 1052 |
+
" break\n",
|
| 1053 |
+
" else:\n",
|
| 1054 |
+
" # loop fell through without finding a factor\n",
|
| 1055 |
+
" print(n, 'is a prime number')"
|
| 1056 |
+
]
|
| 1057 |
+
},
|
| 1058 |
+
{
|
| 1059 |
+
"cell_type": "markdown",
|
| 1060 |
+
"metadata": {
|
| 1061 |
+
"id": "b5pf2Vdbkl0_"
|
| 1062 |
+
},
|
| 1063 |
+
"source": [
|
| 1064 |
+
"The `continue` statement, also borrowed from C, continues with the next iteration of the loop:"
|
| 1065 |
+
]
|
| 1066 |
+
},
|
| 1067 |
+
{
|
| 1068 |
+
"cell_type": "code",
|
| 1069 |
+
"execution_count": null,
|
| 1070 |
+
"metadata": {
|
| 1071 |
+
"id": "swr6-rEwksE2"
|
| 1072 |
+
},
|
| 1073 |
+
"outputs": [],
|
| 1074 |
+
"source": [
|
| 1075 |
+
"for num in range(2, 10):\n",
|
| 1076 |
+
" if num % 2 == 0:\n",
|
| 1077 |
+
" print(\"Found an even number\", num)\n",
|
| 1078 |
+
" continue\n",
|
| 1079 |
+
" print(\"Found an odd number\", num)"
|
| 1080 |
+
]
|
| 1081 |
+
},
|
| 1082 |
+
{
|
| 1083 |
+
"cell_type": "markdown",
|
| 1084 |
+
"metadata": {
|
| 1085 |
+
"id": "JVKNnslx95_d"
|
| 1086 |
+
},
|
| 1087 |
+
"source": [
|
| 1088 |
+
"###`pass` statement"
|
| 1089 |
+
]
|
| 1090 |
+
},
|
| 1091 |
+
{
|
| 1092 |
+
"cell_type": "markdown",
|
| 1093 |
+
"metadata": {
|
| 1094 |
+
"id": "-dZk28jllX6D"
|
| 1095 |
+
},
|
| 1096 |
+
"source": [
|
| 1097 |
+
"The `pass` statement does nothing. It can be used when a statement is required syntactically but the program requires no action. For example:"
|
| 1098 |
+
]
|
| 1099 |
+
},
|
| 1100 |
+
{
|
| 1101 |
+
"cell_type": "code",
|
| 1102 |
+
"execution_count": null,
|
| 1103 |
+
"metadata": {
|
| 1104 |
+
"id": "DnJfJPkTldUz"
|
| 1105 |
+
},
|
| 1106 |
+
"outputs": [],
|
| 1107 |
+
"source": [
|
| 1108 |
+
"while True:\n",
|
| 1109 |
+
" pass # Busy-wait for keyboard interrupt. Please press the stop button to stop execution."
|
| 1110 |
+
]
|
| 1111 |
+
},
|
| 1112 |
+
{
|
| 1113 |
+
"cell_type": "markdown",
|
| 1114 |
+
"metadata": {
|
| 1115 |
+
"id": "-tWb6by_lkn3"
|
| 1116 |
+
},
|
| 1117 |
+
"source": [
|
| 1118 |
+
"This is commonly used for creating minimal classes:"
|
| 1119 |
+
]
|
| 1120 |
+
},
|
| 1121 |
+
{
|
| 1122 |
+
"cell_type": "code",
|
| 1123 |
+
"execution_count": null,
|
| 1124 |
+
"metadata": {
|
| 1125 |
+
"id": "57_9LZkplsSz"
|
| 1126 |
+
},
|
| 1127 |
+
"outputs": [],
|
| 1128 |
+
"source": [
|
| 1129 |
+
"class MyEmptyClass:\n",
|
| 1130 |
+
" pass"
|
| 1131 |
+
]
|
| 1132 |
+
},
|
| 1133 |
+
{
|
| 1134 |
+
"cell_type": "markdown",
|
| 1135 |
+
"metadata": {
|
| 1136 |
+
"id": "45jdssyFlxqo"
|
| 1137 |
+
},
|
| 1138 |
+
"source": [
|
| 1139 |
+
"Another place `pass` can be used is as a place-holder for a function or conditional body when you are working on new code, allowing you to keep thinking at a more abstract level. The `pass` is silently ignored:"
|
| 1140 |
+
]
|
| 1141 |
+
},
|
| 1142 |
+
{
|
| 1143 |
+
"cell_type": "code",
|
| 1144 |
+
"execution_count": null,
|
| 1145 |
+
"metadata": {
|
| 1146 |
+
"id": "0r9Dikptl4-1"
|
| 1147 |
+
},
|
| 1148 |
+
"outputs": [],
|
| 1149 |
+
"source": [
|
| 1150 |
+
"def initlog(*args):\n",
|
| 1151 |
+
" pass # Remember to implement this!"
|
| 1152 |
+
]
|
| 1153 |
+
},
|
| 1154 |
+
{
|
| 1155 |
+
"cell_type": "markdown",
|
| 1156 |
+
"metadata": {
|
| 1157 |
+
"id": "AXA4jrEOL9hM"
|
| 1158 |
+
},
|
| 1159 |
+
"source": [
|
| 1160 |
+
"###Functions"
|
| 1161 |
+
]
|
| 1162 |
+
},
|
| 1163 |
+
{
|
| 1164 |
+
"cell_type": "markdown",
|
| 1165 |
+
"metadata": {
|
| 1166 |
+
"id": "WaRms-QfL9hN"
|
| 1167 |
+
},
|
| 1168 |
+
"source": [
|
| 1169 |
+
"Python functions are defined using the `def` keyword. For example:"
|
| 1170 |
+
]
|
| 1171 |
+
},
|
| 1172 |
+
{
|
| 1173 |
+
"cell_type": "code",
|
| 1174 |
+
"execution_count": null,
|
| 1175 |
+
"metadata": {
|
| 1176 |
+
"id": "kiMDUr58L9hN"
|
| 1177 |
+
},
|
| 1178 |
+
"outputs": [],
|
| 1179 |
+
"source": [
|
| 1180 |
+
"def sign(x):\n",
|
| 1181 |
+
" if x > 0:\n",
|
| 1182 |
+
" return 'positive'\n",
|
| 1183 |
+
" elif x < 0:\n",
|
| 1184 |
+
" return 'negative'\n",
|
| 1185 |
+
" else:\n",
|
| 1186 |
+
" return 'zero'\n",
|
| 1187 |
+
"\n",
|
| 1188 |
+
"for x in [-1, 0, 1]:\n",
|
| 1189 |
+
" print(sign(x))"
|
| 1190 |
+
]
|
| 1191 |
+
},
|
| 1192 |
+
{
|
| 1193 |
+
"cell_type": "markdown",
|
| 1194 |
+
"metadata": {
|
| 1195 |
+
"id": "U-QJFt8TL9hR"
|
| 1196 |
+
},
|
| 1197 |
+
"source": [
|
| 1198 |
+
"We will often define functions to take optional keyword arguments, like this:"
|
| 1199 |
+
]
|
| 1200 |
+
},
|
| 1201 |
+
{
|
| 1202 |
+
"cell_type": "code",
|
| 1203 |
+
"execution_count": null,
|
| 1204 |
+
"metadata": {
|
| 1205 |
+
"id": "PfsZ3DazL9hR"
|
| 1206 |
+
},
|
| 1207 |
+
"outputs": [],
|
| 1208 |
+
"source": [
|
| 1209 |
+
"def hello(name, loud=False):\n",
|
| 1210 |
+
" if loud:\n",
|
| 1211 |
+
" print('HELLO, {}'.format(name.upper()))\n",
|
| 1212 |
+
" else:\n",
|
| 1213 |
+
" print('Hello, {}!'.format(name))\n",
|
| 1214 |
+
"\n",
|
| 1215 |
+
"hello('Bob')\n",
|
| 1216 |
+
"hello('Fred', loud=True)"
|
| 1217 |
+
]
|
| 1218 |
+
},
|
| 1219 |
+
{
|
| 1220 |
+
"cell_type": "markdown",
|
| 1221 |
+
"metadata": {
|
| 1222 |
+
"id": "ObA9PRtQL9hT"
|
| 1223 |
+
},
|
| 1224 |
+
"source": [
|
| 1225 |
+
"###Classes"
|
| 1226 |
+
]
|
| 1227 |
+
},
|
| 1228 |
+
{
|
| 1229 |
+
"cell_type": "markdown",
|
| 1230 |
+
"metadata": {
|
| 1231 |
+
"id": "hAzL_lTkL9hU"
|
| 1232 |
+
},
|
| 1233 |
+
"source": [
|
| 1234 |
+
"In object-oriented programming, a class is a template definition of the methods and variables in a particular kind of object. Thus, an object is a specific instance of a class; it contains real values instead of variables. For more details, on class in object oriented programming, please have a look on this [link](https://www.w3schools.com/java/java_oop.asp). The syntax for defining classes in Python is straightforward and can be done as follows."
|
| 1235 |
+
]
|
| 1236 |
+
},
|
| 1237 |
+
{
|
| 1238 |
+
"cell_type": "code",
|
| 1239 |
+
"execution_count": null,
|
| 1240 |
+
"metadata": {
|
| 1241 |
+
"id": "RWdbaGigL9hU"
|
| 1242 |
+
},
|
| 1243 |
+
"outputs": [],
|
| 1244 |
+
"source": [
|
| 1245 |
+
"class Greeter:\n",
|
| 1246 |
+
"\n",
|
| 1247 |
+
" # Constructor\n",
|
| 1248 |
+
" def __init__(self, name):\n",
|
| 1249 |
+
" self.name = name # Create an instance variable\n",
|
| 1250 |
+
"\n",
|
| 1251 |
+
" # Instance method\n",
|
| 1252 |
+
" def greet(self, loud=False):\n",
|
| 1253 |
+
" if loud:\n",
|
| 1254 |
+
" print('HELLO, {}'.format(self.name.upper()))\n",
|
| 1255 |
+
" else:\n",
|
| 1256 |
+
" print('Hello, {}!'.format(self.name))\n",
|
| 1257 |
+
"\n",
|
| 1258 |
+
"g = Greeter('Fred') # Construct an instance of the Greeter class\n",
|
| 1259 |
+
"g.greet() # Call an instance method; prints \"Hello, Fred\"\n",
|
| 1260 |
+
"g.greet(loud=True) # Call an instance method; prints \"HELLO, FRED!\""
|
| 1261 |
+
]
|
| 1262 |
+
},
|
| 1263 |
+
{
|
| 1264 |
+
"cell_type": "markdown",
|
| 1265 |
+
"metadata": {
|
| 1266 |
+
"id": "3cfrOV4dL9hW"
|
| 1267 |
+
},
|
| 1268 |
+
"source": [
|
| 1269 |
+
"##Numpy"
|
| 1270 |
+
]
|
| 1271 |
+
},
|
| 1272 |
+
{
|
| 1273 |
+
"cell_type": "markdown",
|
| 1274 |
+
"metadata": {
|
| 1275 |
+
"id": "fY12nHhyL9hX"
|
| 1276 |
+
},
|
| 1277 |
+
"source": [
|
| 1278 |
+
"Numpy is the core library for scientific computing in Python. It provides a high-performance multidimensional array object, and tools for working with these arrays. If you are already familiar with MATLAB, you might find this [tutorial](http://wiki.scipy.org/NumPy_for_Matlab_Users) useful to get started with Numpy. To use Numpy, we first need to import the `numpy` package."
|
| 1279 |
+
]
|
| 1280 |
+
},
|
| 1281 |
+
{
|
| 1282 |
+
"cell_type": "markdown",
|
| 1283 |
+
"metadata": {
|
| 1284 |
+
"id": "2_lpLqwZpd-4"
|
| 1285 |
+
},
|
| 1286 |
+
"source": [
|
| 1287 |
+
"### Importing a package"
|
| 1288 |
+
]
|
| 1289 |
+
},
|
| 1290 |
+
{
|
| 1291 |
+
"cell_type": "markdown",
|
| 1292 |
+
"metadata": {
|
| 1293 |
+
"id": "hMmlsjljBbVE"
|
| 1294 |
+
},
|
| 1295 |
+
"source": [
|
| 1296 |
+
"In Python, a package or a module can be imported in many different ways, some of which are shown below. For more details, please have a look on this [documentation](https://docs.python.org/3/tutorial/modules.html#more-on-modules).\n",
|
| 1297 |
+
"\n",
|
| 1298 |
+
"\n",
|
| 1299 |
+
"```\n",
|
| 1300 |
+
"import numpy # import numpy, one can use it as numpy\n",
|
| 1301 |
+
"import numpy as np # import numpy and call it np\n",
|
| 1302 |
+
"from numpy import * # import all the modules from numpy\n",
|
| 1303 |
+
"from numpy import sum # import the \"sum\" function from numpy\n",
|
| 1304 |
+
"```\n",
|
| 1305 |
+
"\n"
|
| 1306 |
+
]
|
| 1307 |
+
},
|
| 1308 |
+
{
|
| 1309 |
+
"cell_type": "code",
|
| 1310 |
+
"execution_count": null,
|
| 1311 |
+
"metadata": {
|
| 1312 |
+
"id": "58QdX8BLL9hZ"
|
| 1313 |
+
},
|
| 1314 |
+
"outputs": [],
|
| 1315 |
+
"source": [
|
| 1316 |
+
"import numpy as np # import numpy and call it np. So the sum function of numpy can be called as np.sum()"
|
| 1317 |
+
]
|
| 1318 |
+
},
|
| 1319 |
+
{
|
| 1320 |
+
"cell_type": "markdown",
|
| 1321 |
+
"metadata": {
|
| 1322 |
+
"id": "DDx6v1EdL9hb"
|
| 1323 |
+
},
|
| 1324 |
+
"source": [
|
| 1325 |
+
"###Arrays"
|
| 1326 |
+
]
|
| 1327 |
+
},
|
| 1328 |
+
{
|
| 1329 |
+
"cell_type": "markdown",
|
| 1330 |
+
"metadata": {
|
| 1331 |
+
"id": "f-Zv3f7LL9hc"
|
| 1332 |
+
},
|
| 1333 |
+
"source": [
|
| 1334 |
+
"A numpy array is a grid of values, all of the same type, and is indexed by a tuple of nonnegative integers. The number of dimensions is the rank of the array; the shape of an array is a tuple of integers giving the size of the array along each dimension."
|
| 1335 |
+
]
|
| 1336 |
+
},
|
| 1337 |
+
{
|
| 1338 |
+
"cell_type": "markdown",
|
| 1339 |
+
"metadata": {
|
| 1340 |
+
"id": "_eMTRnZRL9hc"
|
| 1341 |
+
},
|
| 1342 |
+
"source": [
|
| 1343 |
+
"We can initialize numpy arrays from nested Python lists, and access elements using square brackets:"
|
| 1344 |
+
]
|
| 1345 |
+
},
|
| 1346 |
+
{
|
| 1347 |
+
"cell_type": "code",
|
| 1348 |
+
"execution_count": null,
|
| 1349 |
+
"metadata": {
|
| 1350 |
+
"id": "-l3JrGxCL9hc"
|
| 1351 |
+
},
|
| 1352 |
+
"outputs": [],
|
| 1353 |
+
"source": [
|
| 1354 |
+
"a = np.array([1, 2, 3]) # Create a rank 1 array\n",
|
| 1355 |
+
"print(type(a), a.shape, a[0], a[1], a[2])\n",
|
| 1356 |
+
"a[0] = 5 # Change an element of the array\n",
|
| 1357 |
+
"print(a)"
|
| 1358 |
+
]
|
| 1359 |
+
},
|
| 1360 |
+
{
|
| 1361 |
+
"cell_type": "code",
|
| 1362 |
+
"execution_count": null,
|
| 1363 |
+
"metadata": {
|
| 1364 |
+
"id": "ma6mk-kdL9hh"
|
| 1365 |
+
},
|
| 1366 |
+
"outputs": [],
|
| 1367 |
+
"source": [
|
| 1368 |
+
"b = np.array([[1,2,3],[4,5,6]]) # Create a rank 2 array\n",
|
| 1369 |
+
"print(b)"
|
| 1370 |
+
]
|
| 1371 |
+
},
|
| 1372 |
+
{
|
| 1373 |
+
"cell_type": "code",
|
| 1374 |
+
"execution_count": null,
|
| 1375 |
+
"metadata": {
|
| 1376 |
+
"id": "ymfSHAwtL9hj"
|
| 1377 |
+
},
|
| 1378 |
+
"outputs": [],
|
| 1379 |
+
"source": [
|
| 1380 |
+
"print(b.shape)\n",
|
| 1381 |
+
"print(b[0, 0], b[0, 1], b[1, 0])"
|
| 1382 |
+
]
|
| 1383 |
+
},
|
| 1384 |
+
{
|
| 1385 |
+
"cell_type": "markdown",
|
| 1386 |
+
"metadata": {
|
| 1387 |
+
"id": "F2qwdyvuL9hn"
|
| 1388 |
+
},
|
| 1389 |
+
"source": [
|
| 1390 |
+
"Numpy also provides many functions to create arrays:"
|
| 1391 |
+
]
|
| 1392 |
+
},
|
| 1393 |
+
{
|
| 1394 |
+
"cell_type": "code",
|
| 1395 |
+
"execution_count": null,
|
| 1396 |
+
"metadata": {
|
| 1397 |
+
"id": "mVTN_EBqL9hn"
|
| 1398 |
+
},
|
| 1399 |
+
"outputs": [],
|
| 1400 |
+
"source": [
|
| 1401 |
+
"a = np.zeros((2,2)) # Create an array of all zeros\n",
|
| 1402 |
+
"print(a)"
|
| 1403 |
+
]
|
| 1404 |
+
},
|
| 1405 |
+
{
|
| 1406 |
+
"cell_type": "code",
|
| 1407 |
+
"execution_count": null,
|
| 1408 |
+
"metadata": {
|
| 1409 |
+
"id": "skiKlNmlL9h5"
|
| 1410 |
+
},
|
| 1411 |
+
"outputs": [],
|
| 1412 |
+
"source": [
|
| 1413 |
+
"b = np.ones((1,2)) # Create an array of all ones\n",
|
| 1414 |
+
"print(b)"
|
| 1415 |
+
]
|
| 1416 |
+
},
|
| 1417 |
+
{
|
| 1418 |
+
"cell_type": "code",
|
| 1419 |
+
"execution_count": null,
|
| 1420 |
+
"metadata": {
|
| 1421 |
+
"id": "HtFsr03bL9h7"
|
| 1422 |
+
},
|
| 1423 |
+
"outputs": [],
|
| 1424 |
+
"source": [
|
| 1425 |
+
"c = np.full((2,2), 7) # Create a constant array\n",
|
| 1426 |
+
"print(c)"
|
| 1427 |
+
]
|
| 1428 |
+
},
|
| 1429 |
+
{
|
| 1430 |
+
"cell_type": "code",
|
| 1431 |
+
"execution_count": null,
|
| 1432 |
+
"metadata": {
|
| 1433 |
+
"id": "-QcALHvkL9h9"
|
| 1434 |
+
},
|
| 1435 |
+
"outputs": [],
|
| 1436 |
+
"source": [
|
| 1437 |
+
"d = np.eye(2) # Create a 2x2 identity matrix\n",
|
| 1438 |
+
"print(d)"
|
| 1439 |
+
]
|
| 1440 |
+
},
|
| 1441 |
+
{
|
| 1442 |
+
"cell_type": "code",
|
| 1443 |
+
"execution_count": null,
|
| 1444 |
+
"metadata": {
|
| 1445 |
+
"id": "RCpaYg9qL9iA"
|
| 1446 |
+
},
|
| 1447 |
+
"outputs": [],
|
| 1448 |
+
"source": [
|
| 1449 |
+
"e = np.random.random((2,2)) # Create an array filled with random values\n",
|
| 1450 |
+
"print(e)"
|
| 1451 |
+
]
|
| 1452 |
+
},
|
| 1453 |
+
{
|
| 1454 |
+
"cell_type": "markdown",
|
| 1455 |
+
"metadata": {
|
| 1456 |
+
"id": "jI5qcSDfL9iC"
|
| 1457 |
+
},
|
| 1458 |
+
"source": [
|
| 1459 |
+
"###Array indexing"
|
| 1460 |
+
]
|
| 1461 |
+
},
|
| 1462 |
+
{
|
| 1463 |
+
"cell_type": "markdown",
|
| 1464 |
+
"metadata": {
|
| 1465 |
+
"id": "M-E4MUeVL9iC"
|
| 1466 |
+
},
|
| 1467 |
+
"source": [
|
| 1468 |
+
"Numpy offers several ways to index into arrays.\n",
|
| 1469 |
+
"\n",
|
| 1470 |
+
"Slicing: Similar to Python lists, numpy arrays can be sliced. Since arrays may be multidimensional, you must specify a slice for each dimension of the array:"
|
| 1471 |
+
]
|
| 1472 |
+
},
|
| 1473 |
+
{
|
| 1474 |
+
"cell_type": "code",
|
| 1475 |
+
"execution_count": null,
|
| 1476 |
+
"metadata": {
|
| 1477 |
+
"id": "wLWA0udwL9iD"
|
| 1478 |
+
},
|
| 1479 |
+
"outputs": [],
|
| 1480 |
+
"source": [
|
| 1481 |
+
"import numpy as np\n",
|
| 1482 |
+
"\n",
|
| 1483 |
+
"# Create the following rank 2 array with shape (3, 4)\n",
|
| 1484 |
+
"# [[ 1 2 3 4]\n",
|
| 1485 |
+
"# [ 5 6 7 8]\n",
|
| 1486 |
+
"# [ 9 10 11 12]]\n",
|
| 1487 |
+
"a = np.array([[1,2,3,4], [5,6,7,8], [9,10,11,12]])\n",
|
| 1488 |
+
"\n",
|
| 1489 |
+
"# Use slicing to pull out the subarray consisting of the first 2 rows\n",
|
| 1490 |
+
"# and columns 1 and 2; b is the following array of shape (2, 2):\n",
|
| 1491 |
+
"# [[2 3]\n",
|
| 1492 |
+
"# [6 7]]\n",
|
| 1493 |
+
"b = a[:2, 1:3]\n",
|
| 1494 |
+
"print(b)"
|
| 1495 |
+
]
|
| 1496 |
+
},
|
| 1497 |
+
{
|
| 1498 |
+
"cell_type": "markdown",
|
| 1499 |
+
"metadata": {
|
| 1500 |
+
"id": "KahhtZKYL9iF"
|
| 1501 |
+
},
|
| 1502 |
+
"source": [
|
| 1503 |
+
"A slice of an array is a view into the same data, so modifying it will modify the original array."
|
| 1504 |
+
]
|
| 1505 |
+
},
|
| 1506 |
+
{
|
| 1507 |
+
"cell_type": "code",
|
| 1508 |
+
"execution_count": null,
|
| 1509 |
+
"metadata": {
|
| 1510 |
+
"id": "1kmtaFHuL9iG"
|
| 1511 |
+
},
|
| 1512 |
+
"outputs": [],
|
| 1513 |
+
"source": [
|
| 1514 |
+
"print(a[0, 1])\n",
|
| 1515 |
+
"b[0, 0] = 77 # b[0, 0] is the same piece of data as a[0, 1]\n",
|
| 1516 |
+
"print(a[0, 1])"
|
| 1517 |
+
]
|
| 1518 |
+
},
|
| 1519 |
+
{
|
| 1520 |
+
"cell_type": "markdown",
|
| 1521 |
+
"metadata": {
|
| 1522 |
+
"id": "_Zcf3zi-L9iI"
|
| 1523 |
+
},
|
| 1524 |
+
"source": [
|
| 1525 |
+
"You can also mix integer indexing with slice indexing. However, doing so will yield an array of lower rank than the original array. Note that this is quite different from the way that MATLAB handles array slicing:"
|
| 1526 |
+
]
|
| 1527 |
+
},
|
| 1528 |
+
{
|
| 1529 |
+
"cell_type": "code",
|
| 1530 |
+
"execution_count": null,
|
| 1531 |
+
"metadata": {
|
| 1532 |
+
"id": "G6lfbPuxL9iJ"
|
| 1533 |
+
},
|
| 1534 |
+
"outputs": [],
|
| 1535 |
+
"source": [
|
| 1536 |
+
"# Create the following rank 2 array with shape (3, 4)\n",
|
| 1537 |
+
"a = np.array([[1,2,3,4], [5,6,7,8], [9,10,11,12]])\n",
|
| 1538 |
+
"print(a)"
|
| 1539 |
+
]
|
| 1540 |
+
},
|
| 1541 |
+
{
|
| 1542 |
+
"cell_type": "markdown",
|
| 1543 |
+
"metadata": {
|
| 1544 |
+
"id": "NCye3NXhL9iL"
|
| 1545 |
+
},
|
| 1546 |
+
"source": [
|
| 1547 |
+
"Two ways of accessing the data in the middle row of the array.\n",
|
| 1548 |
+
"Mixing integer indexing with slices yields an array of lower rank,\n",
|
| 1549 |
+
"while using only slices yields an array of the same rank as the\n",
|
| 1550 |
+
"original array:"
|
| 1551 |
+
]
|
| 1552 |
+
},
|
| 1553 |
+
{
|
| 1554 |
+
"cell_type": "code",
|
| 1555 |
+
"execution_count": null,
|
| 1556 |
+
"metadata": {
|
| 1557 |
+
"id": "EOiEMsmNL9iL"
|
| 1558 |
+
},
|
| 1559 |
+
"outputs": [],
|
| 1560 |
+
"source": [
|
| 1561 |
+
"row_r1 = a[1, :] # Rank 1 view of the second row of a\n",
|
| 1562 |
+
"row_r2 = a[1:2, :] # Rank 2 view of the second row of a\n",
|
| 1563 |
+
"row_r3 = a[[1], :] # Rank 2 view of the second row of a\n",
|
| 1564 |
+
"print(row_r1, row_r1.shape)\n",
|
| 1565 |
+
"print(row_r2, row_r2.shape)\n",
|
| 1566 |
+
"print(row_r3, row_r3.shape)"
|
| 1567 |
+
]
|
| 1568 |
+
},
|
| 1569 |
+
{
|
| 1570 |
+
"cell_type": "code",
|
| 1571 |
+
"execution_count": null,
|
| 1572 |
+
"metadata": {
|
| 1573 |
+
"id": "JXu73pfDL9iN"
|
| 1574 |
+
},
|
| 1575 |
+
"outputs": [],
|
| 1576 |
+
"source": [
|
| 1577 |
+
"# We can make the same distinction when accessing columns of an array:\n",
|
| 1578 |
+
"col_r1 = a[:, 1]\n",
|
| 1579 |
+
"col_r2 = a[:, 1:2]\n",
|
| 1580 |
+
"print(col_r1, col_r1.shape)\n",
|
| 1581 |
+
"print()\n",
|
| 1582 |
+
"print(col_r2, col_r2.shape)"
|
| 1583 |
+
]
|
| 1584 |
+
},
|
| 1585 |
+
{
|
| 1586 |
+
"cell_type": "markdown",
|
| 1587 |
+
"metadata": {
|
| 1588 |
+
"id": "VP3916bOL9iP"
|
| 1589 |
+
},
|
| 1590 |
+
"source": [
|
| 1591 |
+
"Integer array indexing: When you index into numpy arrays using slicing, the resulting array view will always be a subarray of the original array. In contrast, integer array indexing allows you to construct arbitrary arrays using the data from another array. Here is an example:"
|
| 1592 |
+
]
|
| 1593 |
+
},
|
| 1594 |
+
{
|
| 1595 |
+
"cell_type": "code",
|
| 1596 |
+
"execution_count": null,
|
| 1597 |
+
"metadata": {
|
| 1598 |
+
"id": "TBnWonIDL9iP"
|
| 1599 |
+
},
|
| 1600 |
+
"outputs": [],
|
| 1601 |
+
"source": [
|
| 1602 |
+
"a = np.array([[1,2], [3, 4], [5, 6]])\n",
|
| 1603 |
+
"\n",
|
| 1604 |
+
"# An example of integer array indexing.\n",
|
| 1605 |
+
"# The returned array will have shape (3,) and\n",
|
| 1606 |
+
"print(a[[0, 1, 2], [0, 1, 0]])\n",
|
| 1607 |
+
"\n",
|
| 1608 |
+
"# The above example of integer array indexing is equivalent to this:\n",
|
| 1609 |
+
"print(np.array([a[0, 0], a[1, 1], a[2, 0]]))"
|
| 1610 |
+
]
|
| 1611 |
+
},
|
| 1612 |
+
{
|
| 1613 |
+
"cell_type": "code",
|
| 1614 |
+
"execution_count": null,
|
| 1615 |
+
"metadata": {
|
| 1616 |
+
"id": "n7vuati-L9iR"
|
| 1617 |
+
},
|
| 1618 |
+
"outputs": [],
|
| 1619 |
+
"source": [
|
| 1620 |
+
"# When using integer array indexing, you can reuse the same\n",
|
| 1621 |
+
"# element from the source array:\n",
|
| 1622 |
+
"print(a[[0, 0], [1, 1]])\n",
|
| 1623 |
+
"\n",
|
| 1624 |
+
"# Equivalent to the previous integer array indexing example\n",
|
| 1625 |
+
"print(np.array([a[0, 1], a[0, 1]]))"
|
| 1626 |
+
]
|
| 1627 |
+
},
|
| 1628 |
+
{
|
| 1629 |
+
"cell_type": "markdown",
|
| 1630 |
+
"metadata": {
|
| 1631 |
+
"id": "kaipSLafL9iU"
|
| 1632 |
+
},
|
| 1633 |
+
"source": [
|
| 1634 |
+
"One useful trick with integer array indexing is selecting or mutating one element from each row of a matrix:"
|
| 1635 |
+
]
|
| 1636 |
+
},
|
| 1637 |
+
{
|
| 1638 |
+
"cell_type": "code",
|
| 1639 |
+
"execution_count": null,
|
| 1640 |
+
"metadata": {
|
| 1641 |
+
"id": "ehqsV7TXL9iU"
|
| 1642 |
+
},
|
| 1643 |
+
"outputs": [],
|
| 1644 |
+
"source": [
|
| 1645 |
+
"# Create a new array from which we will select elements\n",
|
| 1646 |
+
"a = np.array([[1,2,3], [4,5,6], [7,8,9], [10, 11, 12]])\n",
|
| 1647 |
+
"print(a)"
|
| 1648 |
+
]
|
| 1649 |
+
},
|
| 1650 |
+
{
|
| 1651 |
+
"cell_type": "code",
|
| 1652 |
+
"execution_count": null,
|
| 1653 |
+
"metadata": {
|
| 1654 |
+
"id": "pAPOoqy5L9iV"
|
| 1655 |
+
},
|
| 1656 |
+
"outputs": [],
|
| 1657 |
+
"source": [
|
| 1658 |
+
"# Create an array of indices\n",
|
| 1659 |
+
"b = np.array([0, 2, 0, 1])\n",
|
| 1660 |
+
"\n",
|
| 1661 |
+
"# Select one element from each row of a using the indices in b\n",
|
| 1662 |
+
"print(a[np.arange(4), b]) # Prints \"[ 1 6 7 11]\""
|
| 1663 |
+
]
|
| 1664 |
+
},
|
| 1665 |
+
{
|
| 1666 |
+
"cell_type": "code",
|
| 1667 |
+
"execution_count": null,
|
| 1668 |
+
"metadata": {
|
| 1669 |
+
"id": "6v1PdI1DL9ib"
|
| 1670 |
+
},
|
| 1671 |
+
"outputs": [],
|
| 1672 |
+
"source": [
|
| 1673 |
+
"# Mutate one element from each row of a using the indices in b\n",
|
| 1674 |
+
"a[np.arange(4), b] += 10\n",
|
| 1675 |
+
"print(a)"
|
| 1676 |
+
]
|
| 1677 |
+
},
|
| 1678 |
+
{
|
| 1679 |
+
"cell_type": "markdown",
|
| 1680 |
+
"metadata": {
|
| 1681 |
+
"id": "kaE8dBGgL9id"
|
| 1682 |
+
},
|
| 1683 |
+
"source": [
|
| 1684 |
+
"Boolean array indexing: Boolean array indexing lets you pick out arbitrary elements of an array. Frequently this type of indexing is used to select the elements of an array that satisfy some condition. Here is an example:"
|
| 1685 |
+
]
|
| 1686 |
+
},
|
| 1687 |
+
{
|
| 1688 |
+
"cell_type": "code",
|
| 1689 |
+
"execution_count": null,
|
| 1690 |
+
"metadata": {
|
| 1691 |
+
"id": "32PusjtKL9id"
|
| 1692 |
+
},
|
| 1693 |
+
"outputs": [],
|
| 1694 |
+
"source": [
|
| 1695 |
+
"import numpy as np\n",
|
| 1696 |
+
"\n",
|
| 1697 |
+
"a = np.array([[1,2], [3, 4], [5, 6]])\n",
|
| 1698 |
+
"\n",
|
| 1699 |
+
"bool_idx = (a > 2) # Find the elements of a that are bigger than 2;\n",
|
| 1700 |
+
" # this returns a numpy array of Booleans of the same\n",
|
| 1701 |
+
" # shape as a, where each slot of bool_idx tells\n",
|
| 1702 |
+
" # whether that element of a is > 2.\n",
|
| 1703 |
+
"\n",
|
| 1704 |
+
"print(bool_idx)"
|
| 1705 |
+
]
|
| 1706 |
+
},
|
| 1707 |
+
{
|
| 1708 |
+
"cell_type": "code",
|
| 1709 |
+
"execution_count": null,
|
| 1710 |
+
"metadata": {
|
| 1711 |
+
"id": "cb2IRMXaL9if"
|
| 1712 |
+
},
|
| 1713 |
+
"outputs": [],
|
| 1714 |
+
"source": [
|
| 1715 |
+
"# We use boolean array indexing to construct a rank 1 array\n",
|
| 1716 |
+
"# consisting of the elements of a corresponding to the True values\n",
|
| 1717 |
+
"# of bool_idx\n",
|
| 1718 |
+
"print(a[bool_idx])\n",
|
| 1719 |
+
"\n",
|
| 1720 |
+
"# We can do all of the above in a single concise statement:\n",
|
| 1721 |
+
"print(a[a > 2])"
|
| 1722 |
+
]
|
| 1723 |
+
},
|
| 1724 |
+
{
|
| 1725 |
+
"cell_type": "markdown",
|
| 1726 |
+
"metadata": {
|
| 1727 |
+
"id": "CdofMonAL9ih"
|
| 1728 |
+
},
|
| 1729 |
+
"source": [
|
| 1730 |
+
"For brevity we have left out a lot of details about numpy array indexing; if you want to know more you should read the documentation."
|
| 1731 |
+
]
|
| 1732 |
+
},
|
| 1733 |
+
{
|
| 1734 |
+
"cell_type": "markdown",
|
| 1735 |
+
"metadata": {
|
| 1736 |
+
"id": "jTctwqdQL9ih"
|
| 1737 |
+
},
|
| 1738 |
+
"source": [
|
| 1739 |
+
"###Datatypes"
|
| 1740 |
+
]
|
| 1741 |
+
},
|
| 1742 |
+
{
|
| 1743 |
+
"cell_type": "markdown",
|
| 1744 |
+
"metadata": {
|
| 1745 |
+
"id": "kSZQ1WkIL9ih"
|
| 1746 |
+
},
|
| 1747 |
+
"source": [
|
| 1748 |
+
"Every numpy array is a grid of elements of the same type. Numpy provides a large set of numeric datatypes that you can use to construct arrays. Numpy tries to guess a datatype when you create an array, but functions that construct arrays usually also include an optional argument to explicitly specify the datatype. Here is an example:"
|
| 1749 |
+
]
|
| 1750 |
+
},
|
| 1751 |
+
{
|
| 1752 |
+
"cell_type": "code",
|
| 1753 |
+
"execution_count": null,
|
| 1754 |
+
"metadata": {
|
| 1755 |
+
"id": "4za4O0m5L9ih"
|
| 1756 |
+
},
|
| 1757 |
+
"outputs": [],
|
| 1758 |
+
"source": [
|
| 1759 |
+
"x = np.array([1, 2]) # Let numpy choose the datatype\n",
|
| 1760 |
+
"y = np.array([1.0, 2.0]) # Let numpy choose the datatype\n",
|
| 1761 |
+
"z = np.array([1, 2], dtype=np.int64) # Force a particular datatype\n",
|
| 1762 |
+
"\n",
|
| 1763 |
+
"print(x.dtype, y.dtype, z.dtype)"
|
| 1764 |
+
]
|
| 1765 |
+
},
|
| 1766 |
+
{
|
| 1767 |
+
"cell_type": "markdown",
|
| 1768 |
+
"metadata": {
|
| 1769 |
+
"id": "RLVIsZQpL9ik"
|
| 1770 |
+
},
|
| 1771 |
+
"source": [
|
| 1772 |
+
"You can read all about numpy datatypes in the [documentation](http://docs.scipy.org/doc/numpy/reference/arrays.dtypes.html)."
|
| 1773 |
+
]
|
| 1774 |
+
},
|
| 1775 |
+
{
|
| 1776 |
+
"cell_type": "markdown",
|
| 1777 |
+
"metadata": {
|
| 1778 |
+
"id": "TuB-fdhIL9ik"
|
| 1779 |
+
},
|
| 1780 |
+
"source": [
|
| 1781 |
+
"###Array math"
|
| 1782 |
+
]
|
| 1783 |
+
},
|
| 1784 |
+
{
|
| 1785 |
+
"cell_type": "markdown",
|
| 1786 |
+
"metadata": {
|
| 1787 |
+
"id": "18e8V8elL9ik"
|
| 1788 |
+
},
|
| 1789 |
+
"source": [
|
| 1790 |
+
"Basic mathematical functions operate elementwise on arrays, and are available both as operator overloads and as functions in the numpy module:"
|
| 1791 |
+
]
|
| 1792 |
+
},
|
| 1793 |
+
{
|
| 1794 |
+
"cell_type": "code",
|
| 1795 |
+
"execution_count": null,
|
| 1796 |
+
"metadata": {
|
| 1797 |
+
"id": "gHKvBrSKL9il"
|
| 1798 |
+
},
|
| 1799 |
+
"outputs": [],
|
| 1800 |
+
"source": [
|
| 1801 |
+
"x = np.array([[1,2],[3,4]], dtype=np.float64)\n",
|
| 1802 |
+
"y = np.array([[5,6],[7,8]], dtype=np.float64)\n",
|
| 1803 |
+
"\n",
|
| 1804 |
+
"# Elementwise sum; both produce the array\n",
|
| 1805 |
+
"print(x + y)\n",
|
| 1806 |
+
"print(np.add(x, y))"
|
| 1807 |
+
]
|
| 1808 |
+
},
|
| 1809 |
+
{
|
| 1810 |
+
"cell_type": "code",
|
| 1811 |
+
"execution_count": null,
|
| 1812 |
+
"metadata": {
|
| 1813 |
+
"id": "1fZtIAMxL9in"
|
| 1814 |
+
},
|
| 1815 |
+
"outputs": [],
|
| 1816 |
+
"source": [
|
| 1817 |
+
"# Elementwise difference; both produce the array\n",
|
| 1818 |
+
"print(x - y)\n",
|
| 1819 |
+
"print(np.subtract(x, y))"
|
| 1820 |
+
]
|
| 1821 |
+
},
|
| 1822 |
+
{
|
| 1823 |
+
"cell_type": "code",
|
| 1824 |
+
"execution_count": null,
|
| 1825 |
+
"metadata": {
|
| 1826 |
+
"id": "nil4AScML9io"
|
| 1827 |
+
},
|
| 1828 |
+
"outputs": [],
|
| 1829 |
+
"source": [
|
| 1830 |
+
"# Elementwise product; both produce the array\n",
|
| 1831 |
+
"print(x * y)\n",
|
| 1832 |
+
"print(np.multiply(x, y))"
|
| 1833 |
+
]
|
| 1834 |
+
},
|
| 1835 |
+
{
|
| 1836 |
+
"cell_type": "code",
|
| 1837 |
+
"execution_count": null,
|
| 1838 |
+
"metadata": {
|
| 1839 |
+
"id": "0JoA4lH6L9ip"
|
| 1840 |
+
},
|
| 1841 |
+
"outputs": [],
|
| 1842 |
+
"source": [
|
| 1843 |
+
"# Elementwise division; both produce the array\n",
|
| 1844 |
+
"# [[ 0.2 0.33333333]\n",
|
| 1845 |
+
"# [ 0.42857143 0.5 ]]\n",
|
| 1846 |
+
"print(x / y)\n",
|
| 1847 |
+
"print(np.divide(x, y))"
|
| 1848 |
+
]
|
| 1849 |
+
},
|
| 1850 |
+
{
|
| 1851 |
+
"cell_type": "code",
|
| 1852 |
+
"execution_count": null,
|
| 1853 |
+
"metadata": {
|
| 1854 |
+
"id": "g0iZuA6bL9ir"
|
| 1855 |
+
},
|
| 1856 |
+
"outputs": [],
|
| 1857 |
+
"source": [
|
| 1858 |
+
"# Elementwise square root; produces the array\n",
|
| 1859 |
+
"# [[ 1. 1.41421356]\n",
|
| 1860 |
+
"# [ 1.73205081 2. ]]\n",
|
| 1861 |
+
"print(np.sqrt(x))"
|
| 1862 |
+
]
|
| 1863 |
+
},
|
| 1864 |
+
{
|
| 1865 |
+
"cell_type": "markdown",
|
| 1866 |
+
"metadata": {
|
| 1867 |
+
"id": "a5d_uujuL9it"
|
| 1868 |
+
},
|
| 1869 |
+
"source": [
|
| 1870 |
+
"Note that unlike MATLAB, `*` is elementwise multiplication, not matrix multiplication. We instead use the dot function to compute inner products of vectors, to multiply a vector by a matrix, and to multiply matrices. dot is available both as a function in the numpy module and as an instance method of array objects:"
|
| 1871 |
+
]
|
| 1872 |
+
},
|
| 1873 |
+
{
|
| 1874 |
+
"cell_type": "code",
|
| 1875 |
+
"execution_count": null,
|
| 1876 |
+
"metadata": {
|
| 1877 |
+
"id": "I3FnmoSeL9iu"
|
| 1878 |
+
},
|
| 1879 |
+
"outputs": [],
|
| 1880 |
+
"source": [
|
| 1881 |
+
"x = np.array([[1,2],[3,4]])\n",
|
| 1882 |
+
"y = np.array([[5,6],[7,8]])\n",
|
| 1883 |
+
"\n",
|
| 1884 |
+
"v = np.array([9, 10])\n",
|
| 1885 |
+
"w = np.array([11, 12])\n",
|
| 1886 |
+
"\n",
|
| 1887 |
+
"# Inner product of vectors; both produce 219\n",
|
| 1888 |
+
"print(v.dot(w))\n",
|
| 1889 |
+
"print(np.dot(v, w))"
|
| 1890 |
+
]
|
| 1891 |
+
},
|
| 1892 |
+
{
|
| 1893 |
+
"cell_type": "markdown",
|
| 1894 |
+
"metadata": {
|
| 1895 |
+
"id": "vmxPbrHASVeA"
|
| 1896 |
+
},
|
| 1897 |
+
"source": [
|
| 1898 |
+
"You can also use the `@` operator which is equivalent to numpy's `dot` operator."
|
| 1899 |
+
]
|
| 1900 |
+
},
|
| 1901 |
+
{
|
| 1902 |
+
"cell_type": "code",
|
| 1903 |
+
"execution_count": null,
|
| 1904 |
+
"metadata": {
|
| 1905 |
+
"id": "vyrWA-mXSdtt"
|
| 1906 |
+
},
|
| 1907 |
+
"outputs": [],
|
| 1908 |
+
"source": [
|
| 1909 |
+
"print(v @ w)"
|
| 1910 |
+
]
|
| 1911 |
+
},
|
| 1912 |
+
{
|
| 1913 |
+
"cell_type": "code",
|
| 1914 |
+
"execution_count": null,
|
| 1915 |
+
"metadata": {
|
| 1916 |
+
"id": "zvUODeTxL9iw"
|
| 1917 |
+
},
|
| 1918 |
+
"outputs": [],
|
| 1919 |
+
"source": [
|
| 1920 |
+
"# Matrix / vector product; both produce the rank 1 array [29 67]\n",
|
| 1921 |
+
"print(x.dot(v))\n",
|
| 1922 |
+
"print(np.dot(x, v))\n",
|
| 1923 |
+
"print(x @ v)"
|
| 1924 |
+
]
|
| 1925 |
+
},
|
| 1926 |
+
{
|
| 1927 |
+
"cell_type": "code",
|
| 1928 |
+
"execution_count": null,
|
| 1929 |
+
"metadata": {
|
| 1930 |
+
"id": "3V_3NzNEL9iy"
|
| 1931 |
+
},
|
| 1932 |
+
"outputs": [],
|
| 1933 |
+
"source": [
|
| 1934 |
+
"# Matrix / matrix product; both produce the rank 2 array\n",
|
| 1935 |
+
"# [[19 22]\n",
|
| 1936 |
+
"# [43 50]]\n",
|
| 1937 |
+
"print(x.dot(y))\n",
|
| 1938 |
+
"print(np.dot(x, y))\n",
|
| 1939 |
+
"print(x @ y)"
|
| 1940 |
+
]
|
| 1941 |
+
},
|
| 1942 |
+
{
|
| 1943 |
+
"cell_type": "markdown",
|
| 1944 |
+
"metadata": {
|
| 1945 |
+
"id": "FbE-1If_L9i0"
|
| 1946 |
+
},
|
| 1947 |
+
"source": [
|
| 1948 |
+
"Numpy provides many useful functions for performing computations on arrays; one of the most useful is `sum`:"
|
| 1949 |
+
]
|
| 1950 |
+
},
|
| 1951 |
+
{
|
| 1952 |
+
"cell_type": "code",
|
| 1953 |
+
"execution_count": null,
|
| 1954 |
+
"metadata": {
|
| 1955 |
+
"id": "DZUdZvPrL9i0"
|
| 1956 |
+
},
|
| 1957 |
+
"outputs": [],
|
| 1958 |
+
"source": [
|
| 1959 |
+
"x = np.array([[1,2],[3,4]])\n",
|
| 1960 |
+
"\n",
|
| 1961 |
+
"print(np.sum(x)) # Compute sum of all elements; prints \"10\"\n",
|
| 1962 |
+
"print(np.sum(x, axis=0)) # Compute sum of each column; prints \"[4 6]\"\n",
|
| 1963 |
+
"print(np.sum(x, axis=1)) # Compute sum of each row; prints \"[3 7]\""
|
| 1964 |
+
]
|
| 1965 |
+
},
|
| 1966 |
+
{
|
| 1967 |
+
"cell_type": "markdown",
|
| 1968 |
+
"metadata": {
|
| 1969 |
+
"id": "ahdVW4iUL9i3"
|
| 1970 |
+
},
|
| 1971 |
+
"source": [
|
| 1972 |
+
"You can find the full list of mathematical functions provided by numpy in the [documentation](http://docs.scipy.org/doc/numpy/reference/routines.math.html).\n",
|
| 1973 |
+
"\n",
|
| 1974 |
+
"Apart from computing mathematical functions using arrays, we frequently need to reshape or otherwise manipulate data in arrays. The simplest example of this type of operation is transposing a matrix; to transpose a matrix, simply use the T attribute of an array object:"
|
| 1975 |
+
]
|
| 1976 |
+
},
|
| 1977 |
+
{
|
| 1978 |
+
"cell_type": "code",
|
| 1979 |
+
"execution_count": null,
|
| 1980 |
+
"metadata": {
|
| 1981 |
+
"id": "63Yl1f3oL9i3"
|
| 1982 |
+
},
|
| 1983 |
+
"outputs": [],
|
| 1984 |
+
"source": [
|
| 1985 |
+
"print(x)\n",
|
| 1986 |
+
"print(\"transpose\\n\", x.T)"
|
| 1987 |
+
]
|
| 1988 |
+
},
|
| 1989 |
+
{
|
| 1990 |
+
"cell_type": "code",
|
| 1991 |
+
"execution_count": null,
|
| 1992 |
+
"metadata": {
|
| 1993 |
+
"id": "mkk03eNIL9i4"
|
| 1994 |
+
},
|
| 1995 |
+
"outputs": [],
|
| 1996 |
+
"source": [
|
| 1997 |
+
"v = np.array([[1,2,3]])\n",
|
| 1998 |
+
"print(v )\n",
|
| 1999 |
+
"print(\"transpose\\n\", v.T)"
|
| 2000 |
+
]
|
| 2001 |
+
},
|
| 2002 |
+
{
|
| 2003 |
+
"cell_type": "markdown",
|
| 2004 |
+
"metadata": {
|
| 2005 |
+
"id": "REfLrUTcL9i7"
|
| 2006 |
+
},
|
| 2007 |
+
"source": [
|
| 2008 |
+
"###Broadcasting"
|
| 2009 |
+
]
|
| 2010 |
+
},
|
| 2011 |
+
{
|
| 2012 |
+
"cell_type": "markdown",
|
| 2013 |
+
"metadata": {
|
| 2014 |
+
"id": "EygGAMWqL9i7"
|
| 2015 |
+
},
|
| 2016 |
+
"source": [
|
| 2017 |
+
"Broadcasting is a powerful mechanism that allows numpy to work with arrays of different shapes when performing arithmetic operations. Frequently we have a smaller array and a larger array, and we want to use the smaller array multiple times to perform some operation on the larger array.\n",
|
| 2018 |
+
"\n",
|
| 2019 |
+
"For example, suppose that we want to add a constant vector to each row of a matrix. We could do it like this:"
|
| 2020 |
+
]
|
| 2021 |
+
},
|
| 2022 |
+
{
|
| 2023 |
+
"cell_type": "code",
|
| 2024 |
+
"execution_count": null,
|
| 2025 |
+
"metadata": {
|
| 2026 |
+
"id": "WEEvkV1ZL9i7"
|
| 2027 |
+
},
|
| 2028 |
+
"outputs": [],
|
| 2029 |
+
"source": [
|
| 2030 |
+
"# We will add the vector v to each row of the matrix x,\n",
|
| 2031 |
+
"# storing the result in the matrix y\n",
|
| 2032 |
+
"x = np.array([[1,2,3], [4,5,6], [7,8,9], [10, 11, 12]])\n",
|
| 2033 |
+
"v = np.array([1, 0, 1])\n",
|
| 2034 |
+
"y = np.empty_like(x) # Create an empty matrix with the same shape as x\n",
|
| 2035 |
+
"\n",
|
| 2036 |
+
"# Add the vector v to each row of the matrix x with an explicit loop\n",
|
| 2037 |
+
"for i in range(4):\n",
|
| 2038 |
+
" y[i, :] = x[i, :] + v\n",
|
| 2039 |
+
"\n",
|
| 2040 |
+
"print(y)"
|
| 2041 |
+
]
|
| 2042 |
+
},
|
| 2043 |
+
{
|
| 2044 |
+
"cell_type": "markdown",
|
| 2045 |
+
"metadata": {
|
| 2046 |
+
"id": "2OlXXupEL9i-"
|
| 2047 |
+
},
|
| 2048 |
+
"source": [
|
| 2049 |
+
"This works; however when the matrix `x` is very large, computing an explicit loop in Python could be slow. Note that adding the vector v to each row of the matrix `x` is equivalent to forming a matrix `vv` by stacking multiple copies of `v` vertically, then performing elementwise summation of `x` and `vv`. We could implement this approach like this:"
|
| 2050 |
+
]
|
| 2051 |
+
},
|
| 2052 |
+
{
|
| 2053 |
+
"cell_type": "code",
|
| 2054 |
+
"execution_count": null,
|
| 2055 |
+
"metadata": {
|
| 2056 |
+
"id": "vS7UwAQQL9i-"
|
| 2057 |
+
},
|
| 2058 |
+
"outputs": [],
|
| 2059 |
+
"source": [
|
| 2060 |
+
"vv = np.tile(v, (4, 1)) # Stack 4 copies of v on top of each other\n",
|
| 2061 |
+
"print(vv) # Prints \"[[1 0 1]\n",
|
| 2062 |
+
" # [1 0 1]\n",
|
| 2063 |
+
" # [1 0 1]\n",
|
| 2064 |
+
" # [1 0 1]]\""
|
| 2065 |
+
]
|
| 2066 |
+
},
|
| 2067 |
+
{
|
| 2068 |
+
"cell_type": "code",
|
| 2069 |
+
"execution_count": null,
|
| 2070 |
+
"metadata": {
|
| 2071 |
+
"id": "N0hJphSIL9jA"
|
| 2072 |
+
},
|
| 2073 |
+
"outputs": [],
|
| 2074 |
+
"source": [
|
| 2075 |
+
"y = x + vv # Add x and vv elementwise\n",
|
| 2076 |
+
"print(y)"
|
| 2077 |
+
]
|
| 2078 |
+
},
|
| 2079 |
+
{
|
| 2080 |
+
"cell_type": "markdown",
|
| 2081 |
+
"metadata": {
|
| 2082 |
+
"id": "zHos6RJnL9jB"
|
| 2083 |
+
},
|
| 2084 |
+
"source": [
|
| 2085 |
+
"Numpy broadcasting allows us to perform this computation without actually creating multiple copies of v. Consider this version, using broadcasting:"
|
| 2086 |
+
]
|
| 2087 |
+
},
|
| 2088 |
+
{
|
| 2089 |
+
"cell_type": "code",
|
| 2090 |
+
"execution_count": null,
|
| 2091 |
+
"metadata": {
|
| 2092 |
+
"id": "vnYFb-gYL9jC"
|
| 2093 |
+
},
|
| 2094 |
+
"outputs": [],
|
| 2095 |
+
"source": [
|
| 2096 |
+
"import numpy as np\n",
|
| 2097 |
+
"\n",
|
| 2098 |
+
"# We will add the vector v to each row of the matrix x,\n",
|
| 2099 |
+
"# storing the result in the matrix y\n",
|
| 2100 |
+
"x = np.array([[1,2,3], [4,5,6], [7,8,9], [10, 11, 12]])\n",
|
| 2101 |
+
"v = np.array([1, 0, 1])\n",
|
| 2102 |
+
"y = x + v # Add v to each row of x using broadcasting\n",
|
| 2103 |
+
"print(y)"
|
| 2104 |
+
]
|
| 2105 |
+
},
|
| 2106 |
+
{
|
| 2107 |
+
"cell_type": "markdown",
|
| 2108 |
+
"metadata": {
|
| 2109 |
+
"id": "08YyIURKL9jH"
|
| 2110 |
+
},
|
| 2111 |
+
"source": [
|
| 2112 |
+
"The line `y = x + v` works even though `x` has shape `(4, 3)` and `v` has shape `(3,)` due to broadcasting; this line works as if v actually had shape `(4, 3)`, where each row was a copy of `v`, and the sum was performed elementwise.\n",
|
| 2113 |
+
"\n",
|
| 2114 |
+
"Broadcasting two arrays together follows these rules:\n",
|
| 2115 |
+
"\n",
|
| 2116 |
+
"1. If the arrays do not have the same rank, prepend the shape of the lower rank array with 1s until both shapes have the same length.\n",
|
| 2117 |
+
"2. The two arrays are said to be compatible in a dimension if they have the same size in the dimension, or if one of the arrays has size 1 in that dimension.\n",
|
| 2118 |
+
"3. The arrays can be broadcast together if they are compatible in all dimensions.\n",
|
| 2119 |
+
"4. After broadcasting, each array behaves as if it had shape equal to the elementwise maximum of shapes of the two input arrays.\n",
|
| 2120 |
+
"5. In any dimension where one array had size 1 and the other array had size greater than 1, the first array behaves as if it were copied along that dimension\n",
|
| 2121 |
+
"\n",
|
| 2122 |
+
"If this explanation does not make sense, try reading the explanation from the [documentation](http://docs.scipy.org/doc/numpy/user/basics.broadcasting.html) or this [explanation](http://wiki.scipy.org/EricsBroadcastingDoc).\n",
|
| 2123 |
+
"\n",
|
| 2124 |
+
"Functions that support broadcasting are known as universal functions. You can find the list of all universal functions in the [documentation](http://docs.scipy.org/doc/numpy/reference/ufuncs.html#available-ufuncs).\n",
|
| 2125 |
+
"\n",
|
| 2126 |
+
"Here are some applications of broadcasting:"
|
| 2127 |
+
]
|
| 2128 |
+
},
|
| 2129 |
+
{
|
| 2130 |
+
"cell_type": "code",
|
| 2131 |
+
"execution_count": null,
|
| 2132 |
+
"metadata": {
|
| 2133 |
+
"id": "EmQnwoM9L9jH"
|
| 2134 |
+
},
|
| 2135 |
+
"outputs": [],
|
| 2136 |
+
"source": [
|
| 2137 |
+
"# Compute outer product of vectors\n",
|
| 2138 |
+
"v = np.array([1,2,3]) # v has shape (3,)\n",
|
| 2139 |
+
"w = np.array([4,5]) # w has shape (2,)\n",
|
| 2140 |
+
"# To compute an outer product, we first reshape v to be a column\n",
|
| 2141 |
+
"# vector of shape (3, 1); we can then broadcast it against w to yield\n",
|
| 2142 |
+
"# an output of shape (3, 2), which is the outer product of v and w:\n",
|
| 2143 |
+
"\n",
|
| 2144 |
+
"print(np.reshape(v, (3, 1)) * w)"
|
| 2145 |
+
]
|
| 2146 |
+
},
|
| 2147 |
+
{
|
| 2148 |
+
"cell_type": "code",
|
| 2149 |
+
"execution_count": null,
|
| 2150 |
+
"metadata": {
|
| 2151 |
+
"id": "PgotmpcnL9jK"
|
| 2152 |
+
},
|
| 2153 |
+
"outputs": [],
|
| 2154 |
+
"source": [
|
| 2155 |
+
"# Add a vector to each row of a matrix\n",
|
| 2156 |
+
"x = np.array([[1,2,3], [4,5,6]])\n",
|
| 2157 |
+
"# x has shape (2, 3) and v has shape (3,) so they broadcast to (2, 3),\n",
|
| 2158 |
+
"# giving the following matrix:\n",
|
| 2159 |
+
"\n",
|
| 2160 |
+
"print(x + v)"
|
| 2161 |
+
]
|
| 2162 |
+
},
|
| 2163 |
+
{
|
| 2164 |
+
"cell_type": "code",
|
| 2165 |
+
"execution_count": null,
|
| 2166 |
+
"metadata": {
|
| 2167 |
+
"id": "T5hKS1QaL9jK"
|
| 2168 |
+
},
|
| 2169 |
+
"outputs": [],
|
| 2170 |
+
"source": [
|
| 2171 |
+
"# Add a vector to each column of a matrix\n",
|
| 2172 |
+
"# x has shape (2, 3) and w has shape (2,).\n",
|
| 2173 |
+
"# If we transpose x then it has shape (3, 2) and can be broadcast\n",
|
| 2174 |
+
"# against w to yield a result of shape (3, 2); transposing this result\n",
|
| 2175 |
+
"# yields the final result of shape (2, 3) which is the matrix x with\n",
|
| 2176 |
+
"# the vector w added to each column. Gives the following matrix:\n",
|
| 2177 |
+
"\n",
|
| 2178 |
+
"print((x.T + w).T)"
|
| 2179 |
+
]
|
| 2180 |
+
},
|
| 2181 |
+
{
|
| 2182 |
+
"cell_type": "code",
|
| 2183 |
+
"execution_count": null,
|
| 2184 |
+
"metadata": {
|
| 2185 |
+
"id": "JDUrZUl6L9jN"
|
| 2186 |
+
},
|
| 2187 |
+
"outputs": [],
|
| 2188 |
+
"source": [
|
| 2189 |
+
"# Another solution is to reshape w to be a row vector of shape (2, 1);\n",
|
| 2190 |
+
"# we can then broadcast it directly against x to produce the same\n",
|
| 2191 |
+
"# output.\n",
|
| 2192 |
+
"print(x + np.reshape(w, (2, 1)))"
|
| 2193 |
+
]
|
| 2194 |
+
},
|
| 2195 |
+
{
|
| 2196 |
+
"cell_type": "code",
|
| 2197 |
+
"execution_count": null,
|
| 2198 |
+
"metadata": {
|
| 2199 |
+
"id": "VzrEo4KGL9jP"
|
| 2200 |
+
},
|
| 2201 |
+
"outputs": [],
|
| 2202 |
+
"source": [
|
| 2203 |
+
"# Multiply a matrix by a constant:\n",
|
| 2204 |
+
"# x has shape (2, 3). Numpy treats scalars as arrays of shape ();\n",
|
| 2205 |
+
"# these can be broadcast together to shape (2, 3), producing the\n",
|
| 2206 |
+
"# following array:\n",
|
| 2207 |
+
"print(x * 2)"
|
| 2208 |
+
]
|
| 2209 |
+
},
|
| 2210 |
+
{
|
| 2211 |
+
"cell_type": "markdown",
|
| 2212 |
+
"metadata": {
|
| 2213 |
+
"id": "89e2FXxFL9jQ"
|
| 2214 |
+
},
|
| 2215 |
+
"source": [
|
| 2216 |
+
"Broadcasting typically makes your code more concise and faster, so you should strive to use it where possible."
|
| 2217 |
+
]
|
| 2218 |
+
},
|
| 2219 |
+
{
|
| 2220 |
+
"cell_type": "markdown",
|
| 2221 |
+
"metadata": {
|
| 2222 |
+
"id": "yi90439hpLR0"
|
| 2223 |
+
},
|
| 2224 |
+
"source": [
|
| 2225 |
+
"### Numpy documentation"
|
| 2226 |
+
]
|
| 2227 |
+
},
|
| 2228 |
+
{
|
| 2229 |
+
"cell_type": "markdown",
|
| 2230 |
+
"metadata": {
|
| 2231 |
+
"id": "iF3ZtwVNL9jQ"
|
| 2232 |
+
},
|
| 2233 |
+
"source": [
|
| 2234 |
+
"This brief overview has touched on many of the important things that you need to know about numpy, but is far from complete. Check out the [numpy reference](http://docs.scipy.org/doc/numpy/reference/) to find out much more about numpy."
|
| 2235 |
+
]
|
| 2236 |
+
},
|
| 2237 |
+
{
|
| 2238 |
+
"cell_type": "markdown",
|
| 2239 |
+
"metadata": {
|
| 2240 |
+
"id": "tEINf4bEL9jR"
|
| 2241 |
+
},
|
| 2242 |
+
"source": [
|
| 2243 |
+
"##Matplotlib"
|
| 2244 |
+
]
|
| 2245 |
+
},
|
| 2246 |
+
{
|
| 2247 |
+
"cell_type": "markdown",
|
| 2248 |
+
"metadata": {
|
| 2249 |
+
"id": "0hgVWLaXL9jR"
|
| 2250 |
+
},
|
| 2251 |
+
"source": [
|
| 2252 |
+
"Matplotlib is a plotting library. In this section give a brief introduction to the `matplotlib.pyplot` module, which provides a plotting system similar to that of MATLAB."
|
| 2253 |
+
]
|
| 2254 |
+
},
|
| 2255 |
+
{
|
| 2256 |
+
"cell_type": "code",
|
| 2257 |
+
"execution_count": null,
|
| 2258 |
+
"metadata": {
|
| 2259 |
+
"id": "cmh_7c6KL9jR"
|
| 2260 |
+
},
|
| 2261 |
+
"outputs": [],
|
| 2262 |
+
"source": [
|
| 2263 |
+
"import matplotlib.pyplot as plt"
|
| 2264 |
+
]
|
| 2265 |
+
},
|
| 2266 |
+
{
|
| 2267 |
+
"cell_type": "markdown",
|
| 2268 |
+
"metadata": {
|
| 2269 |
+
"id": "jOsaA5hGL9jS"
|
| 2270 |
+
},
|
| 2271 |
+
"source": [
|
| 2272 |
+
"By running this special iPython command, we will be displaying plots inline:"
|
| 2273 |
+
]
|
| 2274 |
+
},
|
| 2275 |
+
{
|
| 2276 |
+
"cell_type": "code",
|
| 2277 |
+
"execution_count": null,
|
| 2278 |
+
"metadata": {
|
| 2279 |
+
"id": "ijpsmwGnL9jT"
|
| 2280 |
+
},
|
| 2281 |
+
"outputs": [],
|
| 2282 |
+
"source": [
|
| 2283 |
+
"%matplotlib inline"
|
| 2284 |
+
]
|
| 2285 |
+
},
|
| 2286 |
+
{
|
| 2287 |
+
"cell_type": "markdown",
|
| 2288 |
+
"metadata": {
|
| 2289 |
+
"id": "U5Z_oMoLL9jV"
|
| 2290 |
+
},
|
| 2291 |
+
"source": [
|
| 2292 |
+
"###Plotting"
|
| 2293 |
+
]
|
| 2294 |
+
},
|
| 2295 |
+
{
|
| 2296 |
+
"cell_type": "markdown",
|
| 2297 |
+
"metadata": {
|
| 2298 |
+
"id": "6QyFJ7dhL9jV"
|
| 2299 |
+
},
|
| 2300 |
+
"source": [
|
| 2301 |
+
"The most important function in `matplotlib` is plot, which allows you to plot 2D data. Here is a simple example:"
|
| 2302 |
+
]
|
| 2303 |
+
},
|
| 2304 |
+
{
|
| 2305 |
+
"cell_type": "code",
|
| 2306 |
+
"execution_count": null,
|
| 2307 |
+
"metadata": {
|
| 2308 |
+
"id": "pua52BGeL9jW"
|
| 2309 |
+
},
|
| 2310 |
+
"outputs": [],
|
| 2311 |
+
"source": [
|
| 2312 |
+
"# Compute the x and y coordinates for points on a sine curve\n",
|
| 2313 |
+
"x = np.arange(0, 3 * np.pi, 0.1)\n",
|
| 2314 |
+
"y = np.sin(x)\n",
|
| 2315 |
+
"\n",
|
| 2316 |
+
"# Plot the points using matplotlib\n",
|
| 2317 |
+
"plt.plot(x, y)"
|
| 2318 |
+
]
|
| 2319 |
+
},
|
| 2320 |
+
{
|
| 2321 |
+
"cell_type": "markdown",
|
| 2322 |
+
"metadata": {
|
| 2323 |
+
"id": "9W2VAcLiL9jX"
|
| 2324 |
+
},
|
| 2325 |
+
"source": [
|
| 2326 |
+
"With just a little bit of extra work we can easily plot multiple lines at once, and add a title, legend, and axis labels:"
|
| 2327 |
+
]
|
| 2328 |
+
},
|
| 2329 |
+
{
|
| 2330 |
+
"cell_type": "code",
|
| 2331 |
+
"execution_count": null,
|
| 2332 |
+
"metadata": {
|
| 2333 |
+
"id": "TfCQHJ5AL9jY"
|
| 2334 |
+
},
|
| 2335 |
+
"outputs": [],
|
| 2336 |
+
"source": [
|
| 2337 |
+
"y_sin = np.sin(x)\n",
|
| 2338 |
+
"y_cos = np.cos(x)\n",
|
| 2339 |
+
"\n",
|
| 2340 |
+
"# Plot the points using matplotlib\n",
|
| 2341 |
+
"plt.plot(x, y_sin)\n",
|
| 2342 |
+
"plt.plot(x, y_cos)\n",
|
| 2343 |
+
"plt.xlabel('x axis label')\n",
|
| 2344 |
+
"plt.ylabel('y axis label')\n",
|
| 2345 |
+
"plt.title('Sine and Cosine')\n",
|
| 2346 |
+
"plt.legend(['Sine', 'Cosine'])"
|
| 2347 |
+
]
|
| 2348 |
+
},
|
| 2349 |
+
{
|
| 2350 |
+
"cell_type": "markdown",
|
| 2351 |
+
"metadata": {
|
| 2352 |
+
"id": "R5IeAY03L9ja"
|
| 2353 |
+
},
|
| 2354 |
+
"source": [
|
| 2355 |
+
"###Subplots"
|
| 2356 |
+
]
|
| 2357 |
+
},
|
| 2358 |
+
{
|
| 2359 |
+
"cell_type": "markdown",
|
| 2360 |
+
"metadata": {
|
| 2361 |
+
"id": "CfUzwJg0L9ja"
|
| 2362 |
+
},
|
| 2363 |
+
"source": [
|
| 2364 |
+
"You can plot different things in the same figure using the subplot function. Here is an example:"
|
| 2365 |
+
]
|
| 2366 |
+
},
|
| 2367 |
+
{
|
| 2368 |
+
"cell_type": "code",
|
| 2369 |
+
"execution_count": null,
|
| 2370 |
+
"metadata": {
|
| 2371 |
+
"id": "dM23yGH9L9ja"
|
| 2372 |
+
},
|
| 2373 |
+
"outputs": [],
|
| 2374 |
+
"source": [
|
| 2375 |
+
"# Compute the x and y coordinates for points on sine and cosine curves\n",
|
| 2376 |
+
"x = np.arange(0, 3 * np.pi, 0.1)\n",
|
| 2377 |
+
"y_sin = np.sin(x)\n",
|
| 2378 |
+
"y_cos = np.cos(x)\n",
|
| 2379 |
+
"\n",
|
| 2380 |
+
"# Set up a subplot grid that has height 2 and width 1,\n",
|
| 2381 |
+
"# and set the first such subplot as active.\n",
|
| 2382 |
+
"plt.subplot(2, 1, 1)\n",
|
| 2383 |
+
"\n",
|
| 2384 |
+
"# Make the first plot\n",
|
| 2385 |
+
"plt.plot(x, y_sin)\n",
|
| 2386 |
+
"plt.title('Sine')\n",
|
| 2387 |
+
"\n",
|
| 2388 |
+
"# Set the second subplot as active, and make the second plot.\n",
|
| 2389 |
+
"plt.subplot(2, 1, 2)\n",
|
| 2390 |
+
"plt.plot(x, y_cos)\n",
|
| 2391 |
+
"plt.title('Cosine')\n",
|
| 2392 |
+
"\n",
|
| 2393 |
+
"# Show the figure.\n",
|
| 2394 |
+
"plt.show()"
|
| 2395 |
+
]
|
| 2396 |
+
},
|
| 2397 |
+
{
|
| 2398 |
+
"cell_type": "markdown",
|
| 2399 |
+
"metadata": {
|
| 2400 |
+
"id": "gLtsST5SL9jc"
|
| 2401 |
+
},
|
| 2402 |
+
"source": [
|
| 2403 |
+
"You can read much more about the `subplot` function in the [documentation](http://matplotlib.org/api/pyplot_api.html#matplotlib.pyplot.subplot)."
|
| 2404 |
+
]
|
| 2405 |
+
},
|
| 2406 |
+
{
|
| 2407 |
+
"cell_type": "markdown",
|
| 2408 |
+
"metadata": {
|
| 2409 |
+
"id": "7Zqndtogsq8J"
|
| 2410 |
+
},
|
| 2411 |
+
"source": [
|
| 2412 |
+
"### Download images"
|
| 2413 |
+
]
|
| 2414 |
+
},
|
| 2415 |
+
{
|
| 2416 |
+
"cell_type": "markdown",
|
| 2417 |
+
"metadata": {
|
| 2418 |
+
"id": "4FFozuUm7OE5"
|
| 2419 |
+
},
|
| 2420 |
+
"source": [
|
| 2421 |
+
"Lets download some images."
|
| 2422 |
+
]
|
| 2423 |
+
},
|
| 2424 |
+
{
|
| 2425 |
+
"cell_type": "code",
|
| 2426 |
+
"execution_count": null,
|
| 2427 |
+
"metadata": {
|
| 2428 |
+
"id": "cOEqkMl47VTP"
|
| 2429 |
+
},
|
| 2430 |
+
"outputs": [],
|
| 2431 |
+
"source": [
|
| 2432 |
+
"import os\n",
|
| 2433 |
+
"if not os.path.exists('images.zip'):\n",
|
| 2434 |
+
" !wget --no-check-certificate https://empslocal.ex.ac.uk/people/staff/ad735/ECMM426/images.zip\n",
|
| 2435 |
+
" !unzip -q images.zip"
|
| 2436 |
+
]
|
| 2437 |
+
},
|
| 2438 |
+
{
|
| 2439 |
+
"cell_type": "markdown",
|
| 2440 |
+
"metadata": {
|
| 2441 |
+
"id": "cjKpKTNYsxcF"
|
| 2442 |
+
},
|
| 2443 |
+
"source": [
|
| 2444 |
+
"You can use the `imread` and the`imshow` function to respectively read and show images. Here is an example:"
|
| 2445 |
+
]
|
| 2446 |
+
},
|
| 2447 |
+
{
|
| 2448 |
+
"cell_type": "code",
|
| 2449 |
+
"execution_count": null,
|
| 2450 |
+
"metadata": {
|
| 2451 |
+
"id": "LEep7wnTs5Si"
|
| 2452 |
+
},
|
| 2453 |
+
"outputs": [],
|
| 2454 |
+
"source": [
|
| 2455 |
+
"import numpy as np\n",
|
| 2456 |
+
"import matplotlib.pyplot as plt\n",
|
| 2457 |
+
"\n",
|
| 2458 |
+
"img = plt.imread('images/lena.png')\n",
|
| 2459 |
+
"img_tinted = img * [1, 0.85, 0.8]\n",
|
| 2460 |
+
"\n",
|
| 2461 |
+
"# Show the original image\n",
|
| 2462 |
+
"plt.subplot(1, 2, 1)\n",
|
| 2463 |
+
"plt.imshow(img)\n",
|
| 2464 |
+
"\n",
|
| 2465 |
+
"# Show the tinted image\n",
|
| 2466 |
+
"plt.subplot(1, 2, 2)\n",
|
| 2467 |
+
"plt.imshow(img_tinted)\n",
|
| 2468 |
+
"plt.show()"
|
| 2469 |
+
]
|
| 2470 |
+
},
|
| 2471 |
+
{
|
| 2472 |
+
"cell_type": "markdown",
|
| 2473 |
+
"metadata": {
|
| 2474 |
+
"id": "wOan27So8lpI"
|
| 2475 |
+
},
|
| 2476 |
+
"source": [
|
| 2477 |
+
"## Scikit-learn"
|
| 2478 |
+
]
|
| 2479 |
+
},
|
| 2480 |
+
{
|
| 2481 |
+
"cell_type": "markdown",
|
| 2482 |
+
"metadata": {
|
| 2483 |
+
"id": "uSki-HDorHgE"
|
| 2484 |
+
},
|
| 2485 |
+
"source": [
|
| 2486 |
+
"[Scikit-learn](https://scikit-learn.org/stable/) is an open source machine learning library that supports supervised and unsupervised learning. It also provides various tools for model fitting, data preprocessing, model selection, model evaluation, and many other utilities."
|
| 2487 |
+
]
|
| 2488 |
+
},
|
| 2489 |
+
{
|
| 2490 |
+
"cell_type": "markdown",
|
| 2491 |
+
"metadata": {
|
| 2492 |
+
"id": "BOwTaL8x2QJ1"
|
| 2493 |
+
},
|
| 2494 |
+
"source": [
|
| 2495 |
+
"### Moon dataset\n",
|
| 2496 |
+
"Below we will consider a toy dataset, such as moon dataset and consider some classifiers from the Scikit-learn library to classify them."
|
| 2497 |
+
]
|
| 2498 |
+
},
|
| 2499 |
+
{
|
| 2500 |
+
"cell_type": "code",
|
| 2501 |
+
"execution_count": null,
|
| 2502 |
+
"metadata": {
|
| 2503 |
+
"id": "-wudgYDp1xO_"
|
| 2504 |
+
},
|
| 2505 |
+
"outputs": [],
|
| 2506 |
+
"source": [
|
| 2507 |
+
"# Create the moon dataset and plot\n",
|
| 2508 |
+
"from sklearn.datasets import make_moons\n",
|
| 2509 |
+
"\n",
|
| 2510 |
+
"X, y = make_moons(n_samples=500, noise=0.30, random_state=42)\n",
|
| 2511 |
+
"\n",
|
| 2512 |
+
"id0 = y == 0\n",
|
| 2513 |
+
"id1 = y == 1\n",
|
| 2514 |
+
"plt.plot(X[id0, 0], X[id0, 1], 'bo', label='0')\n",
|
| 2515 |
+
"plt.plot(X[id1, 0], X[id1, 1], 'ro', label='1')\n",
|
| 2516 |
+
"plt.legend(loc=2)"
|
| 2517 |
+
]
|
| 2518 |
+
},
|
| 2519 |
+
{
|
| 2520 |
+
"cell_type": "markdown",
|
| 2521 |
+
"metadata": {
|
| 2522 |
+
"id": "XbIu90Rg2TvB"
|
| 2523 |
+
},
|
| 2524 |
+
"source": [
|
| 2525 |
+
"### Dataset split"
|
| 2526 |
+
]
|
| 2527 |
+
},
|
| 2528 |
+
{
|
| 2529 |
+
"cell_type": "code",
|
| 2530 |
+
"execution_count": null,
|
| 2531 |
+
"metadata": {
|
| 2532 |
+
"id": "QVDMPmL62EIr"
|
| 2533 |
+
},
|
| 2534 |
+
"outputs": [],
|
| 2535 |
+
"source": [
|
| 2536 |
+
"# Split into train and test sets\n",
|
| 2537 |
+
"from sklearn.model_selection import train_test_split\n",
|
| 2538 |
+
"X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)"
|
| 2539 |
+
]
|
| 2540 |
+
},
|
| 2541 |
+
{
|
| 2542 |
+
"cell_type": "markdown",
|
| 2543 |
+
"metadata": {
|
| 2544 |
+
"id": "uZNqoqgW31rp"
|
| 2545 |
+
},
|
| 2546 |
+
"source": [
|
| 2547 |
+
"### Random forest classifier\n",
|
| 2548 |
+
"Lets now train a random forest classifier from the scikit-learn library on the above training set and test it on the test set, and the compute the classification accuracy. Please check the [documentation](https://scikit-learn.org/stable/modules/generated/sklearn.ensemble.RandomForestClassifier.html#sklearn-ensemble-randomforestclassifier) of `RandomForestClassifier` for more details on its parameters."
|
| 2549 |
+
]
|
| 2550 |
+
},
|
| 2551 |
+
{
|
| 2552 |
+
"cell_type": "code",
|
| 2553 |
+
"execution_count": null,
|
| 2554 |
+
"metadata": {
|
| 2555 |
+
"id": "RC8sl3QE5BjD"
|
| 2556 |
+
},
|
| 2557 |
+
"outputs": [],
|
| 2558 |
+
"source": [
|
| 2559 |
+
"from sklearn.ensemble import RandomForestClassifier\n",
|
| 2560 |
+
"# Define the classifier\n",
|
| 2561 |
+
"rnd_clf = RandomForestClassifier(n_estimators=500, max_leaf_nodes=16, n_jobs=-1, random_state=42)\n",
|
| 2562 |
+
"# Training\n",
|
| 2563 |
+
"rnd_clf.fit(X_train, y_train)\n",
|
| 2564 |
+
"# Test\n",
|
| 2565 |
+
"y_pred_rf = rnd_clf.predict(X_test)\n",
|
| 2566 |
+
"# Classification accuracy\n",
|
| 2567 |
+
"from sklearn.metrics import accuracy_score\n",
|
| 2568 |
+
"print(accuracy_score(y_test, y_pred_rf))"
|
| 2569 |
+
]
|
| 2570 |
+
},
|
| 2571 |
+
{
|
| 2572 |
+
"cell_type": "markdown",
|
| 2573 |
+
"metadata": {
|
| 2574 |
+
"id": "T8JLXEV97i8c"
|
| 2575 |
+
},
|
| 2576 |
+
"source": [
|
| 2577 |
+
"### Non-linear Support Vector Machine\n",
|
| 2578 |
+
"Now lets do the same training and testing with a non-linear [Support Vector Machine (SVM)](https://scikit-learn.org/stable/modules/generated/sklearn.svm.SVC.html) classifier."
|
| 2579 |
+
]
|
| 2580 |
+
},
|
| 2581 |
+
{
|
| 2582 |
+
"cell_type": "code",
|
| 2583 |
+
"execution_count": null,
|
| 2584 |
+
"metadata": {
|
| 2585 |
+
"id": "av46bVVs8w08"
|
| 2586 |
+
},
|
| 2587 |
+
"outputs": [],
|
| 2588 |
+
"source": [
|
| 2589 |
+
"from sklearn.svm import SVC\n",
|
| 2590 |
+
"# Define the non-linear classifier with radial basis function (rbf) kernel\n",
|
| 2591 |
+
"nlin_svm_clf_1 = SVC(kernel=\"rbf\")\n",
|
| 2592 |
+
"# Training\n",
|
| 2593 |
+
"nlin_svm_clf_1.fit(X_train, y_train)\n",
|
| 2594 |
+
"# Test\n",
|
| 2595 |
+
"y_pred = nlin_svm_clf_1.predict(X_test)\n",
|
| 2596 |
+
"# Classification accuracy\n",
|
| 2597 |
+
"from sklearn.metrics import accuracy_score\n",
|
| 2598 |
+
"print(accuracy_score(y_pred, y_test))"
|
| 2599 |
+
]
|
| 2600 |
+
},
|
| 2601 |
+
{
|
| 2602 |
+
"cell_type": "markdown",
|
| 2603 |
+
"metadata": {
|
| 2604 |
+
"id": "74nKsfA99pkb"
|
| 2605 |
+
},
|
| 2606 |
+
"source": [
|
| 2607 |
+
"### Confusion matrix\n",
|
| 2608 |
+
"A confusion matrix is a table that is used to define the performance of a classification algorithm. A confusion matrix visualizes and summarizes the performance of a classification algorithm. More details on how to compute confusion matrix can be found in the [documentation](https://scikit-learn.org/stable/modules/generated/sklearn.metrics.confusion_matrix.html)."
|
| 2609 |
+
]
|
| 2610 |
+
},
|
| 2611 |
+
{
|
| 2612 |
+
"cell_type": "code",
|
| 2613 |
+
"execution_count": null,
|
| 2614 |
+
"metadata": {
|
| 2615 |
+
"id": "aB73WGra-aEF"
|
| 2616 |
+
},
|
| 2617 |
+
"outputs": [],
|
| 2618 |
+
"source": [
|
| 2619 |
+
"from sklearn.metrics import confusion_matrix\n",
|
| 2620 |
+
"confusion_matrix(y_test, y_pred)"
|
| 2621 |
+
]
|
| 2622 |
+
},
|
| 2623 |
+
{
|
| 2624 |
+
"cell_type": "markdown",
|
| 2625 |
+
"metadata": {
|
| 2626 |
+
"id": "Wp2DsEemAOub"
|
| 2627 |
+
},
|
| 2628 |
+
"source": [
|
| 2629 |
+
"### Regression\n",
|
| 2630 |
+
"\n",
|
| 2631 |
+
"Now, lets consider the following function and train an [MLP regressor](https://scikit-learn.org/stable/modules/generated/sklearn.neural_network.MLPRegressor.html) to learn it.\n",
|
| 2632 |
+
"\n",
|
| 2633 |
+
"\\begin{equation}\n",
|
| 2634 |
+
"y = f(x; \\mathbf{w}) = 5x^2 + 3\n",
|
| 2635 |
+
"\\end{equation}"
|
| 2636 |
+
]
|
| 2637 |
+
},
|
| 2638 |
+
{
|
| 2639 |
+
"cell_type": "code",
|
| 2640 |
+
"execution_count": null,
|
| 2641 |
+
"metadata": {
|
| 2642 |
+
"id": "C5Dzr0MLA3h8"
|
| 2643 |
+
},
|
| 2644 |
+
"outputs": [],
|
| 2645 |
+
"source": [
|
| 2646 |
+
"# Create the data that follow uniform distribution\n",
|
| 2647 |
+
"import numpy as np\n",
|
| 2648 |
+
"import matplotlib.pyplot as plt\n",
|
| 2649 |
+
"X = np.random.uniform(-100, 100, 1000)\n",
|
| 2650 |
+
"y = 5*(X*X) + 3\n",
|
| 2651 |
+
"plt.scatter(X, y, s=10);"
|
| 2652 |
+
]
|
| 2653 |
+
},
|
| 2654 |
+
{
|
| 2655 |
+
"cell_type": "code",
|
| 2656 |
+
"execution_count": null,
|
| 2657 |
+
"metadata": {
|
| 2658 |
+
"id": "Lr2qn7zxCKcK"
|
| 2659 |
+
},
|
| 2660 |
+
"outputs": [],
|
| 2661 |
+
"source": [
|
| 2662 |
+
"# Split the dataset into train and test sets\n",
|
| 2663 |
+
"from sklearn.model_selection import train_test_split\n",
|
| 2664 |
+
"X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)"
|
| 2665 |
+
]
|
| 2666 |
+
},
|
| 2667 |
+
{
|
| 2668 |
+
"cell_type": "code",
|
| 2669 |
+
"execution_count": null,
|
| 2670 |
+
"metadata": {
|
| 2671 |
+
"id": "UXpmCrg0CbeA"
|
| 2672 |
+
},
|
| 2673 |
+
"outputs": [],
|
| 2674 |
+
"source": [
|
| 2675 |
+
"from sklearn.neural_network import MLPRegressor\n",
|
| 2676 |
+
"# Define an MLPRegressor\n",
|
| 2677 |
+
"regr = MLPRegressor(hidden_layer_sizes=(10,), solver='lbfgs', activation='relu', max_iter=10000)\n",
|
| 2678 |
+
"# Fit on the training data\n",
|
| 2679 |
+
"regr = regr.fit(X_train.reshape(-1, 1), y_train)\n",
|
| 2680 |
+
"# Predict using the multi-layer perceptron model\n",
|
| 2681 |
+
"y_pred = regr.predict(X_test.reshape(-1, 1))\n",
|
| 2682 |
+
"# Return the coefficient of determination of the prediction. The best score can be 1.\n",
|
| 2683 |
+
"regr.score(X_test.reshape(-1, 1), y_test)"
|
| 2684 |
+
]
|
| 2685 |
+
},
|
| 2686 |
+
{
|
| 2687 |
+
"cell_type": "markdown",
|
| 2688 |
+
"metadata": {
|
| 2689 |
+
"id": "vxAwaKWhq6z5"
|
| 2690 |
+
},
|
| 2691 |
+
"source": [
|
| 2692 |
+
"## OpenCV"
|
| 2693 |
+
]
|
| 2694 |
+
},
|
| 2695 |
+
{
|
| 2696 |
+
"cell_type": "markdown",
|
| 2697 |
+
"metadata": {
|
| 2698 |
+
"id": "7S04s4KFz3EO"
|
| 2699 |
+
},
|
| 2700 |
+
"source": [
|
| 2701 |
+
"OpenCV is a library providing implementation of multitude of algorithms related to image processing, computer vision and machine learning. In this section, we will learn different image processing functions from the OpenCV library. For more details on OpenCV, please see the [OpenCV website](https://opencv.org/)."
|
| 2702 |
+
]
|
| 2703 |
+
},
|
| 2704 |
+
{
|
| 2705 |
+
"cell_type": "markdown",
|
| 2706 |
+
"metadata": {
|
| 2707 |
+
"id": "oTcZ703E2Zxh"
|
| 2708 |
+
},
|
| 2709 |
+
"source": [
|
| 2710 |
+
"### Data structures\n",
|
| 2711 |
+
"\n",
|
| 2712 |
+
"Colour images usually have three channels: red, green and blue and these channels are usually arranged in a certain order. Depending on this arrangement the image is termed in a certain way. For example, if the channels in an image are ordered in red (R), green (G) and blue (B), the image is called as RGB image. In OpenCV an image can be read by `cv2.imread()` function."
|
| 2713 |
+
]
|
| 2714 |
+
},
|
| 2715 |
+
{
|
| 2716 |
+
"cell_type": "code",
|
| 2717 |
+
"execution_count": null,
|
| 2718 |
+
"metadata": {
|
| 2719 |
+
"id": "W_6NRQ762_fP"
|
| 2720 |
+
},
|
| 2721 |
+
"outputs": [],
|
| 2722 |
+
"source": [
|
| 2723 |
+
"# read an image\n",
|
| 2724 |
+
"import cv2\n",
|
| 2725 |
+
"img = cv2.imread('images/lena.png')\n",
|
| 2726 |
+
"\n",
|
| 2727 |
+
"# show image format (basically a 3-d array of pixel colour info, in BGR format)\n",
|
| 2728 |
+
"print('Image shape: {}'.format(img.shape))\n",
|
| 2729 |
+
"print('Image: {}'.format(img))"
|
| 2730 |
+
]
|
| 2731 |
+
},
|
| 2732 |
+
{
|
| 2733 |
+
"cell_type": "markdown",
|
| 2734 |
+
"metadata": {
|
| 2735 |
+
"id": "aZPEVuGh_p8V",
|
| 2736 |
+
"pycharm": {}
|
| 2737 |
+
},
|
| 2738 |
+
"source": [
|
| 2739 |
+
"### Colour conversions\n",
|
| 2740 |
+
"By default, OpenCV loads images in BGR format. This is why the famous image of Lena looks a bit weird. **Note:** we will use imshow function from Matplotlib to display the image."
|
| 2741 |
+
]
|
| 2742 |
+
},
|
| 2743 |
+
{
|
| 2744 |
+
"cell_type": "code",
|
| 2745 |
+
"execution_count": null,
|
| 2746 |
+
"metadata": {
|
| 2747 |
+
"id": "uPBA122WEzXM",
|
| 2748 |
+
"pycharm": {}
|
| 2749 |
+
},
|
| 2750 |
+
"outputs": [],
|
| 2751 |
+
"source": [
|
| 2752 |
+
"# show image with matplotlib\n",
|
| 2753 |
+
"import matplotlib.pyplot as plt\n",
|
| 2754 |
+
"plt.imshow(img)"
|
| 2755 |
+
]
|
| 2756 |
+
},
|
| 2757 |
+
{
|
| 2758 |
+
"cell_type": "markdown",
|
| 2759 |
+
"metadata": {
|
| 2760 |
+
"id": "T89Km6qqQs3t"
|
| 2761 |
+
},
|
| 2762 |
+
"source": [
|
| 2763 |
+
"In OpenCV, a BGR image can be converted to an RGB image by the `cv2.cvtColor()` function as follows"
|
| 2764 |
+
]
|
| 2765 |
+
},
|
| 2766 |
+
{
|
| 2767 |
+
"cell_type": "code",
|
| 2768 |
+
"execution_count": null,
|
| 2769 |
+
"metadata": {
|
| 2770 |
+
"id": "_kIGhwKc_p8V",
|
| 2771 |
+
"pycharm": {},
|
| 2772 |
+
"scrolled": true
|
| 2773 |
+
},
|
| 2774 |
+
"outputs": [],
|
| 2775 |
+
"source": [
|
| 2776 |
+
"# convert image to RGB colour space\n",
|
| 2777 |
+
"img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n",
|
| 2778 |
+
"\n",
|
| 2779 |
+
"# show image with matplotlib\n",
|
| 2780 |
+
"plt.imshow(img)"
|
| 2781 |
+
]
|
| 2782 |
+
},
|
| 2783 |
+
{
|
| 2784 |
+
"cell_type": "markdown",
|
| 2785 |
+
"metadata": {
|
| 2786 |
+
"id": "8WYalGWXQs3u"
|
| 2787 |
+
},
|
| 2788 |
+
"source": [
|
| 2789 |
+
"In a similar way, a BGR image can also be converted to grayscale image which has only a single channel. Converting an RGB image into a grayscale image involves summing up the individual (RGB) components with the weights (0.299, 0.587, 0.114). The OpenCV function `cv2.cvtColor()` can also be used to convert an RGB image into a grayscale image."
|
| 2790 |
+
]
|
| 2791 |
+
},
|
| 2792 |
+
{
|
| 2793 |
+
"cell_type": "code",
|
| 2794 |
+
"execution_count": null,
|
| 2795 |
+
"metadata": {
|
| 2796 |
+
"id": "vpS6RcOV_p8Y",
|
| 2797 |
+
"pycharm": {}
|
| 2798 |
+
},
|
| 2799 |
+
"outputs": [],
|
| 2800 |
+
"source": [
|
| 2801 |
+
"# convert image to grayscale\n",
|
| 2802 |
+
"gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n",
|
| 2803 |
+
"\n",
|
| 2804 |
+
"print('Image shape: {}'.format(gray_img.shape))\n",
|
| 2805 |
+
"# grayscale image represented as a 2-d array\n",
|
| 2806 |
+
"print(gray_img)"
|
| 2807 |
+
]
|
| 2808 |
+
},
|
| 2809 |
+
{
|
| 2810 |
+
"cell_type": "markdown",
|
| 2811 |
+
"metadata": {
|
| 2812 |
+
"id": "E1yoDp2hCihi",
|
| 2813 |
+
"pycharm": {}
|
| 2814 |
+
},
|
| 2815 |
+
"source": [
|
| 2816 |
+
"Gray images have single channel"
|
| 2817 |
+
]
|
| 2818 |
+
},
|
| 2819 |
+
{
|
| 2820 |
+
"cell_type": "code",
|
| 2821 |
+
"execution_count": null,
|
| 2822 |
+
"metadata": {
|
| 2823 |
+
"id": "z2-K1pOh_p8a",
|
| 2824 |
+
"pycharm": {}
|
| 2825 |
+
},
|
| 2826 |
+
"outputs": [],
|
| 2827 |
+
"source": [
|
| 2828 |
+
"# plot the gray image, note the cmap parameter\n",
|
| 2829 |
+
"plt.imshow(gray_img, cmap='gray')"
|
| 2830 |
+
]
|
| 2831 |
+
},
|
| 2832 |
+
{
|
| 2833 |
+
"cell_type": "markdown",
|
| 2834 |
+
"metadata": {
|
| 2835 |
+
"id": "B_EOqF8qQs3v"
|
| 2836 |
+
},
|
| 2837 |
+
"source": [
|
| 2838 |
+
"Colour to grayscale is a lossy conversion. However, in OpenCV a grayscale image can approximately be converted to a colour image using the `cv2.applyColorMap()` function according to the colour maps described at this [link](https://docs.opencv.org/4.x/d3/d50/group__imgproc__colormap.html)."
|
| 2839 |
+
]
|
| 2840 |
+
},
|
| 2841 |
+
{
|
| 2842 |
+
"cell_type": "code",
|
| 2843 |
+
"execution_count": null,
|
| 2844 |
+
"metadata": {
|
| 2845 |
+
"id": "eOs27PQFSYrn",
|
| 2846 |
+
"pycharm": {}
|
| 2847 |
+
},
|
| 2848 |
+
"outputs": [],
|
| 2849 |
+
"source": [
|
| 2850 |
+
"gray_img_col = cv2.applyColorMap(gray_img, cv2.COLORMAP_JET)\n",
|
| 2851 |
+
"plt.imshow(gray_img_col)"
|
| 2852 |
+
]
|
| 2853 |
+
},
|
| 2854 |
+
{
|
| 2855 |
+
"cell_type": "markdown",
|
| 2856 |
+
"metadata": {
|
| 2857 |
+
"id": "cE_h8sgbQs3w"
|
| 2858 |
+
},
|
| 2859 |
+
"source": [
|
| 2860 |
+
"### Conversion from `uint8` to `float64` (`double`) and Normalization"
|
| 2861 |
+
]
|
| 2862 |
+
},
|
| 2863 |
+
{
|
| 2864 |
+
"cell_type": "code",
|
| 2865 |
+
"execution_count": null,
|
| 2866 |
+
"metadata": {
|
| 2867 |
+
"id": "-6aepcnKQs3w",
|
| 2868 |
+
"pycharm": {
|
| 2869 |
+
"name": "#%%\n"
|
| 2870 |
+
}
|
| 2871 |
+
},
|
| 2872 |
+
"outputs": [],
|
| 2873 |
+
"source": [
|
| 2874 |
+
"img_dble = cv2.normalize(img.astype('float64'), None, 0.0, 1.0, cv2.NORM_MINMAX)\n",
|
| 2875 |
+
"print(img_dble)"
|
| 2876 |
+
]
|
| 2877 |
+
},
|
| 2878 |
+
{
|
| 2879 |
+
"cell_type": "code",
|
| 2880 |
+
"execution_count": null,
|
| 2881 |
+
"metadata": {
|
| 2882 |
+
"id": "tOh1DQ5s5Iyg"
|
| 2883 |
+
},
|
| 2884 |
+
"outputs": [],
|
| 2885 |
+
"source": [
|
| 2886 |
+
"plt.imshow(img_dble)"
|
| 2887 |
+
]
|
| 2888 |
+
},
|
| 2889 |
+
{
|
| 2890 |
+
"cell_type": "markdown",
|
| 2891 |
+
"metadata": {
|
| 2892 |
+
"id": "PxNoVeWkwqDF"
|
| 2893 |
+
},
|
| 2894 |
+
"source": [
|
| 2895 |
+
"### Image processing\n",
|
| 2896 |
+
"Below we will review some brief image processing tasks, such as image filtering, binarization, edge detection etc with OpenCV."
|
| 2897 |
+
]
|
| 2898 |
+
},
|
| 2899 |
+
{
|
| 2900 |
+
"cell_type": "markdown",
|
| 2901 |
+
"metadata": {
|
| 2902 |
+
"id": "iIt6E6uC5aMS"
|
| 2903 |
+
},
|
| 2904 |
+
"source": [
|
| 2905 |
+
"#### Box Filtering"
|
| 2906 |
+
]
|
| 2907 |
+
},
|
| 2908 |
+
{
|
| 2909 |
+
"cell_type": "markdown",
|
| 2910 |
+
"metadata": {
|
| 2911 |
+
"id": "dAvdTtHeKQJJ",
|
| 2912 |
+
"pycharm": {}
|
| 2913 |
+
},
|
| 2914 |
+
"source": [
|
| 2915 |
+
"In this filtering, each pixel value in an image is replaced by the weighted average of the neighborhood (defined by the filter mask) intensity values. The most commonly used filter is the Box filter which has equal weights. A 3×3 normalized box filter is shown below\n",
|
| 2916 |
+
"\n",
|
| 2917 |
+
"\n",
|
| 2918 |
+
"\n",
|
| 2919 |
+
"It is a good practice to normalize the filter, this is why the above filter is divided by 9. This is to make sure that the image does not get brighter or darker. You can also use an unnormalized box filter.\n",
|
| 2920 |
+
"\n",
|
| 2921 |
+
"OpenCV provides two inbuilt functions for averaging namely:\n",
|
| 2922 |
+
"\n",
|
| 2923 |
+
"* `cv2.blur()` that blurs an image using only the normalized box filter and\n",
|
| 2924 |
+
"* `cv2.boxFilter()` which is more general, having the option of using either normalized or unnormalized box filter. Just pass an argument normalize=False to the function"
|
| 2925 |
+
]
|
| 2926 |
+
},
|
| 2927 |
+
{
|
| 2928 |
+
"cell_type": "code",
|
| 2929 |
+
"execution_count": null,
|
| 2930 |
+
"metadata": {
|
| 2931 |
+
"id": "HVb5aZSNLaGK",
|
| 2932 |
+
"pycharm": {}
|
| 2933 |
+
},
|
| 2934 |
+
"outputs": [],
|
| 2935 |
+
"source": [
|
| 2936 |
+
"img = cv2.cvtColor(cv2.imread('images/books.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 2937 |
+
"plt.imshow(img)"
|
| 2938 |
+
]
|
| 2939 |
+
},
|
| 2940 |
+
{
|
| 2941 |
+
"cell_type": "code",
|
| 2942 |
+
"execution_count": null,
|
| 2943 |
+
"metadata": {
|
| 2944 |
+
"id": "gS2OY6czd2oX",
|
| 2945 |
+
"pycharm": {}
|
| 2946 |
+
},
|
| 2947 |
+
"outputs": [],
|
| 2948 |
+
"source": [
|
| 2949 |
+
"gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n",
|
| 2950 |
+
"blur_img = cv2.blur(gray_img, (10, 10))\n",
|
| 2951 |
+
"plt.subplot(1, 2, 1); plt.imshow(gray_img, cmap='gray')\n",
|
| 2952 |
+
"plt.subplot(1, 2, 2); plt.imshow(blur_img, cmap='gray')"
|
| 2953 |
+
]
|
| 2954 |
+
},
|
| 2955 |
+
{
|
| 2956 |
+
"cell_type": "markdown",
|
| 2957 |
+
"metadata": {
|
| 2958 |
+
"id": "tOqyh7ZD5nW0"
|
| 2959 |
+
},
|
| 2960 |
+
"source": [
|
| 2961 |
+
"#### Gaussian Filtering"
|
| 2962 |
+
]
|
| 2963 |
+
},
|
| 2964 |
+
{
|
| 2965 |
+
"cell_type": "markdown",
|
| 2966 |
+
"metadata": {
|
| 2967 |
+
"id": "EhoKd7Is_p8y",
|
| 2968 |
+
"pycharm": {}
|
| 2969 |
+
},
|
| 2970 |
+
"source": [
|
| 2971 |
+
"In Gaussian filtering, instead of a box filter, a Gaussian kernel is used. In OpenCV, it is done with the function, `cv2.GaussianBlur()`. We should specify the width and height of the kernel which should be positive and odd. We also should specify the standard deviation in the X and Y directions, sigmaX and sigmaY respectively. If only sigmaX is specified, sigmaY is taken as the same as sigmaX. If both are given as zeros, they are calculated from the kernel size. Gaussian blurring is highly effective in removing Gaussian noise from an image."
|
| 2972 |
+
]
|
| 2973 |
+
},
|
| 2974 |
+
{
|
| 2975 |
+
"cell_type": "code",
|
| 2976 |
+
"execution_count": null,
|
| 2977 |
+
"metadata": {
|
| 2978 |
+
"id": "R2qNZNzB_p8z",
|
| 2979 |
+
"pycharm": {}
|
| 2980 |
+
},
|
| 2981 |
+
"outputs": [],
|
| 2982 |
+
"source": [
|
| 2983 |
+
"img = cv2.cvtColor(cv2.imread('images/oy.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 2984 |
+
"plt.imshow(img)"
|
| 2985 |
+
]
|
| 2986 |
+
},
|
| 2987 |
+
{
|
| 2988 |
+
"cell_type": "code",
|
| 2989 |
+
"execution_count": null,
|
| 2990 |
+
"metadata": {
|
| 2991 |
+
"id": "FcEcRBhT_p81",
|
| 2992 |
+
"pycharm": {}
|
| 2993 |
+
},
|
| 2994 |
+
"outputs": [],
|
| 2995 |
+
"source": [
|
| 2996 |
+
"# preproccess with blurring, with 5x5 kernel (note kernel size should be odd)\n",
|
| 2997 |
+
"img_blur_small = cv2.GaussianBlur(img, (5, 5), 0)\n",
|
| 2998 |
+
"plt.imshow(img_blur_small)"
|
| 2999 |
+
]
|
| 3000 |
+
},
|
| 3001 |
+
{
|
| 3002 |
+
"cell_type": "code",
|
| 3003 |
+
"execution_count": null,
|
| 3004 |
+
"metadata": {
|
| 3005 |
+
"id": "GcpJwBNU_p83",
|
| 3006 |
+
"pycharm": {}
|
| 3007 |
+
},
|
| 3008 |
+
"outputs": [],
|
| 3009 |
+
"source": [
|
| 3010 |
+
"img_blur_small = cv2.GaussianBlur(img, (5, 5), 25)\n",
|
| 3011 |
+
"plt.imshow(img_blur_small)"
|
| 3012 |
+
]
|
| 3013 |
+
},
|
| 3014 |
+
{
|
| 3015 |
+
"cell_type": "code",
|
| 3016 |
+
"execution_count": null,
|
| 3017 |
+
"metadata": {
|
| 3018 |
+
"id": "GW_zbFBx_p85",
|
| 3019 |
+
"pycharm": {}
|
| 3020 |
+
},
|
| 3021 |
+
"outputs": [],
|
| 3022 |
+
"source": [
|
| 3023 |
+
"img_blur_large = cv2.GaussianBlur(img, (15,15), 0)\n",
|
| 3024 |
+
"plt.imshow(img_blur_large)"
|
| 3025 |
+
]
|
| 3026 |
+
},
|
| 3027 |
+
{
|
| 3028 |
+
"cell_type": "markdown",
|
| 3029 |
+
"metadata": {
|
| 3030 |
+
"id": "knYfkspQ5y0F"
|
| 3031 |
+
},
|
| 3032 |
+
"source": [
|
| 3033 |
+
"#### Median Filtering"
|
| 3034 |
+
]
|
| 3035 |
+
},
|
| 3036 |
+
{
|
| 3037 |
+
"cell_type": "markdown",
|
| 3038 |
+
"metadata": {
|
| 3039 |
+
"id": "rFRJL4v1O7kt",
|
| 3040 |
+
"pycharm": {}
|
| 3041 |
+
},
|
| 3042 |
+
"source": [
|
| 3043 |
+
"This is a non-linear filtering technique. As clear from the name, this takes a median of all the pixels under the kernel area and replaces the central element with this median value. This is quite effective in reducing a certain type of noise (like salt-and-pepper noise) with considerably less edge blurring as compared to other linear filters of the same size. First create the function for creating noisy images with \"salt and pepper\" noise."
|
| 3044 |
+
]
|
| 3045 |
+
},
|
| 3046 |
+
{
|
| 3047 |
+
"cell_type": "code",
|
| 3048 |
+
"execution_count": null,
|
| 3049 |
+
"metadata": {
|
| 3050 |
+
"id": "dqMvcYl3u_S8",
|
| 3051 |
+
"pycharm": {}
|
| 3052 |
+
},
|
| 3053 |
+
"outputs": [],
|
| 3054 |
+
"source": [
|
| 3055 |
+
"def add_sp_noise(image, amount=0.1):\n",
|
| 3056 |
+
" row, col, ch = image.shape\n",
|
| 3057 |
+
" s_vs_p = 0.5\n",
|
| 3058 |
+
" out = np.copy(image)\n",
|
| 3059 |
+
" # Salt mode\n",
|
| 3060 |
+
" num_salt = np.ceil(amount * image.size * s_vs_p)\n",
|
| 3061 |
+
" coords = [np.random.randint(0, i - 1, int(num_salt))\n",
|
| 3062 |
+
" for i in image.shape]\n",
|
| 3063 |
+
" out[coords[0], coords[1], coords[2]] = 1\n",
|
| 3064 |
+
"\n",
|
| 3065 |
+
" # Pepper mode\n",
|
| 3066 |
+
" num_pepper = np.ceil(amount* image.size * (1. - s_vs_p))\n",
|
| 3067 |
+
" coords = [np.random.randint(0, i - 1, int(num_pepper))\n",
|
| 3068 |
+
" for i in image.shape]\n",
|
| 3069 |
+
" out[coords[0], coords[1], coords[2]] = 0\n",
|
| 3070 |
+
" return out"
|
| 3071 |
+
]
|
| 3072 |
+
},
|
| 3073 |
+
{
|
| 3074 |
+
"cell_type": "markdown",
|
| 3075 |
+
"metadata": {
|
| 3076 |
+
"id": "NBNSgJgfDgqv",
|
| 3077 |
+
"pycharm": {}
|
| 3078 |
+
},
|
| 3079 |
+
"source": [
|
| 3080 |
+
"Load an image and apply \"salt and pepper\" noise and then try to smooth it with Gaussian and Median filter"
|
| 3081 |
+
]
|
| 3082 |
+
},
|
| 3083 |
+
{
|
| 3084 |
+
"cell_type": "code",
|
| 3085 |
+
"execution_count": null,
|
| 3086 |
+
"metadata": {
|
| 3087 |
+
"id": "ea2GEwfFQLR1",
|
| 3088 |
+
"pycharm": {
|
| 3089 |
+
"is_executing": true
|
| 3090 |
+
}
|
| 3091 |
+
},
|
| 3092 |
+
"outputs": [],
|
| 3093 |
+
"source": [
|
| 3094 |
+
"img = cv2.cvtColor(cv2.imread('images/coins.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 3095 |
+
"noisy_img = add_sp_noise(img, amount=0.1)\n",
|
| 3096 |
+
"img_gaus = cv2.GaussianBlur(noisy_img, (5, 5), 3)\n",
|
| 3097 |
+
"img_med = cv2.medianBlur(noisy_img, 5)\n",
|
| 3098 |
+
"plt.subplot(1, 4, 1); plt.imshow(img); plt.title('Original')\n",
|
| 3099 |
+
"plt.subplot(1, 4, 2); plt.imshow(noisy_img); plt.title('Salt & Pepper Noise')\n",
|
| 3100 |
+
"plt.subplot(1, 4, 3); plt.imshow(img_gaus); plt.title('Gaussian Filtered')\n",
|
| 3101 |
+
"plt.subplot(1, 4, 4); plt.imshow(img_med); plt.title('Median Filtered')"
|
| 3102 |
+
]
|
| 3103 |
+
},
|
| 3104 |
+
{
|
| 3105 |
+
"cell_type": "markdown",
|
| 3106 |
+
"metadata": {
|
| 3107 |
+
"id": "jBSFKfNp58ka"
|
| 3108 |
+
},
|
| 3109 |
+
"source": [
|
| 3110 |
+
"#### Edge Detection"
|
| 3111 |
+
]
|
| 3112 |
+
},
|
| 3113 |
+
{
|
| 3114 |
+
"cell_type": "markdown",
|
| 3115 |
+
"metadata": {
|
| 3116 |
+
"id": "XxeFuSii_p9N",
|
| 3117 |
+
"pycharm": {}
|
| 3118 |
+
},
|
| 3119 |
+
"source": [
|
| 3120 |
+
"Edge detection is an image processing technique for finding the boundaries of objects within images. It works by detecting discontinuities in brightness, colour, surface etc. Edge detection is used for image segmentation and data extraction in areas such as image processing, computer vision, and machine vision. OpenCV provides the `cv2.Canny()` function to compute edges in an image."
|
| 3121 |
+
]
|
| 3122 |
+
},
|
| 3123 |
+
{
|
| 3124 |
+
"cell_type": "code",
|
| 3125 |
+
"execution_count": null,
|
| 3126 |
+
"metadata": {
|
| 3127 |
+
"id": "-utaqZp5SDP7",
|
| 3128 |
+
"pycharm": {}
|
| 3129 |
+
},
|
| 3130 |
+
"outputs": [],
|
| 3131 |
+
"source": [
|
| 3132 |
+
"cups = cv2.cvtColor(cv2.imread('images/cups.jpg'), cv2.COLOR_BGR2RGB)\n",
|
| 3133 |
+
"plt.imshow(cups)"
|
| 3134 |
+
]
|
| 3135 |
+
},
|
| 3136 |
+
{
|
| 3137 |
+
"cell_type": "code",
|
| 3138 |
+
"execution_count": null,
|
| 3139 |
+
"metadata": {
|
| 3140 |
+
"id": "7Ko1a2jmSM-M",
|
| 3141 |
+
"pycharm": {}
|
| 3142 |
+
},
|
| 3143 |
+
"outputs": [],
|
| 3144 |
+
"source": [
|
| 3145 |
+
"# preprocess by blurring and grayscale\n",
|
| 3146 |
+
"cups_preprocessed = cv2.cvtColor(cv2.GaussianBlur(cups, (7,7), 0), cv2.COLOR_RGB2GRAY)"
|
| 3147 |
+
]
|
| 3148 |
+
},
|
| 3149 |
+
{
|
| 3150 |
+
"cell_type": "code",
|
| 3151 |
+
"execution_count": null,
|
| 3152 |
+
"metadata": {
|
| 3153 |
+
"id": "a8-A44piSmwd",
|
| 3154 |
+
"pycharm": {}
|
| 3155 |
+
},
|
| 3156 |
+
"outputs": [],
|
| 3157 |
+
"source": [
|
| 3158 |
+
"# find binary image with thresholding\n",
|
| 3159 |
+
"low_thresh = 120\n",
|
| 3160 |
+
"high_thresh = 200\n",
|
| 3161 |
+
"_, cups_thresh = cv2.threshold(cups_preprocessed, low_thresh, 255, cv2.THRESH_BINARY)\n",
|
| 3162 |
+
"plt.imshow(cv2.cvtColor(cups_thresh, cv2.COLOR_GRAY2RGB))\n",
|
| 3163 |
+
"\n",
|
| 3164 |
+
"_, cups_thresh_hi = cv2.threshold(cups_preprocessed, high_thresh, 255, cv2.THRESH_BINARY)"
|
| 3165 |
+
]
|
| 3166 |
+
},
|
| 3167 |
+
{
|
| 3168 |
+
"cell_type": "code",
|
| 3169 |
+
"execution_count": null,
|
| 3170 |
+
"metadata": {
|
| 3171 |
+
"id": "lVNkDIgDRuci",
|
| 3172 |
+
"pycharm": {}
|
| 3173 |
+
},
|
| 3174 |
+
"outputs": [],
|
| 3175 |
+
"source": [
|
| 3176 |
+
"# find binary image with edges\n",
|
| 3177 |
+
"cups_edges = cv2.Canny(cups_preprocessed, threshold1=90, threshold2=110)\n",
|
| 3178 |
+
"plt.imshow(cv2.cvtColor(cups_edges, cv2.COLOR_GRAY2RGB))"
|
| 3179 |
+
]
|
| 3180 |
+
},
|
| 3181 |
+
{
|
| 3182 |
+
"cell_type": "markdown",
|
| 3183 |
+
"metadata": {
|
| 3184 |
+
"id": "-XOcVQ4hqP4R"
|
| 3185 |
+
},
|
| 3186 |
+
"source": [
|
| 3187 |
+
"## SciPy"
|
| 3188 |
+
]
|
| 3189 |
+
},
|
| 3190 |
+
{
|
| 3191 |
+
"cell_type": "markdown",
|
| 3192 |
+
"metadata": {
|
| 3193 |
+
"id": "tAWDvNu5qn4b"
|
| 3194 |
+
},
|
| 3195 |
+
"source": [
|
| 3196 |
+
"Numpy provides a high-performance multidimensional array and basic tools to compute with and manipulate these arrays. [SciPy](http://docs.scipy.org/doc/scipy/reference/) builds on this, and provides a large number of functions that operate on numpy arrays and are useful for different types of scientific and engineering applications. The best way to get familiar with SciPy is to [browse the documentation](https://docs.scipy.org/doc/scipy/reference/index.html). SciPy provides important functionalities for reading and writing MATLAB files, which show below.\n",
|
| 3197 |
+
"\n",
|
| 3198 |
+
"\n"
|
| 3199 |
+
]
|
| 3200 |
+
},
|
| 3201 |
+
{
|
| 3202 |
+
"cell_type": "markdown",
|
| 3203 |
+
"metadata": {
|
| 3204 |
+
"id": "ajs-UbqSrWk0"
|
| 3205 |
+
},
|
| 3206 |
+
"source": [
|
| 3207 |
+
"###MATLAB files"
|
| 3208 |
+
]
|
| 3209 |
+
},
|
| 3210 |
+
{
|
| 3211 |
+
"cell_type": "markdown",
|
| 3212 |
+
"metadata": {
|
| 3213 |
+
"id": "HoT2zazhrZ5m"
|
| 3214 |
+
},
|
| 3215 |
+
"source": [
|
| 3216 |
+
"The functions `scipy.io.loadmat` and `scipy.io.savemat` allow you to respectively read and write MATLAB files. You can read about them [in the documentation](http://docs.scipy.org/doc/scipy/reference/io.html)."
|
| 3217 |
+
]
|
| 3218 |
+
},
|
| 3219 |
+
{
|
| 3220 |
+
"cell_type": "markdown",
|
| 3221 |
+
"metadata": {
|
| 3222 |
+
"id": "iMtaY6Bzr7-w"
|
| 3223 |
+
},
|
| 3224 |
+
"source": [
|
| 3225 |
+
"###Distance between points"
|
| 3226 |
+
]
|
| 3227 |
+
},
|
| 3228 |
+
{
|
| 3229 |
+
"cell_type": "markdown",
|
| 3230 |
+
"metadata": {
|
| 3231 |
+
"id": "1tq9Mtwkr_s0"
|
| 3232 |
+
},
|
| 3233 |
+
"source": [
|
| 3234 |
+
"SciPy defines some useful functions for computing distances between sets of points.\n",
|
| 3235 |
+
"\n",
|
| 3236 |
+
"The function `scipy.spatial.distance.pdist` computes the distance between all pairs of points in a given set:"
|
| 3237 |
+
]
|
| 3238 |
+
},
|
| 3239 |
+
{
|
| 3240 |
+
"cell_type": "code",
|
| 3241 |
+
"execution_count": null,
|
| 3242 |
+
"metadata": {
|
| 3243 |
+
"id": "EwlHRO0jsJBI"
|
| 3244 |
+
},
|
| 3245 |
+
"outputs": [],
|
| 3246 |
+
"source": [
|
| 3247 |
+
"import numpy as np\n",
|
| 3248 |
+
"from scipy.spatial.distance import pdist, squareform\n",
|
| 3249 |
+
"\n",
|
| 3250 |
+
"# Create the following array where each row is a point in 2D space:\n",
|
| 3251 |
+
"# [[0 1]\n",
|
| 3252 |
+
"# [1 0]\n",
|
| 3253 |
+
"# [2 0]]\n",
|
| 3254 |
+
"x = np.array([[0, 1], [1, 0], [2, 0]])\n",
|
| 3255 |
+
"print(x)\n",
|
| 3256 |
+
"\n",
|
| 3257 |
+
"# Compute the Euclidean distance between all rows of x.\n",
|
| 3258 |
+
"# d[i, j] is the Euclidean distance between x[i, :] and x[j, :],\n",
|
| 3259 |
+
"# and d is the following array:\n",
|
| 3260 |
+
"# [[ 0. 1.41421356 2.23606798]\n",
|
| 3261 |
+
"# [ 1.41421356 0. 1. ]\n",
|
| 3262 |
+
"# [ 2.23606798 1. 0. ]]\n",
|
| 3263 |
+
"d = squareform(pdist(x, 'euclidean'))\n",
|
| 3264 |
+
"print(d)"
|
| 3265 |
+
]
|
| 3266 |
+
},
|
| 3267 |
+
{
|
| 3268 |
+
"cell_type": "markdown",
|
| 3269 |
+
"metadata": {
|
| 3270 |
+
"id": "dzYk_QSfsSXO"
|
| 3271 |
+
},
|
| 3272 |
+
"source": [
|
| 3273 |
+
"A similar function (`scipy.spatial.distance.cdist`) computes the distance between all pairs across two sets of points; you can read about it [in the documentation](https://docs.scipy.org/doc/scipy/reference/generated/scipy.spatial.distance.cdist.html)."
|
| 3274 |
+
]
|
| 3275 |
+
},
|
| 3276 |
+
{
|
| 3277 |
+
"cell_type": "markdown",
|
| 3278 |
+
"metadata": {
|
| 3279 |
+
"id": "d3XD7jVkU9Z3"
|
| 3280 |
+
},
|
| 3281 |
+
"source": [
|
| 3282 |
+
"#### Acknowledgement\n",
|
| 3283 |
+
"This tutorial was originally written by [Justin Johnson](https://web.eecs.umich.edu/~justincj/) for CS231n at the Stanford University. This version has been adapted and modified by [Anjan Dutta](https://www.surrey.ac.uk/people/anjan-dutta) for the Spring 2023 edition of [EEEM068](https://catalogue.surrey.ac.uk/2022-3/module/EEEM068) module at the University of Surrey."
|
| 3284 |
+
]
|
| 3285 |
+
}
|
| 3286 |
+
],
|
| 3287 |
+
"metadata": {
|
| 3288 |
+
"colab": {
|
| 3289 |
+
"include_colab_link": true,
|
| 3290 |
+
"name": "colab-tutorial.ipynb",
|
| 3291 |
+
"provenance": []
|
| 3292 |
+
},
|
| 3293 |
+
"kernelspec": {
|
| 3294 |
+
"display_name": "Python 3 (ipykernel)",
|
| 3295 |
+
"language": "python",
|
| 3296 |
+
"name": "python3"
|
| 3297 |
+
},
|
| 3298 |
+
"language_info": {
|
| 3299 |
+
"codemirror_mode": {
|
| 3300 |
+
"name": "ipython",
|
| 3301 |
+
"version": 3
|
| 3302 |
+
},
|
| 3303 |
+
"file_extension": ".py",
|
| 3304 |
+
"mimetype": "text/x-python",
|
| 3305 |
+
"name": "python",
|
| 3306 |
+
"nbconvert_exporter": "python",
|
| 3307 |
+
"pygments_lexer": "ipython3",
|
| 3308 |
+
"version": "3.12.3"
|
| 3309 |
+
}
|
| 3310 |
+
},
|
| 3311 |
+
"nbformat": 4,
|
| 3312 |
+
"nbformat_minor": 1
|
| 3313 |
+
}
|
Downloads/.~lock.deformation_experiments(3).pptx#
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
,rk01499,otter34.eps.surrey.ac.uk,04.07.2026 12:19,file:///user/HS400/rk01499/.config/libreoffice/4;
|
Downloads/HjxNnnnu.html
ADDED
|
File without changes
|
Downloads/deformation_experiments(1).pptx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:096e84b359a596e84faf7f2756fd674a5ad0d9d43b04bf325e8c647c34ff773a
|
| 3 |
+
size 298687
|
Downloads/deformation_experiments(2).pptx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1813487f48ef15c4b51c6c7bed6d397a3523237c5aaaea6b8bf7b7b76f1c2511
|
| 3 |
+
size 496442
|
Downloads/deformation_experiments(3).pptx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:573cba301f2ebd3051bf7626390220d268b2b1d0356aaf41a500c663b116080c
|
| 3 |
+
size 504332
|
Downloads/deformation_experiments.pptx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:096e84b359a596e84faf7f2756fd674a5ad0d9d43b04bf325e8c647c34ff773a
|
| 3 |
+
size 298687
|
Downloads/gap_zoomed_comparison.png
ADDED
|
Downloads/geometric_solver(1).py
ADDED
|
@@ -0,0 +1,1147 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Geometric Solver for Multi-Object Sketch Animation
|
| 3 |
+
===================================================
|
| 4 |
+
Pipeline:
|
| 5 |
+
SVG + text instruction
|
| 6 |
+
-> Stage 1: Qwen keyframe prompt decomposition
|
| 7 |
+
-> Stage 2: Grounding DINO object segmentation + control point assignment
|
| 8 |
+
-> Stage 3: Qwen semantic motion plan per object per keyframe
|
| 9 |
+
-> Stage 4: Geometric solver (skeleton + skinning + interpolation)
|
| 10 |
+
-> Stage 5: Rasterize frames -> Wan2.2 adjacent pairs -> final video
|
| 11 |
+
|
| 12 |
+
Dependencies:
|
| 13 |
+
pip install controlnet-aux cairosvg scipy numpy pillow svgpathtools
|
| 14 |
+
pip install torch torchvision transformers
|
| 15 |
+
pip install openai # for Qwen/GPT-4 API calls
|
| 16 |
+
"""
|
| 17 |
+
|
| 18 |
+
import os
|
| 19 |
+
import io
|
| 20 |
+
import json
|
| 21 |
+
import numpy as np
|
| 22 |
+
from PIL import Image, ImageDraw
|
| 23 |
+
from scipy.interpolate import CubicSpline
|
| 24 |
+
import xml.etree.ElementTree as ET
|
| 25 |
+
import cairosvg
|
| 26 |
+
import torch
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
# =============================================================================
|
| 30 |
+
# STAGE 1 — QWEN KEYFRAME PROMPT DECOMPOSITION
|
| 31 |
+
# =============================================================================
|
| 32 |
+
|
| 33 |
+
def decompose_keyframe_prompts(image_path, text_instruction, n_keyframes=5, client=None):
|
| 34 |
+
"""
|
| 35 |
+
Give image + text instruction to Qwen/GPT-4
|
| 36 |
+
Returns ordered list of keyframe semantic descriptions
|
| 37 |
+
|
| 38 |
+
Args:
|
| 39 |
+
image_path: path to rasterized sketch image
|
| 40 |
+
text_instruction: e.g. "basketball player takes a jump shot toward hoop"
|
| 41 |
+
n_keyframes: number of keyframes to decompose into
|
| 42 |
+
client: OpenAI-compatible API client
|
| 43 |
+
|
| 44 |
+
Returns:
|
| 45 |
+
list of keyframe description strings
|
| 46 |
+
"""
|
| 47 |
+
import base64
|
| 48 |
+
|
| 49 |
+
with open(image_path, "rb") as f:
|
| 50 |
+
image_b64 = base64.b64encode(f.read()).decode("utf-8")
|
| 51 |
+
|
| 52 |
+
prompt = f"""
|
| 53 |
+
You are given a sketch image and a motion description.
|
| 54 |
+
Decompose the motion into exactly {n_keyframes} ordered keyframe descriptions.
|
| 55 |
+
|
| 56 |
+
Motion: {text_instruction}
|
| 57 |
+
|
| 58 |
+
Rules:
|
| 59 |
+
- Each keyframe must describe the state of ALL objects in the scene
|
| 60 |
+
- Descriptions must be temporally ordered (start to end)
|
| 61 |
+
- Each description should be 1-2 sentences
|
| 62 |
+
- Focus on pose, position, and inter-object relationships
|
| 63 |
+
- Output ONLY a JSON array of {n_keyframes} strings, nothing else
|
| 64 |
+
|
| 65 |
+
Example output format:
|
| 66 |
+
["player standing, ball at hip", "player crouching, ball gripped", ...]
|
| 67 |
+
"""
|
| 68 |
+
|
| 69 |
+
# --- replace with your actual API call ---
|
| 70 |
+
# response = client.chat.completions.create(
|
| 71 |
+
# model="gpt-4o",
|
| 72 |
+
# messages=[{
|
| 73 |
+
# "role": "user",
|
| 74 |
+
# "content": [
|
| 75 |
+
# {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{image_b64}"}},
|
| 76 |
+
# {"type": "text", "text": prompt}
|
| 77 |
+
# ]
|
| 78 |
+
# }]
|
| 79 |
+
# )
|
| 80 |
+
# raw = response.choices[0].message.content
|
| 81 |
+
# keyframe_prompts = json.loads(raw)
|
| 82 |
+
|
| 83 |
+
# --- MOCK OUTPUT for basketball example ---
|
| 84 |
+
keyframe_prompts = [
|
| 85 |
+
"player standing upright, ball held at hip height with both hands, hoop visible in background",
|
| 86 |
+
"player crouching with knees bent, ball pulled back toward chest, preparing to jump",
|
| 87 |
+
"player leaving ground, arm extending upward with ball, body rising",
|
| 88 |
+
"player at peak height, arm fully extended above head, ball releasing from fingertips toward hoop",
|
| 89 |
+
"player descending, arm following through, ball mid-arc toward hoop"
|
| 90 |
+
]
|
| 91 |
+
|
| 92 |
+
return keyframe_prompts
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
# =============================================================================
|
| 96 |
+
# STAGE 2 — OBJECT SEGMENTATION + CONTROL POINT ASSIGNMENT
|
| 97 |
+
# =============================================================================
|
| 98 |
+
|
| 99 |
+
def rasterize_svg(svg_path, width=512, height=512):
|
| 100 |
+
"""Convert SVG to raster numpy array"""
|
| 101 |
+
png_data = cairosvg.svg2png(
|
| 102 |
+
url=svg_path,
|
| 103 |
+
output_width=width,
|
| 104 |
+
output_height=height
|
| 105 |
+
)
|
| 106 |
+
image = Image.open(io.BytesIO(png_data)).convert("RGB")
|
| 107 |
+
return np.array(image)
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
def get_object_bounding_boxes(image_array, object_names):
|
| 111 |
+
"""
|
| 112 |
+
Run Grounding DINO to get bounding boxes per object
|
| 113 |
+
|
| 114 |
+
Args:
|
| 115 |
+
image_array: numpy HxWx3 image
|
| 116 |
+
object_names: list of object name strings e.g. ["player", "basketball", "hoop"]
|
| 117 |
+
|
| 118 |
+
Returns:
|
| 119 |
+
dict {object_name: [x, y, w, h]}
|
| 120 |
+
"""
|
| 121 |
+
# --- real Grounding DINO call ---
|
| 122 |
+
# from groundingdino.util.inference import load_model, predict
|
| 123 |
+
# model = load_model(...)
|
| 124 |
+
# boxes, logits, phrases = predict(model, image, text_prompt, ...)
|
| 125 |
+
|
| 126 |
+
# --- MOCK for basketball example (pixel coordinates at 512x512) ---
|
| 127 |
+
bounding_boxes = {
|
| 128 |
+
"player": [6, 16, 129, 212],
|
| 129 |
+
"basketball": [165, 52, 51, 49],
|
| 130 |
+
"hoop": [380, 80, 60, 40]
|
| 131 |
+
}
|
| 132 |
+
return bounding_boxes
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
def parse_svg_control_points(svg_path):
|
| 136 |
+
"""
|
| 137 |
+
Parse all cubic bezier control points from SVG
|
| 138 |
+
|
| 139 |
+
Returns:
|
| 140 |
+
list of dicts: {id, x, y, path_id, point_type}
|
| 141 |
+
point_type: 'anchor' or 'control'
|
| 142 |
+
"""
|
| 143 |
+
try:
|
| 144 |
+
from svgpathtools import svg2paths
|
| 145 |
+
paths, attributes = svg2paths(svg_path)
|
| 146 |
+
except Exception as e:
|
| 147 |
+
print(f"svgpathtools failed: {e}. Using fallback ET parser.")
|
| 148 |
+
return _parse_svg_fallback(svg_path)
|
| 149 |
+
|
| 150 |
+
control_points = []
|
| 151 |
+
point_id = 0
|
| 152 |
+
|
| 153 |
+
for path_idx, path in enumerate(paths):
|
| 154 |
+
for segment in path:
|
| 155 |
+
# CubicBezier: start, control1, control2, end
|
| 156 |
+
pts = [segment.start, segment.control1,
|
| 157 |
+
segment.control2, segment.end]
|
| 158 |
+
types = ['anchor', 'control', 'control', 'anchor']
|
| 159 |
+
|
| 160 |
+
for pt, ptype in zip(pts, types):
|
| 161 |
+
control_points.append({
|
| 162 |
+
'id': point_id,
|
| 163 |
+
'x': pt.real,
|
| 164 |
+
'y': pt.imag,
|
| 165 |
+
'path_id': path_idx,
|
| 166 |
+
'point_type': ptype
|
| 167 |
+
})
|
| 168 |
+
point_id += 1
|
| 169 |
+
|
| 170 |
+
return control_points
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
def _parse_svg_fallback(svg_path):
|
| 174 |
+
"""Fallback SVG parser using ElementTree"""
|
| 175 |
+
import re
|
| 176 |
+
tree = ET.parse(svg_path)
|
| 177 |
+
root = tree.getroot()
|
| 178 |
+
|
| 179 |
+
ns_map = {'svg': 'http://www.w3.org/2000/svg'}
|
| 180 |
+
control_points = []
|
| 181 |
+
point_id = 0
|
| 182 |
+
|
| 183 |
+
for path_idx, elem in enumerate(
|
| 184 |
+
root.iter('{http://www.w3.org/2000/svg}path')
|
| 185 |
+
):
|
| 186 |
+
d = elem.get('d', '')
|
| 187 |
+
nums = re.findall(r'[-+]?\d*\.?\d+', d)
|
| 188 |
+
coords = [float(n) for n in nums]
|
| 189 |
+
|
| 190 |
+
for i in range(0, len(coords) - 1, 2):
|
| 191 |
+
control_points.append({
|
| 192 |
+
'id': point_id,
|
| 193 |
+
'x': coords[i],
|
| 194 |
+
'y': coords[i + 1],
|
| 195 |
+
'path_id': path_idx,
|
| 196 |
+
'point_type': 'anchor'
|
| 197 |
+
})
|
| 198 |
+
point_id += 1
|
| 199 |
+
|
| 200 |
+
return control_points
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def assign_control_points_to_objects(control_points, bounding_boxes, svg_width=512, svg_height=512):
|
| 204 |
+
"""
|
| 205 |
+
Assign each control point to the object whose bounding box
|
| 206 |
+
center is nearest to the point.
|
| 207 |
+
|
| 208 |
+
Uses MoSketch's nearest-center strategy.
|
| 209 |
+
|
| 210 |
+
Returns:
|
| 211 |
+
dict {object_name: [list of control point dicts]}
|
| 212 |
+
"""
|
| 213 |
+
object_assignments = {name: [] for name in bounding_boxes}
|
| 214 |
+
object_assignments['unassigned'] = []
|
| 215 |
+
|
| 216 |
+
# compute bounding box centers
|
| 217 |
+
centers = {}
|
| 218 |
+
for name, bb in bounding_boxes.items():
|
| 219 |
+
x, y, w, h = bb
|
| 220 |
+
centers[name] = np.array([x + w / 2, y + h / 2])
|
| 221 |
+
|
| 222 |
+
object_names = list(centers.keys())
|
| 223 |
+
center_array = np.array([centers[n] for n in object_names])
|
| 224 |
+
|
| 225 |
+
for pt in control_points:
|
| 226 |
+
pt_pos = np.array([pt['x'], pt['y']])
|
| 227 |
+
|
| 228 |
+
# scale to image coordinates if SVG uses different coordinate space
|
| 229 |
+
pt_pos_scaled = pt_pos * np.array([
|
| 230 |
+
512 / svg_width,
|
| 231 |
+
512 / svg_height
|
| 232 |
+
])
|
| 233 |
+
|
| 234 |
+
distances = np.linalg.norm(center_array - pt_pos_scaled, axis=1)
|
| 235 |
+
nearest_idx = np.argmin(distances)
|
| 236 |
+
nearest_name = object_names[nearest_idx]
|
| 237 |
+
|
| 238 |
+
pt['assigned_object'] = nearest_name
|
| 239 |
+
pt['distance_to_center'] = distances[nearest_idx]
|
| 240 |
+
object_assignments[nearest_name].append(pt)
|
| 241 |
+
|
| 242 |
+
return object_assignments
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
# =============================================================================
|
| 246 |
+
# STAGE 3 — SEMANTIC MOTION PLAN PER OBJECT PER KEYFRAME
|
| 247 |
+
# =============================================================================
|
| 248 |
+
|
| 249 |
+
def get_semantic_motion_plan(
|
| 250 |
+
image_path,
|
| 251 |
+
bounding_boxes,
|
| 252 |
+
keyframe_prompts,
|
| 253 |
+
client=None
|
| 254 |
+
):
|
| 255 |
+
"""
|
| 256 |
+
Give Qwen: image + BB assignments + keyframe prompts
|
| 257 |
+
Returns: per-object per-keyframe semantic pose descriptions
|
| 258 |
+
|
| 259 |
+
Returns:
|
| 260 |
+
dict {object_name: {keyframe_idx: semantic_description}}
|
| 261 |
+
"""
|
| 262 |
+
import base64
|
| 263 |
+
with open(image_path, "rb") as f:
|
| 264 |
+
image_b64 = base64.b64encode(f.read()).decode("utf-8")
|
| 265 |
+
|
| 266 |
+
bb_str = json.dumps(bounding_boxes, indent=2)
|
| 267 |
+
kf_str = json.dumps(
|
| 268 |
+
{i: p for i, p in enumerate(keyframe_prompts)},
|
| 269 |
+
indent=2
|
| 270 |
+
)
|
| 271 |
+
|
| 272 |
+
prompt = f"""
|
| 273 |
+
You are given a sketch image with detected objects at these bounding boxes:
|
| 274 |
+
{bb_str}
|
| 275 |
+
|
| 276 |
+
The animation has these keyframe descriptions:
|
| 277 |
+
{kf_str}
|
| 278 |
+
|
| 279 |
+
For each object at each keyframe, describe its specific pose state.
|
| 280 |
+
Use vocabulary from this list where possible:
|
| 281 |
+
- Arms: "arm at side", "arm extended upward", "arm pulling back", "arm extended forward"
|
| 282 |
+
- Legs: "legs straight", "crouching", "jumping", "kicking", "landing"
|
| 283 |
+
- Torso: "torso upright", "torso leaning forward", "torso rotating right"
|
| 284 |
+
- Position: "stationary", "moving left", "moving right", "rising", "falling"
|
| 285 |
+
|
| 286 |
+
Output ONLY a JSON object, no other text.
|
| 287 |
+
Format:
|
| 288 |
+
{{
|
| 289 |
+
"player": {{
|
| 290 |
+
"0": "torso upright, arm at side, legs straight",
|
| 291 |
+
"1": "crouching, arm pulling back, legs bent",
|
| 292 |
+
...
|
| 293 |
+
}},
|
| 294 |
+
"basketball": {{
|
| 295 |
+
"0": "stationary at hip",
|
| 296 |
+
"1": "moving upward",
|
| 297 |
+
...
|
| 298 |
+
}}
|
| 299 |
+
}}
|
| 300 |
+
"""
|
| 301 |
+
|
| 302 |
+
# --- replace with actual API call ---
|
| 303 |
+
# response = client.chat.completions.create(...)
|
| 304 |
+
# plan = json.loads(response.choices[0].message.content)
|
| 305 |
+
|
| 306 |
+
# --- MOCK for basketball example ---
|
| 307 |
+
plan = {
|
| 308 |
+
"player": {
|
| 309 |
+
"0": "torso upright, arm at side, legs straight, stationary",
|
| 310 |
+
"1": "crouching, arm pulling back, legs bent",
|
| 311 |
+
"2": "torso leaning forward, arm extended upward, jumping",
|
| 312 |
+
"3": "torso upright, arm fully extended upward, peak height",
|
| 313 |
+
"4": "torso leaning forward, arm extended forward, landing"
|
| 314 |
+
},
|
| 315 |
+
"basketball": {
|
| 316 |
+
"0": "stationary at hip level",
|
| 317 |
+
"1": "moving upward, held by player",
|
| 318 |
+
"2": "rising, releasing from hand",
|
| 319 |
+
"3": "peak arc, mid-air",
|
| 320 |
+
"4": "falling, descending toward hoop"
|
| 321 |
+
},
|
| 322 |
+
"hoop": {
|
| 323 |
+
"0": "stationary",
|
| 324 |
+
"1": "stationary",
|
| 325 |
+
"2": "stationary",
|
| 326 |
+
"3": "stationary",
|
| 327 |
+
"4": "stationary"
|
| 328 |
+
}
|
| 329 |
+
}
|
| 330 |
+
|
| 331 |
+
return plan
|
| 332 |
+
|
| 333 |
+
|
| 334 |
+
# =============================================================================
|
| 335 |
+
# STAGE 4 — GEOMETRIC SOLVER
|
| 336 |
+
# =============================================================================
|
| 337 |
+
|
| 338 |
+
# --- 4a: Skeleton Extraction via DWPose ---
|
| 339 |
+
|
| 340 |
+
SKELETON_HIERARCHY = {
|
| 341 |
+
# child -> parent
|
| 342 |
+
"right_elbow": "right_shoulder",
|
| 343 |
+
"right_wrist": "right_elbow",
|
| 344 |
+
"left_elbow": "left_shoulder",
|
| 345 |
+
"left_wrist": "left_elbow",
|
| 346 |
+
"right_knee": "right_hip",
|
| 347 |
+
"right_ankle": "right_knee",
|
| 348 |
+
"left_knee": "left_hip",
|
| 349 |
+
"left_ankle": "left_knee",
|
| 350 |
+
"right_shoulder":"neck",
|
| 351 |
+
"left_shoulder": "neck",
|
| 352 |
+
"neck": "nose",
|
| 353 |
+
"right_hip": "spine",
|
| 354 |
+
"left_hip": "spine",
|
| 355 |
+
}
|
| 356 |
+
|
| 357 |
+
# COCO keypoint order from DWPose
|
| 358 |
+
COCO_KEYPOINTS = [
|
| 359 |
+
"nose", "left_eye", "right_eye", "left_ear", "right_ear",
|
| 360 |
+
"left_shoulder", "right_shoulder", "left_elbow", "right_elbow",
|
| 361 |
+
"left_wrist", "right_wrist", "left_hip", "right_hip",
|
| 362 |
+
"left_knee", "right_knee", "left_ankle", "right_ankle"
|
| 363 |
+
]
|
| 364 |
+
|
| 365 |
+
|
| 366 |
+
def extract_skeleton_dwpose(image_array):
|
| 367 |
+
"""
|
| 368 |
+
Run DWPose on rasterized sketch image
|
| 369 |
+
Returns joint positions dict {joint_name: (x, y)}
|
| 370 |
+
Returns None if no person detected
|
| 371 |
+
"""
|
| 372 |
+
try:
|
| 373 |
+
from controlnet_aux import DWposeDetector
|
| 374 |
+
detector = DWposeDetector()
|
| 375 |
+
pil_image = Image.fromarray(image_array)
|
| 376 |
+
result = detector(pil_image, return_pil=False)
|
| 377 |
+
|
| 378 |
+
# result is dict with 'bodies' containing keypoints
|
| 379 |
+
if result is None or 'bodies' not in result:
|
| 380 |
+
print("DWPose: no person detected")
|
| 381 |
+
return None
|
| 382 |
+
|
| 383 |
+
keypoints = result['bodies']['candidate'] # shape: (N, 2)
|
| 384 |
+
|
| 385 |
+
joints = {}
|
| 386 |
+
for idx, name in enumerate(COCO_KEYPOINTS):
|
| 387 |
+
if idx < len(keypoints):
|
| 388 |
+
kp = keypoints[idx]
|
| 389 |
+
# confidence threshold
|
| 390 |
+
if len(kp) > 2 and kp[2] < 0.3:
|
| 391 |
+
continue
|
| 392 |
+
joints[name] = (float(kp[0]), float(kp[1]))
|
| 393 |
+
|
| 394 |
+
return joints
|
| 395 |
+
|
| 396 |
+
except ImportError:
|
| 397 |
+
print("controlnet_aux not installed. Using mock skeleton.")
|
| 398 |
+
return _mock_skeleton_basketball()
|
| 399 |
+
except Exception as e:
|
| 400 |
+
print(f"DWPose failed: {e}. Using mock skeleton.")
|
| 401 |
+
return _mock_skeleton_basketball()
|
| 402 |
+
|
| 403 |
+
|
| 404 |
+
def _mock_skeleton_basketball():
|
| 405 |
+
"""Mock skeleton for basketball player at 512x512"""
|
| 406 |
+
return {
|
| 407 |
+
"nose": (71, 28),
|
| 408 |
+
"neck": (71, 55),
|
| 409 |
+
"right_shoulder": (50, 75),
|
| 410 |
+
"left_shoulder": (92, 75),
|
| 411 |
+
"right_elbow": (35, 115),
|
| 412 |
+
"left_elbow": (107, 115),
|
| 413 |
+
"right_wrist": (25, 150),
|
| 414 |
+
"left_wrist": (117, 150),
|
| 415 |
+
"right_hip": (55, 155),
|
| 416 |
+
"left_hip": (87, 155),
|
| 417 |
+
"spine": (71, 115),
|
| 418 |
+
"right_knee": (50, 195),
|
| 419 |
+
"left_knee": (92, 195),
|
| 420 |
+
"right_ankle": (45, 228),
|
| 421 |
+
"left_ankle": (97, 228),
|
| 422 |
+
}
|
| 423 |
+
|
| 424 |
+
|
| 425 |
+
def assign_control_points_to_joints(control_points_for_object, joints):
|
| 426 |
+
"""
|
| 427 |
+
For each control point in an object, find nearest joint
|
| 428 |
+
Returns dict {point_id: joint_name}
|
| 429 |
+
"""
|
| 430 |
+
if not joints:
|
| 431 |
+
return {}
|
| 432 |
+
|
| 433 |
+
joint_names = list(joints.keys())
|
| 434 |
+
joint_positions = np.array([joints[j] for j in joint_names])
|
| 435 |
+
|
| 436 |
+
point_to_joint = {}
|
| 437 |
+
for pt in control_points_for_object:
|
| 438 |
+
pt_pos = np.array([pt['x'], pt['y']])
|
| 439 |
+
distances = np.linalg.norm(joint_positions - pt_pos, axis=1)
|
| 440 |
+
nearest_joint = joint_names[np.argmin(distances)]
|
| 441 |
+
point_to_joint[pt['id']] = nearest_joint
|
| 442 |
+
|
| 443 |
+
return point_to_joint
|
| 444 |
+
|
| 445 |
+
|
| 446 |
+
# --- 4b: Semantic to Joint Angles ---
|
| 447 |
+
|
| 448 |
+
SEMANTIC_TO_POSE = {
|
| 449 |
+
# ARM STATES
|
| 450 |
+
"arm at side": {
|
| 451 |
+
"right_shoulder": 0, "right_elbow": 10,
|
| 452 |
+
"left_shoulder": 0, "left_elbow": 10
|
| 453 |
+
},
|
| 454 |
+
"arm pulling back": {
|
| 455 |
+
"right_shoulder": -30, "right_elbow": 90,
|
| 456 |
+
"left_shoulder": 30, "left_elbow": 45
|
| 457 |
+
},
|
| 458 |
+
"arm extended upward": {
|
| 459 |
+
"right_shoulder": -150, "right_elbow": 170,
|
| 460 |
+
"left_shoulder": -150, "left_elbow": 170
|
| 461 |
+
},
|
| 462 |
+
"arm fully extended upward": {
|
| 463 |
+
"right_shoulder": -170, "right_elbow": 175,
|
| 464 |
+
"left_shoulder": -170, "left_elbow": 175
|
| 465 |
+
},
|
| 466 |
+
"arm extended forward": {
|
| 467 |
+
"right_shoulder": -90, "right_elbow": 160,
|
| 468 |
+
"left_shoulder": -90, "left_elbow": 160
|
| 469 |
+
},
|
| 470 |
+
|
| 471 |
+
# LEG STATES
|
| 472 |
+
"legs straight": {
|
| 473 |
+
"right_hip": 0, "right_knee": 0,
|
| 474 |
+
"left_hip": 0, "left_knee": 0
|
| 475 |
+
},
|
| 476 |
+
"legs bent": {
|
| 477 |
+
"right_hip": -30, "right_knee": -60,
|
| 478 |
+
"left_hip": -30, "left_knee": -60
|
| 479 |
+
},
|
| 480 |
+
"crouching": {
|
| 481 |
+
"right_hip": -45, "right_knee": -90,
|
| 482 |
+
"left_hip": -45, "left_knee": -90
|
| 483 |
+
},
|
| 484 |
+
"jumping": {
|
| 485 |
+
"right_hip": 20, "right_knee": 150,
|
| 486 |
+
"left_hip": 20, "left_knee": 150
|
| 487 |
+
},
|
| 488 |
+
"landing": {
|
| 489 |
+
"right_hip": -20, "right_knee": -40,
|
| 490 |
+
"left_hip": -20, "left_knee": -40
|
| 491 |
+
},
|
| 492 |
+
"kicking": {
|
| 493 |
+
"right_hip": -70, "right_knee": 160,
|
| 494 |
+
"left_hip": 10, "left_knee": 0
|
| 495 |
+
},
|
| 496 |
+
|
| 497 |
+
# TORSO STATES
|
| 498 |
+
"torso upright": {
|
| 499 |
+
"spine": 0
|
| 500 |
+
},
|
| 501 |
+
"torso leaning forward": {
|
| 502 |
+
"spine": 25
|
| 503 |
+
},
|
| 504 |
+
"torso rotating right": {
|
| 505 |
+
"spine": 20
|
| 506 |
+
},
|
| 507 |
+
|
| 508 |
+
# VERTICAL POSITION (affects bounding box center y)
|
| 509 |
+
"stationary": {"_translate_y": 0},
|
| 510 |
+
"rising": {"_translate_y": -15},
|
| 511 |
+
"peak height": {"_translate_y": -30},
|
| 512 |
+
"falling": {"_translate_y": -20},
|
| 513 |
+
"moving left": {"_translate_x": -20},
|
| 514 |
+
"moving right": {"_translate_x": 20},
|
| 515 |
+
}
|
| 516 |
+
|
| 517 |
+
|
| 518 |
+
def parse_semantic_to_joint_angles(semantic_description):
|
| 519 |
+
"""
|
| 520 |
+
Match semantic description string to joint angle dict
|
| 521 |
+
Multiple keywords can match and get merged
|
| 522 |
+
|
| 523 |
+
Returns:
|
| 524 |
+
dict {joint_name: angle_degrees}
|
| 525 |
+
plus optional _translate_x, _translate_y keys
|
| 526 |
+
"""
|
| 527 |
+
desc_lower = semantic_description.lower()
|
| 528 |
+
merged = {}
|
| 529 |
+
|
| 530 |
+
for key, angles in SEMANTIC_TO_POSE.items():
|
| 531 |
+
if key in desc_lower:
|
| 532 |
+
# later matches override earlier ones for same joint
|
| 533 |
+
merged.update(angles)
|
| 534 |
+
|
| 535 |
+
return merged
|
| 536 |
+
|
| 537 |
+
|
| 538 |
+
# --- 4c: Compute Joint Transforms ---
|
| 539 |
+
|
| 540 |
+
def compute_rotation_matrix(angle_degrees):
|
| 541 |
+
"""2D rotation matrix"""
|
| 542 |
+
theta = np.radians(angle_degrees)
|
| 543 |
+
return np.array([
|
| 544 |
+
[np.cos(theta), -np.sin(theta)],
|
| 545 |
+
[np.sin(theta), np.cos(theta)]
|
| 546 |
+
])
|
| 547 |
+
|
| 548 |
+
|
| 549 |
+
def compute_joint_transforms(initial_joints, target_joint_angles):
|
| 550 |
+
"""
|
| 551 |
+
Compute per-joint rotation transforms
|
| 552 |
+
relative to initial skeleton pose
|
| 553 |
+
|
| 554 |
+
Returns:
|
| 555 |
+
dict {joint_name: {R: 2x2 matrix, pivot: (x,y)}}
|
| 556 |
+
"""
|
| 557 |
+
transforms = {}
|
| 558 |
+
|
| 559 |
+
for joint_name, target_angle in target_joint_angles.items():
|
| 560 |
+
if joint_name.startswith('_'):
|
| 561 |
+
continue # skip translation keys
|
| 562 |
+
if joint_name not in initial_joints:
|
| 563 |
+
continue
|
| 564 |
+
|
| 565 |
+
# treat target_angle as absolute rotation from neutral
|
| 566 |
+
R = compute_rotation_matrix(target_angle)
|
| 567 |
+
pivot = np.array(initial_joints[joint_name])
|
| 568 |
+
|
| 569 |
+
transforms[joint_name] = {
|
| 570 |
+
'R': R,
|
| 571 |
+
'pivot': pivot
|
| 572 |
+
}
|
| 573 |
+
|
| 574 |
+
return transforms
|
| 575 |
+
|
| 576 |
+
|
| 577 |
+
def apply_skinning(
|
| 578 |
+
control_points_for_object,
|
| 579 |
+
point_to_joint,
|
| 580 |
+
transforms,
|
| 581 |
+
translation,
|
| 582 |
+
initial_positions
|
| 583 |
+
):
|
| 584 |
+
"""
|
| 585 |
+
Apply joint transforms to control points
|
| 586 |
+
Propagates through skeleton hierarchy
|
| 587 |
+
|
| 588 |
+
Args:
|
| 589 |
+
control_points_for_object: list of control point dicts
|
| 590 |
+
point_to_joint: {point_id: joint_name}
|
| 591 |
+
transforms: {joint_name: {R, pivot}}
|
| 592 |
+
translation: (dx, dy) global translation
|
| 593 |
+
initial_positions: {point_id: (x, y)} at frame 0
|
| 594 |
+
|
| 595 |
+
Returns:
|
| 596 |
+
{point_id: (new_x, new_y)}
|
| 597 |
+
"""
|
| 598 |
+
new_positions = {}
|
| 599 |
+
|
| 600 |
+
for pt in control_points_for_object:
|
| 601 |
+
pt_id = pt['id']
|
| 602 |
+
joint = point_to_joint.get(pt_id)
|
| 603 |
+
|
| 604 |
+
if joint is None or pt_id not in initial_positions:
|
| 605 |
+
# no assignment — apply translation only
|
| 606 |
+
init_pos = np.array([pt['x'], pt['y']])
|
| 607 |
+
new_positions[pt_id] = tuple(init_pos + np.array(translation))
|
| 608 |
+
continue
|
| 609 |
+
|
| 610 |
+
pt_pos = np.array(initial_positions[pt_id])
|
| 611 |
+
|
| 612 |
+
# walk up skeleton hierarchy applying transforms
|
| 613 |
+
current_joint = joint
|
| 614 |
+
accumulated_pos = pt_pos.copy()
|
| 615 |
+
visited = set()
|
| 616 |
+
|
| 617 |
+
while current_joint in transforms:
|
| 618 |
+
if current_joint in visited:
|
| 619 |
+
break
|
| 620 |
+
visited.add(current_joint)
|
| 621 |
+
|
| 622 |
+
T = transforms[current_joint]
|
| 623 |
+
accumulated_pos = (
|
| 624 |
+
T['R'] @ (accumulated_pos - T['pivot']) + T['pivot']
|
| 625 |
+
)
|
| 626 |
+
|
| 627 |
+
parent = SKELETON_HIERARCHY.get(current_joint)
|
| 628 |
+
if parent is None:
|
| 629 |
+
break
|
| 630 |
+
current_joint = parent
|
| 631 |
+
|
| 632 |
+
# apply global translation
|
| 633 |
+
accumulated_pos += np.array(translation)
|
| 634 |
+
new_positions[pt_id] = tuple(accumulated_pos)
|
| 635 |
+
|
| 636 |
+
return new_positions
|
| 637 |
+
|
| 638 |
+
|
| 639 |
+
# --- 4d: Full Keyframe Position Computation ---
|
| 640 |
+
|
| 641 |
+
def compute_keyframe_positions(
|
| 642 |
+
object_assignments,
|
| 643 |
+
point_to_joint_map,
|
| 644 |
+
initial_joints,
|
| 645 |
+
semantic_plan,
|
| 646 |
+
bounding_boxes
|
| 647 |
+
):
|
| 648 |
+
"""
|
| 649 |
+
For each keyframe, compute new control point positions
|
| 650 |
+
for all objects based on semantic plan
|
| 651 |
+
|
| 652 |
+
Returns:
|
| 653 |
+
{keyframe_idx: {point_id: (x, y)}}
|
| 654 |
+
"""
|
| 655 |
+
# store initial positions
|
| 656 |
+
all_points = []
|
| 657 |
+
for pts in object_assignments.values():
|
| 658 |
+
all_points.extend(pts)
|
| 659 |
+
|
| 660 |
+
initial_positions = {
|
| 661 |
+
pt['id']: (pt['x'], pt['y'])
|
| 662 |
+
for pt in all_points
|
| 663 |
+
}
|
| 664 |
+
|
| 665 |
+
keyframe_indices = sorted(
|
| 666 |
+
int(k) for k in list(semantic_plan.values())[0].keys()
|
| 667 |
+
)
|
| 668 |
+
|
| 669 |
+
keyframe_positions = {}
|
| 670 |
+
|
| 671 |
+
for kf_idx in keyframe_indices:
|
| 672 |
+
frame_positions = dict(initial_positions) # start from initial
|
| 673 |
+
|
| 674 |
+
for obj_name, obj_points in object_assignments.items():
|
| 675 |
+
if obj_name == 'unassigned' or not obj_points:
|
| 676 |
+
continue
|
| 677 |
+
|
| 678 |
+
obj_plan = semantic_plan.get(obj_name, {})
|
| 679 |
+
description = obj_plan.get(str(kf_idx), "stationary")
|
| 680 |
+
|
| 681 |
+
joint_angles = parse_semantic_to_joint_angles(description)
|
| 682 |
+
|
| 683 |
+
# extract translations
|
| 684 |
+
dx = joint_angles.pop('_translate_x', 0)
|
| 685 |
+
dy = joint_angles.pop('_translate_y', 0)
|
| 686 |
+
translation = (dx, dy)
|
| 687 |
+
|
| 688 |
+
if obj_name == 'player' and initial_joints:
|
| 689 |
+
transforms = compute_joint_transforms(
|
| 690 |
+
initial_joints,
|
| 691 |
+
joint_angles
|
| 692 |
+
)
|
| 693 |
+
pt_to_joint = point_to_joint_map.get(obj_name, {})
|
| 694 |
+
new_pos = apply_skinning(
|
| 695 |
+
obj_points,
|
| 696 |
+
pt_to_joint,
|
| 697 |
+
transforms,
|
| 698 |
+
translation,
|
| 699 |
+
initial_positions
|
| 700 |
+
)
|
| 701 |
+
else:
|
| 702 |
+
# rigid objects: bounding box center translation
|
| 703 |
+
bb = bounding_boxes.get(obj_name, [0, 0, 50, 50])
|
| 704 |
+
new_pos = {}
|
| 705 |
+
for pt in obj_points:
|
| 706 |
+
init = np.array(initial_positions[pt['id']])
|
| 707 |
+
new_pos[pt['id']] = tuple(init + np.array(translation))
|
| 708 |
+
|
| 709 |
+
frame_positions.update(new_pos)
|
| 710 |
+
|
| 711 |
+
keyframe_positions[kf_idx] = frame_positions
|
| 712 |
+
|
| 713 |
+
return keyframe_positions, initial_positions
|
| 714 |
+
|
| 715 |
+
|
| 716 |
+
# --- 4e: Contact Constraints ---
|
| 717 |
+
|
| 718 |
+
def enforce_contact_constraints(keyframe_positions, constraints):
|
| 719 |
+
"""
|
| 720 |
+
Force contact between two points at specified frames
|
| 721 |
+
|
| 722 |
+
constraints: list of dicts:
|
| 723 |
+
[
|
| 724 |
+
{
|
| 725 |
+
"frame": 2,
|
| 726 |
+
"anchor_point_id": 205, # basketball center
|
| 727 |
+
"target_point_id": 117, # player right wrist
|
| 728 |
+
"type": "contact"
|
| 729 |
+
}
|
| 730 |
+
]
|
| 731 |
+
"""
|
| 732 |
+
for c in constraints:
|
| 733 |
+
frame = c['frame']
|
| 734 |
+
if frame not in keyframe_positions:
|
| 735 |
+
continue
|
| 736 |
+
|
| 737 |
+
anchor_id = c['anchor_point_id']
|
| 738 |
+
target_id = c['target_point_id']
|
| 739 |
+
|
| 740 |
+
if target_id in keyframe_positions[frame]:
|
| 741 |
+
ref_pos = keyframe_positions[frame][target_id]
|
| 742 |
+
keyframe_positions[frame][anchor_id] = ref_pos
|
| 743 |
+
|
| 744 |
+
return keyframe_positions
|
| 745 |
+
|
| 746 |
+
|
| 747 |
+
# --- 4f: Cubic Spline Interpolation ---
|
| 748 |
+
|
| 749 |
+
def interpolate_keyframes(keyframe_positions, n_frames=16):
|
| 750 |
+
"""
|
| 751 |
+
Interpolate control point positions between keyframes
|
| 752 |
+
using cubic spline
|
| 753 |
+
|
| 754 |
+
Returns:
|
| 755 |
+
{frame_idx: {point_id: (x, y)}} for all n_frames
|
| 756 |
+
"""
|
| 757 |
+
kf_indices = sorted(keyframe_positions.keys())
|
| 758 |
+
|
| 759 |
+
if len(kf_indices) < 2:
|
| 760 |
+
raise ValueError("Need at least 2 keyframes to interpolate")
|
| 761 |
+
|
| 762 |
+
# scale keyframe indices to cover full frame range
|
| 763 |
+
kf_scaled = [
|
| 764 |
+
int(i * (n_frames - 1) / (len(kf_indices) - 1))
|
| 765 |
+
for i in range(len(kf_indices))
|
| 766 |
+
]
|
| 767 |
+
|
| 768 |
+
point_ids = list(keyframe_positions[kf_indices[0]].keys())
|
| 769 |
+
|
| 770 |
+
all_frames = {f: {} for f in range(n_frames)}
|
| 771 |
+
|
| 772 |
+
for pt_id in point_ids:
|
| 773 |
+
xs = [keyframe_positions[kf][pt_id][0] for kf in kf_indices]
|
| 774 |
+
ys = [keyframe_positions[kf][pt_id][1] for kf in kf_indices]
|
| 775 |
+
|
| 776 |
+
if len(set(xs)) == 1 and len(set(ys)) == 1:
|
| 777 |
+
# stationary point — skip interpolation
|
| 778 |
+
for f in range(n_frames):
|
| 779 |
+
all_frames[f][pt_id] = (xs[0], ys[0])
|
| 780 |
+
continue
|
| 781 |
+
|
| 782 |
+
try:
|
| 783 |
+
cs_x = CubicSpline(kf_scaled, xs, bc_type='not-a-knot')
|
| 784 |
+
cs_y = CubicSpline(kf_scaled, ys, bc_type='not-a-knot')
|
| 785 |
+
|
| 786 |
+
for f in range(n_frames):
|
| 787 |
+
all_frames[f][pt_id] = (
|
| 788 |
+
float(cs_x(f)),
|
| 789 |
+
float(cs_y(f))
|
| 790 |
+
)
|
| 791 |
+
except Exception:
|
| 792 |
+
# fall back to linear
|
| 793 |
+
for f in range(n_frames):
|
| 794 |
+
all_frames[f][pt_id] = (
|
| 795 |
+
float(np.interp(f, kf_scaled, xs)),
|
| 796 |
+
float(np.interp(f, kf_scaled, ys))
|
| 797 |
+
)
|
| 798 |
+
|
| 799 |
+
return all_frames
|
| 800 |
+
|
| 801 |
+
|
| 802 |
+
# =============================================================================
|
| 803 |
+
# STAGE 5 — RASTERIZE FRAME SEQUENCE
|
| 804 |
+
# =============================================================================
|
| 805 |
+
|
| 806 |
+
def update_svg_control_points(svg_path, frame_positions):
|
| 807 |
+
"""
|
| 808 |
+
Update SVG path data with new control point positions
|
| 809 |
+
Returns modified SVG as bytes
|
| 810 |
+
"""
|
| 811 |
+
with open(svg_path, 'r') as f:
|
| 812 |
+
svg_content = f.read()
|
| 813 |
+
|
| 814 |
+
# NOTE: full SVG path rewriting requires svgpathtools
|
| 815 |
+
# This is a simplified version — replace with full
|
| 816 |
+
# path reconstruction for production use
|
| 817 |
+
|
| 818 |
+
tree = ET.parse(svg_path)
|
| 819 |
+
ET.register_namespace('', 'http://www.w3.org/2000/svg')
|
| 820 |
+
|
| 821 |
+
svg_bytes = ET.tostring(
|
| 822 |
+
tree.getroot(),
|
| 823 |
+
encoding='unicode'
|
| 824 |
+
).encode('utf-8')
|
| 825 |
+
|
| 826 |
+
return svg_bytes
|
| 827 |
+
|
| 828 |
+
|
| 829 |
+
def rasterize_frame(svg_path, frame_positions, width=512, height=512):
|
| 830 |
+
"""
|
| 831 |
+
Rasterize one frame with updated control point positions
|
| 832 |
+
Returns PIL Image
|
| 833 |
+
"""
|
| 834 |
+
svg_bytes = update_svg_control_points(svg_path, frame_positions)
|
| 835 |
+
|
| 836 |
+
png_data = cairosvg.svg2png(
|
| 837 |
+
bytestring=svg_bytes,
|
| 838 |
+
output_width=width,
|
| 839 |
+
output_height=height
|
| 840 |
+
)
|
| 841 |
+
|
| 842 |
+
return Image.open(io.BytesIO(png_data)).convert("RGB")
|
| 843 |
+
|
| 844 |
+
|
| 845 |
+
def generate_frame_sequence(svg_path, all_frame_positions, output_dir, width=512, height=512):
|
| 846 |
+
"""
|
| 847 |
+
Generate and save all rasterized frames
|
| 848 |
+
|
| 849 |
+
Returns:
|
| 850 |
+
list of saved frame paths in order
|
| 851 |
+
"""
|
| 852 |
+
os.makedirs(output_dir, exist_ok=True)
|
| 853 |
+
frame_paths = []
|
| 854 |
+
|
| 855 |
+
for frame_idx in sorted(all_frame_positions.keys()):
|
| 856 |
+
positions = all_frame_positions[frame_idx]
|
| 857 |
+
frame_img = rasterize_frame(svg_path, positions, width, height)
|
| 858 |
+
|
| 859 |
+
frame_path = os.path.join(output_dir, f"frame_{frame_idx:04d}.png")
|
| 860 |
+
frame_img.save(frame_path)
|
| 861 |
+
frame_paths.append(frame_path)
|
| 862 |
+
|
| 863 |
+
if frame_idx % 4 == 0:
|
| 864 |
+
print(f" Rasterized frame {frame_idx}/{len(all_frame_positions)}")
|
| 865 |
+
|
| 866 |
+
return frame_paths
|
| 867 |
+
|
| 868 |
+
|
| 869 |
+
# =============================================================================
|
| 870 |
+
# VISUALISATION UTILITIES
|
| 871 |
+
# =============================================================================
|
| 872 |
+
|
| 873 |
+
def visualise_skeleton(image_array, joints, save_path=None):
|
| 874 |
+
"""Draw detected skeleton joints on image"""
|
| 875 |
+
img = Image.fromarray(image_array).convert("RGB")
|
| 876 |
+
draw = ImageDraw.Draw(img)
|
| 877 |
+
|
| 878 |
+
colors = {
|
| 879 |
+
'arm': 'red',
|
| 880 |
+
'leg': 'blue',
|
| 881 |
+
'torso': 'green',
|
| 882 |
+
'head': 'yellow'
|
| 883 |
+
}
|
| 884 |
+
|
| 885 |
+
joint_color_map = {
|
| 886 |
+
'nose': 'yellow', 'neck': 'green', 'spine': 'green',
|
| 887 |
+
'right_shoulder': 'red', 'left_shoulder': 'red',
|
| 888 |
+
'right_elbow': 'red', 'left_elbow': 'red',
|
| 889 |
+
'right_wrist': 'red', 'left_wrist': 'red',
|
| 890 |
+
'right_hip': 'blue', 'left_hip': 'blue',
|
| 891 |
+
'right_knee': 'blue', 'left_knee': 'blue',
|
| 892 |
+
'right_ankle': 'blue', 'left_ankle': 'blue',
|
| 893 |
+
}
|
| 894 |
+
|
| 895 |
+
for joint_name, (x, y) in joints.items():
|
| 896 |
+
color = joint_color_map.get(joint_name, 'white')
|
| 897 |
+
r = 5
|
| 898 |
+
draw.ellipse([x - r, y - r, x + r, y + r], fill=color, outline='black')
|
| 899 |
+
draw.text((x + 6, y - 6), joint_name.split('_')[-1][:3], fill=color)
|
| 900 |
+
|
| 901 |
+
# draw skeleton connections
|
| 902 |
+
connections = [
|
| 903 |
+
('nose', 'neck'), ('neck', 'right_shoulder'), ('neck', 'left_shoulder'),
|
| 904 |
+
('right_shoulder', 'right_elbow'), ('right_elbow', 'right_wrist'),
|
| 905 |
+
('left_shoulder', 'left_elbow'), ('left_elbow', 'left_wrist'),
|
| 906 |
+
('neck', 'spine'), ('spine', 'right_hip'), ('spine', 'left_hip'),
|
| 907 |
+
('right_hip', 'right_knee'), ('right_knee', 'right_ankle'),
|
| 908 |
+
('left_hip', 'left_knee'), ('left_knee', 'left_ankle'),
|
| 909 |
+
]
|
| 910 |
+
|
| 911 |
+
for j1, j2 in connections:
|
| 912 |
+
if j1 in joints and j2 in joints:
|
| 913 |
+
draw.line([joints[j1], joints[j2]], fill='white', width=2)
|
| 914 |
+
|
| 915 |
+
if save_path:
|
| 916 |
+
img.save(save_path)
|
| 917 |
+
print(f"Skeleton visualisation saved: {save_path}")
|
| 918 |
+
|
| 919 |
+
return img
|
| 920 |
+
|
| 921 |
+
|
| 922 |
+
def visualise_point_assignments(image_array, control_points, object_assignments, save_path=None):
|
| 923 |
+
"""Colour control points by object assignment"""
|
| 924 |
+
img = Image.fromarray(image_array).convert("RGB")
|
| 925 |
+
draw = ImageDraw.Draw(img)
|
| 926 |
+
|
| 927 |
+
object_colors = {
|
| 928 |
+
'player': 'red',
|
| 929 |
+
'basketball': 'orange',
|
| 930 |
+
'hoop': 'cyan',
|
| 931 |
+
'unassigned': 'gray'
|
| 932 |
+
}
|
| 933 |
+
|
| 934 |
+
for obj_name, points in object_assignments.items():
|
| 935 |
+
color = object_colors.get(obj_name, 'white')
|
| 936 |
+
for pt in points:
|
| 937 |
+
x, y = pt['x'], pt['y']
|
| 938 |
+
r = 3
|
| 939 |
+
draw.ellipse([x - r, y - r, x + r, y + r], fill=color)
|
| 940 |
+
|
| 941 |
+
if save_path:
|
| 942 |
+
img.save(save_path)
|
| 943 |
+
print(f"Point assignment visualisation saved: {save_path}")
|
| 944 |
+
|
| 945 |
+
return img
|
| 946 |
+
|
| 947 |
+
|
| 948 |
+
# =============================================================================
|
| 949 |
+
# FULL PIPELINE
|
| 950 |
+
# =============================================================================
|
| 951 |
+
|
| 952 |
+
def run_pipeline(
|
| 953 |
+
svg_path,
|
| 954 |
+
text_instruction,
|
| 955 |
+
output_dir="output_frames",
|
| 956 |
+
n_keyframes=5,
|
| 957 |
+
n_frames=16,
|
| 958 |
+
contact_constraints=None,
|
| 959 |
+
client=None
|
| 960 |
+
):
|
| 961 |
+
"""
|
| 962 |
+
Run full geometric solver pipeline
|
| 963 |
+
|
| 964 |
+
Args:
|
| 965 |
+
svg_path: path to input SVG file
|
| 966 |
+
text_instruction: motion description string
|
| 967 |
+
output_dir: where to save rasterized frames
|
| 968 |
+
n_keyframes: number of motion keyframes
|
| 969 |
+
n_frames: total output frames
|
| 970 |
+
contact_constraints: list of contact constraint dicts
|
| 971 |
+
client: LLM API client
|
| 972 |
+
|
| 973 |
+
Returns:
|
| 974 |
+
list of frame image paths
|
| 975 |
+
"""
|
| 976 |
+
print("=" * 60)
|
| 977 |
+
print("GEOMETRIC SOLVER PIPELINE")
|
| 978 |
+
print("=" * 60)
|
| 979 |
+
|
| 980 |
+
os.makedirs(output_dir, exist_ok=True)
|
| 981 |
+
vis_dir = os.path.join(output_dir, "visualisations")
|
| 982 |
+
os.makedirs(vis_dir, exist_ok=True)
|
| 983 |
+
|
| 984 |
+
# ----- Stage 1: Keyframe Prompt Decomposition -----
|
| 985 |
+
print("\n[Stage 1] Decomposing keyframe prompts...")
|
| 986 |
+
raster_path = os.path.join(output_dir, "input_raster.png")
|
| 987 |
+
image_array = rasterize_svg(svg_path)
|
| 988 |
+
Image.fromarray(image_array).save(raster_path)
|
| 989 |
+
|
| 990 |
+
keyframe_prompts = decompose_keyframe_prompts(
|
| 991 |
+
raster_path, text_instruction, n_keyframes, client
|
| 992 |
+
)
|
| 993 |
+
print(f" Got {len(keyframe_prompts)} keyframe prompts")
|
| 994 |
+
for i, kp in enumerate(keyframe_prompts):
|
| 995 |
+
print(f" kf{i}: {kp}")
|
| 996 |
+
|
| 997 |
+
# ----- Stage 2: Object Segmentation -----
|
| 998 |
+
print("\n[Stage 2] Segmenting objects...")
|
| 999 |
+
object_names_from_instruction = ["player", "basketball", "hoop"]
|
| 1000 |
+
bounding_boxes = get_object_bounding_boxes(image_array, object_names_from_instruction)
|
| 1001 |
+
print(f" Bounding boxes: {bounding_boxes}")
|
| 1002 |
+
|
| 1003 |
+
control_points = parse_svg_control_points(svg_path)
|
| 1004 |
+
print(f" Parsed {len(control_points)} control points from SVG")
|
| 1005 |
+
|
| 1006 |
+
object_assignments = assign_control_points_to_objects(
|
| 1007 |
+
control_points, bounding_boxes
|
| 1008 |
+
)
|
| 1009 |
+
for obj, pts in object_assignments.items():
|
| 1010 |
+
print(f" {obj}: {len(pts)} control points assigned")
|
| 1011 |
+
|
| 1012 |
+
vis_path = os.path.join(vis_dir, "point_assignments.png")
|
| 1013 |
+
visualise_point_assignments(image_array, control_points, object_assignments, vis_path)
|
| 1014 |
+
|
| 1015 |
+
# ----- Stage 3: Semantic Motion Plan -----
|
| 1016 |
+
print("\n[Stage 3] Getting semantic motion plan...")
|
| 1017 |
+
semantic_plan = get_semantic_motion_plan(
|
| 1018 |
+
raster_path, bounding_boxes, keyframe_prompts, client
|
| 1019 |
+
)
|
| 1020 |
+
print(" Semantic plan:")
|
| 1021 |
+
for obj, plan in semantic_plan.items():
|
| 1022 |
+
print(f" [{obj}]")
|
| 1023 |
+
for kf, desc in plan.items():
|
| 1024 |
+
print(f" kf{kf}: {desc}")
|
| 1025 |
+
|
| 1026 |
+
# ----- Stage 4: Geometric Solver -----
|
| 1027 |
+
print("\n[Stage 4] Running geometric solver...")
|
| 1028 |
+
|
| 1029 |
+
# 4a: skeleton extraction
|
| 1030 |
+
print(" Extracting skeleton via DWPose...")
|
| 1031 |
+
initial_joints = extract_skeleton_dwpose(image_array)
|
| 1032 |
+
if initial_joints:
|
| 1033 |
+
print(f" Detected {len(initial_joints)} joints")
|
| 1034 |
+
vis_path = os.path.join(vis_dir, "skeleton.png")
|
| 1035 |
+
visualise_skeleton(image_array, initial_joints, vis_path)
|
| 1036 |
+
else:
|
| 1037 |
+
print(" WARNING: No skeleton detected. Falling back to BB translation.")
|
| 1038 |
+
|
| 1039 |
+
# 4b: assign control points to joints
|
| 1040 |
+
point_to_joint_map = {}
|
| 1041 |
+
if initial_joints:
|
| 1042 |
+
player_points = object_assignments.get('player', [])
|
| 1043 |
+
point_to_joint_map['player'] = assign_control_points_to_joints(
|
| 1044 |
+
player_points, initial_joints
|
| 1045 |
+
)
|
| 1046 |
+
print(f" Assigned {len(point_to_joint_map['player'])} player points to joints")
|
| 1047 |
+
|
| 1048 |
+
# 4c: compute keyframe positions
|
| 1049 |
+
print(" Computing keyframe positions...")
|
| 1050 |
+
keyframe_positions, initial_positions = compute_keyframe_positions(
|
| 1051 |
+
object_assignments,
|
| 1052 |
+
point_to_joint_map,
|
| 1053 |
+
initial_joints,
|
| 1054 |
+
semantic_plan,
|
| 1055 |
+
bounding_boxes
|
| 1056 |
+
)
|
| 1057 |
+
print(f" Computed positions for {len(keyframe_positions)} keyframes")
|
| 1058 |
+
|
| 1059 |
+
# 4d: enforce contact constraints
|
| 1060 |
+
if contact_constraints:
|
| 1061 |
+
print(f" Enforcing {len(contact_constraints)} contact constraints...")
|
| 1062 |
+
keyframe_positions = enforce_contact_constraints(
|
| 1063 |
+
keyframe_positions, contact_constraints
|
| 1064 |
+
)
|
| 1065 |
+
|
| 1066 |
+
# 4e: interpolate between keyframes
|
| 1067 |
+
print(" Interpolating between keyframes...")
|
| 1068 |
+
all_frame_positions = interpolate_keyframes(keyframe_positions, n_frames)
|
| 1069 |
+
print(f" Generated positions for {len(all_frame_positions)} frames")
|
| 1070 |
+
|
| 1071 |
+
# ----- Stage 5: Rasterize Frames -----
|
| 1072 |
+
print("\n[Stage 5] Rasterizing frame sequence...")
|
| 1073 |
+
frame_paths = generate_frame_sequence(
|
| 1074 |
+
svg_path, all_frame_positions,
|
| 1075 |
+
os.path.join(output_dir, "frames")
|
| 1076 |
+
)
|
| 1077 |
+
print(f" Saved {len(frame_paths)} frames to {output_dir}/frames/")
|
| 1078 |
+
|
| 1079 |
+
print("\n" + "=" * 60)
|
| 1080 |
+
print("PIPELINE COMPLETE")
|
| 1081 |
+
print(f"Frames saved to: {output_dir}/frames/")
|
| 1082 |
+
print(f"Visualisations: {output_dir}/visualisations/")
|
| 1083 |
+
print("Next step: feed adjacent frame pairs to Wan2.2")
|
| 1084 |
+
print("=" * 60)
|
| 1085 |
+
|
| 1086 |
+
return frame_paths
|
| 1087 |
+
|
| 1088 |
+
|
| 1089 |
+
# =============================================================================
|
| 1090 |
+
# ENTRY POINT — BASKETBALL EXAMPLE
|
| 1091 |
+
# =============================================================================
|
| 1092 |
+
|
| 1093 |
+
if __name__ == "__main__":
|
| 1094 |
+
|
| 1095 |
+
# contact constraints for basketball example:
|
| 1096 |
+
# at keyframe 2 (ball releasing), basketball center
|
| 1097 |
+
# should be at player's right wrist position
|
| 1098 |
+
# NOTE: replace point IDs with actual IDs from your SVG
|
| 1099 |
+
contact_constraints = [
|
| 1100 |
+
{
|
| 1101 |
+
"frame": 2,
|
| 1102 |
+
"anchor_point_id": 200, # basketball center point
|
| 1103 |
+
"target_point_id": 117, # player right wrist point
|
| 1104 |
+
"type": "contact"
|
| 1105 |
+
}
|
| 1106 |
+
]
|
| 1107 |
+
|
| 1108 |
+
# --- set your SVG path here ---
|
| 1109 |
+
SVG_PATH = "basketball_sketch.svg"
|
| 1110 |
+
|
| 1111 |
+
if not os.path.exists(SVG_PATH):
|
| 1112 |
+
print(f"SVG not found at {SVG_PATH}")
|
| 1113 |
+
print("Upload your basketball sketch SVG and set SVG_PATH")
|
| 1114 |
+
print("\nTo test without SVG, running skeleton mock only...\n")
|
| 1115 |
+
|
| 1116 |
+
# test skeleton + semantic parsing without SVG
|
| 1117 |
+
mock_image = np.ones((512, 512, 3), dtype=np.uint8) * 240
|
| 1118 |
+
joints = extract_skeleton_dwpose(mock_image)
|
| 1119 |
+
|
| 1120 |
+
print("Mock skeleton joints:")
|
| 1121 |
+
for j, pos in joints.items():
|
| 1122 |
+
print(f" {j}: {pos}")
|
| 1123 |
+
|
| 1124 |
+
desc = "crouching, arm extended upward, jumping"
|
| 1125 |
+
angles = parse_semantic_to_joint_angles(desc)
|
| 1126 |
+
print(f"\nSemantic parse of '{desc}':")
|
| 1127 |
+
print(json.dumps(angles, indent=2))
|
| 1128 |
+
|
| 1129 |
+
else:
|
| 1130 |
+
frame_paths = run_pipeline(
|
| 1131 |
+
svg_path=SVG_PATH,
|
| 1132 |
+
text_instruction=(
|
| 1133 |
+
"A basketball player takes a jump shot, "
|
| 1134 |
+
"aiming for the hoop, with the basketball "
|
| 1135 |
+
"mid-air and heading towards the hoop."
|
| 1136 |
+
),
|
| 1137 |
+
output_dir="basketball_output",
|
| 1138 |
+
n_keyframes=5,
|
| 1139 |
+
n_frames=16,
|
| 1140 |
+
contact_constraints=contact_constraints
|
| 1141 |
+
)
|
| 1142 |
+
|
| 1143 |
+
print(f"\nGenerated {len(frame_paths)} frames.")
|
| 1144 |
+
print("Feed to Wan2.2 using adjacent pairs:")
|
| 1145 |
+
print(" (frame_0000.png, frame_0001.png) -> clip_0")
|
| 1146 |
+
print(" (frame_0001.png, frame_0002.png) -> clip_1")
|
| 1147 |
+
print(" ...")
|
Downloads/handover_summary(1).md
ADDED
|
@@ -0,0 +1,220 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MoSketch Pipeline Handover Summary
|
| 2 |
+
**Date:** June 2026
|
| 3 |
+
**Project:** Multi-object sketch animation pipeline replacing MoSketch's SDS optimisation
|
| 4 |
+
|
| 5 |
+
---
|
| 6 |
+
|
| 7 |
+
## Research Contribution
|
| 8 |
+
Replace MoSketch's SDS test-time optimisation (~1hr/clip) with:
|
| 9 |
+
1. Feedforward LLM-based motion planning (Qwen2.5-7B)
|
| 10 |
+
2. Geometric solver (rigid translation of SVG control points)
|
| 11 |
+
3. Wan2.2 video synthesis (not yet integrated — VRAM constraint)
|
| 12 |
+
|
| 13 |
+
Uses MoSketch's pre-computed stroke assignments for fair comparison.
|
| 14 |
+
Cite as: "identical segmentation to MoSketch for fair evaluation."
|
| 15 |
+
|
| 16 |
+
---
|
| 17 |
+
|
| 18 |
+
## Server Environment
|
| 19 |
+
- Server: `otter34`, user `rk01499`
|
| 20 |
+
- Conda env: `/scratch/rk01499/anaconda3/envs/mosketch/` (Python 3.8)
|
| 21 |
+
- ALWAYS use full path: `/scratch/rk01499/anaconda3/envs/mosketch/bin/python`
|
| 22 |
+
- GPU: RTX A4000 (16GB VRAM) — only ~2GB free when Qwen loaded
|
| 23 |
+
- `python` alias breaks between sessions — always use full path
|
| 24 |
+
|
| 25 |
+
---
|
| 26 |
+
|
| 27 |
+
## Key File Locations
|
| 28 |
+
|
| 29 |
+
### Pipeline Files
|
| 30 |
+
| File | Location | Status |
|
| 31 |
+
|------|----------|--------|
|
| 32 |
+
| `mosketch_pipeline_v3.py` | `/scratch/rk01499/MoSketch_svg/` | Working baseline |
|
| 33 |
+
| `mosketch_pipeline_v5.py` | `/scratch/rk01499/MoSketch_svg/` | v3 + intent planner (current) |
|
| 34 |
+
| `stage3_intent_planner.py` | `/scratch/rk01499/MoSketch_svg/` | Intent-based Stage 3 module |
|
| 35 |
+
|
| 36 |
+
### Data
|
| 37 |
+
| Resource | Location |
|
| 38 |
+
|----------|----------|
|
| 39 |
+
| SVG files | `/user/HS400/rk01499/my_scratch/MoSketch/data/raw/60sketches/svg/` |
|
| 40 |
+
| PNG files | `/user/HS400/rk01499/my_scratch/MoSketch/data/raw/60sketches/png/` |
|
| 41 |
+
| Captions | `/user/HS400/rk01499/my_scratch/MoSketch/data/raw/60sketches/caption.txt` |
|
| 42 |
+
| Semantic assignments | `/user/HS400/rk01499/my_scratch/MoSketch/data/processed/{name}/{name}_semantic.txt` |
|
| 43 |
+
| v3 outputs | `/user/HS400/rk01499/my_scratch/MoSketch/output_geometric_v3/` |
|
| 44 |
+
| v5 outputs | `/user/HS400/rk01499/my_scratch/MoSketch/output_geometric_v5/` |
|
| 45 |
+
|
| 46 |
+
### Models
|
| 47 |
+
| Model | Location |
|
| 48 |
+
|-------|----------|
|
| 49 |
+
| Qwen2.5-7B | `/user/HS400/rk01499/my_scratch/models/qwen2.5-7b/` |
|
| 50 |
+
| Qwen2.5-14B | `/scratch/rk01499/models/qwen2.5-14b/` |
|
| 51 |
+
| Grounding DINO | `/user/HS400/rk01499/my_scratch/models/groundingdino_swint_ogc.pth` |
|
| 52 |
+
| ModelScope T2V 1.7B | `/scratch/rk01499/MoSketch/text-to-video-ms-1.7b/` |
|
| 53 |
+
| Wan2.2 | NOT downloaded — needs 28GB, 482GB free on /scratch |
|
| 54 |
+
|
| 55 |
+
---
|
| 56 |
+
|
| 57 |
+
## Pipeline Architecture (v5 — current)
|
| 58 |
+
|
| 59 |
+
```
|
| 60 |
+
SVG + caption
|
| 61 |
+
→ Stage 1: Qwen keyframe decomposition (5 descriptions)
|
| 62 |
+
→ Stage 2: Read _semantic.txt stroke assignments
|
| 63 |
+
→ Stage 2b: Parse SVG paths, assign to objects, compute bounding boxes
|
| 64 |
+
→ Stage 3: Intent planner (stage3_intent_planner.py)
|
| 65 |
+
Qwen outputs: endpoint, path, group, contact_kf, group_after_kf
|
| 66 |
+
Geometry computes: actual pixel trajectories
|
| 67 |
+
→ Stage 4: Geometric solver (rigid translation of control points)
|
| 68 |
+
→ Stage 5A: Pipeline A — 5 sparse keyframes
|
| 69 |
+
Stage 5B: Pipeline B — 16 cubic-spline-interpolated frames
|
| 70 |
+
```
|
| 71 |
+
|
| 72 |
+
---
|
| 73 |
+
|
| 74 |
+
## Stage 3 Intent Planner (stage3_intent_planner.py)
|
| 75 |
+
|
| 76 |
+
### Intent Schema
|
| 77 |
+
```json
|
| 78 |
+
{
|
| 79 |
+
"object_name": {
|
| 80 |
+
"endpoint": "other_object | left_edge | right_edge | top | bottom | stationary",
|
| 81 |
+
"path": "straight | arc_up | arc_down | follow | circular | downward",
|
| 82 |
+
"group": "other_object | independent",
|
| 83 |
+
"contact_kf": 3,
|
| 84 |
+
"group_after_kf": "other_object | null"
|
| 85 |
+
}
|
| 86 |
+
}
|
| 87 |
+
```
|
| 88 |
+
|
| 89 |
+
### Path Types
|
| 90 |
+
- `straight` — direct line from start to end
|
| 91 |
+
- `arc_up` — parabolic rise then fall (projectiles, throws, jumps)
|
| 92 |
+
- `arc_down` — dips then recovers (rollercoaster)
|
| 93 |
+
- `follow` — copies another object's displacement (smoke trails shell)
|
| 94 |
+
- `circular` — orbits a center point (satellite)
|
| 95 |
+
- `downward` — falls straight down (liquid)
|
| 96 |
+
|
| 97 |
+
### Key Features
|
| 98 |
+
- **Two-pass trajectory computation**: independent objects first, followers second
|
| 99 |
+
- **Post-contact grouping**: frisbee joins dog's trajectory after contact_kf
|
| 100 |
+
- **Dynamic endpoint**: frisbee targets dog's position at contact_kf, not initial position
|
| 101 |
+
- **Upper quarter contact**: endpoint uses top 25% of target bbox (mouth not belly)
|
| 102 |
+
|
| 103 |
+
### Unit Tests
|
| 104 |
+
```bash
|
| 105 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python stage3_intent_planner.py
|
| 106 |
+
```
|
| 107 |
+
Both cannon1 and dog3 tests pass. Frisbee follows dog at kf4: PASS
|
| 108 |
+
|
| 109 |
+
---
|
| 110 |
+
|
| 111 |
+
## Known Issues and Fixes Applied
|
| 112 |
+
|
| 113 |
+
### Fixed
|
| 114 |
+
- cairosvg 2.7.1 black image bug → `preprocess_svg()` adds white bg + converts rgb() to hex
|
| 115 |
+
- Qwen JSON parsing failures → `parse_qwen_json()` extracts JSON between first { and last }
|
| 116 |
+
- Shell jumping to cannon at kf0 → locked start (kf0 = actual centroid)
|
| 117 |
+
- Smoke moving toward cannon → intent planner `group=shell` (smoke follows shell)
|
| 118 |
+
- Shell falling instead of flying → motion hints added to prompt
|
| 119 |
+
- All paths assigned to one object → using MoSketch _semantic.txt instead of DINO
|
| 120 |
+
|
| 121 |
+
### Known Limitations (document in paper)
|
| 122 |
+
- No video generation yet (Wan2.2 needs 28GB, 16GB VRAM with Qwen loaded)
|
| 123 |
+
- Curved road following not implemented (carfp15, carfp36, carfp38, carfp48, carside13)
|
| 124 |
+
- Non-rigid deformation not implemented (Wan2.2 handles this)
|
| 125 |
+
- Contact events approximate (frisbee reaches dog bbox centroid area)
|
| 126 |
+
- Dog not moving in dog3 (Qwen still plans it stationary despite hints — last known issue)
|
| 127 |
+
- Uses MoSketch pre-computed segmentation (own DINO+SAM deferred to future work)
|
| 128 |
+
|
| 129 |
+
---
|
| 130 |
+
|
| 131 |
+
## Run Commands
|
| 132 |
+
|
| 133 |
+
```bash
|
| 134 |
+
# single sketch
|
| 135 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python mosketch_pipeline_v5.py --sketch cannon1
|
| 136 |
+
|
| 137 |
+
# all 60 sketches
|
| 138 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python mosketch_pipeline_v5.py --all
|
| 139 |
+
|
| 140 |
+
# pipeline A only (faster)
|
| 141 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python mosketch_pipeline_v5.py --all --pipeline A
|
| 142 |
+
|
| 143 |
+
# unit tests for intent planner
|
| 144 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python stage3_intent_planner.py
|
| 145 |
+
```
|
| 146 |
+
|
| 147 |
+
---
|
| 148 |
+
|
| 149 |
+
## Immediate Next Steps
|
| 150 |
+
|
| 151 |
+
1. **Fix dog not moving** — dog3 intent has `endpoint=stationary` despite caption saying "sprints forward". Add stronger hint to prompt or post-process to detect stationary animals described as moving.
|
| 152 |
+
|
| 153 |
+
2. **Download Wan2.2** for video synthesis:
|
| 154 |
+
```bash
|
| 155 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python -c "
|
| 156 |
+
from huggingface_hub import snapshot_download
|
| 157 |
+
snapshot_download(
|
| 158 |
+
repo_id='Wan-AI/Wan2.1-FLF2V-14B-720P-diffusers',
|
| 159 |
+
local_dir='/scratch/rk01499/models/wan-flf2v',
|
| 160 |
+
ignore_patterns=['*.md', '*.txt']
|
| 161 |
+
)
|
| 162 |
+
"
|
| 163 |
+
```
|
| 164 |
+
Need to kill Qwen process first to free VRAM. Run keyframe generation and video synthesis as separate scripts.
|
| 165 |
+
|
| 166 |
+
3. **Run all 60 with v5** and compare against v3 outputs to measure improvement.
|
| 167 |
+
|
| 168 |
+
4. **Modularise code** into:
|
| 169 |
+
- `config.py` — paths and constants
|
| 170 |
+
- `models.py` — Qwen loading
|
| 171 |
+
- `data.py` — caption loader, semantic reader
|
| 172 |
+
- `svg_parser.py` — SVG parsing, path assignment
|
| 173 |
+
- `planner.py` — wraps stage3_intent_planner
|
| 174 |
+
- `solver.py` — geometric solver
|
| 175 |
+
- `renderer.py` — rasterization, pipeline A/B
|
| 176 |
+
- `pipeline.py` — main entry point
|
| 177 |
+
|
| 178 |
+
5. **Implement own DINO+SAM segmentation** for unseen sketch generalisation (deferred — use MoSketch files for now).
|
| 179 |
+
|
| 180 |
+
---
|
| 181 |
+
|
| 182 |
+
## Semantic File Format
|
| 183 |
+
```
|
| 184 |
+
object_name<TAB>stroke_idx1,stroke_idx2,...
|
| 185 |
+
```
|
| 186 |
+
Example (carfp36):
|
| 187 |
+
```
|
| 188 |
+
road 1,2,44,45,...,79
|
| 189 |
+
jeep 0,3,4,...,43
|
| 190 |
+
motorcycle 80,81,...,168
|
| 191 |
+
```
|
| 192 |
+
Background objects (road, ground, sky) never move.
|
| 193 |
+
|
| 194 |
+
---
|
| 195 |
+
|
| 196 |
+
## cairosvg Fix (CRITICAL — without this all frames are black)
|
| 197 |
+
```python
|
| 198 |
+
def preprocess_svg(svg_path=None, svg_content=None):
|
| 199 |
+
if svg_content is None:
|
| 200 |
+
with open(svg_path, 'r') as f:
|
| 201 |
+
svg_content = f.read()
|
| 202 |
+
svg_content = svg_content.replace('<g>', '<g><rect width="256" height="256" fill="white"/>', 1)
|
| 203 |
+
svg_content = re.sub(r'stroke="rgb\(0,\s*0,\s*0\)"', 'stroke="#000000"', svg_content)
|
| 204 |
+
return svg_content.encode('utf-8')
|
| 205 |
+
```
|
| 206 |
+
|
| 207 |
+
---
|
| 208 |
+
|
| 209 |
+
## Motion Types Across 60 Sketches
|
| 210 |
+
| Type | Count | Path | Examples |
|
| 211 |
+
|------|-------|------|---------|
|
| 212 |
+
| Projectile arc | 12 | arc_up | cannon, frisbee, basketball, dolphin |
|
| 213 |
+
| Horizontal approach | 14 | straight | two cars, predator+prey, cat+mouse |
|
| 214 |
+
| Vertical motion | 8 | straight up/down | airplane, shuttle, rappel, ladder |
|
| 215 |
+
| Follow/trail | 7 | follow | smoke+cannon, carriage+horse |
|
| 216 |
+
| Curved road | 5 | arc_road (not impl) | carfp15/36/38/48, carside13 |
|
| 217 |
+
| Stationary | 8 | stationary | eating, grazing, couple |
|
| 218 |
+
| Rotation/orbit | 3 | circular | satellite, rollercoaster |
|
| 219 |
+
| Pour/flow | 3 | downward | bottle, ice splash |
|
| 220 |
+
|
Downloads/handover_summary.md
ADDED
|
@@ -0,0 +1,220 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MoSketch Pipeline Handover Summary
|
| 2 |
+
**Date:** June 2026
|
| 3 |
+
**Project:** Multi-object sketch animation pipeline replacing MoSketch's SDS optimisation
|
| 4 |
+
|
| 5 |
+
---
|
| 6 |
+
|
| 7 |
+
## Research Contribution
|
| 8 |
+
Replace MoSketch's SDS test-time optimisation (~1hr/clip) with:
|
| 9 |
+
1. Feedforward LLM-based motion planning (Qwen2.5-7B)
|
| 10 |
+
2. Geometric solver (rigid translation of SVG control points)
|
| 11 |
+
3. Wan2.2 video synthesis (not yet integrated — VRAM constraint)
|
| 12 |
+
|
| 13 |
+
Uses MoSketch's pre-computed stroke assignments for fair comparison.
|
| 14 |
+
Cite as: "identical segmentation to MoSketch for fair evaluation."
|
| 15 |
+
|
| 16 |
+
---
|
| 17 |
+
|
| 18 |
+
## Server Environment
|
| 19 |
+
- Server: `otter34`, user `rk01499`
|
| 20 |
+
- Conda env: `/scratch/rk01499/anaconda3/envs/mosketch/` (Python 3.8)
|
| 21 |
+
- ALWAYS use full path: `/scratch/rk01499/anaconda3/envs/mosketch/bin/python`
|
| 22 |
+
- GPU: RTX A4000 (16GB VRAM) — only ~2GB free when Qwen loaded
|
| 23 |
+
- `python` alias breaks between sessions — always use full path
|
| 24 |
+
|
| 25 |
+
---
|
| 26 |
+
|
| 27 |
+
## Key File Locations
|
| 28 |
+
|
| 29 |
+
### Pipeline Files
|
| 30 |
+
| File | Location | Status |
|
| 31 |
+
|------|----------|--------|
|
| 32 |
+
| `mosketch_pipeline_v3.py` | `/scratch/rk01499/MoSketch_svg/` | Working baseline |
|
| 33 |
+
| `mosketch_pipeline_v5.py` | `/scratch/rk01499/MoSketch_svg/` | v3 + intent planner (current) |
|
| 34 |
+
| `stage3_intent_planner.py` | `/scratch/rk01499/MoSketch_svg/` | Intent-based Stage 3 module |
|
| 35 |
+
|
| 36 |
+
### Data
|
| 37 |
+
| Resource | Location |
|
| 38 |
+
|----------|----------|
|
| 39 |
+
| SVG files | `/user/HS400/rk01499/my_scratch/MoSketch/data/raw/60sketches/svg/` |
|
| 40 |
+
| PNG files | `/user/HS400/rk01499/my_scratch/MoSketch/data/raw/60sketches/png/` |
|
| 41 |
+
| Captions | `/user/HS400/rk01499/my_scratch/MoSketch/data/raw/60sketches/caption.txt` |
|
| 42 |
+
| Semantic assignments | `/user/HS400/rk01499/my_scratch/MoSketch/data/processed/{name}/{name}_semantic.txt` |
|
| 43 |
+
| v3 outputs | `/user/HS400/rk01499/my_scratch/MoSketch/output_geometric_v3/` |
|
| 44 |
+
| v5 outputs | `/user/HS400/rk01499/my_scratch/MoSketch/output_geometric_v5/` |
|
| 45 |
+
|
| 46 |
+
### Models
|
| 47 |
+
| Model | Location |
|
| 48 |
+
|-------|----------|
|
| 49 |
+
| Qwen2.5-7B | `/user/HS400/rk01499/my_scratch/models/qwen2.5-7b/` |
|
| 50 |
+
| Qwen2.5-14B | `/scratch/rk01499/models/qwen2.5-14b/` |
|
| 51 |
+
| Grounding DINO | `/user/HS400/rk01499/my_scratch/models/groundingdino_swint_ogc.pth` |
|
| 52 |
+
| ModelScope T2V 1.7B | `/scratch/rk01499/MoSketch/text-to-video-ms-1.7b/` |
|
| 53 |
+
| Wan2.2 | NOT downloaded — needs 28GB, 482GB free on /scratch |
|
| 54 |
+
|
| 55 |
+
---
|
| 56 |
+
|
| 57 |
+
## Pipeline Architecture (v5 — current)
|
| 58 |
+
|
| 59 |
+
```
|
| 60 |
+
SVG + caption
|
| 61 |
+
→ Stage 1: Qwen keyframe decomposition (5 descriptions)
|
| 62 |
+
→ Stage 2: Read _semantic.txt stroke assignments
|
| 63 |
+
→ Stage 2b: Parse SVG paths, assign to objects, compute bounding boxes
|
| 64 |
+
→ Stage 3: Intent planner (stage3_intent_planner.py)
|
| 65 |
+
Qwen outputs: endpoint, path, group, contact_kf, group_after_kf
|
| 66 |
+
Geometry computes: actual pixel trajectories
|
| 67 |
+
→ Stage 4: Geometric solver (rigid translation of control points)
|
| 68 |
+
→ Stage 5A: Pipeline A — 5 sparse keyframes
|
| 69 |
+
Stage 5B: Pipeline B — 16 cubic-spline-interpolated frames
|
| 70 |
+
```
|
| 71 |
+
|
| 72 |
+
---
|
| 73 |
+
|
| 74 |
+
## Stage 3 Intent Planner (stage3_intent_planner.py)
|
| 75 |
+
|
| 76 |
+
### Intent Schema
|
| 77 |
+
```json
|
| 78 |
+
{
|
| 79 |
+
"object_name": {
|
| 80 |
+
"endpoint": "other_object | left_edge | right_edge | top | bottom | stationary",
|
| 81 |
+
"path": "straight | arc_up | arc_down | follow | circular | downward",
|
| 82 |
+
"group": "other_object | independent",
|
| 83 |
+
"contact_kf": 3,
|
| 84 |
+
"group_after_kf": "other_object | null"
|
| 85 |
+
}
|
| 86 |
+
}
|
| 87 |
+
```
|
| 88 |
+
|
| 89 |
+
### Path Types
|
| 90 |
+
- `straight` — direct line from start to end
|
| 91 |
+
- `arc_up` — parabolic rise then fall (projectiles, throws, jumps)
|
| 92 |
+
- `arc_down` — dips then recovers (rollercoaster)
|
| 93 |
+
- `follow` — copies another object's displacement (smoke trails shell)
|
| 94 |
+
- `circular` — orbits a center point (satellite)
|
| 95 |
+
- `downward` — falls straight down (liquid)
|
| 96 |
+
|
| 97 |
+
### Key Features
|
| 98 |
+
- **Two-pass trajectory computation**: independent objects first, followers second
|
| 99 |
+
- **Post-contact grouping**: frisbee joins dog's trajectory after contact_kf
|
| 100 |
+
- **Dynamic endpoint**: frisbee targets dog's position at contact_kf, not initial position
|
| 101 |
+
- **Upper quarter contact**: endpoint uses top 25% of target bbox (mouth not belly)
|
| 102 |
+
|
| 103 |
+
### Unit Tests
|
| 104 |
+
```bash
|
| 105 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python stage3_intent_planner.py
|
| 106 |
+
```
|
| 107 |
+
Both cannon1 and dog3 tests pass. Frisbee follows dog at kf4: PASS
|
| 108 |
+
|
| 109 |
+
---
|
| 110 |
+
|
| 111 |
+
## Known Issues and Fixes Applied
|
| 112 |
+
|
| 113 |
+
### Fixed
|
| 114 |
+
- cairosvg 2.7.1 black image bug → `preprocess_svg()` adds white bg + converts rgb() to hex
|
| 115 |
+
- Qwen JSON parsing failures → `parse_qwen_json()` extracts JSON between first { and last }
|
| 116 |
+
- Shell jumping to cannon at kf0 → locked start (kf0 = actual centroid)
|
| 117 |
+
- Smoke moving toward cannon → intent planner `group=shell` (smoke follows shell)
|
| 118 |
+
- Shell falling instead of flying → motion hints added to prompt
|
| 119 |
+
- All paths assigned to one object → using MoSketch _semantic.txt instead of DINO
|
| 120 |
+
|
| 121 |
+
### Known Limitations (document in paper)
|
| 122 |
+
- No video generation yet (Wan2.2 needs 28GB, 16GB VRAM with Qwen loaded)
|
| 123 |
+
- Curved road following not implemented (carfp15, carfp36, carfp38, carfp48, carside13)
|
| 124 |
+
- Non-rigid deformation not implemented (Wan2.2 handles this)
|
| 125 |
+
- Contact events approximate (frisbee reaches dog bbox centroid area)
|
| 126 |
+
- Dog not moving in dog3 (Qwen still plans it stationary despite hints — last known issue)
|
| 127 |
+
- Uses MoSketch pre-computed segmentation (own DINO+SAM deferred to future work)
|
| 128 |
+
|
| 129 |
+
---
|
| 130 |
+
|
| 131 |
+
## Run Commands
|
| 132 |
+
|
| 133 |
+
```bash
|
| 134 |
+
# single sketch
|
| 135 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python mosketch_pipeline_v5.py --sketch cannon1
|
| 136 |
+
|
| 137 |
+
# all 60 sketches
|
| 138 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python mosketch_pipeline_v5.py --all
|
| 139 |
+
|
| 140 |
+
# pipeline A only (faster)
|
| 141 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python mosketch_pipeline_v5.py --all --pipeline A
|
| 142 |
+
|
| 143 |
+
# unit tests for intent planner
|
| 144 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python stage3_intent_planner.py
|
| 145 |
+
```
|
| 146 |
+
|
| 147 |
+
---
|
| 148 |
+
|
| 149 |
+
## Immediate Next Steps
|
| 150 |
+
|
| 151 |
+
1. **Fix dog not moving** — dog3 intent has `endpoint=stationary` despite caption saying "sprints forward". Add stronger hint to prompt or post-process to detect stationary animals described as moving.
|
| 152 |
+
|
| 153 |
+
2. **Download Wan2.2** for video synthesis:
|
| 154 |
+
```bash
|
| 155 |
+
/scratch/rk01499/anaconda3/envs/mosketch/bin/python -c "
|
| 156 |
+
from huggingface_hub import snapshot_download
|
| 157 |
+
snapshot_download(
|
| 158 |
+
repo_id='Wan-AI/Wan2.1-FLF2V-14B-720P-diffusers',
|
| 159 |
+
local_dir='/scratch/rk01499/models/wan-flf2v',
|
| 160 |
+
ignore_patterns=['*.md', '*.txt']
|
| 161 |
+
)
|
| 162 |
+
"
|
| 163 |
+
```
|
| 164 |
+
Need to kill Qwen process first to free VRAM. Run keyframe generation and video synthesis as separate scripts.
|
| 165 |
+
|
| 166 |
+
3. **Run all 60 with v5** and compare against v3 outputs to measure improvement.
|
| 167 |
+
|
| 168 |
+
4. **Modularise code** into:
|
| 169 |
+
- `config.py` — paths and constants
|
| 170 |
+
- `models.py` — Qwen loading
|
| 171 |
+
- `data.py` — caption loader, semantic reader
|
| 172 |
+
- `svg_parser.py` — SVG parsing, path assignment
|
| 173 |
+
- `planner.py` — wraps stage3_intent_planner
|
| 174 |
+
- `solver.py` — geometric solver
|
| 175 |
+
- `renderer.py` — rasterization, pipeline A/B
|
| 176 |
+
- `pipeline.py` — main entry point
|
| 177 |
+
|
| 178 |
+
5. **Implement own DINO+SAM segmentation** for unseen sketch generalisation (deferred — use MoSketch files for now).
|
| 179 |
+
|
| 180 |
+
---
|
| 181 |
+
|
| 182 |
+
## Semantic File Format
|
| 183 |
+
```
|
| 184 |
+
object_name<TAB>stroke_idx1,stroke_idx2,...
|
| 185 |
+
```
|
| 186 |
+
Example (carfp36):
|
| 187 |
+
```
|
| 188 |
+
road 1,2,44,45,...,79
|
| 189 |
+
jeep 0,3,4,...,43
|
| 190 |
+
motorcycle 80,81,...,168
|
| 191 |
+
```
|
| 192 |
+
Background objects (road, ground, sky) never move.
|
| 193 |
+
|
| 194 |
+
---
|
| 195 |
+
|
| 196 |
+
## cairosvg Fix (CRITICAL — without this all frames are black)
|
| 197 |
+
```python
|
| 198 |
+
def preprocess_svg(svg_path=None, svg_content=None):
|
| 199 |
+
if svg_content is None:
|
| 200 |
+
with open(svg_path, 'r') as f:
|
| 201 |
+
svg_content = f.read()
|
| 202 |
+
svg_content = svg_content.replace('<g>', '<g><rect width="256" height="256" fill="white"/>', 1)
|
| 203 |
+
svg_content = re.sub(r'stroke="rgb\(0,\s*0,\s*0\)"', 'stroke="#000000"', svg_content)
|
| 204 |
+
return svg_content.encode('utf-8')
|
| 205 |
+
```
|
| 206 |
+
|
| 207 |
+
---
|
| 208 |
+
|
| 209 |
+
## Motion Types Across 60 Sketches
|
| 210 |
+
| Type | Count | Path | Examples |
|
| 211 |
+
|------|-------|------|---------|
|
| 212 |
+
| Projectile arc | 12 | arc_up | cannon, frisbee, basketball, dolphin |
|
| 213 |
+
| Horizontal approach | 14 | straight | two cars, predator+prey, cat+mouse |
|
| 214 |
+
| Vertical motion | 8 | straight up/down | airplane, shuttle, rappel, ladder |
|
| 215 |
+
| Follow/trail | 7 | follow | smoke+cannon, carriage+horse |
|
| 216 |
+
| Curved road | 5 | arc_road (not impl) | carfp15/36/38/48, carside13 |
|
| 217 |
+
| Stationary | 8 | stationary | eating, grazing, couple |
|
| 218 |
+
| Rotation/orbit | 3 | circular | satellite, rollercoaster |
|
| 219 |
+
| Pour/flow | 3 | downward | bottle, ice splash |
|
| 220 |
+
|
Downloads/lbs_seam_constrained(1).py
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Reference: point-level Linear Blend Skinning with HARD seam constraints.
|
| 3 |
+
|
| 4 |
+
Use this to diff against your existing test_lbs.py. The two things to check
|
| 5 |
+
in your own code:
|
| 6 |
+
|
| 7 |
+
1. Are skinning weights computed PER CONTROL POINT, not per stroke?
|
| 8 |
+
(per-stroke weighting still tears at joints, just less visibly)
|
| 9 |
+
2. Do stroke endpoints that coincide in the rest pose get IDENTICAL
|
| 10 |
+
rest-pose coordinates AND identical weight vectors, forced
|
| 11 |
+
explicitly — not just "close because the kernel is smooth"?
|
| 12 |
+
Matching weights alone is NOT sufficient; the rest positions must
|
| 13 |
+
also be snapped together, or a nonlinear per-joint transform
|
| 14 |
+
(rotation) still maps the two near-identical points to different
|
| 15 |
+
outputs.
|
| 16 |
+
|
| 17 |
+
Smooth inverse-distance weighting alone reduces tearing but does not
|
| 18 |
+
guarantee zero gap at a seam. Hard-constraining seam pairs to share both
|
| 19 |
+
a rest position and a weight vector guarantees it by construction,
|
| 20 |
+
regardless of whether the per-joint transform is rotation, scale, or
|
| 21 |
+
both.
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
import numpy as np
|
| 25 |
+
from scipy.spatial import cKDTree
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
# ---------------------------------------------------------------------------
|
| 29 |
+
# 1. Skinning weights: per POINT, inverse-distance to each joint
|
| 30 |
+
# ---------------------------------------------------------------------------
|
| 31 |
+
|
| 32 |
+
def compute_skinning_weights(points, joints, power=2.0, eps=1e-6):
|
| 33 |
+
"""
|
| 34 |
+
points: (N, 2) array of ALL control points across ALL strokes, flattened.
|
| 35 |
+
Do NOT compute this per-stroke and re-run per stroke — every
|
| 36 |
+
point in the whole sketch must be weighted against every joint
|
| 37 |
+
in one pass so shared/seam points are handled consistently.
|
| 38 |
+
joints: (J, 2) array of joint centers.
|
| 39 |
+
|
| 40 |
+
Returns:
|
| 41 |
+
weights: (N, J) array, each row sums to 1.
|
| 42 |
+
"""
|
| 43 |
+
points = np.asarray(points, dtype=float)
|
| 44 |
+
joints = np.asarray(joints, dtype=float)
|
| 45 |
+
|
| 46 |
+
# (N, J) distance matrix
|
| 47 |
+
diff = points[:, None, :] - joints[None, :, :]
|
| 48 |
+
dist = np.linalg.norm(diff, axis=2)
|
| 49 |
+
dist = np.maximum(dist, eps)
|
| 50 |
+
|
| 51 |
+
inv = 1.0 / (dist ** power)
|
| 52 |
+
weights = inv / inv.sum(axis=1, keepdims=True)
|
| 53 |
+
return weights
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
# ---------------------------------------------------------------------------
|
| 57 |
+
# 2. Seam detection: find stroke endpoints that coincide in rest pose
|
| 58 |
+
# ---------------------------------------------------------------------------
|
| 59 |
+
|
| 60 |
+
def find_seam_groups(stroke_endpoints, tol=2.0):
|
| 61 |
+
"""
|
| 62 |
+
stroke_endpoints: list of (point_index, xy) for every stroke START/END
|
| 63 |
+
point (the points most likely to be shared joints
|
| 64 |
+
between adjacent strokes, e.g. leg-to-torso).
|
| 65 |
+
tol: pixel distance below which two endpoints are considered "the same
|
| 66 |
+
point" and must move together.
|
| 67 |
+
|
| 68 |
+
Returns:
|
| 69 |
+
list of lists, each inner list = point indices that must share
|
| 70 |
+
one weight vector (a seam group).
|
| 71 |
+
"""
|
| 72 |
+
idxs = np.array([i for i, _ in stroke_endpoints])
|
| 73 |
+
coords = np.array([xy for _, xy in stroke_endpoints])
|
| 74 |
+
|
| 75 |
+
tree = cKDTree(coords)
|
| 76 |
+
pairs = tree.query_pairs(r=tol)
|
| 77 |
+
|
| 78 |
+
# union-find to merge transitive seam groups (A-B, B-C => A-B-C)
|
| 79 |
+
parent = {i: i for i in idxs}
|
| 80 |
+
|
| 81 |
+
def find(x):
|
| 82 |
+
while parent[x] != x:
|
| 83 |
+
parent[x] = parent[parent[x]]
|
| 84 |
+
x = parent[x]
|
| 85 |
+
return x
|
| 86 |
+
|
| 87 |
+
def union(a, b):
|
| 88 |
+
ra, rb = find(a), find(b)
|
| 89 |
+
if ra != rb:
|
| 90 |
+
parent[ra] = rb
|
| 91 |
+
|
| 92 |
+
for a, b in pairs:
|
| 93 |
+
union(idxs[a], idxs[b])
|
| 94 |
+
|
| 95 |
+
groups = {}
|
| 96 |
+
for i in idxs:
|
| 97 |
+
root = find(i)
|
| 98 |
+
groups.setdefault(root, []).append(i)
|
| 99 |
+
|
| 100 |
+
return [g for g in groups.values() if len(g) > 1]
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
# ---------------------------------------------------------------------------
|
| 104 |
+
# 3. Force seam groups to share identical weight vectors
|
| 105 |
+
# ---------------------------------------------------------------------------
|
| 106 |
+
|
| 107 |
+
def enforce_seam_constraints(points, weights, seam_groups):
|
| 108 |
+
"""
|
| 109 |
+
points: (N, 2) rest-pose points, modified in place (and returned) so
|
| 110 |
+
every point in a seam group is snapped to the same rest
|
| 111 |
+
position — the average of the group.
|
| 112 |
+
weights: (N, J), modified in place (and returned) so every point in a
|
| 113 |
+
seam group gets the SAME weight row.
|
| 114 |
+
|
| 115 |
+
Both fixes are required. Matching weights alone is not enough: if the
|
| 116 |
+
rest-pose coordinates still differ by even a fraction of a pixel, a
|
| 117 |
+
nonlinear per-joint transform (rotation) maps them to different
|
| 118 |
+
outputs even under identical weights. You need points AND weights to
|
| 119 |
+
agree at a seam, or the "fix" only shrinks the gap instead of
|
| 120 |
+
eliminating it.
|
| 121 |
+
"""
|
| 122 |
+
points = points.copy()
|
| 123 |
+
weights = weights.copy()
|
| 124 |
+
for group in seam_groups:
|
| 125 |
+
avg_point = points[group].mean(axis=0)
|
| 126 |
+
points[group] = avg_point
|
| 127 |
+
|
| 128 |
+
avg_w = weights[group].mean(axis=0)
|
| 129 |
+
avg_w = avg_w / avg_w.sum()
|
| 130 |
+
weights[group] = avg_w
|
| 131 |
+
return points, weights
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
# ---------------------------------------------------------------------------
|
| 135 |
+
# 4. Apply per-joint transforms via the (now seam-safe) weights
|
| 136 |
+
# ---------------------------------------------------------------------------
|
| 137 |
+
|
| 138 |
+
def apply_lbs(points, weights, joint_transforms):
|
| 139 |
+
"""
|
| 140 |
+
points: (N, 2) rest-pose points.
|
| 141 |
+
weights: (N, J) from enforce_seam_constraints.
|
| 142 |
+
joint_transforms: list of J functions, each mapping a point (2,) to
|
| 143 |
+
its transformed position (2,) under that joint's
|
| 144 |
+
rotation/scale/whatever. E.g.:
|
| 145 |
+
|
| 146 |
+
def make_transform(center, angle_deg, scale=1.0):
|
| 147 |
+
theta = np.radians(angle_deg)
|
| 148 |
+
R = np.array([[np.cos(theta), -np.sin(theta)],
|
| 149 |
+
[np.sin(theta), np.cos(theta)]])
|
| 150 |
+
def f(p):
|
| 151 |
+
return center + scale * R @ (p - center)
|
| 152 |
+
return f
|
| 153 |
+
|
| 154 |
+
Returns:
|
| 155 |
+
deformed points, (N, 2).
|
| 156 |
+
"""
|
| 157 |
+
points = np.asarray(points, dtype=float)
|
| 158 |
+
N, J = weights.shape
|
| 159 |
+
out = np.zeros_like(points)
|
| 160 |
+
|
| 161 |
+
for j in range(J):
|
| 162 |
+
transformed_j = np.array([joint_transforms[j](p) for p in points])
|
| 163 |
+
out += weights[:, j:j + 1] * transformed_j
|
| 164 |
+
|
| 165 |
+
return out
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
# ---------------------------------------------------------------------------
|
| 169 |
+
# Example usage / sanity check
|
| 170 |
+
# ---------------------------------------------------------------------------
|
| 171 |
+
|
| 172 |
+
if __name__ == "__main__":
|
| 173 |
+
# Toy example: two strokes meeting at a joint. In real SVG data these
|
| 174 |
+
# "shared" endpoints are almost never bit-identical — they're drawn as
|
| 175 |
+
# two separate strokes with independently-rounded coordinates, e.g.
|
| 176 |
+
# torso ends near (2.0, 0.0) and the leg stroke starts near (2.004, -0.003).
|
| 177 |
+
# That tiny mismatch is enough for a smooth weight kernel to assign
|
| 178 |
+
# slightly different weights to each — and once joints rotate by
|
| 179 |
+
# different amounts, "slightly different weights" becomes a visible gap.
|
| 180 |
+
stroke_a = np.array([[0, 0], [1, 0], [2.0, 0.0]], dtype=float) # torso-ish
|
| 181 |
+
stroke_b = np.array([[2.004, -0.003], [2, -1], [2, -2]], dtype=float) # leg-ish
|
| 182 |
+
|
| 183 |
+
all_points = np.vstack([stroke_a, stroke_b])
|
| 184 |
+
# indices: 0,1,2 = stroke_a ; 3,4,5 = stroke_b ; point 2 and 3 are the "same" joint
|
| 185 |
+
|
| 186 |
+
joints = np.array([[0.5, 0], [2, -1]]) # torso joint, leg joint
|
| 187 |
+
|
| 188 |
+
# seam candidates: the shared endpoint appears twice (index 2 and 3),
|
| 189 |
+
# tol set generously since real sketch data has this kind of slop
|
| 190 |
+
endpoints = [(2, all_points[2]), (3, all_points[3])]
|
| 191 |
+
seams = find_seam_groups(endpoints, tol=0.5)
|
| 192 |
+
print("seam groups:", seams)
|
| 193 |
+
|
| 194 |
+
W = compute_skinning_weights(all_points, joints, power=2.0)
|
| 195 |
+
points_fixed, W_fixed = enforce_seam_constraints(all_points, W, seams)
|
| 196 |
+
|
| 197 |
+
print("naive weights at seam: ", W[2], W[3], " <- not identical")
|
| 198 |
+
print("fixed weights at seam: ", W_fixed[2], W_fixed[3], " <- forced identical")
|
| 199 |
+
print("naive rest points at seam:", all_points[2], all_points[3], " <- not identical")
|
| 200 |
+
print("fixed rest points at seam:", points_fixed[2], points_fixed[3], " <- snapped identical")
|
| 201 |
+
|
| 202 |
+
def make_transform(center, angle_deg, scale=1.0):
|
| 203 |
+
theta = np.radians(angle_deg)
|
| 204 |
+
R = np.array([[np.cos(theta), -np.sin(theta)],
|
| 205 |
+
[np.sin(theta), np.cos(theta)]])
|
| 206 |
+
def f(p):
|
| 207 |
+
return center + scale * (R @ (p - center))
|
| 208 |
+
return f
|
| 209 |
+
|
| 210 |
+
transforms = [
|
| 211 |
+
make_transform(joints[0], angle_deg=10, scale=1.0),
|
| 212 |
+
make_transform(joints[1], angle_deg=-30, scale=1.2), # different rotation AND scale
|
| 213 |
+
]
|
| 214 |
+
|
| 215 |
+
deformed_naive = apply_lbs(all_points, W, transforms)
|
| 216 |
+
deformed_fixed = apply_lbs(points_fixed, W_fixed, transforms)
|
| 217 |
+
|
| 218 |
+
gap_naive = np.linalg.norm(deformed_naive[2] - deformed_naive[3])
|
| 219 |
+
gap_fixed = np.linalg.norm(deformed_fixed[2] - deformed_fixed[3])
|
| 220 |
+
|
| 221 |
+
print(f"seam gap WITHOUT hard constraint: {gap_naive:.6f}")
|
| 222 |
+
print(f"seam gap WITH hard constraint: {gap_fixed:.6f}")
|
Downloads/pipe_3.py
ADDED
|
@@ -0,0 +1,1465 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
mosketch_pipeline.py — the full pipeline as one script with subcommands.
|
| 3 |
+
|
| 4 |
+
Pipeline:
|
| 5 |
+
1. Identify objects -> from the semantic file (no Qwen)
|
| 6 |
+
2. Classify ARAP vs. not -> `classify` subcommand (one Qwen call, all objects)
|
| 7 |
+
3. Narrate + deform -> `narrate` + `deform` subcommands, ONLY for
|
| 8 |
+
objects marked ARAP in step 2
|
| 9 |
+
4. Render -> `render` subcommand; every object gets real
|
| 10 |
+
trajectory translation; ARAP objects
|
| 11 |
+
additionally get deformation on top
|
| 12 |
+
|
| 13 |
+
Run steps individually for debugging, or use `full` to run everything for
|
| 14 |
+
one sketch in one process — this loads the Qwen model ONCE and reuses it
|
| 15 |
+
across classify/narrate/deform, instead of loading it 3 separate times.
|
| 16 |
+
|
| 17 |
+
Examples:
|
| 18 |
+
# step by step
|
| 19 |
+
python mosketch_pipeline.py classify --model M --caption-file C --sketch-name S --semantic SEM --out deformation.json
|
| 20 |
+
python mosketch_pipeline.py narrate --model M --caption-file C --sketch-name S --objects dog --out narratives.json
|
| 21 |
+
python mosketch_pipeline.py deform --model M --svg S.svg --semantic SEM --traj T --narratives narratives.json --deformation deformation.json --out-dir .
|
| 22 |
+
python mosketch_pipeline.py render --svg S.svg --semantic SEM --traj T --handles-dir .
|
| 23 |
+
|
| 24 |
+
# everything at once, one model load
|
| 25 |
+
python mosketch_pipeline.py full --model M --caption-file C --sketch-name S --svg S.svg --semantic SEM --traj T --out-dir .
|
| 26 |
+
"""
|
| 27 |
+
|
| 28 |
+
import argparse
|
| 29 |
+
import json
|
| 30 |
+
import os
|
| 31 |
+
import re
|
| 32 |
+
import sys
|
| 33 |
+
|
| 34 |
+
import numpy as np
|
| 35 |
+
|
| 36 |
+
from lib import (
|
| 37 |
+
load_strokes_from_svg, load_semantic_assignments, filter_strokes, flatten_strokes,
|
| 38 |
+
load_object, deduplicate_points, build_mesh, nearest_mesh_vertex,
|
| 39 |
+
auto_select_handles_deduped, object_bbox_size, arap_deform,
|
| 40 |
+
load_trajectories, bbox_deltas, get_caption, build_stroke_geometry_text,
|
| 41 |
+
)
|
| 42 |
+
|
| 43 |
+
N_KEYFRAMES = 5
|
| 44 |
+
MAX_RETRIES = 5
|
| 45 |
+
PLAUSIBILITY_THRESHOLD = 4
|
| 46 |
+
FAITHFULNESS_THRESHOLD = 4 # both must pass to stop — faithfulness previously only
|
| 47 |
+
# affected feedback text, never actually gated success
|
| 48 |
+
QUALITY_THRESHOLD = 4 # same upgrade applied to the new quality criterion —
|
| 49 |
+
# scored but not gating would repeat the same mistake
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def faithfulness_passed(score):
|
| 53 |
+
"""
|
| 54 |
+
faithfulness_score can be a number 1-5, the string "N/A" (no caption
|
| 55 |
+
was given to compare against, so there's nothing to fail), or missing
|
| 56 |
+
entirely (treated as NOT passed — can't confirm it's good, so err
|
| 57 |
+
toward regenerating the narrative rather than assuming it's fine).
|
| 58 |
+
"""
|
| 59 |
+
if score is None:
|
| 60 |
+
return False
|
| 61 |
+
if isinstance(score, str):
|
| 62 |
+
return score.strip().upper() == "N/A"
|
| 63 |
+
if isinstance(score, (int, float)):
|
| 64 |
+
return score >= FAITHFULNESS_THRESHOLD
|
| 65 |
+
return False
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def unload_model(model):
|
| 69 |
+
"""Frees GPU memory before loading a different model. Necessary because
|
| 70 |
+
the narrate/deform steps use a text Qwen model and the judge step uses
|
| 71 |
+
a separate vision-language model (Qwen3-VL) — on hardware with limited
|
| 72 |
+
VRAM (this project's RTX A4000, 16GB, already documented as a tight
|
| 73 |
+
fit for a single model), loading both at once risks the same OOM issue
|
| 74 |
+
that blocked Wan2.2 integration earlier. Load/unload sequentially
|
| 75 |
+
instead of assuming both fit simultaneously.
|
| 76 |
+
|
| 77 |
+
CONFIRMED BUG (found on real hardware, invisible to all mocked testing
|
| 78 |
+
since no real GPU was available to catch it): `del model` here only
|
| 79 |
+
clears THIS function's own local reference — it does nothing to the
|
| 80 |
+
caller's variable, which stays alive and keeps the whole model
|
| 81 |
+
resident in VRAM. torch.cuda.empty_cache() then has nothing to
|
| 82 |
+
actually free, because the refcount never reaches zero. Fixed by
|
| 83 |
+
returning None — callers MUST reassign their variable to this return
|
| 84 |
+
value (e.g. `model = unload_model(model)`), or the bug reappears."""
|
| 85 |
+
import gc
|
| 86 |
+
del model
|
| 87 |
+
gc.collect()
|
| 88 |
+
try:
|
| 89 |
+
import torch
|
| 90 |
+
if torch.cuda.is_available():
|
| 91 |
+
torch.cuda.empty_cache()
|
| 92 |
+
except ImportError:
|
| 93 |
+
pass
|
| 94 |
+
return None
|
| 95 |
+
|
| 96 |
+
# Dog3's real, human-reviewed narratives — used BOTH as the few-shot example
|
| 97 |
+
# in `narrate` and as the fallback default if --narratives is omitted in
|
| 98 |
+
# `deform`. One constant, one source of truth (previously duplicated across
|
| 99 |
+
# two separate files under two different names with identical content).
|
| 100 |
+
DOG3_CAPTION = ("The person throws a frisbee through the air, and the dog sits poised, "
|
| 101 |
+
"ready to sprint forward and catch it with its mouth in a swift motion.")
|
| 102 |
+
DOG3_NARRATIVES = {
|
| 103 |
+
"dog": [
|
| 104 |
+
"the dog is sitting alert, watching the frisbee as it is thrown",
|
| 105 |
+
"the dog is beginning to rise, weight shifting forward, head reaching toward the frisbee",
|
| 106 |
+
"the dog is mid-leap, body extended, reaching far forward and up toward the frisbee",
|
| 107 |
+
"the dog is at the peak of its jump, reaching as far as possible toward the frisbee",
|
| 108 |
+
"the dog is landing after catching the frisbee, body compacting back down",
|
| 109 |
+
],
|
| 110 |
+
"frisbee": [
|
| 111 |
+
"the frisbee has just left the thrower's hand, angled slightly upward",
|
| 112 |
+
"the frisbee is gliding through the air, tilting slightly as it arcs",
|
| 113 |
+
"the frisbee is near the peak of its arc, angled toward the dog",
|
| 114 |
+
"the frisbee is descending toward the dog, tilting down slightly",
|
| 115 |
+
"the frisbee is at the dog's mouth, being caught",
|
| 116 |
+
],
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
SVG_PATH_DEFAULT = "/mnt/user-data/uploads/dog3.svg"
|
| 120 |
+
SEMANTIC_PATH_DEFAULT = "dog3_semantic.txt"
|
| 121 |
+
DEFAULT_COLOR = "#444444"
|
| 122 |
+
DEFAULT_LINEWIDTH = 1.1
|
| 123 |
+
OBJECT_COLORS = {"dog": "black", "person": "#3F4C57", "frisbee": "#B0463C"}
|
| 124 |
+
OBJECT_LINEWIDTH = {"dog": 1.1, "person": 1.1, "frisbee": 1.4}
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
# =============================================================================
|
| 128 |
+
# shared: Qwen call + response parsing
|
| 129 |
+
# =============================================================================
|
| 130 |
+
|
| 131 |
+
def query_qwen(model, tokenizer, prompt, device, max_new_tokens=500, temperature=0.1):
|
| 132 |
+
import torch
|
| 133 |
+
messages = [{"role": "user", "content": prompt}]
|
| 134 |
+
text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
| 135 |
+
inputs = tokenizer([text], return_tensors="pt").to(device)
|
| 136 |
+
input_token_count = inputs["input_ids"].shape[1]
|
| 137 |
+
with torch.no_grad():
|
| 138 |
+
output_ids = model.generate(**inputs, max_new_tokens=max_new_tokens,
|
| 139 |
+
temperature=temperature, do_sample=True)
|
| 140 |
+
generated = output_ids[0][inputs["input_ids"].shape[1]:]
|
| 141 |
+
output_token_count = generated.shape[0]
|
| 142 |
+
response_text = tokenizer.decode(generated, skip_special_tokens=True)
|
| 143 |
+
return response_text, input_token_count, output_token_count
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def load_qwen_model(model_path):
|
| 147 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 148 |
+
import torch
|
| 149 |
+
print(f"loading model from {model_path} ...")
|
| 150 |
+
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
| 151 |
+
model = AutoModelForCausalLM.from_pretrained(model_path, torch_dtype=torch.float16, device_map="auto")
|
| 152 |
+
device = next(model.parameters()).device
|
| 153 |
+
print("model loaded.")
|
| 154 |
+
return model, tokenizer, device
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def parse_json_response(response_text):
|
| 158 |
+
match = re.search(r"\{.*\}", response_text, re.DOTALL)
|
| 159 |
+
if not match:
|
| 160 |
+
raise ValueError("No JSON object found in response:\n" + response_text)
|
| 161 |
+
return json.loads(match.group(0))
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
# =============================================================================
|
| 165 |
+
# STEP 2: classify — ARAP vs TRAJ_ONLY, one call, all objects
|
| 166 |
+
# =============================================================================
|
| 167 |
+
|
| 168 |
+
def build_deformation_prompt(caption, object_names):
|
| 169 |
+
objects_str = ", ".join(f'"{o}"' for o in object_names)
|
| 170 |
+
lines = "\n".join(
|
| 171 |
+
f' "{o}": "ARAP" or "TRAJ_ONLY"' + ("," if i < len(object_names) - 1 else "")
|
| 172 |
+
for i, o in enumerate(object_names)
|
| 173 |
+
)
|
| 174 |
+
return f"""Scene: "{caption}"
|
| 175 |
+
|
| 176 |
+
Objects in this scene: {objects_str}
|
| 177 |
+
|
| 178 |
+
For each object, decide whether representing it correctly needs NON-RIGID DEFORMATION (its body/shape changes — e.g. limbs moving, a neck reaching, a body crouching or leaning) or whether simple RIGID TRANSLATION (the object moves/rotates as a whole, unchanged in shape, or doesn't move at all) is enough.
|
| 179 |
+
|
| 180 |
+
Answer "ARAP" if the object's shape or body configuration changes at any point in the action, even if its overall position doesn't change. Answer "TRAJ_ONLY" if the object is rigid (a vehicle, tool, projectile, furniture, background element) or is simply carried by its own movement without changing shape.
|
| 181 |
+
|
| 182 |
+
Respond with ONLY a JSON object, no other text, in this exact format:
|
| 183 |
+
{{
|
| 184 |
+
{lines}
|
| 185 |
+
}}
|
| 186 |
+
"""
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
def validate_deformation(parsed, object_names):
|
| 190 |
+
problems = []
|
| 191 |
+
for obj in object_names:
|
| 192 |
+
if obj not in parsed:
|
| 193 |
+
problems.append(f"'{obj}' missing from response")
|
| 194 |
+
continue
|
| 195 |
+
val = str(parsed[obj]).strip().upper()
|
| 196 |
+
if val not in ("ARAP", "TRAJ_ONLY"):
|
| 197 |
+
problems.append(f"'{obj}' has invalid value {parsed[obj]!r}, expected ARAP or TRAJ_ONLY")
|
| 198 |
+
return problems
|
| 199 |
+
|
| 200 |
+
|
| 201 |
+
def run_classify(model, tokenizer, device, caption, semantic_path, out_path):
|
| 202 |
+
assignments = load_semantic_assignments(semantic_path)
|
| 203 |
+
object_names = list(assignments.keys())
|
| 204 |
+
print(f"objects found in {semantic_path}: {object_names}")
|
| 205 |
+
|
| 206 |
+
prompt = build_deformation_prompt(caption, object_names)
|
| 207 |
+
print("\n--- CLASSIFY PROMPT ---")
|
| 208 |
+
print(prompt)
|
| 209 |
+
|
| 210 |
+
response, in_tok, out_tok = query_qwen(model, tokenizer, prompt, device,
|
| 211 |
+
max_new_tokens=250, temperature=0.1)
|
| 212 |
+
print(f"\ntokens: {in_tok} in / {out_tok} out")
|
| 213 |
+
print("--- RAW RESPONSE ---")
|
| 214 |
+
print(response)
|
| 215 |
+
|
| 216 |
+
parsed = parse_json_response(response)
|
| 217 |
+
problems = validate_deformation(parsed, object_names)
|
| 218 |
+
print("\n--- PARSED ---")
|
| 219 |
+
print(json.dumps(parsed, indent=2))
|
| 220 |
+
if problems:
|
| 221 |
+
print("--- VALIDATION PROBLEMS ---")
|
| 222 |
+
for p in problems:
|
| 223 |
+
print(f" - {p}")
|
| 224 |
+
|
| 225 |
+
arap_objs = [o for o in object_names if str(parsed.get(o, "")).strip().upper() == "ARAP"]
|
| 226 |
+
traj_only_objs = [o for o in object_names if o not in arap_objs]
|
| 227 |
+
print(f"\nARAP: {arap_objs}")
|
| 228 |
+
print(f"TRAJ_ONLY: {traj_only_objs}")
|
| 229 |
+
|
| 230 |
+
with open(out_path, "w") as f:
|
| 231 |
+
json.dump(parsed, f, indent=2)
|
| 232 |
+
print(f"wrote {out_path}")
|
| 233 |
+
return parsed, arap_objs
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
def build_objects_info(svg_path, semantic_path, arap_objects):
|
| 237 |
+
"""
|
| 238 |
+
Standalone version of the mesh/joint setup previously embedded inside
|
| 239 |
+
run_deform — factored out so the unified narrate+deform+judge retry
|
| 240 |
+
loop can build this ONCE before the loop (mesh/joints never change
|
| 241 |
+
between attempts) instead of recomputing it every attempt.
|
| 242 |
+
"""
|
| 243 |
+
objects_info = {}
|
| 244 |
+
for obj_name in arap_objects:
|
| 245 |
+
points, slices = load_object(obj_name, svg_path, semantic_path)
|
| 246 |
+
bbox_size = object_bbox_size(points)
|
| 247 |
+
unique_points, p2u = deduplicate_points(points, tol=0.35)
|
| 248 |
+
tri, edges = build_mesh(unique_points)
|
| 249 |
+
strokes = filter_strokes(load_strokes_from_svg(svg_path),
|
| 250 |
+
load_semantic_assignments(semantic_path)[obj_name])
|
| 251 |
+
joints, anchor_idx, handle_idxs, joint_mesh_indices = auto_select_handles_deduped(
|
| 252 |
+
strokes, unique_points, k=4)
|
| 253 |
+
|
| 254 |
+
if len(handle_idxs) == 0:
|
| 255 |
+
print(f"WARNING: '{obj_name}' has no independent handles after dedup, skipping")
|
| 256 |
+
continue
|
| 257 |
+
|
| 258 |
+
print(f"'{obj_name}': {len(joints)} joints, anchor={anchor_idx}, handles={handle_idxs}, "
|
| 259 |
+
f"mesh_indices={joint_mesh_indices}, bbox_size={bbox_size:.1f}")
|
| 260 |
+
|
| 261 |
+
objects_info[obj_name] = {
|
| 262 |
+
"joints": joints, "anchor_idx": anchor_idx, "handle_idxs": handle_idxs,
|
| 263 |
+
"joint_mesh_indices": joint_mesh_indices, "bbox_size": bbox_size, "strokes": strokes,
|
| 264 |
+
"points": points, "slices": slices, "unique_points": unique_points,
|
| 265 |
+
"p2u": p2u, "edges": edges,
|
| 266 |
+
}
|
| 267 |
+
return objects_info
|
| 268 |
+
|
| 269 |
+
|
| 270 |
+
def build_deformed_stroke_geometry_text(object_info, kf_targets, n_points=2):
|
| 271 |
+
"""
|
| 272 |
+
Reconstructs what a SPECIFIC keyframe's ACTUAL DEFORMED shape looks
|
| 273 |
+
like, as compact text — runs the SAME ARAP solve used for real
|
| 274 |
+
rendering, so a stroke's coordinates here are that keyframe's true
|
| 275 |
+
rendered position, not its rest-pose position. This is what makes
|
| 276 |
+
the judge's target_coords anchoring correct PER KEYFRAME instead of
|
| 277 |
+
only correct when a keyframe happens to match rest — e.g. a "mouth"
|
| 278 |
+
stroke's rest coordinate is wrong context if the head itself has
|
| 279 |
+
moved by kf3, this reconstructs where that stroke actually is at kf3.
|
| 280 |
+
|
| 281 |
+
kf_targets: {joint_i_str: [x, y]} for ONE keyframe, in the same
|
| 282 |
+
"joint_i" key format apply_deform_clip_and_write saves (i.e. one
|
| 283 |
+
entry of deform_outputs[obj_name][kf_key]).
|
| 284 |
+
|
| 285 |
+
COST, not glossed over: re-runs arap_deform (10 iterations) once per
|
| 286 |
+
object per keyframe purely to build this text — measured ~2,866
|
| 287 |
+
tokens for a single 64-stroke object across all 5 keyframes, and that
|
| 288 |
+
multiplies per object in a multi-object scene. This is real compute
|
| 289 |
+
and real prompt-length cost on top of everything else already in the
|
| 290 |
+
judge prompt (images, rest geometry, joint legend, targets, bbox).
|
| 291 |
+
"""
|
| 292 |
+
import re as _re
|
| 293 |
+
unique_points = object_info["unique_points"]
|
| 294 |
+
edges = object_info["edges"]
|
| 295 |
+
p2u = object_info["p2u"]
|
| 296 |
+
slices = object_info["slices"]
|
| 297 |
+
anchor_idx = object_info["anchor_idx"]
|
| 298 |
+
joints = object_info["joints"]
|
| 299 |
+
joint_mesh_indices = object_info["joint_mesh_indices"]
|
| 300 |
+
|
| 301 |
+
handle_mesh_indices, handle_targets = [], []
|
| 302 |
+
for key, target in kf_targets.items():
|
| 303 |
+
m = _re.match(r"joint_(\d+)", key)
|
| 304 |
+
if not m:
|
| 305 |
+
continue
|
| 306 |
+
joint_i = int(m.group(1))
|
| 307 |
+
if joint_i >= len(joint_mesh_indices):
|
| 308 |
+
continue
|
| 309 |
+
handle_mesh_indices.append(joint_mesh_indices[joint_i])
|
| 310 |
+
handle_targets.append(target)
|
| 311 |
+
# anchor is always held at rest, same convention as rendering
|
| 312 |
+
handle_mesh_indices.append(joint_mesh_indices[anchor_idx])
|
| 313 |
+
handle_targets.append(joints[anchor_idx].tolist())
|
| 314 |
+
|
| 315 |
+
if not handle_mesh_indices:
|
| 316 |
+
return None
|
| 317 |
+
|
| 318 |
+
deformed_unique = arap_deform(unique_points, edges, handle_mesh_indices,
|
| 319 |
+
np.array(handle_targets), iterations=10)
|
| 320 |
+
deformed_points = deformed_unique[p2u]
|
| 321 |
+
|
| 322 |
+
lines = []
|
| 323 |
+
for i, (start, end) in enumerate(slices):
|
| 324 |
+
stroke_pts = deformed_points[start:end]
|
| 325 |
+
idxs = np.linspace(0, len(stroke_pts) - 1, n_points).astype(int)
|
| 326 |
+
pts = stroke_pts[idxs]
|
| 327 |
+
pts_str = " -> ".join(f"({x:.0f},{y:.0f})" for x, y in pts)
|
| 328 |
+
lines.append(f" stroke_{i}: {pts_str}")
|
| 329 |
+
return "\n".join(lines)
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
def build_all_keyframes_deformed_geometry_text(objects_info, deform_outputs, n_keyframes=N_KEYFRAMES):
|
| 333 |
+
"""
|
| 334 |
+
For EVERY object and EVERY keyframe of the CURRENT attempt, reconstruct
|
| 335 |
+
the actual deformed stroke geometry as text — this is the anchor source
|
| 336 |
+
for the judge's target_coords grounding: the judge is instructed to
|
| 337 |
+
find a visually-identified stroke's coordinate HERE, per keyframe, not
|
| 338 |
+
in the rest-pose geometry (which is only correct when that keyframe
|
| 339 |
+
happens to match rest).
|
| 340 |
+
"""
|
| 341 |
+
sections = []
|
| 342 |
+
for obj_name, info in objects_info.items():
|
| 343 |
+
obj_targets = (deform_outputs or {}).get(obj_name)
|
| 344 |
+
if not obj_targets:
|
| 345 |
+
continue
|
| 346 |
+
kf_parts = []
|
| 347 |
+
for kf in range(n_keyframes):
|
| 348 |
+
kf_key = f"kf{kf}"
|
| 349 |
+
if kf_key not in obj_targets:
|
| 350 |
+
continue
|
| 351 |
+
shape_text = build_deformed_stroke_geometry_text(info, obj_targets[kf_key], n_points=2)
|
| 352 |
+
if shape_text:
|
| 353 |
+
kf_parts.append(f" -- {kf_key} --\n{shape_text}")
|
| 354 |
+
if kf_parts:
|
| 355 |
+
sections.append(f'Object "{obj_name}":\n' + "\n".join(kf_parts))
|
| 356 |
+
return "\n\n".join(sections)
|
| 357 |
+
|
| 358 |
+
|
| 359 |
+
def render_rest_pose_multi(object_names, svg_path, semantic_path, out_path):
|
| 360 |
+
"""
|
| 361 |
+
Renders multiple objects' ORIGINAL strokes together, at their real
|
| 362 |
+
positions in the source SVG (no deformation, no trajectory translation)
|
| 363 |
+
— this is the "what does the sketch actually look like" image every
|
| 364 |
+
attempt is grounded in, so Qwen can see what's actually drawable
|
| 365 |
+
(e.g. whether the dog's back legs even exist as strokes) instead of
|
| 366 |
+
only reasoning from the caption's text description.
|
| 367 |
+
"""
|
| 368 |
+
import matplotlib.pyplot as plt
|
| 369 |
+
fig, ax = plt.subplots(figsize=(6, 6))
|
| 370 |
+
for name in object_names:
|
| 371 |
+
points, slices = load_object(name, svg_path, semantic_path)
|
| 372 |
+
for start, end in slices:
|
| 373 |
+
seg = points[start:end]
|
| 374 |
+
ax.plot(seg[:, 0], seg[:, 1], color="black", linewidth=1.2)
|
| 375 |
+
ax.invert_yaxis()
|
| 376 |
+
ax.set_aspect("equal")
|
| 377 |
+
ax.set_title("rest pose")
|
| 378 |
+
fig.savefig(out_path, dpi=150, bbox_inches="tight")
|
| 379 |
+
plt.close(fig)
|
| 380 |
+
return out_path
|
| 381 |
+
|
| 382 |
+
|
| 383 |
+
def format_joint_feedback_for_object(object_name, joint_feedback):
|
| 384 |
+
"""
|
| 385 |
+
Filters the judge's structured joint_feedback list down to entries for
|
| 386 |
+
THIS object and formats them as explicit, actionable lines. When the
|
| 387 |
+
judge grounded its suggestion in an actual traced coordinate,
|
| 388 |
+
delta_px (computed by vlm_judge.compute_joint_feedback_deltas via
|
| 389 |
+
subtraction, NOT model math) is shown as an exact number the
|
| 390 |
+
generator can apply directly — e.g.
|
| 391 |
+
"joint_1 at kf3: hand not near mouth -> move by (dx=+12.0, dy=-8.0)
|
| 392 |
+
px (target anchored to: mouth stroke in face object)"
|
| 393 |
+
If the judge couldn't ground a target (target_coords was null, or the
|
| 394 |
+
joint/keyframe didn't match anything in deform_outputs), delta_px is
|
| 395 |
+
None and the line falls back to the prose issue only — explicitly
|
| 396 |
+
labeled as ungrounded rather than silently presenting a guess as if
|
| 397 |
+
it were a computed number.
|
| 398 |
+
Returns "" if there's no feedback for this object (nothing to add).
|
| 399 |
+
"""
|
| 400 |
+
if not joint_feedback:
|
| 401 |
+
return ""
|
| 402 |
+
relevant = [f for f in joint_feedback if f.get("object") == object_name]
|
| 403 |
+
if not relevant:
|
| 404 |
+
return ""
|
| 405 |
+
lines = []
|
| 406 |
+
for f in relevant:
|
| 407 |
+
kf = f.get("keyframe")
|
| 408 |
+
kf_str = f"kf{kf}" if kf is not None else "unspecified keyframe"
|
| 409 |
+
issue = f.get("issue", "")
|
| 410 |
+
delta = f.get("delta_px")
|
| 411 |
+
if delta is not None:
|
| 412 |
+
anchor = f.get("anchored_to", "unspecified")
|
| 413 |
+
lines.append(
|
| 414 |
+
f" - joint_{f.get('joint')} at {kf_str}: {issue} "
|
| 415 |
+
f"-> move by (dx={delta['dx']:+.1f}, dy={delta['dy']:+.1f}) px "
|
| 416 |
+
f"(target anchored to: {anchor})"
|
| 417 |
+
)
|
| 418 |
+
else:
|
| 419 |
+
lines.append(
|
| 420 |
+
f" - joint_{f.get('joint')} at {kf_str}: {issue} "
|
| 421 |
+
f"-> [UNGROUNDED — no traceable/disambiguated coordinate was given; use your own judgment]"
|
| 422 |
+
)
|
| 423 |
+
return "\n".join(lines)
|
| 424 |
+
|
| 425 |
+
|
| 426 |
+
def build_combined_object_section(object_name, joints, anchor_idx, handle_idxs, bbox_size,
|
| 427 |
+
previous_narrative=None, feedback=None, n_keyframes=N_KEYFRAMES,
|
| 428 |
+
freeze_narrative=False, strokes=None, joint_feedback=None,
|
| 429 |
+
baseline_targets=None):
|
| 430 |
+
"""
|
| 431 |
+
freeze_narrative: if True, previous_narrative is used as a FIXED target
|
| 432 |
+
pose description (the object's narrative already passed faithfulness —
|
| 433 |
+
only the numeric deformation needs to improve, not the story). If
|
| 434 |
+
False (default), the narrative is regenerated fresh, informed by the
|
| 435 |
+
attached image(s) and feedback, same as before.
|
| 436 |
+
|
| 437 |
+
strokes: if given, the object's actual stroke points are included as
|
| 438 |
+
TEXT (not just the rendered image) — added specifically because
|
| 439 |
+
multimodal LLMs can under-attend to image content relative to text;
|
| 440 |
+
this gives the same geometric information in a text-native form the
|
| 441 |
+
model is more likely to actually use. Kept deliberately sparse (2
|
| 442 |
+
points per stroke, start+end only) since a complex object can have
|
| 443 |
+
60+ strokes — measured on real dog3 data: 2 points/stroke costs
|
| 444 |
+
~570 tokens for a 64-stroke object vs ~2650 for the full 12
|
| 445 |
+
points/stroke used internally for the ARAP mesh.
|
| 446 |
+
|
| 447 |
+
joint_feedback: the FULL joint_feedback list from the judge's verdict
|
| 448 |
+
(all objects) — filtered down to this object's entries and formatted
|
| 449 |
+
as precise per-joint lines. Falls back to nothing (not an error) if
|
| 450 |
+
the judge didn't return joint_feedback (e.g. coordinate context wasn't
|
| 451 |
+
given to build_judge_prompt) — the object still gets the old flat
|
| 452 |
+
`feedback` string via previous_block/feedback_line below.
|
| 453 |
+
|
| 454 |
+
baseline_targets: {kf_key: {joint_i_str: [x,y]}} — this object's EXACT
|
| 455 |
+
numeric targets from the PREVIOUS attempt, given on EVERY retry (not
|
| 456 |
+
just after a pass) so the generator refines real numbers instead of
|
| 457 |
+
reconstructing them from images/prose each time — the "guess
|
| 458 |
+
coordinates from a picture" problem numeric joint_feedback exists to
|
| 459 |
+
avoid elsewhere. Joints with a GROUNDED joint_feedback correction are
|
| 460 |
+
excluded here (that correction takes precedence) — this only shows
|
| 461 |
+
joints without a specific correction, as a "keep unless you have
|
| 462 |
+
reason to change" anchor, not a fixed target — unlike `pose_lines`
|
| 463 |
+
used when freeze_narrative=True, which locks the story, not the
|
| 464 |
+
coordinates.
|
| 465 |
+
"""
|
| 466 |
+
cap = round(bbox_size * 0.25, 1)
|
| 467 |
+
joint_feedback_text = format_joint_feedback_for_object(object_name, joint_feedback)
|
| 468 |
+
joint_feedback_block = (
|
| 469 |
+
f"\n Specific per-joint corrections from the judge (apply these precisely, this is not general "
|
| 470 |
+
f"guidance):\n{joint_feedback_text}\n"
|
| 471 |
+
) if joint_feedback_text else ""
|
| 472 |
+
baseline_block = ""
|
| 473 |
+
if baseline_targets:
|
| 474 |
+
# joints the judge gave a GROUNDED correction for (real delta_px, not ungrounded prose) are
|
| 475 |
+
# excluded from the raw baseline dump below — joint_feedback_block above is the authoritative
|
| 476 |
+
# instruction for those specific joints, and repeating the stale pre-correction number here
|
| 477 |
+
# would be redundant at best and contradictory at worst (two different numbers for the same
|
| 478 |
+
# joint, no clear precedence). Baseline only shows joints WITHOUT a grounded correction, i.e.
|
| 479 |
+
# "keep these as they were unless you have your own reason to change them."
|
| 480 |
+
corrected_joints = {
|
| 481 |
+
f.get("joint") for f in (joint_feedback or [])
|
| 482 |
+
if f.get("object") == object_name and f.get("delta_px") is not None
|
| 483 |
+
}
|
| 484 |
+
kf_parts = []
|
| 485 |
+
for kf, kf_vals in baseline_targets.items():
|
| 486 |
+
shown = {j: v for j, v in kf_vals.items()
|
| 487 |
+
if int(j.replace("joint_", "")) not in corrected_joints}
|
| 488 |
+
if shown:
|
| 489 |
+
kf_parts.append(f"{kf}: {{{', '.join(f'{j}={v}' for j, v in shown.items())}}}")
|
| 490 |
+
if kf_parts:
|
| 491 |
+
baseline_block = (
|
| 492 |
+
f"\n These are the EXACT numeric targets from the PREVIOUS attempt for joints the judge did "
|
| 493 |
+
f"NOT give a specific correction for above — a real, working starting point, not a guess: "
|
| 494 |
+
f"{', '.join(kf_parts)}\n"
|
| 495 |
+
f" Keep these numbers unless you have a genuine reason to change them — do not discard them "
|
| 496 |
+
f"and reinvent from scratch. For any joint listed in the corrections above instead, follow "
|
| 497 |
+
f"that correction, not these numbers (that joint is intentionally omitted here).\n"
|
| 498 |
+
)
|
| 499 |
+
joint_lines = "\n".join(
|
| 500 |
+
f' - joint_{i}: rest position (x={joints[i][0]:.1f}, y={joints[i][1]:.1f})'
|
| 501 |
+
+ (" <-- ANCHOR, must stay at or near this position in EVERY keyframe" if i == anchor_idx else "")
|
| 502 |
+
for i in range(len(joints))
|
| 503 |
+
)
|
| 504 |
+
handle_list_str = ', joint_'.join(str(i) for i in handle_idxs)
|
| 505 |
+
|
| 506 |
+
geometry_block = ""
|
| 507 |
+
if strokes:
|
| 508 |
+
geometry_text = build_stroke_geometry_text(strokes, n_points=2)
|
| 509 |
+
geometry_block = (
|
| 510 |
+
f"\n This object's ACTUAL drawn strokes (start -> end point of each stroke, same coordinate "
|
| 511 |
+
f"space as the joints above) — use this to know exactly what is and isn't actually drawn, don't "
|
| 512 |
+
f"invent motion for parts that have no strokes here:\n{geometry_text}\n"
|
| 513 |
+
)
|
| 514 |
+
|
| 515 |
+
if freeze_narrative and previous_narrative:
|
| 516 |
+
pose_lines = "\n".join(f" kf{i}: {desc}" for i, desc in enumerate(previous_narrative))
|
| 517 |
+
feedback_line = f'\n This pose story already matches the intended action — it is FIXED, do not change it. ' \
|
| 518 |
+
f'Only the numeric target positions need to improve.' \
|
| 519 |
+
+ (f' Previous attempt was judged: "{feedback}"' if feedback else "") + \
|
| 520 |
+
"\n The attached images show exactly what the previous attempt's target positions " \
|
| 521 |
+
"actually looked like when rendered — use them to see specifically what needs to change numerically."
|
| 522 |
+
return f"""Object: "{object_name}"
|
| 523 |
+
Joints:
|
| 524 |
+
{joint_lines}
|
| 525 |
+
{geometry_block}
|
| 526 |
+
Target pose across all {n_keyframes} keyframes (FIXED, already correct — do not rewrite):
|
| 527 |
+
{pose_lines}
|
| 528 |
+
{feedback_line}
|
| 529 |
+
{joint_feedback_block}
|
| 530 |
+
{baseline_block}
|
| 531 |
+
For non-anchor joints (joint_{handle_list_str}), do not move more than {cap} pixels from REST in any keyframe. IMPORTANT: these {n_keyframes} keyframes are SPARSE anchor points spanning the ENTIRE action, NOT consecutive video frames — a large, dramatic difference between consecutive keyframes is NORMAL and EXPECTED, not an error; the actual in-between motion will be generated separately later by a different model. Positions should progress in a DIRECTIONALLY COHERENT way (don't make real progress toward the action and then have a LATER keyframe randomly revert backward without the narrative describing a reason to — e.g. only "landing"/"settling" should move back toward rest). Small, timid, barely-different positions between keyframes are themselves a mistake, not a safe choice."""
|
| 532 |
+
|
| 533 |
+
previous_block = ""
|
| 534 |
+
if previous_narrative or feedback:
|
| 535 |
+
parts = []
|
| 536 |
+
if previous_narrative:
|
| 537 |
+
parts.append(f"Your previous narrative attempt was:\n{json.dumps(previous_narrative, indent=2)}")
|
| 538 |
+
if feedback:
|
| 539 |
+
parts.append(f'That attempt was judged and received this critique: "{feedback}"')
|
| 540 |
+
parts.append("The attached images show exactly what that previous attempt actually looked like when "
|
| 541 |
+
"rendered. Look at them, understand what specifically was wrong, and revise BOTH the "
|
| 542 |
+
"narrative and the target positions to fix it — don't just reword the narrative "
|
| 543 |
+
"superficially while leaving the same underlying problem.")
|
| 544 |
+
previous_block = "\n " + "\n ".join(parts) + "\n"
|
| 545 |
+
|
| 546 |
+
return f"""Object: "{object_name}"
|
| 547 |
+
Joints:
|
| 548 |
+
{joint_lines}
|
| 549 |
+
{geometry_block}
|
| 550 |
+
{previous_block}
|
| 551 |
+
{joint_feedback_block}
|
| 552 |
+
{baseline_block}
|
| 553 |
+
For non-anchor joints (joint_{handle_list_str}), do not move more than {cap} pixels from REST in any keyframe. IMPORTANT: these {n_keyframes} keyframes are SPARSE anchor points spanning the ENTIRE action, NOT consecutive video frames — a large, dramatic difference between consecutive keyframes is NORMAL and EXPECTED, not an error; the actual in-between motion will be generated separately later by a different model. Positions should progress in a DIRECTIONALLY COHERENT way (don't make real progress toward the action and then have a LATER keyframe randomly revert backward without the narrative describing a reason to — e.g. only "landing"/"settling" should move back toward rest). Small, timid, barely-different positions between keyframes are themselves a mistake, not a safe choice."""
|
| 554 |
+
|
| 555 |
+
|
| 556 |
+
DOG3_COMBINED_FEWSHOT_EXAMPLE = {
|
| 557 |
+
"dog": {
|
| 558 |
+
"narrative": [
|
| 559 |
+
"the dog is sitting alert, watching the frisbee as it is thrown",
|
| 560 |
+
"the dog is beginning to rise, weight shifting forward, head reaching toward the frisbee",
|
| 561 |
+
"the dog is mid-leap, body extended, reaching far forward and up toward the frisbee",
|
| 562 |
+
"the dog is at the peak of its jump, reaching as far as possible toward the frisbee",
|
| 563 |
+
"the dog is landing after catching the frisbee, body compacting back down",
|
| 564 |
+
],
|
| 565 |
+
# real dog3 joint rest positions: joint_1=(178.28,136.85) head, joint_2=(228.78,194.94)
|
| 566 |
+
# tail, joint_3=(189.29,151.35) neck — every value below verified to stay within a
|
| 567 |
+
# 27px cap of rest. Notice the progression BUILDS UP through kf0->kf3 (increasing
|
| 568 |
+
# displacement, matching "rising -> leaping -> peak reach") and only SETTLES BACK at
|
| 569 |
+
# kf4 ("landing") — this is the exact monotonic-then-settle shape that was missing
|
| 570 |
+
# when a real run produced a kf2 spike with kf3/kf4 reverting toward rest with no
|
| 571 |
+
# narrative reason to.
|
| 572 |
+
"targets": {
|
| 573 |
+
"kf0": {"joint_1": [178.3, 136.8], "joint_2": [228.8, 194.9], "joint_3": [189.3, 151.3]},
|
| 574 |
+
"kf1": {"joint_1": [168.0, 127.0], "joint_2": [232.0, 191.0], "joint_3": [184.0, 144.0]},
|
| 575 |
+
"kf2": {"joint_1": [160.0, 120.0], "joint_2": [237.0, 186.0], "joint_3": [177.0, 137.0]},
|
| 576 |
+
"kf3": {"joint_1": [159.0, 119.0], "joint_2": [240.0, 183.0], "joint_3": [174.0, 134.0]},
|
| 577 |
+
"kf4": {"joint_1": [168.0, 128.0], "joint_2": [231.0, 192.0], "joint_3": [185.0, 146.0]},
|
| 578 |
+
},
|
| 579 |
+
}
|
| 580 |
+
}
|
| 581 |
+
|
| 582 |
+
|
| 583 |
+
def build_combined_narrate_deform_prompt(objects_info, caption, previous_narratives=None, feedback=None,
|
| 584 |
+
n_keyframes=N_KEYFRAMES, is_retry=False, few_shot=True,
|
| 585 |
+
freeze_narrative=False, joint_feedback=None, baseline_targets=None):
|
| 586 |
+
sections, example_parts = [], []
|
| 587 |
+
for name, info in objects_info.items():
|
| 588 |
+
prev_narrative_for_obj = (previous_narratives or {}).get(name)
|
| 589 |
+
obj_baseline_targets = (baseline_targets or {}).get(name)
|
| 590 |
+
sections.append(build_combined_object_section(
|
| 591 |
+
name, info["joints"], info["anchor_idx"], info["handle_idxs"], info["bbox_size"],
|
| 592 |
+
previous_narrative=prev_narrative_for_obj, feedback=feedback, n_keyframes=n_keyframes,
|
| 593 |
+
freeze_narrative=freeze_narrative, strokes=info.get("strokes"), joint_feedback=joint_feedback,
|
| 594 |
+
baseline_targets=obj_baseline_targets,
|
| 595 |
+
))
|
| 596 |
+
kf_examples = ",\n".join(
|
| 597 |
+
" \"kf%d\": {%s}" % (kf, ", ".join(f'"joint_{i}": [x, y]' for i in info["handle_idxs"]))
|
| 598 |
+
for kf in range(n_keyframes)
|
| 599 |
+
)
|
| 600 |
+
if freeze_narrative:
|
| 601 |
+
example_parts.append(
|
| 602 |
+
f' "{name}": {{\n'
|
| 603 |
+
f' "targets": {{\n{kf_examples}\n }}\n'
|
| 604 |
+
f' }}'
|
| 605 |
+
)
|
| 606 |
+
else:
|
| 607 |
+
example_parts.append(
|
| 608 |
+
f' "{name}": {{\n'
|
| 609 |
+
f' "narrative": [<{n_keyframes} short pose description strings, one per keyframe>],\n'
|
| 610 |
+
f' "targets": {{\n{kf_examples}\n }}\n'
|
| 611 |
+
f' }}'
|
| 612 |
+
)
|
| 613 |
+
|
| 614 |
+
all_sections = "\n\n".join(sections)
|
| 615 |
+
example_json = "{\n" + ",\n".join(example_parts) + "\n}"
|
| 616 |
+
|
| 617 |
+
image_context = ""
|
| 618 |
+
if is_retry:
|
| 619 |
+
image_context = ("The FIRST image attached is the object's original rest pose (undeformed). "
|
| 620 |
+
"The remaining images are the actual rendered result of your PREVIOUS attempt, "
|
| 621 |
+
"one per keyframe, in order.")
|
| 622 |
+
else:
|
| 623 |
+
image_context = ("The attached image shows the object's original rest pose (undeformed) — use this "
|
| 624 |
+
"to understand what strokes actually exist and are available to move; do not "
|
| 625 |
+
"invent motion for body parts that aren't actually drawn.")
|
| 626 |
+
|
| 627 |
+
fewshot_block = ""
|
| 628 |
+
if few_shot:
|
| 629 |
+
fewshot_json = json.dumps(DOG3_COMBINED_FEWSHOT_EXAMPLE, indent=2)
|
| 630 |
+
fewshot_block = f"""Example — for the scene "{DOG3_CAPTION}", a good answer looks like:
|
| 631 |
+
{fewshot_json}
|
| 632 |
+
|
| 633 |
+
Notice: each narrative keyframe reads as a distinct, substantially different stage of the action — not a near-duplicate of its neighbor, and not a small incremental change from it. The joint targets BUILD UP smoothly (kf0 -> kf1 -> kf2 -> kf3 each moving further than the last) and only settle back toward rest at the FINAL keyframe, matching the narrative's "landing" moment — no keyframe overshoots and then has a later keyframe revert back toward rest without a narrative reason to. Match this style and this kind of numeric consistency for the new scene below.
|
| 634 |
+
|
| 635 |
+
"""
|
| 636 |
+
|
| 637 |
+
if freeze_narrative:
|
| 638 |
+
output_instruction = (
|
| 639 |
+
'For EACH object above, the narrative/pose story is already fixed (shown above) — '
|
| 640 |
+
'produce ONLY:\n'
|
| 641 |
+
' "targets": target (x, y) positions for its non-anchor joints, at every keyframe, '
|
| 642 |
+
'consistent with the fixed pose story above.'
|
| 643 |
+
)
|
| 644 |
+
else:
|
| 645 |
+
output_instruction = (
|
| 646 |
+
"For EACH object above, produce BOTH:\n"
|
| 647 |
+
' 1. "narrative": a plain-English pose description for each keyframe. REMEMBER: these are '
|
| 648 |
+
"SPARSE keyframes spanning the WHOLE action, not consecutive video frames — each description "
|
| 649 |
+
"should be a meaningfully, substantially different stage of the action from its neighbors, not "
|
| 650 |
+
"a small incremental change. Write these like 5 distinct captions for 5 different moments spread "
|
| 651 |
+
"across an entire action, not like 5 near-duplicate snapshots a split-second apart. Under 20 "
|
| 652 |
+
"words each.\n"
|
| 653 |
+
' 2. "targets": target (x, y) positions for its non-anchor joints, at every keyframe, '
|
| 654 |
+
"consistent with your own narrative."
|
| 655 |
+
)
|
| 656 |
+
|
| 657 |
+
return f"""{fewshot_block}You are directing a {n_keyframes}-keyframe animated sequence for a hand-drawn sketch, viewed from the side. Coordinate system: x increases rightward, y increases DOWNWARD.
|
| 658 |
+
|
| 659 |
+
IMPORTANT: these {n_keyframes} keyframes are SPARSE anchor points sampled across the ENTIRE action from start to finish — NOT consecutive video frames. Think of them like 5 widely-spaced snapshots of a whole motion, not neighboring frames a fraction of a second apart. Large, dramatic pose changes between consecutive keyframes are normal and expected; a separate model will generate the actual in-between motion frames later. Do not treat these like near-continuous animation frames.
|
| 660 |
+
|
| 661 |
+
Scene: "{caption}"
|
| 662 |
+
|
| 663 |
+
{image_context}
|
| 664 |
+
|
| 665 |
+
{all_sections}
|
| 666 |
+
|
| 667 |
+
{output_instruction}
|
| 668 |
+
|
| 669 |
+
Consider objects together (e.g. a dog reaching toward a frisbee should be spatially consistent with the frisbee's own position) and consider each object's OWN sequence together — these {n_keyframes} keyframes are SPARSE anchor points spanning the WHOLE action, not consecutive video frames, so large differences between consecutive keyframes are expected and correct, not something to avoid. The many actual in-between motion frames will be generated separately later. Only avoid a keyframe making real progress and then a LATER keyframe randomly reverting backward without the narrative describing why.
|
| 670 |
+
|
| 671 |
+
Respond with ONLY one JSON object, no other text, in this exact format:
|
| 672 |
+
{example_json}
|
| 673 |
+
"""
|
| 674 |
+
|
| 675 |
+
|
| 676 |
+
def run_combined_narrate_deform(model, processor, images, prompt, temperature=0.6):
|
| 677 |
+
"""Same multi-image calling pattern as vlm_judge.run_judge — reused
|
| 678 |
+
here since both are Qwen3-VL calls with a list of images + one prompt.
|
| 679 |
+
|
| 680 |
+
temperature default raised from 0.1 to 0.6 — CONFIRMED on real hardware
|
| 681 |
+
(eat2, 3 attempts) that at 0.1 the model reproduced its joint targets
|
| 682 |
+
as an EXACT copy of the rest-position legend values (not just "close
|
| 683 |
+
to rest" — bit-identical to the decimal) on attempt 1, before any
|
| 684 |
+
feedback existed to explain it. Near-zero temperature strongly favors
|
| 685 |
+
the single highest-probability continuation, and copying a number
|
| 686 |
+
already visible in-context (the rest-position legend, formatted in
|
| 687 |
+
the same [x, y] style as the requested targets) is a low-risk, easy
|
| 688 |
+
completion under numeric uncertainty. This is a hypothesis about
|
| 689 |
+
mechanism, not confirmed root cause — raising temperature is the
|
| 690 |
+
cheapest test of it before trying a bigger/different model.
|
| 691 |
+
"""
|
| 692 |
+
import torch
|
| 693 |
+
content = [{"type": "image", "image": img} for img in images]
|
| 694 |
+
content.append({"type": "text", "text": prompt})
|
| 695 |
+
messages = [{"role": "user", "content": content}]
|
| 696 |
+
|
| 697 |
+
text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
| 698 |
+
inputs = processor(text=[text], images=images, return_tensors="pt").to(model.device)
|
| 699 |
+
|
| 700 |
+
with torch.no_grad():
|
| 701 |
+
output_ids = model.generate(**inputs, max_new_tokens=500 * N_KEYFRAMES,
|
| 702 |
+
temperature=temperature, do_sample=True)
|
| 703 |
+
|
| 704 |
+
generated = output_ids[:, inputs["input_ids"].shape[1]:]
|
| 705 |
+
response = processor.batch_decode(generated, skip_special_tokens=True)[0]
|
| 706 |
+
return response
|
| 707 |
+
|
| 708 |
+
|
| 709 |
+
def apply_deform_clip_and_write(parsed, objects_info, out_dir, sketch_name, n_keyframes=N_KEYFRAMES,
|
| 710 |
+
frozen_narratives=None):
|
| 711 |
+
"""
|
| 712 |
+
Shared post-processing for the combined call's "targets" section:
|
| 713 |
+
same hard-clip logic as the old run_deform, applied here instead.
|
| 714 |
+
Returns (narratives_dict, deform_outputs_dict, cap_utilization_dict).
|
| 715 |
+
|
| 716 |
+
frozen_narratives: {obj_name: [...]} — used as a fallback when the
|
| 717 |
+
response doesn't include a "narrative" key for an object, which
|
| 718 |
+
happens when freeze_narrative=True was used in the prompt (the model
|
| 719 |
+
was never asked to produce one, so its absence is expected, not an
|
| 720 |
+
error — carry the frozen one forward instead of losing it).
|
| 721 |
+
|
| 722 |
+
cap_utilization: {obj_name: mean_fraction_of_cap_used} — CONFIRMED on
|
| 723 |
+
real hardware (horsecar5's person) that Qwen can propose displacement
|
| 724 |
+
well within the movement cap without ever being told it did so — the
|
| 725 |
+
handle selection and clipping were both working correctly, but the
|
| 726 |
+
actual output was too timid to be visible (e.g. a leg moving only 18%
|
| 727 |
+
of its allowed range). The hard clip only ever catches OVER the cap;
|
| 728 |
+
nothing previously caught UNDER-using it. This surfaces that as an
|
| 729 |
+
explicit number so it can be fed back to Qwen directly.
|
| 730 |
+
"""
|
| 731 |
+
narratives_out = {}
|
| 732 |
+
deform_outputs = {}
|
| 733 |
+
utilization_by_obj = {} # {obj_name: [fraction, fraction, ...]} across all joints/keyframes
|
| 734 |
+
|
| 735 |
+
for obj_name, info in objects_info.items():
|
| 736 |
+
if obj_name not in parsed:
|
| 737 |
+
print(f" WARNING: '{obj_name}' missing from response entirely, skipping")
|
| 738 |
+
continue
|
| 739 |
+
obj_result = parsed[obj_name]
|
| 740 |
+
utilization_by_obj[obj_name] = []
|
| 741 |
+
|
| 742 |
+
if "narrative" in obj_result:
|
| 743 |
+
narratives_out[obj_name] = obj_result["narrative"]
|
| 744 |
+
elif frozen_narratives and obj_name in frozen_narratives:
|
| 745 |
+
narratives_out[obj_name] = frozen_narratives[obj_name]
|
| 746 |
+
else:
|
| 747 |
+
print(f" WARNING: '{obj_name}' has no narrative in response and no frozen narrative "
|
| 748 |
+
f"to fall back to — narratives.json will be missing this object")
|
| 749 |
+
|
| 750 |
+
deform_outputs[obj_name] = {}
|
| 751 |
+
targets = obj_result.get("targets", {})
|
| 752 |
+
for kf in range(n_keyframes):
|
| 753 |
+
kf_key = f"kf{kf}"
|
| 754 |
+
if kf_key not in targets:
|
| 755 |
+
print(f" WARNING: '{obj_name}' missing {kf_key} targets, skipping this frame")
|
| 756 |
+
continue
|
| 757 |
+
kf_result = targets[kf_key]
|
| 758 |
+
out = {}
|
| 759 |
+
joint_targets_this_kf = {}
|
| 760 |
+
for name, target in kf_result.items():
|
| 761 |
+
m = re.match(r"joint_(\d+)", name)
|
| 762 |
+
if not m:
|
| 763 |
+
print(f" WARNING: unexpected key '{name}' for '{obj_name}' {kf_key}, skipping")
|
| 764 |
+
continue
|
| 765 |
+
joint_i = int(m.group(1))
|
| 766 |
+
if joint_i >= len(info["joint_mesh_indices"]):
|
| 767 |
+
print(f" WARNING: '{obj_name}' joint_{joint_i} out of range, skipping")
|
| 768 |
+
continue
|
| 769 |
+
# Qwen occasionally returns a bare scalar (e.g. 195.4) instead of an [x, y] pair for
|
| 770 |
+
# a joint target. np.array(scalar, dtype=float) does NOT error — it silently makes a
|
| 771 |
+
# 0-d array, which then broadcasts against `rest` (2-d) into a plausible-looking 2-d
|
| 772 |
+
# `disp` with no error anywhere in THIS function. The malformed value then gets
|
| 773 |
+
# written to disk as-is and only crashes 3 steps later in run_render, when it tries
|
| 774 |
+
# to stack this scalar next to properly-shaped [x,y] entries — as an inhomogeneous
|
| 775 |
+
# array error that gives no indication which object/joint/keyframe was actually bad.
|
| 776 |
+
# CONFIRMED on real hardware (eagle4, football7/'person') that this happens on real
|
| 777 |
+
# Qwen output, not just as a theoretical edge case.
|
| 778 |
+
target_list = target if isinstance(target, (list, tuple)) else [target]
|
| 779 |
+
if len(target_list) != 2:
|
| 780 |
+
print(f" WARNING: '{obj_name}' {kf_key} joint_{joint_i} target = {target!r}, expected "
|
| 781 |
+
f"an [x, y] pair (got {len(target_list)} value(s)) — skipping this joint for this "
|
| 782 |
+
f"keyframe rather than writing a malformed value that would crash rendering later")
|
| 783 |
+
continue
|
| 784 |
+
mesh_idx = info["joint_mesh_indices"][joint_i]
|
| 785 |
+
|
| 786 |
+
rest = np.array(info["joints"][joint_i])
|
| 787 |
+
cap = round(info["bbox_size"] * 0.25, 1)
|
| 788 |
+
target_arr = np.array(target_list, dtype=float)
|
| 789 |
+
disp = target_arr - rest
|
| 790 |
+
dist = np.linalg.norm(disp)
|
| 791 |
+
utilization_by_obj[obj_name].append(min(dist / cap, 1.0) if cap > 0 else 0.0)
|
| 792 |
+
if dist > cap:
|
| 793 |
+
clipped = rest + disp / dist * cap
|
| 794 |
+
print(f" CLIPPED '{obj_name}' {kf_key} joint_{joint_i}: requested {dist:.1f}px "
|
| 795 |
+
f"(cap {cap}px) -> clipped to {cap}px, direction preserved")
|
| 796 |
+
target = clipped.tolist()
|
| 797 |
+
else:
|
| 798 |
+
target = target_arr.tolist()
|
| 799 |
+
|
| 800 |
+
out[str(mesh_idx)] = target
|
| 801 |
+
joint_targets_this_kf[f"joint_{joint_i}"] = target
|
| 802 |
+
anchor_mesh_idx = info["joint_mesh_indices"][info["anchor_idx"]]
|
| 803 |
+
out[str(anchor_mesh_idx)] = info["joints"][info["anchor_idx"]].tolist()
|
| 804 |
+
deform_outputs[obj_name][kf_key] = joint_targets_this_kf
|
| 805 |
+
|
| 806 |
+
out_path = os.path.join(out_dir, f"qwen_{sketch_name}_{obj_name}_kf{kf}.json")
|
| 807 |
+
with open(out_path, "w") as f:
|
| 808 |
+
json.dump(out, f, indent=2)
|
| 809 |
+
print(f" wrote {out_path}")
|
| 810 |
+
|
| 811 |
+
cap_utilization = {}
|
| 812 |
+
for obj_name, fractions in utilization_by_obj.items():
|
| 813 |
+
if fractions:
|
| 814 |
+
mean_frac = sum(fractions) / len(fractions)
|
| 815 |
+
cap_utilization[obj_name] = mean_frac
|
| 816 |
+
print(f" '{obj_name}': mean cap utilization = {mean_frac*100:.0f}% "
|
| 817 |
+
f"(across {len(fractions)} joint-keyframe pairs)")
|
| 818 |
+
|
| 819 |
+
return narratives_out, deform_outputs, cap_utilization
|
| 820 |
+
# =============================================================================
|
| 821 |
+
# STEP 4: render — compose the full scene (no Qwen call at all)
|
| 822 |
+
# =============================================================================
|
| 823 |
+
|
| 824 |
+
def run_render(handles_dir, svg_path, semantic_path, traj_path,
|
| 825 |
+
out_path=None, frames_dir=None):
|
| 826 |
+
"""
|
| 827 |
+
out_path: if given, ALSO saves the combined strip image (all 5 keyframes
|
| 828 |
+
side by side) here, outside frames_dir. Optional — pass None
|
| 829 |
+
to keep output confined to frames_dir only.
|
| 830 |
+
frames_dir: if given, saves each keyframe as its own individual PNG
|
| 831 |
+
(kf0.png ... kf4.png) plus the combined strip, named after
|
| 832 |
+
the sketch itself ({sketch_name}.png), all inside this one
|
| 833 |
+
folder.
|
| 834 |
+
"""
|
| 835 |
+
import matplotlib.pyplot as plt
|
| 836 |
+
|
| 837 |
+
sketch_name = os.path.splitext(os.path.basename(svg_path))[0]
|
| 838 |
+
real_trajectories = load_trajectories(traj_path)
|
| 839 |
+
object_names = list(real_trajectories.keys())
|
| 840 |
+
|
| 841 |
+
object_data = {}
|
| 842 |
+
for name in object_names:
|
| 843 |
+
points, slices = load_object(name, svg_path, semantic_path)
|
| 844 |
+
dx_vals, dy_vals = bbox_deltas(real_trajectories[name])
|
| 845 |
+
unique_points, p2u = deduplicate_points(points, tol=0.35)
|
| 846 |
+
tri, edges = build_mesh(unique_points)
|
| 847 |
+
object_data[name] = {
|
| 848 |
+
"points": points, "slices": slices, "dx": dx_vals, "dy": dy_vals,
|
| 849 |
+
"unique_points": unique_points, "p2u": p2u, "edges": edges,
|
| 850 |
+
}
|
| 851 |
+
|
| 852 |
+
if frames_dir:
|
| 853 |
+
os.makedirs(frames_dir, exist_ok=True)
|
| 854 |
+
|
| 855 |
+
xmin, xmax, ymin, ymax = 0, 260, 60, 230
|
| 856 |
+
|
| 857 |
+
# compute each keyframe's drawing data once, reused for both the combined
|
| 858 |
+
# strip and the individual per-keyframe images
|
| 859 |
+
keyframe_lines = [] # list of {name: [(seg_x, seg_y), ...]} per keyframe
|
| 860 |
+
for kf in range(N_KEYFRAMES):
|
| 861 |
+
lines_this_kf = {}
|
| 862 |
+
for name in object_names:
|
| 863 |
+
od = object_data[name]
|
| 864 |
+
handles_path = os.path.join(handles_dir, f"qwen_{sketch_name}_{name}_kf{kf}.json")
|
| 865 |
+
|
| 866 |
+
if os.path.exists(handles_path):
|
| 867 |
+
with open(handles_path) as f:
|
| 868 |
+
spec = json.load(f)
|
| 869 |
+
handle_indices = [int(k) for k in spec.keys()]
|
| 870 |
+
handle_targets = np.array([spec[k] for k in spec.keys()])
|
| 871 |
+
n_verts = len(od["unique_points"])
|
| 872 |
+
bad = [i for i in handle_indices if i >= n_verts]
|
| 873 |
+
if bad:
|
| 874 |
+
print(f" ERROR: {handles_path} has out-of-bounds indices {bad} for '{name}' "
|
| 875 |
+
f"({n_verts} mesh vertices) — likely from a DIFFERENT sketch's mesh. "
|
| 876 |
+
f"Falling back to translation-only.")
|
| 877 |
+
deformed_points = od["points"]
|
| 878 |
+
mode = "translation-only (handles file failed validation)"
|
| 879 |
+
else:
|
| 880 |
+
deformed_unique = arap_deform(od["unique_points"], od["edges"],
|
| 881 |
+
handle_indices, handle_targets, iterations=10)
|
| 882 |
+
deformed_points = deformed_unique[od["p2u"]]
|
| 883 |
+
mode = "ARAP"
|
| 884 |
+
else:
|
| 885 |
+
deformed_points = od["points"]
|
| 886 |
+
mode = "translation-only (no handles file found)"
|
| 887 |
+
|
| 888 |
+
moved = deformed_points + np.array([od["dx"][kf], od["dy"][kf]])
|
| 889 |
+
lines_this_kf[name] = [moved[start:end] for start, end in od["slices"]]
|
| 890 |
+
|
| 891 |
+
print(f"kf{kf} '{name}': {mode}, points_after_move_range="
|
| 892 |
+
f"x[{moved[:,0].min():.1f},{moved[:,0].max():.1f}] "
|
| 893 |
+
f"y[{moved[:,1].min():.1f},{moved[:,1].max():.1f}]")
|
| 894 |
+
|
| 895 |
+
keyframe_lines.append(lines_this_kf)
|
| 896 |
+
|
| 897 |
+
if frames_dir:
|
| 898 |
+
fig_i, ax_i = plt.subplots(figsize=(6, 5.5))
|
| 899 |
+
for name, segs in lines_this_kf.items():
|
| 900 |
+
for seg in segs:
|
| 901 |
+
ax_i.plot(seg[:, 0], seg[:, 1],
|
| 902 |
+
color=OBJECT_COLORS.get(name, DEFAULT_COLOR),
|
| 903 |
+
linewidth=OBJECT_LINEWIDTH.get(name, DEFAULT_LINEWIDTH))
|
| 904 |
+
ax_i.set_xlim(xmin, xmax)
|
| 905 |
+
ax_i.set_ylim(ymax, ymin)
|
| 906 |
+
ax_i.set_aspect("equal")
|
| 907 |
+
ax_i.set_title(f"{sketch_name} — kf{kf}", fontsize=12, fontweight="bold")
|
| 908 |
+
frame_path = os.path.join(frames_dir, f"kf{kf}.png")
|
| 909 |
+
fig_i.savefig(frame_path, dpi=140, bbox_inches="tight")
|
| 910 |
+
plt.close(fig_i)
|
| 911 |
+
print(f" wrote {frame_path}")
|
| 912 |
+
|
| 913 |
+
# combined strip, same as before
|
| 914 |
+
fig, axes = plt.subplots(1, N_KEYFRAMES, figsize=(24, 5))
|
| 915 |
+
for kf in range(N_KEYFRAMES):
|
| 916 |
+
ax = axes[kf]
|
| 917 |
+
for name, segs in keyframe_lines[kf].items():
|
| 918 |
+
for seg in segs:
|
| 919 |
+
ax.plot(seg[:, 0], seg[:, 1],
|
| 920 |
+
color=OBJECT_COLORS.get(name, DEFAULT_COLOR),
|
| 921 |
+
linewidth=OBJECT_LINEWIDTH.get(name, DEFAULT_LINEWIDTH))
|
| 922 |
+
ax.set_xlim(xmin, xmax)
|
| 923 |
+
ax.set_ylim(ymax, ymin)
|
| 924 |
+
ax.set_aspect("equal")
|
| 925 |
+
ax.set_title(f"kf{kf}", fontsize=13, fontweight="bold")
|
| 926 |
+
|
| 927 |
+
plt.tight_layout()
|
| 928 |
+
|
| 929 |
+
if out_path:
|
| 930 |
+
plt.savefig(out_path, dpi=140, bbox_inches="tight")
|
| 931 |
+
print(f"wrote {out_path}")
|
| 932 |
+
|
| 933 |
+
if frames_dir:
|
| 934 |
+
combined_frame_path = os.path.join(frames_dir, f"{sketch_name}.png")
|
| 935 |
+
plt.savefig(combined_frame_path, dpi=140, bbox_inches="tight")
|
| 936 |
+
print(f"wrote {combined_frame_path}")
|
| 937 |
+
|
| 938 |
+
if not out_path and not frames_dir:
|
| 939 |
+
print("WARNING: neither out_path nor frames_dir given, combined strip image not saved anywhere")
|
| 940 |
+
|
| 941 |
+
plt.close(fig)
|
| 942 |
+
|
| 943 |
+
|
| 944 |
+
# =============================================================================
|
| 945 |
+
# CLI — single input: a sketch name or an SVG path. Everything else is
|
| 946 |
+
# derived automatically from the directory conventions used throughout
|
| 947 |
+
# this dataset. Override flags exist for the rare case a path doesn't
|
| 948 |
+
# match convention, but nothing is required beyond the sketch itself.
|
| 949 |
+
# =============================================================================
|
| 950 |
+
|
| 951 |
+
# Confirmed real paths from this dataset, used as defaults so nothing else
|
| 952 |
+
# needs to be typed per run. If your layout differs, override with the
|
| 953 |
+
# corresponding --*-dir / --*-file flag below.
|
| 954 |
+
SVG_DIR_DEFAULT = "/user/HS400/rk01499/my_scratch/sketch/data/raw/60sketches/svg"
|
| 955 |
+
PROCESSED_DIR_DEFAULT = "/user/HS400/rk01499/my_scratch/sketch/data/processed"
|
| 956 |
+
CAPTION_FILE_DEFAULT = "/user/HS400/rk01499/my_scratch/sketch/data/raw/60sketches/caption.txt"
|
| 957 |
+
MODEL_PATH_DEFAULT = "/user/HS400/rk01499/my_scratch/models/qwen2.5-7b/"
|
| 958 |
+
|
| 959 |
+
|
| 960 |
+
def resolve_sketch_paths(sketch, svg_dir, processed_dir, caption_file):
|
| 961 |
+
"""
|
| 962 |
+
sketch: either a bare sketch name ("dog9") or a path to its SVG
|
| 963 |
+
("/path/to/dog9.svg") — either way, everything else (semantic,
|
| 964 |
+
traj, caption) is derived from the same naming convention used
|
| 965 |
+
across this dataset: {name}.svg, {name}/{name}_semantic.txt,
|
| 966 |
+
{name}/{name}_traj.txt, and a lookup in one shared caption.txt.
|
| 967 |
+
"""
|
| 968 |
+
name = os.path.splitext(os.path.basename(sketch))[0]
|
| 969 |
+
svg_path = sketch if sketch.endswith(".svg") else os.path.join(svg_dir, f"{name}.svg")
|
| 970 |
+
semantic_path = os.path.join(processed_dir, name, f"{name}_semantic.txt")
|
| 971 |
+
traj_path = os.path.join(processed_dir, name, f"{name}_traj.txt")
|
| 972 |
+
|
| 973 |
+
missing = [p for p in [svg_path, semantic_path, traj_path, caption_file] if not os.path.exists(p)]
|
| 974 |
+
if missing:
|
| 975 |
+
raise SystemExit(
|
| 976 |
+
f"Could not find these expected files for sketch '{name}':\n " +
|
| 977 |
+
"\n ".join(missing) +
|
| 978 |
+
"\n\nIf your directory layout differs from the default, pass --svg-dir / "
|
| 979 |
+
"--processed-dir / --caption-file explicitly."
|
| 980 |
+
)
|
| 981 |
+
|
| 982 |
+
caption = get_caption(caption_file, name)
|
| 983 |
+
return name, svg_path, semantic_path, traj_path, caption
|
| 984 |
+
|
| 985 |
+
|
| 986 |
+
def main():
|
| 987 |
+
ap = argparse.ArgumentParser(
|
| 988 |
+
description="Run the full sketch deformation pipeline for one image. "
|
| 989 |
+
"The only required input is the sketch — everything else "
|
| 990 |
+
"(semantic assignments, trajectory, caption) is looked up "
|
| 991 |
+
"automatically from the standard dataset layout.")
|
| 992 |
+
ap.add_argument("sketch", type=str,
|
| 993 |
+
help="sketch name (e.g. 'dog9') or path to its .svg file")
|
| 994 |
+
ap.add_argument("--model", type=str, default=MODEL_PATH_DEFAULT)
|
| 995 |
+
ap.add_argument("--svg-dir", type=str, default=SVG_DIR_DEFAULT)
|
| 996 |
+
ap.add_argument("--processed-dir", type=str, default=PROCESSED_DIR_DEFAULT)
|
| 997 |
+
ap.add_argument("--caption-file", type=str, default=CAPTION_FILE_DEFAULT)
|
| 998 |
+
ap.add_argument("--out-dir", type=str, default=".")
|
| 999 |
+
ap.add_argument("--no-fewshot", action="store_true",
|
| 1000 |
+
help="disable the dog3 few-shot example in narrate (for A/B comparison)")
|
| 1001 |
+
ap.add_argument("--narrate-temperature", type=float, default=0.6,
|
| 1002 |
+
help="sampling temperature for the narrate+deform call (default 0.6, raised from the "
|
| 1003 |
+
"original 0.1 — CONFIRMED that 0.1 caused the model to copy rest-position values "
|
| 1004 |
+
"verbatim as targets on attempt 1; use this flag to sweep other values)")
|
| 1005 |
+
ap.add_argument("--deform-only", action="store_true",
|
| 1006 |
+
help="run classify + ONE narrate+deform attempt + render, then STOP — no judge, "
|
| 1007 |
+
"no retries. Prints cap utilization directly and saves the render, so you can "
|
| 1008 |
+
"inspect raw generation quality without the judge's assessment as a confound.")
|
| 1009 |
+
args = ap.parse_args()
|
| 1010 |
+
|
| 1011 |
+
name, svg_path, semantic_path, traj_path, caption = resolve_sketch_paths(
|
| 1012 |
+
args.sketch, args.svg_dir, args.processed_dir, args.caption_file)
|
| 1013 |
+
print(f"sketch: {name}")
|
| 1014 |
+
print(f" svg: {svg_path}")
|
| 1015 |
+
print(f" semantic: {semantic_path}")
|
| 1016 |
+
print(f" traj: {traj_path}")
|
| 1017 |
+
print(f" caption: {caption!r}")
|
| 1018 |
+
|
| 1019 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 1020 |
+
json_dir = os.path.join(args.out_dir, "json", name)
|
| 1021 |
+
os.makedirs(json_dir, exist_ok=True)
|
| 1022 |
+
print(f" json output dir: {json_dir}")
|
| 1023 |
+
|
| 1024 |
+
print("\n########## STEP 1: CLASSIFY ##########")
|
| 1025 |
+
classify_model, classify_tokenizer, classify_device = load_qwen_model(args.model)
|
| 1026 |
+
deformation, arap_objects = run_classify(
|
| 1027 |
+
classify_model, classify_tokenizer, classify_device, caption, semantic_path,
|
| 1028 |
+
os.path.join(json_dir, f"{name}_deformation.json"))
|
| 1029 |
+
classify_model = unload_model(classify_model)
|
| 1030 |
+
|
| 1031 |
+
if arap_objects:
|
| 1032 |
+
import vlm_judge
|
| 1033 |
+
|
| 1034 |
+
temp_dir = os.path.join(args.out_dir, "P_1", name)
|
| 1035 |
+
objects_info = build_objects_info(svg_path, semantic_path, arap_objects)
|
| 1036 |
+
if not objects_info:
|
| 1037 |
+
print("No objects with valid handles found after mesh setup. Nothing to do.")
|
| 1038 |
+
objects_info = None
|
| 1039 |
+
|
| 1040 |
+
# rest joint positions were previously computed once and kept ONLY in memory for the rest
|
| 1041 |
+
# of main()'s lifetime — no file ever recorded them, so a question like "is joint_2's target
|
| 1042 |
+
# actually different from its rest position, or just restating rest" was unanswerable after
|
| 1043 |
+
# a run finished. Saved once per sketch (not per attempt, since rest pose doesn't change
|
| 1044 |
+
# attempt to attempt) so it's always available for exactly this kind of check.
|
| 1045 |
+
if objects_info:
|
| 1046 |
+
rest_joints_dump = {
|
| 1047 |
+
obj_name: {
|
| 1048 |
+
f"joint_{i}": {"x": float(j[0]), "y": float(j[1]),
|
| 1049 |
+
"is_anchor": i == info["anchor_idx"]}
|
| 1050 |
+
for i, j in enumerate(info["joints"])
|
| 1051 |
+
}
|
| 1052 |
+
for obj_name, info in objects_info.items()
|
| 1053 |
+
}
|
| 1054 |
+
with open(os.path.join(json_dir, f"{name}_rest_joints.json"), "w") as f:
|
| 1055 |
+
json.dump(rest_joints_dump, f, indent=2)
|
| 1056 |
+
|
| 1057 |
+
rest_pose_image_path = os.path.join(json_dir, f"{name}_rest_pose.png")
|
| 1058 |
+
render_rest_pose_multi(arap_objects, svg_path, semantic_path, rest_pose_image_path)
|
| 1059 |
+
|
| 1060 |
+
# preprocessed bbox trajectory data — fixed ground truth, given to the
|
| 1061 |
+
# judge as spatial context for EVERY object (ARAP and TRAJ_ONLY alike),
|
| 1062 |
+
# not something the judge critiques or the generator controls
|
| 1063 |
+
real_trajectories = load_trajectories(traj_path)
|
| 1064 |
+
all_object_names = list(real_trajectories.keys())
|
| 1065 |
+
|
| 1066 |
+
import dino_similarity
|
| 1067 |
+
dino_model, dino_processor = dino_similarity.load_dino_model()
|
| 1068 |
+
|
| 1069 |
+
import clip_score
|
| 1070 |
+
clip_model, clip_processor = clip_score.load_clip_model()
|
| 1071 |
+
|
| 1072 |
+
feedback = None
|
| 1073 |
+
previous_narratives = None # {obj_name: [5 descriptions]} from the last attempt
|
| 1074 |
+
previous_temp_dir = None # where the last attempt's rendered kf0..kf4 images live
|
| 1075 |
+
freeze_narrative = False # only frozen once faithfulness has already passed once
|
| 1076 |
+
consecutive_stagnant = 0 # early-stop if DINOv2 confirms no real change 2 attempts in a row
|
| 1077 |
+
joint_feedback = None # judge's structured per-joint corrections from the last attempt
|
| 1078 |
+
final_verdict = None
|
| 1079 |
+
winning_attempt = None
|
| 1080 |
+
all_attempts_summary = []
|
| 1081 |
+
# NUMERIC CONTINUITY: every attempt after the first is given the PREVIOUS attempt's actual
|
| 1082 |
+
# numeric targets (baseline_targets) to refine, not just images/prose to reconstruct numbers
|
| 1083 |
+
# from scratch. This applies to EVERY retry, not just a special round after a pass.
|
| 1084 |
+
baseline_targets = None
|
| 1085 |
+
# CONFIRMATION ROUND: when an attempt first passes all three thresholds, don't stop
|
| 1086 |
+
# immediately — run exactly ONE more attempt (which, same as any retry now, gets the
|
| 1087 |
+
# passing attempt's real baseline_targets to refine) to see if it can be beaten, then keep
|
| 1088 |
+
# whichever actually scores higher. The passing attempt is NEVER at risk of being replaced
|
| 1089 |
+
# by something worse — if the confirmation attempt doesn't beat it, the original passing
|
| 1090 |
+
# attempt is kept exactly as if this mechanism didn't exist. Fires ONCE per sketch.
|
| 1091 |
+
first_pass_attempt = None # attempt number of the FIRST attempt that passed all thresholds
|
| 1092 |
+
first_pass_scores = None # that attempt's (faithfulness_score, plausibility_score)
|
| 1093 |
+
confirmation_used = False # True once the one extra confirmation attempt has been consumed
|
| 1094 |
+
|
| 1095 |
+
for attempt in range(1, MAX_RETRIES + 2): # +2, not +1: room for exactly one confirmation
|
| 1096 |
+
# attempt beyond MAX_RETRIES if the pass happens
|
| 1097 |
+
# on the very last regular attempt
|
| 1098 |
+
if not objects_info:
|
| 1099 |
+
break
|
| 1100 |
+
if attempt > MAX_RETRIES and (first_pass_attempt is None or confirmation_used):
|
| 1101 |
+
# only allowed to exceed MAX_RETRIES for the ONE confirmation attempt — never for an
|
| 1102 |
+
# ordinary failed-retry continuation
|
| 1103 |
+
break
|
| 1104 |
+
label = f"{attempt}/{MAX_RETRIES}" if attempt <= MAX_RETRIES else f"{attempt} (CONFIRMATION, beyond normal {MAX_RETRIES})"
|
| 1105 |
+
print(f"\n########## ATTEMPT {label} ##########")
|
| 1106 |
+
|
| 1107 |
+
attempt_json_dir = os.path.join(json_dir, "attempts", f"attempt_{attempt}")
|
| 1108 |
+
attempt_temp_dir = os.path.join(temp_dir, "attempts", f"attempt_{attempt}")
|
| 1109 |
+
os.makedirs(attempt_json_dir, exist_ok=True)
|
| 1110 |
+
|
| 1111 |
+
# image list: always the rest pose; from attempt 2+, ALSO the
|
| 1112 |
+
# previous attempt's actual rendered keyframes, so Qwen sees
|
| 1113 |
+
# exactly what its last attempt looked like, not just a text
|
| 1114 |
+
# description of it
|
| 1115 |
+
from PIL import Image
|
| 1116 |
+
images = [Image.open(rest_pose_image_path).convert("RGB")]
|
| 1117 |
+
is_retry = attempt > 1
|
| 1118 |
+
if is_retry:
|
| 1119 |
+
images.extend(vlm_judge.load_keyframe_images(previous_temp_dir))
|
| 1120 |
+
|
| 1121 |
+
prompt = build_combined_narrate_deform_prompt(
|
| 1122 |
+
objects_info, caption, previous_narratives=previous_narratives,
|
| 1123 |
+
feedback=feedback, is_retry=is_retry, few_shot=not args.no_fewshot,
|
| 1124 |
+
freeze_narrative=freeze_narrative, joint_feedback=joint_feedback,
|
| 1125 |
+
baseline_targets=baseline_targets)
|
| 1126 |
+
# save the FULL generator prompt too, symmetric with the judge prompt below — otherwise
|
| 1127 |
+
# there's no way to confirm feedback/joint_feedback actually appeared in what Qwen was
|
| 1128 |
+
# shown, only to infer it from whether behavior changed afterward.
|
| 1129 |
+
with open(os.path.join(attempt_json_dir, f"{name}_narrate_deform_prompt.txt"), "w") as f:
|
| 1130 |
+
f.write(prompt)
|
| 1131 |
+
|
| 1132 |
+
print(f"\n---------- STEP 2+3: NARRATE+DEFORM (attempt {attempt}, "
|
| 1133 |
+
f"{len(images)} image{'s' if len(images) != 1 else ''}"
|
| 1134 |
+
f"{', narrative FROZEN' if freeze_narrative else ''}) ----------")
|
| 1135 |
+
vlm_model, vlm_processor = vlm_judge.load_vlm("Qwen/Qwen2.5-VL-3B-Instruct")
|
| 1136 |
+
response = run_combined_narrate_deform(vlm_model, vlm_processor, images, prompt,
|
| 1137 |
+
temperature=args.narrate_temperature)
|
| 1138 |
+
# save the RAW pre-parse response unconditionally, before anything downstream can fail
|
| 1139 |
+
# or silently collapse it — previously this text existed only transiently in memory and
|
| 1140 |
+
# was discarded the moment parsing succeeded, so a case like identical coordinates across
|
| 1141 |
+
# every keyframe couldn't be traced back to "did Qwen write that itself" vs "did something
|
| 1142 |
+
# downstream produce it" after the fact. Written to the SAME attempt_json_dir the parsed
|
| 1143 |
+
# outputs already live in, so raw and parsed are side by side for direct comparison.
|
| 1144 |
+
with open(os.path.join(attempt_json_dir, f"{name}_narrate_deform_raw_response.txt"), "w") as f:
|
| 1145 |
+
f.write(response)
|
| 1146 |
+
try:
|
| 1147 |
+
parsed = parse_json_response(response)
|
| 1148 |
+
except (ValueError, json.JSONDecodeError) as e:
|
| 1149 |
+
print(f"FAILED TO PARSE: {e}\nraw: {response}")
|
| 1150 |
+
vlm_model = unload_model(vlm_model)
|
| 1151 |
+
feedback = "the previous attempt's output could not be parsed; produce valid JSON in the exact requested format"
|
| 1152 |
+
all_attempts_summary.append({"attempt": attempt, "plausibility_score": None, "note": "narrate+deform parse failed"})
|
| 1153 |
+
continue
|
| 1154 |
+
|
| 1155 |
+
narratives_this_attempt, deform_outputs_this_attempt, cap_utilization_this_attempt = apply_deform_clip_and_write(
|
| 1156 |
+
parsed, objects_info, attempt_json_dir, name,
|
| 1157 |
+
frozen_narratives=previous_narratives if freeze_narrative else None)
|
| 1158 |
+
with open(os.path.join(attempt_json_dir, f"{name}_narratives.json"), "w") as f:
|
| 1159 |
+
json.dump(narratives_this_attempt, f, indent=2)
|
| 1160 |
+
|
| 1161 |
+
# print each ARAP object's ACTUAL joint targets per keyframe — this is the real
|
| 1162 |
+
# signal for "did deformation happen", unlike points_after_move_range in the render
|
| 1163 |
+
# log below, which is the whole object's bbox and can stay constant even when a
|
| 1164 |
+
# small joint (e.g. a hand) moves substantially, since the head/torso/limb extremes
|
| 1165 |
+
# usually dominate the bbox regardless of hand position.
|
| 1166 |
+
print(f"\n---------- attempt {attempt}: actual joint targets per keyframe ----------")
|
| 1167 |
+
for obj_name, obj_targets in deform_outputs_this_attempt.items():
|
| 1168 |
+
print(f" '{obj_name}':")
|
| 1169 |
+
for kf in range(N_KEYFRAMES):
|
| 1170 |
+
kf_key = f"kf{kf}"
|
| 1171 |
+
if kf_key in obj_targets:
|
| 1172 |
+
print(f" {kf_key}: {obj_targets[kf_key]}")
|
| 1173 |
+
|
| 1174 |
+
print(f"\n---------- STEP 4: RENDER (attempt {attempt}, no Qwen) ----------")
|
| 1175 |
+
run_render(attempt_json_dir, svg_path, semantic_path, traj_path, frames_dir=attempt_temp_dir)
|
| 1176 |
+
|
| 1177 |
+
if args.deform_only:
|
| 1178 |
+
print(f"\n########## --deform-only: STOPPING after attempt 1, no judge ##########")
|
| 1179 |
+
print(f"cap_utilization (raw, unfiltered by any threshold):")
|
| 1180 |
+
for obj, frac in (cap_utilization_this_attempt or {}).items():
|
| 1181 |
+
print(f" {obj}: {frac*100:.1f}% of allowed movement used")
|
| 1182 |
+
print(f"\nInspect the actual render directly at: {attempt_temp_dir}")
|
| 1183 |
+
print(f"(kf0.png ... kf4.png, plus the combined strip)")
|
| 1184 |
+
sys.exit(0)
|
| 1185 |
+
|
| 1186 |
+
stagnation_result = None
|
| 1187 |
+
if previous_temp_dir:
|
| 1188 |
+
print(f"\n---------- STAGNATION CHECK (attempt {attempt} vs attempt {attempt - 1}) ----------")
|
| 1189 |
+
prev_dino_images = vlm_judge.load_keyframe_images(previous_temp_dir)
|
| 1190 |
+
curr_dino_images = vlm_judge.load_keyframe_images(attempt_temp_dir)
|
| 1191 |
+
stagnation_result = dino_similarity.stagnation_score(
|
| 1192 |
+
dino_model, dino_processor, prev_dino_images, curr_dino_images)
|
| 1193 |
+
stagnant = dino_similarity.is_stagnant(stagnation_result)
|
| 1194 |
+
print(f"mean attempt-to-attempt similarity: {stagnation_result['mean_similarity']:.4f} "
|
| 1195 |
+
f"({'STAGNANT' if stagnant else 'changed'})")
|
| 1196 |
+
consecutive_stagnant = consecutive_stagnant + 1 if stagnant else 0
|
| 1197 |
+
|
| 1198 |
+
print(f"\n---------- TEMPORAL CONSISTENCY (attempt {attempt}, diagnostic only) ----------")
|
| 1199 |
+
temporal_images = vlm_judge.load_keyframe_images(attempt_temp_dir)
|
| 1200 |
+
temporal_result = dino_similarity.temporal_consistency(dino_model, dino_processor, temporal_images)
|
| 1201 |
+
|
| 1202 |
+
clip_result = None
|
| 1203 |
+
if caption:
|
| 1204 |
+
print(f"\n---------- CLIP SCORE (attempt {attempt}) ----------")
|
| 1205 |
+
clip_images = vlm_judge.load_keyframe_images(attempt_temp_dir)
|
| 1206 |
+
clip_result = clip_score.compute_sequence_clip_scores(clip_model, clip_processor, clip_images, caption)
|
| 1207 |
+
|
| 1208 |
+
print(f"\n---------- STEP 5: JUDGE (attempt {attempt}, images + joint/bbox coordinates) ----------")
|
| 1209 |
+
# reuse the SAME already-loaded VLM for judging — no reload
|
| 1210 |
+
# needed, since narrate+deform and judge are both Qwen3-VL calls now.
|
| 1211 |
+
# Judge sees: rest pose + this attempt's 5 rendered keyframes
|
| 1212 |
+
# (images), PLUS rest-pose stroke text + joint legend + this
|
| 1213 |
+
# attempt's joint targets + every object's fixed bbox trajectory
|
| 1214 |
+
# + PER-KEYFRAME DEFORMED stroke geometry (text) — the last one
|
| 1215 |
+
# is what target_coords grounding actually anchors against: a
|
| 1216 |
+
# visually-identified region's TRUE coordinate at that specific
|
| 1217 |
+
# keyframe, not its rest-pose coordinate (which is only correct
|
| 1218 |
+
# when a keyframe happens to match rest).
|
| 1219 |
+
judge_images = [Image.open(rest_pose_image_path).convert("RGB")]
|
| 1220 |
+
judge_images.extend(vlm_judge.load_keyframe_images(attempt_temp_dir))
|
| 1221 |
+
|
| 1222 |
+
deformed_geometry_text = build_all_keyframes_deformed_geometry_text(
|
| 1223 |
+
objects_info, deform_outputs_this_attempt)
|
| 1224 |
+
# save this — it's otherwise invisible after the fact. Needed to verify e.g. whether a
|
| 1225 |
+
# judge's target_coords anchor was itself built from stale/frozen coordinates (a real
|
| 1226 |
+
# failure mode: if deform_outputs_this_attempt is frozen across keyframes, this text
|
| 1227 |
+
# will be too, and a judge "fix" anchored to it would just point back at the same stuck
|
| 1228 |
+
# position rather than actually correcting anything).
|
| 1229 |
+
with open(os.path.join(attempt_json_dir, f"{name}_deformed_geometry_text.txt"), "w") as f:
|
| 1230 |
+
f.write(deformed_geometry_text)
|
| 1231 |
+
|
| 1232 |
+
judge_prompt = vlm_judge.build_judge_prompt(
|
| 1233 |
+
name, caption=caption, dino_stagnation=stagnation_result,
|
| 1234 |
+
dino_temporal=temporal_result, clip_scores=clip_result,
|
| 1235 |
+
objects_info=objects_info, deform_outputs=deform_outputs_this_attempt,
|
| 1236 |
+
real_trajectories=real_trajectories, all_object_names=all_object_names,
|
| 1237 |
+
deformed_geometry_text=deformed_geometry_text, narratives=narratives_this_attempt)
|
| 1238 |
+
# save the FULL prompt too, not just the response — otherwise there's no way to see
|
| 1239 |
+
# exactly what the judge was shown (only what it said back), which makes it impossible
|
| 1240 |
+
# to distinguish "the judge reasoned badly" from "the judge was given bad/stale input".
|
| 1241 |
+
with open(os.path.join(attempt_json_dir, f"{name}_judge_prompt.txt"), "w") as f:
|
| 1242 |
+
f.write(judge_prompt)
|
| 1243 |
+
judge_response = vlm_judge.run_judge(vlm_model, vlm_processor, judge_images, judge_prompt)
|
| 1244 |
+
vlm_model = unload_model(vlm_model)
|
| 1245 |
+
# same rationale as the narrate+deform raw response above — save unconditionally, before
|
| 1246 |
+
# parsing, so the judge's literal output is inspectable after the fact even when parsing
|
| 1247 |
+
# succeeds (previously only visible on parse failure, via print, not saved to disk).
|
| 1248 |
+
with open(os.path.join(attempt_json_dir, f"{name}_judge_raw_response.txt"), "w") as f:
|
| 1249 |
+
f.write(judge_response)
|
| 1250 |
+
|
| 1251 |
+
try:
|
| 1252 |
+
verdict = vlm_judge.parse_judge_response(judge_response)
|
| 1253 |
+
except (ValueError, json.JSONDecodeError) as e:
|
| 1254 |
+
print(f"JUDGE FAILED TO PARSE: {e}\nraw: {judge_response}")
|
| 1255 |
+
print("Treating as a failed attempt, retrying without specific feedback.")
|
| 1256 |
+
feedback = "the previous attempt's evaluation could not be parsed; try a clearer, more varied pose progression"
|
| 1257 |
+
previous_narratives = narratives_this_attempt
|
| 1258 |
+
previous_temp_dir = attempt_temp_dir
|
| 1259 |
+
freeze_narrative = False # unknown state — safest to regenerate rather than assume faithfulness held
|
| 1260 |
+
joint_feedback = None # verdict didn't parse, so any joint_feedback in it is unusable/unknown — don't carry stale feedback forward
|
| 1261 |
+
all_attempts_summary.append({"attempt": attempt, "plausibility_score": None, "note": "judge parse failed"})
|
| 1262 |
+
continue
|
| 1263 |
+
|
| 1264 |
+
problems = vlm_judge.validate_judge_response(verdict, valid_arap_objects=set(objects_info.keys()))
|
| 1265 |
+
print("\n--- JUDGE VERDICT ---")
|
| 1266 |
+
print(json.dumps(verdict, indent=2))
|
| 1267 |
+
if problems:
|
| 1268 |
+
for p in problems:
|
| 1269 |
+
print(f" VALIDATION PROBLEM: {p}")
|
| 1270 |
+
|
| 1271 |
+
with open(os.path.join(attempt_json_dir, f"{name}_judge_verdict.json"), "w") as f:
|
| 1272 |
+
json.dump(verdict, f, indent=2)
|
| 1273 |
+
|
| 1274 |
+
final_verdict = verdict
|
| 1275 |
+
winning_attempt = attempt
|
| 1276 |
+
score = verdict.get("plausibility_score")
|
| 1277 |
+
faith_score = verdict.get("faithfulness_score")
|
| 1278 |
+
quality_score = verdict.get("quality_score")
|
| 1279 |
+
all_attempts_summary.append({"attempt": attempt, "plausibility_score": score,
|
| 1280 |
+
"plausibility_notes": verdict.get("plausibility_notes"),
|
| 1281 |
+
"faithfulness_score": faith_score,
|
| 1282 |
+
"faithfulness_notes": verdict.get("faithfulness_notes"),
|
| 1283 |
+
"quality_score": quality_score,
|
| 1284 |
+
"quality_notes": verdict.get("quality_notes"),
|
| 1285 |
+
"dino_stagnant": dino_similarity.is_stagnant(stagnation_result) if stagnation_result else None})
|
| 1286 |
+
|
| 1287 |
+
plausibility_ok = isinstance(score, (int, float)) and score >= PLAUSIBILITY_THRESHOLD
|
| 1288 |
+
faithfulness_ok = faithfulness_passed(faith_score)
|
| 1289 |
+
quality_ok = isinstance(quality_score, (int, float)) and quality_score >= QUALITY_THRESHOLD
|
| 1290 |
+
print(f"\nplausibility_score = {score} (threshold = {PLAUSIBILITY_THRESHOLD}, "
|
| 1291 |
+
f"{'PASS' if plausibility_ok else 'FAIL'})")
|
| 1292 |
+
print(f"faithfulness_score = {faith_score} (threshold = {FAITHFULNESS_THRESHOLD}, "
|
| 1293 |
+
f"{'PASS' if faithfulness_ok else 'FAIL'})")
|
| 1294 |
+
print(f"quality_score = {quality_score} (threshold = {QUALITY_THRESHOLD}, "
|
| 1295 |
+
f"{'PASS' if quality_ok else 'FAIL'})")
|
| 1296 |
+
|
| 1297 |
+
if plausibility_ok and faithfulness_ok and quality_ok:
|
| 1298 |
+
if first_pass_attempt is None:
|
| 1299 |
+
# FIRST time passing — don't stop yet. Remember this as the safe fallback, run
|
| 1300 |
+
# exactly one more attempt to see if it can be beaten, THEN decide.
|
| 1301 |
+
first_pass_attempt = attempt
|
| 1302 |
+
first_pass_scores = (faith_score, score)
|
| 1303 |
+
print(f"Attempt {attempt} passed all thresholds — running ONE confirmation attempt "
|
| 1304 |
+
f"before finalizing, to see if it can be improved on. If the confirmation "
|
| 1305 |
+
f"attempt is not better, attempt {attempt} is kept exactly as-is. (It also "
|
| 1306 |
+
f"gets this attempt's real numeric targets as its baseline, same as any "
|
| 1307 |
+
f"normal retry now would.)")
|
| 1308 |
+
previous_narratives = narratives_this_attempt
|
| 1309 |
+
previous_temp_dir = attempt_temp_dir
|
| 1310 |
+
baseline_targets = deform_outputs_this_attempt
|
| 1311 |
+
joint_feedback = vlm_judge.compute_joint_feedback_deltas(
|
| 1312 |
+
verdict.get("joint_feedback"), deform_outputs_this_attempt)
|
| 1313 |
+
# feedback stays neutral encouragement, not a correction — this attempt already
|
| 1314 |
+
# passed, there's nothing specifically "wrong" to fix, just seeing if variation
|
| 1315 |
+
# produces something even better
|
| 1316 |
+
feedback = ("This attempt already passed all quality thresholds. This is an "
|
| 1317 |
+
"OPTIONAL confirmation attempt: try to match or improve on it, but do "
|
| 1318 |
+
"not discard what is already working.")
|
| 1319 |
+
freeze_narrative = False
|
| 1320 |
+
continue # do NOT break — proceed to the confirmation attempt
|
| 1321 |
+
else:
|
| 1322 |
+
# this IS the confirmation attempt (first_pass_attempt was already set)
|
| 1323 |
+
confirmation_used = True
|
| 1324 |
+
confirmation_scores = (faith_score, score)
|
| 1325 |
+
if confirmation_scores > first_pass_scores:
|
| 1326 |
+
print(f"Confirmation attempt {attempt} scored better "
|
| 1327 |
+
f"(faithfulness={faith_score}, plausibility={score}) than the original "
|
| 1328 |
+
f"passing attempt {first_pass_attempt} "
|
| 1329 |
+
f"(faithfulness={first_pass_scores[0]}, plausibility={first_pass_scores[1]}) "
|
| 1330 |
+
f"— using attempt {attempt} instead.")
|
| 1331 |
+
winning_attempt = attempt
|
| 1332 |
+
else:
|
| 1333 |
+
print(f"Confirmation attempt {attempt} did NOT score better than the original "
|
| 1334 |
+
f"passing attempt {first_pass_attempt} — keeping attempt {first_pass_attempt} "
|
| 1335 |
+
f"as originally found, discarding the confirmation attempt.")
|
| 1336 |
+
winning_attempt = first_pass_attempt
|
| 1337 |
+
break
|
| 1338 |
+
|
| 1339 |
+
if first_pass_attempt is not None and not confirmation_used:
|
| 1340 |
+
# confirmation attempt FAILED thresholds outright (didn't beat the pass AND didn't
|
| 1341 |
+
# even clear the bar itself) — keep the original passing attempt, don't treat this
|
| 1342 |
+
# as a real failure requiring further retries
|
| 1343 |
+
confirmation_used = True
|
| 1344 |
+
print(f"Confirmation attempt {attempt} did not pass thresholds — keeping the original "
|
| 1345 |
+
f"passing attempt {first_pass_attempt} as the final result.")
|
| 1346 |
+
winning_attempt = first_pass_attempt
|
| 1347 |
+
break
|
| 1348 |
+
|
| 1349 |
+
if consecutive_stagnant >= 2:
|
| 1350 |
+
print(f"\nDINOv2 confirmed NO real change across {consecutive_stagnant} consecutive attempts "
|
| 1351 |
+
f"(attempt {attempt} vs {attempt-1}, and {attempt-1} vs {attempt-2}) — further retries "
|
| 1352 |
+
f"are very unlikely to help. Stopping early and keeping this attempt's result rather "
|
| 1353 |
+
f"than burning through the remaining {MAX_RETRIES - attempt} attempts.")
|
| 1354 |
+
break
|
| 1355 |
+
|
| 1356 |
+
previous_narratives = narratives_this_attempt
|
| 1357 |
+
previous_temp_dir = attempt_temp_dir
|
| 1358 |
+
# NUMERIC CONTINUITY (universal, every retry): carry this attempt's actual numeric
|
| 1359 |
+
# targets forward as the NEXT attempt's starting point to refine, not just images/prose
|
| 1360 |
+
# to reconstruct numbers from scratch. This is the core of the "attempt N+1 works on
|
| 1361 |
+
# attempt N's real output" redesign — replaces relying purely on visual/textual
|
| 1362 |
+
# reconstruction, which this session repeatedly found unreliable (copy-from-rest,
|
| 1363 |
+
# scalar-collapse, frozen joints).
|
| 1364 |
+
baseline_targets = deform_outputs_this_attempt
|
| 1365 |
+
# carry the judge's structured per-joint corrections into the NEXT
|
| 1366 |
+
# attempt's prompt — empty list is valid (judge found nothing
|
| 1367 |
+
# specific to fix), missing key means coordinate context wasn't
|
| 1368 |
+
# given to build_judge_prompt at all; either way, default to None
|
| 1369 |
+
# so format_joint_feedback_for_object just adds nothing.
|
| 1370 |
+
# compute_joint_feedback_deltas turns the judge's target_coords
|
| 1371 |
+
# into an exact pixel delta by SUBTRACTION against this attempt's
|
| 1372 |
+
# actual joint positions — arithmetic, not model math. Entries
|
| 1373 |
+
# where the judge couldn't ground a target (target_coords null)
|
| 1374 |
+
# are left as delta_px=None and fall back to prose-only feedback.
|
| 1375 |
+
joint_feedback = vlm_judge.compute_joint_feedback_deltas(
|
| 1376 |
+
verdict.get("joint_feedback"), deform_outputs_this_attempt)
|
| 1377 |
+
# freeze the narrative on the NEXT attempt only if faithfulness
|
| 1378 |
+
# already passed THIS attempt — no reason to keep regenerating
|
| 1379 |
+
# a story that's already correct, only the numbers need work
|
| 1380 |
+
freeze_narrative = faithfulness_ok
|
| 1381 |
+
|
| 1382 |
+
if attempt == MAX_RETRIES:
|
| 1383 |
+
print("Below threshold on the final attempt — no retry left, skipping feedback construction.")
|
| 1384 |
+
else:
|
| 1385 |
+
# feedback is now ONE unified sentence describing what's wrong and what needs to
|
| 1386 |
+
# change, instead of pasting plausibility_notes/faithfulness_notes/quality_notes
|
| 1387 |
+
# together as three separately-labeled clauses. The three _notes fields (and the
|
| 1388 |
+
# three scores/thresholds) are UNCHANGED and still drive stopping/freeze logic —
|
| 1389 |
+
# this only changes what prose text gets sent back to the generator as feedback.
|
| 1390 |
+
# overall_verdict is the judge's own synthesized summary (already part of the
|
| 1391 |
+
# schema, previously unused for feedback) — using that instead of splicing notes
|
| 1392 |
+
# avoids asking the model to write three separate critiques when one coherent one
|
| 1393 |
+
# is what the generator actually needs to act on.
|
| 1394 |
+
feedback = verdict.get(
|
| 1395 |
+
"overall_verdict",
|
| 1396 |
+
verdict.get("plausibility_notes", "the pose progression needs to look more plausible"))
|
| 1397 |
+
|
| 1398 |
+
# DINOv2 objective override: if the images barely changed from the
|
| 1399 |
+
# previous attempt, say so explicitly and forcefully — this is the
|
| 1400 |
+
# exact failure mode confirmed on real hardware (cannon1 ran all
|
| 1401 |
+
# MAX_RETRIES with no real change), where the judge's own text
|
| 1402 |
+
# critique was never specific enough for Qwen to act on. An
|
| 1403 |
+
# objective embedding-distance number doesn't have that problem.
|
| 1404 |
+
if stagnation_result and dino_similarity.is_stagnant(stagnation_result):
|
| 1405 |
+
feedback = (f"CRITICAL: your last attempt was measured as nearly IDENTICAL to the one "
|
| 1406 |
+
f"before it (DINOv2 similarity {stagnation_result['mean_similarity']:.3f}) — "
|
| 1407 |
+
f"you are NOT making real changes. You MUST produce substantially different "
|
| 1408 |
+
f"target positions this time, not a superficial rewording. Original feedback: {feedback}")
|
| 1409 |
+
|
| 1410 |
+
print(f"Below threshold — retrying with feedback: {feedback!r} "
|
| 1411 |
+
f"(narrative will be {'FROZEN' if freeze_narrative else 'regenerated'})")
|
| 1412 |
+
# the line above only ever showed the flat plausibility/faithfulness/quality string —
|
| 1413 |
+
# joint_feedback is a SEPARATE variable also being carried into the next prompt (see
|
| 1414 |
+
# build_combined_narrate_deform_prompt's joint_feedback= argument), and was never
|
| 1415 |
+
# visible in this log even when it had real content. Printing it explicitly here so
|
| 1416 |
+
# it's possible to verify from the log whether per-joint correction is actually
|
| 1417 |
+
# happening on a given attempt, instead of having to infer it from final results.
|
| 1418 |
+
if joint_feedback:
|
| 1419 |
+
grounded = [f for f in joint_feedback if f.get("delta_px") is not None]
|
| 1420 |
+
ungrounded = [f for f in joint_feedback if f.get("delta_px") is None]
|
| 1421 |
+
print(f" joint_feedback being sent to next attempt: {len(grounded)} grounded "
|
| 1422 |
+
f"(with computed delta_px), {len(ungrounded)} ungrounded")
|
| 1423 |
+
for f in joint_feedback:
|
| 1424 |
+
tag = "GROUNDED" if f.get("delta_px") is not None else "UNGROUNDED"
|
| 1425 |
+
print(f" [{tag}] {f.get('object')}.joint_{f.get('joint')} @ kf{f.get('keyframe')}: "
|
| 1426 |
+
f"{f.get('issue')} (delta_px={f.get('delta_px')})")
|
| 1427 |
+
else:
|
| 1428 |
+
print(" joint_feedback being sent to next attempt: none (empty or not returned by judge)")
|
| 1429 |
+
|
| 1430 |
+
else:
|
| 1431 |
+
print(f"\nReached MAX_RETRIES ({MAX_RETRIES}) without meeting both thresholds. "
|
| 1432 |
+
f"Using the last attempt's result.")
|
| 1433 |
+
|
| 1434 |
+
# print a compact table so the score progression across attempts is
|
| 1435 |
+
# visible in one place, not just scattered through the full log
|
| 1436 |
+
print("\n########## ATTEMPT SUMMARY ##########")
|
| 1437 |
+
for a in all_attempts_summary:
|
| 1438 |
+
print(f" attempt {a['attempt']}: plausibility={a.get('plausibility_score')} "
|
| 1439 |
+
f"faithfulness={a.get('faithfulness_score')} quality={a.get('quality_score')} "
|
| 1440 |
+
f"dino_stagnant={a.get('dino_stagnant')} "
|
| 1441 |
+
f"— {a.get('plausibility_notes') or a.get('note', '')}")
|
| 1442 |
+
|
| 1443 |
+
if winning_attempt:
|
| 1444 |
+
import shutil
|
| 1445 |
+
winning_json = os.path.join(json_dir, "attempts", f"attempt_{winning_attempt}")
|
| 1446 |
+
winning_temp = os.path.join(temp_dir, "attempts", f"attempt_{winning_attempt}")
|
| 1447 |
+
for f in os.listdir(winning_json):
|
| 1448 |
+
shutil.copy2(os.path.join(winning_json, f), os.path.join(json_dir, f))
|
| 1449 |
+
for f in os.listdir(winning_temp):
|
| 1450 |
+
src = os.path.join(winning_temp, f)
|
| 1451 |
+
if os.path.isfile(src):
|
| 1452 |
+
shutil.copy2(src, os.path.join(temp_dir, f))
|
| 1453 |
+
print(f"\ncopied winning attempt ({winning_attempt}) to the top-level "
|
| 1454 |
+
f"json/{name}/ and P_1/{name}/ locations")
|
| 1455 |
+
print(f"all {len(all_attempts_summary)} attempts preserved under "
|
| 1456 |
+
f"json/{name}/attempts/ and P_1/{name}/attempts/ for comparison")
|
| 1457 |
+
else:
|
| 1458 |
+
print("\nNo ARAP objects — skipping narrate/deform/judge entirely.")
|
| 1459 |
+
temp_dir = os.path.join(args.out_dir, "P_1", name)
|
| 1460 |
+
print("\n########## RENDER (no Qwen) ##########")
|
| 1461 |
+
run_render(json_dir, svg_path, semantic_path, traj_path, frames_dir=temp_dir)
|
| 1462 |
+
|
| 1463 |
+
|
| 1464 |
+
if __name__ == "__main__":
|
| 1465 |
+
main()
|
Downloads/sketch_pipeline(1).pptx
ADDED
|
Binary file (94.6 kB). View file
|
|
|
Downloads/sketch_pipeline.pptx
ADDED
|
Binary file (80.4 kB). View file
|
|
|
LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
MIT License
|
| 2 |
+
|
| 3 |
+
Copyright (c) 2025 Jingyu Liu
|
| 4 |
+
|
| 5 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 6 |
+
of this software and associated documentation files (the "Software"), to deal
|
| 7 |
+
in the Software without restriction, including without limitation the rights
|
| 8 |
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 9 |
+
copies of the Software, and to permit persons to whom the Software is
|
| 10 |
+
furnished to do so, subject to the following conditions:
|
| 11 |
+
|
| 12 |
+
The above copyright notice and this permission notice shall be included in all
|
| 13 |
+
copies or substantial portions of the Software.
|
| 14 |
+
|
| 15 |
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 16 |
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 17 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 18 |
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 19 |
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 20 |
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 21 |
+
SOFTWARE.
|
P_1/band6/attempts/attempt_1/band6.png
ADDED
|
Git LFS Details
|
P_1/band6/attempts/attempt_1/kf0.png
ADDED
|
P_1/band6/attempts/attempt_1/kf1.png
ADDED
|
P_1/band6/attempts/attempt_1/kf2.png
ADDED
|
P_1/band6/attempts/attempt_1/kf3.png
ADDED
|
P_1/band6/attempts/attempt_1/kf4.png
ADDED
|
P_1/band6/attempts/attempt_2/band6.png
ADDED
|
Git LFS Details
|
P_1/band6/attempts/attempt_2/kf0.png
ADDED
|
P_1/band6/attempts/attempt_2/kf1.png
ADDED
|
P_1/band6/attempts/attempt_2/kf2.png
ADDED
|
P_1/band6/attempts/attempt_2/kf3.png
ADDED
|
P_1/band6/attempts/attempt_2/kf4.png
ADDED
|
P_1/band6/band6.png
ADDED
|
Git LFS Details
|
P_1/band6/kf0.png
ADDED
|
P_1/band6/kf1.png
ADDED
|
P_1/band6/kf2.png
ADDED
|
P_1/band6/kf3.png
ADDED
|
P_1/band6/kf4.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_1/band6.png
ADDED
|
Git LFS Details
|
P_1/band6__20260826_144524/attempts/attempt_1/kf0.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_1/kf1.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_1/kf2.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_1/kf3.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_1/kf4.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_2/band6.png
ADDED
|
Git LFS Details
|
P_1/band6__20260826_144524/attempts/attempt_2/kf0.png
ADDED
|
P_1/band6__20260826_144524/attempts/attempt_2/kf1.png
ADDED
|