reshma0639 commited on
Commit
179b892
·
verified ·
1 Parent(s): ee5a706

Upload folder using huggingface_hub

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +42 -0
  2. =0.46.1 +25 -0
  3. Downloads/.ipynb_checkpoints/Convolutional Neural Networks Transfer Learning-checkpoint.ipynb +990 -0
  4. Downloads/.ipynb_checkpoints/EEEM071_CourseWork_ipynb-checkpoint.ipynb +379 -0
  5. Downloads/.ipynb_checkpoints/Human Action Recognition Tutorial-checkpoint.ipynb +820 -0
  6. Downloads/.ipynb_checkpoints/PyTorch Tutorial-checkpoint.ipynb +2030 -0
  7. Downloads/.ipynb_checkpoints/Python Tutorial(1)-checkpoint.ipynb +3313 -0
  8. Downloads/.ipynb_checkpoints/Python Tutorial-checkpoint.ipynb +3313 -0
  9. Downloads/.~lock.deformation_experiments(3).pptx# +1 -0
  10. Downloads/HjxNnnnu.html +0 -0
  11. Downloads/deformation_experiments(1).pptx +3 -0
  12. Downloads/deformation_experiments(2).pptx +3 -0
  13. Downloads/deformation_experiments(3).pptx +3 -0
  14. Downloads/deformation_experiments.pptx +3 -0
  15. Downloads/gap_zoomed_comparison.png +0 -0
  16. Downloads/geometric_solver(1).py +1147 -0
  17. Downloads/handover_summary(1).md +220 -0
  18. Downloads/handover_summary.md +220 -0
  19. Downloads/lbs_seam_constrained(1).py +222 -0
  20. Downloads/pipe_3.py +1465 -0
  21. Downloads/sketch_pipeline(1).pptx +0 -0
  22. Downloads/sketch_pipeline.pptx +0 -0
  23. LICENSE +21 -0
  24. P_1/band6/attempts/attempt_1/band6.png +3 -0
  25. P_1/band6/attempts/attempt_1/kf0.png +0 -0
  26. P_1/band6/attempts/attempt_1/kf1.png +0 -0
  27. P_1/band6/attempts/attempt_1/kf2.png +0 -0
  28. P_1/band6/attempts/attempt_1/kf3.png +0 -0
  29. P_1/band6/attempts/attempt_1/kf4.png +0 -0
  30. P_1/band6/attempts/attempt_2/band6.png +3 -0
  31. P_1/band6/attempts/attempt_2/kf0.png +0 -0
  32. P_1/band6/attempts/attempt_2/kf1.png +0 -0
  33. P_1/band6/attempts/attempt_2/kf2.png +0 -0
  34. P_1/band6/attempts/attempt_2/kf3.png +0 -0
  35. P_1/band6/attempts/attempt_2/kf4.png +0 -0
  36. P_1/band6/band6.png +3 -0
  37. P_1/band6/kf0.png +0 -0
  38. P_1/band6/kf1.png +0 -0
  39. P_1/band6/kf2.png +0 -0
  40. P_1/band6/kf3.png +0 -0
  41. P_1/band6/kf4.png +0 -0
  42. P_1/band6__20260826_144524/attempts/attempt_1/band6.png +3 -0
  43. P_1/band6__20260826_144524/attempts/attempt_1/kf0.png +0 -0
  44. P_1/band6__20260826_144524/attempts/attempt_1/kf1.png +0 -0
  45. P_1/band6__20260826_144524/attempts/attempt_1/kf2.png +0 -0
  46. P_1/band6__20260826_144524/attempts/attempt_1/kf3.png +0 -0
  47. P_1/band6__20260826_144524/attempts/attempt_1/kf4.png +0 -0
  48. P_1/band6__20260826_144524/attempts/attempt_2/band6.png +3 -0
  49. P_1/band6__20260826_144524/attempts/attempt_2/kf0.png +0 -0
  50. 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
+ "![TU-Berlin](https://cybertron.cg.tu-berlin.de/eitz/projects/classifysketch/teaser_siggraph.jpg)\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
+ "![box filter](https://i2.wp.com/theailearner.com/wp-content/uploads/2019/05/filter1.png?w=454&ssl=1)\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
+ "![box filter](https://i2.wp.com/theailearner.com/wp-content/uploads/2019/05/filter1.png?w=454&ssl=1)\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

  • SHA256: d94b6b043877a368d61ea86fe39c492b138b3d5a3707aa878fce2ca868fa7fc7
  • Pointer size: 131 Bytes
  • Size of remote file: 311 kB
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

  • SHA256: d4fe6b68f2fdecdf99228b14c81ff67c2b434940f10cf940e203e30cc8079b8f
  • Pointer size: 131 Bytes
  • Size of remote file: 312 kB
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

  • SHA256: d94b6b043877a368d61ea86fe39c492b138b3d5a3707aa878fce2ca868fa7fc7
  • Pointer size: 131 Bytes
  • Size of remote file: 311 kB
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

  • SHA256: f58bae0e692aea079d30a0a33655bd6a9cfda62a1b2567043b6a0cd0722ae7e6
  • Pointer size: 131 Bytes
  • Size of remote file: 347 kB
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

  • SHA256: f58bae0e692aea079d30a0a33655bd6a9cfda62a1b2567043b6a0cd0722ae7e6
  • Pointer size: 131 Bytes
  • Size of remote file: 347 kB
P_1/band6__20260826_144524/attempts/attempt_2/kf0.png ADDED
P_1/band6__20260826_144524/attempts/attempt_2/kf1.png ADDED