caption-flow 0.5.2__tar.gz → 0.5.3__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.
Files changed (58) hide show
  1. {caption_flow-0.5.2/src/caption_flow.egg-info → caption_flow-0.5.3}/PKG-INFO +3 -3
  2. {caption_flow-0.5.2 → caption_flow-0.5.3}/README.md +2 -2
  3. {caption_flow-0.5.2 → caption_flow-0.5.3}/pyproject.toml +1 -1
  4. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/__init__.py +1 -1
  5. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/processors/huggingface.py +118 -20
  6. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/vllm_config.py +13 -1
  7. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/workers/caption.py +49 -2
  8. {caption_flow-0.5.2 → caption_flow-0.5.3/src/caption_flow.egg-info}/PKG-INFO +3 -3
  9. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_processors.py +98 -0
  10. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_vllm_config.py +16 -0
  11. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_worker_caption.py +91 -0
  12. {caption_flow-0.5.2 → caption_flow-0.5.3}/LICENSE +0 -0
  13. {caption_flow-0.5.2 → caption_flow-0.5.3}/setup.cfg +0 -0
  14. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/cli.py +0 -0
  15. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/models.py +0 -0
  16. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/monitor.py +0 -0
  17. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/orchestrator.py +0 -0
  18. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/processors/__init__.py +0 -0
  19. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/processors/base.py +0 -0
  20. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/processors/local_filesystem.py +0 -0
  21. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/processors/webdataset.py +0 -0
  22. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/storage/__init__.py +0 -0
  23. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/storage/exporter.py +0 -0
  24. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/storage/manager.py +0 -0
  25. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/__init__.py +0 -0
  26. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/auth.py +0 -0
  27. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/caption_utils.py +0 -0
  28. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/certificates.py +0 -0
  29. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/checkpoint_tracker.py +0 -0
  30. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/chunk_tracker.py +0 -0
  31. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/image_processor.py +0 -0
  32. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/json_utils.py +0 -0
  33. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/utils/prompt_template.py +0 -0
  34. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/viewer.py +0 -0
  35. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/workers/base.py +0 -0
  36. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow/workers/data.py +0 -0
  37. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow.egg-info/SOURCES.txt +0 -0
  38. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow.egg-info/dependency_links.txt +0 -0
  39. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow.egg-info/entry_points.txt +0 -0
  40. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow.egg-info/requires.txt +0 -0
  41. {caption_flow-0.5.2 → caption_flow-0.5.3}/src/caption_flow.egg-info/top_level.txt +0 -0
  42. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_caption_utils.py +0 -0
  43. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_certificates.py +0 -0
  44. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_cli.py +0 -0
  45. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_config_reload.py +0 -0
  46. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_duplicate_job_assignments.py +0 -0
  47. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_exporter.py +0 -0
  48. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_fix_verification.py +0 -0
  49. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_huggingface_ranges.py +0 -0
  50. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_json_utils.py +0 -0
  51. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_main.py +0 -0
  52. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_monitor.py +0 -0
  53. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_range_level_distribution.py +0 -0
  54. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_storage_components.py +0 -0
  55. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_viewer.py +0 -0
  56. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_webdataset_ranges.py +0 -0
  57. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_worker_reconnection_complete.py +0 -0
  58. {caption_flow-0.5.2 → caption_flow-0.5.3}/tests/test_worker_reconnection_sequence.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: caption-flow
3
- Version: 0.5.2
3
+ Version: 0.5.3
4
4
  Summary: Self-contained distributed community captioning system
5
5
  Author-email: bghira <bghira@users.github.com>
6
6
  License: MIT
@@ -76,7 +76,7 @@ a fast websocket-based orchestrator paired with lightweight gpu workers achieves
76
76
  ```bash
77
77
  python -m venv .venv
78
78
  source .venv/bin/activate # windows: .venv\Scripts\activate
79
- pip install caption-flow
79
+ pip install "caption-flow[vllm]"
80
80
  ```
81
81
 
82
82
  ## quickstart (single box)
@@ -240,7 +240,7 @@ PRs welcome. keep it simple and fast.
240
240
 
241
241
  To contribute compute to a cluster:
242
242
 
243
- 1. Install caption-flow: `pip install caption-flow`
243
+ 1. Install caption-flow: `pip install "caption-flow[vllm]"`
244
244
  2. Get a worker token from the project maintainer
245
245
  3. Run: `caption-flow worker --server wss://project.domain.com:8765 --token YOUR_TOKEN`
246
246
 
