Subsets with Distinct Prime Factors

Last Updated : 31 Jul, 2026

Given an array arr[] of size n. The task is to find a number of subsets whose product can be represented as a product of one or more distinct prime numbers. Modify your answer to the modulo of 109 + 7.

Constraints:

  • 1 <= arr.size() <= 105
  • 1< = arr[i] <= 30

Example:

Input: arr[] = [1, 2, 3, 4]
Output: 6
Explanation: The subsets are:
[2], product = 2 = 2
[3], product = 3 = 3
[1, 2], product = 2 = 2
[1, 3], product = 3 = 3
[2, 3], product = 6 = 2 × 3
[1, 2, 3], product = 6 = 2 × 3
All these products can be expressed as a product of one or more distinct prime numbers. Hence, the count is 6.
Note that [4] or any other subset with 4 are not chosen because products having 4 have repeated primes 2.

Input: arr[] = [2, 2, 3]
Output: 5
Explanation: Since subsets formed using different indices are considered different, the chosen subsets are:
[2] (using the first 2)
[2] (using the second 2)
[3]
[2, 3] (using the first 2)
[2, 3] (using the second 2)
Each subset has a product that can be expressed as a product of one or more distinct prime numbers. Therefore, the answer is 5.

Try It Yourself
redirect icon

[Backtracking Approach] Using Recursion – O(2ⁿ × √m) Time and O(n) Space

The idea is to generate all possible subsets using recursion and backtracking. A subset is considered only if no prime factor repeats across numbers and every selected number is individually square-free. For each element, we either include it in the subset or skip it, while tracking already used prime factors.

  • Recursively try both choices for every element: take or skip
  • Before taking a number, check whether it is square-free
  • Extract distinct prime factors of the current number
  • If any prime factor is already used, skip that number
  • Otherwise, mark primes as used, continue recursion, and backtrack later
  • Count only non-empty valid subsets
C++
#include <bits/stdc++.h>
using namespace std;

// Check whether number is square free or not
// Example:
// 6  -> valid
// 12 -> invalid because 2 repeats
bool validNumber(int num)
{
    for (int i = 2; i * i <= num; i++) {

        int cnt = 0;

        while (num % i == 0) {
            cnt++;
            num /= i;

            // repeated prime factor found
            if (cnt > 1)
                return false;
        }
    }

    return true;
}

// Store distinct prime factors
vector<int> getPrimeFactors(int num)
{
    vector<int> primes;

    for (int i = 2; i * i <= num; i++) {

        if (num % i == 0) {

            primes.push_back(i);

            while (num % i == 0)
                num /= i;
        }
    }

    // remaining prime number
    if (num > 1)
        primes.push_back(num);

    return primes;
}

int solve(int idx, vector<int>& nums,
          map<int, int>& usedPrime,
          bool hasPrime)
{
    // all elements processed
    if (idx == nums.size()) {

        // count only subsets having
        // at least one prime factor
        return hasPrime;
    }

    // option 1 -> skip current element
    int notTake =
        solve(idx + 1,
              nums,
              usedPrime,
              hasPrime);

    int take = 0;

    int num = nums[idx];

    // invalid number can never be used
    if (!validNumber(num))
        return notTake;

    vector<int> primes = getPrimeFactors(num);

    bool possible = true;

    // check whether any prime already used
    for (int p : primes) {

        if (usedPrime[p]) {
            possible = false;
            break;
        }
    }

    // option 2 -> take current element
    if (possible) {

        // mark primes as used
        for (int p : primes)
            usedPrime[p]++;

        take = solve(idx + 1,
                     nums,
                     usedPrime,
                     // subset becomes valid only if
                     // current number contributes
                     // at least one prime factor
                     hasPrime || !primes.empty());

        // backtracking step
        for (int p : primes)
            usedPrime[p]--;
    }

    return take + notTake;
}

int countSubsets(vector<int>& nums)
{
    map<int, int> usedPrime;

    return solve(0, nums, usedPrime, false);
}

int main()
{
    vector<int> arr = {1, 2, 3, 4};

    cout << countSubsets(arr);

    return 0;
}
Java
import java.util.*;

class GFG {

