Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
@@ -1,12 +1,12 @@
repos:
- repo: https://github.com/pre-commit/pre-commit-hooks
rev: v5.0.0
rev: v6.0.0
hooks:
- id: check-yaml
- id: end-of-file-fixer
- id: trailing-whitespace
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.11.0
rev: v0.16.0
hooks:
- id: ruff
args: [--fix]
Expand Down
3 changes: 2 additions & 1 deletion data/test_data_loading.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import pytest
import json

import pytest


@pytest.fixture
def load_data():
Expand Down
11 changes: 6 additions & 5 deletions dataset.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
import json
from typing import Any, Dict, List
from typing import Any

import torch
from loguru import logger
from torch.utils.data import Dataset

from utils.tool_utils import function_formatter


Expand All @@ -18,10 +19,10 @@ def __init__(self, file, tokenizer, max_seq_length, template):
self.observation_format = template["observation_format"]

self.max_seq_length = max_seq_length
logger.info("Loading data: {}".format(file))
logger.info(f"Loading data: {file}")
with open(file, "r", encoding="utf8") as f:
data_list = f.readlines()
logger.info("There are {} data in dataset".format(len(data_list)))
logger.info(f"There are {len(data_list)} data in dataset")
self.data_list = data_list

def __len__(self):
Expand Down Expand Up @@ -93,13 +94,13 @@ def __getitem__(self, index):
return inputs


class SFTDataCollator(object):
class SFTDataCollator:
def __init__(self, tokenizer, max_seq_length):
self.tokenizer = tokenizer
self.max_seq_length = max_seq_length
self.pad_token_id = tokenizer.pad_token_id

def __call__(self, batch: List[Dict[str, Any]]) -> Dict[str, Any]:
def __call__(self, batch: list[dict[str, Any]]) -> dict[str, Any]:
# Find the maximum length in the batch
lengths = [len(x["input_ids"]) for x in batch if x["input_ids"] is not None]
# Take the maximum length in the batch, if it exceeds max_seq_length, take max_seq_length
Expand Down
5 changes: 3 additions & 2 deletions demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,11 @@
from dataclasses import dataclass

import torch
from datasets import Dataset
from peft import LoraConfig
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from trl import SFTTrainer, SFTConfig
from datasets import Dataset
from trl import SFTConfig, SFTTrainer

from dataset import SFTDataCollator, SFTDataset
from utils.constants import model2template

Expand Down
5 changes: 2 additions & 3 deletions full_automation.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,8 @@

import requests
import yaml
from loguru import logger
from huggingface_hub import HfApi
from loguru import logger

from demo import LoraTrainingArguments, train_lora
from utils.constants import model2base_model, model2size
Expand Down Expand Up @@ -36,8 +36,7 @@
# download in chunks
response = requests.get(data_url, stream=True)
with open("data/demo_data.jsonl", "wb") as f:
for chunk in response.iter_content(chunk_size=8192):
f.write(chunk)
f.writelines(response.iter_content(chunk_size=8192))

# train all feasible models and merge
for model_id in all_training_args.keys():
Expand Down
6 changes: 3 additions & 3 deletions utils/tool_utils.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from typing import Dict, Any, List, Tuple
import json
from typing import Any

DEFAULT_TOOL_PROMPT = (
"You have access to the following tools:\n{tool_text}"
Expand All @@ -14,7 +14,7 @@
DEFAULT_FUNCTION_SLOTS = "Action: {name}\nAction Input: {arguments}\n"


def tool_formater(tools: List[Dict[str, Any]]) -> str:
def tool_formater(tools: list[dict[str, Any]]) -> str:
tool_text = ""
tool_names = []
for tool in tools:
Expand Down Expand Up @@ -52,7 +52,7 @@ def tool_formater(tools: List[Dict[str, Any]]) -> str:


def function_formatter(tool_calls, function_slots=DEFAULT_FUNCTION_SLOTS) -> str:
functions: List[Tuple[str, str]] = []
functions: list[tuple[str, str]] = []
if not isinstance(tool_calls, list):
tool_calls = [tool_calls] # parrallel function calls

Expand Down