Skip to content

Commit 6a9b400

Browse files
Merge pull request #867 from KazumaProject/fix/zero-query-leftover
辞書の mmap化
2 parents 3a44964 + b0dbb57 commit 6a9b400

19 files changed

Lines changed: 628 additions & 99 deletions

File tree

app/src/main/java/com/kazumaproject/markdownhelperkeyboard/converter/ConnectionMatrix.kt

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,16 @@
11
package com.kazumaproject.markdownhelperkeyboard.converter
22

3+
import java.nio.ByteBuffer
4+
import java.nio.ByteOrder
35
import kotlin.math.sqrt
46

57
object ConnectionMatrix {
8+
interface CostTable {
9+
val matrixSize: Int
10+
val entryCount: Int
11+
fun cost(rid: Int, lid: Int): Int
12+
}
13+
614
fun inferMatrixSize(connectionIdsSize: Int): Int? {
715
if (connectionIdsSize <= 0) return null
816
val matrixSize = sqrt(connectionIdsSize.toDouble()).toInt()
@@ -15,4 +23,61 @@ object ConnectionMatrix {
1523

1624
fun inferMatrixSize(connectionIds: ShortArray): Int? =
1725
inferMatrixSize(connectionIds.size)
26+
27+
fun fromShortArray(
28+
connectionIds: ShortArray,
29+
matrixSize: Int = inferMatrixSize(connectionIds)
30+
?: error("connectionId.dat size ${connectionIds.size} is not a valid square matrix"),
31+
): CostTable = ShortArrayCostTable(connectionIds, matrixSize)
32+
33+
fun fromByteBuffer(byteBuffer: ByteBuffer): CostTable {
34+
val buffer = byteBuffer.slice().asReadOnlyBuffer().order(ByteOrder.BIG_ENDIAN)
35+
require(buffer.remaining() % Short.SIZE_BYTES == 0) {
36+
"connectionId.dat byte size ${buffer.remaining()} is not divisible by ${Short.SIZE_BYTES}"
37+
}
38+
val entryCount = buffer.remaining() / Short.SIZE_BYTES
39+
val matrixSize = inferMatrixSize(entryCount)
40+
?: error("connectionId.dat entry count $entryCount is not a valid square matrix")
41+
return ByteBufferCostTable(buffer, matrixSize, entryCount)
42+
}
43+
44+
private class ShortArrayCostTable(
45+
private val connectionIds: ShortArray,
46+
override val matrixSize: Int,
47+
) : CostTable {
48+
override val entryCount: Int = connectionIds.size
49+
50+
override fun cost(rid: Int, lid: Int): Int {
51+
val index = indexOf(rid, lid, matrixSize, entryCount)
52+
return connectionIds[index].toInt()
53+
}
54+
}
55+
56+
private class ByteBufferCostTable(
57+
private val buffer: ByteBuffer,
58+
override val matrixSize: Int,
59+
override val entryCount: Int,
60+
) : CostTable {
61+
override fun cost(rid: Int, lid: Int): Int {
62+
val index = indexOf(rid, lid, matrixSize, entryCount)
63+
return buffer.getShort(index * Short.SIZE_BYTES).toInt()
64+
}
65+
}
66+
67+
private fun indexOf(
68+
rid: Int,
69+
lid: Int,
70+
matrixSize: Int,
71+
entryCount: Int,
72+
): Int {
73+
require(matrixSize > 0) { "connectionMatrixSize must be positive: $matrixSize" }
74+
require(rid in 0 until matrixSize && lid in 0 until matrixSize) {
75+
"connection id out of range: rid=$rid, lid=$lid, matrixSize=$matrixSize"
76+
}
77+
val index = rid * matrixSize + lid
78+
require(index in 0 until entryCount) {
79+
"connection index out of range: index=$index, size=$entryCount, matrixSize=$matrixSize"
80+
}
81+
return index
82+
}
1883
}

app/src/main/java/com/kazumaproject/markdownhelperkeyboard/converter/bitset/SuccinctBitVector.kt

Lines changed: 10 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -13,27 +13,19 @@ class SuccinctBitVector(private val bitSet: BitSet) {
1313
private val bigBlockRanks: IntArray
1414

1515
// 各大ブロック内を小ブロックに分割したときの、ブロック開始からの 1 数の差分
16-
private val smallBlockRanks: IntArray
17-
18-
// 8 ビットの数値(0~255)ごとに、その中の 1 のビット数(popcount)を記録するテーブル
19-
private val popCountTable: IntArray = IntArray(256)
16+
private val smallBlockRanks: ShortArray
2017

2118
// BitSet 全体の 1 の総数
2219
private val totalOnes: Int
2320

2421
private val n: Int = bitSet.size()
2522

2623
init {
27-
// 0~255 の各値の 1 の数を求める popcount テーブルの構築
28-
for (i in 0 until 256) {
29-
popCountTable[i] = Integer.bitCount(i)
30-
}
31-
3224
val n = bitSet.size()
3325
val numBigBlocks = (n + bigBlockSize - 1) / bigBlockSize
3426
bigBlockRanks = IntArray(numBigBlocks)
3527
val numSmallBlocks = (n + smallBlockSize - 1) / smallBlockSize
36-
smallBlockRanks = IntArray(numSmallBlocks)
28+
smallBlockRanks = ShortArray(numSmallBlocks)
3729

3830
var rank = 0
3931
// 大ブロックごとに累積値を計算
@@ -48,7 +40,7 @@ class SuccinctBitVector(private val bitSet: BitSet) {
4840
// 範囲外にならないようにチェック
4941
if (globalSmallIndex >= numSmallBlocks) break
5042
// 小ブロック開始時の「大ブロック内での累積 1 数」
51-
smallBlockRanks[globalSmallIndex] = rank - bigBlockRanks[big]
43+
smallBlockRanks[globalSmallIndex] = (rank - bigBlockRanks[big]).toShort()
5244
val smallStart = bigStart + small * smallBlockSize
5345
// 小ブロック内の各ビットを走査して 1 をカウント
5446
for (j in 0 until smallBlockSize) {
@@ -80,7 +72,7 @@ class SuccinctBitVector(private val bitSet: BitSet) {
8072

8173
// 大ブロックの累積値 + 小ブロック内での差分
8274
val rankBase =
83-
bigBlockRanks[bigIndex] + smallBlockRanks[bigIndex * numSmallBlocksPerBig + smallIndex]
75+
bigBlockRanks[bigIndex] + smallBlockRanks[bigIndex * numSmallBlocksPerBig + smallIndex].toInt()
8476
var additional = 0
8577
val smallBlockStart = bigIndex * bigBlockSize + smallIndex * smallBlockSize
8678
for (i in 0..offsetInSmall) {
@@ -130,7 +122,7 @@ class SuccinctBitVector(private val bitSet: BitSet) {
130122
val nextGlobalSmallIndex = baseSmallIndex + smallBlock + 1
131123
if (nextGlobalSmallIndex >= smallBlockRanks.size) break
132124

133-
if (smallBlockRanks[nextGlobalSmallIndex] < localTarget) {
125+
if (smallBlockRanks[nextGlobalSmallIndex].toInt() < localTarget) {
134126
smallBlock++
135127

136128
} else {
@@ -140,7 +132,7 @@ class SuccinctBitVector(private val bitSet: BitSet) {
140132

141133
// 小ブロック内をビット単位に走査して正確な位置を求める
142134
val globalSmallIndex = baseSmallIndex + smallBlock
143-
val offsetInSmallBlock = localTarget - smallBlockRanks[globalSmallIndex]
135+
val offsetInSmallBlock = localTarget - smallBlockRanks[globalSmallIndex].toInt()
144136
val smallBlockStart = bigBlock * bigBlockSize + smallBlock * smallBlockSize
145137
var count = 0
146138
for (i in 0 until smallBlockSize) {
@@ -182,22 +174,22 @@ class SuccinctBitVector(private val bitSet: BitSet) {
182174
val zerosBeforeBlock = bigBlock * bigBlockSize - bigBlockRanks[bigBlock]
183175
val localTarget = nodeId - zerosBeforeBlock
184176

185-
// 大ブロック内での小ブロック線形探索:
186177
// 大ブロック内での小ブロック線形探索:
187178
val baseSmallIndex = bigBlock * numSmallBlocksPerBig
188179
val smallBlocksInThisBig = min(numSmallBlocksPerBig, smallBlockRanks.size - baseSmallIndex)
189180
var smallBlock = 0
190181
while (smallBlock < smallBlocksInThisBig - 1) {
191182
val nextZeros =
192-
(smallBlock + 1) * smallBlockSize - smallBlockRanks[baseSmallIndex + smallBlock + 1]
183+
(smallBlock + 1) * smallBlockSize -
184+
smallBlockRanks[baseSmallIndex + smallBlock + 1].toInt()
193185
if (nextZeros < localTarget) {
194186
smallBlock++
195187
} else {
196188
break
197189
}
198190
}
199191
val globalSmallIndex = baseSmallIndex + smallBlock
200-
val zerosBeforeSmall = (smallBlock * smallBlockSize) - smallBlockRanks[globalSmallIndex]
192+
val zerosBeforeSmall = (smallBlock * smallBlockSize) - smallBlockRanks[globalSmallIndex].toInt()
201193
val offsetInSmallBlock = localTarget - zerosBeforeSmall
202194

203195
// 小ブロック内を 1 ビットずつ走査して目的の 0 を探す
@@ -215,4 +207,4 @@ class SuccinctBitVector(private val bitSet: BitSet) {
215207
}
216208
return -1 // 見つからなかった場合
217209
}
218-
}
210+
}

app/src/main/java/com/kazumaproject/markdownhelperkeyboard/converter/dictionary/TokenArray.kt

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,8 +25,10 @@ class TokenArray {
2525
private val nodeIdListTemp: MutableList<Int> = arrayListOf()
2626
private var bitListTemp: MutableList<Boolean> = arrayListOf()
2727
var bitvector: BitSet = BitSet()
28-
var leftIds: List<Short> = listOf()
29-
var rightIds: List<Short> = listOf()
28+
var leftIds: ShortArray = shortArrayOf()
29+
private set
30+
var rightIds: ShortArray = shortArrayOf()
31+
private set
3032

3133
fun getNodeIds(): IntArray {
3234
return nodeIdList
@@ -329,11 +331,18 @@ class TokenArray {
329331
objectInputStream: ObjectInputStream
330332
) {
331333
objectInputStream.apply {
332-
leftIds = (readObject() as ShortArray).toList()
333-
rightIds = (readObject() as ShortArray).toList()
334+
setPOSTable(
335+
leftIds = readObject() as ShortArray,
336+
rightIds = readObject() as ShortArray,
337+
)
334338
}
335339
}
336340

341+
fun setPOSTable(leftIds: ShortArray, rightIds: ShortArray) {
342+
this.leftIds = leftIds
343+
this.rightIds = rightIds
344+
}
345+
337346
fun readPOSTableWithIndex(
338347
objectInputStream: ObjectInputStream
339348
): Map<Pair<Short, Short>, Int> {

app/src/main/java/com/kazumaproject/markdownhelperkeyboard/converter/engine/KanaKanjiEngine.kt

Lines changed: 16 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -80,8 +80,7 @@ class KanaKanjiEngine {
8080
private lateinit var findPath: FindPath
8181
private var dictionaryBinaryReader: DictionaryBinaryReader? = null
8282

83-
private lateinit var connectionIds: ShortArray
84-
private var connectionMatrixSize: Int = 0
83+
private lateinit var connectionMatrix: ConnectionMatrix.CostTable
8584

8685
private lateinit var systemYomiTrie: LOUDSWithTermId
8786
private lateinit var systemTangoTrie: LOUDS
@@ -263,18 +262,16 @@ class KanaKanjiEngine {
263262
)
264263

265264
private data class ConnectionMatrixSnapshot(
266-
val connectionIds: ShortArray,
267-
val matrixSize: Int,
265+
val costTable: ConnectionMatrix.CostTable,
268266
)
269267

270268
private fun connectionMatrixSnapshot(): ConnectionMatrixSnapshot =
271269
synchronized(this) {
272-
check(connectionMatrixSize > 0) {
273-
"connectionMatrixSize must be initialized from connectionId.dat before use"
270+
check(::connectionMatrix.isInitialized) {
271+
"connectionMatrix must be initialized from connectionId.dat before use"
274272
}
275273
ConnectionMatrixSnapshot(
276-
connectionIds = connectionIds,
277-
matrixSize = connectionMatrixSize,
274+
costTable = connectionMatrix,
278275
)
279276
}
280277

@@ -338,8 +335,7 @@ class KanaKanjiEngine {
338335
findPath.backwardAStarWithBunsetsu(
339336
graph = graph,
340337
length = input.length,
341-
connectionIds = connectionMatrix.connectionIds,
342-
connectionMatrixSize = connectionMatrix.matrixSize,
338+
connectionMatrix = connectionMatrix.costTable,
343339
n = 1000,
344340
penaltyTrace = penaltyTrace,
345341
forwardDpTrace = forwardDpTrace,
@@ -364,9 +360,7 @@ class KanaKanjiEngine {
364360
val store = DictionaryOverrideStore(appContext, DictionaryOverrideValidator())
365361
DictionaryCompatibilityValidator(DictionarySourceResolver(appContext, store))
366362
.requireActiveStateCompatible()
367-
val newConnectionIds = reader.loadConnectionIds(DictionaryFileKey.CONNECTION_ID)
368-
val newConnectionMatrixSize = ConnectionMatrix.inferMatrixSize(newConnectionIds)
369-
?: error("connectionId.dat size ${newConnectionIds.size} is not a valid square matrix")
363+
val newConnectionMatrix = reader.loadConnectionMatrix(DictionaryFileKey.CONNECTION_ID)
370364
val newSystem = loadTripleDictionary(reader, DictionaryCategory.SYSTEM)
371365
val newSingleKanji = loadTripleDictionary(reader, DictionaryCategory.SINGLE_KANJI)
372366
val newEmoji = loadTripleDictionary(reader, DictionaryCategory.EMOJI)
@@ -397,8 +391,7 @@ class KanaKanjiEngine {
397391
reader.resolveCategoryLoadState(DictionaryCategory.SYSTEM) == DictionaryCategoryLoadState.Bundled
398392

399393
synchronized(this) {
400-
connectionIds = newConnectionIds
401-
connectionMatrixSize = newConnectionMatrixSize
394+
connectionMatrix = newConnectionMatrix
402395
assignSystemDictionary(newSystem)
403396
assignSingleKanjiDictionary(newSingleKanji)
404397
assignEmojiDictionary(newEmoji)
@@ -463,7 +456,7 @@ class KanaKanjiEngine {
463456
fun buildEngine(
464457
graphBuilder: GraphBuilder,
465458
findPath: FindPath,
466-
connectionIdList: ShortArray,
459+
connectionMatrix: ConnectionMatrix.CostTable,
467460

468461
systemTangoTrie: LOUDS,
469462
systemYomiTrie: LOUDSWithTermId,
@@ -542,10 +535,7 @@ class KanaKanjiEngine {
542535
)
543536

544537
// System
545-
val inferredConnectionMatrixSize = ConnectionMatrix.inferMatrixSize(connectionIdList)
546-
?: error("connectionId.dat size ${connectionIdList.size} is not a valid square matrix")
547-
this@KanaKanjiEngine.connectionIds = connectionIdList
548-
this@KanaKanjiEngine.connectionMatrixSize = inferredConnectionMatrixSize
538+
this@KanaKanjiEngine.connectionMatrix = connectionMatrix
549539
this@KanaKanjiEngine.systemTangoTrie = systemTangoTrie
550540
this@KanaKanjiEngine.systemTokenArray = systemTokenArray
551541
this@KanaKanjiEngine.systemYomiTrie = systemYomiTrie
@@ -1046,8 +1036,7 @@ class KanaKanjiEngine {
10461036
findPath.backwardAStar(
10471037
graph = graph,
10481038
length = input.length,
1049-
connectionIds = connectionMatrix.connectionIds,
1050-
connectionMatrixSize = connectionMatrix.matrixSize,
1039+
connectionMatrix = connectionMatrix.costTable,
10511040
n = n,
10521041
beamWidth = beamWidth,
10531042
)
@@ -1532,8 +1521,7 @@ class KanaKanjiEngine {
15321521
findPath.backwardAStarWithBunsetsu(
15331522
graph = graph,
15341523
length = input.length,
1535-
connectionIds = connectionMatrix.connectionIds,
1536-
connectionMatrixSize = connectionMatrix.matrixSize,
1524+
connectionMatrix = connectionMatrix.costTable,
15371525
n = n,
15381526
beamWidth = beamWidth,
15391527
)
@@ -2041,8 +2029,7 @@ class KanaKanjiEngine {
20412029
findPath.backwardAStarWithBunsetsu(
20422030
graph = graph,
20432031
length = input.length,
2044-
connectionIds = connectionMatrix.connectionIds,
2045-
connectionMatrixSize = connectionMatrix.matrixSize,
2032+
connectionMatrix = connectionMatrix.costTable,
20462033
n = n,
20472034
beamWidth = beamWidth,
20482035
)
@@ -2492,8 +2479,7 @@ class KanaKanjiEngine {
24922479
findPath.backwardAStar(
24932480
graph = graph,
24942481
length = input.length,
2495-
connectionIds = connectionMatrix.connectionIds,
2496-
connectionMatrixSize = connectionMatrix.matrixSize,
2482+
connectionMatrix = connectionMatrix.costTable,
24972483
n = n,
24982484
beamWidth = beamWidth,
24992485
)
@@ -2930,8 +2916,7 @@ class KanaKanjiEngine {
29302916
findPath.backwardAStar(
29312917
graph = graph,
29322918
length = input.length,
2933-
connectionIds = connectionMatrix.connectionIds,
2934-
connectionMatrixSize = connectionMatrix.matrixSize,
2919+
connectionMatrix = connectionMatrix.costTable,
29352920
n = n,
29362921
beamWidth = beamWidth,
29372922
)
@@ -3400,8 +3385,7 @@ class KanaKanjiEngine {
34003385
findPath.backwardAStarWithBunsetsu(
34013386
graph = graph,
34023387
length = input.length,
3403-
connectionIds = connectionMatrix.connectionIds,
3404-
connectionMatrixSize = connectionMatrix.matrixSize,
3388+
connectionMatrix = connectionMatrix.costTable,
34053389
n = n,
34063390
beamWidth = beamWidth,
34073391
)

0 commit comments

Comments
 (0)