mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-08 11:35:40 -04:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
67deec88ca | ||
|
|
209d618d5a |
No files matched your search
@@ -0,0 +1,12 @@
|
||||
name: Type Check
|
||||
|
||||
description: "Run type checker"
|
||||
|
||||
runs:
|
||||
using: "composite"
|
||||
steps:
|
||||
- name: Run type checker
|
||||
run: |
|
||||
nix --extra-experimental-features nix-command --extra-experimental-features flakes develop -c just sync
|
||||
nix --extra-experimental-features nix-command --extra-experimental-features flakes develop -c just check
|
||||
shell: bash
|
||||
@@ -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}')
|
||||
@@ -516,7 +396,7 @@ jobs:
|
||||
path: output/EXO-${{ env.RELEASE_VERSION }}.dmg
|
||||
|
||||
- name: Upload to S3
|
||||
if: env.SPARKLE_S3_BUCKET != ''
|
||||
if: env.SPARKLE_S3_BUCKET != '' && github.ref_type == 'tag'
|
||||
env:
|
||||
AWS_ACCESS_KEY_ID: ${{ secrets.AWS_ACCESS_KEY_ID }}
|
||||
AWS_SECRET_ACCESS_KEY: ${{ secrets.AWS_SECRET_ACCESS_KEY }}
|
||||
@@ -532,12 +412,6 @@ jobs:
|
||||
PREFIX="${PREFIX}/"
|
||||
fi
|
||||
DMG_NAME="EXO-${RELEASE_VERSION}.dmg"
|
||||
|
||||
if [[ "${{ github.ref_type }}" != "tag" ]]; then
|
||||
aws s3 cp "$DMG_NAME" "s3://${SPARKLE_S3_BUCKET}/${PREFIX}EXO-${GITHUB_SHA}.dmg"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
aws s3 cp "$DMG_NAME" "s3://${SPARKLE_S3_BUCKET}/${PREFIX}${DMG_NAME}"
|
||||
if [[ "$IS_ALPHA" != "true" ]]; then
|
||||
aws s3 cp "$DMG_NAME" "s3://${SPARKLE_S3_BUCKET}/${PREFIX}EXO-latest.dmg"
|
||||
|
||||
@@ -8,6 +8,92 @@ on:
|
||||
- main
|
||||
|
||||
jobs:
|
||||
typecheck:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
lfs: false
|
||||
|
||||
- uses: cachix/install-nix-action@v31
|
||||
with:
|
||||
nix_path: nixpkgs=channel:nixos-unstable
|
||||
|
||||
- uses: cachix/cachix-action@v14
|
||||
name: Configure Cachix
|
||||
with:
|
||||
name: exo
|
||||
authToken: "${{ secrets.CACHIX_AUTH_TOKEN }}"
|
||||
|
||||
- name: Configure git user
|
||||
run: |
|
||||
git config --local user.email "github-actions@users.noreply.github.com"
|
||||
git config --local user.name "github-actions bot"
|
||||
shell: bash
|
||||
|
||||
- name: Pull LFS files
|
||||
run: |
|
||||
echo "Pulling Git LFS files..."
|
||||
git lfs pull
|
||||
shell: bash
|
||||
|
||||
- name: Setup Nix Environment
|
||||
run: |
|
||||
echo "Checking for nix installation..."
|
||||
|
||||
# Check if nix binary exists directly
|
||||
if [ -f /nix/var/nix/profiles/default/bin/nix ]; then
|
||||
echo "Found nix binary at /nix/var/nix/profiles/default/bin/nix"
|
||||
export PATH="/nix/var/nix/profiles/default/bin:$PATH"
|
||||
echo "PATH=$PATH" >> $GITHUB_ENV
|
||||
nix --version
|
||||
elif [ -f /nix/var/nix/profiles/default/etc/profile.d/nix-daemon.sh ]; then
|
||||
echo "Found nix profile script, sourcing..."
|
||||
source /nix/var/nix/profiles/default/etc/profile.d/nix-daemon.sh
|
||||
nix --version
|
||||
elif command -v nix >/dev/null 2>&1; then
|
||||
echo "Nix already in PATH"
|
||||
nix --version
|
||||
else
|
||||
echo "Nix not found. Debugging info:"
|
||||
echo "Contents of /nix/var/nix/profiles/default/:"
|
||||
ls -la /nix/var/nix/profiles/default/ 2>/dev/null || echo "Directory not found"
|
||||
echo "Contents of /nix/var/nix/profiles/default/bin/:"
|
||||
ls -la /nix/var/nix/profiles/default/bin/ 2>/dev/null || echo "Directory not found"
|
||||
exit 1
|
||||
fi
|
||||
shell: bash
|
||||
|
||||
- name: Configure basedpyright include for local MLX
|
||||
run: |
|
||||
RUNNER_LABELS='${{ toJSON(runner.labels) }}'
|
||||
if echo "$RUNNER_LABELS" | grep -q "local_mlx"; then
|
||||
if [ -d "/Users/Shared/mlx" ]; then
|
||||
echo "Updating [tool.basedpyright].include to use /Users/Shared/mlx"
|
||||
awk '
|
||||
BEGIN { in=0 }
|
||||
/^\[tool\.basedpyright\]/ { in=1; print; next }
|
||||
in && /^\[/ { in=0 } # next section
|
||||
in && /^[ \t]*include[ \t]*=/ {
|
||||
print "include = [\"/Users/Shared/mlx\"]"
|
||||
next
|
||||
}
|
||||
{ print }
|
||||
' pyproject.toml > pyproject.toml.tmp && mv pyproject.toml.tmp pyproject.toml
|
||||
|
||||
echo "New [tool.basedpyright] section:"
|
||||
sed -n '/^\[tool\.basedpyright\]/,/^\[/p' pyproject.toml | sed '$d' || true
|
||||
else
|
||||
echo "local_mlx tag present but /Users/Shared/mlx not found; leaving pyproject unchanged."
|
||||
fi
|
||||
else
|
||||
echo "Runner does not have 'local_mlx' tag; leaving pyproject unchanged."
|
||||
fi
|
||||
shell: bash
|
||||
|
||||
- uses: ./.github/actions/typecheck
|
||||
|
||||
nix:
|
||||
name: Build and check (${{ matrix.system }})
|
||||
runs-on: ${{ matrix.runner }}
|
||||
@@ -37,60 +123,6 @@ jobs:
|
||||
name: exo
|
||||
authToken: "${{ secrets.CACHIX_AUTH_TOKEN }}"
|
||||
|
||||
- name: Build Metal packages (macOS only)
|
||||
if: runner.os == 'macOS'
|
||||
run: |
|
||||
# Try to build metal-toolchain first (may succeed via cachix cache hit)
|
||||
if nix build .#metal-toolchain 2>/dev/null; then
|
||||
echo "metal-toolchain built successfully (likely cache hit)"
|
||||
else
|
||||
echo "metal-toolchain build failed, extracting from Xcode..."
|
||||
|
||||
NAR_HASH="sha256-ayR5mXN4sZAddwKEG2OszGRF93k9ZFc7H0yi2xbylQw="
|
||||
NAR_NAME="metal-toolchain-17C48.nar"
|
||||
|
||||
# Use RUNNER_TEMP to avoid /tmp symlink issues on macOS
|
||||
WORK_DIR="${RUNNER_TEMP}/metal-work"
|
||||
mkdir -p "$WORK_DIR"
|
||||
|
||||
# Download the Metal toolchain component
|
||||
xcodebuild -downloadComponent MetalToolchain
|
||||
|
||||
# Find and mount the DMG
|
||||
DMG_PATH=$(find /System/Library/AssetsV2/com_apple_MobileAsset_MetalToolchain -name '*.dmg' 2>/dev/null | head -1)
|
||||
if [ -z "$DMG_PATH" ]; then
|
||||
echo "Error: Could not find Metal toolchain DMG"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Found DMG at: $DMG_PATH"
|
||||
hdiutil attach "$DMG_PATH" -mountpoint "${WORK_DIR}/metal-dmg"
|
||||
|
||||
# Copy the toolchain
|
||||
cp -R "${WORK_DIR}/metal-dmg/Metal.xctoolchain" "${WORK_DIR}/metal-export"
|
||||
hdiutil detach "${WORK_DIR}/metal-dmg"
|
||||
|
||||
# Create NAR and add to store
|
||||
nix nar pack "${WORK_DIR}/metal-export" > "${WORK_DIR}/${NAR_NAME}"
|
||||
STORE_PATH=$(nix store add --mode flat "${WORK_DIR}/${NAR_NAME}")
|
||||
echo "Added NAR to store: $STORE_PATH"
|
||||
|
||||
# Verify the hash matches
|
||||
ACTUAL_HASH=$(nix hash file "${WORK_DIR}/${NAR_NAME}")
|
||||
if [ "$ACTUAL_HASH" != "$NAR_HASH" ]; then
|
||||
echo "Warning: NAR hash mismatch!"
|
||||
echo "Expected: $NAR_HASH"
|
||||
echo "Actual: $ACTUAL_HASH"
|
||||
echo "The metal-toolchain.nix may need updating"
|
||||
fi
|
||||
|
||||
# Clean up
|
||||
rm -rf "$WORK_DIR"
|
||||
|
||||
# Retry the build now that NAR is in store
|
||||
nix build .#metal-toolchain
|
||||
fi
|
||||
|
||||
- name: Build all Nix outputs
|
||||
run: |
|
||||
nix flake show --json | jq -r '
|
||||
@@ -102,16 +134,3 @@ jobs:
|
||||
|
||||
- name: Run nix flake check
|
||||
run: nix flake check
|
||||
|
||||
- name: Run pytest (macOS only)
|
||||
if: runner.os == 'macOS'
|
||||
run: |
|
||||
# Build the test environment (requires relaxed sandbox for uv2nix on macOS)
|
||||
TEST_ENV=$(nix build '.#exo-test-env' --option sandbox relaxed --print-out-paths)
|
||||
|
||||
# Run pytest outside sandbox (needs GPU access for MLX)
|
||||
export HOME="$RUNNER_TEMP"
|
||||
export EXO_TESTS=1
|
||||
export EXO_DASHBOARD_DIR="$PWD/dashboard/"
|
||||
export EXO_RESOURCES_DIR="$PWD/resources"
|
||||
$TEST_ENV/bin/python -m pytest src -m "not slow" --import-mode=importlib
|
||||
-12
@@ -28,15 +28,3 @@ target/
|
||||
dashboard/build/
|
||||
dashboard/node_modules/
|
||||
dashboard/.svelte-kit/
|
||||
|
||||
# host config snapshots
|
||||
hosts_*.json
|
||||
.swp
|
||||
|
||||
# bench files
|
||||
bench/**/*.json
|
||||
|
||||
# tmp
|
||||
tmp/models
|
||||
/build/exo
|
||||
/.claude/skills
|
||||
Generated
+31
@@ -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>
|
||||
Generated
+1
-1
@@ -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
@@ -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>
|
||||
Generated
+3
@@ -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>
|
||||
@@ -1,7 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
if "TOKENIZERS_PARALLELISM" not in os.environ: ...
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,47 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
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):
|
||||
def call_before_loop(
|
||||
self,
|
||||
seed: int,
|
||||
prompt: str,
|
||||
latents: mx.array,
|
||||
config: Config,
|
||||
canny_image: PIL.Image.Image | None = ...,
|
||||
depth_image: PIL.Image.Image | None = ...,
|
||||
) -> None: ...
|
||||
|
||||
class InLoopCallback(Protocol):
|
||||
def call_in_loop(
|
||||
self,
|
||||
t: int,
|
||||
seed: int,
|
||||
prompt: str,
|
||||
latents: mx.array,
|
||||
config: Config,
|
||||
time_steps: tqdm,
|
||||
) -> None: ...
|
||||
|
||||
class AfterLoopCallback(Protocol):
|
||||
def call_after_loop(
|
||||
self, seed: int, prompt: str, latents: mx.array, config: Config
|
||||
) -> None: ...
|
||||
|
||||
class InterruptCallback(Protocol):
|
||||
def call_interrupt(
|
||||
self,
|
||||
t: int,
|
||||
seed: int,
|
||||
prompt: str,
|
||||
latents: mx.array,
|
||||
config: Config,
|
||||
time_steps: tqdm,
|
||||
) -> None: ...
|
||||
@@ -1,24 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from mflux.callbacks.callback import (
|
||||
AfterLoopCallback,
|
||||
BeforeLoopCallback,
|
||||
InLoopCallback,
|
||||
InterruptCallback,
|
||||
)
|
||||
from mflux.callbacks.generation_context import GenerationContext
|
||||
from mflux.models.common.config.config import Config
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class CallbackRegistry:
|
||||
def __init__(self) -> None: ...
|
||||
def register(self, callback) -> None: ...
|
||||
def start(self, seed: int, prompt: str, config: Config) -> GenerationContext: ...
|
||||
def before_loop_callbacks(self) -> list[BeforeLoopCallback]: ...
|
||||
def in_loop_callbacks(self) -> list[InLoopCallback]: ...
|
||||
def after_loop_callbacks(self) -> list[AfterLoopCallback]: ...
|
||||
def interrupt_callbacks(self) -> list[InterruptCallback]: ...
|
||||
@@ -1,29 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
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
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class GenerationContext:
|
||||
def __init__(
|
||||
self, registry: CallbackRegistry, seed: int, prompt: str, config: Config
|
||||
) -> None: ...
|
||||
def before_loop(
|
||||
self,
|
||||
latents: mx.array,
|
||||
*,
|
||||
canny_image: PIL.Image.Image | None = ...,
|
||||
depth_image: PIL.Image.Image | None = ...,
|
||||
) -> None: ...
|
||||
def in_loop(self, t: int, latents: mx.array, time_steps: tqdm = ...) -> None: ...
|
||||
def after_loop(self, latents: mx.array) -> None: ...
|
||||
def interruption(
|
||||
self, t: int, latents: mx.array, time_steps: tqdm = ...
|
||||
) -> None: ...
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,22 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
BATTERY_PERCENTAGE_STOP_LIMIT = ...
|
||||
CONTROLNET_STRENGTH = ...
|
||||
DEFAULT_DEV_FILL_GUIDANCE = ...
|
||||
DEFAULT_DEPTH_GUIDANCE = ...
|
||||
DIMENSION_STEP_PIXELS = ...
|
||||
GUIDANCE_SCALE = ...
|
||||
GUIDANCE_SCALE_KONTEXT = ...
|
||||
IMAGE_STRENGTH = ...
|
||||
MODEL_CHOICES = ...
|
||||
MODEL_INFERENCE_STEPS = ...
|
||||
QUANTIZE_CHOICES = ...
|
||||
if os.environ.get("MFLUX_CACHE_DIR"):
|
||||
MFLUX_CACHE_DIR = ...
|
||||
else:
|
||||
MFLUX_CACHE_DIR = ...
|
||||
MFLUX_LORA_CACHE_DIR = ...
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,8 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.common.config.config import Config
|
||||
from mflux.models.common.config.model_config import ModelConfig
|
||||
|
||||
__all__ = ["Config", "ModelConfig"]
|
||||
@@ -1,66 +0,0 @@
|
||||
"""
|
||||
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
|
||||
from mflux.models.common.config.model_config import ModelConfig
|
||||
|
||||
logger = ...
|
||||
|
||||
class Config:
|
||||
def __init__(
|
||||
self,
|
||||
model_config: ModelConfig,
|
||||
num_inference_steps: int = ...,
|
||||
height: int = ...,
|
||||
width: int = ...,
|
||||
guidance: float = ...,
|
||||
image_path: Path | str | None = ...,
|
||||
image_strength: float | None = ...,
|
||||
depth_image_path: Path | str | None = ...,
|
||||
redux_image_paths: list[Path | str] | None = ...,
|
||||
redux_image_strengths: list[float] | None = ...,
|
||||
masked_image_path: Path | str | None = ...,
|
||||
controlnet_strength: float | None = ...,
|
||||
scheduler: str = ...,
|
||||
) -> None: ...
|
||||
@property
|
||||
def height(self) -> int: ...
|
||||
@property
|
||||
def width(self) -> int: ...
|
||||
@width.setter
|
||||
def width(self, value): # -> None:
|
||||
...
|
||||
@property
|
||||
def image_seq_len(self) -> int: ...
|
||||
@property
|
||||
def guidance(self) -> float: ...
|
||||
@property
|
||||
def num_inference_steps(self) -> int: ...
|
||||
@property
|
||||
def precision(self) -> mx.Dtype: ...
|
||||
@property
|
||||
def num_train_steps(self) -> int: ...
|
||||
@property
|
||||
def image_path(self) -> Path | None: ...
|
||||
@property
|
||||
def image_strength(self) -> float | None: ...
|
||||
@property
|
||||
def depth_image_path(self) -> Path | None: ...
|
||||
@property
|
||||
def redux_image_paths(self) -> list[Path] | None: ...
|
||||
@property
|
||||
def redux_image_strengths(self) -> list[float] | None: ...
|
||||
@property
|
||||
def masked_image_path(self) -> Path | None: ...
|
||||
@property
|
||||
def init_time_step(self) -> int: ...
|
||||
@property
|
||||
def time_steps(self) -> tqdm: ...
|
||||
@property
|
||||
def controlnet_strength(self) -> float | None: ...
|
||||
@property
|
||||
def scheduler(self) -> Any: ...
|
||||
@@ -1,86 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from functools import lru_cache
|
||||
from typing import Literal
|
||||
|
||||
class ModelConfig:
|
||||
precision: mx.Dtype = ...
|
||||
def __init__(
|
||||
self,
|
||||
priority: int,
|
||||
aliases: list[str],
|
||||
model_name: str,
|
||||
base_model: str | None,
|
||||
controlnet_model: str | None,
|
||||
custom_transformer_model: str | None,
|
||||
num_train_steps: int | None,
|
||||
max_sequence_length: int | None,
|
||||
supports_guidance: bool | None,
|
||||
requires_sigma_shift: bool | None,
|
||||
transformer_overrides: dict | None = ...,
|
||||
) -> None: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def schnell() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_kontext() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_fill() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_redux() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_depth() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_controlnet_canny() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def schnell_controlnet_canny() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_controlnet_upscaler() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def dev_fill_catvton() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def krea_dev() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def flux2_klein_4b() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def flux2_klein_9b() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def qwen_image() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def qwen_image_edit() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def fibo() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def z_image_turbo() -> ModelConfig: ...
|
||||
@staticmethod
|
||||
@lru_cache
|
||||
def seedvr2_3b() -> ModelConfig: ...
|
||||
def x_embedder_input_dim(self) -> int: ...
|
||||
def is_canny(self) -> bool: ...
|
||||
@staticmethod
|
||||
def from_name(
|
||||
model_name: str, base_model: Literal["dev", "schnell", "krea-dev"] | None = ...
|
||||
) -> ModelConfig: ...
|
||||
|
||||
AVAILABLE_MODELS = ...
|
||||
@@ -1,7 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,49 +0,0 @@
|
||||
"""
|
||||
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
|
||||
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
|
||||
from mflux.models.qwen.latent_creator.qwen_latent_creator import QwenLatentCreator
|
||||
from mflux.models.z_image.latent_creator.z_image_latent_creator import (
|
||||
ZImageLatentCreator,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
LatentCreatorType: TypeAlias = type[
|
||||
FiboLatentCreator | FluxLatentCreator | QwenLatentCreator | ZImageLatentCreator
|
||||
]
|
||||
|
||||
class Img2Img:
|
||||
def __init__(
|
||||
self,
|
||||
vae: nn.Module,
|
||||
latent_creator: LatentCreatorType,
|
||||
sigmas: mx.array,
|
||||
init_time_step: int,
|
||||
image_path: str | Path | None,
|
||||
tiling_config: TilingConfig | None = ...,
|
||||
) -> None: ...
|
||||
|
||||
class LatentCreator:
|
||||
@staticmethod
|
||||
def create_for_txt2img_or_img2img(
|
||||
seed: int, height: int, width: int, img2img: Img2Img
|
||||
) -> mx.array: ...
|
||||
@staticmethod
|
||||
def encode_image(
|
||||
vae: nn.Module,
|
||||
image_path: str | Path,
|
||||
height: int,
|
||||
width: int,
|
||||
tiling_config: TilingConfig | None = ...,
|
||||
) -> mx.array: ...
|
||||
@staticmethod
|
||||
def add_noise_by_interpolation(
|
||||
clean: mx.array, noise: mx.array, sigma: float
|
||||
) -> mx.array: ...
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,13 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mlx import nn
|
||||
from mflux.models.common.lora.layer.linear_lora_layer import LoRALinear
|
||||
|
||||
class FusedLoRALinear(nn.Module):
|
||||
def __init__(
|
||||
self, base_linear: nn.Linear | nn.QuantizedLinear, loras: list[LoRALinear]
|
||||
) -> None: ...
|
||||
def __call__(self, x): # -> array:
|
||||
...
|
||||
@@ -1,22 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mlx import nn
|
||||
|
||||
class LoRALinear(nn.Module):
|
||||
@staticmethod
|
||||
def from_linear(
|
||||
linear: nn.Linear | nn.QuantizedLinear, r: int = ..., scale: float = ...
|
||||
): # -> LoRALinear:
|
||||
...
|
||||
def __init__(
|
||||
self,
|
||||
input_dims: int,
|
||||
output_dims: int,
|
||||
r: int = ...,
|
||||
scale: float = ...,
|
||||
bias: bool = ...,
|
||||
) -> None: ...
|
||||
def __call__(self, x): # -> array:
|
||||
...
|
||||
@@ -1,26 +0,0 @@
|
||||
"""
|
||||
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
|
||||
from mflux.models.common.lora.mapping.lora_mapping import LoRATarget
|
||||
|
||||
@dataclass
|
||||
class PatternMatch:
|
||||
source_pattern: str
|
||||
target_path: str
|
||||
matrix_name: str
|
||||
transpose: bool
|
||||
transform: Callable[[mx.array], mx.array] | None = ...
|
||||
|
||||
class LoRALoader:
|
||||
@staticmethod
|
||||
def load_and_apply_lora(
|
||||
lora_mapping: list[LoRATarget],
|
||||
transformer: nn.Module,
|
||||
lora_paths: list[str] | None = ...,
|
||||
lora_scales: list[float] | None = ...,
|
||||
) -> tuple[list[str], list[float]]: ...
|
||||
@@ -1,21 +0,0 @@
|
||||
"""
|
||||
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
|
||||
|
||||
@dataclass
|
||||
class LoRATarget:
|
||||
model_path: str
|
||||
possible_up_patterns: List[str]
|
||||
possible_down_patterns: List[str]
|
||||
possible_alpha_patterns: List[str] = ...
|
||||
up_transform: Callable[[mx.array], mx.array] | None = ...
|
||||
down_transform: Callable[[mx.array], mx.array] | None = ...
|
||||
|
||||
class LoRAMapping(Protocol):
|
||||
@staticmethod
|
||||
def get_mapping() -> List[LoRATarget]: ...
|
||||
@@ -1,9 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.nn as nn
|
||||
|
||||
class LoRASaver:
|
||||
@staticmethod
|
||||
def bake_and_strip_lora(module: nn.Module) -> nn.Module: ...
|
||||
@@ -1,35 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
class LoraTransforms:
|
||||
@staticmethod
|
||||
def split_q_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_k_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_v_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_q_down(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_k_down(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_v_down(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_q_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_k_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_v_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_mlp_up(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_q_down(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_k_down(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_v_down(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def split_single_mlp_down(tensor: mx.array) -> mx.array: ...
|
||||
@@ -1,17 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.common.resolution.config_resolution import ConfigResolution
|
||||
from mflux.models.common.resolution.lora_resolution import LoraResolution
|
||||
from mflux.models.common.resolution.path_resolution import PathResolution
|
||||
from mflux.models.common.resolution.quantization_resolution import (
|
||||
QuantizationResolution,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ConfigResolution",
|
||||
"LoraResolution",
|
||||
"PathResolution",
|
||||
"QuantizationResolution",
|
||||
]
|
||||
@@ -1,39 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
from typing import NamedTuple
|
||||
|
||||
class QuantizationAction(Enum):
|
||||
NONE = ...
|
||||
STORED = ...
|
||||
REQUESTED = ...
|
||||
|
||||
class PathAction(Enum):
|
||||
LOCAL = ...
|
||||
HUGGINGFACE_CACHED = ...
|
||||
HUGGINGFACE = ...
|
||||
ERROR = ...
|
||||
|
||||
class LoraAction(Enum):
|
||||
LOCAL = ...
|
||||
REGISTRY = ...
|
||||
HUGGINGFACE_COLLECTION_CACHED = ...
|
||||
HUGGINGFACE_COLLECTION = ...
|
||||
HUGGINGFACE_REPO_CACHED = ...
|
||||
HUGGINGFACE_REPO = ...
|
||||
ERROR = ...
|
||||
|
||||
class ConfigAction(Enum):
|
||||
EXACT_MATCH = ...
|
||||
EXPLICIT_BASE = ...
|
||||
INFER_SUBSTRING = ...
|
||||
ERROR = ...
|
||||
|
||||
class Rule(NamedTuple):
|
||||
priority: int
|
||||
name: str
|
||||
check: str
|
||||
action: QuantizationAction | PathAction | LoraAction | ConfigAction
|
||||
...
|
||||
@@ -1,14 +0,0 @@
|
||||
"""
|
||||
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: ...
|
||||
logger = ...
|
||||
|
||||
class ConfigResolution:
|
||||
RULES = ...
|
||||
@staticmethod
|
||||
def resolve(model_name: str, base_model: str | None = ...) -> ModelConfig: ...
|
||||
@@ -1,21 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
logger = ...
|
||||
|
||||
class LoraResolution:
|
||||
RULES = ...
|
||||
_registry: dict[str, Path] = ...
|
||||
@staticmethod
|
||||
def resolve(path: str) -> str: ...
|
||||
@staticmethod
|
||||
def resolve_paths(paths: list[str] | None) -> list[str]: ...
|
||||
@staticmethod
|
||||
def resolve_scales(scales: list[float] | None, num_paths: int) -> list[float]: ...
|
||||
@staticmethod
|
||||
def get_registry() -> dict[str, Path]: ...
|
||||
@staticmethod
|
||||
def discover_files(library_paths: list[Path]) -> dict[str, Path]: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
logger = ...
|
||||
|
||||
class PathResolution:
|
||||
RULES = ...
|
||||
@staticmethod
|
||||
def resolve(path: str | None, patterns: list[str] | None = ...) -> Path | None: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
logger = ...
|
||||
|
||||
class QuantizationResolution:
|
||||
RULES = ...
|
||||
@staticmethod
|
||||
def resolve(
|
||||
stored: int | None, requested: int | None
|
||||
) -> tuple[int | None, str | None]: ...
|
||||
@@ -1,26 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from .flow_match_euler_discrete_scheduler import FlowMatchEulerDiscreteScheduler
|
||||
from .linear_scheduler import LinearScheduler
|
||||
from .seedvr2_euler_scheduler import SeedVR2EulerScheduler
|
||||
|
||||
__all__ = [
|
||||
"LinearScheduler",
|
||||
"FlowMatchEulerDiscreteScheduler",
|
||||
"SeedVR2EulerScheduler",
|
||||
]
|
||||
|
||||
class SchedulerModuleNotFound(ValueError): ...
|
||||
class SchedulerClassNotFound(ValueError): ...
|
||||
class InvalidSchedulerType(TypeError): ...
|
||||
|
||||
SCHEDULER_REGISTRY = ...
|
||||
|
||||
def register_contrib(scheduler_object, scheduler_name=...): # -> None:
|
||||
...
|
||||
def try_import_external_scheduler(
|
||||
scheduler_object_path: str,
|
||||
): # -> type[BaseScheduler]:
|
||||
...
|
||||
@@ -1,16 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
class BaseScheduler(ABC):
|
||||
@property
|
||||
@abstractmethod
|
||||
def sigmas(self) -> mx.array: ...
|
||||
@abstractmethod
|
||||
def step(
|
||||
self, noise: mx.array, timestep: int, latents: mx.array, **kwargs
|
||||
) -> mx.array: ...
|
||||
def scale_model_input(self, latents: mx.array, t: int) -> mx.array: ...
|
||||
@@ -1,26 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from typing import TYPE_CHECKING
|
||||
from mflux.models.common.config.config import Config
|
||||
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class FlowMatchEulerDiscreteScheduler(BaseScheduler):
|
||||
def __init__(self, config: Config) -> None: ...
|
||||
@property
|
||||
def sigmas(self) -> mx.array: ...
|
||||
@property
|
||||
def timesteps(self) -> mx.array: ...
|
||||
def set_image_seq_len(self, image_seq_len: int) -> None: ...
|
||||
@staticmethod
|
||||
def get_timesteps_and_sigmas(
|
||||
image_seq_len: int, num_inference_steps: int, num_train_timesteps: int = ...
|
||||
) -> tuple[mx.array, mx.array]: ...
|
||||
def step(
|
||||
self, noise: mx.array, timestep: int, latents: mx.array, **kwargs
|
||||
) -> mx.array: ...
|
||||
def scale_model_input(self, latents: mx.array, t: int) -> mx.array: ...
|
||||
@@ -1,20 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from typing import TYPE_CHECKING
|
||||
from mflux.models.common.config.config import Config
|
||||
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class LinearScheduler(BaseScheduler):
|
||||
def __init__(self, config: Config) -> None: ...
|
||||
@property
|
||||
def sigmas(self) -> mx.array: ...
|
||||
@property
|
||||
def timesteps(self) -> mx.array: ...
|
||||
def step(
|
||||
self, noise: mx.array, timestep: int, latents: mx.array, **kwargs
|
||||
) -> mx.array: ...
|
||||
@@ -1,20 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from typing import TYPE_CHECKING
|
||||
from mflux.models.common.config.config import Config
|
||||
from mflux.models.common.schedulers.base_scheduler import BaseScheduler
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class SeedVR2EulerScheduler(BaseScheduler):
|
||||
def __init__(self, config: Config) -> None: ...
|
||||
@property
|
||||
def timesteps(self) -> mx.array: ...
|
||||
@property
|
||||
def sigmas(self) -> mx.array: ...
|
||||
def step(
|
||||
self, noise: mx.array, timestep: int, latents: mx.array, **kwargs
|
||||
) -> mx.array: ...
|
||||
@@ -1,24 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.common.tokenizer.tokenizer import (
|
||||
BaseTokenizer,
|
||||
LanguageTokenizer,
|
||||
Tokenizer,
|
||||
VisionLanguageTokenizer,
|
||||
)
|
||||
from mflux.models.common.tokenizer.tokenizer_loader import TokenizerLoader
|
||||
from mflux.models.common.tokenizer.tokenizer_output import TokenizerOutput
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
__all__ = [
|
||||
"Tokenizer",
|
||||
"BaseTokenizer",
|
||||
"LanguageTokenizer",
|
||||
"VisionLanguageTokenizer",
|
||||
"TokenizerLoader",
|
||||
"TokenizerOutput",
|
||||
]
|
||||
@@ -1,74 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Protocol, runtime_checkable
|
||||
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.
|
||||
"""
|
||||
|
||||
@runtime_checkable
|
||||
class Tokenizer(Protocol):
|
||||
tokenizer: PreTrainedTokenizer
|
||||
def tokenize(
|
||||
self,
|
||||
prompt: str | list[str],
|
||||
images: list[Image.Image] | None = ...,
|
||||
max_length: int | None = ...,
|
||||
**kwargs,
|
||||
) -> TokenizerOutput: ...
|
||||
|
||||
class BaseTokenizer(ABC):
|
||||
def __init__(
|
||||
self, tokenizer: PreTrainedTokenizer, max_length: int = ...
|
||||
) -> None: ...
|
||||
@abstractmethod
|
||||
def tokenize(
|
||||
self,
|
||||
prompt: str | list[str],
|
||||
images: list[Image.Image] | None = ...,
|
||||
max_length: int | None = ...,
|
||||
**kwargs,
|
||||
) -> TokenizerOutput: ...
|
||||
|
||||
class LanguageTokenizer(BaseTokenizer):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
max_length: int = ...,
|
||||
padding: str = ...,
|
||||
return_attention_mask: bool = ...,
|
||||
template: str | None = ...,
|
||||
use_chat_template: bool = ...,
|
||||
chat_template_kwargs: dict | None = ...,
|
||||
add_special_tokens: bool = ...,
|
||||
) -> None: ...
|
||||
def tokenize(
|
||||
self,
|
||||
prompt: str | list[str],
|
||||
images: list[Image.Image] | None = ...,
|
||||
max_length: int | None = ...,
|
||||
**kwargs,
|
||||
) -> TokenizerOutput: ...
|
||||
|
||||
class VisionLanguageTokenizer(BaseTokenizer):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
processor,
|
||||
max_length: int = ...,
|
||||
template: str | None = ...,
|
||||
image_token: str = ...,
|
||||
) -> None: ...
|
||||
def tokenize(
|
||||
self,
|
||||
prompt: str | list[str],
|
||||
images: list[Image.Image] | None = ...,
|
||||
max_length: int | None = ...,
|
||||
**kwargs,
|
||||
) -> TokenizerOutput: ...
|
||||
@@ -1,22 +0,0 @@
|
||||
"""
|
||||
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
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class TokenizerLoader:
|
||||
@staticmethod
|
||||
def load(definition: TokenizerDefinition, model_path: str) -> BaseTokenizer: ...
|
||||
@staticmethod
|
||||
def load_all(
|
||||
definitions: list[TokenizerDefinition],
|
||||
model_path: str,
|
||||
max_length_overrides: dict[str, int] | None = ...,
|
||||
) -> dict[str, BaseTokenizer]: ...
|
||||
@@ -1,17 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from dataclasses import dataclass
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
@dataclass
|
||||
class TokenizerOutput:
|
||||
input_ids: mx.array
|
||||
attention_mask: mx.array
|
||||
pixel_values: mx.array | None = ...
|
||||
image_grid_thw: mx.array | None = ...
|
||||
@@ -1,8 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.common.vae.tiling_config import TilingConfig
|
||||
from mflux.models.common.vae.vae_tiler import VAETiler
|
||||
|
||||
__all__ = ["TilingConfig", "VAETiler"]
|
||||
@@ -1,13 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TilingConfig:
|
||||
vae_decode_tiles_per_dim: int | None = ...
|
||||
vae_decode_overlap: int = ...
|
||||
vae_encode_tiled: bool = ...
|
||||
vae_encode_tile_size: int = ...
|
||||
vae_encode_tile_overlap: int = ...
|
||||
@@ -1,27 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from typing import Callable
|
||||
|
||||
class VAETiler:
|
||||
@staticmethod
|
||||
def encode_image_tiled(
|
||||
*,
|
||||
image: mx.array,
|
||||
encode_fn: Callable[[mx.array], mx.array],
|
||||
latent_channels: int,
|
||||
tile_size: tuple[int, int] = ...,
|
||||
tile_overlap: tuple[int, int] = ...,
|
||||
spatial_scale: int = ...,
|
||||
) -> mx.array: ...
|
||||
@staticmethod
|
||||
def decode_image_tiled(
|
||||
*,
|
||||
latent: mx.array,
|
||||
decode_fn: Callable[[mx.array], mx.array],
|
||||
tile_size: tuple[int, int] = ...,
|
||||
tile_overlap: tuple[int, int] = ...,
|
||||
spatial_scale: int = ...,
|
||||
) -> mx.array: ...
|
||||
@@ -1,17 +0,0 @@
|
||||
"""
|
||||
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
|
||||
|
||||
class VAEUtil:
|
||||
@staticmethod
|
||||
def encode(
|
||||
vae: nn.Module, image: mx.array, tiling_config: TilingConfig | None = ...
|
||||
) -> mx.array: ...
|
||||
@staticmethod
|
||||
def decode(
|
||||
vae: nn.Module, latent: mx.array, tiling_config: TilingConfig | None = ...
|
||||
) -> mx.array: ...
|
||||
@@ -1,18 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.common.weights.loading.loaded_weights import LoadedWeights, MetaData
|
||||
from mflux.models.common.weights.loading.weight_applier import WeightApplier
|
||||
from mflux.models.common.weights.loading.weight_definition import ComponentDefinition
|
||||
from mflux.models.common.weights.loading.weight_loader import WeightLoader
|
||||
from mflux.models.common.weights.saving.model_saver import ModelSaver
|
||||
|
||||
__all__ = [
|
||||
"ComponentDefinition",
|
||||
"LoadedWeights",
|
||||
"MetaData",
|
||||
"ModelSaver",
|
||||
"WeightApplier",
|
||||
"WeightLoader",
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
@dataclass
|
||||
class MetaData:
|
||||
quantization_level: int | None = ...
|
||||
mflux_version: str | None = ...
|
||||
|
||||
@dataclass
|
||||
class LoadedWeights:
|
||||
components: dict[str, dict]
|
||||
meta_data: MetaData
|
||||
def __getattr__(self, name: str) -> dict | None: ...
|
||||
def num_transformer_blocks(self, component_name: str = ...) -> int: ...
|
||||
def num_single_transformer_blocks(self, component_name: str = ...) -> int: ...
|
||||
@@ -1,30 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.nn as nn
|
||||
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,
|
||||
WeightDefinitionType,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class WeightApplier:
|
||||
@staticmethod
|
||||
def apply_and_quantize_single(
|
||||
weights: LoadedWeights,
|
||||
model: nn.Module,
|
||||
component: ComponentDefinition,
|
||||
quantize_arg: int | None,
|
||||
quantization_predicate=...,
|
||||
) -> int | None: ...
|
||||
@staticmethod
|
||||
def apply_and_quantize(
|
||||
weights: LoadedWeights,
|
||||
models: dict[str, nn.Module],
|
||||
quantize_arg: int | None,
|
||||
weight_definition: WeightDefinitionType,
|
||||
) -> int | None: ...
|
||||
@@ -1,73 +0,0 @@
|
||||
"""
|
||||
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 mflux.models.common.tokenizer.tokenizer import BaseTokenizer
|
||||
from mflux.models.depth_pro.weights.depth_pro_weight_definition import (
|
||||
DepthProWeightDefinition,
|
||||
)
|
||||
from mflux.models.fibo.weights.fibo_weight_definition import FIBOWeightDefinition
|
||||
from mflux.models.fibo_vlm.weights.fibo_vlm_weight_definition import (
|
||||
FIBOVLMWeightDefinition,
|
||||
)
|
||||
from mflux.models.flux.weights.flux_weight_definition import FluxWeightDefinition
|
||||
from mflux.models.qwen.weights.qwen_weight_definition import QwenWeightDefinition
|
||||
from mflux.models.seedvr2.weights.seedvr2_weight_definition import (
|
||||
SeedVR2WeightDefinition,
|
||||
)
|
||||
from mflux.models.z_image.weights.z_image_weight_definition import (
|
||||
ZImageWeightDefinition,
|
||||
)
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
if TYPE_CHECKING:
|
||||
WeightDefinitionType: TypeAlias = type[
|
||||
FluxWeightDefinition
|
||||
| FIBOWeightDefinition
|
||||
| FIBOVLMWeightDefinition
|
||||
| QwenWeightDefinition
|
||||
| ZImageWeightDefinition
|
||||
| SeedVR2WeightDefinition
|
||||
| DepthProWeightDefinition
|
||||
]
|
||||
|
||||
@dataclass
|
||||
class ComponentDefinition:
|
||||
name: str
|
||||
hf_subdir: str
|
||||
mapping_getter: Callable[[], List[WeightTarget]] | None = ...
|
||||
model_attr: str | None = ...
|
||||
num_blocks: int | None = ...
|
||||
num_layers: int | None = ...
|
||||
loading_mode: str = ...
|
||||
precision: mx.Dtype | None = ...
|
||||
skip_quantization: bool = ...
|
||||
bulk_transform: Callable[[mx.array], mx.array] | None = ...
|
||||
weight_subkey: str | None = ...
|
||||
download_url: str | None = ...
|
||||
weight_prefix_filters: List[str] | None = ...
|
||||
weight_files: List[str] | None = ...
|
||||
|
||||
@dataclass
|
||||
class TokenizerDefinition:
|
||||
name: str
|
||||
hf_subdir: str
|
||||
tokenizer_class: str = ...
|
||||
fallback_subdirs: List[str] | None = ...
|
||||
download_patterns: List[str] | None = ...
|
||||
encoder_class: type[BaseTokenizer] | None = ...
|
||||
max_length: int = ...
|
||||
padding: str = ...
|
||||
template: str | None = ...
|
||||
use_chat_template: bool = ...
|
||||
chat_template_kwargs: dict | None = ...
|
||||
add_special_tokens: bool = ...
|
||||
processor_class: type | None = ...
|
||||
image_token: str = ...
|
||||
chat_template: str | None = ...
|
||||
@@ -1,23 +0,0 @@
|
||||
"""
|
||||
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,
|
||||
WeightDefinitionType,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
logger = ...
|
||||
|
||||
class WeightLoader:
|
||||
@staticmethod
|
||||
def load_single(
|
||||
component: ComponentDefinition, repo_id: str, file_pattern: str = ...
|
||||
) -> LoadedWeights: ...
|
||||
@staticmethod
|
||||
def load(
|
||||
weight_definition: WeightDefinitionType, model_path: str | None = ...
|
||||
) -> LoadedWeights: ...
|
||||
@@ -1,16 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from typing import Dict, List, Optional
|
||||
from mflux.models.common.weights.mapping.weight_mapping import WeightTarget
|
||||
|
||||
class WeightMapper:
|
||||
@staticmethod
|
||||
def apply_mapping(
|
||||
hf_weights: Dict[str, mx.array],
|
||||
mapping: List[WeightTarget],
|
||||
num_blocks: Optional[int] = ...,
|
||||
num_layers: Optional[int] = ...,
|
||||
) -> Dict: ...
|
||||
@@ -1,23 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, List, Optional, Protocol
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
@dataclass
|
||||
class WeightTarget:
|
||||
to_pattern: str
|
||||
from_pattern: List[str]
|
||||
transform: Optional[Callable[[mx.array], mx.array]] = ...
|
||||
required: bool = ...
|
||||
max_blocks: Optional[int] = ...
|
||||
|
||||
class WeightMapping(Protocol):
|
||||
@staticmethod
|
||||
def get_mapping() -> List[WeightTarget]: ...
|
||||
@@ -1,17 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
class WeightTransforms:
|
||||
@staticmethod
|
||||
def reshape_gamma_to_1d(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def transpose_patch_embed(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def transpose_conv3d_weight(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def transpose_conv2d_weight(tensor: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def transpose_conv_transpose2d_weight(tensor: mx.array) -> mx.array: ...
|
||||
@@ -1,14 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import Any, TYPE_CHECKING
|
||||
from mflux.models.common.weights.loading.weight_definition import WeightDefinitionType
|
||||
|
||||
if TYPE_CHECKING: ...
|
||||
|
||||
class ModelSaver:
|
||||
@staticmethod
|
||||
def save_model(
|
||||
model: Any, bits: int, base_path: str, weight_definition: WeightDefinitionType
|
||||
) -> None: ...
|
||||
@@ -1,9 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.depth_pro.model.depth_pro_model import DepthProModel
|
||||
|
||||
class DepthProInitializer:
|
||||
@staticmethod
|
||||
def init(model: DepthProModel, quantize: int | None = ...) -> None: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class FeatureFusionBlock2d(nn.Module):
|
||||
def __init__(self, num_features: int, deconv: bool = ...) -> None: ...
|
||||
def __call__(self, x0: mx.array, x1: mx.array | None = ...) -> mx.array: ...
|
||||
@@ -1,17 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class MultiresConvDecoder(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
x0_latent: mx.array,
|
||||
x1_latent: mx.array,
|
||||
x0_features: mx.array,
|
||||
x1_features: mx.array,
|
||||
x_global_features: mx.array,
|
||||
) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
def __init__(self, num_features: int) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,20 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from PIL import Image
|
||||
|
||||
@dataclass
|
||||
class DepthResult:
|
||||
depth_image: Image.Image
|
||||
depth_array: mx.array
|
||||
min_depth: float
|
||||
max_depth: float
|
||||
...
|
||||
|
||||
class DepthPro:
|
||||
def __init__(self, quantize: int | None = ...) -> None: ...
|
||||
def create_depth_map(self, image_path: str | Path) -> DepthResult: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class DepthProModel(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(
|
||||
self, x0: mx.array, x1: mx.array, x2: mx.array
|
||||
) -> tuple[mx.array, mx.array]: ...
|
||||
@@ -1,15 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class DepthProUtil:
|
||||
@staticmethod
|
||||
def split(x: mx.array, overlap_ratio: float = ...) -> mx.array: ...
|
||||
@staticmethod
|
||||
def interpolate(x: mx.array, size=..., scale_factor=...): # -> array:
|
||||
...
|
||||
@staticmethod
|
||||
def apply_conv(x: mx.array, conv_module: nn.Module) -> mx.array: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self, dim: int = ..., head_dim: int = ..., num_heads: int = ...
|
||||
) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class DinoVisionTransformer(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, x: mx.array) -> tuple[mx.array, mx.array, mx.array]: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class LayerScale(nn.Module):
|
||||
def __init__(self, dims: int, init_values: float = ...) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class DepthProEncoder(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(
|
||||
self, x0: mx.array, x1: mx.array, x2: mx.array
|
||||
) -> tuple[mx.array, mx.array, mx.array, mx.array, mx.array]: ...
|
||||
@@ -1,16 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class UpSampleBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim_in: int = ...,
|
||||
dim_int: int = ...,
|
||||
dim_out: int = ...,
|
||||
upsample_layers: int = ...,
|
||||
) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
class FOVHead(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
@@ -1,23 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from mflux.models.common.weights.loading.weight_definition import (
|
||||
ComponentDefinition,
|
||||
TokenizerDefinition,
|
||||
)
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
class DepthProWeightDefinition:
|
||||
@staticmethod
|
||||
def get_components() -> List[ComponentDefinition]: ...
|
||||
@staticmethod
|
||||
def get_tokenizers() -> List[TokenizerDefinition]: ...
|
||||
@staticmethod
|
||||
def get_download_patterns() -> List[str]: ...
|
||||
@staticmethod
|
||||
def quantization_predicate(path: str, module) -> bool: ...
|
||||
@@ -1,13 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from mflux.models.common.weights.mapping.weight_mapping import (
|
||||
WeightMapping,
|
||||
WeightTarget,
|
||||
)
|
||||
|
||||
class DepthProWeightMapping(WeightMapping):
|
||||
@staticmethod
|
||||
def get_mapping() -> List[WeightTarget]: ...
|
||||
@@ -1,13 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
class FiboLatentCreator:
|
||||
@staticmethod
|
||||
def create_noise(seed: int, height: int, width: int) -> mx.array: ...
|
||||
@staticmethod
|
||||
def pack_latents(latents: mx.array, height: int, width: int) -> mx.array: ...
|
||||
@staticmethod
|
||||
def unpack_latents(latents: mx.array, height: int, width: int) -> mx.array: ...
|
||||
@@ -1,23 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from mflux.models.common.weights.loading.weight_definition import (
|
||||
ComponentDefinition,
|
||||
TokenizerDefinition,
|
||||
)
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
class FIBOWeightDefinition:
|
||||
@staticmethod
|
||||
def get_components() -> List[ComponentDefinition]: ...
|
||||
@staticmethod
|
||||
def get_tokenizers() -> List[TokenizerDefinition]: ...
|
||||
@staticmethod
|
||||
def get_download_patterns() -> List[str]: ...
|
||||
@staticmethod
|
||||
def quantization_predicate(path: str, module) -> bool: ...
|
||||
@@ -1,17 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from mflux.models.common.weights.mapping.weight_mapping import (
|
||||
WeightMapping,
|
||||
WeightTarget,
|
||||
)
|
||||
|
||||
class FIBOWeightMapping(WeightMapping):
|
||||
@staticmethod
|
||||
def get_transformer_mapping() -> List[WeightTarget]: ...
|
||||
@staticmethod
|
||||
def get_text_encoder_mapping() -> List[WeightTarget]: ...
|
||||
@staticmethod
|
||||
def get_vae_mapping() -> List[WeightTarget]: ...
|
||||
@@ -1,8 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.qwen.tokenizer.qwen_image_processor import QwenImageProcessor
|
||||
|
||||
class Qwen2VLImageProcessor(QwenImageProcessor):
|
||||
def __init__(self) -> None: ...
|
||||
@@ -1,28 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import Optional, Union
|
||||
from PIL import Image
|
||||
|
||||
class Qwen2VLProcessor:
|
||||
def __init__(self, tokenizer) -> None: ...
|
||||
def apply_chat_template(
|
||||
self,
|
||||
messages,
|
||||
tokenize: bool = ...,
|
||||
add_generation_prompt: bool = ...,
|
||||
return_tensors: Optional[str] = ...,
|
||||
return_dict: bool = ...,
|
||||
**kwargs,
|
||||
): # -> dict[Any, Any]:
|
||||
...
|
||||
def __call__(
|
||||
self,
|
||||
text: Optional[Union[str, list[str]]] = ...,
|
||||
images: Optional[Union[Image.Image, list[Image.Image]]] = ...,
|
||||
padding: bool = ...,
|
||||
return_tensors: Optional[str] = ...,
|
||||
**kwargs,
|
||||
): # -> dict[Any, Any]:
|
||||
...
|
||||
@@ -1,24 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from mflux.models.common.weights.loading.weight_definition import (
|
||||
ComponentDefinition,
|
||||
TokenizerDefinition,
|
||||
)
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
QWEN2VL_CHAT_TEMPLATE = ...
|
||||
|
||||
class FIBOVLMWeightDefinition:
|
||||
@staticmethod
|
||||
def get_components() -> List[ComponentDefinition]: ...
|
||||
@staticmethod
|
||||
def get_tokenizers() -> List[TokenizerDefinition]: ...
|
||||
@staticmethod
|
||||
def get_download_patterns() -> List[str]: ...
|
||||
@staticmethod
|
||||
def quantization_predicate(path: str, module) -> bool: ...
|
||||
@@ -1,15 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from mflux.models.common.weights.mapping.weight_mapping import (
|
||||
WeightMapping,
|
||||
WeightTarget,
|
||||
)
|
||||
|
||||
class FIBOVLMWeightMapping(WeightMapping):
|
||||
@staticmethod
|
||||
def get_vlm_decoder_mapping(num_layers: int = ...) -> List[WeightTarget]: ...
|
||||
@staticmethod
|
||||
def get_vlm_visual_mapping(depth: int = ...) -> List[WeightTarget]: ...
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,3 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,53 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from mflux.models.common.config import ModelConfig
|
||||
|
||||
class FluxInitializer:
|
||||
@staticmethod
|
||||
def init(
|
||||
model,
|
||||
model_config: ModelConfig,
|
||||
quantize: int | None,
|
||||
model_path: str | None = ...,
|
||||
lora_paths: list[str] | None = ...,
|
||||
lora_scales: list[float] | None = ...,
|
||||
custom_transformer=...,
|
||||
) -> None: ...
|
||||
@staticmethod
|
||||
def init_depth(
|
||||
model,
|
||||
model_config: ModelConfig,
|
||||
quantize: int | None,
|
||||
model_path: str | None = ...,
|
||||
lora_paths: list[str] | None = ...,
|
||||
lora_scales: list[float] | None = ...,
|
||||
) -> None: ...
|
||||
@staticmethod
|
||||
def init_redux(
|
||||
model,
|
||||
model_config: ModelConfig,
|
||||
quantize: int | None,
|
||||
model_path: str | None = ...,
|
||||
lora_paths: list[str] | None = ...,
|
||||
lora_scales: list[float] | None = ...,
|
||||
) -> None: ...
|
||||
@staticmethod
|
||||
def init_controlnet(
|
||||
model,
|
||||
model_config: ModelConfig,
|
||||
quantize: int | None,
|
||||
model_path: str | None = ...,
|
||||
lora_paths: list[str] | None = ...,
|
||||
lora_scales: list[float] | None = ...,
|
||||
) -> None: ...
|
||||
@staticmethod
|
||||
def init_concept(
|
||||
model,
|
||||
model_config: ModelConfig,
|
||||
quantize: int | None,
|
||||
model_path: str | None = ...,
|
||||
lora_paths: list[str] | None = ...,
|
||||
lora_scales: list[float] | None = ...,
|
||||
) -> None: ...
|
||||
@@ -1,7 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
@@ -1,19 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
class FluxLatentCreator:
|
||||
@staticmethod
|
||||
def create_noise(seed: int, height: int, width: int) -> mx.array: ...
|
||||
@staticmethod
|
||||
def pack_latents(
|
||||
latents: mx.array, height: int, width: int, num_channels_latents: int = ...
|
||||
) -> mx.array: ...
|
||||
@staticmethod
|
||||
def unpack_latents(latents: mx.array, height: int, width: int) -> mx.array: ...
|
||||
@@ -1,7 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
-10
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class CLIPEmbeddings(nn.Module):
|
||||
def __init__(self, dims: int) -> None: ...
|
||||
def __call__(self, tokens: mx.array) -> mx.array: ...
|
||||
@@ -1,14 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
class CLIPEncoder(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, tokens: mx.array) -> mx.array: ...
|
||||
-12
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class CLIPEncoderLayer(nn.Module):
|
||||
def __init__(self, layer: int) -> None: ...
|
||||
def __call__(
|
||||
self, hidden_states: mx.array, causal_attention_mask: mx.array
|
||||
) -> mx.array: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class CLIPMLP(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, hidden_states: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def quick_gelu(input_array: mx.array) -> mx.array: ...
|
||||
-18
@@ -1,18 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class CLIPSdpaAttention(nn.Module):
|
||||
head_dimension = ...
|
||||
batch_size = ...
|
||||
num_heads = ...
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(
|
||||
self, hidden_states: mx.array, causal_attention_mask: mx.array
|
||||
) -> mx.array: ...
|
||||
@staticmethod
|
||||
def reshape_and_transpose(x, batch_size, num_heads, head_dim): # -> array:
|
||||
...
|
||||
-12
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class CLIPTextModel(nn.Module):
|
||||
def __init__(self, dims: int, num_encoder_layers: int) -> None: ...
|
||||
def __call__(self, tokens: mx.array) -> tuple[mx.array, mx.array]: ...
|
||||
@staticmethod
|
||||
def create_causal_attention_mask(input_shape: tuple) -> mx.array: ...
|
||||
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class EncoderCLIP(nn.Module):
|
||||
def __init__(self, num_encoder_layers: int) -> None: ...
|
||||
def __call__(
|
||||
self, tokens: mx.array, causal_attention_mask: mx.array
|
||||
) -> mx.array: ...
|
||||
@@ -1,25 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mflux.models.common.tokenizer import Tokenizer
|
||||
from mflux.models.flux.model.flux_text_encoder.clip_encoder.clip_encoder import (
|
||||
CLIPEncoder,
|
||||
)
|
||||
from mflux.models.flux.model.flux_text_encoder.t5_encoder.t5_encoder import T5Encoder
|
||||
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
class PromptEncoder:
|
||||
@staticmethod
|
||||
def encode_prompt(
|
||||
prompt: str,
|
||||
prompt_cache: dict[str, tuple[mx.array, mx.array]],
|
||||
t5_tokenizer: Tokenizer,
|
||||
clip_tokenizer: Tokenizer,
|
||||
t5_text_encoder: T5Encoder,
|
||||
clip_text_encoder: CLIPEncoder,
|
||||
) -> tuple[mx.array, mx.array]: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class T5Attention(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, hidden_states: mx.array) -> mx.array: ...
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class T5Block(nn.Module):
|
||||
def __init__(self, layer: int) -> None: ...
|
||||
def __call__(self, hidden_states: mx.array) -> mx.array: ...
|
||||
-12
@@ -1,12 +0,0 @@
|
||||
"""
|
||||
This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx import nn
|
||||
|
||||
class T5DenseReluDense(nn.Module):
|
||||
def __init__(self) -> None: ...
|
||||
def __call__(self, hidden_states: mx.array) -> mx.array: ...
|
||||
@staticmethod
|
||||
def new_gelu(input_array: mx.array) -> mx.array: ...
|
||||
Loaded 100 of 768 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user