@@ -21,7 +21,7 @@ a fast websocket-based orchestrator paired with lightweight gpu workers achieves
21
21
  ```bash
22
22
  python -m venv .venv
23
23
  source .venv/bin/activate # windows: .venv\Scripts\activate
24
- pip install caption-flow
24
+ pip install "caption-flow[vllm]"
25
25
  ```
26
26
 
27
27
  ## quickstart (single box)
@@ -185,7 +185,7 @@ PRs welcome. keep it simple and fast.
185
185
 
186
186
  To contribute compute to a cluster:
187
187
 
188
- 1. Install caption-flow: `pip install caption-flow`
188
+ 1. Install caption-flow: `pip install "caption-flow[vllm]"`
189
189
  2. Get a worker token from the project maintainer
190
190
  3. Run: `caption-flow worker --server wss://project.domain.com:8765 --token YOUR_TOKEN`
191
191
 
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "caption-flow"
3
- version = "0.5.2"
3
+ version = "0.5.3"
4
4
  description = "Self-contained distributed community captioning system"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.11,<3.14"
@@ -1,6 +1,6 @@
1
1
  """CaptionFlow - Distributed community captioning system."""
2
2
 
3
- __version__ = "0.5.2"
3
+ __version__ = "0.5.3"
4
4
 
5
5
  from .monitor import Monitor
6
6
  from .orchestrator import Orchestrator
@@ -32,6 +32,8 @@ from .base import OrchestratorProcessor, ProcessorConfig, WorkerProcessor, WorkR
32
32
  logger = logging.getLogger(__name__)
33
33
  logger.setLevel(os.environ.get("CAPTIONFLOW_LOG_LEVEL", "INFO").upper())
34
34
 
35
+ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp", ".gif", ".tiff", ".tif"}
36
+
35
37
 
36
38
  def log_memory(location: str):
37
39
  """Log memory usage at specific location."""
@@ -244,6 +246,10 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
244
246
  return match.group(1)
245
247
  return url.split("/")[-1]
246
248
 
249
+ def _is_image_file(self, filename: str) -> bool:
250
+ """Return whether a Hugging Face dataset file is a supported raw image."""
251
+ return Path(filename).suffix.lower() in IMAGE_EXTENSIONS
252
+
247
253
  def _get_data_files_from_builder(self) -> List[str]:
248
254
  """Get data files using dataset builder with minimal memory usage."""
249
255
  # Load builder to get correct file structure
@@ -317,21 +323,29 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
317
323
  token=self.token,
318
324
  )
319
325
 
320
- # Read only metadata
321
- metadata = pq.read_metadata(local_path)
322
- size = metadata.num_rows
326
+ file_type = "parquet"
327
+ try:
328
+ # Read only metadata for parquet-backed datasets.
329
+ metadata = pq.read_metadata(local_path)
330
+ size = metadata.num_rows
331
+ except Exception:
332
+ if not self._is_image_file(filename):
333
+ raise
334
+ file_type = "image"
335
+ size = 1
323
336
 
324
337
  self.shard_info[i] = {
325
338
  "shard_id": i,
326
339
  "file_url": file_url,
327
340
  "filename": filename,
341
+ "file_type": file_type,
328
342
  "start_offset": cumulative_offset,
329
343
  "size": size,
330
344
  "end_offset": cumulative_offset + size - 1,
331
345
  }
332
346
 
333
347
  cumulative_offset += size
334
- logger.info(f"Shard {i} ({filename}): {size} rows")
348
+ logger.info(f"Shard {i} ({filename}): {size} {file_type} item(s)")
335
349
 
336
350
  except Exception as e:
337
351
  logger.error(f"Failed to discover shard {i}: {e}")
@@ -371,6 +385,14 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
371
385
  return shard_id, local_index
372
386
  raise ValueError(f"Global index {global_index} not found in any shard")
373
387
 
388
+ def _get_shards_for_range(self, start_index: int, end_index: int) -> List[int]:
389
+ """Get all shard IDs intersecting an absolute index range."""
390
+ shard_ids = []
391
+ for shard_id, sinfo in self.shard_info.items():
392
+ if sinfo["start_offset"] <= end_index and sinfo["end_offset"] >= start_index:
393
+ shard_ids.append(shard_id)
394
+ return shard_ids
395
+
374
396
  def _restore_state(self, storage: StorageManager) -> None:
375
397
  """Restore state from chunk tracker and synchronize with storage."""
376
398
  logger.debug("Restoring state from chunk tracker and synchronizing with storage")
@@ -468,9 +490,10 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
468
490
  current_index = chunk_index * self.chunk_size
