Commit 16dc09c7 authored by opitz's avatar opitz
Browse files

added training and testing scripts and a pretrained model

parent 74494e33
Loading
Loading
Loading
Loading

src/amr_helpers.py

0 → 100644
+636 −0
Changes for src/amr_helpers.py: 636 added lines, 0 removed lines.
Original line number Diff line number Diff line
#!/usr/bin/env python3
from copy import deepcopy
import os
import numpy as np
from collections import Counter
import re
import sys 
#import matplotlib.pyplot as plt
import networkx as nx
from pyparsing import Forward,Word,nestedExpr,alphanums,ParseException
#from networkx.drawing.nx_agraph import  graphviz_layout
from logger import SimpleLogger
import numpy as np
import spacy
from nltk import Tree
import json
import pprint

def tok_format(node):
    return " ".join([":"+node.dep_,node.orth_+"#ID#"+str(node.i)])#,node.pos_])

def to_nltk_tree(node):
    if node.n_lefts + node.n_rights > 0:
        return Tree(tok_format(node), [to_nltk_tree(child) for child in node.children])
    else:
        return tok_format(node)

nlp = spacy.load("en_core_web_sm")

def split_simple(string,superchars=[]):
    if not superchars:
        return string.replace("-","%").replace("_","%").split("%")
    else:
        ls=string.replace("_","%").split("%")
        new = []
        #print("a",ls)
        for string in ls:
            if string in superchars:
                new+=[string]
            else:
                new+=list(string)#.split('')
            new+=["-#-"]
        #print("b",new)
        return new



def splitchars(scs=[]):
    return None


class AmrDataSet:

    def __init__(self,path):
        self._build(self.path)

    def _build(self,path):
        amrs = [get_amr_meta_fromstring(string) for string in amrs][1:]
        print(amrs)
        amrs = [(nested_fromstring(amrstring),meta) for amrstring,meta in amrs]
        ls = []
        for nested,meta in amrs:
            print(nested)
            graph=amr_fromnested(nested)
            ls.append((graph,meta))



logger = SimpleLogger(0)

def read_file(p):
    with open(p,'r') as f:
        return f.read()

def read_amr_file(path):
    content = read_file(path)
    return [l for l in content.split("\n\n") if l]




def get_amr_meta_fromstring(string):
    logger.log(["parsing string, extracting amr and meta..."])
    splitted = string.split("\n")
    meta = {}
    amr = ""
    for spl in splitted:
        if spl and "#" == spl[0]:
            tokens = spl.split(" ")
            meta[tokens[1]] = " ".join(tokens[2:])
        else:
            amr+=spl
    logger.log(["successful"])
    #amr=amr.replace("\'","\\'")
    amr=re.sub(r"[^\\]\'"," \\'",amr)
    return amr,meta


def nested_fromstring(string):
    logger.log(["parsing amr to nested structure from string",string])
    nested_parens = nestedExpr('(', ')')#i, content=enclosed) 
    #print(string+"\n\n")
    string=" ".join(string.split())
    try:
        res = nested_parens.parseString(string).asList()
    except ParseException:
        #desparate tries to repair some borken AMRs generated from seq2seq models
        print("bracket missing, correcting...")
        print(string)
        nbr = int(len(string.split("(")))-int(len(string.split(")")))
        if int(len(string.split("("))) == 1 and int(len(string.split(")"))) == 1:
            string ="()"
        try:
            res = nested_parens.parseString(string+")"*nbr).asList()
        except ParseException:
            string="(e / error-01)"
            res = nested_parens.parseString(string).asList()
    logger.log(["successful",res])
    return res

def add_recursively(head="",rel="",this=[],graph=""):
    nextchilds = [(this[i-1],this[i]) for i in range(1,len(this)+1) if  ":" == this[i-1][0]]
    #single nodes consisting of strings
    if type(this) == str:
        tmpvar = this
        tmpconcept = this
    # else variable is at 1st index concept at 3rd
    else:
        tmpvar = this[0]
        #if len(this) > 1 and type(this) != str:
        tmpconcept=this[2]
    #basecase
    if not nextchilds:
        tmpvar 
        if tmpvar not in graph:
            if tmpconcept:
                graph.add_node(tmpvar,concept=tmpconcept)
        graph.add_edge(head,tmpvar,label=rel)
        return None
    #recursion case
    for child in nextchilds:
        if tmpvar not in graph:
            if tmpconcept:
                graph.add_node(tmpvar,concept=tmpconcept)
        graph.add_edge(head,tmpvar,label=rel)
        add_recursively(head=tmpvar,rel=child[0],this=child[1],graph=graph)

