Added errors on parse prolem instead of getting NoneType error in the next node over.
This commit is contained in:
@@ -82,8 +82,6 @@ class AudioMathNode:
|
||||
bv = b if b else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
|
||||
cv = c if c else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
|
||||
dv = d if d else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
|
||||
print("AudioMathNode: a shape:", a['waveform'].shape)
|
||||
|
||||
|
||||
B = getIndexTensorAlongDim(a['waveform'], 0)
|
||||
C = getIndexTensorAlongDim(a['waveform'], 1)
|
||||
@@ -105,6 +103,10 @@ class AudioMathNode:
|
||||
visitor = TensorEvalVisitor(variables, a['waveform'].shape)
|
||||
result_tensor = visitor.visit(tree)
|
||||
|
||||
print("Audio Tree\n" + tree.toStringTree(recog=parser))
|
||||
print("Result Tensor Shape: ", result_tensor.shape)
|
||||
print("Result Tensor: ", result_tensor)
|
||||
|
||||
# Create output dictionary with the same sample rate
|
||||
output = {
|
||||
'waveform': result_tensor,
|
||||
|
||||
@@ -3,6 +3,8 @@ from math import e
|
||||
|
||||
from antlr4 import CommonTokenStream, InputStream
|
||||
|
||||
from .helper_functions import ThrowingErrorListener
|
||||
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.FloatEvalVisitor import FloatEvalVisitor
|
||||
@@ -98,6 +100,7 @@ class FloatMathNode:
|
||||
lexer = MathExprLexer(input_stream)
|
||||
stream = CommonTokenStream(lexer)
|
||||
parser = MathExprParser(stream)
|
||||
parser.addErrorListener(ThrowingErrorListener())
|
||||
tree = parser.expr()
|
||||
print("Tree\n"+tree.toStringTree(recog=parser))
|
||||
visitor = FloatEvalVisitor(variables)
|
||||
|
||||
@@ -3,7 +3,7 @@ from inspect import cleandoc
|
||||
from antlr4 import CommonTokenStream
|
||||
from antlr4.atn.LexerActionExecutor import InputStream
|
||||
import torch
|
||||
from .helper_functions import getIndexTensorAlongDim
|
||||
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim
|
||||
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
@@ -111,6 +111,7 @@ class ImageMathNode:
|
||||
lexer = MathExprLexer(input_stream)
|
||||
stream = CommonTokenStream(lexer)
|
||||
parser = MathExprParser(stream)
|
||||
parser.addErrorListener(ThrowingErrorListener())
|
||||
tree = parser.expr()
|
||||
print("Tree\n"+tree.toStringTree(recog=parser))
|
||||
visitor = TensorEvalVisitor(variables,a.shape)
|
||||
|
||||
@@ -4,7 +4,7 @@ from math import e
|
||||
from antlr4 import CommonTokenStream, InputStream
|
||||
import torch
|
||||
|
||||
from .helper_functions import getIndexTensorAlongDim
|
||||
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim
|
||||
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
@@ -104,6 +104,7 @@ class LatentMathNode:
|
||||
lexer = MathExprLexer(input_stream)
|
||||
stream = CommonTokenStream(lexer)
|
||||
parser = MathExprParser(stream)
|
||||
parser.addErrorListener(ThrowingErrorListener())
|
||||
tree = parser.expr()
|
||||
print("Tree\n"+tree.toStringTree(recog=parser))
|
||||
visitor = TensorEvalVisitor(variables,a.shape)
|
||||
|
||||
@@ -4,7 +4,7 @@ from math import e
|
||||
from antlr4 import CommonTokenStream, InputStream
|
||||
import torch
|
||||
|
||||
from .helper_functions import getIndexTensorAlongDim
|
||||
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim
|
||||
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
@@ -145,6 +145,7 @@ class NoiseMathNode:
|
||||
lexer = MathExprLexer(input_stream)
|
||||
stream = CommonTokenStream(lexer)
|
||||
parser = MathExprParser(stream)
|
||||
parser.addErrorListener(ThrowingErrorListener())
|
||||
tree = parser.expr()
|
||||
print("Tree\n"+tree.toStringTree(recog=parser))
|
||||
visitor = TensorEvalVisitor(variables,variables['a'].shape)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from antlr4.error import ErrorListener
|
||||
import torch
|
||||
|
||||
def getIndexTensorAlongDim(tensor, dim):
|
||||
@@ -85,3 +86,6 @@ def freq_to_time(freq_dict: torch.Tensor, n_fft: int = 512, hop_length: int = 25
|
||||
waveform[b, c] = istft_result
|
||||
|
||||
return waveform
|
||||
class ThrowingErrorListener(ErrorListener):
|
||||
def syntaxError(self, recognizer, offendingSymbol, line, column, msg, e):
|
||||
raise ValueError(f"Syntax error in AudioExpr at line {line}, col {column}: {msg}")
|
||||
|
||||
Reference in New Issue
Block a user