From ed80f41c5023f84f2b3079b0c1459bf5cdf91f05 Mon Sep 17 00:00:00 2001 From: cthotti Date: Sat, 4 Jul 2026 14:15:50 +0530 Subject: [PATCH 1/5] DFlash Phase 2: draft model definition, HF weight loading, and .pte export for Qwen3-4B --- .../mlx/examples/llm/dflash_draft_model.py | 227 ++++++++++++++++++ examples/models/qwen3/export_dflash_draft.py | 67 ++++++ .../qwen3/tests/test_dflash_draft_pte.py | 19 ++ 3 files changed, 313 insertions(+) create mode 100644 backends/mlx/examples/llm/dflash_draft_model.py create mode 100644 examples/models/qwen3/export_dflash_draft.py create mode 100644 examples/models/qwen3/tests/test_dflash_draft_pte.py diff --git a/backends/mlx/examples/llm/dflash_draft_model.py b/backends/mlx/examples/llm/dflash_draft_model.py new file mode 100644 index 00000000000..a573e64142b --- /dev/null +++ b/backends/mlx/examples/llm/dflash_draft_model.py @@ -0,0 +1,227 @@ +"""PyTorch DFlash draft model, structured for ExecuTorch export. + +Model-agnostic: the same code exports a valid draft .pte for any standard- +attention target (Qwen3, Gemma-4, Llama-3.1) by reading a DFlashConfig loaded +from the z-lab draft checkpoint. Per-model differences — RoPE base/scaling, +sliding-window layers, final-logit softcap, embedding scale — are config values, +so they resolve at trace time to one model-specific graph. The universal +branches add no ops to the exported program. + +Two deliberate deviations from the reference forward, both for ET (design doc +"Option A", self-contained draft): + - embed_tokens / lm_head are owned here (filled from the target at export) + instead of referenced live off the target module. + - forward returns draft logits (norm -> lm_head -> [:, 1:]) rather than the + bare normed hidden state; the reference does lm_head + logits_start=1 in its + generate loop. + +References: + z-lab/dflash dflash/model_mlx.py — universal MLX reference (Qwen3 / Qwen3.5 / + Gemma-4); source of the sliding-window, softcap, single-rope, QK-norm design. + z-lab/dflash dflash/model.py — PyTorch reference; exact weight layout for + load_state_dict. + transformers RoPE init — config-driven inv_freq below. +""" + +from dataclasses import dataclass, field +from typing import Any, Dict, Optional, Tuple + +import torch +from torch import nn + + +@dataclass +class DFlashConfig: + hidden_size: int + num_hidden_layers: int + num_attention_heads: int + num_key_value_heads: int + head_dim: int + intermediate_size: int + vocab_size: int + rms_norm_eps: float + rope_theta: float + max_position_embeddings: int + target_layer_ids: Tuple[int, ...] + block_size: int = 16 + mask_token_id: int = 0 + rope_scaling: Optional[Dict[str, Any]] = None + layer_types: Tuple[str, ...] = field(default_factory=tuple) + sliding_window: Optional[int] = None + final_logit_softcapping: Optional[float] = None + embed_scale: float = 1.0 # 1.0 for Qwen3/Llama; sqrt(hidden_size) for Gemma + + +def _rope_inv_freq(config: DFlashConfig) -> torch.Tensor: + # Covers the two rope types these drafts use: 'default' (Qwen3, Gemma-4) and + # 'linear' position-scaling. YaRN/longrope aren't handled — no current z-lab + # draft uses them. + dim = config.head_dim + inv_freq = 1.0 / (config.rope_theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) + scaling = config.rope_scaling or {} + if scaling.get("rope_type", scaling.get("type")) == "linear": + inv_freq = inv_freq / float(scaling["factor"]) + return inv_freq + + +class DFlashRotaryEmbedding(nn.Module): + def __init__(self, config: DFlashConfig): + super().__init__() + self.register_buffer("inv_freq", _rope_inv_freq(config), persistent=False) + + def forward(self, position_ids: torch.Tensor): + inv = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) + pos = position_ids[:, None, :].float() + freqs = (inv @ pos).transpose(1, 2) # [B, S, head_dim/2] + emb = torch.cat((freqs, freqs), dim=-1) # [B, S, head_dim] + return emb.cos(), emb.sin() + + +def rotate_half(x: torch.Tensor) -> torch.Tensor: + x1, x2 = x.chunk(2, dim=-1) + return torch.cat((-x2, x1), dim=-1) + + +def apply_rotary_pos_emb(q, k, cos, sin): + # q rotates over its own (last q_len) positions; k rotates over its full length. + q_len = q.shape[-2] + cq, sq = cos[:, None, -q_len:, :], sin[:, None, -q_len:, :] + ck, sk = cos[:, None, :, :], sin[:, None, :, :] + return (q * cq) + (rotate_half(q) * sq), (k * ck) + (rotate_half(k) * sk) + + +class DFlashRMSNorm(nn.Module): + def __init__(self, dim: int, eps: float): + super().__init__() + self.weight = nn.Parameter(torch.ones(dim)) + self.eps = eps + + def forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x = x.float() + x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) + return self.weight * x.to(dtype) + + +class DFlashMLP(nn.Module): + def __init__(self, hidden_size: int, intermediate_size: int): + super().__init__() + self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) + self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) + self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.down_proj(nn.functional.silu(self.gate_proj(x)) * self.up_proj(x)) + + +class DFlashAttention(nn.Module): + def __init__(self, config: DFlashConfig, layer_idx: int): + super().__init__() + h, hd = config.hidden_size, config.head_dim + self.n_heads = config.num_attention_heads + self.n_kv = config.num_key_value_heads + self.head_dim = hd + self.scaling = hd ** -0.5 + self.n_rep = self.n_heads // self.n_kv + lt = config.layer_types + self.is_sliding = bool(lt) and lt[layer_idx] == "sliding_attention" + self.sliding_window = config.sliding_window if self.is_sliding else None + self.q_proj = nn.Linear(h, self.n_heads * hd, bias=False) + self.k_proj = nn.Linear(h, self.n_kv * hd, bias=False) + self.v_proj = nn.Linear(h, self.n_kv * hd, bias=False) + self.o_proj = nn.Linear(self.n_heads * hd, h, bias=False) + self.q_norm = DFlashRMSNorm(hd, config.rms_norm_eps) + self.k_norm = DFlashRMSNorm(hd, config.rms_norm_eps) + + def forward(self, x, x_ctx, cos, sin): + B, L, _ = x.shape + S = x_ctx.shape[1] + q = self.q_norm(self.q_proj(x).view(B, L, self.n_heads, self.head_dim)).transpose(1, 2) + k = torch.cat([self.k_proj(x_ctx), self.k_proj(x)], dim=1).view(B, S + L, self.n_kv, self.head_dim) + v = torch.cat([self.v_proj(x_ctx), self.v_proj(x)], dim=1).view(B, S + L, self.n_kv, self.head_dim) + k = self.k_norm(k).transpose(1, 2) + v = v.transpose(1, 2) + q, k = apply_rotary_pos_emb(q, k, cos, sin) + if self.n_rep > 1: + k = k.repeat_interleave(self.n_rep, dim=1) + v = v.repeat_interleave(self.n_rep, dim=1) + mask = self._sliding_mask(L, S, q.device, q.dtype) if self.is_sliding else None + out = torch.nn.functional.scaled_dot_product_attention( + q, k, v, attn_mask=mask, is_causal=False, scale=self.scaling) + return self.o_proj(out.transpose(1, 2).reshape(B, L, -1)) + + def _sliding_mask(self, L, S, device, dtype): + total = S + L + q_pos = torch.arange(S, total, device=device)[:, None] + k_pos = torch.arange(total, device=device)[None, :] + allowed = (k_pos <= q_pos) & (k_pos > q_pos - self.sliding_window) + return torch.where(allowed, 0.0, float("-inf")).to(dtype)[None, None] + + +class DFlashDecoderLayer(nn.Module): + def __init__(self, config: DFlashConfig, layer_idx: int): + super().__init__() + self.self_attn = DFlashAttention(config, layer_idx) + self.mlp = DFlashMLP(config.hidden_size, config.intermediate_size) + self.input_layernorm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) + self.post_attention_layernorm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) + + def forward(self, x, x_ctx, cos, sin): + x = x + self.self_attn(self.input_layernorm(x), x_ctx, cos, sin) + return x + self.mlp(self.post_attention_layernorm(x)) + + +class DFlashDraftModel(nn.Module): + def __init__(self, config: DFlashConfig): + super().__init__() + self.config = config + concat_dim = len(config.target_layer_ids) * config.hidden_size + self.fc = nn.Linear(concat_dim, config.hidden_size, bias=False) + self.hidden_norm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) + self.layers = nn.ModuleList( + [DFlashDecoderLayer(config, i) for i in range(config.num_hidden_layers)]) + self.norm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) + self.rotary_emb = DFlashRotaryEmbedding(config) + # Option A: owned copies, filled from the target at export time. + self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) + self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) + + def forward(self, tokens, target_hidden, position_ids): + h = self.embed_tokens(tokens) * self.config.embed_scale + h_ctx = self.hidden_norm(self.fc(target_hidden)) + cos, sin = self.rotary_emb(position_ids) + for layer in self.layers: + h = layer(h, h_ctx, cos, sin) + h = self.norm(h) + logits = self.lm_head(h[:, 1:, :]) # logits_start=1: drop the known first token + cap = self.config.final_logit_softcapping + if cap is not None: + logits = torch.tanh(logits / cap) * cap + return logits + +def load_dflash_config(checkpoint_dir) -> "DFlashConfig": + """Build a DFlashConfig from a z-lab DFlash checkpoint's config.json.""" + import json + from pathlib import Path + + cfg = json.loads((Path(checkpoint_dir) / "config.json").read_text()) + dcfg = cfg["dflash_config"] + return DFlashConfig( + hidden_size=cfg["hidden_size"], + num_hidden_layers=cfg["num_hidden_layers"], + num_attention_heads=cfg["num_attention_heads"], + num_key_value_heads=cfg["num_key_value_heads"], + head_dim=cfg["head_dim"], + intermediate_size=cfg["intermediate_size"], + vocab_size=cfg["vocab_size"], + rms_norm_eps=cfg["rms_norm_eps"], + rope_theta=cfg["rope_theta"], + max_position_embeddings=cfg["max_position_embeddings"], + target_layer_ids=tuple(dcfg["target_layer_ids"]), + block_size=cfg["block_size"], + mask_token_id=dcfg["mask_token_id"], + rope_scaling=cfg.get("rope_scaling"), + layer_types=tuple(cfg.get("layer_types") or ["full_attention"] * cfg["num_hidden_layers"]), + sliding_window=cfg.get("sliding_window"), + final_logit_softcapping=cfg.get("final_logit_softcapping"), + ) diff --git a/examples/models/qwen3/export_dflash_draft.py b/examples/models/qwen3/export_dflash_draft.py new file mode 100644 index 00000000000..838ab14b6a5 --- /dev/null +++ b/examples/models/qwen3/export_dflash_draft.py @@ -0,0 +1,67 @@ +import argparse +from pathlib import Path + +import torch +from huggingface_hub import snapshot_download +from safetensors.torch import load_file +from transformers import AutoModelForCausalLM + +from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashDraftModel, load_dflash_config + + +def load_draft_model(draft_id: str, target_state_dict: dict) -> DFlashDraftModel: + path = Path(snapshot_download(draft_id, allow_patterns=["*.safetensors", "*.json"])) + config = load_dflash_config(path) + model = DFlashDraftModel(config) + + draft_weights = {} + for f in path.glob("*.safetensors"): + draft_weights.update(load_file(str(f))) + + missing, unexpected = model.load_state_dict(draft_weights, strict=False) + assert not unexpected, f"Unexpected draft checkpoint keys: {unexpected}" + still_missing = [k for k in missing if not k.startswith(("embed_tokens.", "lm_head."))] + assert not still_missing, f"Missing draft checkpoint keys: {still_missing}" + + model.embed_tokens.weight.data.copy_(target_state_dict["model.embed_tokens.weight"]) + lm_head_key = "lm_head.weight" if "lm_head.weight" in target_state_dict else "model.embed_tokens.weight" + model.lm_head.weight.data.copy_(target_state_dict[lm_head_key]) + return model + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--target-model", default="Qwen/Qwen3-4B") + parser.add_argument("--draft-model", default="z-lab/Qwen3-4B-DFlash-b16") + parser.add_argument("--output", default="qwen3_4b_dflash_draft.pte") + parser.add_argument("--block-size", type=int, default=16) + parser.add_argument("--ctx-len", type=int, default=8) + args = parser.parse_args() + + target = AutoModelForCausalLM.from_pretrained(args.target_model, dtype="auto") + model = load_draft_model(args.draft_model, target.state_dict()) + model.eval() + model = model.float() + del target + + block_size, ctx_len = args.block_size, args.ctx_len + hidden_size = model.fc.in_features + tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) + target_hidden = torch.randn(1, ctx_len, hidden_size) + position_ids = torch.arange(ctx_len + block_size).unsqueeze(0) + + exported = torch.export.export(model, (tokens, target_hidden, position_ids)) + + from executorch.exir import to_edge_transform_and_lower + from executorch.backends.mlx.partitioner import MLXPartitioner + + edge = to_edge_transform_and_lower(exported, partitioner=[MLXPartitioner()]) + et_program = edge.to_executorch() + + with open(args.output, "wb") as f: + f.write(et_program.buffer) + print(f"Saved draft model to: {args.output}") + + +if __name__ == "__main__": + main() diff --git a/examples/models/qwen3/tests/test_dflash_draft_pte.py b/examples/models/qwen3/tests/test_dflash_draft_pte.py new file mode 100644 index 00000000000..374bb0e2aa0 --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_draft_pte.py @@ -0,0 +1,19 @@ +import sys +import torch +from executorch.runtime import Runtime, Verification + +pte_path = sys.argv[1] if len(sys.argv) > 1 else "qwen3_4b_dflash_draft.pte" +et_runtime = Runtime.get() +program = et_runtime.load_program(pte_path, verification=Verification.Minimal) +method = program.load_method("forward") + +# Must match the exact static shapes used at export time: ctx_len=8, block_size=16 +block_size, ctx_len, hidden_size, vocab_size = 16, 8, 12800, 151936 +tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) +target_hidden = torch.randn(1, ctx_len, hidden_size) +position_ids = torch.arange(ctx_len + block_size).unsqueeze(0).long() + +(draft_logits,) = method.execute([tokens, target_hidden, position_ids]) +assert draft_logits.shape == (1, block_size - 1, vocab_size), draft_logits.shape +assert not torch.isnan(draft_logits).any() and not torch.isinf(draft_logits).any() +print(f"OK: draft_logits {tuple(draft_logits.shape)}") From e0420a374629ae644830e0a888ffd7d72232acf9 Mon Sep 17 00:00:00 2001 From: cthotti Date: Sun, 5 Jul 2026 21:15:14 +0530 Subject: [PATCH 2/5] DFlash: consolidate tests --- .gitignore | 1 + Makefile | 9 + .../mlx/examples/llm/dflash_draft_model.py | 10 +- backends/mlx/examples/llm/export_llm_hf.py | 37 +++- examples/models/qwen3/export_dflash_draft.py | 32 +++- .../qwen3/mlx_source_transformations.py | 110 ++++++++++++ examples/models/qwen3/run_baseline.py | 69 ++++++++ examples/models/qwen3/run_dflash.py | 162 ++++++++++++++++++ .../models/qwen3/tests/test_dflash_draft.py | 28 +++ .../qwen3/tests/test_dflash_draft_dynamic.py | 18 ++ .../qwen3/tests/test_dflash_draft_eager.py | 34 ++++ .../qwen3/tests/test_dflash_draft_forward.py | 46 +++++ .../tests/test_dflash_draft_load_weights.py | 50 ++++++ .../models/qwen3/tests/test_dflash_export.py | 32 ++++ .../qwen3/tests/test_dflash_lossless.py | 32 ++++ .../models/qwen3/tests/test_dflash_target.py | 32 ++++ 16 files changed, 695 insertions(+), 7 deletions(-) create mode 100644 examples/models/qwen3/mlx_source_transformations.py create mode 100644 examples/models/qwen3/run_baseline.py create mode 100644 examples/models/qwen3/run_dflash.py create mode 100644 examples/models/qwen3/tests/test_dflash_draft.py create mode 100644 examples/models/qwen3/tests/test_dflash_draft_dynamic.py create mode 100644 examples/models/qwen3/tests/test_dflash_draft_eager.py create mode 100644 examples/models/qwen3/tests/test_dflash_draft_forward.py create mode 100644 examples/models/qwen3/tests/test_dflash_draft_load_weights.py create mode 100644 examples/models/qwen3/tests/test_dflash_export.py create mode 100644 examples/models/qwen3/tests/test_dflash_lossless.py create mode 100644 examples/models/qwen3/tests/test_dflash_target.py diff --git a/.gitignore b/.gitignore index ee206e23d94..10b048c17b5 100644 --- a/.gitignore +++ b/.gitignore @@ -69,6 +69,7 @@ xcuserdata/ /src/executorch/include/ /src/executorch/share/ /src/executorch/version.py +/dflash_benchmarks.md *_etdump # Android diff --git a/Makefile b/Makefile index 969b53644cd..26167aeb2b2 100644 --- a/Makefile +++ b/Makefile @@ -484,3 +484,12 @@ clean: rm -rf cmake-out \ extension/llm/tokenizers/build \ extension/llm/tokenizers/pytorch_tokenizers.egg-info + +qwen3_dflash-mlx: + @echo "==> Building and installing ExecuTorch with MLX..." + cmake --workflow --preset mlx-release + @echo "==> Building Qwen3 DFlash speculative decoding runner with MLX..." + cd examples/models/qwen3 && cmake --workflow --preset qwen3-dflash-mlx + @echo "" + @echo "✓ Build complete!" + @echo " Runner: cmake-out/examples/models/qwen3/qwen3_dflash_runner" diff --git a/backends/mlx/examples/llm/dflash_draft_model.py b/backends/mlx/examples/llm/dflash_draft_model.py index a573e64142b..913f89e2d57 100644 --- a/backends/mlx/examples/llm/dflash_draft_model.py +++ b/backends/mlx/examples/llm/dflash_draft_model.py @@ -81,6 +81,12 @@ def rotate_half(x: torch.Tensor) -> torch.Tensor: x1, x2 = x.chunk(2, dim=-1) return torch.cat((-x2, x1), dim=-1) +def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor: + if n_rep == 1: + return x + b, h, s, d = x.shape + x = x[:, :, None, :, :].expand(b, h, n_rep, s, d) + return x.reshape(b, h * n_rep, s, d) def apply_rotary_pos_emb(q, k, cos, sin): # q rotates over its own (last q_len) positions; k rotates over its full length. @@ -143,8 +149,8 @@ def forward(self, x, x_ctx, cos, sin): v = v.transpose(1, 2) q, k = apply_rotary_pos_emb(q, k, cos, sin) if self.n_rep > 1: - k = k.repeat_interleave(self.n_rep, dim=1) - v = v.repeat_interleave(self.n_rep, dim=1) + k = repeat_kv(k, self.n_rep) + v = repeat_kv(v, self.n_rep) mask = self._sliding_mask(L, S, q.device, q.dtype) if self.is_sliding else None out = torch.nn.functional.scaled_dot_product_attention( q, k, v, attn_mask=mask, is_causal=False, scale=self.scaling) diff --git a/backends/mlx/examples/llm/export_llm_hf.py b/backends/mlx/examples/llm/export_llm_hf.py index fe6b8094f6b..45fd8bfad28 100644 --- a/backends/mlx/examples/llm/export_llm_hf.py +++ b/backends/mlx/examples/llm/export_llm_hf.py @@ -137,6 +137,7 @@ def _export_with_custom_components( no_tie_word_embeddings: bool = False, qlinear_group_size: Optional[int] = None, qembedding_group_size: Optional[int] = None, + dflash_layers: Optional[list[int]] = None, ) -> None: """ Export using direct HF model with custom MLX components. @@ -219,6 +220,24 @@ def _export_with_custom_components( batch_size=1, max_cache_len=effective_cache_len, ) + elif dflash_layers is not None: + # Qwen3-specific for now + # Generalize the import if/when another model needs DFlash tapping. + # + # Stateless (no persistent KV cache), NOT TorchExportableModuleWithStaticCacheAndHidden: + # DFlash's speculative-decode loop re-verifies overlapping/non-contiguous token + # ranges every round, which corrupts a persistent StaticCache (confirmed at the + # eager PyTorch level -- see StatelessQwen3WithHidden's docstring). The driver + # always passes the full accumulated sequence instead of relying on a cache. + from executorch.examples.models.qwen3.mlx_source_transformations import ( + StatelessQwen3WithHidden, + ) + + logger.info(f"Creating stateless DFlash hidden-state-tapping wrapper, layers={dflash_layers}") + exportable = StatelessQwen3WithHidden( + model=model, + layer_ids=dflash_layers, + ) else: logger.info("Creating TorchExportableModuleWithStaticCache wrapper...") exportable = TorchExportableModuleWithStaticCache( @@ -227,7 +246,7 @@ def _export_with_custom_components( max_cache_len=effective_cache_len, ) - if use_custom_kv_cache: + if use_custom_kv_cache and dflash_layers is None: from executorch.backends.mlx.llm.source_transformation import ( replace_hf_cache_with_mlx, ) @@ -335,6 +354,7 @@ def export_llama_hf( no_tie_word_embeddings: bool = False, qlinear_group_size: Optional[int] = None, qembedding_group_size: Optional[int] = None, + dflash_layers: Optional[list[int]] = None, ) -> None: """ Export a HuggingFace Llama model to ExecuTorch with MLX backend. @@ -349,10 +369,10 @@ def export_llama_hf( use_custom_sdpa: Use MLX custom SDPA (mlx::custom_sdpa) use_custom_kv_cache: Use MLX custom KV cache (mlx::kv_cache_update) """ - if use_custom_sdpa or use_custom_kv_cache: + if use_custom_sdpa or use_custom_kv_cache or dflash_layers is not None: logger.info( f"Using custom components: sdpa={use_custom_sdpa}, " - f"kv_cache={use_custom_kv_cache}" + f"kv_cache={use_custom_kv_cache}, dflash_layers={dflash_layers}" ) _export_with_custom_components( model_id=model_id, @@ -367,6 +387,7 @@ def export_llama_hf( no_tie_word_embeddings=no_tie_word_embeddings, qlinear_group_size=qlinear_group_size, qembedding_group_size=qembedding_group_size, + dflash_layers=dflash_layers, ) else: logger.info("Using optimum-executorch pipeline (no custom components)") @@ -434,8 +455,17 @@ def main(): default=False, help="Use MLX custom KV cache (mlx::kv_cache_update)", ) + parser.add_argument( + "--dflash-layers", + type=str, + default=None, + help="Comma-separated layer indices to tap for DFlash hidden-state output, e.g. '2,18,33'", + ) args = parser.parse_args() + dflash_layers = ( + [int(x) for x in args.dflash_layers.split(",")] if args.dflash_layers else None + ) export_llama_hf( model_id=args.model_id, @@ -450,6 +480,7 @@ def main(): no_tie_word_embeddings=args.no_tie_word_embeddings, qlinear_group_size=args.qlinear_group_size, qembedding_group_size=args.qembedding_group_size, + dflash_layers=dflash_layers, ) diff --git a/examples/models/qwen3/export_dflash_draft.py b/examples/models/qwen3/export_dflash_draft.py index 838ab14b6a5..75b72bd49cb 100644 --- a/examples/models/qwen3/export_dflash_draft.py +++ b/examples/models/qwen3/export_dflash_draft.py @@ -4,6 +4,7 @@ import torch from huggingface_hub import snapshot_download from safetensors.torch import load_file +from torch.export import Dim from transformers import AutoModelForCausalLM from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashDraftModel, load_dflash_config @@ -36,21 +37,47 @@ def main(): parser.add_argument("--output", default="qwen3_4b_dflash_draft.pte") parser.add_argument("--block-size", type=int, default=16) parser.add_argument("--ctx-len", type=int, default=8) + parser.add_argument("--max-ctx-len", type=int, default=4096) args = parser.parse_args() target = AutoModelForCausalLM.from_pretrained(args.target_model, dtype="auto") model = load_draft_model(args.draft_model, target.state_dict()) model.eval() - model = model.float() del target + # Quantize to 4-bit to match the target export (--qlinear 4w --qembedding 4w, + # group_size=32). Without this the draft is full float32 (~3.5GB) with its + # shared embed/lm_head dominating memory; quantizing brings it to ~1GB and, + # critically, keeps the shared embed/lm_head at the SAME precision as the + # target so their logits stay consistent (acceptance depends on this). + from executorch.backends.mlx.llm.quantization import quantize_model_ + quantize_model_( + model, + qlinear_config="4w", + qlinear_group_size=32, + qembedding_config="4w", + qembedding_group_size=32, + tie_word_embeddings=False, + ) + block_size, ctx_len = args.block_size, args.ctx_len hidden_size = model.fc.in_features tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) target_hidden = torch.randn(1, ctx_len, hidden_size) position_ids = torch.arange(ctx_len + block_size).unsqueeze(0) - exported = torch.export.export(model, (tokens, target_hidden, position_ids)) + ctx_dim = Dim("ctx_len", min=1, max=args.max_ctx_len) + dynamic_shapes = { + "tokens": None, + "target_hidden": {1: ctx_dim}, + "position_ids": {1: ctx_dim + block_size}, + } + + import torch.fx.experimental._config as fx_config + with fx_config.patch(backed_size_oblivious=True): + exported = torch.export.export( + model, (tokens, target_hidden, position_ids), dynamic_shapes=dynamic_shapes + ) from executorch.exir import to_edge_transform_and_lower from executorch.backends.mlx.partitioner import MLXPartitioner @@ -61,6 +88,7 @@ def main(): with open(args.output, "wb") as f: f.write(et_program.buffer) print(f"Saved draft model to: {args.output}") + print(f"Dynamic ctx_len supported: 1 to {args.max_ctx_len}, block_size fixed at {block_size}.") if __name__ == "__main__": diff --git a/examples/models/qwen3/mlx_source_transformations.py b/examples/models/qwen3/mlx_source_transformations.py new file mode 100644 index 00000000000..ee39af7bcda --- /dev/null +++ b/examples/models/qwen3/mlx_source_transformations.py @@ -0,0 +1,110 @@ +"""Extracting Qwen3 hidden-state for DFlash. + +Same idea as examples/models/gemma4_31b/mlx_source_transformations.py -- +tap layers [2, N//2, N-3] and return them concatenated alongside logits. +Gemma 4 does this by patching its own hand-written forward(). Qwen3 goes +through the generic HF export path instead (export_llm_hf.py), which wraps +the model in transformers' TorchExportableModuleWithStaticCache before +torch.export. So we subclass that wrapper and add output_hidden_states +to its forward rather than patching Qwen3 itself. + +Base class signature/behavior confirmed via: + inspect.getsource(transformers.integrations.executorch.TorchExportableModuleWithStaticCache) +""" + +from typing import List, Optional, Sequence + +import torch +from transformers.integrations.executorch import TorchExportableModuleWithStaticCache + + +class TorchExportableModuleWithStaticCacheAndHidden(TorchExportableModuleWithStaticCache): + """forward() also returns tapped hidden states. + forward() -> (logits, hidden), where hidden is [B, T, len(layer_ids) * H]. + """ + + def __init__( + self, + model, + batch_size: Optional[int] = None, + max_cache_len: Optional[int] = None, + device: Optional[torch.device] = None, + layer_ids: Sequence[int] = (), + ): + super().__init__(model, batch_size=batch_size, max_cache_len=max_cache_len, device=device) + if not layer_ids: + raise ValueError("layer_ids must be non-empty") + self.layer_ids: List[int] = list(layer_ids) + + def forward( + self, + input_ids: Optional[torch.LongTensor] = None, + inputs_embeds: Optional[torch.Tensor] = None, + cache_position: Optional[torch.Tensor] = None, + ): + outs = self.model( + input_ids=input_ids, + inputs_embeds=inputs_embeds, + cache_position=cache_position, + attention_mask=None, + past_key_values=self.static_cache, + use_cache=True, + output_hidden_states=True, + ) + + # hidden_states[0] is the embedding output, hidden_states[i+1] is decoder layer i's output + captured = [outs.hidden_states[i + 1] for i in self.layer_ids] + hidden = torch.cat(captured, dim=-1) + + if hasattr(outs, "logits"): + return outs.logits, hidden + return outs.last_hidden_state, hidden + + +def default_dflash_layer_ids(num_layers: int) -> List[int]: + """[2, N//2, N-3] tap pattern, same as Gemma 4. For Qwen3-4B (36 layers): [2, 18, 33].""" + return [2, num_layers // 2, num_layers - 3] + +class StatelessQwen3WithHidden(torch.nn.Module): + """Cache-free counterpart to TorchExportableModuleWithStaticCacheAndHidden. + + DFlash's speculative-decode loop re-verifies overlapping/non-contiguous + token ranges every round (draft tokens get proposed, verified, some + rejected). TorchExportableModuleWithStaticCache's persistent internal + cache is built for strictly-sequential autoregressive decoding and + produces corrupted hidden states under this access pattern (confirmed at + the eager PyTorch level: two calls at non-contiguous cache_position values + on the same cached wrapper differ by ~8694 vs. ~0.0003 for a correctly + stateless model). This class sidesteps the whole problem: every forward() + call recomputes attention over exactly the tokens/positions given, with no + persistent state, matching the accumulate-full-context approach already + used by DFlashDraftModel in dflash_draft_model.py. + + forward() -> (logits, hidden), where hidden is [B, T, len(layer_ids) * H]. + """ + + def __init__(self, model, layer_ids: Sequence[int] = ()): + super().__init__() + if not layer_ids: + raise ValueError("layer_ids must be non-empty") + self.model = model + self.layer_ids: List[int] = list(layer_ids) + + def forward( + self, + input_ids: Optional[torch.LongTensor] = None, + cache_position: Optional[torch.Tensor] = None, + ): + outs = self.model( + input_ids=input_ids, + cache_position=cache_position, + attention_mask=None, + past_key_values=None, + use_cache=False, + output_hidden_states=True, + ) + captured = [outs.hidden_states[i + 1] for i in self.layer_ids] + hidden = torch.cat(captured, dim=-1) + if hasattr(outs, "logits"): + return outs.logits, hidden + return outs.last_hidden_state, hidden diff --git a/examples/models/qwen3/run_baseline.py b/examples/models/qwen3/run_baseline.py new file mode 100644 index 00000000000..84a82ef8d5d --- /dev/null +++ b/examples/models/qwen3/run_baseline.py @@ -0,0 +1,69 @@ +"""Plain autoregressive baseline using the SAME target .pte, tokenizer, and +chat-template settings as run_dflash.py -- for an apples-to-apples comparison, +not the old Phase 0 number (measured under different conditions/quant pass). +""" +import argparse +import time + +import torch +from transformers import AutoTokenizer +from executorch.runtime import Runtime, Verification + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--target-pte", default="qwen3_4b_dflash_target.pte") + p.add_argument("--tokenizer", default="Qwen/Qwen3-4B") + p.add_argument("--prompt", required=True) + p.add_argument("--max-new-tokens", type=int, default=128) + p.add_argument("--chat-template", action="store_true", default=True) + p.add_argument("--no-chat-template", dest="chat_template", action="store_false") + p.add_argument("--enable-thinking", action="store_true", default=False) + args = p.parse_args() + + tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, local_files_only=True) + eos_id = tokenizer.eos_token_id + + if args.chat_template: + messages = [{"role": "user", "content": args.prompt}] + chat_out = tokenizer.apply_chat_template( + messages, add_generation_prompt=True, + enable_thinking=args.enable_thinking, return_tensors="pt", + ) + prompt_ids = chat_out.input_ids if hasattr(chat_out, "input_ids") else chat_out + else: + prompt_ids = tokenizer(args.prompt, return_tensors="pt").input_ids + + rt = Runtime.get() + target = rt.load_program(args.target_pte, verification=Verification.Minimal).load_method("forward") + + prompt_len = prompt_ids.shape[1] + input_pos = torch.arange(prompt_len, dtype=torch.long) + + t0 = time.time() + logits, _hidden = target.execute([prompt_ids, input_pos]) + pos = prompt_len + token = int(logits[0, -1].argmax()) + generated = [token] + + while len(generated) < args.max_new_tokens: + tok_input = torch.tensor([[token]], dtype=torch.long) + pos_input = torch.tensor([pos], dtype=torch.long) + logits, _hidden = target.execute([tok_input, pos_input]) + token = int(logits[0, -1].argmax()) + generated.append(token) + pos += 1 + if token == eos_id: + break + + dt = time.time() - t0 + text = tokenizer.decode(generated) + n = len(generated) + print(f"Prompt: {args.prompt}") + print(f"Generated ({n} tokens): {text}") + print(f"\n--- baseline stats ---") + print(f"time: {dt:.2f}s tokens/s: {n / dt:.2f}") + + +if __name__ == "__main__": + main() diff --git a/examples/models/qwen3/run_dflash.py b/examples/models/qwen3/run_dflash.py new file mode 100644 index 00000000000..fd42c86aac5 --- /dev/null +++ b/examples/models/qwen3/run_dflash.py @@ -0,0 +1,162 @@ +"""DFlash speculative decoding driver for the ExecuTorch MLX backend (Python). + +Same four Phase 3 pieces as qwen3_dflash_engine.cpp, driven through the ET +Python runtime so it runs on machines that can't build the C++ core: + 1. draft block construction [last_token, mask, mask, ...] + 2. target verification run target on [last_token] + draft_tokens + 3. acceptance keep prefix up to first mismatch, + bonus token + 4. position-based rollback pos += accepted + 1 + +V1 scope (per design doc): greedy, single batch, chain drafting, standard attn. +""" + +import argparse +import time +from pathlib import Path + +import torch +from huggingface_hub import snapshot_download +from transformers import AutoTokenizer +from executorch.runtime import Runtime, Verification + +from executorch.backends.mlx.examples.llm.dflash_draft_model import load_dflash_config + + +def first_mismatch(draft_ids, target_ids): + """Number of leading draft tokens the target agrees with (greedy accept).""" + for i in range(len(draft_ids)): + if draft_ids[i] != target_ids[i]: + return i + return len(draft_ids) + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--target-pte", default="qwen3_4b_dflash_target.pte") + p.add_argument("--draft-pte", default="qwen3_4b_dflash_draft.pte") + p.add_argument("--draft-model", default="z-lab/Qwen3-4B-DFlash-b16") + p.add_argument("--tokenizer", default="Qwen/Qwen3-4B") + p.add_argument("--prompt", default="The capital of France is") + p.add_argument("--max-new-tokens", type=int, default=64) + p.add_argument("--chat-template", action="store_true", default=True, + help="Apply Qwen3's chat template (paper's eval setup). Default on.") + p.add_argument("--no-chat-template", dest="chat_template", action="store_false") + p.add_argument("--enable-thinking", action="store_true", default=False, + help="Qwen3 thinking mode. Paper's Table 1 uses thinking mode DISABLED.") + p.add_argument("--block-size", type=int, default=None, + help="Override the draft checkpoint config's block_size -- needed when " + "--draft-pte was exported with a different block_size than the " + "z-lab checkpoint's native config (e.g. our block_size=8 test export).") + args = p.parse_args() + + config = load_dflash_config(Path(snapshot_download( + args.draft_model, allow_patterns=["*.json"], local_files_only=True))) + mask_id = config.mask_token_id + block_size = args.block_size if args.block_size is not None else config.block_size + layer_ids = config.target_layer_ids + + tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, local_files_only=True) + eos_id = tokenizer.eos_token_id + + # Paper's Table 1 evaluates with the chat template applied and thinking mode + # disabled ("Q3-4B ... thinking mode disabled" -- Section 5.1), not raw + # completion text. The draft was also trained on prompt+response pairs + # (Section 5 "Datasets"; Figure 4 shows clean prompt p / response r), so + # feeding it untemplated text is a real distribution mismatch, not a + # cosmetic difference. + + rt = Runtime.get() + target = rt.load_program(args.target_pte, verification=Verification.Minimal).load_method("forward") + draft = rt.load_program(args.draft_pte, verification=Verification.Minimal).load_method("forward") + + if args.chat_template: + messages = [{"role": "user", "content": args.prompt}] + chat_out = tokenizer.apply_chat_template( + messages, + add_generation_prompt=True, + enable_thinking=args.enable_thinking, + return_tensors="pt", + ) + # Some transformers versions return a BatchEncoding (dict-like) here + # instead of a raw tensor; normalize either way. + prompt_ids = chat_out.input_ids if hasattr(chat_out, "input_ids") else chat_out + else: + prompt_ids = tokenizer(args.prompt, return_tensors="pt").input_ids + prompt_len = prompt_ids.shape[1] + + # --- Prefill: target over the full prompt -> (logits, hidden) --- + input_pos = torch.arange(prompt_len, dtype=torch.long) + logits, hidden = target.execute([prompt_ids, input_pos]) + hidden = hidden.float() + pos = prompt_len + last_token = int(logits[0, -1].argmax()) + + generated = [last_token] + rounds = 0 + accepted_total = 0 + t0 = time.time() + + while len(generated) < args.max_new_tokens: + rounds += 1 + + # 1. Draft block: [last_token, mask, mask, ...] + draft_input = torch.cat( + [torch.tensor([[last_token]], dtype=torch.long), + torch.full((1, block_size - 1), mask_id, dtype=torch.long)], dim=1) + draft_pos = torch.arange(hidden.shape[1] + block_size, dtype=torch.long).unsqueeze(0) + _t0 = time.time() + (draft_logits,) = draft.execute([draft_input, hidden, draft_pos]) + _draft_time = time.time() - _t0 + draft_ids = draft_logits[0].argmax(-1).tolist() # block_size - 1 tokens + + # 2. Verify: target on [last_token] + draft_ids + verify_input = torch.cat( + [torch.tensor([[last_token]], dtype=torch.long), + torch.tensor([draft_ids], dtype=torch.long)], dim=1) + verify_pos = torch.arange(pos, pos + verify_input.shape[1], dtype=torch.long) + _t1 = time.time() + target_logits, new_hidden = target.execute([verify_input, verify_pos]) + _verify_time = time.time() - _t1 + if rounds <= 10: + print(f" timing: draft={_draft_time*1000:.1f}ms verify={_verify_time*1000:.1f}ms ctx_len={hidden.shape[1]}") + target_ids = target_logits[0].argmax(-1).tolist() # block_size tokens + + # 3. Accept: matching prefix + the target's bonus token at the mismatch + accepted = first_mismatch(draft_ids, target_ids) + if rounds <= 5: + print(f"round {rounds}: pos={pos} hidden_ctx={hidden.shape[1]} " + f"draft_ids[:5]={draft_ids[:5]} target_ids[:5]={target_ids[:5]} accepted={accepted}") + new_tokens = draft_ids[:accepted] + [target_ids[accepted]] + accepted_total += accepted + + # Trim at EOS if it appears in the accepted run + if eos_id in new_tokens: + new_tokens = new_tokens[:new_tokens.index(eos_id) + 1] + + generated.extend(new_tokens) + + # 4. Position-based rollback + pos += len(new_tokens) + last_token = new_tokens[-1] + # Accumulate: append this round's newly-generated tokens' hidden onto + # the running context, don't replace it. The target context feature + # must span the whole sequence generated so far (paper Figure 2), not + # just the latest round. + hidden = torch.cat([hidden, new_hidden[:, :len(new_tokens), :].float()], dim=1) + + if eos_id in new_tokens: + break + + dt = time.time() - t0 + text = tokenizer.decode(generated) + n = len(generated) + print(f"\nPrompt: {args.prompt}") + print(f"Generated ({n} tokens): {text}") + print(f"\n--- stats ---") + print(f"rounds: {rounds}") + print(f"avg accepted/round (tau proxy): {accepted_total / rounds:.2f}") + print(f"time: {dt:.2f}s tokens/s: {n / dt:.2f}") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_draft.py b/examples/models/qwen3/tests/test_dflash_draft.py new file mode 100644 index 00000000000..77dbd3d5fa6 --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_draft.py @@ -0,0 +1,28 @@ +"""Phase 2 / Part 8 verification: the exported draft .pte loads, executes, and +supports dynamic ctx_len (since DFlash accumulated target-hidden context grows +every speculative round -- see run_dflash.py). Covers export correctness and +weight correctness implicitly: a shape or weight mismatch would have failed at +export time (export_dflash_draft.py asserts checkpoint keys match before +saving), so a successful load and execute here is sufficient proof. +""" +import sys +import torch +from executorch.runtime import Runtime, Verification + +pte_path = sys.argv[1] if len(sys.argv) > 1 else "qwen3_4b_dflash_draft.pte" +et_runtime = Runtime.get() +method = et_runtime.load_program(pte_path, verification=Verification.Minimal).load_method("forward") + +block_size, hidden_size, vocab_size = 16, 12800, 151936 + +for ctx_len in (8, 20, 1): + tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) + target_hidden = torch.randn(1, ctx_len, hidden_size) + position_ids = torch.arange(ctx_len + block_size).unsqueeze(0).long() + + (draft_logits,) = method.execute([tokens, target_hidden, position_ids]) + assert draft_logits.shape == (1, block_size - 1, vocab_size), (ctx_len, draft_logits.shape) + assert not torch.isnan(draft_logits).any() and not torch.isinf(draft_logits).any() + print(f"ctx_len={ctx_len}: OK {tuple(draft_logits.shape)}") + +print("PASS: draft .pte loads, executes, and supports dynamic ctx_len") diff --git a/examples/models/qwen3/tests/test_dflash_draft_dynamic.py b/examples/models/qwen3/tests/test_dflash_draft_dynamic.py new file mode 100644 index 00000000000..e1c5cf02b8a --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_draft_dynamic.py @@ -0,0 +1,18 @@ +import torch +from executorch.runtime import Runtime, Verification + +rt = Runtime.get() +method = rt.load_program("qwen3_4b_dflash_draft.pte", + verification=Verification.Minimal).load_method("forward") + +block_size, hidden_size, vocab_size = 16, 12800, 151936 +for ctx_len in (8, 20, 1): + tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) + target_hidden = torch.randn(1, ctx_len, hidden_size) + position_ids = torch.arange(ctx_len + block_size).unsqueeze(0).long() + (draft_logits,) = method.execute([tokens, target_hidden, position_ids]) + assert draft_logits.shape == (1, block_size - 1, vocab_size), (ctx_len, draft_logits.shape) + assert not torch.isnan(draft_logits).any() and not torch.isinf(draft_logits).any() + print(f"ctx_len={ctx_len}: OK {tuple(draft_logits.shape)}") + +print("PASS: dynamic ctx_len verified at multiple lengths") diff --git a/examples/models/qwen3/tests/test_dflash_draft_eager.py b/examples/models/qwen3/tests/test_dflash_draft_eager.py new file mode 100644 index 00000000000..ee3f0c06f8d --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_draft_eager.py @@ -0,0 +1,34 @@ +import torch +from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashConfig, DFlashDraftModel + +config = DFlashConfig( + hidden_size=2560, + num_hidden_layers=5, + num_attention_heads=32, + num_key_value_heads=8, + head_dim=128, + intermediate_size=9728, + vocab_size=151936, + rms_norm_eps=1e-6, + rope_theta=1_000_000.0, + max_position_embeddings=40960, + target_layer_ids=(1, 9, 17, 25, 33), + block_size=16, + mask_token_id=151669, + layer_types=("full_attention",) * 5, +) + +model = DFlashDraftModel(config) +model.eval() + +block_size, ctx_len = 16, 12 +tokens = torch.randint(0, config.vocab_size, (1, block_size), dtype=torch.long) +target_hidden = torch.randn(1, ctx_len, len(config.target_layer_ids) * config.hidden_size) +position_ids = torch.arange(ctx_len + block_size).unsqueeze(0) + +with torch.no_grad(): + logits = model(tokens, target_hidden, position_ids) + +assert logits.shape == (1, block_size - 1, config.vocab_size), logits.shape +assert not torch.isnan(logits).any() and not torch.isinf(logits).any() +print(f"OK: draft logits {tuple(logits.shape)}, no NaN/Inf") \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_draft_forward.py b/examples/models/qwen3/tests/test_dflash_draft_forward.py new file mode 100644 index 00000000000..bfcc8895374 --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_draft_forward.py @@ -0,0 +1,46 @@ +from pathlib import Path + +import torch +from huggingface_hub import snapshot_download +from safetensors.torch import load_file +from transformers import AutoModelForCausalLM, AutoTokenizer + +from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashDraftModel, load_dflash_config + +path = Path(snapshot_download("z-lab/Qwen3-4B-DFlash-b16", allow_patterns=["*.safetensors", "*.json"])) +config = load_dflash_config(path) + +model = DFlashDraftModel(config) +weights = {} +for f in path.glob("*.safetensors"): + weights.update(load_file(str(f))) +model.load_state_dict(weights, strict=False) + +tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-4B") +target = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-4B", dtype="auto") +model.embed_tokens.weight.data.copy_(target.model.embed_tokens.weight) +model.lm_head.weight.data.copy_(target.lm_head.weight) +model.eval() +target.eval() + +prompt = "The capital of France is" +input_ids = tokenizer(prompt, return_tensors="pt").input_ids + +with torch.no_grad(): + out = target(input_ids, output_hidden_states=True) + tapped = [out.hidden_states[i + 1] for i in config.target_layer_ids] + target_hidden = torch.cat(tapped, dim=-1).float() + + last_token = input_ids[:, -1:] + block_size = 8 + draft_tokens = torch.cat( + [last_token, torch.full((1, block_size - 1), config.mask_token_id, dtype=torch.long)], dim=1 + ) + position_ids = torch.arange(target_hidden.shape[1] + block_size).unsqueeze(0) + + draft_logits = model(draft_tokens, target_hidden, position_ids) + predicted = draft_logits.argmax(-1) + +print("Prompt:", prompt) +print("Predicted continuation:", tokenizer.decode(predicted[0])) +print("Shape:", draft_logits.shape) diff --git a/examples/models/qwen3/tests/test_dflash_draft_load_weights.py b/examples/models/qwen3/tests/test_dflash_draft_load_weights.py new file mode 100644 index 00000000000..75af1d702db --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_draft_load_weights.py @@ -0,0 +1,50 @@ +import json +from pathlib import Path + +import torch +from huggingface_hub import snapshot_download +from safetensors.torch import load_file +from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashConfig, DFlashDraftModel + +path = Path(snapshot_download("z-lab/Qwen3-4B-DFlash-b16", allow_patterns=["*.safetensors", "*.json"])) +cfg = json.loads((path / "config.json").read_text()) +dcfg = cfg["dflash_config"] + +config = DFlashConfig( + hidden_size=cfg["hidden_size"], + num_hidden_layers=cfg["num_hidden_layers"], + num_attention_heads=cfg["num_attention_heads"], + num_key_value_heads=cfg["num_key_value_heads"], + head_dim=cfg["head_dim"], + intermediate_size=cfg["intermediate_size"], + vocab_size=cfg["vocab_size"], + rms_norm_eps=cfg["rms_norm_eps"], + rope_theta=cfg["rope_theta"], + max_position_embeddings=cfg["max_position_embeddings"], + target_layer_ids=tuple(dcfg["target_layer_ids"]), + block_size=cfg["block_size"], + mask_token_id=dcfg["mask_token_id"], + layer_types=tuple(cfg.get("layer_types") or ["full_attention"] * cfg["num_hidden_layers"]), + sliding_window=cfg.get("sliding_window"), + final_logit_softcapping=cfg.get("final_logit_softcapping"), +) + +model = DFlashDraftModel(config) + +draft_weights = {} +for f in path.glob("*.safetensors"): + draft_weights.update(load_file(str(f))) + +print(f"Checkpoint has {len(draft_weights)} tensors") +print("First 10 checkpoint keys:", list(draft_weights.keys())[:10]) +print("First 10 model keys: ", list(model.state_dict().keys())[:10]) + +missing, unexpected = model.load_state_dict(draft_weights, strict=False) +still_missing = [k for k in missing if not k.startswith(("embed_tokens.", "lm_head."))] + +print(f"\nMissing (excl. embed/lm_head, expected empty): {still_missing}") +print(f"Unexpected (expected empty): {unexpected}") + +assert not still_missing, "Architecture mismatch — key names don't match the real checkpoint" +assert not unexpected, "Checkpoint has tensors our model doesn't define — architecture mismatch" +print("\nOK: state_dict loaded cleanly, structure matches the real checkpoint") \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_export.py b/examples/models/qwen3/tests/test_dflash_export.py new file mode 100644 index 00000000000..17c6d7b0013 --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_export.py @@ -0,0 +1,32 @@ +"""Sanity check: qwen3_4b_dflash_target.pte returns (logits, hidden) with the +expected shapes at runtime. + +Run after exporting with --dflash-layers, e.g.: + python3 export_llm_hf.py --model-id Qwen/Qwen3-4B --dflash-layers 2,18,33 ... + python3 test_dflash_export.py qwen3_4b_dflash_target.pte +""" + +import sys +import torch +from executorch.runtime import Runtime, Verification + +DFLASH_LAYERS = [1, 9, 17, 25, 33] +HIDDEN_SIZE = 2560 +EXPECTED_HIDDEN_DIM = len(DFLASH_LAYERS) * HIDDEN_SIZE # 12800 +VOCAB_SIZE = 151936 + +pte_path = sys.argv[1] +et_runtime = Runtime.get() +program = et_runtime.load_program(pte_path, verification=Verification.Minimal) +method = program.load_method("forward") + +tokens = torch.tensor([[1, 2, 3]], dtype=torch.long) +input_pos = torch.tensor([0], dtype=torch.long) +logits, hidden = method.execute([tokens, input_pos]) + +assert logits.shape == (1, 3, VOCAB_SIZE), logits.shape +assert hidden.shape == (1, 3, EXPECTED_HIDDEN_DIM), hidden.shape +assert not torch.isnan(logits).any() and not torch.isinf(logits).any() +assert not torch.isnan(hidden).any() and not torch.isinf(hidden).any() + +print(f"OK: logits {tuple(logits.shape)}, hidden {tuple(hidden.shape)}") \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_lossless.py b/examples/models/qwen3/tests/test_dflash_lossless.py new file mode 100644 index 00000000000..86c81de5ece --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_lossless.py @@ -0,0 +1,32 @@ +"""Verification (doc Part 8, criterion 1): greedy DFlash must be LOSSLESS -- +identical tokens to greedy autoregressive baseline, since the target verifies +every accepted token. Any divergence means the speculative loop is broken.""" +import subprocess, sys, re + +PROMPT = "Write a Python function that takes a list of integers and returns the second largest number in the list." +N = 96 + +def run(script, extra): + out = subprocess.run( + [sys.executable, f"examples/models/qwen3/{script}", + "--prompt", PROMPT, "--max-new-tokens", str(N)] + extra, + capture_output=True, text=True, cwd=".", + ).stdout + m = re.search(r"Generated \([^)]*\): (.*?)\n\n", out, re.DOTALL) + return m.group(1) if m else out + +baseline = run("run_baseline.py", []) +dflash = run("run_dflash.py", []) + +print("=== BASELINE ===\n", baseline[:400]) +print("\n=== DFLASH ===\n", dflash[:400]) +print("\n=== RESULT ===") +if baseline.strip() == dflash.strip(): + print("PASS: DFlash output is token-for-token identical to baseline (LOSSLESS)") +else: + # find first divergence + for i, (a, b) in enumerate(zip(baseline, dflash)): + if a != b: + print(f"DIVERGE at char {i}: baseline={baseline[i:i+30]!r} dflash={dflash[i:i+30]!r}") + break + print("FAIL: outputs differ -- speculative loop is not lossless") diff --git a/examples/models/qwen3/tests/test_dflash_target.py b/examples/models/qwen3/tests/test_dflash_target.py new file mode 100644 index 00000000000..17c6d7b0013 --- /dev/null +++ b/examples/models/qwen3/tests/test_dflash_target.py @@ -0,0 +1,32 @@ +"""Sanity check: qwen3_4b_dflash_target.pte returns (logits, hidden) with the +expected shapes at runtime. + +Run after exporting with --dflash-layers, e.g.: + python3 export_llm_hf.py --model-id Qwen/Qwen3-4B --dflash-layers 2,18,33 ... + python3 test_dflash_export.py qwen3_4b_dflash_target.pte +""" + +import sys +import torch +from executorch.runtime import Runtime, Verification + +DFLASH_LAYERS = [1, 9, 17, 25, 33] +HIDDEN_SIZE = 2560 +EXPECTED_HIDDEN_DIM = len(DFLASH_LAYERS) * HIDDEN_SIZE # 12800 +VOCAB_SIZE = 151936 + +pte_path = sys.argv[1] +et_runtime = Runtime.get() +program = et_runtime.load_program(pte_path, verification=Verification.Minimal) +method = program.load_method("forward") + +tokens = torch.tensor([[1, 2, 3]], dtype=torch.long) +input_pos = torch.tensor([0], dtype=torch.long) +logits, hidden = method.execute([tokens, input_pos]) + +assert logits.shape == (1, 3, VOCAB_SIZE), logits.shape +assert hidden.shape == (1, 3, EXPECTED_HIDDEN_DIM), hidden.shape +assert not torch.isnan(logits).any() and not torch.isinf(logits).any() +assert not torch.isnan(hidden).any() and not torch.isinf(hidden).any() + +print(f"OK: logits {tuple(logits.shape)}, hidden {tuple(hidden.shape)}") \ No newline at end of file From 7be25cd611d7d43ac10934a53f7b3f4d0897d5be Mon Sep 17 00:00:00 2001 From: cthotti Date: Fri, 10 Jul 2026 10:48:44 +0530 Subject: [PATCH 3/5] Add DFlash speculative decoding support for Qwen3 --- .gitignore | 7 + .../mlx/examples/llm/dflash_draft_model.py | 118 ++++++++---- backends/mlx/examples/llm/export_llm_hf.py | 27 +-- examples/models/qwen3/export_dflash_draft.py | 38 ++-- .../qwen3/mlx_source_transformations.py | 64 ++----- examples/models/qwen3/run_baseline.py | 19 +- examples/models/qwen3/run_dflash.py | 177 ++++++++++++------ .../models/qwen3/tests/test_dflash_draft.py | 21 ++- .../qwen3/tests/test_dflash_draft_dynamic.py | 18 -- .../qwen3/tests/test_dflash_draft_eager.py | 34 ---- .../qwen3/tests/test_dflash_draft_forward.py | 46 ----- .../tests/test_dflash_draft_load_weights.py | 50 ----- .../qwen3/tests/test_dflash_draft_pte.py | 19 -- .../models/qwen3/tests/test_dflash_export.py | 32 ---- .../qwen3/tests/test_dflash_lossless.py | 39 ++-- .../models/qwen3/tests/test_dflash_target.py | 11 +- 16 files changed, 308 insertions(+), 412 deletions(-) delete mode 100644 examples/models/qwen3/tests/test_dflash_draft_dynamic.py delete mode 100644 examples/models/qwen3/tests/test_dflash_draft_eager.py delete mode 100644 examples/models/qwen3/tests/test_dflash_draft_forward.py delete mode 100644 examples/models/qwen3/tests/test_dflash_draft_load_weights.py delete mode 100644 examples/models/qwen3/tests/test_dflash_draft_pte.py delete mode 100644 examples/models/qwen3/tests/test_dflash_export.py diff --git a/.gitignore b/.gitignore index 10b048c17b5..69c6bcdc115 100644 --- a/.gitignore +++ b/.gitignore @@ -86,3 +86,10 @@ zephyr_dev_root.backup.*/ # Agents .claude/*.local.* extension/pybindings/mlx.metallib + +# DFlash C++ engine (Qwen3) -- incomplete, broken build, follow-up PR +examples/models/qwen3/CMakeLists.txt +examples/models/qwen3/CMakePresets.json +examples/models/qwen3/qwen3_dflash_engine.h +examples/models/qwen3/qwen3_dflash_engine.cpp +examples/models/qwen3/main.cpp diff --git a/backends/mlx/examples/llm/dflash_draft_model.py b/backends/mlx/examples/llm/dflash_draft_model.py index 913f89e2d57..225d330b9b9 100644 --- a/backends/mlx/examples/llm/dflash_draft_model.py +++ b/backends/mlx/examples/llm/dflash_draft_model.py @@ -1,26 +1,25 @@ -"""PyTorch DFlash draft model, structured for ExecuTorch export. - -Model-agnostic: the same code exports a valid draft .pte for any standard- -attention target (Qwen3, Gemma-4, Llama-3.1) by reading a DFlashConfig loaded -from the z-lab draft checkpoint. Per-model differences — RoPE base/scaling, -sliding-window layers, final-logit softcap, embedding scale — are config values, -so they resolve at trace time to one model-specific graph. The universal -branches add no ops to the exported program. - -Two deliberate deviations from the reference forward, both for ET (design doc -"Option A", self-contained draft): - - embed_tokens / lm_head are owned here (filled from the target at export) - instead of referenced live off the target module. - - forward returns draft logits (norm -> lm_head -> [:, 1:]) rather than the - bare normed hidden state; the reference does lm_head + logits_start=1 in its - generate loop. - -References: - z-lab/dflash dflash/model_mlx.py — universal MLX reference (Qwen3 / Qwen3.5 / - Gemma-4); source of the sliding-window, softcap, single-rope, QK-norm design. - z-lab/dflash dflash/model.py — PyTorch reference; exact weight layout for - load_state_dict. - transformers RoPE init — config-driven inv_freq below. +"""PyTorch implementation of the DFlash draft model for ExecuTorch export. + +This model is the lightweight "draft" network used in DFlash speculative +decoding. Instead of generating one token at a time like the target LLM, it +predicts an entire block of future tokens in parallel. To do this, it takes: + - proposal tokens (the draft block, beginning with the last accepted token), + - hidden states extracted from the target model (Phase 1), and + - position IDs for the draft block. + +The target hidden states are first projected into the draft model's hidden +space, then every draft transformer layer attends to both the projected target +context and the proposal tokens. The result is a fast approximation of what +the target model is likely to generate next. + +The implementation is model-agnostic. Architectural differences such as RoPE, +sliding-window attention, embedding scale, and logit softcapping come from +DFlashConfig, allowing the same code to export draft models for Qwen3, Gemma, +Llama, and other standard-attention architectures. + +For ExecuTorch export, the draft model owns its own embedding and LM head +weights (copied from the target during export) and returns final draft logits +directly rather than intermediate hidden states. """ from dataclasses import dataclass, field @@ -49,15 +48,18 @@ class DFlashConfig: layer_types: Tuple[str, ...] = field(default_factory=tuple) sliding_window: Optional[int] = None final_logit_softcapping: Optional[float] = None - embed_scale: float = 1.0 # 1.0 for Qwen3/Llama; sqrt(hidden_size) for Gemma + # Some models scale token embeddings before entering transformer. + # Qwen3/Llama use 1.0, while Gemma scales by sqrt(hidden_size). + embed_scale: float = 1.0 def _rope_inv_freq(config: DFlashConfig) -> torch.Tensor: - # Covers the two rope types these drafts use: 'default' (Qwen3, Gemma-4) and - # 'linear' position-scaling. YaRN/longrope aren't handled — no current z-lab - # draft uses them. + # Build the RoPE frequencies expected by draft checkpoint. + # Different model families use different RoPE scaling strategies, so this is driven entirely from the checkpoint config rather than hardcoded. dim = config.head_dim - inv_freq = 1.0 / (config.rope_theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) + inv_freq = 1.0 / ( + config.rope_theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim) + ) scaling = config.rope_scaling or {} if scaling.get("rope_type", scaling.get("type")) == "linear": inv_freq = inv_freq / float(scaling["factor"]) @@ -72,7 +74,7 @@ def __init__(self, config: DFlashConfig): def forward(self, position_ids: torch.Tensor): inv = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) pos = position_ids[:, None, :].float() - freqs = (inv @ pos).transpose(1, 2) # [B, S, head_dim/2] + freqs = (inv @ pos).transpose(1, 2) # [B, S, head_dim/2] emb = torch.cat((freqs, freqs), dim=-1) # [B, S, head_dim] return emb.cos(), emb.sin() @@ -81,15 +83,19 @@ def rotate_half(x: torch.Tensor) -> torch.Tensor: x1, x2 = x.chunk(2, dim=-1) return torch.cat((-x2, x1), dim=-1) + def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor: + # Repeat KV heads when the model has fewer KV heads than attention heads. if n_rep == 1: return x b, h, s, d = x.shape x = x[:, :, None, :, :].expand(b, h, n_rep, s, d) return x.reshape(b, h * n_rep, s, d) + def apply_rotary_pos_emb(q, k, cos, sin): - # q rotates over its own (last q_len) positions; k rotates over its full length. + # The proposal block only contains the current draft tokens, so queries rotate over those positions only. + # Keys contain both the target context and proposal block, so they rotate over the full combined sequence. q_len = q.shape[-2] cq, sq = cos[:, None, -q_len:, :], sin[:, None, -q_len:, :] ck, sk = cos[:, None, :, :], sin[:, None, :, :] @@ -121,13 +127,15 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class DFlashAttention(nn.Module): + """Proposal tokens generate the queries. These represent the positions whose contents we are trying to predict.""" + def __init__(self, config: DFlashConfig, layer_idx: int): super().__init__() h, hd = config.hidden_size, config.head_dim self.n_heads = config.num_attention_heads self.n_kv = config.num_key_value_heads self.head_dim = hd - self.scaling = hd ** -0.5 + self.scaling = hd**-0.5 self.n_rep = self.n_heads // self.n_kv lt = config.layer_types self.is_sliding = bool(lt) and lt[layer_idx] == "sliding_attention" @@ -140,11 +148,18 @@ def __init__(self, config: DFlashConfig, layer_idx: int): self.k_norm = DFlashRMSNorm(hd, config.rms_norm_eps) def forward(self, x, x_ctx, cos, sin): + """Keys and values come frmo both the projected target context and the proposal block itself. This lets the draft attend to what the target model already understands while also allowing predictions within the proposal block to interact with one another.""" B, L, _ = x.shape S = x_ctx.shape[1] - q = self.q_norm(self.q_proj(x).view(B, L, self.n_heads, self.head_dim)).transpose(1, 2) - k = torch.cat([self.k_proj(x_ctx), self.k_proj(x)], dim=1).view(B, S + L, self.n_kv, self.head_dim) - v = torch.cat([self.v_proj(x_ctx), self.v_proj(x)], dim=1).view(B, S + L, self.n_kv, self.head_dim) + q = self.q_norm( + self.q_proj(x).view(B, L, self.n_heads, self.head_dim) + ).transpose(1, 2) + k = torch.cat([self.k_proj(x_ctx), self.k_proj(x)], dim=1).view( + B, S + L, self.n_kv, self.head_dim + ) + v = torch.cat([self.v_proj(x_ctx), self.v_proj(x)], dim=1).view( + B, S + L, self.n_kv, self.head_dim + ) k = self.k_norm(k).transpose(1, 2) v = v.transpose(1, 2) q, k = apply_rotary_pos_emb(q, k, cos, sin) @@ -152,11 +167,14 @@ def forward(self, x, x_ctx, cos, sin): k = repeat_kv(k, self.n_rep) v = repeat_kv(v, self.n_rep) mask = self._sliding_mask(L, S, q.device, q.dtype) if self.is_sliding else None + # Each proposal position attends over the combined context to build a richer representation before predicting its token. out = torch.nn.functional.scaled_dot_product_attention( - q, k, v, attn_mask=mask, is_causal=False, scale=self.scaling) + q, k, v, attn_mask=mask, is_causal=False, scale=self.scaling + ) return self.o_proj(out.transpose(1, 2).reshape(B, L, -1)) def _sliding_mask(self, L, S, device, dtype): + """Restrict attention to the configured sliding window for models that use sliding-window attention, like Gemma.""" total = S + L q_pos = torch.arange(S, total, device=device)[:, None] k_pos = torch.arange(total, device=device)[None, :] @@ -170,9 +188,14 @@ def __init__(self, config: DFlashConfig, layer_idx: int): self.self_attn = DFlashAttention(config, layer_idx) self.mlp = DFlashMLP(config.hidden_size, config.intermediate_size) self.input_layernorm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) - self.post_attention_layernorm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) + self.post_attention_layernorm = DFlashRMSNorm( + config.hidden_size, config.rms_norm_eps + ) def forward(self, x, x_ctx, cos, sin): + # Standard transformer decoder block: + # RMSNorm --> Attention --> Residual + # RMSNorm --> MLP --> Residual x = x + self.self_attn(self.input_layernorm(x), x_ctx, cos, sin) return x + self.mlp(self.post_attention_layernorm(x)) @@ -185,28 +208,39 @@ def __init__(self, config: DFlashConfig): self.fc = nn.Linear(concat_dim, config.hidden_size, bias=False) self.hidden_norm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) self.layers = nn.ModuleList( - [DFlashDecoderLayer(config, i) for i in range(config.num_hidden_layers)]) + [DFlashDecoderLayer(config, i) for i in range(config.num_hidden_layers)] + ) self.norm = DFlashRMSNorm(config.hidden_size, config.rms_norm_eps) self.rotary_emb = DFlashRotaryEmbedding(config) - # Option A: owned copies, filled from the target at export time. + # The draft owns its own embedding and LM head weights. + # During export these are copied from the target model, making the draft .pte self-contained. self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) def forward(self, tokens, target_hidden, position_ids): + # Embed the proposal block (last accepted token + masked future positions). h = self.embed_tokens(tokens) * self.config.embed_scale + # Translate the concatenated target hidden states into the draft model's hidden space. h_ctx = self.hidden_norm(self.fc(target_hidden)) + # Positional information for both the proposal block and target context. cos, sin = self.rotary_emb(position_ids) for layer in self.layers: h = layer(h, h_ctx, cos, sin) h = self.norm(h) - logits = self.lm_head(h[:, 1:, :]) # logits_start=1: drop the known first token + # Only return predictions for the future positions. + logits = self.lm_head(h[:, 1:, :]) + # logits_start=1: drop the known first token cap = self.config.final_logit_softcapping if cap is not None: logits = torch.tanh(logits / cap) * cap return logits + def load_dflash_config(checkpoint_dir) -> "DFlashConfig": - """Build a DFlashConfig from a z-lab DFlash checkpoint's config.json.""" + """Load the architecture needed to reconstruct a DFlash draft model. + + The checkpoint config describes both underlying transformer architecture (hidden size, attention heads, RoPE, etc.) and the DFlash-specific settings such as the tapped target layers and mask token. + """ import json from pathlib import Path @@ -227,7 +261,9 @@ def load_dflash_config(checkpoint_dir) -> "DFlashConfig": block_size=cfg["block_size"], mask_token_id=dcfg["mask_token_id"], rope_scaling=cfg.get("rope_scaling"), - layer_types=tuple(cfg.get("layer_types") or ["full_attention"] * cfg["num_hidden_layers"]), + layer_types=tuple( + cfg.get("layer_types") or ["full_attention"] * cfg["num_hidden_layers"] + ), sliding_window=cfg.get("sliding_window"), final_logit_softcapping=cfg.get("final_logit_softcapping"), ) diff --git a/backends/mlx/examples/llm/export_llm_hf.py b/backends/mlx/examples/llm/export_llm_hf.py index 45fd8bfad28..91b081696c8 100644 --- a/backends/mlx/examples/llm/export_llm_hf.py +++ b/backends/mlx/examples/llm/export_llm_hf.py @@ -221,21 +221,18 @@ def _export_with_custom_components( max_cache_len=effective_cache_len, ) elif dflash_layers is not None: - # Qwen3-specific for now - # Generalize the import if/when another model needs DFlash tapping. - # - # Stateless (no persistent KV cache), NOT TorchExportableModuleWithStaticCacheAndHidden: - # DFlash's speculative-decode loop re-verifies overlapping/non-contiguous token - # ranges every round, which corrupts a persistent StaticCache (confirmed at the - # eager PyTorch level -- see StatelessQwen3WithHidden's docstring). The driver - # always passes the full accumulated sequence instead of relying on a cache. + # Qwen3-specific for now. from executorch.examples.models.qwen3.mlx_source_transformations import ( - StatelessQwen3WithHidden, + TorchExportableModuleWithStaticCacheAndHidden, ) - logger.info(f"Creating stateless DFlash hidden-state-tapping wrapper, layers={dflash_layers}") - exportable = StatelessQwen3WithHidden( + logger.info( + f"Creating DFlash hidden-state-tapping wrapper, layers={dflash_layers}" + ) + exportable = TorchExportableModuleWithStaticCacheAndHidden( model=model, + batch_size=1, + max_cache_len=effective_cache_len, layer_ids=dflash_layers, ) else: @@ -246,7 +243,7 @@ def _export_with_custom_components( max_cache_len=effective_cache_len, ) - if use_custom_kv_cache and dflash_layers is None: + if use_custom_kv_cache: from executorch.backends.mlx.llm.source_transformation import ( replace_hf_cache_with_mlx, ) @@ -318,6 +315,9 @@ def _export_with_custom_components( transform_passes=get_default_passes(), partitioner=[MLXPartitioner()], compile_config=edge_config, + # Required by the C++ LLMEngine metadata contract (get_llm_metadata in + # llm_runner_helper.cpp) -- this export path (used for --dflash-layers) + constant_methods={"get_max_seq_len": max_seq_len}, ) logger.info("Exporting to ExecuTorch...") @@ -459,10 +459,11 @@ def main(): "--dflash-layers", type=str, default=None, - help="Comma-separated layer indices to tap for DFlash hidden-state output, e.g. '2,18,33'", + help="Comma-separated transformer layer indices whose hidden states are concatenated and returned alongside logits for DFlash. E.g. '1,9,17,25,33'", ) args = parser.parse_args() + # Convert "1,9,17,25,33" -> [1, 9, 17, 25, 33] dflash_layers = ( [int(x) for x in args.dflash_layers.split(",")] if args.dflash_layers else None ) diff --git a/examples/models/qwen3/export_dflash_draft.py b/examples/models/qwen3/export_dflash_draft.py index 75b72bd49cb..8b5fdc73291 100644 --- a/examples/models/qwen3/export_dflash_draft.py +++ b/examples/models/qwen3/export_dflash_draft.py @@ -1,14 +1,23 @@ +"""Exports the DFlash draft model to an .pte program. + +This script loads the pretrained DFlash draft checkpoint, copies the shared embedding and output project weights from target model, applies same 4-bit quantization used by target, and exports the draft model for MLX inference. +The exported model is used alongside the target model during speculative decoding. +""" + import argparse from pathlib import Path import torch + +from executorch.backends.mlx.examples.llm.dflash_draft_model import ( + DFlashDraftModel, + load_dflash_config, +) from huggingface_hub import snapshot_download from safetensors.torch import load_file from torch.export import Dim from transformers import AutoModelForCausalLM -from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashDraftModel, load_dflash_config - def load_draft_model(draft_id: str, target_state_dict: dict) -> DFlashDraftModel: path = Path(snapshot_download(draft_id, allow_patterns=["*.safetensors", "*.json"])) @@ -21,11 +30,17 @@ def load_draft_model(draft_id: str, target_state_dict: dict) -> DFlashDraftModel missing, unexpected = model.load_state_dict(draft_weights, strict=False) assert not unexpected, f"Unexpected draft checkpoint keys: {unexpected}" - still_missing = [k for k in missing if not k.startswith(("embed_tokens.", "lm_head."))] + still_missing = [ + k for k in missing if not k.startswith(("embed_tokens.", "lm_head.")) + ] assert not still_missing, f"Missing draft checkpoint keys: {still_missing}" model.embed_tokens.weight.data.copy_(target_state_dict["model.embed_tokens.weight"]) - lm_head_key = "lm_head.weight" if "lm_head.weight" in target_state_dict else "model.embed_tokens.weight" + lm_head_key = ( + "lm_head.weight" + if "lm_head.weight" in target_state_dict + else "model.embed_tokens.weight" + ) model.lm_head.weight.data.copy_(target_state_dict[lm_head_key]) return model @@ -45,12 +60,10 @@ def main(): model.eval() del target - # Quantize to 4-bit to match the target export (--qlinear 4w --qembedding 4w, - # group_size=32). Without this the draft is full float32 (~3.5GB) with its - # shared embed/lm_head dominating memory; quantizing brings it to ~1GB and, - # critically, keeps the shared embed/lm_head at the SAME precision as the - # target so their logits stay consistent (acceptance depends on this). + # Quantize the draft model to match the target model. + # Keeping both models at the same precision reduces memory usage and helps keep their predictions consistent, which is important for achieving a high draft acceptance rate. from executorch.backends.mlx.llm.quantization import quantize_model_ + quantize_model_( model, qlinear_config="4w", @@ -74,13 +87,14 @@ def main(): } import torch.fx.experimental._config as fx_config + with fx_config.patch(backed_size_oblivious=True): exported = torch.export.export( model, (tokens, target_hidden, position_ids), dynamic_shapes=dynamic_shapes ) - from executorch.exir import to_edge_transform_and_lower from executorch.backends.mlx.partitioner import MLXPartitioner + from executorch.exir import to_edge_transform_and_lower edge = to_edge_transform_and_lower(exported, partitioner=[MLXPartitioner()]) et_program = edge.to_executorch() @@ -88,7 +102,9 @@ def main(): with open(args.output, "wb") as f: f.write(et_program.buffer) print(f"Saved draft model to: {args.output}") - print(f"Dynamic ctx_len supported: 1 to {args.max_ctx_len}, block_size fixed at {block_size}.") + print( + f"Dynamic ctx_len supported: 1 to {args.max_ctx_len}, block_size fixed at {block_size}." + ) if __name__ == "__main__": diff --git a/examples/models/qwen3/mlx_source_transformations.py b/examples/models/qwen3/mlx_source_transformations.py index ee39af7bcda..edfc41eff61 100644 --- a/examples/models/qwen3/mlx_source_transformations.py +++ b/examples/models/qwen3/mlx_source_transformations.py @@ -1,7 +1,9 @@ """Extracting Qwen3 hidden-state for DFlash. Same idea as examples/models/gemma4_31b/mlx_source_transformations.py -- -tap layers [2, N//2, N-3] and return them concatenated alongside logits. +extract layers and return them concatenated alongside logits. In Qwen3, the +layer ids from z-lab Qwen3 DFlash draft config is [1, 9, 17, 25, 33] + Gemma 4 does this by patching its own hand-written forward(). Qwen3 goes through the generic HF export path instead (export_llm_hf.py), which wraps the model in transformers' TorchExportableModuleWithStaticCache before @@ -18,10 +20,9 @@ from transformers.integrations.executorch import TorchExportableModuleWithStaticCache -class TorchExportableModuleWithStaticCacheAndHidden(TorchExportableModuleWithStaticCache): - """forward() also returns tapped hidden states. - forward() -> (logits, hidden), where hidden is [B, T, len(layer_ids) * H]. - """ +class TorchExportableModuleWithStaticCacheAndHidden( + TorchExportableModuleWithStaticCache +): def __init__( self, @@ -31,7 +32,9 @@ def __init__( device: Optional[torch.device] = None, layer_ids: Sequence[int] = (), ): - super().__init__(model, batch_size=batch_size, max_cache_len=max_cache_len, device=device) + super().__init__( + model, batch_size=batch_size, max_cache_len=max_cache_len, device=device + ) if not layer_ids: raise ValueError("layer_ids must be non-empty") self.layer_ids: List[int] = list(layer_ids) @@ -52,7 +55,6 @@ def forward( output_hidden_states=True, ) - # hidden_states[0] is the embedding output, hidden_states[i+1] is decoder layer i's output captured = [outs.hidden_states[i + 1] for i in self.layer_ids] hidden = torch.cat(captured, dim=-1) @@ -62,49 +64,7 @@ def forward( def default_dflash_layer_ids(num_layers: int) -> List[int]: - """[2, N//2, N-3] tap pattern, same as Gemma 4. For Qwen3-4B (36 layers): [2, 18, 33].""" - return [2, num_layers // 2, num_layers - 3] - -class StatelessQwen3WithHidden(torch.nn.Module): - """Cache-free counterpart to TorchExportableModuleWithStaticCacheAndHidden. - - DFlash's speculative-decode loop re-verifies overlapping/non-contiguous - token ranges every round (draft tokens get proposed, verified, some - rejected). TorchExportableModuleWithStaticCache's persistent internal - cache is built for strictly-sequential autoregressive decoding and - produces corrupted hidden states under this access pattern (confirmed at - the eager PyTorch level: two calls at non-contiguous cache_position values - on the same cached wrapper differ by ~8694 vs. ~0.0003 for a correctly - stateless model). This class sidesteps the whole problem: every forward() - call recomputes attention over exactly the tokens/positions given, with no - persistent state, matching the accumulate-full-context approach already - used by DFlashDraftModel in dflash_draft_model.py. - - forward() -> (logits, hidden), where hidden is [B, T, len(layer_ids) * H]. + """This is simply the default dflash layer selection for Qwen3. + [2, N//2, N-3] tap pattern, same as Gemma 4. For Qwen3-4B (36 layers): [2, 18, 33]. """ - - def __init__(self, model, layer_ids: Sequence[int] = ()): - super().__init__() - if not layer_ids: - raise ValueError("layer_ids must be non-empty") - self.model = model - self.layer_ids: List[int] = list(layer_ids) - - def forward( - self, - input_ids: Optional[torch.LongTensor] = None, - cache_position: Optional[torch.Tensor] = None, - ): - outs = self.model( - input_ids=input_ids, - cache_position=cache_position, - attention_mask=None, - past_key_values=None, - use_cache=False, - output_hidden_states=True, - ) - captured = [outs.hidden_states[i + 1] for i in self.layer_ids] - hidden = torch.cat(captured, dim=-1) - if hasattr(outs, "logits"): - return outs.logits, hidden - return outs.last_hidden_state, hidden + return [2, num_layers // 2, num_layers - 3] diff --git a/examples/models/qwen3/run_baseline.py b/examples/models/qwen3/run_baseline.py index 84a82ef8d5d..c37b12b3b6d 100644 --- a/examples/models/qwen3/run_baseline.py +++ b/examples/models/qwen3/run_baseline.py @@ -1,13 +1,12 @@ -"""Plain autoregressive baseline using the SAME target .pte, tokenizer, and -chat-template settings as run_dflash.py -- for an apples-to-apples comparison, -not the old Phase 0 number (measured under different conditions/quant pass). +"""Standard autoregressive decoding used as the baseline for the comparison. """ + import argparse import time import torch -from transformers import AutoTokenizer from executorch.runtime import Runtime, Verification +from transformers import AutoTokenizer def main(): @@ -27,15 +26,19 @@ def main(): if args.chat_template: messages = [{"role": "user", "content": args.prompt}] chat_out = tokenizer.apply_chat_template( - messages, add_generation_prompt=True, - enable_thinking=args.enable_thinking, return_tensors="pt", + messages, + add_generation_prompt=True, + enable_thinking=args.enable_thinking, + return_tensors="pt", ) prompt_ids = chat_out.input_ids if hasattr(chat_out, "input_ids") else chat_out else: prompt_ids = tokenizer(args.prompt, return_tensors="pt").input_ids rt = Runtime.get() - target = rt.load_program(args.target_pte, verification=Verification.Minimal).load_method("forward") + target = rt.load_program( + args.target_pte, verification=Verification.Minimal + ).load_method("forward") prompt_len = prompt_ids.shape[1] input_pos = torch.arange(prompt_len, dtype=torch.long) @@ -61,7 +64,7 @@ def main(): n = len(generated) print(f"Prompt: {args.prompt}") print(f"Generated ({n} tokens): {text}") - print(f"\n--- baseline stats ---") + print("\n--baseline stats--") print(f"time: {dt:.2f}s tokens/s: {n / dt:.2f}") diff --git a/examples/models/qwen3/run_dflash.py b/examples/models/qwen3/run_dflash.py index fd42c86aac5..849afcb874d 100644 --- a/examples/models/qwen3/run_dflash.py +++ b/examples/models/qwen3/run_dflash.py @@ -1,13 +1,18 @@ -"""DFlash speculative decoding driver for the ExecuTorch MLX backend (Python). +"""Python implementation of the DFlash speculative decoding loop for the ExecuTorch MLX backend. -Same four Phase 3 pieces as qwen3_dflash_engine.cpp, driven through the ET -Python runtime so it runs on machines that can't build the C++ core: - 1. draft block construction [last_token, mask, mask, ...] - 2. target verification run target on [last_token] + draft_tokens - 3. acceptance keep prefix up to first mismatch, + bonus token - 4. position-based rollback pos += accepted + 1 +This file coordinates the interaction between the target model and the draft model during inference. Instead of asking the target model to generate one token at a time, DFlash first lets the lightweight draft model predict a block of future tokens, then asks the target model to verify those predictions in a single forward pass. Any matching draft tokens are accepted, while the first incorrect prediction is replaced with the target model's token. The process then repeats from the updated position. -V1 scope (per design doc): greedy, single batch, chain drafting, standard attn. +Each speculation round consists of four steps: + 1. Build a draft block: [last_token, , , ...] + 2. Run draft model to predict all masked tokens in parallel + 3. Verify those predictions with the target model, keeping matching prefix and replacing the first mismatch with target's prediction. + 4. advance the sequence position to the newly accepted prefix and repeat. + +V1 scope (per the issue discussion): + - Greedy decoding + - Single-batch inference + - Chain drafting + - Standard attention models """ import argparse @@ -15,15 +20,15 @@ from pathlib import Path import torch -from huggingface_hub import snapshot_download -from transformers import AutoTokenizer -from executorch.runtime import Runtime, Verification from executorch.backends.mlx.examples.llm.dflash_draft_model import load_dflash_config +from executorch.runtime import Runtime, Verification +from huggingface_hub import snapshot_download +from transformers import AutoTokenizer def first_mismatch(draft_ids, target_ids): - """Number of leading draft tokens the target agrees with (greedy accept).""" + """Returns the number of consecutive draft predictions that match the target.""" for i in range(len(draft_ids)): if draft_ids[i] != target_ids[i]: return i @@ -38,36 +43,57 @@ def main(): p.add_argument("--tokenizer", default="Qwen/Qwen3-4B") p.add_argument("--prompt", default="The capital of France is") p.add_argument("--max-new-tokens", type=int, default=64) - p.add_argument("--chat-template", action="store_true", default=True, - help="Apply Qwen3's chat template (paper's eval setup). Default on.") + p.add_argument( + "--chat-template", + action="store_true", + default=True, + help="Apply Qwen3's chat template (paper's eval setup). Default on.", + ) p.add_argument("--no-chat-template", dest="chat_template", action="store_false") - p.add_argument("--enable-thinking", action="store_true", default=False, - help="Qwen3 thinking mode. Paper's Table 1 uses thinking mode DISABLED.") - p.add_argument("--block-size", type=int, default=None, - help="Override the draft checkpoint config's block_size -- needed when " - "--draft-pte was exported with a different block_size than the " - "z-lab checkpoint's native config (e.g. our block_size=8 test export).") + p.add_argument( + "--enable-thinking", + action="store_true", + default=False, + help="Qwen3 thinking mode. Paper's Table 1 uses thinking mode DISABLED.", + ) + p.add_argument( + "--verbose", + action="store_true", + help="Print per-round timing/acceptance debug output.", + ) + p.add_argument( + "--block-size", + type=int, + default=None, + help="Override the draft checkpoint config's block_size -- needed when " + "--draft-pte was exported with a different block_size than the " + "z-lab checkpoint's native config (e.g. our block_size=8 test export).", + ) args = p.parse_args() - config = load_dflash_config(Path(snapshot_download( - args.draft_model, allow_patterns=["*.json"], local_files_only=True))) + config = load_dflash_config( + Path( + snapshot_download( + args.draft_model, allow_patterns=["*.json"], local_files_only=True + ) + ) + ) mask_id = config.mask_token_id block_size = args.block_size if args.block_size is not None else config.block_size - layer_ids = config.target_layer_ids tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, local_files_only=True) eos_id = tokenizer.eos_token_id - # Paper's Table 1 evaluates with the chat template applied and thinking mode - # disabled ("Q3-4B ... thinking mode disabled" -- Section 5.1), not raw - # completion text. The draft was also trained on prompt+response pairs - # (Section 5 "Datasets"; Figure 4 shows clean prompt p / response r), so - # feeding it untemplated text is a real distribution mismatch, not a - # cosmetic difference. + # The draft model was trained on Qwen3 chat-formatted prompt/response pairs,so applying the same chat template during inference keeps the input distribution consistent with training. + # Using raw completion text noticeably reduces acceptance rates. rt = Runtime.get() - target = rt.load_program(args.target_pte, verification=Verification.Minimal).load_method("forward") - draft = rt.load_program(args.draft_pte, verification=Verification.Minimal).load_method("forward") + target = rt.load_program( + args.target_pte, verification=Verification.Minimal + ).load_method("forward") + draft = rt.load_program( + args.draft_pte, verification=Verification.Minimal + ).load_method("forward") if args.chat_template: messages = [{"role": "user", "content": args.prompt}] @@ -77,14 +103,15 @@ def main(): enable_thinking=args.enable_thinking, return_tensors="pt", ) - # Some transformers versions return a BatchEncoding (dict-like) here - # instead of a raw tensor; normalize either way. + # Different Transformers versions return either a BatchEncoding or a tensor. + # Normalize both cases to a tensor. prompt_ids = chat_out.input_ids if hasattr(chat_out, "input_ids") else chat_out else: prompt_ids = tokenizer(args.prompt, return_tensors="pt").input_ids prompt_len = prompt_ids.shape[1] - # --- Prefill: target over the full prompt -> (logits, hidden) --- + # Run the target model over the prompt once to initialize generation. + # This produces the first next-token prediction and the hidden states that condition the draft model during speculative decoding. input_pos = torch.arange(prompt_len, dtype=torch.long) logits, hidden = target.execute([prompt_ids, input_pos]) hidden = hidden.float() @@ -94,55 +121,84 @@ def main(): generated = [last_token] rounds = 0 accepted_total = 0 + emitted_total = 0 t0 = time.time() while len(generated) < args.max_new_tokens: rounds += 1 + # The exported draft model expects a fixed input shape for its token block. + # Although the hidden-state and position inputs support dynamic lengths, the token input does not. + # Reducing the block size near the end of generation causes the runtime to reject the input. + # Supporting truly dynamic block sizes would require exporting the draft model with a dynamic token dimension. + bs = block_size - # 1. Draft block: [last_token, mask, mask, ...] + # 1. Build draft input block. draft_input = torch.cat( - [torch.tensor([[last_token]], dtype=torch.long), - torch.full((1, block_size - 1), mask_id, dtype=torch.long)], dim=1) - draft_pos = torch.arange(hidden.shape[1] + block_size, dtype=torch.long).unsqueeze(0) + [ + torch.tensor([[last_token]], dtype=torch.long), + torch.full((1, bs - 1), mask_id, dtype=torch.long), + ], + dim=1, + ) + draft_pos = torch.arange(hidden.shape[1] + bs, dtype=torch.long).unsqueeze(0) _t0 = time.time() (draft_logits,) = draft.execute([draft_input, hidden, draft_pos]) - _draft_time = time.time() - _t0 + _draft_exec_time = time.time() - _t0 + _t0b = time.time() draft_ids = draft_logits[0].argmax(-1).tolist() # block_size - 1 tokens + _draft_argmax_time = time.time() - _t0b - # 2. Verify: target on [last_token] + draft_ids + # 2. Verify the draft predictions. Target model predicts the next token after every position in the block in a single forward pass. verify_input = torch.cat( - [torch.tensor([[last_token]], dtype=torch.long), - torch.tensor([draft_ids], dtype=torch.long)], dim=1) + [ + torch.tensor([[last_token]], dtype=torch.long), + torch.tensor([draft_ids], dtype=torch.long), + ], + dim=1, + ) verify_pos = torch.arange(pos, pos + verify_input.shape[1], dtype=torch.long) _t1 = time.time() target_logits, new_hidden = target.execute([verify_input, verify_pos]) - _verify_time = time.time() - _t1 - if rounds <= 10: - print(f" timing: draft={_draft_time*1000:.1f}ms verify={_verify_time*1000:.1f}ms ctx_len={hidden.shape[1]}") + _target_exec_time = time.time() - _t1 + _t1b = time.time() target_ids = target_logits[0].argmax(-1).tolist() # block_size tokens + _target_argmax_time = time.time() - _t1b - # 3. Accept: matching prefix + the target's bonus token at the mismatch + # 3. Keep every drafting token that matches the target. At the first mismatch, stop accepting draft predictions and use the target model's token instead. + _t2 = time.time() accepted = first_mismatch(draft_ids, target_ids) - if rounds <= 5: - print(f"round {rounds}: pos={pos} hidden_ctx={hidden.shape[1]} " - f"draft_ids[:5]={draft_ids[:5]} target_ids[:5]={target_ids[:5]} accepted={accepted}") + _fm_time = time.time() - _t2 + if args.verbose and rounds <= 10: + print( + f" timing: draft_exec={_draft_exec_time*1000:.1f}ms draft_argmax={_draft_argmax_time*1000:.2f}ms " + f"target_exec={_target_exec_time*1000:.1f}ms target_argmax={_target_argmax_time*1000:.2f}ms " + f"first_mismatch={_fm_time*1000:.3f}ms ctx_len={hidden.shape[1]}" + ) + if args.verbose and rounds <= 5: + print( + f"round {rounds}: pos={pos} hidden_ctx={hidden.shape[1]} " + f"draft_ids[:5]={draft_ids[:5]} target_ids[:5]={target_ids[:5]} accepted={accepted}" + ) new_tokens = draft_ids[:accepted] + [target_ids[accepted]] accepted_total += accepted + emitted_total += len(new_tokens) - # Trim at EOS if it appears in the accepted run + # Stop generation once an EOS token becomes part of the accepted sequence. if eos_id in new_tokens: - new_tokens = new_tokens[:new_tokens.index(eos_id) + 1] + new_tokens = new_tokens[: new_tokens.index(eos_id) + 1] generated.extend(new_tokens) - # 4. Position-based rollback + # 4. Advance the accepted sequence. Rejected draft tokens are discarded, and the next round starts from the updated position. pos += len(new_tokens) last_token = new_tokens[-1] - # Accumulate: append this round's newly-generated tokens' hidden onto - # the running context, don't replace it. The target context feature - # must span the whole sequence generated so far (paper Figure 2), not - # just the latest round. - hidden = torch.cat([hidden, new_hidden[:, :len(new_tokens), :].float()], dim=1) + # Append the hidden states for the newly accepted tokens to the running target context. + # The draft model conditions on the hidden states of the entire sequence, so this context grows as generation progresses rather than being replaced each round. + _t3 = time.time() + hidden = torch.cat([hidden, new_hidden[:, : len(new_tokens), :].float()], dim=1) + _cat_time = time.time() - _t3 + if args.verbose and rounds <= 10: + print(f" timing: hidden_cat={_cat_time*1000:.2f}ms") if eos_id in new_tokens: break @@ -152,11 +208,12 @@ def main(): n = len(generated) print(f"\nPrompt: {args.prompt}") print(f"Generated ({n} tokens): {text}") - print(f"\n--- stats ---") + print("\n--stats--") print(f"rounds: {rounds}") - print(f"avg accepted/round (tau proxy): {accepted_total / rounds:.2f}") + print(f"avg accepted/round (draft-only): {accepted_total / rounds:.2f}") + print(f"avg emitted/round (tau, incl. bonus): {emitted_total / rounds:.2f}") print(f"time: {dt:.2f}s tokens/s: {n / dt:.2f}") if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/examples/models/qwen3/tests/test_dflash_draft.py b/examples/models/qwen3/tests/test_dflash_draft.py index 77dbd3d5fa6..f01ea0d254e 100644 --- a/examples/models/qwen3/tests/test_dflash_draft.py +++ b/examples/models/qwen3/tests/test_dflash_draft.py @@ -1,17 +1,17 @@ -"""Phase 2 / Part 8 verification: the exported draft .pte loads, executes, and -supports dynamic ctx_len (since DFlash accumulated target-hidden context grows -every speculative round -- see run_dflash.py). Covers export correctness and -weight correctness implicitly: a shape or weight mismatch would have failed at -export time (export_dflash_draft.py asserts checkpoint keys match before -saving), so a successful load and execute here is sufficient proof. """ +Verifies that the exported draft .pte loads, runs correctly, and supports dynamic context lengths as the accumulated target hidden-state context grows during speculative decoding. A successful export and execution also confirms that the checkpoint weights and exported model are compatible. +""" + import sys + import torch from executorch.runtime import Runtime, Verification pte_path = sys.argv[1] if len(sys.argv) > 1 else "qwen3_4b_dflash_draft.pte" et_runtime = Runtime.get() -method = et_runtime.load_program(pte_path, verification=Verification.Minimal).load_method("forward") +method = et_runtime.load_program( + pte_path, verification=Verification.Minimal +).load_method("forward") block_size, hidden_size, vocab_size = 16, 12800, 151936 @@ -21,8 +21,11 @@ position_ids = torch.arange(ctx_len + block_size).unsqueeze(0).long() (draft_logits,) = method.execute([tokens, target_hidden, position_ids]) - assert draft_logits.shape == (1, block_size - 1, vocab_size), (ctx_len, draft_logits.shape) + assert draft_logits.shape == (1, block_size - 1, vocab_size), ( + ctx_len, + draft_logits.shape, + ) assert not torch.isnan(draft_logits).any() and not torch.isinf(draft_logits).any() print(f"ctx_len={ctx_len}: OK {tuple(draft_logits.shape)}") -print("PASS: draft .pte loads, executes, and supports dynamic ctx_len") +print("PASS- draft .pte loads, executes, and supports dynamic ctx_len") diff --git a/examples/models/qwen3/tests/test_dflash_draft_dynamic.py b/examples/models/qwen3/tests/test_dflash_draft_dynamic.py deleted file mode 100644 index e1c5cf02b8a..00000000000 --- a/examples/models/qwen3/tests/test_dflash_draft_dynamic.py +++ /dev/null @@ -1,18 +0,0 @@ -import torch -from executorch.runtime import Runtime, Verification - -rt = Runtime.get() -method = rt.load_program("qwen3_4b_dflash_draft.pte", - verification=Verification.Minimal).load_method("forward") - -block_size, hidden_size, vocab_size = 16, 12800, 151936 -for ctx_len in (8, 20, 1): - tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) - target_hidden = torch.randn(1, ctx_len, hidden_size) - position_ids = torch.arange(ctx_len + block_size).unsqueeze(0).long() - (draft_logits,) = method.execute([tokens, target_hidden, position_ids]) - assert draft_logits.shape == (1, block_size - 1, vocab_size), (ctx_len, draft_logits.shape) - assert not torch.isnan(draft_logits).any() and not torch.isinf(draft_logits).any() - print(f"ctx_len={ctx_len}: OK {tuple(draft_logits.shape)}") - -print("PASS: dynamic ctx_len verified at multiple lengths") diff --git a/examples/models/qwen3/tests/test_dflash_draft_eager.py b/examples/models/qwen3/tests/test_dflash_draft_eager.py deleted file mode 100644 index ee3f0c06f8d..00000000000 --- a/examples/models/qwen3/tests/test_dflash_draft_eager.py +++ /dev/null @@ -1,34 +0,0 @@ -import torch -from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashConfig, DFlashDraftModel - -config = DFlashConfig( - hidden_size=2560, - num_hidden_layers=5, - num_attention_heads=32, - num_key_value_heads=8, - head_dim=128, - intermediate_size=9728, - vocab_size=151936, - rms_norm_eps=1e-6, - rope_theta=1_000_000.0, - max_position_embeddings=40960, - target_layer_ids=(1, 9, 17, 25, 33), - block_size=16, - mask_token_id=151669, - layer_types=("full_attention",) * 5, -) - -model = DFlashDraftModel(config) -model.eval() - -block_size, ctx_len = 16, 12 -tokens = torch.randint(0, config.vocab_size, (1, block_size), dtype=torch.long) -target_hidden = torch.randn(1, ctx_len, len(config.target_layer_ids) * config.hidden_size) -position_ids = torch.arange(ctx_len + block_size).unsqueeze(0) - -with torch.no_grad(): - logits = model(tokens, target_hidden, position_ids) - -assert logits.shape == (1, block_size - 1, config.vocab_size), logits.shape -assert not torch.isnan(logits).any() and not torch.isinf(logits).any() -print(f"OK: draft logits {tuple(logits.shape)}, no NaN/Inf") \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_draft_forward.py b/examples/models/qwen3/tests/test_dflash_draft_forward.py deleted file mode 100644 index bfcc8895374..00000000000 --- a/examples/models/qwen3/tests/test_dflash_draft_forward.py +++ /dev/null @@ -1,46 +0,0 @@ -from pathlib import Path - -import torch -from huggingface_hub import snapshot_download -from safetensors.torch import load_file -from transformers import AutoModelForCausalLM, AutoTokenizer - -from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashDraftModel, load_dflash_config - -path = Path(snapshot_download("z-lab/Qwen3-4B-DFlash-b16", allow_patterns=["*.safetensors", "*.json"])) -config = load_dflash_config(path) - -model = DFlashDraftModel(config) -weights = {} -for f in path.glob("*.safetensors"): - weights.update(load_file(str(f))) -model.load_state_dict(weights, strict=False) - -tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-4B") -target = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-4B", dtype="auto") -model.embed_tokens.weight.data.copy_(target.model.embed_tokens.weight) -model.lm_head.weight.data.copy_(target.lm_head.weight) -model.eval() -target.eval() - -prompt = "The capital of France is" -input_ids = tokenizer(prompt, return_tensors="pt").input_ids - -with torch.no_grad(): - out = target(input_ids, output_hidden_states=True) - tapped = [out.hidden_states[i + 1] for i in config.target_layer_ids] - target_hidden = torch.cat(tapped, dim=-1).float() - - last_token = input_ids[:, -1:] - block_size = 8 - draft_tokens = torch.cat( - [last_token, torch.full((1, block_size - 1), config.mask_token_id, dtype=torch.long)], dim=1 - ) - position_ids = torch.arange(target_hidden.shape[1] + block_size).unsqueeze(0) - - draft_logits = model(draft_tokens, target_hidden, position_ids) - predicted = draft_logits.argmax(-1) - -print("Prompt:", prompt) -print("Predicted continuation:", tokenizer.decode(predicted[0])) -print("Shape:", draft_logits.shape) diff --git a/examples/models/qwen3/tests/test_dflash_draft_load_weights.py b/examples/models/qwen3/tests/test_dflash_draft_load_weights.py deleted file mode 100644 index 75af1d702db..00000000000 --- a/examples/models/qwen3/tests/test_dflash_draft_load_weights.py +++ /dev/null @@ -1,50 +0,0 @@ -import json -from pathlib import Path - -import torch -from huggingface_hub import snapshot_download -from safetensors.torch import load_file -from executorch.backends.mlx.examples.llm.dflash_draft_model import DFlashConfig, DFlashDraftModel - -path = Path(snapshot_download("z-lab/Qwen3-4B-DFlash-b16", allow_patterns=["*.safetensors", "*.json"])) -cfg = json.loads((path / "config.json").read_text()) -dcfg = cfg["dflash_config"] - -config = DFlashConfig( - hidden_size=cfg["hidden_size"], - num_hidden_layers=cfg["num_hidden_layers"], - num_attention_heads=cfg["num_attention_heads"], - num_key_value_heads=cfg["num_key_value_heads"], - head_dim=cfg["head_dim"], - intermediate_size=cfg["intermediate_size"], - vocab_size=cfg["vocab_size"], - rms_norm_eps=cfg["rms_norm_eps"], - rope_theta=cfg["rope_theta"], - max_position_embeddings=cfg["max_position_embeddings"], - target_layer_ids=tuple(dcfg["target_layer_ids"]), - block_size=cfg["block_size"], - mask_token_id=dcfg["mask_token_id"], - layer_types=tuple(cfg.get("layer_types") or ["full_attention"] * cfg["num_hidden_layers"]), - sliding_window=cfg.get("sliding_window"), - final_logit_softcapping=cfg.get("final_logit_softcapping"), -) - -model = DFlashDraftModel(config) - -draft_weights = {} -for f in path.glob("*.safetensors"): - draft_weights.update(load_file(str(f))) - -print(f"Checkpoint has {len(draft_weights)} tensors") -print("First 10 checkpoint keys:", list(draft_weights.keys())[:10]) -print("First 10 model keys: ", list(model.state_dict().keys())[:10]) - -missing, unexpected = model.load_state_dict(draft_weights, strict=False) -still_missing = [k for k in missing if not k.startswith(("embed_tokens.", "lm_head."))] - -print(f"\nMissing (excl. embed/lm_head, expected empty): {still_missing}") -print(f"Unexpected (expected empty): {unexpected}") - -assert not still_missing, "Architecture mismatch — key names don't match the real checkpoint" -assert not unexpected, "Checkpoint has tensors our model doesn't define — architecture mismatch" -print("\nOK: state_dict loaded cleanly, structure matches the real checkpoint") \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_draft_pte.py b/examples/models/qwen3/tests/test_dflash_draft_pte.py deleted file mode 100644 index 374bb0e2aa0..00000000000 --- a/examples/models/qwen3/tests/test_dflash_draft_pte.py +++ /dev/null @@ -1,19 +0,0 @@ -import sys -import torch -from executorch.runtime import Runtime, Verification - -pte_path = sys.argv[1] if len(sys.argv) > 1 else "qwen3_4b_dflash_draft.pte" -et_runtime = Runtime.get() -program = et_runtime.load_program(pte_path, verification=Verification.Minimal) -method = program.load_method("forward") - -# Must match the exact static shapes used at export time: ctx_len=8, block_size=16 -block_size, ctx_len, hidden_size, vocab_size = 16, 8, 12800, 151936 -tokens = torch.randint(0, 1000, (1, block_size), dtype=torch.long) -target_hidden = torch.randn(1, ctx_len, hidden_size) -position_ids = torch.arange(ctx_len + block_size).unsqueeze(0).long() - -(draft_logits,) = method.execute([tokens, target_hidden, position_ids]) -assert draft_logits.shape == (1, block_size - 1, vocab_size), draft_logits.shape -assert not torch.isnan(draft_logits).any() and not torch.isinf(draft_logits).any() -print(f"OK: draft_logits {tuple(draft_logits.shape)}") diff --git a/examples/models/qwen3/tests/test_dflash_export.py b/examples/models/qwen3/tests/test_dflash_export.py deleted file mode 100644 index 17c6d7b0013..00000000000 --- a/examples/models/qwen3/tests/test_dflash_export.py +++ /dev/null @@ -1,32 +0,0 @@ -"""Sanity check: qwen3_4b_dflash_target.pte returns (logits, hidden) with the -expected shapes at runtime. - -Run after exporting with --dflash-layers, e.g.: - python3 export_llm_hf.py --model-id Qwen/Qwen3-4B --dflash-layers 2,18,33 ... - python3 test_dflash_export.py qwen3_4b_dflash_target.pte -""" - -import sys -import torch -from executorch.runtime import Runtime, Verification - -DFLASH_LAYERS = [1, 9, 17, 25, 33] -HIDDEN_SIZE = 2560 -EXPECTED_HIDDEN_DIM = len(DFLASH_LAYERS) * HIDDEN_SIZE # 12800 -VOCAB_SIZE = 151936 - -pte_path = sys.argv[1] -et_runtime = Runtime.get() -program = et_runtime.load_program(pte_path, verification=Verification.Minimal) -method = program.load_method("forward") - -tokens = torch.tensor([[1, 2, 3]], dtype=torch.long) -input_pos = torch.tensor([0], dtype=torch.long) -logits, hidden = method.execute([tokens, input_pos]) - -assert logits.shape == (1, 3, VOCAB_SIZE), logits.shape -assert hidden.shape == (1, 3, EXPECTED_HIDDEN_DIM), hidden.shape -assert not torch.isnan(logits).any() and not torch.isinf(logits).any() -assert not torch.isnan(hidden).any() and not torch.isinf(hidden).any() - -print(f"OK: logits {tuple(logits.shape)}, hidden {tuple(hidden.shape)}") \ No newline at end of file diff --git a/examples/models/qwen3/tests/test_dflash_lossless.py b/examples/models/qwen3/tests/test_dflash_lossless.py index 86c81de5ece..35d42f66084 100644 --- a/examples/models/qwen3/tests/test_dflash_lossless.py +++ b/examples/models/qwen3/tests/test_dflash_lossless.py @@ -1,32 +1,45 @@ -"""Verification (doc Part 8, criterion 1): greedy DFlash must be LOSSLESS -- -identical tokens to greedy autoregressive baseline, since the target verifies -every accepted token. Any divergence means the speculative loop is broken.""" -import subprocess, sys, re +"""Checks that DFlash prodcues the exact same output as normal greedy decoding.""" + +import re +import subprocess +import sys PROMPT = "Write a Python function that takes a list of integers and returns the second largest number in the list." N = 96 + def run(script, extra): out = subprocess.run( - [sys.executable, f"examples/models/qwen3/{script}", - "--prompt", PROMPT, "--max-new-tokens", str(N)] + extra, - capture_output=True, text=True, cwd=".", + [ + sys.executable, + f"examples/models/qwen3/{script}", + "--prompt", + PROMPT, + "--max-new-tokens", + str(N), + ] + + extra, + capture_output=True, + text=True, + cwd=".", ).stdout m = re.search(r"Generated \([^)]*\): (.*?)\n\n", out, re.DOTALL) return m.group(1) if m else out + baseline = run("run_baseline.py", []) dflash = run("run_dflash.py", []) -print("=== BASELINE ===\n", baseline[:400]) -print("\n=== DFLASH ===\n", dflash[:400]) -print("\n=== RESULT ===") +print("BASELINE:\n", baseline[:400]) +print("\nDFLASH:\n", dflash[:400]) +print("\nRESULT:") if baseline.strip() == dflash.strip(): print("PASS: DFlash output is token-for-token identical to baseline (LOSSLESS)") else: - # find first divergence for i, (a, b) in enumerate(zip(baseline, dflash)): if a != b: - print(f"DIVERGE at char {i}: baseline={baseline[i:i+30]!r} dflash={dflash[i:i+30]!r}") + print( + f"DIVERGE at char {i}: baseline={baseline[i:i+30]!r} dflash={dflash[i:i+30]!r}" + ) break - print("FAIL: outputs differ -- speculative loop is not lossless") + print("FAIL- outputs differ: speculative loop is not lossless") diff --git a/examples/models/qwen3/tests/test_dflash_target.py b/examples/models/qwen3/tests/test_dflash_target.py index 17c6d7b0013..1cba7b803c9 100644 --- a/examples/models/qwen3/tests/test_dflash_target.py +++ b/examples/models/qwen3/tests/test_dflash_target.py @@ -1,12 +1,11 @@ -"""Sanity check: qwen3_4b_dflash_target.pte returns (logits, hidden) with the -expected shapes at runtime. +""" +Verifies that the exported DFlash target model runs correctly and returns both logits and concatenated hidden states with the expected shapes. -Run after exporting with --dflash-layers, e.g.: - python3 export_llm_hf.py --model-id Qwen/Qwen3-4B --dflash-layers 2,18,33 ... - python3 test_dflash_export.py qwen3_4b_dflash_target.pte +Run this after exporting the target model with --dflash-layers. """ import sys + import torch from executorch.runtime import Runtime, Verification @@ -29,4 +28,4 @@ assert not torch.isnan(logits).any() and not torch.isinf(logits).any() assert not torch.isnan(hidden).any() and not torch.isinf(hidden).any() -print(f"OK: logits {tuple(logits.shape)}, hidden {tuple(hidden.shape)}") \ No newline at end of file +print(f"OK- logits {tuple(logits.shape)}, hidden {tuple(hidden.shape)}") From 01c4e8275d1cd8e576e4e38977b19d045a72ae9c Mon Sep 17 00:00:00 2001 From: Chet Hotti Date: Tue, 14 Jul 2026 05:57:12 +0000 Subject: [PATCH 4/5] Address code review: drop dead Makefile target, fix pytest-collection hazard, trim profiling scaffolding, add license headers, fix typos - Remove qwen3_dflash-mlx Makefile target + .gitignore block: depended on C++ engine files that are gitignored/not yet landed (follow-up PR) - Rename tests/test_dflash_*.py -> check_dflash_*.py so pytest's test_* glob never collects these manual, hardware-gated driver scripts - Trim sub-millisecond profiling scaffolding in run_dflash.py, keep only draft_exec/target_exec timing under --verbose (those dominate wall time) - Add missing BSD license headers to 8 files - Fix no-op --chat-template flag (only --no-chat-template had any effect) - Fix typos: frmo->from, prodcues->produces, an .pte->a .pte, output project->output projection weights - Add one-line comment on first_mismatch's draft/target length asymmetry - Remove unused default_dflash_layer_ids (layer ids come from --dflash-layers / draft config, not this helper) - Document DFlash's check_dflash_*.py as manual/CI-exempt in README --- .gitignore | 8 ++--- Makefile | 12 +++---- .../mlx/examples/llm/dflash_draft_model.py | 8 ++++- examples/models/qwen3/README.md | 22 +++++++++++++ examples/models/qwen3/export_dflash_draft.py | 10 ++++-- .../qwen3/mlx_source_transformations.py | 13 ++++---- examples/models/qwen3/run_baseline.py | 9 ++++-- examples/models/qwen3/run_dflash.py | 32 +++++++++---------- ..._dflash_draft.py => check_dflash_draft.py} | 6 ++++ ...h_lossless.py => check_dflash_lossless.py} | 8 ++++- ...flash_target.py => check_dflash_target.py} | 6 ++++ 11 files changed, 90 insertions(+), 44 deletions(-) rename examples/models/qwen3/tests/{test_dflash_draft.py => check_dflash_draft.py} (86%) rename examples/models/qwen3/tests/{test_dflash_lossless.py => check_dflash_lossless.py} (81%) rename examples/models/qwen3/tests/{test_dflash_target.py => check_dflash_target.py} (84%) diff --git a/.gitignore b/.gitignore index 69c6bcdc115..487b1d9e3ed 100644 --- a/.gitignore +++ b/.gitignore @@ -87,9 +87,5 @@ zephyr_dev_root.backup.*/ .claude/*.local.* extension/pybindings/mlx.metallib -# DFlash C++ engine (Qwen3) -- incomplete, broken build, follow-up PR -examples/models/qwen3/CMakeLists.txt -examples/models/qwen3/CMakePresets.json -examples/models/qwen3/qwen3_dflash_engine.h -examples/models/qwen3/qwen3_dflash_engine.cpp -examples/models/qwen3/main.cpp +# Scratch/WIP work not ready for review -- never committed +/wip/ diff --git a/Makefile b/Makefile index 26167aeb2b2..0891e06e1ba 100644 --- a/Makefile +++ b/Makefile @@ -485,11 +485,7 @@ clean: extension/llm/tokenizers/build \ extension/llm/tokenizers/pytorch_tokenizers.egg-info -qwen3_dflash-mlx: - @echo "==> Building and installing ExecuTorch with MLX..." - cmake --workflow --preset mlx-release - @echo "==> Building Qwen3 DFlash speculative decoding runner with MLX..." - cd examples/models/qwen3 && cmake --workflow --preset qwen3-dflash-mlx - @echo "" - @echo "✓ Build complete!" - @echo " Runner: cmake-out/examples/models/qwen3/qwen3_dflash_runner" +# qwen3_dflash-mlx target removed: it depended on the C++ engine sources +# (CMakeLists.txt, CMakePresets.json, qwen3_dflash_engine.*), which are +# gitignored/not yet landed. Restore this target in the follow-up PR that +# actually lands the C++ engine. diff --git a/backends/mlx/examples/llm/dflash_draft_model.py b/backends/mlx/examples/llm/dflash_draft_model.py index 225d330b9b9..3562d840b72 100644 --- a/backends/mlx/examples/llm/dflash_draft_model.py +++ b/backends/mlx/examples/llm/dflash_draft_model.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + """PyTorch implementation of the DFlash draft model for ExecuTorch export. This model is the lightweight "draft" network used in DFlash speculative @@ -148,7 +154,7 @@ def __init__(self, config: DFlashConfig, layer_idx: int): self.k_norm = DFlashRMSNorm(hd, config.rms_norm_eps) def forward(self, x, x_ctx, cos, sin): - """Keys and values come frmo both the projected target context and the proposal block itself. This lets the draft attend to what the target model already understands while also allowing predictions within the proposal block to interact with one another.""" + """Keys and values come from both the projected target context and the proposal block itself. This lets the draft attend to what the target model already understands while also allowing predictions within the proposal block to interact with one another.""" B, L, _ = x.shape S = x_ctx.shape[1] q = self.q_norm( diff --git a/examples/models/qwen3/README.md b/examples/models/qwen3/README.md index 123e65f16c5..2b3e3d94230 100644 --- a/examples/models/qwen3/README.md +++ b/examples/models/qwen3/README.md @@ -68,5 +68,27 @@ Note that you have to apply the chat template manually for the C++ runner. To run the model on an example iOS or Android app, see the Llama README's [Step 5: Build Mobile apps](../llama/README.md#step-5-build-mobile-apps) section. +### DFlash speculative decoding (MLX delegate) + +`export_dflash_draft.py`, `run_dflash.py`, and `run_baseline.py` implement +block-diffusion speculative decoding (DFlash) for Qwen3 on the MLX delegate. +See `mlx_source_transformations.py` for the hidden-state-tapping wrapper used +during export. + +The `check_dflash_*.py` scripts under `tests/` are manual driver scripts, not +pytest tests -- they require exported `qwen3_4b_dflash_target.pte` / +`_draft.pte` files (multi-GB, not checked in), HF downloads, and Apple +M-series hardware with the MLX delegate, so they cannot run in this repo's +CI. Run them by hand after exporting: + +```bash +python examples/models/qwen3/tests/check_dflash_target.py qwen3_4b_dflash_target.pte +python examples/models/qwen3/tests/check_dflash_draft.py qwen3_4b_dflash_draft.pte +python examples/models/qwen3/tests/check_dflash_lossless.py +``` + +The "lossless" guarantee (DFlash output is token-for-token identical to +greedy baseline decoding) is currently only verified this way, manually. + ### FAQ For more help with exporting or running this model, feel free to ask in our [discord channel](https://discord.gg/UEjkY9Zs). diff --git a/examples/models/qwen3/export_dflash_draft.py b/examples/models/qwen3/export_dflash_draft.py index 8b5fdc73291..a3c18866764 100644 --- a/examples/models/qwen3/export_dflash_draft.py +++ b/examples/models/qwen3/export_dflash_draft.py @@ -1,6 +1,12 @@ -"""Exports the DFlash draft model to an .pte program. +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. -This script loads the pretrained DFlash draft checkpoint, copies the shared embedding and output project weights from target model, applies same 4-bit quantization used by target, and exports the draft model for MLX inference. +"""Exports the DFlash draft model to a .pte program. + +This script loads the pretrained DFlash draft checkpoint, copies the shared embedding and output projection weights from target model, applies same 4-bit quantization used by target, and exports the draft model for MLX inference. The exported model is used alongside the target model during speculative decoding. """ diff --git a/examples/models/qwen3/mlx_source_transformations.py b/examples/models/qwen3/mlx_source_transformations.py index edfc41eff61..16a293c34e7 100644 --- a/examples/models/qwen3/mlx_source_transformations.py +++ b/examples/models/qwen3/mlx_source_transformations.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + """Extracting Qwen3 hidden-state for DFlash. Same idea as examples/models/gemma4_31b/mlx_source_transformations.py -- @@ -61,10 +67,3 @@ def forward( if hasattr(outs, "logits"): return outs.logits, hidden return outs.last_hidden_state, hidden - - -def default_dflash_layer_ids(num_layers: int) -> List[int]: - """This is simply the default dflash layer selection for Qwen3. - [2, N//2, N-3] tap pattern, same as Gemma 4. For Qwen3-4B (36 layers): [2, 18, 33]. - """ - return [2, num_layers // 2, num_layers - 3] diff --git a/examples/models/qwen3/run_baseline.py b/examples/models/qwen3/run_baseline.py index c37b12b3b6d..594dada473a 100644 --- a/examples/models/qwen3/run_baseline.py +++ b/examples/models/qwen3/run_baseline.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + """Standard autoregressive decoding used as the baseline for the comparison. """ @@ -15,8 +21,7 @@ def main(): p.add_argument("--tokenizer", default="Qwen/Qwen3-4B") p.add_argument("--prompt", required=True) p.add_argument("--max-new-tokens", type=int, default=128) - p.add_argument("--chat-template", action="store_true", default=True) - p.add_argument("--no-chat-template", dest="chat_template", action="store_false") + p.add_argument("--no-chat-template", dest="chat_template", action="store_false", default=True) p.add_argument("--enable-thinking", action="store_true", default=False) args = p.parse_args() diff --git a/examples/models/qwen3/run_dflash.py b/examples/models/qwen3/run_dflash.py index 849afcb874d..d4c683ea86d 100644 --- a/examples/models/qwen3/run_dflash.py +++ b/examples/models/qwen3/run_dflash.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + """Python implementation of the DFlash speculative decoding loop for the ExecuTorch MLX backend. This file coordinates the interaction between the target model and the draft model during inference. Instead of asking the target model to generate one token at a time, DFlash first lets the lightweight draft model predict a block of future tokens, then asks the target model to verify those predictions in a single forward pass. Any matching draft tokens are accepted, while the first incorrect prediction is replaced with the target model's token. The process then repeats from the updated position. @@ -44,12 +50,12 @@ def main(): p.add_argument("--prompt", default="The capital of France is") p.add_argument("--max-new-tokens", type=int, default=64) p.add_argument( - "--chat-template", - action="store_true", + "--no-chat-template", + dest="chat_template", + action="store_false", default=True, - help="Apply Qwen3's chat template (paper's eval setup). Default on.", + help="Disable Qwen3's chat template. On by default (paper's eval setup).", ) - p.add_argument("--no-chat-template", dest="chat_template", action="store_false") p.add_argument( "--enable-thinking", action="store_true", @@ -144,9 +150,7 @@ def main(): _t0 = time.time() (draft_logits,) = draft.execute([draft_input, hidden, draft_pos]) _draft_exec_time = time.time() - _t0 - _t0b = time.time() draft_ids = draft_logits[0].argmax(-1).tolist() # block_size - 1 tokens - _draft_argmax_time = time.time() - _t0b # 2. Verify the draft predictions. Target model predicts the next token after every position in the block in a single forward pass. verify_input = torch.cat( @@ -160,19 +164,17 @@ def main(): _t1 = time.time() target_logits, new_hidden = target.execute([verify_input, verify_pos]) _target_exec_time = time.time() - _t1 - _t1b = time.time() target_ids = target_logits[0].argmax(-1).tolist() # block_size tokens - _target_argmax_time = time.time() - _t1b # 3. Keep every drafting token that matches the target. At the first mismatch, stop accepting draft predictions and use the target model's token instead. - _t2 = time.time() + # (target_ids has block_size entries vs draft_ids' block_size - 1, so + # target_ids[accepted] is always in-bounds, including the all-accepted + # bonus-token case.) accepted = first_mismatch(draft_ids, target_ids) - _fm_time = time.time() - _t2 if args.verbose and rounds <= 10: print( - f" timing: draft_exec={_draft_exec_time*1000:.1f}ms draft_argmax={_draft_argmax_time*1000:.2f}ms " - f"target_exec={_target_exec_time*1000:.1f}ms target_argmax={_target_argmax_time*1000:.2f}ms " - f"first_mismatch={_fm_time*1000:.3f}ms ctx_len={hidden.shape[1]}" + f" timing: draft_exec={_draft_exec_time*1000:.1f}ms " + f"target_exec={_target_exec_time*1000:.1f}ms ctx_len={hidden.shape[1]}" ) if args.verbose and rounds <= 5: print( @@ -194,11 +196,7 @@ def main(): last_token = new_tokens[-1] # Append the hidden states for the newly accepted tokens to the running target context. # The draft model conditions on the hidden states of the entire sequence, so this context grows as generation progresses rather than being replaced each round. - _t3 = time.time() hidden = torch.cat([hidden, new_hidden[:, : len(new_tokens), :].float()], dim=1) - _cat_time = time.time() - _t3 - if args.verbose and rounds <= 10: - print(f" timing: hidden_cat={_cat_time*1000:.2f}ms") if eos_id in new_tokens: break diff --git a/examples/models/qwen3/tests/test_dflash_draft.py b/examples/models/qwen3/tests/check_dflash_draft.py similarity index 86% rename from examples/models/qwen3/tests/test_dflash_draft.py rename to examples/models/qwen3/tests/check_dflash_draft.py index f01ea0d254e..a0c66b21986 100644 --- a/examples/models/qwen3/tests/test_dflash_draft.py +++ b/examples/models/qwen3/tests/check_dflash_draft.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + """ Verifies that the exported draft .pte loads, runs correctly, and supports dynamic context lengths as the accumulated target hidden-state context grows during speculative decoding. A successful export and execution also confirms that the checkpoint weights and exported model are compatible. """ diff --git a/examples/models/qwen3/tests/test_dflash_lossless.py b/examples/models/qwen3/tests/check_dflash_lossless.py similarity index 81% rename from examples/models/qwen3/tests/test_dflash_lossless.py rename to examples/models/qwen3/tests/check_dflash_lossless.py index 35d42f66084..ea09a6bd196 100644 --- a/examples/models/qwen3/tests/test_dflash_lossless.py +++ b/examples/models/qwen3/tests/check_dflash_lossless.py @@ -1,4 +1,10 @@ -"""Checks that DFlash prodcues the exact same output as normal greedy decoding.""" +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Checks that DFlash produces the exact same output as normal greedy decoding.""" import re import subprocess diff --git a/examples/models/qwen3/tests/test_dflash_target.py b/examples/models/qwen3/tests/check_dflash_target.py similarity index 84% rename from examples/models/qwen3/tests/test_dflash_target.py rename to examples/models/qwen3/tests/check_dflash_target.py index 1cba7b803c9..51f7e7e7f25 100644 --- a/examples/models/qwen3/tests/test_dflash_target.py +++ b/examples/models/qwen3/tests/check_dflash_target.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + """ Verifies that the exported DFlash target model runs correctly and returns both logits and concatenated hidden states with the expected shapes. From 9f86c900689073de9822944dd58aa11046e2995b Mon Sep 17 00:00:00 2001 From: cthotti Date: Tue, 21 Jul 2026 18:36:19 +0530 Subject: [PATCH 5/5] Implementing Dflash_experiments.md for next users. --- backends/mlx/_generated_inspector.py | 966 ++ backends/mlx/runtime/MLXLoader.cpp | 2473 +++ backends/mlx/runtime/MLXLoader.h | 2743 +++ backends/mlx/runtime/schema_generated.h | 14154 ++++++++++++++++ .../mlx/serialization/_generated/__init__.py | 153 + .../_generated/mlx_delegate/ARangeNode.py | 118 + .../_generated/mlx_delegate/AbsNode.py | 71 + .../_generated/mlx_delegate/AddIntNode.py | 88 + .../_generated/mlx_delegate/AddNode.py | 88 + .../_generated/mlx_delegate/AddmmNode.py | 131 + .../_generated/mlx_delegate/AllNode.py | 123 + .../_generated/mlx_delegate/AnyNode.py | 123 + .../_generated/mlx_delegate/ArccosNode.py | 71 + .../_generated/mlx_delegate/ArccoshNode.py | 71 + .../_generated/mlx_delegate/ArcsinNode.py | 71 + .../_generated/mlx_delegate/ArcsinhNode.py | 71 + .../_generated/mlx_delegate/ArctanNode.py | 71 + .../_generated/mlx_delegate/ArctanhNode.py | 71 + .../mlx_delegate/ArgPartitionNode.py | 101 + .../_generated/mlx_delegate/ArgmaxNode.py | 97 + .../_generated/mlx_delegate/ArgminNode.py | 97 + .../_generated/mlx_delegate/ArgsortNode.py | 84 + .../_generated/mlx_delegate/AsStridedNode.py | 158 + .../_generated/mlx_delegate/AsTypeNode.py | 84 + .../_generated/mlx_delegate/Atan2Node.py | 88 + .../_generated/mlx_delegate/BitwiseAndNode.py | 88 + .../mlx_delegate/BitwiseInvertNode.py | 71 + .../_generated/mlx_delegate/BitwiseOrNode.py | 88 + .../_generated/mlx_delegate/BitwiseXorNode.py | 88 + .../mlx_delegate/BroadcastToNode.py | 108 + .../_generated/mlx_delegate/CeilNode.py | 71 + .../_generated/mlx_delegate/ClipNode.py | 105 + .../mlx_delegate/ConcatenateNode.py | 103 + .../_generated/mlx_delegate/ContiguousNode.py | 71 + .../_generated/mlx_delegate/Conv1DNode.py | 140 + .../_generated/mlx_delegate/Conv2DNode.py | 179 + .../_generated/mlx_delegate/Conv3DNode.py | 218 + .../mlx_delegate/ConvTranspose1DNode.py | 153 + .../mlx_delegate/ConvTranspose2DNode.py | 205 + .../mlx_delegate/ConvTranspose3DNode.py | 257 + .../_generated/mlx_delegate/CosNode.py | 71 + .../_generated/mlx_delegate/CoshNode.py | 71 + .../_generated/mlx_delegate/CumsumNode.py | 110 + .../_generated/mlx_delegate/DequantizeNode.py | 174 + .../_generated/mlx_delegate/DivideNode.py | 88 + .../_generated/mlx_delegate/EqualNode.py | 88 + .../_generated/mlx_delegate/ErfNode.py | 71 + .../_generated/mlx_delegate/ExpNode.py | 71 + .../_generated/mlx_delegate/ExpandDimsNode.py | 84 + .../_generated/mlx_delegate/Expm1Node.py | 71 + .../_generated/mlx_delegate/FloatOrVid.py | 80 + .../mlx_delegate/FloorDivideIntNode.py | 88 + .../mlx_delegate/FloorDivideNode.py | 88 + .../_generated/mlx_delegate/FloorNode.py | 71 + .../_generated/mlx_delegate/FullLikeNode.py | 101 + .../_generated/mlx_delegate/FullNode.py | 121 + .../_generated/mlx_delegate/GatherMmNode.py | 135 + .../_generated/mlx_delegate/GatherNode.py | 185 + .../_generated/mlx_delegate/GatherQmmNode.py | 221 + .../_generated/mlx_delegate/GeluNode.py | 84 + .../mlx_delegate/GreaterEqualNode.py | 88 + .../_generated/mlx_delegate/GreaterNode.py | 88 + .../_generated/mlx_delegate/IdCopyNode.py | 71 + .../_generated/mlx_delegate/IfNode.py | 80 + .../_generated/mlx_delegate/IndexCopyNode.py | 118 + .../_generated/mlx_delegate/Instruction.py | 66 + .../mlx_delegate/InstructionChain.py | 74 + .../_generated/mlx_delegate/IntOrVid.py | 80 + .../_generated/mlx_delegate/IntOrVidOrTid.py | 97 + .../_generated/mlx_delegate/ItemIntNode.py | 71 + .../_generated/mlx_delegate/LayerNormNode.py | 118 + .../_generated/mlx_delegate/LessEqualNode.py | 88 + .../_generated/mlx_delegate/LessNode.py | 88 + .../_generated/mlx_delegate/Log10Node.py | 71 + .../_generated/mlx_delegate/Log1pNode.py | 71 + .../_generated/mlx_delegate/Log2Node.py | 71 + .../_generated/mlx_delegate/LogAddExpNode.py | 88 + .../_generated/mlx_delegate/LogNode.py | 71 + .../_generated/mlx_delegate/LogSumExpNode.py | 123 + .../_generated/mlx_delegate/LogicalAndNode.py | 88 + .../_generated/mlx_delegate/LogicalNotNode.py | 71 + .../_generated/mlx_delegate/LogicalOrNode.py | 88 + .../_generated/mlx_delegate/MLXGraph.py | 376 + .../_generated/mlx_delegate/MaxNode.py | 123 + .../_generated/mlx_delegate/MaximumNode.py | 88 + .../_generated/mlx_delegate/MeanNode.py | 123 + .../_generated/mlx_delegate/MedianNode.py | 123 + .../mlx_delegate/MetalKernelNode.py | 550 + .../_generated/mlx_delegate/MinNode.py | 123 + .../_generated/mlx_delegate/MinimumNode.py | 88 + .../_generated/mlx_delegate/ModIntNode.py | 88 + .../mlx_delegate/MultiplyIntNode.py | 88 + .../_generated/mlx_delegate/MultiplyNode.py | 88 + .../_generated/mlx_delegate/NamedSlot.py | 67 + .../_generated/mlx_delegate/NegNode.py | 71 + .../_generated/mlx_delegate/NoopNode.py | 37 + .../_generated/mlx_delegate/NotEqualNode.py | 88 + .../_generated/mlx_delegate/OpNode.py | 139 + .../_generated/mlx_delegate/PadNode.py | 134 + .../_generated/mlx_delegate/PartitionNode.py | 101 + .../_generated/mlx_delegate/PowerNode.py | 88 + .../_generated/mlx_delegate/ProdNode.py | 123 + .../mlx_delegate/QuantizedMatmulNode.py | 174 + .../_generated/mlx_delegate/RMSNormNode.py | 101 + .../_generated/mlx_delegate/RandomBitsNode.py | 121 + .../_generated/mlx_delegate/ReciprocalNode.py | 71 + .../_generated/mlx_delegate/RemainderNode.py | 88 + .../_generated/mlx_delegate/RepeatNode.py | 101 + .../_generated/mlx_delegate/ReshapeNode.py | 108 + .../_generated/mlx_delegate/RollNode.py | 147 + .../_generated/mlx_delegate/RopeNode.py | 157 + .../_generated/mlx_delegate/RoundNode.py | 84 + .../_generated/mlx_delegate/RsqrtNode.py | 71 + .../_generated/mlx_delegate/ScanNode.py | 207 + .../_generated/mlx_delegate/ScatterAddNode.py | 118 + .../_generated/mlx_delegate/SdpaNode.py | 148 + .../_generated/mlx_delegate/ShapeDim.py | 76 + .../_generated/mlx_delegate/SigmoidNode.py | 71 + .../_generated/mlx_delegate/SignNode.py | 71 + .../_generated/mlx_delegate/SiluNode.py | 71 + .../_generated/mlx_delegate/SinNode.py | 71 + .../_generated/mlx_delegate/SinhNode.py | 71 + .../_generated/mlx_delegate/SliceNode.py | 135 + .../mlx_delegate/SliceUpdateNode.py | 152 + .../_generated/mlx_delegate/SlotType.py | 9 + .../_generated/mlx_delegate/SlotVariant.py | 63 + .../_generated/mlx_delegate/SoftmaxNode.py | 97 + .../_generated/mlx_delegate/SortNode.py | 84 + .../_generated/mlx_delegate/SplitNode.py | 140 + .../_generated/mlx_delegate/SqrtNode.py | 71 + .../_generated/mlx_delegate/SquareNode.py | 71 + .../_generated/mlx_delegate/SqueezeNode.py | 110 + .../_generated/mlx_delegate/StackNode.py | 103 + .../_generated/mlx_delegate/StdNode.py | 136 + .../mlx_delegate/SubtractIntNode.py | 88 + .../_generated/mlx_delegate/SubtractNode.py | 88 + .../_generated/mlx_delegate/SumNode.py | 123 + .../_generated/mlx_delegate/SymSizeNode.py | 84 + .../mlx_delegate/TakeAlongAxisNode.py | 101 + .../_generated/mlx_delegate/TakeNode.py | 101 + .../_generated/mlx_delegate/TanNode.py | 71 + .../_generated/mlx_delegate/TanhNode.py | 71 + .../_generated/mlx_delegate/TensorMeta.py | 126 + .../_generated/mlx_delegate/Tid.py | 26 + .../_generated/mlx_delegate/TileNode.py | 108 + .../_generated/mlx_delegate/TransposeNode.py | 110 + .../_generated/mlx_delegate/TriNode.py | 114 + .../_generated/mlx_delegate/TrilNode.py | 84 + .../_generated/mlx_delegate/TriuNode.py | 84 + .../_generated/mlx_delegate/VarNode.py | 136 + .../_generated/mlx_delegate/Vid.py | 26 + .../_generated/mlx_delegate/VidOrTid.py | 84 + .../_generated/mlx_delegate/WhereNode.py | 105 + .../_generated/mlx_delegate/__init__.py | 0 .../serialization/_generated_serializers.py | 2987 ++++ .../mlx/serialization/mlx_graph_schema.py | 1384 ++ backends/mlx/third-party/mlx | 2 +- examples/models/qwen3/DFLASH_EXPERIMENTS.md | 81 + examples/models/qwen3/SUMMARY.txt | 14 + 159 files changed, 40454 insertions(+), 1 deletion(-) create mode 100644 backends/mlx/_generated_inspector.py create mode 100644 backends/mlx/runtime/MLXLoader.cpp create mode 100644 backends/mlx/runtime/MLXLoader.h create mode 100644 backends/mlx/runtime/schema_generated.h create mode 100644 backends/mlx/serialization/_generated/__init__.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ARangeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AbsNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AddIntNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AddNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AddmmNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AllNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AnyNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArccosNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArccoshNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArcsinNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArcsinhNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArctanNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArctanhNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArgPartitionNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArgmaxNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArgminNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ArgsortNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AsStridedNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/AsTypeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Atan2Node.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/BitwiseAndNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/BitwiseInvertNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/BitwiseOrNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/BitwiseXorNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/BroadcastToNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/CeilNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ClipNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ConcatenateNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ContiguousNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Conv1DNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Conv2DNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Conv3DNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose1DNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose2DNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose3DNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/CosNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/CoshNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/CumsumNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/DequantizeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/DivideNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/EqualNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ErfNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ExpNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ExpandDimsNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Expm1Node.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/FloatOrVid.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/FloorDivideIntNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/FloorDivideNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/FloorNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/FullLikeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/FullNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/GatherMmNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/GatherNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/GatherQmmNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/GeluNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/GreaterEqualNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/GreaterNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/IdCopyNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/IfNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/IndexCopyNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Instruction.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/InstructionChain.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/IntOrVid.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/IntOrVidOrTid.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ItemIntNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LayerNormNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LessEqualNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LessNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Log10Node.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Log1pNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Log2Node.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LogAddExpNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LogNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LogSumExpNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LogicalAndNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LogicalNotNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/LogicalOrNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MLXGraph.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MaxNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MaximumNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MeanNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MedianNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MetalKernelNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MinNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MinimumNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ModIntNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MultiplyIntNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/MultiplyNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/NamedSlot.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/NegNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/NoopNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/NotEqualNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/OpNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/PadNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/PartitionNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/PowerNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ProdNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/QuantizedMatmulNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RMSNormNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RandomBitsNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ReciprocalNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RemainderNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RepeatNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ReshapeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RollNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RopeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RoundNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/RsqrtNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ScanNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ScatterAddNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SdpaNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/ShapeDim.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SigmoidNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SignNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SiluNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SinNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SinhNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SliceNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SliceUpdateNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SlotType.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SlotVariant.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SoftmaxNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SortNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SplitNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SqrtNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SquareNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SqueezeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/StackNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/StdNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SubtractIntNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SubtractNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SumNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/SymSizeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TakeAlongAxisNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TakeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TanNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TanhNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TensorMeta.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Tid.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TileNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TransposeNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TriNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TrilNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/TriuNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/VarNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/Vid.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/VidOrTid.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/WhereNode.py create mode 100644 backends/mlx/serialization/_generated/mlx_delegate/__init__.py create mode 100644 backends/mlx/serialization/_generated_serializers.py create mode 100644 backends/mlx/serialization/mlx_graph_schema.py create mode 100644 examples/models/qwen3/DFLASH_EXPERIMENTS.md create mode 100644 examples/models/qwen3/SUMMARY.txt diff --git a/backends/mlx/_generated_inspector.py b/backends/mlx/_generated_inspector.py new file mode 100644 index 00000000000..380837874d2 --- /dev/null +++ b/backends/mlx/_generated_inspector.py @@ -0,0 +1,966 @@ +# +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. +# +# ============================================================================ +# AUTO-GENERATED FILE - DO NOT EDIT MANUALLY +# ============================================================================ +# +# This file was generated from schema.fbs by the MLX delegate code generator. +# +# Source: backends/mlx/serialization/schema.fbs +# Generator: backends/mlx/serialization/generate.py +# +# To regenerate, run from the executorch root: +# python backends/mlx/serialization/generate.py +# +# ============================================================================ + +""" +Auto-generated inspector field mappings for MLX delegate. + +This module provides field metadata for each op node type, enabling +the pte_inspector to parse FlatBuffer op nodes without manually +maintaining field mappings. +""" + +from __future__ import annotations + +from typing import Dict, List, Tuple + + +# Field kinds and their extractors +# Each field is a tuple of (display_name, accessor_name, kind) +# where kind is one of: 'tid', 'vid', 'int_or_vid', 'float_or_vid', +# 'int_list', 'int_or_vid_list', 'tid_list', 'string_list', 'scalar', 'string' + +FieldSpec = Tuple[str, str, str] # (display_name, accessor_name, kind) + + +# Mapping from op node name to list of field specs +OP_NODE_FIELDS: Dict[str, List[FieldSpec]] = { + "NoopNode": [ + ], + "IdCopyNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "AddmmNode": [ + ("mat1", "Mat1", "tid"), + ("mat2", "Mat2", "tid"), + ("out", "Out", "tid"), + ("bias", "Bias", "tid"), + ("alpha", "Alpha", "scalar"), + ("beta", "Beta", "scalar"), + ], + "ItemIntNode": [ + ("x", "X", "tid"), + ("out", "Out", "vid"), + ], + "ExpandDimsNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "TileNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("reps", "Reps", "int_or_vid_list"), + ], + "TakeAlongAxisNode": [ + ("x", "X", "tid"), + ("indices", "Indices", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "TakeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("index", "Index", "int_or_vid_or_tid"), + ("axis", "Axis", "scalar"), + ], + "RMSNormNode": [ + ("x", "X", "tid"), + ("weight", "Weight", "tid"), + ("out", "Out", "tid"), + ("eps", "Eps", "scalar"), + ], + "LayerNormNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("weight", "Weight", "tid"), + ("bias", "Bias", "tid"), + ("eps", "Eps", "scalar"), + ], + "RopeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("dims", "Dims", "scalar"), + ("offset", "Offset", "vid_or_tid"), + ("freqs", "Freqs", "tid"), + ("traditional", "Traditional", "scalar"), + ("base", "Base", "scalar"), + ("scale", "Scale", "scalar"), + ], + "SdpaNode": [ + ("q", "Q", "tid"), + ("k", "K", "tid"), + ("v", "V", "tid"), + ("out", "Out", "tid"), + ("scale", "Scale", "scalar"), + ("mask", "Mask", "tid"), + ("causal", "Causal", "scalar"), + ], + "AddNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "AddIntNode": [ + ("a", "A", "int_or_vid"), + ("b", "B", "int_or_vid"), + ("out", "Out", "vid"), + ], + "SubtractIntNode": [ + ("a", "A", "int_or_vid"), + ("b", "B", "int_or_vid"), + ("out", "Out", "vid"), + ], + "MultiplyIntNode": [ + ("a", "A", "int_or_vid"), + ("b", "B", "int_or_vid"), + ("out", "Out", "vid"), + ], + "FloorDivideIntNode": [ + ("a", "A", "int_or_vid"), + ("b", "B", "int_or_vid"), + ("out", "Out", "vid"), + ], + "ModIntNode": [ + ("a", "A", "int_or_vid"), + ("b", "B", "int_or_vid"), + ("out", "Out", "vid"), + ], + "SymSizeNode": [ + ("a", "A", "tid"), + ("dim", "Dim", "scalar"), + ("out", "Out", "vid"), + ], + "MultiplyNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "DivideNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "SubtractNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "Conv1DNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("out", "Out", "tid"), + ("stride", "Stride", "scalar"), + ("padding", "Padding", "scalar"), + ("dilation", "Dilation", "scalar"), + ("groups", "Groups", "scalar"), + ], + "Conv2DNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("out", "Out", "tid"), + ("stride_h", "StrideH", "scalar"), + ("stride_w", "StrideW", "scalar"), + ("padding_h", "PaddingH", "scalar"), + ("padding_w", "PaddingW", "scalar"), + ("dilation_h", "DilationH", "scalar"), + ("dilation_w", "DilationW", "scalar"), + ("groups", "Groups", "scalar"), + ], + "Conv3DNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("out", "Out", "tid"), + ("stride_d", "StrideD", "scalar"), + ("stride_h", "StrideH", "scalar"), + ("stride_w", "StrideW", "scalar"), + ("padding_d", "PaddingD", "scalar"), + ("padding_h", "PaddingH", "scalar"), + ("padding_w", "PaddingW", "scalar"), + ("dilation_d", "DilationD", "scalar"), + ("dilation_h", "DilationH", "scalar"), + ("dilation_w", "DilationW", "scalar"), + ("groups", "Groups", "scalar"), + ], + "ConvTranspose1DNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("out", "Out", "tid"), + ("stride", "Stride", "scalar"), + ("padding", "Padding", "scalar"), + ("dilation", "Dilation", "scalar"), + ("output_padding", "OutputPadding", "scalar"), + ("groups", "Groups", "scalar"), + ], + "ConvTranspose2DNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("out", "Out", "tid"), + ("stride_h", "StrideH", "scalar"), + ("stride_w", "StrideW", "scalar"), + ("padding_h", "PaddingH", "scalar"), + ("padding_w", "PaddingW", "scalar"), + ("dilation_h", "DilationH", "scalar"), + ("dilation_w", "DilationW", "scalar"), + ("output_padding_h", "OutputPaddingH", "scalar"), + ("output_padding_w", "OutputPaddingW", "scalar"), + ("groups", "Groups", "scalar"), + ], + "ConvTranspose3DNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("out", "Out", "tid"), + ("stride_d", "StrideD", "scalar"), + ("stride_h", "StrideH", "scalar"), + ("stride_w", "StrideW", "scalar"), + ("padding_d", "PaddingD", "scalar"), + ("padding_h", "PaddingH", "scalar"), + ("padding_w", "PaddingW", "scalar"), + ("dilation_d", "DilationD", "scalar"), + ("dilation_h", "DilationH", "scalar"), + ("dilation_w", "DilationW", "scalar"), + ("output_padding_d", "OutputPaddingD", "scalar"), + ("output_padding_h", "OutputPaddingH", "scalar"), + ("output_padding_w", "OutputPaddingW", "scalar"), + ("groups", "Groups", "scalar"), + ], + "GeluNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("approximate", "Approximate", "string"), + ], + "ARangeNode": [ + ("out", "Out", "tid"), + ("start", "Start", "int_or_vid"), + ("stop", "Stop", "int_or_vid"), + ("step", "Step", "int_or_vid"), + ("scalar_type", "ScalarType", "scalar"), + ], + "SiluNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SigmoidNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "TanhNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SqueezeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("dims", "Dims", "int_list"), + ], + "SplitNode": [ + ("x", "X", "tid"), + ("outs", "Outs", "tid_list"), + ("sizes", "Sizes", "int_or_vid_list"), + ("axis", "Axis", "scalar"), + ], + "RsqrtNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "MaximumNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "MinimumNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "LogNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SoftmaxNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ("precise", "Precise", "scalar"), + ], + "BroadcastToNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("shape", "Shape", "int_or_vid_list"), + ], + "PadNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("pad_width", "PadWidth", "int_or_vid_list"), + ("mode", "Mode", "string"), + ("constant_value", "ConstantValue", "scalar"), + ], + "WhereNode": [ + ("condition", "Condition", "tid"), + ("x", "X", "tid"), + ("y", "Y", "tid"), + ("out", "Out", "tid"), + ], + "ReshapeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("shape", "Shape", "int_or_vid_list"), + ], + "TransposeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("perm", "Perm", "int_list"), + ], + "AsStridedNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("shape", "Shape", "int_or_vid_list"), + ("strides", "Strides", "int_or_vid_list"), + ("offset", "Offset", "scalar"), + ], + "ContiguousNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "GatherNode": [ + ("x", "X", "tid"), + ("indices", "Indices", "tid_list"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("slice_sizes", "SliceSizes", "int_list"), + ], + "SliceNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "int_or_vid"), + ("start", "Start", "int_or_vid"), + ("stop", "Stop", "int_or_vid"), + ("step", "Step", "scalar"), + ], + "AsTypeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("scalar_type", "ScalarType", "scalar"), + ], + "QuantizedMatmulNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("scales", "Scales", "tid"), + ("out", "Out", "tid"), + ("biases", "Biases", "tid"), + ("group_size", "GroupSize", "scalar"), + ("bits", "Bits", "scalar"), + ("mode", "Mode", "string"), + ("transpose", "Transpose", "scalar"), + ], + "ScatterAddNode": [ + ("x", "X", "tid"), + ("indices", "Indices", "tid"), + ("updates", "Updates", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "ConcatenateNode": [ + ("tensors", "Tensors", "tid_list"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "FullNode": [ + ("out", "Out", "tid"), + ("shape", "Shape", "int_or_vid_list"), + ("v", "V", "float_or_vid"), + ("scalar_type", "ScalarType", "scalar"), + ], + "FullLikeNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("v", "V", "float_or_vid"), + ("scalar_type", "ScalarType", "scalar"), + ], + "ArgmaxNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ("keepdims", "Keepdims", "scalar"), + ], + "SliceUpdateNode": [ + ("dst", "Dst", "tid"), + ("update", "Update", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "int_or_vid"), + ("start", "Start", "int_or_vid"), + ("stop", "Stop", "int_or_vid"), + ("step", "Step", "scalar"), + ], + "IndexCopyNode": [ + ("dst", "Dst", "tid"), + ("update", "Update", "tid"), + ("indices", "Indices", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "DequantizeNode": [ + ("w", "W", "tid"), + ("scales", "Scales", "tid"), + ("out", "Out", "tid"), + ("biases", "Biases", "tid"), + ("group_size", "GroupSize", "scalar"), + ("bits", "Bits", "scalar"), + ("mode", "Mode", "string"), + ("global_scale", "GlobalScale", "tid"), + ("dtype", "Dtype", "scalar"), + ], + "LessNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "LessEqualNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "GreaterNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "GreaterEqualNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "EqualNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "NotEqualNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "LogicalNotNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "BitwiseInvertNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "LogicalAndNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "LogicalOrNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "BitwiseAndNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "BitwiseOrNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "BitwiseXorNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "TriNode": [ + ("out", "Out", "tid"), + ("n", "N", "int_or_vid"), + ("m", "M", "int_or_vid"), + ("k", "K", "scalar"), + ("scalar_type", "ScalarType", "scalar"), + ], + "TrilNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("k", "K", "scalar"), + ], + "TriuNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("k", "K", "scalar"), + ], + "ClipNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("a_min", "AMin", "tid"), + ("a_max", "AMax", "tid"), + ], + "CumsumNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ("reverse", "Reverse", "scalar"), + ("inclusive", "Inclusive", "scalar"), + ], + "StackNode": [ + ("tensors", "Tensors", "tid_list"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "SignNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "AnyNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "AllNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "RepeatNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("repeats", "Repeats", "int_or_vid"), + ("axis", "Axis", "scalar"), + ], + "SortNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "ArgsortNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ], + "PartitionNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("kth", "Kth", "int_or_vid"), + ("axis", "Axis", "scalar"), + ], + "ArgPartitionNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("kth", "Kth", "int_or_vid"), + ("axis", "Axis", "scalar"), + ], + "RollNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("shift", "Shift", "int_or_vid_list"), + ("axes", "Axes", "int_list"), + ], + "FloorNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "CeilNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SquareNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ExpNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SinNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "CosNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "TanNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ArcsinNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ArccosNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ArctanNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SinhNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "CoshNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ArcsinhNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ArccoshNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ArctanhNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "Log2Node": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "Log10Node": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "Log1pNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "ErfNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "Expm1Node": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "RoundNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("decimals", "Decimals", "scalar"), + ], + "ReciprocalNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "SqrtNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "AbsNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "NegNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ], + "Atan2Node": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "LogAddExpNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "FloorDivideNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "RemainderNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "PowerNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ], + "LogSumExpNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "SumNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "MeanNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "VarNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ("ddof", "Ddof", "scalar"), + ], + "StdNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ("ddof", "Ddof", "scalar"), + ], + "ProdNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "MaxNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "MinNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "ArgminNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axis", "Axis", "scalar"), + ("keepdims", "Keepdims", "scalar"), + ], + "MedianNode": [ + ("x", "X", "tid"), + ("out", "Out", "tid"), + ("axes", "Axes", "int_list"), + ("keepdims", "Keepdims", "scalar"), + ], + "GatherMmNode": [ + ("a", "A", "tid"), + ("b", "B", "tid"), + ("out", "Out", "tid"), + ("lhs_indices", "LhsIndices", "tid"), + ("rhs_indices", "RhsIndices", "tid"), + ("sorted_indices", "SortedIndices", "scalar"), + ], + "GatherQmmNode": [ + ("x", "X", "tid"), + ("w", "W", "tid"), + ("scales", "Scales", "tid"), + ("out", "Out", "tid"), + ("mode", "Mode", "string"), + ("biases", "Biases", "tid"), + ("lhs_indices", "LhsIndices", "tid"), + ("rhs_indices", "RhsIndices", "tid"), + ("transpose", "Transpose", "scalar"), + ("group_size", "GroupSize", "scalar"), + ("bits", "Bits", "scalar"), + ("sorted_indices", "SortedIndices", "scalar"), + ], + "ScanNode": [ + ("originals", "Originals", "tid_list"), + ("sliced", "Sliced", "tid_list"), + ("outputs", "Outputs", "tid_list"), + ("carry", "Carry", "tid_list"), + ("body_chain_idx", "BodyChainIdx", "scalar"), + ("scan_axis", "ScanAxis", "scalar"), + ], + "IfNode": [ + ("cond", "Cond", "int_or_vid"), + ("then_chain_idx", "ThenChainIdx", "scalar"), + ("else_chain_idx", "ElseChainIdx", "scalar"), + ], + "RandomBitsNode": [ + ("out", "Out", "tid"), + ("shape", "Shape", "int_or_vid_list"), + ("seed", "Seed", "vid"), + ("width", "Width", "scalar"), + ], + "MetalKernelNode": [ + ("name", "Name", "string"), + ("source", "Source", "string"), + ("inputs", "Inputs", "tid_list"), + ("outputs", "Outputs", "tid_list"), + ("grid", "Grid", "int_or_vid_list"), + ("threadgroup", "Threadgroup", "int_or_vid_list"), + ("header", "Header", "string"), + ("input_names", "InputNames", "string_list"), + ("output_names", "OutputNames", "string_list"), + ("ensure_row_contiguous", "EnsureRowContiguous", "scalar"), + ("atomic_outputs", "AtomicOutputs", "scalar"), + ("output_shapes_flat", "OutputShapesFlat", "int_or_vid_list"), + ("output_shape_lengths", "OutputShapeLengths", "int_list"), + ("output_dtypes", "OutputDtypes", "int_list"), + ("template_arg_names", "TemplateArgNames", "string_list"), + ("template_arg_kinds", "TemplateArgKinds", "int_list"), + ("template_arg_values", "TemplateArgValues", "int_list"), + ("init_value", "InitValue", "scalar"), + ], +} + + +# List of all op node names (for dynamic imports) +OP_NODE_NAMES: List[str] = [ + "NoopNode", + "IdCopyNode", + "AddmmNode", + "ItemIntNode", + "ExpandDimsNode", + "TileNode", + "TakeAlongAxisNode", + "TakeNode", + "RMSNormNode", + "LayerNormNode", + "RopeNode", + "SdpaNode", + "AddNode", + "AddIntNode", + "SubtractIntNode", + "MultiplyIntNode", + "FloorDivideIntNode", + "ModIntNode", + "SymSizeNode", + "MultiplyNode", + "DivideNode", + "SubtractNode", + "Conv1DNode", + "Conv2DNode", + "Conv3DNode", + "ConvTranspose1DNode", + "ConvTranspose2DNode", + "ConvTranspose3DNode", + "GeluNode", + "ARangeNode", + "SiluNode", + "SigmoidNode", + "TanhNode", + "SqueezeNode", + "SplitNode", + "RsqrtNode", + "MaximumNode", + "MinimumNode", + "LogNode", + "SoftmaxNode", + "BroadcastToNode", + "PadNode", + "WhereNode", + "ReshapeNode", + "TransposeNode", + "AsStridedNode", + "ContiguousNode", + "GatherNode", + "SliceNode", + "AsTypeNode", + "QuantizedMatmulNode", + "ScatterAddNode", + "ConcatenateNode", + "FullNode", + "FullLikeNode", + "ArgmaxNode", + "SliceUpdateNode", + "IndexCopyNode", + "DequantizeNode", + "LessNode", + "LessEqualNode", + "GreaterNode", + "GreaterEqualNode", + "EqualNode", + "NotEqualNode", + "LogicalNotNode", + "BitwiseInvertNode", + "LogicalAndNode", + "LogicalOrNode", + "BitwiseAndNode", + "BitwiseOrNode", + "BitwiseXorNode", + "TriNode", + "TrilNode", + "TriuNode", + "ClipNode", + "CumsumNode", + "StackNode", + "SignNode", + "AnyNode", + "AllNode", + "RepeatNode", + "SortNode", + "ArgsortNode", + "PartitionNode", + "ArgPartitionNode", + "RollNode", + "FloorNode", + "CeilNode", + "SquareNode", + "ExpNode", + "SinNode", + "CosNode", + "TanNode", + "ArcsinNode", + "ArccosNode", + "ArctanNode", + "SinhNode", + "CoshNode", + "ArcsinhNode", + "ArccoshNode", + "ArctanhNode", + "Log2Node", + "Log10Node", + "Log1pNode", + "ErfNode", + "Expm1Node", + "RoundNode", + "ReciprocalNode", + "SqrtNode", + "AbsNode", + "NegNode", + "Atan2Node", + "LogAddExpNode", + "FloorDivideNode", + "RemainderNode", + "PowerNode", + "LogSumExpNode", + "SumNode", + "MeanNode", + "VarNode", + "StdNode", + "ProdNode", + "MaxNode", + "MinNode", + "ArgminNode", + "MedianNode", + "GatherMmNode", + "GatherQmmNode", + "ScanNode", + "IfNode", + "RandomBitsNode", + "MetalKernelNode", +] diff --git a/backends/mlx/runtime/MLXLoader.cpp b/backends/mlx/runtime/MLXLoader.cpp new file mode 100644 index 00000000000..98ef787a7bc --- /dev/null +++ b/backends/mlx/runtime/MLXLoader.cpp @@ -0,0 +1,2473 @@ +// +// Copyright (c) Meta Platforms, Inc. and affiliates. +// All rights reserved. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. +// +// ============================================================================ +// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY +// ============================================================================ +// +// This file was generated from schema.fbs by the MLX delegate code generator. +// +// Source: backends/mlx/serialization/schema.fbs +// Generator: backends/mlx/serialization/generate.py +// +// To regenerate, run from the executorch root: +// python backends/mlx/serialization/generate.py +// +// ============================================================================ +// -*- c++ -*- + +#include "MLXLoader.h" + +#include +#include + +namespace executorch { +namespace backends { +namespace mlx { +namespace loader { + +namespace { + +// Header structure for MLX payload +constexpr size_t kHeaderSize = 24; +constexpr uint32_t kMagic = 0x30584C4D; // "MLX0" in little-endian + +struct MLXHeader { + uint32_t padding; + uint32_t magic; + uint64_t data_offset; + uint64_t data_size; +}; +static_assert(sizeof(MLXHeader) == kHeaderSize, "MLXHeader size mismatch"); + +bool parse_header(const void* data, size_t size, MLXHeader& header) { + if (size < kHeaderSize) { + return false; + } + std::memcpy(&header, data, sizeof(MLXHeader)); + if (header.magic != kMagic) { + return false; + } + // Validate data_offset: must be strictly greater than kHeaderSize (so the + // FlatBuffer region is non-empty) and must not exceed the total buffer size. + if (header.data_offset <= kHeaderSize || header.data_offset > size) { + return false; + } + return true; +} + +// Helper to convert FlatBuffer vectors to std::vector. +// Caps size to prevent unbounded allocations from malformed payloads. +template +std::vector to_vector(const flatbuffers::Vector* fb_vec) { + if (!fb_vec) { + return {}; + } + constexpr size_t kMaxVectorSize = 1'000'000; + if (fb_vec->size() > kMaxVectorSize) { + throw std::runtime_error( + "FlatBuffer vector size " + std::to_string(fb_vec->size()) + + " exceeds maximum of " + std::to_string(kMaxVectorSize)); + } + return std::vector(fb_vec->begin(), fb_vec->end()); +} + +} // namespace + +// ============================================================================= +// load_instruction - AUTO-GENERATED switch statement +// ============================================================================= + +Instruction load_instruction( + const mlx_delegate::Instruction* fb_instr, StringPool& strpool) { + Instruction instr; + + if (!fb_instr || !fb_instr->op()) { + instr.op = OpCode::NOOP; + instr.node = NoopNode{}; + return instr; + } + + auto op_type = fb_instr->op_type(); + + switch (op_type) { + case mlx_delegate::OpNode_NoopNode: { + instr.op = OpCode::NOOP; + instr.node = NoopNode{}; + break; + } + + case mlx_delegate::OpNode_IdCopyNode: { + auto fb = fb_instr->op_as_IdCopyNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + IdCopyNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ID_COPY; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AddmmNode: { + auto fb = fb_instr->op_as_AddmmNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AddmmNode node; + node.mat1 = convert_tid(fb->mat1()); + node.mat2 = convert_tid(fb->mat2()); + node.out = convert_tid(fb->out()); + if (fb->bias()) { + node.bias = convert_tid(fb->bias()); + } + node.alpha = fb->alpha(); + node.beta = fb->beta(); + instr.op = OpCode::ADDMM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ItemIntNode: { + auto fb = fb_instr->op_as_ItemIntNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ItemIntNode node; + node.x = convert_tid(fb->x()); + node.out = convert_vid(fb->out()); + instr.op = OpCode::ITEM_INT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ExpandDimsNode: { + auto fb = fb_instr->op_as_ExpandDimsNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ExpandDimsNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::EXPAND_DIMS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TileNode: { + auto fb = fb_instr->op_as_TileNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TileNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->reps()) { + for (size_t i = 0; i < fb->reps()->size(); ++i) { + node.reps.push_back(convert_int_or_vid(fb->reps()->Get(static_cast(i)))); + } + } + instr.op = OpCode::TILE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TakeAlongAxisNode: { + auto fb = fb_instr->op_as_TakeAlongAxisNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TakeAlongAxisNode node; + node.x = convert_tid(fb->x()); + node.indices = convert_tid(fb->indices()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::TAKE_ALONG_AXIS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TakeNode: { + auto fb = fb_instr->op_as_TakeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TakeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.index = convert_int_or_vid_or_tid(fb->index()); + node.axis = fb->axis(); + instr.op = OpCode::TAKE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RMSNormNode: { + auto fb = fb_instr->op_as_RMSNormNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RMSNormNode node; + node.x = convert_tid(fb->x()); + if (fb->weight()) { + node.weight = convert_tid(fb->weight()); + } + node.out = convert_tid(fb->out()); + node.eps = fb->eps(); + instr.op = OpCode::RMS_NORM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LayerNormNode: { + auto fb = fb_instr->op_as_LayerNormNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LayerNormNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->weight()) { + node.weight = convert_tid(fb->weight()); + } + if (fb->bias()) { + node.bias = convert_tid(fb->bias()); + } + node.eps = fb->eps(); + instr.op = OpCode::LAYER_NORM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RopeNode: { + auto fb = fb_instr->op_as_RopeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RopeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.dims = fb->dims(); + node.offset = convert_vid_or_tid(fb->offset()); + if (fb->freqs()) { + node.freqs = convert_tid(fb->freqs()); + } + node.traditional = fb->traditional(); + node.base = fb->base(); + node.scale = fb->scale(); + instr.op = OpCode::ROPE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SdpaNode: { + auto fb = fb_instr->op_as_SdpaNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SdpaNode node; + node.q = convert_tid(fb->q()); + node.k = convert_tid(fb->k()); + node.v = convert_tid(fb->v()); + node.out = convert_tid(fb->out()); + node.scale = fb->scale(); + if (fb->mask()) { + node.mask = convert_tid(fb->mask()); + } + node.causal = fb->causal(); + instr.op = OpCode::SDPA; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AddNode: { + auto fb = fb_instr->op_as_AddNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AddNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ADD; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AddIntNode: { + auto fb = fb_instr->op_as_AddIntNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AddIntNode node; + node.a = convert_int_or_vid(fb->a()); + node.b = convert_int_or_vid(fb->b()); + node.out = convert_vid(fb->out()); + instr.op = OpCode::ADD_INT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SubtractIntNode: { + auto fb = fb_instr->op_as_SubtractIntNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SubtractIntNode node; + node.a = convert_int_or_vid(fb->a()); + node.b = convert_int_or_vid(fb->b()); + node.out = convert_vid(fb->out()); + instr.op = OpCode::SUBTRACT_INT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MultiplyIntNode: { + auto fb = fb_instr->op_as_MultiplyIntNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MultiplyIntNode node; + node.a = convert_int_or_vid(fb->a()); + node.b = convert_int_or_vid(fb->b()); + node.out = convert_vid(fb->out()); + instr.op = OpCode::MULTIPLY_INT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_FloorDivideIntNode: { + auto fb = fb_instr->op_as_FloorDivideIntNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + FloorDivideIntNode node; + node.a = convert_int_or_vid(fb->a()); + node.b = convert_int_or_vid(fb->b()); + node.out = convert_vid(fb->out()); + instr.op = OpCode::FLOOR_DIVIDE_INT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ModIntNode: { + auto fb = fb_instr->op_as_ModIntNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ModIntNode node; + node.a = convert_int_or_vid(fb->a()); + node.b = convert_int_or_vid(fb->b()); + node.out = convert_vid(fb->out()); + instr.op = OpCode::MOD_INT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SymSizeNode: { + auto fb = fb_instr->op_as_SymSizeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SymSizeNode node; + node.a = convert_tid(fb->a()); + node.dim = fb->dim(); + node.out = convert_vid(fb->out()); + instr.op = OpCode::SYM_SIZE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MultiplyNode: { + auto fb = fb_instr->op_as_MultiplyNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MultiplyNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::MULTIPLY; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_DivideNode: { + auto fb = fb_instr->op_as_DivideNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + DivideNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::DIVIDE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SubtractNode: { + auto fb = fb_instr->op_as_SubtractNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SubtractNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SUBTRACT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Conv1DNode: { + auto fb = fb_instr->op_as_Conv1DNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Conv1DNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.out = convert_tid(fb->out()); + node.stride = fb->stride(); + node.padding = fb->padding(); + node.dilation = fb->dilation(); + node.groups = fb->groups(); + instr.op = OpCode::CONV1D; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Conv2DNode: { + auto fb = fb_instr->op_as_Conv2DNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Conv2DNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.out = convert_tid(fb->out()); + node.stride_h = fb->stride_h(); + node.stride_w = fb->stride_w(); + node.padding_h = fb->padding_h(); + node.padding_w = fb->padding_w(); + node.dilation_h = fb->dilation_h(); + node.dilation_w = fb->dilation_w(); + node.groups = fb->groups(); + instr.op = OpCode::CONV2D; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Conv3DNode: { + auto fb = fb_instr->op_as_Conv3DNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Conv3DNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.out = convert_tid(fb->out()); + node.stride_d = fb->stride_d(); + node.stride_h = fb->stride_h(); + node.stride_w = fb->stride_w(); + node.padding_d = fb->padding_d(); + node.padding_h = fb->padding_h(); + node.padding_w = fb->padding_w(); + node.dilation_d = fb->dilation_d(); + node.dilation_h = fb->dilation_h(); + node.dilation_w = fb->dilation_w(); + node.groups = fb->groups(); + instr.op = OpCode::CONV3D; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ConvTranspose1DNode: { + auto fb = fb_instr->op_as_ConvTranspose1DNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ConvTranspose1DNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.out = convert_tid(fb->out()); + node.stride = fb->stride(); + node.padding = fb->padding(); + node.dilation = fb->dilation(); + node.output_padding = fb->output_padding(); + node.groups = fb->groups(); + instr.op = OpCode::CONV_TRANSPOSE1D; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ConvTranspose2DNode: { + auto fb = fb_instr->op_as_ConvTranspose2DNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ConvTranspose2DNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.out = convert_tid(fb->out()); + node.stride_h = fb->stride_h(); + node.stride_w = fb->stride_w(); + node.padding_h = fb->padding_h(); + node.padding_w = fb->padding_w(); + node.dilation_h = fb->dilation_h(); + node.dilation_w = fb->dilation_w(); + node.output_padding_h = fb->output_padding_h(); + node.output_padding_w = fb->output_padding_w(); + node.groups = fb->groups(); + instr.op = OpCode::CONV_TRANSPOSE2D; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ConvTranspose3DNode: { + auto fb = fb_instr->op_as_ConvTranspose3DNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ConvTranspose3DNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.out = convert_tid(fb->out()); + node.stride_d = fb->stride_d(); + node.stride_h = fb->stride_h(); + node.stride_w = fb->stride_w(); + node.padding_d = fb->padding_d(); + node.padding_h = fb->padding_h(); + node.padding_w = fb->padding_w(); + node.dilation_d = fb->dilation_d(); + node.dilation_h = fb->dilation_h(); + node.dilation_w = fb->dilation_w(); + node.output_padding_d = fb->output_padding_d(); + node.output_padding_h = fb->output_padding_h(); + node.output_padding_w = fb->output_padding_w(); + node.groups = fb->groups(); + instr.op = OpCode::CONV_TRANSPOSE3D; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_GeluNode: { + auto fb = fb_instr->op_as_GeluNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + GeluNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.approximate = fb->approximate() ? fb->approximate()->str() : ""; + instr.op = OpCode::GELU; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ARangeNode: { + auto fb = fb_instr->op_as_ARangeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ARangeNode node; + node.out = convert_tid(fb->out()); + node.start = convert_int_or_vid(fb->start()); + node.stop = convert_int_or_vid(fb->stop()); + node.step = convert_int_or_vid(fb->step()); + auto scalar_type_opt = fb->scalar_type(); + if (scalar_type_opt.has_value()) { + node.scalar_type = scalar_type_opt.value(); + } + instr.op = OpCode::ARANGE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SiluNode: { + auto fb = fb_instr->op_as_SiluNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SiluNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SILU; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SigmoidNode: { + auto fb = fb_instr->op_as_SigmoidNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SigmoidNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SIGMOID; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TanhNode: { + auto fb = fb_instr->op_as_TanhNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TanhNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::TANH; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SqueezeNode: { + auto fb = fb_instr->op_as_SqueezeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SqueezeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.dims = to_vector(fb->dims()); + instr.op = OpCode::SQUEEZE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SplitNode: { + auto fb = fb_instr->op_as_SplitNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SplitNode node; + node.x = convert_tid(fb->x()); + if (fb->outs()) { + for (auto fb_tid : *fb->outs()) { + node.outs.push_back(convert_tid(fb_tid)); + } + } + if (fb->sizes()) { + for (size_t i = 0; i < fb->sizes()->size(); ++i) { + node.sizes.push_back(convert_int_or_vid(fb->sizes()->Get(static_cast(i)))); + } + } + node.axis = fb->axis(); + instr.op = OpCode::SPLIT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RsqrtNode: { + auto fb = fb_instr->op_as_RsqrtNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RsqrtNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::RSQRT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MaximumNode: { + auto fb = fb_instr->op_as_MaximumNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MaximumNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::MAXIMUM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MinimumNode: { + auto fb = fb_instr->op_as_MinimumNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MinimumNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::MINIMUM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LogNode: { + auto fb = fb_instr->op_as_LogNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LogNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOG; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SoftmaxNode: { + auto fb = fb_instr->op_as_SoftmaxNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SoftmaxNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + node.precise = fb->precise(); + instr.op = OpCode::SOFTMAX; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_BroadcastToNode: { + auto fb = fb_instr->op_as_BroadcastToNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + BroadcastToNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->shape()) { + for (size_t i = 0; i < fb->shape()->size(); ++i) { + node.shape.push_back(convert_int_or_vid(fb->shape()->Get(static_cast(i)))); + } + } + instr.op = OpCode::BROADCAST_TO; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_PadNode: { + auto fb = fb_instr->op_as_PadNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + PadNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->pad_width()) { + for (size_t i = 0; i < fb->pad_width()->size(); ++i) { + node.pad_width.push_back(convert_int_or_vid(fb->pad_width()->Get(static_cast(i)))); + } + } + node.mode = fb->mode() ? fb->mode()->str() : ""; + node.constant_value = fb->constant_value(); + instr.op = OpCode::PAD; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_WhereNode: { + auto fb = fb_instr->op_as_WhereNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + WhereNode node; + node.condition = convert_tid(fb->condition()); + node.x = convert_tid(fb->x()); + node.y = convert_tid(fb->y()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::WHERE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ReshapeNode: { + auto fb = fb_instr->op_as_ReshapeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ReshapeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->shape()) { + for (size_t i = 0; i < fb->shape()->size(); ++i) { + node.shape.push_back(convert_int_or_vid(fb->shape()->Get(static_cast(i)))); + } + } + instr.op = OpCode::RESHAPE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TransposeNode: { + auto fb = fb_instr->op_as_TransposeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TransposeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.perm = to_vector(fb->perm()); + instr.op = OpCode::TRANSPOSE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AsStridedNode: { + auto fb = fb_instr->op_as_AsStridedNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AsStridedNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->shape()) { + for (size_t i = 0; i < fb->shape()->size(); ++i) { + node.shape.push_back(convert_int_or_vid(fb->shape()->Get(static_cast(i)))); + } + } + if (fb->strides()) { + for (size_t i = 0; i < fb->strides()->size(); ++i) { + node.strides.push_back(convert_int_or_vid(fb->strides()->Get(static_cast(i)))); + } + } + node.offset = fb->offset(); + instr.op = OpCode::AS_STRIDED; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ContiguousNode: { + auto fb = fb_instr->op_as_ContiguousNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ContiguousNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::CONTIGUOUS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_GatherNode: { + auto fb = fb_instr->op_as_GatherNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + GatherNode node; + node.x = convert_tid(fb->x()); + if (fb->indices()) { + for (auto fb_tid : *fb->indices()) { + node.indices.push_back(convert_tid(fb_tid)); + } + } + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.slice_sizes = to_vector(fb->slice_sizes()); + instr.op = OpCode::GATHER; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SliceNode: { + auto fb = fb_instr->op_as_SliceNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SliceNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = convert_int_or_vid(fb->axis()); + node.start = convert_int_or_vid(fb->start()); + node.stop = convert_int_or_vid(fb->stop()); + node.step = fb->step(); + instr.op = OpCode::SLICE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AsTypeNode: { + auto fb = fb_instr->op_as_AsTypeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AsTypeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.scalar_type = fb->scalar_type(); + instr.op = OpCode::ASTYPE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_QuantizedMatmulNode: { + auto fb = fb_instr->op_as_QuantizedMatmulNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + QuantizedMatmulNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.scales = convert_tid(fb->scales()); + node.out = convert_tid(fb->out()); + if (fb->biases()) { + node.biases = convert_tid(fb->biases()); + } + node.group_size = fb->group_size(); + node.bits = fb->bits(); + node.mode = fb->mode() ? fb->mode()->str() : ""; + node.transpose = fb->transpose(); + instr.op = OpCode::QUANTIZED_MATMUL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ScatterAddNode: { + auto fb = fb_instr->op_as_ScatterAddNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ScatterAddNode node; + node.x = convert_tid(fb->x()); + node.indices = convert_tid(fb->indices()); + node.updates = convert_tid(fb->updates()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::SCATTER_ADD; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ConcatenateNode: { + auto fb = fb_instr->op_as_ConcatenateNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ConcatenateNode node; + if (fb->tensors()) { + for (auto fb_tid : *fb->tensors()) { + node.tensors.push_back(convert_tid(fb_tid)); + } + } + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::CONCATENATE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_FullNode: { + auto fb = fb_instr->op_as_FullNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + FullNode node; + node.out = convert_tid(fb->out()); + if (fb->shape()) { + for (size_t i = 0; i < fb->shape()->size(); ++i) { + node.shape.push_back(convert_int_or_vid(fb->shape()->Get(static_cast(i)))); + } + } + node.v = convert_float_or_vid(fb->v()); + node.scalar_type = fb->scalar_type(); + instr.op = OpCode::FULL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_FullLikeNode: { + auto fb = fb_instr->op_as_FullLikeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + FullLikeNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.v = convert_float_or_vid(fb->v()); + auto scalar_type_opt = fb->scalar_type(); + if (scalar_type_opt.has_value()) { + node.scalar_type = scalar_type_opt.value(); + } + instr.op = OpCode::FULL_LIKE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArgmaxNode: { + auto fb = fb_instr->op_as_ArgmaxNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArgmaxNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + node.keepdims = fb->keepdims(); + instr.op = OpCode::ARGMAX; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SliceUpdateNode: { + auto fb = fb_instr->op_as_SliceUpdateNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SliceUpdateNode node; + node.dst = convert_tid(fb->dst()); + node.update = convert_tid(fb->update()); + node.out = convert_tid(fb->out()); + node.axis = convert_int_or_vid(fb->axis()); + node.start = convert_int_or_vid(fb->start()); + node.stop = convert_int_or_vid(fb->stop()); + node.step = fb->step(); + instr.op = OpCode::SLICE_UPDATE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_IndexCopyNode: { + auto fb = fb_instr->op_as_IndexCopyNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + IndexCopyNode node; + node.dst = convert_tid(fb->dst()); + node.update = convert_tid(fb->update()); + node.indices = convert_tid(fb->indices()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::INDEX_COPY; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_DequantizeNode: { + auto fb = fb_instr->op_as_DequantizeNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + DequantizeNode node; + node.w = convert_tid(fb->w()); + node.scales = convert_tid(fb->scales()); + node.out = convert_tid(fb->out()); + if (fb->biases()) { + node.biases = convert_tid(fb->biases()); + } + node.group_size = fb->group_size(); + node.bits = fb->bits(); + node.mode = fb->mode() ? fb->mode()->str() : ""; + if (fb->global_scale()) { + node.global_scale = convert_tid(fb->global_scale()); + } + auto dtype_opt = fb->dtype(); + if (dtype_opt.has_value()) { + node.dtype = dtype_opt.value(); + } + instr.op = OpCode::DEQUANTIZE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LessNode: { + auto fb = fb_instr->op_as_LessNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LessNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LESS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LessEqualNode: { + auto fb = fb_instr->op_as_LessEqualNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LessEqualNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LESS_EQUAL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_GreaterNode: { + auto fb = fb_instr->op_as_GreaterNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + GreaterNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::GREATER; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_GreaterEqualNode: { + auto fb = fb_instr->op_as_GreaterEqualNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + GreaterEqualNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::GREATER_EQUAL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_EqualNode: { + auto fb = fb_instr->op_as_EqualNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + EqualNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::EQUAL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_NotEqualNode: { + auto fb = fb_instr->op_as_NotEqualNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + NotEqualNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::NOT_EQUAL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LogicalNotNode: { + auto fb = fb_instr->op_as_LogicalNotNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LogicalNotNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOGICAL_NOT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_BitwiseInvertNode: { + auto fb = fb_instr->op_as_BitwiseInvertNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + BitwiseInvertNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::BITWISE_INVERT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LogicalAndNode: { + auto fb = fb_instr->op_as_LogicalAndNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LogicalAndNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOGICAL_AND; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LogicalOrNode: { + auto fb = fb_instr->op_as_LogicalOrNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LogicalOrNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOGICAL_OR; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_BitwiseAndNode: { + auto fb = fb_instr->op_as_BitwiseAndNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + BitwiseAndNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::BITWISE_AND; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_BitwiseOrNode: { + auto fb = fb_instr->op_as_BitwiseOrNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + BitwiseOrNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::BITWISE_OR; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_BitwiseXorNode: { + auto fb = fb_instr->op_as_BitwiseXorNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + BitwiseXorNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::BITWISE_XOR; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TriNode: { + auto fb = fb_instr->op_as_TriNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TriNode node; + node.out = convert_tid(fb->out()); + node.n = convert_int_or_vid(fb->n()); + node.m = convert_int_or_vid(fb->m()); + node.k = fb->k(); + node.scalar_type = fb->scalar_type(); + instr.op = OpCode::TRI; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TrilNode: { + auto fb = fb_instr->op_as_TrilNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TrilNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.k = fb->k(); + instr.op = OpCode::TRIL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TriuNode: { + auto fb = fb_instr->op_as_TriuNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TriuNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.k = fb->k(); + instr.op = OpCode::TRIU; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ClipNode: { + auto fb = fb_instr->op_as_ClipNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ClipNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->a_min()) { + node.a_min = convert_tid(fb->a_min()); + } + if (fb->a_max()) { + node.a_max = convert_tid(fb->a_max()); + } + instr.op = OpCode::CLIP; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_CumsumNode: { + auto fb = fb_instr->op_as_CumsumNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + CumsumNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + node.reverse = fb->reverse(); + node.inclusive = fb->inclusive(); + instr.op = OpCode::CUMSUM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_StackNode: { + auto fb = fb_instr->op_as_StackNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + StackNode node; + if (fb->tensors()) { + for (auto fb_tid : *fb->tensors()) { + node.tensors.push_back(convert_tid(fb_tid)); + } + } + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::STACK; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SignNode: { + auto fb = fb_instr->op_as_SignNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SignNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SIGN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AnyNode: { + auto fb = fb_instr->op_as_AnyNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AnyNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::ANY; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AllNode: { + auto fb = fb_instr->op_as_AllNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AllNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::ALL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RepeatNode: { + auto fb = fb_instr->op_as_RepeatNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RepeatNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.repeats = convert_int_or_vid(fb->repeats()); + node.axis = fb->axis(); + instr.op = OpCode::REPEAT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SortNode: { + auto fb = fb_instr->op_as_SortNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SortNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::SORT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArgsortNode: { + auto fb = fb_instr->op_as_ArgsortNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArgsortNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + instr.op = OpCode::ARGSORT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_PartitionNode: { + auto fb = fb_instr->op_as_PartitionNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + PartitionNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.kth = convert_int_or_vid(fb->kth()); + node.axis = fb->axis(); + instr.op = OpCode::PARTITION; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArgPartitionNode: { + auto fb = fb_instr->op_as_ArgPartitionNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArgPartitionNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.kth = convert_int_or_vid(fb->kth()); + node.axis = fb->axis(); + instr.op = OpCode::ARG_PARTITION; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RollNode: { + auto fb = fb_instr->op_as_RollNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RollNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + if (fb->shift()) { + for (size_t i = 0; i < fb->shift()->size(); ++i) { + node.shift.push_back(convert_int_or_vid(fb->shift()->Get(static_cast(i)))); + } + } + node.axes = to_vector(fb->axes()); + instr.op = OpCode::ROLL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_FloorNode: { + auto fb = fb_instr->op_as_FloorNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + FloorNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::FLOOR; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_CeilNode: { + auto fb = fb_instr->op_as_CeilNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + CeilNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::CEIL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SquareNode: { + auto fb = fb_instr->op_as_SquareNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SquareNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SQUARE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ExpNode: { + auto fb = fb_instr->op_as_ExpNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ExpNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::EXP; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SinNode: { + auto fb = fb_instr->op_as_SinNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SinNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SIN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_CosNode: { + auto fb = fb_instr->op_as_CosNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + CosNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::COS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_TanNode: { + auto fb = fb_instr->op_as_TanNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + TanNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::TAN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArcsinNode: { + auto fb = fb_instr->op_as_ArcsinNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArcsinNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ARCSIN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArccosNode: { + auto fb = fb_instr->op_as_ArccosNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArccosNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ARCCOS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArctanNode: { + auto fb = fb_instr->op_as_ArctanNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArctanNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ARCTAN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SinhNode: { + auto fb = fb_instr->op_as_SinhNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SinhNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SINH; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_CoshNode: { + auto fb = fb_instr->op_as_CoshNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + CoshNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::COSH; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArcsinhNode: { + auto fb = fb_instr->op_as_ArcsinhNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArcsinhNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ARCSINH; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArccoshNode: { + auto fb = fb_instr->op_as_ArccoshNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArccoshNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ARCCOSH; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArctanhNode: { + auto fb = fb_instr->op_as_ArctanhNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArctanhNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ARCTANH; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Log2Node: { + auto fb = fb_instr->op_as_Log2Node(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Log2Node node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOG2; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Log10Node: { + auto fb = fb_instr->op_as_Log10Node(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Log10Node node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOG10; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Log1pNode: { + auto fb = fb_instr->op_as_Log1pNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Log1pNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOG1P; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ErfNode: { + auto fb = fb_instr->op_as_ErfNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ErfNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ERF; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Expm1Node: { + auto fb = fb_instr->op_as_Expm1Node(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Expm1Node node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::EXPM1; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RoundNode: { + auto fb = fb_instr->op_as_RoundNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RoundNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.decimals = fb->decimals(); + instr.op = OpCode::ROUND; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ReciprocalNode: { + auto fb = fb_instr->op_as_ReciprocalNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ReciprocalNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::RECIPROCAL; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SqrtNode: { + auto fb = fb_instr->op_as_SqrtNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SqrtNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::SQRT; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_AbsNode: { + auto fb = fb_instr->op_as_AbsNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + AbsNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ABS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_NegNode: { + auto fb = fb_instr->op_as_NegNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + NegNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::NEG; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_Atan2Node: { + auto fb = fb_instr->op_as_Atan2Node(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + Atan2Node node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::ATAN2; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LogAddExpNode: { + auto fb = fb_instr->op_as_LogAddExpNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LogAddExpNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::LOG_ADD_EXP; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_FloorDivideNode: { + auto fb = fb_instr->op_as_FloorDivideNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + FloorDivideNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::FLOOR_DIVIDE; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RemainderNode: { + auto fb = fb_instr->op_as_RemainderNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RemainderNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::REMAINDER; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_PowerNode: { + auto fb = fb_instr->op_as_PowerNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + PowerNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + instr.op = OpCode::POWER; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_LogSumExpNode: { + auto fb = fb_instr->op_as_LogSumExpNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + LogSumExpNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::LOG_SUM_EXP; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_SumNode: { + auto fb = fb_instr->op_as_SumNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + SumNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::SUM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MeanNode: { + auto fb = fb_instr->op_as_MeanNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MeanNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::MEAN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_VarNode: { + auto fb = fb_instr->op_as_VarNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + VarNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + node.ddof = fb->ddof(); + instr.op = OpCode::VAR; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_StdNode: { + auto fb = fb_instr->op_as_StdNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + StdNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + node.ddof = fb->ddof(); + instr.op = OpCode::STD; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ProdNode: { + auto fb = fb_instr->op_as_ProdNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ProdNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::PROD; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MaxNode: { + auto fb = fb_instr->op_as_MaxNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MaxNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::MAX; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MinNode: { + auto fb = fb_instr->op_as_MinNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MinNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::MIN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ArgminNode: { + auto fb = fb_instr->op_as_ArgminNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ArgminNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axis = fb->axis(); + node.keepdims = fb->keepdims(); + instr.op = OpCode::ARGMIN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MedianNode: { + auto fb = fb_instr->op_as_MedianNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MedianNode node; + node.x = convert_tid(fb->x()); + node.out = convert_tid(fb->out()); + node.axes = to_vector(fb->axes()); + node.keepdims = fb->keepdims(); + instr.op = OpCode::MEDIAN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_GatherMmNode: { + auto fb = fb_instr->op_as_GatherMmNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + GatherMmNode node; + node.a = convert_tid(fb->a()); + node.b = convert_tid(fb->b()); + node.out = convert_tid(fb->out()); + if (fb->lhs_indices()) { + node.lhs_indices = convert_tid(fb->lhs_indices()); + } + if (fb->rhs_indices()) { + node.rhs_indices = convert_tid(fb->rhs_indices()); + } + node.sorted_indices = fb->sorted_indices(); + instr.op = OpCode::GATHER_MM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_GatherQmmNode: { + auto fb = fb_instr->op_as_GatherQmmNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + GatherQmmNode node; + node.x = convert_tid(fb->x()); + node.w = convert_tid(fb->w()); + node.scales = convert_tid(fb->scales()); + node.out = convert_tid(fb->out()); + node.mode = fb->mode() ? fb->mode()->str() : ""; + if (fb->biases()) { + node.biases = convert_tid(fb->biases()); + } + if (fb->lhs_indices()) { + node.lhs_indices = convert_tid(fb->lhs_indices()); + } + if (fb->rhs_indices()) { + node.rhs_indices = convert_tid(fb->rhs_indices()); + } + node.transpose = fb->transpose(); + node.group_size = fb->group_size(); + node.bits = fb->bits(); + node.sorted_indices = fb->sorted_indices(); + instr.op = OpCode::GATHER_QMM; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_ScanNode: { + auto fb = fb_instr->op_as_ScanNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + ScanNode node; + if (fb->originals()) { + for (auto fb_tid : *fb->originals()) { + node.originals.push_back(convert_tid(fb_tid)); + } + } + if (fb->sliced()) { + for (auto fb_tid : *fb->sliced()) { + node.sliced.push_back(convert_tid(fb_tid)); + } + } + if (fb->outputs()) { + for (auto fb_tid : *fb->outputs()) { + node.outputs.push_back(convert_tid(fb_tid)); + } + } + if (fb->carry()) { + for (auto fb_tid : *fb->carry()) { + node.carry.push_back(convert_tid(fb_tid)); + } + } + node.body_chain_idx = fb->body_chain_idx(); + node.scan_axis = fb->scan_axis(); + instr.op = OpCode::SCAN; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_IfNode: { + auto fb = fb_instr->op_as_IfNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + IfNode node; + node.cond = convert_int_or_vid(fb->cond()); + node.then_chain_idx = fb->then_chain_idx(); + node.else_chain_idx = fb->else_chain_idx(); + instr.op = OpCode::IF; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_RandomBitsNode: { + auto fb = fb_instr->op_as_RandomBitsNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + RandomBitsNode node; + node.out = convert_tid(fb->out()); + if (fb->shape()) { + for (size_t i = 0; i < fb->shape()->size(); ++i) { + node.shape.push_back(convert_int_or_vid(fb->shape()->Get(static_cast(i)))); + } + } + if (fb->seed()) { + node.seed = convert_vid(fb->seed()); + } + node.width = fb->width(); + instr.op = OpCode::RANDOM_BITS; + instr.node = std::move(node); + break; + } + + case mlx_delegate::OpNode_MetalKernelNode: { + auto fb = fb_instr->op_as_MetalKernelNode(); + if (!fb) {{ + throw std::runtime_error("FlatBuffer op_type/payload mismatch for {class_name}"); + }} + MetalKernelNode node; + node.name = fb->name() ? fb->name()->str() : ""; + node.source = strpool.intern(fb->source()); + if (fb->inputs()) { + for (auto fb_tid : *fb->inputs()) { + node.inputs.push_back(convert_tid(fb_tid)); + } + } + if (fb->outputs()) { + for (auto fb_tid : *fb->outputs()) { + node.outputs.push_back(convert_tid(fb_tid)); + } + } + if (fb->grid()) { + for (size_t i = 0; i < fb->grid()->size(); ++i) { + node.grid.push_back(convert_int_or_vid(fb->grid()->Get(static_cast(i)))); + } + } + if (fb->threadgroup()) { + for (size_t i = 0; i < fb->threadgroup()->size(); ++i) { + node.threadgroup.push_back(convert_int_or_vid(fb->threadgroup()->Get(static_cast(i)))); + } + } + node.header = strpool.intern(fb->header()); + if (fb->input_names()) { + for (const auto* s : *fb->input_names()) { + node.input_names.push_back(s ? s->str() : std::string{}); + } + } + if (fb->output_names()) { + for (const auto* s : *fb->output_names()) { + node.output_names.push_back(s ? s->str() : std::string{}); + } + } + node.ensure_row_contiguous = fb->ensure_row_contiguous(); + node.atomic_outputs = fb->atomic_outputs(); + if (fb->output_shapes_flat()) { + for (size_t i = 0; i < fb->output_shapes_flat()->size(); ++i) { + node.output_shapes_flat.push_back(convert_int_or_vid(fb->output_shapes_flat()->Get(static_cast(i)))); + } + } + node.output_shape_lengths = to_vector(fb->output_shape_lengths()); + node.output_dtypes = to_vector(fb->output_dtypes()); + if (fb->template_arg_names()) { + for (const auto* s : *fb->template_arg_names()) { + node.template_arg_names.push_back(s ? s->str() : std::string{}); + } + } + node.template_arg_kinds = to_vector(fb->template_arg_kinds()); + node.template_arg_values = to_vector(fb->template_arg_values()); + auto init_value_opt = fb->init_value(); + if (init_value_opt.has_value()) { + node.init_value = init_value_opt.value(); + } + instr.op = OpCode::METAL_KERNEL; + instr.node = std::move(node); + break; + } + + default: + throw std::runtime_error( + "Unknown op_type in load_instruction: " + + std::to_string(static_cast(op_type)) + + ". The .pte was built with a newer schema than this binary. " + "Rebuild with the latest runtime."); + } + + return instr; +} + +// ============================================================================= +// load_program +// ============================================================================= + +MLXProgram load_program(const void* data, size_t size) { + MLXHeader header; + if (!parse_header(data, size, header)) { + throw std::runtime_error("Invalid MLX header"); + } + + // Defense-in-depth: parse_header already validates this, but guard the + // unsigned subtraction against underflow in case the call site ever changes. + if (header.data_offset <= kHeaderSize || header.data_offset > size) { + throw std::runtime_error("data_offset out of range"); + } + const uint8_t* fb_data = static_cast(data) + kHeaderSize; + size_t fb_size = header.data_offset - kHeaderSize; + + flatbuffers::Verifier verifier(fb_data, fb_size); + if (!mlx_delegate::VerifyMLXGraphBuffer(verifier)) { + throw std::runtime_error("Invalid FlatBuffer data"); + } + + const auto* fb_graph = mlx_delegate::GetMLXGraph(fb_data); + if (!fb_graph) { + throw std::runtime_error("Failed to parse MLXGraph"); + } + + MLXProgram program; + + if (fb_graph->version()) { + program.version = fb_graph->version()->str(); + } + + program.num_constant_tensors = fb_graph->num_constant_tensors(); + program.num_input_tensors = fb_graph->num_input_tensors(); + program.num_output_tensors = fb_graph->num_output_tensors(); + program.num_mutable_buffer_tensors = fb_graph->num_mutable_buffer_tensors(); + program.num_temp_tensors = fb_graph->num_temp_tensors(); + program.num_values = fb_graph->num_values(); + + // Cap all counts/collection sizes to prevent unbounded allocations from + // malformed FlatBuffer payloads + constexpr size_t kMaxCollectionSize = 1'000'000; + auto check_collection_size = [](size_t sz, const char* name) { + if (sz > kMaxCollectionSize) { + throw std::runtime_error( + std::string("Malformed program: ") + name + " size " + + std::to_string(sz) + " exceeds maximum of " + + std::to_string(kMaxCollectionSize)); + } + }; + + check_collection_size(program.num_tensors(), "num_tensors()"); + check_collection_size(program.num_values, "num_values"); + + // Pool shared across all chains so identical kernel source/header blobs are + // interned once for the whole program. + StringPool strpool; + + if (fb_graph->instruction_chains()) { + check_collection_size(fb_graph->instruction_chains()->size(), "instruction_chains"); + program.instruction_chains.reserve(fb_graph->instruction_chains()->size()); + for (size_t c = 0; c < fb_graph->instruction_chains()->size(); ++c) { + const auto* fb_chain = fb_graph->instruction_chains()->Get(static_cast(c)); + std::vector chain; + if (fb_chain && fb_chain->instructions()) { + check_collection_size(fb_chain->instructions()->size(), "instructions in chain"); + chain.reserve(fb_chain->instructions()->size()); + for (size_t i = 0; i < fb_chain->instructions()->size(); ++i) { + chain.push_back(load_instruction(fb_chain->instructions()->Get(static_cast(i)), strpool)); + } + } + program.instruction_chains.push_back(std::move(chain)); + } + } + + program.main_chain_idx = fb_graph->main_chain_idx(); + program.init_chain_idx = fb_graph->init_chain_idx(); + + // Validate chain indices against actual instruction_chains size. + if (program.main_chain_idx >= program.instruction_chains.size()) { + throw std::runtime_error( + "Invalid main_chain_idx " + + std::to_string(program.main_chain_idx) + + " (only " + std::to_string(program.instruction_chains.size()) + + " chains loaded)"); + } + if (program.init_chain_idx >= 0 && + static_cast(program.init_chain_idx) >= + program.instruction_chains.size()) { + throw std::runtime_error( + "Invalid init_chain_idx " + + std::to_string(program.init_chain_idx) + + " (only " + std::to_string(program.instruction_chains.size()) + + " chains loaded)"); + } + + if (fb_graph->input_map()) { + check_collection_size(fb_graph->input_map()->size(), "input_map"); + for (size_t i = 0; i < fb_graph->input_map()->size(); ++i) { + const auto* slot = fb_graph->input_map()->Get(static_cast(i)); + auto sv = convert_slot_variant(slot); + if (sv.slot_type == SlotType::TensorSlot && + sv.idx >= program.num_tensors()) { + throw std::runtime_error( + "input_map: slot index " + std::to_string(sv.idx) + + " exceeds num_tensors " + + std::to_string(program.num_tensors())); + } + program.input_map.push_back(sv); + } + } + + if (fb_graph->output_map()) { + check_collection_size(fb_graph->output_map()->size(), "output_map"); + for (size_t i = 0; i < fb_graph->output_map()->size(); ++i) { + const auto* slot = fb_graph->output_map()->Get(static_cast(i)); + auto sv = convert_slot_variant(slot); + if (sv.slot_type == SlotType::TensorSlot && + sv.idx >= program.num_tensors()) { + throw std::runtime_error( + "output_map: slot index " + std::to_string(sv.idx) + + " exceeds num_tensors " + + std::to_string(program.num_tensors())); + } + program.output_map.push_back(sv); + } + } + + if (fb_graph->mutable_buffer_map()) { + check_collection_size(fb_graph->mutable_buffer_map()->size(), "mutable_buffer_map"); + for (size_t i = 0; i < fb_graph->mutable_buffer_map()->size(); ++i) { + const auto* slot = fb_graph->mutable_buffer_map()->Get(static_cast(i)); + auto sv = convert_slot_variant(slot); + if (sv.slot_type == SlotType::TensorSlot && + sv.idx >= program.num_tensors()) { + throw std::runtime_error( + "mutable_buffer_map: slot index " + std::to_string(sv.idx) + + " exceeds num_tensors " + + std::to_string(program.num_tensors())); + } + program.mutable_buffer_map.push_back(sv); + } + } + + if (fb_graph->named_slots()) { + check_collection_size(fb_graph->named_slots()->size(), "named_slots"); + for (size_t i = 0; i < fb_graph->named_slots()->size(); ++i) { + const auto* fb_slot = fb_graph->named_slots()->Get(static_cast(i)); + if (!fb_slot || !fb_slot->name()) { + throw std::runtime_error( + "Malformed program: named_slot at index " + std::to_string(i) + + " is null or has null name"); + } + NamedSlot slot; + slot.name = fb_slot->name()->str(); + slot.slot = convert_slot_variant(fb_slot->slot()); + program.named_slots.push_back(std::move(slot)); + } + } + + if (fb_graph->tensor_meta()) { + check_collection_size(fb_graph->tensor_meta()->size(), "tensor_meta"); + for (size_t i = 0; i < fb_graph->tensor_meta()->size(); ++i) { + const auto* fb_meta = fb_graph->tensor_meta()->Get(static_cast(i)); + if (fb_meta) { + TensorMeta meta; + if (fb_meta->shape()) { + // Validate tensor rank against kTensorDimensionLimit to prevent + // stack overflows from unchecked rank + constexpr size_t kTensorDimensionLimit = 16; + if (fb_meta->shape()->size() > kTensorDimensionLimit) { + throw std::runtime_error( + "Tensor at index " + std::to_string(i) + + " has rank " + std::to_string(fb_meta->shape()->size()) + + " exceeding kTensorDimensionLimit (" + + std::to_string(kTensorDimensionLimit) + ")"); + } + for (size_t j = 0; j < fb_meta->shape()->size(); ++j) { + const auto* fb_dim = fb_meta->shape()->Get(static_cast(j)); + if (!fb_dim) { + throw std::runtime_error( + "Null ShapeDim at index " + std::to_string(j) + + " in tensor_meta " + std::to_string(i)); + } + ShapeDim dim; + dim.value = fb_dim->value(); + dim.min_value = fb_dim->min_value(); + dim.max_value = fb_dim->max_value(); + if (dim.value < -1) { + throw std::runtime_error( + "Invalid ShapeDim value " + std::to_string(dim.value) + + " at index " + std::to_string(j) + + " in tensor_meta " + std::to_string(i)); + } + if (dim.is_dynamic()) { + if (dim.min_value < 0) { + throw std::runtime_error( + "Invalid ShapeDim min_value " + std::to_string(dim.min_value) + + " at index " + std::to_string(j) + + " in tensor_meta " + std::to_string(i)); + } + if (dim.max_value != -1 && dim.max_value < dim.min_value) { + throw std::runtime_error( + "ShapeDim max_value " + std::to_string(dim.max_value) + + " < min_value " + std::to_string(dim.min_value) + + " at index " + std::to_string(j) + + " in tensor_meta " + std::to_string(i)); + } + } + meta.shape.push_back(dim); + } + } + auto raw_scalar_type = fb_meta->scalar_type(); + if (raw_scalar_type < 0 || + raw_scalar_type >= + static_cast(ScalarType::NumOptions)) { + throw std::runtime_error( + "Invalid scalar_type " + std::to_string(raw_scalar_type) + + " in tensor_meta at index " + std::to_string(i)); + } + meta.scalar_type = static_cast(raw_scalar_type); + if (fb_meta->dim_order()) { + meta.dim_order = to_vector(fb_meta->dim_order()); + } + program.tensor_meta.push_back(std::move(meta)); + } else { + program.tensor_meta.push_back(std::nullopt); + } + } + } + + return program; +} + +} // namespace loader +} // namespace mlx +} // namespace backends +} // namespace executorch diff --git a/backends/mlx/runtime/MLXLoader.h b/backends/mlx/runtime/MLXLoader.h new file mode 100644 index 00000000000..57fba11ec78 --- /dev/null +++ b/backends/mlx/runtime/MLXLoader.h @@ -0,0 +1,2743 @@ +// +// Copyright (c) Meta Platforms, Inc. and affiliates. +// All rights reserved. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. +// +// ============================================================================ +// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY +// ============================================================================ +// +// This file was generated from schema.fbs by the MLX delegate code generator. +// +// Source: backends/mlx/serialization/schema.fbs +// Generator: backends/mlx/serialization/generate.py +// +// To regenerate, run from the executorch root: +// python backends/mlx/serialization/generate.py +// +// ============================================================================ +// +// -*- c++ -*- + +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "schema_generated.h" + +// ExecuTorch scalar type for dtype representation +#include + +namespace executorch { +namespace backends { +namespace mlx { + +// ============================================================================= +// Core types matching the Python side +// ============================================================================= + +struct Tid { + uint32_t idx{}; +}; + +struct Vid { + uint32_t idx{}; +}; + +// ============================================================================= +// Tensor metadata +// ============================================================================= + +// Import ScalarType from ExecuTorch +using ScalarType = ::executorch::runtime::etensor::ScalarType; + +struct ShapeDim { + int32_t value{-1}; // Static dim (>= 0), or -1 for dynamic + int32_t min_value{0}; // Lower bound (when value == -1) + int32_t max_value{-1}; // Upper bound (-1 = unbounded, when value == -1) + + bool is_dynamic() const { return value < 0; } +}; + +struct TensorMeta { + std::vector shape; + ScalarType scalar_type{ScalarType::Float}; // ET ScalarType + std::vector dim_order; +}; + +// VidOrTid: either a scalar value (Vid) or a tensor (Tid) +struct VidOrTid { + Vid vid{}; + Tid tid{}; + bool is_vid{false}; // false = use tid, true = use vid +}; + +// IntOrVidOrTid: a literal int, a runtime Vid, or a tensor (Tid) +struct IntOrVidOrTid { + int64_t literal{0}; + Vid vid{}; + Tid tid{}; + uint8_t kind{0}; // 0 = literal int, 1 = vid, 2 = tid +}; + +// ============================================================================= +// Op node types (AUTO-GENERATED from schema.fbs) +// ============================================================================= + +struct NoopNode { +}; + +struct IdCopyNode { + Tid x; + Tid out; +}; + +struct AddmmNode { + Tid mat1; + Tid mat2; + Tid out; + std::optional bias; + float alpha; + float beta; +}; + +struct ItemIntNode { + Tid x; + Vid out; +}; + +struct ExpandDimsNode { + Tid x; + Tid out; + int32_t axis; +}; + +struct TileNode { + Tid x; + Tid out; + std::vector> reps; +}; + +struct TakeAlongAxisNode { + Tid x; + Tid indices; + Tid out; + int32_t axis; +}; + +struct TakeNode { + Tid x; + Tid out; + IntOrVidOrTid index; + int32_t axis; +}; + +struct RMSNormNode { + Tid x; + std::optional weight; + Tid out; + float eps; +}; + +struct LayerNormNode { + Tid x; + Tid out; + std::optional weight; + std::optional bias; + float eps; +}; + +struct RopeNode { + Tid x; + Tid out; + int32_t dims; + VidOrTid offset; + std::optional freqs; + bool traditional; + float base; + float scale; +}; + +struct SdpaNode { + Tid q; + Tid k; + Tid v; + Tid out; + float scale; + std::optional mask; + bool causal; +}; + +struct AddNode { + Tid a; + Tid b; + Tid out; +}; + +struct AddIntNode { + std::variant a; + std::variant b; + Vid out; +}; + +struct SubtractIntNode { + std::variant a; + std::variant b; + Vid out; +}; + +struct MultiplyIntNode { + std::variant a; + std::variant b; + Vid out; +}; + +struct FloorDivideIntNode { + std::variant a; + std::variant b; + Vid out; +}; + +struct ModIntNode { + std::variant a; + std::variant b; + Vid out; +}; + +struct SymSizeNode { + Tid a; + int32_t dim; + Vid out; +}; + +struct MultiplyNode { + Tid a; + Tid b; + Tid out; +}; + +struct DivideNode { + Tid a; + Tid b; + Tid out; +}; + +struct SubtractNode { + Tid a; + Tid b; + Tid out; +}; + +struct Conv1DNode { + Tid x; + Tid w; + Tid out; + int32_t stride; + int32_t padding; + int32_t dilation; + int32_t groups; +}; + +struct Conv2DNode { + Tid x; + Tid w; + Tid out; + int32_t stride_h; + int32_t stride_w; + int32_t padding_h; + int32_t padding_w; + int32_t dilation_h; + int32_t dilation_w; + int32_t groups; +}; + +struct Conv3DNode { + Tid x; + Tid w; + Tid out; + int32_t stride_d; + int32_t stride_h; + int32_t stride_w; + int32_t padding_d; + int32_t padding_h; + int32_t padding_w; + int32_t dilation_d; + int32_t dilation_h; + int32_t dilation_w; + int32_t groups; +}; + +struct ConvTranspose1DNode { + Tid x; + Tid w; + Tid out; + int32_t stride; + int32_t padding; + int32_t dilation; + int32_t output_padding; + int32_t groups; +}; + +struct ConvTranspose2DNode { + Tid x; + Tid w; + Tid out; + int32_t stride_h; + int32_t stride_w; + int32_t padding_h; + int32_t padding_w; + int32_t dilation_h; + int32_t dilation_w; + int32_t output_padding_h; + int32_t output_padding_w; + int32_t groups; +}; + +struct ConvTranspose3DNode { + Tid x; + Tid w; + Tid out; + int32_t stride_d; + int32_t stride_h; + int32_t stride_w; + int32_t padding_d; + int32_t padding_h; + int32_t padding_w; + int32_t dilation_d; + int32_t dilation_h; + int32_t dilation_w; + int32_t output_padding_d; + int32_t output_padding_h; + int32_t output_padding_w; + int32_t groups; +}; + +struct GeluNode { + Tid x; + Tid out; + std::string approximate; +}; + +struct ARangeNode { + Tid out; + std::variant start; + std::variant stop; + std::variant step; + std::optional scalar_type; +}; + +struct SiluNode { + Tid x; + Tid out; +}; + +struct SigmoidNode { + Tid x; + Tid out; +}; + +struct TanhNode { + Tid x; + Tid out; +}; + +struct SqueezeNode { + Tid x; + Tid out; + std::vector dims; +}; + +struct SplitNode { + Tid x; + std::vector outs; + std::vector> sizes; + int32_t axis; +}; + +struct RsqrtNode { + Tid x; + Tid out; +}; + +struct MaximumNode { + Tid a; + Tid b; + Tid out; +}; + +struct MinimumNode { + Tid a; + Tid b; + Tid out; +}; + +struct LogNode { + Tid x; + Tid out; +}; + +struct SoftmaxNode { + Tid x; + Tid out; + int32_t axis; + bool precise; +}; + +struct BroadcastToNode { + Tid x; + Tid out; + std::vector> shape; +}; + +struct PadNode { + Tid x; + Tid out; + std::vector> pad_width; + std::string mode; + float constant_value; +}; + +struct WhereNode { + Tid condition; + Tid x; + Tid y; + Tid out; +}; + +struct ReshapeNode { + Tid x; + Tid out; + std::vector> shape; +}; + +struct TransposeNode { + Tid x; + Tid out; + std::vector perm; +}; + +struct AsStridedNode { + Tid x; + Tid out; + std::vector> shape; + std::vector> strides; + uint64_t offset; +}; + +struct ContiguousNode { + Tid x; + Tid out; +}; + +struct GatherNode { + Tid x; + std::vector indices; + Tid out; + std::vector axes; + std::vector slice_sizes; +}; + +struct SliceNode { + Tid x; + Tid out; + std::variant axis; + std::variant start; + std::variant stop; + int32_t step; +}; + +struct AsTypeNode { + Tid x; + Tid out; + int8_t scalar_type; +}; + +struct QuantizedMatmulNode { + Tid x; + Tid w; + Tid scales; + Tid out; + std::optional biases; + int32_t group_size; + int32_t bits; + std::string mode; + bool transpose; +}; + +struct ScatterAddNode { + Tid x; + Tid indices; + Tid updates; + Tid out; + int32_t axis; +}; + +struct ConcatenateNode { + std::vector tensors; + Tid out; + int32_t axis; +}; + +struct FullNode { + Tid out; + std::vector> shape; + std::variant v; + int8_t scalar_type; +}; + +struct FullLikeNode { + Tid x; + Tid out; + std::variant v; + std::optional scalar_type; +}; + +struct ArgmaxNode { + Tid x; + Tid out; + int32_t axis; + bool keepdims; +}; + +struct SliceUpdateNode { + Tid dst; + Tid update; + Tid out; + std::variant axis; + std::variant start; + std::variant stop; + int32_t step; +}; + +struct IndexCopyNode { + Tid dst; + Tid update; + Tid indices; + Tid out; + int32_t axis; +}; + +struct DequantizeNode { + Tid w; + Tid scales; + Tid out; + std::optional biases; + int32_t group_size; + int32_t bits; + std::string mode; + std::optional global_scale; + std::optional dtype; +}; + +struct LessNode { + Tid a; + Tid b; + Tid out; +}; + +struct LessEqualNode { + Tid a; + Tid b; + Tid out; +}; + +struct GreaterNode { + Tid a; + Tid b; + Tid out; +}; + +struct GreaterEqualNode { + Tid a; + Tid b; + Tid out; +}; + +struct EqualNode { + Tid a; + Tid b; + Tid out; +}; + +struct NotEqualNode { + Tid a; + Tid b; + Tid out; +}; + +struct LogicalNotNode { + Tid x; + Tid out; +}; + +struct BitwiseInvertNode { + Tid x; + Tid out; +}; + +struct LogicalAndNode { + Tid a; + Tid b; + Tid out; +}; + +struct LogicalOrNode { + Tid a; + Tid b; + Tid out; +}; + +struct BitwiseAndNode { + Tid a; + Tid b; + Tid out; +}; + +struct BitwiseOrNode { + Tid a; + Tid b; + Tid out; +}; + +struct BitwiseXorNode { + Tid a; + Tid b; + Tid out; +}; + +struct TriNode { + Tid out; + std::variant n; + std::variant m; + int32_t k; + int8_t scalar_type; +}; + +struct TrilNode { + Tid x; + Tid out; + int32_t k; +}; + +struct TriuNode { + Tid x; + Tid out; + int32_t k; +}; + +struct ClipNode { + Tid x; + Tid out; + std::optional a_min; + std::optional a_max; +}; + +struct CumsumNode { + Tid x; + Tid out; + int32_t axis; + bool reverse; + bool inclusive; +}; + +struct StackNode { + std::vector tensors; + Tid out; + int32_t axis; +}; + +struct SignNode { + Tid x; + Tid out; +}; + +struct AnyNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct AllNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct RepeatNode { + Tid x; + Tid out; + std::variant repeats; + int32_t axis; +}; + +struct SortNode { + Tid x; + Tid out; + int32_t axis; +}; + +struct ArgsortNode { + Tid x; + Tid out; + int32_t axis; +}; + +struct PartitionNode { + Tid x; + Tid out; + std::variant kth; + int32_t axis; +}; + +struct ArgPartitionNode { + Tid x; + Tid out; + std::variant kth; + int32_t axis; +}; + +struct RollNode { + Tid x; + Tid out; + std::vector> shift; + std::vector axes; +}; + +struct FloorNode { + Tid x; + Tid out; +}; + +struct CeilNode { + Tid x; + Tid out; +}; + +struct SquareNode { + Tid x; + Tid out; +}; + +struct ExpNode { + Tid x; + Tid out; +}; + +struct SinNode { + Tid x; + Tid out; +}; + +struct CosNode { + Tid x; + Tid out; +}; + +struct TanNode { + Tid x; + Tid out; +}; + +struct ArcsinNode { + Tid x; + Tid out; +}; + +struct ArccosNode { + Tid x; + Tid out; +}; + +struct ArctanNode { + Tid x; + Tid out; +}; + +struct SinhNode { + Tid x; + Tid out; +}; + +struct CoshNode { + Tid x; + Tid out; +}; + +struct ArcsinhNode { + Tid x; + Tid out; +}; + +struct ArccoshNode { + Tid x; + Tid out; +}; + +struct ArctanhNode { + Tid x; + Tid out; +}; + +struct Log2Node { + Tid x; + Tid out; +}; + +struct Log10Node { + Tid x; + Tid out; +}; + +struct Log1pNode { + Tid x; + Tid out; +}; + +struct ErfNode { + Tid x; + Tid out; +}; + +struct Expm1Node { + Tid x; + Tid out; +}; + +struct RoundNode { + Tid x; + Tid out; + int32_t decimals; +}; + +struct ReciprocalNode { + Tid x; + Tid out; +}; + +struct SqrtNode { + Tid x; + Tid out; +}; + +struct AbsNode { + Tid x; + Tid out; +}; + +struct NegNode { + Tid x; + Tid out; +}; + +struct Atan2Node { + Tid a; + Tid b; + Tid out; +}; + +struct LogAddExpNode { + Tid a; + Tid b; + Tid out; +}; + +struct FloorDivideNode { + Tid a; + Tid b; + Tid out; +}; + +struct RemainderNode { + Tid a; + Tid b; + Tid out; +}; + +struct PowerNode { + Tid a; + Tid b; + Tid out; +}; + +struct LogSumExpNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct SumNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct MeanNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct VarNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; + int32_t ddof; +}; + +struct StdNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; + int32_t ddof; +}; + +struct ProdNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct MaxNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct MinNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct ArgminNode { + Tid x; + Tid out; + int32_t axis; + bool keepdims; +}; + +struct MedianNode { + Tid x; + Tid out; + std::vector axes; + bool keepdims; +}; + +struct GatherMmNode { + Tid a; + Tid b; + Tid out; + std::optional lhs_indices; + std::optional rhs_indices; + bool sorted_indices; +}; + +struct GatherQmmNode { + Tid x; + Tid w; + Tid scales; + Tid out; + std::string mode; + std::optional biases; + std::optional lhs_indices; + std::optional rhs_indices; + bool transpose; + int32_t group_size; + int32_t bits; + bool sorted_indices; +}; + +struct ScanNode { + std::vector originals; + std::vector sliced; + std::vector outputs; + std::vector carry; + int32_t body_chain_idx; + int32_t scan_axis; +}; + +struct IfNode { + std::variant cond; + uint32_t then_chain_idx; + uint32_t else_chain_idx; +}; + +struct RandomBitsNode { + Tid out; + std::vector> shape; + std::optional seed; + int32_t width; +}; + +struct MetalKernelNode { + std::string name; + std::shared_ptr source; + std::vector inputs; + std::vector outputs; + std::vector> grid; + std::vector> threadgroup; + std::shared_ptr header; + std::vector input_names; + std::vector output_names; + bool ensure_row_contiguous; + bool atomic_outputs; + std::vector> output_shapes_flat; + std::vector output_shape_lengths; + std::vector output_dtypes; + std::vector template_arg_names; + std::vector template_arg_kinds; + std::vector template_arg_values; + std::optional init_value; +}; + + +// ============================================================================= +// OpCode enum (AUTO-GENERATED from schema.fbs) +// ============================================================================= + +enum class OpCode : uint8_t { + NOOP, + ID_COPY, + ADDMM, + ITEM_INT, + EXPAND_DIMS, + TILE, + TAKE_ALONG_AXIS, + TAKE, + RMS_NORM, + LAYER_NORM, + ROPE, + SDPA, + ADD, + ADD_INT, + SUBTRACT_INT, + MULTIPLY_INT, + FLOOR_DIVIDE_INT, + MOD_INT, + SYM_SIZE, + MULTIPLY, + DIVIDE, + SUBTRACT, + CONV1D, + CONV2D, + CONV3D, + CONV_TRANSPOSE1D, + CONV_TRANSPOSE2D, + CONV_TRANSPOSE3D, + GELU, + ARANGE, + SILU, + SIGMOID, + TANH, + SQUEEZE, + SPLIT, + RSQRT, + MAXIMUM, + MINIMUM, + LOG, + SOFTMAX, + BROADCAST_TO, + PAD, + WHERE, + RESHAPE, + TRANSPOSE, + AS_STRIDED, + CONTIGUOUS, + GATHER, + SLICE, + ASTYPE, + QUANTIZED_MATMUL, + SCATTER_ADD, + CONCATENATE, + FULL, + FULL_LIKE, + ARGMAX, + SLICE_UPDATE, + INDEX_COPY, + DEQUANTIZE, + LESS, + LESS_EQUAL, + GREATER, + GREATER_EQUAL, + EQUAL, + NOT_EQUAL, + LOGICAL_NOT, + BITWISE_INVERT, + LOGICAL_AND, + LOGICAL_OR, + BITWISE_AND, + BITWISE_OR, + BITWISE_XOR, + TRI, + TRIL, + TRIU, + CLIP, + CUMSUM, + STACK, + SIGN, + ANY, + ALL, + REPEAT, + SORT, + ARGSORT, + PARTITION, + ARG_PARTITION, + ROLL, + FLOOR, + CEIL, + SQUARE, + EXP, + SIN, + COS, + TAN, + ARCSIN, + ARCCOS, + ARCTAN, + SINH, + COSH, + ARCSINH, + ARCCOSH, + ARCTANH, + LOG2, + LOG10, + LOG1P, + ERF, + EXPM1, + ROUND, + RECIPROCAL, + SQRT, + ABS, + NEG, + ATAN2, + LOG_ADD_EXP, + FLOOR_DIVIDE, + REMAINDER, + POWER, + LOG_SUM_EXP, + SUM, + MEAN, + VAR, + STD, + PROD, + MAX, + MIN, + ARGMIN, + MEDIAN, + GATHER_MM, + GATHER_QMM, + SCAN, + IF, + RANDOM_BITS, + METAL_KERNEL, +}; + +// OpCode to string conversion (for logging) +inline const char* op_name(OpCode op) { + switch (op) { + case OpCode::NOOP: + return "NOOP"; + case OpCode::ID_COPY: + return "ID_COPY"; + case OpCode::ADDMM: + return "ADDMM"; + case OpCode::ITEM_INT: + return "ITEM_INT"; + case OpCode::EXPAND_DIMS: + return "EXPAND_DIMS"; + case OpCode::TILE: + return "TILE"; + case OpCode::TAKE_ALONG_AXIS: + return "TAKE_ALONG_AXIS"; + case OpCode::TAKE: + return "TAKE"; + case OpCode::RMS_NORM: + return "RMS_NORM"; + case OpCode::LAYER_NORM: + return "LAYER_NORM"; + case OpCode::ROPE: + return "ROPE"; + case OpCode::SDPA: + return "SDPA"; + case OpCode::ADD: + return "ADD"; + case OpCode::ADD_INT: + return "ADD_INT"; + case OpCode::SUBTRACT_INT: + return "SUBTRACT_INT"; + case OpCode::MULTIPLY_INT: + return "MULTIPLY_INT"; + case OpCode::FLOOR_DIVIDE_INT: + return "FLOOR_DIVIDE_INT"; + case OpCode::MOD_INT: + return "MOD_INT"; + case OpCode::SYM_SIZE: + return "SYM_SIZE"; + case OpCode::MULTIPLY: + return "MULTIPLY"; + case OpCode::DIVIDE: + return "DIVIDE"; + case OpCode::SUBTRACT: + return "SUBTRACT"; + case OpCode::CONV1D: + return "CONV1D"; + case OpCode::CONV2D: + return "CONV2D"; + case OpCode::CONV3D: + return "CONV3D"; + case OpCode::CONV_TRANSPOSE1D: + return "CONV_TRANSPOSE1D"; + case OpCode::CONV_TRANSPOSE2D: + return "CONV_TRANSPOSE2D"; + case OpCode::CONV_TRANSPOSE3D: + return "CONV_TRANSPOSE3D"; + case OpCode::GELU: + return "GELU"; + case OpCode::ARANGE: + return "ARANGE"; + case OpCode::SILU: + return "SILU"; + case OpCode::SIGMOID: + return "SIGMOID"; + case OpCode::TANH: + return "TANH"; + case OpCode::SQUEEZE: + return "SQUEEZE"; + case OpCode::SPLIT: + return "SPLIT"; + case OpCode::RSQRT: + return "RSQRT"; + case OpCode::MAXIMUM: + return "MAXIMUM"; + case OpCode::MINIMUM: + return "MINIMUM"; + case OpCode::LOG: + return "LOG"; + case OpCode::SOFTMAX: + return "SOFTMAX"; + case OpCode::BROADCAST_TO: + return "BROADCAST_TO"; + case OpCode::PAD: + return "PAD"; + case OpCode::WHERE: + return "WHERE"; + case OpCode::RESHAPE: + return "RESHAPE"; + case OpCode::TRANSPOSE: + return "TRANSPOSE"; + case OpCode::AS_STRIDED: + return "AS_STRIDED"; + case OpCode::CONTIGUOUS: + return "CONTIGUOUS"; + case OpCode::GATHER: + return "GATHER"; + case OpCode::SLICE: + return "SLICE"; + case OpCode::ASTYPE: + return "ASTYPE"; + case OpCode::QUANTIZED_MATMUL: + return "QUANTIZED_MATMUL"; + case OpCode::SCATTER_ADD: + return "SCATTER_ADD"; + case OpCode::CONCATENATE: + return "CONCATENATE"; + case OpCode::FULL: + return "FULL"; + case OpCode::FULL_LIKE: + return "FULL_LIKE"; + case OpCode::ARGMAX: + return "ARGMAX"; + case OpCode::SLICE_UPDATE: + return "SLICE_UPDATE"; + case OpCode::INDEX_COPY: + return "INDEX_COPY"; + case OpCode::DEQUANTIZE: + return "DEQUANTIZE"; + case OpCode::LESS: + return "LESS"; + case OpCode::LESS_EQUAL: + return "LESS_EQUAL"; + case OpCode::GREATER: + return "GREATER"; + case OpCode::GREATER_EQUAL: + return "GREATER_EQUAL"; + case OpCode::EQUAL: + return "EQUAL"; + case OpCode::NOT_EQUAL: + return "NOT_EQUAL"; + case OpCode::LOGICAL_NOT: + return "LOGICAL_NOT"; + case OpCode::BITWISE_INVERT: + return "BITWISE_INVERT"; + case OpCode::LOGICAL_AND: + return "LOGICAL_AND"; + case OpCode::LOGICAL_OR: + return "LOGICAL_OR"; + case OpCode::BITWISE_AND: + return "BITWISE_AND"; + case OpCode::BITWISE_OR: + return "BITWISE_OR"; + case OpCode::BITWISE_XOR: + return "BITWISE_XOR"; + case OpCode::TRI: + return "TRI"; + case OpCode::TRIL: + return "TRIL"; + case OpCode::TRIU: + return "TRIU"; + case OpCode::CLIP: + return "CLIP"; + case OpCode::CUMSUM: + return "CUMSUM"; + case OpCode::STACK: + return "STACK"; + case OpCode::SIGN: + return "SIGN"; + case OpCode::ANY: + return "ANY"; + case OpCode::ALL: + return "ALL"; + case OpCode::REPEAT: + return "REPEAT"; + case OpCode::SORT: + return "SORT"; + case OpCode::ARGSORT: + return "ARGSORT"; + case OpCode::PARTITION: + return "PARTITION"; + case OpCode::ARG_PARTITION: + return "ARG_PARTITION"; + case OpCode::ROLL: + return "ROLL"; + case OpCode::FLOOR: + return "FLOOR"; + case OpCode::CEIL: + return "CEIL"; + case OpCode::SQUARE: + return "SQUARE"; + case OpCode::EXP: + return "EXP"; + case OpCode::SIN: + return "SIN"; + case OpCode::COS: + return "COS"; + case OpCode::TAN: + return "TAN"; + case OpCode::ARCSIN: + return "ARCSIN"; + case OpCode::ARCCOS: + return "ARCCOS"; + case OpCode::ARCTAN: + return "ARCTAN"; + case OpCode::SINH: + return "SINH"; + case OpCode::COSH: + return "COSH"; + case OpCode::ARCSINH: + return "ARCSINH"; + case OpCode::ARCCOSH: + return "ARCCOSH"; + case OpCode::ARCTANH: + return "ARCTANH"; + case OpCode::LOG2: + return "LOG2"; + case OpCode::LOG10: + return "LOG10"; + case OpCode::LOG1P: + return "LOG1P"; + case OpCode::ERF: + return "ERF"; + case OpCode::EXPM1: + return "EXPM1"; + case OpCode::ROUND: + return "ROUND"; + case OpCode::RECIPROCAL: + return "RECIPROCAL"; + case OpCode::SQRT: + return "SQRT"; + case OpCode::ABS: + return "ABS"; + case OpCode::NEG: + return "NEG"; + case OpCode::ATAN2: + return "ATAN2"; + case OpCode::LOG_ADD_EXP: + return "LOG_ADD_EXP"; + case OpCode::FLOOR_DIVIDE: + return "FLOOR_DIVIDE"; + case OpCode::REMAINDER: + return "REMAINDER"; + case OpCode::POWER: + return "POWER"; + case OpCode::LOG_SUM_EXP: + return "LOG_SUM_EXP"; + case OpCode::SUM: + return "SUM"; + case OpCode::MEAN: + return "MEAN"; + case OpCode::VAR: + return "VAR"; + case OpCode::STD: + return "STD"; + case OpCode::PROD: + return "PROD"; + case OpCode::MAX: + return "MAX"; + case OpCode::MIN: + return "MIN"; + case OpCode::ARGMIN: + return "ARGMIN"; + case OpCode::MEDIAN: + return "MEDIAN"; + case OpCode::GATHER_MM: + return "GATHER_MM"; + case OpCode::GATHER_QMM: + return "GATHER_QMM"; + case OpCode::SCAN: + return "SCAN"; + case OpCode::IF: + return "IF"; + case OpCode::RANDOM_BITS: + return "RANDOM_BITS"; + case OpCode::METAL_KERNEL: + return "METAL_KERNEL"; + } + return "UNKNOWN"; +} + +// ============================================================================= +// NodeVariant for type-erased op storage (AUTO-GENERATED) +// ============================================================================= + +using NodeVariant = std::variant< + NoopNode, + IdCopyNode, + AddmmNode, + ItemIntNode, + ExpandDimsNode, + TileNode, + TakeAlongAxisNode, + TakeNode, + RMSNormNode, + LayerNormNode, + RopeNode, + SdpaNode, + AddNode, + AddIntNode, + SubtractIntNode, + MultiplyIntNode, + FloorDivideIntNode, + ModIntNode, + SymSizeNode, + MultiplyNode, + DivideNode, + SubtractNode, + Conv1DNode, + Conv2DNode, + Conv3DNode, + ConvTranspose1DNode, + ConvTranspose2DNode, + ConvTranspose3DNode, + GeluNode, + ARangeNode, + SiluNode, + SigmoidNode, + TanhNode, + SqueezeNode, + SplitNode, + RsqrtNode, + MaximumNode, + MinimumNode, + LogNode, + SoftmaxNode, + BroadcastToNode, + PadNode, + WhereNode, + ReshapeNode, + TransposeNode, + AsStridedNode, + ContiguousNode, + GatherNode, + SliceNode, + AsTypeNode, + QuantizedMatmulNode, + ScatterAddNode, + ConcatenateNode, + FullNode, + FullLikeNode, + ArgmaxNode, + SliceUpdateNode, + IndexCopyNode, + DequantizeNode, + LessNode, + LessEqualNode, + GreaterNode, + GreaterEqualNode, + EqualNode, + NotEqualNode, + LogicalNotNode, + BitwiseInvertNode, + LogicalAndNode, + LogicalOrNode, + BitwiseAndNode, + BitwiseOrNode, + BitwiseXorNode, + TriNode, + TrilNode, + TriuNode, + ClipNode, + CumsumNode, + StackNode, + SignNode, + AnyNode, + AllNode, + RepeatNode, + SortNode, + ArgsortNode, + PartitionNode, + ArgPartitionNode, + RollNode, + FloorNode, + CeilNode, + SquareNode, + ExpNode, + SinNode, + CosNode, + TanNode, + ArcsinNode, + ArccosNode, + ArctanNode, + SinhNode, + CoshNode, + ArcsinhNode, + ArccoshNode, + ArctanhNode, + Log2Node, + Log10Node, + Log1pNode, + ErfNode, + Expm1Node, + RoundNode, + ReciprocalNode, + SqrtNode, + AbsNode, + NegNode, + Atan2Node, + LogAddExpNode, + FloorDivideNode, + RemainderNode, + PowerNode, + LogSumExpNode, + SumNode, + MeanNode, + VarNode, + StdNode, + ProdNode, + MaxNode, + MinNode, + ArgminNode, + MedianNode, + GatherMmNode, + GatherQmmNode, + ScanNode, + IfNode, + RandomBitsNode, + MetalKernelNode +>; + +// ============================================================================= +// Instruction +// ============================================================================= + +struct Instruction { + OpCode op{OpCode::NOOP}; + NodeVariant node; + + template + T& get() { + return std::get(node); + } + + template + const T& get() const { + return std::get(node); + } +}; + +// ============================================================================= +// for_each_tid - invokes cb(Tid) for every tensor id referenced by an +// instruction's node (AUTO-GENERATED from schema.fbs). +// +// NOTE: nested chain-index fields (ScanNode::body_chain_idx, +// IfNode::then_chain_idx/else_chain_idx) are NOT followed here; callers that +// need transitive coverage must recurse into those chains themselves. +// ============================================================================= + +template +inline void for_each_tid(const Instruction& instr, F&& cb) { + switch (instr.op) { + case OpCode::NOOP: { + break; + } + case OpCode::ID_COPY: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ADDMM: { + const auto& n = std::get(instr.node); + cb(n.mat1); + cb(n.mat2); + cb(n.out); + if (n.bias.has_value()) cb(*n.bias); + break; + } + case OpCode::ITEM_INT: { + const auto& n = std::get(instr.node); + cb(n.x); + break; + } + case OpCode::EXPAND_DIMS: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::TILE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::TAKE_ALONG_AXIS: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.indices); + cb(n.out); + break; + } + case OpCode::TAKE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + if (n.index.kind == 2) cb(n.index.tid); + break; + } + case OpCode::RMS_NORM: { + const auto& n = std::get(instr.node); + cb(n.x); + if (n.weight.has_value()) cb(*n.weight); + cb(n.out); + break; + } + case OpCode::LAYER_NORM: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + if (n.weight.has_value()) cb(*n.weight); + if (n.bias.has_value()) cb(*n.bias); + break; + } + case OpCode::ROPE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + if (!n.offset.is_vid) cb(n.offset.tid); + if (n.freqs.has_value()) cb(*n.freqs); + break; + } + case OpCode::SDPA: { + const auto& n = std::get(instr.node); + cb(n.q); + cb(n.k); + cb(n.v); + cb(n.out); + if (n.mask.has_value()) cb(*n.mask); + break; + } + case OpCode::ADD: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::ADD_INT: { + break; + } + case OpCode::SUBTRACT_INT: { + break; + } + case OpCode::MULTIPLY_INT: { + break; + } + case OpCode::FLOOR_DIVIDE_INT: { + break; + } + case OpCode::MOD_INT: { + break; + } + case OpCode::SYM_SIZE: { + const auto& n = std::get(instr.node); + cb(n.a); + break; + } + case OpCode::MULTIPLY: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::DIVIDE: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::SUBTRACT: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::CONV1D: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.out); + break; + } + case OpCode::CONV2D: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.out); + break; + } + case OpCode::CONV3D: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.out); + break; + } + case OpCode::CONV_TRANSPOSE1D: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.out); + break; + } + case OpCode::CONV_TRANSPOSE2D: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.out); + break; + } + case OpCode::CONV_TRANSPOSE3D: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.out); + break; + } + case OpCode::GELU: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARANGE: { + const auto& n = std::get(instr.node); + cb(n.out); + break; + } + case OpCode::SILU: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SIGMOID: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::TANH: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SQUEEZE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SPLIT: { + const auto& n = std::get(instr.node); + cb(n.x); + for (const auto& tid_elem : n.outs) cb(tid_elem); + break; + } + case OpCode::RSQRT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::MAXIMUM: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::MINIMUM: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::LOG: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SOFTMAX: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::BROADCAST_TO: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::PAD: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::WHERE: { + const auto& n = std::get(instr.node); + cb(n.condition); + cb(n.x); + cb(n.y); + cb(n.out); + break; + } + case OpCode::RESHAPE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::TRANSPOSE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::AS_STRIDED: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::CONTIGUOUS: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::GATHER: { + const auto& n = std::get(instr.node); + cb(n.x); + for (const auto& tid_elem : n.indices) cb(tid_elem); + cb(n.out); + break; + } + case OpCode::SLICE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ASTYPE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::QUANTIZED_MATMUL: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.scales); + cb(n.out); + if (n.biases.has_value()) cb(*n.biases); + break; + } + case OpCode::SCATTER_ADD: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.indices); + cb(n.updates); + cb(n.out); + break; + } + case OpCode::CONCATENATE: { + const auto& n = std::get(instr.node); + for (const auto& tid_elem : n.tensors) cb(tid_elem); + cb(n.out); + break; + } + case OpCode::FULL: { + const auto& n = std::get(instr.node); + cb(n.out); + break; + } + case OpCode::FULL_LIKE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARGMAX: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SLICE_UPDATE: { + const auto& n = std::get(instr.node); + cb(n.dst); + cb(n.update); + cb(n.out); + break; + } + case OpCode::INDEX_COPY: { + const auto& n = std::get(instr.node); + cb(n.dst); + cb(n.update); + cb(n.indices); + cb(n.out); + break; + } + case OpCode::DEQUANTIZE: { + const auto& n = std::get(instr.node); + cb(n.w); + cb(n.scales); + cb(n.out); + if (n.biases.has_value()) cb(*n.biases); + if (n.global_scale.has_value()) cb(*n.global_scale); + break; + } + case OpCode::LESS: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::LESS_EQUAL: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::GREATER: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::GREATER_EQUAL: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::EQUAL: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::NOT_EQUAL: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::LOGICAL_NOT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::BITWISE_INVERT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::LOGICAL_AND: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::LOGICAL_OR: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::BITWISE_AND: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::BITWISE_OR: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::BITWISE_XOR: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::TRI: { + const auto& n = std::get(instr.node); + cb(n.out); + break; + } + case OpCode::TRIL: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::TRIU: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::CLIP: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + if (n.a_min.has_value()) cb(*n.a_min); + if (n.a_max.has_value()) cb(*n.a_max); + break; + } + case OpCode::CUMSUM: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::STACK: { + const auto& n = std::get(instr.node); + for (const auto& tid_elem : n.tensors) cb(tid_elem); + cb(n.out); + break; + } + case OpCode::SIGN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ANY: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ALL: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::REPEAT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SORT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARGSORT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::PARTITION: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARG_PARTITION: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ROLL: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::FLOOR: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::CEIL: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SQUARE: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::EXP: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SIN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::COS: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::TAN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARCSIN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARCCOS: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARCTAN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SINH: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::COSH: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARCSINH: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARCCOSH: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARCTANH: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::LOG2: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::LOG10: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::LOG1P: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ERF: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::EXPM1: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ROUND: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::RECIPROCAL: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SQRT: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ABS: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::NEG: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ATAN2: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::LOG_ADD_EXP: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::FLOOR_DIVIDE: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::REMAINDER: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::POWER: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + break; + } + case OpCode::LOG_SUM_EXP: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::SUM: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::MEAN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::VAR: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::STD: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::PROD: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::MAX: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::MIN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::ARGMIN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::MEDIAN: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.out); + break; + } + case OpCode::GATHER_MM: { + const auto& n = std::get(instr.node); + cb(n.a); + cb(n.b); + cb(n.out); + if (n.lhs_indices.has_value()) cb(*n.lhs_indices); + if (n.rhs_indices.has_value()) cb(*n.rhs_indices); + break; + } + case OpCode::GATHER_QMM: { + const auto& n = std::get(instr.node); + cb(n.x); + cb(n.w); + cb(n.scales); + cb(n.out); + if (n.biases.has_value()) cb(*n.biases); + if (n.lhs_indices.has_value()) cb(*n.lhs_indices); + if (n.rhs_indices.has_value()) cb(*n.rhs_indices); + break; + } + case OpCode::SCAN: { + const auto& n = std::get(instr.node); + for (const auto& tid_elem : n.originals) cb(tid_elem); + for (const auto& tid_elem : n.sliced) cb(tid_elem); + for (const auto& tid_elem : n.outputs) cb(tid_elem); + for (const auto& tid_elem : n.carry) cb(tid_elem); + break; + } + case OpCode::IF: { + break; + } + case OpCode::RANDOM_BITS: { + const auto& n = std::get(instr.node); + cb(n.out); + break; + } + case OpCode::METAL_KERNEL: { + const auto& n = std::get(instr.node); + for (const auto& tid_elem : n.inputs) cb(tid_elem); + for (const auto& tid_elem : n.outputs) cb(tid_elem); + break; + } + default: + throw std::runtime_error( + std::string("for_each_tid: unhandled op ") + op_name(instr.op)); + } +} + +// ============================================================================= +// SlotVariant for I/O mapping +// ============================================================================= + +enum class SlotType : uint8_t { + TensorSlot = 0, + IntValueSlot = 1, + FloatValueSlot = 2, + BoolValueSlot = 3, +}; + +struct SlotVariant { + uint32_t idx; + SlotType slot_type; +}; + +// ============================================================================= +// Named slot (name -> slot mapping) +// ============================================================================= + +struct NamedSlot { + std::string name; + SlotVariant slot; +}; + +// ============================================================================= +// MLXProgram - the loaded program ready for execution +// ============================================================================= + +struct MLXProgram { + std::string version; + + // Tensor/value slot counts (in Tid assignment order) + uint32_t num_constant_tensors{0}; + uint32_t num_input_tensors{0}; + uint32_t num_output_tensors{0}; + uint32_t num_mutable_buffer_tensors{0}; + uint32_t num_temp_tensors{0}; + uint32_t num_values{0}; + + // Instruction chains + std::vector> instruction_chains; + uint32_t main_chain_idx{0}; + int32_t init_chain_idx{-1}; // -1 = no init chain + + // I/O mappings + std::vector input_map; + std::vector output_map; + std::vector mutable_buffer_map; + + // Name to slot lookup + std::vector named_slots; + + // Tensor metadata + std::vector> tensor_meta; + + // Helper methods + inline uint64_t num_tensors() const { + return static_cast(num_constant_tensors) + + num_input_tensors + num_output_tensors + + num_mutable_buffer_tensors + num_temp_tensors; + } + + inline bool is_constant_tensor(Tid id) const { + return id.idx < num_constant_tensors; + } + + inline size_t num_inputs() const { + return input_map.size(); + } + + inline size_t num_outputs() const { + return output_map.size(); + } +}; + +// ============================================================================= +// init_chain_references_mutable_buffer - returns true if the program's init +// chain references any mutable-buffer Tid, following nested SCAN/IF chains. +// +// Used to decide whether the handle's default mutable-buffer initialization can +// be safely skipped (e.g. when buffers are managed per-session). If the init +// chain reads/writes mutable buffers, skipping would be unsafe. +// ============================================================================= +inline bool init_chain_references_mutable_buffer(const MLXProgram& program) { + if (program.init_chain_idx < 0 || program.num_mutable_buffer_tensors == 0) { + return false; + } + // Tid layout: Constant -> Input -> Output -> MutableBuffer -> Temp. + const uint32_t output_end = program.num_constant_tensors + + program.num_input_tensors + program.num_output_tensors; + const uint32_t mutable_buffer_end = + output_end + program.num_mutable_buffer_tensors; + auto is_mutable_buffer = [&](Tid t) { + return t.idx >= output_end && t.idx < mutable_buffer_end; + }; + + const size_t num_chains = program.instruction_chains.size(); + std::vector visited(num_chains, false); + std::vector worklist; + worklist.push_back(static_cast(program.init_chain_idx)); + + while (!worklist.empty()) { + const uint32_t chain_idx = worklist.back(); + worklist.pop_back(); + if (chain_idx >= num_chains || visited[chain_idx]) { + continue; + } + visited[chain_idx] = true; + for (const auto& instr : program.instruction_chains[chain_idx]) { + bool hit = false; + for_each_tid(instr, [&](Tid t) { + if (is_mutable_buffer(t)) { + hit = true; + } + }); + if (hit) { + return true; + } + // Follow nested chains (not covered by for_each_tid). + if (instr.op == OpCode::SCAN) { + const auto& n = std::get(instr.node); + if (n.body_chain_idx >= 0) { + worklist.push_back(static_cast(n.body_chain_idx)); + } + } else if (instr.op == OpCode::IF) { + const auto& n = std::get(instr.node); + worklist.push_back(n.then_chain_idx); + worklist.push_back(n.else_chain_idx); + } + } + } + return false; +} + +// ============================================================================= +// FlatBuffer loading functions +// ============================================================================= + +namespace loader { + +// Convert FlatBuffer SlotType to our SlotType +inline SlotType convert_slot_type(mlx_delegate::SlotType fb_type) { + switch (fb_type) { + case mlx_delegate::SlotType_TensorSlot: + return SlotType::TensorSlot; + case mlx_delegate::SlotType_IntValueSlot: + return SlotType::IntValueSlot; + case mlx_delegate::SlotType_FloatValueSlot: + return SlotType::FloatValueSlot; + case mlx_delegate::SlotType_BoolValueSlot: + return SlotType::BoolValueSlot; + default: + throw std::runtime_error("Unknown SlotType: " + + std::to_string(static_cast(fb_type))); + } +} + +// Convert FlatBuffer Tid +inline Tid convert_tid(const mlx_delegate::Tid* fb_tid) { + if (!fb_tid) { + throw std::runtime_error("Null Tid in FlatBuffer"); + } + return Tid{fb_tid->idx()}; +} + +// Convert FlatBuffer Vid +inline Vid convert_vid(const mlx_delegate::Vid* fb_vid) { + if (!fb_vid) { + throw std::runtime_error("Null Vid in FlatBuffer"); + } + return Vid{fb_vid->idx()}; +} + +// Convert FlatBuffer IntOrVid +inline std::variant convert_int_or_vid( + const mlx_delegate::IntOrVid* fb) { + if (!fb) { + throw std::runtime_error("Null IntOrVid in FlatBuffer"); + } + if (!fb->is_vid()) { + return fb->literal(); + } + const auto* vid_ptr = fb->vid(); + if (!vid_ptr) { + throw std::runtime_error("IntOrVid has is_vid=true but vid pointer is null"); + } + return Vid{vid_ptr->idx()}; +} + +// Convert FlatBuffer FloatOrVid +inline std::variant convert_float_or_vid( + const mlx_delegate::FloatOrVid* fb) { + if (!fb) { + throw std::runtime_error("Null FloatOrVid in FlatBuffer"); + } + if (!fb->is_vid()) { + return fb->literal(); + } + const auto* vid_ptr = fb->vid(); + if (!vid_ptr) { + throw std::runtime_error("FloatOrVid has is_vid=true but vid pointer is null"); + } + return Vid{vid_ptr->idx()}; +} + +// Convert FlatBuffer VidOrTid (scalar value or tensor) +inline VidOrTid convert_vid_or_tid( + const mlx_delegate::VidOrTid* fb) { + if (!fb) { + throw std::runtime_error("Null VidOrTid in FlatBuffer"); + } + VidOrTid result; + result.is_vid = fb->is_vid(); + if (result.is_vid) { + if (!fb->vid()) { + throw std::runtime_error("VidOrTid has is_vid=true but vid pointer is null"); + } + result.vid = Vid{fb->vid()->idx()}; + } else { + if (!fb->tid()) { + throw std::runtime_error("VidOrTid has is_vid=false but tid pointer is null"); + } + result.tid = Tid{fb->tid()->idx()}; + } + return result; +} + +// Convert FlatBuffer IntOrVidOrTid (literal int, Vid, or Tid) +inline IntOrVidOrTid convert_int_or_vid_or_tid( + const mlx_delegate::IntOrVidOrTid* fb) { + if (!fb) { + throw std::runtime_error("Null IntOrVidOrTid in FlatBuffer"); + } + IntOrVidOrTid result; + result.kind = fb->kind(); + switch (result.kind) { + case 0: // literal int + result.literal = fb->literal(); + break; + case 1: { // Vid + const auto* vid_ptr = fb->vid(); + if (!vid_ptr) { + throw std::runtime_error( + "IntOrVidOrTid has kind=1 (Vid) but vid pointer is null"); + } + result.vid = Vid{vid_ptr->idx()}; + break; + } + case 2: { // Tid + const auto* tid_ptr = fb->tid(); + if (!tid_ptr) { + throw std::runtime_error( + "IntOrVidOrTid has kind=2 (Tid) but tid pointer is null"); + } + result.tid = Tid{tid_ptr->idx()}; + break; + } + default: + throw std::runtime_error( + "IntOrVidOrTid has invalid kind: " + std::to_string(result.kind)); + } + return result; +} + +// Convert FlatBuffer SlotVariant +inline SlotVariant convert_slot_variant(const mlx_delegate::SlotVariant* fb) { + if (!fb) { + throw std::runtime_error("Null SlotVariant in FlatBuffer"); + } + return SlotVariant{fb->idx(), convert_slot_type(fb->slot_type())}; +} + +// Interns FlatBuffer strings by pointer so identical kernel source/header +// blobs (deduplicated to a single offset by the serializer) share one +// std::string in memory. Buffers written without string sharing simply get +// one entry per node — correct, just not deduplicated. +struct StringPool { + std::unordered_map> map; + std::shared_ptr intern(const flatbuffers::String* s) { + if (!s) { + return nullptr; + } + auto& slot = map[static_cast(s)]; + if (!slot) { + slot = std::make_shared(s->str()); + } + return slot; + } +}; + +// Load an instruction from FlatBuffer +Instruction load_instruction( + const mlx_delegate::Instruction* fb_instr, StringPool& strpool); + +// Load the full MLXProgram from FlatBuffer data +MLXProgram load_program(const void* data, size_t size); + +} // namespace loader + +} // namespace mlx +} // namespace backends +} // namespace executorch diff --git a/backends/mlx/runtime/schema_generated.h b/backends/mlx/runtime/schema_generated.h new file mode 100644 index 00000000000..35e7489514c --- /dev/null +++ b/backends/mlx/runtime/schema_generated.h @@ -0,0 +1,14154 @@ +// automatically generated by the FlatBuffers compiler, do not modify + + +#ifndef FLATBUFFERS_GENERATED_SCHEMA_MLX_DELEGATE_H_ +#define FLATBUFFERS_GENERATED_SCHEMA_MLX_DELEGATE_H_ + +#include "flatbuffers/flatbuffers.h" + +// Ensure the included flatbuffers.h is the same version as when this file was +// generated, otherwise it may not be compatible. +static_assert(FLATBUFFERS_VERSION_MAJOR == 24 && + FLATBUFFERS_VERSION_MINOR == 3 && + FLATBUFFERS_VERSION_REVISION == 25, + "Non-compatible flatbuffers version included"); + +namespace mlx_delegate { + +struct Tid; + +struct Vid; + +struct IntOrVid; +struct IntOrVidBuilder; + +struct FloatOrVid; +struct FloatOrVidBuilder; + +struct VidOrTid; +struct VidOrTidBuilder; + +struct IntOrVidOrTid; +struct IntOrVidOrTidBuilder; + +struct NoopNode; +struct NoopNodeBuilder; + +struct IdCopyNode; +struct IdCopyNodeBuilder; + +struct AddmmNode; +struct AddmmNodeBuilder; + +struct ItemIntNode; +struct ItemIntNodeBuilder; + +struct ExpandDimsNode; +struct ExpandDimsNodeBuilder; + +struct TileNode; +struct TileNodeBuilder; + +struct TakeAlongAxisNode; +struct TakeAlongAxisNodeBuilder; + +struct TakeNode; +struct TakeNodeBuilder; + +struct RMSNormNode; +struct RMSNormNodeBuilder; + +struct LayerNormNode; +struct LayerNormNodeBuilder; + +struct RopeNode; +struct RopeNodeBuilder; + +struct SdpaNode; +struct SdpaNodeBuilder; + +struct AddNode; +struct AddNodeBuilder; + +struct AddIntNode; +struct AddIntNodeBuilder; + +struct SubtractIntNode; +struct SubtractIntNodeBuilder; + +struct MultiplyIntNode; +struct MultiplyIntNodeBuilder; + +struct FloorDivideIntNode; +struct FloorDivideIntNodeBuilder; + +struct ModIntNode; +struct ModIntNodeBuilder; + +struct SymSizeNode; +struct SymSizeNodeBuilder; + +struct MultiplyNode; +struct MultiplyNodeBuilder; + +struct DivideNode; +struct DivideNodeBuilder; + +struct SubtractNode; +struct SubtractNodeBuilder; + +struct Conv1DNode; +struct Conv1DNodeBuilder; + +struct Conv2DNode; +struct Conv2DNodeBuilder; + +struct Conv3DNode; +struct Conv3DNodeBuilder; + +struct ConvTranspose1DNode; +struct ConvTranspose1DNodeBuilder; + +struct ConvTranspose2DNode; +struct ConvTranspose2DNodeBuilder; + +struct ConvTranspose3DNode; +struct ConvTranspose3DNodeBuilder; + +struct GeluNode; +struct GeluNodeBuilder; + +struct ARangeNode; +struct ARangeNodeBuilder; + +struct SiluNode; +struct SiluNodeBuilder; + +struct SigmoidNode; +struct SigmoidNodeBuilder; + +struct TanhNode; +struct TanhNodeBuilder; + +struct SqueezeNode; +struct SqueezeNodeBuilder; + +struct SplitNode; +struct SplitNodeBuilder; + +struct RsqrtNode; +struct RsqrtNodeBuilder; + +struct MaximumNode; +struct MaximumNodeBuilder; + +struct MinimumNode; +struct MinimumNodeBuilder; + +struct LogNode; +struct LogNodeBuilder; + +struct SoftmaxNode; +struct SoftmaxNodeBuilder; + +struct BroadcastToNode; +struct BroadcastToNodeBuilder; + +struct PadNode; +struct PadNodeBuilder; + +struct WhereNode; +struct WhereNodeBuilder; + +struct ReshapeNode; +struct ReshapeNodeBuilder; + +struct TransposeNode; +struct TransposeNodeBuilder; + +struct AsStridedNode; +struct AsStridedNodeBuilder; + +struct ContiguousNode; +struct ContiguousNodeBuilder; + +struct GatherNode; +struct GatherNodeBuilder; + +struct SliceNode; +struct SliceNodeBuilder; + +struct AsTypeNode; +struct AsTypeNodeBuilder; + +struct QuantizedMatmulNode; +struct QuantizedMatmulNodeBuilder; + +struct ScatterAddNode; +struct ScatterAddNodeBuilder; + +struct ConcatenateNode; +struct ConcatenateNodeBuilder; + +struct FullNode; +struct FullNodeBuilder; + +struct FullLikeNode; +struct FullLikeNodeBuilder; + +struct ArgmaxNode; +struct ArgmaxNodeBuilder; + +struct SliceUpdateNode; +struct SliceUpdateNodeBuilder; + +struct IndexCopyNode; +struct IndexCopyNodeBuilder; + +struct DequantizeNode; +struct DequantizeNodeBuilder; + +struct LessNode; +struct LessNodeBuilder; + +struct LessEqualNode; +struct LessEqualNodeBuilder; + +struct GreaterNode; +struct GreaterNodeBuilder; + +struct GreaterEqualNode; +struct GreaterEqualNodeBuilder; + +struct EqualNode; +struct EqualNodeBuilder; + +struct NotEqualNode; +struct NotEqualNodeBuilder; + +struct LogicalNotNode; +struct LogicalNotNodeBuilder; + +struct BitwiseInvertNode; +struct BitwiseInvertNodeBuilder; + +struct LogicalAndNode; +struct LogicalAndNodeBuilder; + +struct LogicalOrNode; +struct LogicalOrNodeBuilder; + +struct BitwiseAndNode; +struct BitwiseAndNodeBuilder; + +struct BitwiseOrNode; +struct BitwiseOrNodeBuilder; + +struct BitwiseXorNode; +struct BitwiseXorNodeBuilder; + +struct TriNode; +struct TriNodeBuilder; + +struct TrilNode; +struct TrilNodeBuilder; + +struct TriuNode; +struct TriuNodeBuilder; + +struct ClipNode; +struct ClipNodeBuilder; + +struct CumsumNode; +struct CumsumNodeBuilder; + +struct StackNode; +struct StackNodeBuilder; + +struct SignNode; +struct SignNodeBuilder; + +struct AnyNode; +struct AnyNodeBuilder; + +struct AllNode; +struct AllNodeBuilder; + +struct RepeatNode; +struct RepeatNodeBuilder; + +struct SortNode; +struct SortNodeBuilder; + +struct ArgsortNode; +struct ArgsortNodeBuilder; + +struct PartitionNode; +struct PartitionNodeBuilder; + +struct ArgPartitionNode; +struct ArgPartitionNodeBuilder; + +struct RollNode; +struct RollNodeBuilder; + +struct FloorNode; +struct FloorNodeBuilder; + +struct CeilNode; +struct CeilNodeBuilder; + +struct SquareNode; +struct SquareNodeBuilder; + +struct ExpNode; +struct ExpNodeBuilder; + +struct SinNode; +struct SinNodeBuilder; + +struct CosNode; +struct CosNodeBuilder; + +struct TanNode; +struct TanNodeBuilder; + +struct ArcsinNode; +struct ArcsinNodeBuilder; + +struct ArccosNode; +struct ArccosNodeBuilder; + +struct ArctanNode; +struct ArctanNodeBuilder; + +struct SinhNode; +struct SinhNodeBuilder; + +struct CoshNode; +struct CoshNodeBuilder; + +struct ArcsinhNode; +struct ArcsinhNodeBuilder; + +struct ArccoshNode; +struct ArccoshNodeBuilder; + +struct ArctanhNode; +struct ArctanhNodeBuilder; + +struct Log2Node; +struct Log2NodeBuilder; + +struct Log10Node; +struct Log10NodeBuilder; + +struct Log1pNode; +struct Log1pNodeBuilder; + +struct ErfNode; +struct ErfNodeBuilder; + +struct Expm1Node; +struct Expm1NodeBuilder; + +struct RoundNode; +struct RoundNodeBuilder; + +struct ReciprocalNode; +struct ReciprocalNodeBuilder; + +struct SqrtNode; +struct SqrtNodeBuilder; + +struct AbsNode; +struct AbsNodeBuilder; + +struct NegNode; +struct NegNodeBuilder; + +struct Atan2Node; +struct Atan2NodeBuilder; + +struct LogAddExpNode; +struct LogAddExpNodeBuilder; + +struct FloorDivideNode; +struct FloorDivideNodeBuilder; + +struct RemainderNode; +struct RemainderNodeBuilder; + +struct PowerNode; +struct PowerNodeBuilder; + +struct LogSumExpNode; +struct LogSumExpNodeBuilder; + +struct SumNode; +struct SumNodeBuilder; + +struct MeanNode; +struct MeanNodeBuilder; + +struct VarNode; +struct VarNodeBuilder; + +struct StdNode; +struct StdNodeBuilder; + +struct ProdNode; +struct ProdNodeBuilder; + +struct MaxNode; +struct MaxNodeBuilder; + +struct MinNode; +struct MinNodeBuilder; + +struct ArgminNode; +struct ArgminNodeBuilder; + +struct MedianNode; +struct MedianNodeBuilder; + +struct GatherMmNode; +struct GatherMmNodeBuilder; + +struct GatherQmmNode; +struct GatherQmmNodeBuilder; + +struct ScanNode; +struct ScanNodeBuilder; + +struct IfNode; +struct IfNodeBuilder; + +struct RandomBitsNode; +struct RandomBitsNodeBuilder; + +struct MetalKernelNode; +struct MetalKernelNodeBuilder; + +struct Instruction; +struct InstructionBuilder; + +struct InstructionChain; +struct InstructionChainBuilder; + +struct ShapeDim; +struct ShapeDimBuilder; + +struct TensorMeta; +struct TensorMetaBuilder; + +struct SlotVariant; +struct SlotVariantBuilder; + +struct NamedSlot; +struct NamedSlotBuilder; + +struct MLXGraph; +struct MLXGraphBuilder; + +enum OpNode : uint8_t { + OpNode_NONE = 0, + OpNode_NoopNode = 1, + OpNode_IdCopyNode = 2, + OpNode_AddmmNode = 3, + OpNode_ItemIntNode = 4, + OpNode_ExpandDimsNode = 5, + OpNode_TileNode = 6, + OpNode_TakeAlongAxisNode = 7, + OpNode_TakeNode = 8, + OpNode_RMSNormNode = 9, + OpNode_LayerNormNode = 10, + OpNode_RopeNode = 11, + OpNode_SdpaNode = 12, + OpNode_AddNode = 13, + OpNode_AddIntNode = 14, + OpNode_SubtractIntNode = 15, + OpNode_MultiplyIntNode = 16, + OpNode_FloorDivideIntNode = 17, + OpNode_SymSizeNode = 18, + OpNode_MultiplyNode = 19, + OpNode_DivideNode = 20, + OpNode_SubtractNode = 21, + OpNode_Conv1DNode = 22, + OpNode_Conv2DNode = 23, + OpNode_Conv3DNode = 24, + OpNode_GeluNode = 25, + OpNode_ARangeNode = 26, + OpNode_SiluNode = 27, + OpNode_SigmoidNode = 28, + OpNode_TanhNode = 29, + OpNode_SqueezeNode = 30, + OpNode_SplitNode = 31, + OpNode_RsqrtNode = 32, + OpNode_MaximumNode = 33, + OpNode_MinimumNode = 34, + OpNode_LogNode = 35, + OpNode_SoftmaxNode = 36, + OpNode_BroadcastToNode = 37, + OpNode_PadNode = 38, + OpNode_WhereNode = 39, + OpNode_ReshapeNode = 40, + OpNode_TransposeNode = 41, + OpNode_AsStridedNode = 42, + OpNode_ContiguousNode = 43, + OpNode_GatherNode = 44, + OpNode_SliceNode = 45, + OpNode_AsTypeNode = 46, + OpNode_ConcatenateNode = 47, + OpNode_FullNode = 48, + OpNode_FullLikeNode = 49, + OpNode_ArgmaxNode = 50, + OpNode_SliceUpdateNode = 51, + OpNode_IndexCopyNode = 52, + OpNode_DequantizeNode = 53, + OpNode_LessNode = 54, + OpNode_LessEqualNode = 55, + OpNode_GreaterNode = 56, + OpNode_GreaterEqualNode = 57, + OpNode_EqualNode = 58, + OpNode_NotEqualNode = 59, + OpNode_LogicalNotNode = 60, + OpNode_LogicalAndNode = 61, + OpNode_LogicalOrNode = 62, + OpNode_TriNode = 63, + OpNode_TrilNode = 64, + OpNode_TriuNode = 65, + OpNode_FloorNode = 66, + OpNode_CeilNode = 67, + OpNode_SquareNode = 68, + OpNode_ExpNode = 69, + OpNode_SinNode = 70, + OpNode_CosNode = 71, + OpNode_TanNode = 72, + OpNode_ArcsinNode = 73, + OpNode_ArccosNode = 74, + OpNode_ArctanNode = 75, + OpNode_SinhNode = 76, + OpNode_CoshNode = 77, + OpNode_ArcsinhNode = 78, + OpNode_ArccoshNode = 79, + OpNode_ArctanhNode = 80, + OpNode_Log2Node = 81, + OpNode_Log10Node = 82, + OpNode_Log1pNode = 83, + OpNode_ErfNode = 84, + OpNode_Expm1Node = 85, + OpNode_RoundNode = 86, + OpNode_ReciprocalNode = 87, + OpNode_SqrtNode = 88, + OpNode_AbsNode = 89, + OpNode_NegNode = 90, + OpNode_Atan2Node = 91, + OpNode_LogAddExpNode = 92, + OpNode_FloorDivideNode = 93, + OpNode_PowerNode = 94, + OpNode_LogSumExpNode = 95, + OpNode_SumNode = 96, + OpNode_MeanNode = 97, + OpNode_VarNode = 98, + OpNode_StdNode = 99, + OpNode_ProdNode = 100, + OpNode_MaxNode = 101, + OpNode_MinNode = 102, + OpNode_ArgminNode = 103, + OpNode_MedianNode = 104, + OpNode_ModIntNode = 105, + OpNode_RemainderNode = 106, + OpNode_ConvTranspose1DNode = 107, + OpNode_ConvTranspose2DNode = 108, + OpNode_ConvTranspose3DNode = 109, + OpNode_ClipNode = 110, + OpNode_CumsumNode = 111, + OpNode_StackNode = 112, + OpNode_SignNode = 113, + OpNode_AnyNode = 114, + OpNode_AllNode = 115, + OpNode_RepeatNode = 116, + OpNode_SortNode = 117, + OpNode_ArgsortNode = 118, + OpNode_PartitionNode = 119, + OpNode_ArgPartitionNode = 120, + OpNode_QuantizedMatmulNode = 121, + OpNode_ScatterAddNode = 122, + OpNode_GatherMmNode = 123, + OpNode_GatherQmmNode = 124, + OpNode_ScanNode = 125, + OpNode_MetalKernelNode = 126, + OpNode_BitwiseInvertNode = 127, + OpNode_RollNode = 128, + OpNode_BitwiseAndNode = 129, + OpNode_BitwiseOrNode = 130, + OpNode_BitwiseXorNode = 131, + OpNode_IfNode = 132, + OpNode_RandomBitsNode = 133, + OpNode_MIN = OpNode_NONE, + OpNode_MAX = OpNode_RandomBitsNode +}; + +inline const OpNode (&EnumValuesOpNode())[134] { + static const OpNode values[] = { + OpNode_NONE, + OpNode_NoopNode, + OpNode_IdCopyNode, + OpNode_AddmmNode, + OpNode_ItemIntNode, + OpNode_ExpandDimsNode, + OpNode_TileNode, + OpNode_TakeAlongAxisNode, + OpNode_TakeNode, + OpNode_RMSNormNode, + OpNode_LayerNormNode, + OpNode_RopeNode, + OpNode_SdpaNode, + OpNode_AddNode, + OpNode_AddIntNode, + OpNode_SubtractIntNode, + OpNode_MultiplyIntNode, + OpNode_FloorDivideIntNode, + OpNode_SymSizeNode, + OpNode_MultiplyNode, + OpNode_DivideNode, + OpNode_SubtractNode, + OpNode_Conv1DNode, + OpNode_Conv2DNode, + OpNode_Conv3DNode, + OpNode_GeluNode, + OpNode_ARangeNode, + OpNode_SiluNode, + OpNode_SigmoidNode, + OpNode_TanhNode, + OpNode_SqueezeNode, + OpNode_SplitNode, + OpNode_RsqrtNode, + OpNode_MaximumNode, + OpNode_MinimumNode, + OpNode_LogNode, + OpNode_SoftmaxNode, + OpNode_BroadcastToNode, + OpNode_PadNode, + OpNode_WhereNode, + OpNode_ReshapeNode, + OpNode_TransposeNode, + OpNode_AsStridedNode, + OpNode_ContiguousNode, + OpNode_GatherNode, + OpNode_SliceNode, + OpNode_AsTypeNode, + OpNode_ConcatenateNode, + OpNode_FullNode, + OpNode_FullLikeNode, + OpNode_ArgmaxNode, + OpNode_SliceUpdateNode, + OpNode_IndexCopyNode, + OpNode_DequantizeNode, + OpNode_LessNode, + OpNode_LessEqualNode, + OpNode_GreaterNode, + OpNode_GreaterEqualNode, + OpNode_EqualNode, + OpNode_NotEqualNode, + OpNode_LogicalNotNode, + OpNode_LogicalAndNode, + OpNode_LogicalOrNode, + OpNode_TriNode, + OpNode_TrilNode, + OpNode_TriuNode, + OpNode_FloorNode, + OpNode_CeilNode, + OpNode_SquareNode, + OpNode_ExpNode, + OpNode_SinNode, + OpNode_CosNode, + OpNode_TanNode, + OpNode_ArcsinNode, + OpNode_ArccosNode, + OpNode_ArctanNode, + OpNode_SinhNode, + OpNode_CoshNode, + OpNode_ArcsinhNode, + OpNode_ArccoshNode, + OpNode_ArctanhNode, + OpNode_Log2Node, + OpNode_Log10Node, + OpNode_Log1pNode, + OpNode_ErfNode, + OpNode_Expm1Node, + OpNode_RoundNode, + OpNode_ReciprocalNode, + OpNode_SqrtNode, + OpNode_AbsNode, + OpNode_NegNode, + OpNode_Atan2Node, + OpNode_LogAddExpNode, + OpNode_FloorDivideNode, + OpNode_PowerNode, + OpNode_LogSumExpNode, + OpNode_SumNode, + OpNode_MeanNode, + OpNode_VarNode, + OpNode_StdNode, + OpNode_ProdNode, + OpNode_MaxNode, + OpNode_MinNode, + OpNode_ArgminNode, + OpNode_MedianNode, + OpNode_ModIntNode, + OpNode_RemainderNode, + OpNode_ConvTranspose1DNode, + OpNode_ConvTranspose2DNode, + OpNode_ConvTranspose3DNode, + OpNode_ClipNode, + OpNode_CumsumNode, + OpNode_StackNode, + OpNode_SignNode, + OpNode_AnyNode, + OpNode_AllNode, + OpNode_RepeatNode, + OpNode_SortNode, + OpNode_ArgsortNode, + OpNode_PartitionNode, + OpNode_ArgPartitionNode, + OpNode_QuantizedMatmulNode, + OpNode_ScatterAddNode, + OpNode_GatherMmNode, + OpNode_GatherQmmNode, + OpNode_ScanNode, + OpNode_MetalKernelNode, + OpNode_BitwiseInvertNode, + OpNode_RollNode, + OpNode_BitwiseAndNode, + OpNode_BitwiseOrNode, + OpNode_BitwiseXorNode, + OpNode_IfNode, + OpNode_RandomBitsNode + }; + return values; +} + +inline const char * const *EnumNamesOpNode() { + static const char * const names[135] = { + "NONE", + "NoopNode", + "IdCopyNode", + "AddmmNode", + "ItemIntNode", + "ExpandDimsNode", + "TileNode", + "TakeAlongAxisNode", + "TakeNode", + "RMSNormNode", + "LayerNormNode", + "RopeNode", + "SdpaNode", + "AddNode", + "AddIntNode", + "SubtractIntNode", + "MultiplyIntNode", + "FloorDivideIntNode", + "SymSizeNode", + "MultiplyNode", + "DivideNode", + "SubtractNode", + "Conv1DNode", + "Conv2DNode", + "Conv3DNode", + "GeluNode", + "ARangeNode", + "SiluNode", + "SigmoidNode", + "TanhNode", + "SqueezeNode", + "SplitNode", + "RsqrtNode", + "MaximumNode", + "MinimumNode", + "LogNode", + "SoftmaxNode", + "BroadcastToNode", + "PadNode", + "WhereNode", + "ReshapeNode", + "TransposeNode", + "AsStridedNode", + "ContiguousNode", + "GatherNode", + "SliceNode", + "AsTypeNode", + "ConcatenateNode", + "FullNode", + "FullLikeNode", + "ArgmaxNode", + "SliceUpdateNode", + "IndexCopyNode", + "DequantizeNode", + "LessNode", + "LessEqualNode", + "GreaterNode", + "GreaterEqualNode", + "EqualNode", + "NotEqualNode", + "LogicalNotNode", + "LogicalAndNode", + "LogicalOrNode", + "TriNode", + "TrilNode", + "TriuNode", + "FloorNode", + "CeilNode", + "SquareNode", + "ExpNode", + "SinNode", + "CosNode", + "TanNode", + "ArcsinNode", + "ArccosNode", + "ArctanNode", + "SinhNode", + "CoshNode", + "ArcsinhNode", + "ArccoshNode", + "ArctanhNode", + "Log2Node", + "Log10Node", + "Log1pNode", + "ErfNode", + "Expm1Node", + "RoundNode", + "ReciprocalNode", + "SqrtNode", + "AbsNode", + "NegNode", + "Atan2Node", + "LogAddExpNode", + "FloorDivideNode", + "PowerNode", + "LogSumExpNode", + "SumNode", + "MeanNode", + "VarNode", + "StdNode", + "ProdNode", + "MaxNode", + "MinNode", + "ArgminNode", + "MedianNode", + "ModIntNode", + "RemainderNode", + "ConvTranspose1DNode", + "ConvTranspose2DNode", + "ConvTranspose3DNode", + "ClipNode", + "CumsumNode", + "StackNode", + "SignNode", + "AnyNode", + "AllNode", + "RepeatNode", + "SortNode", + "ArgsortNode", + "PartitionNode", + "ArgPartitionNode", + "QuantizedMatmulNode", + "ScatterAddNode", + "GatherMmNode", + "GatherQmmNode", + "ScanNode", + "MetalKernelNode", + "BitwiseInvertNode", + "RollNode", + "BitwiseAndNode", + "BitwiseOrNode", + "BitwiseXorNode", + "IfNode", + "RandomBitsNode", + nullptr + }; + return names; +} + +inline const char *EnumNameOpNode(OpNode e) { + if (::flatbuffers::IsOutRange(e, OpNode_NONE, OpNode_RandomBitsNode)) return ""; + const size_t index = static_cast(e); + return EnumNamesOpNode()[index]; +} + +template struct OpNodeTraits { + static const OpNode enum_value = OpNode_NONE; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_NoopNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_IdCopyNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AddmmNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ItemIntNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ExpandDimsNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TileNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TakeAlongAxisNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TakeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RMSNormNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LayerNormNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RopeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SdpaNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AddNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AddIntNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SubtractIntNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MultiplyIntNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_FloorDivideIntNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SymSizeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MultiplyNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_DivideNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SubtractNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Conv1DNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Conv2DNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Conv3DNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_GeluNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ARangeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SiluNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SigmoidNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TanhNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SqueezeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SplitNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RsqrtNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MaximumNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MinimumNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LogNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SoftmaxNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_BroadcastToNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_PadNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_WhereNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ReshapeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TransposeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AsStridedNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ContiguousNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_GatherNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SliceNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AsTypeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ConcatenateNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_FullNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_FullLikeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArgmaxNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SliceUpdateNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_IndexCopyNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_DequantizeNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LessNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LessEqualNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_GreaterNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_GreaterEqualNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_EqualNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_NotEqualNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LogicalNotNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LogicalAndNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LogicalOrNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TriNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TrilNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TriuNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_FloorNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_CeilNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SquareNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ExpNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SinNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_CosNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_TanNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArcsinNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArccosNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArctanNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SinhNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_CoshNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArcsinhNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArccoshNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArctanhNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Log2Node; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Log10Node; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Log1pNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ErfNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Expm1Node; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RoundNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ReciprocalNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SqrtNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AbsNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_NegNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_Atan2Node; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LogAddExpNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_FloorDivideNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_PowerNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_LogSumExpNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SumNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MeanNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_VarNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_StdNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ProdNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MaxNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MinNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArgminNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MedianNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ModIntNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RemainderNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ConvTranspose1DNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ConvTranspose2DNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ConvTranspose3DNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ClipNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_CumsumNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_StackNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SignNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AnyNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_AllNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RepeatNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_SortNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArgsortNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_PartitionNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ArgPartitionNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_QuantizedMatmulNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ScatterAddNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_GatherMmNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_GatherQmmNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_ScanNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_MetalKernelNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_BitwiseInvertNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RollNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_BitwiseAndNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_BitwiseOrNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_BitwiseXorNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_IfNode; +}; + +template<> struct OpNodeTraits { + static const OpNode enum_value = OpNode_RandomBitsNode; +}; + +bool VerifyOpNode(::flatbuffers::Verifier &verifier, const void *obj, OpNode type); +bool VerifyOpNodeVector(::flatbuffers::Verifier &verifier, const ::flatbuffers::Vector<::flatbuffers::Offset> *values, const ::flatbuffers::Vector *types); + +enum SlotType : int8_t { + SlotType_TensorSlot = 0, + SlotType_IntValueSlot = 1, + SlotType_FloatValueSlot = 2, + SlotType_BoolValueSlot = 3, + SlotType_MIN = SlotType_TensorSlot, + SlotType_MAX = SlotType_BoolValueSlot +}; + +inline const SlotType (&EnumValuesSlotType())[4] { + static const SlotType values[] = { + SlotType_TensorSlot, + SlotType_IntValueSlot, + SlotType_FloatValueSlot, + SlotType_BoolValueSlot + }; + return values; +} + +inline const char * const *EnumNamesSlotType() { + static const char * const names[5] = { + "TensorSlot", + "IntValueSlot", + "FloatValueSlot", + "BoolValueSlot", + nullptr + }; + return names; +} + +inline const char *EnumNameSlotType(SlotType e) { + if (::flatbuffers::IsOutRange(e, SlotType_TensorSlot, SlotType_BoolValueSlot)) return ""; + const size_t index = static_cast(e); + return EnumNamesSlotType()[index]; +} + +FLATBUFFERS_MANUALLY_ALIGNED_STRUCT(4) Tid FLATBUFFERS_FINAL_CLASS { + private: + uint32_t idx_; + + public: + Tid() + : idx_(0) { + } + Tid(uint32_t _idx) + : idx_(::flatbuffers::EndianScalar(_idx)) { + } + uint32_t idx() const { + return ::flatbuffers::EndianScalar(idx_); + } +}; +FLATBUFFERS_STRUCT_END(Tid, 4); + +FLATBUFFERS_MANUALLY_ALIGNED_STRUCT(4) Vid FLATBUFFERS_FINAL_CLASS { + private: + uint32_t idx_; + + public: + Vid() + : idx_(0) { + } + Vid(uint32_t _idx) + : idx_(::flatbuffers::EndianScalar(_idx)) { + } + uint32_t idx() const { + return ::flatbuffers::EndianScalar(idx_); + } +}; +FLATBUFFERS_STRUCT_END(Vid, 4); + +struct IntOrVid FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef IntOrVidBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_LITERAL = 4, + VT_VID = 6, + VT_IS_VID = 8 + }; + int64_t literal() const { + return GetField(VT_LITERAL, 0); + } + const mlx_delegate::Vid *vid() const { + return GetStruct(VT_VID); + } + bool is_vid() const { + return GetField(VT_IS_VID, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_LITERAL, 8) && + VerifyField(verifier, VT_VID, 4) && + VerifyField(verifier, VT_IS_VID, 1) && + verifier.EndTable(); + } +}; + +struct IntOrVidBuilder { + typedef IntOrVid Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_literal(int64_t literal) { + fbb_.AddElement(IntOrVid::VT_LITERAL, literal, 0); + } + void add_vid(const mlx_delegate::Vid *vid) { + fbb_.AddStruct(IntOrVid::VT_VID, vid); + } + void add_is_vid(bool is_vid) { + fbb_.AddElement(IntOrVid::VT_IS_VID, static_cast(is_vid), 0); + } + explicit IntOrVidBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateIntOrVid( + ::flatbuffers::FlatBufferBuilder &_fbb, + int64_t literal = 0, + const mlx_delegate::Vid *vid = nullptr, + bool is_vid = false) { + IntOrVidBuilder builder_(_fbb); + builder_.add_literal(literal); + builder_.add_vid(vid); + builder_.add_is_vid(is_vid); + return builder_.Finish(); +} + +struct FloatOrVid FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef FloatOrVidBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_LITERAL = 4, + VT_VID = 6, + VT_IS_VID = 8 + }; + double literal() const { + return GetField(VT_LITERAL, 0.0); + } + const mlx_delegate::Vid *vid() const { + return GetStruct(VT_VID); + } + bool is_vid() const { + return GetField(VT_IS_VID, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_LITERAL, 8) && + VerifyField(verifier, VT_VID, 4) && + VerifyField(verifier, VT_IS_VID, 1) && + verifier.EndTable(); + } +}; + +struct FloatOrVidBuilder { + typedef FloatOrVid Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_literal(double literal) { + fbb_.AddElement(FloatOrVid::VT_LITERAL, literal, 0.0); + } + void add_vid(const mlx_delegate::Vid *vid) { + fbb_.AddStruct(FloatOrVid::VT_VID, vid); + } + void add_is_vid(bool is_vid) { + fbb_.AddElement(FloatOrVid::VT_IS_VID, static_cast(is_vid), 0); + } + explicit FloatOrVidBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateFloatOrVid( + ::flatbuffers::FlatBufferBuilder &_fbb, + double literal = 0.0, + const mlx_delegate::Vid *vid = nullptr, + bool is_vid = false) { + FloatOrVidBuilder builder_(_fbb); + builder_.add_literal(literal); + builder_.add_vid(vid); + builder_.add_is_vid(is_vid); + return builder_.Finish(); +} + +struct VidOrTid FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef VidOrTidBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_VID = 4, + VT_TID = 6, + VT_IS_VID = 8 + }; + const mlx_delegate::Vid *vid() const { + return GetStruct(VT_VID); + } + const mlx_delegate::Tid *tid() const { + return GetStruct(VT_TID); + } + bool is_vid() const { + return GetField(VT_IS_VID, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_VID, 4) && + VerifyField(verifier, VT_TID, 4) && + VerifyField(verifier, VT_IS_VID, 1) && + verifier.EndTable(); + } +}; + +struct VidOrTidBuilder { + typedef VidOrTid Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_vid(const mlx_delegate::Vid *vid) { + fbb_.AddStruct(VidOrTid::VT_VID, vid); + } + void add_tid(const mlx_delegate::Tid *tid) { + fbb_.AddStruct(VidOrTid::VT_TID, tid); + } + void add_is_vid(bool is_vid) { + fbb_.AddElement(VidOrTid::VT_IS_VID, static_cast(is_vid), 0); + } + explicit VidOrTidBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateVidOrTid( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Vid *vid = nullptr, + const mlx_delegate::Tid *tid = nullptr, + bool is_vid = false) { + VidOrTidBuilder builder_(_fbb); + builder_.add_tid(tid); + builder_.add_vid(vid); + builder_.add_is_vid(is_vid); + return builder_.Finish(); +} + +struct IntOrVidOrTid FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef IntOrVidOrTidBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_LITERAL = 4, + VT_VID = 6, + VT_TID = 8, + VT_KIND = 10 + }; + int64_t literal() const { + return GetField(VT_LITERAL, 0); + } + const mlx_delegate::Vid *vid() const { + return GetStruct(VT_VID); + } + const mlx_delegate::Tid *tid() const { + return GetStruct(VT_TID); + } + uint8_t kind() const { + return GetField(VT_KIND, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_LITERAL, 8) && + VerifyField(verifier, VT_VID, 4) && + VerifyField(verifier, VT_TID, 4) && + VerifyField(verifier, VT_KIND, 1) && + verifier.EndTable(); + } +}; + +struct IntOrVidOrTidBuilder { + typedef IntOrVidOrTid Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_literal(int64_t literal) { + fbb_.AddElement(IntOrVidOrTid::VT_LITERAL, literal, 0); + } + void add_vid(const mlx_delegate::Vid *vid) { + fbb_.AddStruct(IntOrVidOrTid::VT_VID, vid); + } + void add_tid(const mlx_delegate::Tid *tid) { + fbb_.AddStruct(IntOrVidOrTid::VT_TID, tid); + } + void add_kind(uint8_t kind) { + fbb_.AddElement(IntOrVidOrTid::VT_KIND, kind, 0); + } + explicit IntOrVidOrTidBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateIntOrVidOrTid( + ::flatbuffers::FlatBufferBuilder &_fbb, + int64_t literal = 0, + const mlx_delegate::Vid *vid = nullptr, + const mlx_delegate::Tid *tid = nullptr, + uint8_t kind = 0) { + IntOrVidOrTidBuilder builder_(_fbb); + builder_.add_literal(literal); + builder_.add_tid(tid); + builder_.add_vid(vid); + builder_.add_kind(kind); + return builder_.Finish(); +} + +struct NoopNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef NoopNodeBuilder Builder; + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + verifier.EndTable(); + } +}; + +struct NoopNodeBuilder { + typedef NoopNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + explicit NoopNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateNoopNode( + ::flatbuffers::FlatBufferBuilder &_fbb) { + NoopNodeBuilder builder_(_fbb); + return builder_.Finish(); +} + +struct IdCopyNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef IdCopyNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct IdCopyNodeBuilder { + typedef IdCopyNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(IdCopyNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(IdCopyNode::VT_OUT, out); + } + explicit IdCopyNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, IdCopyNode::VT_X); + fbb_.Required(o, IdCopyNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateIdCopyNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + IdCopyNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct AddmmNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AddmmNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_MAT1 = 4, + VT_MAT2 = 6, + VT_OUT = 8, + VT_BIAS = 10, + VT_ALPHA = 12, + VT_BETA = 14 + }; + const mlx_delegate::Tid *mat1() const { + return GetStruct(VT_MAT1); + } + const mlx_delegate::Tid *mat2() const { + return GetStruct(VT_MAT2); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::Tid *bias() const { + return GetStruct(VT_BIAS); + } + float alpha() const { + return GetField(VT_ALPHA, 1.0f); + } + float beta() const { + return GetField(VT_BETA, 1.0f); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_MAT1, 4) && + VerifyFieldRequired(verifier, VT_MAT2, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_BIAS, 4) && + VerifyField(verifier, VT_ALPHA, 4) && + VerifyField(verifier, VT_BETA, 4) && + verifier.EndTable(); + } +}; + +struct AddmmNodeBuilder { + typedef AddmmNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_mat1(const mlx_delegate::Tid *mat1) { + fbb_.AddStruct(AddmmNode::VT_MAT1, mat1); + } + void add_mat2(const mlx_delegate::Tid *mat2) { + fbb_.AddStruct(AddmmNode::VT_MAT2, mat2); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AddmmNode::VT_OUT, out); + } + void add_bias(const mlx_delegate::Tid *bias) { + fbb_.AddStruct(AddmmNode::VT_BIAS, bias); + } + void add_alpha(float alpha) { + fbb_.AddElement(AddmmNode::VT_ALPHA, alpha, 1.0f); + } + void add_beta(float beta) { + fbb_.AddElement(AddmmNode::VT_BETA, beta, 1.0f); + } + explicit AddmmNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AddmmNode::VT_MAT1); + fbb_.Required(o, AddmmNode::VT_MAT2); + fbb_.Required(o, AddmmNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAddmmNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *mat1 = nullptr, + const mlx_delegate::Tid *mat2 = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *bias = nullptr, + float alpha = 1.0f, + float beta = 1.0f) { + AddmmNodeBuilder builder_(_fbb); + builder_.add_beta(beta); + builder_.add_alpha(alpha); + builder_.add_bias(bias); + builder_.add_out(out); + builder_.add_mat2(mat2); + builder_.add_mat1(mat1); + return builder_.Finish(); +} + +struct ItemIntNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ItemIntNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ItemIntNodeBuilder { + typedef ItemIntNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ItemIntNode::VT_X, x); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(ItemIntNode::VT_OUT, out); + } + explicit ItemIntNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ItemIntNode::VT_X); + fbb_.Required(o, ItemIntNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateItemIntNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Vid *out = nullptr) { + ItemIntNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ExpandDimsNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ExpandDimsNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct ExpandDimsNodeBuilder { + typedef ExpandDimsNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ExpandDimsNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ExpandDimsNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ExpandDimsNode::VT_AXIS, axis, 0); + } + explicit ExpandDimsNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ExpandDimsNode::VT_X); + fbb_.Required(o, ExpandDimsNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateExpandDimsNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + ExpandDimsNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct TileNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TileNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_REPS = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *reps() const { + return GetPointer> *>(VT_REPS); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_REPS) && + verifier.VerifyVector(reps()) && + verifier.VerifyVectorOfTables(reps()) && + verifier.EndTable(); + } +}; + +struct TileNodeBuilder { + typedef TileNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TileNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TileNode::VT_OUT, out); + } + void add_reps(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> reps) { + fbb_.AddOffset(TileNode::VT_REPS, reps); + } + explicit TileNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TileNode::VT_X); + fbb_.Required(o, TileNode::VT_OUT); + fbb_.Required(o, TileNode::VT_REPS); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTileNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> reps = 0) { + TileNodeBuilder builder_(_fbb); + builder_.add_reps(reps); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateTileNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *reps = nullptr) { + auto reps__ = reps ? _fbb.CreateVector<::flatbuffers::Offset>(*reps) : 0; + return mlx_delegate::CreateTileNode( + _fbb, + x, + out, + reps__); +} + +struct TakeAlongAxisNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TakeAlongAxisNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_INDICES = 6, + VT_OUT = 8, + VT_AXIS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *indices() const { + return GetStruct(VT_INDICES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_INDICES, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct TakeAlongAxisNodeBuilder { + typedef TakeAlongAxisNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TakeAlongAxisNode::VT_X, x); + } + void add_indices(const mlx_delegate::Tid *indices) { + fbb_.AddStruct(TakeAlongAxisNode::VT_INDICES, indices); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TakeAlongAxisNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(TakeAlongAxisNode::VT_AXIS, axis, 0); + } + explicit TakeAlongAxisNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TakeAlongAxisNode::VT_X); + fbb_.Required(o, TakeAlongAxisNode::VT_INDICES); + fbb_.Required(o, TakeAlongAxisNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTakeAlongAxisNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *indices = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + TakeAlongAxisNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_indices(indices); + builder_.add_x(x); + return builder_.Finish(); +} + +struct TakeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TakeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_INDEX = 8, + VT_AXIS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVidOrTid *index() const { + return GetPointer(VT_INDEX); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_INDEX) && + verifier.VerifyTable(index()) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct TakeNodeBuilder { + typedef TakeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TakeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TakeNode::VT_OUT, out); + } + void add_index(::flatbuffers::Offset index) { + fbb_.AddOffset(TakeNode::VT_INDEX, index); + } + void add_axis(int32_t axis) { + fbb_.AddElement(TakeNode::VT_AXIS, axis, 0); + } + explicit TakeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TakeNode::VT_X); + fbb_.Required(o, TakeNode::VT_OUT); + fbb_.Required(o, TakeNode::VT_INDEX); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTakeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset index = 0, + int32_t axis = 0) { + TakeNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_index(index); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct RMSNormNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RMSNormNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_WEIGHT = 6, + VT_OUT = 8, + VT_EPS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *weight() const { + return GetStruct(VT_WEIGHT); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + float eps() const { + return GetField(VT_EPS, 0.0f); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyField(verifier, VT_WEIGHT, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_EPS, 4) && + verifier.EndTable(); + } +}; + +struct RMSNormNodeBuilder { + typedef RMSNormNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(RMSNormNode::VT_X, x); + } + void add_weight(const mlx_delegate::Tid *weight) { + fbb_.AddStruct(RMSNormNode::VT_WEIGHT, weight); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RMSNormNode::VT_OUT, out); + } + void add_eps(float eps) { + fbb_.AddElement(RMSNormNode::VT_EPS, eps, 0.0f); + } + explicit RMSNormNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RMSNormNode::VT_X); + fbb_.Required(o, RMSNormNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRMSNormNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *weight = nullptr, + const mlx_delegate::Tid *out = nullptr, + float eps = 0.0f) { + RMSNormNodeBuilder builder_(_fbb); + builder_.add_eps(eps); + builder_.add_out(out); + builder_.add_weight(weight); + builder_.add_x(x); + return builder_.Finish(); +} + +struct LayerNormNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LayerNormNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_WEIGHT = 8, + VT_BIAS = 10, + VT_EPS = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::Tid *weight() const { + return GetStruct(VT_WEIGHT); + } + const mlx_delegate::Tid *bias() const { + return GetStruct(VT_BIAS); + } + float eps() const { + return GetField(VT_EPS, 0.0f); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_WEIGHT, 4) && + VerifyField(verifier, VT_BIAS, 4) && + VerifyField(verifier, VT_EPS, 4) && + verifier.EndTable(); + } +}; + +struct LayerNormNodeBuilder { + typedef LayerNormNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(LayerNormNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LayerNormNode::VT_OUT, out); + } + void add_weight(const mlx_delegate::Tid *weight) { + fbb_.AddStruct(LayerNormNode::VT_WEIGHT, weight); + } + void add_bias(const mlx_delegate::Tid *bias) { + fbb_.AddStruct(LayerNormNode::VT_BIAS, bias); + } + void add_eps(float eps) { + fbb_.AddElement(LayerNormNode::VT_EPS, eps, 0.0f); + } + explicit LayerNormNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LayerNormNode::VT_X); + fbb_.Required(o, LayerNormNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLayerNormNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *weight = nullptr, + const mlx_delegate::Tid *bias = nullptr, + float eps = 0.0f) { + LayerNormNodeBuilder builder_(_fbb); + builder_.add_eps(eps); + builder_.add_bias(bias); + builder_.add_weight(weight); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct RopeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RopeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_DIMS = 8, + VT_OFFSET = 10, + VT_FREQS = 12, + VT_TRADITIONAL = 14, + VT_BASE = 16, + VT_SCALE = 18 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t dims() const { + return GetField(VT_DIMS, 0); + } + const mlx_delegate::VidOrTid *offset() const { + return GetPointer(VT_OFFSET); + } + const mlx_delegate::Tid *freqs() const { + return GetStruct(VT_FREQS); + } + bool traditional() const { + return GetField(VT_TRADITIONAL, 0) != 0; + } + float base() const { + return GetField(VT_BASE, 500000.0f); + } + float scale() const { + return GetField(VT_SCALE, 1.0f); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_DIMS, 4) && + VerifyOffsetRequired(verifier, VT_OFFSET) && + verifier.VerifyTable(offset()) && + VerifyField(verifier, VT_FREQS, 4) && + VerifyField(verifier, VT_TRADITIONAL, 1) && + VerifyField(verifier, VT_BASE, 4) && + VerifyField(verifier, VT_SCALE, 4) && + verifier.EndTable(); + } +}; + +struct RopeNodeBuilder { + typedef RopeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(RopeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RopeNode::VT_OUT, out); + } + void add_dims(int32_t dims) { + fbb_.AddElement(RopeNode::VT_DIMS, dims, 0); + } + void add_offset(::flatbuffers::Offset offset) { + fbb_.AddOffset(RopeNode::VT_OFFSET, offset); + } + void add_freqs(const mlx_delegate::Tid *freqs) { + fbb_.AddStruct(RopeNode::VT_FREQS, freqs); + } + void add_traditional(bool traditional) { + fbb_.AddElement(RopeNode::VT_TRADITIONAL, static_cast(traditional), 0); + } + void add_base(float base) { + fbb_.AddElement(RopeNode::VT_BASE, base, 500000.0f); + } + void add_scale(float scale) { + fbb_.AddElement(RopeNode::VT_SCALE, scale, 1.0f); + } + explicit RopeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RopeNode::VT_X); + fbb_.Required(o, RopeNode::VT_OUT); + fbb_.Required(o, RopeNode::VT_OFFSET); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRopeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t dims = 0, + ::flatbuffers::Offset offset = 0, + const mlx_delegate::Tid *freqs = nullptr, + bool traditional = false, + float base = 500000.0f, + float scale = 1.0f) { + RopeNodeBuilder builder_(_fbb); + builder_.add_scale(scale); + builder_.add_base(base); + builder_.add_freqs(freqs); + builder_.add_offset(offset); + builder_.add_dims(dims); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_traditional(traditional); + return builder_.Finish(); +} + +struct SdpaNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SdpaNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_Q = 4, + VT_K = 6, + VT_V = 8, + VT_OUT = 10, + VT_SCALE = 12, + VT_MASK = 14, + VT_CAUSAL = 16 + }; + const mlx_delegate::Tid *q() const { + return GetStruct(VT_Q); + } + const mlx_delegate::Tid *k() const { + return GetStruct(VT_K); + } + const mlx_delegate::Tid *v() const { + return GetStruct(VT_V); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + float scale() const { + return GetField(VT_SCALE, 0.0f); + } + const mlx_delegate::Tid *mask() const { + return GetStruct(VT_MASK); + } + bool causal() const { + return GetField(VT_CAUSAL, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_Q, 4) && + VerifyFieldRequired(verifier, VT_K, 4) && + VerifyFieldRequired(verifier, VT_V, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_SCALE, 4) && + VerifyField(verifier, VT_MASK, 4) && + VerifyField(verifier, VT_CAUSAL, 1) && + verifier.EndTable(); + } +}; + +struct SdpaNodeBuilder { + typedef SdpaNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_q(const mlx_delegate::Tid *q) { + fbb_.AddStruct(SdpaNode::VT_Q, q); + } + void add_k(const mlx_delegate::Tid *k) { + fbb_.AddStruct(SdpaNode::VT_K, k); + } + void add_v(const mlx_delegate::Tid *v) { + fbb_.AddStruct(SdpaNode::VT_V, v); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SdpaNode::VT_OUT, out); + } + void add_scale(float scale) { + fbb_.AddElement(SdpaNode::VT_SCALE, scale, 0.0f); + } + void add_mask(const mlx_delegate::Tid *mask) { + fbb_.AddStruct(SdpaNode::VT_MASK, mask); + } + void add_causal(bool causal) { + fbb_.AddElement(SdpaNode::VT_CAUSAL, static_cast(causal), 0); + } + explicit SdpaNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SdpaNode::VT_Q); + fbb_.Required(o, SdpaNode::VT_K); + fbb_.Required(o, SdpaNode::VT_V); + fbb_.Required(o, SdpaNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSdpaNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *q = nullptr, + const mlx_delegate::Tid *k = nullptr, + const mlx_delegate::Tid *v = nullptr, + const mlx_delegate::Tid *out = nullptr, + float scale = 0.0f, + const mlx_delegate::Tid *mask = nullptr, + bool causal = false) { + SdpaNodeBuilder builder_(_fbb); + builder_.add_mask(mask); + builder_.add_scale(scale); + builder_.add_out(out); + builder_.add_v(v); + builder_.add_k(k); + builder_.add_q(q); + builder_.add_causal(causal); + return builder_.Finish(); +} + +struct AddNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AddNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct AddNodeBuilder { + typedef AddNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(AddNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(AddNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AddNode::VT_OUT, out); + } + explicit AddNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AddNode::VT_A); + fbb_.Required(o, AddNode::VT_B); + fbb_.Required(o, AddNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAddNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + AddNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct AddIntNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AddIntNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::IntOrVid *a() const { + return GetPointer(VT_A); + } + const mlx_delegate::IntOrVid *b() const { + return GetPointer(VT_B); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_A) && + verifier.VerifyTable(a()) && + VerifyOffsetRequired(verifier, VT_B) && + verifier.VerifyTable(b()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct AddIntNodeBuilder { + typedef AddIntNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(::flatbuffers::Offset a) { + fbb_.AddOffset(AddIntNode::VT_A, a); + } + void add_b(::flatbuffers::Offset b) { + fbb_.AddOffset(AddIntNode::VT_B, b); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(AddIntNode::VT_OUT, out); + } + explicit AddIntNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AddIntNode::VT_A); + fbb_.Required(o, AddIntNode::VT_B); + fbb_.Required(o, AddIntNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAddIntNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset a = 0, + ::flatbuffers::Offset b = 0, + const mlx_delegate::Vid *out = nullptr) { + AddIntNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct SubtractIntNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SubtractIntNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::IntOrVid *a() const { + return GetPointer(VT_A); + } + const mlx_delegate::IntOrVid *b() const { + return GetPointer(VT_B); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_A) && + verifier.VerifyTable(a()) && + VerifyOffsetRequired(verifier, VT_B) && + verifier.VerifyTable(b()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SubtractIntNodeBuilder { + typedef SubtractIntNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(::flatbuffers::Offset a) { + fbb_.AddOffset(SubtractIntNode::VT_A, a); + } + void add_b(::flatbuffers::Offset b) { + fbb_.AddOffset(SubtractIntNode::VT_B, b); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(SubtractIntNode::VT_OUT, out); + } + explicit SubtractIntNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SubtractIntNode::VT_A); + fbb_.Required(o, SubtractIntNode::VT_B); + fbb_.Required(o, SubtractIntNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSubtractIntNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset a = 0, + ::flatbuffers::Offset b = 0, + const mlx_delegate::Vid *out = nullptr) { + SubtractIntNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct MultiplyIntNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MultiplyIntNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::IntOrVid *a() const { + return GetPointer(VT_A); + } + const mlx_delegate::IntOrVid *b() const { + return GetPointer(VT_B); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_A) && + verifier.VerifyTable(a()) && + VerifyOffsetRequired(verifier, VT_B) && + verifier.VerifyTable(b()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct MultiplyIntNodeBuilder { + typedef MultiplyIntNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(::flatbuffers::Offset a) { + fbb_.AddOffset(MultiplyIntNode::VT_A, a); + } + void add_b(::flatbuffers::Offset b) { + fbb_.AddOffset(MultiplyIntNode::VT_B, b); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(MultiplyIntNode::VT_OUT, out); + } + explicit MultiplyIntNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MultiplyIntNode::VT_A); + fbb_.Required(o, MultiplyIntNode::VT_B); + fbb_.Required(o, MultiplyIntNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMultiplyIntNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset a = 0, + ::flatbuffers::Offset b = 0, + const mlx_delegate::Vid *out = nullptr) { + MultiplyIntNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct FloorDivideIntNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef FloorDivideIntNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::IntOrVid *a() const { + return GetPointer(VT_A); + } + const mlx_delegate::IntOrVid *b() const { + return GetPointer(VT_B); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_A) && + verifier.VerifyTable(a()) && + VerifyOffsetRequired(verifier, VT_B) && + verifier.VerifyTable(b()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct FloorDivideIntNodeBuilder { + typedef FloorDivideIntNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(::flatbuffers::Offset a) { + fbb_.AddOffset(FloorDivideIntNode::VT_A, a); + } + void add_b(::flatbuffers::Offset b) { + fbb_.AddOffset(FloorDivideIntNode::VT_B, b); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(FloorDivideIntNode::VT_OUT, out); + } + explicit FloorDivideIntNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, FloorDivideIntNode::VT_A); + fbb_.Required(o, FloorDivideIntNode::VT_B); + fbb_.Required(o, FloorDivideIntNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateFloorDivideIntNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset a = 0, + ::flatbuffers::Offset b = 0, + const mlx_delegate::Vid *out = nullptr) { + FloorDivideIntNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct ModIntNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ModIntNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::IntOrVid *a() const { + return GetPointer(VT_A); + } + const mlx_delegate::IntOrVid *b() const { + return GetPointer(VT_B); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_A) && + verifier.VerifyTable(a()) && + VerifyOffsetRequired(verifier, VT_B) && + verifier.VerifyTable(b()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ModIntNodeBuilder { + typedef ModIntNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(::flatbuffers::Offset a) { + fbb_.AddOffset(ModIntNode::VT_A, a); + } + void add_b(::flatbuffers::Offset b) { + fbb_.AddOffset(ModIntNode::VT_B, b); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(ModIntNode::VT_OUT, out); + } + explicit ModIntNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ModIntNode::VT_A); + fbb_.Required(o, ModIntNode::VT_B); + fbb_.Required(o, ModIntNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateModIntNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset a = 0, + ::flatbuffers::Offset b = 0, + const mlx_delegate::Vid *out = nullptr) { + ModIntNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct SymSizeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SymSizeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_DIM = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + int32_t dim() const { + return GetField(VT_DIM, 0); + } + const mlx_delegate::Vid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyField(verifier, VT_DIM, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SymSizeNodeBuilder { + typedef SymSizeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(SymSizeNode::VT_A, a); + } + void add_dim(int32_t dim) { + fbb_.AddElement(SymSizeNode::VT_DIM, dim, 0); + } + void add_out(const mlx_delegate::Vid *out) { + fbb_.AddStruct(SymSizeNode::VT_OUT, out); + } + explicit SymSizeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SymSizeNode::VT_A); + fbb_.Required(o, SymSizeNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSymSizeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + int32_t dim = 0, + const mlx_delegate::Vid *out = nullptr) { + SymSizeNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_dim(dim); + builder_.add_a(a); + return builder_.Finish(); +} + +struct MultiplyNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MultiplyNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct MultiplyNodeBuilder { + typedef MultiplyNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(MultiplyNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(MultiplyNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MultiplyNode::VT_OUT, out); + } + explicit MultiplyNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MultiplyNode::VT_A); + fbb_.Required(o, MultiplyNode::VT_B); + fbb_.Required(o, MultiplyNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMultiplyNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + MultiplyNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct DivideNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef DivideNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct DivideNodeBuilder { + typedef DivideNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(DivideNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(DivideNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(DivideNode::VT_OUT, out); + } + explicit DivideNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, DivideNode::VT_A); + fbb_.Required(o, DivideNode::VT_B); + fbb_.Required(o, DivideNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateDivideNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + DivideNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct SubtractNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SubtractNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SubtractNodeBuilder { + typedef SubtractNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(SubtractNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(SubtractNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SubtractNode::VT_OUT, out); + } + explicit SubtractNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SubtractNode::VT_A); + fbb_.Required(o, SubtractNode::VT_B); + fbb_.Required(o, SubtractNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSubtractNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SubtractNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct Conv1DNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Conv1DNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_OUT = 8, + VT_STRIDE = 10, + VT_PADDING = 12, + VT_DILATION = 14, + VT_GROUPS = 16 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t stride() const { + return GetField(VT_STRIDE, 1); + } + int32_t padding() const { + return GetField(VT_PADDING, 0); + } + int32_t dilation() const { + return GetField(VT_DILATION, 1); + } + int32_t groups() const { + return GetField(VT_GROUPS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_STRIDE, 4) && + VerifyField(verifier, VT_PADDING, 4) && + VerifyField(verifier, VT_DILATION, 4) && + VerifyField(verifier, VT_GROUPS, 4) && + verifier.EndTable(); + } +}; + +struct Conv1DNodeBuilder { + typedef Conv1DNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Conv1DNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(Conv1DNode::VT_W, w); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Conv1DNode::VT_OUT, out); + } + void add_stride(int32_t stride) { + fbb_.AddElement(Conv1DNode::VT_STRIDE, stride, 1); + } + void add_padding(int32_t padding) { + fbb_.AddElement(Conv1DNode::VT_PADDING, padding, 0); + } + void add_dilation(int32_t dilation) { + fbb_.AddElement(Conv1DNode::VT_DILATION, dilation, 1); + } + void add_groups(int32_t groups) { + fbb_.AddElement(Conv1DNode::VT_GROUPS, groups, 1); + } + explicit Conv1DNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Conv1DNode::VT_X); + fbb_.Required(o, Conv1DNode::VT_W); + fbb_.Required(o, Conv1DNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConv1DNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t stride = 1, + int32_t padding = 0, + int32_t dilation = 1, + int32_t groups = 1) { + Conv1DNodeBuilder builder_(_fbb); + builder_.add_groups(groups); + builder_.add_dilation(dilation); + builder_.add_padding(padding); + builder_.add_stride(stride); + builder_.add_out(out); + builder_.add_w(w); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Conv2DNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Conv2DNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_OUT = 8, + VT_STRIDE_H = 10, + VT_STRIDE_W = 12, + VT_PADDING_H = 14, + VT_PADDING_W = 16, + VT_DILATION_H = 18, + VT_DILATION_W = 20, + VT_GROUPS = 22 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t stride_h() const { + return GetField(VT_STRIDE_H, 1); + } + int32_t stride_w() const { + return GetField(VT_STRIDE_W, 1); + } + int32_t padding_h() const { + return GetField(VT_PADDING_H, 0); + } + int32_t padding_w() const { + return GetField(VT_PADDING_W, 0); + } + int32_t dilation_h() const { + return GetField(VT_DILATION_H, 1); + } + int32_t dilation_w() const { + return GetField(VT_DILATION_W, 1); + } + int32_t groups() const { + return GetField(VT_GROUPS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_STRIDE_H, 4) && + VerifyField(verifier, VT_STRIDE_W, 4) && + VerifyField(verifier, VT_PADDING_H, 4) && + VerifyField(verifier, VT_PADDING_W, 4) && + VerifyField(verifier, VT_DILATION_H, 4) && + VerifyField(verifier, VT_DILATION_W, 4) && + VerifyField(verifier, VT_GROUPS, 4) && + verifier.EndTable(); + } +}; + +struct Conv2DNodeBuilder { + typedef Conv2DNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Conv2DNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(Conv2DNode::VT_W, w); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Conv2DNode::VT_OUT, out); + } + void add_stride_h(int32_t stride_h) { + fbb_.AddElement(Conv2DNode::VT_STRIDE_H, stride_h, 1); + } + void add_stride_w(int32_t stride_w) { + fbb_.AddElement(Conv2DNode::VT_STRIDE_W, stride_w, 1); + } + void add_padding_h(int32_t padding_h) { + fbb_.AddElement(Conv2DNode::VT_PADDING_H, padding_h, 0); + } + void add_padding_w(int32_t padding_w) { + fbb_.AddElement(Conv2DNode::VT_PADDING_W, padding_w, 0); + } + void add_dilation_h(int32_t dilation_h) { + fbb_.AddElement(Conv2DNode::VT_DILATION_H, dilation_h, 1); + } + void add_dilation_w(int32_t dilation_w) { + fbb_.AddElement(Conv2DNode::VT_DILATION_W, dilation_w, 1); + } + void add_groups(int32_t groups) { + fbb_.AddElement(Conv2DNode::VT_GROUPS, groups, 1); + } + explicit Conv2DNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Conv2DNode::VT_X); + fbb_.Required(o, Conv2DNode::VT_W); + fbb_.Required(o, Conv2DNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConv2DNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t stride_h = 1, + int32_t stride_w = 1, + int32_t padding_h = 0, + int32_t padding_w = 0, + int32_t dilation_h = 1, + int32_t dilation_w = 1, + int32_t groups = 1) { + Conv2DNodeBuilder builder_(_fbb); + builder_.add_groups(groups); + builder_.add_dilation_w(dilation_w); + builder_.add_dilation_h(dilation_h); + builder_.add_padding_w(padding_w); + builder_.add_padding_h(padding_h); + builder_.add_stride_w(stride_w); + builder_.add_stride_h(stride_h); + builder_.add_out(out); + builder_.add_w(w); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Conv3DNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Conv3DNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_OUT = 8, + VT_STRIDE_D = 10, + VT_STRIDE_H = 12, + VT_STRIDE_W = 14, + VT_PADDING_D = 16, + VT_PADDING_H = 18, + VT_PADDING_W = 20, + VT_DILATION_D = 22, + VT_DILATION_H = 24, + VT_DILATION_W = 26, + VT_GROUPS = 28 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t stride_d() const { + return GetField(VT_STRIDE_D, 1); + } + int32_t stride_h() const { + return GetField(VT_STRIDE_H, 1); + } + int32_t stride_w() const { + return GetField(VT_STRIDE_W, 1); + } + int32_t padding_d() const { + return GetField(VT_PADDING_D, 0); + } + int32_t padding_h() const { + return GetField(VT_PADDING_H, 0); + } + int32_t padding_w() const { + return GetField(VT_PADDING_W, 0); + } + int32_t dilation_d() const { + return GetField(VT_DILATION_D, 1); + } + int32_t dilation_h() const { + return GetField(VT_DILATION_H, 1); + } + int32_t dilation_w() const { + return GetField(VT_DILATION_W, 1); + } + int32_t groups() const { + return GetField(VT_GROUPS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_STRIDE_D, 4) && + VerifyField(verifier, VT_STRIDE_H, 4) && + VerifyField(verifier, VT_STRIDE_W, 4) && + VerifyField(verifier, VT_PADDING_D, 4) && + VerifyField(verifier, VT_PADDING_H, 4) && + VerifyField(verifier, VT_PADDING_W, 4) && + VerifyField(verifier, VT_DILATION_D, 4) && + VerifyField(verifier, VT_DILATION_H, 4) && + VerifyField(verifier, VT_DILATION_W, 4) && + VerifyField(verifier, VT_GROUPS, 4) && + verifier.EndTable(); + } +}; + +struct Conv3DNodeBuilder { + typedef Conv3DNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Conv3DNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(Conv3DNode::VT_W, w); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Conv3DNode::VT_OUT, out); + } + void add_stride_d(int32_t stride_d) { + fbb_.AddElement(Conv3DNode::VT_STRIDE_D, stride_d, 1); + } + void add_stride_h(int32_t stride_h) { + fbb_.AddElement(Conv3DNode::VT_STRIDE_H, stride_h, 1); + } + void add_stride_w(int32_t stride_w) { + fbb_.AddElement(Conv3DNode::VT_STRIDE_W, stride_w, 1); + } + void add_padding_d(int32_t padding_d) { + fbb_.AddElement(Conv3DNode::VT_PADDING_D, padding_d, 0); + } + void add_padding_h(int32_t padding_h) { + fbb_.AddElement(Conv3DNode::VT_PADDING_H, padding_h, 0); + } + void add_padding_w(int32_t padding_w) { + fbb_.AddElement(Conv3DNode::VT_PADDING_W, padding_w, 0); + } + void add_dilation_d(int32_t dilation_d) { + fbb_.AddElement(Conv3DNode::VT_DILATION_D, dilation_d, 1); + } + void add_dilation_h(int32_t dilation_h) { + fbb_.AddElement(Conv3DNode::VT_DILATION_H, dilation_h, 1); + } + void add_dilation_w(int32_t dilation_w) { + fbb_.AddElement(Conv3DNode::VT_DILATION_W, dilation_w, 1); + } + void add_groups(int32_t groups) { + fbb_.AddElement(Conv3DNode::VT_GROUPS, groups, 1); + } + explicit Conv3DNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Conv3DNode::VT_X); + fbb_.Required(o, Conv3DNode::VT_W); + fbb_.Required(o, Conv3DNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConv3DNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t stride_d = 1, + int32_t stride_h = 1, + int32_t stride_w = 1, + int32_t padding_d = 0, + int32_t padding_h = 0, + int32_t padding_w = 0, + int32_t dilation_d = 1, + int32_t dilation_h = 1, + int32_t dilation_w = 1, + int32_t groups = 1) { + Conv3DNodeBuilder builder_(_fbb); + builder_.add_groups(groups); + builder_.add_dilation_w(dilation_w); + builder_.add_dilation_h(dilation_h); + builder_.add_dilation_d(dilation_d); + builder_.add_padding_w(padding_w); + builder_.add_padding_h(padding_h); + builder_.add_padding_d(padding_d); + builder_.add_stride_w(stride_w); + builder_.add_stride_h(stride_h); + builder_.add_stride_d(stride_d); + builder_.add_out(out); + builder_.add_w(w); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ConvTranspose1DNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ConvTranspose1DNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_OUT = 8, + VT_STRIDE = 10, + VT_PADDING = 12, + VT_DILATION = 14, + VT_OUTPUT_PADDING = 16, + VT_GROUPS = 18 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t stride() const { + return GetField(VT_STRIDE, 1); + } + int32_t padding() const { + return GetField(VT_PADDING, 0); + } + int32_t dilation() const { + return GetField(VT_DILATION, 1); + } + int32_t output_padding() const { + return GetField(VT_OUTPUT_PADDING, 0); + } + int32_t groups() const { + return GetField(VT_GROUPS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_STRIDE, 4) && + VerifyField(verifier, VT_PADDING, 4) && + VerifyField(verifier, VT_DILATION, 4) && + VerifyField(verifier, VT_OUTPUT_PADDING, 4) && + VerifyField(verifier, VT_GROUPS, 4) && + verifier.EndTable(); + } +}; + +struct ConvTranspose1DNodeBuilder { + typedef ConvTranspose1DNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ConvTranspose1DNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(ConvTranspose1DNode::VT_W, w); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ConvTranspose1DNode::VT_OUT, out); + } + void add_stride(int32_t stride) { + fbb_.AddElement(ConvTranspose1DNode::VT_STRIDE, stride, 1); + } + void add_padding(int32_t padding) { + fbb_.AddElement(ConvTranspose1DNode::VT_PADDING, padding, 0); + } + void add_dilation(int32_t dilation) { + fbb_.AddElement(ConvTranspose1DNode::VT_DILATION, dilation, 1); + } + void add_output_padding(int32_t output_padding) { + fbb_.AddElement(ConvTranspose1DNode::VT_OUTPUT_PADDING, output_padding, 0); + } + void add_groups(int32_t groups) { + fbb_.AddElement(ConvTranspose1DNode::VT_GROUPS, groups, 1); + } + explicit ConvTranspose1DNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ConvTranspose1DNode::VT_X); + fbb_.Required(o, ConvTranspose1DNode::VT_W); + fbb_.Required(o, ConvTranspose1DNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConvTranspose1DNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t stride = 1, + int32_t padding = 0, + int32_t dilation = 1, + int32_t output_padding = 0, + int32_t groups = 1) { + ConvTranspose1DNodeBuilder builder_(_fbb); + builder_.add_groups(groups); + builder_.add_output_padding(output_padding); + builder_.add_dilation(dilation); + builder_.add_padding(padding); + builder_.add_stride(stride); + builder_.add_out(out); + builder_.add_w(w); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ConvTranspose2DNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ConvTranspose2DNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_OUT = 8, + VT_STRIDE_H = 10, + VT_STRIDE_W = 12, + VT_PADDING_H = 14, + VT_PADDING_W = 16, + VT_DILATION_H = 18, + VT_DILATION_W = 20, + VT_OUTPUT_PADDING_H = 22, + VT_OUTPUT_PADDING_W = 24, + VT_GROUPS = 26 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t stride_h() const { + return GetField(VT_STRIDE_H, 1); + } + int32_t stride_w() const { + return GetField(VT_STRIDE_W, 1); + } + int32_t padding_h() const { + return GetField(VT_PADDING_H, 0); + } + int32_t padding_w() const { + return GetField(VT_PADDING_W, 0); + } + int32_t dilation_h() const { + return GetField(VT_DILATION_H, 1); + } + int32_t dilation_w() const { + return GetField(VT_DILATION_W, 1); + } + int32_t output_padding_h() const { + return GetField(VT_OUTPUT_PADDING_H, 0); + } + int32_t output_padding_w() const { + return GetField(VT_OUTPUT_PADDING_W, 0); + } + int32_t groups() const { + return GetField(VT_GROUPS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_STRIDE_H, 4) && + VerifyField(verifier, VT_STRIDE_W, 4) && + VerifyField(verifier, VT_PADDING_H, 4) && + VerifyField(verifier, VT_PADDING_W, 4) && + VerifyField(verifier, VT_DILATION_H, 4) && + VerifyField(verifier, VT_DILATION_W, 4) && + VerifyField(verifier, VT_OUTPUT_PADDING_H, 4) && + VerifyField(verifier, VT_OUTPUT_PADDING_W, 4) && + VerifyField(verifier, VT_GROUPS, 4) && + verifier.EndTable(); + } +}; + +struct ConvTranspose2DNodeBuilder { + typedef ConvTranspose2DNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ConvTranspose2DNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(ConvTranspose2DNode::VT_W, w); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ConvTranspose2DNode::VT_OUT, out); + } + void add_stride_h(int32_t stride_h) { + fbb_.AddElement(ConvTranspose2DNode::VT_STRIDE_H, stride_h, 1); + } + void add_stride_w(int32_t stride_w) { + fbb_.AddElement(ConvTranspose2DNode::VT_STRIDE_W, stride_w, 1); + } + void add_padding_h(int32_t padding_h) { + fbb_.AddElement(ConvTranspose2DNode::VT_PADDING_H, padding_h, 0); + } + void add_padding_w(int32_t padding_w) { + fbb_.AddElement(ConvTranspose2DNode::VT_PADDING_W, padding_w, 0); + } + void add_dilation_h(int32_t dilation_h) { + fbb_.AddElement(ConvTranspose2DNode::VT_DILATION_H, dilation_h, 1); + } + void add_dilation_w(int32_t dilation_w) { + fbb_.AddElement(ConvTranspose2DNode::VT_DILATION_W, dilation_w, 1); + } + void add_output_padding_h(int32_t output_padding_h) { + fbb_.AddElement(ConvTranspose2DNode::VT_OUTPUT_PADDING_H, output_padding_h, 0); + } + void add_output_padding_w(int32_t output_padding_w) { + fbb_.AddElement(ConvTranspose2DNode::VT_OUTPUT_PADDING_W, output_padding_w, 0); + } + void add_groups(int32_t groups) { + fbb_.AddElement(ConvTranspose2DNode::VT_GROUPS, groups, 1); + } + explicit ConvTranspose2DNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ConvTranspose2DNode::VT_X); + fbb_.Required(o, ConvTranspose2DNode::VT_W); + fbb_.Required(o, ConvTranspose2DNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConvTranspose2DNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t stride_h = 1, + int32_t stride_w = 1, + int32_t padding_h = 0, + int32_t padding_w = 0, + int32_t dilation_h = 1, + int32_t dilation_w = 1, + int32_t output_padding_h = 0, + int32_t output_padding_w = 0, + int32_t groups = 1) { + ConvTranspose2DNodeBuilder builder_(_fbb); + builder_.add_groups(groups); + builder_.add_output_padding_w(output_padding_w); + builder_.add_output_padding_h(output_padding_h); + builder_.add_dilation_w(dilation_w); + builder_.add_dilation_h(dilation_h); + builder_.add_padding_w(padding_w); + builder_.add_padding_h(padding_h); + builder_.add_stride_w(stride_w); + builder_.add_stride_h(stride_h); + builder_.add_out(out); + builder_.add_w(w); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ConvTranspose3DNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ConvTranspose3DNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_OUT = 8, + VT_STRIDE_D = 10, + VT_STRIDE_H = 12, + VT_STRIDE_W = 14, + VT_PADDING_D = 16, + VT_PADDING_H = 18, + VT_PADDING_W = 20, + VT_DILATION_D = 22, + VT_DILATION_H = 24, + VT_DILATION_W = 26, + VT_OUTPUT_PADDING_D = 28, + VT_OUTPUT_PADDING_H = 30, + VT_OUTPUT_PADDING_W = 32, + VT_GROUPS = 34 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t stride_d() const { + return GetField(VT_STRIDE_D, 1); + } + int32_t stride_h() const { + return GetField(VT_STRIDE_H, 1); + } + int32_t stride_w() const { + return GetField(VT_STRIDE_W, 1); + } + int32_t padding_d() const { + return GetField(VT_PADDING_D, 0); + } + int32_t padding_h() const { + return GetField(VT_PADDING_H, 0); + } + int32_t padding_w() const { + return GetField(VT_PADDING_W, 0); + } + int32_t dilation_d() const { + return GetField(VT_DILATION_D, 1); + } + int32_t dilation_h() const { + return GetField(VT_DILATION_H, 1); + } + int32_t dilation_w() const { + return GetField(VT_DILATION_W, 1); + } + int32_t output_padding_d() const { + return GetField(VT_OUTPUT_PADDING_D, 0); + } + int32_t output_padding_h() const { + return GetField(VT_OUTPUT_PADDING_H, 0); + } + int32_t output_padding_w() const { + return GetField(VT_OUTPUT_PADDING_W, 0); + } + int32_t groups() const { + return GetField(VT_GROUPS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_STRIDE_D, 4) && + VerifyField(verifier, VT_STRIDE_H, 4) && + VerifyField(verifier, VT_STRIDE_W, 4) && + VerifyField(verifier, VT_PADDING_D, 4) && + VerifyField(verifier, VT_PADDING_H, 4) && + VerifyField(verifier, VT_PADDING_W, 4) && + VerifyField(verifier, VT_DILATION_D, 4) && + VerifyField(verifier, VT_DILATION_H, 4) && + VerifyField(verifier, VT_DILATION_W, 4) && + VerifyField(verifier, VT_OUTPUT_PADDING_D, 4) && + VerifyField(verifier, VT_OUTPUT_PADDING_H, 4) && + VerifyField(verifier, VT_OUTPUT_PADDING_W, 4) && + VerifyField(verifier, VT_GROUPS, 4) && + verifier.EndTable(); + } +}; + +struct ConvTranspose3DNodeBuilder { + typedef ConvTranspose3DNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ConvTranspose3DNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(ConvTranspose3DNode::VT_W, w); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ConvTranspose3DNode::VT_OUT, out); + } + void add_stride_d(int32_t stride_d) { + fbb_.AddElement(ConvTranspose3DNode::VT_STRIDE_D, stride_d, 1); + } + void add_stride_h(int32_t stride_h) { + fbb_.AddElement(ConvTranspose3DNode::VT_STRIDE_H, stride_h, 1); + } + void add_stride_w(int32_t stride_w) { + fbb_.AddElement(ConvTranspose3DNode::VT_STRIDE_W, stride_w, 1); + } + void add_padding_d(int32_t padding_d) { + fbb_.AddElement(ConvTranspose3DNode::VT_PADDING_D, padding_d, 0); + } + void add_padding_h(int32_t padding_h) { + fbb_.AddElement(ConvTranspose3DNode::VT_PADDING_H, padding_h, 0); + } + void add_padding_w(int32_t padding_w) { + fbb_.AddElement(ConvTranspose3DNode::VT_PADDING_W, padding_w, 0); + } + void add_dilation_d(int32_t dilation_d) { + fbb_.AddElement(ConvTranspose3DNode::VT_DILATION_D, dilation_d, 1); + } + void add_dilation_h(int32_t dilation_h) { + fbb_.AddElement(ConvTranspose3DNode::VT_DILATION_H, dilation_h, 1); + } + void add_dilation_w(int32_t dilation_w) { + fbb_.AddElement(ConvTranspose3DNode::VT_DILATION_W, dilation_w, 1); + } + void add_output_padding_d(int32_t output_padding_d) { + fbb_.AddElement(ConvTranspose3DNode::VT_OUTPUT_PADDING_D, output_padding_d, 0); + } + void add_output_padding_h(int32_t output_padding_h) { + fbb_.AddElement(ConvTranspose3DNode::VT_OUTPUT_PADDING_H, output_padding_h, 0); + } + void add_output_padding_w(int32_t output_padding_w) { + fbb_.AddElement(ConvTranspose3DNode::VT_OUTPUT_PADDING_W, output_padding_w, 0); + } + void add_groups(int32_t groups) { + fbb_.AddElement(ConvTranspose3DNode::VT_GROUPS, groups, 1); + } + explicit ConvTranspose3DNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ConvTranspose3DNode::VT_X); + fbb_.Required(o, ConvTranspose3DNode::VT_W); + fbb_.Required(o, ConvTranspose3DNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConvTranspose3DNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t stride_d = 1, + int32_t stride_h = 1, + int32_t stride_w = 1, + int32_t padding_d = 0, + int32_t padding_h = 0, + int32_t padding_w = 0, + int32_t dilation_d = 1, + int32_t dilation_h = 1, + int32_t dilation_w = 1, + int32_t output_padding_d = 0, + int32_t output_padding_h = 0, + int32_t output_padding_w = 0, + int32_t groups = 1) { + ConvTranspose3DNodeBuilder builder_(_fbb); + builder_.add_groups(groups); + builder_.add_output_padding_w(output_padding_w); + builder_.add_output_padding_h(output_padding_h); + builder_.add_output_padding_d(output_padding_d); + builder_.add_dilation_w(dilation_w); + builder_.add_dilation_h(dilation_h); + builder_.add_dilation_d(dilation_d); + builder_.add_padding_w(padding_w); + builder_.add_padding_h(padding_h); + builder_.add_padding_d(padding_d); + builder_.add_stride_w(stride_w); + builder_.add_stride_h(stride_h); + builder_.add_stride_d(stride_d); + builder_.add_out(out); + builder_.add_w(w); + builder_.add_x(x); + return builder_.Finish(); +} + +struct GeluNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef GeluNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_APPROXIMATE = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::String *approximate() const { + return GetPointer(VT_APPROXIMATE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_APPROXIMATE) && + verifier.VerifyString(approximate()) && + verifier.EndTable(); + } +}; + +struct GeluNodeBuilder { + typedef GeluNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(GeluNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(GeluNode::VT_OUT, out); + } + void add_approximate(::flatbuffers::Offset<::flatbuffers::String> approximate) { + fbb_.AddOffset(GeluNode::VT_APPROXIMATE, approximate); + } + explicit GeluNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, GeluNode::VT_X); + fbb_.Required(o, GeluNode::VT_OUT); + fbb_.Required(o, GeluNode::VT_APPROXIMATE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateGeluNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::String> approximate = 0) { + GeluNodeBuilder builder_(_fbb); + builder_.add_approximate(approximate); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateGeluNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const char *approximate = nullptr) { + auto approximate__ = approximate ? _fbb.CreateString(approximate) : 0; + return mlx_delegate::CreateGeluNode( + _fbb, + x, + out, + approximate__); +} + +struct ARangeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ARangeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_OUT = 4, + VT_START = 6, + VT_STOP = 8, + VT_STEP = 10, + VT_SCALAR_TYPE = 12 + }; + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *start() const { + return GetPointer(VT_START); + } + const mlx_delegate::IntOrVid *stop() const { + return GetPointer(VT_STOP); + } + const mlx_delegate::IntOrVid *step() const { + return GetPointer(VT_STEP); + } + ::flatbuffers::Optional scalar_type() const { + return GetOptional(VT_SCALAR_TYPE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_START) && + verifier.VerifyTable(start()) && + VerifyOffsetRequired(verifier, VT_STOP) && + verifier.VerifyTable(stop()) && + VerifyOffsetRequired(verifier, VT_STEP) && + verifier.VerifyTable(step()) && + VerifyField(verifier, VT_SCALAR_TYPE, 1) && + verifier.EndTable(); + } +}; + +struct ARangeNodeBuilder { + typedef ARangeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ARangeNode::VT_OUT, out); + } + void add_start(::flatbuffers::Offset start) { + fbb_.AddOffset(ARangeNode::VT_START, start); + } + void add_stop(::flatbuffers::Offset stop) { + fbb_.AddOffset(ARangeNode::VT_STOP, stop); + } + void add_step(::flatbuffers::Offset step) { + fbb_.AddOffset(ARangeNode::VT_STEP, step); + } + void add_scalar_type(int8_t scalar_type) { + fbb_.AddElement(ARangeNode::VT_SCALAR_TYPE, scalar_type); + } + explicit ARangeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ARangeNode::VT_OUT); + fbb_.Required(o, ARangeNode::VT_START); + fbb_.Required(o, ARangeNode::VT_STOP); + fbb_.Required(o, ARangeNode::VT_STEP); + return o; + } +}; + +inline ::flatbuffers::Offset CreateARangeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset start = 0, + ::flatbuffers::Offset stop = 0, + ::flatbuffers::Offset step = 0, + ::flatbuffers::Optional scalar_type = ::flatbuffers::nullopt) { + ARangeNodeBuilder builder_(_fbb); + builder_.add_step(step); + builder_.add_stop(stop); + builder_.add_start(start); + builder_.add_out(out); + if(scalar_type) { builder_.add_scalar_type(*scalar_type); } + return builder_.Finish(); +} + +struct SiluNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SiluNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SiluNodeBuilder { + typedef SiluNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SiluNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SiluNode::VT_OUT, out); + } + explicit SiluNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SiluNode::VT_X); + fbb_.Required(o, SiluNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSiluNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SiluNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SigmoidNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SigmoidNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SigmoidNodeBuilder { + typedef SigmoidNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SigmoidNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SigmoidNode::VT_OUT, out); + } + explicit SigmoidNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SigmoidNode::VT_X); + fbb_.Required(o, SigmoidNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSigmoidNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SigmoidNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct TanhNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TanhNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct TanhNodeBuilder { + typedef TanhNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TanhNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TanhNode::VT_OUT, out); + } + explicit TanhNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TanhNode::VT_X); + fbb_.Required(o, TanhNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTanhNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + TanhNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SqueezeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SqueezeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_DIMS = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *dims() const { + return GetPointer *>(VT_DIMS); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_DIMS) && + verifier.VerifyVector(dims()) && + verifier.EndTable(); + } +}; + +struct SqueezeNodeBuilder { + typedef SqueezeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SqueezeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SqueezeNode::VT_OUT, out); + } + void add_dims(::flatbuffers::Offset<::flatbuffers::Vector> dims) { + fbb_.AddOffset(SqueezeNode::VT_DIMS, dims); + } + explicit SqueezeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SqueezeNode::VT_X); + fbb_.Required(o, SqueezeNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSqueezeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> dims = 0) { + SqueezeNodeBuilder builder_(_fbb); + builder_.add_dims(dims); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateSqueezeNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *dims = nullptr) { + auto dims__ = dims ? _fbb.CreateVector(*dims) : 0; + return mlx_delegate::CreateSqueezeNode( + _fbb, + x, + out, + dims__); +} + +struct SplitNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SplitNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUTS = 6, + VT_SIZES = 8, + VT_AXIS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const ::flatbuffers::Vector *outs() const { + return GetPointer *>(VT_OUTS); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *sizes() const { + return GetPointer> *>(VT_SIZES); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyOffsetRequired(verifier, VT_OUTS) && + verifier.VerifyVector(outs()) && + VerifyOffsetRequired(verifier, VT_SIZES) && + verifier.VerifyVector(sizes()) && + verifier.VerifyVectorOfTables(sizes()) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct SplitNodeBuilder { + typedef SplitNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SplitNode::VT_X, x); + } + void add_outs(::flatbuffers::Offset<::flatbuffers::Vector> outs) { + fbb_.AddOffset(SplitNode::VT_OUTS, outs); + } + void add_sizes(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> sizes) { + fbb_.AddOffset(SplitNode::VT_SIZES, sizes); + } + void add_axis(int32_t axis) { + fbb_.AddElement(SplitNode::VT_AXIS, axis, 0); + } + explicit SplitNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SplitNode::VT_X); + fbb_.Required(o, SplitNode::VT_OUTS); + fbb_.Required(o, SplitNode::VT_SIZES); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSplitNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> outs = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> sizes = 0, + int32_t axis = 0) { + SplitNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_sizes(sizes); + builder_.add_outs(outs); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateSplitNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const std::vector *outs = nullptr, + const std::vector<::flatbuffers::Offset> *sizes = nullptr, + int32_t axis = 0) { + auto outs__ = outs ? _fbb.CreateVectorOfStructs(*outs) : 0; + auto sizes__ = sizes ? _fbb.CreateVector<::flatbuffers::Offset>(*sizes) : 0; + return mlx_delegate::CreateSplitNode( + _fbb, + x, + outs__, + sizes__, + axis); +} + +struct RsqrtNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RsqrtNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct RsqrtNodeBuilder { + typedef RsqrtNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(RsqrtNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RsqrtNode::VT_OUT, out); + } + explicit RsqrtNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RsqrtNode::VT_X); + fbb_.Required(o, RsqrtNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRsqrtNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + RsqrtNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct MaximumNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MaximumNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct MaximumNodeBuilder { + typedef MaximumNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(MaximumNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(MaximumNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MaximumNode::VT_OUT, out); + } + explicit MaximumNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MaximumNode::VT_A); + fbb_.Required(o, MaximumNode::VT_B); + fbb_.Required(o, MaximumNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMaximumNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + MaximumNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct MinimumNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MinimumNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct MinimumNodeBuilder { + typedef MinimumNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(MinimumNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(MinimumNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MinimumNode::VT_OUT, out); + } + explicit MinimumNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MinimumNode::VT_A); + fbb_.Required(o, MinimumNode::VT_B); + fbb_.Required(o, MinimumNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMinimumNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + MinimumNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct LogNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LogNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LogNodeBuilder { + typedef LogNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(LogNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LogNode::VT_OUT, out); + } + explicit LogNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LogNode::VT_X); + fbb_.Required(o, LogNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLogNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LogNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SoftmaxNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SoftmaxNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8, + VT_PRECISE = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool precise() const { + return GetField(VT_PRECISE, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + VerifyField(verifier, VT_PRECISE, 1) && + verifier.EndTable(); + } +}; + +struct SoftmaxNodeBuilder { + typedef SoftmaxNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SoftmaxNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SoftmaxNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(SoftmaxNode::VT_AXIS, axis, 0); + } + void add_precise(bool precise) { + fbb_.AddElement(SoftmaxNode::VT_PRECISE, static_cast(precise), 0); + } + explicit SoftmaxNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SoftmaxNode::VT_X); + fbb_.Required(o, SoftmaxNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSoftmaxNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0, + bool precise = false) { + SoftmaxNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_precise(precise); + return builder_.Finish(); +} + +struct BroadcastToNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef BroadcastToNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_SHAPE = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *shape() const { + return GetPointer> *>(VT_SHAPE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_SHAPE) && + verifier.VerifyVector(shape()) && + verifier.VerifyVectorOfTables(shape()) && + verifier.EndTable(); + } +}; + +struct BroadcastToNodeBuilder { + typedef BroadcastToNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(BroadcastToNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(BroadcastToNode::VT_OUT, out); + } + void add_shape(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape) { + fbb_.AddOffset(BroadcastToNode::VT_SHAPE, shape); + } + explicit BroadcastToNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, BroadcastToNode::VT_X); + fbb_.Required(o, BroadcastToNode::VT_OUT); + fbb_.Required(o, BroadcastToNode::VT_SHAPE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateBroadcastToNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape = 0) { + BroadcastToNodeBuilder builder_(_fbb); + builder_.add_shape(shape); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateBroadcastToNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *shape = nullptr) { + auto shape__ = shape ? _fbb.CreateVector<::flatbuffers::Offset>(*shape) : 0; + return mlx_delegate::CreateBroadcastToNode( + _fbb, + x, + out, + shape__); +} + +struct PadNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef PadNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_PAD_WIDTH = 8, + VT_MODE = 10, + VT_CONSTANT_VALUE = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *pad_width() const { + return GetPointer> *>(VT_PAD_WIDTH); + } + const ::flatbuffers::String *mode() const { + return GetPointer(VT_MODE); + } + float constant_value() const { + return GetField(VT_CONSTANT_VALUE, 0.0f); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_PAD_WIDTH) && + verifier.VerifyVector(pad_width()) && + verifier.VerifyVectorOfTables(pad_width()) && + VerifyOffsetRequired(verifier, VT_MODE) && + verifier.VerifyString(mode()) && + VerifyField(verifier, VT_CONSTANT_VALUE, 4) && + verifier.EndTable(); + } +}; + +struct PadNodeBuilder { + typedef PadNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(PadNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(PadNode::VT_OUT, out); + } + void add_pad_width(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> pad_width) { + fbb_.AddOffset(PadNode::VT_PAD_WIDTH, pad_width); + } + void add_mode(::flatbuffers::Offset<::flatbuffers::String> mode) { + fbb_.AddOffset(PadNode::VT_MODE, mode); + } + void add_constant_value(float constant_value) { + fbb_.AddElement(PadNode::VT_CONSTANT_VALUE, constant_value, 0.0f); + } + explicit PadNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, PadNode::VT_X); + fbb_.Required(o, PadNode::VT_OUT); + fbb_.Required(o, PadNode::VT_PAD_WIDTH); + fbb_.Required(o, PadNode::VT_MODE); + return o; + } +}; + +inline ::flatbuffers::Offset CreatePadNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> pad_width = 0, + ::flatbuffers::Offset<::flatbuffers::String> mode = 0, + float constant_value = 0.0f) { + PadNodeBuilder builder_(_fbb); + builder_.add_constant_value(constant_value); + builder_.add_mode(mode); + builder_.add_pad_width(pad_width); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreatePadNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *pad_width = nullptr, + const char *mode = nullptr, + float constant_value = 0.0f) { + auto pad_width__ = pad_width ? _fbb.CreateVector<::flatbuffers::Offset>(*pad_width) : 0; + auto mode__ = mode ? _fbb.CreateString(mode) : 0; + return mlx_delegate::CreatePadNode( + _fbb, + x, + out, + pad_width__, + mode__, + constant_value); +} + +struct WhereNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef WhereNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_CONDITION = 4, + VT_X = 6, + VT_Y = 8, + VT_OUT = 10 + }; + const mlx_delegate::Tid *condition() const { + return GetStruct(VT_CONDITION); + } + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *y() const { + return GetStruct(VT_Y); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_CONDITION, 4) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_Y, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct WhereNodeBuilder { + typedef WhereNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_condition(const mlx_delegate::Tid *condition) { + fbb_.AddStruct(WhereNode::VT_CONDITION, condition); + } + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(WhereNode::VT_X, x); + } + void add_y(const mlx_delegate::Tid *y) { + fbb_.AddStruct(WhereNode::VT_Y, y); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(WhereNode::VT_OUT, out); + } + explicit WhereNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, WhereNode::VT_CONDITION); + fbb_.Required(o, WhereNode::VT_X); + fbb_.Required(o, WhereNode::VT_Y); + fbb_.Required(o, WhereNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateWhereNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *condition = nullptr, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *y = nullptr, + const mlx_delegate::Tid *out = nullptr) { + WhereNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_y(y); + builder_.add_x(x); + builder_.add_condition(condition); + return builder_.Finish(); +} + +struct ReshapeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ReshapeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_SHAPE = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *shape() const { + return GetPointer> *>(VT_SHAPE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_SHAPE) && + verifier.VerifyVector(shape()) && + verifier.VerifyVectorOfTables(shape()) && + verifier.EndTable(); + } +}; + +struct ReshapeNodeBuilder { + typedef ReshapeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ReshapeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ReshapeNode::VT_OUT, out); + } + void add_shape(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape) { + fbb_.AddOffset(ReshapeNode::VT_SHAPE, shape); + } + explicit ReshapeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ReshapeNode::VT_X); + fbb_.Required(o, ReshapeNode::VT_OUT); + fbb_.Required(o, ReshapeNode::VT_SHAPE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateReshapeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape = 0) { + ReshapeNodeBuilder builder_(_fbb); + builder_.add_shape(shape); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateReshapeNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *shape = nullptr) { + auto shape__ = shape ? _fbb.CreateVector<::flatbuffers::Offset>(*shape) : 0; + return mlx_delegate::CreateReshapeNode( + _fbb, + x, + out, + shape__); +} + +struct TransposeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TransposeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_PERM = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *perm() const { + return GetPointer *>(VT_PERM); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_PERM) && + verifier.VerifyVector(perm()) && + verifier.EndTable(); + } +}; + +struct TransposeNodeBuilder { + typedef TransposeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TransposeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TransposeNode::VT_OUT, out); + } + void add_perm(::flatbuffers::Offset<::flatbuffers::Vector> perm) { + fbb_.AddOffset(TransposeNode::VT_PERM, perm); + } + explicit TransposeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TransposeNode::VT_X); + fbb_.Required(o, TransposeNode::VT_OUT); + fbb_.Required(o, TransposeNode::VT_PERM); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTransposeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> perm = 0) { + TransposeNodeBuilder builder_(_fbb); + builder_.add_perm(perm); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateTransposeNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *perm = nullptr) { + auto perm__ = perm ? _fbb.CreateVector(*perm) : 0; + return mlx_delegate::CreateTransposeNode( + _fbb, + x, + out, + perm__); +} + +struct AsStridedNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AsStridedNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_SHAPE = 8, + VT_STRIDES = 10, + VT_OFFSET = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *shape() const { + return GetPointer> *>(VT_SHAPE); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *strides() const { + return GetPointer> *>(VT_STRIDES); + } + uint64_t offset() const { + return GetField(VT_OFFSET, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_SHAPE) && + verifier.VerifyVector(shape()) && + verifier.VerifyVectorOfTables(shape()) && + VerifyOffsetRequired(verifier, VT_STRIDES) && + verifier.VerifyVector(strides()) && + verifier.VerifyVectorOfTables(strides()) && + VerifyField(verifier, VT_OFFSET, 8) && + verifier.EndTable(); + } +}; + +struct AsStridedNodeBuilder { + typedef AsStridedNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(AsStridedNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AsStridedNode::VT_OUT, out); + } + void add_shape(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape) { + fbb_.AddOffset(AsStridedNode::VT_SHAPE, shape); + } + void add_strides(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> strides) { + fbb_.AddOffset(AsStridedNode::VT_STRIDES, strides); + } + void add_offset(uint64_t offset) { + fbb_.AddElement(AsStridedNode::VT_OFFSET, offset, 0); + } + explicit AsStridedNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AsStridedNode::VT_X); + fbb_.Required(o, AsStridedNode::VT_OUT); + fbb_.Required(o, AsStridedNode::VT_SHAPE); + fbb_.Required(o, AsStridedNode::VT_STRIDES); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAsStridedNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> strides = 0, + uint64_t offset = 0) { + AsStridedNodeBuilder builder_(_fbb); + builder_.add_offset(offset); + builder_.add_strides(strides); + builder_.add_shape(shape); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateAsStridedNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *shape = nullptr, + const std::vector<::flatbuffers::Offset> *strides = nullptr, + uint64_t offset = 0) { + auto shape__ = shape ? _fbb.CreateVector<::flatbuffers::Offset>(*shape) : 0; + auto strides__ = strides ? _fbb.CreateVector<::flatbuffers::Offset>(*strides) : 0; + return mlx_delegate::CreateAsStridedNode( + _fbb, + x, + out, + shape__, + strides__, + offset); +} + +struct ContiguousNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ContiguousNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ContiguousNodeBuilder { + typedef ContiguousNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ContiguousNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ContiguousNode::VT_OUT, out); + } + explicit ContiguousNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ContiguousNode::VT_X); + fbb_.Required(o, ContiguousNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateContiguousNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ContiguousNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct GatherNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef GatherNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_INDICES = 6, + VT_OUT = 8, + VT_AXES = 10, + VT_SLICE_SIZES = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const ::flatbuffers::Vector *indices() const { + return GetPointer *>(VT_INDICES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + const ::flatbuffers::Vector *slice_sizes() const { + return GetPointer *>(VT_SLICE_SIZES); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyOffsetRequired(verifier, VT_INDICES) && + verifier.VerifyVector(indices()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyOffsetRequired(verifier, VT_SLICE_SIZES) && + verifier.VerifyVector(slice_sizes()) && + verifier.EndTable(); + } +}; + +struct GatherNodeBuilder { + typedef GatherNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(GatherNode::VT_X, x); + } + void add_indices(::flatbuffers::Offset<::flatbuffers::Vector> indices) { + fbb_.AddOffset(GatherNode::VT_INDICES, indices); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(GatherNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(GatherNode::VT_AXES, axes); + } + void add_slice_sizes(::flatbuffers::Offset<::flatbuffers::Vector> slice_sizes) { + fbb_.AddOffset(GatherNode::VT_SLICE_SIZES, slice_sizes); + } + explicit GatherNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, GatherNode::VT_X); + fbb_.Required(o, GatherNode::VT_INDICES); + fbb_.Required(o, GatherNode::VT_OUT); + fbb_.Required(o, GatherNode::VT_AXES); + fbb_.Required(o, GatherNode::VT_SLICE_SIZES); + return o; + } +}; + +inline ::flatbuffers::Offset CreateGatherNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> indices = 0, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> slice_sizes = 0) { + GatherNodeBuilder builder_(_fbb); + builder_.add_slice_sizes(slice_sizes); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_indices(indices); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateGatherNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const std::vector *indices = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + const std::vector *slice_sizes = nullptr) { + auto indices__ = indices ? _fbb.CreateVectorOfStructs(*indices) : 0; + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + auto slice_sizes__ = slice_sizes ? _fbb.CreateVector(*slice_sizes) : 0; + return mlx_delegate::CreateGatherNode( + _fbb, + x, + indices__, + out, + axes__, + slice_sizes__); +} + +struct SliceNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SliceNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8, + VT_START = 10, + VT_STOP = 12, + VT_STEP = 14 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *axis() const { + return GetPointer(VT_AXIS); + } + const mlx_delegate::IntOrVid *start() const { + return GetPointer(VT_START); + } + const mlx_delegate::IntOrVid *stop() const { + return GetPointer(VT_STOP); + } + int32_t step() const { + return GetField(VT_STEP, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_AXIS) && + verifier.VerifyTable(axis()) && + VerifyOffsetRequired(verifier, VT_START) && + verifier.VerifyTable(start()) && + VerifyOffsetRequired(verifier, VT_STOP) && + verifier.VerifyTable(stop()) && + VerifyField(verifier, VT_STEP, 4) && + verifier.EndTable(); + } +}; + +struct SliceNodeBuilder { + typedef SliceNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SliceNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SliceNode::VT_OUT, out); + } + void add_axis(::flatbuffers::Offset axis) { + fbb_.AddOffset(SliceNode::VT_AXIS, axis); + } + void add_start(::flatbuffers::Offset start) { + fbb_.AddOffset(SliceNode::VT_START, start); + } + void add_stop(::flatbuffers::Offset stop) { + fbb_.AddOffset(SliceNode::VT_STOP, stop); + } + void add_step(int32_t step) { + fbb_.AddElement(SliceNode::VT_STEP, step, 1); + } + explicit SliceNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SliceNode::VT_X); + fbb_.Required(o, SliceNode::VT_OUT); + fbb_.Required(o, SliceNode::VT_AXIS); + fbb_.Required(o, SliceNode::VT_START); + fbb_.Required(o, SliceNode::VT_STOP); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSliceNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset axis = 0, + ::flatbuffers::Offset start = 0, + ::flatbuffers::Offset stop = 0, + int32_t step = 1) { + SliceNodeBuilder builder_(_fbb); + builder_.add_step(step); + builder_.add_stop(stop); + builder_.add_start(start); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct AsTypeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AsTypeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_SCALAR_TYPE = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int8_t scalar_type() const { + return GetField(VT_SCALAR_TYPE, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_SCALAR_TYPE, 1) && + verifier.EndTable(); + } +}; + +struct AsTypeNodeBuilder { + typedef AsTypeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(AsTypeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AsTypeNode::VT_OUT, out); + } + void add_scalar_type(int8_t scalar_type) { + fbb_.AddElement(AsTypeNode::VT_SCALAR_TYPE, scalar_type, 0); + } + explicit AsTypeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AsTypeNode::VT_X); + fbb_.Required(o, AsTypeNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAsTypeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int8_t scalar_type = 0) { + AsTypeNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_scalar_type(scalar_type); + return builder_.Finish(); +} + +struct QuantizedMatmulNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef QuantizedMatmulNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_SCALES = 8, + VT_OUT = 10, + VT_BIASES = 12, + VT_GROUP_SIZE = 14, + VT_BITS = 16, + VT_MODE = 18, + VT_TRANSPOSE = 20 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *scales() const { + return GetStruct(VT_SCALES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::Tid *biases() const { + return GetStruct(VT_BIASES); + } + int32_t group_size() const { + return GetField(VT_GROUP_SIZE, 0); + } + int32_t bits() const { + return GetField(VT_BITS, 0); + } + const ::flatbuffers::String *mode() const { + return GetPointer(VT_MODE); + } + bool transpose() const { + return GetField(VT_TRANSPOSE, 1) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_SCALES, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_BIASES, 4) && + VerifyField(verifier, VT_GROUP_SIZE, 4) && + VerifyField(verifier, VT_BITS, 4) && + VerifyOffsetRequired(verifier, VT_MODE) && + verifier.VerifyString(mode()) && + VerifyField(verifier, VT_TRANSPOSE, 1) && + verifier.EndTable(); + } +}; + +struct QuantizedMatmulNodeBuilder { + typedef QuantizedMatmulNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(QuantizedMatmulNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(QuantizedMatmulNode::VT_W, w); + } + void add_scales(const mlx_delegate::Tid *scales) { + fbb_.AddStruct(QuantizedMatmulNode::VT_SCALES, scales); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(QuantizedMatmulNode::VT_OUT, out); + } + void add_biases(const mlx_delegate::Tid *biases) { + fbb_.AddStruct(QuantizedMatmulNode::VT_BIASES, biases); + } + void add_group_size(int32_t group_size) { + fbb_.AddElement(QuantizedMatmulNode::VT_GROUP_SIZE, group_size, 0); + } + void add_bits(int32_t bits) { + fbb_.AddElement(QuantizedMatmulNode::VT_BITS, bits, 0); + } + void add_mode(::flatbuffers::Offset<::flatbuffers::String> mode) { + fbb_.AddOffset(QuantizedMatmulNode::VT_MODE, mode); + } + void add_transpose(bool transpose) { + fbb_.AddElement(QuantizedMatmulNode::VT_TRANSPOSE, static_cast(transpose), 1); + } + explicit QuantizedMatmulNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, QuantizedMatmulNode::VT_X); + fbb_.Required(o, QuantizedMatmulNode::VT_W); + fbb_.Required(o, QuantizedMatmulNode::VT_SCALES); + fbb_.Required(o, QuantizedMatmulNode::VT_OUT); + fbb_.Required(o, QuantizedMatmulNode::VT_MODE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateQuantizedMatmulNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *scales = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *biases = nullptr, + int32_t group_size = 0, + int32_t bits = 0, + ::flatbuffers::Offset<::flatbuffers::String> mode = 0, + bool transpose = true) { + QuantizedMatmulNodeBuilder builder_(_fbb); + builder_.add_mode(mode); + builder_.add_bits(bits); + builder_.add_group_size(group_size); + builder_.add_biases(biases); + builder_.add_out(out); + builder_.add_scales(scales); + builder_.add_w(w); + builder_.add_x(x); + builder_.add_transpose(transpose); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateQuantizedMatmulNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *scales = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *biases = nullptr, + int32_t group_size = 0, + int32_t bits = 0, + const char *mode = nullptr, + bool transpose = true) { + auto mode__ = mode ? _fbb.CreateString(mode) : 0; + return mlx_delegate::CreateQuantizedMatmulNode( + _fbb, + x, + w, + scales, + out, + biases, + group_size, + bits, + mode__, + transpose); +} + +struct ScatterAddNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ScatterAddNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_INDICES = 6, + VT_UPDATES = 8, + VT_OUT = 10, + VT_AXIS = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *indices() const { + return GetStruct(VT_INDICES); + } + const mlx_delegate::Tid *updates() const { + return GetStruct(VT_UPDATES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_INDICES, 4) && + VerifyFieldRequired(verifier, VT_UPDATES, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct ScatterAddNodeBuilder { + typedef ScatterAddNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ScatterAddNode::VT_X, x); + } + void add_indices(const mlx_delegate::Tid *indices) { + fbb_.AddStruct(ScatterAddNode::VT_INDICES, indices); + } + void add_updates(const mlx_delegate::Tid *updates) { + fbb_.AddStruct(ScatterAddNode::VT_UPDATES, updates); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ScatterAddNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ScatterAddNode::VT_AXIS, axis, 0); + } + explicit ScatterAddNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ScatterAddNode::VT_X); + fbb_.Required(o, ScatterAddNode::VT_INDICES); + fbb_.Required(o, ScatterAddNode::VT_UPDATES); + fbb_.Required(o, ScatterAddNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateScatterAddNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *indices = nullptr, + const mlx_delegate::Tid *updates = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + ScatterAddNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_updates(updates); + builder_.add_indices(indices); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ConcatenateNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ConcatenateNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_TENSORS = 4, + VT_OUT = 6, + VT_AXIS = 8 + }; + const ::flatbuffers::Vector *tensors() const { + return GetPointer *>(VT_TENSORS); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_TENSORS) && + verifier.VerifyVector(tensors()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct ConcatenateNodeBuilder { + typedef ConcatenateNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_tensors(::flatbuffers::Offset<::flatbuffers::Vector> tensors) { + fbb_.AddOffset(ConcatenateNode::VT_TENSORS, tensors); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ConcatenateNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ConcatenateNode::VT_AXIS, axis, 0); + } + explicit ConcatenateNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ConcatenateNode::VT_TENSORS); + fbb_.Required(o, ConcatenateNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateConcatenateNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::Vector> tensors = 0, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + ConcatenateNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_tensors(tensors); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateConcatenateNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const std::vector *tensors = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + auto tensors__ = tensors ? _fbb.CreateVectorOfStructs(*tensors) : 0; + return mlx_delegate::CreateConcatenateNode( + _fbb, + tensors__, + out, + axis); +} + +struct FullNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef FullNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_OUT = 4, + VT_SHAPE = 6, + VT_V = 8, + VT_SCALAR_TYPE = 10 + }; + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *shape() const { + return GetPointer> *>(VT_SHAPE); + } + const mlx_delegate::FloatOrVid *v() const { + return GetPointer(VT_V); + } + int8_t scalar_type() const { + return GetField(VT_SCALAR_TYPE, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_SHAPE) && + verifier.VerifyVector(shape()) && + verifier.VerifyVectorOfTables(shape()) && + VerifyOffsetRequired(verifier, VT_V) && + verifier.VerifyTable(v()) && + VerifyField(verifier, VT_SCALAR_TYPE, 1) && + verifier.EndTable(); + } +}; + +struct FullNodeBuilder { + typedef FullNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(FullNode::VT_OUT, out); + } + void add_shape(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape) { + fbb_.AddOffset(FullNode::VT_SHAPE, shape); + } + void add_v(::flatbuffers::Offset v) { + fbb_.AddOffset(FullNode::VT_V, v); + } + void add_scalar_type(int8_t scalar_type) { + fbb_.AddElement(FullNode::VT_SCALAR_TYPE, scalar_type, 0); + } + explicit FullNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, FullNode::VT_OUT); + fbb_.Required(o, FullNode::VT_SHAPE); + fbb_.Required(o, FullNode::VT_V); + return o; + } +}; + +inline ::flatbuffers::Offset CreateFullNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape = 0, + ::flatbuffers::Offset v = 0, + int8_t scalar_type = 0) { + FullNodeBuilder builder_(_fbb); + builder_.add_v(v); + builder_.add_shape(shape); + builder_.add_out(out); + builder_.add_scalar_type(scalar_type); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateFullNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *shape = nullptr, + ::flatbuffers::Offset v = 0, + int8_t scalar_type = 0) { + auto shape__ = shape ? _fbb.CreateVector<::flatbuffers::Offset>(*shape) : 0; + return mlx_delegate::CreateFullNode( + _fbb, + out, + shape__, + v, + scalar_type); +} + +struct FullLikeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef FullLikeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_V = 8, + VT_SCALAR_TYPE = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::FloatOrVid *v() const { + return GetPointer(VT_V); + } + ::flatbuffers::Optional scalar_type() const { + return GetOptional(VT_SCALAR_TYPE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_V) && + verifier.VerifyTable(v()) && + VerifyField(verifier, VT_SCALAR_TYPE, 1) && + verifier.EndTable(); + } +}; + +struct FullLikeNodeBuilder { + typedef FullLikeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(FullLikeNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(FullLikeNode::VT_OUT, out); + } + void add_v(::flatbuffers::Offset v) { + fbb_.AddOffset(FullLikeNode::VT_V, v); + } + void add_scalar_type(int8_t scalar_type) { + fbb_.AddElement(FullLikeNode::VT_SCALAR_TYPE, scalar_type); + } + explicit FullLikeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, FullLikeNode::VT_X); + fbb_.Required(o, FullLikeNode::VT_OUT); + fbb_.Required(o, FullLikeNode::VT_V); + return o; + } +}; + +inline ::flatbuffers::Offset CreateFullLikeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset v = 0, + ::flatbuffers::Optional scalar_type = ::flatbuffers::nullopt) { + FullLikeNodeBuilder builder_(_fbb); + builder_.add_v(v); + builder_.add_out(out); + builder_.add_x(x); + if(scalar_type) { builder_.add_scalar_type(*scalar_type); } + return builder_.Finish(); +} + +struct ArgmaxNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArgmaxNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct ArgmaxNodeBuilder { + typedef ArgmaxNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArgmaxNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArgmaxNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ArgmaxNode::VT_AXIS, axis, 0); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(ArgmaxNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit ArgmaxNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArgmaxNode::VT_X); + fbb_.Required(o, ArgmaxNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArgmaxNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0, + bool keepdims = false) { + ArgmaxNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +struct SliceUpdateNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SliceUpdateNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_DST = 4, + VT_UPDATE = 6, + VT_OUT = 8, + VT_AXIS = 10, + VT_START = 12, + VT_STOP = 14, + VT_STEP = 16 + }; + const mlx_delegate::Tid *dst() const { + return GetStruct(VT_DST); + } + const mlx_delegate::Tid *update() const { + return GetStruct(VT_UPDATE); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *axis() const { + return GetPointer(VT_AXIS); + } + const mlx_delegate::IntOrVid *start() const { + return GetPointer(VT_START); + } + const mlx_delegate::IntOrVid *stop() const { + return GetPointer(VT_STOP); + } + int32_t step() const { + return GetField(VT_STEP, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_DST, 4) && + VerifyFieldRequired(verifier, VT_UPDATE, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_AXIS) && + verifier.VerifyTable(axis()) && + VerifyOffsetRequired(verifier, VT_START) && + verifier.VerifyTable(start()) && + VerifyOffsetRequired(verifier, VT_STOP) && + verifier.VerifyTable(stop()) && + VerifyField(verifier, VT_STEP, 4) && + verifier.EndTable(); + } +}; + +struct SliceUpdateNodeBuilder { + typedef SliceUpdateNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_dst(const mlx_delegate::Tid *dst) { + fbb_.AddStruct(SliceUpdateNode::VT_DST, dst); + } + void add_update(const mlx_delegate::Tid *update) { + fbb_.AddStruct(SliceUpdateNode::VT_UPDATE, update); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SliceUpdateNode::VT_OUT, out); + } + void add_axis(::flatbuffers::Offset axis) { + fbb_.AddOffset(SliceUpdateNode::VT_AXIS, axis); + } + void add_start(::flatbuffers::Offset start) { + fbb_.AddOffset(SliceUpdateNode::VT_START, start); + } + void add_stop(::flatbuffers::Offset stop) { + fbb_.AddOffset(SliceUpdateNode::VT_STOP, stop); + } + void add_step(int32_t step) { + fbb_.AddElement(SliceUpdateNode::VT_STEP, step, 1); + } + explicit SliceUpdateNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SliceUpdateNode::VT_DST); + fbb_.Required(o, SliceUpdateNode::VT_UPDATE); + fbb_.Required(o, SliceUpdateNode::VT_OUT); + fbb_.Required(o, SliceUpdateNode::VT_AXIS); + fbb_.Required(o, SliceUpdateNode::VT_START); + fbb_.Required(o, SliceUpdateNode::VT_STOP); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSliceUpdateNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *dst = nullptr, + const mlx_delegate::Tid *update = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset axis = 0, + ::flatbuffers::Offset start = 0, + ::flatbuffers::Offset stop = 0, + int32_t step = 1) { + SliceUpdateNodeBuilder builder_(_fbb); + builder_.add_step(step); + builder_.add_stop(stop); + builder_.add_start(start); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_update(update); + builder_.add_dst(dst); + return builder_.Finish(); +} + +struct IndexCopyNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef IndexCopyNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_DST = 4, + VT_UPDATE = 6, + VT_INDICES = 8, + VT_OUT = 10, + VT_AXIS = 12 + }; + const mlx_delegate::Tid *dst() const { + return GetStruct(VT_DST); + } + const mlx_delegate::Tid *update() const { + return GetStruct(VT_UPDATE); + } + const mlx_delegate::Tid *indices() const { + return GetStruct(VT_INDICES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_DST, 4) && + VerifyFieldRequired(verifier, VT_UPDATE, 4) && + VerifyFieldRequired(verifier, VT_INDICES, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct IndexCopyNodeBuilder { + typedef IndexCopyNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_dst(const mlx_delegate::Tid *dst) { + fbb_.AddStruct(IndexCopyNode::VT_DST, dst); + } + void add_update(const mlx_delegate::Tid *update) { + fbb_.AddStruct(IndexCopyNode::VT_UPDATE, update); + } + void add_indices(const mlx_delegate::Tid *indices) { + fbb_.AddStruct(IndexCopyNode::VT_INDICES, indices); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(IndexCopyNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(IndexCopyNode::VT_AXIS, axis, 0); + } + explicit IndexCopyNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, IndexCopyNode::VT_DST); + fbb_.Required(o, IndexCopyNode::VT_UPDATE); + fbb_.Required(o, IndexCopyNode::VT_INDICES); + fbb_.Required(o, IndexCopyNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateIndexCopyNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *dst = nullptr, + const mlx_delegate::Tid *update = nullptr, + const mlx_delegate::Tid *indices = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + IndexCopyNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_indices(indices); + builder_.add_update(update); + builder_.add_dst(dst); + return builder_.Finish(); +} + +struct DequantizeNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef DequantizeNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_W = 4, + VT_SCALES = 6, + VT_OUT = 8, + VT_BIASES = 10, + VT_GROUP_SIZE = 12, + VT_BITS = 14, + VT_MODE = 16, + VT_GLOBAL_SCALE = 18, + VT_DTYPE = 20 + }; + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *scales() const { + return GetStruct(VT_SCALES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::Tid *biases() const { + return GetStruct(VT_BIASES); + } + int32_t group_size() const { + return GetField(VT_GROUP_SIZE, 0); + } + int32_t bits() const { + return GetField(VT_BITS, 0); + } + const ::flatbuffers::String *mode() const { + return GetPointer(VT_MODE); + } + const mlx_delegate::Tid *global_scale() const { + return GetStruct(VT_GLOBAL_SCALE); + } + ::flatbuffers::Optional dtype() const { + return GetOptional(VT_DTYPE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_SCALES, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_BIASES, 4) && + VerifyField(verifier, VT_GROUP_SIZE, 4) && + VerifyField(verifier, VT_BITS, 4) && + VerifyOffsetRequired(verifier, VT_MODE) && + verifier.VerifyString(mode()) && + VerifyField(verifier, VT_GLOBAL_SCALE, 4) && + VerifyField(verifier, VT_DTYPE, 1) && + verifier.EndTable(); + } +}; + +struct DequantizeNodeBuilder { + typedef DequantizeNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(DequantizeNode::VT_W, w); + } + void add_scales(const mlx_delegate::Tid *scales) { + fbb_.AddStruct(DequantizeNode::VT_SCALES, scales); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(DequantizeNode::VT_OUT, out); + } + void add_biases(const mlx_delegate::Tid *biases) { + fbb_.AddStruct(DequantizeNode::VT_BIASES, biases); + } + void add_group_size(int32_t group_size) { + fbb_.AddElement(DequantizeNode::VT_GROUP_SIZE, group_size, 0); + } + void add_bits(int32_t bits) { + fbb_.AddElement(DequantizeNode::VT_BITS, bits, 0); + } + void add_mode(::flatbuffers::Offset<::flatbuffers::String> mode) { + fbb_.AddOffset(DequantizeNode::VT_MODE, mode); + } + void add_global_scale(const mlx_delegate::Tid *global_scale) { + fbb_.AddStruct(DequantizeNode::VT_GLOBAL_SCALE, global_scale); + } + void add_dtype(int8_t dtype) { + fbb_.AddElement(DequantizeNode::VT_DTYPE, dtype); + } + explicit DequantizeNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, DequantizeNode::VT_W); + fbb_.Required(o, DequantizeNode::VT_SCALES); + fbb_.Required(o, DequantizeNode::VT_OUT); + fbb_.Required(o, DequantizeNode::VT_MODE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateDequantizeNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *scales = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *biases = nullptr, + int32_t group_size = 0, + int32_t bits = 0, + ::flatbuffers::Offset<::flatbuffers::String> mode = 0, + const mlx_delegate::Tid *global_scale = nullptr, + ::flatbuffers::Optional dtype = ::flatbuffers::nullopt) { + DequantizeNodeBuilder builder_(_fbb); + builder_.add_global_scale(global_scale); + builder_.add_mode(mode); + builder_.add_bits(bits); + builder_.add_group_size(group_size); + builder_.add_biases(biases); + builder_.add_out(out); + builder_.add_scales(scales); + builder_.add_w(w); + if(dtype) { builder_.add_dtype(*dtype); } + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateDequantizeNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *scales = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *biases = nullptr, + int32_t group_size = 0, + int32_t bits = 0, + const char *mode = nullptr, + const mlx_delegate::Tid *global_scale = nullptr, + ::flatbuffers::Optional dtype = ::flatbuffers::nullopt) { + auto mode__ = mode ? _fbb.CreateString(mode) : 0; + return mlx_delegate::CreateDequantizeNode( + _fbb, + w, + scales, + out, + biases, + group_size, + bits, + mode__, + global_scale, + dtype); +} + +struct LessNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LessNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LessNodeBuilder { + typedef LessNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(LessNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(LessNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LessNode::VT_OUT, out); + } + explicit LessNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LessNode::VT_A); + fbb_.Required(o, LessNode::VT_B); + fbb_.Required(o, LessNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLessNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LessNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct LessEqualNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LessEqualNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LessEqualNodeBuilder { + typedef LessEqualNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(LessEqualNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(LessEqualNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LessEqualNode::VT_OUT, out); + } + explicit LessEqualNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LessEqualNode::VT_A); + fbb_.Required(o, LessEqualNode::VT_B); + fbb_.Required(o, LessEqualNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLessEqualNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LessEqualNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct GreaterNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef GreaterNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct GreaterNodeBuilder { + typedef GreaterNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(GreaterNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(GreaterNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(GreaterNode::VT_OUT, out); + } + explicit GreaterNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, GreaterNode::VT_A); + fbb_.Required(o, GreaterNode::VT_B); + fbb_.Required(o, GreaterNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateGreaterNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + GreaterNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct GreaterEqualNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef GreaterEqualNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct GreaterEqualNodeBuilder { + typedef GreaterEqualNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(GreaterEqualNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(GreaterEqualNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(GreaterEqualNode::VT_OUT, out); + } + explicit GreaterEqualNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, GreaterEqualNode::VT_A); + fbb_.Required(o, GreaterEqualNode::VT_B); + fbb_.Required(o, GreaterEqualNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateGreaterEqualNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + GreaterEqualNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct EqualNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef EqualNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct EqualNodeBuilder { + typedef EqualNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(EqualNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(EqualNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(EqualNode::VT_OUT, out); + } + explicit EqualNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, EqualNode::VT_A); + fbb_.Required(o, EqualNode::VT_B); + fbb_.Required(o, EqualNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateEqualNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + EqualNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct NotEqualNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef NotEqualNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct NotEqualNodeBuilder { + typedef NotEqualNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(NotEqualNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(NotEqualNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(NotEqualNode::VT_OUT, out); + } + explicit NotEqualNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, NotEqualNode::VT_A); + fbb_.Required(o, NotEqualNode::VT_B); + fbb_.Required(o, NotEqualNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateNotEqualNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + NotEqualNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct LogicalNotNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LogicalNotNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LogicalNotNodeBuilder { + typedef LogicalNotNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(LogicalNotNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LogicalNotNode::VT_OUT, out); + } + explicit LogicalNotNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LogicalNotNode::VT_X); + fbb_.Required(o, LogicalNotNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLogicalNotNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LogicalNotNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct BitwiseInvertNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef BitwiseInvertNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct BitwiseInvertNodeBuilder { + typedef BitwiseInvertNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(BitwiseInvertNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(BitwiseInvertNode::VT_OUT, out); + } + explicit BitwiseInvertNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, BitwiseInvertNode::VT_X); + fbb_.Required(o, BitwiseInvertNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateBitwiseInvertNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + BitwiseInvertNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct LogicalAndNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LogicalAndNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LogicalAndNodeBuilder { + typedef LogicalAndNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(LogicalAndNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(LogicalAndNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LogicalAndNode::VT_OUT, out); + } + explicit LogicalAndNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LogicalAndNode::VT_A); + fbb_.Required(o, LogicalAndNode::VT_B); + fbb_.Required(o, LogicalAndNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLogicalAndNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LogicalAndNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct LogicalOrNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LogicalOrNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LogicalOrNodeBuilder { + typedef LogicalOrNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(LogicalOrNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(LogicalOrNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LogicalOrNode::VT_OUT, out); + } + explicit LogicalOrNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LogicalOrNode::VT_A); + fbb_.Required(o, LogicalOrNode::VT_B); + fbb_.Required(o, LogicalOrNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLogicalOrNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LogicalOrNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct BitwiseAndNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef BitwiseAndNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct BitwiseAndNodeBuilder { + typedef BitwiseAndNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(BitwiseAndNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(BitwiseAndNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(BitwiseAndNode::VT_OUT, out); + } + explicit BitwiseAndNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, BitwiseAndNode::VT_A); + fbb_.Required(o, BitwiseAndNode::VT_B); + fbb_.Required(o, BitwiseAndNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateBitwiseAndNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + BitwiseAndNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct BitwiseOrNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef BitwiseOrNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct BitwiseOrNodeBuilder { + typedef BitwiseOrNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(BitwiseOrNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(BitwiseOrNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(BitwiseOrNode::VT_OUT, out); + } + explicit BitwiseOrNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, BitwiseOrNode::VT_A); + fbb_.Required(o, BitwiseOrNode::VT_B); + fbb_.Required(o, BitwiseOrNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateBitwiseOrNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + BitwiseOrNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct BitwiseXorNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef BitwiseXorNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct BitwiseXorNodeBuilder { + typedef BitwiseXorNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(BitwiseXorNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(BitwiseXorNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(BitwiseXorNode::VT_OUT, out); + } + explicit BitwiseXorNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, BitwiseXorNode::VT_A); + fbb_.Required(o, BitwiseXorNode::VT_B); + fbb_.Required(o, BitwiseXorNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateBitwiseXorNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + BitwiseXorNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct TriNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TriNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_OUT = 4, + VT_N = 6, + VT_M = 8, + VT_K = 10, + VT_SCALAR_TYPE = 12 + }; + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *n() const { + return GetPointer(VT_N); + } + const mlx_delegate::IntOrVid *m() const { + return GetPointer(VT_M); + } + int32_t k() const { + return GetField(VT_K, 0); + } + int8_t scalar_type() const { + return GetField(VT_SCALAR_TYPE, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_N) && + verifier.VerifyTable(n()) && + VerifyOffsetRequired(verifier, VT_M) && + verifier.VerifyTable(m()) && + VerifyField(verifier, VT_K, 4) && + VerifyField(verifier, VT_SCALAR_TYPE, 1) && + verifier.EndTable(); + } +}; + +struct TriNodeBuilder { + typedef TriNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TriNode::VT_OUT, out); + } + void add_n(::flatbuffers::Offset n) { + fbb_.AddOffset(TriNode::VT_N, n); + } + void add_m(::flatbuffers::Offset m) { + fbb_.AddOffset(TriNode::VT_M, m); + } + void add_k(int32_t k) { + fbb_.AddElement(TriNode::VT_K, k, 0); + } + void add_scalar_type(int8_t scalar_type) { + fbb_.AddElement(TriNode::VT_SCALAR_TYPE, scalar_type, 0); + } + explicit TriNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TriNode::VT_OUT); + fbb_.Required(o, TriNode::VT_N); + fbb_.Required(o, TriNode::VT_M); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTriNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset n = 0, + ::flatbuffers::Offset m = 0, + int32_t k = 0, + int8_t scalar_type = 0) { + TriNodeBuilder builder_(_fbb); + builder_.add_k(k); + builder_.add_m(m); + builder_.add_n(n); + builder_.add_out(out); + builder_.add_scalar_type(scalar_type); + return builder_.Finish(); +} + +struct TrilNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TrilNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_K = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t k() const { + return GetField(VT_K, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_K, 4) && + verifier.EndTable(); + } +}; + +struct TrilNodeBuilder { + typedef TrilNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TrilNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TrilNode::VT_OUT, out); + } + void add_k(int32_t k) { + fbb_.AddElement(TrilNode::VT_K, k, 0); + } + explicit TrilNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TrilNode::VT_X); + fbb_.Required(o, TrilNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTrilNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t k = 0) { + TrilNodeBuilder builder_(_fbb); + builder_.add_k(k); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct TriuNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TriuNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_K = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t k() const { + return GetField(VT_K, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_K, 4) && + verifier.EndTable(); + } +}; + +struct TriuNodeBuilder { + typedef TriuNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TriuNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TriuNode::VT_OUT, out); + } + void add_k(int32_t k) { + fbb_.AddElement(TriuNode::VT_K, k, 0); + } + explicit TriuNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TriuNode::VT_X); + fbb_.Required(o, TriuNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTriuNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t k = 0) { + TriuNodeBuilder builder_(_fbb); + builder_.add_k(k); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ClipNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ClipNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_A_MIN = 8, + VT_A_MAX = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::Tid *a_min() const { + return GetStruct(VT_A_MIN); + } + const mlx_delegate::Tid *a_max() const { + return GetStruct(VT_A_MAX); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_A_MIN, 4) && + VerifyField(verifier, VT_A_MAX, 4) && + verifier.EndTable(); + } +}; + +struct ClipNodeBuilder { + typedef ClipNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ClipNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ClipNode::VT_OUT, out); + } + void add_a_min(const mlx_delegate::Tid *a_min) { + fbb_.AddStruct(ClipNode::VT_A_MIN, a_min); + } + void add_a_max(const mlx_delegate::Tid *a_max) { + fbb_.AddStruct(ClipNode::VT_A_MAX, a_max); + } + explicit ClipNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ClipNode::VT_X); + fbb_.Required(o, ClipNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateClipNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *a_min = nullptr, + const mlx_delegate::Tid *a_max = nullptr) { + ClipNodeBuilder builder_(_fbb); + builder_.add_a_max(a_max); + builder_.add_a_min(a_min); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct CumsumNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef CumsumNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8, + VT_REVERSE = 10, + VT_INCLUSIVE = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool reverse() const { + return GetField(VT_REVERSE, 0) != 0; + } + bool inclusive() const { + return GetField(VT_INCLUSIVE, 1) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + VerifyField(verifier, VT_REVERSE, 1) && + VerifyField(verifier, VT_INCLUSIVE, 1) && + verifier.EndTable(); + } +}; + +struct CumsumNodeBuilder { + typedef CumsumNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(CumsumNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(CumsumNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(CumsumNode::VT_AXIS, axis, 0); + } + void add_reverse(bool reverse) { + fbb_.AddElement(CumsumNode::VT_REVERSE, static_cast(reverse), 0); + } + void add_inclusive(bool inclusive) { + fbb_.AddElement(CumsumNode::VT_INCLUSIVE, static_cast(inclusive), 1); + } + explicit CumsumNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, CumsumNode::VT_X); + fbb_.Required(o, CumsumNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateCumsumNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0, + bool reverse = false, + bool inclusive = true) { + CumsumNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_inclusive(inclusive); + builder_.add_reverse(reverse); + return builder_.Finish(); +} + +struct StackNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef StackNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_TENSORS = 4, + VT_OUT = 6, + VT_AXIS = 8 + }; + const ::flatbuffers::Vector *tensors() const { + return GetPointer *>(VT_TENSORS); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_TENSORS) && + verifier.VerifyVector(tensors()) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct StackNodeBuilder { + typedef StackNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_tensors(::flatbuffers::Offset<::flatbuffers::Vector> tensors) { + fbb_.AddOffset(StackNode::VT_TENSORS, tensors); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(StackNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(StackNode::VT_AXIS, axis, 0); + } + explicit StackNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, StackNode::VT_TENSORS); + fbb_.Required(o, StackNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateStackNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::Vector> tensors = 0, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + StackNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_tensors(tensors); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateStackNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const std::vector *tensors = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + auto tensors__ = tensors ? _fbb.CreateVectorOfStructs(*tensors) : 0; + return mlx_delegate::CreateStackNode( + _fbb, + tensors__, + out, + axis); +} + +struct SignNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SignNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SignNodeBuilder { + typedef SignNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SignNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SignNode::VT_OUT, out); + } + explicit SignNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SignNode::VT_X); + fbb_.Required(o, SignNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSignNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SignNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct AnyNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AnyNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct AnyNodeBuilder { + typedef AnyNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(AnyNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AnyNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(AnyNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(AnyNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit AnyNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AnyNode::VT_X); + fbb_.Required(o, AnyNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAnyNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + AnyNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateAnyNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateAnyNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct AllNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AllNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct AllNodeBuilder { + typedef AllNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(AllNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AllNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(AllNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(AllNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit AllNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AllNode::VT_X); + fbb_.Required(o, AllNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAllNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + AllNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateAllNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateAllNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct RepeatNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RepeatNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_REPEATS = 8, + VT_AXIS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *repeats() const { + return GetPointer(VT_REPEATS); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_REPEATS) && + verifier.VerifyTable(repeats()) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct RepeatNodeBuilder { + typedef RepeatNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(RepeatNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RepeatNode::VT_OUT, out); + } + void add_repeats(::flatbuffers::Offset repeats) { + fbb_.AddOffset(RepeatNode::VT_REPEATS, repeats); + } + void add_axis(int32_t axis) { + fbb_.AddElement(RepeatNode::VT_AXIS, axis, 0); + } + explicit RepeatNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RepeatNode::VT_X); + fbb_.Required(o, RepeatNode::VT_OUT); + fbb_.Required(o, RepeatNode::VT_REPEATS); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRepeatNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset repeats = 0, + int32_t axis = 0) { + RepeatNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_repeats(repeats); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SortNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SortNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct SortNodeBuilder { + typedef SortNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SortNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SortNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(SortNode::VT_AXIS, axis, 0); + } + explicit SortNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SortNode::VT_X); + fbb_.Required(o, SortNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSortNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + SortNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArgsortNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArgsortNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct ArgsortNodeBuilder { + typedef ArgsortNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArgsortNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArgsortNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ArgsortNode::VT_AXIS, axis, 0); + } + explicit ArgsortNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArgsortNode::VT_X); + fbb_.Required(o, ArgsortNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArgsortNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0) { + ArgsortNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct PartitionNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef PartitionNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_KTH = 8, + VT_AXIS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *kth() const { + return GetPointer(VT_KTH); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_KTH) && + verifier.VerifyTable(kth()) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct PartitionNodeBuilder { + typedef PartitionNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(PartitionNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(PartitionNode::VT_OUT, out); + } + void add_kth(::flatbuffers::Offset kth) { + fbb_.AddOffset(PartitionNode::VT_KTH, kth); + } + void add_axis(int32_t axis) { + fbb_.AddElement(PartitionNode::VT_AXIS, axis, 0); + } + explicit PartitionNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, PartitionNode::VT_X); + fbb_.Required(o, PartitionNode::VT_OUT); + fbb_.Required(o, PartitionNode::VT_KTH); + return o; + } +}; + +inline ::flatbuffers::Offset CreatePartitionNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset kth = 0, + int32_t axis = 0) { + PartitionNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_kth(kth); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArgPartitionNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArgPartitionNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_KTH = 8, + VT_AXIS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::IntOrVid *kth() const { + return GetPointer(VT_KTH); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_KTH) && + verifier.VerifyTable(kth()) && + VerifyField(verifier, VT_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct ArgPartitionNodeBuilder { + typedef ArgPartitionNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArgPartitionNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArgPartitionNode::VT_OUT, out); + } + void add_kth(::flatbuffers::Offset kth) { + fbb_.AddOffset(ArgPartitionNode::VT_KTH, kth); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ArgPartitionNode::VT_AXIS, axis, 0); + } + explicit ArgPartitionNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArgPartitionNode::VT_X); + fbb_.Required(o, ArgPartitionNode::VT_OUT); + fbb_.Required(o, ArgPartitionNode::VT_KTH); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArgPartitionNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset kth = 0, + int32_t axis = 0) { + ArgPartitionNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_kth(kth); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct RollNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RollNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_SHIFT = 8, + VT_AXES = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *shift() const { + return GetPointer> *>(VT_SHIFT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_SHIFT) && + verifier.VerifyVector(shift()) && + verifier.VerifyVectorOfTables(shift()) && + VerifyOffsetRequired(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + verifier.EndTable(); + } +}; + +struct RollNodeBuilder { + typedef RollNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(RollNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RollNode::VT_OUT, out); + } + void add_shift(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shift) { + fbb_.AddOffset(RollNode::VT_SHIFT, shift); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(RollNode::VT_AXES, axes); + } + explicit RollNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RollNode::VT_X); + fbb_.Required(o, RollNode::VT_OUT); + fbb_.Required(o, RollNode::VT_SHIFT); + fbb_.Required(o, RollNode::VT_AXES); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRollNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shift = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0) { + RollNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_shift(shift); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateRollNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *shift = nullptr, + const std::vector *axes = nullptr) { + auto shift__ = shift ? _fbb.CreateVector<::flatbuffers::Offset>(*shift) : 0; + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateRollNode( + _fbb, + x, + out, + shift__, + axes__); +} + +struct FloorNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef FloorNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct FloorNodeBuilder { + typedef FloorNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(FloorNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(FloorNode::VT_OUT, out); + } + explicit FloorNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, FloorNode::VT_X); + fbb_.Required(o, FloorNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateFloorNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + FloorNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct CeilNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef CeilNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct CeilNodeBuilder { + typedef CeilNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(CeilNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(CeilNode::VT_OUT, out); + } + explicit CeilNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, CeilNode::VT_X); + fbb_.Required(o, CeilNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateCeilNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + CeilNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SquareNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SquareNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SquareNodeBuilder { + typedef SquareNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SquareNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SquareNode::VT_OUT, out); + } + explicit SquareNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SquareNode::VT_X); + fbb_.Required(o, SquareNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSquareNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SquareNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ExpNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ExpNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ExpNodeBuilder { + typedef ExpNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ExpNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ExpNode::VT_OUT, out); + } + explicit ExpNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ExpNode::VT_X); + fbb_.Required(o, ExpNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateExpNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ExpNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SinNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SinNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SinNodeBuilder { + typedef SinNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SinNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SinNode::VT_OUT, out); + } + explicit SinNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SinNode::VT_X); + fbb_.Required(o, SinNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSinNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SinNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct CosNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef CosNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct CosNodeBuilder { + typedef CosNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(CosNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(CosNode::VT_OUT, out); + } + explicit CosNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, CosNode::VT_X); + fbb_.Required(o, CosNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateCosNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + CosNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct TanNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TanNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct TanNodeBuilder { + typedef TanNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(TanNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(TanNode::VT_OUT, out); + } + explicit TanNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TanNode::VT_X); + fbb_.Required(o, TanNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTanNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + TanNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArcsinNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArcsinNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ArcsinNodeBuilder { + typedef ArcsinNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArcsinNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArcsinNode::VT_OUT, out); + } + explicit ArcsinNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArcsinNode::VT_X); + fbb_.Required(o, ArcsinNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArcsinNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ArcsinNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArccosNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArccosNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ArccosNodeBuilder { + typedef ArccosNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArccosNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArccosNode::VT_OUT, out); + } + explicit ArccosNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArccosNode::VT_X); + fbb_.Required(o, ArccosNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArccosNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ArccosNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArctanNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArctanNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ArctanNodeBuilder { + typedef ArctanNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArctanNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArctanNode::VT_OUT, out); + } + explicit ArctanNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArctanNode::VT_X); + fbb_.Required(o, ArctanNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArctanNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ArctanNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SinhNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SinhNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SinhNodeBuilder { + typedef SinhNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SinhNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SinhNode::VT_OUT, out); + } + explicit SinhNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SinhNode::VT_X); + fbb_.Required(o, SinhNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSinhNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SinhNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct CoshNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef CoshNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct CoshNodeBuilder { + typedef CoshNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(CoshNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(CoshNode::VT_OUT, out); + } + explicit CoshNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, CoshNode::VT_X); + fbb_.Required(o, CoshNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateCoshNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + CoshNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArcsinhNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArcsinhNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ArcsinhNodeBuilder { + typedef ArcsinhNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArcsinhNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArcsinhNode::VT_OUT, out); + } + explicit ArcsinhNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArcsinhNode::VT_X); + fbb_.Required(o, ArcsinhNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArcsinhNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ArcsinhNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArccoshNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArccoshNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ArccoshNodeBuilder { + typedef ArccoshNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArccoshNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArccoshNode::VT_OUT, out); + } + explicit ArccoshNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArccoshNode::VT_X); + fbb_.Required(o, ArccoshNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArccoshNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ArccoshNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ArctanhNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArctanhNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ArctanhNodeBuilder { + typedef ArctanhNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArctanhNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArctanhNode::VT_OUT, out); + } + explicit ArctanhNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArctanhNode::VT_X); + fbb_.Required(o, ArctanhNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArctanhNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ArctanhNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Log2Node FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Log2NodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct Log2NodeBuilder { + typedef Log2Node Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Log2Node::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Log2Node::VT_OUT, out); + } + explicit Log2NodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Log2Node::VT_X); + fbb_.Required(o, Log2Node::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLog2Node( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + Log2NodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Log10Node FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Log10NodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct Log10NodeBuilder { + typedef Log10Node Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Log10Node::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Log10Node::VT_OUT, out); + } + explicit Log10NodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Log10Node::VT_X); + fbb_.Required(o, Log10Node::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLog10Node( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + Log10NodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Log1pNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Log1pNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct Log1pNodeBuilder { + typedef Log1pNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Log1pNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Log1pNode::VT_OUT, out); + } + explicit Log1pNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Log1pNode::VT_X); + fbb_.Required(o, Log1pNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLog1pNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + Log1pNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ErfNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ErfNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ErfNodeBuilder { + typedef ErfNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ErfNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ErfNode::VT_OUT, out); + } + explicit ErfNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ErfNode::VT_X); + fbb_.Required(o, ErfNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateErfNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ErfNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Expm1Node FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Expm1NodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct Expm1NodeBuilder { + typedef Expm1Node Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(Expm1Node::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Expm1Node::VT_OUT, out); + } + explicit Expm1NodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Expm1Node::VT_X); + fbb_.Required(o, Expm1Node::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateExpm1Node( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + Expm1NodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct RoundNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RoundNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_DECIMALS = 8 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t decimals() const { + return GetField(VT_DECIMALS, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_DECIMALS, 4) && + verifier.EndTable(); + } +}; + +struct RoundNodeBuilder { + typedef RoundNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(RoundNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RoundNode::VT_OUT, out); + } + void add_decimals(int32_t decimals) { + fbb_.AddElement(RoundNode::VT_DECIMALS, decimals, 0); + } + explicit RoundNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RoundNode::VT_X); + fbb_.Required(o, RoundNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRoundNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t decimals = 0) { + RoundNodeBuilder builder_(_fbb); + builder_.add_decimals(decimals); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct ReciprocalNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ReciprocalNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct ReciprocalNodeBuilder { + typedef ReciprocalNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ReciprocalNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ReciprocalNode::VT_OUT, out); + } + explicit ReciprocalNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ReciprocalNode::VT_X); + fbb_.Required(o, ReciprocalNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateReciprocalNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + ReciprocalNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct SqrtNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SqrtNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct SqrtNodeBuilder { + typedef SqrtNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SqrtNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SqrtNode::VT_OUT, out); + } + explicit SqrtNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SqrtNode::VT_X); + fbb_.Required(o, SqrtNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSqrtNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + SqrtNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct AbsNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef AbsNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct AbsNodeBuilder { + typedef AbsNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(AbsNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(AbsNode::VT_OUT, out); + } + explicit AbsNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, AbsNode::VT_X); + fbb_.Required(o, AbsNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAbsNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + AbsNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct NegNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef NegNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct NegNodeBuilder { + typedef NegNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(NegNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(NegNode::VT_OUT, out); + } + explicit NegNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, NegNode::VT_X); + fbb_.Required(o, NegNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateNegNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr) { + NegNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_x(x); + return builder_.Finish(); +} + +struct Atan2Node FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef Atan2NodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct Atan2NodeBuilder { + typedef Atan2Node Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(Atan2Node::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(Atan2Node::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(Atan2Node::VT_OUT, out); + } + explicit Atan2NodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Atan2Node::VT_A); + fbb_.Required(o, Atan2Node::VT_B); + fbb_.Required(o, Atan2Node::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateAtan2Node( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + Atan2NodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct LogAddExpNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LogAddExpNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct LogAddExpNodeBuilder { + typedef LogAddExpNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(LogAddExpNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(LogAddExpNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LogAddExpNode::VT_OUT, out); + } + explicit LogAddExpNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LogAddExpNode::VT_A); + fbb_.Required(o, LogAddExpNode::VT_B); + fbb_.Required(o, LogAddExpNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLogAddExpNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + LogAddExpNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct FloorDivideNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef FloorDivideNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct FloorDivideNodeBuilder { + typedef FloorDivideNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(FloorDivideNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(FloorDivideNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(FloorDivideNode::VT_OUT, out); + } + explicit FloorDivideNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, FloorDivideNode::VT_A); + fbb_.Required(o, FloorDivideNode::VT_B); + fbb_.Required(o, FloorDivideNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateFloorDivideNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + FloorDivideNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct RemainderNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RemainderNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct RemainderNodeBuilder { + typedef RemainderNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(RemainderNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(RemainderNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RemainderNode::VT_OUT, out); + } + explicit RemainderNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RemainderNode::VT_A); + fbb_.Required(o, RemainderNode::VT_B); + fbb_.Required(o, RemainderNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRemainderNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + RemainderNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct PowerNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef PowerNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + verifier.EndTable(); + } +}; + +struct PowerNodeBuilder { + typedef PowerNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(PowerNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(PowerNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(PowerNode::VT_OUT, out); + } + explicit PowerNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, PowerNode::VT_A); + fbb_.Required(o, PowerNode::VT_B); + fbb_.Required(o, PowerNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreatePowerNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr) { + PowerNodeBuilder builder_(_fbb); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + return builder_.Finish(); +} + +struct LogSumExpNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef LogSumExpNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct LogSumExpNodeBuilder { + typedef LogSumExpNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(LogSumExpNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(LogSumExpNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(LogSumExpNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(LogSumExpNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit LogSumExpNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, LogSumExpNode::VT_X); + fbb_.Required(o, LogSumExpNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateLogSumExpNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + LogSumExpNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateLogSumExpNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateLogSumExpNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct SumNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SumNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct SumNodeBuilder { + typedef SumNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(SumNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(SumNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(SumNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(SumNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit SumNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, SumNode::VT_X); + fbb_.Required(o, SumNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSumNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + SumNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateSumNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateSumNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct MeanNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MeanNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct MeanNodeBuilder { + typedef MeanNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(MeanNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MeanNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(MeanNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(MeanNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit MeanNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MeanNode::VT_X); + fbb_.Required(o, MeanNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMeanNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + MeanNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateMeanNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateMeanNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct VarNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef VarNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10, + VT_DDOF = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + int32_t ddof() const { + return GetField(VT_DDOF, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + VerifyField(verifier, VT_DDOF, 4) && + verifier.EndTable(); + } +}; + +struct VarNodeBuilder { + typedef VarNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(VarNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(VarNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(VarNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(VarNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + void add_ddof(int32_t ddof) { + fbb_.AddElement(VarNode::VT_DDOF, ddof, 0); + } + explicit VarNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, VarNode::VT_X); + fbb_.Required(o, VarNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateVarNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false, + int32_t ddof = 0) { + VarNodeBuilder builder_(_fbb); + builder_.add_ddof(ddof); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateVarNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false, + int32_t ddof = 0) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateVarNode( + _fbb, + x, + out, + axes__, + keepdims, + ddof); +} + +struct StdNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef StdNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10, + VT_DDOF = 12 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + int32_t ddof() const { + return GetField(VT_DDOF, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + VerifyField(verifier, VT_DDOF, 4) && + verifier.EndTable(); + } +}; + +struct StdNodeBuilder { + typedef StdNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(StdNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(StdNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(StdNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(StdNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + void add_ddof(int32_t ddof) { + fbb_.AddElement(StdNode::VT_DDOF, ddof, 0); + } + explicit StdNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, StdNode::VT_X); + fbb_.Required(o, StdNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateStdNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false, + int32_t ddof = 0) { + StdNodeBuilder builder_(_fbb); + builder_.add_ddof(ddof); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateStdNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false, + int32_t ddof = 0) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateStdNode( + _fbb, + x, + out, + axes__, + keepdims, + ddof); +} + +struct ProdNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ProdNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct ProdNodeBuilder { + typedef ProdNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ProdNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ProdNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(ProdNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(ProdNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit ProdNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ProdNode::VT_X); + fbb_.Required(o, ProdNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateProdNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + ProdNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateProdNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateProdNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct MaxNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MaxNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct MaxNodeBuilder { + typedef MaxNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(MaxNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MaxNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(MaxNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(MaxNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit MaxNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MaxNode::VT_X); + fbb_.Required(o, MaxNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMaxNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + MaxNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateMaxNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateMaxNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct MinNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MinNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct MinNodeBuilder { + typedef MinNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(MinNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MinNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(MinNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(MinNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit MinNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MinNode::VT_X); + fbb_.Required(o, MinNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMinNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + MinNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateMinNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateMinNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct ArgminNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ArgminNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXIS = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + int32_t axis() const { + return GetField(VT_AXIS, 0); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_AXIS, 4) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct ArgminNodeBuilder { + typedef ArgminNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(ArgminNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(ArgminNode::VT_OUT, out); + } + void add_axis(int32_t axis) { + fbb_.AddElement(ArgminNode::VT_AXIS, axis, 0); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(ArgminNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit ArgminNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ArgminNode::VT_X); + fbb_.Required(o, ArgminNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateArgminNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + int32_t axis = 0, + bool keepdims = false) { + ArgminNodeBuilder builder_(_fbb); + builder_.add_axis(axis); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +struct MedianNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MedianNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_OUT = 6, + VT_AXES = 8, + VT_KEEPDIMS = 10 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector *axes() const { + return GetPointer *>(VT_AXES); + } + bool keepdims() const { + return GetField(VT_KEEPDIMS, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffset(verifier, VT_AXES) && + verifier.VerifyVector(axes()) && + VerifyField(verifier, VT_KEEPDIMS, 1) && + verifier.EndTable(); + } +}; + +struct MedianNodeBuilder { + typedef MedianNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(MedianNode::VT_X, x); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(MedianNode::VT_OUT, out); + } + void add_axes(::flatbuffers::Offset<::flatbuffers::Vector> axes) { + fbb_.AddOffset(MedianNode::VT_AXES, axes); + } + void add_keepdims(bool keepdims) { + fbb_.AddElement(MedianNode::VT_KEEPDIMS, static_cast(keepdims), 0); + } + explicit MedianNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MedianNode::VT_X); + fbb_.Required(o, MedianNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMedianNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector> axes = 0, + bool keepdims = false) { + MedianNodeBuilder builder_(_fbb); + builder_.add_axes(axes); + builder_.add_out(out); + builder_.add_x(x); + builder_.add_keepdims(keepdims); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateMedianNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *out = nullptr, + const std::vector *axes = nullptr, + bool keepdims = false) { + auto axes__ = axes ? _fbb.CreateVector(*axes) : 0; + return mlx_delegate::CreateMedianNode( + _fbb, + x, + out, + axes__, + keepdims); +} + +struct GatherMmNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef GatherMmNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_A = 4, + VT_B = 6, + VT_OUT = 8, + VT_LHS_INDICES = 10, + VT_RHS_INDICES = 12, + VT_SORTED_INDICES = 14 + }; + const mlx_delegate::Tid *a() const { + return GetStruct(VT_A); + } + const mlx_delegate::Tid *b() const { + return GetStruct(VT_B); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const mlx_delegate::Tid *lhs_indices() const { + return GetStruct(VT_LHS_INDICES); + } + const mlx_delegate::Tid *rhs_indices() const { + return GetStruct(VT_RHS_INDICES); + } + bool sorted_indices() const { + return GetField(VT_SORTED_INDICES, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_A, 4) && + VerifyFieldRequired(verifier, VT_B, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyField(verifier, VT_LHS_INDICES, 4) && + VerifyField(verifier, VT_RHS_INDICES, 4) && + VerifyField(verifier, VT_SORTED_INDICES, 1) && + verifier.EndTable(); + } +}; + +struct GatherMmNodeBuilder { + typedef GatherMmNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_a(const mlx_delegate::Tid *a) { + fbb_.AddStruct(GatherMmNode::VT_A, a); + } + void add_b(const mlx_delegate::Tid *b) { + fbb_.AddStruct(GatherMmNode::VT_B, b); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(GatherMmNode::VT_OUT, out); + } + void add_lhs_indices(const mlx_delegate::Tid *lhs_indices) { + fbb_.AddStruct(GatherMmNode::VT_LHS_INDICES, lhs_indices); + } + void add_rhs_indices(const mlx_delegate::Tid *rhs_indices) { + fbb_.AddStruct(GatherMmNode::VT_RHS_INDICES, rhs_indices); + } + void add_sorted_indices(bool sorted_indices) { + fbb_.AddElement(GatherMmNode::VT_SORTED_INDICES, static_cast(sorted_indices), 0); + } + explicit GatherMmNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, GatherMmNode::VT_A); + fbb_.Required(o, GatherMmNode::VT_B); + fbb_.Required(o, GatherMmNode::VT_OUT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateGatherMmNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *a = nullptr, + const mlx_delegate::Tid *b = nullptr, + const mlx_delegate::Tid *out = nullptr, + const mlx_delegate::Tid *lhs_indices = nullptr, + const mlx_delegate::Tid *rhs_indices = nullptr, + bool sorted_indices = false) { + GatherMmNodeBuilder builder_(_fbb); + builder_.add_rhs_indices(rhs_indices); + builder_.add_lhs_indices(lhs_indices); + builder_.add_out(out); + builder_.add_b(b); + builder_.add_a(a); + builder_.add_sorted_indices(sorted_indices); + return builder_.Finish(); +} + +struct GatherQmmNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef GatherQmmNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_X = 4, + VT_W = 6, + VT_SCALES = 8, + VT_OUT = 10, + VT_MODE = 12, + VT_BIASES = 14, + VT_LHS_INDICES = 16, + VT_RHS_INDICES = 18, + VT_TRANSPOSE = 20, + VT_GROUP_SIZE = 22, + VT_BITS = 24, + VT_SORTED_INDICES = 26 + }; + const mlx_delegate::Tid *x() const { + return GetStruct(VT_X); + } + const mlx_delegate::Tid *w() const { + return GetStruct(VT_W); + } + const mlx_delegate::Tid *scales() const { + return GetStruct(VT_SCALES); + } + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::String *mode() const { + return GetPointer(VT_MODE); + } + const mlx_delegate::Tid *biases() const { + return GetStruct(VT_BIASES); + } + const mlx_delegate::Tid *lhs_indices() const { + return GetStruct(VT_LHS_INDICES); + } + const mlx_delegate::Tid *rhs_indices() const { + return GetStruct(VT_RHS_INDICES); + } + bool transpose() const { + return GetField(VT_TRANSPOSE, 1) != 0; + } + int32_t group_size() const { + return GetField(VT_GROUP_SIZE, 0); + } + int32_t bits() const { + return GetField(VT_BITS, 0); + } + bool sorted_indices() const { + return GetField(VT_SORTED_INDICES, 0) != 0; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_X, 4) && + VerifyFieldRequired(verifier, VT_W, 4) && + VerifyFieldRequired(verifier, VT_SCALES, 4) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_MODE) && + verifier.VerifyString(mode()) && + VerifyField(verifier, VT_BIASES, 4) && + VerifyField(verifier, VT_LHS_INDICES, 4) && + VerifyField(verifier, VT_RHS_INDICES, 4) && + VerifyField(verifier, VT_TRANSPOSE, 1) && + VerifyField(verifier, VT_GROUP_SIZE, 4) && + VerifyField(verifier, VT_BITS, 4) && + VerifyField(verifier, VT_SORTED_INDICES, 1) && + verifier.EndTable(); + } +}; + +struct GatherQmmNodeBuilder { + typedef GatherQmmNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_x(const mlx_delegate::Tid *x) { + fbb_.AddStruct(GatherQmmNode::VT_X, x); + } + void add_w(const mlx_delegate::Tid *w) { + fbb_.AddStruct(GatherQmmNode::VT_W, w); + } + void add_scales(const mlx_delegate::Tid *scales) { + fbb_.AddStruct(GatherQmmNode::VT_SCALES, scales); + } + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(GatherQmmNode::VT_OUT, out); + } + void add_mode(::flatbuffers::Offset<::flatbuffers::String> mode) { + fbb_.AddOffset(GatherQmmNode::VT_MODE, mode); + } + void add_biases(const mlx_delegate::Tid *biases) { + fbb_.AddStruct(GatherQmmNode::VT_BIASES, biases); + } + void add_lhs_indices(const mlx_delegate::Tid *lhs_indices) { + fbb_.AddStruct(GatherQmmNode::VT_LHS_INDICES, lhs_indices); + } + void add_rhs_indices(const mlx_delegate::Tid *rhs_indices) { + fbb_.AddStruct(GatherQmmNode::VT_RHS_INDICES, rhs_indices); + } + void add_transpose(bool transpose) { + fbb_.AddElement(GatherQmmNode::VT_TRANSPOSE, static_cast(transpose), 1); + } + void add_group_size(int32_t group_size) { + fbb_.AddElement(GatherQmmNode::VT_GROUP_SIZE, group_size, 0); + } + void add_bits(int32_t bits) { + fbb_.AddElement(GatherQmmNode::VT_BITS, bits, 0); + } + void add_sorted_indices(bool sorted_indices) { + fbb_.AddElement(GatherQmmNode::VT_SORTED_INDICES, static_cast(sorted_indices), 0); + } + explicit GatherQmmNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, GatherQmmNode::VT_X); + fbb_.Required(o, GatherQmmNode::VT_W); + fbb_.Required(o, GatherQmmNode::VT_SCALES); + fbb_.Required(o, GatherQmmNode::VT_OUT); + fbb_.Required(o, GatherQmmNode::VT_MODE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateGatherQmmNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *scales = nullptr, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::String> mode = 0, + const mlx_delegate::Tid *biases = nullptr, + const mlx_delegate::Tid *lhs_indices = nullptr, + const mlx_delegate::Tid *rhs_indices = nullptr, + bool transpose = true, + int32_t group_size = 0, + int32_t bits = 0, + bool sorted_indices = false) { + GatherQmmNodeBuilder builder_(_fbb); + builder_.add_bits(bits); + builder_.add_group_size(group_size); + builder_.add_rhs_indices(rhs_indices); + builder_.add_lhs_indices(lhs_indices); + builder_.add_biases(biases); + builder_.add_mode(mode); + builder_.add_out(out); + builder_.add_scales(scales); + builder_.add_w(w); + builder_.add_x(x); + builder_.add_sorted_indices(sorted_indices); + builder_.add_transpose(transpose); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateGatherQmmNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *x = nullptr, + const mlx_delegate::Tid *w = nullptr, + const mlx_delegate::Tid *scales = nullptr, + const mlx_delegate::Tid *out = nullptr, + const char *mode = nullptr, + const mlx_delegate::Tid *biases = nullptr, + const mlx_delegate::Tid *lhs_indices = nullptr, + const mlx_delegate::Tid *rhs_indices = nullptr, + bool transpose = true, + int32_t group_size = 0, + int32_t bits = 0, + bool sorted_indices = false) { + auto mode__ = mode ? _fbb.CreateString(mode) : 0; + return mlx_delegate::CreateGatherQmmNode( + _fbb, + x, + w, + scales, + out, + mode__, + biases, + lhs_indices, + rhs_indices, + transpose, + group_size, + bits, + sorted_indices); +} + +struct ScanNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ScanNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_ORIGINALS = 4, + VT_SLICED = 6, + VT_OUTPUTS = 8, + VT_CARRY = 10, + VT_BODY_CHAIN_IDX = 12, + VT_SCAN_AXIS = 14 + }; + const ::flatbuffers::Vector *originals() const { + return GetPointer *>(VT_ORIGINALS); + } + const ::flatbuffers::Vector *sliced() const { + return GetPointer *>(VT_SLICED); + } + const ::flatbuffers::Vector *outputs() const { + return GetPointer *>(VT_OUTPUTS); + } + const ::flatbuffers::Vector *carry() const { + return GetPointer *>(VT_CARRY); + } + int32_t body_chain_idx() const { + return GetField(VT_BODY_CHAIN_IDX, 0); + } + int32_t scan_axis() const { + return GetField(VT_SCAN_AXIS, 1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_ORIGINALS) && + verifier.VerifyVector(originals()) && + VerifyOffsetRequired(verifier, VT_SLICED) && + verifier.VerifyVector(sliced()) && + VerifyOffsetRequired(verifier, VT_OUTPUTS) && + verifier.VerifyVector(outputs()) && + VerifyOffsetRequired(verifier, VT_CARRY) && + verifier.VerifyVector(carry()) && + VerifyField(verifier, VT_BODY_CHAIN_IDX, 4) && + VerifyField(verifier, VT_SCAN_AXIS, 4) && + verifier.EndTable(); + } +}; + +struct ScanNodeBuilder { + typedef ScanNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_originals(::flatbuffers::Offset<::flatbuffers::Vector> originals) { + fbb_.AddOffset(ScanNode::VT_ORIGINALS, originals); + } + void add_sliced(::flatbuffers::Offset<::flatbuffers::Vector> sliced) { + fbb_.AddOffset(ScanNode::VT_SLICED, sliced); + } + void add_outputs(::flatbuffers::Offset<::flatbuffers::Vector> outputs) { + fbb_.AddOffset(ScanNode::VT_OUTPUTS, outputs); + } + void add_carry(::flatbuffers::Offset<::flatbuffers::Vector> carry) { + fbb_.AddOffset(ScanNode::VT_CARRY, carry); + } + void add_body_chain_idx(int32_t body_chain_idx) { + fbb_.AddElement(ScanNode::VT_BODY_CHAIN_IDX, body_chain_idx, 0); + } + void add_scan_axis(int32_t scan_axis) { + fbb_.AddElement(ScanNode::VT_SCAN_AXIS, scan_axis, 1); + } + explicit ScanNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, ScanNode::VT_ORIGINALS); + fbb_.Required(o, ScanNode::VT_SLICED); + fbb_.Required(o, ScanNode::VT_OUTPUTS); + fbb_.Required(o, ScanNode::VT_CARRY); + return o; + } +}; + +inline ::flatbuffers::Offset CreateScanNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::Vector> originals = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> sliced = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> outputs = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> carry = 0, + int32_t body_chain_idx = 0, + int32_t scan_axis = 1) { + ScanNodeBuilder builder_(_fbb); + builder_.add_scan_axis(scan_axis); + builder_.add_body_chain_idx(body_chain_idx); + builder_.add_carry(carry); + builder_.add_outputs(outputs); + builder_.add_sliced(sliced); + builder_.add_originals(originals); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateScanNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const std::vector *originals = nullptr, + const std::vector *sliced = nullptr, + const std::vector *outputs = nullptr, + const std::vector *carry = nullptr, + int32_t body_chain_idx = 0, + int32_t scan_axis = 1) { + auto originals__ = originals ? _fbb.CreateVectorOfStructs(*originals) : 0; + auto sliced__ = sliced ? _fbb.CreateVectorOfStructs(*sliced) : 0; + auto outputs__ = outputs ? _fbb.CreateVectorOfStructs(*outputs) : 0; + auto carry__ = carry ? _fbb.CreateVectorOfStructs(*carry) : 0; + return mlx_delegate::CreateScanNode( + _fbb, + originals__, + sliced__, + outputs__, + carry__, + body_chain_idx, + scan_axis); +} + +struct IfNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef IfNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_COND = 4, + VT_THEN_CHAIN_IDX = 6, + VT_ELSE_CHAIN_IDX = 8 + }; + const mlx_delegate::IntOrVid *cond() const { + return GetPointer(VT_COND); + } + uint32_t then_chain_idx() const { + return GetField(VT_THEN_CHAIN_IDX, 0); + } + uint32_t else_chain_idx() const { + return GetField(VT_ELSE_CHAIN_IDX, 0); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_COND) && + verifier.VerifyTable(cond()) && + VerifyField(verifier, VT_THEN_CHAIN_IDX, 4) && + VerifyField(verifier, VT_ELSE_CHAIN_IDX, 4) && + verifier.EndTable(); + } +}; + +struct IfNodeBuilder { + typedef IfNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_cond(::flatbuffers::Offset cond) { + fbb_.AddOffset(IfNode::VT_COND, cond); + } + void add_then_chain_idx(uint32_t then_chain_idx) { + fbb_.AddElement(IfNode::VT_THEN_CHAIN_IDX, then_chain_idx, 0); + } + void add_else_chain_idx(uint32_t else_chain_idx) { + fbb_.AddElement(IfNode::VT_ELSE_CHAIN_IDX, else_chain_idx, 0); + } + explicit IfNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, IfNode::VT_COND); + return o; + } +}; + +inline ::flatbuffers::Offset CreateIfNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset cond = 0, + uint32_t then_chain_idx = 0, + uint32_t else_chain_idx = 0) { + IfNodeBuilder builder_(_fbb); + builder_.add_else_chain_idx(else_chain_idx); + builder_.add_then_chain_idx(then_chain_idx); + builder_.add_cond(cond); + return builder_.Finish(); +} + +struct RandomBitsNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef RandomBitsNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_OUT = 4, + VT_SHAPE = 6, + VT_SEED = 8, + VT_WIDTH = 10 + }; + const mlx_delegate::Tid *out() const { + return GetStruct(VT_OUT); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *shape() const { + return GetPointer> *>(VT_SHAPE); + } + const mlx_delegate::Vid *seed() const { + return GetStruct(VT_SEED); + } + int32_t width() const { + return GetField(VT_WIDTH, 4); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyFieldRequired(verifier, VT_OUT, 4) && + VerifyOffsetRequired(verifier, VT_SHAPE) && + verifier.VerifyVector(shape()) && + verifier.VerifyVectorOfTables(shape()) && + VerifyField(verifier, VT_SEED, 4) && + VerifyField(verifier, VT_WIDTH, 4) && + verifier.EndTable(); + } +}; + +struct RandomBitsNodeBuilder { + typedef RandomBitsNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_out(const mlx_delegate::Tid *out) { + fbb_.AddStruct(RandomBitsNode::VT_OUT, out); + } + void add_shape(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape) { + fbb_.AddOffset(RandomBitsNode::VT_SHAPE, shape); + } + void add_seed(const mlx_delegate::Vid *seed) { + fbb_.AddStruct(RandomBitsNode::VT_SEED, seed); + } + void add_width(int32_t width) { + fbb_.AddElement(RandomBitsNode::VT_WIDTH, width, 4); + } + explicit RandomBitsNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, RandomBitsNode::VT_OUT); + fbb_.Required(o, RandomBitsNode::VT_SHAPE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateRandomBitsNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *out = nullptr, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape = 0, + const mlx_delegate::Vid *seed = nullptr, + int32_t width = 4) { + RandomBitsNodeBuilder builder_(_fbb); + builder_.add_width(width); + builder_.add_seed(seed); + builder_.add_shape(shape); + builder_.add_out(out); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateRandomBitsNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const mlx_delegate::Tid *out = nullptr, + const std::vector<::flatbuffers::Offset> *shape = nullptr, + const mlx_delegate::Vid *seed = nullptr, + int32_t width = 4) { + auto shape__ = shape ? _fbb.CreateVector<::flatbuffers::Offset>(*shape) : 0; + return mlx_delegate::CreateRandomBitsNode( + _fbb, + out, + shape__, + seed, + width); +} + +struct MetalKernelNode FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MetalKernelNodeBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_NAME = 4, + VT_SOURCE = 6, + VT_INPUTS = 8, + VT_OUTPUTS = 10, + VT_GRID = 12, + VT_THREADGROUP = 14, + VT_HEADER = 16, + VT_INPUT_NAMES = 18, + VT_OUTPUT_NAMES = 20, + VT_ENSURE_ROW_CONTIGUOUS = 22, + VT_ATOMIC_OUTPUTS = 24, + VT_OUTPUT_SHAPES_FLAT = 26, + VT_OUTPUT_SHAPE_LENGTHS = 28, + VT_OUTPUT_DTYPES = 30, + VT_TEMPLATE_ARG_NAMES = 32, + VT_TEMPLATE_ARG_KINDS = 34, + VT_TEMPLATE_ARG_VALUES = 36, + VT_INIT_VALUE = 38 + }; + const ::flatbuffers::String *name() const { + return GetPointer(VT_NAME); + } + const ::flatbuffers::String *source() const { + return GetPointer(VT_SOURCE); + } + const ::flatbuffers::Vector *inputs() const { + return GetPointer *>(VT_INPUTS); + } + const ::flatbuffers::Vector *outputs() const { + return GetPointer *>(VT_OUTPUTS); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *grid() const { + return GetPointer> *>(VT_GRID); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *threadgroup() const { + return GetPointer> *>(VT_THREADGROUP); + } + const ::flatbuffers::String *header() const { + return GetPointer(VT_HEADER); + } + const ::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>> *input_names() const { + return GetPointer> *>(VT_INPUT_NAMES); + } + const ::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>> *output_names() const { + return GetPointer> *>(VT_OUTPUT_NAMES); + } + bool ensure_row_contiguous() const { + return GetField(VT_ENSURE_ROW_CONTIGUOUS, 1) != 0; + } + bool atomic_outputs() const { + return GetField(VT_ATOMIC_OUTPUTS, 0) != 0; + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *output_shapes_flat() const { + return GetPointer> *>(VT_OUTPUT_SHAPES_FLAT); + } + const ::flatbuffers::Vector *output_shape_lengths() const { + return GetPointer *>(VT_OUTPUT_SHAPE_LENGTHS); + } + const ::flatbuffers::Vector *output_dtypes() const { + return GetPointer *>(VT_OUTPUT_DTYPES); + } + const ::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>> *template_arg_names() const { + return GetPointer> *>(VT_TEMPLATE_ARG_NAMES); + } + const ::flatbuffers::Vector *template_arg_kinds() const { + return GetPointer *>(VT_TEMPLATE_ARG_KINDS); + } + const ::flatbuffers::Vector *template_arg_values() const { + return GetPointer *>(VT_TEMPLATE_ARG_VALUES); + } + ::flatbuffers::Optional init_value() const { + return GetOptional(VT_INIT_VALUE); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_NAME) && + verifier.VerifyString(name()) && + VerifyOffsetRequired(verifier, VT_SOURCE) && + verifier.VerifyString(source()) && + VerifyOffsetRequired(verifier, VT_INPUTS) && + verifier.VerifyVector(inputs()) && + VerifyOffsetRequired(verifier, VT_OUTPUTS) && + verifier.VerifyVector(outputs()) && + VerifyOffsetRequired(verifier, VT_GRID) && + verifier.VerifyVector(grid()) && + verifier.VerifyVectorOfTables(grid()) && + VerifyOffsetRequired(verifier, VT_THREADGROUP) && + verifier.VerifyVector(threadgroup()) && + verifier.VerifyVectorOfTables(threadgroup()) && + VerifyOffset(verifier, VT_HEADER) && + verifier.VerifyString(header()) && + VerifyOffset(verifier, VT_INPUT_NAMES) && + verifier.VerifyVector(input_names()) && + verifier.VerifyVectorOfStrings(input_names()) && + VerifyOffset(verifier, VT_OUTPUT_NAMES) && + verifier.VerifyVector(output_names()) && + verifier.VerifyVectorOfStrings(output_names()) && + VerifyField(verifier, VT_ENSURE_ROW_CONTIGUOUS, 1) && + VerifyField(verifier, VT_ATOMIC_OUTPUTS, 1) && + VerifyOffset(verifier, VT_OUTPUT_SHAPES_FLAT) && + verifier.VerifyVector(output_shapes_flat()) && + verifier.VerifyVectorOfTables(output_shapes_flat()) && + VerifyOffset(verifier, VT_OUTPUT_SHAPE_LENGTHS) && + verifier.VerifyVector(output_shape_lengths()) && + VerifyOffset(verifier, VT_OUTPUT_DTYPES) && + verifier.VerifyVector(output_dtypes()) && + VerifyOffset(verifier, VT_TEMPLATE_ARG_NAMES) && + verifier.VerifyVector(template_arg_names()) && + verifier.VerifyVectorOfStrings(template_arg_names()) && + VerifyOffset(verifier, VT_TEMPLATE_ARG_KINDS) && + verifier.VerifyVector(template_arg_kinds()) && + VerifyOffset(verifier, VT_TEMPLATE_ARG_VALUES) && + verifier.VerifyVector(template_arg_values()) && + VerifyField(verifier, VT_INIT_VALUE, 4) && + verifier.EndTable(); + } +}; + +struct MetalKernelNodeBuilder { + typedef MetalKernelNode Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_name(::flatbuffers::Offset<::flatbuffers::String> name) { + fbb_.AddOffset(MetalKernelNode::VT_NAME, name); + } + void add_source(::flatbuffers::Offset<::flatbuffers::String> source) { + fbb_.AddOffset(MetalKernelNode::VT_SOURCE, source); + } + void add_inputs(::flatbuffers::Offset<::flatbuffers::Vector> inputs) { + fbb_.AddOffset(MetalKernelNode::VT_INPUTS, inputs); + } + void add_outputs(::flatbuffers::Offset<::flatbuffers::Vector> outputs) { + fbb_.AddOffset(MetalKernelNode::VT_OUTPUTS, outputs); + } + void add_grid(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> grid) { + fbb_.AddOffset(MetalKernelNode::VT_GRID, grid); + } + void add_threadgroup(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> threadgroup) { + fbb_.AddOffset(MetalKernelNode::VT_THREADGROUP, threadgroup); + } + void add_header(::flatbuffers::Offset<::flatbuffers::String> header) { + fbb_.AddOffset(MetalKernelNode::VT_HEADER, header); + } + void add_input_names(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>>> input_names) { + fbb_.AddOffset(MetalKernelNode::VT_INPUT_NAMES, input_names); + } + void add_output_names(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>>> output_names) { + fbb_.AddOffset(MetalKernelNode::VT_OUTPUT_NAMES, output_names); + } + void add_ensure_row_contiguous(bool ensure_row_contiguous) { + fbb_.AddElement(MetalKernelNode::VT_ENSURE_ROW_CONTIGUOUS, static_cast(ensure_row_contiguous), 1); + } + void add_atomic_outputs(bool atomic_outputs) { + fbb_.AddElement(MetalKernelNode::VT_ATOMIC_OUTPUTS, static_cast(atomic_outputs), 0); + } + void add_output_shapes_flat(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> output_shapes_flat) { + fbb_.AddOffset(MetalKernelNode::VT_OUTPUT_SHAPES_FLAT, output_shapes_flat); + } + void add_output_shape_lengths(::flatbuffers::Offset<::flatbuffers::Vector> output_shape_lengths) { + fbb_.AddOffset(MetalKernelNode::VT_OUTPUT_SHAPE_LENGTHS, output_shape_lengths); + } + void add_output_dtypes(::flatbuffers::Offset<::flatbuffers::Vector> output_dtypes) { + fbb_.AddOffset(MetalKernelNode::VT_OUTPUT_DTYPES, output_dtypes); + } + void add_template_arg_names(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>>> template_arg_names) { + fbb_.AddOffset(MetalKernelNode::VT_TEMPLATE_ARG_NAMES, template_arg_names); + } + void add_template_arg_kinds(::flatbuffers::Offset<::flatbuffers::Vector> template_arg_kinds) { + fbb_.AddOffset(MetalKernelNode::VT_TEMPLATE_ARG_KINDS, template_arg_kinds); + } + void add_template_arg_values(::flatbuffers::Offset<::flatbuffers::Vector> template_arg_values) { + fbb_.AddOffset(MetalKernelNode::VT_TEMPLATE_ARG_VALUES, template_arg_values); + } + void add_init_value(float init_value) { + fbb_.AddElement(MetalKernelNode::VT_INIT_VALUE, init_value); + } + explicit MetalKernelNodeBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MetalKernelNode::VT_NAME); + fbb_.Required(o, MetalKernelNode::VT_SOURCE); + fbb_.Required(o, MetalKernelNode::VT_INPUTS); + fbb_.Required(o, MetalKernelNode::VT_OUTPUTS); + fbb_.Required(o, MetalKernelNode::VT_GRID); + fbb_.Required(o, MetalKernelNode::VT_THREADGROUP); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMetalKernelNode( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::String> name = 0, + ::flatbuffers::Offset<::flatbuffers::String> source = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> inputs = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> outputs = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> grid = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> threadgroup = 0, + ::flatbuffers::Offset<::flatbuffers::String> header = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>>> input_names = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>>> output_names = 0, + bool ensure_row_contiguous = true, + bool atomic_outputs = false, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> output_shapes_flat = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> output_shape_lengths = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> output_dtypes = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset<::flatbuffers::String>>> template_arg_names = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> template_arg_kinds = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> template_arg_values = 0, + ::flatbuffers::Optional init_value = ::flatbuffers::nullopt) { + MetalKernelNodeBuilder builder_(_fbb); + if(init_value) { builder_.add_init_value(*init_value); } + builder_.add_template_arg_values(template_arg_values); + builder_.add_template_arg_kinds(template_arg_kinds); + builder_.add_template_arg_names(template_arg_names); + builder_.add_output_dtypes(output_dtypes); + builder_.add_output_shape_lengths(output_shape_lengths); + builder_.add_output_shapes_flat(output_shapes_flat); + builder_.add_output_names(output_names); + builder_.add_input_names(input_names); + builder_.add_header(header); + builder_.add_threadgroup(threadgroup); + builder_.add_grid(grid); + builder_.add_outputs(outputs); + builder_.add_inputs(inputs); + builder_.add_source(source); + builder_.add_name(name); + builder_.add_atomic_outputs(atomic_outputs); + builder_.add_ensure_row_contiguous(ensure_row_contiguous); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateMetalKernelNodeDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const char *name = nullptr, + const char *source = nullptr, + const std::vector *inputs = nullptr, + const std::vector *outputs = nullptr, + const std::vector<::flatbuffers::Offset> *grid = nullptr, + const std::vector<::flatbuffers::Offset> *threadgroup = nullptr, + const char *header = nullptr, + const std::vector<::flatbuffers::Offset<::flatbuffers::String>> *input_names = nullptr, + const std::vector<::flatbuffers::Offset<::flatbuffers::String>> *output_names = nullptr, + bool ensure_row_contiguous = true, + bool atomic_outputs = false, + const std::vector<::flatbuffers::Offset> *output_shapes_flat = nullptr, + const std::vector *output_shape_lengths = nullptr, + const std::vector *output_dtypes = nullptr, + const std::vector<::flatbuffers::Offset<::flatbuffers::String>> *template_arg_names = nullptr, + const std::vector *template_arg_kinds = nullptr, + const std::vector *template_arg_values = nullptr, + ::flatbuffers::Optional init_value = ::flatbuffers::nullopt) { + auto name__ = name ? _fbb.CreateString(name) : 0; + auto source__ = source ? _fbb.CreateString(source) : 0; + auto inputs__ = inputs ? _fbb.CreateVectorOfStructs(*inputs) : 0; + auto outputs__ = outputs ? _fbb.CreateVectorOfStructs(*outputs) : 0; + auto grid__ = grid ? _fbb.CreateVector<::flatbuffers::Offset>(*grid) : 0; + auto threadgroup__ = threadgroup ? _fbb.CreateVector<::flatbuffers::Offset>(*threadgroup) : 0; + auto header__ = header ? _fbb.CreateString(header) : 0; + auto input_names__ = input_names ? _fbb.CreateVector<::flatbuffers::Offset<::flatbuffers::String>>(*input_names) : 0; + auto output_names__ = output_names ? _fbb.CreateVector<::flatbuffers::Offset<::flatbuffers::String>>(*output_names) : 0; + auto output_shapes_flat__ = output_shapes_flat ? _fbb.CreateVector<::flatbuffers::Offset>(*output_shapes_flat) : 0; + auto output_shape_lengths__ = output_shape_lengths ? _fbb.CreateVector(*output_shape_lengths) : 0; + auto output_dtypes__ = output_dtypes ? _fbb.CreateVector(*output_dtypes) : 0; + auto template_arg_names__ = template_arg_names ? _fbb.CreateVector<::flatbuffers::Offset<::flatbuffers::String>>(*template_arg_names) : 0; + auto template_arg_kinds__ = template_arg_kinds ? _fbb.CreateVector(*template_arg_kinds) : 0; + auto template_arg_values__ = template_arg_values ? _fbb.CreateVector(*template_arg_values) : 0; + return mlx_delegate::CreateMetalKernelNode( + _fbb, + name__, + source__, + inputs__, + outputs__, + grid__, + threadgroup__, + header__, + input_names__, + output_names__, + ensure_row_contiguous, + atomic_outputs, + output_shapes_flat__, + output_shape_lengths__, + output_dtypes__, + template_arg_names__, + template_arg_kinds__, + template_arg_values__, + init_value); +} + +struct Instruction FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef InstructionBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_OP_TYPE = 4, + VT_OP = 6 + }; + mlx_delegate::OpNode op_type() const { + return static_cast(GetField(VT_OP_TYPE, 0)); + } + const void *op() const { + return GetPointer(VT_OP); + } + template const T *op_as() const; + const mlx_delegate::NoopNode *op_as_NoopNode() const { + return op_type() == mlx_delegate::OpNode_NoopNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::IdCopyNode *op_as_IdCopyNode() const { + return op_type() == mlx_delegate::OpNode_IdCopyNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AddmmNode *op_as_AddmmNode() const { + return op_type() == mlx_delegate::OpNode_AddmmNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ItemIntNode *op_as_ItemIntNode() const { + return op_type() == mlx_delegate::OpNode_ItemIntNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ExpandDimsNode *op_as_ExpandDimsNode() const { + return op_type() == mlx_delegate::OpNode_ExpandDimsNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TileNode *op_as_TileNode() const { + return op_type() == mlx_delegate::OpNode_TileNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TakeAlongAxisNode *op_as_TakeAlongAxisNode() const { + return op_type() == mlx_delegate::OpNode_TakeAlongAxisNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TakeNode *op_as_TakeNode() const { + return op_type() == mlx_delegate::OpNode_TakeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RMSNormNode *op_as_RMSNormNode() const { + return op_type() == mlx_delegate::OpNode_RMSNormNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LayerNormNode *op_as_LayerNormNode() const { + return op_type() == mlx_delegate::OpNode_LayerNormNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RopeNode *op_as_RopeNode() const { + return op_type() == mlx_delegate::OpNode_RopeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SdpaNode *op_as_SdpaNode() const { + return op_type() == mlx_delegate::OpNode_SdpaNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AddNode *op_as_AddNode() const { + return op_type() == mlx_delegate::OpNode_AddNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AddIntNode *op_as_AddIntNode() const { + return op_type() == mlx_delegate::OpNode_AddIntNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SubtractIntNode *op_as_SubtractIntNode() const { + return op_type() == mlx_delegate::OpNode_SubtractIntNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MultiplyIntNode *op_as_MultiplyIntNode() const { + return op_type() == mlx_delegate::OpNode_MultiplyIntNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::FloorDivideIntNode *op_as_FloorDivideIntNode() const { + return op_type() == mlx_delegate::OpNode_FloorDivideIntNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SymSizeNode *op_as_SymSizeNode() const { + return op_type() == mlx_delegate::OpNode_SymSizeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MultiplyNode *op_as_MultiplyNode() const { + return op_type() == mlx_delegate::OpNode_MultiplyNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::DivideNode *op_as_DivideNode() const { + return op_type() == mlx_delegate::OpNode_DivideNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SubtractNode *op_as_SubtractNode() const { + return op_type() == mlx_delegate::OpNode_SubtractNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::Conv1DNode *op_as_Conv1DNode() const { + return op_type() == mlx_delegate::OpNode_Conv1DNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::Conv2DNode *op_as_Conv2DNode() const { + return op_type() == mlx_delegate::OpNode_Conv2DNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::Conv3DNode *op_as_Conv3DNode() const { + return op_type() == mlx_delegate::OpNode_Conv3DNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::GeluNode *op_as_GeluNode() const { + return op_type() == mlx_delegate::OpNode_GeluNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ARangeNode *op_as_ARangeNode() const { + return op_type() == mlx_delegate::OpNode_ARangeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SiluNode *op_as_SiluNode() const { + return op_type() == mlx_delegate::OpNode_SiluNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SigmoidNode *op_as_SigmoidNode() const { + return op_type() == mlx_delegate::OpNode_SigmoidNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TanhNode *op_as_TanhNode() const { + return op_type() == mlx_delegate::OpNode_TanhNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SqueezeNode *op_as_SqueezeNode() const { + return op_type() == mlx_delegate::OpNode_SqueezeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SplitNode *op_as_SplitNode() const { + return op_type() == mlx_delegate::OpNode_SplitNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RsqrtNode *op_as_RsqrtNode() const { + return op_type() == mlx_delegate::OpNode_RsqrtNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MaximumNode *op_as_MaximumNode() const { + return op_type() == mlx_delegate::OpNode_MaximumNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MinimumNode *op_as_MinimumNode() const { + return op_type() == mlx_delegate::OpNode_MinimumNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LogNode *op_as_LogNode() const { + return op_type() == mlx_delegate::OpNode_LogNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SoftmaxNode *op_as_SoftmaxNode() const { + return op_type() == mlx_delegate::OpNode_SoftmaxNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::BroadcastToNode *op_as_BroadcastToNode() const { + return op_type() == mlx_delegate::OpNode_BroadcastToNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::PadNode *op_as_PadNode() const { + return op_type() == mlx_delegate::OpNode_PadNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::WhereNode *op_as_WhereNode() const { + return op_type() == mlx_delegate::OpNode_WhereNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ReshapeNode *op_as_ReshapeNode() const { + return op_type() == mlx_delegate::OpNode_ReshapeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TransposeNode *op_as_TransposeNode() const { + return op_type() == mlx_delegate::OpNode_TransposeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AsStridedNode *op_as_AsStridedNode() const { + return op_type() == mlx_delegate::OpNode_AsStridedNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ContiguousNode *op_as_ContiguousNode() const { + return op_type() == mlx_delegate::OpNode_ContiguousNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::GatherNode *op_as_GatherNode() const { + return op_type() == mlx_delegate::OpNode_GatherNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SliceNode *op_as_SliceNode() const { + return op_type() == mlx_delegate::OpNode_SliceNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AsTypeNode *op_as_AsTypeNode() const { + return op_type() == mlx_delegate::OpNode_AsTypeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ConcatenateNode *op_as_ConcatenateNode() const { + return op_type() == mlx_delegate::OpNode_ConcatenateNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::FullNode *op_as_FullNode() const { + return op_type() == mlx_delegate::OpNode_FullNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::FullLikeNode *op_as_FullLikeNode() const { + return op_type() == mlx_delegate::OpNode_FullLikeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArgmaxNode *op_as_ArgmaxNode() const { + return op_type() == mlx_delegate::OpNode_ArgmaxNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SliceUpdateNode *op_as_SliceUpdateNode() const { + return op_type() == mlx_delegate::OpNode_SliceUpdateNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::IndexCopyNode *op_as_IndexCopyNode() const { + return op_type() == mlx_delegate::OpNode_IndexCopyNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::DequantizeNode *op_as_DequantizeNode() const { + return op_type() == mlx_delegate::OpNode_DequantizeNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LessNode *op_as_LessNode() const { + return op_type() == mlx_delegate::OpNode_LessNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LessEqualNode *op_as_LessEqualNode() const { + return op_type() == mlx_delegate::OpNode_LessEqualNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::GreaterNode *op_as_GreaterNode() const { + return op_type() == mlx_delegate::OpNode_GreaterNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::GreaterEqualNode *op_as_GreaterEqualNode() const { + return op_type() == mlx_delegate::OpNode_GreaterEqualNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::EqualNode *op_as_EqualNode() const { + return op_type() == mlx_delegate::OpNode_EqualNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::NotEqualNode *op_as_NotEqualNode() const { + return op_type() == mlx_delegate::OpNode_NotEqualNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LogicalNotNode *op_as_LogicalNotNode() const { + return op_type() == mlx_delegate::OpNode_LogicalNotNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LogicalAndNode *op_as_LogicalAndNode() const { + return op_type() == mlx_delegate::OpNode_LogicalAndNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LogicalOrNode *op_as_LogicalOrNode() const { + return op_type() == mlx_delegate::OpNode_LogicalOrNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TriNode *op_as_TriNode() const { + return op_type() == mlx_delegate::OpNode_TriNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TrilNode *op_as_TrilNode() const { + return op_type() == mlx_delegate::OpNode_TrilNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TriuNode *op_as_TriuNode() const { + return op_type() == mlx_delegate::OpNode_TriuNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::FloorNode *op_as_FloorNode() const { + return op_type() == mlx_delegate::OpNode_FloorNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::CeilNode *op_as_CeilNode() const { + return op_type() == mlx_delegate::OpNode_CeilNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SquareNode *op_as_SquareNode() const { + return op_type() == mlx_delegate::OpNode_SquareNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ExpNode *op_as_ExpNode() const { + return op_type() == mlx_delegate::OpNode_ExpNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SinNode *op_as_SinNode() const { + return op_type() == mlx_delegate::OpNode_SinNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::CosNode *op_as_CosNode() const { + return op_type() == mlx_delegate::OpNode_CosNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::TanNode *op_as_TanNode() const { + return op_type() == mlx_delegate::OpNode_TanNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArcsinNode *op_as_ArcsinNode() const { + return op_type() == mlx_delegate::OpNode_ArcsinNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArccosNode *op_as_ArccosNode() const { + return op_type() == mlx_delegate::OpNode_ArccosNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArctanNode *op_as_ArctanNode() const { + return op_type() == mlx_delegate::OpNode_ArctanNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SinhNode *op_as_SinhNode() const { + return op_type() == mlx_delegate::OpNode_SinhNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::CoshNode *op_as_CoshNode() const { + return op_type() == mlx_delegate::OpNode_CoshNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArcsinhNode *op_as_ArcsinhNode() const { + return op_type() == mlx_delegate::OpNode_ArcsinhNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArccoshNode *op_as_ArccoshNode() const { + return op_type() == mlx_delegate::OpNode_ArccoshNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArctanhNode *op_as_ArctanhNode() const { + return op_type() == mlx_delegate::OpNode_ArctanhNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::Log2Node *op_as_Log2Node() const { + return op_type() == mlx_delegate::OpNode_Log2Node ? static_cast(op()) : nullptr; + } + const mlx_delegate::Log10Node *op_as_Log10Node() const { + return op_type() == mlx_delegate::OpNode_Log10Node ? static_cast(op()) : nullptr; + } + const mlx_delegate::Log1pNode *op_as_Log1pNode() const { + return op_type() == mlx_delegate::OpNode_Log1pNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ErfNode *op_as_ErfNode() const { + return op_type() == mlx_delegate::OpNode_ErfNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::Expm1Node *op_as_Expm1Node() const { + return op_type() == mlx_delegate::OpNode_Expm1Node ? static_cast(op()) : nullptr; + } + const mlx_delegate::RoundNode *op_as_RoundNode() const { + return op_type() == mlx_delegate::OpNode_RoundNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ReciprocalNode *op_as_ReciprocalNode() const { + return op_type() == mlx_delegate::OpNode_ReciprocalNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SqrtNode *op_as_SqrtNode() const { + return op_type() == mlx_delegate::OpNode_SqrtNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AbsNode *op_as_AbsNode() const { + return op_type() == mlx_delegate::OpNode_AbsNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::NegNode *op_as_NegNode() const { + return op_type() == mlx_delegate::OpNode_NegNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::Atan2Node *op_as_Atan2Node() const { + return op_type() == mlx_delegate::OpNode_Atan2Node ? static_cast(op()) : nullptr; + } + const mlx_delegate::LogAddExpNode *op_as_LogAddExpNode() const { + return op_type() == mlx_delegate::OpNode_LogAddExpNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::FloorDivideNode *op_as_FloorDivideNode() const { + return op_type() == mlx_delegate::OpNode_FloorDivideNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::PowerNode *op_as_PowerNode() const { + return op_type() == mlx_delegate::OpNode_PowerNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::LogSumExpNode *op_as_LogSumExpNode() const { + return op_type() == mlx_delegate::OpNode_LogSumExpNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SumNode *op_as_SumNode() const { + return op_type() == mlx_delegate::OpNode_SumNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MeanNode *op_as_MeanNode() const { + return op_type() == mlx_delegate::OpNode_MeanNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::VarNode *op_as_VarNode() const { + return op_type() == mlx_delegate::OpNode_VarNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::StdNode *op_as_StdNode() const { + return op_type() == mlx_delegate::OpNode_StdNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ProdNode *op_as_ProdNode() const { + return op_type() == mlx_delegate::OpNode_ProdNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MaxNode *op_as_MaxNode() const { + return op_type() == mlx_delegate::OpNode_MaxNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MinNode *op_as_MinNode() const { + return op_type() == mlx_delegate::OpNode_MinNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArgminNode *op_as_ArgminNode() const { + return op_type() == mlx_delegate::OpNode_ArgminNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MedianNode *op_as_MedianNode() const { + return op_type() == mlx_delegate::OpNode_MedianNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ModIntNode *op_as_ModIntNode() const { + return op_type() == mlx_delegate::OpNode_ModIntNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RemainderNode *op_as_RemainderNode() const { + return op_type() == mlx_delegate::OpNode_RemainderNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ConvTranspose1DNode *op_as_ConvTranspose1DNode() const { + return op_type() == mlx_delegate::OpNode_ConvTranspose1DNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ConvTranspose2DNode *op_as_ConvTranspose2DNode() const { + return op_type() == mlx_delegate::OpNode_ConvTranspose2DNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ConvTranspose3DNode *op_as_ConvTranspose3DNode() const { + return op_type() == mlx_delegate::OpNode_ConvTranspose3DNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ClipNode *op_as_ClipNode() const { + return op_type() == mlx_delegate::OpNode_ClipNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::CumsumNode *op_as_CumsumNode() const { + return op_type() == mlx_delegate::OpNode_CumsumNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::StackNode *op_as_StackNode() const { + return op_type() == mlx_delegate::OpNode_StackNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SignNode *op_as_SignNode() const { + return op_type() == mlx_delegate::OpNode_SignNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AnyNode *op_as_AnyNode() const { + return op_type() == mlx_delegate::OpNode_AnyNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::AllNode *op_as_AllNode() const { + return op_type() == mlx_delegate::OpNode_AllNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RepeatNode *op_as_RepeatNode() const { + return op_type() == mlx_delegate::OpNode_RepeatNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::SortNode *op_as_SortNode() const { + return op_type() == mlx_delegate::OpNode_SortNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArgsortNode *op_as_ArgsortNode() const { + return op_type() == mlx_delegate::OpNode_ArgsortNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::PartitionNode *op_as_PartitionNode() const { + return op_type() == mlx_delegate::OpNode_PartitionNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ArgPartitionNode *op_as_ArgPartitionNode() const { + return op_type() == mlx_delegate::OpNode_ArgPartitionNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::QuantizedMatmulNode *op_as_QuantizedMatmulNode() const { + return op_type() == mlx_delegate::OpNode_QuantizedMatmulNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ScatterAddNode *op_as_ScatterAddNode() const { + return op_type() == mlx_delegate::OpNode_ScatterAddNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::GatherMmNode *op_as_GatherMmNode() const { + return op_type() == mlx_delegate::OpNode_GatherMmNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::GatherQmmNode *op_as_GatherQmmNode() const { + return op_type() == mlx_delegate::OpNode_GatherQmmNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::ScanNode *op_as_ScanNode() const { + return op_type() == mlx_delegate::OpNode_ScanNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::MetalKernelNode *op_as_MetalKernelNode() const { + return op_type() == mlx_delegate::OpNode_MetalKernelNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::BitwiseInvertNode *op_as_BitwiseInvertNode() const { + return op_type() == mlx_delegate::OpNode_BitwiseInvertNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RollNode *op_as_RollNode() const { + return op_type() == mlx_delegate::OpNode_RollNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::BitwiseAndNode *op_as_BitwiseAndNode() const { + return op_type() == mlx_delegate::OpNode_BitwiseAndNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::BitwiseOrNode *op_as_BitwiseOrNode() const { + return op_type() == mlx_delegate::OpNode_BitwiseOrNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::BitwiseXorNode *op_as_BitwiseXorNode() const { + return op_type() == mlx_delegate::OpNode_BitwiseXorNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::IfNode *op_as_IfNode() const { + return op_type() == mlx_delegate::OpNode_IfNode ? static_cast(op()) : nullptr; + } + const mlx_delegate::RandomBitsNode *op_as_RandomBitsNode() const { + return op_type() == mlx_delegate::OpNode_RandomBitsNode ? static_cast(op()) : nullptr; + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_OP_TYPE, 1) && + VerifyOffsetRequired(verifier, VT_OP) && + VerifyOpNode(verifier, op(), op_type()) && + verifier.EndTable(); + } +}; + +template<> inline const mlx_delegate::NoopNode *Instruction::op_as() const { + return op_as_NoopNode(); +} + +template<> inline const mlx_delegate::IdCopyNode *Instruction::op_as() const { + return op_as_IdCopyNode(); +} + +template<> inline const mlx_delegate::AddmmNode *Instruction::op_as() const { + return op_as_AddmmNode(); +} + +template<> inline const mlx_delegate::ItemIntNode *Instruction::op_as() const { + return op_as_ItemIntNode(); +} + +template<> inline const mlx_delegate::ExpandDimsNode *Instruction::op_as() const { + return op_as_ExpandDimsNode(); +} + +template<> inline const mlx_delegate::TileNode *Instruction::op_as() const { + return op_as_TileNode(); +} + +template<> inline const mlx_delegate::TakeAlongAxisNode *Instruction::op_as() const { + return op_as_TakeAlongAxisNode(); +} + +template<> inline const mlx_delegate::TakeNode *Instruction::op_as() const { + return op_as_TakeNode(); +} + +template<> inline const mlx_delegate::RMSNormNode *Instruction::op_as() const { + return op_as_RMSNormNode(); +} + +template<> inline const mlx_delegate::LayerNormNode *Instruction::op_as() const { + return op_as_LayerNormNode(); +} + +template<> inline const mlx_delegate::RopeNode *Instruction::op_as() const { + return op_as_RopeNode(); +} + +template<> inline const mlx_delegate::SdpaNode *Instruction::op_as() const { + return op_as_SdpaNode(); +} + +template<> inline const mlx_delegate::AddNode *Instruction::op_as() const { + return op_as_AddNode(); +} + +template<> inline const mlx_delegate::AddIntNode *Instruction::op_as() const { + return op_as_AddIntNode(); +} + +template<> inline const mlx_delegate::SubtractIntNode *Instruction::op_as() const { + return op_as_SubtractIntNode(); +} + +template<> inline const mlx_delegate::MultiplyIntNode *Instruction::op_as() const { + return op_as_MultiplyIntNode(); +} + +template<> inline const mlx_delegate::FloorDivideIntNode *Instruction::op_as() const { + return op_as_FloorDivideIntNode(); +} + +template<> inline const mlx_delegate::SymSizeNode *Instruction::op_as() const { + return op_as_SymSizeNode(); +} + +template<> inline const mlx_delegate::MultiplyNode *Instruction::op_as() const { + return op_as_MultiplyNode(); +} + +template<> inline const mlx_delegate::DivideNode *Instruction::op_as() const { + return op_as_DivideNode(); +} + +template<> inline const mlx_delegate::SubtractNode *Instruction::op_as() const { + return op_as_SubtractNode(); +} + +template<> inline const mlx_delegate::Conv1DNode *Instruction::op_as() const { + return op_as_Conv1DNode(); +} + +template<> inline const mlx_delegate::Conv2DNode *Instruction::op_as() const { + return op_as_Conv2DNode(); +} + +template<> inline const mlx_delegate::Conv3DNode *Instruction::op_as() const { + return op_as_Conv3DNode(); +} + +template<> inline const mlx_delegate::GeluNode *Instruction::op_as() const { + return op_as_GeluNode(); +} + +template<> inline const mlx_delegate::ARangeNode *Instruction::op_as() const { + return op_as_ARangeNode(); +} + +template<> inline const mlx_delegate::SiluNode *Instruction::op_as() const { + return op_as_SiluNode(); +} + +template<> inline const mlx_delegate::SigmoidNode *Instruction::op_as() const { + return op_as_SigmoidNode(); +} + +template<> inline const mlx_delegate::TanhNode *Instruction::op_as() const { + return op_as_TanhNode(); +} + +template<> inline const mlx_delegate::SqueezeNode *Instruction::op_as() const { + return op_as_SqueezeNode(); +} + +template<> inline const mlx_delegate::SplitNode *Instruction::op_as() const { + return op_as_SplitNode(); +} + +template<> inline const mlx_delegate::RsqrtNode *Instruction::op_as() const { + return op_as_RsqrtNode(); +} + +template<> inline const mlx_delegate::MaximumNode *Instruction::op_as() const { + return op_as_MaximumNode(); +} + +template<> inline const mlx_delegate::MinimumNode *Instruction::op_as() const { + return op_as_MinimumNode(); +} + +template<> inline const mlx_delegate::LogNode *Instruction::op_as() const { + return op_as_LogNode(); +} + +template<> inline const mlx_delegate::SoftmaxNode *Instruction::op_as() const { + return op_as_SoftmaxNode(); +} + +template<> inline const mlx_delegate::BroadcastToNode *Instruction::op_as() const { + return op_as_BroadcastToNode(); +} + +template<> inline const mlx_delegate::PadNode *Instruction::op_as() const { + return op_as_PadNode(); +} + +template<> inline const mlx_delegate::WhereNode *Instruction::op_as() const { + return op_as_WhereNode(); +} + +template<> inline const mlx_delegate::ReshapeNode *Instruction::op_as() const { + return op_as_ReshapeNode(); +} + +template<> inline const mlx_delegate::TransposeNode *Instruction::op_as() const { + return op_as_TransposeNode(); +} + +template<> inline const mlx_delegate::AsStridedNode *Instruction::op_as() const { + return op_as_AsStridedNode(); +} + +template<> inline const mlx_delegate::ContiguousNode *Instruction::op_as() const { + return op_as_ContiguousNode(); +} + +template<> inline const mlx_delegate::GatherNode *Instruction::op_as() const { + return op_as_GatherNode(); +} + +template<> inline const mlx_delegate::SliceNode *Instruction::op_as() const { + return op_as_SliceNode(); +} + +template<> inline const mlx_delegate::AsTypeNode *Instruction::op_as() const { + return op_as_AsTypeNode(); +} + +template<> inline const mlx_delegate::ConcatenateNode *Instruction::op_as() const { + return op_as_ConcatenateNode(); +} + +template<> inline const mlx_delegate::FullNode *Instruction::op_as() const { + return op_as_FullNode(); +} + +template<> inline const mlx_delegate::FullLikeNode *Instruction::op_as() const { + return op_as_FullLikeNode(); +} + +template<> inline const mlx_delegate::ArgmaxNode *Instruction::op_as() const { + return op_as_ArgmaxNode(); +} + +template<> inline const mlx_delegate::SliceUpdateNode *Instruction::op_as() const { + return op_as_SliceUpdateNode(); +} + +template<> inline const mlx_delegate::IndexCopyNode *Instruction::op_as() const { + return op_as_IndexCopyNode(); +} + +template<> inline const mlx_delegate::DequantizeNode *Instruction::op_as() const { + return op_as_DequantizeNode(); +} + +template<> inline const mlx_delegate::LessNode *Instruction::op_as() const { + return op_as_LessNode(); +} + +template<> inline const mlx_delegate::LessEqualNode *Instruction::op_as() const { + return op_as_LessEqualNode(); +} + +template<> inline const mlx_delegate::GreaterNode *Instruction::op_as() const { + return op_as_GreaterNode(); +} + +template<> inline const mlx_delegate::GreaterEqualNode *Instruction::op_as() const { + return op_as_GreaterEqualNode(); +} + +template<> inline const mlx_delegate::EqualNode *Instruction::op_as() const { + return op_as_EqualNode(); +} + +template<> inline const mlx_delegate::NotEqualNode *Instruction::op_as() const { + return op_as_NotEqualNode(); +} + +template<> inline const mlx_delegate::LogicalNotNode *Instruction::op_as() const { + return op_as_LogicalNotNode(); +} + +template<> inline const mlx_delegate::LogicalAndNode *Instruction::op_as() const { + return op_as_LogicalAndNode(); +} + +template<> inline const mlx_delegate::LogicalOrNode *Instruction::op_as() const { + return op_as_LogicalOrNode(); +} + +template<> inline const mlx_delegate::TriNode *Instruction::op_as() const { + return op_as_TriNode(); +} + +template<> inline const mlx_delegate::TrilNode *Instruction::op_as() const { + return op_as_TrilNode(); +} + +template<> inline const mlx_delegate::TriuNode *Instruction::op_as() const { + return op_as_TriuNode(); +} + +template<> inline const mlx_delegate::FloorNode *Instruction::op_as() const { + return op_as_FloorNode(); +} + +template<> inline const mlx_delegate::CeilNode *Instruction::op_as() const { + return op_as_CeilNode(); +} + +template<> inline const mlx_delegate::SquareNode *Instruction::op_as() const { + return op_as_SquareNode(); +} + +template<> inline const mlx_delegate::ExpNode *Instruction::op_as() const { + return op_as_ExpNode(); +} + +template<> inline const mlx_delegate::SinNode *Instruction::op_as() const { + return op_as_SinNode(); +} + +template<> inline const mlx_delegate::CosNode *Instruction::op_as() const { + return op_as_CosNode(); +} + +template<> inline const mlx_delegate::TanNode *Instruction::op_as() const { + return op_as_TanNode(); +} + +template<> inline const mlx_delegate::ArcsinNode *Instruction::op_as() const { + return op_as_ArcsinNode(); +} + +template<> inline const mlx_delegate::ArccosNode *Instruction::op_as() const { + return op_as_ArccosNode(); +} + +template<> inline const mlx_delegate::ArctanNode *Instruction::op_as() const { + return op_as_ArctanNode(); +} + +template<> inline const mlx_delegate::SinhNode *Instruction::op_as() const { + return op_as_SinhNode(); +} + +template<> inline const mlx_delegate::CoshNode *Instruction::op_as() const { + return op_as_CoshNode(); +} + +template<> inline const mlx_delegate::ArcsinhNode *Instruction::op_as() const { + return op_as_ArcsinhNode(); +} + +template<> inline const mlx_delegate::ArccoshNode *Instruction::op_as() const { + return op_as_ArccoshNode(); +} + +template<> inline const mlx_delegate::ArctanhNode *Instruction::op_as() const { + return op_as_ArctanhNode(); +} + +template<> inline const mlx_delegate::Log2Node *Instruction::op_as() const { + return op_as_Log2Node(); +} + +template<> inline const mlx_delegate::Log10Node *Instruction::op_as() const { + return op_as_Log10Node(); +} + +template<> inline const mlx_delegate::Log1pNode *Instruction::op_as() const { + return op_as_Log1pNode(); +} + +template<> inline const mlx_delegate::ErfNode *Instruction::op_as() const { + return op_as_ErfNode(); +} + +template<> inline const mlx_delegate::Expm1Node *Instruction::op_as() const { + return op_as_Expm1Node(); +} + +template<> inline const mlx_delegate::RoundNode *Instruction::op_as() const { + return op_as_RoundNode(); +} + +template<> inline const mlx_delegate::ReciprocalNode *Instruction::op_as() const { + return op_as_ReciprocalNode(); +} + +template<> inline const mlx_delegate::SqrtNode *Instruction::op_as() const { + return op_as_SqrtNode(); +} + +template<> inline const mlx_delegate::AbsNode *Instruction::op_as() const { + return op_as_AbsNode(); +} + +template<> inline const mlx_delegate::NegNode *Instruction::op_as() const { + return op_as_NegNode(); +} + +template<> inline const mlx_delegate::Atan2Node *Instruction::op_as() const { + return op_as_Atan2Node(); +} + +template<> inline const mlx_delegate::LogAddExpNode *Instruction::op_as() const { + return op_as_LogAddExpNode(); +} + +template<> inline const mlx_delegate::FloorDivideNode *Instruction::op_as() const { + return op_as_FloorDivideNode(); +} + +template<> inline const mlx_delegate::PowerNode *Instruction::op_as() const { + return op_as_PowerNode(); +} + +template<> inline const mlx_delegate::LogSumExpNode *Instruction::op_as() const { + return op_as_LogSumExpNode(); +} + +template<> inline const mlx_delegate::SumNode *Instruction::op_as() const { + return op_as_SumNode(); +} + +template<> inline const mlx_delegate::MeanNode *Instruction::op_as() const { + return op_as_MeanNode(); +} + +template<> inline const mlx_delegate::VarNode *Instruction::op_as() const { + return op_as_VarNode(); +} + +template<> inline const mlx_delegate::StdNode *Instruction::op_as() const { + return op_as_StdNode(); +} + +template<> inline const mlx_delegate::ProdNode *Instruction::op_as() const { + return op_as_ProdNode(); +} + +template<> inline const mlx_delegate::MaxNode *Instruction::op_as() const { + return op_as_MaxNode(); +} + +template<> inline const mlx_delegate::MinNode *Instruction::op_as() const { + return op_as_MinNode(); +} + +template<> inline const mlx_delegate::ArgminNode *Instruction::op_as() const { + return op_as_ArgminNode(); +} + +template<> inline const mlx_delegate::MedianNode *Instruction::op_as() const { + return op_as_MedianNode(); +} + +template<> inline const mlx_delegate::ModIntNode *Instruction::op_as() const { + return op_as_ModIntNode(); +} + +template<> inline const mlx_delegate::RemainderNode *Instruction::op_as() const { + return op_as_RemainderNode(); +} + +template<> inline const mlx_delegate::ConvTranspose1DNode *Instruction::op_as() const { + return op_as_ConvTranspose1DNode(); +} + +template<> inline const mlx_delegate::ConvTranspose2DNode *Instruction::op_as() const { + return op_as_ConvTranspose2DNode(); +} + +template<> inline const mlx_delegate::ConvTranspose3DNode *Instruction::op_as() const { + return op_as_ConvTranspose3DNode(); +} + +template<> inline const mlx_delegate::ClipNode *Instruction::op_as() const { + return op_as_ClipNode(); +} + +template<> inline const mlx_delegate::CumsumNode *Instruction::op_as() const { + return op_as_CumsumNode(); +} + +template<> inline const mlx_delegate::StackNode *Instruction::op_as() const { + return op_as_StackNode(); +} + +template<> inline const mlx_delegate::SignNode *Instruction::op_as() const { + return op_as_SignNode(); +} + +template<> inline const mlx_delegate::AnyNode *Instruction::op_as() const { + return op_as_AnyNode(); +} + +template<> inline const mlx_delegate::AllNode *Instruction::op_as() const { + return op_as_AllNode(); +} + +template<> inline const mlx_delegate::RepeatNode *Instruction::op_as() const { + return op_as_RepeatNode(); +} + +template<> inline const mlx_delegate::SortNode *Instruction::op_as() const { + return op_as_SortNode(); +} + +template<> inline const mlx_delegate::ArgsortNode *Instruction::op_as() const { + return op_as_ArgsortNode(); +} + +template<> inline const mlx_delegate::PartitionNode *Instruction::op_as() const { + return op_as_PartitionNode(); +} + +template<> inline const mlx_delegate::ArgPartitionNode *Instruction::op_as() const { + return op_as_ArgPartitionNode(); +} + +template<> inline const mlx_delegate::QuantizedMatmulNode *Instruction::op_as() const { + return op_as_QuantizedMatmulNode(); +} + +template<> inline const mlx_delegate::ScatterAddNode *Instruction::op_as() const { + return op_as_ScatterAddNode(); +} + +template<> inline const mlx_delegate::GatherMmNode *Instruction::op_as() const { + return op_as_GatherMmNode(); +} + +template<> inline const mlx_delegate::GatherQmmNode *Instruction::op_as() const { + return op_as_GatherQmmNode(); +} + +template<> inline const mlx_delegate::ScanNode *Instruction::op_as() const { + return op_as_ScanNode(); +} + +template<> inline const mlx_delegate::MetalKernelNode *Instruction::op_as() const { + return op_as_MetalKernelNode(); +} + +template<> inline const mlx_delegate::BitwiseInvertNode *Instruction::op_as() const { + return op_as_BitwiseInvertNode(); +} + +template<> inline const mlx_delegate::RollNode *Instruction::op_as() const { + return op_as_RollNode(); +} + +template<> inline const mlx_delegate::BitwiseAndNode *Instruction::op_as() const { + return op_as_BitwiseAndNode(); +} + +template<> inline const mlx_delegate::BitwiseOrNode *Instruction::op_as() const { + return op_as_BitwiseOrNode(); +} + +template<> inline const mlx_delegate::BitwiseXorNode *Instruction::op_as() const { + return op_as_BitwiseXorNode(); +} + +template<> inline const mlx_delegate::IfNode *Instruction::op_as() const { + return op_as_IfNode(); +} + +template<> inline const mlx_delegate::RandomBitsNode *Instruction::op_as() const { + return op_as_RandomBitsNode(); +} + +struct InstructionBuilder { + typedef Instruction Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_op_type(mlx_delegate::OpNode op_type) { + fbb_.AddElement(Instruction::VT_OP_TYPE, static_cast(op_type), 0); + } + void add_op(::flatbuffers::Offset op) { + fbb_.AddOffset(Instruction::VT_OP, op); + } + explicit InstructionBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, Instruction::VT_OP); + return o; + } +}; + +inline ::flatbuffers::Offset CreateInstruction( + ::flatbuffers::FlatBufferBuilder &_fbb, + mlx_delegate::OpNode op_type = mlx_delegate::OpNode_NONE, + ::flatbuffers::Offset op = 0) { + InstructionBuilder builder_(_fbb); + builder_.add_op(op); + builder_.add_op_type(op_type); + return builder_.Finish(); +} + +struct InstructionChain FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef InstructionChainBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_INSTRUCTIONS = 4 + }; + const ::flatbuffers::Vector<::flatbuffers::Offset> *instructions() const { + return GetPointer> *>(VT_INSTRUCTIONS); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_INSTRUCTIONS) && + verifier.VerifyVector(instructions()) && + verifier.VerifyVectorOfTables(instructions()) && + verifier.EndTable(); + } +}; + +struct InstructionChainBuilder { + typedef InstructionChain Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_instructions(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> instructions) { + fbb_.AddOffset(InstructionChain::VT_INSTRUCTIONS, instructions); + } + explicit InstructionChainBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, InstructionChain::VT_INSTRUCTIONS); + return o; + } +}; + +inline ::flatbuffers::Offset CreateInstructionChain( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> instructions = 0) { + InstructionChainBuilder builder_(_fbb); + builder_.add_instructions(instructions); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateInstructionChainDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const std::vector<::flatbuffers::Offset> *instructions = nullptr) { + auto instructions__ = instructions ? _fbb.CreateVector<::flatbuffers::Offset>(*instructions) : 0; + return mlx_delegate::CreateInstructionChain( + _fbb, + instructions__); +} + +struct ShapeDim FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef ShapeDimBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_VALUE = 4, + VT_MIN_VALUE = 6, + VT_MAX_VALUE = 8 + }; + int32_t value() const { + return GetField(VT_VALUE, -1); + } + int32_t min_value() const { + return GetField(VT_MIN_VALUE, 0); + } + int32_t max_value() const { + return GetField(VT_MAX_VALUE, -1); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_VALUE, 4) && + VerifyField(verifier, VT_MIN_VALUE, 4) && + VerifyField(verifier, VT_MAX_VALUE, 4) && + verifier.EndTable(); + } +}; + +struct ShapeDimBuilder { + typedef ShapeDim Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_value(int32_t value) { + fbb_.AddElement(ShapeDim::VT_VALUE, value, -1); + } + void add_min_value(int32_t min_value) { + fbb_.AddElement(ShapeDim::VT_MIN_VALUE, min_value, 0); + } + void add_max_value(int32_t max_value) { + fbb_.AddElement(ShapeDim::VT_MAX_VALUE, max_value, -1); + } + explicit ShapeDimBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateShapeDim( + ::flatbuffers::FlatBufferBuilder &_fbb, + int32_t value = -1, + int32_t min_value = 0, + int32_t max_value = -1) { + ShapeDimBuilder builder_(_fbb); + builder_.add_max_value(max_value); + builder_.add_min_value(min_value); + builder_.add_value(value); + return builder_.Finish(); +} + +struct TensorMeta FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef TensorMetaBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_SHAPE = 4, + VT_SCALAR_TYPE = 6, + VT_DIM_ORDER = 8 + }; + const ::flatbuffers::Vector<::flatbuffers::Offset> *shape() const { + return GetPointer> *>(VT_SHAPE); + } + int8_t scalar_type() const { + return GetField(VT_SCALAR_TYPE, 0); + } + const ::flatbuffers::Vector *dim_order() const { + return GetPointer *>(VT_DIM_ORDER); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_SHAPE) && + verifier.VerifyVector(shape()) && + verifier.VerifyVectorOfTables(shape()) && + VerifyField(verifier, VT_SCALAR_TYPE, 1) && + VerifyOffset(verifier, VT_DIM_ORDER) && + verifier.VerifyVector(dim_order()) && + verifier.EndTable(); + } +}; + +struct TensorMetaBuilder { + typedef TensorMeta Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_shape(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape) { + fbb_.AddOffset(TensorMeta::VT_SHAPE, shape); + } + void add_scalar_type(int8_t scalar_type) { + fbb_.AddElement(TensorMeta::VT_SCALAR_TYPE, scalar_type, 0); + } + void add_dim_order(::flatbuffers::Offset<::flatbuffers::Vector> dim_order) { + fbb_.AddOffset(TensorMeta::VT_DIM_ORDER, dim_order); + } + explicit TensorMetaBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, TensorMeta::VT_SHAPE); + return o; + } +}; + +inline ::flatbuffers::Offset CreateTensorMeta( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> shape = 0, + int8_t scalar_type = 0, + ::flatbuffers::Offset<::flatbuffers::Vector> dim_order = 0) { + TensorMetaBuilder builder_(_fbb); + builder_.add_dim_order(dim_order); + builder_.add_shape(shape); + builder_.add_scalar_type(scalar_type); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateTensorMetaDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const std::vector<::flatbuffers::Offset> *shape = nullptr, + int8_t scalar_type = 0, + const std::vector *dim_order = nullptr) { + auto shape__ = shape ? _fbb.CreateVector<::flatbuffers::Offset>(*shape) : 0; + auto dim_order__ = dim_order ? _fbb.CreateVector(*dim_order) : 0; + return mlx_delegate::CreateTensorMeta( + _fbb, + shape__, + scalar_type, + dim_order__); +} + +struct SlotVariant FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef SlotVariantBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_IDX = 4, + VT_SLOT_TYPE = 6 + }; + uint32_t idx() const { + return GetField(VT_IDX, 0); + } + mlx_delegate::SlotType slot_type() const { + return static_cast(GetField(VT_SLOT_TYPE, 0)); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyField(verifier, VT_IDX, 4) && + VerifyField(verifier, VT_SLOT_TYPE, 1) && + verifier.EndTable(); + } +}; + +struct SlotVariantBuilder { + typedef SlotVariant Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_idx(uint32_t idx) { + fbb_.AddElement(SlotVariant::VT_IDX, idx, 0); + } + void add_slot_type(mlx_delegate::SlotType slot_type) { + fbb_.AddElement(SlotVariant::VT_SLOT_TYPE, static_cast(slot_type), 0); + } + explicit SlotVariantBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + return o; + } +}; + +inline ::flatbuffers::Offset CreateSlotVariant( + ::flatbuffers::FlatBufferBuilder &_fbb, + uint32_t idx = 0, + mlx_delegate::SlotType slot_type = mlx_delegate::SlotType_TensorSlot) { + SlotVariantBuilder builder_(_fbb); + builder_.add_idx(idx); + builder_.add_slot_type(slot_type); + return builder_.Finish(); +} + +struct NamedSlot FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef NamedSlotBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_NAME = 4, + VT_SLOT = 6 + }; + const ::flatbuffers::String *name() const { + return GetPointer(VT_NAME); + } + const mlx_delegate::SlotVariant *slot() const { + return GetPointer(VT_SLOT); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffsetRequired(verifier, VT_NAME) && + verifier.VerifyString(name()) && + VerifyOffsetRequired(verifier, VT_SLOT) && + verifier.VerifyTable(slot()) && + verifier.EndTable(); + } +}; + +struct NamedSlotBuilder { + typedef NamedSlot Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_name(::flatbuffers::Offset<::flatbuffers::String> name) { + fbb_.AddOffset(NamedSlot::VT_NAME, name); + } + void add_slot(::flatbuffers::Offset slot) { + fbb_.AddOffset(NamedSlot::VT_SLOT, slot); + } + explicit NamedSlotBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, NamedSlot::VT_NAME); + fbb_.Required(o, NamedSlot::VT_SLOT); + return o; + } +}; + +inline ::flatbuffers::Offset CreateNamedSlot( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::String> name = 0, + ::flatbuffers::Offset slot = 0) { + NamedSlotBuilder builder_(_fbb); + builder_.add_slot(slot); + builder_.add_name(name); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateNamedSlotDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const char *name = nullptr, + ::flatbuffers::Offset slot = 0) { + auto name__ = name ? _fbb.CreateString(name) : 0; + return mlx_delegate::CreateNamedSlot( + _fbb, + name__, + slot); +} + +struct MLXGraph FLATBUFFERS_FINAL_CLASS : private ::flatbuffers::Table { + typedef MLXGraphBuilder Builder; + enum FlatBuffersVTableOffset FLATBUFFERS_VTABLE_UNDERLYING_TYPE { + VT_VERSION = 4, + VT_NUM_CONSTANT_TENSORS = 6, + VT_NUM_INPUT_TENSORS = 8, + VT_NUM_OUTPUT_TENSORS = 10, + VT_NUM_MUTABLE_BUFFER_TENSORS = 12, + VT_NUM_TEMP_TENSORS = 14, + VT_NUM_VALUES = 16, + VT_INSTRUCTION_CHAINS = 18, + VT_MAIN_CHAIN_IDX = 20, + VT_INIT_CHAIN_IDX = 22, + VT_INPUT_MAP = 24, + VT_OUTPUT_MAP = 26, + VT_MUTABLE_BUFFER_MAP = 28, + VT_NAMED_SLOTS = 30, + VT_TENSOR_META = 32 + }; + const ::flatbuffers::String *version() const { + return GetPointer(VT_VERSION); + } + uint32_t num_constant_tensors() const { + return GetField(VT_NUM_CONSTANT_TENSORS, 0); + } + uint32_t num_input_tensors() const { + return GetField(VT_NUM_INPUT_TENSORS, 0); + } + uint32_t num_output_tensors() const { + return GetField(VT_NUM_OUTPUT_TENSORS, 0); + } + uint32_t num_mutable_buffer_tensors() const { + return GetField(VT_NUM_MUTABLE_BUFFER_TENSORS, 0); + } + uint32_t num_temp_tensors() const { + return GetField(VT_NUM_TEMP_TENSORS, 0); + } + uint32_t num_values() const { + return GetField(VT_NUM_VALUES, 0); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *instruction_chains() const { + return GetPointer> *>(VT_INSTRUCTION_CHAINS); + } + uint32_t main_chain_idx() const { + return GetField(VT_MAIN_CHAIN_IDX, 0); + } + int32_t init_chain_idx() const { + return GetField(VT_INIT_CHAIN_IDX, -1); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *input_map() const { + return GetPointer> *>(VT_INPUT_MAP); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *output_map() const { + return GetPointer> *>(VT_OUTPUT_MAP); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *mutable_buffer_map() const { + return GetPointer> *>(VT_MUTABLE_BUFFER_MAP); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *named_slots() const { + return GetPointer> *>(VT_NAMED_SLOTS); + } + const ::flatbuffers::Vector<::flatbuffers::Offset> *tensor_meta() const { + return GetPointer> *>(VT_TENSOR_META); + } + bool Verify(::flatbuffers::Verifier &verifier) const { + return VerifyTableStart(verifier) && + VerifyOffset(verifier, VT_VERSION) && + verifier.VerifyString(version()) && + VerifyField(verifier, VT_NUM_CONSTANT_TENSORS, 4) && + VerifyField(verifier, VT_NUM_INPUT_TENSORS, 4) && + VerifyField(verifier, VT_NUM_OUTPUT_TENSORS, 4) && + VerifyField(verifier, VT_NUM_MUTABLE_BUFFER_TENSORS, 4) && + VerifyField(verifier, VT_NUM_TEMP_TENSORS, 4) && + VerifyField(verifier, VT_NUM_VALUES, 4) && + VerifyOffsetRequired(verifier, VT_INSTRUCTION_CHAINS) && + verifier.VerifyVector(instruction_chains()) && + verifier.VerifyVectorOfTables(instruction_chains()) && + VerifyField(verifier, VT_MAIN_CHAIN_IDX, 4) && + VerifyField(verifier, VT_INIT_CHAIN_IDX, 4) && + VerifyOffset(verifier, VT_INPUT_MAP) && + verifier.VerifyVector(input_map()) && + verifier.VerifyVectorOfTables(input_map()) && + VerifyOffset(verifier, VT_OUTPUT_MAP) && + verifier.VerifyVector(output_map()) && + verifier.VerifyVectorOfTables(output_map()) && + VerifyOffset(verifier, VT_MUTABLE_BUFFER_MAP) && + verifier.VerifyVector(mutable_buffer_map()) && + verifier.VerifyVectorOfTables(mutable_buffer_map()) && + VerifyOffset(verifier, VT_NAMED_SLOTS) && + verifier.VerifyVector(named_slots()) && + verifier.VerifyVectorOfTables(named_slots()) && + VerifyOffset(verifier, VT_TENSOR_META) && + verifier.VerifyVector(tensor_meta()) && + verifier.VerifyVectorOfTables(tensor_meta()) && + verifier.EndTable(); + } +}; + +struct MLXGraphBuilder { + typedef MLXGraph Table; + ::flatbuffers::FlatBufferBuilder &fbb_; + ::flatbuffers::uoffset_t start_; + void add_version(::flatbuffers::Offset<::flatbuffers::String> version) { + fbb_.AddOffset(MLXGraph::VT_VERSION, version); + } + void add_num_constant_tensors(uint32_t num_constant_tensors) { + fbb_.AddElement(MLXGraph::VT_NUM_CONSTANT_TENSORS, num_constant_tensors, 0); + } + void add_num_input_tensors(uint32_t num_input_tensors) { + fbb_.AddElement(MLXGraph::VT_NUM_INPUT_TENSORS, num_input_tensors, 0); + } + void add_num_output_tensors(uint32_t num_output_tensors) { + fbb_.AddElement(MLXGraph::VT_NUM_OUTPUT_TENSORS, num_output_tensors, 0); + } + void add_num_mutable_buffer_tensors(uint32_t num_mutable_buffer_tensors) { + fbb_.AddElement(MLXGraph::VT_NUM_MUTABLE_BUFFER_TENSORS, num_mutable_buffer_tensors, 0); + } + void add_num_temp_tensors(uint32_t num_temp_tensors) { + fbb_.AddElement(MLXGraph::VT_NUM_TEMP_TENSORS, num_temp_tensors, 0); + } + void add_num_values(uint32_t num_values) { + fbb_.AddElement(MLXGraph::VT_NUM_VALUES, num_values, 0); + } + void add_instruction_chains(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> instruction_chains) { + fbb_.AddOffset(MLXGraph::VT_INSTRUCTION_CHAINS, instruction_chains); + } + void add_main_chain_idx(uint32_t main_chain_idx) { + fbb_.AddElement(MLXGraph::VT_MAIN_CHAIN_IDX, main_chain_idx, 0); + } + void add_init_chain_idx(int32_t init_chain_idx) { + fbb_.AddElement(MLXGraph::VT_INIT_CHAIN_IDX, init_chain_idx, -1); + } + void add_input_map(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> input_map) { + fbb_.AddOffset(MLXGraph::VT_INPUT_MAP, input_map); + } + void add_output_map(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> output_map) { + fbb_.AddOffset(MLXGraph::VT_OUTPUT_MAP, output_map); + } + void add_mutable_buffer_map(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> mutable_buffer_map) { + fbb_.AddOffset(MLXGraph::VT_MUTABLE_BUFFER_MAP, mutable_buffer_map); + } + void add_named_slots(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> named_slots) { + fbb_.AddOffset(MLXGraph::VT_NAMED_SLOTS, named_slots); + } + void add_tensor_meta(::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> tensor_meta) { + fbb_.AddOffset(MLXGraph::VT_TENSOR_META, tensor_meta); + } + explicit MLXGraphBuilder(::flatbuffers::FlatBufferBuilder &_fbb) + : fbb_(_fbb) { + start_ = fbb_.StartTable(); + } + ::flatbuffers::Offset Finish() { + const auto end = fbb_.EndTable(start_); + auto o = ::flatbuffers::Offset(end); + fbb_.Required(o, MLXGraph::VT_INSTRUCTION_CHAINS); + return o; + } +}; + +inline ::flatbuffers::Offset CreateMLXGraph( + ::flatbuffers::FlatBufferBuilder &_fbb, + ::flatbuffers::Offset<::flatbuffers::String> version = 0, + uint32_t num_constant_tensors = 0, + uint32_t num_input_tensors = 0, + uint32_t num_output_tensors = 0, + uint32_t num_mutable_buffer_tensors = 0, + uint32_t num_temp_tensors = 0, + uint32_t num_values = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> instruction_chains = 0, + uint32_t main_chain_idx = 0, + int32_t init_chain_idx = -1, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> input_map = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> output_map = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> mutable_buffer_map = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> named_slots = 0, + ::flatbuffers::Offset<::flatbuffers::Vector<::flatbuffers::Offset>> tensor_meta = 0) { + MLXGraphBuilder builder_(_fbb); + builder_.add_tensor_meta(tensor_meta); + builder_.add_named_slots(named_slots); + builder_.add_mutable_buffer_map(mutable_buffer_map); + builder_.add_output_map(output_map); + builder_.add_input_map(input_map); + builder_.add_init_chain_idx(init_chain_idx); + builder_.add_main_chain_idx(main_chain_idx); + builder_.add_instruction_chains(instruction_chains); + builder_.add_num_values(num_values); + builder_.add_num_temp_tensors(num_temp_tensors); + builder_.add_num_mutable_buffer_tensors(num_mutable_buffer_tensors); + builder_.add_num_output_tensors(num_output_tensors); + builder_.add_num_input_tensors(num_input_tensors); + builder_.add_num_constant_tensors(num_constant_tensors); + builder_.add_version(version); + return builder_.Finish(); +} + +inline ::flatbuffers::Offset CreateMLXGraphDirect( + ::flatbuffers::FlatBufferBuilder &_fbb, + const char *version = nullptr, + uint32_t num_constant_tensors = 0, + uint32_t num_input_tensors = 0, + uint32_t num_output_tensors = 0, + uint32_t num_mutable_buffer_tensors = 0, + uint32_t num_temp_tensors = 0, + uint32_t num_values = 0, + const std::vector<::flatbuffers::Offset> *instruction_chains = nullptr, + uint32_t main_chain_idx = 0, + int32_t init_chain_idx = -1, + const std::vector<::flatbuffers::Offset> *input_map = nullptr, + const std::vector<::flatbuffers::Offset> *output_map = nullptr, + const std::vector<::flatbuffers::Offset> *mutable_buffer_map = nullptr, + const std::vector<::flatbuffers::Offset> *named_slots = nullptr, + const std::vector<::flatbuffers::Offset> *tensor_meta = nullptr) { + auto version__ = version ? _fbb.CreateString(version) : 0; + auto instruction_chains__ = instruction_chains ? _fbb.CreateVector<::flatbuffers::Offset>(*instruction_chains) : 0; + auto input_map__ = input_map ? _fbb.CreateVector<::flatbuffers::Offset>(*input_map) : 0; + auto output_map__ = output_map ? _fbb.CreateVector<::flatbuffers::Offset>(*output_map) : 0; + auto mutable_buffer_map__ = mutable_buffer_map ? _fbb.CreateVector<::flatbuffers::Offset>(*mutable_buffer_map) : 0; + auto named_slots__ = named_slots ? _fbb.CreateVector<::flatbuffers::Offset>(*named_slots) : 0; + auto tensor_meta__ = tensor_meta ? _fbb.CreateVector<::flatbuffers::Offset>(*tensor_meta) : 0; + return mlx_delegate::CreateMLXGraph( + _fbb, + version__, + num_constant_tensors, + num_input_tensors, + num_output_tensors, + num_mutable_buffer_tensors, + num_temp_tensors, + num_values, + instruction_chains__, + main_chain_idx, + init_chain_idx, + input_map__, + output_map__, + mutable_buffer_map__, + named_slots__, + tensor_meta__); +} + +inline bool VerifyOpNode(::flatbuffers::Verifier &verifier, const void *obj, OpNode type) { + switch (type) { + case OpNode_NONE: { + return true; + } + case OpNode_NoopNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_IdCopyNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AddmmNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ItemIntNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ExpandDimsNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TileNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TakeAlongAxisNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TakeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RMSNormNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LayerNormNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RopeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SdpaNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AddNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AddIntNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SubtractIntNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MultiplyIntNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_FloorDivideIntNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SymSizeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MultiplyNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_DivideNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SubtractNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Conv1DNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Conv2DNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Conv3DNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_GeluNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ARangeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SiluNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SigmoidNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TanhNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SqueezeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SplitNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RsqrtNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MaximumNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MinimumNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LogNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SoftmaxNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_BroadcastToNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_PadNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_WhereNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ReshapeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TransposeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AsStridedNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ContiguousNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_GatherNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SliceNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AsTypeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ConcatenateNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_FullNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_FullLikeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArgmaxNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SliceUpdateNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_IndexCopyNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_DequantizeNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LessNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LessEqualNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_GreaterNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_GreaterEqualNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_EqualNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_NotEqualNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LogicalNotNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LogicalAndNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LogicalOrNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TriNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TrilNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TriuNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_FloorNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_CeilNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SquareNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ExpNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SinNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_CosNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_TanNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArcsinNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArccosNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArctanNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SinhNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_CoshNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArcsinhNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArccoshNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArctanhNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Log2Node: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Log10Node: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Log1pNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ErfNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Expm1Node: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RoundNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ReciprocalNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SqrtNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AbsNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_NegNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_Atan2Node: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LogAddExpNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_FloorDivideNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_PowerNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_LogSumExpNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SumNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MeanNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_VarNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_StdNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ProdNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MaxNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MinNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArgminNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MedianNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ModIntNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RemainderNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ConvTranspose1DNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ConvTranspose2DNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ConvTranspose3DNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ClipNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_CumsumNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_StackNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SignNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AnyNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_AllNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RepeatNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_SortNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArgsortNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_PartitionNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ArgPartitionNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_QuantizedMatmulNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ScatterAddNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_GatherMmNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_GatherQmmNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_ScanNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_MetalKernelNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_BitwiseInvertNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RollNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_BitwiseAndNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_BitwiseOrNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_BitwiseXorNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_IfNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + case OpNode_RandomBitsNode: { + auto ptr = reinterpret_cast(obj); + return verifier.VerifyTable(ptr); + } + default: return true; + } +} + +inline bool VerifyOpNodeVector(::flatbuffers::Verifier &verifier, const ::flatbuffers::Vector<::flatbuffers::Offset> *values, const ::flatbuffers::Vector *types) { + if (!values || !types) return !values && !types; + if (values->size() != types->size()) return false; + for (::flatbuffers::uoffset_t i = 0; i < values->size(); ++i) { + if (!VerifyOpNode( + verifier, values->Get(i), types->GetEnum(i))) { + return false; + } + } + return true; +} + +inline const mlx_delegate::MLXGraph *GetMLXGraph(const void *buf) { + return ::flatbuffers::GetRoot(buf); +} + +inline const mlx_delegate::MLXGraph *GetSizePrefixedMLXGraph(const void *buf) { + return ::flatbuffers::GetSizePrefixedRoot(buf); +} + +inline bool VerifyMLXGraphBuffer( + ::flatbuffers::Verifier &verifier) { + return verifier.VerifyBuffer(nullptr); +} + +inline bool VerifySizePrefixedMLXGraphBuffer( + ::flatbuffers::Verifier &verifier) { + return verifier.VerifySizePrefixedBuffer(nullptr); +} + +inline void FinishMLXGraphBuffer( + ::flatbuffers::FlatBufferBuilder &fbb, + ::flatbuffers::Offset root) { + fbb.Finish(root); +} + +inline void FinishSizePrefixedMLXGraphBuffer( + ::flatbuffers::FlatBufferBuilder &fbb, + ::flatbuffers::Offset root) { + fbb.FinishSizePrefixed(root); +} + +} // namespace mlx_delegate + +#endif // FLATBUFFERS_GENERATED_SCHEMA_MLX_DELEGATE_H_ diff --git a/backends/mlx/serialization/_generated/__init__.py b/backends/mlx/serialization/_generated/__init__.py new file mode 100644 index 00000000000..bcd49fbe30b --- /dev/null +++ b/backends/mlx/serialization/_generated/__init__.py @@ -0,0 +1,153 @@ +# Auto-generated FlatBuffer bindings +# Re-exports from mlx_delegate namespace for convenient imports + +from executorch.backends.mlx.serialization._generated.mlx_delegate.ARangeNode import ARangeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AbsNode import AbsNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AddIntNode import AddIntNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AddNode import AddNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AddmmNode import AddmmNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AllNode import AllNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AnyNode import AnyNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArccosNode import ArccosNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArccoshNode import ArccoshNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArcsinNode import ArcsinNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArcsinhNode import ArcsinhNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArctanNode import ArctanNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArctanhNode import ArctanhNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArgPartitionNode import ArgPartitionNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArgmaxNode import ArgmaxNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArgminNode import ArgminNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ArgsortNode import ArgsortNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AsStridedNode import AsStridedNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.AsTypeNode import AsTypeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Atan2Node import Atan2Node +from executorch.backends.mlx.serialization._generated.mlx_delegate.BitwiseAndNode import BitwiseAndNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.BitwiseInvertNode import BitwiseInvertNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.BitwiseOrNode import BitwiseOrNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.BitwiseXorNode import BitwiseXorNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.BroadcastToNode import BroadcastToNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.CeilNode import CeilNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ClipNode import ClipNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ConcatenateNode import ConcatenateNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ContiguousNode import ContiguousNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Conv1DNode import Conv1DNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Conv2DNode import Conv2DNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Conv3DNode import Conv3DNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ConvTranspose1DNode import ConvTranspose1DNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ConvTranspose2DNode import ConvTranspose2DNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ConvTranspose3DNode import ConvTranspose3DNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.CosNode import CosNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.CoshNode import CoshNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.CumsumNode import CumsumNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.DequantizeNode import DequantizeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.DivideNode import DivideNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.EqualNode import EqualNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ErfNode import ErfNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ExpNode import ExpNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ExpandDimsNode import ExpandDimsNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Expm1Node import Expm1Node +from executorch.backends.mlx.serialization._generated.mlx_delegate.FloatOrVid import FloatOrVid +from executorch.backends.mlx.serialization._generated.mlx_delegate.FloorDivideIntNode import FloorDivideIntNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.FloorDivideNode import FloorDivideNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.FloorNode import FloorNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.FullLikeNode import FullLikeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.FullNode import FullNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.GatherMmNode import GatherMmNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.GatherNode import GatherNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.GatherQmmNode import GatherQmmNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.GeluNode import GeluNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.GreaterEqualNode import GreaterEqualNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.GreaterNode import GreaterNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.IdCopyNode import IdCopyNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.IfNode import IfNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.IndexCopyNode import IndexCopyNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Instruction import Instruction +from executorch.backends.mlx.serialization._generated.mlx_delegate.InstructionChain import InstructionChain +from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid +from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVidOrTid import IntOrVidOrTid +from executorch.backends.mlx.serialization._generated.mlx_delegate.ItemIntNode import ItemIntNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LayerNormNode import LayerNormNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LessEqualNode import LessEqualNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LessNode import LessNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Log10Node import Log10Node +from executorch.backends.mlx.serialization._generated.mlx_delegate.Log1pNode import Log1pNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Log2Node import Log2Node +from executorch.backends.mlx.serialization._generated.mlx_delegate.LogAddExpNode import LogAddExpNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LogNode import LogNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LogSumExpNode import LogSumExpNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LogicalAndNode import LogicalAndNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LogicalNotNode import LogicalNotNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.LogicalOrNode import LogicalOrNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MLXGraph import MLXGraph +from executorch.backends.mlx.serialization._generated.mlx_delegate.MaxNode import MaxNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MaximumNode import MaximumNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MeanNode import MeanNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MedianNode import MedianNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MetalKernelNode import MetalKernelNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MinNode import MinNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MinimumNode import MinimumNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ModIntNode import ModIntNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MultiplyIntNode import MultiplyIntNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.MultiplyNode import MultiplyNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.NamedSlot import NamedSlot +from executorch.backends.mlx.serialization._generated.mlx_delegate.NegNode import NegNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.NoopNode import NoopNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.NotEqualNode import NotEqualNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.OpNode import OpNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.PadNode import PadNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.PartitionNode import PartitionNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.PowerNode import PowerNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ProdNode import ProdNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.QuantizedMatmulNode import QuantizedMatmulNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RMSNormNode import RMSNormNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RandomBitsNode import RandomBitsNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ReciprocalNode import ReciprocalNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RemainderNode import RemainderNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RepeatNode import RepeatNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ReshapeNode import ReshapeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RollNode import RollNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RopeNode import RopeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RoundNode import RoundNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.RsqrtNode import RsqrtNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ScanNode import ScanNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ScatterAddNode import ScatterAddNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SdpaNode import SdpaNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.ShapeDim import ShapeDim +from executorch.backends.mlx.serialization._generated.mlx_delegate.SigmoidNode import SigmoidNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SignNode import SignNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SiluNode import SiluNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SinNode import SinNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SinhNode import SinhNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SliceNode import SliceNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SliceUpdateNode import SliceUpdateNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SlotType import SlotType +from executorch.backends.mlx.serialization._generated.mlx_delegate.SlotVariant import SlotVariant +from executorch.backends.mlx.serialization._generated.mlx_delegate.SoftmaxNode import SoftmaxNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SortNode import SortNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SplitNode import SplitNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SqrtNode import SqrtNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SquareNode import SquareNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SqueezeNode import SqueezeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.StackNode import StackNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.StdNode import StdNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SubtractIntNode import SubtractIntNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SubtractNode import SubtractNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SumNode import SumNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.SymSizeNode import SymSizeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TakeAlongAxisNode import TakeAlongAxisNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TakeNode import TakeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TanNode import TanNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TanhNode import TanhNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TensorMeta import TensorMeta +from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid +from executorch.backends.mlx.serialization._generated.mlx_delegate.TileNode import TileNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TransposeNode import TransposeNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TriNode import TriNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TrilNode import TrilNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.TriuNode import TriuNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.VarNode import VarNode +from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid +from executorch.backends.mlx.serialization._generated.mlx_delegate.VidOrTid import VidOrTid +from executorch.backends.mlx.serialization._generated.mlx_delegate.WhereNode import WhereNode + +__all__ = ['ARangeNode', 'AbsNode', 'AddIntNode', 'AddNode', 'AddmmNode', 'AllNode', 'AnyNode', 'ArccosNode', 'ArccoshNode', 'ArcsinNode', 'ArcsinhNode', 'ArctanNode', 'ArctanhNode', 'ArgPartitionNode', 'ArgmaxNode', 'ArgminNode', 'ArgsortNode', 'AsStridedNode', 'AsTypeNode', 'Atan2Node', 'BitwiseAndNode', 'BitwiseInvertNode', 'BitwiseOrNode', 'BitwiseXorNode', 'BroadcastToNode', 'CeilNode', 'ClipNode', 'ConcatenateNode', 'ContiguousNode', 'Conv1DNode', 'Conv2DNode', 'Conv3DNode', 'ConvTranspose1DNode', 'ConvTranspose2DNode', 'ConvTranspose3DNode', 'CosNode', 'CoshNode', 'CumsumNode', 'DequantizeNode', 'DivideNode', 'EqualNode', 'ErfNode', 'ExpNode', 'ExpandDimsNode', 'Expm1Node', 'FloatOrVid', 'FloorDivideIntNode', 'FloorDivideNode', 'FloorNode', 'FullLikeNode', 'FullNode', 'GatherMmNode', 'GatherNode', 'GatherQmmNode', 'GeluNode', 'GreaterEqualNode', 'GreaterNode', 'IdCopyNode', 'IfNode', 'IndexCopyNode', 'Instruction', 'InstructionChain', 'IntOrVid', 'IntOrVidOrTid', 'ItemIntNode', 'LayerNormNode', 'LessEqualNode', 'LessNode', 'Log10Node', 'Log1pNode', 'Log2Node', 'LogAddExpNode', 'LogNode', 'LogSumExpNode', 'LogicalAndNode', 'LogicalNotNode', 'LogicalOrNode', 'MLXGraph', 'MaxNode', 'MaximumNode', 'MeanNode', 'MedianNode', 'MetalKernelNode', 'MinNode', 'MinimumNode', 'ModIntNode', 'MultiplyIntNode', 'MultiplyNode', 'NamedSlot', 'NegNode', 'NoopNode', 'NotEqualNode', 'OpNode', 'PadNode', 'PartitionNode', 'PowerNode', 'ProdNode', 'QuantizedMatmulNode', 'RMSNormNode', 'RandomBitsNode', 'ReciprocalNode', 'RemainderNode', 'RepeatNode', 'ReshapeNode', 'RollNode', 'RopeNode', 'RoundNode', 'RsqrtNode', 'ScanNode', 'ScatterAddNode', 'SdpaNode', 'ShapeDim', 'SigmoidNode', 'SignNode', 'SiluNode', 'SinNode', 'SinhNode', 'SliceNode', 'SliceUpdateNode', 'SlotType', 'SlotVariant', 'SoftmaxNode', 'SortNode', 'SplitNode', 'SqrtNode', 'SquareNode', 'SqueezeNode', 'StackNode', 'StdNode', 'SubtractIntNode', 'SubtractNode', 'SumNode', 'SymSizeNode', 'TakeAlongAxisNode', 'TakeNode', 'TanNode', 'TanhNode', 'TensorMeta', 'Tid', 'TileNode', 'TransposeNode', 'TriNode', 'TrilNode', 'TriuNode', 'VarNode', 'Vid', 'VidOrTid', 'WhereNode'] diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ARangeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ARangeNode.py new file mode 100644 index 00000000000..4eb905226de --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ARangeNode.py @@ -0,0 +1,118 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ARangeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ARangeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsARangeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ARangeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ARangeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ARangeNode + def Start(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ARangeNode + def Stop(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ARangeNode + def Step(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ARangeNode + def ScalarType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return None + +def ARangeNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + ARangeNodeStart(builder) + +def ARangeNodeAddOut(builder, out): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ARangeNodeAddOut(builder, out) + +def ARangeNodeAddStart(builder, start): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(start), 0) + +def AddStart(builder, start): + ARangeNodeAddStart(builder, start) + +def ARangeNodeAddStop(builder, stop): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(stop), 0) + +def AddStop(builder, stop): + ARangeNodeAddStop(builder, stop) + +def ARangeNodeAddStep(builder, step): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(step), 0) + +def AddStep(builder, step): + ARangeNodeAddStep(builder, step) + +def ARangeNodeAddScalarType(builder, scalarType): + builder.PrependInt8Slot(4, scalarType, None) + +def AddScalarType(builder, scalarType): + ARangeNodeAddScalarType(builder, scalarType) + +def ARangeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ARangeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AbsNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AbsNode.py new file mode 100644 index 00000000000..7e24a21be8c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AbsNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AbsNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AbsNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAbsNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AbsNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AbsNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AbsNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def AbsNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + AbsNodeStart(builder) + +def AbsNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + AbsNodeAddX(builder, x) + +def AbsNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AbsNodeAddOut(builder, out) + +def AbsNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AbsNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AddIntNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AddIntNode.py new file mode 100644 index 00000000000..a437f7dd694 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AddIntNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AddIntNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AddIntNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAddIntNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AddIntNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AddIntNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddIntNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddIntNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def AddIntNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + AddIntNodeStart(builder) + +def AddIntNodeAddA(builder, a): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + AddIntNodeAddA(builder, a) + +def AddIntNodeAddB(builder, b): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + AddIntNodeAddB(builder, b) + +def AddIntNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AddIntNodeAddOut(builder, out) + +def AddIntNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AddIntNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AddNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AddNode.py new file mode 100644 index 00000000000..0af5493f001 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AddNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AddNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AddNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAddNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AddNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AddNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def AddNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + AddNodeStart(builder) + +def AddNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + AddNodeAddA(builder, a) + +def AddNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + AddNodeAddB(builder, b) + +def AddNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AddNodeAddOut(builder, out) + +def AddNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AddNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AddmmNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AddmmNode.py new file mode 100644 index 00000000000..269ec0d0d3e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AddmmNode.py @@ -0,0 +1,131 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AddmmNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AddmmNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAddmmNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AddmmNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AddmmNode + def Mat1(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddmmNode + def Mat2(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddmmNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddmmNode + def Bias(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AddmmNode + def Alpha(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 1.0 + + # AddmmNode + def Beta(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 1.0 + +def AddmmNodeStart(builder): + builder.StartObject(6) + +def Start(builder): + AddmmNodeStart(builder) + +def AddmmNodeAddMat1(builder, mat1): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(mat1), 0) + +def AddMat1(builder, mat1): + AddmmNodeAddMat1(builder, mat1) + +def AddmmNodeAddMat2(builder, mat2): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(mat2), 0) + +def AddMat2(builder, mat2): + AddmmNodeAddMat2(builder, mat2) + +def AddmmNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AddmmNodeAddOut(builder, out) + +def AddmmNodeAddBias(builder, bias): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(bias), 0) + +def AddBias(builder, bias): + AddmmNodeAddBias(builder, bias) + +def AddmmNodeAddAlpha(builder, alpha): + builder.PrependFloat32Slot(4, alpha, 1.0) + +def AddAlpha(builder, alpha): + AddmmNodeAddAlpha(builder, alpha) + +def AddmmNodeAddBeta(builder, beta): + builder.PrependFloat32Slot(5, beta, 1.0) + +def AddBeta(builder, beta): + AddmmNodeAddBeta(builder, beta) + +def AddmmNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AddmmNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AllNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AllNode.py new file mode 100644 index 00000000000..f836968be6c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AllNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AllNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AllNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAllNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AllNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AllNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AllNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AllNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # AllNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # AllNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # AllNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # AllNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def AllNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + AllNodeStart(builder) + +def AllNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + AllNodeAddX(builder, x) + +def AllNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AllNodeAddOut(builder, out) + +def AllNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + AllNodeAddAxes(builder, axes) + +def AllNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return AllNodeStartAxesVector(builder, numElems) + +def AllNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + AllNodeAddKeepdims(builder, keepdims) + +def AllNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AllNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AnyNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AnyNode.py new file mode 100644 index 00000000000..7476789f01a --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AnyNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AnyNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AnyNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAnyNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AnyNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AnyNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AnyNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AnyNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # AnyNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # AnyNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # AnyNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # AnyNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def AnyNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + AnyNodeStart(builder) + +def AnyNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + AnyNodeAddX(builder, x) + +def AnyNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AnyNodeAddOut(builder, out) + +def AnyNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + AnyNodeAddAxes(builder, axes) + +def AnyNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return AnyNodeStartAxesVector(builder, numElems) + +def AnyNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + AnyNodeAddKeepdims(builder, keepdims) + +def AnyNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AnyNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArccosNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArccosNode.py new file mode 100644 index 00000000000..36c0dd1d02b --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArccosNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArccosNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArccosNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArccosNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArccosNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArccosNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArccosNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ArccosNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ArccosNodeStart(builder) + +def ArccosNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArccosNodeAddX(builder, x) + +def ArccosNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArccosNodeAddOut(builder, out) + +def ArccosNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArccosNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArccoshNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArccoshNode.py new file mode 100644 index 00000000000..b4138a35ab3 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArccoshNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArccoshNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArccoshNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArccoshNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArccoshNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArccoshNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArccoshNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ArccoshNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ArccoshNodeStart(builder) + +def ArccoshNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArccoshNodeAddX(builder, x) + +def ArccoshNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArccoshNodeAddOut(builder, out) + +def ArccoshNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArccoshNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArcsinNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArcsinNode.py new file mode 100644 index 00000000000..d1f5c4f0dfa --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArcsinNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArcsinNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArcsinNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArcsinNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArcsinNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArcsinNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArcsinNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ArcsinNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ArcsinNodeStart(builder) + +def ArcsinNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArcsinNodeAddX(builder, x) + +def ArcsinNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArcsinNodeAddOut(builder, out) + +def ArcsinNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArcsinNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArcsinhNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArcsinhNode.py new file mode 100644 index 00000000000..62402c2610a --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArcsinhNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArcsinhNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArcsinhNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArcsinhNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArcsinhNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArcsinhNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArcsinhNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ArcsinhNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ArcsinhNodeStart(builder) + +def ArcsinhNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArcsinhNodeAddX(builder, x) + +def ArcsinhNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArcsinhNodeAddOut(builder, out) + +def ArcsinhNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArcsinhNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArctanNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArctanNode.py new file mode 100644 index 00000000000..e004684e0cb --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArctanNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArctanNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArctanNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArctanNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArctanNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArctanNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArctanNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ArctanNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ArctanNodeStart(builder) + +def ArctanNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArctanNodeAddX(builder, x) + +def ArctanNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArctanNodeAddOut(builder, out) + +def ArctanNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArctanNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArctanhNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArctanhNode.py new file mode 100644 index 00000000000..491ebd2d401 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArctanhNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArctanhNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArctanhNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArctanhNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArctanhNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArctanhNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArctanhNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ArctanhNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ArctanhNodeStart(builder) + +def ArctanhNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArctanhNodeAddX(builder, x) + +def ArctanhNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArctanhNodeAddOut(builder, out) + +def ArctanhNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArctanhNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArgPartitionNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArgPartitionNode.py new file mode 100644 index 00000000000..6fa6ceecbb2 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArgPartitionNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArgPartitionNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArgPartitionNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArgPartitionNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArgPartitionNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArgPartitionNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgPartitionNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgPartitionNode + def Kth(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgPartitionNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def ArgPartitionNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + ArgPartitionNodeStart(builder) + +def ArgPartitionNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArgPartitionNodeAddX(builder, x) + +def ArgPartitionNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArgPartitionNodeAddOut(builder, out) + +def ArgPartitionNodeAddKth(builder, kth): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(kth), 0) + +def AddKth(builder, kth): + ArgPartitionNodeAddKth(builder, kth) + +def ArgPartitionNodeAddAxis(builder, axis): + builder.PrependInt32Slot(3, axis, 0) + +def AddAxis(builder, axis): + ArgPartitionNodeAddAxis(builder, axis) + +def ArgPartitionNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArgPartitionNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArgmaxNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArgmaxNode.py new file mode 100644 index 00000000000..81cbf423a18 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArgmaxNode.py @@ -0,0 +1,97 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArgmaxNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArgmaxNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArgmaxNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArgmaxNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArgmaxNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgmaxNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgmaxNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ArgmaxNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def ArgmaxNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + ArgmaxNodeStart(builder) + +def ArgmaxNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArgmaxNodeAddX(builder, x) + +def ArgmaxNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArgmaxNodeAddOut(builder, out) + +def ArgmaxNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + ArgmaxNodeAddAxis(builder, axis) + +def ArgmaxNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + ArgmaxNodeAddKeepdims(builder, keepdims) + +def ArgmaxNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArgmaxNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArgminNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArgminNode.py new file mode 100644 index 00000000000..acca75a883d --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArgminNode.py @@ -0,0 +1,97 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArgminNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArgminNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArgminNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArgminNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArgminNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgminNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgminNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ArgminNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def ArgminNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + ArgminNodeStart(builder) + +def ArgminNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArgminNodeAddX(builder, x) + +def ArgminNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArgminNodeAddOut(builder, out) + +def ArgminNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + ArgminNodeAddAxis(builder, axis) + +def ArgminNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + ArgminNodeAddKeepdims(builder, keepdims) + +def ArgminNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArgminNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ArgsortNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ArgsortNode.py new file mode 100644 index 00000000000..6bcfbb2d1c1 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ArgsortNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ArgsortNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ArgsortNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsArgsortNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ArgsortNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ArgsortNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgsortNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ArgsortNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def ArgsortNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + ArgsortNodeStart(builder) + +def ArgsortNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ArgsortNodeAddX(builder, x) + +def ArgsortNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ArgsortNodeAddOut(builder, out) + +def ArgsortNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + ArgsortNodeAddAxis(builder, axis) + +def ArgsortNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ArgsortNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AsStridedNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AsStridedNode.py new file mode 100644 index 00000000000..64fdb9730b5 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AsStridedNode.py @@ -0,0 +1,158 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AsStridedNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AsStridedNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAsStridedNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AsStridedNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AsStridedNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AsStridedNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AsStridedNode + def Shape(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AsStridedNode + def ShapeLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # AsStridedNode + def ShapeIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # AsStridedNode + def Strides(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AsStridedNode + def StridesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # AsStridedNode + def StridesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + return o == 0 + + # AsStridedNode + def Offset(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint64Flags, o + self._tab.Pos) + return 0 + +def AsStridedNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + AsStridedNodeStart(builder) + +def AsStridedNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + AsStridedNodeAddX(builder, x) + +def AsStridedNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AsStridedNodeAddOut(builder, out) + +def AsStridedNodeAddShape(builder, shape): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(shape), 0) + +def AddShape(builder, shape): + AsStridedNodeAddShape(builder, shape) + +def AsStridedNodeStartShapeVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShapeVector(builder, numElems): + return AsStridedNodeStartShapeVector(builder, numElems) + +def AsStridedNodeAddStrides(builder, strides): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(strides), 0) + +def AddStrides(builder, strides): + AsStridedNodeAddStrides(builder, strides) + +def AsStridedNodeStartStridesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartStridesVector(builder, numElems): + return AsStridedNodeStartStridesVector(builder, numElems) + +def AsStridedNodeAddOffset(builder, offset): + builder.PrependUint64Slot(4, offset, 0) + +def AddOffset(builder, offset): + AsStridedNodeAddOffset(builder, offset) + +def AsStridedNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AsStridedNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/AsTypeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/AsTypeNode.py new file mode 100644 index 00000000000..603634f6fed --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/AsTypeNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class AsTypeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = AsTypeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAsTypeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # AsTypeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # AsTypeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AsTypeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # AsTypeNode + def ScalarType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return 0 + +def AsTypeNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + AsTypeNodeStart(builder) + +def AsTypeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + AsTypeNodeAddX(builder, x) + +def AsTypeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + AsTypeNodeAddOut(builder, out) + +def AsTypeNodeAddScalarType(builder, scalarType): + builder.PrependInt8Slot(2, scalarType, 0) + +def AddScalarType(builder, scalarType): + AsTypeNodeAddScalarType(builder, scalarType) + +def AsTypeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return AsTypeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Atan2Node.py b/backends/mlx/serialization/_generated/mlx_delegate/Atan2Node.py new file mode 100644 index 00000000000..aaa19cf8dcc --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Atan2Node.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Atan2Node(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Atan2Node() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsAtan2Node(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Atan2Node + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Atan2Node + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Atan2Node + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Atan2Node + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def Atan2NodeStart(builder): + builder.StartObject(3) + +def Start(builder): + Atan2NodeStart(builder) + +def Atan2NodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + Atan2NodeAddA(builder, a) + +def Atan2NodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + Atan2NodeAddB(builder, b) + +def Atan2NodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Atan2NodeAddOut(builder, out) + +def Atan2NodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Atan2NodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/BitwiseAndNode.py b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseAndNode.py new file mode 100644 index 00000000000..029ad7ee76e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseAndNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class BitwiseAndNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = BitwiseAndNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsBitwiseAndNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # BitwiseAndNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # BitwiseAndNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseAndNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseAndNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def BitwiseAndNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + BitwiseAndNodeStart(builder) + +def BitwiseAndNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + BitwiseAndNodeAddA(builder, a) + +def BitwiseAndNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + BitwiseAndNodeAddB(builder, b) + +def BitwiseAndNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + BitwiseAndNodeAddOut(builder, out) + +def BitwiseAndNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return BitwiseAndNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/BitwiseInvertNode.py b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseInvertNode.py new file mode 100644 index 00000000000..b6ce6f5a90b --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseInvertNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class BitwiseInvertNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = BitwiseInvertNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsBitwiseInvertNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # BitwiseInvertNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # BitwiseInvertNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseInvertNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def BitwiseInvertNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + BitwiseInvertNodeStart(builder) + +def BitwiseInvertNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + BitwiseInvertNodeAddX(builder, x) + +def BitwiseInvertNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + BitwiseInvertNodeAddOut(builder, out) + +def BitwiseInvertNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return BitwiseInvertNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/BitwiseOrNode.py b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseOrNode.py new file mode 100644 index 00000000000..a6b57d68bd6 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseOrNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class BitwiseOrNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = BitwiseOrNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsBitwiseOrNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # BitwiseOrNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # BitwiseOrNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseOrNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseOrNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def BitwiseOrNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + BitwiseOrNodeStart(builder) + +def BitwiseOrNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + BitwiseOrNodeAddA(builder, a) + +def BitwiseOrNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + BitwiseOrNodeAddB(builder, b) + +def BitwiseOrNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + BitwiseOrNodeAddOut(builder, out) + +def BitwiseOrNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return BitwiseOrNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/BitwiseXorNode.py b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseXorNode.py new file mode 100644 index 00000000000..a3fb4685c92 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/BitwiseXorNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class BitwiseXorNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = BitwiseXorNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsBitwiseXorNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # BitwiseXorNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # BitwiseXorNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseXorNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BitwiseXorNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def BitwiseXorNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + BitwiseXorNodeStart(builder) + +def BitwiseXorNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + BitwiseXorNodeAddA(builder, a) + +def BitwiseXorNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + BitwiseXorNodeAddB(builder, b) + +def BitwiseXorNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + BitwiseXorNodeAddOut(builder, out) + +def BitwiseXorNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return BitwiseXorNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/BroadcastToNode.py b/backends/mlx/serialization/_generated/mlx_delegate/BroadcastToNode.py new file mode 100644 index 00000000000..e05d836b584 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/BroadcastToNode.py @@ -0,0 +1,108 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class BroadcastToNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = BroadcastToNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsBroadcastToNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # BroadcastToNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # BroadcastToNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BroadcastToNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BroadcastToNode + def Shape(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # BroadcastToNode + def ShapeLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # BroadcastToNode + def ShapeIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + +def BroadcastToNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + BroadcastToNodeStart(builder) + +def BroadcastToNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + BroadcastToNodeAddX(builder, x) + +def BroadcastToNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + BroadcastToNodeAddOut(builder, out) + +def BroadcastToNodeAddShape(builder, shape): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(shape), 0) + +def AddShape(builder, shape): + BroadcastToNodeAddShape(builder, shape) + +def BroadcastToNodeStartShapeVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShapeVector(builder, numElems): + return BroadcastToNodeStartShapeVector(builder, numElems) + +def BroadcastToNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return BroadcastToNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/CeilNode.py b/backends/mlx/serialization/_generated/mlx_delegate/CeilNode.py new file mode 100644 index 00000000000..ee189a9cb9e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/CeilNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class CeilNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = CeilNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsCeilNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # CeilNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # CeilNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # CeilNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def CeilNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + CeilNodeStart(builder) + +def CeilNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + CeilNodeAddX(builder, x) + +def CeilNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + CeilNodeAddOut(builder, out) + +def CeilNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return CeilNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ClipNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ClipNode.py new file mode 100644 index 00000000000..0c5f1c87c7d --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ClipNode.py @@ -0,0 +1,105 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ClipNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ClipNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsClipNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ClipNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ClipNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ClipNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ClipNode + def AMin(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ClipNode + def AMax(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ClipNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + ClipNodeStart(builder) + +def ClipNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ClipNodeAddX(builder, x) + +def ClipNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ClipNodeAddOut(builder, out) + +def ClipNodeAddAMin(builder, aMin): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(aMin), 0) + +def AddAMin(builder, aMin): + ClipNodeAddAMin(builder, aMin) + +def ClipNodeAddAMax(builder, aMax): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(aMax), 0) + +def AddAMax(builder, aMax): + ClipNodeAddAMax(builder, aMax) + +def ClipNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ClipNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ConcatenateNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ConcatenateNode.py new file mode 100644 index 00000000000..dd8a9d7a853 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ConcatenateNode.py @@ -0,0 +1,103 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ConcatenateNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ConcatenateNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConcatenateNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ConcatenateNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ConcatenateNode + def Tensors(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConcatenateNode + def TensorsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ConcatenateNode + def TensorsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + return o == 0 + + # ConcatenateNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConcatenateNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def ConcatenateNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + ConcatenateNodeStart(builder) + +def ConcatenateNodeAddTensors(builder, tensors): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(tensors), 0) + +def AddTensors(builder, tensors): + ConcatenateNodeAddTensors(builder, tensors) + +def ConcatenateNodeStartTensorsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartTensorsVector(builder, numElems): + return ConcatenateNodeStartTensorsVector(builder, numElems) + +def ConcatenateNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ConcatenateNodeAddOut(builder, out) + +def ConcatenateNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + ConcatenateNodeAddAxis(builder, axis) + +def ConcatenateNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ConcatenateNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ContiguousNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ContiguousNode.py new file mode 100644 index 00000000000..26b759fd44b --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ContiguousNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ContiguousNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ContiguousNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsContiguousNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ContiguousNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ContiguousNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ContiguousNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ContiguousNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ContiguousNodeStart(builder) + +def ContiguousNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ContiguousNodeAddX(builder, x) + +def ContiguousNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ContiguousNodeAddOut(builder, out) + +def ContiguousNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ContiguousNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Conv1DNode.py b/backends/mlx/serialization/_generated/mlx_delegate/Conv1DNode.py new file mode 100644 index 00000000000..5f24f40215b --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Conv1DNode.py @@ -0,0 +1,140 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Conv1DNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Conv1DNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConv1DNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Conv1DNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Conv1DNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv1DNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv1DNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv1DNode + def Stride(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv1DNode + def Padding(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # Conv1DNode + def Dilation(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv1DNode + def Groups(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def Conv1DNodeStart(builder): + builder.StartObject(7) + +def Start(builder): + Conv1DNodeStart(builder) + +def Conv1DNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Conv1DNodeAddX(builder, x) + +def Conv1DNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + Conv1DNodeAddW(builder, w) + +def Conv1DNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Conv1DNodeAddOut(builder, out) + +def Conv1DNodeAddStride(builder, stride): + builder.PrependInt32Slot(3, stride, 1) + +def AddStride(builder, stride): + Conv1DNodeAddStride(builder, stride) + +def Conv1DNodeAddPadding(builder, padding): + builder.PrependInt32Slot(4, padding, 0) + +def AddPadding(builder, padding): + Conv1DNodeAddPadding(builder, padding) + +def Conv1DNodeAddDilation(builder, dilation): + builder.PrependInt32Slot(5, dilation, 1) + +def AddDilation(builder, dilation): + Conv1DNodeAddDilation(builder, dilation) + +def Conv1DNodeAddGroups(builder, groups): + builder.PrependInt32Slot(6, groups, 1) + +def AddGroups(builder, groups): + Conv1DNodeAddGroups(builder, groups) + +def Conv1DNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Conv1DNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Conv2DNode.py b/backends/mlx/serialization/_generated/mlx_delegate/Conv2DNode.py new file mode 100644 index 00000000000..9a49a19c862 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Conv2DNode.py @@ -0,0 +1,179 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Conv2DNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Conv2DNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConv2DNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Conv2DNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Conv2DNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv2DNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv2DNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv2DNode + def StrideH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv2DNode + def StrideW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv2DNode + def PaddingH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # Conv2DNode + def PaddingW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # Conv2DNode + def DilationH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv2DNode + def DilationW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv2DNode + def Groups(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def Conv2DNodeStart(builder): + builder.StartObject(10) + +def Start(builder): + Conv2DNodeStart(builder) + +def Conv2DNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Conv2DNodeAddX(builder, x) + +def Conv2DNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + Conv2DNodeAddW(builder, w) + +def Conv2DNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Conv2DNodeAddOut(builder, out) + +def Conv2DNodeAddStrideH(builder, strideH): + builder.PrependInt32Slot(3, strideH, 1) + +def AddStrideH(builder, strideH): + Conv2DNodeAddStrideH(builder, strideH) + +def Conv2DNodeAddStrideW(builder, strideW): + builder.PrependInt32Slot(4, strideW, 1) + +def AddStrideW(builder, strideW): + Conv2DNodeAddStrideW(builder, strideW) + +def Conv2DNodeAddPaddingH(builder, paddingH): + builder.PrependInt32Slot(5, paddingH, 0) + +def AddPaddingH(builder, paddingH): + Conv2DNodeAddPaddingH(builder, paddingH) + +def Conv2DNodeAddPaddingW(builder, paddingW): + builder.PrependInt32Slot(6, paddingW, 0) + +def AddPaddingW(builder, paddingW): + Conv2DNodeAddPaddingW(builder, paddingW) + +def Conv2DNodeAddDilationH(builder, dilationH): + builder.PrependInt32Slot(7, dilationH, 1) + +def AddDilationH(builder, dilationH): + Conv2DNodeAddDilationH(builder, dilationH) + +def Conv2DNodeAddDilationW(builder, dilationW): + builder.PrependInt32Slot(8, dilationW, 1) + +def AddDilationW(builder, dilationW): + Conv2DNodeAddDilationW(builder, dilationW) + +def Conv2DNodeAddGroups(builder, groups): + builder.PrependInt32Slot(9, groups, 1) + +def AddGroups(builder, groups): + Conv2DNodeAddGroups(builder, groups) + +def Conv2DNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Conv2DNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Conv3DNode.py b/backends/mlx/serialization/_generated/mlx_delegate/Conv3DNode.py new file mode 100644 index 00000000000..90dc9dad150 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Conv3DNode.py @@ -0,0 +1,218 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Conv3DNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Conv3DNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConv3DNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Conv3DNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Conv3DNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv3DNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv3DNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Conv3DNode + def StrideD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv3DNode + def StrideH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv3DNode + def StrideW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv3DNode + def PaddingD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # Conv3DNode + def PaddingH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # Conv3DNode + def PaddingW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # Conv3DNode + def DilationD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv3DNode + def DilationH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv3DNode + def DilationW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # Conv3DNode + def Groups(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def Conv3DNodeStart(builder): + builder.StartObject(13) + +def Start(builder): + Conv3DNodeStart(builder) + +def Conv3DNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Conv3DNodeAddX(builder, x) + +def Conv3DNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + Conv3DNodeAddW(builder, w) + +def Conv3DNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Conv3DNodeAddOut(builder, out) + +def Conv3DNodeAddStrideD(builder, strideD): + builder.PrependInt32Slot(3, strideD, 1) + +def AddStrideD(builder, strideD): + Conv3DNodeAddStrideD(builder, strideD) + +def Conv3DNodeAddStrideH(builder, strideH): + builder.PrependInt32Slot(4, strideH, 1) + +def AddStrideH(builder, strideH): + Conv3DNodeAddStrideH(builder, strideH) + +def Conv3DNodeAddStrideW(builder, strideW): + builder.PrependInt32Slot(5, strideW, 1) + +def AddStrideW(builder, strideW): + Conv3DNodeAddStrideW(builder, strideW) + +def Conv3DNodeAddPaddingD(builder, paddingD): + builder.PrependInt32Slot(6, paddingD, 0) + +def AddPaddingD(builder, paddingD): + Conv3DNodeAddPaddingD(builder, paddingD) + +def Conv3DNodeAddPaddingH(builder, paddingH): + builder.PrependInt32Slot(7, paddingH, 0) + +def AddPaddingH(builder, paddingH): + Conv3DNodeAddPaddingH(builder, paddingH) + +def Conv3DNodeAddPaddingW(builder, paddingW): + builder.PrependInt32Slot(8, paddingW, 0) + +def AddPaddingW(builder, paddingW): + Conv3DNodeAddPaddingW(builder, paddingW) + +def Conv3DNodeAddDilationD(builder, dilationD): + builder.PrependInt32Slot(9, dilationD, 1) + +def AddDilationD(builder, dilationD): + Conv3DNodeAddDilationD(builder, dilationD) + +def Conv3DNodeAddDilationH(builder, dilationH): + builder.PrependInt32Slot(10, dilationH, 1) + +def AddDilationH(builder, dilationH): + Conv3DNodeAddDilationH(builder, dilationH) + +def Conv3DNodeAddDilationW(builder, dilationW): + builder.PrependInt32Slot(11, dilationW, 1) + +def AddDilationW(builder, dilationW): + Conv3DNodeAddDilationW(builder, dilationW) + +def Conv3DNodeAddGroups(builder, groups): + builder.PrependInt32Slot(12, groups, 1) + +def AddGroups(builder, groups): + Conv3DNodeAddGroups(builder, groups) + +def Conv3DNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Conv3DNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose1DNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose1DNode.py new file mode 100644 index 00000000000..75e79eb2922 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose1DNode.py @@ -0,0 +1,153 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ConvTranspose1DNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ConvTranspose1DNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConvTranspose1DNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ConvTranspose1DNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ConvTranspose1DNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose1DNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose1DNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose1DNode + def Stride(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose1DNode + def Padding(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose1DNode + def Dilation(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose1DNode + def OutputPadding(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose1DNode + def Groups(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def ConvTranspose1DNodeStart(builder): + builder.StartObject(8) + +def Start(builder): + ConvTranspose1DNodeStart(builder) + +def ConvTranspose1DNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ConvTranspose1DNodeAddX(builder, x) + +def ConvTranspose1DNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + ConvTranspose1DNodeAddW(builder, w) + +def ConvTranspose1DNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ConvTranspose1DNodeAddOut(builder, out) + +def ConvTranspose1DNodeAddStride(builder, stride): + builder.PrependInt32Slot(3, stride, 1) + +def AddStride(builder, stride): + ConvTranspose1DNodeAddStride(builder, stride) + +def ConvTranspose1DNodeAddPadding(builder, padding): + builder.PrependInt32Slot(4, padding, 0) + +def AddPadding(builder, padding): + ConvTranspose1DNodeAddPadding(builder, padding) + +def ConvTranspose1DNodeAddDilation(builder, dilation): + builder.PrependInt32Slot(5, dilation, 1) + +def AddDilation(builder, dilation): + ConvTranspose1DNodeAddDilation(builder, dilation) + +def ConvTranspose1DNodeAddOutputPadding(builder, outputPadding): + builder.PrependInt32Slot(6, outputPadding, 0) + +def AddOutputPadding(builder, outputPadding): + ConvTranspose1DNodeAddOutputPadding(builder, outputPadding) + +def ConvTranspose1DNodeAddGroups(builder, groups): + builder.PrependInt32Slot(7, groups, 1) + +def AddGroups(builder, groups): + ConvTranspose1DNodeAddGroups(builder, groups) + +def ConvTranspose1DNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ConvTranspose1DNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose2DNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose2DNode.py new file mode 100644 index 00000000000..9e4b9fb48b9 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose2DNode.py @@ -0,0 +1,205 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ConvTranspose2DNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ConvTranspose2DNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConvTranspose2DNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ConvTranspose2DNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ConvTranspose2DNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose2DNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose2DNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose2DNode + def StrideH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose2DNode + def StrideW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose2DNode + def PaddingH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose2DNode + def PaddingW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose2DNode + def DilationH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose2DNode + def DilationW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose2DNode + def OutputPaddingH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose2DNode + def OutputPaddingW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose2DNode + def Groups(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def ConvTranspose2DNodeStart(builder): + builder.StartObject(12) + +def Start(builder): + ConvTranspose2DNodeStart(builder) + +def ConvTranspose2DNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ConvTranspose2DNodeAddX(builder, x) + +def ConvTranspose2DNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + ConvTranspose2DNodeAddW(builder, w) + +def ConvTranspose2DNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ConvTranspose2DNodeAddOut(builder, out) + +def ConvTranspose2DNodeAddStrideH(builder, strideH): + builder.PrependInt32Slot(3, strideH, 1) + +def AddStrideH(builder, strideH): + ConvTranspose2DNodeAddStrideH(builder, strideH) + +def ConvTranspose2DNodeAddStrideW(builder, strideW): + builder.PrependInt32Slot(4, strideW, 1) + +def AddStrideW(builder, strideW): + ConvTranspose2DNodeAddStrideW(builder, strideW) + +def ConvTranspose2DNodeAddPaddingH(builder, paddingH): + builder.PrependInt32Slot(5, paddingH, 0) + +def AddPaddingH(builder, paddingH): + ConvTranspose2DNodeAddPaddingH(builder, paddingH) + +def ConvTranspose2DNodeAddPaddingW(builder, paddingW): + builder.PrependInt32Slot(6, paddingW, 0) + +def AddPaddingW(builder, paddingW): + ConvTranspose2DNodeAddPaddingW(builder, paddingW) + +def ConvTranspose2DNodeAddDilationH(builder, dilationH): + builder.PrependInt32Slot(7, dilationH, 1) + +def AddDilationH(builder, dilationH): + ConvTranspose2DNodeAddDilationH(builder, dilationH) + +def ConvTranspose2DNodeAddDilationW(builder, dilationW): + builder.PrependInt32Slot(8, dilationW, 1) + +def AddDilationW(builder, dilationW): + ConvTranspose2DNodeAddDilationW(builder, dilationW) + +def ConvTranspose2DNodeAddOutputPaddingH(builder, outputPaddingH): + builder.PrependInt32Slot(9, outputPaddingH, 0) + +def AddOutputPaddingH(builder, outputPaddingH): + ConvTranspose2DNodeAddOutputPaddingH(builder, outputPaddingH) + +def ConvTranspose2DNodeAddOutputPaddingW(builder, outputPaddingW): + builder.PrependInt32Slot(10, outputPaddingW, 0) + +def AddOutputPaddingW(builder, outputPaddingW): + ConvTranspose2DNodeAddOutputPaddingW(builder, outputPaddingW) + +def ConvTranspose2DNodeAddGroups(builder, groups): + builder.PrependInt32Slot(11, groups, 1) + +def AddGroups(builder, groups): + ConvTranspose2DNodeAddGroups(builder, groups) + +def ConvTranspose2DNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ConvTranspose2DNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose3DNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose3DNode.py new file mode 100644 index 00000000000..21d1ae0b8df --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ConvTranspose3DNode.py @@ -0,0 +1,257 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ConvTranspose3DNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ConvTranspose3DNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsConvTranspose3DNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ConvTranspose3DNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ConvTranspose3DNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose3DNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose3DNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ConvTranspose3DNode + def StrideD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose3DNode + def StrideH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose3DNode + def StrideW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose3DNode + def PaddingD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose3DNode + def PaddingH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose3DNode + def PaddingW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose3DNode + def DilationD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose3DNode + def DilationH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose3DNode + def DilationW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + + # ConvTranspose3DNode + def OutputPaddingD(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose3DNode + def OutputPaddingH(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose3DNode + def OutputPaddingW(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ConvTranspose3DNode + def Groups(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(34)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def ConvTranspose3DNodeStart(builder): + builder.StartObject(16) + +def Start(builder): + ConvTranspose3DNodeStart(builder) + +def ConvTranspose3DNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ConvTranspose3DNodeAddX(builder, x) + +def ConvTranspose3DNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + ConvTranspose3DNodeAddW(builder, w) + +def ConvTranspose3DNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ConvTranspose3DNodeAddOut(builder, out) + +def ConvTranspose3DNodeAddStrideD(builder, strideD): + builder.PrependInt32Slot(3, strideD, 1) + +def AddStrideD(builder, strideD): + ConvTranspose3DNodeAddStrideD(builder, strideD) + +def ConvTranspose3DNodeAddStrideH(builder, strideH): + builder.PrependInt32Slot(4, strideH, 1) + +def AddStrideH(builder, strideH): + ConvTranspose3DNodeAddStrideH(builder, strideH) + +def ConvTranspose3DNodeAddStrideW(builder, strideW): + builder.PrependInt32Slot(5, strideW, 1) + +def AddStrideW(builder, strideW): + ConvTranspose3DNodeAddStrideW(builder, strideW) + +def ConvTranspose3DNodeAddPaddingD(builder, paddingD): + builder.PrependInt32Slot(6, paddingD, 0) + +def AddPaddingD(builder, paddingD): + ConvTranspose3DNodeAddPaddingD(builder, paddingD) + +def ConvTranspose3DNodeAddPaddingH(builder, paddingH): + builder.PrependInt32Slot(7, paddingH, 0) + +def AddPaddingH(builder, paddingH): + ConvTranspose3DNodeAddPaddingH(builder, paddingH) + +def ConvTranspose3DNodeAddPaddingW(builder, paddingW): + builder.PrependInt32Slot(8, paddingW, 0) + +def AddPaddingW(builder, paddingW): + ConvTranspose3DNodeAddPaddingW(builder, paddingW) + +def ConvTranspose3DNodeAddDilationD(builder, dilationD): + builder.PrependInt32Slot(9, dilationD, 1) + +def AddDilationD(builder, dilationD): + ConvTranspose3DNodeAddDilationD(builder, dilationD) + +def ConvTranspose3DNodeAddDilationH(builder, dilationH): + builder.PrependInt32Slot(10, dilationH, 1) + +def AddDilationH(builder, dilationH): + ConvTranspose3DNodeAddDilationH(builder, dilationH) + +def ConvTranspose3DNodeAddDilationW(builder, dilationW): + builder.PrependInt32Slot(11, dilationW, 1) + +def AddDilationW(builder, dilationW): + ConvTranspose3DNodeAddDilationW(builder, dilationW) + +def ConvTranspose3DNodeAddOutputPaddingD(builder, outputPaddingD): + builder.PrependInt32Slot(12, outputPaddingD, 0) + +def AddOutputPaddingD(builder, outputPaddingD): + ConvTranspose3DNodeAddOutputPaddingD(builder, outputPaddingD) + +def ConvTranspose3DNodeAddOutputPaddingH(builder, outputPaddingH): + builder.PrependInt32Slot(13, outputPaddingH, 0) + +def AddOutputPaddingH(builder, outputPaddingH): + ConvTranspose3DNodeAddOutputPaddingH(builder, outputPaddingH) + +def ConvTranspose3DNodeAddOutputPaddingW(builder, outputPaddingW): + builder.PrependInt32Slot(14, outputPaddingW, 0) + +def AddOutputPaddingW(builder, outputPaddingW): + ConvTranspose3DNodeAddOutputPaddingW(builder, outputPaddingW) + +def ConvTranspose3DNodeAddGroups(builder, groups): + builder.PrependInt32Slot(15, groups, 1) + +def AddGroups(builder, groups): + ConvTranspose3DNodeAddGroups(builder, groups) + +def ConvTranspose3DNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ConvTranspose3DNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/CosNode.py b/backends/mlx/serialization/_generated/mlx_delegate/CosNode.py new file mode 100644 index 00000000000..5d29ac79c79 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/CosNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class CosNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = CosNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsCosNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # CosNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # CosNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # CosNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def CosNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + CosNodeStart(builder) + +def CosNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + CosNodeAddX(builder, x) + +def CosNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + CosNodeAddOut(builder, out) + +def CosNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return CosNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/CoshNode.py b/backends/mlx/serialization/_generated/mlx_delegate/CoshNode.py new file mode 100644 index 00000000000..54b74dd0018 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/CoshNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class CoshNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = CoshNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsCoshNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # CoshNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # CoshNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # CoshNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def CoshNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + CoshNodeStart(builder) + +def CoshNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + CoshNodeAddX(builder, x) + +def CoshNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + CoshNodeAddOut(builder, out) + +def CoshNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return CoshNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/CumsumNode.py b/backends/mlx/serialization/_generated/mlx_delegate/CumsumNode.py new file mode 100644 index 00000000000..8d29deda1a3 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/CumsumNode.py @@ -0,0 +1,110 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class CumsumNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = CumsumNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsCumsumNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # CumsumNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # CumsumNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # CumsumNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # CumsumNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # CumsumNode + def Reverse(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + + # CumsumNode + def Inclusive(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return True + +def CumsumNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + CumsumNodeStart(builder) + +def CumsumNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + CumsumNodeAddX(builder, x) + +def CumsumNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + CumsumNodeAddOut(builder, out) + +def CumsumNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + CumsumNodeAddAxis(builder, axis) + +def CumsumNodeAddReverse(builder, reverse): + builder.PrependBoolSlot(3, reverse, 0) + +def AddReverse(builder, reverse): + CumsumNodeAddReverse(builder, reverse) + +def CumsumNodeAddInclusive(builder, inclusive): + builder.PrependBoolSlot(4, inclusive, 1) + +def AddInclusive(builder, inclusive): + CumsumNodeAddInclusive(builder, inclusive) + +def CumsumNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return CumsumNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/DequantizeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/DequantizeNode.py new file mode 100644 index 00000000000..9c53d71bc5e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/DequantizeNode.py @@ -0,0 +1,174 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class DequantizeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = DequantizeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsDequantizeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # DequantizeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # DequantizeNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DequantizeNode + def Scales(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DequantizeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DequantizeNode + def Biases(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DequantizeNode + def GroupSize(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # DequantizeNode + def Bits(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # DequantizeNode + def Mode(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # DequantizeNode + def GlobalScale(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DequantizeNode + def Dtype(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return None + +def DequantizeNodeStart(builder): + builder.StartObject(9) + +def Start(builder): + DequantizeNodeStart(builder) + +def DequantizeNodeAddW(builder, w): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + DequantizeNodeAddW(builder, w) + +def DequantizeNodeAddScales(builder, scales): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(scales), 0) + +def AddScales(builder, scales): + DequantizeNodeAddScales(builder, scales) + +def DequantizeNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + DequantizeNodeAddOut(builder, out) + +def DequantizeNodeAddBiases(builder, biases): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(biases), 0) + +def AddBiases(builder, biases): + DequantizeNodeAddBiases(builder, biases) + +def DequantizeNodeAddGroupSize(builder, groupSize): + builder.PrependInt32Slot(4, groupSize, 0) + +def AddGroupSize(builder, groupSize): + DequantizeNodeAddGroupSize(builder, groupSize) + +def DequantizeNodeAddBits(builder, bits): + builder.PrependInt32Slot(5, bits, 0) + +def AddBits(builder, bits): + DequantizeNodeAddBits(builder, bits) + +def DequantizeNodeAddMode(builder, mode): + builder.PrependUOffsetTRelativeSlot(6, flatbuffers.number_types.UOffsetTFlags.py_type(mode), 0) + +def AddMode(builder, mode): + DequantizeNodeAddMode(builder, mode) + +def DequantizeNodeAddGlobalScale(builder, globalScale): + builder.PrependStructSlot(7, flatbuffers.number_types.UOffsetTFlags.py_type(globalScale), 0) + +def AddGlobalScale(builder, globalScale): + DequantizeNodeAddGlobalScale(builder, globalScale) + +def DequantizeNodeAddDtype(builder, dtype): + builder.PrependInt8Slot(8, dtype, None) + +def AddDtype(builder, dtype): + DequantizeNodeAddDtype(builder, dtype) + +def DequantizeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return DequantizeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/DivideNode.py b/backends/mlx/serialization/_generated/mlx_delegate/DivideNode.py new file mode 100644 index 00000000000..54efa9e79b0 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/DivideNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class DivideNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = DivideNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsDivideNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # DivideNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # DivideNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DivideNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # DivideNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def DivideNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + DivideNodeStart(builder) + +def DivideNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + DivideNodeAddA(builder, a) + +def DivideNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + DivideNodeAddB(builder, b) + +def DivideNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + DivideNodeAddOut(builder, out) + +def DivideNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return DivideNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/EqualNode.py b/backends/mlx/serialization/_generated/mlx_delegate/EqualNode.py new file mode 100644 index 00000000000..a1f941b5f90 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/EqualNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class EqualNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = EqualNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsEqualNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # EqualNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # EqualNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # EqualNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # EqualNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def EqualNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + EqualNodeStart(builder) + +def EqualNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + EqualNodeAddA(builder, a) + +def EqualNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + EqualNodeAddB(builder, b) + +def EqualNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + EqualNodeAddOut(builder, out) + +def EqualNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return EqualNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ErfNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ErfNode.py new file mode 100644 index 00000000000..531423bc1dd --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ErfNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ErfNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ErfNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsErfNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ErfNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ErfNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ErfNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ErfNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ErfNodeStart(builder) + +def ErfNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ErfNodeAddX(builder, x) + +def ErfNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ErfNodeAddOut(builder, out) + +def ErfNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ErfNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ExpNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ExpNode.py new file mode 100644 index 00000000000..378c98507be --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ExpNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ExpNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ExpNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsExpNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ExpNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ExpNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ExpNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ExpNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ExpNodeStart(builder) + +def ExpNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ExpNodeAddX(builder, x) + +def ExpNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ExpNodeAddOut(builder, out) + +def ExpNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ExpNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ExpandDimsNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ExpandDimsNode.py new file mode 100644 index 00000000000..b10e05e8077 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ExpandDimsNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ExpandDimsNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ExpandDimsNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsExpandDimsNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ExpandDimsNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ExpandDimsNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ExpandDimsNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ExpandDimsNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def ExpandDimsNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + ExpandDimsNodeStart(builder) + +def ExpandDimsNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ExpandDimsNodeAddX(builder, x) + +def ExpandDimsNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ExpandDimsNodeAddOut(builder, out) + +def ExpandDimsNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + ExpandDimsNodeAddAxis(builder, axis) + +def ExpandDimsNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ExpandDimsNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Expm1Node.py b/backends/mlx/serialization/_generated/mlx_delegate/Expm1Node.py new file mode 100644 index 00000000000..fa6590845ce --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Expm1Node.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Expm1Node(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Expm1Node() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsExpm1Node(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Expm1Node + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Expm1Node + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Expm1Node + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def Expm1NodeStart(builder): + builder.StartObject(2) + +def Start(builder): + Expm1NodeStart(builder) + +def Expm1NodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Expm1NodeAddX(builder, x) + +def Expm1NodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Expm1NodeAddOut(builder, out) + +def Expm1NodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Expm1NodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/FloatOrVid.py b/backends/mlx/serialization/_generated/mlx_delegate/FloatOrVid.py new file mode 100644 index 00000000000..78d90bc0b8e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/FloatOrVid.py @@ -0,0 +1,80 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class FloatOrVid(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = FloatOrVid() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsFloatOrVid(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # FloatOrVid + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # FloatOrVid + def Literal(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float64Flags, o + self._tab.Pos) + return 0.0 + + # FloatOrVid + def Vid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FloatOrVid + def IsVid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def FloatOrVidStart(builder): + builder.StartObject(3) + +def Start(builder): + FloatOrVidStart(builder) + +def FloatOrVidAddLiteral(builder, literal): + builder.PrependFloat64Slot(0, literal, 0.0) + +def AddLiteral(builder, literal): + FloatOrVidAddLiteral(builder, literal) + +def FloatOrVidAddVid(builder, vid): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(vid), 0) + +def AddVid(builder, vid): + FloatOrVidAddVid(builder, vid) + +def FloatOrVidAddIsVid(builder, isVid): + builder.PrependBoolSlot(2, isVid, 0) + +def AddIsVid(builder, isVid): + FloatOrVidAddIsVid(builder, isVid) + +def FloatOrVidEnd(builder): + return builder.EndObject() + +def End(builder): + return FloatOrVidEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/FloorDivideIntNode.py b/backends/mlx/serialization/_generated/mlx_delegate/FloorDivideIntNode.py new file mode 100644 index 00000000000..e73ecc83488 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/FloorDivideIntNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class FloorDivideIntNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = FloorDivideIntNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsFloorDivideIntNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # FloorDivideIntNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # FloorDivideIntNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FloorDivideIntNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FloorDivideIntNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def FloorDivideIntNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + FloorDivideIntNodeStart(builder) + +def FloorDivideIntNodeAddA(builder, a): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + FloorDivideIntNodeAddA(builder, a) + +def FloorDivideIntNodeAddB(builder, b): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + FloorDivideIntNodeAddB(builder, b) + +def FloorDivideIntNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + FloorDivideIntNodeAddOut(builder, out) + +def FloorDivideIntNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return FloorDivideIntNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/FloorDivideNode.py b/backends/mlx/serialization/_generated/mlx_delegate/FloorDivideNode.py new file mode 100644 index 00000000000..aec86dc8e3a --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/FloorDivideNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class FloorDivideNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = FloorDivideNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsFloorDivideNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # FloorDivideNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # FloorDivideNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FloorDivideNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FloorDivideNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def FloorDivideNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + FloorDivideNodeStart(builder) + +def FloorDivideNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + FloorDivideNodeAddA(builder, a) + +def FloorDivideNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + FloorDivideNodeAddB(builder, b) + +def FloorDivideNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + FloorDivideNodeAddOut(builder, out) + +def FloorDivideNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return FloorDivideNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/FloorNode.py b/backends/mlx/serialization/_generated/mlx_delegate/FloorNode.py new file mode 100644 index 00000000000..d92b6a48151 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/FloorNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class FloorNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = FloorNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsFloorNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # FloorNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # FloorNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FloorNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def FloorNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + FloorNodeStart(builder) + +def FloorNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + FloorNodeAddX(builder, x) + +def FloorNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + FloorNodeAddOut(builder, out) + +def FloorNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return FloorNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/FullLikeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/FullLikeNode.py new file mode 100644 index 00000000000..252b8ee686c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/FullLikeNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class FullLikeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = FullLikeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsFullLikeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # FullLikeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # FullLikeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FullLikeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FullLikeNode + def V(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.FloatOrVid import FloatOrVid + obj = FloatOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FullLikeNode + def ScalarType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return None + +def FullLikeNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + FullLikeNodeStart(builder) + +def FullLikeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + FullLikeNodeAddX(builder, x) + +def FullLikeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + FullLikeNodeAddOut(builder, out) + +def FullLikeNodeAddV(builder, v): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(v), 0) + +def AddV(builder, v): + FullLikeNodeAddV(builder, v) + +def FullLikeNodeAddScalarType(builder, scalarType): + builder.PrependInt8Slot(3, scalarType, None) + +def AddScalarType(builder, scalarType): + FullLikeNodeAddScalarType(builder, scalarType) + +def FullLikeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return FullLikeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/FullNode.py b/backends/mlx/serialization/_generated/mlx_delegate/FullNode.py new file mode 100644 index 00000000000..bdf9eaa838e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/FullNode.py @@ -0,0 +1,121 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class FullNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = FullNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsFullNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # FullNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # FullNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FullNode + def Shape(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FullNode + def ShapeLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # FullNode + def ShapeIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + return o == 0 + + # FullNode + def V(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.FloatOrVid import FloatOrVid + obj = FloatOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # FullNode + def ScalarType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return 0 + +def FullNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + FullNodeStart(builder) + +def FullNodeAddOut(builder, out): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + FullNodeAddOut(builder, out) + +def FullNodeAddShape(builder, shape): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(shape), 0) + +def AddShape(builder, shape): + FullNodeAddShape(builder, shape) + +def FullNodeStartShapeVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShapeVector(builder, numElems): + return FullNodeStartShapeVector(builder, numElems) + +def FullNodeAddV(builder, v): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(v), 0) + +def AddV(builder, v): + FullNodeAddV(builder, v) + +def FullNodeAddScalarType(builder, scalarType): + builder.PrependInt8Slot(3, scalarType, 0) + +def AddScalarType(builder, scalarType): + FullNodeAddScalarType(builder, scalarType) + +def FullNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return FullNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/GatherMmNode.py b/backends/mlx/serialization/_generated/mlx_delegate/GatherMmNode.py new file mode 100644 index 00000000000..169e49872ab --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/GatherMmNode.py @@ -0,0 +1,135 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class GatherMmNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = GatherMmNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsGatherMmNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # GatherMmNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # GatherMmNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherMmNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherMmNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherMmNode + def LhsIndices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherMmNode + def RhsIndices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherMmNode + def SortedIndices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def GatherMmNodeStart(builder): + builder.StartObject(6) + +def Start(builder): + GatherMmNodeStart(builder) + +def GatherMmNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + GatherMmNodeAddA(builder, a) + +def GatherMmNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + GatherMmNodeAddB(builder, b) + +def GatherMmNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + GatherMmNodeAddOut(builder, out) + +def GatherMmNodeAddLhsIndices(builder, lhsIndices): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(lhsIndices), 0) + +def AddLhsIndices(builder, lhsIndices): + GatherMmNodeAddLhsIndices(builder, lhsIndices) + +def GatherMmNodeAddRhsIndices(builder, rhsIndices): + builder.PrependStructSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(rhsIndices), 0) + +def AddRhsIndices(builder, rhsIndices): + GatherMmNodeAddRhsIndices(builder, rhsIndices) + +def GatherMmNodeAddSortedIndices(builder, sortedIndices): + builder.PrependBoolSlot(5, sortedIndices, 0) + +def AddSortedIndices(builder, sortedIndices): + GatherMmNodeAddSortedIndices(builder, sortedIndices) + +def GatherMmNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return GatherMmNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/GatherNode.py b/backends/mlx/serialization/_generated/mlx_delegate/GatherNode.py new file mode 100644 index 00000000000..9fba7dbd4b4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/GatherNode.py @@ -0,0 +1,185 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class GatherNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = GatherNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsGatherNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # GatherNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # GatherNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherNode + def Indices(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherNode + def IndicesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # GatherNode + def IndicesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + return o == 0 + + # GatherNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # GatherNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # GatherNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # GatherNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + return o == 0 + + # GatherNode + def SliceSizes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # GatherNode + def SliceSizesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # GatherNode + def SliceSizesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # GatherNode + def SliceSizesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + return o == 0 + +def GatherNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + GatherNodeStart(builder) + +def GatherNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + GatherNodeAddX(builder, x) + +def GatherNodeAddIndices(builder, indices): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(indices), 0) + +def AddIndices(builder, indices): + GatherNodeAddIndices(builder, indices) + +def GatherNodeStartIndicesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartIndicesVector(builder, numElems): + return GatherNodeStartIndicesVector(builder, numElems) + +def GatherNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + GatherNodeAddOut(builder, out) + +def GatherNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + GatherNodeAddAxes(builder, axes) + +def GatherNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return GatherNodeStartAxesVector(builder, numElems) + +def GatherNodeAddSliceSizes(builder, sliceSizes): + builder.PrependUOffsetTRelativeSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(sliceSizes), 0) + +def AddSliceSizes(builder, sliceSizes): + GatherNodeAddSliceSizes(builder, sliceSizes) + +def GatherNodeStartSliceSizesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartSliceSizesVector(builder, numElems): + return GatherNodeStartSliceSizesVector(builder, numElems) + +def GatherNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return GatherNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/GatherQmmNode.py b/backends/mlx/serialization/_generated/mlx_delegate/GatherQmmNode.py new file mode 100644 index 00000000000..dc2be5d2174 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/GatherQmmNode.py @@ -0,0 +1,221 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class GatherQmmNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = GatherQmmNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsGatherQmmNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # GatherQmmNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # GatherQmmNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def Scales(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def Mode(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # GatherQmmNode + def Biases(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def LhsIndices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def RhsIndices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GatherQmmNode + def Transpose(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return True + + # GatherQmmNode + def GroupSize(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # GatherQmmNode + def Bits(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # GatherQmmNode + def SortedIndices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def GatherQmmNodeStart(builder): + builder.StartObject(12) + +def Start(builder): + GatherQmmNodeStart(builder) + +def GatherQmmNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + GatherQmmNodeAddX(builder, x) + +def GatherQmmNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + GatherQmmNodeAddW(builder, w) + +def GatherQmmNodeAddScales(builder, scales): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(scales), 0) + +def AddScales(builder, scales): + GatherQmmNodeAddScales(builder, scales) + +def GatherQmmNodeAddOut(builder, out): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + GatherQmmNodeAddOut(builder, out) + +def GatherQmmNodeAddMode(builder, mode): + builder.PrependUOffsetTRelativeSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(mode), 0) + +def AddMode(builder, mode): + GatherQmmNodeAddMode(builder, mode) + +def GatherQmmNodeAddBiases(builder, biases): + builder.PrependStructSlot(5, flatbuffers.number_types.UOffsetTFlags.py_type(biases), 0) + +def AddBiases(builder, biases): + GatherQmmNodeAddBiases(builder, biases) + +def GatherQmmNodeAddLhsIndices(builder, lhsIndices): + builder.PrependStructSlot(6, flatbuffers.number_types.UOffsetTFlags.py_type(lhsIndices), 0) + +def AddLhsIndices(builder, lhsIndices): + GatherQmmNodeAddLhsIndices(builder, lhsIndices) + +def GatherQmmNodeAddRhsIndices(builder, rhsIndices): + builder.PrependStructSlot(7, flatbuffers.number_types.UOffsetTFlags.py_type(rhsIndices), 0) + +def AddRhsIndices(builder, rhsIndices): + GatherQmmNodeAddRhsIndices(builder, rhsIndices) + +def GatherQmmNodeAddTranspose(builder, transpose): + builder.PrependBoolSlot(8, transpose, 1) + +def AddTranspose(builder, transpose): + GatherQmmNodeAddTranspose(builder, transpose) + +def GatherQmmNodeAddGroupSize(builder, groupSize): + builder.PrependInt32Slot(9, groupSize, 0) + +def AddGroupSize(builder, groupSize): + GatherQmmNodeAddGroupSize(builder, groupSize) + +def GatherQmmNodeAddBits(builder, bits): + builder.PrependInt32Slot(10, bits, 0) + +def AddBits(builder, bits): + GatherQmmNodeAddBits(builder, bits) + +def GatherQmmNodeAddSortedIndices(builder, sortedIndices): + builder.PrependBoolSlot(11, sortedIndices, 0) + +def AddSortedIndices(builder, sortedIndices): + GatherQmmNodeAddSortedIndices(builder, sortedIndices) + +def GatherQmmNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return GatherQmmNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/GeluNode.py b/backends/mlx/serialization/_generated/mlx_delegate/GeluNode.py new file mode 100644 index 00000000000..8af22a5a1b4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/GeluNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class GeluNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = GeluNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsGeluNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # GeluNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # GeluNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GeluNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GeluNode + def Approximate(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + +def GeluNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + GeluNodeStart(builder) + +def GeluNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + GeluNodeAddX(builder, x) + +def GeluNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + GeluNodeAddOut(builder, out) + +def GeluNodeAddApproximate(builder, approximate): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(approximate), 0) + +def AddApproximate(builder, approximate): + GeluNodeAddApproximate(builder, approximate) + +def GeluNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return GeluNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/GreaterEqualNode.py b/backends/mlx/serialization/_generated/mlx_delegate/GreaterEqualNode.py new file mode 100644 index 00000000000..bf7682d5610 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/GreaterEqualNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class GreaterEqualNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = GreaterEqualNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsGreaterEqualNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # GreaterEqualNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # GreaterEqualNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GreaterEqualNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GreaterEqualNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def GreaterEqualNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + GreaterEqualNodeStart(builder) + +def GreaterEqualNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + GreaterEqualNodeAddA(builder, a) + +def GreaterEqualNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + GreaterEqualNodeAddB(builder, b) + +def GreaterEqualNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + GreaterEqualNodeAddOut(builder, out) + +def GreaterEqualNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return GreaterEqualNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/GreaterNode.py b/backends/mlx/serialization/_generated/mlx_delegate/GreaterNode.py new file mode 100644 index 00000000000..35105c267c1 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/GreaterNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class GreaterNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = GreaterNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsGreaterNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # GreaterNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # GreaterNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GreaterNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # GreaterNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def GreaterNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + GreaterNodeStart(builder) + +def GreaterNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + GreaterNodeAddA(builder, a) + +def GreaterNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + GreaterNodeAddB(builder, b) + +def GreaterNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + GreaterNodeAddOut(builder, out) + +def GreaterNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return GreaterNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/IdCopyNode.py b/backends/mlx/serialization/_generated/mlx_delegate/IdCopyNode.py new file mode 100644 index 00000000000..2b289b57772 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/IdCopyNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class IdCopyNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = IdCopyNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsIdCopyNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # IdCopyNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # IdCopyNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IdCopyNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def IdCopyNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + IdCopyNodeStart(builder) + +def IdCopyNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + IdCopyNodeAddX(builder, x) + +def IdCopyNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + IdCopyNodeAddOut(builder, out) + +def IdCopyNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return IdCopyNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/IfNode.py b/backends/mlx/serialization/_generated/mlx_delegate/IfNode.py new file mode 100644 index 00000000000..1d2898cfc20 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/IfNode.py @@ -0,0 +1,80 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class IfNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = IfNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsIfNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # IfNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # IfNode + def Cond(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IfNode + def ThenChainIdx(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # IfNode + def ElseChainIdx(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + +def IfNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + IfNodeStart(builder) + +def IfNodeAddCond(builder, cond): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(cond), 0) + +def AddCond(builder, cond): + IfNodeAddCond(builder, cond) + +def IfNodeAddThenChainIdx(builder, thenChainIdx): + builder.PrependUint32Slot(1, thenChainIdx, 0) + +def AddThenChainIdx(builder, thenChainIdx): + IfNodeAddThenChainIdx(builder, thenChainIdx) + +def IfNodeAddElseChainIdx(builder, elseChainIdx): + builder.PrependUint32Slot(2, elseChainIdx, 0) + +def AddElseChainIdx(builder, elseChainIdx): + IfNodeAddElseChainIdx(builder, elseChainIdx) + +def IfNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return IfNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/IndexCopyNode.py b/backends/mlx/serialization/_generated/mlx_delegate/IndexCopyNode.py new file mode 100644 index 00000000000..6c6394786cb --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/IndexCopyNode.py @@ -0,0 +1,118 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class IndexCopyNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = IndexCopyNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsIndexCopyNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # IndexCopyNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # IndexCopyNode + def Dst(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IndexCopyNode + def Update(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IndexCopyNode + def Indices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IndexCopyNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IndexCopyNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def IndexCopyNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + IndexCopyNodeStart(builder) + +def IndexCopyNodeAddDst(builder, dst): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(dst), 0) + +def AddDst(builder, dst): + IndexCopyNodeAddDst(builder, dst) + +def IndexCopyNodeAddUpdate(builder, update): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(update), 0) + +def AddUpdate(builder, update): + IndexCopyNodeAddUpdate(builder, update) + +def IndexCopyNodeAddIndices(builder, indices): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(indices), 0) + +def AddIndices(builder, indices): + IndexCopyNodeAddIndices(builder, indices) + +def IndexCopyNodeAddOut(builder, out): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + IndexCopyNodeAddOut(builder, out) + +def IndexCopyNodeAddAxis(builder, axis): + builder.PrependInt32Slot(4, axis, 0) + +def AddAxis(builder, axis): + IndexCopyNodeAddAxis(builder, axis) + +def IndexCopyNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return IndexCopyNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Instruction.py b/backends/mlx/serialization/_generated/mlx_delegate/Instruction.py new file mode 100644 index 00000000000..7c6b9e0f08f --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Instruction.py @@ -0,0 +1,66 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Instruction(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Instruction() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsInstruction(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Instruction + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Instruction + def OpType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint8Flags, o + self._tab.Pos) + return 0 + + # Instruction + def Op(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + from flatbuffers.table import Table + obj = Table(bytearray(), 0) + self._tab.Union(obj, o) + return obj + return None + +def InstructionStart(builder): + builder.StartObject(2) + +def Start(builder): + InstructionStart(builder) + +def InstructionAddOpType(builder, opType): + builder.PrependUint8Slot(0, opType, 0) + +def AddOpType(builder, opType): + InstructionAddOpType(builder, opType) + +def InstructionAddOp(builder, op): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(op), 0) + +def AddOp(builder, op): + InstructionAddOp(builder, op) + +def InstructionEnd(builder): + return builder.EndObject() + +def End(builder): + return InstructionEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/InstructionChain.py b/backends/mlx/serialization/_generated/mlx_delegate/InstructionChain.py new file mode 100644 index 00000000000..d5f021bbc33 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/InstructionChain.py @@ -0,0 +1,74 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class InstructionChain(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = InstructionChain() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsInstructionChain(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # InstructionChain + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # InstructionChain + def Instructions(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.Instruction import Instruction + obj = Instruction() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # InstructionChain + def InstructionsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # InstructionChain + def InstructionsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + return o == 0 + +def InstructionChainStart(builder): + builder.StartObject(1) + +def Start(builder): + InstructionChainStart(builder) + +def InstructionChainAddInstructions(builder, instructions): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(instructions), 0) + +def AddInstructions(builder, instructions): + InstructionChainAddInstructions(builder, instructions) + +def InstructionChainStartInstructionsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartInstructionsVector(builder, numElems): + return InstructionChainStartInstructionsVector(builder, numElems) + +def InstructionChainEnd(builder): + return builder.EndObject() + +def End(builder): + return InstructionChainEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/IntOrVid.py b/backends/mlx/serialization/_generated/mlx_delegate/IntOrVid.py new file mode 100644 index 00000000000..8b0719c8ec4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/IntOrVid.py @@ -0,0 +1,80 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class IntOrVid(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = IntOrVid() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsIntOrVid(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # IntOrVid + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # IntOrVid + def Literal(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int64Flags, o + self._tab.Pos) + return 0 + + # IntOrVid + def Vid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IntOrVid + def IsVid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def IntOrVidStart(builder): + builder.StartObject(3) + +def Start(builder): + IntOrVidStart(builder) + +def IntOrVidAddLiteral(builder, literal): + builder.PrependInt64Slot(0, literal, 0) + +def AddLiteral(builder, literal): + IntOrVidAddLiteral(builder, literal) + +def IntOrVidAddVid(builder, vid): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(vid), 0) + +def AddVid(builder, vid): + IntOrVidAddVid(builder, vid) + +def IntOrVidAddIsVid(builder, isVid): + builder.PrependBoolSlot(2, isVid, 0) + +def AddIsVid(builder, isVid): + IntOrVidAddIsVid(builder, isVid) + +def IntOrVidEnd(builder): + return builder.EndObject() + +def End(builder): + return IntOrVidEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/IntOrVidOrTid.py b/backends/mlx/serialization/_generated/mlx_delegate/IntOrVidOrTid.py new file mode 100644 index 00000000000..770f0ac4daf --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/IntOrVidOrTid.py @@ -0,0 +1,97 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class IntOrVidOrTid(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = IntOrVidOrTid() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsIntOrVidOrTid(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # IntOrVidOrTid + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # IntOrVidOrTid + def Literal(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int64Flags, o + self._tab.Pos) + return 0 + + # IntOrVidOrTid + def Vid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IntOrVidOrTid + def Tid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # IntOrVidOrTid + def Kind(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint8Flags, o + self._tab.Pos) + return 0 + +def IntOrVidOrTidStart(builder): + builder.StartObject(4) + +def Start(builder): + IntOrVidOrTidStart(builder) + +def IntOrVidOrTidAddLiteral(builder, literal): + builder.PrependInt64Slot(0, literal, 0) + +def AddLiteral(builder, literal): + IntOrVidOrTidAddLiteral(builder, literal) + +def IntOrVidOrTidAddVid(builder, vid): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(vid), 0) + +def AddVid(builder, vid): + IntOrVidOrTidAddVid(builder, vid) + +def IntOrVidOrTidAddTid(builder, tid): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(tid), 0) + +def AddTid(builder, tid): + IntOrVidOrTidAddTid(builder, tid) + +def IntOrVidOrTidAddKind(builder, kind): + builder.PrependUint8Slot(3, kind, 0) + +def AddKind(builder, kind): + IntOrVidOrTidAddKind(builder, kind) + +def IntOrVidOrTidEnd(builder): + return builder.EndObject() + +def End(builder): + return IntOrVidOrTidEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ItemIntNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ItemIntNode.py new file mode 100644 index 00000000000..76873ae7d9c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ItemIntNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ItemIntNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ItemIntNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsItemIntNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ItemIntNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ItemIntNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ItemIntNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ItemIntNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ItemIntNodeStart(builder) + +def ItemIntNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ItemIntNodeAddX(builder, x) + +def ItemIntNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ItemIntNodeAddOut(builder, out) + +def ItemIntNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ItemIntNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LayerNormNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LayerNormNode.py new file mode 100644 index 00000000000..472cce91f11 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LayerNormNode.py @@ -0,0 +1,118 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LayerNormNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LayerNormNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLayerNormNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LayerNormNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LayerNormNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LayerNormNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LayerNormNode + def Weight(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LayerNormNode + def Bias(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LayerNormNode + def Eps(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 0.0 + +def LayerNormNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + LayerNormNodeStart(builder) + +def LayerNormNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + LayerNormNodeAddX(builder, x) + +def LayerNormNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LayerNormNodeAddOut(builder, out) + +def LayerNormNodeAddWeight(builder, weight): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(weight), 0) + +def AddWeight(builder, weight): + LayerNormNodeAddWeight(builder, weight) + +def LayerNormNodeAddBias(builder, bias): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(bias), 0) + +def AddBias(builder, bias): + LayerNormNodeAddBias(builder, bias) + +def LayerNormNodeAddEps(builder, eps): + builder.PrependFloat32Slot(4, eps, 0.0) + +def AddEps(builder, eps): + LayerNormNodeAddEps(builder, eps) + +def LayerNormNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LayerNormNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LessEqualNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LessEqualNode.py new file mode 100644 index 00000000000..63f9a9625c8 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LessEqualNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LessEqualNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LessEqualNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLessEqualNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LessEqualNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LessEqualNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LessEqualNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LessEqualNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LessEqualNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + LessEqualNodeStart(builder) + +def LessEqualNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + LessEqualNodeAddA(builder, a) + +def LessEqualNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + LessEqualNodeAddB(builder, b) + +def LessEqualNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LessEqualNodeAddOut(builder, out) + +def LessEqualNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LessEqualNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LessNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LessNode.py new file mode 100644 index 00000000000..e6c8dce524c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LessNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LessNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LessNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLessNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LessNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LessNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LessNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LessNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LessNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + LessNodeStart(builder) + +def LessNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + LessNodeAddA(builder, a) + +def LessNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + LessNodeAddB(builder, b) + +def LessNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LessNodeAddOut(builder, out) + +def LessNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LessNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Log10Node.py b/backends/mlx/serialization/_generated/mlx_delegate/Log10Node.py new file mode 100644 index 00000000000..63c6ef57dde --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Log10Node.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Log10Node(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Log10Node() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLog10Node(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Log10Node + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Log10Node + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Log10Node + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def Log10NodeStart(builder): + builder.StartObject(2) + +def Start(builder): + Log10NodeStart(builder) + +def Log10NodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Log10NodeAddX(builder, x) + +def Log10NodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Log10NodeAddOut(builder, out) + +def Log10NodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Log10NodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Log1pNode.py b/backends/mlx/serialization/_generated/mlx_delegate/Log1pNode.py new file mode 100644 index 00000000000..3876bb74db9 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Log1pNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Log1pNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Log1pNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLog1pNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Log1pNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Log1pNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Log1pNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def Log1pNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + Log1pNodeStart(builder) + +def Log1pNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Log1pNodeAddX(builder, x) + +def Log1pNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Log1pNodeAddOut(builder, out) + +def Log1pNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Log1pNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Log2Node.py b/backends/mlx/serialization/_generated/mlx_delegate/Log2Node.py new file mode 100644 index 00000000000..84c245c2984 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Log2Node.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Log2Node(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = Log2Node() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLog2Node(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # Log2Node + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Log2Node + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # Log2Node + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def Log2NodeStart(builder): + builder.StartObject(2) + +def Start(builder): + Log2NodeStart(builder) + +def Log2NodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + Log2NodeAddX(builder, x) + +def Log2NodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + Log2NodeAddOut(builder, out) + +def Log2NodeEnd(builder): + return builder.EndObject() + +def End(builder): + return Log2NodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LogAddExpNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LogAddExpNode.py new file mode 100644 index 00000000000..eb91393548e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LogAddExpNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LogAddExpNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LogAddExpNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLogAddExpNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LogAddExpNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LogAddExpNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogAddExpNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogAddExpNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LogAddExpNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + LogAddExpNodeStart(builder) + +def LogAddExpNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + LogAddExpNodeAddA(builder, a) + +def LogAddExpNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + LogAddExpNodeAddB(builder, b) + +def LogAddExpNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LogAddExpNodeAddOut(builder, out) + +def LogAddExpNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LogAddExpNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LogNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LogNode.py new file mode 100644 index 00000000000..3f442ae49e4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LogNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LogNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LogNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLogNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LogNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LogNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LogNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + LogNodeStart(builder) + +def LogNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + LogNodeAddX(builder, x) + +def LogNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LogNodeAddOut(builder, out) + +def LogNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LogNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LogSumExpNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LogSumExpNode.py new file mode 100644 index 00000000000..2153a9837f4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LogSumExpNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LogSumExpNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LogSumExpNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLogSumExpNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LogSumExpNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LogSumExpNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogSumExpNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogSumExpNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # LogSumExpNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # LogSumExpNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # LogSumExpNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # LogSumExpNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def LogSumExpNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + LogSumExpNodeStart(builder) + +def LogSumExpNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + LogSumExpNodeAddX(builder, x) + +def LogSumExpNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LogSumExpNodeAddOut(builder, out) + +def LogSumExpNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + LogSumExpNodeAddAxes(builder, axes) + +def LogSumExpNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return LogSumExpNodeStartAxesVector(builder, numElems) + +def LogSumExpNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + LogSumExpNodeAddKeepdims(builder, keepdims) + +def LogSumExpNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LogSumExpNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LogicalAndNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LogicalAndNode.py new file mode 100644 index 00000000000..97c14595962 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LogicalAndNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LogicalAndNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LogicalAndNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLogicalAndNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LogicalAndNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LogicalAndNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogicalAndNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogicalAndNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LogicalAndNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + LogicalAndNodeStart(builder) + +def LogicalAndNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + LogicalAndNodeAddA(builder, a) + +def LogicalAndNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + LogicalAndNodeAddB(builder, b) + +def LogicalAndNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LogicalAndNodeAddOut(builder, out) + +def LogicalAndNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LogicalAndNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LogicalNotNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LogicalNotNode.py new file mode 100644 index 00000000000..671d36232c6 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LogicalNotNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LogicalNotNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LogicalNotNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLogicalNotNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LogicalNotNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LogicalNotNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogicalNotNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LogicalNotNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + LogicalNotNodeStart(builder) + +def LogicalNotNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + LogicalNotNodeAddX(builder, x) + +def LogicalNotNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LogicalNotNodeAddOut(builder, out) + +def LogicalNotNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LogicalNotNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/LogicalOrNode.py b/backends/mlx/serialization/_generated/mlx_delegate/LogicalOrNode.py new file mode 100644 index 00000000000..47d2084d1dc --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/LogicalOrNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class LogicalOrNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = LogicalOrNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsLogicalOrNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # LogicalOrNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # LogicalOrNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogicalOrNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # LogicalOrNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def LogicalOrNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + LogicalOrNodeStart(builder) + +def LogicalOrNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + LogicalOrNodeAddA(builder, a) + +def LogicalOrNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + LogicalOrNodeAddB(builder, b) + +def LogicalOrNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + LogicalOrNodeAddOut(builder, out) + +def LogicalOrNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return LogicalOrNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MLXGraph.py b/backends/mlx/serialization/_generated/mlx_delegate/MLXGraph.py new file mode 100644 index 00000000000..4b4acf85b0a --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MLXGraph.py @@ -0,0 +1,376 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MLXGraph(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MLXGraph() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMLXGraph(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MLXGraph + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MLXGraph + def Version(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # MLXGraph + def NumConstantTensors(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def NumInputTensors(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def NumOutputTensors(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def NumMutableBufferTensors(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def NumTempTensors(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def NumValues(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def InstructionChains(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.InstructionChain import InstructionChain + obj = InstructionChain() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MLXGraph + def InstructionChainsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MLXGraph + def InstructionChainsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + return o == 0 + + # MLXGraph + def MainChainIdx(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # MLXGraph + def InitChainIdx(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return -1 + + # MLXGraph + def InputMap(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.SlotVariant import SlotVariant + obj = SlotVariant() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MLXGraph + def InputMapLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MLXGraph + def InputMapIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + return o == 0 + + # MLXGraph + def OutputMap(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.SlotVariant import SlotVariant + obj = SlotVariant() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MLXGraph + def OutputMapLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MLXGraph + def OutputMapIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + return o == 0 + + # MLXGraph + def MutableBufferMap(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.SlotVariant import SlotVariant + obj = SlotVariant() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MLXGraph + def MutableBufferMapLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MLXGraph + def MutableBufferMapIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + return o == 0 + + # MLXGraph + def NamedSlots(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.NamedSlot import NamedSlot + obj = NamedSlot() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MLXGraph + def NamedSlotsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MLXGraph + def NamedSlotsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + return o == 0 + + # MLXGraph + def TensorMeta(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.TensorMeta import TensorMeta + obj = TensorMeta() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MLXGraph + def TensorMetaLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MLXGraph + def TensorMetaIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + return o == 0 + +def MLXGraphStart(builder): + builder.StartObject(15) + +def Start(builder): + MLXGraphStart(builder) + +def MLXGraphAddVersion(builder, version): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(version), 0) + +def AddVersion(builder, version): + MLXGraphAddVersion(builder, version) + +def MLXGraphAddNumConstantTensors(builder, numConstantTensors): + builder.PrependUint32Slot(1, numConstantTensors, 0) + +def AddNumConstantTensors(builder, numConstantTensors): + MLXGraphAddNumConstantTensors(builder, numConstantTensors) + +def MLXGraphAddNumInputTensors(builder, numInputTensors): + builder.PrependUint32Slot(2, numInputTensors, 0) + +def AddNumInputTensors(builder, numInputTensors): + MLXGraphAddNumInputTensors(builder, numInputTensors) + +def MLXGraphAddNumOutputTensors(builder, numOutputTensors): + builder.PrependUint32Slot(3, numOutputTensors, 0) + +def AddNumOutputTensors(builder, numOutputTensors): + MLXGraphAddNumOutputTensors(builder, numOutputTensors) + +def MLXGraphAddNumMutableBufferTensors(builder, numMutableBufferTensors): + builder.PrependUint32Slot(4, numMutableBufferTensors, 0) + +def AddNumMutableBufferTensors(builder, numMutableBufferTensors): + MLXGraphAddNumMutableBufferTensors(builder, numMutableBufferTensors) + +def MLXGraphAddNumTempTensors(builder, numTempTensors): + builder.PrependUint32Slot(5, numTempTensors, 0) + +def AddNumTempTensors(builder, numTempTensors): + MLXGraphAddNumTempTensors(builder, numTempTensors) + +def MLXGraphAddNumValues(builder, numValues): + builder.PrependUint32Slot(6, numValues, 0) + +def AddNumValues(builder, numValues): + MLXGraphAddNumValues(builder, numValues) + +def MLXGraphAddInstructionChains(builder, instructionChains): + builder.PrependUOffsetTRelativeSlot(7, flatbuffers.number_types.UOffsetTFlags.py_type(instructionChains), 0) + +def AddInstructionChains(builder, instructionChains): + MLXGraphAddInstructionChains(builder, instructionChains) + +def MLXGraphStartInstructionChainsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartInstructionChainsVector(builder, numElems): + return MLXGraphStartInstructionChainsVector(builder, numElems) + +def MLXGraphAddMainChainIdx(builder, mainChainIdx): + builder.PrependUint32Slot(8, mainChainIdx, 0) + +def AddMainChainIdx(builder, mainChainIdx): + MLXGraphAddMainChainIdx(builder, mainChainIdx) + +def MLXGraphAddInitChainIdx(builder, initChainIdx): + builder.PrependInt32Slot(9, initChainIdx, -1) + +def AddInitChainIdx(builder, initChainIdx): + MLXGraphAddInitChainIdx(builder, initChainIdx) + +def MLXGraphAddInputMap(builder, inputMap): + builder.PrependUOffsetTRelativeSlot(10, flatbuffers.number_types.UOffsetTFlags.py_type(inputMap), 0) + +def AddInputMap(builder, inputMap): + MLXGraphAddInputMap(builder, inputMap) + +def MLXGraphStartInputMapVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartInputMapVector(builder, numElems): + return MLXGraphStartInputMapVector(builder, numElems) + +def MLXGraphAddOutputMap(builder, outputMap): + builder.PrependUOffsetTRelativeSlot(11, flatbuffers.number_types.UOffsetTFlags.py_type(outputMap), 0) + +def AddOutputMap(builder, outputMap): + MLXGraphAddOutputMap(builder, outputMap) + +def MLXGraphStartOutputMapVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutputMapVector(builder, numElems): + return MLXGraphStartOutputMapVector(builder, numElems) + +def MLXGraphAddMutableBufferMap(builder, mutableBufferMap): + builder.PrependUOffsetTRelativeSlot(12, flatbuffers.number_types.UOffsetTFlags.py_type(mutableBufferMap), 0) + +def AddMutableBufferMap(builder, mutableBufferMap): + MLXGraphAddMutableBufferMap(builder, mutableBufferMap) + +def MLXGraphStartMutableBufferMapVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartMutableBufferMapVector(builder, numElems): + return MLXGraphStartMutableBufferMapVector(builder, numElems) + +def MLXGraphAddNamedSlots(builder, namedSlots): + builder.PrependUOffsetTRelativeSlot(13, flatbuffers.number_types.UOffsetTFlags.py_type(namedSlots), 0) + +def AddNamedSlots(builder, namedSlots): + MLXGraphAddNamedSlots(builder, namedSlots) + +def MLXGraphStartNamedSlotsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartNamedSlotsVector(builder, numElems): + return MLXGraphStartNamedSlotsVector(builder, numElems) + +def MLXGraphAddTensorMeta(builder, tensorMeta): + builder.PrependUOffsetTRelativeSlot(14, flatbuffers.number_types.UOffsetTFlags.py_type(tensorMeta), 0) + +def AddTensorMeta(builder, tensorMeta): + MLXGraphAddTensorMeta(builder, tensorMeta) + +def MLXGraphStartTensorMetaVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartTensorMetaVector(builder, numElems): + return MLXGraphStartTensorMetaVector(builder, numElems) + +def MLXGraphEnd(builder): + return builder.EndObject() + +def End(builder): + return MLXGraphEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MaxNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MaxNode.py new file mode 100644 index 00000000000..84c1109dc57 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MaxNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MaxNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MaxNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMaxNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MaxNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MaxNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MaxNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MaxNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # MaxNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # MaxNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MaxNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # MaxNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def MaxNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + MaxNodeStart(builder) + +def MaxNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + MaxNodeAddX(builder, x) + +def MaxNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MaxNodeAddOut(builder, out) + +def MaxNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + MaxNodeAddAxes(builder, axes) + +def MaxNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return MaxNodeStartAxesVector(builder, numElems) + +def MaxNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + MaxNodeAddKeepdims(builder, keepdims) + +def MaxNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MaxNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MaximumNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MaximumNode.py new file mode 100644 index 00000000000..42354b5f834 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MaximumNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MaximumNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MaximumNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMaximumNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MaximumNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MaximumNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MaximumNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MaximumNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def MaximumNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + MaximumNodeStart(builder) + +def MaximumNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + MaximumNodeAddA(builder, a) + +def MaximumNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + MaximumNodeAddB(builder, b) + +def MaximumNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MaximumNodeAddOut(builder, out) + +def MaximumNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MaximumNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MeanNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MeanNode.py new file mode 100644 index 00000000000..16cd84c2b2b --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MeanNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MeanNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MeanNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMeanNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MeanNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MeanNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MeanNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MeanNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # MeanNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # MeanNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MeanNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # MeanNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def MeanNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + MeanNodeStart(builder) + +def MeanNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + MeanNodeAddX(builder, x) + +def MeanNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MeanNodeAddOut(builder, out) + +def MeanNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + MeanNodeAddAxes(builder, axes) + +def MeanNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return MeanNodeStartAxesVector(builder, numElems) + +def MeanNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + MeanNodeAddKeepdims(builder, keepdims) + +def MeanNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MeanNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MedianNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MedianNode.py new file mode 100644 index 00000000000..c4bc88cf341 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MedianNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MedianNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MedianNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMedianNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MedianNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MedianNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MedianNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MedianNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # MedianNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # MedianNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MedianNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # MedianNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def MedianNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + MedianNodeStart(builder) + +def MedianNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + MedianNodeAddX(builder, x) + +def MedianNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MedianNodeAddOut(builder, out) + +def MedianNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + MedianNodeAddAxes(builder, axes) + +def MedianNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return MedianNodeStartAxesVector(builder, numElems) + +def MedianNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + MedianNodeAddKeepdims(builder, keepdims) + +def MedianNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MedianNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MetalKernelNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MetalKernelNode.py new file mode 100644 index 00000000000..c6b0df5362e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MetalKernelNode.py @@ -0,0 +1,550 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MetalKernelNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MetalKernelNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMetalKernelNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MetalKernelNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MetalKernelNode + def Name(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # MetalKernelNode + def Source(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # MetalKernelNode + def Inputs(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MetalKernelNode + def InputsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def InputsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # MetalKernelNode + def Outputs(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MetalKernelNode + def OutputsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def OutputsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + return o == 0 + + # MetalKernelNode + def Grid(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MetalKernelNode + def GridLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def GridIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + return o == 0 + + # MetalKernelNode + def Threadgroup(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MetalKernelNode + def ThreadgroupLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def ThreadgroupIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + return o == 0 + + # MetalKernelNode + def Header(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # MetalKernelNode + def InputNames(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.String(a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return "" + + # MetalKernelNode + def InputNamesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def InputNamesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + return o == 0 + + # MetalKernelNode + def OutputNames(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.String(a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return "" + + # MetalKernelNode + def OutputNamesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def OutputNamesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + return o == 0 + + # MetalKernelNode + def EnsureRowContiguous(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(22)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return True + + # MetalKernelNode + def AtomicOutputs(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(24)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + + # MetalKernelNode + def OutputShapesFlat(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MetalKernelNode + def OutputShapesFlatLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def OutputShapesFlatIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(26)) + return o == 0 + + # MetalKernelNode + def OutputShapeLengths(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # MetalKernelNode + def OutputShapeLengthsAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # MetalKernelNode + def OutputShapeLengthsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def OutputShapeLengthsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(28)) + return o == 0 + + # MetalKernelNode + def OutputDtypes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int8Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 1)) + return 0 + + # MetalKernelNode + def OutputDtypesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int8Flags, o) + return 0 + + # MetalKernelNode + def OutputDtypesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def OutputDtypesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(30)) + return o == 0 + + # MetalKernelNode + def TemplateArgNames(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.String(a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return "" + + # MetalKernelNode + def TemplateArgNamesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def TemplateArgNamesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(32)) + return o == 0 + + # MetalKernelNode + def TemplateArgKinds(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(34)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int8Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 1)) + return 0 + + # MetalKernelNode + def TemplateArgKindsAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(34)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int8Flags, o) + return 0 + + # MetalKernelNode + def TemplateArgKindsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(34)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def TemplateArgKindsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(34)) + return o == 0 + + # MetalKernelNode + def TemplateArgValues(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(36)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # MetalKernelNode + def TemplateArgValuesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(36)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # MetalKernelNode + def TemplateArgValuesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(36)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MetalKernelNode + def TemplateArgValuesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(36)) + return o == 0 + + # MetalKernelNode + def InitValue(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(38)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return None + +def MetalKernelNodeStart(builder): + builder.StartObject(18) + +def Start(builder): + MetalKernelNodeStart(builder) + +def MetalKernelNodeAddName(builder, name): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(name), 0) + +def AddName(builder, name): + MetalKernelNodeAddName(builder, name) + +def MetalKernelNodeAddSource(builder, source): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(source), 0) + +def AddSource(builder, source): + MetalKernelNodeAddSource(builder, source) + +def MetalKernelNodeAddInputs(builder, inputs): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(inputs), 0) + +def AddInputs(builder, inputs): + MetalKernelNodeAddInputs(builder, inputs) + +def MetalKernelNodeStartInputsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartInputsVector(builder, numElems): + return MetalKernelNodeStartInputsVector(builder, numElems) + +def MetalKernelNodeAddOutputs(builder, outputs): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(outputs), 0) + +def AddOutputs(builder, outputs): + MetalKernelNodeAddOutputs(builder, outputs) + +def MetalKernelNodeStartOutputsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutputsVector(builder, numElems): + return MetalKernelNodeStartOutputsVector(builder, numElems) + +def MetalKernelNodeAddGrid(builder, grid): + builder.PrependUOffsetTRelativeSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(grid), 0) + +def AddGrid(builder, grid): + MetalKernelNodeAddGrid(builder, grid) + +def MetalKernelNodeStartGridVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartGridVector(builder, numElems): + return MetalKernelNodeStartGridVector(builder, numElems) + +def MetalKernelNodeAddThreadgroup(builder, threadgroup): + builder.PrependUOffsetTRelativeSlot(5, flatbuffers.number_types.UOffsetTFlags.py_type(threadgroup), 0) + +def AddThreadgroup(builder, threadgroup): + MetalKernelNodeAddThreadgroup(builder, threadgroup) + +def MetalKernelNodeStartThreadgroupVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartThreadgroupVector(builder, numElems): + return MetalKernelNodeStartThreadgroupVector(builder, numElems) + +def MetalKernelNodeAddHeader(builder, header): + builder.PrependUOffsetTRelativeSlot(6, flatbuffers.number_types.UOffsetTFlags.py_type(header), 0) + +def AddHeader(builder, header): + MetalKernelNodeAddHeader(builder, header) + +def MetalKernelNodeAddInputNames(builder, inputNames): + builder.PrependUOffsetTRelativeSlot(7, flatbuffers.number_types.UOffsetTFlags.py_type(inputNames), 0) + +def AddInputNames(builder, inputNames): + MetalKernelNodeAddInputNames(builder, inputNames) + +def MetalKernelNodeStartInputNamesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartInputNamesVector(builder, numElems): + return MetalKernelNodeStartInputNamesVector(builder, numElems) + +def MetalKernelNodeAddOutputNames(builder, outputNames): + builder.PrependUOffsetTRelativeSlot(8, flatbuffers.number_types.UOffsetTFlags.py_type(outputNames), 0) + +def AddOutputNames(builder, outputNames): + MetalKernelNodeAddOutputNames(builder, outputNames) + +def MetalKernelNodeStartOutputNamesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutputNamesVector(builder, numElems): + return MetalKernelNodeStartOutputNamesVector(builder, numElems) + +def MetalKernelNodeAddEnsureRowContiguous(builder, ensureRowContiguous): + builder.PrependBoolSlot(9, ensureRowContiguous, 1) + +def AddEnsureRowContiguous(builder, ensureRowContiguous): + MetalKernelNodeAddEnsureRowContiguous(builder, ensureRowContiguous) + +def MetalKernelNodeAddAtomicOutputs(builder, atomicOutputs): + builder.PrependBoolSlot(10, atomicOutputs, 0) + +def AddAtomicOutputs(builder, atomicOutputs): + MetalKernelNodeAddAtomicOutputs(builder, atomicOutputs) + +def MetalKernelNodeAddOutputShapesFlat(builder, outputShapesFlat): + builder.PrependUOffsetTRelativeSlot(11, flatbuffers.number_types.UOffsetTFlags.py_type(outputShapesFlat), 0) + +def AddOutputShapesFlat(builder, outputShapesFlat): + MetalKernelNodeAddOutputShapesFlat(builder, outputShapesFlat) + +def MetalKernelNodeStartOutputShapesFlatVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutputShapesFlatVector(builder, numElems): + return MetalKernelNodeStartOutputShapesFlatVector(builder, numElems) + +def MetalKernelNodeAddOutputShapeLengths(builder, outputShapeLengths): + builder.PrependUOffsetTRelativeSlot(12, flatbuffers.number_types.UOffsetTFlags.py_type(outputShapeLengths), 0) + +def AddOutputShapeLengths(builder, outputShapeLengths): + MetalKernelNodeAddOutputShapeLengths(builder, outputShapeLengths) + +def MetalKernelNodeStartOutputShapeLengthsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutputShapeLengthsVector(builder, numElems): + return MetalKernelNodeStartOutputShapeLengthsVector(builder, numElems) + +def MetalKernelNodeAddOutputDtypes(builder, outputDtypes): + builder.PrependUOffsetTRelativeSlot(13, flatbuffers.number_types.UOffsetTFlags.py_type(outputDtypes), 0) + +def AddOutputDtypes(builder, outputDtypes): + MetalKernelNodeAddOutputDtypes(builder, outputDtypes) + +def MetalKernelNodeStartOutputDtypesVector(builder, numElems): + return builder.StartVector(1, numElems, 1) + +def StartOutputDtypesVector(builder, numElems): + return MetalKernelNodeStartOutputDtypesVector(builder, numElems) + +def MetalKernelNodeAddTemplateArgNames(builder, templateArgNames): + builder.PrependUOffsetTRelativeSlot(14, flatbuffers.number_types.UOffsetTFlags.py_type(templateArgNames), 0) + +def AddTemplateArgNames(builder, templateArgNames): + MetalKernelNodeAddTemplateArgNames(builder, templateArgNames) + +def MetalKernelNodeStartTemplateArgNamesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartTemplateArgNamesVector(builder, numElems): + return MetalKernelNodeStartTemplateArgNamesVector(builder, numElems) + +def MetalKernelNodeAddTemplateArgKinds(builder, templateArgKinds): + builder.PrependUOffsetTRelativeSlot(15, flatbuffers.number_types.UOffsetTFlags.py_type(templateArgKinds), 0) + +def AddTemplateArgKinds(builder, templateArgKinds): + MetalKernelNodeAddTemplateArgKinds(builder, templateArgKinds) + +def MetalKernelNodeStartTemplateArgKindsVector(builder, numElems): + return builder.StartVector(1, numElems, 1) + +def StartTemplateArgKindsVector(builder, numElems): + return MetalKernelNodeStartTemplateArgKindsVector(builder, numElems) + +def MetalKernelNodeAddTemplateArgValues(builder, templateArgValues): + builder.PrependUOffsetTRelativeSlot(16, flatbuffers.number_types.UOffsetTFlags.py_type(templateArgValues), 0) + +def AddTemplateArgValues(builder, templateArgValues): + MetalKernelNodeAddTemplateArgValues(builder, templateArgValues) + +def MetalKernelNodeStartTemplateArgValuesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartTemplateArgValuesVector(builder, numElems): + return MetalKernelNodeStartTemplateArgValuesVector(builder, numElems) + +def MetalKernelNodeAddInitValue(builder, initValue): + builder.PrependFloat32Slot(17, initValue, None) + +def AddInitValue(builder, initValue): + MetalKernelNodeAddInitValue(builder, initValue) + +def MetalKernelNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MetalKernelNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MinNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MinNode.py new file mode 100644 index 00000000000..5f83030aaa2 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MinNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MinNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MinNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMinNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MinNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MinNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MinNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MinNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # MinNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # MinNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # MinNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # MinNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def MinNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + MinNodeStart(builder) + +def MinNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + MinNodeAddX(builder, x) + +def MinNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MinNodeAddOut(builder, out) + +def MinNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + MinNodeAddAxes(builder, axes) + +def MinNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return MinNodeStartAxesVector(builder, numElems) + +def MinNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + MinNodeAddKeepdims(builder, keepdims) + +def MinNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MinNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MinimumNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MinimumNode.py new file mode 100644 index 00000000000..44abafdb1c6 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MinimumNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MinimumNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MinimumNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMinimumNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MinimumNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MinimumNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MinimumNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MinimumNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def MinimumNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + MinimumNodeStart(builder) + +def MinimumNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + MinimumNodeAddA(builder, a) + +def MinimumNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + MinimumNodeAddB(builder, b) + +def MinimumNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MinimumNodeAddOut(builder, out) + +def MinimumNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MinimumNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ModIntNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ModIntNode.py new file mode 100644 index 00000000000..bda293d0edc --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ModIntNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ModIntNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ModIntNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsModIntNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ModIntNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ModIntNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ModIntNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ModIntNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ModIntNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + ModIntNodeStart(builder) + +def ModIntNodeAddA(builder, a): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + ModIntNodeAddA(builder, a) + +def ModIntNodeAddB(builder, b): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + ModIntNodeAddB(builder, b) + +def ModIntNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ModIntNodeAddOut(builder, out) + +def ModIntNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ModIntNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MultiplyIntNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MultiplyIntNode.py new file mode 100644 index 00000000000..d74974f16c6 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MultiplyIntNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MultiplyIntNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MultiplyIntNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMultiplyIntNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MultiplyIntNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MultiplyIntNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MultiplyIntNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MultiplyIntNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def MultiplyIntNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + MultiplyIntNodeStart(builder) + +def MultiplyIntNodeAddA(builder, a): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + MultiplyIntNodeAddA(builder, a) + +def MultiplyIntNodeAddB(builder, b): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + MultiplyIntNodeAddB(builder, b) + +def MultiplyIntNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MultiplyIntNodeAddOut(builder, out) + +def MultiplyIntNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MultiplyIntNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/MultiplyNode.py b/backends/mlx/serialization/_generated/mlx_delegate/MultiplyNode.py new file mode 100644 index 00000000000..f97aa77bf01 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/MultiplyNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class MultiplyNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = MultiplyNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsMultiplyNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # MultiplyNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # MultiplyNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MultiplyNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # MultiplyNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def MultiplyNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + MultiplyNodeStart(builder) + +def MultiplyNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + MultiplyNodeAddA(builder, a) + +def MultiplyNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + MultiplyNodeAddB(builder, b) + +def MultiplyNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + MultiplyNodeAddOut(builder, out) + +def MultiplyNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return MultiplyNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/NamedSlot.py b/backends/mlx/serialization/_generated/mlx_delegate/NamedSlot.py new file mode 100644 index 00000000000..c8ac6f953e7 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/NamedSlot.py @@ -0,0 +1,67 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class NamedSlot(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = NamedSlot() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsNamedSlot(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # NamedSlot + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # NamedSlot + def Name(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # NamedSlot + def Slot(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.SlotVariant import SlotVariant + obj = SlotVariant() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def NamedSlotStart(builder): + builder.StartObject(2) + +def Start(builder): + NamedSlotStart(builder) + +def NamedSlotAddName(builder, name): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(name), 0) + +def AddName(builder, name): + NamedSlotAddName(builder, name) + +def NamedSlotAddSlot(builder, slot): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(slot), 0) + +def AddSlot(builder, slot): + NamedSlotAddSlot(builder, slot) + +def NamedSlotEnd(builder): + return builder.EndObject() + +def End(builder): + return NamedSlotEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/NegNode.py b/backends/mlx/serialization/_generated/mlx_delegate/NegNode.py new file mode 100644 index 00000000000..be18e619b55 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/NegNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class NegNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = NegNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsNegNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # NegNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # NegNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # NegNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def NegNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + NegNodeStart(builder) + +def NegNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + NegNodeAddX(builder, x) + +def NegNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + NegNodeAddOut(builder, out) + +def NegNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return NegNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/NoopNode.py b/backends/mlx/serialization/_generated/mlx_delegate/NoopNode.py new file mode 100644 index 00000000000..44adda50594 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/NoopNode.py @@ -0,0 +1,37 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class NoopNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = NoopNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsNoopNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # NoopNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + +def NoopNodeStart(builder): + builder.StartObject(0) + +def Start(builder): + NoopNodeStart(builder) + +def NoopNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return NoopNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/NotEqualNode.py b/backends/mlx/serialization/_generated/mlx_delegate/NotEqualNode.py new file mode 100644 index 00000000000..675e4fe17e6 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/NotEqualNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class NotEqualNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = NotEqualNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsNotEqualNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # NotEqualNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # NotEqualNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # NotEqualNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # NotEqualNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def NotEqualNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + NotEqualNodeStart(builder) + +def NotEqualNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + NotEqualNodeAddA(builder, a) + +def NotEqualNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + NotEqualNodeAddB(builder, b) + +def NotEqualNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + NotEqualNodeAddOut(builder, out) + +def NotEqualNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return NotEqualNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/OpNode.py b/backends/mlx/serialization/_generated/mlx_delegate/OpNode.py new file mode 100644 index 00000000000..23d54b847f5 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/OpNode.py @@ -0,0 +1,139 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +class OpNode(object): + NONE = 0 + NoopNode = 1 + IdCopyNode = 2 + AddmmNode = 3 + ItemIntNode = 4 + ExpandDimsNode = 5 + TileNode = 6 + TakeAlongAxisNode = 7 + TakeNode = 8 + RMSNormNode = 9 + LayerNormNode = 10 + RopeNode = 11 + SdpaNode = 12 + AddNode = 13 + AddIntNode = 14 + SubtractIntNode = 15 + MultiplyIntNode = 16 + FloorDivideIntNode = 17 + SymSizeNode = 18 + MultiplyNode = 19 + DivideNode = 20 + SubtractNode = 21 + Conv1DNode = 22 + Conv2DNode = 23 + Conv3DNode = 24 + GeluNode = 25 + ARangeNode = 26 + SiluNode = 27 + SigmoidNode = 28 + TanhNode = 29 + SqueezeNode = 30 + SplitNode = 31 + RsqrtNode = 32 + MaximumNode = 33 + MinimumNode = 34 + LogNode = 35 + SoftmaxNode = 36 + BroadcastToNode = 37 + PadNode = 38 + WhereNode = 39 + ReshapeNode = 40 + TransposeNode = 41 + AsStridedNode = 42 + ContiguousNode = 43 + GatherNode = 44 + SliceNode = 45 + AsTypeNode = 46 + ConcatenateNode = 47 + FullNode = 48 + FullLikeNode = 49 + ArgmaxNode = 50 + SliceUpdateNode = 51 + IndexCopyNode = 52 + DequantizeNode = 53 + LessNode = 54 + LessEqualNode = 55 + GreaterNode = 56 + GreaterEqualNode = 57 + EqualNode = 58 + NotEqualNode = 59 + LogicalNotNode = 60 + LogicalAndNode = 61 + LogicalOrNode = 62 + TriNode = 63 + TrilNode = 64 + TriuNode = 65 + FloorNode = 66 + CeilNode = 67 + SquareNode = 68 + ExpNode = 69 + SinNode = 70 + CosNode = 71 + TanNode = 72 + ArcsinNode = 73 + ArccosNode = 74 + ArctanNode = 75 + SinhNode = 76 + CoshNode = 77 + ArcsinhNode = 78 + ArccoshNode = 79 + ArctanhNode = 80 + Log2Node = 81 + Log10Node = 82 + Log1pNode = 83 + ErfNode = 84 + Expm1Node = 85 + RoundNode = 86 + ReciprocalNode = 87 + SqrtNode = 88 + AbsNode = 89 + NegNode = 90 + Atan2Node = 91 + LogAddExpNode = 92 + FloorDivideNode = 93 + PowerNode = 94 + LogSumExpNode = 95 + SumNode = 96 + MeanNode = 97 + VarNode = 98 + StdNode = 99 + ProdNode = 100 + MaxNode = 101 + MinNode = 102 + ArgminNode = 103 + MedianNode = 104 + ModIntNode = 105 + RemainderNode = 106 + ConvTranspose1DNode = 107 + ConvTranspose2DNode = 108 + ConvTranspose3DNode = 109 + ClipNode = 110 + CumsumNode = 111 + StackNode = 112 + SignNode = 113 + AnyNode = 114 + AllNode = 115 + RepeatNode = 116 + SortNode = 117 + ArgsortNode = 118 + PartitionNode = 119 + ArgPartitionNode = 120 + QuantizedMatmulNode = 121 + ScatterAddNode = 122 + GatherMmNode = 123 + GatherQmmNode = 124 + ScanNode = 125 + MetalKernelNode = 126 + BitwiseInvertNode = 127 + RollNode = 128 + BitwiseAndNode = 129 + BitwiseOrNode = 130 + BitwiseXorNode = 131 + IfNode = 132 + RandomBitsNode = 133 diff --git a/backends/mlx/serialization/_generated/mlx_delegate/PadNode.py b/backends/mlx/serialization/_generated/mlx_delegate/PadNode.py new file mode 100644 index 00000000000..824c53795fd --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/PadNode.py @@ -0,0 +1,134 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class PadNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = PadNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsPadNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # PadNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # PadNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PadNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PadNode + def PadWidth(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PadNode + def PadWidthLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # PadNode + def PadWidthIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # PadNode + def Mode(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # PadNode + def ConstantValue(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 0.0 + +def PadNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + PadNodeStart(builder) + +def PadNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + PadNodeAddX(builder, x) + +def PadNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + PadNodeAddOut(builder, out) + +def PadNodeAddPadWidth(builder, padWidth): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(padWidth), 0) + +def AddPadWidth(builder, padWidth): + PadNodeAddPadWidth(builder, padWidth) + +def PadNodeStartPadWidthVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartPadWidthVector(builder, numElems): + return PadNodeStartPadWidthVector(builder, numElems) + +def PadNodeAddMode(builder, mode): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(mode), 0) + +def AddMode(builder, mode): + PadNodeAddMode(builder, mode) + +def PadNodeAddConstantValue(builder, constantValue): + builder.PrependFloat32Slot(4, constantValue, 0.0) + +def AddConstantValue(builder, constantValue): + PadNodeAddConstantValue(builder, constantValue) + +def PadNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return PadNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/PartitionNode.py b/backends/mlx/serialization/_generated/mlx_delegate/PartitionNode.py new file mode 100644 index 00000000000..78f1ea26ffc --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/PartitionNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class PartitionNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = PartitionNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsPartitionNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # PartitionNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # PartitionNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PartitionNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PartitionNode + def Kth(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PartitionNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def PartitionNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + PartitionNodeStart(builder) + +def PartitionNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + PartitionNodeAddX(builder, x) + +def PartitionNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + PartitionNodeAddOut(builder, out) + +def PartitionNodeAddKth(builder, kth): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(kth), 0) + +def AddKth(builder, kth): + PartitionNodeAddKth(builder, kth) + +def PartitionNodeAddAxis(builder, axis): + builder.PrependInt32Slot(3, axis, 0) + +def AddAxis(builder, axis): + PartitionNodeAddAxis(builder, axis) + +def PartitionNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return PartitionNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/PowerNode.py b/backends/mlx/serialization/_generated/mlx_delegate/PowerNode.py new file mode 100644 index 00000000000..df84d893513 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/PowerNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class PowerNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = PowerNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsPowerNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # PowerNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # PowerNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PowerNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # PowerNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def PowerNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + PowerNodeStart(builder) + +def PowerNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + PowerNodeAddA(builder, a) + +def PowerNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + PowerNodeAddB(builder, b) + +def PowerNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + PowerNodeAddOut(builder, out) + +def PowerNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return PowerNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ProdNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ProdNode.py new file mode 100644 index 00000000000..c40d9e32b82 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ProdNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ProdNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ProdNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsProdNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ProdNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ProdNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ProdNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ProdNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # ProdNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # ProdNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ProdNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # ProdNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def ProdNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + ProdNodeStart(builder) + +def ProdNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ProdNodeAddX(builder, x) + +def ProdNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ProdNodeAddOut(builder, out) + +def ProdNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + ProdNodeAddAxes(builder, axes) + +def ProdNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return ProdNodeStartAxesVector(builder, numElems) + +def ProdNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + ProdNodeAddKeepdims(builder, keepdims) + +def ProdNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ProdNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/QuantizedMatmulNode.py b/backends/mlx/serialization/_generated/mlx_delegate/QuantizedMatmulNode.py new file mode 100644 index 00000000000..2a36f12c73f --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/QuantizedMatmulNode.py @@ -0,0 +1,174 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class QuantizedMatmulNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = QuantizedMatmulNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsQuantizedMatmulNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # QuantizedMatmulNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # QuantizedMatmulNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # QuantizedMatmulNode + def W(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # QuantizedMatmulNode + def Scales(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # QuantizedMatmulNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # QuantizedMatmulNode + def Biases(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # QuantizedMatmulNode + def GroupSize(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # QuantizedMatmulNode + def Bits(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # QuantizedMatmulNode + def Mode(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.String(o + self._tab.Pos) + return None + + # QuantizedMatmulNode + def Transpose(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(20)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return True + +def QuantizedMatmulNodeStart(builder): + builder.StartObject(9) + +def Start(builder): + QuantizedMatmulNodeStart(builder) + +def QuantizedMatmulNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + QuantizedMatmulNodeAddX(builder, x) + +def QuantizedMatmulNodeAddW(builder, w): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(w), 0) + +def AddW(builder, w): + QuantizedMatmulNodeAddW(builder, w) + +def QuantizedMatmulNodeAddScales(builder, scales): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(scales), 0) + +def AddScales(builder, scales): + QuantizedMatmulNodeAddScales(builder, scales) + +def QuantizedMatmulNodeAddOut(builder, out): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + QuantizedMatmulNodeAddOut(builder, out) + +def QuantizedMatmulNodeAddBiases(builder, biases): + builder.PrependStructSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(biases), 0) + +def AddBiases(builder, biases): + QuantizedMatmulNodeAddBiases(builder, biases) + +def QuantizedMatmulNodeAddGroupSize(builder, groupSize): + builder.PrependInt32Slot(5, groupSize, 0) + +def AddGroupSize(builder, groupSize): + QuantizedMatmulNodeAddGroupSize(builder, groupSize) + +def QuantizedMatmulNodeAddBits(builder, bits): + builder.PrependInt32Slot(6, bits, 0) + +def AddBits(builder, bits): + QuantizedMatmulNodeAddBits(builder, bits) + +def QuantizedMatmulNodeAddMode(builder, mode): + builder.PrependUOffsetTRelativeSlot(7, flatbuffers.number_types.UOffsetTFlags.py_type(mode), 0) + +def AddMode(builder, mode): + QuantizedMatmulNodeAddMode(builder, mode) + +def QuantizedMatmulNodeAddTranspose(builder, transpose): + builder.PrependBoolSlot(8, transpose, 1) + +def AddTranspose(builder, transpose): + QuantizedMatmulNodeAddTranspose(builder, transpose) + +def QuantizedMatmulNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return QuantizedMatmulNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RMSNormNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RMSNormNode.py new file mode 100644 index 00000000000..028ba834e12 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RMSNormNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RMSNormNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RMSNormNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRMSNormNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RMSNormNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RMSNormNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RMSNormNode + def Weight(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RMSNormNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RMSNormNode + def Eps(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 0.0 + +def RMSNormNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + RMSNormNodeStart(builder) + +def RMSNormNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + RMSNormNodeAddX(builder, x) + +def RMSNormNodeAddWeight(builder, weight): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(weight), 0) + +def AddWeight(builder, weight): + RMSNormNodeAddWeight(builder, weight) + +def RMSNormNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RMSNormNodeAddOut(builder, out) + +def RMSNormNodeAddEps(builder, eps): + builder.PrependFloat32Slot(3, eps, 0.0) + +def AddEps(builder, eps): + RMSNormNodeAddEps(builder, eps) + +def RMSNormNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RMSNormNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RandomBitsNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RandomBitsNode.py new file mode 100644 index 00000000000..9edab48550e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RandomBitsNode.py @@ -0,0 +1,121 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RandomBitsNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RandomBitsNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRandomBitsNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RandomBitsNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RandomBitsNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RandomBitsNode + def Shape(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RandomBitsNode + def ShapeLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # RandomBitsNode + def ShapeIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + return o == 0 + + # RandomBitsNode + def Seed(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RandomBitsNode + def Width(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 4 + +def RandomBitsNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + RandomBitsNodeStart(builder) + +def RandomBitsNodeAddOut(builder, out): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RandomBitsNodeAddOut(builder, out) + +def RandomBitsNodeAddShape(builder, shape): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(shape), 0) + +def AddShape(builder, shape): + RandomBitsNodeAddShape(builder, shape) + +def RandomBitsNodeStartShapeVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShapeVector(builder, numElems): + return RandomBitsNodeStartShapeVector(builder, numElems) + +def RandomBitsNodeAddSeed(builder, seed): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(seed), 0) + +def AddSeed(builder, seed): + RandomBitsNodeAddSeed(builder, seed) + +def RandomBitsNodeAddWidth(builder, width): + builder.PrependInt32Slot(3, width, 4) + +def AddWidth(builder, width): + RandomBitsNodeAddWidth(builder, width) + +def RandomBitsNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RandomBitsNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ReciprocalNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ReciprocalNode.py new file mode 100644 index 00000000000..816390edc90 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ReciprocalNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ReciprocalNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ReciprocalNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsReciprocalNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ReciprocalNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ReciprocalNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ReciprocalNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def ReciprocalNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + ReciprocalNodeStart(builder) + +def ReciprocalNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ReciprocalNodeAddX(builder, x) + +def ReciprocalNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ReciprocalNodeAddOut(builder, out) + +def ReciprocalNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ReciprocalNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RemainderNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RemainderNode.py new file mode 100644 index 00000000000..7cd1c7ab2d1 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RemainderNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RemainderNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RemainderNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRemainderNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RemainderNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RemainderNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RemainderNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RemainderNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def RemainderNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + RemainderNodeStart(builder) + +def RemainderNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + RemainderNodeAddA(builder, a) + +def RemainderNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + RemainderNodeAddB(builder, b) + +def RemainderNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RemainderNodeAddOut(builder, out) + +def RemainderNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RemainderNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RepeatNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RepeatNode.py new file mode 100644 index 00000000000..1aed6c83867 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RepeatNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RepeatNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RepeatNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRepeatNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RepeatNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RepeatNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RepeatNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RepeatNode + def Repeats(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RepeatNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def RepeatNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + RepeatNodeStart(builder) + +def RepeatNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + RepeatNodeAddX(builder, x) + +def RepeatNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RepeatNodeAddOut(builder, out) + +def RepeatNodeAddRepeats(builder, repeats): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(repeats), 0) + +def AddRepeats(builder, repeats): + RepeatNodeAddRepeats(builder, repeats) + +def RepeatNodeAddAxis(builder, axis): + builder.PrependInt32Slot(3, axis, 0) + +def AddAxis(builder, axis): + RepeatNodeAddAxis(builder, axis) + +def RepeatNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RepeatNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ReshapeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ReshapeNode.py new file mode 100644 index 00000000000..8b78bff24e2 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ReshapeNode.py @@ -0,0 +1,108 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ReshapeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ReshapeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsReshapeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ReshapeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ReshapeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ReshapeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ReshapeNode + def Shape(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ReshapeNode + def ShapeLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ReshapeNode + def ShapeIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + +def ReshapeNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + ReshapeNodeStart(builder) + +def ReshapeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ReshapeNodeAddX(builder, x) + +def ReshapeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ReshapeNodeAddOut(builder, out) + +def ReshapeNodeAddShape(builder, shape): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(shape), 0) + +def AddShape(builder, shape): + ReshapeNodeAddShape(builder, shape) + +def ReshapeNodeStartShapeVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShapeVector(builder, numElems): + return ReshapeNodeStartShapeVector(builder, numElems) + +def ReshapeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ReshapeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RollNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RollNode.py new file mode 100644 index 00000000000..04ca366c2ee --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RollNode.py @@ -0,0 +1,147 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RollNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RollNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRollNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RollNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RollNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RollNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RollNode + def Shift(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RollNode + def ShiftLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # RollNode + def ShiftIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # RollNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # RollNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # RollNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # RollNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + return o == 0 + +def RollNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + RollNodeStart(builder) + +def RollNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + RollNodeAddX(builder, x) + +def RollNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RollNodeAddOut(builder, out) + +def RollNodeAddShift(builder, shift): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(shift), 0) + +def AddShift(builder, shift): + RollNodeAddShift(builder, shift) + +def RollNodeStartShiftVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShiftVector(builder, numElems): + return RollNodeStartShiftVector(builder, numElems) + +def RollNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + RollNodeAddAxes(builder, axes) + +def RollNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return RollNodeStartAxesVector(builder, numElems) + +def RollNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RollNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RopeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RopeNode.py new file mode 100644 index 00000000000..1d2498986b2 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RopeNode.py @@ -0,0 +1,157 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RopeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RopeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRopeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RopeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RopeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RopeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RopeNode + def Dims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # RopeNode + def Offset(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.VidOrTid import VidOrTid + obj = VidOrTid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RopeNode + def Freqs(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RopeNode + def Traditional(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + + # RopeNode + def Base(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 500000.0 + + # RopeNode + def Scale(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(18)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 1.0 + +def RopeNodeStart(builder): + builder.StartObject(8) + +def Start(builder): + RopeNodeStart(builder) + +def RopeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + RopeNodeAddX(builder, x) + +def RopeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RopeNodeAddOut(builder, out) + +def RopeNodeAddDims(builder, dims): + builder.PrependInt32Slot(2, dims, 0) + +def AddDims(builder, dims): + RopeNodeAddDims(builder, dims) + +def RopeNodeAddOffset(builder, offset): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(offset), 0) + +def AddOffset(builder, offset): + RopeNodeAddOffset(builder, offset) + +def RopeNodeAddFreqs(builder, freqs): + builder.PrependStructSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(freqs), 0) + +def AddFreqs(builder, freqs): + RopeNodeAddFreqs(builder, freqs) + +def RopeNodeAddTraditional(builder, traditional): + builder.PrependBoolSlot(5, traditional, 0) + +def AddTraditional(builder, traditional): + RopeNodeAddTraditional(builder, traditional) + +def RopeNodeAddBase(builder, base): + builder.PrependFloat32Slot(6, base, 500000.0) + +def AddBase(builder, base): + RopeNodeAddBase(builder, base) + +def RopeNodeAddScale(builder, scale): + builder.PrependFloat32Slot(7, scale, 1.0) + +def AddScale(builder, scale): + RopeNodeAddScale(builder, scale) + +def RopeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RopeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RoundNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RoundNode.py new file mode 100644 index 00000000000..0731d1ced1e --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RoundNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RoundNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RoundNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRoundNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RoundNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RoundNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RoundNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RoundNode + def Decimals(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def RoundNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + RoundNodeStart(builder) + +def RoundNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + RoundNodeAddX(builder, x) + +def RoundNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RoundNodeAddOut(builder, out) + +def RoundNodeAddDecimals(builder, decimals): + builder.PrependInt32Slot(2, decimals, 0) + +def AddDecimals(builder, decimals): + RoundNodeAddDecimals(builder, decimals) + +def RoundNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RoundNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/RsqrtNode.py b/backends/mlx/serialization/_generated/mlx_delegate/RsqrtNode.py new file mode 100644 index 00000000000..41d9ec2deaf --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/RsqrtNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class RsqrtNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = RsqrtNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsRsqrtNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # RsqrtNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # RsqrtNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # RsqrtNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def RsqrtNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + RsqrtNodeStart(builder) + +def RsqrtNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + RsqrtNodeAddX(builder, x) + +def RsqrtNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + RsqrtNodeAddOut(builder, out) + +def RsqrtNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return RsqrtNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ScanNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ScanNode.py new file mode 100644 index 00000000000..e6998120889 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ScanNode.py @@ -0,0 +1,207 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ScanNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ScanNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsScanNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ScanNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ScanNode + def Originals(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScanNode + def OriginalsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ScanNode + def OriginalsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + return o == 0 + + # ScanNode + def Sliced(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScanNode + def SlicedLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ScanNode + def SlicedIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + return o == 0 + + # ScanNode + def Outputs(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScanNode + def OutputsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ScanNode + def OutputsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # ScanNode + def Carry(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScanNode + def CarryLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # ScanNode + def CarryIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + return o == 0 + + # ScanNode + def BodyChainIdx(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ScanNode + def ScanAxis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def ScanNodeStart(builder): + builder.StartObject(6) + +def Start(builder): + ScanNodeStart(builder) + +def ScanNodeAddOriginals(builder, originals): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(originals), 0) + +def AddOriginals(builder, originals): + ScanNodeAddOriginals(builder, originals) + +def ScanNodeStartOriginalsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOriginalsVector(builder, numElems): + return ScanNodeStartOriginalsVector(builder, numElems) + +def ScanNodeAddSliced(builder, sliced): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(sliced), 0) + +def AddSliced(builder, sliced): + ScanNodeAddSliced(builder, sliced) + +def ScanNodeStartSlicedVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartSlicedVector(builder, numElems): + return ScanNodeStartSlicedVector(builder, numElems) + +def ScanNodeAddOutputs(builder, outputs): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(outputs), 0) + +def AddOutputs(builder, outputs): + ScanNodeAddOutputs(builder, outputs) + +def ScanNodeStartOutputsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutputsVector(builder, numElems): + return ScanNodeStartOutputsVector(builder, numElems) + +def ScanNodeAddCarry(builder, carry): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(carry), 0) + +def AddCarry(builder, carry): + ScanNodeAddCarry(builder, carry) + +def ScanNodeStartCarryVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartCarryVector(builder, numElems): + return ScanNodeStartCarryVector(builder, numElems) + +def ScanNodeAddBodyChainIdx(builder, bodyChainIdx): + builder.PrependInt32Slot(4, bodyChainIdx, 0) + +def AddBodyChainIdx(builder, bodyChainIdx): + ScanNodeAddBodyChainIdx(builder, bodyChainIdx) + +def ScanNodeAddScanAxis(builder, scanAxis): + builder.PrependInt32Slot(5, scanAxis, 1) + +def AddScanAxis(builder, scanAxis): + ScanNodeAddScanAxis(builder, scanAxis) + +def ScanNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ScanNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ScatterAddNode.py b/backends/mlx/serialization/_generated/mlx_delegate/ScatterAddNode.py new file mode 100644 index 00000000000..723fd776402 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ScatterAddNode.py @@ -0,0 +1,118 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ScatterAddNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ScatterAddNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsScatterAddNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ScatterAddNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ScatterAddNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScatterAddNode + def Indices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScatterAddNode + def Updates(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScatterAddNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # ScatterAddNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def ScatterAddNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + ScatterAddNodeStart(builder) + +def ScatterAddNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + ScatterAddNodeAddX(builder, x) + +def ScatterAddNodeAddIndices(builder, indices): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(indices), 0) + +def AddIndices(builder, indices): + ScatterAddNodeAddIndices(builder, indices) + +def ScatterAddNodeAddUpdates(builder, updates): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(updates), 0) + +def AddUpdates(builder, updates): + ScatterAddNodeAddUpdates(builder, updates) + +def ScatterAddNodeAddOut(builder, out): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + ScatterAddNodeAddOut(builder, out) + +def ScatterAddNodeAddAxis(builder, axis): + builder.PrependInt32Slot(4, axis, 0) + +def AddAxis(builder, axis): + ScatterAddNodeAddAxis(builder, axis) + +def ScatterAddNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return ScatterAddNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SdpaNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SdpaNode.py new file mode 100644 index 00000000000..23ffc653100 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SdpaNode.py @@ -0,0 +1,148 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SdpaNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SdpaNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSdpaNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SdpaNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SdpaNode + def Q(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SdpaNode + def K(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SdpaNode + def V(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SdpaNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SdpaNode + def Scale(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Float32Flags, o + self._tab.Pos) + return 0.0 + + # SdpaNode + def Mask(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SdpaNode + def Causal(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def SdpaNodeStart(builder): + builder.StartObject(7) + +def Start(builder): + SdpaNodeStart(builder) + +def SdpaNodeAddQ(builder, q): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(q), 0) + +def AddQ(builder, q): + SdpaNodeAddQ(builder, q) + +def SdpaNodeAddK(builder, k): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(k), 0) + +def AddK(builder, k): + SdpaNodeAddK(builder, k) + +def SdpaNodeAddV(builder, v): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(v), 0) + +def AddV(builder, v): + SdpaNodeAddV(builder, v) + +def SdpaNodeAddOut(builder, out): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SdpaNodeAddOut(builder, out) + +def SdpaNodeAddScale(builder, scale): + builder.PrependFloat32Slot(4, scale, 0.0) + +def AddScale(builder, scale): + SdpaNodeAddScale(builder, scale) + +def SdpaNodeAddMask(builder, mask): + builder.PrependStructSlot(5, flatbuffers.number_types.UOffsetTFlags.py_type(mask), 0) + +def AddMask(builder, mask): + SdpaNodeAddMask(builder, mask) + +def SdpaNodeAddCausal(builder, causal): + builder.PrependBoolSlot(6, causal, 0) + +def AddCausal(builder, causal): + SdpaNodeAddCausal(builder, causal) + +def SdpaNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SdpaNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/ShapeDim.py b/backends/mlx/serialization/_generated/mlx_delegate/ShapeDim.py new file mode 100644 index 00000000000..47362b73b05 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/ShapeDim.py @@ -0,0 +1,76 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class ShapeDim(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = ShapeDim() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsShapeDim(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # ShapeDim + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # ShapeDim + def Value(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return -1 + + # ShapeDim + def MinValue(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # ShapeDim + def MaxValue(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return -1 + +def ShapeDimStart(builder): + builder.StartObject(3) + +def Start(builder): + ShapeDimStart(builder) + +def ShapeDimAddValue(builder, value): + builder.PrependInt32Slot(0, value, -1) + +def AddValue(builder, value): + ShapeDimAddValue(builder, value) + +def ShapeDimAddMinValue(builder, minValue): + builder.PrependInt32Slot(1, minValue, 0) + +def AddMinValue(builder, minValue): + ShapeDimAddMinValue(builder, minValue) + +def ShapeDimAddMaxValue(builder, maxValue): + builder.PrependInt32Slot(2, maxValue, -1) + +def AddMaxValue(builder, maxValue): + ShapeDimAddMaxValue(builder, maxValue) + +def ShapeDimEnd(builder): + return builder.EndObject() + +def End(builder): + return ShapeDimEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SigmoidNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SigmoidNode.py new file mode 100644 index 00000000000..51d80fd89b4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SigmoidNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SigmoidNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SigmoidNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSigmoidNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SigmoidNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SigmoidNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SigmoidNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SigmoidNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SigmoidNodeStart(builder) + +def SigmoidNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SigmoidNodeAddX(builder, x) + +def SigmoidNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SigmoidNodeAddOut(builder, out) + +def SigmoidNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SigmoidNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SignNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SignNode.py new file mode 100644 index 00000000000..84dbddae300 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SignNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SignNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SignNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSignNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SignNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SignNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SignNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SignNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SignNodeStart(builder) + +def SignNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SignNodeAddX(builder, x) + +def SignNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SignNodeAddOut(builder, out) + +def SignNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SignNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SiluNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SiluNode.py new file mode 100644 index 00000000000..acc4df0e4ec --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SiluNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SiluNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SiluNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSiluNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SiluNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SiluNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SiluNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SiluNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SiluNodeStart(builder) + +def SiluNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SiluNodeAddX(builder, x) + +def SiluNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SiluNodeAddOut(builder, out) + +def SiluNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SiluNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SinNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SinNode.py new file mode 100644 index 00000000000..f7819283f9c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SinNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SinNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SinNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSinNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SinNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SinNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SinNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SinNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SinNodeStart(builder) + +def SinNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SinNodeAddX(builder, x) + +def SinNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SinNodeAddOut(builder, out) + +def SinNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SinNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SinhNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SinhNode.py new file mode 100644 index 00000000000..a4408479207 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SinhNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SinhNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SinhNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSinhNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SinhNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SinhNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SinhNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SinhNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SinhNodeStart(builder) + +def SinhNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SinhNodeAddX(builder, x) + +def SinhNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SinhNodeAddOut(builder, out) + +def SinhNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SinhNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SliceNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SliceNode.py new file mode 100644 index 00000000000..c28c646388d --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SliceNode.py @@ -0,0 +1,135 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SliceNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SliceNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSliceNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SliceNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SliceNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceNode + def Start(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceNode + def Stop(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceNode + def Step(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def SliceNodeStart(builder): + builder.StartObject(6) + +def Start(builder): + SliceNodeStart(builder) + +def SliceNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SliceNodeAddX(builder, x) + +def SliceNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SliceNodeAddOut(builder, out) + +def SliceNodeAddAxis(builder, axis): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axis), 0) + +def AddAxis(builder, axis): + SliceNodeAddAxis(builder, axis) + +def SliceNodeAddStart(builder, start): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(start), 0) + +def AddStart(builder, start): + SliceNodeAddStart(builder, start) + +def SliceNodeAddStop(builder, stop): + builder.PrependUOffsetTRelativeSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(stop), 0) + +def AddStop(builder, stop): + SliceNodeAddStop(builder, stop) + +def SliceNodeAddStep(builder, step): + builder.PrependInt32Slot(5, step, 1) + +def AddStep(builder, step): + SliceNodeAddStep(builder, step) + +def SliceNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SliceNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SliceUpdateNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SliceUpdateNode.py new file mode 100644 index 00000000000..b3a52b4fa43 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SliceUpdateNode.py @@ -0,0 +1,152 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SliceUpdateNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SliceUpdateNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSliceUpdateNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SliceUpdateNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SliceUpdateNode + def Dst(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceUpdateNode + def Update(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceUpdateNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceUpdateNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceUpdateNode + def Start(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceUpdateNode + def Stop(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(14)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SliceUpdateNode + def Step(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(16)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 1 + +def SliceUpdateNodeStart(builder): + builder.StartObject(7) + +def Start(builder): + SliceUpdateNodeStart(builder) + +def SliceUpdateNodeAddDst(builder, dst): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(dst), 0) + +def AddDst(builder, dst): + SliceUpdateNodeAddDst(builder, dst) + +def SliceUpdateNodeAddUpdate(builder, update): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(update), 0) + +def AddUpdate(builder, update): + SliceUpdateNodeAddUpdate(builder, update) + +def SliceUpdateNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SliceUpdateNodeAddOut(builder, out) + +def SliceUpdateNodeAddAxis(builder, axis): + builder.PrependUOffsetTRelativeSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(axis), 0) + +def AddAxis(builder, axis): + SliceUpdateNodeAddAxis(builder, axis) + +def SliceUpdateNodeAddStart(builder, start): + builder.PrependUOffsetTRelativeSlot(4, flatbuffers.number_types.UOffsetTFlags.py_type(start), 0) + +def AddStart(builder, start): + SliceUpdateNodeAddStart(builder, start) + +def SliceUpdateNodeAddStop(builder, stop): + builder.PrependUOffsetTRelativeSlot(5, flatbuffers.number_types.UOffsetTFlags.py_type(stop), 0) + +def AddStop(builder, stop): + SliceUpdateNodeAddStop(builder, stop) + +def SliceUpdateNodeAddStep(builder, step): + builder.PrependInt32Slot(6, step, 1) + +def AddStep(builder, step): + SliceUpdateNodeAddStep(builder, step) + +def SliceUpdateNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SliceUpdateNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SlotType.py b/backends/mlx/serialization/_generated/mlx_delegate/SlotType.py new file mode 100644 index 00000000000..9ab785d0ac2 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SlotType.py @@ -0,0 +1,9 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +class SlotType(object): + TensorSlot = 0 + IntValueSlot = 1 + FloatValueSlot = 2 + BoolValueSlot = 3 diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SlotVariant.py b/backends/mlx/serialization/_generated/mlx_delegate/SlotVariant.py new file mode 100644 index 00000000000..e41e2d02ded --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SlotVariant.py @@ -0,0 +1,63 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SlotVariant(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SlotVariant() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSlotVariant(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SlotVariant + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SlotVariant + def Idx(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Uint32Flags, o + self._tab.Pos) + return 0 + + # SlotVariant + def SlotType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return 0 + +def SlotVariantStart(builder): + builder.StartObject(2) + +def Start(builder): + SlotVariantStart(builder) + +def SlotVariantAddIdx(builder, idx): + builder.PrependUint32Slot(0, idx, 0) + +def AddIdx(builder, idx): + SlotVariantAddIdx(builder, idx) + +def SlotVariantAddSlotType(builder, slotType): + builder.PrependInt8Slot(1, slotType, 0) + +def AddSlotType(builder, slotType): + SlotVariantAddSlotType(builder, slotType) + +def SlotVariantEnd(builder): + return builder.EndObject() + +def End(builder): + return SlotVariantEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SoftmaxNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SoftmaxNode.py new file mode 100644 index 00000000000..6f8c760671c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SoftmaxNode.py @@ -0,0 +1,97 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SoftmaxNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SoftmaxNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSoftmaxNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SoftmaxNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SoftmaxNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SoftmaxNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SoftmaxNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # SoftmaxNode + def Precise(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def SoftmaxNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + SoftmaxNodeStart(builder) + +def SoftmaxNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SoftmaxNodeAddX(builder, x) + +def SoftmaxNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SoftmaxNodeAddOut(builder, out) + +def SoftmaxNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + SoftmaxNodeAddAxis(builder, axis) + +def SoftmaxNodeAddPrecise(builder, precise): + builder.PrependBoolSlot(3, precise, 0) + +def AddPrecise(builder, precise): + SoftmaxNodeAddPrecise(builder, precise) + +def SoftmaxNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SoftmaxNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SortNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SortNode.py new file mode 100644 index 00000000000..48f98464124 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SortNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SortNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SortNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSortNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SortNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SortNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SortNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SortNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def SortNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + SortNodeStart(builder) + +def SortNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SortNodeAddX(builder, x) + +def SortNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SortNodeAddOut(builder, out) + +def SortNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + SortNodeAddAxis(builder, axis) + +def SortNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SortNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SplitNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SplitNode.py new file mode 100644 index 00000000000..5c6b3dc9f36 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SplitNode.py @@ -0,0 +1,140 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SplitNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SplitNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSplitNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SplitNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SplitNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SplitNode + def Outs(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SplitNode + def OutsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # SplitNode + def OutsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + return o == 0 + + # SplitNode + def Sizes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SplitNode + def SizesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # SplitNode + def SizesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # SplitNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def SplitNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + SplitNodeStart(builder) + +def SplitNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SplitNodeAddX(builder, x) + +def SplitNodeAddOuts(builder, outs): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(outs), 0) + +def AddOuts(builder, outs): + SplitNodeAddOuts(builder, outs) + +def SplitNodeStartOutsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartOutsVector(builder, numElems): + return SplitNodeStartOutsVector(builder, numElems) + +def SplitNodeAddSizes(builder, sizes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(sizes), 0) + +def AddSizes(builder, sizes): + SplitNodeAddSizes(builder, sizes) + +def SplitNodeStartSizesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartSizesVector(builder, numElems): + return SplitNodeStartSizesVector(builder, numElems) + +def SplitNodeAddAxis(builder, axis): + builder.PrependInt32Slot(3, axis, 0) + +def AddAxis(builder, axis): + SplitNodeAddAxis(builder, axis) + +def SplitNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SplitNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SqrtNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SqrtNode.py new file mode 100644 index 00000000000..16837e31db3 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SqrtNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SqrtNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SqrtNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSqrtNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SqrtNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SqrtNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SqrtNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SqrtNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SqrtNodeStart(builder) + +def SqrtNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SqrtNodeAddX(builder, x) + +def SqrtNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SqrtNodeAddOut(builder, out) + +def SqrtNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SqrtNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SquareNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SquareNode.py new file mode 100644 index 00000000000..c221895a73c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SquareNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SquareNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SquareNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSquareNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SquareNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SquareNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SquareNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SquareNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + SquareNodeStart(builder) + +def SquareNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SquareNodeAddX(builder, x) + +def SquareNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SquareNodeAddOut(builder, out) + +def SquareNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SquareNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SqueezeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SqueezeNode.py new file mode 100644 index 00000000000..d2ffd8dad57 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SqueezeNode.py @@ -0,0 +1,110 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SqueezeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SqueezeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSqueezeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SqueezeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SqueezeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SqueezeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SqueezeNode + def Dims(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # SqueezeNode + def DimsAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # SqueezeNode + def DimsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # SqueezeNode + def DimsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + +def SqueezeNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + SqueezeNodeStart(builder) + +def SqueezeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SqueezeNodeAddX(builder, x) + +def SqueezeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SqueezeNodeAddOut(builder, out) + +def SqueezeNodeAddDims(builder, dims): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(dims), 0) + +def AddDims(builder, dims): + SqueezeNodeAddDims(builder, dims) + +def SqueezeNodeStartDimsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartDimsVector(builder, numElems): + return SqueezeNodeStartDimsVector(builder, numElems) + +def SqueezeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SqueezeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/StackNode.py b/backends/mlx/serialization/_generated/mlx_delegate/StackNode.py new file mode 100644 index 00000000000..1eb4f5ec007 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/StackNode.py @@ -0,0 +1,103 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class StackNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = StackNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsStackNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # StackNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # StackNode + def Tensors(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # StackNode + def TensorsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # StackNode + def TensorsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + return o == 0 + + # StackNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # StackNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def StackNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + StackNodeStart(builder) + +def StackNodeAddTensors(builder, tensors): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(tensors), 0) + +def AddTensors(builder, tensors): + StackNodeAddTensors(builder, tensors) + +def StackNodeStartTensorsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartTensorsVector(builder, numElems): + return StackNodeStartTensorsVector(builder, numElems) + +def StackNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + StackNodeAddOut(builder, out) + +def StackNodeAddAxis(builder, axis): + builder.PrependInt32Slot(2, axis, 0) + +def AddAxis(builder, axis): + StackNodeAddAxis(builder, axis) + +def StackNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return StackNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/StdNode.py b/backends/mlx/serialization/_generated/mlx_delegate/StdNode.py new file mode 100644 index 00000000000..82a71a14c38 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/StdNode.py @@ -0,0 +1,136 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class StdNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = StdNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsStdNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # StdNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # StdNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # StdNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # StdNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # StdNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # StdNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # StdNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # StdNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + + # StdNode + def Ddof(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def StdNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + StdNodeStart(builder) + +def StdNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + StdNodeAddX(builder, x) + +def StdNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + StdNodeAddOut(builder, out) + +def StdNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + StdNodeAddAxes(builder, axes) + +def StdNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return StdNodeStartAxesVector(builder, numElems) + +def StdNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + StdNodeAddKeepdims(builder, keepdims) + +def StdNodeAddDdof(builder, ddof): + builder.PrependInt32Slot(4, ddof, 0) + +def AddDdof(builder, ddof): + StdNodeAddDdof(builder, ddof) + +def StdNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return StdNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SubtractIntNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SubtractIntNode.py new file mode 100644 index 00000000000..851f658c6cc --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SubtractIntNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SubtractIntNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SubtractIntNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSubtractIntNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SubtractIntNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SubtractIntNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SubtractIntNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SubtractIntNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SubtractIntNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + SubtractIntNodeStart(builder) + +def SubtractIntNodeAddA(builder, a): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + SubtractIntNodeAddA(builder, a) + +def SubtractIntNodeAddB(builder, b): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + SubtractIntNodeAddB(builder, b) + +def SubtractIntNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SubtractIntNodeAddOut(builder, out) + +def SubtractIntNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SubtractIntNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SubtractNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SubtractNode.py new file mode 100644 index 00000000000..04bd26fb71a --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SubtractNode.py @@ -0,0 +1,88 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SubtractNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SubtractNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSubtractNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SubtractNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SubtractNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SubtractNode + def B(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SubtractNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SubtractNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + SubtractNodeStart(builder) + +def SubtractNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + SubtractNodeAddA(builder, a) + +def SubtractNodeAddB(builder, b): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(b), 0) + +def AddB(builder, b): + SubtractNodeAddB(builder, b) + +def SubtractNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SubtractNodeAddOut(builder, out) + +def SubtractNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SubtractNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SumNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SumNode.py new file mode 100644 index 00000000000..5e04f3dba18 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SumNode.py @@ -0,0 +1,123 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SumNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SumNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSumNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SumNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SumNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SumNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SumNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # SumNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # SumNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # SumNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # SumNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def SumNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + SumNodeStart(builder) + +def SumNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + SumNodeAddX(builder, x) + +def SumNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SumNodeAddOut(builder, out) + +def SumNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + SumNodeAddAxes(builder, axes) + +def SumNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return SumNodeStartAxesVector(builder, numElems) + +def SumNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + SumNodeAddKeepdims(builder, keepdims) + +def SumNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SumNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/SymSizeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/SymSizeNode.py new file mode 100644 index 00000000000..5bb26b43e78 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/SymSizeNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class SymSizeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = SymSizeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsSymSizeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # SymSizeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # SymSizeNode + def A(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # SymSizeNode + def Dim(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # SymSizeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def SymSizeNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + SymSizeNodeStart(builder) + +def SymSizeNodeAddA(builder, a): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(a), 0) + +def AddA(builder, a): + SymSizeNodeAddA(builder, a) + +def SymSizeNodeAddDim(builder, dim): + builder.PrependInt32Slot(1, dim, 0) + +def AddDim(builder, dim): + SymSizeNodeAddDim(builder, dim) + +def SymSizeNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + SymSizeNodeAddOut(builder, out) + +def SymSizeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return SymSizeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TakeAlongAxisNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TakeAlongAxisNode.py new file mode 100644 index 00000000000..795146b763b --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TakeAlongAxisNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TakeAlongAxisNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TakeAlongAxisNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTakeAlongAxisNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TakeAlongAxisNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TakeAlongAxisNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TakeAlongAxisNode + def Indices(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TakeAlongAxisNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TakeAlongAxisNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def TakeAlongAxisNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + TakeAlongAxisNodeStart(builder) + +def TakeAlongAxisNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TakeAlongAxisNodeAddX(builder, x) + +def TakeAlongAxisNodeAddIndices(builder, indices): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(indices), 0) + +def AddIndices(builder, indices): + TakeAlongAxisNodeAddIndices(builder, indices) + +def TakeAlongAxisNodeAddOut(builder, out): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TakeAlongAxisNodeAddOut(builder, out) + +def TakeAlongAxisNodeAddAxis(builder, axis): + builder.PrependInt32Slot(3, axis, 0) + +def AddAxis(builder, axis): + TakeAlongAxisNodeAddAxis(builder, axis) + +def TakeAlongAxisNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TakeAlongAxisNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TakeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TakeNode.py new file mode 100644 index 00000000000..b49a9279e03 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TakeNode.py @@ -0,0 +1,101 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TakeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TakeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTakeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TakeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TakeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TakeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TakeNode + def Index(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVidOrTid import IntOrVidOrTid + obj = IntOrVidOrTid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TakeNode + def Axis(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def TakeNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + TakeNodeStart(builder) + +def TakeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TakeNodeAddX(builder, x) + +def TakeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TakeNodeAddOut(builder, out) + +def TakeNodeAddIndex(builder, index): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(index), 0) + +def AddIndex(builder, index): + TakeNodeAddIndex(builder, index) + +def TakeNodeAddAxis(builder, axis): + builder.PrependInt32Slot(3, axis, 0) + +def AddAxis(builder, axis): + TakeNodeAddAxis(builder, axis) + +def TakeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TakeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TanNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TanNode.py new file mode 100644 index 00000000000..397c2e37766 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TanNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TanNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TanNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTanNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TanNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TanNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TanNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def TanNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + TanNodeStart(builder) + +def TanNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TanNodeAddX(builder, x) + +def TanNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TanNodeAddOut(builder, out) + +def TanNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TanNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TanhNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TanhNode.py new file mode 100644 index 00000000000..806893d4954 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TanhNode.py @@ -0,0 +1,71 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TanhNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TanhNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTanhNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TanhNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TanhNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TanhNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def TanhNodeStart(builder): + builder.StartObject(2) + +def Start(builder): + TanhNodeStart(builder) + +def TanhNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TanhNodeAddX(builder, x) + +def TanhNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TanhNodeAddOut(builder, out) + +def TanhNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TanhNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TensorMeta.py b/backends/mlx/serialization/_generated/mlx_delegate/TensorMeta.py new file mode 100644 index 00000000000..2e0516b8006 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TensorMeta.py @@ -0,0 +1,126 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TensorMeta(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TensorMeta() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTensorMeta(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TensorMeta + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TensorMeta + def Shape(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.ShapeDim import ShapeDim + obj = ShapeDim() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TensorMeta + def ShapeLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # TensorMeta + def ShapeIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + return o == 0 + + # TensorMeta + def ScalarType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return 0 + + # TensorMeta + def DimOrder(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Uint8Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 1)) + return 0 + + # TensorMeta + def DimOrderAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Uint8Flags, o) + return 0 + + # TensorMeta + def DimOrderLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # TensorMeta + def DimOrderIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + +def TensorMetaStart(builder): + builder.StartObject(3) + +def Start(builder): + TensorMetaStart(builder) + +def TensorMetaAddShape(builder, shape): + builder.PrependUOffsetTRelativeSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(shape), 0) + +def AddShape(builder, shape): + TensorMetaAddShape(builder, shape) + +def TensorMetaStartShapeVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartShapeVector(builder, numElems): + return TensorMetaStartShapeVector(builder, numElems) + +def TensorMetaAddScalarType(builder, scalarType): + builder.PrependInt8Slot(1, scalarType, 0) + +def AddScalarType(builder, scalarType): + TensorMetaAddScalarType(builder, scalarType) + +def TensorMetaAddDimOrder(builder, dimOrder): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(dimOrder), 0) + +def AddDimOrder(builder, dimOrder): + TensorMetaAddDimOrder(builder, dimOrder) + +def TensorMetaStartDimOrderVector(builder, numElems): + return builder.StartVector(1, numElems, 1) + +def StartDimOrderVector(builder, numElems): + return TensorMetaStartDimOrderVector(builder, numElems) + +def TensorMetaEnd(builder): + return builder.EndObject() + +def End(builder): + return TensorMetaEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Tid.py b/backends/mlx/serialization/_generated/mlx_delegate/Tid.py new file mode 100644 index 00000000000..5189bb402b4 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Tid.py @@ -0,0 +1,26 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Tid(object): + __slots__ = ['_tab'] + + @classmethod + def SizeOf(cls): + return 4 + + # Tid + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Tid + def Idx(self): return self._tab.Get(flatbuffers.number_types.Uint32Flags, self._tab.Pos + flatbuffers.number_types.UOffsetTFlags.py_type(0)) + +def CreateTid(builder, idx): + builder.Prep(4, 4) + builder.PrependUint32(idx) + return builder.Offset() diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TileNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TileNode.py new file mode 100644 index 00000000000..22e4a7c2ace --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TileNode.py @@ -0,0 +1,108 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TileNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TileNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTileNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TileNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TileNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TileNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TileNode + def Reps(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Vector(o) + x += flatbuffers.number_types.UOffsetTFlags.py_type(j) * 4 + x = self._tab.Indirect(x) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TileNode + def RepsLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # TileNode + def RepsIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + +def TileNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + TileNodeStart(builder) + +def TileNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TileNodeAddX(builder, x) + +def TileNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TileNodeAddOut(builder, out) + +def TileNodeAddReps(builder, reps): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(reps), 0) + +def AddReps(builder, reps): + TileNodeAddReps(builder, reps) + +def TileNodeStartRepsVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartRepsVector(builder, numElems): + return TileNodeStartRepsVector(builder, numElems) + +def TileNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TileNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TransposeNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TransposeNode.py new file mode 100644 index 00000000000..7034105b099 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TransposeNode.py @@ -0,0 +1,110 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TransposeNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TransposeNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTransposeNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TransposeNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TransposeNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TransposeNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TransposeNode + def Perm(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # TransposeNode + def PermAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # TransposeNode + def PermLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # TransposeNode + def PermIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + +def TransposeNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + TransposeNodeStart(builder) + +def TransposeNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TransposeNodeAddX(builder, x) + +def TransposeNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TransposeNodeAddOut(builder, out) + +def TransposeNodeAddPerm(builder, perm): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(perm), 0) + +def AddPerm(builder, perm): + TransposeNodeAddPerm(builder, perm) + +def TransposeNodeStartPermVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartPermVector(builder, numElems): + return TransposeNodeStartPermVector(builder, numElems) + +def TransposeNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TransposeNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TriNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TriNode.py new file mode 100644 index 00000000000..694a46a8ee9 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TriNode.py @@ -0,0 +1,114 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TriNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TriNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTriNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TriNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TriNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TriNode + def N(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TriNode + def M(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = self._tab.Indirect(o + self._tab.Pos) + from executorch.backends.mlx.serialization._generated.mlx_delegate.IntOrVid import IntOrVid + obj = IntOrVid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TriNode + def K(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + + # TriNode + def ScalarType(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int8Flags, o + self._tab.Pos) + return 0 + +def TriNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + TriNodeStart(builder) + +def TriNodeAddOut(builder, out): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TriNodeAddOut(builder, out) + +def TriNodeAddN(builder, n): + builder.PrependUOffsetTRelativeSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(n), 0) + +def AddN(builder, n): + TriNodeAddN(builder, n) + +def TriNodeAddM(builder, m): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(m), 0) + +def AddM(builder, m): + TriNodeAddM(builder, m) + +def TriNodeAddK(builder, k): + builder.PrependInt32Slot(3, k, 0) + +def AddK(builder, k): + TriNodeAddK(builder, k) + +def TriNodeAddScalarType(builder, scalarType): + builder.PrependInt8Slot(4, scalarType, 0) + +def AddScalarType(builder, scalarType): + TriNodeAddScalarType(builder, scalarType) + +def TriNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TriNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TrilNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TrilNode.py new file mode 100644 index 00000000000..7c8a5a70362 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TrilNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TrilNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TrilNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTrilNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TrilNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TrilNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TrilNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TrilNode + def K(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def TrilNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + TrilNodeStart(builder) + +def TrilNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TrilNodeAddX(builder, x) + +def TrilNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TrilNodeAddOut(builder, out) + +def TrilNodeAddK(builder, k): + builder.PrependInt32Slot(2, k, 0) + +def AddK(builder, k): + TrilNodeAddK(builder, k) + +def TrilNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TrilNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/TriuNode.py b/backends/mlx/serialization/_generated/mlx_delegate/TriuNode.py new file mode 100644 index 00000000000..9910c554f9c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/TriuNode.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class TriuNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = TriuNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsTriuNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # TriuNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # TriuNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TriuNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # TriuNode + def K(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def TriuNodeStart(builder): + builder.StartObject(3) + +def Start(builder): + TriuNodeStart(builder) + +def TriuNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + TriuNodeAddX(builder, x) + +def TriuNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + TriuNodeAddOut(builder, out) + +def TriuNodeAddK(builder, k): + builder.PrependInt32Slot(2, k, 0) + +def AddK(builder, k): + TriuNodeAddK(builder, k) + +def TriuNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return TriuNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/VarNode.py b/backends/mlx/serialization/_generated/mlx_delegate/VarNode.py new file mode 100644 index 00000000000..02ebd687415 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/VarNode.py @@ -0,0 +1,136 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class VarNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = VarNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsVarNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # VarNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # VarNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # VarNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # VarNode + def Axes(self, j): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + a = self._tab.Vector(o) + return self._tab.Get(flatbuffers.number_types.Int32Flags, a + flatbuffers.number_types.UOffsetTFlags.py_type(j * 4)) + return 0 + + # VarNode + def AxesAsNumpy(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.GetVectorAsNumpy(flatbuffers.number_types.Int32Flags, o) + return 0 + + # VarNode + def AxesLength(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return self._tab.VectorLen(o) + return 0 + + # VarNode + def AxesIsNone(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + return o == 0 + + # VarNode + def Keepdims(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + + # VarNode + def Ddof(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(12)) + if o != 0: + return self._tab.Get(flatbuffers.number_types.Int32Flags, o + self._tab.Pos) + return 0 + +def VarNodeStart(builder): + builder.StartObject(5) + +def Start(builder): + VarNodeStart(builder) + +def VarNodeAddX(builder, x): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + VarNodeAddX(builder, x) + +def VarNodeAddOut(builder, out): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + VarNodeAddOut(builder, out) + +def VarNodeAddAxes(builder, axes): + builder.PrependUOffsetTRelativeSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(axes), 0) + +def AddAxes(builder, axes): + VarNodeAddAxes(builder, axes) + +def VarNodeStartAxesVector(builder, numElems): + return builder.StartVector(4, numElems, 4) + +def StartAxesVector(builder, numElems): + return VarNodeStartAxesVector(builder, numElems) + +def VarNodeAddKeepdims(builder, keepdims): + builder.PrependBoolSlot(3, keepdims, 0) + +def AddKeepdims(builder, keepdims): + VarNodeAddKeepdims(builder, keepdims) + +def VarNodeAddDdof(builder, ddof): + builder.PrependInt32Slot(4, ddof, 0) + +def AddDdof(builder, ddof): + VarNodeAddDdof(builder, ddof) + +def VarNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return VarNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/Vid.py b/backends/mlx/serialization/_generated/mlx_delegate/Vid.py new file mode 100644 index 00000000000..6525d23b92c --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/Vid.py @@ -0,0 +1,26 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class Vid(object): + __slots__ = ['_tab'] + + @classmethod + def SizeOf(cls): + return 4 + + # Vid + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # Vid + def Idx(self): return self._tab.Get(flatbuffers.number_types.Uint32Flags, self._tab.Pos + flatbuffers.number_types.UOffsetTFlags.py_type(0)) + +def CreateVid(builder, idx): + builder.Prep(4, 4) + builder.PrependUint32(idx) + return builder.Offset() diff --git a/backends/mlx/serialization/_generated/mlx_delegate/VidOrTid.py b/backends/mlx/serialization/_generated/mlx_delegate/VidOrTid.py new file mode 100644 index 00000000000..68bdb586e25 --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/VidOrTid.py @@ -0,0 +1,84 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class VidOrTid(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = VidOrTid() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsVidOrTid(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # VidOrTid + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # VidOrTid + def Vid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import Vid + obj = Vid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # VidOrTid + def Tid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # VidOrTid + def IsVid(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + return bool(self._tab.Get(flatbuffers.number_types.BoolFlags, o + self._tab.Pos)) + return False + +def VidOrTidStart(builder): + builder.StartObject(3) + +def Start(builder): + VidOrTidStart(builder) + +def VidOrTidAddVid(builder, vid): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(vid), 0) + +def AddVid(builder, vid): + VidOrTidAddVid(builder, vid) + +def VidOrTidAddTid(builder, tid): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(tid), 0) + +def AddTid(builder, tid): + VidOrTidAddTid(builder, tid) + +def VidOrTidAddIsVid(builder, isVid): + builder.PrependBoolSlot(2, isVid, 0) + +def AddIsVid(builder, isVid): + VidOrTidAddIsVid(builder, isVid) + +def VidOrTidEnd(builder): + return builder.EndObject() + +def End(builder): + return VidOrTidEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/WhereNode.py b/backends/mlx/serialization/_generated/mlx_delegate/WhereNode.py new file mode 100644 index 00000000000..5338b51d66f --- /dev/null +++ b/backends/mlx/serialization/_generated/mlx_delegate/WhereNode.py @@ -0,0 +1,105 @@ +# automatically generated by the FlatBuffers compiler, do not modify + +# namespace: mlx_delegate + +import flatbuffers +from flatbuffers.compat import import_numpy +np = import_numpy() + +class WhereNode(object): + __slots__ = ['_tab'] + + @classmethod + def GetRootAs(cls, buf, offset=0): + n = flatbuffers.encode.Get(flatbuffers.packer.uoffset, buf, offset) + x = WhereNode() + x.Init(buf, n + offset) + return x + + @classmethod + def GetRootAsWhereNode(cls, buf, offset=0): + """This method is deprecated. Please switch to GetRootAs.""" + return cls.GetRootAs(buf, offset) + # WhereNode + def Init(self, buf, pos): + self._tab = flatbuffers.table.Table(buf, pos) + + # WhereNode + def Condition(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(4)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # WhereNode + def X(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(6)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # WhereNode + def Y(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(8)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + + # WhereNode + def Out(self): + o = flatbuffers.number_types.UOffsetTFlags.py_type(self._tab.Offset(10)) + if o != 0: + x = o + self._tab.Pos + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import Tid + obj = Tid() + obj.Init(self._tab.Bytes, x) + return obj + return None + +def WhereNodeStart(builder): + builder.StartObject(4) + +def Start(builder): + WhereNodeStart(builder) + +def WhereNodeAddCondition(builder, condition): + builder.PrependStructSlot(0, flatbuffers.number_types.UOffsetTFlags.py_type(condition), 0) + +def AddCondition(builder, condition): + WhereNodeAddCondition(builder, condition) + +def WhereNodeAddX(builder, x): + builder.PrependStructSlot(1, flatbuffers.number_types.UOffsetTFlags.py_type(x), 0) + +def AddX(builder, x): + WhereNodeAddX(builder, x) + +def WhereNodeAddY(builder, y): + builder.PrependStructSlot(2, flatbuffers.number_types.UOffsetTFlags.py_type(y), 0) + +def AddY(builder, y): + WhereNodeAddY(builder, y) + +def WhereNodeAddOut(builder, out): + builder.PrependStructSlot(3, flatbuffers.number_types.UOffsetTFlags.py_type(out), 0) + +def AddOut(builder, out): + WhereNodeAddOut(builder, out) + +def WhereNodeEnd(builder): + return builder.EndObject() + +def End(builder): + return WhereNodeEnd(builder) diff --git a/backends/mlx/serialization/_generated/mlx_delegate/__init__.py b/backends/mlx/serialization/_generated/mlx_delegate/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/backends/mlx/serialization/_generated_serializers.py b/backends/mlx/serialization/_generated_serializers.py new file mode 100644 index 00000000000..40681ceb02f --- /dev/null +++ b/backends/mlx/serialization/_generated_serializers.py @@ -0,0 +1,2987 @@ +# +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. +# +# ============================================================================ +# AUTO-GENERATED FILE - DO NOT EDIT MANUALLY +# ============================================================================ +# +# This file was generated from schema.fbs by the MLX delegate code generator. +# +# Source: backends/mlx/serialization/schema.fbs +# Generator: backends/mlx/serialization/generate.py +# +# To regenerate, run from the executorch root: +# python backends/mlx/serialization/generate.py +# +# ============================================================================ +# +# This file contains auto-generated serializer methods for all op types. + +from __future__ import annotations + +from typing import List, Tuple, Dict + +import flatbuffers + +# FlatBuffer union indices: 0 = NONE, then 1-indexed from union order +MLX_OP_TYPE_NAMES = { + 0: "NONE", + 1: "NoopNode", + 2: "IdCopyNode", + 3: "AddmmNode", + 4: "ItemIntNode", + 5: "ExpandDimsNode", + 6: "TileNode", + 7: "TakeAlongAxisNode", + 8: "TakeNode", + 9: "RMSNormNode", + 10: "LayerNormNode", + 11: "RopeNode", + 12: "SdpaNode", + 13: "AddNode", + 14: "AddIntNode", + 15: "SubtractIntNode", + 16: "MultiplyIntNode", + 17: "FloorDivideIntNode", + 18: "SymSizeNode", + 19: "MultiplyNode", + 20: "DivideNode", + 21: "SubtractNode", + 22: "Conv1DNode", + 23: "Conv2DNode", + 24: "Conv3DNode", + 25: "GeluNode", + 26: "ARangeNode", + 27: "SiluNode", + 28: "SigmoidNode", + 29: "TanhNode", + 30: "SqueezeNode", + 31: "SplitNode", + 32: "RsqrtNode", + 33: "MaximumNode", + 34: "MinimumNode", + 35: "LogNode", + 36: "SoftmaxNode", + 37: "BroadcastToNode", + 38: "PadNode", + 39: "WhereNode", + 40: "ReshapeNode", + 41: "TransposeNode", + 42: "AsStridedNode", + 43: "ContiguousNode", + 44: "GatherNode", + 45: "SliceNode", + 46: "AsTypeNode", + 47: "ConcatenateNode", + 48: "FullNode", + 49: "FullLikeNode", + 50: "ArgmaxNode", + 51: "SliceUpdateNode", + 52: "IndexCopyNode", + 53: "DequantizeNode", + 54: "LessNode", + 55: "LessEqualNode", + 56: "GreaterNode", + 57: "GreaterEqualNode", + 58: "EqualNode", + 59: "NotEqualNode", + 60: "LogicalNotNode", + 61: "LogicalAndNode", + 62: "LogicalOrNode", + 63: "TriNode", + 64: "TrilNode", + 65: "TriuNode", + 66: "FloorNode", + 67: "CeilNode", + 68: "SquareNode", + 69: "ExpNode", + 70: "SinNode", + 71: "CosNode", + 72: "TanNode", + 73: "ArcsinNode", + 74: "ArccosNode", + 75: "ArctanNode", + 76: "SinhNode", + 77: "CoshNode", + 78: "ArcsinhNode", + 79: "ArccoshNode", + 80: "ArctanhNode", + 81: "Log2Node", + 82: "Log10Node", + 83: "Log1pNode", + 84: "ErfNode", + 85: "Expm1Node", + 86: "RoundNode", + 87: "ReciprocalNode", + 88: "SqrtNode", + 89: "AbsNode", + 90: "NegNode", + 91: "Atan2Node", + 92: "LogAddExpNode", + 93: "FloorDivideNode", + 94: "PowerNode", + 95: "LogSumExpNode", + 96: "SumNode", + 97: "MeanNode", + 98: "VarNode", + 99: "StdNode", + 100: "ProdNode", + 101: "MaxNode", + 102: "MinNode", + 103: "ArgminNode", + 104: "MedianNode", + 105: "ModIntNode", + 106: "RemainderNode", + 107: "ConvTranspose1DNode", + 108: "ConvTranspose2DNode", + 109: "ConvTranspose3DNode", + 110: "ClipNode", + 111: "CumsumNode", + 112: "StackNode", + 113: "SignNode", + 114: "AnyNode", + 115: "AllNode", + 116: "RepeatNode", + 117: "SortNode", + 118: "ArgsortNode", + 119: "PartitionNode", + 120: "ArgPartitionNode", + 121: "QuantizedMatmulNode", + 122: "ScatterAddNode", + 123: "GatherMmNode", + 124: "GatherQmmNode", + 125: "ScanNode", + 126: "MetalKernelNode", + 127: "BitwiseInvertNode", + 128: "RollNode", + 129: "BitwiseAndNode", + 130: "BitwiseOrNode", + 131: "BitwiseXorNode", + 132: "IfNode", + 133: "RandomBitsNode", +} + +from executorch.backends.mlx.serialization.mlx_graph_schema import ( + NoopNode, + IdCopyNode, + AddmmNode, + ItemIntNode, + ExpandDimsNode, + TileNode, + TakeAlongAxisNode, + TakeNode, + RMSNormNode, + LayerNormNode, + RopeNode, + SdpaNode, + AddNode, + AddIntNode, + SubtractIntNode, + MultiplyIntNode, + FloorDivideIntNode, + ModIntNode, + SymSizeNode, + MultiplyNode, + DivideNode, + SubtractNode, + Conv1DNode, + Conv2DNode, + Conv3DNode, + ConvTranspose1DNode, + ConvTranspose2DNode, + ConvTranspose3DNode, + GeluNode, + ARangeNode, + SiluNode, + SigmoidNode, + TanhNode, + SqueezeNode, + SplitNode, + RsqrtNode, + MaximumNode, + MinimumNode, + LogNode, + SoftmaxNode, + BroadcastToNode, + PadNode, + WhereNode, + ReshapeNode, + TransposeNode, + AsStridedNode, + ContiguousNode, + GatherNode, + SliceNode, + AsTypeNode, + QuantizedMatmulNode, + ScatterAddNode, + ConcatenateNode, + FullNode, + FullLikeNode, + ArgmaxNode, + SliceUpdateNode, + IndexCopyNode, + DequantizeNode, + LessNode, + LessEqualNode, + GreaterNode, + GreaterEqualNode, + EqualNode, + NotEqualNode, + LogicalNotNode, + BitwiseInvertNode, + LogicalAndNode, + LogicalOrNode, + BitwiseAndNode, + BitwiseOrNode, + BitwiseXorNode, + TriNode, + TrilNode, + TriuNode, + ClipNode, + CumsumNode, + StackNode, + SignNode, + AnyNode, + AllNode, + RepeatNode, + SortNode, + ArgsortNode, + PartitionNode, + ArgPartitionNode, + RollNode, + FloorNode, + CeilNode, + SquareNode, + ExpNode, + SinNode, + CosNode, + TanNode, + ArcsinNode, + ArccosNode, + ArctanNode, + SinhNode, + CoshNode, + ArcsinhNode, + ArccoshNode, + ArctanhNode, + Log2Node, + Log10Node, + Log1pNode, + ErfNode, + Expm1Node, + RoundNode, + ReciprocalNode, + SqrtNode, + AbsNode, + NegNode, + Atan2Node, + LogAddExpNode, + FloorDivideNode, + RemainderNode, + PowerNode, + LogSumExpNode, + SumNode, + MeanNode, + VarNode, + StdNode, + ProdNode, + MaxNode, + MinNode, + ArgminNode, + MedianNode, + GatherMmNode, + GatherQmmNode, + ScanNode, + IfNode, + RandomBitsNode, + MetalKernelNode, + IntOrVid, + FloatOrVid, + VidOrTid, + IntOrVidOrTid, + Tid, + Vid, +) + + +def _build_int_vector(builder: flatbuffers.Builder, vec: List[int]) -> int: + """Pre-build a vector of int32 values (must be called before table Start).""" + builder.StartVector(4, len(vec), 4) + for v in reversed(vec): + builder.PrependInt32(v) + return builder.EndVector() + + +def _build_int8_vector(builder: flatbuffers.Builder, vec: List[int]) -> int: + """Pre-build a vector of int8 values (must be called before table Start).""" + builder.StartVector(1, len(vec), 1) + for v in reversed(vec): + builder.PrependInt8(v) + return builder.EndVector() + + +def _build_uint8_vector(builder: flatbuffers.Builder, vec: List[int]) -> int: + """Pre-build a vector of uint8 values (must be called before table Start).""" + builder.StartVector(1, len(vec), 1) + for v in reversed(vec): + builder.PrependUint8(v) + return builder.EndVector() + + +def _shared_string(builder: flatbuffers.Builder, s): + """CreateString with per-buffer dedup so identical strings share one offset.""" + if s is None: + return None + # flatbuffers' Builder dedups identical strings via its built-in + # sharedStrings cache; fall back to CreateString on old flatbuffers. + create = getattr(builder, "CreateSharedString", None) or builder.CreateString + return create(s) + + +class GeneratedOpBuilders: + """Mixin class with auto-generated op builder methods.""" + + def _build_int_or_vid(self, builder: flatbuffers.Builder, iov: IntOrVid) -> int: + """Build an IntOrVid table.""" + from executorch.backends.mlx.serialization._generated.mlx_delegate import IntOrVid as FBIntOrVidModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBIntOrVidModule.Start(builder) + FBIntOrVidModule.AddLiteral(builder, iov.literal) + FBIntOrVidModule.AddIsVid(builder, iov.is_vid) + if iov.vid is not None: + # Vid is an inline struct - must be added last for proper FlatBuffer layout + FBIntOrVidModule.AddVid(builder, CreateVid(builder, iov.vid.idx)) + return FBIntOrVidModule.End(builder) + + def _build_float_or_vid(self, builder: flatbuffers.Builder, fov: FloatOrVid) -> int: + """Build a FloatOrVid table.""" + from executorch.backends.mlx.serialization._generated.mlx_delegate import FloatOrVid as FBFloatOrVidModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBFloatOrVidModule.Start(builder) + FBFloatOrVidModule.AddLiteral(builder, fov.literal) + FBFloatOrVidModule.AddIsVid(builder, fov.is_vid) + if fov.vid is not None: + FBFloatOrVidModule.AddVid(builder, CreateVid(builder, fov.vid.idx)) + return FBFloatOrVidModule.End(builder) + + def _build_vid_or_tid(self, builder: flatbuffers.Builder, vot: VidOrTid) -> int: + """Build a TidOrVid table.""" + from executorch.backends.mlx.serialization._generated.mlx_delegate import VidOrTid as FBVidOrTidModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBVidOrTidModule.Start(builder) + FBVidOrTidModule.AddIsVid(builder, vot.is_vid) + if vot.tid is not None: + FBVidOrTidModule.AddTid(builder, CreateTid(builder, vot.tid.idx)) + if vot.vid is not None: + FBVidOrTidModule.AddVid(builder, CreateVid(builder, vot.vid.idx)) + return FBVidOrTidModule.End(builder) + + def _build_int_or_vid_or_tid(self, builder: flatbuffers.Builder, ivt: IntOrVidOrTid) -> int: + """Build an IntOrVidOrTid table.""" + from executorch.backends.mlx.serialization._generated.mlx_delegate import IntOrVidOrTid as FBIntOrVidOrTidModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBIntOrVidOrTidModule.Start(builder) + FBIntOrVidOrTidModule.AddLiteral(builder, ivt.literal) + FBIntOrVidOrTidModule.AddKind(builder, ivt.kind) + if ivt.tid is not None: + FBIntOrVidOrTidModule.AddTid(builder, CreateTid(builder, ivt.tid.idx)) + if ivt.vid is not None: + FBIntOrVidOrTidModule.AddVid(builder, CreateVid(builder, ivt.vid.idx)) + return FBIntOrVidOrTidModule.End(builder) + + def _build_int_or_vid_vector( + self, builder: flatbuffers.Builder, vec: List[IntOrVid] + ) -> int: + """Build a vector of IntOrVid tables.""" + offsets = [] + for iov in vec: + offsets.append(self._build_int_or_vid(builder, iov)) + builder.StartVector(4, len(offsets), 4) + for off in reversed(offsets): + builder.PrependUOffsetTRelative(off) + return builder.EndVector() + + def _build_tid_vector( + self, builder: flatbuffers.Builder, vec: List[Tid] + ) -> int: + """Build a vector of Tid structs.""" + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + + # For vectors of structs, we need to build the vector differently + # Each Tid struct is 4 bytes (uint32), so we manually write them + builder.StartVector(4, len(vec), 4) + for tid in reversed(vec): + builder.Prep(4, 0) # Align for struct + builder.PrependUint32(tid.idx) + return builder.EndVector() + + def _build_string_vector( + self, builder: flatbuffers.Builder, vec: List[str] + ) -> int: + """Pre-build a vector of strings (offsets must be created before table Start).""" + offsets = [_shared_string(builder, s) for s in vec] + builder.StartVector(4, len(offsets), 4) + for off in reversed(offsets): + builder.PrependUOffsetTRelative(off) + return builder.EndVector() + + def _build_NoopNode( + self, builder: flatbuffers.Builder, op: NoopNode + ) -> Tuple[int, int]: + """Auto-generated builder for NoopNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import NoopNode as FBNoopNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBNoopNodeModule.Start(builder) + offset = FBNoopNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.NoopNode + + def _build_IdCopyNode( + self, builder: flatbuffers.Builder, op: IdCopyNode + ) -> Tuple[int, int]: + """Auto-generated builder for IdCopyNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import IdCopyNode as FBIdCopyNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBIdCopyNodeModule.Start(builder) + FBIdCopyNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBIdCopyNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBIdCopyNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.IdCopyNode + + def _build_AddmmNode( + self, builder: flatbuffers.Builder, op: AddmmNode + ) -> Tuple[int, int]: + """Auto-generated builder for AddmmNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AddmmNode as FBAddmmNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBAddmmNodeModule.Start(builder) + FBAddmmNodeModule.AddMat1(builder, CreateTid(builder, op.mat1.idx)) + FBAddmmNodeModule.AddMat2(builder, CreateTid(builder, op.mat2.idx)) + FBAddmmNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if op.bias is not None: + FBAddmmNodeModule.AddBias(builder, CreateTid(builder, op.bias.idx)) + FBAddmmNodeModule.AddAlpha(builder, op.alpha) + FBAddmmNodeModule.AddBeta(builder, op.beta) + offset = FBAddmmNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AddmmNode + + def _build_ItemIntNode( + self, builder: flatbuffers.Builder, op: ItemIntNode + ) -> Tuple[int, int]: + """Auto-generated builder for ItemIntNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ItemIntNode as FBItemIntNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBItemIntNodeModule.Start(builder) + FBItemIntNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBItemIntNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBItemIntNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ItemIntNode + + def _build_ExpandDimsNode( + self, builder: flatbuffers.Builder, op: ExpandDimsNode + ) -> Tuple[int, int]: + """Auto-generated builder for ExpandDimsNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ExpandDimsNode as FBExpandDimsNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBExpandDimsNodeModule.Start(builder) + FBExpandDimsNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBExpandDimsNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBExpandDimsNodeModule.AddAxis(builder, op.axis) + offset = FBExpandDimsNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ExpandDimsNode + + def _build_TileNode( + self, builder: flatbuffers.Builder, op: TileNode + ) -> Tuple[int, int]: + """Auto-generated builder for TileNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TileNode as FBTileNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + reps_vec = self._build_int_or_vid_vector(builder, op.reps) + + FBTileNodeModule.Start(builder) + FBTileNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTileNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTileNodeModule.AddReps(builder, reps_vec) + offset = FBTileNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TileNode + + def _build_TakeAlongAxisNode( + self, builder: flatbuffers.Builder, op: TakeAlongAxisNode + ) -> Tuple[int, int]: + """Auto-generated builder for TakeAlongAxisNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TakeAlongAxisNode as FBTakeAlongAxisNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBTakeAlongAxisNodeModule.Start(builder) + FBTakeAlongAxisNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTakeAlongAxisNodeModule.AddIndices(builder, CreateTid(builder, op.indices.idx)) + FBTakeAlongAxisNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTakeAlongAxisNodeModule.AddAxis(builder, op.axis) + offset = FBTakeAlongAxisNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TakeAlongAxisNode + + def _build_TakeNode( + self, builder: flatbuffers.Builder, op: TakeNode + ) -> Tuple[int, int]: + """Auto-generated builder for TakeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TakeNode as FBTakeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + index_off = self._build_int_or_vid_or_tid(builder, op.index) + + FBTakeNodeModule.Start(builder) + FBTakeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTakeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTakeNodeModule.AddIndex(builder, index_off) + FBTakeNodeModule.AddAxis(builder, op.axis) + offset = FBTakeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TakeNode + + def _build_RMSNormNode( + self, builder: flatbuffers.Builder, op: RMSNormNode + ) -> Tuple[int, int]: + """Auto-generated builder for RMSNormNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RMSNormNode as FBRMSNormNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBRMSNormNodeModule.Start(builder) + FBRMSNormNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + if op.weight is not None: + FBRMSNormNodeModule.AddWeight(builder, CreateTid(builder, op.weight.idx)) + FBRMSNormNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBRMSNormNodeModule.AddEps(builder, op.eps) + offset = FBRMSNormNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RMSNormNode + + def _build_LayerNormNode( + self, builder: flatbuffers.Builder, op: LayerNormNode + ) -> Tuple[int, int]: + """Auto-generated builder for LayerNormNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LayerNormNode as FBLayerNormNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLayerNormNodeModule.Start(builder) + FBLayerNormNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLayerNormNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if op.weight is not None: + FBLayerNormNodeModule.AddWeight(builder, CreateTid(builder, op.weight.idx)) + if op.bias is not None: + FBLayerNormNodeModule.AddBias(builder, CreateTid(builder, op.bias.idx)) + FBLayerNormNodeModule.AddEps(builder, op.eps) + offset = FBLayerNormNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LayerNormNode + + def _build_RopeNode( + self, builder: flatbuffers.Builder, op: RopeNode + ) -> Tuple[int, int]: + """Auto-generated builder for RopeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RopeNode as FBRopeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + offset_off = self._build_vid_or_tid(builder, op.offset) + + FBRopeNodeModule.Start(builder) + FBRopeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBRopeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBRopeNodeModule.AddDims(builder, op.dims) + FBRopeNodeModule.AddOffset(builder, offset_off) + if op.freqs is not None: + FBRopeNodeModule.AddFreqs(builder, CreateTid(builder, op.freqs.idx)) + FBRopeNodeModule.AddTraditional(builder, op.traditional) + FBRopeNodeModule.AddBase(builder, op.base) + FBRopeNodeModule.AddScale(builder, op.scale) + offset = FBRopeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RopeNode + + def _build_SdpaNode( + self, builder: flatbuffers.Builder, op: SdpaNode + ) -> Tuple[int, int]: + """Auto-generated builder for SdpaNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SdpaNode as FBSdpaNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSdpaNodeModule.Start(builder) + FBSdpaNodeModule.AddQ(builder, CreateTid(builder, op.q.idx)) + FBSdpaNodeModule.AddK(builder, CreateTid(builder, op.k.idx)) + FBSdpaNodeModule.AddV(builder, CreateTid(builder, op.v.idx)) + FBSdpaNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBSdpaNodeModule.AddScale(builder, op.scale) + if op.mask is not None: + FBSdpaNodeModule.AddMask(builder, CreateTid(builder, op.mask.idx)) + FBSdpaNodeModule.AddCausal(builder, op.causal) + offset = FBSdpaNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SdpaNode + + def _build_AddNode( + self, builder: flatbuffers.Builder, op: AddNode + ) -> Tuple[int, int]: + """Auto-generated builder for AddNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AddNode as FBAddNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBAddNodeModule.Start(builder) + FBAddNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBAddNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBAddNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBAddNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AddNode + + def _build_AddIntNode( + self, builder: flatbuffers.Builder, op: AddIntNode + ) -> Tuple[int, int]: + """Auto-generated builder for AddIntNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AddIntNode as FBAddIntNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + a_off = self._build_int_or_vid(builder, op.a) + b_off = self._build_int_or_vid(builder, op.b) + + FBAddIntNodeModule.Start(builder) + FBAddIntNodeModule.AddA(builder, a_off) + FBAddIntNodeModule.AddB(builder, b_off) + FBAddIntNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBAddIntNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AddIntNode + + def _build_SubtractIntNode( + self, builder: flatbuffers.Builder, op: SubtractIntNode + ) -> Tuple[int, int]: + """Auto-generated builder for SubtractIntNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SubtractIntNode as FBSubtractIntNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + a_off = self._build_int_or_vid(builder, op.a) + b_off = self._build_int_or_vid(builder, op.b) + + FBSubtractIntNodeModule.Start(builder) + FBSubtractIntNodeModule.AddA(builder, a_off) + FBSubtractIntNodeModule.AddB(builder, b_off) + FBSubtractIntNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBSubtractIntNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SubtractIntNode + + def _build_MultiplyIntNode( + self, builder: flatbuffers.Builder, op: MultiplyIntNode + ) -> Tuple[int, int]: + """Auto-generated builder for MultiplyIntNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MultiplyIntNode as FBMultiplyIntNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + a_off = self._build_int_or_vid(builder, op.a) + b_off = self._build_int_or_vid(builder, op.b) + + FBMultiplyIntNodeModule.Start(builder) + FBMultiplyIntNodeModule.AddA(builder, a_off) + FBMultiplyIntNodeModule.AddB(builder, b_off) + FBMultiplyIntNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBMultiplyIntNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MultiplyIntNode + + def _build_FloorDivideIntNode( + self, builder: flatbuffers.Builder, op: FloorDivideIntNode + ) -> Tuple[int, int]: + """Auto-generated builder for FloorDivideIntNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import FloorDivideIntNode as FBFloorDivideIntNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + a_off = self._build_int_or_vid(builder, op.a) + b_off = self._build_int_or_vid(builder, op.b) + + FBFloorDivideIntNodeModule.Start(builder) + FBFloorDivideIntNodeModule.AddA(builder, a_off) + FBFloorDivideIntNodeModule.AddB(builder, b_off) + FBFloorDivideIntNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBFloorDivideIntNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.FloorDivideIntNode + + def _build_ModIntNode( + self, builder: flatbuffers.Builder, op: ModIntNode + ) -> Tuple[int, int]: + """Auto-generated builder for ModIntNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ModIntNode as FBModIntNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + a_off = self._build_int_or_vid(builder, op.a) + b_off = self._build_int_or_vid(builder, op.b) + + FBModIntNodeModule.Start(builder) + FBModIntNodeModule.AddA(builder, a_off) + FBModIntNodeModule.AddB(builder, b_off) + FBModIntNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBModIntNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ModIntNode + + def _build_SymSizeNode( + self, builder: flatbuffers.Builder, op: SymSizeNode + ) -> Tuple[int, int]: + """Auto-generated builder for SymSizeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SymSizeNode as FBSymSizeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSymSizeNodeModule.Start(builder) + FBSymSizeNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBSymSizeNodeModule.AddDim(builder, op.dim) + FBSymSizeNodeModule.AddOut(builder, CreateVid(builder, op.out.idx)) + offset = FBSymSizeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SymSizeNode + + def _build_MultiplyNode( + self, builder: flatbuffers.Builder, op: MultiplyNode + ) -> Tuple[int, int]: + """Auto-generated builder for MultiplyNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MultiplyNode as FBMultiplyNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBMultiplyNodeModule.Start(builder) + FBMultiplyNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBMultiplyNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBMultiplyNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBMultiplyNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MultiplyNode + + def _build_DivideNode( + self, builder: flatbuffers.Builder, op: DivideNode + ) -> Tuple[int, int]: + """Auto-generated builder for DivideNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import DivideNode as FBDivideNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBDivideNodeModule.Start(builder) + FBDivideNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBDivideNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBDivideNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBDivideNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.DivideNode + + def _build_SubtractNode( + self, builder: flatbuffers.Builder, op: SubtractNode + ) -> Tuple[int, int]: + """Auto-generated builder for SubtractNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SubtractNode as FBSubtractNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSubtractNodeModule.Start(builder) + FBSubtractNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBSubtractNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBSubtractNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSubtractNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SubtractNode + + def _build_Conv1DNode( + self, builder: flatbuffers.Builder, op: Conv1DNode + ) -> Tuple[int, int]: + """Auto-generated builder for Conv1DNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Conv1DNode as FBConv1DNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBConv1DNodeModule.Start(builder) + FBConv1DNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBConv1DNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBConv1DNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConv1DNodeModule.AddStride(builder, op.stride) + FBConv1DNodeModule.AddPadding(builder, op.padding) + FBConv1DNodeModule.AddDilation(builder, op.dilation) + FBConv1DNodeModule.AddGroups(builder, op.groups) + offset = FBConv1DNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Conv1DNode + + def _build_Conv2DNode( + self, builder: flatbuffers.Builder, op: Conv2DNode + ) -> Tuple[int, int]: + """Auto-generated builder for Conv2DNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Conv2DNode as FBConv2DNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBConv2DNodeModule.Start(builder) + FBConv2DNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBConv2DNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBConv2DNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConv2DNodeModule.AddStrideH(builder, op.stride_h) + FBConv2DNodeModule.AddStrideW(builder, op.stride_w) + FBConv2DNodeModule.AddPaddingH(builder, op.padding_h) + FBConv2DNodeModule.AddPaddingW(builder, op.padding_w) + FBConv2DNodeModule.AddDilationH(builder, op.dilation_h) + FBConv2DNodeModule.AddDilationW(builder, op.dilation_w) + FBConv2DNodeModule.AddGroups(builder, op.groups) + offset = FBConv2DNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Conv2DNode + + def _build_Conv3DNode( + self, builder: flatbuffers.Builder, op: Conv3DNode + ) -> Tuple[int, int]: + """Auto-generated builder for Conv3DNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Conv3DNode as FBConv3DNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBConv3DNodeModule.Start(builder) + FBConv3DNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBConv3DNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBConv3DNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConv3DNodeModule.AddStrideD(builder, op.stride_d) + FBConv3DNodeModule.AddStrideH(builder, op.stride_h) + FBConv3DNodeModule.AddStrideW(builder, op.stride_w) + FBConv3DNodeModule.AddPaddingD(builder, op.padding_d) + FBConv3DNodeModule.AddPaddingH(builder, op.padding_h) + FBConv3DNodeModule.AddPaddingW(builder, op.padding_w) + FBConv3DNodeModule.AddDilationD(builder, op.dilation_d) + FBConv3DNodeModule.AddDilationH(builder, op.dilation_h) + FBConv3DNodeModule.AddDilationW(builder, op.dilation_w) + FBConv3DNodeModule.AddGroups(builder, op.groups) + offset = FBConv3DNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Conv3DNode + + def _build_ConvTranspose1DNode( + self, builder: flatbuffers.Builder, op: ConvTranspose1DNode + ) -> Tuple[int, int]: + """Auto-generated builder for ConvTranspose1DNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ConvTranspose1DNode as FBConvTranspose1DNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBConvTranspose1DNodeModule.Start(builder) + FBConvTranspose1DNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBConvTranspose1DNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBConvTranspose1DNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConvTranspose1DNodeModule.AddStride(builder, op.stride) + FBConvTranspose1DNodeModule.AddPadding(builder, op.padding) + FBConvTranspose1DNodeModule.AddDilation(builder, op.dilation) + FBConvTranspose1DNodeModule.AddOutputPadding(builder, op.output_padding) + FBConvTranspose1DNodeModule.AddGroups(builder, op.groups) + offset = FBConvTranspose1DNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ConvTranspose1DNode + + def _build_ConvTranspose2DNode( + self, builder: flatbuffers.Builder, op: ConvTranspose2DNode + ) -> Tuple[int, int]: + """Auto-generated builder for ConvTranspose2DNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ConvTranspose2DNode as FBConvTranspose2DNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBConvTranspose2DNodeModule.Start(builder) + FBConvTranspose2DNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBConvTranspose2DNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBConvTranspose2DNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConvTranspose2DNodeModule.AddStrideH(builder, op.stride_h) + FBConvTranspose2DNodeModule.AddStrideW(builder, op.stride_w) + FBConvTranspose2DNodeModule.AddPaddingH(builder, op.padding_h) + FBConvTranspose2DNodeModule.AddPaddingW(builder, op.padding_w) + FBConvTranspose2DNodeModule.AddDilationH(builder, op.dilation_h) + FBConvTranspose2DNodeModule.AddDilationW(builder, op.dilation_w) + FBConvTranspose2DNodeModule.AddOutputPaddingH(builder, op.output_padding_h) + FBConvTranspose2DNodeModule.AddOutputPaddingW(builder, op.output_padding_w) + FBConvTranspose2DNodeModule.AddGroups(builder, op.groups) + offset = FBConvTranspose2DNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ConvTranspose2DNode + + def _build_ConvTranspose3DNode( + self, builder: flatbuffers.Builder, op: ConvTranspose3DNode + ) -> Tuple[int, int]: + """Auto-generated builder for ConvTranspose3DNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ConvTranspose3DNode as FBConvTranspose3DNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBConvTranspose3DNodeModule.Start(builder) + FBConvTranspose3DNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBConvTranspose3DNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBConvTranspose3DNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConvTranspose3DNodeModule.AddStrideD(builder, op.stride_d) + FBConvTranspose3DNodeModule.AddStrideH(builder, op.stride_h) + FBConvTranspose3DNodeModule.AddStrideW(builder, op.stride_w) + FBConvTranspose3DNodeModule.AddPaddingD(builder, op.padding_d) + FBConvTranspose3DNodeModule.AddPaddingH(builder, op.padding_h) + FBConvTranspose3DNodeModule.AddPaddingW(builder, op.padding_w) + FBConvTranspose3DNodeModule.AddDilationD(builder, op.dilation_d) + FBConvTranspose3DNodeModule.AddDilationH(builder, op.dilation_h) + FBConvTranspose3DNodeModule.AddDilationW(builder, op.dilation_w) + FBConvTranspose3DNodeModule.AddOutputPaddingD(builder, op.output_padding_d) + FBConvTranspose3DNodeModule.AddOutputPaddingH(builder, op.output_padding_h) + FBConvTranspose3DNodeModule.AddOutputPaddingW(builder, op.output_padding_w) + FBConvTranspose3DNodeModule.AddGroups(builder, op.groups) + offset = FBConvTranspose3DNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ConvTranspose3DNode + + def _build_GeluNode( + self, builder: flatbuffers.Builder, op: GeluNode + ) -> Tuple[int, int]: + """Auto-generated builder for GeluNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import GeluNode as FBGeluNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + approximate_off = _shared_string(builder, op.approximate) + + FBGeluNodeModule.Start(builder) + FBGeluNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBGeluNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBGeluNodeModule.AddApproximate(builder, approximate_off) + offset = FBGeluNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.GeluNode + + def _build_ARangeNode( + self, builder: flatbuffers.Builder, op: ARangeNode + ) -> Tuple[int, int]: + """Auto-generated builder for ARangeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ARangeNode as FBARangeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + start_off = self._build_int_or_vid(builder, op.start) + stop_off = self._build_int_or_vid(builder, op.stop) + step_off = self._build_int_or_vid(builder, op.step) + + FBARangeNodeModule.Start(builder) + FBARangeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBARangeNodeModule.AddStart(builder, start_off) + FBARangeNodeModule.AddStop(builder, stop_off) + FBARangeNodeModule.AddStep(builder, step_off) + if op.scalar_type is not None: + FBARangeNodeModule.AddScalarType(builder, op.scalar_type) + offset = FBARangeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ARangeNode + + def _build_SiluNode( + self, builder: flatbuffers.Builder, op: SiluNode + ) -> Tuple[int, int]: + """Auto-generated builder for SiluNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SiluNode as FBSiluNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSiluNodeModule.Start(builder) + FBSiluNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSiluNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSiluNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SiluNode + + def _build_SigmoidNode( + self, builder: flatbuffers.Builder, op: SigmoidNode + ) -> Tuple[int, int]: + """Auto-generated builder for SigmoidNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SigmoidNode as FBSigmoidNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSigmoidNodeModule.Start(builder) + FBSigmoidNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSigmoidNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSigmoidNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SigmoidNode + + def _build_TanhNode( + self, builder: flatbuffers.Builder, op: TanhNode + ) -> Tuple[int, int]: + """Auto-generated builder for TanhNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TanhNode as FBTanhNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBTanhNodeModule.Start(builder) + FBTanhNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTanhNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBTanhNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TanhNode + + def _build_SqueezeNode( + self, builder: flatbuffers.Builder, op: SqueezeNode + ) -> Tuple[int, int]: + """Auto-generated builder for SqueezeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SqueezeNode as FBSqueezeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + dims_vec = _build_int_vector(builder, op.dims) if op.dims is not None else None + + FBSqueezeNodeModule.Start(builder) + FBSqueezeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSqueezeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if dims_vec is not None: + FBSqueezeNodeModule.AddDims(builder, dims_vec) + offset = FBSqueezeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SqueezeNode + + def _build_SplitNode( + self, builder: flatbuffers.Builder, op: SplitNode + ) -> Tuple[int, int]: + """Auto-generated builder for SplitNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SplitNode as FBSplitNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + outs_vec = self._build_tid_vector(builder, op.outs) + sizes_vec = self._build_int_or_vid_vector(builder, op.sizes) + + FBSplitNodeModule.Start(builder) + FBSplitNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSplitNodeModule.AddOuts(builder, outs_vec) + FBSplitNodeModule.AddSizes(builder, sizes_vec) + FBSplitNodeModule.AddAxis(builder, op.axis) + offset = FBSplitNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SplitNode + + def _build_RsqrtNode( + self, builder: flatbuffers.Builder, op: RsqrtNode + ) -> Tuple[int, int]: + """Auto-generated builder for RsqrtNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RsqrtNode as FBRsqrtNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBRsqrtNodeModule.Start(builder) + FBRsqrtNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBRsqrtNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBRsqrtNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RsqrtNode + + def _build_MaximumNode( + self, builder: flatbuffers.Builder, op: MaximumNode + ) -> Tuple[int, int]: + """Auto-generated builder for MaximumNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MaximumNode as FBMaximumNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBMaximumNodeModule.Start(builder) + FBMaximumNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBMaximumNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBMaximumNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBMaximumNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MaximumNode + + def _build_MinimumNode( + self, builder: flatbuffers.Builder, op: MinimumNode + ) -> Tuple[int, int]: + """Auto-generated builder for MinimumNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MinimumNode as FBMinimumNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBMinimumNodeModule.Start(builder) + FBMinimumNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBMinimumNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBMinimumNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBMinimumNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MinimumNode + + def _build_LogNode( + self, builder: flatbuffers.Builder, op: LogNode + ) -> Tuple[int, int]: + """Auto-generated builder for LogNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LogNode as FBLogNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLogNodeModule.Start(builder) + FBLogNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLogNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLogNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LogNode + + def _build_SoftmaxNode( + self, builder: flatbuffers.Builder, op: SoftmaxNode + ) -> Tuple[int, int]: + """Auto-generated builder for SoftmaxNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SoftmaxNode as FBSoftmaxNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSoftmaxNodeModule.Start(builder) + FBSoftmaxNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSoftmaxNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBSoftmaxNodeModule.AddAxis(builder, op.axis) + FBSoftmaxNodeModule.AddPrecise(builder, op.precise) + offset = FBSoftmaxNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SoftmaxNode + + def _build_BroadcastToNode( + self, builder: flatbuffers.Builder, op: BroadcastToNode + ) -> Tuple[int, int]: + """Auto-generated builder for BroadcastToNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import BroadcastToNode as FBBroadcastToNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + shape_vec = self._build_int_or_vid_vector(builder, op.shape) + + FBBroadcastToNodeModule.Start(builder) + FBBroadcastToNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBBroadcastToNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBBroadcastToNodeModule.AddShape(builder, shape_vec) + offset = FBBroadcastToNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.BroadcastToNode + + def _build_PadNode( + self, builder: flatbuffers.Builder, op: PadNode + ) -> Tuple[int, int]: + """Auto-generated builder for PadNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import PadNode as FBPadNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + pad_width_vec = self._build_int_or_vid_vector(builder, op.pad_width) + mode_off = _shared_string(builder, op.mode) + + FBPadNodeModule.Start(builder) + FBPadNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBPadNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBPadNodeModule.AddPadWidth(builder, pad_width_vec) + FBPadNodeModule.AddMode(builder, mode_off) + FBPadNodeModule.AddConstantValue(builder, op.constant_value) + offset = FBPadNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.PadNode + + def _build_WhereNode( + self, builder: flatbuffers.Builder, op: WhereNode + ) -> Tuple[int, int]: + """Auto-generated builder for WhereNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import WhereNode as FBWhereNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBWhereNodeModule.Start(builder) + FBWhereNodeModule.AddCondition(builder, CreateTid(builder, op.condition.idx)) + FBWhereNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBWhereNodeModule.AddY(builder, CreateTid(builder, op.y.idx)) + FBWhereNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBWhereNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.WhereNode + + def _build_ReshapeNode( + self, builder: flatbuffers.Builder, op: ReshapeNode + ) -> Tuple[int, int]: + """Auto-generated builder for ReshapeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ReshapeNode as FBReshapeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + shape_vec = self._build_int_or_vid_vector(builder, op.shape) + + FBReshapeNodeModule.Start(builder) + FBReshapeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBReshapeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBReshapeNodeModule.AddShape(builder, shape_vec) + offset = FBReshapeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ReshapeNode + + def _build_TransposeNode( + self, builder: flatbuffers.Builder, op: TransposeNode + ) -> Tuple[int, int]: + """Auto-generated builder for TransposeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TransposeNode as FBTransposeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + perm_vec = _build_int_vector(builder, op.perm) + + FBTransposeNodeModule.Start(builder) + FBTransposeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTransposeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTransposeNodeModule.AddPerm(builder, perm_vec) + offset = FBTransposeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TransposeNode + + def _build_AsStridedNode( + self, builder: flatbuffers.Builder, op: AsStridedNode + ) -> Tuple[int, int]: + """Auto-generated builder for AsStridedNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AsStridedNode as FBAsStridedNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + shape_vec = self._build_int_or_vid_vector(builder, op.shape) + strides_vec = self._build_int_or_vid_vector(builder, op.strides) + + FBAsStridedNodeModule.Start(builder) + FBAsStridedNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBAsStridedNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBAsStridedNodeModule.AddShape(builder, shape_vec) + FBAsStridedNodeModule.AddStrides(builder, strides_vec) + FBAsStridedNodeModule.AddOffset(builder, op.offset) + offset = FBAsStridedNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AsStridedNode + + def _build_ContiguousNode( + self, builder: flatbuffers.Builder, op: ContiguousNode + ) -> Tuple[int, int]: + """Auto-generated builder for ContiguousNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ContiguousNode as FBContiguousNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBContiguousNodeModule.Start(builder) + FBContiguousNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBContiguousNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBContiguousNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ContiguousNode + + def _build_GatherNode( + self, builder: flatbuffers.Builder, op: GatherNode + ) -> Tuple[int, int]: + """Auto-generated builder for GatherNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import GatherNode as FBGatherNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + indices_vec = self._build_tid_vector(builder, op.indices) + axes_vec = _build_int_vector(builder, op.axes) + slice_sizes_vec = _build_int_vector(builder, op.slice_sizes) + + FBGatherNodeModule.Start(builder) + FBGatherNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBGatherNodeModule.AddIndices(builder, indices_vec) + FBGatherNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBGatherNodeModule.AddAxes(builder, axes_vec) + FBGatherNodeModule.AddSliceSizes(builder, slice_sizes_vec) + offset = FBGatherNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.GatherNode + + def _build_SliceNode( + self, builder: flatbuffers.Builder, op: SliceNode + ) -> Tuple[int, int]: + """Auto-generated builder for SliceNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SliceNode as FBSliceNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axis_off = self._build_int_or_vid(builder, op.axis) + start_off = self._build_int_or_vid(builder, op.start) + stop_off = self._build_int_or_vid(builder, op.stop) + + FBSliceNodeModule.Start(builder) + FBSliceNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSliceNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBSliceNodeModule.AddAxis(builder, axis_off) + FBSliceNodeModule.AddStart(builder, start_off) + FBSliceNodeModule.AddStop(builder, stop_off) + FBSliceNodeModule.AddStep(builder, op.step) + offset = FBSliceNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SliceNode + + def _build_AsTypeNode( + self, builder: flatbuffers.Builder, op: AsTypeNode + ) -> Tuple[int, int]: + """Auto-generated builder for AsTypeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AsTypeNode as FBAsTypeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBAsTypeNodeModule.Start(builder) + FBAsTypeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBAsTypeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBAsTypeNodeModule.AddScalarType(builder, op.scalar_type) + offset = FBAsTypeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AsTypeNode + + def _build_QuantizedMatmulNode( + self, builder: flatbuffers.Builder, op: QuantizedMatmulNode + ) -> Tuple[int, int]: + """Auto-generated builder for QuantizedMatmulNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import QuantizedMatmulNode as FBQuantizedMatmulNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + mode_off = _shared_string(builder, op.mode) + + FBQuantizedMatmulNodeModule.Start(builder) + FBQuantizedMatmulNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBQuantizedMatmulNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBQuantizedMatmulNodeModule.AddScales(builder, CreateTid(builder, op.scales.idx)) + FBQuantizedMatmulNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if op.biases is not None: + FBQuantizedMatmulNodeModule.AddBiases(builder, CreateTid(builder, op.biases.idx)) + FBQuantizedMatmulNodeModule.AddGroupSize(builder, op.group_size) + FBQuantizedMatmulNodeModule.AddBits(builder, op.bits) + FBQuantizedMatmulNodeModule.AddMode(builder, mode_off) + FBQuantizedMatmulNodeModule.AddTranspose(builder, op.transpose) + offset = FBQuantizedMatmulNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.QuantizedMatmulNode + + def _build_ScatterAddNode( + self, builder: flatbuffers.Builder, op: ScatterAddNode + ) -> Tuple[int, int]: + """Auto-generated builder for ScatterAddNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ScatterAddNode as FBScatterAddNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBScatterAddNodeModule.Start(builder) + FBScatterAddNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBScatterAddNodeModule.AddIndices(builder, CreateTid(builder, op.indices.idx)) + FBScatterAddNodeModule.AddUpdates(builder, CreateTid(builder, op.updates.idx)) + FBScatterAddNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBScatterAddNodeModule.AddAxis(builder, op.axis) + offset = FBScatterAddNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ScatterAddNode + + def _build_ConcatenateNode( + self, builder: flatbuffers.Builder, op: ConcatenateNode + ) -> Tuple[int, int]: + """Auto-generated builder for ConcatenateNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ConcatenateNode as FBConcatenateNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + tensors_vec = self._build_tid_vector(builder, op.tensors) + + FBConcatenateNodeModule.Start(builder) + FBConcatenateNodeModule.AddTensors(builder, tensors_vec) + FBConcatenateNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBConcatenateNodeModule.AddAxis(builder, op.axis) + offset = FBConcatenateNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ConcatenateNode + + def _build_FullNode( + self, builder: flatbuffers.Builder, op: FullNode + ) -> Tuple[int, int]: + """Auto-generated builder for FullNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import FullNode as FBFullNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + shape_vec = self._build_int_or_vid_vector(builder, op.shape) + v_off = self._build_float_or_vid(builder, op.v) + + FBFullNodeModule.Start(builder) + FBFullNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBFullNodeModule.AddShape(builder, shape_vec) + FBFullNodeModule.AddV(builder, v_off) + FBFullNodeModule.AddScalarType(builder, op.scalar_type) + offset = FBFullNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.FullNode + + def _build_FullLikeNode( + self, builder: flatbuffers.Builder, op: FullLikeNode + ) -> Tuple[int, int]: + """Auto-generated builder for FullLikeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import FullLikeNode as FBFullLikeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + v_off = self._build_float_or_vid(builder, op.v) + + FBFullLikeNodeModule.Start(builder) + FBFullLikeNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBFullLikeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBFullLikeNodeModule.AddV(builder, v_off) + if op.scalar_type is not None: + FBFullLikeNodeModule.AddScalarType(builder, op.scalar_type) + offset = FBFullLikeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.FullLikeNode + + def _build_ArgmaxNode( + self, builder: flatbuffers.Builder, op: ArgmaxNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArgmaxNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArgmaxNode as FBArgmaxNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArgmaxNodeModule.Start(builder) + FBArgmaxNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArgmaxNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBArgmaxNodeModule.AddAxis(builder, op.axis) + FBArgmaxNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBArgmaxNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArgmaxNode + + def _build_SliceUpdateNode( + self, builder: flatbuffers.Builder, op: SliceUpdateNode + ) -> Tuple[int, int]: + """Auto-generated builder for SliceUpdateNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SliceUpdateNode as FBSliceUpdateNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axis_off = self._build_int_or_vid(builder, op.axis) + start_off = self._build_int_or_vid(builder, op.start) + stop_off = self._build_int_or_vid(builder, op.stop) + + FBSliceUpdateNodeModule.Start(builder) + FBSliceUpdateNodeModule.AddDst(builder, CreateTid(builder, op.dst.idx)) + FBSliceUpdateNodeModule.AddUpdate(builder, CreateTid(builder, op.update.idx)) + FBSliceUpdateNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBSliceUpdateNodeModule.AddAxis(builder, axis_off) + FBSliceUpdateNodeModule.AddStart(builder, start_off) + FBSliceUpdateNodeModule.AddStop(builder, stop_off) + FBSliceUpdateNodeModule.AddStep(builder, op.step) + offset = FBSliceUpdateNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SliceUpdateNode + + def _build_IndexCopyNode( + self, builder: flatbuffers.Builder, op: IndexCopyNode + ) -> Tuple[int, int]: + """Auto-generated builder for IndexCopyNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import IndexCopyNode as FBIndexCopyNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBIndexCopyNodeModule.Start(builder) + FBIndexCopyNodeModule.AddDst(builder, CreateTid(builder, op.dst.idx)) + FBIndexCopyNodeModule.AddUpdate(builder, CreateTid(builder, op.update.idx)) + FBIndexCopyNodeModule.AddIndices(builder, CreateTid(builder, op.indices.idx)) + FBIndexCopyNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBIndexCopyNodeModule.AddAxis(builder, op.axis) + offset = FBIndexCopyNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.IndexCopyNode + + def _build_DequantizeNode( + self, builder: flatbuffers.Builder, op: DequantizeNode + ) -> Tuple[int, int]: + """Auto-generated builder for DequantizeNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import DequantizeNode as FBDequantizeNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + mode_off = _shared_string(builder, op.mode) + + FBDequantizeNodeModule.Start(builder) + FBDequantizeNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBDequantizeNodeModule.AddScales(builder, CreateTid(builder, op.scales.idx)) + FBDequantizeNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if op.biases is not None: + FBDequantizeNodeModule.AddBiases(builder, CreateTid(builder, op.biases.idx)) + FBDequantizeNodeModule.AddGroupSize(builder, op.group_size) + FBDequantizeNodeModule.AddBits(builder, op.bits) + FBDequantizeNodeModule.AddMode(builder, mode_off) + if op.global_scale is not None: + FBDequantizeNodeModule.AddGlobalScale(builder, CreateTid(builder, op.global_scale.idx)) + if op.dtype is not None: + FBDequantizeNodeModule.AddDtype(builder, op.dtype) + offset = FBDequantizeNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.DequantizeNode + + def _build_LessNode( + self, builder: flatbuffers.Builder, op: LessNode + ) -> Tuple[int, int]: + """Auto-generated builder for LessNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LessNode as FBLessNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLessNodeModule.Start(builder) + FBLessNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBLessNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBLessNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLessNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LessNode + + def _build_LessEqualNode( + self, builder: flatbuffers.Builder, op: LessEqualNode + ) -> Tuple[int, int]: + """Auto-generated builder for LessEqualNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LessEqualNode as FBLessEqualNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLessEqualNodeModule.Start(builder) + FBLessEqualNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBLessEqualNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBLessEqualNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLessEqualNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LessEqualNode + + def _build_GreaterNode( + self, builder: flatbuffers.Builder, op: GreaterNode + ) -> Tuple[int, int]: + """Auto-generated builder for GreaterNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import GreaterNode as FBGreaterNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBGreaterNodeModule.Start(builder) + FBGreaterNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBGreaterNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBGreaterNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBGreaterNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.GreaterNode + + def _build_GreaterEqualNode( + self, builder: flatbuffers.Builder, op: GreaterEqualNode + ) -> Tuple[int, int]: + """Auto-generated builder for GreaterEqualNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import GreaterEqualNode as FBGreaterEqualNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBGreaterEqualNodeModule.Start(builder) + FBGreaterEqualNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBGreaterEqualNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBGreaterEqualNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBGreaterEqualNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.GreaterEqualNode + + def _build_EqualNode( + self, builder: flatbuffers.Builder, op: EqualNode + ) -> Tuple[int, int]: + """Auto-generated builder for EqualNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import EqualNode as FBEqualNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBEqualNodeModule.Start(builder) + FBEqualNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBEqualNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBEqualNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBEqualNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.EqualNode + + def _build_NotEqualNode( + self, builder: flatbuffers.Builder, op: NotEqualNode + ) -> Tuple[int, int]: + """Auto-generated builder for NotEqualNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import NotEqualNode as FBNotEqualNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBNotEqualNodeModule.Start(builder) + FBNotEqualNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBNotEqualNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBNotEqualNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBNotEqualNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.NotEqualNode + + def _build_LogicalNotNode( + self, builder: flatbuffers.Builder, op: LogicalNotNode + ) -> Tuple[int, int]: + """Auto-generated builder for LogicalNotNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LogicalNotNode as FBLogicalNotNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLogicalNotNodeModule.Start(builder) + FBLogicalNotNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLogicalNotNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLogicalNotNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LogicalNotNode + + def _build_BitwiseInvertNode( + self, builder: flatbuffers.Builder, op: BitwiseInvertNode + ) -> Tuple[int, int]: + """Auto-generated builder for BitwiseInvertNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import BitwiseInvertNode as FBBitwiseInvertNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBBitwiseInvertNodeModule.Start(builder) + FBBitwiseInvertNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBBitwiseInvertNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBBitwiseInvertNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.BitwiseInvertNode + + def _build_LogicalAndNode( + self, builder: flatbuffers.Builder, op: LogicalAndNode + ) -> Tuple[int, int]: + """Auto-generated builder for LogicalAndNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LogicalAndNode as FBLogicalAndNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLogicalAndNodeModule.Start(builder) + FBLogicalAndNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBLogicalAndNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBLogicalAndNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLogicalAndNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LogicalAndNode + + def _build_LogicalOrNode( + self, builder: flatbuffers.Builder, op: LogicalOrNode + ) -> Tuple[int, int]: + """Auto-generated builder for LogicalOrNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LogicalOrNode as FBLogicalOrNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLogicalOrNodeModule.Start(builder) + FBLogicalOrNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBLogicalOrNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBLogicalOrNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLogicalOrNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LogicalOrNode + + def _build_BitwiseAndNode( + self, builder: flatbuffers.Builder, op: BitwiseAndNode + ) -> Tuple[int, int]: + """Auto-generated builder for BitwiseAndNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import BitwiseAndNode as FBBitwiseAndNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBBitwiseAndNodeModule.Start(builder) + FBBitwiseAndNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBBitwiseAndNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBBitwiseAndNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBBitwiseAndNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.BitwiseAndNode + + def _build_BitwiseOrNode( + self, builder: flatbuffers.Builder, op: BitwiseOrNode + ) -> Tuple[int, int]: + """Auto-generated builder for BitwiseOrNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import BitwiseOrNode as FBBitwiseOrNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBBitwiseOrNodeModule.Start(builder) + FBBitwiseOrNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBBitwiseOrNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBBitwiseOrNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBBitwiseOrNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.BitwiseOrNode + + def _build_BitwiseXorNode( + self, builder: flatbuffers.Builder, op: BitwiseXorNode + ) -> Tuple[int, int]: + """Auto-generated builder for BitwiseXorNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import BitwiseXorNode as FBBitwiseXorNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBBitwiseXorNodeModule.Start(builder) + FBBitwiseXorNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBBitwiseXorNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBBitwiseXorNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBBitwiseXorNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.BitwiseXorNode + + def _build_TriNode( + self, builder: flatbuffers.Builder, op: TriNode + ) -> Tuple[int, int]: + """Auto-generated builder for TriNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TriNode as FBTriNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + n_off = self._build_int_or_vid(builder, op.n) + m_off = self._build_int_or_vid(builder, op.m) + + FBTriNodeModule.Start(builder) + FBTriNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTriNodeModule.AddN(builder, n_off) + FBTriNodeModule.AddM(builder, m_off) + FBTriNodeModule.AddK(builder, op.k) + FBTriNodeModule.AddScalarType(builder, op.scalar_type) + offset = FBTriNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TriNode + + def _build_TrilNode( + self, builder: flatbuffers.Builder, op: TrilNode + ) -> Tuple[int, int]: + """Auto-generated builder for TrilNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TrilNode as FBTrilNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBTrilNodeModule.Start(builder) + FBTrilNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTrilNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTrilNodeModule.AddK(builder, op.k) + offset = FBTrilNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TrilNode + + def _build_TriuNode( + self, builder: flatbuffers.Builder, op: TriuNode + ) -> Tuple[int, int]: + """Auto-generated builder for TriuNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TriuNode as FBTriuNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBTriuNodeModule.Start(builder) + FBTriuNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTriuNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBTriuNodeModule.AddK(builder, op.k) + offset = FBTriuNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TriuNode + + def _build_ClipNode( + self, builder: flatbuffers.Builder, op: ClipNode + ) -> Tuple[int, int]: + """Auto-generated builder for ClipNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ClipNode as FBClipNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBClipNodeModule.Start(builder) + FBClipNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBClipNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if op.a_min is not None: + FBClipNodeModule.AddAMin(builder, CreateTid(builder, op.a_min.idx)) + if op.a_max is not None: + FBClipNodeModule.AddAMax(builder, CreateTid(builder, op.a_max.idx)) + offset = FBClipNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ClipNode + + def _build_CumsumNode( + self, builder: flatbuffers.Builder, op: CumsumNode + ) -> Tuple[int, int]: + """Auto-generated builder for CumsumNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import CumsumNode as FBCumsumNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBCumsumNodeModule.Start(builder) + FBCumsumNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBCumsumNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBCumsumNodeModule.AddAxis(builder, op.axis) + FBCumsumNodeModule.AddReverse(builder, op.reverse) + FBCumsumNodeModule.AddInclusive(builder, op.inclusive) + offset = FBCumsumNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.CumsumNode + + def _build_StackNode( + self, builder: flatbuffers.Builder, op: StackNode + ) -> Tuple[int, int]: + """Auto-generated builder for StackNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import StackNode as FBStackNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + tensors_vec = self._build_tid_vector(builder, op.tensors) + + FBStackNodeModule.Start(builder) + FBStackNodeModule.AddTensors(builder, tensors_vec) + FBStackNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBStackNodeModule.AddAxis(builder, op.axis) + offset = FBStackNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.StackNode + + def _build_SignNode( + self, builder: flatbuffers.Builder, op: SignNode + ) -> Tuple[int, int]: + """Auto-generated builder for SignNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SignNode as FBSignNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSignNodeModule.Start(builder) + FBSignNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSignNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSignNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SignNode + + def _build_AnyNode( + self, builder: flatbuffers.Builder, op: AnyNode + ) -> Tuple[int, int]: + """Auto-generated builder for AnyNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AnyNode as FBAnyNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBAnyNodeModule.Start(builder) + FBAnyNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBAnyNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBAnyNodeModule.AddAxes(builder, axes_vec) + FBAnyNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBAnyNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AnyNode + + def _build_AllNode( + self, builder: flatbuffers.Builder, op: AllNode + ) -> Tuple[int, int]: + """Auto-generated builder for AllNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AllNode as FBAllNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBAllNodeModule.Start(builder) + FBAllNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBAllNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBAllNodeModule.AddAxes(builder, axes_vec) + FBAllNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBAllNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AllNode + + def _build_RepeatNode( + self, builder: flatbuffers.Builder, op: RepeatNode + ) -> Tuple[int, int]: + """Auto-generated builder for RepeatNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RepeatNode as FBRepeatNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + repeats_off = self._build_int_or_vid(builder, op.repeats) + + FBRepeatNodeModule.Start(builder) + FBRepeatNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBRepeatNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBRepeatNodeModule.AddRepeats(builder, repeats_off) + FBRepeatNodeModule.AddAxis(builder, op.axis) + offset = FBRepeatNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RepeatNode + + def _build_SortNode( + self, builder: flatbuffers.Builder, op: SortNode + ) -> Tuple[int, int]: + """Auto-generated builder for SortNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SortNode as FBSortNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSortNodeModule.Start(builder) + FBSortNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSortNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBSortNodeModule.AddAxis(builder, op.axis) + offset = FBSortNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SortNode + + def _build_ArgsortNode( + self, builder: flatbuffers.Builder, op: ArgsortNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArgsortNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArgsortNode as FBArgsortNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArgsortNodeModule.Start(builder) + FBArgsortNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArgsortNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBArgsortNodeModule.AddAxis(builder, op.axis) + offset = FBArgsortNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArgsortNode + + def _build_PartitionNode( + self, builder: flatbuffers.Builder, op: PartitionNode + ) -> Tuple[int, int]: + """Auto-generated builder for PartitionNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import PartitionNode as FBPartitionNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + kth_off = self._build_int_or_vid(builder, op.kth) + + FBPartitionNodeModule.Start(builder) + FBPartitionNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBPartitionNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBPartitionNodeModule.AddKth(builder, kth_off) + FBPartitionNodeModule.AddAxis(builder, op.axis) + offset = FBPartitionNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.PartitionNode + + def _build_ArgPartitionNode( + self, builder: flatbuffers.Builder, op: ArgPartitionNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArgPartitionNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArgPartitionNode as FBArgPartitionNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + kth_off = self._build_int_or_vid(builder, op.kth) + + FBArgPartitionNodeModule.Start(builder) + FBArgPartitionNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArgPartitionNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBArgPartitionNodeModule.AddKth(builder, kth_off) + FBArgPartitionNodeModule.AddAxis(builder, op.axis) + offset = FBArgPartitionNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArgPartitionNode + + def _build_RollNode( + self, builder: flatbuffers.Builder, op: RollNode + ) -> Tuple[int, int]: + """Auto-generated builder for RollNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RollNode as FBRollNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + shift_vec = self._build_int_or_vid_vector(builder, op.shift) + axes_vec = _build_int_vector(builder, op.axes) + + FBRollNodeModule.Start(builder) + FBRollNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBRollNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBRollNodeModule.AddShift(builder, shift_vec) + FBRollNodeModule.AddAxes(builder, axes_vec) + offset = FBRollNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RollNode + + def _build_FloorNode( + self, builder: flatbuffers.Builder, op: FloorNode + ) -> Tuple[int, int]: + """Auto-generated builder for FloorNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import FloorNode as FBFloorNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBFloorNodeModule.Start(builder) + FBFloorNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBFloorNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBFloorNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.FloorNode + + def _build_CeilNode( + self, builder: flatbuffers.Builder, op: CeilNode + ) -> Tuple[int, int]: + """Auto-generated builder for CeilNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import CeilNode as FBCeilNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBCeilNodeModule.Start(builder) + FBCeilNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBCeilNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBCeilNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.CeilNode + + def _build_SquareNode( + self, builder: flatbuffers.Builder, op: SquareNode + ) -> Tuple[int, int]: + """Auto-generated builder for SquareNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SquareNode as FBSquareNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSquareNodeModule.Start(builder) + FBSquareNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSquareNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSquareNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SquareNode + + def _build_ExpNode( + self, builder: flatbuffers.Builder, op: ExpNode + ) -> Tuple[int, int]: + """Auto-generated builder for ExpNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ExpNode as FBExpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBExpNodeModule.Start(builder) + FBExpNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBExpNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBExpNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ExpNode + + def _build_SinNode( + self, builder: flatbuffers.Builder, op: SinNode + ) -> Tuple[int, int]: + """Auto-generated builder for SinNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SinNode as FBSinNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSinNodeModule.Start(builder) + FBSinNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSinNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSinNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SinNode + + def _build_CosNode( + self, builder: flatbuffers.Builder, op: CosNode + ) -> Tuple[int, int]: + """Auto-generated builder for CosNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import CosNode as FBCosNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBCosNodeModule.Start(builder) + FBCosNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBCosNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBCosNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.CosNode + + def _build_TanNode( + self, builder: flatbuffers.Builder, op: TanNode + ) -> Tuple[int, int]: + """Auto-generated builder for TanNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import TanNode as FBTanNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBTanNodeModule.Start(builder) + FBTanNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBTanNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBTanNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.TanNode + + def _build_ArcsinNode( + self, builder: flatbuffers.Builder, op: ArcsinNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArcsinNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArcsinNode as FBArcsinNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArcsinNodeModule.Start(builder) + FBArcsinNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArcsinNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBArcsinNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArcsinNode + + def _build_ArccosNode( + self, builder: flatbuffers.Builder, op: ArccosNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArccosNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArccosNode as FBArccosNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArccosNodeModule.Start(builder) + FBArccosNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArccosNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBArccosNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArccosNode + + def _build_ArctanNode( + self, builder: flatbuffers.Builder, op: ArctanNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArctanNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArctanNode as FBArctanNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArctanNodeModule.Start(builder) + FBArctanNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArctanNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBArctanNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArctanNode + + def _build_SinhNode( + self, builder: flatbuffers.Builder, op: SinhNode + ) -> Tuple[int, int]: + """Auto-generated builder for SinhNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SinhNode as FBSinhNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSinhNodeModule.Start(builder) + FBSinhNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSinhNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSinhNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SinhNode + + def _build_CoshNode( + self, builder: flatbuffers.Builder, op: CoshNode + ) -> Tuple[int, int]: + """Auto-generated builder for CoshNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import CoshNode as FBCoshNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBCoshNodeModule.Start(builder) + FBCoshNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBCoshNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBCoshNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.CoshNode + + def _build_ArcsinhNode( + self, builder: flatbuffers.Builder, op: ArcsinhNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArcsinhNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArcsinhNode as FBArcsinhNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArcsinhNodeModule.Start(builder) + FBArcsinhNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArcsinhNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBArcsinhNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArcsinhNode + + def _build_ArccoshNode( + self, builder: flatbuffers.Builder, op: ArccoshNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArccoshNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArccoshNode as FBArccoshNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArccoshNodeModule.Start(builder) + FBArccoshNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArccoshNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBArccoshNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArccoshNode + + def _build_ArctanhNode( + self, builder: flatbuffers.Builder, op: ArctanhNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArctanhNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArctanhNode as FBArctanhNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArctanhNodeModule.Start(builder) + FBArctanhNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArctanhNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBArctanhNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArctanhNode + + def _build_Log2Node( + self, builder: flatbuffers.Builder, op: Log2Node + ) -> Tuple[int, int]: + """Auto-generated builder for Log2Node.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Log2Node as FBLog2NodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLog2NodeModule.Start(builder) + FBLog2NodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLog2NodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLog2NodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Log2Node + + def _build_Log10Node( + self, builder: flatbuffers.Builder, op: Log10Node + ) -> Tuple[int, int]: + """Auto-generated builder for Log10Node.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Log10Node as FBLog10NodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLog10NodeModule.Start(builder) + FBLog10NodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLog10NodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLog10NodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Log10Node + + def _build_Log1pNode( + self, builder: flatbuffers.Builder, op: Log1pNode + ) -> Tuple[int, int]: + """Auto-generated builder for Log1pNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Log1pNode as FBLog1pNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLog1pNodeModule.Start(builder) + FBLog1pNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLog1pNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLog1pNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Log1pNode + + def _build_ErfNode( + self, builder: flatbuffers.Builder, op: ErfNode + ) -> Tuple[int, int]: + """Auto-generated builder for ErfNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ErfNode as FBErfNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBErfNodeModule.Start(builder) + FBErfNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBErfNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBErfNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ErfNode + + def _build_Expm1Node( + self, builder: flatbuffers.Builder, op: Expm1Node + ) -> Tuple[int, int]: + """Auto-generated builder for Expm1Node.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Expm1Node as FBExpm1NodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBExpm1NodeModule.Start(builder) + FBExpm1NodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBExpm1NodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBExpm1NodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Expm1Node + + def _build_RoundNode( + self, builder: flatbuffers.Builder, op: RoundNode + ) -> Tuple[int, int]: + """Auto-generated builder for RoundNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RoundNode as FBRoundNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBRoundNodeModule.Start(builder) + FBRoundNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBRoundNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBRoundNodeModule.AddDecimals(builder, op.decimals) + offset = FBRoundNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RoundNode + + def _build_ReciprocalNode( + self, builder: flatbuffers.Builder, op: ReciprocalNode + ) -> Tuple[int, int]: + """Auto-generated builder for ReciprocalNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ReciprocalNode as FBReciprocalNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBReciprocalNodeModule.Start(builder) + FBReciprocalNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBReciprocalNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBReciprocalNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ReciprocalNode + + def _build_SqrtNode( + self, builder: flatbuffers.Builder, op: SqrtNode + ) -> Tuple[int, int]: + """Auto-generated builder for SqrtNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SqrtNode as FBSqrtNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBSqrtNodeModule.Start(builder) + FBSqrtNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSqrtNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBSqrtNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SqrtNode + + def _build_AbsNode( + self, builder: flatbuffers.Builder, op: AbsNode + ) -> Tuple[int, int]: + """Auto-generated builder for AbsNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import AbsNode as FBAbsNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBAbsNodeModule.Start(builder) + FBAbsNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBAbsNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBAbsNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.AbsNode + + def _build_NegNode( + self, builder: flatbuffers.Builder, op: NegNode + ) -> Tuple[int, int]: + """Auto-generated builder for NegNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import NegNode as FBNegNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBNegNodeModule.Start(builder) + FBNegNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBNegNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBNegNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.NegNode + + def _build_Atan2Node( + self, builder: flatbuffers.Builder, op: Atan2Node + ) -> Tuple[int, int]: + """Auto-generated builder for Atan2Node.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import Atan2Node as FBAtan2NodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBAtan2NodeModule.Start(builder) + FBAtan2NodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBAtan2NodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBAtan2NodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBAtan2NodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.Atan2Node + + def _build_LogAddExpNode( + self, builder: flatbuffers.Builder, op: LogAddExpNode + ) -> Tuple[int, int]: + """Auto-generated builder for LogAddExpNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LogAddExpNode as FBLogAddExpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBLogAddExpNodeModule.Start(builder) + FBLogAddExpNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBLogAddExpNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBLogAddExpNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBLogAddExpNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LogAddExpNode + + def _build_FloorDivideNode( + self, builder: flatbuffers.Builder, op: FloorDivideNode + ) -> Tuple[int, int]: + """Auto-generated builder for FloorDivideNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import FloorDivideNode as FBFloorDivideNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBFloorDivideNodeModule.Start(builder) + FBFloorDivideNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBFloorDivideNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBFloorDivideNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBFloorDivideNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.FloorDivideNode + + def _build_RemainderNode( + self, builder: flatbuffers.Builder, op: RemainderNode + ) -> Tuple[int, int]: + """Auto-generated builder for RemainderNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RemainderNode as FBRemainderNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBRemainderNodeModule.Start(builder) + FBRemainderNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBRemainderNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBRemainderNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBRemainderNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RemainderNode + + def _build_PowerNode( + self, builder: flatbuffers.Builder, op: PowerNode + ) -> Tuple[int, int]: + """Auto-generated builder for PowerNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import PowerNode as FBPowerNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBPowerNodeModule.Start(builder) + FBPowerNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBPowerNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBPowerNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + offset = FBPowerNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.PowerNode + + def _build_LogSumExpNode( + self, builder: flatbuffers.Builder, op: LogSumExpNode + ) -> Tuple[int, int]: + """Auto-generated builder for LogSumExpNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import LogSumExpNode as FBLogSumExpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBLogSumExpNodeModule.Start(builder) + FBLogSumExpNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBLogSumExpNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBLogSumExpNodeModule.AddAxes(builder, axes_vec) + FBLogSumExpNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBLogSumExpNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.LogSumExpNode + + def _build_SumNode( + self, builder: flatbuffers.Builder, op: SumNode + ) -> Tuple[int, int]: + """Auto-generated builder for SumNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import SumNode as FBSumNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBSumNodeModule.Start(builder) + FBSumNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBSumNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBSumNodeModule.AddAxes(builder, axes_vec) + FBSumNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBSumNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.SumNode + + def _build_MeanNode( + self, builder: flatbuffers.Builder, op: MeanNode + ) -> Tuple[int, int]: + """Auto-generated builder for MeanNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MeanNode as FBMeanNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBMeanNodeModule.Start(builder) + FBMeanNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBMeanNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBMeanNodeModule.AddAxes(builder, axes_vec) + FBMeanNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBMeanNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MeanNode + + def _build_VarNode( + self, builder: flatbuffers.Builder, op: VarNode + ) -> Tuple[int, int]: + """Auto-generated builder for VarNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import VarNode as FBVarNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBVarNodeModule.Start(builder) + FBVarNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBVarNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBVarNodeModule.AddAxes(builder, axes_vec) + FBVarNodeModule.AddKeepdims(builder, op.keepdims) + FBVarNodeModule.AddDdof(builder, op.ddof) + offset = FBVarNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.VarNode + + def _build_StdNode( + self, builder: flatbuffers.Builder, op: StdNode + ) -> Tuple[int, int]: + """Auto-generated builder for StdNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import StdNode as FBStdNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBStdNodeModule.Start(builder) + FBStdNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBStdNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBStdNodeModule.AddAxes(builder, axes_vec) + FBStdNodeModule.AddKeepdims(builder, op.keepdims) + FBStdNodeModule.AddDdof(builder, op.ddof) + offset = FBStdNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.StdNode + + def _build_ProdNode( + self, builder: flatbuffers.Builder, op: ProdNode + ) -> Tuple[int, int]: + """Auto-generated builder for ProdNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ProdNode as FBProdNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBProdNodeModule.Start(builder) + FBProdNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBProdNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBProdNodeModule.AddAxes(builder, axes_vec) + FBProdNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBProdNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ProdNode + + def _build_MaxNode( + self, builder: flatbuffers.Builder, op: MaxNode + ) -> Tuple[int, int]: + """Auto-generated builder for MaxNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MaxNode as FBMaxNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBMaxNodeModule.Start(builder) + FBMaxNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBMaxNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBMaxNodeModule.AddAxes(builder, axes_vec) + FBMaxNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBMaxNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MaxNode + + def _build_MinNode( + self, builder: flatbuffers.Builder, op: MinNode + ) -> Tuple[int, int]: + """Auto-generated builder for MinNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MinNode as FBMinNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBMinNodeModule.Start(builder) + FBMinNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBMinNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBMinNodeModule.AddAxes(builder, axes_vec) + FBMinNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBMinNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MinNode + + def _build_ArgminNode( + self, builder: flatbuffers.Builder, op: ArgminNode + ) -> Tuple[int, int]: + """Auto-generated builder for ArgminNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ArgminNode as FBArgminNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBArgminNodeModule.Start(builder) + FBArgminNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBArgminNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBArgminNodeModule.AddAxis(builder, op.axis) + FBArgminNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBArgminNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ArgminNode + + def _build_MedianNode( + self, builder: flatbuffers.Builder, op: MedianNode + ) -> Tuple[int, int]: + """Auto-generated builder for MedianNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MedianNode as FBMedianNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + axes_vec = _build_int_vector(builder, op.axes) if op.axes is not None else None + + FBMedianNodeModule.Start(builder) + FBMedianNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBMedianNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if axes_vec is not None: + FBMedianNodeModule.AddAxes(builder, axes_vec) + FBMedianNodeModule.AddKeepdims(builder, op.keepdims) + offset = FBMedianNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MedianNode + + def _build_GatherMmNode( + self, builder: flatbuffers.Builder, op: GatherMmNode + ) -> Tuple[int, int]: + """Auto-generated builder for GatherMmNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import GatherMmNode as FBGatherMmNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + FBGatherMmNodeModule.Start(builder) + FBGatherMmNodeModule.AddA(builder, CreateTid(builder, op.a.idx)) + FBGatherMmNodeModule.AddB(builder, CreateTid(builder, op.b.idx)) + FBGatherMmNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + if op.lhs_indices is not None: + FBGatherMmNodeModule.AddLhsIndices(builder, CreateTid(builder, op.lhs_indices.idx)) + if op.rhs_indices is not None: + FBGatherMmNodeModule.AddRhsIndices(builder, CreateTid(builder, op.rhs_indices.idx)) + FBGatherMmNodeModule.AddSortedIndices(builder, op.sorted_indices) + offset = FBGatherMmNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.GatherMmNode + + def _build_GatherQmmNode( + self, builder: flatbuffers.Builder, op: GatherQmmNode + ) -> Tuple[int, int]: + """Auto-generated builder for GatherQmmNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import GatherQmmNode as FBGatherQmmNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + mode_off = _shared_string(builder, op.mode) + + FBGatherQmmNodeModule.Start(builder) + FBGatherQmmNodeModule.AddX(builder, CreateTid(builder, op.x.idx)) + FBGatherQmmNodeModule.AddW(builder, CreateTid(builder, op.w.idx)) + FBGatherQmmNodeModule.AddScales(builder, CreateTid(builder, op.scales.idx)) + FBGatherQmmNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBGatherQmmNodeModule.AddMode(builder, mode_off) + if op.biases is not None: + FBGatherQmmNodeModule.AddBiases(builder, CreateTid(builder, op.biases.idx)) + if op.lhs_indices is not None: + FBGatherQmmNodeModule.AddLhsIndices(builder, CreateTid(builder, op.lhs_indices.idx)) + if op.rhs_indices is not None: + FBGatherQmmNodeModule.AddRhsIndices(builder, CreateTid(builder, op.rhs_indices.idx)) + FBGatherQmmNodeModule.AddTranspose(builder, op.transpose) + FBGatherQmmNodeModule.AddGroupSize(builder, op.group_size) + FBGatherQmmNodeModule.AddBits(builder, op.bits) + FBGatherQmmNodeModule.AddSortedIndices(builder, op.sorted_indices) + offset = FBGatherQmmNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.GatherQmmNode + + def _build_ScanNode( + self, builder: flatbuffers.Builder, op: ScanNode + ) -> Tuple[int, int]: + """Auto-generated builder for ScanNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import ScanNode as FBScanNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + originals_vec = self._build_tid_vector(builder, op.originals) + sliced_vec = self._build_tid_vector(builder, op.sliced) + outputs_vec = self._build_tid_vector(builder, op.outputs) + carry_vec = self._build_tid_vector(builder, op.carry) + + FBScanNodeModule.Start(builder) + FBScanNodeModule.AddOriginals(builder, originals_vec) + FBScanNodeModule.AddSliced(builder, sliced_vec) + FBScanNodeModule.AddOutputs(builder, outputs_vec) + FBScanNodeModule.AddCarry(builder, carry_vec) + FBScanNodeModule.AddBodyChainIdx(builder, op.body_chain_idx) + FBScanNodeModule.AddScanAxis(builder, op.scan_axis) + offset = FBScanNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.ScanNode + + def _build_IfNode( + self, builder: flatbuffers.Builder, op: IfNode + ) -> Tuple[int, int]: + """Auto-generated builder for IfNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import IfNode as FBIfNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + cond_off = self._build_int_or_vid(builder, op.cond) + + FBIfNodeModule.Start(builder) + FBIfNodeModule.AddCond(builder, cond_off) + FBIfNodeModule.AddThenChainIdx(builder, op.then_chain_idx) + FBIfNodeModule.AddElseChainIdx(builder, op.else_chain_idx) + offset = FBIfNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.IfNode + + def _build_RandomBitsNode( + self, builder: flatbuffers.Builder, op: RandomBitsNode + ) -> Tuple[int, int]: + """Auto-generated builder for RandomBitsNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import RandomBitsNode as FBRandomBitsNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + shape_vec = self._build_int_or_vid_vector(builder, op.shape) + + FBRandomBitsNodeModule.Start(builder) + FBRandomBitsNodeModule.AddOut(builder, CreateTid(builder, op.out.idx)) + FBRandomBitsNodeModule.AddShape(builder, shape_vec) + if op.seed is not None: + FBRandomBitsNodeModule.AddSeed(builder, CreateVid(builder, op.seed.idx)) + FBRandomBitsNodeModule.AddWidth(builder, op.width) + offset = FBRandomBitsNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.RandomBitsNode + + def _build_MetalKernelNode( + self, builder: flatbuffers.Builder, op: MetalKernelNode + ) -> Tuple[int, int]: + """Auto-generated builder for MetalKernelNode.""" + # Import the MODULE (not class) to access builder functions like Start(), Add*(), End() + from executorch.backends.mlx.serialization._generated.mlx_delegate import MetalKernelNode as FBMetalKernelNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate import OpNode as FBOpNodeModule + from executorch.backends.mlx.serialization._generated.mlx_delegate.Tid import CreateTid + from executorch.backends.mlx.serialization._generated.mlx_delegate.Vid import CreateVid + + name_off = _shared_string(builder, op.name) + source_off = _shared_string(builder, op.source) + inputs_vec = self._build_tid_vector(builder, op.inputs) + outputs_vec = self._build_tid_vector(builder, op.outputs) + grid_vec = self._build_int_or_vid_vector(builder, op.grid) + threadgroup_vec = self._build_int_or_vid_vector(builder, op.threadgroup) + header_off = _shared_string(builder, op.header) + input_names_vec = self._build_string_vector(builder, op.input_names) if op.input_names is not None else None + output_names_vec = self._build_string_vector(builder, op.output_names) if op.output_names is not None else None + output_shapes_flat_vec = self._build_int_or_vid_vector(builder, op.output_shapes_flat) if op.output_shapes_flat is not None else None + output_shape_lengths_vec = _build_int_vector(builder, op.output_shape_lengths) if op.output_shape_lengths is not None else None + output_dtypes_vec = _build_int8_vector(builder, op.output_dtypes) if op.output_dtypes is not None else None + template_arg_names_vec = self._build_string_vector(builder, op.template_arg_names) if op.template_arg_names is not None else None + template_arg_kinds_vec = _build_int8_vector(builder, op.template_arg_kinds) if op.template_arg_kinds is not None else None + template_arg_values_vec = _build_int_vector(builder, op.template_arg_values) if op.template_arg_values is not None else None + + FBMetalKernelNodeModule.Start(builder) + FBMetalKernelNodeModule.AddName(builder, name_off) + FBMetalKernelNodeModule.AddSource(builder, source_off) + FBMetalKernelNodeModule.AddInputs(builder, inputs_vec) + FBMetalKernelNodeModule.AddOutputs(builder, outputs_vec) + FBMetalKernelNodeModule.AddGrid(builder, grid_vec) + FBMetalKernelNodeModule.AddThreadgroup(builder, threadgroup_vec) + if header_off is not None: + FBMetalKernelNodeModule.AddHeader(builder, header_off) + if input_names_vec is not None: + FBMetalKernelNodeModule.AddInputNames(builder, input_names_vec) + if output_names_vec is not None: + FBMetalKernelNodeModule.AddOutputNames(builder, output_names_vec) + FBMetalKernelNodeModule.AddEnsureRowContiguous(builder, op.ensure_row_contiguous) + FBMetalKernelNodeModule.AddAtomicOutputs(builder, op.atomic_outputs) + if output_shapes_flat_vec is not None: + FBMetalKernelNodeModule.AddOutputShapesFlat(builder, output_shapes_flat_vec) + if output_shape_lengths_vec is not None: + FBMetalKernelNodeModule.AddOutputShapeLengths(builder, output_shape_lengths_vec) + if output_dtypes_vec is not None: + FBMetalKernelNodeModule.AddOutputDtypes(builder, output_dtypes_vec) + if template_arg_names_vec is not None: + FBMetalKernelNodeModule.AddTemplateArgNames(builder, template_arg_names_vec) + if template_arg_kinds_vec is not None: + FBMetalKernelNodeModule.AddTemplateArgKinds(builder, template_arg_kinds_vec) + if template_arg_values_vec is not None: + FBMetalKernelNodeModule.AddTemplateArgValues(builder, template_arg_values_vec) + if op.init_value is not None: + FBMetalKernelNodeModule.AddInitValue(builder, op.init_value) + offset = FBMetalKernelNodeModule.End(builder) + return offset, FBOpNodeModule.OpNode.MetalKernelNode diff --git a/backends/mlx/serialization/mlx_graph_schema.py b/backends/mlx/serialization/mlx_graph_schema.py new file mode 100644 index 00000000000..4c30e75fdb2 --- /dev/null +++ b/backends/mlx/serialization/mlx_graph_schema.py @@ -0,0 +1,1384 @@ +# +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. +# +# ============================================================================ +# AUTO-GENERATED FILE - DO NOT EDIT MANUALLY +# ============================================================================ +# +# This file was generated from schema.fbs by the MLX delegate code generator. +# +# Source: backends/mlx/serialization/schema.fbs +# Generator: backends/mlx/serialization/generate.py +# +# To regenerate, run from the executorch root: +# python backends/mlx/serialization/generate.py +# +# ============================================================================ + +from __future__ import annotations + +from dataclasses import dataclass, field +from enum import IntEnum +from typing import List, Optional, Union + + +# ============================================================================ +# Enums +# ============================================================================ + +class SlotType(IntEnum): + TensorSlot = 0 + IntValueSlot = 1 + FloatValueSlot = 2 + BoolValueSlot = 3 + + +# ============================================================================ +# Core types +# ============================================================================ + +@dataclass +class Tid: + idx: Optional[int] + + +@dataclass +class Vid: + idx: Optional[int] + + +@dataclass +class FloatOrVid: + """Represents either a literal float or a runtime Vid reference.""" + literal: float = 0.0 + vid: Optional[Vid] = None + is_vid: bool = False + + @classmethod + def from_literal(cls, value: float) -> "FloatOrVid": + """Create a FloatOrVid from a literal float.""" + return cls(literal=value, is_vid=False) + + @classmethod + def from_vid(cls, vid: Vid) -> "FloatOrVid": + """Create a FloatOrVid from a Vid reference.""" + return cls(vid=vid, is_vid=True) + + +@dataclass +class IntOrVid: + """Represents either a literal integer or a runtime Vid reference.""" + literal: int = 0 + vid: Optional[Vid] = None + is_vid: bool = False + + @classmethod + def from_literal(cls, value: int) -> "IntOrVid": + """Create a IntOrVid from a literal integer.""" + return cls(literal=value, is_vid=False) + + @classmethod + def from_vid(cls, vid: Vid) -> "IntOrVid": + """Create a IntOrVid from a Vid reference.""" + return cls(vid=vid, is_vid=True) + + +@dataclass +class IntOrVidOrTid: + """Represents either a literal integer or a runtime Vid reference.""" + literal: int = 0 + vid: Optional[Vid] = None + tid: Optional[Tid] = None + kind: int = 0 + + @classmethod + def from_literal(cls, value: int) -> "IntOrVidOrTid": + """Create a IntOrVidOrTid from a literal integer.""" + return cls(literal=value, kind=0) + + @classmethod + def from_vid(cls, vid: Vid) -> "IntOrVidOrTid": + """Create a IntOrVidOrTid from a Vid reference.""" + return cls(vid=vid, kind=1) + + @classmethod + def from_tid(cls, tid: Tid) -> "IntOrVidOrTid": + """Create a IntOrVidOrTid from a Tid tensor reference.""" + return cls(tid=tid, kind=2) + + +@dataclass +class VidOrTid: + """Represents either a tensor reference or a runtime Vid reference.""" + vid: Optional[Vid] = None + tid: Optional[Tid] = None + is_vid: bool = False + + @classmethod + def from_tid(cls, value: Tid) -> "VidOrTid": + """Create a VidOrTid from a tensor reference.""" + return cls(tid=value, is_vid=False) + + @classmethod + def from_vid(cls, vid: Vid) -> "VidOrTid": + """Create a VidOrTid from a Vid reference.""" + return cls(vid=vid, is_vid=True) + + @classmethod + def from_tid(cls, tid: Tid) -> "VidOrTid": + """Create a VidOrTid from a Tid tensor reference.""" + return cls(tid=tid, is_vid=False) + + +@dataclass +class ShapeDim: + value: int = -1 + min_value: int = 0 + max_value: int = -1 + + +@dataclass +class SlotVariant: + slot_type: SlotType = SlotType.TensorSlot + idx: Optional[int] = None + + +@dataclass +class NamedSlot: + name: str + slot: SlotVariant + + +@dataclass +class TensorMeta: + shape: List[ShapeDim] + scalar_type: Optional[int] = None + dim_order: Optional[List[int]] = None + + +# ============================================================================ +# Op nodes +# ============================================================================ + +@dataclass +class NoopNode: + pass + + +@dataclass +class IdCopyNode: + x: Tid + out: Tid + + +@dataclass +class AddmmNode: + mat1: Tid + mat2: Tid + out: Tid + alpha: float = 1.0 + beta: float = 1.0 + bias: Optional[Tid] = None + + +@dataclass +class ItemIntNode: + x: Tid + out: Vid + + +@dataclass +class ExpandDimsNode: + x: Tid + out: Tid + axis: Optional[int] = None + + +@dataclass +class TileNode: + x: Tid + out: Tid + reps: List[IntOrVid] + + +@dataclass +class TakeAlongAxisNode: + x: Tid + indices: Tid + out: Tid + axis: Optional[int] = None + + +@dataclass +class TakeNode: + x: Tid + out: Tid + index: IntOrVidOrTid + axis: Optional[int] = None + + +@dataclass +class RMSNormNode: + x: Tid + out: Tid + weight: Optional[Tid] = None + eps: Optional[float] = None + + +@dataclass +class LayerNormNode: + x: Tid + out: Tid + weight: Optional[Tid] = None + bias: Optional[Tid] = None + eps: Optional[float] = None + + +@dataclass +class RopeNode: + x: Tid + out: Tid + offset: VidOrTid + traditional: bool = False + base: float = 500000.0 + scale: float = 1.0 + dims: Optional[int] = None + freqs: Optional[Tid] = None + + +@dataclass +class SdpaNode: + q: Tid + k: Tid + v: Tid + out: Tid + causal: bool = False + scale: Optional[float] = None + mask: Optional[Tid] = None + + +@dataclass +class AddNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class AddIntNode: + a: IntOrVid + b: IntOrVid + out: Vid + + +@dataclass +class SubtractIntNode: + a: IntOrVid + b: IntOrVid + out: Vid + + +@dataclass +class MultiplyIntNode: + a: IntOrVid + b: IntOrVid + out: Vid + + +@dataclass +class FloorDivideIntNode: + a: IntOrVid + b: IntOrVid + out: Vid + + +@dataclass +class ModIntNode: + a: IntOrVid + b: IntOrVid + out: Vid + + +@dataclass +class SymSizeNode: + a: Tid + out: Vid + dim: Optional[int] = None + + +@dataclass +class MultiplyNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class DivideNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class SubtractNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class Conv1DNode: + x: Tid + w: Tid + out: Tid + stride: int = 1 + padding: int = 0 + dilation: int = 1 + groups: int = 1 + + +@dataclass +class Conv2DNode: + x: Tid + w: Tid + out: Tid + stride_h: int = 1 + stride_w: int = 1 + padding_h: int = 0 + padding_w: int = 0 + dilation_h: int = 1 + dilation_w: int = 1 + groups: int = 1 + + +@dataclass +class Conv3DNode: + x: Tid + w: Tid + out: Tid + stride_d: int = 1 + stride_h: int = 1 + stride_w: int = 1 + padding_d: int = 0 + padding_h: int = 0 + padding_w: int = 0 + dilation_d: int = 1 + dilation_h: int = 1 + dilation_w: int = 1 + groups: int = 1 + + +@dataclass +class ConvTranspose1DNode: + x: Tid + w: Tid + out: Tid + stride: int = 1 + padding: int = 0 + dilation: int = 1 + output_padding: int = 0 + groups: int = 1 + + +@dataclass +class ConvTranspose2DNode: + x: Tid + w: Tid + out: Tid + stride_h: int = 1 + stride_w: int = 1 + padding_h: int = 0 + padding_w: int = 0 + dilation_h: int = 1 + dilation_w: int = 1 + output_padding_h: int = 0 + output_padding_w: int = 0 + groups: int = 1 + + +@dataclass +class ConvTranspose3DNode: + x: Tid + w: Tid + out: Tid + stride_d: int = 1 + stride_h: int = 1 + stride_w: int = 1 + padding_d: int = 0 + padding_h: int = 0 + padding_w: int = 0 + dilation_d: int = 1 + dilation_h: int = 1 + dilation_w: int = 1 + output_padding_d: int = 0 + output_padding_h: int = 0 + output_padding_w: int = 0 + groups: int = 1 + + +@dataclass +class GeluNode: + x: Tid + out: Tid + approximate: str + + +@dataclass +class ARangeNode: + out: Tid + start: IntOrVid + stop: IntOrVid + step: IntOrVid + scalar_type: int = None + + +@dataclass +class SiluNode: + x: Tid + out: Tid + + +@dataclass +class SigmoidNode: + x: Tid + out: Tid + + +@dataclass +class TanhNode: + x: Tid + out: Tid + + +@dataclass +class SqueezeNode: + x: Tid + out: Tid + dims: Optional[List[int]] = None + + +@dataclass +class SplitNode: + x: Tid + outs: List[Tid] + sizes: List[IntOrVid] + axis: Optional[int] = None + + +@dataclass +class RsqrtNode: + x: Tid + out: Tid + + +@dataclass +class MaximumNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class MinimumNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class LogNode: + x: Tid + out: Tid + + +@dataclass +class SoftmaxNode: + x: Tid + out: Tid + precise: bool = False + axis: Optional[int] = None + + +@dataclass +class BroadcastToNode: + x: Tid + out: Tid + shape: List[IntOrVid] + + +@dataclass +class PadNode: + x: Tid + out: Tid + pad_width: List[IntOrVid] + mode: str + constant_value: float = 0.0 + + +@dataclass +class WhereNode: + condition: Tid + x: Tid + y: Tid + out: Tid + + +@dataclass +class ReshapeNode: + x: Tid + out: Tid + shape: List[IntOrVid] + + +@dataclass +class TransposeNode: + x: Tid + out: Tid + perm: List[int] + + +@dataclass +class AsStridedNode: + x: Tid + out: Tid + shape: List[IntOrVid] + strides: List[IntOrVid] + offset: int = 0 + + +@dataclass +class ContiguousNode: + x: Tid + out: Tid + + +@dataclass +class GatherNode: + x: Tid + indices: List[Tid] + out: Tid + axes: List[int] + slice_sizes: List[int] + + +@dataclass +class SliceNode: + x: Tid + out: Tid + axis: IntOrVid + start: IntOrVid + stop: IntOrVid + step: int = 1 + + +@dataclass +class AsTypeNode: + x: Tid + out: Tid + scalar_type: Optional[int] = None + + +@dataclass +class QuantizedMatmulNode: + x: Tid + w: Tid + scales: Tid + out: Tid + mode: str + transpose: bool = True + biases: Optional[Tid] = None + group_size: Optional[int] = None + bits: Optional[int] = None + + +@dataclass +class ScatterAddNode: + x: Tid + indices: Tid + updates: Tid + out: Tid + axis: Optional[int] = None + + +@dataclass +class ConcatenateNode: + tensors: List[Tid] + out: Tid + axis: Optional[int] = None + + +@dataclass +class FullNode: + out: Tid + shape: List[IntOrVid] + v: FloatOrVid + scalar_type: Optional[int] = None + + +@dataclass +class FullLikeNode: + x: Tid + out: Tid + v: FloatOrVid + scalar_type: int = None + + +@dataclass +class ArgmaxNode: + x: Tid + out: Tid + keepdims: bool = False + axis: Optional[int] = None + + +@dataclass +class SliceUpdateNode: + dst: Tid + update: Tid + out: Tid + axis: IntOrVid + start: IntOrVid + stop: IntOrVid + step: int = 1 + + +@dataclass +class IndexCopyNode: + dst: Tid + update: Tid + indices: Tid + out: Tid + axis: Optional[int] = None + + +@dataclass +class DequantizeNode: + w: Tid + scales: Tid + out: Tid + mode: str + dtype: int = None + biases: Optional[Tid] = None + group_size: Optional[int] = None + bits: Optional[int] = None + global_scale: Optional[Tid] = None + + +@dataclass +class LessNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class LessEqualNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class GreaterNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class GreaterEqualNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class EqualNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class NotEqualNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class LogicalNotNode: + x: Tid + out: Tid + + +@dataclass +class BitwiseInvertNode: + x: Tid + out: Tid + + +@dataclass +class LogicalAndNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class LogicalOrNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class BitwiseAndNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class BitwiseOrNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class BitwiseXorNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class TriNode: + out: Tid + n: IntOrVid + m: IntOrVid + k: int = 0 + scalar_type: Optional[int] = None + + +@dataclass +class TrilNode: + x: Tid + out: Tid + k: int = 0 + + +@dataclass +class TriuNode: + x: Tid + out: Tid + k: int = 0 + + +@dataclass +class ClipNode: + x: Tid + out: Tid + a_min: Optional[Tid] = None + a_max: Optional[Tid] = None + + +@dataclass +class CumsumNode: + x: Tid + out: Tid + reverse: bool = False + inclusive: bool = True + axis: Optional[int] = None + + +@dataclass +class StackNode: + tensors: List[Tid] + out: Tid + axis: int = 0 + + +@dataclass +class SignNode: + x: Tid + out: Tid + + +@dataclass +class AnyNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class AllNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class RepeatNode: + x: Tid + out: Tid + repeats: IntOrVid + axis: Optional[int] = None + + +@dataclass +class SortNode: + x: Tid + out: Tid + axis: Optional[int] = None + + +@dataclass +class ArgsortNode: + x: Tid + out: Tid + axis: Optional[int] = None + + +@dataclass +class PartitionNode: + x: Tid + out: Tid + kth: IntOrVid + axis: Optional[int] = None + + +@dataclass +class ArgPartitionNode: + x: Tid + out: Tid + kth: IntOrVid + axis: Optional[int] = None + + +@dataclass +class RollNode: + x: Tid + out: Tid + shift: List[IntOrVid] + axes: List[int] + + +@dataclass +class FloorNode: + x: Tid + out: Tid + + +@dataclass +class CeilNode: + x: Tid + out: Tid + + +@dataclass +class SquareNode: + x: Tid + out: Tid + + +@dataclass +class ExpNode: + x: Tid + out: Tid + + +@dataclass +class SinNode: + x: Tid + out: Tid + + +@dataclass +class CosNode: + x: Tid + out: Tid + + +@dataclass +class TanNode: + x: Tid + out: Tid + + +@dataclass +class ArcsinNode: + x: Tid + out: Tid + + +@dataclass +class ArccosNode: + x: Tid + out: Tid + + +@dataclass +class ArctanNode: + x: Tid + out: Tid + + +@dataclass +class SinhNode: + x: Tid + out: Tid + + +@dataclass +class CoshNode: + x: Tid + out: Tid + + +@dataclass +class ArcsinhNode: + x: Tid + out: Tid + + +@dataclass +class ArccoshNode: + x: Tid + out: Tid + + +@dataclass +class ArctanhNode: + x: Tid + out: Tid + + +@dataclass +class Log2Node: + x: Tid + out: Tid + + +@dataclass +class Log10Node: + x: Tid + out: Tid + + +@dataclass +class Log1pNode: + x: Tid + out: Tid + + +@dataclass +class ErfNode: + x: Tid + out: Tid + + +@dataclass +class Expm1Node: + x: Tid + out: Tid + + +@dataclass +class RoundNode: + x: Tid + out: Tid + decimals: int = 0 + + +@dataclass +class ReciprocalNode: + x: Tid + out: Tid + + +@dataclass +class SqrtNode: + x: Tid + out: Tid + + +@dataclass +class AbsNode: + x: Tid + out: Tid + + +@dataclass +class NegNode: + x: Tid + out: Tid + + +@dataclass +class Atan2Node: + a: Tid + b: Tid + out: Tid + + +@dataclass +class LogAddExpNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class FloorDivideNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class RemainderNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class PowerNode: + a: Tid + b: Tid + out: Tid + + +@dataclass +class LogSumExpNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class SumNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class MeanNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class VarNode: + x: Tid + out: Tid + keepdims: bool = False + ddof: int = 0 + axes: Optional[List[int]] = None + + +@dataclass +class StdNode: + x: Tid + out: Tid + keepdims: bool = False + ddof: int = 0 + axes: Optional[List[int]] = None + + +@dataclass +class ProdNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class MaxNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class MinNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class ArgminNode: + x: Tid + out: Tid + keepdims: bool = False + axis: Optional[int] = None + + +@dataclass +class MedianNode: + x: Tid + out: Tid + keepdims: bool = False + axes: Optional[List[int]] = None + + +@dataclass +class GatherMmNode: + a: Tid + b: Tid + out: Tid + sorted_indices: bool = False + lhs_indices: Optional[Tid] = None + rhs_indices: Optional[Tid] = None + + +@dataclass +class GatherQmmNode: + x: Tid + w: Tid + scales: Tid + out: Tid + mode: str + transpose: bool = True + sorted_indices: bool = False + biases: Optional[Tid] = None + lhs_indices: Optional[Tid] = None + rhs_indices: Optional[Tid] = None + group_size: Optional[int] = None + bits: Optional[int] = None + + +@dataclass +class ScanNode: + originals: List[Tid] + sliced: List[Tid] + outputs: List[Tid] + carry: List[Tid] + scan_axis: int = 1 + body_chain_idx: Optional[int] = None + + +@dataclass +class IfNode: + cond: IntOrVid + then_chain_idx: Optional[int] = None + else_chain_idx: Optional[int] = None + + +@dataclass +class RandomBitsNode: + out: Tid + shape: List[IntOrVid] + width: int = 4 + seed: Optional[Vid] = None + + +@dataclass +class MetalKernelNode: + name: str + source: str + inputs: List[Tid] + outputs: List[Tid] + grid: List[IntOrVid] + threadgroup: List[IntOrVid] + ensure_row_contiguous: bool = True + atomic_outputs: bool = False + init_value: float = None + header: Optional[str] = None + input_names: Optional[List[str]] = None + output_names: Optional[List[str]] = None + output_shapes_flat: Optional[List[IntOrVid]] = None + output_shape_lengths: Optional[List[int]] = None + output_dtypes: Optional[List[int]] = None + template_arg_names: Optional[List[str]] = None + template_arg_kinds: Optional[List[int]] = None + template_arg_values: Optional[List[int]] = None + + +# Union of all op types +OpNodeUnion = Union[ + NoopNode, + IdCopyNode, + AddmmNode, + ItemIntNode, + ExpandDimsNode, + TileNode, + TakeAlongAxisNode, + TakeNode, + RMSNormNode, + LayerNormNode, + RopeNode, + SdpaNode, + AddNode, + AddIntNode, + SubtractIntNode, + MultiplyIntNode, + FloorDivideIntNode, + ModIntNode, + SymSizeNode, + MultiplyNode, + DivideNode, + SubtractNode, + Conv1DNode, + Conv2DNode, + Conv3DNode, + ConvTranspose1DNode, + ConvTranspose2DNode, + ConvTranspose3DNode, + GeluNode, + ARangeNode, + SiluNode, + SigmoidNode, + TanhNode, + SqueezeNode, + SplitNode, + RsqrtNode, + MaximumNode, + MinimumNode, + LogNode, + SoftmaxNode, + BroadcastToNode, + PadNode, + WhereNode, + ReshapeNode, + TransposeNode, + AsStridedNode, + ContiguousNode, + GatherNode, + SliceNode, + AsTypeNode, + QuantizedMatmulNode, + ScatterAddNode, + ConcatenateNode, + FullNode, + FullLikeNode, + ArgmaxNode, + SliceUpdateNode, + IndexCopyNode, + DequantizeNode, + LessNode, + LessEqualNode, + GreaterNode, + GreaterEqualNode, + EqualNode, + NotEqualNode, + LogicalNotNode, + BitwiseInvertNode, + LogicalAndNode, + LogicalOrNode, + BitwiseAndNode, + BitwiseOrNode, + BitwiseXorNode, + TriNode, + TrilNode, + TriuNode, + ClipNode, + CumsumNode, + StackNode, + SignNode, + AnyNode, + AllNode, + RepeatNode, + SortNode, + ArgsortNode, + PartitionNode, + ArgPartitionNode, + RollNode, + FloorNode, + CeilNode, + SquareNode, + ExpNode, + SinNode, + CosNode, + TanNode, + ArcsinNode, + ArccosNode, + ArctanNode, + SinhNode, + CoshNode, + ArcsinhNode, + ArccoshNode, + ArctanhNode, + Log2Node, + Log10Node, + Log1pNode, + ErfNode, + Expm1Node, + RoundNode, + ReciprocalNode, + SqrtNode, + AbsNode, + NegNode, + Atan2Node, + LogAddExpNode, + FloorDivideNode, + RemainderNode, + PowerNode, + LogSumExpNode, + SumNode, + MeanNode, + VarNode, + StdNode, + ProdNode, + MaxNode, + MinNode, + ArgminNode, + MedianNode, + GatherMmNode, + GatherQmmNode, + ScanNode, + IfNode, + RandomBitsNode, + MetalKernelNode, +] + +# ============================================================================ +# Container types (reference OpNodeUnion) +# ============================================================================ + +@dataclass +class Instruction: + op: OpNodeUnion + + +@dataclass +class InstructionChain: + instructions: List[Instruction] + + +@dataclass +class MLXGraph: + instruction_chains: List[InstructionChain] + version: Optional[str] = None + num_constant_tensors: int = 0 + num_input_tensors: int = 0 + num_output_tensors: int = 0 + num_mutable_buffer_tensors: int = 0 + num_temp_tensors: int = 0 + num_values: int = 0 + main_chain_idx: int = 0 + init_chain_idx: int = -1 + input_map: Optional[List[SlotVariant]] = None + output_map: Optional[List[SlotVariant]] = None + mutable_buffer_map: Optional[List[SlotVariant]] = None + named_slots: Optional[List[NamedSlot]] = None + tensor_meta: Optional[List[TensorMeta]] = None diff --git a/backends/mlx/third-party/mlx b/backends/mlx/third-party/mlx index 7a1d4f5c12a..ce45c52505c 160000 --- a/backends/mlx/third-party/mlx +++ b/backends/mlx/third-party/mlx @@ -1 +1 @@ -Subproject commit 7a1d4f5c12ac82f4b4d0a6e71538d89ca0605247 +Subproject commit ce45c52505c8158ea48d2a54e8caae05efd86bfe diff --git a/examples/models/qwen3/DFLASH_EXPERIMENTS.md b/examples/models/qwen3/DFLASH_EXPERIMENTS.md new file mode 100644 index 00000000000..e5f8907c7b5 --- /dev/null +++ b/examples/models/qwen3/DFLASH_EXPERIMENTS.md @@ -0,0 +1,81 @@ +Written By: Chetan Thotti (cthotti) +Date: 08/16/2026 + +This is a record of the benchmarking we did on DFlash speculative decoding +for Qwen3-4B, across three Apple Silicon machines. The short +version: **DFlash's speedup depends on GPU architecture generation, not on +how big or fast the chip otherwise is.** A base M4 clearly beats an M2 Pro +here, even though the M2 Pro has more GPU cores and more memory bandwidth. +If you're benchmarking DFlash on new hardware, read this first so you +don't waste time on a chip that was never going to show a speedup. + +## The setup + +Model was `Qwen/Qwen3-4B` with the `z-lab/Qwen3-4B-DFlash-b16` draft +checkpoint, exported with: +--dflash-layers 1,9,17,25,33 --qlinear 4w --qembedding 4w --use-custom-sdpa --use-custom-kv-cache + +We tested three chips: the M2 in a MacBook Air (8 GPU cores), an M2 Pro +rental (16 GPU cores), and a base M4 rental (10 GPU cores). Same `.pte` +files got copied across the M2 Pro and M4 runs rather than re-exported +separately, so any difference we saw was purely hardware, not export +drift. + +## What we expected vs. what we found + +Going in, the assumption was that a "bigger" chip -- more GPU cores, more +memory bandwidth -- would just be faster across the board, M2 Pro +included. That's not what happened. Baseline (plain, one-token-at-a-time) +decoding did scale the way you'd expect: M2 Pro's extra bandwidth made it +faster than the M4 at baseline decoding, in every category we tested. +But DFlash flipped that around entirely, only the M4 ever beat its own +baseline. The M2 Pro was slower with DFlash turned on than without it, +every single time. + +| Chip | Category | Baseline tok/s | DFlash tok/s | Speedup | +|--------|------|-------|-------|-------| +| M2 Pro | Math | 48.30 | 42.63 | 0.88x | +| M2 Pro | Code | 51.65 | 42.84 | 0.83x | +| M2 Pro | Chat | 51.29 | 18.22 | 0.36x | +| M4 | Math | 31.71 | 51.44 | 1.62x | +| M4 | Code | 31.33 | 53.39 | 1.70x | +| M4 | Chat | 31.42 | 22.73 | 0.72x | + +(Math/Code/Chat here are three different prompts, run 3 times each and +averaged.) + +## The main difference between M2 and M4 + +The performance difference comes down to GPU architecture, not the CPU. +While I initially suspected SME2, the MLX backend doesn't use it. Instead, +M3/M4 GPUs (Apple9) introduced Dynamic Caching and an improved SIMD matrix- +multiply pipeline, which are much better suited for DFlash's verification +stage. Dflash verifies an entire block of tokens at once using large matrix- +matrix operations, allowing the M4 to execute this workload far more efficiently +than the M2's Apple8 GPU, which lacks these architectural improvements. + +**Practical takeaway: DFlash is worth using on M3 or M4 generation Macs, +any tier, but not on an M1/M2.** + +## The draft model is also just bad at chat + +Separately from all the hardware stuff: math and code prompts got a tau +(average tokens accepted per speculative round) around 6.6-6.8 on both +chips. Chat-style prompts landed around 2.9 -- consistently, across three +different chat prompts we tried, not just one unlucky example. Since tau +was identical across both chips for the same prompt, this isn't a +hardware issue at all, it's the draft model itself being noticeably +worse at predicting open-ended conversational text than it is at +structured math or code. Worth knowing if you're deciding whether DFlash +is worth turning on for a particular kind of workload, independent of +what hardware you're running it on. + +## What we didn't get to + +- **8-bit target quantization** (`--qlinear 8w`) fails to export with + `RuntimeError: Missing out variants: {'torchao::dequantize_affine'}`. + That's a real gap in the MLX partitioner's op coverage, not a flag + mistake. +- **M3-generation chips** were never tested directly. Based on the + mechanism above they should behave like the M4 (same Apple9 GPU + family), but that's inference, not something we confirmed ourselves. diff --git a/examples/models/qwen3/SUMMARY.txt b/examples/models/qwen3/SUMMARY.txt new file mode 100644 index 00000000000..ebccf33fd33 --- /dev/null +++ b/examples/models/qwen3/SUMMARY.txt @@ -0,0 +1,14 @@ +Running on math prompt: + baseline_runs=[25.73, 24.63, 25.65] + baseline_median=25.65 + dflash_runs=[19.44, 19.47, 19.44] + dflash_median=19.44 tau=5.59 speedup=0.76x +Running on code prompt: + baseline_runs=[28.28, 28.19, 26.1] + baseline_median=28.19 dflash_runs=[23.84, 23.81, 22.7] + dflash_median=23.81 tau=6.81 speedup=0.84x +Running on chat prompt: + baseline_runs=[26.05, 26.11, 27.98] + baseline_median=26.11 + dflash_runs=[14.78, 14.77, 14.75] + dflash_median=14.77 tau=4.17 speedup=0.57x