def amr_fromnested(ls):
    logger.log(["recursively building graph"])
    G=nx.DiGraph()
    root = [0,0]
    G.add_node("root",concept="root")
    add_recursively(head="root",rel="root",this=ls[0],graph=G)
    logger.log(["successful",nx.to_edgelist(G)])
    return G

def replace_vars_with_concepts(G):
    nodes_data = list(G.nodes(data=True))
    logger.log(["replacing variables in amr with concepts","nodes data:"+str(nodes_data)])
    G = nx.relabel_nodes(G,dict([(var,attr_dict["concept"]) for var,attr_dict in nodes_data]))
    logger.log(["successful",nx.to_edgelist(G)])
    return G

def explicit_polarity(G):
    pattern = re.compile(".+-[0-9]+$")
    newnodes = []
    newedges= []
    for node in G.nodes(data=True):
        if pattern.match(node[1]["concept"]):
            #check if already negated
            if G.has_edge(node[0],"-"):
                continue
            else:
                newnodes.append("+")
                newedges.append(node[0])
    for n in newnodes:
        G.add_node(n,concept="+")
    for e in newedges:
        G.add_edge(e,"+",label=":polarity")
    return None

def symbols(Gs):
    edge_labels = []
    node_labels = []
    for G in Gs:
        for (u,v) in G.edges():
            node_labels.append(G.node[u]["concept"])
            node_labels.append(G.node[v]["concept"])
            edge_labels.append(G.get_edge_data(u,v)["label"])
    return Counter(node_labels), Counter(edge_labels)


def get_prior_pol(syms):
    den = syms["-"]+syms["+"]
    return {"-":syms["-"]/den, "+": syms["+"]/den}

def norm_dist(dic):
    den = sum(dic.values())
    return {k:v/den for k,v in dic.items()}

def get_prior_wsd(syms):
    dic = {}
    pattern = re.compile(".+-[0-9]+$")
    ls = []
    for concept in syms:
        print(concept)
        if pattern.match(concept):
            ls.append("-".join(concept.split("-")[0:-1]))
    ls = set(ls)
    dic = {k:{} for k in ls}
    for concept in syms:
        if pattern.match(concept):
            tc = "-".join(concept.split("-")[0:-1])
            ts = concept.split("-")[-1]
            if ts in dic[tc]:
                dic[tc][ts] += 1
            else:
                dic[tc] = {ts:1}
    for key in dic:
        dic[key] = norm_dist(dic[key])
    return dic
"""
def amrs2nodefeas_and_adj(path="/home/students/opitz//Data/amr_deft_phase2/LDC2015E86_DEFT_Phase_2_AMR_Annotation_R1/data/amrs/split/test/deft-p2-amr-r1-amrs-test-consensus.txt"):
    amrs = read_amr_file(path)
    amrs = [get_amr_meta_fromstring(string) for string in amrs][1:]
    #print(amrs)
    amrs = [(nested_fromstring(amrstring),meta) for amrstring,meta in amrs if amrstring]
    Gs = []
    sents = []
    for nested,meta in amrs:
        #print(nested)
        graph=amr_fromnested(nested)
        explicit_polarity(graph)
        Gs.append(graph)
        sents.append(meta["::snt"])
    Ns = [nx.convert_node_labels_to_integers(x,label_attribute="stringlabel") for x in Gs]
    As = [nx.to_scipy_sparse_matrix(x) for x in Gs]
    syms = symbols(Gs)
    Gs = [[ x.node[u]["concept"] for u in x.nodes()] for x in Gs]
    return Gs, As, sents, syms
"""

def get_vocabulary_tokens(listoflistofstrings,splitf=list):
    lss=[[splitf(x) for x in l] for l in listoflistofstrings]
    vocab= {"PAD":1,"NONE":0}
    idx=2
    tokens=[]
    for ls in lss: 
        l1 = []
        for l in ls:
            l2 = []
            for string in l:
                if string in vocab:
                    l2.append(string)
                    continue
                else:
                    vocab[string]= idx
                    idx+=1
                    l2.append(string)
            l1.append(l2)
        tokens.append(l1)
    return vocab,tokens