469
491
  chunk_size = min(self.chunk_size, self.total_items - current_index)
470
492
 
471
- # Find shard for this chunk
493
+ # Find all shards touched by this chunk/range set
472
494
  shard_id, local_idx = self._get_shard_for_index(current_index)
473
495
  shard_name = Path(self.shard_info[shard_id]["filename"]).stem
496
+ shard_ids = self._get_shards_for_range(ranges[0][0], ranges[-1][1])
474
497
 
475
498
  # Create unique unit ID that includes range info
476
499
  range_suffix = f"r{len(ranges)}" # r2 = 2 ranges, etc.
@@ -497,8 +520,9 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
497
520
  "unprocessed_ranges": ranges,
498
521
  "range_based": True,
499
522
  "is_split_unit": True, # Flag to indicate this is a split from larger chunk
500
- "shard_ids": [shard_id],
523
+ "shard_ids": shard_ids,
501
524
  "data_files": self.data_files,
525
+ "shard_info": self.shard_info,
502
526
  },
503
527
  metadata={
504
528
  "dataset": self.dataset_name,
@@ -519,9 +543,10 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
519
543
 
520
544
  chunk_size = min(self.chunk_size, self.total_items - current_index)
521
545
 
522
- # Find shard for this chunk
546
+ # Find all shards touched by this chunk
523
547
  shard_id, local_idx = self._get_shard_for_index(current_index)
524
548
  shard_name = Path(self.shard_info[shard_id]["filename"]).stem
549
+ shard_ids = self._get_shards_for_range(current_index, current_index + chunk_size - 1)
525
550
 
526
551
  # Calculate RELATIVE chunk index within the shard
527
552
  job_id_obj = JobId(
@@ -556,13 +581,6 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
556
581
  # Calculate actual unprocessed items and total work to be assigned
557
582
  unprocessed_items = sum(end - start + 1 for start, end in unprocessed_ranges)
558
583
 
559
- # Skip assignment if there are very few unprocessed items (< 10 items)
560
- if unprocessed_items < 10:
561
- logger.debug(
562
- f"Chunk {unit_id} has only {unprocessed_items} unprocessed items, skipping assignment"
563
- )
564
- return None
565
-
566
584
  # Create work unit that represents ONLY the unprocessed ranges
567
585
  # This is the key fix: don't assign the full chunk, assign only unprocessed parts
568
586
  unit = WorkUnit(
@@ -579,8 +597,9 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
579
597
  "actual_work_size": unprocessed_items, # NEW: actual work to be done
580
598
  "unprocessed_ranges": unprocessed_ranges, # The specific ranges to process
581
599
  "range_based": True, # NEW: flag to indicate this is range-based
582
- "shard_ids": [shard_id],
600
+ "shard_ids": shard_ids,
583
601
  "data_files": self.data_files,
602
+ "shard_info": self.shard_info,
584
603
  },
585
604
  metadata={
586
605
  "dataset": self.dataset_name,
@@ -629,7 +648,7 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
629
648
  )
630
649
  break
631
650
  # Get shard info for proper unit_id
632
- current_index = self.current_chunk_index
651
+ current_index = self.current_chunk_index * self.chunk_size
633
652
  if current_index < self.total_items:
634
653
  shard_id, _ = self._get_shard_for_index(current_index)
635
654
  shard_name = Path(self.shard_info[shard_id]["filename"]).stem
@@ -980,6 +999,9 @@ class HuggingFaceDatasetOrchestratorProcessor(OrchestratorProcessor):
980
999
 
981
1000
  for start_idx, end_idx in ranges:
982
1001
  self.chunk_tracker.mark_items_processed(result.chunk_id, start_idx, end_idx)
1002
+ elif "_item_index" in result.metadata and result.metadata["_item_index"] is not None:
1003
+ item_index = int(result.metadata["_item_index"])
1004
+ self.chunk_tracker.mark_items_processed(result.chunk_id, item_index, item_index)
983
1005
 
984
1006
  return base_result
985
1007
 
@@ -1070,6 +1092,10 @@ class HuggingFaceDatasetWorkerProcessor(WorkerProcessor):
1070
1092
  return match.group(1)
1071
1093
  return url.split("/")[-1]
1072
1094
 
1095
+ def _is_image_file(self, filename: str) -> bool:
1096
+ """Return whether a Hugging Face dataset file is a supported raw image."""
1097
+ return Path(filename).suffix.lower() in IMAGE_EXTENSIONS
1098
+
1073
1099
  def _create_dummy_image(self, index: int, metadata: Dict[str, Any]) -> Image.Image:
1074
1100
  """Create a dummy image."""
1075
1101
  color = (0, 0, 0)
@@ -1091,13 +1117,41 @@ class HuggingFaceDatasetWorkerProcessor(WorkerProcessor):
1091
1117
  )
1092
1118
  shard_ids = unit.data.get("shard_ids", [])
1093
1119
  data_files = unit.data.get("data_files", [])
1120
+ unit_shard_info = unit.data.get("shard_info", {})
1121
+ unit_shard_info = {int(k): v for k, v in unit_shard_info.items()}
1094
1122
 
1095
1123
  logger.info(f"Processing unit {unit.unit_id} with ranges: {unprocessed_ranges}")
1096
1124
 
1097
1125
  # Build shard info from provided data files (no dataset builder needed)
1098
1126
  shard_info = {}
1099
1127
 
1100
- if data_files:
1128
+ if unit_shard_info:
1129
+ for shard_id in shard_ids:
1130
+ shard_id = int(shard_id)
1131
+ if shard_id not in unit_shard_info:
1132
+ continue
1133
+ metadata = unit_shard_info[shard_id]
1134
+ filename = metadata["filename"]
1135
+ shard_path = self._get_shard_path(dataset_name, filename)
1136
+ file_type = metadata.get("file_type") or (
1137
+ "image" if self._is_image_file(filename) else "parquet"
1138
+ )
1139
+
1140
+ shard_info[shard_id] = {
1141
+ "path": shard_path,
1142
+ "filename": filename,
1143
+ "file_type": file_type,
1144
+ "start_offset": metadata["start_offset"],
1145
+ "end_offset": metadata["end_offset"],
1146
+ "size": metadata["size"],
1147
+ "metadata": None,
1148
+ }
1149
+
1150
+ if file_type == "parquet":
1151
+ shard_info[shard_id]["metadata"] = pq.read_metadata(shard_path)
1152
+
1153
+ elif data_files:
1154
+ shard_ids = {int(shard_id) for shard_id in shard_ids}
1101
1155
  # Use provided data files
1102
1156
  for i, file_url in enumerate(data_files):
1103
1157
  if i in shard_ids:
@@ -1110,6 +1164,8 @@ class HuggingFaceDatasetWorkerProcessor(WorkerProcessor):
1110
1164
 
1111
1165
  shard_info[i] = {
1112
1166
  "path": shard_path,
1167
+ "filename": filename,
1168
+ "file_type": "parquet",
1113
1169
  "start_offset": 0, # Will be set below
1114
1170
  "end_offset": 0, # Will be set below
1115
1171
  "size": size,
@@ -1152,6 +1208,49 @@ class HuggingFaceDatasetWorkerProcessor(WorkerProcessor):
1152
1208
  # Process items shard by shard
1153
1209
  for shard_id, idx_pairs in indices_by_shard.items():
1154
1210
  shard_path = shard_info[shard_id]["path"]
1211
+ file_type = shard_info[shard_id].get("file_type", "parquet")
1212
+ shard_name = Path(shard_info[shard_id]["filename"]).stem
1213
+
1214
+ if file_type == "image":
1215
+ for global_idx, local_idx in idx_pairs:
1216
+ try:
1217
+ with Image.open(shard_path) as img:
1218
+ image = img.copy()
1219
+
1220
+ chunk_index = unit.metadata["chunk_index"]
1221
+ job_id_obj = JobId(
1222
+ shard_id=shard_name,
1223
+ chunk_id=str(chunk_index),
1224
+ sample_id=str(global_idx),
1225
+ )
1226
+ job_id = job_id_obj.get_sample_str()
1227
+
1228
+ clean_metadata = {
1229
+ "_item_index": global_idx,
1230
+ "_chunk_relative_index": global_idx - start_index,
1231
+ "_job_id": job_id,
1232
+ "_shard_id": shard_id,
1233
+ "_local_index": local_idx,
1234
+ "_filename": shard_info[shard_id]["filename"],
1235
+ "_url": None,
1236
+ "_mock": self.mock_results,
1237
+ }
1238
+
1239
+ yield {
1240
+ "image": image,
1241
+ "item_key": str(global_idx),
1242
+ "item_index": global_idx,
1243
+ "metadata": clean_metadata,
1244
+ "job_id": job_id,
1245
+ "_processed_indices": processed_indices,
1246
+ }
1247
+
1248
+ processed_indices.append(global_idx)
1249
+
1250
+ except Exception as e:
1251
+ logger.error(f"Error processing image file at index {global_idx}: {e}")
1252
+
1253
+ continue
1155
1254
 
1156
1255
  # Process in batches to avoid loading entire table
1157
1256
  batch_size = 100
@@ -1325,11 +1424,10 @@ class HuggingFaceDatasetWorkerProcessor(WorkerProcessor):
1325
1424
 
1326
1425
  # Build job ID
1327
1426
  chunk_index = unit.metadata["chunk_index"]
1328
- shard_name = unit.metadata["shard_name"]
1329
1427
  job_id_obj = JobId(
1330
1428
  shard_id=shard_name,
1331
1429
  chunk_id=str(chunk_index),
1332
- sample_id=str(local_idx),
1430
+ sample_id=str(global_idx),
1333
1431
  )
1334
1432
  job_id = job_id_obj.get_sample_str()
1335
1433
 
@@ -1362,7 +1460,7 @@ class HuggingFaceDatasetWorkerProcessor(WorkerProcessor):
1362
1460
  "_processed_indices": processed_indices,
1363
1461
  }
1364
1462
 
1365
- processed_indices.append(local_idx)
1463
+ processed_indices.append(global_idx)
1366
1464
 
1367
1465
  except Exception as e:
1368
1466
  logger.error(f"Error processing item at index {global_idx}: {e}")
@@ -31,6 +31,9 @@ class VLLMConfigManager:
31
31
  "enforce_eager",
32
32
  "limit_mm_per_prompt",
33
33
  "disable_mm_preprocessor_cache",
34
+ "mm_processor_cache_gb",
35
+ "mm_processor_cache_type",
36
+ "mm_shm_cache_max_object_size_mb",
34
37
  }
