|
4 | 4 | import re |
5 | 5 |
|
6 | 6 |
|
| 7 | +_DETAIL_PREFIX = "infini/rt/detail" |
| 8 | + |
7 | 9 | _DEVICE_HEADERS = { |
8 | 10 | "cpu": ( |
9 | 11 | ("cpu", "data_type_.h", "native/cpu/data_type_.h"), |
@@ -82,20 +84,83 @@ def _write_wrapper(include_root, device, header_name, target): |
82 | 84 | f"""#ifndef {guard} |
83 | 85 | #define {guard} |
84 | 86 |
|
85 | | -#include "{target}" |
| 87 | +#include <{_DETAIL_PREFIX}/{target}> |
86 | 88 |
|
87 | 89 | #endif |
88 | 90 | """ |
89 | 91 | ) |
90 | 92 |
|
91 | 93 |
|
| 94 | +def _detail_include(path): |
| 95 | + return f"<{_DETAIL_PREFIX}/{path}>" |
| 96 | + |
| 97 | + |
| 98 | +def _rewrite_detail_include(match): |
| 99 | + target = match.group(1) |
| 100 | + return f"#include {_detail_include(target)}" |
| 101 | + |
| 102 | + |
| 103 | +_DETAIL_INCLUDE_PATTERN = re.compile( |
| 104 | + r'#include "((?:common|native)/[^"]+|data_type\.h|device\.h|dispatcher\.h|hash\.h|runtime\.h|tensor_view\.h)"' |
| 105 | +) |
| 106 | + |
| 107 | + |
| 108 | +def _detail_header_dependencies(source_root, relative_path): |
| 109 | + source_path = source_root / relative_path |
| 110 | + text = source_path.read_text() |
| 111 | + |
| 112 | + return { |
| 113 | + target |
| 114 | + for target in _DETAIL_INCLUDE_PATTERN.findall(text) |
| 115 | + if (source_root / target).exists() |
| 116 | + } |
| 117 | + |
| 118 | + |
| 119 | +def _write_detail_header(include_root, source_root, relative_path): |
| 120 | + source_path = source_root / relative_path |
| 121 | + output_path = include_root / _DETAIL_PREFIX / relative_path |
| 122 | + text = source_path.read_text() |
| 123 | + text = _DETAIL_INCLUDE_PATTERN.sub(_rewrite_detail_include, text) |
| 124 | + |
| 125 | + output_path.parent.mkdir(parents=True, exist_ok=True) |
| 126 | + output_path.write_text(text) |
| 127 | + |
| 128 | + |
| 129 | +def _write_detail_headers(include_root, source_root, devices): |
| 130 | + detail_headers = { |
| 131 | + "common/constexpr_map.h", |
| 132 | + "common/traits.h", |
| 133 | + "data_type.h", |
| 134 | + "device.h", |
| 135 | + "dispatcher.h", |
| 136 | + "hash.h", |
| 137 | + "runtime.h", |
| 138 | + "tensor_view.h", |
| 139 | + } |
| 140 | + |
| 141 | + for device in devices: |
| 142 | + for _, _, target in _DEVICE_HEADERS[device]: |
| 143 | + detail_headers.add(target) |
| 144 | + |
| 145 | + pending_headers = list(detail_headers) |
| 146 | + while pending_headers: |
| 147 | + relative_path = pending_headers.pop() |
| 148 | + for dependency in _detail_header_dependencies(source_root, relative_path): |
| 149 | + if dependency not in detail_headers: |
| 150 | + detail_headers.add(dependency) |
| 151 | + pending_headers.append(dependency) |
| 152 | + |
| 153 | + for relative_path in sorted(detail_headers): |
| 154 | + _write_detail_header(include_root, source_root, relative_path) |
| 155 | + |
| 156 | + |
92 | 157 | def _write_generated_header(include_root, devices): |
93 | 158 | includes = [ |
94 | | - '#include "data_type.h"', |
95 | | - '#include "device.h"', |
96 | | - '#include "hash.h"', |
97 | | - '#include "runtime.h"', |
98 | | - '#include "tensor_view.h"', |
| 159 | + f"#include {_detail_include('data_type.h')}", |
| 160 | + f"#include {_detail_include('device.h')}", |
| 161 | + f"#include {_detail_include('hash.h')}", |
| 162 | + f"#include {_detail_include('runtime.h')}", |
| 163 | + f"#include {_detail_include('tensor_view.h')}", |
99 | 164 | ] |
100 | 165 |
|
101 | 166 | for device in devices: |
@@ -349,7 +414,9 @@ def main(): |
349 | 414 | devices.append(device) |
350 | 415 |
|
351 | 416 | include_root = pathlib.Path(args.output_dir) |
| 417 | + source_root = pathlib.Path(args.runtime_header).parent |
352 | 418 |
|
| 419 | + _write_detail_headers(include_root, source_root, devices) |
353 | 420 | for device in devices: |
354 | 421 | for wrapper_device, header_name, target in _DEVICE_HEADERS[device]: |
355 | 422 | _write_wrapper(include_root, wrapper_device, header_name, target) |
|
0 commit comments