def padcrop(seqs,length=25):
    seqs=[(length-len(seq))*["PAD"]+seq for seq in seqs]
    seqs = [seq[:length] for seq in seqs]
    return seqs


def idxmap(listoflistofstrings,vocab):
    out = []
    for listofstrings in listoflistofstrings:
        tmp=[]
        for string in listofstrings:
            if string in vocab:
                tmp.append(vocab[string])
            else:
                tmp.append(vocab["NONE"])
        out.append(tmp)
    return out
"""
def get_toy_data():
    Gs,As,sents,syms = amrs2nodefeas_and_adj("/home/students/opitz//Data/amr_deft_phase2/LDC2015E86_DEFT_Phase_2_AMR_Annotation_R1/data/amrs/split/training/deft-p2-amr-r1-amrs-training-xinhua.txt")
    pwsd,ppol = get_prior_wsd(syms[0]),get_prior_pol(syms[0])
    
    sents = [sents[i].split(" ")*len(Gs[i]) for i in range(len(Gs))]
    voc,tokenized = get_vocabulary_tokens(sents,splitf=lambda x:x)
    tokenized=[padcrop(toks) for toks in tokenized]
    idseqs_sents = [ np.array(idxmap(toks,voc)) for toks in tokenized]
    nwords = len(voc)

    voc,tokenized = get_vocabulary_tokens(Gs)
    tokenized=[padcrop(toks) for toks in tokenized]
    idseqs_nodes = [ np.array(idxmap(toks,voc)) for toks in tokenized]


    targets = []
    pattern = re.compile(".+-[0-9]+$")
    for G in Gs:
        tmp=[]
        for i,concept in enumerate(G):
            if concept in ["-","+"]:
                true = concept
                id_c = list(enumerate(ppol.keys()))
                probs = [ppol[key] for _,key in id_c]
                other = id_c[np.argmax(np.random.multinomial(1,probs))][1]
                G[i] = other
                if other == true:
                    tmp.append(1)
                else:
                    tmp.append(2)
                    print("changed to false",true,other)
            elif pattern.match(concept):
                tc = "-".join(concept.split("-")[0:-1])
                ts = concept.split("-")[-1]
                true = ts
                id_c = list(enumerate(pwsd[tc].keys()))
                probs = [pwsd[tc][key] for _,key in id_c]
                other = id_c[np.argmax(np.random.multinomial(1,probs))][1]
                G[i] = tc+"-"+other
                if other == true:
                    tmp.append(1)
                else:
                    tmp.append(2)
                    print("changed to false",tc,true,other)
            else:
                tmp.append(0)
        targets.append(tmp)
    
    #target = [np.eye(2)[np.random.randint(0,2,size=len(G))] for G in Gs]
    target = [np.eye(3)[targets[i]] for i in range(len(targets))]
    nchars = len(voc)
    return idseqs_nodes,As,target,idseqs_sents,nchars,nwords

"""

"""
def testing():
    amrs = read_amr_file("/home/students/opitz//Data/amr_deft_phase2/LDC2015E86_DEFT_Phase_2_AMR_Annotation_R1/data/amrs/split/test/deft-p2-amr-r1-amrs-test-consensus.txt")
    amrs = [get_amr_meta_fromstring(string) for string in amrs][1:]
    print(amrs)
    amrs = [(nested_fromstring(amrstring),meta) for amrstring,meta in amrs]
    ls = []
    for nested,meta in amrs:
        print(nested)
        graph=amr_fromnested(nested)
        explicit_polarity(graph)
        ls.append((graph,meta))
    print(ls[4],nx.to_dict_of_dicts(ls[4][0]))#,weight="label"))
    print(ls[4][0].nodes(data=True))
    print(nx.convert_node_labels_to_integers(ls[-1][0],label_attribute="stringlabel").nodes(data=True))
    print(nx.node_link_data(nx.convert_node_labels_to_integers(ls[-1][0],label_attribute="stringlabel")))
    print(nx.to_numpy_matrix(nx.convert_node_labels_to_integers(ls[-1][0],label_attribute="stringlabel")))
    
    print([nx.to_numpy_matrix(nx.convert_node_labels_to_integers(x,label_attribute="stringlabel")).shape for x,_ in ls] )
    print(symbols([g for g,_ in ls]))
    print(amrs2nodefeas_and_adj())
"""

