{"id":15655676,"url":"https://github.com/rasbt/pytorch-fabric-demo","last_synced_at":"2025-05-05T14:42:29.046Z","repository":{"id":143086591,"uuid":"614415893","full_name":"rasbt/pytorch-fabric-demo","owner":"rasbt","description":null,"archived":false,"fork":false,"pushed_at":"2023-03-15T22:31:33.000Z","size":22,"stargazers_count":24,"open_issues_count":0,"forks_count":4,"subscribers_count":3,"default_branch":"main","last_synced_at":"2025-03-30T21:51:12.638Z","etag":null,"topics":[],"latest_commit_sha":null,"homepage":null,"language":"Python","has_issues":true,"has_wiki":null,"has_pages":null,"mirror_url":null,"source_name":null,"license":"apache-2.0","status":null,"scm":"git","pull_requests_enabled":true,"icon_url":"https://github.com/rasbt.png","metadata":{"files":{"readme":"README.md","changelog":null,"contributing":null,"funding":null,"license":"LICENSE","code_of_conduct":null,"threat_model":null,"audit":null,"citation":null,"codeowners":null,"security":null,"support":null,"governance":null,"roadmap":null,"authors":null,"dei":null,"publiccode":null,"codemeta":null}},"created_at":"2023-03-15T14:40:51.000Z","updated_at":"2023-09-08T18:41:54.000Z","dependencies_parsed_at":null,"dependency_job_id":"6f6dcac3-b586-43a2-9c08-2ee254df2736","html_url":"https://github.com/rasbt/pytorch-fabric-demo","commit_stats":null,"previous_names":[],"tags_count":0,"template":false,"template_full_name":null,"repository_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/rasbt%2Fpytorch-fabric-demo","tags_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/rasbt%2Fpytorch-fabric-demo/tags","releases_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/rasbt%2Fpytorch-fabric-demo/releases","manifests_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/rasbt%2Fpytorch-fabric-demo/manifests","owner_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/owners/rasbt","download_url":"https://codeload.github.com/rasbt/pytorch-fabric-demo/tar.gz/refs/heads/main","host":{"name":"GitHub","url":"https://github.com","kind":"github","repositories_count":252516180,"owners_count":21760742,"icon_url":"https://github.com/github.png","version":null,"created_at":"2022-05-30T11:31:42.601Z","updated_at":"2022-07-04T15:15:14.044Z","host_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub","repositories_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories","repository_names_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repository_names","owners_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/owners"}},"keywords":[],"created_at":"2024-10-03T13:00:19.874Z","updated_at":"2025-05-05T14:42:29.030Z","avatar_url":"https://github.com/rasbt.png","language":"Python","funding_links":[],"categories":[],"sub_categories":[],"readme":"# Modifying PyTorch Code to Run With Fabric\n\n\n\nThis repository shows a quick demo for how to modify PyTorch code (here: finetuning a DistilBERT model for 3 epochs to reach 93% accuracy on the IMDB movie review dataset) to make it run faster in Fabric.\n\n\n\nOn a single A100 GPU, the PyTorch code in [src/1_pytorch-distilbert.py](src/1_pytorch-distilbert.py) takes about 24.8 min to run. After adding a few lines for [Fabric](https://lightning.ai/docs/fabric/stable/) as shown in [src/2_pytorch-fabric-distilbert.py](src/2_pytorch-fabric-distilbert.py), it now runs in 1.78 min on 4 A100 GPUs. That's a 14x speed-up!\n\nYou can install `Lightning` + `Fabric` via\n\n    pip install lightning\n\nBelow is the file diff for reference.\n\n```diff\n\nimport os\nimport os.path as op\nimport time\n\n+ from lightning import Fabric\n\nfrom datasets import load_dataset\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\nimport torchmetrics\nfrom transformers import AutoTokenizer\nfrom transformers import AutoModelForSequenceClassification\nfrom watermark import watermark\n\nfrom local_dataset_utilities import download_dataset, load_dataset_into_to_dataframe, partition_dataset\nfrom local_dataset_utilities import IMDBDataset\n\n\ndef tokenize_text(batch):\n    return tokenizer(batch[\"text\"], truncation=True, padding=True)\n\n\ndef plot_logs(log_dir):\n    metrics = pd.read_csv(op.join(log_dir, \"metrics.csv\"))\n\n    aggreg_metrics = []\n    agg_col = \"epoch\"\n    for i, dfg in metrics.groupby(agg_col):\n        agg = dict(dfg.mean())\n        agg[agg_col] = i\n        aggreg_metrics.append(agg)\n\n    df_metrics = pd.DataFrame(aggreg_metrics)\n    df_metrics[[\"train_loss\", \"val_loss\"]].plot(\n        grid=True, legend=True, xlabel=\"Epoch\", ylabel=\"Loss\"\n    )\n    plt.savefig(op.join(log_dir, \"loss.pdf\"))\n\n    df_metrics[[\"train_acc\", \"val_acc\"]].plot(\n        grid=True, legend=True, xlabel=\"Epoch\", ylabel=\"Accuracy\"\n    )\n    plt.savefig(op.join(log_dir, \"acc.pdf\"))\n\n\n- def train(num_epochs, model, optimizer, train_loader, val_loader, device):\n+ def train(num_epochs, model, optimizer, train_loader, val_loader, fabric):\n\n      for epoch in range(num_epochs):\n-         train_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=2).to(device)\n+         train_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=2).to(fabric.device)\n\n        model.train()\n        for batch_idx, batch in enumerate(train_loader):\n\n-             for s in [\"input_ids\", \"attention_mask\", \"label\"]:\n-                 batch[s] = batch[s].to(device)\n\n            outputs = model(batch[\"input_ids\"], attention_mask=batch[\"attention_mask\"], labels=batch[\"label\"]) \n            optimizer.zero_grad()\n-            outputs[\"loss\"].backward()\n+            fabric.backward(outputs[\"loss\"])\n\n            ### UPDATE MODEL PARAMETERS\n            optimizer.step()\n\n            ### LOGGING\n            if not batch_idx % 300:\n                print(f\"Epoch: {epoch+1:04d}/{num_epochs:04d} | Batch {batch_idx:04d}/{len(train_loader):04d} | Loss: {outputs['loss']:.4f}\")\n\n            model.eval()\n            with torch.no_grad():\n                predicted_labels = torch.argmax(outputs[\"logits\"], 1)\n                train_acc.update(predicted_labels, batch[\"label\"])\n\n        ### MORE LOGGING\n        model.eval()\n        with torch.no_grad():\n-            val_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=2).to(device)\n+            val_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=2).to(fabric.device)\n            for batch in val_loader:\n-                for s in [\"input_ids\", \"attention_mask\", \"label\"]:\n-                    batch[s] = batch[s].to(device)\n                outputs = model(batch[\"input_ids\"], attention_mask=batch[\"attention_mask\"], labels=batch[\"label\"])\n                predicted_labels = torch.argmax(outputs[\"logits\"], 1)\n                val_acc.update(predicted_labels, batch[\"label\"])\n\n            print(f\"Epoch: {epoch+1:04d}/{num_epochs:04d} | Train acc.: {train_acc.compute()*100:.2f}% | Val acc.: {val_acc.compute()*100:.2f}%\")\n            train_acc.reset(), val_acc.reset()\n\n\nif __name__ == \"__main__\":\n\n    print(watermark(packages=\"torch,lightning,transformers\", python=True))\n    print(\"Torch CUDA available?\", torch.cuda.is_available())    \n-   device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    torch.manual_seed(123)\n\n    ##########################\n    ### 1 Loading the Dataset\n    ##########################\n    download_dataset()\n    df = load_dataset_into_to_dataframe()\n    if not (op.exists(\"train.csv\") and op.exists(\"val.csv\") and op.exists(\"test.csv\")):\n        partition_dataset(df)\n\n    imdb_dataset = load_dataset(\n        \"csv\",\n        data_files={\n            \"train\": \"train.csv\",\n            \"validation\": \"val.csv\",\n            \"test\": \"test.csv\",\n        },\n    )\n\n    #########################################\n    ### 2 Tokenization and Numericalization\n    #########################################\n\n    tokenizer = AutoTokenizer.from_pretrained(\"distilbert-base-uncased\")\n    print(\"Tokenizer input max length:\", tokenizer.model_max_length, flush=True)\n    print(\"Tokenizer vocabulary size:\", tokenizer.vocab_size, flush=True)\n\n    print(\"Tokenizing ...\", flush=True)\n    imdb_tokenized = imdb_dataset.map(tokenize_text, batched=True, batch_size=None)\n    del imdb_dataset\n    imdb_tokenized.set_format(\"torch\", columns=[\"input_ids\", \"attention_mask\", \"label\"])\n    os.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\n    #########################################\n    ### 3 Set Up DataLoaders\n    #########################################\n\n    train_dataset = IMDBDataset(imdb_tokenized, partition_key=\"train\")\n    val_dataset = IMDBDataset(imdb_tokenized, partition_key=\"validation\")\n    test_dataset = IMDBDataset(imdb_tokenized, partition_key=\"test\")\n\n    train_loader = DataLoader(\n        dataset=train_dataset,\n        batch_size=12,\n        shuffle=True, \n        num_workers=2,\n        drop_last=True,\n    )\n\n    val_loader = DataLoader(\n        dataset=val_dataset,\n        batch_size=12,\n        num_workers=2,\n        drop_last=True,\n    )\n\n    test_loader = DataLoader(\n        dataset=test_dataset,\n        batch_size=12,\n        num_workers=2,\n        drop_last=True,\n    )\n\n\n    #########################################\n    ### 4 Initializing the Model\n    #########################################\n\n+    fabric = Fabric(accelerator=\"cuda\", devices=4, strategy=\"deepspeed_stage_2\", precision=\"16-mixed\")\n+    fabric.launch()\n\n    model = AutoModelForSequenceClassification.from_pretrained(\n        \"distilbert-base-uncased\", num_labels=2)\n\n-   model.to(device)\n    optimizer = torch.optim.Adam(model.parameters(), lr=5e-5)\n\n+    model, optimizer = fabric.setup(model, optimizer)\n+    train_loader, val_loader, test_loader = fabric.setup_dataloaders(train_loader, val_loader, test_loader)\n\n    #########################################\n    ### 5 Finetuning\n    #########################################\n\n    start = time.time()\n    train(\n        num_epochs=3,\n        model=model,\n        optimizer=optimizer,\n        train_loader=train_loader,\n        val_loader=val_loader,\n-       device=device\n+       fabric=fabric\n    )\n\n    end = time.time()\n    elapsed = end-start\n    print(f\"Time elapsed {elapsed/60:.2f} min\")\n\n    with torch.no_grad():\n        model.eval()\n-       test_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=2).to(device)\n+       test_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=2).to(fabric.device)\n        for batch in test_loader:\n-           for s in [\"input_ids\", \"attention_mask\", \"label\"]:\n-               batch[s] = batch[s].to(device)\n            outputs = model(batch[\"input_ids\"], attention_mask=batch[\"attention_mask\"], labels=batch[\"label\"])\n            predicted_labels = torch.argmax(outputs[\"logits\"], 1)\n            test_acc.update(predicted_labels, batch[\"label\"])\n\n    print(f\"Test accuracy {test_acc.compute()*100:.2f}%\")\n\n```\n\n","project_url":"https://awesome.ecosyste.ms/api/v1/projects/github.com%2Frasbt%2Fpytorch-fabric-demo","html_url":"https://awesome.ecosyste.ms/projects/github.com%2Frasbt%2Fpytorch-fabric-demo","lists_url":"https://awesome.ecosyste.ms/api/v1/projects/github.com%2Frasbt%2Fpytorch-fabric-demo/lists"}