Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
72 changes: 48 additions & 24 deletions emlx/lib/emlx.ex
Original file line number Diff line number Diff line change
Expand Up @@ -1851,46 +1851,57 @@ defmodule EMLX do

@behaviour Nx.Defn.Compiler

# Known EMLX-specific compiler opts. `:command_queue` is injected by
# `__partitions_options__/1` but may also be passed directly by callers
# that manage their own queues (equivalent to a manual `with_queue`).
@valid_compiler_keys [:device, :max_concurrency, :command_queue, :hooks]

# Process-lifetime dispatch cache backing `dispatch_key/3` +
# `get_or_compile_program/6` (see their docs) — a compiled program is keyed
# by a *structural* signature of its `Expr` (not object identity), so
# it survives across `Nx.Defn.Graph.run/3`'s per-call re-tracing and is
# shared across structurally-identical call sites (e.g. every one of
# Qwen3's 28 attention layers), not just within one closure's lifetime.
# Structural program cache for `get_or_compile_program/6`.
@native_dispatch_cache_table :emlx_native_dispatch_cache

# A second, process-lifetime cache in front of `dispatch_key/3`'s own
# (expensive — O(nodes), plus per-opaque-scope SHA256 hashing) structural
# walk, keyed by `output_expr`'s own node identity rather than its
# structural signature. This matters for `run_while_loop/3`'s host-driven
# `cond_fn`/`body_fn`: `Nx.Defn.jit/2` retraces `fn _ -> body_expr end`
# once and caches *that* trace by argument template, so every subsequent
# call re-enters `__jit__`/`build_eval_fn` with the *exact same* `Expr`
# (identical ids, not just structurally identical) — walking it again on
# every decode step is pure waste. (This is unlike `Nx.Defn.Graph.run/3`'s
# per-stage re-tracing, which *does* mint fresh ids each call — hence
# `dispatch_key/3` still has to fall back to the structural walk on a miss.)
# Expr-id cache in front of `dispatch_key/3`.
@dispatch_key_by_id_table :emlx_dispatch_key_by_id

# `native_compile/3` closures. Created by `init/0`.
@compile_closure_table :emlx_compile_closures

@doc false
def init do
case :ets.whereis(@compile_closure_table) do
:undefined ->
:ets.new(@compile_closure_table, [
:named_table,
:public,
:set,
read_concurrency: true,
write_concurrency: true
])

_ ->
:ok
end
end

@impl Nx.Defn.Compiler
def __jit__(key, vars, fun, args_list, opts) do
__compile__(key, vars, fun, opts).(args_list)
end

defp compile_cache_key(key, vars, hooks, device) do
templates =
Enum.map(vars, fn var -> Nx.Defn.Composite.traverse(var, &Nx.to_template/1) end)

{key, templates, hooks, device}
end

@impl Nx.Defn.Compiler
def __compile__(_key, vars, fun, opts) do
def __compile__(key, vars, fun, opts) do
Keyword.validate!(opts, @valid_compiler_keys)

case Keyword.get(opts, :hooks, %{}) do
hooks = Keyword.get(opts, :hooks, %{})

case hooks do
empty when empty == %{} ->
:ok

hooks ->
_ ->
raise ArgumentError,
"EMLX does not support the :hooks named-override map (got callbacks for " <>
"#{inspect(Map.keys(hooks))}) — :hook/:io_call expr nodes lower natively " <>
Expand All @@ -1901,7 +1912,20 @@ defmodule EMLX do
queue = Keyword.get(opts, :command_queue)
device = Keyword.get(opts, :device, default_device())

wrap_with_queue(queue, native_compile(vars, fun, device))
cache_key = compile_cache_key(key, vars, hooks, device)

eval_fn =
case :ets.lookup(@compile_closure_table, cache_key) do
[{^cache_key, cached}] ->
cached

[] ->
built = native_compile(vars, fun, device)
:ets.insert(@compile_closure_table, {cache_key, built})
built
end

wrap_with_queue(queue, eval_fn)
end

# Attempts to lower `fun.(vars)` to an `EMLX.Native.Expr` program and build a
Expand Down
1 change: 1 addition & 0 deletions emlx/lib/emlx/application.ex
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@ defmodule EMLX.Application do
@doc false
def start(_type, _args) do
EMLX.Profiling.init()
EMLX.init()
ensure_default_worker!(:cpu, _gpu_optional? = false)
ensure_default_worker!(:gpu, _gpu_optional? = true)
ensure_worker!(:runtime_call_worker, :cpu, _gpu_optional? = false)
Expand Down
Loading