Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pykokkos/core/visitors/constructor_visitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -392,4 +392,4 @@ def _infer_view_type_from_array_call(
return view_type

def error(self, node: ast.AST, message: str):
visitors_util.error(self.src, self.debug, node, message)
raise visitors_util.TranslationError(self.src, self.debug, node, message)
8 changes: 3 additions & 5 deletions pykokkos/core/visitors/parameter_visitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@

from pykokkos.core import cppast

from . import visitors_util
from .visitors_util import TranslationError, get_type


class ParameterVisitor(ast.NodeVisitor):
Expand Down Expand Up @@ -87,9 +87,7 @@ def visit_arg(self, node: ast.arg) -> None:
annotation: Union[ast.Name, ast.Attribute] = node.annotation

declref = cppast.DeclRefExpr(node.arg)
decltype: Optional[cppast.Type] = visitors_util.get_type(
annotation, self.pk_import
)
decltype: Optional[cppast.Type] = get_type(annotation, self.pk_import)

if decltype is None:
self.error(node, "Type is not supported")
Expand All @@ -114,4 +112,4 @@ def visit_arg(self, node: ast.arg) -> None:
self.views[declref] = decltype

def error(self, node: ast.AST, message: str):
visitors_util.error(self.src, self.debug, node, message)
raise TranslationError(self.src, self.debug, node, message)
4 changes: 3 additions & 1 deletion pykokkos/core/visitors/pykokkos_visitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -742,7 +742,9 @@ def get_scratch_view_type(self, view_type: ast.Subscript) -> Optional[str]:
return cpp_view_type

def error(self, node, message):
visitors_util.error(self.src, self.debug, node, message, self.path)
raise visitors_util.TranslationError(
self.src, self.debug, node, message, self.path
)

def generic_error(self, node):
self.error(node, "Not supported for translation")
68 changes: 42 additions & 26 deletions pykokkos/core/visitors/visitors_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,8 @@
from pykokkos.interface import Layout, MemorySpace, Trait


def pretty_print(node):
print(ast.dump(node, indent=4))
def pretty_format_ast(node: Union[ast.Attribute, ast.Name]) -> str:
return ast.dump(node, indent=4)


allowed_types: Dict[str, str] = {
Expand Down Expand Up @@ -132,32 +132,48 @@ def pretty_print(node):
}


def error(src, debug: bool, node, message, path: Optional[str] = None) -> None:
if hasattr(node, "lineno"):
if path:
print(f"\n\033[31m\033[01mError in {path}:{node.lineno}\033[0m: {message}")
else:
print(f"\n\033[31m\033[01mError:{node.lineno} \033[0m: {message}")
else:
if path:
print(f"\n\033[31m\033[01mError in {path}\033[0m: {message}")
else:
print(f"\n\033[31m\033[01mError\033[0m: {message}")

if debug and node is not None:
print("DEBUG AST:")
pretty_print(node)

if hasattr(node, "lineno"):
print(src[0][node.lineno - src[1] - 1], end="")
err_len = node.end_col_offset - node.col_offset if node.end_col_offset else 1
print(" " * node.col_offset + "^" * err_len)

sys.exit("PyKokkos: Translation failed")
class TranslationError(Exception):
"""
PyKokkos code is not translatable to Kokkos
"""

def __init__(
self, src, debug: bool, node, message: str, path: Optional[str] = None
) -> None:
self.src = src
self.debug = debug
self.node = node
self.message = message
self.path = path

def __str__(self):
msg = ""
if hasattr(self.node, "lineno"):
if self.path:
msg += f"\n\033[31m\033[01mError in {self.path}:{self.node.lineno}\033[0m: {self.message}"
else:
msg += f"\n\033[31m\033[01mError:{self.node.lineno} \033[0m: {self.message}"
else:
if self.path:
msg += f"\n\033[31m\033[01mError in {self.path}\033[0m: {self.message}"
else:
msg += f"\n\033[31m\033[01mError\033[0m: {self.message}"

if self.debug and self.node is not None:
msg += "\nDEBUG AST:\n"
msg += pretty_format_ast(self.node)

if hasattr(self.node, "lineno"):
msg += self.src[0][self.node.lineno - self.src[1] - 1] + "\n"
err_len = (
self.node.end_col_offset - self.node.col_offset
if self.node.end_col_offset
else 1
)
msg += " " * self.node.col_offset + "^" * err_len

def generic_error(src, debug: bool, node) -> None:
error(src, debug, node, "Not supported for translation")
msg += "\nPyKokkos translation failed"
return msg


def get_op_str(op: ast.expr) -> str:
Expand Down
Loading