{"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":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport os\nimport random\nimport gc\nimport time\nimport time\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\nfrom tqdm.notebook import tqdm\nimport torch.optim.lr_scheduler as lr_scheduler\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader, Dataset\nfrom concurrent.futures import ThreadPoolExecutor\nimport gc\nimport math","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-21T15:34:23.681051Z","iopub.execute_input":"2023-04-21T15:34:23.681668Z","iopub.status.idle":"2023-04-21T15:34:27.176530Z","shell.execute_reply.started":"2023-04-21T15:34:23.681626Z","shell.execute_reply":"2023-04-21T15:34:27.175411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=3):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\nseed_everything()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:34:27.182811Z","iopub.execute_input":"2023-04-21T15:34:27.185691Z","iopub.status.idle":"2023-04-21T15:34:27.205401Z","shell.execute_reply.started":"2023-04-21T15:34:27.185649Z","shell.execute_reply":"2023-04-21T15:34:27.203736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_x = np.load(\"/kaggle/input/dataextraction/feature_data.npz\",allow_pickle=True)['arr_0']\ntrain_x_head = np.load('/kaggle/input/dataextraction/ratio.npz',allow_pickle=True)['arr_0']\ntrain_y = np.load('/kaggle/input/dataextraction/feature_labels.npz',allow_pickle=True)['arr_0']","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:34:27.206668Z","iopub.execute_input":"2023-04-21T15:34:27.207021Z","iopub.status.idle":"2023-04-21T15:34:58.944469Z","shell.execute_reply.started":"2023-04-21T15:34:27.206985Z","shell.execute_reply":"2023-04-21T15:34:58.943342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class InputNet(nn.Module):\n    def __init__(self , sample = False):\n        super().__init__()\n        self.max_length = CONF['max_length'] \n        self.sample  =sample\n  \n    def forward(self, xyz):\n        xyz = xyz[:self.max_length]\n        if xyz.shape[0]> 20 and self.sample:\n            random_frames = random.sample(list(range(xyz.shape[0])), (xyz.shape[0]//10)  * 9)\n            xyz = xyz[random_frames]\n            \n        xyz = xyz - xyz[~torch.isnan(xyz)].mean(0 , keepdim=True)\n        xyz = xyz / xyz[~torch.isnan(xyz)].std(0, keepdim=True)\n\n        xyz[torch.isnan(xyz)] = 0\n        \n        \n        \n        return xyz","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:34:58.949336Z","iopub.execute_input":"2023-04-21T15:34:58.949710Z","iopub.status.idle":"2023-04-21T15:34:58.957466Z","shell.execute_reply.started":"2023-04-21T15:34:58.949678Z","shell.execute_reply":"2023-04-21T15:34:58.956285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start = [0,0,0,0,0]\nend = [4,8,12,16,20]\nfor e in range(21,79):\n    start += [0,4,8,12,16,20]\n    end += [e]*6\nclass Distance(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n    def __call__(self,x):\n        st = x[:,start]\n        en = x[:,end]\n        \n        return torch.sum((en-st)**2,-1)**.5\n        ","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:34:58.959027Z","iopub.execute_input":"2023-04-21T15:34:58.959768Z","iopub.status.idle":"2023-04-21T15:34:58.970610Z","shell.execute_reply.started":"2023-04-21T15:34:58.959730Z","shell.execute_reply":"2023-04-21T15:34:58.969573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Distance()(torch.tensor(train_x[0]))","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:34:58.972166Z","iopub.execute_input":"2023-04-21T15:34:58.973003Z","iopub.status.idle":"2023-04-21T15:34:59.094245Z","shell.execute_reply.started":"2023-04-21T15:34:58.972971Z","shell.execute_reply":"2023-04-21T15:34:59.093334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lh_idx = list(range(21))\nrh_idx = list(range(21,21+21))\nupper_lip_index = list(range(21+21,21+21+11))\nlower_lip_index = list(range(21+21+11,21+21+11+11))\nleft_eye_index = list(range(21+21+11+11,21+21+11+11+14))\nright_eye_index = list(range(21+21+11+11+14,21+21+11+11+14+14))\ns_nose = list(range(21+21+11+11+14+14,21+21+11+11+14+14+3))\nright_chick_index = list(range(21+21+11+11+14+14+3,21+21+11+11+14+14+3+2))\nleft_chick_index = list(range(21+21+11+11+14+14+3+2,21+21+11+11+14+14+3+2+2))\n\nreversed_upper_lip = list(reversed(upper_lip_index))\nreversed_lower_lip = list(reversed(lower_lip_index))\n\n\nreverse_index = rh_idx+lh_idx+reversed_upper_lip+reversed_lower_lip+right_eye_index+left_eye_index+s_nose+left_chick_index+right_chick_index+[99]\nlen(reverse_index)\n\nfor i in range(len(train_x)):\n    ratio = train_x_head[i]\n    if ratio <0.5:\n        train_x[i] = train_x[i][:,reverse_index]\n    train_x[i] = train_x[i][:,21:]\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:34:59.095876Z","iopub.execute_input":"2023-04-21T15:34:59.096338Z","iopub.status.idle":"2023-04-21T15:35:01.893962Z","shell.execute_reply.started":"2023-04-21T15:34:59.096263Z","shell.execute_reply":"2023-04-21T15:35:01.892685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONF = {\n    \"Padding\" : 2,\n    \"max_length\" : 256,\n    \"fold\" : 0,\n    \"num_points\" : 62,\n    \"num_channel\" : 3,\n    \"num_class\" : 250,\n    \"vect_dim\" : 128 ,\n    \"vect_hidden_dim\" : 256,\n    \"xyz_dim\" : 6,\n    \"xyz_hidden_dim\" : 12,\n    \"emb_dim\" : 384,\n    \"batch_size\" : 256\n}","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:35:01.898843Z","iopub.execute_input":"2023-04-21T15:35:01.903205Z","iopub.status.idle":"2023-04-21T15:35:01.910721Z","shell.execute_reply.started":"2023-04-21T15:35:01.903157Z","shell.execute_reply":"2023-04-21T15:35:01.909594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_class = 250\nprep = InputNet(False)\n# \n\n    \n    \nclass ASLDataset(torch.utils.data.Dataset):\n    def __init__(self  ,indexs  ,sample = False ,transform=None,aug=True):\n        \n        self.transform = transform\n        self.indexs = indexs ; self.l = len(indexs)\n        self.input_prep = InputNet(sample)\n        self.aug=aug\n        self.fact = 1.5\n    def __len__(self):\n        return self.l\n\n    \n    def __getitem__(self,idx):\n\n        idx = self.indexs[idx]\n\n        X = train_x[idx][:,:,:2]\n        ratio = train_x_head[idx]\n        Y = train_y[idx]\n        \n        X = (X-X[:,[-1]])[:,:-1,:2]\n        \n        full_hand = train_x[idx][:,:21]\n        full_hand = self.input_prep(torch.tensor(full_hand))\n        \n        full_hand[torch.isnan(full_hand)] = 0\n        \n        motion = torch.cat([torch.zeros((1,full_hand.shape[1],full_hand.shape[2])),full_hand[1:,:,:]-full_hand[:-1,:,:]],0)\n#         print((full_hand[1:,:,:]-full_hand[:-1,:,:]).shape)\n        full_hand = torch.cat([full_hand,motion],-1)\n        \n#         X = np.concatenate([X,head],1)\n#         full_hand = torch.cat([full_hand,motion],-1)\n        \n        X = self.input_prep(torch.tensor(X))\n        \n       \n        if ratio<0.5:\n            X[:,:,0] *= -1\n            full_hand[:,:,0] *= -1\n            \n        if self.aug:  \n            X = X + X*(torch.rand(X.shape)-0.5)/(100*self.fact)\n            X = X * ((torch.rand(1)-0.5)/(60*self.fact)+1)\n#         print(torch.mean(X[:,reversed_lip],1).unsqueeze(1).shape)\n#         print(X[:,lh_idx+rh_idx].shape)\n\n\n#\n        X[torch.isnan(X)] = 0\n        \n        motion = torch.cat([torch.zeros((1,X.shape[1],X.shape[2])),X[1:,:,:]-X[:-1,:,:]],0)\n#         print((X[1:,:,:]-X[:-1,:,:]).shape)\n        X = torch.cat([X,motion],-1)\n        return X , full_hand ,torch.tensor(Y).to(torch.long)\n    def get_labels(self):\n        return train_y[self.indexs]","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:37:45.778263Z","iopub.execute_input":"2023-04-21T15:37:45.778908Z","iopub.status.idle":"2023-04-21T15:37:45.811343Z","shell.execute_reply.started":"2023-04-21T15:37:45.778862Z","shell.execute_reply":"2023-04-21T15:37:45.810325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/asl-folded/train_prepared.csv\")\nto_take = [True if (train_x_head[i]<0.4 or train_x_head[i]>0.6) else False for i in range(len(train_x))]\ndf['to_take'] = to_take\ndf = df[df['to_take']==True]\nindexs =df [['fold']].reset_index()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:37:46.424379Z","iopub.execute_input":"2023-04-21T15:37:46.424746Z","iopub.status.idle":"2023-04-21T15:37:47.093832Z","shell.execute_reply.started":"2023-04-21T15:37:46.424711Z","shell.execute_reply":"2023-04-21T15:37:47.092763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dense(nn.Module):\n    def __init__(\n        self , \n        in_,\n        out_,\n        bias,\n        dp=0.1,\n    ):\n        super(Dense, self).__init__()\n        \n        self.ff = nn.Sequential(*[\n            nn.Linear(in_,out_,bias),\n            nn.LayerNorm(out_),\n            nn.SiLU(inplace=True)\n        ])\n        \n    def forward  (self,x):\n        return self.ff(x)\n    \n\n    \n#feed forward for transformer\nclass FeedForward(nn.Module):\n    def __init__(self, embed_dim, hidden_dim):\n        super().__init__()\n        self.mlp = nn.Sequential(\n            nn.Linear(embed_dim, hidden_dim),\n            nn.SiLU(inplace=True),\n            nn.Linear(hidden_dim, embed_dim),\n        )\n    def forward(self, x):\n        return self.mlp(x)\n    \ndef positional_encoding(length, embed_dim):\n    dim = embed_dim//2\n\n    position = np.arange(length)[:, np.newaxis]     # (seq, 1)\n    dim = np.arange(dim)[np.newaxis, :]/dim   # (1, dim)\n\n    angle = 1 / (10000**dim)         # (1, dim)\n    angle = position * angle    # (pos, dim)\n\n    pos_embed = np.concatenate(\n        [np.sin(angle), np.cos(angle)],\n        axis=-1\n    )\n    pos_embed = torch.from_numpy(pos_embed).float()\n    return pos_embed","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:37:48.005082Z","iopub.execute_input":"2023-04-21T15:37:48.005453Z","iopub.status.idle":"2023-04-21T15:37:48.016611Z","shell.execute_reply.started":"2023-04-21T15:37:48.005419Z","shell.execute_reply":"2023-04-21T15:37:48.015538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## pacuing for training\ndef pack_seq(\n    seq,\n    \n):\n    length = [len(s) for s in seq]\n    batch_size = len(seq)\n    \n    num_landmark = seq.shape[2]\n    shape = list(seq.shape)\n    shape[1] = max(length)\n    x = torch.zeros(shape).to(seq[0].device)\n        \n    x_mask = torch.zeros((batch_size, max(length))).to(seq[0].device)\n    \n    for b in range(batch_size):\n        L = length[b]\n        x[b, :L] = seq[b][:L]\n        x_mask[b, L:] = 1\n    x_mask = (x_mask>0.5)\n    \n    return x, x_mask\n\n\n\n### multi head\nclass MultiHeadAttention(nn.Module):\n    def __init__(self,\n            embed_dim,\n            num_head,\n            batch_first,\n        ):\n        super().__init__()\n        self.mha = nn.MultiheadAttention(\n            embed_dim,\n            num_heads=num_head,\n            bias=True,\n            add_bias_kv=False,\n            kdim=None,\n            vdim=None,\n            dropout=0.0,\n            batch_first=batch_first,\n        )\n\n    def forward(self, x, x_mask):\n        out, _ = self.mha(x,x,x, key_padding_mask=x_mask)\n        return out\n    \n\n\n\nclass TransformerBlock(nn.Module):\n    def __init__(self,\n        embed_dim,\n        num_head,\n        out_dim,\n        batch_first=True,\n    ):\n        super().__init__()\n        self.attn  = MultiHeadAttention(embed_dim, num_head,batch_first)\n        self.ffn   = FeedForward(embed_dim, out_dim)\n        self.norm1 = nn.LayerNorm(embed_dim)\n        self.norm2 = nn.LayerNorm(out_dim)\n\n    def forward(self, x, x_mask=None):\n        x = x + self.attn((self.norm1(x)), x_mask)\n        x = x + self.ffn((self.norm2(x)))\n        return x\n    \n\n\n    \nclass PointsEmbedder(nn.Module):\n    def __init__(\n        self , \n        vect_dim , \n        vect_hidden_dim,\n        emb_dim,\n        dp = 0.25,\n    ):\n        super(PointsEmbedder, self).__init__()\n\n        \n        self.vect_dim = vect_dim\n        self.vect_hidden_dim = vect_hidden_dim\n        \n        self.emb_dim = emb_dim\n        self.input_prep = InputNet(CONF['max_length'])\n        \n#         self.vector_extractor = Dense(CONF[\"num_points\"] , vect_dim , False , dp)\n        \n        self.vectorizer = nn.Sequential(*[\n            Dense(CONF['num_points'] , CONF['num_points']*4 ,True ),\n            Dense(CONF['num_points']*4 , CONF['num_points']*4 ,True ),\n        ])\n        \n        self.hand = nn.Sequential(*[\n            Dense(21 ,42 ,True ),\n            Dense(42 , 42 ,True ),\n        ])\n        \n        self.embedder = nn.Sequential(*[\n            Dense(CONF['num_points']* 4 * CONF['num_channel']  + 42*6, emb_dim*2 ,True ),\n            Dense(emb_dim*2 , emb_dim ,True ),\n        ])\n        \n        \n        \n    def forward (self,x , hand):\n        \n        x = self.vectorizer(x.transpose(-1,-2))\n        hand = self.hand(hand.transpose(-1,-2))\n        hand = hand.reshape(hand.shape[0],-1,42*6)\n        x = x.reshape(x.shape[0],-1,CONF['num_points']*4 * CONF['num_channel'])\n        \n        x = torch.cat([x,hand],-1)\n        return self.embedder(x)\n    \nclass LandmarkTransformerEmbedder(nn.Module):\n    def __init__(self, dp = 0.2,in_features=256,num_block=1, num_head=8):\n        super().__init__()\n        self.num_block = num_block\n        self.num_head  = num_head\n        self.emb_dim = in_features\n        self.dp = nn.Dropout(dp)\n        \n        pos_embed = positional_encoding(CONF['max_length'], self.emb_dim)\n        # self.register_buffer('pos_embed', pos_embed)\n        self.pos_embed = nn.Parameter(pos_embed)\n        self.cls_embed = nn.Parameter(torch.zeros((1, self.emb_dim)))\n        \n        \n        self.encoder = nn.ModuleList([\n            TransformerBlock(\n                self.emb_dim,\n                self.num_head,\n                self.emb_dim,\n            ) for i in range(self.num_block)\n        ])\n        \n        \n    def forward(self,batch):\n        \n        length = [len(x) for x in batch]\n        xyz = batch\n        x, x_mask = pack_seq(xyz)\n        B,L,_ = x.shape\n        \n        x = x + self.pos_embed[:L].unsqueeze(0)\n        \n        \n\n        x = torch.cat([\n            self.cls_embed.unsqueeze(0).repeat(B,1,1),\n            x\n        ],1)\n        x_mask = torch.cat([\n            torch.zeros(B,1).to(x_mask),\n            x_mask\n        ],1)\n   \n\n\n        #x = F.dropout(x,p=0.25,training=self.training)\n        for block in self.encoder:\n            x = block(x,x_mask)\n\n#         cls = x[:,0]\n        x = self.dp(x)\n        x_mask = x_mask.unsqueeze(-1)\n        x_mask = 1-x_mask.float()\n        last = (x*x_mask).sum(1)/x_mask.sum(1)\n        \n\n\n        return last\n\nclass Net(nn.Module):\n\n    def __init__(self, num_class=CONF['num_class']):\n        super().__init__()\n        \n        \n        self.marks_feat_r = PointsEmbedder(CONF['vect_dim'] ,CONF['vect_hidden_dim'], CONF['emb_dim'])\n        #self.lip_net = LandMarksLinearNet(41,256)\n        \n        self.hand_transformer = LandmarkTransformerEmbedder(0.5 , CONF['emb_dim'] )\n        \n        #self.hand_masker_transformer = LandmarkTransformerEmbedder(1 , emb_dim)\n        #self.mid_embedder = LinearN(emb_dim,emb_dim-1)\n        \n        \n        \n        self.logit = nn.Sequential(*[\n            nn.Linear(CONF['emb_dim'] ,CONF['num_class'] ),\n        ])\n    \n    def reset_dp(self,dp):\n        self.hand_transformer.dp = nn.Dropout(dp)\n    \n    def embed(self, batch_r,hand):\n        emb_r = self.marks_feat_r(batch_r,hand)\n        \n        emb = self.hand_transformer(emb_r)\n        \n        return emb\n    def forward(self,batch_r,hand):\n        return self.logit(self.embed(batch_r,hand))","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:34.786514Z","iopub.execute_input":"2023-04-21T15:39:34.786936Z","iopub.status.idle":"2023-04-21T15:39:35.027917Z","shell.execute_reply.started":"2023-04-21T15:39:34.786898Z","shell.execute_reply":"2023-04-21T15:39:35.026733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_dataset = ASLDataset(list(indexs[indexs['fold'] != 0]['index']))\n# train_dataset[0][0].shape[2]","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:35.030731Z","iopub.execute_input":"2023-04-21T15:39:35.031448Z","iopub.status.idle":"2023-04-21T15:39:35.044343Z","shell.execute_reply.started":"2023-04-21T15:39:35.031405Z","shell.execute_reply":"2023-04-21T15:39:35.043331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X = train_x[0]\nX = (X-X[:,[-1]])[:,:42,:2]\nprint(X.shape)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:35.188692Z","iopub.execute_input":"2023-04-21T15:39:35.189465Z","iopub.status.idle":"2023-04-21T15:39:35.196553Z","shell.execute_reply.started":"2023-04-21T15:39:35.189424Z","shell.execute_reply":"2023-04-21T15:39:35.195422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = ASLDataset(list(indexs[indexs['fold'] != 0]['index']))\ntrain_dataset[0][1].shape","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:35.365937Z","iopub.execute_input":"2023-04-21T15:39:35.366266Z","iopub.status.idle":"2023-04-21T15:39:35.385318Z","shell.execute_reply.started":"2023-04-21T15:39:35.366222Z","shell.execute_reply":"2023-04-21T15:39:35.384202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONF['num_points'] = train_dataset[0][0].shape[1]\nCONF['num_channel'] = train_dataset[0][0].shape[2]","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:35.503874Z","iopub.execute_input":"2023-04-21T15:39:35.504148Z","iopub.status.idle":"2023-04-21T15:39:35.511900Z","shell.execute_reply.started":"2023-04-21T15:39:35.504122Z","shell.execute_reply":"2023-04-21T15:39:35.510652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn.utils.rnn import pad_sequence\n\ndef collate_fn(batch):\n    \"\"\"\n    Collate function to pad variable-length sequences with zeros.\n    \"\"\"\n    \n    X_left = pad_sequence([batch[i][0] for i in range(len(batch))]  , True)\n    hand = pad_sequence([batch[i][1] for i in range(len(batch))]  , True)\n    y = [batch[i][2] for i in range(len(batch))]\n    # Stack the padded sequences and the labels\n    \n    return X_left,hand ,torch.tensor(y)\nfrom torch.utils.data import DataLoader, Dataset\n\n\ntry :\n    from torchsampler import ImbalancedDatasetSampler\nexcept :\n    !pip install torchsampler\n    from torchsampler import ImbalancedDatasetSampler\n\ntrain_dataset = ASLDataset(list(indexs[indexs['fold'] != 0]['index']))\nval_dataset = ASLDataset(list(indexs[indexs['fold'] == 0]['index']),aug = False)\n\n\n\nval_dataloader = DataLoader(val_dataset,collate_fn = collate_fn ,batch_size=CONF[\"batch_size\"])\ntrain_dataloader = DataLoader(train_dataset,collate_fn = collate_fn ,batch_size=64, sampler=ImbalancedDatasetSampler(train_dataset),)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:35.715965Z","iopub.execute_input":"2023-04-21T15:39:35.716346Z","iopub.status.idle":"2023-04-21T15:39:35.778996Z","shell.execute_reply.started":"2023-04-21T15:39:35.716289Z","shell.execute_reply":"2023-04-21T15:39:35.777866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try :\n    import torch_optimizer as optim\nexcept :\n    !pip install torch_optimizer\n    import torch_optimizer as optim","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:35.955602Z","iopub.execute_input":"2023-04-21T15:39:35.956618Z","iopub.status.idle":"2023-04-21T15:39:35.962084Z","shell.execute_reply.started":"2023-04-21T15:39:35.956576Z","shell.execute_reply":"2023-04-21T15:39:35.961070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LabelSmoothingLoss(torch.nn.Module):\n    def __init__(self, smoothing: float = 0.1, \n                 reduction=\"mean\", weight=None):\n        super(LabelSmoothingLoss, self).__init__()\n        self.smoothing   = smoothing\n        self.reduction = reduction\n        self.weight    = weight\n\n    def reduce_loss(self, loss):\n        return loss.mean() if self.reduction == 'mean' else loss.sum() \\\n         if self.reduction == 'sum' else loss\n\n    def linear_combination(self, x, y):\n        return self.smoothing * x + (1 - self.smoothing) * y\n\n    def forward(self, preds, target):\n        assert 0 <= self.smoothing < 1\n\n        if self.weight is not None:\n            self.weight = self.weight.to(preds.device)\n\n        n = preds.size(-1)\n        log_preds = F.log_softmax(preds, dim=-1)\n        loss = self.reduce_loss(-log_preds.sum(dim=-1))\n        nll = F.nll_loss(\n            log_preds, target, reduction=self.reduction, weight=self.weight\n        )\n        return self.linear_combination(loss / n, nll)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:36.213874Z","iopub.execute_input":"2023-04-21T15:39:36.214215Z","iopub.status.idle":"2023-04-21T15:39:36.225595Z","shell.execute_reply.started":"2023-04-21T15:39:36.214183Z","shell.execute_reply":"2023-04-21T15:39:36.224541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dps =( np.array(list(range(0,10)))/210).tolist() + ( np.array(list(range(35,70)))/120).tolist()+[1/3]*200","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:36.447372Z","iopub.execute_input":"2023-04-21T15:39:36.447654Z","iopub.status.idle":"2023-04-21T15:39:36.453874Z","shell.execute_reply.started":"2023-04-21T15:39:36.447627Z","shell.execute_reply":"2023-04-21T15:39:36.452602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"time_didnt_get_better = 0\nlearning_rate = 0.0025\ndef train_epoch(model , epoch , scheduler):\n    global learning_rate\n    model.train()\n#     if time_didnt_get_better>10:\n#         return False\n    if epoch<30:\n        optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate , weight_decay = 0.00001)\n        optimizer = optim.Lookahead(optimizer, k=5, alpha=0.5)\n    else :\n        optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate , weight_decay = 0.0001)\n        optimizer = optim.Lookahead(optimizer, k=5, alpha=0.5)\n    losses = []\n    if (epoch %2 == 0 and epoch>10):\n#         train_dataloader.dataset.augment()\n        learning_rate = learning_rate *.95\n    for i , d in tqdm(enumerate(train_dataloader) , total = len(train_dataloader) , postfix = np.array(losses).mean()):\n        \n        sequences_l = d[0].to(device)\n        hand = d[1].to(device)\n\n        labels = d[2].to(device)\n        outputs = model(sequences_l,hand)\n        loss = criterion(outputs, torch.tensor(labels,dtype=torch.long))\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        losses.append(loss.item())\n    print(f'Epoch [{epoch+1}/{num_epochs}] , Loss: {np.array(losses).mean():.4f}')\n    scheduler.step()\n    return np.array(losses).mean()\n        \n    \n#     return True\ndef val_epoch(model , epoch):\n    global time_didnt_get_better,best,train_dataloader\n    model.eval()      \n    losses = []\n    with torch.no_grad():\n            correct = 0\n            total = 0\n            for i , d in tqdm(enumerate(val_dataloader) , total = len(val_dataloader)):\n                \n                sequences_l = d[0].to(device)\n                hand = d[1].to(device)\n\n                labels = d[2].to(device)\n                outputs = model(sequences_l,hand)\n                loss = criterion(outputs, torch.tensor(labels,dtype=torch.long))\n                losses.append(loss.item())\n                _, predicted = torch.max(outputs, 1)\n                total += labels.size(0)\n                correct += (predicted == labels).sum().item()\n            if (correct / total)>=best:\n                time_didnt_get_better =0 \n                torch.save(model.state_dict(),\"model_0.pth\")\n                best = (correct / total)\n            else : \n                time_didnt_get_better+=1\n            print(f'Epoch [{epoch+1}/{num_epochs}] , Test Accuracy: {(correct / total) * 100:.2f} %     {np.array(losses).mean()}')\n    return np.array(losses).mean()\ndef do_epoch(model,epoch,scheduler):\n    c = train_epoch(model , epoch,scheduler)\n    global train_dataloader,train_dataset\n    model.reset_dp(dps[epoch])\n    if epoch==20:\n        rain_dataloader = DataLoader(train_dataset,collate_fn = collate_fn ,batch_size=64, sampler=ImbalancedDatasetSampler(train_dataset),)\n        train_dataloader.dataset.fact = 2\n\n    if epoch==40:\n        train_dataloader = DataLoader(train_dataset,collate_fn = collate_fn ,batch_size=32, sampler=ImbalancedDatasetSampler(train_dataset),)\n        train_dataloader.dataset.fact = 4\n    if epoch==50:\n        train_dataloader.dataset.fact = 5\n    if epoch == 60:\n        train_dataloader = DataLoader(train_dataset,collate_fn = collate_fn ,batch_size=16, sampler=ImbalancedDatasetSampler(train_dataset),)\n        train_dataloader.dataset.fact = 8\n        \n    if epoch == 80:\n        train_dataloader = DataLoader(train_dataset,collate_fn = collate_fn ,batch_size=32, sampler=ImbalancedDatasetSampler(train_dataset),)\n        train_dataloader.dataset.fact = 9\n    if epoch == 90:\n        train_dataloader = DataLoader(train_dataset,collate_fn = collate_fn ,batch_size=16, sampler=ImbalancedDatasetSampler(train_dataset),)\n        train_dataloader.dataset.fact = 10   \n        \n    if epoch%3==0 and epoch >0:\n        c_val = val_epoch(model , epoch)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:49:34.758526Z","iopub.execute_input":"2023-04-21T15:49:34.758913Z","iopub.status.idle":"2023-04-21T15:49:34.780990Z","shell.execute_reply.started":"2023-04-21T15:49:34.758881Z","shell.execute_reply":"2023-04-21T15:49:34.777269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:49:47.538700Z","iopub.execute_input":"2023-04-21T15:49:47.539250Z","iopub.status.idle":"2023-04-21T15:50:22.128248Z","shell.execute_reply.started":"2023-04-21T15:49:47.539203Z","shell.execute_reply":"2023-04-21T15:50:22.127090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset[0][0].shape[2]\ntrain_dataset[0][0].shape[1]","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:37.875913Z","iopub.execute_input":"2023-04-21T15:39:37.876486Z","iopub.status.idle":"2023-04-21T15:39:37.889259Z","shell.execute_reply.started":"2023-04-21T15:39:37.876439Z","shell.execute_reply":"2023-04-21T15:39:37.888355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nCONF[\"emb_dim\"] = 256+64","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:38.185209Z","iopub.execute_input":"2023-04-21T15:39:38.185556Z","iopub.status.idle":"2023-04-21T15:39:38.190733Z","shell.execute_reply.started":"2023-04-21T15:39:38.185525Z","shell.execute_reply":"2023-04-21T15:39:38.189481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONF['num_points']","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:38.231535Z","iopub.execute_input":"2023-04-21T15:39:38.232461Z","iopub.status.idle":"2023-04-21T15:39:38.238610Z","shell.execute_reply.started":"2023-04-21T15:39:38.232429Z","shell.execute_reply":"2023-04-21T15:39:38.237494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_dataloader.dataset.augment()\nnum_classes = 250\nnum_epochs = 100\nbatch_size = 64\nlearning_rate = 0.00033\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nbest = 0\n\nmodel = Net(CONF[\"emb_dim\"] ).to(device)\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=0.001 , weight_decay = 0.00001)\nscheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\ncriterion = LabelSmoothingLoss(0.65)\nratio = 0.6\nfor epoch in range(num_epochs):\n    do_epoch(model,epoch,scheduler)\n    \n# 60.45 %     4.782241629458022","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:39:38.931191Z","iopub.execute_input":"2023-04-21T15:39:38.933574Z","iopub.status.idle":"2023-04-21T15:48:59.389507Z","shell.execute_reply.started":"2023-04-21T15:39:38.933535Z","shell.execute_reply":"2023-04-21T15:48:59.387896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for epoch in range(num_epochs):\n#     do_epoch(model,19,scheduler)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:35:32.795830Z","iopub.status.idle":"2023-04-21T15:35:32.796454Z","shell.execute_reply.started":"2023-04-21T15:35:32.796156Z","shell.execute_reply":"2023-04-21T15:35:32.796183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Epoch [1/70] , Test Accuracy: 42.00 %     5.0306192485765475\n# 100%\n# 282/282 [02:37<00:00, 1.74it/snan]\n# Epoch [2/70] , Loss: 4.8729\n# Epoch [2/70] , Loss: 4.8413\n# 100%\n# 87/87 [00:27<00:00, 3.25it/s]\n# Epoch [2/70] , Test Accuracy: 50.55 %     4.916942601916434\n# 100%\n# 282/282 [02:37<00:00, 1.81it/snan]\n# Epoch [3/70] , Loss: 4.7459\n# Epoch [3/70] , Loss: 4.7282\n# 100%\n# 87/87 [00:28<00:00, 3.19it/s]\n# Epoch [3/70] , Test Accuracy: 55.81 %     4.853009722698694\n# 100%\n# 282/282 [02:35<00:00, 2.17it/snan]\n# Epoch [4/70] , Loss: 4.6843\n# Epoch [4/70] , Loss: 4.6746\n# 100%\n# 87/87 [00:27<00:00, 3.23it/s]\n# Epoch [4/70] , Test Accuracy: 57.20 %     4.826726107761778\n# 100%\n# 282/282 [02:35<00:00, 1.95it/snan]\n# Epoch [5/70] , Loss: 4.6485\n# Epoch [5/70] , Loss: 4.6402\n# 100%\n# 87/87 [00:27<00:00, 3.16it/s]\n# Epoch [5/70] , Test Accuracy: 58.11 %     4.816454903832797\n# 100%\n# 282/282 [02:38<00:00, 1.61it/snan]\n# Epoch [6/70] , Loss: 4.6196\n# Epoch [6/70] , Loss: 4.6135\n# 100%\n# 87/87 [00:27<00:00, 3.07it/s]\n# Epoch [6/70] , Test Accuracy: 58.90 %     4.801272490928913\n# 63%\n# 177/282 [01:40<00:59, 1.77it/snan]\n# Epoch [7/70] , Loss: 4.6023","metadata":{"execution":{"iopub.status.busy":"2023-04-21T15:35:32.798334Z","iopub.status.idle":"2023-04-21T15:35:32.799204Z","shell.execute_reply.started":"2023-04-21T15:35:32.798931Z","shell.execute_reply":"2023-04-21T15:35:32.798959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}