From 441ca51af96ec83f642f120b5b91bb3ef5875d31 Mon Sep 17 00:00:00 2001 From: XianBW <36835909+XianBW@users.noreply.github.com> Date: Fri, 7 Jun 2024 00:23:11 +0800 Subject: [PATCH] CI tool refinement (#13) * adjust interact and log for CI fix tool * change method to split large code files * refine statistics table after fixing one file * process default --fix option when ruff check * remove hardcode path * add tree-sitter requirements * replace <| |> to < > in comments, for GPT * update mypy version * Add Mypy Evaluator --- Makefile | 4 +- constraints/3.10.txt | 2 +- constraints/3.11.txt | 2 +- constraints/3.8.txt | 2 +- constraints/3.9.txt | 2 +- pyproject.toml | 4 +- rdagent/app/CI/prompts.py | 126 ---- rdagent/app/CI/prompts.yaml | 117 ++++ rdagent/app/CI/run.py | 647 ++++++++++++------ rdagent/core/prompts.py | 2 +- .../factor_implementation_evolving_cli.py | 6 +- .../share_modules/evaluator.py | 2 +- .../share_modules/factor.py | 3 +- rdagent/knowledge_management/graph.py | 34 +- rdagent/oai/llm_utils.py | 13 +- requirements.txt | 6 +- 16 files changed, 606 insertions(+), 366 deletions(-) delete mode 100644 rdagent/app/CI/prompts.py create mode 100644 rdagent/app/CI/prompts.yaml diff --git a/Makefile b/Makefile index cbd5b8e8..47e73614 100644 --- a/Makefile +++ b/Makefile @@ -93,14 +93,14 @@ mypy: # Check lint with ruff. ruff: - $(PIPRUN) ruff check . --exclude FinCo,finco,rdagent/scripts,test/scripts,git_ignore_folder --line-length 120 + $(PIPRUN) ruff check . --exclude FinCo,finco,rdagent/scripts,test/scripts,git_ignore_folder # Check lint with toml-sort. toml-sort: $(PIPRUN) toml-sort --check pyproject.toml # Check lint with all linters. -lint: black isort mypy ruff toml-sort +lint: mypy ruff toml-sort # Run pre-commit with autofix against all files. pre-commit: diff --git a/constraints/3.10.txt b/constraints/3.10.txt index 9441e17c..0eb3965a 100644 --- a/constraints/3.10.txt +++ b/constraints/3.10.txt @@ -65,7 +65,7 @@ more-itertools==10.1.0 mpmath==1.3.0 msrest==0.7.1 multidict==6.0.4 -mypy==1.8.0 +mypy==1.10.0 mypy-extensions==1.0.0 myst-parser==2.0.0 networkx==3.2.1 diff --git a/constraints/3.11.txt b/constraints/3.11.txt index d4914162..4ac9c55c 100644 --- a/constraints/3.11.txt +++ b/constraints/3.11.txt @@ -63,7 +63,7 @@ more-itertools==10.1.0 mpmath==1.3.0 msrest==0.7.1 multidict==6.0.4 -mypy==1.8.0 +mypy==1.10.0 mypy-extensions==1.0.0 myst-parser==2.0.0 networkx==3.2.1 diff --git a/constraints/3.8.txt b/constraints/3.8.txt index f022bd1f..b1dc0ef5 100644 --- a/constraints/3.8.txt +++ b/constraints/3.8.txt @@ -66,7 +66,7 @@ more-itertools==10.1.0 mpmath==1.3.0 msrest==0.7.1 multidict==6.0.4 -mypy==1.8.0 +mypy==1.10.0 mypy-extensions==1.0.0 myst-parser==2.0.0 networkx==3.1 diff --git a/constraints/3.9.txt b/constraints/3.9.txt index 9441e17c..0eb3965a 100644 --- a/constraints/3.9.txt +++ b/constraints/3.9.txt @@ -65,7 +65,7 @@ more-itertools==10.1.0 mpmath==1.3.0 msrest==0.7.1 multidict==6.0.4 -mypy==1.8.0 +mypy==1.10.0 mypy-extensions==1.0.0 myst-parser==2.0.0 networkx==3.2.1 diff --git a/pyproject.toml b/pyproject.toml index f497268d..9f7a46b4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,14 +52,13 @@ color_output = true profile = "black" [tool.mypy] +explicit_package_bases = true check_untyped_defs = true disallow_any_unimported = true disallow_untyped_defs = true enable_error_code = [ "ignore-without-code", ] -no_implicit_optional = true -show_error_codes = true warn_return_any = true warn_unused_ignores = true @@ -72,6 +71,7 @@ log_format = "%(asctime)s %(levelname)s %(message)s" minversion = "6.0" [tool.ruff] +fix = true line-length = 120 src = ["rdagent"] diff --git a/rdagent/app/CI/prompts.py b/rdagent/app/CI/prompts.py deleted file mode 100644 index cfb9548e..00000000 --- a/rdagent/app/CI/prompts.py +++ /dev/null @@ -1,126 +0,0 @@ -linting_system_prompt_template = "You are a software engineer. \ -You can write code to a high standard and are adept at solving {language} linting problems." - - -user_get_makefile_lint_commands_template = """ -You get a Makefile which contains some linting rules. Here are its content ```{file_text}``` -Please find executable commands about linting from it. - -Please response with following json template: -{{ - "commands": , -}} -""" - -user_get_files_contain_lint_commands_template = """ -You get a file list of a repository. \ -Some file maybe contain linting rules or linting commands which defined by repo authors. -Here are the file list: -``` -{file_list} -``` - -Please find all files maybe correspond to linting from it. -Please response with following json template: -{{ - "files": , -}} -""" - - -generate_lint_command_template = """ -Please generate a command to lint or format a {language} repository. -Here are some information about different linting tools ```{linting_tools}``` -""" - - -suffix2language_template = """ -Here are the files suffix in one code repo: {suffix}. -Please tell me the programming language used in this repo and which language has linting-tools. -Your response should follow this template: -{{ - "languages": , - "languages_with_linting_tools": -}} -""" - - -session_start_template = """ -Please modify the Python code based on the lint info. -Due to the length of the code, I will first tell you the entire code, and then each time I ask a question, \ -I will extract a portion of the code and tell you the error information contained in this code segment. -You need to fix the corresponding error in the code segment \ -and return the code that can replace the corresponding code segment. - -The Python code is from a complete Python project file. Each line of the code is annotated with a line number, \ -separated from the original code by three characters ("|"). The vertical bars are aligned. -Here is the complete code, please be prepared to fix it: -```Python -{code} -``` -""" - - -session_normal_template = """Please modify this code snippet based on the lint info. Here is the code snippet: -```Python -{code} -``` - ------Lint info----- -{lint_info} -------------------- - -The lint info contains one or more errors. \ -Different errors are separated by blank lines. Each error follows this format: ------Lint info format----- -: - - --------------------------- -The error code is an abbreviation set by the checker for ease of describing the error. \ -The error context includes the relevant code around the error, and the helpful information suggests possible fixes. - -Please simply reply the code after you fix all linting errors. -The code you return does not require line numbers, \ -and should just replace the code I provided you, and does not require comments. -Please wrap your code with following format: - -```python - -``` -""" - - -user_template_for_code_snippet = """Please modify the Python code based on the lint info. ------Python Code----- -{code} ---------------------- - ------Lint info----- -{lint_info} -------------------- - -The Python code is a snippet from a complete Python project file. \ -Each line of the code is annotated with a line number, \ -separated from the original code by three characters ("|"). \ -The vertical bars are aligned. - -The lint info contains one or more errors. Different errors are separated by blank lines. \ -Each error follows this format: ------Lint info format----- -: - - --------------------------- -The error code is an abbreviation set by the checker for ease of describing the error. \ -The error context includes the relevant code around the error, and the helpful information suggests possible fixes. - -Please simply reply the code after you fix all linting errors. -The code you return does not require line numbers, \ -and should just replace the code I provided you, and does not require comments. -Please wrap your code with following format: - -```python - -``` -""" diff --git a/rdagent/app/CI/prompts.yaml b/rdagent/app/CI/prompts.yaml new file mode 100644 index 00000000..4499d54a --- /dev/null +++ b/rdagent/app/CI/prompts.yaml @@ -0,0 +1,117 @@ +generate_lint_command_template: | + Please generate a command to lint or format a {language} repository. + Here are some information about different linting tools ```{linting_tools}``` +linting_system_prompt_template: | + You are a software engineer. You can write code to a high standard and are adept at solving {language} linting problems. +session_manual_template: | + There are some problems with the code you provided, please modify the code again according to the instruction and return the errors list you modified. + + Instruction: + {operation} + + Your response format should be like this: + + ```python + + ``` + + ```json + {{ + "errors": [": ", ...] + }} + ``` +session_normal_template: | + Please modify this code snippet based on the lint info. Here is the code snippet: + ```Python + {code} + ``` + + -----Lint info----- + {lint_info} + ------------------- + + The lint info contains one or more errors. Different errors are separated by blank lines. Each error follows this format: + -----Lint info format----- + : + + + -------------------------- + The error code is an abbreviation set by the checker for ease of describing the error. The error position includes the relevant code around the error, and the helpful information provides useful information or possible fix method. + + Please simply reply the code after you fix all linting errors. You should be aware of the following: + 1. The indentation of the code should be consistent with the original code. + 2. You should just replace the code I provided you, which starts from line {start_line} to line {end_line}. + 3. You'll need to add line numbers to the modified code which starts from {start_lineno}. + 4. You don't need to add comments to explain your changes. + Please wrap your code with following format: + + ```python + + ``` +session_start_template: | + Please modify the Python code based on the lint info. + Due to the length of the code, I will first tell you the entire code, and then each time I ask a question, I will extract a portion of the code and tell you the error information contained in this code segment. + You need to fix the corresponding error in the code segment and return the code that can replace the corresponding code segment. + + The Python code is from a complete Python project file. Each line of the code is annotated with a line number, separated from the original code by three characters ("|"). The vertical bars are aligned. + Here is the complete code, please be prepared to fix it: + ```Python + {code} + ``` +suffix2language_template: | + Here are the files suffix in one code repo: {suffix}. + Please tell me the programming language used in this repo and which language has linting-tools. + Your response should follow this template: + {{ + "languages": , + "languages_with_linting_tools": + }} +user_get_files_contain_lint_commands_template: | + You get a file list of a repository. Some files may contain linting rules or linting commands defined by repo authors. + Here are the file list: + ``` + {file_list} + ``` + + Please find all files that may correspond to linting from it. + Please respond with the following JSON template: + {{ + "files": , + }} +user_get_makefile_lint_commands_template: | + You get a Makefile which contains some linting rules. Here are its content: + ``` + {file_text} + ``` + Please find executable commands about linting from it. + Please respond with the following JSON template: + {{ + "commands": ["python -m xxx --params"...], + }} +user_template_for_code_snippet: | + Please modify the Python code based on the lint info. + -----Python Code----- + {code} + --------------------- + + -----Lint info----- + {lint_info} + ------------------- + + The Python code is a snippet from a complete Python project file. Each line of the code is annotated with a line number, separated from the original code by three characters ("|"). The vertical bars are aligned. + + The lint info contains one or more errors. Different errors are separated by blank lines. Each error follows this format: + -----Lint info format----- + : + + + -------------------------- + The error code is an abbreviation set by the checker for ease of describing the error. The error context includes the relevant code around the error, and the helpful information suggests possible fixes. + + Please simply reply the code after you fix all linting errors. + The code you return does not require line numbers, and should just replace the code I provided you, and does not require comments. + Please wrap your code with following format: + + ```python + + ``` \ No newline at end of file diff --git a/rdagent/app/CI/run.py b/rdagent/app/CI/run.py index 2a89835c..b3667373 100644 --- a/rdagent/app/CI/run.py +++ b/rdagent/app/CI/run.py @@ -1,24 +1,17 @@ -""" +from __future__ import annotations -""" +import datetime import json import re import subprocess import time from collections import defaultdict from dataclasses import dataclass -from difflib import IS_LINE_JUNK, ndiff +from difflib import ndiff from pathlib import Path -from typing import Dict, List, Tuple, Union, cast - -from rich import print -from rich.panel import Panel -from rich.prompt import Prompt -from rich.rule import Rule -from rich.syntax import Syntax -from rich.table import Table -from rich.text import Text +from typing import Any, Literal +import tree_sitter_python from rdagent.core.evolving_framework import ( Evaluator, EvoAgent, @@ -28,115 +21,229 @@ from rdagent.core.evolving_framework import ( Feedback, Knowledge, ) +from rdagent.core.prompts import Prompts from rdagent.oai.llm_utils import APIBackend +from rich import print +from rich.panel import Panel +from rich.progress import Progress, SpinnerColumn, TimeElapsedColumn +from rich.prompt import Prompt +from rich.rule import Rule +from rich.syntax import Syntax +from rich.table import Table +from rich.text import Text +from tree_sitter import Language, Node, Parser -from .prompts import ( - linting_system_prompt_template, - session_normal_template, - session_start_template, -) - +py_parser = Parser(Language(tree_sitter_python.language())) +CI_prompts = Prompts(file_path=Path(__file__).parent / "prompts.yaml") @dataclass class CIError: raw_str: str - file_path: Union[Path, str] + file_path: Path | str line: int column: int code: str msg: str hint: str + checker: Literal["ruff", "mypy"] + + def to_dict(self) -> dict[str, object]: + return self.__dict__ + + def __str__(self) -> str: + return f"{self.file_path}:{self.line}:{self.column}: {self.code} {self.msg}\n{self.hint}".strip() @dataclass class CIFeedback(Feedback): - errors: Dict[str, List[CIError]] + errors: dict[str, list[CIError]] + + def statistics(self) -> dict[Literal["ruff", "mypy"], dict[str, int]]: + error_counts = defaultdict(lambda: defaultdict(int)) + for file_errors in self.errors.values(): + for error in file_errors: + error_counts[error.checker][error.code] += 1 + return error_counts @dataclass class FixRecord: - skipped_errors: List[CIError] - directly_fixed_errors: List[CIError] - manually_fixed_errors: List[CIError] + skipped_errors: list[CIError] + directly_fixed_errors: list[CIError] + manually_fixed_errors: list[CIError] + manual_instructions: dict[str, list[CIError]] + + def to_dict(self) -> dict[str, Any]: + return { + "skipped_errors": [error.to_dict() for error in self.skipped_errors], + "directly_fixed_errors": [error.to_dict() for error in self.directly_fixed_errors], + "manually_fixed_errors": [error.to_dict() for error in self.manually_fixed_errors], + "manual_instructions": { + key: [error.to_dict() for error in errors] + for key, errors in self.manual_instructions.items() + }, + } class CodeFile: - def __init__(self, path: Union[Path, str]): + def __init__(self, path: Path | str) -> None: self.path = Path(path) self.load() + @classmethod + def add_line_number(cls: CodeFile, code: list[str] | str, start: int = 1) -> list[str] | str: + code_lines = code.split("\n") if isinstance(code, str) else code + + lineno_width = len(str(start - 1 + len(code_lines))) + code_with_lineno = [] + for i, code_line in enumerate(code_lines): + code_with_lineno.append(f"{i+start: >{lineno_width}} | {code_line}") + + return code_with_lineno if isinstance(code, list) else "\n".join(code_with_lineno) + + + @classmethod + def remove_line_number(cls: CodeFile, code: list[str] | str) -> list[str] | str: + code_lines = code.split("\n") if isinstance(code, str) else code + + try: + code_without_lineno = [re.split(r"\| ", code_line, maxsplit=1)[1] for code_line in code_lines] + except IndexError: + code_without_lineno = ["something went wrong when remove line numbers", *code_lines] + + return code_without_lineno if isinstance(code, list) else "\n".join(code_without_lineno) + def load(self) -> None: code = self.path.read_text(encoding="utf-8") self.code_lines = code.split("\n") - # add line number + # line numbers self.lineno = len(self.code_lines) self.lineno_width = len(str(self.lineno)) - self.code_lines_with_lineno = [] - for i, code_line in enumerate(self.code_lines): - self.code_lines_with_lineno.append(f"{i+1: >{self.lineno_width}} | {code_line}") + self.code_lines_with_lineno = self.add_line_number(self.code_lines) - def get(self, start=0, end=None, add_line_number: bool = False, return_list: bool = False) -> Union[List[str], str]: + def get(self, start: int = 1, end: int | None = None, add_line_number: bool = False, return_list: bool = False) -> list[str] | str: + """ + Retrieves a portion of the code lines. + line number starts from 1, return codes in [start, end]. + + Args: + start (int): The starting line number (inclusive). Defaults to 1. + end (int): The ending line number (inclusive). Defaults to None, which means the last line. + add_line_number (bool): Whether to include line numbers in the result. Defaults to False. + return_list (bool): Whether to return the result as a list of lines + or as a single string. Defaults to False. + + Returns: + Union[List[str], str]: The code lines as a list of strings or as a + single string, depending on the value of `return_list`. + """ start -= 1 if start < 0: start = 0 - end = self.lineno if end is None else end - 1 - + end = self.lineno if end is None else end + if end <= start: + res = [] res = self.code_lines_with_lineno[start:end] if add_line_number else self.code_lines[start:end] return res if return_list else "\n".join(res) - def apply_changes(self, changes: List[Tuple[int, int, str]]) -> None: + def apply_changes(self, changes: list[tuple[int, int, str]]) -> None: + """ + Applies the given changes to the code lines. + + Args: + changes (List[Tuple[int, int, str]]): A list of tuples representing the changes to be applied. + Each tuple contains the start line number, end line number, and the new code to be inserted. + + Returns: + None + """ offset = 0 for start, end, code in changes: - start -= 1 - if start < 0: - start = 0 - end -= 1 + adjusted_start = max(start - 1, 0) new_code = code.split("\n") - self.code_lines[start + offset : end + offset] = new_code + self.code_lines[adjusted_start+offset:end+offset] = new_code offset += len(new_code) - (end - start) self.path.write_text("\n".join(self.code_lines), encoding="utf-8") self.load() - def __str__(self): + def get_code_blocks(self, max_lines: int = 30) -> list[tuple[int, int]]: + tree = py_parser.parse(bytes("\n".join(self.code_lines), "utf8")) + + def get_blocks_in_node(node: Node, max_lines: int) -> list[tuple[int, int]]: + if node.type == "assignment": + return [(node.start_point.row, node.end_point.row + 1)] + + blocks: list[tuple[int, int]] = [] + block: tuple[int, int] | None = None # [start, end), line number starts from 0 + + for child in node.children: + if child.end_point.row + 1 - child.start_point.row > max_lines: + if block is not None: + blocks.append(block) + block = None + blocks.extend(get_blocks_in_node(child, max_lines)) + elif block is None: + block = (child.start_point.row, child.end_point.row + 1) + elif child.end_point.row + 1 - block[0] <= max_lines: + block = (block[0], child.end_point.row + 1) + else: + blocks.append(block) + block = (child.start_point.row, child.end_point.row + 1) + + if block is not None: + blocks.append(block) + + return blocks + + # change line number to start from 1 and [start, end) to [start, end] + return [(a+1,b) for a,b in get_blocks_in_node(tree.root_node, max_lines)] + + def __str__(self) -> str: return f"{self.path}" class Repo(EvolvableSubjects): - def __init__(self, project_path: Union[Path, str], **kwargs): + def __init__(self, project_path: Path | str, excludes: list[Path] = [], **kwargs: Any) -> None: self.params = kwargs self.project_path = Path(project_path) + + excludes = [self.project_path / path for path in excludes] + git_ignored_output = subprocess.check_output( - "git status --ignored -s", - shell=True, + ["git", "status", "--ignored", "-s"], cwd=project_path, stderr=subprocess.STDOUT, - ).decode("utf-8") + text=True, + ) git_ignored_files = [ (self.project_path / Path(line[3:])).resolve() for line in git_ignored_output.split("\n") if line.startswith("!!") ] + excludes.extend(git_ignored_files) + files = [ file for file in self.project_path.glob("**/*") if file.is_file() - and not any(str(file).startswith(str(path)) for path in git_ignored_files) + and not any(str(file).startswith(str(path)) for path in excludes) and ".git/" not in str(file) and file.suffix == ".py" ] self.files = {file: CodeFile(file) for file in files} - self.fix_records: Dict[str, FixRecord] | None = None + self.fix_records: dict[str, FixRecord] | None = None @dataclass class RuffRule: """ + Example: { "name": "missing-trailing-comma", "code": "COM812", @@ -146,53 +253,54 @@ class RuffRule: "Trailing comma missing" ], "fix": "Fix is always available.", - "explanation": "## What it does\nChecks for the absence of trailing commas.\n\n## Why is this bad?\nThe presence of a trailing comma can reduce diff size when parameters or\nelements are added or removed from function calls, function definitions,\nliterals, etc.\n\n## Example\n```python\nfoo = {\n \"bar\": 1,\n \"baz\": 2\n}\n```\n\nUse instead:\n```python\nfoo = {\n \"bar\": 1,\n \"baz\": 2,\n}\n```\n", + "explanation": "...", "preview": false } """ - name: str code: str linter: str summary: str - message_formats: List[str] + message_formats: list[str] fix: str explanation: str preview: bool class RuffEvaluator(Evaluator): - """The error message are generated by - `python -m ruff . --exclude FinCo,finco,fincov1 --ignore ANN101,TCH003,D,ERA001` + """ + The error message are generated by command """ - def __init__(self, command: str = None): + def __init__(self, command: str | None = None) -> None: if command is None: - self.command = "ruff check . --no-fix --output-format full" + self.command = "ruff check . --output-format full" else: self.command = command - def explain_rule(self, error_code: str) -> RuffRule: - explain_command = "ruff rule {error_code} --output-format json" + @staticmethod + def explain_rule(error_code: str) -> RuffRule: + explain_command = f"ruff rule {error_code} --output-format json" try: out = subprocess.check_output( - explain_command.format(error_code=error_code), - shell=True, + explain_command, stderr=subprocess.STDOUT, + text=True, ) except subprocess.CalledProcessError as e: out = e.output - return json.loads(out.decode()) + return RuffRule(**json.loads(out)) - def evaluate(self, evo: Repo, **kwargs) -> CIFeedback: + + def evaluate(self, evo: Repo, **kwargs: Any) -> CIFeedback: """Simply run ruff to get the feedbacks.""" try: out = subprocess.check_output( - self.command, - shell=True, + self.command.split(), cwd=evo.project_path, stderr=subprocess.STDOUT, + text=True, ) except subprocess.CalledProcessError as e: out = e.output @@ -210,182 +318,344 @@ class RuffEvaluator(Evaluator): # extract error info pattern = r"(([^\n]*):(\d+):(\d+): (\w+) ([^\n]*)\n(.*?))\n\n" - matches = re.findall(pattern, out.decode(), re.DOTALL) + matches = re.findall(pattern, out, re.DOTALL) errors = defaultdict(list) + for match in matches: raw_str, file_path, line_number, column_number, error_code, error_message, error_hint = match - error = CIError( - raw_str=raw_str, - file_path=file_path, - line=int(line_number), - column=int(column_number), - code=error_code, - msg=error_message, - hint=error_hint, + + # TODO @bowen: filter these files when running the check command + if evo.project_path / Path(file_path) not in evo.files: + continue + error = CIError(raw_str=raw_str, + file_path=file_path, + line=int(line_number), + column=int(column_number), + code=error_code, + msg=error_message, + hint=error_hint, + checker="ruff") + + errors[file_path].append(error) + + return CIFeedback(errors=errors) + +class MypyEvaluator(Evaluator): + + def __init__(self, command: str | None = None) -> None: + if command is None: + self.command = "mypy . --pretty --no-error-summary --show-column-numbers" + else: + self.command = command + + def evaluate(self, evo: Repo, **kwargs: Any) -> CIFeedback: + try: + out = subprocess.check_output( + self.command.split(), + cwd=evo.project_path, + stderr=subprocess.STDOUT, + text=True, ) + except subprocess.CalledProcessError as e: + out = e.output + + errors = defaultdict(list) + + out = re.sub(r"([^\n]*?:\d+:\d+): error:", r"\n\1: error:", out) + out += "\n" + pattern = r"(([^\n]*?):(\d+):(\d+): error:(.*?)\s\[([\w-]*?)\]\s(.*?))\n\n" + for match in re.findall(pattern, out, re.DOTALL): + raw_str, file_path, line_number, column_number, error_message, error_code, error_hint = match + error_message = error_message.strip().replace("\n", " ") + if re.match(r".*[^\n]*?:\d+:\d+: note:.*", error_hint, re.DOTALL) is not None: + error_hint_position = re.split(r"[^\n]*?:\d+:\d+: note:", error_hint, re.DOTALL)[0] + error_hint_help = re.findall(r"^.*?:\d+:\d+: note: (.*)$", error_hint, re.MULTILINE) + error_hint_help = "\n".join(error_hint_help) + error_hint = f"{error_hint_position}\nHelp:\n{error_hint_help}" + + if evo.project_path / Path(file_path) not in evo.files: + continue + error = CIError(raw_str=raw_str, + file_path=file_path, + line=int(line_number), + column=int(column_number), + code=error_code, + msg=error_message, + hint=error_hint, + checker="mypy") + errors[file_path].append(error) return CIFeedback(errors=errors) -class MypyEvaluator(Evaluator): - def __init__(self, command: str = None): - if command is None: - self.command = "mypy . --explicit-package-bases" - else: - self.command = command +class MultiEvaluator(Evaluator): - def evaluate(self, evo: Repo, **kwargs) -> CIFeedback: - try: - out = subprocess.check_output( - self.command, - shell=True, - cwd=evo.project_path, - stderr=subprocess.STDOUT, - ) - except subprocess.CalledProcessError as e: - out = e.output + def __init__(self, *evaluators: Evaluator) -> None: + self.evaluators = evaluators - return CIFeedback(cast(str, out).decode("utf-8")) + def evaluate(self, evo: Repo, **kwargs: Any) -> CIFeedback: + all_errors = defaultdict(list) + for evaluator in self.evaluators: + feedback: CIFeedback = evaluator.evaluate(evo, **kwargs) + for file_path, errors in feedback.errors.items(): + all_errors[file_path].extend(errors) + + # sort errors by position + for file_path in all_errors: + all_errors[file_path].sort(key=lambda x: (x.line, x.column)) + + return CIFeedback(errors=all_errors) class CIEvoStr(EvolvingStrategy): def evolve( self, evo: Repo, - evolving_trace: List[EvoStep] = [], - knowledge_l: List[Knowledge] = [], - **kwargs, + evolving_trace: list[EvoStep] | None = None, + knowledge_l: list[Knowledge] | None = None, + **kwargs: Any, ) -> Repo: + + @dataclass + class CodeFixGroup: + start_line: int + end_line: int + errors: list[CIError] + session_id: str + responses: list[str] + api = APIBackend() - system_prompt = linting_system_prompt_template.format(language="Python") + system_prompt = CI_prompts["linting_system_prompt_template"].format(language="Python") if len(evolving_trace) > 0: last_feedback: CIFeedback = evolving_trace[-1].feedback - fix_records: Dict[str, FixRecord] = defaultdict(lambda: FixRecord([], [], [])) - # iterate by file + + # print statistics + checker_error_counts = {checker: sum(c_statistics.values()) for checker, c_statistics in last_feedback.statistics().items()} + print(f"Found [red]{sum(checker_error_counts.values())}[/red] errors, including: " + + ", ".join(f"[red]{count}[/red] [magenta]{checker}[/magenta] errors" for checker, count in checker_error_counts.items())) + + fix_records: dict[str, FixRecord] = defaultdict(lambda: FixRecord([], [], [], defaultdict(list))) + + # Group errors by code blocks + fix_groups: dict[str, list[CodeFixGroup]] = defaultdict(list) + changes: dict[str, list[tuple[int, int, str]]] = defaultdict(list) for file_path, errors in last_feedback.errors.items(): - print(Rule(f"[cyan]Fixing {file_path}[/cyan]", style="bold cyan", align="left", characters=".")) + file = evo.files[evo.project_path / Path(file_path)] + + # check if the file needs to add `from __future__ import annotations` + # need to add rules here for different languages/tools + # TODO @bowen: current way of handling errors like 'Add import statement' may be not good + for error in errors: + if error.code in ("FA100", "FA102"): + changes[file_path].append((0, 0, "from __future__ import annotations\n")) + break + + # Group errors by code blocks + error_p = 0 + for start_line, end_line in file.get_code_blocks(max_lines=30): + group_errors: list[CIError] = [] + + # collect errors in the same code block + while error_p < len(errors) and start_line <= errors[error_p].line <= end_line: + if errors[error_p].code not in ("FA100", "FA102"): + group_errors.append(errors[error_p]) + error_p += 1 + + # process errors in the code block + if group_errors: + session = api.build_chat_session(session_system_prompt=system_prompt) + session_id = session.get_conversation_id() + session.build_chat_completion( + CI_prompts["session_start_template"].format(code=file.get(add_line_number=True)), + ) + + fix_groups[file_path].append( + CodeFixGroup(start_line, end_line, group_errors, session_id, []), + ) + + + # Fix errors in each code block + with Progress(SpinnerColumn(), *Progress.get_default_columns(), TimeElapsedColumn()) as progress: + group_counts = sum([len(groups) for groups in fix_groups.values()]) + task_id = progress.add_task("Fixing repo...", total=group_counts) + + for file_path in fix_groups: + file = evo.files[evo.project_path / Path(file_path)] + for code_fix_g in fix_groups[file_path]: + start_line, end_line, group_errors = code_fix_g.start_line, code_fix_g.end_line, code_fix_g.errors + code_snippet_with_lineno = file.get( + start_line, end_line, add_line_number=True, return_list=False, + ) + errors_str = "\n\n".join(str(e) for e in group_errors) + + # ask LLM to repair current code snippet + user_prompt = CI_prompts["session_normal_template"].format( + code=code_snippet_with_lineno, + lint_info=errors_str, + start_line=start_line, + end_line=end_line, + start_lineno=start_line, + ) + + session = api.build_chat_session(conversation_id=code_fix_g.session_id) + res = session.build_chat_completion(user_prompt) + + code_fix_g.responses.append(res) + progress.update(task_id, description=f"[green]Fixing[/green] [cyan]{file_path}[/cyan]...", advance=1) + + + # Manual inspection and repair + for file_path in last_feedback.errors: + print(Rule(f"[bright_blue]Checking[/bright_blue] [cyan]{file_path}[/cyan]", style="bright_blue", align="left", characters=".")) file = evo.files[evo.project_path / Path(file_path)] - # Group errors based on position - # TODO @bowen: Crossover between different groups after adding 3 lines of context - groups: List[List[CIError]] = [] - near_errors = [errors[0]] - for error in errors[1:]: - if error.line - near_errors[-1].line <= 6: - near_errors.append(error) - else: - groups.append(near_errors) - near_errors = [error] - groups.append(near_errors) - - changes = [] - # generate changes - for group_id, group in enumerate(groups, start=1): - session = api.build_chat_session(session_system_prompt=system_prompt) - session.build_chat_completion(session_start_template.format(code=file.get(add_line_number=True))) + for group_id, code_fix_g in enumerate(fix_groups[file_path], start=1): + start_line, end_line, group_errors = code_fix_g.start_line, code_fix_g.end_line, code_fix_g.errors + session = api.build_chat_session(conversation_id=code_fix_g.session_id) - print(f"[yellow]Fixing part {group_id}...[/yellow]\n") + print(f"[yellow]Checking part {group_id}...[/yellow]") + + front_context = file.get(start_line-3, start_line-1) + rear_context = file.get(end_line+1, end_line+3) + front_context_with_lineno = file.get(start_line-3, start_line-1, add_line_number=True) + rear_context_with_lineno = file.get(end_line+1, end_line+3, add_line_number=True) - start_line = group[0].line - 3 - end_line = group[-1].line + 3 + 1 code_snippet_with_lineno = file.get(start_line, end_line, add_line_number=True, return_list=False) + + # print errors + printed_errors_str = "\n".join( + [f"[{error.checker}] {error.line: >{file.lineno_width}}:{error.column: <4} {error.code} {error.msg}" for error in group_errors], + ) + print( + Panel.fit( + Syntax(printed_errors_str, lexer="python", background_color="default"), + title=f"{len(group_errors)} Errors", + ), + ) + + # print original code + table = Table(show_header=False, box=None) + table.add_column() + table.add_row(Syntax(front_context_with_lineno, lexer="python", background_color="default")) + table.add_row(Rule(style="dark_orange")) + table.add_row(Syntax(code_snippet_with_lineno, lexer="python", background_color="default")) + table.add_row(Rule(style="dark_orange")) + table.add_row(Syntax(rear_context_with_lineno, lexer="python", background_color="default")) + print(Panel.fit(table, title="Original Code")) + + res = code_fix_g.responses[0] code_snippet_lines = file.get(start_line, end_line, add_line_number=False, return_list=True) - # front_anchor_code = file.get(start_line-3, start_line, add_line_number=False, return_list=False) - # rear_anchor_code = file.get(end_line+1, end_line+3+1, add_line_number=False, return_list=False) - - errors_str = "\n".join([f"{error.raw_str}\n" for error in group]) - - print( - Panel.fit( - Syntax( - "\n".join([f"{error.line}: {error.msg}" for error in group]), - lexer="python", - background_color="default", - ), - title=f"{len(group)} Errors", - ) - ) - # print(f"[bold yellow]original code:[/bold yellow]\n\n{code_snippet_with_lineno}") - print( - Panel.fit( - Syntax(code_snippet_with_lineno, lexer="python", background_color="default"), - title="Original Code", - ) - ) - user_prompt = session_normal_template.format( - code=code_snippet_with_lineno, - lint_info=errors_str, - ) - res = session.build_chat_completion(user_prompt) - - manual_fix_flag = False - while True: - new_code = re.search(r".*```[Pp]ython\n(.*)\n```.*", res, re.DOTALL).group(1) + try: + new_code = re.search(r".*```[Pp]ython\n(.*?)\n```.*", res, re.DOTALL).group(1) + except Exception: + print(f"[red]Error when extract codes[/red]:\n {res}") + try: + fixed_errors_info = re.search(r".*```[Jj]son\n(.*?)\n```.*", res, re.DOTALL).group(1) + fixed_errors_info = json.loads(fixed_errors_info) + except Exception: + fixed_errors_info = None + + new_code = CodeFile.remove_line_number(new_code) # print repair status (code diff) - diff = ndiff(code_snippet_lines, new_code.split("\n"), linejunk=IS_LINE_JUNK) + diff = ndiff(code_snippet_lines, new_code.split("\n")) + + # add 2 spaces to align with diff format + front_context = re.sub(r"^", " ", front_context, flags=re.MULTILINE) + rear_context = re.sub(r"^", " ", rear_context, flags=re.MULTILINE) + table = Table(show_header=False, box=None) table.add_column() + table.add_column() + table.add_column() + table.add_row("", "", Syntax(front_context, lexer="python", background_color="default")) + table.add_row("", "", Rule(style="dark_orange")) + diff_original_lineno = start_line + diff_new_lineno = start_line for i in diff: if i.startswith("+"): - table.add_row(Text(i, style="green")) + table.add_row("", Text(str(diff_new_lineno), style="green bold"), + Text(i, style="green")) + diff_new_lineno += 1 elif i.startswith("-"): - table.add_row(Text(i, style="red")) + table.add_row(Text(str(diff_original_lineno), style="red bold"), "", + Text(i, style="red")) + diff_original_lineno += 1 elif i.startswith("?"): - table.add_row(Text(i, style="yellow")) + table.add_row("", "", Text(i, style="yellow")) else: - table.add_row(Syntax(i, lexer="python", background_color="default")) + table.add_row(str(diff_original_lineno), str(diff_new_lineno), + Syntax(i, lexer="python", background_color="default")) + diff_original_lineno += 1 + diff_new_lineno += 1 + table.add_row("", "", Rule(style="dark_orange")) + table.add_row("", "", Syntax(rear_context, lexer="python", background_color="default")) print(Panel.fit(table, title="Repair Status")) - operation = input("Input your operation: ") - if operation == "s" or operation == "skip": - fix_records[file_path].skipped_errors.extend(group) + operation = Prompt.ask("Input your operation [ [red]([bold]s[/bold])kip[/red] / " + "[green]([bold]a[/bold])pply[/green] / [yellow]manual instruction[/yellow] ]") + print() + if operation in ("s", "skip"): + fix_records[file_path].skipped_errors.extend(group_errors) break - if operation == "a" or operation == "apply": - if manual_fix_flag: - fix_records[file_path].manually_fixed_errors.extend(group) + if operation in ("a", "apply"): + if fixed_errors_info: + fixed_errors_str = "\n".join(fixed_errors_info["errors"]) + for error in group_errors: + if f"{error.line}:{error.column}" in fixed_errors_str: + fix_records[file_path].manually_fixed_errors.append(error) + else: + fix_records[file_path].skipped_errors.append(error) else: - fix_records[file_path].directly_fixed_errors.extend(group) + fix_records[file_path].directly_fixed_errors.extend(group_errors) - changes.append((start_line, end_line, new_code)) + changes[file_path].append((start_line, end_line, new_code)) break - manual_fix_flag = True - res = session.build_chat_completion(operation) + fix_records[file_path].manual_instructions[operation].extend(group_errors) + res = session.build_chat_completion(CI_prompts["session_manual_template"].format(operation=operation)) + code_fix_g.responses.append(res) # apply changes - file.apply_changes(changes) + file.apply_changes(changes[file_path]) evo.fix_records = fix_records return evo -# DIR = "/home/bowen/workspace/fincov2_test/" -DIR = "/home/bowen/workspace/RD-Agent/" -PY = "/home/bowen/miniconda3/envs/cr/bin/python" +DIR = None +while DIR is None or not DIR.exists(): + DIR = Prompt.ask("Please input the [cyan]project directory[/cyan]") + DIR = Path(DIR) + +excludes = Prompt.ask("Input the [dark_orange]excluded directories[/dark_orange] (relative to [cyan]project path[/cyan] and separated by whitespace)").split(" ") +excludes = [Path(exclude.strip()) for exclude in excludes if exclude.strip() != ""] start_time = time.time() +start_timestamp = datetime.datetime.now(datetime.timezone.utc).strftime("%m%d%H%M") -evo = Repo(DIR, python_path=PY) - -eval = RuffEvaluator() +repo = Repo(DIR, excludes=excludes) +evaluator = MultiEvaluator(MypyEvaluator(), RuffEvaluator()) estr = CIEvoStr() rag = None # RAG is not enable firstly. ea = EvoAgent(estr, rag=rag) -ea.step_evolving(evo, eval) +ea.step_evolving(repo, evaluator) while True: print(Rule(f"Round {len(ea.evolving_trace)} repair", style="blue")) - evo: Repo = ea.step_evolving(evo, eval) + repo: Repo = ea.step_evolving(repo, evaluator) - fix_records = evo.fix_records + fix_records = repo.fix_records + filename = f"{DIR.name}_{start_timestamp}_round_{len(ea.evolving_trace)}_fix_records.json" + with Path(filename).open("w") as file: + json.dump({k:v.to_dict() for k,v in fix_records.items()}, file, indent=4) # Count the number of skipped errors skipped_errors_count = 0 @@ -394,26 +664,30 @@ while True: skipped_errors_code_count = defaultdict(int) directly_fixed_errors_code_count = defaultdict(int) manually_fixed_errors_code_count = defaultdict(int) + code_message = defaultdict(str) for record in fix_records.values(): skipped_errors_count += len(record.skipped_errors) directly_fixed_errors_count += len(record.directly_fixed_errors) manually_fixed_errors_count += len(record.manually_fixed_errors) for error in record.skipped_errors: skipped_errors_code_count[error.code] += 1 + code_message[error.code] = error.msg for error in record.directly_fixed_errors: directly_fixed_errors_code_count[error.code] += 1 + code_message[error.code] = error.msg for error in record.manually_fixed_errors: manually_fixed_errors_code_count[error.code] += 1 + code_message[error.code] = error.msg skipped_errors_statistics = "" directly_fixed_errors_statistics = "" manually_fixed_errors_statistics = "" for code, count in sorted(skipped_errors_code_count.items(), key=lambda x: x[1], reverse=True): - skipped_errors_statistics += f"{count: >5} {code: >10} {eval.explain_rule(code).summary}\n" + skipped_errors_statistics += f"{count: >5} {code: >10} {code_message[code]}\n" for code, count in sorted(directly_fixed_errors_code_count.items(), key=lambda x: x[1], reverse=True): - directly_fixed_errors_statistics += f"{count: >5} {code: >10} {eval.explain_rule(code).summary}\n" + directly_fixed_errors_statistics += f"{count: >5} {code: >10} {code_message[code]}\n" for code, count in sorted(manually_fixed_errors_code_count.items(), key=lambda x: x[1], reverse=True): - manually_fixed_errors_statistics += f"{count: >5} {code: >10} {eval.explain_rule(code).summary}\n" + manually_fixed_errors_statistics += f"{count: >5} {code: >10} {code_message[code]}\n" # Create a table to display the counts and ratios table = Table(title="Error Fix Statistics") @@ -423,28 +697,19 @@ while True: table.add_column("Ratio") total_errors_count = skipped_errors_count + directly_fixed_errors_count + manually_fixed_errors_count - table.add_row("Total Errors", "", str(total_errors_count), "") - table.add_row( - "Skipped Errors", - skipped_errors_statistics, - str(skipped_errors_count), - f"{skipped_errors_count / total_errors_count:.2%}", - ) - table.add_row( - "Directly Fixed Errors", - directly_fixed_errors_statistics, - str(directly_fixed_errors_count), - f"{directly_fixed_errors_count / total_errors_count:.2%}", - ) - table.add_row( - "Manually Fixed Errors", - manually_fixed_errors_statistics, - str(manually_fixed_errors_count), - f"{manually_fixed_errors_count / total_errors_count:.2%}", - ) + table.add_row("Total Errors", "", Text(str(total_errors_count), style="cyan"), "") + table.add_row(Text("Skipped Errors", style="red"), skipped_errors_statistics, + Text(str(skipped_errors_count), style="red"), + Text(f"{skipped_errors_count / total_errors_count:.2%}"), style="red") + table.add_row(Text("Directly Fixed Errors", style="green"), directly_fixed_errors_statistics, + Text(str(directly_fixed_errors_count), style="green"), + Text(f"{directly_fixed_errors_count / total_errors_count:.2%}"), style="green") + table.add_row(Text("Manually Fixed Errors", style="yellow"), manually_fixed_errors_statistics, + Text(str(manually_fixed_errors_count), style="yellow"), + Text(f"{manually_fixed_errors_count / total_errors_count:.2%}"), style="yellow") print(table) - operation = Prompt.ask("Start next round? (y/n): ", choices=["y", "n"]) + operation = Prompt.ask("Start next round? (y/n)", choices=["y", "n"]) if operation == "n": break diff --git a/rdagent/core/prompts.py b/rdagent/core/prompts.py index e138d953..4b6b8cef 100644 --- a/rdagent/core/prompts.py +++ b/rdagent/core/prompts.py @@ -5,7 +5,7 @@ import yaml from rdagent.core.utils import SingletonBaseClass -class Prompts(Dict, SingletonBaseClass): +class Prompts(Dict[str, str], SingletonBaseClass): def __init__(self, file_path: Path): prompt_yaml_dict = yaml.load( open( diff --git a/rdagent/factor_implementation/evolving/factor_implementation_evolving_cli.py b/rdagent/factor_implementation/evolving/factor_implementation_evolving_cli.py index 2b410eab..ba5eb97f 100644 --- a/rdagent/factor_implementation/evolving/factor_implementation_evolving_cli.py +++ b/rdagent/factor_implementation/evolving/factor_implementation_evolving_cli.py @@ -262,11 +262,7 @@ class FactorImplementationEvolvingCli: print([feedback.final_decision if feedback is not None else None for feedback in feedbacks].count(True)) def implement_amc( - self, - evo_sub_path_str, - former_knowledge_base_path_str, - implementation_dump_path_str, - slice_index, + self, evo_sub_path_str, former_knowledge_base_path_str, implementation_dump_path_str, slice_index, ): factor_implementations: FactorImplementationList = pickle.load(open(evo_sub_path_str, "rb")) factor_implementations.target_factor_tasks = factor_implementations.target_factor_tasks[ diff --git a/rdagent/factor_implementation/share_modules/evaluator.py b/rdagent/factor_implementation/share_modules/evaluator.py index 41c14210..3820e0a7 100644 --- a/rdagent/factor_implementation/share_modules/evaluator.py +++ b/rdagent/factor_implementation/share_modules/evaluator.py @@ -10,7 +10,7 @@ from factor_implementation.share_modules.factor import ( from factor_implementation.share_modules.prompt import FactorImplementationPrompts from finco.log import FinCoLog from jinja2 import Template -from oai.llm_utils import APIBackend +from rdagent.oai.llm_utils import APIBackend from rdagent.factor_implementation.share_modules.factor_implementation_config import ( FactorImplementSettings, diff --git a/rdagent/factor_implementation/share_modules/factor.py b/rdagent/factor_implementation/share_modules/factor.py index 29c8eae1..5d211944 100644 --- a/rdagent/factor_implementation/share_modules/factor.py +++ b/rdagent/factor_implementation/share_modules/factor.py @@ -124,8 +124,7 @@ class FileBasedFactorImplementation(FactorImplementation): raise ValueError(self.FB_CODE_NOT_SET) with FileLock(self.workspace_path / "execution.lock"): (Path.cwd() / "git_ignore_folder" / "factor_implementation_execution_cache").mkdir( - exist_ok=True, - parents=True, + exist_ok=True, parents=True, ) if FactorImplementSettings().enable_execution_cache: # NOTE: cache the result for the same code diff --git a/rdagent/knowledge_management/graph.py b/rdagent/knowledge_management/graph.py index a30e5390..b3f4bf19 100644 --- a/rdagent/knowledge_management/graph.py +++ b/rdagent/knowledge_management/graph.py @@ -238,8 +238,7 @@ class UndirectedGraph(Graph): result.append(node) for neighbor in sorted( - list(self.get_node(node.id).neighbors), - key=lambda x: x.content, + list(self.get_node(node.id).neighbors), key=lambda x: x.content, ): # to make sure the result is deterministic if neighbor not in visited: if not (block and neighbor.label not in constraint_labels): @@ -276,16 +275,12 @@ class UndirectedGraph(Graph): for node in nodes: if intersection is None: intersection = self.get_nodes_within_steps( - node, - steps=steps, - constraint_labels=constraint_labels, + node, steps=steps, constraint_labels=constraint_labels, ) intersection = self.intersection( nodes1=intersection, nodes2=self.get_nodes_within_steps( - node, - steps=steps, - constraint_labels=constraint_labels, + node, steps=steps, constraint_labels=constraint_labels, ), ) @@ -398,9 +393,7 @@ class UndirectedGraph(Graph): res_list = [] for query in content: similar_nodes = self.semantic_search( - content=query, - topk_k=topk_k, - similarity_threshold=similarity_threshold, + content=query, topk_k=topk_k, similarity_threshold=similarity_threshold, ) connected_nodes = [] @@ -414,7 +407,11 @@ class UndirectedGraph(Graph): block=block, ) connected_nodes.extend( - [node for node in graph_query_node_res if node not in connected_nodes], + [ + node + for node in graph_query_node_res + if node not in connected_nodes + ], ) if len(connected_nodes) >= topk_k: break @@ -458,9 +455,7 @@ def graph_to_edges(graph: Dict[str, List[str]]): def assign_random_coordinate_to_node( - nodes: List, - scope: float = 1.0, - origin: Tuple = (0.0, 0.0), + nodes: List, scope: float = 1.0, origin: Tuple = (0.0, 0.0), ) -> Dict: coordinates = {} @@ -473,10 +468,7 @@ def assign_random_coordinate_to_node( def assign_isometric_coordinate_to_node( - nodes: List, - x_step: float = 1.0, - x_origin: float = 0.0, - y_origin: float = 0.0, + nodes: List, x_step: float = 1.0, x_origin: float = 0.0, y_origin: float = 0.0, ) -> Dict: coordinates = {} @@ -489,9 +481,7 @@ def assign_isometric_coordinate_to_node( def curly_node_coordinate( - coordinates: Dict, - center_y: float = 1.0, - r: float = 1.0, + coordinates: Dict, center_y: float = 1.0, r: float = 1.0, ) -> Dict: # noto: this method can only curly < 90 degree, and the curl line is circle. # the original funtion is: x**2 + (y-m)**2 = r**2 diff --git a/rdagent/oai/llm_utils.py b/rdagent/oai/llm_utils.py index 71d7d859..f20b990b 100644 --- a/rdagent/oai/llm_utils.py +++ b/rdagent/oai/llm_utils.py @@ -15,7 +15,6 @@ from typing import List, Optional, Tuple, Union import numpy as np import tiktoken - from rdagent.core.conf import FincoSettings as Config from rdagent.core.log import FinCoLog, LogColors from rdagent.core.utils import SingletonBaseClass @@ -205,9 +204,7 @@ class ChatSession: messages = self.build_chat_completion_message(user_prompt, **kwargs) response = self.api_backend._try_create_chat_completion_or_embedding( - messages=messages, - chat_completion=True, - **kwargs, + messages=messages, chat_completion=True, **kwargs, ) messages.append( { @@ -422,9 +419,7 @@ class APIBackend: elif isinstance(input_content, list): input_content_list = input_content resp = self._try_create_chat_completion_or_embedding( - input_content_list=input_content_list, - embedding=True, - **kwargs, + input_content_list=input_content_list, embedding=True, **kwargs, ) if isinstance(input_content, str): return resp[0] @@ -669,7 +664,7 @@ class APIBackend: tokens_per_message = 3 tokens_per_name = 1 else: - tokens_per_message = 4 # every message follows <|start|>{role/name}\n{content}<|end|>\n + tokens_per_message = 4 # every message follows {role/name}\n{content}\n tokens_per_name = -1 # if there's a name, the role is omitted num_tokens = 0 for message in messages: @@ -678,7 +673,7 @@ class APIBackend: num_tokens += len(self.encoder.encode(value)) if key == "name": num_tokens += tokens_per_name - num_tokens += 3 # every reply is primed with <|start|>assistant<|message|> + num_tokens += 3 # every reply is primed with assistant return num_tokens def build_messages_and_calculate_token( diff --git a/requirements.txt b/requirements.txt index 13c937bf..c4661959 100644 --- a/requirements.txt +++ b/requirements.txt @@ -26,4 +26,8 @@ azure-ai-formrecognizer tables # azure identity related -azure.identity \ No newline at end of file +azure.identity + +# CI Fix Tool +tree-sitter-python +tree-sitter