35
38
 
36
39
  # Fields that can be updated without reload
@@ -115,7 +118,7 @@ class VLLMConfigManager:
115
118
 
116
119
  def get_vllm_init_params(self, vllm_config: Dict[str, Any]) -> Dict[str, Any]:
117
120
  """Extract vLLM initialization parameters from config."""
118
- return {
121
+ params = {
119
122
  "model": vllm_config["model"],
120
123
  "trust_remote_code": True,
121
124
  "tensor_parallel_size": vllm_config.get("tensor_parallel_size", 1),
@@ -126,6 +129,15 @@ class VLLMConfigManager:
126
129
  "limit_mm_per_prompt": vllm_config.get("limit_mm_per_prompt", {"image": 1}),
127
130
  "disable_mm_preprocessor_cache": vllm_config.get("disable_mm_preprocessor_cache", True),
128
131
  }
132
+ for optional_key in (
133
+ "mm_processor_cache_gb",
134
+ "mm_processor_cache_type",
135
+ "mm_shm_cache_max_object_size_mb",
136
+ ):
137
+ if vllm_config.get(optional_key) is not None:
138
+ params[optional_key] = vllm_config[optional_key]
139
+
140
+ return params
129
141
 
130
142
  def requires_tokenizer_reload(
131
143
  self, old_config: Optional[Dict[str, Any]], new_config: Dict[str, Any]
@@ -5,6 +5,7 @@ import os
5
5
  os.environ["VLLM_ENABLE_V1_MULTIPROCESSING"] = "0"
6
6
 
7
7
  import asyncio
8
+ import inspect
8
9
  import json
9
10
  import logging
10
11
  import time
@@ -66,6 +67,39 @@ class MultiStageVLLMManager:
66
67
  self.tokenizers: Dict[str, Any] = {} # model_name -> tokenizer
67
68
  self.sampling_params: Dict[str, Any] = {} # stage_name -> SamplingParams
68
69
 
70
+ @staticmethod
71
+ def _filter_engine_args(params: Dict[str, Any], engine_args_cls: Any) -> Dict[str, Any]:
72
+ """Filter vLLM init kwargs to those supported by the installed EngineArgs."""
73
+ signature = inspect.signature(engine_args_cls)
74
+ supported_params = set(signature.parameters)
75
+ filtered = params.copy()
76
+
77
+ disable_mm_cache = filtered.pop("disable_mm_preprocessor_cache", None)
78
+ if (
79
+ disable_mm_cache is True
80
+ and "disable_mm_preprocessor_cache" not in supported_params
81
+ and "mm_processor_cache_gb" in supported_params
82
+ and "mm_processor_cache_gb" not in filtered
83
+ ):
84
+ filtered["mm_processor_cache_gb"] = 0
85
+ logger.info(
86
+ "Mapped disable_mm_preprocessor_cache=True to mm_processor_cache_gb=0 "
87
+ "for this vLLM version"
88
+ )
89
+ elif "disable_mm_preprocessor_cache" in supported_params and disable_mm_cache is not None:
90
+ filtered["disable_mm_preprocessor_cache"] = disable_mm_cache
91
+
92
+ unsupported = sorted(key for key in filtered if key not in supported_params)
93
+ if unsupported:
94
+ logger.warning(
95
+ "Dropping unsupported vLLM initialization parameter(s): %s",
96
+ ", ".join(unsupported),
97
+ )
98
+ for key in unsupported:
99
+ filtered.pop(key, None)
100
+
101
+ return filtered
102
+
69
103
  def load_model(self, model_name: str, stage: ProcessingStage, base_config: Dict[str, Any]):
70
104
  """Load a model if not already loaded."""
71
105
  if model_name in self.models:
@@ -74,6 +108,7 @@ class MultiStageVLLMManager:
74
108
 
75
109
  from transformers import AutoProcessor, AutoTokenizer
76
110
  from vllm import LLM
111
+ from vllm.engine.arg_utils import EngineArgs
77
112
 
78
113
  logger.info(f"Loading model {model_name} for stage {stage.name}")
79
114
 
@@ -113,8 +148,15 @@ class MultiStageVLLMManager:
113
148
  "disable_mm_preprocessor_cache", True
114
149
  ),
115
150
  }