def retrieve_eval(string):
    ls = [x for x in string.split("\n") if x]
    if len(ls) < 7:
        return None
    ls = [string.split("->") for string in ls]
    keys = [s[0] for s in ls]
    dic={}
    #print("---")
    for i,key in enumerate(keys):
        tmp=ls[i][1].split(",")
        #print(tmp)
        try:
            p=float(tmp[0].split()[1])
            r=float(tmp[1].split()[1])
            f=float(tmp[2].split()[1])
        except IndexError:
            print("warning... could not extract scores",tmp,"setting tmp scores to 0")
            p,r,f = 0.0,0.0,0.0
        dic[key] = [p,r,f]
    #print("---")
    return dic

def rec_gather_vars(nested=[],amrvars={}):
    if len(nested) > 1 and type(nested) != str:
        var = nested[0]
        try:
            con = nested[2]
        except IndexError:
            print("error",nested)
            sys.exit(1)
        if var not in amrvars:
            amrvars[var] = con
        rels = []
        childlists=[]
        for i,elm in enumerate(nested):
            if elm[0] == ":":
                rels.append(elm)
                #print(i,nested)
                childlists.append(nested[i+1])
        for j,rel in enumerate(rels):
            rec_gather_vars(childlists[j],amrvars)
    return None


def linearize_nested_amr_df(nested=[],string="",amrvars={}):
    #print(nested)
    try:
        tmpvar = nested[0]
        tmpconcept = nested[2]
    except IndexError:
        tmpvar=nested[0]
        tmpconcept=nested[2]
    if tmpconcept in amrvars:
        tmpconcept=amrvars[tmpconcept]
    nested.pop(0)
    nested.pop(0)
    rels = []
    childlists=[]
    for i,elm in enumerate(nested):
        if elm[0] == ":":
            if type(nested[i+1]) == str:
                if nested[i+1] in amrvars:
                    nested[i+1] = amrvars[nested[i+1]]
                elif nested[i+1] == "-":
                    nested[i+1] = "!"
                else:
                    continue
            else:
                rels.append(elm)
                childlists.append(nested[i+1])
    if not rels:
        return None
    else:
        for j,rel in enumerate(rels):
            if type(childlists[j]) == str:
                if childlists[j] in amrvars:
                    nested[j+1] = amrvars[childlists[j]]
                    print (childlists[j],"done",amrvars)
            else:
                linearize_nested_amr_df(childlists[j],amrvars=amrvars)



def super_chars_amr(amrstrings):
    ar = [[y for y in x.split("%") if y and y[0]==":"] for x in amrstrings]
    scs =[j for i in ar for j in i]
    return set(scs)

def super_chars_dep(depstrings):
    ar = [[y.split("_") for y in x.split("%")] for x in depstrings]
    ar = [[[w[0],w[2],w[3]] for w in y if len(w) > 3] for y in ar]
    ar = [[j for i in xs for j in i] for xs in ar]
    scs =[j for i in ar for j in i]
    return set(scs)


def treelinspacy(tmp=[],token=None,openbrackets=0):
    chs = list(token.children)
    if not chs:
        tmp.append(token.orth_)
        for i in range(openbrackets):
            tmp.append("]")
    string=""
    for ch in chs:
        tmp+=[ch.dep_,"["]
        openbrackets+=1
        treelinspacy(tmp,ch,openbrackets)


def alignamr(amrtokenstrings,tokenstrings):
    out=[]
    for at in amrtokenstrings:
        s = at.split("-")[0]
        if s in tokenstrings:
            i=tokenstrings.index(s)
        elif s == "[":
            i=550
        elif s== "]":
            i=551
        else:
            i=552
        out.append(str(i))
    return out

def rels(tokenstrings):
    out=[]
    for i,at in enumerate(tokenstrings):
        if at.startswith(":"):
            out.append(at)
        else:
            out.append("NOREL")
    return out

