{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":22962,"databundleVersionId":3171193,"sourceType":"competition"},{"sourceId":38601,"sourceType":"datasetVersion","datasetId":30279}],"dockerImageVersionId":30762,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install --upgrade jax\nimport flax.linen as nn\nimport jax\nimport jax.numpy as jnp\n# from flax.training import train_state\nimport matplotlib.pyplot as plt\nimport random\nimport tensorflow as tf\nimport optax\n\nimport numpy\nimport os\nimport tensorflow_datasets as tfds\nimport sys\nimport cv2\nimport numpy as np\ndef load_and_preprocess_image(filenames, img_height, img_width):\n    dataset = tf.keras.preprocessing.image_dataset_from_directory(\n        filenames,\n        image_size=(256, 256),  # Adjust image size as needed\n        batch_size=100,\n        color_mode=\"rgb\",\n    )\n\n    # Prefetch the dataset for performance\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset\n\nclass Encoder(nn.Module):\n    @nn.compact\n    def __call__(self, inputs):\n        x_16_feature_skip = nn.leaky_relu(nn.Conv(features = 16, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(inputs), negative_slope = 0.2)\n        x_16_feature_skip = nn.leaky_relu(nn.Conv(features = 16, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_16_feature_skip), negative_slope = 0.2)\n        # print(x_16_feature_skip.shape, \"pre-pool-16\")\n        x_pool_16 = nn.max_pool(x_16_feature_skip, window_shape = (2, 2), strides = (2, 2))\n\n        # \t\tprint(x_pool_16.shape, \"post-pool-16\")\n        x_32_feature_skip = nn.leaky_relu(nn.Conv(features = 32, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_pool_16), negative_slope = 0.2)\n        x_32_feature_skip = nn.leaky_relu(nn.Conv(features = 32, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_32_feature_skip), negative_slope = 0.2)\n        # \t\tprint(x_32_feature_skip.shape, \"pre-pool-32\")\n        x_pool_32 = nn.max_pool(x_32_feature_skip, window_shape = (2, 2), strides = (2, 2))\n\n        # \t\tprint(x_pool_32.shape, \"post-pool-32\")\n        x_64_feature_skip = nn.leaky_relu(nn.Conv(features = 64, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_pool_32), negative_slope = 0.2)\n        x_64_feature_skip = nn.leaky_relu(nn.Conv(features = 64, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_64_feature_skip), negative_slope = 0.2)\n        # \t\tprint(x_64_feature_skip.shape, \"pre-pool-64\")\n        x_pool_64 = nn.max_pool(x_64_feature_skip, window_shape = (2, 2), strides = (2, 2))\n\n        # \t\tprint(x_pool_64.shape, \"post-pool-64\")\n        x_128_feature_skip = nn.leaky_relu(nn.Conv(features = 128, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_pool_64), negative_slope = 0.2)\n        x_128_feature_skip = nn.leaky_relu(nn.Conv(features = 128, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_128_feature_skip), negative_slope = 0.2)\n        # \t\tprint(x_128_feature_skip.shape, \"pre-pool-128\")\n        x_pool_128 = nn.max_pool(x_128_feature_skip, window_shape = (2, 2), strides = (2, 2))\n\n        # \t\tprint(x_pool_128.shape, \"post-pool-128\")\n        x_256_feature_skip = nn.leaky_relu(nn.Conv(features = 256, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_pool_128), negative_slope = 0.2)\n        x_256_feature_skip = nn.leaky_relu(nn.Conv(features = 256, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_256_feature_skip), negative_slope = 0.2)\n        # \t\tprint(x_256_feature_skip.shape, \"pre-pool-256\")\n        x_pool_256 = nn.max_pool(x_256_feature_skip, window_shape = (2, 2), strides = (2, 2))\n\n        # \t\tprint(x_pool_256.shape, \"post-pool-256\")\n        x_pool_256_1 = nn.max_pool(x_pool_256, window_shape = (2, 2), strides = (2, 2))\n\n        mu = jnp.mean(x_pool_256_1)\n        log_var = jnp.log(jnp.var(x_pool_256_1) ** 2)\n        randint = random.randint(1, 1000)\n        parameterization = mu + jax.random.normal(jax.random.PRNGKey(randint), x_pool_256_1.shape) * jnp.exp(log_var)\n        return parameterization, mu, log_var, [inputs, x_pool_16, x_pool_32, x_pool_64, x_pool_128, x_pool_256]\n\n\nclass Decoder(nn.Module):\n    @nn.compact\n    def __call__(self, inputs: nn.Module, skip_connections: list[nn.Module], input_channels: int, output_channels: int):\n        # print(inputs.shape, \"hello\")\n        x_transpose_1 = nn.ConvTranspose(features = input_channels, kernel_size = (3, 3), strides = (2, 2), padding = 'SAME')(inputs)\n        x_added_5 = x_transpose_1 + skip_connections[5]\n        x_convoluted = nn.Conv(features = input_channels, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_added_5)\n        x_convoluted = nn.relu(x_convoluted)\n        x_convoluted = nn.Conv(features = input_channels, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_convoluted)\n        x_convoluted = nn.relu(x_convoluted)\n\n        x_transpose_2 = nn.ConvTranspose(features = input_channels // 2, kernel_size = (3, 3), strides = (2, 2), padding = 'SAME')(x_convoluted)\n        x_added_4 = x_transpose_2 + skip_connections[4]\n        x_convoluted = nn.Conv(features = input_channels // 2, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_added_4)\n        x_convoluted = nn.relu(x_convoluted)\n        x_convoluted = nn.Conv(features = input_channels // 2, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_convoluted)\n        x_convoluted = nn.relu(x_convoluted)\n\n        x_transpose_3 = nn.ConvTranspose(features = input_channels // 4, kernel_size = (3, 3), strides = (2, 2), padding = 'SAME')(x_convoluted)\n        x_added_3 = x_transpose_3 + skip_connections[3]\n        x_convoluted = nn.Conv(features = input_channels // 4, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_added_3)\n        x_convoluted = nn.relu(x_convoluted)\n        x_convoluted = nn.Conv(features = input_channels // 4, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_convoluted)\n        x_convoluted = nn.relu(x_convoluted)\n\n        x_transpose_4 = nn.ConvTranspose(features = input_channels // 8, kernel_size = (3, 3), strides = (2, 2), padding = 'SAME')(x_convoluted)\n        x_added_2 = x_transpose_4 + skip_connections[2]\n        x_convoluted = nn.Conv(features = input_channels // 8, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_added_2)\n        x_convoluted = nn.relu(x_convoluted)\n        x_convoluted = nn.Conv(features = input_channels // 8, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_convoluted)\n        x_convoluted = nn.relu(x_convoluted)\n\n        x_transpose_5 = nn.ConvTranspose(features = input_channels // 16, kernel_size = (3, 3), strides = (2, 2), padding = 'SAME')(x_convoluted)\n        x_added_1 = x_transpose_5 + skip_connections[1]\n        x_convoluted = nn.Conv(features = input_channels // 16, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_added_1)\n        x_convoluted = nn.relu(x_convoluted)\n        x_convoluted = nn.Conv(features = input_channels // 16, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_convoluted)\n        x_convoluted = nn.relu(x_convoluted)\n\n        x_transpose_6 = nn.ConvTranspose(features = output_channels, kernel_size = (3, 3), strides = (2, 2), padding = 'SAME')(x_convoluted)\n        x_added_0 = x_transpose_6 + skip_connections[0]\n        x_convoluted = nn.Conv(features = output_channels, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_added_0)\n        x_convoluted = nn.relu(x_convoluted)\n        x_convoluted = nn.Conv(features = output_channels, kernel_size = (3, 3), strides = (1, 1), padding = 'SAME')(x_convoluted)\n        x_convoluted = nn.relu(nn.sigmoid(x_convoluted))\n        # print(x_convoluted.shape, \"post-conv\")\n        return x_convoluted\n\n\nclass EncoderDecoder(nn.Module):\n    @nn.compact\n    def __call__(self, input1):\n        encoder = Encoder()\n        decoder = Decoder()\n        latent, mu, log_var, skips = encoder(input1)\n        reconstructed = decoder(latent, skips, 256, 3)\n        return nn.relu(reconstructed), mu, log_var\n\nencoder_decoder = EncoderDecoder()\nparams = encoder_decoder.init(jax.random.PRNGKey(42), jnp.ones([1, 256, 256, 3]))\nprint('now')\n\ndef loss(params, x, y):\n    x_reconstructed, mu, log_var = encoder_decoder.apply(params, x)\n    return jnp.mean((x_reconstructed - y)) + (-0.5 * jnp.sum(1 + log_var - jnp.square(mu) - jnp.exp(log_var)))\n\n\noptimizer = optax.adamw(0.001)\nopt_state = optimizer.init(params)\n\n\ndef update(params, opt_state, x, y):\n    loss_val, grads = jax.value_and_grad(loss, argnums=0)(params, x, y)\n    updates, opt_state = optimizer.update(grads, opt_state, params)\n    new_params = optax.apply_updates(params, updates)\n    return new_params, opt_state, loss_val\n\ni = 0\nmnop = (load_and_preprocess_image(\"/kaggle/input/caltech256/256_ObjectCategories\", 256, 256))\n# for m, n in mnop:\n#     x = jnp.array(m)\n# print(x.shape)\n# exit(0)\n# print(\"hello\")\n# y = x\n# print(jnp.max(x), \"max of x\")\nprint('starting')\ni = 0\nfor x, n in mnop:\n    x = x.numpy()\n    if i >= 300:\n        break\n    params, opt_state, loss_val = update(params, opt_state, x / 255., x / 255.)\n    if i % 10 == 0:\n        print(loss_val)\n    i += 1\n# print(x.shape)\n# print(jnp.max(x[0] / 255.))\n# print(jnp.min(x[0]))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-14T15:42:07.900112Z","iopub.execute_input":"2024-09-14T15:42:07.900570Z","iopub.status.idle":"2024-09-14T15:49:13.631381Z","shell.execute_reply.started":"2024-09-14T15:42:07.900533Z","shell.execute_reply":"2024-09-14T15:49:13.630313Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"jnp.savez(\"vae.npz\", x=params)","metadata":{"execution":{"iopub.status.busy":"2024-09-14T16:00:43.167422Z","iopub.execute_input":"2024-09-14T16:00:43.168195Z","iopub.status.idle":"2024-09-14T16:00:43.174403Z","shell.execute_reply.started":"2024-09-14T16:00:43.168156Z","shell.execute_reply":"2024-09-14T16:00:43.172823Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"params","metadata":{"execution":{"iopub.status.busy":"2024-09-14T15:59:13.252805Z","iopub.execute_input":"2024-09-14T15:59:13.253229Z","iopub.status.idle":"2024-09-14T15:59:13.427475Z","shell.execute_reply.started":"2024-09-14T15:59:13.253191Z","shell.execute_reply":"2024-09-14T15:59:13.426479Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import jax\nimport jax.numpy as jnp\nimport optax\nfrom jax import random, grad, pmap\nfrom functools import partial\n\n# Define a simple loss function\n@partial(jax.pmap, axis_name='batch')\ndef loss_fn(params, batch):\n    predictions = jnp.dot(batch, params)\n    loss = jax.lax.pmean(jnp.mean((predictions - 1.0) ** 2), axis_name = 'batch')\n    return loss\n\n# Define an update function\ndef update(params, opt_state, batch):\n    loss, grads = jax.value_and_grad(loss_fn)(params, batch)\n    updates, opt_state = optimizer.update(grads, opt_state)\n    params = optax.apply_updates(params, updates)\n    return params, opt_state, loss\n\n# Create an optimizer (Optax SGD for example)\noptimizer = optax.sgd(learning_rate=0.01)\n\n# Initialize model parameters and optimizer state\nparams = jnp.ones((100,))  # example parameters\nopt_state = optimizer.init(params)\n\n# Create some dummy data\nbatch = jax.device_put(jnp.ones((2, 320000, 100)))\n\n# Replicate params and optimizer state across devices\nparams_replicated = jax.device_put_replicated(params, jax.devices())\nopt_state_replicated = jax.device_put_replicated(opt_state, jax.devices())\n\n# pmap the update function across devices\n\n# Perform the update in parallel\nfor i in range(10000):\n    params_replicated, opt_state_replicated, loss = update(\n        params_replicated, opt_state_replicated, batch\n    )\n    if i % 1000 == 0:\n        print(loss)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-17T13:27:14.109570Z","iopub.execute_input":"2024-09-17T13:27:14.110061Z","iopub.status.idle":"2024-09-17T13:27:14.306689Z","shell.execute_reply.started":"2024-09-17T13:27:14.110023Z","shell.execute_reply":"2024-09-17T13:27:14.305412Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import jax\nimport jax.numpy as jnp\nimport optax\nfrom jax import random, grad, pmap\nfrom functools import partial\n\n# Define a simple loss function\ndef loss_fn(params, batch):\n    predictions = jnp.dot(batch, params[0])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[1])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[2])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[3])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[4])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[5])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[6])\n    predictions /= jnp.max(predictions)\n    predictions = jnp.dot(predictions, params[7])\n    # Compute the mean loss across the batch axis (reduce to a scalar)\n    per_device_loss = jnp.mean((predictions - 1.0) ** 2, axis=0)\n    # Synchronize the loss across devices\n    loss = jax.lax.pmean(per_device_loss, axis_name='device')\n    return loss\n\n# Define an update function\ndef update(params, opt_state, batch):\n    # Compute loss and gradients\n    loss, grads = jax.value_and_grad(loss_fn)(params, batch)\n    # Update parameters using the optimizer\n    updates, opt_state = optimizer.update(grads, opt_state)\n    params = optax.apply_updates(params, updates)\n    return params, opt_state, loss\n\n# Create an optimizer (Optax SGD for example)\noptimizer = optax.sgd(learning_rate=0.01)\n\n# Initialize model parameters and optimizer state\nparams = [jnp.ones((100, 100)), jnp.ones((100, 100)), jnp.ones((100, 100)), jnp.ones((100, 100)), jnp.ones((100, 100)), jnp.ones((100, 100)), jnp.ones((100, 100)), jnp.ones(100,)]  # example parameters\nopt_state = optimizer.init(params)\n\n# Create some dummy data\nbatch = jax.device_put(jnp.ones((2, 32000, 100)))\nbatch /= jnp.max(batch)\n\n# Replicate params and optimizer state across devices\nparams_replicated = jax.device_put_replicated(params, jax.devices())\nopt_state_replicated = jax.device_put_replicated(opt_state, jax.devices())\n\n# Perform the update in parallel\nfor i in range(10000):\n    params_replicated, opt_state_replicated, loss = pmap(update, axis_name = 'device')(\n        params_replicated, opt_state_replicated, batch\n    )\n    if i % 1000 == 0:\n        print(loss)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-17T13:40:43.063855Z","iopub.execute_input":"2024-09-17T13:40:43.064237Z"},"trusted":true},"outputs":[],"execution_count":null}]}