{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59575,"databundleVersionId":8060720,"sourceType":"competition"},{"sourceId":8246447,"sourceType":"datasetVersion","datasetId":4892374},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":9019440,"sourceType":"datasetVersion","datasetId":5435081},{"sourceId":9019578,"sourceType":"datasetVersion","datasetId":5435180},{"sourceId":9019591,"sourceType":"datasetVersion","datasetId":5435193},{"sourceId":174185912,"sourceType":"kernelVersion"}],"dockerImageVersionId":30664,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!mkdir cppinput\n!ln -s /kaggle/input/uspto-100000/detd.txt cppinput/detd.txt \n# !ln -s /kaggle/input/uspto-all/ab.txt cppinput/ab.txt\n!ln -s /kaggle/input/uspto-all/clm.txt cppinput/clm.txt\n!ln -s /kaggle/input/uspto-all/cpc.txt cppinput/cpc.txt\n!ln -s /kaggle/input/uspto-all/ti.txt cppinput/ti.txt\n!ln -s /kaggle/input/uspto-100000/ab.txt cppinput/ab.txt\n# !ln -s /kaggle/input/uspto-100000/clm.txt cppinput/clm.txt\n# !ln -s /kaggle/input/uspto-100000/cpc.txt cppinput/cpc.txt\n# !ln -s /kaggle/input/uspto-100000/ti.txt cppinput/ti.txt","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:05:34.190063Z","iopub.execute_input":"2024-07-24T02:05:34.190485Z","iopub.status.idle":"2024-07-24T02:05:41.288817Z","shell.execute_reply.started":"2024-07-24T02:05:34.190452Z","shell.execute_reply":"2024-07-24T02:05:41.287277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile main.cpp\n\n#line 1 \"main.cpp\"\n\n#undef _GLIBCXX_DEBUG  // disable run-time bound checking, etc\n#define NDEBUG\n\n// #define _GLIBCXX_DEBUG\n\n// #define DEBUG\n#define USE_CPC\n#define USE_TI\n#define USE_CLM\n#define USE_AB\n#define USE_DETD\n\n#define USE_ALL_SAMPLES\n\n#define TIME_LIMIT_INPUT 1900\n#define TIME_RATIO_INPUT 1.5\n#define LOG_TABLE_SIZE_INPUT (1 << 13)\n// #define PRINT_MAX_STACK_SIZE\n\n#line 2 \"library/include.hpp\"\n\n#include <algorithm>\n#include <array>\n#include <bitset>\n#include <cassert>\n#include <chrono>\n#include <cmath>\n#include <cstring>\n// #include <format>\n#include <filesystem>\n#include <fstream>\n#include <iomanip>\n#include <iostream>\n#include <iterator>\n#include <map>\n#include <memory>\n#include <mutex>\n#include <numeric>\n#include <queue>\n#include <random>\n#include <set>\n#include <sstream>\n#include <string>\n#include <thread>\n#include <unordered_map>\n#include <unordered_set>\n#include <vector>\n#line 2 \"library/variable.cpp\"\n\n#ifdef ONLINE_JUDGE\nconstexpr double TIME_RATIO = 1.0;\n#else\n#ifdef TIME_RATIO_INPUT\nconstexpr double TIME_RATIO = TIME_RATIO_INPUT;\n#else\nconstexpr double TIME_RATIO = 1.0;\n#endif\n#endif\n\nconstexpr int INF = 1e9;\nconstexpr long long INFLL = 1ll << 60;\n#ifdef TIME_LIMIT_INPUT\nconstexpr int TIME_LIMIT = TIME_LIMIT_INPUT;\n#else\nconstexpr int TIME_LIMIT = 1900;  // 問題固有\n#endif\n#line 4 \"library/util.cpp\"\n\n#line 6 \"library/util.cpp\"\n\nusing namespace std;\nusing std::cin;\nusing std::cout;\nusing uint = unsigned int;\nusing i8 = signed char;\nusing u8 = unsigned char;\nusing i16 = signed short;\nusing u16 = unsigned short;\nusing i32 = signed int;\nusing u32 = unsigned int;\nusing i64 = signed long long;\nusing u64 = unsigned long long;\nusing ll = long long;\nusing ull = unsigned long long;\nusing pii = pair<int, int>;\n\n\nnamespace ethy {\n// basic\ntemplate <class T>\nbool chmin(T& a, const T& b) {\n    if (a > b) {\n        a = b;\n        return 1;\n    } else\n        return 0;\n}\ntemplate <class T>\nbool chmax(T& a, const T& b) {\n    if (a < b) {\n        a = b;\n        return 1;\n    } else\n        return 0;\n}\n\n// time\nauto time_start = chrono::system_clock::now();\nint get_time() {\n    int elapsed\n        = (int)(TIME_RATIO\n                * double(\n                    chrono::duration_cast<chrono::milliseconds>(chrono::system_clock::now() - time_start).count()));\n    return elapsed;\n}\nint get_time_remaining() {\n    int time_remaining = TIME_LIMIT - get_time();\n    return time_remaining;\n}\ndouble get_time_remaining_ratio() { return (double)get_time_remaining() / TIME_LIMIT; }\nvoid print_time() {\n    int elapsed = get_time();\n    cerr << \"[DATA] totaltime = \" << elapsed << endl;\n}\n\n// random\nclass RandomGenerator {\npublic:\n    RandomGenerator() {\n        random_mw_ = 1999;\n        random_mz_ = 2024;\n    }\n    void reset() {\n        random_mw_ = 1999;\n        random_mz_ = 2024;\n    }\n    inline int int_random() {\n        update_mw();\n        update_mz();\n        unsigned result = (random_mz_ << 16) + random_mw_;\n        return result >> 1;\n    }\n    inline uint uint_random_65535() {\n        update_mw();\n        return random_mw_ & 65535;\n    }\n    inline uint uint_random_65535_unsafe() {\n        const unsigned tmp = random_mw_ & 65535;\n        random_mw_ = 18000 * tmp + (random_mw_ >> 16);\n        return tmp;\n    }\n    inline uint uint_random() {\n        update_mw();\n        update_mz();\n        unsigned result = (random_mz_ << 16) + random_mw_;\n        return result;\n    }\n    inline double double_random() {\n        // 0 <= result < 1\n        unsigned result = uint_random();\n        return (result + 1.0) * 2.328306435454494e-10;\n    }\n    inline int range_random_safe(const int& r) {\n        assert(r > 0);\n        return uint_random() % r;\n    }\n    inline int range_random_safe(const int& l, const int& r) {\n        assert(l < r);\n        return l + range_random_safe(r - l);\n    }\n    inline int range_random(const int& r) {\n        assert(r > 0);\n        return (ull(uint_random()) * r) >> 32;\n    }\n    inline int range_random(const int& l, const int& r) {\n        assert(l < r);\n        return l + range_random(r - l);\n    }\n    inline int range_random_65535(const int& r) {\n        assert(r > 0);\n        assert(r < 65536 / 2);\n        return (uint_random_65535() * r) >> 16;\n    }\n    inline int range_random_65535(const int& l, const int& r) {\n        assert(l < r);\n        return l + range_random_65535(r - l);\n    }\n    inline int range_random_65535_unsafe(const int& r) {\n        assert(r > 0);\n        assert(r < 65536 / 2);\n        return (uint_random_65535_unsafe() * r) >> 16;\n    }\n    inline int range_random_65535_unsafe(const int& l, const int& r) {\n        assert(l < r);\n        return l + range_random_65535_unsafe(r - l);\n    }\n    template <auto R>\n    inline int range_random() {\n        assert(R > 0);\n        return (ull(uint_random()) * R) >> 32;\n    }\n    template <auto L, auto R>\n    inline int range_random() {\n        assert(L < R);\n        return L + range_random(R - L);\n    }\n    template <auto R>\n    inline int range_random_65535() {\n        assert(R > 0);\n        assert(R < 65536 / 2);\n        return (uint_random_65535() * R) >> 16;\n    }\n    template <auto L, auto R>\n    inline int range_random_65535() {\n        assert(L < R);\n        return L + range_random_65535(R - L);\n    }\n    template <auto R>\n    inline int range_random_65535_unsafe() {\n        assert(R > 0);\n        assert(R < 65536 / 2);\n        return (uint_random_65535_unsafe() * R) >> 16;\n    }\n    template <auto L, auto R>\n    inline int range_random_65535_unsafe() {\n        assert(L < R);\n        return L + range_random_65535_unsafe(R - L);\n    }\n    template <class T>\n    inline bool bool_random(const T& p) {\n        assert(p >= 1. / 65536);\n        return uint_random_65535() < (p * 65536);\n    }\n    template <class T>\n    inline bool bool_random_unsafe(const T& p) {\n        assert(p >= 1. / 65536);\n        return uint_random_65535_unsafe() < (p * 65536);\n    }\n    template <class T>\n    inline bool bool_random_safe(const T& p) {\n        return uint_random() < (p * 4294967296);\n    }\n    template <auto P>\n    inline bool bool_random() {\n        assert(P >= 1. / 65536);\n        return uint_random_65535() < (P * 65536);\n    }\n    template <auto P>\n    inline bool bool_random_unsafe() {\n        assert(P >= 1. / 65536);\n        return uint_random_65535_unsafe() < (P * 65536);\n    }\n    template <auto P>\n    inline bool bool_random_safe() {\n        return uint_random() < (P * 4294967296);\n    }\n    template <class T>\n    inline void shuffle_vector(vector<T>& vec) {\n        for (int i = (int)vec.size() - 1; i > 0; i--) {\n            swap(vec[i], vec[range_random(i + 1)]);\n        }\n    }\n\nprivate:\n    uint random_mw_;\n    uint random_mz_;\n\n    inline void update_mw() { random_mw_ = 18000 * (random_mw_ & 65535) + (random_mw_ >> 16); }\n    inline void update_mz() { random_mz_ = 36969 * (random_mz_ & 65535) + (random_mz_ >> 16); }\n};\nRandomGenerator rng;\n\n\n}  // namespace ethy\n#line 23 \"main.cpp\"\n\nusing namespace std;\nusing namespace ethy;\nnamespace fs = std::filesystem;\n\nstd::mutex mtx;\n\n// variables\nint N_total = -1;\nint N_files;\nvector<int> num_tokens;                       // number of tokens in each file\nvector<vector<int>> data_len;                 // data_len[file][row]\nvector<vector<vector<int>>> data_tokens;      // data[file][row][token], sorted\nvector<vector<vector<int>>> data_tokens_inv;  // data[file][token][:]\n\nvector<string> token_files;\nvector<string> token_names;\n\nint N_neighbors;\nconstexpr int Neighbor_size = 50;\nvector<vector<int>> data_neighbors;  // data_neighbors[query][:]\n\nstring dir_input;\nstring dir_output;\n\n// TODO\nint seed_solve_begin;\nint seed_solve_step;\n\nvector<int> file_costs_one;\n\n// cutoff for available tokens\nconstexpr int MAX_USE_ALL_MIN = 100000;\nconstexpr int MAX_USE_ALL_MAX = 300000;\nconstexpr int MAX_USE_ALL_DEFAULT = MAX_USE_ALL_MAX;\nconstexpr int MIN_USE_NEIGHBOR = 2;\nconstexpr int TIME_LIMIT_ALL = 500;\nconstexpr double ITER_RATIO = 4;\nconstexpr int ITER_BASE_MAX = 1000;\n\nconstexpr int ALLOW_INVALID = 10;\nconstexpr int MAX_DEPTH = 10;\nconstexpr int ITER_MAX = 1250;\nconstexpr int QUERIES_MAX = 1250;\n\nvector<int> vector_intersection_fast(const vector<int>& a, const vector<int>& b) {\n    vector<int> res;\n\n    const int A = (int)a.size();\n    const int B = (int)b.size();\n\n    if (A == 0 || B == 0) {\n        return {};\n    }\n\n    int ApB = A + B;\n    int AlogB = int(A * int(log(B) + 3));\n    int BlogA = int(B * int(log(A) + 3));\n    int MIN = min({ApB, AlogB, BlogA});\n\n    if (MIN == ApB) {\n        int idx_a = 0;\n        int idx_b = 0;\n        while (idx_a < (int)a.size() && idx_b < (int)b.size()) {\n            if (a[idx_a] == b[idx_b]) {\n                res.emplace_back(a[idx_a]);\n                idx_a++;\n                idx_b++;\n            } else if (a[idx_a] < b[idx_b]) {\n                idx_a++;\n            } else {\n                idx_b++;\n            }\n        }\n    } else if (MIN == AlogB) {\n        for (int i = 0; i < (int)a.size(); i++) {\n            if (binary_search(b.begin(), b.end(), a[i])) {\n                res.emplace_back(a[i]);\n            }\n        }\n    } else if (MIN == BlogA) {\n        for (int i = 0; i < (int)b.size(); i++) {\n            if (binary_search(a.begin(), a.end(), b[i])) {\n                res.emplace_back(b[i]);\n            }\n        }\n    } else {\n        assert(false);\n    }\n    return res;\n}\n\nclass Query {\npublic:\n    // int score; // スコアは外部で保持\n    int token_type;\n    vector<int> tokens;\n    vector<int> neighbors;\n    int score_invalid;\n    Query() = default;\n    Query(int n_file_processing) {\n        token_type = n_file_processing;\n        score_invalid = -1;\n    }\n    Query(const vector<int>& tokens, const vector<int>& neighbors, int score_invalid, int n_file_processing)\n        : tokens(tokens), neighbors(neighbors), score_invalid(score_invalid) {\n        token_type = n_file_processing;\n    }\n    void print(ostream& os = cerr) const {\n        os << token_names[token_type] << \" \";\n        os << (int)tokens.size() << \" \";\n        for (int i = 0; i < (int)tokens.size(); i++) {\n            os << tokens[i] << \" \";\n        }\n        os << (int)neighbors.size() << \" \";\n        for (int i = 0; i < (int)neighbors.size(); i++) {\n            os << neighbors[i] << \" \";\n        }\n        os << score_invalid;\n        os << endl;\n    }\n\n    void sort() {\n        std::sort(tokens.begin(), tokens.end());\n        std::sort(neighbors.begin(), neighbors.end());\n    }\n};\n\nvoid remove_empty_query(vector<Query> queries) {\n    int N = (int)queries.size();\n    for (int i = N - 1; i >= 0; i--) {\n        Query& query = queries[i];\n        if (query.tokens.empty() || query.neighbors.empty()) {\n            swap(query, queries.back());\n            queries.pop_back();\n        }\n    }\n}\n\nclass Solver {\npublic:\n    int seed_solve;\n    string file_output;\n\n    vector<Query> queries;\n\n    vector<int> neighbors;\n\n    vector<int> neighbors_processing;\n    vector<int> neighbors_flag;\n\n    tuple<int, int, vector<int>> calc_valid_invalid_neighbors(\n        const vector<int>& docs, const bool flag_no_upper_limit_invalid = false) {\n        int valid = 0;\n        int invalid = 0;\n        vector<int> neighbors;\n        for (const auto& doc : docs) {\n            if (neighbors_flag[doc] == -1) {\n                invalid++;\n            } else {\n                valid++;\n                neighbors.emplace_back(neighbors_flag[doc]);\n            }\n            if (!flag_no_upper_limit_invalid && invalid > ALLOW_INVALID) {\n                break;\n            }\n        }\n        if (!flag_no_upper_limit_invalid && valid == 0) {\n            // invalid = ALLOW_INVALID + 1;\n        }\n        return {valid, invalid, neighbors};\n    }\n\n    vector<Query> get_queries_from_tokens(\n        const int n_file_processing, const vector<int>& tokens_raw, const bool flag_no_upper_limit_invalid = false) {\n        // flag_no_upper_limit_invalid: true の場合、invalid の上限を無視する\n        // 一番良いクエリのみを返す\n        if (tokens_raw.empty()) {\n            return {};\n        }\n        if (tokens_raw.size() == 1) {\n            const int token_raw = tokens_raw[0];\n            auto [valid, invalid, neighbors] = calc_valid_invalid_neighbors(\n                data_tokens_inv[n_file_processing][token_raw], flag_no_upper_limit_invalid);\n            if ((valid == 0) || (!flag_no_upper_limit_invalid && (invalid > ALLOW_INVALID))) {\n                return {};\n            }\n            // always valid\n            Query query(n_file_processing);\n            query.neighbors = neighbors;\n            query.score_invalid = invalid;\n            query.tokens = tokens_raw;\n            return {query};\n        }\n        // number of tokens >= 2;\n        vector<Query> queries;\n        vector<int> docs = vector_intersection_fast(\n            data_tokens_inv[n_file_processing][tokens_raw[0]], data_tokens_inv[n_file_processing][tokens_raw[1]]);\n        vector<int> tokens_raw_current = {tokens_raw[0], tokens_raw[1]};\n        auto [valid, invalid, neighbors] = calc_valid_invalid_neighbors(docs, flag_no_upper_limit_invalid);\n        int invalid_bef = invalid;\n        if ((valid != 0) && (flag_no_upper_limit_invalid || (invalid <= ALLOW_INVALID))) {\n            Query query(tokens_raw_current, neighbors, invalid, n_file_processing);\n            queries.emplace_back(query);\n        }\n        for (int i = 2; i < (int)tokens_raw.size(); i++) {\n            int token = tokens_raw[i];\n            docs = vector_intersection_fast(docs, data_tokens_inv[n_file_processing][token]);\n            auto [valid, invalid, neighbors] = calc_valid_invalid_neighbors(docs, flag_no_upper_limit_invalid);\n            if (valid == 0)\n                break;\n            tokens_raw_current.emplace_back(token);\n            if (chmin(invalid_bef, invalid)) {\n                if (flag_no_upper_limit_invalid || (invalid <= ALLOW_INVALID)) {\n                    Query query(tokens_raw_current, neighbors, invalid, n_file_processing);\n                    queries.emplace_back(query);\n                }\n            }\n            if (invalid == 0) {\n                break;\n            }\n        }\n        if (flag_no_upper_limit_invalid && !queries.empty()) {\n            return {queries.back()};\n        }\n        return queries;\n    }\n\n\n    vector<Query> get_queries_from_neighbors_indicies(\n        const vector<int>& neighbor_indicies, const int n_file_processing) {\n        if (neighbor_indicies.empty()) {\n            return {};\n        }\n        if (neighbor_indicies.size() == 1) {\n            const int nei_0 = neighbors_processing[neighbor_indicies[0]];\n            const vector<int>& tokens = data_tokens[n_file_processing][nei_0];\n            vector<Query> queries = get_queries_from_tokens(n_file_processing, tokens);\n            remove_empty_query(queries);\n            if (queries.empty()) {\n                queries = get_queries_from_tokens(n_file_processing, tokens, true);\n            }\n            return queries;\n        }\n        const int nei_0 = neighbors_processing[neighbor_indicies[0]];\n        const int nei_1 = neighbors_processing[neighbor_indicies[1]];\n        vector<int> tokens\n            = vector_intersection_fast(data_tokens[n_file_processing][nei_0], data_tokens[n_file_processing][nei_1]);\n        for (int i = 2; i < (int)neighbor_indicies.size(); i++) {\n            int nei = neighbors_processing[neighbor_indicies[i]];\n            tokens = vector_intersection_fast(tokens, data_tokens[n_file_processing][nei]);\n        }\n        vector<Query> queries = get_queries_from_tokens(n_file_processing, tokens);\n        remove_empty_query(queries);\n        if (queries.empty() && neighbor_indicies.size() == 2) {\n            queries = get_queries_from_tokens(n_file_processing, tokens, true);\n        }\n        return queries;\n    }\n\n    void init(int _seed_solve) {\n        queries.clear();\n        seed_solve = _seed_solve;\n        cerr << \"seed_solve = \" << seed_solve << endl;\n        file_output = fs::path(dir_output) / (to_string(seed_solve) + \".txt\");\n#ifdef DEBUG\n        cerr << \"file_output = \" << file_output << endl;\n#endif\n        vector<int> neighbors = data_neighbors[seed_solve];\n        neighbors_flag.resize(N_total);\n        fill(neighbors_flag.begin(), neighbors_flag.end(), -1);\n        for (int i = 0; i < (int)neighbors.size(); i++) {\n            neighbors_flag[neighbors[i]] = i;\n        }\n        neighbors_processing = data_neighbors[seed_solve];\n    }\n\n    Solver(int seed_solve) : seed_solve(seed_solve) { init(seed_solve); }\n\n    void solve2_dfs(const int depth_target, const int iterations_max, int& iterations_current,\n        vector<int>& neighbors_current, set<vector<int>>& set_ok, const int n_file_processing) {\n        if (depth_target == 0) {\n            return;\n        }\n        if (iterations_current > iterations_max) {\n            return;\n        }\n        int depth_current = (int)neighbors_current.size();\n        if (depth_current == depth_target) {\n            iterations_current++;\n            vector<Query> queries_new = get_queries_from_neighbors_indicies(neighbors_current, n_file_processing);\n            if (!queries_new.empty()) {\n                vector<int> vec = neighbors_current;\n                set_ok.insert(vec);\n                copy(queries_new.begin(), queries_new.end(), back_inserter(queries));\n            }\n            return;\n        }\n        // assert(depth_current == 0 || set_ok.count(neighbors_current) > 0);\n        if (depth_current >= 0 && set_ok.count(neighbors_current) == 0) {\n            cerr << \"invalid neighbors count\" << endl;\n            for (auto& n : neighbors_current) {\n                cerr << n << \" \";\n            }\n            cerr << endl;\n            exit(1);\n        }\n        int nei_start;\n        if (neighbors_current.empty()) {\n            nei_start = 0;\n        } else {\n            nei_start = neighbors_current.back() + 1;\n        }\n        for (int nei_new = nei_start; nei_new < Neighbor_size; nei_new++) {\n            neighbors_current.emplace_back(nei_new);\n\n            // check if all subsets are ok\n            bool flag_ok = true;\n            if ((int)neighbors_current.size() < depth_target) {\n                if (set_ok.count(neighbors_current) == 0) {\n                    flag_ok = false;\n                }\n            }\n            if (flag_ok) {\n                for (int nei : neighbors_current) {\n                    vector<int> neighbors_tmp = neighbors_current;\n                    // delete\n                    neighbors_tmp.erase(find(neighbors_tmp.begin(), neighbors_tmp.end(), nei));\n                    if (set_ok.count(neighbors_tmp) == 0) {\n                        flag_ok = false;\n                        break;\n                    }\n                }\n            }\n            // flag_ok = true;\n            if (flag_ok)\n                solve2_dfs(\n                    depth_target, iterations_max, iterations_current, neighbors_current, set_ok, n_file_processing);\n            neighbors_current.pop_back();\n        }\n    }\n\n    void solve2() {\n        // int time_start = get_time();\n        // int loop = 0;\n        for (int n_file = 0; n_file < N_files; n_file++) {\n            cerr << \"n_file = \" << n_file << endl;\n            int queries_size_bef = (int)queries.size();\n            const int iterations_max = ITER_MAX;\n            int iterations_current = 0;\n            vector<int> neighbors_current;\n            set<vector<int>> set_ok;\n            set_ok.insert(vector<int>());\n            for (int depth_target = 0; depth_target <= MAX_DEPTH; depth_target++) {\n                solve2_dfs(depth_target, iterations_max, iterations_current, neighbors_current, set_ok, n_file);\n                int queries_size_aft = (int)queries.size();\n                int dqueries_size = queries_size_aft - queries_size_bef;\n                if (dqueries_size > QUERIES_MAX)\n                    break;\n            }\n        }\n        cerr << (int)queries.size() << endl;\n    }\n\n    void print(ostream& os = cerr) const {\n        os << (int)queries.size() << endl;\n        for (const auto& query : queries) {\n            query.print(os);\n        }\n    }\n\n    void print_to_file() const {\n        ofstream ofs(file_output);\n        print(ofs);\n    }\n\n    void sieve_queries() {\n        // remove queries with same tokens\n        vector<Query> queries_new;\n        vector<set<vector<int>>> st(N_files);\n        for (const auto& query : queries) {\n            if (query.tokens.empty()) {\n                continue;\n            }\n            if (query.neighbors.empty()) {\n                continue;\n            }\n            // if (query.neighbors.size() <= 1) {\n            //     // TODO for testing\n            //     continue;\n            // }\n            if (st[query.token_type].count(query.tokens) == 0) {\n                queries_new.push_back(query);\n                st[query.token_type].insert(query.tokens);\n            }\n        }\n        queries = queries_new;\n    }\n\n    void sieve_queries2() {\n        // consider tree structure\n        // map<vector<int>, vector<Query>> mp;\n        cerr << \"sieve_queries2 start!  queries.size() = \" << queries.size() << endl;\n        map<vector<int>, vector<Query>> mp;\n        map<vector<int>, vector<Query>> mp_cpc;\n        for (auto& query : queries) {\n            if (file_costs_one[query.token_type]) {\n                mp[query.neighbors].emplace_back(query);\n            } else {\n                mp_cpc[query.neighbors].emplace_back(query);\n            }\n        }\n\n        vector<Query> queries_new1;\n        for (auto& [neighbors, queries_mp] : mp) {\n            // sort by invalid and number of tokens\n            sort(queries_mp.begin(), queries_mp.end(), [&](const Query& a, const Query& b) {\n                if (a.score_invalid != b.score_invalid) {\n                    return a.score_invalid < b.score_invalid;\n                }\n                return a.tokens.size() < b.tokens.size();\n            });\n            int invalid_bef = -1;\n            int tokens_bef = 1000000;\n            for (auto& query : queries_mp) {\n                if (invalid_bef < query.score_invalid) {\n                    invalid_bef = query.score_invalid;\n                    if ((int)query.tokens.size() < tokens_bef) {\n                        tokens_bef = query.tokens.size();\n                        queries_new1.emplace_back(query);\n                    }\n                }\n            }\n        }\n\n        vector<vector<Query>> queries_new2_arr(N_files);\n        sort(queries_new1.begin(), queries_new1.end(),\n            [&](const Query& a, const Query& b) { return a.neighbors.size() > b.neighbors.size(); });\n        for (auto& query : queries_new1) {\n            bool flag = true;\n            for (auto& query2 : queries_new2_arr[query.token_type]) {\n                if (query.score_invalid >= query2.score_invalid\n                    && vector_intersection_fast(query.neighbors, query2.neighbors).size() == query.neighbors.size()) {\n                    flag = false;\n                    break;\n                }\n            }\n            if (flag) {\n                queries_new2_arr[query.token_type].emplace_back(query);\n            }\n        }\n\n        vector<Query> queries_new;\n        for (int i = 0; i < N_files; i++) {\n            copy(queries_new2_arr[i].begin(), queries_new2_arr[i].end(), back_inserter(queries_new));\n        }\n        for (auto& [neighbors, queries_mp] : mp_cpc) {\n            queries_new.insert(queries_new.end(), queries_mp.begin(), queries_mp.end());\n        }\n\n        cerr << \"sieve_queries2 end! queries_new.size() = \" << queries_new.size() << endl;\n        queries = queries_new;\n    }\n\n    void summary() {\n        vector<int> neighbor_count(Neighbor_size, 0);\n        for (const auto& query : queries) {\n            for (int neighbor : query.neighbors) {\n                neighbor_count[neighbor]++;\n            }\n        }\n        int neighbor_cover = 0;\n        for (int i = 0; i < Neighbor_size; i++) {\n            if (neighbor_count[i] > 0) {\n                neighbor_cover++;\n            }\n        }\n        print_time();\n        cerr << \"seed = \" << seed_solve << \" neighbor_cover = \" << neighbor_cover << endl;\n    }\n};\n\nvoid solve() {\n    for (int seed = seed_solve_begin; seed < N_neighbors; seed += seed_solve_step) {\n        Solver solver(seed);\n        if (data_neighbors[seed].size() != Neighbor_size) {\n            cerr << \"seed = \" << seed << \" has invalid number of neighbors of \" << data_neighbors[seed].size() << endl;\n        } else {\n            // solver.solve();\n            solver.solve2();\n            solver.sieve_queries();\n            solver.sieve_queries2();\n            solver.summary();\n        }\n        solver.print_to_file();\n    }\n}\n\nvoid solve_parallel_order() {\n    constexpr int N_THREAD = 4;\n    vector<thread> threads;\n    int seed_global = seed_solve_begin;\n    for (int i = 0; i < N_THREAD; i++) {\n        threads.emplace_back([&]() {\n            Solver solver(seed_solve_begin);\n            int seed_thread;\n            while (true) {\n                {\n                    std::lock_guard<std::mutex> lock(mtx);\n                    seed_thread = seed_global;\n                    seed_global += seed_solve_step;\n                }\n                if (seed_thread >= N_neighbors) {\n                    break;\n                }\n                solver.init(seed_thread);\n                if (data_neighbors[seed_thread].size() != Neighbor_size) {\n                    cerr << \"seed = \" << seed_thread << \" has invalid number of neighbors of \"\n                         << data_neighbors[seed_thread].size() << endl;\n                } else {\n                    solver.solve2();\n                    solver.sieve_queries();\n                    solver.sieve_queries2();\n                    solver.summary();\n                }\n                solver.print_to_file();\n            }\n        });\n    }\n    for (auto& th : threads) {\n        th.join();\n    }\n}\n\n\nvoid init() {}\n\nvoid read_input(int argc, char* argv[]) {\n#ifdef DEBUG\n    cerr << \"[start] read_input\" << endl;\n#endif\n    // read argv\n    if (argc >= 5) {\n        dir_input = argv[1];\n        dir_output = argv[2];\n        seed_solve_begin = atoi(argv[3]);\n        seed_solve_step = atoi(argv[4]);\n    } else {\n        cout << \"Usage: ./main.out <dir_input> <dir_output>\" << endl;\n        exit(1);\n    }\n#ifdef DEBUG\n    cerr << \"dir_input = \" << dir_input << endl;\n    cerr << \"dir_output = \" << dir_output << endl;\n#endif\n\n    // read all files in dir_input\n    vector<string> files_input;\n    for (const auto& p : fs::directory_iterator(dir_input)) {\n        cerr << p.path().string() << endl;\n        if (p.path().string().find(\".txt\") == string::npos) {\n            cerr << \"skip this file because not .txt file\" << endl;\n        }\n#ifndef USE_CPC\n        if (p.path().string().find(\"cpc.txt\") != string::npos) {\n            continue;\n        }\n#endif\n#ifndef USE_TI\n        if (p.path().string().find(\"ti.txt\") != string::npos) {\n            continue;\n        }\n#endif\n#ifndef USE_AB\n        if (p.path().string().find(\"ab.txt\") != string::npos) {\n            continue;\n        }\n#endif\n#ifndef USE_CLM\n        if (p.path().string().find(\"clm.txt\") != string::npos) {\n            continue;\n        }\n#endif\n#ifndef USE_DETD\n        if (p.path().string().find(\"detd.txt\") != string::npos) {\n            continue;\n        }\n#endif\n        cerr << \"use this file\" << endl;\n        files_input.push_back(p.path().string());\n    }\n    N_files = (int)files_input.size() - 1;\n    cerr << \"N_files = \" << N_files << endl;\n\n    // resize\n    num_tokens.resize(N_files);\n    data_len.resize(N_files);\n    data_tokens.resize(N_files);\n    data_tokens_inv.resize(N_files);\n    token_files.resize(N_files);\n    token_names.resize(N_files);\n\n    // all neighbors\n    set<int> set_data_neighbors;\n\n\n    // read neighbors.txt\n    for (const auto& file_name : files_input) {\n        if (file_name.find(\"neighbors.txt\") == string::npos) {\n            continue;\n        }\n        cerr << \"reading \" << file_name << endl;\n        ifstream ifs(file_name);\n        if (!ifs) {\n            cerr << \"Error: file not found: \" << file_name << endl;\n            exit(1);\n        }\n        ifs >> N_neighbors;\n        data_neighbors.resize(N_neighbors);\n        for (int j = 0; j < N_neighbors; j++) {\n            int neighbor_size;\n            ifs >> neighbor_size;\n            data_neighbors[j].resize(neighbor_size);\n            // assert(neighbor_size == Neighbor_size);\n            for (int k = 0; k < neighbor_size; k++) {\n                ifs >> data_neighbors[j][k];\n                // cerr << \"data_neighbors[\" << j << \"][\" << k << \"] = \" << data_neighbors[j][k] << endl;\n                set_data_neighbors.insert(data_neighbors[j][k]);\n            }\n        }\n    }\n\n\n    int n_file = 0;\n    for (int i = 0; i <= N_files; i++) {\n        const auto file_name = files_input[i];\n        // read file\n        ifstream ifs(file_name);\n        if (!ifs) {\n            cerr << \"Error: file not found: \" << file_name << endl;\n            exit(1);\n        }\n        cerr << \"reading \" << file_name << endl;\n        if (file_name.find(\"neighbors.txt\") != string::npos) {\n            continue;\n        } else {\n            // read input\n            int N_row, N_tokens;\n            ifs >> N_row >> N_tokens;\n            token_files[n_file] = file_name;\n            token_names[n_file] = fs::path(file_name).filename().stem().string();\n#ifdef DEBUG\n            cerr << \"token_files[\" << n_file << \"] = \" << token_files[n_file] << endl;\n            cerr << \"token_names[\" << n_file << \"] = \" << token_names[n_file] << endl;\n#endif\n\n            if (file_name.find(\"cpc.txt\") != string::npos) {\n                file_costs_one.emplace_back(0);\n            } else {\n                file_costs_one.emplace_back(1);\n            }\n\n\n            if (N_total == -1) {\n                N_total = N_row;\n            }\n            assert(N_total == N_row);\n            num_tokens[n_file] = N_tokens;\n            data_len[n_file].resize(N_row);\n            data_tokens[n_file].resize(N_row);\n            data_tokens_inv[n_file].resize(N_tokens);\n            for (int j = 0; j < N_row; j++) {\n                ifs >> data_len[n_file][j];\n                data_tokens[n_file][j].resize(data_len[n_file][j]);\n                for (int k = 0; k < data_len[n_file][j]; k++) {\n                    ifs >> data_tokens[n_file][j][k];\n                    // add data to inv\n                    data_tokens_inv[n_file][data_tokens[n_file][j][k]].emplace_back(j);\n                }\n                if (set_data_neighbors.count(j)) {\n                    sort(data_tokens[n_file][j].begin(), data_tokens[n_file][j].end());\n                } else {\n                    data_tokens[n_file][j].clear();\n                    data_tokens[n_file][j].shrink_to_fit();\n                }\n            }\n\n            for (int j = 0; j < N_tokens; j++) {\n                sort(data_tokens_inv[n_file][j].begin(), data_tokens_inv[n_file][j].end());\n            }\n\n            n_file++;\n        }\n    }\n#ifdef DEBUG\n    cerr << \"N_total = \" << N_total << endl;\n    cerr << \"N_files = \" << N_files << endl;\n    cerr << \"N_neighbors = \" << N_neighbors << endl;\n\n    cerr << \"[end] read_input\" << endl;\n#endif\n}\n\nint main(int argc, char* argv[]) {\n    read_input(argc, argv);\n    print_time();\n    init();\n\n    solve_parallel_order();\n\n    print_time();\n    return 0;\n}\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:05:41.292577Z","iopub.execute_input":"2024-07-24T02:05:41.292997Z","iopub.status.idle":"2024-07-24T02:05:41.367135Z","shell.execute_reply.started":"2024-07-24T02:05:41.292958Z","shell.execute_reply":"2024-07-24T02:05:41.365816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!g++ main.cpp -o main -std=c++17 -pthread -Wall -Wextra -pedantic -O3","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:05:41.368954Z","iopub.execute_input":"2024-07-24T02:05:41.369431Z","iopub.status.idle":"2024-07-24T02:05:48.58012Z","shell.execute_reply.started":"2024-07-24T02:05:41.36939Z","shell.execute_reply":"2024-07-24T02:05:48.578566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://github.com/hitonanode/linprog-modeler?tab=readme-ov-file\n\nfrom enum import Enum\nfrom typing import SupportsFloat\n\nfrom scipy.optimize import linprog  # type: ignore\n\n\nclass LPVarType(Enum):\n    CONTINUOUS = 0\n    INTEGER = 1\n    SEMICONTINUOUS = 2  # 0 or lb <= x <= ub\n    SEMIINTEGER = 3  # 0 or lb <= x <= ub, integer\n\n\nclass LPStatus(Enum):\n    OPTIMAL = 0\n    ITERATION_LIMIT = 1\n    INFEASIBLE = 2\n    UNBOUNDED = 3\n    NUMERICAL_ERROR = 4\n\n\nclass LPSense(Enum):\n    MINIMIZE = 0\n    MAXIMIZE = 1\n\n\nclass _LPBase:\n    def __add__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        return LPExpression.build(self) + other\n\n    def __radd__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        return other + LPExpression.build(self)\n\n    def __sub__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        return LPExpression.build(self) - other\n\n    def __rsub__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        return other - LPExpression.build(self)\n\n    def __mul__(self, other: SupportsFloat) -> \"LPExpression\":\n        return LPExpression.build(self) * other\n\n    def __rmul__(self, other: SupportsFloat) -> \"LPExpression\":\n        return LPExpression.build(self) * other\n\n    def __truediv__(self, other: SupportsFloat) -> \"LPExpression\":\n        return LPExpression.build(self) / other\n\n    def __neg__(self) -> \"LPExpression\":\n        return -LPExpression.build(self)\n\n    def __le__(\n        self,\n        other: \"LPExpressionLike\",\n    ) -> \"LPInequality\":\n        return LPInequality(\n            lhs=LPExpression.build(self),\n            rhs=other,\n            inequality_type=LPInequalityType.LESSEQ,\n        )\n\n    def __ge__(\n        self,\n        other: \"LPExpressionLike\",\n    ) -> \"LPInequality\":\n        return LPInequality(\n            lhs=other,\n            rhs=LPExpression.build(self),\n            inequality_type=LPInequalityType.LESSEQ,\n        )\n\n    def __eq__(  # type: ignore\n        self,\n        other: \"LPExpressionLike\",\n    ) -> \"LPInequality\":\n        return LPInequality(\n            lhs=LPExpression.build(self),\n            rhs=other,\n            inequality_type=LPInequalityType.EQUAL,\n        )\n\n\nclass LPVar(_LPBase):\n    def __init__(\n        self,\n        name: str,\n        lower_bound: float | None = None,\n        upper_bound: float | None = None,\n        variable_type: LPVarType = LPVarType.CONTINUOUS,\n    ) -> None:\n        self.name = name\n        self.lower_bound: float | None = lower_bound\n        self.upper_bound: float | None = upper_bound\n        self.variable_type: LPVarType = variable_type\n        self._value: float | None = None\n\n    def __str__(self) -> str:\n        s = \"{}(lb={}, ub={}, type={})\".format(\n            self.name,\n            self.lower_bound,\n            self.upper_bound,\n            self.variable_type.name,\n        )\n        if self._value is not None:\n            s += \": {}\".format(self._value)\n        return s\n\n    def __repr__(self) -> str:\n        return str(self)\n\n    def value(self) -> float | None:\n        return self._value\n\n\nclass LPInequalityType(Enum):\n    LESSEQ = 1  # lhs <= rhs\n    EQUAL = 2  # lhs == rhs\n\n\nclass _LPTerm(_LPBase):\n    \"\"\"\n    A term of the form coefficient * variable.\n    \"\"\"\n\n    def __init__(self, coefficient: float, variable: LPVar) -> None:\n        self.coefficient: float = coefficient\n        self.variable: LPVar = variable\n\n\nclass LPExpression:\n    def __init__(\n        self,\n        const: SupportsFloat,\n        terms: list[_LPTerm],\n    ) -> None:\n        self.const = float(const)\n        self.terms = terms.copy()\n        self._value: float | None = None\n\n    @classmethod\n    def build(self, x: \"LPExpressionLike | _LPBase\") -> \"LPExpression\":\n        if isinstance(x, LPExpression):\n            return LPExpression(x.const, x.terms)\n        elif isinstance(x, _LPTerm):\n            return LPExpression(0, [x])\n        elif isinstance(x, LPVar):\n            return LPExpression(0, [_LPTerm(1.0, x)])\n        elif isinstance(x, SupportsFloat):\n            return LPExpression(x, [])\n        else:\n            raise TypeError(\"Invalid type for LPExpression\")\n\n    def __add__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        rhs = LPExpression.build(other)\n        return LPExpression(\n            self.const + rhs.const,\n            self.terms + rhs.terms,\n        )\n\n    def __radd__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        return self + other\n\n    def __sub__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        rhs = LPExpression.build(other)\n        return self + (-rhs)\n\n    def __rsub__(self, other: \"LPExpressionLike\") -> \"LPExpression\":\n        return LPExpression.build(other) + (-self)\n\n    def __mul__(self, other: SupportsFloat) -> \"LPExpression\":\n        return LPExpression(\n            self.const * float(other),\n            [\n                _LPTerm(\n                    t.coefficient * float(other),\n                    t.variable,\n                )\n                for t in self.terms\n            ],\n        )\n\n    def __rmul__(self, other: SupportsFloat) -> \"LPExpression\":\n        return self * other\n\n    def __truediv__(self, other: SupportsFloat) -> \"LPExpression\":\n        return LPExpression(\n            self.const / float(other),\n            [\n                _LPTerm(\n                    t.coefficient / float(other),\n                    t.variable,\n                )\n                for t in self.terms\n            ],\n        )\n\n    def __neg__(self) -> \"LPExpression\":\n        return LPExpression(\n            -self.const,\n            [_LPTerm(-t.coefficient, t.variable) for t in self.terms],\n        )\n\n    def __le__(\n        self,\n        other: \"LPExpressionLike\",\n    ) -> \"LPInequality\":\n        return LPInequality(self, other, LPInequalityType.LESSEQ)\n\n    def __ge__(\n        self,\n        other: \"LPExpressionLike\",\n    ) -> \"LPInequality\":\n        return LPInequality(other, self, LPInequalityType.LESSEQ)\n\n    def __eq__(  # type: ignore\n        self,\n        other: \"LPExpressionLike\",\n    ) -> \"LPInequality\":\n        return LPInequality(self, other, LPInequalityType.EQUAL)\n\n    def __str__(self) -> str:\n        ret = [\"{}\".format(self.const)]\n\n        for t in self.terms:\n            sgn = \"+\" if t.coefficient >= 0 else \"\"\n            ret.append(\"{}{}{}\".format(sgn, t.coefficient, t.variable.name))\n\n        return \" \".join(ret)\n\n    def value(self) -> float | None:\n        return self._value\n\n\nLPExpressionLike = LPExpression | _LPTerm | LPVar | SupportsFloat\n\n\nclass LPInequality:\n    def __init__(\n        self,\n        lhs: LPExpressionLike,\n        rhs: LPExpressionLike,\n        inequality_type: LPInequalityType,\n    ) -> None:\n        \"\"\"\n        lhs <= rhs\n        -> terms + const (<= or ==) 0\n        \"\"\"\n\n        self.lhs = LPExpression.build(lhs) - LPExpression.build(rhs)\n        self.inequality_type = inequality_type\n\n    def __str__(self) -> str:\n        if self.inequality_type == LPInequalityType.LESSEQ:\n            return \"{} <= 0\".format(self.lhs)\n        elif self.inequality_type == LPInequalityType.EQUAL:\n            return \"{} == 0\".format(self.lhs)\n        else:\n            raise ValueError(\"Invalid inequality type\")\n\n    def __repr__(self) -> str:\n        return str(self)\n\n\nclass LPModel:\n    def __init__(self, sense: LPSense = LPSense.MINIMIZE) -> None:\n        self.sense = sense\n        self.has_impossible_constraints = False\n        self.constraints: list[LPInequality] = []\n        self.objective: LPExpression = LPExpression(0, [])\n        self.status: LPStatus | None = None\n\n    def add_constraint(self, constraint: LPInequality | bool) -> None:\n        if isinstance(constraint, bool):\n            if not constraint:\n                self.has_impossible_constraints = True\n        else:\n            self.constraints.append(constraint)\n\n    def set_objective(self, objective: LPExpressionLike) -> None:\n        self.objective = LPExpression.build(objective)\n\n    def solve(self) -> None:\n        var_dict: dict[int, LPVar] = {}\n        for constraint in self.constraints:\n            for term in constraint.lhs.terms:\n                var_dict.setdefault(id(term.variable), term.variable)\n\n        for term in self.objective.terms:\n            var_dict.setdefault(id(term.variable), term.variable)\n\n        # Reset status\n        self.objective._value = None\n        self.status = None\n        for var in var_dict.values():\n            var._value = None\n\n        # Obviously infeasible\n        if self.has_impossible_constraints:\n            self.status = LPStatus.INFEASIBLE\n            return\n\n        # Obviously optimal\n        if not var_dict:\n            self.status = LPStatus.OPTIMAL\n            self.objective._value = self.objective.const\n            return\n\n        id_to_idx = {id(v): i for i, v in enumerate(var_dict.values())}\n\n        A_ub: list[list[float]] = []\n        b_ub: list[float] = []\n        A_eq: list[list[float]] = []\n        b_eq: list[float] = []\n\n        for constraint in self.constraints:\n            lhs: list[float] = [0.0] * len(var_dict)\n            rhs = -constraint.lhs.const\n\n            for term in constraint.lhs.terms:\n                lhs[id_to_idx[id(term.variable)]] += term.coefficient\n\n            if constraint.inequality_type == LPInequalityType.LESSEQ:\n                A_ub.append(lhs)\n                b_ub.append(rhs)\n            elif constraint.inequality_type == LPInequalityType.EQUAL:\n                A_eq.append(lhs)\n                b_eq.append(rhs)\n            else:\n                raise ValueError(\"Invalid inequality type\")\n\n        bounds = [(v.lower_bound, v.upper_bound) for v in var_dict.values()]\n\n        integrality = [v.variable_type.value for v in var_dict.values()]\n\n        c: list[float] = [0.0] * len(var_dict)\n\n        func_weight = 1 if self.sense == LPSense.MINIMIZE else -1\n\n        for term in self.objective.terms:\n            c[id_to_idx[id(term.variable)]] += term.coefficient * func_weight\n\n        res = linprog(\n            c,\n            A_ub=A_ub or None,\n            b_ub=b_ub or None,\n            A_eq=A_eq or None,\n            b_eq=b_eq or None,\n            bounds=bounds,\n            integrality=integrality if sum(integrality) else None,\n            # options={\"time_limit\": 4}\n        )\n\n        if res.status == 0:\n            self.status = LPStatus.OPTIMAL\n\n            for i, variable in enumerate(var_dict.values()):\n                variable._value = res.x[i]\n\n            self.objective._value = res.fun * func_weight + self.objective.const\n\n        elif res.status == 1:\n            self.status = LPStatus.ITERATION_LIMIT\n\n        elif res.status == 2:\n            self.status = LPStatus.INFEASIBLE\n\n        elif res.status == 3:\n            self.status = LPStatus.UNBOUNDED\n\n        elif res.status == 4:\n            self.status = LPStatus.NUMERICAL_ERROR\n\n        else:\n            raise ValueError(\"Invalid status code\")\n","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:05:48.583515Z","iopub.execute_input":"2024-07-24T02:05:48.583885Z","iopub.status.idle":"2024-07-24T02:05:48.856975Z","shell.execute_reply.started":"2024-07-24T02:05:48.583851Z","shell.execute_reply":"2024-07-24T02:05:48.855511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%python\n\n\nfrom pathlib import Path\n\nimport polars as pl\nimport gc\n\npath_input_comp_meta = Path(\"/kaggle/input/uspto-explainable-ai/patent_metadata.parquet\")\n\ndf_meta_simple = pl.read_parquet(path_input_comp_meta)\nprint(df_meta_simple.shape)\nprint(df_meta_simple.head())\n\n# df_meta_simple から publication_number と行番号のマッピングを作成\nindex_mapping = pl.DataFrame({\n    \"publication_number\": df_meta_simple[\"publication_number\"],\n    # \"row_number\": pl.arange(0, len(df_meta_simple))\n    \"row_number\": pl.Series(\"row_number\", list(range(len(df_meta_simple))), pl.Int32)\n})\nindex_mapping = index_mapping.with_columns([\n    pl.col(\"publication_number\").cast(pl.Utf8).alias(\"publication_number\"),\n])\n\nmapping_dict = dict(zip(index_mapping[\"publication_number\"], index_mapping[\"row_number\"]))\nprint(len(mapping_dict))\nprint(list(mapping_dict.items())[:5])\n\ndel index_mapping\ngc.collect()\n\n\ndef map_to_row_number(column):\n    return column.replace(\n        mapping_dict,\n    )\n\npath_val_nn = \"/kaggle/input/uspto-explainable-ai/test.csv\"\ndf_val_nn = pl.read_csv(path_val_nn)\n# path_val_nn = \"/kaggle/input/uspto-explainable-ai-validation-index/neighbors_small.csv\"\n# df_val_nn = pl.read_csv(path_val_nn).head(100)\n\ncols_original = df_val_nn.columns\ndf_val_nn_id = df_val_nn.with_columns([\n    map_to_row_number(pl.col(column_name)).alias(f\"{column_name}_id\")\n    for column_name in df_val_nn.columns\n])\ndf_val_nn_id = df_val_nn_id.drop(cols_original)\n\n\npath_val_nn_id_txt = \"cppinput/neighbors.txt\"\nwith open(path_val_nn_id_txt, \"w\") as f:\n    N_row = df_val_nn_id.shape[0]\n    f.write(f\"{N_row}\\n\")\n    for row in df_val_nn_id.rows():\n        neighbors = list(row)[1:]\n        assert len(neighbors) == 50\n        f.write(f\"{len(neighbors)} {' '.join(map(str, neighbors))}\\n\")\n","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:06:39.392684Z","iopub.execute_input":"2024-07-24T02:06:39.39325Z","iopub.status.idle":"2024-07-24T02:13:10.50996Z","shell.execute_reply.started":"2024-07-24T02:06:39.393187Z","shell.execute_reply":"2024-07-24T02:13:10.508503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir cppoutput","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:13:10.513273Z","iopub.execute_input":"2024-07-24T02:13:10.513752Z","iopub.status.idle":"2024-07-24T02:13:11.680291Z","shell.execute_reply.started":"2024-07-24T02:13:10.51371Z","shell.execute_reply":"2024-07-24T02:13:11.678656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!./main ./cppinput ./cppoutput 0 1","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:13:11.681973Z","iopub.execute_input":"2024-07-24T02:13:11.682405Z","iopub.status.idle":"2024-07-24T02:34:05.447976Z","shell.execute_reply.started":"2024-07-24T02:13:11.682368Z","shell.execute_reply":"2024-07-24T02:34:05.445948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import whoosh_utils\nimport whoosh","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-24T02:34:05.453351Z","iopub.execute_input":"2024-07-24T02:34:05.453971Z","iopub.status.idle":"2024-07-24T02:34:41.487486Z","shell.execute_reply.started":"2024-07-24T02:34:05.453915Z","shell.execute_reply":"2024-07-24T02:34:41.486267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import polars as pl\ndf_test_nn = pl.read_csv(\"/kaggle/input/uspto-explainable-ai/test.csv\")\n\n# path_test_nn = \"/kaggle/input/uspto-explainable-ai-validation-index/neighbors_small.csv\"\n# df_test_nn = pl.read_csv(path_val_nn).head(100)","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:36:08.24473Z","iopub.execute_input":"2024-07-24T02:36:08.245232Z","iopub.status.idle":"2024-07-24T02:36:08.274891Z","shell.execute_reply.started":"2024-07-24T02:36:08.245157Z","shell.execute_reply":"2024-07-24T02:36:08.273635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport tqdm\nfrom pathlib import Path\nfrom multiprocessing import Pool, Manager\nfrom functools import partial\n\ndef process_file(file_idx, df_test_nn, shared_results):\n    path_cpp_output_dir = Path(\"./cppoutput\")\n    path_cpp_output_file = path_cpp_output_dir / f\"{file_idx}.txt\"\n    \n    # read the file\n    data_raw = []\n    with open(path_cpp_output_file, \"r\") as f:\n        lines = f.readlines()\n        n_lines = int(lines[0].strip())\n        assert len(lines) == n_lines + 1\n        for l in lines[1:]:\n            lis = l.strip().split(\" \")\n            col_ab = lis[0]\n            n_tokens = int(lis[1])\n            idx = 2\n            tokens = list(map(int, lis[idx:idx+n_tokens]))\n            n_neighbors = int(lis[idx+n_tokens])\n            idx += n_tokens + 1\n            neighbors = list(map(int, lis[idx:idx+n_neighbors]))\n            n_invalid = int(lis[idx+n_neighbors])\n            data_raw.append((col_ab, tokens, neighbors, n_invalid))\n\n    # create the LP model\n    N = len(data_raw)\n    M = 50\n\n    max_ORs = 25;\n    max_cost_true = 350\n    max_n_invalid = 1000\n\n    costs = []\n    costs_true = []\n    n_invalids = []\n    related = [[] for _ in range(M)]\n    for i, (col_ab, tokens, neighbors, n_invalid) in enumerate(data_raw):\n        if col_ab == 'cpc':\n            costs.append(len(tokens))\n        else:\n            costs.append(1) # 後で変えるかも\n            # costs.append(len(tokens))\n        costs_true.append(len(tokens))\n        for n in neighbors:\n            related[n].append(i)\n        n_invalids.append(n_invalid)\n\n    xs = [LPVar(f\"x{i}\", 0, 1, LPVarType.INTEGER) for i in range(N)] # どの制約を使うか\n    ys = [LPVar(f\"y{i}\", 0, 1, LPVarType.INTEGER) for i in range(M)] # どのpatentを包含できているか\n\n    problem = LPModel(LPSense.MAXIMIZE)\n\n    for i in range(M):\n        problem.add_constraint(sum([xs[j] for j in related[i]]) >= ys[i]) # 1つでも関連があれば1になる\n\n    problem.add_constraint(sum([xs[i]*costs[i] for i in range(N)]) <= max_ORs) # ORを取る対象が25以下\n    problem.add_constraint(sum([xs[i]*costs_true[i] for i in range(N)]) <= max_cost_true) # コストが一定数以下\n    \n    # invalid\n    problem.add_constraint(sum([xs[i]*n_invalids[i] for i in range(N)]) <= max_n_invalid) # invalidが一定数以下\n    ratio = 198104 / 13307751\n    problem.add_constraint(sum([xs[i]*n_invalids[i] for i in range(N)])/50 + sum([ys[i] for i in range(M)]) <= M + 1) # invalidが一定数以下\n\n\n    problem.set_objective(sum([ys[i] for i in range(M)])*10000\n                          - sum([xs[i] * n_invalids[i] for i in range(N)])\n    )\n\n\n    problem.solve()\n\n    \"\"\"\n    # クエリの作成\n    query_text = \"\"\n    for x, (col_ab, tokens, neighbors, n_invalid) in zip(xs, data_raw):\n        if round(x.value()) == 1:\n            if query_text != \"\":\n                query_text += \" OR \"\n            \n            if col_ab == 'cpc':\n                tmp = [\"cpc:\" + t for t in tokens]\n                query_text += f\"({' '.join(tmp)})\"\n            else:\n                query_text += f\"({col_ab}:{'-'.join(tokens)})\"\n    \n\n    if query_text == \"\":\n        query_text = \"hoge\"\n\n    \"\"\"\n        \n    # 結果の保存\n    \n    shared_results['list_queries'].append((file_idx, data_raw))\n    shared_results['list_xs'].append((file_idx, [int(round(x.value())) for x in xs]))\n\n\ndef main():\n    # 共有データの初期化\n    manager = Manager()\n    shared_results = manager.dict({\n        'list_queries': manager.list(),\n        'list_xs': manager.list(),\n    })\n\n    # パラメータの設定\n    \n    # 並列処理の実行\n    with Pool() as pool:\n        process_func = partial(process_file, df_test_nn=df_test_nn, shared_results=shared_results)\n        list(tqdm.tqdm(pool.imap(process_func, range(len(df_test_nn))), total=len(df_test_nn)))\n\n    # 結果の取得\n    \"\"\"\n    list_queries_raw = list(shared_results['list_queries'])\n    list_xs =  list(shared_results['list_xs'])\n    \n    list_queries_raw.sort(key=lambda x:x[0])\n    list_xs.sort(key=lambda x:x[0])\n    print([x[0] for x in list_queries_raw[:10]])\n    print([x[0] for x in list_queries_raw[:10]])\n\n    list_queries_raw = [x[1] for x in list_queries_raw]\n    list_xs = [x[1] for x in list_xs]\n    \"\"\"\n    \n    list_queries_raw_tmp = list(shared_results['list_queries'])\n    list_xs_tmp =  list(shared_results['list_xs'])\n    \n    print([x[0] for x in list_queries_raw_tmp[:10]])\n    print([x[0] for x in list_xs_tmp[:10]])\n    \n    N = len(df_test_nn)\n    list_queries_raw = [None] * N\n    list_xs = [None] * N\n    for i, x in list_queries_raw_tmp:\n        list_queries_raw[i] = x\n    for i, x in list_xs_tmp:\n        list_xs[i] = x\n    \n    # 結果の表示や保存など、必要な処理を行う\n    print(f\"Total results: {len(list_xs)}, {len(list_queries_raw)}\")\n    \n    return list_queries_raw, list_xs\n\nlist_queries_raw, list_xs = main()","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:36:09.992532Z","iopub.execute_input":"2024-07-24T02:36:09.992942Z","iopub.status.idle":"2024-07-24T02:36:44.738828Z","shell.execute_reply.started":"2024-07-24T02:36:09.992912Z","shell.execute_reply":"2024-07-24T02:36:44.737304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\nimport os\nimport pickle\n\ntokens_sorted = {}\n\ncols = [\n    (\"title_tokens\", \"ti\"),\n    (\"abstract_tokens\", \"ab\"),\n    (\"claims_tokens\", \"clm\"),\n    (\"description_tokens\", \"detd\"),\n    (\"cpc_codes\", \"cpc\"),\n]\n\npath_tokens_dir = Path(\"/kaggle/input/list-all\")\n\nfor col, col_ab in cols:\n    pth_list = path_tokens_dir / f\"{col}_list_all.pkl\"\n    with open(pth_list, \"rb\") as f:\n        list_tokens_flatten = pickle.load(f)\n    tokens_sorted[col_ab] = list_tokens_flatten","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:36:44.741086Z","iopub.execute_input":"2024-07-24T02:36:44.741535Z","iopub.status.idle":"2024-07-24T02:37:27.08852Z","shell.execute_reply.started":"2024-07-24T02:36:44.741497Z","shell.execute_reply":"2024-07-24T02:37:27.087221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_queries = []\nfor xs, data_raw in zip(list_xs, list_queries_raw):\n    if (xs is None) or (data_raw is None):\n        print(\"found none\")\n        query_text = \"\"\n    else:\n        query_text = \"\"\n        for x, (col_ab, tokens, neighbors, n_invalid) in zip(xs, data_raw):\n            tokens = [tokens_sorted[col_ab][i] for i in tokens]\n            if round(x) == 1:\n                if query_text != \"\":\n                    query_text += \" OR \"\n\n                if col_ab == 'cpc':\n                    tmp = [\"cpc:\" + t for t in tokens]\n                    query_text += f\"({' '.join(tmp)})\"\n                else:\n                    query_text += f\"({col_ab}:{'-'.join(tokens)})\"\n\n\n    if query_text == \"\":\n        query_text = \"ti:hoge\"\n    \n    list_queries.append(query_text)","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:37:27.090331Z","iopub.execute_input":"2024-07-24T02:37:27.090809Z","iopub.status.idle":"2024-07-24T02:37:27.644176Z","shell.execute_reply.started":"2024-07-24T02:37:27.090766Z","shell.execute_reply":"2024-07-24T02:37:27.642908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(list_queries)\nprint(list_queries[:2])","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:37:27.646989Z","iopub.execute_input":"2024-07-24T02:37:27.64777Z","iopub.status.idle":"2024-07-24T02:37:27.654634Z","shell.execute_reply.started":"2024-07-24T02:37:27.647724Z","shell.execute_reply":"2024-07-24T02:37:27.653349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tqdm\n\nresults = []\n\nfor i in tqdm.trange(len(df_test_nn)):\n    results.append(\n        {\"publication_number\": df_test_nn[i, \"publication_number\"], \"query\": list_queries[i]}\n    )","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:37:27.656433Z","iopub.execute_input":"2024-07-24T02:37:27.656872Z","iopub.status.idle":"2024-07-24T02:37:27.67472Z","shell.execute_reply.started":"2024-07-24T02:37:27.656835Z","shell.execute_reply":"2024-07-24T02:37:27.673249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ndef test_results():\n    test_idx = whoosh_utils.load_index(\"/kaggle/input/uspto-explainable-ai-validation-index/validation/validation_index\")\n    test_searcher = whoosh_utils.get_searcher(test_idx)\n    test_qp = whoosh_utils.get_query_parser()\n    \n    assert len(list_queries) == len(df_test_nn)\n    \n    for i, row in enumerate(df_test_nn.rows()):\n        neighbors = list(row)[1:]\n        if i >= len(list_queries):\n            break\n        query_text = list_queries[i]\n        if query_text is None:\n            print(f\"skip row {i} due to query_text=None\")\n            continue\n        # print(query_text)\n        res_test = whoosh_utils.execute_query(query_text, test_qp, test_searcher, 100)\n        count = 0\n        for found in res_test:\n            if found in neighbors:\n                count += 1\n        print(f\"{i}\\t{count}\")\n# gc.collect()\n# test_results()\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-07-24T02:42:29.958796Z","iopub.execute_input":"2024-07-24T02:42:29.959293Z","iopub.status.idle":"2024-07-24T02:43:11.580653Z","shell.execute_reply.started":"2024-07-24T02:42:29.959255Z","shell.execute_reply":"2024-07-24T02:43:11.579422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf /kaggle/working/*","metadata":{"execution":{"iopub.status.busy":"2024-07-23T19:53:57.656602Z","iopub.execute_input":"2024-07-23T19:53:57.656998Z","iopub.status.idle":"2024-07-23T19:53:58.871099Z","shell.execute_reply.started":"2024-07-23T19:53:57.656967Z","shell.execute_reply":"2024-07-23T19:53:58.869661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pl.DataFrame(results)\nsubmission.write_csv(\"submission.csv\")\n\nsubmission","metadata":{"execution":{"iopub.status.busy":"2024-07-23T19:53:59.420654Z","iopub.execute_input":"2024-07-23T19:53:59.421095Z","iopub.status.idle":"2024-07-23T19:53:59.440639Z","shell.execute_reply.started":"2024-07-23T19:53:59.421058Z","shell.execute_reply":"2024-07-23T19:53:59.439332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}