llm-annotator 0.3.2__tar.gz → 0.3.4__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/PKG-INFO +1 -1
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/examples/sentiment.py +7 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/src/llm_annotator/annotator.py +59 -19
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/src/llm_annotator/utils.py +8 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/.gitignore +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/.pre-commit-config.yaml +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/LICENSE +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/Makefile +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/README.md +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/pyproject.toml +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/src/llm_annotator/__init__.py +0 -0
- {llm_annotator-0.3.2 → llm_annotator-0.3.4}/tests/conftest.py +0 -0
|
@@ -35,6 +35,12 @@ Classification:"""
|
|
|
35
35
|
|
|
36
36
|
def random_validitity(sample):
|
|
37
37
|
return random.random() < 0.5
|
|
38
|
+
|
|
39
|
+
def postprocess_fn(sample):
|
|
40
|
+
# Example postprocessing: strip whitespace from sentiment
|
|
41
|
+
if "sentiment" in sample and isinstance(sample["sentiment"], str):
|
|
42
|
+
sample["sentiment"] = sample["sentiment"].strip()
|
|
43
|
+
return sample
|
|
38
44
|
|
|
39
45
|
with Annotator(model="Qwen/Qwen2.5-0.5B-Instruct", max_model_len=4096, verbose=True) as anno:
|
|
40
46
|
ds = anno.annotate_dataset(
|
|
@@ -55,6 +61,7 @@ Classification:"""
|
|
|
55
61
|
sort_by_length=True, # Sort by prompt length for more efficient batching -- final dataset will be re-ordered to original
|
|
56
62
|
validate_fn=random_validitity,
|
|
57
63
|
num_retries_invalid=3,
|
|
64
|
+
postprocess_fn=postprocess_fn,
|
|
58
65
|
)
|
|
59
66
|
print(ds)
|
|
60
67
|
shutil.rmtree("outputs/sentiment-imdb-qwen")
|
|
@@ -17,7 +17,7 @@ from vllm import LLM, RequestOutput, SamplingParams
|
|
|
17
17
|
from vllm.distributed import destroy_distributed_environment, destroy_model_parallel
|
|
18
18
|
from vllm.sampling_params import StructuredOutputsParams
|
|
19
19
|
|
|
20
|
-
from llm_annotator.utils import ensure_returns_bool, remove_empty_jsonl_files, retry
|
|
20
|
+
from llm_annotator.utils import ensure_returns_bool, ensure_returns_dict, remove_empty_jsonl_files, retry
|
|
21
21
|
|
|
22
22
|
|
|
23
23
|
def destroy_model_on_error(func):
|
|
@@ -385,15 +385,30 @@ class Annotator:
|
|
|
385
385
|
"""Load and configure the tokenizer for the model.
|
|
386
386
|
|
|
387
387
|
Sets up the tokenizer with appropriate padding settings and ensures
|
|
388
|
-
a pad token is available.
|
|
388
|
+
a pad token is available. Some models (like recent Mistral) do not have
|
|
389
|
+
a chat template defined in their tokenizer, so we attempt to load
|
|
390
|
+
a processor instead in that case.
|
|
389
391
|
"""
|
|
390
|
-
|
|
391
|
-
|
|
392
|
-
|
|
392
|
+
try:
|
|
393
|
+
self.tokenizer = AutoTokenizer.from_pretrained(self.model)
|
|
394
|
+
if self.tokenizer.chat_template is None:
|
|
395
|
+
raise ValueError("No chat template defined in the tokenizer.")
|
|
396
|
+
except ValueError as tokenizer_error:
|
|
397
|
+
try:
|
|
398
|
+
self.tokenizer = self.AutoProcessor.from_pretrained(self.model)
|
|
399
|
+
if self.tokenizer.chat_template is None:
|
|
400
|
+
raise ValueError("No chat template defined in the processor either.") from tokenizer_error
|
|
401
|
+
except Exception as exc:
|
|
402
|
+
raise ValueError(
|
|
403
|
+
"Failed to load tokenizer or processor with chat template."
|
|
404
|
+
" Are you correctly using an instruct model?"
|
|
405
|
+
) from exc
|
|
406
|
+
else:
|
|
407
|
+
self.tokenizer.padding_side = "left"
|
|
393
408
|
|
|
394
|
-
|
|
395
|
-
|
|
396
|
-
|
|
409
|
+
if not self.tokenizer.pad_token_id:
|
|
410
|
+
self.tokenizer.pad_token = self.tokenizer.eos_token
|
|
411
|
+
self.tokenizer.pad_token_id = self.tokenizer.convert_tokens_to_ids(self.tokenizer.pad_token)
|
|
397
412
|
|
|
398
413
|
def _load_pipeline(self) -> None:
|
|
399
414
|
"""Load and initialize the vLLM pipeline for inference.
|
|
@@ -489,6 +504,7 @@ class Annotator:
|
|
|
489
504
|
sampling_params: SamplingParams,
|
|
490
505
|
task_prefix: str = "",
|
|
491
506
|
validate_fn: Callable | None = None,
|
|
507
|
+
postprocess_fn: Callable | None = None,
|
|
492
508
|
) -> list[dict[str, Any]]:
|
|
493
509
|
"""Process a batch of samples through the model.
|
|
494
510
|
|
|
@@ -506,23 +522,28 @@ class Annotator:
|
|
|
506
522
|
with keys as produced by `_process_output`: 1. the raw response ("{prefix}response"),
|
|
507
523
|
2. the finish reason ("{prefix}finish_reason"), 3. the number of tokens ("{prefix}num_tokens").
|
|
508
524
|
When an output schema was provided, also the parsed JSON fields and the "{prefix}valid_fields" key.
|
|
525
|
+
postprocess_fn: Optional function to postprocess each sample after annotation. Must return a modified
|
|
526
|
+
sample dictionary.
|
|
509
527
|
|
|
510
528
|
Returns:
|
|
511
529
|
List of processed output dictionaries for each sample in the batch.
|
|
512
530
|
"""
|
|
513
531
|
output_schema = sampling_params.structured_outputs.json if sampling_params.structured_outputs else None
|
|
514
532
|
outputs = self.pipe.generate(batch[f"{task_prefix}prompted"], sampling_params, use_tqdm=False)
|
|
515
|
-
results = [
|
|
516
|
-
self._process_output(output=outp, output_schema=output_schema, task_prefix=task_prefix) for outp in outputs
|
|
517
|
-
]
|
|
518
533
|
|
|
519
|
-
|
|
520
|
-
|
|
534
|
+
results = []
|
|
535
|
+
for outp in outputs:
|
|
536
|
+
res = self._process_output(output=outp, output_schema=output_schema, task_prefix=task_prefix)
|
|
537
|
+
if postprocess_fn:
|
|
538
|
+
res = ensure_returns_dict(postprocess_fn, res)
|
|
539
|
+
|
|
540
|
+
if validate_fn:
|
|
521
541
|
if f"{task_prefix}valid_fields" in res and res[f"{task_prefix}valid_fields"] is False:
|
|
522
542
|
is_valid = False
|
|
523
543
|
else:
|
|
524
544
|
is_valid = ensure_returns_bool(validate_fn, res)
|
|
525
545
|
res[f"{task_prefix}valid"] = is_valid
|
|
546
|
+
results.append(res)
|
|
526
547
|
|
|
527
548
|
if f"{task_prefix}valid_fields" in results[0]:
|
|
528
549
|
n_invalid = sum([1 for res in results if not res[f"{task_prefix}valid_fields"]])
|
|
@@ -576,7 +597,9 @@ class Annotator:
|
|
|
576
597
|
task_prefix: str = "",
|
|
577
598
|
sort_by_length: bool = False,
|
|
578
599
|
validate_fn: Callable | None = None,
|
|
579
|
-
|
|
600
|
+
postprocess_fn: Callable | None = None,
|
|
601
|
+
num_retries_invalid: int = 5,
|
|
602
|
+
keep_idx_column: bool = False,
|
|
580
603
|
) -> Dataset:
|
|
581
604
|
"""Annotate an entire dataset using the configured model and prompt.
|
|
582
605
|
|
|
@@ -631,11 +654,16 @@ class Annotator:
|
|
|
631
654
|
with keys as produced by `_process_output`: 1. the raw response ("{prefix}response"),
|
|
632
655
|
2. the finish reason ("{prefix}finish_reason"), 3. the number of tokens ("{prefix}num_tokens").
|
|
633
656
|
When an output schema was provided, also the parsed JSON fields and the "{prefix}valid_fields" key.
|
|
657
|
+
postprocess_fn: Optional function to postprocess each sample after annotation. Must return a modified
|
|
658
|
+
sample dictionary.
|
|
634
659
|
num_retries_invalid: Number of retries for samples that produce invalid outputs (when
|
|
635
660
|
a JSON schema is given and {prefix}valid_fields is False, or when a validate_fn is given
|
|
636
661
|
and it {prefix}valid is False).
|
|
662
|
+
keep_idx_column: Whether to keep the idx_column in the final dataset before uploading and returning.
|
|
663
|
+
|
|
664
|
+
Returns:
|
|
665
|
+
The concatenated dataset of all annotation results (JSON-invalid samples are NOT removed)
|
|
637
666
|
"""
|
|
638
|
-
# Verify shared_prompt_template
|
|
639
667
|
if prompt_template_prefix:
|
|
640
668
|
if prompt_template_prefix not in full_prompt_template:
|
|
641
669
|
raise ValueError(
|
|
@@ -763,7 +791,11 @@ class Annotator:
|
|
|
763
791
|
unit="batch",
|
|
764
792
|
):
|
|
765
793
|
results = self._process_batch(
|
|
766
|
-
batch=batch,
|
|
794
|
+
batch=batch,
|
|
795
|
+
sampling_params=sampling_params,
|
|
796
|
+
task_prefix=task_prefix,
|
|
797
|
+
validate_fn=validate_fn,
|
|
798
|
+
postprocess_fn=postprocess_fn,
|
|
767
799
|
)
|
|
768
800
|
|
|
769
801
|
if num_retries_invalid > 0:
|
|
@@ -846,9 +878,13 @@ class Annotator:
|
|
|
846
878
|
if new_hub_id and upload_every_n_samples > 0:
|
|
847
879
|
self.push_dir_to_hub(pdout, new_hub_id=new_hub_id)
|
|
848
880
|
|
|
849
|
-
return self._post_annotate(
|
|
881
|
+
return self._post_annotate(
|
|
882
|
+
pdout=pdout, idx_column=idx_column, new_hub_id=new_hub_id, keep_idx_column=keep_idx_column
|
|
883
|
+
)
|
|
850
884
|
|
|
851
|
-
def _post_annotate(
|
|
885
|
+
def _post_annotate(
|
|
886
|
+
self, *, pdout: Path, idx_column: str, new_hub_id: str | None = None, keep_idx_column: bool = False
|
|
887
|
+
) -> Dataset:
|
|
852
888
|
"""Clean up after annotation is complete.
|
|
853
889
|
|
|
854
890
|
Removes empty output files and performs any final cleanup operations.
|
|
@@ -856,6 +892,8 @@ class Annotator:
|
|
|
856
892
|
Args:
|
|
857
893
|
pdout: Output directory path to clean up.
|
|
858
894
|
new_hub_id: Optional Hugging Face dataset ID for uploads.
|
|
895
|
+
idx_column: Column name used as unique identifier.
|
|
896
|
+
keep_idx_column: Whether to keep the idx_column in the final dataset before uploading and returning
|
|
859
897
|
|
|
860
898
|
Returns:
|
|
861
899
|
The concatenated dataset of all annotation results (JSON-invalid samples are NOT removed)
|
|
@@ -865,7 +903,9 @@ class Annotator:
|
|
|
865
903
|
if pfin.stat().st_size > 0:
|
|
866
904
|
ds_parts.append(Dataset.from_json(str(pfin)))
|
|
867
905
|
|
|
868
|
-
ds: Dataset = concatenate_datasets(ds_parts).sort(idx_column)
|
|
906
|
+
ds: Dataset = concatenate_datasets(ds_parts).sort(idx_column)
|
|
907
|
+
if not keep_idx_column:
|
|
908
|
+
ds = ds.remove_columns([idx_column])
|
|
869
909
|
|
|
870
910
|
if new_hub_id:
|
|
871
911
|
ds.push_to_hub(new_hub_id, private=True)
|
|
@@ -155,3 +155,11 @@ def ensure_returns_bool(func, *args, **kwargs):
|
|
|
155
155
|
if not isinstance(result, bool):
|
|
156
156
|
raise TypeError(f"{func.__name__} should return a bool, got {type(result).__name__}")
|
|
157
157
|
return result
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
def ensure_returns_dict(func, *args, **kwargs):
|
|
161
|
+
"""Ensure that the given function returns a dict value. If not, raise a TypeError."""
|
|
162
|
+
result = func(*args, **kwargs)
|
|
163
|
+
if not isinstance(result, dict):
|
|
164
|
+
raise TypeError(f"{func.__name__} should return a dict, got {type(result).__name__}")
|
|
165
|
+
return result
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|