{"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":"# Fingerspelling Experiment\n\n## Intrinsic Learning\n\nI am unaware of another term for the way I train the network, but the general concept is that I train a model to predict the distribution of characters in a phrase, by summing and normalizing the output of the model for each input frame of the signed phrase. This allows us to remove the sequential components from both the input and the output. Because the output distribution is a linear combination of normally distributed embeddings, the only way to predict the distribution correctly is to predict the correct embedding for each input frame. Because the model needs to be consistent / coherent internally it will automatically (intrinsically) map the correct embeddings to each frame, given enough input (with different distributions of characters). Because the bias of the input pretty much matches the bias of the output, it produces quite balanced results.\n\n## Auto Segmentation\n\nI created an algorithm that should optimally segment any sequence of embeddings into a smaller set, where subsequent frames are merged if they are correlated. The algorithm automatically tries to find the optimal boundaries of the segments by selecting the 2 segments that divide any window into 2 higher density windows and applies this recursively.\n\n## Complex Transformer\n\nA complex transformer doesn't work if you implement it to the letter, simply substituting real layers / functions for their complex counterparts. This is because the softmax over the self attention matrix is technically defined for complex numbers but doesn't produce results that are appropriate for what it needs to do. We therefore substitute the `softmax(z)` with a `softmax(sign(z.real) * z.abs()) * (z/z.abs())`. This calculates the softmax as if each value first gets rotated into the real plane, and then rotated back after the softmax.\n\n**edit:** At this time I don't think the transformer is helping much, because I removed a loss that created a pseudo ground truth by mapping the phrase to the predicted output (which had its own problems). There is no incentive for the transformer to keep the characters in order. So I turned it off and running solely on the conv / resnet.","metadata":{"execution":{"iopub.status.busy":"2023-07-30T15:50:33.66185Z","iopub.execute_input":"2023-07-30T15:50:33.662251Z","iopub.status.idle":"2023-07-30T15:50:33.682994Z","shell.execute_reply.started":"2023-07-30T15:50:33.662215Z","shell.execute_reply":"2023-07-30T15:50:33.682115Z"}}},{"cell_type":"code","source":"# File config.py\n# Created by E/S Pronk\n\nimport os\nimport glob\nimport json\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.multiprocessing as mp\nimport torch.distributed as dist\nimport numpy as np\nimport pandas as pd\nimport pyarrow.parquet as pq\nimport matplotlib.pyplot as plt\nimport tqdm.notebook as tqdm\nimport complex as Z\nimport algorithm as alg\nimport reduce\nimport inline\nimport scheduler as sched\nfrom collections.abc import Iterable\n\ntorch.manual_seed(13)\ntorch.cuda.manual_seed(13)\n\nclass Config(dict):\n    def __init__(self):\n        super().__init__(\n            dict(\n                max_duration = 60 * 5,\n                input_path = \"/kaggle/input/asl-fingerspelling\",\n                output_path = \"/kaggle/working\",\n                checkpoint_path = \"cp-{hidden_dim}-{expand_dim}-L{num_project_layers:02d}{num_encoder_layers:02d}-{num_heads:01d}.cuda.{epoch:03d}.pth\",\n                hidden_dim = 80, \n                sequence_dim=1024,\n                epochs = 10, \n                expand_dim = 160,\n                num_project_layers=3,\n                num_encoder_layers=0,\n                num_heads=2,\n                num_prototypes=59,\n                cpu_count = mp.cpu_count(),\n                gpu_count = torch.cuda.device_count() if torch.cuda.is_available() else 0,\n                device = \"cuda\" if torch.cuda.is_available() else \"cpu\",\n                cleanup = False,\n                flush_cache = False,\n                dataset_repeat = 1,\n                num_train_shards = 48,\n                num_test_shards = 16,\n                shard_size = 960,\n                phrases = {},\n                sequence_to_file_id = {},\n                num_workers=2, # 12 per GPU if enough cpus available\n                batch_size=24,\n                worker_yields = 960,\n                dist_enabled = False, # set to true for best performance using DistributedDataParallel\n                dist_worldsize = os.environ[\"DIST_WORLDSIZE\"] if \"DIST_WORLDSIZE\" in os.environ else (torch.cuda.device_count() if torch.cuda.is_available() else 1),\n                dist_port = os.environ[\"DIST_PORT\"] if \"DIST_PORT\" in os.environ else 23456,\n                dist_hostname = os.environ[\"DIST_HOSTNAME\"] if \"DIST_HOSTNAME\" in os.environ else \"localhost\",\n                rank = dist.get_rank() if dist.is_initialized() else None,\n                localrank = dist.get_rank() % torch.cuda.device_count() if (torch.cuda.is_available() and dist.is_initialized()) else None,\n                columns = [\"\".join(c) for c in inline.all_combinations(\n                    [\"y_\",\"x_\"], \n                    [\"left_hand_\",\"right_hand_\"], \n                    [str(i) for i in range(21)])]\n            ))\n\n        with open(os.path.join(self[\"input_path\"], \"character_to_prediction_index.json\")) as stream:\n            self[\"index\"] = json.load(stream)\n            self[\"index_rev\"] = { value: key for key,value in self[\"index\"].items() }\n\n        df = pd.read_csv(os.path.join(self[\"input_path\"], \"train.csv\"))\n        for name, row in df.iterrows():\n            self[\"phrases\"][row[\"sequence_id\"]] = row[\"phrase\"]\n            self[\"sequence_to_file_id\"][row[\"sequence_id\"]] = row[\"file_id\"]\n\n    def decode(self,ind):\n        return \"\".join([self[\"index_rev\"].get(c.item(),\"\") for c in ind.cpu()])\n\n    def encode(self,s):\n        return torch.tensor([self[\"index\"].get(c,\"\") for c in s], device=self[\"device\"])\n","metadata":{"_kg_hide-input":false,"_kg_hide-output":false,"execution":{"iopub.status.busy":"2023-08-24T13:57:44.642604Z","iopub.execute_input":"2023-08-24T13:57:44.643303Z","iopub.status.idle":"2023-08-24T13:57:51.573442Z","shell.execute_reply.started":"2023-08-24T13:57:44.643256Z","shell.execute_reply":"2023-08-24T13:57:51.571706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset\n\nWhen you have a lot of workers and a very large dataset, split into different shards, that needs a lot of memory per shard: it is not optimal to use the default route of having a single dataset that is loaded by each workers, and a sampler that selects only relevant entries from the dataset for each worker. If each worker loads the same dataset, and you have N workers, the data loaded in memory by the dataset is N times redundant. Each worker will only return 1/Nth of the data that is in the dataset, and after the shard is handled by the worker group, a new shard needs to be loaded by each worker. This is a lot of overhead.\n\nIt woud be much better if each worker loads a separate shard, and return each entry in the shard, with the dataloader aggegrating the data for multple workers. That way each shard is only loaded into memory once, there is no redundancy and no need to reload shards as often.\n\nDownside is that it is a bit more complex to set up, but nothing that isn't worth the benefit on large multi-gpu systems with 100+ cpus","metadata":{}},{"cell_type":"code","source":"# File: dataset.py\n# Created by E/S Pronk\n\nimport os\nimport glob\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport pyarrow.parquet as pq\nimport tqdm as tqdm\nimport inline\nimport reduce\nimport complex as Z\n\nclass ASLDataset(torch.utils.data.IterableDataset):\n    def __init__(self, config, subset):\n        self.config=config\n        self.worker_yields =  int(config[\"shard_size\"] * (subset[1] - subset[0]) / config[\"num_workers\"]) # fix this\n        self.subset = subset\n        self.num_workers = config[\"num_workers\"]\n        self.repeat = config[\"dataset_repeat\"]\n    \n    \n    def create_cache(self):\n        parquet_paths = sorted(glob.glob(os.path.join(self.config[\"input_path\"], \"train_landmarks\", \"*.parquet\")))[self.subset[0]:self.subset[1]]\n        cache_paths = [os.path.join(self.config[\"output_path\"], os.path.split(path)[1] + \".cache\") for path in parquet_paths]\n        for path, fn in zip(parquet_paths,cache_paths):\n            if os.path.exists(fn) and not self.config[\"flush_cache\"]:\n                continue\n            else:\n                part = {}\n                columns = self.config[\"columns\"]\n                df = pq.read_table(path, columns=[\"sequence_id\"] + columns).to_pandas()\n                for name, row in df.iterrows():\n                    t = torch.from_numpy(row.to_numpy())\n                    # Remove nan frames\n                    if torch.isnan(t).all(): continue\n                    if name not in part: part[name] = [t]\n                    else: part[name].append(t)\n\n                part = {key: torch.stack(p) for key,p in part.items()}\n                torch.save(part, fn)\n                print(f\"created cache {fn}\")\n    \n    def find(self, name):\n        file_id = self.config[\"sequence_to_file_id\"][name]\n        cache_path = os.path.join(self.config[\"output_path\"], str(file_id) + \".parquet.cache\")\n        shard = torch.load(cache_path)\n        phrase = self.config[\"phrases\"][name]\n        target = torch.tensor([self.config[\"index\"][c] for c in phrase])\n        sequence = (name, phrase, shard[name], target)\n        return sequence \n    \n    def __len__(self):\n        return self.num_workers * self.worker_yields\n\n    def __iter__(self):\n        pattern = os.path.join(self.config[\"output_path\"],\"*.parquet.cache\")\n        cache_paths = sorted(glob.glob(pattern))[self.subset[0]:self.subset[1]] * max(self.repeat,1)\n        info = torch.utils.data.get_worker_info()\n        assert info is not None, \"for now\"\n        assert self.num_workers == info.num_workers\n\n        i = self.worker_yields\n        while i > 0:\n            if info.id >= len(cache_paths): return\n            for shard_idx in range(info.id, len(cache_paths), info.num_workers):\n                cache_path = cache_paths[shard_idx]\n                shard = torch.load(cache_path)\n                keys = list(shard.keys())\n                perm = torch.randperm(len(shard))[:i]\n                i -= len(perm)\n                for idx in perm:\n                    name = keys[idx]\n                    phrase = self.config[\"phrases\"][name]\n                    target = torch.tensor([self.config[\"index\"][c] for c in phrase])\n                    sequence = (name, phrase, shard[name], target)\n                    yield sequence\n\ndef collate_pad(x):\n    return [v[0] for v in x], \\\n           [v[1] for v in x], \\\n           inline.padstack([v[2] for v in x], pad_value=math.nan), \\\n           inline.padstack([v[3] for v in x], pad_value=-1)\n\nclass Preprocess(nn.Module):\n    '''\n    Turn the input row into a tensor where the x and y values represent complex numbers y+xi,\n    both hands are normalized, and the conjugate of the right hand is added to the left.\n    '''\n    def __init__(self, in_size=512, mode=\"none\", pad_value=0, interpolation_mode=\"linear\"):\n        super().__init__()\n        self.in_size = in_size\n        self.mode = mode\n        self.pad_value = pad_value\n        self.interpolation_mode = interpolation_mode\n\n    def forward(self, x):\n        if not isinstance(x, torch.Tensor):\n            if not isinstance(x, np.ndarray):\n                if not isinstance(x, pd.DataFrame):\n                    x = x.to_pandas()\n                x = x.to_numpy()\n            x = torch.from_numpy(x)\n        assert x.shape[-1] == 84\n        if x.dim() < 3: x = x[None]\n\n        if self.mode == \"none\":\n            in_size = x.shape[-2]\n            pad = 16 - (in_size % 16)\n            x = F.pad(x, (0,0,0,pad), \"constant\", math.nan)\n            in_size += pad\n\n        elif self.mode == \"batch_interpolate\":\n            x = F.interpolate(x.mT, (self.in_size,), mode=self.interpolation_mode).mT\n            in_size = self.in_size\n        else:\n            r = []\n            for v in x:\n                v = v[~torch.isnan(v).all(-1)]\n                if self.mode == \"interpolate\":\n                    v = F.interpolate(v[None].mT, (self.in_size,), mode=self.interpolation_mode)[0].mT\n                elif self.mode == \"pad\":\n                    if v.shape[-2] > self.in_size:\n                        print(f\"\\033[33mWarning: sequence length ({v.shape[-2]}) > in_size ({self.in_size}), cropping!\\033[0m\")\n                        v = v[:self.in_size]\n                    else:\n                        v = F.pad(v, (0,0,0,self.in_size - v.shape[-2]), \"constant\", self.pad_value)\n                else:\n                    assert False, \"invalid preprocessing mode\"\n                r.append(v)\n            x = torch.stack(r)\n            in_size = self.in_size\n\n        # Separate the hands into left and right, move x,y to the last dim\n        sls = alg.sequence_lengths(x)\n        hand = x.view(-1,in_size,2,2,21).permute(0,1,3,4,2)\n        nans = torch.isnan(hand)\n        hand = torch.nan_to_num(hand)\n\n        # Calculate the mean and std for each landmark that isn't nan, then normalize\n        # the mean is calculated per frame, the std per sequence\n        mean = hand.sum(-2, keepdim=True) / (~nans).sum(-2, keepdim=True).clamp(1)\n        #stds = ((hand-mean).pow(2).sum((-2,-3), keepdim=True) / (~nans).sum((-2,-3), keepdim=True).clamp(1e-16)).sqrt()\n        stds = ((hand-mean).pow(2).sum((-1,-2,-4), keepdim=True) / (~nans).sum((-1,-2,-4), keepdim=True).div(2).clamp(1e-16)).sqrt()\n        norm = (hand-mean) / (stds + 1e-8)\n\n        # if there are 2 hands active at the same time, divide by 2, anticipating the addition that follows\n        norm = norm / (~nans).sum(-3, keepdim=True).clamp(1)\n\n        # reflect the right hand and add it to the left\n        norm = torch.stack((norm[:,:,0,:,0] + norm[:,:,1,:,0], norm[:,:,0,:,1] - norm[:,:,1,:,1]), -1)\n        assert not torch.isnan(norm).any()\n\n        for v,sl in zip(norm,sls):\n            v[sl:] = math.nan\n        return torch.view_as_complex(norm.contiguous())\n","metadata":{"execution":{"iopub.status.busy":"2023-08-24T13:57:56.598671Z","iopub.execute_input":"2023-08-24T13:57:56.599308Z","iopub.status.idle":"2023-08-24T13:57:56.644888Z","shell.execute_reply.started":"2023-08-24T13:57:56.599272Z","shell.execute_reply":"2023-08-24T13:57:56.644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Complex Embeddings\n\nThe landmarks for the hands are mapped to **complex numbers**. Each x,y coordinate becomes (y + ix), so that the conjugation is effectively a horizontal reflection (mapping the right to the left hand). The landmarks for each hand are normalized to still allow for relative scale within the sequence. This could facilitate detection of movement in the z axis (towards the camera) when a sign is accented to indicate character separation. \n\nI visualize complex embeddings by _\"looking down the x axis\"_, a real valued embedding is usually displayed with the channels on the x axis and the value in the y axis. I imagine a complex embedding to have each value rotated around the x axis. By looking down the x axis, you see each channel as a angle + distance from its center (polar)","metadata":{}},{"cell_type":"code","source":"CONFIG = Config()\nDATASET = ASLDataset(CONFIG, (0,48))\n\ndef get_sample(name=None, in_size=512, mode=\"none\", batch_only=True, preprocessed=False):\n    if name is None: name = sorted(list(CONFIG[\"phrases\"].keys()))[0]\n    name, phrase, x, target = DATASET.find(name)\n    target = target.unsqueeze(0)\n    if preprocessed: x = Preprocess(in_size, mode)(x.unsqueeze(0))[0]\n    if batch_only:return x\n    else: return name, phrase, x, target\n","metadata":{"execution":{"iopub.status.busy":"2023-08-24T13:57:57.970908Z","iopub.execute_input":"2023-08-24T13:57:57.97203Z","iopub.status.idle":"2023-08-24T13:58:03.285512Z","shell.execute_reply.started":"2023-08-24T13:57:57.971992Z","shell.execute_reply":"2023-08-24T13:58:03.284123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib.collections import LineCollection\n\ndef plot_complex_embedding_h(handle, z):\n    linear = torch.arange(z.shape[-1]).float()\n    p0 = torch.complex(linear, torch.zeros_like(z.real))\n    p1 = torch.complex(linear, z.real)\n    segm = torch.view_as_real(torch.stack((p0,p1),1).view(-1,2))\n    lc = LineCollection(segm, colors=inline.n_colors(z.shape[-1]))\n    handle.set_xlim(-1,z.shape[-1]+1)\n    handle.set_ylim(-2,2)\n    handle.add_collection(lc)\n    handle.plot((0,z.shape[-1]-1),(0,0), c=\"lightgray\")\n    handle.scatter(linear, z.real, c=inline.n_colors(z.shape[-1]))\n      \ndef plot_complex_embedding_polar(handle, z):\n    p0 = torch.complex(torch.linspace(0,0,z.shape[-1]), torch.linspace(0,0,z.shape[-1]))\n    segm = torch.view_as_real(torch.stack((p0,z),1).view(-1,2)).flip(-1)\n    lc = LineCollection(segm, colors=inline.n_colors(z.shape[-1]))\n    handle.set_xlim(-2,2)\n    handle.set_ylim(-2,2)\n    handle.add_collection(lc)\n    handle.scatter(z.imag, z.real, c=inline.n_colors(z.shape[-1]))\n\ndef plot_complex_embedding(z, title):\n    fig,(ax,bx,cx) = plt.subplots(1,3, figsize=(15,4))\n    fig.suptitle(title)\n    ax.set_title(\"Z\")\n    bx.set_title(\"Z.real\")\n    cx.set_title(\"Z.imag\")\n    \n    plot_complex_embedding_polar(ax, z)\n    plot_complex_embedding_h(bx, z.real)\n    plot_complex_embedding_h(cx, z.imag)\n    plt.show()\n\nplot_complex_embedding(get_sample(preprocessed=True)[0], \"COMPLEX EMBEDDING\")","metadata":{"execution":{"iopub.status.busy":"2023-08-24T13:58:03.287747Z","iopub.execute_input":"2023-08-24T13:58:03.288174Z","iopub.status.idle":"2023-08-24T13:58:04.496765Z","shell.execute_reply.started":"2023-08-24T13:58:03.288135Z","shell.execute_reply":"2023-08-24T13:58:04.495427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The Signogram\n\n_(I'm coining the term :) unless there is already something out there that does something similar)_\n\nMuch like a spectogram shows the amplitude of different frequencies over time, we can get something similar by taking the complex valued landmarks, put them on the y axis and plot them over time using a complex to YUV color mapping. This clearly shows there is movement between the frames, which I would compare to articulation. It gave rise to the idea to use a 2d convolutional network on the (2d) signogram.","metadata":{}},{"cell_type":"code","source":"\nname, phrase, sample, _ = get_sample(preprocessed=True, mode=\"interpolate\", batch_only=False)\nsample = inline.interpolate(sample[None,None].mT, (128, 512), mode=\"bilinear\")\ninline.plot(sample / 1.4, title=f\"Signogram for: '{phrase}'\")","metadata":{"execution":{"iopub.status.busy":"2023-08-24T13:58:04.498186Z","iopub.execute_input":"2023-08-24T13:58:04.498561Z","iopub.status.idle":"2023-08-24T13:58:04.935628Z","shell.execute_reply.started":"2023-08-24T13:58:04.49853Z","shell.execute_reply":"2023-08-24T13:58:04.934703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Automatic Segmentation\n\nI want to automatically separate the chars in the strings based on the auto correlation of adjacent frames. I'll try to explain what I'm doing:\n\nLet's say we have a sliding window of a dynamic size. When we slide over a sequence, whenever the window contains only a single character (spanning the window), the mean of the value for all frames in the window is close to the value of a single frame. If we normalize the sequence, we can say that the sum of all frames will be 0, and we assume that the embeddings are relatively statistically independent. The more different embeddings inside the window, the lower the total absolute value of the embedding will be, untill it reaches zero when all embeddings are in the window that spans the entire sequence.","metadata":{}},{"cell_type":"code","source":"# Start with a sequence, normalize it\nt = get_sample(preprocessed=True, in_size=512, mode=\"interpolate\").mT\nt = 0.5 * (t - t.mean(-1,keepdim=True)) / t.std(-1,keepdim=True).add(1e-8)\ninline.plot(t[None,None], cols=1, title=\"normalized original\")\n\n# get the result for any window size by subtracting  shifted versions of the cumulative sum\nall_sums = (t.cumsum(-1)[None])\nfor i in [int(v**2) for v in range(2,10,2)]:\n    d = all_sums[...,i:] - all_sums[...,:-i]\n    inline.plot(d / i, title=f\"window-size {i}\")","metadata":{"execution":{"iopub.status.busy":"2023-08-24T13:58:04.937427Z","iopub.execute_input":"2023-08-24T13:58:04.938595Z","iopub.status.idle":"2023-08-24T13:58:05.804689Z","shell.execute_reply.started":"2023-08-24T13:58:04.938551Z","shell.execute_reply":"2023-08-24T13:58:05.803882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can create a matrix of size NxN where each row `a` in (0...N-1) denotes the position of the left side of the window, and each column `b` in (0...N-1) denotes the postion of the right side. Such a matrix can easily be created using the cumulative sum `CS` of the sequence, where the sum for a window `W(a,b)` is calculated as `W(a,b) = CS(b) - CS(a)`. Using broadcasting we subtract a row vector of `CS` from a column vector of `CS` which will expand into a matrix `M`, where `M(a,b) = W(a,b)`.","metadata":{}},{"cell_type":"code","source":"d = ((all_sums[...,:,None] - all_sums[...,None,:]).contiguous()) \nd = d.transpose(0,1)\n\n# divide by the window size to get the mean value for the window\nx = torch.linspace(0,d.shape[-1],d.shape[-1])\ny = torch.linspace(0,d.shape[-1],d.shape[-1])\nv = ((x[None,:] - y[:,None]).abs()) + 1e-8\nM = d  / v\ninline.plot(M.sum(0), title=\"Integrals of every possible window over a complex valued sequence\")\n\nPRETTY = M","metadata":{"execution":{"iopub.status.busy":"2023-08-24T14:17:59.910805Z","iopub.execute_input":"2023-08-24T14:17:59.911218Z","iopub.status.idle":"2023-08-24T14:18:00.654078Z","shell.execute_reply.started":"2023-08-24T14:17:59.911186Z","shell.execute_reply":"2023-08-24T14:18:00.652539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can make this matrix symmetric along the diagonal axis, multiplying with a diagonally shifted conjugate for instance, this was just an improvisation and acts as a sort of complex cosine similarity measure. The important thing is that it gives us areas of interest (high brightness) along the diagonal that we would like to separate / segment. Where the embedding stays the same, multiplying along the diagonal with the conjugate will give you a high positive real number, whereas you will get imaginary results on the boundaries.","metadata":{}},{"cell_type":"code","source":"Ma = M[:,:,:-1,:-1]\nMb = M[:,:,1:,1:].conj()\nM = (Ma * Mb)\nM = M.transpose(0,1)\ninline.plot(torch.cat((M.real, M.imag)),cols=7,width=20, title=\"real / imag part of all possible windows (A,B)\")","metadata":{"execution":{"iopub.status.busy":"2023-08-24T14:00:06.680888Z","iopub.execute_input":"2023-08-24T14:00:06.681298Z","iopub.status.idle":"2023-08-24T14:00:07.76884Z","shell.execute_reply.started":"2023-08-24T14:00:06.681265Z","shell.execute_reply":"2023-08-24T14:00:07.767484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Before we continue, let's define a square by using top-left, bottom-right notation, `(t,l,b,r)`. Notice that a square with its center on the diagonal of the matrix can be described as `(a,a,b,b)`. The matrix described above has the window activations along the diagonal, and `(a,a,b,b)` is the total density of the activation in window `W(a,b)`. We divide this total density by the window size `(b-a)` in order to get a relative density for a window.\n\nIf we normalize the matrix, there is an optimal way you can put squares along the diagonal (each next to another without overlap) and get the highest relative density out of all sets of possible combinations of squares. This optimal solution segments the bright spots the best and each square is likely to envelop a single bright spot (character). We can calculate the solution by first creating a summed-area table `SAT` (cumulative sum in 2d). With such a table you can calculate the area for rectangle `(t,l,b,r)` by: `area = SAT(b,r) + SAT(t,l) - SAT(b,l)) - SAT(t,r)`. If we constrain the problem and only want to calculate squares with their origin on the diagonal and take into account the symmetry of the matrix we can simplify because `SAT(b,l) == SAT(t,r)` and `t == l, b == r`. Last thing we do is take the diagonal `Diag` of the summed area table `SAT`, so that `Diag(a) == SAT(a,a)` and therefore the square `(a,a,b,b)`, which represents window `W(a,b)` becomes:\n```\nDiag(a) + Diag(b) - 2 * CS(a,b)\n```\nDoing the same thing as before  we add a column vector containing `Diag` to a row vector containing `Diag`. This is then added to the matrix `CS * -2` giving us matrix `Q`, where `Q(a,b)` contains the integral of the values in the square representing window `W(a,b)`","metadata":{}},{"cell_type":"code","source":"CS = M.cumsum(-1).cumsum(-2)\nDiag = CS.diagonal(0, dim1=-1, dim2=-2)\nQ = -2 * CS + Diag[...,None,:] + Diag[...,:,None]\nQ = Q.sum(1)","metadata":{"execution":{"iopub.status.busy":"2023-08-24T03:03:40.785016Z","iopub.execute_input":"2023-08-24T03:03:40.788514Z","iopub.status.idle":"2023-08-24T03:03:41.062765Z","shell.execute_reply.started":"2023-08-24T03:03:40.788473Z","shell.execute_reply":"2023-08-24T03:03:41.061669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"In order to get the optimal solution I use a recursive algorithm that tries to split a window `(a,b)` into 2 smaller windows `(a,x),(x,b)` that maximize the relative density. It then does the same for the 2 smaller windows ,recursively, ideally until there are no more splits that can be made that give a better solution, but in practice it will most likely need to be unrolled and have a fixed maximum depth of recursion. We also add a hyper parameter `gamma` which defines if the algorithm should be more relaxed or more strict (0.9 ... 1.1)","metadata":{}},{"cell_type":"code","source":"def split(d, a, b, gamma):\n    if a == b: return a\n    v1 = torch.nan_to_num(d[...,a:b,a].real / torch.arange(0,b-a,1).pow(gamma))\n    v2 = torch.nan_to_num(d[...,a:b,b].real / torch.arange(b-a,0,-1).pow(gamma))\n    return a+(v1+v2).argmax()\n    \ndef split_all(d,a,b, gamma=1):\n    c = split(d,a,b, gamma)\n    if c > a and c < b:\n        va = split_all(d,a,c,gamma)\n        vb = split_all(d,c,b,gamma)\n        return va + vb\n    else:\n        return [(a,b)]\n\nr = []\nfor g in [0.5, 0.75, 1.0]:\n    c = split_all(Q, 0, Q.shape[-1]-1, gamma=g)\n    m = M.clone()\n    for a,b in c:\n        m[...,a:b,b] = 1\n        m[...,b,a:b] = 1\n        m[...,a:b,a] = 1\n        m[...,a,a:b] = 1\n    r.append(m.abs())\n    \ninline.plot(torch.cat(r), title=\"gammas 0.5, 0.75, 1.0\")","metadata":{"execution":{"iopub.status.busy":"2023-08-24T03:03:41.0668Z","iopub.execute_input":"2023-08-24T03:03:41.067131Z","iopub.status.idle":"2023-08-24T03:03:42.322015Z","shell.execute_reply.started":"2023-08-24T03:03:41.067103Z","shell.execute_reply":"2023-08-24T03:03:42.321088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Putting it together","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport pyarrow.parquet as pq\nimport tqdm as tqdm\nimport inline\nimport reduce\nimport complex as Z\nimport algorithm as alg\n\nclass AutoSegment(nn.Module):\n    def __init__(self, gamma=1, mode=\"box\"):\n        super().__init__()\n        self.gamma = gamma\n        self.mode = mode\n\n    def split(self,d,a,b):\n        if a == b: return a\n        v1 = torch.nan_to_num(d[...,a:b+1,a].real / torch.arange(0,b-a+1,1,device=d.device).add(1e-8).pow(self.gamma))\n        v2 = torch.nan_to_num(d[...,a:b+1,b].real / torch.arange(b-a+1,0,-1,device=d.device).add(1e-8).pow(self.gamma))\n        return a+(v1+v2).argmax().item()\n    \n    def split_all(self,d,a,b):\n        c = self.split(d,a,b)\n        if c > a and c < b:\n            va = self.split_all(d,a,c)\n            vb = self.split_all(d,c,b)\n            return va + vb\n        else:\n            return [(a,b)]\n\n    def forward(self, source, target=None,  return_segments=False):\n        target = target if target is not None else source\n        rs = torch.full_like(target, math.nan)\n        ss = []\n        for i, (x, xl) in enumerate(zip(source, alg.sequence_lengths(source))):\n            if not xl: continue\n            mean = x[:xl].mean(-2, keepdim=True)\n            c = x[:xl] - mean\n            c = c.transpose(-1,-2)\n            c = c.cumsum(-1)\n            m = c[:,:,None] - c[:,None,:] # E,L,L\n            lx = torch.linspace(0,xl,xl, device=x.device)\n            ly = torch.linspace(0,xl,xl, device=x.device)\n            v = (lx[None,:] - ly[:,None]).abs()\n            m = m / v.add(1e-8)\n            m = m[...,:-1,:-1] * m[...,1:,1:].conj()\n            m = F.pad(m, (1,0,1,0), \"constant\", 0)\n            m = m.cumsum(-1).cumsum(-2)\n            d = m.diagonal(0, dim1=-1, dim2=-2)\n            m = d[:,None,:] + d[:,:,None] - 2 * m\n            m = m.sum(0).abs()\n            segments = self.split_all(m, 0, m.shape[-1]-1)\n            ss.append(segments)\n            for j,(a,b) in enumerate(segments):\n                if a >= b: continue\n                if self.mode == \"box\": rs[i,j] = (target[i,a:b].mean(0)) \n                if self.mode ==\"gaussian\": rs[i,j] = (target[i,a:b] * inline.gaussian1d(b-a).to(target.device)[:,None]).sum(0)\n        if return_segments:\n            return rs, ss\n        else:\n            return rs\n        \nAS = AutoSegment(mode=\"gaussian\")\ns = get_sample(preprocessed=True, mode=\"interpolate\", in_size=512)[None]\ns = inline.interpolate(s, (64,), mode=\"linear\")\ninline.plot(s.mT/1.4, title=\"Unsegmented\")\n\nr,segments = AS(s, return_segments=True)\ns[:] = 0\nfor v,segment in zip(r[0],segments[0]):\n    s[0,1+segment[0]:segment[1]] = v\ninline.plot(s.mT/1.4, cols=1, title=\"Auto-segmented\")","metadata":{"execution":{"iopub.status.busy":"2023-08-24T14:08:46.517084Z","iopub.execute_input":"2023-08-24T14:08:46.51791Z","iopub.status.idle":"2023-08-24T14:08:48.203863Z","shell.execute_reply.started":"2023-08-24T14:08:46.517866Z","shell.execute_reply":"2023-08-24T14:08:48.202642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The Model","metadata":{}},{"cell_type":"code","source":"def make_interpolation_matrix(in_range, out_size):\n    r = torch.zeros(in_range, out_size, in_range)\n    for i in range(1,in_range):\n        identity = torch.eye(i)\n        r[i,:,:i] = F.interpolate(identity[None,None], (out_size, i), mode=\"bilinear\")[0,0]\n    return r\nIM = make_interpolation_matrix(1024,300)\ninline.plot(IM, \"Interpolation Matrix\")\n\nfor i in [120,240,480]:\n    t = get_sample(in_size=120, mode=\"interpolate\")\n    t = IM[t.shape[0],:,:t.shape[0]] @ t\n    inline.plot(t.unsqueeze(0).mT, title=t.shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-24T03:03:42.324054Z","iopub.execute_input":"2023-08-24T03:03:42.324954Z","iopub.status.idle":"2023-08-24T03:03:51.088177Z","shell.execute_reply.started":"2023-08-24T03:03:42.324916Z","shell.execute_reply":"2023-08-24T03:03:51.087132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Translator(nn.Module):\n    def __init__(self, config): #hidden_dim, expand_dim, sequence_dim, num_heads, num_prototypes, num_project_layers, num_encoder_layers):\n        super().__init__()\n        self.hidden_dim = config[\"hidden_dim\"]\n        self.expand_dim = config[\"expand_dim\"]\n        \n        self.preprocess = Preprocess(mode=\"interpolate\", in_size=160)\n        self.segmenter = AutoSegment(gamma=1)\n        \n        self.project = nn.Sequential(\n            inline.Unsqueeze(1),\n            Z.Conv2d(1, self.expand_dim, 5, padding=2),\n            Z.ReLU(),\n            Z.Conv2d(self.expand_dim, self.expand_dim, 5, padding=2),\n            Z.ReLU(),\n            Z.Conv2d(self.expand_dim, self.expand_dim, 5, padding=2),\n            Z.ReLU(),\n            Z.Conv2d(self.expand_dim, self.expand_dim, 5, padding=2),\n            \n            *[nn.Sequential(\n                inline.Skip(\n                    Z.Conv2d(self.expand_dim, self.expand_dim//2, 3, padding=1),\n                    Z.InstanceNorm2d(self.expand_dim//2),\n                    Z.ReLU(),\n                    Z.Conv2d(self.expand_dim//2, self.expand_dim, 3, padding=1)),\n                ) for _ in range(config[\"num_project_layers\"])],\n            inline.Mean(-1),\n            inline.Transpose(1,2),\n            inline.Contiguous(),\n            Z.ReLU(),\n            Z.Linear(self.expand_dim, self.hidden_dim))\n        \n        self.encoder = nn.ModuleList(\n            [Z.TransformerEncoderLayer(self.hidden_dim, config[\"num_heads\"]) \n              for _ in range(config[\"num_encoder_layers\"])])\n        \n        self.unproject = nn.Sequential(\n            Z.MLP(self.hidden_dim, self.hidden_dim),\n            Z.SignedAbs(),\n            nn.LayerNorm(self.hidden_dim))\n\n        # NCCL doesn't work with complex, use view_as_complex in forward instead\n        self._positional = nn.parameter.Parameter(torch.zeros(config[\"sequence_dim\"],self.hidden_dim,2))\n        self._prototypes = nn.parameter.Parameter(torch.randn(config[\"num_prototypes\"], self.hidden_dim))\n        \n    @property\n    def prototypes(self):\n        return self._prototypes.detach().clone()\n    \n    @property\n    def positional(self):\n        return self._positional.detach().clone()\n    \n    def embed(self, indices, requires_grad=False):\n        padding = (indices == -1)\n        indices[padding] = 0\n        proto = self._prototypes if requires_grad else self.prototypes\n        r = F.one_hot(indices, proto.shape[0]).float() @ proto #, proto.shape[-1:]\n        #r = F.layer_norm(r, r.shape[-1:])\n        r[padding] = math.nan\n        return r\n    \n    def unembed(self, embeddings, requires_grad=False, return_scores=False, nan=-1):\n        nans = torch.isnan(embeddings)\n        embeddings = torch.where(nans, torch.zeros_like(embeddings), embeddings)\n        proto = self._prototypes if requires_grad else self.prototypes\n        la = embeddings.pow(2).sum(-1).unsqueeze(-1)\n        lb = proto.pow(2).sum(-1).unsqueeze(-2)\n        v = (embeddings @ proto.mT) / (embeddings.shape[-1] ** 0.5) #(la*lb).clamp(1e-23).sqrt()\n        v, arg = v.softmax(-1).max(-1)\n        arg[nans.any(-1)] = nan\n        return (arg, v) if return_scores else arg\n    \n    def postprocess(self, embeddings, requires_grad=False, min_prob=0.001):\n        nans = torch.isnan(embeddings)\n        embeddings = torch.where(nans, torch.zeros_like(embeddings), embeddings)\n        proto = self._prototypes if requires_grad else self.prototypes\n        v = (embeddings @ proto.mT) / (embeddings.shape[-1] ** 0.5) #(la*lb).clamp(1e-23).sqrt()\n        v, arg = v.softmax(-1).max(-1)\n        arg[nans.any(-1)] = -1\n        arg[v < min_prob] = -1\n        duplicates = F.pad(arg[...,1:] == arg[...,:-1], (1,0),\"constant\",0)\n        arg[duplicates] = -1\n        for a in arg:\n            r = a[a != -1]\n            a[:len(r)] = r\n            a[len(r):] = -1\n\n        padding = arg == -1\n        arg[padding] = 0\n        r = F.one_hot(arg, proto.shape[0]).float() @ proto\n        r[padding] = math.nan\n        return r\n\n    def parameter_groups(self, prototypes_lr_factor=0.01):\n        return [\n            dict(params=[p for name, p in self.named_parameters() if not name.endswith(\"_prototypes\") and \"bias\" not in name and \"bn\" not in name]),\n            dict(params=[p for name, p in self.named_parameters() if \"bias\" in name or \"bn\" in name], wd_factor=0),\n            dict(params=[self._prototypes], lr_factor=prototypes_lr_factor, wd_factor=0)\n        ]\n\n    def forward(self, x):\n        x = self.preprocess(x)\n        preprocessed = x.clone()\n        sls = alg.sequence_lengths(x)\n        nans = torch.isnan(x).any(-1)\n        x[nans] = 0\n        x = self.project(x)\n        \n        x_padding_mask = torch.zeros(x.shape[0], x.shape[1], device=x.device, dtype=torch.bool)\n        if len(self.encoder):\n            for i,sl in enumerate(sls):\n                x_padding_mask[i,sl:] = True\n                if sl <= 0: continue\n                if sl > len(self.positional): sl = len(self.positional)\n                x[i,:sl] = x[i,:sl] + torch.view_as_complex(self.positional[:sl] + self.positional[-sl:])\n            for encoder in self.encoder:\n                x = encoder(x, x_padding_mask)\n        \n        x = self.unproject(x)\n        \n        # Just improvising here, keeping correlation low and distribution\n        # close to normal. kl-ish\n        proto_loss = self._prototypes.mean() ** 2 \\\n                   - self._prototypes.var().clamp(1e-5).log() \\\n                   + self._prototypes.var() - 1\n        corr = (self._prototypes @ self._prototypes.mT) / self._prototypes.shape[-1]\n        corr_loss = (corr.pow(2).sum() - corr.pow(2).trace()) / corr.numel()\n        \n        x[nans] = math.nan\n        self.segmenter.mode = \"box\" if self.training else  \"gaussian\"\n        if not self.training:\n            x = self.segmenter(x)\n            x = self.postprocess(x)\n            \n        return x, 10 * proto_loss[None] + corr_loss[None]\n\n    \nclass SingleTranslator(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.module = Translator(config).to(config[\"device\"])\n    def embed(self, *args, **kwargs):\n        return self.module.embed(*args, **kwargs)\n    def unembed(self, *args, **kwargs):\n        return self.module.unembed(*args, **kwargs)\n    def parameter_groups(self, prototypes_lr_factor=0.01):\n        return [\n            dict(params=[p for name, p in self.named_parameters() if not name.endswith(\"_prototypes\") and \"bias\" not in name and \"bn\" not in name]),\n            dict(params=[p for name, p in self.named_parameters() if \"bias\" in name or \"bn\" in name], wd_factor=0),\n            dict(params=[self.module._prototypes], lr_factor=prototypes_lr_factor)\n        ]\n    def forward(self, x):\n            return self.module(x)\n    @property\n    def prototypes(self): \n        return self.module.prototypes\n\n    @property\n    def positional(self): \n        return self.positional\n    \nclass DataParallelTranslator(nn.DataParallel):\n    def __init__(self, config):\n        super().__init__(Translator(config).to(config[\"device\"]))\n    def embed(self, *args, **kwargs):\n        return self.module.embed(*args, **kwargs)\n    def unembed(self, *args, **kwargs):\n        return self.module.unembed(*args, **kwargs)\n    def parameter_groups(self, prototypes_lr_factor=0.01):\n        return [\n            dict(params=[p for name, p in self.named_parameters() if not name.endswith(\"_prototypes\") and \"bias\" not in name and \"bn\" not in name]),\n            dict(params=[p for name, p in self.named_parameters() if \"bias\" in name or \"bn\" in name], wd_factor=0),\n            dict(params=[self.module._prototypes], lr_factor=prototypes_lr_factor)\n        ]\n\n    @property\n    def prototypes(self): \n        return self.module.prototypes\n\n    @property\n    def positional(self): \n        return self.positional\n\nclass DistributedDataParallelTranslator(nn.parallel.DistributedDataParallel):\n    def __init__(self, config, **kwargs):\n        super().__init__(Translator(config).to(config[\"device\"]), **kwargs)\n    def embed(self, *args, **kwargs):\n        return self.module.embed(*args, **kwargs)\n    def unembed(self, *args, **kwargs):\n        return self.module.unembed(*args, **kwargs)\n    def parameter_groups(self, prototypes_lr_factor=0.01):\n        return [\n            dict(params=[p for name, p in self.named_parameters() if not name.endswith(\"_prototypes\") and \"bias\" not in name and \"bn\" not in name]),\n            dict(params=[p for name, p in self.named_parameters() if \"bias\" in name or \"bn\" in name], wd_factor=0),\n            dict(params=[self.module._prototypes], lr_factor=prototypes_lr_factor)\n        ]\n\n    @property\n    def prototypes(self): \n        return self.module.prototypes\n    @property\n    def positional(self): \n        return self.positional","metadata":{"execution":{"iopub.status.busy":"2023-08-24T03:05:10.746333Z","iopub.execute_input":"2023-08-24T03:05:10.746693Z","iopub.status.idle":"2023-08-24T03:05:10.796686Z","shell.execute_reply.started":"2023-08-24T03:05:10.746664Z","shell.execute_reply":"2023-08-24T03:05:10.795575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Code\n\nPretraining is done without the transformer encoder layers and without the auto segmentation. The model trains fully to match the output distribution to the target distribution. The finetuning done here adds the transformer layers and autosegmentation in order to form phrases. The finetuning code is very similar to this code, and is added in the dataset.","metadata":{}},{"cell_type":"code","source":"# File main.py\n# Created by E/S Pronk\n\nimport os\nimport time\nimport glob\nimport json\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.multiprocessing as mp\nimport torch.distributed as dist\nfrom torch.utils.data import DataLoader\n\nimport numpy as np\nimport pandas as pd\nimport pyarrow.parquet as pq\nimport matplotlib.pyplot as plt\nimport tqdm.notebook as tqdm\nimport complex as Z\nimport reduce\nimport inline\nimport algorithm as alg\nimport scheduler as sched\nfrom collections.abc import Iterable\n\n# Uncomment these if saved outside notebook\n#from config import Config\n#from dataset import ASLDataset, collate_pad\n#from model import Translator, DataParallelTranslator, DistributedDataParallelTranslator\n\ndef main(config, model):\n    \n    if not dist.is_initialized() or config[\"localrank\"] == 0:\n        # clean up old cache files when config[\"columns\"] changes for instance\n        if config[\"cleanup\"]:\n            pattern = config[\"cleanup\"] if isinstance(config[\"cleanup\"], Iterable) else { config[\"cleanuo\"] }\n            for pat in pattern:\n                for path in glob.glob(os.path.join(config[\"output_path\"], pat)):\n                    os.remove(path)\n\n    if dist.is_initialized():\n        dist.barrier()\n\n    # data\n    if dist.is_initialized():\n        # this needs to be fixed\n        n1= max(1, config[\"num_train_shards\"] // (config[\"dist_worldsize\"]))\n        n2= max(1, config[\"num_test_shards\"] // (config[\"dist_worldsize\"]))\n        dataset_train = ASLDataset(config, (config[\"rank\"]*n1, (config[\"rank\"]+1)*n1))\n        dataset_test = ASLDataset(config, (config[\"num_train_shards\"] + config[\"rank\"]*n2, config[\"num_train_shards\"] + (config[\"rank\"]+1)*n2))\n    else:\n        dataset_train = ASLDataset(config, (0,config[\"num_train_shards\"]))\n        dataset_test = ASLDataset(config, (config[\"num_train_shards\"], config[\"num_train_shards\"] + config[\"num_test_shards\"]))\n        \n    dataset_train.create_cache()\n    dataset_test.create_cache()\n    if dist.is_initialized():\n        dist.barrier()\n\n    dataloader_train = DataLoader(dataset_train, num_workers=config[\"num_workers\"], batch_size=config[\"batch_size\"], collate_fn=collate_pad, pin_memory=config[\"device\"] != \"cpu\", persistent_workers=True)\n    dataloader_test = DataLoader(dataset_test, num_workers=config[\"num_workers\"], batch_size=config[\"batch_size\"], collate_fn=collate_pad, pin_memory=config[\"device\"] != \"cpu\", persistent_workers=True)\n\n    # training \n    nit = len(dataloader_train) * config[\"epochs\"]\n    opt = torch.optim.AdamW(model.parameter_groups())\n    lfn = nn.CosineEmbeddingLoss(margin=0.2, reduction=\"none\")\n\n    sch_lr = sched.LR(opt, sched.Cosine, start=1e-6, value=1e-3, final=1e-6, iterations=nit, warmup=1000, name=\"lr\")\n    sch_wd = sched.WD(opt, sched.Cosine, start=0, value=1e-6, final=0, iterations=nit, warmup=1000, name=\"wd\")\n    schedulers = [sch_lr, sch_wd]\n\n    # try to load the latest checkpoint\n    for start_epoch in range(config[\"epochs\"],0,-1):\n        path = os.path.join(config[\"output_path\"], config[\"checkpoint_path\"].format(epoch=start_epoch, **config))\n        if os.path.exists(path):\n            checkpoint = torch.load(path, map_location=config[\"device\"])\n            model.load_state_dict(checkpoint[\"model\"])\n            opt.load_state_dict(checkpoint[\"optim\"])\n            break\n    else:\n        start_epoch = 0\n    for sch in schedulers: sch.step(start_epoch * len(dataloader_train))\n\n\n    for e in range(start_epoch, config[\"epochs\"]):\n        for training, dl, log_interval in [\n                (True, dataloader_train, 1000), \n                (False, dataloader_test, 1000)]:\n            torch.set_grad_enabled(training)\n            model.train(training)\n    \n            window_size = {\"loss\":100, \"alt\":100}\n            window = {\"loss\":[],\"alt\":[]}\n            \n            for i, (names, phrases, sequences, target_ind) in enumerate(dl):\n                sequences = sequences.to(config[\"device\"], non_blocking=True)\n                target_ind = target_ind.to(config[\"device\"], non_blocking=True)\n                pred, proto_loss = model(sequences)\n                target = model.embed(target_ind, requires_grad=True)\n                \n                # notice I apply loss to the means of the prediction and targets\n                norm_pred = F.layer_norm(torch.nan_to_num(pred).mean(-2), pred.shape[-1:])\n                norm_target = F.layer_norm(torch.nan_to_num(target).mean(-2), target.shape[-1:])\n                loss = lfn(norm_pred, norm_target, torch.ones(norm_pred.shape[0], device=norm_pred.device)).mean() \n                \n                # We just took a step forward towards the target, now we take a step back from the wrong predictions\n                # get the actual embeddings for the predicted values, and use that as a negative target\n                pred_ind = model.unembed(pred)\n                pred_emb = model.embed(pred_ind)\n                norm_emb = F.layer_norm(torch.nan_to_num(pred_emb).mean(-2), pred_emb.shape[-1:])\n                error = (norm_emb - norm_target)\n\n                # With the error being in the form emb - target = sum(emb[i] - target[i]) where emb[i] != target[i]\n                # we can rewrite as sum(emb[i]) - sum(target[i]) = sum(emb[i]) + factor * target, factor being negative\n                # Because they are N(0,1), factor = sum(error * target) / sum(target^2)\n                # So, error - factor * target = sum(emb[i]) where emb[i] != target. \n                factor = (norm_target * error).sum(-1,keepdim=True) / norm_target.pow(2).sum(-1,keepdim=True)\n                error = error - factor * norm_target\n                norm_error = F.layer_norm(error, error.shape[-1:])\n                alt_loss = lfn(norm_pred, norm_error,-torch.ones(norm_pred.shape[0], device=norm_pred.device)).mean()\n                total_loss = loss + alt_loss + proto_loss.mean()\n                \n                \n                if i % log_interval == 0 and (not dist.is_initialized() or config[\"rank\"] == 0):\n                    print()    \n                    print()\n                    for phrase, p, t in zip(phrases[:4], pred, target):\n                        print(phrase)  #total_loss.item(), sch_state\n                        t = config.decode(model.unembed(t))\n                        p = config.decode(model.unembed(p))\n                        #p = \"\".join([f\"\\033[31m{pc}\" if tc != pc else f\"\\033[32m{pc}\" for tc,pc in zip(t,p)]) + \"\\033[0m\"\n                        print(f\"\\033[36mtarget:\\033[0m \",t)\n                        print(f\"\\033[36mprediction:\\033[0m\",p)\n                        print()\n\n                if training:\n                    opt.zero_grad()\n                    total_loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), 3)\n                    opt.step()\n                    sch_state = { sch.name: sch.step() for sch in schedulers }\n                \n                if dist.is_initialized():\n                    dist.all_reduce(loss, dist.ReduceOp.AVG)\n                    dist.all_reduce(alt_loss, dist.ReduceOp.AVG)\n                    dist.all_reduce(proto_loss, dist.ReduceOp.AVG)\n                    #dist.all_reduce(erode_loss, dist.ReduceOp.AVG)\n                    \n                window[\"loss\"].append(loss.item())\n                window[\"alt\"].append(alt_loss.item())\n                \n                if not dist.is_initialized() or config[\"rank\"] == 0:\n                    stage = \"\\033[33mTRAIN\" if training else \"\\033[32mTEST\"\n                    print(f\"{stage} \", end=\"\")\n                    print(f\"ep:{e+1:03d} it:{i+1:05d} \\033[0m[\", end=\"\")\n                    print(\" \".join([f\"{k}:{sum(w)/max(len(w),1):0.06f}\" for k,w in window.items()]), end=\"\")\n                    print(\"] \\033[66m[\", end=\"\")\n                    print(\" \".join([f\"{s.name}:{s.current:0.06f}\" for s in schedulers]), end=\"]\\033[0m    \\r\")\n\n                for k,w in window.items():\n                    while len(w) > window_size[k]: \n                        w.pop(0)\n                \n\n        if not dist.is_initialized() or config[\"localrank\"] == 0:\n            checkpoint_path = os.path.join(config[\"output_path\"], config[\"checkpoint_path\"].format(epoch=e+1,**config))\n            print()\n            print(f\"Saving checkpoint as {checkpoint_path}\")\n            print()\n            checkpoint = {\n                    \"model\": model.state_dict(),\n                    \"optim\": opt.state_dict(),\n                    \"epoch\": e+1 }\n            torch.save(checkpoint, checkpoint_path)\n    \n    if dist.is_initialized():\n        dist.barrier()\n\ndef run_distributed(rank):\n\n    config = Config()\n    localrank = rank % config[\"gpu_count\"]\n    torch.cuda.set_device(\"cuda:\" + str(localrank))\n    config[\"rank\"] = rank\n    config[\"localrank\"] = localrank\n    print(f\"[{rank}] Waiting for peers...\")\n    dist.init_process_group(\n            backend=\"nccl\",\n            init_method=\"tcp://\" + config[\"dist_hostname\"] + \":\" + str(config[\"dist_port\"]),\n            rank=rank,\n            world_size= config[\"dist_worldsize\"])\n    if localrank == 0:\n        print(\"starting...\")\n    assert dist.is_initialized()\n    translator = DistributedDataParallelTranslator(config, find_unused_parameters=True)\n    main(config, translator)\n\nif __name__ == \"__main__\":\n    config = Config()\n    if config[\"device\"] == \"cpu\":\n        # CPU path\n        print(\"Running on the CPU\")\n        translator = SingleTranslator(config)\n        main(config, translator)\n\n    elif config[\"gpu_count\"] > 1:\n        if config[\"dist_enabled\"]:\n            # DistributedDataParallel path\n            print(\"Running with DistributedDataParallel\")\n            mp.spawn(run_distributed, nprocs=min(config[\"dist_worldsize\"], config[\"gpu_count\"]))\n        else:\n            # DataParallel path\n            print(\"Running with DataParallel\")\n            translator = DataParallelTranslator(config)\n            main(config, translator)\n    else:\n        # single GPU CUDA path \n        print(\"Running with single GPU\")\n        translator = SingleTranslator(config)\n        main(config, translator)","metadata":{"execution":{"iopub.status.busy":"2023-08-24T03:05:12.702992Z","iopub.execute_input":"2023-08-24T03:05:12.703363Z","iopub.status.idle":"2023-08-24T09:00:53.081462Z","shell.execute_reply.started":"2023-08-24T03:05:12.703333Z","shell.execute_reply":"2023-08-24T09:00:53.080314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inline.plot(PRETTY.sum(0))","metadata":{"execution":{"iopub.status.busy":"2023-08-24T14:18:09.395181Z","iopub.execute_input":"2023-08-24T14:18:09.396011Z","iopub.status.idle":"2023-08-24T14:18:10.013453Z","shell.execute_reply.started":"2023-08-24T14:18:09.395951Z","shell.execute_reply":"2023-08-24T14:18:10.01213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}}]}