[
  {
    "number": 2837,
    "url": "https://api.github.com/repos/ml-explore/mlx/issues/2837/comments?per_page=100",
    "comments": [
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/3583254650",
        "html_url": "https://github.com/ml-explore/mlx/issues/2837#issuecomment-3583254650",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/2837",
        "id": 3583254650,
        "node_id": "IC_kwDOKzRn187VlCB6",
        "user": {
          "login": "CC-Yeh",
          "id": 46629671,
          "node_id": "MDQ6VXNlcjQ2NjI5Njcx",
          "avatar_url": "https://avatars.githubusercontent.com/u/46629671?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/CC-Yeh",
          "html_url": "https://github.com/CC-Yeh",
          "followers_url": "https://api.github.com/users/CC-Yeh/followers",
          "following_url": "https://api.github.com/users/CC-Yeh/following{/other_user}",
          "gists_url": "https://api.github.com/users/CC-Yeh/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/CC-Yeh/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/CC-Yeh/subscriptions",
          "organizations_url": "https://api.github.com/users/CC-Yeh/orgs",
          "repos_url": "https://api.github.com/users/CC-Yeh/repos",
          "events_url": "https://api.github.com/users/CC-Yeh/events{/privacy}",
          "received_events_url": "https://api.github.com/users/CC-Yeh/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2025-11-26T21:23:06Z",
        "updated_at": "2025-11-26T21:24:21Z",
        "body": "Maybe `MultiOptimizer` can help in your case?\n\nhttps://ml-explore.github.io/mlx/build/html/python/optimizers/_autosummary/mlx.optimizers.MultiOptimizer.html#mlx.optimizers.MultiOptimizer",
        "author_association": "CONTRIBUTOR",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/3583254650/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/3584734540",
        "html_url": "https://github.com/ml-explore/mlx/issues/2837#issuecomment-3584734540",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/2837",
        "id": 3584734540,
        "node_id": "IC_kwDOKzRn187VqrVM",
        "user": {
          "login": "yuchaoran2011",
          "id": 1168769,
          "node_id": "MDQ6VXNlcjExNjg3Njk=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1168769?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/yuchaoran2011",
          "html_url": "https://github.com/yuchaoran2011",
          "followers_url": "https://api.github.com/users/yuchaoran2011/followers",
          "following_url": "https://api.github.com/users/yuchaoran2011/following{/other_user}",
          "gists_url": "https://api.github.com/users/yuchaoran2011/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/yuchaoran2011/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/yuchaoran2011/subscriptions",
          "organizations_url": "https://api.github.com/users/yuchaoran2011/orgs",
          "repos_url": "https://api.github.com/users/yuchaoran2011/repos",
          "events_url": "https://api.github.com/users/yuchaoran2011/events{/privacy}",
          "received_events_url": "https://api.github.com/users/yuchaoran2011/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2025-11-27T08:37:57Z",
        "updated_at": "2025-11-27T08:37:57Z",
        "body": "@CC-Yeh How come I didn't find it earlier. That's exactly what I need. Thanks!",
        "author_association": "CONTRIBUTOR",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/3584734540/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      }
    ]
  },
  {
    "number": 1622,
    "url": "https://api.github.com/repos/ml-explore/mlx/issues/1622/comments?per_page=100",
    "comments": [
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2496460186",
        "html_url": "https://github.com/ml-explore/mlx/issues/1622#issuecomment-2496460186",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1622",
        "id": 2496460186,
        "node_id": "IC_kwDOKzRn186UzPWa",
        "user": {
          "login": "awni",
          "id": 1542805,
          "node_id": "MDQ6VXNlcjE1NDI4MDU=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1542805?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/awni",
          "html_url": "https://github.com/awni",
          "followers_url": "https://api.github.com/users/awni/followers",
          "following_url": "https://api.github.com/users/awni/following{/other_user}",
          "gists_url": "https://api.github.com/users/awni/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/awni/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/awni/subscriptions",
          "organizations_url": "https://api.github.com/users/awni/orgs",
          "repos_url": "https://api.github.com/users/awni/repos",
          "events_url": "https://api.github.com/users/awni/events{/privacy}",
          "received_events_url": "https://api.github.com/users/awni/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-11-25T00:41:59Z",
        "updated_at": "2024-11-25T00:41:59Z",
        "body": "Yes, exactly like that :).",
        "author_association": "MEMBER",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2496460186/reactions",
          "total_count": 1,
          "+1": 1,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      }
    ]
  },
  {
    "number": 1043,
    "url": "https://api.github.com/repos/ml-explore/mlx/issues/1043/comments?per_page=100",
    "comments": [
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081271463",
        "html_url": "https://github.com/ml-explore/mlx/pull/1043#issuecomment-2081271463",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1043",
        "id": 2081271463,
        "node_id": "IC_kwDOKzRn1858Da6n",
        "user": {
          "login": "NripeshN",
          "id": 86844847,
          "node_id": "MDQ6VXNlcjg2ODQ0ODQ3",
          "avatar_url": "https://avatars.githubusercontent.com/u/86844847?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/NripeshN",
          "html_url": "https://github.com/NripeshN",
          "followers_url": "https://api.github.com/users/NripeshN/followers",
          "following_url": "https://api.github.com/users/NripeshN/following{/other_user}",
          "gists_url": "https://api.github.com/users/NripeshN/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/NripeshN/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/NripeshN/subscriptions",
          "organizations_url": "https://api.github.com/users/NripeshN/orgs",
          "repos_url": "https://api.github.com/users/NripeshN/repos",
          "events_url": "https://api.github.com/users/NripeshN/events{/privacy}",
          "received_events_url": "https://api.github.com/users/NripeshN/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-28T01:02:06Z",
        "updated_at": "2024-04-28T01:02:06Z",
        "body": "@awni \r\nTests pass locally, ready to merge",
        "author_association": "CONTRIBUTOR",
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081271463/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      }
    ]
  },
  {
    "number": 1049,
    "url": "https://api.github.com/repos/ml-explore/mlx/issues/1049/comments?per_page=100",
    "comments": [
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081562088",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2081562088",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2081562088,
        "node_id": "IC_kwDOKzRn1858Eh3o",
        "user": {
          "login": "awni",
          "id": 1542805,
          "node_id": "MDQ6VXNlcjE1NDI4MDU=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1542805?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/awni",
          "html_url": "https://github.com/awni",
          "followers_url": "https://api.github.com/users/awni/followers",
          "following_url": "https://api.github.com/users/awni/following{/other_user}",
          "gists_url": "https://api.github.com/users/awni/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/awni/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/awni/subscriptions",
          "organizations_url": "https://api.github.com/users/awni/orgs",
          "repos_url": "https://api.github.com/users/awni/repos",
          "events_url": "https://api.github.com/users/awni/events{/privacy}",
          "received_events_url": "https://api.github.com/users/awni/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-28T17:23:28Z",
        "updated_at": "2024-04-28T17:23:28Z",
        "body": "Actually, that's not true. `LayerNorm` only and always normalizes normalizes over the last axis.\r\n\r\n```python\r\nimport mlx.nn as nn\r\nimport mlx.core as mx\r\n\r\nln = nn.LayerNorm(32)\r\nx = mx.random.uniform(shape=(10, 32))\r\nprint(ln(x).sum(axis=-1)) # close to 0\r\nprint(ln(x).sum(axis=0)) # not close to 0\r\n```\r\n\r\nSince MLX NN standardizes on the feature dimension being last, we don't have plans to include an axis parameter in our LayerNorm. It sounds like the current behavior works for what you want since it is consistent with Keras?\r\n\r\nIf there is something more here, let me know and we can reopen/discuss further.",
        "author_association": "MEMBER",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081562088/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081581363",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2081581363",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2081581363,
        "node_id": "IC_kwDOKzRn1858Emkz",
        "user": {
          "login": "thegodone",
          "id": 1186658,
          "node_id": "MDQ6VXNlcjExODY2NTg=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1186658?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/thegodone",
          "html_url": "https://github.com/thegodone",
          "followers_url": "https://api.github.com/users/thegodone/followers",
          "following_url": "https://api.github.com/users/thegodone/following{/other_user}",
          "gists_url": "https://api.github.com/users/thegodone/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/thegodone/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/thegodone/subscriptions",
          "organizations_url": "https://api.github.com/users/thegodone/orgs",
          "repos_url": "https://api.github.com/users/thegodone/repos",
          "events_url": "https://api.github.com/users/thegodone/events{/privacy}",
          "received_events_url": "https://api.github.com/users/thegodone/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-28T18:08:35Z",
        "updated_at": "2024-04-28T18:13:23Z",
        "body": "But If I use the weights and bias from Keras in LayerNorm I don't get the same result , why ? \r\n\r\nIf I comment remove LayerNorm on both models assert is working.\r\n\r\n````\r\nimport tensorflow as tf\r\nfrom tensorflow.keras.models import Sequential\r\nfrom tensorflow.keras.layers import Dense, Embedding, LayerNormalization\r\nimport numpy as np\r\nimport mlx.nn as nn\r\nfrom mlx.utils import tree_flatten\r\nimport mlx.core as mx\r\n\r\n# Define the model\r\nmodel = Sequential([\r\n    Embedding(input_dim=40,  output_dim=32, input_length=20,name=\"Embedding\"),\r\n    LayerNormalization(name='LN1'),\r\n\r\n])\r\n\r\n# Here input_dim is the number of input features, and output_dim is the number of output features.\r\nmodel.compile(\r\n    optimizer='adam',  # Optimizer\r\n    loss='mse',  # Mean Squared Error for regression tasks\r\n    metrics=['mae']  # Mean Absolute Error for regression metrics\r\n)\r\nmodel.summary()\r\n\r\ndef exposeweights(model):\r\n    w = {}\r\n    j =0\r\n    for layer in model.layers:\r\n        weights = layer.get_weights()  # returns a list of all weight tensors in the layer\r\n        print(layer.name)\r\n        for i, weight in enumerate(weights):\r\n            # is the Dense / linear are opposite array (ie Transpose) ?\r\n            if layer.name+\".\"+str(i) in ['Output.0','Proj.0','TimeDistributed.0'] :\r\n                w[layer.name+\".\"+str(i)]=weight.T\r\n            else:\r\n                w[layer.name+\".\"+str(i)]=weight\r\n            j+=1\r\n    return w\r\n\r\nw = exposeweights(model)\r\nnp.savez('w.npz', **w)\r\n\r\nclass mlxduplicate(nn.Module):\r\n    def __init__(\r\n        self):\r\n        super().__init__()\r\n        self.Embedding = nn.Embedding(num_embeddings=40, dims=32)\r\n        self.LN1 = nn.LayerNorm(32)\r\n\r\n    def __call__(self, x):\r\n        x = self.Embedding (x)\r\n        x = self.LN1(x)\r\n        return x \r\n\r\n\r\nmodel_mlx = mlxduplicate()\r\nwe = 0\r\nfor k, x in tree_flatten(model_mlx.parameters()):\r\n    we+=x.size\r\n    print(x.size,k)\r\nprint(we)\r\n\r\ntensor_loaded = np.load('w.npz')\r\n\r\ndef replace_key(key: str) -> str:\r\n    key = key.replace(\"Embedding.0\", \"Embedding.weight\")\r\n    key = key.replace(\"LN1.0\", \"LN1.weight\")\r\n    key = key.replace(\"LN1.1\", \"LN1.bias\")\r\n    return key\r\n\r\n# switch layer names of saved keras tensors\r\ntensors_mlx = {\r\n    replace_key(key): tensor for key, tensor in tensor_loaded.items()\r\n}\r\n\r\nfor k,v in tensor_loaded.items():\r\n    print(k,v.shape)\r\n\r\nfor k,v in tensors_mlx.items():\r\n    print(k,v.shape)\r\n\r\nnp.savez('w_convert_to_mlx.npz', **tensors_mlx)\r\n\r\nmodel_mlx.load_weights('w_convert_to_mlx.npz')\r\n\r\nx_train = np.random.randint(0,39, (2,20))\r\nkeras_output = model.predict(x_train)\r\nkeras_output.shape\r\n\r\nmlx_output = model_mlx(mx.array(x_train))\r\nassert mlx_output.shape == keras_output.shape\r\n\r\n\r\nassert np.max(np.abs(mlx_output-mx.array(keras_output))) < 1e-6\r\n``` \r\n\r\n2024-04-28 20:12:54.329876: I metal_plugin[/src/device/metal_device.cc:1154](http://localhost:8888/src/device/metal_device.cc#line=1153)] Metal device set to: Apple M3 Max\r\n2024-04-28 20:12:54.329898: I metal_plugin[/src/device/metal_device.cc:296](http://localhost:8888/src/device/metal_device.cc#line=295)] systemMemory: 128.00 GB\r\n2024-04-28 20:12:54.329900: I metal_plugin[/src/device/metal_device.cc:313](http://localhost:8888/src/device/metal_device.cc#line=312)] maxCacheSize: 48.00 GB\r\n2024-04-28 20:12:54.329932: I tensorflow[/core/common_runtime/pluggable_device/pluggable_device_factory.cc:306](http://localhost:8888/core/common_runtime/pluggable_device/pluggable_device_factory.cc#line=305)] Could not identify NUMA node of platform GPU ID 0, defaulting to 0. Your kernel may not have been built with NUMA support.\r\n2024-04-28 20:12:54.329949: I tensorflow[/core/common_runtime/pluggable_device/pluggable_device_factory.cc:272](http://localhost:8888/core/common_runtime/pluggable_device/pluggable_device_factory.cc#line=271)] Created TensorFlow device ([/job](http://localhost:8888/job):localhost[/replica:0](http://localhost:8888/replica#line=-1)[/task:0](http://localhost:8888/task#line=-1)[/device](http://localhost:8888/device):GPU:0 with 0 MB memory) -> physical PluggableDevice (device: 0, name: METAL, pci bus id: <undefined>)\r\nModel: \"sequential\"\r\n_________________________________________________________________\r\n Layer (type)                Output Shape              Param #   \r\n=================================================================\r\n Embedding (Embedding)       (None, 20, 32)            1280      \r\n                                                                 \r\n LN1 (LayerNormalization)    (None, 20, 32)            64        \r\n                                                                 \r\n=================================================================\r\nTotal params: 1344 (5.25 KB)\r\nTrainable params: 1344 (5.25 KB)\r\nNon-trainable params: 0 (0.00 Byte)\r\n_________________________________________________________________\r\nEmbedding\r\nLN1\r\n1280 Embedding.weight\r\n32 LN1.bias\r\n32 LN1.weight\r\n1344\r\nEmbedding.0 (40, 32)\r\nLN1.0 (32,)\r\nLN1.1 (32,)\r\nEmbedding.weight (40, 32)\r\nLN1.weight (32,)\r\nLN1.bias (32,)\r\n2024-04-28 20:12:54.785211: I tensorflow[/core/grappler/optimizers/custom_graph_optimizer_registry.cc:117](http://localhost:8888/core/grappler/optimizers/custom_graph_optimizer_registry.cc#line=116)] Plugin optimizer for device_type GPU is enabled.\r\n1/1 [==============================] - 1s 832ms/step\r\n---------------------------------------------------------------------------\r\nAssertionError                            Traceback (most recent call last)\r\nCell In[2], line 93\r\n     89 mlx_output = model_mlx(mx.array(x_train))\r\n     90 assert mlx_output.shape == keras_output.shape\r\n---> 93 assert np.max(np.abs(mlx_output-mx.array(keras_output))) < 1e-6\r\n\r\nAssertionError:",
        "author_association": "NONE",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081581363/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081735401",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2081735401",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2081735401,
        "node_id": "IC_kwDOKzRn1858FMLp",
        "user": {
          "login": "awni",
          "id": 1542805,
          "node_id": "MDQ6VXNlcjE1NDI4MDU=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1542805?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/awni",
          "html_url": "https://github.com/awni",
          "followers_url": "https://api.github.com/users/awni/followers",
          "following_url": "https://api.github.com/users/awni/following{/other_user}",
          "gists_url": "https://api.github.com/users/awni/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/awni/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/awni/subscriptions",
          "organizations_url": "https://api.github.com/users/awni/orgs",
          "repos_url": "https://api.github.com/users/awni/repos",
          "events_url": "https://api.github.com/users/awni/events{/privacy}",
          "received_events_url": "https://api.github.com/users/awni/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T01:00:38Z",
        "updated_at": "2024-04-29T01:00:38Z",
        "body": "It looks like Keras uses a different default (and much higher) `epsilon` for numerical stability. You can set this in the MLX `LayerNorm` constructor. The following passes:\r\n\r\n```Python\r\nfrom tensorflow.keras.layers import LayerNormalization\r\nimport numpy as np\r\nimport mlx.nn as nn\r\nimport mlx.core as mx\r\n\r\n# Define the model\r\nln = LayerNormalization(name='LN1')\r\nx = np.random.uniform(size=(10, 32))\r\nout_keras = np.array(ln(x))\r\n\r\nln_mlx = nn.LayerNorm(32, eps=1e-3) # note setting epsilon here\r\nout_mlx = np.array(ln_mlx(mx.array(x)))\r\nassert np.abs((out_keras - out_mlx)).max() < 1e-6\r\n```",
        "author_association": "MEMBER",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2081735401/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082008513",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2082008513",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2082008513,
        "node_id": "IC_kwDOKzRn1858GO3B",
        "user": {
          "login": "thegodone",
          "id": 1186658,
          "node_id": "MDQ6VXNlcjExODY2NTg=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1186658?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/thegodone",
          "html_url": "https://github.com/thegodone",
          "followers_url": "https://api.github.com/users/thegodone/followers",
          "following_url": "https://api.github.com/users/thegodone/following{/other_user}",
          "gists_url": "https://api.github.com/users/thegodone/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/thegodone/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/thegodone/subscriptions",
          "organizations_url": "https://api.github.com/users/thegodone/orgs",
          "repos_url": "https://api.github.com/users/thegodone/repos",
          "events_url": "https://api.github.com/users/thegodone/events{/privacy}",
          "received_events_url": "https://api.github.com/users/thegodone/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T06:57:58Z",
        "updated_at": "2024-04-29T07:01:17Z",
        "body": "thanks @awni  very much appreciate your help on that: I still have one question can you explain me this error for very large dataset I have a strange behaviour ? \r\n![image](https://github.com/ml-explore/mlx/assets/1186658/8e7f6424-1641-4db0-a63d-3c1d48c1cc91)\r\nI use this code : \r\n```\r\nimport math\r\nfrom typing import Any\r\nimport mlx.nn as nn\r\nfrom mlx.utils import tree_flatten\r\nimport numpy as np\r\nimport mlx.core as mx\r\n\r\nfrom mlx.nn.layers.base import Module\r\n\r\nclass AttentionM_(Module):\r\n    def __init__(self, input_dims: int, output_dims: int, bias: bool = True) -> None:\r\n        super().__init__()\r\n        self.output_dims = output_dims\r\n        scale = math.sqrt(1.0 / input_dims)\r\n        self.weight = mx.random.uniform(\r\n            low=-scale,\r\n            high=scale,\r\n            shape=(input_dims, 1),\r\n        )\r\n        if bias:\r\n            self.bias = mx.zeros(shape=(output_dims, 1))\r\n\r\n    def _extra_repr(self) -> str:\r\n        return f\"input_dims={self.weight.shape[0]}, output_dims={self.output_dims}, bias={'bias' in self}\"\r\n\r\n    def __call__(self, x: mx.array) -> mx.array:\r\n        if \"bias\" in self:\r\n            x_ = mx.addmm(self[\"bias\"], x, self[\"weight\"])\r\n        else:\r\n            x_ = x @ self[\"weight\"]\r\n        x_ = mx.tanh(x_)\r\n        x_ = mx.expand_dims(mx.softmax(mx.squeeze(x_,axis=-1), axis=-1),axis=-1)\r\n        x = mx.sum(x*x_,axis=1)\r\n        return x\r\n        \r\nclass TimeDistributed_(nn.Module):\r\n    def __init__(\r\n        self,\r\n        func : nn.Module\r\n    ):\r\n        super().__init__()\r\n        self.func = func\r\n\r\n    def __call__(self, x):\r\n        b_, t_ = x.shape[:2]\r\n        c_ = self.func(x.flatten(0,1))\r\n        return c_.reshape(b_, t_, *c_.shape[1:])\r\n        \r\nclass Bidirectionnal_(nn.Module):\r\n    def __init__(\r\n        self,\r\n        func1 : nn.Module,\r\n        func2 : nn.Module\r\n    ):\r\n        super().__init__()\r\n        self.func1 = func1\r\n        self.func2 = func2\r\n\r\n    def __call__(self, x):\r\n\r\n        h_f, h_b = self.func1(x), self.func2(x[:, ::-1, :]) \r\n        return  mx.stack([h_f[0], h_b[0]], axis=-1).flatten(-2,-1)\r\n    \r\n\r\n\r\nclass SmilesX(nn.Module):\r\n    def __init__(\r\n        self,\r\n        vocab_size: int,\r\n        inputdim: int,\r\n        embdim: int,\r\n        lstmdim: int,\r\n        densedim1: int,\r\n        densedim2: int,\r\n        checkpoint: bool,\r\n        debug: bool,\r\n    ):\r\n        super().__init__()\r\n\r\n        self.Embedding = nn.Embedding(num_embeddings=vocab_size, dims=embdim)\r\n        \r\n        self.Image = Bidirectionnal_(nn.LSTM(embdim, lstmdim, bias=True),\r\n                                    nn.LSTM(embdim, lstmdim, bias=True))\r\n        self.TimeDistributed = TimeDistributed_(nn.Linear(2*lstmdim,densedim1))\r\n        self.AttentionM  = AttentionM_(densedim1,inputdim, bias=True)\r\n        self.Layernorm1 =  nn.LayerNorm(densedim1,eps=0.001)\r\n        self.Proj = nn.Linear(densedim1,densedim2)\r\n        self.Layernorm2 =  nn.LayerNorm(densedim2,eps=0.001)\r\n        self.Output = nn.Linear(densedim2,1)\r\n        self.lk = nn.LeakyReLU(0.1)\r\n        self.debug = debug\r\n    \r\n    def __call__(self, x):\r\n        if self.debug:\r\n            print('Input:',x.shape)\r\n\r\n        # embedding\r\n        x = self.Embedding(x)\r\n        if self.debug:\r\n\r\n            print('Embedding:',x.shape)\r\n        # Bidirectional \r\n        x = self.Image(x)\r\n        if self.debug:\r\n            print('BiLSTM:',x.shape)\r\n\r\n        #  TimeDistributed \r\n        x = self.TimeDistributed(x)\r\n        if self.debug:\r\n            print('TimeDistributed:',x.shape)\r\n       # self attention\r\n        x = self.AttentionM(x)\r\n        if self.debug:\r\n            print('AttentionM:',x.shape)\r\n\r\n        # Layer norm \r\n        x = self.Layernorm1(x)\r\n        if self.debug:\r\n            print('LayerNorm 1:',x.shape)        \r\n        x = self.Proj(x)        \r\n        if self.debug:\r\n            print('proj:',x.shape)\r\n\r\n        x = self.lk(x)\r\n        # Layer norm \r\n        x = self.Layernorm2(x)\r\n        if self.debug:\r\n            print('LayerNorm 2:',x.shape)\r\n\r\n        x = self.Output(x)\r\n        if self.debug:\r\n            print('Output:',x.shape)\r\n        return x\r\n\r\n\r\nmodel = SmilesX(vocab_size=42, \r\n                inputdim = 128,\r\n                embdim = 32,\r\n                lstmdim =  32,\r\n                densedim1 = 64,\r\n                densedim2 = 64,\r\n                checkpoint=False,\r\n                debug=False)\r\n\r\n\r\n\r\n# Initialize model:\r\nnparams = sum(\r\n    x.size for k, x in tree_flatten(model.parameters()))\r\nprint(f\"Training a SMILES-X Model with {nparams}  parameters\")\r\n\r\nxt = 0\r\nfor k, x in tree_flatten(model.parameters()):\r\n    print(x.size,k)\r\n    xt+=x.size\r\nprint(xt)\r\n\r\n# test for big dataset:\r\nX = mx.random.randint(0,42,[600000,128])\r\n\r\n#\u00a0evaluate by data size the results\r\nY1 = model(X[:10,:])\r\nY3 = model(X[:100,:])\r\nY2 = model(X[:200,:])\r\nY4 = model(X[:1000,:])\r\nY5 = model(X[:10000,:])\r\nY6 = model(X[:100000,:])\r\nY7 = model(X[:600000,:])\r\n\r\n# validate the results \r\nassert np.max(np.abs(Y1 - Y3[:10]))<1e-8\r\nassert np.max(np.abs(Y1 - Y2[:10]))<1e-8\r\nassert np.max(np.abs(Y1 - Y4[:10]))<1e-8\r\nassert np.max(np.abs(Y1 - Y5[:10]))<1e-8\r\nassert np.max(np.abs(Y1 - Y6[:10]))<1e-8\r\nassert np.max(np.abs(Y1 - Y7[:10]))<1e-8\r\n\r\n``` ",
        "author_association": "NONE",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082008513/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082750857",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2082750857",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2082750857,
        "node_id": "IC_kwDOKzRn1858JEGJ",
        "user": {
          "login": "awni",
          "id": 1542805,
          "node_id": "MDQ6VXNlcjE1NDI4MDU=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1542805?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/awni",
          "html_url": "https://github.com/awni",
          "followers_url": "https://api.github.com/users/awni/followers",
          "following_url": "https://api.github.com/users/awni/following{/other_user}",
          "gists_url": "https://api.github.com/users/awni/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/awni/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/awni/subscriptions",
          "organizations_url": "https://api.github.com/users/awni/orgs",
          "repos_url": "https://api.github.com/users/awni/repos",
          "events_url": "https://api.github.com/users/awni/events{/privacy}",
          "received_events_url": "https://api.github.com/users/awni/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T13:27:27Z",
        "updated_at": "2024-04-29T13:27:27Z",
        "body": "The size is really large, I think some matrices are well over 4B entries. My guess is it's overflowing an integer index somewhere but I'm not sure where. I'll look into where that is to see if we can put a error message or fix it. For now I would stick to smaller sizes.",
        "author_association": "MEMBER",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082750857/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082797075",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2082797075",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2082797075,
        "node_id": "IC_kwDOKzRn1858JPYT",
        "user": {
          "login": "awni",
          "id": 1542805,
          "node_id": "MDQ6VXNlcjE1NDI4MDU=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1542805?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/awni",
          "html_url": "https://github.com/awni",
          "followers_url": "https://api.github.com/users/awni/followers",
          "following_url": "https://api.github.com/users/awni/following{/other_user}",
          "gists_url": "https://api.github.com/users/awni/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/awni/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/awni/subscriptions",
          "organizations_url": "https://api.github.com/users/awni/orgs",
          "repos_url": "https://api.github.com/users/awni/repos",
          "events_url": "https://api.github.com/users/awni/events{/privacy}",
          "received_events_url": "https://api.github.com/users/awni/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T13:45:00Z",
        "updated_at": "2024-04-29T13:45:00Z",
        "body": "I filed a separate issue about this https://github.com/ml-explore/mlx/issues/1051",
        "author_association": "MEMBER",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082797075/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082847542",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2082847542",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2082847542,
        "node_id": "IC_kwDOKzRn1858Jbs2",
        "user": {
          "login": "thegodone",
          "id": 1186658,
          "node_id": "MDQ6VXNlcjExODY2NTg=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1186658?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/thegodone",
          "html_url": "https://github.com/thegodone",
          "followers_url": "https://api.github.com/users/thegodone/followers",
          "following_url": "https://api.github.com/users/thegodone/following{/other_user}",
          "gists_url": "https://api.github.com/users/thegodone/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/thegodone/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/thegodone/subscriptions",
          "organizations_url": "https://api.github.com/users/thegodone/orgs",
          "repos_url": "https://api.github.com/users/thegodone/repos",
          "events_url": "https://api.github.com/users/thegodone/events{/privacy}",
          "received_events_url": "https://api.github.com/users/thegodone/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T14:04:12Z",
        "updated_at": "2024-04-29T14:04:12Z",
        "body": "Does it happens during evaluation too ? Would be nice to add batch size for inference\u00a0Envoy\u00e9 de mon iPhoneLe 29 avr. 2024 \u00e0 15:45, Awni Hannun ***@***.***> a \u00e9crit\u00a0:\ufeff\r\nI filed a separate issue about this #1051\r\n\r\n\u2014Reply to this email directly, view it on GitHub, or unsubscribe.You are receiving this because you authored the thread.Message ID: ***@***.***>",
        "author_association": "NONE",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082847542/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082883219",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2082883219",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2082883219,
        "node_id": "IC_kwDOKzRn1858JkaT",
        "user": {
          "login": "awni",
          "id": 1542805,
          "node_id": "MDQ6VXNlcjE1NDI4MDU=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1542805?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/awni",
          "html_url": "https://github.com/awni",
          "followers_url": "https://api.github.com/users/awni/followers",
          "following_url": "https://api.github.com/users/awni/following{/other_user}",
          "gists_url": "https://api.github.com/users/awni/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/awni/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/awni/subscriptions",
          "organizations_url": "https://api.github.com/users/awni/orgs",
          "repos_url": "https://api.github.com/users/awni/repos",
          "events_url": "https://api.github.com/users/awni/events{/privacy}",
          "received_events_url": "https://api.github.com/users/awni/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T14:20:08Z",
        "updated_at": "2024-04-29T14:20:08Z",
        "body": "> Does it happens during evaluation too ? Would be nice to add batch size for inference\r\n\r\nThe problem is the large matmul in the LSTM. So if the batch size is larger than about 131k (for the LSTM / model dimensions you provided) then it will break regardless of inference / training modes.",
        "author_association": "MEMBER",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082883219/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      },
      {
        "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082926425",
        "html_url": "https://github.com/ml-explore/mlx/issues/1049#issuecomment-2082926425",
        "issue_url": "https://api.github.com/repos/ml-explore/mlx/issues/1049",
        "id": 2082926425,
        "node_id": "IC_kwDOKzRn1858Ju9Z",
        "user": {
          "login": "thegodone",
          "id": 1186658,
          "node_id": "MDQ6VXNlcjExODY2NTg=",
          "avatar_url": "https://avatars.githubusercontent.com/u/1186658?v=4",
          "gravatar_id": "",
          "url": "https://api.github.com/users/thegodone",
          "html_url": "https://github.com/thegodone",
          "followers_url": "https://api.github.com/users/thegodone/followers",
          "following_url": "https://api.github.com/users/thegodone/following{/other_user}",
          "gists_url": "https://api.github.com/users/thegodone/gists{/gist_id}",
          "starred_url": "https://api.github.com/users/thegodone/starred{/owner}{/repo}",
          "subscriptions_url": "https://api.github.com/users/thegodone/subscriptions",
          "organizations_url": "https://api.github.com/users/thegodone/orgs",
          "repos_url": "https://api.github.com/users/thegodone/repos",
          "events_url": "https://api.github.com/users/thegodone/events{/privacy}",
          "received_events_url": "https://api.github.com/users/thegodone/received_events",
          "type": "User",
          "user_view_type": "public",
          "site_admin": false
        },
        "created_at": "2024-04-29T14:39:16Z",
        "updated_at": "2024-04-29T14:39:16Z",
        "body": "Thanks for clarification.Envoy\u00e9 de mon iPhoneLe 29 avr. 2024 \u00e0 16:20, Awni Hannun ***@***.***> a \u00e9crit\u00a0:\ufeff\r\n\r\nDoes it happens during evaluation too ? Would be nice to add batch size for inference\r\n\r\nThe problem is the large matmul in the LSTM. So if the batch size is larger than about 131k (for the LSTM / model dimensions you provided) then it will break regardless of inference / training modes.\r\n\r\n\u2014Reply to this email directly, view it on GitHub, or unsubscribe.You are receiving this because you authored the thread.Message ID: ***@***.***>",
        "author_association": "NONE",
        "pin": null,
        "reactions": {
          "url": "https://api.github.com/repos/ml-explore/mlx/issues/comments/2082926425/reactions",
          "total_count": 0,
          "+1": 0,
          "-1": 0,
          "laugh": 0,
          "hooray": 0,
          "confused": 0,
          "heart": 0,
          "rocket": 0,
          "eyes": 0
        },
        "performed_via_github_app": null,
        "minimized": null
      }
    ]
  }
]