    // Check whether number is square free or not
    // Example:
    // 6  -> valid
    // 12 -> invalid because 2 repeats
    static boolean validNumber(int num) {
        for (int i = 2; i * i <= num; i++) {

            int cnt = 0;

            while (num % i == 0) {
                cnt++;
                num /= i;

                // repeated prime factor found
                if (cnt > 1)
                    return false;
            }
        }

        return true;
    }

    // Store distinct prime factors
    static List<Integer> getPrimeFactors(int num) {
        List<Integer> primes = new ArrayList<>();

        for (int i = 2; i * i <= num; i++) {

            if (num % i == 0) {

                primes.add(i);

                while (num % i == 0)
                    num /= i;
            }
        }

        // remaining prime number
        if (num > 1)
            primes.add(num);

        return primes;
    }

    static int solve(int idx, int[] nums,
                     Map<Integer, Integer> usedPrime,
                     boolean hasPrime) {

        // all elements processed
        if (idx == nums.length) {

            // count only subsets having
            // at least one prime factor
            return hasPrime ? 1 : 0;
        }

        // option 1 -> skip current element
        int notTake = solve(idx + 1,
                            nums,
                            usedPrime,
                            hasPrime);

        int take = 0;

        int num = nums[idx];

        // invalid number can never be used
        if (!validNumber(num))
            return notTake;

        List<Integer> primes = getPrimeFactors(num);

        boolean possible = true;

        // check whether any prime already used
        for (int p : primes) {
            if (usedPrime.getOrDefault(p, 0) > 0) {
                possible = false;
                break;
            }
        }

        // option 2 -> take current element
        if (possible) {

            // mark primes as used
            for (int p : primes)
                usedPrime.put(p, usedPrime.getOrDefault(p, 0) + 1);

            take = solve(idx + 1,
                         nums,
                         usedPrime,
                         // subset becomes valid only if
                         // current number contributes
                         // at least one prime factor
                         hasPrime || !primes.isEmpty());

            // backtracking step
            for (int p : primes) {
                int cnt = usedPrime.get(p) - 1;
                if (cnt == 0)
                    usedPrime.remove(p);
                else
                    usedPrime.put(p, cnt);
            }
        }

        return take + notTake;
    }

    static int countSubsets(int[] nums) {
        Map<Integer, Integer> usedPrime = new HashMap<>();
        return solve(0, nums, usedPrime, false);
    }

    public static void main(String[] args) {
        int[] arr = {1, 2, 3, 4};

        System.out.println(countSubsets(arr));
    }
}
Python
def isValidNumber(num):
    i = 2
    while i * i <= num:
        cnt = 0
        while num % i == 0:
            cnt += 1
            num //= i
            if cnt > 1:
                return False
        i += 1
    return True

def getPrimeFactors(num):
    primes = []
    i = 2
    while i * i <= num:
        if num % i == 0:
            primes.append(i)
            while num % i == 0:
                num //= i
        i += 1
    if num > 1:
        primes.append(num)
    return primes

def solve(idx, nums, usedPrime, hasPrime):
    if idx == len(nums):
        return 1 if hasPrime else 0
    notTake = solve(idx + 1, nums, usedPrime, hasPrime)
    take = 0
    num = nums[idx]
    if not isValidNumber(num):
        return notTake
    primes = getPrimeFactors(num)
    possible = True
    for p in primes:
        if usedPrime.get(p, 0):
            possible = False
            break
    if possible:
        for p in primes:
            usedPrime[p] = usedPrime.get(p, 0) + 1
        take = solve(idx + 1, nums, usedPrime, hasPrime or bool(primes))
        for p in primes:
            usedPrime[p] -= 1
            if usedPrime[p] == 0:
                del usedPrime[p]
    return take + notTake

def countSubsets(nums):
    usedPrime = {}
    return solve(0, nums, usedPrime, False)

if __name__ == "__main__":
    arr = [1, 2, 3, 4]
    print(countSubsets(arr))
C#
using System;
using System.Collections.Generic;

class GFG
{
    // Check whether number is square free or not
    // Example:
    // 6  -> valid
    // 12 -> invalid because 2 repeats
    static bool validNumber(int num)
    {
        for (int i = 2; i * i <= num; i++)
        {
            int cnt = 0;

            while (num % i == 0)
            {
                cnt++;
                num /= i;

                // repeated prime factor found
                if (cnt > 1)
                    return false;
            }
        }

        return true;
    }

