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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: llm-annotator
3
- Version: 0.3.2
3
+ Version: 0.3.4
4
4
  Summary: An easy-to-extend LLM annotator for robust, resumable data annotation.
5
5
  Author-email: Bram Vanroy <2779410+BramVanroy@users.noreply.github.com>
6
6
  License-Expression: Apache-2.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
- self.tokenizer = AutoTokenizer.from_pretrained(self.model)
391
-
392
- self.tokenizer.padding_side = "left"
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
- if not self.tokenizer.pad_token_id:
395
- self.tokenizer.pad_token = self.tokenizer.eos_token
396
- self.tokenizer.pad_token_id = self.tokenizer.convert_tokens_to_ids(self.tokenizer.pad_token)
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
- if validate_fn:
520
- for res in results:
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
- num_retries_invalid: int = 0,
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, sampling_params=sampling_params, task_prefix=task_prefix, validate_fn=validate_fn
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(pdout=pdout, idx_column=idx_column, new_hub_id=new_hub_id)
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(self, *, pdout: Path, idx_column: str, new_hub_id: str | None = None) -> Dataset:
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).remove_columns([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