def aligndep(deptokenstrings):
    out=[]
    for at in deptokenstrings:
        s = at.split("#ID#")
        if len(s)==2:
            i=int(s[1])
            s=s[0]
        elif s == "[":
            i=550
        elif s== "]":
            i=551
        else:
            i=552
        out.append(str(i))
    return out

def create_dat(folder,prefixpred):
    dat = {}
    for f in os.listdir(folder):
        print(f)
        if "-" in f and ".txt" == f[-4:]:
            try:
                num = int(f.split(".")[0].split("-")[-1])
            except ValueError:
                print("Error, filename not allowed",f.split("."), "assuming unecessary file, continuing...")
                #sys.exit(1)
                continue
            if num not in dat:
                dat[num] = {}
        print("A")
        if prefixpred in f:
            print("A")
            if "eval" in f:
                ev = read_file(folder+f)
                res = retrieve_eval(ev)
                if not res:
                    dat[num]["eval"] = "NA"
                dat[num]["eval"] = res
            else:
                amr,meta = get_amr_meta_fromstring(read_amr_file(folder+f)[0])
                amrpure = read_file(folder+f)
                amr = nested_fromstring(amr)[0]
                amr_c = deepcopy(amr)
                amrvars = {}
                rec_gather_vars(amr_c,amrvars)
                linearize_nested_amr_df(amr_c,amrvars=amrvars)
                amr_c =  str(amr_c).replace("\'","").replace("[","[ ").replace("]"," ]").replace(",","")
                dat[num]["amr_lin"] = amr_c
                dat[num]["amr_pure"] = amrpure
                sent = meta["::snt"]
                doc = nlp(sent)
                dat[num]["tokens"] = [tok.text for tok in doc]
                dat[num]["lemmas"] = [tok.lemma_ for tok in doc]
                dat[num]["pos"] = [tok.pos_ for tok in doc]
                dl = " ".join(str(to_nltk_tree(list(doc.sents)[0].root)).split()).replace("(","[ ").replace(")"," ]")
                #dl = dl.replace("[ ","")
                #dl = dl.split("_")
                #dl = " [ ".join(dl)
                #print(dl)
                dl=dl.split(" ")
                for i in range(len(dl)):
                    if dl[i] == "[" and dl[i+1].startswith(":"):
                        tmp=dl[i]
                        dl[i] = dl[i+1]
                        dl[i+1] = tmp


                dat[num]["dep_lin"] = " ".join(dl)#" ".join(str(to_nltk_tree(list(doc.sents)[0].root)).split()).replace("(","[ ").replace(")"," ]")
    for num in dat:
        print(num,dat[num])
        toks = dat[num]["lemmas"]
        dat[num]["amralign"] = " ".join(alignamr(dat[num]["amr_lin"].split(" "),toks))
        dat[num]["depalign"] = " ".join(aligndep(dat[num]["dep_lin"].split(" ")))
        dat[num]["amrrels"] = " ".join(rels(dat[num]["amr_lin"].split(" ")))
        dat[num]["deprels"] = " ".join(rels(dat[num]["dep_lin"].split(" ")))
        dat[num]["dep_lin_no_ids"] = re.sub(r"#ID#[0-9]+","",dat[num]["dep_lin"])
    return dat

def make_dat(folder,prefixpred,p="exampledict-jamr.json.train"):
    dat=create_dat(folder,prefixpred)
    pprint.pprint(list(dat.items())[:10])
    dat={key:dat[key] for key in dat if dat[key]["eval"] != None}
    with open(p, "w") as f:
        f.write(json.dumps(dat,sort_keys=True,indent=4))

def linearize_amr_file(fp):
    for elm in read_amr_file(fp):
        amr,meta = get_amr_meta_fromstring(elm)
        amr = nested_fromstring(amr)[0]
        amr_c = deepcopy(amr)
        amrvars = {}
        rec_gather_vars(amr_c,amrvars)
        linearize_nested_amr_df(amr_c,amrvars=amrvars)
        amr_c =  str(amr_c).replace("\'","").replace("[","[ ").replace("]"," ]").replace(",","")
        print(amr_c)