    // Store distinct prime factors
    static List<int> getPrimeFactors(int num)
    {
        List<int> primes = new List<int>();

        for (int i = 2; i * i <= num; i++)
        {
            if (num % i == 0)
            {
                primes.Add(i);

                while (num % i == 0)
                    num /= i;
            }
        }

        // remaining prime number
        if (num > 1)
            primes.Add(num);

        return primes;
    }

    static int solve(int idx, int[] nums,
                     Dictionary<int, int> usedPrime,
                     bool hasPrime)
    {
        // all elements processed
        if (idx == nums.Length)
        {
            // count only subsets having
            // at least one prime factor
            return hasPrime ? 1 : 0;
        }

        // option 1 -> skip current element
        int notTake = solve(idx + 1,
                            nums,
                            usedPrime,
                            hasPrime);

        int take = 0;

        int num = nums[idx];

        // invalid number can never be used
        if (!validNumber(num))
            return notTake;

        List<int> primes = getPrimeFactors(num);

        bool possible = true;

        // check whether any prime already used
        foreach (int p in primes)
        {
            if (usedPrime.ContainsKey(p))
            {
                possible = false;
                break;
            }
        }

        // option 2 -> take current element
        if (possible)
        {
            // mark primes as used
            foreach (int p in primes)
                usedPrime[p] = 1;

            take = solve(idx + 1,
                         nums,
                         usedPrime,
                         hasPrime || primes.Count > 0);

            // backtracking step
            foreach (int p in primes)
                usedPrime.Remove(p);
        }

        return take + notTake;
    }

    static int countSubsets(int[] nums)
    {
        Dictionary<int, int> usedPrime = new Dictionary<int, int>();

        return solve(0, nums, usedPrime, false);
    }

    static void Main()
    {
        int[] arr = { 1, 2, 3, 4 };

        Console.WriteLine(countSubsets(arr));
    }
}
JavaScript
// Check whether number is square free or not
// Example:
// 6  -> valid
// 12 -> invalid because 2 repeats
function validNumber(num) {
    for (let i = 2; i * i <= num; i++) {
        let cnt = 0;

        while (num % i === 0) {
            cnt++;
            num = Math.floor(num / i);

            // repeated prime factor found
            if (cnt > 1)
                return false;
        }
    }

    return true;
}

// Store distinct prime factors
function getPrimeFactors(num) {
    const primes = [];

    for (let i = 2; i * i <= num; i++) {
        if (num % i === 0) {
            primes.push(i);

            while (num % i === 0)
                num = Math.floor(num / i);
        }
    }

    // remaining prime number
    if (num > 1)
        primes.push(num);

    return primes;
}

function solve(idx, nums, usedPrime, hasPrime) {
    // all elements processed
    if (idx === nums.length) {
        // count only subsets having
        // at least one prime factor
        return hasPrime ? 1 : 0;
    }

    // option 1 -> skip current element
    const notTake = solve(idx + 1,
                          nums,
                          usedPrime,
                          hasPrime);

    let take = 0;

    const num = nums[idx];

    // invalid number can never be used
    if (!validNumber(num))
        return notTake;

    const primes = getPrimeFactors(num);

    let possible = true;

    // check whether any prime already used
    for (const p of primes) {
        if ((usedPrime.get(p) || 0) > 0) {
            possible = false;
            break;
        }
    }

    // option 2 -> take current element
    if (possible) {
        // mark primes as used
        for (const p of primes)
            usedPrime.set(p, (usedPrime.get(p) || 0) + 1);

        take = solve(idx + 1,
                     nums,
                     usedPrime,
                     hasPrime || primes.length > 0);

        // backtracking step
        for (const p of primes) {
            usedPrime.set(p, usedPrime.get(p) - 1);

            if (usedPrime.get(p) === 0)
                usedPrime.delete(p);
        }
    }

    return take + notTake;
}

function countSubsets(nums) {
    const usedPrime = new Map();
    return solve(0, nums, usedPrime, false);
}

// Driver code
const arr = [1, 2, 3, 4];
console.log(countSubsets(arr));

Output
6

[Expected Approach] Using Bitmask DP – O(30 × 2¹⁰) Time and O(2¹⁰) Space

