{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":8575494,"sourceType":"datasetVersion","datasetId":5127884}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torch torchvision pillow > /dev/null","metadata":{"execution":{"iopub.status.busy":"2024-06-01T09:47:59.238023Z","iopub.execute_input":"2024-06-01T09:47:59.238828Z","iopub.status.idle":"2024-06-01T09:48:12.475584Z","shell.execute_reply.started":"2024-06-01T09:47:59.238796Z","shell.execute_reply":"2024-06-01T09:48:12.474325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import transforms, models\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\ndef load_image(image_path, transform=None, max_size=400, shape=None):\n    image = Image.open(image_path)\n    \n    if max_size:\n        size = max(image.size)\n        if size > max_size:\n            size = max_size\n        image = image.resize((size, int(size * image.height / image.width)))\n    \n    if shape:\n        image = image.resize(shape)\n    \n    if transform:\n        image = transform(image).unsqueeze(0)\n    \n    return image\n\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n])\n\ncontent_image_path = \"/kaggle/input/images/image.png\"\nstyle_image_path = \"/kaggle/input/images/filter.png\"\n\ncontent_image = load_image(content_image_path, transform, max_size=400)\nstyle_image = load_image(style_image_path, transform, shape=[content_image.size(2), content_image.size(3)])\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ncontent_image = content_image.to(device)\nstyle_image = style_image.to(device)\n\ndef imshow(tensor, title=None):\n    image = tensor.cpu().clone().detach().squeeze(0)\n    image = image.clamp_(0, 1)\n    plt.imshow(image.permute(1, 2, 0))\n    if title:\n        plt.title(title)\n    plt.show()\n\nplt.figure(figsize=(10, 5))\nplt.subplot(1, 2, 1)\nimshow(content_image, title='Контентное изображение')\n\nplt.subplot(1, 2, 2)\nimshow(style_image, title='Стильное изображение')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T09:48:48.686273Z","iopub.execute_input":"2024-06-01T09:48:48.687098Z","iopub.status.idle":"2024-06-01T09:48:49.423974Z","shell.execute_reply.started":"2024-06-01T09:48:48.687062Z","shell.execute_reply":"2024-06-01T09:48:49.423076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VGGFeatureExtractor(nn.Module):\n    def __init__(self):\n        super(VGGFeatureExtractor, self).__init__()\n        self.model = models.vgg19(pretrained=True).features\n        self.layers = {\n            '0': 'conv1_1',\n            '5': 'conv2_1',\n            '10': 'conv3_1',\n            '19': 'conv4_1',\n            '21': 'conv4_2',\n            '28': 'conv5_1'\n        }\n    \n    def forward(self, x):\n        features = {}\n        for name, layer in self.model._modules.items():\n            x = layer(x)\n            if name in self.layers:\n                features[self.layers[name]] = x\n        return features\n\nvgg_extractor = VGGFeatureExtractor().to(device).eval()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T11:08:36.501775Z","iopub.execute_input":"2024-06-01T11:08:36.502780Z","iopub.status.idle":"2024-06-01T11:08:38.275674Z","shell.execute_reply.started":"2024-06-01T11:08:36.502737Z","shell.execute_reply":"2024-06-01T11:08:38.274694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"style_weights = {'conv1_1': 1.0, 'conv2_1': 0.8, 'conv3_1': 0.5, 'conv4_1': 0.3, 'conv5_1': 0.2}\ncontent_weight = 1\nstyle_weight = 1e6\n\ndef compute_content_loss(predicted, target):\n    return torch.mean((predicted - target)**2)\n\ndef compute_gram_matrix(feature_map):\n    batch_size, num_features, h, w = feature_map.size()\n    features = feature_map.view(batch_size * num_features, h * w)\n    gram_matrix = torch.mm(features, features.t())\n    return gram_matrix / (batch_size * num_features * h * w)\n\ndef compute_style_loss(style_grams, target_features):\n    loss = 0\n    for layer in style_weights:\n        target_gram = compute_gram_matrix(target_features[layer])\n        style_gram = style_grams[layer]\n        layer_loss = style_weights[layer] * torch.mean((target_gram - style_gram)**2)\n        loss += layer_loss\n    return loss\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T11:08:38.277192Z","iopub.execute_input":"2024-06-01T11:08:38.277472Z","iopub.status.idle":"2024-06-01T11:08:38.285485Z","shell.execute_reply.started":"2024-06-01T11:08:38.277448Z","shell.execute_reply":"2024-06-01T11:08:38.284587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_image = content_image.clone().requires_grad_(True).to(device)\n\noptimizer = optim.Adam([target_image], lr=0.003)\n\nstyle_features = vgg_extractor(style_image)\nstyle_grams = {layer: compute_gram_matrix(style_features[layer]) for layer in style_features}\ncontent_features = vgg_extractor(content_image)\n\nsteps = 5000\n\nfor step in range(steps):\n    target_features = vgg_extractor(target_image)\n    content_loss = compute_content_loss(target_features['conv4_2'], content_features['conv4_2'])\n    style_loss = compute_style_loss(style_grams, target_features)\n    total_loss = content_weight * content_loss + style_weight * style_loss\n    \n    optimizer.zero_grad()\n    total_loss.backward(retain_graph=True)\n    optimizer.step()\n\n    if step % 500 == 0:\n        print(f\"Step {step}, Total loss: {total_loss.item()}\")\n        imshow(target_image, title=f'Step {step}')\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T11:20:01.108425Z","iopub.execute_input":"2024-06-01T11:20:01.109314Z","iopub.status.idle":"2024-06-01T11:33:35.519045Z","shell.execute_reply.started":"2024-06-01T11:20:01.109279Z","shell.execute_reply":"2024-06-01T11:33:35.518207Z"},"trusted":true},"execution_count":null,"outputs":[]}]}