| 4870 | |
| 4871 | |
| 4872 | def codegen_func_forward(adj, func_type="kernel", device="cpu"): |
| 4873 | if device == "cpu": |
| 4874 | indent = 4 |
| 4875 | elif device == "cuda": |
| 4876 | if func_type == "kernel": |
| 4877 | indent = 8 |
| 4878 | else: |
| 4879 | indent = 4 |
| 4880 | else: |
| 4881 | raise ValueError(f"Device {device} not supported for codegen") |
| 4882 | |
| 4883 | indent_block = " " * indent |
| 4884 | |
| 4885 | lines = [] |
| 4886 | |
| 4887 | # argument vars |
| 4888 | if device == "cpu" and func_type == "kernel": |
| 4889 | lines += ["//---------\n"] |
| 4890 | lines += ["// argument vars\n"] |
| 4891 | |
| 4892 | for var in adj.args: |
| 4893 | lines += [f"{var.ctype()} {var.emit()} = _wp_args->{var.label};\n"] |
| 4894 | |
| 4895 | # primal vars |
| 4896 | lines += ["//---------\n"] |
| 4897 | lines += ["// primal vars\n"] |
| 4898 | |
| 4899 | for var in adj.variables: |
| 4900 | if is_tile(var.type): |
| 4901 | lines += [f"{var.ctype()} {var.emit()} = {var.type.cinit(requires_grad=False)};\n"] |
| 4902 | elif is_tile_stack(var.type): |
| 4903 | lines += [f"{var.ctype()} {var.emit()} = {var.type.cinit()};\n"] |
| 4904 | elif var.constant is None: |
| 4905 | lines += [f"{var.ctype()} {var.emit()};\n"] |
| 4906 | else: |
| 4907 | lines += [f"const {var.ctype()} {var.emit()} = {constant_str(var.constant)};\n"] |
| 4908 | |
| 4909 | if line_directive := adj.get_line_directive(lines[-1], var.relative_lineno): |
| 4910 | lines.insert(-1, f"{line_directive}\n") |
| 4911 | |
| 4912 | # forward pass |
| 4913 | lines += ["//---------\n"] |
| 4914 | lines += ["// forward\n"] |
| 4915 | |
| 4916 | for f in adj.blocks[0].body_forward: |
| 4917 | if func_type == "kernel" and device == "cuda" and f.lstrip().startswith("return;"): |
| 4918 | # Use of grid-stride loops in CUDA kernels requires that we convert return; to continue; |
| 4919 | lines += [f.replace("return;", "continue;") + "\n"] |
| 4920 | else: |
| 4921 | lines += [f + "\n"] |
| 4922 | |
| 4923 | return "".join(l.lstrip() if l.lstrip().startswith("#line") else indent_block + l for l in lines) |
| 4924 | |
| 4925 | |
| 4926 | def codegen_func_reverse(adj, func_type="kernel", device="cpu"): |