As per the question constraints, the array elements are in range from 1 to 30, so the idea is to represent the prime factors of every number using a bitmask. Each bit corresponds to a distinct prime number. Since numbers are limited from 1 to 30, only the first 10 primes are needed. A number is valid only if it is square-free. Dynamic Programming is then used to count all possible good subsets while ensuring no prime factor overlaps.

  • Count frequency of every number from 1 to 30.
  • Generate a prime-factor bitmask for each number.
  • Ignore numbers having repeated prime factors.
  • Use dp[mask] where mask represents used prime factors.
  • Traverse masks in reverse to avoid overwriting current states.
  • Add current number only if its mask does not overlap with the existing mask.
  • Sum all valid states and subtract the empty subset.

Let us understand with an example:
consider an array arr[] = [1, 2, 3, 4]

  • Frequency array: freq[1]=1, freq[2]=1, freq[3]=1, freq[4]=1, all others 0; ones = 1.
  • Initialize DP with dp[0] = 1 representing the empty subset.
  • Process 2; getMask(2) = 0000000001; update dp[1] = 1.
  • Process 3; getMask(3) = 0000000010; update dp[2] = 1 and dp[3] = 1 using the existing state dp[1].
  • Process 4; getMask(4) = -1 because prime factor 2 repeats, so ignore this number.
  • Non-zero DP states are dp[1]=1 for {2}, dp[2]=1 for {3}, and dp[3]=1 for {2,3}.
  • Add all non-empty DP states to get ans = 1 + 1 + 1 = 3.
  • Multiply by 2^1 = 2 to account for including or excluding 1; ans = 3 × 2 = 6.
  • Valid good subsets are {2}, {3}, {2,3}, {1,2}, {1,3}, and {1,2,3}.

Final answer is 6.

C++
#include <bits/stdc++.h>
using namespace std;

const int MOD = 1e9 + 7;

// Fast exponentiation
long long power(long long a, long long b)
{
    long long ans = 1;

    while (b) {

        if (b & 1)
            ans = (ans * a) % MOD;

        a = (a * a) % MOD;
        b >>= 1;
    }

    return ans;
}

// Create prime mask for current number
int getMask(int num)
{
    vector<int> primes = {2, 3, 5, 7, 11, 13, 17, 19, 23, 29};

    int mask = 0;

    for (int i = 0; i < 10; i++) {

        int p = primes[i];
        int cnt = 0;

        while (num % p == 0) {
            cnt++;
            num /= p;
        }

        // repeated prime factor
        // invalid number
        if (cnt > 1)
            return -1;

        // store current prime in mask
        if (cnt == 1)
            mask |= (1 << i);
    }

    return mask;
}

int countSubsets(vector<int>& nums)
{
    // frequency of every number
    vector<int> freq(31, 0);

    for (int x : nums)
        freq[x]++;

    // number of ones
    int ones = freq[1];

    // dp[mask]
    // mask represents used prime factors
    vector<long long> dp(1024, 0);

    // empty subset
    dp[0] = 1;

    // process numbers from 2 to 30
    for (int num = 2; num <= 30; num++) {

        // number not present
        if (freq[num] == 0)
            continue;

        int currMask = getMask(num);

        // invalid number
        if (currMask == -1)
            continue;

        // reverse traversal
        // prevents overwriting current states
        for (int mask = 1023; mask >= 0; mask--) {

            // overlapping prime factor
            // cannot take together
            if ((mask & currMask) != 0)
                continue;

            dp[mask | currMask] =
                (dp[mask | currMask] +
                 dp[mask] * freq[num]) % MOD;
        }
    }

    long long ans = 0;

    // add all good subsets
    // skip empty subset (mask = 0)
    for (int mask = 1; mask < 1024; mask++)
        ans = (ans + dp[mask]) % MOD;

    // every valid subset can be combined
    // with any subset of ones
    ans = (ans * power(2, ones)) % MOD;

    return ans;
}

int main()
{
    vector<int> arr = {1, 2, 3, 4};

    cout << countSubsets(arr);

    return 0;
}
Java
import java.util.Arrays;

public class GFG {
    static final int MOD = 1000000007;
    
    // Fast exponentiation
    static long power(long a, long b) {
        long ans = 1;
        while (b > 0) {
            if ((b & 1)!= 0)
                ans = (ans * a) % MOD;
            a = (a * a) % MOD;
            b >>= 1;
        }
        return ans;
    }
    