116
-
117
- self.models[model_name] = LLM(**vllm_params)
151
+ for optional_key in (
152
+ "mm_processor_cache_gb",
153
+ "mm_processor_cache_type",
154
+ "mm_shm_cache_max_object_size_mb",
155
+ ):
156
+ if model_config.get(optional_key) is not None:
157
+ vllm_params[optional_key] = model_config[optional_key]
158
+
159
+ self.models[model_name] = LLM(**self._filter_engine_args(vllm_params, EngineArgs))
118
160
  logger.info(f"Model {model_name} loaded successfully")
119
161
 
120
162
  def create_sampling_params(self, stage: ProcessingStage, base_sampling: Dict[str, Any]):
@@ -544,6 +586,11 @@ class CaptionWorker(BaseWorker):
544
586
  "disable_mm_preprocessor_cache", True
545
587
  ),
546
588
  "limit_mm_per_prompt": self.vllm_config.get("limit_mm_per_prompt", {"image": 1}),
589
+ "mm_processor_cache_gb": self.vllm_config.get("mm_processor_cache_gb"),
590
+ "mm_processor_cache_type": self.vllm_config.get("mm_processor_cache_type"),
591
+ "mm_shm_cache_max_object_size_mb": self.vllm_config.get(
592
+ "mm_shm_cache_max_object_size_mb"
593
+ ),
547
594
  }
