Compare commits

..
Author SHA1 Message Date
rltakashige fd17a3bc09 vllm loads! 2026-03-27 12:02:36 +00:00
rltakashige 7ca9edaf42 wohohoh 2026-03-25 10:35:46 +00:00
Evan d1f13f21a1 time to fiddle 2026-03-24 11:54:10 +00:00
Evan a6cdc93d3c disable 2026-03-24 11:54:04 +00:00
Evan a9c7b1c68a urgk 2026-03-24 11:54:04 +00:00
Evan 6355f3d8fb man i gotta commit 2026-03-23 16:00:57 +00:00
Evan 8cd6191c70 gotta move the cache soon 2026-03-21 22:22:25 +00:00
Evan 532a8f0b07 we got it to BUILD 2026-03-20 15:38:16 +00:00
Evan 581a1fcd79 move chunks to core 2026-03-19 16:18:40 +00:00
Evan ac7194dd90 mupdate 2026-03-19 16:01:35 +00:00
Evan f538b44211 shmovin 2026-03-19 14:51:08 +00:00
Evan a3ce437fd4 first 2026-03-19 11:10:58 +00:00
Ryuichi Leo Takashige be731d3a85 Merge main 2026-03-17 19:05:55 +00:00
Ryuichi Leo Takashige 655185cfe7 Address comments 4 - defer to the warmup into the exo batch generator and vllm batch engine and don't store model on the generators. 2026-03-17 19:00:12 +00:00
Ryuichi Leo Takashige 1dd9c28842 Address comments 3, mainly refactors 2026-03-17 18:31:41 +00:00
Ryuichi Leo Takashige cacd26e63c Type error lol 2026-03-17 18:00:39 +00:00
Ryuichi Leo Takashige 6a3eb2f37d close() 2026-03-17 17:31:15 +00:00
Ryuichi Leo Takashige e1df77bc4c No more future annotations 2026-03-17 17:07:32 +00:00
Ryuichi Leo Takashige e78e53df6e Address comments 2 including vllm capability in state 2026-03-17 16:55:08 +00:00
Ryuichi Leo Takashige c70d9006e8 Address comments including task id interface 2026-03-17 15:31:55 +00:00
Ryuichi Leo Takashige 72cd8552ae Merge branch 'main' into leo/dgx-spark-integrations 2026-03-17 13:58:49 +00:00
Ryuichi Leo Takashige 8cd1308336 Tidy pass 1 2026-03-17 00:38:16 +00:00
Ryuichi Leo Takashige 04dcdbd127 Merge main 2026-03-16 23:01:48 +00:00
Ryuichi Leo Takashige ec5d62f935 Strip vllm generator 2026-03-16 22:11:35 +00:00
Ryuichi Leo Takashige e96f084051 Distributed callbacks 2026-03-16 21:17:14 +00:00
Ryuichi Leo Takashige dc68ddbac0 Prompt formatting 2026-03-16 20:37:38 +00:00
Ryuichi Leo Takashige 073f8c1690 add batching 2026-03-16 19:25:45 +00:00
Ryuichi Leo Takashige 3c29d0dd4c test prefix caching 2026-03-16 16:43:51 +00:00
rltakashige 594ed99734 new uv lock for fastsafetensors 2026-03-13 17:08:49 +00:00
Ryuichi Leo Takashige e9e23e556e Have loading progress 2026-03-13 12:58:34 +00:00
Ryuichi Leo Takashige 169ea2a5e8 Fix GPT OSS by not retokenizing prompts 2026-03-12 18:02:43 +00:00
Ryuichi Leo Takashige 4a7901c548 Fix patches 2026-03-12 17:42:02 +00:00
Ryuichi Leo Takashige 7bb5cb4fc7 Allow memory profiling to be unstable 2026-03-12 17:34:53 +00:00
Ryuichi Leo Takashige 493e342f83 Skip impossible shardings 2026-03-12 17:33:12 +00:00
Ryuichi Leo Takashige 283b1809c9 Move VLLM runner into VLLM engine 2026-03-12 17:20:30 +00:00
Ryuichi Leo Takashige 35030119e3 ExoBench and ExoEval for CUDA 2026-03-12 17:09:51 +00:00
rltakashige 585dfe3549 Merge branch 'main' into leo/dgx-spark-integrations 2026-03-12 15:41:29 +00:00
Ryuichi Leo Takashige 5d7a005a13 Destroy process group on Keyboard Interrupt 2026-03-12 15:33:32 +00:00
Ryuichi Leo Takashige 87b7c5ef8b Pass CI 2026-03-12 15:24:49 +00:00
Ryuichi Leo Takashige 957ebbd21f Set max token length as max context length if no max tokens set 2026-03-12 15:15:39 +00:00
rltakashige 4b6dd7588f lockgit status 2026-03-12 14:41:33 +00:00
Ryuichi Leo Takashige 3f4f7c9ba6 Only do for aarch64 linux 2026-03-12 13:42:23 +00:00
Ryuichi Leo Takashige 1331465ba0 Add missing runner features 2026-03-12 10:51:37 +00:00
Ryuichi Leo Takashige 8f94727f14 Ignore missing modules if type stubs exist 2026-03-12 00:20:28 +00:00
Ryuichi Leo Takashige 9ee23ee0d3 Make vllm inference runner closer to the normal inference runner 2026-03-11 23:41:34 +00:00
Ryuichi Leo Takashige f75d36cbe0 Fix cache patch 2026-03-11 22:28:45 +00:00
Ryuichi Leo Takashige 2683ac7b61 Add Torch typings 2026-03-11 21:44:10 +00:00
Ryuichi Leo Takashige 404b9769ac Download models without model.safetensors.index 2026-03-11 21:32:07 +00:00
Ryuichi Leo Takashige 3e097f7243 only download a single copy of the model. 2026-03-11 20:57:29 +00:00
Ryuichi Leo Takashige 0c8615f25c Only import VLLM once.. 2026-03-11 20:49:19 +00:00
Ryuichi Leo Takashige ba35a4ba13 Move VLLM into the runner and add type stubs 2026-03-11 19:15:55 +00:00
Ryuichi Leo Takashige 9a83fa6cdf Patch VLLM to load multiple models dynamically 2026-03-11 18:17:59 +00:00
Ryuichi Leo Takashige 659c1bc737 Progress: Run EXO-CUDA through nix! 2026-03-11 18:17:59 +00:00
Ryuichi Leo Takashige 34df811b92 Vibe coding design baby 2026-03-11 18:17:59 +00:00
Ryuichi Leo Takashige ca5870a2e8 Some Linux Laptop/Desktop detection and goodbye penguin 2026-03-11 18:17:59 +00:00
Ryuichi Leo Takashige 6be6ea5fd2 Fix placement preview 2026-03-11 18:17:59 +00:00
Ryuichi Leo Takashige 5e9d27b753 Show Sparks and Linux in topology 2026-03-11 18:17:59 +00:00
Ryuichi Leo Takashige cfc8f09004 Fast direct USB connectivity 2026-03-11 18:17:59 +00:00
750 changed files with 14497 additions and 32965 deletions

No files matched your search