    // Create prime mask for current number
    static int getMask(int num) {
        int[] primes = {2, 3, 5, 7, 11, 13, 17, 19, 23, 29};
        int mask = 0;
        for (int i = 0; i < 10; i++) {
            int p = primes[i];
            int cnt = 0;
            while (num % p == 0) {
                cnt++;
                num /= p;
            }
            // repeated prime factor
            // invalid number
            if (cnt > 1)
                return -1;
            // store current prime in mask
            if (cnt == 1)
                mask |= (1 << i);
        }
        return mask;
    }
    
    static int countSubsets(int[] nums) {
        
        // frequency of every number
        int[] freq = new int[31];
        Arrays.fill(freq, 0);
        for (int x : nums)
            freq[x]++;
            
        // number of ones
        int ones = freq[1];
        
        // dp[mask]
        // mask represents used prime factors
        long[] dp = new long[1024];
        Arrays.fill(dp, 0);
        
        // empty subset
        dp[0] = 1;
        
        // process numbers from 2 to 30
        for (int num = 2; num <= 30; num++) {
            
            // number not present
            if (freq[num] == 0)
                continue;
            int currMask = getMask(num);
            
            // invalid number
            if (currMask == -1)
                continue;
                
            // reverse traversal
            // prevents overwriting current states
            for (int mask = 1023; mask >= 0; mask--) {
                
                // overlapping prime factor
                // cannot take together
                if ((mask & currMask)!= 0)
                    continue;
                dp[mask | currMask] = (dp[mask | currMask] + dp[mask] * freq[num]) % MOD;
            }
        }
        long ans = 0;
        
        // add all good subsets
        // skip empty subset (mask = 0)
        for (int mask = 1; mask < 1024; mask++)
            ans = (ans + dp[mask]) % MOD;
            
        // every valid subset can be combined
        // with any subset of ones
        ans = (ans * power(2, ones)) % MOD;
        return (int) ans;
    }
    
    public static void main(String[] args) {
        int[] arr = {1, 2, 3, 4};
        System.out.println(countSubsets(arr));
    }
}
Python
MOD = 10**9 + 7

# Fast exponentiation
def power(a, b):
    ans = 1

    while b:

        if b & 1:
            ans = (ans * a) % MOD

        a = (a * a) % MOD
        b >>= 1

    return ans

# Create prime mask for current number
def getMask(num):
    primes = [2, 3, 5, 7, 11, 13, 17, 19, 23, 29]

    mask = 0

    for i in range(10):

        p = primes[i]
        cnt = 0

        while num % p == 0:
            cnt += 1
            num //= p

        # repeated prime factor
        # invalid number
        if cnt > 1:
            return -1

        # store current prime in mask
        if cnt == 1:
            mask |= (1 << i)

    return mask

def countSubsets(nums):
    
    # frequency of every number
    freq = [0] * 31

    for x in nums:
        freq[x] += 1

    # number of ones
    ones = freq[1]

    # dp[mask]
    # mask represents used prime factors
    dp = [0] * 1024

    # empty subset
    dp[0] = 1

    # process numbers from 2 to 30
    for num in range(2, 31):

        # number not present
        if freq[num] == 0:
            continue

        currMask = getMask(num)

        # invalid number
        if currMask == -1:
            continue

        # reverse traversal
        # prevents overwriting current states
        for mask in range(1023, -1, -1):

            # overlapping prime factor
            # cannot take together
            if (mask & currMask)!= 0:
                continue

            dp[mask | currMask] = (
                dp[mask | currMask] +
                dp[mask] * freq[num]
            ) % MOD

    ans = 0

    # add all good subsets
    # skip empty subset (mask = 0)
    for mask in range(1, 1024):
        ans = (ans + dp[mask]) % MOD

    # every valid subset can be combined
    # with any subset of ones
    ans = (ans * power(2, ones)) % MOD

    return ans

# Driver code
print(countSubsets([1, 2, 3, 4]))
C#
using System;
using System.Collections.Generic;

class GFG
{
    const int MOD = 1000000007;

    // Fast exponentiation
    static long power(long a, long b)
    {
        long ans = 1;

        while (b > 0)
        {
            if ((b & 1) == 1)
                ans = (ans * a) % MOD;

            a = (a * a) % MOD;
            b >>= 1;
        }

        return ans;
    }