548
595
 
549
596
  base_sampling = self.vllm_config.get("sampling", {})
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: caption-flow
3
- Version: 0.5.2
3
+ Version: 0.5.3
4
4
  Summary: Self-contained distributed community captioning system
5
5
  Author-email: bghira <bghira@users.github.com>
6
6
  License: MIT
@@ -76,7 +76,7 @@ a fast websocket-based orchestrator paired with lightweight gpu workers achieves
76
76
  ```bash
77
77
  python -m venv .venv
78
78
  source .venv/bin/activate # windows: .venv\Scripts\activate
79
- pip install caption-flow
79
+ pip install "caption-flow[vllm]"
80
80
  ```
81
81
 
82
82
  ## quickstart (single box)
@@ -240,7 +240,7 @@ PRs welcome. keep it simple and fast.
240
240
 
241
241
  To contribute compute to a cluster:
242
242
 
243
- 1. Install caption-flow: `pip install caption-flow`
243
+ 1. Install caption-flow: `pip install "caption-flow[vllm]"`
244
244
  2. Get a worker token from the project maintainer
245
245
  3. Run: `caption-flow worker --server wss://project.domain.com:8765 --token YOUR_TOKEN`
246
246
 
@@ -455,6 +455,42 @@ class TestHuggingFaceDatasetProcessors(ProcessorTestBase):
455
455
  finally:
456
456
  orchestrator.cleanup()
457
457
 
