Title: 가지치기 기법(Pruning) 튜토리얼 — 파이토치 한국어 튜토리얼 (PyTorch tutorials in Korean)
Open Graph Title: 가지치기 기법(Pruning) 튜토리얼
Description: Author: Michela Paganini, 번역: 안상준,. 최첨단 딥러닝 모델들은 굉장히 많은 수의 파라미터값들로 구성되기 때문에, 쉽게 배포하기가 어렵습니다. 이와 반대로, 생물학적 신경망들은 효율적으로 희소하게 연결된 것으로 알려져 있습니다. 모델의 정확도를 훼손하지 않으면서 모델에 포함된 파라미터 수를 줄여 압축하는 최적의 기법을 파악하는 것은 메모리, 배터리, 하드웨어 소비량을 줄일 수 있기 때문에 중요합니다. 그럼으로서 기기에 경량화된 모델을 배포하여 개개인이 사용하고 있는 기기에서 연산을 수행하여 프라이버시를 ...
Open Graph Description: Author: Michela Paganini, 번역: 안상준,. 최첨단 딥러닝 모델들은 굉장히 많은 수의 파라미터값들로 구성되기 때문에, 쉽게 배포하기가 어렵습니다. 이와 반대로, 생물학적 신경망들은 효율적으로 희소하게 연결된 것으로 알려져 있습니다. 모델의 정확도를 훼손하지 않으면서 모델에 포함된 파라미터 수를 줄여 압축하는 최적의 기법을 파악하는 것은 메모리, 배터리, 하드웨어 소비량을 줄일 수 있기 때문에 중요합니다. 그럼으로서 기기에 경량화된 모델을 배포하여 개개인이 사용하고 있는 기기에서 연산을 수행하여 프라이버시를 ...
Opengraph URL: https://tutorials.pytorch.kr/intermediate/pruning_tutorial.html
Domain: tutorials.pytorch.kr
{
"@context": "https://schema.org",
"@type": "Article",
"name": "\uac00\uc9c0\uce58\uae30 \uae30\ubc95(Pruning) \ud29c\ud1a0\ub9ac\uc5bc",
"headline": "\uac00\uc9c0\uce58\uae30 \uae30\ubc95(Pruning) \ud29c\ud1a0\ub9ac\uc5bc",
"description": "PyTorch Documentation. Explore PyTorch, an open-source machine learning library that accelerates the path from research prototyping to production deployment. Discover tutorials, API references, and guides to help you build and deploy deep learning models efficiently.",
"url": "/intermediate/pruning_tutorial.html",
"articleBody": "\ucc38\uace0 Go to the end to download the full example code. \uac00\uc9c0\uce58\uae30 \uae30\ubc95(Pruning) \ud29c\ud1a0\ub9ac\uc5bc# Author: Michela Paganini\ubc88\uc5ed: \uc548\uc0c1\uc900 \ucd5c\ucca8\ub2e8 \ub525\ub7ec\ub2dd \ubaa8\ub378\ub4e4\uc740 \uad49\uc7a5\ud788 \ub9ce\uc740 \uc218\uc758 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\ub85c \uad6c\uc131\ub418\uae30 \ub54c\ubb38\uc5d0, \uc27d\uac8c \ubc30\ud3ec\ud558\uae30\uac00 \uc5b4\ub835\uc2b5\ub2c8\ub2e4. \uc774\uc640 \ubc18\ub300\ub85c, \uc0dd\ubb3c\ud559\uc801 \uc2e0\uacbd\ub9dd\ub4e4\uc740 \ud6a8\uc728\uc801\uc73c\ub85c \ud76c\uc18c\ud558\uac8c \uc5f0\uacb0\ub41c \uac83\uc73c\ub85c \uc54c\ub824\uc838 \uc788\uc2b5\ub2c8\ub2e4. \ubaa8\ub378\uc758 \uc815\ud655\ub3c4\ub97c \ud6fc\uc190\ud558\uc9c0 \uc54a\uc73c\uba74\uc11c \ubaa8\ub378\uc5d0 \ud3ec\ud568\ub41c \ud30c\ub77c\ubbf8\ud130 \uc218\ub97c \uc904\uc5ec \uc555\ucd95\ud558\ub294 \ucd5c\uc801\uc758 \uae30\ubc95\uc744 \ud30c\uc545\ud558\ub294 \uac83\uc740 \uba54\ubaa8\ub9ac, \ubc30\ud130\ub9ac, \ud558\ub4dc\uc6e8\uc5b4 \uc18c\ube44\ub7c9\uc744 \uc904\uc77c \uc218 \uc788\uae30 \ub54c\ubb38\uc5d0 \uc911\uc694\ud569\ub2c8\ub2e4. \uadf8\ub7fc\uc73c\ub85c\uc11c \uae30\uae30\uc5d0 \uacbd\ub7c9\ud654\ub41c \ubaa8\ub378\uc744 \ubc30\ud3ec\ud558\uc5ec \uac1c\uac1c\uc778\uc774 \uc0ac\uc6a9\ud558\uace0 \uc788\ub294 \uae30\uae30\uc5d0\uc11c \uc5f0\uc0b0\uc744 \uc218\ud589\ud558\uc5ec \ud504\ub77c\uc774\ubc84\uc2dc\ub97c \ubcf4\uc7a5\ud560 \uc218 \uc788\uae30 \ub54c\ubb38\uc785\ub2c8\ub2e4. \uc5f0\uad6c \uce21\uba74\uc5d0\uc11c\ub294, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc740 \uad49\uc7a5\ud788 \ub9ce\uc740 \uc218\uc758 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\ub85c \uad6c\uc131\ub41c \ubaa8\ub378\uacfc \uad49\uc7a5\ud788 \uc801\uc740 \uc218\uc758 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\ub85c \uad6c\uc131\ub41c \ubaa8\ub378 \uac04 \ud559\uc2b5 \uc5ed\ud559 \ucc28\uc774\ub97c \uc870\uc0ac\ud558\ub294\ub370 \uc8fc\ub85c \uc774\uc6a9\ub418\uae30\ub3c4 \ud558\uba70, \ud558\uc704 \uc2e0\uacbd\ub9dd \ubaa8\ub378\uacfc \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\uc758 \ucd08\uae30\ud654\uac00 \uc6b4\uc774 \uc88b\uac8c \uc798 \ub41c \ucf00\uc774\uc2a4\ub97c \ubc14\ud0d5\uc73c\ub85c (\u201dlottery tickets\u201d) \uc2e0\uacbd\ub9dd \uad6c\uc870\ub97c \ucc3e\ub294 \uae30\uc220\ub4e4\uc5d0 \ub300\ud574 \ubc18\ub300 \uc758\uacac\uc744 \uc81c\uc2dc\ud558\uae30\ub3c4 \ud569\ub2c8\ub2e4. \uc774\ubc88 \ud29c\ud1a0\ub9ac\uc5bc\uc5d0\uc11c\ub294, torch.nn.utils.prune \uc744 \uc0ac\uc6a9\ud558\uc5ec \uc5ec\ub7ec\ubd84\uc774 \uc124\uacc4\ud55c \ub525\ub7ec\ub2dd \ubaa8\ub378\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud574\ubcf4\ub294 \uac83\uc744 \ubc30\uc6cc\ubcf4\uace0, \uc2ec\ud654\uc801\uc73c\ub85c \uc5ec\ub7ec\ubd84\uc758 \ub9de\ucda4\ud615 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uad6c\ud604\ud558\ub294 \ubc29\ubc95\uc5d0 \ub300\ud574 \ubc30\uc6cc\ubcf4\ub3c4\ub85d \ud558\uaca0\uc2b5\ub2c8\ub2e4. \uc694\uad6c\uc0ac\ud56d# \"torch\u003e=1.4\" import torch from torch import nn import torch.nn.utils.prune as prune import torch.nn.functional as F \ub525\ub7ec\ub2dd \ubaa8\ub378 \uc0dd\uc131# \uc774\ubc88 \ud29c\ud1a0\ub9ac\uc5bc\uc5d0\uc11c\ub294, \uc580 \ub974\ucfe4 \uad50\uc218\ub2d8\uc758 \uc5f0\uad6c\uc9c4\ub4e4\uc774 1998\ub144\ub3c4\uc5d0 \ubc1c\ud45c\ud55c LeNet \uc758 \ubaa8\ub378 \uad6c\uc870\ub97c \uc774\uc6a9\ud569\ub2c8\ub2e4. device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\") class LeNet(nn.Module): def __init__(self): super(LeNet, self).__init__() # 1\uac1c \ucc44\ub110 \uc218\uc758 \uc774\ubbf8\uc9c0\ub97c \uc785\ub825\uac12\uc73c\ub85c \uc774\uc6a9\ud558\uc5ec 6\uac1c \ucc44\ub110 \uc218\uc758 \ucd9c\ub825\uac12\uc744 \uacc4\uc0b0\ud558\ub294 \ubc29\uc2dd # Convolution \uc5f0\uc0b0\uc744 \uc9c4\ud589\ud558\ub294 \ucee4\ub110(\ud544\ud130)\uc758 \ud06c\uae30\ub294 5x5 \uc744 \uc774\uc6a9 self.conv1 = nn.Conv2d(1, 6, 5) self.conv2 = nn.Conv2d(6, 16, 5) self.fc1 = nn.Linear(16 * 5 * 5, 120) # Convolution \uc5f0\uc0b0 \uacb0\uacfc 5x5 \ud06c\uae30\uc758 16 \ucc44\ub110 \uc218\uc758 \uc774\ubbf8\uc9c0 self.fc2 = nn.Linear(120, 84) self.fc3 = nn.Linear(84, 10) def forward(self, x): x = F.max_pool2d(F.relu(self.conv1(x)), (2, 2)) x = F.max_pool2d(F.relu(self.conv2(x)), 2) x = x.view(-1, int(x.nelement() / x.shape[0])) x = F.relu(self.fc1(x)) x = F.relu(self.fc2(x)) x = self.fc3(x) return x model = LeNet().to(device=device) \ubaa8\ub4c8 \uc810\uac80# \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc9c0 \uc54a\uc740 LeNet \ubaa8\ub378\uc758 conv1 \uce35\uc744 \uc810\uac80\ud574\ubd05\uc2dc\ub2e4. \uc5ec\uae30\uc5d0\ub294 2\uac1c\uc758 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\uc778 \uac00\uc911\uce58 \uac12\uacfc \ud3b8\ud5a5 \uac12\uc774 \ud3ec\ud568\ub420 \uac83\uc774\uba70, \ubc84\ud37c\ub294 \uc874\uc7ac\ud558\uc9c0 \uc54a\uc744 \uac83\uc785\ub2c8\ub2e4. module = model.conv1 print(list(module.named_parameters())) [(\u0027weight\u0027, Parameter containing: tensor([[[[ 0.1232, -0.0642, -0.0623, -0.1485, 0.0055], [ 0.0314, -0.1857, -0.0403, 0.1223, 0.1689], [-0.0559, -0.1542, 0.0177, 0.0988, 0.0749], [ 0.0634, -0.1869, 0.1766, 0.1120, 0.0219], [-0.1110, -0.0066, -0.0742, 0.1935, -0.1882]]], [[[-0.1891, 0.0554, -0.1982, 0.1154, -0.0523], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0445], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0206, -0.1911, -0.0794, 0.1364, 0.0037], [ 0.0427, -0.1558, -0.1073, -0.0763, 0.0211]]], [[[ 0.0116, -0.1011, -0.1792, -0.0166, -0.1940], [-0.1058, -0.0902, -0.0587, 0.0361, -0.0123], [-0.1890, -0.0632, 0.0668, -0.0883, -0.1008], [-0.0702, 0.1404, 0.0646, -0.1084, -0.0797], [-0.0942, -0.0567, -0.1763, 0.0473, -0.1682]]], [[[-0.0958, 0.0936, 0.1754, -0.0095, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0554, 0.0111, 0.1206, 0.0263], [ 0.1758, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0716, 0.0628, 0.0945]]], [[[ 0.1180, -0.0116, 0.1336, -0.0599, -0.0110], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0822], [-0.1528, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.1081, 0.0753, -0.0684], [ 0.0304, 0.0038, -0.0709, 0.0481, -0.0312]]], [[[ 0.0692, -0.1867, -0.0930, -0.0373, -0.1380], [-0.0196, 0.1388, 0.0801, -0.1948, 0.0013], [ 0.0478, -0.1248, -0.0969, -0.1181, 0.1294], [ 0.0343, -0.0799, -0.0200, -0.1351, -0.0577], [-0.0725, 0.0144, -0.0758, 0.0333, 0.0219]]]], device=\u0027cuda:0\u0027, requires_grad=True)), (\u0027bias\u0027, Parameter containing: tensor([-0.0229, -0.0985, -0.0547, -0.0255, -0.1460, -0.1458], device=\u0027cuda:0\u0027, requires_grad=True))] print(list(module.named_buffers())) [] \ubaa8\ub4c8 \uac00\uc9c0\uce58\uae30 \uae30\ubc95 \uc801\uc6a9 \uc608\uc81c# \ubaa8\ub4c8\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uae30 \uc704\ud574 (\uc774\ubc88 \uc608\uc81c\uc5d0\uc11c\ub294, LeNet \ubaa8\ub378\uc758 conv1 \uce35) \uccab \ubc88\uc9f8\ub85c\ub294, torch.nn.utils.prune (\ub610\ub294 BasePruningMethod \uc758 \uc11c\ube0c \ud074\ub798\uc2a4\ub85c \uc9c1\uc811 \uad6c\ud604 ) \ub0b4 \uc874\uc7ac\ud558\ub294 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc120\ud0dd\ud569\ub2c8\ub2e4. \uadf8 \ud6c4, \ud574\ub2f9 \ubaa8\ub4c8 \ub0b4\uc5d0\uc11c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uace0\uc790 \ud558\ub294 \ubaa8\ub4c8\uacfc \ud30c\ub77c\ubbf8\ud130\ub97c \uc9c0\uc815\ud569\ub2c8\ub2e4. \ub9c8\uc9c0\ub9c9\uc73c\ub85c, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc5d0 \uc801\ub2f9\ud55c \ud0a4\uc6cc\ub4dc \uc778\uc790\uac12\uc744 \uc774\uc6a9\ud558\uc5ec \uac00\uc9c0\uce58\uae30 \ub9e4\uac1c\ubcc0\uc218\ub97c \uc9c0\uc815\ud569\ub2c8\ub2e4. \uc774\ubc88 \uc608\uc81c\uc5d0\uc11c\ub294, conv1 \uce35\uc758 \uac00\uc911\uce58\uc758 30%\uac12\ub4e4\uc744 \ub79c\ub364\uc73c\ub85c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud574\ubcf4\uaca0\uc2b5\ub2c8\ub2e4. \ubaa8\ub4c8\uc740 \ud568\uc218\uc5d0 \ub300\ud55c \uccab \ubc88\uc9f8 \uc778\uc790\uac12\uc73c\ub85c \uc804\ub2ec\ub418\uba70, name \uc740 \ubb38\uc790\uc5f4 \uc2dd\ubcc4\uc790\ub97c \uc774\uc6a9\ud558\uc5ec \ud574\ub2f9 \ubaa8\ub4c8 \ub0b4 \ub9e4\uac1c\ubcc0\uc218\ub97c \uad6c\ubd84\ud569\ub2c8\ub2e4. \uadf8\ub9ac\uace0, amount \ub294 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uae30 \uc704\ud55c \ub300\uc0c1 \uac00\uc911\uce58\uac12\ub4e4\uc758 \ubc31\ubd84\uc728 (0\uacfc 1\uc0ac\uc774\uc758 \uc2e4\uc218\uac12), \ud639\uc740 \uac00\uc911\uce58\uac12\uc758 \uc5f0\uacb0\uc758 \uac1c\uc218 (\uc74c\uc218\uac00 \uc544\ub2cc \uc815\uc218) \ub97c \uc9c0\uc815\ud569\ub2c8\ub2e4. prune.random_unstructured(module, name=\"weight\", amount=0.3) Conv2d(1, 6, kernel_size=(5, 5), stride=(1, 1)) \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc740 \uac00\uc911\uce58\uac12\ub4e4\uc744 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\ub85c\ubd80\ud130 \uc81c\uac70\ud558\uace0 weight_orig (\uc989, \ucd08\uae30 \uac00\uc911\uce58 \uc774\ub984\uc5d0 \u201c_orig\u201d\uc744 \ubd99\uc778) \uc774\ub77c\ub294 \uc0c8\ub85c\uc6b4 \ud30c\ub77c\ubbf8\ud130\uac12\uc73c\ub85c \ub300\uccb4\ud558\ub294 \uac83\uc73c\ub85c \uc2e4\ud589\ub429\ub2c8\ub2e4. weight_orig \uc740 \ud150\uc11c\uac12\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc9c0 \uc54a\uc740 \uc0c1\ud0dc\ub97c \uc800\uc7a5\ud569\ub2c8\ub2e4. bias \uc740 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc9c0 \uc54a\uc558\uae30 \ub54c\ubb38\uc5d0 \uadf8\ub300\ub85c \ub0a8\uc544 \uc788\uc2b5\ub2c8\ub2e4. print(list(module.named_parameters())) [(\u0027bias\u0027, Parameter containing: tensor([-0.0229, -0.0985, -0.0547, -0.0255, -0.1460, -0.1458], device=\u0027cuda:0\u0027, requires_grad=True)), (\u0027weight_orig\u0027, Parameter containing: tensor([[[[ 0.1232, -0.0642, -0.0623, -0.1485, 0.0055], [ 0.0314, -0.1857, -0.0403, 0.1223, 0.1689], [-0.0559, -0.1542, 0.0177, 0.0988, 0.0749], [ 0.0634, -0.1869, 0.1766, 0.1120, 0.0219], [-0.1110, -0.0066, -0.0742, 0.1935, -0.1882]]], [[[-0.1891, 0.0554, -0.1982, 0.1154, -0.0523], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0445], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0206, -0.1911, -0.0794, 0.1364, 0.0037], [ 0.0427, -0.1558, -0.1073, -0.0763, 0.0211]]], [[[ 0.0116, -0.1011, -0.1792, -0.0166, -0.1940], [-0.1058, -0.0902, -0.0587, 0.0361, -0.0123], [-0.1890, -0.0632, 0.0668, -0.0883, -0.1008], [-0.0702, 0.1404, 0.0646, -0.1084, -0.0797], [-0.0942, -0.0567, -0.1763, 0.0473, -0.1682]]], [[[-0.0958, 0.0936, 0.1754, -0.0095, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0554, 0.0111, 0.1206, 0.0263], [ 0.1758, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0716, 0.0628, 0.0945]]], [[[ 0.1180, -0.0116, 0.1336, -0.0599, -0.0110], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0822], [-0.1528, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.1081, 0.0753, -0.0684], [ 0.0304, 0.0038, -0.0709, 0.0481, -0.0312]]], [[[ 0.0692, -0.1867, -0.0930, -0.0373, -0.1380], [-0.0196, 0.1388, 0.0801, -0.1948, 0.0013], [ 0.0478, -0.1248, -0.0969, -0.1181, 0.1294], [ 0.0343, -0.0799, -0.0200, -0.1351, -0.0577], [-0.0725, 0.0144, -0.0758, 0.0333, 0.0219]]]], device=\u0027cuda:0\u0027, requires_grad=True))] \uc704\uc5d0\uc11c \uc120\ud0dd\ud55c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc5d0 \uc758\ud574 \uc0dd\uc131\ub418\ub294 \uac00\uc9c0\uce58\uae30 \ub9c8\uc2a4\ud06c\ub294 \ucd08\uae30 \ud30c\ub77c\ubbf8\ud130 name \uc5d0 weight_mask (\uc989, \ucd08\uae30 \uac00\uc911\uce58 \uc774\ub984\uc5d0 \u201c_mask\u201d\ub97c \ubd99\uc778) \uc774\ub984\uc758 \ubaa8\ub4c8 \ubc84\ud37c\ub85c \uc800\uc7a5\ub429\ub2c8\ub2e4. print(list(module.named_buffers())) [(\u0027weight_mask\u0027, tensor([[[[0., 0., 1., 0., 1.], [1., 1., 0., 1., 1.], [0., 0., 1., 1., 1.], [1., 0., 1., 0., 1.], [0., 0., 1., 1., 1.]]], [[[1., 1., 0., 1., 0.], [1., 1., 1., 1., 0.], [1., 1., 1., 1., 1.], [0., 1., 1., 0., 0.], [1., 1., 1., 0., 1.]]], [[[1., 1., 1., 1., 0.], [1., 1., 1., 1., 0.], [1., 1., 1., 0., 1.], [1., 1., 0., 0., 0.], [1., 0., 1., 0., 0.]]], [[[1., 1., 1., 0., 1.], [1., 1., 1., 1., 1.], [1., 0., 1., 1., 0.], [0., 1., 1., 1., 1.], [1., 1., 0., 1., 0.]]], [[[1., 0., 1., 1., 0.], [1., 1., 1., 1., 0.], [0., 1., 1., 1., 1.], [1., 1., 0., 1., 1.], [1., 0., 1., 1., 0.]]], [[[0., 1., 1., 1., 1.], [1., 1., 0., 1., 1.], [1., 0., 1., 1., 1.], [1., 1., 0., 0., 0.], [1., 1., 1., 1., 1.]]]], device=\u0027cuda:0\u0027))] \uc218\uc815\uc774 \ub418\uc9c0 \uc54a\uc740 \uc0c1\ud0dc\uc5d0\uc11c \uc21c\uc804\ud30c\ub97c \uc9c4\ud589\ud558\uae30 \uc704\ud574\uc11c\ub294 \uac00\uc911\uce58 \uac12 \uc18d\uc131\uc774 \uc874\uc7ac\ud574\uc57c \ud569\ub2c8\ub2e4. torch.nn.utils.prune \ub0b4 \uad6c\ud604\ub41c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc740 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uac00\uc911\uce58\uac12\ub4e4\uc744 \uc774\uc6a9\ud558\uc5ec (\uae30\uc874\uc758 \uac00\uc911\uce58\uac12\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c) \uc21c\uc804\ud30c\ub97c \uc9c4\ud589\ud558\uace0, weight \uc18d\uc131\uac12\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uac00\uc911\uce58\uac12\ub4e4\uc744 \uc800\uc7a5\ud569\ub2c8\ub2e4. \uc774\uc81c \uac00\uc911\uce58\uac12\ub4e4\uc740 module \uc758 \ub9e4\uac1c\ubcc0\uc218\uac00 \uc544\ub2c8\ub77c \ud558\ub098\uc758 \uc18d\uc131\uac12\uc73c\ub85c \ucde8\uae09\ub418\ub294 \uc810\uc744 \uc8fc\uc758\ud558\uc138\uc694. print(module.weight) tensor([[[[ 0.0000, -0.0000, -0.0623, -0.0000, 0.0055], [ 0.0314, -0.1857, -0.0000, 0.1223, 0.1689], [-0.0000, -0.0000, 0.0177, 0.0988, 0.0749], [ 0.0634, -0.0000, 0.1766, 0.0000, 0.0219], [-0.0000, -0.0000, -0.0742, 0.1935, -0.1882]]], [[[-0.1891, 0.0554, -0.0000, 0.1154, -0.0000], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0000], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0000, -0.1911, -0.0794, 0.0000, 0.0000], [ 0.0427, -0.1558, -0.1073, -0.0000, 0.0211]]], [[[ 0.0116, -0.1011, -0.1792, -0.0166, -0.0000], [-0.1058, -0.0902, -0.0587, 0.0361, -0.0000], [-0.1890, -0.0632, 0.0668, -0.0000, -0.1008], [-0.0702, 0.1404, 0.0000, -0.0000, -0.0000], [-0.0942, -0.0000, -0.1763, 0.0000, -0.0000]]], [[[-0.0958, 0.0936, 0.1754, -0.0000, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0000, 0.0111, 0.1206, 0.0000], [ 0.0000, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0000, 0.0628, 0.0000]]], [[[ 0.1180, -0.0000, 0.1336, -0.0599, -0.0000], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0000], [-0.0000, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.0000, 0.0753, -0.0684], [ 0.0304, 0.0000, -0.0709, 0.0481, -0.0000]]], [[[ 0.0000, -0.1867, -0.0930, -0.0373, -0.1380], [-0.0196, 0.1388, 0.0000, -0.1948, 0.0013], [ 0.0478, -0.0000, -0.0969, -0.1181, 0.1294], [ 0.0343, -0.0799, -0.0000, -0.0000, -0.0000], [-0.0725, 0.0144, -0.0758, 0.0333, 0.0219]]]], device=\u0027cuda:0\u0027, grad_fn=\u003cMulBackward0\u003e) \ucd5c\uc885\uc801\uc73c\ub85c, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc740 \ud30c\uc774\ud1a0\uce58\uc758 forward_pre_hooks \ub97c \uc774\uc6a9\ud558\uc5ec \uac01 \uc21c\uc804\ud30c\uac00 \uc9c4\ud589\ub418\uae30 \uc804\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub429\ub2c8\ub2e4. \uad6c\uccb4\uc801\uc73c\ub85c, \uc9c0\uae08\uae4c\uc9c0 \uc9c4\ud589\ud55c \uac83 \ucc98\ub7fc, \ubaa8\ub4c8\uc774 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc5c8\uc744 \ub54c, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uac01 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\uc774 forward_pre_hook \ub97c \uc5bb\uac8c\ub429\ub2c8\ub2e4. \uc774\ub7ec\ud55c \uacbd\uc6b0, weight \uc774\ub984\uc778 \uae30\uc874 \ud30c\ub77c\ubbf8\ud130\uac12\uc5d0 \ub300\ud574\uc11c\ub9cc \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uc600\uae30 \ub54c\ubb38\uc5d0, \ud6c5\uc740 \uc624\uc9c1 1\uac1c\ub9cc \uc874\uc7ac\ud560 \uac83\uc785\ub2c8\ub2e4. print(module._forward_pre_hooks) OrderedDict([(0, \u003ctorch.nn.utils.prune.RandomUnstructured object at 0x7f0dab246050\u003e)]) \uc644\uacb0\uc131\uc744 \uc704\ud574, \ud3b8\ud5a5\uac12\uc5d0 \ub300\ud574\uc11c\ub3c4 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud560 \uc218 \uc788\uc73c\uba70, \ubaa8\ub4c8\uc758 \ud30c\ub77c\ubbf8\ud130, \ubc84\ud37c, \ud6c5, \uc18d\uc131\uac12\ub4e4\uc774 \uc5b4\ub5bb\uac8c \ubcc0\uacbd\ub418\ub294\uc9c0 \ud655\uc778\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \ub610 \ub2e4\ub978 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud574\ubcf4\uae30 \uc704\ud574, l1_unstructured \uac00\uc9c0\uce58\uae30 \ud568\uc218\uc5d0\uc11c \uad6c\ud604\ub41c \ub0b4\uc6a9\uacfc \uac19\uc774, L1 Norm \uac12\uc774 \uac00\uc7a5 \uc791\uc740 \ud3b8\ud5a5\uac12 3\uac1c\ub97c \uac00\uc9c0\uce58\uae30\ub97c \uc2dc\ub3c4\ud574\ubd05\uc2dc\ub2e4. prune.l1_unstructured(module, name=\"bias\", amount=3) Conv2d(1, 6, kernel_size=(5, 5), stride=(1, 1)) \uc774\uc804\uc5d0\uc11c \uc2e4\uc2b5\ud55c \ub0b4\uc6a9\uc744 \ud1a0\ub300\ub85c, \uba85\uba85\ub41c \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\uc774 weight_orig, bias_orig 2\uac1c\ub97c \ubaa8\ub450 \ud3ec\ud568\ud560 \uac83\uc774\ub77c \uc608\uc0c1\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \ubc84\ud37c\ub4e4\uc740 weight_mask, bias_mask 2\uac1c\ub97c \ud3ec\ud568\ud560 \uac83\uc785\ub2c8\ub2e4. \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c 2\uac1c\uc758 \ud150\uc11c\uac12\ub4e4\uc740 \ubaa8\ub4c8\uc758 \uc18d\uc131\uac12\uc73c\ub85c \uc874\uc7ac\ud560 \uac83\uc774\uba70, \ubaa8\ub4c8\uc740 2\uac1c\uc758 forward_pre_hooks \uc744 \uac16\uac8c \ub420 \uac83\uc785\ub2c8\ub2e4. print(list(module.named_parameters())) [(\u0027weight_orig\u0027, Parameter containing: tensor([[[[ 0.1232, -0.0642, -0.0623, -0.1485, 0.0055], [ 0.0314, -0.1857, -0.0403, 0.1223, 0.1689], [-0.0559, -0.1542, 0.0177, 0.0988, 0.0749], [ 0.0634, -0.1869, 0.1766, 0.1120, 0.0219], [-0.1110, -0.0066, -0.0742, 0.1935, -0.1882]]], [[[-0.1891, 0.0554, -0.1982, 0.1154, -0.0523], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0445], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0206, -0.1911, -0.0794, 0.1364, 0.0037], [ 0.0427, -0.1558, -0.1073, -0.0763, 0.0211]]], [[[ 0.0116, -0.1011, -0.1792, -0.0166, -0.1940], [-0.1058, -0.0902, -0.0587, 0.0361, -0.0123], [-0.1890, -0.0632, 0.0668, -0.0883, -0.1008], [-0.0702, 0.1404, 0.0646, -0.1084, -0.0797], [-0.0942, -0.0567, -0.1763, 0.0473, -0.1682]]], [[[-0.0958, 0.0936, 0.1754, -0.0095, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0554, 0.0111, 0.1206, 0.0263], [ 0.1758, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0716, 0.0628, 0.0945]]], [[[ 0.1180, -0.0116, 0.1336, -0.0599, -0.0110], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0822], [-0.1528, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.1081, 0.0753, -0.0684], [ 0.0304, 0.0038, -0.0709, 0.0481, -0.0312]]], [[[ 0.0692, -0.1867, -0.0930, -0.0373, -0.1380], [-0.0196, 0.1388, 0.0801, -0.1948, 0.0013], [ 0.0478, -0.1248, -0.0969, -0.1181, 0.1294], [ 0.0343, -0.0799, -0.0200, -0.1351, -0.0577], [-0.0725, 0.0144, -0.0758, 0.0333, 0.0219]]]], device=\u0027cuda:0\u0027, requires_grad=True)), (\u0027bias_orig\u0027, Parameter containing: tensor([-0.0229, -0.0985, -0.0547, -0.0255, -0.1460, -0.1458], device=\u0027cuda:0\u0027, requires_grad=True))] print(list(module.named_buffers())) [(\u0027weight_mask\u0027, tensor([[[[0., 0., 1., 0., 1.], [1., 1., 0., 1., 1.], [0., 0., 1., 1., 1.], [1., 0., 1., 0., 1.], [0., 0., 1., 1., 1.]]], [[[1., 1., 0., 1., 0.], [1., 1., 1., 1., 0.], [1., 1., 1., 1., 1.], [0., 1., 1., 0., 0.], [1., 1., 1., 0., 1.]]], [[[1., 1., 1., 1., 0.], [1., 1., 1., 1., 0.], [1., 1., 1., 0., 1.], [1., 1., 0., 0., 0.], [1., 0., 1., 0., 0.]]], [[[1., 1., 1., 0., 1.], [1., 1., 1., 1., 1.], [1., 0., 1., 1., 0.], [0., 1., 1., 1., 1.], [1., 1., 0., 1., 0.]]], [[[1., 0., 1., 1., 0.], [1., 1., 1., 1., 0.], [0., 1., 1., 1., 1.], [1., 1., 0., 1., 1.], [1., 0., 1., 1., 0.]]], [[[0., 1., 1., 1., 1.], [1., 1., 0., 1., 1.], [1., 0., 1., 1., 1.], [1., 1., 0., 0., 0.], [1., 1., 1., 1., 1.]]]], device=\u0027cuda:0\u0027)), (\u0027bias_mask\u0027, tensor([0., 1., 0., 0., 1., 1.], device=\u0027cuda:0\u0027))] print(module.bias) tensor([-0.0000, -0.0985, -0.0000, -0.0000, -0.1460, -0.1458], device=\u0027cuda:0\u0027, grad_fn=\u003cMulBackward0\u003e) print(module._forward_pre_hooks) OrderedDict([(0, \u003ctorch.nn.utils.prune.RandomUnstructured object at 0x7f0dab246050\u003e), (1, \u003ctorch.nn.utils.prune.L1Unstructured object at 0x7f0dab183890\u003e)]) \uac00\uc9c0\uce58\uae30 \uae30\ubc95 \ubc18\ubcf5 \uc801\uc6a9# \ubaa8\ub4c8 \ub0b4 \uac19\uc740 \ud30c\ub77c\ubbf8\ud130\uac12\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc5ec\ub7ec\ubc88 \uc801\uc6a9\ub420 \uc218 \uc788\uc73c\uba70, \ub2e4\uc591\ud55c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc758 \uc870\ud569\uc774 \uc801\uc6a9\ub41c \uac83\uacfc \ub3d9\uc77c\ud558\uac8c \uc801\uc6a9\ub420 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \uc0c8\ub85c\uc6b4 \ub9c8\uc2a4\ud06c\uc640 \uc774\uc804\uc758 \ub9c8\uc2a4\ud06c\uc758 \uacb0\ud569\uc740 PruningContainer \uc758 compute_mask \uba54\uc18c\ub4dc\ub97c \ud1b5\ud574 \ucc98\ub9ac\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \uc608\ub97c \ub4e4\uc5b4, \ub9cc\uc57d module.weight \uac12\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uace0 \uc2f6\uc744 \ub54c, \ud150\uc11c\uc758 0\ubc88\uc9f8 \ucd95\uc758 L2 norm\uac12\uc744 \uae30\uc900\uc73c\ub85c \uad6c\uc870\ud654\ub41c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud569\ub2c8\ub2e4. (\uc5ec\uae30\uc11c 0\ubc88\uc9f8 \ucd95\uc774\ub780, \ud569\uc131\uacf1 \uc5f0\uc0b0\uc744 \ud1b5\ud574 \uacc4\uc0b0\ub41c \ucd9c\ub825\uac12\uc5d0 \ub300\ud574 \uac01 \ucc44\ub110\ubcc4\ub85c \uc801\uc6a9\ub41c\ub2e4\ub294 \uac83\uc744 \uc758\ubbf8\ud569\ub2c8\ub2e4.) \uc774 \ubc29\uc2dd\uc740 ln_structured \ud568\uc218\uc640 n=2 \uc640 dim=0 \uc758 \uc778\uc790\uac12\uc744 \ubc14\ud0d5\uc73c\ub85c \uad6c\ud604\ub420 \uc218 \uc788\uc2b5\ub2c8\ub2e4. prune.ln_structured(module, name=\"weight\", amount=0.5, n=2, dim=0) Conv2d(1, 6, kernel_size=(5, 5), stride=(1, 1)) \uc6b0\ub9ac\uac00 \ud655\uc778\ud560 \uc218 \uc788\ub4ef\uc774, \uc774\uc804 \ub9c8\uc2a4\ud06c\uc758 \uc791\uc6a9\uc744 \uc720\uc9c0\ud558\uba74\uc11c \ucc44\ub110\uc758 50% (6\uac1c \uc911 3\uac1c) \uc5d0 \ud574\ub2f9\ub418\ub294 \ubaa8\ub4e0 \uc5f0\uacb0\uc744 0\uc73c\ub85c \ubcc0\uacbd\ud569\ub2c8\ub2e4. print(module.weight) tensor([[[[ 0.0000, -0.0000, -0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, 0.0000, 0.0000], [-0.0000, -0.0000, 0.0000, 0.0000, 0.0000], [ 0.0000, -0.0000, 0.0000, 0.0000, 0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000]]], [[[-0.1891, 0.0554, -0.0000, 0.1154, -0.0000], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0000], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0000, -0.1911, -0.0794, 0.0000, 0.0000], [ 0.0427, -0.1558, -0.1073, -0.0000, 0.0211]]], [[[ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000], [-0.0000, -0.0000, 0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, 0.0000, -0.0000, -0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000]]], [[[-0.0958, 0.0936, 0.1754, -0.0000, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0000, 0.0111, 0.1206, 0.0000], [ 0.0000, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0000, 0.0628, 0.0000]]], [[[ 0.1180, -0.0000, 0.1336, -0.0599, -0.0000], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0000], [-0.0000, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.0000, 0.0753, -0.0684], [ 0.0304, 0.0000, -0.0709, 0.0481, -0.0000]]], [[[ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, 0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, -0.0000, 0.0000, 0.0000]]]], device=\u0027cuda:0\u0027, grad_fn=\u003cMulBackward0\u003e) \uc774\uc5d0 \ud574\ub2f9\ud558\ub294 \ud6c5\uc740 torch.nn.utils.prune.PruningContainer \ud615\ud0dc\ub85c \uc874\uc7ac\ud558\uba70, \uac00\uc911\uce58\uc5d0 \uc801\uc6a9\ub41c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc758 \uc774\ub825\uc744 \uc800\uc7a5\ud569\ub2c8\ub2e4. for hook in module._forward_pre_hooks.values(): if hook._tensor_name == \"weight\": # \uac00\uc911\uce58\uc5d0 \ud574\ub2f9\ud558\ub294 \ud6c5\uc744 \uc120\ud0dd break print(list(hook)) # \ucee8\ud14c\uc774\ub108 \ub0b4 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc758 \uc774\ub825 [\u003ctorch.nn.utils.prune.RandomUnstructured object at 0x7f0dab246050\u003e, \u003ctorch.nn.utils.prune.LnStructured object at 0x7f0dab0edd10\u003e] \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \ubaa8\ub378\uc758 \uc9c1\ub82c\ud654# \ub9c8\uc2a4\ud06c \ubc84\ud37c\ub4e4\uacfc \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \ud150\uc11c \uacc4\uc0b0\uc5d0 \uc0ac\uc6a9\ub41c \uae30\uc874\uc758 \ud30c\ub77c\ubbf8\ud130\ub97c \ud3ec\ud568\ud558\uc5ec \uad00\ub828\ub41c \ubaa8\ub4e0 \ud150\uc11c\uac12\ub4e4\uc740 \ud544\uc694\ud55c \uacbd\uc6b0 \ubaa8\ub378\uc758 state_dict \uc5d0 \uc800\uc7a5\ub418\uae30 \ub54c\ubb38\uc5d0, \uc27d\uac8c \uc9c1\ub82c\ud654\ud558\uc5ec \uc800\uc7a5\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. print(model.state_dict().keys()) odict_keys([\u0027conv1.weight_orig\u0027, \u0027conv1.bias_orig\u0027, \u0027conv1.weight_mask\u0027, \u0027conv1.bias_mask\u0027, \u0027conv2.weight\u0027, \u0027conv2.bias\u0027, \u0027fc1.weight\u0027, \u0027fc1.bias\u0027, \u0027fc2.weight\u0027, \u0027fc2.bias\u0027, \u0027fc3.weight\u0027, \u0027fc3.bias\u0027]) \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc758 \uc7ac-\ud30c\ub77c\ubbf8\ud130\ud654 \uc81c\uac70# \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uac83\uc744 \uc601\uad6c\uc801\uc73c\ub85c \ub9cc\ub4e4\uae30 \uc704\ud574\uc11c, \uc7ac-\ud30c\ub77c\ubbf8\ud130\ud654 \uad00\uc810\uc758 weight_orig \uc640 weight_mask \uac12\uc744 \uc81c\uac70\ud558\uace0, forward_pre_hook \uac12\uc744 \uc81c\uac70\ud569\ub2c8\ub2e4. \uc81c\uac70\ud558\uae30 \uc704\ud574 torch.nn.utils.prune \ub0b4 remove \ud568\uc218\ub97c \uc774\uc6a9\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc9c0 \uc54a\uc740 \uac83\ucc98\ub7fc \uc2e4\ud589\ub418\ub294 \uac83\uc774 \uc544\ub2cc \uc810\uc744 \uc8fc\uc758\ud558\uc138\uc694. \uc774\ub294 \ub2e8\uc9c0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uc0c1\ud0dc\uc5d0\uc11c \uac00\uc911\uce58 \ud30c\ub77c\ubbf8\ud130\uac12\uc744 \ubaa8\ub378 \ud30c\ub77c\ubbf8\ud130\uac12\uc73c\ub85c \uc7ac\ud560\ub2f9\ud558\ub294 \uac83\uc744 \ud1b5\ud574 \uc601\uad6c\uc801\uc73c\ub85c \ub9cc\ub4dc\ub294 \uac83\uc77c \ubfd0\uc785\ub2c8\ub2e4. \uc7ac-\ud30c\ub77c\ubbf8\ud130\ud654\ub97c \uc81c\uac70\ud558\uae30 \uc804 \uc0c1\ud0dc print(list(module.named_parameters())) [(\u0027weight_orig\u0027, Parameter containing: tensor([[[[ 0.1232, -0.0642, -0.0623, -0.1485, 0.0055], [ 0.0314, -0.1857, -0.0403, 0.1223, 0.1689], [-0.0559, -0.1542, 0.0177, 0.0988, 0.0749], [ 0.0634, -0.1869, 0.1766, 0.1120, 0.0219], [-0.1110, -0.0066, -0.0742, 0.1935, -0.1882]]], [[[-0.1891, 0.0554, -0.1982, 0.1154, -0.0523], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0445], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0206, -0.1911, -0.0794, 0.1364, 0.0037], [ 0.0427, -0.1558, -0.1073, -0.0763, 0.0211]]], [[[ 0.0116, -0.1011, -0.1792, -0.0166, -0.1940], [-0.1058, -0.0902, -0.0587, 0.0361, -0.0123], [-0.1890, -0.0632, 0.0668, -0.0883, -0.1008], [-0.0702, 0.1404, 0.0646, -0.1084, -0.0797], [-0.0942, -0.0567, -0.1763, 0.0473, -0.1682]]], [[[-0.0958, 0.0936, 0.1754, -0.0095, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0554, 0.0111, 0.1206, 0.0263], [ 0.1758, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0716, 0.0628, 0.0945]]], [[[ 0.1180, -0.0116, 0.1336, -0.0599, -0.0110], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0822], [-0.1528, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.1081, 0.0753, -0.0684], [ 0.0304, 0.0038, -0.0709, 0.0481, -0.0312]]], [[[ 0.0692, -0.1867, -0.0930, -0.0373, -0.1380], [-0.0196, 0.1388, 0.0801, -0.1948, 0.0013], [ 0.0478, -0.1248, -0.0969, -0.1181, 0.1294], [ 0.0343, -0.0799, -0.0200, -0.1351, -0.0577], [-0.0725, 0.0144, -0.0758, 0.0333, 0.0219]]]], device=\u0027cuda:0\u0027, requires_grad=True)), (\u0027bias_orig\u0027, Parameter containing: tensor([-0.0229, -0.0985, -0.0547, -0.0255, -0.1460, -0.1458], device=\u0027cuda:0\u0027, requires_grad=True))] print(list(module.named_buffers())) [(\u0027weight_mask\u0027, tensor([[[[0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.]]], [[[1., 1., 0., 1., 0.], [1., 1., 1., 1., 0.], [1., 1., 1., 1., 1.], [0., 1., 1., 0., 0.], [1., 1., 1., 0., 1.]]], [[[0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.]]], [[[1., 1., 1., 0., 1.], [1., 1., 1., 1., 1.], [1., 0., 1., 1., 0.], [0., 1., 1., 1., 1.], [1., 1., 0., 1., 0.]]], [[[1., 0., 1., 1., 0.], [1., 1., 1., 1., 0.], [0., 1., 1., 1., 1.], [1., 1., 0., 1., 1.], [1., 0., 1., 1., 0.]]], [[[0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.]]]], device=\u0027cuda:0\u0027)), (\u0027bias_mask\u0027, tensor([0., 1., 0., 0., 1., 1.], device=\u0027cuda:0\u0027))] print(module.weight) tensor([[[[ 0.0000, -0.0000, -0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, 0.0000, 0.0000], [-0.0000, -0.0000, 0.0000, 0.0000, 0.0000], [ 0.0000, -0.0000, 0.0000, 0.0000, 0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000]]], [[[-0.1891, 0.0554, -0.0000, 0.1154, -0.0000], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0000], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0000, -0.1911, -0.0794, 0.0000, 0.0000], [ 0.0427, -0.1558, -0.1073, -0.0000, 0.0211]]], [[[ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000], [-0.0000, -0.0000, 0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, 0.0000, -0.0000, -0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000]]], [[[-0.0958, 0.0936, 0.1754, -0.0000, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0000, 0.0111, 0.1206, 0.0000], [ 0.0000, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0000, 0.0628, 0.0000]]], [[[ 0.1180, -0.0000, 0.1336, -0.0599, -0.0000], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0000], [-0.0000, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.0000, 0.0753, -0.0684], [ 0.0304, 0.0000, -0.0709, 0.0481, -0.0000]]], [[[ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, 0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, -0.0000, 0.0000, 0.0000]]]], device=\u0027cuda:0\u0027, grad_fn=\u003cMulBackward0\u003e) \uc7ac-\ud30c\ub77c\ubbf8\ud130\ub97c \uc81c\uac70\ud55c \ud6c4 \uc0c1\ud0dc prune.remove(module, \u0027weight\u0027) print(list(module.named_parameters())) [(\u0027bias_orig\u0027, Parameter containing: tensor([-0.0229, -0.0985, -0.0547, -0.0255, -0.1460, -0.1458], device=\u0027cuda:0\u0027, requires_grad=True)), (\u0027weight\u0027, Parameter containing: tensor([[[[ 0.0000, -0.0000, -0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, 0.0000, 0.0000], [-0.0000, -0.0000, 0.0000, 0.0000, 0.0000], [ 0.0000, -0.0000, 0.0000, 0.0000, 0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000]]], [[[-0.1891, 0.0554, -0.0000, 0.1154, -0.0000], [ 0.0580, 0.1063, 0.1315, -0.1787, -0.0000], [-0.1216, 0.1807, 0.1527, -0.1806, -0.0789], [-0.0000, -0.1911, -0.0794, 0.0000, 0.0000], [ 0.0427, -0.1558, -0.1073, -0.0000, 0.0211]]], [[[ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000], [-0.0000, -0.0000, 0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, 0.0000, -0.0000, -0.0000], [-0.0000, -0.0000, -0.0000, 0.0000, -0.0000]]], [[[-0.0958, 0.0936, 0.1754, -0.0000, 0.0009], [-0.1752, -0.1877, 0.1632, -0.0735, 0.1270], [-0.0448, -0.0000, 0.0111, 0.1206, 0.0000], [ 0.0000, -0.1420, 0.1933, -0.1722, -0.1062], [-0.0772, 0.0547, 0.0000, 0.0628, 0.0000]]], [[[ 0.1180, -0.0000, 0.1336, -0.0599, -0.0000], [ 0.1084, 0.1545, -0.0840, -0.1709, -0.0000], [-0.0000, 0.1098, 0.1429, -0.0835, -0.1162], [-0.1901, 0.0091, 0.0000, 0.0753, -0.0684], [ 0.0304, 0.0000, -0.0709, 0.0481, -0.0000]]], [[[ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, 0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, -0.0000, 0.0000], [ 0.0000, -0.0000, -0.0000, -0.0000, -0.0000], [-0.0000, 0.0000, -0.0000, 0.0000, 0.0000]]]], device=\u0027cuda:0\u0027, requires_grad=True))] print(list(module.named_buffers())) [(\u0027bias_mask\u0027, tensor([0., 1., 0., 0., 1., 1.], device=\u0027cuda:0\u0027))] \ubaa8\ub378 \ub0b4 \uc5ec\ub7ec \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\uc5d0 \ub300\ud558\uc5ec \uac00\uc9c0\uce58\uae30 \uae30\ubc95 \uc801\uc6a9# \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uace0 \uc2f6\uc740 \ud30c\ub77c\ubbf8\ud130\uac12\ub4e4\uc744 \uc9c0\uc815\ud568\uc73c\ub85c\uc368, \uc774\ubc88 \uc608\uc81c\uc5d0\uc11c \ubcfc \uc218 \uc788\ub294 \uac83 \ucc98\ub7fc, \uc2e0\uacbd\ub9dd \ubaa8\ub378 \ub0b4 \uc5ec\ub7ec \ud150\uc11c\uac12\ub4e4\uc5d0 \ub300\ud574\uc11c \uc27d\uac8c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. new_model = LeNet() for name, module in new_model.named_modules(): # \ubaa8\ub4e0 2D-conv \uce35\uc758 20% \uc5f0\uacb0\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9 if isinstance(module, torch.nn.Conv2d): prune.l1_unstructured(module, name=\u0027weight\u0027, amount=0.2) # \ubaa8\ub4e0 \uc120\ud615 \uce35\uc758 40% \uc5f0\uacb0\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9 elif isinstance(module, torch.nn.Linear): prune.l1_unstructured(module, name=\u0027weight\u0027, amount=0.4) print(dict(new_model.named_buffers()).keys()) # \uc874\uc7ac\ud558\ub294 \ubaa8\ub4e0 \ub9c8\uc2a4\ud06c\ub4e4\uc744 \ud655\uc778 dict_keys([\u0027conv1.weight_mask\u0027, \u0027conv2.weight_mask\u0027, \u0027fc1.weight_mask\u0027, \u0027fc2.weight_mask\u0027, \u0027fc3.weight_mask\u0027]) \uc804\uc5ed \ubc94\uc704\uc5d0 \ub300\ud55c \uac00\uc9c0\uce58\uae30 \uae30\ubc95 \uc801\uc6a9# \uc9c0\uae08\uae4c\uc9c0, \u201c\uc9c0\uc5ed \ubcc0\uc218\u201d \uc5d0 \ub300\ud574\uc11c\ub9cc \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\ub294 \ubc29\ubc95\uc744 \uc0b4\ud3b4\ubcf4\uc558\uc2b5\ub2c8\ub2e4. (\uc989, \uac00\uc911\uce58 \uaddc\ubaa8, \ud65c\uc131\ud654 \uc815\ub3c4, \uacbd\uc0ac\uac12 \ub4f1\uc758 \uac01 \ud56d\ubaa9\uc758 \ud1b5\uacc4\ub7c9\uc744 \ubc14\ud0d5\uc73c\ub85c \ubaa8\ub378 \ub0b4 \ud150\uc11c\uac12 \ud558\ub098\uc529 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\ub294 \ubc29\uc2dd) \uadf8\ub7ec\ub098, \ubc94\uc6a9\uc801\uc774\uace0 \uc544\ub9c8 \ub354 \uac15\ub825\ud55c \ubc29\ubc95\uc740 \uac01 \uce35\uc5d0\uc11c \uac00\uc7a5 \ub0ae\uc740 20%\uc758 \uc5f0\uacb0\uc744 \uc81c\uac70\ud558\ub294 \uac83 \ub300\uc2e0\uc5d0, \uc804\uccb4 \ubaa8\ub378\uc5d0 \ub300\ud574\uc11c \uac00\uc7a5 \ub0ae\uc740 20% \uc5f0\uacb0\uc744 \ud55c\ubc88\uc5d0 \uc81c\uac70\ud558\ub294 \uac83\uc785\ub2c8\ub2e4. \uc774\uac83\uc740 \uac01 \uce35\uc5d0 \ub300\ud574\uc11c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\ub294 \uc5f0\uacb0\uc758 \ubc31\ubd84\uc728\uac12\uc744 \ub2e4\ub974\uac8c \ub9cc\ub4e4 \uac00\ub2a5\uc131\uc774 \uc788\uc2b5\ub2c8\ub2e4. torch.nn.utils.prune \ub0b4 global_unstructured \uc744 \uc774\uc6a9\ud558\uc5ec \uc5b4\ub5bb\uac8c \uc804\uc5ed \ubc94\uc704\uc5d0 \ub300\ud55c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\ub294\uc9c0 \uc0b4\ud3b4\ubd05\uc2dc\ub2e4. model = LeNet() parameters_to_prune = ( (model.conv1, \u0027weight\u0027), (model.conv2, \u0027weight\u0027), (model.fc1, \u0027weight\u0027), (model.fc2, \u0027weight\u0027), (model.fc3, \u0027weight\u0027), ) prune.global_unstructured( parameters_to_prune, pruning_method=prune.L1Unstructured, amount=0.2, ) \uc774\uc81c \uac01 \uce35\uc5d0 \uc874\uc7ac\ud558\ub294 \uc5f0\uacb0\ub4e4\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uc815\ub3c4\uac00 20%\uac00 \uc544\ub2cc \uac83\uc744 \ud655\uc778\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \uadf8\ub7ec\ub098, \uc804\uccb4 \uac00\uc9c0\uce58\uae30 \uc801\uc6a9 \ubc94\uc704\ub294 \uc57d 20%\uac00 \ub420 \uac83\uc785\ub2c8\ub2e4. print( \"Sparsity in conv1.weight: {:.2f}%\".format( 100. * float(torch.sum(model.conv1.weight == 0)) / float(model.conv1.weight.nelement()) ) ) print( \"Sparsity in conv2.weight: {:.2f}%\".format( 100. * float(torch.sum(model.conv2.weight == 0)) / float(model.conv2.weight.nelement()) ) ) print( \"Sparsity in fc1.weight: {:.2f}%\".format( 100. * float(torch.sum(model.fc1.weight == 0)) / float(model.fc1.weight.nelement()) ) ) print( \"Sparsity in fc2.weight: {:.2f}%\".format( 100. * float(torch.sum(model.fc2.weight == 0)) / float(model.fc2.weight.nelement()) ) ) print( \"Sparsity in fc3.weight: {:.2f}%\".format( 100. * float(torch.sum(model.fc3.weight == 0)) / float(model.fc3.weight.nelement()) ) ) print( \"Global sparsity: {:.2f}%\".format( 100. * float( torch.sum(model.conv1.weight == 0) + torch.sum(model.conv2.weight == 0) + torch.sum(model.fc1.weight == 0) + torch.sum(model.fc2.weight == 0) + torch.sum(model.fc3.weight == 0) ) / float( model.conv1.weight.nelement() + model.conv2.weight.nelement() + model.fc1.weight.nelement() + model.fc2.weight.nelement() + model.fc3.weight.nelement() ) ) ) Sparsity in conv1.weight: 4.67% Sparsity in conv2.weight: 13.12% Sparsity in fc1.weight: 22.33% Sparsity in fc2.weight: 11.69% Sparsity in fc3.weight: 9.17% Global sparsity: 20.00% torch.nn.utils.prune \uc5d0\uc11c \ud655\uc7a5\ub41c \ub9de\ucda4\ud615 \uac00\uc9c0\uce58\uae30 \uae30\ubc95# \ub9de\ucda4\ud615 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc740, \ub2e4\ub978 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\ub294 \uac83\uacfc \uac19\uc740 \ubc29\uc2dd\uc73c\ub85c, BasePruningMethod \uc758 \uae30\ubcf8 \ud074\ub798\uc2a4\uc778 nn.utils.prune \ubaa8\ub4c8\uc744 \ud65c\uc6a9\ud558\uc5ec \uad6c\ud604\ud560 \uc218 \uc788\uc2b5\ub2c8\ub2e4. \uae30\ubcf8 \ud074\ub798\uc2a4\ub294 __call__, apply_mask, apply, prune, remove \uba54\uc18c\ub4dc\ub4e4\uc744 \ub0b4\ud3ec\ud558\uace0 \uc788\uc2b5\ub2c8\ub2e4. \ud2b9\ubcc4\ud55c \ucf00\uc774\uc2a4\uac00 \uc544\ub2cc \uacbd\uc6b0, \uae30\ubcf8\uc801\uc73c\ub85c \uad6c\uc131\ub41c \uba54\uc18c\ub4dc\ub4e4\uc744 \uc7ac\uad6c\uc131\ud560 \ud544\uc694\uac00 \uc5c6\uc2b5\ub2c8\ub2e4. \uadf8\ub7ec\ub098, __init__ (\uad6c\uc131\uc694\uc18c), compute_mask (\uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc758 \ub17c\ub9ac\uc5d0 \ub530\ub77c \uc8fc\uc5b4\uc9c4 \ud150\uc11c\uac12\uc5d0 \ub9c8\uc2a4\ud06c\ub97c \uc801\uc6a9\ud558\ub294 \ubc29\ubc95) \uc744 \uace0\ub824\ud558\uc5ec \uad6c\uc131\ud574\uc57c \ud569\ub2c8\ub2e4. \uac8c\ub2e4\uac00, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc5b4\ub5a0\ud55c \ubc29\uc2dd\uc73c\ub85c \uc801\uc6a9\ud558\ub294\uc9c0 \uba85\ud655\ud558\uac8c \uad6c\uc131\ud574\uc57c \ud569\ub2c8\ub2e4. (\uc9c0\uc6d0\ub418\ub294 \uc635\uc158\uc740 global, structured, unstructured \uc785\ub2c8\ub2e4.) \uc774\ub7ec\ud55c \ubc29\uc2dd\uc740, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \ubc18\ubcf5\uc801\uc73c\ub85c \uc801\uc6a9\ud574\uc57c \ud558\ub294 \uacbd\uc6b0 \ub9c8\uc2a4\ud06c\ub97c \uacb0\ud569\ud558\ub294 \ubc29\ubc95\uc744 \uacb0\uc815\ud558\uae30 \uc704\ud574 \ud544\uc694\ud569\ub2c8\ub2e4. \uc989, \uc774\ubbf8 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \ubaa8\ub378\uc5d0 \ub300\ud574\uc11c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud560 \ub54c, \uae30\uc874\uc758 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc9c0 \uc54a\uc740 \ud30c\ub77c\ubbf8\ud130 \uac12\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc601\ud5a5\uc744 \ubbf8\uce60 \uac83\uc73c\ub85c \uc608\uc0c1\ub429\ub2c8\ub2e4. PRUNING_TYPE \uc744 \uc9c0\uc815\ud55c\ub2e4\uba74, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud558\uae30 \uc704\ud574 \ud30c\ub77c\ubbf8\ud130 \uac12\uc744 \uc62c\ubc14\ub974\uac8c \uc81c\uac70\ud558\ub294 PruningContainer (\ub9c8\uc2a4\ud06c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \ubc18\ubcf5\uc801\uc73c\ub85c \uc801\uc6a9\ud558\ub294 \uac83\uc744 \ucc98\ub9ac\ud558\ub294)\ub97c \uac00\ub2a5\ud558\uac8c \ud569\ub2c8\ub2e4. \uc608\ub97c \ub4e4\uc5b4, \ub2e4\ub978 \ubaa8\ub4e0 \ud56d\ubaa9\uc774 \uc874\uc7ac\ud558\ub294 \ud150\uc11c\ub97c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uad6c\ud604\ud558\uace0 \uc2f6\uc744 \ub54c, (\ub610\ub294, \ud150\uc11c\uac00 \uc774\uc804\uc5d0 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc5d0 \uc758\ud574 \uc81c\uac70\ub418\uc5c8\uac70\ub098 \ub0a8\uc544\uc788\ub294 \ud150\uc11c\uc5d0 \ub300\ud574) \ud55c \uce35\uc758 \uac1c\ubcc4 \uc5f0\uacb0\uc5d0 \uc791\uc6a9\ud558\uba70 \uc804\uccb4 \uc720\ub2db/\ucc44\ub110 (\u0027structured\u0027), \ub610\ub294 \ub2e4\ub978 \ud30c\ub77c\ubbf8\ud130 \uac04 (\u0027global\u0027) \uc5f0\uacb0\uc5d0\ub294 \uc791\uc6a9\ud558\uc9c0 \uc54a\uae30 \ub54c\ubb38\uc5d0 PRUNING_TYPE=\u0027unstructured\u0027 \ubc29\uc2dd\uc73c\ub85c \uc9c4\ud589\ub429\ub2c8\ub2e4. class FooBarPruningMethod(prune.BasePruningMethod): \"\"\" \ud150\uc11c \ub0b4 \ub2e4\ub978 \ud56d\ubaa9\ub4e4\uc5d0 \ub300\ud574 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9 \"\"\" PRUNING_TYPE = \u0027unstructured\u0027 def compute_mask(self, t, default_mask): mask = default_mask.clone() mask.view(-1)[::2] = 0 return mask nn.Module \uc758 \ub9e4\uac1c\ubcc0\uc218\uc5d0 \uc801\uc6a9\ud558\uae30 \uc704\ud574 \uc778\uc2a4\ud134\uc2a4\ud654\ud558\uace0 \uc801\uc6a9\ud558\ub294 \uac04\ub2e8\ud55c \uae30\ub2a5\uc744 \uad6c\ud604\ud574\ubd05\ub2c8\ub2e4. def foobar_unstructured(module, name): \"\"\" \ud150\uc11c \ub0b4 \ub2e4\ub978 \ubaa8\ub4e0 \ud56d\ubaa9\ub4e4\uc744 \uc81c\uac70\ud558\uc5ec `module` \uc5d0\uc11c `name` \uc774\ub77c\ub294 \ud30c\ub77c\ubbf8\ud130\uc5d0 \ub300\ud574 \uac00\uc790\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9 \ub2e4\uc74c \ub0b4\uc6a9\uc5d0 \ub530\ub77c \ubaa8\ub4c8\uc744 \uc218\uc815 (\ub610\ub294 \uc218\uc815\ub41c \ubaa8\ub4c8\uc744 \ubc18\ud658): 1) \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc5d0 \uc758\ud574 \ub9e4\uac1c\ubcc0\uc218 `name` \uc5d0 \uc801\uc6a9\ub41c \uc774\uc9c4 \ub9c8\uc2a4\ud06c\uc5d0 \ud574\ub2f9\ud558\ub294 \uba85\uba85\ub41c \ubc84\ud37c `name+\u0027_mask\u0027` \ub97c \ucd94\uac00\ud569\ub2c8\ub2e4. `name` \ud30c\ub77c\ubbf8\ud130\ub294 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \uac83\uc73c\ub85c \ub300\uccb4\ub418\uba70, \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub418\uc9c0 \uc54a\uc740 \uae30\uc874\uc758 \ud30c\ub77c\ubbf8\ud130\ub294 `name+\u0027_orig\u0027` \ub77c\ub294 \uc774\ub984\uc758 \uc0c8\ub85c\uc6b4 \ub9e4\uac1c\ubcc0\uc218\uc5d0 \uc800\uc7a5\ub429\ub2c8\ub2e4. \uc778\uc790\uac12: module (nn.Module): \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc744 \uc801\uc6a9\ud574\uc57c \ud558\ub294 \ud150\uc11c\ub97c \ud3ec\ud568\ud558\ub294 \ubaa8\ub4c8 name (string): \ubaa8\ub4c8 \ub0b4 \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub420 \ud30c\ub77c\ubbf8\ud130\uc758 \uc774\ub984 \ubc18\ud658\uac12: module (nn.Module): \uc785\ub825 \ubaa8\ub4c8\uc5d0 \ub300\ud574\uc11c \uac00\uc9c0\uce58\uae30 \uae30\ubc95\uc774 \uc801\uc6a9\ub41c \ubaa8\ub4c8 \uc608\uc2dc: \u003e\u003e\u003e m = nn.Linear(3, 4) \u003e\u003e\u003e foobar_unstructured(m, name=\u0027bias\u0027) \"\"\" FooBarPruningMethod.apply(module, name) return module \ud55c\ubc88 \ud574\ubd05\uc2dc\ub2e4! model = LeNet() foobar_unstructured(model.fc3, name=\u0027bias\u0027) print(model.fc3.bias_mask) tensor([0., 1., 0., 1., 0., 1., 0., 1., 0., 1.]) Total running time of the script: (0 minutes 4.590 seconds) Download Jupyter notebook: pruning_tutorial.ipynb Download Python source code: pruning_tutorial.py Download zipped: pruning_tutorial.zip",
"author": {
"@type": "Organization",
"name": "PyTorch Contributors",
"url": "https://pytorch.org"
},
"image": "../_static/img/pytorch_seo.png",
"mainEntityOfPage": {
"@type": "WebPage",
"@id": "/intermediate/pruning_tutorial.html"
},
"datePublished": "2023-01-01T00:00:00Z",
"dateModified": "2023-01-01T00:00:00Z"
}
| article:modified_time | 2022-11-30T07:09:41+00:00 |
| og:type | article |
| og:site_name | PyTorch Tutorials KR |
| og:image | ../_static/img/pytorch_seo.png |
| og:image:alt | PyTorch Tutorials KR |
| og:ignore_canonical | true |
| docsearch:language | ko |
| docbuild:last-update | 2022년 11월 30일 |
| None | 2 |
| pytorch_project | tutorials |
Links:
Viewport: width=device-width, initial-scale=1