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
41 changes: 23 additions & 18 deletions openseek/competition/LongContext-ICL-Annotation/src/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,20 +4,20 @@

# from method import build_prompt, select_examples, annotate

from method import build_prompt, select_examples
from method import build_prompt, select_examples, build_prompt_cot, build_prompt_by_task_type

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The imports build_prompt and build_prompt_cot are not used anywhere in main.py. Removing them keeps the imports clean and avoids confusion.

Suggested change
from method import build_prompt, select_examples, build_prompt_cot, build_prompt_by_task_type
from method import select_examples, build_prompt_by_task_type


from method import annotate_nvidia as annotate # For Nvidia GPU
# from method import annotate_ascend as annotate # For Huawei Ascend
# from method import annotate_nvidia as annotate # For Nvidia GPU
from method import annotate_ascend as annotate # For Huawei Ascend

TASK_FILES = {
1: './data/openseek-1_closest_integers.json',
2: './data/openseek-2_count_nouns_verbs.json',
3: './data/openseek-3_collatz_conjecture.json',
4: './data/openseek-4_conala_concat_strings.json',
5: './data/openseek-5_semeval_2018_task1_tweet_sadness_detection.json',
6: './data/openseek-6_mnli_same_genre_classification.json',
7: './data/openseek-7_jeopardy_answer_generation_all.json',
8: '../data/openseek-8_kernel_generation.json',
1: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-1_closest_integers.json',
2: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-2_count_nouns_verbs.json',
3: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-3_collatz_conjecture.json',
4: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-4_conala_concat_strings.json',
5: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-5_semeval_2018_task1_tweet_sadness_detection.json',
6: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-6_mnli_same_genre_classification.json',
7: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-7_jeopardy_answer_generation_all.json',
8: '/root/OpenSeek/openseek/competition/LongContext-ICL-Annotation/data/openseek-8_kernel_generation.json',
}

def parser_args():
Expand All @@ -30,7 +30,7 @@ def parser_args():
default='../outputs/',
help='Prefix path to save the evaluation logs.')
parser.add_argument('--tokenizer_path', type=str,
default='/share/project/wuhaiming/spaces/data_agent/OpenSeek-main/openseek/competition/LongContext-ICL-Annotation/src/Qwen3-4B')
default='/root/Qwen3-4B')
args = parser.parse_args()
return args

Expand All @@ -48,7 +48,7 @@ def evaluate(task_id:int,

task_name = task_dict['task_name']
task_description = task_dict['Definition'][0]
icl_examples = task_dict['examples'][:100]
icl_examples = task_dict['examples'][:50]
test_samples = task_dict['test_samples']

version = 1
Expand All @@ -62,26 +62,31 @@ def evaluate(task_id:int,
pass

examples_str = None
for test_sample in tqdm(test_samples, desc=f'Evaluation on Task {task_id}: {task_name}'):
for sample_idx, test_sample in enumerate(tqdm(test_samples, desc=f'Evaluation on Task {task_id}: {task_name}')):
test_record = dict()

test_sample_id = test_sample['id']
test_record['test_sample_id'] = test_sample_id


text2annotate = test_sample['input']
prompt = build_prompt(task_description, text2annotate)

# M03优化:使用任务分型Prompt路由系统
# 根据任务类型自动选择最合适的prompt策略
prompt = build_prompt_by_task_type(task_id, task_description, text2annotate)

if examples_str is None:
examples_str = select_examples(icl_examples, task_description, text2annotate)
# 使用混合上下文长度策略:传递task_id和sample_idx
examples_str = select_examples(icl_examples, task_description, text2annotate, task_id, sample_idx)
input_prompt = prompt.replace("[[EXAMPLES]]\n\n", examples_str+'\n\n')
Comment on lines 78 to 81

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The mixed context length strategy is currently broken because examples_str is cached as a single string initialized to None outside the loop. As a result, select_examples is only called once (at sample_idx == 0), meaning the 30k context length is used for all subsequent samples, and the 8k optimization is never triggered.

To fix this, we can repurpose examples_str as a dictionary cache to store both the large (30k/16k) and small (8k) versions of the few-shot examples, and select the appropriate one based on the current sample_idx.

Suggested change
if examples_str is None:
examples_str = select_examples(icl_examples, task_description, text2annotate)
# 使用混合上下文长度策略:传递task_id和sample_idx
examples_str = select_examples(icl_examples, task_description, text2annotate, task_id, sample_idx)
input_prompt = prompt.replace("[[EXAMPLES]]\n\n", examples_str+'\n\n')
if not isinstance(examples_str, dict):
examples_str = {}
cache_key = "large" if (task_id == 8 or sample_idx < 50) else "small"
if cache_key not in examples_str:
examples_str[cache_key] = select_examples(icl_examples, task_description, text2annotate, task_id, sample_idx, qwen_tokenizer)
input_prompt = prompt.replace("[[EXAMPLES]]\n\n", examples_str[cache_key] + '\n\n')


# tokenized_input = qwen_tokenizer(input_prompt, return_tensors="pt")
# if tokenized_input['input_ids'].shape[1] > max_input_length:
# test_record['prediction'] = None
# else:
# prediction = annotate(input_prompt)
# prediction = annotate(input_prompt, task_id)
# test_record['prediction'] = prediction
prediction = annotate(input_prompt)
prediction = annotate(input_prompt, task_id)
test_record['prediction'] = prediction
with open(output_file, 'a') as f:
f.write(json.dumps(test_record)+'\n')
Expand Down
Loading