458
+ def test_raw_image_file_discovery(self, orchestrator_config, storage_manager, temp_dir):
459
+ """Test Hugging Face repos that expose raw image files instead of parquet shards."""
460
+ image_paths = []
461
+ data_files = []
462
+ for i in range(3):
463
+ image_path = temp_dir / f"image-{i}.jpg"
464
+ Image.new("RGB", (10, 10), color=(i, i, i)).save(image_path)
465
+ image_paths.append(str(image_path))
466
+ data_files.append(f"image-{i}.jpg")
467
+
468
+ with patch(
469
+ "caption_flow.processors.huggingface.HuggingFaceDatasetOrchestratorProcessor._get_data_files_from_builder",
470
+ return_value=data_files,
471
+ ):
472
+ with patch(
473
+ "caption_flow.processors.huggingface.hf_hub_download",
474
+ side_effect=image_paths,
475
+ ):
476
+ orchestrator = HuggingFaceDatasetOrchestratorProcessor()
477
+ try:
478
+ orchestrator.initialize(orchestrator_config, storage_manager)
479
+
480
+ assert orchestrator.total_items == 3
481
+ assert len(orchestrator.shard_info) == 3
482
+ assert all(
483
+ shard["file_type"] == "image"
484
+ for shard in orchestrator.shard_info.values()
485
+ )
486
+
487
+ unit = orchestrator._create_work_unit(0)
488
+ assert unit is not None
489
+ assert unit.unit_size == 3
490
+ assert unit.data["shard_ids"] == [0, 1, 2]
491
+ finally:
492
+ orchestrator.cleanup()
493
+
458
494
  # For test_worker_processing_with_ranges, replace with:
459
495
  def test_worker_processing_with_ranges(self, worker_config, temp_dir):
460
496
  """Test worker processing with specific unprocessed ranges."""
@@ -517,6 +553,68 @@ class TestHuggingFaceDatasetProcessors(ProcessorTestBase):
517
553
  processed_indices = context.get("_processed_indices", [])
518
554
  assert len(processed_indices) == 20
519
555
 
556
+ def test_worker_processing_raw_image_files(self, worker_config, temp_dir):
557
+ """Test worker processing for raw image files from Hugging Face repos."""
558
+ worker = HuggingFaceDatasetWorkerProcessor()
559
+ worker.gpu_id = 0
560
+
561
+ image_paths = []
562
+ for i in range(2):
563
+ image_path = temp_dir / f"raw-{i}.webp"
564
+ Image.new("RGB", (10, 10), color=(i, i, i)).save(image_path)
565
+ image_paths.append(str(image_path))
566
+
567
+ with patch(
568
+ "caption_flow.processors.huggingface.hf_hub_download",
569
+ side_effect=image_paths,
570
+ ):
571
+ worker.initialize(worker_config)
572
+
573
+ unit = WorkUnit(
574
+ unit_id="raw-0:chunk:0",
575
+ chunk_id="raw-0:chunk:0",
576
+ source_id="raw-0",
577
+ unit_size=2,
578
+ data={
579
+ "dataset_name": "test/dataset",
580
+ "config": "default",
581
+ "split": "train",
582
+ "start_index": 0,
583
+ "chunk_size": 2,
584
+ "unprocessed_ranges": [(0, 1)],
585
+ "shard_ids": [0, 1],
586
+ "data_files": ["raw-0.webp", "raw-1.webp"],
587
+ "shard_info": {
588
+ 0: {
589
+ "filename": "raw-0.webp",
590
+ "file_type": "image",
591
+ "start_offset": 0,
592
+ "end_offset": 0,
593
+ "size": 1,
594
+ },
595
+ 1: {
596
+ "filename": "raw-1.webp",
597
+ "file_type": "image",
598
+ "start_offset": 1,
599
+ "end_offset": 1,
600
+ "size": 1,
601
+ },
602
+ },
603
+ },
604
+ metadata={"chunk_index": 0, "shard_name": "raw-0"},
605
+ )
606
+
607
+ context = {}
608
+ items = list(worker.process_unit(unit, context))
609
+
610
+ assert len(items) == 2
611
+ assert [item["item_index"] for item in items] == [0, 1]
612
+ assert [item["job_id"] for item in items] == [
613
+ "raw-0:chunk:0:idx:0",
614
+ "raw-1:chunk:0:idx:1",
615
+ ]
616
+ assert context["_processed_indices"] == [0, 1]
617
+
520
618
  @pytest.mark.asyncio
521
619
  async def test_storage_update_flow(self, orchestrator_config, storage_manager, temp_dir):
522
620
  """Test updating chunk tracker from storage."""
@@ -268,6 +268,22 @@ class TestVLLMConfigManagerInitParams:
268
268
 
269
269
  assert params == expected
270
270
 
