1#!/usr/bin/env python3 2# -*- coding: utf-8 -*- 3 4# Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. 5# See https://llvm.org/LICENSE.txt for license information. 6# SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception 7 8# Script for updating SPIR-V dialect by scraping information from SPIR-V 9# HTML and JSON specs from the Internet. 10# 11# For example, to define the enum attribute for SPIR-V memory model: 12# 13# ./gen_spirv_dialect.py --base_td_path /path/to/SPIRVBase.td \ 14# --new-enum MemoryModel 15# 16# The 'operand_kinds' dict of spirv.core.grammar.json contains all supported 17# SPIR-V enum classes. 18 19import itertools 20import re 21import requests 22import textwrap 23 24SPIRV_HTML_SPEC_URL = 'https://www.khronos.org/registry/spir-v/specs/unified1/SPIRV.html' 25SPIRV_JSON_SPEC_URL = 'https://raw.githubusercontent.com/KhronosGroup/SPIRV-Headers/master/include/spirv/unified1/spirv.core.grammar.json' 26 27AUTOGEN_OP_DEF_SEPARATOR = '\n// -----\n\n' 28AUTOGEN_ENUM_SECTION_MARKER = 'enum section. Generated from SPIR-V spec; DO NOT MODIFY!' 29AUTOGEN_OPCODE_SECTION_MARKER = ( 30 'opcode section. Generated from SPIR-V spec; DO NOT MODIFY!') 31 32 33def get_spirv_doc_from_html_spec(): 34 """Extracts instruction documentation from SPIR-V HTML spec. 35 36 Returns: 37 - A dict mapping from instruction opcode to documentation. 38 """ 39 response = requests.get(SPIRV_HTML_SPEC_URL) 40 spec = response.content 41 42 from bs4 import BeautifulSoup 43 spirv = BeautifulSoup(spec, 'html.parser') 44 45 section_anchor = spirv.find('h3', {'id': '_a_id_instructions_a_instructions'}) 46 47 doc = {} 48 49 for section in section_anchor.parent.find_all('div', {'class': 'sect3'}): 50 for table in section.find_all('table'): 51 inst_html = table.tbody.tr.td.p 52 opname = inst_html.a['id'] 53 # Ignore the first line, which is just the opname. 54 doc[opname] = inst_html.text.split('\n', 1)[1].strip() 55 56 return doc 57 58 59def get_spirv_grammar_from_json_spec(): 60 """Extracts operand kind and instruction grammar from SPIR-V JSON spec. 61 62 Returns: 63 - A list containing all operand kinds' grammar 64 - A list containing all instructions' grammar 65 """ 66 response = requests.get(SPIRV_JSON_SPEC_URL) 67 spec = response.content 68 69 import json 70 spirv = json.loads(spec) 71 72 return spirv['operand_kinds'], spirv['instructions'] 73 74 75def split_list_into_sublists(items, offset): 76 """Split the list of items into multiple sublists. 77 78 This is to make sure the string composed from each sublist won't exceed 79 80 characters. 80 81 Arguments: 82 - items: a list of strings 83 - offset: the offset in calculating each sublist's length 84 """ 85 chuncks = [] 86 chunk = [] 87 chunk_len = 0 88 89 for item in items: 90 chunk_len += len(item) + 2 91 if chunk_len > 80: 92 chuncks.append(chunk) 93 chunk = [] 94 chunk_len = len(item) + 2 95 chunk.append(item) 96 97 if len(chunk) != 0: 98 chuncks.append(chunk) 99 100 return chuncks 101 102 103def uniquify_enum_cases(lst): 104 """Prunes duplicate enum cases from the list. 105 106 Arguments: 107 - lst: List whose elements are to be uniqued. Assumes each element is a 108 (symbol, value) pair and elements already sorted according to value. 109 110 Returns: 111 - A list with all duplicates removed. The elements are sorted according to 112 value and, for each value, uniqued according to symbol. 113 original list, 114 - A map from deduplicated cases to the uniqued case. 115 """ 116 cases = lst 117 uniqued_cases = [] 118 duplicated_cases = {} 119 120 # First sort according to the value 121 cases.sort(key=lambda x: x[1]) 122 123 # Then group them according to the value 124 for _, groups in itertools.groupby(cases, key=lambda x: x[1]): 125 # For each value, sort according to the enumerant symbol. 126 sorted_group = sorted(groups, key=lambda x: x[0]) 127 # Keep the "smallest" case, which is typically the symbol without extension 128 # suffix. But we have special cases that we want to fix. 129 case = sorted_group[0] 130 for i in range(1, len(sorted_group)): 131 duplicated_cases[sorted_group[i][0]] = case[0] 132 if case[0] == 'HlslSemanticGOOGLE': 133 assert len(sorted_group) == 2, 'unexpected new variant for HlslSemantic' 134 case = sorted_group[1] 135 duplicated_cases[sorted_group[0][0]] = case[0] 136 uniqued_cases.append(case) 137 138 return uniqued_cases, duplicated_cases 139 140 141def toposort(dag, sort_fn): 142 """Topologically sorts the given dag. 143 144 Arguments: 145 - dag: a dict mapping from a node to its incoming nodes. 146 - sort_fn: a function for sorting nodes in the same batch. 147 148 Returns: 149 A list containing topologically sorted nodes. 150 """ 151 152 # Returns the next batch of nodes without incoming edges 153 def get_next_batch(dag): 154 while True: 155 no_prev_nodes = set(node for node, prev in dag.items() if not prev) 156 if not no_prev_nodes: 157 break 158 yield sorted(no_prev_nodes, key=sort_fn) 159 dag = { 160 node: (prev - no_prev_nodes) 161 for node, prev in dag.items() 162 if node not in no_prev_nodes 163 } 164 assert not dag, 'found cyclic dependency' 165 166 sorted_nodes = [] 167 for batch in get_next_batch(dag): 168 sorted_nodes.extend(batch) 169 170 return sorted_nodes 171 172 173def toposort_capabilities(all_cases, capability_mapping): 174 """Returns topologically sorted capability (symbol, value) pairs. 175 176 Arguments: 177 - all_cases: all capability cases (containing symbol, value, and implied 178 capabilities). 179 - capability_mapping: mapping from duplicated capability symbols to the 180 canonicalized symbol chosen for SPIRVBase.td. 181 182 Returns: 183 A list containing topologically sorted capability (symbol, value) pairs. 184 """ 185 dag = {} 186 name_to_value = {} 187 for case in all_cases: 188 # Get the current capability. 189 cur = case['enumerant'] 190 name_to_value[cur] = case['value'] 191 # Ignore duplicated symbols. 192 if cur in capability_mapping: 193 continue 194 195 # Get capabilities implied by the current capability. 196 prev = case.get('capabilities', []) 197 uniqued_prev = set([capability_mapping.get(c, c) for c in prev]) 198 dag[cur] = uniqued_prev 199 200 sorted_caps = toposort(dag, lambda x: name_to_value[x]) 201 # Attach the capability's value as the second component of the pair. 202 return [(c, name_to_value[c]) for c in sorted_caps] 203 204 205def get_capability_mapping(operand_kinds): 206 """Returns the capability mapping from duplicated cases to canonicalized ones. 207 208 Arguments: 209 - operand_kinds: all operand kinds' grammar spec 210 211 Returns: 212 - A map mapping from duplicated capability symbols to the canonicalized 213 symbol chosen for SPIRVBase.td. 214 """ 215 # Find the operand kind for capability 216 cap_kind = {} 217 for kind in operand_kinds: 218 if kind['kind'] == 'Capability': 219 cap_kind = kind 220 221 kind_cases = [ 222 (case['enumerant'], case['value']) for case in cap_kind['enumerants'] 223 ] 224 _, capability_mapping = uniquify_enum_cases(kind_cases) 225 226 return capability_mapping 227 228 229def get_availability_spec(enum_case, capability_mapping, for_op, for_cap): 230 """Returns the availability specification string for the given enum case. 231 232 Arguments: 233 - enum_case: the enum case to generate availability spec for. It may contain 234 'version', 'lastVersion', 'extensions', or 'capabilities'. 235 - capability_mapping: mapping from duplicated capability symbols to the 236 canonicalized symbol chosen for SPIRVBase.td. 237 - for_op: bool value indicating whether this is the availability spec for an 238 op itself. 239 - for_cap: bool value indicating whether this is the availability spec for 240 capabilities themselves. 241 242 Returns: 243 - A `let availability = [...];` string if with availability spec or 244 empty string if without availability spec 245 """ 246 assert not (for_op and for_cap), 'cannot set both for_op and for_cap' 247 248 DEFAULT_MIN_VERSION = 'MinVersion<SPV_V_1_0>' 249 DEFAULT_MAX_VERSION = 'MaxVersion<SPV_V_1_5>' 250 DEFAULT_CAP = 'Capability<[]>' 251 DEFAULT_EXT = 'Extension<[]>' 252 253 min_version = enum_case.get('version', '') 254 if min_version == 'None': 255 min_version = '' 256 elif min_version: 257 min_version = 'MinVersion<SPV_V_{}>'.format(min_version.replace('.', '_')) 258 # TODO(antiagainst): delete this once ODS can support dialect-specific content 259 # and we can use omission to mean no requirements. 260 if for_op and not min_version: 261 min_version = DEFAULT_MIN_VERSION 262 263 max_version = enum_case.get('lastVersion', '') 264 if max_version: 265 max_version = 'MaxVersion<SPV_V_{}>'.format(max_version.replace('.', '_')) 266 # TODO(antiagainst): delete this once ODS can support dialect-specific content 267 # and we can use omission to mean no requirements. 268 if for_op and not max_version: 269 max_version = DEFAULT_MAX_VERSION 270 271 exts = enum_case.get('extensions', []) 272 if exts: 273 exts = 'Extension<[{}]>'.format(', '.join(sorted(set(exts)))) 274 # We need to strip the minimal version requirement if this symbol is 275 # available via an extension, which means *any* SPIR-V version can support 276 # it as long as the extension is provided. The grammar's 'version' field 277 # under such case should be interpreted as this symbol is introduced as 278 # a core symbol since the given version, rather than a minimal version 279 # requirement. 280 min_version = DEFAULT_MIN_VERSION if for_op else '' 281 # TODO(antiagainst): delete this once ODS can support dialect-specific content 282 # and we can use omission to mean no requirements. 283 if for_op and not exts: 284 exts = DEFAULT_EXT 285 286 caps = enum_case.get('capabilities', []) 287 implies = '' 288 if caps: 289 canonicalized_caps = [] 290 for c in caps: 291 if c in capability_mapping: 292 canonicalized_caps.append(capability_mapping[c]) 293 else: 294 canonicalized_caps.append(c) 295 prefixed_caps = [ 296 'SPV_C_{}'.format(c) for c in sorted(set(canonicalized_caps)) 297 ] 298 if for_cap: 299 # If this is generating the availability for capabilities, we need to 300 # put the capability "requirements" in implies field because now 301 # the "capabilities" field in the source grammar means so. 302 caps = '' 303 implies = 'list<I32EnumAttrCase> implies = [{}];'.format( 304 ', '.join(prefixed_caps)) 305 else: 306 caps = 'Capability<[{}]>'.format(', '.join(prefixed_caps)) 307 implies = '' 308 # TODO(antiagainst): delete this once ODS can support dialect-specific content 309 # and we can use omission to mean no requirements. 310 if for_op and not caps: 311 caps = DEFAULT_CAP 312 313 avail = '' 314 # Compose availability spec if any of the requirements is not empty. 315 # For ops, because we have a default in SPV_Op class, omit if the spec 316 # is the same. 317 if (min_version or max_version or caps or exts) and not ( 318 for_op and min_version == DEFAULT_MIN_VERSION and 319 max_version == DEFAULT_MAX_VERSION and caps == DEFAULT_CAP and 320 exts == DEFAULT_EXT): 321 joined_spec = ',\n '.join( 322 [e for e in [min_version, max_version, exts, caps] if e]) 323 avail = '{} availability = [\n {}\n ];'.format( 324 'let' if for_op else 'list<Availability>', joined_spec) 325 326 return '{}{}{}'.format(implies, '\n ' if implies and avail else '', avail) 327 328 329def gen_operand_kind_enum_attr(operand_kind, capability_mapping): 330 """Generates the TableGen EnumAttr definition for the given operand kind. 331 332 Returns: 333 - The operand kind's name 334 - A string containing the TableGen EnumAttr definition 335 """ 336 if 'enumerants' not in operand_kind: 337 return '', '' 338 339 # Returns a symbol for the given case in the given kind. This function 340 # handles Dim specially to avoid having numbers as the start of symbols, 341 # which does not play well with C++ and the MLIR parser. 342 def get_case_symbol(kind_name, case_name): 343 if kind_name == 'Dim': 344 if case_name == '1D' or case_name == '2D' or case_name == '3D': 345 return 'Dim{}'.format(case_name) 346 return case_name 347 348 kind_name = operand_kind['kind'] 349 is_bit_enum = operand_kind['category'] == 'BitEnum' 350 kind_category = 'Bit' if is_bit_enum else 'I32' 351 kind_acronym = ''.join([c for c in kind_name if c >= 'A' and c <= 'Z']) 352 353 name_to_case_dict = {} 354 for case in operand_kind['enumerants']: 355 name_to_case_dict[case['enumerant']] = case 356 357 if kind_name == 'Capability': 358 # Special treatment for capability cases: we need to sort them topologically 359 # because a capability can refer to another via the 'implies' field. 360 kind_cases = toposort_capabilities(operand_kind['enumerants'], 361 capability_mapping) 362 else: 363 kind_cases = [(case['enumerant'], case['value']) 364 for case in operand_kind['enumerants']] 365 kind_cases, _ = uniquify_enum_cases(kind_cases) 366 max_len = max([len(symbol) for (symbol, _) in kind_cases]) 367 368 # Generate the definition for each enum case 369 fmt_str = 'def SPV_{acronym}_{case} {colon:>{offset}} '\ 370 '{category}EnumAttrCase<"{symbol}", {value}>{avail}' 371 case_defs = [] 372 for case in kind_cases: 373 avail = get_availability_spec(name_to_case_dict[case[0]], 374 capability_mapping, 375 False, kind_name == 'Capability') 376 case_def = fmt_str.format( 377 category=kind_category, 378 acronym=kind_acronym, 379 case=case[0], 380 symbol=get_case_symbol(kind_name, case[0]), 381 value=case[1], 382 avail=' {{\n {}\n}}'.format(avail) if avail else ';', 383 colon=':', 384 offset=(max_len + 1 - len(case[0]))) 385 case_defs.append(case_def) 386 case_defs = '\n'.join(case_defs) 387 388 # Generate the list of enum case names 389 fmt_str = 'SPV_{acronym}_{symbol}'; 390 case_names = [fmt_str.format(acronym=kind_acronym,symbol=case[0]) 391 for case in kind_cases] 392 393 # Split them into sublists and concatenate into multiple lines 394 case_names = split_list_into_sublists(case_names, 6) 395 case_names = ['{:6}'.format('') + ', '.join(sublist) 396 for sublist in case_names] 397 case_names = ',\n'.join(case_names) 398 399 # Generate the enum attribute definition 400 enum_attr = '''def SPV_{name}Attr : 401 SPV_{category}EnumAttr<"{name}", "valid SPIR-V {name}", [ 402{cases} 403 ]>;'''.format( 404 name=kind_name, category=kind_category, cases=case_names) 405 return kind_name, case_defs + '\n\n' + enum_attr 406 407 408def gen_opcode(instructions): 409 """ Generates the TableGen definition to map opname to opcode 410 411 Returns: 412 - A string containing the TableGen SPV_OpCode definition 413 """ 414 415 max_len = max([len(inst['opname']) for inst in instructions]) 416 def_fmt_str = 'def SPV_OC_{name} {colon:>{offset}} '\ 417 'I32EnumAttrCase<"{name}", {value}>;' 418 opcode_defs = [ 419 def_fmt_str.format( 420 name=inst['opname'], 421 value=inst['opcode'], 422 colon=':', 423 offset=(max_len + 1 - len(inst['opname']))) for inst in instructions 424 ] 425 opcode_str = '\n'.join(opcode_defs) 426 427 decl_fmt_str = 'SPV_OC_{name}' 428 opcode_list = [ 429 decl_fmt_str.format(name=inst['opname']) for inst in instructions 430 ] 431 opcode_list = split_list_into_sublists(opcode_list, 6) 432 opcode_list = [ 433 '{:6}'.format('') + ', '.join(sublist) for sublist in opcode_list 434 ] 435 opcode_list = ',\n'.join(opcode_list) 436 enum_attr = 'def SPV_OpcodeAttr :\n'\ 437 ' SPV_I32EnumAttr<"{name}", "valid SPIR-V instructions", [\n'\ 438 '{lst}\n'\ 439 ' ]>;'.format(name='Opcode', lst=opcode_list) 440 return opcode_str + '\n\n' + enum_attr 441 442 443def update_td_opcodes(path, instructions, filter_list): 444 """Updates SPIRBase.td with new generated opcode cases. 445 446 Arguments: 447 - path: the path to SPIRBase.td 448 - instructions: a list containing all SPIR-V instructions' grammar 449 - filter_list: a list containing new opnames to add 450 """ 451 452 with open(path, 'r') as f: 453 content = f.read() 454 455 content = content.split(AUTOGEN_OPCODE_SECTION_MARKER) 456 assert len(content) == 3 457 458 # Extend opcode list with existing list 459 existing_opcodes = [k[11:] for k in re.findall('def SPV_OC_\w+', content[1])] 460 filter_list.extend(existing_opcodes) 461 filter_list = list(set(filter_list)) 462 463 # Generate the opcode for all instructions in SPIR-V 464 filter_instrs = list( 465 filter(lambda inst: (inst['opname'] in filter_list), instructions)) 466 # Sort instruction based on opcode 467 filter_instrs.sort(key=lambda inst: inst['opcode']) 468 opcode = gen_opcode(filter_instrs) 469 470 # Substitute the opcode 471 content = content[0] + AUTOGEN_OPCODE_SECTION_MARKER + '\n\n' + \ 472 opcode + '\n\n// End ' + AUTOGEN_OPCODE_SECTION_MARKER \ 473 + content[2] 474 475 with open(path, 'w') as f: 476 f.write(content) 477 478 479def update_td_enum_attrs(path, operand_kinds, filter_list): 480 """Updates SPIRBase.td with new generated enum definitions. 481 482 Arguments: 483 - path: the path to SPIRBase.td 484 - operand_kinds: a list containing all operand kinds' grammar 485 - filter_list: a list containing new enums to add 486 """ 487 with open(path, 'r') as f: 488 content = f.read() 489 490 content = content.split(AUTOGEN_ENUM_SECTION_MARKER) 491 assert len(content) == 3 492 493 # Extend filter list with existing enum definitions 494 existing_kinds = [ 495 k[8:-4] for k in re.findall('def SPV_\w+Attr', content[1])] 496 filter_list.extend(existing_kinds) 497 498 capability_mapping = get_capability_mapping(operand_kinds) 499 500 # Generate definitions for all enums in filter list 501 defs = [ 502 gen_operand_kind_enum_attr(kind, capability_mapping) 503 for kind in operand_kinds 504 if kind['kind'] in filter_list 505 ] 506 # Sort alphabetically according to enum name 507 defs.sort(key=lambda enum : enum[0]) 508 # Only keep the definitions from now on 509 # Put Capability's definition at the very beginning because capability cases 510 # will be referenced later 511 defs = [enum[1] for enum in defs if enum[0] == 'Capability' 512 ] + [enum[1] for enum in defs if enum[0] != 'Capability'] 513 514 # Substitute the old section 515 content = content[0] + AUTOGEN_ENUM_SECTION_MARKER + '\n\n' + \ 516 '\n\n'.join(defs) + "\n\n// End " + AUTOGEN_ENUM_SECTION_MARKER \ 517 + content[2]; 518 519 with open(path, 'w') as f: 520 f.write(content) 521 522 523def snake_casify(name): 524 """Turns the given name to follow snake_case convention.""" 525 name = re.sub('\W+', '', name).split() 526 name = [s.lower() for s in name] 527 return '_'.join(name) 528 529 530def map_spec_operand_to_ods_argument(operand): 531 """Maps an operand in SPIR-V JSON spec to an op argument in ODS. 532 533 Arguments: 534 - A dict containing the operand's kind, quantifier, and name 535 536 Returns: 537 - A string containing both the type and name for the argument 538 """ 539 kind = operand['kind'] 540 quantifier = operand.get('quantifier', '') 541 542 # These instruction "operands" are for encoding the results; they should 543 # not be handled here. 544 assert kind != 'IdResultType', 'unexpected to handle "IdResultType" kind' 545 assert kind != 'IdResult', 'unexpected to handle "IdResult" kind' 546 547 if kind == 'IdRef': 548 if quantifier == '': 549 arg_type = 'SPV_Type' 550 elif quantifier == '?': 551 arg_type = 'SPV_Optional<SPV_Type>' 552 else: 553 arg_type = 'Variadic<SPV_Type>' 554 elif kind == 'IdMemorySemantics' or kind == 'IdScope': 555 # TODO(antiagainst): Need to further constrain 'IdMemorySemantics' 556 # and 'IdScope' given that they should be generated from OpConstant. 557 assert quantifier == '', ('unexpected to have optional/variadic memory ' 558 'semantics or scope <id>') 559 arg_type = 'SPV_' + kind[2:] + 'Attr' 560 elif kind == 'LiteralInteger': 561 if quantifier == '': 562 arg_type = 'I32Attr' 563 elif quantifier == '?': 564 arg_type = 'OptionalAttr<I32Attr>' 565 else: 566 arg_type = 'OptionalAttr<I32ArrayAttr>' 567 elif kind == 'LiteralString' or \ 568 kind == 'LiteralContextDependentNumber' or \ 569 kind == 'LiteralExtInstInteger' or \ 570 kind == 'LiteralSpecConstantOpInteger' or \ 571 kind == 'PairLiteralIntegerIdRef' or \ 572 kind == 'PairIdRefLiteralInteger' or \ 573 kind == 'PairIdRefIdRef': 574 assert False, '"{}" kind unimplemented'.format(kind) 575 else: 576 # The rest are all enum operands that we represent with op attributes. 577 assert quantifier != '*', 'unexpected to have variadic enum attribute' 578 arg_type = 'SPV_{}Attr'.format(kind) 579 if quantifier == '?': 580 arg_type = 'OptionalAttr<{}>'.format(arg_type) 581 582 name = operand.get('name', '') 583 name = snake_casify(name) if name else kind.lower() 584 585 return '{}:${}'.format(arg_type, name) 586 587 588def get_description(text, appendix): 589 """Generates the description for the given SPIR-V instruction. 590 591 Arguments: 592 - text: Textual description of the operation as string. 593 - appendix: Additional contents to attach in description as string, 594 includking IR examples, and others. 595 596 Returns: 597 - A string that corresponds to the description of the Tablegen op. 598 """ 599 fmt_str = '{text}\n\n <!-- End of AutoGen section -->\n{appendix}\n ' 600 return fmt_str.format(text=text, appendix=appendix) 601 602 603def get_op_definition(instruction, doc, existing_info, capability_mapping): 604 """Generates the TableGen op definition for the given SPIR-V instruction. 605 606 Arguments: 607 - instruction: the instruction's SPIR-V JSON grammar 608 - doc: the instruction's SPIR-V HTML doc 609 - existing_info: a dict containing potential manually specified sections for 610 this instruction 611 - capability_mapping: mapping from duplicated capability symbols to the 612 canonicalized symbol chosen for SPIRVBase.td 613 614 Returns: 615 - A string containing the TableGen op definition 616 """ 617 fmt_str = ('def SPV_{opname}Op : ' 618 'SPV_{inst_category}<"{opname}"{category_args}[{traits}]> ' 619 '{{\n let summary = {summary};\n\n let description = ' 620 '[{{\n{description}}}];{availability}\n') 621 inst_category = existing_info.get('inst_category', 'Op') 622 if inst_category == 'Op': 623 fmt_str +='\n let arguments = (ins{args});\n\n'\ 624 ' let results = (outs{results});\n' 625 626 fmt_str +='{extras}'\ 627 '}}\n' 628 629 opname = instruction['opname'][2:] 630 category_args = existing_info.get('category_args', '') 631 632 if '\n' in doc: 633 summary, text = doc.split('\n', 1) 634 else: 635 summary = doc 636 text = '' 637 wrapper = textwrap.TextWrapper( 638 width=76, initial_indent=' ', subsequent_indent=' ') 639 640 # Format summary. If the summary can fit in the same line, we print it out 641 # as a "-quoted string; otherwise, wrap the lines using "[{...}]". 642 summary = summary.strip(); 643 if len(summary) + len(' let summary = "";') <= 80: 644 summary = '"{}"'.format(summary) 645 else: 646 summary = '[{{\n{}\n }}]'.format(wrapper.fill(summary)) 647 648 # Wrap text 649 text = text.split('\n') 650 text = [wrapper.fill(line) for line in text if line] 651 text = '\n\n'.join(text) 652 653 operands = instruction.get('operands', []) 654 655 # Op availability 656 avail = '' 657 # We assume other instruction categories has a base availability spec, so 658 # only add this if this is directly using SPV_Op as the base. 659 if inst_category == 'Op': 660 avail = get_availability_spec(instruction, capability_mapping, True, False) 661 if avail: 662 avail = '\n\n {0}'.format(avail) 663 664 # Set op's result 665 results = '' 666 if len(operands) > 0 and operands[0]['kind'] == 'IdResultType': 667 results = '\n SPV_Type:$result\n ' 668 operands = operands[1:] 669 if 'results' in existing_info: 670 results = existing_info['results'] 671 672 # Ignore the operand standing for the result <id> 673 if len(operands) > 0 and operands[0]['kind'] == 'IdResult': 674 operands = operands[1:] 675 676 # Set op' argument 677 arguments = existing_info.get('arguments', None) 678 if arguments is None: 679 arguments = [map_spec_operand_to_ods_argument(o) for o in operands] 680 arguments = ',\n '.join(arguments) 681 if arguments: 682 # Prepend and append whitespace for formatting 683 arguments = '\n {}\n '.format(arguments) 684 685 description = existing_info.get('description', None) 686 if description is None: 687 assembly = '\n ```\n'\ 688 ' [TODO]\n'\ 689 ' ```mlir\n\n'\ 690 ' #### Example:\n\n'\ 691 ' ```\n'\ 692 ' [TODO]\n' \ 693 ' ```' 694 description = get_description(text, assembly) 695 696 return fmt_str.format( 697 opname=opname, 698 category_args=category_args, 699 inst_category=inst_category, 700 traits=existing_info.get('traits', ''), 701 summary=summary, 702 description=description, 703 availability=avail, 704 args=arguments, 705 results=results, 706 extras=existing_info.get('extras', '')) 707 708 709def get_string_between(base, start, end): 710 """Extracts a substring with a specified start and end from a string. 711 712 Arguments: 713 - base: string to extract from. 714 - start: string to use as the start of the substring. 715 - end: string to use as the end of the substring. 716 717 Returns: 718 - The substring if found 719 - The part of the base after end of the substring. Is the base string itself 720 if the substring wasnt found. 721 """ 722 split = base.split(start, 1) 723 if len(split) == 2: 724 rest = split[1].split(end, 1) 725 assert len(rest) == 2, \ 726 'cannot find end "{end}" while extracting substring '\ 727 'starting with {start}'.format(start=start, end=end) 728 return rest[0].rstrip(end), rest[1] 729 return '', split[0] 730 731 732def get_string_between_nested(base, start, end): 733 """Extracts a substring with a nested start and end from a string. 734 735 Arguments: 736 - base: string to extract from. 737 - start: string to use as the start of the substring. 738 - end: string to use as the end of the substring. 739 740 Returns: 741 - The substring if found 742 - The part of the base after end of the substring. Is the base string itself 743 if the substring wasn't found. 744 """ 745 split = base.split(start, 1) 746 if len(split) == 2: 747 # Handle nesting delimiters 748 rest = split[1] 749 unmatched_start = 1 750 index = 0 751 while unmatched_start > 0 and index < len(rest): 752 if rest[index:].startswith(end): 753 unmatched_start -= 1 754 if unmatched_start == 0: 755 break 756 index += len(end) 757 elif rest[index:].startswith(start): 758 unmatched_start += 1 759 index += len(start) 760 else: 761 index += 1 762 763 assert index < len(rest), \ 764 'cannot find end "{end}" while extracting substring '\ 765 'starting with "{start}"'.format(start=start, end=end) 766 return rest[:index], rest[index + len(end):] 767 return '', split[0] 768 769 770def extract_td_op_info(op_def): 771 """Extracts potentially manually specified sections in op's definition. 772 773 Arguments: - A string containing the op's TableGen definition 774 - doc: the instruction's SPIR-V HTML doc 775 776 Returns: 777 - A dict containing potential manually specified sections 778 """ 779 # Get opname 780 opname = [o[8:-2] for o in re.findall('def SPV_\w+Op', op_def)] 781 assert len(opname) == 1, 'more than one ops in the same section!' 782 opname = opname[0] 783 784 # Get instruction category 785 inst_category = [ 786 o[4:] for o in re.findall('SPV_\w+Op', 787 op_def.split(':', 1)[1]) 788 ] 789 assert len(inst_category) <= 1, 'more than one ops in the same section!' 790 inst_category = inst_category[0] if len(inst_category) == 1 else 'Op' 791 792 # Get category_args 793 op_tmpl_params, _ = get_string_between_nested(op_def, '<', '>') 794 opstringname, rest = get_string_between(op_tmpl_params, '"', '"') 795 category_args = rest.split('[', 1)[0] 796 797 # Get traits 798 traits, _ = get_string_between_nested(rest, '[', ']') 799 800 # Get description 801 description, rest = get_string_between(op_def, 'let description = [{\n', 802 '}];\n') 803 804 # Get arguments 805 args, rest = get_string_between(rest, ' let arguments = (ins', ');\n') 806 807 # Get results 808 results, rest = get_string_between(rest, ' let results = (outs', ');\n') 809 810 extras = rest.strip(' }\n') 811 if extras: 812 extras = '\n {}\n'.format(extras) 813 814 return { 815 # Prefix with 'Op' to make it consistent with SPIR-V spec 816 'opname': 'Op{}'.format(opname), 817 'inst_category': inst_category, 818 'category_args': category_args, 819 'traits': traits, 820 'description': description, 821 'arguments': args, 822 'results': results, 823 'extras': extras 824 } 825 826 827def update_td_op_definitions(path, instructions, docs, filter_list, 828 inst_category, capability_mapping): 829 """Updates SPIRVOps.td with newly generated op definition. 830 831 Arguments: 832 - path: path to SPIRVOps.td 833 - instructions: SPIR-V JSON grammar for all instructions 834 - docs: SPIR-V HTML doc for all instructions 835 - filter_list: a list containing new opnames to include 836 - capability_mapping: mapping from duplicated capability symbols to the 837 canonicalized symbol chosen for SPIRVBase.td. 838 839 Returns: 840 - A string containing all the TableGen op definitions 841 """ 842 with open(path, 'r') as f: 843 content = f.read() 844 845 # Split the file into chunks, each containing one op. 846 ops = content.split(AUTOGEN_OP_DEF_SEPARATOR) 847 header = ops[0] 848 footer = ops[-1] 849 ops = ops[1:-1] 850 851 # For each existing op, extract the manually-written sections out to retain 852 # them when re-generating the ops. Also append the existing ops to filter 853 # list. 854 name_op_map = {} # Map from opname to its existing ODS definition 855 op_info_dict = {} 856 for op in ops: 857 info_dict = extract_td_op_info(op) 858 opname = info_dict['opname'] 859 name_op_map[opname] = op 860 op_info_dict[opname] = info_dict 861 filter_list.append(opname) 862 filter_list = sorted(list(set(filter_list))) 863 864 op_defs = [] 865 for opname in filter_list: 866 # Find the grammar spec for this op 867 try: 868 instruction = next( 869 inst for inst in instructions if inst['opname'] == opname) 870 op_defs.append( 871 get_op_definition( 872 instruction, docs[opname], 873 op_info_dict.get(opname, {'inst_category': inst_category}), 874 capability_mapping)) 875 except StopIteration: 876 # This is an op added by us; use the existing ODS definition. 877 op_defs.append(name_op_map[opname]) 878 879 # Substitute the old op definitions 880 op_defs = [header] + op_defs + [footer] 881 content = AUTOGEN_OP_DEF_SEPARATOR.join(op_defs) 882 883 with open(path, 'w') as f: 884 f.write(content) 885 886 887if __name__ == '__main__': 888 import argparse 889 890 cli_parser = argparse.ArgumentParser( 891 description='Update SPIR-V dialect definitions using SPIR-V spec') 892 893 cli_parser.add_argument( 894 '--base-td-path', 895 dest='base_td_path', 896 type=str, 897 default=None, 898 help='Path to SPIRVBase.td') 899 cli_parser.add_argument( 900 '--op-td-path', 901 dest='op_td_path', 902 type=str, 903 default=None, 904 help='Path to SPIRVOps.td') 905 906 cli_parser.add_argument( 907 '--new-enum', 908 dest='new_enum', 909 type=str, 910 default=None, 911 help='SPIR-V enum to be added to SPIRVBase.td') 912 cli_parser.add_argument( 913 '--new-opcodes', 914 dest='new_opcodes', 915 type=str, 916 default=None, 917 nargs='*', 918 help='update SPIR-V opcodes in SPIRVBase.td') 919 cli_parser.add_argument( 920 '--new-inst', 921 dest='new_inst', 922 type=str, 923 default=None, 924 nargs='*', 925 help='SPIR-V instruction to be added to ops file') 926 cli_parser.add_argument( 927 '--inst-category', 928 dest='inst_category', 929 type=str, 930 default='Op', 931 help='SPIR-V instruction category used for choosing '\ 932 'the TableGen base class to define this op') 933 934 args = cli_parser.parse_args() 935 936 operand_kinds, instructions = get_spirv_grammar_from_json_spec() 937 938 # Define new enum attr 939 if args.new_enum is not None: 940 assert args.base_td_path is not None 941 filter_list = [args.new_enum] if args.new_enum else [] 942 update_td_enum_attrs(args.base_td_path, operand_kinds, filter_list) 943 944 # Define new opcode 945 if args.new_opcodes is not None: 946 assert args.base_td_path is not None 947 update_td_opcodes(args.base_td_path, instructions, args.new_opcodes) 948 949 # Define new op 950 if args.new_inst is not None: 951 assert args.op_td_path is not None 952 docs = get_spirv_doc_from_html_spec() 953 capability_mapping = get_capability_mapping(operand_kinds) 954 update_td_op_definitions(args.op_td_path, instructions, docs, args.new_inst, 955 args.inst_category, capability_mapping) 956 print('Done. Note that this script just generates a template; ', end='') 957 print('please read the spec and update traits, arguments, and ', end='') 958 print('results accordingly.') 959