{"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":"**Introduction:**\n\nIn ref[1], Joint Variational Autoencoder(VAE) learns disentangled continuous and discrete representations in an unsupervised manner. In image data, distinct objects or entities would most naturally be represented by discrete variables, while their position or scale might be represented by continuous variables.  Joint VAE disentagles MNIST digit type (discrete) from slant, width and stroke thickness (continuous).\n\nSample from a categorical distribution is non-differentiable, therefore inability to backpropagate through sample, hence, unable to capture the discrete latent variables.  In ref[2], present an efficient gradient estimator that replaces the non-differentiable sample from categorical distribution with a differentiable sample from a Gumbel-Softmax distribution.\n\nFrom the research papers that I read. The manifold of dimensionality reduction hypothesis suggests that high-dimensional image datas reside on a lower-dimensional manifold. The low-level features are without semantic meaning(smoothness) common to all natural images. \nSupervised learning encourages semantic, high-level features.\n\nIn this notebook,\n\n- MNIST, 1D and 2D (Section 1-6)\n\n- investigate Supervised vs. Semi-Supervised learning latents distributions. Curious about Supervised learning contributing to distributions. \n\n- in Gumbel loss function, target reconstruction loss (r_tgt) is given by:\n\n   r_tgt  = self.bce(y_true_tgt, y_pred_tgt) *MASK \n   \n   by setting MASK = 1 implies Supervised learning, 0 implies Semi-Supervised learning.\n   \n- Comparing latents distributions (Section 5)\n\n- Simultaneously Combo Supervised/Semi-Supervised learning (Section 6)\n\n- UltraMNIST, 1D and 2D demo codes (Sections 7-9)\n","metadata":{}},{"cell_type":"markdown","source":"1. MNIST 1D Supervised\n\n2. MNIST 1D Semi-Supervised\n\n3. MNIST 2D Supervised\n\n4. MNIST 2D Semi-Supervised\n\n5. MNIST Comparing Latents 1D/2D, Supervised/Semi-Supervised\n\n6. MNIST Combo Supervised and Semi-Supervised\n\n7. UltraMNIST\n\n8. UltraMNIST 1D Supervised\n\n9. UltraMNIST 2D Supervised\n\n10. Summary\n\n11. References","metadata":{}},{"cell_type":"code","source":"##https://github.com/ericjang/gumbel-softmax/blob/master/gumbel_softmax_vae_v2.ipynb\n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.utils import to_categorical\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n\nimport numpy as np \nimport pandas as pd \nimport matplotlib.pyplot as plt\nimport cv2,os,gc\n\nfrom IPython.display import Image\n\nimport tensorflow_probability as tfp\ntfd = tfp.distributions\ntfpl = tfp.layers\ntfk = tf.keras\ntfkl = tf.keras.layers\nOneHotCategorical = tfd.OneHotCategorical\n\ndef reset_random_seeds(seed_num):\n    os.environ['PYTHONHASHSEED'] = str(seed_num)\n    tf.random.set_seed(seed_num)\n    np.random.seed(seed_num)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-03-28T10:15:53.751654Z","iopub.execute_input":"2022-03-28T10:15:53.752043Z","iopub.status.idle":"2022-03-28T10:15:55.950916Z","shell.execute_reply.started":"2022-03-28T10:15:53.751950Z","shell.execute_reply":"2022-03-28T10:15:55.950171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 128  \n\nepsilon_std = 0.01\nANNEAL_RATE = 0.0003\nMIN_TEMP = 0.5\nTAU0 = 5.0 # initial temperature\nlearn_temp = False \n\ntau = tf.Variable(TAU0, dtype=tf.float32, name = \"temperature\", trainable=learn_temp)\n\nM = 10  # number of classes \nN = 30  # number of categorical distributions\nH1 = 512 #int(data_dim/2)\nH2 = 256 #int(H1/2)\n\nIMG_SZ = 28\ndata_dim = 784  #len(feature_cols)\nH1_cat = H1 #int(data_dim/2)\nH2_cat = H2 #int(H1/2)\nM_cat = M  # \nN_cat = N  #","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:55.952448Z","iopub.execute_input":"2022-03-28T10:15:55.952721Z","iopub.status.idle":"2022-03-28T10:15:56.923528Z","shell.execute_reply.started":"2022-03-28T10:15:55.952687Z","shell.execute_reply":"2022-03-28T10:15:56.922410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**1. MNIST 1D Supervised**","metadata":{}},{"cell_type":"code","source":"tf.keras.backend.clear_session()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:56.925190Z","iopub.execute_input":"2022-03-28T10:15:56.925701Z","iopub.status.idle":"2022-03-28T10:15:56.939312Z","shell.execute_reply.started":"2022-03-28T10:15:56.925663Z","shell.execute_reply":"2022-03-28T10:15:56.938650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"################################################################################\n# Load data\n################################################################################\n#([60k,28,28],[60k,]), ([10k,28,28],[10k,])\n#(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data(path=\"mnist.npz\")  #y is label\n(x_train, y_train),(_,_) = tf.keras.datasets.mnist.load_data(path=\"mnist.npz\")  #y is label\n\nx_train = x_train.astype(\"float32\") / 255.0\n#x_test = x_test.astype(\"float32\") / 255.0\nx_train = x_train.reshape((len(x_train), np.prod(x_train.shape[1:])))  #60000,784\n#x_test = x_test.reshape((len(x_test), np.prod(x_test.shape[1:])))  #10000,784\n\nimg = x_train[0,].copy()  #for visualize\n# binarize mnist pixels to 1 or 0\nx_train[x_train >=0.5] = 1\nx_train[x_train < 0.5] = 0\n\n# Train/Val split\nsplit = int(0.8 * len(x_train))\nx_train, x_val = x_train[:split], x_train[split:]\ny_train, y_val = y_train[:split], y_train[split:]\n\ntarget_trn_bin = tf.cast(tf.one_hot(y_train,M), y_val.dtype)\ntarget_val_bin = tf.cast(tf.one_hot(y_val,M), y_val.dtype)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:56.941596Z","iopub.execute_input":"2022-03-28T10:15:56.942091Z","iopub.status.idle":"2022-03-28T10:15:57.482145Z","shell.execute_reply.started":"2022-03-28T10:15:56.942055Z","shell.execute_reply":"2022-03-28T10:15:57.481281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img2= img.copy()\n\nk = 0.5 \nimg3 = (img > k).reshape(IMG_SZ, IMG_SZ)\nnet = tf.cast(tf.random.uniform(tf.shape(img2)) < img2, img2.dtype) # dynamic binarization\nnet2 = tf.cast(tf.random.uniform(tf.shape(img2)) < img2, img2.dtype) # dynamic binarization\n\nplt.subplot(1,4,1)\nplt.imshow(img.reshape(IMG_SZ,IMG_SZ), cmap='Greys_r')\nplt.subplot(1,4,2)\nplt.imshow(img3, cmap='Greys_r')\nplt.subplot(1,4,3)\nplt.imshow(tf.reshape(net,(IMG_SZ,IMG_SZ)), cmap='Greys_r')\nplt.subplot(1,4,4)\nplt.imshow(tf.reshape(net2,(IMG_SZ,IMG_SZ)), cmap='Greys_r');","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:57.483508Z","iopub.execute_input":"2022-03-28T10:15:57.483778Z","iopub.status.idle":"2022-03-28T10:15:57.909298Z","shell.execute_reply.started":"2022-03-28T10:15:57.483741Z","shell.execute_reply":"2022-03-28T10:15:57.908580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sampling(keras.layers.Layer):\n    def call(self, logits_y):\n        u = tf.random.uniform(tf.shape(logits_y), 0, 1)\n        y = logits_y - tf.math.log(-tf.math.log(u + 1e-20) + 1e-20)  # logits + gumbel noise\n        y = tf.nn.softmax(tf.reshape(y, (-1, N, M)) / tau)\n        y = tf.reshape(y, (-1, N * M))\n        return y","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:57.910531Z","iopub.execute_input":"2022-03-28T10:15:57.911302Z","iopub.status.idle":"2022-03-28T10:15:57.918282Z","shell.execute_reply.started":"2022-03-28T10:15:57.911262Z","shell.execute_reply":"2022-03-28T10:15:57.917428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"H1_tgt = H1\nH2_tgt = H2\ndata_dim_tgt = 10\ndata_dim_bin = data_dim\n\n############################################################################################## \n# 1D MODELS\n############################################################################################## \n\ndef model_1d():\n        #Encoder_binary\n        encoder_bin_inputs = keras.Input(shape=(data_dim_bin))\n        x = keras.layers.Dense(H1, activation=\"relu\")(encoder_bin_inputs)  #512\n        x = keras.layers.Dense(H2, activation=\"relu\")(x)               #256\n        logits_bin_y = keras.layers.Dense(M * N, name=\"logits_bin_y\")(x)\n        z_bin = Sampling()(logits_bin_y)\n\n        encoder_bin = keras.Model(inputs=encoder_bin_inputs, \n                                  outputs=z_bin, \n                                  name  =\"encoder_bin\" )\n        encoder_bin.build(encoder_bin_inputs)\n\n        #Decoder_binary\n        decoder_bin_inputs = keras.Input(shape=(N * M))\n        x = keras.layers.Dense(H2, activation=\"relu\")(decoder_bin_inputs)  #256\n        x = keras.layers.Dense(H1, activation=\"relu\")(x)                   #512\n        decoder_bin_outputs = keras.layers.Dense(data_dim_bin, activation=\"sigmoid\")(x)\n        #Decoder_target\n        x_tgt = keras.layers.Dense(H2_tgt, activation=\"relu\")(decoder_bin_inputs)  #256\n        x_tgt = keras.layers.Dense(H1_tgt, activation=\"relu\")(x_tgt)               #512\n        x_tgt = keras.layers.Dense(128, activation=\"relu\")(x_tgt)               #128\n\n        decoder_tgt_outputs = keras.layers.Dense(data_dim_tgt, activation=\"sigmoid\")(x_tgt)\n\n        decoder_bin_tgt = keras.Model(inputs=decoder_bin_inputs,\n                                      outputs=[decoder_bin_outputs, decoder_tgt_outputs], \n                                      name  =\"decoder_bin_tgt\" )\n        decoder_bin_tgt.build(decoder_bin_inputs)\n\n        return encoder_bin, decoder_bin_tgt\n    \n","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:57.919601Z","iopub.execute_input":"2022-03-28T10:15:57.920328Z","iopub.status.idle":"2022-03-28T10:15:57.932991Z","shell.execute_reply.started":"2022-03-28T10:15:57.920280Z","shell.execute_reply":"2022-03-28T10:15:57.932249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VAE_TARGET(keras.Model):\n    def __init__(self, encoder_1, decoder_1, **kwargs):\n        super(VAE_TARGET, self).__init__(**kwargs)\n        self.encoder_bin = tf.keras.Model(encoder_1.input, \n                                          [encoder_1.get_layer(name='logits_bin_y').output, \n                                          encoder_1.output])\n\n        self.decoder_bin_tgt = decoder_1\n        self.bce = tf.keras.losses.BinaryCrossentropy()\n        self.loss_tracker = keras.metrics.Mean(name=\"loss\")         #elbo_feature\n        self.r1_tracker   = keras.metrics.Mean(name=\"r_feat\")       #reconstruction_loss_feature\n        self.k1_tracker   = keras.metrics.Mean(name=\"KL feat\")      #KL feature\n        self.r2_tracker   = keras.metrics.Mean(name=\"r_target\")     #reconstruction_loss_target\n\n        \n    @property\n    def metrics(self): \n        # `Metric` objects so that `reset_states()` automatically at the start of each epoch\n        return [self.loss_tracker,\n                self.r1_tracker,\n                self.k1_tracker,\n                self.r2_tracker,               \n                ]\n\n    def call(self, x):\n        _, z = self.encoder_bin(x)\n        x_hat,x_hat_tgt = self.decoder_bin_tgt(z)     #feature/target reconstruction\n        return x_hat,x_hat_tgt\n\n    @tf.function\n    def gumbel_loss(self, y_true_feat, y_pred_feat, logits_y, y_true_tgt, y_pred_tgt):\n        q_y = tf.reshape(logits_y, (-1, N, M))\n        q_y = tf.nn.softmax(q_y)\n        log_q_y = tf.math.log(q_y + 1e-20)\n        kl_tmp = q_y * (log_q_y - tf.math.log(1.0 / M))\n        kl = tf.math.reduce_sum(kl_tmp, axis=(1, 2))    #KL loss\n        kl = tf.math.reduce_sum(kl,-1)  #scalar per batch\n        \n        #MASK=0 ->SEMI Supervised, 1->Supervised learning\n        r_feat = self.bce(y_true_feat, y_pred_feat)     #features reconstruction loss\n        r_tgt  = self.bce(y_true_tgt, y_pred_tgt) *MASK #target reconstruction loss\n        elbo   = data_dim_bin *(r_feat + r_tgt)  - kl\n        return elbo,r_feat,kl,r_tgt\n    \n    def train_step(self, data):\n        # Unpack data. Structure depends on model and on what pass to `fit()`.  \n        print('step', np.shape(data[0]))\n        x, x_tgt = data   #(features,target)\n\n        with tf.GradientTape(persistent=True) as tape:\n          logits_y, z = self.encoder_bin(x, training=True)          #for bin KL loss\n          x_hat,x_hat_tgt = self.decoder_bin_tgt(z, training=True)  #for feat/tgt reconstruct loss\n          loss,r1,k1,r2 = self.gumbel_loss(x, x_hat, logits_y, x_tgt, x_hat_tgt)  \n          r1 = tf.reduce_mean(r1)  #reconstruction_loss features\n          k1 = tf.reduce_mean(k1)  #KL features\n          r2 = tf.reduce_mean(r2)  #reconstruction_loss target\n        \n        # Compute gradients\n        grads = tape.gradient(loss, self.trainable_weights)\n        # Update weights\n        self.optimizer.apply_gradients(zip(grads, self.trainable_weights))\n\n        # Update metrics (includes the metric that tracks the loss)\n        self.loss_tracker.update_state(loss)  #elbo\n        self.r1_tracker.update_state(r1)\n        self.k1_tracker.update_state(k1)    \n        self.r2_tracker.update_state(r2)\n        \n        # Return a dict mapping metric names to current value\n        return {\"loss\": self.loss_tracker.result(),       # elbo \n                \"r_loss_feature\": self.r1_tracker.result(),  #reconstruction_loss_feat\n                \"KL\": self.k1_tracker.result(),              #KL feat\n                \"r_loss_target\": self.r2_tracker.result(),   #reconstruction_loss_target\n                } ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:57.934269Z","iopub.execute_input":"2022-03-28T10:15:57.934727Z","iopub.status.idle":"2022-03-28T10:15:57.955465Z","shell.execute_reply.started":"2022-03-28T10:15:57.934686Z","shell.execute_reply":"2022-03-28T10:15:57.954644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomLearningRateScheduler(keras.callbacks.Callback):\n    \"\"\"Learning rate scheduler which sets the learning rate according to schedule.\n\n    Arguments:\n            schedule: a function that takes an epoch index\n            (integer, indexed from 0) and current learning rate\n            as inputs and returns a new learning rate as output (float).\n    \"\"\"\n\n    def __init__(self, schedule):\n        super(CustomLearningRateScheduler, self).__init__()\n        self.schedule = schedule\n\n    def on_epoch_begin(self, epoch, logs=None):\n        if not hasattr(self.model.optimizer, \"lr\"):\n            raise ValueError('Optimizer must have a \"lr\" attribute.')\n        # Get the current learning rate from model's optimizer.\n        lr = float(tf.keras.backend.get_value(self.model.optimizer.learning_rate))\n        # Call schedule function to get the scheduled learning rate.\n        scheduled_lr = self.schedule(epoch, lr)\n        # Set the value back to the optimizer before this epoch starts\n        tf.keras.backend.set_value(self.model.optimizer.lr, scheduled_lr)\n        #print(f\"e={epoch}, tau={tau}, lr={scheduled_lr}\")\n\nLR_SCHEDULE = [\n              # (epoch to start, learning rate) tuples\n              (3, 0.0015),\n              (6, 0.001),\n              (16, 0.0005),  #9\n              (26, 0.0002),  #12\n              ]\n\ndef lr_schedule(epoch, lr):\n    \"\"\"Helper function to retrieve the scheduled learning rate based on epoch.\"\"\"\n    if epoch < LR_SCHEDULE[0][0] or epoch > LR_SCHEDULE[-1][0]:\n        return lr\n    for i in range(len(LR_SCHEDULE)):\n        if epoch == LR_SCHEDULE[i][0]:\n            return LR_SCHEDULE[i][1]\n    return lr","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:57.956772Z","iopub.execute_input":"2022-03-28T10:15:57.957226Z","iopub.status.idle":"2022-03-28T10:15:57.968978Z","shell.execute_reply.started":"2022-03-28T10:15:57.957189Z","shell.execute_reply":"2022-03-28T10:15:57.968192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK = 1.  #1->Supervise Learning, 0->SemiSupervised\n\nencoder_bin, decoder_bin_tgt = model_1d()\nvae_target = VAE_TARGET(encoder_bin, decoder_bin_tgt, name=\"vae_target-model\")\nvae_target_inputs = (None, IMG_SZ*IMG_SZ)  \n\nvae_target.build(vae_target_inputs)\nvae_target.compile(optimizer=\"adam\", loss=None)\n\n#reset all models coeffs and biases\nencoder_bin.reset_states()\ndecoder_bin_tgt.reset_states()\nvae_target.reset_states()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:57.972353Z","iopub.execute_input":"2022-03-28T10:15:57.972946Z","iopub.status.idle":"2022-03-28T10:15:58.368839Z","shell.execute_reply.started":"2022-03-28T10:15:57.972892Z","shell.execute_reply":"2022-03-28T10:15:58.368112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_bin.summary()\ndecoder_bin_tgt.summary()\nvae_target.summary()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:58.370236Z","iopub.execute_input":"2022-03-28T10:15:58.370490Z","iopub.status.idle":"2022-03-28T10:15:58.386795Z","shell.execute_reply.started":"2022-03-28T10:15:58.370456Z","shell.execute_reply":"2022-03-28T10:15:58.386128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ep=20\n\n#%%time\nhist = vae_target.fit(\n        x_train,target_trn_bin,\n        shuffle=True,\n        epochs=ep,\n        callbacks=[CustomLearningRateScheduler(lr_schedule)],\n        ) ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:15:58.387925Z","iopub.execute_input":"2022-03-28T10:15:58.388240Z","iopub.status.idle":"2022-03-28T10:18:21.960704Z","shell.execute_reply.started":"2022-03-28T10:15:58.388205Z","shell.execute_reply":"2022-03-28T10:18:21.959797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_hist():  \n    plt.figure(figsize=(20,4))\n    plt.subplot(141)\n    plt.plot(hist.history['loss'])\n    plt.legend(['elbo'])\n    plt.subplot(142)\n    plt.plot(hist.history['r_loss_feature'])\n    plt.legend(['rec_loss_feature'])\n    plt.subplot(143)\n    plt.plot(hist.history['r_loss_target'])\n    plt.legend(['rec_loss_target'])\n    plt.subplot(144)\n    plt.plot(hist.history['KL'])\n    plt.legend(['KL_loss'])\n    \ndef rmse(y_true, y_pred):\n    return tf.sqrt(tf.reduce_mean(tf.square(y_true - y_pred)))  \n\ndef accuracy(y_true, y_pred):\n    return 100*(len(y_true) - tf.math.count_nonzero(y_true-y_pred, dtype=tf.float32))/len(y_true)\n    \ndef plot_target_predict():\n    #Target Prediction\n    a = plt.axes(aspect='equal')\n    plt.scatter(y_val, yphat_val)\n    plt.xlabel('True Values [sum]')\n    plt.ylabel('Predictions [sum]')\n    lims = [0, np.max(y_val)]\n    plt.xlim(lims)\n    plt.ylim(lims)\n    _ = plt.plot(lims, lims) \n    \n    y_true = tf.reshape(tf.convert_to_tensor(y_val, dtype=tf.float32),[-1,]) \n    y_pred = tf.cast(yphat_val, tf.float32)\n\n    idx = np.random.randint(len(y_true))\n    print('y_true = ', y_true[idx:idx+20].numpy())\n    print('y_pred = ', y_pred[idx:idx+20].numpy(),'\\n')\n\n    #tf.print('rmse = ', rmse(y_true, y_pred))\n    print(f\"RMSE: {rmse(y_true, y_pred):.5f}\\n\")\n    print(f'accuracy: {accuracy(y_true, y_pred):.5f} %\\n')    \n\ndef plot_latent():\n#Features Prediction and Latents\n\n    plt.figure(figsize=(22,20))\n    for i in range(4):\n      j=5*i+1\n      z_2D = tf.reshape(z_latent[i],(N, M))\n      proj_Z_c = tf.math.reduce_sum(z_2D, axis=0)  #proj on class\n      proj_Z_d = tf.math.reduce_sum(z_2D, axis=1)  #proj on distribution\n      plt.subplot(4,5,j)\n      plt.imshow( x_val[i].reshape(IMG_SZ, IMG_SZ), cmap='gray') , plt.axis('off')         #x_test(1,784)\n      plt.title('true')\n\n      plt.subplot(4,5,j+1)\n      plt.bar(range(0,N),proj_Z_d);\n      plt.xlabel('dist')\n      plt.ylabel('sum')\n      plt.title('latent')\n\n      plt.subplot(4,5,j+2)\n      plt.imshow(z_2D, cmap='gray'), plt.axis('off')     #z=argmax_y(1,30,10)\n      plt.xlabel('class')\n      plt.ylabel('dist')\n      plt.title('latent')\n\n      plt.subplot(4,5,j+3)\n      plt.bar(range(0,M),proj_Z_c/tf.math.reduce_sum(proj_Z_c));\n      plt.xlabel('class')\n      plt.ylabel('avg prob')\n      plt.title('latent')\n\n      plt.subplot(4,5,j+4)\n      plt.imshow(tf.reshape(xhat_val[i],(IMG_SZ, IMG_SZ)), cmap='gray'), plt.axis('off')   #data_hat(1,784)\n      plt.title('predict')\n    \ndef plot_features():\n    #Features Prediction\n    plt.figure(figsize=(12,18))\n    for i in range(4):\n        tmp1 = np.reshape(x_val[i*10:i*10+10,:],(-1,IMG_SZ,IMG_SZ))   #true\n        img1 = np.hstack([tmp1[k] for k in range(10)])\n        tmp = np.reshape(xhat_val[i*10:i*10+10,:],(-1,IMG_SZ,IMG_SZ)) #predict\n        img = np.hstack([tmp[k] for k in range(10)])\n        j=2*i+1\n        plt.subplot(8,1,j)\n        plt.imshow(img1)\n        plt.title('true')\n        plt.subplot(8,1,j+1)\n        plt.imshow(img)\n        plt.title('predict') \n        \ndef plot_histogram():\n    # histogram of distribution\n    plt.figure(figsize=(12,4))\n    plt.subplot(121)\n    plt.hist(y_train, M, density=False, facecolor='g');\n    plt.subplot(122)\n    plt.hist(y_val, M, density=False, facecolor='g');    ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:21.963579Z","iopub.execute_input":"2022-03-28T10:18:21.964025Z","iopub.status.idle":"2022-03-28T10:18:21.988977Z","shell.execute_reply.started":"2022-03-28T10:18:21.963987Z","shell.execute_reply":"2022-03-28T10:18:21.988147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_hist()\n\nz_latent = encoder_bin(x_val)\nxhat_val, yhat_val = decoder_bin_tgt(z_latent)\n\nyphat_val = tf.argmax(yhat_val,-1)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:21.990528Z","iopub.execute_input":"2022-03-28T10:18:21.990844Z","iopub.status.idle":"2022-03-28T10:18:22.550383Z","shell.execute_reply.started":"2022-03-28T10:18:21.990801Z","shell.execute_reply":"2022-03-28T10:18:22.549713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Target Prediction\nplot_target_predict()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:22.551850Z","iopub.execute_input":"2022-03-28T10:18:22.552329Z","iopub.status.idle":"2022-03-28T10:18:22.789464Z","shell.execute_reply.started":"2022-03-28T10:18:22.552292Z","shell.execute_reply":"2022-03-28T10:18:22.788775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Features Prediction and Latents\nplot_latent()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:22.790791Z","iopub.execute_input":"2022-03-28T10:18:22.791244Z","iopub.status.idle":"2022-03-28T10:18:25.295885Z","shell.execute_reply.started":"2022-03-28T10:18:22.791204Z","shell.execute_reply":"2022-03-28T10:18:25.293789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" First column: True validation feature\n \n Second column: Latent distributions sum across classes, sum of each distribution is one.\n \n Third column: 2D latent distributions\n \n Fourth column: Latent distributions average sum across distributions, average prob. for each class.  \n Re-run it on the same set of features got different results. One can use PCA to analyze the latents.\n \n Fifth column: Predict validation feature","metadata":{}},{"cell_type":"code","source":"#Features Prediction\nplot_features()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:25.297576Z","iopub.execute_input":"2022-03-28T10:18:25.297968Z","iopub.status.idle":"2022-03-28T10:18:26.391087Z","shell.execute_reply.started":"2022-03-28T10:18:25.297894Z","shell.execute_reply":"2022-03-28T10:18:26.388780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Histogram\nplot_histogram()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:26.392761Z","iopub.execute_input":"2022-03-28T10:18:26.393036Z","iopub.status.idle":"2022-03-28T10:18:26.700944Z","shell.execute_reply.started":"2022-03-28T10:18:26.393003Z","shell.execute_reply":"2022-03-28T10:18:26.700303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent_1d, xhat_val_1d, yhat_val_1d, yphat_val_1d = z_latent, xhat_val, yhat_val, yphat_val\n\ndel x_train, x_val, target_trn_bin\ndel y_train, y_val, target_val_bin \ndel z_latent, xhat_val, yhat_val, yphat_val\ndel encoder_bin, decoder_bin_tgt, vae_target\ngc.collect()\n\ntf.keras.backend.clear_session()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:26.703069Z","iopub.execute_input":"2022-03-28T10:18:26.703520Z","iopub.status.idle":"2022-03-28T10:18:26.936936Z","shell.execute_reply.started":"2022-03-28T10:18:26.703481Z","shell.execute_reply":"2022-03-28T10:18:26.936185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**2. MNIST 1D Semi-Supervised**\n\nBy setting MASK = 0.  The target error is cleared, only reconstruction error is used for update, which implies semi-supervised training.","metadata":{}},{"cell_type":"code","source":"################################################################################\n# Load data\n################################################################################\n#([60k,28,28],[60k,]), ([10k,28,28],[10k,])\n(x_train, y_train),(_,_) = tf.keras.datasets.mnist.load_data(path=\"mnist.npz\")  #y is label\n\nx_train = x_train.astype(\"float32\") / 255.0\n#x_test = x_test.astype(\"float32\") / 255.0\nx_train = x_train.reshape((len(x_train), np.prod(x_train.shape[1:])))  #60000,784\n#x_test = x_test.reshape((len(x_test), np.prod(x_test.shape[1:])))  #10000,784\n\nimg = x_train[0,].copy()  #for visualize\n# binarize mnist pixels to 1 or 0\nx_train[x_train >=0.5] = 1\nx_train[x_train < 0.5] = 0\n\n# Train/Val split\nsplit = int(0.8 * len(x_train))\nx_train, x_val = x_train[:split], x_train[split:]\ny_train, y_val = y_train[:split], y_train[split:]\n\ntarget_trn_bin = tf.cast(tf.one_hot(y_train,M), y_val.dtype)\ntarget_val_bin = tf.cast(tf.one_hot(y_val,M), y_val.dtype)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:26.940190Z","iopub.execute_input":"2022-03-28T10:18:26.940426Z","iopub.status.idle":"2022-03-28T10:18:27.466229Z","shell.execute_reply.started":"2022-03-28T10:18:26.940398Z","shell.execute_reply":"2022-03-28T10:18:27.465437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK = 0.  #1->Supervise Learning, 0->SemiSupervised\n\nencoder_bin, decoder_bin_tgt = model_1d()\nvae_target = VAE_TARGET(encoder_bin, decoder_bin_tgt, name=\"vae_target-model\")\nvae_target_inputs = (None, IMG_SZ*IMG_SZ)  \n\nvae_target.build(vae_target_inputs)\nvae_target.compile(optimizer=\"adam\", loss=None)\n\n#reset all models coeffs and biases\nencoder_bin.reset_states()\ndecoder_bin_tgt.reset_states()\nvae_target.reset_states()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:27.467596Z","iopub.execute_input":"2022-03-28T10:18:27.467854Z","iopub.status.idle":"2022-03-28T10:18:27.599688Z","shell.execute_reply.started":"2022-03-28T10:18:27.467819Z","shell.execute_reply":"2022-03-28T10:18:27.598913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nhist = vae_target.fit(\n        x_train,target_trn_bin,\n        shuffle=True,\n        epochs=ep,\n        callbacks=[CustomLearningRateScheduler(lr_schedule)],\n        ) ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:18:27.601183Z","iopub.execute_input":"2022-03-28T10:18:27.601445Z","iopub.status.idle":"2022-03-28T10:20:50.920633Z","shell.execute_reply.started":"2022-03-28T10:18:27.601411Z","shell.execute_reply":"2022-03-28T10:20:50.919863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" Since MASK =0. The target reconstruction loss is always zero.\n \n r_tgt  = self.bce(y_true_tgt, y_pred_tgt) *MASK \n \n The target prediction is random and invalid.  Hence, RMSE and accuracy are random.\n \n Reconstruction features are valid because it's semi-supervised.","metadata":{}},{"cell_type":"code","source":"plot_hist()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:50.921882Z","iopub.execute_input":"2022-03-28T10:20:50.922472Z","iopub.status.idle":"2022-03-28T10:20:51.441535Z","shell.execute_reply.started":"2022-03-28T10:20:50.922429Z","shell.execute_reply":"2022-03-28T10:20:51.440768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent = encoder_bin(x_val)\nxhat_val, yhat_val = decoder_bin_tgt(z_latent)\n\nyphat_val = tf.argmax(yhat_val,-1)\n\n#Target Prediction\nplot_target_predict()\n\n#Features Prediction and Latents\nplot_latent()\n\n#Features Prediction\nplot_features()\n\n#Histogram\nplot_histogram()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:51.442724Z","iopub.execute_input":"2022-03-28T10:20:51.443607Z","iopub.status.idle":"2022-03-28T10:20:55.248605Z","shell.execute_reply.started":"2022-03-28T10:20:51.443564Z","shell.execute_reply":"2022-03-28T10:20:55.247763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent_1d_semi, xhat_val_1d_semi, yhat_val_1d_semi, yphat_val_1d_semi = z_latent, xhat_val, yhat_val, yphat_val\n\ndel x_train, x_val, target_trn_bin\ndel y_train, y_val, target_val_bin \ndel z_latent, xhat_val, yhat_val, yphat_val\ndel encoder_bin, decoder_bin_tgt, vae_target\ngc.collect()\n\ntf.keras.backend.clear_session()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:55.249871Z","iopub.execute_input":"2022-03-28T10:20:55.250155Z","iopub.status.idle":"2022-03-28T10:20:55.505548Z","shell.execute_reply.started":"2022-03-28T10:20:55.250116Z","shell.execute_reply":"2022-03-28T10:20:55.504749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**3. MNIST 2D Supervised**","metadata":{}},{"cell_type":"code","source":"############################################################################################## \n# 2D MODELS\n############################################################################################## \n\ndef model_2d():\n        \n        fo=32\n        \n        #Encoder_binary\n        encoder_bin_inputs = tf.keras.layers.Input(shape=(IMG_SZ,IMG_SZ,1))\n\n        x = tf.keras.layers.Conv2D(fo, (3,3), padding=\"same\", strides=2)(encoder_bin_inputs)\n        x = tf.keras.layers.ReLU()(x)    \n        x = tf.keras.layers.Conv2D(fo, (3,3), padding=\"same\", strides=2)(x)\n        x = tf.keras.layers.ReLU()(x)\n        x = tf.keras.layers.Flatten()(x)\n\n        logits_bin_y = keras.layers.Dense(M * N, name=\"logits_bin_y\")(x)\n        z_bin = Sampling()(logits_bin_y)\n\n        encoder_bin = keras.Model(inputs=encoder_bin_inputs, \n                                  outputs=z_bin, \n                                  name  =\"encoder_bin\" )\n        encoder_bin.build(encoder_bin_inputs)\n\n        #Decoder_feature\n        decoder_bin_inputs = keras.Input(shape=(N * M))\n        x     = tf.keras.layers.Dense(7*7*fo)(decoder_bin_inputs)\n        x     = tf.keras.layers.Reshape(target_shape = (7,7,fo))(x)\n        x     = tf.keras.layers.Conv2DTranspose(filters=fo, kernel_size=(3,3),padding=\"same\", strides=2)(x)\n        x     = tf.keras.layers.ReLU()(x)  #LeakyReLU\n        x     = tf.keras.layers.Conv2DTranspose(filters=1, kernel_size=(3,3), padding=\"same\", strides=2)(x)  \n        decoder_bin_outputs = tf.keras.layers.ReLU()(x)\n\n        #Decoder_target\n        x     = tf.keras.layers.Dense(7*7*fo)(decoder_bin_inputs)\n        x     = tf.keras.layers.Reshape(target_shape = (7,7,fo))(x)\n        x     = tf.keras.layers.Conv2DTranspose(filters=1, kernel_size=(3,3), activation='leaky_relu', padding=\"same\", strides=1)(x)\n        x     = tf.keras.layers.Flatten()(x)\n        decoder_tgt_outputs = keras.layers.Dense(data_dim_tgt, activation=\"sigmoid\")(x)\n\n        #model\n        decoder_bin_tgt = keras.Model(inputs=decoder_bin_inputs,\n                                      outputs=[decoder_bin_outputs, decoder_tgt_outputs], \n                                      name  =\"decoder_bin_tgt\" )\n        decoder_bin_tgt.build(decoder_bin_inputs)\n\n        return encoder_bin,decoder_bin_tgt","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:55.506936Z","iopub.execute_input":"2022-03-28T10:20:55.507345Z","iopub.status.idle":"2022-03-28T10:20:55.522777Z","shell.execute_reply.started":"2022-03-28T10:20:55.507304Z","shell.execute_reply":"2022-03-28T10:20:55.522088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"################################################################################\n# Load data\n################################################################################\n#([60k,28,28],[60k,]), ([10k,28,28],[10k,])\n(x_train, y_train),(_,_) = tf.keras.datasets.mnist.load_data(path=\"mnist.npz\")  #y is label\n\nx_train = x_train.astype(\"float32\") / 255.0\n#x_test = x_test.astype(\"float32\") / 255.0\n\ndata_dim = (IMG_SZ,IMG_SZ,1)  #len(feature_cols)\n\nx_train = x_train.reshape((-1, IMG_SZ, IMG_SZ, 1))  #integer 0, 255\n#x_test = x_test.reshape((-1, IMG_SZ, IMG_SZ, 1))  #integer 0, 255\n\n# binarize mnist pixels to 1 or 0\nx_train[x_train >=0.5] = 1\nx_train[x_train < 0.5] = 0\n\n# Train/Val split\nsplit = int(0.8 * len(x_train))\nx_train, x_val = x_train[:split], x_train[split:]\ny_train, y_val = y_train[:split], y_train[split:]\n\ntarget_trn_bin = tf.cast(tf.one_hot(y_train,M), y_val.dtype)\ntarget_val_bin = tf.cast(tf.one_hot(y_val,M), y_val.dtype)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:55.524064Z","iopub.execute_input":"2022-03-28T10:20:55.524391Z","iopub.status.idle":"2022-03-28T10:20:56.079964Z","shell.execute_reply.started":"2022-03-28T10:20:55.524351Z","shell.execute_reply":"2022-03-28T10:20:56.079209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK = 1.  #1->Supervise Learning, 0->SemiSupervised\n\nencoder_bin, decoder_bin_tgt = model_2d()\nvae_target = VAE_TARGET(encoder_bin, decoder_bin_tgt, name=\"vae_target-model\")  \nvae_target_inputs = (None,IMG_SZ,IMG_SZ,1)  #data_dim_bin\n\nvae_target.build(vae_target_inputs)\nvae_target.compile(optimizer=\"adam\", loss=None)\n\nencoder_bin.reset_states()\ndecoder_bin_tgt.reset_states()\nvae_target.reset_states()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:56.084881Z","iopub.execute_input":"2022-03-28T10:20:56.085098Z","iopub.status.idle":"2022-03-28T10:20:56.277025Z","shell.execute_reply.started":"2022-03-28T10:20:56.085073Z","shell.execute_reply":"2022-03-28T10:20:56.276253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_bin.summary()\ndecoder_bin_tgt.summary()\nvae_target.summary()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:56.278430Z","iopub.execute_input":"2022-03-28T10:20:56.278654Z","iopub.status.idle":"2022-03-28T10:20:56.297239Z","shell.execute_reply.started":"2022-03-28T10:20:56.278621Z","shell.execute_reply":"2022-03-28T10:20:56.296562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nhist = vae_target.fit(\n        x_train,target_trn_bin,\n        shuffle=True,\n        epochs=ep,\n        callbacks=[CustomLearningRateScheduler(lr_schedule)],\n        ) ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:20:56.298484Z","iopub.execute_input":"2022-03-28T10:20:56.298709Z","iopub.status.idle":"2022-03-28T10:23:12.922220Z","shell.execute_reply.started":"2022-03-28T10:20:56.298677Z","shell.execute_reply":"2022-03-28T10:23:12.921454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_hist()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:23:12.923874Z","iopub.execute_input":"2022-03-28T10:23:12.924200Z","iopub.status.idle":"2022-03-28T10:23:13.569950Z","shell.execute_reply.started":"2022-03-28T10:23:12.924164Z","shell.execute_reply":"2022-03-28T10:23:13.568918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent = encoder_bin(x_val)\nxhat_val, yhat_val = decoder_bin_tgt(z_latent)\n\nyphat_val = tf.argmax(yhat_val,-1)\n\n#Target Prediction\nplot_target_predict()\n\n#Features Prediction and Latents\nplot_latent()\n\n#Features Prediction\nplot_features()\n\n#Histogram\nplot_histogram()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:23:13.571293Z","iopub.execute_input":"2022-03-28T10:23:13.571693Z","iopub.status.idle":"2022-03-28T10:23:18.003588Z","shell.execute_reply.started":"2022-03-28T10:23:13.571652Z","shell.execute_reply":"2022-03-28T10:23:18.002780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent_2d, xhat_val_2d, yhat_val_2d, yphat_val_2d = z_latent, xhat_val, yhat_val, yphat_val\ndel x_train, x_val, target_trn_bin\ndel y_train, y_val, target_val_bin \ndel z_latent, xhat_val, yhat_val, yphat_val\ndel encoder_bin, decoder_bin_tgt, vae_target  #models\ngc.collect()\n\ntf.keras.backend.clear_session()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:23:18.005002Z","iopub.execute_input":"2022-03-28T10:23:18.005843Z","iopub.status.idle":"2022-03-28T10:23:18.286510Z","shell.execute_reply.started":"2022-03-28T10:23:18.005799Z","shell.execute_reply":"2022-03-28T10:23:18.285407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**4. MNIST 2D Semi-Supervised**","metadata":{}},{"cell_type":"code","source":"################################################################################\n# Load data\n################################################################################\n#([60k,28,28],[60k,]), ([10k,28,28],[10k,])\n(x_train, y_train),(_,_) = tf.keras.datasets.mnist.load_data(path=\"mnist.npz\")  #y is label\n\nx_train = x_train.astype(\"float32\") / 255.0\n#x_test = x_test.astype(\"float32\") / 255.0\n\ndata_dim = (IMG_SZ,IMG_SZ,1)  #len(feature_cols)\n\nx_train = x_train.reshape((-1, IMG_SZ, IMG_SZ, 1))  #integer 0, 255\n#x_test = x_test.reshape((-1, IMG_SZ, IMG_SZ, 1))  #integer 0, 255\n\n# binarize mnist pixels to 1 or 0\nx_train[x_train >=0.5] = 1\nx_train[x_train < 0.5] = 0\n\n# Train/Val split\nsplit = int(0.8 * len(x_train))\nx_train, x_val = x_train[:split], x_train[split:]\ny_train, y_val = y_train[:split], y_train[split:]\n\ntarget_trn_bin = tf.cast(tf.one_hot(y_train,M), y_val.dtype)\ntarget_val_bin = tf.cast(tf.one_hot(y_val,M), y_val.dtype)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:23:18.287804Z","iopub.execute_input":"2022-03-28T10:23:18.288172Z","iopub.status.idle":"2022-03-28T10:23:18.856184Z","shell.execute_reply.started":"2022-03-28T10:23:18.288134Z","shell.execute_reply":"2022-03-28T10:23:18.855370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK = 0.  #1->Supervise Learning, 0->SemiSupervised\n\nencoder_bin, decoder_bin_tgt = model_2d()\nvae_target = VAE_TARGET(encoder_bin, decoder_bin_tgt, name=\"vae_target-model\")  \nvae_target_inputs = (None,IMG_SZ,IMG_SZ,1)  #data_dim_bin\n\nvae_target.build(vae_target_inputs)\nvae_target.compile(optimizer=\"adam\", loss=None)\n\nencoder_bin.reset_states()\ndecoder_bin_tgt.reset_states()\nvae_target.reset_states()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:23:18.857554Z","iopub.execute_input":"2022-03-28T10:23:18.859342Z","iopub.status.idle":"2022-03-28T10:23:19.066344Z","shell.execute_reply.started":"2022-03-28T10:23:18.859299Z","shell.execute_reply":"2022-03-28T10:23:19.065642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nhist = vae_target.fit(\n        x_train,target_trn_bin,\n        shuffle=True,\n        epochs=ep,\n        callbacks=[CustomLearningRateScheduler(lr_schedule)],\n        ) ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:23:19.067627Z","iopub.execute_input":"2022-03-28T10:23:19.067882Z","iopub.status.idle":"2022-03-28T10:25:42.580322Z","shell.execute_reply.started":"2022-03-28T10:23:19.067848Z","shell.execute_reply":"2022-03-28T10:25:42.579575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_hist()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:42.581890Z","iopub.execute_input":"2022-03-28T10:25:42.582174Z","iopub.status.idle":"2022-03-28T10:25:43.085403Z","shell.execute_reply.started":"2022-03-28T10:25:42.582136Z","shell.execute_reply":"2022-03-28T10:25:43.084735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent = encoder_bin(x_val)\nxhat_val, yhat_val = decoder_bin_tgt(z_latent)\n\nyphat_val = tf.argmax(yhat_val,-1)\n\n#Target Prediction\nplot_target_predict()\n\n#Features Prediction and Latents\nplot_latent()\n\n#Features Prediction\nplot_features()\n\n#Histogram\nplot_histogram()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:43.086642Z","iopub.execute_input":"2022-03-28T10:25:43.088406Z","iopub.status.idle":"2022-03-28T10:25:46.972616Z","shell.execute_reply.started":"2022-03-28T10:25:43.088367Z","shell.execute_reply":"2022-03-28T10:25:46.971929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**5. MNIST Comparing Latents 1D/2D, Supervised/Semi-Supervised**","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(22,20))\nfor i in range(4):\n        j=5*i+1\n      \n        #1D Supervised latents  \n        z1d = tf.reshape(z_latent_1d[i],(N, M))\n        proj_Z1d_c = tf.math.reduce_sum(z1d, axis=0)  #proj onto class\n        #1D Semi-Supervised latents   \n        z1d_semi = tf.reshape(z_latent_1d_semi[i],(N, M))    \n        proj_Z1d_semi_c = tf.math.reduce_sum(z1d_semi, axis=0)  #proj onto class\n      \n        #2D-Supervised latents \n        z2d = tf.reshape(z_latent_2d[i],(N, M))    \n        proj_Z2d_c = tf.math.reduce_sum(z2d, axis=0)  #proj onto class        \n        #2D Semi-Supervised latents \n        z = tf.reshape(z_latent[i],(N, M))    \n        proj_Z_c = tf.math.reduce_sum(z, axis=0)  #proj onto class    \n    \n        #Features  \n        plt.subplot(4,5,j)\n        plt.imshow( x_val[i].reshape(IMG_SZ, IMG_SZ), cmap='gray') , plt.axis('off')         #x_test(1,784)\n        plt.title('true')\n      \n        #Latent 1D Supervised\n        plt.subplot(4,5,j+1)\n        plt.bar(range(0,M),proj_Z1d_c/tf.math.reduce_sum(proj_Z1d_c));\n        plt.xlabel('class')\n        plt.ylabel('avg prob')\n        plt.title('latent 1D Super')\n        #Latent 1D Semi-Supervised\n        plt.subplot(4,5,j+2)\n        plt.bar(range(0,M),proj_Z1d_semi_c/tf.math.reduce_sum(proj_Z1d_semi_c));\n        plt.xlabel('class')\n        plt.ylabel('avg prob')\n        plt.title('latent 1D Semi')\n        #Latent 1D Supervised\n        plt.subplot(4,5,j+3)\n        plt.bar(range(0,M),proj_Z2d_c/tf.math.reduce_sum(proj_Z2d_c));\n        plt.xlabel('class')\n        plt.ylabel('avg prob')\n        plt.title('latent 2D Super')\n        #Latent 1D Semi-Supervised\n        plt.subplot(4,5,j+4)\n        plt.bar(range(0,M),proj_Z_c/tf.math.reduce_sum(proj_Z_c));\n        plt.xlabel('class')\n        plt.ylabel('avg prob')\n        plt.title('latent 2D Semi')","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:46.973975Z","iopub.execute_input":"2022-03-28T10:25:46.974374Z","iopub.status.idle":"2022-03-28T10:25:49.310508Z","shell.execute_reply.started":"2022-03-28T10:25:46.974336Z","shell.execute_reply":"2022-03-28T10:25:49.309797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"First Column: True Validation features\n\nSecond Column: 1D Supervised, very little variations across digits, class 3 is dominant avg prob.\n\nThird Column: 1D Semi-Supervised, more equal dominant classes ( class 2 and 7)\n\nFourth Column: 2D Supervised, can't really tell the difference.\n\nFifth Column: 2D Semi-Supervised, can't really tell the difference.\n\nComparing 1D Second and Third Columns, Supervised learning has less dominant distribution.\n\nComparing 2D Fourth and Fifth Columns, can't really tell the difference.\n\nThe algorithm uses all the distributions.  Can try to use threshold criteria  e.g. threshold = sum(top average prob) >= 0.5 are the dominant classes for a particular digit. \n\nPCA will choose the top principal components. in Ref[3]","metadata":{}},{"cell_type":"code","source":"del x_train, x_val, target_trn_bin\ndel y_train, y_val, target_val_bin \ndel z_latent, xhat_val, yhat_val, yphat_val\n\ndel z_latent_1d, xhat_val_1d, yhat_val_1d, yphat_val_1d\ndel z_latent_1d_semi, xhat_val_1d_semi, yhat_val_1d_semi, yphat_val_1d_semi\ndel z_latent_2d, xhat_val_2d, yhat_val_2d, yphat_val_2d\n\ndel encoder_bin, decoder_bin_tgt, vae_target\ngc.collect()\n\ntf.keras.backend.clear_session()\nreset_random_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:49.312626Z","iopub.execute_input":"2022-03-28T10:25:49.313216Z","iopub.status.idle":"2022-03-28T10:25:49.650673Z","shell.execute_reply.started":"2022-03-28T10:25:49.313174Z","shell.execute_reply":"2022-03-28T10:25:49.649843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**6. MNIST Combo Supervised and Semi-Supervised**\n\nCurrent code does not support combo supervised and semi-supervised training simultaneously. Pseudo code as follows:\n\nmask_train = np.ones(shape=(len(x_train),1),dtype=np.float32)  #np.bool_\n\nmask_test = np.zeros(shape=(len(x_test),1), dtype=np.float32)  \n\ny_test = np.zeros(shape=(len(x_test),1), dtype=np.float32)  #dummy target\n\nx_all = np.vstack([x_train, x_test])\n\ny_all = np.vstack([y_train, y_test])  #y_test is a dummy target\n\nmask_all = np.vstack([mask_train, mask_test])\n\ndata = np.hstack([x_all, y_all, mask_all])\n\nnp.random.shuffle(data)\n\nx,y,mask <- unpack(data)\n\n#target reconstruction loss\n\nr_tgt  = self.bce(y_true_tgt, y_pred_tgt) * mask ","metadata":{}},{"cell_type":"markdown","source":"**7. UltraMNIST**","metadata":{}},{"cell_type":"code","source":"WORK_DIR = '../input/ultra-mnist/'\n\ntrain = pd.read_csv(WORK_DIR + \"train.csv\")\ntrain.loc[:, \"image_path\"] = WORK_DIR + \"train/\" + train.id.values + \".jpeg\"\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:49.652244Z","iopub.execute_input":"2022-03-28T10:25:49.652526Z","iopub.status.idle":"2022-03-28T10:25:49.693797Z","shell.execute_reply.started":"2022-03-28T10:25:49.652488Z","shell.execute_reply":"2022-03-28T10:25:49.692915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(WORK_DIR + \"sample_submission.csv\")\nsubmission.loc[:, \"image_path\"] = WORK_DIR + \"train/\" + submission.id.values + \".jpeg\"\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:49.695491Z","iopub.execute_input":"2022-03-28T10:25:49.695780Z","iopub.status.idle":"2022-03-28T10:25:49.729231Z","shell.execute_reply.started":"2022-03-28T10:25:49.695741Z","shell.execute_reply":"2022-03-28T10:25:49.728265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = True\n\nif DEBUG:\n    #train_df = train[train[\"digit_sum\"] < 12]  #debug\n    train_df = train[:400].copy()\n    test_df  = submission[:20].copy()\nelse:\n    train_df = train.copy()\n    test_df  = submission.copy()\n    \ndel train\ngc.collect()    ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:49.730947Z","iopub.execute_input":"2022-03-28T10:25:49.731264Z","iopub.status.idle":"2022-03-28T10:25:49.926945Z","shell.execute_reply.started":"2022-03-28T10:25:49.731226Z","shell.execute_reply":"2022-03-28T10:25:49.926125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(2,4,figsize=(20,10))\n\nfid = np.random.choice(train_df.index.values, replace=False, size=4)\n\nfor n in range(4):\n    path = train_df.loc[fid[n]].image_path\n    target = train_df.loc[fid[n]].digit_sum\n    image = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    plt.subplot(2,4,n+1)\n    plt.imshow(image)\n    plt.title(\"Target: {}\".format(target))\n    #binarize\n    img = image.astype(\"float32\") / 255.0 \n    img = tf.cast(tf.random.uniform(tf.shape(img)) < img, img.dtype) # dynamic binarization\n    plt.subplot(2,4,n+5)\n    plt.imshow(img) \n    plt.title(\"Target: {}\".format(target))","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:25:49.928634Z","iopub.execute_input":"2022-03-28T10:25:49.928974Z","iopub.status.idle":"2022-03-28T10:26:04.742385Z","shell.execute_reply.started":"2022-03-28T10:25:49.928933Z","shell.execute_reply":"2022-03-28T10:26:04.737095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_SZ = 512\nN=28\nM=28\n\nIMG_SHAPE = (IMG_SZ,IMG_SZ)\nINTERPOLATION = cv2.INTER_AREA\n\ndef load_img(df):\n    img_lst=[]\n    #tgt_lst=[]\n    mask_lst=[]\n    \n    path     = df.image_path\n    target   = df.digit_sum\n\n    for idx in range(len(df)):\n        load_img = cv2.imread(path[idx], cv2.IMREAD_GRAYSCALE)\n        resized_img = cv2.resize(load_img, IMG_SHAPE , interpolation =INTERPOLATION)\n        img_lst.append(resized_img)\n        #tgt_lst.append(tgt_lst)\n        mask_lst.append(1.)\n        \n    return np.asarray(img_lst), np.asarray(target), np.asarray(mask_lst)\n\ndef preprocess_img(df):\n    img, target, mask = load_img(df)\n    img = img.astype(\"float32\") / 255.0   \n    # binarize mnist pixels to 1 or 0\n    img[img >=0.5] = 1\n    img[img < 0.5] = 0\n    #img = tf.cast(tf.random.uniform(tf.shape(img)) < img, img.dtype) # dynamic binarization\n    #return img[0].reshape(-1,IMG_SZ*IMG_SZ), target, mask\n    return img, target, mask\n\n#generator = ImageDataGenerator(preprocessing_function = preprocess_img)  #iterator needs rank 4 = (batch,x,y,z)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:04.743832Z","iopub.execute_input":"2022-03-28T10:26:04.744115Z","iopub.status.idle":"2022-03-28T10:26:04.754698Z","shell.execute_reply.started":"2022-03-28T10:26:04.744077Z","shell.execute_reply":"2022-03-28T10:26:04.753952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**8. UltraMNIST 1D**","metadata":{}},{"cell_type":"code","source":"x_train,y_train,mask_train = preprocess_img(train_df)\n#x_test, y_test, mask_test = preprocess_img(test_df)\n#x_test, y_test, mask_test = x_train[:10,],y_train[:10,],mask_train[:10,]","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:04.756173Z","iopub.execute_input":"2022-03-28T10:26:04.756439Z","iopub.status.idle":"2022-03-28T10:26:32.184183Z","shell.execute_reply.started":"2022-03-28T10:26:04.756399Z","shell.execute_reply":"2022-03-28T10:26:32.183401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_train = x_train.reshape((len(x_train), np.prod(x_train.shape[1:])))  #batch,IMG_SZ*IMG_SZ\n#x_test = x_test.reshape((len(x_test), np.prod(x_test.shape[1:])))      #batch,IMG_SZ*IMG_SZ\n\n# Train/Val split\nsplit = int(0.8 * len(x_train))\nx_train, x_val = x_train[:split], x_train[split:]\ny_train, y_val = y_train[:split], y_train[split:]\n\ntarget_trn_bin = tf.cast(tf.one_hot(y_train,M), y_val.dtype)\ntarget_val_bin = tf.cast(tf.one_hot(y_val,M), y_val.dtype)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:32.185697Z","iopub.execute_input":"2022-03-28T10:26:32.185967Z","iopub.status.idle":"2022-03-28T10:26:32.195521Z","shell.execute_reply.started":"2022-03-28T10:26:32.185931Z","shell.execute_reply":"2022-03-28T10:26:32.194583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"H1=128\nH2=64\nH1_tgt=H1\nH2_tgt=H2","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:32.197073Z","iopub.execute_input":"2022-03-28T10:26:32.197511Z","iopub.status.idle":"2022-03-28T10:26:32.203186Z","shell.execute_reply.started":"2022-03-28T10:26:32.197473Z","shell.execute_reply":"2022-03-28T10:26:32.202334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_1d_ultra():\n        #Encoder_binary\n        encoder_bin_inputs = keras.Input(shape=(data_dim_bin))\n        x = keras.layers.Dense(H1, activation=\"relu\")(encoder_bin_inputs)  #512\n        x = keras.layers.Dense(H2, activation=\"relu\")(x)                   #256\n        logits_bin_y = keras.layers.Dense(M * N, name=\"logits_bin_y\")(x) \n        z_bin = Sampling()(logits_bin_y)\n\n        encoder_bin = keras.Model(inputs=encoder_bin_inputs, \n                                  outputs=z_bin, \n                                  name  =\"encoder_bin\" )\n        encoder_bin.build(encoder_bin_inputs)\n\n        #Decoder_binary_features\n        decoder_bin_inputs = keras.Input(shape=(N * M))\n        x = keras.layers.Dense(H2, activation=\"relu\")(decoder_bin_inputs)  #256\n        x = keras.layers.Dense(H1, activation=\"relu\")(x)                   #512\n        decoder_bin_outputs = keras.layers.Dense(data_dim_bin, activation=\"sigmoid\")(x)\n        #Decoder_target\n        x_tgt = keras.layers.Dense(H2_tgt, activation=\"relu\")(decoder_bin_inputs)  #256\n        x_tgt = keras.layers.Dense(H1_tgt, activation=\"relu\")(x_tgt)               #512\n        x_tgt = keras.layers.Dense(128, activation=\"relu\")(x_tgt)                  #128\n\n        decoder_tgt_outputs = keras.layers.Dense(data_dim_tgt, activation=\"sigmoid\")(x_tgt)\n\n        decoder_bin_tgt = keras.Model(inputs=decoder_bin_inputs,\n                                      outputs=[decoder_bin_outputs, decoder_tgt_outputs], \n                                      name  =\"decoder_bin_tgt\" )\n        decoder_bin_tgt.build(decoder_bin_inputs)\n\n        return encoder_bin, decoder_bin_tgt\n    \n############################################################################################## \n# 2D MODELS\n############################################################################################## \n\ndef model_2d_ultra():\n        \n        fo=8\n        f2=128\n        \n        #Encoder_binary\n        encoder_bin_inputs = tf.keras.layers.Input(shape=(IMG_SZ,IMG_SZ,1))\n\n        x = tf.keras.layers.Conv2D(fo, (3,3), padding=\"same\", strides=2)(encoder_bin_inputs)\n        x = tf.keras.layers.ReLU()(x)    \n        x = tf.keras.layers.Conv2D(fo, (3,3), padding=\"same\", strides=2)(x)\n        x = tf.keras.layers.ReLU()(x)\n        x = tf.keras.layers.Flatten()(x)\n\n        logits_bin_y = keras.layers.Dense(M * N, name=\"logits_bin_y\")(x)\n        z_bin = Sampling()(logits_bin_y)\n\n        encoder_bin = keras.Model(inputs=encoder_bin_inputs, \n                                  outputs=z_bin, \n                                  name  =\"encoder_bin\" )\n        encoder_bin.build(encoder_bin_inputs)\n\n        #Decoder_feature\n        decoder_bin_inputs = keras.Input(shape=(N * M))\n        x     = tf.keras.layers.Dense(f2*f2*fo)(decoder_bin_inputs)\n        x     = tf.keras.layers.Reshape(target_shape = (f2,f2,fo))(x)\n        x     = tf.keras.layers.Conv2DTranspose(filters=fo, kernel_size=(3,3),padding=\"same\", strides=2)(x)\n        x     = tf.keras.layers.ReLU()(x)  #LeakyReLU\n        x     = tf.keras.layers.Conv2DTranspose(filters=1, kernel_size=(3,3), padding=\"same\", strides=2)(x)  \n        decoder_bin_outputs = tf.keras.layers.ReLU()(x)\n\n        #Decoder_target\n        x     = tf.keras.layers.Dense(f2*f2*fo)(decoder_bin_inputs)\n        x     = tf.keras.layers.Reshape(target_shape = (f2,f2,fo))(x)\n        x     = tf.keras.layers.Conv2DTranspose(filters=1, kernel_size=(3,3), activation='leaky_relu', padding=\"same\", strides=1)(x)\n        x     = tf.keras.layers.Flatten()(x)\n        decoder_tgt_outputs = keras.layers.Dense(data_dim_tgt, activation=\"sigmoid\")(x)\n\n        #model\n        decoder_bin_tgt = keras.Model(inputs=decoder_bin_inputs,\n                                      outputs=[decoder_bin_outputs, decoder_tgt_outputs], \n                                      name  =\"decoder_bin_tgt\" )\n        decoder_bin_tgt.build(decoder_bin_inputs)\n\n        return encoder_bin,decoder_bin_tgt    ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:32.205221Z","iopub.execute_input":"2022-03-28T10:26:32.205558Z","iopub.status.idle":"2022-03-28T10:26:32.226795Z","shell.execute_reply.started":"2022-03-28T10:26:32.205469Z","shell.execute_reply":"2022-03-28T10:26:32.226053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK = 1.  #1->Supervise Learning, 0->SemiSupervised\n\ndata_dim = IMG_SZ*IMG_SZ  \ndata_dim_bin = data_dim\ndata_dim_tgt = M     #28\n\nencoder_bin, decoder_bin_tgt = model_1d_ultra()\nvae_target = VAE_TARGET(encoder_bin, decoder_bin_tgt, name=\"vae_target-model\")\nvae_target_inputs = (None, IMG_SZ*IMG_SZ)  \n\nvae_target.build(vae_target_inputs)\nvae_target.compile(optimizer=\"adam\", loss=None)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:32.229463Z","iopub.execute_input":"2022-03-28T10:26:32.229874Z","iopub.status.idle":"2022-03-28T10:26:32.370106Z","shell.execute_reply.started":"2022-03-28T10:26:32.229845Z","shell.execute_reply":"2022-03-28T10:26:32.369308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_bin.summary()\ndecoder_bin_tgt.summary()\nvae_target.summary()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:32.371331Z","iopub.execute_input":"2022-03-28T10:26:32.371594Z","iopub.status.idle":"2022-03-28T10:26:32.391039Z","shell.execute_reply.started":"2022-03-28T10:26:32.371558Z","shell.execute_reply":"2022-03-28T10:26:32.390239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nhist = vae_target.fit(\n        x_train,target_trn_bin,\n        shuffle=True,\n        epochs=5, #  ep,\n        callbacks=[CustomLearningRateScheduler(lr_schedule)],\n        ) ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:32.393891Z","iopub.execute_input":"2022-03-28T10:26:32.394365Z","iopub.status.idle":"2022-03-28T10:26:36.002858Z","shell.execute_reply.started":"2022-03-28T10:26:32.394335Z","shell.execute_reply":"2022-03-28T10:26:36.002183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_hist()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:36.004097Z","iopub.execute_input":"2022-03-28T10:26:36.004347Z","iopub.status.idle":"2022-03-28T10:26:36.573341Z","shell.execute_reply.started":"2022-03-28T10:26:36.004312Z","shell.execute_reply":"2022-03-28T10:26:36.572651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent = encoder_bin(x_val)\nxhat_val, yhat_val = decoder_bin_tgt(z_latent)\n\nyphat_val = tf.argmax(yhat_val,-1)\n\n#Target Prediction\nplot_target_predict()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:36.574545Z","iopub.execute_input":"2022-03-28T10:26:36.574866Z","iopub.status.idle":"2022-03-28T10:26:36.855352Z","shell.execute_reply.started":"2022-03-28T10:26:36.574828Z","shell.execute_reply":"2022-03-28T10:26:36.854658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Features Prediction and Latents\nplot_latent()\n\n#Features Prediction\nplot_features()\n\n#Histogram\nplot_histogram()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:36.856722Z","iopub.execute_input":"2022-03-28T10:26:36.856998Z","iopub.status.idle":"2022-03-28T10:26:43.526065Z","shell.execute_reply.started":"2022-03-28T10:26:36.856962Z","shell.execute_reply":"2022-03-28T10:26:43.525222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#del x_train, x_val, target_trn_bin\n#del y_train, y_val, target_val_bin \ndel z_latent, xhat_val, yhat_val, yphat_val\ndel encoder_bin, decoder_bin_tgt, vae_target\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:43.527522Z","iopub.execute_input":"2022-03-28T10:26:43.527786Z","iopub.status.idle":"2022-03-28T10:26:43.782589Z","shell.execute_reply.started":"2022-03-28T10:26:43.527750Z","shell.execute_reply":"2022-03-28T10:26:43.781782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**9. UltraMNIST 2D**","metadata":{}},{"cell_type":"code","source":"data_dim = (IMG_SZ,IMG_SZ,1)\ndata_dim_bin = IMG_SZ*IMG_SZ  #weight\ndata_dim_tgt = M     #28\n\nx_train = x_train.reshape((-1, IMG_SZ, IMG_SZ, 1))\nx_val = x_val.reshape((-1, IMG_SZ, IMG_SZ, 1))\n#x_test = x_test.reshape((-1, IMG_SZ, IMG_SZ, 1))  ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:43.783974Z","iopub.execute_input":"2022-03-28T10:26:43.784213Z","iopub.status.idle":"2022-03-28T10:26:43.789344Z","shell.execute_reply.started":"2022-03-28T10:26:43.784180Z","shell.execute_reply":"2022-03-28T10:26:43.788333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.shape(x_train),np.shape(y_train),np.shape(x_val),np.shape(y_val), #np.shape(x_test)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:43.790856Z","iopub.execute_input":"2022-03-28T10:26:43.791127Z","iopub.status.idle":"2022-03-28T10:26:43.801051Z","shell.execute_reply.started":"2022-03-28T10:26:43.791093Z","shell.execute_reply":"2022-03-28T10:26:43.800289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK = 1.  #1->Supervise Learning, 0->SemiSupervised\n\nencoder_bin, decoder_bin_tgt = model_2d_ultra()\nvae_target = VAE_TARGET(encoder_bin, decoder_bin_tgt, name=\"vae_target-model\")  \nvae_target_inputs = (None,IMG_SZ,IMG_SZ,1)  \n\nvae_target.build(vae_target_inputs)\nvae_target.compile(optimizer=\"adam\", loss=None)","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:43.803624Z","iopub.execute_input":"2022-03-28T10:26:43.804235Z","iopub.status.idle":"2022-03-28T10:26:43.989756Z","shell.execute_reply.started":"2022-03-28T10:26:43.804157Z","shell.execute_reply":"2022-03-28T10:26:43.988967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_bin.summary()\ndecoder_bin_tgt.summary()\nvae_target.summary()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:43.991926Z","iopub.execute_input":"2022-03-28T10:26:43.992394Z","iopub.status.idle":"2022-03-28T10:26:44.010776Z","shell.execute_reply.started":"2022-03-28T10:26:43.992357Z","shell.execute_reply":"2022-03-28T10:26:44.010064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nhist = vae_target.fit(\n        x_train,target_trn_bin,\n        shuffle=True,\n        epochs=5, #  ep,\n        callbacks=[CustomLearningRateScheduler(lr_schedule)],\n        ) ","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:44.012812Z","iopub.execute_input":"2022-03-28T10:26:44.013080Z","iopub.status.idle":"2022-03-28T10:26:51.034769Z","shell.execute_reply.started":"2022-03-28T10:26:44.013045Z","shell.execute_reply":"2022-03-28T10:26:51.033981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_hist()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:51.036359Z","iopub.execute_input":"2022-03-28T10:26:51.036880Z","iopub.status.idle":"2022-03-28T10:26:51.558472Z","shell.execute_reply.started":"2022-03-28T10:26:51.036840Z","shell.execute_reply":"2022-03-28T10:26:51.557763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_latent = encoder_bin(x_val)\nxhat_val, yhat_val = decoder_bin_tgt(z_latent)\n\nyphat_val = tf.argmax(yhat_val,-1)\n\n#Target Prediction\nplot_target_predict()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:51.559539Z","iopub.execute_input":"2022-03-28T10:26:51.560029Z","iopub.status.idle":"2022-03-28T10:26:51.968460Z","shell.execute_reply.started":"2022-03-28T10:26:51.559878Z","shell.execute_reply":"2022-03-28T10:26:51.967756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Features Prediction and Latents\nplot_latent()\n\n#Features Prediction\nplot_features()\n\n#Histogram\nplot_histogram()","metadata":{"execution":{"iopub.status.busy":"2022-03-28T10:26:51.969555Z","iopub.execute_input":"2022-03-28T10:26:51.970371Z","iopub.status.idle":"2022-03-28T10:26:58.675996Z","shell.execute_reply.started":"2022-03-28T10:26:51.970331Z","shell.execute_reply":"2022-03-28T10:26:58.675163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**10. Summary**\n\n- investigated MNIST, 1D and 2D, Supervised vs. Semi-Supervised learning latents distributions results in Section 5. \n\nSupervised 1D, Re-run it on the same set of features got different results. ( with reset models, reset seeds)\n\nComparing 1D, Supervised and Semi-Supervised, Supervised learning has less dominant distribution.\n\nComparing 2D, Supervised and Semi-Supervised, can't really tell the difference. Run to run variations.\n\nThe algorithm uses all the distributions.  Can try to use threshold criteria  e.g. threshold = sum(top average prob) >= 0.5 are the dominant classes for a particular digit. \n\n- Simultaneously Combo Supervised/Semi-Supervised learning pseudo code.  The code structure need to modify.\n\n- UltraMNIST, 1D and 2D demo codes","metadata":{}},{"cell_type":"markdown","source":"**11. References**\n\n[1] Emilien Dupont, Learning Disentangled Joint Continuous and Discrete Representations. 2018 https://arxiv.org/abs/1804.00104, \n\n[2] Eric Jang et al, Categorical Reparameterization with Gumbel-Softmax, ICLR 5 Aug 2017 https://arxiv.org/abs/1611.01144, https://github.com/ericjang/gumbel-softmax/blob/master/gumbel_softmax_vae_v2.ipynb\n\n[3] Zheng Ding et al, Guided Variational Autoencoder for Disentanglement Learning, CVPR 2020 https://arxiv.org/abs/2004.01255","metadata":{}}]}