271
+ def test_get_vllm_init_params_new_mm_cache_config(self):
272
+ """Test getting vLLM init params with current multimodal cache config."""
273
+ manager = VLLMConfigManager()
274
+ vllm_config = {
275
+ "model": "custom-model",
276
+ "mm_processor_cache_gb": 8,
277
+ "mm_processor_cache_type": "shm",
278
+ "mm_shm_cache_max_object_size_mb": 256,
279
+ }
280
+
281
+ params = manager.get_vllm_init_params(vllm_config)
282
+
283
+ assert params["mm_processor_cache_gb"] == 8
284
+ assert params["mm_processor_cache_type"] == "shm"
285
+ assert params["mm_shm_cache_max_object_size_mb"] == 256
286
+
271
287
 
272
288
  class TestVLLMConfigManagerTokenizerReload:
273
289
  """Test VLLMConfigManager tokenizer reload checking."""
@@ -19,6 +19,7 @@ from caption_flow.storage import StorageManager
19
19
  # Import the modules to test
20
20
  from caption_flow.workers.caption import (
21
21
  CaptionWorker,
22
+ MultiStageVLLMManager,
22
23
  ProcessingItem,
23
24
  )
24
25
 
@@ -37,6 +38,96 @@ def create_fast_caption_worker(worker_config):
37
38
  return worker
38
39
 
39
40
 
41
+ class TestMultiStageVLLMManager:
42
+ """Test vLLM manager compatibility helpers."""
43
+
44
+ def test_filter_engine_args_maps_removed_mm_cache_flag(self):
45
+ """Test newer vLLM EngineArgs compatibility for removed cache flag."""
46
+
47
+ class EngineArgs:
48
+ def __init__(
49
+ self,
50
+ model,
51
+ trust_remote_code=False,
52
+ limit_mm_per_prompt=None,
53
+ mm_processor_cache_gb=4,
54
+ ):
55
+ pass
56
+
57
+ params = {
58
+ "model": "test-model",
59
+ "trust_remote_code": True,
60
+ "limit_mm_per_prompt": {"image": 1},
61
+ "disable_mm_preprocessor_cache": True,
62
+ "unsupported": "drop-me",
63
+ }
64
+
65
+ filtered = MultiStageVLLMManager._filter_engine_args(params, EngineArgs)
66
+
67
+ assert filtered == {
68
+ "model": "test-model",
69
+ "trust_remote_code": True,
70
+ "limit_mm_per_prompt": {"image": 1},
71
+ "mm_processor_cache_gb": 0,
72
+ }
73
+
74
+ def test_filter_engine_args_preserves_explicit_mm_cache_size(self):
75
+ """Test explicit newer vLLM cache config takes precedence over legacy alias."""
76
+
77
+ class EngineArgs:
78
+ def __init__(
79
+ self,
80
+ model,
81
+ trust_remote_code=False,
82
+ mm_processor_cache_gb=4,
83
+ mm_processor_cache_type="lru",
84
+ ):
85
+ pass
86
+
87
+ params = {
88
+ "model": "test-model",
89
+ "trust_remote_code": True,
90
+ "disable_mm_preprocessor_cache": True,
91
+ "mm_processor_cache_gb": 8,
92
+ "mm_processor_cache_type": "shm",
93
+ }
94
+
95
+ filtered = MultiStageVLLMManager._filter_engine_args(params, EngineArgs)
96
+
97
+ assert filtered == {
98
+ "model": "test-model",
99
+ "trust_remote_code": True,
100
+ "mm_processor_cache_gb": 8,
101
+ "mm_processor_cache_type": "shm",
102
+ }
103
+
104
+ def test_filter_engine_args_preserves_supported_mm_cache_flag(self):
105
+ """Test older vLLM EngineArgs compatibility for the legacy cache flag."""
106
+
107
+ class EngineArgs:
108
+ def __init__(
109
+ self,
110
+ model,
111
+ trust_remote_code=False,
112
+ disable_mm_preprocessor_cache=True,
113
+ ):
114
+ pass
115
+
116
+ params = {
117
+ "model": "test-model",
118
+ "trust_remote_code": True,
119
+ "disable_mm_preprocessor_cache": False,
120
+ }
121
+
122
+ filtered = MultiStageVLLMManager._filter_engine_args(params, EngineArgs)
123
+
124
+ assert filtered == {
125
+ "model": "test-model",
126
+ "trust_remote_code": True,
127
+ "disable_mm_preprocessor_cache": False,
128
+ }
129
+
130
+
40
131
  class TestCaptionWorker:
41
132
  """Test suite for CaptionWorker."""
42
133
 
File without changes
File without changes