def prepare4prediction(fp):
    dat={}
    for num,elm in enumerate(read_amr_file(fp)):
        amr,meta = get_amr_meta_fromstring(elm)
        #print(amr)
        if not amr:
            print("empty amr line...skipping....")
            continue
        dat[num] = {}
        amr = nested_fromstring(amr)[0]
        amr_c = deepcopy(amr)
        amrvars = {}
        rec_gather_vars(amr_c,amrvars)
        
        linearize_nested_amr_df(amr_c,amrvars=amrvars)
        amr_c =  str(amr_c).replace("\'","").replace("[","[ ").replace("]"," ]").replace(",","")
        dat[num]["amr_lin"] = amr_c
        dat[num]["amr_pure"] = elm
        sent = meta["::snt"]
        doc = nlp(sent)
        dat[num]["tokens"] = [tok.text for tok in doc]
        dat[num]["lemmas"] = [tok.lemma_ for tok in doc]
        dat[num]["pos"] = [tok.pos_ for tok in doc]
        dl = " ".join(str(to_nltk_tree(list(doc.sents)[0].root)).split()).replace("(","[ ").replace(")"," ]")
        #dl = dl.replace("[ ","")
        #dl = dl.split("_")
        #dl = " [ ".join(dl)
        #print(dl)
        dl=dl.split(" ")
        for i in range(len(dl)):
            if dl[i] == "[" and dl[i+1].startswith(":"):
                tmp=dl[i]
                dl[i] = dl[i+1]
                dl[i+1] = tmp


       #print(" ".join(dl))
        dat[num]["dep_lin"] = " ".join(dl)#" ".join(str(to_nltk_tree(list(doc.sents)[0].root)).split()).replace("(","[ ").replace(")"," ]")
        #asdasd 
    for num in dat:
        toks = dat[num]["lemmas"]
        dat[num]["amralign"] = " ".join(alignamr(dat[num]["amr_lin"].split(" "),toks))
        dat[num]["depalign"] = " ".join(aligndep(dat[num]["dep_lin"].split(" ")))
        dat[num]["amrrels"] = " ".join(rels(dat[num]["amr_lin"].split(" ")))
        dat[num]["deprels"] = " ".join(rels(dat[num]["dep_lin"].split(" ")))
        dat[num]["dep_lin_no_ids"] = re.sub(r"#ID#[0-9]+","",dat[num]["dep_lin"])
    return dat

src/config.py

0 → 100644
+70 −0
Changes for src/config.py: 70 added lines, 0 removed lines.
Original line number Diff line number Diff line
CONFIG = {
    "1": { 
         "SHAREW":False
        ,"HIER":True
        ,"LSTMLAYERS":2
        ,"DEP":True
        ,"POINTERS":True
        ,"SENSES":True
        }

    ,"2": { 
         "SHAREW":False
        ,"HIER":False
        ,"LSTMLAYERS":2
        ,"DEP":True
        ,"POINTERS":True
        ,"SENSES":True
        }

        
        
        
    ,"3": { 
         "SHAREW":False
        ,"HIER":True
        ,"LSTMLAYERS":1
        ,"DEP":True
        ,"POINTERS":True
        ,"SENSES":True
        }

        
        
        
    
    ,"4": { 
         "SHAREW":False
        ,"HIER":True
        ,"LSTMLAYERS":2
        ,"DEP":False
        ,"POINTERS":True
        ,"SENSES":True
        }

        
        
        
    
    ,"5": { 
         "SHAREW":False
        ,"HIER":True
        ,"LSTMLAYERS":2
        ,"DEP":True
        ,"POINTERS":False
        ,"SENSES":True
        

        }
    ,"6": { 
         "SHAREW":False
        ,"HIER":True
        ,"LSTMLAYERS":2
        ,"DEP":True
        ,"POINTERS":True
        ,"SENSES":False
        }

        

}

src/logger.py

0 → 100644
+12 −0
Changes for src/logger.py: 12 added lines, 0 removed lines.
Original line number Diff line number Diff line
#!/usr/bin/env python3

class SimpleLogger:

    def __init__(self,loglevel=0):
        self.loglevel=loglevel

    def log(self,msgs):
        for i in range(self.loglevel):
            if i < len(msgs):
                print("loglevel:",i+1," --> ",msgs[i])

src/model.py

0 → 100644
+99 −0

File added.

Preview size limit exceeded, changes collapsed.

src/models/0-model.h5

0 → 100644
+13.2 MiB

File added.

No diff preview for this file type.

Loading