+20
View File
@@ -0,0 +1,20 @@
from enum import Enum
class HarmonyEncodingName(Enum):
HARMONY_GPT_OSS = ...
class HarmonyEncoding: ...
class HarmonyError(Exception): ...
class Role(Enum):
ASSISTANT = ...
class StreamableParser:
last_content_delta: str
current_channel: str | None
current_recipient: str | None
def __init__(self, encoding: HarmonyEncoding, role: Role = ...) -> None: ...
def process(self, token_id: int) -> None: ...
def load_harmony_encoding(name: HarmonyEncodingName) -> HarmonyEncoding: ...
+17
View File
@@ -0,0 +1,17 @@
class NvmlMemoryInfo:
used: int
total: int
free: int
class NvmlUtilizationRates:
gpu: int
memory: int
def nvmlInit() -> None: ...
def nvmlShutdown() -> None: ...
def nvmlDeviceGetCount() -> int: ...
def nvmlDeviceGetHandleByIndex(index: int) -> object: ...
def nvmlDeviceGetUtilizationRates(handle: object) -> NvmlUtilizationRates: ...
def nvmlDeviceGetTemperature(handle: object, sensor_type: int) -> int: ...
def nvmlDeviceGetPowerUsage(handle: object) -> int: ...
def nvmlDeviceGetMemoryInfo(handle: object) -> NvmlMemoryInfo: ...
+61
View File
@@ -0,0 +1,61 @@
from typing import Any, Sequence
from torch import backends as backends
from torch import cuda as cuda
from torch import distributed as distributed
__version__: str
class version:
cuda: str
class dtype: ...
bfloat16: dtype
float16: dtype
float32: dtype
int8: dtype
int32: dtype
int64: dtype
long: dtype
float8_e4m3fn: dtype
class Tensor:
shape: Sequence[int]
dtype: dtype
def __getitem__(self, key: Any) -> Tensor: ...
def __setitem__(self, key: Any, value: Any) -> None: ...
def to(self, *args: Any, **kwargs: Any) -> Tensor: ...
def cpu(self) -> Tensor: ...
def detach(self) -> Tensor: ...
def clone(self) -> Tensor: ...
def flatten(self, start_dim: int = 0, end_dim: int = -1) -> Tensor: ...
def view(self, *shape: Any) -> Tensor: ...
def squeeze(self, dim: int = ...) -> Tensor: ...
def unsqueeze(self, dim: int) -> Tensor: ...
def permute(self, *dims: int) -> Tensor: ...
def float(self) -> Tensor: ...
def numpy(self) -> Any: ...
def numel(self) -> int: ...
def nelement(self) -> int: ...
@property
def is_cuda(self) -> bool: ...
@property
def device(self) -> device: ...
def __len__(self) -> int: ...
def data_ptr(self) -> int: ...
def tolist(self) -> Any: ...
def abs(self) -> Tensor: ...
def max(self) -> Tensor: ...
def mean(self) -> Tensor: ...
def sum(self, dim: int = ...) -> Tensor: ...
def item(self) -> float: ...
def tensor(data: Any, dtype: dtype | None = None, device: Any = None) -> Tensor: ...
def zeros(*size: Any, dtype: dtype | None = None, device: Any = None) -> Tensor: ...
def empty(*size: Any, dtype: dtype | None = None, device: Any = None) -> Tensor: ...
def from_numpy(ndarray: Any) -> Tensor: ...
def inference_mode() -> Any: ...
class device:
def __init__(self, type: str, index: int = ...) -> None: ...
@@ -0,0 +1 @@
from torch.backends import cuda as cuda
@@ -0,0 +1 @@
def is_built() -> bool: ...
+10
View File
@@ -0,0 +1,10 @@
class _DeviceProperties:
total_memory: int
def is_available() -> bool: ...
def get_device_name(device: int) -> str: ...
def get_device_properties(device: int) -> _DeviceProperties: ...
def empty_cache() -> None: ...
def mem_get_info() -> tuple[int, int]: ...
def synchronize() -> None: ...
def max_memory_allocated() -> int: ...
@@ -0,0 +1,2 @@
def is_initialized() -> bool: ...
def destroy_process_group() -> None: ...
+1
View File
@@ -0,0 +1 @@
__version__: str
+2
View File
@@ -0,0 +1,2 @@
class ModelConfig:
max_model_len: int
File renamed without changes.
+18
View File
@@ -0,0 +1,18 @@
from dataclasses import dataclass
@dataclass
class EngineArgs:
model: str = ...
served_model_name: str | list[str] | None = ...
tokenizer: str | None = ...
trust_remote_code: bool = ...
dtype: str = ...
seed: int = ...
max_model_len: int | None = ...
gpu_memory_utilization: float = ...
enforce_eager: bool = ...
tensor_parallel_size: int = ...
pipeline_parallel_size: int = ...
quantization: str | None = ...
load_format: str = ...
enable_sleep_mode: bool = ...
+17
View File
@@ -0,0 +1,17 @@
class CompletionOutput:
index: int
text: str
token_ids: list[int]
cumulative_logprob: float | None
logprobs: object | None
finish_reason: str | None
stop_reason: int | str | None
def finished(self) -> bool: ...
class RequestOutput:
request_id: str
prompt: str | None
prompt_token_ids: list[int] | None
outputs: list[CompletionOutput]
finished: bool
+11
View File
@@ -0,0 +1,11 @@
class SamplingParams:
n: int
temperature: float
top_p: float
top_k: int
min_p: float
seed: int | None
stop: str | list[str] | None
max_tokens: int | None
logprobs: int | None
repetition_penalty: float
@@ -0,0 +1,3 @@
from vllm.tokenizers.protocol import TokenizerLike
__all__ = ["TokenizerLike"]
@@ -0,0 +1,15 @@
from typing import Protocol
class TokenizerLike(Protocol):
@property
def eos_token_id(self) -> int: ...
@property
def vocab_size(self) -> int: ...
def encode(self, text: str, add_special_tokens: bool = ...) -> list[int]: ...
def decode(self, ids: list[int] | int, skip_special_tokens: bool = ...) -> str: ...
def apply_chat_template(
self,
messages: list[dict[str, str]],
tools: list[dict[str, object]] | None = ...,
**kwargs: object,
) -> str | list[int]: ...
+1
View File
@@ -0,0 +1 @@
+1
View File
@@ -0,0 +1 @@
@@ -0,0 +1,24 @@
from collections.abc import Sequence
from vllm.v1.core.kv_cache_utils import BlockPool, KVCacheBlock
from vllm.v1.kv_cache_interface import KVCacheConfig
class KVCacheBlocks:
blocks: tuple[Sequence[KVCacheBlock], ...]
def __init__(self, blocks: tuple[Sequence[KVCacheBlock], ...]) -> None: ...
def get_block_ids(self) -> tuple[list[int], ...]: ...
class KVCacheManager:
block_pool: BlockPool
kv_cache_config: KVCacheConfig
enable_caching: bool
num_kv_cache_groups: int
coordinator: object
def __init__(self, *args: object, **kwargs: object) -> None: ...
def allocate_slots(
self, request: object, num_new_tokens: int, *args: object, **kwargs: object
) -> KVCacheBlocks | None: ...
def get_computed_blocks(self, request: object) -> tuple[KVCacheBlocks, int]: ...
def create_kv_cache_blocks(
self, blocks: tuple[list[KVCacheBlock], ...]
) -> KVCacheBlocks: ...
@@ -0,0 +1,16 @@
class KVCacheBlock:
block_id: int
ref_cnt: int
def __init__(self, block_id: int) -> None: ...
class FreeKVCacheBlockQueue:
def append_n(self, blocks: list[KVCacheBlock]) -> None: ...
def popleft_n(self, n: int) -> list[KVCacheBlock]: ...
class BlockPool:
blocks: list[KVCacheBlock]
free_block_queue: FreeKVCacheBlockQueue
num_gpu_blocks: int
enable_caching: bool
def get_num_free_blocks(self) -> int: ...
def get_new_blocks(self, num_blocks: int) -> list[KVCacheBlock]: ...
File renamed without changes.
@@ -0,0 +1,22 @@
from vllm.config import ModelConfig
from vllm.engine.arg_utils import EngineArgs
from vllm.outputs import RequestOutput
from vllm.sampling_params import SamplingParams
from vllm.tokenizers import TokenizerLike
class LLMEngine:
tokenizer: TokenizerLike | None
model_config: ModelConfig
@classmethod
def from_engine_args(cls, engine_args: EngineArgs) -> LLMEngine: ...
def add_request(
self,
request_id: str,
prompt: str,
params: SamplingParams,
arrival_time: float | None = ...,
) -> None: ...
def step(self) -> list[RequestOutput]: ...
def has_unfinished_requests(self) -> bool: ...
def get_tokenizer(self) -> TokenizerLike: ...
@@ -0,0 +1,23 @@
from dataclasses import dataclass
@dataclass
class KVCacheSpec:
block_size: int
num_kv_heads: int
head_size: int
@dataclass
class KVCacheGroupSpec:
layer_names: list[str]
kv_cache_spec: KVCacheSpec
@dataclass
class KVCacheTensorSpec:
shared_by: list[str]
size: int
@dataclass
class KVCacheConfig:
num_blocks: int
kv_cache_groups: list[KVCacheGroupSpec]
kv_cache_tensors: list[KVCacheTensorSpec]
+6
View File
@@ -0,0 +1,6 @@
class Request:
request_id: str
prompt_token_ids: list[int] | None
num_prompt_tokens: int
num_computed_tokens: int
num_tokens: int
@@ -0,0 +1 @@
@@ -0,0 +1,24 @@
import torch
class _CompilationConfig:
static_forward_context: dict[str, object]
class _ModelConfig:
hf_config: object
class GPUModelRunner:
kv_caches: list[torch.Tensor]
compilation_config: _CompilationConfig
model_config: _ModelConfig | None
def _allocate_kv_cache_tensors(
self, kv_cache_config: object
) -> dict[str, torch.Tensor]: ...
def initialize_kv_cache_tensors(
self, kv_cache_config: object, kernel_block_sizes: list[int]
) -> dict[str, torch.Tensor]: ...
def _reshape_kv_cache_tensors(
self,
kv_cache_config: object,
raw_tensors: dict[str, torch.Tensor],
kernel_block_sizes: list[int],
) -> dict[str, torch.Tensor]: ...
@@ -0,0 +1,6 @@
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
class Worker:
model_runner: GPUModelRunner
def determine_available_memory(self) -> int: ...
def initialize_from_config(self, kv_cache_config: object) -> None: ...
+1
View File
@@ -0,0 +1 @@
def extract_layer_index(layer_name: str, num_attn_module: int) -> int: ...
-7
View File
@@ -1,8 +1 @@
use flake
# creates .venv if doesn't exist and loads its environment
export VIRTUAL_ENV=".venv"
if ! [ -d "./$VIRTUAL_ENV" ]; then
uv venv
fi
layout python
+4 -124
View File
@@ -32,6 +32,7 @@ jobs:
SPARKLE_ED25519_PRIVATE: ${{ secrets.SPARKLE_ED25519_PRIVATE }}
SPARKLE_S3_BUCKET: ${{ secrets.SPARKLE_S3_BUCKET }}
SPARKLE_S3_PREFIX: ${{ secrets.SPARKLE_S3_PREFIX }}
EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT: ${{ secrets.EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT }}
AWS_REGION: ${{ secrets.AWS_REGION }}
EXO_BUILD_NUMBER: ${{ github.run_number }}
EXO_LIBP2P_NAMESPACE: ${{ github.ref_name }}
@@ -158,7 +159,7 @@ jobs:
fi
- name: Install Homebrew packages
run: brew install just awscli
run: brew install just awscli macmon
- name: Install UV
uses: astral-sh/setup-uv@v6
@@ -238,92 +239,10 @@ jobs:
# Export keychain path for other steps
echo "BUILD_KEYCHAIN_PATH=$KEYCHAIN_PATH" >> $GITHUB_ENV
# ============================================================
# Pre-flight credential / profile validation
# Runs BEFORE the ~16 min build so auth/expiry failures surface in <1 min.
# ============================================================
- name: Validate Apple notarization credentials
env:
APPLE_NOTARIZATION_USERNAME: ${{ secrets.APPLE_NOTARIZATION_USERNAME }}
APPLE_NOTARIZATION_PASSWORD: ${{ secrets.APPLE_NOTARIZATION_PASSWORD }}
APPLE_NOTARIZATION_TEAM: ${{ secrets.APPLE_NOTARIZATION_TEAM }}
run: |
# All-or-nothing: either all three creds are set, or none are.
CRED_COUNT=0
for v in "$APPLE_NOTARIZATION_USERNAME" "$APPLE_NOTARIZATION_PASSWORD" "$APPLE_NOTARIZATION_TEAM"; do
[[ -n "$v" ]] && CRED_COUNT=$((CRED_COUNT + 1))
done
if [[ "$CRED_COUNT" -eq 0 ]]; then
echo "No notarization credentials configured — skipping notarization for this build."
exit 0
fi
if [[ "$CRED_COUNT" -ne 3 ]]; then
echo "ERROR: partial notarization credentials set ($CRED_COUNT/3). Aborting before build."
exit 1
fi
# Cheap, ~5s, auth-only call. Fails instantly with a clear message if
# the app-specific password is stale, wrong team-id, etc.
echo "Verifying Apple notarization credentials via notarytool history..."
if ! xcrun notarytool history \
--apple-id "$APPLE_NOTARIZATION_USERNAME" \
--password "$APPLE_NOTARIZATION_PASSWORD" \
--team-id "$APPLE_NOTARIZATION_TEAM" >/dev/null; then
echo "ERROR: notarytool rejected the provided credentials. Fix before rerunning."
echo "Common causes: app-specific password expired/revoked, wrong team-id,"
echo "Apple ID not on the team, or 2FA not configured for this Apple ID."
exit 1
fi
echo "Apple notarization credentials OK."
- name: Validate provisioning profile expiry
run: |
PROFILE="$HOME/Library/Developer/Xcode/UserData/Provisioning Profiles/EXO.provisionprofile"
if [[ ! -f "$PROFILE" ]]; then
echo "ERROR: provisioning profile not found at $PROFILE"
exit 1
fi
EXPIRY=$(security cms -D -i "$PROFILE" | plutil -extract ExpirationDate raw -o - - 2>/dev/null || true)
if [[ -z "$EXPIRY" ]]; then
echo "WARNING: could not read ExpirationDate from provisioning profile; skipping expiry check."
exit 0
fi
# Try a couple of known plutil date formats. If none parse, skip the check rather
# than risk a false-positive "expired" block on a format we didn't anticipate.
EXPIRY_EPOCH=""
for fmt in "%Y-%m-%dT%H:%M:%SZ" "%Y-%m-%d %H:%M:%S %z" "%Y-%m-%d %H:%M:%S +0000"; do
if parsed=$(date -j -f "$fmt" "$EXPIRY" +%s 2>/dev/null); then
EXPIRY_EPOCH="$parsed"
break
fi
done
if [[ -z "$EXPIRY_EPOCH" ]]; then
echo "WARNING: could not parse ExpirationDate '$EXPIRY'; skipping expiry check."
exit 0
fi
NOW_EPOCH=$(date +%s)
if [[ "$EXPIRY_EPOCH" -le "$NOW_EPOCH" ]]; then
echo "ERROR: provisioning profile expired on $EXPIRY. Regenerate it before rerunning."
exit 1
fi
DAYS_LEFT=$(( (EXPIRY_EPOCH - NOW_EPOCH) / 86400 ))
echo "Provisioning profile valid until $EXPIRY ($DAYS_LEFT days remaining)."
if [[ "$DAYS_LEFT" -lt 14 ]]; then
echo "WARNING: profile expires in under 14 days — regenerate soon."
fi
# ============================================================
# Build the bundle
# ============================================================
- name: Add pinned macmon to PATH
run: |
MACMON_DIR=$(nix develop --command sh -c 'dirname $(which macmon)')
echo "Using macmon from: $MACMON_DIR"
echo "$MACMON_DIR" >> $GITHUB_PATH
# Remove any Homebrew macmon so PyInstaller can't accidentally pick it up
brew uninstall macmon 2>/dev/null || true
- name: Build PyInstaller bundle
run: uv run pyinstaller packaging/pyinstaller/exo.spec
@@ -346,6 +265,7 @@ jobs:
EXO_BUILD_COMMIT="$GITHUB_SHA" \
SPARKLE_FEED_URL="$SPARKLE_FEED_URL" \
SPARKLE_ED25519_PUBLIC="$SPARKLE_ED25519_PUBLIC" \
EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT="$EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT" \
CODE_SIGNING_IDENTITY="$SIGNING_IDENTITY" \
CODE_SIGN_INJECT_BASE_ENTITLEMENTS=YES
mkdir -p ../../output
@@ -378,41 +298,11 @@ jobs:
APPLE_NOTARIZATION_PASSWORD: ${{ secrets.APPLE_NOTARIZATION_PASSWORD }}
APPLE_NOTARIZATION_TEAM: ${{ secrets.APPLE_NOTARIZATION_TEAM }}
run: |
set -o pipefail
cd output
security unlock-keychain -p "$MACOS_CERTIFICATE_PASSWORD" "$BUILD_KEYCHAIN_PATH"
SIGNING_IDENTITY=$(security find-identity -v -p codesigning "$BUILD_KEYCHAIN_PATH" | awk -F '"' '{print $2}')
# Fail fast if notarization creds are partial. All-or-nothing.
CRED_COUNT=0
for v in "$APPLE_NOTARIZATION_USERNAME" "$APPLE_NOTARIZATION_PASSWORD" "$APPLE_NOTARIZATION_TEAM"; do
[[ -n "$v" ]] && CRED_COUNT=$((CRED_COUNT + 1))
done
if [[ "$CRED_COUNT" -ne 0 && "$CRED_COUNT" -ne 3 ]]; then
echo "ERROR: partial Apple notarization credentials set ($CRED_COUNT/3). Aborting."
exit 1
fi
/usr/bin/codesign --deep --force --timestamp --options runtime \
--sign "$SIGNING_IDENTITY" EXO.app
# Pre-flight: verify the signed app BEFORE building DMG and submitting to Apple.
# If this fails, notarization will fail too — cheap way to fail in seconds, not 15 minutes.
echo "===== codesign --verify EXO.app ====="
if ! /usr/bin/codesign --verify --deep --strict --verbose=2 EXO.app; then
echo "ERROR: EXO.app failed codesign verification. Dumping signing status of every executable:"
find EXO.app -type f \( -perm -111 -o -name "*.dylib" -o -name "*.so" -o -name "*.framework" \) -print0 |
while IFS= read -r -d '' f; do
printf -- '--- %s\n' "$f"
/usr/bin/codesign -dv --verbose=2 "$f" 2>&1 | sed 's/^/ /' || true
done
exit 1
fi
# Gatekeeper assessment. A failure here strongly predicts notarization rejection.
echo "===== spctl assessment (predicts notarization outcome) ====="
/usr/bin/spctl -a -vvv -t install EXO.app || echo "WARNING: spctl assessment failed — notarization is likely to fail too."
mkdir -p dmg-root
cp -R EXO.app dmg-root/
ln -s /Applications dmg-root/Applications
@@ -420,22 +310,12 @@ jobs:
hdiutil create -volname "EXO" -srcfolder dmg-root -ov -format UDZO "$DMG_NAME"
/usr/bin/codesign --force --timestamp --options runtime \
--sign "$SIGNING_IDENTITY" "$DMG_NAME"
echo "===== codesign --verify DMG ====="
if ! /usr/bin/codesign --verify --verbose=2 "$DMG_NAME"; then
echo "ERROR: DMG failed codesign verification."
exit 1
fi
if [[ -n "$APPLE_NOTARIZATION_USERNAME" ]]; then
echo "===== notarytool submit ====="
# `|| true` so set -e doesn't abort before we can echo output / fetch the log.
# We rely on the parsed STATUS below to decide pass/fail.
SUBMISSION_OUTPUT=$(xcrun notarytool submit "$DMG_NAME" \
--apple-id "$APPLE_NOTARIZATION_USERNAME" \
--password "$APPLE_NOTARIZATION_PASSWORD" \
--team-id "$APPLE_NOTARIZATION_TEAM" \
--wait --timeout 15m 2>&1) || true
--wait --timeout 15m 2>&1)
echo "$SUBMISSION_OUTPUT"
SUBMISSION_ID=$(echo "$SUBMISSION_OUTPUT" | awk 'tolower($1)=="id:" && $2 ~ /^[0-9a-fA-F-]+$/ {print $2; exit}')
+3
View File
@@ -91,6 +91,9 @@ jobs:
nix build .#metal-toolchain
fi
# Build mlx (depends on metal-toolchain)
nix build .#mlx
- name: Build all Nix outputs
run: |
nix flake show --json | jq -r '
-3
View File
@@ -38,6 +38,3 @@ bench/**/*.json
# tmp
tmp/models
/build/exo
/.claude/skills
/.claude
+31
View File
@@ -0,0 +1,31 @@
<?xml version="1.0" encoding="UTF-8"?>
<module type="EMPTY_MODULE" version="4">
<component name="FacetManager">
<facet type="Python" name="Python facet">
<configuration sdkName="Python 3.13 virtualenv at ~/Desktop/exo/.venv" />
</facet>
</component>
<component name="Go" enabled="true" />
<component name="NewModuleRootManager">
<content url="file://$MODULE_DIR$">
<sourceFolder url="file://$MODULE_DIR$/scripts/src" isTestSource="false" />
<sourceFolder url="file://$MODULE_DIR$/src" isTestSource="false" />
<sourceFolder url="file://$MODULE_DIR$/rust/exo_pyo3_bindings/src" isTestSource="false" />
<sourceFolder url="file://$MODULE_DIR$/rust/exo_pyo3_bindings/tests" isTestSource="true" />
<sourceFolder url="file://$MODULE_DIR$/rust/util/src" isTestSource="false" />
<sourceFolder url="file://$MODULE_DIR$/rust/networking/examples" isTestSource="false" />
<sourceFolder url="file://$MODULE_DIR$/rust/networking/src" isTestSource="false" />
<sourceFolder url="file://$MODULE_DIR$/rust/networking/tests" isTestSource="true" />
<sourceFolder url="file://$MODULE_DIR$/rust/system_custodian/src" isTestSource="false" />
<excludeFolder url="file://$MODULE_DIR$/.venv" />
<excludeFolder url="file://$MODULE_DIR$/.direnv" />
<excludeFolder url="file://$MODULE_DIR$/build" />
<excludeFolder url="file://$MODULE_DIR$/dist" />
<excludeFolder url="file://$MODULE_DIR$/.go_cache" />
<excludeFolder url="file://$MODULE_DIR$/rust/target" />
</content>
<orderEntry type="jdk" jdkName="Python 3.13 (exo)" jdkType="Python SDK" />
<orderEntry type="sourceFolder" forTests="false" />
<orderEntry type="library" name="Python 3.13 virtualenv at ~/Desktop/exo/.venv interpreter library" level="application" />
</component>
</module>
+1 -1
View File
@@ -1,6 +1,6 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="ExternalDependencies">
<plugin id="al.aoli.intellijdirenv" />
<plugin id="systems.fehn.intellijdirenv" />
</component>
</project>
+14
View File
@@ -0,0 +1,14 @@
<component name="InspectionProjectProfileManager">
<profile version="1.0">
<option name="myName" value="Project Default" />
<inspection_tool class="PyCompatibilityInspection" enabled="true" level="WARNING" enabled_by_default="true">
<option name="ourVersions">
<value>
<list size="1">
<item index="0" class="java.lang.String" itemvalue="3.14" />
</list>
</value>
</option>
</inspection_tool>
</profile>
</component>
+3
View File
@@ -4,4 +4,7 @@
<option name="sdkName" value="Python 3.13 (exo)" />
</component>
<component name="ProjectRootManager" version="2" project-jdk-name="Python 3.13 (exo)" project-jdk-type="Python SDK" />
<component name="PythonCompatibilityInspectionAdvertiser">
<option name="version" value="3" />
</component>
</project>
-5
View File
@@ -1,5 +0,0 @@
"""
This type stub file was generated by pyright.
"""
__version__ = ...
@@ -1,19 +0,0 @@
"""
This type stub file was generated by pyright.
"""
import mlx.core as mx
import mlx.nn as nn
from functools import partial
@partial(mx.compile, shapeless=True)
def swiglu(gate, x): ...
@partial(mx.compile, shapeless=True)
def xielu(x, alpha_p, alpha_n, beta, eps): # -> array:
...
class XieLU(nn.Module):
def __init__(
self, alpha_p_init=..., alpha_n_init=..., beta=..., eps=...
) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
-280
View File
@@ -1,280 +0,0 @@
"""Type stubs for mlx_lm.models.deepseek_v4"""
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
import mlx.core as mx
import mlx.nn as nn
from .base import BaseModelArgs
from .cache import ArraysCache, RotatingKVCache
from .switch_layers import SwitchGLU
@dataclass
class ModelArgs(BaseModelArgs):
model_type: str
vocab_size: int
hidden_size: int
intermediate_size: int
moe_intermediate_size: int
num_hidden_layers: int
num_attention_heads: int
num_key_value_heads: int
n_shared_experts: Optional[int]
n_routed_experts: int
num_experts_per_tok: int
head_dim: int
qk_rope_head_dim: int
q_lora_rank: int
o_lora_rank: int
o_groups: int
sliding_window: int
hc_mult: int
hc_sinkhorn_iters: int
hc_eps: float
compress_ratios: Optional[List[int]]
compress_rope_theta: float
rope_theta: float
rope_scaling: Optional[Dict[str, Any]]
rms_norm_eps: float
swiglu_limit: float
attention_bias: bool
max_position_embeddings: int
class DeepseekV4RoPE(nn.Module):
dims: int
freqs: mx.array
def __init__(
self,
dims: int,
base: float,
scaling_config: Optional[Dict[str, Any]] = None,
) -> None: ...
def __call__(
self,
x: mx.array,
offset: int = 0,
inverse: bool = False,
) -> mx.array: ...
class HyperConnection(nn.Module):
dim: int
hc_mult: int
norm_eps: float
def __init__(
self,
dim: int,
hc_mult: int,
norm_eps: float,
sinkhorn_iters: int,
hc_eps: float,
) -> None: ...
class HyperHead(nn.Module):
dim: int
hc_mult: int
def __init__(
self,
dim: int,
hc_mult: int,
norm_eps: float,
hc_eps: float,
) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
class Compressor(nn.Module):
dim: int
head_dim: int
rope_head_dim: int
compress_ratio: int
overlap: bool
wkv_gate: nn.Linear
ape: mx.array
norm: nn.RMSNorm
rope: DeepseekV4RoPE
def __init__(
self,
dim: int,
compress_ratio: int,
head_dim: int,
rope_head_dim: int,
rms_norm_eps: float,
rope: DeepseekV4RoPE,
) -> None: ...
def __call__(
self,
x: mx.array,
cache: "DeepseekV4Cache",
offset: Any,
key: str = ...,
) -> mx.array: ...
class Indexer(nn.Module):
def __init__(
self,
args: ModelArgs,
compress_ratio: int,
rope: DeepseekV4RoPE,
) -> None: ...
class _CompressorBranch:
buffer_kv: Optional[mx.array]
buffer_gate: Optional[mx.array]
prev_kv: Optional[mx.array]
prev_gate: Optional[mx.array]
pool: Optional[mx.array]
buffer_lengths: Optional[List[int]]
pool_lengths: Optional[List[int]]
buffer_count: int
_new_pool_lengths: Optional[List[int]]
def __init__(self) -> None: ...
class DeepseekV4Cache:
local: RotatingKVCache
offset: int
keys: Optional[mx.array]
values: Optional[mx.array]
state: Any
meta_state: Any
nbytes: int
_branches: Dict[str, _CompressorBranch]
_pending_lengths: Optional[List[int]]
def __init__(self, sliding_window: int) -> None: ...
def update_and_fetch(
self, keys: mx.array, values: mx.array
) -> tuple[mx.array, mx.array]: ...
def is_trimmable(self) -> bool: ...
def trim(self, n: int) -> int: ...
def empty(self) -> bool: ...
def size(self) -> int: ...
def prepare(
self,
*,
left_padding: Optional[List[int]] = None,
lengths: Optional[List[int]] = None,
right_padding: Optional[List[int]] = None,
) -> None: ...
def finalize(self) -> None: ...
def filter(self, batch_indices: mx.array) -> None: ...
def extend(self, other: "DeepseekV4Cache") -> None: ...
def extract(self, idx: int) -> "DeepseekV4Cache": ...
@classmethod
def merge(cls, caches: List["DeepseekV4Cache"]) -> "DeepseekV4Cache": ...
class V4Attention(nn.Module):
args: ModelArgs
layer_id: int
dim: int
n_heads: int
head_dim: int
rope_head_dim: int
nope_head_dim: int
n_groups: int
q_lora_rank: int
o_lora_rank: int
window: int
eps: float
scale: float
compress_ratio: int
wqkv_a: nn.Linear
q_norm: nn.RMSNorm
wq_b: nn.Linear
kv_norm: nn.RMSNorm
attn_sink: mx.array
wo_a: nn.Linear
wo_b: nn.Linear
rope: DeepseekV4RoPE
compressor: Compressor
indexer: Indexer
def __init__(self, args: ModelArgs, layer_id: int) -> None: ...
def __call__(
self,
x: mx.array,
mask: Optional[mx.array] = None,
cache: Optional[Any] = None,
) -> mx.array: ...
class DeepseekV4MLP(nn.Module):
gate_proj: nn.Linear
up_proj: nn.Linear
down_proj: nn.Linear
def __init__(
self,
hidden_size: int,
intermediate_size: int,
swiglu_limit: float = 0.0,
) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
class MoEGate(nn.Module):
weight: mx.array
def __init__(self, args: ModelArgs, layer_id: int) -> None: ...
def __call__(
self, x: mx.array, input_ids: mx.array
) -> tuple[mx.array, mx.array]: ...
class DeepseekV4MoE(nn.Module):
num_experts_per_tok: int
switch_mlp: SwitchGLU
gate: MoEGate
shared_experts: DeepseekV4MLP
def __init__(self, args: ModelArgs, layer_id: int) -> None: ...
def __call__(self, x: mx.array, input_ids: mx.array) -> mx.array: ...
class DeepseekV4Block(nn.Module):
attn_norm: nn.RMSNorm
attn: V4Attention
hc_attn: HyperConnection
ffn_norm: nn.RMSNorm
ffn: DeepseekV4MoE
hc_ffn: HyperConnection
def __init__(self, args: ModelArgs, layer_id: int) -> None: ...
def __call__(
self,
h: mx.array,
cache: Optional[Any],
input_ids: mx.array,
) -> mx.array: ...
class DeepseekV4Model(nn.Module):
args: ModelArgs
vocab_size: int
embed_tokens: nn.Embedding
layers: list[DeepseekV4Block]
norm: nn.RMSNorm
hc_head: HyperHead
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
inputs: mx.array,
cache: Optional[List[Any]] = None,
) -> mx.array: ...
class Model(nn.Module):
args: ModelArgs
model_type: str
model: DeepseekV4Model
lm_head: nn.Linear
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
inputs: mx.array,
cache: Optional[List[Any]] = None,
) -> mx.array: ...
def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
def make_cache(self) -> list[RotatingKVCache | DeepseekV4Cache]: ...
@property
def layers(self) -> list[DeepseekV4Block]: ...
@@ -1,35 +0,0 @@
from typing import Optional
import mlx.core as mx
def compute_g(A_log: mx.array, a: mx.array, dt_bias: mx.array) -> mx.array: ...
def gated_delta_update(
q: mx.array,
k: mx.array,
v: mx.array,
a: mx.array,
b: mx.array,
A_log: mx.array,
dt_bias: mx.array,
state: Optional[mx.array] = ...,
mask: Optional[mx.array] = ...,
use_kernel: bool = ...,
) -> tuple[mx.array, mx.array]: ...
def gated_delta_ops(
q: mx.array,
k: mx.array,
v: mx.array,
g: mx.array,
beta: mx.array,
state: Optional[mx.array] = ...,
mask: Optional[mx.array] = ...,
) -> tuple[mx.array, mx.array]: ...
def gated_delta_kernel(
q: mx.array,
k: mx.array,
v: mx.array,
g: mx.array,
beta: mx.array,
state: mx.array,
mask: Optional[mx.array] = ...,
) -> tuple[mx.array, mx.array]: ...
-31
View File
@@ -1,31 +0,0 @@
from dataclasses import dataclass
from typing import Any, Optional
import mlx.core as mx
import mlx.nn as nn
from . import gemma4_text
from .base import BaseModelArgs
from .cache import KVCache, RotatingKVCache
@dataclass
class ModelArgs(BaseModelArgs):
model_type: str
text_config: Optional[dict[str, Any]]
vocab_size: int
def __post_init__(self) -> None: ...
class Model(nn.Module):
args: ModelArgs
model_type: str
language_model: gemma4_text.Model
def __init__(self, args: ModelArgs) -> None: ...
def __call__(self, *args: Any, **kwargs: Any) -> mx.array: ...
def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
@property
def layers(self) -> list[gemma4_text.DecoderLayer]: ...
@property
def quant_predicate(self) -> Any: ...
def make_cache(self) -> list[KVCache | RotatingKVCache]: ...
-179
View File
@@ -1,179 +0,0 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
import mlx.core as mx
import mlx.nn as nn
from .base import BaseModelArgs
from .cache import KVCache, RotatingKVCache
from .switch_layers import SwitchGLU
@dataclass
class ModelArgs(BaseModelArgs):
model_type: str
hidden_size: int
num_hidden_layers: int
intermediate_size: int
num_attention_heads: int
head_dim: int
global_head_dim: int
global_partial_rotary_factor: float
rms_norm_eps: float
vocab_size: int
vocab_size_per_layer_input: int
num_key_value_heads: int
num_global_key_value_heads: Optional[int]
num_kv_shared_layers: int
pad_token_id: int
hidden_size_per_layer_input: int
rope_traditional: bool
partial_rotary_factor: float
rope_parameters: Optional[Dict[str, Any]]
sliding_window: int
sliding_window_pattern: int
max_position_embeddings: int
attention_k_eq_v: bool
final_logit_softcapping: float
use_double_wide_mlp: bool
enable_moe_block: bool
num_experts: Optional[int]
top_k_experts: Optional[int]
moe_intermediate_size: Optional[int]
layer_types: Optional[List[str]]
tie_word_embeddings: bool
def __post_init__(self) -> None: ...
class MLP(nn.Module):
gate_proj: nn.Linear
down_proj: nn.Linear
up_proj: nn.Linear
def __init__(self, config: ModelArgs, layer_idx: int = 0) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
class Router(nn.Module):
proj: nn.Linear
scale: mx.array
per_expert_scale: mx.array
def __init__(self, config: ModelArgs) -> None: ...
def __call__(self, x: mx.array) -> tuple[mx.array, mx.array]: ...
class Experts(nn.Module):
switch_glu: SwitchGLU
def __init__(self, config: ModelArgs) -> None: ...
def __call__(
self, x: mx.array, top_k_indices: mx.array, top_k_weights: mx.array
) -> mx.array: ...
class Attention(nn.Module):
layer_idx: int
layer_type: str
is_sliding: bool
head_dim: int
n_heads: int
n_kv_heads: int
use_k_eq_v: bool
scale: float
q_proj: nn.Linear
k_proj: nn.Linear
v_proj: nn.Linear
o_proj: nn.Linear
q_norm: nn.Module
k_norm: nn.Module
v_norm: nn.Module
rope: nn.Module
def __init__(self, config: ModelArgs, layer_idx: int) -> None: ...
def __call__(self, *args: Any, **kwargs: Any) -> Any: ...
class DecoderLayer(nn.Module):
layer_idx: int
layer_type: str
self_attn: Attention
mlp: MLP
enable_moe: bool
router: Router
experts: Experts
input_layernorm: nn.Module
post_attention_layernorm: nn.Module
pre_feedforward_layernorm: nn.Module
post_feedforward_layernorm: nn.Module
post_feedforward_layernorm_1: nn.Module
post_feedforward_layernorm_2: nn.Module
pre_feedforward_layernorm_2: nn.Module
hidden_size_per_layer_input: int
per_layer_input_gate: Optional[nn.Linear]
per_layer_projection: Optional[nn.Linear]
post_per_layer_input_norm: Optional[nn.Module]
layer_scalar: mx.array
def __init__(self, config: ModelArgs, layer_idx: int) -> None: ...
def __call__(
self,
x: mx.array,
mask: Optional[mx.array] = ...,
cache: Optional[Any] = ...,
per_layer_input: Optional[mx.array] = ...,
shared_kv: Optional[tuple[mx.array, mx.array]] = ...,
offset: Optional[mx.array] = ...,
) -> tuple[mx.array, tuple[mx.array, mx.array], mx.array]: ...
class Gemma4TextModel(nn.Module):
config: ModelArgs
vocab_size: int
window_size: int
sliding_window_pattern: int
num_hidden_layers: int
embed_tokens: nn.Embedding
embed_scale: float
layers: list[DecoderLayer]
norm: nn.Module
hidden_size_per_layer_input: int
embed_tokens_per_layer: Optional[nn.Embedding]
per_layer_model_projection: Optional[nn.Linear]
per_layer_projection_norm: Optional[nn.Module]
previous_kvs: list[int]
def __init__(self, config: ModelArgs) -> None: ...
def __call__(
self,
inputs: Optional[mx.array] = ...,
cache: Optional[list[Any]] = ...,
input_embeddings: Optional[mx.array] = ...,
per_layer_inputs: Optional[mx.array] = ...,
) -> mx.array: ...
def _get_per_layer_inputs(
self,
input_ids: Optional[mx.array],
input_embeddings: Optional[mx.array] = ...,
) -> mx.array: ...
def _project_per_layer_inputs(
self,
input_embeddings: mx.array,
per_layer_inputs: Optional[mx.array] = ...,
) -> mx.array: ...
def _make_masks(self, h: mx.array, cache: list[Any]) -> list[Any]: ...
class Model(nn.Module):
args: ModelArgs
model_type: str
model: Gemma4TextModel
final_logit_softcapping: float
tie_word_embeddings: bool
lm_head: nn.Linear
def __init__(self, args: ModelArgs) -> None: ...
def __call__(self, *args: Any, **kwargs: Any) -> mx.array: ...
def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
@property
def layers(self) -> list[DecoderLayer]: ...
@property
def head_dim(self) -> int: ...
@property
def n_kv_heads(self) -> int: ...
@property
def quant_predicate(self) -> Any: ...
def make_cache(self) -> list[KVCache | RotatingKVCache]: ...
-103
View File
@@ -1,103 +0,0 @@
"""Type stubs for mlx_lm.models.gpt_oss"""
from dataclasses import dataclass
from typing import Any, List, Optional
import mlx.core as mx
import mlx.nn as nn
from .base import BaseModelArgs
from .cache import KVCache
from .switch_layers import SwitchGLU
@dataclass
class ModelArgs(BaseModelArgs):
model_type: str
hidden_size: int
intermediate_size: int
num_hidden_layers: int
num_attention_heads: int
num_key_value_heads: int
num_local_experts: int
num_experts_per_tok: int
vocab_size: int
rms_norm_eps: float
sliding_window: int
layer_types: Optional[List[str]]
def mlx_topk(a: mx.array, k: int, axis: int = -1) -> tuple[mx.array, mx.array]: ...
class AttentionBlock(nn.Module):
head_dim: int
num_attention_heads: int
num_key_value_heads: int
num_key_value_groups: int
sinks: mx.array
q_proj: nn.Linear
k_proj: nn.Linear
v_proj: nn.Linear
o_proj: nn.Linear
sm_scale: float
rope: nn.Module
def __init__(self, config: ModelArgs) -> None: ...
def __call__(
self,
x: mx.array,
mask: Optional[mx.array] = None,
cache: Optional[Any] = None,
) -> mx.array: ...
class TransformerBlock(nn.Module):
self_attn: AttentionBlock
mlp: MLPBlock
def __init__(self, config: ModelArgs) -> None: ...
def __call__(
self,
x: mx.array,
mask: Optional[mx.array] = None,
cache: Optional[Any] = None,
) -> mx.array: ...
class MLPBlock(nn.Module):
hidden_size: int
num_local_experts: int
num_experts_per_tok: int
experts: SwitchGLU
router: nn.Linear
sharding_group: Optional[mx.distributed.Group]
def __init__(self, config: ModelArgs) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
class GptOssMoeModel(nn.Module):
embed_tokens: nn.Embedding
norm: nn.RMSNorm
layer_types: List[str]
layers: list[TransformerBlock]
window_size: int
swa_idx: int
ga_idx: int
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
inputs: mx.array,
cache: Optional[Any] = None,
) -> mx.array: ...
class Model(nn.Module):
model_type: str
model: GptOssMoeModel
lm_head: nn.Linear
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
inputs: mx.array,
cache: Optional[Any] = None,
) -> mx.array: ...
@property
def layers(self) -> list[nn.Module]: ...
def make_cache(self) -> list[KVCache]: ...
-94
View File
@@ -1,94 +0,0 @@
"""Type stubs for mlx_lm.models.minimax"""
from dataclasses import dataclass
from typing import Any, Optional
import mlx.core as mx
import mlx.nn as nn
from .base import BaseModelArgs
from .switch_layers import SwitchGLU
@dataclass
class ModelArgs(BaseModelArgs):
model_type: str
hidden_size: int
intermediate_size: int
num_hidden_layers: int
num_attention_heads: int
num_key_value_heads: int
num_local_experts: int
num_experts_per_tok: int
max_position_embeddings: int
class MiniMaxAttention(nn.Module):
num_heads: int
num_attention_heads: int
num_key_value_heads: int
head_dim: int
scale: float
q_proj: nn.Linear
k_proj: nn.Linear
v_proj: nn.Linear
o_proj: nn.Linear
q_norm: nn.Module
k_norm: nn.Module
rope: nn.Module
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
x: mx.array,
mask: Optional[mx.array] = None,
cache: Optional[Any] = None,
) -> mx.array: ...
class MiniMaxSparseMoeBlock(nn.Module):
num_experts_per_tok: int
gate: nn.Linear
switch_mlp: SwitchGLU
e_score_correction_bias: mx.array
sharding_group: Optional[mx.distributed.Group]
def __init__(self, args: ModelArgs) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
class MiniMaxDecoderLayer(nn.Module):
self_attn: MiniMaxAttention
block_sparse_moe: MiniMaxSparseMoeBlock
input_layernorm: nn.RMSNorm
post_attention_layernorm: nn.RMSNorm
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
x: mx.array,
mask: Optional[mx.array] = None,
cache: Optional[Any] = None,
) -> mx.array: ...
class MiniMaxModel(nn.Module):
embed_tokens: nn.Embedding
layers: list[MiniMaxDecoderLayer]
norm: nn.RMSNorm
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
inputs: mx.array,
cache: Optional[Any] = None,
) -> mx.array: ...
class Model(nn.Module):
model_type: str
model: MiniMaxModel
lm_head: nn.Linear
def __init__(self, args: ModelArgs) -> None: ...
def __call__(
self,
inputs: mx.array,
cache: Optional[Any] = None,
) -> mx.array: ...
@property
def layers(self) -> list[MiniMaxDecoderLayer]: ...
-51
View File
@@ -1,51 +0,0 @@
from typing import Any, Optional
import mlx.nn as nn
class YarnRoPE(nn.Module):
def __init__(
self,
dims: int,
traditional: bool = ...,
max_position_embeddings: int = ...,
base: float = ...,
scaling_factor: float = ...,
original_max_position_embeddings: int = ...,
beta_fast: float = ...,
beta_slow: float = ...,
mscale: float = ...,
mscale_all_dim: float = ...,
) -> None: ...
class Llama3RoPE(nn.Module):
def __init__(
self,
dims: int,
traditional: bool = ...,
max_position_embeddings: int = ...,
base: float = ...,
scaling_factor: float = ...,
original_max_position_embeddings: int = ...,
low_freq_factor: float = ...,
high_freq_factor: float = ...,
) -> None: ...
class SuScaledRoPE(nn.Module):
def __init__(
self,
dims: int,
traditional: bool = ...,
max_position_embeddings: int = ...,
base: float = ...,
short_factor: Any = ...,
long_factor: Any = ...,
original_max_position_embeddings: int = ...,
) -> None: ...
def initialize_rope(
dims: int,
base: float = ...,
traditional: bool = ...,
scaling_config: Optional[dict[str, Any]] = ...,
max_position_embeddings: Optional[int] = ...,
) -> nn.Module: ...
-50
View File
@@ -1,50 +0,0 @@
"""
This type stub file was generated by pyright.
"""
import mlx.nn as nn
class DoRALinear(nn.Module):
@staticmethod
def from_base(
linear: nn.Linear, r: int = ..., dropout: float = ..., scale: float = ...
): # -> DoRALinear:
...
def fuse(self, dequantize: bool = ...): # -> QuantizedLinear | Linear:
...
def __init__(
self,
input_dims: int,
output_dims: int,
r: int = ...,
dropout: float = ...,
scale: float = ...,
bias: bool = ...,
) -> None: ...
def set_linear(self, linear): # -> None:
"""
Set the self.linear layer and recompute self.m.
"""
...
def __call__(self, x): ...
class DoRAEmbedding(nn.Module):
def from_base(
embedding: nn.Embedding, r: int = ..., dropout: float = ..., scale: float = ...
): # -> DoRAEmbedding:
...
def fuse(self, dequantize: bool = ...): # -> Embedding:
...
def __init__(
self,
num_embeddings: int,
dims: int,
r: int = ...,
dropout: float = ...,
scale: float = ...,
) -> None: ...
def set_embedding(self, embedding: nn.Module): # -> None:
...
def __call__(self, x): ...
def as_linear(self, x): ...
-66
View File
@@ -1,66 +0,0 @@
"""
This type stub file was generated by pyright.
"""
import mlx.nn as nn
class LoRALinear(nn.Module):
@staticmethod
def from_base(
linear: nn.Linear, r: int = ..., dropout: float = ..., scale: float = ...
): # -> LoRALinear:
...
def fuse(self, dequantize: bool = ...): # -> QuantizedLinear | Linear:
...
def __init__(
self,
input_dims: int,
output_dims: int,
r: int = ...,
dropout: float = ...,
scale: float = ...,
bias: bool = ...,
) -> None: ...
def __call__(self, x): # -> array:
...
class LoRASwitchLinear(nn.Module):
@staticmethod
def from_base(
linear: nn.Module, r: int = ..., dropout: float = ..., scale: float = ...
): # -> LoRASwitchLinear:
...
def fuse(self, dequantize: bool = ...): # -> QuantizedSwitchLinear | SwitchLinear:
...
def __init__(
self,
input_dims: int,
output_dims: int,
num_experts: int,
r: int = ...,
dropout: float = ...,
scale: float = ...,
bias: bool = ...,
) -> None: ...
def __call__(self, x, indices, sorted_indices=...): ...
class LoRAEmbedding(nn.Module):
@staticmethod
def from_base(
embedding: nn.Embedding, r: int = ..., dropout: float = ..., scale: float = ...
): # -> LoRAEmbedding:
...
def fuse(self, dequantize: bool = ...): # -> QuantizedEmbedding | Embedding:
...
def __init__(
self,
num_embeddings: int,
dims: int,
r: int = ...,
dropout: float = ...,
scale: float = ...,
) -> None: ...
def __call__(self, x): # -> array:
...
def as_linear(self, x): # -> array:
...
-57
View File
@@ -1,57 +0,0 @@
"""
This type stub file was generated by pyright.
"""
import mlx.nn as nn
from typing import Dict
def build_schedule(schedule_config: Dict): # -> Any:
"""
Build a learning rate schedule from the given config.
"""
...
def linear_to_lora_layers(
model: nn.Module, num_layers: int, config: Dict, use_dora: bool = ...
): # -> None:
"""
Convert some of the models linear layers to lora layers.
Args:
model (nn.Module): The neural network model.
num_layers (int): The number of blocks to convert to lora layers
starting from the last layer.
config (dict): More configuration parameters for LoRA, including the
rank, scale, and optional layer keys.
use_dora (bool): If True, uses DoRA instead of LoRA.
Default: ``False``
"""
...
def load_adapters(model: nn.Module, adapter_path: str) -> nn.Module:
"""
Load any fine-tuned adapters / layers.
Args:
model (nn.Module): The neural network model.
adapter_path (str): Path to the adapter configuration file.
Returns:
nn.Module: The updated model with LoRA layers applied.
"""
...
def remove_lora_layers(model: nn.Module) -> nn.Module:
"""
Remove the LoRA layers from the model.
Args:
model (nn.Module): The model with LoRA layers.
Returns:
nn.Module: The model without LoRA layers.
"""
...
def print_trainable_parameters(model): # -> None:
...
-12
View File
@@ -1,12 +0,0 @@
from typing import Any
def get_message_json(
model_name: str,
prompt: str,
role: str = "user",
skip_image_token: bool = False,
skip_audio_token: bool = False,
num_images: int = 0,
num_audios: int = 0,
**kwargs: Any,
) -> dict[str, Any]: ...
-15
View File
@@ -1,15 +0,0 @@
from pathlib import Path
from typing import Any
class ImageProcessor:
def preprocess(
self, images: list[dict[str, Any]], **kwargs: Any
) -> dict[str, Any]: ...
def __call__(self, **kwargs: Any) -> dict[str, Any]: ...
def load_image_processor(
model_path: str | Path, **kwargs: Any
) -> ImageProcessor | None: ...
def load_processor(
model_path: str | Path, add_detokenizer: bool = ..., **kwargs: Any
) -> ImageProcessor: ...
-8
View File
@@ -1,8 +0,0 @@
from typing import Any, Self
class safe_open:
def __init__(self, filename: str, framework: str = "pt") -> None: ...
def __enter__(self) -> Self: ...
def __exit__(self, *args: Any) -> None: ...
def keys(self) -> list[str]: ...
def get_tensor(self, name: str) -> Any: ...
File renamed without changes.
File renamed without changes.
@@ -2,10 +2,11 @@
This type stub file was generated by pyright.
"""
from typing import Protocol
import mlx.core as mx
import PIL.Image
import tqdm
from typing import Protocol
from mflux.models.common.config.config import Config
class BeforeLoopCallback(Protocol):
@@ -3,6 +3,7 @@ This type stub file was generated by pyright.
"""
from typing import TYPE_CHECKING
from mflux.callbacks.callback import (
AfterLoopCallback,
BeforeLoopCallback,
@@ -2,10 +2,11 @@
This type stub file was generated by pyright.
"""
from typing import TYPE_CHECKING
import mlx.core as mx
import PIL.Image
import tqdm
from typing import TYPE_CHECKING
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.models.common.config.config import Config
File renamed without changes.
File renamed without changes.
@@ -2,11 +2,12 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from pathlib import Path
from typing import Any
from tqdm import tqdm
import mlx.core as mx
from mflux.models.common.config.model_config import ModelConfig
from tqdm import tqdm
logger = ...
@@ -2,10 +2,11 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from functools import lru_cache
from typing import Literal
import mlx.core as mx
class ModelConfig:
precision: mx.Dtype = ...
def __init__(
@@ -2,10 +2,10 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from pathlib import Path
from typing import TYPE_CHECKING, TypeAlias
from mlx import nn
import mlx.core as mx
from mflux.models.common.vae.tiling_config import TilingConfig
from mflux.models.fibo.latent_creator.fibo_latent_creator import FiboLatentCreator
from mflux.models.flux.latent_creator.flux_latent_creator import FluxLatentCreator
@@ -13,6 +13,7 @@ from mflux.models.qwen.latent_creator.qwen_latent_creator import QwenLatentCreat
from mflux.models.z_image.latent_creator.z_image_latent_creator import (
ZImageLatentCreator,
)
from mlx import nn
if TYPE_CHECKING:
LatentCreatorType: TypeAlias = type[
@@ -2,8 +2,8 @@
This type stub file was generated by pyright.
"""
from mlx import nn
from mflux.models.common.lora.layer.linear_lora_layer import LoRALinear
from mlx import nn
class FusedLoRALinear(nn.Module):
def __init__(
@@ -2,10 +2,11 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
import mlx.nn as nn
from collections.abc import Callable
from dataclasses import dataclass
import mlx.core as mx
import mlx.nn as nn
from mflux.models.common.lora.mapping.lora_mapping import LoRATarget
@dataclass
@@ -2,11 +2,12 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from collections.abc import Callable
from dataclasses import dataclass
from typing import List, Protocol
import mlx.core as mx
@dataclass
class LoRATarget:
model_path: str
@@ -36,4 +36,3 @@ class Rule(NamedTuple):
name: str
check: str
action: QuantizationAction | PathAction | LoraAction | ConfigAction
...
@@ -3,6 +3,7 @@ This type stub file was generated by pyright.
"""
from typing import TYPE_CHECKING
from mflux.models.common.config.model_config import ModelConfig
if TYPE_CHECKING: ...
@@ -2,9 +2,10 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from abc import ABC, abstractmethod
import mlx.core as mx
class BaseScheduler(ABC):
@property
@abstractmethod
@@ -2,8 +2,9 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from typing import TYPE_CHECKING
import mlx.core as mx
from mflux.models.common.config.config import Config
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
@@ -2,8 +2,9 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from typing import TYPE_CHECKING
import mlx.core as mx
from mflux.models.common.config.config import Config
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
@@ -2,8 +2,9 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from typing import TYPE_CHECKING
import mlx.core as mx
from mflux.models.common.config.config import Config
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
@@ -4,9 +4,10 @@ This type stub file was generated by pyright.
from abc import ABC, abstractmethod
from typing import Protocol, runtime_checkable
from mflux.models.common.tokenizer.tokenizer_output import TokenizerOutput
from PIL import Image
from transformers import PreTrainedTokenizer
from mflux.models.common.tokenizer.tokenizer_output import TokenizerOutput
"""
This type stub file was generated by pyright.
@@ -3,6 +3,7 @@ This type stub file was generated by pyright.
"""
from typing import TYPE_CHECKING
from mflux.models.common.tokenizer.tokenizer import BaseTokenizer
from mflux.models.common.weights.loading.weight_definition import TokenizerDefinition
@@ -2,9 +2,10 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from dataclasses import dataclass
import mlx.core as mx
"""
This type stub file was generated by pyright.
"""
@@ -2,9 +2,10 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from typing import Callable
import mlx.core as mx
class VAETiler:
@staticmethod
def encode_image_tiled(
@@ -3,8 +3,8 @@ This type stub file was generated by pyright.
"""
import mlx.core as mx
from mlx import nn
from mflux.models.common.vae.tiling_config import TilingConfig
from mlx import nn
class VAEUtil:
@staticmethod
@@ -2,8 +2,9 @@
This type stub file was generated by pyright.
"""
import mlx.nn as nn
from typing import TYPE_CHECKING
import mlx.nn as nn
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights
from mflux.models.common.weights.loading.weight_definition import (
ComponentDefinition,
@@ -2,11 +2,12 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from dataclasses import dataclass
from typing import Callable, List, TYPE_CHECKING, TypeAlias
from mflux.models.common.weights.mapping.weight_mapping import WeightTarget
from typing import TYPE_CHECKING, Callable, List, TypeAlias
import mlx.core as mx
from mflux.models.common.tokenizer.tokenizer import BaseTokenizer
from mflux.models.common.weights.mapping.weight_mapping import WeightTarget
from mflux.models.depth_pro.weights.depth_pro_weight_definition import (
DepthProWeightDefinition,
)
@@ -3,6 +3,7 @@ This type stub file was generated by pyright.
"""
from typing import TYPE_CHECKING
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights
from mflux.models.common.weights.loading.weight_definition import (
ComponentDefinition,
@@ -2,8 +2,9 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from typing import Dict, List, Optional
import mlx.core as mx
from mflux.models.common.weights.mapping.weight_mapping import WeightTarget
class WeightMapper:
@@ -2,10 +2,11 @@
This type stub file was generated by pyright.
"""
import mlx.core as mx
from dataclasses import dataclass
from typing import Callable, List, Optional, Protocol
import mlx.core as mx
"""
This type stub file was generated by pyright.
"""
Loaded 100 of 750 files, more files were not shown because too many files have changed in this diff. Show more