diff --git a/scripts/convert_rdf_to_hf_dataset_streaming.py b/scripts/convert_rdf_to_hf_dataset_streaming.py index b7f7ac6..663d4bb 100755 --- a/scripts/convert_rdf_to_hf_dataset_streaming.py +++ b/scripts/convert_rdf_to_hf_dataset_streaming.py @@ -74,6 +74,78 @@ from rich.progress import ( logger = logging.getLogger(__name__) +def _process_ntriples_file( + file_obj, chunk_size: int +) -> Iterator[list[dict[str, str]]]: + """Process N-Triples file and yield chunks of triples. + + This helper function processes the file object line by line, parsing each + line as an N-Triple and accumulating triples into chunks. + + Args: + file_obj: File-like object (already opened) + chunk_size: Number of triples per chunk + + Yields: + Chunks of triple dictionaries + """ + from rdflib import Graph + + current_chunk = [] + + for line_no, line in enumerate(file_obj, 1): + line = line.strip() + # Skip empty lines and comments + if not line or line.startswith("#"): + continue + + try: + # Parse single line as N-Triple + mini_graph = Graph() + mini_graph.parse(data=line, format="nt") + + for s, p, o in mini_graph: + # Convert RDFLib objects to dictionary format + if isinstance(o, Literal): + object_type = "literal" + object_datatype = ( + str(o.datatype) if o.datatype else None + ) + object_language = o.language if o.language else None + elif isinstance(o, URIRef): + object_type = "uri" + object_datatype = None + object_language = None + else: + object_type = "blank_node" + object_datatype = None + object_language = None + + triple = { + "subject": str(s), + "predicate": str(p), + "object": str(o), + "object_type": object_type, + "object_datatype": object_datatype, + "object_language": object_language, + } + current_chunk.append(triple) + + # Yield chunk when it reaches the target size + if len(current_chunk) >= chunk_size: + yield current_chunk + current_chunk = [] + + except Exception as e: + # Skip malformed lines and continue processing + logger.debug(f"Skipping malformed line {line_no}: {e}") + continue + + # Yield any remaining triples that didn't fill a complete chunk + if current_chunk: + yield current_chunk + + def stream_ntriples( file_path: Path, chunk_size: int = 10000 ) -> Iterator[list[dict[str, str]]]: @@ -86,119 +158,14 @@ def stream_ntriples( Yields: Chunks of triple dictionaries """ - from rdflib import Graph - - current_chunk = [] - - # Handle compressed files + # Open file with appropriate context manager based on compression + # Context managers ensure files are properly closed even if exceptions occur if file_path.suffix == ".gz": with gzip.open(file_path, "rt", encoding="utf-8") as file_obj: - # Create a small graph for parsing individual triples - mini_graph = Graph() - - for line_no, line in enumerate(file_obj, 1): - line = line.strip() - if not line or line.startswith("#"): - continue - - try: - # Parse single line as N-Triple - mini_graph = Graph() - mini_graph.parse(data=line, format="nt") - - for s, p, o in mini_graph: - # Convert to dictionary format - if isinstance(o, Literal): - object_type = "literal" - object_datatype = ( - str(o.datatype) if o.datatype else None - ) - object_language = o.language if o.language else None - elif isinstance(o, URIRef): - object_type = "uri" - object_datatype = None - object_language = None - else: - object_type = "blank_node" - object_datatype = None - object_language = None - - triple = { - "subject": str(s), - "predicate": str(p), - "object": str(o), - "object_type": object_type, - "object_datatype": object_datatype, - "object_language": object_language, - } - current_chunk.append(triple) - - if len(current_chunk) >= chunk_size: - yield current_chunk - current_chunk = [] - - except Exception as e: - # Skip malformed lines - logger.debug(f"Skipping malformed line {line_no}: {e}") - continue - - # Yield remaining triples - if current_chunk: - yield current_chunk + yield from _process_ntriples_file(file_obj, chunk_size) else: with open(file_path, encoding="utf-8") as file_obj: - # Create a small graph for parsing individual triples - mini_graph = Graph() - - for line_no, line in enumerate(file_obj, 1): - line = line.strip() - if not line or line.startswith("#"): - continue - - try: - # Parse single line as N-Triple - mini_graph = Graph() - mini_graph.parse(data=line, format="nt") - - for s, p, o in mini_graph: - # Convert to dictionary format - if isinstance(o, Literal): - object_type = "literal" - object_datatype = ( - str(o.datatype) if o.datatype else None - ) - object_language = o.language if o.language else None - elif isinstance(o, URIRef): - object_type = "uri" - object_datatype = None - object_language = None - else: - object_type = "blank_node" - object_datatype = None - object_language = None - - triple = { - "subject": str(s), - "predicate": str(p), - "object": str(o), - "object_type": object_type, - "object_datatype": object_datatype, - "object_language": object_language, - } - current_chunk.append(triple) - - if len(current_chunk) >= chunk_size: - yield current_chunk - current_chunk = [] - - except Exception as e: - # Skip malformed lines - logger.debug(f"Skipping malformed line {line_no}: {e}") - continue - - # Yield remaining triples - if current_chunk: - yield current_chunk + yield from _process_ntriples_file(file_obj, chunk_size) def stream_turtle_chunks( diff --git a/scripts/convert_rdf_to_hf_dataset_streaming_parallel.py b/scripts/convert_rdf_to_hf_dataset_streaming_parallel.py index bda8a1e..0111910 100755 --- a/scripts/convert_rdf_to_hf_dataset_streaming_parallel.py +++ b/scripts/convert_rdf_to_hf_dataset_streaming_parallel.py @@ -200,6 +200,64 @@ def process_ntriples_lines(lines): return triples +def _process_batched_lines(file_obj, batch_size: int, format: str): + """Process file lines and yield batches. + + Helper function to avoid code duplication between compressed and + uncompressed file handling. + + Args: + file_obj: File-like object (already opened) + batch_size: Number of lines per batch + format: File format (geonames or ntriples) + + Yields: + Batches of lines + """ + current_batch = [] + current_doc = [] + + if format == "geonames": + # For GeoNames, keep documents together + for line in file_obj: + # Check if this is a document boundary + if line.startswith("http://") or line.startswith("https://"): + if current_doc: + # Add completed document to batch + current_batch.extend(current_doc) + current_doc = [line] + + # Yield batch if large enough + if len(current_batch) >= batch_size: + yield current_batch + current_batch = [] + else: + current_doc = [line] + else: + current_doc.append(line) + + # Add final document + if current_doc: + current_batch.extend(current_doc) + + # Yield final batch + if current_batch: + yield current_batch + + else: + # For N-Triples, just batch lines + for line in file_obj: + current_batch.append(line) + + if len(current_batch) >= batch_size: + yield current_batch + current_batch = [] + + # Yield remaining lines + if current_batch: + yield current_batch + + def batch_file_lines(file_path: Path, batch_size: int = 1000, format: str = "geonames"): """Generator that yields batches of lines from file. @@ -213,7 +271,7 @@ def batch_file_lines(file_path: Path, batch_size: int = 1000, format: str = "geo Yields: Batches of lines """ - # Handle compressed files + # Handle compressed files - test validity first if file_path.suffix == ".gz": try: with gzip.open(file_path, "rt", encoding="utf-8") as test_file: @@ -237,93 +295,13 @@ def batch_file_lines(file_path: Path, batch_size: int = 1000, format: str = "geo ) raise + # Process compressed file with context manager with gzip.open(file_path, "rt", encoding="utf-8") as file_obj: - current_batch = [] - current_doc = [] - - if format == "geonames": - # For GeoNames, keep documents together - for line in file_obj: - # Check if this is a document boundary - if line.startswith("http://") or line.startswith("https://"): - if current_doc: - # Add completed document to batch - current_batch.extend(current_doc) - current_doc = [line] - - # Yield batch if large enough - if len(current_batch) >= batch_size: - yield current_batch - current_batch = [] - else: - current_doc = [line] - else: - current_doc.append(line) - - # Add final document - if current_doc: - current_batch.extend(current_doc) - - # Yield final batch - if current_batch: - yield current_batch - - else: - # For N-Triples, just batch lines - for line in file_obj: - current_batch.append(line) - - if len(current_batch) >= batch_size: - yield current_batch - current_batch = [] - - # Yield remaining lines - if current_batch: - yield current_batch + yield from _process_batched_lines(file_obj, batch_size, format) else: + # Process uncompressed file with context manager with open(file_path, encoding="utf-8") as file_obj: - current_batch = [] - current_doc = [] - - if format == "geonames": - # For GeoNames, keep documents together - for line in file_obj: - # Check if this is a document boundary - if line.startswith("http://") or line.startswith("https://"): - if current_doc: - # Add completed document to batch - current_batch.extend(current_doc) - current_doc = [line] - - # Yield batch if large enough - if len(current_batch) >= batch_size: - yield current_batch - current_batch = [] - else: - current_doc = [line] - else: - current_doc.append(line) - - # Add final document - if current_doc: - current_batch.extend(current_doc) - - # Yield final batch - if current_batch: - yield current_batch - - else: - # For N-Triples, just batch lines - for line in file_obj: - current_batch.append(line) - - if len(current_batch) >= batch_size: - yield current_batch - current_batch = [] - - # Yield remaining lines - if current_batch: - yield current_batch + yield from _process_batched_lines(file_obj, batch_size, format) def stream_geonames_parallel(