From 9638894a416339fc573f77da472d1f0aa452c29d Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 23:25:44 +0000 Subject: [PATCH 1/2] [pre-commit.ci] pre-commit suggestions MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit updates: - [github.com/codespell-project/codespell: v2.4.2 → v2.4.3](https://github.com/codespell-project/codespell/compare/v2.4.2...v2.4.3) - [github.com/astral-sh/ruff-pre-commit: v0.15.9 → v0.16.10](https://github.com/astral-sh/ruff-pre-commit/compare/v0.15.9...v0.16.10) - [github.com/tox-dev/pyproject-fmt: v2.21.0 → v2.30.1](https://github.com/tox-dev/pyproject-fmt/compare/v2.21.0...v2.30.1) - [github.com/abravalheri/validate-pyproject: v0.25 → 0.26](https://github.com/abravalheri/validate-pyproject/compare/v0.25...0.26) --- .pre-commit-config.yaml | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 09bfb1d52a..fd60b48cac 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -38,7 +38,7 @@ repos: - id: detect-private-key - repo: https://github.com/codespell-project/codespell - rev: v2.4.2 + rev: v2.4.3 hooks: - id: codespell additional_dependencies: [tomli] @@ -71,7 +71,7 @@ repos: args: ["--print-width=140"] - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.9 + rev: v0.16.10 hooks: - id: ruff args: ["--fix"] @@ -79,11 +79,11 @@ repos: - id: ruff - repo: https://github.com/tox-dev/pyproject-fmt - rev: v2.21.0 + rev: v2.30.1 hooks: - id: pyproject-fmt additional_dependencies: [tox] - repo: https://github.com/abravalheri/validate-pyproject - rev: v0.25 + rev: '0.26' hooks: - id: validate-pyproject From 831290da394b32c0dab8e7a0778a3bc19abff1ef Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 23:26:01 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .pre-commit-config.yaml | 2 +- README.md | 4 +- extensions/thunder/README.md | 922 +++++++++++++--------- pyproject.toml | 25 +- tutorials/convert_lit_models.md | 6 +- tutorials/deploy.md | 14 +- tutorials/developer-docs/adding-models.md | 66 +- tutorials/developer-docs/python-api.md | 10 +- tutorials/evaluation.md | 12 +- tutorials/python-api.md | 25 +- 10 files changed, 631 insertions(+), 455 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index fd60b48cac..51d0ca94d8 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -84,6 +84,6 @@ repos: - id: pyproject-fmt additional_dependencies: [tox] - repo: https://github.com/abravalheri/validate-pyproject - rev: '0.26' + rev: "0.26" hooks: - id: validate-pyproject diff --git a/README.md b/README.md index 724d3ca734..7210c001be 100644 --- a/README.md +++ b/README.md @@ -282,9 +282,9 @@ Test the server in a separate terminal and integrate the model API into your AI ```python # 3) Use the server (in a separate Python session) import requests, json + response = requests.post( - "http://127.0.0.1:8000/predict", - json={"prompt": "Fix typos in the following sentence: Example input"} + "http://127.0.0.1:8000/predict", json={"prompt": "Fix typos in the following sentence: Example input"} ) print(response.json()["output"]) ``` diff --git a/extensions/thunder/README.md b/extensions/thunder/README.md index 713cbaf2e7..9ab21b1aac 100644 --- a/extensions/thunder/README.md +++ b/extensions/thunder/README.md @@ -45,149 +45,254 @@ print(forward_trace) @torch.no_grad() @no_autocast() def augmented_forward_fn(*args): - # args: "Collection" - t0, t1, t2, t3, t4, t5, t6, t7, t8, t9, t10, t11, t12, t13, t14, t15, t16, t17, \ - t18, t19, = args - del args - t24 = torch.nn.functional.embedding(t0, t19, None, None, 2.0, False, False) # t24: "cuda:0 f32[2, 5, 4096]" - t20 = torch_slice_prim_impl(t1, [0, 0], [5, 128], [1, 1]) # t20: "cuda:0 f32[5, 128]" - t21 = torch_slice_prim_impl(t2, [0, 0], [5, 128], [1, 1]) # t21: "cuda:0 f32[5, 128]" - t200 = torch.unsqueeze(t11, 0) # t200: "cuda:0 f32[1, 4096]" - t201 = torch.unsqueeze(t200, 1) # t201: "cuda:0 f32[1, 1, 4096]" - del t200 - t33 = Tensor.expand(t201, (2, 5, 4096)) # t33: "cuda:0 f32[2, 5, 4096]" - del t201 - t229 = torch.unsqueeze(t13, 0) # t229: "cuda:0 f32[1, 4096]" - t230 = torch.unsqueeze(t229, 1) # t230: "cuda:0 f32[1, 1, 4096]" - del t229 - t84 = Tensor.expand(t230, (2, 5, 4096)) # t84: "cuda:0 f32[2, 5, 4096]" - del t230 - t232 = torch.unsqueeze(t12, 0) # t232: "cuda:0 f32[1, 4096]" - t233 = torch.unsqueeze(t232, 1) # t233: "cuda:0 f32[1, 1, 4096]" - del t232 - t104 = Tensor.expand(t233, (2, 5, 4096)) # t104: "cuda:0 f32[2, 5, 4096]" - del t233 - t253 = torch.unsqueeze(t14, 0) # t253: "cuda:0 f32[1, 4096]" - t254 = torch.unsqueeze(t253, 1) # t254: "cuda:0 f32[1, 1, 4096]" - del t253 - t155 = Tensor.expand(t254, (2, 5, 4096)) # t155: "cuda:0 f32[2, 5, 4096]" - del t254 - t256 = torch.unsqueeze(t10, 0) # t256: "cuda:0 f32[1, 4096]" - t257 = torch.unsqueeze(t256, 1) # t257: "cuda:0 f32[1, 1, 4096]" - del t256 - t175 = Tensor.expand(t257, (2, 5, 4096)) # t175: "cuda:0 f32[2, 5, 4096]" - del t257 - t221 = torch.unsqueeze(t20, 0) # t221: "cuda:0 f32[1, 5, 128]" - del t20 - t222 = torch.unsqueeze(t221, 1) # t222: "cuda:0 f32[1, 1, 5, 128]" - del t221 - t49 = Tensor.expand(t222, (2, 32, 5, 128)) # t49: "cuda:0 f32[2, 32, 5, 128]" - del t222 - t224 = torch.unsqueeze(t21, 0) # t224: "cuda:0 f32[1, 5, 128]" - del t21 - t225 = torch.unsqueeze(t224, 1) # t225: "cuda:0 f32[1, 1, 5, 128]" - del t224 - t51 = Tensor.expand(t225, (2, 32, 5, 128)) # t51: "cuda:0 f32[2, 32, 5, 128]" - del t225 - [t30, t34] = nvFusion0(t24, t33) - t35 = torch.nn.functional.linear(t34, t3, None) # t35: "cuda:0 f32[2, 5, 12288]" - t36 = torch.reshape(t35, (2, 5, 32, 3, 128)) # t36: "cuda:0 f32[2, 5, 32, 3, 128]" - del t35 - t37 = torch.permute(t36, (0, 2, 3, 1, 4)) # t37: "cuda:0 f32[2, 32, 3, 5, 128]" - del t36 - (t38, t39, t40) = torch.split(t37, (1, 1, 1), 2) - del t37 - t41 = torch.reshape(t38, (2, 32, 5, 128)) # t41: "cuda:0 f32[2, 32, 5, 128]" - del t38 - t42 = torch.reshape(t39, (2, 32, 5, 128)) # t42: "cuda:0 f32[2, 32, 5, 128]" - del t39 - t43 = torch.reshape(t40, (2, 32, 5, 128)) # t43: "cuda:0 f32[2, 32, 5, 128]" - del t40 - t44 = torch_slice_prim_impl(t41, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t44: "cuda:0 f32[2, 32, 5, 128]" - t54 = torch_slice_prim_impl(t42, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t54: "cuda:0 f32[2, 32, 5, 128]" - t64 = torch_slice_prim_impl(t41, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t64: "cuda:0 f32[2, 32, 5, 0]" - del t41 - t66 = torch_slice_prim_impl(t42, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t66: "cuda:0 f32[2, 32, 5, 0]" - del t42 - t46 = torch_slice_prim_impl(t44, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t46: "cuda:0 f32[2, 32, 5, 64]" - t45 = torch_slice_prim_impl(t44, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t45: "cuda:0 f32[2, 32, 5, 64]" - t55 = torch_slice_prim_impl(t54, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t55: "cuda:0 f32[2, 32, 5, 64]" - t56 = torch_slice_prim_impl(t54, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t56: "cuda:0 f32[2, 32, 5, 64]" - [t47, t57] = nvFusion1(t46, t56) - del t46, t56 - t48 = torch.cat((t47, t45), -1) # t48: "cuda:0 f32[2, 32, 5, 128]" - del t47, t45 - t58 = torch.cat((t57, t55), -1) # t58: "cuda:0 f32[2, 32, 5, 128]" - del t57, t55 - [t53, t63] = nvFusion2(t44, t48, t49, t51, t54, t58) - del t44, t48, t54, t58 - t65 = torch.cat((t53, t64), -1) # t65: "cuda:0 f32[2, 32, 5, 128]" - del t53, t64 - t67 = torch.cat((t63, t66), -1) # t67: "cuda:0 f32[2, 32, 5, 128]" - del t63, t66 - (t68, t69, t70, t71) = sdpaex_grad_forward_scaled_dot_product_efficient_attention(t65, t67, t43, None, 0.0, True, 0.08838834764831843) - t72 = torch.permute(t68, (0, 2, 1, 3)) # t72: "cuda:0 f32[2, 5, 32, 128]" - t73 = torch.reshape(t72, (2, 5, 4096)) # t73: "cuda:0 f32[2, 5, 4096]" - del t72 - t74 = torch.nn.functional.linear(t73, t15, None) # t74: "cuda:0 f32[2, 5, 4096]" - [t75, t81, t85] = nvFusion3(t24, t74, t84) - del t74 - t86 = torch.nn.functional.linear(t85, t5, None) # t86: "cuda:0 f32[2, 5, 11008]" - t87 = torch.nn.functional.linear(t85, t7, None) # t87: "cuda:0 f32[2, 5, 11008]" - [t93] = nvFusion4(t86, t87) - t94 = torch.nn.functional.linear(t93, t16, None) # t94: "cuda:0 f32[2, 5, 4096]" - [t101, t105, t95] = nvFusion5(t104, t75, t94) - del t94 - t106 = torch.nn.functional.linear(t105, t4, None) # t106: "cuda:0 f32[2, 5, 12288]" - t107 = torch.reshape(t106, (2, 5, 32, 3, 128)) # t107: "cuda:0 f32[2, 5, 32, 3, 128]" - del t106 - t108 = torch.permute(t107, (0, 2, 3, 1, 4)) # t108: "cuda:0 f32[2, 32, 3, 5, 128]" - del t107 - (t109, t110, t111) = torch.split(t108, (1, 1, 1), 2) - del t108 - t112 = torch.reshape(t109, (2, 32, 5, 128)) # t112: "cuda:0 f32[2, 32, 5, 128]" - del t109 - t113 = torch.reshape(t110, (2, 32, 5, 128)) # t113: "cuda:0 f32[2, 32, 5, 128]" - del t110 - t114 = torch.reshape(t111, (2, 32, 5, 128)) # t114: "cuda:0 f32[2, 32, 5, 128]" - del t111 - t135 = torch_slice_prim_impl(t112, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t135: "cuda:0 f32[2, 32, 5, 0]" - t137 = torch_slice_prim_impl(t113, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t137: "cuda:0 f32[2, 32, 5, 0]" - t115 = torch_slice_prim_impl(t112, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t115: "cuda:0 f32[2, 32, 5, 128]" - del t112 - t125 = torch_slice_prim_impl(t113, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t125: "cuda:0 f32[2, 32, 5, 128]" - del t113 - t116 = torch_slice_prim_impl(t115, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t116: "cuda:0 f32[2, 32, 5, 64]" - t117 = torch_slice_prim_impl(t115, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t117: "cuda:0 f32[2, 32, 5, 64]" - t127 = torch_slice_prim_impl(t125, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t127: "cuda:0 f32[2, 32, 5, 64]" - t126 = torch_slice_prim_impl(t125, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t126: "cuda:0 f32[2, 32, 5, 64]" - [t118, t128] = nvFusion6(t117, t127) - del t117, t127 - t129 = torch.cat((t128, t126), -1) # t129: "cuda:0 f32[2, 32, 5, 128]" - del t128, t126 - t119 = torch.cat((t118, t116), -1) # t119: "cuda:0 f32[2, 32, 5, 128]" - del t118, t116 - [t124, t134] = nvFusion7(t115, t119, t125, t129, t49, t51) - del t115, t119, t125, t129 - t136 = torch.cat((t124, t135), -1) # t136: "cuda:0 f32[2, 32, 5, 128]" - del t124, t135 - t138 = torch.cat((t134, t137), -1) # t138: "cuda:0 f32[2, 32, 5, 128]" - del t134, t137 - (t139, t140, t141, t142) = sdpaex_grad_forward_scaled_dot_product_efficient_attention(t136, t138, t114, None, 0.0, True, 0.08838834764831843) - t143 = torch.permute(t139, (0, 2, 1, 3)) # t143: "cuda:0 f32[2, 5, 32, 128]" - t144 = torch.reshape(t143, (2, 5, 4096)) # t144: "cuda:0 f32[2, 5, 4096]" - del t143 - t145 = torch.nn.functional.linear(t144, t17, None) # t145: "cuda:0 f32[2, 5, 4096]" - [t146, t152, t156] = nvFusion8(t145, t155, t95) - del t145 - t158 = torch.nn.functional.linear(t156, t8, None) # t158: "cuda:0 f32[2, 5, 11008]" - t157 = torch.nn.functional.linear(t156, t6, None) # t157: "cuda:0 f32[2, 5, 11008]" - [t164] = nvFusion9(t157, t158) - t165 = torch.nn.functional.linear(t164, t18, None) # t165: "cuda:0 f32[2, 5, 4096]" - [t166, t172, t176] = nvFusion10(t146, t165, t175) - del t165 - t177 = torch.nn.functional.linear(t176, t9, None) # t177: "cuda:0 f32[2, 5, 32000]" - return {'output': t177, 'flat_args': [t0, t1, t2, t3, t4, t5, t6, t7, t8, t9, t10, t11, t12, t13, t14, t15, t16, t17, t18, t19], 'flat_output': (t177,)}, ((t0, t101, t104, t105, t114, t136, t138, t139, t140, t141, t142, t144, t146, t15, t152, t155, t156, t157, t158, t16, t164, t166, t17, t172, t175, t176, t18, t24, t3, t30, t33, t34, t4, t43, t49, t5, t51, t6, t65, t67, t68, t69, t7, t70, t71, t73, t75, t8, t81, t84, t85, t86, t87, t9, t93, t95), (False, False, True, True, 4096.0, 4096.0, 0.0, 0.08838834764831843, 4096.0, 4096.0, 4096.0, 0.0, 0.08838834764831843, 32000, 2, 2)) + # args: "Collection" + ( + t0, + t1, + t2, + t3, + t4, + t5, + t6, + t7, + t8, + t9, + t10, + t11, + t12, + t13, + t14, + t15, + t16, + t17, + t18, + t19, + ) = args + del args + t24 = torch.nn.functional.embedding(t0, t19, None, None, 2.0, False, False) # t24: "cuda:0 f32[2, 5, 4096]" + t20 = torch_slice_prim_impl(t1, [0, 0], [5, 128], [1, 1]) # t20: "cuda:0 f32[5, 128]" + t21 = torch_slice_prim_impl(t2, [0, 0], [5, 128], [1, 1]) # t21: "cuda:0 f32[5, 128]" + t200 = torch.unsqueeze(t11, 0) # t200: "cuda:0 f32[1, 4096]" + t201 = torch.unsqueeze(t200, 1) # t201: "cuda:0 f32[1, 1, 4096]" + del t200 + t33 = Tensor.expand(t201, (2, 5, 4096)) # t33: "cuda:0 f32[2, 5, 4096]" + del t201 + t229 = torch.unsqueeze(t13, 0) # t229: "cuda:0 f32[1, 4096]" + t230 = torch.unsqueeze(t229, 1) # t230: "cuda:0 f32[1, 1, 4096]" + del t229 + t84 = Tensor.expand(t230, (2, 5, 4096)) # t84: "cuda:0 f32[2, 5, 4096]" + del t230 + t232 = torch.unsqueeze(t12, 0) # t232: "cuda:0 f32[1, 4096]" + t233 = torch.unsqueeze(t232, 1) # t233: "cuda:0 f32[1, 1, 4096]" + del t232 + t104 = Tensor.expand(t233, (2, 5, 4096)) # t104: "cuda:0 f32[2, 5, 4096]" + del t233 + t253 = torch.unsqueeze(t14, 0) # t253: "cuda:0 f32[1, 4096]" + t254 = torch.unsqueeze(t253, 1) # t254: "cuda:0 f32[1, 1, 4096]" + del t253 + t155 = Tensor.expand(t254, (2, 5, 4096)) # t155: "cuda:0 f32[2, 5, 4096]" + del t254 + t256 = torch.unsqueeze(t10, 0) # t256: "cuda:0 f32[1, 4096]" + t257 = torch.unsqueeze(t256, 1) # t257: "cuda:0 f32[1, 1, 4096]" + del t256 + t175 = Tensor.expand(t257, (2, 5, 4096)) # t175: "cuda:0 f32[2, 5, 4096]" + del t257 + t221 = torch.unsqueeze(t20, 0) # t221: "cuda:0 f32[1, 5, 128]" + del t20 + t222 = torch.unsqueeze(t221, 1) # t222: "cuda:0 f32[1, 1, 5, 128]" + del t221 + t49 = Tensor.expand(t222, (2, 32, 5, 128)) # t49: "cuda:0 f32[2, 32, 5, 128]" + del t222 + t224 = torch.unsqueeze(t21, 0) # t224: "cuda:0 f32[1, 5, 128]" + del t21 + t225 = torch.unsqueeze(t224, 1) # t225: "cuda:0 f32[1, 1, 5, 128]" + del t224 + t51 = Tensor.expand(t225, (2, 32, 5, 128)) # t51: "cuda:0 f32[2, 32, 5, 128]" + del t225 + [t30, t34] = nvFusion0(t24, t33) + t35 = torch.nn.functional.linear(t34, t3, None) # t35: "cuda:0 f32[2, 5, 12288]" + t36 = torch.reshape(t35, (2, 5, 32, 3, 128)) # t36: "cuda:0 f32[2, 5, 32, 3, 128]" + del t35 + t37 = torch.permute(t36, (0, 2, 3, 1, 4)) # t37: "cuda:0 f32[2, 32, 3, 5, 128]" + del t36 + (t38, t39, t40) = torch.split(t37, (1, 1, 1), 2) + del t37 + t41 = torch.reshape(t38, (2, 32, 5, 128)) # t41: "cuda:0 f32[2, 32, 5, 128]" + del t38 + t42 = torch.reshape(t39, (2, 32, 5, 128)) # t42: "cuda:0 f32[2, 32, 5, 128]" + del t39 + t43 = torch.reshape(t40, (2, 32, 5, 128)) # t43: "cuda:0 f32[2, 32, 5, 128]" + del t40 + t44 = torch_slice_prim_impl(t41, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t44: "cuda:0 f32[2, 32, 5, 128]" + t54 = torch_slice_prim_impl(t42, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t54: "cuda:0 f32[2, 32, 5, 128]" + t64 = torch_slice_prim_impl(t41, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t64: "cuda:0 f32[2, 32, 5, 0]" + del t41 + t66 = torch_slice_prim_impl(t42, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t66: "cuda:0 f32[2, 32, 5, 0]" + del t42 + t46 = torch_slice_prim_impl(t44, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t46: "cuda:0 f32[2, 32, 5, 64]" + t45 = torch_slice_prim_impl(t44, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t45: "cuda:0 f32[2, 32, 5, 64]" + t55 = torch_slice_prim_impl(t54, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t55: "cuda:0 f32[2, 32, 5, 64]" + t56 = torch_slice_prim_impl(t54, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t56: "cuda:0 f32[2, 32, 5, 64]" + [t47, t57] = nvFusion1(t46, t56) + del t46, t56 + t48 = torch.cat((t47, t45), -1) # t48: "cuda:0 f32[2, 32, 5, 128]" + del t47, t45 + t58 = torch.cat((t57, t55), -1) # t58: "cuda:0 f32[2, 32, 5, 128]" + del t57, t55 + [t53, t63] = nvFusion2(t44, t48, t49, t51, t54, t58) + del t44, t48, t54, t58 + t65 = torch.cat((t53, t64), -1) # t65: "cuda:0 f32[2, 32, 5, 128]" + del t53, t64 + t67 = torch.cat((t63, t66), -1) # t67: "cuda:0 f32[2, 32, 5, 128]" + del t63, t66 + (t68, t69, t70, t71) = sdpaex_grad_forward_scaled_dot_product_efficient_attention( + t65, t67, t43, None, 0.0, True, 0.08838834764831843 + ) + t72 = torch.permute(t68, (0, 2, 1, 3)) # t72: "cuda:0 f32[2, 5, 32, 128]" + t73 = torch.reshape(t72, (2, 5, 4096)) # t73: "cuda:0 f32[2, 5, 4096]" + del t72 + t74 = torch.nn.functional.linear(t73, t15, None) # t74: "cuda:0 f32[2, 5, 4096]" + [t75, t81, t85] = nvFusion3(t24, t74, t84) + del t74 + t86 = torch.nn.functional.linear(t85, t5, None) # t86: "cuda:0 f32[2, 5, 11008]" + t87 = torch.nn.functional.linear(t85, t7, None) # t87: "cuda:0 f32[2, 5, 11008]" + [t93] = nvFusion4(t86, t87) + t94 = torch.nn.functional.linear(t93, t16, None) # t94: "cuda:0 f32[2, 5, 4096]" + [t101, t105, t95] = nvFusion5(t104, t75, t94) + del t94 + t106 = torch.nn.functional.linear(t105, t4, None) # t106: "cuda:0 f32[2, 5, 12288]" + t107 = torch.reshape(t106, (2, 5, 32, 3, 128)) # t107: "cuda:0 f32[2, 5, 32, 3, 128]" + del t106 + t108 = torch.permute(t107, (0, 2, 3, 1, 4)) # t108: "cuda:0 f32[2, 32, 3, 5, 128]" + del t107 + (t109, t110, t111) = torch.split(t108, (1, 1, 1), 2) + del t108 + t112 = torch.reshape(t109, (2, 32, 5, 128)) # t112: "cuda:0 f32[2, 32, 5, 128]" + del t109 + t113 = torch.reshape(t110, (2, 32, 5, 128)) # t113: "cuda:0 f32[2, 32, 5, 128]" + del t110 + t114 = torch.reshape(t111, (2, 32, 5, 128)) # t114: "cuda:0 f32[2, 32, 5, 128]" + del t111 + t135 = torch_slice_prim_impl(t112, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t135: "cuda:0 f32[2, 32, 5, 0]" + t137 = torch_slice_prim_impl(t113, [0, 0, 0, 0], [2, 32, 5, 0], [1, 1, 1, 1]) # t137: "cuda:0 f32[2, 32, 5, 0]" + t115 = torch_slice_prim_impl(t112, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t115: "cuda:0 f32[2, 32, 5, 128]" + del t112 + t125 = torch_slice_prim_impl(t113, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t125: "cuda:0 f32[2, 32, 5, 128]" + del t113 + t116 = torch_slice_prim_impl(t115, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t116: "cuda:0 f32[2, 32, 5, 64]" + t117 = torch_slice_prim_impl(t115, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t117: "cuda:0 f32[2, 32, 5, 64]" + t127 = torch_slice_prim_impl(t125, [0, 0, 0, 64], [2, 32, 5, 128], [1, 1, 1, 1]) # t127: "cuda:0 f32[2, 32, 5, 64]" + t126 = torch_slice_prim_impl(t125, [0, 0, 0, 0], [2, 32, 5, 64], [1, 1, 1, 1]) # t126: "cuda:0 f32[2, 32, 5, 64]" + [t118, t128] = nvFusion6(t117, t127) + del t117, t127 + t129 = torch.cat((t128, t126), -1) # t129: "cuda:0 f32[2, 32, 5, 128]" + del t128, t126 + t119 = torch.cat((t118, t116), -1) # t119: "cuda:0 f32[2, 32, 5, 128]" + del t118, t116 + [t124, t134] = nvFusion7(t115, t119, t125, t129, t49, t51) + del t115, t119, t125, t129 + t136 = torch.cat((t124, t135), -1) # t136: "cuda:0 f32[2, 32, 5, 128]" + del t124, t135 + t138 = torch.cat((t134, t137), -1) # t138: "cuda:0 f32[2, 32, 5, 128]" + del t134, t137 + (t139, t140, t141, t142) = sdpaex_grad_forward_scaled_dot_product_efficient_attention( + t136, t138, t114, None, 0.0, True, 0.08838834764831843 + ) + t143 = torch.permute(t139, (0, 2, 1, 3)) # t143: "cuda:0 f32[2, 5, 32, 128]" + t144 = torch.reshape(t143, (2, 5, 4096)) # t144: "cuda:0 f32[2, 5, 4096]" + del t143 + t145 = torch.nn.functional.linear(t144, t17, None) # t145: "cuda:0 f32[2, 5, 4096]" + [t146, t152, t156] = nvFusion8(t145, t155, t95) + del t145 + t158 = torch.nn.functional.linear(t156, t8, None) # t158: "cuda:0 f32[2, 5, 11008]" + t157 = torch.nn.functional.linear(t156, t6, None) # t157: "cuda:0 f32[2, 5, 11008]" + [t164] = nvFusion9(t157, t158) + t165 = torch.nn.functional.linear(t164, t18, None) # t165: "cuda:0 f32[2, 5, 4096]" + [t166, t172, t176] = nvFusion10(t146, t165, t175) + del t165 + t177 = torch.nn.functional.linear(t176, t9, None) # t177: "cuda:0 f32[2, 5, 32000]" + return { + "output": t177, + "flat_args": [t0, t1, t2, t3, t4, t5, t6, t7, t8, t9, t10, t11, t12, t13, t14, t15, t16, t17, t18, t19], + "flat_output": (t177,), + }, ( + ( + t0, + t101, + t104, + t105, + t114, + t136, + t138, + t139, + t140, + t141, + t142, + t144, + t146, + t15, + t152, + t155, + t156, + t157, + t158, + t16, + t164, + t166, + t17, + t172, + t175, + t176, + t18, + t24, + t3, + t30, + t33, + t34, + t4, + t43, + t49, + t5, + t51, + t6, + t65, + t67, + t68, + t69, + t7, + t70, + t71, + t73, + t75, + t8, + t81, + t84, + t85, + t86, + t87, + t9, + t93, + t95, + ), + ( + False, + False, + True, + True, + 4096.0, + 4096.0, + 0.0, + 0.08838834764831843, + 4096.0, + 4096.0, + 4096.0, + 0.0, + 0.08838834764831843, + 32000, + 2, + 2, + ), + ) ``` This is a straight-lined version of `GPT.forward` that has been optimized. Since it's running on CUDA, the [NvFuser](https://github.com/NVIDIA/Fuser) executor has created regions (look for "nvFusion") that fuse multiple operators together. @@ -207,17 +312,17 @@ print(forward_trace) We can see as comments the primitives that compose the fusion regions. For instance, this is the region associated to [the `RMSNorm` implementation](https://github.com/Lightning-AI/litgpt/blob/9b6475dabf90c7acee506a026bd9fa86251835bf/litgpt/model.py#L409-L420) ```python - [t146, t152, t156] = nvFusion8(t145, t155, t95) - # t146 = prims.add(t145, t95) # t146: "cuda:0 f32[2, 5, 4096]" - # t147 = prims.mul(t146, t146) # t147: "cuda:0 f32[2, 5, 4096]" - # t148 = prims.sum(t147, (2,)) # t148: "cuda:0 f32[2, 5]" - # t149 = prims.broadcast_in_dim(t148, [2, 5, 1], [0, 1]) # t149: "cuda:0 f32[2, 5, 1]" - # t150 = prims.div(t149, 4096.0) # t150: "cuda:0 f32[2, 5, 1]" - # t151 = prims.add(t150, 1e-05) # t151: "cuda:0 f32[2, 5, 1]" - # t152 = prims.rsqrt(t151) # t152: "cuda:0 f32[2, 5, 1]" - # t153 = prims.broadcast_in_dim(t152, (2, 5, 4096), (0, 1, 2)) # t153: "cuda:0 f32[2, 5, 4096]" - # t154 = prims.mul(t146, t153) # t154: "cuda:0 f32[2, 5, 4096]" - # t156 = prims.mul(t154, t155) # t156: "cuda:0 f32[2, 5, 4096]" +[t146, t152, t156] = nvFusion8(t145, t155, t95) +# t146 = prims.add(t145, t95) # t146: "cuda:0 f32[2, 5, 4096]" +# t147 = prims.mul(t146, t146) # t147: "cuda:0 f32[2, 5, 4096]" +# t148 = prims.sum(t147, (2,)) # t148: "cuda:0 f32[2, 5]" +# t149 = prims.broadcast_in_dim(t148, [2, 5, 1], [0, 1]) # t149: "cuda:0 f32[2, 5, 1]" +# t150 = prims.div(t149, 4096.0) # t150: "cuda:0 f32[2, 5, 1]" +# t151 = prims.add(t150, 1e-05) # t151: "cuda:0 f32[2, 5, 1]" +# t152 = prims.rsqrt(t151) # t152: "cuda:0 f32[2, 5, 1]" +# t153 = prims.broadcast_in_dim(t152, (2, 5, 4096), (0, 1, 2)) # t153: "cuda:0 f32[2, 5, 4096]" +# t154 = prims.mul(t146, t153) # t154: "cuda:0 f32[2, 5, 4096]" +# t156 = prims.mul(t154, t155) # t156: "cuda:0 f32[2, 5, 4096]" ``` Similarly, we can visualize the backward trace: @@ -231,202 +336,300 @@ print(backward_trace) @torch.no_grad() @no_autocast() def backward_fn(saved_for_backward, cotangents): - # saved_for_backward: "Collection" - # cotangents: "Collection" - C0, C1, = saved_for_backward - clear_collection(saved_for_backward) - del saved_for_backward - t178, = cotangents - clear_collection(cotangents) - del cotangents - t0, t101, t104, t105, t114, t136, t138, t139, t140, t141, t142, t144, t146, \ - t15, t152, t155, t156, t157, t158, t16, t164, t166, t17, t172, t175, t176, t18, \ - t24, t3, t30, t33, t34, t4, t43, t49, t5, t51, t6, t65, t67, t68, t69, t7, t70, \ - t71, t73, t75, t8, t81, t84, t85, t86, t87, t9, t93, t95, = C0 - clear_collection(C0) - del C0 - b1, b2, b41, b91, f101, f106, f40, f42, f51, f56, f6, f90, f92, i0, i23, i73, \ - = C1 - clear_collection(C1) - del C1 - t639 = torch.reshape(t178, (-1, 32000)) # t639: "cuda:0 f32[10, 32000]" - del t178 - t643 = torch.permute(t639, (1, 0)) # t643: "cuda:0 f32[32000, 10]" - t644 = torch.reshape(t176, (-1, 4096)) # t644: "cuda:0 f32[10, 4096]" - del t176 - t669 = torch.reshape(t164, (-1, 11008)) # t669: "cuda:0 f32[10, 11008]" - del t164 - t686 = torch.reshape(t156, (-1, 4096)) # t686: "cuda:0 f32[10, 4096]" - del t156 - t720 = torch.reshape(t144, (-1, 4096)) # t720: "cuda:0 f32[10, 4096]" - del t144 - t776 = torch.reshape(t105, (-1, 4096)) # t776: "cuda:0 f32[10, 4096]" - del t105 - t802 = torch.reshape(t93, (-1, 11008)) # t802: "cuda:0 f32[10, 11008]" - del t93 - t819 = torch.reshape(t85, (-1, 4096)) # t819: "cuda:0 f32[10, 4096]" - del t85 - t853 = torch.reshape(t73, (-1, 4096)) # t853: "cuda:0 f32[10, 4096]" - del t73 - t911 = torch.reshape(t34, (-1, 4096)) # t911: "cuda:0 f32[10, 4096]" - del t34 - t640 = torch.matmul(t639, t9) # t640: "cuda:0 f32[10, 4096]" - del t639, t9 - t645 = torch.matmul(t643, t644) # t645: "cuda:0 f32[32000, 4096]" - del t643, t644 - t641 = torch.reshape(t640, (2, 5, 4096)) # t641: "cuda:0 f32[2, 5, 4096]" - del t640 - [t648, t663] = nvFusion0(f106, t166, t172, t175, t641) - del f106, t166, t172, t175, t641 - t664 = torch.reshape(t663, (-1, 4096)) # t664: "cuda:0 f32[10, 4096]" - t668 = torch.permute(t664, (1, 0)) # t668: "cuda:0 f32[4096, 10]" - t665 = torch.matmul(t664, t18) # t665: "cuda:0 f32[10, 11008]" - del t664, t18 - t670 = torch.matmul(t668, t669) # t670: "cuda:0 f32[4096, 11008]" - del t668, t669 - t666 = torch.reshape(t665, (2, 5, 11008)) # t666: "cuda:0 f32[2, 5, 11008]" - del t665 - [t672, t680] = nvFusion1(t157, t158, t666) - del t157, t158, t666 - t681 = torch.reshape(t672, (-1, 11008)) # t681: "cuda:0 f32[10, 11008]" - del t672 - t685 = torch.permute(t681, (1, 0)) # t685: "cuda:0 f32[11008, 10]" - t688 = torch.reshape(t680, (-1, 11008)) # t688: "cuda:0 f32[10, 11008]" - del t680 - t692 = torch.permute(t688, (1, 0)) # t692: "cuda:0 f32[11008, 10]" - t689 = torch.matmul(t688, t6) # t689: "cuda:0 f32[10, 4096]" - del t688, t6 - t682 = torch.matmul(t681, t8) # t682: "cuda:0 f32[10, 4096]" - del t681, t8 - t694 = torch.matmul(t692, t686) # t694: "cuda:0 f32[11008, 4096]" - del t692 - t687 = torch.matmul(t685, t686) # t687: "cuda:0 f32[11008, 4096]" - del t685, t686 - t683 = torch.reshape(t682, (2, 5, 4096)) # t683: "cuda:0 f32[2, 5, 4096]" - del t682 - t690 = torch.reshape(t689, (2, 5, 4096)) # t690: "cuda:0 f32[2, 5, 4096]" - del t689 - [t698, t714] = nvFusion2(f101, t146, t152, t155, t663, t683, t690) - del f101, t146, t152, t155, t663, t683, t690 - t715 = torch.reshape(t714, (-1, 4096)) # t715: "cuda:0 f32[10, 4096]" - t719 = torch.permute(t715, (1, 0)) # t719: "cuda:0 f32[4096, 10]" - t716 = torch.matmul(t715, t17) # t716: "cuda:0 f32[10, 4096]" - del t715, t17 - t721 = torch.matmul(t719, t720) # t721: "cuda:0 f32[4096, 4096]" - del t719, t720 - t717 = torch.reshape(t716, (2, 5, 4096)) # t717: "cuda:0 f32[2, 5, 4096]" - del t716 - t722 = torch.reshape(t717, (2, 5, 32, 128)) # t722: "cuda:0 f32[2, 5, 32, 128]" - del t717 - t723 = torch.permute(t722, (0, 2, 1, 3)) # t723: "cuda:0 f32[2, 32, 5, 128]" - del t722 - (t724, t725, t726, _) = sdpaex_scaled_dot_product_efficient_attention_backward(t723, t136, t138, t114, None, t139, t140, t141, t142, f90, b91, scale=f92) - del t723, t136, t138, t114, t139, t140, t141, t142, f90, b91, f92 - t765 = torch.reshape(t726, (2, 32, 1, 5, 128)) # t765: "cuda:0 f32[2, 32, 1, 5, 128]" - del t726 - t727 = torch_slice_prim_impl(t725, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t727: "cuda:0 f32[2, 32, 5, 128]" - del t725 - t730 = torch_slice_prim_impl(t724, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t730: "cuda:0 f32[2, 32, 5, 128]" - del t724 - [t747, t764] = nvFusion3(t49, t51, t727, t730) - del t727, t730 - t766 = torch.reshape(t747, (2, 32, 1, 5, 128)) # t766: "cuda:0 f32[2, 32, 1, 5, 128]" - del t747 - t767 = torch.reshape(t764, (2, 32, 1, 5, 128)) # t767: "cuda:0 f32[2, 32, 1, 5, 128]" - del t764 - t768 = torch.cat((t767, t766, t765), i73) # t768: "cuda:0 f32[2, 32, 3, 5, 128]" - del t767, t766, t765, i73 - t769 = torch.permute(t768, (0, 3, 1, 2, 4)) # t769: "cuda:0 f32[2, 5, 32, 3, 128]" - del t768 - t770 = torch.reshape(t769, (2, 5, 12288)) # t770: "cuda:0 f32[2, 5, 12288]" - del t769 - t771 = torch.reshape(t770, (-1, 12288)) # t771: "cuda:0 f32[10, 12288]" - del t770 - t775 = torch.permute(t771, (1, 0)) # t775: "cuda:0 f32[12288, 10]" - t777 = torch.matmul(t775, t776) # t777: "cuda:0 f32[12288, 4096]" - del t775, t776 - t772 = torch.matmul(t771, t4) # t772: "cuda:0 f32[10, 4096]" - del t771, t4 - t773 = torch.reshape(t772, (2, 5, 4096)) # t773: "cuda:0 f32[2, 5, 4096]" - del t772 - [t780, t796] = nvFusion4(f56, t101, t104, t714, t773, t95) - del f56, t101, t104, t714, t773, t95 - t797 = torch.reshape(t796, (-1, 4096)) # t797: "cuda:0 f32[10, 4096]" - t801 = torch.permute(t797, (1, 0)) # t801: "cuda:0 f32[4096, 10]" - t798 = torch.matmul(t797, t16) # t798: "cuda:0 f32[10, 11008]" - del t797, t16 - t803 = torch.matmul(t801, t802) # t803: "cuda:0 f32[4096, 11008]" - del t801, t802 - t799 = torch.reshape(t798, (2, 5, 11008)) # t799: "cuda:0 f32[2, 5, 11008]" - del t798 - [t805, t813] = nvFusion5(t799, t86, t87) - del t799, t86, t87 - t814 = torch.reshape(t805, (-1, 11008)) # t814: "cuda:0 f32[10, 11008]" - del t805 - t818 = torch.permute(t814, (1, 0)) # t818: "cuda:0 f32[11008, 10]" - t821 = torch.reshape(t813, (-1, 11008)) # t821: "cuda:0 f32[10, 11008]" - del t813 - t825 = torch.permute(t821, (1, 0)) # t825: "cuda:0 f32[11008, 10]" - t822 = torch.matmul(t821, t5) # t822: "cuda:0 f32[10, 4096]" - del t821, t5 - t815 = torch.matmul(t814, t7) # t815: "cuda:0 f32[10, 4096]" - del t814, t7 - t827 = torch.matmul(t825, t819) # t827: "cuda:0 f32[11008, 4096]" - del t825 - t820 = torch.matmul(t818, t819) # t820: "cuda:0 f32[11008, 4096]" - del t818, t819 - t816 = torch.reshape(t815, (2, 5, 4096)) # t816: "cuda:0 f32[2, 5, 4096]" - del t815 - t823 = torch.reshape(t822, (2, 5, 4096)) # t823: "cuda:0 f32[2, 5, 4096]" - del t822 - [t831, t847] = nvFusion6(f51, t75, t796, t81, t816, t823, t84) - del f51, t75, t796, t81, t816, t823, t84 - t848 = torch.reshape(t847, (-1, 4096)) # t848: "cuda:0 f32[10, 4096]" - t852 = torch.permute(t848, (1, 0)) # t852: "cuda:0 f32[4096, 10]" - t849 = torch.matmul(t848, t15) # t849: "cuda:0 f32[10, 4096]" - del t848, t15 - t854 = torch.matmul(t852, t853) # t854: "cuda:0 f32[4096, 4096]" - del t852, t853 - t850 = torch.reshape(t849, (2, 5, 4096)) # t850: "cuda:0 f32[2, 5, 4096]" - del t849 - t855 = torch.reshape(t850, (2, 5, 32, 128)) # t855: "cuda:0 f32[2, 5, 32, 128]" - del t850 - t856 = torch.permute(t855, (0, 2, 1, 3)) # t856: "cuda:0 f32[2, 32, 5, 128]" - del t855 - (t857, t858, t859, _) = sdpaex_scaled_dot_product_efficient_attention_backward(t856, t65, t67, t43, None, t68, t69, t70, t71, f40, b41, scale=f42) - del t856, t65, t67, t43, t68, t69, t70, t71, f40, b41, f42 - t900 = torch.reshape(t859, (2, 32, 1, 5, 128)) # t900: "cuda:0 f32[2, 32, 1, 5, 128]" - del t859 - t863 = torch_slice_prim_impl(t857, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t863: "cuda:0 f32[2, 32, 5, 128]" - del t857 - t860 = torch_slice_prim_impl(t858, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t860: "cuda:0 f32[2, 32, 5, 128]" - del t858 - [t882, t899] = nvFusion7(t49, t51, t860, t863) - del t49, t51, t860, t863 - t902 = torch.reshape(t899, (2, 32, 1, 5, 128)) # t902: "cuda:0 f32[2, 32, 1, 5, 128]" - del t899 - t901 = torch.reshape(t882, (2, 32, 1, 5, 128)) # t901: "cuda:0 f32[2, 32, 1, 5, 128]" - del t882 - t903 = torch.cat((t902, t901, t900), i23) # t903: "cuda:0 f32[2, 32, 3, 5, 128]" - del t902, t901, t900, i23 - t904 = torch.permute(t903, (0, 3, 1, 2, 4)) # t904: "cuda:0 f32[2, 5, 32, 3, 128]" - del t903 - t905 = torch.reshape(t904, (2, 5, 12288)) # t905: "cuda:0 f32[2, 5, 12288]" - del t904 - t906 = torch.reshape(t905, (-1, 12288)) # t906: "cuda:0 f32[10, 12288]" - del t905 - t910 = torch.permute(t906, (1, 0)) # t910: "cuda:0 f32[12288, 10]" - t907 = torch.matmul(t906, t3) # t907: "cuda:0 f32[10, 4096]" - del t906, t3 - t912 = torch.matmul(t910, t911) # t912: "cuda:0 f32[12288, 4096]" - del t910, t911 - t908 = torch.reshape(t907, (2, 5, 4096)) # t908: "cuda:0 f32[2, 5, 4096]" - del t907 - [t915, t931] = nvFusion8(f6, t24, t30, t33, t847, t908) - del f6, t24, t30, t33, t847, t908 - t932 = torch.torch.ops.aten.embedding_backward(t931, t0, i0, -1, b1, b2) # t932: "cuda:0 f32[32000, 4096]" - del t931, t0, i0, b1, b2 - return (None, None, None, t912, t777, t827, t694, t820, t687, t645, t648, t915, t780, t831, t698, t854, t803, t721, t670, t932) + # saved_for_backward: "Collection" + # cotangents: "Collection" + ( + C0, + C1, + ) = saved_for_backward + clear_collection(saved_for_backward) + del saved_for_backward + (t178,) = cotangents + clear_collection(cotangents) + del cotangents + ( + t0, + t101, + t104, + t105, + t114, + t136, + t138, + t139, + t140, + t141, + t142, + t144, + t146, + t15, + t152, + t155, + t156, + t157, + t158, + t16, + t164, + t166, + t17, + t172, + t175, + t176, + t18, + t24, + t3, + t30, + t33, + t34, + t4, + t43, + t49, + t5, + t51, + t6, + t65, + t67, + t68, + t69, + t7, + t70, + t71, + t73, + t75, + t8, + t81, + t84, + t85, + t86, + t87, + t9, + t93, + t95, + ) = C0 + clear_collection(C0) + del C0 + ( + b1, + b2, + b41, + b91, + f101, + f106, + f40, + f42, + f51, + f56, + f6, + f90, + f92, + i0, + i23, + i73, + ) = C1 + clear_collection(C1) + del C1 + t639 = torch.reshape(t178, (-1, 32000)) # t639: "cuda:0 f32[10, 32000]" + del t178 + t643 = torch.permute(t639, (1, 0)) # t643: "cuda:0 f32[32000, 10]" + t644 = torch.reshape(t176, (-1, 4096)) # t644: "cuda:0 f32[10, 4096]" + del t176 + t669 = torch.reshape(t164, (-1, 11008)) # t669: "cuda:0 f32[10, 11008]" + del t164 + t686 = torch.reshape(t156, (-1, 4096)) # t686: "cuda:0 f32[10, 4096]" + del t156 + t720 = torch.reshape(t144, (-1, 4096)) # t720: "cuda:0 f32[10, 4096]" + del t144 + t776 = torch.reshape(t105, (-1, 4096)) # t776: "cuda:0 f32[10, 4096]" + del t105 + t802 = torch.reshape(t93, (-1, 11008)) # t802: "cuda:0 f32[10, 11008]" + del t93 + t819 = torch.reshape(t85, (-1, 4096)) # t819: "cuda:0 f32[10, 4096]" + del t85 + t853 = torch.reshape(t73, (-1, 4096)) # t853: "cuda:0 f32[10, 4096]" + del t73 + t911 = torch.reshape(t34, (-1, 4096)) # t911: "cuda:0 f32[10, 4096]" + del t34 + t640 = torch.matmul(t639, t9) # t640: "cuda:0 f32[10, 4096]" + del t639, t9 + t645 = torch.matmul(t643, t644) # t645: "cuda:0 f32[32000, 4096]" + del t643, t644 + t641 = torch.reshape(t640, (2, 5, 4096)) # t641: "cuda:0 f32[2, 5, 4096]" + del t640 + [t648, t663] = nvFusion0(f106, t166, t172, t175, t641) + del f106, t166, t172, t175, t641 + t664 = torch.reshape(t663, (-1, 4096)) # t664: "cuda:0 f32[10, 4096]" + t668 = torch.permute(t664, (1, 0)) # t668: "cuda:0 f32[4096, 10]" + t665 = torch.matmul(t664, t18) # t665: "cuda:0 f32[10, 11008]" + del t664, t18 + t670 = torch.matmul(t668, t669) # t670: "cuda:0 f32[4096, 11008]" + del t668, t669 + t666 = torch.reshape(t665, (2, 5, 11008)) # t666: "cuda:0 f32[2, 5, 11008]" + del t665 + [t672, t680] = nvFusion1(t157, t158, t666) + del t157, t158, t666 + t681 = torch.reshape(t672, (-1, 11008)) # t681: "cuda:0 f32[10, 11008]" + del t672 + t685 = torch.permute(t681, (1, 0)) # t685: "cuda:0 f32[11008, 10]" + t688 = torch.reshape(t680, (-1, 11008)) # t688: "cuda:0 f32[10, 11008]" + del t680 + t692 = torch.permute(t688, (1, 0)) # t692: "cuda:0 f32[11008, 10]" + t689 = torch.matmul(t688, t6) # t689: "cuda:0 f32[10, 4096]" + del t688, t6 + t682 = torch.matmul(t681, t8) # t682: "cuda:0 f32[10, 4096]" + del t681, t8 + t694 = torch.matmul(t692, t686) # t694: "cuda:0 f32[11008, 4096]" + del t692 + t687 = torch.matmul(t685, t686) # t687: "cuda:0 f32[11008, 4096]" + del t685, t686 + t683 = torch.reshape(t682, (2, 5, 4096)) # t683: "cuda:0 f32[2, 5, 4096]" + del t682 + t690 = torch.reshape(t689, (2, 5, 4096)) # t690: "cuda:0 f32[2, 5, 4096]" + del t689 + [t698, t714] = nvFusion2(f101, t146, t152, t155, t663, t683, t690) + del f101, t146, t152, t155, t663, t683, t690 + t715 = torch.reshape(t714, (-1, 4096)) # t715: "cuda:0 f32[10, 4096]" + t719 = torch.permute(t715, (1, 0)) # t719: "cuda:0 f32[4096, 10]" + t716 = torch.matmul(t715, t17) # t716: "cuda:0 f32[10, 4096]" + del t715, t17 + t721 = torch.matmul(t719, t720) # t721: "cuda:0 f32[4096, 4096]" + del t719, t720 + t717 = torch.reshape(t716, (2, 5, 4096)) # t717: "cuda:0 f32[2, 5, 4096]" + del t716 + t722 = torch.reshape(t717, (2, 5, 32, 128)) # t722: "cuda:0 f32[2, 5, 32, 128]" + del t717 + t723 = torch.permute(t722, (0, 2, 1, 3)) # t723: "cuda:0 f32[2, 32, 5, 128]" + del t722 + (t724, t725, t726, _) = sdpaex_scaled_dot_product_efficient_attention_backward( + t723, t136, t138, t114, None, t139, t140, t141, t142, f90, b91, scale=f92 + ) + del t723, t136, t138, t114, t139, t140, t141, t142, f90, b91, f92 + t765 = torch.reshape(t726, (2, 32, 1, 5, 128)) # t765: "cuda:0 f32[2, 32, 1, 5, 128]" + del t726 + t727 = torch_slice_prim_impl(t725, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t727: "cuda:0 f32[2, 32, 5, 128]" + del t725 + t730 = torch_slice_prim_impl(t724, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t730: "cuda:0 f32[2, 32, 5, 128]" + del t724 + [t747, t764] = nvFusion3(t49, t51, t727, t730) + del t727, t730 + t766 = torch.reshape(t747, (2, 32, 1, 5, 128)) # t766: "cuda:0 f32[2, 32, 1, 5, 128]" + del t747 + t767 = torch.reshape(t764, (2, 32, 1, 5, 128)) # t767: "cuda:0 f32[2, 32, 1, 5, 128]" + del t764 + t768 = torch.cat((t767, t766, t765), i73) # t768: "cuda:0 f32[2, 32, 3, 5, 128]" + del t767, t766, t765, i73 + t769 = torch.permute(t768, (0, 3, 1, 2, 4)) # t769: "cuda:0 f32[2, 5, 32, 3, 128]" + del t768 + t770 = torch.reshape(t769, (2, 5, 12288)) # t770: "cuda:0 f32[2, 5, 12288]" + del t769 + t771 = torch.reshape(t770, (-1, 12288)) # t771: "cuda:0 f32[10, 12288]" + del t770 + t775 = torch.permute(t771, (1, 0)) # t775: "cuda:0 f32[12288, 10]" + t777 = torch.matmul(t775, t776) # t777: "cuda:0 f32[12288, 4096]" + del t775, t776 + t772 = torch.matmul(t771, t4) # t772: "cuda:0 f32[10, 4096]" + del t771, t4 + t773 = torch.reshape(t772, (2, 5, 4096)) # t773: "cuda:0 f32[2, 5, 4096]" + del t772 + [t780, t796] = nvFusion4(f56, t101, t104, t714, t773, t95) + del f56, t101, t104, t714, t773, t95 + t797 = torch.reshape(t796, (-1, 4096)) # t797: "cuda:0 f32[10, 4096]" + t801 = torch.permute(t797, (1, 0)) # t801: "cuda:0 f32[4096, 10]" + t798 = torch.matmul(t797, t16) # t798: "cuda:0 f32[10, 11008]" + del t797, t16 + t803 = torch.matmul(t801, t802) # t803: "cuda:0 f32[4096, 11008]" + del t801, t802 + t799 = torch.reshape(t798, (2, 5, 11008)) # t799: "cuda:0 f32[2, 5, 11008]" + del t798 + [t805, t813] = nvFusion5(t799, t86, t87) + del t799, t86, t87 + t814 = torch.reshape(t805, (-1, 11008)) # t814: "cuda:0 f32[10, 11008]" + del t805 + t818 = torch.permute(t814, (1, 0)) # t818: "cuda:0 f32[11008, 10]" + t821 = torch.reshape(t813, (-1, 11008)) # t821: "cuda:0 f32[10, 11008]" + del t813 + t825 = torch.permute(t821, (1, 0)) # t825: "cuda:0 f32[11008, 10]" + t822 = torch.matmul(t821, t5) # t822: "cuda:0 f32[10, 4096]" + del t821, t5 + t815 = torch.matmul(t814, t7) # t815: "cuda:0 f32[10, 4096]" + del t814, t7 + t827 = torch.matmul(t825, t819) # t827: "cuda:0 f32[11008, 4096]" + del t825 + t820 = torch.matmul(t818, t819) # t820: "cuda:0 f32[11008, 4096]" + del t818, t819 + t816 = torch.reshape(t815, (2, 5, 4096)) # t816: "cuda:0 f32[2, 5, 4096]" + del t815 + t823 = torch.reshape(t822, (2, 5, 4096)) # t823: "cuda:0 f32[2, 5, 4096]" + del t822 + [t831, t847] = nvFusion6(f51, t75, t796, t81, t816, t823, t84) + del f51, t75, t796, t81, t816, t823, t84 + t848 = torch.reshape(t847, (-1, 4096)) # t848: "cuda:0 f32[10, 4096]" + t852 = torch.permute(t848, (1, 0)) # t852: "cuda:0 f32[4096, 10]" + t849 = torch.matmul(t848, t15) # t849: "cuda:0 f32[10, 4096]" + del t848, t15 + t854 = torch.matmul(t852, t853) # t854: "cuda:0 f32[4096, 4096]" + del t852, t853 + t850 = torch.reshape(t849, (2, 5, 4096)) # t850: "cuda:0 f32[2, 5, 4096]" + del t849 + t855 = torch.reshape(t850, (2, 5, 32, 128)) # t855: "cuda:0 f32[2, 5, 32, 128]" + del t850 + t856 = torch.permute(t855, (0, 2, 1, 3)) # t856: "cuda:0 f32[2, 32, 5, 128]" + del t855 + (t857, t858, t859, _) = sdpaex_scaled_dot_product_efficient_attention_backward( + t856, t65, t67, t43, None, t68, t69, t70, t71, f40, b41, scale=f42 + ) + del t856, t65, t67, t43, t68, t69, t70, t71, f40, b41, f42 + t900 = torch.reshape(t859, (2, 32, 1, 5, 128)) # t900: "cuda:0 f32[2, 32, 1, 5, 128]" + del t859 + t863 = torch_slice_prim_impl(t857, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t863: "cuda:0 f32[2, 32, 5, 128]" + del t857 + t860 = torch_slice_prim_impl(t858, [0, 0, 0, 0], [2, 32, 5, 128], [1, 1, 1, 1]) # t860: "cuda:0 f32[2, 32, 5, 128]" + del t858 + [t882, t899] = nvFusion7(t49, t51, t860, t863) + del t49, t51, t860, t863 + t902 = torch.reshape(t899, (2, 32, 1, 5, 128)) # t902: "cuda:0 f32[2, 32, 1, 5, 128]" + del t899 + t901 = torch.reshape(t882, (2, 32, 1, 5, 128)) # t901: "cuda:0 f32[2, 32, 1, 5, 128]" + del t882 + t903 = torch.cat((t902, t901, t900), i23) # t903: "cuda:0 f32[2, 32, 3, 5, 128]" + del t902, t901, t900, i23 + t904 = torch.permute(t903, (0, 3, 1, 2, 4)) # t904: "cuda:0 f32[2, 5, 32, 3, 128]" + del t903 + t905 = torch.reshape(t904, (2, 5, 12288)) # t905: "cuda:0 f32[2, 5, 12288]" + del t904 + t906 = torch.reshape(t905, (-1, 12288)) # t906: "cuda:0 f32[10, 12288]" + del t905 + t910 = torch.permute(t906, (1, 0)) # t910: "cuda:0 f32[12288, 10]" + t907 = torch.matmul(t906, t3) # t907: "cuda:0 f32[10, 4096]" + del t906, t3 + t912 = torch.matmul(t910, t911) # t912: "cuda:0 f32[12288, 4096]" + del t910, t911 + t908 = torch.reshape(t907, (2, 5, 4096)) # t908: "cuda:0 f32[2, 5, 4096]" + del t907 + [t915, t931] = nvFusion8(f6, t24, t30, t33, t847, t908) + del f6, t24, t30, t33, t847, t908 + t932 = torch.torch.ops.aten.embedding_backward(t931, t0, i0, -1, b1, b2) # t932: "cuda:0 f32[32000, 4096]" + del t931, t0, i0, b1, b2 + return ( + None, + None, + None, + t912, + t777, + t827, + t694, + t820, + t687, + t645, + t648, + t915, + t780, + t831, + t698, + t854, + t803, + t721, + t670, + t932, + ) ``` These traces are long, and require some familiarity with the model implementation to follow them, but they allow you to: @@ -452,9 +655,11 @@ model = thunder.jit(model) After applying the DDP transformation, the backward trace will include the expected all-reduce collectives: ```python - p1022 = torch_all_reduce_prim_impl(t1021, _DistributedReduceOps_0, _torch_distributed_distributed_c10d_ProcessGroup_1, True, False) # p1022: "FUTURE cuda:0 f32[16797696]" - ... - t1059 = torch_wait_prim_impl(p1025) # t1059: "cuda:0 f32[131072000]" +p1022 = torch_all_reduce_prim_impl( + t1021, _DistributedReduceOps_0, _torch_distributed_distributed_c10d_ProcessGroup_1, True, False +) # p1022: "FUTURE cuda:0 f32[16797696]" +... +t1059 = torch_wait_prim_impl(p1025) # t1059: "cuda:0 f32[131072000]" ``` With `L.Fabric`, this is how to use them: @@ -488,10 +693,7 @@ Thunder allows you to define a priority list of executors that can map operators ```python import thunder -model = thunder.jit( - model, - executors=["sdpa", "torchcompile_cat", "nvfuser", "torch"] -) +model = thunder.jit(model, executors=["sdpa", "torchcompile_cat", "nvfuser", "torch"]) ``` Notice how `torch.compile` is a valid executor. This executor registers a few operators with improved performance so that you can utilize the fastest set of operator implementations possible. @@ -512,10 +714,7 @@ We can enable this executor by passing it to the list of executors available. Th ```python import thunder -model = thunder.jit( - model, - executors=["sdpa", "unsloth", "torchcompile_cat", "nvfuser", "torch"] -) +model = thunder.jit(model, executors=["sdpa", "unsloth", "torchcompile_cat", "nvfuser", "torch"]) ``` Doing this, the model trace now includes the Unsloth kernel calls: @@ -528,6 +727,7 @@ def augmented_forward_fn(*args): (t189, t190) = unsloth_cross_entropy(t187, t188) ... + def backward_fn(saved_for_backward, cotangents): ... t652 = unsloth_cross_entropy_backward(t651, t187, t188, t190) # t652: "cuda:0 f32[6, 320]" diff --git a/pyproject.toml b/pyproject.toml index 98166717bc..00058a6be4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,6 +22,7 @@ classifiers = [ "Programming Language :: Python :: 3.12", "Programming Language :: Python :: 3.13", "Programming Language :: Python :: 3.14", + "Programming Language :: Python :: 3.15", ] dependencies = [ # download models: @@ -89,15 +90,15 @@ urls.homepage = "https://github.com/lightning-AI/litgpt" scripts.litgpt = "litgpt.__main__:main" [tool.setuptools] -package-data.litgpt = [ - "LICENSE.md", - "README.md", -] packages.find.include = [ "litgpt", "litgpt.*", ] packages.find.exclude = [] +package-data.litgpt = [ + "LICENSE.md", + "README.md", +] [tool.ruff] target-version = "py310" @@ -114,12 +115,6 @@ lint.select = [ "UP", # see: https://docs.astral.sh/ruff/rules/#pyupgrade-up "W", # see: https://pypi.org/project/pycodestyle ] -# extend-select = [ -# "C4", # see: https://pypi.org/project/flake8-comprehensions -# "PT", # see: https://pypi.org/project/flake8-pytest-style -# "RET", # see: https://pypi.org/project/flake8-return -# "SIM", # see: https://pypi.org/project/flake8-simplify -# ] lint.ignore = [ "E501", # Line too long "E731", # Do not assign a lambda expression, use a def @@ -128,14 +123,20 @@ lint.ignore = [ ] # Use Google-style docstrings. lint.pydocstyle.convention = "google" +# extend-select = [ +# "C4", # see: https://pypi.org/project/flake8-comprehensions +# "PT", # see: https://pypi.org/project/flake8-pytest-style +# "RET", # see: https://pypi.org/project/flake8-return +# "SIM", # see: https://pypi.org/project/flake8-simplify +# ] [tool.codespell] -# skip = '*.py' -quiet-level = 3 ignore-words-list = """ tral, \ Rockerfeller """ +# skip = "*.py" +quiet-level = 3 [tool.pytest] ini_options.addopts = [ diff --git a/tutorials/convert_lit_models.md b/tutorials/convert_lit_models.md index 53b24c2bdb..4322a53024 100644 --- a/tutorials/convert_lit_models.md +++ b/tutorials/convert_lit_models.md @@ -27,9 +27,7 @@ from transformers import AutoModel state_dict = torch.load("output_dir/model.pth") -model = AutoModel.from_pretrained( - "output_dir/", local_files_only=True, state_dict=state_dict -) +model = AutoModel.from_pretrained("output_dir/", local_files_only=True, state_dict=state_dict) ``` Alternatively, you can also load the model without copying the `config.json` file as follows: @@ -107,7 +105,7 @@ litgpt convert_from_litgpt $finetuned_dir/final/ out/hf-tinyllama/converted import torch from transformers import AutoModel -state_dict = torch.load('out/hf-tinyllama/converted/model.pth') +state_dict = torch.load("out/hf-tinyllama/converted/model.pth") model = AutoModel.from_pretrained("TinyLlama/TinyLlama-1.1B-intermediate-step-1431k-3T", state_dict=state_dict) ``` diff --git a/tutorials/deploy.md b/tutorials/deploy.md index b58a08ac03..01e496bba6 100644 --- a/tutorials/deploy.md +++ b/tutorials/deploy.md @@ -35,8 +35,7 @@ You can now send requests to the inference server you started in step 2. For exa import requests, json response = requests.post( - "http://127.0.0.1:8000/predict", - json={"prompt": "Fix typos in the following sentence: Example input"} + "http://127.0.0.1:8000/predict", json={"prompt": "Fix typos in the following sentence: Example input"} ) print(response.json()["output"]) @@ -63,9 +62,7 @@ Then, use the following updated code to query the inference server: import requests, json response = requests.post( - "http://127.0.0.1:8000/predict", - json={"prompt": "Fix typos in the following sentence: Example input"}, - stream=True + "http://127.0.0.1:8000/predict", json={"prompt": "Fix typos in the following sentence: Example input"}, stream=True ) # stream the response @@ -121,14 +118,11 @@ from openai import OpenAI # Configure the client to use your local LitGPT server client = OpenAI( base_url="http://127.0.0.1:8000/v1", - api_key="not-needed" # LitGPT doesn't require authentication by default + api_key="not-needed", # LitGPT doesn't require authentication by default ) response = client.chat.completions.create( - model="SmolLM2-135M-Instruct", - messages=[ - {"role": "user", "content": "Hello! How are you?"} - ] + model="SmolLM2-135M-Instruct", messages=[{"role": "user", "content": "Hello! How are you?"}] ) print(response.choices[0].message.content) diff --git a/tutorials/developer-docs/adding-models.md b/tutorials/developer-docs/adding-models.md index 0cd128ce14..578f6c0e95 100644 --- a/tutorials/developer-docs/adding-models.md +++ b/tutorials/developer-docs/adding-models.md @@ -41,23 +41,25 @@ For example, suppose an entry for Llama 3 8B already exists and you want to add Copy the Llama 3 8B entry: ```python - # https://huggingface.co/meta-llama/Meta-Llama-3-8B/blob/main/config.json - dict( - name="Llama-3-8B{}", - hf_config=dict(org="meta-llama", name="Meta-Llama-3-8B{}"), - vocab_size=128256, - padding_multiple=64, - n_layer=32, - n_head=32, - n_query_groups=8, - rotary_percentage=1.0, - parallel_residual=False, - bias=False, - norm_class_name="RMSNorm", - mlp_class_name="LLaMAMLP", - intermediate_size=14336, - rope_base=500000, - ), +# https://huggingface.co/meta-llama/Meta-Llama-3-8B/blob/main/config.json +( + dict( + name="Llama-3-8B{}", + hf_config=dict(org="meta-llama", name="Meta-Llama-3-8B{}"), + vocab_size=128256, + padding_multiple=64, + n_layer=32, + n_head=32, + n_query_groups=8, + rotary_percentage=1.0, + parallel_residual=False, + bias=False, + norm_class_name="RMSNorm", + mlp_class_name="LLaMAMLP", + intermediate_size=14336, + rope_base=500000, + ), +) ``` Then create the entry for the 70B model. Here, make sure you update the values according to the `config.json` file available on the HF hub: @@ -130,21 +132,21 @@ If you are adding a new model class, find out its prompt style. First, check [li ```python class Llama3(PromptStyle): - def apply(self, prompt: str, **kwargs: str) -> str: - # https://github.com/meta-llama/llama3/blob/359887376f0aaf30e433f23e25df858d8c2a9833/llama/tokenizer.py#L202-L229 - return ( - "<|begin_of_text|><|start_header_id|>system<|end_header_id|>\n\n" - "You are a helpful assistant.<|eot_id|>\n" # The system prompt is optional - "<|start_header_id|>user<|end_header_id|>\n\n" - f"{prompt}<|eot_id|>\n" - "<|start_header_id|>assistant<|end_header_id|>\n\n" - ) - - def stop_tokens(self, tokenizer: "Tokenizer") -> Tuple[List[int], ...]: - return ( - [tokenizer.eos_id], - [tokenizer.token_to_id("<|eot_id|>")], - ) + def apply(self, prompt: str, **kwargs: str) -> str: + # https://github.com/meta-llama/llama3/blob/359887376f0aaf30e433f23e25df858d8c2a9833/llama/tokenizer.py#L202-L229 + return ( + "<|begin_of_text|><|start_header_id|>system<|end_header_id|>\n\n" + "You are a helpful assistant.<|eot_id|>\n" # The system prompt is optional + "<|start_header_id|>user<|end_header_id|>\n\n" + f"{prompt}<|eot_id|>\n" + "<|start_header_id|>assistant<|end_header_id|>\n\n" + ) + + def stop_tokens(self, tokenizer: "Tokenizer") -> Tuple[List[int], ...]: + return ( + [tokenizer.eos_id], + [tokenizer.token_to_id("<|eot_id|>")], + ) ``` If your model requires a different prompt template, create a new `PromptStyle` class. diff --git a/tutorials/developer-docs/python-api.md b/tutorials/developer-docs/python-api.md index df98efaf2d..07fddcafcc 100644 --- a/tutorials/developer-docs/python-api.md +++ b/tutorials/developer-docs/python-api.md @@ -74,12 +74,7 @@ dataset = llm.prepare_dataset( ```python -llm.instruction_finetune( - config=None, - dataset=dataset, - max_iter=10, - method="full | lora | adapter | adapter_v2" -) +llm.instruction_finetune(config=None, dataset=dataset, max_iter=10, method="full | lora | adapter | adapter_v2") ``` ```python @@ -100,8 +95,7 @@ Then in another Python session: import requests, json response = requests.post( - "http://127.0.0.1:8000/predict", - json={"prompt": "Fix typos in the following sentence: Example input"} + "http://127.0.0.1:8000/predict", json={"prompt": "Fix typos in the following sentence: Example input"} ) print(response.json()["output"]) diff --git a/tutorials/evaluation.md b/tutorials/evaluation.md index d0e2b876b0..0b164e8f73 100644 --- a/tutorials/evaluation.md +++ b/tutorials/evaluation.md @@ -112,15 +112,11 @@ Suppose you have a test dataset with the following structure: ```python test_data = [ - { - "instruction": "Name the author of 'Pride and Prejudice'.", - "input": "", - "output": "Jane Austen." - }, + {"instruction": "Name the author of 'Pride and Prejudice'.", "input": "", "output": "Jane Austen."}, { "instruction": "Pick out the adjective from the following list.", "input": "run, tall, quickly", - "output": "The correct adjective from the list is 'tall.'" + "output": "The correct adjective from the list is 'tall.'", }, ] ``` @@ -173,7 +169,7 @@ Next, we use a second LLM to calculate the response quality on a scale from 0 to ```python -del llm # delete previous `llm` to free up GPU memory +del llm # delete previous `llm` to free up GPU memory scorer = LLM.load("meta-llama/Meta-Llama-3-8B-Instruct", access_token="...") ``` @@ -207,7 +203,7 @@ def generate_model_scores(data_dict, model, response_field="response", target_fi scores = generate_model_scores(test_data, model=scorer) print(f"\n{llm}") print(f"Number of scores: {len(scores)} of {len(test_data)}") -print(f"Average score: {sum(scores)/len(scores):.2f}\n") +print(f"Average score: {sum(scores) / len(scores):.2f}\n") ``` This will print out the average score on all test set entries: diff --git a/tutorials/python-api.md b/tutorials/python-api.md index 52f9a7698b..3db62be146 100644 --- a/tutorials/python-api.md +++ b/tutorials/python-api.md @@ -104,6 +104,7 @@ To start with random weights, for example, if you plan a pretraining script, ini ```python from litgpt.api import LLM + llm = LLM.load("pythia-160m", init="random", tokenizer_dir="EleutherAI/pythia-160m") ``` @@ -121,15 +122,12 @@ The `generate_strategy="sequential"` setting loads different parts of the models ```python from litgpt.api import LLM -llm = LLM.load( - "microsoft/phi-2", - distribute=None -) +llm = LLM.load("microsoft/phi-2", distribute=None) llm.distribute( generate_strategy="sequential", devices=4, # Optional setting, otherwise uses all available GPUs - fixed_kv_cache_size=256 # Optionally use a small kv-cache to further reduce memory usage + fixed_kv_cache_size=256, # Optionally use a small kv-cache to further reduce memory usage ) ``` @@ -161,11 +159,7 @@ from litgpt.api import LLM if __name__ == "__main__": - - llm = LLM.load( - model="meta-llama/Meta-Llama-3.1-8B-Instruct", - distribute=None - ) + llm = LLM.load(model="meta-llama/Meta-Llama-3.1-8B-Instruct", distribute=None) llm.distribute(generate_strategy="tensor_parallel", devices=4) @@ -183,10 +177,7 @@ Use the `.benchmark()` method to compare the computational performance of differ from litgpt.api import LLM from pprint import pprint -llm = LLM.load( - model="microsoft/phi-2", - distribute=None -) +llm = LLM.load(model="microsoft/phi-2", distribute=None) llm.distribute(fixed_kv_cache_size=500) @@ -355,7 +346,6 @@ lit_model.llm.generate("hello world") The continued pretraining or finetuning from a downloaded model checkpoint is similar to the example above, except that we can skip the initial steps of instantiating a model with random weights. ```python - lit_model = LitLLM(checkpoint_dir="EleutherAI/pythia-160m") data = Alpaca2k() @@ -380,16 +370,16 @@ lit_model.llm.generate("hello world") Suppose you trained a model and decide to follow up with a few additional training rounds. This can be achieved as follows by loading an existing Trainer checkpoint: ```python - import os + def find_latest_checkpoint(directory): latest_checkpoint = None latest_time = 0 for root, _, files in os.walk(directory): for file in files: - if file.endswith('.ckpt'): + if file.endswith(".ckpt"): file_path = os.path.join(root, file) file_time = os.path.getmtime(file_path) if file_time > latest_time: @@ -398,6 +388,7 @@ def find_latest_checkpoint(directory): return latest_checkpoint + lit_model = LitLLM(checkpoint_dir="EleutherAI/pythia-160m", trainer_ckpt_path=find_latest_checkpoint("lightning_logs")) data.connect(lit_model.llm.tokenizer, batch_size=batch_size, max_seq_length=512)