    // Create prime mask for current number
    static int getMask(int num)
    {
        int[] primes = { 2, 3, 5, 7, 11, 13, 17, 19, 23, 29 };

        int mask = 0;

        for (int i = 0; i < 10; i++)
        {
            int p = primes[i];
            int cnt = 0;

            while (num % p == 0)
            {
                cnt++;
                num /= p;
            }

            // repeated prime factor
            // invalid number
            if (cnt > 1)
                return -1;

            // store current prime in mask
            if (cnt == 1)
                mask |= (1 << i);
        }

        return mask;
    }

    static int countSubsets(int[] nums)
    {
        // frequency of every number
        int[] freq = new int[31];

        foreach (int x in nums)
            freq[x]++;

        // number of ones
        int ones = freq[1];

        // dp[mask]
        // mask represents used prime factors
        long[] dp = new long[1024];

        // empty subset
        dp[0] = 1;

        // process numbers from 2 to 30
        for (int num = 2; num <= 30; num++)
        {
            // number not present
            if (freq[num] == 0)
                continue;

            int currMask = getMask(num);

            // invalid number
            if (currMask == -1)
                continue;

            // reverse traversal
            // prevents overwriting current states
            for (int mask = 1023; mask >= 0; mask--)
            {
                // overlapping prime factor
                // cannot take together
                if ((mask & currMask) != 0)
                    continue;

                dp[mask | currMask] =
                    (dp[mask | currMask] +
                     dp[mask] * freq[num]) % MOD;
            }
        }

        long ans = 0;

        // add all good subsets
        // skip empty subset (mask = 0)
        for (int mask = 1; mask < 1024; mask++)
            ans = (ans + dp[mask]) % MOD;

        // every valid subset can be combined
        // with any subset of ones
        ans = (ans * power(2, ones)) % MOD;

        return (int)ans;
    }

    static void Main()
    {
        int[] arr = { 1, 2, 3, 4 };

        Console.WriteLine(countSubsets(arr));
    }
}
JavaScript
const MOD = 1000000007n;

// Fast exponentiation
function power(a, b) {
    let ans = 1n;
    a = BigInt(a);

    while (b > 0) {

        if (b & 1)
            ans = (ans * a) % MOD;

        a = (a * a) % MOD;
        b >>= 1;
    }

    return ans;
}

// Create prime mask for current number
function getMask(num) {
    const primes = [2, 3, 5, 7, 11, 13, 17, 19, 23, 29];

    let mask = 0;

    for (let i = 0; i < 10; i++) {

        let p = primes[i];
        let cnt = 0;

        while (num % p === 0) {
            cnt++;
            num = Math.floor(num / p);
        }

        // repeated prime factor
        // invalid number
        if (cnt > 1)
            return -1;

        // store current prime in mask
        if (cnt === 1)
            mask |= (1 << i);
    }

    return mask;
}

function countSubsets(nums) {

    // frequency of every number
    let freq = new Array(31).fill(0);

    for (let x of nums)
        freq[x]++;

    // number of ones
    let ones = freq[1];

    // dp[mask]
    // mask represents used prime factors
    let dp = new Array(1024).fill(0n);

    // empty subset
    dp[0] = 1n;

    // process numbers from 2 to 30
    for (let num = 2; num <= 30; num++) {

        // number not present
        if (freq[num] === 0)
            continue;

        let currMask = getMask(num);

        // invalid number
        if (currMask === -1)
            continue;

        // reverse traversal
        // prevents overwriting current states
        for (let mask = 1023; mask >= 0; mask--) {

            // overlapping prime factor
            // cannot take together
            if ((mask & currMask)!== 0)
                continue;

            dp[mask | currMask] =
                (dp[mask | currMask] +
                dp[mask] * BigInt(freq[num])) % MOD;
        }
    }

    let ans = 0n;

    // add all good subsets
    // skip empty subset (mask = 0)
    for (let mask = 1; mask < 1024; mask++)
        ans = (ans + dp[mask]) % MOD;

    // every valid subset can be combined
    // with any subset of ones
    ans = (ans * power(2, ones)) % MOD;

    return Number(ans);
}

// Driver code
console.log(countSubsets([1, 2, 3, 4])); 

Output
6
Comment