diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index 387ef1d..82b6669 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -87,6 +87,13 @@ jobs: if: steps.release.outputs.publish == 'true' run: uv tool run twine check dist/* + - name: publish package + if: steps.release.outputs.publish == 'true' + env: + TWINE_USERNAME: __token__ + TWINE_PASSWORD: ${{ secrets.PYPI_API_TOKEN }} + run: uv tool run twine upload --skip-existing --verbose dist/* + - name: create tag if: steps.release.outputs.publish == 'true' && steps.release.outputs.create_tag == 'true' env: @@ -95,13 +102,6 @@ jobs: git tag "$TAG" "$GITHUB_SHA" git push origin "$TAG" - - name: publish package - if: steps.release.outputs.publish == 'true' - env: - TWINE_USERNAME: __token__ - TWINE_PASSWORD: ${{ secrets.PYPI_API_TOKEN }} - run: uv tool run twine upload dist/* - - name: create github release if: steps.release.outputs.publish == 'true' env: diff --git a/.gitignore b/.gitignore index 073f444..dd0bded 100644 --- a/.gitignore +++ b/.gitignore @@ -9,3 +9,6 @@ examples/tinystories-llm/TinyStoriesV2-GPT4-valid.txt # Virtual environments .venv + +# Environment files +.env diff --git a/README.md b/README.md index c48c9e4..9183084 100644 --- a/README.md +++ b/README.md @@ -1,15 +1,14 @@ # torchmlx -torchmlx is a pytorch-shaped compatibility layer that uses mlx on apple silicon and pytorch elsewhere. +torchmlx is a pytorch-shaped compatibility layer that uses mlx on apple silicon and pytorch elsewhere. it is experimental and built first for educational use. -NOTE: it is experimental and built first for educational use. If you find a bug, please create an issue on the repo. +```bash +pip install pytorchmlx +``` ```python import torchmlx as torch from torchmlx import nn, optim - -model = nn.Linear(4, 2) -optimizer = optim.AdamW(model.parameters(), lr=3e-4) ``` see the [tinystories example](examples/tinystories-llm/train.py) and [compatibility details](docs/compatibility.md). diff --git a/pyproject.toml b/pyproject.toml index f84b31e..23a69cd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] -name = "torchmlx" -version = "0.1.0" +name = "pytorchmlx" +version = "0.0.1" description = "An educational PyTorch-shaped interface for MLX and PyTorch" readme = "README.md" requires-python = ">=3.10" diff --git a/uv.lock b/uv.lock index 8a6d2d5..b45b1b6 100644 --- a/uv.lock +++ b/uv.lock @@ -686,6 +686,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878 }, ] +[[package]] +name = "pytorchmlx" +version = "0.0.1" +source = { editable = "." } +dependencies = [ + { name = "mlx" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "2.5.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "torch" }, +] + +[package.metadata] +requires-dist = [ + { name = "mlx", specifier = ">=0.32.2,<0.33" }, + { name = "numpy" }, + { name = "torch", specifier = ">=2.4,<3" }, +] + [[package]] name = "setuptools" version = "84.0.0" @@ -755,25 +774,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f8/c3/72ae1f02747b1f012e1975743e48cd608f83095d7f9ce58de78b79248b35/torch-2.14.0-cp314-cp314t-win_amd64.whl", hash = "sha256:731784e3914843c6bcc7aba3987ff7610ac57dbbc816a5d6b9b62e04c240a641", size = 124400194 }, ] -[[package]] -name = "torchmlx" -version = "0.1.0" -source = { editable = "." } -dependencies = [ - { name = "mlx" }, - { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, - { name = "numpy", version = "2.5.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, - { name = "torch" }, -] - -[package.metadata] -requires-dist = [ - { name = "mlx", specifier = ">=0.32.2,<0.33" }, - { name = "numpy" }, - { name = "torch", specifier = ">=2.4,<3" }, -] - [[package]] name = "triton" version = "3.8.0"