diff --git a/python/cmusphinx/s3dict.py b/python/cmusphinx/s3dict.py index 1f80ecf4..ea7bfafe 100644 --- a/python/cmusphinx/s3dict.py +++ b/python/cmusphinx/s3dict.py @@ -12,6 +12,7 @@ SphinxTrain, Sphinx-III, and PocketSphinx. __author__ = "David Huggins-Daines " __version__ = "$Revision $" +from collections import defaultdict import re def open(file): @@ -21,11 +22,36 @@ class S3Dict(dict): """ Class for reading / processing Sphinx format dictionary files. """ - def __init__(self, infile=None): - self.phoneset = set() + def __init__(self, infile=None, preserve_alts=False): + self.preserve_alts = preserve_alts + self.phoneset = defaultdict(int) + self.maxalt = defaultdict(int) if infile != None: self.read(infile) + def __getitem__(self, key): + if not isinstance(key, tuple): + m = self.altre.match(key) + if m: + word, alt = m.groups() + return self.get_alt_phones(word, int(alt)) + else: + return self.get_phones(key) + else: + return self.get_alt_phones(*key) + + def __putitem__(self, key, val): + if not isinstance(key, tuple): + m = self.altre.match(key) + if m: + word, alt = m.groups() + return self.set_alt_phones(word, int(alt), val) + else: + self.set_phones(key, val) + else: + w, p = key + self.set_alt_phones(w, p, val) + altre = re.compile(r'(.*)\(([^\)]+)\)') def read(self, infile): """ @@ -45,14 +71,8 @@ class S3Dict(dict): spam = line.split() word = unicode(spam[0], 'utf8') phones = spam[1:] - m = self.altre.match(word) - if m: - word, alt = m.groups() - self[word,int(alt)] = phones - else: - self[word,1] = phones - for ph in phones: - self.phoneset.add(ph) + # Why can't we just say self[word] = phones? + self.__putitem__(word, phones) def write(self, outfile): """ @@ -70,24 +90,112 @@ class S3Dict(dict): word = "%s(%d)" % (word, alt) fh.write("%-30s %s\n" % (word, " ".join(self[k]))) - def __getitem__(self, key): - if not isinstance(key, tuple): - m = self.altre.match(key) - if m: - word, alt = m.groups() - return self.get((word,int(alt))) - else: - return self.get((key,1)) - else: - return self.get(key) + def get_phones(self, word): + """ + Get default pronunciation for word. - def getalts(self, key): - i = 1 - alts = [] - while True: - if (key,i) in self: - alts.append(self.get((key,i))) - else: - break - i += 1 - return alts + If word is not present, KeyError will be raised. + """ + return dict.__getitem__(self, (word, 1)) + + def get_alt_phones(self, word, alt): + """ + Get alternate pronunciaition #alt for word. + + Alternate pronunciations are numbered from 1, where 1 is the + default pronunciation. If word is not present, KeyError will + be raised. If no alternate pronunciation alt exists, + IndexError will be raised. + """ + + if (word, 1) not in self: + raise KeyError + elif (word,alt) not in self: + raise IndexError, "Alternate pronunciation index %d does not exist" % alt + return dict.__getitem__(self, (word, alt)) + + def set_phones(self, word, phones): + """ + Set default pronunciation for word. + """ + dict.__setitem__(self, (word, 1), phones) + self.maxalt[word] = 1 + for ph in phones: + self.phoneset[ph] += 1 # FIXME: should make a class for this + + def set_alt_phones(self, word, alt, phones): + """ + Set alternate pronunciaition #alt for word. + + If alt is greater than the maximum alternate pronunciation + index plus one for this dictionary, IndexError will be raised. + + """ + if alt > self.maxalt[word] + 1: + raise IndexError, "Alternate pronunciation index %d too high" % alt + dict.__setitem__(self, (word, alt), phones) + self.maxalt[word] = max(alt, self.maxalt[word]) + for ph in phones: + self.phoneset[ph] += 1 + + def add_alt_phones(self, word, phones): + """ + Add a new alternate pronunciation for word. + """ + dict.__setitem__(self, (word, self.maxalt[word] + 1), phones) + self.maxalt[word] += 1 + for ph in phones: + self.phoneset[ph] += 1 + + def del_alt_phones(self, word, alt): + """ + Delete alternate pronunciation alt for word. + + If no such alternate pronunciation exists, IndexError will be + raised. If this S3Dict was created with preserve_alts=True, + the indices of remaining alternate pronunciations will be + preserved (you can still this alternative index with + set_alt_phones()). Otherwise, the remaining alternate + pronunciations will be renumbered accordingly. + + """ + if (word, alt) not in self: + raise IndexError, "Alternate pronunciation index %d does not exist" % alt + for ph in self[word, alt]: + self.phoneset[ph] -= 1 + if self.phoneset[ph] == 0: # FIXME: make a class + del self.phoneset[ph] + del self[word, alt] + if alt == self.maxalt[word]: + self.maxalt[word] -= 1 + if not self.preserve_alts: + alts = list(self.alts(word)) + self.del_phones(word) + for i, phones in enumerate(alts): + dict.__setitem__(self, (word, i + 1), phones) + self.maxalt[word] = len(alts) + + def del_phones(self, word): + """ + Delete all pronunciations for word. + + If you only wish to delete the default pronunciation (it is + strongly suggested that you don't do this), use + del_alt_phones(word, 1). + """ + for i in range(1, self.maxalt[word] + 1): + if (word,i) in self: + for ph in self[word,i]: + self.phoneset[ph] -= 1 + if self.phoneset[ph] == 0: # FIXME: make a class + del self.phoneset[ph] + dict.__delitem__(self, (word, i)) + del self.maxalt[word] + + def alts(self, word): + """ + Iterate over alternative pronunciations for a word. + """ + for i in range(1, self.maxalt[word] + 1): + if (word,i) in self: + yield self[word,i] diff --git a/python/cmusphinx/test_s3dict.py b/python/cmusphinx/test_s3dict.py index 283eb9a7..76ca8b65 100644 --- a/python/cmusphinx/test_s3dict.py +++ b/python/cmusphinx/test_s3dict.py @@ -9,9 +9,55 @@ class TestS3Dict(unittest.TestCase): def setUp(self): self.basedir = os.path.dirname(__file__) - def testCreate(self): + def testRead(self): foodict = s3dict.open(os.path.join(self.basedir, "foo.dict")) + self.assert_('AH' in foodict.phoneset) + self.assertEquals(foodict.get_phones('A'), ['AH']) + self.assertEquals(foodict.get_alt_phones('A', 2), ['EY']) + self.assertEquals(foodict.get_phones('ZSWANG'), ['S', 'W', 'AE', 'NG']) + try: + foo = foodict.get_phones('QRXG') + print foo + except KeyError: + pass # Expected fail + else: + self.fail() + try: + foo = foodict.get_alt_phones('A',3) + except IndexError: + pass # Expected fail + else: + self.fail() + try: + foo = foodict.get_alt_phones('!@#$!@',3) + except KeyError: + pass # Expected fail + else: + self.fail() + self.assertEquals(foodict['A'], ['AH']) + self.assertEquals(foodict['A',2], ['EY']) + self.assertEquals(foodict['A(2)'], ['EY']) + self.assertEquals(foodict['ZSWANG'], ['S', 'W', 'AE', 'NG']) + def testCreate(self): + mydict = s3dict.S3Dict() + mydict.set_phones('A', ['AH']) + mydict.add_alt_phones('A', ['EY']) + mydict.set_phones('ZSWANG', ['S', 'W', 'AE', 'NG']) + mydict.set_alt_phones('A', 2, ['EY']) + try: + mydict.set_alt_phones('A', 5, ['AX']) + except IndexError: + pass # Expected fail + else: + self.fail() + self.assertEquals(mydict.get_phones('A'), ['AH']) + self.assertEquals(mydict.get_alt_phones('A', 2), ['EY']) + self.assertEquals(mydict.get_phones('ZSWANG'), ['S', 'W', 'AE', 'NG']) + mydict.set_alt_phones('A', 2, ['AA']) + self.assertEquals(mydict.get_alt_phones('A', 2), ['AA']) + mydict.del_phones('ZSWANG') + self.assert_('NG' not in mydict.phoneset) if __name__ == '__main__': unittest.main()