Source code for hlsfactory.opt_dsl_v2.opt_dsl

import re
from collections import defaultdict


[docs] class OptDSL: def __init__(self, opt_dsl_source: str) -> None: self.static_lines: list[str] = [] self.grouped_directives: defaultdict[str, defaultdict[int, list[str]]] = ( defaultdict(lambda: defaultdict(list)) ) self.pipeline_directives: defaultdict[str, list[str | None]] = defaultdict(list) self.ungrouped_partition_directives: defaultdict[str, list[str]] = defaultdict( list ) self.ungrouped_unroll_directives: defaultdict[str, list[str]] = defaultdict( list ) self.opt_dsl_error: bool = False self.error_message: str = "" self.parse_opt_dsl_lang(opt_dsl_source)
[docs] def parse_opt_dsl_lang(self, opt_dsl_source) -> None: lines = opt_dsl_source.splitlines() dynamic_lines = [] for line in lines: stripped_line = line.strip() if ( # Extract static directives stripped_line.startswith("set_directive_") and not any( kw in stripped_line for kw in ["pipeline", "unroll", "array_partition"] ) ): self.static_lines.append(stripped_line) else: dynamic_lines.append(line) self.opt_dsl_error = self._validate_syntax(dynamic_lines) if not self.opt_dsl_error: opt_dsl_source = "\n".join(dynamic_lines) opt_dsl_source = re.sub(r"pipeline\s*\(", "self.pipeline(", opt_dsl_source) opt_dsl_source = re.sub(r"unroll\s*\(", "self.unroll(", opt_dsl_source) opt_dsl_source = re.sub( r"partition\s*\(", "self.partition(", opt_dsl_source ) ############################################################################# # exec is potentially unsafe, but seems like the easiest way to go about this ############################################################################# exec(opt_dsl_source) self.opt_dsl_error = self._validate_grouped_directives()
[docs] def pipeline(self, label: str, function: str, optional: bool = False) -> None: directive = f"set_directive_pipeline {function}/{label}" if optional: self.pipeline_directives[f"{function}/{label}"] = [directive, None] else: self.pipeline_directives[f"{function}/{label}"] = [directive]
[docs] def unroll( self, label: str, function: str, factor: list[int] | int, group: str | None = None, ) -> None: factors = factor if isinstance(factor, list) else [factor] for f in factors: directive = f"set_directive_unroll -factor {f} {function}/{label}" if group: self.grouped_directives[group][f].append(directive) else: self.ungrouped_unroll_directives[f"{function}/{label}"].append( directive )
[docs] def partition( self, array_var: str, function: str, partition_type: str, factor: list[int] | int, dim: int, group: str | None = None, ) -> None: factors = factor if isinstance(factor, list) else [factor] for f in factors: directive = f"set_directive_array_partition -type {partition_type} -factor {f} -dim {dim} {function} {array_var}" if group: self.grouped_directives[group][f].append(directive) else: self.ungrouped_partition_directives[f"{function}/{array_var}"].append( directive )
[docs] def get_directives(self): return ( self.static_lines, self.grouped_directives, self.pipeline_directives, self.ungrouped_partition_directives, self.ungrouped_unroll_directives, )
# Check syntax of dynamic lines
[docs] def _validate_syntax(self, dynamic_lines: list) -> bool: str_pat = r'(?:"[^"]+"|\'[^\']+\')' int_list_pat = r"(?:\d+|\[\s*\d+(?:\s*,\s*\d+)*\s*\])" bool_pat = r"(?:True|False)" check_pipeline = re.compile( rf"^\s*pipeline\(\s*{str_pat}\s*,\s*{str_pat}\s*" rf"(?:,\s*(?:{bool_pat}|optional\s*=\s*{bool_pat}))?\)\s*$" ) check_unroll = re.compile( rf"^\s*unroll\(\s*{str_pat}\s*,\s*{str_pat}\s*,\s*{int_list_pat}\s*" rf"(?:,\s*(?:{str_pat}|group\s*=\s*{str_pat}))?\)\s*$" ) check_partition = re.compile( rf"^\s*partition\(\s*{str_pat}\s*,\s*{str_pat}\s*,\s*{str_pat}\s*,\s*{int_list_pat}\s*,\s*\d+\s*" rf"(?:,\s*(?:{str_pat}|group\s*=\s*{str_pat}))?\)\s*$" ) for i, line in enumerate(dynamic_lines): if not line or line.strip().startswith("#"): continue # Skip empty lines and comments if not ( check_pipeline.match(line) or check_unroll.match(line) or check_partition.match(line) ): self.error_message = f"Syntax error: {line}" return True return False
# Check consistency of grouped factor lists
[docs] def _validate_grouped_directives(self) -> bool: for group, directives_by_factor in self.grouped_directives.items(): if len(directives_by_factor) < 2: continue factor_list = list(directives_by_factor.keys()) reference_len = len(directives_by_factor[factor_list[0]]) for factor, directives in directives_by_factor.items(): if len(directives) != reference_len: self.error_message = f"Inconsistent factor list in group '{group}'" return True return False