Skip to content

Commit e7fe9ec

Browse files
committed
refactor: refactor llm generation code for localization
1 parent cd75efd commit e7fe9ec

3 files changed

Lines changed: 63 additions & 74 deletions

File tree

agentless/fl/FL.py

Lines changed: 49 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -211,24 +211,24 @@ class LLMFL(FL):
211211
Return just the locations.
212212
"""
213213

214-
def __init__(self, instance_id, structure, problem_statement, **kwargs):
214+
def __init__(
215+
self, instance_id, structure, problem_statement, model_name, backend, **kwargs
216+
):
215217
super().__init__(instance_id, structure, problem_statement)
216218
self.max_tokens = 300
219+
self.model_name = model_name
220+
self.backend = backend
217221

218222
def _parse_model_return_lines(self, content: str) -> list[str]:
219223
return content.strip().split("\n")
220224

221225
def localize(self, top_n=1, mock=False) -> tuple[list, list, list, any]:
226+
# lazy import, not sure if this is actually better?
227+
from agentless.util.api_requests import num_tokens_from_messages
228+
from agentless.util.model import make_model
222229

223230
found_files = []
224231

225-
# lazy import, not sure if this is actually better?
226-
from agentless.util.api_requests import (
227-
create_chatgpt_config,
228-
num_tokens_from_messages,
229-
request_chatgpt_engine,
230-
)
231-
232232
message = self.obtain_relevant_files_prompt.format(
233233
problem_statement=self.problem_statement,
234234
structure=show_project_structure(self.structure).strip(),
@@ -246,23 +246,16 @@ def localize(self, top_n=1, mock=False) -> tuple[list, list, list, any]:
246246
}
247247
return [], {"raw_output_loc": ""}, traj
248248

249-
config = create_chatgpt_config(
250-
message=message,
249+
model = make_model(
250+
model=self.model_name,
251+
backend=self.backend,
251252
max_tokens=self.max_tokens,
252253
temperature=0,
253254
batch_size=1,
254-
model="gpt-4o-2024-05-13", # use gpt-4o for now.
255255
)
256-
ret = request_chatgpt_engine(config)
257-
raw_output = ret.choices[0].message.content
258-
traj = {
259-
"prompt": message,
260-
"response": raw_output,
261-
"usage": {
262-
"prompt_tokens": ret.usage.prompt_tokens,
263-
"completion_tokens": ret.usage.completion_tokens,
264-
},
265-
}
256+
traj = model.codegen(message, num_samples=1)[0]
257+
traj["prompt"] = message
258+
raw_output = traj["response"]
266259
model_found_files = self._parse_model_return_lines(raw_output)
267260

268261
files, classes, functions = get_full_file_paths_and_classes_and_functions(
@@ -288,11 +281,8 @@ def localize(self, top_n=1, mock=False) -> tuple[list, list, list, any]:
288281
def localize_function_for_files(
289282
self, file_names, mock=False
290283
) -> tuple[list, dict, dict]:
291-
from agentless.util.api_requests import (
292-
create_chatgpt_config,
293-
num_tokens_from_messages,
294-
request_chatgpt_engine,
295-
)
284+
from agentless.util.api_requests import num_tokens_from_messages
285+
from agentless.util.model import make_model
296286

297287
files, classes, functions = get_full_file_paths_and_classes_and_functions(
298288
self.structure
@@ -339,23 +329,16 @@ def localize_function_for_files(
339329
}
340330
return [], {"raw_output_loc": ""}, traj
341331

342-
config = create_chatgpt_config(
343-
message=message,
332+
model = make_model(
333+
model=self.model_name,
334+
backend=self.backend,
344335
max_tokens=self.max_tokens,
345336
temperature=0,
346337
batch_size=1,
347-
model="gpt-4o-2024-05-13", # use gpt-4o for now.
348338
)
349-
ret = request_chatgpt_engine(config)
350-
raw_output = ret.choices[0].message.content
351-
traj = {
352-
"prompt": message,
353-
"response": raw_output,
354-
"usage": {
355-
"prompt_tokens": ret.usage.prompt_tokens,
356-
"completion_tokens": ret.usage.completion_tokens,
357-
},
358-
}
339+
traj = model.codegen(message, num_samples=1)[0]
340+
traj["prompt"] = message
341+
raw_output = traj["response"]
359342

360343
model_found_locs = extract_code_blocks(raw_output)
361344
model_found_locs_separated = extract_locs_for_files(
@@ -367,11 +350,8 @@ def localize_function_for_files(
367350
return model_found_locs_separated, {"raw_output_loc": raw_output}, traj
368351

369352
def localize_function_from_compressed_files(self, file_names, mock=False):
370-
from agentless.util.api_requests import (
371-
create_chatgpt_config,
372-
num_tokens_from_messages,
373-
request_chatgpt_engine,
374-
)
353+
from agentless.util.api_requests import num_tokens_from_messages
354+
from agentless.util.model import make_model
375355

376356
file_contents = get_repo_files(self.structure, file_names)
377357
compressed_file_contents = {
@@ -397,30 +377,23 @@ def localize_function_from_compressed_files(self, file_names, mock=False):
397377
"prompt": message,
398378
"usage": {
399379
"prompt_tokens": num_tokens_from_messages(
400-
message, "gpt-4o-2024-05-13"
380+
message,
381+
self.model_name,
401382
),
402383
},
403384
}
404385
return [], {"raw_output_loc": ""}, traj
405386

406-
config = create_chatgpt_config(
407-
message=message,
387+
model = make_model(
388+
model=self.model_name,
389+
backend=self.backend,
408390
max_tokens=self.max_tokens,
409391
temperature=0,
410392
batch_size=1,
411-
model="gpt-4o-2024-05-13", # use gpt-4o for now.
412393
)
413-
ret = request_chatgpt_engine(config)
414-
raw_output = ret.choices[0].message.content
415-
traj = {
416-
"prompt": message,
417-
"response": raw_output,
418-
"usage": {
419-
"prompt_tokens": ret.usage.prompt_tokens,
420-
"completion_tokens": ret.usage.completion_tokens,
421-
},
422-
}
423-
394+
traj = model.codegen(message, num_samples=1)[0]
395+
traj["prompt"] = message
396+
raw_output = traj["response"]
424397
model_found_locs = extract_code_blocks(raw_output)
425398
model_found_locs_separated = extract_locs_for_files(
426399
model_found_locs, file_names
@@ -450,11 +423,8 @@ def localize_line_from_coarse_function_locs(
450423
num_samples: int = 1,
451424
mock=False,
452425
):
453-
from agentless.util.api_requests import (
454-
create_chatgpt_config,
455-
num_tokens_from_messages,
456-
request_chatgpt_engine,
457-
)
426+
from agentless.util.api_requests import num_tokens_from_messages
427+
from agentless.util.model import make_model
458428

459429
file_contents = get_repo_files(self.structure, file_names)
460430
topn_content, file_loc_intervals = construct_topn_file_context(
@@ -488,21 +458,28 @@ def localize_line_from_coarse_function_locs(
488458
},
489459
}
490460
return [], {"raw_output_loc": ""}, traj
491-
config = create_chatgpt_config(
492-
message=message,
461+
462+
model = make_model(
463+
model=self.model_name,
464+
backend=self.backend,
493465
max_tokens=self.max_tokens,
494466
temperature=temperature,
495467
batch_size=num_samples,
496-
model="gpt-4o-2024-05-13", # use gpt-4o for now.
497468
)
498-
ret = request_chatgpt_engine(config)
499-
raw_outputs = [choice.message.content for choice in ret.choices]
469+
raw_trajs = model.codegen(message, num_samples=num_samples)
470+
471+
# Merge trajectories
472+
raw_outputs = [raw_traj["response"] for raw_traj in raw_trajs]
500473
traj = {
501474
"prompt": message,
502475
"response": raw_outputs,
503-
"usage": {
504-
"prompt_tokens": ret.usage.prompt_tokens,
505-
"completion_tokens": ret.usage.completion_tokens,
476+
"usage": { # merge token usage
477+
"completion_tokens": sum(
478+
raw_traj["usage"]["completion_tokens"] for raw_traj in raw_trajs
479+
),
480+
"prompt_tokens": sum(
481+
raw_traj["usage"]["prompt_tokens"] for raw_traj in raw_trajs
482+
),
506483
},
507484
}
508485
model_found_locs_separated_in_samples = []

agentless/fl/localize.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,8 @@ def localize(args):
7373
d["instance_id"],
7474
structure,
7575
problem_statement,
76+
args.model,
77+
args.backend,
7678
)
7779
found_files, additional_artifact_loc_file, file_traj = fl.localize(
7880
mock=args.mock
@@ -101,6 +103,8 @@ def localize(args):
101103
d["instance_id"],
102104
structure,
103105
problem_statement,
106+
args.model,
107+
args.backend,
104108
)
105109

106110
additional_artifact_loc_related = []
@@ -128,6 +132,8 @@ def localize(args):
128132
instance_id,
129133
structure,
130134
problem_statement,
135+
args.model,
136+
args.backend,
131137
)
132138
coarse_found_locs = {}
133139
for i, pred_file in enumerate(pred_files):
@@ -264,6 +270,10 @@ def main():
264270
parser.add_argument(
265271
"--mock", action="store_true", help="Mock run to compute prompt tokens."
266272
)
273+
parser.add_argument(
274+
"--model", type=str, default="gpt-4o-2024-05-13", choices=["gpt-4o-2024-05-13"]
275+
)
276+
parser.add_argument("--backend", type=str, default="openai", choices=["openai"])
267277

268278
args = parser.parse_args()
269279

agentless/util/model.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@ def __init__(
1919
self.max_new_tokens = max_new_tokens
2020

2121
@abstractmethod
22-
def codegen(self, message: str, num_samples: int = 1) -> List[str]:
22+
def codegen(self, message: str, num_samples: int = 1) -> List[dict]:
2323
pass
2424

2525
@abstractmethod
@@ -37,7 +37,9 @@ class OpenAIChatDecoder(DecoderBase):
3737
def __init__(self, name: str, **kwargs) -> None:
3838
super().__init__(name, **kwargs)
3939

40-
def codegen(self, message: str, num_samples: int = 1) -> List[str]:
40+
def codegen(self, message: str, num_samples: int = 1) -> List[dict]:
41+
if self.temperature == 0:
42+
assert num_samples == 1
4143
batch_size = min(self.batch_size, num_samples)
4244

4345
config = create_chatgpt_config(

0 commit comments

Comments
 (0)