{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# Multilingual CLIP with XLM-RoBERTa as Text Encoder\n\nImplementation with Pytorch Lighting\n","metadata":{}},{"cell_type":"code","source":"###########################\n#Author Marten Rogall\n#References:\n#https://colab.research.google.com/github/sachinruk/blog/blob/master/_notebooks/2021-03-07-CLIP.ipynb#scrollTo=3HbcfpJSdXjZ\n#https://github.com/moein-shariatnia/OpenAI-CLIP\n#https://github.com/openai/CLIP/issues/83\n#https://github.com/gzomer/clip-multilingual\n############################\n\n!pip install transformers pytorch-lightning\n!pip install ftfy regex tqdm\n!pip install git+https://github.com/openai/CLIP.git\n\n############################","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"############################\n#Imports\n############################\nimport os\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nimport multiprocessing\nimport random\nfrom typing import Dict, Tuple, List\nimport gc\nimport encodings\n\n#Images and Plotting\nimport uuid\nfrom urllib import request\nfrom urllib.request import urlopen\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport pickle\n\n\nimport pandas as pd\nimport numpy as np\nimport dask.dataframe as dd\nfrom sklearn.model_selection import train_test_split\n\n#CLIP\nimport clip\n\n#Transformer\nfrom transformers import AutoTokenizer, AutoModel\n\n#PyTorch Lighting/PyTorch\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer, seed_everything\nfrom torch.utils.data import random_split, DataLoader","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config():\n    seed = 42\n    batch_size = 64\n    epochs = 5\n    num_workers = multiprocessing.cpu_count()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    #learning rates\n    image_encoder_lr = 1e-4\n    text_encoder_lr = 1e-5\n    head_lr = 1e-3\n    #optimizer = optim.Adam(clip_model.parameters(), lr=5e-5,betas=(0.9,0.98),eps=1e-6,weight_decay=0.2) \n    #Parameter from the paper; lr for finetuning set smaller than the paper value\n    \n    #CLIP\n    clip_visual_model = 'RN50x4'\n    clip_embed_dim = 640\n    embedding_dim = 512\n    \n    #TextEncoder\n    text_encoder_model = \"xlm-roberta-base\"\n    text_embedding = 768\n    context_length = 77\n    content_length_unicode = 60\n    \n    \n    #Projection Head\n    num_layer = 3\n    dropout = 0.5\n    projection_dim = 256\n\n    n_faktor = 20","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Seed for reproductivity\nseed_everything(Config.seed)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Globals\nDATASET_PATH = '../input/wikitrain/'\nD_NAME = 'https://upload.wikimedia.org'\nIMAGE_DIR = './images/'","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_image(link):\n    ###\n    #fetching the Images and returning them as PIL Images\n    ###\n    try:\n        URL = link\n        filename = str(uuid.uuid4())\n        path = f'./images/{filename}'\n        req = request.Request(URL)\n        req.add_header('User-Agent', 'User-bot-abc')\n        response = request.urlopen(req)\n        \n        with open(path, 'wb') as f:\n            f.write(response.read())\n        \n        image = Image.open(path).convert(\"RGB\")\n        \n        os.remove(path)\n        \n        return image\n    \n    except Exception as e:\n        print(e)\n        return None","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Erstellen der Image Directory\nPath('./images').mkdir(exist_ok=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"incr_encoder = encodings.search_function('utf8').incrementalencoder()\n\ndef utf8_byte_truncate(text, max_bytes):\n    ###\n    #truncating the utf captions to the clip content length\n    ###\n    byte_len = 0\n    incr_encoder.reset()\n    for index,ch in enumerate(text):\n        byte_len += len(incr_encoder.encode(ch))\n        if byte_len > max_bytes:\n            break\n    else:\n        return text\n    return text[:index]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_wiki = pd.read_csv(f'{DATASET_PATH}wiki_train.csv')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Tokenizer:\n    def __init__(self, tokenizer: AutoTokenizer):\n        self.tokenizer = tokenizer\n\n    def __call__(self, caption: str) -> AutoTokenizer:\n        return self.tokenizer(\n            caption,\n            add_special_tokens = True,#adding <s> and </s> tokens\n            max_length = Config.context_length,\n            truncation = True,\n            padding = 'max_length',\n            return_tensors = 'pt',\n        )\n\n    def decode(self, x: Dict[str, torch.LongTensor]):\n        return [self.tokenizer.decode(sentence[:sentence_len]) for sentence, sentence_len in\n                zip(x['input_ids'], x['attention_mask'].sum(axis=-1))]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WikiClipDataset(torch.utils.data.Dataset):\n    def __init__(self, df, transform_img = None, transform_text = None):\n\n        self.images = df[\"image_url\"].tolist()\n        self.captions = df['caption_title_and_reference_description'].tolist()\n        self.transform_img = transform_img\n        self.transform_text = transform_text\n        \n\n    def __len__(self):\n        return len(self.captions)\n\n    def __getitem__(self, idx):\n        \n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n        \n        \n        image_raw = load_image(self.images[idx])\n        #if the requested images are still 404\n        while image_raw == None:\n            idx = random.randint(0, len(self.images)-1)\n            image_raw = load_image(self.images[idx])\n        image = self.transform_img(image_raw)\n        caption_raw = self.captions[idx]\n        caption = self.transform_text(caption_raw)\n        \n        return image, caption","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_dl(df, image_transform, tokenizer):\n    \n    wikiData = WikiClipDataset(df, image_transform, tokenizer)\n    train_len = int(0.7*len(wikiData))\n    train_data, valid_data = random_split(wikiData, [train_len, len(wikiData) - train_len], generator = torch.Generator().manual_seed(Config.seed))\n\n    train_dl = DataLoader(\n        train_data,\n        Config.batch_size,\n        pin_memory = True,\n        shuffle = True,\n        num_workers = Config.num_workers,\n        drop_last = True\n    )\n\n    valid_dl = DataLoader(\n        valid_data,\n        Config.batch_size,\n        pin_memory = True,\n        shuffle = False,\n        num_workers = Config.num_workers,\n        drop_last = False\n    )\n    \n    return train_dl, valid_dl","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ProjectionHead(nn.Module):\n    ######\n    #Changing both Vectors to the same dimension\n    ######\n    def __init__(self, dim_in, dim_out, p = Config.dropout):\n        super().__init__()\n        self.linear1 = nn.Linear(dim_in, dim_out, bias=False)\n        self.linear2 = nn.Linear(dim_out, dim_out, bias=False)\n        self.layer_norm = nn.LayerNorm(dim_out)\n        self.drop = nn.Dropout(p)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        embed1 = self.linear1(x)\n        embed2 = self.drop(self.linear2(F.gelu(embed1)))\n        embeds = self.layer_norm(embed1 + embed2)\n        return embeds","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageEncoder(nn.Module):\n    def __init__(self, clip_vision_model, dim_in, dim_out):\n        super().__init__()\n        model = clip_vision_model\n        self.model = model\n        self.projection = ProjectionHead(dim_in, dim_out)\n        for p in self.model.parameters():#freezing backbone\n            p.requires_grad = False\n\n    def forward(self, x):\n        projected_vec = self.projection(self.model(x))\n        projection_len = torch.norm(projected_vec, dim = -1, keepdim = True)\n        return projected_vec / projection_len","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TextEncoder(nn.Module):\n    ######\n    #The new Text Encoder is changed to the multilingual xlm-roberta-base model\n    ######\n    def __init__(self, text_encoder_model, dim_out):\n        super().__init__()\n        self.model = AutoModel.from_pretrained(text_encoder_model)\n        self.projection = ProjectionHead(Config.text_embedding, dim_out)\n        for p in self.model.parameters():\n            p.requires_grad = False\n        \n    def forward(self, x):\n        out = self.model(**x)[0]\n        out = out[:, 0, :] #</s> token\n        projected_vec = self.projection(out)\n\n        projection_len = torch.norm(projected_vec, dim = -1, keepdim = True)\n        return projected_vec / projection_len","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def contrastive_loss(logits, dim):\n    neg_ce = torch.diag(F.log_softmax(logits, dim=dim))\n    return -neg_ce.mean()\n    \ndef clip_loss(similarity: torch.Tensor) -> torch.Tensor:\n    caption_loss = contrastive_loss(similarity, dim=0)\n    image_loss = contrastive_loss(similarity, dim=1)\n    return (caption_loss + image_loss) / 2.0\n\ndef metrics(similarity: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n    y = torch.arange(len(similarity)).to(similarity.device)\n    img2cap_match_idx = similarity.argmax(dim=1)\n    cap2img_match_idx = similarity.argmax(dim=0)\n\n    img_acc = (img2cap_match_idx == y).float().mean()\n    cap_acc = (cap2img_match_idx == y).float().mean()\n\n    return img_acc, cap_acc","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WikiCLIPMultilingual(pl.LightningModule):\n    def __init__(self, \n                 clip_model,\n                 transform_img,\n                 tokenizer,\n                 text_model_name,\n                 clip_embed_dim = Config.clip_embed_dim,\n                 embed_dim = Config.embedding_dim,\n                ):\n        super().__init__()\n        self.clip_model = clip_model\n        self.transform_img = transform_img\n        self.tokenizer = tokenizer\n        self.image_encoder = ImageEncoder(\n            clip_model.visual,\n            clip_embed_dim,\n            embed_dim,\n        )\n        self.text_encoder = TextEncoder(text_model_name, embed_dim)\n        self.save_hyperparameters()\n        \n    \n    def common_step(self, batch: Tuple[torch.Tensor, List[str]]) -> torch.Tensor:\n        images, text = batch\n        text = {k: torch.squeeze(v, 1).to(Config.device) for k, v in text.items()}\n\n        image_embed = self.image_encoder(images)\n        caption_embed = self.text_encoder(text)\n        similarity = caption_embed @ image_embed.T\n\n        loss = clip_loss(similarity)\n        img_acc, text_acc = metrics(similarity)\n        return loss, img_acc, text_acc\n\n    def training_step(self, batch: Tuple[torch.Tensor, List[str]], *args: list\n                     ) -> torch.Tensor:\n        loss, img_acc, text_acc = self.common_step(batch)\n        self.log('training_loss', loss, on_step=True)\n        self.log('training_img_acc', img_acc, on_step=True, prog_bar=True)\n        self.log('training_text_acc', text_acc, on_step=True, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch: Tuple[torch.Tensor, List[str]], *args: list\n                       ) -> torch.Tensor:\n        loss, img_acc, text_acc = self.common_step(batch)\n        self.log('validation_loss', loss, on_step=True)\n        self.log('validation_img_acc', img_acc, on_step=True, prog_bar=True)\n        self.log('validation_text_acc', text_acc, on_step=True, prog_bar=True)\n        return loss\n\n    def configure_optimizers(self) -> torch.optim.Optimizer:\n        vision_params = {'params': self.image_encoder.projection.parameters(), 'lr': Config.image_encoder_lr}\n        text_params = {'params': self.text_encoder.projection.parameters() , 'lr': Config.text_encoder_lr}\n        optimizer = torch.optim.Adam([vision_params, text_params])\n        return optimizer","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Instantiating the CLIP-Modell and Tokenizer\nclip_model, compose = clip.load(Config.clip_visual_model, device = Config.device, jit = False)\ntokenizer = Tokenizer(AutoTokenizer.from_pretrained(Config.text_encoder_model))\n\n#Creating the new model\nmodel = WikiCLIPMultilingual(\n    clip_model = clip_model,\n    transform_img = compose,\n    tokenizer = tokenizer,\n    text_model_name = Config.text_encoder_model,\n    clip_embed_dim = Config.clip_embed_dim,\n    embed_dim = Config.embedding_dim,\n)\n\n\ntrainer = pl.Trainer(\n    max_epochs = Config.epochs,\n    deterministic = True, #because of seed\n    gpus = torch.cuda.device_count(),\n    gradient_clip_val = 1.0,\n    accelerator = \"auto\",\n    precision = 16,\n  )\n\n\ntrain_dl, valid_dl = create_dl(\n    df_wiki,\n    image_transform = compose,\n    tokenizer = tokenizer\n)\n\n\ntrainer.fit(\n    model,\n    train_dl,\n    valid_dl\n)\n\ntrainer.save_checkpoint(\"clip_multi_step5.ckpt\")","metadata":{},"execution_count":null,"outputs":[]}]}